Compare commits

..
Author SHA1 Message Date
-LAN- 0d3aab5901 refactor(api): move TokenBufferMemory to model_runtime 2026-02-28 18:02:39 +08:00
dependabot[bot]GitHubdependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com>
48d8667c4f chore(deps): bump pypdf from 6.7.1 to 6.7.4 in /api (#32736)
Signed-off-by: dependabot[bot] <support@github.com>
Co-authored-by: dependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com>
2026-02-28 16:42:03 +09:00
Tyson CungandGitHub 91dfdd87e3 fix: replace unreachable yield expression with yield from () (#32727) 2026-02-28 15:27:32 +09:00
Tyson CungandGitHub e4316a9bf6 fix(ci): fix invalid workflow file pyrefly-diff.yml (#32728) 2026-02-28 15:26:48 +09:00
hj24andGitHub 87bf7401f1 feat: add backend-code-review skill (#32719) 2026-02-28 14:17:48 +08:00
33242697ce test: migrate document_service_status SQL tests to testcontainers (#32536)
Co-authored-by: KinomotoMio <200703522+KinomotoMio@users.noreply.github.com>
2026-02-28 01:50:55 +09:00
Niels KaspersandGitHub 24fe95308a fix: YAML syntax error in pyrefly-diff-comment workflow (#32718) 2026-02-28 00:09:56 +09:00
yyhandGitHub d8f8b8cd07 chore(deps-dev): align all @storybook/* packages to 10.2.13 (#32714) 2026-02-27 22:55:53 +09:00
ad600f0827 test: migrate test_dataset_service SQL tests to testcontainers (#32535)
Co-authored-by: KinomotoMio <200703522+KinomotoMio@users.noreply.github.com>
2026-02-27 22:40:20 +09:00
yyhandGitHub 35b31d0cdd ci(web): parallelize web tests with 4-shard Vitest sharding (#32713) 2026-02-27 21:33:12 +08:00
dependabot[bot]GitHubdependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com>
592ad04818 chore(deps-dev): bump storybook from 10.2.0 to 10.2.10 in /web (#32659)
Signed-off-by: dependabot[bot] <support@github.com>
Co-authored-by: dependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com>
2026-02-27 21:53:36 +09:00
71ff135927 fix: add return type to abstract _publish method (#32701)
Co-authored-by: root <root@DESKTOP-KQLO90N>
Co-authored-by: Sisyphus <clio-agent@sisyphuslabs.ai>
2026-02-27 21:52:49 +09:00
f73be8d69e feat(web): add hover clear button for provider search (#32707)
Signed-off-by: -LAN- <laipz8200@outlook.com>
Co-authored-by: yyh <yuanyouhuilyz@gmail.com>
Co-authored-by: yyh <92089059+lyzno1@users.noreply.github.com>
2026-02-27 20:42:30 +08:00
f9196f7bea test: migrate document_indexing_sync_task SQL tests to testcontainers (#32534)
Co-authored-by: KinomotoMio <200703522+KinomotoMio@users.noreply.github.com>
2026-02-27 21:36:32 +09:00
Stephen ZhouandGitHub 439ff3775d chore: update to eslint 10 (#32646) 2026-02-27 19:44:54 +08:00
Varun ChawlaandGitHub 233e12e631 fix: correct mock return type in CodeBasedExtension test (#32058) 2026-02-27 20:40:51 +09:00
wangxiaoleiGitHubautofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
eccb67d5b6 refactor: decouple the business logic from datasource_node (#32515)
Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
2026-02-27 18:49:14 +08:00
-LAN-andGitHub 1e6de0e6ad docs(api): simplify setup README and worker guidance (#32704) 2026-02-27 18:12:52 +08:00
非法操作andGitHub 9f0ee5c145 fix: the action button of structure output modal should align right (#32700) 2026-02-27 17:28:41 +08:00
dependabot[bot]GitHubdependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com>
6c66e11cac chore(deps-dev): bump nltk from 3.9.2 to 3.9.3 in /api (#32691)
Signed-off-by: dependabot[bot] <support@github.com>
Co-authored-by: dependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com>
2026-02-27 17:20:55 +09:00
-LAN-andGitHub 149a7870bc test: align file preview mimetype expectation (#32688) 2026-02-27 15:27:30 +08:00
-LAN-andGitHub 661af404e9 chore(ci): fold pyrefly diff comments (#32685) 2026-02-27 16:23:59 +09:00
8ff51a58fd refactor(web): remove mouseup listener in use-resize-panel cleanup (#32636)
Co-authored-by: 非法操作 <hjlarry@163.com>
2026-02-27 15:06:10 +08:00
LeileiandGitHub f17c234a92 chore: update README.md (#32680) 2026-02-27 14:39:15 +08:00
-LAN-andGitHub a694533fc9 refactor(workflow): inject credential/model access ports into LLM nodes (#32569)
Signed-off-by: -LAN- <laipz8200@outlook.com>
2026-02-27 14:36:41 +08:00
-LAN-andGitHub d20880d102 revert: "fix: image preview triggers binary download" (#32683) 2026-02-27 14:28:30 +08:00
-LAN-andGitHub eea1cf17ef refactor(workflow): inject redis into graph engine manager (#32622) 2026-02-27 13:29:52 +08:00
-LAN-andGitHub 700a4029c6 refactor(api): inject code executor from node factory (#32618) 2026-02-27 13:29:00 +08:00
PoojanandGitHub 5b45b62994 test: improve coverage for header components (#32628) 2026-02-27 10:27:46 +08:00
不做了睡大觉GitHubUserautofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
349d2d8e4e fix: replace deprecated SpanAttributes and ResourceAttributes with new semconv imports (#32661)
Co-authored-by: User <user@example.com>
Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
2026-02-27 08:53:45 +09:00
edvatarGitHubgemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com>
2eefb585f9 fix: add type annotations to BaseStorage.exists and BaseStorage.download (#32652)
Signed-off-by: edvatar <88481784+toroleapinc@users.noreply.github.com>
Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com>
2026-02-27 07:35:30 +09:00
木之本澪GitHubKinomotoMiogemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com>autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>Copilot
5cb1b53b47 test: migrate dataset service update-dataset SQL tests to testcontainers (#32533)
Co-authored-by: KinomotoMio <200703522+KinomotoMio@users.noreply.github.com>
Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com>
Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
Co-authored-by: Copilot <175728472+Copilot@users.noreply.github.com>
2026-02-27 07:10:15 +09:00
edvatarandGitHub b48f36a4e5 fix: replace dict() merge with dict unpacking to resolve overload error (#32653)
Signed-off-by: edvatar <88481784+toroleapinc@users.noreply.github.com>
2026-02-27 06:15:17 +09:00
木之本澪GitHubKinomotoMioautofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
0bf5f4df3b test: migrate dataset_indexing_task SQL tests to testcontainers (#32531)
Co-authored-by: KinomotoMio <200703522+KinomotoMio@users.noreply.github.com>
Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
2026-02-27 06:06:42 +09:00
56759c03b7 test: migrate clean_dataset_task SQL tests to testcontainers (#32529)
Co-authored-by: KinomotoMio <200703522+KinomotoMio@users.noreply.github.com>
Co-authored-by: Copilot <175728472+Copilot@users.noreply.github.com>
2026-02-26 18:59:36 +09:00
cec6d82650 fix: add None checks for tenant.id in dataset vector index tests (#32603)
Co-authored-by: User <user@example.com>
2026-02-26 17:15:45 +09:00
141 changed files with 9179 additions and 8116 deletions
+168
View File
@@ -0,0 +1,168 @@
---
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.
---
# Backend Code Review
## When to use this skill
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:
- **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`).
Do NOT use this skill when:
- 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/`).
## How to use this skill
Follow these steps when using this skill:
1. **Identify the review mode** (pending-change vs snippet vs file-focused) based on the users 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**.
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.
## Checklist
- 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 <X> critical issues need to be fixed:
## 🔴 Critical (Must Fix)
### 1. <brief description of the issue>
FilePath: <path> line <line>
<relevant code snippet or pointer>
#### Explanation
<detailed explanation and references of the issue>
#### Suggested Fix
1. <brief description of suggested fix>
2. <code example> (optional, omit if not applicable)
---
... (repeat for each critical issue) ...
Found <Y> suggestions for improvement:
## 🟡 Suggestions (Should Consider)
### 1. <brief description of the suggestion>
FilePath: <path> line <line>
<relevant code snippet or pointer>
#### Explanation
<detailed explanation and references of the suggestion>
#### Suggested Fix
1. <brief description of suggested fix>
2. <code example> (optional, omit if not applicable)
---
... (repeat for each suggestion) ...
Found <Z> optional nits:
## 🟢 Nits (Optional)
### 1. <brief description of the nit>
FilePath: <path> line <line>
<relevant code snippet or pointer>
#### Explanation
<explanation and references of the optional nit>
#### Suggested Fix
- <minor suggestions>
---
... (repeat for each nits) ...
## ✅ What's Good
- <Positive feedback on good patterns>
```
- 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.
```
@@ -0,0 +1,91 @@
# Rule Catalog — Architecture
## Scope
- Covers: controller/service/core-domain/libs/model layering, dependency direction, responsibility placement, observability-friendly flow.
## Rules
### 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:
- Bad:
```python
@bp.post("/apps/<app_id>/publish")
def publish_app(app_id: str):
payload = request.get_json() or {}
if payload.get("force") and current_user.role != "admin":
raise ValueError("only admin can force publish")
app = App.query.get(app_id)
app.status = "published"
db.session.commit()
return {"result": "ok"}
```
- Good:
```python
@bp.post("/apps/<app_id>/publish")
def publish_app(app_id: str):
payload = PublishRequest.model_validate(request.get_json() or {})
app_service.publish_app(app_id=app_id, force=payload.force, actor_id=current_user.id)
return {"result": "ok"}
```
### 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:
- Bad:
```python
# core/policy/publish_policy.py
from controllers.console.app import request_context
def can_publish() -> bool:
return request_context.current_user.is_admin
```
- Good:
```python
# core/policy/publish_policy.py
def can_publish(role: str) -> bool:
return role == "admin"
# service layer adapts web/user context to domain input
allowed = can_publish(role=current_user.role)
```
### 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.
- Keep `libs` dependencies clean: avoid importing service/controller/domain-specific modules into `api/libs/`.
- Example:
- Bad:
```python
# api/libs/conversation_filter.py
from services.conversation_service import ConversationService
def should_archive_conversation(conversation, tenant_id: str) -> bool:
# Domain policy and service dependency are leaking into libs.
service = ConversationService()
if service.has_paid_plan(tenant_id):
return conversation.idle_days > 90
return conversation.idle_days > 30
```
- Good:
```python
# api/libs/datetime_utils.py (business-agnostic helper)
def older_than_days(idle_days: int, threshold_days: int) -> bool:
return idle_days > threshold_days
# services/conversation_service.py (business logic stays in service/core)
from libs.datetime_utils import older_than_days
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)
```
@@ -0,0 +1,157 @@
# Rule Catalog — DB Schema Design
## Scope
- Covers: model/base inheritance, schema boundaries in model properties, tenant-aware schema design, index redundancy checks, dialect portability in models, and cross-database compatibility in migrations.
- Does NOT cover: session lifecycle, transaction boundaries, and query execution patterns (handled by `sqlalchemy-rule.md`).
## Rules
### 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.
- Move cross-table data fetching to service/repository methods.
- For list/batch reads, fetch required related data explicitly (join/preload/bulk query) before rendering derived values.
- Example:
- Bad:
```python
class Conversation(TypeBase):
__tablename__ = "conversations"
@property
def app_name(self) -> str:
with Session(db.engine, expire_on_commit=False) as session:
app = session.execute(select(App).where(App.id == self.app_id)).scalar_one()
return app.name
```
- Good:
```python
class Conversation(TypeBase):
__tablename__ = "conversations"
@property
def display_title(self) -> str:
return self.name or "Untitled"
# Service/repository layer performs explicit batch fetch for related App rows.
```
### 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.
- Propagate `tenant_id` through service/repository contracts to keep access paths tenant-aware.
- Exception: if a table is explicitly designed as non-tenant-scoped global metadata, document that design decision clearly.
- Example:
- Bad:
```python
from sqlalchemy.orm import Mapped
class Dataset(TypeBase):
__tablename__ = "datasets"
id: Mapped[str] = mapped_column(StringUUID, primary_key=True)
name: Mapped[str] = mapped_column(sa.String(255), nullable=False)
```
- Good:
```python
from sqlalchemy.orm import Mapped
class Dataset(TypeBase):
__tablename__ = "datasets"
id: Mapped[str] = mapped_column(StringUUID, primary_key=True)
tenant_id: Mapped[str] = mapped_column(StringUUID, nullable=False, index=True)
name: Mapped[str] = mapped_column(sa.String(255), nullable=False)
```
### 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.
- Drop or avoid creating redundant prefixes unless there is a proven query-pattern need.
- Apply the same review standard in both model `__table_args__` and migration index DDL.
- Example:
- Bad:
```python
__table_args__ = (
sa.Index("idx_msg_tenant_app", "tenant_id", "app_id"),
sa.Index("idx_msg_tenant_app_created", "tenant_id", "app_id", "created_at"),
)
```
- Good:
```python
__table_args__ = (
# Keep the wider index unless profiling proves a dedicated short index is needed.
sa.Index("idx_msg_tenant_app_created", "tenant_id", "app_id", "created_at"),
)
```
### 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.
- Add or extend wrappers in `models.types` (for example, `AdjustedJSON`, `LongText`, `BinaryData`) to normalize behavior across PostgreSQL and MySQL.
- Example:
- Bad:
```python
from sqlalchemy.dialects.postgresql import JSONB
from sqlalchemy.orm import Mapped
class ToolConfig(TypeBase):
__tablename__ = "tool_configs"
config: Mapped[dict] = mapped_column(JSONB, nullable=False)
```
- Good:
```python
from sqlalchemy.orm import Mapped
from models.types import AdjustedJSON
class ToolConfig(TypeBase):
__tablename__ = "tool_configs"
config: Mapped[dict] = mapped_column(AdjustedJSON(), nullable=False)
```
### 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.
- Reuse `models.types` wrappers in column definitions when that keeps behavior aligned with runtime models.
- Avoid one-dialect-only migration logic unless there is a documented, deliberate compatibility exception.
- Example:
- Bad:
```python
with op.batch_alter_table("dataset_keyword_tables") as batch_op:
batch_op.add_column(
sa.Column(
"data_source_type",
sa.String(255),
server_default=sa.text("'database'::character varying"),
nullable=False,
)
)
```
- Good:
```python
def _is_pg(conn) -> bool:
return conn.dialect.name == "postgresql"
conn = op.get_bind()
default_expr = sa.text("'database'::character varying") if _is_pg(conn) else sa.text("'database'")
with op.batch_alter_table("dataset_keyword_tables") as batch_op:
batch_op.add_column(
sa.Column("data_source_type", sa.String(255), server_default=default_expr, nullable=False)
)
```
@@ -0,0 +1,61 @@
# Rule Catalog - Repositories Abstraction
## Scope
- Covers: when to reuse existing repository abstractions, when to introduce new repositories, and how to preserve dependency direction between service/core and infrastructure implementations.
- Does NOT cover: SQLAlchemy session lifecycle and query-shape specifics (handled by `sqlalchemy-rule.md`), and table schema/migration design (handled by `db-schema-rule.md`).
## Rules
### 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.
- If no repository exists, add one only when complexity warrants it (for example, repeated complex queries, large data domains, or multiple storage strategies), while preserving dependency direction (service/core depends on abstraction; infra provides implementation).
- Example:
- Bad:
```python
# Existing repository is ignored and service uses ad-hoc table queries.
class AppService:
def archive_app(self, app_id: str, tenant_id: str) -> None:
app = self.session.execute(
select(App).where(App.id == app_id, App.tenant_id == tenant_id)
).scalar_one()
app.archived = True
self.session.commit()
```
- Good:
```python
# Case A: Existing repository must be reused for all table operations.
class AppService:
def archive_app(self, app_id: str, tenant_id: str) -> None:
app = self.app_repo.get_by_id(app_id=app_id, tenant_id=tenant_id)
app.archived = True
self.app_repo.save(app)
# If the query is missing, extend the existing abstraction.
active_apps = self.app_repo.list_active_for_tenant(tenant_id=tenant_id)
```
- Bad:
```python
# No repository exists, but large-domain query logic is scattered in service code.
class ConversationService:
def list_recent_for_app(self, app_id: str, tenant_id: str, limit: int) -> list[Conversation]:
...
# many filters/joins/pagination variants duplicated across services
```
- Good:
```python
# Case B: Introduce repository for large/complex domains or storage variation.
class ConversationRepository(Protocol):
def list_recent_for_app(self, app_id: str, tenant_id: str, limit: int) -> list[Conversation]: ...
class SqlAlchemyConversationRepository:
def list_recent_for_app(self, app_id: str, tenant_id: str, limit: int) -> list[Conversation]:
...
class ConversationService:
def __init__(self, conversation_repo: ConversationRepository):
self.conversation_repo = conversation_repo
```
@@ -0,0 +1,139 @@
# Rule Catalog — SQLAlchemy Patterns
## Scope
- Covers: SQLAlchemy session and transaction lifecycle, query construction, tenant scoping, raw SQL boundaries, and write-path concurrency safeguards.
- Does NOT cover: table/model schema and migration design details (handled by `db-schema-rule.md`).
## Rules
### 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.
- Or use **`session.begin()` context manager** for automatic commit/rollback on a scoped block.
- Keep transaction windows short: avoid network I/O, heavy computation, or unrelated work inside the transaction.
- Example:
- Bad:
```python
# Missing commit: write may never be persisted.
with Session(db.engine, expire_on_commit=False) as session:
run = session.get(WorkflowRun, run_id)
run.status = "cancelled"
# Long transaction: external I/O inside a DB transaction.
with Session(db.engine, expire_on_commit=False) as session, session.begin():
run = session.get(WorkflowRun, run_id)
run.status = "cancelled"
call_external_api()
```
- Good:
```python
# Option 1: explicit commit.
with Session(db.engine, expire_on_commit=False) as session:
run = session.get(WorkflowRun, run_id)
run.status = "cancelled"
session.commit()
# Option 2: scoped transaction with automatic commit/rollback.
with Session(db.engine, expire_on_commit=False) as session, session.begin():
run = session.get(WorkflowRun, run_id)
run.status = "cancelled"
# Keep non-DB work outside transaction scope.
call_external_api()
```
### 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:
- Bad:
```python
stmt = select(Workflow).where(Workflow.id == workflow_id)
workflow = session.execute(stmt).scalar_one_or_none()
```
- Good:
```python
stmt = select(Workflow).where(
Workflow.id == workflow_id,
Workflow.tenant_id == tenant_id,
)
workflow = session.execute(stmt).scalar_one_or_none()
```
### 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:
- Bad:
```python
row = session.execute(
text("SELECT * FROM workflows WHERE id = :id AND tenant_id = :tenant_id"),
{"id": workflow_id, "tenant_id": tenant_id},
).first()
```
- Good:
```python
stmt = select(Workflow).where(
Workflow.id == workflow_id,
Workflow.tenant_id == tenant_id,
)
row = session.execute(stmt).scalar_one_or_none()
```
### 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.
- **Redis distributed lock**: Use when the critical section spans multiple steps/processes (or includes non-DB side effects) and you need cross-worker mutual exclusion.
- **SELECT ... FOR UPDATE**: Use when contention is high on the same rows and strict in-transaction serialization is required. Keep transactions short to reduce lock wait/deadlock risk.
- In all cases, scope by `tenant_id` and verify affected row counts for conditional writes.
- Example:
- Bad:
```python
# No tenant scope, no conflict detection, and no lock on a contested write path.
session.execute(update(WorkflowRun).where(WorkflowRun.id == run_id).values(status="cancelled"))
session.commit() # silently overwrites concurrent updates
```
- Good:
```python
# 1) Optimistic lock (low contention, retry on conflict)
result = session.execute(
update(WorkflowRun)
.where(
WorkflowRun.id == run_id,
WorkflowRun.tenant_id == tenant_id,
WorkflowRun.version == expected_version,
)
.values(status="cancelled", version=WorkflowRun.version + 1)
)
if result.rowcount == 0:
raise WorkflowStateConflictError("stale version, retry")
# 2) Redis distributed lock (cross-worker critical section)
lock_name = f"workflow_run_lock:{tenant_id}:{run_id}"
with redis_client.lock(lock_name, timeout=20):
session.execute(
update(WorkflowRun)
.where(WorkflowRun.id == run_id, WorkflowRun.tenant_id == tenant_id)
.values(status="cancelled")
)
session.commit()
# 3) Pessimistic lock with SELECT ... FOR UPDATE (high contention)
run = session.execute(
select(WorkflowRun)
.where(WorkflowRun.id == run_id, WorkflowRun.tenant_id == tenant_id)
.with_for_update()
).scalar_one()
run.status = "cancelled"
session.commit()
```
+1
View File
@@ -0,0 +1 @@
../../.agents/skills/backend-code-review
+2 -2
View File
@@ -77,8 +77,8 @@ jobs:
}
const body = diff.trim()
? `### Pyrefly Diff (base → PR)\\n\\`\\`\\`diff\\n${diff}\\n\\`\\`\\``
: '### Pyrefly Diff\\nNo changes detected.';
? '### Pyrefly Diff\n<details>\n<summary>base → PR</summary>\n\n```diff\n' + diff + '\n```\n</details>'
: '### Pyrefly Diff\nNo changes detected.';
await github.rest.issues.createComment({
issue_number: prNumber,
+10 -1
View File
@@ -74,7 +74,16 @@ jobs:
}
const body = diff.trim()
? `### Pyrefly Diff (base → PR)\n\`\`\`diff\n${diff}\n\`\`\``
? [
'### Pyrefly Diff',
'<details>',
'<summary>base → PR</summary>',
'',
'```diff',
diff,
'```',
'</details>',
].join('\n')
: '### Pyrefly Diff\nNo changes detected.';
await github.rest.issues.createComment({
+61 -2
View File
@@ -3,14 +3,22 @@ name: Web Tests
on:
workflow_call:
permissions:
contents: read
concurrency:
group: web-tests-${{ github.head_ref || github.run_id }}
cancel-in-progress: true
jobs:
test:
name: Web Tests
name: Web Tests (${{ matrix.shardIndex }}/${{ matrix.shardTotal }})
runs-on: ubuntu-latest
strategy:
fail-fast: false
matrix:
shardIndex: [1, 2, 3, 4]
shardTotal: [4]
defaults:
run:
shell: bash
@@ -39,7 +47,58 @@ jobs:
run: pnpm install --frozen-lockfile
- name: Run tests
run: pnpm test:ci
run: pnpm vitest run --reporter=blob --shard=${{ matrix.shardIndex }}/${{ matrix.shardTotal }} --coverage
- name: Upload blob report
if: ${{ !cancelled() }}
uses: actions/upload-artifact@v6
with:
name: blob-report-${{ matrix.shardIndex }}
path: web/.vitest-reports/*
include-hidden-files: true
retention-days: 1
merge-reports:
name: Merge Test Reports
if: ${{ !cancelled() }}
needs: [test]
runs-on: ubuntu-latest
defaults:
run:
shell: bash
working-directory: ./web
steps:
- name: Checkout code
uses: actions/checkout@v6
with:
persist-credentials: false
- name: Install pnpm
uses: pnpm/action-setup@v4
with:
package_json_file: web/package.json
run_install: false
- name: Setup Node.js
uses: actions/setup-node@v6
with:
node-version: 24
cache: pnpm
cache-dependency-path: ./web/pnpm-lock.yaml
- name: Install dependencies
run: pnpm install --frozen-lockfile
- name: Download blob reports
uses: actions/download-artifact@v6
with:
path: web/.vitest-reports
pattern: blob-report-*
merge-multiple: true
- name: Merge reports
run: pnpm vitest --merge-reports --coverage --silent=passed-only
- name: Coverage Summary
if: always()
-4
View File
@@ -1,9 +1,5 @@
![cover-v5-optimized](./images/GitHub_README_if.png)
<p align="center">
📌 <a href="https://dify.ai/blog/introducing-dify-workflow-file-upload-a-demo-on-ai-podcast">Introducing Dify Workflow File Upload: Recreate Google NotebookLM Podcast</a>
</p>
<p align="center">
<a href="https://cloud.dify.ai">Dify Cloud</a> ·
<a href="https://docs.dify.ai/getting-started/install-self-hosted">Self-hosting</a> ·
+13 -15
View File
@@ -50,14 +50,11 @@ forbidden_modules =
allow_indirect_imports = True
ignore_imports =
core.workflow.nodes.agent.agent_node -> extensions.ext_database
core.workflow.nodes.datasource.datasource_node -> extensions.ext_database
core.workflow.nodes.knowledge_index.knowledge_index_node -> extensions.ext_database
core.workflow.nodes.llm.file_saver -> extensions.ext_database
core.workflow.nodes.llm.llm_utils -> extensions.ext_database
core.workflow.nodes.llm.node -> extensions.ext_database
core.workflow.nodes.tool.tool_node -> extensions.ext_database
core.workflow.graph_engine.command_channels.redis_channel -> extensions.ext_redis
core.workflow.graph_engine.manager -> extensions.ext_redis
# TODO(QuantumGhost): use DI to avoid depending on global DB.
core.workflow.nodes.human_input.human_input_node -> extensions.ext_database
@@ -91,7 +88,6 @@ forbidden_modules =
core.logging
core.mcp
core.memory
core.model_manager
core.moderation
core.ops
core.plugin
@@ -105,15 +101,10 @@ forbidden_modules =
core.variables
ignore_imports =
core.workflow.nodes.loop.loop_node -> core.app.workflow.node_factory
core.workflow.graph_engine.command_channels.redis_channel -> extensions.ext_redis
core.workflow.workflow_entry -> core.app.workflow.layers.observability
core.workflow.nodes.agent.agent_node -> core.model_manager
core.workflow.nodes.agent.agent_node -> core.provider_manager
core.workflow.nodes.agent.agent_node -> core.tools.tool_manager
core.workflow.nodes.code.code_node -> core.helper.code_executor.code_executor
core.workflow.nodes.datasource.datasource_node -> models.model
core.workflow.nodes.datasource.datasource_node -> models.tools
core.workflow.nodes.datasource.datasource_node -> services.datasource_provider_service
core.workflow.nodes.document_extractor.node -> core.helper.ssrf_proxy
core.workflow.nodes.http_request.node -> core.tools.tool_file_manager
core.workflow.nodes.iteration.iteration_node -> core.app.workflow.node_factory
@@ -121,6 +112,7 @@ ignore_imports =
core.workflow.nodes.llm.llm_utils -> configs
core.workflow.nodes.llm.llm_utils -> core.app.entities.app_invoke_entities
core.workflow.nodes.llm.llm_utils -> core.model_manager
core.workflow.nodes.llm.protocols -> core.model_manager
core.workflow.nodes.llm.llm_utils -> core.model_runtime.model_providers.__base.large_language_model
core.workflow.nodes.llm.llm_utils -> models.model
core.workflow.nodes.llm.llm_utils -> models.provider
@@ -150,8 +142,6 @@ ignore_imports =
core.workflow.workflow_entry -> core.app.apps.exc
core.workflow.workflow_entry -> core.app.entities.app_invoke_entities
core.workflow.workflow_entry -> core.app.workflow.node_factory
core.workflow.nodes.datasource.datasource_node -> core.datasource.datasource_manager
core.workflow.nodes.datasource.datasource_node -> core.datasource.utils.message_transformer
core.workflow.nodes.llm.llm_utils -> core.entities.provider_entities
core.workflow.nodes.parameter_extractor.parameter_extractor_node -> core.model_manager
core.workflow.nodes.question_classifier.question_classifier_node -> core.model_manager
@@ -164,7 +154,6 @@ ignore_imports =
core.workflow.nodes.code.code_node -> core.helper.code_executor.javascript.javascript_code_provider
core.workflow.nodes.code.code_node -> core.helper.code_executor.python3.python3_code_provider
core.workflow.nodes.code.entities -> core.helper.code_executor.code_executor
core.workflow.nodes.datasource.datasource_node -> core.variables.variables
core.workflow.nodes.http_request.executor -> core.helper.ssrf_proxy
core.workflow.nodes.http_request.node -> core.helper.ssrf_proxy
core.workflow.nodes.llm.file_saver -> core.helper.ssrf_proxy
@@ -201,7 +190,6 @@ ignore_imports =
core.workflow.nodes.code.code_node -> core.variables.segments
core.workflow.nodes.code.code_node -> core.variables.types
core.workflow.nodes.code.entities -> core.variables.types
core.workflow.nodes.datasource.datasource_node -> core.variables.segments
core.workflow.nodes.document_extractor.node -> core.variables
core.workflow.nodes.document_extractor.node -> core.variables.segments
core.workflow.nodes.http_request.executor -> core.variables.segments
@@ -243,9 +231,7 @@ ignore_imports =
core.workflow.variable_loader -> core.variables
core.workflow.variable_loader -> core.variables.consts
core.workflow.workflow_type_encoder -> core.variables
core.workflow.graph_engine.manager -> extensions.ext_redis
core.workflow.nodes.agent.agent_node -> extensions.ext_database
core.workflow.nodes.datasource.datasource_node -> extensions.ext_database
core.workflow.nodes.knowledge_index.knowledge_index_node -> extensions.ext_database
core.workflow.nodes.llm.file_saver -> extensions.ext_database
core.workflow.nodes.llm.llm_utils -> extensions.ext_database
@@ -261,6 +247,11 @@ ignore_imports =
core.workflow.workflow_entry -> models.enums
core.workflow.nodes.agent.agent_node -> services
core.workflow.nodes.tool.tool_node -> services
core.workflow.nodes.agent.agent_node -> core.model_runtime.token_buffer_memory
core.workflow.nodes.llm.llm_utils -> core.model_runtime.token_buffer_memory
core.workflow.nodes.llm.node -> core.model_runtime.token_buffer_memory
core.workflow.nodes.parameter_extractor.parameter_extractor_node -> core.model_runtime.token_buffer_memory
core.workflow.nodes.question_classifier.question_classifier_node -> core.model_runtime.token_buffer_memory
[importlinter:contract:model-runtime-no-internal-imports]
name = Model Runtime Internal Imports
@@ -313,6 +304,13 @@ ignore_imports =
core.model_runtime.model_providers.model_provider_factory -> configs
core.model_runtime.model_providers.model_provider_factory -> extensions.ext_redis
core.model_runtime.model_providers.model_provider_factory -> models.provider_ids
core.model_runtime.token_buffer_memory -> core.app.app_config.features.file_upload.manager
core.model_runtime.token_buffer_memory -> core.model_manager
core.model_runtime.token_buffer_memory -> core.prompt.utils.extract_thread_messages
core.model_runtime.token_buffer_memory -> core.workflow.file.file_manager
core.model_runtime.token_buffer_memory -> extensions.ext_database
core.model_runtime.token_buffer_memory -> models.model
core.model_runtime.token_buffer_memory -> models.workflow
[importlinter:contract:rsc]
name = RSC
+1 -81
View File
@@ -42,7 +42,7 @@ The scripts resolve paths relative to their location, so you can run them from a
1. Set up your application by visiting `http://localhost:3000`.
1. Optional: start the worker service (async tasks, runs from `api`).
1. Start the worker service (async and scheduler tasks, runs from `api`).
```bash
./dev/start-worker
@@ -54,86 +54,6 @@ The scripts resolve paths relative to their location, so you can run them from a
./dev/start-beat
```
### Manual commands
<details>
<summary>Show manual setup and run steps</summary>
These commands assume you start from the repository root.
1. Start the docker-compose stack.
The backend requires middleware, including PostgreSQL, Redis, and Weaviate, which can be started together using `docker-compose`.
```bash
cp docker/middleware.env.example docker/middleware.env
# Use mysql or another vector database profile if you are not using postgres/weaviate.
docker compose -f docker/docker-compose.middleware.yaml --profile postgresql --profile weaviate -p dify up -d
```
1. Copy env files.
```bash
cp api/.env.example api/.env
cp web/.env.example web/.env.local
```
1. Install UV if needed.
```bash
pip install uv
# Or on macOS
brew install uv
```
1. Install API dependencies.
```bash
cd api
uv sync --group dev
```
1. Install web dependencies.
```bash
cd web
pnpm install
cd ..
```
1. Start backend (runs migrations first, in a new terminal).
```bash
cd api
uv run flask db upgrade
uv run flask run --host 0.0.0.0 --port=5001 --debug
```
1. Start Dify [web](../web) service (in a new terminal).
```bash
cd web
pnpm dev:inspect
```
1. Set up your application by visiting `http://localhost:3000`.
1. Optional: start the worker service (async tasks, in a new terminal).
```bash
cd api
uv run celery -A app.celery worker -P threads -c 2 --loglevel INFO -Q api_token,dataset,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
```
1. Optional: start Celery Beat (scheduled tasks, in a new terminal).
```bash
cd api
uv run celery -A app.celery beat
```
</details>
### Environment notes
> [!IMPORTANT]
+2 -1
View File
@@ -33,6 +33,7 @@ from core.workflow.enums import NodeType
from core.workflow.file.models import File
from core.workflow.graph_engine.manager import GraphEngineManager
from extensions.ext_database import db
from extensions.ext_redis import redis_client
from factories import file_factory, variable_factory
from fields.member_fields import simple_account_fields
from fields.workflow_fields import workflow_fields, workflow_pagination_fields
@@ -740,7 +741,7 @@ class WorkflowTaskStopApi(Resource):
AppQueueManager.set_stop_flag_no_user_check(task_id)
# New graph engine command channel mechanism
GraphEngineManager.send_stop_command(task_id)
GraphEngineManager(redis_client).send_stop_command(task_id)
return {"result": "success"}
@@ -112,11 +112,11 @@ _WORKFLOW_DRAFT_VARIABLE_WITHOUT_VALUE_FIELDS = {
"is_truncated": fields.Boolean(attribute=lambda model: model.file_id is not None),
}
_WORKFLOW_DRAFT_VARIABLE_FIELDS = dict(
_WORKFLOW_DRAFT_VARIABLE_WITHOUT_VALUE_FIELDS,
value=fields.Raw(attribute=_serialize_var_value),
full_content=fields.Raw(attribute=_serialize_full_content),
)
_WORKFLOW_DRAFT_VARIABLE_FIELDS = {
**_WORKFLOW_DRAFT_VARIABLE_WITHOUT_VALUE_FIELDS,
"value": fields.Raw(attribute=_serialize_var_value),
"full_content": fields.Raw(attribute=_serialize_full_content),
}
_WORKFLOW_DRAFT_ENV_VARIABLE_FIELDS = {
"id": fields.String,
+2 -1
View File
@@ -44,6 +44,7 @@ from core.errors.error import (
from core.model_runtime.errors.invoke import InvokeError
from core.workflow.graph_engine.manager import GraphEngineManager
from extensions.ext_database import db
from extensions.ext_redis import redis_client
from fields.app_fields import (
app_detail_fields_with_site,
deleted_tool_fields,
@@ -225,7 +226,7 @@ class TrialAppWorkflowTaskStopApi(TrialAppResource):
AppQueueManager.set_stop_flag_no_user_check(task_id)
# New graph engine command channel mechanism
GraphEngineManager.send_stop_command(task_id)
GraphEngineManager(redis_client).send_stop_command(task_id)
return {"result": "success"}
+2 -1
View File
@@ -23,6 +23,7 @@ from core.errors.error import (
)
from core.model_runtime.errors.invoke import InvokeError
from core.workflow.graph_engine.manager import GraphEngineManager
from extensions.ext_redis import redis_client
from libs import helper
from libs.login import current_account_with_tenant
from models.model import AppMode, InstalledApp
@@ -100,6 +101,6 @@ class InstalledAppWorkflowTaskStopApi(InstalledAppResource):
AppQueueManager.set_stop_flag_no_user_check(task_id)
# New graph engine command channel mechanism
GraphEngineManager.send_stop_command(task_id)
GraphEngineManager(redis_client).send_stop_command(task_id)
return {"result": "success"}
+1 -1
View File
@@ -137,7 +137,7 @@ class FilePreviewApi(Resource):
if args.as_attachment:
encoded_filename = quote(upload_file.name)
response.headers["Content-Disposition"] = f"attachment; filename*=UTF-8''{encoded_filename}"
response.headers["Content-Type"] = "application/octet-stream"
response.headers["Content-Type"] = "application/octet-stream"
enforce_download_for_html(
response,
+2 -1
View File
@@ -31,6 +31,7 @@ from core.model_runtime.errors.invoke import InvokeError
from core.workflow.enums import WorkflowExecutionStatus
from core.workflow.graph_engine.manager import GraphEngineManager
from extensions.ext_database import db
from extensions.ext_redis import redis_client
from fields.workflow_app_log_fields import build_workflow_app_log_pagination_model
from libs import helper
from libs.helper import OptionalTimestampField, TimestampField
@@ -280,7 +281,7 @@ class WorkflowTaskStopApi(Resource):
AppQueueManager.set_stop_flag_no_user_check(task_id)
# New graph engine command channel mechanism
GraphEngineManager.send_stop_command(task_id)
GraphEngineManager(redis_client).send_stop_command(task_id)
return {"result": "success"}
+2 -1
View File
@@ -24,6 +24,7 @@ from core.errors.error import (
)
from core.model_runtime.errors.invoke import InvokeError
from core.workflow.graph_engine.manager import GraphEngineManager
from extensions.ext_redis import redis_client
from libs import helper
from models.model import App, AppMode, EndUser
from services.app_generate_service import AppGenerateService
@@ -121,6 +122,6 @@ class WorkflowTaskStopApi(WebApiResource):
AppQueueManager.set_stop_flag_no_user_check(task_id)
# New graph engine command channel mechanism
GraphEngineManager.send_stop_command(task_id)
GraphEngineManager(redis_client).send_stop_command(task_id)
return {"result": "success"}
+2 -2
View File
@@ -17,7 +17,6 @@ from core.app.entities.app_invoke_entities import (
)
from core.callback_handler.agent_tool_callback_handler import DifyAgentCallbackHandler
from core.callback_handler.index_tool_callback_handler import DatasetIndexToolCallbackHandler
from core.memory.token_buffer_memory import TokenBufferMemory
from core.model_manager import ModelInstance
from core.model_runtime.entities import (
AssistantPromptMessage,
@@ -32,6 +31,7 @@ from core.model_runtime.entities import (
from core.model_runtime.entities.message_entities import ImagePromptMessageContent, PromptMessageContentUnionTypes
from core.model_runtime.entities.model_entities import ModelFeature
from core.model_runtime.model_providers.__base.large_language_model import LargeLanguageModel
from core.model_runtime.token_buffer_memory import TokenBufferMemory
from core.prompt.utils.extract_thread_messages import extract_thread_messages
from core.tools.__base.tool import Tool
from core.tools.entities.tool_entities import (
@@ -112,7 +112,7 @@ class BaseAgentRunner(AppRunner):
# check if model supports stream tool call
llm_model = cast(LargeLanguageModel, model_instance.model_type_instance)
model_schema = llm_model.get_model_schema(model_instance.model, model_instance.credentials)
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 []
+2 -2
View File
@@ -245,7 +245,7 @@ class CotAgentRunner(BaseAgentRunner, ABC):
iteration_step += 1
yield LLMResultChunk(
model=model_instance.model,
model=model_instance.model_name,
prompt_messages=prompt_messages,
delta=LLMResultChunkDelta(
index=0, message=AssistantPromptMessage(content=final_answer), usage=llm_usage["usage"]
@@ -268,7 +268,7 @@ class CotAgentRunner(BaseAgentRunner, ABC):
self.queue_manager.publish(
QueueMessageEndEvent(
llm_result=LLMResult(
model=model_instance.model,
model=model_instance.model_name,
prompt_messages=prompt_messages,
message=AssistantPromptMessage(content=final_answer),
usage=llm_usage["usage"] or LLMUsage.empty_usage(),
+2 -2
View File
@@ -178,7 +178,7 @@ class FunctionCallAgentRunner(BaseAgentRunner):
)
yield LLMResultChunk(
model=model_instance.model,
model=model_instance.model_name,
prompt_messages=result.prompt_messages,
system_fingerprint=result.system_fingerprint,
delta=LLMResultChunkDelta(
@@ -308,7 +308,7 @@ class FunctionCallAgentRunner(BaseAgentRunner):
self.queue_manager.publish(
QueueMessageEndEvent(
llm_result=LLMResult(
model=model_instance.model,
model=model_instance.model_name,
prompt_messages=prompt_messages,
message=AssistantPromptMessage(content=final_answer),
usage=llm_usage["usage"] or LLMUsage.empty_usage(),
@@ -669,16 +669,14 @@ class AdvancedChatAppGenerateTaskPipeline(GraphRuntimeStateSupport):
) -> Generator[StreamResponse, None, None]:
"""Handle retriever resources events."""
self._message_cycle_manager.handle_retriever_resources(event)
return
yield # Make this a generator
yield from ()
def _handle_annotation_reply_event(
self, event: QueueAnnotationReplyEvent, **kwargs
) -> Generator[StreamResponse, None, None]:
"""Handle annotation reply events."""
self._message_cycle_manager.handle_annotation_reply(event)
return
yield # Make this a generator
yield from ()
def _handle_message_replace_event(
self, event: QueueMessageReplaceEvent, **kwargs
+2 -2
View File
@@ -12,11 +12,11 @@ from core.app.apps.base_app_queue_manager import AppQueueManager, PublishFrom
from core.app.apps.base_app_runner import AppRunner
from core.app.entities.app_invoke_entities import AgentChatAppGenerateEntity
from core.app.entities.queue_entities import QueueAnnotationReplyEvent
from core.memory.token_buffer_memory import TokenBufferMemory
from core.model_manager import ModelInstance
from core.model_runtime.entities.llm_entities import LLMMode
from core.model_runtime.entities.model_entities import ModelFeature, ModelPropertyKey
from core.model_runtime.model_providers.__base.large_language_model import LargeLanguageModel
from core.model_runtime.token_buffer_memory import TokenBufferMemory
from core.moderation.base import ModerationError
from extensions.ext_database import db
from models.model import App, Conversation, Message
@@ -178,7 +178,7 @@ class AgentChatAppRunner(AppRunner):
# change function call strategy based on LLM model
llm_model = cast(LargeLanguageModel, model_instance.model_type_instance)
model_schema = llm_model.get_model_schema(model_instance.model, model_instance.credentials)
model_schema = llm_model.get_model_schema(model_instance.model_name, model_instance.credentials)
if not model_schema:
raise ValueError("Model schema not found")
+1 -1
View File
@@ -122,7 +122,7 @@ class AppQueueManager(ABC):
"""Attach the live graph runtime state reference for downstream consumers."""
self._graph_runtime_state = graph_runtime_state
def publish(self, event: AppQueueEvent, pub_from: PublishFrom):
def publish(self, event: AppQueueEvent, pub_from: PublishFrom) -> None:
"""
Publish event to queue
:param event:
+1 -1
View File
@@ -22,7 +22,6 @@ from core.app.entities.queue_entities import (
from core.app.features.annotation_reply.annotation_reply import AnnotationReplyFeature
from core.app.features.hosting_moderation.hosting_moderation import HostingModerationFeature
from core.external_data_tool.external_data_fetch import ExternalDataFetch
from core.memory.token_buffer_memory import TokenBufferMemory
from core.model_manager import ModelInstance
from core.model_runtime.entities.llm_entities import LLMResult, LLMResultChunk, LLMResultChunkDelta, LLMUsage
from core.model_runtime.entities.message_entities import (
@@ -33,6 +32,7 @@ from core.model_runtime.entities.message_entities import (
)
from core.model_runtime.entities.model_entities import ModelPropertyKey
from core.model_runtime.errors.invoke import InvokeBadRequestError
from core.model_runtime.token_buffer_memory import TokenBufferMemory
from core.moderation.input_moderation import InputModeration
from core.prompt.advanced_prompt_transform import AdvancedPromptTransform
from core.prompt.entities.advanced_prompt_entities import ChatModelMessage, CompletionModelPromptTemplate, MemoryConfig
+1 -1
View File
@@ -11,9 +11,9 @@ from core.app.entities.app_invoke_entities import (
)
from core.app.entities.queue_entities import QueueAnnotationReplyEvent
from core.callback_handler.index_tool_callback_handler import DatasetIndexToolCallbackHandler
from core.memory.token_buffer_memory import TokenBufferMemory
from core.model_manager import ModelInstance
from core.model_runtime.entities.message_entities import ImagePromptMessageContent
from core.model_runtime.token_buffer_memory import TokenBufferMemory
from core.moderation.base import ModerationError
from core.rag.retrieval.dataset_retrieval import DatasetRetrieval
from core.workflow.file import File
+1
View File
@@ -0,0 +1 @@
"""LLM-related application services."""
+103
View File
@@ -0,0 +1,103 @@
from __future__ import annotations
from typing import Any
from core.app.entities.app_invoke_entities import ModelConfigWithCredentialsEntity
from core.errors.error import ProviderTokenNotInitError
from core.model_manager import ModelInstance, ModelManager
from core.model_runtime.entities.model_entities import ModelType
from core.provider_manager import ProviderManager
from core.workflow.nodes.llm.entities import ModelConfig
from core.workflow.nodes.llm.exc import LLMModeRequiredError, ModelNotExistError
from core.workflow.nodes.llm.protocols import CredentialsProvider, ModelFactory
class DifyCredentialsProvider:
tenant_id: str
provider_manager: ProviderManager
def __init__(self, tenant_id: str, provider_manager: ProviderManager | None = None) -> None:
self.tenant_id = tenant_id
self.provider_manager = provider_manager or ProviderManager()
def fetch(self, provider_name: str, model_name: str) -> dict[str, Any]:
provider_configurations = self.provider_manager.get_configurations(self.tenant_id)
provider_configuration = provider_configurations.get(provider_name)
if not provider_configuration:
raise ValueError(f"Provider {provider_name} does not exist.")
provider_model = provider_configuration.get_provider_model(model_type=ModelType.LLM, model=model_name)
if provider_model is None:
raise ModelNotExistError(f"Model {model_name} not exist.")
provider_model.raise_for_status()
credentials = provider_configuration.get_current_credentials(model_type=ModelType.LLM, model=model_name)
if credentials is None:
raise ProviderTokenNotInitError(f"Model {model_name} credentials is not initialized.")
return credentials
class DifyModelFactory:
tenant_id: str
model_manager: ModelManager
def __init__(self, tenant_id: str, model_manager: ModelManager | None = None) -> None:
self.tenant_id = tenant_id
self.model_manager = model_manager or ModelManager()
def init_model_instance(self, provider_name: str, model_name: str) -> ModelInstance:
return self.model_manager.get_model_instance(
tenant_id=self.tenant_id,
provider=provider_name,
model_type=ModelType.LLM,
model=model_name,
)
def build_dify_model_access(tenant_id: str) -> tuple[CredentialsProvider, ModelFactory]:
return (
DifyCredentialsProvider(tenant_id=tenant_id),
DifyModelFactory(tenant_id=tenant_id),
)
def fetch_model_config(
*,
node_data_model: ModelConfig,
credentials_provider: CredentialsProvider,
model_factory: ModelFactory,
) -> tuple[ModelInstance, ModelConfigWithCredentialsEntity]:
if not node_data_model.mode:
raise LLMModeRequiredError("LLM mode is required.")
credentials = credentials_provider.fetch(node_data_model.provider, node_data_model.name)
model_instance = model_factory.init_model_instance(node_data_model.provider, node_data_model.name)
provider_model_bundle = model_instance.provider_model_bundle
provider_model = provider_model_bundle.configuration.get_provider_model(
model=node_data_model.name,
model_type=ModelType.LLM,
)
if provider_model is None:
raise ModelNotExistError(f"Model {node_data_model.name} not exist.")
provider_model.raise_for_status()
stop: list[str] = []
if "stop" in node_data_model.completion_params:
stop = node_data_model.completion_params.pop("stop")
model_schema = model_instance.model_type_instance.get_model_schema(node_data_model.name, credentials)
if not model_schema:
raise ModelNotExistError(f"Model {node_data_model.name} not exist.")
return model_instance, ModelConfigWithCredentialsEntity(
provider=node_data_model.provider,
model=node_data_model.name,
model_schema=model_schema,
mode=node_data_model.mode,
provider_model_bundle=provider_model_bundle,
credentials=credentials,
parameters=node_data_model.completion_params,
stop=stop,
)
+74 -5
View File
@@ -1,9 +1,12 @@
from typing import TYPE_CHECKING, final
from collections.abc import Mapping
from typing import TYPE_CHECKING, Any, final
from typing_extensions import override
from configs import dify_config
from core.helper.code_executor.code_executor import CodeExecutor
from core.app.llm.model_access import build_dify_model_access
from core.datasource.datasource_manager import DatasourceManager
from core.helper.code_executor.code_executor import CodeExecutionError, CodeExecutor
from core.helper.code_executor.code_node_provider import CodeNodeProvider
from core.helper.ssrf_proxy import ssrf_proxy
from core.rag.retrieval.dataset_retrieval import DatasetRetrieval
@@ -13,13 +16,20 @@ from core.workflow.enums import NodeType
from core.workflow.file.file_manager import file_manager
from core.workflow.graph.graph import NodeFactory
from core.workflow.nodes.base.node import Node
from core.workflow.nodes.code.code_node import CodeNode
from core.workflow.nodes.code.code_node import CodeNode, WorkflowCodeExecutor
from core.workflow.nodes.code.entities import CodeLanguage
from core.workflow.nodes.code.limits import CodeNodeLimits
from core.workflow.nodes.datasource import DatasourceNode
from core.workflow.nodes.document_extractor import DocumentExtractorNode, UnstructuredApiConfig
from core.workflow.nodes.http_request import HttpRequestNode, build_http_request_config
from core.workflow.nodes.knowledge_retrieval.knowledge_retrieval_node import KnowledgeRetrievalNode
from core.workflow.nodes.llm.node import LLMNode
from core.workflow.nodes.node_mapping import LATEST_VERSION, NODE_TYPE_CLASSES_MAPPING
from core.workflow.nodes.template_transform.template_renderer import CodeExecutorJinja2TemplateRenderer
from core.workflow.nodes.parameter_extractor.parameter_extractor_node import ParameterExtractorNode
from core.workflow.nodes.question_classifier.question_classifier_node import QuestionClassifierNode
from core.workflow.nodes.template_transform.template_renderer import (
CodeExecutorJinja2TemplateRenderer,
)
from core.workflow.nodes.template_transform.template_transform_node import TemplateTransformNode
if TYPE_CHECKING:
@@ -27,6 +37,24 @@ if TYPE_CHECKING:
from core.workflow.runtime import GraphRuntimeState
class DefaultWorkflowCodeExecutor:
def execute(
self,
*,
language: CodeLanguage,
code: str,
inputs: Mapping[str, Any],
) -> Mapping[str, Any]:
return CodeExecutor.execute_workflow_code_template(
language=language,
code=code,
inputs=inputs,
)
def is_execution_error(self, error: Exception) -> bool:
return isinstance(error, CodeExecutionError)
@final
class DifyNodeFactory(NodeFactory):
"""
@@ -43,7 +71,7 @@ class DifyNodeFactory(NodeFactory):
) -> None:
self.graph_init_params = graph_init_params
self.graph_runtime_state = graph_runtime_state
self._code_executor: type[CodeExecutor] = CodeExecutor
self._code_executor: WorkflowCodeExecutor = DefaultWorkflowCodeExecutor()
self._code_providers: tuple[type[CodeNodeProvider], ...] = CodeNode.default_code_providers()
self._code_limits = CodeNodeLimits(
max_string_length=dify_config.CODE_MAX_STRING_LENGTH,
@@ -75,6 +103,8 @@ class DifyNodeFactory(NodeFactory):
ssrf_default_max_retries=dify_config.SSRF_DEFAULT_MAX_RETRIES,
)
self._llm_credentials_provider, self._llm_model_factory = build_dify_model_access(graph_init_params.tenant_id)
@override
def create_node(self, node_config: NodeConfigDict) -> Node:
"""
@@ -140,6 +170,25 @@ class DifyNodeFactory(NodeFactory):
file_manager=self._http_request_file_manager,
)
if node_type == NodeType.LLM:
return LLMNode(
id=node_id,
config=node_config,
graph_init_params=self.graph_init_params,
graph_runtime_state=self.graph_runtime_state,
credentials_provider=self._llm_credentials_provider,
model_factory=self._llm_model_factory,
)
if node_type == NodeType.DATASOURCE:
return DatasourceNode(
id=node_id,
config=node_config,
graph_init_params=self.graph_init_params,
graph_runtime_state=self.graph_runtime_state,
datasource_manager=DatasourceManager,
)
if node_type == NodeType.KNOWLEDGE_RETRIEVAL:
return KnowledgeRetrievalNode(
id=node_id,
@@ -158,6 +207,26 @@ class DifyNodeFactory(NodeFactory):
unstructured_api_config=self._document_extractor_unstructured_api_config,
)
if node_type == NodeType.QUESTION_CLASSIFIER:
return QuestionClassifierNode(
id=node_id,
config=node_config,
graph_init_params=self.graph_init_params,
graph_runtime_state=self.graph_runtime_state,
credentials_provider=self._llm_credentials_provider,
model_factory=self._llm_model_factory,
)
if node_type == NodeType.PARAMETER_EXTRACTOR:
return ParameterExtractorNode(
id=node_id,
config=node_config,
graph_init_params=self.graph_init_params,
graph_runtime_state=self.graph_runtime_state,
credentials_provider=self._llm_credentials_provider,
model_factory=self._llm_model_factory,
)
return node_class(
id=node_id,
config=node_config,
+259 -1
View File
@@ -1,16 +1,39 @@
import logging
from collections.abc import Generator
from threading import Lock
from typing import Any, cast
from sqlalchemy import select
import contexts
from core.datasource.__base.datasource_plugin import DatasourcePlugin
from core.datasource.__base.datasource_provider import DatasourcePluginProviderController
from core.datasource.entities.datasource_entities import DatasourceProviderType
from core.datasource.entities.datasource_entities import (
DatasourceMessage,
DatasourceProviderType,
GetOnlineDocumentPageContentRequest,
OnlineDriveDownloadFileRequest,
)
from core.datasource.errors import DatasourceProviderNotFoundError
from core.datasource.local_file.local_file_provider import LocalFileDatasourcePluginProviderController
from core.datasource.online_document.online_document_plugin import OnlineDocumentDatasourcePlugin
from core.datasource.online_document.online_document_provider import OnlineDocumentDatasourcePluginProviderController
from core.datasource.online_drive.online_drive_plugin import OnlineDriveDatasourcePlugin
from core.datasource.online_drive.online_drive_provider import OnlineDriveDatasourcePluginProviderController
from core.datasource.utils.message_transformer import DatasourceFileMessageTransformer
from core.datasource.website_crawl.website_crawl_provider import WebsiteCrawlDatasourcePluginProviderController
from core.db.session_factory import session_factory
from core.plugin.impl.datasource import PluginDatasourceManager
from core.workflow.entities.workflow_node_execution import WorkflowNodeExecutionStatus
from core.workflow.enums import WorkflowNodeExecutionMetadataKey
from core.workflow.file import File
from core.workflow.file.enums import FileTransferMethod, FileType
from core.workflow.node_events import NodeRunResult, StreamChunkEvent, StreamCompletedEvent
from core.workflow.repositories.datasource_manager_protocol import DatasourceParameter, OnlineDriveDownloadFileParam
from factories import file_factory
from models.model import UploadFile
from models.tools import ToolFile
from services.datasource_provider_service import DatasourceProviderService
logger = logging.getLogger(__name__)
@@ -103,3 +126,238 @@ class DatasourceManager:
tenant_id,
datasource_type,
).get_datasource(datasource_name)
@classmethod
def get_icon_url(cls, provider_id: str, tenant_id: str, datasource_name: str, datasource_type: str) -> str:
datasource_runtime = cls.get_datasource_runtime(
provider_id=provider_id,
datasource_name=datasource_name,
tenant_id=tenant_id,
datasource_type=DatasourceProviderType.value_of(datasource_type),
)
return datasource_runtime.get_icon_url(tenant_id)
@classmethod
def stream_online_results(
cls,
*,
user_id: str,
datasource_name: str,
datasource_type: str,
provider_id: str,
tenant_id: str,
provider: str,
plugin_id: str,
credential_id: str,
datasource_param: DatasourceParameter | None = None,
online_drive_request: OnlineDriveDownloadFileParam | None = None,
) -> Generator[DatasourceMessage, None, Any]:
"""
Pull-based streaming of domain messages from datasource plugins.
Returns a generator that yields DatasourceMessage and finally returns a minimal final payload.
Only ONLINE_DOCUMENT and ONLINE_DRIVE are streamable here; other types are handled by nodes directly.
"""
ds_type = DatasourceProviderType.value_of(datasource_type)
runtime = cls.get_datasource_runtime(
provider_id=provider_id,
datasource_name=datasource_name,
tenant_id=tenant_id,
datasource_type=ds_type,
)
dsp_service = DatasourceProviderService()
credentials = dsp_service.get_datasource_credentials(
tenant_id=tenant_id,
provider=provider,
plugin_id=plugin_id,
credential_id=credential_id,
)
if ds_type == DatasourceProviderType.ONLINE_DOCUMENT:
doc_runtime = cast(OnlineDocumentDatasourcePlugin, runtime)
if credentials:
doc_runtime.runtime.credentials = credentials
if datasource_param is None:
raise ValueError("datasource_param is required for ONLINE_DOCUMENT streaming")
inner_gen: Generator[DatasourceMessage, None, None] = doc_runtime.get_online_document_page_content(
user_id=user_id,
datasource_parameters=GetOnlineDocumentPageContentRequest(
workspace_id=datasource_param.workspace_id,
page_id=datasource_param.page_id,
type=datasource_param.type,
),
provider_type=ds_type,
)
elif ds_type == DatasourceProviderType.ONLINE_DRIVE:
drive_runtime = cast(OnlineDriveDatasourcePlugin, runtime)
if credentials:
drive_runtime.runtime.credentials = credentials
if online_drive_request is None:
raise ValueError("online_drive_request is required for ONLINE_DRIVE streaming")
inner_gen = drive_runtime.online_drive_download_file(
user_id=user_id,
request=OnlineDriveDownloadFileRequest(
id=online_drive_request.id,
bucket=online_drive_request.bucket,
),
provider_type=ds_type,
)
else:
raise ValueError(f"Unsupported datasource type for streaming: {ds_type}")
# Bridge through to caller while preserving generator return contract
yield from inner_gen
# No structured final data here; node/adapter will assemble outputs
return {}
@classmethod
def stream_node_events(
cls,
*,
node_id: str,
user_id: str,
datasource_name: str,
datasource_type: str,
provider_id: str,
tenant_id: str,
provider: str,
plugin_id: str,
credential_id: str,
parameters_for_log: dict[str, Any],
datasource_info: dict[str, Any],
variable_pool: Any,
datasource_param: DatasourceParameter | None = None,
online_drive_request: OnlineDriveDownloadFileParam | None = None,
) -> Generator[StreamChunkEvent | StreamCompletedEvent, None, None]:
ds_type = DatasourceProviderType.value_of(datasource_type)
messages = cls.stream_online_results(
user_id=user_id,
datasource_name=datasource_name,
datasource_type=datasource_type,
provider_id=provider_id,
tenant_id=tenant_id,
provider=provider,
plugin_id=plugin_id,
credential_id=credential_id,
datasource_param=datasource_param,
online_drive_request=online_drive_request,
)
transformed = DatasourceFileMessageTransformer.transform_datasource_invoke_messages(
messages=messages, user_id=user_id, tenant_id=tenant_id, conversation_id=None
)
variables: dict[str, Any] = {}
file_out: File | None = None
for message in transformed:
mtype = message.type
if mtype in {
DatasourceMessage.MessageType.IMAGE_LINK,
DatasourceMessage.MessageType.BINARY_LINK,
DatasourceMessage.MessageType.IMAGE,
}:
wanted_ds_type = ds_type in {
DatasourceProviderType.ONLINE_DRIVE,
DatasourceProviderType.ONLINE_DOCUMENT,
}
if wanted_ds_type and isinstance(message.message, DatasourceMessage.TextMessage):
url = message.message.text
datasource_file_id = str(url).split("/")[-1].split(".")[0]
with session_factory.create_session() as session:
stmt = select(ToolFile).where(
ToolFile.id == datasource_file_id, ToolFile.tenant_id == tenant_id
)
datasource_file = session.scalar(stmt)
if not datasource_file:
raise ValueError(
f"ToolFile not found for file_id={datasource_file_id}, tenant_id={tenant_id}"
)
mime_type = datasource_file.mimetype
if datasource_file is not None:
mapping = {
"tool_file_id": datasource_file_id,
"type": file_factory.get_file_type_by_mime_type(mime_type),
"transfer_method": FileTransferMethod.TOOL_FILE,
"url": url,
}
file_out = file_factory.build_from_mapping(mapping=mapping, tenant_id=tenant_id)
elif mtype == DatasourceMessage.MessageType.TEXT:
assert isinstance(message.message, DatasourceMessage.TextMessage)
yield StreamChunkEvent(selector=[node_id, "text"], chunk=message.message.text, is_final=False)
elif mtype == DatasourceMessage.MessageType.LINK:
assert isinstance(message.message, DatasourceMessage.TextMessage)
yield StreamChunkEvent(
selector=[node_id, "text"], chunk=f"Link: {message.message.text}\n", is_final=False
)
elif mtype == DatasourceMessage.MessageType.VARIABLE:
assert isinstance(message.message, DatasourceMessage.VariableMessage)
name = message.message.variable_name
value = message.message.variable_value
if message.message.stream:
assert isinstance(value, str), "stream variable_value must be str"
variables[name] = variables.get(name, "") + value
yield StreamChunkEvent(selector=[node_id, name], chunk=value, is_final=False)
else:
variables[name] = value
elif mtype == DatasourceMessage.MessageType.FILE:
if ds_type == DatasourceProviderType.ONLINE_DRIVE and message.meta:
f = message.meta.get("file")
if isinstance(f, File):
file_out = f
else:
pass
yield StreamChunkEvent(selector=[node_id, "text"], chunk="", is_final=True)
if ds_type == DatasourceProviderType.ONLINE_DRIVE and file_out is not None:
variable_pool.add([node_id, "file"], file_out)
if ds_type == DatasourceProviderType.ONLINE_DOCUMENT:
yield StreamCompletedEvent(
node_run_result=NodeRunResult(
status=WorkflowNodeExecutionStatus.SUCCEEDED,
inputs=parameters_for_log,
metadata={WorkflowNodeExecutionMetadataKey.DATASOURCE_INFO: datasource_info},
outputs={**variables},
)
)
else:
yield StreamCompletedEvent(
node_run_result=NodeRunResult(
status=WorkflowNodeExecutionStatus.SUCCEEDED,
inputs=parameters_for_log,
metadata={WorkflowNodeExecutionMetadataKey.DATASOURCE_INFO: datasource_info},
outputs={
"file": file_out,
"datasource_type": ds_type,
},
)
)
@classmethod
def get_upload_file_by_id(cls, file_id: str, tenant_id: str) -> File:
with session_factory.create_session() as session:
upload_file = (
session.query(UploadFile).where(UploadFile.id == file_id, UploadFile.tenant_id == tenant_id).first()
)
if not upload_file:
raise ValueError(f"UploadFile not found for file_id={file_id}, tenant_id={tenant_id}")
file_info = File(
id=upload_file.id,
filename=upload_file.name,
extension="." + upload_file.extension,
mime_type=upload_file.mime_type,
tenant_id=tenant_id,
type=FileType.CUSTOM,
transfer_method=FileTransferMethod.LOCAL_FILE,
remote_url=upload_file.source_url,
related_id=upload_file.id,
size=upload_file.size,
storage_key=upload_file.key,
url=upload_file.source_url,
)
return file_info
@@ -379,4 +379,11 @@ class OnlineDriveDownloadFileRequest(BaseModel):
"""
id: str = Field(..., description="The id of the file")
bucket: str | None = Field(None, description="The name of the bucket")
bucket: str = Field("", description="The name of the bucket")
@field_validator("bucket", mode="before")
@classmethod
def _coerce_bucket(cls, v) -> str:
if v is None:
return ""
return str(v)
+12 -12
View File
@@ -35,7 +35,7 @@ class ModelInstance:
def __init__(self, provider_model_bundle: ProviderModelBundle, model: str):
self.provider_model_bundle = provider_model_bundle
self.model = model
self.model_name = model
self.provider = provider_model_bundle.configuration.provider.provider
self.credentials = self._fetch_credentials_from_bundle(provider_model_bundle, model)
self.model_type_instance = self.provider_model_bundle.model_type_instance
@@ -163,7 +163,7 @@ class ModelInstance:
Union[LLMResult, Generator],
self._round_robin_invoke(
function=self.model_type_instance.invoke,
model=self.model,
model=self.model_name,
credentials=self.credentials,
prompt_messages=prompt_messages,
model_parameters=model_parameters,
@@ -191,7 +191,7 @@ class ModelInstance:
int,
self._round_robin_invoke(
function=self.model_type_instance.get_num_tokens,
model=self.model,
model=self.model_name,
credentials=self.credentials,
prompt_messages=prompt_messages,
tools=tools,
@@ -215,7 +215,7 @@ class ModelInstance:
EmbeddingResult,
self._round_robin_invoke(
function=self.model_type_instance.invoke,
model=self.model,
model=self.model_name,
credentials=self.credentials,
texts=texts,
user=user,
@@ -243,7 +243,7 @@ class ModelInstance:
EmbeddingResult,
self._round_robin_invoke(
function=self.model_type_instance.invoke,
model=self.model,
model=self.model_name,
credentials=self.credentials,
multimodel_documents=multimodel_documents,
user=user,
@@ -264,7 +264,7 @@ class ModelInstance:
list[int],
self._round_robin_invoke(
function=self.model_type_instance.get_num_tokens,
model=self.model,
model=self.model_name,
credentials=self.credentials,
texts=texts,
),
@@ -294,7 +294,7 @@ class ModelInstance:
RerankResult,
self._round_robin_invoke(
function=self.model_type_instance.invoke,
model=self.model,
model=self.model_name,
credentials=self.credentials,
query=query,
docs=docs,
@@ -328,7 +328,7 @@ class ModelInstance:
RerankResult,
self._round_robin_invoke(
function=self.model_type_instance.invoke_multimodal_rerank,
model=self.model,
model=self.model_name,
credentials=self.credentials,
query=query,
docs=docs,
@@ -352,7 +352,7 @@ class ModelInstance:
bool,
self._round_robin_invoke(
function=self.model_type_instance.invoke,
model=self.model,
model=self.model_name,
credentials=self.credentials,
text=text,
user=user,
@@ -373,7 +373,7 @@ class ModelInstance:
str,
self._round_robin_invoke(
function=self.model_type_instance.invoke,
model=self.model,
model=self.model_name,
credentials=self.credentials,
file=file,
user=user,
@@ -396,7 +396,7 @@ class ModelInstance:
Iterable[bytes],
self._round_robin_invoke(
function=self.model_type_instance.invoke,
model=self.model,
model=self.model_name,
credentials=self.credentials,
content_text=content_text,
user=user,
@@ -469,7 +469,7 @@ class ModelInstance:
if not isinstance(self.model_type_instance, TTSModel):
raise Exception("Model type instance is not TTSModel")
return self.model_type_instance.get_tts_model_voices(
model=self.model, credentials=self.credentials, language=language
model=self.model_name, credentials=self.credentials, language=language
)
+1 -1
View File
@@ -3,7 +3,6 @@ from typing import cast
from core.app.entities.app_invoke_entities import ModelConfigWithCredentialsEntity
from core.helper.code_executor.jinja2.jinja2_formatter import Jinja2Formatter
from core.memory.token_buffer_memory import TokenBufferMemory
from core.model_runtime.entities import (
AssistantPromptMessage,
PromptMessage,
@@ -13,6 +12,7 @@ from core.model_runtime.entities import (
UserPromptMessage,
)
from core.model_runtime.entities.message_entities import ImagePromptMessageContent, PromptMessageContentUnionTypes
from core.model_runtime.token_buffer_memory import TokenBufferMemory
from core.prompt.entities.advanced_prompt_entities import ChatModelMessage, CompletionModelPromptTemplate, MemoryConfig
from core.prompt.prompt_transform import PromptTransform
from core.prompt.utils.prompt_template_parser import PromptTemplateParser
@@ -3,13 +3,13 @@ from typing import cast
from core.app.entities.app_invoke_entities import (
ModelConfigWithCredentialsEntity,
)
from core.memory.token_buffer_memory import TokenBufferMemory
from core.model_runtime.entities.message_entities import (
PromptMessage,
SystemPromptMessage,
UserPromptMessage,
)
from core.model_runtime.model_providers.__base.large_language_model import LargeLanguageModel
from core.model_runtime.token_buffer_memory import TokenBufferMemory
from core.prompt.prompt_transform import PromptTransform
@@ -47,7 +47,9 @@ class AgentHistoryPromptTransform(PromptTransform):
model_type_instance = cast(LargeLanguageModel, model_type_instance)
curr_message_tokens = model_type_instance.get_num_tokens(
self.memory.model_instance.model, self.memory.model_instance.credentials, self.history_messages
self.model_config.model,
self.model_config.credentials,
self.history_messages,
)
if curr_message_tokens <= max_token_limit:
return self.history_messages
@@ -63,7 +65,9 @@ class AgentHistoryPromptTransform(PromptTransform):
# a message is start with UserPromptMessage
if isinstance(prompt_message, UserPromptMessage):
curr_message_tokens = model_type_instance.get_num_tokens(
self.memory.model_instance.model, self.memory.model_instance.credentials, prompt_messages
self.model_config.model,
self.model_config.credentials,
prompt_messages,
)
# if current message token is overflow, drop all the prompts in current message and break
if curr_message_tokens > max_token_limit:
+1 -1
View File
@@ -1,10 +1,10 @@
from typing import Any
from core.app.entities.app_invoke_entities import ModelConfigWithCredentialsEntity
from core.memory.token_buffer_memory import TokenBufferMemory
from core.model_manager import ModelInstance
from core.model_runtime.entities.message_entities import PromptMessage
from core.model_runtime.entities.model_entities import ModelPropertyKey
from core.model_runtime.token_buffer_memory import TokenBufferMemory
from core.prompt.entities.advanced_prompt_entities import MemoryConfig
+1 -1
View File
@@ -6,7 +6,6 @@ from typing import TYPE_CHECKING, Any, cast
from core.app.app_config.entities import PromptTemplateEntity
from core.app.entities.app_invoke_entities import ModelConfigWithCredentialsEntity
from core.memory.token_buffer_memory import TokenBufferMemory
from core.model_runtime.entities.message_entities import (
ImagePromptMessageContent,
PromptMessage,
@@ -15,6 +14,7 @@ from core.model_runtime.entities.message_entities import (
TextPromptMessageContent,
UserPromptMessage,
)
from core.model_runtime.token_buffer_memory import TokenBufferMemory
from core.prompt.entities.advanced_prompt_entities import MemoryConfig
from core.prompt.prompt_transform import PromptTransform
from core.prompt.utils.prompt_template_parser import PromptTemplateParser
+12 -8
View File
@@ -35,7 +35,9 @@ class CacheEmbedding(Embeddings):
embedding = (
db.session.query(Embedding)
.filter_by(
model_name=self._model_instance.model, hash=hash, provider_name=self._model_instance.provider
model_name=self._model_instance.model_name,
hash=hash,
provider_name=self._model_instance.provider,
)
.first()
)
@@ -52,7 +54,7 @@ class CacheEmbedding(Embeddings):
try:
model_type_instance = cast(TextEmbeddingModel, self._model_instance.model_type_instance)
model_schema = model_type_instance.get_model_schema(
self._model_instance.model, self._model_instance.credentials
self._model_instance.model_name, self._model_instance.credentials
)
max_chunks = (
model_schema.model_properties[ModelPropertyKey.MAX_CHUNKS]
@@ -87,7 +89,7 @@ class CacheEmbedding(Embeddings):
hash = helper.generate_text_hash(texts[i])
if hash not in cache_embeddings:
embedding_cache = Embedding(
model_name=self._model_instance.model,
model_name=self._model_instance.model_name,
hash=hash,
provider_name=self._model_instance.provider,
embedding=pickle.dumps(n_embedding, protocol=pickle.HIGHEST_PROTOCOL),
@@ -114,7 +116,9 @@ class CacheEmbedding(Embeddings):
embedding = (
db.session.query(Embedding)
.filter_by(
model_name=self._model_instance.model, hash=file_id, provider_name=self._model_instance.provider
model_name=self._model_instance.model_name,
hash=file_id,
provider_name=self._model_instance.provider,
)
.first()
)
@@ -131,7 +135,7 @@ class CacheEmbedding(Embeddings):
try:
model_type_instance = cast(TextEmbeddingModel, self._model_instance.model_type_instance)
model_schema = model_type_instance.get_model_schema(
self._model_instance.model, self._model_instance.credentials
self._model_instance.model_name, self._model_instance.credentials
)
max_chunks = (
model_schema.model_properties[ModelPropertyKey.MAX_CHUNKS]
@@ -168,7 +172,7 @@ class CacheEmbedding(Embeddings):
file_id = multimodel_documents[i]["file_id"]
if file_id not in cache_embeddings:
embedding_cache = Embedding(
model_name=self._model_instance.model,
model_name=self._model_instance.model_name,
hash=file_id,
provider_name=self._model_instance.provider,
embedding=pickle.dumps(n_embedding, protocol=pickle.HIGHEST_PROTOCOL),
@@ -190,7 +194,7 @@ class CacheEmbedding(Embeddings):
"""Embed query text."""
# use doc embedding cache or store if not exists
hash = helper.generate_text_hash(text)
embedding_cache_key = f"{self._model_instance.provider}_{self._model_instance.model}_{hash}"
embedding_cache_key = f"{self._model_instance.provider}_{self._model_instance.model_name}_{hash}"
embedding = redis_client.get(embedding_cache_key)
if embedding:
redis_client.expire(embedding_cache_key, 600)
@@ -233,7 +237,7 @@ class CacheEmbedding(Embeddings):
"""Embed multimodal documents."""
# use doc embedding cache or store if not exists
file_id = multimodel_document["file_id"]
embedding_cache_key = f"{self._model_instance.provider}_{self._model_instance.model}_{file_id}"
embedding_cache_key = f"{self._model_instance.provider}_{self._model_instance.model_name}_{file_id}"
embedding = redis_client.get(embedding_cache_key)
if embedding:
redis_client.expire(embedding_cache_key, 600)
+1 -1
View File
@@ -38,7 +38,7 @@ class RerankModelRunner(BaseRerankRunner):
is_support_vision = model_manager.check_model_support_vision(
tenant_id=self.rerank_model_instance.provider_model_bundle.configuration.tenant_id,
provider=self.rerank_model_instance.provider,
model=self.rerank_model_instance.model,
model=self.rerank_model_instance.model_name,
model_type=ModelType.RERANK,
)
if not is_support_vision:
+1 -1
View File
@@ -23,12 +23,12 @@ from core.callback_handler.index_tool_callback_handler import DatasetIndexToolCa
from core.db.session_factory import session_factory
from core.entities.agent_entities import PlanningStrategy
from core.entities.model_entities import ModelStatus
from core.memory.token_buffer_memory import TokenBufferMemory
from core.model_manager import ModelInstance, ModelManager
from core.model_runtime.entities.llm_entities import LLMResult, LLMUsage
from core.model_runtime.entities.message_entities import PromptMessage, PromptMessageRole, PromptMessageTool
from core.model_runtime.entities.model_entities import ModelFeature, ModelType
from core.model_runtime.model_providers.__base.large_language_model import LargeLanguageModel
from core.model_runtime.token_buffer_memory import TokenBufferMemory
from core.ops.entities.trace_entity import TraceTaskName
from core.ops.ops_trace_manager import TraceQueueManager, TraceTask
from core.ops.utils import measure_time
@@ -47,7 +47,7 @@ class ModelInvocationUtils:
raise InvokeModelError("Model not found")
llm_model = cast(LargeLanguageModel, model_instance.model_type_instance)
schema = llm_model.get_model_schema(model_instance.model, model_instance.credentials)
schema = llm_model.get_model_schema(model_instance.model_name, model_instance.credentials)
if not schema:
raise InvokeModelError("No model schema found")
@@ -7,12 +7,28 @@ Each instance uses a unique key for its command queue.
"""
import json
from typing import TYPE_CHECKING, Any, final
from contextlib import AbstractContextManager
from typing import Any, Protocol, final
from ..entities.commands import AbortCommand, CommandType, GraphEngineCommand, PauseCommand, UpdateVariablesCommand
if TYPE_CHECKING:
from extensions.ext_redis import RedisClientWrapper
class RedisPipelineProtocol(Protocol):
"""Minimal Redis pipeline contract used by the command channel."""
def lrange(self, name: str, start: int, end: int) -> Any: ...
def delete(self, *names: str) -> Any: ...
def execute(self) -> list[Any]: ...
def rpush(self, name: str, *values: str) -> Any: ...
def expire(self, name: str, time: int) -> Any: ...
def set(self, name: str, value: str, ex: int | None = None) -> Any: ...
def get(self, name: str) -> Any: ...
class RedisClientProtocol(Protocol):
"""Redis client contract required by the command channel."""
def pipeline(self) -> AbstractContextManager[RedisPipelineProtocol]: ...
@final
@@ -26,7 +42,7 @@ class RedisChannel:
def __init__(
self,
redis_client: "RedisClientWrapper",
redis_client: RedisClientProtocol,
channel_key: str,
command_ttl: int = 3600,
) -> None:
+15 -14
View File
@@ -3,13 +3,14 @@ GraphEngine Manager for sending control commands via Redis channel.
This module provides a simplified interface for controlling workflow executions
using the new Redis command channel, without requiring user permission checks.
Callers must provide a Redis client dependency from outside the workflow package.
"""
import logging
from collections.abc import Sequence
from typing import final
from core.workflow.graph_engine.command_channels.redis_channel import RedisChannel
from core.workflow.graph_engine.command_channels.redis_channel import RedisChannel, RedisClientProtocol
from core.workflow.graph_engine.entities.commands import (
AbortCommand,
GraphEngineCommand,
@@ -17,7 +18,6 @@ from core.workflow.graph_engine.entities.commands import (
UpdateVariablesCommand,
VariableUpdate,
)
from extensions.ext_redis import redis_client
logger = logging.getLogger(__name__)
@@ -31,8 +31,12 @@ class GraphEngineManager:
by sending commands through Redis channels, without user validation.
"""
@staticmethod
def send_stop_command(task_id: str, reason: str | None = None) -> None:
_redis_client: RedisClientProtocol
def __init__(self, redis_client: RedisClientProtocol) -> None:
self._redis_client = redis_client
def send_stop_command(self, task_id: str, reason: str | None = None) -> None:
"""
Send a stop command to a running workflow.
@@ -41,34 +45,31 @@ class GraphEngineManager:
reason: Optional reason for stopping (defaults to "User requested stop")
"""
abort_command = AbortCommand(reason=reason or "User requested stop")
GraphEngineManager._send_command(task_id, abort_command)
self._send_command(task_id, abort_command)
@staticmethod
def send_pause_command(task_id: str, reason: str | None = None) -> None:
def send_pause_command(self, task_id: str, reason: str | None = None) -> None:
"""Send a pause command to a running workflow."""
pause_command = PauseCommand(reason=reason or "User requested pause")
GraphEngineManager._send_command(task_id, pause_command)
self._send_command(task_id, pause_command)
@staticmethod
def send_update_variables_command(task_id: str, updates: Sequence[VariableUpdate]) -> None:
def send_update_variables_command(self, task_id: str, updates: Sequence[VariableUpdate]) -> None:
"""Send a command to update variables in a running workflow."""
if not updates:
return
update_command = UpdateVariablesCommand(updates=updates)
GraphEngineManager._send_command(task_id, update_command)
self._send_command(task_id, update_command)
@staticmethod
def _send_command(task_id: str, command: GraphEngineCommand) -> None:
def _send_command(self, task_id: str, command: GraphEngineCommand) -> None:
"""Send a command to the workflow-specific Redis channel."""
if not task_id:
return
channel_key = f"workflow:{task_id}:commands"
channel = RedisChannel(redis_client, channel_key)
channel = RedisChannel(self._redis_client, channel_key)
try:
channel.send_command(command)
+1 -1
View File
@@ -11,10 +11,10 @@ from sqlalchemy.orm import Session
from core.agent.entities import AgentToolEntity
from core.agent.plugin_entities import AgentStrategyParameter
from core.memory.token_buffer_memory import TokenBufferMemory
from core.model_manager import ModelInstance, ModelManager
from core.model_runtime.entities.llm_entities import LLMUsage, LLMUsageMetadata
from core.model_runtime.entities.model_entities import AIModelEntity, ModelType
from core.model_runtime.token_buffer_memory import TokenBufferMemory
from core.model_runtime.utils.encoders import jsonable_encoder
from core.provider_manager import ProviderManager
from core.tools.entities.tool_entities import (
+24 -7
View File
@@ -1,8 +1,7 @@
from collections.abc import Mapping, Sequence
from decimal import Decimal
from typing import TYPE_CHECKING, Any, ClassVar, cast
from typing import TYPE_CHECKING, Any, ClassVar, Protocol, cast
from core.helper.code_executor.code_executor import CodeExecutionError, CodeExecutor, CodeLanguage
from core.helper.code_executor.code_node_provider import CodeNodeProvider
from core.helper.code_executor.javascript.javascript_code_provider import JavascriptCodeProvider
from core.helper.code_executor.python3.python3_code_provider import Python3CodeProvider
@@ -11,7 +10,7 @@ from core.variables.types import SegmentType
from core.workflow.enums import NodeType, WorkflowNodeExecutionStatus
from core.workflow.node_events import NodeRunResult
from core.workflow.nodes.base.node import Node
from core.workflow.nodes.code.entities import CodeNodeData
from core.workflow.nodes.code.entities import CodeLanguage, CodeNodeData
from core.workflow.nodes.code.limits import CodeNodeLimits
from .exc import (
@@ -25,6 +24,18 @@ if TYPE_CHECKING:
from core.workflow.runtime import GraphRuntimeState
class WorkflowCodeExecutor(Protocol):
def execute(
self,
*,
language: CodeLanguage,
code: str,
inputs: Mapping[str, Any],
) -> Mapping[str, Any]: ...
def is_execution_error(self, error: Exception) -> bool: ...
class CodeNode(Node[CodeNodeData]):
node_type = NodeType.CODE
_DEFAULT_CODE_PROVIDERS: ClassVar[tuple[type[CodeNodeProvider], ...]] = (
@@ -40,7 +51,7 @@ class CodeNode(Node[CodeNodeData]):
graph_init_params: "GraphInitParams",
graph_runtime_state: "GraphRuntimeState",
*,
code_executor: type[CodeExecutor] | None = None,
code_executor: WorkflowCodeExecutor,
code_providers: Sequence[type[CodeNodeProvider]] | None = None,
code_limits: CodeNodeLimits,
) -> None:
@@ -50,7 +61,7 @@ class CodeNode(Node[CodeNodeData]):
graph_init_params=graph_init_params,
graph_runtime_state=graph_runtime_state,
)
self._code_executor: type[CodeExecutor] = code_executor or CodeExecutor
self._code_executor: WorkflowCodeExecutor = code_executor
self._code_providers: tuple[type[CodeNodeProvider], ...] = (
tuple(code_providers) if code_providers else self._DEFAULT_CODE_PROVIDERS
)
@@ -98,7 +109,7 @@ class CodeNode(Node[CodeNodeData]):
# Run code
try:
_ = self._select_code_provider(code_language)
result = self._code_executor.execute_workflow_code_template(
result = self._code_executor.execute(
language=code_language,
code=code,
inputs=variables,
@@ -106,7 +117,13 @@ class CodeNode(Node[CodeNodeData]):
# Transform result
result = self._transform_result(result=result, output_schema=self.node_data.outputs)
except (CodeExecutionError, CodeNodeError) as e:
except CodeNodeError as e:
return NodeRunResult(
status=WorkflowNodeExecutionStatus.FAILED, inputs=variables, error=str(e), error_type=type(e).__name__
)
except Exception as e:
if not self._code_executor.is_execution_error(e):
raise
return NodeRunResult(
status=WorkflowNodeExecutionStatus.FAILED, inputs=variables, error=str(e), error_type=type(e).__name__
)
@@ -1,40 +1,26 @@
from collections.abc import Generator, Mapping, Sequence
from typing import Any, cast
from typing import TYPE_CHECKING, Any
from sqlalchemy import select
from sqlalchemy.orm import Session
from core.datasource.entities.datasource_entities import (
DatasourceMessage,
DatasourceParameter,
DatasourceProviderType,
GetOnlineDocumentPageContentRequest,
OnlineDriveDownloadFileRequest,
)
from core.datasource.online_document.online_document_plugin import OnlineDocumentDatasourcePlugin
from core.datasource.online_drive.online_drive_plugin import OnlineDriveDatasourcePlugin
from core.datasource.utils.message_transformer import DatasourceFileMessageTransformer
from core.datasource.entities.datasource_entities import DatasourceProviderType
from core.plugin.impl.exc import PluginDaemonClientSideError
from core.variables.segments import ArrayAnySegment
from core.variables.variables import ArrayAnyVariable
from core.workflow.entities.workflow_node_execution import WorkflowNodeExecutionStatus
from core.workflow.enums import NodeExecutionType, NodeType, SystemVariableKey
from core.workflow.file import File
from core.workflow.file.enums import FileTransferMethod, FileType
from core.workflow.node_events import NodeRunResult, StreamChunkEvent, StreamCompletedEvent
from core.workflow.node_events import NodeRunResult, StreamCompletedEvent
from core.workflow.nodes.base.node import Node
from core.workflow.nodes.base.variable_template_parser import VariableTemplateParser
from core.workflow.nodes.tool.exc import ToolFileError
from core.workflow.runtime import VariablePool
from extensions.ext_database import db
from factories import file_factory
from models.model import UploadFile
from models.tools import ToolFile
from services.datasource_provider_service import DatasourceProviderService
from core.workflow.repositories.datasource_manager_protocol import (
DatasourceManagerProtocol,
DatasourceParameter,
OnlineDriveDownloadFileParam,
)
from ...entities.workflow_node_execution import WorkflowNodeExecutionMetadataKey
from .entities import DatasourceNodeData
from .exc import DatasourceNodeError, DatasourceParameterError
from .exc import DatasourceNodeError
if TYPE_CHECKING:
from core.workflow.entities import GraphInitParams
from core.workflow.runtime import GraphRuntimeState
class DatasourceNode(Node[DatasourceNodeData]):
@@ -45,6 +31,22 @@ class DatasourceNode(Node[DatasourceNodeData]):
node_type = NodeType.DATASOURCE
execution_type = NodeExecutionType.ROOT
def __init__(
self,
id: str,
config: Mapping[str, Any],
graph_init_params: "GraphInitParams",
graph_runtime_state: "GraphRuntimeState",
datasource_manager: DatasourceManagerProtocol,
):
super().__init__(
id=id,
config=config,
graph_init_params=graph_init_params,
graph_runtime_state=graph_runtime_state,
)
self.datasource_manager = datasource_manager
def _run(self) -> Generator:
"""
Run the datasource node
@@ -52,84 +54,69 @@ class DatasourceNode(Node[DatasourceNodeData]):
node_data = self.node_data
variable_pool = self.graph_runtime_state.variable_pool
datasource_type_segement = variable_pool.get(["sys", SystemVariableKey.DATASOURCE_TYPE])
if not datasource_type_segement:
datasource_type_segment = variable_pool.get(["sys", SystemVariableKey.DATASOURCE_TYPE])
if not datasource_type_segment:
raise DatasourceNodeError("Datasource type is not set")
datasource_type = str(datasource_type_segement.value) if datasource_type_segement.value else None
datasource_info_segement = variable_pool.get(["sys", SystemVariableKey.DATASOURCE_INFO])
if not datasource_info_segement:
datasource_type = str(datasource_type_segment.value) if datasource_type_segment.value else None
datasource_info_segment = variable_pool.get(["sys", SystemVariableKey.DATASOURCE_INFO])
if not datasource_info_segment:
raise DatasourceNodeError("Datasource info is not set")
datasource_info_value = datasource_info_segement.value
datasource_info_value = datasource_info_segment.value
if not isinstance(datasource_info_value, dict):
raise DatasourceNodeError("Invalid datasource info format")
datasource_info: dict[str, Any] = datasource_info_value
# get datasource runtime
from core.datasource.datasource_manager import DatasourceManager
if datasource_type is None:
raise DatasourceNodeError("Datasource type is not set")
datasource_type = DatasourceProviderType.value_of(datasource_type)
provider_id = f"{node_data.plugin_id}/{node_data.provider_name}"
datasource_runtime = DatasourceManager.get_datasource_runtime(
provider_id=f"{node_data.plugin_id}/{node_data.provider_name}",
datasource_info["icon"] = self.datasource_manager.get_icon_url(
provider_id=provider_id,
datasource_name=node_data.datasource_name or "",
tenant_id=self.tenant_id,
datasource_type=datasource_type,
datasource_type=datasource_type.value,
)
datasource_info["icon"] = datasource_runtime.get_icon_url(self.tenant_id)
parameters_for_log = datasource_info
try:
datasource_provider_service = DatasourceProviderService()
credentials = datasource_provider_service.get_datasource_credentials(
tenant_id=self.tenant_id,
provider=node_data.provider_name,
plugin_id=node_data.plugin_id,
credential_id=datasource_info.get("credential_id", ""),
)
match datasource_type:
case DatasourceProviderType.ONLINE_DOCUMENT:
datasource_runtime = cast(OnlineDocumentDatasourcePlugin, datasource_runtime)
if credentials:
datasource_runtime.runtime.credentials = credentials
online_document_result: Generator[DatasourceMessage, None, None] = (
datasource_runtime.get_online_document_page_content(
user_id=self.user_id,
datasource_parameters=GetOnlineDocumentPageContentRequest(
workspace_id=datasource_info.get("workspace_id", ""),
page_id=datasource_info.get("page", {}).get("page_id", ""),
type=datasource_info.get("page", {}).get("type", ""),
),
provider_type=datasource_type,
case DatasourceProviderType.ONLINE_DOCUMENT | DatasourceProviderType.ONLINE_DRIVE:
# Build typed request objects
datasource_parameters = None
if datasource_type == DatasourceProviderType.ONLINE_DOCUMENT:
datasource_parameters = DatasourceParameter(
workspace_id=datasource_info.get("workspace_id", ""),
page_id=datasource_info.get("page", {}).get("page_id", ""),
type=datasource_info.get("page", {}).get("type", ""),
)
)
yield from self._transform_message(
messages=online_document_result,
parameters_for_log=parameters_for_log,
datasource_info=datasource_info,
)
case DatasourceProviderType.ONLINE_DRIVE:
datasource_runtime = cast(OnlineDriveDatasourcePlugin, datasource_runtime)
if credentials:
datasource_runtime.runtime.credentials = credentials
online_drive_result: Generator[DatasourceMessage, None, None] = (
datasource_runtime.online_drive_download_file(
user_id=self.user_id,
request=OnlineDriveDownloadFileRequest(
id=datasource_info.get("id", ""),
bucket=datasource_info.get("bucket"),
),
provider_type=datasource_type,
online_drive_request = None
if datasource_type == DatasourceProviderType.ONLINE_DRIVE:
online_drive_request = OnlineDriveDownloadFileParam(
id=datasource_info.get("id", ""),
bucket=datasource_info.get("bucket", ""),
)
)
yield from self._transform_datasource_file_message(
messages=online_drive_result,
credential_id = datasource_info.get("credential_id", "")
yield from self.datasource_manager.stream_node_events(
node_id=self._node_id,
user_id=self.user_id,
datasource_name=node_data.datasource_name or "",
datasource_type=datasource_type.value,
provider_id=provider_id,
tenant_id=self.tenant_id,
provider=node_data.provider_name,
plugin_id=node_data.plugin_id,
credential_id=credential_id,
parameters_for_log=parameters_for_log,
datasource_info=datasource_info,
variable_pool=variable_pool,
datasource_type=datasource_type,
datasource_param=datasource_parameters,
online_drive_request=online_drive_request,
)
case DatasourceProviderType.WEBSITE_CRAWL:
yield StreamCompletedEvent(
@@ -147,23 +134,9 @@ class DatasourceNode(Node[DatasourceNodeData]):
related_id = datasource_info.get("related_id")
if not related_id:
raise DatasourceNodeError("File is not exist")
upload_file = db.session.query(UploadFile).where(UploadFile.id == related_id).first()
if not upload_file:
raise ValueError("Invalid upload file Info")
file_info = File(
id=upload_file.id,
filename=upload_file.name,
extension="." + upload_file.extension,
mime_type=upload_file.mime_type,
tenant_id=self.tenant_id,
type=FileType.CUSTOM,
transfer_method=FileTransferMethod.LOCAL_FILE,
remote_url=upload_file.source_url,
related_id=upload_file.id,
size=upload_file.size,
storage_key=upload_file.key,
url=upload_file.source_url,
file_info = self.datasource_manager.get_upload_file_by_id(
file_id=related_id, tenant_id=self.tenant_id
)
variable_pool.add([self._node_id, "file"], file_info)
# variable_pool.add([self.node_id, "file"], file_info.to_dict())
@@ -201,55 +174,6 @@ class DatasourceNode(Node[DatasourceNodeData]):
)
)
def _generate_parameters(
self,
*,
datasource_parameters: Sequence[DatasourceParameter],
variable_pool: VariablePool,
node_data: DatasourceNodeData,
for_log: bool = False,
) -> dict[str, Any]:
"""
Generate parameters based on the given tool parameters, variable pool, and node data.
Args:
tool_parameters (Sequence[ToolParameter]): The list of tool parameters.
variable_pool (VariablePool): The variable pool containing the variables.
node_data (ToolNodeData): The data associated with the tool node.
Returns:
Mapping[str, Any]: A dictionary containing the generated parameters.
"""
datasource_parameters_dictionary = {parameter.name: parameter for parameter in datasource_parameters}
result: dict[str, Any] = {}
if node_data.datasource_parameters:
for parameter_name in node_data.datasource_parameters:
parameter = datasource_parameters_dictionary.get(parameter_name)
if not parameter:
result[parameter_name] = None
continue
datasource_input = node_data.datasource_parameters[parameter_name]
if datasource_input.type == "variable":
variable = variable_pool.get(datasource_input.value)
if variable is None:
raise DatasourceParameterError(f"Variable {datasource_input.value} does not exist")
parameter_value = variable.value
elif datasource_input.type in {"mixed", "constant"}:
segment_group = variable_pool.convert_template(str(datasource_input.value))
parameter_value = segment_group.log if for_log else segment_group.text
else:
raise DatasourceParameterError(f"Unknown datasource input type '{datasource_input.type}'")
result[parameter_name] = parameter_value
return result
def _fetch_files(self, variable_pool: VariablePool) -> list[File]:
variable = variable_pool.get(["sys", SystemVariableKey.FILES])
assert isinstance(variable, ArrayAnyVariable | ArrayAnySegment)
return list(variable.value) if variable else []
@classmethod
def _extract_variable_selector_to_variable_mapping(
cls,
@@ -287,206 +211,6 @@ class DatasourceNode(Node[DatasourceNodeData]):
return result
def _transform_message(
self,
messages: Generator[DatasourceMessage, None, None],
parameters_for_log: dict[str, Any],
datasource_info: dict[str, Any],
) -> Generator:
"""
Convert ToolInvokeMessages into tuple[plain_text, files]
"""
# transform message and handle file storage
message_stream = DatasourceFileMessageTransformer.transform_datasource_invoke_messages(
messages=messages,
user_id=self.user_id,
tenant_id=self.tenant_id,
conversation_id=None,
)
text = ""
files: list[File] = []
json: list[dict | list] = []
variables: dict[str, Any] = {}
for message in message_stream:
match message.type:
case (
DatasourceMessage.MessageType.IMAGE_LINK
| DatasourceMessage.MessageType.BINARY_LINK
| DatasourceMessage.MessageType.IMAGE
):
assert isinstance(message.message, DatasourceMessage.TextMessage)
url = message.message.text
transfer_method = FileTransferMethod.TOOL_FILE
datasource_file_id = str(url).split("/")[-1].split(".")[0]
with Session(db.engine) as session:
stmt = select(ToolFile).where(ToolFile.id == datasource_file_id)
datasource_file = session.scalar(stmt)
if datasource_file is None:
raise ToolFileError(f"Tool file {datasource_file_id} does not exist")
mapping = {
"tool_file_id": datasource_file_id,
"type": file_factory.get_file_type_by_mime_type(datasource_file.mimetype),
"transfer_method": transfer_method,
"url": url,
}
file = file_factory.build_from_mapping(
mapping=mapping,
tenant_id=self.tenant_id,
)
files.append(file)
case DatasourceMessage.MessageType.BLOB:
# get tool file id
assert isinstance(message.message, DatasourceMessage.TextMessage)
assert message.meta
datasource_file_id = message.message.text.split("/")[-1].split(".")[0]
with Session(db.engine) as session:
stmt = select(ToolFile).where(ToolFile.id == datasource_file_id)
datasource_file = session.scalar(stmt)
if datasource_file is None:
raise ToolFileError(f"datasource file {datasource_file_id} not exists")
mapping = {
"tool_file_id": datasource_file_id,
"transfer_method": FileTransferMethod.TOOL_FILE,
}
files.append(
file_factory.build_from_mapping(
mapping=mapping,
tenant_id=self.tenant_id,
)
)
case DatasourceMessage.MessageType.TEXT:
assert isinstance(message.message, DatasourceMessage.TextMessage)
text += message.message.text
yield StreamChunkEvent(
selector=[self._node_id, "text"],
chunk=message.message.text,
is_final=False,
)
case DatasourceMessage.MessageType.JSON:
assert isinstance(message.message, DatasourceMessage.JsonMessage)
json.append(message.message.json_object)
case DatasourceMessage.MessageType.LINK:
assert isinstance(message.message, DatasourceMessage.TextMessage)
stream_text = f"Link: {message.message.text}\n"
text += stream_text
yield StreamChunkEvent(
selector=[self._node_id, "text"],
chunk=stream_text,
is_final=False,
)
case DatasourceMessage.MessageType.VARIABLE:
assert isinstance(message.message, DatasourceMessage.VariableMessage)
variable_name = message.message.variable_name
variable_value = message.message.variable_value
if message.message.stream:
if not isinstance(variable_value, str):
raise ValueError("When 'stream' is True, 'variable_value' must be a string.")
if variable_name not in variables:
variables[variable_name] = ""
variables[variable_name] += variable_value
yield StreamChunkEvent(
selector=[self._node_id, variable_name],
chunk=variable_value,
is_final=False,
)
else:
variables[variable_name] = variable_value
case DatasourceMessage.MessageType.FILE:
assert message.meta is not None
files.append(message.meta["file"])
case (
DatasourceMessage.MessageType.BLOB_CHUNK
| DatasourceMessage.MessageType.LOG
| DatasourceMessage.MessageType.RETRIEVER_RESOURCES
):
pass
# mark the end of the stream
yield StreamChunkEvent(
selector=[self._node_id, "text"],
chunk="",
is_final=True,
)
yield StreamCompletedEvent(
node_run_result=NodeRunResult(
status=WorkflowNodeExecutionStatus.SUCCEEDED,
outputs={**variables},
metadata={
WorkflowNodeExecutionMetadataKey.DATASOURCE_INFO: datasource_info,
},
inputs=parameters_for_log,
)
)
@classmethod
def version(cls) -> str:
return "1"
def _transform_datasource_file_message(
self,
messages: Generator[DatasourceMessage, None, None],
parameters_for_log: dict[str, Any],
datasource_info: dict[str, Any],
variable_pool: VariablePool,
datasource_type: DatasourceProviderType,
) -> Generator:
"""
Convert ToolInvokeMessages into tuple[plain_text, files]
"""
# transform message and handle file storage
message_stream = DatasourceFileMessageTransformer.transform_datasource_invoke_messages(
messages=messages,
user_id=self.user_id,
tenant_id=self.tenant_id,
conversation_id=None,
)
file = None
for message in message_stream:
if message.type == DatasourceMessage.MessageType.BINARY_LINK:
assert isinstance(message.message, DatasourceMessage.TextMessage)
url = message.message.text
transfer_method = FileTransferMethod.TOOL_FILE
datasource_file_id = str(url).split("/")[-1].split(".")[0]
with Session(db.engine) as session:
stmt = select(ToolFile).where(ToolFile.id == datasource_file_id)
datasource_file = session.scalar(stmt)
if datasource_file is None:
raise ToolFileError(f"Tool file {datasource_file_id} does not exist")
mapping = {
"tool_file_id": datasource_file_id,
"type": file_factory.get_file_type_by_mime_type(datasource_file.mimetype),
"transfer_method": transfer_method,
"url": url,
}
file = file_factory.build_from_mapping(
mapping=mapping,
tenant_id=self.tenant_id,
)
if file:
variable_pool.add([self._node_id, "file"], file)
yield StreamCompletedEvent(
node_run_result=NodeRunResult(
status=WorkflowNodeExecutionStatus.SUCCEEDED,
inputs=parameters_for_log,
metadata={WorkflowNodeExecutionMetadataKey.DATASOURCE_INFO: datasource_info},
outputs={
"file": file,
"datasource_type": datasource_type,
},
)
)
+21 -22
View File
@@ -7,16 +7,18 @@ from sqlalchemy.orm import Session
from configs import dify_config
from core.app.entities.app_invoke_entities import ModelConfigWithCredentialsEntity
from core.entities.provider_entities import ProviderQuotaType, QuotaUnit
from core.memory.token_buffer_memory import TokenBufferMemory
from core.model_manager import ModelInstance, ModelManager
from core.model_manager import ModelInstance
from core.model_runtime.entities.llm_entities import LLMUsage
from core.model_runtime.entities.model_entities import ModelType
from core.model_runtime.model_providers.__base.large_language_model import LargeLanguageModel
from core.model_runtime.token_buffer_memory import TokenBufferMemory
from core.prompt.entities.advanced_prompt_entities import MemoryConfig
from core.variables.segments import ArrayAnySegment, ArrayFileSegment, FileSegment, NoneSegment, StringSegment
from core.workflow.enums import SystemVariableKey
from core.workflow.file.models import File
from core.workflow.nodes.llm.entities import ModelConfig
from core.workflow.nodes.llm.exc import LLMModeRequiredError, ModelNotExistError
from core.workflow.nodes.llm.protocols import CredentialsProvider, ModelFactory
from core.workflow.runtime import VariablePool
from extensions.ext_database import db
from libs.datetime_utils import naive_utc_now
@@ -24,49 +26,46 @@ from models.model import Conversation
from models.provider import Provider, ProviderType
from models.provider_ids import ModelProviderID
from .exc import InvalidVariableTypeError, LLMModeRequiredError, ModelNotExistError
from .exc import InvalidVariableTypeError
def fetch_model_config(
tenant_id: str, node_data_model: ModelConfig
*,
node_data_model: ModelConfig,
credentials_provider: CredentialsProvider,
model_factory: ModelFactory,
) -> tuple[ModelInstance, ModelConfigWithCredentialsEntity]:
if not node_data_model.mode:
raise LLMModeRequiredError("LLM mode is required.")
model = ModelManager().get_model_instance(
tenant_id=tenant_id,
model_type=ModelType.LLM,
provider=node_data_model.provider,
credentials = credentials_provider.fetch(node_data_model.provider, node_data_model.name)
model_instance = model_factory.init_model_instance(node_data_model.provider, node_data_model.name)
provider_model_bundle = model_instance.provider_model_bundle
provider_model = provider_model_bundle.configuration.get_provider_model(
model=node_data_model.name,
model_type=ModelType.LLM,
)
model.model_type_instance = cast(LargeLanguageModel, model.model_type_instance)
# check model
provider_model = model.provider_model_bundle.configuration.get_provider_model(
model=node_data_model.name, model_type=ModelType.LLM
)
if provider_model is None:
raise ModelNotExistError(f"Model {node_data_model.name} not exist.")
provider_model.raise_for_status()
# model config
stop: list[str] = []
if "stop" in node_data_model.completion_params:
stop = node_data_model.completion_params.pop("stop")
model_schema = model.model_type_instance.get_model_schema(node_data_model.name, model.credentials)
model_schema = model_instance.model_type_instance.get_model_schema(node_data_model.name, credentials)
if not model_schema:
raise ModelNotExistError(f"Model {node_data_model.name} not exist.")
return model, ModelConfigWithCredentialsEntity(
model_instance.model_type_instance = cast(LargeLanguageModel, model_instance.model_type_instance)
return model_instance, ModelConfigWithCredentialsEntity(
provider=node_data_model.provider,
model=node_data_model.name,
model_schema=model_schema,
mode=node_data_model.mode,
provider_model_bundle=model.provider_model_bundle,
credentials=model.credentials,
provider_model_bundle=provider_model_bundle,
credentials=credentials,
parameters=node_data_model.completion_params,
stop=stop,
)
@@ -131,7 +130,7 @@ def deduct_llm_quota(tenant_id: str, model_instance: ModelInstance, usage: LLMUs
if quota_unit == QuotaUnit.TOKENS:
used_quota = usage.total_tokens
elif quota_unit == QuotaUnit.CREDITS:
used_quota = dify_config.get_model_credits(model_instance.model)
used_quota = dify_config.get_model_credits(model_instance.model_name)
else:
used_quota = 1
+79 -58
View File
@@ -15,8 +15,7 @@ from core.app.entities.app_invoke_entities import ModelConfigWithCredentialsEnti
from core.helper.code_executor import CodeExecutor, CodeLanguage
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.memory.token_buffer_memory import TokenBufferMemory
from core.model_manager import ModelInstance, ModelManager
from core.model_manager import ModelInstance
from core.model_runtime.entities import (
ImagePromptMessageContent,
PromptMessage,
@@ -38,11 +37,8 @@ from core.model_runtime.entities.message_entities import (
SystemPromptMessage,
UserPromptMessage,
)
from core.model_runtime.entities.model_entities import (
ModelFeature,
ModelPropertyKey,
ModelType,
)
from core.model_runtime.entities.model_entities import AIModelEntity, ModelFeature, ModelPropertyKey
from core.model_runtime.token_buffer_memory import TokenBufferMemory
from core.model_runtime.utils.encoders import jsonable_encoder
from core.prompt.entities.advanced_prompt_entities import CompletionModelPromptTemplate, MemoryConfig
from core.prompt.utils.prompt_message_util import PromptMessageUtil
@@ -76,6 +72,7 @@ from core.workflow.node_events import (
from core.workflow.nodes.base.entities import VariableSelector
from core.workflow.nodes.base.node import Node
from core.workflow.nodes.base.variable_template_parser import VariableTemplateParser
from core.workflow.nodes.llm.protocols import CredentialsProvider, ModelFactory
from core.workflow.runtime import VariablePool
from extensions.ext_database import db
from models.dataset import SegmentAttachmentBinding
@@ -93,7 +90,6 @@ from .exc import (
InvalidVariableTypeError,
LLMNodeError,
MemoryRolePrefixRequiredError,
ModelNotExistError,
NoPromptFoundError,
TemplateTypeNotSupportError,
VariableNotFoundError,
@@ -118,6 +114,8 @@ class LLMNode(Node[LLMNodeData]):
_file_outputs: list[File]
_llm_file_saver: LLMFileSaver
_credentials_provider: CredentialsProvider
_model_factory: ModelFactory
def __init__(
self,
@@ -126,6 +124,8 @@ class LLMNode(Node[LLMNodeData]):
graph_init_params: GraphInitParams,
graph_runtime_state: GraphRuntimeState,
*,
credentials_provider: CredentialsProvider,
model_factory: ModelFactory,
llm_file_saver: LLMFileSaver | None = None,
):
super().__init__(
@@ -137,6 +137,9 @@ class LLMNode(Node[LLMNodeData]):
# LLM file outputs, used for MultiModal outputs.
self._file_outputs = []
self._credentials_provider = credentials_provider
self._model_factory = model_factory
if llm_file_saver is None:
llm_file_saver = FileSaverImpl(
user_id=graph_init_params.user_id,
@@ -199,10 +202,21 @@ class LLMNode(Node[LLMNodeData]):
node_inputs["#context_files#"] = [file.model_dump() for file in context_files]
# fetch model config
model_instance, model_config = LLMNode._fetch_model_config(
model_instance, model_config = self._fetch_model_config(
node_data_model=self.node_data.model,
tenant_id=self.tenant_id,
)
model_name = getattr(model_instance, "model_name", None)
if not isinstance(model_name, str):
model_name = model_config.model
model_provider = getattr(model_instance, "provider", None)
if not isinstance(model_provider, str):
model_provider = model_config.provider
model_schema = model_instance.model_type_instance.get_model_schema(
model_name,
model_instance.credentials,
)
if not model_schema:
raise ValueError(f"Model schema not found for {model_name}")
# fetch memory
memory = llm_utils.fetch_memory(
@@ -225,14 +239,16 @@ class LLMNode(Node[LLMNodeData]):
sys_files=files,
context=context,
memory=memory,
model_config=model_config,
model_instance=model_instance,
model_schema=model_schema,
model_parameters=self.node_data.model.completion_params,
stop=model_config.stop,
prompt_template=self.node_data.prompt_template,
memory_config=self.node_data.memory,
vision_enabled=self.node_data.vision.enabled,
vision_detail=self.node_data.vision.configs.detail,
variable_pool=variable_pool,
jinja2_variables=self.node_data.prompt_config.jinja2_variables,
tenant_id=self.tenant_id,
context_files=context_files,
)
@@ -286,14 +302,14 @@ class LLMNode(Node[LLMNodeData]):
structured_output = event
process_data = {
"model_mode": model_config.mode,
"model_mode": self.node_data.model.mode,
"prompts": PromptMessageUtil.prompt_messages_to_prompt_for_saving(
model_mode=model_config.mode, prompt_messages=prompt_messages
model_mode=self.node_data.model.mode, prompt_messages=prompt_messages
),
"usage": jsonable_encoder(usage),
"finish_reason": finish_reason,
"model_provider": model_config.provider,
"model_name": model_config.model,
"model_provider": model_provider,
"model_name": model_name,
}
outputs = {
@@ -755,21 +771,18 @@ class LLMNode(Node[LLMNodeData]):
return None
@staticmethod
def _fetch_model_config(
self,
*,
node_data_model: ModelConfig,
tenant_id: str,
) -> tuple[ModelInstance, ModelConfigWithCredentialsEntity]:
model, model_config_with_cred = llm_utils.fetch_model_config(
tenant_id=tenant_id, node_data_model=node_data_model
node_data_model=node_data_model,
credentials_provider=self._credentials_provider,
model_factory=self._model_factory,
)
completion_params = model_config_with_cred.parameters
model_schema = model.model_type_instance.get_model_schema(node_data_model.name, model.credentials)
if not model_schema:
raise ModelNotExistError(f"Model {node_data_model.name} not exist.")
model_config_with_cred.parameters = completion_params
# NOTE(-LAN-): This line modify the `self.node_data.model`, which is used in `_invoke_llm()`.
node_data_model.completion_params = completion_params
@@ -782,14 +795,16 @@ class LLMNode(Node[LLMNodeData]):
sys_files: Sequence[File],
context: str | None = None,
memory: TokenBufferMemory | None = None,
model_config: ModelConfigWithCredentialsEntity,
model_instance: ModelInstance,
model_schema: AIModelEntity,
model_parameters: Mapping[str, Any],
prompt_template: Sequence[LLMNodeChatModelMessage] | LLMNodeCompletionModelPromptTemplate,
stop: Sequence[str] | None = None,
memory_config: MemoryConfig | None = None,
vision_enabled: bool = False,
vision_detail: ImagePromptMessageContent.DETAIL,
variable_pool: VariablePool,
jinja2_variables: Sequence[VariableSelector],
tenant_id: str,
context_files: list[File] | None = None,
) -> tuple[Sequence[PromptMessage], Sequence[str] | None]:
prompt_messages: list[PromptMessage] = []
@@ -810,7 +825,9 @@ class LLMNode(Node[LLMNodeData]):
memory_messages = _handle_memory_chat_mode(
memory=memory,
memory_config=memory_config,
model_config=model_config,
model_instance=model_instance,
model_schema=model_schema,
model_parameters=model_parameters,
)
# Extend prompt_messages with memory messages
prompt_messages.extend(memory_messages)
@@ -847,7 +864,9 @@ class LLMNode(Node[LLMNodeData]):
memory_text = _handle_memory_completion_mode(
memory=memory,
memory_config=memory_config,
model_config=model_config,
model_instance=model_instance,
model_schema=model_schema,
model_parameters=model_parameters,
)
# Insert histories into the prompt
prompt_content = prompt_messages[0].content
@@ -924,7 +943,7 @@ class LLMNode(Node[LLMNodeData]):
prompt_message_content: list[PromptMessageContentUnionTypes] = []
for content_item in prompt_message.content:
# Skip content if features are not defined
if not model_config.model_schema.features:
if not model_schema.features:
if content_item.type != PromptMessageContentType.TEXT:
continue
prompt_message_content.append(content_item)
@@ -934,19 +953,19 @@ class LLMNode(Node[LLMNodeData]):
if (
(
content_item.type == PromptMessageContentType.IMAGE
and ModelFeature.VISION not in model_config.model_schema.features
and ModelFeature.VISION not in model_schema.features
)
or (
content_item.type == PromptMessageContentType.DOCUMENT
and ModelFeature.DOCUMENT not in model_config.model_schema.features
and ModelFeature.DOCUMENT not in model_schema.features
)
or (
content_item.type == PromptMessageContentType.VIDEO
and ModelFeature.VIDEO not in model_config.model_schema.features
and ModelFeature.VIDEO not in model_schema.features
)
or (
content_item.type == PromptMessageContentType.AUDIO
and ModelFeature.AUDIO not in model_config.model_schema.features
and ModelFeature.AUDIO not in model_schema.features
)
):
continue
@@ -965,19 +984,7 @@ class LLMNode(Node[LLMNodeData]):
"Please ensure a prompt is properly configured before proceeding."
)
model = ModelManager().get_model_instance(
tenant_id=tenant_id,
model_type=ModelType.LLM,
provider=model_config.provider,
model=model_config.model,
)
model_schema = model.model_type_instance.get_model_schema(
model=model_config.model,
credentials=model.credentials,
)
if not model_schema:
raise ModelNotExistError(f"Model {model_config.model} not exist.")
return filtered_prompt_messages, model_config.stop
return filtered_prompt_messages, stop
@classmethod
def _extract_variable_selector_to_variable_mapping(
@@ -1306,26 +1313,26 @@ def _render_jinja2_message(
def _calculate_rest_token(
*, prompt_messages: list[PromptMessage], model_config: ModelConfigWithCredentialsEntity
*,
prompt_messages: list[PromptMessage],
model_instance: ModelInstance,
model_schema: AIModelEntity,
model_parameters: Mapping[str, Any],
) -> int:
rest_tokens = 2000
model_context_tokens = model_config.model_schema.model_properties.get(ModelPropertyKey.CONTEXT_SIZE)
model_context_tokens = model_schema.model_properties.get(ModelPropertyKey.CONTEXT_SIZE)
if model_context_tokens:
model_instance = ModelInstance(
provider_model_bundle=model_config.provider_model_bundle, model=model_config.model
)
curr_message_tokens = model_instance.get_llm_num_tokens(prompt_messages)
max_tokens = 0
for parameter_rule in model_config.model_schema.parameter_rules:
for parameter_rule in model_schema.parameter_rules:
if parameter_rule.name == "max_tokens" or (
parameter_rule.use_template and parameter_rule.use_template == "max_tokens"
):
max_tokens = (
model_config.parameters.get(parameter_rule.name)
or model_config.parameters.get(str(parameter_rule.use_template))
model_parameters.get(parameter_rule.name)
or model_parameters.get(str(parameter_rule.use_template))
or 0
)
@@ -1339,12 +1346,19 @@ def _handle_memory_chat_mode(
*,
memory: TokenBufferMemory | None,
memory_config: MemoryConfig | None,
model_config: ModelConfigWithCredentialsEntity,
model_instance: ModelInstance,
model_schema: AIModelEntity,
model_parameters: Mapping[str, Any],
) -> Sequence[PromptMessage]:
memory_messages: Sequence[PromptMessage] = []
# Get messages from memory for chat model
if memory and memory_config:
rest_tokens = _calculate_rest_token(prompt_messages=[], model_config=model_config)
rest_tokens = _calculate_rest_token(
prompt_messages=[],
model_instance=model_instance,
model_schema=model_schema,
model_parameters=model_parameters,
)
memory_messages = memory.get_history_prompt_messages(
max_token_limit=rest_tokens,
message_limit=memory_config.window.size if memory_config.window.enabled else None,
@@ -1356,12 +1370,19 @@ def _handle_memory_completion_mode(
*,
memory: TokenBufferMemory | None,
memory_config: MemoryConfig | None,
model_config: ModelConfigWithCredentialsEntity,
model_instance: ModelInstance,
model_schema: AIModelEntity,
model_parameters: Mapping[str, Any],
) -> str:
memory_text = ""
# Get history text from memory for completion model
if memory and memory_config:
rest_tokens = _calculate_rest_token(prompt_messages=[], model_config=model_config)
rest_tokens = _calculate_rest_token(
prompt_messages=[],
model_instance=model_instance,
model_schema=model_schema,
model_parameters=model_parameters,
)
if not memory_config.role_prefix:
raise MemoryRolePrefixRequiredError("Memory role prefix is required for completion model.")
memory_text = memory.get_history_prompt_text(
+21
View File
@@ -0,0 +1,21 @@
from __future__ import annotations
from typing import Any, Protocol
from core.model_manager import ModelInstance
class CredentialsProvider(Protocol):
"""Port for loading runtime credentials for a provider/model pair."""
def fetch(self, provider_name: str, model_name: str) -> dict[str, Any]:
"""Return credentials for the target provider/model or raise a domain error."""
...
class ModelFactory(Protocol):
"""Port for creating initialized LLM model instances for execution."""
def init_model_instance(self, provider_name: str, model_name: str) -> ModelInstance:
"""Create a model instance that is ready for schema lookup and invocation."""
...
@@ -3,10 +3,9 @@ import json
import logging
import uuid
from collections.abc import Mapping, Sequence
from typing import Any, cast
from typing import TYPE_CHECKING, Any, cast
from core.app.entities.app_invoke_entities import ModelConfigWithCredentialsEntity
from core.memory.token_buffer_memory import TokenBufferMemory
from core.model_manager import ModelInstance
from core.model_runtime.entities import ImagePromptMessageContent
from core.model_runtime.entities.llm_entities import LLMUsage
@@ -20,6 +19,7 @@ from core.model_runtime.entities.message_entities import (
)
from core.model_runtime.entities.model_entities import ModelFeature, ModelPropertyKey
from core.model_runtime.model_providers.__base.large_language_model import LargeLanguageModel
from core.model_runtime.token_buffer_memory import TokenBufferMemory
from core.model_runtime.utils.encoders import jsonable_encoder
from core.prompt.advanced_prompt_transform import AdvancedPromptTransform
from core.prompt.entities.advanced_prompt_entities import ChatModelMessage, CompletionModelPromptTemplate
@@ -60,6 +60,11 @@ from .prompts import (
logger = logging.getLogger(__name__)
if TYPE_CHECKING:
from core.workflow.entities import GraphInitParams
from core.workflow.nodes.llm.protocols import CredentialsProvider, ModelFactory
from core.workflow.runtime import GraphRuntimeState
def extract_json(text):
"""
@@ -92,6 +97,27 @@ class ParameterExtractorNode(Node[ParameterExtractorNodeData]):
_model_instance: ModelInstance | None = None
_model_config: ModelConfigWithCredentialsEntity | None = None
_credentials_provider: "CredentialsProvider"
_model_factory: "ModelFactory"
def __init__(
self,
id: str,
config: Mapping[str, Any],
graph_init_params: "GraphInitParams",
graph_runtime_state: "GraphRuntimeState",
*,
credentials_provider: "CredentialsProvider",
model_factory: "ModelFactory",
) -> None:
super().__init__(
id=id,
config=config,
graph_init_params=graph_init_params,
graph_runtime_state=graph_runtime_state,
)
self._credentials_provider = credentials_provider
self._model_factory = model_factory
@classmethod
def get_default_config(cls, filters: Mapping[str, object] | None = None) -> Mapping[str, object]:
@@ -806,7 +832,9 @@ class ParameterExtractorNode(Node[ParameterExtractorNodeData]):
"""
if not self._model_instance or not self._model_config:
self._model_instance, self._model_config = llm_utils.fetch_model_config(
tenant_id=self.tenant_id, node_data_model=node_data_model
node_data_model=node_data_model,
credentials_provider=self._credentials_provider,
model_factory=self._model_factory,
)
return self._model_instance, self._model_config
@@ -4,9 +4,9 @@ from collections.abc import Mapping, Sequence
from typing import TYPE_CHECKING, Any
from core.app.entities.app_invoke_entities import ModelConfigWithCredentialsEntity
from core.memory.token_buffer_memory import TokenBufferMemory
from core.model_manager import ModelInstance
from core.model_runtime.entities import LLMUsage, ModelPropertyKey, PromptMessageRole
from core.model_runtime.token_buffer_memory import TokenBufferMemory
from core.model_runtime.utils.encoders import jsonable_encoder
from core.prompt.advanced_prompt_transform import AdvancedPromptTransform
from core.prompt.simple_prompt_transform import ModelMode
@@ -24,6 +24,7 @@ from core.workflow.nodes.base.node import Node
from core.workflow.nodes.base.variable_template_parser import VariableTemplateParser
from core.workflow.nodes.llm import LLMNode, LLMNodeChatModelMessage, LLMNodeCompletionModelPromptTemplate, llm_utils
from core.workflow.nodes.llm.file_saver import FileSaverImpl, LLMFileSaver
from core.workflow.nodes.llm.protocols import CredentialsProvider, ModelFactory
from libs.json_in_md_parser import parse_and_check_json_markdown
from .entities import QuestionClassifierNodeData
@@ -49,6 +50,8 @@ class QuestionClassifierNode(Node[QuestionClassifierNodeData]):
_file_outputs: list["File"]
_llm_file_saver: LLMFileSaver
_credentials_provider: "CredentialsProvider"
_model_factory: "ModelFactory"
def __init__(
self,
@@ -57,6 +60,8 @@ class QuestionClassifierNode(Node[QuestionClassifierNodeData]):
graph_init_params: "GraphInitParams",
graph_runtime_state: "GraphRuntimeState",
*,
credentials_provider: "CredentialsProvider",
model_factory: "ModelFactory",
llm_file_saver: LLMFileSaver | None = None,
):
super().__init__(
@@ -68,6 +73,9 @@ class QuestionClassifierNode(Node[QuestionClassifierNodeData]):
# LLM file outputs, used for MultiModal outputs.
self._file_outputs = []
self._credentials_provider = credentials_provider
self._model_factory = model_factory
if llm_file_saver is None:
llm_file_saver = FileSaverImpl(
user_id=graph_init_params.user_id,
@@ -89,9 +97,16 @@ class QuestionClassifierNode(Node[QuestionClassifierNodeData]):
variables = {"query": query}
# fetch model config
model_instance, model_config = llm_utils.fetch_model_config(
tenant_id=self.tenant_id,
node_data_model=node_data.model,
credentials_provider=self._credentials_provider,
model_factory=self._model_factory,
)
model_schema = model_instance.model_type_instance.get_model_schema(
model_instance.model_name,
model_instance.credentials,
)
if not model_schema:
raise ValueError(f"Model schema not found for {model_instance.model_name}")
# fetch memory
memory = llm_utils.fetch_memory(
variable_pool=variable_pool,
@@ -133,13 +148,15 @@ class QuestionClassifierNode(Node[QuestionClassifierNodeData]):
prompt_template=prompt_template,
sys_query="",
memory=memory,
model_config=model_config,
model_instance=model_instance,
model_schema=model_schema,
model_parameters=node_data.model.completion_params,
stop=model_config.stop,
sys_files=files,
vision_enabled=node_data.vision.enabled,
vision_detail=node_data.vision.configs.detail,
variable_pool=variable_pool,
jinja2_variables=[],
tenant_id=self.tenant_id,
)
result_text = ""
@@ -0,0 +1,50 @@
from collections.abc import Generator
from typing import Any, Protocol
from pydantic import BaseModel
from core.workflow.file import File
from core.workflow.node_events import StreamChunkEvent, StreamCompletedEvent
class DatasourceParameter(BaseModel):
workspace_id: str
page_id: str
type: str
class OnlineDriveDownloadFileParam(BaseModel):
id: str
bucket: str
class DatasourceFinal(BaseModel):
data: dict[str, Any] | None = None
class DatasourceManagerProtocol(Protocol):
@classmethod
def get_icon_url(cls, provider_id: str, tenant_id: str, datasource_name: str, datasource_type: str) -> str: ...
@classmethod
def stream_node_events(
cls,
*,
node_id: str,
user_id: str,
datasource_name: str,
datasource_type: str,
provider_id: str,
tenant_id: str,
provider: str,
plugin_id: str,
credential_id: str,
parameters_for_log: dict[str, Any],
datasource_info: dict[str, Any],
variable_pool: Any,
datasource_param: DatasourceParameter | None = None,
online_drive_request: OnlineDriveDownloadFileParam | None = None,
) -> Generator[StreamChunkEvent | StreamCompletedEvent, None, None]: ...
@classmethod
def get_upload_file_by_id(cls, file_id: str, tenant_id: str) -> File: ...
+9 -9
View File
@@ -1,8 +1,7 @@
import logging
import time
import uuid
from collections.abc import Generator, Mapping, Sequence
from typing import Any
from typing import Any, cast
from configs import dify_config
from core.app.apps.exc import GenerateTaskStoppedError
@@ -11,6 +10,7 @@ from core.app.workflow.layers.observability import ObservabilityLayer
from core.app.workflow.node_factory import DifyNodeFactory
from core.workflow.constants import ENVIRONMENT_VARIABLE_NODE_ID
from core.workflow.entities import GraphInitParams
from core.workflow.entities.graph_config import NodeConfigData, NodeConfigDict
from core.workflow.errors import WorkflowNodeRunFailedError
from core.workflow.file.models import File
from core.workflow.graph import Graph
@@ -168,7 +168,8 @@ class WorkflowEntry:
graph_init_params=graph_init_params,
graph_runtime_state=graph_runtime_state,
)
node = node_factory.create_node(node_config)
typed_node_config = cast(dict[str, object], node_config)
node = cast(Any, node_factory).create_node(typed_node_config)
node_cls = type(node)
try:
@@ -256,7 +257,7 @@ class WorkflowEntry:
@classmethod
def run_free_node(
cls, node_data: dict, node_id: str, tenant_id: str, user_id: str, user_inputs: dict[str, Any]
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]]:
"""
Run free node
@@ -302,16 +303,15 @@ class WorkflowEntry:
graph_runtime_state = GraphRuntimeState(variable_pool=variable_pool, start_at=time.perf_counter())
# init workflow run state
node_config = {
node_config: NodeConfigDict = {
"id": node_id,
"data": node_data,
"data": cast(NodeConfigData, node_data),
}
node: Node = node_cls(
id=str(uuid.uuid4()),
config=node_config,
node_factory = DifyNodeFactory(
graph_init_params=graph_init_params,
graph_runtime_state=graph_runtime_state,
)
node = node_factory.create_node(node_config)
try:
# variable selector to variable mapping
+30 -11
View File
@@ -26,7 +26,26 @@ def init_app(app: DifyApp):
ConsoleSpanExporter,
)
from opentelemetry.sdk.trace.sampling import ParentBasedTraceIdRatio
from opentelemetry.semconv.resource import ResourceAttributes
from opentelemetry.semconv._incubating.attributes.deployment_attributes import ( # type: ignore[import-untyped]
DEPLOYMENT_ENVIRONMENT_NAME,
)
from opentelemetry.semconv._incubating.attributes.host_attributes import ( # type: ignore[import-untyped]
HOST_ARCH,
HOST_ID,
HOST_NAME,
)
from opentelemetry.semconv._incubating.attributes.os_attributes import ( # type: ignore[import-untyped]
OS_DESCRIPTION,
OS_TYPE,
OS_VERSION,
)
from opentelemetry.semconv._incubating.attributes.process_attributes import ( # type: ignore[import-untyped]
PROCESS_PID,
)
from opentelemetry.semconv.attributes.service_attributes import ( # type: ignore[import-untyped]
SERVICE_NAME,
SERVICE_VERSION,
)
from opentelemetry.trace import set_tracer_provider
from extensions.otel.instrumentation import init_instruments
@@ -37,17 +56,17 @@ def init_app(app: DifyApp):
# Follow Semantic Convertions 1.32.0 to define resource attributes
resource = Resource(
attributes={
ResourceAttributes.SERVICE_NAME: dify_config.APPLICATION_NAME,
ResourceAttributes.SERVICE_VERSION: f"dify-{dify_config.project.version}-{dify_config.COMMIT_SHA}",
ResourceAttributes.PROCESS_PID: os.getpid(),
ResourceAttributes.DEPLOYMENT_ENVIRONMENT: f"{dify_config.DEPLOY_ENV}-{dify_config.EDITION}",
ResourceAttributes.HOST_NAME: socket.gethostname(),
ResourceAttributes.HOST_ARCH: platform.machine(),
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}",
HOST_NAME: socket.gethostname(),
HOST_ARCH: platform.machine(),
"custom.deployment.git_commit": dify_config.COMMIT_SHA,
ResourceAttributes.HOST_ID: platform.node(),
ResourceAttributes.OS_TYPE: platform.system().lower(),
ResourceAttributes.OS_DESCRIPTION: platform.platform(),
ResourceAttributes.OS_VERSION: platform.version(),
HOST_ID: platform.node(),
OS_TYPE: platform.system().lower(),
OS_DESCRIPTION: platform.platform(),
OS_VERSION: platform.version(),
}
)
sampler = ParentBasedTraceIdRatio(dify_config.OTEL_SAMPLING_RATE)
+1
View File
@@ -111,6 +111,7 @@ class RedisClientWrapper:
def zcard(self, name: str | bytes) -> Any: ...
def getdel(self, name: str | bytes) -> Any: ...
def pubsub(self) -> PubSub: ...
def pipeline(self, transaction: bool = True, shard_hint: str | None = None) -> Any: ...
def __getattr__(self, item: str) -> Any:
if self._client is None:
+2 -2
View File
@@ -110,10 +110,10 @@ class Storage:
def load_stream(self, filename: str) -> Generator:
return self.storage_runner.load_stream(filename)
def download(self, filename: str, target_filepath):
def download(self, filename, target_filepath):
self.storage_runner.download(filename, target_filepath)
def exists(self, filename: str):
def exists(self, filename):
return self.storage_runner.exists(filename)
def delete(self, filename: str):
+6 -3
View File
@@ -7,7 +7,10 @@ from opentelemetry.instrumentation.httpx import HTTPXClientInstrumentor
from opentelemetry.instrumentation.redis import RedisInstrumentor
from opentelemetry.instrumentation.sqlalchemy import SQLAlchemyInstrumentor
from opentelemetry.metrics import get_meter, get_meter_provider
from opentelemetry.semconv.trace import SpanAttributes
from opentelemetry.semconv.attributes.http_attributes import ( # type: ignore[import-untyped]
HTTP_REQUEST_METHOD,
HTTP_ROUTE,
)
from opentelemetry.trace import Span, get_tracer_provider
from opentelemetry.trace.status import StatusCode
@@ -85,9 +88,9 @@ def init_flask_instrumentor(app: DifyApp) -> None:
attributes: dict[str, str | int] = {"status_code": status_code, "status_class": status_class}
request = flask.request
if request and request.url_rule:
attributes[SpanAttributes.HTTP_TARGET] = str(request.url_rule.rule)
attributes[HTTP_ROUTE] = str(request.url_rule.rule)
if request and request.method:
attributes[SpanAttributes.HTTP_METHOD] = str(request.method)
attributes[HTTP_REQUEST_METHOD] = str(request.method)
_http_response_counter.add(1, attributes)
except Exception:
logger.exception("Error setting status and attributes")
+2 -2
View File
@@ -20,11 +20,11 @@ class BaseStorage(ABC):
raise NotImplementedError
@abstractmethod
def download(self, filename, target_filepath):
def download(self, filename: str, target_filepath: str) -> None:
raise NotImplementedError
@abstractmethod
def exists(self, filename):
def exists(self, filename: str) -> bool:
raise NotImplementedError
@abstractmethod
+2
View File
@@ -1,7 +1,9 @@
project-includes = ["."]
project-excludes = [
"tests/",
".venv",
"migrations/",
"core/rag",
]
python-platform = "linux"
python-version = "3.11.0"
+4 -4
View File
@@ -107,19 +107,19 @@ class AppService:
if model_instance:
if (
model_instance.model == default_model_config["model"]["name"]
model_instance.model_name == default_model_config["model"]["name"]
and model_instance.provider == default_model_config["model"]["provider"]
):
default_model_dict = default_model_config["model"]
else:
llm_model = cast(LargeLanguageModel, model_instance.model_type_instance)
model_schema = llm_model.get_model_schema(model_instance.model, model_instance.credentials)
model_schema = llm_model.get_model_schema(model_instance.model_name, model_instance.credentials)
if model_schema is None:
raise ValueError(f"model schema not found for model {model_instance.model}")
raise ValueError(f"model schema not found for model {model_instance.model_name}")
default_model_dict = {
"provider": model_instance.provider,
"name": model_instance.model,
"name": model_instance.model_name,
"mode": model_schema.model_properties.get(ModelPropertyKey.MODE),
"completion_params": {},
}
+2 -1
View File
@@ -8,6 +8,7 @@ new GraphEngine command channel mechanism.
from core.app.apps.base_app_queue_manager import AppQueueManager
from core.app.entities.app_invoke_entities import InvokeFrom
from core.workflow.graph_engine.manager import GraphEngineManager
from extensions.ext_redis import redis_client
from models.model import AppMode
@@ -42,4 +43,4 @@ class AppTaskService:
# New mechanism: Send stop command via GraphEngine for workflow-based apps
# This ensures proper workflow status recording in the persistence layer
if app_mode in (AppMode.ADVANCED_CHAT, AppMode.WORKFLOW):
GraphEngineManager.send_stop_command(task_id)
GraphEngineManager(redis_client).send_stop_command(task_id)
+23 -13
View File
@@ -252,7 +252,7 @@ class DatasetService:
dataset.updated_by = account.id
dataset.tenant_id = tenant_id
dataset.embedding_model_provider = embedding_model.provider if embedding_model else None
dataset.embedding_model = embedding_model.model if embedding_model else None
dataset.embedding_model = embedding_model.model_name if embedding_model else None
dataset.retrieval_model = retrieval_model.model_dump() if retrieval_model else None
dataset.permission = permission or DatasetPermissionEnum.ONLY_ME
dataset.provider = provider
@@ -384,7 +384,7 @@ class DatasetService:
model=model,
)
text_embedding_model = cast(TextEmbeddingModel, model_instance.model_type_instance)
model_schema = text_embedding_model.get_model_schema(model_instance.model, model_instance.credentials)
model_schema = text_embedding_model.get_model_schema(model_instance.model_name, model_instance.credentials)
if not model_schema:
raise ValueError("Model schema not found")
if model_schema.features and ModelFeature.VISION in model_schema.features:
@@ -743,10 +743,12 @@ class DatasetService:
model_type=ModelType.TEXT_EMBEDDING,
model=data["embedding_model"],
)
filtered_data["embedding_model"] = embedding_model.model
embedding_model_name = embedding_model.model_name
filtered_data["embedding_model"] = embedding_model_name
filtered_data["embedding_model_provider"] = embedding_model.provider
dataset_collection_binding = DatasetCollectionBindingService.get_dataset_collection_binding(
embedding_model.provider, embedding_model.model
embedding_model.provider,
embedding_model_name,
)
filtered_data["collection_binding_id"] = dataset_collection_binding.id
except LLMBadRequestError:
@@ -876,10 +878,12 @@ class DatasetService:
return
# Apply new embedding model settings
filtered_data["embedding_model"] = embedding_model.model
embedding_model_name = embedding_model.model_name
filtered_data["embedding_model"] = embedding_model_name
filtered_data["embedding_model_provider"] = embedding_model.provider
dataset_collection_binding = DatasetCollectionBindingService.get_dataset_collection_binding(
embedding_model.provider, embedding_model.model
embedding_model.provider,
embedding_model_name,
)
filtered_data["collection_binding_id"] = dataset_collection_binding.id
@@ -955,10 +959,12 @@ class DatasetService:
knowledge_configuration.embedding_model,
)
dataset.is_multimodal = is_multimodal
dataset.embedding_model = embedding_model.model
embedding_model_name = embedding_model.model_name
dataset.embedding_model = embedding_model_name
dataset.embedding_model_provider = embedding_model.provider
dataset_collection_binding = DatasetCollectionBindingService.get_dataset_collection_binding(
embedding_model.provider, embedding_model.model
embedding_model.provider,
embedding_model_name,
)
dataset.collection_binding_id = dataset_collection_binding.id
elif knowledge_configuration.indexing_technique == "economy":
@@ -989,10 +995,12 @@ class DatasetService:
model_type=ModelType.TEXT_EMBEDDING,
model=knowledge_configuration.embedding_model,
)
dataset.embedding_model = embedding_model.model
embedding_model_name = embedding_model.model_name
dataset.embedding_model = embedding_model_name
dataset.embedding_model_provider = embedding_model.provider
dataset_collection_binding = DatasetCollectionBindingService.get_dataset_collection_binding(
embedding_model.provider, embedding_model.model
embedding_model.provider,
embedding_model_name,
)
is_multimodal = DatasetService.check_is_multimodal_model(
current_user.current_tenant_id,
@@ -1049,11 +1057,13 @@ class DatasetService:
skip_embedding_update = True
if not skip_embedding_update:
if embedding_model:
dataset.embedding_model = embedding_model.model
embedding_model_name = embedding_model.model_name
dataset.embedding_model = embedding_model_name
dataset.embedding_model_provider = embedding_model.provider
dataset_collection_binding = (
DatasetCollectionBindingService.get_dataset_collection_binding(
embedding_model.provider, embedding_model.model
embedding_model.provider,
embedding_model_name,
)
)
dataset.collection_binding_id = dataset_collection_binding.id
@@ -1884,7 +1894,7 @@ class DocumentService:
embedding_model = model_manager.get_default_model_instance(
tenant_id=current_user.current_tenant_id, model_type=ModelType.TEXT_EMBEDDING
)
dataset_embedding_model = embedding_model.model
dataset_embedding_model = embedding_model.model_name
dataset_embedding_model_provider = embedding_model.provider
dataset.embedding_model = dataset_embedding_model
dataset.embedding_model_provider = dataset_embedding_model_provider
+1 -1
View File
@@ -7,9 +7,9 @@ from sqlalchemy.orm import sessionmaker
from core.app.apps.advanced_chat.app_config_manager import AdvancedChatAppConfigManager
from core.app.entities.app_invoke_entities import InvokeFrom
from core.llm_generator.llm_generator import LLMGenerator
from core.memory.token_buffer_memory import TokenBufferMemory
from core.model_manager import ModelManager
from core.model_runtime.entities.model_entities import ModelType
from core.model_runtime.token_buffer_memory import TokenBufferMemory
from core.ops.entities.trace_entity import TraceTaskName
from core.ops.ops_trace_manager import TraceQueueManager, TraceTask
from core.ops.utils import measure_time
@@ -0,0 +1,42 @@
from collections.abc import Generator
from core.datasource.datasource_manager import DatasourceManager
from core.datasource.entities.datasource_entities import DatasourceMessage
from core.workflow.node_events import StreamCompletedEvent
def _gen_var_stream() -> Generator[DatasourceMessage, None, None]:
# produce a streamed variable "a"="xy"
yield DatasourceMessage(
type=DatasourceMessage.MessageType.VARIABLE,
message=DatasourceMessage.VariableMessage(variable_name="a", variable_value="x", stream=True),
meta=None,
)
yield DatasourceMessage(
type=DatasourceMessage.MessageType.VARIABLE,
message=DatasourceMessage.VariableMessage(variable_name="a", variable_value="y", stream=True),
meta=None,
)
def test_stream_node_events_accumulates_variables(mocker):
mocker.patch.object(DatasourceManager, "stream_online_results", return_value=_gen_var_stream())
events = list(
DatasourceManager.stream_node_events(
node_id="A",
user_id="u",
datasource_name="ds",
datasource_type="online_document",
provider_id="p/x",
tenant_id="t",
provider="prov",
plugin_id="plug",
credential_id="",
parameters_for_log={},
datasource_info={"user_id": "u"},
variable_pool=mocker.Mock(),
datasource_param=type("P", (), {"workspace_id": "w", "page_id": "pg", "type": "t"})(),
online_drive_request=None,
)
)
assert isinstance(events[-1], StreamCompletedEvent)
@@ -0,0 +1,84 @@
from core.workflow.entities.workflow_node_execution import WorkflowNodeExecutionStatus
from core.workflow.node_events import NodeRunResult, StreamCompletedEvent
from core.workflow.nodes.datasource.datasource_node import DatasourceNode
class _Seg:
def __init__(self, v):
self.value = v
class _VarPool:
def __init__(self, data):
self.data = data
def get(self, path):
d = self.data
for k in path:
d = d[k]
return _Seg(d)
def add(self, *_a, **_k):
pass
class _GS:
def __init__(self, vp):
self.variable_pool = vp
class _GP:
tenant_id = "t1"
app_id = "app-1"
workflow_id = "wf-1"
graph_config = {}
user_id = "u1"
user_from = "account"
invoke_from = "debugger"
call_depth = 0
def test_node_integration_minimal_stream(mocker):
sys_d = {
"sys": {
"datasource_type": "online_document",
"datasource_info": {"workspace_id": "w", "page": {"page_id": "pg", "type": "t"}, "credential_id": ""},
}
}
vp = _VarPool(sys_d)
class _Mgr:
@classmethod
def get_icon_url(cls, **_):
return "icon"
@classmethod
def stream_node_events(cls, **_):
yield from ()
yield StreamCompletedEvent(node_run_result=NodeRunResult(status=WorkflowNodeExecutionStatus.SUCCEEDED))
@classmethod
def get_upload_file_by_id(cls, **_):
raise AssertionError
node = DatasourceNode(
id="n",
config={
"id": "n",
"data": {
"type": "datasource",
"version": "1",
"title": "Datasource",
"provider_type": "plugin",
"provider_name": "p",
"plugin_id": "plug",
"datasource_name": "ds",
},
},
graph_init_params=_GP(),
graph_runtime_state=_GS(vp),
datasource_manager=_Mgr,
)
out = list(node._run())
assert isinstance(out[-1], StreamCompletedEvent)
@@ -68,6 +68,7 @@ def init_code_node(code_config: dict):
config=code_config,
graph_init_params=init_params,
graph_runtime_state=graph_runtime_state,
code_executor=node_factory._code_executor,
code_limits=CodeNodeLimits(
max_string_length=dify_config.CODE_MAX_STRING_LENGTH,
max_number=dify_config.CODE_MAX_NUMBER,
@@ -80,6 +80,8 @@ def init_llm_node(config: dict) -> LLMNode:
config=config,
graph_init_params=init_params,
graph_runtime_state=graph_runtime_state,
credentials_provider=MagicMock(),
model_factory=MagicMock(),
)
return node
@@ -115,7 +117,7 @@ def test_execute_llm():
db.session.close = MagicMock()
# Mock the _fetch_model_config to avoid database calls
def mock_fetch_model_config(**_kwargs):
def mock_fetch_model_config(*_args, **_kwargs):
from decimal import Decimal
from unittest.mock import MagicMock
@@ -227,7 +229,7 @@ def test_execute_llm_with_jinja2():
db.session.close = MagicMock()
# Mock the _fetch_model_config method
def mock_fetch_model_config(**_kwargs):
def mock_fetch_model_config(*_args, **_kwargs):
from decimal import Decimal
from unittest.mock import MagicMock
@@ -9,6 +9,7 @@ from core.model_runtime.entities import AssistantPromptMessage
from core.workflow.entities import GraphInitParams
from core.workflow.enums import WorkflowNodeExecutionStatus
from core.workflow.graph import Graph
from core.workflow.nodes.llm.protocols import CredentialsProvider, ModelFactory
from core.workflow.nodes.parameter_extractor.parameter_extractor_node import ParameterExtractorNode
from core.workflow.runtime import GraphRuntimeState, VariablePool
from core.workflow.system_variable import SystemVariable
@@ -84,6 +85,8 @@ def init_parameter_extractor_node(config: dict):
config=config,
graph_init_params=init_params,
graph_runtime_state=graph_runtime_state,
credentials_provider=MagicMock(spec=CredentialsProvider),
model_factory=MagicMock(spec=ModelFactory),
)
return node
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,418 @@
"""Integration tests for SQL-oriented DatasetService scenarios.
This suite migrates SQL-backed behaviors from the old unit suite to real
container-backed integration tests. The tests exercise real ORM persistence and
only patch non-DB collaborators when needed.
"""
from unittest.mock import Mock, patch
from uuid import uuid4
import pytest
from core.model_runtime.entities.model_entities import ModelType
from core.rag.retrieval.retrieval_methods import RetrievalMethod
from extensions.ext_database import db
from models.account import Account, Tenant, TenantAccountJoin, TenantAccountRole
from models.dataset import Dataset, DatasetPermissionEnum, Document, ExternalKnowledgeBindings
from services.dataset_service import DatasetService
from services.entities.knowledge_entities.knowledge_entities import RerankingModel, RetrievalModel
from services.errors.dataset import DatasetNameDuplicateError
class DatasetServiceIntegrationDataFactory:
"""Factory for creating real database entities used by integration tests."""
@staticmethod
def create_account_with_tenant(role: TenantAccountRole = TenantAccountRole.OWNER) -> tuple[Account, Tenant]:
"""Create an account and tenant, then bind the account as current tenant member."""
account = Account(
email=f"{uuid4()}@example.com",
name=f"user-{uuid4()}",
interface_language="en-US",
status="active",
)
tenant = Tenant(name=f"tenant-{uuid4()}", status="normal")
db.session.add_all([account, tenant])
db.session.flush()
join = TenantAccountJoin(
tenant_id=tenant.id,
account_id=account.id,
role=role,
current=True,
)
db.session.add(join)
db.session.flush()
# Keep tenant context on the in-memory user without opening a separate session.
account.role = role
account._current_tenant = tenant
return account, tenant
@staticmethod
def create_dataset(
tenant_id: str,
created_by: str,
name: str = "Test Dataset",
description: str | None = "Test description",
provider: str = "vendor",
indexing_technique: str | None = "high_quality",
permission: str = DatasetPermissionEnum.ONLY_ME,
retrieval_model: dict | None = None,
embedding_model_provider: str | None = None,
embedding_model: str | None = None,
collection_binding_id: str | None = None,
chunk_structure: str | None = None,
) -> Dataset:
"""Create a dataset record with configurable SQL fields."""
dataset = Dataset(
tenant_id=tenant_id,
name=name,
description=description,
data_source_type="upload_file",
indexing_technique=indexing_technique,
created_by=created_by,
provider=provider,
permission=permission,
retrieval_model=retrieval_model,
embedding_model_provider=embedding_model_provider,
embedding_model=embedding_model,
collection_binding_id=collection_binding_id,
chunk_structure=chunk_structure,
)
db.session.add(dataset)
db.session.flush()
return dataset
@staticmethod
def create_document(dataset: Dataset, created_by: str, name: str = "doc.txt") -> Document:
"""Create a document row belonging to the given dataset."""
document = Document(
tenant_id=dataset.tenant_id,
dataset_id=dataset.id,
position=1,
data_source_type="upload_file",
data_source_info='{"upload_file_id": "upload-file-id"}',
batch=str(uuid4()),
name=name,
created_from="web",
created_by=created_by,
indexing_status="completed",
doc_form="text_model",
)
db.session.add(document)
db.session.flush()
return document
@staticmethod
def create_embedding_model(provider: str = "openai", model_name: str = "text-embedding-ada-002") -> Mock:
"""Create a fake embedding model object for external provider boundary patching."""
embedding_model = Mock()
embedding_model.provider = provider
embedding_model.model_name = model_name
return embedding_model
class TestDatasetServiceCreateDataset:
"""Integration coverage for DatasetService.create_empty_dataset."""
def test_create_internal_dataset_basic_success(self, db_session_with_containers):
"""Create a basic internal dataset with minimal configuration."""
# Arrange
account, tenant = DatasetServiceIntegrationDataFactory.create_account_with_tenant()
# Act
result = DatasetService.create_empty_dataset(
tenant_id=tenant.id,
name="Basic Internal Dataset",
description="Test description",
indexing_technique=None,
account=account,
)
# Assert
created_dataset = db.session.get(Dataset, result.id)
assert created_dataset is not None
assert created_dataset.provider == "vendor"
assert created_dataset.permission == DatasetPermissionEnum.ONLY_ME
assert created_dataset.embedding_model_provider is None
assert created_dataset.embedding_model is None
def test_create_internal_dataset_with_economy_indexing(self, db_session_with_containers):
"""Create an internal dataset with economy indexing and no embedding model."""
# Arrange
account, tenant = DatasetServiceIntegrationDataFactory.create_account_with_tenant()
# Act
result = DatasetService.create_empty_dataset(
tenant_id=tenant.id,
name="Economy Dataset",
description=None,
indexing_technique="economy",
account=account,
)
# Assert
db.session.refresh(result)
assert result.indexing_technique == "economy"
assert result.embedding_model_provider is None
assert result.embedding_model is None
def test_create_internal_dataset_with_high_quality_indexing(self, db_session_with_containers):
"""Create a high-quality dataset and persist embedding model settings."""
# Arrange
account, tenant = DatasetServiceIntegrationDataFactory.create_account_with_tenant()
embedding_model = DatasetServiceIntegrationDataFactory.create_embedding_model()
# Act
with patch("services.dataset_service.ModelManager") as mock_model_manager:
mock_model_manager.return_value.get_default_model_instance.return_value = embedding_model
result = DatasetService.create_empty_dataset(
tenant_id=tenant.id,
name="High Quality Dataset",
description=None,
indexing_technique="high_quality",
account=account,
)
# Assert
db.session.refresh(result)
assert result.indexing_technique == "high_quality"
assert result.embedding_model_provider == embedding_model.provider
assert result.embedding_model == embedding_model.model_name
mock_model_manager.return_value.get_default_model_instance.assert_called_once_with(
tenant_id=tenant.id,
model_type=ModelType.TEXT_EMBEDDING,
)
def test_create_dataset_duplicate_name_error(self, db_session_with_containers):
"""Raise duplicate-name error when the same tenant already has the name."""
# Arrange
account, tenant = DatasetServiceIntegrationDataFactory.create_account_with_tenant()
DatasetServiceIntegrationDataFactory.create_dataset(
tenant_id=tenant.id,
created_by=account.id,
name="Duplicate Dataset",
indexing_technique=None,
)
# Act / Assert
with pytest.raises(DatasetNameDuplicateError):
DatasetService.create_empty_dataset(
tenant_id=tenant.id,
name="Duplicate Dataset",
description=None,
indexing_technique=None,
account=account,
)
def test_create_external_dataset_success(self, db_session_with_containers):
"""Create an external dataset and persist external knowledge binding."""
# Arrange
account, tenant = DatasetServiceIntegrationDataFactory.create_account_with_tenant()
external_knowledge_api_id = str(uuid4())
external_knowledge_id = "knowledge-123"
# Act
with patch("services.dataset_service.ExternalDatasetService.get_external_knowledge_api") as mock_get_api:
mock_get_api.return_value = Mock(id=external_knowledge_api_id)
result = DatasetService.create_empty_dataset(
tenant_id=tenant.id,
name="External Dataset",
description=None,
indexing_technique=None,
account=account,
provider="external",
external_knowledge_api_id=external_knowledge_api_id,
external_knowledge_id=external_knowledge_id,
)
# Assert
binding = db.session.query(ExternalKnowledgeBindings).filter_by(dataset_id=result.id).first()
assert result.provider == "external"
assert binding is not None
assert binding.external_knowledge_id == external_knowledge_id
assert binding.external_knowledge_api_id == external_knowledge_api_id
def test_create_dataset_with_retrieval_model_and_reranking(self, db_session_with_containers):
"""Create a high-quality dataset with retrieval/reranking settings."""
# Arrange
account, tenant = DatasetServiceIntegrationDataFactory.create_account_with_tenant()
embedding_model = DatasetServiceIntegrationDataFactory.create_embedding_model()
retrieval_model = RetrievalModel(
search_method=RetrievalMethod.SEMANTIC_SEARCH,
reranking_enable=True,
reranking_model=RerankingModel(
reranking_provider_name="cohere",
reranking_model_name="rerank-english-v2.0",
),
top_k=3,
score_threshold_enabled=True,
score_threshold=0.6,
)
# Act
with (
patch("services.dataset_service.ModelManager") as mock_model_manager,
patch("services.dataset_service.DatasetService.check_reranking_model_setting") as mock_check_reranking,
):
mock_model_manager.return_value.get_default_model_instance.return_value = embedding_model
result = DatasetService.create_empty_dataset(
tenant_id=tenant.id,
name="Dataset With Reranking",
description=None,
indexing_technique="high_quality",
account=account,
retrieval_model=retrieval_model,
)
# Assert
db.session.refresh(result)
assert result.retrieval_model == retrieval_model.model_dump()
mock_check_reranking.assert_called_once_with(tenant.id, "cohere", "rerank-english-v2.0")
class TestDatasetServiceUpdateAndDeleteDataset:
"""Integration coverage for SQL-backed update and delete behavior."""
def test_update_dataset_duplicate_name_error(self, db_session_with_containers):
"""Reject update when target name already exists within the same tenant."""
# Arrange
account, tenant = DatasetServiceIntegrationDataFactory.create_account_with_tenant()
source_dataset = DatasetServiceIntegrationDataFactory.create_dataset(
tenant_id=tenant.id,
created_by=account.id,
name="Source Dataset",
)
DatasetServiceIntegrationDataFactory.create_dataset(
tenant_id=tenant.id,
created_by=account.id,
name="Existing Dataset",
)
# Act / Assert
with pytest.raises(ValueError, match="Dataset name already exists"):
DatasetService.update_dataset(source_dataset.id, {"name": "Existing Dataset"}, account)
def test_delete_dataset_with_documents_success(self, db_session_with_containers):
"""Delete a dataset that already has documents."""
# Arrange
account, tenant = DatasetServiceIntegrationDataFactory.create_account_with_tenant()
dataset = DatasetServiceIntegrationDataFactory.create_dataset(
tenant_id=tenant.id,
created_by=account.id,
indexing_technique="high_quality",
chunk_structure="text_model",
)
DatasetServiceIntegrationDataFactory.create_document(dataset=dataset, created_by=account.id)
# Act
with patch("services.dataset_service.dataset_was_deleted") as dataset_deleted_signal:
result = DatasetService.delete_dataset(dataset.id, account)
# Assert
assert result is True
assert db.session.get(Dataset, dataset.id) is None
dataset_deleted_signal.send.assert_called_once_with(dataset)
def test_delete_empty_dataset_success(self, db_session_with_containers):
"""Delete a dataset that has no documents and no indexing technique."""
# Arrange
account, tenant = DatasetServiceIntegrationDataFactory.create_account_with_tenant()
dataset = DatasetServiceIntegrationDataFactory.create_dataset(
tenant_id=tenant.id,
created_by=account.id,
indexing_technique=None,
chunk_structure=None,
)
# Act
with patch("services.dataset_service.dataset_was_deleted") as dataset_deleted_signal:
result = DatasetService.delete_dataset(dataset.id, account)
# Assert
assert result is True
assert db.session.get(Dataset, dataset.id) is None
dataset_deleted_signal.send.assert_called_once_with(dataset)
def test_delete_dataset_with_partial_none_values(self, db_session_with_containers):
"""Delete dataset when indexing_technique is None but doc_form path still exists."""
# Arrange
account, tenant = DatasetServiceIntegrationDataFactory.create_account_with_tenant()
dataset = DatasetServiceIntegrationDataFactory.create_dataset(
tenant_id=tenant.id,
created_by=account.id,
indexing_technique=None,
chunk_structure="text_model",
)
# Act
with patch("services.dataset_service.dataset_was_deleted") as dataset_deleted_signal:
result = DatasetService.delete_dataset(dataset.id, account)
# Assert
assert result is True
assert db.session.get(Dataset, dataset.id) is None
dataset_deleted_signal.send.assert_called_once_with(dataset)
class TestDatasetServiceRetrievalConfiguration:
"""Integration coverage for retrieval configuration persistence."""
def test_get_dataset_retrieval_configuration(self, db_session_with_containers):
"""Return retrieval configuration that is persisted in SQL."""
# Arrange
account, tenant = DatasetServiceIntegrationDataFactory.create_account_with_tenant()
retrieval_model = {
"search_method": "semantic_search",
"top_k": 5,
"score_threshold": 0.5,
"reranking_enable": True,
}
dataset = DatasetServiceIntegrationDataFactory.create_dataset(
tenant_id=tenant.id,
created_by=account.id,
retrieval_model=retrieval_model,
)
# Act
result = DatasetService.get_dataset(dataset.id)
# Assert
assert result is not None
assert result.retrieval_model == retrieval_model
assert result.retrieval_model["search_method"] == "semantic_search"
assert result.retrieval_model["top_k"] == 5
def test_update_dataset_retrieval_configuration(self, db_session_with_containers):
"""Persist retrieval configuration updates through DatasetService.update_dataset."""
# Arrange
account, tenant = DatasetServiceIntegrationDataFactory.create_account_with_tenant()
dataset = DatasetServiceIntegrationDataFactory.create_dataset(
tenant_id=tenant.id,
created_by=account.id,
indexing_technique="high_quality",
retrieval_model={"search_method": "semantic_search", "top_k": 2, "score_threshold": 0.0},
embedding_model_provider="openai",
embedding_model="text-embedding-ada-002",
collection_binding_id=str(uuid4()),
)
update_data = {
"indexing_technique": "high_quality",
"retrieval_model": {
"search_method": "full_text_search",
"top_k": 10,
"score_threshold": 0.7,
},
}
# Act
result = DatasetService.update_dataset(dataset.id, update_data, account)
# Assert
db.session.refresh(dataset)
assert result.id == dataset.id
assert dataset.retrieval_model == update_data["retrieval_model"]
@@ -0,0 +1,529 @@
from unittest.mock import Mock, patch
from uuid import uuid4
import pytest
from core.model_runtime.entities.model_entities import ModelType
from extensions.ext_database import db
from models.account import Account, Tenant, TenantAccountJoin, TenantAccountRole
from models.dataset import Dataset, ExternalKnowledgeBindings
from services.dataset_service import DatasetService
from services.errors.account import NoPermissionError
class DatasetUpdateTestDataFactory:
"""Factory class for creating real test data for dataset update integration tests."""
@staticmethod
def create_account_with_tenant(role: TenantAccountRole = TenantAccountRole.OWNER) -> tuple[Account, Tenant]:
"""Create a real account and tenant with the given role."""
account = Account(
email=f"{uuid4()}@example.com",
name=f"user-{uuid4()}",
interface_language="en-US",
status="active",
)
db.session.add(account)
db.session.commit()
tenant = Tenant(name=f"tenant-{account.id}", status="normal")
db.session.add(tenant)
db.session.commit()
join = TenantAccountJoin(
tenant_id=tenant.id,
account_id=account.id,
role=role,
current=True,
)
db.session.add(join)
db.session.commit()
account.current_tenant = tenant
return account, tenant
@staticmethod
def create_dataset(
tenant_id: str,
created_by: str,
provider: str = "vendor",
name: str = "old_name",
description: str = "old_description",
indexing_technique: str = "high_quality",
retrieval_model: str = "old_model",
permission: str = "only_me",
embedding_model_provider: str | None = None,
embedding_model: str | None = None,
collection_binding_id: str | None = None,
) -> Dataset:
"""Create a real dataset."""
dataset = Dataset(
tenant_id=tenant_id,
name=name,
description=description,
data_source_type="upload_file",
indexing_technique=indexing_technique,
created_by=created_by,
provider=provider,
retrieval_model=retrieval_model,
permission=permission,
embedding_model_provider=embedding_model_provider,
embedding_model=embedding_model,
collection_binding_id=collection_binding_id,
)
db.session.add(dataset)
db.session.commit()
return dataset
@staticmethod
def create_external_binding(
tenant_id: str,
dataset_id: str,
created_by: str,
external_knowledge_id: str = "old_knowledge_id",
external_knowledge_api_id: str | None = None,
) -> ExternalKnowledgeBindings:
"""Create a real external knowledge binding."""
if external_knowledge_api_id is None:
external_knowledge_api_id = str(uuid4())
binding = ExternalKnowledgeBindings(
tenant_id=tenant_id,
dataset_id=dataset_id,
created_by=created_by,
external_knowledge_id=external_knowledge_id,
external_knowledge_api_id=external_knowledge_api_id,
)
db.session.add(binding)
db.session.commit()
return binding
class TestDatasetServiceUpdateDataset:
"""
Comprehensive integration tests for DatasetService.update_dataset method.
This test suite covers all supported scenarios including:
- External dataset updates
- Internal dataset updates with different indexing techniques
- Embedding model updates
- Permission checks
- Error conditions and edge cases
"""
# ==================== External Dataset Tests ====================
def test_update_external_dataset_success(self, db_session_with_containers):
"""Test successful update of external dataset."""
user, tenant = DatasetUpdateTestDataFactory.create_account_with_tenant()
dataset = DatasetUpdateTestDataFactory.create_dataset(
tenant_id=tenant.id,
created_by=user.id,
provider="external",
name="old_name",
description="old_description",
retrieval_model="old_model",
)
binding = DatasetUpdateTestDataFactory.create_external_binding(
tenant_id=tenant.id,
dataset_id=dataset.id,
created_by=user.id,
)
binding_id = binding.id
db.session.expunge(binding)
update_data = {
"name": "new_name",
"description": "new_description",
"external_retrieval_model": "new_model",
"permission": "only_me",
"external_knowledge_id": "new_knowledge_id",
"external_knowledge_api_id": str(uuid4()),
}
result = DatasetService.update_dataset(dataset.id, update_data, user)
db.session.refresh(dataset)
updated_binding = db.session.query(ExternalKnowledgeBindings).filter_by(id=binding_id).first()
assert dataset.name == "new_name"
assert dataset.description == "new_description"
assert dataset.retrieval_model == "new_model"
assert updated_binding is not None
assert updated_binding.external_knowledge_id == "new_knowledge_id"
assert updated_binding.external_knowledge_api_id == update_data["external_knowledge_api_id"]
assert result.id == dataset.id
def test_update_external_dataset_missing_knowledge_id_error(self, db_session_with_containers):
"""Test error when external knowledge id is missing."""
user, tenant = DatasetUpdateTestDataFactory.create_account_with_tenant()
dataset = DatasetUpdateTestDataFactory.create_dataset(
tenant_id=tenant.id,
created_by=user.id,
provider="external",
)
DatasetUpdateTestDataFactory.create_external_binding(
tenant_id=tenant.id,
dataset_id=dataset.id,
created_by=user.id,
)
update_data = {"name": "new_name", "external_knowledge_api_id": str(uuid4())}
with pytest.raises(ValueError) as context:
DatasetService.update_dataset(dataset.id, update_data, user)
assert "External knowledge id is required" in str(context.value)
db.session.rollback()
def test_update_external_dataset_missing_api_id_error(self, db_session_with_containers):
"""Test error when external knowledge api id is missing."""
user, tenant = DatasetUpdateTestDataFactory.create_account_with_tenant()
dataset = DatasetUpdateTestDataFactory.create_dataset(
tenant_id=tenant.id,
created_by=user.id,
provider="external",
)
DatasetUpdateTestDataFactory.create_external_binding(
tenant_id=tenant.id,
dataset_id=dataset.id,
created_by=user.id,
)
update_data = {"name": "new_name", "external_knowledge_id": "knowledge_id"}
with pytest.raises(ValueError) as context:
DatasetService.update_dataset(dataset.id, update_data, user)
assert "External knowledge api id is required" in str(context.value)
db.session.rollback()
def test_update_external_dataset_binding_not_found_error(self, db_session_with_containers):
"""Test error when external knowledge binding is not found."""
user, tenant = DatasetUpdateTestDataFactory.create_account_with_tenant()
dataset = DatasetUpdateTestDataFactory.create_dataset(
tenant_id=tenant.id,
created_by=user.id,
provider="external",
)
update_data = {
"name": "new_name",
"external_knowledge_id": "knowledge_id",
"external_knowledge_api_id": str(uuid4()),
}
with pytest.raises(ValueError) as context:
DatasetService.update_dataset(dataset.id, update_data, user)
assert "External knowledge binding not found" in str(context.value)
db.session.rollback()
# ==================== Internal Dataset Basic Tests ====================
def test_update_internal_dataset_basic_success(self, db_session_with_containers):
"""Test successful update of internal dataset with basic fields."""
user, tenant = DatasetUpdateTestDataFactory.create_account_with_tenant()
existing_binding_id = str(uuid4())
dataset = DatasetUpdateTestDataFactory.create_dataset(
tenant_id=tenant.id,
created_by=user.id,
provider="vendor",
indexing_technique="high_quality",
embedding_model_provider="openai",
embedding_model="text-embedding-ada-002",
collection_binding_id=existing_binding_id,
)
update_data = {
"name": "new_name",
"description": "new_description",
"indexing_technique": "high_quality",
"retrieval_model": "new_model",
"embedding_model_provider": "openai",
"embedding_model": "text-embedding-ada-002",
}
result = DatasetService.update_dataset(dataset.id, update_data, user)
db.session.refresh(dataset)
assert dataset.name == "new_name"
assert dataset.description == "new_description"
assert dataset.indexing_technique == "high_quality"
assert dataset.retrieval_model == "new_model"
assert dataset.embedding_model_provider == "openai"
assert dataset.embedding_model == "text-embedding-ada-002"
assert result.id == dataset.id
def test_update_internal_dataset_filter_none_values(self, db_session_with_containers):
"""Test that None values are filtered out except for description field."""
user, tenant = DatasetUpdateTestDataFactory.create_account_with_tenant()
existing_binding_id = str(uuid4())
dataset = DatasetUpdateTestDataFactory.create_dataset(
tenant_id=tenant.id,
created_by=user.id,
provider="vendor",
indexing_technique="high_quality",
embedding_model_provider="openai",
embedding_model="text-embedding-ada-002",
collection_binding_id=existing_binding_id,
)
update_data = {
"name": "new_name",
"description": None,
"indexing_technique": "high_quality",
"retrieval_model": "new_model",
"embedding_model_provider": None,
"embedding_model": None,
}
result = DatasetService.update_dataset(dataset.id, update_data, user)
db.session.refresh(dataset)
assert dataset.name == "new_name"
assert dataset.description is None
assert dataset.embedding_model_provider == "openai"
assert dataset.embedding_model == "text-embedding-ada-002"
assert dataset.retrieval_model == "new_model"
assert result.id == dataset.id
# ==================== Indexing Technique Switch Tests ====================
def test_update_internal_dataset_indexing_technique_to_economy(self, db_session_with_containers):
"""Test updating internal dataset indexing technique to economy."""
user, tenant = DatasetUpdateTestDataFactory.create_account_with_tenant()
existing_binding_id = str(uuid4())
dataset = DatasetUpdateTestDataFactory.create_dataset(
tenant_id=tenant.id,
created_by=user.id,
provider="vendor",
indexing_technique="high_quality",
embedding_model_provider="openai",
embedding_model="text-embedding-ada-002",
collection_binding_id=existing_binding_id,
)
update_data = {
"indexing_technique": "economy",
"retrieval_model": "new_model",
}
with patch("services.dataset_service.deal_dataset_vector_index_task") as mock_task:
result = DatasetService.update_dataset(dataset.id, update_data, user)
mock_task.delay.assert_called_once_with(dataset.id, "remove")
db.session.refresh(dataset)
assert dataset.indexing_technique == "economy"
assert dataset.embedding_model is None
assert dataset.embedding_model_provider is None
assert dataset.collection_binding_id is None
assert dataset.retrieval_model == "new_model"
assert result.id == dataset.id
def test_update_internal_dataset_indexing_technique_to_high_quality(self, db_session_with_containers):
"""Test updating internal dataset indexing technique to high_quality."""
user, tenant = DatasetUpdateTestDataFactory.create_account_with_tenant()
dataset = DatasetUpdateTestDataFactory.create_dataset(
tenant_id=tenant.id,
created_by=user.id,
provider="vendor",
indexing_technique="economy",
)
embedding_model = Mock()
embedding_model.model_name = "text-embedding-ada-002"
embedding_model.provider = "openai"
binding = Mock()
binding.id = str(uuid4())
update_data = {
"indexing_technique": "high_quality",
"embedding_model_provider": "openai",
"embedding_model": "text-embedding-ada-002",
"retrieval_model": "new_model",
}
with (
patch("services.dataset_service.current_user", user),
patch("services.dataset_service.ModelManager") as mock_model_manager,
patch(
"services.dataset_service.DatasetCollectionBindingService.get_dataset_collection_binding"
) as mock_get_binding,
patch("services.dataset_service.deal_dataset_vector_index_task") as mock_task,
):
mock_model_manager.return_value.get_model_instance.return_value = embedding_model
mock_get_binding.return_value = binding
result = DatasetService.update_dataset(dataset.id, update_data, user)
mock_model_manager.return_value.get_model_instance.assert_called_once_with(
tenant_id=tenant.id,
provider="openai",
model_type=ModelType.TEXT_EMBEDDING,
model="text-embedding-ada-002",
)
mock_get_binding.assert_called_once_with("openai", "text-embedding-ada-002")
mock_task.delay.assert_called_once_with(dataset.id, "add")
db.session.refresh(dataset)
assert dataset.indexing_technique == "high_quality"
assert dataset.embedding_model == "text-embedding-ada-002"
assert dataset.embedding_model_provider == "openai"
assert dataset.collection_binding_id == binding.id
assert dataset.retrieval_model == "new_model"
assert result.id == dataset.id
# ==================== Embedding Model Update Tests ====================
def test_update_internal_dataset_keep_existing_embedding_model_when_indexing_technique_unchanged(
self, db_session_with_containers
):
"""Test preserving embedding settings when indexing technique remains unchanged."""
user, tenant = DatasetUpdateTestDataFactory.create_account_with_tenant()
existing_binding_id = str(uuid4())
dataset = DatasetUpdateTestDataFactory.create_dataset(
tenant_id=tenant.id,
created_by=user.id,
provider="vendor",
indexing_technique="high_quality",
embedding_model_provider="openai",
embedding_model="text-embedding-ada-002",
collection_binding_id=existing_binding_id,
)
update_data = {
"name": "new_name",
"indexing_technique": "high_quality",
"retrieval_model": "new_model",
}
result = DatasetService.update_dataset(dataset.id, update_data, user)
db.session.refresh(dataset)
assert dataset.name == "new_name"
assert dataset.indexing_technique == "high_quality"
assert dataset.embedding_model_provider == "openai"
assert dataset.embedding_model == "text-embedding-ada-002"
assert dataset.collection_binding_id == existing_binding_id
assert dataset.retrieval_model == "new_model"
assert result.id == dataset.id
def test_update_internal_dataset_embedding_model_update(self, db_session_with_containers):
"""Test updating internal dataset with new embedding model."""
user, tenant = DatasetUpdateTestDataFactory.create_account_with_tenant()
existing_binding_id = str(uuid4())
dataset = DatasetUpdateTestDataFactory.create_dataset(
tenant_id=tenant.id,
created_by=user.id,
provider="vendor",
indexing_technique="high_quality",
embedding_model_provider="openai",
embedding_model="text-embedding-ada-002",
collection_binding_id=existing_binding_id,
)
embedding_model = Mock()
embedding_model.model_name = "text-embedding-3-small"
embedding_model.provider = "openai"
binding = Mock()
binding.id = str(uuid4())
update_data = {
"indexing_technique": "high_quality",
"embedding_model_provider": "openai",
"embedding_model": "text-embedding-3-small",
"retrieval_model": "new_model",
}
with (
patch("services.dataset_service.current_user", user),
patch("services.dataset_service.ModelManager") as mock_model_manager,
patch(
"services.dataset_service.DatasetCollectionBindingService.get_dataset_collection_binding"
) as mock_get_binding,
patch("services.dataset_service.deal_dataset_vector_index_task") as mock_task,
patch("services.dataset_service.regenerate_summary_index_task") as mock_regenerate_task,
):
mock_model_manager.return_value.get_model_instance.return_value = embedding_model
mock_get_binding.return_value = binding
result = DatasetService.update_dataset(dataset.id, update_data, user)
mock_model_manager.return_value.get_model_instance.assert_called_once_with(
tenant_id=tenant.id,
provider="openai",
model_type=ModelType.TEXT_EMBEDDING,
model="text-embedding-3-small",
)
mock_get_binding.assert_called_once_with("openai", "text-embedding-3-small")
mock_task.delay.assert_called_once_with(dataset.id, "update")
mock_regenerate_task.delay.assert_called_once_with(
dataset.id,
regenerate_reason="embedding_model_changed",
regenerate_vectors_only=True,
)
db.session.refresh(dataset)
assert dataset.embedding_model == "text-embedding-3-small"
assert dataset.embedding_model_provider == "openai"
assert dataset.collection_binding_id == binding.id
assert dataset.retrieval_model == "new_model"
assert result.id == dataset.id
# ==================== Error Handling Tests ====================
def test_update_dataset_not_found_error(self, db_session_with_containers):
"""Test error when dataset is not found."""
user, _ = DatasetUpdateTestDataFactory.create_account_with_tenant()
update_data = {"name": "new_name"}
with pytest.raises(ValueError) as context:
DatasetService.update_dataset(str(uuid4()), update_data, user)
assert "Dataset not found" in str(context.value)
def test_update_dataset_permission_error(self, db_session_with_containers):
"""Test error when user doesn't have permission."""
owner, tenant = DatasetUpdateTestDataFactory.create_account_with_tenant(role=TenantAccountRole.OWNER)
outsider, _ = DatasetUpdateTestDataFactory.create_account_with_tenant(role=TenantAccountRole.NORMAL)
dataset = DatasetUpdateTestDataFactory.create_dataset(
tenant_id=tenant.id,
created_by=owner.id,
provider="vendor",
permission="only_me",
)
update_data = {"name": "new_name"}
with pytest.raises(NoPermissionError):
DatasetService.update_dataset(dataset.id, update_data, outsider)
def test_update_internal_dataset_embedding_model_error(self, db_session_with_containers):
"""Test error when embedding model is not available."""
user, tenant = DatasetUpdateTestDataFactory.create_account_with_tenant()
dataset = DatasetUpdateTestDataFactory.create_dataset(
tenant_id=tenant.id,
created_by=user.id,
provider="vendor",
indexing_technique="economy",
)
update_data = {
"indexing_technique": "high_quality",
"embedding_model_provider": "invalid_provider",
"embedding_model": "invalid_model",
"retrieval_model": "new_model",
}
with (
patch("services.dataset_service.current_user", user),
patch("services.dataset_service.ModelManager") as mock_model_manager,
):
mock_model_manager.return_value.get_model_instance.side_effect = Exception("No Embedding Model available")
with pytest.raises(Exception) as context:
DatasetService.update_dataset(dataset.id, update_data, user)
assert "No Embedding Model available".lower() in str(context.value).lower()
@@ -273,9 +273,10 @@ class TestWebAppAuthService:
# Arrange: Create banned account
fake = Faker()
password = fake.password(length=12)
unique_email = f"test_{uuid.uuid4().hex[:8]}@example.com"
account = Account(
email=fake.email(),
email=unique_email,
name=fake.name(),
interface_language="en-US",
status=AccountStatus.BANNED,
@@ -426,8 +427,7 @@ class TestWebAppAuthService:
- Correct return value (None)
"""
# Arrange: Use non-existent email
fake = Faker()
non_existent_email = fake.email()
non_existent_email = f"nonexistent_{uuid.uuid4().hex}@example.com"
# Act: Execute user retrieval
result = WebAppAuthService.get_user_through_email(non_existent_email)
@@ -0,0 +1,704 @@
"""Integration tests for dataset indexing task SQL behaviors using testcontainers."""
import uuid
from collections.abc import Sequence
from unittest.mock import MagicMock, patch
import pytest
from faker import Faker
from core.indexing_runner import DocumentIsPausedError
from enums.cloud_plan import CloudPlan
from models import Account, Tenant, TenantAccountJoin, TenantAccountRole
from models.dataset import Dataset, Document
from tasks.document_indexing_task import (
_document_indexing,
_document_indexing_with_tenant_queue,
document_indexing_task,
normal_document_indexing_task,
priority_document_indexing_task,
)
class _TrackedSessionContext:
def __init__(self, original_context_manager, opened_sessions: list, closed_sessions: list):
self._original_context_manager = original_context_manager
self._opened_sessions = opened_sessions
self._closed_sessions = closed_sessions
self._close_patcher = None
self._session = None
def __enter__(self):
self._session = self._original_context_manager.__enter__()
self._opened_sessions.append(self._session)
original_close = self._session.close
def _tracked_close(*args, **kwargs):
self._closed_sessions.append(self._session)
return original_close(*args, **kwargs)
self._close_patcher = patch.object(self._session, "close", side_effect=_tracked_close)
self._close_patcher.start()
return self._session
def __exit__(self, exc_type, exc_val, exc_tb):
try:
return self._original_context_manager.__exit__(exc_type, exc_val, exc_tb)
finally:
if self._close_patcher is not None:
self._close_patcher.stop()
@pytest.fixture(autouse=True)
def _ensure_testcontainers_db(db_session_with_containers):
"""Ensure this suite always runs on testcontainers infrastructure."""
return db_session_with_containers
@pytest.fixture
def session_close_tracker():
"""Track all sessions opened by session_factory and which were closed."""
opened_sessions = []
closed_sessions = []
from tasks import document_indexing_task as task_module
original_create_session = task_module.session_factory.create_session
def _tracked_create_session(*args, **kwargs):
original_context_manager = original_create_session(*args, **kwargs)
return _TrackedSessionContext(original_context_manager, opened_sessions, closed_sessions)
with patch.object(task_module.session_factory, "create_session", side_effect=_tracked_create_session):
yield {"opened_sessions": opened_sessions, "closed_sessions": closed_sessions}
@pytest.fixture
def patched_external_dependencies():
"""Patch non-DB collaborators while keeping database behavior real."""
with (
patch("tasks.document_indexing_task.IndexingRunner") as mock_indexing_runner,
patch("tasks.document_indexing_task.FeatureService") as mock_feature_service,
patch("tasks.document_indexing_task.generate_summary_index_task") as mock_summary_task,
):
mock_runner_instance = MagicMock()
mock_indexing_runner.return_value = mock_runner_instance
mock_features = MagicMock()
mock_features.billing.enabled = False
mock_features.billing.subscription.plan = CloudPlan.PROFESSIONAL
mock_features.vector_space.limit = 100
mock_features.vector_space.size = 0
mock_feature_service.get_features.return_value = mock_features
yield {
"indexing_runner": mock_indexing_runner,
"indexing_runner_instance": mock_runner_instance,
"feature_service": mock_feature_service,
"features": mock_features,
"summary_task": mock_summary_task,
}
class TestDatasetIndexingTaskIntegration:
"""1:1 SQL test migration from unit tests to testcontainers integration tests."""
def _create_test_dataset_and_documents(
self,
db_session_with_containers,
*,
document_count: int = 3,
document_ids: Sequence[str] | None = None,
) -> tuple[Dataset, list[Document]]:
"""Create a tenant dataset and waiting documents used by indexing tests."""
fake = Faker()
account = Account(
email=fake.email(),
name=fake.name(),
interface_language="en-US",
status="active",
)
db_session_with_containers.add(account)
db_session_with_containers.flush()
tenant = Tenant(name=fake.company(), status="normal")
db_session_with_containers.add(tenant)
db_session_with_containers.flush()
join = TenantAccountJoin(
tenant_id=tenant.id,
account_id=account.id,
role=TenantAccountRole.OWNER,
current=True,
)
db_session_with_containers.add(join)
dataset = Dataset(
id=fake.uuid4(),
tenant_id=tenant.id,
name=fake.company(),
description=fake.text(max_nb_chars=100),
data_source_type="upload_file",
indexing_technique="high_quality",
created_by=account.id,
)
db_session_with_containers.add(dataset)
if document_ids is None:
document_ids = [str(uuid.uuid4()) for _ in range(document_count)]
documents = []
for position, document_id in enumerate(document_ids):
document = Document(
id=document_id,
tenant_id=tenant.id,
dataset_id=dataset.id,
position=position,
data_source_type="upload_file",
batch="test_batch",
name=f"doc-{position}.txt",
created_from="upload_file",
created_by=account.id,
indexing_status="waiting",
enabled=True,
)
db_session_with_containers.add(document)
documents.append(document)
db_session_with_containers.commit()
db_session_with_containers.refresh(dataset)
return dataset, documents
def _query_document(self, db_session_with_containers, document_id: str) -> Document | None:
"""Return the latest persisted document state."""
return db_session_with_containers.query(Document).where(Document.id == document_id).first()
def _assert_documents_parsing(self, db_session_with_containers, document_ids: Sequence[str]) -> None:
"""Assert all target documents are persisted in parsing status."""
db_session_with_containers.expire_all()
for document_id in document_ids:
updated = self._query_document(db_session_with_containers, document_id)
assert updated is not None
assert updated.indexing_status == "parsing"
assert updated.processing_started_at is not None
def _assert_documents_error_contains(
self,
db_session_with_containers,
document_ids: Sequence[str],
expected_error_substring: str,
) -> None:
"""Assert all target documents are persisted in error status with message."""
db_session_with_containers.expire_all()
for document_id in document_ids:
updated = self._query_document(db_session_with_containers, document_id)
assert updated is not None
assert updated.indexing_status == "error"
assert updated.error is not None
assert expected_error_substring in updated.error
assert updated.stopped_at is not None
def _assert_all_opened_sessions_closed(self, session_close_tracker: dict) -> None:
"""Assert that every opened session is eventually closed."""
opened = session_close_tracker["opened_sessions"]
closed = session_close_tracker["closed_sessions"]
opened_ids = {id(session) for session in opened}
closed_ids = {id(session) for session in closed}
assert len(opened) >= 2
assert opened_ids <= closed_ids
def test_legacy_document_indexing_task_still_works(self, db_session_with_containers, patched_external_dependencies):
"""Ensure the legacy task entrypoint still updates parsing status."""
# Arrange
dataset, documents = self._create_test_dataset_and_documents(db_session_with_containers, document_count=2)
document_ids = [doc.id for doc in documents]
# Act
document_indexing_task(dataset.id, document_ids)
# Assert
patched_external_dependencies["indexing_runner_instance"].run.assert_called_once()
self._assert_documents_parsing(db_session_with_containers, document_ids)
def test_batch_processing_multiple_documents(self, db_session_with_containers, patched_external_dependencies):
"""Process multiple documents in one batch."""
# Arrange
dataset, documents = self._create_test_dataset_and_documents(db_session_with_containers, document_count=3)
document_ids = [doc.id for doc in documents]
# Act
_document_indexing(dataset.id, document_ids)
# Assert
patched_external_dependencies["indexing_runner_instance"].run.assert_called_once()
run_args = patched_external_dependencies["indexing_runner_instance"].run.call_args[0][0]
assert len(run_args) == len(document_ids)
self._assert_documents_parsing(db_session_with_containers, document_ids)
def test_batch_processing_with_limit_check(self, db_session_with_containers, patched_external_dependencies):
"""Reject batches larger than configured upload limit.
This test patches config only to force a deterministic limit branch while keeping SQL writes real.
"""
# Arrange
dataset, documents = self._create_test_dataset_and_documents(db_session_with_containers, document_count=3)
document_ids = [doc.id for doc in documents]
features = patched_external_dependencies["features"]
features.billing.enabled = True
features.billing.subscription.plan = CloudPlan.PROFESSIONAL
features.vector_space.limit = 100
features.vector_space.size = 50
# Act
with patch("tasks.document_indexing_task.dify_config.BATCH_UPLOAD_LIMIT", "2"):
_document_indexing(dataset.id, document_ids)
# Assert
patched_external_dependencies["indexing_runner_instance"].run.assert_not_called()
self._assert_documents_error_contains(db_session_with_containers, document_ids, "batch upload limit")
def test_batch_processing_sandbox_plan_single_document_only(
self, db_session_with_containers, patched_external_dependencies
):
"""Reject multi-document upload under sandbox plan."""
# Arrange
dataset, documents = self._create_test_dataset_and_documents(db_session_with_containers, document_count=2)
document_ids = [doc.id for doc in documents]
features = patched_external_dependencies["features"]
features.billing.enabled = True
features.billing.subscription.plan = CloudPlan.SANDBOX
# Act
_document_indexing(dataset.id, document_ids)
# Assert
patched_external_dependencies["indexing_runner_instance"].run.assert_not_called()
self._assert_documents_error_contains(db_session_with_containers, document_ids, "does not support batch upload")
def test_batch_processing_empty_document_list(self, db_session_with_containers, patched_external_dependencies):
"""Handle empty list input without failing."""
# Arrange
dataset, _ = self._create_test_dataset_and_documents(db_session_with_containers, document_count=0)
# Act
_document_indexing(dataset.id, [])
# Assert
patched_external_dependencies["indexing_runner_instance"].run.assert_called_once_with([])
def test_tenant_queue_dispatches_next_task_after_completion(
self, db_session_with_containers, patched_external_dependencies
):
"""Dispatch the next queued task after current tenant task completes.
Queue APIs are patched to isolate dispatch side effects while preserving DB assertions.
"""
# Arrange
dataset, documents = self._create_test_dataset_and_documents(db_session_with_containers, document_count=1)
document_ids = [doc.id for doc in documents]
next_task = {
"tenant_id": dataset.tenant_id,
"dataset_id": dataset.id,
"document_ids": [str(uuid.uuid4())],
}
task_dispatch_spy = MagicMock()
# Act
with (
patch("tasks.document_indexing_task.TenantIsolatedTaskQueue.pull_tasks", return_value=[next_task]),
patch("tasks.document_indexing_task.TenantIsolatedTaskQueue.set_task_waiting_time") as set_waiting_spy,
patch("tasks.document_indexing_task.TenantIsolatedTaskQueue.delete_task_key") as delete_key_spy,
):
_document_indexing_with_tenant_queue(dataset.tenant_id, dataset.id, document_ids, task_dispatch_spy)
# Assert
task_dispatch_spy.delay.assert_called_once_with(
tenant_id=next_task["tenant_id"],
dataset_id=next_task["dataset_id"],
document_ids=next_task["document_ids"],
)
set_waiting_spy.assert_called_once()
delete_key_spy.assert_not_called()
def test_tenant_queue_deletes_running_key_when_no_follow_up_tasks(
self, db_session_with_containers, patched_external_dependencies
):
"""Delete tenant running flag when queue has no pending tasks.
Queue APIs are patched to isolate dispatch side effects while preserving DB assertions.
"""
# Arrange
dataset, documents = self._create_test_dataset_and_documents(db_session_with_containers, document_count=1)
document_ids = [doc.id for doc in documents]
task_dispatch_spy = MagicMock()
# Act
with (
patch("tasks.document_indexing_task.TenantIsolatedTaskQueue.pull_tasks", return_value=[]),
patch("tasks.document_indexing_task.TenantIsolatedTaskQueue.delete_task_key") as delete_key_spy,
):
_document_indexing_with_tenant_queue(dataset.tenant_id, dataset.id, document_ids, task_dispatch_spy)
# Assert
task_dispatch_spy.delay.assert_not_called()
delete_key_spy.assert_called_once()
def test_validation_failure_sets_error_status_when_vector_space_at_limit(
self, db_session_with_containers, patched_external_dependencies
):
"""Set error status when vector space validation fails before runner phase."""
# Arrange
dataset, documents = self._create_test_dataset_and_documents(db_session_with_containers, document_count=3)
document_ids = [doc.id for doc in documents]
features = patched_external_dependencies["features"]
features.billing.enabled = True
features.billing.subscription.plan = CloudPlan.PROFESSIONAL
features.vector_space.limit = 100
features.vector_space.size = 100
# Act
_document_indexing(dataset.id, document_ids)
# Assert
patched_external_dependencies["indexing_runner_instance"].run.assert_not_called()
self._assert_documents_error_contains(db_session_with_containers, document_ids, "over the limit")
def test_runner_exception_does_not_crash_indexing_task(
self, db_session_with_containers, patched_external_dependencies
):
"""Catch generic runner exceptions without crashing the task."""
# Arrange
dataset, documents = self._create_test_dataset_and_documents(db_session_with_containers, document_count=2)
document_ids = [doc.id for doc in documents]
patched_external_dependencies["indexing_runner_instance"].run.side_effect = Exception("runner failed")
# Act
_document_indexing(dataset.id, document_ids)
# Assert
patched_external_dependencies["indexing_runner_instance"].run.assert_called_once()
self._assert_documents_parsing(db_session_with_containers, document_ids)
def test_document_paused_error_handling(self, db_session_with_containers, patched_external_dependencies):
"""Handle DocumentIsPausedError and keep persisted state consistent."""
# Arrange
dataset, documents = self._create_test_dataset_and_documents(db_session_with_containers, document_count=2)
document_ids = [doc.id for doc in documents]
patched_external_dependencies["indexing_runner_instance"].run.side_effect = DocumentIsPausedError("paused")
# Act
_document_indexing(dataset.id, document_ids)
# Assert
patched_external_dependencies["indexing_runner_instance"].run.assert_called_once()
self._assert_documents_parsing(db_session_with_containers, document_ids)
def test_dataset_not_found_error_handling(self, patched_external_dependencies):
"""Exit gracefully when dataset does not exist."""
# Arrange
missing_dataset_id = str(uuid.uuid4())
missing_document_id = str(uuid.uuid4())
# Act
_document_indexing(missing_dataset_id, [missing_document_id])
# Assert
patched_external_dependencies["indexing_runner_instance"].run.assert_not_called()
def test_tenant_queue_error_handling_still_processes_next_task(
self, db_session_with_containers, patched_external_dependencies
):
"""Even on current task failure, enqueue the next waiting tenant task.
Queue APIs are patched to isolate dispatch side effects while preserving DB assertions.
"""
# Arrange
dataset, documents = self._create_test_dataset_and_documents(db_session_with_containers, document_count=1)
document_ids = [doc.id for doc in documents]
next_task = {
"tenant_id": dataset.tenant_id,
"dataset_id": dataset.id,
"document_ids": [str(uuid.uuid4())],
}
task_dispatch_spy = MagicMock()
# Act
with (
patch("tasks.document_indexing_task._document_indexing", side_effect=Exception("failed")),
patch("tasks.document_indexing_task.TenantIsolatedTaskQueue.pull_tasks", return_value=[next_task]),
patch("tasks.document_indexing_task.TenantIsolatedTaskQueue.set_task_waiting_time"),
):
_document_indexing_with_tenant_queue(dataset.tenant_id, dataset.id, document_ids, task_dispatch_spy)
# Assert
task_dispatch_spy.delay.assert_called_once()
def test_sessions_close_on_successful_indexing(
self,
db_session_with_containers,
patched_external_dependencies,
session_close_tracker,
):
"""Close all opened sessions in successful indexing path."""
# Arrange
dataset, documents = self._create_test_dataset_and_documents(db_session_with_containers, document_count=2)
document_ids = [doc.id for doc in documents]
# Act
_document_indexing(dataset.id, document_ids)
# Assert
self._assert_all_opened_sessions_closed(session_close_tracker)
def test_sessions_close_when_runner_raises(
self,
db_session_with_containers,
patched_external_dependencies,
session_close_tracker,
):
"""Close opened sessions even when runner fails."""
# Arrange
dataset, documents = self._create_test_dataset_and_documents(db_session_with_containers, document_count=2)
document_ids = [doc.id for doc in documents]
patched_external_dependencies["indexing_runner_instance"].run.side_effect = Exception("boom")
# Act
_document_indexing(dataset.id, document_ids)
# Assert
self._assert_all_opened_sessions_closed(session_close_tracker)
def test_multiple_documents_with_mixed_success_and_failure(
self, db_session_with_containers, patched_external_dependencies
):
"""Process only existing documents when request includes missing ids."""
# Arrange
dataset, documents = self._create_test_dataset_and_documents(db_session_with_containers, document_count=2)
existing_ids = [doc.id for doc in documents]
mixed_ids = [existing_ids[0], str(uuid.uuid4()), existing_ids[1]]
# Act
_document_indexing(dataset.id, mixed_ids)
# Assert
run_args = patched_external_dependencies["indexing_runner_instance"].run.call_args[0][0]
assert len(run_args) == 2
self._assert_documents_parsing(db_session_with_containers, existing_ids)
def test_tenant_queue_dispatches_up_to_concurrency_limit(
self, db_session_with_containers, patched_external_dependencies
):
"""Dispatch only up to configured concurrency under queued backlog burst.
Queue APIs are patched to isolate dispatch side effects while preserving DB assertions.
"""
# Arrange
dataset, documents = self._create_test_dataset_and_documents(db_session_with_containers, document_count=1)
document_ids = [doc.id for doc in documents]
concurrency_limit = 3
backlog_size = 20
pending_tasks = [
{"tenant_id": dataset.tenant_id, "dataset_id": dataset.id, "document_ids": [f"doc_{idx}"]}
for idx in range(backlog_size)
]
task_dispatch_spy = MagicMock()
# Act
with (
patch("tasks.document_indexing_task.dify_config.TENANT_ISOLATED_TASK_CONCURRENCY", concurrency_limit),
patch(
"tasks.document_indexing_task.TenantIsolatedTaskQueue.pull_tasks",
return_value=pending_tasks[:concurrency_limit],
),
patch("tasks.document_indexing_task.TenantIsolatedTaskQueue.set_task_waiting_time") as set_waiting_spy,
):
_document_indexing_with_tenant_queue(dataset.tenant_id, dataset.id, document_ids, task_dispatch_spy)
# Assert
assert task_dispatch_spy.delay.call_count == concurrency_limit
assert set_waiting_spy.call_count == concurrency_limit
def test_task_queue_fifo_ordering(self, db_session_with_containers, patched_external_dependencies):
"""Keep FIFO ordering when dispatching next queued tasks.
Queue APIs are patched to isolate dispatch side effects while preserving DB assertions.
"""
# Arrange
dataset, documents = self._create_test_dataset_and_documents(db_session_with_containers, document_count=1)
document_ids = [doc.id for doc in documents]
ordered_tasks = [
{"tenant_id": dataset.tenant_id, "dataset_id": dataset.id, "document_ids": ["task_A"]},
{"tenant_id": dataset.tenant_id, "dataset_id": dataset.id, "document_ids": ["task_B"]},
{"tenant_id": dataset.tenant_id, "dataset_id": dataset.id, "document_ids": ["task_C"]},
]
task_dispatch_spy = MagicMock()
# Act
with (
patch("tasks.document_indexing_task.dify_config.TENANT_ISOLATED_TASK_CONCURRENCY", 3),
patch("tasks.document_indexing_task.TenantIsolatedTaskQueue.pull_tasks", return_value=ordered_tasks),
patch("tasks.document_indexing_task.TenantIsolatedTaskQueue.set_task_waiting_time"),
):
_document_indexing_with_tenant_queue(dataset.tenant_id, dataset.id, document_ids, task_dispatch_spy)
# Assert
assert task_dispatch_spy.delay.call_count == 3
for index, expected_task in enumerate(ordered_tasks):
assert task_dispatch_spy.delay.call_args_list[index].kwargs["document_ids"] == expected_task["document_ids"]
def test_billing_disabled_skips_limit_checks(self, db_session_with_containers, patched_external_dependencies):
"""Skip limit checks when billing feature is disabled."""
# Arrange
large_document_ids = [str(uuid.uuid4()) for _ in range(100)]
dataset, _ = self._create_test_dataset_and_documents(
db_session_with_containers,
document_ids=large_document_ids,
)
features = patched_external_dependencies["features"]
features.billing.enabled = False
# Act
_document_indexing(dataset.id, large_document_ids)
# Assert
run_args = patched_external_dependencies["indexing_runner_instance"].run.call_args[0][0]
assert len(run_args) == 100
self._assert_documents_parsing(db_session_with_containers, large_document_ids)
def test_complete_workflow_normal_task(self, db_session_with_containers, patched_external_dependencies):
"""Run end-to-end normal queue workflow with tenant queue cleanup.
Queue APIs are patched to isolate dispatch side effects while preserving DB assertions.
"""
# Arrange
dataset, documents = self._create_test_dataset_and_documents(db_session_with_containers, document_count=2)
document_ids = [doc.id for doc in documents]
# Act
with (
patch("tasks.document_indexing_task.TenantIsolatedTaskQueue.pull_tasks", return_value=[]),
patch("tasks.document_indexing_task.TenantIsolatedTaskQueue.delete_task_key") as delete_key_spy,
):
normal_document_indexing_task(dataset.tenant_id, dataset.id, document_ids)
# Assert
patched_external_dependencies["indexing_runner_instance"].run.assert_called_once()
self._assert_documents_parsing(db_session_with_containers, document_ids)
delete_key_spy.assert_called_once()
def test_complete_workflow_priority_task(self, db_session_with_containers, patched_external_dependencies):
"""Run end-to-end priority queue workflow with tenant queue cleanup.
Queue APIs are patched to isolate dispatch side effects while preserving DB assertions.
"""
# Arrange
dataset, documents = self._create_test_dataset_and_documents(db_session_with_containers, document_count=2)
document_ids = [doc.id for doc in documents]
# Act
with (
patch("tasks.document_indexing_task.TenantIsolatedTaskQueue.pull_tasks", return_value=[]),
patch("tasks.document_indexing_task.TenantIsolatedTaskQueue.delete_task_key") as delete_key_spy,
):
priority_document_indexing_task(dataset.tenant_id, dataset.id, document_ids)
# Assert
patched_external_dependencies["indexing_runner_instance"].run.assert_called_once()
self._assert_documents_parsing(db_session_with_containers, document_ids)
delete_key_spy.assert_called_once()
def test_single_document_processing(self, db_session_with_containers, patched_external_dependencies):
"""Process the minimum batch size (single document)."""
# Arrange
dataset, documents = self._create_test_dataset_and_documents(db_session_with_containers, document_count=1)
document_id = documents[0].id
# Act
_document_indexing(dataset.id, [document_id])
# Assert
run_args = patched_external_dependencies["indexing_runner_instance"].run.call_args[0][0]
assert len(run_args) == 1
self._assert_documents_parsing(db_session_with_containers, [document_id])
def test_document_with_special_characters_in_id(self, db_session_with_containers, patched_external_dependencies):
"""Handle standard UUID ids with hyphen characters safely."""
# Arrange
special_document_id = str(uuid.uuid4())
dataset, _ = self._create_test_dataset_and_documents(
db_session_with_containers,
document_ids=[special_document_id],
)
# Act
_document_indexing(dataset.id, [special_document_id])
# Assert
self._assert_documents_parsing(db_session_with_containers, [special_document_id])
def test_zero_vector_space_limit_allows_unlimited(self, db_session_with_containers, patched_external_dependencies):
"""Treat vector limit 0 as unlimited and continue indexing."""
# Arrange
dataset, documents = self._create_test_dataset_and_documents(db_session_with_containers, document_count=3)
document_ids = [doc.id for doc in documents]
features = patched_external_dependencies["features"]
features.billing.enabled = True
features.billing.subscription.plan = CloudPlan.PROFESSIONAL
features.vector_space.limit = 0
features.vector_space.size = 1000
# Act
_document_indexing(dataset.id, document_ids)
# Assert
patched_external_dependencies["indexing_runner_instance"].run.assert_called_once()
self._assert_documents_parsing(db_session_with_containers, document_ids)
def test_negative_vector_space_values_handled_gracefully(
self, db_session_with_containers, patched_external_dependencies
):
"""Treat negative vector limits as non-blocking and continue indexing."""
# Arrange
dataset, documents = self._create_test_dataset_and_documents(db_session_with_containers, document_count=3)
document_ids = [doc.id for doc in documents]
features = patched_external_dependencies["features"]
features.billing.enabled = True
features.billing.subscription.plan = CloudPlan.PROFESSIONAL
features.vector_space.limit = -1
features.vector_space.size = 100
# Act
_document_indexing(dataset.id, document_ids)
# Assert
patched_external_dependencies["indexing_runner_instance"].run.assert_called_once()
self._assert_documents_parsing(db_session_with_containers, document_ids)
def test_large_document_batch_processing(self, db_session_with_containers, patched_external_dependencies):
"""Process a batch exactly at configured upload limit.
This test patches config only to force a deterministic limit branch while keeping SQL writes real.
"""
# Arrange
batch_limit = 50
document_ids = [str(uuid.uuid4()) for _ in range(batch_limit)]
dataset, _ = self._create_test_dataset_and_documents(
db_session_with_containers,
document_ids=document_ids,
)
features = patched_external_dependencies["features"]
features.billing.enabled = True
features.billing.subscription.plan = CloudPlan.PROFESSIONAL
features.vector_space.limit = 10000
features.vector_space.size = 0
# Act
with patch("tasks.document_indexing_task.dify_config.BATCH_UPLOAD_LIMIT", str(batch_limit)):
_document_indexing(dataset.id, document_ids)
# Assert
run_args = patched_external_dependencies["indexing_runner_instance"].run.call_args[0][0]
assert len(run_args) == batch_limit
self._assert_documents_parsing(db_session_with_containers, document_ids)
@@ -50,8 +50,26 @@ class TestDealDatasetVectorIndexTask:
mock_factory.return_value = mock_instance
yield mock_factory
@pytest.fixture
def account_and_tenant(self, db_session_with_containers, mock_external_service_dependencies):
"""Create an account with an owner tenant for testing.
Returns a tuple of (account, tenant) where tenant is guaranteed to be non-None.
"""
fake = Faker()
account = AccountService.create_account(
email=fake.email(),
name=fake.name(),
interface_language="en-US",
password=fake.password(length=12),
)
TenantService.create_owner_tenant_if_not_exist(account, name=fake.company())
tenant = account.current_tenant
assert tenant is not None
return account, tenant
def test_deal_dataset_vector_index_task_remove_action_success(
self, db_session_with_containers, mock_index_processor_factory, mock_external_service_dependencies
self, db_session_with_containers, mock_index_processor_factory, account_and_tenant
):
"""
Test successful removal of dataset vector index.
@@ -63,16 +81,7 @@ class TestDealDatasetVectorIndexTask:
4. Completes without errors
"""
fake = Faker()
# Create test data
account = AccountService.create_account(
email=fake.email(),
name=fake.name(),
interface_language="en-US",
password=fake.password(length=12),
)
TenantService.create_owner_tenant_if_not_exist(account, name=fake.company())
tenant = account.current_tenant
account, tenant = account_and_tenant
# Create dataset
dataset = Dataset(
@@ -118,7 +127,7 @@ class TestDealDatasetVectorIndexTask:
assert mock_processor.clean.call_count >= 0 # For now, just check it doesn't fail
def test_deal_dataset_vector_index_task_add_action_success(
self, db_session_with_containers, mock_index_processor_factory, mock_external_service_dependencies
self, db_session_with_containers, mock_index_processor_factory, account_and_tenant
):
"""
Test successful addition of dataset vector index.
@@ -132,16 +141,7 @@ class TestDealDatasetVectorIndexTask:
6. Updates document status to completed
"""
fake = Faker()
# Create test data
account = AccountService.create_account(
email=fake.email(),
name=fake.name(),
interface_language="en-US",
password=fake.password(length=12),
)
TenantService.create_owner_tenant_if_not_exist(account, name=fake.company())
tenant = account.current_tenant
account, tenant = account_and_tenant
# Create dataset
dataset = Dataset(
@@ -227,7 +227,7 @@ class TestDealDatasetVectorIndexTask:
mock_processor.load.assert_called_once()
def test_deal_dataset_vector_index_task_update_action_success(
self, db_session_with_containers, mock_index_processor_factory, mock_external_service_dependencies
self, db_session_with_containers, mock_index_processor_factory, account_and_tenant
):
"""
Test successful update of dataset vector index.
@@ -242,16 +242,7 @@ class TestDealDatasetVectorIndexTask:
7. Updates document status to completed
"""
fake = Faker()
# Create test data
account = AccountService.create_account(
email=fake.email(),
name=fake.name(),
interface_language="en-US",
password=fake.password(length=12),
)
TenantService.create_owner_tenant_if_not_exist(account, name=fake.company())
tenant = account.current_tenant
account, tenant = account_and_tenant
# Create dataset with parent-child index
dataset = Dataset(
@@ -338,7 +329,7 @@ class TestDealDatasetVectorIndexTask:
mock_processor.load.assert_called_once()
def test_deal_dataset_vector_index_task_dataset_not_found_error(
self, db_session_with_containers, mock_index_processor_factory, mock_external_service_dependencies
self, db_session_with_containers, mock_index_processor_factory, account_and_tenant
):
"""
Test task behavior when dataset is not found.
@@ -358,7 +349,7 @@ class TestDealDatasetVectorIndexTask:
mock_processor.load.assert_not_called()
def test_deal_dataset_vector_index_task_add_action_no_documents(
self, db_session_with_containers, mock_index_processor_factory, mock_external_service_dependencies
self, db_session_with_containers, mock_index_processor_factory, account_and_tenant
):
"""
Test add action when no documents exist for the dataset.
@@ -367,16 +358,7 @@ class TestDealDatasetVectorIndexTask:
a dataset exists but has no documents to process.
"""
fake = Faker()
# Create test data
account = AccountService.create_account(
email=fake.email(),
name=fake.name(),
interface_language="en-US",
password=fake.password(length=12),
)
TenantService.create_owner_tenant_if_not_exist(account, name=fake.company())
tenant = account.current_tenant
account, tenant = account_and_tenant
# Create dataset without documents
dataset = Dataset(
@@ -399,7 +381,7 @@ class TestDealDatasetVectorIndexTask:
mock_processor.load.assert_not_called()
def test_deal_dataset_vector_index_task_add_action_no_segments(
self, db_session_with_containers, mock_index_processor_factory, mock_external_service_dependencies
self, db_session_with_containers, mock_index_processor_factory, account_and_tenant
):
"""
Test add action when documents exist but have no segments.
@@ -408,16 +390,7 @@ class TestDealDatasetVectorIndexTask:
documents exist but contain no segments to process.
"""
fake = Faker()
# Create test data
account = AccountService.create_account(
email=fake.email(),
name=fake.name(),
interface_language="en-US",
password=fake.password(length=12),
)
TenantService.create_owner_tenant_if_not_exist(account, name=fake.company())
tenant = account.current_tenant
account, tenant = account_and_tenant
# Create dataset
dataset = Dataset(
@@ -464,7 +437,7 @@ class TestDealDatasetVectorIndexTask:
mock_processor.load.assert_not_called()
def test_deal_dataset_vector_index_task_update_action_no_documents(
self, db_session_with_containers, mock_index_processor_factory, mock_external_service_dependencies
self, db_session_with_containers, mock_index_processor_factory, account_and_tenant
):
"""
Test update action when no documents exist for the dataset.
@@ -473,16 +446,7 @@ class TestDealDatasetVectorIndexTask:
a dataset exists but has no documents to process during update.
"""
fake = Faker()
# Create test data
account = AccountService.create_account(
email=fake.email(),
name=fake.name(),
interface_language="en-US",
password=fake.password(length=12),
)
TenantService.create_owner_tenant_if_not_exist(account, name=fake.company())
tenant = account.current_tenant
account, tenant = account_and_tenant
# Create dataset without documents
dataset = Dataset(
@@ -506,7 +470,7 @@ class TestDealDatasetVectorIndexTask:
mock_processor.load.assert_not_called()
def test_deal_dataset_vector_index_task_add_action_with_exception_handling(
self, db_session_with_containers, mock_index_processor_factory, mock_external_service_dependencies
self, db_session_with_containers, mock_index_processor_factory, account_and_tenant
):
"""
Test add action with exception handling during processing.
@@ -515,16 +479,7 @@ class TestDealDatasetVectorIndexTask:
during document processing and updates document status to error.
"""
fake = Faker()
# Create test data
account = AccountService.create_account(
email=fake.email(),
name=fake.name(),
interface_language="en-US",
password=fake.password(length=12),
)
TenantService.create_owner_tenant_if_not_exist(account, name=fake.company())
tenant = account.current_tenant
account, tenant = account_and_tenant
# Create dataset
dataset = Dataset(
@@ -611,7 +566,7 @@ class TestDealDatasetVectorIndexTask:
assert "Test exception during indexing" in updated_document.error
def test_deal_dataset_vector_index_task_with_custom_index_type(
self, db_session_with_containers, mock_index_processor_factory, mock_external_service_dependencies
self, db_session_with_containers, mock_index_processor_factory, account_and_tenant
):
"""
Test task behavior with custom index type (QA_INDEX).
@@ -620,16 +575,7 @@ class TestDealDatasetVectorIndexTask:
and initializes the appropriate index processor.
"""
fake = Faker()
# Create test data
account = AccountService.create_account(
email=fake.email(),
name=fake.name(),
interface_language="en-US",
password=fake.password(length=12),
)
TenantService.create_owner_tenant_if_not_exist(account, name=fake.company())
tenant = account.current_tenant
account, tenant = account_and_tenant
# Create dataset with custom index type
dataset = Dataset(
@@ -696,7 +642,7 @@ class TestDealDatasetVectorIndexTask:
mock_processor.load.assert_called_once()
def test_deal_dataset_vector_index_task_with_default_index_type(
self, db_session_with_containers, mock_index_processor_factory, mock_external_service_dependencies
self, db_session_with_containers, mock_index_processor_factory, account_and_tenant
):
"""
Test task behavior with default index type (PARAGRAPH_INDEX).
@@ -705,16 +651,7 @@ class TestDealDatasetVectorIndexTask:
when dataset.doc_form is None.
"""
fake = Faker()
# Create test data
account = AccountService.create_account(
email=fake.email(),
name=fake.name(),
interface_language="en-US",
password=fake.password(length=12),
)
TenantService.create_owner_tenant_if_not_exist(account, name=fake.company())
tenant = account.current_tenant
account, tenant = account_and_tenant
# Create dataset without doc_form (should use default)
dataset = Dataset(
@@ -781,7 +718,7 @@ class TestDealDatasetVectorIndexTask:
mock_processor.load.assert_called_once()
def test_deal_dataset_vector_index_task_multiple_documents_processing(
self, db_session_with_containers, mock_index_processor_factory, mock_external_service_dependencies
self, db_session_with_containers, mock_index_processor_factory, account_and_tenant
):
"""
Test task processing with multiple documents and segments.
@@ -790,16 +727,7 @@ class TestDealDatasetVectorIndexTask:
and their segments in sequence.
"""
fake = Faker()
# Create test data
account = AccountService.create_account(
email=fake.email(),
name=fake.name(),
interface_language="en-US",
password=fake.password(length=12),
)
TenantService.create_owner_tenant_if_not_exist(account, name=fake.company())
tenant = account.current_tenant
account, tenant = account_and_tenant
# Create dataset
dataset = Dataset(
@@ -893,7 +821,7 @@ class TestDealDatasetVectorIndexTask:
assert mock_processor.load.call_count == 3
def test_deal_dataset_vector_index_task_document_status_transitions(
self, db_session_with_containers, mock_index_processor_factory, mock_external_service_dependencies
self, db_session_with_containers, mock_index_processor_factory, account_and_tenant
):
"""
Test document status transitions during task execution.
@@ -902,16 +830,7 @@ class TestDealDatasetVectorIndexTask:
'completed' to 'indexing' and back to 'completed' during processing.
"""
fake = Faker()
# Create test data
account = AccountService.create_account(
email=fake.email(),
name=fake.name(),
interface_language="en-US",
password=fake.password(length=12),
)
TenantService.create_owner_tenant_if_not_exist(account, name=fake.company())
tenant = account.current_tenant
account, tenant = account_and_tenant
# Create dataset
dataset = Dataset(
@@ -999,7 +918,7 @@ class TestDealDatasetVectorIndexTask:
assert updated_document.indexing_status == "completed"
def test_deal_dataset_vector_index_task_with_disabled_documents(
self, db_session_with_containers, mock_index_processor_factory, mock_external_service_dependencies
self, db_session_with_containers, mock_index_processor_factory, account_and_tenant
):
"""
Test task behavior with disabled documents.
@@ -1008,16 +927,7 @@ class TestDealDatasetVectorIndexTask:
during processing.
"""
fake = Faker()
# Create test data
account = AccountService.create_account(
email=fake.email(),
name=fake.name(),
interface_language="en-US",
password=fake.password(length=12),
)
TenantService.create_owner_tenant_if_not_exist(account, name=fake.company())
tenant = account.current_tenant
account, tenant = account_and_tenant
# Create dataset
dataset = Dataset(
@@ -1129,7 +1039,7 @@ class TestDealDatasetVectorIndexTask:
mock_processor.load.assert_called_once()
def test_deal_dataset_vector_index_task_with_archived_documents(
self, db_session_with_containers, mock_index_processor_factory, mock_external_service_dependencies
self, db_session_with_containers, mock_index_processor_factory, account_and_tenant
):
"""
Test task behavior with archived documents.
@@ -1138,16 +1048,7 @@ class TestDealDatasetVectorIndexTask:
during processing.
"""
fake = Faker()
# Create test data
account = AccountService.create_account(
email=fake.email(),
name=fake.name(),
interface_language="en-US",
password=fake.password(length=12),
)
TenantService.create_owner_tenant_if_not_exist(account, name=fake.company())
tenant = account.current_tenant
account, tenant = account_and_tenant
# Create dataset
dataset = Dataset(
@@ -1259,7 +1160,7 @@ class TestDealDatasetVectorIndexTask:
mock_processor.load.assert_called_once()
def test_deal_dataset_vector_index_task_with_incomplete_documents(
self, db_session_with_containers, mock_index_processor_factory, mock_external_service_dependencies
self, db_session_with_containers, mock_index_processor_factory, account_and_tenant
):
"""
Test task behavior with documents that have incomplete indexing status.
@@ -1268,16 +1169,7 @@ class TestDealDatasetVectorIndexTask:
incomplete indexing status during processing.
"""
fake = Faker()
# Create test data
account = AccountService.create_account(
email=fake.email(),
name=fake.name(),
interface_language="en-US",
password=fake.password(length=12),
)
TenantService.create_owner_tenant_if_not_exist(account, name=fake.company())
tenant = account.current_tenant
account, tenant = account_and_tenant
# Create dataset
dataset = Dataset(
@@ -0,0 +1,464 @@
"""
Integration tests for document_indexing_sync_task using testcontainers.
This module validates SQL-backed behavior for document sync flows:
- Notion sync precondition checks
- Segment cleanup and document state updates
- Credential and indexing error handling
"""
import json
from unittest.mock import Mock, patch
from uuid import uuid4
import pytest
from psycopg2.extensions import register_adapter
from psycopg2.extras import Json
from core.indexing_runner import DocumentIsPausedError, IndexingRunner
from models import Account, Tenant, TenantAccountJoin, TenantAccountRole
from models.dataset import Dataset, Document, DocumentSegment
from tasks.document_indexing_sync_task import document_indexing_sync_task
@pytest.fixture(autouse=True)
def _register_dict_adapter_for_psycopg2():
"""Align test DB adapter behavior with dict payloads used in task update flow."""
register_adapter(dict, Json)
class DocumentIndexingSyncTaskTestDataFactory:
"""Create real DB entities for document indexing sync integration tests."""
@staticmethod
def create_account_with_tenant(db_session_with_containers) -> tuple[Account, Tenant]:
account = Account(
email=f"{uuid4()}@example.com",
name=f"user-{uuid4()}",
interface_language="en-US",
status="active",
)
db_session_with_containers.add(account)
db_session_with_containers.flush()
tenant = Tenant(name=f"tenant-{account.id}", status="normal")
db_session_with_containers.add(tenant)
db_session_with_containers.flush()
join = TenantAccountJoin(
tenant_id=tenant.id,
account_id=account.id,
role=TenantAccountRole.OWNER,
current=True,
)
db_session_with_containers.add(join)
db_session_with_containers.commit()
return account, tenant
@staticmethod
def create_dataset(db_session_with_containers, tenant_id: str, created_by: str) -> Dataset:
dataset = Dataset(
tenant_id=tenant_id,
name=f"dataset-{uuid4()}",
description="sync test dataset",
data_source_type="notion_import",
indexing_technique="high_quality",
created_by=created_by,
)
db_session_with_containers.add(dataset)
db_session_with_containers.commit()
return dataset
@staticmethod
def create_document(
db_session_with_containers,
*,
tenant_id: str,
dataset_id: str,
created_by: str,
data_source_info: dict | None,
indexing_status: str = "completed",
) -> Document:
document = Document(
tenant_id=tenant_id,
dataset_id=dataset_id,
position=0,
data_source_type="notion_import",
data_source_info=json.dumps(data_source_info) if data_source_info is not None else None,
batch="test-batch",
name=f"doc-{uuid4()}",
created_from="notion_import",
created_by=created_by,
indexing_status=indexing_status,
enabled=True,
doc_form="text_model",
doc_language="en",
)
db_session_with_containers.add(document)
db_session_with_containers.commit()
return document
@staticmethod
def create_segments(
db_session_with_containers,
*,
tenant_id: str,
dataset_id: str,
document_id: str,
created_by: str,
count: int = 3,
) -> list[DocumentSegment]:
segments: list[DocumentSegment] = []
for i in range(count):
segment = DocumentSegment(
tenant_id=tenant_id,
dataset_id=dataset_id,
document_id=document_id,
position=i,
content=f"segment-{i}",
answer=None,
word_count=10,
tokens=5,
index_node_id=f"node-{document_id}-{i}",
status="completed",
created_by=created_by,
)
db_session_with_containers.add(segment)
segments.append(segment)
db_session_with_containers.commit()
return segments
class TestDocumentIndexingSyncTask:
"""Integration tests for document_indexing_sync_task with real database assertions."""
@pytest.fixture
def mock_external_dependencies(self):
"""Patch only external collaborators; keep DB access real."""
with (
patch("tasks.document_indexing_sync_task.DatasourceProviderService") as mock_datasource_service_class,
patch("tasks.document_indexing_sync_task.NotionExtractor") as mock_notion_extractor_class,
patch("tasks.document_indexing_sync_task.IndexProcessorFactory") as mock_index_processor_factory,
patch("tasks.document_indexing_sync_task.IndexingRunner") as mock_indexing_runner_class,
):
datasource_service = Mock()
datasource_service.get_datasource_credentials.return_value = {"integration_secret": "test_token"}
mock_datasource_service_class.return_value = datasource_service
notion_extractor = Mock()
notion_extractor.get_notion_last_edited_time.return_value = "2024-01-02T00:00:00Z"
mock_notion_extractor_class.return_value = notion_extractor
index_processor = Mock()
index_processor.clean = Mock()
mock_index_processor_factory.return_value.init_index_processor.return_value = index_processor
indexing_runner = Mock(spec=IndexingRunner)
indexing_runner.run = Mock()
mock_indexing_runner_class.return_value = indexing_runner
yield {
"datasource_service": datasource_service,
"notion_extractor": notion_extractor,
"notion_extractor_class": mock_notion_extractor_class,
"index_processor": index_processor,
"index_processor_factory": mock_index_processor_factory,
"indexing_runner": indexing_runner,
}
def _create_notion_sync_context(self, db_session_with_containers, *, data_source_info: dict | None = None):
account, tenant = DocumentIndexingSyncTaskTestDataFactory.create_account_with_tenant(db_session_with_containers)
dataset = DocumentIndexingSyncTaskTestDataFactory.create_dataset(
db_session_with_containers,
tenant_id=tenant.id,
created_by=account.id,
)
notion_info = data_source_info or {
"notion_workspace_id": str(uuid4()),
"notion_page_id": str(uuid4()),
"type": "page",
"last_edited_time": "2024-01-01T00:00:00Z",
"credential_id": str(uuid4()),
}
document = DocumentIndexingSyncTaskTestDataFactory.create_document(
db_session_with_containers,
tenant_id=tenant.id,
dataset_id=dataset.id,
created_by=account.id,
data_source_info=notion_info,
indexing_status="completed",
)
segments = DocumentIndexingSyncTaskTestDataFactory.create_segments(
db_session_with_containers,
tenant_id=tenant.id,
dataset_id=dataset.id,
document_id=document.id,
created_by=account.id,
count=3,
)
return {
"account": account,
"tenant": tenant,
"dataset": dataset,
"document": document,
"segments": segments,
"node_ids": [segment.index_node_id for segment in segments],
"notion_info": notion_info,
}
def test_document_not_found(self, db_session_with_containers, mock_external_dependencies):
"""Test that task handles missing document gracefully."""
# Arrange
dataset_id = str(uuid4())
document_id = str(uuid4())
# Act
document_indexing_sync_task(dataset_id, document_id)
# Assert
mock_external_dependencies["datasource_service"].get_datasource_credentials.assert_not_called()
mock_external_dependencies["indexing_runner"].run.assert_not_called()
def test_missing_notion_workspace_id(self, db_session_with_containers, mock_external_dependencies):
"""Test that task raises error when notion_workspace_id is missing."""
# Arrange
context = self._create_notion_sync_context(
db_session_with_containers,
data_source_info={
"notion_page_id": str(uuid4()),
"type": "page",
"last_edited_time": "2024-01-01T00:00:00Z",
},
)
# Act & Assert
with pytest.raises(ValueError, match="no notion page found"):
document_indexing_sync_task(context["dataset"].id, context["document"].id)
def test_missing_notion_page_id(self, db_session_with_containers, mock_external_dependencies):
"""Test that task raises error when notion_page_id is missing."""
# Arrange
context = self._create_notion_sync_context(
db_session_with_containers,
data_source_info={
"notion_workspace_id": str(uuid4()),
"type": "page",
"last_edited_time": "2024-01-01T00:00:00Z",
},
)
# Act & Assert
with pytest.raises(ValueError, match="no notion page found"):
document_indexing_sync_task(context["dataset"].id, context["document"].id)
def test_empty_data_source_info(self, db_session_with_containers, mock_external_dependencies):
"""Test that task raises error when data_source_info is empty."""
# Arrange
context = self._create_notion_sync_context(db_session_with_containers, data_source_info=None)
db_session_with_containers.query(Document).where(Document.id == context["document"].id).update(
{"data_source_info": None}
)
db_session_with_containers.commit()
# Act & Assert
with pytest.raises(ValueError, match="no notion page found"):
document_indexing_sync_task(context["dataset"].id, context["document"].id)
def test_credential_not_found(self, db_session_with_containers, mock_external_dependencies):
"""Test that task sets document error state when credential is missing."""
# Arrange
context = self._create_notion_sync_context(db_session_with_containers)
mock_external_dependencies["datasource_service"].get_datasource_credentials.return_value = None
# Act
document_indexing_sync_task(context["dataset"].id, context["document"].id)
# Assert
db_session_with_containers.expire_all()
updated_document = (
db_session_with_containers.query(Document).where(Document.id == context["document"].id).first()
)
assert updated_document is not None
assert updated_document.indexing_status == "error"
assert "Datasource credential not found" in updated_document.error
assert updated_document.stopped_at is not None
mock_external_dependencies["indexing_runner"].run.assert_not_called()
def test_page_not_updated(self, db_session_with_containers, mock_external_dependencies):
"""Test that task exits early when notion page is unchanged."""
# Arrange
context = self._create_notion_sync_context(db_session_with_containers)
mock_external_dependencies["notion_extractor"].get_notion_last_edited_time.return_value = "2024-01-01T00:00:00Z"
# Act
document_indexing_sync_task(context["dataset"].id, context["document"].id)
# Assert
db_session_with_containers.expire_all()
updated_document = (
db_session_with_containers.query(Document).where(Document.id == context["document"].id).first()
)
remaining_segments = (
db_session_with_containers.query(DocumentSegment)
.where(DocumentSegment.document_id == context["document"].id)
.count()
)
assert updated_document is not None
assert updated_document.indexing_status == "completed"
assert updated_document.processing_started_at is None
assert remaining_segments == 3
mock_external_dependencies["index_processor"].clean.assert_not_called()
mock_external_dependencies["indexing_runner"].run.assert_not_called()
def test_successful_sync_when_page_updated(self, db_session_with_containers, mock_external_dependencies):
"""Test full successful sync flow with SQL state updates and side effects."""
# Arrange
context = self._create_notion_sync_context(db_session_with_containers)
# Act
document_indexing_sync_task(context["dataset"].id, context["document"].id)
# Assert
db_session_with_containers.expire_all()
updated_document = (
db_session_with_containers.query(Document).where(Document.id == context["document"].id).first()
)
remaining_segments = (
db_session_with_containers.query(DocumentSegment)
.where(DocumentSegment.document_id == context["document"].id)
.count()
)
assert updated_document is not None
assert updated_document.indexing_status == "parsing"
assert updated_document.processing_started_at is not None
assert updated_document.data_source_info_dict.get("last_edited_time") == "2024-01-02T00:00:00Z"
assert remaining_segments == 0
clean_call_args = mock_external_dependencies["index_processor"].clean.call_args
assert clean_call_args is not None
clean_args, clean_kwargs = clean_call_args
assert getattr(clean_args[0], "id", None) == context["dataset"].id
assert set(clean_args[1]) == set(context["node_ids"])
assert clean_kwargs.get("with_keywords") is True
assert clean_kwargs.get("delete_child_chunks") is True
run_call_args = mock_external_dependencies["indexing_runner"].run.call_args
assert run_call_args is not None
run_documents = run_call_args[0][0]
assert len(run_documents) == 1
assert getattr(run_documents[0], "id", None) == context["document"].id
def test_dataset_not_found_during_cleaning(self, db_session_with_containers, mock_external_dependencies):
"""Test that task still updates document and reindexes if dataset vanishes before clean."""
# Arrange
context = self._create_notion_sync_context(db_session_with_containers)
def _delete_dataset_before_clean() -> str:
db_session_with_containers.query(Dataset).where(Dataset.id == context["dataset"].id).delete()
db_session_with_containers.commit()
return "2024-01-02T00:00:00Z"
mock_external_dependencies[
"notion_extractor"
].get_notion_last_edited_time.side_effect = _delete_dataset_before_clean
# Act
document_indexing_sync_task(context["dataset"].id, context["document"].id)
# Assert
db_session_with_containers.expire_all()
updated_document = (
db_session_with_containers.query(Document).where(Document.id == context["document"].id).first()
)
assert updated_document is not None
assert updated_document.indexing_status == "parsing"
mock_external_dependencies["index_processor"].clean.assert_not_called()
mock_external_dependencies["indexing_runner"].run.assert_called_once()
def test_cleaning_error_continues_to_indexing(self, db_session_with_containers, mock_external_dependencies):
"""Test that indexing continues when index cleanup fails."""
# Arrange
context = self._create_notion_sync_context(db_session_with_containers)
mock_external_dependencies["index_processor"].clean.side_effect = Exception("Cleaning error")
# Act
document_indexing_sync_task(context["dataset"].id, context["document"].id)
# Assert
db_session_with_containers.expire_all()
updated_document = (
db_session_with_containers.query(Document).where(Document.id == context["document"].id).first()
)
remaining_segments = (
db_session_with_containers.query(DocumentSegment)
.where(DocumentSegment.document_id == context["document"].id)
.count()
)
assert updated_document is not None
assert updated_document.indexing_status == "parsing"
assert remaining_segments == 0
mock_external_dependencies["indexing_runner"].run.assert_called_once()
def test_indexing_runner_document_paused_error(self, db_session_with_containers, mock_external_dependencies):
"""Test that DocumentIsPausedError does not flip document into error state."""
# Arrange
context = self._create_notion_sync_context(db_session_with_containers)
mock_external_dependencies["indexing_runner"].run.side_effect = DocumentIsPausedError("Document paused")
# Act
document_indexing_sync_task(context["dataset"].id, context["document"].id)
# Assert
db_session_with_containers.expire_all()
updated_document = (
db_session_with_containers.query(Document).where(Document.id == context["document"].id).first()
)
assert updated_document is not None
assert updated_document.indexing_status == "parsing"
assert updated_document.error is None
def test_indexing_runner_general_error(self, db_session_with_containers, mock_external_dependencies):
"""Test that indexing errors are persisted to document state."""
# Arrange
context = self._create_notion_sync_context(db_session_with_containers)
mock_external_dependencies["indexing_runner"].run.side_effect = Exception("Indexing error")
# Act
document_indexing_sync_task(context["dataset"].id, context["document"].id)
# Assert
db_session_with_containers.expire_all()
updated_document = (
db_session_with_containers.query(Document).where(Document.id == context["document"].id).first()
)
assert updated_document is not None
assert updated_document.indexing_status == "error"
assert "Indexing error" in updated_document.error
assert updated_document.stopped_at is not None
def test_index_processor_clean_called_with_correct_params(
self,
db_session_with_containers,
mock_external_dependencies,
):
"""Test that clean is called with dataset instance and collected node ids."""
# Arrange
context = self._create_notion_sync_context(db_session_with_containers)
# Act
document_indexing_sync_task(context["dataset"].id, context["document"].id)
# Assert
clean_call_args = mock_external_dependencies["index_processor"].clean.call_args
assert clean_call_args is not None
clean_args, clean_kwargs = clean_call_args
assert getattr(clean_args[0], "id", None) == context["dataset"].id
assert set(clean_args[1]) == set(context["node_ids"])
assert clean_kwargs.get("with_keywords") is True
assert clean_kwargs.get("delete_child_chunks") is True
@@ -77,7 +77,7 @@ def _restx_mask_defaults(app: Flask):
def test_code_based_extension_get_returns_service_data(app: Flask, monkeypatch: pytest.MonkeyPatch):
service_result = {"entrypoint": "main:agent"}
service_result = [{"entrypoint": "main:agent"}]
service_mock = MagicMock(return_value=service_result)
monkeypatch.setattr(
"controllers.console.extension.CodeBasedExtensionService.get_code_based_extension",
@@ -107,7 +107,7 @@ class TestFilePreviewApi:
response = get_fn("file-id")
assert response.mimetype == "text/plain"
assert response.mimetype == "application/octet-stream"
assert response.headers["Content-Length"] == "100"
assert "Accept-Ranges" not in response.headers
mock_enforce.assert_called_once()
@@ -596,7 +596,8 @@ class TestWorkflowTaskStopApiPost:
assert result == {"result": "success"}
mock_queue_mgr.set_stop_flag_no_user_check.assert_called_once_with("task-1")
mock_graph_mgr.send_stop_command.assert_called_once_with("task-1")
mock_graph_mgr.assert_called_once()
mock_graph_mgr.return_value.send_stop_command.assert_called_once_with("task-1")
def test_stop_workflow_task_wrong_app_mode(self, app):
"""Test NotWorkflowAppError when app mode is not workflow."""
@@ -0,0 +1,135 @@
import types
from collections.abc import Generator
from core.datasource.datasource_manager import DatasourceManager
from core.datasource.entities.datasource_entities import DatasourceMessage
from core.workflow.entities.workflow_node_execution import WorkflowNodeExecutionStatus
from core.workflow.node_events import StreamChunkEvent, StreamCompletedEvent
def _gen_messages_text_only(text: str) -> Generator[DatasourceMessage, None, None]:
yield DatasourceMessage(
type=DatasourceMessage.MessageType.TEXT,
message=DatasourceMessage.TextMessage(text=text),
meta=None,
)
def test_get_icon_url_calls_runtime(mocker):
fake_runtime = mocker.Mock()
fake_runtime.get_icon_url.return_value = "https://icon"
mocker.patch.object(DatasourceManager, "get_datasource_runtime", return_value=fake_runtime)
url = DatasourceManager.get_icon_url(
provider_id="p/x",
tenant_id="t1",
datasource_name="ds",
datasource_type="online_document",
)
assert url == "https://icon"
DatasourceManager.get_datasource_runtime.assert_called_once()
def test_stream_online_results_yields_messages_online_document(mocker):
# stub runtime to yield a text message
def _doc_messages(**_):
yield from _gen_messages_text_only("hello")
fake_runtime = mocker.Mock()
fake_runtime.get_online_document_page_content.side_effect = _doc_messages
mocker.patch.object(DatasourceManager, "get_datasource_runtime", return_value=fake_runtime)
mocker.patch(
"core.datasource.datasource_manager.DatasourceProviderService.get_datasource_credentials",
return_value=None,
)
gen = DatasourceManager.stream_online_results(
user_id="u1",
datasource_name="ds",
datasource_type="online_document",
provider_id="p/x",
tenant_id="t1",
provider="prov",
plugin_id="plug",
credential_id="",
datasource_param=types.SimpleNamespace(workspace_id="w", page_id="pg", type="t"),
online_drive_request=None,
)
msgs = list(gen)
assert len(msgs) == 1
assert msgs[0].message.text == "hello"
def test_stream_node_events_emits_events_online_document(mocker):
# make manager's low-level stream produce TEXT only
mocker.patch.object(
DatasourceManager,
"stream_online_results",
return_value=_gen_messages_text_only("hello"),
)
events = list(
DatasourceManager.stream_node_events(
node_id="nodeA",
user_id="u1",
datasource_name="ds",
datasource_type="online_document",
provider_id="p/x",
tenant_id="t1",
provider="prov",
plugin_id="plug",
credential_id="",
parameters_for_log={"k": "v"},
datasource_info={"user_id": "u1"},
variable_pool=mocker.Mock(),
datasource_param=types.SimpleNamespace(workspace_id="w", page_id="pg", type="t"),
online_drive_request=None,
)
)
# should contain one StreamChunkEvent then a final chunk (empty) and a completed event
assert isinstance(events[0], StreamChunkEvent)
assert events[0].chunk == "hello"
assert isinstance(events[-1], StreamCompletedEvent)
assert events[-1].node_run_result.status == WorkflowNodeExecutionStatus.SUCCEEDED
def test_get_upload_file_by_id_builds_file(mocker):
# fake UploadFile row
fake_row = types.SimpleNamespace(
id="fid",
name="f",
extension="txt",
mime_type="text/plain",
size=1,
key="k",
source_url="http://x",
)
class _Q:
def __init__(self, row):
self._row = row
def where(self, *_args, **_kwargs):
return self
def first(self):
return self._row
class _S:
def __init__(self, row):
self._row = row
def __enter__(self):
return self
def __exit__(self, *exc):
return False
def query(self, *_):
return _Q(self._row)
mocker.patch("core.datasource.datasource_manager.session_factory.create_session", return_value=_S(fake_row))
f = DatasourceManager.get_upload_file_by_id(file_id="fid", tenant_id="t1")
assert f.related_id == "fid"
assert f.extension == ".txt"
@@ -4,13 +4,13 @@ import pytest
from configs import dify_config
from core.app.app_config.entities import ModelConfigEntity
from core.memory.token_buffer_memory import TokenBufferMemory
from core.model_runtime.entities.message_entities import (
AssistantPromptMessage,
ImagePromptMessageContent,
PromptMessageRole,
UserPromptMessage,
)
from core.model_runtime.token_buffer_memory import TokenBufferMemory
from core.prompt.advanced_prompt_transform import AdvancedPromptTransform
from core.prompt.entities.advanced_prompt_entities import ChatModelMessage, CompletionModelPromptTemplate, MemoryConfig
from core.prompt.utils.prompt_template_parser import PromptTemplateParser
@@ -4,7 +4,6 @@ from core.app.entities.app_invoke_entities import (
ModelConfigWithCredentialsEntity,
)
from core.entities.provider_configuration import ProviderModelBundle
from core.memory.token_buffer_memory import TokenBufferMemory
from core.model_runtime.entities.message_entities import (
AssistantPromptMessage,
SystemPromptMessage,
@@ -12,6 +11,7 @@ from core.model_runtime.entities.message_entities import (
UserPromptMessage,
)
from core.model_runtime.model_providers.__base.large_language_model import LargeLanguageModel
from core.model_runtime.token_buffer_memory import TokenBufferMemory
from core.prompt.agent_history_prompt_transform import AgentHistoryPromptTransform
from models.model import Conversation
@@ -1,8 +1,8 @@
from unittest.mock import MagicMock
from core.app.entities.app_invoke_entities import ModelConfigWithCredentialsEntity
from core.memory.token_buffer_memory import TokenBufferMemory
from core.model_runtime.entities.message_entities import AssistantPromptMessage, UserPromptMessage
from core.model_runtime.token_buffer_memory import TokenBufferMemory
from core.prompt.simple_prompt_transform import SimplePromptTransform
from models.model import AppMode, Conversation
@@ -82,7 +82,7 @@ class TestCacheEmbeddingDocuments:
Mock: Configured ModelInstance with text embedding capabilities
"""
model_instance = Mock()
model_instance.model = "text-embedding-ada-002"
model_instance.model_name = "text-embedding-ada-002"
model_instance.provider = "openai"
model_instance.credentials = {"api_key": "test-key"}
@@ -597,7 +597,7 @@ class TestCacheEmbeddingQuery:
def mock_model_instance(self):
"""Create a mock ModelInstance for testing."""
model_instance = Mock()
model_instance.model = "text-embedding-ada-002"
model_instance.model_name = "text-embedding-ada-002"
model_instance.provider = "openai"
model_instance.credentials = {"api_key": "test-key"}
return model_instance
@@ -830,7 +830,7 @@ class TestEmbeddingModelSwitching:
"""
# Arrange
model_instance_ada = Mock()
model_instance_ada.model = "text-embedding-ada-002"
model_instance_ada.model_name = "text-embedding-ada-002"
model_instance_ada.provider = "openai"
# Mock model type instance for ada
@@ -841,7 +841,7 @@ class TestEmbeddingModelSwitching:
model_type_instance_ada.get_model_schema.return_value = model_schema_ada
model_instance_3_small = Mock()
model_instance_3_small.model = "text-embedding-3-small"
model_instance_3_small.model_name = "text-embedding-3-small"
model_instance_3_small.provider = "openai"
# Mock model type instance for 3-small
@@ -914,11 +914,11 @@ class TestEmbeddingModelSwitching:
"""
# Arrange
model_instance_openai = Mock()
model_instance_openai.model = "text-embedding-ada-002"
model_instance_openai.model_name = "text-embedding-ada-002"
model_instance_openai.provider = "openai"
model_instance_cohere = Mock()
model_instance_cohere.model = "embed-english-v3.0"
model_instance_cohere.model_name = "embed-english-v3.0"
model_instance_cohere.provider = "cohere"
cache_openai = CacheEmbedding(model_instance_openai)
@@ -1001,7 +1001,7 @@ class TestEmbeddingDimensionValidation:
def mock_model_instance(self):
"""Create a mock ModelInstance for testing."""
model_instance = Mock()
model_instance.model = "text-embedding-ada-002"
model_instance.model_name = "text-embedding-ada-002"
model_instance.provider = "openai"
model_instance.credentials = {"api_key": "test-key"}
@@ -1123,7 +1123,7 @@ class TestEmbeddingDimensionValidation:
"""
# Arrange - OpenAI ada-002 (1536 dimensions)
model_instance_ada = Mock()
model_instance_ada.model = "text-embedding-ada-002"
model_instance_ada.model_name = "text-embedding-ada-002"
model_instance_ada.provider = "openai"
# Mock model type instance for ada
@@ -1156,7 +1156,7 @@ class TestEmbeddingDimensionValidation:
# Arrange - Cohere embed-english-v3.0 (1024 dimensions)
model_instance_cohere = Mock()
model_instance_cohere.model = "embed-english-v3.0"
model_instance_cohere.model_name = "embed-english-v3.0"
model_instance_cohere.provider = "cohere"
# Mock model type instance for cohere
@@ -1225,7 +1225,7 @@ class TestEmbeddingEdgeCases:
- MAX_CHUNKS: 10
"""
model_instance = Mock()
model_instance.model = "text-embedding-ada-002"
model_instance.model_name = "text-embedding-ada-002"
model_instance.provider = "openai"
model_type_instance = Mock()
@@ -1702,7 +1702,7 @@ class TestEmbeddingCachePerformance:
- MAX_CHUNKS: 10
"""
model_instance = Mock()
model_instance.model = "text-embedding-ada-002"
model_instance.model_name = "text-embedding-ada-002"
model_instance.provider = "openai"
model_type_instance = Mock()
@@ -34,7 +34,7 @@ def create_mock_model_instance():
mock_instance.provider_model_bundle.configuration = Mock()
mock_instance.provider_model_bundle.configuration.tenant_id = "test-tenant-id"
mock_instance.provider = "test-provider"
mock_instance.model = "test-model"
mock_instance.model_name = "test-model"
return mock_instance
@@ -65,7 +65,7 @@ class TestRerankModelRunner:
mock_instance.provider_model_bundle.configuration = Mock()
mock_instance.provider_model_bundle.configuration.tenant_id = "test-tenant-id"
mock_instance.provider = "test-provider"
mock_instance.model = "test-model"
mock_instance.model_name = "test-model"
return mock_instance
@pytest.fixture
@@ -199,11 +199,32 @@ def test_mock_config_builder():
def test_mock_factory_node_type_detection():
"""Test that MockNodeFactory correctly identifies nodes to mock."""
from core.app.entities.app_invoke_entities import InvokeFrom
from core.workflow.entities import GraphInitParams
from core.workflow.runtime import GraphRuntimeState, VariablePool
from models.enums import UserFrom
from .test_mock_factory import MockNodeFactory
graph_init_params = GraphInitParams(
tenant_id="test",
app_id="test",
workflow_id="test",
graph_config={},
user_id="test",
user_from=UserFrom.ACCOUNT,
invoke_from=InvokeFrom.SERVICE_API,
call_depth=0,
)
graph_runtime_state = GraphRuntimeState(
variable_pool=VariablePool(environment_variables=[], conversation_variables=[], user_inputs={}),
start_at=0,
total_tokens=0,
node_run_steps=0,
)
factory = MockNodeFactory(
graph_init_params=None, # Will be set by test
graph_runtime_state=None, # Will be set by test
graph_init_params=graph_init_params,
graph_runtime_state=graph_runtime_state,
mock_config=None,
)
@@ -288,7 +309,11 @@ def test_workflow_without_auto_mock():
def test_register_custom_mock_node():
"""Test registering a custom mock implementation for a node type."""
from core.app.entities.app_invoke_entities import InvokeFrom
from core.workflow.entities import GraphInitParams
from core.workflow.nodes.template_transform import TemplateTransformNode
from core.workflow.runtime import GraphRuntimeState, VariablePool
from models.enums import UserFrom
from .test_mock_factory import MockNodeFactory
@@ -298,9 +323,25 @@ def test_register_custom_mock_node():
# Custom mock implementation
pass
graph_init_params = GraphInitParams(
tenant_id="test",
app_id="test",
workflow_id="test",
graph_config={},
user_id="test",
user_from=UserFrom.ACCOUNT,
invoke_from=InvokeFrom.SERVICE_API,
call_depth=0,
)
graph_runtime_state = GraphRuntimeState(
variable_pool=VariablePool(environment_variables=[], conversation_variables=[], user_inputs={}),
start_at=0,
total_tokens=0,
node_run_steps=0,
)
factory = MockNodeFactory(
graph_init_params=None,
graph_runtime_state=None,
graph_init_params=graph_init_params,
graph_runtime_state=graph_runtime_state,
mock_config=None,
)
@@ -1,9 +1,9 @@
import datetime
import time
from collections.abc import Iterable
from unittest import mock
from unittest.mock import MagicMock
from core.model_runtime.entities.llm_entities import LLMMode
from core.model_runtime.entities.message_entities import PromptMessageRole
from core.workflow.entities import GraphInitParams
from core.workflow.graph import Graph
@@ -82,7 +82,7 @@ def _build_branching_graph(
def _create_llm_node(node_id: str, title: str, prompt_text: str) -> MockLLMNode:
llm_data = LLMNodeData(
title=title,
model=ModelConfig(provider="openai", name="gpt-3.5-turbo", mode=LLMMode.CHAT, completion_params={}),
model=ModelConfig(provider="openai", name="gpt-3.5-turbo", mode="chat", completion_params={}),
prompt_template=[
LLMNodeChatModelMessage(
text=prompt_text,
@@ -101,6 +101,8 @@ def _build_branching_graph(
graph_init_params=graph_init_params,
graph_runtime_state=graph_runtime_state,
mock_config=mock_config,
credentials_provider=mock.Mock(),
model_factory=mock.Mock(),
)
return llm_node
@@ -1,8 +1,8 @@
import datetime
import time
from unittest import mock
from unittest.mock import MagicMock
from core.model_runtime.entities.llm_entities import LLMMode
from core.model_runtime.entities.message_entities import PromptMessageRole
from core.workflow.entities import GraphInitParams
from core.workflow.graph import Graph
@@ -78,7 +78,7 @@ def _build_llm_human_llm_graph(
def _create_llm_node(node_id: str, title: str, prompt_text: str) -> MockLLMNode:
llm_data = LLMNodeData(
title=title,
model=ModelConfig(provider="openai", name="gpt-3.5-turbo", mode=LLMMode.CHAT, completion_params={}),
model=ModelConfig(provider="openai", name="gpt-3.5-turbo", mode="chat", completion_params={}),
prompt_template=[
LLMNodeChatModelMessage(
text=prompt_text,
@@ -97,6 +97,8 @@ def _build_llm_human_llm_graph(
graph_init_params=graph_init_params,
graph_runtime_state=graph_runtime_state,
mock_config=mock_config,
credentials_provider=mock.Mock(),
model_factory=mock.Mock(),
)
return llm_node
@@ -1,4 +1,5 @@
import time
from unittest import mock
from core.model_runtime.entities.llm_entities import LLMMode
from core.model_runtime.entities.message_entities import PromptMessageRole
@@ -85,6 +86,8 @@ def _build_if_else_graph(branch_value: str, mock_config: MockConfig) -> tuple[Gr
graph_init_params=graph_init_params,
graph_runtime_state=graph_runtime_state,
mock_config=mock_config,
credentials_provider=mock.Mock(),
model_factory=mock.Mock(),
)
return llm_node
@@ -5,6 +5,7 @@ This module provides a MockNodeFactory that automatically detects and mocks node
requiring external services (LLM, Agent, Tool, Knowledge Retrieval, HTTP Request).
"""
from collections.abc import Mapping
from typing import TYPE_CHECKING, Any
from core.app.workflow.node_factory import DifyNodeFactory
@@ -74,7 +75,7 @@ class MockNodeFactory(DifyNodeFactory):
NodeType.CODE: MockCodeNode,
}
def create_node(self, node_config: dict[str, Any]) -> Node:
def create_node(self, node_config: Mapping[str, Any]) -> Node:
"""
Create a node instance, using mock implementations for third-party service nodes.
@@ -123,6 +124,16 @@ class MockNodeFactory(DifyNodeFactory):
mock_config=self.mock_config,
http_request_config=self._http_request_config,
)
elif node_type in {NodeType.LLM, NodeType.QUESTION_CLASSIFIER, NodeType.PARAMETER_EXTRACTOR}:
mock_instance = mock_class(
id=node_id,
config=node_config,
graph_init_params=self.graph_init_params,
graph_runtime_state=self.graph_runtime_state,
mock_config=self.mock_config,
credentials_provider=self._llm_credentials_provider,
model_factory=self._llm_model_factory,
)
else:
mock_instance = mock_class(
id=node_id,
@@ -16,9 +16,33 @@ from tests.unit_tests.core.workflow.graph_engine.test_mock_factory import MockNo
def test_mock_factory_registers_iteration_node():
"""Test that MockNodeFactory has iteration node registered."""
from core.app.entities.app_invoke_entities import InvokeFrom
from core.workflow.entities import GraphInitParams
from core.workflow.runtime import GraphRuntimeState, VariablePool
from models.enums import UserFrom
# Create a MockNodeFactory instance
factory = MockNodeFactory(graph_init_params=None, graph_runtime_state=None, mock_config=None)
graph_init_params = GraphInitParams(
tenant_id="test",
app_id="test",
workflow_id="test",
graph_config={"nodes": [], "edges": []},
user_id="test",
user_from=UserFrom.ACCOUNT,
invoke_from=InvokeFrom.SERVICE_API,
call_depth=0,
)
graph_runtime_state = GraphRuntimeState(
variable_pool=VariablePool(environment_variables=[], conversation_variables=[], user_inputs={}),
start_at=0,
total_tokens=0,
node_run_steps=0,
)
factory = MockNodeFactory(
graph_init_params=graph_init_params,
graph_runtime_state=graph_runtime_state,
mock_config=None,
)
# Check that iteration node is registered
assert NodeType.ITERATION in factory._mock_node_types
@@ -8,6 +8,7 @@ allowing tests to run without external dependencies.
import time
from collections.abc import Generator, Mapping
from typing import TYPE_CHECKING, Any, Optional
from unittest.mock import MagicMock
from core.model_runtime.entities.llm_entities import LLMUsage
from core.workflow.enums import WorkflowNodeExecutionMetadataKey, WorkflowNodeExecutionStatus
@@ -18,6 +19,7 @@ from core.workflow.nodes.document_extractor import DocumentExtractorNode
from core.workflow.nodes.http_request import HttpRequestNode
from core.workflow.nodes.knowledge_retrieval import KnowledgeRetrievalNode
from core.workflow.nodes.llm import LLMNode
from core.workflow.nodes.llm.protocols import CredentialsProvider, ModelFactory
from core.workflow.nodes.parameter_extractor import ParameterExtractorNode
from core.workflow.nodes.question_classifier import QuestionClassifierNode
from core.workflow.nodes.template_transform import TemplateTransformNode
@@ -42,6 +44,10 @@ class MockNodeMixin:
mock_config: Optional["MockConfig"] = None,
**kwargs: Any,
):
if isinstance(self, (LLMNode, QuestionClassifierNode)):
kwargs.setdefault("credentials_provider", MagicMock(spec=CredentialsProvider))
kwargs.setdefault("model_factory", MagicMock(spec=ModelFactory))
super().__init__(
id=id,
config=config,
@@ -24,6 +24,16 @@ DEFAULT_CODE_LIMITS = CodeNodeLimits(
)
class _NoopCodeExecutor:
def execute(self, *, language: object, code: str, inputs: dict[str, object]) -> dict[str, object]:
_ = (language, code, inputs)
return {}
def is_execution_error(self, error: Exception) -> bool:
_ = error
return False
class TestMockTemplateTransformNode:
"""Test cases for MockTemplateTransformNode."""
@@ -319,6 +329,7 @@ class TestMockCodeNode:
graph_init_params=graph_init_params,
graph_runtime_state=graph_runtime_state,
mock_config=mock_config,
code_executor=_NoopCodeExecutor(),
code_limits=DEFAULT_CODE_LIMITS,
)
@@ -384,6 +395,7 @@ class TestMockCodeNode:
graph_init_params=graph_init_params,
graph_runtime_state=graph_runtime_state,
mock_config=mock_config,
code_executor=_NoopCodeExecutor(),
code_limits=DEFAULT_CODE_LIMITS,
)
@@ -453,6 +465,7 @@ class TestMockCodeNode:
graph_init_params=graph_init_params,
graph_runtime_state=graph_runtime_state,
mock_config=mock_config,
code_executor=_NoopCodeExecutor(),
code_limits=DEFAULT_CODE_LIMITS,
)
@@ -101,11 +101,32 @@ def test_node_mock_config():
def test_mock_factory_detection():
"""Test MockNodeFactory node type detection."""
from core.app.entities.app_invoke_entities import InvokeFrom
from core.workflow.entities import GraphInitParams
from core.workflow.runtime import GraphRuntimeState, VariablePool
from models.enums import UserFrom
print("Testing MockNodeFactory detection...")
graph_init_params = GraphInitParams(
tenant_id="test",
app_id="test",
workflow_id="test",
graph_config={},
user_id="test",
user_from=UserFrom.ACCOUNT,
invoke_from=InvokeFrom.SERVICE_API,
call_depth=0,
)
graph_runtime_state = GraphRuntimeState(
variable_pool=VariablePool(environment_variables=[], conversation_variables=[], user_inputs={}),
start_at=0,
total_tokens=0,
node_run_steps=0,
)
factory = MockNodeFactory(
graph_init_params=None,
graph_runtime_state=None,
graph_init_params=graph_init_params,
graph_runtime_state=graph_runtime_state,
mock_config=None,
)
@@ -133,11 +154,32 @@ def test_mock_factory_detection():
def test_mock_factory_registration():
"""Test registering and unregistering mock node types."""
from core.app.entities.app_invoke_entities import InvokeFrom
from core.workflow.entities import GraphInitParams
from core.workflow.runtime import GraphRuntimeState, VariablePool
from models.enums import UserFrom
print("Testing MockNodeFactory registration...")
graph_init_params = GraphInitParams(
tenant_id="test",
app_id="test",
workflow_id="test",
graph_config={},
user_id="test",
user_from=UserFrom.ACCOUNT,
invoke_from=InvokeFrom.SERVICE_API,
call_depth=0,
)
graph_runtime_state = GraphRuntimeState(
variable_pool=VariablePool(environment_variables=[], conversation_variables=[], user_inputs={}),
start_at=0,
total_tokens=0,
node_run_steps=0,
)
factory = MockNodeFactory(
graph_init_params=None,
graph_runtime_state=None,
graph_init_params=graph_init_params,
graph_runtime_state=graph_runtime_state,
mock_config=None,
)
@@ -32,25 +32,26 @@ class TestRedisStopIntegration:
mock_redis.pipeline.return_value.__enter__ = Mock(return_value=mock_pipeline)
mock_redis.pipeline.return_value.__exit__ = Mock(return_value=None)
with patch("core.workflow.graph_engine.manager.redis_client", mock_redis):
# Execute
GraphEngineManager.send_stop_command(task_id, reason="Test stop")
manager = GraphEngineManager(mock_redis)
# Verify
mock_redis.pipeline.assert_called_once()
# Execute
manager.send_stop_command(task_id, reason="Test stop")
# Check that rpush was called with correct arguments
calls = mock_pipeline.rpush.call_args_list
assert len(calls) == 1
# Verify
mock_redis.pipeline.assert_called_once()
# Verify the channel key
assert calls[0][0][0] == expected_channel_key
# Check that rpush was called with correct arguments
calls = mock_pipeline.rpush.call_args_list
assert len(calls) == 1
# Verify the command data
command_json = calls[0][0][1]
command_data = json.loads(command_json)
assert command_data["command_type"] == CommandType.ABORT
assert command_data["reason"] == "Test stop"
# Verify the channel key
assert calls[0][0][0] == expected_channel_key
# Verify the command data
command_json = calls[0][0][1]
command_data = json.loads(command_json)
assert command_data["command_type"] == CommandType.ABORT
assert command_data["reason"] == "Test stop"
def test_graph_engine_manager_sends_pause_command(self):
"""Test that GraphEngineManager correctly sends pause command through Redis."""
@@ -62,18 +63,18 @@ class TestRedisStopIntegration:
mock_redis.pipeline.return_value.__enter__ = Mock(return_value=mock_pipeline)
mock_redis.pipeline.return_value.__exit__ = Mock(return_value=None)
with patch("core.workflow.graph_engine.manager.redis_client", mock_redis):
GraphEngineManager.send_pause_command(task_id, reason="Awaiting resources")
manager = GraphEngineManager(mock_redis)
manager.send_pause_command(task_id, reason="Awaiting resources")
mock_redis.pipeline.assert_called_once()
calls = mock_pipeline.rpush.call_args_list
assert len(calls) == 1
assert calls[0][0][0] == expected_channel_key
mock_redis.pipeline.assert_called_once()
calls = mock_pipeline.rpush.call_args_list
assert len(calls) == 1
assert calls[0][0][0] == expected_channel_key
command_json = calls[0][0][1]
command_data = json.loads(command_json)
assert command_data["command_type"] == CommandType.PAUSE.value
assert command_data["reason"] == "Awaiting resources"
command_json = calls[0][0][1]
command_data = json.loads(command_json)
assert command_data["command_type"] == CommandType.PAUSE.value
assert command_data["reason"] == "Awaiting resources"
def test_graph_engine_manager_handles_redis_failure_gracefully(self):
"""Test that GraphEngineManager handles Redis failures without raising exceptions."""
@@ -82,13 +83,13 @@ class TestRedisStopIntegration:
# Mock redis client to raise exception
mock_redis = MagicMock()
mock_redis.pipeline.side_effect = redis.ConnectionError("Redis connection failed")
manager = GraphEngineManager(mock_redis)
with patch("core.workflow.graph_engine.manager.redis_client", mock_redis):
# Should not raise exception
try:
GraphEngineManager.send_stop_command(task_id)
except Exception as e:
pytest.fail(f"GraphEngineManager.send_stop_command raised {e} unexpectedly")
# Should not raise exception
try:
manager.send_stop_command(task_id)
except Exception as e:
pytest.fail(f"GraphEngineManager.send_stop_command raised {e} unexpectedly")
def test_app_queue_manager_no_user_check(self):
"""Test that AppQueueManager.set_stop_flag_no_user_check works without user validation."""
@@ -251,13 +252,10 @@ class TestRedisStopIntegration:
mock_redis.pipeline.return_value.__enter__ = Mock(return_value=mock_pipeline)
mock_redis.pipeline.return_value.__exit__ = Mock(return_value=None)
with (
patch("core.app.apps.base_app_queue_manager.redis_client", mock_redis),
patch("core.workflow.graph_engine.manager.redis_client", mock_redis),
):
with patch("core.app.apps.base_app_queue_manager.redis_client", mock_redis):
# Execute both stop mechanisms
AppQueueManager.set_stop_flag_no_user_check(task_id)
GraphEngineManager.send_stop_command(task_id)
GraphEngineManager(mock_redis).send_stop_command(task_id)
# Verify legacy stop flag was set
expected_stop_flag_key = f"generate_task_stopped:{task_id}"
@@ -0,0 +1,93 @@
from core.workflow.entities.workflow_node_execution import WorkflowNodeExecutionStatus
from core.workflow.node_events import NodeRunResult, StreamChunkEvent, StreamCompletedEvent
from core.workflow.nodes.datasource.datasource_node import DatasourceNode
class _VarSeg:
def __init__(self, v):
self.value = v
class _VarPool:
def __init__(self, mapping):
self._m = mapping
def get(self, selector):
d = self._m
for k in selector:
d = d[k]
return _VarSeg(d)
def add(self, *_args, **_kwargs):
pass
class _GraphState:
def __init__(self, var_pool):
self.variable_pool = var_pool
class _GraphParams:
tenant_id = "t1"
app_id = "app-1"
workflow_id = "wf-1"
graph_config = {}
user_id = "u1"
user_from = "account"
invoke_from = "debugger"
call_depth = 0
def test_datasource_node_delegates_to_manager_stream(mocker):
# prepare sys variables
sys_vars = {
"sys": {
"datasource_type": "online_document",
"datasource_info": {
"workspace_id": "w",
"page": {"page_id": "pg", "type": "t"},
"credential_id": "",
},
}
}
var_pool = _VarPool(sys_vars)
gs = _GraphState(var_pool)
gp = _GraphParams()
# stub manager class
class _Mgr:
@classmethod
def get_icon_url(cls, **_):
return "icon"
@classmethod
def stream_node_events(cls, **_):
yield StreamChunkEvent(selector=["n", "text"], chunk="hi", is_final=False)
yield StreamCompletedEvent(node_run_result=NodeRunResult(status=WorkflowNodeExecutionStatus.SUCCEEDED))
@classmethod
def get_upload_file_by_id(cls, **_):
raise AssertionError("not called")
node = DatasourceNode(
id="n",
config={
"id": "n",
"data": {
"type": "datasource",
"version": "1",
"title": "Datasource",
"provider_type": "plugin",
"provider_name": "p",
"plugin_id": "plug",
"datasource_name": "ds",
},
},
graph_init_params=gp,
graph_runtime_state=gs,
datasource_manager=_Mgr,
)
evts = list(node._run())
assert isinstance(evts[0], StreamChunkEvent)
assert isinstance(evts[-1], StreamCompletedEvent)
@@ -6,6 +6,7 @@ from unittest import mock
import pytest
from core.app.entities.app_invoke_entities import InvokeFrom, ModelConfigWithCredentialsEntity
from core.app.llm.model_access import DifyCredentialsProvider, DifyModelFactory, fetch_model_config
from core.entities.provider_configuration import ProviderConfiguration, ProviderModelBundle
from core.entities.provider_entities import CustomConfiguration, SystemConfiguration
from core.model_runtime.entities.common_entities import I18nObject
@@ -32,6 +33,7 @@ from core.workflow.nodes.llm.entities import (
)
from core.workflow.nodes.llm.file_saver import LLMFileSaver
from core.workflow.nodes.llm.node import LLMNode
from core.workflow.nodes.llm.protocols import CredentialsProvider, ModelFactory
from core.workflow.runtime import GraphRuntimeState, VariablePool
from core.workflow.system_variable import SystemVariable
from models.enums import UserFrom
@@ -100,6 +102,8 @@ def llm_node(
llm_node_data: LLMNodeData, graph_init_params: GraphInitParams, graph_runtime_state: GraphRuntimeState
) -> LLMNode:
mock_file_saver = mock.MagicMock(spec=LLMFileSaver)
mock_credentials_provider = mock.MagicMock(spec=CredentialsProvider)
mock_model_factory = mock.MagicMock(spec=ModelFactory)
node_config = {
"id": "1",
"data": llm_node_data.model_dump(),
@@ -109,13 +113,29 @@ def llm_node(
config=node_config,
graph_init_params=graph_init_params,
graph_runtime_state=graph_runtime_state,
credentials_provider=mock_credentials_provider,
model_factory=mock_model_factory,
llm_file_saver=mock_file_saver,
)
return node
@pytest.fixture
def model_config():
def model_config(monkeypatch):
from tests.integration_tests.model_runtime.__mock.plugin_model import MockModelClass
def mock_plugin_model_providers(_self):
providers = MockModelClass().fetch_model_providers("test")
for provider in providers:
provider.declaration.provider = f"{provider.plugin_id}/{provider.declaration.provider}"
return providers
monkeypatch.setattr(
ModelProviderFactory,
"get_plugin_model_providers",
mock_plugin_model_providers,
)
# Create actual provider and model type instances
model_provider_factory = ModelProviderFactory(tenant_id="test")
provider_instance = model_provider_factory.get_plugin_model_provider("openai")
@@ -125,7 +145,7 @@ def model_config():
provider_model_bundle = ProviderModelBundle(
configuration=ProviderConfiguration(
tenant_id="1",
provider=provider_instance,
provider=provider_instance.declaration,
preferred_provider_type=ProviderType.CUSTOM,
using_provider_type=ProviderType.CUSTOM,
system_configuration=SystemConfiguration(enabled=False),
@@ -153,6 +173,89 @@ def model_config():
)
def test_fetch_model_config_uses_ports(model_config: ModelConfigWithCredentialsEntity):
mock_credentials_provider = mock.MagicMock(spec=CredentialsProvider)
mock_model_factory = mock.MagicMock(spec=ModelFactory)
provider_model_bundle = model_config.provider_model_bundle
model_type_instance = provider_model_bundle.model_type_instance
provider_model = mock.MagicMock()
model_instance = mock.MagicMock(
model_type_instance=model_type_instance,
provider_model_bundle=provider_model_bundle,
)
mock_credentials_provider.fetch.return_value = {"api_key": "test"}
mock_model_factory.init_model_instance.return_value = model_instance
with (
mock.patch.object(
provider_model_bundle.configuration.__class__,
"get_provider_model",
return_value=provider_model,
),
mock.patch.object(
model_type_instance.__class__,
"get_model_schema",
return_value=model_config.model_schema,
),
):
fetch_model_config(
node_data_model=ModelConfig(provider="openai", name="gpt-3.5-turbo", mode="chat", completion_params={}),
credentials_provider=mock_credentials_provider,
model_factory=mock_model_factory,
)
mock_credentials_provider.fetch.assert_called_once_with("openai", "gpt-3.5-turbo")
mock_model_factory.init_model_instance.assert_called_once_with("openai", "gpt-3.5-turbo")
provider_model.raise_for_status.assert_called_once()
def test_dify_model_access_adapters_call_managers():
mock_provider_manager = mock.MagicMock()
mock_model_manager = mock.MagicMock()
mock_configurations = mock.MagicMock()
mock_provider_configuration = mock.MagicMock()
mock_provider_model = mock.MagicMock()
mock_configurations.get.return_value = mock_provider_configuration
mock_provider_configuration.get_provider_model.return_value = mock_provider_model
mock_provider_configuration.get_current_credentials.return_value = {"api_key": "test"}
credentials_provider = DifyCredentialsProvider(
tenant_id="tenant",
provider_manager=mock_provider_manager,
)
model_factory = DifyModelFactory(
tenant_id="tenant",
model_manager=mock_model_manager,
)
mock_provider_manager.get_configurations.return_value = mock_configurations
credentials_provider.fetch("openai", "gpt-3.5-turbo")
model_factory.init_model_instance("openai", "gpt-3.5-turbo")
mock_provider_manager.get_configurations.assert_called_once_with("tenant")
mock_configurations.get.assert_called_once_with("openai")
mock_provider_configuration.get_provider_model.assert_called_once_with(
model_type=ModelType.LLM,
model="gpt-3.5-turbo",
)
mock_provider_configuration.get_current_credentials.assert_called_once_with(
model_type=ModelType.LLM,
model="gpt-3.5-turbo",
)
mock_provider_model.raise_for_status.assert_called_once()
mock_model_manager.get_model_instance.assert_called_once_with(
tenant_id="tenant",
provider="openai",
model_type=ModelType.LLM,
model="gpt-3.5-turbo",
)
def test_fetch_files_with_file_segment():
file = File(
id="1",
@@ -485,6 +588,8 @@ def test_handle_list_messages_basic(llm_node):
@pytest.fixture
def llm_node_for_multimodal(llm_node_data, graph_init_params, graph_runtime_state) -> tuple[LLMNode, LLMFileSaver]:
mock_file_saver: LLMFileSaver = mock.MagicMock(spec=LLMFileSaver)
mock_credentials_provider = mock.MagicMock(spec=CredentialsProvider)
mock_model_factory = mock.MagicMock(spec=ModelFactory)
node_config = {
"id": "1",
"data": llm_node_data.model_dump(),
@@ -494,6 +599,8 @@ def llm_node_for_multimodal(llm_node_data, graph_init_params, graph_runtime_stat
config=node_config,
graph_init_params=graph_init_params,
graph_runtime_state=graph_runtime_state,
credentials_provider=mock_credentials_provider,
model_factory=mock_model_factory,
llm_file_saver=mock_file_saver,
)
return node, mock_file_saver
@@ -642,8 +642,16 @@ class TestDatasetServiceUpdateRagPipelineDatasetSettings:
# Mock embedding model
mock_embedding_model = Mock()
mock_embedding_model.model = "text-embedding-ada-002"
mock_embedding_model.model_name = "text-embedding-ada-002"
mock_embedding_model.provider = "openai"
mock_embedding_model.credentials = {}
mock_model_schema = Mock()
mock_model_schema.features = []
mock_text_embedding_model = Mock()
mock_text_embedding_model.get_model_schema.return_value = mock_model_schema
mock_embedding_model.model_type_instance = mock_text_embedding_model
mock_model_instance = Mock()
mock_model_instance.get_model_instance.return_value = mock_embedding_model
File diff suppressed because it is too large Load Diff
@@ -44,9 +44,10 @@ class TestAppTaskService:
# Assert
mock_app_queue_manager.set_stop_flag.assert_called_once_with(task_id, invoke_from, user_id)
if should_call_graph_engine:
mock_graph_engine_manager.send_stop_command.assert_called_once_with(task_id)
mock_graph_engine_manager.assert_called_once()
mock_graph_engine_manager.return_value.send_stop_command.assert_called_once_with(task_id)
else:
mock_graph_engine_manager.send_stop_command.assert_not_called()
mock_graph_engine_manager.assert_not_called()
@pytest.mark.parametrize(
"invoke_from",
@@ -76,7 +77,8 @@ class TestAppTaskService:
# Assert
mock_app_queue_manager.set_stop_flag.assert_called_once_with(task_id, invoke_from, user_id)
mock_graph_engine_manager.send_stop_command.assert_called_once_with(task_id)
mock_graph_engine_manager.assert_called_once()
mock_graph_engine_manager.return_value.send_stop_command.assert_called_once_with(task_id)
@patch("services.app_task_service.GraphEngineManager")
@patch("services.app_task_service.AppQueueManager")
@@ -96,7 +98,7 @@ class TestAppTaskService:
app_mode = AppMode.ADVANCED_CHAT
# Simulate GraphEngine failure
mock_graph_engine_manager.send_stop_command.side_effect = Exception("GraphEngine error")
mock_graph_engine_manager.return_value.send_stop_command.side_effect = Exception("GraphEngine error")
# Act & Assert - should raise the exception since it's not caught
with pytest.raises(Exception, match="GraphEngine error"):

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