mirror of
https://github.com/langgenius/dify.git
synced 2026-09-08 11:04:27 +08:00
Merge branch 'origin/main' into feat/evaluation
This commit is contained in:
commit
ef3973f188
13
.gemini/config.yaml
Normal file
13
.gemini/config.yaml
Normal 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
2
.github/CODEOWNERS
vendored
@ -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
|
||||||
|
|||||||
9
.github/actions/setup-web/action.yml
vendored
9
.github/actions/setup-web/action.yml
vendored
@ -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
|
|
||||||
|
|||||||
29
.github/workflows/style.yml
vendored
29
.github/workflows/style.yml
vendored
@ -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
|
||||||
|
|||||||
2
.github/workflows/translate-i18n-claude.yml
vendored
2
.github/workflows/translate-i18n-claude.yml
vendored
@ -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 }}
|
||||||
|
|||||||
@ -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
|
||||||
|
|||||||
@ -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
|
||||||
|
|||||||
@ -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
|
||||||
|
|||||||
@ -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 = []
|
||||||
|
|||||||
@ -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,
|
||||||
|
)
|
||||||
|
|||||||
@ -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",
|
||||||
]
|
]
|
||||||
|
|||||||
@ -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",
|
||||||
|
]
|
||||||
@ -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
|
||||||
|
|||||||
@ -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"))
|
||||||
)
|
)
|
||||||
|
|||||||
@ -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]
|
||||||
|
|||||||
@ -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"
|
||||||
|
|||||||
@ -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,
|
||||||
|
|||||||
@ -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
|
||||||
|
|||||||
@ -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
|
||||||
|
|||||||
@ -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:
|
||||||
|
|||||||
@ -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)
|
||||||
|
|||||||
@ -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()
|
||||||
|
|||||||
@ -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
|
||||||
|
|||||||
@ -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
|
||||||
|
|||||||
@ -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
|
||||||
|
|||||||
@ -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
|
||||||
|
|
||||||
|
|||||||
@ -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
|
||||||
|
|||||||
@ -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()
|
||||||
|
|||||||
@ -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)
|
||||||
|
),
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
)
|
)
|
||||||
|
|||||||
@ -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
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@ -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:
|
||||||
|
|||||||
@ -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 (
|
||||||
|
|||||||
@ -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
|
||||||
|
|||||||
@ -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
|
||||||
|
|||||||
@ -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
|
||||||
|
|||||||
@ -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}"
|
||||||
|
|||||||
@ -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.")
|
||||||
|
|||||||
@ -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
|
||||||
|
|
||||||
|
|||||||
@ -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
|
||||||
|
|||||||
@ -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
|
||||||
|
|||||||
@ -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()
|
||||||
|
|||||||
@ -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
|
||||||
|
|||||||
@ -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:
|
||||||
|
|||||||
@ -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,
|
||||||
|
|||||||
@ -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
|
||||||
|
|||||||
@ -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
|
||||||
|
|||||||
@ -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
|
||||||
|
|||||||
@ -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
|
||||||
|
|||||||
@ -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
|
||||||
|
|
||||||
|
|||||||
@ -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
|
||||||
|
|
||||||
|
|||||||
@ -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
|
||||||
|
|
||||||
|
|||||||
@ -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
|
||||||
|
|||||||
@ -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
|
||||||
|
|||||||
@ -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(
|
||||||
|
|||||||
@ -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
|
||||||
|
|||||||
@ -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
|
||||||
|
|||||||
@ -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
|
||||||
|
|||||||
@ -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,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@ -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 (
|
||||||
|
|||||||
@ -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",
|
||||||
|
|||||||
1
api/controllers/inner_api/app/__init__.py
Normal file
1
api/controllers/inner_api/app/__init__.py
Normal file
@ -0,0 +1 @@
|
|||||||
|
|
||||||
110
api/controllers/inner_api/app/dsl.py
Normal file
110
api/controllers/inner_api/app/dsl.py
Normal 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
|
||||||
@ -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
|
||||||
|
|||||||
@ -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
|
||||||
|
|||||||
@ -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 (
|
||||||
|
|||||||
@ -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
|
||||||
|
|||||||
@ -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")
|
||||||
|
|
||||||
|
|||||||
@ -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()
|
||||||
|
|||||||
@ -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
|
||||||
|
|||||||
@ -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
|
||||||
|
|||||||
@ -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.")
|
||||||
|
|||||||
@ -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.")
|
||||||
|
|
||||||
|
|||||||
@ -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
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@ -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.")
|
||||||
|
|||||||
@ -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
|
||||||
|
|||||||
@ -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
|
||||||
|
|||||||
@ -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
|
||||||
|
|||||||
@ -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
|
||||||
|
|||||||
@ -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
|
||||||
|
|||||||
@ -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)
|
||||||
|
|||||||
@ -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=[],
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@ -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):
|
||||||
|
|||||||
@ -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):
|
||||||
|
|||||||
@ -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=[],
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@ -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:
|
||||||
|
|||||||
@ -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
|
||||||
)
|
)
|
||||||
|
|||||||
@ -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
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@ -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
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@ -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(
|
||||||
|
|||||||
@ -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
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@ -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:
|
||||||
|
|||||||
@ -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
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@ -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,
|
||||||
|
|||||||
@ -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()
|
||||||
|
|||||||
@ -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)
|
|
||||||
|
|||||||
@ -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,
|
||||||
|
|||||||
@ -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__)
|
||||||
|
|||||||
@ -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__)
|
||||||
|
|
||||||
|
|||||||
@ -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
Loading…
Reference in New Issue
Block a user