Merge branch 'origin/main' into feat/evaluation

This commit is contained in:
FFXN 2026-03-26 16:46:10 +08:00
commit ef3973f188
1766 changed files with 101202 additions and 34610 deletions

13
.gemini/config.yaml Normal file
View File

@ -0,0 +1,13 @@
have_fun: false
memory_config:
disabled: false
code_review:
disable: true
comment_severity_threshold: MEDIUM
max_review_comments: -1
pull_request_opened:
help: false
summary: false
code_review: false
include_drafts: false
ignore_patterns: []

2
.github/CODEOWNERS vendored
View File

@ -36,7 +36,7 @@
/api/core/workflow/graph/ @laipz8200 @QuantumGhost /api/core/workflow/graph/ @laipz8200 @QuantumGhost
/api/core/workflow/graph_events/ @laipz8200 @QuantumGhost /api/core/workflow/graph_events/ @laipz8200 @QuantumGhost
/api/core/workflow/node_events/ @laipz8200 @QuantumGhost /api/core/workflow/node_events/ @laipz8200 @QuantumGhost
/api/dify_graph/model_runtime/ @laipz8200 @QuantumGhost /api/graphon/model_runtime/ @laipz8200 @WH-2099
# Backend - Workflow - Nodes (Agent, Iteration, Loop, LLM) # Backend - Workflow - Nodes (Agent, Iteration, Loop, LLM)
/api/core/workflow/nodes/agent/ @Nov1c444 /api/core/workflow/nodes/agent/ @Nov1c444

View File

@ -4,10 +4,9 @@ runs:
using: composite using: composite
steps: steps:
- name: Setup Vite+ - name: Setup Vite+
uses: voidzero-dev/setup-vp@4a524139920f87f9f7080d3b8545acac019e1852 # v1.0.0 uses: voidzero-dev/setup-vp@20553a7a7429c429a74894104a2835d7fed28a72 # v1.3.0
with: with:
node-version-file: web/.nvmrc working-directory: web
node-version-file: .nvmrc
cache: true cache: true
cache-dependency-path: web/pnpm-lock.yaml run-install: true
run-install: |
cwd: ./web

View File

@ -84,20 +84,20 @@ jobs:
if: steps.changed-files.outputs.any_changed == 'true' if: steps.changed-files.outputs.any_changed == 'true'
uses: ./.github/actions/setup-web uses: ./.github/actions/setup-web
- name: Restore ESLint cache
if: steps.changed-files.outputs.any_changed == 'true'
id: eslint-cache-restore
uses: actions/cache/restore@668228422ae6a00e4ad889ee87cd7109ec5666a7 # v5.0.4
with:
path: web/.eslintcache
key: ${{ runner.os }}-web-eslint-${{ hashFiles('web/package.json', 'web/pnpm-lock.yaml', 'web/eslint.config.mjs', 'web/eslint.constants.mjs', 'web/plugins/eslint/**') }}-${{ github.sha }}
restore-keys: |
${{ runner.os }}-web-eslint-${{ hashFiles('web/package.json', 'web/pnpm-lock.yaml', 'web/eslint.config.mjs', 'web/eslint.constants.mjs', 'web/plugins/eslint/**') }}-
- name: Web style check - name: Web style check
if: steps.changed-files.outputs.any_changed == 'true' if: steps.changed-files.outputs.any_changed == 'true'
working-directory: ./web working-directory: ./web
run: | run: vp run lint:ci
vp run lint:ci
# pnpm run lint:report
# continue-on-error: true
# - name: Annotate Code
# if: steps.changed-files.outputs.any_changed == 'true' && github.event_name == 'pull_request'
# uses: DerLev/eslint-annotations@51347b3a0abfb503fc8734d5ae31c4b151297fae
# with:
# eslint-report: web/eslint_report.json
# github-token: ${{ secrets.GITHUB_TOKEN }}
- name: Web tsslint - name: Web tsslint
if: steps.changed-files.outputs.any_changed == 'true' if: steps.changed-files.outputs.any_changed == 'true'
@ -114,6 +114,13 @@ jobs:
working-directory: ./web working-directory: ./web
run: vp run knip run: vp run knip
- name: Save ESLint cache
if: steps.changed-files.outputs.any_changed == 'true' && success() && steps.eslint-cache-restore.outputs.cache-hit != 'true'
uses: actions/cache/save@668228422ae6a00e4ad889ee87cd7109ec5666a7 # v5.0.4
with:
path: web/.eslintcache
key: ${{ steps.eslint-cache-restore.outputs.cache-primary-key }}
superlinter: superlinter:
name: SuperLinter name: SuperLinter
runs-on: ubuntu-latest runs-on: ubuntu-latest

View File

@ -120,7 +120,7 @@ jobs:
- name: Run Claude Code for Translation Sync - name: Run Claude Code for Translation Sync
if: steps.detect_changes.outputs.CHANGED_FILES != '' if: steps.detect_changes.outputs.CHANGED_FILES != ''
uses: anthropics/claude-code-action@6062f3709600659be5e47fcddf2cf76993c235c2 # v1.0.76 uses: anthropics/claude-code-action@ff9acae5886d41a99ed4ec14b7dc147d55834722 # v1.0.77
with: with:
anthropic_api_key: ${{ secrets.ANTHROPIC_API_KEY }} anthropic_api_key: ${{ secrets.ANTHROPIC_API_KEY }}
github_token: ${{ secrets.GITHUB_TOKEN }} github_token: ${{ secrets.GITHUB_TOKEN }}

View File

@ -353,6 +353,9 @@ BAIDU_VECTOR_DB_SHARD=1
BAIDU_VECTOR_DB_REPLICAS=3 BAIDU_VECTOR_DB_REPLICAS=3
BAIDU_VECTOR_DB_INVERTED_INDEX_ANALYZER=DEFAULT_ANALYZER BAIDU_VECTOR_DB_INVERTED_INDEX_ANALYZER=DEFAULT_ANALYZER
BAIDU_VECTOR_DB_INVERTED_INDEX_PARSER_MODE=COARSE_MODE BAIDU_VECTOR_DB_INVERTED_INDEX_PARSER_MODE=COARSE_MODE
BAIDU_VECTOR_DB_AUTO_BUILD_ROW_COUNT_INCREMENT=500
BAIDU_VECTOR_DB_AUTO_BUILD_ROW_COUNT_INCREMENT_RATIO=0.05
BAIDU_VECTOR_DB_REBUILD_INDEX_TIMEOUT_IN_SECONDS=300
# Upstash configuration # Upstash configuration
UPSTASH_VECTOR_URL=your-server-url UPSTASH_VECTOR_URL=your-server-url

View File

@ -1,10 +1,14 @@
[importlinter] [importlinter]
root_packages = root_packages =
core core
dify_graph constants
context
graphon
configs configs
controllers controllers
extensions extensions
factories
libs
models models
tasks tasks
services services
@ -22,40 +26,30 @@ layers =
runtime runtime
entities entities
containers = containers =
dify_graph graphon
ignore_imports = ignore_imports =
dify_graph.nodes.base.node -> dify_graph.graph_events graphon.nodes.base.node -> graphon.graph_events
dify_graph.nodes.iteration.iteration_node -> dify_graph.graph_events graphon.nodes.iteration.iteration_node -> graphon.graph_events
dify_graph.nodes.loop.loop_node -> dify_graph.graph_events graphon.nodes.loop.loop_node -> graphon.graph_events
dify_graph.nodes.iteration.iteration_node -> dify_graph.graph_engine graphon.nodes.iteration.iteration_node -> graphon.graph_engine
dify_graph.nodes.loop.loop_node -> dify_graph.graph_engine graphon.nodes.loop.loop_node -> graphon.graph_engine
# TODO(QuantumGhost): fix the import violation later # TODO(QuantumGhost): fix the import violation later
dify_graph.entities.pause_reason -> dify_graph.nodes.human_input.entities graphon.entities.pause_reason -> graphon.nodes.human_input.entities
[importlinter:contract:workflow-infrastructure-dependencies]
name = Workflow Infrastructure Dependencies
type = forbidden
source_modules =
dify_graph
forbidden_modules =
extensions.ext_database
extensions.ext_redis
allow_indirect_imports = True
ignore_imports =
dify_graph.nodes.llm.node -> extensions.ext_database
dify_graph.model_runtime.model_providers.__base.ai_model -> extensions.ext_redis
dify_graph.model_runtime.model_providers.model_provider_factory -> extensions.ext_redis
[importlinter:contract:workflow-external-imports] [importlinter:contract:workflow-external-imports]
name = Workflow External Imports name = Workflow External Imports
type = forbidden type = forbidden
source_modules = source_modules =
dify_graph graphon
forbidden_modules = forbidden_modules =
constants
configs configs
context
controllers controllers
extensions extensions
factories
libs
models models
services services
tasks tasks
@ -88,46 +82,14 @@ forbidden_modules =
core.tools core.tools
core.trigger core.trigger
core.variables core.variables
ignore_imports =
dify_graph.nodes.llm.llm_utils -> core.model_manager [importlinter:contract:workflow-third-party-imports]
dify_graph.nodes.llm.protocols -> core.model_manager name = Workflow Third-Party Imports
dify_graph.nodes.llm.llm_utils -> dify_graph.model_runtime.model_providers.__base.large_language_model type = forbidden
dify_graph.nodes.llm.node -> core.tools.signature source_modules =
dify_graph.nodes.tool.tool_node -> core.callback_handler.workflow_tool_callback_handler graphon
dify_graph.nodes.tool.tool_node -> core.tools.tool_engine forbidden_modules =
dify_graph.nodes.tool.tool_node -> core.tools.tool_manager sqlalchemy
dify_graph.nodes.parameter_extractor.parameter_extractor_node -> core.prompt.advanced_prompt_transform
dify_graph.nodes.parameter_extractor.parameter_extractor_node -> core.prompt.simple_prompt_transform
dify_graph.nodes.parameter_extractor.parameter_extractor_node -> dify_graph.model_runtime.model_providers.__base.large_language_model
dify_graph.nodes.question_classifier.question_classifier_node -> core.prompt.simple_prompt_transform
dify_graph.nodes.parameter_extractor.parameter_extractor_node -> core.model_manager
dify_graph.nodes.question_classifier.question_classifier_node -> core.model_manager
dify_graph.nodes.tool.tool_node -> core.tools.utils.message_transformer
dify_graph.nodes.llm.node -> core.llm_generator.output_parser.errors
dify_graph.nodes.llm.node -> core.llm_generator.output_parser.structured_output
dify_graph.nodes.llm.node -> core.model_manager
dify_graph.nodes.llm.entities -> core.prompt.entities.advanced_prompt_entities
dify_graph.nodes.llm.node -> core.prompt.entities.advanced_prompt_entities
dify_graph.nodes.llm.node -> core.prompt.utils.prompt_message_util
dify_graph.nodes.parameter_extractor.entities -> core.prompt.entities.advanced_prompt_entities
dify_graph.nodes.parameter_extractor.parameter_extractor_node -> core.prompt.entities.advanced_prompt_entities
dify_graph.nodes.parameter_extractor.parameter_extractor_node -> core.prompt.utils.prompt_message_util
dify_graph.nodes.question_classifier.entities -> core.prompt.entities.advanced_prompt_entities
dify_graph.nodes.question_classifier.question_classifier_node -> core.prompt.utils.prompt_message_util
dify_graph.nodes.llm.node -> models.dataset
dify_graph.nodes.llm.file_saver -> core.tools.signature
dify_graph.nodes.llm.file_saver -> core.tools.tool_file_manager
dify_graph.nodes.tool.tool_node -> core.tools.errors
dify_graph.nodes.llm.node -> extensions.ext_database
dify_graph.nodes.llm.node -> models.model
dify_graph.nodes.tool.tool_node -> services
dify_graph.model_runtime.model_providers.__base.ai_model -> configs
dify_graph.model_runtime.model_providers.__base.ai_model -> extensions.ext_redis
dify_graph.model_runtime.model_providers.__base.large_language_model -> configs
dify_graph.model_runtime.model_providers.__base.text_embedding_model -> core.entities.embedding_type
dify_graph.model_runtime.model_providers.model_provider_factory -> configs
dify_graph.model_runtime.model_providers.model_provider_factory -> extensions.ext_redis
dify_graph.model_runtime.model_providers.model_provider_factory -> models.provider_ids
[importlinter:contract:rsc] [importlinter:contract:rsc]
name = RSC name = RSC
@ -136,7 +98,7 @@ layers =
graph_engine graph_engine
response_coordinator response_coordinator
containers = containers =
dify_graph.graph_engine graphon.graph_engine
[importlinter:contract:worker] [importlinter:contract:worker]
name = Worker name = Worker
@ -145,7 +107,7 @@ layers =
graph_engine graph_engine
worker worker
containers = containers =
dify_graph.graph_engine graphon.graph_engine
[importlinter:contract:graph-engine-architecture] [importlinter:contract:graph-engine-architecture]
name = Graph Engine Architecture name = Graph Engine Architecture
@ -161,28 +123,28 @@ layers =
worker_management worker_management
domain domain
containers = containers =
dify_graph.graph_engine graphon.graph_engine
[importlinter:contract:domain-isolation] [importlinter:contract:domain-isolation]
name = Domain Model Isolation name = Domain Model Isolation
type = forbidden type = forbidden
source_modules = source_modules =
dify_graph.graph_engine.domain graphon.graph_engine.domain
forbidden_modules = forbidden_modules =
dify_graph.graph_engine.worker_management graphon.graph_engine.worker_management
dify_graph.graph_engine.command_channels graphon.graph_engine.command_channels
dify_graph.graph_engine.layers graphon.graph_engine.layers
dify_graph.graph_engine.protocols graphon.graph_engine.protocols
[importlinter:contract:worker-management] [importlinter:contract:worker-management]
name = Worker Management name = Worker Management
type = forbidden type = forbidden
source_modules = source_modules =
dify_graph.graph_engine.worker_management graphon.graph_engine.worker_management
forbidden_modules = forbidden_modules =
dify_graph.graph_engine.orchestration graphon.graph_engine.orchestration
dify_graph.graph_engine.command_processing graphon.graph_engine.command_processing
dify_graph.graph_engine.event_management graphon.graph_engine.event_management
[importlinter:contract:graph-traversal-components] [importlinter:contract:graph-traversal-components]
@ -192,11 +154,11 @@ layers =
edge_processor edge_processor
skip_propagator skip_propagator
containers = containers =
dify_graph.graph_engine.graph_traversal graphon.graph_engine.graph_traversal
[importlinter:contract:command-channels] [importlinter:contract:command-channels]
name = Command Channels Independence name = Command Channels Independence
type = independence type = independence
modules = modules =
dify_graph.graph_engine.command_channels.in_memory_channel graphon.graph_engine.command_channels.in_memory_channel
dify_graph.graph_engine.command_channels.redis_channel graphon.graph_engine.command_channels.redis_channel

View File

@ -100,7 +100,7 @@ ignore = [
"configs/*" = [ "configs/*" = [
"N802", # invalid-function-name "N802", # invalid-function-name
] ]
"dify_graph/model_runtime/callbacks/base_callback.py" = ["T201"] "graphon/model_runtime/callbacks/base_callback.py" = ["T201"]
"core/workflow/callbacks/workflow_logging_callback.py" = ["T201"] "core/workflow/callbacks/workflow_logging_callback.py" = ["T201"]
"libs/gmpy2_pkcs10aep_cipher.py" = [ "libs/gmpy2_pkcs10aep_cipher.py" = [
"N803", # invalid-argument-name "N803", # invalid-argument-name

View File

@ -10,6 +10,7 @@ from configs import dify_config
from core.rag.datasource.vdb.vector_factory import Vector from core.rag.datasource.vdb.vector_factory import Vector
from core.rag.datasource.vdb.vector_type import VectorType from core.rag.datasource.vdb.vector_type import VectorType
from core.rag.index_processor.constant.built_in_field import BuiltInField from core.rag.index_processor.constant.built_in_field import BuiltInField
from core.rag.index_processor.constant.index_type import IndexStructureType, IndexTechniqueType
from core.rag.models.document import ChildDocument, Document from core.rag.models.document import ChildDocument, Document
from extensions.ext_database import db from extensions.ext_database import db
from models.dataset import Dataset, DatasetCollectionBinding, DatasetMetadata, DatasetMetadataBinding, DocumentSegment from models.dataset import Dataset, DatasetCollectionBinding, DatasetMetadata, DatasetMetadataBinding, DocumentSegment
@ -85,7 +86,7 @@ def migrate_annotation_vector_database():
dataset = Dataset( dataset = Dataset(
id=app.id, id=app.id,
tenant_id=app.tenant_id, tenant_id=app.tenant_id,
indexing_technique="high_quality", indexing_technique=IndexTechniqueType.HIGH_QUALITY,
embedding_model_provider=dataset_collection_binding.provider_name, embedding_model_provider=dataset_collection_binding.provider_name,
embedding_model=dataset_collection_binding.model_name, embedding_model=dataset_collection_binding.model_name,
collection_binding_id=dataset_collection_binding.id, collection_binding_id=dataset_collection_binding.id,
@ -177,7 +178,9 @@ def migrate_knowledge_vector_database():
while True: while True:
try: try:
stmt = ( stmt = (
select(Dataset).where(Dataset.indexing_technique == "high_quality").order_by(Dataset.created_at.desc()) select(Dataset)
.where(Dataset.indexing_technique == IndexTechniqueType.HIGH_QUALITY)
.order_by(Dataset.created_at.desc())
) )
datasets = db.paginate(select=stmt, page=page, per_page=50, max_per_page=50, error_out=False) datasets = db.paginate(select=stmt, page=page, per_page=50, max_per_page=50, error_out=False)
@ -269,7 +272,7 @@ def migrate_knowledge_vector_database():
"dataset_id": segment.dataset_id, "dataset_id": segment.dataset_id,
}, },
) )
if dataset_document.doc_form == "hierarchical_model": if dataset_document.doc_form == IndexStructureType.PARENT_CHILD_INDEX:
child_chunks = segment.get_child_chunks() child_chunks = segment.get_child_chunks()
if child_chunks: if child_chunks:
child_documents = [] child_documents = []

View File

@ -51,3 +51,18 @@ class BaiduVectorDBConfig(BaseSettings):
description="Parser mode for inverted index in Baidu Vector Database (default is COARSE_MODE)", description="Parser mode for inverted index in Baidu Vector Database (default is COARSE_MODE)",
default="COARSE_MODE", default="COARSE_MODE",
) )
BAIDU_VECTOR_DB_AUTO_BUILD_ROW_COUNT_INCREMENT: int = Field(
description="Auto build row count increment threshold (default is 500)",
default=500,
)
BAIDU_VECTOR_DB_AUTO_BUILD_ROW_COUNT_INCREMENT_RATIO: float = Field(
description="Auto build row count increment ratio threshold (default is 0.05)",
default=0.05,
)
BAIDU_VECTOR_DB_REBUILD_INDEX_TIMEOUT_IN_SECONDS: int = Field(
description="Timeout in seconds for rebuilding the index in Baidu Vector Database (default is 3600 seconds)",
default=300,
)

View File

@ -1,74 +1,36 @@
""" """
Core Context - Framework-agnostic context management. Application-layer context adapters.
This module provides context management that is independent of any specific Concrete execution-context implementations live here so `graphon` only
web framework. Framework-specific implementations register their context depends on injected context managers rather than framework state capture.
capture functions at application initialization time.
This ensures the workflow layer remains completely decoupled from Flask
or any other web framework.
""" """
import contextvars from context.execution_context import (
from collections.abc import Callable AppContext,
ContextProviderNotFoundError,
from dify_graph.context.execution_context import (
ExecutionContext, ExecutionContext,
ExecutionContextBuilder,
IExecutionContext, IExecutionContext,
NullAppContext, NullAppContext,
capture_current_context,
read_context,
register_context,
register_context_capturer,
reset_context_provider,
) )
from context.models import SandboxContext
# Global capturer function - set by framework-specific modules
_capturer: Callable[[], IExecutionContext] | None = None
def register_context_capturer(capturer: Callable[[], IExecutionContext]) -> None:
"""
Register a context capture function.
This should be called by framework-specific modules (e.g., Flask)
during application initialization.
Args:
capturer: Function that captures current context and returns IExecutionContext
"""
global _capturer
_capturer = capturer
def capture_current_context() -> IExecutionContext:
"""
Capture current execution context.
This function uses the registered context capturer. If no capturer
is registered, it returns a minimal context with only contextvars
(suitable for non-framework environments like tests or standalone scripts).
Returns:
IExecutionContext with captured context
"""
if _capturer is None:
# No framework registered - return minimal context
return ExecutionContext(
app_context=NullAppContext(),
context_vars=contextvars.copy_context(),
)
return _capturer()
def reset_context_provider() -> None:
"""
Reset the context capturer.
This is primarily useful for testing to ensure a clean state.
"""
global _capturer
_capturer = None
__all__ = [ __all__ = [
"AppContext",
"ContextProviderNotFoundError",
"ExecutionContext",
"ExecutionContextBuilder",
"IExecutionContext",
"NullAppContext",
"SandboxContext",
"capture_current_context", "capture_current_context",
"read_context",
"register_context",
"register_context_capturer", "register_context_capturer",
"reset_context_provider", "reset_context_provider",
] ]

View File

@ -1,5 +1,8 @@
""" """
Execution Context - Abstracted context management for workflow execution. Application-layer execution context adapters.
Concrete context capture lives outside `graphon` so the graph package only
consumes injected context managers when it needs to preserve thread-local state.
""" """
import contextvars import contextvars
@ -16,33 +19,33 @@ class AppContext(ABC):
""" """
Abstract application context interface. Abstract application context interface.
This abstraction allows workflow execution to work with or without Flask Application adapters can implement this to restore framework-specific state
by providing a common interface for application context management. such as Flask app context around worker execution.
""" """
@abstractmethod @abstractmethod
def get_config(self, key: str, default: Any = None) -> Any: def get_config(self, key: str, default: Any = None) -> Any:
"""Get configuration value by key.""" """Get configuration value by key."""
pass raise NotImplementedError
@abstractmethod @abstractmethod
def get_extension(self, name: str) -> Any: def get_extension(self, name: str) -> Any:
"""Get Flask extension by name (e.g., 'db', 'cache').""" """Get application extension by name."""
pass raise NotImplementedError
@abstractmethod @abstractmethod
def enter(self) -> AbstractContextManager[None]: def enter(self) -> AbstractContextManager[None]:
"""Enter the application context.""" """Enter the application context."""
pass raise NotImplementedError
@runtime_checkable @runtime_checkable
class IExecutionContext(Protocol): class IExecutionContext(Protocol):
""" """
Protocol for execution context. Protocol for enterable execution context objects.
This protocol defines the interface that all execution contexts must implement, Concrete implementations may carry extra framework state, but callers only
allowing both ExecutionContext and FlaskExecutionContext to be used interchangeably. depend on standard context-manager behavior plus optional user metadata.
""" """
def __enter__(self) -> "IExecutionContext": def __enter__(self) -> "IExecutionContext":
@ -62,14 +65,10 @@ class IExecutionContext(Protocol):
@final @final
class ExecutionContext: class ExecutionContext:
""" """
Execution context for workflow execution in worker threads. Generic execution context used by application-layer adapters.
This class encapsulates all context needed for workflow execution: It restores captured `contextvars` and optionally enters an application
- Application context (Flask app or standalone) context before the worker executes graph logic.
- Context variables for Python contextvars
- User information (optional)
It is designed to be serializable and passable to worker threads.
""" """
def __init__( def __init__(
@ -78,14 +77,6 @@ class ExecutionContext:
context_vars: contextvars.Context | None = None, context_vars: contextvars.Context | None = None,
user: Any = None, user: Any = None,
) -> None: ) -> None:
"""
Initialize execution context.
Args:
app_context: Application context (Flask or standalone)
context_vars: Python contextvars to preserve
user: User object (optional)
"""
self._app_context = app_context self._app_context = app_context
self._context_vars = context_vars self._context_vars = context_vars
self._user = user self._user = user
@ -98,27 +89,21 @@ class ExecutionContext:
@property @property
def context_vars(self) -> contextvars.Context | None: def context_vars(self) -> contextvars.Context | None:
"""Get context variables.""" """Get captured context variables."""
return self._context_vars return self._context_vars
@property @property
def user(self) -> Any: def user(self) -> Any:
"""Get user object.""" """Get captured user object."""
return self._user return self._user
@contextmanager @contextmanager
def enter(self) -> Generator[None, None, None]: def enter(self) -> Generator[None, None, None]:
""" """Enter this execution context."""
Enter this execution context.
This is a convenience method that creates a context manager.
"""
# Restore context variables if provided
if self._context_vars: if self._context_vars:
for var, val in self._context_vars.items(): for var, val in self._context_vars.items():
var.set(val) var.set(val)
# Enter app context if available
if self._app_context is not None: if self._app_context is not None:
with self._app_context.enter(): with self._app_context.enter():
yield yield
@ -141,18 +126,10 @@ class ExecutionContext:
class NullAppContext(AppContext): class NullAppContext(AppContext):
""" """
Null implementation of AppContext for non-Flask environments. Null application context for non-framework environments.
This is used when running without Flask (e.g., in tests or standalone mode).
""" """
def __init__(self, config: dict[str, Any] | None = None) -> None: def __init__(self, config: dict[str, Any] | None = None) -> None:
"""
Initialize null app context.
Args:
config: Optional configuration dictionary
"""
self._config = config or {} self._config = config or {}
self._extensions: dict[str, Any] = {} self._extensions: dict[str, Any] = {}
@ -165,7 +142,7 @@ class NullAppContext(AppContext):
return self._extensions.get(name) return self._extensions.get(name)
def set_extension(self, name: str, extension: Any) -> None: def set_extension(self, name: str, extension: Any) -> None:
"""Set extension by name.""" """Register an extension for tests or standalone execution."""
self._extensions[name] = extension self._extensions[name] = extension
@contextmanager @contextmanager
@ -176,9 +153,7 @@ class NullAppContext(AppContext):
class ExecutionContextBuilder: class ExecutionContextBuilder:
""" """
Builder for creating ExecutionContext instances. Builder for creating `ExecutionContext` instances.
This provides a fluent API for building execution contexts.
""" """
def __init__(self) -> None: def __init__(self) -> None:
@ -211,63 +186,42 @@ class ExecutionContextBuilder:
_capturer: Callable[[], IExecutionContext] | None = None _capturer: Callable[[], IExecutionContext] | None = None
# Tenant-scoped providers using tuple keys for clarity and constant-time lookup.
# Key mapping:
# (name, tenant_id) -> provider
# - name: namespaced identifier (recommend prefixing, e.g. "workflow.sandbox")
# - tenant_id: tenant identifier string
# Value:
# provider: Callable[[], BaseModel] returning the typed context value
# Type-safety note:
# - This registry cannot enforce that all providers for a given name return the same BaseModel type.
# - Implementors SHOULD provide typed wrappers around register/read (like Go's context best practice),
# e.g. def register_sandbox_ctx(tenant_id: str, p: Callable[[], SandboxContext]) and
# def read_sandbox_ctx(tenant_id: str) -> SandboxContext.
_tenant_context_providers: dict[tuple[str, str], Callable[[], BaseModel]] = {} _tenant_context_providers: dict[tuple[str, str], Callable[[], BaseModel]] = {}
T = TypeVar("T", bound=BaseModel) T = TypeVar("T", bound=BaseModel)
class ContextProviderNotFoundError(KeyError): class ContextProviderNotFoundError(KeyError):
"""Raised when a tenant-scoped context provider is missing for a given (name, tenant_id).""" """Raised when a tenant-scoped context provider is missing."""
pass pass
def register_context_capturer(capturer: Callable[[], IExecutionContext]) -> None: def register_context_capturer(capturer: Callable[[], IExecutionContext]) -> None:
"""Register a single enterable execution context capturer (e.g., Flask).""" """Register an enterable execution context capturer."""
global _capturer global _capturer
_capturer = capturer _capturer = capturer
def register_context(name: str, tenant_id: str, provider: Callable[[], BaseModel]) -> None: def register_context(name: str, tenant_id: str, provider: Callable[[], BaseModel]) -> None:
"""Register a tenant-specific provider for a named context. """Register a tenant-specific provider for a named context."""
Tip: use a namespaced "name" (e.g., "workflow.sandbox") to avoid key collisions.
Consider adding a typed wrapper for this registration in your feature module.
"""
_tenant_context_providers[(name, tenant_id)] = provider _tenant_context_providers[(name, tenant_id)] = provider
def read_context(name: str, *, tenant_id: str) -> BaseModel: def read_context(name: str, *, tenant_id: str) -> BaseModel:
""" """Read a context value for a specific tenant."""
Read a context value for a specific tenant. provider = _tenant_context_providers.get((name, tenant_id))
if provider is None:
Raises KeyError if the provider for (name, tenant_id) is not registered.
"""
prov = _tenant_context_providers.get((name, tenant_id))
if prov is None:
raise ContextProviderNotFoundError(f"Context provider '{name}' not registered for tenant '{tenant_id}'") raise ContextProviderNotFoundError(f"Context provider '{name}' not registered for tenant '{tenant_id}'")
return prov() return provider()
def capture_current_context() -> IExecutionContext: def capture_current_context() -> IExecutionContext:
""" """
Capture current execution context from the calling environment. Capture current execution context from the calling environment.
If a capturer is registered (e.g., Flask), use it. Otherwise, return a minimal If no framework adapter is registered, return a minimal context that only
context with NullAppContext + copy of current contextvars. restores `contextvars`.
""" """
if _capturer is None: if _capturer is None:
return ExecutionContext( return ExecutionContext(
@ -278,7 +232,22 @@ def capture_current_context() -> IExecutionContext:
def reset_context_provider() -> None: def reset_context_provider() -> None:
"""Reset the capturer and all tenant-scoped context providers (primarily for tests).""" """Reset the capturer and tenant-scoped providers."""
global _capturer global _capturer
_capturer = None _capturer = None
_tenant_context_providers.clear() _tenant_context_providers.clear()
__all__ = [
"AppContext",
"ContextProviderNotFoundError",
"ExecutionContext",
"ExecutionContextBuilder",
"IExecutionContext",
"NullAppContext",
"capture_current_context",
"read_context",
"register_context",
"register_context_capturer",
"reset_context_provider",
]

View File

@ -10,11 +10,7 @@ from typing import Any, final
from flask import Flask, current_app, g from flask import Flask, current_app, g
from dify_graph.context import register_context_capturer from context.execution_context import AppContext, IExecutionContext, register_context_capturer
from dify_graph.context.execution_context import (
AppContext,
IExecutionContext,
)
@final @final

View File

@ -6,7 +6,6 @@ from contexts.wrapper import RecyclableContextVar
if TYPE_CHECKING: if TYPE_CHECKING:
from core.datasource.__base.datasource_provider import DatasourcePluginProviderController from core.datasource.__base.datasource_provider import DatasourcePluginProviderController
from core.plugin.entities.plugin_daemon import PluginModelProviderEntity
from core.tools.plugin_tool.provider import PluginToolProviderController from core.tools.plugin_tool.provider import PluginToolProviderController
from core.trigger.provider import PluginTriggerProviderController from core.trigger.provider import PluginTriggerProviderController
@ -20,14 +19,6 @@ plugin_tool_providers: RecyclableContextVar[dict[str, "PluginToolProviderControl
plugin_tool_providers_lock: RecyclableContextVar[Lock] = RecyclableContextVar(ContextVar("plugin_tool_providers_lock")) plugin_tool_providers_lock: RecyclableContextVar[Lock] = RecyclableContextVar(ContextVar("plugin_tool_providers_lock"))
plugin_model_providers: RecyclableContextVar[list["PluginModelProviderEntity"] | None] = RecyclableContextVar(
ContextVar("plugin_model_providers")
)
plugin_model_providers_lock: RecyclableContextVar[Lock] = RecyclableContextVar(
ContextVar("plugin_model_providers_lock")
)
datasource_plugin_providers: RecyclableContextVar[dict[str, "DatasourcePluginProviderController"]] = ( datasource_plugin_providers: RecyclableContextVar[dict[str, "DatasourcePluginProviderController"]] = (
RecyclableContextVar(ContextVar("datasource_plugin_providers")) RecyclableContextVar(ContextVar("datasource_plugin_providers"))
) )

View File

@ -4,7 +4,7 @@ from typing import Any, TypeAlias
from pydantic import BaseModel, ConfigDict, computed_field from pydantic import BaseModel, ConfigDict, computed_field
from dify_graph.file import helpers as file_helpers from graphon.file import helpers as file_helpers
from models.model import IconType from models.model import IconType
JSONValue: TypeAlias = str | int | float | bool | None | dict[str, Any] | list[Any] JSONValue: TypeAlias = str | int | float | bool | None | dict[str, Any] | list[Any]

View File

@ -9,6 +9,7 @@ from extensions.ext_database import db
from libs.helper import TimestampField from libs.helper import TimestampField
from libs.login import current_account_with_tenant, login_required from libs.login import current_account_with_tenant, login_required
from models.dataset import Dataset from models.dataset import Dataset
from models.enums import ApiTokenType
from models.model import ApiToken, App from models.model import ApiToken, App
from services.api_token_service import ApiTokenCache from services.api_token_service import ApiTokenCache
@ -47,7 +48,7 @@ def _get_resource(resource_id, tenant_id, resource_model):
class BaseApiKeyListResource(Resource): class BaseApiKeyListResource(Resource):
method_decorators = [account_initialization_required, login_required, setup_required] method_decorators = [account_initialization_required, login_required, setup_required]
resource_type: str | None = None resource_type: ApiTokenType | None = None
resource_model: type | None = None resource_model: type | None = None
resource_id_field: str | None = None resource_id_field: str | None = None
token_prefix: str | None = None token_prefix: str | None = None
@ -91,6 +92,7 @@ class BaseApiKeyListResource(Resource):
) )
key = ApiToken.generate_api_key(self.token_prefix or "", 24) key = ApiToken.generate_api_key(self.token_prefix or "", 24)
assert self.resource_type is not None, "resource_type must be set"
api_token = ApiToken() api_token = ApiToken()
setattr(api_token, self.resource_id_field, resource_id) setattr(api_token, self.resource_id_field, resource_id)
api_token.tenant_id = current_tenant_id api_token.tenant_id = current_tenant_id
@ -104,7 +106,7 @@ class BaseApiKeyListResource(Resource):
class BaseApiKeyResource(Resource): class BaseApiKeyResource(Resource):
method_decorators = [account_initialization_required, login_required, setup_required] method_decorators = [account_initialization_required, login_required, setup_required]
resource_type: str | None = None resource_type: ApiTokenType | None = None
resource_model: type | None = None resource_model: type | None = None
resource_id_field: str | None = None resource_id_field: str | None = None
@ -159,7 +161,7 @@ class AppApiKeyListResource(BaseApiKeyListResource):
"""Create a new API key for an app""" """Create a new API key for an app"""
return super().post(resource_id) return super().post(resource_id)
resource_type = "app" resource_type = ApiTokenType.APP
resource_model = App resource_model = App
resource_id_field = "app_id" resource_id_field = "app_id"
token_prefix = "app-" token_prefix = "app-"
@ -175,7 +177,7 @@ class AppApiKeyResource(BaseApiKeyResource):
"""Delete an API key for an app""" """Delete an API key for an app"""
return super().delete(resource_id, api_key_id) return super().delete(resource_id, api_key_id)
resource_type = "app" resource_type = ApiTokenType.APP
resource_model = App resource_model = App
resource_id_field = "app_id" resource_id_field = "app_id"
@ -199,7 +201,7 @@ class DatasetApiKeyListResource(BaseApiKeyListResource):
"""Create a new API key for a dataset""" """Create a new API key for a dataset"""
return super().post(resource_id) return super().post(resource_id)
resource_type = "dataset" resource_type = ApiTokenType.DATASET
resource_model = Dataset resource_model = Dataset
resource_id_field = "dataset_id" resource_id_field = "dataset_id"
token_prefix = "ds-" token_prefix = "ds-"
@ -215,6 +217,6 @@ class DatasetApiKeyResource(BaseApiKeyResource):
"""Delete an API key for a dataset""" """Delete an API key for a dataset"""
return super().delete(resource_id, api_key_id) return super().delete(resource_id, api_key_id)
resource_type = "dataset" resource_type = ApiTokenType.DATASET
resource_model = Dataset resource_model = Dataset
resource_id_field = "dataset_id" resource_id_field = "dataset_id"

View File

@ -26,9 +26,9 @@ from controllers.console.wraps import (
from core.ops.ops_trace_manager import OpsTraceManager from core.ops.ops_trace_manager import OpsTraceManager
from core.rag.retrieval.retrieval_methods import RetrievalMethod from core.rag.retrieval.retrieval_methods import RetrievalMethod
from core.trigger.constants import TRIGGER_NODE_TYPES from core.trigger.constants import TRIGGER_NODE_TYPES
from dify_graph.enums import WorkflowExecutionStatus
from dify_graph.file import helpers as file_helpers
from extensions.ext_database import db from extensions.ext_database import db
from graphon.enums import WorkflowExecutionStatus
from graphon.file import helpers as file_helpers
from libs.login import current_account_with_tenant, login_required from libs.login import current_account_with_tenant, login_required
from models import App, DatasetPermissionEnum, Workflow from models import App, DatasetPermissionEnum, Workflow
from models.model import IconType from models.model import IconType
@ -95,7 +95,7 @@ class CreateAppPayload(BaseModel):
name: str = Field(..., min_length=1, description="App name") name: str = Field(..., min_length=1, description="App name")
description: str | None = Field(default=None, description="App description (max 400 chars)", max_length=400) description: str | None = Field(default=None, description="App description (max 400 chars)", max_length=400)
mode: Literal["chat", "agent-chat", "advanced-chat", "workflow", "completion"] = Field(..., description="App mode") mode: Literal["chat", "agent-chat", "advanced-chat", "workflow", "completion"] = Field(..., description="App mode")
icon_type: str | None = Field(default=None, description="Icon type") icon_type: IconType | None = Field(default=None, description="Icon type")
icon: str | None = Field(default=None, description="Icon") icon: str | None = Field(default=None, description="Icon")
icon_background: str | None = Field(default=None, description="Icon background color") icon_background: str | None = Field(default=None, description="Icon background color")
@ -103,7 +103,7 @@ class CreateAppPayload(BaseModel):
class UpdateAppPayload(BaseModel): class UpdateAppPayload(BaseModel):
name: str = Field(..., min_length=1, description="App name") name: str = Field(..., min_length=1, description="App name")
description: str | None = Field(default=None, description="App description (max 400 chars)", max_length=400) description: str | None = Field(default=None, description="App description (max 400 chars)", max_length=400)
icon_type: str | None = Field(default=None, description="Icon type") icon_type: IconType | None = Field(default=None, description="Icon type")
icon: str | None = Field(default=None, description="Icon") icon: str | None = Field(default=None, description="Icon")
icon_background: str | None = Field(default=None, description="Icon background color") icon_background: str | None = Field(default=None, description="Icon background color")
use_icon_as_answer_icon: bool | None = Field(default=None, description="Use icon as answer icon") use_icon_as_answer_icon: bool | None = Field(default=None, description="Use icon as answer icon")
@ -113,7 +113,7 @@ class UpdateAppPayload(BaseModel):
class CopyAppPayload(BaseModel): class CopyAppPayload(BaseModel):
name: str | None = Field(default=None, description="Name for the copied app") name: str | None = Field(default=None, description="Name for the copied app")
description: str | None = Field(default=None, description="Description for the copied app", max_length=400) description: str | None = Field(default=None, description="Description for the copied app", max_length=400)
icon_type: str | None = Field(default=None, description="Icon type") icon_type: IconType | None = Field(default=None, description="Icon type")
icon: str | None = Field(default=None, description="Icon") icon: str | None = Field(default=None, description="Icon")
icon_background: str | None = Field(default=None, description="Icon background color") icon_background: str | None = Field(default=None, description="Icon background color")
@ -594,7 +594,7 @@ class AppApi(Resource):
args_dict: AppService.ArgsDict = { args_dict: AppService.ArgsDict = {
"name": args.name, "name": args.name,
"description": args.description or "", "description": args.description or "",
"icon_type": args.icon_type or "", "icon_type": args.icon_type,
"icon": args.icon or "", "icon": args.icon or "",
"icon_background": args.icon_background or "", "icon_background": args.icon_background or "",
"use_icon_as_answer_icon": args.use_icon_as_answer_icon or False, "use_icon_as_answer_icon": args.use_icon_as_answer_icon or False,

View File

@ -22,7 +22,7 @@ from controllers.console.app.error import (
from controllers.console.app.wraps import get_app_model from controllers.console.app.wraps import get_app_model
from controllers.console.wraps import account_initialization_required, setup_required from controllers.console.wraps import account_initialization_required, setup_required
from core.errors.error import ModelCurrentlyNotSupportError, ProviderTokenNotInitError, QuotaExceededError from core.errors.error import ModelCurrentlyNotSupportError, ProviderTokenNotInitError, QuotaExceededError
from dify_graph.model_runtime.errors.invoke import InvokeError from graphon.model_runtime.errors.invoke import InvokeError
from libs.login import login_required from libs.login import login_required
from models import App, AppMode from models import App, AppMode
from services.audio_service import AudioService from services.audio_service import AudioService

View File

@ -26,7 +26,7 @@ from core.errors.error import (
QuotaExceededError, QuotaExceededError,
) )
from core.helper.trace_id_helper import get_external_trace_id from core.helper.trace_id_helper import get_external_trace_id
from dify_graph.model_runtime.errors.invoke import InvokeError from graphon.model_runtime.errors.invoke import InvokeError
from libs import helper from libs import helper
from libs.helper import uuid_value from libs.helper import uuid_value
from libs.login import current_user, login_required from libs.login import current_user, login_required

View File

@ -458,9 +458,7 @@ class ChatConversationApi(Resource):
args = ChatConversationQuery.model_validate(request.args.to_dict(flat=True)) # type: ignore args = ChatConversationQuery.model_validate(request.args.to_dict(flat=True)) # type: ignore
subquery = ( subquery = (
db.session.query( sa.select(Conversation.id.label("conversation_id"), EndUser.session_id.label("from_end_user_session_id"))
Conversation.id.label("conversation_id"), EndUser.session_id.label("from_end_user_session_id")
)
.outerjoin(EndUser, Conversation.from_end_user_id == EndUser.id) .outerjoin(EndUser, Conversation.from_end_user_id == EndUser.id)
.subquery() .subquery()
) )
@ -595,10 +593,8 @@ class ChatConversationDetailApi(Resource):
def _get_conversation(app_model, conversation_id): def _get_conversation(app_model, conversation_id):
current_user, _ = current_account_with_tenant() current_user, _ = current_account_with_tenant()
conversation = ( conversation = db.session.scalar(
db.session.query(Conversation) sa.select(Conversation).where(Conversation.id == conversation_id, Conversation.app_id == app_model.id).limit(1)
.where(Conversation.id == conversation_id, Conversation.app_id == app_model.id)
.first()
) )
if not conversation: if not conversation:

View File

@ -18,8 +18,8 @@ from core.helper.code_executor.javascript.javascript_code_provider import Javasc
from core.helper.code_executor.python3.python3_code_provider import Python3CodeProvider from core.helper.code_executor.python3.python3_code_provider import Python3CodeProvider
from core.llm_generator.entities import RuleCodeGeneratePayload, RuleGeneratePayload, RuleStructuredOutputPayload from core.llm_generator.entities import RuleCodeGeneratePayload, RuleGeneratePayload, RuleStructuredOutputPayload
from core.llm_generator.llm_generator import LLMGenerator from core.llm_generator.llm_generator import LLMGenerator
from dify_graph.model_runtime.errors.invoke import InvokeError
from extensions.ext_database import db from extensions.ext_database import db
from graphon.model_runtime.errors.invoke import InvokeError
from libs.login import current_account_with_tenant, login_required from libs.login import current_account_with_tenant, login_required
from models import App from models import App
from services.workflow_service import WorkflowService from services.workflow_service import WorkflowService
@ -168,7 +168,7 @@ class InstructionGenerateApi(Resource):
try: try:
# Generate from nothing for a workflow node # Generate from nothing for a workflow node
if (args.current in (code_template, "")) and args.node_id != "": if (args.current in (code_template, "")) and args.node_id != "":
app = db.session.query(App).where(App.id == args.flow_id).first() app = db.session.get(App, args.flow_id)
if not app: if not app:
return {"error": f"app {args.flow_id} not found"}, 400 return {"error": f"app {args.flow_id} not found"}, 400
workflow = WorkflowService().get_draft_workflow(app_model=app) workflow = WorkflowService().get_draft_workflow(app_model=app)

View File

@ -2,6 +2,7 @@ import json
from flask_restx import Resource, marshal_with from flask_restx import Resource, marshal_with
from pydantic import BaseModel, Field from pydantic import BaseModel, Field
from sqlalchemy import select
from werkzeug.exceptions import NotFound from werkzeug.exceptions import NotFound
from controllers.console import console_ns from controllers.console import console_ns
@ -47,7 +48,7 @@ class AppMCPServerController(Resource):
@get_app_model @get_app_model
@marshal_with(app_server_model) @marshal_with(app_server_model)
def get(self, app_model): def get(self, app_model):
server = db.session.query(AppMCPServer).where(AppMCPServer.app_id == app_model.id).first() server = db.session.scalar(select(AppMCPServer).where(AppMCPServer.app_id == app_model.id).limit(1))
return server return server
@console_ns.doc("create_app_mcp_server") @console_ns.doc("create_app_mcp_server")
@ -98,7 +99,7 @@ class AppMCPServerController(Resource):
@edit_permission_required @edit_permission_required
def put(self, app_model): def put(self, app_model):
payload = MCPServerUpdatePayload.model_validate(console_ns.payload or {}) payload = MCPServerUpdatePayload.model_validate(console_ns.payload or {})
server = db.session.query(AppMCPServer).where(AppMCPServer.id == payload.id).first() server = db.session.get(AppMCPServer, payload.id)
if not server: if not server:
raise NotFound() raise NotFound()
@ -135,11 +136,10 @@ class AppMCPServerRefreshController(Resource):
@edit_permission_required @edit_permission_required
def get(self, server_id): def get(self, server_id):
_, current_tenant_id = current_account_with_tenant() _, current_tenant_id = current_account_with_tenant()
server = ( server = db.session.scalar(
db.session.query(AppMCPServer) select(AppMCPServer)
.where(AppMCPServer.id == server_id) .where(AppMCPServer.id == server_id, AppMCPServer.tenant_id == current_tenant_id)
.where(AppMCPServer.tenant_id == current_tenant_id) .limit(1)
.first()
) )
if not server: if not server:
raise NotFound() raise NotFound()

View File

@ -24,9 +24,9 @@ from controllers.console.wraps import (
) )
from core.app.entities.app_invoke_entities import InvokeFrom from core.app.entities.app_invoke_entities import InvokeFrom
from core.errors.error import ModelCurrentlyNotSupportError, ProviderTokenNotInitError, QuotaExceededError from core.errors.error import ModelCurrentlyNotSupportError, ProviderTokenNotInitError, QuotaExceededError
from dify_graph.model_runtime.errors.invoke import InvokeError
from extensions.ext_database import db from extensions.ext_database import db
from fields.raws import FilesContainedField from fields.raws import FilesContainedField
from graphon.model_runtime.errors.invoke import InvokeError
from libs.helper import TimestampField, uuid_value from libs.helper import TimestampField, uuid_value
from libs.infinite_scroll_pagination import InfiniteScrollPagination from libs.infinite_scroll_pagination import InfiniteScrollPagination
from libs.login import current_account_with_tenant, login_required from libs.login import current_account_with_tenant, login_required

View File

@ -69,9 +69,7 @@ class ModelConfigResource(Resource):
if app_model.mode == AppMode.AGENT_CHAT or app_model.is_agent: if app_model.mode == AppMode.AGENT_CHAT or app_model.is_agent:
# get original app model config # get original app model config
original_app_model_config = ( original_app_model_config = db.session.get(AppModelConfig, app_model.app_model_config_id)
db.session.query(AppModelConfig).where(AppModelConfig.id == app_model.app_model_config_id).first()
)
if original_app_model_config is None: if original_app_model_config is None:
raise ValueError("Original app model config not found") raise ValueError("Original app model config not found")
agent_mode = original_app_model_config.agent_mode_dict agent_mode = original_app_model_config.agent_mode_dict
@ -90,6 +88,7 @@ class ModelConfigResource(Resource):
tenant_id=current_tenant_id, tenant_id=current_tenant_id,
app_id=app_model.id, app_id=app_model.id,
agent_tool=agent_tool_entity, agent_tool=agent_tool_entity,
user_id=current_user.id,
) )
manager = ToolParameterConfigurationManager( manager = ToolParameterConfigurationManager(
tenant_id=current_tenant_id, tenant_id=current_tenant_id,
@ -129,6 +128,7 @@ class ModelConfigResource(Resource):
tenant_id=current_tenant_id, tenant_id=current_tenant_id,
app_id=app_model.id, app_id=app_model.id,
agent_tool=agent_tool_entity, agent_tool=agent_tool_entity,
user_id=current_user.id,
) )
except Exception: except Exception:
continue continue

View File

@ -2,6 +2,7 @@ from typing import Literal
from flask_restx import Resource, marshal_with from flask_restx import Resource, marshal_with
from pydantic import BaseModel, Field, field_validator from pydantic import BaseModel, Field, field_validator
from sqlalchemy import select
from werkzeug.exceptions import NotFound from werkzeug.exceptions import NotFound
from constants.languages import supported_language from constants.languages import supported_language
@ -75,7 +76,7 @@ class AppSite(Resource):
def post(self, app_model): def post(self, app_model):
args = AppSiteUpdatePayload.model_validate(console_ns.payload or {}) args = AppSiteUpdatePayload.model_validate(console_ns.payload or {})
current_user, _ = current_account_with_tenant() current_user, _ = current_account_with_tenant()
site = db.session.query(Site).where(Site.app_id == app_model.id).first() site = db.session.scalar(select(Site).where(Site.app_id == app_model.id).limit(1))
if not site: if not site:
raise NotFound raise NotFound
@ -124,7 +125,7 @@ class AppSiteAccessTokenReset(Resource):
@marshal_with(app_site_model) @marshal_with(app_site_model)
def post(self, app_model): def post(self, app_model):
current_user, _ = current_account_with_tenant() current_user, _ = current_account_with_tenant()
site = db.session.query(Site).where(Site.app_id == app_model.id).first() site = db.session.scalar(select(Site).where(Site.app_id == app_model.id).limit(1))
if not site: if not site:
raise NotFound raise NotFound

View File

@ -20,6 +20,7 @@ from core.app.app_config.features.file_upload.manager import FileUploadConfigMan
from core.app.apps.base_app_queue_manager import AppQueueManager from core.app.apps.base_app_queue_manager import AppQueueManager
from core.app.apps.workflow.app_generator import SKIP_PREPARE_USER_INPUTS_KEY from core.app.apps.workflow.app_generator import SKIP_PREPARE_USER_INPUTS_KEY
from core.app.entities.app_invoke_entities import InvokeFrom from core.app.entities.app_invoke_entities import InvokeFrom
from core.app.file_access import DatabaseFileAccessController
from core.helper.trace_id_helper import get_external_trace_id from core.helper.trace_id_helper import get_external_trace_id
from core.plugin.impl.exc import PluginInvokeError from core.plugin.impl.exc import PluginInvokeError
from core.trigger.constants import TRIGGER_SCHEDULE_NODE_TYPE from core.trigger.constants import TRIGGER_SCHEDULE_NODE_TYPE
@ -29,15 +30,15 @@ from core.trigger.debug.event_selectors import (
create_event_poller, create_event_poller,
select_trigger_debug_events, select_trigger_debug_events,
) )
from dify_graph.enums import NodeType
from dify_graph.file.models import File
from dify_graph.graph_engine.manager import GraphEngineManager
from dify_graph.model_runtime.utils.encoders import jsonable_encoder
from extensions.ext_database import db from extensions.ext_database import db
from extensions.ext_redis import redis_client from extensions.ext_redis import redis_client
from factories import file_factory, variable_factory from factories import file_factory, variable_factory
from fields.member_fields import simple_account_fields from fields.member_fields import simple_account_fields
from fields.workflow_fields import workflow_fields, workflow_pagination_fields from fields.workflow_fields import workflow_fields, workflow_pagination_fields
from graphon.enums import NodeType
from graphon.file.models import File
from graphon.graph_engine.manager import GraphEngineManager
from graphon.model_runtime.utils.encoders import jsonable_encoder
from libs import helper from libs import helper
from libs.datetime_utils import naive_utc_now from libs.datetime_utils import naive_utc_now
from libs.helper import TimestampField, uuid_value from libs.helper import TimestampField, uuid_value
@ -51,6 +52,7 @@ from services.errors.llm import InvokeRateLimitError
from services.workflow_service import DraftWorkflowDeletionError, WorkflowInUseError, WorkflowService from services.workflow_service import DraftWorkflowDeletionError, WorkflowInUseError, WorkflowService
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
_file_access_controller = DatabaseFileAccessController()
LISTENING_RETRY_IN = 2000 LISTENING_RETRY_IN = 2000
DEFAULT_REF_TEMPLATE_SWAGGER_2_0 = "#/definitions/{model}" DEFAULT_REF_TEMPLATE_SWAGGER_2_0 = "#/definitions/{model}"
RESTORE_SOURCE_WORKFLOW_MUST_BE_PUBLISHED_MESSAGE = "source workflow must be published" RESTORE_SOURCE_WORKFLOW_MUST_BE_PUBLISHED_MESSAGE = "source workflow must be published"
@ -204,6 +206,7 @@ def _parse_file(workflow: Workflow, files: list[dict] | None = None) -> Sequence
mappings=files, mappings=files,
tenant_id=workflow.tenant_id, tenant_id=workflow.tenant_id,
config=file_extra_config, config=file_extra_config,
access_controller=_file_access_controller,
) )
return file_objs return file_objs

View File

@ -9,12 +9,12 @@ from sqlalchemy.orm import Session
from controllers.console import console_ns from controllers.console import console_ns
from controllers.console.app.wraps import get_app_model from controllers.console.app.wraps import get_app_model
from controllers.console.wraps import account_initialization_required, setup_required from controllers.console.wraps import account_initialization_required, setup_required
from dify_graph.enums import WorkflowExecutionStatus
from extensions.ext_database import db from extensions.ext_database import db
from fields.workflow_app_log_fields import ( from fields.workflow_app_log_fields import (
build_workflow_app_log_pagination_model, build_workflow_app_log_pagination_model,
build_workflow_archived_log_pagination_model, build_workflow_archived_log_pagination_model,
) )
from graphon.enums import WorkflowExecutionStatus
from libs.login import login_required from libs.login import login_required
from models import App from models import App
from models.model import AppMode from models.model import AppMode

View File

@ -15,14 +15,15 @@ from controllers.console.app.error import (
from controllers.console.app.wraps import get_app_model from controllers.console.app.wraps import get_app_model
from controllers.console.wraps import account_initialization_required, edit_permission_required, setup_required from controllers.console.wraps import account_initialization_required, edit_permission_required, setup_required
from controllers.web.error import InvalidArgumentError, NotFoundError from controllers.web.error import InvalidArgumentError, NotFoundError
from dify_graph.constants import CONVERSATION_VARIABLE_NODE_ID, SYSTEM_VARIABLE_NODE_ID from core.app.file_access import DatabaseFileAccessController
from dify_graph.file import helpers as file_helpers from core.workflow.variable_prefixes import CONVERSATION_VARIABLE_NODE_ID, SYSTEM_VARIABLE_NODE_ID
from dify_graph.variables.segment_group import SegmentGroup
from dify_graph.variables.segments import ArrayFileSegment, FileSegment, Segment
from dify_graph.variables.types import SegmentType
from extensions.ext_database import db from extensions.ext_database import db
from factories.file_factory import build_from_mapping, build_from_mappings from factories.file_factory import build_from_mapping, build_from_mappings
from factories.variable_factory import build_segment_with_type from factories.variable_factory import build_segment_with_type
from graphon.file import helpers as file_helpers
from graphon.variables.segment_group import SegmentGroup
from graphon.variables.segments import ArrayFileSegment, FileSegment, Segment
from graphon.variables.types import SegmentType
from libs.login import current_user, login_required from libs.login import current_user, login_required
from models import App, AppMode from models import App, AppMode
from models.workflow import WorkflowDraftVariable from models.workflow import WorkflowDraftVariable
@ -30,6 +31,7 @@ from services.workflow_draft_variable_service import WorkflowDraftVariableList,
from services.workflow_service import WorkflowService from services.workflow_service import WorkflowService
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
_file_access_controller = DatabaseFileAccessController()
DEFAULT_REF_TEMPLATE_SWAGGER_2_0 = "#/definitions/{model}" DEFAULT_REF_TEMPLATE_SWAGGER_2_0 = "#/definitions/{model}"
@ -389,13 +391,21 @@ class VariableApi(Resource):
if variable.value_type == SegmentType.FILE: if variable.value_type == SegmentType.FILE:
if not isinstance(raw_value, dict): if not isinstance(raw_value, dict):
raise InvalidArgumentError(description=f"expected dict for file, got {type(raw_value)}") raise InvalidArgumentError(description=f"expected dict for file, got {type(raw_value)}")
raw_value = build_from_mapping(mapping=raw_value, tenant_id=app_model.tenant_id) raw_value = build_from_mapping(
mapping=raw_value,
tenant_id=app_model.tenant_id,
access_controller=_file_access_controller,
)
elif variable.value_type == SegmentType.ARRAY_FILE: elif variable.value_type == SegmentType.ARRAY_FILE:
if not isinstance(raw_value, list): if not isinstance(raw_value, list):
raise InvalidArgumentError(description=f"expected list for files, got {type(raw_value)}") raise InvalidArgumentError(description=f"expected list for files, got {type(raw_value)}")
if len(raw_value) > 0 and not isinstance(raw_value[0], dict): if len(raw_value) > 0 and not isinstance(raw_value[0], dict):
raise InvalidArgumentError(description=f"expected dict for files[0], got {type(raw_value)}") raise InvalidArgumentError(description=f"expected dict for files[0], got {type(raw_value)}")
raw_value = build_from_mappings(mappings=raw_value, tenant_id=app_model.tenant_id) raw_value = build_from_mappings(
mappings=raw_value,
tenant_id=app_model.tenant_id,
access_controller=_file_access_controller,
)
new_value = build_segment_with_type(variable.value_type, raw_value) new_value = build_segment_with_type(variable.value_type, raw_value)
draft_var_srv.update_variable(variable, name=new_name, value=new_value) draft_var_srv.update_variable(variable, name=new_name, value=new_value)
db.session.commit() db.session.commit()

View File

@ -12,8 +12,7 @@ from controllers.console import console_ns
from controllers.console.app.wraps import get_app_model from controllers.console.app.wraps import get_app_model
from controllers.console.wraps import account_initialization_required, setup_required from controllers.console.wraps import account_initialization_required, setup_required
from controllers.web.error import NotFoundError from controllers.web.error import NotFoundError
from dify_graph.entities.pause_reason import HumanInputRequired from core.workflow.human_input_forms import load_form_tokens_by_form_id as _load_form_tokens_by_form_id
from dify_graph.enums import WorkflowExecutionStatus
from extensions.ext_database import db from extensions.ext_database import db
from fields.end_user_fields import simple_end_user_fields from fields.end_user_fields import simple_end_user_fields
from fields.member_fields import simple_account_fields from fields.member_fields import simple_account_fields
@ -27,6 +26,8 @@ from fields.workflow_run_fields import (
workflow_run_node_execution_list_fields, workflow_run_node_execution_list_fields,
workflow_run_pagination_fields, workflow_run_pagination_fields,
) )
from graphon.entities.pause_reason import HumanInputRequired
from graphon.enums import WorkflowExecutionStatus
from libs.archive_storage import ArchiveStorageNotConfiguredError, get_archive_storage from libs.archive_storage import ArchiveStorageNotConfiguredError, get_archive_storage
from libs.custom_inputs import time_duration from libs.custom_inputs import time_duration
from libs.helper import uuid_value from libs.helper import uuid_value
@ -496,6 +497,9 @@ class ConsoleWorkflowPauseDetailsApi(Resource):
pause_entity = workflow_run_repo.get_workflow_pause(workflow_run_id) pause_entity = workflow_run_repo.get_workflow_pause(workflow_run_id)
pause_reasons = pause_entity.get_pause_reasons() if pause_entity else [] pause_reasons = pause_entity.get_pause_reasons() if pause_entity else []
form_tokens_by_form_id = _load_form_tokens_by_form_id(
[reason.form_id for reason in pause_reasons if isinstance(reason, HumanInputRequired)]
)
# Build response # Build response
paused_at = pause_entity.paused_at if pause_entity else None paused_at = pause_entity.paused_at if pause_entity else None
@ -514,7 +518,9 @@ class ConsoleWorkflowPauseDetailsApi(Resource):
"pause_type": { "pause_type": {
"type": "human_input", "type": "human_input",
"form_id": reason.form_id, "form_id": reason.form_id,
"backstage_input_url": _build_backstage_input_url(reason.form_token), "backstage_input_url": _build_backstage_input_url(
form_tokens_by_form_id.get(reason.form_id)
),
}, },
} }
) )

View File

@ -2,6 +2,8 @@ from collections.abc import Callable
from functools import wraps from functools import wraps
from typing import ParamSpec, TypeVar, Union from typing import ParamSpec, TypeVar, Union
from sqlalchemy import select
from controllers.console.app.error import AppNotFoundError from controllers.console.app.error import AppNotFoundError
from extensions.ext_database import db from extensions.ext_database import db
from libs.login import current_account_with_tenant from libs.login import current_account_with_tenant
@ -15,16 +17,14 @@ R1 = TypeVar("R1")
def _load_app_model(app_id: str) -> App | None: def _load_app_model(app_id: str) -> App | None:
_, current_tenant_id = current_account_with_tenant() _, current_tenant_id = current_account_with_tenant()
app_model = ( app_model = db.session.scalar(
db.session.query(App) select(App).where(App.id == app_id, App.tenant_id == current_tenant_id, App.status == "normal").limit(1)
.where(App.id == app_id, App.tenant_id == current_tenant_id, App.status == "normal")
.first()
) )
return app_model return app_model
def _load_app_model_with_trial(app_id: str) -> App | None: def _load_app_model_with_trial(app_id: str) -> App | None:
app_model = db.session.query(App).where(App.id == app_id, App.status == "normal").first() app_model = db.session.scalar(select(App).where(App.id == app_id, App.status == "normal").limit(1))
return app_model return app_model

View File

@ -1,7 +1,7 @@
from flask import request from flask import request
from flask_restx import Resource from flask_restx import Resource
from pydantic import BaseModel, Field, field_validator from pydantic import BaseModel, Field, field_validator
from sqlalchemy.orm import Session from sqlalchemy.orm import sessionmaker
from configs import dify_config from configs import dify_config
from constants.languages import languages from constants.languages import languages
@ -73,7 +73,7 @@ class EmailRegisterSendEmailApi(Resource):
if dify_config.BILLING_ENABLED and BillingService.is_email_in_freeze(normalized_email): if dify_config.BILLING_ENABLED and BillingService.is_email_in_freeze(normalized_email):
raise AccountInFreezeError() raise AccountInFreezeError()
with Session(db.engine) as session: with sessionmaker(db.engine).begin() as session:
account = AccountService.get_account_by_email_with_case_fallback(args.email, session=session) account = AccountService.get_account_by_email_with_case_fallback(args.email, session=session)
token = AccountService.send_email_register_email(email=normalized_email, account=account, language=language) token = AccountService.send_email_register_email(email=normalized_email, account=account, language=language)
return {"result": "success", "data": token} return {"result": "success", "data": token}
@ -145,7 +145,7 @@ class EmailRegisterResetApi(Resource):
email = register_data.get("email", "") email = register_data.get("email", "")
normalized_email = email.lower() normalized_email = email.lower()
with Session(db.engine) as session: with sessionmaker(db.engine).begin() as session:
account = AccountService.get_account_by_email_with_case_fallback(email, session=session) account = AccountService.get_account_by_email_with_case_fallback(email, session=session)
if account: if account:

View File

@ -4,7 +4,7 @@ import secrets
from flask import request from flask import request
from flask_restx import Resource from flask_restx import Resource
from pydantic import BaseModel, Field, field_validator from pydantic import BaseModel, Field, field_validator
from sqlalchemy.orm import Session from sqlalchemy.orm import sessionmaker
from controllers.common.schema import register_schema_models from controllers.common.schema import register_schema_models
from controllers.console import console_ns from controllers.console import console_ns
@ -102,7 +102,7 @@ class ForgotPasswordSendEmailApi(Resource):
else: else:
language = "en-US" language = "en-US"
with Session(db.engine) as session: with sessionmaker(db.engine).begin() as session:
account = AccountService.get_account_by_email_with_case_fallback(args.email, session=session) account = AccountService.get_account_by_email_with_case_fallback(args.email, session=session)
token = AccountService.send_reset_password_email( token = AccountService.send_reset_password_email(
@ -201,7 +201,7 @@ class ForgotPasswordResetApi(Resource):
password_hashed = hash_password(args.new_password, salt) password_hashed = hash_password(args.new_password, salt)
email = reset_data.get("email", "") email = reset_data.get("email", "")
with Session(db.engine) as session: with sessionmaker(db.engine).begin() as session:
account = AccountService.get_account_by_email_with_case_fallback(email, session=session) account = AccountService.get_account_by_email_with_case_fallback(email, session=session)
if account: if account:
@ -215,7 +215,6 @@ class ForgotPasswordResetApi(Resource):
# Update existing account credentials # Update existing account credentials
account.password = base64.b64encode(password_hashed).decode() account.password = base64.b64encode(password_hashed).decode()
account.password_salt = base64.b64encode(salt).decode() account.password_salt = base64.b64encode(salt).decode()
session.commit()
# Create workspace if needed # Create workspace if needed
if ( if (

View File

@ -4,7 +4,7 @@ import urllib.parse
import httpx import httpx
from flask import current_app, redirect, request from flask import current_app, redirect, request
from flask_restx import Resource from flask_restx import Resource
from sqlalchemy.orm import Session from sqlalchemy.orm import sessionmaker
from werkzeug.exceptions import Unauthorized from werkzeug.exceptions import Unauthorized
from configs import dify_config from configs import dify_config
@ -180,7 +180,7 @@ def _get_account_by_openid_or_email(provider: str, user_info: OAuthUserInfo) ->
account: Account | None = Account.get_by_openid(provider, user_info.id) account: Account | None = Account.get_by_openid(provider, user_info.id)
if not account: if not account:
with Session(db.engine) as session: with sessionmaker(db.engine).begin() as session:
account = AccountService.get_account_by_email_with_case_fallback(user_info.email, session=session) account = AccountService.get_account_by_email_with_case_fallback(user_info.email, session=session)
return account return account

View File

@ -8,7 +8,7 @@ from pydantic import BaseModel
from werkzeug.exceptions import BadRequest, NotFound from werkzeug.exceptions import BadRequest, NotFound
from controllers.console.wraps import account_initialization_required, setup_required from controllers.console.wraps import account_initialization_required, setup_required
from dify_graph.model_runtime.utils.encoders import jsonable_encoder from graphon.model_runtime.utils.encoders import jsonable_encoder
from libs.login import current_account_with_tenant, login_required from libs.login import current_account_with_tenant, login_required
from models import Account from models import Account
from models.model import OAuthProviderApp from models.model import OAuthProviderApp

View File

@ -5,7 +5,7 @@ from urllib.parse import quote
from flask import Response, request from flask import Response, request
from flask_restx import Resource, fields, marshal, marshal_with from flask_restx import Resource, fields, marshal, marshal_with
from pydantic import BaseModel, Field, field_validator from pydantic import BaseModel, Field, field_validator
from sqlalchemy import select from sqlalchemy import func, select
from sqlalchemy.orm import Session from sqlalchemy.orm import Session
from werkzeug.exceptions import BadRequest, Forbidden, NotFound from werkzeug.exceptions import BadRequest, Forbidden, NotFound
@ -29,12 +29,12 @@ from controllers.console.wraps import (
from core.errors.error import LLMBadRequestError, ProviderTokenNotInitError from core.errors.error import LLMBadRequestError, ProviderTokenNotInitError
from core.evaluation.entities.evaluation_entity import EvaluationCategory, EvaluationConfigData, EvaluationRunRequest from core.evaluation.entities.evaluation_entity import EvaluationCategory, EvaluationConfigData, EvaluationRunRequest
from core.indexing_runner import IndexingRunner from core.indexing_runner import IndexingRunner
from core.provider_manager import ProviderManager from core.plugin.impl.model_runtime_factory import create_plugin_provider_manager
from core.rag.datasource.vdb.vector_type import VectorType from core.rag.datasource.vdb.vector_type import VectorType
from core.rag.extractor.entity.datasource_type import DatasourceType from core.rag.extractor.entity.datasource_type import DatasourceType
from core.rag.extractor.entity.extract_setting import ExtractSetting, NotionInfo, WebsiteInfo from core.rag.extractor.entity.extract_setting import ExtractSetting, NotionInfo, WebsiteInfo
from core.rag.index_processor.constant.index_type import IndexTechniqueType
from core.rag.retrieval.retrieval_methods import RetrievalMethod from core.rag.retrieval.retrieval_methods import RetrievalMethod
from dify_graph.model_runtime.entities.model_entities import ModelType
from extensions.ext_database import db from extensions.ext_database import db
from extensions.ext_storage import storage from extensions.ext_storage import storage
from fields.app_fields import app_detail_kernel_fields, related_app_list from fields.app_fields import app_detail_kernel_fields, related_app_list
@ -56,10 +56,11 @@ from fields.dataset_fields import (
weighted_score_fields, weighted_score_fields,
) )
from fields.document_fields import document_status_fields from fields.document_fields import document_status_fields
from graphon.model_runtime.entities.model_entities import ModelType
from libs.login import current_account_with_tenant, login_required from libs.login import current_account_with_tenant, login_required
from models import ApiToken, Dataset, Document, DocumentSegment, EvaluationRun, EvaluationTargetType, UploadFile from models import ApiToken, Dataset, Document, DocumentSegment, EvaluationRun, EvaluationTargetType, UploadFile
from models.dataset import DatasetPermission, DatasetPermissionEnum from models.dataset import DatasetPermission, DatasetPermissionEnum
from models.enums import SegmentStatus from models.enums import ApiTokenType, SegmentStatus
from models.provider_ids import ModelProviderID from models.provider_ids import ModelProviderID
from services.api_token_service import ApiTokenCache from services.api_token_service import ApiTokenCache
from services.dataset_service import DatasetPermissionService, DatasetService, DocumentService from services.dataset_service import DatasetPermissionService, DatasetService, DocumentService
@ -343,7 +344,7 @@ class DatasetListApi(Resource):
) )
# check embedding setting # check embedding setting
provider_manager = ProviderManager() provider_manager = create_plugin_provider_manager(tenant_id=current_tenant_id)
configurations = provider_manager.get_configurations(tenant_id=current_tenant_id) configurations = provider_manager.get_configurations(tenant_id=current_tenant_id)
embedding_models = configurations.get_models(model_type=ModelType.TEXT_EMBEDDING, only_active=True) embedding_models = configurations.get_models(model_type=ModelType.TEXT_EMBEDDING, only_active=True)
@ -367,7 +368,7 @@ class DatasetListApi(Resource):
for item in data: for item in data:
# convert embedding_model_provider to plugin standard format # convert embedding_model_provider to plugin standard format
if item["indexing_technique"] == "high_quality" and item["embedding_model_provider"]: if item["indexing_technique"] == IndexTechniqueType.HIGH_QUALITY and item["embedding_model_provider"]:
item["embedding_model_provider"] = str(ModelProviderID(item["embedding_model_provider"])) item["embedding_model_provider"] = str(ModelProviderID(item["embedding_model_provider"]))
item_model = f"{item['embedding_model']}:{item['embedding_model_provider']}" item_model = f"{item['embedding_model']}:{item['embedding_model_provider']}"
if item_model in model_names: if item_model in model_names:
@ -448,7 +449,7 @@ class DatasetApi(Resource):
except services.errors.account.NoPermissionError as e: except services.errors.account.NoPermissionError as e:
raise Forbidden(str(e)) raise Forbidden(str(e))
data = cast(dict[str, Any], marshal(dataset, dataset_detail_fields)) data = cast(dict[str, Any], marshal(dataset, dataset_detail_fields))
if dataset.indexing_technique == "high_quality": if dataset.indexing_technique == IndexTechniqueType.HIGH_QUALITY:
if dataset.embedding_model_provider: if dataset.embedding_model_provider:
provider_id = ModelProviderID(dataset.embedding_model_provider) provider_id = ModelProviderID(dataset.embedding_model_provider)
data["embedding_model_provider"] = str(provider_id) data["embedding_model_provider"] = str(provider_id)
@ -457,7 +458,7 @@ class DatasetApi(Resource):
data.update({"partial_member_list": part_users_list}) data.update({"partial_member_list": part_users_list})
# check embedding setting # check embedding setting
provider_manager = ProviderManager() provider_manager = create_plugin_provider_manager(tenant_id=current_tenant_id)
configurations = provider_manager.get_configurations(tenant_id=current_tenant_id) configurations = provider_manager.get_configurations(tenant_id=current_tenant_id)
embedding_models = configurations.get_models(model_type=ModelType.TEXT_EMBEDDING, only_active=True) embedding_models = configurations.get_models(model_type=ModelType.TEXT_EMBEDDING, only_active=True)
@ -466,7 +467,7 @@ class DatasetApi(Resource):
for embedding_model in embedding_models: for embedding_model in embedding_models:
model_names.append(f"{embedding_model.model}:{embedding_model.provider.provider}") model_names.append(f"{embedding_model.model}:{embedding_model.provider.provider}")
if data["indexing_technique"] == "high_quality": if data["indexing_technique"] == IndexTechniqueType.HIGH_QUALITY:
item_model = f"{data['embedding_model']}:{data['embedding_model_provider']}" item_model = f"{data['embedding_model']}:{data['embedding_model_provider']}"
if item_model in model_names: if item_model in model_names:
data["embedding_available"] = True data["embedding_available"] = True
@ -497,7 +498,7 @@ class DatasetApi(Resource):
current_user, current_tenant_id = current_account_with_tenant() current_user, current_tenant_id = current_account_with_tenant()
# check embedding model setting # check embedding model setting
if ( if (
payload.indexing_technique == "high_quality" payload.indexing_technique == IndexTechniqueType.HIGH_QUALITY
and payload.embedding_model_provider is not None and payload.embedding_model_provider is not None
and payload.embedding_model is not None and payload.embedding_model is not None
): ):
@ -750,20 +751,23 @@ class DatasetIndexingStatusApi(Resource):
documents_status = [] documents_status = []
for document in documents: for document in documents:
completed_segments = ( completed_segments = (
db.session.query(DocumentSegment) db.session.scalar(
.where( select(func.count(DocumentSegment.id)).where(
DocumentSegment.completed_at.isnot(None), DocumentSegment.completed_at.isnot(None),
DocumentSegment.document_id == str(document.id), DocumentSegment.document_id == str(document.id),
DocumentSegment.status != SegmentStatus.RE_SEGMENT, DocumentSegment.status != SegmentStatus.RE_SEGMENT,
)
) )
.count() or 0
) )
total_segments = ( total_segments = (
db.session.query(DocumentSegment) db.session.scalar(
.where( select(func.count(DocumentSegment.id)).where(
DocumentSegment.document_id == str(document.id), DocumentSegment.status != SegmentStatus.RE_SEGMENT DocumentSegment.document_id == str(document.id),
DocumentSegment.status != SegmentStatus.RE_SEGMENT,
)
) )
.count() or 0
) )
# Create a dictionary with document attributes and additional fields # Create a dictionary with document attributes and additional fields
document_dict = { document_dict = {
@ -789,7 +793,7 @@ class DatasetIndexingStatusApi(Resource):
class DatasetApiKeyApi(Resource): class DatasetApiKeyApi(Resource):
max_keys = 10 max_keys = 10
token_prefix = "dataset-" token_prefix = "dataset-"
resource_type = "dataset" resource_type = ApiTokenType.DATASET
@console_ns.doc("get_dataset_api_keys") @console_ns.doc("get_dataset_api_keys")
@console_ns.doc(description="Get dataset API keys") @console_ns.doc(description="Get dataset API keys")
@ -814,9 +818,12 @@ class DatasetApiKeyApi(Resource):
_, current_tenant_id = current_account_with_tenant() _, current_tenant_id = current_account_with_tenant()
current_key_count = ( current_key_count = (
db.session.query(ApiToken) db.session.scalar(
.where(ApiToken.type == self.resource_type, ApiToken.tenant_id == current_tenant_id) select(func.count(ApiToken.id)).where(
.count() ApiToken.type == self.resource_type, ApiToken.tenant_id == current_tenant_id
)
)
or 0
) )
if current_key_count >= self.max_keys: if current_key_count >= self.max_keys:
@ -838,7 +845,7 @@ class DatasetApiKeyApi(Resource):
@console_ns.route("/datasets/api-keys/<uuid:api_key_id>") @console_ns.route("/datasets/api-keys/<uuid:api_key_id>")
class DatasetApiDeleteApi(Resource): class DatasetApiDeleteApi(Resource):
resource_type = "dataset" resource_type = ApiTokenType.DATASET
@console_ns.doc("delete_dataset_api_key") @console_ns.doc("delete_dataset_api_key")
@console_ns.doc(description="Delete dataset API key") @console_ns.doc(description="Delete dataset API key")
@ -851,14 +858,14 @@ class DatasetApiDeleteApi(Resource):
def delete(self, api_key_id): def delete(self, api_key_id):
_, current_tenant_id = current_account_with_tenant() _, current_tenant_id = current_account_with_tenant()
api_key_id = str(api_key_id) api_key_id = str(api_key_id)
key = ( key = db.session.scalar(
db.session.query(ApiToken) select(ApiToken)
.where( .where(
ApiToken.tenant_id == current_tenant_id, ApiToken.tenant_id == current_tenant_id,
ApiToken.type == self.resource_type, ApiToken.type == self.resource_type,
ApiToken.id == api_key_id, ApiToken.id == api_key_id,
) )
.first() .limit(1)
) )
if key is None: if key is None:
@ -869,7 +876,7 @@ class DatasetApiDeleteApi(Resource):
assert key is not None # nosec - for type checker only assert key is not None # nosec - for type checker only
ApiTokenCache.delete(key.token, key.type) ApiTokenCache.delete(key.token, key.type)
db.session.query(ApiToken).where(ApiToken.id == api_key_id).delete() db.session.delete(key)
db.session.commit() db.session.commit()
return {"result": "success"}, 204 return {"result": "success"}, 204

View File

@ -10,7 +10,7 @@ import sqlalchemy as sa
from flask import request, send_file from flask import request, send_file
from flask_restx import Resource, fields, marshal, marshal_with from flask_restx import Resource, fields, marshal, marshal_with
from pydantic import BaseModel, Field from pydantic import BaseModel, Field
from sqlalchemy import asc, desc, select from sqlalchemy import asc, desc, func, select
from werkzeug.exceptions import Forbidden, NotFound from werkzeug.exceptions import Forbidden, NotFound
import services import services
@ -27,8 +27,7 @@ from core.model_manager import ModelManager
from core.plugin.impl.exc import PluginDaemonClientSideError from core.plugin.impl.exc import PluginDaemonClientSideError
from core.rag.extractor.entity.datasource_type import DatasourceType from core.rag.extractor.entity.datasource_type import DatasourceType
from core.rag.extractor.entity.extract_setting import ExtractSetting, NotionInfo, WebsiteInfo from core.rag.extractor.entity.extract_setting import ExtractSetting, NotionInfo, WebsiteInfo
from dify_graph.model_runtime.entities.model_entities import ModelType from core.rag.index_processor.constant.index_type import IndexTechniqueType
from dify_graph.model_runtime.errors.invoke import InvokeAuthorizationError
from extensions.ext_database import db from extensions.ext_database import db
from fields.dataset_fields import dataset_fields from fields.dataset_fields import dataset_fields
from fields.document_fields import ( from fields.document_fields import (
@ -38,6 +37,8 @@ from fields.document_fields import (
document_status_fields, document_status_fields,
document_with_segments_fields, document_with_segments_fields,
) )
from graphon.model_runtime.entities.model_entities import ModelType
from graphon.model_runtime.errors.invoke import InvokeAuthorizationError
from libs.datetime_utils import naive_utc_now from libs.datetime_utils import naive_utc_now
from libs.login import current_account_with_tenant, login_required from libs.login import current_account_with_tenant, login_required
from models import DatasetProcessRule, Document, DocumentSegment, UploadFile from models import DatasetProcessRule, Document, DocumentSegment, UploadFile
@ -211,12 +212,11 @@ class GetProcessRuleApi(Resource):
raise Forbidden(str(e)) raise Forbidden(str(e))
# get the latest process rule # get the latest process rule
dataset_process_rule = ( dataset_process_rule = db.session.scalar(
db.session.query(DatasetProcessRule) select(DatasetProcessRule)
.where(DatasetProcessRule.dataset_id == document.dataset_id) .where(DatasetProcessRule.dataset_id == document.dataset_id)
.order_by(DatasetProcessRule.created_at.desc()) .order_by(DatasetProcessRule.created_at.desc())
.limit(1) .limit(1)
.one_or_none()
) )
if dataset_process_rule: if dataset_process_rule:
mode = dataset_process_rule.mode mode = dataset_process_rule.mode
@ -330,21 +330,23 @@ class DatasetDocumentListApi(Resource):
if fetch: if fetch:
for document in documents: for document in documents:
completed_segments = ( completed_segments = (
db.session.query(DocumentSegment) db.session.scalar(
.where( select(func.count(DocumentSegment.id)).where(
DocumentSegment.completed_at.isnot(None), DocumentSegment.completed_at.isnot(None),
DocumentSegment.document_id == str(document.id), DocumentSegment.document_id == str(document.id),
DocumentSegment.status != SegmentStatus.RE_SEGMENT, DocumentSegment.status != SegmentStatus.RE_SEGMENT,
)
) )
.count() or 0
) )
total_segments = ( total_segments = (
db.session.query(DocumentSegment) db.session.scalar(
.where( select(func.count(DocumentSegment.id)).where(
DocumentSegment.document_id == str(document.id), DocumentSegment.document_id == str(document.id),
DocumentSegment.status != SegmentStatus.RE_SEGMENT, DocumentSegment.status != SegmentStatus.RE_SEGMENT,
)
) )
.count() or 0
) )
document.completed_segments = completed_segments document.completed_segments = completed_segments
document.total_segments = total_segments document.total_segments = total_segments
@ -448,11 +450,11 @@ class DatasetInitApi(Resource):
raise Forbidden() raise Forbidden()
knowledge_config = KnowledgeConfig.model_validate(console_ns.payload or {}) knowledge_config = KnowledgeConfig.model_validate(console_ns.payload or {})
if knowledge_config.indexing_technique == "high_quality": if knowledge_config.indexing_technique == IndexTechniqueType.HIGH_QUALITY:
if knowledge_config.embedding_model is None or knowledge_config.embedding_model_provider is None: if knowledge_config.embedding_model is None or knowledge_config.embedding_model_provider is None:
raise ValueError("embedding model and embedding model provider are required for high quality indexing.") raise ValueError("embedding model and embedding model provider are required for high quality indexing.")
try: try:
model_manager = ModelManager() model_manager = ModelManager.for_tenant(tenant_id=current_tenant_id)
model_manager.get_model_instance( model_manager.get_model_instance(
tenant_id=current_tenant_id, tenant_id=current_tenant_id,
provider=knowledge_config.embedding_model_provider, provider=knowledge_config.embedding_model_provider,
@ -462,7 +464,7 @@ class DatasetInitApi(Resource):
is_multimodal = DatasetService.check_is_multimodal_model( is_multimodal = DatasetService.check_is_multimodal_model(
current_tenant_id, knowledge_config.embedding_model_provider, knowledge_config.embedding_model current_tenant_id, knowledge_config.embedding_model_provider, knowledge_config.embedding_model
) )
knowledge_config.is_multimodal = is_multimodal knowledge_config.is_multimodal = is_multimodal # pyrefly: ignore[bad-assignment]
except InvokeAuthorizationError: except InvokeAuthorizationError:
raise ProviderNotInitializeError( raise ProviderNotInitializeError(
"No Embedding Model available. Please configure a valid provider in the Settings -> Model Provider." "No Embedding Model available. Please configure a valid provider in the Settings -> Model Provider."
@ -521,10 +523,10 @@ class DocumentIndexingEstimateApi(DocumentResource):
if data_source_info and "upload_file_id" in data_source_info: if data_source_info and "upload_file_id" in data_source_info:
file_id = data_source_info["upload_file_id"] file_id = data_source_info["upload_file_id"]
file = ( file = db.session.scalar(
db.session.query(UploadFile) select(UploadFile)
.where(UploadFile.tenant_id == document.tenant_id, UploadFile.id == file_id) .where(UploadFile.tenant_id == document.tenant_id, UploadFile.id == file_id)
.first() .limit(1)
) )
# raise error if file not found # raise error if file not found
@ -586,10 +588,10 @@ class DocumentBatchIndexingEstimateApi(DocumentResource):
if not data_source_info: if not data_source_info:
continue continue
file_id = data_source_info["upload_file_id"] file_id = data_source_info["upload_file_id"]
file_detail = ( file_detail = db.session.scalar(
db.session.query(UploadFile) select(UploadFile)
.where(UploadFile.tenant_id == current_tenant_id, UploadFile.id == file_id) .where(UploadFile.tenant_id == current_tenant_id, UploadFile.id == file_id)
.first() .limit(1)
) )
if file_detail is None: if file_detail is None:
@ -672,20 +674,23 @@ class DocumentBatchIndexingStatusApi(DocumentResource):
documents_status = [] documents_status = []
for document in documents: for document in documents:
completed_segments = ( completed_segments = (
db.session.query(DocumentSegment) db.session.scalar(
.where( select(func.count(DocumentSegment.id)).where(
DocumentSegment.completed_at.isnot(None), DocumentSegment.completed_at.isnot(None),
DocumentSegment.document_id == str(document.id), DocumentSegment.document_id == str(document.id),
DocumentSegment.status != SegmentStatus.RE_SEGMENT, DocumentSegment.status != SegmentStatus.RE_SEGMENT,
)
) )
.count() or 0
) )
total_segments = ( total_segments = (
db.session.query(DocumentSegment) db.session.scalar(
.where( select(func.count(DocumentSegment.id)).where(
DocumentSegment.document_id == str(document.id), DocumentSegment.status != SegmentStatus.RE_SEGMENT DocumentSegment.document_id == str(document.id),
DocumentSegment.status != SegmentStatus.RE_SEGMENT,
)
) )
.count() or 0
) )
# Create a dictionary with document attributes and additional fields # Create a dictionary with document attributes and additional fields
document_dict = { document_dict = {
@ -723,18 +728,23 @@ class DocumentIndexingStatusApi(DocumentResource):
document = self.get_document(dataset_id, document_id) document = self.get_document(dataset_id, document_id)
completed_segments = ( completed_segments = (
db.session.query(DocumentSegment) db.session.scalar(
.where( select(func.count(DocumentSegment.id)).where(
DocumentSegment.completed_at.isnot(None), DocumentSegment.completed_at.isnot(None),
DocumentSegment.document_id == str(document_id), DocumentSegment.document_id == str(document_id),
DocumentSegment.status != SegmentStatus.RE_SEGMENT, DocumentSegment.status != SegmentStatus.RE_SEGMENT,
)
) )
.count() or 0
) )
total_segments = ( total_segments = (
db.session.query(DocumentSegment) db.session.scalar(
.where(DocumentSegment.document_id == str(document_id), DocumentSegment.status != SegmentStatus.RE_SEGMENT) select(func.count(DocumentSegment.id)).where(
.count() DocumentSegment.document_id == str(document_id),
DocumentSegment.status != SegmentStatus.RE_SEGMENT,
)
)
or 0
) )
# Create a dictionary with document attributes and additional fields # Create a dictionary with document attributes and additional fields
@ -1258,11 +1268,11 @@ class DocumentPipelineExecutionLogApi(DocumentResource):
document = DocumentService.get_document(dataset.id, document_id) document = DocumentService.get_document(dataset.id, document_id)
if not document: if not document:
raise NotFound("Document not found.") raise NotFound("Document not found.")
log = ( log = db.session.scalar(
db.session.query(DocumentPipelineExecutionLog) select(DocumentPipelineExecutionLog)
.filter_by(document_id=document_id) .where(DocumentPipelineExecutionLog.document_id == document_id)
.order_by(DocumentPipelineExecutionLog.created_at.desc()) .order_by(DocumentPipelineExecutionLog.created_at.desc())
.first() .limit(1)
) )
if not log: if not log:
return { return {
@ -1328,7 +1338,7 @@ class DocumentGenerateSummaryApi(Resource):
raise BadRequest("document_list cannot be empty.") raise BadRequest("document_list cannot be empty.")
# Check if dataset configuration supports summary generation # Check if dataset configuration supports summary generation
if dataset.indexing_technique != "high_quality": if dataset.indexing_technique != IndexTechniqueType.HIGH_QUALITY:
raise ValueError( raise ValueError(
f"Summary generation is only available for 'high_quality' indexing technique. " f"Summary generation is only available for 'high_quality' indexing technique. "
f"Current indexing technique: {dataset.indexing_technique}" f"Current indexing technique: {dataset.indexing_technique}"

View File

@ -26,10 +26,11 @@ from controllers.console.wraps import (
) )
from core.errors.error import LLMBadRequestError, ProviderTokenNotInitError from core.errors.error import LLMBadRequestError, ProviderTokenNotInitError
from core.model_manager import ModelManager from core.model_manager import ModelManager
from dify_graph.model_runtime.entities.model_entities import ModelType from core.rag.index_processor.constant.index_type import IndexTechniqueType
from extensions.ext_database import db from extensions.ext_database import db
from extensions.ext_redis import redis_client from extensions.ext_redis import redis_client
from fields.segment_fields import child_chunk_fields, segment_fields from fields.segment_fields import child_chunk_fields, segment_fields
from graphon.model_runtime.entities.model_entities import ModelType
from libs.helper import escape_like_pattern from libs.helper import escape_like_pattern
from libs.login import current_account_with_tenant, login_required from libs.login import current_account_with_tenant, login_required
from models.dataset import ChildChunk, DocumentSegment from models.dataset import ChildChunk, DocumentSegment
@ -45,7 +46,7 @@ def _get_segment_with_summary(segment, dataset_id):
"""Helper function to marshal segment and add summary information.""" """Helper function to marshal segment and add summary information."""
from services.summary_index_service import SummaryIndexService from services.summary_index_service import SummaryIndexService
segment_dict = dict(marshal(segment, segment_fields)) segment_dict = dict(marshal(segment, segment_fields)) # type: ignore
# Query summary for this segment (only enabled summaries) # Query summary for this segment (only enabled summaries)
summary = SummaryIndexService.get_segment_summary(segment_id=segment.id, dataset_id=dataset_id) summary = SummaryIndexService.get_segment_summary(segment_id=segment.id, dataset_id=dataset_id)
segment_dict["summary"] = summary.summary_content if summary else None segment_dict["summary"] = summary.summary_content if summary else None
@ -206,7 +207,7 @@ class DatasetDocumentSegmentListApi(Resource):
# Add summary to each segment # Add summary to each segment
segments_with_summary = [] segments_with_summary = []
for segment in segments.items: for segment in segments.items:
segment_dict = dict(marshal(segment, segment_fields)) segment_dict = dict(marshal(segment, segment_fields)) # type: ignore
segment_dict["summary"] = summaries.get(segment.id) segment_dict["summary"] = summaries.get(segment.id)
segments_with_summary.append(segment_dict) segments_with_summary.append(segment_dict)
@ -279,10 +280,10 @@ class DatasetDocumentSegmentApi(Resource):
DatasetService.check_dataset_permission(dataset, current_user) DatasetService.check_dataset_permission(dataset, current_user)
except services.errors.account.NoPermissionError as e: except services.errors.account.NoPermissionError as e:
raise Forbidden(str(e)) raise Forbidden(str(e))
if dataset.indexing_technique == "high_quality": if dataset.indexing_technique == IndexTechniqueType.HIGH_QUALITY:
# check embedding model setting # check embedding model setting
try: try:
model_manager = ModelManager() model_manager = ModelManager.for_tenant(tenant_id=current_tenant_id)
model_manager.get_model_instance( model_manager.get_model_instance(
tenant_id=current_tenant_id, tenant_id=current_tenant_id,
provider=dataset.embedding_model_provider, provider=dataset.embedding_model_provider,
@ -333,9 +334,9 @@ class DatasetDocumentSegmentAddApi(Resource):
if not current_user.is_dataset_editor: if not current_user.is_dataset_editor:
raise Forbidden() raise Forbidden()
# check embedding model setting # check embedding model setting
if dataset.indexing_technique == "high_quality": if dataset.indexing_technique == IndexTechniqueType.HIGH_QUALITY:
try: try:
model_manager = ModelManager() model_manager = ModelManager.for_tenant(tenant_id=current_tenant_id)
model_manager.get_model_instance( model_manager.get_model_instance(
tenant_id=current_tenant_id, tenant_id=current_tenant_id,
provider=dataset.embedding_model_provider, provider=dataset.embedding_model_provider,
@ -383,10 +384,10 @@ class DatasetDocumentSegmentUpdateApi(Resource):
document = DocumentService.get_document(dataset_id, document_id) document = DocumentService.get_document(dataset_id, document_id)
if not document: if not document:
raise NotFound("Document not found.") raise NotFound("Document not found.")
if dataset.indexing_technique == "high_quality": if dataset.indexing_technique == IndexTechniqueType.HIGH_QUALITY:
# check embedding model setting # check embedding model setting
try: try:
model_manager = ModelManager() model_manager = ModelManager.for_tenant(tenant_id=current_tenant_id)
model_manager.get_model_instance( model_manager.get_model_instance(
tenant_id=current_tenant_id, tenant_id=current_tenant_id,
provider=dataset.embedding_model_provider, provider=dataset.embedding_model_provider,
@ -401,10 +402,10 @@ class DatasetDocumentSegmentUpdateApi(Resource):
raise ProviderNotInitializeError(ex.description) raise ProviderNotInitializeError(ex.description)
# check segment # check segment
segment_id = str(segment_id) segment_id = str(segment_id)
segment = ( segment = db.session.scalar(
db.session.query(DocumentSegment) select(DocumentSegment)
.where(DocumentSegment.id == str(segment_id), DocumentSegment.tenant_id == current_tenant_id) .where(DocumentSegment.id == str(segment_id), DocumentSegment.tenant_id == current_tenant_id)
.first() .limit(1)
) )
if not segment: if not segment:
raise NotFound("Segment not found.") raise NotFound("Segment not found.")
@ -447,10 +448,10 @@ class DatasetDocumentSegmentUpdateApi(Resource):
raise NotFound("Document not found.") raise NotFound("Document not found.")
# check segment # check segment
segment_id = str(segment_id) segment_id = str(segment_id)
segment = ( segment = db.session.scalar(
db.session.query(DocumentSegment) select(DocumentSegment)
.where(DocumentSegment.id == str(segment_id), DocumentSegment.tenant_id == current_tenant_id) .where(DocumentSegment.id == str(segment_id), DocumentSegment.tenant_id == current_tenant_id)
.first() .limit(1)
) )
if not segment: if not segment:
raise NotFound("Segment not found.") raise NotFound("Segment not found.")
@ -494,7 +495,7 @@ class DatasetDocumentSegmentBatchImportApi(Resource):
payload = BatchImportPayload.model_validate(console_ns.payload or {}) payload = BatchImportPayload.model_validate(console_ns.payload or {})
upload_file_id = payload.upload_file_id upload_file_id = payload.upload_file_id
upload_file = db.session.query(UploadFile).where(UploadFile.id == upload_file_id).first() upload_file = db.session.scalar(select(UploadFile).where(UploadFile.id == upload_file_id).limit(1))
if not upload_file: if not upload_file:
raise NotFound("UploadFile not found.") raise NotFound("UploadFile not found.")
@ -559,19 +560,19 @@ class ChildChunkAddApi(Resource):
raise NotFound("Document not found.") raise NotFound("Document not found.")
# check segment # check segment
segment_id = str(segment_id) segment_id = str(segment_id)
segment = ( segment = db.session.scalar(
db.session.query(DocumentSegment) select(DocumentSegment)
.where(DocumentSegment.id == str(segment_id), DocumentSegment.tenant_id == current_tenant_id) .where(DocumentSegment.id == str(segment_id), DocumentSegment.tenant_id == current_tenant_id)
.first() .limit(1)
) )
if not segment: if not segment:
raise NotFound("Segment not found.") raise NotFound("Segment not found.")
if not current_user.is_dataset_editor: if not current_user.is_dataset_editor:
raise Forbidden() raise Forbidden()
# check embedding model setting # check embedding model setting
if dataset.indexing_technique == "high_quality": if dataset.indexing_technique == IndexTechniqueType.HIGH_QUALITY:
try: try:
model_manager = ModelManager() model_manager = ModelManager.for_tenant(tenant_id=current_tenant_id)
model_manager.get_model_instance( model_manager.get_model_instance(
tenant_id=current_tenant_id, tenant_id=current_tenant_id,
provider=dataset.embedding_model_provider, provider=dataset.embedding_model_provider,
@ -616,10 +617,10 @@ class ChildChunkAddApi(Resource):
raise NotFound("Document not found.") raise NotFound("Document not found.")
# check segment # check segment
segment_id = str(segment_id) segment_id = str(segment_id)
segment = ( segment = db.session.scalar(
db.session.query(DocumentSegment) select(DocumentSegment)
.where(DocumentSegment.id == str(segment_id), DocumentSegment.tenant_id == current_tenant_id) .where(DocumentSegment.id == str(segment_id), DocumentSegment.tenant_id == current_tenant_id)
.first() .limit(1)
) )
if not segment: if not segment:
raise NotFound("Segment not found.") raise NotFound("Segment not found.")
@ -666,10 +667,10 @@ class ChildChunkAddApi(Resource):
raise NotFound("Document not found.") raise NotFound("Document not found.")
# check segment # check segment
segment_id = str(segment_id) segment_id = str(segment_id)
segment = ( segment = db.session.scalar(
db.session.query(DocumentSegment) select(DocumentSegment)
.where(DocumentSegment.id == str(segment_id), DocumentSegment.tenant_id == current_tenant_id) .where(DocumentSegment.id == str(segment_id), DocumentSegment.tenant_id == current_tenant_id)
.first() .limit(1)
) )
if not segment: if not segment:
raise NotFound("Segment not found.") raise NotFound("Segment not found.")
@ -714,24 +715,24 @@ class ChildChunkUpdateApi(Resource):
raise NotFound("Document not found.") raise NotFound("Document not found.")
# check segment # check segment
segment_id = str(segment_id) segment_id = str(segment_id)
segment = ( segment = db.session.scalar(
db.session.query(DocumentSegment) select(DocumentSegment)
.where(DocumentSegment.id == str(segment_id), DocumentSegment.tenant_id == current_tenant_id) .where(DocumentSegment.id == str(segment_id), DocumentSegment.tenant_id == current_tenant_id)
.first() .limit(1)
) )
if not segment: if not segment:
raise NotFound("Segment not found.") raise NotFound("Segment not found.")
# check child chunk # check child chunk
child_chunk_id = str(child_chunk_id) child_chunk_id = str(child_chunk_id)
child_chunk = ( child_chunk = db.session.scalar(
db.session.query(ChildChunk) select(ChildChunk)
.where( .where(
ChildChunk.id == str(child_chunk_id), ChildChunk.id == str(child_chunk_id),
ChildChunk.tenant_id == current_tenant_id, ChildChunk.tenant_id == current_tenant_id,
ChildChunk.segment_id == segment.id, ChildChunk.segment_id == segment.id,
ChildChunk.document_id == document_id, ChildChunk.document_id == document_id,
) )
.first() .limit(1)
) )
if not child_chunk: if not child_chunk:
raise NotFound("Child chunk not found.") raise NotFound("Child chunk not found.")
@ -771,24 +772,24 @@ class ChildChunkUpdateApi(Resource):
raise NotFound("Document not found.") raise NotFound("Document not found.")
# check segment # check segment
segment_id = str(segment_id) segment_id = str(segment_id)
segment = ( segment = db.session.scalar(
db.session.query(DocumentSegment) select(DocumentSegment)
.where(DocumentSegment.id == str(segment_id), DocumentSegment.tenant_id == current_tenant_id) .where(DocumentSegment.id == str(segment_id), DocumentSegment.tenant_id == current_tenant_id)
.first() .limit(1)
) )
if not segment: if not segment:
raise NotFound("Segment not found.") raise NotFound("Segment not found.")
# check child chunk # check child chunk
child_chunk_id = str(child_chunk_id) child_chunk_id = str(child_chunk_id)
child_chunk = ( child_chunk = db.session.scalar(
db.session.query(ChildChunk) select(ChildChunk)
.where( .where(
ChildChunk.id == str(child_chunk_id), ChildChunk.id == str(child_chunk_id),
ChildChunk.tenant_id == current_tenant_id, ChildChunk.tenant_id == current_tenant_id,
ChildChunk.segment_id == segment.id, ChildChunk.segment_id == segment.id,
ChildChunk.document_id == document_id, ChildChunk.document_id == document_id,
) )
.first() .limit(1)
) )
if not child_chunk: if not child_chunk:
raise NotFound("Child chunk not found.") raise NotFound("Child chunk not found.")

View File

@ -25,7 +25,7 @@ from libs.login import current_account_with_tenant, login_required
from services.dataset_service import DatasetService from services.dataset_service import DatasetService
from services.external_knowledge_service import ExternalDatasetService from services.external_knowledge_service import ExternalDatasetService
from services.hit_testing_service import HitTestingService from services.hit_testing_service import HitTestingService
from services.knowledge_service import ExternalDatasetTestService from services.knowledge_service import BedrockRetrievalSetting, ExternalDatasetTestService
def _build_dataset_detail_model(): def _build_dataset_detail_model():
@ -86,7 +86,7 @@ class ExternalHitTestingPayload(BaseModel):
class BedrockRetrievalPayload(BaseModel): class BedrockRetrievalPayload(BaseModel):
retrieval_setting: dict[str, object] retrieval_setting: "BedrockRetrievalSetting"
query: str query: str
knowledge_id: str knowledge_id: str

View File

@ -19,8 +19,8 @@ from core.errors.error import (
ProviderTokenNotInitError, ProviderTokenNotInitError,
QuotaExceededError, QuotaExceededError,
) )
from dify_graph.model_runtime.errors.invoke import InvokeError
from fields.hit_testing_fields import hit_testing_record_fields from fields.hit_testing_fields import hit_testing_record_fields
from graphon.model_runtime.errors.invoke import InvokeError
from libs.login import current_user from libs.login import current_user
from models.account import Account from models.account import Account
from services.dataset_service import DatasetService from services.dataset_service import DatasetService

View File

@ -10,8 +10,8 @@ from controllers.common.schema import register_schema_models
from controllers.console import console_ns from controllers.console import console_ns
from controllers.console.wraps import account_initialization_required, edit_permission_required, setup_required from controllers.console.wraps import account_initialization_required, edit_permission_required, setup_required
from core.plugin.impl.oauth import OAuthHandler from core.plugin.impl.oauth import OAuthHandler
from dify_graph.model_runtime.errors.validate import CredentialsValidateFailedError from graphon.model_runtime.errors.validate import CredentialsValidateFailedError
from dify_graph.model_runtime.utils.encoders import jsonable_encoder from graphon.model_runtime.utils.encoders import jsonable_encoder
from libs.login import current_account_with_tenant, login_required from libs.login import current_account_with_tenant, login_required
from models.provider_ids import DatasourceProviderID from models.provider_ids import DatasourceProviderID
from services.datasource_provider_service import DatasourceProviderService from services.datasource_provider_service import DatasourceProviderService

View File

@ -21,11 +21,12 @@ from controllers.console.app.workflow_draft_variable import (
from controllers.console.datasets.wraps import get_rag_pipeline from controllers.console.datasets.wraps import get_rag_pipeline
from controllers.console.wraps import account_initialization_required, setup_required from controllers.console.wraps import account_initialization_required, setup_required
from controllers.web.error import InvalidArgumentError, NotFoundError from controllers.web.error import InvalidArgumentError, NotFoundError
from dify_graph.constants import CONVERSATION_VARIABLE_NODE_ID, SYSTEM_VARIABLE_NODE_ID from core.app.file_access import DatabaseFileAccessController
from dify_graph.variables.types import SegmentType from core.workflow.variable_prefixes import CONVERSATION_VARIABLE_NODE_ID, SYSTEM_VARIABLE_NODE_ID
from extensions.ext_database import db from extensions.ext_database import db
from factories.file_factory import build_from_mapping, build_from_mappings from factories.file_factory import build_from_mapping, build_from_mappings
from factories.variable_factory import build_segment_with_type from factories.variable_factory import build_segment_with_type
from graphon.variables.types import SegmentType
from libs.login import current_user, login_required from libs.login import current_user, login_required
from models import Account from models import Account
from models.dataset import Pipeline from models.dataset import Pipeline
@ -33,6 +34,7 @@ from services.rag_pipeline.rag_pipeline import RagPipelineService
from services.workflow_draft_variable_service import WorkflowDraftVariableList, WorkflowDraftVariableService from services.workflow_draft_variable_service import WorkflowDraftVariableList, WorkflowDraftVariableService
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
_file_access_controller = DatabaseFileAccessController()
def _create_pagination_parser(): def _create_pagination_parser():
@ -223,13 +225,21 @@ class RagPipelineVariableApi(Resource):
if variable.value_type == SegmentType.FILE: if variable.value_type == SegmentType.FILE:
if not isinstance(raw_value, dict): if not isinstance(raw_value, dict):
raise InvalidArgumentError(description=f"expected dict for file, got {type(raw_value)}") raise InvalidArgumentError(description=f"expected dict for file, got {type(raw_value)}")
raw_value = build_from_mapping(mapping=raw_value, tenant_id=pipeline.tenant_id) raw_value = build_from_mapping(
mapping=raw_value,
tenant_id=pipeline.tenant_id,
access_controller=_file_access_controller,
)
elif variable.value_type == SegmentType.ARRAY_FILE: elif variable.value_type == SegmentType.ARRAY_FILE:
if not isinstance(raw_value, list): if not isinstance(raw_value, list):
raise InvalidArgumentError(description=f"expected list for files, got {type(raw_value)}") raise InvalidArgumentError(description=f"expected list for files, got {type(raw_value)}")
if len(raw_value) > 0 and not isinstance(raw_value[0], dict): if len(raw_value) > 0 and not isinstance(raw_value[0], dict):
raise InvalidArgumentError(description=f"expected dict for files[0], got {type(raw_value)}") raise InvalidArgumentError(description=f"expected dict for files[0], got {type(raw_value)}")
raw_value = build_from_mappings(mappings=raw_value, tenant_id=pipeline.tenant_id) raw_value = build_from_mappings(
mappings=raw_value,
tenant_id=pipeline.tenant_id,
access_controller=_file_access_controller,
)
new_value = build_segment_with_type(variable.value_type, raw_value) new_value = build_segment_with_type(variable.value_type, raw_value)
draft_var_srv.update_variable(variable, name=new_name, value=new_value) draft_var_srv.update_variable(variable, name=new_name, value=new_value)
db.session.commit() db.session.commit()

View File

@ -37,9 +37,9 @@ from controllers.web.error import InvokeRateLimitError as InvokeRateLimitHttpErr
from core.app.apps.base_app_queue_manager import AppQueueManager from core.app.apps.base_app_queue_manager import AppQueueManager
from core.app.apps.pipeline.pipeline_generator import PipelineGenerator from core.app.apps.pipeline.pipeline_generator import PipelineGenerator
from core.app.entities.app_invoke_entities import InvokeFrom from core.app.entities.app_invoke_entities import InvokeFrom
from dify_graph.model_runtime.utils.encoders import jsonable_encoder
from extensions.ext_database import db from extensions.ext_database import db
from factories import variable_factory from factories import variable_factory
from graphon.model_runtime.utils.encoders import jsonable_encoder
from libs import helper from libs import helper
from libs.helper import TimestampField, UUIDStrOrEmpty from libs.helper import TimestampField, UUIDStrOrEmpty
from libs.login import current_account_with_tenant, current_user, login_required from libs.login import current_account_with_tenant, current_user, login_required

View File

@ -2,6 +2,8 @@ from collections.abc import Callable
from functools import wraps from functools import wraps
from typing import ParamSpec, TypeVar from typing import ParamSpec, TypeVar
from sqlalchemy import select
from controllers.console.datasets.error import PipelineNotFoundError from controllers.console.datasets.error import PipelineNotFoundError
from extensions.ext_database import db from extensions.ext_database import db
from libs.login import current_account_with_tenant from libs.login import current_account_with_tenant
@ -24,10 +26,8 @@ def get_rag_pipeline(view_func: Callable[P, R]):
del kwargs["pipeline_id"] del kwargs["pipeline_id"]
pipeline = ( pipeline = db.session.scalar(
db.session.query(Pipeline) select(Pipeline).where(Pipeline.id == pipeline_id, Pipeline.tenant_id == current_tenant_id).limit(1)
.where(Pipeline.id == pipeline_id, Pipeline.tenant_id == current_tenant_id)
.first()
) )
if not pipeline: if not pipeline:

View File

@ -19,7 +19,7 @@ from controllers.console.app.error import (
) )
from controllers.console.explore.wraps import InstalledAppResource from controllers.console.explore.wraps import InstalledAppResource
from core.errors.error import ModelCurrentlyNotSupportError, ProviderTokenNotInitError, QuotaExceededError from core.errors.error import ModelCurrentlyNotSupportError, ProviderTokenNotInitError, QuotaExceededError
from dify_graph.model_runtime.errors.invoke import InvokeError from graphon.model_runtime.errors.invoke import InvokeError
from services.audio_service import AudioService from services.audio_service import AudioService
from services.errors.audio import ( from services.errors.audio import (
AudioTooLargeServiceError, AudioTooLargeServiceError,

View File

@ -24,8 +24,8 @@ from core.errors.error import (
ProviderTokenNotInitError, ProviderTokenNotInitError,
QuotaExceededError, QuotaExceededError,
) )
from dify_graph.model_runtime.errors.invoke import InvokeError
from extensions.ext_database import db from extensions.ext_database import db
from graphon.model_runtime.errors.invoke import InvokeError
from libs import helper from libs import helper
from libs.datetime_utils import naive_utc_now from libs.datetime_utils import naive_utc_now
from libs.login import current_user from libs.login import current_user

View File

@ -21,9 +21,9 @@ from controllers.console.explore.error import (
from controllers.console.explore.wraps import InstalledAppResource from controllers.console.explore.wraps import InstalledAppResource
from core.app.entities.app_invoke_entities import InvokeFrom from core.app.entities.app_invoke_entities import InvokeFrom
from core.errors.error import ModelCurrentlyNotSupportError, ProviderTokenNotInitError, QuotaExceededError from core.errors.error import ModelCurrentlyNotSupportError, ProviderTokenNotInitError, QuotaExceededError
from dify_graph.model_runtime.errors.invoke import InvokeError
from fields.conversation_fields import ResultResponse from fields.conversation_fields import ResultResponse
from fields.message_fields import MessageInfiniteScrollPagination, MessageListItem, SuggestedQuestionsResponse from fields.message_fields import MessageInfiniteScrollPagination, MessageListItem, SuggestedQuestionsResponse
from graphon.model_runtime.errors.invoke import InvokeError
from libs import helper from libs import helper
from libs.helper import UUIDStrOrEmpty from libs.helper import UUIDStrOrEmpty
from libs.login import current_account_with_tenant from libs.login import current_account_with_tenant

View File

@ -42,8 +42,6 @@ from core.errors.error import (
ProviderTokenNotInitError, ProviderTokenNotInitError,
QuotaExceededError, QuotaExceededError,
) )
from dify_graph.graph_engine.manager import GraphEngineManager
from dify_graph.model_runtime.errors.invoke import InvokeError
from extensions.ext_database import db from extensions.ext_database import db
from extensions.ext_redis import redis_client from extensions.ext_redis import redis_client
from fields.app_fields import ( from fields.app_fields import (
@ -61,6 +59,8 @@ from fields.workflow_fields import (
workflow_fields, workflow_fields,
workflow_partial_fields, workflow_partial_fields,
) )
from graphon.graph_engine.manager import GraphEngineManager
from graphon.model_runtime.errors.invoke import InvokeError
from libs import helper from libs import helper
from libs.helper import uuid_value from libs.helper import uuid_value
from libs.login import current_user from libs.login import current_user

View File

@ -21,9 +21,9 @@ from core.errors.error import (
ProviderTokenNotInitError, ProviderTokenNotInitError,
QuotaExceededError, QuotaExceededError,
) )
from dify_graph.graph_engine.manager import GraphEngineManager
from dify_graph.model_runtime.errors.invoke import InvokeError
from extensions.ext_redis import redis_client from extensions.ext_redis import redis_client
from graphon.graph_engine.manager import GraphEngineManager
from graphon.model_runtime.errors.invoke import InvokeError
from libs import helper from libs import helper
from libs.login import current_account_with_tenant from libs.login import current_account_with_tenant
from models.model import AppMode, InstalledApp from models.model import AppMode, InstalledApp

View File

@ -13,9 +13,9 @@ from controllers.common.errors import (
) )
from controllers.console import console_ns from controllers.console import console_ns
from core.helper import ssrf_proxy from core.helper import ssrf_proxy
from dify_graph.file import helpers as file_helpers
from extensions.ext_database import db from extensions.ext_database import db
from fields.file_fields import FileWithSignedUrl, RemoteFileInfo from fields.file_fields import FileWithSignedUrl, RemoteFileInfo
from graphon.file import helpers as file_helpers
from libs.login import current_account_with_tenant, login_required from libs.login import current_account_with_tenant, login_required
from services.file_service import FileService from services.file_service import FileService

View File

@ -2,7 +2,7 @@ from flask_restx import Resource, fields
from controllers.console import console_ns from controllers.console import console_ns
from controllers.console.wraps import account_initialization_required, setup_required from controllers.console.wraps import account_initialization_required, setup_required
from dify_graph.model_runtime.utils.encoders import jsonable_encoder from graphon.model_runtime.utils.encoders import jsonable_encoder
from libs.login import current_account_with_tenant, login_required from libs.login import current_account_with_tenant, login_required
from services.agent_service import AgentService from services.agent_service import AgentService

View File

@ -8,7 +8,7 @@ from controllers.common.schema import register_schema_models
from controllers.console import console_ns from controllers.console import console_ns
from controllers.console.wraps import account_initialization_required, is_admin_or_owner_required, setup_required from controllers.console.wraps import account_initialization_required, is_admin_or_owner_required, setup_required
from core.plugin.impl.exc import PluginPermissionDeniedError from core.plugin.impl.exc import PluginPermissionDeniedError
from dify_graph.model_runtime.utils.encoders import jsonable_encoder from graphon.model_runtime.utils.encoders import jsonable_encoder
from libs.login import current_account_with_tenant, login_required from libs.login import current_account_with_tenant, login_required
from services.plugin.endpoint_service import EndpointService from services.plugin.endpoint_service import EndpointService

View File

@ -5,8 +5,8 @@ from werkzeug.exceptions import Forbidden
from controllers.common.schema import register_schema_models from controllers.common.schema import register_schema_models
from controllers.console import console_ns from controllers.console import console_ns
from controllers.console.wraps import account_initialization_required, setup_required from controllers.console.wraps import account_initialization_required, setup_required
from dify_graph.model_runtime.entities.model_entities import ModelType from graphon.model_runtime.entities.model_entities import ModelType
from dify_graph.model_runtime.errors.validate import CredentialsValidateFailedError from graphon.model_runtime.errors.validate import CredentialsValidateFailedError
from libs.login import current_account_with_tenant, login_required from libs.login import current_account_with_tenant, login_required
from models import TenantAccountRole from models import TenantAccountRole
from services.model_load_balancing_service import ModelLoadBalancingService from services.model_load_balancing_service import ModelLoadBalancingService

View File

@ -7,9 +7,9 @@ from pydantic import BaseModel, Field, field_validator
from controllers.console import console_ns from controllers.console import console_ns
from controllers.console.wraps import account_initialization_required, is_admin_or_owner_required, setup_required from controllers.console.wraps import account_initialization_required, is_admin_or_owner_required, setup_required
from dify_graph.model_runtime.entities.model_entities import ModelType from graphon.model_runtime.entities.model_entities import ModelType
from dify_graph.model_runtime.errors.validate import CredentialsValidateFailedError from graphon.model_runtime.errors.validate import CredentialsValidateFailedError
from dify_graph.model_runtime.utils.encoders import jsonable_encoder from graphon.model_runtime.utils.encoders import jsonable_encoder
from libs.helper import uuid_value from libs.helper import uuid_value
from libs.login import current_account_with_tenant, login_required from libs.login import current_account_with_tenant, login_required
from services.billing_service import BillingService from services.billing_service import BillingService

View File

@ -8,9 +8,9 @@ from pydantic import BaseModel, Field, field_validator
from controllers.common.schema import register_enum_models, register_schema_models from controllers.common.schema import register_enum_models, register_schema_models
from controllers.console import console_ns from controllers.console import console_ns
from controllers.console.wraps import account_initialization_required, is_admin_or_owner_required, setup_required from controllers.console.wraps import account_initialization_required, is_admin_or_owner_required, setup_required
from dify_graph.model_runtime.entities.model_entities import ModelType from graphon.model_runtime.entities.model_entities import ModelType
from dify_graph.model_runtime.errors.validate import CredentialsValidateFailedError from graphon.model_runtime.errors.validate import CredentialsValidateFailedError
from dify_graph.model_runtime.utils.encoders import jsonable_encoder from graphon.model_runtime.utils.encoders import jsonable_encoder
from libs.helper import uuid_value from libs.helper import uuid_value
from libs.login import current_account_with_tenant, login_required from libs.login import current_account_with_tenant, login_required
from services.model_load_balancing_service import ModelLoadBalancingService from services.model_load_balancing_service import ModelLoadBalancingService
@ -282,14 +282,18 @@ class ModelProviderModelCredentialApi(Resource):
) )
if args.config_from == "predefined-model": if args.config_from == "predefined-model":
available_credentials = model_provider_service.provider_manager.get_provider_available_credentials( available_credentials = model_provider_service.get_provider_available_credentials(
tenant_id=tenant_id, provider_name=provider tenant_id=tenant_id,
provider=provider,
) )
else: else:
# Normalize model_type to the origin value stored in DB (e.g., "text-generation" for LLM) # Normalize model_type to the origin value stored in DB (e.g., "text-generation" for LLM)
normalized_model_type = args.model_type.to_origin_model_type() normalized_model_type = args.model_type.to_origin_model_type()
available_credentials = model_provider_service.provider_manager.get_provider_model_available_credentials( available_credentials = model_provider_service.get_provider_model_available_credentials(
tenant_id=tenant_id, provider_name=provider, model_type=normalized_model_type, model_name=args.model tenant_id=tenant_id,
provider=provider,
model_type=normalized_model_type,
model=args.model,
) )
return jsonable_encoder( return jsonable_encoder(

View File

@ -14,7 +14,7 @@ from controllers.console import console_ns
from controllers.console.workspace import plugin_permission_required from controllers.console.workspace import plugin_permission_required
from controllers.console.wraps import account_initialization_required, is_admin_or_owner_required, setup_required from controllers.console.wraps import account_initialization_required, is_admin_or_owner_required, setup_required
from core.plugin.impl.exc import PluginDaemonClientSideError from core.plugin.impl.exc import PluginDaemonClientSideError
from dify_graph.model_runtime.utils.encoders import jsonable_encoder from graphon.model_runtime.utils.encoders import jsonable_encoder
from libs.login import current_account_with_tenant, login_required from libs.login import current_account_with_tenant, login_required
from models.account import TenantPluginAutoUpgradeStrategy, TenantPluginPermission from models.account import TenantPluginAutoUpgradeStrategy, TenantPluginPermission
from services.plugin.plugin_auto_upgrade_service import PluginAutoUpgradeService from services.plugin.plugin_auto_upgrade_service import PluginAutoUpgradeService

View File

@ -26,8 +26,8 @@ from core.mcp.mcp_client import MCPClient
from core.plugin.entities.plugin_daemon import CredentialType from core.plugin.entities.plugin_daemon import CredentialType
from core.plugin.impl.oauth import OAuthHandler from core.plugin.impl.oauth import OAuthHandler
from core.tools.entities.tool_entities import ApiProviderSchemaType, WorkflowToolParameterConfiguration from core.tools.entities.tool_entities import ApiProviderSchemaType, WorkflowToolParameterConfiguration
from dify_graph.model_runtime.utils.encoders import jsonable_encoder
from extensions.ext_database import db from extensions.ext_database import db
from graphon.model_runtime.utils.encoders import jsonable_encoder
from libs.helper import alphanumeric, uuid_value from libs.helper import alphanumeric, uuid_value
from libs.login import current_account_with_tenant, login_required from libs.login import current_account_with_tenant, login_required
from models.provider_ids import ToolProviderID from models.provider_ids import ToolProviderID

View File

@ -14,8 +14,8 @@ from core.plugin.entities.plugin_daemon import CredentialType
from core.plugin.impl.oauth import OAuthHandler from core.plugin.impl.oauth import OAuthHandler
from core.trigger.entities.entities import SubscriptionBuilderUpdater from core.trigger.entities.entities import SubscriptionBuilderUpdater
from core.trigger.trigger_manager import TriggerManager from core.trigger.trigger_manager import TriggerManager
from dify_graph.model_runtime.utils.encoders import jsonable_encoder
from extensions.ext_database import db from extensions.ext_database import db
from graphon.model_runtime.utils.encoders import jsonable_encoder
from libs.login import current_user, login_required from libs.login import current_user, login_required
from models.account import Account from models.account import Account
from models.provider_ids import TriggerProviderID from models.provider_ids import TriggerProviderID

View File

@ -70,22 +70,25 @@ class ToolFileApi(Resource):
except Exception: except Exception:
raise UnsupportedFileTypeError() raise UnsupportedFileTypeError()
mime_type = tool_file.mime_type
filename = tool_file.filename
response = Response( response = Response(
stream, stream,
mimetype=tool_file.mimetype, mimetype=mime_type,
direct_passthrough=True, direct_passthrough=True,
headers={}, headers={},
) )
if tool_file.size > 0: if tool_file.size > 0:
response.headers["Content-Length"] = str(tool_file.size) response.headers["Content-Length"] = str(tool_file.size)
if args.as_attachment: if args.as_attachment and filename:
encoded_filename = quote(tool_file.name) encoded_filename = quote(filename)
response.headers["Content-Disposition"] = f"attachment; filename*=UTF-8''{encoded_filename}" response.headers["Content-Disposition"] = f"attachment; filename*=UTF-8''{encoded_filename}"
enforce_download_for_html( enforce_download_for_html(
response, response,
mime_type=tool_file.mimetype, mime_type=mime_type,
filename=tool_file.name, filename=filename,
extension=extension, extension=extension,
) )

View File

@ -7,8 +7,8 @@ from pydantic import BaseModel, Field
from werkzeug.exceptions import Forbidden from werkzeug.exceptions import Forbidden
import services import services
from core.tools.signature import verify_plugin_file_signature
from core.tools.tool_file_manager import ToolFileManager from core.tools.tool_file_manager import ToolFileManager
from dify_graph.file.helpers import verify_plugin_file_signature
from fields.file_fields import FileResponse from fields.file_fields import FileResponse
from ..common.errors import ( from ..common.errors import (

View File

@ -16,12 +16,14 @@ api = ExternalApi(
inner_api_ns = Namespace("inner_api", description="Internal API operations", path="/") inner_api_ns = Namespace("inner_api", description="Internal API operations", path="/")
from . import mail as _mail from . import mail as _mail
from .app import dsl as _app_dsl
from .plugin import plugin as _plugin from .plugin import plugin as _plugin
from .workspace import workspace as _workspace from .workspace import workspace as _workspace
api.add_namespace(inner_api_ns) api.add_namespace(inner_api_ns)
__all__ = [ __all__ = [
"_app_dsl",
"_mail", "_mail",
"_plugin", "_plugin",
"_workspace", "_workspace",

View File

@ -0,0 +1 @@

View File

@ -0,0 +1,110 @@
"""Inner API endpoints for app DSL import/export.
Called by the enterprise admin-api service. Import requires ``creator_email``
to attribute the created app; workspace/membership validation is done by the
Go admin-api caller.
"""
from flask import request
from flask_restx import Resource
from pydantic import BaseModel, Field
from sqlalchemy.orm import Session
from controllers.common.schema import register_schema_model
from controllers.console.wraps import setup_required
from controllers.inner_api import inner_api_ns
from controllers.inner_api.wraps import enterprise_inner_api_only
from extensions.ext_database import db
from models import Account, App
from models.account import AccountStatus
from services.app_dsl_service import AppDslService, ImportMode, ImportStatus
class InnerAppDSLImportPayload(BaseModel):
yaml_content: str = Field(description="YAML DSL content")
creator_email: str = Field(description="Email of the workspace member who will own the imported app")
name: str | None = Field(default=None, description="Override app name from DSL")
description: str | None = Field(default=None, description="Override app description from DSL")
register_schema_model(inner_api_ns, InnerAppDSLImportPayload)
@inner_api_ns.route("/enterprise/workspaces/<string:workspace_id>/dsl/import")
class EnterpriseAppDSLImport(Resource):
@setup_required
@enterprise_inner_api_only
@inner_api_ns.doc("enterprise_app_dsl_import")
@inner_api_ns.expect(inner_api_ns.models[InnerAppDSLImportPayload.__name__])
@inner_api_ns.doc(
responses={
200: "Import completed",
202: "Import pending (DSL version mismatch requires confirmation)",
400: "Import failed (business error)",
404: "Creator account not found or inactive",
}
)
def post(self, workspace_id: str):
"""Import a DSL into a workspace on behalf of a specified creator."""
args = InnerAppDSLImportPayload.model_validate(inner_api_ns.payload or {})
account = _get_active_account(args.creator_email)
if account is None:
return {"message": f"account '{args.creator_email}' not found or inactive"}, 404
account.set_tenant_id(workspace_id)
with Session(db.engine) as session:
dsl_service = AppDslService(session)
result = dsl_service.import_app(
account=account,
import_mode=ImportMode.YAML_CONTENT,
yaml_content=args.yaml_content,
name=args.name,
description=args.description,
)
session.commit()
if result.status == ImportStatus.FAILED:
return result.model_dump(mode="json"), 400
if result.status == ImportStatus.PENDING:
return result.model_dump(mode="json"), 202
return result.model_dump(mode="json"), 200
@inner_api_ns.route("/enterprise/apps/<string:app_id>/dsl")
class EnterpriseAppDSLExport(Resource):
@setup_required
@enterprise_inner_api_only
@inner_api_ns.doc(
"enterprise_app_dsl_export",
responses={
200: "Export successful",
404: "App not found",
},
)
def get(self, app_id: str):
"""Export an app's DSL as YAML."""
include_secret = request.args.get("include_secret", "false").lower() == "true"
app_model = db.session.query(App).filter_by(id=app_id).first()
if not app_model:
return {"message": "app not found"}, 404
data = AppDslService.export_dsl(
app_model=app_model,
include_secret=include_secret,
)
return {"data": data}, 200
def _get_active_account(email: str) -> Account | None:
"""Look up an active account by email.
Workspace membership is already validated by the Go admin-api caller.
"""
account = db.session.query(Account).filter_by(email=email).first()
if account is None or account.status != AccountStatus.ACTIVE:
return None
return account

View File

@ -28,8 +28,8 @@ from core.plugin.entities.request import (
RequestRequestUploadFile, RequestRequestUploadFile,
) )
from core.tools.entities.tool_entities import ToolProviderType from core.tools.entities.tool_entities import ToolProviderType
from dify_graph.file.helpers import get_signed_file_url_for_plugin from core.tools.signature import get_signed_file_url_for_plugin
from dify_graph.model_runtime.utils.encoders import jsonable_encoder from graphon.model_runtime.utils.encoders import jsonable_encoder
from libs.helper import length_prefixed_response from libs.helper import length_prefixed_response
from models import Account, Tenant from models import Account, Tenant
from models.model import EndUser from models.model import EndUser

View File

@ -9,8 +9,8 @@ from controllers.common.schema import register_schema_model
from controllers.mcp import mcp_ns from controllers.mcp import mcp_ns
from core.mcp import types as mcp_types from core.mcp import types as mcp_types
from core.mcp.server.streamable_http import handle_mcp_request from core.mcp.server.streamable_http import handle_mcp_request
from dify_graph.variables.input_entities import VariableEntity
from extensions.ext_database import db from extensions.ext_database import db
from graphon.variables.input_entities import VariableEntity
from libs import helper from libs import helper
from models.enums import AppMCPServerStatus from models.enums import AppMCPServerStatus
from models.model import App, AppMCPServer, AppMode, EndUser from models.model import App, AppMCPServer, AppMode, EndUser

View File

@ -21,7 +21,7 @@ from controllers.service_api.app.error import (
) )
from controllers.service_api.wraps import FetchUserArg, WhereisUserArg, validate_app_token from controllers.service_api.wraps import FetchUserArg, WhereisUserArg, validate_app_token
from core.errors.error import ModelCurrentlyNotSupportError, ProviderTokenNotInitError, QuotaExceededError from core.errors.error import ModelCurrentlyNotSupportError, ProviderTokenNotInitError, QuotaExceededError
from dify_graph.model_runtime.errors.invoke import InvokeError from graphon.model_runtime.errors.invoke import InvokeError
from models.model import App, EndUser from models.model import App, EndUser
from services.audio_service import AudioService from services.audio_service import AudioService
from services.errors.audio import ( from services.errors.audio import (

View File

@ -28,7 +28,7 @@ from core.errors.error import (
QuotaExceededError, QuotaExceededError,
) )
from core.helper.trace_id_helper import get_external_trace_id from core.helper.trace_id_helper import get_external_trace_id
from dify_graph.model_runtime.errors.invoke import InvokeError from graphon.model_runtime.errors.invoke import InvokeError
from libs import helper from libs import helper
from libs.helper import UUIDStrOrEmpty from libs.helper import UUIDStrOrEmpty
from models.model import App, AppMode, EndUser from models.model import App, AppMode, EndUser

View File

@ -4,6 +4,7 @@ from urllib.parse import quote
from flask import Response, request from flask import Response, request
from flask_restx import Resource from flask_restx import Resource
from pydantic import BaseModel, Field from pydantic import BaseModel, Field
from sqlalchemy import select
from controllers.common.file_response import enforce_download_for_html from controllers.common.file_response import enforce_download_for_html
from controllers.common.schema import register_schema_model from controllers.common.schema import register_schema_model
@ -102,27 +103,27 @@ class FilePreviewApi(Resource):
raise FileAccessDeniedError("Invalid file or app identifier") raise FileAccessDeniedError("Invalid file or app identifier")
# First, find the MessageFile that references this upload file # First, find the MessageFile that references this upload file
message_file = db.session.query(MessageFile).where(MessageFile.upload_file_id == file_id).first() message_file = db.session.scalar(select(MessageFile).where(MessageFile.upload_file_id == file_id).limit(1))
if not message_file: if not message_file:
raise FileNotFoundError("File not found in message context") raise FileNotFoundError("File not found in message context")
# Get the message and verify it belongs to the requesting app # Get the message and verify it belongs to the requesting app
message = ( message = db.session.scalar(
db.session.query(Message).where(Message.id == message_file.message_id, Message.app_id == app_id).first() select(Message).where(Message.id == message_file.message_id, Message.app_id == app_id).limit(1)
) )
if not message: if not message:
raise FileAccessDeniedError("File access denied: not owned by requesting app") raise FileAccessDeniedError("File access denied: not owned by requesting app")
# Get the actual upload file record # Get the actual upload file record
upload_file = db.session.query(UploadFile).where(UploadFile.id == file_id).first() upload_file = db.session.get(UploadFile, file_id)
if not upload_file: if not upload_file:
raise FileNotFoundError("Upload file record not found") raise FileNotFoundError("Upload file record not found")
# Additional security: verify tenant isolation # Additional security: verify tenant isolation
app = db.session.query(App).where(App.id == app_id).first() app = db.session.get(App, app_id)
if app and upload_file.tenant_id != app.tenant_id: if app and upload_file.tenant_id != app.tenant_id:
raise FileAccessDeniedError("File access denied: tenant mismatch") raise FileAccessDeniedError("File access denied: tenant mismatch")

View File

@ -1,4 +1,5 @@
from flask_restx import Resource from flask_restx import Resource
from sqlalchemy import select
from werkzeug.exceptions import Forbidden from werkzeug.exceptions import Forbidden
from controllers.common.fields import Site as SiteResponse from controllers.common.fields import Site as SiteResponse
@ -28,7 +29,7 @@ class AppSiteApi(Resource):
Returns the site configuration for the application including theme, icons, and text. Returns the site configuration for the application including theme, icons, and text.
""" """
site = db.session.query(Site).where(Site.app_id == app_model.id).first() site = db.session.scalar(select(Site).where(Site.app_id == app_model.id).limit(1))
if not site: if not site:
raise Forbidden() raise Forbidden()

View File

@ -27,12 +27,12 @@ from core.errors.error import (
QuotaExceededError, QuotaExceededError,
) )
from core.helper.trace_id_helper import get_external_trace_id from core.helper.trace_id_helper import get_external_trace_id
from dify_graph.enums import WorkflowExecutionStatus
from dify_graph.graph_engine.manager import GraphEngineManager
from dify_graph.model_runtime.errors.invoke import InvokeError
from extensions.ext_database import db from extensions.ext_database import db
from extensions.ext_redis import redis_client from extensions.ext_redis import redis_client
from fields.workflow_app_log_fields import build_workflow_app_log_pagination_model from fields.workflow_app_log_fields import build_workflow_app_log_pagination_model
from graphon.enums import WorkflowExecutionStatus
from graphon.graph_engine.manager import GraphEngineManager
from graphon.model_runtime.errors.invoke import InvokeError
from libs import helper from libs import helper
from libs.helper import OptionalTimestampField, TimestampField from libs.helper import OptionalTimestampField, TimestampField
from models.model import App, AppMode, EndUser from models.model import App, AppMode, EndUser

View File

@ -14,10 +14,11 @@ from controllers.service_api.wraps import (
DatasetApiResource, DatasetApiResource,
cloud_edition_billing_rate_limit_check, cloud_edition_billing_rate_limit_check,
) )
from core.provider_manager import ProviderManager from core.plugin.impl.model_runtime_factory import create_plugin_provider_manager
from dify_graph.model_runtime.entities.model_entities import ModelType from core.rag.index_processor.constant.index_type import IndexTechniqueType
from fields.dataset_fields import dataset_detail_fields from fields.dataset_fields import dataset_detail_fields
from fields.tag_fields import DataSetTag from fields.tag_fields import DataSetTag
from graphon.model_runtime.entities.model_entities import ModelType
from libs.login import current_user from libs.login import current_user
from models.account import Account from models.account import Account
from models.dataset import DatasetPermissionEnum from models.dataset import DatasetPermissionEnum
@ -139,10 +140,10 @@ class DatasetListApi(DatasetApiResource):
query.page, query.limit, tenant_id, current_user, query.keyword, query.tag_ids, query.include_all query.page, query.limit, tenant_id, current_user, query.keyword, query.tag_ids, query.include_all
) )
# check embedding setting # check embedding setting
provider_manager = ProviderManager()
assert isinstance(current_user, Account) assert isinstance(current_user, Account)
cid = current_user.current_tenant_id cid = current_user.current_tenant_id
assert cid is not None assert cid is not None
provider_manager = create_plugin_provider_manager(tenant_id=cid)
configurations = provider_manager.get_configurations(tenant_id=cid) configurations = provider_manager.get_configurations(tenant_id=cid)
embedding_models = configurations.get_models(model_type=ModelType.TEXT_EMBEDDING, only_active=True) embedding_models = configurations.get_models(model_type=ModelType.TEXT_EMBEDDING, only_active=True)
@ -153,15 +154,20 @@ class DatasetListApi(DatasetApiResource):
data = marshal(datasets, dataset_detail_fields) data = marshal(datasets, dataset_detail_fields)
for item in data: for item in data:
if item["indexing_technique"] == "high_quality" and item["embedding_model_provider"]: if (
item["embedding_model_provider"] = str(ModelProviderID(item["embedding_model_provider"])) item["indexing_technique"] == IndexTechniqueType.HIGH_QUALITY # pyrefly: ignore[bad-index]
item_model = f"{item['embedding_model']}:{item['embedding_model_provider']}" and item["embedding_model_provider"] # pyrefly: ignore[bad-index]
):
item["embedding_model_provider"] = str( # pyrefly: ignore[unsupported-operation]
ModelProviderID(item["embedding_model_provider"]) # pyrefly: ignore[bad-index]
)
item_model = f"{item['embedding_model']}:{item['embedding_model_provider']}" # pyrefly: ignore[bad-index]
if item_model in model_names: if item_model in model_names:
item["embedding_available"] = True item["embedding_available"] = True # type: ignore
else: else:
item["embedding_available"] = False item["embedding_available"] = False # type: ignore
else: else:
item["embedding_available"] = True item["embedding_available"] = True # type: ignore
response = { response = {
"data": data, "data": data,
"has_more": len(datasets) == query.limit, "has_more": len(datasets) == query.limit,
@ -253,10 +259,10 @@ class DatasetApi(DatasetApiResource):
raise Forbidden(str(e)) raise Forbidden(str(e))
data = cast(dict[str, Any], marshal(dataset, dataset_detail_fields)) data = cast(dict[str, Any], marshal(dataset, dataset_detail_fields))
# check embedding setting # check embedding setting
provider_manager = ProviderManager()
assert isinstance(current_user, Account) assert isinstance(current_user, Account)
cid = current_user.current_tenant_id cid = current_user.current_tenant_id
assert cid is not None assert cid is not None
provider_manager = create_plugin_provider_manager(tenant_id=cid)
configurations = provider_manager.get_configurations(tenant_id=cid) configurations = provider_manager.get_configurations(tenant_id=cid)
embedding_models = configurations.get_models(model_type=ModelType.TEXT_EMBEDDING, only_active=True) embedding_models = configurations.get_models(model_type=ModelType.TEXT_EMBEDDING, only_active=True)
@ -265,7 +271,7 @@ class DatasetApi(DatasetApiResource):
for embedding_model in embedding_models: for embedding_model in embedding_models:
model_names.append(f"{embedding_model.model}:{embedding_model.provider.provider}") model_names.append(f"{embedding_model.model}:{embedding_model.provider.provider}")
if data.get("indexing_technique") == "high_quality": if data.get("indexing_technique") == IndexTechniqueType.HIGH_QUALITY:
item_model = f"{data.get('embedding_model')}:{data.get('embedding_model_provider')}" item_model = f"{data.get('embedding_model')}:{data.get('embedding_model_provider')}"
if item_model in model_names: if item_model in model_names:
data["embedding_available"] = True data["embedding_available"] = True
@ -315,7 +321,7 @@ class DatasetApi(DatasetApiResource):
# check embedding model setting # check embedding model setting
embedding_model_provider = payload.embedding_model_provider embedding_model_provider = payload.embedding_model_provider
embedding_model = payload.embedding_model embedding_model = payload.embedding_model
if payload.indexing_technique == "high_quality" or embedding_model_provider: if payload.indexing_technique == IndexTechniqueType.HIGH_QUALITY or embedding_model_provider:
if embedding_model_provider and embedding_model: if embedding_model_provider and embedding_model:
DatasetService.check_embedding_model_setting( DatasetService.check_embedding_model_setting(
dataset.tenant_id, embedding_model_provider, embedding_model dataset.tenant_id, embedding_model_provider, embedding_model

View File

@ -6,7 +6,7 @@ from uuid import UUID
from flask import request, send_file from flask import request, send_file
from flask_restx import marshal from flask_restx import marshal
from pydantic import BaseModel, Field, field_validator, model_validator from pydantic import BaseModel, Field, field_validator, model_validator
from sqlalchemy import desc, select from sqlalchemy import desc, func, select
from werkzeug.exceptions import Forbidden, NotFound from werkzeug.exceptions import Forbidden, NotFound
import services import services
@ -155,7 +155,9 @@ class DocumentAddByTextApi(DatasetApiResource):
dataset_id = str(dataset_id) dataset_id = str(dataset_id)
tenant_id = str(tenant_id) tenant_id = str(tenant_id)
dataset = db.session.query(Dataset).where(Dataset.tenant_id == tenant_id, Dataset.id == dataset_id).first() dataset = db.session.scalar(
select(Dataset).where(Dataset.tenant_id == tenant_id, Dataset.id == dataset_id).limit(1)
)
if not dataset: if not dataset:
raise ValueError("Dataset does not exist.") raise ValueError("Dataset does not exist.")
@ -238,7 +240,9 @@ class DocumentUpdateByTextApi(DatasetApiResource):
def post(self, tenant_id: str, dataset_id: UUID, document_id: UUID): def post(self, tenant_id: str, dataset_id: UUID, document_id: UUID):
"""Update document by text.""" """Update document by text."""
payload = DocumentTextUpdate.model_validate(service_api_ns.payload or {}) payload = DocumentTextUpdate.model_validate(service_api_ns.payload or {})
dataset = db.session.query(Dataset).where(Dataset.tenant_id == tenant_id, Dataset.id == str(dataset_id)).first() dataset = db.session.scalar(
select(Dataset).where(Dataset.tenant_id == tenant_id, Dataset.id == str(dataset_id)).limit(1)
)
args = payload.model_dump(exclude_none=True) args = payload.model_dump(exclude_none=True)
if not dataset: if not dataset:
raise ValueError("Dataset does not exist.") raise ValueError("Dataset does not exist.")
@ -315,7 +319,9 @@ class DocumentAddByFileApi(DatasetApiResource):
@cloud_edition_billing_rate_limit_check("knowledge", "dataset") @cloud_edition_billing_rate_limit_check("knowledge", "dataset")
def post(self, tenant_id, dataset_id): def post(self, tenant_id, dataset_id):
"""Create document by upload file.""" """Create document by upload file."""
dataset = db.session.query(Dataset).where(Dataset.tenant_id == tenant_id, Dataset.id == dataset_id).first() dataset = db.session.scalar(
select(Dataset).where(Dataset.tenant_id == tenant_id, Dataset.id == dataset_id).limit(1)
)
if not dataset: if not dataset:
raise ValueError("Dataset does not exist.") raise ValueError("Dataset does not exist.")
@ -425,7 +431,9 @@ class DocumentUpdateByFileApi(DatasetApiResource):
@cloud_edition_billing_rate_limit_check("knowledge", "dataset") @cloud_edition_billing_rate_limit_check("knowledge", "dataset")
def post(self, tenant_id, dataset_id, document_id): def post(self, tenant_id, dataset_id, document_id):
"""Update document by upload file.""" """Update document by upload file."""
dataset = db.session.query(Dataset).where(Dataset.tenant_id == tenant_id, Dataset.id == dataset_id).first() dataset = db.session.scalar(
select(Dataset).where(Dataset.tenant_id == tenant_id, Dataset.id == dataset_id).limit(1)
)
if not dataset: if not dataset:
raise ValueError("Dataset does not exist.") raise ValueError("Dataset does not exist.")
@ -515,7 +523,9 @@ class DocumentListApi(DatasetApiResource):
dataset_id = str(dataset_id) dataset_id = str(dataset_id)
tenant_id = str(tenant_id) tenant_id = str(tenant_id)
query_params = DocumentListQuery.model_validate(request.args.to_dict()) query_params = DocumentListQuery.model_validate(request.args.to_dict())
dataset = db.session.query(Dataset).where(Dataset.tenant_id == tenant_id, Dataset.id == dataset_id).first() dataset = db.session.scalar(
select(Dataset).where(Dataset.tenant_id == tenant_id, Dataset.id == dataset_id).limit(1)
)
if not dataset: if not dataset:
raise NotFound("Dataset not found.") raise NotFound("Dataset not found.")
@ -609,7 +619,9 @@ class DocumentIndexingStatusApi(DatasetApiResource):
batch = str(batch) batch = str(batch)
tenant_id = str(tenant_id) tenant_id = str(tenant_id)
# get dataset # get dataset
dataset = db.session.query(Dataset).where(Dataset.tenant_id == tenant_id, Dataset.id == dataset_id).first() dataset = db.session.scalar(
select(Dataset).where(Dataset.tenant_id == tenant_id, Dataset.id == dataset_id).limit(1)
)
if not dataset: if not dataset:
raise NotFound("Dataset not found.") raise NotFound("Dataset not found.")
# get documents # get documents
@ -619,20 +631,23 @@ class DocumentIndexingStatusApi(DatasetApiResource):
documents_status = [] documents_status = []
for document in documents: for document in documents:
completed_segments = ( completed_segments = (
db.session.query(DocumentSegment) db.session.scalar(
.where( select(func.count(DocumentSegment.id)).where(
DocumentSegment.completed_at.isnot(None), DocumentSegment.completed_at.isnot(None),
DocumentSegment.document_id == str(document.id), DocumentSegment.document_id == str(document.id),
DocumentSegment.status != SegmentStatus.RE_SEGMENT, DocumentSegment.status != SegmentStatus.RE_SEGMENT,
)
) )
.count() or 0
) )
total_segments = ( total_segments = (
db.session.query(DocumentSegment) db.session.scalar(
.where( select(func.count(DocumentSegment.id)).where(
DocumentSegment.document_id == str(document.id), DocumentSegment.status != SegmentStatus.RE_SEGMENT DocumentSegment.document_id == str(document.id),
DocumentSegment.status != SegmentStatus.RE_SEGMENT,
)
) )
.count() or 0
) )
# Create a dictionary with document attributes and additional fields # Create a dictionary with document attributes and additional fields
document_dict = { document_dict = {
@ -822,7 +837,9 @@ class DocumentApi(DatasetApiResource):
tenant_id = str(tenant_id) tenant_id = str(tenant_id)
# get dataset info # get dataset info
dataset = db.session.query(Dataset).where(Dataset.tenant_id == tenant_id, Dataset.id == dataset_id).first() dataset = db.session.scalar(
select(Dataset).where(Dataset.tenant_id == tenant_id, Dataset.id == dataset_id).limit(1)
)
if not dataset: if not dataset:
raise ValueError("Dataset does not exist.") raise ValueError("Dataset does not exist.")

View File

@ -3,6 +3,7 @@ from typing import Any
from flask import request from flask import request
from flask_restx import marshal from flask_restx import marshal
from pydantic import BaseModel, Field from pydantic import BaseModel, Field
from sqlalchemy import select
from werkzeug.exceptions import NotFound from werkzeug.exceptions import NotFound
from configs import dify_config from configs import dify_config
@ -17,9 +18,10 @@ from controllers.service_api.wraps import (
) )
from core.errors.error import LLMBadRequestError, ProviderTokenNotInitError from core.errors.error import LLMBadRequestError, ProviderTokenNotInitError
from core.model_manager import ModelManager from core.model_manager import ModelManager
from dify_graph.model_runtime.entities.model_entities import ModelType from core.rag.index_processor.constant.index_type import IndexTechniqueType
from extensions.ext_database import db from extensions.ext_database import db
from fields.segment_fields import child_chunk_fields, segment_fields from fields.segment_fields import child_chunk_fields, segment_fields
from graphon.model_runtime.entities.model_entities import ModelType
from libs.login import current_account_with_tenant from libs.login import current_account_with_tenant
from models.dataset import Dataset from models.dataset import Dataset
from services.dataset_service import DatasetService, DocumentService, SegmentService from services.dataset_service import DatasetService, DocumentService, SegmentService
@ -91,7 +93,9 @@ class SegmentApi(DatasetApiResource):
_, current_tenant_id = current_account_with_tenant() _, current_tenant_id = current_account_with_tenant()
"""Create single segment.""" """Create single segment."""
# check dataset # check dataset
dataset = db.session.query(Dataset).where(Dataset.tenant_id == tenant_id, Dataset.id == dataset_id).first() dataset = db.session.scalar(
select(Dataset).where(Dataset.tenant_id == tenant_id, Dataset.id == dataset_id).limit(1)
)
if not dataset: if not dataset:
raise NotFound("Dataset not found.") raise NotFound("Dataset not found.")
# check document # check document
@ -103,9 +107,9 @@ class SegmentApi(DatasetApiResource):
if not document.enabled: if not document.enabled:
raise NotFound("Document is disabled.") raise NotFound("Document is disabled.")
# check embedding model setting # check embedding model setting
if dataset.indexing_technique == "high_quality": if dataset.indexing_technique == IndexTechniqueType.HIGH_QUALITY:
try: try:
model_manager = ModelManager() model_manager = ModelManager.for_tenant(tenant_id=current_tenant_id)
model_manager.get_model_instance( model_manager.get_model_instance(
tenant_id=current_tenant_id, tenant_id=current_tenant_id,
provider=dataset.embedding_model_provider, provider=dataset.embedding_model_provider,
@ -149,7 +153,9 @@ class SegmentApi(DatasetApiResource):
# check dataset # check dataset
page = request.args.get("page", default=1, type=int) page = request.args.get("page", default=1, type=int)
limit = request.args.get("limit", default=20, type=int) limit = request.args.get("limit", default=20, type=int)
dataset = db.session.query(Dataset).where(Dataset.tenant_id == tenant_id, Dataset.id == dataset_id).first() dataset = db.session.scalar(
select(Dataset).where(Dataset.tenant_id == tenant_id, Dataset.id == dataset_id).limit(1)
)
if not dataset: if not dataset:
raise NotFound("Dataset not found.") raise NotFound("Dataset not found.")
# check document # check document
@ -157,9 +163,9 @@ class SegmentApi(DatasetApiResource):
if not document: if not document:
raise NotFound("Document not found.") raise NotFound("Document not found.")
# check embedding model setting # check embedding model setting
if dataset.indexing_technique == "high_quality": if dataset.indexing_technique == IndexTechniqueType.HIGH_QUALITY:
try: try:
model_manager = ModelManager() model_manager = ModelManager.for_tenant(tenant_id=current_tenant_id)
model_manager.get_model_instance( model_manager.get_model_instance(
tenant_id=current_tenant_id, tenant_id=current_tenant_id,
provider=dataset.embedding_model_provider, provider=dataset.embedding_model_provider,
@ -219,7 +225,9 @@ class DatasetSegmentApi(DatasetApiResource):
def delete(self, tenant_id: str, dataset_id: str, document_id: str, segment_id: str): def delete(self, tenant_id: str, dataset_id: str, document_id: str, segment_id: str):
_, current_tenant_id = current_account_with_tenant() _, current_tenant_id = current_account_with_tenant()
# check dataset # check dataset
dataset = db.session.query(Dataset).where(Dataset.tenant_id == tenant_id, Dataset.id == dataset_id).first() dataset = db.session.scalar(
select(Dataset).where(Dataset.tenant_id == tenant_id, Dataset.id == dataset_id).limit(1)
)
if not dataset: if not dataset:
raise NotFound("Dataset not found.") raise NotFound("Dataset not found.")
# check user's model setting # check user's model setting
@ -253,7 +261,9 @@ class DatasetSegmentApi(DatasetApiResource):
def post(self, tenant_id: str, dataset_id: str, document_id: str, segment_id: str): def post(self, tenant_id: str, dataset_id: str, document_id: str, segment_id: str):
_, current_tenant_id = current_account_with_tenant() _, current_tenant_id = current_account_with_tenant()
# check dataset # check dataset
dataset = db.session.query(Dataset).where(Dataset.tenant_id == tenant_id, Dataset.id == dataset_id).first() dataset = db.session.scalar(
select(Dataset).where(Dataset.tenant_id == tenant_id, Dataset.id == dataset_id).limit(1)
)
if not dataset: if not dataset:
raise NotFound("Dataset not found.") raise NotFound("Dataset not found.")
# check user's model setting # check user's model setting
@ -262,10 +272,10 @@ class DatasetSegmentApi(DatasetApiResource):
document = DocumentService.get_document(dataset_id, document_id) document = DocumentService.get_document(dataset_id, document_id)
if not document: if not document:
raise NotFound("Document not found.") raise NotFound("Document not found.")
if dataset.indexing_technique == "high_quality": if dataset.indexing_technique == IndexTechniqueType.HIGH_QUALITY:
# check embedding model setting # check embedding model setting
try: try:
model_manager = ModelManager() model_manager = ModelManager.for_tenant(tenant_id=current_tenant_id)
model_manager.get_model_instance( model_manager.get_model_instance(
tenant_id=current_tenant_id, tenant_id=current_tenant_id,
provider=dataset.embedding_model_provider, provider=dataset.embedding_model_provider,
@ -300,7 +310,9 @@ class DatasetSegmentApi(DatasetApiResource):
def get(self, tenant_id: str, dataset_id: str, document_id: str, segment_id: str): def get(self, tenant_id: str, dataset_id: str, document_id: str, segment_id: str):
_, current_tenant_id = current_account_with_tenant() _, current_tenant_id = current_account_with_tenant()
# check dataset # check dataset
dataset = db.session.query(Dataset).where(Dataset.tenant_id == tenant_id, Dataset.id == dataset_id).first() dataset = db.session.scalar(
select(Dataset).where(Dataset.tenant_id == tenant_id, Dataset.id == dataset_id).limit(1)
)
if not dataset: if not dataset:
raise NotFound("Dataset not found.") raise NotFound("Dataset not found.")
# check user's model setting # check user's model setting
@ -343,7 +355,9 @@ class ChildChunkApi(DatasetApiResource):
_, current_tenant_id = current_account_with_tenant() _, current_tenant_id = current_account_with_tenant()
"""Create child chunk.""" """Create child chunk."""
# check dataset # check dataset
dataset = db.session.query(Dataset).where(Dataset.tenant_id == tenant_id, Dataset.id == dataset_id).first() dataset = db.session.scalar(
select(Dataset).where(Dataset.tenant_id == tenant_id, Dataset.id == dataset_id).limit(1)
)
if not dataset: if not dataset:
raise NotFound("Dataset not found.") raise NotFound("Dataset not found.")
@ -358,9 +372,9 @@ class ChildChunkApi(DatasetApiResource):
raise NotFound("Segment not found.") raise NotFound("Segment not found.")
# check embedding model setting # check embedding model setting
if dataset.indexing_technique == "high_quality": if dataset.indexing_technique == IndexTechniqueType.HIGH_QUALITY:
try: try:
model_manager = ModelManager() model_manager = ModelManager.for_tenant(tenant_id=current_tenant_id)
model_manager.get_model_instance( model_manager.get_model_instance(
tenant_id=current_tenant_id, tenant_id=current_tenant_id,
provider=dataset.embedding_model_provider, provider=dataset.embedding_model_provider,
@ -401,7 +415,9 @@ class ChildChunkApi(DatasetApiResource):
_, current_tenant_id = current_account_with_tenant() _, current_tenant_id = current_account_with_tenant()
"""Get child chunks.""" """Get child chunks."""
# check dataset # check dataset
dataset = db.session.query(Dataset).where(Dataset.tenant_id == tenant_id, Dataset.id == dataset_id).first() dataset = db.session.scalar(
select(Dataset).where(Dataset.tenant_id == tenant_id, Dataset.id == dataset_id).limit(1)
)
if not dataset: if not dataset:
raise NotFound("Dataset not found.") raise NotFound("Dataset not found.")
@ -467,7 +483,9 @@ class DatasetChildChunkApi(DatasetApiResource):
_, current_tenant_id = current_account_with_tenant() _, current_tenant_id = current_account_with_tenant()
"""Delete child chunk.""" """Delete child chunk."""
# check dataset # check dataset
dataset = db.session.query(Dataset).where(Dataset.tenant_id == tenant_id, Dataset.id == dataset_id).first() dataset = db.session.scalar(
select(Dataset).where(Dataset.tenant_id == tenant_id, Dataset.id == dataset_id).limit(1)
)
if not dataset: if not dataset:
raise NotFound("Dataset not found.") raise NotFound("Dataset not found.")
@ -526,7 +544,9 @@ class DatasetChildChunkApi(DatasetApiResource):
_, current_tenant_id = current_account_with_tenant() _, current_tenant_id = current_account_with_tenant()
"""Update child chunk.""" """Update child chunk."""
# check dataset # check dataset
dataset = db.session.query(Dataset).where(Dataset.tenant_id == tenant_id, Dataset.id == dataset_id).first() dataset = db.session.scalar(
select(Dataset).where(Dataset.tenant_id == tenant_id, Dataset.id == dataset_id).limit(1)
)
if not dataset: if not dataset:
raise NotFound("Dataset not found.") raise NotFound("Dataset not found.")

View File

@ -3,7 +3,7 @@ from flask_restx import Resource
from controllers.service_api import service_api_ns from controllers.service_api import service_api_ns
from controllers.service_api.wraps import validate_dataset_token from controllers.service_api.wraps import validate_dataset_token
from dify_graph.model_runtime.utils.encoders import jsonable_encoder from graphon.model_runtime.utils.encoders import jsonable_encoder
from services.model_provider_service import ModelProviderService from services.model_provider_service import ModelProviderService

View File

@ -9,6 +9,7 @@ from flask import current_app, request
from flask_login import user_logged_in from flask_login import user_logged_in
from flask_restx import Resource from flask_restx import Resource
from pydantic import BaseModel from pydantic import BaseModel
from sqlalchemy import select
from werkzeug.exceptions import Forbidden, NotFound, Unauthorized from werkzeug.exceptions import Forbidden, NotFound, Unauthorized
from enums.cloud_plan import CloudPlan from enums.cloud_plan import CloudPlan
@ -62,7 +63,7 @@ def validate_app_token(
def decorated_view(*args: P.args, **kwargs: P.kwargs) -> R: def decorated_view(*args: P.args, **kwargs: P.kwargs) -> R:
api_token = validate_and_get_api_token("app") api_token = validate_and_get_api_token("app")
app_model = db.session.query(App).where(App.id == api_token.app_id).first() app_model = db.session.get(App, api_token.app_id)
if not app_model: if not app_model:
raise Forbidden("The app no longer exists.") raise Forbidden("The app no longer exists.")
@ -72,7 +73,7 @@ def validate_app_token(
if not app_model.enable_api: if not app_model.enable_api:
raise Forbidden("The app's API service has been disabled.") raise Forbidden("The app's API service has been disabled.")
tenant = db.session.query(Tenant).where(Tenant.id == app_model.tenant_id).first() tenant = db.session.get(Tenant, app_model.tenant_id)
if tenant is None: if tenant is None:
raise ValueError("Tenant does not exist.") raise ValueError("Tenant does not exist.")
if tenant.status == TenantStatus.ARCHIVE: if tenant.status == TenantStatus.ARCHIVE:
@ -106,8 +107,8 @@ def validate_app_token(
else: else:
# For service API without end-user context, ensure an Account is logged in # For service API without end-user context, ensure an Account is logged in
# so services relying on current_account_with_tenant() work correctly. # so services relying on current_account_with_tenant() work correctly.
tenant_owner_info = ( tenant_owner_info = db.session.execute(
db.session.query(Tenant, Account) select(Tenant, Account)
.join(TenantAccountJoin, Tenant.id == TenantAccountJoin.tenant_id) .join(TenantAccountJoin, Tenant.id == TenantAccountJoin.tenant_id)
.join(Account, TenantAccountJoin.account_id == Account.id) .join(Account, TenantAccountJoin.account_id == Account.id)
.where( .where(
@ -115,8 +116,7 @@ def validate_app_token(
TenantAccountJoin.role == "owner", TenantAccountJoin.role == "owner",
Tenant.status == TenantStatus.NORMAL, Tenant.status == TenantStatus.NORMAL,
) )
.one_or_none() ).one_or_none()
)
if tenant_owner_info: if tenant_owner_info:
tenant_model, account = tenant_owner_info tenant_model, account = tenant_owner_info
@ -277,29 +277,28 @@ def validate_dataset_token(
# Validate dataset if dataset_id is provided # Validate dataset if dataset_id is provided
if dataset_id: if dataset_id:
dataset_id = str(dataset_id) dataset_id = str(dataset_id)
dataset = ( dataset = db.session.scalar(
db.session.query(Dataset) select(Dataset)
.where( .where(
Dataset.id == dataset_id, Dataset.id == dataset_id,
Dataset.tenant_id == api_token.tenant_id, Dataset.tenant_id == api_token.tenant_id,
) )
.first() .limit(1)
) )
if not dataset: if not dataset:
raise NotFound("Dataset not found.") raise NotFound("Dataset not found.")
if not dataset.enable_api: if not dataset.enable_api:
raise Forbidden("Dataset api access is not enabled.") raise Forbidden("Dataset api access is not enabled.")
tenant_account_join = ( tenant_account_join = db.session.execute(
db.session.query(Tenant, TenantAccountJoin) select(Tenant, TenantAccountJoin)
.where(Tenant.id == api_token.tenant_id) .where(Tenant.id == api_token.tenant_id)
.where(TenantAccountJoin.tenant_id == Tenant.id) .where(TenantAccountJoin.tenant_id == Tenant.id)
.where(TenantAccountJoin.role.in_(["owner"])) .where(TenantAccountJoin.role.in_(["owner"]))
.where(Tenant.status == TenantStatus.NORMAL) .where(Tenant.status == TenantStatus.NORMAL)
.one_or_none() ).one_or_none() # TODO: only owner information is required, so only one is returned.
) # TODO: only owner information is required, so only one is returned.
if tenant_account_join: if tenant_account_join:
tenant, ta = tenant_account_join tenant, ta = tenant_account_join
account = db.session.query(Account).where(Account.id == ta.account_id).first() account = db.session.get(Account, ta.account_id)
# Login admin # Login admin
if account: if account:
account.current_tenant = tenant account.current_tenant = tenant
@ -360,7 +359,9 @@ class DatasetApiResource(Resource):
method_decorators = [validate_dataset_token] method_decorators = [validate_dataset_token]
def get_dataset(self, dataset_id: str, tenant_id: str) -> Dataset: def get_dataset(self, dataset_id: str, tenant_id: str) -> Dataset:
dataset = db.session.query(Dataset).where(Dataset.id == dataset_id, Dataset.tenant_id == tenant_id).first() dataset = db.session.scalar(
select(Dataset).where(Dataset.id == dataset_id, Dataset.tenant_id == tenant_id).limit(1)
)
if not dataset: if not dataset:
raise NotFound("Dataset not found.") raise NotFound("Dataset not found.")

View File

@ -20,7 +20,7 @@ from controllers.web.error import (
) )
from controllers.web.wraps import WebApiResource from controllers.web.wraps import WebApiResource
from core.errors.error import ModelCurrentlyNotSupportError, ProviderTokenNotInitError, QuotaExceededError from core.errors.error import ModelCurrentlyNotSupportError, ProviderTokenNotInitError, QuotaExceededError
from dify_graph.model_runtime.errors.invoke import InvokeError from graphon.model_runtime.errors.invoke import InvokeError
from libs.helper import uuid_value from libs.helper import uuid_value
from models.model import App from models.model import App
from services.audio_service import AudioService from services.audio_service import AudioService

View File

@ -25,7 +25,7 @@ from core.errors.error import (
ProviderTokenNotInitError, ProviderTokenNotInitError,
QuotaExceededError, QuotaExceededError,
) )
from dify_graph.model_runtime.errors.invoke import InvokeError from graphon.model_runtime.errors.invoke import InvokeError
from libs import helper from libs import helper
from libs.helper import uuid_value from libs.helper import uuid_value
from models.model import AppMode from models.model import AppMode

View File

@ -20,9 +20,9 @@ from controllers.web.error import (
from controllers.web.wraps import WebApiResource from controllers.web.wraps import WebApiResource
from core.app.entities.app_invoke_entities import InvokeFrom from core.app.entities.app_invoke_entities import InvokeFrom
from core.errors.error import ModelCurrentlyNotSupportError, ProviderTokenNotInitError, QuotaExceededError from core.errors.error import ModelCurrentlyNotSupportError, ProviderTokenNotInitError, QuotaExceededError
from dify_graph.model_runtime.errors.invoke import InvokeError
from fields.conversation_fields import ResultResponse from fields.conversation_fields import ResultResponse
from fields.message_fields import SuggestedQuestionsResponse, WebMessageInfiniteScrollPagination, WebMessageListItem from fields.message_fields import SuggestedQuestionsResponse, WebMessageInfiniteScrollPagination, WebMessageListItem
from graphon.model_runtime.errors.invoke import InvokeError
from libs import helper from libs import helper
from libs.helper import uuid_value from libs.helper import uuid_value
from models.enums import FeedbackRating from models.enums import FeedbackRating

View File

@ -11,9 +11,9 @@ from controllers.common.errors import (
UnsupportedFileTypeError, UnsupportedFileTypeError,
) )
from core.helper import ssrf_proxy from core.helper import ssrf_proxy
from dify_graph.file import helpers as file_helpers
from extensions.ext_database import db from extensions.ext_database import db
from fields.file_fields import FileWithSignedUrl, RemoteFileInfo from fields.file_fields import FileWithSignedUrl, RemoteFileInfo
from graphon.file import helpers as file_helpers
from services.file_service import FileService from services.file_service import FileService
from ..common.schema import register_schema_models from ..common.schema import register_schema_models

View File

@ -22,9 +22,9 @@ from core.errors.error import (
ProviderTokenNotInitError, ProviderTokenNotInitError,
QuotaExceededError, QuotaExceededError,
) )
from dify_graph.graph_engine.manager import GraphEngineManager
from dify_graph.model_runtime.errors.invoke import InvokeError
from extensions.ext_redis import redis_client from extensions.ext_redis import redis_client
from graphon.graph_engine.manager import GraphEngineManager
from graphon.model_runtime.errors.invoke import InvokeError
from libs import helper from libs import helper
from models.model import App, AppMode, EndUser from models.model import App, AppMode, EndUser
from services.app_generate_service import AppGenerateService from services.app_generate_service import AppGenerateService

View File

@ -15,6 +15,7 @@ from core.app.entities.app_invoke_entities import (
AgentChatAppGenerateEntity, AgentChatAppGenerateEntity,
ModelConfigWithCredentialsEntity, ModelConfigWithCredentialsEntity,
) )
from core.app.file_access import DatabaseFileAccessController
from core.callback_handler.agent_tool_callback_handler import DifyAgentCallbackHandler from core.callback_handler.agent_tool_callback_handler import DifyAgentCallbackHandler
from core.callback_handler.index_tool_callback_handler import DatasetIndexToolCallbackHandler from core.callback_handler.index_tool_callback_handler import DatasetIndexToolCallbackHandler
from core.memory.token_buffer_memory import TokenBufferMemory from core.memory.token_buffer_memory import TokenBufferMemory
@ -26,8 +27,10 @@ from core.tools.entities.tool_entities import (
) )
from core.tools.tool_manager import ToolManager from core.tools.tool_manager import ToolManager
from core.tools.utils.dataset_retriever_tool import DatasetRetrieverTool from core.tools.utils.dataset_retriever_tool import DatasetRetrieverTool
from dify_graph.file import file_manager from extensions.ext_database import db
from dify_graph.model_runtime.entities import ( from factories import file_factory
from graphon.file import file_manager
from graphon.model_runtime.entities import (
AssistantPromptMessage, AssistantPromptMessage,
LLMUsage, LLMUsage,
PromptMessage, PromptMessage,
@ -37,15 +40,14 @@ from dify_graph.model_runtime.entities import (
ToolPromptMessage, ToolPromptMessage,
UserPromptMessage, UserPromptMessage,
) )
from dify_graph.model_runtime.entities.message_entities import ImagePromptMessageContent, PromptMessageContentUnionTypes from graphon.model_runtime.entities.message_entities import ImagePromptMessageContent, PromptMessageContentUnionTypes
from dify_graph.model_runtime.entities.model_entities import ModelFeature from graphon.model_runtime.entities.model_entities import ModelFeature
from dify_graph.model_runtime.model_providers.__base.large_language_model import LargeLanguageModel from graphon.model_runtime.model_providers.__base.large_language_model import LargeLanguageModel
from extensions.ext_database import db
from factories import file_factory
from models.enums import CreatorUserRole from models.enums import CreatorUserRole
from models.model import Conversation, Message, MessageAgentThought, MessageFile from models.model import Conversation, Message, MessageAgentThought, MessageFile
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
_file_access_controller = DatabaseFileAccessController()
class BaseAgentRunner(AppRunner): class BaseAgentRunner(AppRunner):
@ -138,6 +140,7 @@ class BaseAgentRunner(AppRunner):
tenant_id=self.tenant_id, tenant_id=self.tenant_id,
app_id=self.app_config.app_id, app_id=self.app_config.app_id,
agent_tool=tool, agent_tool=tool,
user_id=self.user_id,
invoke_from=self.application_generate_entity.invoke_from, invoke_from=self.application_generate_entity.invoke_from,
) )
assert tool_entity.entity.description assert tool_entity.entity.description
@ -524,7 +527,10 @@ class BaseAgentRunner(AppRunner):
image_detail_config = image_detail_config or ImagePromptMessageContent.DETAIL.LOW image_detail_config = image_detail_config or ImagePromptMessageContent.DETAIL.LOW
file_objs = file_factory.build_from_message_files( file_objs = file_factory.build_from_message_files(
message_files=files, tenant_id=self.tenant_id, config=file_extra_config message_files=files,
tenant_id=self.tenant_id,
config=file_extra_config,
access_controller=_file_access_controller,
) )
if not file_objs: if not file_objs:
return UserPromptMessage(content=message.query) return UserPromptMessage(content=message.query)

View File

@ -15,8 +15,8 @@ from core.prompt.agent_history_prompt_transform import AgentHistoryPromptTransfo
from core.tools.__base.tool import Tool from core.tools.__base.tool import Tool
from core.tools.entities.tool_entities import ToolInvokeMeta from core.tools.entities.tool_entities import ToolInvokeMeta
from core.tools.tool_engine import ToolEngine from core.tools.tool_engine import ToolEngine
from dify_graph.model_runtime.entities.llm_entities import LLMResult, LLMResultChunk, LLMResultChunkDelta, LLMUsage from graphon.model_runtime.entities.llm_entities import LLMResult, LLMResultChunk, LLMResultChunkDelta, LLMUsage
from dify_graph.model_runtime.entities.message_entities import ( from graphon.model_runtime.entities.message_entities import (
AssistantPromptMessage, AssistantPromptMessage,
PromptMessage, PromptMessage,
PromptMessageTool, PromptMessageTool,
@ -122,7 +122,6 @@ class CotAgentRunner(BaseAgentRunner, ABC):
tools=[], tools=[],
stop=app_generate_entity.model_conf.stop, stop=app_generate_entity.model_conf.stop,
stream=True, stream=True,
user=self.user_id,
callbacks=[], callbacks=[],
) )

View File

@ -1,16 +1,16 @@
import json import json
from core.agent.cot_agent_runner import CotAgentRunner from core.agent.cot_agent_runner import CotAgentRunner
from dify_graph.file import file_manager from graphon.file import file_manager
from dify_graph.model_runtime.entities import ( from graphon.model_runtime.entities import (
AssistantPromptMessage, AssistantPromptMessage,
PromptMessage, PromptMessage,
SystemPromptMessage, SystemPromptMessage,
TextPromptMessageContent, TextPromptMessageContent,
UserPromptMessage, UserPromptMessage,
) )
from dify_graph.model_runtime.entities.message_entities import ImagePromptMessageContent, PromptMessageContentUnionTypes from graphon.model_runtime.entities.message_entities import ImagePromptMessageContent, PromptMessageContentUnionTypes
from dify_graph.model_runtime.utils.encoders import jsonable_encoder from graphon.model_runtime.utils.encoders import jsonable_encoder
class CotChatAgentRunner(CotAgentRunner): class CotChatAgentRunner(CotAgentRunner):

View File

@ -1,13 +1,13 @@
import json import json
from core.agent.cot_agent_runner import CotAgentRunner from core.agent.cot_agent_runner import CotAgentRunner
from dify_graph.model_runtime.entities.message_entities import ( from graphon.model_runtime.entities.message_entities import (
AssistantPromptMessage, AssistantPromptMessage,
PromptMessage, PromptMessage,
TextPromptMessageContent, TextPromptMessageContent,
UserPromptMessage, UserPromptMessage,
) )
from dify_graph.model_runtime.utils.encoders import jsonable_encoder from graphon.model_runtime.utils.encoders import jsonable_encoder
class CotCompletionAgentRunner(CotAgentRunner): class CotCompletionAgentRunner(CotAgentRunner):

View File

@ -11,8 +11,8 @@ from core.app.entities.queue_entities import QueueAgentThoughtEvent, QueueMessag
from core.prompt.agent_history_prompt_transform import AgentHistoryPromptTransform from core.prompt.agent_history_prompt_transform import AgentHistoryPromptTransform
from core.tools.entities.tool_entities import ToolInvokeMeta from core.tools.entities.tool_entities import ToolInvokeMeta
from core.tools.tool_engine import ToolEngine from core.tools.tool_engine import ToolEngine
from dify_graph.file import file_manager from graphon.file import file_manager
from dify_graph.model_runtime.entities import ( from graphon.model_runtime.entities import (
AssistantPromptMessage, AssistantPromptMessage,
LLMResult, LLMResult,
LLMResultChunk, LLMResultChunk,
@ -25,7 +25,7 @@ from dify_graph.model_runtime.entities import (
ToolPromptMessage, ToolPromptMessage,
UserPromptMessage, UserPromptMessage,
) )
from dify_graph.model_runtime.entities.message_entities import ImagePromptMessageContent, PromptMessageContentUnionTypes from graphon.model_runtime.entities.message_entities import ImagePromptMessageContent, PromptMessageContentUnionTypes
from models.model import Message from models.model import Message
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@ -96,7 +96,6 @@ class FunctionCallAgentRunner(BaseAgentRunner):
tools=prompt_messages_tools, tools=prompt_messages_tools,
stop=app_generate_entity.model_conf.stop, stop=app_generate_entity.model_conf.stop,
stream=self.stream_tool_call, stream=self.stream_tool_call,
user=self.user_id,
callbacks=[], callbacks=[],
) )

View File

@ -4,7 +4,7 @@ from collections.abc import Generator
from typing import Union from typing import Union
from core.agent.entities import AgentScratchpadUnit from core.agent.entities import AgentScratchpadUnit
from dify_graph.model_runtime.entities.llm_entities import LLMResultChunk from graphon.model_runtime.entities.llm_entities import LLMResultChunk
class CotAgentOutputParser: class CotAgentOutputParser:

View File

@ -4,10 +4,10 @@ from core.app.app_config.entities import EasyUIBasedAppConfig
from core.app.entities.app_invoke_entities import ModelConfigWithCredentialsEntity from core.app.entities.app_invoke_entities import ModelConfigWithCredentialsEntity
from core.entities.model_entities import ModelStatus from core.entities.model_entities import ModelStatus
from core.errors.error import ModelCurrentlyNotSupportError, ProviderTokenNotInitError, QuotaExceededError from core.errors.error import ModelCurrentlyNotSupportError, ProviderTokenNotInitError, QuotaExceededError
from core.provider_manager import ProviderManager from core.plugin.impl.model_runtime_factory import create_plugin_provider_manager
from dify_graph.model_runtime.entities.llm_entities import LLMMode from graphon.model_runtime.entities.llm_entities import LLMMode
from dify_graph.model_runtime.entities.model_entities import ModelPropertyKey, ModelType from graphon.model_runtime.entities.model_entities import ModelPropertyKey, ModelType
from dify_graph.model_runtime.model_providers.__base.large_language_model import LargeLanguageModel from graphon.model_runtime.model_providers.__base.large_language_model import LargeLanguageModel
class ModelConfigConverter: class ModelConfigConverter:
@ -21,7 +21,7 @@ class ModelConfigConverter:
""" """
model_config = app_config.model model_config = app_config.model
provider_manager = ProviderManager() provider_manager = create_plugin_provider_manager(tenant_id=app_config.tenant_id)
provider_model_bundle = provider_manager.get_provider_model_bundle( provider_model_bundle = provider_manager.get_provider_model_bundle(
tenant_id=app_config.tenant_id, provider=model_config.provider, model_type=ModelType.LLM tenant_id=app_config.tenant_id, provider=model_config.provider, model_type=ModelType.LLM
) )

View File

@ -2,9 +2,8 @@ from collections.abc import Mapping
from typing import Any from typing import Any
from core.app.app_config.entities import ModelConfigEntity from core.app.app_config.entities import ModelConfigEntity
from core.provider_manager import ProviderManager from core.plugin.impl.model_runtime_factory import create_plugin_model_assembly
from dify_graph.model_runtime.entities.model_entities import ModelPropertyKey, ModelType from graphon.model_runtime.entities.model_entities import ModelPropertyKey, ModelType
from dify_graph.model_runtime.model_providers.model_provider_factory import ModelProviderFactory
from models.model import AppModelConfigDict from models.model import AppModelConfigDict
from models.provider_ids import ModelProviderID from models.provider_ids import ModelProviderID
@ -54,9 +53,12 @@ class ModelConfigManager:
if not isinstance(config["model"], dict): if not isinstance(config["model"], dict):
raise ValueError("model must be of object type") raise ValueError("model must be of object type")
# Keep provider discovery and provider-backed model listing on the same
# request-scoped runtime so caller scope and provider caches stay aligned.
assembly = create_plugin_model_assembly(tenant_id=tenant_id)
# model.provider # model.provider
model_provider_factory = ModelProviderFactory(tenant_id) provider_entities = assembly.model_provider_factory.get_providers()
provider_entities = model_provider_factory.get_providers()
model_provider_names = [provider.provider for provider in provider_entities] model_provider_names = [provider.provider for provider in provider_entities]
if "provider" not in config["model"]: if "provider" not in config["model"]:
raise ValueError(f"model.provider is required and must be in {str(model_provider_names)}") raise ValueError(f"model.provider is required and must be in {str(model_provider_names)}")
@ -71,8 +73,7 @@ class ModelConfigManager:
if "name" not in config["model"]: if "name" not in config["model"]:
raise ValueError("model.name is required") raise ValueError("model.name is required")
provider_manager = ProviderManager() models = assembly.provider_manager.get_configurations(tenant_id).get_models(
models = provider_manager.get_configurations(tenant_id).get_models(
provider=config["model"]["provider"], model_type=ModelType.LLM provider=config["model"]["provider"], model_type=ModelType.LLM
) )

View File

@ -7,7 +7,7 @@ from core.app.app_config.entities import (
PromptTemplateEntity, PromptTemplateEntity,
) )
from core.prompt.simple_prompt_transform import ModelMode from core.prompt.simple_prompt_transform import ModelMode
from dify_graph.model_runtime.entities.message_entities import PromptMessageRole from graphon.model_runtime.entities.message_entities import PromptMessageRole
from models.model import AppMode, AppModelConfigDict from models.model import AppMode, AppModelConfigDict

View File

@ -3,7 +3,7 @@ from typing import cast
from core.app.app_config.entities import ExternalDataVariableEntity from core.app.app_config.entities import ExternalDataVariableEntity
from core.external_data_tool.factory import ExternalDataToolFactory from core.external_data_tool.factory import ExternalDataToolFactory
from dify_graph.variables.input_entities import VariableEntity, VariableEntityType from graphon.variables.input_entities import VariableEntity, VariableEntityType
from models.model import AppModelConfigDict from models.model import AppModelConfigDict
_ALLOWED_VARIABLE_ENTITY_TYPE = frozenset( _ALLOWED_VARIABLE_ENTITY_TYPE = frozenset(

View File

@ -5,10 +5,10 @@ from typing import Any, Literal
from pydantic import BaseModel, Field from pydantic import BaseModel, Field
from core.rag.data_post_processor.data_post_processor import RerankingModelDict, WeightsDict from core.rag.data_post_processor.data_post_processor import RerankingModelDict, WeightsDict
from dify_graph.file import FileUploadConfig from graphon.file import FileUploadConfig
from dify_graph.model_runtime.entities.llm_entities import LLMMode from graphon.model_runtime.entities.llm_entities import LLMMode
from dify_graph.model_runtime.entities.message_entities import PromptMessageRole from graphon.model_runtime.entities.message_entities import PromptMessageRole
from dify_graph.variables.input_entities import VariableEntity as WorkflowVariableEntity from graphon.variables.input_entities import VariableEntity as WorkflowVariableEntity
from models.model import AppMode from models.model import AppMode

View File

@ -2,7 +2,7 @@ from collections.abc import Mapping
from typing import Any from typing import Any
from constants import DEFAULT_FILE_NUMBER_LIMITS from constants import DEFAULT_FILE_NUMBER_LIMITS
from dify_graph.file import FileUploadConfig from graphon.file import FileUploadConfig
class FileUploadConfigManager: class FileUploadConfigManager:

View File

@ -1,7 +1,7 @@
import re import re
from core.app.app_config.entities import RagPipelineVariableEntity from core.app.app_config.entities import RagPipelineVariableEntity
from dify_graph.variables.input_entities import VariableEntity from graphon.variables.input_entities import VariableEntity
from models.workflow import Workflow from models.workflow import Workflow

View File

@ -24,6 +24,7 @@ from core.app.apps.advanced_chat.app_runner import AdvancedChatAppRunner
from core.app.apps.advanced_chat.generate_response_converter import AdvancedChatAppGenerateResponseConverter from core.app.apps.advanced_chat.generate_response_converter import AdvancedChatAppGenerateResponseConverter
from core.app.apps.advanced_chat.generate_task_pipeline import AdvancedChatAppGenerateTaskPipeline from core.app.apps.advanced_chat.generate_task_pipeline import AdvancedChatAppGenerateTaskPipeline
from core.app.apps.base_app_queue_manager import AppQueueManager, PublishFrom from core.app.apps.base_app_queue_manager import AppQueueManager, PublishFrom
from core.app.apps.draft_variable_saver import DraftVariableSaverFactory
from core.app.apps.exc import GenerateTaskStoppedError from core.app.apps.exc import GenerateTaskStoppedError
from core.app.apps.message_based_app_generator import MessageBasedAppGenerator from core.app.apps.message_based_app_generator import MessageBasedAppGenerator
from core.app.apps.message_based_app_queue_manager import MessageBasedAppQueueManager from core.app.apps.message_based_app_queue_manager import MessageBasedAppQueueManager
@ -34,17 +35,13 @@ from core.helper.trace_id_helper import extract_external_trace_id_from_args
from core.ops.ops_trace_manager import TraceQueueManager from core.ops.ops_trace_manager import TraceQueueManager
from core.prompt.utils.get_thread_messages_length import get_thread_messages_length from core.prompt.utils.get_thread_messages_length import get_thread_messages_length
from core.repositories import DifyCoreRepositoryFactory from core.repositories import DifyCoreRepositoryFactory
from dify_graph.graph_engine.layers.base import GraphEngineLayer from core.repositories.factory import WorkflowExecutionRepository, WorkflowNodeExecutionRepository
from dify_graph.model_runtime.errors.invoke import InvokeAuthorizationError
from dify_graph.repositories.draft_variable_repository import (
DraftVariableSaverFactory,
)
from dify_graph.repositories.workflow_execution_repository import WorkflowExecutionRepository
from dify_graph.repositories.workflow_node_execution_repository import WorkflowNodeExecutionRepository
from dify_graph.runtime import GraphRuntimeState
from dify_graph.variable_loader import DUMMY_VARIABLE_LOADER, VariableLoader
from extensions.ext_database import db from extensions.ext_database import db
from factories import file_factory from factories import file_factory
from graphon.graph_engine.layers.base import GraphEngineLayer
from graphon.model_runtime.errors.invoke import InvokeAuthorizationError
from graphon.runtime import GraphRuntimeState
from graphon.variable_loader import DUMMY_VARIABLE_LOADER, VariableLoader
from libs.flask_utils import preserve_flask_contexts from libs.flask_utils import preserve_flask_contexts
from models import Account, App, Conversation, EndUser, Message, Workflow, WorkflowNodeExecutionTriggeredFrom from models import Account, App, Conversation, EndUser, Message, Workflow, WorkflowNodeExecutionTriggeredFrom
from models.base import Base from models.base import Base
@ -150,85 +147,87 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator):
# #
# For implementation reference, see the `_parse_file` function and # For implementation reference, see the `_parse_file` function and
# `DraftWorkflowNodeRunApi` class which handle this properly. # `DraftWorkflowNodeRunApi` class which handle this properly.
files = args["files"] if args.get("files") else [] with self._bind_file_access_scope(tenant_id=app_model.tenant_id, user=user, invoke_from=invoke_from):
file_extra_config = FileUploadConfigManager.convert(workflow.features_dict, is_vision=False) files = args["files"] if args.get("files") else []
if file_extra_config: file_extra_config = FileUploadConfigManager.convert(workflow.features_dict, is_vision=False)
file_objs = file_factory.build_from_mappings( if file_extra_config:
mappings=files, file_objs = file_factory.build_from_mappings(
tenant_id=app_model.tenant_id, mappings=files,
config=file_extra_config, tenant_id=app_model.tenant_id,
config=file_extra_config,
access_controller=self._file_access_controller,
)
else:
file_objs = []
# convert to app config
app_config = AdvancedChatAppConfigManager.get_app_config(app_model=app_model, workflow=workflow)
# get tracing instance
trace_manager = TraceQueueManager(
app_id=app_model.id, user_id=user.id if isinstance(user, Account) else user.session_id
) )
else:
file_objs = []
# convert to app config if invoke_from == InvokeFrom.DEBUGGER:
app_config = AdvancedChatAppConfigManager.get_app_config(app_model=app_model, workflow=workflow) # always enable retriever resource in debugger mode
app_config.additional_features.show_retrieve_source = True # type: ignore
# get tracing instance # init application generate entity
trace_manager = TraceQueueManager( application_generate_entity = AdvancedChatAppGenerateEntity(
app_id=app_model.id, user_id=user.id if isinstance(user, Account) else user.session_id task_id=str(uuid.uuid4()),
) app_config=app_config,
file_upload_config=file_extra_config,
conversation_id=conversation.id if conversation else None,
inputs=self._prepare_user_inputs(
user_inputs=inputs, variables=app_config.variables, tenant_id=app_model.tenant_id
),
query=query,
files=list(file_objs),
parent_message_id=args.get("parent_message_id") if invoke_from != InvokeFrom.SERVICE_API else UUID_NIL,
user_id=user.id,
stream=streaming,
invoke_from=invoke_from,
extras=extras,
trace_manager=trace_manager,
workflow_run_id=str(workflow_run_id),
)
contexts.plugin_tool_providers.set({})
contexts.plugin_tool_providers_lock.set(threading.Lock())
if invoke_from == InvokeFrom.DEBUGGER: # Create repositories
# always enable retriever resource in debugger mode #
app_config.additional_features.show_retrieve_source = True # type: ignore # Create session factory
session_factory = sessionmaker(bind=db.engine, expire_on_commit=False)
# Create workflow execution(aka workflow run) repository
if invoke_from == InvokeFrom.DEBUGGER:
workflow_triggered_from = WorkflowRunTriggeredFrom.DEBUGGING
else:
workflow_triggered_from = WorkflowRunTriggeredFrom.APP_RUN
workflow_execution_repository = DifyCoreRepositoryFactory.create_workflow_execution_repository(
session_factory=session_factory,
user=user,
app_id=application_generate_entity.app_config.app_id,
triggered_from=workflow_triggered_from,
)
# Create workflow node execution repository
workflow_node_execution_repository = DifyCoreRepositoryFactory.create_workflow_node_execution_repository(
session_factory=session_factory,
user=user,
app_id=application_generate_entity.app_config.app_id,
triggered_from=WorkflowNodeExecutionTriggeredFrom.WORKFLOW_RUN,
)
# init application generate entity return self._generate(
application_generate_entity = AdvancedChatAppGenerateEntity( workflow=workflow,
task_id=str(uuid.uuid4()), user=user,
app_config=app_config, invoke_from=invoke_from,
file_upload_config=file_extra_config, application_generate_entity=application_generate_entity,
conversation_id=conversation.id if conversation else None, workflow_execution_repository=workflow_execution_repository,
inputs=self._prepare_user_inputs( workflow_node_execution_repository=workflow_node_execution_repository,
user_inputs=inputs, variables=app_config.variables, tenant_id=app_model.tenant_id conversation=conversation,
), stream=streaming,
query=query, pause_state_config=pause_state_config,
files=list(file_objs), )
parent_message_id=args.get("parent_message_id") if invoke_from != InvokeFrom.SERVICE_API else UUID_NIL,
user_id=user.id,
stream=streaming,
invoke_from=invoke_from,
extras=extras,
trace_manager=trace_manager,
workflow_run_id=str(workflow_run_id),
)
contexts.plugin_tool_providers.set({})
contexts.plugin_tool_providers_lock.set(threading.Lock())
# Create repositories
#
# Create session factory
session_factory = sessionmaker(bind=db.engine, expire_on_commit=False)
# Create workflow execution(aka workflow run) repository
if invoke_from == InvokeFrom.DEBUGGER:
workflow_triggered_from = WorkflowRunTriggeredFrom.DEBUGGING
else:
workflow_triggered_from = WorkflowRunTriggeredFrom.APP_RUN
workflow_execution_repository = DifyCoreRepositoryFactory.create_workflow_execution_repository(
session_factory=session_factory,
user=user,
app_id=application_generate_entity.app_config.app_id,
triggered_from=workflow_triggered_from,
)
# Create workflow node execution repository
workflow_node_execution_repository = DifyCoreRepositoryFactory.create_workflow_node_execution_repository(
session_factory=session_factory,
user=user,
app_id=application_generate_entity.app_config.app_id,
triggered_from=WorkflowNodeExecutionTriggeredFrom.WORKFLOW_RUN,
)
return self._generate(
workflow=workflow,
user=user,
invoke_from=invoke_from,
application_generate_entity=application_generate_entity,
workflow_execution_repository=workflow_execution_repository,
workflow_node_execution_repository=workflow_node_execution_repository,
conversation=conversation,
stream=streaming,
pause_state_config=pause_state_config,
)
def resume( def resume(
self, self,
@ -460,94 +459,90 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator):
:param conversation: conversation :param conversation: conversation
:param stream: is stream :param stream: is stream
""" """
is_first_conversation = conversation is None with self._bind_file_access_scope(
tenant_id=application_generate_entity.app_config.tenant_id,
user=user,
invoke_from=invoke_from,
):
is_first_conversation = conversation is None
if conversation is not None and message is not None: if conversation is not None and message is not None:
pass pass
else: else:
conversation, message = self._init_generate_records(application_generate_entity, conversation) conversation, message = self._init_generate_records(application_generate_entity, conversation)
if is_first_conversation: if is_first_conversation:
# update conversation features # update conversation features
conversation.override_model_configs = workflow.features conversation.override_model_configs = workflow.features
db.session.commit() db.session.commit()
db.session.refresh(conversation) db.session.refresh(conversation)
# get conversation dialogue count # get conversation dialogue count
# NOTE: dialogue_count should not start from 0, # NOTE: dialogue_count should not start from 0,
# because during the first conversation, dialogue_count should be 1. # because during the first conversation, dialogue_count should be 1.
self._dialogue_count = get_thread_messages_length(conversation.id) + 1 self._dialogue_count = get_thread_messages_length(conversation.id) + 1
# init queue manager # init queue manager
queue_manager = MessageBasedAppQueueManager( queue_manager = MessageBasedAppQueueManager(
task_id=application_generate_entity.task_id, task_id=application_generate_entity.task_id,
user_id=application_generate_entity.user_id, user_id=application_generate_entity.user_id,
invoke_from=application_generate_entity.invoke_from, invoke_from=application_generate_entity.invoke_from,
conversation_id=conversation.id, conversation_id=conversation.id,
app_mode=conversation.mode, app_mode=conversation.mode,
message_id=message.id, message_id=message.id,
)
graph_layers: list[GraphEngineLayer] = list(graph_engine_layers)
if pause_state_config is not None:
graph_layers.append(
PauseStatePersistenceLayer(
session_factory=pause_state_config.session_factory,
generate_entity=application_generate_entity,
state_owner_user_id=pause_state_config.state_owner_user_id,
)
) )
# new thread with request context and contextvars graph_layers: list[GraphEngineLayer] = list(graph_engine_layers)
context = contextvars.copy_context() if pause_state_config is not None:
graph_layers.append(
PauseStatePersistenceLayer(
session_factory=pause_state_config.session_factory,
generate_entity=application_generate_entity,
state_owner_user_id=pause_state_config.state_owner_user_id,
)
)
worker_thread = threading.Thread( # new thread with request context and contextvars
target=self._generate_worker, context = contextvars.copy_context()
kwargs={
"flask_app": current_app._get_current_object(), # type: ignore
"application_generate_entity": application_generate_entity,
"queue_manager": queue_manager,
"conversation_id": conversation.id,
"message_id": message.id,
"context": context,
"variable_loader": variable_loader,
"workflow_execution_repository": workflow_execution_repository,
"workflow_node_execution_repository": workflow_node_execution_repository,
"graph_engine_layers": tuple(graph_layers),
"graph_runtime_state": graph_runtime_state,
},
)
worker_thread.start() worker_thread = threading.Thread(
target=self._generate_worker,
kwargs={
"flask_app": current_app._get_current_object(), # type: ignore
"application_generate_entity": application_generate_entity,
"queue_manager": queue_manager,
"conversation_id": conversation.id,
"message_id": message.id,
"context": context,
"variable_loader": variable_loader,
"workflow_execution_repository": workflow_execution_repository,
"workflow_node_execution_repository": workflow_node_execution_repository,
"graph_engine_layers": tuple(graph_layers),
"graph_runtime_state": graph_runtime_state,
},
)
# release database connection, because the following new thread operations may take a long time worker_thread.start()
with Session(bind=db.engine, expire_on_commit=False) as session:
workflow = _refresh_model(session, workflow)
message = _refresh_model(session, message)
# workflow_ = session.get(Workflow, workflow.id)
# assert workflow_ is not None
# workflow = workflow_
# message_ = session.get(Message, message.id)
# assert message_ is not None
# message = message_
# db.session.refresh(workflow)
# db.session.refresh(message)
# db.session.refresh(user)
db.session.close()
# return response or stream generator # release database connection, because the following new thread operations may take a long time
response = self._handle_advanced_chat_response( with Session(bind=db.engine, expire_on_commit=False) as session:
application_generate_entity=application_generate_entity, workflow = _refresh_model(session, workflow)
workflow=workflow, message = _refresh_model(session, message)
queue_manager=queue_manager, db.session.close()
conversation=conversation,
message=message,
user=user,
stream=stream,
draft_var_saver_factory=self._get_draft_var_saver_factory(invoke_from, account=user),
)
return AdvancedChatAppGenerateResponseConverter.convert(response=response, invoke_from=invoke_from) # return response or stream generator
response = self._handle_advanced_chat_response(
application_generate_entity=application_generate_entity,
workflow=workflow,
queue_manager=queue_manager,
conversation=conversation,
message=message,
user=user,
stream=stream,
draft_var_saver_factory=self._get_draft_var_saver_factory(invoke_from, account=user),
)
return AdvancedChatAppGenerateResponseConverter.convert(response=response, invoke_from=invoke_from)
def _generate_worker( def _generate_worker(
self, self,

View File

@ -25,19 +25,24 @@ from core.app.workflow.layers.persistence import PersistenceWorkflowInfo, Workfl
from core.db.session_factory import session_factory from core.db.session_factory import session_factory
from core.moderation.base import ModerationError from core.moderation.base import ModerationError
from core.moderation.input_moderation import InputModeration from core.moderation.input_moderation import InputModeration
from core.repositories.factory import WorkflowExecutionRepository, WorkflowNodeExecutionRepository
from core.workflow.node_factory import get_default_root_node_id
from core.workflow.system_variables import (
build_bootstrap_variables,
build_system_variables,
system_variables_to_mapping,
)
from core.workflow.variable_pool_initializer import add_node_inputs_to_pool, add_variables_to_pool
from core.workflow.workflow_entry import WorkflowEntry from core.workflow.workflow_entry import WorkflowEntry
from dify_graph.enums import WorkflowType
from dify_graph.graph_engine.command_channels.redis_channel import RedisChannel
from dify_graph.graph_engine.layers.base import GraphEngineLayer
from dify_graph.repositories.workflow_execution_repository import WorkflowExecutionRepository
from dify_graph.repositories.workflow_node_execution_repository import WorkflowNodeExecutionRepository
from dify_graph.runtime import GraphRuntimeState, VariablePool
from dify_graph.system_variable import SystemVariable
from dify_graph.variable_loader import VariableLoader
from dify_graph.variables.variables import Variable
from extensions.ext_database import db from extensions.ext_database import db
from extensions.ext_redis import redis_client from extensions.ext_redis import redis_client
from extensions.otel import WorkflowAppRunnerHandler, trace_span from extensions.otel import WorkflowAppRunnerHandler, trace_span
from graphon.enums import WorkflowType
from graphon.graph_engine.command_channels.redis_channel import RedisChannel
from graphon.graph_engine.layers.base import GraphEngineLayer
from graphon.runtime import GraphRuntimeState, VariablePool
from graphon.variable_loader import VariableLoader
from graphon.variables.variables import Variable
from models import Workflow from models import Workflow
from models.model import App, Conversation, Message, MessageAnnotation from models.model import App, Conversation, Message, MessageAnnotation
from models.workflow import ConversationVariable from models.workflow import ConversationVariable
@ -90,7 +95,7 @@ class AdvancedChatAppRunner(WorkflowBasedAppRunner):
app_config = self.application_generate_entity.app_config app_config = self.application_generate_entity.app_config
app_config = cast(AdvancedChatAppConfig, app_config) app_config = cast(AdvancedChatAppConfig, app_config)
system_inputs = SystemVariable( system_inputs = build_system_variables(
query=self.application_generate_entity.query, query=self.application_generate_entity.query,
files=self.application_generate_entity.files, files=self.application_generate_entity.files,
conversation_id=self.conversation.id, conversation_id=self.conversation.id,
@ -132,6 +137,7 @@ class AdvancedChatAppRunner(WorkflowBasedAppRunner):
workflow=self._workflow, workflow=self._workflow,
single_iteration_run=self.application_generate_entity.single_iteration_run, single_iteration_run=self.application_generate_entity.single_iteration_run,
single_loop_run=self.application_generate_entity.single_loop_run, single_loop_run=self.application_generate_entity.single_loop_run,
user_id=self.application_generate_entity.user_id,
) )
else: else:
inputs = self.application_generate_entity.inputs inputs = self.application_generate_entity.inputs
@ -150,7 +156,10 @@ class AdvancedChatAppRunner(WorkflowBasedAppRunner):
self.application_generate_entity.inputs = new_inputs self.application_generate_entity.inputs = new_inputs
self.application_generate_entity.query = new_query self.application_generate_entity.query = new_query
system_inputs.query = new_query system_inputs = build_system_variables(
system_variables_to_mapping(system_inputs),
query=new_query,
)
# annotation reply # annotation reply
if self.handle_annotation_reply( if self.handle_annotation_reply(
@ -166,14 +175,17 @@ class AdvancedChatAppRunner(WorkflowBasedAppRunner):
# Create a variable pool. # Create a variable pool.
# init variable pool # init variable pool
variable_pool = VariablePool( variable_pool = VariablePool()
system_variables=system_inputs, add_variables_to_pool(
user_inputs=new_inputs, variable_pool,
environment_variables=self._workflow.environment_variables, build_bootstrap_variables(
# Based on the definition of `Variable`, system_variables=system_inputs,
# `VariableBase` instances can be safely used as `Variable` since they are compatible. environment_variables=self._workflow.environment_variables,
conversation_variables=conversation_variables, conversation_variables=conversation_variables,
),
) )
root_node_id = get_default_root_node_id(self._workflow.graph_dict)
add_node_inputs_to_pool(variable_pool, node_id=root_node_id, inputs=new_inputs)
# init graph # init graph
graph_runtime_state = GraphRuntimeState(variable_pool=variable_pool, start_at=time.time()) graph_runtime_state = GraphRuntimeState(variable_pool=variable_pool, start_at=time.time())
@ -185,6 +197,7 @@ class AdvancedChatAppRunner(WorkflowBasedAppRunner):
user_id=self.application_generate_entity.user_id, user_id=self.application_generate_entity.user_id,
user_from=user_from, user_from=user_from,
invoke_from=invoke_from, invoke_from=invoke_from,
root_node_id=root_node_id,
) )
db.session.close() db.session.close()

View File

@ -14,6 +14,7 @@ from constants.tts_auto_play_timeout import TTS_AUTO_PLAY_TIMEOUT, TTS_AUTO_PLAY
from core.app.apps.base_app_queue_manager import AppQueueManager, PublishFrom from core.app.apps.base_app_queue_manager import AppQueueManager, PublishFrom
from core.app.apps.common.graph_runtime_state_support import GraphRuntimeStateSupport from core.app.apps.common.graph_runtime_state_support import GraphRuntimeStateSupport
from core.app.apps.common.workflow_response_converter import WorkflowResponseConverter from core.app.apps.common.workflow_response_converter import WorkflowResponseConverter
from core.app.apps.draft_variable_saver import DraftVariableSaverFactory
from core.app.entities.app_invoke_entities import ( from core.app.entities.app_invoke_entities import (
AdvancedChatAppGenerateEntity, AdvancedChatAppGenerateEntity,
InvokeFrom, InvokeFrom,
@ -65,15 +66,15 @@ from core.app.task_pipeline.message_cycle_manager import MessageCycleManager
from core.base.tts import AppGeneratorTTSPublisher, AudioTrunk from core.base.tts import AppGeneratorTTSPublisher, AudioTrunk
from core.ops.ops_trace_manager import TraceQueueManager from core.ops.ops_trace_manager import TraceQueueManager
from core.repositories.human_input_repository import HumanInputFormRepositoryImpl from core.repositories.human_input_repository import HumanInputFormRepositoryImpl
from dify_graph.entities.pause_reason import HumanInputRequired from core.workflow.file_reference import resolve_file_record_id
from dify_graph.enums import WorkflowExecutionStatus from core.workflow.system_variables import build_system_variables
from dify_graph.model_runtime.entities.llm_entities import LLMUsage
from dify_graph.model_runtime.utils.encoders import jsonable_encoder
from dify_graph.nodes import BuiltinNodeTypes
from dify_graph.repositories.draft_variable_repository import DraftVariableSaverFactory
from dify_graph.runtime import GraphRuntimeState
from dify_graph.system_variable import SystemVariable
from extensions.ext_database import db from extensions.ext_database import db
from graphon.entities.pause_reason import HumanInputRequired
from graphon.enums import WorkflowExecutionStatus
from graphon.model_runtime.entities.llm_entities import LLMUsage
from graphon.model_runtime.utils.encoders import jsonable_encoder
from graphon.nodes import BuiltinNodeTypes
from graphon.runtime import GraphRuntimeState
from libs.datetime_utils import naive_utc_now from libs.datetime_utils import naive_utc_now
from models import Account, Conversation, EndUser, Message, MessageFile from models import Account, Conversation, EndUser, Message, MessageFile
from models.enums import CreatorUserRole, MessageFileBelongsTo, MessageStatus from models.enums import CreatorUserRole, MessageFileBelongsTo, MessageStatus
@ -117,7 +118,7 @@ class AdvancedChatAppGenerateTaskPipeline(GraphRuntimeStateSupport):
else: else:
raise NotImplementedError(f"User type not supported: {type(user)}") raise NotImplementedError(f"User type not supported: {type(user)}")
self._workflow_system_variables = SystemVariable( self._workflow_system_variables = build_system_variables(
query=message.query, query=message.query,
files=application_generate_entity.files, files=application_generate_entity.files,
conversation_id=conversation.id, conversation_id=conversation.id,
@ -741,8 +742,9 @@ class AdvancedChatAppGenerateTaskPipeline(GraphRuntimeStateSupport):
def _load_human_input_form_id(self, *, node_id: str) -> str | None: def _load_human_input_form_id(self, *, node_id: str) -> str | None:
form_repository = HumanInputFormRepositoryImpl( form_repository = HumanInputFormRepositoryImpl(
tenant_id=self._workflow_tenant_id, tenant_id=self._workflow_tenant_id,
workflow_execution_id=self._workflow_run_id,
) )
form = form_repository.get_form(self._workflow_run_id, node_id) form = form_repository.get_form(node_id)
if form is None: if form is None:
return None return None
return form.id return form.id
@ -933,21 +935,23 @@ class AdvancedChatAppGenerateTaskPipeline(GraphRuntimeStateSupport):
metadata = self._task_state.metadata.model_dump() metadata = self._task_state.metadata.model_dump()
message.message_metadata = json.dumps(jsonable_encoder(metadata)) message.message_metadata = json.dumps(jsonable_encoder(metadata))
message_files = [ message_files: list[MessageFile] = []
MessageFile( for file in self._recorded_files:
message_id=message.id, reference = file.get("reference") or file.get("related_id")
type=file["type"], message_files.append(
transfer_method=file["transfer_method"], MessageFile(
url=file["remote_url"], message_id=message.id,
belongs_to=MessageFileBelongsTo.ASSISTANT, type=file["type"],
upload_file_id=file["related_id"], transfer_method=file["transfer_method"],
created_by_role=CreatorUserRole.ACCOUNT url=file["remote_url"],
if message.invoke_from in {InvokeFrom.EXPLORE, InvokeFrom.DEBUGGER} belongs_to=MessageFileBelongsTo.ASSISTANT,
else CreatorUserRole.END_USER, upload_file_id=resolve_file_record_id(reference if isinstance(reference, str) else None),
created_by=message.from_account_id or message.from_end_user_id or "", created_by_role=CreatorUserRole.ACCOUNT
if message.invoke_from in {InvokeFrom.EXPLORE, InvokeFrom.DEBUGGER}
else CreatorUserRole.END_USER,
created_by=message.from_account_id or message.from_end_user_id or "",
)
) )
for file in self._recorded_files
]
session.add_all(message_files) session.add_all(message_files)
def _seed_graph_runtime_state_from_queue_manager(self) -> None: def _seed_graph_runtime_state_from_queue_manager(self) -> None:
@ -1003,13 +1007,11 @@ class AdvancedChatAppGenerateTaskPipeline(GraphRuntimeStateSupport):
return message return message
def _save_output_for_event(self, event: QueueNodeSucceededEvent | QueueNodeExceptionEvent, node_execution_id: str): def _save_output_for_event(self, event: QueueNodeSucceededEvent | QueueNodeExceptionEvent, node_execution_id: str):
with Session(db.engine) as session, session.begin(): saver = self._draft_var_saver_factory(
saver = self._draft_var_saver_factory( app_id=self._application_generate_entity.app_config.app_id,
session=session, node_id=event.node_id,
app_id=self._application_generate_entity.app_config.app_id, node_type=event.node_type,
node_id=event.node_id, node_execution_id=node_execution_id,
node_type=event.node_type, enclosing_node_id=event.in_loop_id or event.in_iteration_id,
node_execution_id=node_execution_id, )
enclosing_node_id=event.in_loop_id or event.in_iteration_id, saver.save(event.process_data, event.outputs)
)
saver.save(event.process_data, event.outputs)

View File

@ -21,9 +21,9 @@ from core.app.apps.message_based_app_generator import MessageBasedAppGenerator
from core.app.apps.message_based_app_queue_manager import MessageBasedAppQueueManager from core.app.apps.message_based_app_queue_manager import MessageBasedAppQueueManager
from core.app.entities.app_invoke_entities import AgentChatAppGenerateEntity, InvokeFrom from core.app.entities.app_invoke_entities import AgentChatAppGenerateEntity, InvokeFrom
from core.ops.ops_trace_manager import TraceQueueManager from core.ops.ops_trace_manager import TraceQueueManager
from dify_graph.model_runtime.errors.invoke import InvokeAuthorizationError
from extensions.ext_database import db from extensions.ext_database import db
from factories import file_factory from factories import file_factory
from graphon.model_runtime.errors.invoke import InvokeAuthorizationError
from libs.flask_utils import preserve_flask_contexts from libs.flask_utils import preserve_flask_contexts
from models import Account, App, EndUser from models import Account, App, EndUser
from services.conversation_service import ConversationService from services.conversation_service import ConversationService
@ -129,89 +129,93 @@ class AgentChatAppGenerator(MessageBasedAppGenerator):
# #
# For implementation reference, see the `_parse_file` function and # For implementation reference, see the `_parse_file` function and
# `DraftWorkflowNodeRunApi` class which handle this properly. # `DraftWorkflowNodeRunApi` class which handle this properly.
files = args.get("files") or [] with self._bind_file_access_scope(tenant_id=app_model.tenant_id, user=user, invoke_from=invoke_from):
file_extra_config = FileUploadConfigManager.convert(override_model_config_dict or app_model_config.to_dict()) files = args.get("files") or []
if file_extra_config: file_extra_config = FileUploadConfigManager.convert(
file_objs = file_factory.build_from_mappings( override_model_config_dict or app_model_config.to_dict()
mappings=files,
tenant_id=app_model.tenant_id,
config=file_extra_config,
) )
else: if file_extra_config:
file_objs = [] file_objs = file_factory.build_from_mappings(
mappings=files,
tenant_id=app_model.tenant_id,
config=file_extra_config,
access_controller=self._file_access_controller,
)
else:
file_objs = []
# convert to app config # convert to app config
app_config = AgentChatAppConfigManager.get_app_config( app_config = AgentChatAppConfigManager.get_app_config(
app_model=app_model, app_model=app_model,
app_model_config=app_model_config, app_model_config=app_model_config,
conversation=conversation, conversation=conversation,
override_config_dict=override_model_config_dict, override_config_dict=override_model_config_dict,
) )
# get tracing instance # get tracing instance
trace_manager = TraceQueueManager(app_model.id, user.id if isinstance(user, Account) else user.session_id) trace_manager = TraceQueueManager(app_model.id, user.id if isinstance(user, Account) else user.session_id)
# init application generate entity # init application generate entity
application_generate_entity = AgentChatAppGenerateEntity( application_generate_entity = AgentChatAppGenerateEntity(
task_id=str(uuid.uuid4()), task_id=str(uuid.uuid4()),
app_config=app_config, app_config=app_config,
model_conf=ModelConfigConverter.convert(app_config), model_conf=ModelConfigConverter.convert(app_config),
file_upload_config=file_extra_config, file_upload_config=file_extra_config,
conversation_id=conversation.id if conversation else None, conversation_id=conversation.id if conversation else None,
inputs=self._prepare_user_inputs( inputs=self._prepare_user_inputs(
user_inputs=inputs, variables=app_config.variables, tenant_id=app_model.tenant_id user_inputs=inputs, variables=app_config.variables, tenant_id=app_model.tenant_id
), ),
query=query, query=query,
files=list(file_objs), files=list(file_objs),
parent_message_id=args.get("parent_message_id") if invoke_from != InvokeFrom.SERVICE_API else UUID_NIL, parent_message_id=args.get("parent_message_id") if invoke_from != InvokeFrom.SERVICE_API else UUID_NIL,
user_id=user.id, user_id=user.id,
stream=streaming, stream=streaming,
invoke_from=invoke_from, invoke_from=invoke_from,
extras=extras, extras=extras,
call_depth=0, call_depth=0,
trace_manager=trace_manager, trace_manager=trace_manager,
) )
# init generate records # init generate records
(conversation, message) = self._init_generate_records(application_generate_entity, conversation) (conversation, message) = self._init_generate_records(application_generate_entity, conversation)
# init queue manager # init queue manager
queue_manager = MessageBasedAppQueueManager( queue_manager = MessageBasedAppQueueManager(
task_id=application_generate_entity.task_id, task_id=application_generate_entity.task_id,
user_id=application_generate_entity.user_id, user_id=application_generate_entity.user_id,
invoke_from=application_generate_entity.invoke_from, invoke_from=application_generate_entity.invoke_from,
conversation_id=conversation.id, conversation_id=conversation.id,
app_mode=conversation.mode, app_mode=conversation.mode,
message_id=message.id, message_id=message.id,
) )
# new thread with request context and contextvars # new thread with request context and contextvars
context = contextvars.copy_context() context = contextvars.copy_context()
worker_thread = threading.Thread( worker_thread = threading.Thread(
target=self._generate_worker, target=self._generate_worker,
kwargs={ kwargs={
"flask_app": current_app._get_current_object(), # type: ignore "flask_app": current_app._get_current_object(), # type: ignore
"context": context, "context": context,
"application_generate_entity": application_generate_entity, "application_generate_entity": application_generate_entity,
"queue_manager": queue_manager, "queue_manager": queue_manager,
"conversation_id": conversation.id, "conversation_id": conversation.id,
"message_id": message.id, "message_id": message.id,
}, },
) )
worker_thread.start() worker_thread.start()
# return response or stream generator # return response or stream generator
response = self._handle_response( response = self._handle_response(
application_generate_entity=application_generate_entity, application_generate_entity=application_generate_entity,
queue_manager=queue_manager, queue_manager=queue_manager,
conversation=conversation, conversation=conversation,
message=message, message=message,
user=user, user=user,
stream=streaming, stream=streaming,
) )
return AgentChatAppGenerateResponseConverter.convert(response=response, invoke_from=invoke_from) return AgentChatAppGenerateResponseConverter.convert(response=response, invoke_from=invoke_from)
def _generate_worker( def _generate_worker(
self, self,

View File

@ -15,10 +15,10 @@ from core.app.entities.queue_entities import QueueAnnotationReplyEvent
from core.memory.token_buffer_memory import TokenBufferMemory from core.memory.token_buffer_memory import TokenBufferMemory
from core.model_manager import ModelInstance from core.model_manager import ModelInstance
from core.moderation.base import ModerationError from core.moderation.base import ModerationError
from dify_graph.model_runtime.entities.llm_entities import LLMMode
from dify_graph.model_runtime.entities.model_entities import ModelFeature, ModelPropertyKey
from dify_graph.model_runtime.model_providers.__base.large_language_model import LargeLanguageModel
from extensions.ext_database import db from extensions.ext_database import db
from graphon.model_runtime.entities.llm_entities import LLMMode
from graphon.model_runtime.entities.model_entities import ModelFeature, ModelPropertyKey
from graphon.model_runtime.model_providers.__base.large_language_model import LargeLanguageModel
from models.model import App, Conversation, Message from models.model import App, Conversation, Message
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)

View File

@ -6,7 +6,7 @@ from typing import Any, Union
from core.app.entities.app_invoke_entities import InvokeFrom from core.app.entities.app_invoke_entities import InvokeFrom
from core.app.entities.task_entities import AppBlockingResponse, AppStreamResponse from core.app.entities.task_entities import AppBlockingResponse, AppStreamResponse
from core.errors.error import ModelCurrentlyNotSupportError, ProviderTokenNotInitError, QuotaExceededError from core.errors.error import ModelCurrentlyNotSupportError, ProviderTokenNotInitError, QuotaExceededError
from dify_graph.model_runtime.errors.invoke import InvokeError from graphon.model_runtime.errors.invoke import InvokeError
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)

View File

@ -1,27 +1,89 @@
from collections.abc import Generator, Mapping, Sequence from collections.abc import Generator, Mapping, Sequence
from contextlib import AbstractContextManager, nullcontext
from typing import TYPE_CHECKING, Any, Union, final from typing import TYPE_CHECKING, Any, Union, final
from sqlalchemy.orm import Session from sqlalchemy.orm import Session
from core.app.entities.app_invoke_entities import InvokeFrom from core.app.apps.draft_variable_saver import (
from dify_graph.enums import NodeType
from dify_graph.file import File, FileUploadConfig
from dify_graph.repositories.draft_variable_repository import (
DraftVariableSaver, DraftVariableSaver,
DraftVariableSaverFactory, DraftVariableSaverFactory,
NoopDraftVariableSaver, NoopDraftVariableSaver,
) )
from dify_graph.variables.input_entities import VariableEntityType from core.app.entities.app_invoke_entities import InvokeFrom, UserFrom
from core.app.file_access import DatabaseFileAccessController, FileAccessScope, bind_file_access_scope
from extensions.ext_database import db
from factories import file_factory from factories import file_factory
from graphon.enums import NodeType
from graphon.file import File, FileUploadConfig
from graphon.variables.input_entities import VariableEntityType
from libs.orjson import orjson_dumps from libs.orjson import orjson_dumps
from models import Account, EndUser from models import Account, EndUser
from services.workflow_draft_variable_service import DraftVariableSaver as DraftVariableSaverImpl from services.workflow_draft_variable_service import DraftVariableSaver as DraftVariableSaverImpl
if TYPE_CHECKING: if TYPE_CHECKING:
from dify_graph.variables.input_entities import VariableEntity from graphon.variables.input_entities import VariableEntity
@final
class _DebuggerDraftVariableSaver:
"""Adapter that binds SQLAlchemy session setup outside the saver port."""
def __init__(
self,
*,
account: Account,
app_id: str,
node_id: str,
node_type: NodeType,
node_execution_id: str,
enclosing_node_id: str | None = None,
) -> None:
self._account = account
self._app_id = app_id
self._node_id = node_id
self._node_type = node_type
self._node_execution_id = node_execution_id
self._enclosing_node_id = enclosing_node_id
def save(self, process_data: Mapping[str, Any] | None, outputs: Mapping[str, Any] | None) -> None:
with Session(db.engine) as session, session.begin():
DraftVariableSaverImpl(
session=session,
app_id=self._app_id,
node_id=self._node_id,
node_type=self._node_type,
node_execution_id=self._node_execution_id,
enclosing_node_id=self._enclosing_node_id,
user=self._account,
).save(process_data, outputs)
class BaseAppGenerator: class BaseAppGenerator:
_file_access_controller: DatabaseFileAccessController = DatabaseFileAccessController()
@staticmethod
def _bind_file_access_scope(
*,
tenant_id: str,
user: Account | EndUser,
invoke_from: InvokeFrom,
) -> AbstractContextManager[None]:
"""Bind request-scoped file ownership markers for downstream file lookups."""
user_id = getattr(user, "id", None)
if not isinstance(user_id, str) or not user_id:
return nullcontext()
user_from = UserFrom.ACCOUNT if isinstance(user, Account) else UserFrom.END_USER
return bind_file_access_scope(
FileAccessScope(
tenant_id=tenant_id,
user_id=user_id,
user_from=user_from,
invoke_from=invoke_from,
)
)
def _prepare_user_inputs( def _prepare_user_inputs(
self, self,
*, *,
@ -50,6 +112,7 @@ class BaseAppGenerator:
allowed_file_upload_methods=entity_dictionary[k].allowed_file_upload_methods or [], allowed_file_upload_methods=entity_dictionary[k].allowed_file_upload_methods or [],
), ),
strict_type_validation=strict_type_validation, strict_type_validation=strict_type_validation,
access_controller=self._file_access_controller,
) )
for k, v in user_inputs.items() for k, v in user_inputs.items()
if isinstance(v, dict) and entity_dictionary[k].type == VariableEntityType.FILE if isinstance(v, dict) and entity_dictionary[k].type == VariableEntityType.FILE
@ -64,6 +127,7 @@ class BaseAppGenerator:
allowed_file_extensions=entity_dictionary[k].allowed_file_extensions or [], allowed_file_extensions=entity_dictionary[k].allowed_file_extensions or [],
allowed_file_upload_methods=entity_dictionary[k].allowed_file_upload_methods or [], allowed_file_upload_methods=entity_dictionary[k].allowed_file_upload_methods or [],
), ),
access_controller=self._file_access_controller,
) )
for k, v in user_inputs.items() for k, v in user_inputs.items()
if isinstance(v, list) if isinstance(v, list)
@ -226,32 +290,30 @@ class BaseAppGenerator:
assert isinstance(account, Account) assert isinstance(account, Account)
def draft_var_saver_factory( def draft_var_saver_factory(
session: Session,
app_id: str, app_id: str,
node_id: str, node_id: str,
node_type: NodeType, node_type: NodeType,
node_execution_id: str, node_execution_id: str,
enclosing_node_id: str | None = None, enclosing_node_id: str | None = None,
) -> DraftVariableSaver: ) -> DraftVariableSaver:
return DraftVariableSaverImpl( return _DebuggerDraftVariableSaver(
session=session, account=account,
app_id=app_id, app_id=app_id,
node_id=node_id, node_id=node_id,
node_type=node_type, node_type=node_type,
node_execution_id=node_execution_id, node_execution_id=node_execution_id,
enclosing_node_id=enclosing_node_id, enclosing_node_id=enclosing_node_id,
user=account,
) )
else: else:
def draft_var_saver_factory( def draft_var_saver_factory(
session: Session,
app_id: str, app_id: str,
node_id: str, node_id: str,
node_type: NodeType, node_type: NodeType,
node_execution_id: str, node_execution_id: str,
enclosing_node_id: str | None = None, enclosing_node_id: str | None = None,
) -> DraftVariableSaver: ) -> DraftVariableSaver:
_ = app_id, node_id, node_type, node_execution_id, enclosing_node_id
return NoopDraftVariableSaver() return NoopDraftVariableSaver()
return draft_var_saver_factory return draft_var_saver_factory

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