Compare commits
110
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
c2fc4f4822 | ||
|
|
67e0eeefd2 | ||
|
|
cf58035fa2 | ||
|
|
097f10d920 | ||
|
|
c34d05141e | ||
|
|
302701b303 | ||
|
|
8080159eaf | ||
|
|
c730fec1e4 | ||
|
|
b4fec9b7aa | ||
|
|
7e0bccbbf0 | ||
|
|
2f87ecc0ce | ||
|
|
5b4c7b2a40 | ||
|
|
378a1d7d08 | ||
|
|
ce0192620d | ||
|
|
e9feeedc01 | ||
|
|
e32490f54e | ||
|
|
e9db50f781 | ||
|
|
0310f631ee | ||
|
|
abc5a61e98 | ||
|
|
5f1698add6 | ||
|
|
36e50f277f | ||
|
|
704ee40caa | ||
|
|
3119c99979 | ||
|
|
16b8733886 | ||
|
|
83f64104fd | ||
|
|
5077879886 | ||
|
|
697b57631a | ||
|
|
6015f23e79 | ||
|
|
f355c8d595 | ||
|
|
0142001fc2 | ||
|
|
4058e9ae23 | ||
|
|
95310561ec | ||
|
|
de33561a52 | ||
|
|
6d9665578b | ||
|
|
18f14c04dc | ||
|
|
14251b249d | ||
|
|
1819bd72ef | ||
|
|
7dabc03a08 | ||
|
|
1a050c9f86 | ||
|
|
7fb6e0cdfe | ||
|
|
e0fcf33979 | ||
|
|
898e09264b | ||
|
|
4ac461d882 | ||
|
|
fa763216d0 | ||
|
|
d546210040 | ||
|
|
4e0a7a7f9e | ||
|
|
e4ab6e0919 | ||
|
|
6fa943fe75 | ||
|
|
a1fc280102 | ||
|
|
56e3a55023 | ||
|
|
6c63c6a221 | ||
|
|
5b06203ef5 | ||
|
|
3348b89436 | ||
|
|
0428ac5f3a | ||
|
|
aead4fe65c | ||
|
|
bdf6739b86 | ||
|
|
483db22b97 | ||
|
|
aa800d838d | ||
|
|
4bd80683a4 | ||
|
|
c185a51bad | ||
|
|
4430a1b3da | ||
|
|
2c9430313d | ||
|
|
552ee369b2 | ||
|
|
d5b9a7b2f8 | ||
|
|
c2a3f459c7 | ||
|
|
4971e11734 | ||
|
|
a297b06aac | ||
|
|
e988266f53 | ||
|
|
d9530f7bb7 | ||
|
|
b24e6edada | ||
|
|
59a9cbbf78 | ||
|
|
45164ce33e | ||
|
|
095b3ee234 | ||
|
|
cb970e54da | ||
|
|
e04f2a0786 | ||
|
|
7202a24bcf | ||
|
|
be8f265e43 | ||
|
|
9e54f086dc | ||
|
|
8c31b69c8e | ||
|
|
b886b3f6c8 | ||
|
|
ef0d18bb61 | ||
|
|
c56ad8e323 | ||
|
|
365f749ed5 | ||
|
|
f686197589 | ||
|
|
f584be9cf0 | ||
|
|
3bd228ddb7 | ||
|
|
0dfa59b1db | ||
|
|
1e344f773b | ||
|
|
bba2040a05 | ||
|
|
ad3be1e4d0 | ||
|
|
297dd832aa | ||
|
|
cc5705cb71 | ||
|
|
74b027c41a | ||
|
|
5f69470ebf | ||
|
|
ec7ccd800c | ||
|
|
0d74ac634b | ||
|
|
468990cc39 | ||
|
|
64e769f96e | ||
|
|
778aabb485 | ||
|
|
d8402f686e | ||
|
|
8bd8dee767 | ||
|
|
05f2764d7c | ||
|
|
f5d6c250ed | ||
|
|
45daec7541 | ||
|
|
c14a8bb437 | ||
|
|
b76c8fa853 | ||
|
|
8c3e77cd0c | ||
|
|
476946f122 | ||
|
|
62a698a883 | ||
|
|
ebca36ffbb |
@@ -1 +0,0 @@
|
|||||||
../../.agents/skills/component-refactoring
|
|
||||||
@@ -1 +0,0 @@
|
|||||||
../../.agents/skills/frontend-code-review
|
|
||||||
@@ -1 +0,0 @@
|
|||||||
../../.agents/skills/frontend-testing
|
|
||||||
@@ -1 +0,0 @@
|
|||||||
../../.agents/skills/orpc-contract-first
|
|
||||||
@@ -24,6 +24,10 @@
|
|||||||
/api/services/tools/mcp_tools_manage_service.py @Nov1c444
|
/api/services/tools/mcp_tools_manage_service.py @Nov1c444
|
||||||
/api/controllers/mcp/ @Nov1c444
|
/api/controllers/mcp/ @Nov1c444
|
||||||
/api/controllers/console/app/mcp_server.py @Nov1c444
|
/api/controllers/console/app/mcp_server.py @Nov1c444
|
||||||
|
|
||||||
|
# Backend - Tests
|
||||||
|
/api/tests/ @laipz8200 @QuantumGhost
|
||||||
|
|
||||||
/api/tests/**/*mcp* @Nov1c444
|
/api/tests/**/*mcp* @Nov1c444
|
||||||
|
|
||||||
# Backend - Workflow - Engine (Core graph execution engine)
|
# Backend - Workflow - Engine (Core graph execution engine)
|
||||||
@@ -234,6 +238,9 @@
|
|||||||
# Frontend - Base Components
|
# Frontend - Base Components
|
||||||
/web/app/components/base/ @iamjoel @zxhlyh
|
/web/app/components/base/ @iamjoel @zxhlyh
|
||||||
|
|
||||||
|
# Frontend - Base Components Tests
|
||||||
|
/web/app/components/base/**/*.spec.tsx @hyoban @CodingOnStar
|
||||||
|
|
||||||
# Frontend - Utils and Hooks
|
# Frontend - Utils and Hooks
|
||||||
/web/utils/classnames.ts @iamjoel @zxhlyh
|
/web/utils/classnames.ts @iamjoel @zxhlyh
|
||||||
/web/utils/time.ts @iamjoel @zxhlyh
|
/web/utils/time.ts @iamjoel @zxhlyh
|
||||||
|
|||||||
@@ -79,29 +79,6 @@ jobs:
|
|||||||
find . -name "*.py" -type f -exec sed -i.bak -E 's/"([^"]+)" \| None/Optional["\1"]/g; s/'"'"'([^'"'"']+)'"'"' \| None/Optional['"'"'\1'"'"']/g' {} \;
|
find . -name "*.py" -type f -exec sed -i.bak -E 's/"([^"]+)" \| None/Optional["\1"]/g; s/'"'"'([^'"'"']+)'"'"' \| None/Optional['"'"'\1'"'"']/g' {} \;
|
||||||
find . -name "*.py.bak" -type f -delete
|
find . -name "*.py.bak" -type f -delete
|
||||||
|
|
||||||
- name: Install pnpm
|
|
||||||
uses: pnpm/action-setup@v4
|
|
||||||
with:
|
|
||||||
package_json_file: web/package.json
|
|
||||||
run_install: false
|
|
||||||
|
|
||||||
- name: Setup Node.js
|
|
||||||
uses: actions/setup-node@v6
|
|
||||||
with:
|
|
||||||
node-version: 24
|
|
||||||
cache: pnpm
|
|
||||||
cache-dependency-path: ./web/pnpm-lock.yaml
|
|
||||||
|
|
||||||
- name: Install web dependencies
|
|
||||||
run: |
|
|
||||||
cd web
|
|
||||||
pnpm install --frozen-lockfile
|
|
||||||
|
|
||||||
- name: ESLint autofix
|
|
||||||
run: |
|
|
||||||
cd web
|
|
||||||
pnpm lint:fix || true
|
|
||||||
|
|
||||||
# mdformat breaks YAML front matter in markdown files. Add --exclude for directories containing YAML front matter.
|
# mdformat breaks YAML front matter in markdown files. Add --exclude for directories containing YAML front matter.
|
||||||
- name: mdformat
|
- name: mdformat
|
||||||
run: |
|
run: |
|
||||||
|
|||||||
@@ -8,6 +8,7 @@ on:
|
|||||||
- "build/**"
|
- "build/**"
|
||||||
- "release/e-*"
|
- "release/e-*"
|
||||||
- "hotfix/**"
|
- "hotfix/**"
|
||||||
|
- "feat/hitl-backend"
|
||||||
tags:
|
tags:
|
||||||
- "*"
|
- "*"
|
||||||
|
|
||||||
|
|||||||
@@ -4,8 +4,7 @@ on:
|
|||||||
workflow_run:
|
workflow_run:
|
||||||
workflows: ["Build and Push API & Web"]
|
workflows: ["Build and Push API & Web"]
|
||||||
branches:
|
branches:
|
||||||
- "feat/hitl-frontend"
|
- "build/feat/hitl"
|
||||||
- "feat/hitl-backend"
|
|
||||||
types:
|
types:
|
||||||
- completed
|
- completed
|
||||||
|
|
||||||
@@ -14,10 +13,7 @@ jobs:
|
|||||||
runs-on: ubuntu-latest
|
runs-on: ubuntu-latest
|
||||||
if: |
|
if: |
|
||||||
github.event.workflow_run.conclusion == 'success' &&
|
github.event.workflow_run.conclusion == 'success' &&
|
||||||
(
|
github.event.workflow_run.head_branch == 'build/feat/hitl'
|
||||||
github.event.workflow_run.head_branch == 'feat/hitl-frontend' ||
|
|
||||||
github.event.workflow_run.head_branch == 'feat/hitl-backend'
|
|
||||||
)
|
|
||||||
steps:
|
steps:
|
||||||
- name: Deploy to server
|
- name: Deploy to server
|
||||||
uses: appleboy/ssh-action@v1
|
uses: appleboy/ssh-action@v1
|
||||||
|
|||||||
@@ -39,7 +39,7 @@ jobs:
|
|||||||
run: pnpm install --frozen-lockfile
|
run: pnpm install --frozen-lockfile
|
||||||
|
|
||||||
- name: Run tests
|
- name: Run tests
|
||||||
run: pnpm test:coverage
|
run: pnpm test:ci
|
||||||
|
|
||||||
- name: Coverage Summary
|
- name: Coverage Summary
|
||||||
if: always()
|
if: always()
|
||||||
|
|||||||
Vendored
+1
-1
@@ -37,7 +37,7 @@
|
|||||||
"-c",
|
"-c",
|
||||||
"1",
|
"1",
|
||||||
"-Q",
|
"-Q",
|
||||||
"dataset,priority_dataset,priority_pipeline,pipeline,mail,ops_trace,app_deletion,plugin,workflow_storage,conversation,workflow,schedule_poller,schedule_executor,triggered_workflow_dispatcher,trigger_refresh_executor,retention",
|
"dataset,priority_dataset,priority_pipeline,pipeline,mail,ops_trace,app_deletion,plugin,workflow_storage,conversation,workflow,schedule_poller,schedule_executor,triggered_workflow_dispatcher,trigger_refresh_executor,retention,workflow_based_app_execution",
|
||||||
"--loglevel",
|
"--loglevel",
|
||||||
"INFO"
|
"INFO"
|
||||||
],
|
],
|
||||||
|
|||||||
@@ -715,5 +715,31 @@ ANNOTATION_IMPORT_MAX_CONCURRENT=5
|
|||||||
# Sandbox expired records clean configuration
|
# Sandbox expired records clean configuration
|
||||||
SANDBOX_EXPIRED_RECORDS_CLEAN_GRACEFUL_PERIOD=21
|
SANDBOX_EXPIRED_RECORDS_CLEAN_GRACEFUL_PERIOD=21
|
||||||
SANDBOX_EXPIRED_RECORDS_CLEAN_BATCH_SIZE=1000
|
SANDBOX_EXPIRED_RECORDS_CLEAN_BATCH_SIZE=1000
|
||||||
|
SANDBOX_EXPIRED_RECORDS_CLEAN_BATCH_MAX_INTERVAL=200
|
||||||
SANDBOX_EXPIRED_RECORDS_RETENTION_DAYS=30
|
SANDBOX_EXPIRED_RECORDS_RETENTION_DAYS=30
|
||||||
SANDBOX_EXPIRED_RECORDS_CLEAN_TASK_LOCK_TTL=90000
|
SANDBOX_EXPIRED_RECORDS_CLEAN_TASK_LOCK_TTL=90000
|
||||||
|
|
||||||
|
|
||||||
|
# Redis URL used for PubSub between API and
|
||||||
|
# celery worker
|
||||||
|
# defaults to url constructed from `REDIS_*`
|
||||||
|
# configurations
|
||||||
|
PUBSUB_REDIS_URL=
|
||||||
|
# Pub/sub channel type for streaming events.
|
||||||
|
# valid options are:
|
||||||
|
#
|
||||||
|
# - pubsub: for normal Pub/Sub
|
||||||
|
# - sharded: for sharded Pub/Sub
|
||||||
|
#
|
||||||
|
# It's highly recommended to use sharded Pub/Sub AND redis cluster
|
||||||
|
# for large deployments.
|
||||||
|
PUBSUB_REDIS_CHANNEL_TYPE=pubsub
|
||||||
|
# Whether to use Redis cluster mode while running
|
||||||
|
# PubSub.
|
||||||
|
# It's highly recommended to enable this for large deployments.
|
||||||
|
PUBSUB_REDIS_USE_CLUSTERS=false
|
||||||
|
|
||||||
|
# Whether to Enable human input timeout check task
|
||||||
|
ENABLE_HUMAN_INPUT_TIMEOUT_TASK=true
|
||||||
|
# Human input timeout check interval in minutes
|
||||||
|
HUMAN_INPUT_TIMEOUT_TASK_INTERVAL=1
|
||||||
|
|||||||
+8
-16
@@ -36,6 +36,8 @@ ignore_imports =
|
|||||||
core.workflow.nodes.loop.loop_node -> core.workflow.graph_engine
|
core.workflow.nodes.loop.loop_node -> core.workflow.graph_engine
|
||||||
core.workflow.nodes.loop.loop_node -> core.workflow.graph
|
core.workflow.nodes.loop.loop_node -> core.workflow.graph
|
||||||
core.workflow.nodes.loop.loop_node -> core.workflow.graph_engine.command_channels
|
core.workflow.nodes.loop.loop_node -> core.workflow.graph_engine.command_channels
|
||||||
|
# TODO(QuantumGhost): fix the import violation later
|
||||||
|
core.workflow.entities.pause_reason -> core.workflow.nodes.human_input.entities
|
||||||
|
|
||||||
[importlinter:contract:workflow-infrastructure-dependencies]
|
[importlinter:contract:workflow-infrastructure-dependencies]
|
||||||
name = Workflow Infrastructure Dependencies
|
name = Workflow Infrastructure Dependencies
|
||||||
@@ -50,14 +52,14 @@ ignore_imports =
|
|||||||
core.workflow.nodes.agent.agent_node -> extensions.ext_database
|
core.workflow.nodes.agent.agent_node -> extensions.ext_database
|
||||||
core.workflow.nodes.datasource.datasource_node -> extensions.ext_database
|
core.workflow.nodes.datasource.datasource_node -> extensions.ext_database
|
||||||
core.workflow.nodes.knowledge_index.knowledge_index_node -> extensions.ext_database
|
core.workflow.nodes.knowledge_index.knowledge_index_node -> extensions.ext_database
|
||||||
core.workflow.nodes.knowledge_retrieval.knowledge_retrieval_node -> extensions.ext_database
|
|
||||||
core.workflow.nodes.llm.file_saver -> extensions.ext_database
|
core.workflow.nodes.llm.file_saver -> extensions.ext_database
|
||||||
core.workflow.nodes.llm.llm_utils -> extensions.ext_database
|
core.workflow.nodes.llm.llm_utils -> extensions.ext_database
|
||||||
core.workflow.nodes.llm.node -> extensions.ext_database
|
core.workflow.nodes.llm.node -> extensions.ext_database
|
||||||
core.workflow.nodes.tool.tool_node -> extensions.ext_database
|
core.workflow.nodes.tool.tool_node -> extensions.ext_database
|
||||||
core.workflow.graph_engine.command_channels.redis_channel -> extensions.ext_redis
|
core.workflow.graph_engine.command_channels.redis_channel -> extensions.ext_redis
|
||||||
core.workflow.graph_engine.manager -> extensions.ext_redis
|
core.workflow.graph_engine.manager -> extensions.ext_redis
|
||||||
core.workflow.nodes.knowledge_retrieval.knowledge_retrieval_node -> extensions.ext_redis
|
# TODO(QuantumGhost): use DI to avoid depending on global DB.
|
||||||
|
core.workflow.nodes.human_input.human_input_node -> extensions.ext_database
|
||||||
|
|
||||||
[importlinter:contract:workflow-external-imports]
|
[importlinter:contract:workflow-external-imports]
|
||||||
name = Workflow External Imports
|
name = Workflow External Imports
|
||||||
@@ -122,11 +124,6 @@ ignore_imports =
|
|||||||
core.workflow.nodes.http_request.node -> core.tools.tool_file_manager
|
core.workflow.nodes.http_request.node -> core.tools.tool_file_manager
|
||||||
core.workflow.nodes.iteration.iteration_node -> core.app.workflow.node_factory
|
core.workflow.nodes.iteration.iteration_node -> core.app.workflow.node_factory
|
||||||
core.workflow.nodes.knowledge_index.knowledge_index_node -> core.rag.index_processor.index_processor_factory
|
core.workflow.nodes.knowledge_index.knowledge_index_node -> core.rag.index_processor.index_processor_factory
|
||||||
core.workflow.nodes.knowledge_retrieval.knowledge_retrieval_node -> core.rag.datasource.retrieval_service
|
|
||||||
core.workflow.nodes.knowledge_retrieval.knowledge_retrieval_node -> core.rag.retrieval.dataset_retrieval
|
|
||||||
core.workflow.nodes.knowledge_retrieval.knowledge_retrieval_node -> models.dataset
|
|
||||||
core.workflow.nodes.knowledge_retrieval.knowledge_retrieval_node -> services.feature_service
|
|
||||||
core.workflow.nodes.knowledge_retrieval.knowledge_retrieval_node -> core.model_runtime.model_providers.__base.large_language_model
|
|
||||||
core.workflow.nodes.llm.llm_utils -> configs
|
core.workflow.nodes.llm.llm_utils -> configs
|
||||||
core.workflow.nodes.llm.llm_utils -> core.app.entities.app_invoke_entities
|
core.workflow.nodes.llm.llm_utils -> core.app.entities.app_invoke_entities
|
||||||
core.workflow.nodes.llm.llm_utils -> core.file.models
|
core.workflow.nodes.llm.llm_utils -> core.file.models
|
||||||
@@ -136,7 +133,6 @@ ignore_imports =
|
|||||||
core.workflow.nodes.llm.llm_utils -> models.provider
|
core.workflow.nodes.llm.llm_utils -> models.provider
|
||||||
core.workflow.nodes.llm.llm_utils -> services.credit_pool_service
|
core.workflow.nodes.llm.llm_utils -> services.credit_pool_service
|
||||||
core.workflow.nodes.llm.node -> core.tools.signature
|
core.workflow.nodes.llm.node -> core.tools.signature
|
||||||
core.workflow.nodes.template_transform.template_transform_node -> configs
|
|
||||||
core.workflow.nodes.tool.tool_node -> core.callback_handler.workflow_tool_callback_handler
|
core.workflow.nodes.tool.tool_node -> core.callback_handler.workflow_tool_callback_handler
|
||||||
core.workflow.nodes.tool.tool_node -> core.tools.tool_engine
|
core.workflow.nodes.tool.tool_node -> core.tools.tool_engine
|
||||||
core.workflow.nodes.tool.tool_node -> core.tools.tool_manager
|
core.workflow.nodes.tool.tool_node -> core.tools.tool_manager
|
||||||
@@ -145,9 +141,9 @@ ignore_imports =
|
|||||||
core.workflow.nodes.agent.agent_node -> core.agent.entities
|
core.workflow.nodes.agent.agent_node -> core.agent.entities
|
||||||
core.workflow.nodes.agent.agent_node -> core.agent.plugin_entities
|
core.workflow.nodes.agent.agent_node -> core.agent.plugin_entities
|
||||||
core.workflow.nodes.base.node -> core.app.entities.app_invoke_entities
|
core.workflow.nodes.base.node -> core.app.entities.app_invoke_entities
|
||||||
|
core.workflow.nodes.human_input.human_input_node -> core.app.entities.app_invoke_entities
|
||||||
core.workflow.nodes.knowledge_index.knowledge_index_node -> core.app.entities.app_invoke_entities
|
core.workflow.nodes.knowledge_index.knowledge_index_node -> core.app.entities.app_invoke_entities
|
||||||
core.workflow.nodes.knowledge_retrieval.knowledge_retrieval_node -> core.app.app_config.entities
|
core.workflow.nodes.knowledge_retrieval.knowledge_retrieval_node -> core.app.app_config.entities
|
||||||
core.workflow.nodes.knowledge_retrieval.knowledge_retrieval_node -> core.app.entities.app_invoke_entities
|
|
||||||
core.workflow.nodes.llm.node -> core.app.entities.app_invoke_entities
|
core.workflow.nodes.llm.node -> core.app.entities.app_invoke_entities
|
||||||
core.workflow.nodes.parameter_extractor.parameter_extractor_node -> core.app.entities.app_invoke_entities
|
core.workflow.nodes.parameter_extractor.parameter_extractor_node -> core.app.entities.app_invoke_entities
|
||||||
core.workflow.nodes.parameter_extractor.parameter_extractor_node -> core.prompt.advanced_prompt_transform
|
core.workflow.nodes.parameter_extractor.parameter_extractor_node -> core.prompt.advanced_prompt_transform
|
||||||
@@ -163,9 +159,6 @@ ignore_imports =
|
|||||||
core.workflow.workflow_entry -> core.app.workflow.node_factory
|
core.workflow.workflow_entry -> core.app.workflow.node_factory
|
||||||
core.workflow.nodes.datasource.datasource_node -> core.datasource.datasource_manager
|
core.workflow.nodes.datasource.datasource_node -> core.datasource.datasource_manager
|
||||||
core.workflow.nodes.datasource.datasource_node -> core.datasource.utils.message_transformer
|
core.workflow.nodes.datasource.datasource_node -> core.datasource.utils.message_transformer
|
||||||
core.workflow.nodes.knowledge_retrieval.knowledge_retrieval_node -> core.entities.agent_entities
|
|
||||||
core.workflow.nodes.knowledge_retrieval.knowledge_retrieval_node -> core.entities.model_entities
|
|
||||||
core.workflow.nodes.knowledge_retrieval.knowledge_retrieval_node -> core.model_manager
|
|
||||||
core.workflow.nodes.llm.llm_utils -> core.entities.provider_entities
|
core.workflow.nodes.llm.llm_utils -> core.entities.provider_entities
|
||||||
core.workflow.nodes.parameter_extractor.parameter_extractor_node -> core.model_manager
|
core.workflow.nodes.parameter_extractor.parameter_extractor_node -> core.model_manager
|
||||||
core.workflow.nodes.question_classifier.question_classifier_node -> core.model_manager
|
core.workflow.nodes.question_classifier.question_classifier_node -> core.model_manager
|
||||||
@@ -214,7 +207,6 @@ ignore_imports =
|
|||||||
core.workflow.nodes.llm.node -> core.llm_generator.output_parser.structured_output
|
core.workflow.nodes.llm.node -> core.llm_generator.output_parser.structured_output
|
||||||
core.workflow.nodes.llm.node -> core.model_manager
|
core.workflow.nodes.llm.node -> core.model_manager
|
||||||
core.workflow.nodes.agent.entities -> core.prompt.entities.advanced_prompt_entities
|
core.workflow.nodes.agent.entities -> core.prompt.entities.advanced_prompt_entities
|
||||||
core.workflow.nodes.knowledge_retrieval.knowledge_retrieval_node -> core.prompt.simple_prompt_transform
|
|
||||||
core.workflow.nodes.llm.entities -> core.prompt.entities.advanced_prompt_entities
|
core.workflow.nodes.llm.entities -> core.prompt.entities.advanced_prompt_entities
|
||||||
core.workflow.nodes.llm.llm_utils -> core.prompt.entities.advanced_prompt_entities
|
core.workflow.nodes.llm.llm_utils -> core.prompt.entities.advanced_prompt_entities
|
||||||
core.workflow.nodes.llm.node -> core.prompt.entities.advanced_prompt_entities
|
core.workflow.nodes.llm.node -> core.prompt.entities.advanced_prompt_entities
|
||||||
@@ -230,7 +222,6 @@ ignore_imports =
|
|||||||
core.workflow.nodes.knowledge_index.knowledge_index_node -> services.summary_index_service
|
core.workflow.nodes.knowledge_index.knowledge_index_node -> services.summary_index_service
|
||||||
core.workflow.nodes.knowledge_index.knowledge_index_node -> tasks.generate_summary_index_task
|
core.workflow.nodes.knowledge_index.knowledge_index_node -> tasks.generate_summary_index_task
|
||||||
core.workflow.nodes.knowledge_index.knowledge_index_node -> core.rag.index_processor.processor.paragraph_index_processor
|
core.workflow.nodes.knowledge_index.knowledge_index_node -> core.rag.index_processor.processor.paragraph_index_processor
|
||||||
core.workflow.nodes.knowledge_retrieval.knowledge_retrieval_node -> core.rag.retrieval.retrieval_methods
|
|
||||||
core.workflow.nodes.llm.node -> models.dataset
|
core.workflow.nodes.llm.node -> models.dataset
|
||||||
core.workflow.nodes.agent.agent_node -> core.tools.utils.message_transformer
|
core.workflow.nodes.agent.agent_node -> core.tools.utils.message_transformer
|
||||||
core.workflow.nodes.llm.file_saver -> core.tools.signature
|
core.workflow.nodes.llm.file_saver -> core.tools.signature
|
||||||
@@ -248,6 +239,7 @@ ignore_imports =
|
|||||||
core.workflow.nodes.document_extractor.node -> core.variables.segments
|
core.workflow.nodes.document_extractor.node -> core.variables.segments
|
||||||
core.workflow.nodes.http_request.executor -> core.variables.segments
|
core.workflow.nodes.http_request.executor -> core.variables.segments
|
||||||
core.workflow.nodes.http_request.node -> core.variables.segments
|
core.workflow.nodes.http_request.node -> core.variables.segments
|
||||||
|
core.workflow.nodes.human_input.entities -> core.variables.consts
|
||||||
core.workflow.nodes.iteration.iteration_node -> core.variables
|
core.workflow.nodes.iteration.iteration_node -> core.variables
|
||||||
core.workflow.nodes.iteration.iteration_node -> core.variables.segments
|
core.workflow.nodes.iteration.iteration_node -> core.variables.segments
|
||||||
core.workflow.nodes.iteration.iteration_node -> core.variables.variables
|
core.workflow.nodes.iteration.iteration_node -> core.variables.variables
|
||||||
@@ -288,12 +280,12 @@ ignore_imports =
|
|||||||
core.workflow.nodes.agent.agent_node -> extensions.ext_database
|
core.workflow.nodes.agent.agent_node -> extensions.ext_database
|
||||||
core.workflow.nodes.datasource.datasource_node -> extensions.ext_database
|
core.workflow.nodes.datasource.datasource_node -> extensions.ext_database
|
||||||
core.workflow.nodes.knowledge_index.knowledge_index_node -> extensions.ext_database
|
core.workflow.nodes.knowledge_index.knowledge_index_node -> extensions.ext_database
|
||||||
core.workflow.nodes.knowledge_retrieval.knowledge_retrieval_node -> extensions.ext_database
|
|
||||||
core.workflow.nodes.knowledge_retrieval.knowledge_retrieval_node -> extensions.ext_redis
|
|
||||||
core.workflow.nodes.llm.file_saver -> extensions.ext_database
|
core.workflow.nodes.llm.file_saver -> extensions.ext_database
|
||||||
core.workflow.nodes.llm.llm_utils -> extensions.ext_database
|
core.workflow.nodes.llm.llm_utils -> extensions.ext_database
|
||||||
core.workflow.nodes.llm.node -> extensions.ext_database
|
core.workflow.nodes.llm.node -> extensions.ext_database
|
||||||
core.workflow.nodes.tool.tool_node -> extensions.ext_database
|
core.workflow.nodes.tool.tool_node -> extensions.ext_database
|
||||||
|
core.workflow.nodes.human_input.human_input_node -> extensions.ext_database
|
||||||
|
core.workflow.nodes.human_input.human_input_node -> core.repositories.human_input_repository
|
||||||
core.workflow.workflow_entry -> extensions.otel.runtime
|
core.workflow.workflow_entry -> extensions.otel.runtime
|
||||||
core.workflow.nodes.agent.agent_node -> models
|
core.workflow.nodes.agent.agent_node -> models
|
||||||
core.workflow.nodes.base.node -> models.enums
|
core.workflow.nodes.base.node -> models.enums
|
||||||
|
|||||||
Vendored
+1
-1
@@ -54,7 +54,7 @@
|
|||||||
"--loglevel",
|
"--loglevel",
|
||||||
"DEBUG",
|
"DEBUG",
|
||||||
"-Q",
|
"-Q",
|
||||||
"dataset,priority_pipeline,pipeline,mail,ops_trace,app_deletion,plugin,workflow_storage,conversation,workflow,schedule_poller,schedule_executor,triggered_workflow_dispatcher,trigger_refresh_executor"
|
"dataset,priority_pipeline,pipeline,mail,ops_trace,app_deletion,plugin,workflow_storage,conversation,workflow,workflow_based_app_execution,schedule_poller,schedule_executor,triggered_workflow_dispatcher,trigger_refresh_executor"
|
||||||
]
|
]
|
||||||
}
|
}
|
||||||
]
|
]
|
||||||
|
|||||||
+1
-1
@@ -122,7 +122,7 @@ These commands assume you start from the repository root.
|
|||||||
|
|
||||||
```bash
|
```bash
|
||||||
cd api
|
cd api
|
||||||
uv run celery -A app.celery worker -P threads -c 2 --loglevel INFO -Q dataset,priority_dataset,priority_pipeline,pipeline,mail,ops_trace,app_deletion,plugin,workflow_storage,conversation,workflow,schedule_poller,schedule_executor,triggered_workflow_dispatcher,trigger_refresh_executor,retention
|
uv run celery -A app.celery worker -P threads -c 2 --loglevel INFO -Q api_token,dataset,priority_dataset,priority_pipeline,pipeline,mail,ops_trace,app_deletion,plugin,workflow_storage,conversation,workflow,schedule_poller,schedule_executor,triggered_workflow_dispatcher,trigger_refresh_executor,retention
|
||||||
```
|
```
|
||||||
|
|
||||||
1. Optional: start Celery Beat (scheduled tasks, in a new terminal).
|
1. Optional: start Celery Beat (scheduled tasks, in a new terminal).
|
||||||
|
|||||||
+3
-1
@@ -739,8 +739,10 @@ def upgrade_db():
|
|||||||
|
|
||||||
click.echo(click.style("Database migration successful!", fg="green"))
|
click.echo(click.style("Database migration successful!", fg="green"))
|
||||||
|
|
||||||
except Exception:
|
except Exception as e:
|
||||||
logger.exception("Failed to execute database migration")
|
logger.exception("Failed to execute database migration")
|
||||||
|
click.echo(click.style(f"Database migration failed: {e}", fg="red"))
|
||||||
|
raise SystemExit(1)
|
||||||
finally:
|
finally:
|
||||||
lock.release()
|
lock.release()
|
||||||
else:
|
else:
|
||||||
|
|||||||
@@ -1,3 +1,4 @@
|
|||||||
|
from datetime import timedelta
|
||||||
from enum import StrEnum
|
from enum import StrEnum
|
||||||
from typing import Literal
|
from typing import Literal
|
||||||
|
|
||||||
@@ -48,6 +49,16 @@ class SecurityConfig(BaseSettings):
|
|||||||
default=5,
|
default=5,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
WEB_FORM_SUBMIT_RATE_LIMIT_MAX_ATTEMPTS: PositiveInt = Field(
|
||||||
|
description="Maximum number of web form submissions allowed per IP within the rate limit window",
|
||||||
|
default=30,
|
||||||
|
)
|
||||||
|
|
||||||
|
WEB_FORM_SUBMIT_RATE_LIMIT_WINDOW_SECONDS: PositiveInt = Field(
|
||||||
|
description="Time window in seconds for web form submission rate limiting",
|
||||||
|
default=60,
|
||||||
|
)
|
||||||
|
|
||||||
LOGIN_DISABLED: bool = Field(
|
LOGIN_DISABLED: bool = Field(
|
||||||
description="Whether to disable login checks",
|
description="Whether to disable login checks",
|
||||||
default=False,
|
default=False,
|
||||||
@@ -82,6 +93,12 @@ class AppExecutionConfig(BaseSettings):
|
|||||||
default=0,
|
default=0,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
HUMAN_INPUT_GLOBAL_TIMEOUT_SECONDS: PositiveInt = Field(
|
||||||
|
description="Maximum seconds a workflow run can stay paused waiting for human input before global timeout.",
|
||||||
|
default=int(timedelta(days=7).total_seconds()),
|
||||||
|
ge=1,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class CodeExecutionSandboxConfig(BaseSettings):
|
class CodeExecutionSandboxConfig(BaseSettings):
|
||||||
"""
|
"""
|
||||||
@@ -1134,6 +1151,14 @@ class CeleryScheduleTasksConfig(BaseSettings):
|
|||||||
description="Enable queue monitor task",
|
description="Enable queue monitor task",
|
||||||
default=False,
|
default=False,
|
||||||
)
|
)
|
||||||
|
ENABLE_HUMAN_INPUT_TIMEOUT_TASK: bool = Field(
|
||||||
|
description="Enable human input timeout check task",
|
||||||
|
default=True,
|
||||||
|
)
|
||||||
|
HUMAN_INPUT_TIMEOUT_TASK_INTERVAL: PositiveInt = Field(
|
||||||
|
description="Human input timeout check interval in minutes",
|
||||||
|
default=1,
|
||||||
|
)
|
||||||
ENABLE_CHECK_UPGRADABLE_PLUGIN_TASK: bool = Field(
|
ENABLE_CHECK_UPGRADABLE_PLUGIN_TASK: bool = Field(
|
||||||
description="Enable check upgradable plugin task",
|
description="Enable check upgradable plugin task",
|
||||||
default=True,
|
default=True,
|
||||||
@@ -1155,6 +1180,16 @@ class CeleryScheduleTasksConfig(BaseSettings):
|
|||||||
default=0,
|
default=0,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# API token last_used_at batch update
|
||||||
|
ENABLE_API_TOKEN_LAST_USED_UPDATE_TASK: bool = Field(
|
||||||
|
description="Enable periodic batch update of API token last_used_at timestamps",
|
||||||
|
default=True,
|
||||||
|
)
|
||||||
|
API_TOKEN_LAST_USED_UPDATE_INTERVAL: int = Field(
|
||||||
|
description="Interval in minutes for batch updating API token last_used_at (default 30)",
|
||||||
|
default=30,
|
||||||
|
)
|
||||||
|
|
||||||
# Trigger provider refresh (simple version)
|
# Trigger provider refresh (simple version)
|
||||||
ENABLE_TRIGGER_PROVIDER_REFRESH_TASK: bool = Field(
|
ENABLE_TRIGGER_PROVIDER_REFRESH_TASK: bool = Field(
|
||||||
description="Enable trigger provider refresh poller",
|
description="Enable trigger provider refresh poller",
|
||||||
@@ -1309,6 +1344,10 @@ class SandboxExpiredRecordsCleanConfig(BaseSettings):
|
|||||||
description="Maximum number of records to process in each batch",
|
description="Maximum number of records to process in each batch",
|
||||||
default=1000,
|
default=1000,
|
||||||
)
|
)
|
||||||
|
SANDBOX_EXPIRED_RECORDS_CLEAN_BATCH_MAX_INTERVAL: PositiveInt = Field(
|
||||||
|
description="Maximum interval in milliseconds between batches",
|
||||||
|
default=200,
|
||||||
|
)
|
||||||
SANDBOX_EXPIRED_RECORDS_RETENTION_DAYS: PositiveInt = Field(
|
SANDBOX_EXPIRED_RECORDS_RETENTION_DAYS: PositiveInt = Field(
|
||||||
description="Retention days for sandbox expired workflow_run records and message records",
|
description="Retention days for sandbox expired workflow_run records and message records",
|
||||||
default=30,
|
default=30,
|
||||||
|
|||||||
@@ -6,6 +6,7 @@ from pydantic import Field, NonNegativeFloat, NonNegativeInt, PositiveFloat, Pos
|
|||||||
from pydantic_settings import BaseSettings
|
from pydantic_settings import BaseSettings
|
||||||
|
|
||||||
from .cache.redis_config import RedisConfig
|
from .cache.redis_config import RedisConfig
|
||||||
|
from .cache.redis_pubsub_config import RedisPubSubConfig
|
||||||
from .storage.aliyun_oss_storage_config import AliyunOSSStorageConfig
|
from .storage.aliyun_oss_storage_config import AliyunOSSStorageConfig
|
||||||
from .storage.amazon_s3_storage_config import S3StorageConfig
|
from .storage.amazon_s3_storage_config import S3StorageConfig
|
||||||
from .storage.azure_blob_storage_config import AzureBlobStorageConfig
|
from .storage.azure_blob_storage_config import AzureBlobStorageConfig
|
||||||
@@ -258,11 +259,20 @@ class CeleryConfig(DatabaseConfig):
|
|||||||
description="Password of the Redis Sentinel master.",
|
description="Password of the Redis Sentinel master.",
|
||||||
default=None,
|
default=None,
|
||||||
)
|
)
|
||||||
|
|
||||||
CELERY_SENTINEL_SOCKET_TIMEOUT: PositiveFloat | None = Field(
|
CELERY_SENTINEL_SOCKET_TIMEOUT: PositiveFloat | None = Field(
|
||||||
description="Timeout for Redis Sentinel socket operations in seconds.",
|
description="Timeout for Redis Sentinel socket operations in seconds.",
|
||||||
default=0.1,
|
default=0.1,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
CELERY_TASK_ANNOTATIONS: dict[str, Any] | None = Field(
|
||||||
|
description=(
|
||||||
|
"Annotations for Celery tasks as a JSON mapping of task name -> options "
|
||||||
|
"(for example, rate limits or other task-specific settings)."
|
||||||
|
),
|
||||||
|
default=None,
|
||||||
|
)
|
||||||
|
|
||||||
@computed_field
|
@computed_field
|
||||||
def CELERY_RESULT_BACKEND(self) -> str | None:
|
def CELERY_RESULT_BACKEND(self) -> str | None:
|
||||||
if self.CELERY_BACKEND in ("database", "rabbitmq"):
|
if self.CELERY_BACKEND in ("database", "rabbitmq"):
|
||||||
@@ -317,6 +327,7 @@ class MiddlewareConfig(
|
|||||||
CeleryConfig, # Note: CeleryConfig already inherits from DatabaseConfig
|
CeleryConfig, # Note: CeleryConfig already inherits from DatabaseConfig
|
||||||
KeywordStoreConfig,
|
KeywordStoreConfig,
|
||||||
RedisConfig,
|
RedisConfig,
|
||||||
|
RedisPubSubConfig,
|
||||||
# configs of storage and storage providers
|
# configs of storage and storage providers
|
||||||
StorageConfig,
|
StorageConfig,
|
||||||
AliyunOSSStorageConfig,
|
AliyunOSSStorageConfig,
|
||||||
|
|||||||
@@ -0,0 +1,96 @@
|
|||||||
|
from typing import Literal, Protocol
|
||||||
|
from urllib.parse import quote_plus, urlunparse
|
||||||
|
|
||||||
|
from pydantic import Field
|
||||||
|
from pydantic_settings import BaseSettings
|
||||||
|
|
||||||
|
|
||||||
|
class RedisConfigDefaults(Protocol):
|
||||||
|
REDIS_HOST: str
|
||||||
|
REDIS_PORT: int
|
||||||
|
REDIS_USERNAME: str | None
|
||||||
|
REDIS_PASSWORD: str | None
|
||||||
|
REDIS_DB: int
|
||||||
|
REDIS_USE_SSL: bool
|
||||||
|
REDIS_USE_SENTINEL: bool | None
|
||||||
|
REDIS_USE_CLUSTERS: bool
|
||||||
|
|
||||||
|
|
||||||
|
class RedisConfigDefaultsMixin:
|
||||||
|
def _redis_defaults(self: RedisConfigDefaults) -> RedisConfigDefaults:
|
||||||
|
return self
|
||||||
|
|
||||||
|
|
||||||
|
class RedisPubSubConfig(BaseSettings, RedisConfigDefaultsMixin):
|
||||||
|
"""
|
||||||
|
Configuration settings for Redis pub/sub streaming.
|
||||||
|
"""
|
||||||
|
|
||||||
|
PUBSUB_REDIS_URL: str | None = Field(
|
||||||
|
alias="PUBSUB_REDIS_URL",
|
||||||
|
description=(
|
||||||
|
"Redis connection URL for pub/sub streaming events between API "
|
||||||
|
"and celery worker, defaults to url constructed from "
|
||||||
|
"`REDIS_*` configurations"
|
||||||
|
),
|
||||||
|
default=None,
|
||||||
|
)
|
||||||
|
|
||||||
|
PUBSUB_REDIS_USE_CLUSTERS: bool = Field(
|
||||||
|
description=(
|
||||||
|
"Enable Redis Cluster mode for pub/sub streaming. It's highly "
|
||||||
|
"recommended to enable this for large deployments."
|
||||||
|
),
|
||||||
|
default=False,
|
||||||
|
)
|
||||||
|
|
||||||
|
PUBSUB_REDIS_CHANNEL_TYPE: Literal["pubsub", "sharded"] = Field(
|
||||||
|
description=(
|
||||||
|
"Pub/sub channel type for streaming events. "
|
||||||
|
"Valid options are:\n"
|
||||||
|
"\n"
|
||||||
|
" - pubsub: for normal Pub/Sub\n"
|
||||||
|
" - sharded: for sharded Pub/Sub\n"
|
||||||
|
"\n"
|
||||||
|
"It's highly recommended to use sharded Pub/Sub AND redis cluster "
|
||||||
|
"for large deployments."
|
||||||
|
),
|
||||||
|
default="pubsub",
|
||||||
|
)
|
||||||
|
|
||||||
|
def _build_default_pubsub_url(self) -> str:
|
||||||
|
defaults = self._redis_defaults()
|
||||||
|
if not defaults.REDIS_HOST or not defaults.REDIS_PORT:
|
||||||
|
raise ValueError("PUBSUB_REDIS_URL must be set when default Redis URL cannot be constructed")
|
||||||
|
|
||||||
|
scheme = "rediss" if defaults.REDIS_USE_SSL else "redis"
|
||||||
|
username = defaults.REDIS_USERNAME or None
|
||||||
|
password = defaults.REDIS_PASSWORD or None
|
||||||
|
|
||||||
|
userinfo = ""
|
||||||
|
if username:
|
||||||
|
userinfo = quote_plus(username)
|
||||||
|
if password:
|
||||||
|
password_part = quote_plus(password)
|
||||||
|
userinfo = f"{userinfo}:{password_part}" if userinfo else f":{password_part}"
|
||||||
|
if userinfo:
|
||||||
|
userinfo = f"{userinfo}@"
|
||||||
|
|
||||||
|
host = defaults.REDIS_HOST
|
||||||
|
port = defaults.REDIS_PORT
|
||||||
|
db = defaults.REDIS_DB
|
||||||
|
|
||||||
|
netloc = f"{userinfo}{host}:{port}"
|
||||||
|
return urlunparse((scheme, netloc, f"/{db}", "", "", ""))
|
||||||
|
|
||||||
|
@property
|
||||||
|
def normalized_pubsub_redis_url(self) -> str:
|
||||||
|
pubsub_redis_url = self.PUBSUB_REDIS_URL
|
||||||
|
if pubsub_redis_url:
|
||||||
|
cleaned = pubsub_redis_url.strip()
|
||||||
|
pubsub_redis_url = cleaned or None
|
||||||
|
|
||||||
|
if pubsub_redis_url:
|
||||||
|
return pubsub_redis_url
|
||||||
|
|
||||||
|
return self._build_default_pubsub_url()
|
||||||
@@ -21,6 +21,7 @@ language_timezone_mapping = {
|
|||||||
"th-TH": "Asia/Bangkok",
|
"th-TH": "Asia/Bangkok",
|
||||||
"id-ID": "Asia/Jakarta",
|
"id-ID": "Asia/Jakarta",
|
||||||
"ar-TN": "Africa/Tunis",
|
"ar-TN": "Africa/Tunis",
|
||||||
|
"nl-NL": "Europe/Amsterdam",
|
||||||
}
|
}
|
||||||
|
|
||||||
languages = list(language_timezone_mapping.keys())
|
languages = list(language_timezone_mapping.keys())
|
||||||
|
|||||||
@@ -5,8 +5,6 @@ from enum import StrEnum
|
|||||||
from flask_restx import Namespace
|
from flask_restx import Namespace
|
||||||
from pydantic import BaseModel, TypeAdapter
|
from pydantic import BaseModel, TypeAdapter
|
||||||
|
|
||||||
from controllers.console import console_ns
|
|
||||||
|
|
||||||
DEFAULT_REF_TEMPLATE_SWAGGER_2_0 = "#/definitions/{model}"
|
DEFAULT_REF_TEMPLATE_SWAGGER_2_0 = "#/definitions/{model}"
|
||||||
|
|
||||||
|
|
||||||
@@ -24,6 +22,9 @@ def register_schema_models(namespace: Namespace, *models: type[BaseModel]) -> No
|
|||||||
|
|
||||||
|
|
||||||
def get_or_create_model(model_name: str, field_def):
|
def get_or_create_model(model_name: str, field_def):
|
||||||
|
# Import lazily to avoid circular imports between console controllers and schema helpers.
|
||||||
|
from controllers.console import console_ns
|
||||||
|
|
||||||
existing = console_ns.models.get(model_name)
|
existing = console_ns.models.get(model_name)
|
||||||
if existing is None:
|
if existing is None:
|
||||||
existing = console_ns.model(model_name, field_def)
|
existing = console_ns.model(model_name, field_def)
|
||||||
|
|||||||
@@ -37,6 +37,7 @@ from . import (
|
|||||||
apikey,
|
apikey,
|
||||||
extension,
|
extension,
|
||||||
feature,
|
feature,
|
||||||
|
human_input_form,
|
||||||
init_validate,
|
init_validate,
|
||||||
ping,
|
ping,
|
||||||
setup,
|
setup,
|
||||||
@@ -171,6 +172,7 @@ __all__ = [
|
|||||||
"forgot_password",
|
"forgot_password",
|
||||||
"generator",
|
"generator",
|
||||||
"hit_testing",
|
"hit_testing",
|
||||||
|
"human_input_form",
|
||||||
"init_validate",
|
"init_validate",
|
||||||
"installed_app",
|
"installed_app",
|
||||||
"load_balancing_config",
|
"load_balancing_config",
|
||||||
|
|||||||
@@ -10,6 +10,7 @@ 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.model import ApiToken, App
|
from models.model import ApiToken, App
|
||||||
|
from services.api_token_service import ApiTokenCache
|
||||||
|
|
||||||
from . import console_ns
|
from . import console_ns
|
||||||
from .wraps import account_initialization_required, edit_permission_required, setup_required
|
from .wraps import account_initialization_required, edit_permission_required, setup_required
|
||||||
@@ -131,6 +132,11 @@ class BaseApiKeyResource(Resource):
|
|||||||
if key is None:
|
if key is None:
|
||||||
flask_restx.abort(HTTPStatus.NOT_FOUND, message="API key not found")
|
flask_restx.abort(HTTPStatus.NOT_FOUND, message="API key not found")
|
||||||
|
|
||||||
|
# Invalidate cache before deleting from database
|
||||||
|
# Type assertion: key is guaranteed to be non-None here because abort() raises
|
||||||
|
assert key is not None # nosec - for type checker only
|
||||||
|
ApiTokenCache.delete(key.token, key.type)
|
||||||
|
|
||||||
db.session.query(ApiToken).where(ApiToken.id == api_key_id).delete()
|
db.session.query(ApiToken).where(ApiToken.id == api_key_id).delete()
|
||||||
db.session.commit()
|
db.session.commit()
|
||||||
|
|
||||||
|
|||||||
@@ -1,3 +1,4 @@
|
|||||||
|
import logging
|
||||||
import uuid
|
import uuid
|
||||||
from datetime import datetime
|
from datetime import datetime
|
||||||
from typing import Any, Literal, TypeAlias
|
from typing import Any, Literal, TypeAlias
|
||||||
@@ -54,6 +55,8 @@ ALLOW_CREATE_APP_MODES = ["chat", "agent-chat", "advanced-chat", "workflow", "co
|
|||||||
|
|
||||||
register_enum_models(console_ns, IconType)
|
register_enum_models(console_ns, IconType)
|
||||||
|
|
||||||
|
_logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
class AppListQuery(BaseModel):
|
class AppListQuery(BaseModel):
|
||||||
page: int = Field(default=1, ge=1, le=99999, description="Page number (1-99999)")
|
page: int = Field(default=1, ge=1, le=99999, description="Page number (1-99999)")
|
||||||
@@ -499,6 +502,7 @@ class AppListApi(Resource):
|
|||||||
select(Workflow).where(
|
select(Workflow).where(
|
||||||
Workflow.version == Workflow.VERSION_DRAFT,
|
Workflow.version == Workflow.VERSION_DRAFT,
|
||||||
Workflow.app_id.in_(workflow_capable_app_ids),
|
Workflow.app_id.in_(workflow_capable_app_ids),
|
||||||
|
Workflow.tenant_id == current_tenant_id,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
.scalars()
|
.scalars()
|
||||||
@@ -510,12 +514,14 @@ class AppListApi(Resource):
|
|||||||
NodeType.TRIGGER_PLUGIN,
|
NodeType.TRIGGER_PLUGIN,
|
||||||
}
|
}
|
||||||
for workflow in draft_workflows:
|
for workflow in draft_workflows:
|
||||||
|
node_id = None
|
||||||
try:
|
try:
|
||||||
for _, node_data in workflow.walk_nodes():
|
for node_id, node_data in workflow.walk_nodes():
|
||||||
if node_data.get("type") in trigger_node_types:
|
if node_data.get("type") in trigger_node_types:
|
||||||
draft_trigger_app_ids.add(str(workflow.app_id))
|
draft_trigger_app_ids.add(str(workflow.app_id))
|
||||||
break
|
break
|
||||||
except Exception:
|
except Exception:
|
||||||
|
_logger.exception("error while walking nodes, workflow_id=%s, node_id=%s", workflow.id, node_id)
|
||||||
continue
|
continue
|
||||||
|
|
||||||
for app in app_pagination.items:
|
for app in app_pagination.items:
|
||||||
|
|||||||
@@ -89,6 +89,7 @@ status_count_model = console_ns.model(
|
|||||||
"success": fields.Integer,
|
"success": fields.Integer,
|
||||||
"failed": fields.Integer,
|
"failed": fields.Integer,
|
||||||
"partial_success": fields.Integer,
|
"partial_success": fields.Integer,
|
||||||
|
"paused": fields.Integer,
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -598,7 +599,12 @@ def _get_conversation(app_model, conversation_id):
|
|||||||
db.session.execute(
|
db.session.execute(
|
||||||
sa.update(Conversation)
|
sa.update(Conversation)
|
||||||
.where(Conversation.id == conversation_id, Conversation.read_at.is_(None))
|
.where(Conversation.id == conversation_id, Conversation.read_at.is_(None))
|
||||||
.values(read_at=naive_utc_now(), read_account_id=current_user.id)
|
# Keep updated_at unchanged when only marking a conversation as read.
|
||||||
|
.values(
|
||||||
|
read_at=naive_utc_now(),
|
||||||
|
read_account_id=current_user.id,
|
||||||
|
updated_at=Conversation.updated_at,
|
||||||
|
)
|
||||||
)
|
)
|
||||||
db.session.commit()
|
db.session.commit()
|
||||||
db.session.refresh(conversation)
|
db.session.refresh(conversation)
|
||||||
|
|||||||
@@ -33,7 +33,7 @@ from libs.login import current_account_with_tenant, login_required
|
|||||||
from models.model import AppMode, Conversation, Message, MessageAnnotation, MessageFeedback
|
from models.model import AppMode, Conversation, Message, MessageAnnotation, MessageFeedback
|
||||||
from services.errors.conversation import ConversationNotExistsError
|
from services.errors.conversation import ConversationNotExistsError
|
||||||
from services.errors.message import MessageNotExistsError, SuggestedQuestionsAfterAnswerDisabledError
|
from services.errors.message import MessageNotExistsError, SuggestedQuestionsAfterAnswerDisabledError
|
||||||
from services.message_service import MessageService
|
from services.message_service import MessageService, attach_message_extra_contents
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
@@ -207,6 +207,7 @@ message_detail_model = console_ns.model(
|
|||||||
"created_at": TimestampField,
|
"created_at": TimestampField,
|
||||||
"agent_thoughts": fields.List(fields.Nested(agent_thought_model)),
|
"agent_thoughts": fields.List(fields.Nested(agent_thought_model)),
|
||||||
"message_files": fields.List(fields.Nested(message_file_model)),
|
"message_files": fields.List(fields.Nested(message_file_model)),
|
||||||
|
"extra_contents": fields.List(fields.Raw),
|
||||||
"metadata": fields.Raw(attribute="message_metadata_dict"),
|
"metadata": fields.Raw(attribute="message_metadata_dict"),
|
||||||
"status": fields.String,
|
"status": fields.String,
|
||||||
"error": fields.String,
|
"error": fields.String,
|
||||||
@@ -299,6 +300,7 @@ class ChatMessageListApi(Resource):
|
|||||||
has_more = False
|
has_more = False
|
||||||
|
|
||||||
history_messages = list(reversed(history_messages))
|
history_messages = list(reversed(history_messages))
|
||||||
|
attach_message_extra_contents(history_messages)
|
||||||
|
|
||||||
return InfiniteScrollPagination(data=history_messages, limit=args.limit, has_more=has_more)
|
return InfiniteScrollPagination(data=history_messages, limit=args.limit, has_more=has_more)
|
||||||
|
|
||||||
@@ -481,4 +483,5 @@ class MessageApi(Resource):
|
|||||||
if not message:
|
if not message:
|
||||||
raise NotFound("Message Not Exists.")
|
raise NotFound("Message Not Exists.")
|
||||||
|
|
||||||
|
attach_message_extra_contents([message])
|
||||||
return message
|
return message
|
||||||
|
|||||||
@@ -507,6 +507,179 @@ class WorkflowDraftRunLoopNodeApi(Resource):
|
|||||||
raise InternalServerError()
|
raise InternalServerError()
|
||||||
|
|
||||||
|
|
||||||
|
class HumanInputFormPreviewPayload(BaseModel):
|
||||||
|
inputs: dict[str, Any] = Field(
|
||||||
|
default_factory=dict,
|
||||||
|
description="Values used to fill missing upstream variables referenced in form_content",
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class HumanInputFormSubmitPayload(BaseModel):
|
||||||
|
form_inputs: dict[str, Any] = Field(..., description="Values the user provides for the form's own fields")
|
||||||
|
inputs: dict[str, Any] = Field(
|
||||||
|
...,
|
||||||
|
description="Values used to fill missing upstream variables referenced in form_content",
|
||||||
|
)
|
||||||
|
action: str = Field(..., description="Selected action ID")
|
||||||
|
|
||||||
|
|
||||||
|
class HumanInputDeliveryTestPayload(BaseModel):
|
||||||
|
delivery_method_id: str = Field(..., description="Delivery method ID")
|
||||||
|
inputs: dict[str, Any] = Field(
|
||||||
|
default_factory=dict,
|
||||||
|
description="Values used to fill missing upstream variables referenced in form_content",
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
reg(HumanInputFormPreviewPayload)
|
||||||
|
reg(HumanInputFormSubmitPayload)
|
||||||
|
reg(HumanInputDeliveryTestPayload)
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/apps/<uuid:app_id>/advanced-chat/workflows/draft/human-input/nodes/<string:node_id>/form/preview")
|
||||||
|
class AdvancedChatDraftHumanInputFormPreviewApi(Resource):
|
||||||
|
@console_ns.doc("get_advanced_chat_draft_human_input_form")
|
||||||
|
@console_ns.doc(description="Get human input form preview for advanced chat workflow")
|
||||||
|
@console_ns.doc(params={"app_id": "Application ID", "node_id": "Node ID"})
|
||||||
|
@console_ns.expect(console_ns.models[HumanInputFormPreviewPayload.__name__])
|
||||||
|
@setup_required
|
||||||
|
@login_required
|
||||||
|
@account_initialization_required
|
||||||
|
@get_app_model(mode=[AppMode.ADVANCED_CHAT])
|
||||||
|
@edit_permission_required
|
||||||
|
def post(self, app_model: App, node_id: str):
|
||||||
|
"""
|
||||||
|
Preview human input form content and placeholders
|
||||||
|
"""
|
||||||
|
current_user, _ = current_account_with_tenant()
|
||||||
|
args = HumanInputFormPreviewPayload.model_validate(console_ns.payload or {})
|
||||||
|
inputs = args.inputs
|
||||||
|
|
||||||
|
workflow_service = WorkflowService()
|
||||||
|
preview = workflow_service.get_human_input_form_preview(
|
||||||
|
app_model=app_model,
|
||||||
|
account=current_user,
|
||||||
|
node_id=node_id,
|
||||||
|
inputs=inputs,
|
||||||
|
)
|
||||||
|
return jsonable_encoder(preview)
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/apps/<uuid:app_id>/advanced-chat/workflows/draft/human-input/nodes/<string:node_id>/form/run")
|
||||||
|
class AdvancedChatDraftHumanInputFormRunApi(Resource):
|
||||||
|
@console_ns.doc("submit_advanced_chat_draft_human_input_form")
|
||||||
|
@console_ns.doc(description="Submit human input form preview for advanced chat workflow")
|
||||||
|
@console_ns.doc(params={"app_id": "Application ID", "node_id": "Node ID"})
|
||||||
|
@console_ns.expect(console_ns.models[HumanInputFormSubmitPayload.__name__])
|
||||||
|
@setup_required
|
||||||
|
@login_required
|
||||||
|
@account_initialization_required
|
||||||
|
@get_app_model(mode=[AppMode.ADVANCED_CHAT])
|
||||||
|
@edit_permission_required
|
||||||
|
def post(self, app_model: App, node_id: str):
|
||||||
|
"""
|
||||||
|
Submit human input form preview
|
||||||
|
"""
|
||||||
|
current_user, _ = current_account_with_tenant()
|
||||||
|
args = HumanInputFormSubmitPayload.model_validate(console_ns.payload or {})
|
||||||
|
workflow_service = WorkflowService()
|
||||||
|
result = workflow_service.submit_human_input_form_preview(
|
||||||
|
app_model=app_model,
|
||||||
|
account=current_user,
|
||||||
|
node_id=node_id,
|
||||||
|
form_inputs=args.form_inputs,
|
||||||
|
inputs=args.inputs,
|
||||||
|
action=args.action,
|
||||||
|
)
|
||||||
|
return jsonable_encoder(result)
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/apps/<uuid:app_id>/workflows/draft/human-input/nodes/<string:node_id>/form/preview")
|
||||||
|
class WorkflowDraftHumanInputFormPreviewApi(Resource):
|
||||||
|
@console_ns.doc("get_workflow_draft_human_input_form")
|
||||||
|
@console_ns.doc(description="Get human input form preview for workflow")
|
||||||
|
@console_ns.doc(params={"app_id": "Application ID", "node_id": "Node ID"})
|
||||||
|
@console_ns.expect(console_ns.models[HumanInputFormPreviewPayload.__name__])
|
||||||
|
@setup_required
|
||||||
|
@login_required
|
||||||
|
@account_initialization_required
|
||||||
|
@get_app_model(mode=[AppMode.WORKFLOW])
|
||||||
|
@edit_permission_required
|
||||||
|
def post(self, app_model: App, node_id: str):
|
||||||
|
"""
|
||||||
|
Preview human input form content and placeholders
|
||||||
|
"""
|
||||||
|
current_user, _ = current_account_with_tenant()
|
||||||
|
args = HumanInputFormPreviewPayload.model_validate(console_ns.payload or {})
|
||||||
|
inputs = args.inputs
|
||||||
|
|
||||||
|
workflow_service = WorkflowService()
|
||||||
|
preview = workflow_service.get_human_input_form_preview(
|
||||||
|
app_model=app_model,
|
||||||
|
account=current_user,
|
||||||
|
node_id=node_id,
|
||||||
|
inputs=inputs,
|
||||||
|
)
|
||||||
|
return jsonable_encoder(preview)
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/apps/<uuid:app_id>/workflows/draft/human-input/nodes/<string:node_id>/form/run")
|
||||||
|
class WorkflowDraftHumanInputFormRunApi(Resource):
|
||||||
|
@console_ns.doc("submit_workflow_draft_human_input_form")
|
||||||
|
@console_ns.doc(description="Submit human input form preview for workflow")
|
||||||
|
@console_ns.doc(params={"app_id": "Application ID", "node_id": "Node ID"})
|
||||||
|
@console_ns.expect(console_ns.models[HumanInputFormSubmitPayload.__name__])
|
||||||
|
@setup_required
|
||||||
|
@login_required
|
||||||
|
@account_initialization_required
|
||||||
|
@get_app_model(mode=[AppMode.WORKFLOW])
|
||||||
|
@edit_permission_required
|
||||||
|
def post(self, app_model: App, node_id: str):
|
||||||
|
"""
|
||||||
|
Submit human input form preview
|
||||||
|
"""
|
||||||
|
current_user, _ = current_account_with_tenant()
|
||||||
|
workflow_service = WorkflowService()
|
||||||
|
args = HumanInputFormSubmitPayload.model_validate(console_ns.payload or {})
|
||||||
|
result = workflow_service.submit_human_input_form_preview(
|
||||||
|
app_model=app_model,
|
||||||
|
account=current_user,
|
||||||
|
node_id=node_id,
|
||||||
|
form_inputs=args.form_inputs,
|
||||||
|
inputs=args.inputs,
|
||||||
|
action=args.action,
|
||||||
|
)
|
||||||
|
return jsonable_encoder(result)
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/apps/<uuid:app_id>/workflows/draft/human-input/nodes/<string:node_id>/delivery-test")
|
||||||
|
class WorkflowDraftHumanInputDeliveryTestApi(Resource):
|
||||||
|
@console_ns.doc("test_workflow_draft_human_input_delivery")
|
||||||
|
@console_ns.doc(description="Test human input delivery for workflow")
|
||||||
|
@console_ns.doc(params={"app_id": "Application ID", "node_id": "Node ID"})
|
||||||
|
@console_ns.expect(console_ns.models[HumanInputDeliveryTestPayload.__name__])
|
||||||
|
@setup_required
|
||||||
|
@login_required
|
||||||
|
@account_initialization_required
|
||||||
|
@get_app_model(mode=[AppMode.WORKFLOW, AppMode.ADVANCED_CHAT])
|
||||||
|
@edit_permission_required
|
||||||
|
def post(self, app_model: App, node_id: str):
|
||||||
|
"""
|
||||||
|
Test human input delivery
|
||||||
|
"""
|
||||||
|
current_user, _ = current_account_with_tenant()
|
||||||
|
workflow_service = WorkflowService()
|
||||||
|
args = HumanInputDeliveryTestPayload.model_validate(console_ns.payload or {})
|
||||||
|
workflow_service.test_human_input_delivery(
|
||||||
|
app_model=app_model,
|
||||||
|
account=current_user,
|
||||||
|
node_id=node_id,
|
||||||
|
delivery_method_id=args.delivery_method_id,
|
||||||
|
inputs=args.inputs,
|
||||||
|
)
|
||||||
|
return jsonable_encoder({})
|
||||||
|
|
||||||
|
|
||||||
@console_ns.route("/apps/<uuid:app_id>/workflows/draft/run")
|
@console_ns.route("/apps/<uuid:app_id>/workflows/draft/run")
|
||||||
class DraftWorkflowRunApi(Resource):
|
class DraftWorkflowRunApi(Resource):
|
||||||
@console_ns.doc("run_draft_workflow")
|
@console_ns.doc("run_draft_workflow")
|
||||||
|
|||||||
@@ -5,10 +5,15 @@ from flask import request
|
|||||||
from flask_restx import Resource, fields, marshal_with
|
from flask_restx import Resource, fields, marshal_with
|
||||||
from pydantic import BaseModel, Field, field_validator
|
from pydantic import BaseModel, Field, field_validator
|
||||||
from sqlalchemy import select
|
from sqlalchemy import select
|
||||||
|
from sqlalchemy.orm import sessionmaker
|
||||||
|
|
||||||
|
from configs import dify_config
|
||||||
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 controllers.web.error import NotFoundError
|
||||||
|
from core.workflow.entities.pause_reason import HumanInputRequired
|
||||||
|
from core.workflow.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,9 +32,21 @@ from libs.custom_inputs import time_duration
|
|||||||
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
|
||||||
from models import Account, App, AppMode, EndUser, WorkflowArchiveLog, WorkflowRunTriggeredFrom
|
from models import Account, App, AppMode, EndUser, WorkflowArchiveLog, WorkflowRunTriggeredFrom
|
||||||
|
from models.workflow import WorkflowRun
|
||||||
|
from repositories.factory import DifyAPIRepositoryFactory
|
||||||
from services.retention.workflow_run.constants import ARCHIVE_BUNDLE_NAME
|
from services.retention.workflow_run.constants import ARCHIVE_BUNDLE_NAME
|
||||||
from services.workflow_run_service import WorkflowRunService
|
from services.workflow_run_service import WorkflowRunService
|
||||||
|
|
||||||
|
|
||||||
|
def _build_backstage_input_url(form_token: str | None) -> str | None:
|
||||||
|
if not form_token:
|
||||||
|
return None
|
||||||
|
base_url = dify_config.APP_WEB_URL
|
||||||
|
if not base_url:
|
||||||
|
return None
|
||||||
|
return f"{base_url.rstrip('/')}/form/{form_token}"
|
||||||
|
|
||||||
|
|
||||||
# Workflow run status choices for filtering
|
# Workflow run status choices for filtering
|
||||||
WORKFLOW_RUN_STATUS_CHOICES = ["running", "succeeded", "failed", "stopped", "partial-succeeded"]
|
WORKFLOW_RUN_STATUS_CHOICES = ["running", "succeeded", "failed", "stopped", "partial-succeeded"]
|
||||||
EXPORT_SIGNED_URL_EXPIRE_SECONDS = 3600
|
EXPORT_SIGNED_URL_EXPIRE_SECONDS = 3600
|
||||||
@@ -440,3 +457,68 @@ class WorkflowRunNodeExecutionListApi(Resource):
|
|||||||
)
|
)
|
||||||
|
|
||||||
return {"data": node_executions}
|
return {"data": node_executions}
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workflow/<string:workflow_run_id>/pause-details")
|
||||||
|
class ConsoleWorkflowPauseDetailsApi(Resource):
|
||||||
|
"""Console API for getting workflow pause details."""
|
||||||
|
|
||||||
|
@setup_required
|
||||||
|
@login_required
|
||||||
|
@account_initialization_required
|
||||||
|
def get(self, workflow_run_id: str):
|
||||||
|
"""
|
||||||
|
Get workflow pause details.
|
||||||
|
|
||||||
|
GET /console/api/workflow/<workflow_run_id>/pause-details
|
||||||
|
|
||||||
|
Returns information about why and where the workflow is paused.
|
||||||
|
"""
|
||||||
|
|
||||||
|
# Query WorkflowRun to determine if workflow is suspended
|
||||||
|
session_maker = sessionmaker(bind=db.engine)
|
||||||
|
workflow_run_repo = DifyAPIRepositoryFactory.create_api_workflow_run_repository(session_maker=session_maker)
|
||||||
|
|
||||||
|
workflow_run = db.session.get(WorkflowRun, workflow_run_id)
|
||||||
|
if not workflow_run:
|
||||||
|
raise NotFoundError("Workflow run not found")
|
||||||
|
|
||||||
|
if workflow_run.tenant_id != current_user.current_tenant_id:
|
||||||
|
raise NotFoundError("Workflow run not found")
|
||||||
|
|
||||||
|
# Check if workflow is suspended
|
||||||
|
is_paused = workflow_run.status == WorkflowExecutionStatus.PAUSED
|
||||||
|
if not is_paused:
|
||||||
|
return {
|
||||||
|
"paused_at": None,
|
||||||
|
"paused_nodes": [],
|
||||||
|
}, 200
|
||||||
|
|
||||||
|
pause_entity = workflow_run_repo.get_workflow_pause(workflow_run_id)
|
||||||
|
pause_reasons = pause_entity.get_pause_reasons() if pause_entity else []
|
||||||
|
|
||||||
|
# Build response
|
||||||
|
paused_at = pause_entity.paused_at if pause_entity else None
|
||||||
|
paused_nodes = []
|
||||||
|
response = {
|
||||||
|
"paused_at": paused_at.isoformat() + "Z" if paused_at else None,
|
||||||
|
"paused_nodes": paused_nodes,
|
||||||
|
}
|
||||||
|
|
||||||
|
for reason in pause_reasons:
|
||||||
|
if isinstance(reason, HumanInputRequired):
|
||||||
|
paused_nodes.append(
|
||||||
|
{
|
||||||
|
"node_id": reason.node_id,
|
||||||
|
"node_title": reason.node_title,
|
||||||
|
"pause_type": {
|
||||||
|
"type": "human_input",
|
||||||
|
"form_id": reason.form_id,
|
||||||
|
"backstage_input_url": _build_backstage_input_url(reason.form_token),
|
||||||
|
},
|
||||||
|
}
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
raise AssertionError("unimplemented.")
|
||||||
|
|
||||||
|
return response, 200
|
||||||
|
|||||||
@@ -55,6 +55,7 @@ from libs.login import current_account_with_tenant, login_required
|
|||||||
from models import ApiToken, Dataset, Document, DocumentSegment, UploadFile
|
from models import ApiToken, Dataset, Document, DocumentSegment, UploadFile
|
||||||
from models.dataset import DatasetPermissionEnum
|
from models.dataset import DatasetPermissionEnum
|
||||||
from models.provider_ids import ModelProviderID
|
from models.provider_ids import ModelProviderID
|
||||||
|
from services.api_token_service import ApiTokenCache
|
||||||
from services.dataset_service import DatasetPermissionService, DatasetService, DocumentService
|
from services.dataset_service import DatasetPermissionService, DatasetService, DocumentService
|
||||||
|
|
||||||
# Register models for flask_restx to avoid dict type issues in Swagger
|
# Register models for flask_restx to avoid dict type issues in Swagger
|
||||||
@@ -820,6 +821,11 @@ class DatasetApiDeleteApi(Resource):
|
|||||||
if key is None:
|
if key is None:
|
||||||
console_ns.abort(404, message="API key not found")
|
console_ns.abort(404, message="API key not found")
|
||||||
|
|
||||||
|
# Invalidate cache before deleting from database
|
||||||
|
# Type assertion: key is guaranteed to be non-None here because abort() raises
|
||||||
|
assert key is not None # nosec - for type checker only
|
||||||
|
ApiTokenCache.delete(key.token, key.type)
|
||||||
|
|
||||||
db.session.query(ApiToken).where(ApiToken.id == api_key_id).delete()
|
db.session.query(ApiToken).where(ApiToken.id == api_key_id).delete()
|
||||||
db.session.commit()
|
db.session.commit()
|
||||||
|
|
||||||
|
|||||||
@@ -1,15 +1,16 @@
|
|||||||
import logging
|
import logging
|
||||||
from typing import Any, cast
|
from typing import Any, Literal, cast
|
||||||
|
|
||||||
from flask import request
|
from flask import request
|
||||||
from flask_restx import Resource, fields, marshal, marshal_with, reqparse
|
from flask_restx import Resource, fields, marshal, marshal_with
|
||||||
|
from pydantic import BaseModel
|
||||||
from werkzeug.exceptions import Forbidden, InternalServerError, NotFound
|
from werkzeug.exceptions import Forbidden, InternalServerError, NotFound
|
||||||
|
|
||||||
import services
|
import services
|
||||||
from controllers.common.fields import Parameters as ParametersResponse
|
from controllers.common.fields import Parameters as ParametersResponse
|
||||||
from controllers.common.fields import Site as SiteResponse
|
from controllers.common.fields import Site as SiteResponse
|
||||||
from controllers.common.schema import get_or_create_model
|
from controllers.common.schema import get_or_create_model
|
||||||
from controllers.console import api
|
from controllers.console import api, console_ns
|
||||||
from controllers.console.app.error import (
|
from controllers.console.app.error import (
|
||||||
AppUnavailableError,
|
AppUnavailableError,
|
||||||
AudioTooLargeError,
|
AudioTooLargeError,
|
||||||
@@ -117,7 +118,56 @@ workflow_fields_copy["rag_pipeline_variables"] = fields.List(fields.Nested(pipel
|
|||||||
workflow_model = get_or_create_model("TrialWorkflow", workflow_fields_copy)
|
workflow_model = get_or_create_model("TrialWorkflow", workflow_fields_copy)
|
||||||
|
|
||||||
|
|
||||||
|
# Pydantic models for request validation
|
||||||
|
DEFAULT_REF_TEMPLATE_SWAGGER_2_0 = "#/definitions/{model}"
|
||||||
|
|
||||||
|
|
||||||
|
class WorkflowRunRequest(BaseModel):
|
||||||
|
inputs: dict
|
||||||
|
files: list | None = None
|
||||||
|
|
||||||
|
|
||||||
|
class ChatRequest(BaseModel):
|
||||||
|
inputs: dict
|
||||||
|
query: str
|
||||||
|
files: list | None = None
|
||||||
|
conversation_id: str | None = None
|
||||||
|
parent_message_id: str | None = None
|
||||||
|
retriever_from: str = "explore_app"
|
||||||
|
|
||||||
|
|
||||||
|
class TextToSpeechRequest(BaseModel):
|
||||||
|
message_id: str | None = None
|
||||||
|
voice: str | None = None
|
||||||
|
text: str | None = None
|
||||||
|
streaming: bool | None = None
|
||||||
|
|
||||||
|
|
||||||
|
class CompletionRequest(BaseModel):
|
||||||
|
inputs: dict
|
||||||
|
query: str = ""
|
||||||
|
files: list | None = None
|
||||||
|
response_mode: Literal["blocking", "streaming"] | None = None
|
||||||
|
retriever_from: str = "explore_app"
|
||||||
|
|
||||||
|
|
||||||
|
# Register schemas for Swagger documentation
|
||||||
|
console_ns.schema_model(
|
||||||
|
WorkflowRunRequest.__name__, WorkflowRunRequest.model_json_schema(ref_template=DEFAULT_REF_TEMPLATE_SWAGGER_2_0)
|
||||||
|
)
|
||||||
|
console_ns.schema_model(
|
||||||
|
ChatRequest.__name__, ChatRequest.model_json_schema(ref_template=DEFAULT_REF_TEMPLATE_SWAGGER_2_0)
|
||||||
|
)
|
||||||
|
console_ns.schema_model(
|
||||||
|
TextToSpeechRequest.__name__, TextToSpeechRequest.model_json_schema(ref_template=DEFAULT_REF_TEMPLATE_SWAGGER_2_0)
|
||||||
|
)
|
||||||
|
console_ns.schema_model(
|
||||||
|
CompletionRequest.__name__, CompletionRequest.model_json_schema(ref_template=DEFAULT_REF_TEMPLATE_SWAGGER_2_0)
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class TrialAppWorkflowRunApi(TrialAppResource):
|
class TrialAppWorkflowRunApi(TrialAppResource):
|
||||||
|
@console_ns.expect(console_ns.models[WorkflowRunRequest.__name__])
|
||||||
def post(self, trial_app):
|
def post(self, trial_app):
|
||||||
"""
|
"""
|
||||||
Run workflow
|
Run workflow
|
||||||
@@ -129,10 +179,8 @@ class TrialAppWorkflowRunApi(TrialAppResource):
|
|||||||
if app_mode != AppMode.WORKFLOW:
|
if app_mode != AppMode.WORKFLOW:
|
||||||
raise NotWorkflowAppError()
|
raise NotWorkflowAppError()
|
||||||
|
|
||||||
parser = reqparse.RequestParser()
|
request_data = WorkflowRunRequest.model_validate(console_ns.payload)
|
||||||
parser.add_argument("inputs", type=dict, required=True, nullable=False, location="json")
|
args = request_data.model_dump()
|
||||||
parser.add_argument("files", type=list, required=False, location="json")
|
|
||||||
args = parser.parse_args()
|
|
||||||
assert current_user is not None
|
assert current_user is not None
|
||||||
try:
|
try:
|
||||||
app_id = app_model.id
|
app_id = app_model.id
|
||||||
@@ -183,6 +231,7 @@ class TrialAppWorkflowTaskStopApi(TrialAppResource):
|
|||||||
|
|
||||||
|
|
||||||
class TrialChatApi(TrialAppResource):
|
class TrialChatApi(TrialAppResource):
|
||||||
|
@console_ns.expect(console_ns.models[ChatRequest.__name__])
|
||||||
@trial_feature_enable
|
@trial_feature_enable
|
||||||
def post(self, trial_app):
|
def post(self, trial_app):
|
||||||
app_model = trial_app
|
app_model = trial_app
|
||||||
@@ -190,14 +239,14 @@ class TrialChatApi(TrialAppResource):
|
|||||||
if app_mode not in {AppMode.CHAT, AppMode.AGENT_CHAT, AppMode.ADVANCED_CHAT}:
|
if app_mode not in {AppMode.CHAT, AppMode.AGENT_CHAT, AppMode.ADVANCED_CHAT}:
|
||||||
raise NotChatAppError()
|
raise NotChatAppError()
|
||||||
|
|
||||||
parser = reqparse.RequestParser()
|
request_data = ChatRequest.model_validate(console_ns.payload)
|
||||||
parser.add_argument("inputs", type=dict, required=True, location="json")
|
args = request_data.model_dump()
|
||||||
parser.add_argument("query", type=str, required=True, location="json")
|
|
||||||
parser.add_argument("files", type=list, required=False, location="json")
|
# Validate UUID values if provided
|
||||||
parser.add_argument("conversation_id", type=uuid_value, location="json")
|
if args.get("conversation_id"):
|
||||||
parser.add_argument("parent_message_id", type=uuid_value, required=False, location="json")
|
args["conversation_id"] = uuid_value(args["conversation_id"])
|
||||||
parser.add_argument("retriever_from", type=str, required=False, default="explore_app", location="json")
|
if args.get("parent_message_id"):
|
||||||
args = parser.parse_args()
|
args["parent_message_id"] = uuid_value(args["parent_message_id"])
|
||||||
|
|
||||||
args["auto_generate_name"] = False
|
args["auto_generate_name"] = False
|
||||||
|
|
||||||
@@ -320,20 +369,16 @@ class TrialChatAudioApi(TrialAppResource):
|
|||||||
|
|
||||||
|
|
||||||
class TrialChatTextApi(TrialAppResource):
|
class TrialChatTextApi(TrialAppResource):
|
||||||
|
@console_ns.expect(console_ns.models[TextToSpeechRequest.__name__])
|
||||||
@trial_feature_enable
|
@trial_feature_enable
|
||||||
def post(self, trial_app):
|
def post(self, trial_app):
|
||||||
app_model = trial_app
|
app_model = trial_app
|
||||||
try:
|
try:
|
||||||
parser = reqparse.RequestParser()
|
request_data = TextToSpeechRequest.model_validate(console_ns.payload)
|
||||||
parser.add_argument("message_id", type=str, required=False, location="json")
|
|
||||||
parser.add_argument("voice", type=str, location="json")
|
|
||||||
parser.add_argument("text", type=str, location="json")
|
|
||||||
parser.add_argument("streaming", type=bool, location="json")
|
|
||||||
args = parser.parse_args()
|
|
||||||
|
|
||||||
message_id = args.get("message_id", None)
|
message_id = request_data.message_id
|
||||||
text = args.get("text", None)
|
text = request_data.text
|
||||||
voice = args.get("voice", None)
|
voice = request_data.voice
|
||||||
if not isinstance(current_user, Account):
|
if not isinstance(current_user, Account):
|
||||||
raise ValueError("current_user must be an Account instance")
|
raise ValueError("current_user must be an Account instance")
|
||||||
|
|
||||||
@@ -371,19 +416,15 @@ class TrialChatTextApi(TrialAppResource):
|
|||||||
|
|
||||||
|
|
||||||
class TrialCompletionApi(TrialAppResource):
|
class TrialCompletionApi(TrialAppResource):
|
||||||
|
@console_ns.expect(console_ns.models[CompletionRequest.__name__])
|
||||||
@trial_feature_enable
|
@trial_feature_enable
|
||||||
def post(self, trial_app):
|
def post(self, trial_app):
|
||||||
app_model = trial_app
|
app_model = trial_app
|
||||||
if app_model.mode != "completion":
|
if app_model.mode != "completion":
|
||||||
raise NotCompletionAppError()
|
raise NotCompletionAppError()
|
||||||
|
|
||||||
parser = reqparse.RequestParser()
|
request_data = CompletionRequest.model_validate(console_ns.payload)
|
||||||
parser.add_argument("inputs", type=dict, required=True, location="json")
|
args = request_data.model_dump()
|
||||||
parser.add_argument("query", type=str, location="json", default="")
|
|
||||||
parser.add_argument("files", type=list, required=False, location="json")
|
|
||||||
parser.add_argument("response_mode", type=str, choices=["blocking", "streaming"], location="json")
|
|
||||||
parser.add_argument("retriever_from", type=str, required=False, default="explore_app", location="json")
|
|
||||||
args = parser.parse_args()
|
|
||||||
|
|
||||||
streaming = args["response_mode"] == "streaming"
|
streaming = args["response_mode"] == "streaming"
|
||||||
args["auto_generate_name"] = False
|
args["auto_generate_name"] = False
|
||||||
|
|||||||
@@ -0,0 +1,217 @@
|
|||||||
|
"""
|
||||||
|
Console/Studio Human Input Form APIs.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import json
|
||||||
|
import logging
|
||||||
|
from collections.abc import Generator
|
||||||
|
|
||||||
|
from flask import Response, jsonify, request
|
||||||
|
from flask_restx import Resource, reqparse
|
||||||
|
from sqlalchemy import select
|
||||||
|
from sqlalchemy.orm import Session, sessionmaker
|
||||||
|
|
||||||
|
from controllers.console import console_ns
|
||||||
|
from controllers.console.wraps import account_initialization_required, setup_required
|
||||||
|
from controllers.web.error import InvalidArgumentError, NotFoundError
|
||||||
|
from core.app.apps.advanced_chat.app_generator import AdvancedChatAppGenerator
|
||||||
|
from core.app.apps.common.workflow_response_converter import WorkflowResponseConverter
|
||||||
|
from core.app.apps.message_generator import MessageGenerator
|
||||||
|
from core.app.apps.workflow.app_generator import WorkflowAppGenerator
|
||||||
|
from extensions.ext_database import db
|
||||||
|
from libs.login import current_account_with_tenant, login_required
|
||||||
|
from models import App
|
||||||
|
from models.enums import CreatorUserRole
|
||||||
|
from models.human_input import RecipientType
|
||||||
|
from models.model import AppMode
|
||||||
|
from models.workflow import WorkflowRun
|
||||||
|
from repositories.factory import DifyAPIRepositoryFactory
|
||||||
|
from services.human_input_service import Form, HumanInputService
|
||||||
|
from services.workflow_event_snapshot_service import build_workflow_event_stream
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
def _jsonify_form_definition(form: Form) -> Response:
|
||||||
|
payload = form.get_definition().model_dump()
|
||||||
|
payload["expiration_time"] = int(form.expiration_time.timestamp())
|
||||||
|
return Response(json.dumps(payload, ensure_ascii=False), mimetype="application/json")
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/form/human_input/<string:form_token>")
|
||||||
|
class ConsoleHumanInputFormApi(Resource):
|
||||||
|
"""Console API for getting human input form definition."""
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _ensure_console_access(form: Form):
|
||||||
|
_, current_tenant_id = current_account_with_tenant()
|
||||||
|
|
||||||
|
if form.tenant_id != current_tenant_id:
|
||||||
|
raise NotFoundError("App not found")
|
||||||
|
|
||||||
|
@setup_required
|
||||||
|
@login_required
|
||||||
|
@account_initialization_required
|
||||||
|
def get(self, form_token: str):
|
||||||
|
"""
|
||||||
|
Get human input form definition by form token.
|
||||||
|
|
||||||
|
GET /console/api/form/human_input/<form_token>
|
||||||
|
"""
|
||||||
|
service = HumanInputService(db.engine)
|
||||||
|
form = service.get_form_definition_by_token_for_console(form_token)
|
||||||
|
if form is None:
|
||||||
|
raise NotFoundError(f"form not found, token={form_token}")
|
||||||
|
|
||||||
|
self._ensure_console_access(form)
|
||||||
|
|
||||||
|
return _jsonify_form_definition(form)
|
||||||
|
|
||||||
|
@account_initialization_required
|
||||||
|
@login_required
|
||||||
|
def post(self, form_token: str):
|
||||||
|
"""
|
||||||
|
Submit human input form by form token.
|
||||||
|
|
||||||
|
POST /console/api/form/human_input/<form_token>
|
||||||
|
|
||||||
|
Request body:
|
||||||
|
{
|
||||||
|
"inputs": {
|
||||||
|
"content": "User input content"
|
||||||
|
},
|
||||||
|
"action": "Approve"
|
||||||
|
}
|
||||||
|
"""
|
||||||
|
parser = reqparse.RequestParser()
|
||||||
|
parser.add_argument("inputs", type=dict, required=True, location="json")
|
||||||
|
parser.add_argument("action", type=str, required=True, location="json")
|
||||||
|
args = parser.parse_args()
|
||||||
|
current_user, _ = current_account_with_tenant()
|
||||||
|
|
||||||
|
service = HumanInputService(db.engine)
|
||||||
|
form = service.get_form_by_token(form_token)
|
||||||
|
if form is None:
|
||||||
|
raise NotFoundError(f"form not found, token={form_token}")
|
||||||
|
|
||||||
|
self._ensure_console_access(form)
|
||||||
|
|
||||||
|
recipient_type = form.recipient_type
|
||||||
|
if recipient_type not in {RecipientType.CONSOLE, RecipientType.BACKSTAGE}:
|
||||||
|
raise NotFoundError(f"form not found, token={form_token}")
|
||||||
|
# The type checker is not smart enought to validate the following invariant.
|
||||||
|
# So we need to assert it manually.
|
||||||
|
assert recipient_type is not None, "recipient_type cannot be None here."
|
||||||
|
|
||||||
|
service.submit_form_by_token(
|
||||||
|
recipient_type=recipient_type,
|
||||||
|
form_token=form_token,
|
||||||
|
selected_action_id=args["action"],
|
||||||
|
form_data=args["inputs"],
|
||||||
|
submission_user_id=current_user.id,
|
||||||
|
)
|
||||||
|
|
||||||
|
return jsonify({})
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workflow/<string:workflow_run_id>/events")
|
||||||
|
class ConsoleWorkflowEventsApi(Resource):
|
||||||
|
"""Console API for getting workflow execution events after resume."""
|
||||||
|
|
||||||
|
@account_initialization_required
|
||||||
|
@login_required
|
||||||
|
def get(self, workflow_run_id: str):
|
||||||
|
"""
|
||||||
|
Get workflow execution events stream after resume.
|
||||||
|
|
||||||
|
GET /console/api/workflow/<workflow_run_id>/events
|
||||||
|
|
||||||
|
Returns Server-Sent Events stream.
|
||||||
|
"""
|
||||||
|
|
||||||
|
user, tenant_id = current_account_with_tenant()
|
||||||
|
session_maker = sessionmaker(db.engine)
|
||||||
|
repo = DifyAPIRepositoryFactory.create_api_workflow_run_repository(session_maker)
|
||||||
|
workflow_run = repo.get_workflow_run_by_id_and_tenant_id(
|
||||||
|
tenant_id=tenant_id,
|
||||||
|
run_id=workflow_run_id,
|
||||||
|
)
|
||||||
|
if workflow_run is None:
|
||||||
|
raise NotFoundError(f"WorkflowRun not found, id={workflow_run_id}")
|
||||||
|
|
||||||
|
if workflow_run.created_by_role != CreatorUserRole.ACCOUNT:
|
||||||
|
raise NotFoundError(f"WorkflowRun not created by account, id={workflow_run_id}")
|
||||||
|
|
||||||
|
if workflow_run.created_by != user.id:
|
||||||
|
raise NotFoundError(f"WorkflowRun not created by the current account, id={workflow_run_id}")
|
||||||
|
|
||||||
|
with Session(expire_on_commit=False, bind=db.engine) as session:
|
||||||
|
app = _retrieve_app_for_workflow_run(session, workflow_run)
|
||||||
|
|
||||||
|
if workflow_run.finished_at is not None:
|
||||||
|
# TODO(QuantumGhost): should we modify the handling for finished workflow run here?
|
||||||
|
response = WorkflowResponseConverter.workflow_run_result_to_finish_response(
|
||||||
|
task_id=workflow_run.id,
|
||||||
|
workflow_run=workflow_run,
|
||||||
|
creator_user=user,
|
||||||
|
)
|
||||||
|
|
||||||
|
payload = response.model_dump(mode="json")
|
||||||
|
payload["event"] = response.event.value
|
||||||
|
|
||||||
|
def _generate_finished_events() -> Generator[str, None, None]:
|
||||||
|
yield f"data: {json.dumps(payload)}\n\n"
|
||||||
|
|
||||||
|
event_generator = _generate_finished_events
|
||||||
|
|
||||||
|
else:
|
||||||
|
msg_generator = MessageGenerator()
|
||||||
|
if app.mode == AppMode.ADVANCED_CHAT:
|
||||||
|
generator = AdvancedChatAppGenerator()
|
||||||
|
elif app.mode == AppMode.WORKFLOW:
|
||||||
|
generator = WorkflowAppGenerator()
|
||||||
|
else:
|
||||||
|
raise InvalidArgumentError(f"cannot subscribe to workflow run, workflow_run_id={workflow_run.id}")
|
||||||
|
|
||||||
|
include_state_snapshot = request.args.get("include_state_snapshot", "false").lower() == "true"
|
||||||
|
|
||||||
|
def _generate_stream_events():
|
||||||
|
if include_state_snapshot:
|
||||||
|
return generator.convert_to_event_stream(
|
||||||
|
build_workflow_event_stream(
|
||||||
|
app_mode=AppMode(app.mode),
|
||||||
|
workflow_run=workflow_run,
|
||||||
|
tenant_id=workflow_run.tenant_id,
|
||||||
|
app_id=workflow_run.app_id,
|
||||||
|
session_maker=session_maker,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
return generator.convert_to_event_stream(
|
||||||
|
msg_generator.retrieve_events(AppMode(app.mode), workflow_run.id),
|
||||||
|
)
|
||||||
|
|
||||||
|
event_generator = _generate_stream_events
|
||||||
|
|
||||||
|
return Response(
|
||||||
|
event_generator(),
|
||||||
|
mimetype="text/event-stream",
|
||||||
|
headers={
|
||||||
|
"Cache-Control": "no-cache",
|
||||||
|
"Connection": "keep-alive",
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _retrieve_app_for_workflow_run(session: Session, workflow_run: WorkflowRun):
|
||||||
|
query = select(App).where(
|
||||||
|
App.id == workflow_run.app_id,
|
||||||
|
App.tenant_id == workflow_run.tenant_id,
|
||||||
|
)
|
||||||
|
app = session.scalars(query).first()
|
||||||
|
if app is None:
|
||||||
|
raise AssertionError(
|
||||||
|
f"App not found for WorkflowRun, workflow_run_id={workflow_run.id}, "
|
||||||
|
f"app_id={workflow_run.app_id}, tenant_id={workflow_run.tenant_id}"
|
||||||
|
)
|
||||||
|
|
||||||
|
return app
|
||||||
@@ -1,6 +1,7 @@
|
|||||||
import urllib.parse
|
import urllib.parse
|
||||||
|
|
||||||
import httpx
|
import httpx
|
||||||
|
from flask_restx import Resource
|
||||||
from pydantic import BaseModel, Field
|
from pydantic import BaseModel, Field
|
||||||
|
|
||||||
import services
|
import services
|
||||||
@@ -10,12 +11,12 @@ from controllers.common.errors import (
|
|||||||
RemoteFileUploadError,
|
RemoteFileUploadError,
|
||||||
UnsupportedFileTypeError,
|
UnsupportedFileTypeError,
|
||||||
)
|
)
|
||||||
from controllers.fastopenapi import console_router
|
from controllers.console import console_ns
|
||||||
from core.file import helpers as file_helpers
|
from core.file import helpers as file_helpers
|
||||||
from core.helper import ssrf_proxy
|
from core.helper import ssrf_proxy
|
||||||
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 libs.login import current_account_with_tenant
|
from libs.login import current_account_with_tenant, login_required
|
||||||
from services.file_service import FileService
|
from services.file_service import FileService
|
||||||
|
|
||||||
|
|
||||||
@@ -23,69 +24,73 @@ class RemoteFileUploadPayload(BaseModel):
|
|||||||
url: str = Field(..., description="URL to fetch")
|
url: str = Field(..., description="URL to fetch")
|
||||||
|
|
||||||
|
|
||||||
@console_router.get(
|
@console_ns.route("/remote-files/<path:url>")
|
||||||
"/remote-files/<path:url>",
|
class GetRemoteFileInfo(Resource):
|
||||||
response_model=RemoteFileInfo,
|
@login_required
|
||||||
tags=["console"],
|
def get(self, url: str):
|
||||||
)
|
decoded_url = urllib.parse.unquote(url)
|
||||||
def get_remote_file_info(url: str) -> RemoteFileInfo:
|
resp = ssrf_proxy.head(decoded_url)
|
||||||
decoded_url = urllib.parse.unquote(url)
|
|
||||||
resp = ssrf_proxy.head(decoded_url)
|
|
||||||
if resp.status_code != httpx.codes.OK:
|
|
||||||
resp = ssrf_proxy.get(decoded_url, timeout=3)
|
|
||||||
resp.raise_for_status()
|
|
||||||
return RemoteFileInfo(
|
|
||||||
file_type=resp.headers.get("Content-Type", "application/octet-stream"),
|
|
||||||
file_length=int(resp.headers.get("Content-Length", 0)),
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
@console_router.post(
|
|
||||||
"/remote-files/upload",
|
|
||||||
response_model=FileWithSignedUrl,
|
|
||||||
tags=["console"],
|
|
||||||
status_code=201,
|
|
||||||
)
|
|
||||||
def upload_remote_file(payload: RemoteFileUploadPayload) -> FileWithSignedUrl:
|
|
||||||
url = payload.url
|
|
||||||
|
|
||||||
try:
|
|
||||||
resp = ssrf_proxy.head(url=url)
|
|
||||||
if resp.status_code != httpx.codes.OK:
|
if resp.status_code != httpx.codes.OK:
|
||||||
resp = ssrf_proxy.get(url=url, timeout=3, follow_redirects=True)
|
resp = ssrf_proxy.get(decoded_url, timeout=3)
|
||||||
if resp.status_code != httpx.codes.OK:
|
resp.raise_for_status()
|
||||||
raise RemoteFileUploadError(f"Failed to fetch file from {url}: {resp.text}")
|
return RemoteFileInfo(
|
||||||
except httpx.RequestError as e:
|
file_type=resp.headers.get("Content-Type", "application/octet-stream"),
|
||||||
raise RemoteFileUploadError(f"Failed to fetch file from {url}: {str(e)}")
|
file_length=int(resp.headers.get("Content-Length", 0)),
|
||||||
|
).model_dump(mode="json")
|
||||||
|
|
||||||
file_info = helpers.guess_file_info_from_response(resp)
|
|
||||||
|
|
||||||
if not FileService.is_file_size_within_limit(extension=file_info.extension, file_size=file_info.size):
|
@console_ns.route("/remote-files/upload")
|
||||||
raise FileTooLargeError
|
class RemoteFileUpload(Resource):
|
||||||
|
@login_required
|
||||||
|
def post(self):
|
||||||
|
payload = RemoteFileUploadPayload.model_validate(console_ns.payload)
|
||||||
|
url = payload.url
|
||||||
|
|
||||||
content = resp.content if resp.request.method == "GET" else ssrf_proxy.get(url).content
|
# Try to fetch remote file metadata/content first
|
||||||
|
try:
|
||||||
|
resp = ssrf_proxy.head(url=url)
|
||||||
|
if resp.status_code != httpx.codes.OK:
|
||||||
|
resp = ssrf_proxy.get(url=url, timeout=3, follow_redirects=True)
|
||||||
|
if resp.status_code != httpx.codes.OK:
|
||||||
|
# Normalize into a user-friendly error message expected by tests
|
||||||
|
raise RemoteFileUploadError(f"Failed to fetch file from {url}: {resp.text}")
|
||||||
|
except httpx.RequestError as e:
|
||||||
|
raise RemoteFileUploadError(f"Failed to fetch file from {url}: {str(e)}")
|
||||||
|
|
||||||
try:
|
file_info = helpers.guess_file_info_from_response(resp)
|
||||||
user, _ = current_account_with_tenant()
|
|
||||||
upload_file = FileService(db.engine).upload_file(
|
# Enforce file size limit with 400 (Bad Request) per tests' expectation
|
||||||
filename=file_info.filename,
|
if not FileService.is_file_size_within_limit(extension=file_info.extension, file_size=file_info.size):
|
||||||
content=content,
|
raise FileTooLargeError()
|
||||||
mimetype=file_info.mimetype,
|
|
||||||
user=user,
|
# Load content if needed
|
||||||
source_url=url,
|
content = resp.content if resp.request.method == "GET" else ssrf_proxy.get(url).content
|
||||||
|
|
||||||
|
try:
|
||||||
|
user, _ = current_account_with_tenant()
|
||||||
|
upload_file = FileService(db.engine).upload_file(
|
||||||
|
filename=file_info.filename,
|
||||||
|
content=content,
|
||||||
|
mimetype=file_info.mimetype,
|
||||||
|
user=user,
|
||||||
|
source_url=url,
|
||||||
|
)
|
||||||
|
except services.errors.file.FileTooLargeError as file_too_large_error:
|
||||||
|
raise FileTooLargeError(file_too_large_error.description)
|
||||||
|
except services.errors.file.UnsupportedFileTypeError:
|
||||||
|
raise UnsupportedFileTypeError()
|
||||||
|
|
||||||
|
# Success: return created resource with 201 status
|
||||||
|
return (
|
||||||
|
FileWithSignedUrl(
|
||||||
|
id=upload_file.id,
|
||||||
|
name=upload_file.name,
|
||||||
|
size=upload_file.size,
|
||||||
|
extension=upload_file.extension,
|
||||||
|
url=file_helpers.get_signed_file_url(upload_file_id=upload_file.id),
|
||||||
|
mime_type=upload_file.mime_type,
|
||||||
|
created_by=upload_file.created_by,
|
||||||
|
created_at=int(upload_file.created_at.timestamp()),
|
||||||
|
).model_dump(mode="json"),
|
||||||
|
201,
|
||||||
)
|
)
|
||||||
except services.errors.file.FileTooLargeError as file_too_large_error:
|
|
||||||
raise FileTooLargeError(file_too_large_error.description)
|
|
||||||
except services.errors.file.UnsupportedFileTypeError:
|
|
||||||
raise UnsupportedFileTypeError()
|
|
||||||
|
|
||||||
return FileWithSignedUrl(
|
|
||||||
id=upload_file.id,
|
|
||||||
name=upload_file.name,
|
|
||||||
size=upload_file.size,
|
|
||||||
extension=upload_file.extension,
|
|
||||||
url=file_helpers.get_signed_file_url(upload_file_id=upload_file.id),
|
|
||||||
mime_type=upload_file.mime_type,
|
|
||||||
created_by=upload_file.created_by,
|
|
||||||
created_at=int(upload_file.created_at.timestamp()),
|
|
||||||
)
|
|
||||||
|
|||||||
@@ -42,7 +42,15 @@ class SetupResponse(BaseModel):
|
|||||||
tags=["console"],
|
tags=["console"],
|
||||||
)
|
)
|
||||||
def get_setup_status_api() -> SetupStatusResponse:
|
def get_setup_status_api() -> SetupStatusResponse:
|
||||||
"""Get system setup status."""
|
"""Get system setup status.
|
||||||
|
|
||||||
|
NOTE: This endpoint is unauthenticated by design.
|
||||||
|
|
||||||
|
During first-time bootstrap there is no admin account yet, so frontend initialization must be
|
||||||
|
able to query setup progress before any login flow exists.
|
||||||
|
|
||||||
|
Only bootstrap-safe status information should be returned by this endpoint.
|
||||||
|
"""
|
||||||
if dify_config.EDITION == "SELF_HOSTED":
|
if dify_config.EDITION == "SELF_HOSTED":
|
||||||
setup_status = get_setup_status()
|
setup_status = get_setup_status()
|
||||||
if setup_status and not isinstance(setup_status, bool):
|
if setup_status and not isinstance(setup_status, bool):
|
||||||
@@ -61,7 +69,12 @@ def get_setup_status_api() -> SetupStatusResponse:
|
|||||||
)
|
)
|
||||||
@only_edition_self_hosted
|
@only_edition_self_hosted
|
||||||
def setup_system(payload: SetupRequestPayload) -> SetupResponse:
|
def setup_system(payload: SetupRequestPayload) -> SetupResponse:
|
||||||
"""Initialize system setup with admin account."""
|
"""Initialize system setup with admin account.
|
||||||
|
|
||||||
|
NOTE: This endpoint is unauthenticated by design for first-time bootstrap.
|
||||||
|
Access is restricted by deployment mode (`SELF_HOSTED`), one-time setup guards,
|
||||||
|
and init-password validation rather than user session authentication.
|
||||||
|
"""
|
||||||
if get_setup_status():
|
if get_setup_status():
|
||||||
raise AlreadySetupError()
|
raise AlreadySetupError()
|
||||||
|
|
||||||
|
|||||||
+110
-111
@@ -1,14 +1,27 @@
|
|||||||
from typing import Literal
|
from typing import Literal
|
||||||
from uuid import UUID
|
|
||||||
|
|
||||||
|
from flask import request
|
||||||
|
from flask_restx import Namespace, Resource, fields, marshal_with
|
||||||
from pydantic import BaseModel, Field
|
from pydantic import BaseModel, Field
|
||||||
from werkzeug.exceptions import Forbidden
|
from werkzeug.exceptions import Forbidden
|
||||||
|
|
||||||
|
from controllers.common.schema import register_schema_models
|
||||||
|
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 controllers.fastopenapi import console_router
|
|
||||||
from libs.login import current_account_with_tenant, login_required
|
from libs.login import current_account_with_tenant, login_required
|
||||||
from services.tag_service import TagService
|
from services.tag_service import TagService
|
||||||
|
|
||||||
|
dataset_tag_fields = {
|
||||||
|
"id": fields.String,
|
||||||
|
"name": fields.String,
|
||||||
|
"type": fields.String,
|
||||||
|
"binding_count": fields.String,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def build_dataset_tag_fields(api_or_ns: Namespace):
|
||||||
|
return api_or_ns.model("DataSetTag", dataset_tag_fields)
|
||||||
|
|
||||||
|
|
||||||
class TagBasePayload(BaseModel):
|
class TagBasePayload(BaseModel):
|
||||||
name: str = Field(description="Tag name", min_length=1, max_length=50)
|
name: str = Field(description="Tag name", min_length=1, max_length=50)
|
||||||
@@ -32,129 +45,115 @@ class TagListQueryParam(BaseModel):
|
|||||||
keyword: str | None = Field(None, description="Search keyword")
|
keyword: str | None = Field(None, description="Search keyword")
|
||||||
|
|
||||||
|
|
||||||
class TagResponse(BaseModel):
|
register_schema_models(
|
||||||
id: str = Field(description="Tag ID")
|
console_ns,
|
||||||
name: str = Field(description="Tag name")
|
TagBasePayload,
|
||||||
type: str = Field(description="Tag type")
|
TagBindingPayload,
|
||||||
binding_count: int = Field(description="Number of bindings")
|
TagBindingRemovePayload,
|
||||||
|
TagListQueryParam,
|
||||||
|
|
||||||
class TagBindingResult(BaseModel):
|
|
||||||
result: Literal["success"] = Field(description="Operation result", examples=["success"])
|
|
||||||
|
|
||||||
|
|
||||||
@console_router.get(
|
|
||||||
"/tags",
|
|
||||||
response_model=list[TagResponse],
|
|
||||||
tags=["console"],
|
|
||||||
)
|
)
|
||||||
@setup_required
|
|
||||||
@login_required
|
|
||||||
@account_initialization_required
|
|
||||||
def list_tags(query: TagListQueryParam) -> list[TagResponse]:
|
|
||||||
_, current_tenant_id = current_account_with_tenant()
|
|
||||||
tags = TagService.get_tags(query.type, current_tenant_id, query.keyword)
|
|
||||||
|
|
||||||
return [
|
|
||||||
TagResponse(
|
|
||||||
id=tag.id,
|
|
||||||
name=tag.name,
|
|
||||||
type=tag.type,
|
|
||||||
binding_count=int(tag.binding_count),
|
|
||||||
)
|
|
||||||
for tag in tags
|
|
||||||
]
|
|
||||||
|
|
||||||
|
|
||||||
@console_router.post(
|
@console_ns.route("/tags")
|
||||||
"/tags",
|
class TagListApi(Resource):
|
||||||
response_model=TagResponse,
|
@setup_required
|
||||||
tags=["console"],
|
@login_required
|
||||||
)
|
@account_initialization_required
|
||||||
@setup_required
|
@console_ns.doc(
|
||||||
@login_required
|
params={"type": 'Tag type filter. Can be "knowledge" or "app".', "keyword": "Search keyword for tag name."}
|
||||||
@account_initialization_required
|
)
|
||||||
def create_tag(payload: TagBasePayload) -> TagResponse:
|
@marshal_with(dataset_tag_fields)
|
||||||
current_user, _ = current_account_with_tenant()
|
def get(self):
|
||||||
# The role of the current user in the tag table must be admin, owner, or editor
|
_, current_tenant_id = current_account_with_tenant()
|
||||||
if not (current_user.has_edit_permission or current_user.is_dataset_editor):
|
raw_args = request.args.to_dict()
|
||||||
raise Forbidden()
|
param = TagListQueryParam.model_validate(raw_args)
|
||||||
|
tags = TagService.get_tags(param.type, current_tenant_id, param.keyword)
|
||||||
|
|
||||||
tag = TagService.save_tags(payload.model_dump())
|
return tags, 200
|
||||||
|
|
||||||
return TagResponse(id=tag.id, name=tag.name, type=tag.type, binding_count=0)
|
@console_ns.expect(console_ns.models[TagBasePayload.__name__])
|
||||||
|
@setup_required
|
||||||
|
@login_required
|
||||||
|
@account_initialization_required
|
||||||
|
def post(self):
|
||||||
|
current_user, _ = current_account_with_tenant()
|
||||||
|
# The role of the current user in the ta table must be admin, owner, or editor
|
||||||
|
if not (current_user.has_edit_permission or current_user.is_dataset_editor):
|
||||||
|
raise Forbidden()
|
||||||
|
|
||||||
|
payload = TagBasePayload.model_validate(console_ns.payload or {})
|
||||||
|
tag = TagService.save_tags(payload.model_dump())
|
||||||
|
|
||||||
|
response = {"id": tag.id, "name": tag.name, "type": tag.type, "binding_count": 0}
|
||||||
|
|
||||||
|
return response, 200
|
||||||
|
|
||||||
|
|
||||||
@console_router.patch(
|
@console_ns.route("/tags/<uuid:tag_id>")
|
||||||
"/tags/<uuid:tag_id>",
|
class TagUpdateDeleteApi(Resource):
|
||||||
response_model=TagResponse,
|
@console_ns.expect(console_ns.models[TagBasePayload.__name__])
|
||||||
tags=["console"],
|
@setup_required
|
||||||
)
|
@login_required
|
||||||
@setup_required
|
@account_initialization_required
|
||||||
@login_required
|
def patch(self, tag_id):
|
||||||
@account_initialization_required
|
current_user, _ = current_account_with_tenant()
|
||||||
def update_tag(tag_id: UUID, payload: TagBasePayload) -> TagResponse:
|
tag_id = str(tag_id)
|
||||||
current_user, _ = current_account_with_tenant()
|
# The role of the current user in the ta table must be admin, owner, or editor
|
||||||
tag_id_str = str(tag_id)
|
if not (current_user.has_edit_permission or current_user.is_dataset_editor):
|
||||||
# The role of the current user in the ta table must be admin, owner, or editor
|
raise Forbidden()
|
||||||
if not (current_user.has_edit_permission or current_user.is_dataset_editor):
|
|
||||||
raise Forbidden()
|
|
||||||
|
|
||||||
tag = TagService.update_tags(payload.model_dump(), tag_id_str)
|
payload = TagBasePayload.model_validate(console_ns.payload or {})
|
||||||
|
tag = TagService.update_tags(payload.model_dump(), tag_id)
|
||||||
|
|
||||||
binding_count = TagService.get_tag_binding_count(tag_id_str)
|
binding_count = TagService.get_tag_binding_count(tag_id)
|
||||||
|
|
||||||
return TagResponse(id=tag.id, name=tag.name, type=tag.type, binding_count=binding_count)
|
response = {"id": tag.id, "name": tag.name, "type": tag.type, "binding_count": binding_count}
|
||||||
|
|
||||||
|
return response, 200
|
||||||
|
|
||||||
|
@setup_required
|
||||||
|
@login_required
|
||||||
|
@account_initialization_required
|
||||||
|
@edit_permission_required
|
||||||
|
def delete(self, tag_id):
|
||||||
|
tag_id = str(tag_id)
|
||||||
|
|
||||||
|
TagService.delete_tag(tag_id)
|
||||||
|
|
||||||
|
return "", 204
|
||||||
|
|
||||||
|
|
||||||
@console_router.delete(
|
@console_ns.route("/tag-bindings/create")
|
||||||
"/tags/<uuid:tag_id>",
|
class TagBindingCreateApi(Resource):
|
||||||
tags=["console"],
|
@console_ns.expect(console_ns.models[TagBindingPayload.__name__])
|
||||||
status_code=204,
|
@setup_required
|
||||||
)
|
@login_required
|
||||||
@setup_required
|
@account_initialization_required
|
||||||
@login_required
|
def post(self):
|
||||||
@account_initialization_required
|
current_user, _ = current_account_with_tenant()
|
||||||
@edit_permission_required
|
# The role of the current user in the ta table must be admin, owner, editor, or dataset_operator
|
||||||
def delete_tag(tag_id: UUID) -> None:
|
if not (current_user.has_edit_permission or current_user.is_dataset_editor):
|
||||||
tag_id_str = str(tag_id)
|
raise Forbidden()
|
||||||
|
|
||||||
TagService.delete_tag(tag_id_str)
|
payload = TagBindingPayload.model_validate(console_ns.payload or {})
|
||||||
|
TagService.save_tag_binding(payload.model_dump())
|
||||||
|
|
||||||
|
return {"result": "success"}, 200
|
||||||
|
|
||||||
|
|
||||||
@console_router.post(
|
@console_ns.route("/tag-bindings/remove")
|
||||||
"/tag-bindings/create",
|
class TagBindingDeleteApi(Resource):
|
||||||
response_model=TagBindingResult,
|
@console_ns.expect(console_ns.models[TagBindingRemovePayload.__name__])
|
||||||
tags=["console"],
|
@setup_required
|
||||||
)
|
@login_required
|
||||||
@setup_required
|
@account_initialization_required
|
||||||
@login_required
|
def post(self):
|
||||||
@account_initialization_required
|
current_user, _ = current_account_with_tenant()
|
||||||
def create_tag_binding(payload: TagBindingPayload) -> TagBindingResult:
|
# The role of the current user in the ta table must be admin, owner, editor, or dataset_operator
|
||||||
current_user, _ = current_account_with_tenant()
|
if not (current_user.has_edit_permission or current_user.is_dataset_editor):
|
||||||
# The role of the current user in the tag table must be admin, owner, editor, or dataset_operator
|
raise Forbidden()
|
||||||
if not (current_user.has_edit_permission or current_user.is_dataset_editor):
|
|
||||||
raise Forbidden()
|
|
||||||
|
|
||||||
TagService.save_tag_binding(payload.model_dump())
|
payload = TagBindingRemovePayload.model_validate(console_ns.payload or {})
|
||||||
|
TagService.delete_tag_binding(payload.model_dump())
|
||||||
|
|
||||||
return TagBindingResult(result="success")
|
return {"result": "success"}, 200
|
||||||
|
|
||||||
|
|
||||||
@console_router.post(
|
|
||||||
"/tag-bindings/remove",
|
|
||||||
response_model=TagBindingResult,
|
|
||||||
tags=["console"],
|
|
||||||
)
|
|
||||||
@setup_required
|
|
||||||
@login_required
|
|
||||||
@account_initialization_required
|
|
||||||
def delete_tag_binding(payload: TagBindingRemovePayload) -> TagBindingResult:
|
|
||||||
current_user, _ = current_account_with_tenant()
|
|
||||||
# The role of the current user in the tag table must be admin, owner, editor, or dataset_operator
|
|
||||||
if not (current_user.has_edit_permission or current_user.is_dataset_editor):
|
|
||||||
raise Forbidden()
|
|
||||||
|
|
||||||
TagService.delete_tag_binding(payload.model_dump())
|
|
||||||
|
|
||||||
return TagBindingResult(result="success")
|
|
||||||
|
|||||||
@@ -34,6 +34,8 @@ from .dataset import (
|
|||||||
metadata,
|
metadata,
|
||||||
segment,
|
segment,
|
||||||
)
|
)
|
||||||
|
from .dataset.rag_pipeline import rag_pipeline_workflow
|
||||||
|
from .end_user import end_user
|
||||||
from .workspace import models
|
from .workspace import models
|
||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
@@ -44,6 +46,7 @@ __all__ = [
|
|||||||
"conversation",
|
"conversation",
|
||||||
"dataset",
|
"dataset",
|
||||||
"document",
|
"document",
|
||||||
|
"end_user",
|
||||||
"file",
|
"file",
|
||||||
"file_preview",
|
"file_preview",
|
||||||
"hit_testing",
|
"hit_testing",
|
||||||
@@ -51,6 +54,7 @@ __all__ = [
|
|||||||
"message",
|
"message",
|
||||||
"metadata",
|
"metadata",
|
||||||
"models",
|
"models",
|
||||||
|
"rag_pipeline_workflow",
|
||||||
"segment",
|
"segment",
|
||||||
"site",
|
"site",
|
||||||
"workflow",
|
"workflow",
|
||||||
|
|||||||
@@ -33,8 +33,9 @@ from core.workflow.graph_engine.manager import GraphEngineManager
|
|||||||
from extensions.ext_database import db
|
from extensions.ext_database import db
|
||||||
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 libs import helper
|
from libs import helper
|
||||||
from libs.helper import TimestampField
|
from libs.helper import OptionalTimestampField, TimestampField
|
||||||
from models.model import App, AppMode, EndUser
|
from models.model import App, AppMode, EndUser
|
||||||
|
from models.workflow import WorkflowRun
|
||||||
from repositories.factory import DifyAPIRepositoryFactory
|
from repositories.factory import DifyAPIRepositoryFactory
|
||||||
from services.app_generate_service import AppGenerateService
|
from services.app_generate_service import AppGenerateService
|
||||||
from services.errors.app import IsDraftWorkflowError, WorkflowIdFormatError, WorkflowNotFoundError
|
from services.errors.app import IsDraftWorkflowError, WorkflowIdFormatError, WorkflowNotFoundError
|
||||||
@@ -63,17 +64,32 @@ class WorkflowLogQuery(BaseModel):
|
|||||||
|
|
||||||
register_schema_models(service_api_ns, WorkflowRunPayload, WorkflowLogQuery)
|
register_schema_models(service_api_ns, WorkflowRunPayload, WorkflowLogQuery)
|
||||||
|
|
||||||
|
|
||||||
|
class WorkflowRunStatusField(fields.Raw):
|
||||||
|
def output(self, key, obj: WorkflowRun, **kwargs):
|
||||||
|
return obj.status.value
|
||||||
|
|
||||||
|
|
||||||
|
class WorkflowRunOutputsField(fields.Raw):
|
||||||
|
def output(self, key, obj: WorkflowRun, **kwargs):
|
||||||
|
if obj.status == WorkflowExecutionStatus.PAUSED:
|
||||||
|
return {}
|
||||||
|
|
||||||
|
outputs = obj.outputs_dict
|
||||||
|
return outputs or {}
|
||||||
|
|
||||||
|
|
||||||
workflow_run_fields = {
|
workflow_run_fields = {
|
||||||
"id": fields.String,
|
"id": fields.String,
|
||||||
"workflow_id": fields.String,
|
"workflow_id": fields.String,
|
||||||
"status": fields.String,
|
"status": WorkflowRunStatusField,
|
||||||
"inputs": fields.Raw,
|
"inputs": fields.Raw,
|
||||||
"outputs": fields.Raw,
|
"outputs": WorkflowRunOutputsField,
|
||||||
"error": fields.String,
|
"error": fields.String,
|
||||||
"total_steps": fields.Integer,
|
"total_steps": fields.Integer,
|
||||||
"total_tokens": fields.Integer,
|
"total_tokens": fields.Integer,
|
||||||
"created_at": TimestampField,
|
"created_at": TimestampField,
|
||||||
"finished_at": TimestampField,
|
"finished_at": OptionalTimestampField,
|
||||||
"elapsed_time": fields.Float,
|
"elapsed_time": fields.Float,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -396,7 +396,7 @@ class DatasetApi(DatasetApiResource):
|
|||||||
try:
|
try:
|
||||||
if DatasetService.delete_dataset(dataset_id_str, current_user):
|
if DatasetService.delete_dataset(dataset_id_str, current_user):
|
||||||
DatasetPermissionService.clear_partial_member_list(dataset_id_str)
|
DatasetPermissionService.clear_partial_member_list(dataset_id_str)
|
||||||
return 204
|
return "", 204
|
||||||
else:
|
else:
|
||||||
raise NotFound("Dataset not found.")
|
raise NotFound("Dataset not found.")
|
||||||
except services.errors.dataset.DatasetInUseError:
|
except services.errors.dataset.DatasetInUseError:
|
||||||
@@ -557,7 +557,7 @@ class DatasetTagsApi(DatasetApiResource):
|
|||||||
payload = TagDeletePayload.model_validate(service_api_ns.payload or {})
|
payload = TagDeletePayload.model_validate(service_api_ns.payload or {})
|
||||||
TagService.delete_tag(payload.tag_id)
|
TagService.delete_tag(payload.tag_id)
|
||||||
|
|
||||||
return 204
|
return "", 204
|
||||||
|
|
||||||
|
|
||||||
@service_api_ns.route("/datasets/tags/binding")
|
@service_api_ns.route("/datasets/tags/binding")
|
||||||
@@ -581,7 +581,7 @@ class DatasetTagBindingApi(DatasetApiResource):
|
|||||||
payload = TagBindingPayload.model_validate(service_api_ns.payload or {})
|
payload = TagBindingPayload.model_validate(service_api_ns.payload or {})
|
||||||
TagService.save_tag_binding({"tag_ids": payload.tag_ids, "target_id": payload.target_id, "type": "knowledge"})
|
TagService.save_tag_binding({"tag_ids": payload.tag_ids, "target_id": payload.target_id, "type": "knowledge"})
|
||||||
|
|
||||||
return 204
|
return "", 204
|
||||||
|
|
||||||
|
|
||||||
@service_api_ns.route("/datasets/tags/unbinding")
|
@service_api_ns.route("/datasets/tags/unbinding")
|
||||||
@@ -605,7 +605,7 @@ class DatasetTagUnbindingApi(DatasetApiResource):
|
|||||||
payload = TagUnbindingPayload.model_validate(service_api_ns.payload or {})
|
payload = TagUnbindingPayload.model_validate(service_api_ns.payload or {})
|
||||||
TagService.delete_tag_binding({"tag_id": payload.tag_id, "target_id": payload.target_id, "type": "knowledge"})
|
TagService.delete_tag_binding({"tag_id": payload.tag_id, "target_id": payload.target_id, "type": "knowledge"})
|
||||||
|
|
||||||
return 204
|
return "", 204
|
||||||
|
|
||||||
|
|
||||||
@service_api_ns.route("/datasets/<uuid:dataset_id>/tags")
|
@service_api_ns.route("/datasets/<uuid:dataset_id>/tags")
|
||||||
|
|||||||
@@ -746,4 +746,4 @@ class DocumentApi(DatasetApiResource):
|
|||||||
except services.errors.document.DocumentIndexingError:
|
except services.errors.document.DocumentIndexingError:
|
||||||
raise DocumentIndexingError("Cannot delete document during indexing.")
|
raise DocumentIndexingError("Cannot delete document during indexing.")
|
||||||
|
|
||||||
return 204
|
return "", 204
|
||||||
|
|||||||
@@ -128,7 +128,7 @@ class DatasetMetadataServiceApi(DatasetApiResource):
|
|||||||
DatasetService.check_dataset_permission(dataset, current_user)
|
DatasetService.check_dataset_permission(dataset, current_user)
|
||||||
|
|
||||||
MetadataService.delete_metadata(dataset_id_str, metadata_id_str)
|
MetadataService.delete_metadata(dataset_id_str, metadata_id_str)
|
||||||
return 204
|
return "", 204
|
||||||
|
|
||||||
|
|
||||||
@service_api_ns.route("/datasets/<uuid:dataset_id>/metadata/built-in")
|
@service_api_ns.route("/datasets/<uuid:dataset_id>/metadata/built-in")
|
||||||
|
|||||||
@@ -1,5 +1,3 @@
|
|||||||
import string
|
|
||||||
import uuid
|
|
||||||
from collections.abc import Generator
|
from collections.abc import Generator
|
||||||
from typing import Any
|
from typing import Any
|
||||||
|
|
||||||
@@ -12,6 +10,7 @@ from controllers.common.errors import FilenameNotExistsError, NoFileUploadedErro
|
|||||||
from controllers.common.schema import register_schema_model
|
from controllers.common.schema import register_schema_model
|
||||||
from controllers.service_api import service_api_ns
|
from controllers.service_api import service_api_ns
|
||||||
from controllers.service_api.dataset.error import PipelineRunError
|
from controllers.service_api.dataset.error import PipelineRunError
|
||||||
|
from controllers.service_api.dataset.rag_pipeline.serializers import serialize_upload_file
|
||||||
from controllers.service_api.wraps import DatasetApiResource
|
from controllers.service_api.wraps import DatasetApiResource
|
||||||
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
|
||||||
@@ -41,7 +40,7 @@ register_schema_model(service_api_ns, DatasourceNodeRunPayload)
|
|||||||
register_schema_model(service_api_ns, PipelineRunApiEntity)
|
register_schema_model(service_api_ns, PipelineRunApiEntity)
|
||||||
|
|
||||||
|
|
||||||
@service_api_ns.route(f"/datasets/{uuid:dataset_id}/pipeline/datasource-plugins")
|
@service_api_ns.route("/datasets/<uuid:dataset_id>/pipeline/datasource-plugins")
|
||||||
class DatasourcePluginsApi(DatasetApiResource):
|
class DatasourcePluginsApi(DatasetApiResource):
|
||||||
"""Resource for datasource plugins."""
|
"""Resource for datasource plugins."""
|
||||||
|
|
||||||
@@ -76,7 +75,7 @@ class DatasourcePluginsApi(DatasetApiResource):
|
|||||||
return datasource_plugins, 200
|
return datasource_plugins, 200
|
||||||
|
|
||||||
|
|
||||||
@service_api_ns.route(f"/datasets/{uuid:dataset_id}/pipeline/datasource/nodes/{string:node_id}/run")
|
@service_api_ns.route("/datasets/<uuid:dataset_id>/pipeline/datasource/nodes/<string:node_id>/run")
|
||||||
class DatasourceNodeRunApi(DatasetApiResource):
|
class DatasourceNodeRunApi(DatasetApiResource):
|
||||||
"""Resource for datasource node run."""
|
"""Resource for datasource node run."""
|
||||||
|
|
||||||
@@ -131,7 +130,7 @@ class DatasourceNodeRunApi(DatasetApiResource):
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@service_api_ns.route(f"/datasets/{uuid:dataset_id}/pipeline/run")
|
@service_api_ns.route("/datasets/<uuid:dataset_id>/pipeline/run")
|
||||||
class PipelineRunApi(DatasetApiResource):
|
class PipelineRunApi(DatasetApiResource):
|
||||||
"""Resource for datasource node run."""
|
"""Resource for datasource node run."""
|
||||||
|
|
||||||
@@ -232,12 +231,4 @@ class KnowledgebasePipelineFileUploadApi(DatasetApiResource):
|
|||||||
except services.errors.file.UnsupportedFileTypeError:
|
except services.errors.file.UnsupportedFileTypeError:
|
||||||
raise UnsupportedFileTypeError()
|
raise UnsupportedFileTypeError()
|
||||||
|
|
||||||
return {
|
return serialize_upload_file(upload_file), 201
|
||||||
"id": upload_file.id,
|
|
||||||
"name": upload_file.name,
|
|
||||||
"size": upload_file.size,
|
|
||||||
"extension": upload_file.extension,
|
|
||||||
"mime_type": upload_file.mime_type,
|
|
||||||
"created_by": upload_file.created_by,
|
|
||||||
"created_at": upload_file.created_at,
|
|
||||||
}, 201
|
|
||||||
|
|||||||
@@ -0,0 +1,22 @@
|
|||||||
|
"""
|
||||||
|
Serialization helpers for Service API knowledge pipeline endpoints.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from typing import TYPE_CHECKING, Any
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from models.model import UploadFile
|
||||||
|
|
||||||
|
|
||||||
|
def serialize_upload_file(upload_file: UploadFile) -> dict[str, Any]:
|
||||||
|
return {
|
||||||
|
"id": upload_file.id,
|
||||||
|
"name": upload_file.name,
|
||||||
|
"size": upload_file.size,
|
||||||
|
"extension": upload_file.extension,
|
||||||
|
"mime_type": upload_file.mime_type,
|
||||||
|
"created_by": upload_file.created_by,
|
||||||
|
"created_at": upload_file.created_at.isoformat() if upload_file.created_at else None,
|
||||||
|
}
|
||||||
@@ -233,7 +233,7 @@ class DatasetSegmentApi(DatasetApiResource):
|
|||||||
if not segment:
|
if not segment:
|
||||||
raise NotFound("Segment not found.")
|
raise NotFound("Segment not found.")
|
||||||
SegmentService.delete_segment(segment, document, dataset)
|
SegmentService.delete_segment(segment, document, dataset)
|
||||||
return 204
|
return "", 204
|
||||||
|
|
||||||
@service_api_ns.expect(service_api_ns.models[SegmentUpdatePayload.__name__])
|
@service_api_ns.expect(service_api_ns.models[SegmentUpdatePayload.__name__])
|
||||||
@service_api_ns.doc("update_segment")
|
@service_api_ns.doc("update_segment")
|
||||||
@@ -499,7 +499,7 @@ class DatasetChildChunkApi(DatasetApiResource):
|
|||||||
except ChildChunkDeleteIndexServiceError as e:
|
except ChildChunkDeleteIndexServiceError as e:
|
||||||
raise ChildChunkDeleteIndexError(str(e))
|
raise ChildChunkDeleteIndexError(str(e))
|
||||||
|
|
||||||
return 204
|
return "", 204
|
||||||
|
|
||||||
@service_api_ns.expect(service_api_ns.models[ChildChunkUpdatePayload.__name__])
|
@service_api_ns.expect(service_api_ns.models[ChildChunkUpdatePayload.__name__])
|
||||||
@service_api_ns.doc("update_child_chunk")
|
@service_api_ns.doc("update_child_chunk")
|
||||||
|
|||||||
@@ -0,0 +1,3 @@
|
|||||||
|
from . import end_user
|
||||||
|
|
||||||
|
__all__ = ["end_user"]
|
||||||
@@ -0,0 +1,41 @@
|
|||||||
|
from uuid import UUID
|
||||||
|
|
||||||
|
from flask_restx import Resource
|
||||||
|
|
||||||
|
from controllers.service_api import service_api_ns
|
||||||
|
from controllers.service_api.end_user.error import EndUserNotFoundError
|
||||||
|
from controllers.service_api.wraps import validate_app_token
|
||||||
|
from fields.end_user_fields import EndUserDetail
|
||||||
|
from models.model import App
|
||||||
|
from services.end_user_service import EndUserService
|
||||||
|
|
||||||
|
|
||||||
|
@service_api_ns.route("/end-users/<uuid:end_user_id>")
|
||||||
|
class EndUserApi(Resource):
|
||||||
|
"""Resource for retrieving end user details by ID."""
|
||||||
|
|
||||||
|
@service_api_ns.doc("get_end_user")
|
||||||
|
@service_api_ns.doc(description="Get an end user by ID")
|
||||||
|
@service_api_ns.doc(
|
||||||
|
params={"end_user_id": "End user ID"},
|
||||||
|
responses={
|
||||||
|
200: "End user retrieved successfully",
|
||||||
|
401: "Unauthorized - invalid API token",
|
||||||
|
404: "End user not found",
|
||||||
|
},
|
||||||
|
)
|
||||||
|
@validate_app_token
|
||||||
|
def get(self, app_model: App, end_user_id: UUID):
|
||||||
|
"""Get end user detail.
|
||||||
|
|
||||||
|
This endpoint is scoped to the current app token's tenant/app to prevent
|
||||||
|
cross-tenant/app access when an end-user ID is known.
|
||||||
|
"""
|
||||||
|
|
||||||
|
end_user = EndUserService.get_end_user_by_id(
|
||||||
|
tenant_id=app_model.tenant_id, app_id=app_model.id, end_user_id=str(end_user_id)
|
||||||
|
)
|
||||||
|
if end_user is None:
|
||||||
|
raise EndUserNotFoundError()
|
||||||
|
|
||||||
|
return EndUserDetail.model_validate(end_user).model_dump(mode="json")
|
||||||
@@ -0,0 +1,7 @@
|
|||||||
|
from libs.exception import BaseHTTPException
|
||||||
|
|
||||||
|
|
||||||
|
class EndUserNotFoundError(BaseHTTPException):
|
||||||
|
error_code = "end_user_not_found"
|
||||||
|
description = "End user not found."
|
||||||
|
code = 404
|
||||||
@@ -1,27 +1,24 @@
|
|||||||
import logging
|
import logging
|
||||||
import time
|
import time
|
||||||
from collections.abc import Callable
|
from collections.abc import Callable
|
||||||
from datetime import timedelta
|
|
||||||
from enum import StrEnum, auto
|
from enum import StrEnum, auto
|
||||||
from functools import wraps
|
from functools import wraps
|
||||||
from typing import Concatenate, ParamSpec, TypeVar
|
from typing import Concatenate, ParamSpec, TypeVar, cast
|
||||||
|
|
||||||
from flask import current_app, request
|
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, update
|
|
||||||
from sqlalchemy.orm import Session
|
|
||||||
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
|
||||||
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 libs.datetime_utils import naive_utc_now
|
|
||||||
from libs.login import current_user
|
from libs.login import current_user
|
||||||
from models import Account, Tenant, TenantAccountJoin, TenantStatus
|
from models import Account, Tenant, TenantAccountJoin, TenantStatus
|
||||||
from models.dataset import Dataset, RateLimitLog
|
from models.dataset import Dataset, RateLimitLog
|
||||||
from models.model import ApiToken, App
|
from models.model import ApiToken, App
|
||||||
|
from services.api_token_service import ApiTokenCache, fetch_token_with_single_flight, record_token_usage
|
||||||
from services.end_user_service import EndUserService
|
from services.end_user_service import EndUserService
|
||||||
from services.feature_service import FeatureService
|
from services.feature_service import FeatureService
|
||||||
|
|
||||||
@@ -220,6 +217,8 @@ def validate_dataset_token(view: Callable[Concatenate[T, P], R] | None = None):
|
|||||||
def decorator(view: Callable[Concatenate[T, P], R]):
|
def decorator(view: Callable[Concatenate[T, P], R]):
|
||||||
@wraps(view)
|
@wraps(view)
|
||||||
def decorated(*args: P.args, **kwargs: P.kwargs):
|
def decorated(*args: P.args, **kwargs: P.kwargs):
|
||||||
|
api_token = validate_and_get_api_token("dataset")
|
||||||
|
|
||||||
# get url path dataset_id from positional args or kwargs
|
# get url path dataset_id from positional args or kwargs
|
||||||
# Flask passes URL path parameters as positional arguments
|
# Flask passes URL path parameters as positional arguments
|
||||||
dataset_id = None
|
dataset_id = None
|
||||||
@@ -256,12 +255,18 @@ def validate_dataset_token(view: Callable[Concatenate[T, P], R] | None = None):
|
|||||||
# 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 = db.session.query(Dataset).where(Dataset.id == dataset_id).first()
|
dataset = (
|
||||||
|
db.session.query(Dataset)
|
||||||
|
.where(
|
||||||
|
Dataset.id == dataset_id,
|
||||||
|
Dataset.tenant_id == api_token.tenant_id,
|
||||||
|
)
|
||||||
|
.first()
|
||||||
|
)
|
||||||
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.")
|
||||||
api_token = validate_and_get_api_token("dataset")
|
|
||||||
tenant_account_join = (
|
tenant_account_join = (
|
||||||
db.session.query(Tenant, TenantAccountJoin)
|
db.session.query(Tenant, TenantAccountJoin)
|
||||||
.where(Tenant.id == api_token.tenant_id)
|
.where(Tenant.id == api_token.tenant_id)
|
||||||
@@ -296,7 +301,14 @@ def validate_dataset_token(view: Callable[Concatenate[T, P], R] | None = None):
|
|||||||
|
|
||||||
def validate_and_get_api_token(scope: str | None = None):
|
def validate_and_get_api_token(scope: str | None = None):
|
||||||
"""
|
"""
|
||||||
Validate and get API token.
|
Validate and get API token with Redis caching.
|
||||||
|
|
||||||
|
This function uses a two-tier approach:
|
||||||
|
1. First checks Redis cache for the token
|
||||||
|
2. If not cached, queries database and caches the result
|
||||||
|
|
||||||
|
The last_used_at field is updated asynchronously via Celery task
|
||||||
|
to avoid blocking the request.
|
||||||
"""
|
"""
|
||||||
auth_header = request.headers.get("Authorization")
|
auth_header = request.headers.get("Authorization")
|
||||||
if auth_header is None or " " not in auth_header:
|
if auth_header is None or " " not in auth_header:
|
||||||
@@ -308,29 +320,18 @@ def validate_and_get_api_token(scope: str | None = None):
|
|||||||
if auth_scheme != "bearer":
|
if auth_scheme != "bearer":
|
||||||
raise Unauthorized("Authorization scheme must be 'Bearer'")
|
raise Unauthorized("Authorization scheme must be 'Bearer'")
|
||||||
|
|
||||||
current_time = naive_utc_now()
|
# Try to get token from cache first
|
||||||
cutoff_time = current_time - timedelta(minutes=1)
|
# Returns a CachedApiToken (plain Python object), not a SQLAlchemy model
|
||||||
with Session(db.engine, expire_on_commit=False) as session:
|
cached_token = ApiTokenCache.get(auth_token, scope)
|
||||||
update_stmt = (
|
if cached_token is not None:
|
||||||
update(ApiToken)
|
logger.debug("Token validation served from cache for scope: %s", scope)
|
||||||
.where(
|
# Record usage in Redis for later batch update (no Celery task per request)
|
||||||
ApiToken.token == auth_token,
|
record_token_usage(auth_token, scope)
|
||||||
(ApiToken.last_used_at.is_(None) | (ApiToken.last_used_at < cutoff_time)),
|
return cast(ApiToken, cached_token)
|
||||||
ApiToken.type == scope,
|
|
||||||
)
|
|
||||||
.values(last_used_at=current_time)
|
|
||||||
)
|
|
||||||
stmt = select(ApiToken).where(ApiToken.token == auth_token, ApiToken.type == scope)
|
|
||||||
result = session.execute(update_stmt)
|
|
||||||
api_token = session.scalar(stmt)
|
|
||||||
|
|
||||||
if hasattr(result, "rowcount") and result.rowcount > 0:
|
# Cache miss - use Redis lock for single-flight mode
|
||||||
session.commit()
|
# This ensures only one request queries DB for the same token concurrently
|
||||||
|
return fetch_token_with_single_flight(auth_token, scope)
|
||||||
if not api_token:
|
|
||||||
raise Unauthorized("Access token is invalid")
|
|
||||||
|
|
||||||
return api_token
|
|
||||||
|
|
||||||
|
|
||||||
class DatasetApiResource(Resource):
|
class DatasetApiResource(Resource):
|
||||||
|
|||||||
@@ -23,6 +23,7 @@ from . import (
|
|||||||
feature,
|
feature,
|
||||||
files,
|
files,
|
||||||
forgot_password,
|
forgot_password,
|
||||||
|
human_input_form,
|
||||||
login,
|
login,
|
||||||
message,
|
message,
|
||||||
passport,
|
passport,
|
||||||
@@ -30,6 +31,7 @@ from . import (
|
|||||||
saved_message,
|
saved_message,
|
||||||
site,
|
site,
|
||||||
workflow,
|
workflow,
|
||||||
|
workflow_events,
|
||||||
)
|
)
|
||||||
|
|
||||||
api.add_namespace(web_ns)
|
api.add_namespace(web_ns)
|
||||||
@@ -44,6 +46,7 @@ __all__ = [
|
|||||||
"feature",
|
"feature",
|
||||||
"files",
|
"files",
|
||||||
"forgot_password",
|
"forgot_password",
|
||||||
|
"human_input_form",
|
||||||
"login",
|
"login",
|
||||||
"message",
|
"message",
|
||||||
"passport",
|
"passport",
|
||||||
@@ -52,4 +55,5 @@ __all__ = [
|
|||||||
"site",
|
"site",
|
||||||
"web_ns",
|
"web_ns",
|
||||||
"workflow",
|
"workflow",
|
||||||
|
"workflow_events",
|
||||||
]
|
]
|
||||||
|
|||||||
@@ -117,6 +117,12 @@ class InvokeRateLimitError(BaseHTTPException):
|
|||||||
code = 429
|
code = 429
|
||||||
|
|
||||||
|
|
||||||
|
class WebFormRateLimitExceededError(BaseHTTPException):
|
||||||
|
error_code = "web_form_rate_limit_exceeded"
|
||||||
|
description = "Too many form requests. Please try again later."
|
||||||
|
code = 429
|
||||||
|
|
||||||
|
|
||||||
class NotFoundError(BaseHTTPException):
|
class NotFoundError(BaseHTTPException):
|
||||||
error_code = "not_found"
|
error_code = "not_found"
|
||||||
code = 404
|
code = 404
|
||||||
|
|||||||
@@ -0,0 +1,161 @@
|
|||||||
|
"""
|
||||||
|
Web App Human Input Form APIs.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import json
|
||||||
|
import logging
|
||||||
|
from datetime import datetime
|
||||||
|
|
||||||
|
from flask import Response, request
|
||||||
|
from flask_restx import Resource, reqparse
|
||||||
|
from werkzeug.exceptions import Forbidden
|
||||||
|
|
||||||
|
from configs import dify_config
|
||||||
|
from controllers.web import web_ns
|
||||||
|
from controllers.web.error import NotFoundError, WebFormRateLimitExceededError
|
||||||
|
from controllers.web.site import serialize_app_site_payload
|
||||||
|
from extensions.ext_database import db
|
||||||
|
from libs.helper import RateLimiter, extract_remote_ip
|
||||||
|
from models.account import TenantStatus
|
||||||
|
from models.model import App, Site
|
||||||
|
from services.human_input_service import Form, FormNotFoundError, HumanInputService
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
_FORM_SUBMIT_RATE_LIMITER = RateLimiter(
|
||||||
|
prefix="web_form_submit_rate_limit",
|
||||||
|
max_attempts=dify_config.WEB_FORM_SUBMIT_RATE_LIMIT_MAX_ATTEMPTS,
|
||||||
|
time_window=dify_config.WEB_FORM_SUBMIT_RATE_LIMIT_WINDOW_SECONDS,
|
||||||
|
)
|
||||||
|
_FORM_ACCESS_RATE_LIMITER = RateLimiter(
|
||||||
|
prefix="web_form_access_rate_limit",
|
||||||
|
max_attempts=dify_config.WEB_FORM_SUBMIT_RATE_LIMIT_MAX_ATTEMPTS,
|
||||||
|
time_window=dify_config.WEB_FORM_SUBMIT_RATE_LIMIT_WINDOW_SECONDS,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _stringify_default_values(values: dict[str, object]) -> dict[str, str]:
|
||||||
|
result: dict[str, str] = {}
|
||||||
|
for key, value in values.items():
|
||||||
|
if value is None:
|
||||||
|
result[key] = ""
|
||||||
|
elif isinstance(value, (dict, list)):
|
||||||
|
result[key] = json.dumps(value, ensure_ascii=False)
|
||||||
|
else:
|
||||||
|
result[key] = str(value)
|
||||||
|
return result
|
||||||
|
|
||||||
|
|
||||||
|
def _to_timestamp(value: datetime) -> int:
|
||||||
|
return int(value.timestamp())
|
||||||
|
|
||||||
|
|
||||||
|
def _jsonify_form_definition(form: Form, site_payload: dict | None = None) -> Response:
|
||||||
|
"""Return the form payload (optionally with site) as a JSON response."""
|
||||||
|
definition_payload = form.get_definition().model_dump()
|
||||||
|
payload = {
|
||||||
|
"form_content": definition_payload["rendered_content"],
|
||||||
|
"inputs": definition_payload["inputs"],
|
||||||
|
"resolved_default_values": _stringify_default_values(definition_payload["default_values"]),
|
||||||
|
"user_actions": definition_payload["user_actions"],
|
||||||
|
"expiration_time": _to_timestamp(form.expiration_time),
|
||||||
|
}
|
||||||
|
if site_payload is not None:
|
||||||
|
payload["site"] = site_payload
|
||||||
|
return Response(json.dumps(payload, ensure_ascii=False), mimetype="application/json")
|
||||||
|
|
||||||
|
|
||||||
|
@web_ns.route("/form/human_input/<string:form_token>")
|
||||||
|
class HumanInputFormApi(Resource):
|
||||||
|
"""API for getting and submitting human input forms via the web app."""
|
||||||
|
|
||||||
|
# NOTE(QuantumGhost): this endpoint is unauthenticated on purpose for now.
|
||||||
|
|
||||||
|
# def get(self, _app_model: App, _end_user: EndUser, form_token: str):
|
||||||
|
def get(self, form_token: str):
|
||||||
|
"""
|
||||||
|
Get human input form definition by token.
|
||||||
|
|
||||||
|
GET /api/form/human_input/<form_token>
|
||||||
|
"""
|
||||||
|
ip_address = extract_remote_ip(request)
|
||||||
|
if _FORM_ACCESS_RATE_LIMITER.is_rate_limited(ip_address):
|
||||||
|
raise WebFormRateLimitExceededError()
|
||||||
|
_FORM_ACCESS_RATE_LIMITER.increment_rate_limit(ip_address)
|
||||||
|
|
||||||
|
service = HumanInputService(db.engine)
|
||||||
|
# TODO(QuantumGhost): forbid submision for form tokens
|
||||||
|
# that are only for console.
|
||||||
|
form = service.get_form_by_token(form_token)
|
||||||
|
|
||||||
|
if form is None:
|
||||||
|
raise NotFoundError("Form not found")
|
||||||
|
|
||||||
|
service.ensure_form_active(form)
|
||||||
|
app_model, site = _get_app_site_from_form(form)
|
||||||
|
|
||||||
|
return _jsonify_form_definition(form, site_payload=serialize_app_site_payload(app_model, site, None))
|
||||||
|
|
||||||
|
# def post(self, _app_model: App, _end_user: EndUser, form_token: str):
|
||||||
|
def post(self, form_token: str):
|
||||||
|
"""
|
||||||
|
Submit human input form by token.
|
||||||
|
|
||||||
|
POST /api/form/human_input/<form_token>
|
||||||
|
|
||||||
|
Request body:
|
||||||
|
{
|
||||||
|
"inputs": {
|
||||||
|
"content": "User input content"
|
||||||
|
},
|
||||||
|
"action": "Approve"
|
||||||
|
}
|
||||||
|
"""
|
||||||
|
parser = reqparse.RequestParser()
|
||||||
|
parser.add_argument("inputs", type=dict, required=True, location="json")
|
||||||
|
parser.add_argument("action", type=str, required=True, location="json")
|
||||||
|
args = parser.parse_args()
|
||||||
|
|
||||||
|
ip_address = extract_remote_ip(request)
|
||||||
|
if _FORM_SUBMIT_RATE_LIMITER.is_rate_limited(ip_address):
|
||||||
|
raise WebFormRateLimitExceededError()
|
||||||
|
_FORM_SUBMIT_RATE_LIMITER.increment_rate_limit(ip_address)
|
||||||
|
|
||||||
|
service = HumanInputService(db.engine)
|
||||||
|
form = service.get_form_by_token(form_token)
|
||||||
|
if form is None:
|
||||||
|
raise NotFoundError("Form not found")
|
||||||
|
|
||||||
|
if (recipient_type := form.recipient_type) is None:
|
||||||
|
logger.warning("Recipient type is None for form, form_id=%", form.id)
|
||||||
|
raise AssertionError("Recipient type is None")
|
||||||
|
|
||||||
|
try:
|
||||||
|
service.submit_form_by_token(
|
||||||
|
recipient_type=recipient_type,
|
||||||
|
form_token=form_token,
|
||||||
|
selected_action_id=args["action"],
|
||||||
|
form_data=args["inputs"],
|
||||||
|
submission_end_user_id=None,
|
||||||
|
# submission_end_user_id=_end_user.id,
|
||||||
|
)
|
||||||
|
except FormNotFoundError:
|
||||||
|
raise NotFoundError("Form not found")
|
||||||
|
|
||||||
|
return {}, 200
|
||||||
|
|
||||||
|
|
||||||
|
def _get_app_site_from_form(form: Form) -> tuple[App, Site]:
|
||||||
|
"""Resolve App/Site for the form's app and validate tenant status."""
|
||||||
|
app_model = db.session.query(App).where(App.id == form.app_id).first()
|
||||||
|
if app_model is None or app_model.tenant_id != form.tenant_id:
|
||||||
|
raise NotFoundError("Form not found")
|
||||||
|
|
||||||
|
site = db.session.query(Site).where(Site.app_id == app_model.id).first()
|
||||||
|
if site is None:
|
||||||
|
raise Forbidden()
|
||||||
|
|
||||||
|
if app_model.tenant and app_model.tenant.status == TenantStatus.ARCHIVE:
|
||||||
|
raise Forbidden()
|
||||||
|
|
||||||
|
return app_model, site
|
||||||
@@ -1,4 +1,6 @@
|
|||||||
from flask_restx import fields, marshal_with
|
from typing import cast
|
||||||
|
|
||||||
|
from flask_restx import fields, marshal, marshal_with
|
||||||
from werkzeug.exceptions import Forbidden
|
from werkzeug.exceptions import Forbidden
|
||||||
|
|
||||||
from configs import dify_config
|
from configs import dify_config
|
||||||
@@ -7,7 +9,7 @@ from controllers.web.wraps import WebApiResource
|
|||||||
from extensions.ext_database import db
|
from extensions.ext_database import db
|
||||||
from libs.helper import AppIconUrlField
|
from libs.helper import AppIconUrlField
|
||||||
from models.account import TenantStatus
|
from models.account import TenantStatus
|
||||||
from models.model import Site
|
from models.model import App, Site
|
||||||
from services.feature_service import FeatureService
|
from services.feature_service import FeatureService
|
||||||
|
|
||||||
|
|
||||||
@@ -108,3 +110,14 @@ class AppSiteInfo:
|
|||||||
"remove_webapp_brand": remove_webapp_brand,
|
"remove_webapp_brand": remove_webapp_brand,
|
||||||
"replace_webapp_logo": replace_webapp_logo,
|
"replace_webapp_logo": replace_webapp_logo,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def serialize_site(site: Site) -> dict:
|
||||||
|
"""Serialize Site model using the same schema as AppSiteApi."""
|
||||||
|
return cast(dict, marshal(site, AppSiteApi.site_fields))
|
||||||
|
|
||||||
|
|
||||||
|
def serialize_app_site_payload(app_model: App, site: Site, end_user_id: str | None) -> dict:
|
||||||
|
can_replace_logo = FeatureService.get_features(app_model.tenant_id).can_replace_logo
|
||||||
|
app_site_info = AppSiteInfo(app_model.tenant, app_model, site, end_user_id, can_replace_logo)
|
||||||
|
return cast(dict, marshal(app_site_info, AppSiteApi.app_fields))
|
||||||
|
|||||||
@@ -0,0 +1,112 @@
|
|||||||
|
"""
|
||||||
|
Web App Workflow Resume APIs.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import json
|
||||||
|
from collections.abc import Generator
|
||||||
|
|
||||||
|
from flask import Response, request
|
||||||
|
from sqlalchemy.orm import sessionmaker
|
||||||
|
|
||||||
|
from controllers.web import api
|
||||||
|
from controllers.web.error import InvalidArgumentError, NotFoundError
|
||||||
|
from controllers.web.wraps import WebApiResource
|
||||||
|
from core.app.apps.advanced_chat.app_generator import AdvancedChatAppGenerator
|
||||||
|
from core.app.apps.base_app_generator import BaseAppGenerator
|
||||||
|
from core.app.apps.common.workflow_response_converter import WorkflowResponseConverter
|
||||||
|
from core.app.apps.message_generator import MessageGenerator
|
||||||
|
from core.app.apps.workflow.app_generator import WorkflowAppGenerator
|
||||||
|
from extensions.ext_database import db
|
||||||
|
from models.enums import CreatorUserRole
|
||||||
|
from models.model import App, AppMode, EndUser
|
||||||
|
from repositories.factory import DifyAPIRepositoryFactory
|
||||||
|
from services.workflow_event_snapshot_service import build_workflow_event_stream
|
||||||
|
|
||||||
|
|
||||||
|
class WorkflowEventsApi(WebApiResource):
|
||||||
|
"""API for getting workflow execution events after resume."""
|
||||||
|
|
||||||
|
def get(self, app_model: App, end_user: EndUser, task_id: str):
|
||||||
|
"""
|
||||||
|
Get workflow execution events stream after resume.
|
||||||
|
|
||||||
|
GET /api/workflow/<task_id>/events
|
||||||
|
|
||||||
|
Returns Server-Sent Events stream.
|
||||||
|
"""
|
||||||
|
workflow_run_id = task_id
|
||||||
|
session_maker = sessionmaker(db.engine)
|
||||||
|
repo = DifyAPIRepositoryFactory.create_api_workflow_run_repository(session_maker)
|
||||||
|
workflow_run = repo.get_workflow_run_by_id_and_tenant_id(
|
||||||
|
tenant_id=app_model.tenant_id,
|
||||||
|
run_id=workflow_run_id,
|
||||||
|
)
|
||||||
|
|
||||||
|
if workflow_run is None:
|
||||||
|
raise NotFoundError(f"WorkflowRun not found, id={workflow_run_id}")
|
||||||
|
|
||||||
|
if workflow_run.app_id != app_model.id:
|
||||||
|
raise NotFoundError(f"WorkflowRun not found, id={workflow_run_id}")
|
||||||
|
|
||||||
|
if workflow_run.created_by_role != CreatorUserRole.END_USER:
|
||||||
|
raise NotFoundError(f"WorkflowRun not created by end user, id={workflow_run_id}")
|
||||||
|
|
||||||
|
if workflow_run.created_by != end_user.id:
|
||||||
|
raise NotFoundError(f"WorkflowRun not created by the current end user, id={workflow_run_id}")
|
||||||
|
|
||||||
|
if workflow_run.finished_at is not None:
|
||||||
|
response = WorkflowResponseConverter.workflow_run_result_to_finish_response(
|
||||||
|
task_id=workflow_run.id,
|
||||||
|
workflow_run=workflow_run,
|
||||||
|
creator_user=end_user,
|
||||||
|
)
|
||||||
|
|
||||||
|
payload = response.model_dump(mode="json")
|
||||||
|
payload["event"] = response.event.value
|
||||||
|
|
||||||
|
def _generate_finished_events() -> Generator[str, None, None]:
|
||||||
|
yield f"data: {json.dumps(payload)}\n\n"
|
||||||
|
|
||||||
|
event_generator = _generate_finished_events
|
||||||
|
else:
|
||||||
|
app_mode = AppMode.value_of(app_model.mode)
|
||||||
|
msg_generator = MessageGenerator()
|
||||||
|
generator: BaseAppGenerator
|
||||||
|
if app_mode == AppMode.ADVANCED_CHAT:
|
||||||
|
generator = AdvancedChatAppGenerator()
|
||||||
|
elif app_mode == AppMode.WORKFLOW:
|
||||||
|
generator = WorkflowAppGenerator()
|
||||||
|
else:
|
||||||
|
raise InvalidArgumentError(f"cannot subscribe to workflow run, workflow_run_id={workflow_run.id}")
|
||||||
|
|
||||||
|
include_state_snapshot = request.args.get("include_state_snapshot", "false").lower() == "true"
|
||||||
|
|
||||||
|
def _generate_stream_events():
|
||||||
|
if include_state_snapshot:
|
||||||
|
return generator.convert_to_event_stream(
|
||||||
|
build_workflow_event_stream(
|
||||||
|
app_mode=app_mode,
|
||||||
|
workflow_run=workflow_run,
|
||||||
|
tenant_id=app_model.tenant_id,
|
||||||
|
app_id=app_model.id,
|
||||||
|
session_maker=session_maker,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
return generator.convert_to_event_stream(
|
||||||
|
msg_generator.retrieve_events(app_mode, workflow_run.id),
|
||||||
|
)
|
||||||
|
|
||||||
|
event_generator = _generate_stream_events
|
||||||
|
|
||||||
|
return Response(
|
||||||
|
event_generator(),
|
||||||
|
mimetype="text/event-stream",
|
||||||
|
headers={
|
||||||
|
"Cache-Control": "no-cache",
|
||||||
|
"Connection": "keep-alive",
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
# Register the APIs
|
||||||
|
api.add_resource(WorkflowEventsApi, "/workflow/<string:task_id>/events")
|
||||||
@@ -4,8 +4,8 @@ import contextvars
|
|||||||
import logging
|
import logging
|
||||||
import threading
|
import threading
|
||||||
import uuid
|
import uuid
|
||||||
from collections.abc import Generator, Mapping
|
from collections.abc import Generator, Mapping, Sequence
|
||||||
from typing import TYPE_CHECKING, Any, Literal, Union, overload
|
from typing import TYPE_CHECKING, Any, Literal, TypeVar, Union, overload
|
||||||
|
|
||||||
from flask import Flask, current_app
|
from flask import Flask, current_app
|
||||||
from pydantic import ValidationError
|
from pydantic import ValidationError
|
||||||
@@ -29,21 +29,25 @@ 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 AdvancedChatAppGenerateEntity, InvokeFrom
|
from core.app.entities.app_invoke_entities import AdvancedChatAppGenerateEntity, InvokeFrom
|
||||||
from core.app.entities.task_entities import ChatbotAppBlockingResponse, ChatbotAppStreamResponse
|
from core.app.entities.task_entities import ChatbotAppBlockingResponse, ChatbotAppStreamResponse
|
||||||
|
from core.app.layers.pause_state_persist_layer import PauseStateLayerConfig, PauseStatePersistenceLayer
|
||||||
from core.helper.trace_id_helper import extract_external_trace_id_from_args
|
from core.helper.trace_id_helper import extract_external_trace_id_from_args
|
||||||
from core.model_runtime.errors.invoke import InvokeAuthorizationError
|
from core.model_runtime.errors.invoke import InvokeAuthorizationError
|
||||||
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 core.workflow.graph_engine.layers.base import GraphEngineLayer
|
||||||
from core.workflow.repositories.draft_variable_repository import (
|
from core.workflow.repositories.draft_variable_repository import (
|
||||||
DraftVariableSaverFactory,
|
DraftVariableSaverFactory,
|
||||||
)
|
)
|
||||||
from core.workflow.repositories.workflow_execution_repository import WorkflowExecutionRepository
|
from core.workflow.repositories.workflow_execution_repository import WorkflowExecutionRepository
|
||||||
from core.workflow.repositories.workflow_node_execution_repository import WorkflowNodeExecutionRepository
|
from core.workflow.repositories.workflow_node_execution_repository import WorkflowNodeExecutionRepository
|
||||||
|
from core.workflow.runtime import GraphRuntimeState
|
||||||
from core.workflow.variable_loader import DUMMY_VARIABLE_LOADER, VariableLoader
|
from core.workflow.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 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.enums import WorkflowRunTriggeredFrom
|
from models.enums import WorkflowRunTriggeredFrom
|
||||||
from services.conversation_service import ConversationService
|
from services.conversation_service import ConversationService
|
||||||
from services.workflow_draft_variable_service import (
|
from services.workflow_draft_variable_service import (
|
||||||
@@ -65,7 +69,9 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator):
|
|||||||
user: Union[Account, EndUser],
|
user: Union[Account, EndUser],
|
||||||
args: Mapping[str, Any],
|
args: Mapping[str, Any],
|
||||||
invoke_from: InvokeFrom,
|
invoke_from: InvokeFrom,
|
||||||
|
workflow_run_id: str,
|
||||||
streaming: Literal[False],
|
streaming: Literal[False],
|
||||||
|
pause_state_config: PauseStateLayerConfig | None = None,
|
||||||
) -> Mapping[str, Any]: ...
|
) -> Mapping[str, Any]: ...
|
||||||
|
|
||||||
@overload
|
@overload
|
||||||
@@ -74,9 +80,11 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator):
|
|||||||
app_model: App,
|
app_model: App,
|
||||||
workflow: Workflow,
|
workflow: Workflow,
|
||||||
user: Union[Account, EndUser],
|
user: Union[Account, EndUser],
|
||||||
args: Mapping,
|
args: Mapping[str, Any],
|
||||||
invoke_from: InvokeFrom,
|
invoke_from: InvokeFrom,
|
||||||
|
workflow_run_id: str,
|
||||||
streaming: Literal[True],
|
streaming: Literal[True],
|
||||||
|
pause_state_config: PauseStateLayerConfig | None = None,
|
||||||
) -> Generator[Mapping | str, None, None]: ...
|
) -> Generator[Mapping | str, None, None]: ...
|
||||||
|
|
||||||
@overload
|
@overload
|
||||||
@@ -85,9 +93,11 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator):
|
|||||||
app_model: App,
|
app_model: App,
|
||||||
workflow: Workflow,
|
workflow: Workflow,
|
||||||
user: Union[Account, EndUser],
|
user: Union[Account, EndUser],
|
||||||
args: Mapping,
|
args: Mapping[str, Any],
|
||||||
invoke_from: InvokeFrom,
|
invoke_from: InvokeFrom,
|
||||||
|
workflow_run_id: str,
|
||||||
streaming: bool,
|
streaming: bool,
|
||||||
|
pause_state_config: PauseStateLayerConfig | None = None,
|
||||||
) -> Mapping[str, Any] | Generator[str | Mapping, None, None]: ...
|
) -> Mapping[str, Any] | Generator[str | Mapping, None, None]: ...
|
||||||
|
|
||||||
def generate(
|
def generate(
|
||||||
@@ -95,9 +105,11 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator):
|
|||||||
app_model: App,
|
app_model: App,
|
||||||
workflow: Workflow,
|
workflow: Workflow,
|
||||||
user: Union[Account, EndUser],
|
user: Union[Account, EndUser],
|
||||||
args: Mapping,
|
args: Mapping[str, Any],
|
||||||
invoke_from: InvokeFrom,
|
invoke_from: InvokeFrom,
|
||||||
|
workflow_run_id: str,
|
||||||
streaming: bool = True,
|
streaming: bool = True,
|
||||||
|
pause_state_config: PauseStateLayerConfig | None = None,
|
||||||
) -> Mapping[str, Any] | Generator[str | Mapping, None, None]:
|
) -> Mapping[str, Any] | Generator[str | Mapping, None, None]:
|
||||||
"""
|
"""
|
||||||
Generate App response.
|
Generate App response.
|
||||||
@@ -161,7 +173,6 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator):
|
|||||||
# always enable retriever resource in debugger mode
|
# always enable retriever resource in debugger mode
|
||||||
app_config.additional_features.show_retrieve_source = True # type: ignore
|
app_config.additional_features.show_retrieve_source = True # type: ignore
|
||||||
|
|
||||||
workflow_run_id = str(uuid.uuid4())
|
|
||||||
# init application generate entity
|
# init application generate entity
|
||||||
application_generate_entity = AdvancedChatAppGenerateEntity(
|
application_generate_entity = AdvancedChatAppGenerateEntity(
|
||||||
task_id=str(uuid.uuid4()),
|
task_id=str(uuid.uuid4()),
|
||||||
@@ -179,7 +190,7 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator):
|
|||||||
invoke_from=invoke_from,
|
invoke_from=invoke_from,
|
||||||
extras=extras,
|
extras=extras,
|
||||||
trace_manager=trace_manager,
|
trace_manager=trace_manager,
|
||||||
workflow_run_id=workflow_run_id,
|
workflow_run_id=str(workflow_run_id),
|
||||||
)
|
)
|
||||||
contexts.plugin_tool_providers.set({})
|
contexts.plugin_tool_providers.set({})
|
||||||
contexts.plugin_tool_providers_lock.set(threading.Lock())
|
contexts.plugin_tool_providers_lock.set(threading.Lock())
|
||||||
@@ -216,6 +227,38 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator):
|
|||||||
workflow_node_execution_repository=workflow_node_execution_repository,
|
workflow_node_execution_repository=workflow_node_execution_repository,
|
||||||
conversation=conversation,
|
conversation=conversation,
|
||||||
stream=streaming,
|
stream=streaming,
|
||||||
|
pause_state_config=pause_state_config,
|
||||||
|
)
|
||||||
|
|
||||||
|
def resume(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
app_model: App,
|
||||||
|
workflow: Workflow,
|
||||||
|
user: Union[Account, EndUser],
|
||||||
|
conversation: Conversation,
|
||||||
|
message: Message,
|
||||||
|
application_generate_entity: AdvancedChatAppGenerateEntity,
|
||||||
|
workflow_execution_repository: WorkflowExecutionRepository,
|
||||||
|
workflow_node_execution_repository: WorkflowNodeExecutionRepository,
|
||||||
|
graph_runtime_state: GraphRuntimeState,
|
||||||
|
pause_state_config: PauseStateLayerConfig | None = None,
|
||||||
|
):
|
||||||
|
"""
|
||||||
|
Resume a paused advanced chat execution.
|
||||||
|
"""
|
||||||
|
return self._generate(
|
||||||
|
workflow=workflow,
|
||||||
|
user=user,
|
||||||
|
invoke_from=application_generate_entity.invoke_from,
|
||||||
|
application_generate_entity=application_generate_entity,
|
||||||
|
workflow_execution_repository=workflow_execution_repository,
|
||||||
|
workflow_node_execution_repository=workflow_node_execution_repository,
|
||||||
|
conversation=conversation,
|
||||||
|
message=message,
|
||||||
|
stream=application_generate_entity.stream,
|
||||||
|
pause_state_config=pause_state_config,
|
||||||
|
graph_runtime_state=graph_runtime_state,
|
||||||
)
|
)
|
||||||
|
|
||||||
def single_iteration_generate(
|
def single_iteration_generate(
|
||||||
@@ -396,8 +439,12 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator):
|
|||||||
workflow_execution_repository: WorkflowExecutionRepository,
|
workflow_execution_repository: WorkflowExecutionRepository,
|
||||||
workflow_node_execution_repository: WorkflowNodeExecutionRepository,
|
workflow_node_execution_repository: WorkflowNodeExecutionRepository,
|
||||||
conversation: Conversation | None = None,
|
conversation: Conversation | None = None,
|
||||||
|
message: Message | None = None,
|
||||||
stream: bool = True,
|
stream: bool = True,
|
||||||
variable_loader: VariableLoader = DUMMY_VARIABLE_LOADER,
|
variable_loader: VariableLoader = DUMMY_VARIABLE_LOADER,
|
||||||
|
pause_state_config: PauseStateLayerConfig | None = None,
|
||||||
|
graph_runtime_state: GraphRuntimeState | None = None,
|
||||||
|
graph_engine_layers: Sequence[GraphEngineLayer] = (),
|
||||||
) -> Mapping[str, Any] | Generator[str | Mapping[str, Any], Any, None]:
|
) -> Mapping[str, Any] | Generator[str | Mapping[str, Any], Any, None]:
|
||||||
"""
|
"""
|
||||||
Generate App response.
|
Generate App response.
|
||||||
@@ -411,12 +458,12 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator):
|
|||||||
:param conversation: conversation
|
:param conversation: conversation
|
||||||
:param stream: is stream
|
:param stream: is stream
|
||||||
"""
|
"""
|
||||||
is_first_conversation = False
|
is_first_conversation = conversation is None
|
||||||
if not conversation:
|
|
||||||
is_first_conversation = True
|
|
||||||
|
|
||||||
# init generate records
|
if conversation is not None and message is not None:
|
||||||
(conversation, message) = self._init_generate_records(application_generate_entity, conversation)
|
pass
|
||||||
|
else:
|
||||||
|
conversation, message = self._init_generate_records(application_generate_entity, conversation)
|
||||||
|
|
||||||
if is_first_conversation:
|
if is_first_conversation:
|
||||||
# update conversation features
|
# update conversation features
|
||||||
@@ -439,6 +486,16 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator):
|
|||||||
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
|
# new thread with request context and contextvars
|
||||||
context = contextvars.copy_context()
|
context = contextvars.copy_context()
|
||||||
|
|
||||||
@@ -454,14 +511,25 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator):
|
|||||||
"variable_loader": variable_loader,
|
"variable_loader": variable_loader,
|
||||||
"workflow_execution_repository": workflow_execution_repository,
|
"workflow_execution_repository": workflow_execution_repository,
|
||||||
"workflow_node_execution_repository": workflow_node_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.start()
|
||||||
|
|
||||||
# release database connection, because the following new thread operations may take a long time
|
# release database connection, because the following new thread operations may take a long time
|
||||||
db.session.refresh(workflow)
|
with Session(bind=db.engine, expire_on_commit=False) as session:
|
||||||
db.session.refresh(message)
|
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.refresh(user)
|
||||||
db.session.close()
|
db.session.close()
|
||||||
|
|
||||||
@@ -490,6 +558,8 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator):
|
|||||||
variable_loader: VariableLoader,
|
variable_loader: VariableLoader,
|
||||||
workflow_execution_repository: WorkflowExecutionRepository,
|
workflow_execution_repository: WorkflowExecutionRepository,
|
||||||
workflow_node_execution_repository: WorkflowNodeExecutionRepository,
|
workflow_node_execution_repository: WorkflowNodeExecutionRepository,
|
||||||
|
graph_engine_layers: Sequence[GraphEngineLayer] = (),
|
||||||
|
graph_runtime_state: GraphRuntimeState | None = None,
|
||||||
):
|
):
|
||||||
"""
|
"""
|
||||||
Generate worker in a new thread.
|
Generate worker in a new thread.
|
||||||
@@ -547,6 +617,8 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator):
|
|||||||
app=app,
|
app=app,
|
||||||
workflow_execution_repository=workflow_execution_repository,
|
workflow_execution_repository=workflow_execution_repository,
|
||||||
workflow_node_execution_repository=workflow_node_execution_repository,
|
workflow_node_execution_repository=workflow_node_execution_repository,
|
||||||
|
graph_engine_layers=graph_engine_layers,
|
||||||
|
graph_runtime_state=graph_runtime_state,
|
||||||
)
|
)
|
||||||
|
|
||||||
try:
|
try:
|
||||||
@@ -614,3 +686,13 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator):
|
|||||||
else:
|
else:
|
||||||
logger.exception("Failed to process generate task pipeline, conversation_id: %s", conversation.id)
|
logger.exception("Failed to process generate task pipeline, conversation_id: %s", conversation.id)
|
||||||
raise e
|
raise e
|
||||||
|
|
||||||
|
|
||||||
|
_T = TypeVar("_T", bound=Base)
|
||||||
|
|
||||||
|
|
||||||
|
def _refresh_model(session, model: _T) -> _T:
|
||||||
|
with Session(bind=db.engine, expire_on_commit=False) as session:
|
||||||
|
detach_model = session.get(type(model), model.id)
|
||||||
|
assert detach_model is not None
|
||||||
|
return detach_model
|
||||||
|
|||||||
@@ -66,6 +66,7 @@ class AdvancedChatAppRunner(WorkflowBasedAppRunner):
|
|||||||
workflow_execution_repository: WorkflowExecutionRepository,
|
workflow_execution_repository: WorkflowExecutionRepository,
|
||||||
workflow_node_execution_repository: WorkflowNodeExecutionRepository,
|
workflow_node_execution_repository: WorkflowNodeExecutionRepository,
|
||||||
graph_engine_layers: Sequence[GraphEngineLayer] = (),
|
graph_engine_layers: Sequence[GraphEngineLayer] = (),
|
||||||
|
graph_runtime_state: GraphRuntimeState | None = None,
|
||||||
):
|
):
|
||||||
super().__init__(
|
super().__init__(
|
||||||
queue_manager=queue_manager,
|
queue_manager=queue_manager,
|
||||||
@@ -82,6 +83,7 @@ class AdvancedChatAppRunner(WorkflowBasedAppRunner):
|
|||||||
self._app = app
|
self._app = app
|
||||||
self._workflow_execution_repository = workflow_execution_repository
|
self._workflow_execution_repository = workflow_execution_repository
|
||||||
self._workflow_node_execution_repository = workflow_node_execution_repository
|
self._workflow_node_execution_repository = workflow_node_execution_repository
|
||||||
|
self._resume_graph_runtime_state = graph_runtime_state
|
||||||
|
|
||||||
@trace_span(WorkflowAppRunnerHandler)
|
@trace_span(WorkflowAppRunnerHandler)
|
||||||
def run(self):
|
def run(self):
|
||||||
@@ -110,7 +112,21 @@ class AdvancedChatAppRunner(WorkflowBasedAppRunner):
|
|||||||
invoke_from = InvokeFrom.DEBUGGER
|
invoke_from = InvokeFrom.DEBUGGER
|
||||||
user_from = self._resolve_user_from(invoke_from)
|
user_from = self._resolve_user_from(invoke_from)
|
||||||
|
|
||||||
if self.application_generate_entity.single_iteration_run or self.application_generate_entity.single_loop_run:
|
resume_state = self._resume_graph_runtime_state
|
||||||
|
|
||||||
|
if resume_state is not None:
|
||||||
|
graph_runtime_state = resume_state
|
||||||
|
variable_pool = graph_runtime_state.variable_pool
|
||||||
|
graph = self._init_graph(
|
||||||
|
graph_config=self._workflow.graph_dict,
|
||||||
|
graph_runtime_state=graph_runtime_state,
|
||||||
|
workflow_id=self._workflow.id,
|
||||||
|
tenant_id=self._workflow.tenant_id,
|
||||||
|
user_id=self.application_generate_entity.user_id,
|
||||||
|
invoke_from=invoke_from,
|
||||||
|
user_from=user_from,
|
||||||
|
)
|
||||||
|
elif self.application_generate_entity.single_iteration_run or self.application_generate_entity.single_loop_run:
|
||||||
# Handle single iteration or single loop run
|
# Handle single iteration or single loop run
|
||||||
graph, variable_pool, graph_runtime_state = self._prepare_single_node_execution(
|
graph, variable_pool, graph_runtime_state = self._prepare_single_node_execution(
|
||||||
workflow=self._workflow,
|
workflow=self._workflow,
|
||||||
|
|||||||
@@ -24,6 +24,8 @@ from core.app.entities.queue_entities import (
|
|||||||
QueueAgentLogEvent,
|
QueueAgentLogEvent,
|
||||||
QueueAnnotationReplyEvent,
|
QueueAnnotationReplyEvent,
|
||||||
QueueErrorEvent,
|
QueueErrorEvent,
|
||||||
|
QueueHumanInputFormFilledEvent,
|
||||||
|
QueueHumanInputFormTimeoutEvent,
|
||||||
QueueIterationCompletedEvent,
|
QueueIterationCompletedEvent,
|
||||||
QueueIterationNextEvent,
|
QueueIterationNextEvent,
|
||||||
QueueIterationStartEvent,
|
QueueIterationStartEvent,
|
||||||
@@ -42,6 +44,7 @@ from core.app.entities.queue_entities import (
|
|||||||
QueueTextChunkEvent,
|
QueueTextChunkEvent,
|
||||||
QueueWorkflowFailedEvent,
|
QueueWorkflowFailedEvent,
|
||||||
QueueWorkflowPartialSuccessEvent,
|
QueueWorkflowPartialSuccessEvent,
|
||||||
|
QueueWorkflowPausedEvent,
|
||||||
QueueWorkflowStartedEvent,
|
QueueWorkflowStartedEvent,
|
||||||
QueueWorkflowSucceededEvent,
|
QueueWorkflowSucceededEvent,
|
||||||
WorkflowQueueMessage,
|
WorkflowQueueMessage,
|
||||||
@@ -63,6 +66,8 @@ from core.base.tts import AppGeneratorTTSPublisher, AudioTrunk
|
|||||||
from core.model_runtime.entities.llm_entities import LLMUsage
|
from core.model_runtime.entities.llm_entities import LLMUsage
|
||||||
from core.model_runtime.utils.encoders import jsonable_encoder
|
from core.model_runtime.utils.encoders import jsonable_encoder
|
||||||
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.workflow.entities.pause_reason import HumanInputRequired
|
||||||
from core.workflow.enums import WorkflowExecutionStatus
|
from core.workflow.enums import WorkflowExecutionStatus
|
||||||
from core.workflow.nodes import NodeType
|
from core.workflow.nodes import NodeType
|
||||||
from core.workflow.repositories.draft_variable_repository import DraftVariableSaverFactory
|
from core.workflow.repositories.draft_variable_repository import DraftVariableSaverFactory
|
||||||
@@ -71,7 +76,8 @@ from core.workflow.system_variable import SystemVariable
|
|||||||
from extensions.ext_database import db
|
from extensions.ext_database import db
|
||||||
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
|
from models.enums import CreatorUserRole, MessageStatus
|
||||||
|
from models.execution_extra_content import HumanInputContent
|
||||||
from models.workflow import Workflow
|
from models.workflow import Workflow
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
@@ -128,6 +134,7 @@ class AdvancedChatAppGenerateTaskPipeline(GraphRuntimeStateSupport):
|
|||||||
)
|
)
|
||||||
|
|
||||||
self._task_state = WorkflowTaskState()
|
self._task_state = WorkflowTaskState()
|
||||||
|
self._seed_task_state_from_message(message)
|
||||||
self._message_cycle_manager = MessageCycleManager(
|
self._message_cycle_manager = MessageCycleManager(
|
||||||
application_generate_entity=application_generate_entity, task_state=self._task_state
|
application_generate_entity=application_generate_entity, task_state=self._task_state
|
||||||
)
|
)
|
||||||
@@ -135,6 +142,7 @@ class AdvancedChatAppGenerateTaskPipeline(GraphRuntimeStateSupport):
|
|||||||
self._application_generate_entity = application_generate_entity
|
self._application_generate_entity = application_generate_entity
|
||||||
self._workflow_id = workflow.id
|
self._workflow_id = workflow.id
|
||||||
self._workflow_features_dict = workflow.features_dict
|
self._workflow_features_dict = workflow.features_dict
|
||||||
|
self._workflow_tenant_id = workflow.tenant_id
|
||||||
self._conversation_id = conversation.id
|
self._conversation_id = conversation.id
|
||||||
self._conversation_mode = conversation.mode
|
self._conversation_mode = conversation.mode
|
||||||
self._message_id = message.id
|
self._message_id = message.id
|
||||||
@@ -144,8 +152,13 @@ class AdvancedChatAppGenerateTaskPipeline(GraphRuntimeStateSupport):
|
|||||||
self._workflow_run_id: str = ""
|
self._workflow_run_id: str = ""
|
||||||
self._draft_var_saver_factory = draft_var_saver_factory
|
self._draft_var_saver_factory = draft_var_saver_factory
|
||||||
self._graph_runtime_state: GraphRuntimeState | None = None
|
self._graph_runtime_state: GraphRuntimeState | None = None
|
||||||
|
self._message_saved_on_pause = False
|
||||||
self._seed_graph_runtime_state_from_queue_manager()
|
self._seed_graph_runtime_state_from_queue_manager()
|
||||||
|
|
||||||
|
def _seed_task_state_from_message(self, message: Message) -> None:
|
||||||
|
if message.status == MessageStatus.PAUSED and message.answer:
|
||||||
|
self._task_state.answer = message.answer
|
||||||
|
|
||||||
def process(self) -> Union[ChatbotAppBlockingResponse, Generator[ChatbotAppStreamResponse, None, None]]:
|
def process(self) -> Union[ChatbotAppBlockingResponse, Generator[ChatbotAppStreamResponse, None, None]]:
|
||||||
"""
|
"""
|
||||||
Process generate task pipeline.
|
Process generate task pipeline.
|
||||||
@@ -308,6 +321,7 @@ class AdvancedChatAppGenerateTaskPipeline(GraphRuntimeStateSupport):
|
|||||||
task_id=self._application_generate_entity.task_id,
|
task_id=self._application_generate_entity.task_id,
|
||||||
workflow_run_id=run_id,
|
workflow_run_id=run_id,
|
||||||
workflow_id=self._workflow_id,
|
workflow_id=self._workflow_id,
|
||||||
|
reason=event.reason,
|
||||||
)
|
)
|
||||||
|
|
||||||
yield workflow_start_resp
|
yield workflow_start_resp
|
||||||
@@ -525,6 +539,35 @@ class AdvancedChatAppGenerateTaskPipeline(GraphRuntimeStateSupport):
|
|||||||
)
|
)
|
||||||
|
|
||||||
yield workflow_finish_resp
|
yield workflow_finish_resp
|
||||||
|
|
||||||
|
def _handle_workflow_paused_event(
|
||||||
|
self,
|
||||||
|
event: QueueWorkflowPausedEvent,
|
||||||
|
**kwargs,
|
||||||
|
) -> Generator[StreamResponse, None, None]:
|
||||||
|
"""Handle workflow paused events."""
|
||||||
|
validated_state = self._ensure_graph_runtime_initialized()
|
||||||
|
responses = self._workflow_response_converter.workflow_pause_to_stream_response(
|
||||||
|
event=event,
|
||||||
|
task_id=self._application_generate_entity.task_id,
|
||||||
|
graph_runtime_state=validated_state,
|
||||||
|
)
|
||||||
|
for reason in event.reasons:
|
||||||
|
if isinstance(reason, HumanInputRequired):
|
||||||
|
self._persist_human_input_extra_content(form_id=reason.form_id, node_id=reason.node_id)
|
||||||
|
yield from responses
|
||||||
|
resolved_state: GraphRuntimeState | None = None
|
||||||
|
try:
|
||||||
|
resolved_state = self._ensure_graph_runtime_initialized()
|
||||||
|
except ValueError:
|
||||||
|
resolved_state = None
|
||||||
|
|
||||||
|
with self._database_session() as session:
|
||||||
|
self._save_message(session=session, graph_runtime_state=resolved_state)
|
||||||
|
message = self._get_message(session=session)
|
||||||
|
if message is not None:
|
||||||
|
message.status = MessageStatus.PAUSED
|
||||||
|
self._message_saved_on_pause = True
|
||||||
self._base_task_pipeline.queue_manager.publish(QueueAdvancedChatMessageEndEvent(), PublishFrom.TASK_PIPELINE)
|
self._base_task_pipeline.queue_manager.publish(QueueAdvancedChatMessageEndEvent(), PublishFrom.TASK_PIPELINE)
|
||||||
|
|
||||||
def _handle_workflow_failed_event(
|
def _handle_workflow_failed_event(
|
||||||
@@ -614,9 +657,10 @@ class AdvancedChatAppGenerateTaskPipeline(GraphRuntimeStateSupport):
|
|||||||
reason=QueueMessageReplaceEvent.MessageReplaceReason.OUTPUT_MODERATION,
|
reason=QueueMessageReplaceEvent.MessageReplaceReason.OUTPUT_MODERATION,
|
||||||
)
|
)
|
||||||
|
|
||||||
# Save message
|
# Save message unless it has already been persisted on pause.
|
||||||
with self._database_session() as session:
|
if not self._message_saved_on_pause:
|
||||||
self._save_message(session=session, graph_runtime_state=resolved_state)
|
with self._database_session() as session:
|
||||||
|
self._save_message(session=session, graph_runtime_state=resolved_state)
|
||||||
|
|
||||||
yield self._message_end_to_stream_response()
|
yield self._message_end_to_stream_response()
|
||||||
|
|
||||||
@@ -642,6 +686,65 @@ class AdvancedChatAppGenerateTaskPipeline(GraphRuntimeStateSupport):
|
|||||||
"""Handle message replace events."""
|
"""Handle message replace events."""
|
||||||
yield self._message_cycle_manager.message_replace_to_stream_response(answer=event.text, reason=event.reason)
|
yield self._message_cycle_manager.message_replace_to_stream_response(answer=event.text, reason=event.reason)
|
||||||
|
|
||||||
|
def _handle_human_input_form_filled_event(
|
||||||
|
self, event: QueueHumanInputFormFilledEvent, **kwargs
|
||||||
|
) -> Generator[StreamResponse, None, None]:
|
||||||
|
"""Handle human input form filled events."""
|
||||||
|
self._persist_human_input_extra_content(node_id=event.node_id)
|
||||||
|
yield self._workflow_response_converter.human_input_form_filled_to_stream_response(
|
||||||
|
event=event, task_id=self._application_generate_entity.task_id
|
||||||
|
)
|
||||||
|
|
||||||
|
def _handle_human_input_form_timeout_event(
|
||||||
|
self, event: QueueHumanInputFormTimeoutEvent, **kwargs
|
||||||
|
) -> Generator[StreamResponse, None, None]:
|
||||||
|
"""Handle human input form timeout events."""
|
||||||
|
yield self._workflow_response_converter.human_input_form_timeout_to_stream_response(
|
||||||
|
event=event, task_id=self._application_generate_entity.task_id
|
||||||
|
)
|
||||||
|
|
||||||
|
def _persist_human_input_extra_content(self, *, node_id: str | None = None, form_id: str | None = None) -> None:
|
||||||
|
if not self._workflow_run_id or not self._message_id:
|
||||||
|
return
|
||||||
|
|
||||||
|
if form_id is None:
|
||||||
|
if node_id is None:
|
||||||
|
return
|
||||||
|
form_id = self._load_human_input_form_id(node_id=node_id)
|
||||||
|
if form_id is None:
|
||||||
|
logger.warning(
|
||||||
|
"HumanInput form not found for workflow run %s node %s",
|
||||||
|
self._workflow_run_id,
|
||||||
|
node_id,
|
||||||
|
)
|
||||||
|
return
|
||||||
|
|
||||||
|
with self._database_session() as session:
|
||||||
|
exists_stmt = select(HumanInputContent).where(
|
||||||
|
HumanInputContent.workflow_run_id == self._workflow_run_id,
|
||||||
|
HumanInputContent.message_id == self._message_id,
|
||||||
|
HumanInputContent.form_id == form_id,
|
||||||
|
)
|
||||||
|
if session.scalar(exists_stmt) is not None:
|
||||||
|
return
|
||||||
|
|
||||||
|
content = HumanInputContent(
|
||||||
|
workflow_run_id=self._workflow_run_id,
|
||||||
|
message_id=self._message_id,
|
||||||
|
form_id=form_id,
|
||||||
|
)
|
||||||
|
session.add(content)
|
||||||
|
|
||||||
|
def _load_human_input_form_id(self, *, node_id: str) -> str | None:
|
||||||
|
form_repository = HumanInputFormRepositoryImpl(
|
||||||
|
session_factory=db.engine,
|
||||||
|
tenant_id=self._workflow_tenant_id,
|
||||||
|
)
|
||||||
|
form = form_repository.get_form(self._workflow_run_id, node_id)
|
||||||
|
if form is None:
|
||||||
|
return None
|
||||||
|
return form.id
|
||||||
|
|
||||||
def _handle_agent_log_event(self, event: QueueAgentLogEvent, **kwargs) -> Generator[StreamResponse, None, None]:
|
def _handle_agent_log_event(self, event: QueueAgentLogEvent, **kwargs) -> Generator[StreamResponse, None, None]:
|
||||||
"""Handle agent log events."""
|
"""Handle agent log events."""
|
||||||
yield self._workflow_response_converter.handle_agent_log(
|
yield self._workflow_response_converter.handle_agent_log(
|
||||||
@@ -659,6 +762,7 @@ class AdvancedChatAppGenerateTaskPipeline(GraphRuntimeStateSupport):
|
|||||||
QueueWorkflowStartedEvent: self._handle_workflow_started_event,
|
QueueWorkflowStartedEvent: self._handle_workflow_started_event,
|
||||||
QueueWorkflowSucceededEvent: self._handle_workflow_succeeded_event,
|
QueueWorkflowSucceededEvent: self._handle_workflow_succeeded_event,
|
||||||
QueueWorkflowPartialSuccessEvent: self._handle_workflow_partial_success_event,
|
QueueWorkflowPartialSuccessEvent: self._handle_workflow_partial_success_event,
|
||||||
|
QueueWorkflowPausedEvent: self._handle_workflow_paused_event,
|
||||||
QueueWorkflowFailedEvent: self._handle_workflow_failed_event,
|
QueueWorkflowFailedEvent: self._handle_workflow_failed_event,
|
||||||
# Node events
|
# Node events
|
||||||
QueueNodeRetryEvent: self._handle_node_retry_event,
|
QueueNodeRetryEvent: self._handle_node_retry_event,
|
||||||
@@ -680,6 +784,8 @@ class AdvancedChatAppGenerateTaskPipeline(GraphRuntimeStateSupport):
|
|||||||
QueueMessageReplaceEvent: self._handle_message_replace_event,
|
QueueMessageReplaceEvent: self._handle_message_replace_event,
|
||||||
QueueAdvancedChatMessageEndEvent: self._handle_advanced_chat_message_end_event,
|
QueueAdvancedChatMessageEndEvent: self._handle_advanced_chat_message_end_event,
|
||||||
QueueAgentLogEvent: self._handle_agent_log_event,
|
QueueAgentLogEvent: self._handle_agent_log_event,
|
||||||
|
QueueHumanInputFormFilledEvent: self._handle_human_input_form_filled_event,
|
||||||
|
QueueHumanInputFormTimeoutEvent: self._handle_human_input_form_timeout_event,
|
||||||
}
|
}
|
||||||
|
|
||||||
def _dispatch_event(
|
def _dispatch_event(
|
||||||
@@ -747,6 +853,9 @@ class AdvancedChatAppGenerateTaskPipeline(GraphRuntimeStateSupport):
|
|||||||
case QueueWorkflowFailedEvent():
|
case QueueWorkflowFailedEvent():
|
||||||
yield from self._handle_workflow_failed_event(event, trace_manager=trace_manager)
|
yield from self._handle_workflow_failed_event(event, trace_manager=trace_manager)
|
||||||
break
|
break
|
||||||
|
case QueueWorkflowPausedEvent():
|
||||||
|
yield from self._handle_workflow_paused_event(event)
|
||||||
|
break
|
||||||
|
|
||||||
case QueueStopEvent():
|
case QueueStopEvent():
|
||||||
yield from self._handle_stop_event(event, graph_runtime_state=None, trace_manager=trace_manager)
|
yield from self._handle_stop_event(event, graph_runtime_state=None, trace_manager=trace_manager)
|
||||||
@@ -772,6 +881,11 @@ class AdvancedChatAppGenerateTaskPipeline(GraphRuntimeStateSupport):
|
|||||||
|
|
||||||
def _save_message(self, *, session: Session, graph_runtime_state: GraphRuntimeState | None = None):
|
def _save_message(self, *, session: Session, graph_runtime_state: GraphRuntimeState | None = None):
|
||||||
message = self._get_message(session=session)
|
message = self._get_message(session=session)
|
||||||
|
if message is None:
|
||||||
|
return
|
||||||
|
|
||||||
|
if message.status == MessageStatus.PAUSED:
|
||||||
|
message.status = MessageStatus.NORMAL
|
||||||
|
|
||||||
# If there are assistant files, remove markdown image links from answer
|
# If there are assistant files, remove markdown image links from answer
|
||||||
answer_text = self._task_state.answer
|
answer_text = self._task_state.answer
|
||||||
|
|||||||
@@ -5,9 +5,14 @@ from dataclasses import dataclass
|
|||||||
from datetime import datetime
|
from datetime import datetime
|
||||||
from typing import Any, NewType, Union
|
from typing import Any, NewType, Union
|
||||||
|
|
||||||
|
from sqlalchemy import select
|
||||||
|
from sqlalchemy.orm import Session
|
||||||
|
|
||||||
from core.app.entities.app_invoke_entities import AdvancedChatAppGenerateEntity, InvokeFrom, WorkflowAppGenerateEntity
|
from core.app.entities.app_invoke_entities import AdvancedChatAppGenerateEntity, InvokeFrom, WorkflowAppGenerateEntity
|
||||||
from core.app.entities.queue_entities import (
|
from core.app.entities.queue_entities import (
|
||||||
QueueAgentLogEvent,
|
QueueAgentLogEvent,
|
||||||
|
QueueHumanInputFormFilledEvent,
|
||||||
|
QueueHumanInputFormTimeoutEvent,
|
||||||
QueueIterationCompletedEvent,
|
QueueIterationCompletedEvent,
|
||||||
QueueIterationNextEvent,
|
QueueIterationNextEvent,
|
||||||
QueueIterationStartEvent,
|
QueueIterationStartEvent,
|
||||||
@@ -19,9 +24,13 @@ from core.app.entities.queue_entities import (
|
|||||||
QueueNodeRetryEvent,
|
QueueNodeRetryEvent,
|
||||||
QueueNodeStartedEvent,
|
QueueNodeStartedEvent,
|
||||||
QueueNodeSucceededEvent,
|
QueueNodeSucceededEvent,
|
||||||
|
QueueWorkflowPausedEvent,
|
||||||
)
|
)
|
||||||
from core.app.entities.task_entities import (
|
from core.app.entities.task_entities import (
|
||||||
AgentLogStreamResponse,
|
AgentLogStreamResponse,
|
||||||
|
HumanInputFormFilledResponse,
|
||||||
|
HumanInputFormTimeoutResponse,
|
||||||
|
HumanInputRequiredResponse,
|
||||||
IterationNodeCompletedStreamResponse,
|
IterationNodeCompletedStreamResponse,
|
||||||
IterationNodeNextStreamResponse,
|
IterationNodeNextStreamResponse,
|
||||||
IterationNodeStartStreamResponse,
|
IterationNodeStartStreamResponse,
|
||||||
@@ -31,7 +40,9 @@ from core.app.entities.task_entities import (
|
|||||||
NodeFinishStreamResponse,
|
NodeFinishStreamResponse,
|
||||||
NodeRetryStreamResponse,
|
NodeRetryStreamResponse,
|
||||||
NodeStartStreamResponse,
|
NodeStartStreamResponse,
|
||||||
|
StreamResponse,
|
||||||
WorkflowFinishStreamResponse,
|
WorkflowFinishStreamResponse,
|
||||||
|
WorkflowPauseStreamResponse,
|
||||||
WorkflowStartStreamResponse,
|
WorkflowStartStreamResponse,
|
||||||
)
|
)
|
||||||
from core.file import FILE_MODEL_IDENTITY, File
|
from core.file import FILE_MODEL_IDENTITY, File
|
||||||
@@ -40,6 +51,8 @@ from core.tools.entities.tool_entities import ToolProviderType
|
|||||||
from core.tools.tool_manager import ToolManager
|
from core.tools.tool_manager import ToolManager
|
||||||
from core.trigger.trigger_manager import TriggerManager
|
from core.trigger.trigger_manager import TriggerManager
|
||||||
from core.variables.segments import ArrayFileSegment, FileSegment, Segment
|
from core.variables.segments import ArrayFileSegment, FileSegment, Segment
|
||||||
|
from core.workflow.entities.pause_reason import HumanInputRequired
|
||||||
|
from core.workflow.entities.workflow_start_reason import WorkflowStartReason
|
||||||
from core.workflow.enums import (
|
from core.workflow.enums import (
|
||||||
NodeType,
|
NodeType,
|
||||||
SystemVariableKey,
|
SystemVariableKey,
|
||||||
@@ -51,8 +64,11 @@ from core.workflow.runtime import GraphRuntimeState
|
|||||||
from core.workflow.system_variable import SystemVariable
|
from core.workflow.system_variable import SystemVariable
|
||||||
from core.workflow.workflow_entry import WorkflowEntry
|
from core.workflow.workflow_entry import WorkflowEntry
|
||||||
from core.workflow.workflow_type_encoder import WorkflowRuntimeTypeConverter
|
from core.workflow.workflow_type_encoder import WorkflowRuntimeTypeConverter
|
||||||
|
from extensions.ext_database import db
|
||||||
from libs.datetime_utils import naive_utc_now
|
from libs.datetime_utils import naive_utc_now
|
||||||
from models import Account, EndUser
|
from models import Account, EndUser
|
||||||
|
from models.human_input import HumanInputForm
|
||||||
|
from models.workflow import WorkflowRun
|
||||||
from services.variable_truncator import BaseTruncator, DummyVariableTruncator, VariableTruncator
|
from services.variable_truncator import BaseTruncator, DummyVariableTruncator, VariableTruncator
|
||||||
|
|
||||||
NodeExecutionId = NewType("NodeExecutionId", str)
|
NodeExecutionId = NewType("NodeExecutionId", str)
|
||||||
@@ -191,6 +207,7 @@ class WorkflowResponseConverter:
|
|||||||
task_id: str,
|
task_id: str,
|
||||||
workflow_run_id: str,
|
workflow_run_id: str,
|
||||||
workflow_id: str,
|
workflow_id: str,
|
||||||
|
reason: WorkflowStartReason,
|
||||||
) -> WorkflowStartStreamResponse:
|
) -> WorkflowStartStreamResponse:
|
||||||
run_id = self._ensure_workflow_run_id(workflow_run_id)
|
run_id = self._ensure_workflow_run_id(workflow_run_id)
|
||||||
started_at = naive_utc_now()
|
started_at = naive_utc_now()
|
||||||
@@ -204,6 +221,7 @@ class WorkflowResponseConverter:
|
|||||||
workflow_id=workflow_id,
|
workflow_id=workflow_id,
|
||||||
inputs=self._workflow_inputs,
|
inputs=self._workflow_inputs,
|
||||||
created_at=int(started_at.timestamp()),
|
created_at=int(started_at.timestamp()),
|
||||||
|
reason=reason,
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -264,6 +282,160 @@ class WorkflowResponseConverter:
|
|||||||
),
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def workflow_pause_to_stream_response(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
event: QueueWorkflowPausedEvent,
|
||||||
|
task_id: str,
|
||||||
|
graph_runtime_state: GraphRuntimeState,
|
||||||
|
) -> list[StreamResponse]:
|
||||||
|
run_id = self._ensure_workflow_run_id()
|
||||||
|
started_at = self._workflow_started_at
|
||||||
|
if started_at is None:
|
||||||
|
raise ValueError(
|
||||||
|
"workflow_pause_to_stream_response called before workflow_start_to_stream_response",
|
||||||
|
)
|
||||||
|
paused_at = naive_utc_now()
|
||||||
|
elapsed_time = (paused_at - started_at).total_seconds()
|
||||||
|
encoded_outputs = self._encode_outputs(event.outputs) or {}
|
||||||
|
if self._application_generate_entity.invoke_from == InvokeFrom.SERVICE_API:
|
||||||
|
encoded_outputs = {}
|
||||||
|
pause_reasons = [reason.model_dump(mode="json") for reason in event.reasons]
|
||||||
|
human_input_form_ids = [reason.form_id for reason in event.reasons if isinstance(reason, HumanInputRequired)]
|
||||||
|
expiration_times_by_form_id: dict[str, datetime] = {}
|
||||||
|
if human_input_form_ids:
|
||||||
|
stmt = select(HumanInputForm.id, HumanInputForm.expiration_time).where(
|
||||||
|
HumanInputForm.id.in_(human_input_form_ids)
|
||||||
|
)
|
||||||
|
with Session(bind=db.engine) as session:
|
||||||
|
for form_id, expiration_time in session.execute(stmt):
|
||||||
|
expiration_times_by_form_id[str(form_id)] = expiration_time
|
||||||
|
|
||||||
|
responses: list[StreamResponse] = []
|
||||||
|
|
||||||
|
for reason in event.reasons:
|
||||||
|
if isinstance(reason, HumanInputRequired):
|
||||||
|
expiration_time = expiration_times_by_form_id.get(reason.form_id)
|
||||||
|
if expiration_time is None:
|
||||||
|
raise ValueError(f"HumanInputForm not found for pause reason, form_id={reason.form_id}")
|
||||||
|
responses.append(
|
||||||
|
HumanInputRequiredResponse(
|
||||||
|
task_id=task_id,
|
||||||
|
workflow_run_id=run_id,
|
||||||
|
data=HumanInputRequiredResponse.Data(
|
||||||
|
form_id=reason.form_id,
|
||||||
|
node_id=reason.node_id,
|
||||||
|
node_title=reason.node_title,
|
||||||
|
form_content=reason.form_content,
|
||||||
|
inputs=reason.inputs,
|
||||||
|
actions=reason.actions,
|
||||||
|
display_in_ui=reason.display_in_ui,
|
||||||
|
form_token=reason.form_token,
|
||||||
|
resolved_default_values=reason.resolved_default_values,
|
||||||
|
expiration_time=int(expiration_time.timestamp()),
|
||||||
|
),
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
responses.append(
|
||||||
|
WorkflowPauseStreamResponse(
|
||||||
|
task_id=task_id,
|
||||||
|
workflow_run_id=run_id,
|
||||||
|
data=WorkflowPauseStreamResponse.Data(
|
||||||
|
workflow_run_id=run_id,
|
||||||
|
paused_nodes=list(event.paused_nodes),
|
||||||
|
outputs=encoded_outputs,
|
||||||
|
reasons=pause_reasons,
|
||||||
|
status=WorkflowExecutionStatus.PAUSED,
|
||||||
|
created_at=int(started_at.timestamp()),
|
||||||
|
elapsed_time=elapsed_time,
|
||||||
|
total_tokens=graph_runtime_state.total_tokens,
|
||||||
|
total_steps=graph_runtime_state.node_run_steps,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
return responses
|
||||||
|
|
||||||
|
def human_input_form_filled_to_stream_response(
|
||||||
|
self, *, event: QueueHumanInputFormFilledEvent, task_id: str
|
||||||
|
) -> HumanInputFormFilledResponse:
|
||||||
|
run_id = self._ensure_workflow_run_id()
|
||||||
|
return HumanInputFormFilledResponse(
|
||||||
|
task_id=task_id,
|
||||||
|
workflow_run_id=run_id,
|
||||||
|
data=HumanInputFormFilledResponse.Data(
|
||||||
|
node_id=event.node_id,
|
||||||
|
node_title=event.node_title,
|
||||||
|
rendered_content=event.rendered_content,
|
||||||
|
action_id=event.action_id,
|
||||||
|
action_text=event.action_text,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
def human_input_form_timeout_to_stream_response(
|
||||||
|
self, *, event: QueueHumanInputFormTimeoutEvent, task_id: str
|
||||||
|
) -> HumanInputFormTimeoutResponse:
|
||||||
|
run_id = self._ensure_workflow_run_id()
|
||||||
|
return HumanInputFormTimeoutResponse(
|
||||||
|
task_id=task_id,
|
||||||
|
workflow_run_id=run_id,
|
||||||
|
data=HumanInputFormTimeoutResponse.Data(
|
||||||
|
node_id=event.node_id,
|
||||||
|
node_title=event.node_title,
|
||||||
|
expiration_time=int(event.expiration_time.timestamp()),
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def workflow_run_result_to_finish_response(
|
||||||
|
cls,
|
||||||
|
*,
|
||||||
|
task_id: str,
|
||||||
|
workflow_run: WorkflowRun,
|
||||||
|
creator_user: Account | EndUser,
|
||||||
|
) -> WorkflowFinishStreamResponse:
|
||||||
|
run_id = workflow_run.id
|
||||||
|
elapsed_time = workflow_run.elapsed_time
|
||||||
|
|
||||||
|
encoded_outputs = workflow_run.outputs_dict
|
||||||
|
finished_at = workflow_run.finished_at
|
||||||
|
assert finished_at is not None
|
||||||
|
|
||||||
|
created_by: Mapping[str, object]
|
||||||
|
user = creator_user
|
||||||
|
if isinstance(user, Account):
|
||||||
|
created_by = {
|
||||||
|
"id": user.id,
|
||||||
|
"name": user.name,
|
||||||
|
"email": user.email,
|
||||||
|
}
|
||||||
|
else:
|
||||||
|
created_by = {
|
||||||
|
"id": user.id,
|
||||||
|
"user": user.session_id,
|
||||||
|
}
|
||||||
|
|
||||||
|
return WorkflowFinishStreamResponse(
|
||||||
|
task_id=task_id,
|
||||||
|
workflow_run_id=run_id,
|
||||||
|
data=WorkflowFinishStreamResponse.Data(
|
||||||
|
id=run_id,
|
||||||
|
workflow_id=workflow_run.workflow_id,
|
||||||
|
status=workflow_run.status,
|
||||||
|
outputs=encoded_outputs,
|
||||||
|
error=workflow_run.error,
|
||||||
|
elapsed_time=elapsed_time,
|
||||||
|
total_tokens=workflow_run.total_tokens,
|
||||||
|
total_steps=workflow_run.total_steps,
|
||||||
|
created_by=created_by,
|
||||||
|
created_at=int(workflow_run.created_at.timestamp()),
|
||||||
|
finished_at=int(finished_at.timestamp()),
|
||||||
|
files=cls.fetch_files_from_node_outputs(encoded_outputs),
|
||||||
|
exceptions_count=workflow_run.exceptions_count,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
def workflow_node_start_to_stream_response(
|
def workflow_node_start_to_stream_response(
|
||||||
self,
|
self,
|
||||||
*,
|
*,
|
||||||
@@ -592,7 +764,8 @@ class WorkflowResponseConverter:
|
|||||||
),
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
def fetch_files_from_node_outputs(self, outputs_dict: Mapping[str, Any] | None) -> Sequence[Mapping[str, Any]]:
|
@classmethod
|
||||||
|
def fetch_files_from_node_outputs(cls, outputs_dict: Mapping[str, Any] | None) -> Sequence[Mapping[str, Any]]:
|
||||||
"""
|
"""
|
||||||
Fetch files from node outputs
|
Fetch files from node outputs
|
||||||
:param outputs_dict: node outputs dict
|
:param outputs_dict: node outputs dict
|
||||||
@@ -601,7 +774,7 @@ class WorkflowResponseConverter:
|
|||||||
if not outputs_dict:
|
if not outputs_dict:
|
||||||
return []
|
return []
|
||||||
|
|
||||||
files = [self._fetch_files_from_variable_value(output_value) for output_value in outputs_dict.values()]
|
files = [cls._fetch_files_from_variable_value(output_value) for output_value in outputs_dict.values()]
|
||||||
# Remove None
|
# Remove None
|
||||||
files = [file for file in files if file]
|
files = [file for file in files if file]
|
||||||
# Flatten list
|
# Flatten list
|
||||||
|
|||||||
@@ -1,6 +1,6 @@
|
|||||||
import json
|
import json
|
||||||
import logging
|
import logging
|
||||||
from collections.abc import Generator
|
from collections.abc import Callable, Generator, Mapping
|
||||||
from typing import Union, cast
|
from typing import Union, cast
|
||||||
|
|
||||||
from sqlalchemy import select
|
from sqlalchemy import select
|
||||||
@@ -10,12 +10,14 @@ from core.app.app_config.entities import EasyUIBasedAppConfig, EasyUIBasedAppMod
|
|||||||
from core.app.apps.base_app_generator import BaseAppGenerator
|
from core.app.apps.base_app_generator import BaseAppGenerator
|
||||||
from core.app.apps.base_app_queue_manager import AppQueueManager
|
from core.app.apps.base_app_queue_manager import AppQueueManager
|
||||||
from core.app.apps.exc import GenerateTaskStoppedError
|
from core.app.apps.exc import GenerateTaskStoppedError
|
||||||
|
from core.app.apps.streaming_utils import stream_topic_events
|
||||||
from core.app.entities.app_invoke_entities import (
|
from core.app.entities.app_invoke_entities import (
|
||||||
AdvancedChatAppGenerateEntity,
|
AdvancedChatAppGenerateEntity,
|
||||||
AgentChatAppGenerateEntity,
|
AgentChatAppGenerateEntity,
|
||||||
AppGenerateEntity,
|
AppGenerateEntity,
|
||||||
ChatAppGenerateEntity,
|
ChatAppGenerateEntity,
|
||||||
CompletionAppGenerateEntity,
|
CompletionAppGenerateEntity,
|
||||||
|
ConversationAppGenerateEntity,
|
||||||
InvokeFrom,
|
InvokeFrom,
|
||||||
)
|
)
|
||||||
from core.app.entities.task_entities import (
|
from core.app.entities.task_entities import (
|
||||||
@@ -27,6 +29,8 @@ from core.app.entities.task_entities import (
|
|||||||
from core.app.task_pipeline.easy_ui_based_generate_task_pipeline import EasyUIBasedGenerateTaskPipeline
|
from core.app.task_pipeline.easy_ui_based_generate_task_pipeline import EasyUIBasedGenerateTaskPipeline
|
||||||
from core.prompt.utils.prompt_template_parser import PromptTemplateParser
|
from core.prompt.utils.prompt_template_parser import PromptTemplateParser
|
||||||
from extensions.ext_database import db
|
from extensions.ext_database import db
|
||||||
|
from extensions.ext_redis import get_pubsub_broadcast_channel
|
||||||
|
from libs.broadcast_channel.channel import Topic
|
||||||
from libs.datetime_utils import naive_utc_now
|
from libs.datetime_utils import naive_utc_now
|
||||||
from models import Account
|
from models import Account
|
||||||
from models.enums import CreatorUserRole
|
from models.enums import CreatorUserRole
|
||||||
@@ -156,6 +160,7 @@ class MessageBasedAppGenerator(BaseAppGenerator):
|
|||||||
query = application_generate_entity.query or "New conversation"
|
query = application_generate_entity.query or "New conversation"
|
||||||
conversation_name = (query[:20] + "…") if len(query) > 20 else query
|
conversation_name = (query[:20] + "…") if len(query) > 20 else query
|
||||||
|
|
||||||
|
created_new_conversation = conversation is None
|
||||||
try:
|
try:
|
||||||
if not conversation:
|
if not conversation:
|
||||||
conversation = Conversation(
|
conversation = Conversation(
|
||||||
@@ -232,6 +237,10 @@ class MessageBasedAppGenerator(BaseAppGenerator):
|
|||||||
db.session.add_all(message_files)
|
db.session.add_all(message_files)
|
||||||
|
|
||||||
db.session.commit()
|
db.session.commit()
|
||||||
|
|
||||||
|
if isinstance(application_generate_entity, ConversationAppGenerateEntity):
|
||||||
|
application_generate_entity.conversation_id = conversation.id
|
||||||
|
application_generate_entity.is_new_conversation = created_new_conversation
|
||||||
return conversation, message
|
return conversation, message
|
||||||
except Exception:
|
except Exception:
|
||||||
db.session.rollback()
|
db.session.rollback()
|
||||||
@@ -284,3 +293,29 @@ class MessageBasedAppGenerator(BaseAppGenerator):
|
|||||||
raise MessageNotExistsError("Message not exists")
|
raise MessageNotExistsError("Message not exists")
|
||||||
|
|
||||||
return message
|
return message
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _make_channel_key(app_mode: AppMode, workflow_run_id: str):
|
||||||
|
return f"channel:{app_mode}:{workflow_run_id}"
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def get_response_topic(cls, app_mode: AppMode, workflow_run_id: str) -> Topic:
|
||||||
|
key = cls._make_channel_key(app_mode, workflow_run_id)
|
||||||
|
channel = get_pubsub_broadcast_channel()
|
||||||
|
topic = channel.topic(key)
|
||||||
|
return topic
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def retrieve_events(
|
||||||
|
cls,
|
||||||
|
app_mode: AppMode,
|
||||||
|
workflow_run_id: str,
|
||||||
|
idle_timeout=300,
|
||||||
|
on_subscribe: Callable[[], None] | None = None,
|
||||||
|
) -> Generator[Mapping | str, None, None]:
|
||||||
|
topic = cls.get_response_topic(app_mode, workflow_run_id)
|
||||||
|
return stream_topic_events(
|
||||||
|
topic=topic,
|
||||||
|
idle_timeout=idle_timeout,
|
||||||
|
on_subscribe=on_subscribe,
|
||||||
|
)
|
||||||
|
|||||||
@@ -0,0 +1,36 @@
|
|||||||
|
from collections.abc import Callable, Generator, Mapping
|
||||||
|
|
||||||
|
from core.app.apps.streaming_utils import stream_topic_events
|
||||||
|
from extensions.ext_redis import get_pubsub_broadcast_channel
|
||||||
|
from libs.broadcast_channel.channel import Topic
|
||||||
|
from models.model import AppMode
|
||||||
|
|
||||||
|
|
||||||
|
class MessageGenerator:
|
||||||
|
@staticmethod
|
||||||
|
def _make_channel_key(app_mode: AppMode, workflow_run_id: str):
|
||||||
|
return f"channel:{app_mode}:{str(workflow_run_id)}"
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def get_response_topic(cls, app_mode: AppMode, workflow_run_id: str) -> Topic:
|
||||||
|
key = cls._make_channel_key(app_mode, workflow_run_id)
|
||||||
|
channel = get_pubsub_broadcast_channel()
|
||||||
|
topic = channel.topic(key)
|
||||||
|
return topic
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def retrieve_events(
|
||||||
|
cls,
|
||||||
|
app_mode: AppMode,
|
||||||
|
workflow_run_id: str,
|
||||||
|
idle_timeout=300,
|
||||||
|
ping_interval: float = 10.0,
|
||||||
|
on_subscribe: Callable[[], None] | None = None,
|
||||||
|
) -> Generator[Mapping | str, None, None]:
|
||||||
|
topic = cls.get_response_topic(app_mode, workflow_run_id)
|
||||||
|
return stream_topic_events(
|
||||||
|
topic=topic,
|
||||||
|
idle_timeout=idle_timeout,
|
||||||
|
ping_interval=ping_interval,
|
||||||
|
on_subscribe=on_subscribe,
|
||||||
|
)
|
||||||
@@ -0,0 +1,70 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import json
|
||||||
|
import time
|
||||||
|
from collections.abc import Callable, Generator, Iterable, Mapping
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
from core.app.entities.task_entities import StreamEvent
|
||||||
|
from libs.broadcast_channel.channel import Topic
|
||||||
|
from libs.broadcast_channel.exc import SubscriptionClosedError
|
||||||
|
|
||||||
|
|
||||||
|
def stream_topic_events(
|
||||||
|
*,
|
||||||
|
topic: Topic,
|
||||||
|
idle_timeout: float,
|
||||||
|
ping_interval: float | None = None,
|
||||||
|
on_subscribe: Callable[[], None] | None = None,
|
||||||
|
terminal_events: Iterable[str | StreamEvent] | None = None,
|
||||||
|
) -> Generator[Mapping[str, Any] | str, None, None]:
|
||||||
|
# send a PING event immediately to prevent the connection staying in pending state for a long time.
|
||||||
|
#
|
||||||
|
# This simplify the debugging process as the DevTools in Chrome does not
|
||||||
|
# provide complete curl command for pending connections.
|
||||||
|
yield StreamEvent.PING.value
|
||||||
|
|
||||||
|
terminal_values = _normalize_terminal_events(terminal_events)
|
||||||
|
last_msg_time = time.time()
|
||||||
|
last_ping_time = last_msg_time
|
||||||
|
with topic.subscribe() as sub:
|
||||||
|
# on_subscribe fires only after the Redis subscription is active.
|
||||||
|
# This is used to gate task start and reduce pub/sub race for the first event.
|
||||||
|
if on_subscribe is not None:
|
||||||
|
on_subscribe()
|
||||||
|
while True:
|
||||||
|
try:
|
||||||
|
msg = sub.receive(timeout=1)
|
||||||
|
except SubscriptionClosedError:
|
||||||
|
return
|
||||||
|
if msg is None:
|
||||||
|
current_time = time.time()
|
||||||
|
if current_time - last_msg_time > idle_timeout:
|
||||||
|
return
|
||||||
|
if ping_interval is not None and current_time - last_ping_time >= ping_interval:
|
||||||
|
yield StreamEvent.PING.value
|
||||||
|
last_ping_time = current_time
|
||||||
|
continue
|
||||||
|
|
||||||
|
last_msg_time = time.time()
|
||||||
|
last_ping_time = last_msg_time
|
||||||
|
event = json.loads(msg)
|
||||||
|
yield event
|
||||||
|
if not isinstance(event, dict):
|
||||||
|
continue
|
||||||
|
|
||||||
|
event_type = event.get("event")
|
||||||
|
if event_type in terminal_values:
|
||||||
|
return
|
||||||
|
|
||||||
|
|
||||||
|
def _normalize_terminal_events(terminal_events: Iterable[str | StreamEvent] | None) -> set[str]:
|
||||||
|
if not terminal_events:
|
||||||
|
return {StreamEvent.WORKFLOW_FINISHED.value, StreamEvent.WORKFLOW_PAUSED.value}
|
||||||
|
values: set[str] = set()
|
||||||
|
for item in terminal_events:
|
||||||
|
if isinstance(item, StreamEvent):
|
||||||
|
values.add(item.value)
|
||||||
|
else:
|
||||||
|
values.add(str(item))
|
||||||
|
return values
|
||||||
@@ -25,6 +25,7 @@ from core.app.apps.workflow.generate_response_converter import WorkflowAppGenera
|
|||||||
from core.app.apps.workflow.generate_task_pipeline import WorkflowAppGenerateTaskPipeline
|
from core.app.apps.workflow.generate_task_pipeline import WorkflowAppGenerateTaskPipeline
|
||||||
from core.app.entities.app_invoke_entities import InvokeFrom, WorkflowAppGenerateEntity
|
from core.app.entities.app_invoke_entities import InvokeFrom, WorkflowAppGenerateEntity
|
||||||
from core.app.entities.task_entities import WorkflowAppBlockingResponse, WorkflowAppStreamResponse
|
from core.app.entities.task_entities import WorkflowAppBlockingResponse, WorkflowAppStreamResponse
|
||||||
|
from core.app.layers.pause_state_persist_layer import PauseStateLayerConfig, PauseStatePersistenceLayer
|
||||||
from core.db.session_factory import session_factory
|
from core.db.session_factory import session_factory
|
||||||
from core.helper.trace_id_helper import extract_external_trace_id_from_args
|
from core.helper.trace_id_helper import extract_external_trace_id_from_args
|
||||||
from core.model_runtime.errors.invoke import InvokeAuthorizationError
|
from core.model_runtime.errors.invoke import InvokeAuthorizationError
|
||||||
@@ -34,12 +35,15 @@ from core.workflow.graph_engine.layers.base import GraphEngineLayer
|
|||||||
from core.workflow.repositories.draft_variable_repository import DraftVariableSaverFactory
|
from core.workflow.repositories.draft_variable_repository import DraftVariableSaverFactory
|
||||||
from core.workflow.repositories.workflow_execution_repository import WorkflowExecutionRepository
|
from core.workflow.repositories.workflow_execution_repository import WorkflowExecutionRepository
|
||||||
from core.workflow.repositories.workflow_node_execution_repository import WorkflowNodeExecutionRepository
|
from core.workflow.repositories.workflow_node_execution_repository import WorkflowNodeExecutionRepository
|
||||||
|
from core.workflow.runtime import GraphRuntimeState
|
||||||
from core.workflow.variable_loader import DUMMY_VARIABLE_LOADER, VariableLoader
|
from core.workflow.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 libs.flask_utils import preserve_flask_contexts
|
from libs.flask_utils import preserve_flask_contexts
|
||||||
from models import Account, App, EndUser, Workflow, WorkflowNodeExecutionTriggeredFrom
|
from models.account import Account
|
||||||
from models.enums import WorkflowRunTriggeredFrom
|
from models.enums import WorkflowRunTriggeredFrom
|
||||||
|
from models.model import App, EndUser
|
||||||
|
from models.workflow import Workflow, WorkflowNodeExecutionTriggeredFrom
|
||||||
from services.workflow_draft_variable_service import DraftVarLoader, WorkflowDraftVariableService
|
from services.workflow_draft_variable_service import DraftVarLoader, WorkflowDraftVariableService
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
@@ -66,9 +70,11 @@ class WorkflowAppGenerator(BaseAppGenerator):
|
|||||||
invoke_from: InvokeFrom,
|
invoke_from: InvokeFrom,
|
||||||
streaming: Literal[True],
|
streaming: Literal[True],
|
||||||
call_depth: int,
|
call_depth: int,
|
||||||
|
workflow_run_id: str | uuid.UUID | None = None,
|
||||||
triggered_from: WorkflowRunTriggeredFrom | None = None,
|
triggered_from: WorkflowRunTriggeredFrom | None = None,
|
||||||
root_node_id: str | None = None,
|
root_node_id: str | None = None,
|
||||||
graph_engine_layers: Sequence[GraphEngineLayer] = (),
|
graph_engine_layers: Sequence[GraphEngineLayer] = (),
|
||||||
|
pause_state_config: PauseStateLayerConfig | None = None,
|
||||||
) -> Generator[Mapping[str, Any] | str, None, None]: ...
|
) -> Generator[Mapping[str, Any] | str, None, None]: ...
|
||||||
|
|
||||||
@overload
|
@overload
|
||||||
@@ -82,9 +88,11 @@ class WorkflowAppGenerator(BaseAppGenerator):
|
|||||||
invoke_from: InvokeFrom,
|
invoke_from: InvokeFrom,
|
||||||
streaming: Literal[False],
|
streaming: Literal[False],
|
||||||
call_depth: int,
|
call_depth: int,
|
||||||
|
workflow_run_id: str | uuid.UUID | None = None,
|
||||||
triggered_from: WorkflowRunTriggeredFrom | None = None,
|
triggered_from: WorkflowRunTriggeredFrom | None = None,
|
||||||
root_node_id: str | None = None,
|
root_node_id: str | None = None,
|
||||||
graph_engine_layers: Sequence[GraphEngineLayer] = (),
|
graph_engine_layers: Sequence[GraphEngineLayer] = (),
|
||||||
|
pause_state_config: PauseStateLayerConfig | None = None,
|
||||||
) -> Mapping[str, Any]: ...
|
) -> Mapping[str, Any]: ...
|
||||||
|
|
||||||
@overload
|
@overload
|
||||||
@@ -98,9 +106,11 @@ class WorkflowAppGenerator(BaseAppGenerator):
|
|||||||
invoke_from: InvokeFrom,
|
invoke_from: InvokeFrom,
|
||||||
streaming: bool,
|
streaming: bool,
|
||||||
call_depth: int,
|
call_depth: int,
|
||||||
|
workflow_run_id: str | uuid.UUID | None = None,
|
||||||
triggered_from: WorkflowRunTriggeredFrom | None = None,
|
triggered_from: WorkflowRunTriggeredFrom | None = None,
|
||||||
root_node_id: str | None = None,
|
root_node_id: str | None = None,
|
||||||
graph_engine_layers: Sequence[GraphEngineLayer] = (),
|
graph_engine_layers: Sequence[GraphEngineLayer] = (),
|
||||||
|
pause_state_config: PauseStateLayerConfig | None = None,
|
||||||
) -> Union[Mapping[str, Any], Generator[Mapping[str, Any] | str, None, None]]: ...
|
) -> Union[Mapping[str, Any], Generator[Mapping[str, Any] | str, None, None]]: ...
|
||||||
|
|
||||||
def generate(
|
def generate(
|
||||||
@@ -113,9 +123,11 @@ class WorkflowAppGenerator(BaseAppGenerator):
|
|||||||
invoke_from: InvokeFrom,
|
invoke_from: InvokeFrom,
|
||||||
streaming: bool = True,
|
streaming: bool = True,
|
||||||
call_depth: int = 0,
|
call_depth: int = 0,
|
||||||
|
workflow_run_id: str | uuid.UUID | None = None,
|
||||||
triggered_from: WorkflowRunTriggeredFrom | None = None,
|
triggered_from: WorkflowRunTriggeredFrom | None = None,
|
||||||
root_node_id: str | None = None,
|
root_node_id: str | None = None,
|
||||||
graph_engine_layers: Sequence[GraphEngineLayer] = (),
|
graph_engine_layers: Sequence[GraphEngineLayer] = (),
|
||||||
|
pause_state_config: PauseStateLayerConfig | None = None,
|
||||||
) -> Union[Mapping[str, Any], Generator[Mapping[str, Any] | str, None, None]]:
|
) -> Union[Mapping[str, Any], Generator[Mapping[str, Any] | str, None, None]]:
|
||||||
files: Sequence[Mapping[str, Any]] = args.get("files") or []
|
files: Sequence[Mapping[str, Any]] = args.get("files") or []
|
||||||
|
|
||||||
@@ -150,7 +162,7 @@ class WorkflowAppGenerator(BaseAppGenerator):
|
|||||||
extras = {
|
extras = {
|
||||||
**extract_external_trace_id_from_args(args),
|
**extract_external_trace_id_from_args(args),
|
||||||
}
|
}
|
||||||
workflow_run_id = str(uuid.uuid4())
|
workflow_run_id = str(workflow_run_id or uuid.uuid4())
|
||||||
# FIXME (Yeuoly): we need to remove the SKIP_PREPARE_USER_INPUTS_KEY from the args
|
# FIXME (Yeuoly): we need to remove the SKIP_PREPARE_USER_INPUTS_KEY from the args
|
||||||
# trigger shouldn't prepare user inputs
|
# trigger shouldn't prepare user inputs
|
||||||
if self._should_prepare_user_inputs(args):
|
if self._should_prepare_user_inputs(args):
|
||||||
@@ -216,13 +228,40 @@ class WorkflowAppGenerator(BaseAppGenerator):
|
|||||||
streaming=streaming,
|
streaming=streaming,
|
||||||
root_node_id=root_node_id,
|
root_node_id=root_node_id,
|
||||||
graph_engine_layers=graph_engine_layers,
|
graph_engine_layers=graph_engine_layers,
|
||||||
|
pause_state_config=pause_state_config,
|
||||||
)
|
)
|
||||||
|
|
||||||
def resume(self, *, workflow_run_id: str) -> None:
|
def resume(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
app_model: App,
|
||||||
|
workflow: Workflow,
|
||||||
|
user: Union[Account, EndUser],
|
||||||
|
application_generate_entity: WorkflowAppGenerateEntity,
|
||||||
|
graph_runtime_state: GraphRuntimeState,
|
||||||
|
workflow_execution_repository: WorkflowExecutionRepository,
|
||||||
|
workflow_node_execution_repository: WorkflowNodeExecutionRepository,
|
||||||
|
graph_engine_layers: Sequence[GraphEngineLayer] = (),
|
||||||
|
pause_state_config: PauseStateLayerConfig | None = None,
|
||||||
|
variable_loader: VariableLoader = DUMMY_VARIABLE_LOADER,
|
||||||
|
) -> Union[Mapping[str, Any], Generator[str | Mapping[str, Any], None, None]]:
|
||||||
"""
|
"""
|
||||||
@TBD
|
Resume a paused workflow execution using the persisted runtime state.
|
||||||
"""
|
"""
|
||||||
pass
|
return self._generate(
|
||||||
|
app_model=app_model,
|
||||||
|
workflow=workflow,
|
||||||
|
user=user,
|
||||||
|
application_generate_entity=application_generate_entity,
|
||||||
|
invoke_from=application_generate_entity.invoke_from,
|
||||||
|
workflow_execution_repository=workflow_execution_repository,
|
||||||
|
workflow_node_execution_repository=workflow_node_execution_repository,
|
||||||
|
streaming=application_generate_entity.stream,
|
||||||
|
variable_loader=variable_loader,
|
||||||
|
graph_engine_layers=graph_engine_layers,
|
||||||
|
graph_runtime_state=graph_runtime_state,
|
||||||
|
pause_state_config=pause_state_config,
|
||||||
|
)
|
||||||
|
|
||||||
def _generate(
|
def _generate(
|
||||||
self,
|
self,
|
||||||
@@ -238,6 +277,8 @@ class WorkflowAppGenerator(BaseAppGenerator):
|
|||||||
variable_loader: VariableLoader = DUMMY_VARIABLE_LOADER,
|
variable_loader: VariableLoader = DUMMY_VARIABLE_LOADER,
|
||||||
root_node_id: str | None = None,
|
root_node_id: str | None = None,
|
||||||
graph_engine_layers: Sequence[GraphEngineLayer] = (),
|
graph_engine_layers: Sequence[GraphEngineLayer] = (),
|
||||||
|
graph_runtime_state: GraphRuntimeState | None = None,
|
||||||
|
pause_state_config: PauseStateLayerConfig | None = None,
|
||||||
) -> Union[Mapping[str, Any], Generator[str | Mapping[str, Any], None, None]]:
|
) -> Union[Mapping[str, Any], Generator[str | Mapping[str, Any], None, None]]:
|
||||||
"""
|
"""
|
||||||
Generate App response.
|
Generate App response.
|
||||||
@@ -251,6 +292,8 @@ class WorkflowAppGenerator(BaseAppGenerator):
|
|||||||
:param workflow_node_execution_repository: repository for workflow node execution
|
:param workflow_node_execution_repository: repository for workflow node execution
|
||||||
:param streaming: is stream
|
:param streaming: is stream
|
||||||
"""
|
"""
|
||||||
|
graph_layers: list[GraphEngineLayer] = list(graph_engine_layers)
|
||||||
|
|
||||||
# init queue manager
|
# init queue manager
|
||||||
queue_manager = WorkflowAppQueueManager(
|
queue_manager = WorkflowAppQueueManager(
|
||||||
task_id=application_generate_entity.task_id,
|
task_id=application_generate_entity.task_id,
|
||||||
@@ -259,6 +302,15 @@ class WorkflowAppGenerator(BaseAppGenerator):
|
|||||||
app_mode=app_model.mode,
|
app_mode=app_model.mode,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
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
|
# new thread with request context and contextvars
|
||||||
context = contextvars.copy_context()
|
context = contextvars.copy_context()
|
||||||
|
|
||||||
@@ -276,7 +328,8 @@ class WorkflowAppGenerator(BaseAppGenerator):
|
|||||||
"root_node_id": root_node_id,
|
"root_node_id": root_node_id,
|
||||||
"workflow_execution_repository": workflow_execution_repository,
|
"workflow_execution_repository": workflow_execution_repository,
|
||||||
"workflow_node_execution_repository": workflow_node_execution_repository,
|
"workflow_node_execution_repository": workflow_node_execution_repository,
|
||||||
"graph_engine_layers": graph_engine_layers,
|
"graph_engine_layers": tuple(graph_layers),
|
||||||
|
"graph_runtime_state": graph_runtime_state,
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -378,6 +431,7 @@ class WorkflowAppGenerator(BaseAppGenerator):
|
|||||||
workflow_node_execution_repository=workflow_node_execution_repository,
|
workflow_node_execution_repository=workflow_node_execution_repository,
|
||||||
streaming=streaming,
|
streaming=streaming,
|
||||||
variable_loader=var_loader,
|
variable_loader=var_loader,
|
||||||
|
pause_state_config=None,
|
||||||
)
|
)
|
||||||
|
|
||||||
def single_loop_generate(
|
def single_loop_generate(
|
||||||
@@ -459,6 +513,7 @@ class WorkflowAppGenerator(BaseAppGenerator):
|
|||||||
workflow_node_execution_repository=workflow_node_execution_repository,
|
workflow_node_execution_repository=workflow_node_execution_repository,
|
||||||
streaming=streaming,
|
streaming=streaming,
|
||||||
variable_loader=var_loader,
|
variable_loader=var_loader,
|
||||||
|
pause_state_config=None,
|
||||||
)
|
)
|
||||||
|
|
||||||
def _generate_worker(
|
def _generate_worker(
|
||||||
@@ -472,6 +527,7 @@ class WorkflowAppGenerator(BaseAppGenerator):
|
|||||||
workflow_node_execution_repository: WorkflowNodeExecutionRepository,
|
workflow_node_execution_repository: WorkflowNodeExecutionRepository,
|
||||||
root_node_id: str | None = None,
|
root_node_id: str | None = None,
|
||||||
graph_engine_layers: Sequence[GraphEngineLayer] = (),
|
graph_engine_layers: Sequence[GraphEngineLayer] = (),
|
||||||
|
graph_runtime_state: GraphRuntimeState | None = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""
|
"""
|
||||||
Generate worker in a new thread.
|
Generate worker in a new thread.
|
||||||
@@ -517,6 +573,7 @@ class WorkflowAppGenerator(BaseAppGenerator):
|
|||||||
workflow_node_execution_repository=workflow_node_execution_repository,
|
workflow_node_execution_repository=workflow_node_execution_repository,
|
||||||
root_node_id=root_node_id,
|
root_node_id=root_node_id,
|
||||||
graph_engine_layers=graph_engine_layers,
|
graph_engine_layers=graph_engine_layers,
|
||||||
|
graph_runtime_state=graph_runtime_state,
|
||||||
)
|
)
|
||||||
|
|
||||||
try:
|
try:
|
||||||
|
|||||||
@@ -42,6 +42,7 @@ class WorkflowAppRunner(WorkflowBasedAppRunner):
|
|||||||
workflow_execution_repository: WorkflowExecutionRepository,
|
workflow_execution_repository: WorkflowExecutionRepository,
|
||||||
workflow_node_execution_repository: WorkflowNodeExecutionRepository,
|
workflow_node_execution_repository: WorkflowNodeExecutionRepository,
|
||||||
graph_engine_layers: Sequence[GraphEngineLayer] = (),
|
graph_engine_layers: Sequence[GraphEngineLayer] = (),
|
||||||
|
graph_runtime_state: GraphRuntimeState | None = None,
|
||||||
):
|
):
|
||||||
super().__init__(
|
super().__init__(
|
||||||
queue_manager=queue_manager,
|
queue_manager=queue_manager,
|
||||||
@@ -55,6 +56,7 @@ class WorkflowAppRunner(WorkflowBasedAppRunner):
|
|||||||
self._root_node_id = root_node_id
|
self._root_node_id = root_node_id
|
||||||
self._workflow_execution_repository = workflow_execution_repository
|
self._workflow_execution_repository = workflow_execution_repository
|
||||||
self._workflow_node_execution_repository = workflow_node_execution_repository
|
self._workflow_node_execution_repository = workflow_node_execution_repository
|
||||||
|
self._resume_graph_runtime_state = graph_runtime_state
|
||||||
|
|
||||||
@trace_span(WorkflowAppRunnerHandler)
|
@trace_span(WorkflowAppRunnerHandler)
|
||||||
def run(self):
|
def run(self):
|
||||||
@@ -63,23 +65,28 @@ class WorkflowAppRunner(WorkflowBasedAppRunner):
|
|||||||
"""
|
"""
|
||||||
app_config = self.application_generate_entity.app_config
|
app_config = self.application_generate_entity.app_config
|
||||||
app_config = cast(WorkflowAppConfig, app_config)
|
app_config = cast(WorkflowAppConfig, app_config)
|
||||||
|
|
||||||
system_inputs = SystemVariable(
|
|
||||||
files=self.application_generate_entity.files,
|
|
||||||
user_id=self._sys_user_id,
|
|
||||||
app_id=app_config.app_id,
|
|
||||||
timestamp=int(naive_utc_now().timestamp()),
|
|
||||||
workflow_id=app_config.workflow_id,
|
|
||||||
workflow_execution_id=self.application_generate_entity.workflow_execution_id,
|
|
||||||
)
|
|
||||||
|
|
||||||
invoke_from = self.application_generate_entity.invoke_from
|
invoke_from = self.application_generate_entity.invoke_from
|
||||||
# if only single iteration or single loop run is requested
|
# if only single iteration or single loop run is requested
|
||||||
if self.application_generate_entity.single_iteration_run or self.application_generate_entity.single_loop_run:
|
if self.application_generate_entity.single_iteration_run or self.application_generate_entity.single_loop_run:
|
||||||
invoke_from = InvokeFrom.DEBUGGER
|
invoke_from = InvokeFrom.DEBUGGER
|
||||||
user_from = self._resolve_user_from(invoke_from)
|
user_from = self._resolve_user_from(invoke_from)
|
||||||
|
|
||||||
if self.application_generate_entity.single_iteration_run or self.application_generate_entity.single_loop_run:
|
resume_state = self._resume_graph_runtime_state
|
||||||
|
|
||||||
|
if resume_state is not None:
|
||||||
|
graph_runtime_state = resume_state
|
||||||
|
variable_pool = graph_runtime_state.variable_pool
|
||||||
|
graph = self._init_graph(
|
||||||
|
graph_config=self._workflow.graph_dict,
|
||||||
|
graph_runtime_state=graph_runtime_state,
|
||||||
|
workflow_id=self._workflow.id,
|
||||||
|
tenant_id=self._workflow.tenant_id,
|
||||||
|
user_id=self.application_generate_entity.user_id,
|
||||||
|
user_from=user_from,
|
||||||
|
invoke_from=invoke_from,
|
||||||
|
root_node_id=self._root_node_id,
|
||||||
|
)
|
||||||
|
elif self.application_generate_entity.single_iteration_run or self.application_generate_entity.single_loop_run:
|
||||||
graph, variable_pool, graph_runtime_state = self._prepare_single_node_execution(
|
graph, variable_pool, graph_runtime_state = self._prepare_single_node_execution(
|
||||||
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,
|
||||||
@@ -89,7 +96,14 @@ class WorkflowAppRunner(WorkflowBasedAppRunner):
|
|||||||
inputs = self.application_generate_entity.inputs
|
inputs = self.application_generate_entity.inputs
|
||||||
|
|
||||||
# Create a variable pool.
|
# Create a variable pool.
|
||||||
|
system_inputs = SystemVariable(
|
||||||
|
files=self.application_generate_entity.files,
|
||||||
|
user_id=self._sys_user_id,
|
||||||
|
app_id=app_config.app_id,
|
||||||
|
timestamp=int(naive_utc_now().timestamp()),
|
||||||
|
workflow_id=app_config.workflow_id,
|
||||||
|
workflow_execution_id=self.application_generate_entity.workflow_execution_id,
|
||||||
|
)
|
||||||
variable_pool = VariablePool(
|
variable_pool = VariablePool(
|
||||||
system_variables=system_inputs,
|
system_variables=system_inputs,
|
||||||
user_inputs=inputs,
|
user_inputs=inputs,
|
||||||
@@ -98,8 +112,6 @@ class WorkflowAppRunner(WorkflowBasedAppRunner):
|
|||||||
)
|
)
|
||||||
|
|
||||||
graph_runtime_state = GraphRuntimeState(variable_pool=variable_pool, start_at=time.perf_counter())
|
graph_runtime_state = GraphRuntimeState(variable_pool=variable_pool, start_at=time.perf_counter())
|
||||||
|
|
||||||
# init graph
|
|
||||||
graph = self._init_graph(
|
graph = self._init_graph(
|
||||||
graph_config=self._workflow.graph_dict,
|
graph_config=self._workflow.graph_dict,
|
||||||
graph_runtime_state=graph_runtime_state,
|
graph_runtime_state=graph_runtime_state,
|
||||||
|
|||||||
@@ -0,0 +1,7 @@
|
|||||||
|
from libs.exception import BaseHTTPException
|
||||||
|
|
||||||
|
|
||||||
|
class WorkflowPausedInBlockingModeError(BaseHTTPException):
|
||||||
|
error_code = "workflow_paused_in_blocking_mode"
|
||||||
|
description = "Workflow execution paused for human input; blocking response mode is not supported."
|
||||||
|
code = 400
|
||||||
@@ -16,6 +16,8 @@ from core.app.entities.queue_entities import (
|
|||||||
MessageQueueMessage,
|
MessageQueueMessage,
|
||||||
QueueAgentLogEvent,
|
QueueAgentLogEvent,
|
||||||
QueueErrorEvent,
|
QueueErrorEvent,
|
||||||
|
QueueHumanInputFormFilledEvent,
|
||||||
|
QueueHumanInputFormTimeoutEvent,
|
||||||
QueueIterationCompletedEvent,
|
QueueIterationCompletedEvent,
|
||||||
QueueIterationNextEvent,
|
QueueIterationNextEvent,
|
||||||
QueueIterationStartEvent,
|
QueueIterationStartEvent,
|
||||||
@@ -32,6 +34,7 @@ from core.app.entities.queue_entities import (
|
|||||||
QueueTextChunkEvent,
|
QueueTextChunkEvent,
|
||||||
QueueWorkflowFailedEvent,
|
QueueWorkflowFailedEvent,
|
||||||
QueueWorkflowPartialSuccessEvent,
|
QueueWorkflowPartialSuccessEvent,
|
||||||
|
QueueWorkflowPausedEvent,
|
||||||
QueueWorkflowStartedEvent,
|
QueueWorkflowStartedEvent,
|
||||||
QueueWorkflowSucceededEvent,
|
QueueWorkflowSucceededEvent,
|
||||||
WorkflowQueueMessage,
|
WorkflowQueueMessage,
|
||||||
@@ -46,11 +49,13 @@ from core.app.entities.task_entities import (
|
|||||||
WorkflowAppBlockingResponse,
|
WorkflowAppBlockingResponse,
|
||||||
WorkflowAppStreamResponse,
|
WorkflowAppStreamResponse,
|
||||||
WorkflowFinishStreamResponse,
|
WorkflowFinishStreamResponse,
|
||||||
|
WorkflowPauseStreamResponse,
|
||||||
WorkflowStartStreamResponse,
|
WorkflowStartStreamResponse,
|
||||||
)
|
)
|
||||||
from core.app.task_pipeline.based_generate_task_pipeline import BasedGenerateTaskPipeline
|
from core.app.task_pipeline.based_generate_task_pipeline import BasedGenerateTaskPipeline
|
||||||
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.workflow.entities.workflow_start_reason import WorkflowStartReason
|
||||||
from core.workflow.enums import WorkflowExecutionStatus
|
from core.workflow.enums import WorkflowExecutionStatus
|
||||||
from core.workflow.repositories.draft_variable_repository import DraftVariableSaverFactory
|
from core.workflow.repositories.draft_variable_repository import DraftVariableSaverFactory
|
||||||
from core.workflow.runtime import GraphRuntimeState
|
from core.workflow.runtime import GraphRuntimeState
|
||||||
@@ -132,6 +137,25 @@ class WorkflowAppGenerateTaskPipeline(GraphRuntimeStateSupport):
|
|||||||
for stream_response in generator:
|
for stream_response in generator:
|
||||||
if isinstance(stream_response, ErrorStreamResponse):
|
if isinstance(stream_response, ErrorStreamResponse):
|
||||||
raise stream_response.err
|
raise stream_response.err
|
||||||
|
elif isinstance(stream_response, WorkflowPauseStreamResponse):
|
||||||
|
response = WorkflowAppBlockingResponse(
|
||||||
|
task_id=self._application_generate_entity.task_id,
|
||||||
|
workflow_run_id=stream_response.data.workflow_run_id,
|
||||||
|
data=WorkflowAppBlockingResponse.Data(
|
||||||
|
id=stream_response.data.workflow_run_id,
|
||||||
|
workflow_id=self._workflow.id,
|
||||||
|
status=stream_response.data.status,
|
||||||
|
outputs=stream_response.data.outputs or {},
|
||||||
|
error=None,
|
||||||
|
elapsed_time=stream_response.data.elapsed_time,
|
||||||
|
total_tokens=stream_response.data.total_tokens,
|
||||||
|
total_steps=stream_response.data.total_steps,
|
||||||
|
created_at=stream_response.data.created_at,
|
||||||
|
finished_at=None,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
return response
|
||||||
elif isinstance(stream_response, WorkflowFinishStreamResponse):
|
elif isinstance(stream_response, WorkflowFinishStreamResponse):
|
||||||
response = WorkflowAppBlockingResponse(
|
response = WorkflowAppBlockingResponse(
|
||||||
task_id=self._application_generate_entity.task_id,
|
task_id=self._application_generate_entity.task_id,
|
||||||
@@ -146,7 +170,7 @@ class WorkflowAppGenerateTaskPipeline(GraphRuntimeStateSupport):
|
|||||||
total_tokens=stream_response.data.total_tokens,
|
total_tokens=stream_response.data.total_tokens,
|
||||||
total_steps=stream_response.data.total_steps,
|
total_steps=stream_response.data.total_steps,
|
||||||
created_at=int(stream_response.data.created_at),
|
created_at=int(stream_response.data.created_at),
|
||||||
finished_at=int(stream_response.data.finished_at),
|
finished_at=int(stream_response.data.finished_at) if stream_response.data.finished_at else None,
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -259,13 +283,15 @@ class WorkflowAppGenerateTaskPipeline(GraphRuntimeStateSupport):
|
|||||||
run_id = self._extract_workflow_run_id(runtime_state)
|
run_id = self._extract_workflow_run_id(runtime_state)
|
||||||
self._workflow_execution_id = run_id
|
self._workflow_execution_id = run_id
|
||||||
|
|
||||||
with self._database_session() as session:
|
if event.reason == WorkflowStartReason.INITIAL:
|
||||||
self._save_workflow_app_log(session=session, workflow_run_id=self._workflow_execution_id)
|
with self._database_session() as session:
|
||||||
|
self._save_workflow_app_log(session=session, workflow_run_id=self._workflow_execution_id)
|
||||||
|
|
||||||
start_resp = self._workflow_response_converter.workflow_start_to_stream_response(
|
start_resp = self._workflow_response_converter.workflow_start_to_stream_response(
|
||||||
task_id=self._application_generate_entity.task_id,
|
task_id=self._application_generate_entity.task_id,
|
||||||
workflow_run_id=run_id,
|
workflow_run_id=run_id,
|
||||||
workflow_id=self._workflow.id,
|
workflow_id=self._workflow.id,
|
||||||
|
reason=event.reason,
|
||||||
)
|
)
|
||||||
yield start_resp
|
yield start_resp
|
||||||
|
|
||||||
@@ -440,6 +466,21 @@ class WorkflowAppGenerateTaskPipeline(GraphRuntimeStateSupport):
|
|||||||
)
|
)
|
||||||
yield workflow_finish_resp
|
yield workflow_finish_resp
|
||||||
|
|
||||||
|
def _handle_workflow_paused_event(
|
||||||
|
self,
|
||||||
|
event: QueueWorkflowPausedEvent,
|
||||||
|
**kwargs,
|
||||||
|
) -> Generator[StreamResponse, None, None]:
|
||||||
|
"""Handle workflow paused events."""
|
||||||
|
self._ensure_workflow_initialized()
|
||||||
|
validated_state = self._ensure_graph_runtime_initialized()
|
||||||
|
responses = self._workflow_response_converter.workflow_pause_to_stream_response(
|
||||||
|
event=event,
|
||||||
|
task_id=self._application_generate_entity.task_id,
|
||||||
|
graph_runtime_state=validated_state,
|
||||||
|
)
|
||||||
|
yield from responses
|
||||||
|
|
||||||
def _handle_workflow_failed_and_stop_events(
|
def _handle_workflow_failed_and_stop_events(
|
||||||
self,
|
self,
|
||||||
event: Union[QueueWorkflowFailedEvent, QueueStopEvent],
|
event: Union[QueueWorkflowFailedEvent, QueueStopEvent],
|
||||||
@@ -495,6 +536,22 @@ class WorkflowAppGenerateTaskPipeline(GraphRuntimeStateSupport):
|
|||||||
task_id=self._application_generate_entity.task_id, event=event
|
task_id=self._application_generate_entity.task_id, event=event
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def _handle_human_input_form_filled_event(
|
||||||
|
self, event: QueueHumanInputFormFilledEvent, **kwargs
|
||||||
|
) -> Generator[StreamResponse, None, None]:
|
||||||
|
"""Handle human input form filled events."""
|
||||||
|
yield self._workflow_response_converter.human_input_form_filled_to_stream_response(
|
||||||
|
event=event, task_id=self._application_generate_entity.task_id
|
||||||
|
)
|
||||||
|
|
||||||
|
def _handle_human_input_form_timeout_event(
|
||||||
|
self, event: QueueHumanInputFormTimeoutEvent, **kwargs
|
||||||
|
) -> Generator[StreamResponse, None, None]:
|
||||||
|
"""Handle human input form timeout events."""
|
||||||
|
yield self._workflow_response_converter.human_input_form_timeout_to_stream_response(
|
||||||
|
event=event, task_id=self._application_generate_entity.task_id
|
||||||
|
)
|
||||||
|
|
||||||
def _get_event_handlers(self) -> dict[type, Callable]:
|
def _get_event_handlers(self) -> dict[type, Callable]:
|
||||||
"""Get mapping of event types to their handlers using fluent pattern."""
|
"""Get mapping of event types to their handlers using fluent pattern."""
|
||||||
return {
|
return {
|
||||||
@@ -506,6 +563,7 @@ class WorkflowAppGenerateTaskPipeline(GraphRuntimeStateSupport):
|
|||||||
QueueWorkflowStartedEvent: self._handle_workflow_started_event,
|
QueueWorkflowStartedEvent: self._handle_workflow_started_event,
|
||||||
QueueWorkflowSucceededEvent: self._handle_workflow_succeeded_event,
|
QueueWorkflowSucceededEvent: self._handle_workflow_succeeded_event,
|
||||||
QueueWorkflowPartialSuccessEvent: self._handle_workflow_partial_success_event,
|
QueueWorkflowPartialSuccessEvent: self._handle_workflow_partial_success_event,
|
||||||
|
QueueWorkflowPausedEvent: self._handle_workflow_paused_event,
|
||||||
# Node events
|
# Node events
|
||||||
QueueNodeRetryEvent: self._handle_node_retry_event,
|
QueueNodeRetryEvent: self._handle_node_retry_event,
|
||||||
QueueNodeStartedEvent: self._handle_node_started_event,
|
QueueNodeStartedEvent: self._handle_node_started_event,
|
||||||
@@ -520,6 +578,8 @@ class WorkflowAppGenerateTaskPipeline(GraphRuntimeStateSupport):
|
|||||||
QueueLoopCompletedEvent: self._handle_loop_completed_event,
|
QueueLoopCompletedEvent: self._handle_loop_completed_event,
|
||||||
# Agent events
|
# Agent events
|
||||||
QueueAgentLogEvent: self._handle_agent_log_event,
|
QueueAgentLogEvent: self._handle_agent_log_event,
|
||||||
|
QueueHumanInputFormFilledEvent: self._handle_human_input_form_filled_event,
|
||||||
|
QueueHumanInputFormTimeoutEvent: self._handle_human_input_form_timeout_event,
|
||||||
}
|
}
|
||||||
|
|
||||||
def _dispatch_event(
|
def _dispatch_event(
|
||||||
@@ -602,6 +662,9 @@ class WorkflowAppGenerateTaskPipeline(GraphRuntimeStateSupport):
|
|||||||
case QueueWorkflowFailedEvent():
|
case QueueWorkflowFailedEvent():
|
||||||
yield from self._handle_workflow_failed_and_stop_events(event)
|
yield from self._handle_workflow_failed_and_stop_events(event)
|
||||||
break
|
break
|
||||||
|
case QueueWorkflowPausedEvent():
|
||||||
|
yield from self._handle_workflow_paused_event(event)
|
||||||
|
break
|
||||||
|
|
||||||
case QueueStopEvent():
|
case QueueStopEvent():
|
||||||
yield from self._handle_workflow_failed_and_stop_events(event)
|
yield from self._handle_workflow_failed_and_stop_events(event)
|
||||||
|
|||||||
@@ -1,3 +1,4 @@
|
|||||||
|
import logging
|
||||||
import time
|
import time
|
||||||
from collections.abc import Mapping, Sequence
|
from collections.abc import Mapping, Sequence
|
||||||
from typing import Any, cast
|
from typing import Any, cast
|
||||||
@@ -7,6 +8,8 @@ from core.app.entities.app_invoke_entities import InvokeFrom
|
|||||||
from core.app.entities.queue_entities import (
|
from core.app.entities.queue_entities import (
|
||||||
AppQueueEvent,
|
AppQueueEvent,
|
||||||
QueueAgentLogEvent,
|
QueueAgentLogEvent,
|
||||||
|
QueueHumanInputFormFilledEvent,
|
||||||
|
QueueHumanInputFormTimeoutEvent,
|
||||||
QueueIterationCompletedEvent,
|
QueueIterationCompletedEvent,
|
||||||
QueueIterationNextEvent,
|
QueueIterationNextEvent,
|
||||||
QueueIterationStartEvent,
|
QueueIterationStartEvent,
|
||||||
@@ -22,22 +25,27 @@ from core.app.entities.queue_entities import (
|
|||||||
QueueTextChunkEvent,
|
QueueTextChunkEvent,
|
||||||
QueueWorkflowFailedEvent,
|
QueueWorkflowFailedEvent,
|
||||||
QueueWorkflowPartialSuccessEvent,
|
QueueWorkflowPartialSuccessEvent,
|
||||||
|
QueueWorkflowPausedEvent,
|
||||||
QueueWorkflowStartedEvent,
|
QueueWorkflowStartedEvent,
|
||||||
QueueWorkflowSucceededEvent,
|
QueueWorkflowSucceededEvent,
|
||||||
)
|
)
|
||||||
from core.app.workflow.node_factory import DifyNodeFactory
|
from core.app.workflow.node_factory import DifyNodeFactory
|
||||||
from core.workflow.entities import GraphInitParams
|
from core.workflow.entities import GraphInitParams
|
||||||
|
from core.workflow.entities.pause_reason import HumanInputRequired
|
||||||
from core.workflow.graph import Graph
|
from core.workflow.graph import Graph
|
||||||
from core.workflow.graph_engine.layers.base import GraphEngineLayer
|
from core.workflow.graph_engine.layers.base import GraphEngineLayer
|
||||||
from core.workflow.graph_events import (
|
from core.workflow.graph_events import (
|
||||||
GraphEngineEvent,
|
GraphEngineEvent,
|
||||||
GraphRunFailedEvent,
|
GraphRunFailedEvent,
|
||||||
GraphRunPartialSucceededEvent,
|
GraphRunPartialSucceededEvent,
|
||||||
|
GraphRunPausedEvent,
|
||||||
GraphRunStartedEvent,
|
GraphRunStartedEvent,
|
||||||
GraphRunSucceededEvent,
|
GraphRunSucceededEvent,
|
||||||
NodeRunAgentLogEvent,
|
NodeRunAgentLogEvent,
|
||||||
NodeRunExceptionEvent,
|
NodeRunExceptionEvent,
|
||||||
NodeRunFailedEvent,
|
NodeRunFailedEvent,
|
||||||
|
NodeRunHumanInputFormFilledEvent,
|
||||||
|
NodeRunHumanInputFormTimeoutEvent,
|
||||||
NodeRunIterationFailedEvent,
|
NodeRunIterationFailedEvent,
|
||||||
NodeRunIterationNextEvent,
|
NodeRunIterationNextEvent,
|
||||||
NodeRunIterationStartedEvent,
|
NodeRunIterationStartedEvent,
|
||||||
@@ -61,6 +69,9 @@ from core.workflow.variable_loader import DUMMY_VARIABLE_LOADER, VariableLoader,
|
|||||||
from core.workflow.workflow_entry import WorkflowEntry
|
from core.workflow.workflow_entry import WorkflowEntry
|
||||||
from models.enums import UserFrom
|
from models.enums import UserFrom
|
||||||
from models.workflow import Workflow
|
from models.workflow import Workflow
|
||||||
|
from tasks.mail_human_input_delivery_task import dispatch_human_input_email_task
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
class WorkflowBasedAppRunner:
|
class WorkflowBasedAppRunner:
|
||||||
@@ -327,7 +338,7 @@ class WorkflowBasedAppRunner:
|
|||||||
:param event: event
|
:param event: event
|
||||||
"""
|
"""
|
||||||
if isinstance(event, GraphRunStartedEvent):
|
if isinstance(event, GraphRunStartedEvent):
|
||||||
self._publish_event(QueueWorkflowStartedEvent())
|
self._publish_event(QueueWorkflowStartedEvent(reason=event.reason))
|
||||||
elif isinstance(event, GraphRunSucceededEvent):
|
elif isinstance(event, GraphRunSucceededEvent):
|
||||||
self._publish_event(QueueWorkflowSucceededEvent(outputs=event.outputs))
|
self._publish_event(QueueWorkflowSucceededEvent(outputs=event.outputs))
|
||||||
elif isinstance(event, GraphRunPartialSucceededEvent):
|
elif isinstance(event, GraphRunPartialSucceededEvent):
|
||||||
@@ -338,6 +349,38 @@ class WorkflowBasedAppRunner:
|
|||||||
self._publish_event(QueueWorkflowFailedEvent(error=event.error, exceptions_count=event.exceptions_count))
|
self._publish_event(QueueWorkflowFailedEvent(error=event.error, exceptions_count=event.exceptions_count))
|
||||||
elif isinstance(event, GraphRunAbortedEvent):
|
elif isinstance(event, GraphRunAbortedEvent):
|
||||||
self._publish_event(QueueWorkflowFailedEvent(error=event.reason or "Unknown error", exceptions_count=0))
|
self._publish_event(QueueWorkflowFailedEvent(error=event.reason or "Unknown error", exceptions_count=0))
|
||||||
|
elif isinstance(event, GraphRunPausedEvent):
|
||||||
|
runtime_state = workflow_entry.graph_engine.graph_runtime_state
|
||||||
|
paused_nodes = runtime_state.get_paused_nodes()
|
||||||
|
self._enqueue_human_input_notifications(event.reasons)
|
||||||
|
self._publish_event(
|
||||||
|
QueueWorkflowPausedEvent(
|
||||||
|
reasons=event.reasons,
|
||||||
|
outputs=event.outputs,
|
||||||
|
paused_nodes=paused_nodes,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
elif isinstance(event, NodeRunHumanInputFormFilledEvent):
|
||||||
|
self._publish_event(
|
||||||
|
QueueHumanInputFormFilledEvent(
|
||||||
|
node_execution_id=event.id,
|
||||||
|
node_id=event.node_id,
|
||||||
|
node_type=event.node_type,
|
||||||
|
node_title=event.node_title,
|
||||||
|
rendered_content=event.rendered_content,
|
||||||
|
action_id=event.action_id,
|
||||||
|
action_text=event.action_text,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
elif isinstance(event, NodeRunHumanInputFormTimeoutEvent):
|
||||||
|
self._publish_event(
|
||||||
|
QueueHumanInputFormTimeoutEvent(
|
||||||
|
node_id=event.node_id,
|
||||||
|
node_type=event.node_type,
|
||||||
|
node_title=event.node_title,
|
||||||
|
expiration_time=event.expiration_time,
|
||||||
|
)
|
||||||
|
)
|
||||||
elif isinstance(event, NodeRunRetryEvent):
|
elif isinstance(event, NodeRunRetryEvent):
|
||||||
node_run_result = event.node_run_result
|
node_run_result = event.node_run_result
|
||||||
inputs = node_run_result.inputs
|
inputs = node_run_result.inputs
|
||||||
@@ -544,5 +587,19 @@ class WorkflowBasedAppRunner:
|
|||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def _enqueue_human_input_notifications(self, reasons: Sequence[object]) -> None:
|
||||||
|
for reason in reasons:
|
||||||
|
if not isinstance(reason, HumanInputRequired):
|
||||||
|
continue
|
||||||
|
if not reason.form_id:
|
||||||
|
continue
|
||||||
|
try:
|
||||||
|
dispatch_human_input_email_task.apply_async(
|
||||||
|
kwargs={"form_id": reason.form_id, "node_title": reason.node_title},
|
||||||
|
queue="mail",
|
||||||
|
)
|
||||||
|
except Exception: # pragma: no cover - defensive logging
|
||||||
|
logger.exception("Failed to enqueue human input email task for form %s", reason.form_id)
|
||||||
|
|
||||||
def _publish_event(self, event: AppQueueEvent):
|
def _publish_event(self, event: AppQueueEvent):
|
||||||
self._queue_manager.publish(event, PublishFrom.APPLICATION_MANAGER)
|
self._queue_manager.publish(event, PublishFrom.APPLICATION_MANAGER)
|
||||||
|
|||||||
@@ -132,7 +132,7 @@ class AppGenerateEntity(BaseModel):
|
|||||||
extras: dict[str, Any] = Field(default_factory=dict)
|
extras: dict[str, Any] = Field(default_factory=dict)
|
||||||
|
|
||||||
# tracing instance
|
# tracing instance
|
||||||
trace_manager: Optional["TraceQueueManager"] = None
|
trace_manager: Optional["TraceQueueManager"] = Field(default=None, exclude=True, repr=False)
|
||||||
|
|
||||||
|
|
||||||
class EasyUIBasedAppGenerateEntity(AppGenerateEntity):
|
class EasyUIBasedAppGenerateEntity(AppGenerateEntity):
|
||||||
@@ -156,6 +156,7 @@ class ConversationAppGenerateEntity(AppGenerateEntity):
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
conversation_id: str | None = None
|
conversation_id: str | None = None
|
||||||
|
is_new_conversation: bool = False
|
||||||
parent_message_id: str | None = Field(
|
parent_message_id: str | None = Field(
|
||||||
default=None,
|
default=None,
|
||||||
description=(
|
description=(
|
||||||
|
|||||||
@@ -8,6 +8,8 @@ from pydantic import BaseModel, ConfigDict, Field
|
|||||||
from core.model_runtime.entities.llm_entities import LLMResult, LLMResultChunk
|
from core.model_runtime.entities.llm_entities import LLMResult, LLMResultChunk
|
||||||
from core.rag.entities.citation_metadata import RetrievalSourceMetadata
|
from core.rag.entities.citation_metadata import RetrievalSourceMetadata
|
||||||
from core.workflow.entities import AgentNodeStrategyInit
|
from core.workflow.entities import AgentNodeStrategyInit
|
||||||
|
from core.workflow.entities.pause_reason import PauseReason
|
||||||
|
from core.workflow.entities.workflow_start_reason import WorkflowStartReason
|
||||||
from core.workflow.enums import WorkflowNodeExecutionMetadataKey
|
from core.workflow.enums import WorkflowNodeExecutionMetadataKey
|
||||||
from core.workflow.nodes import NodeType
|
from core.workflow.nodes import NodeType
|
||||||
|
|
||||||
@@ -46,6 +48,9 @@ class QueueEvent(StrEnum):
|
|||||||
PING = "ping"
|
PING = "ping"
|
||||||
STOP = "stop"
|
STOP = "stop"
|
||||||
RETRY = "retry"
|
RETRY = "retry"
|
||||||
|
PAUSE = "pause"
|
||||||
|
HUMAN_INPUT_FORM_FILLED = "human_input_form_filled"
|
||||||
|
HUMAN_INPUT_FORM_TIMEOUT = "human_input_form_timeout"
|
||||||
|
|
||||||
|
|
||||||
class AppQueueEvent(BaseModel):
|
class AppQueueEvent(BaseModel):
|
||||||
@@ -261,6 +266,8 @@ class QueueWorkflowStartedEvent(AppQueueEvent):
|
|||||||
"""QueueWorkflowStartedEvent entity."""
|
"""QueueWorkflowStartedEvent entity."""
|
||||||
|
|
||||||
event: QueueEvent = QueueEvent.WORKFLOW_STARTED
|
event: QueueEvent = QueueEvent.WORKFLOW_STARTED
|
||||||
|
# Always present; mirrors GraphRunStartedEvent.reason for downstream consumers.
|
||||||
|
reason: WorkflowStartReason = WorkflowStartReason.INITIAL
|
||||||
|
|
||||||
|
|
||||||
class QueueWorkflowSucceededEvent(AppQueueEvent):
|
class QueueWorkflowSucceededEvent(AppQueueEvent):
|
||||||
@@ -484,6 +491,35 @@ class QueueStopEvent(AppQueueEvent):
|
|||||||
return reason_mapping.get(self.stopped_by, "Stopped by unknown reason.")
|
return reason_mapping.get(self.stopped_by, "Stopped by unknown reason.")
|
||||||
|
|
||||||
|
|
||||||
|
class QueueHumanInputFormFilledEvent(AppQueueEvent):
|
||||||
|
"""
|
||||||
|
QueueHumanInputFormFilledEvent entity
|
||||||
|
"""
|
||||||
|
|
||||||
|
event: QueueEvent = QueueEvent.HUMAN_INPUT_FORM_FILLED
|
||||||
|
|
||||||
|
node_execution_id: str
|
||||||
|
node_id: str
|
||||||
|
node_type: NodeType
|
||||||
|
node_title: str
|
||||||
|
rendered_content: str
|
||||||
|
action_id: str
|
||||||
|
action_text: str
|
||||||
|
|
||||||
|
|
||||||
|
class QueueHumanInputFormTimeoutEvent(AppQueueEvent):
|
||||||
|
"""
|
||||||
|
QueueHumanInputFormTimeoutEvent entity
|
||||||
|
"""
|
||||||
|
|
||||||
|
event: QueueEvent = QueueEvent.HUMAN_INPUT_FORM_TIMEOUT
|
||||||
|
|
||||||
|
node_id: str
|
||||||
|
node_type: NodeType
|
||||||
|
node_title: str
|
||||||
|
expiration_time: datetime
|
||||||
|
|
||||||
|
|
||||||
class QueueMessage(BaseModel):
|
class QueueMessage(BaseModel):
|
||||||
"""
|
"""
|
||||||
QueueMessage abstract entity
|
QueueMessage abstract entity
|
||||||
@@ -509,3 +545,14 @@ class WorkflowQueueMessage(QueueMessage):
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
pass
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
class QueueWorkflowPausedEvent(AppQueueEvent):
|
||||||
|
"""
|
||||||
|
QueueWorkflowPausedEvent entity
|
||||||
|
"""
|
||||||
|
|
||||||
|
event: QueueEvent = QueueEvent.PAUSE
|
||||||
|
reasons: Sequence[PauseReason] = Field(default_factory=list)
|
||||||
|
outputs: Mapping[str, object] = Field(default_factory=dict)
|
||||||
|
paused_nodes: Sequence[str] = Field(default_factory=list)
|
||||||
|
|||||||
@@ -7,7 +7,9 @@ from pydantic import BaseModel, ConfigDict, Field
|
|||||||
from core.model_runtime.entities.llm_entities import LLMResult, LLMUsage
|
from core.model_runtime.entities.llm_entities import LLMResult, LLMUsage
|
||||||
from core.rag.entities.citation_metadata import RetrievalSourceMetadata
|
from core.rag.entities.citation_metadata import RetrievalSourceMetadata
|
||||||
from core.workflow.entities import AgentNodeStrategyInit
|
from core.workflow.entities import AgentNodeStrategyInit
|
||||||
|
from core.workflow.entities.workflow_start_reason import WorkflowStartReason
|
||||||
from core.workflow.enums import WorkflowExecutionStatus, WorkflowNodeExecutionMetadataKey, WorkflowNodeExecutionStatus
|
from core.workflow.enums import WorkflowExecutionStatus, WorkflowNodeExecutionMetadataKey, WorkflowNodeExecutionStatus
|
||||||
|
from core.workflow.nodes.human_input.entities import FormInput, UserAction
|
||||||
|
|
||||||
|
|
||||||
class AnnotationReplyAccount(BaseModel):
|
class AnnotationReplyAccount(BaseModel):
|
||||||
@@ -69,6 +71,7 @@ class StreamEvent(StrEnum):
|
|||||||
AGENT_THOUGHT = "agent_thought"
|
AGENT_THOUGHT = "agent_thought"
|
||||||
AGENT_MESSAGE = "agent_message"
|
AGENT_MESSAGE = "agent_message"
|
||||||
WORKFLOW_STARTED = "workflow_started"
|
WORKFLOW_STARTED = "workflow_started"
|
||||||
|
WORKFLOW_PAUSED = "workflow_paused"
|
||||||
WORKFLOW_FINISHED = "workflow_finished"
|
WORKFLOW_FINISHED = "workflow_finished"
|
||||||
NODE_STARTED = "node_started"
|
NODE_STARTED = "node_started"
|
||||||
NODE_FINISHED = "node_finished"
|
NODE_FINISHED = "node_finished"
|
||||||
@@ -82,6 +85,9 @@ class StreamEvent(StrEnum):
|
|||||||
TEXT_CHUNK = "text_chunk"
|
TEXT_CHUNK = "text_chunk"
|
||||||
TEXT_REPLACE = "text_replace"
|
TEXT_REPLACE = "text_replace"
|
||||||
AGENT_LOG = "agent_log"
|
AGENT_LOG = "agent_log"
|
||||||
|
HUMAN_INPUT_REQUIRED = "human_input_required"
|
||||||
|
HUMAN_INPUT_FORM_FILLED = "human_input_form_filled"
|
||||||
|
HUMAN_INPUT_FORM_TIMEOUT = "human_input_form_timeout"
|
||||||
|
|
||||||
|
|
||||||
class StreamResponse(BaseModel):
|
class StreamResponse(BaseModel):
|
||||||
@@ -205,6 +211,8 @@ class WorkflowStartStreamResponse(StreamResponse):
|
|||||||
workflow_id: str
|
workflow_id: str
|
||||||
inputs: Mapping[str, Any]
|
inputs: Mapping[str, Any]
|
||||||
created_at: int
|
created_at: int
|
||||||
|
# Always present; mirrors QueueWorkflowStartedEvent.reason for SSE clients.
|
||||||
|
reason: WorkflowStartReason = WorkflowStartReason.INITIAL
|
||||||
|
|
||||||
event: StreamEvent = StreamEvent.WORKFLOW_STARTED
|
event: StreamEvent = StreamEvent.WORKFLOW_STARTED
|
||||||
workflow_run_id: str
|
workflow_run_id: str
|
||||||
@@ -231,7 +239,7 @@ class WorkflowFinishStreamResponse(StreamResponse):
|
|||||||
total_steps: int
|
total_steps: int
|
||||||
created_by: Mapping[str, object] = Field(default_factory=dict)
|
created_by: Mapping[str, object] = Field(default_factory=dict)
|
||||||
created_at: int
|
created_at: int
|
||||||
finished_at: int
|
finished_at: int | None
|
||||||
exceptions_count: int | None = 0
|
exceptions_count: int | None = 0
|
||||||
files: Sequence[Mapping[str, Any]] | None = []
|
files: Sequence[Mapping[str, Any]] | None = []
|
||||||
|
|
||||||
@@ -240,6 +248,85 @@ class WorkflowFinishStreamResponse(StreamResponse):
|
|||||||
data: Data
|
data: Data
|
||||||
|
|
||||||
|
|
||||||
|
class WorkflowPauseStreamResponse(StreamResponse):
|
||||||
|
"""
|
||||||
|
WorkflowPauseStreamResponse entity
|
||||||
|
"""
|
||||||
|
|
||||||
|
class Data(BaseModel):
|
||||||
|
"""
|
||||||
|
Data entity
|
||||||
|
"""
|
||||||
|
|
||||||
|
workflow_run_id: str
|
||||||
|
paused_nodes: Sequence[str] = Field(default_factory=list)
|
||||||
|
outputs: Mapping[str, Any] = Field(default_factory=dict)
|
||||||
|
reasons: Sequence[Mapping[str, Any]] = Field(default_factory=list)
|
||||||
|
status: WorkflowExecutionStatus
|
||||||
|
created_at: int
|
||||||
|
elapsed_time: float
|
||||||
|
total_tokens: int
|
||||||
|
total_steps: int
|
||||||
|
|
||||||
|
event: StreamEvent = StreamEvent.WORKFLOW_PAUSED
|
||||||
|
workflow_run_id: str
|
||||||
|
data: Data
|
||||||
|
|
||||||
|
|
||||||
|
class HumanInputRequiredResponse(StreamResponse):
|
||||||
|
class Data(BaseModel):
|
||||||
|
"""
|
||||||
|
Data entity
|
||||||
|
"""
|
||||||
|
|
||||||
|
form_id: str
|
||||||
|
node_id: str
|
||||||
|
node_title: str
|
||||||
|
form_content: str
|
||||||
|
inputs: Sequence[FormInput] = Field(default_factory=list)
|
||||||
|
actions: Sequence[UserAction] = Field(default_factory=list)
|
||||||
|
display_in_ui: bool = False
|
||||||
|
form_token: str | None = None
|
||||||
|
resolved_default_values: Mapping[str, Any] = Field(default_factory=dict)
|
||||||
|
expiration_time: int = Field(..., description="Unix timestamp in seconds")
|
||||||
|
|
||||||
|
event: StreamEvent = StreamEvent.HUMAN_INPUT_REQUIRED
|
||||||
|
workflow_run_id: str
|
||||||
|
data: Data
|
||||||
|
|
||||||
|
|
||||||
|
class HumanInputFormFilledResponse(StreamResponse):
|
||||||
|
class Data(BaseModel):
|
||||||
|
"""
|
||||||
|
Data entity
|
||||||
|
"""
|
||||||
|
|
||||||
|
node_id: str
|
||||||
|
node_title: str
|
||||||
|
rendered_content: str
|
||||||
|
action_id: str
|
||||||
|
action_text: str
|
||||||
|
|
||||||
|
event: StreamEvent = StreamEvent.HUMAN_INPUT_FORM_FILLED
|
||||||
|
workflow_run_id: str
|
||||||
|
data: Data
|
||||||
|
|
||||||
|
|
||||||
|
class HumanInputFormTimeoutResponse(StreamResponse):
|
||||||
|
class Data(BaseModel):
|
||||||
|
"""
|
||||||
|
Data entity
|
||||||
|
"""
|
||||||
|
|
||||||
|
node_id: str
|
||||||
|
node_title: str
|
||||||
|
expiration_time: int
|
||||||
|
|
||||||
|
event: StreamEvent = StreamEvent.HUMAN_INPUT_FORM_TIMEOUT
|
||||||
|
workflow_run_id: str
|
||||||
|
data: Data
|
||||||
|
|
||||||
|
|
||||||
class NodeStartStreamResponse(StreamResponse):
|
class NodeStartStreamResponse(StreamResponse):
|
||||||
"""
|
"""
|
||||||
NodeStartStreamResponse entity
|
NodeStartStreamResponse entity
|
||||||
@@ -726,7 +813,7 @@ class WorkflowAppBlockingResponse(AppBlockingResponse):
|
|||||||
total_tokens: int
|
total_tokens: int
|
||||||
total_steps: int
|
total_steps: int
|
||||||
created_at: int
|
created_at: int
|
||||||
finished_at: int
|
finished_at: int | None
|
||||||
|
|
||||||
workflow_run_id: str
|
workflow_run_id: str
|
||||||
data: Data
|
data: Data
|
||||||
|
|||||||
@@ -1,3 +1,4 @@
|
|||||||
|
import contextlib
|
||||||
import logging
|
import logging
|
||||||
import time
|
import time
|
||||||
import uuid
|
import uuid
|
||||||
@@ -103,6 +104,14 @@ class RateLimit:
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@contextlib.contextmanager
|
||||||
|
def rate_limit_context(rate_limit: RateLimit, request_id: str | None):
|
||||||
|
request_id = rate_limit.enter(request_id)
|
||||||
|
yield
|
||||||
|
if request_id is not None:
|
||||||
|
rate_limit.exit(request_id)
|
||||||
|
|
||||||
|
|
||||||
class RateLimitGenerator:
|
class RateLimitGenerator:
|
||||||
def __init__(self, rate_limit: RateLimit, generator: Generator[str, None, None], request_id: str):
|
def __init__(self, rate_limit: RateLimit, generator: Generator[str, None, None], request_id: str):
|
||||||
self.rate_limit = rate_limit
|
self.rate_limit = rate_limit
|
||||||
|
|||||||
@@ -1,3 +1,4 @@
|
|||||||
|
from dataclasses import dataclass
|
||||||
from typing import Annotated, Literal, Self, TypeAlias
|
from typing import Annotated, Literal, Self, TypeAlias
|
||||||
|
|
||||||
from pydantic import BaseModel, Field
|
from pydantic import BaseModel, Field
|
||||||
@@ -52,6 +53,14 @@ class WorkflowResumptionContext(BaseModel):
|
|||||||
return self.generate_entity.entity
|
return self.generate_entity.entity
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True)
|
||||||
|
class PauseStateLayerConfig:
|
||||||
|
"""Configuration container for instantiating pause persistence layers."""
|
||||||
|
|
||||||
|
session_factory: Engine | sessionmaker[Session]
|
||||||
|
state_owner_user_id: str
|
||||||
|
|
||||||
|
|
||||||
class PauseStatePersistenceLayer(GraphEngineLayer):
|
class PauseStatePersistenceLayer(GraphEngineLayer):
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
|
|||||||
@@ -45,6 +45,8 @@ from core.app.entities.task_entities import (
|
|||||||
from core.app.task_pipeline.based_generate_task_pipeline import BasedGenerateTaskPipeline
|
from core.app.task_pipeline.based_generate_task_pipeline import BasedGenerateTaskPipeline
|
||||||
from core.app.task_pipeline.message_cycle_manager import MessageCycleManager
|
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.file import helpers as file_helpers
|
||||||
|
from core.file.enums import FileTransferMethod
|
||||||
from core.model_manager import ModelInstance
|
from core.model_manager import ModelInstance
|
||||||
from core.model_runtime.entities.llm_entities import LLMResult, LLMResultChunk, LLMResultChunkDelta, LLMUsage
|
from core.model_runtime.entities.llm_entities import LLMResult, LLMResultChunk, LLMResultChunkDelta, LLMUsage
|
||||||
from core.model_runtime.entities.message_entities import (
|
from core.model_runtime.entities.message_entities import (
|
||||||
@@ -56,10 +58,11 @@ from core.ops.entities.trace_entity import TraceTaskName
|
|||||||
from core.ops.ops_trace_manager import TraceQueueManager, TraceTask
|
from core.ops.ops_trace_manager import TraceQueueManager, TraceTask
|
||||||
from core.prompt.utils.prompt_message_util import PromptMessageUtil
|
from core.prompt.utils.prompt_message_util import PromptMessageUtil
|
||||||
from core.prompt.utils.prompt_template_parser import PromptTemplateParser
|
from core.prompt.utils.prompt_template_parser import PromptTemplateParser
|
||||||
|
from core.tools.signature import sign_tool_file
|
||||||
from events.message_event import message_was_created
|
from events.message_event import message_was_created
|
||||||
from extensions.ext_database import db
|
from extensions.ext_database import db
|
||||||
from libs.datetime_utils import naive_utc_now
|
from libs.datetime_utils import naive_utc_now
|
||||||
from models.model import AppMode, Conversation, Message, MessageAgentThought
|
from models.model import AppMode, Conversation, Message, MessageAgentThought, MessageFile, UploadFile
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
@@ -154,7 +157,7 @@ class EasyUIBasedGenerateTaskPipeline(BasedGenerateTaskPipeline):
|
|||||||
id=self._message_id,
|
id=self._message_id,
|
||||||
mode=self._conversation_mode,
|
mode=self._conversation_mode,
|
||||||
message_id=self._message_id,
|
message_id=self._message_id,
|
||||||
answer=cast(str, self._task_state.llm_result.message.content),
|
answer=self._task_state.llm_result.message.get_text_content(),
|
||||||
created_at=self._message_created_at,
|
created_at=self._message_created_at,
|
||||||
**extras,
|
**extras,
|
||||||
),
|
),
|
||||||
@@ -167,7 +170,7 @@ class EasyUIBasedGenerateTaskPipeline(BasedGenerateTaskPipeline):
|
|||||||
mode=self._conversation_mode,
|
mode=self._conversation_mode,
|
||||||
conversation_id=self._conversation_id,
|
conversation_id=self._conversation_id,
|
||||||
message_id=self._message_id,
|
message_id=self._message_id,
|
||||||
answer=cast(str, self._task_state.llm_result.message.content),
|
answer=self._task_state.llm_result.message.get_text_content(),
|
||||||
created_at=self._message_created_at,
|
created_at=self._message_created_at,
|
||||||
**extras,
|
**extras,
|
||||||
),
|
),
|
||||||
@@ -280,7 +283,7 @@ class EasyUIBasedGenerateTaskPipeline(BasedGenerateTaskPipeline):
|
|||||||
|
|
||||||
# handle output moderation
|
# handle output moderation
|
||||||
output_moderation_answer = self.handle_output_moderation_when_task_finished(
|
output_moderation_answer = self.handle_output_moderation_when_task_finished(
|
||||||
cast(str, self._task_state.llm_result.message.content)
|
self._task_state.llm_result.message.get_text_content()
|
||||||
)
|
)
|
||||||
if output_moderation_answer:
|
if output_moderation_answer:
|
||||||
self._task_state.llm_result.message.content = output_moderation_answer
|
self._task_state.llm_result.message.content = output_moderation_answer
|
||||||
@@ -394,7 +397,7 @@ class EasyUIBasedGenerateTaskPipeline(BasedGenerateTaskPipeline):
|
|||||||
message.message_unit_price = usage.prompt_unit_price
|
message.message_unit_price = usage.prompt_unit_price
|
||||||
message.message_price_unit = usage.prompt_price_unit
|
message.message_price_unit = usage.prompt_price_unit
|
||||||
message.answer = (
|
message.answer = (
|
||||||
PromptTemplateParser.remove_template_variables(cast(str, llm_result.message.content).strip())
|
PromptTemplateParser.remove_template_variables(llm_result.message.get_text_content().strip())
|
||||||
if llm_result.message.content
|
if llm_result.message.content
|
||||||
else ""
|
else ""
|
||||||
)
|
)
|
||||||
@@ -463,6 +466,85 @@ class EasyUIBasedGenerateTaskPipeline(BasedGenerateTaskPipeline):
|
|||||||
metadata=metadata_dict,
|
metadata=metadata_dict,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def _record_files(self):
|
||||||
|
with Session(db.engine, expire_on_commit=False) as session:
|
||||||
|
message_files = session.scalars(select(MessageFile).where(MessageFile.message_id == self._message_id)).all()
|
||||||
|
if not message_files:
|
||||||
|
return None
|
||||||
|
|
||||||
|
files_list = []
|
||||||
|
upload_file_ids = [
|
||||||
|
mf.upload_file_id
|
||||||
|
for mf in message_files
|
||||||
|
if mf.transfer_method == FileTransferMethod.LOCAL_FILE and mf.upload_file_id
|
||||||
|
]
|
||||||
|
upload_files_map = {}
|
||||||
|
if upload_file_ids:
|
||||||
|
upload_files = session.scalars(select(UploadFile).where(UploadFile.id.in_(upload_file_ids))).all()
|
||||||
|
upload_files_map = {uf.id: uf for uf in upload_files}
|
||||||
|
|
||||||
|
for message_file in message_files:
|
||||||
|
upload_file = None
|
||||||
|
if message_file.transfer_method == FileTransferMethod.LOCAL_FILE and message_file.upload_file_id:
|
||||||
|
upload_file = upload_files_map.get(message_file.upload_file_id)
|
||||||
|
|
||||||
|
url = None
|
||||||
|
filename = "file"
|
||||||
|
mime_type = "application/octet-stream"
|
||||||
|
size = 0
|
||||||
|
extension = ""
|
||||||
|
|
||||||
|
if message_file.transfer_method == FileTransferMethod.REMOTE_URL:
|
||||||
|
url = message_file.url
|
||||||
|
if message_file.url:
|
||||||
|
filename = message_file.url.split("/")[-1].split("?")[0] # Remove query params
|
||||||
|
elif message_file.transfer_method == FileTransferMethod.LOCAL_FILE:
|
||||||
|
if upload_file:
|
||||||
|
url = file_helpers.get_signed_file_url(upload_file_id=str(upload_file.id))
|
||||||
|
filename = upload_file.name
|
||||||
|
mime_type = upload_file.mime_type or "application/octet-stream"
|
||||||
|
size = upload_file.size or 0
|
||||||
|
extension = f".{upload_file.extension}" if upload_file.extension else ""
|
||||||
|
elif message_file.upload_file_id:
|
||||||
|
# Fallback: generate URL even if upload_file not found
|
||||||
|
url = file_helpers.get_signed_file_url(upload_file_id=str(message_file.upload_file_id))
|
||||||
|
elif message_file.transfer_method == FileTransferMethod.TOOL_FILE and message_file.url:
|
||||||
|
# For tool files, use URL directly if it's HTTP, otherwise sign it
|
||||||
|
if message_file.url.startswith("http"):
|
||||||
|
url = message_file.url
|
||||||
|
filename = message_file.url.split("/")[-1].split("?")[0]
|
||||||
|
else:
|
||||||
|
# Extract tool file id and extension from URL
|
||||||
|
url_parts = message_file.url.split("/")
|
||||||
|
if url_parts:
|
||||||
|
file_part = url_parts[-1].split("?")[0] # Remove query params first
|
||||||
|
# Use rsplit to correctly handle filenames with multiple dots
|
||||||
|
if "." in file_part:
|
||||||
|
tool_file_id, ext = file_part.rsplit(".", 1)
|
||||||
|
extension = f".{ext}"
|
||||||
|
else:
|
||||||
|
tool_file_id = file_part
|
||||||
|
extension = ".bin"
|
||||||
|
url = sign_tool_file(tool_file_id=tool_file_id, extension=extension)
|
||||||
|
filename = file_part
|
||||||
|
|
||||||
|
transfer_method_value = message_file.transfer_method
|
||||||
|
remote_url = message_file.url if message_file.transfer_method == FileTransferMethod.REMOTE_URL else ""
|
||||||
|
file_dict = {
|
||||||
|
"related_id": message_file.id,
|
||||||
|
"extension": extension,
|
||||||
|
"filename": filename,
|
||||||
|
"size": size,
|
||||||
|
"mime_type": mime_type,
|
||||||
|
"transfer_method": transfer_method_value,
|
||||||
|
"type": message_file.type,
|
||||||
|
"url": url or "",
|
||||||
|
"upload_file_id": message_file.upload_file_id or message_file.id,
|
||||||
|
"remote_url": remote_url,
|
||||||
|
}
|
||||||
|
files_list.append(file_dict)
|
||||||
|
return files_list or None
|
||||||
|
|
||||||
def _agent_message_to_stream_response(self, answer: str, message_id: str) -> AgentMessageStreamResponse:
|
def _agent_message_to_stream_response(self, answer: str, message_id: str) -> AgentMessageStreamResponse:
|
||||||
"""
|
"""
|
||||||
Agent message to stream response.
|
Agent message to stream response.
|
||||||
|
|||||||
@@ -64,7 +64,13 @@ class MessageCycleManager:
|
|||||||
|
|
||||||
# Use SQLAlchemy 2.x style session.scalar(select(...))
|
# Use SQLAlchemy 2.x style session.scalar(select(...))
|
||||||
with session_factory.create_session() as session:
|
with session_factory.create_session() as session:
|
||||||
message_file = session.scalar(select(MessageFile).where(MessageFile.message_id == message_id))
|
message_file = session.scalar(
|
||||||
|
select(MessageFile)
|
||||||
|
.where(
|
||||||
|
MessageFile.message_id == message_id,
|
||||||
|
)
|
||||||
|
.where(MessageFile.belongs_to == "assistant")
|
||||||
|
)
|
||||||
|
|
||||||
if message_file:
|
if message_file:
|
||||||
self._message_has_file.add(message_id)
|
self._message_has_file.add(message_id)
|
||||||
@@ -82,10 +88,11 @@ class MessageCycleManager:
|
|||||||
if isinstance(self._application_generate_entity, CompletionAppGenerateEntity):
|
if isinstance(self._application_generate_entity, CompletionAppGenerateEntity):
|
||||||
return None
|
return None
|
||||||
|
|
||||||
is_first_message = self._application_generate_entity.conversation_id is None
|
is_first_message = self._application_generate_entity.is_new_conversation
|
||||||
extras = self._application_generate_entity.extras
|
extras = self._application_generate_entity.extras
|
||||||
auto_generate_conversation_name = extras.get("auto_generate_conversation_name", True)
|
auto_generate_conversation_name = extras.get("auto_generate_conversation_name", True)
|
||||||
|
|
||||||
|
thread: Thread | None = None
|
||||||
if auto_generate_conversation_name and is_first_message:
|
if auto_generate_conversation_name and is_first_message:
|
||||||
# start generate thread
|
# start generate thread
|
||||||
# time.sleep not block other logic
|
# time.sleep not block other logic
|
||||||
@@ -101,9 +108,10 @@ class MessageCycleManager:
|
|||||||
thread.daemon = True
|
thread.daemon = True
|
||||||
thread.start()
|
thread.start()
|
||||||
|
|
||||||
return thread
|
if is_first_message:
|
||||||
|
self._application_generate_entity.is_new_conversation = False
|
||||||
|
|
||||||
return None
|
return thread
|
||||||
|
|
||||||
def _generate_conversation_name_worker(self, flask_app: Flask, conversation_id: str, query: str):
|
def _generate_conversation_name_worker(self, flask_app: Flask, conversation_id: str, query: str):
|
||||||
with flask_app.app_context():
|
with flask_app.app_context():
|
||||||
|
|||||||
@@ -8,6 +8,7 @@ from core.file.file_manager import file_manager
|
|||||||
from core.helper.code_executor.code_executor import CodeExecutor
|
from core.helper.code_executor.code_executor import CodeExecutor
|
||||||
from core.helper.code_executor.code_node_provider import CodeNodeProvider
|
from core.helper.code_executor.code_node_provider import CodeNodeProvider
|
||||||
from core.helper.ssrf_proxy import ssrf_proxy
|
from core.helper.ssrf_proxy import ssrf_proxy
|
||||||
|
from core.rag.retrieval.dataset_retrieval import DatasetRetrieval
|
||||||
from core.tools.tool_file_manager import ToolFileManager
|
from core.tools.tool_file_manager import ToolFileManager
|
||||||
from core.workflow.entities.graph_config import NodeConfigDict
|
from core.workflow.entities.graph_config import NodeConfigDict
|
||||||
from core.workflow.enums import NodeType
|
from core.workflow.enums import NodeType
|
||||||
@@ -16,6 +17,7 @@ from core.workflow.nodes.base.node import Node
|
|||||||
from core.workflow.nodes.code.code_node import CodeNode
|
from core.workflow.nodes.code.code_node import CodeNode
|
||||||
from core.workflow.nodes.code.limits import CodeNodeLimits
|
from core.workflow.nodes.code.limits import CodeNodeLimits
|
||||||
from core.workflow.nodes.http_request.node import HttpRequestNode
|
from core.workflow.nodes.http_request.node import HttpRequestNode
|
||||||
|
from core.workflow.nodes.knowledge_retrieval.knowledge_retrieval_node import KnowledgeRetrievalNode
|
||||||
from core.workflow.nodes.node_mapping import LATEST_VERSION, NODE_TYPE_CLASSES_MAPPING
|
from core.workflow.nodes.node_mapping import LATEST_VERSION, NODE_TYPE_CLASSES_MAPPING
|
||||||
from core.workflow.nodes.protocols import FileManagerProtocol, HttpClientProtocol
|
from core.workflow.nodes.protocols import FileManagerProtocol, HttpClientProtocol
|
||||||
from core.workflow.nodes.template_transform.template_renderer import (
|
from core.workflow.nodes.template_transform.template_renderer import (
|
||||||
@@ -47,6 +49,7 @@ class DifyNodeFactory(NodeFactory):
|
|||||||
code_providers: Sequence[type[CodeNodeProvider]] | None = None,
|
code_providers: Sequence[type[CodeNodeProvider]] | None = None,
|
||||||
code_limits: CodeNodeLimits | None = None,
|
code_limits: CodeNodeLimits | None = None,
|
||||||
template_renderer: Jinja2TemplateRenderer | None = None,
|
template_renderer: Jinja2TemplateRenderer | None = None,
|
||||||
|
template_transform_max_output_length: int | None = None,
|
||||||
http_request_http_client: HttpClientProtocol | None = None,
|
http_request_http_client: HttpClientProtocol | None = None,
|
||||||
http_request_tool_file_manager_factory: Callable[[], ToolFileManager] = ToolFileManager,
|
http_request_tool_file_manager_factory: Callable[[], ToolFileManager] = ToolFileManager,
|
||||||
http_request_file_manager: FileManagerProtocol | None = None,
|
http_request_file_manager: FileManagerProtocol | None = None,
|
||||||
@@ -68,9 +71,13 @@ class DifyNodeFactory(NodeFactory):
|
|||||||
max_object_array_length=dify_config.CODE_MAX_OBJECT_ARRAY_LENGTH,
|
max_object_array_length=dify_config.CODE_MAX_OBJECT_ARRAY_LENGTH,
|
||||||
)
|
)
|
||||||
self._template_renderer = template_renderer or CodeExecutorJinja2TemplateRenderer()
|
self._template_renderer = template_renderer or CodeExecutorJinja2TemplateRenderer()
|
||||||
|
self._template_transform_max_output_length = (
|
||||||
|
template_transform_max_output_length or dify_config.TEMPLATE_TRANSFORM_MAX_LENGTH
|
||||||
|
)
|
||||||
self._http_request_http_client = http_request_http_client or ssrf_proxy
|
self._http_request_http_client = http_request_http_client or ssrf_proxy
|
||||||
self._http_request_tool_file_manager_factory = http_request_tool_file_manager_factory
|
self._http_request_tool_file_manager_factory = http_request_tool_file_manager_factory
|
||||||
self._http_request_file_manager = http_request_file_manager or file_manager
|
self._http_request_file_manager = http_request_file_manager or file_manager
|
||||||
|
self._rag_retrieval = DatasetRetrieval()
|
||||||
|
|
||||||
@override
|
@override
|
||||||
def create_node(self, node_config: NodeConfigDict) -> Node:
|
def create_node(self, node_config: NodeConfigDict) -> Node:
|
||||||
@@ -122,6 +129,7 @@ class DifyNodeFactory(NodeFactory):
|
|||||||
graph_init_params=self.graph_init_params,
|
graph_init_params=self.graph_init_params,
|
||||||
graph_runtime_state=self.graph_runtime_state,
|
graph_runtime_state=self.graph_runtime_state,
|
||||||
template_renderer=self._template_renderer,
|
template_renderer=self._template_renderer,
|
||||||
|
max_output_length=self._template_transform_max_output_length,
|
||||||
)
|
)
|
||||||
|
|
||||||
if node_type == NodeType.HTTP_REQUEST:
|
if node_type == NodeType.HTTP_REQUEST:
|
||||||
@@ -135,6 +143,15 @@ class DifyNodeFactory(NodeFactory):
|
|||||||
file_manager=self._http_request_file_manager,
|
file_manager=self._http_request_file_manager,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
if node_type == NodeType.KNOWLEDGE_RETRIEVAL:
|
||||||
|
return KnowledgeRetrievalNode(
|
||||||
|
id=node_id,
|
||||||
|
config=node_config,
|
||||||
|
graph_init_params=self.graph_init_params,
|
||||||
|
graph_runtime_state=self.graph_runtime_state,
|
||||||
|
rag_retrieval=self._rag_retrieval,
|
||||||
|
)
|
||||||
|
|
||||||
return node_class(
|
return node_class(
|
||||||
id=node_id,
|
id=node_id,
|
||||||
config=node_config,
|
config=node_config,
|
||||||
|
|||||||
@@ -0,0 +1,54 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from collections.abc import Mapping, Sequence
|
||||||
|
from typing import Any, TypeAlias
|
||||||
|
|
||||||
|
from pydantic import BaseModel, ConfigDict, Field
|
||||||
|
|
||||||
|
from core.workflow.nodes.human_input.entities import FormInput, UserAction
|
||||||
|
from models.execution_extra_content import ExecutionContentType
|
||||||
|
|
||||||
|
|
||||||
|
class HumanInputFormDefinition(BaseModel):
|
||||||
|
model_config = ConfigDict(frozen=True)
|
||||||
|
|
||||||
|
form_id: str
|
||||||
|
node_id: str
|
||||||
|
node_title: str
|
||||||
|
form_content: str
|
||||||
|
inputs: Sequence[FormInput] = Field(default_factory=list)
|
||||||
|
actions: Sequence[UserAction] = Field(default_factory=list)
|
||||||
|
display_in_ui: bool = False
|
||||||
|
form_token: str | None = None
|
||||||
|
resolved_default_values: Mapping[str, Any] = Field(default_factory=dict)
|
||||||
|
expiration_time: int
|
||||||
|
|
||||||
|
|
||||||
|
class HumanInputFormSubmissionData(BaseModel):
|
||||||
|
model_config = ConfigDict(frozen=True)
|
||||||
|
|
||||||
|
node_id: str
|
||||||
|
node_title: str
|
||||||
|
rendered_content: str
|
||||||
|
action_id: str
|
||||||
|
action_text: str
|
||||||
|
|
||||||
|
|
||||||
|
class HumanInputContent(BaseModel):
|
||||||
|
model_config = ConfigDict(frozen=True)
|
||||||
|
|
||||||
|
workflow_run_id: str
|
||||||
|
submitted: bool
|
||||||
|
form_definition: HumanInputFormDefinition | None = None
|
||||||
|
form_submission_data: HumanInputFormSubmissionData | None = None
|
||||||
|
type: ExecutionContentType = Field(default=ExecutionContentType.HUMAN_INPUT)
|
||||||
|
|
||||||
|
|
||||||
|
ExecutionExtraContentDomainModel: TypeAlias = HumanInputContent
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
"ExecutionExtraContentDomainModel",
|
||||||
|
"HumanInputContent",
|
||||||
|
"HumanInputFormDefinition",
|
||||||
|
"HumanInputFormSubmissionData",
|
||||||
|
]
|
||||||
@@ -28,8 +28,8 @@ from core.model_runtime.entities.provider_entities import (
|
|||||||
)
|
)
|
||||||
from core.model_runtime.model_providers.__base.ai_model import AIModel
|
from core.model_runtime.model_providers.__base.ai_model import AIModel
|
||||||
from core.model_runtime.model_providers.model_provider_factory import ModelProviderFactory
|
from core.model_runtime.model_providers.model_provider_factory import ModelProviderFactory
|
||||||
from extensions.ext_database import db
|
|
||||||
from libs.datetime_utils import naive_utc_now
|
from libs.datetime_utils import naive_utc_now
|
||||||
|
from models.engine import db
|
||||||
from models.provider import (
|
from models.provider import (
|
||||||
LoadBalancingModelConfig,
|
LoadBalancingModelConfig,
|
||||||
Provider,
|
Provider,
|
||||||
|
|||||||
@@ -6,7 +6,8 @@ from yarl import URL
|
|||||||
|
|
||||||
from configs import dify_config
|
from configs import dify_config
|
||||||
from core.helper.download import download_with_size_limit
|
from core.helper.download import download_with_size_limit
|
||||||
from core.plugin.entities.marketplace import MarketplacePluginDeclaration
|
from core.plugin.entities.marketplace import MarketplacePluginDeclaration, MarketplacePluginSnapshot
|
||||||
|
from extensions.ext_redis import redis_client
|
||||||
|
|
||||||
marketplace_api_url = URL(str(dify_config.MARKETPLACE_API_URL))
|
marketplace_api_url = URL(str(dify_config.MARKETPLACE_API_URL))
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
@@ -43,28 +44,37 @@ def batch_fetch_plugin_by_ids(plugin_ids: list[str]) -> list[dict]:
|
|||||||
return data.get("data", {}).get("plugins", [])
|
return data.get("data", {}).get("plugins", [])
|
||||||
|
|
||||||
|
|
||||||
def batch_fetch_plugin_manifests_ignore_deserialization_error(
|
|
||||||
plugin_ids: list[str],
|
|
||||||
) -> Sequence[MarketplacePluginDeclaration]:
|
|
||||||
if len(plugin_ids) == 0:
|
|
||||||
return []
|
|
||||||
|
|
||||||
url = str(marketplace_api_url / "api/v1/plugins/batch")
|
|
||||||
response = httpx.post(url, json={"plugin_ids": plugin_ids}, headers={"X-Dify-Version": dify_config.project.version})
|
|
||||||
response.raise_for_status()
|
|
||||||
result: list[MarketplacePluginDeclaration] = []
|
|
||||||
for plugin in response.json()["data"]["plugins"]:
|
|
||||||
try:
|
|
||||||
result.append(MarketplacePluginDeclaration.model_validate(plugin))
|
|
||||||
except Exception:
|
|
||||||
logger.exception(
|
|
||||||
"Failed to deserialize marketplace plugin manifest for %s", plugin.get("plugin_id", "unknown")
|
|
||||||
)
|
|
||||||
|
|
||||||
return result
|
|
||||||
|
|
||||||
|
|
||||||
def record_install_plugin_event(plugin_unique_identifier: str):
|
def record_install_plugin_event(plugin_unique_identifier: str):
|
||||||
url = str(marketplace_api_url / "api/v1/stats/plugins/install_count")
|
url = str(marketplace_api_url / "api/v1/stats/plugins/install_count")
|
||||||
response = httpx.post(url, json={"unique_identifier": plugin_unique_identifier})
|
response = httpx.post(url, json={"unique_identifier": plugin_unique_identifier})
|
||||||
response.raise_for_status()
|
response.raise_for_status()
|
||||||
|
|
||||||
|
|
||||||
|
def fetch_global_plugin_manifest(cache_key_prefix: str, cache_ttl: int) -> None:
|
||||||
|
"""
|
||||||
|
Fetch all plugin manifests from marketplace and cache them in Redis.
|
||||||
|
This should be called once per check cycle to populate the instance-level cache.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
cache_key_prefix: Redis key prefix for caching plugin manifests
|
||||||
|
cache_ttl: Cache TTL in seconds
|
||||||
|
|
||||||
|
Raises:
|
||||||
|
httpx.HTTPError: If the HTTP request fails
|
||||||
|
Exception: If any other error occurs during fetching or caching
|
||||||
|
"""
|
||||||
|
url = str(marketplace_api_url / "api/v1/dist/plugins/manifest.json")
|
||||||
|
response = httpx.get(url, headers={"X-Dify-Version": dify_config.project.version}, timeout=30)
|
||||||
|
response.raise_for_status()
|
||||||
|
|
||||||
|
raw_json = response.json()
|
||||||
|
plugins_data = raw_json.get("plugins", [])
|
||||||
|
|
||||||
|
# Parse and cache all plugin snapshots
|
||||||
|
for plugin_data in plugins_data:
|
||||||
|
plugin_snapshot = MarketplacePluginSnapshot.model_validate(plugin_data)
|
||||||
|
redis_client.setex(
|
||||||
|
name=f"{cache_key_prefix}{plugin_snapshot.plugin_id}",
|
||||||
|
time=cache_ttl,
|
||||||
|
value=plugin_snapshot.model_dump_json(),
|
||||||
|
)
|
||||||
|
|||||||
@@ -15,10 +15,7 @@ from sqlalchemy import select
|
|||||||
from sqlalchemy.orm import Session, sessionmaker
|
from sqlalchemy.orm import Session, sessionmaker
|
||||||
|
|
||||||
from core.helper.encrypter import batch_decrypt_token, encrypt_token, obfuscated_token
|
from core.helper.encrypter import batch_decrypt_token, encrypt_token, obfuscated_token
|
||||||
from core.ops.entities.config_entity import (
|
from core.ops.entities.config_entity import OPS_FILE_PATH, TracingProviderEnum
|
||||||
OPS_FILE_PATH,
|
|
||||||
TracingProviderEnum,
|
|
||||||
)
|
|
||||||
from core.ops.entities.trace_entity import (
|
from core.ops.entities.trace_entity import (
|
||||||
DatasetRetrievalTraceInfo,
|
DatasetRetrievalTraceInfo,
|
||||||
GenerateNameTraceInfo,
|
GenerateNameTraceInfo,
|
||||||
@@ -31,8 +28,8 @@ from core.ops.entities.trace_entity import (
|
|||||||
WorkflowTraceInfo,
|
WorkflowTraceInfo,
|
||||||
)
|
)
|
||||||
from core.ops.utils import get_message_data
|
from core.ops.utils import get_message_data
|
||||||
from extensions.ext_database import db
|
|
||||||
from extensions.ext_storage import storage
|
from extensions.ext_storage import storage
|
||||||
|
from models.engine import db
|
||||||
from models.model import App, AppModelConfig, Conversation, Message, MessageFile, TraceAppConfig
|
from models.model import App, AppModelConfig, Conversation, Message, MessageFile, TraceAppConfig
|
||||||
from models.workflow import WorkflowAppLog
|
from models.workflow import WorkflowAppLog
|
||||||
from tasks.ops_trace_task import process_trace_tasks
|
from tasks.ops_trace_task import process_trace_tasks
|
||||||
@@ -469,6 +466,8 @@ class TraceTask:
|
|||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def _get_workflow_run_repo(cls):
|
def _get_workflow_run_repo(cls):
|
||||||
|
from repositories.factory import DifyAPIRepositoryFactory
|
||||||
|
|
||||||
if cls._workflow_run_repo is None:
|
if cls._workflow_run_repo is None:
|
||||||
with cls._repo_lock:
|
with cls._repo_lock:
|
||||||
if cls._workflow_run_repo is None:
|
if cls._workflow_run_repo is None:
|
||||||
|
|||||||
@@ -5,7 +5,7 @@ from urllib.parse import urlparse
|
|||||||
|
|
||||||
from sqlalchemy import select
|
from sqlalchemy import select
|
||||||
|
|
||||||
from extensions.ext_database import db
|
from models.engine import db
|
||||||
from models.model import Message
|
from models.model import Message
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -1,3 +1,4 @@
|
|||||||
|
import uuid
|
||||||
from collections.abc import Generator, Mapping
|
from collections.abc import Generator, Mapping
|
||||||
from typing import Union
|
from typing import Union
|
||||||
|
|
||||||
@@ -11,6 +12,7 @@ from core.app.apps.chat.app_generator import ChatAppGenerator
|
|||||||
from core.app.apps.completion.app_generator import CompletionAppGenerator
|
from core.app.apps.completion.app_generator import CompletionAppGenerator
|
||||||
from core.app.apps.workflow.app_generator import WorkflowAppGenerator
|
from core.app.apps.workflow.app_generator import WorkflowAppGenerator
|
||||||
from core.app.entities.app_invoke_entities import InvokeFrom
|
from core.app.entities.app_invoke_entities import InvokeFrom
|
||||||
|
from core.app.layers.pause_state_persist_layer import PauseStateLayerConfig
|
||||||
from core.plugin.backwards_invocation.base import BaseBackwardsInvocation
|
from core.plugin.backwards_invocation.base import BaseBackwardsInvocation
|
||||||
from extensions.ext_database import db
|
from extensions.ext_database import db
|
||||||
from models import Account
|
from models import Account
|
||||||
@@ -101,6 +103,11 @@ class PluginAppBackwardsInvocation(BaseBackwardsInvocation):
|
|||||||
if not workflow:
|
if not workflow:
|
||||||
raise ValueError("unexpected app type")
|
raise ValueError("unexpected app type")
|
||||||
|
|
||||||
|
pause_config = PauseStateLayerConfig(
|
||||||
|
session_factory=db.engine,
|
||||||
|
state_owner_user_id=workflow.created_by,
|
||||||
|
)
|
||||||
|
|
||||||
return AdvancedChatAppGenerator().generate(
|
return AdvancedChatAppGenerator().generate(
|
||||||
app_model=app,
|
app_model=app,
|
||||||
workflow=workflow,
|
workflow=workflow,
|
||||||
@@ -112,7 +119,9 @@ class PluginAppBackwardsInvocation(BaseBackwardsInvocation):
|
|||||||
"conversation_id": conversation_id,
|
"conversation_id": conversation_id,
|
||||||
},
|
},
|
||||||
invoke_from=InvokeFrom.SERVICE_API,
|
invoke_from=InvokeFrom.SERVICE_API,
|
||||||
|
workflow_run_id=str(uuid.uuid4()),
|
||||||
streaming=stream,
|
streaming=stream,
|
||||||
|
pause_state_config=pause_config,
|
||||||
)
|
)
|
||||||
elif app.mode == AppMode.AGENT_CHAT:
|
elif app.mode == AppMode.AGENT_CHAT:
|
||||||
return AgentChatAppGenerator().generate(
|
return AgentChatAppGenerator().generate(
|
||||||
@@ -159,6 +168,11 @@ class PluginAppBackwardsInvocation(BaseBackwardsInvocation):
|
|||||||
if not workflow:
|
if not workflow:
|
||||||
raise ValueError("unexpected app type")
|
raise ValueError("unexpected app type")
|
||||||
|
|
||||||
|
pause_config = PauseStateLayerConfig(
|
||||||
|
session_factory=db.engine,
|
||||||
|
state_owner_user_id=workflow.created_by,
|
||||||
|
)
|
||||||
|
|
||||||
return WorkflowAppGenerator().generate(
|
return WorkflowAppGenerator().generate(
|
||||||
app_model=app,
|
app_model=app,
|
||||||
workflow=workflow,
|
workflow=workflow,
|
||||||
@@ -167,6 +181,7 @@ class PluginAppBackwardsInvocation(BaseBackwardsInvocation):
|
|||||||
invoke_from=InvokeFrom.SERVICE_API,
|
invoke_from=InvokeFrom.SERVICE_API,
|
||||||
streaming=stream,
|
streaming=stream,
|
||||||
call_depth=1,
|
call_depth=1,
|
||||||
|
pause_state_config=pause_config,
|
||||||
)
|
)
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
|
|||||||
@@ -1,4 +1,4 @@
|
|||||||
from pydantic import BaseModel, Field, model_validator
|
from pydantic import BaseModel, Field, computed_field, model_validator
|
||||||
|
|
||||||
from core.model_runtime.entities.provider_entities import ProviderEntity
|
from core.model_runtime.entities.provider_entities import ProviderEntity
|
||||||
from core.plugin.entities.endpoint import EndpointProviderDeclaration
|
from core.plugin.entities.endpoint import EndpointProviderDeclaration
|
||||||
@@ -48,3 +48,15 @@ class MarketplacePluginDeclaration(BaseModel):
|
|||||||
if "tool" in data and not data["tool"]:
|
if "tool" in data and not data["tool"]:
|
||||||
del data["tool"]
|
del data["tool"]
|
||||||
return data
|
return data
|
||||||
|
|
||||||
|
|
||||||
|
class MarketplacePluginSnapshot(BaseModel):
|
||||||
|
org: str
|
||||||
|
name: str
|
||||||
|
latest_version: str
|
||||||
|
latest_package_identifier: str
|
||||||
|
latest_package_url: str
|
||||||
|
|
||||||
|
@computed_field
|
||||||
|
def plugin_id(self) -> str:
|
||||||
|
return f"{self.org}/{self.name}"
|
||||||
|
|||||||
@@ -1,13 +1,15 @@
|
|||||||
import json
|
import json
|
||||||
|
import logging
|
||||||
import math
|
import math
|
||||||
import re
|
import re
|
||||||
import threading
|
import threading
|
||||||
|
import time
|
||||||
from collections import Counter, defaultdict
|
from collections import Counter, defaultdict
|
||||||
from collections.abc import Generator, Mapping
|
from collections.abc import Generator, Mapping
|
||||||
from typing import Any, Union, cast
|
from typing import Any, Union, cast
|
||||||
|
|
||||||
from flask import Flask, current_app
|
from flask import Flask, current_app
|
||||||
from sqlalchemy import and_, literal, or_, select
|
from sqlalchemy import and_, func, literal, or_, select
|
||||||
from sqlalchemy.orm import Session
|
from sqlalchemy.orm import Session
|
||||||
|
|
||||||
from core.app.app_config.entities import (
|
from core.app.app_config.entities import (
|
||||||
@@ -18,6 +20,7 @@ from core.app.app_config.entities import (
|
|||||||
)
|
)
|
||||||
from core.app.entities.app_invoke_entities import InvokeFrom, ModelConfigWithCredentialsEntity
|
from core.app.entities.app_invoke_entities import InvokeFrom, ModelConfigWithCredentialsEntity
|
||||||
from core.callback_handler.index_tool_callback_handler import DatasetIndexToolCallbackHandler
|
from core.callback_handler.index_tool_callback_handler import DatasetIndexToolCallbackHandler
|
||||||
|
from core.db.session_factory import session_factory
|
||||||
from core.entities.agent_entities import PlanningStrategy
|
from core.entities.agent_entities import PlanningStrategy
|
||||||
from core.entities.model_entities import ModelStatus
|
from core.entities.model_entities import ModelStatus
|
||||||
from core.file import File, FileTransferMethod, FileType
|
from core.file import File, FileTransferMethod, FileType
|
||||||
@@ -58,12 +61,30 @@ from core.rag.retrieval.template_prompts import (
|
|||||||
)
|
)
|
||||||
from core.tools.signature import sign_upload_file
|
from core.tools.signature import sign_upload_file
|
||||||
from core.tools.utils.dataset_retriever.dataset_retriever_base_tool import DatasetRetrieverBaseTool
|
from core.tools.utils.dataset_retriever.dataset_retriever_base_tool import DatasetRetrieverBaseTool
|
||||||
|
from core.workflow.nodes.knowledge_retrieval import exc
|
||||||
|
from core.workflow.repositories.rag_retrieval_protocol import (
|
||||||
|
KnowledgeRetrievalRequest,
|
||||||
|
Source,
|
||||||
|
SourceChildChunk,
|
||||||
|
SourceMetadata,
|
||||||
|
)
|
||||||
from extensions.ext_database import db
|
from extensions.ext_database import db
|
||||||
|
from extensions.ext_redis import redis_client
|
||||||
from libs.json_in_md_parser import parse_and_check_json_markdown
|
from libs.json_in_md_parser import parse_and_check_json_markdown
|
||||||
from models import UploadFile
|
from models import UploadFile
|
||||||
from models.dataset import ChildChunk, Dataset, DatasetMetadata, DatasetQuery, DocumentSegment, SegmentAttachmentBinding
|
from models.dataset import (
|
||||||
|
ChildChunk,
|
||||||
|
Dataset,
|
||||||
|
DatasetMetadata,
|
||||||
|
DatasetQuery,
|
||||||
|
DocumentSegment,
|
||||||
|
RateLimitLog,
|
||||||
|
SegmentAttachmentBinding,
|
||||||
|
)
|
||||||
from models.dataset import Document as DatasetDocument
|
from models.dataset import Document as DatasetDocument
|
||||||
|
from models.dataset import Document as DocumentModel
|
||||||
from services.external_knowledge_service import ExternalDatasetService
|
from services.external_knowledge_service import ExternalDatasetService
|
||||||
|
from services.feature_service import FeatureService
|
||||||
|
|
||||||
default_retrieval_model: dict[str, Any] = {
|
default_retrieval_model: dict[str, Any] = {
|
||||||
"search_method": RetrievalMethod.SEMANTIC_SEARCH,
|
"search_method": RetrievalMethod.SEMANTIC_SEARCH,
|
||||||
@@ -73,6 +94,8 @@ default_retrieval_model: dict[str, Any] = {
|
|||||||
"score_threshold_enabled": False,
|
"score_threshold_enabled": False,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
class DatasetRetrieval:
|
class DatasetRetrieval:
|
||||||
def __init__(self, application_generate_entity=None):
|
def __init__(self, application_generate_entity=None):
|
||||||
@@ -91,6 +114,233 @@ class DatasetRetrieval:
|
|||||||
else:
|
else:
|
||||||
self._llm_usage = self._llm_usage.plus(usage)
|
self._llm_usage = self._llm_usage.plus(usage)
|
||||||
|
|
||||||
|
def knowledge_retrieval(self, request: KnowledgeRetrievalRequest) -> list[Source]:
|
||||||
|
self._check_knowledge_rate_limit(request.tenant_id)
|
||||||
|
available_datasets = self._get_available_datasets(request.tenant_id, request.dataset_ids)
|
||||||
|
available_datasets_ids = [i.id for i in available_datasets]
|
||||||
|
if not available_datasets_ids:
|
||||||
|
return []
|
||||||
|
|
||||||
|
if not request.query:
|
||||||
|
return []
|
||||||
|
|
||||||
|
metadata_filter_document_ids, metadata_condition = None, None
|
||||||
|
|
||||||
|
if request.metadata_filtering_mode != "disabled":
|
||||||
|
# Convert workflow layer types to app_config layer types
|
||||||
|
if not request.metadata_model_config:
|
||||||
|
raise ValueError("metadata_model_config is required for this method")
|
||||||
|
|
||||||
|
app_metadata_model_config = ModelConfig.model_validate(request.metadata_model_config.model_dump())
|
||||||
|
|
||||||
|
app_metadata_filtering_conditions = None
|
||||||
|
if request.metadata_filtering_conditions is not None:
|
||||||
|
app_metadata_filtering_conditions = MetadataFilteringCondition.model_validate(
|
||||||
|
request.metadata_filtering_conditions.model_dump()
|
||||||
|
)
|
||||||
|
|
||||||
|
query = request.query if request.query is not None else ""
|
||||||
|
|
||||||
|
metadata_filter_document_ids, metadata_condition = self.get_metadata_filter_condition(
|
||||||
|
dataset_ids=available_datasets_ids,
|
||||||
|
query=query,
|
||||||
|
tenant_id=request.tenant_id,
|
||||||
|
user_id=request.user_id,
|
||||||
|
metadata_filtering_mode=request.metadata_filtering_mode,
|
||||||
|
metadata_model_config=app_metadata_model_config,
|
||||||
|
metadata_filtering_conditions=app_metadata_filtering_conditions,
|
||||||
|
inputs={},
|
||||||
|
)
|
||||||
|
|
||||||
|
if request.retrieval_mode == DatasetRetrieveConfigEntity.RetrieveStrategy.SINGLE:
|
||||||
|
planning_strategy = PlanningStrategy.REACT_ROUTER
|
||||||
|
# Ensure required fields are not None for single retrieval mode
|
||||||
|
if request.model_provider is None or request.model_name is None or request.query is None:
|
||||||
|
raise ValueError("model_provider, model_name, and query are required for single retrieval mode")
|
||||||
|
|
||||||
|
model_manager = ModelManager()
|
||||||
|
model_instance = model_manager.get_model_instance(
|
||||||
|
tenant_id=request.tenant_id,
|
||||||
|
model_type=ModelType.LLM,
|
||||||
|
provider=request.model_provider,
|
||||||
|
model=request.model_name,
|
||||||
|
)
|
||||||
|
|
||||||
|
provider_model_bundle = model_instance.provider_model_bundle
|
||||||
|
model_type_instance = model_instance.model_type_instance
|
||||||
|
model_type_instance = cast(LargeLanguageModel, model_type_instance)
|
||||||
|
|
||||||
|
model_credentials = model_instance.credentials
|
||||||
|
|
||||||
|
# check model
|
||||||
|
provider_model = provider_model_bundle.configuration.get_provider_model(
|
||||||
|
model=request.model_name, model_type=ModelType.LLM
|
||||||
|
)
|
||||||
|
|
||||||
|
if provider_model is None:
|
||||||
|
raise exc.ModelNotExistError(f"Model {request.model_name} not exist.")
|
||||||
|
|
||||||
|
if provider_model.status == ModelStatus.NO_CONFIGURE:
|
||||||
|
raise exc.ModelCredentialsNotInitializedError(
|
||||||
|
f"Model {request.model_name} credentials is not initialized."
|
||||||
|
)
|
||||||
|
elif provider_model.status == ModelStatus.NO_PERMISSION:
|
||||||
|
raise exc.ModelNotSupportedError(f"Dify Hosted OpenAI {request.model_name} currently not support.")
|
||||||
|
elif provider_model.status == ModelStatus.QUOTA_EXCEEDED:
|
||||||
|
raise exc.ModelQuotaExceededError(f"Model provider {request.model_provider} quota exceeded.")
|
||||||
|
|
||||||
|
stop = []
|
||||||
|
completion_params = (request.completion_params or {}).copy()
|
||||||
|
if "stop" in completion_params:
|
||||||
|
stop = completion_params["stop"]
|
||||||
|
del completion_params["stop"]
|
||||||
|
|
||||||
|
model_schema = model_type_instance.get_model_schema(request.model_name, model_credentials)
|
||||||
|
|
||||||
|
if not model_schema:
|
||||||
|
raise exc.ModelNotExistError(f"Model {request.model_name} not exist.")
|
||||||
|
|
||||||
|
model_config = ModelConfigWithCredentialsEntity(
|
||||||
|
provider=request.model_provider,
|
||||||
|
model=request.model_name,
|
||||||
|
model_schema=model_schema,
|
||||||
|
mode=request.model_mode or "chat",
|
||||||
|
provider_model_bundle=provider_model_bundle,
|
||||||
|
credentials=model_credentials,
|
||||||
|
parameters=completion_params,
|
||||||
|
stop=stop,
|
||||||
|
)
|
||||||
|
all_documents = self.single_retrieve(
|
||||||
|
request.app_id,
|
||||||
|
request.tenant_id,
|
||||||
|
request.user_id,
|
||||||
|
request.user_from,
|
||||||
|
request.query,
|
||||||
|
available_datasets,
|
||||||
|
model_instance,
|
||||||
|
model_config,
|
||||||
|
planning_strategy,
|
||||||
|
None, # message_id
|
||||||
|
metadata_filter_document_ids,
|
||||||
|
metadata_condition,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
all_documents = self.multiple_retrieve(
|
||||||
|
app_id=request.app_id,
|
||||||
|
tenant_id=request.tenant_id,
|
||||||
|
user_id=request.user_id,
|
||||||
|
user_from=request.user_from,
|
||||||
|
available_datasets=available_datasets,
|
||||||
|
query=request.query,
|
||||||
|
top_k=request.top_k,
|
||||||
|
score_threshold=request.score_threshold,
|
||||||
|
reranking_mode=request.reranking_mode,
|
||||||
|
reranking_model=request.reranking_model,
|
||||||
|
weights=request.weights,
|
||||||
|
reranking_enable=request.reranking_enable,
|
||||||
|
metadata_filter_document_ids=metadata_filter_document_ids,
|
||||||
|
metadata_condition=metadata_condition,
|
||||||
|
attachment_ids=request.attachment_ids,
|
||||||
|
)
|
||||||
|
|
||||||
|
dify_documents = [item for item in all_documents if item.provider == "dify"]
|
||||||
|
external_documents = [item for item in all_documents if item.provider == "external"]
|
||||||
|
retrieval_resource_list = []
|
||||||
|
# deal with external documents
|
||||||
|
for item in external_documents:
|
||||||
|
source = Source(
|
||||||
|
metadata=SourceMetadata(
|
||||||
|
source="knowledge",
|
||||||
|
dataset_id=item.metadata.get("dataset_id"),
|
||||||
|
dataset_name=item.metadata.get("dataset_name"),
|
||||||
|
document_id=item.metadata.get("document_id"),
|
||||||
|
document_name=item.metadata.get("title"),
|
||||||
|
data_source_type="external",
|
||||||
|
retriever_from="workflow",
|
||||||
|
score=item.metadata.get("score"),
|
||||||
|
doc_metadata=item.metadata,
|
||||||
|
),
|
||||||
|
title=item.metadata.get("title"),
|
||||||
|
content=item.page_content,
|
||||||
|
)
|
||||||
|
retrieval_resource_list.append(source)
|
||||||
|
# deal with dify documents
|
||||||
|
if dify_documents:
|
||||||
|
records = RetrievalService.format_retrieval_documents(dify_documents)
|
||||||
|
dataset_ids = [i.segment.dataset_id for i in records]
|
||||||
|
document_ids = [i.segment.document_id for i in records]
|
||||||
|
|
||||||
|
with session_factory.create_session() as session:
|
||||||
|
datasets = session.query(Dataset).where(Dataset.id.in_(dataset_ids)).all()
|
||||||
|
documents = session.query(DatasetDocument).where(DatasetDocument.id.in_(document_ids)).all()
|
||||||
|
|
||||||
|
dataset_map = {i.id: i for i in datasets}
|
||||||
|
document_map = {i.id: i for i in documents}
|
||||||
|
|
||||||
|
if records:
|
||||||
|
for record in records:
|
||||||
|
segment = record.segment
|
||||||
|
dataset = dataset_map.get(segment.dataset_id)
|
||||||
|
document = document_map.get(segment.document_id)
|
||||||
|
|
||||||
|
if dataset and document:
|
||||||
|
source = Source(
|
||||||
|
metadata=SourceMetadata(
|
||||||
|
source="knowledge",
|
||||||
|
dataset_id=dataset.id,
|
||||||
|
dataset_name=dataset.name,
|
||||||
|
document_id=document.id,
|
||||||
|
document_name=document.name,
|
||||||
|
data_source_type=document.data_source_type,
|
||||||
|
segment_id=segment.id,
|
||||||
|
retriever_from="workflow",
|
||||||
|
score=record.score or 0.0,
|
||||||
|
segment_hit_count=segment.hit_count,
|
||||||
|
segment_word_count=segment.word_count,
|
||||||
|
segment_position=segment.position,
|
||||||
|
segment_index_node_hash=segment.index_node_hash,
|
||||||
|
doc_metadata=document.doc_metadata,
|
||||||
|
child_chunks=[
|
||||||
|
SourceChildChunk(
|
||||||
|
id=str(getattr(chunk, "id", "")),
|
||||||
|
content=str(getattr(chunk, "content", "")),
|
||||||
|
position=int(getattr(chunk, "position", 0)),
|
||||||
|
score=float(getattr(chunk, "score", 0.0)),
|
||||||
|
)
|
||||||
|
for chunk in (record.child_chunks or [])
|
||||||
|
],
|
||||||
|
position=None,
|
||||||
|
),
|
||||||
|
title=document.name,
|
||||||
|
files=list(record.files) if record.files else None,
|
||||||
|
content=segment.get_sign_content(),
|
||||||
|
)
|
||||||
|
if segment.answer:
|
||||||
|
source.content = f"question:{segment.get_sign_content()} \nanswer:{segment.answer}"
|
||||||
|
|
||||||
|
if record.summary:
|
||||||
|
source.summary = record.summary
|
||||||
|
|
||||||
|
retrieval_resource_list.append(source)
|
||||||
|
|
||||||
|
if retrieval_resource_list:
|
||||||
|
|
||||||
|
def _score(item: Source) -> float:
|
||||||
|
meta = item.metadata
|
||||||
|
score = meta.score
|
||||||
|
if isinstance(score, (int, float)):
|
||||||
|
return float(score)
|
||||||
|
return 0.0
|
||||||
|
|
||||||
|
retrieval_resource_list = sorted(
|
||||||
|
retrieval_resource_list,
|
||||||
|
key=_score, # type: ignore[arg-type, return-value]
|
||||||
|
reverse=True,
|
||||||
|
)
|
||||||
|
for position, item in enumerate(retrieval_resource_list, start=1):
|
||||||
|
item.metadata.position = position # type: ignore[index]
|
||||||
|
return retrieval_resource_list
|
||||||
|
|
||||||
def retrieve(
|
def retrieve(
|
||||||
self,
|
self,
|
||||||
app_id: str,
|
app_id: str,
|
||||||
@@ -150,14 +400,7 @@ class DatasetRetrieval:
|
|||||||
if features:
|
if features:
|
||||||
if ModelFeature.TOOL_CALL in features or ModelFeature.MULTI_TOOL_CALL in features:
|
if ModelFeature.TOOL_CALL in features or ModelFeature.MULTI_TOOL_CALL in features:
|
||||||
planning_strategy = PlanningStrategy.ROUTER
|
planning_strategy = PlanningStrategy.ROUTER
|
||||||
available_datasets = []
|
available_datasets = self._get_available_datasets(tenant_id, dataset_ids)
|
||||||
|
|
||||||
dataset_stmt = select(Dataset).where(Dataset.tenant_id == tenant_id, Dataset.id.in_(dataset_ids))
|
|
||||||
datasets: list[Dataset] = db.session.execute(dataset_stmt).scalars().all() # type: ignore
|
|
||||||
for dataset in datasets:
|
|
||||||
if dataset.available_document_count == 0 and dataset.provider != "external":
|
|
||||||
continue
|
|
||||||
available_datasets.append(dataset)
|
|
||||||
|
|
||||||
if inputs:
|
if inputs:
|
||||||
inputs = {key: str(value) for key, value in inputs.items()}
|
inputs = {key: str(value) for key, value in inputs.items()}
|
||||||
@@ -1161,7 +1404,6 @@ class DatasetRetrieval:
|
|||||||
query=query or "",
|
query=query or "",
|
||||||
)
|
)
|
||||||
|
|
||||||
result_text = ""
|
|
||||||
try:
|
try:
|
||||||
# handle invoke result
|
# handle invoke result
|
||||||
invoke_result = cast(
|
invoke_result = cast(
|
||||||
@@ -1192,7 +1434,8 @@ class DatasetRetrieval:
|
|||||||
"condition": item.get("comparison_operator"),
|
"condition": item.get("comparison_operator"),
|
||||||
}
|
}
|
||||||
)
|
)
|
||||||
except Exception:
|
except Exception as e:
|
||||||
|
logger.warning(e, exc_info=True)
|
||||||
return None
|
return None
|
||||||
return automatic_metadata_filters
|
return automatic_metadata_filters
|
||||||
|
|
||||||
@@ -1406,7 +1649,12 @@ class DatasetRetrieval:
|
|||||||
usage = None
|
usage = None
|
||||||
for result in invoke_result:
|
for result in invoke_result:
|
||||||
text = result.delta.message.content
|
text = result.delta.message.content
|
||||||
full_text += text
|
if isinstance(text, str):
|
||||||
|
full_text += text
|
||||||
|
elif isinstance(text, list):
|
||||||
|
for i in text:
|
||||||
|
if i.data:
|
||||||
|
full_text += i.data
|
||||||
|
|
||||||
if not model:
|
if not model:
|
||||||
model = result.model
|
model = result.model
|
||||||
@@ -1524,3 +1772,53 @@ class DatasetRetrieval:
|
|||||||
cancel_event.set()
|
cancel_event.set()
|
||||||
if thread_exceptions is not None:
|
if thread_exceptions is not None:
|
||||||
thread_exceptions.append(e)
|
thread_exceptions.append(e)
|
||||||
|
|
||||||
|
def _get_available_datasets(self, tenant_id: str, dataset_ids: list[str]) -> list[Dataset]:
|
||||||
|
with session_factory.create_session() as session:
|
||||||
|
subquery = (
|
||||||
|
session.query(DocumentModel.dataset_id, func.count(DocumentModel.id).label("available_document_count"))
|
||||||
|
.where(
|
||||||
|
DocumentModel.indexing_status == "completed",
|
||||||
|
DocumentModel.enabled == True,
|
||||||
|
DocumentModel.archived == False,
|
||||||
|
DocumentModel.dataset_id.in_(dataset_ids),
|
||||||
|
)
|
||||||
|
.group_by(DocumentModel.dataset_id)
|
||||||
|
.having(func.count(DocumentModel.id) > 0)
|
||||||
|
.subquery()
|
||||||
|
)
|
||||||
|
|
||||||
|
results = (
|
||||||
|
session.query(Dataset)
|
||||||
|
.outerjoin(subquery, Dataset.id == subquery.c.dataset_id)
|
||||||
|
.where(Dataset.tenant_id == tenant_id, Dataset.id.in_(dataset_ids))
|
||||||
|
.where((subquery.c.available_document_count > 0) | (Dataset.provider == "external"))
|
||||||
|
.all()
|
||||||
|
)
|
||||||
|
|
||||||
|
available_datasets = []
|
||||||
|
for dataset in results:
|
||||||
|
if not dataset:
|
||||||
|
continue
|
||||||
|
available_datasets.append(dataset)
|
||||||
|
return available_datasets
|
||||||
|
|
||||||
|
def _check_knowledge_rate_limit(self, tenant_id: str):
|
||||||
|
knowledge_rate_limit = FeatureService.get_knowledge_rate_limit(tenant_id)
|
||||||
|
if knowledge_rate_limit.enabled:
|
||||||
|
current_time = int(time.time() * 1000)
|
||||||
|
key = f"rate_limit_{tenant_id}"
|
||||||
|
redis_client.zadd(key, {current_time: current_time})
|
||||||
|
redis_client.zremrangebyscore(key, 0, current_time - 60000)
|
||||||
|
request_count = redis_client.zcard(key)
|
||||||
|
if request_count > knowledge_rate_limit.limit:
|
||||||
|
with session_factory.create_session() as session:
|
||||||
|
rate_limit_log = RateLimitLog(
|
||||||
|
tenant_id=tenant_id,
|
||||||
|
subscription_plan=knowledge_rate_limit.subscription_plan,
|
||||||
|
operation="knowledge",
|
||||||
|
)
|
||||||
|
session.add(rate_limit_log)
|
||||||
|
raise exc.RateLimitExceededError(
|
||||||
|
"you have reached the knowledge base request rate limit of your subscription."
|
||||||
|
)
|
||||||
|
|||||||
@@ -1,19 +1,18 @@
|
|||||||
"""
|
"""Repository implementations for data access."""
|
||||||
Repository implementations for data access.
|
|
||||||
|
|
||||||
This package contains concrete implementations of the repository interfaces
|
from __future__ import annotations
|
||||||
defined in the core.workflow.repository package.
|
|
||||||
"""
|
|
||||||
|
|
||||||
from core.repositories.celery_workflow_execution_repository import CeleryWorkflowExecutionRepository
|
from .celery_workflow_execution_repository import CeleryWorkflowExecutionRepository
|
||||||
from core.repositories.celery_workflow_node_execution_repository import CeleryWorkflowNodeExecutionRepository
|
from .celery_workflow_node_execution_repository import CeleryWorkflowNodeExecutionRepository
|
||||||
from core.repositories.factory import DifyCoreRepositoryFactory, RepositoryImportError
|
from .factory import DifyCoreRepositoryFactory, RepositoryImportError
|
||||||
from core.repositories.sqlalchemy_workflow_node_execution_repository import SQLAlchemyWorkflowNodeExecutionRepository
|
from .sqlalchemy_workflow_execution_repository import SQLAlchemyWorkflowExecutionRepository
|
||||||
|
from .sqlalchemy_workflow_node_execution_repository import SQLAlchemyWorkflowNodeExecutionRepository
|
||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
"CeleryWorkflowExecutionRepository",
|
"CeleryWorkflowExecutionRepository",
|
||||||
"CeleryWorkflowNodeExecutionRepository",
|
"CeleryWorkflowNodeExecutionRepository",
|
||||||
"DifyCoreRepositoryFactory",
|
"DifyCoreRepositoryFactory",
|
||||||
"RepositoryImportError",
|
"RepositoryImportError",
|
||||||
|
"SQLAlchemyWorkflowExecutionRepository",
|
||||||
"SQLAlchemyWorkflowNodeExecutionRepository",
|
"SQLAlchemyWorkflowNodeExecutionRepository",
|
||||||
]
|
]
|
||||||
|
|||||||
@@ -0,0 +1,553 @@
|
|||||||
|
import dataclasses
|
||||||
|
import json
|
||||||
|
from collections.abc import Mapping, Sequence
|
||||||
|
from datetime import datetime
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
from sqlalchemy import Engine, select
|
||||||
|
from sqlalchemy.orm import Session, selectinload, sessionmaker
|
||||||
|
|
||||||
|
from core.workflow.nodes.human_input.entities import (
|
||||||
|
DeliveryChannelConfig,
|
||||||
|
EmailDeliveryMethod,
|
||||||
|
EmailRecipients,
|
||||||
|
ExternalRecipient,
|
||||||
|
FormDefinition,
|
||||||
|
HumanInputNodeData,
|
||||||
|
MemberRecipient,
|
||||||
|
WebAppDeliveryMethod,
|
||||||
|
)
|
||||||
|
from core.workflow.nodes.human_input.enums import (
|
||||||
|
DeliveryMethodType,
|
||||||
|
HumanInputFormKind,
|
||||||
|
HumanInputFormStatus,
|
||||||
|
)
|
||||||
|
from core.workflow.repositories.human_input_form_repository import (
|
||||||
|
FormCreateParams,
|
||||||
|
FormNotFoundError,
|
||||||
|
HumanInputFormEntity,
|
||||||
|
HumanInputFormRecipientEntity,
|
||||||
|
)
|
||||||
|
from libs.datetime_utils import naive_utc_now
|
||||||
|
from libs.uuid_utils import uuidv7
|
||||||
|
from models.account import Account, TenantAccountJoin
|
||||||
|
from models.human_input import (
|
||||||
|
BackstageRecipientPayload,
|
||||||
|
ConsoleDeliveryPayload,
|
||||||
|
ConsoleRecipientPayload,
|
||||||
|
EmailExternalRecipientPayload,
|
||||||
|
EmailMemberRecipientPayload,
|
||||||
|
HumanInputDelivery,
|
||||||
|
HumanInputForm,
|
||||||
|
HumanInputFormRecipient,
|
||||||
|
RecipientType,
|
||||||
|
StandaloneWebAppRecipientPayload,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclasses.dataclass(frozen=True)
|
||||||
|
class _DeliveryAndRecipients:
|
||||||
|
delivery: HumanInputDelivery
|
||||||
|
recipients: Sequence[HumanInputFormRecipient]
|
||||||
|
|
||||||
|
|
||||||
|
@dataclasses.dataclass(frozen=True)
|
||||||
|
class _WorkspaceMemberInfo:
|
||||||
|
user_id: str
|
||||||
|
email: str
|
||||||
|
|
||||||
|
|
||||||
|
class _HumanInputFormRecipientEntityImpl(HumanInputFormRecipientEntity):
|
||||||
|
def __init__(self, recipient_model: HumanInputFormRecipient):
|
||||||
|
self._recipient_model = recipient_model
|
||||||
|
|
||||||
|
@property
|
||||||
|
def id(self) -> str:
|
||||||
|
return self._recipient_model.id
|
||||||
|
|
||||||
|
@property
|
||||||
|
def token(self) -> str:
|
||||||
|
if self._recipient_model.access_token is None:
|
||||||
|
raise AssertionError(f"access_token should not be None for recipient {self._recipient_model.id}")
|
||||||
|
return self._recipient_model.access_token
|
||||||
|
|
||||||
|
|
||||||
|
class _HumanInputFormEntityImpl(HumanInputFormEntity):
|
||||||
|
def __init__(self, form_model: HumanInputForm, recipient_models: Sequence[HumanInputFormRecipient]):
|
||||||
|
self._form_model = form_model
|
||||||
|
self._recipients = [_HumanInputFormRecipientEntityImpl(recipient) for recipient in recipient_models]
|
||||||
|
self._web_app_recipient = next(
|
||||||
|
(
|
||||||
|
recipient
|
||||||
|
for recipient in recipient_models
|
||||||
|
if recipient.recipient_type == RecipientType.STANDALONE_WEB_APP
|
||||||
|
),
|
||||||
|
None,
|
||||||
|
)
|
||||||
|
self._console_recipient = next(
|
||||||
|
(recipient for recipient in recipient_models if recipient.recipient_type == RecipientType.CONSOLE),
|
||||||
|
None,
|
||||||
|
)
|
||||||
|
self._submitted_data: Mapping[str, Any] | None = (
|
||||||
|
json.loads(form_model.submitted_data) if form_model.submitted_data is not None else None
|
||||||
|
)
|
||||||
|
|
||||||
|
@property
|
||||||
|
def id(self) -> str:
|
||||||
|
return self._form_model.id
|
||||||
|
|
||||||
|
@property
|
||||||
|
def web_app_token(self):
|
||||||
|
if self._console_recipient is not None:
|
||||||
|
return self._console_recipient.access_token
|
||||||
|
if self._web_app_recipient is None:
|
||||||
|
return None
|
||||||
|
return self._web_app_recipient.access_token
|
||||||
|
|
||||||
|
@property
|
||||||
|
def recipients(self) -> list[HumanInputFormRecipientEntity]:
|
||||||
|
return list(self._recipients)
|
||||||
|
|
||||||
|
@property
|
||||||
|
def rendered_content(self) -> str:
|
||||||
|
return self._form_model.rendered_content
|
||||||
|
|
||||||
|
@property
|
||||||
|
def selected_action_id(self) -> str | None:
|
||||||
|
return self._form_model.selected_action_id
|
||||||
|
|
||||||
|
@property
|
||||||
|
def submitted_data(self) -> Mapping[str, Any] | None:
|
||||||
|
return self._submitted_data
|
||||||
|
|
||||||
|
@property
|
||||||
|
def submitted(self) -> bool:
|
||||||
|
return self._form_model.submitted_at is not None
|
||||||
|
|
||||||
|
@property
|
||||||
|
def status(self) -> HumanInputFormStatus:
|
||||||
|
return self._form_model.status
|
||||||
|
|
||||||
|
@property
|
||||||
|
def expiration_time(self) -> datetime:
|
||||||
|
return self._form_model.expiration_time
|
||||||
|
|
||||||
|
|
||||||
|
@dataclasses.dataclass(frozen=True)
|
||||||
|
class HumanInputFormRecord:
|
||||||
|
form_id: str
|
||||||
|
workflow_run_id: str | None
|
||||||
|
node_id: str
|
||||||
|
tenant_id: str
|
||||||
|
app_id: str
|
||||||
|
form_kind: HumanInputFormKind
|
||||||
|
definition: FormDefinition
|
||||||
|
rendered_content: str
|
||||||
|
created_at: datetime
|
||||||
|
expiration_time: datetime
|
||||||
|
status: HumanInputFormStatus
|
||||||
|
selected_action_id: str | None
|
||||||
|
submitted_data: Mapping[str, Any] | None
|
||||||
|
submitted_at: datetime | None
|
||||||
|
submission_user_id: str | None
|
||||||
|
submission_end_user_id: str | None
|
||||||
|
completed_by_recipient_id: str | None
|
||||||
|
recipient_id: str | None
|
||||||
|
recipient_type: RecipientType | None
|
||||||
|
access_token: str | None
|
||||||
|
|
||||||
|
@property
|
||||||
|
def submitted(self) -> bool:
|
||||||
|
return self.submitted_at is not None
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_models(
|
||||||
|
cls, form_model: HumanInputForm, recipient_model: HumanInputFormRecipient | None
|
||||||
|
) -> "HumanInputFormRecord":
|
||||||
|
definition_payload = json.loads(form_model.form_definition)
|
||||||
|
if "expiration_time" not in definition_payload:
|
||||||
|
definition_payload["expiration_time"] = form_model.expiration_time
|
||||||
|
return cls(
|
||||||
|
form_id=form_model.id,
|
||||||
|
workflow_run_id=form_model.workflow_run_id,
|
||||||
|
node_id=form_model.node_id,
|
||||||
|
tenant_id=form_model.tenant_id,
|
||||||
|
app_id=form_model.app_id,
|
||||||
|
form_kind=form_model.form_kind,
|
||||||
|
definition=FormDefinition.model_validate(definition_payload),
|
||||||
|
rendered_content=form_model.rendered_content,
|
||||||
|
created_at=form_model.created_at,
|
||||||
|
expiration_time=form_model.expiration_time,
|
||||||
|
status=form_model.status,
|
||||||
|
selected_action_id=form_model.selected_action_id,
|
||||||
|
submitted_data=json.loads(form_model.submitted_data) if form_model.submitted_data else None,
|
||||||
|
submitted_at=form_model.submitted_at,
|
||||||
|
submission_user_id=form_model.submission_user_id,
|
||||||
|
submission_end_user_id=form_model.submission_end_user_id,
|
||||||
|
completed_by_recipient_id=form_model.completed_by_recipient_id,
|
||||||
|
recipient_id=recipient_model.id if recipient_model else None,
|
||||||
|
recipient_type=recipient_model.recipient_type if recipient_model else None,
|
||||||
|
access_token=recipient_model.access_token if recipient_model else None,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class _InvalidTimeoutStatusError(ValueError):
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
class HumanInputFormRepositoryImpl:
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
session_factory: sessionmaker | Engine,
|
||||||
|
tenant_id: str,
|
||||||
|
):
|
||||||
|
if isinstance(session_factory, Engine):
|
||||||
|
session_factory = sessionmaker(bind=session_factory)
|
||||||
|
self._session_factory = session_factory
|
||||||
|
self._tenant_id = tenant_id
|
||||||
|
|
||||||
|
def _delivery_method_to_model(
|
||||||
|
self,
|
||||||
|
session: Session,
|
||||||
|
form_id: str,
|
||||||
|
delivery_method: DeliveryChannelConfig,
|
||||||
|
) -> _DeliveryAndRecipients:
|
||||||
|
delivery_id = str(uuidv7())
|
||||||
|
delivery_model = HumanInputDelivery(
|
||||||
|
id=delivery_id,
|
||||||
|
form_id=form_id,
|
||||||
|
delivery_method_type=delivery_method.type,
|
||||||
|
delivery_config_id=delivery_method.id,
|
||||||
|
channel_payload=delivery_method.model_dump_json(),
|
||||||
|
)
|
||||||
|
recipients: list[HumanInputFormRecipient] = []
|
||||||
|
if isinstance(delivery_method, WebAppDeliveryMethod):
|
||||||
|
recipient_model = HumanInputFormRecipient(
|
||||||
|
form_id=form_id,
|
||||||
|
delivery_id=delivery_id,
|
||||||
|
recipient_type=RecipientType.STANDALONE_WEB_APP,
|
||||||
|
recipient_payload=StandaloneWebAppRecipientPayload().model_dump_json(),
|
||||||
|
)
|
||||||
|
recipients.append(recipient_model)
|
||||||
|
elif isinstance(delivery_method, EmailDeliveryMethod):
|
||||||
|
email_recipients_config = delivery_method.config.recipients
|
||||||
|
recipients.extend(
|
||||||
|
self._build_email_recipients(
|
||||||
|
session=session,
|
||||||
|
form_id=form_id,
|
||||||
|
delivery_id=delivery_id,
|
||||||
|
recipients_config=email_recipients_config,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
return _DeliveryAndRecipients(delivery=delivery_model, recipients=recipients)
|
||||||
|
|
||||||
|
def _build_email_recipients(
|
||||||
|
self,
|
||||||
|
session: Session,
|
||||||
|
form_id: str,
|
||||||
|
delivery_id: str,
|
||||||
|
recipients_config: EmailRecipients,
|
||||||
|
) -> list[HumanInputFormRecipient]:
|
||||||
|
member_user_ids = [
|
||||||
|
recipient.user_id for recipient in recipients_config.items if isinstance(recipient, MemberRecipient)
|
||||||
|
]
|
||||||
|
external_emails = [
|
||||||
|
recipient.email for recipient in recipients_config.items if isinstance(recipient, ExternalRecipient)
|
||||||
|
]
|
||||||
|
if recipients_config.whole_workspace:
|
||||||
|
members = self._query_all_workspace_members(session=session)
|
||||||
|
else:
|
||||||
|
members = self._query_workspace_members_by_ids(session=session, restrict_to_user_ids=member_user_ids)
|
||||||
|
|
||||||
|
return self._create_email_recipients_from_resolved(
|
||||||
|
form_id=form_id,
|
||||||
|
delivery_id=delivery_id,
|
||||||
|
members=members,
|
||||||
|
external_emails=external_emails,
|
||||||
|
)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _create_email_recipients_from_resolved(
|
||||||
|
*,
|
||||||
|
form_id: str,
|
||||||
|
delivery_id: str,
|
||||||
|
members: Sequence[_WorkspaceMemberInfo],
|
||||||
|
external_emails: Sequence[str],
|
||||||
|
) -> list[HumanInputFormRecipient]:
|
||||||
|
recipient_models: list[HumanInputFormRecipient] = []
|
||||||
|
seen_emails: set[str] = set()
|
||||||
|
|
||||||
|
for member in members:
|
||||||
|
if not member.email:
|
||||||
|
continue
|
||||||
|
if member.email in seen_emails:
|
||||||
|
continue
|
||||||
|
seen_emails.add(member.email)
|
||||||
|
payload = EmailMemberRecipientPayload(user_id=member.user_id, email=member.email)
|
||||||
|
recipient_models.append(
|
||||||
|
HumanInputFormRecipient.new(
|
||||||
|
form_id=form_id,
|
||||||
|
delivery_id=delivery_id,
|
||||||
|
payload=payload,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
for email in external_emails:
|
||||||
|
if not email:
|
||||||
|
continue
|
||||||
|
if email in seen_emails:
|
||||||
|
continue
|
||||||
|
seen_emails.add(email)
|
||||||
|
recipient_models.append(
|
||||||
|
HumanInputFormRecipient.new(
|
||||||
|
form_id=form_id,
|
||||||
|
delivery_id=delivery_id,
|
||||||
|
payload=EmailExternalRecipientPayload(email=email),
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
return recipient_models
|
||||||
|
|
||||||
|
def _query_all_workspace_members(
|
||||||
|
self,
|
||||||
|
session: Session,
|
||||||
|
) -> list[_WorkspaceMemberInfo]:
|
||||||
|
stmt = (
|
||||||
|
select(Account.id, Account.email)
|
||||||
|
.join(TenantAccountJoin, TenantAccountJoin.account_id == Account.id)
|
||||||
|
.where(TenantAccountJoin.tenant_id == self._tenant_id)
|
||||||
|
)
|
||||||
|
rows = session.execute(stmt).all()
|
||||||
|
return [_WorkspaceMemberInfo(user_id=account_id, email=email) for account_id, email in rows]
|
||||||
|
|
||||||
|
def _query_workspace_members_by_ids(
|
||||||
|
self,
|
||||||
|
session: Session,
|
||||||
|
restrict_to_user_ids: Sequence[str],
|
||||||
|
) -> list[_WorkspaceMemberInfo]:
|
||||||
|
unique_ids = {user_id for user_id in restrict_to_user_ids if user_id}
|
||||||
|
if not unique_ids:
|
||||||
|
return []
|
||||||
|
|
||||||
|
stmt = (
|
||||||
|
select(Account.id, Account.email)
|
||||||
|
.join(TenantAccountJoin, TenantAccountJoin.account_id == Account.id)
|
||||||
|
.where(TenantAccountJoin.tenant_id == self._tenant_id)
|
||||||
|
)
|
||||||
|
stmt = stmt.where(Account.id.in_(unique_ids))
|
||||||
|
|
||||||
|
rows = session.execute(stmt).all()
|
||||||
|
return [_WorkspaceMemberInfo(user_id=account_id, email=email) for account_id, email in rows]
|
||||||
|
|
||||||
|
def create_form(self, params: FormCreateParams) -> HumanInputFormEntity:
|
||||||
|
form_config: HumanInputNodeData = params.form_config
|
||||||
|
|
||||||
|
with self._session_factory(expire_on_commit=False) as session, session.begin():
|
||||||
|
# Generate unique form ID
|
||||||
|
form_id = str(uuidv7())
|
||||||
|
start_time = naive_utc_now()
|
||||||
|
node_expiration = form_config.expiration_time(start_time)
|
||||||
|
form_definition = FormDefinition(
|
||||||
|
form_content=form_config.form_content,
|
||||||
|
inputs=form_config.inputs,
|
||||||
|
user_actions=form_config.user_actions,
|
||||||
|
rendered_content=params.rendered_content,
|
||||||
|
expiration_time=node_expiration,
|
||||||
|
default_values=dict(params.resolved_default_values),
|
||||||
|
display_in_ui=params.display_in_ui,
|
||||||
|
node_title=form_config.title,
|
||||||
|
)
|
||||||
|
form_model = HumanInputForm(
|
||||||
|
id=form_id,
|
||||||
|
tenant_id=self._tenant_id,
|
||||||
|
app_id=params.app_id,
|
||||||
|
workflow_run_id=params.workflow_execution_id,
|
||||||
|
form_kind=params.form_kind,
|
||||||
|
node_id=params.node_id,
|
||||||
|
form_definition=form_definition.model_dump_json(),
|
||||||
|
rendered_content=params.rendered_content,
|
||||||
|
expiration_time=node_expiration,
|
||||||
|
created_at=start_time,
|
||||||
|
)
|
||||||
|
session.add(form_model)
|
||||||
|
recipient_models: list[HumanInputFormRecipient] = []
|
||||||
|
for delivery in params.delivery_methods:
|
||||||
|
delivery_and_recipients = self._delivery_method_to_model(
|
||||||
|
session=session,
|
||||||
|
form_id=form_id,
|
||||||
|
delivery_method=delivery,
|
||||||
|
)
|
||||||
|
session.add(delivery_and_recipients.delivery)
|
||||||
|
session.add_all(delivery_and_recipients.recipients)
|
||||||
|
recipient_models.extend(delivery_and_recipients.recipients)
|
||||||
|
if params.console_recipient_required and not any(
|
||||||
|
recipient.recipient_type == RecipientType.CONSOLE for recipient in recipient_models
|
||||||
|
):
|
||||||
|
console_delivery_id = str(uuidv7())
|
||||||
|
console_delivery = HumanInputDelivery(
|
||||||
|
id=console_delivery_id,
|
||||||
|
form_id=form_id,
|
||||||
|
delivery_method_type=DeliveryMethodType.WEBAPP,
|
||||||
|
delivery_config_id=None,
|
||||||
|
channel_payload=ConsoleDeliveryPayload().model_dump_json(),
|
||||||
|
)
|
||||||
|
console_recipient = HumanInputFormRecipient(
|
||||||
|
form_id=form_id,
|
||||||
|
delivery_id=console_delivery_id,
|
||||||
|
recipient_type=RecipientType.CONSOLE,
|
||||||
|
recipient_payload=ConsoleRecipientPayload(
|
||||||
|
account_id=params.console_creator_account_id,
|
||||||
|
).model_dump_json(),
|
||||||
|
)
|
||||||
|
session.add(console_delivery)
|
||||||
|
session.add(console_recipient)
|
||||||
|
recipient_models.append(console_recipient)
|
||||||
|
if params.backstage_recipient_required and not any(
|
||||||
|
recipient.recipient_type == RecipientType.BACKSTAGE for recipient in recipient_models
|
||||||
|
):
|
||||||
|
backstage_delivery_id = str(uuidv7())
|
||||||
|
backstage_delivery = HumanInputDelivery(
|
||||||
|
id=backstage_delivery_id,
|
||||||
|
form_id=form_id,
|
||||||
|
delivery_method_type=DeliveryMethodType.WEBAPP,
|
||||||
|
delivery_config_id=None,
|
||||||
|
channel_payload=ConsoleDeliveryPayload().model_dump_json(),
|
||||||
|
)
|
||||||
|
backstage_recipient = HumanInputFormRecipient(
|
||||||
|
form_id=form_id,
|
||||||
|
delivery_id=backstage_delivery_id,
|
||||||
|
recipient_type=RecipientType.BACKSTAGE,
|
||||||
|
recipient_payload=BackstageRecipientPayload(
|
||||||
|
account_id=params.console_creator_account_id,
|
||||||
|
).model_dump_json(),
|
||||||
|
)
|
||||||
|
session.add(backstage_delivery)
|
||||||
|
session.add(backstage_recipient)
|
||||||
|
recipient_models.append(backstage_recipient)
|
||||||
|
session.flush()
|
||||||
|
|
||||||
|
return _HumanInputFormEntityImpl(form_model=form_model, recipient_models=recipient_models)
|
||||||
|
|
||||||
|
def get_form(self, workflow_execution_id: str, node_id: str) -> HumanInputFormEntity | None:
|
||||||
|
form_query = select(HumanInputForm).where(
|
||||||
|
HumanInputForm.workflow_run_id == workflow_execution_id,
|
||||||
|
HumanInputForm.node_id == node_id,
|
||||||
|
HumanInputForm.tenant_id == self._tenant_id,
|
||||||
|
)
|
||||||
|
with self._session_factory(expire_on_commit=False) as session:
|
||||||
|
form_model: HumanInputForm | None = session.scalars(form_query).first()
|
||||||
|
if form_model is None:
|
||||||
|
return None
|
||||||
|
|
||||||
|
recipient_query = select(HumanInputFormRecipient).where(HumanInputFormRecipient.form_id == form_model.id)
|
||||||
|
recipient_models = session.scalars(recipient_query).all()
|
||||||
|
return _HumanInputFormEntityImpl(form_model=form_model, recipient_models=recipient_models)
|
||||||
|
|
||||||
|
|
||||||
|
class HumanInputFormSubmissionRepository:
|
||||||
|
"""Repository for fetching and submitting human input forms."""
|
||||||
|
|
||||||
|
def __init__(self, session_factory: sessionmaker | Engine):
|
||||||
|
if isinstance(session_factory, Engine):
|
||||||
|
session_factory = sessionmaker(bind=session_factory)
|
||||||
|
self._session_factory = session_factory
|
||||||
|
|
||||||
|
def get_by_token(self, form_token: str) -> HumanInputFormRecord | None:
|
||||||
|
query = (
|
||||||
|
select(HumanInputFormRecipient)
|
||||||
|
.options(selectinload(HumanInputFormRecipient.form))
|
||||||
|
.where(HumanInputFormRecipient.access_token == form_token)
|
||||||
|
)
|
||||||
|
with self._session_factory(expire_on_commit=False) as session:
|
||||||
|
recipient_model = session.scalars(query).first()
|
||||||
|
if recipient_model is None or recipient_model.form is None:
|
||||||
|
return None
|
||||||
|
return HumanInputFormRecord.from_models(recipient_model.form, recipient_model)
|
||||||
|
|
||||||
|
def get_by_form_id_and_recipient_type(
|
||||||
|
self,
|
||||||
|
form_id: str,
|
||||||
|
recipient_type: RecipientType,
|
||||||
|
) -> HumanInputFormRecord | None:
|
||||||
|
query = (
|
||||||
|
select(HumanInputFormRecipient)
|
||||||
|
.options(selectinload(HumanInputFormRecipient.form))
|
||||||
|
.where(
|
||||||
|
HumanInputFormRecipient.form_id == form_id,
|
||||||
|
HumanInputFormRecipient.recipient_type == recipient_type,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
with self._session_factory(expire_on_commit=False) as session:
|
||||||
|
recipient_model = session.scalars(query).first()
|
||||||
|
if recipient_model is None or recipient_model.form is None:
|
||||||
|
return None
|
||||||
|
return HumanInputFormRecord.from_models(recipient_model.form, recipient_model)
|
||||||
|
|
||||||
|
def mark_submitted(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
form_id: str,
|
||||||
|
recipient_id: str | None,
|
||||||
|
selected_action_id: str,
|
||||||
|
form_data: Mapping[str, Any],
|
||||||
|
submission_user_id: str | None,
|
||||||
|
submission_end_user_id: str | None,
|
||||||
|
) -> HumanInputFormRecord:
|
||||||
|
with self._session_factory(expire_on_commit=False) as session, session.begin():
|
||||||
|
form_model = session.get(HumanInputForm, form_id)
|
||||||
|
if form_model is None:
|
||||||
|
raise FormNotFoundError(f"form not found, id={form_id}")
|
||||||
|
|
||||||
|
recipient_model = session.get(HumanInputFormRecipient, recipient_id) if recipient_id else None
|
||||||
|
|
||||||
|
form_model.selected_action_id = selected_action_id
|
||||||
|
form_model.submitted_data = json.dumps(form_data)
|
||||||
|
form_model.submitted_at = naive_utc_now()
|
||||||
|
form_model.status = HumanInputFormStatus.SUBMITTED
|
||||||
|
form_model.submission_user_id = submission_user_id
|
||||||
|
form_model.submission_end_user_id = submission_end_user_id
|
||||||
|
form_model.completed_by_recipient_id = recipient_id
|
||||||
|
|
||||||
|
session.add(form_model)
|
||||||
|
session.flush()
|
||||||
|
session.refresh(form_model)
|
||||||
|
if recipient_model is not None:
|
||||||
|
session.refresh(recipient_model)
|
||||||
|
|
||||||
|
return HumanInputFormRecord.from_models(form_model, recipient_model)
|
||||||
|
|
||||||
|
def mark_timeout(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
form_id: str,
|
||||||
|
timeout_status: HumanInputFormStatus,
|
||||||
|
reason: str | None = None,
|
||||||
|
) -> HumanInputFormRecord:
|
||||||
|
with self._session_factory(expire_on_commit=False) as session, session.begin():
|
||||||
|
form_model = session.get(HumanInputForm, form_id)
|
||||||
|
if form_model is None:
|
||||||
|
raise FormNotFoundError(f"form not found, id={form_id}")
|
||||||
|
|
||||||
|
if timeout_status not in {HumanInputFormStatus.TIMEOUT, HumanInputFormStatus.EXPIRED}:
|
||||||
|
raise _InvalidTimeoutStatusError(f"invalid timeout status: {timeout_status}")
|
||||||
|
|
||||||
|
# already handled or submitted
|
||||||
|
if form_model.status in {HumanInputFormStatus.TIMEOUT, HumanInputFormStatus.EXPIRED}:
|
||||||
|
return HumanInputFormRecord.from_models(form_model, None)
|
||||||
|
|
||||||
|
if form_model.submitted_at is not None or form_model.status == HumanInputFormStatus.SUBMITTED:
|
||||||
|
raise FormNotFoundError(f"form already submitted, id={form_id}")
|
||||||
|
|
||||||
|
form_model.status = timeout_status
|
||||||
|
form_model.selected_action_id = None
|
||||||
|
form_model.submitted_data = None
|
||||||
|
form_model.submission_user_id = None
|
||||||
|
form_model.submission_end_user_id = None
|
||||||
|
form_model.completed_by_recipient_id = None
|
||||||
|
# Reason is recorded in status/error downstream; not stored on form.
|
||||||
|
session.add(form_model)
|
||||||
|
session.flush()
|
||||||
|
session.refresh(form_model)
|
||||||
|
|
||||||
|
return HumanInputFormRecord.from_models(form_model, None)
|
||||||
@@ -488,6 +488,7 @@ class SQLAlchemyWorkflowNodeExecutionRepository(WorkflowNodeExecutionRepository)
|
|||||||
WorkflowNodeExecutionModel.workflow_run_id == workflow_run_id,
|
WorkflowNodeExecutionModel.workflow_run_id == workflow_run_id,
|
||||||
WorkflowNodeExecutionModel.tenant_id == self._tenant_id,
|
WorkflowNodeExecutionModel.tenant_id == self._tenant_id,
|
||||||
WorkflowNodeExecutionModel.triggered_from == triggered_from,
|
WorkflowNodeExecutionModel.triggered_from == triggered_from,
|
||||||
|
WorkflowNodeExecutionModel.status != WorkflowNodeExecutionStatus.PAUSED,
|
||||||
)
|
)
|
||||||
|
|
||||||
if self._app_id:
|
if self._app_id:
|
||||||
|
|||||||
@@ -5,7 +5,7 @@ from collections.abc import Generator
|
|||||||
from copy import deepcopy
|
from copy import deepcopy
|
||||||
from typing import TYPE_CHECKING, Any
|
from typing import TYPE_CHECKING, Any
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING: # pragma: no cover
|
||||||
from models.model import File
|
from models.model import File
|
||||||
|
|
||||||
from core.tools.__base.tool_runtime import ToolRuntime
|
from core.tools.__base.tool_runtime import ToolRuntime
|
||||||
@@ -171,7 +171,7 @@ class Tool(ABC):
|
|||||||
def create_file_message(self, file: File) -> ToolInvokeMessage:
|
def create_file_message(self, file: File) -> ToolInvokeMessage:
|
||||||
return ToolInvokeMessage(
|
return ToolInvokeMessage(
|
||||||
type=ToolInvokeMessage.MessageType.FILE,
|
type=ToolInvokeMessage.MessageType.FILE,
|
||||||
message=ToolInvokeMessage.FileMessage(),
|
message=ToolInvokeMessage.FileMessage(file_marker="file_marker"),
|
||||||
meta={"file": file},
|
meta={"file": file},
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -1,4 +1,5 @@
|
|||||||
from core.tools.entities.tool_entities import ToolInvokeMeta
|
from core.tools.entities.tool_entities import ToolInvokeMeta
|
||||||
|
from libs.exception import BaseHTTPException
|
||||||
|
|
||||||
|
|
||||||
class ToolProviderNotFoundError(ValueError):
|
class ToolProviderNotFoundError(ValueError):
|
||||||
@@ -37,6 +38,12 @@ class ToolCredentialPolicyViolationError(ValueError):
|
|||||||
pass
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
class WorkflowToolHumanInputNotSupportedError(BaseHTTPException):
|
||||||
|
error_code = "workflow_tool_human_input_not_supported"
|
||||||
|
description = "Workflow with Human Input nodes cannot be published as a workflow tool."
|
||||||
|
code = 400
|
||||||
|
|
||||||
|
|
||||||
class ToolEngineInvokeError(Exception):
|
class ToolEngineInvokeError(Exception):
|
||||||
meta: ToolInvokeMeta
|
meta: ToolInvokeMeta
|
||||||
|
|
||||||
|
|||||||
@@ -3,8 +3,8 @@ from __future__ import annotations
|
|||||||
import base64
|
import base64
|
||||||
import json
|
import json
|
||||||
import logging
|
import logging
|
||||||
from collections.abc import Generator
|
from collections.abc import Generator, Mapping
|
||||||
from typing import Any
|
from typing import Any, cast
|
||||||
|
|
||||||
from core.mcp.auth_client import MCPClientWithAuthRetry
|
from core.mcp.auth_client import MCPClientWithAuthRetry
|
||||||
from core.mcp.error import MCPConnectionError
|
from core.mcp.error import MCPConnectionError
|
||||||
@@ -17,6 +17,7 @@ from core.mcp.types import (
|
|||||||
TextContent,
|
TextContent,
|
||||||
TextResourceContents,
|
TextResourceContents,
|
||||||
)
|
)
|
||||||
|
from core.model_runtime.entities.llm_entities import LLMUsage, LLMUsageMetadata
|
||||||
from core.tools.__base.tool import Tool
|
from core.tools.__base.tool import Tool
|
||||||
from core.tools.__base.tool_runtime import ToolRuntime
|
from core.tools.__base.tool_runtime import ToolRuntime
|
||||||
from core.tools.entities.tool_entities import ToolEntity, ToolInvokeMessage, ToolProviderType
|
from core.tools.entities.tool_entities import ToolEntity, ToolInvokeMessage, ToolProviderType
|
||||||
@@ -46,6 +47,7 @@ class MCPTool(Tool):
|
|||||||
self.headers = headers or {}
|
self.headers = headers or {}
|
||||||
self.timeout = timeout
|
self.timeout = timeout
|
||||||
self.sse_read_timeout = sse_read_timeout
|
self.sse_read_timeout = sse_read_timeout
|
||||||
|
self._latest_usage = LLMUsage.empty_usage()
|
||||||
|
|
||||||
def tool_provider_type(self) -> ToolProviderType:
|
def tool_provider_type(self) -> ToolProviderType:
|
||||||
return ToolProviderType.MCP
|
return ToolProviderType.MCP
|
||||||
@@ -59,6 +61,10 @@ class MCPTool(Tool):
|
|||||||
message_id: str | None = None,
|
message_id: str | None = None,
|
||||||
) -> Generator[ToolInvokeMessage, None, None]:
|
) -> Generator[ToolInvokeMessage, None, None]:
|
||||||
result = self.invoke_remote_mcp_tool(tool_parameters)
|
result = self.invoke_remote_mcp_tool(tool_parameters)
|
||||||
|
|
||||||
|
# Extract usage metadata from MCP protocol's _meta field
|
||||||
|
self._latest_usage = self._derive_usage_from_result(result)
|
||||||
|
|
||||||
# handle dify tool output
|
# handle dify tool output
|
||||||
for content in result.content:
|
for content in result.content:
|
||||||
if isinstance(content, TextContent):
|
if isinstance(content, TextContent):
|
||||||
@@ -120,6 +126,99 @@ class MCPTool(Tool):
|
|||||||
for item in json_list:
|
for item in json_list:
|
||||||
yield self.create_json_message(item)
|
yield self.create_json_message(item)
|
||||||
|
|
||||||
|
@property
|
||||||
|
def latest_usage(self) -> LLMUsage:
|
||||||
|
return self._latest_usage
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def _derive_usage_from_result(cls, result: CallToolResult) -> LLMUsage:
|
||||||
|
"""
|
||||||
|
Extract usage metadata from MCP tool result's _meta field.
|
||||||
|
|
||||||
|
The MCP protocol's _meta field (aliased as 'meta' in Python) can contain
|
||||||
|
usage information such as token counts, costs, and other metadata.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
result: The CallToolResult from MCP tool invocation
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
LLMUsage instance with values from meta or empty_usage if not found
|
||||||
|
"""
|
||||||
|
# Extract usage from the meta field if present
|
||||||
|
if result.meta:
|
||||||
|
usage_dict = cls._extract_usage_dict(result.meta)
|
||||||
|
if usage_dict is not None:
|
||||||
|
return LLMUsage.from_metadata(cast(LLMUsageMetadata, cast(object, dict(usage_dict))))
|
||||||
|
|
||||||
|
return LLMUsage.empty_usage()
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def _extract_usage_dict(cls, payload: Mapping[str, Any]) -> Mapping[str, Any] | None:
|
||||||
|
"""
|
||||||
|
Recursively search for usage dictionary in the payload.
|
||||||
|
|
||||||
|
The MCP protocol's _meta field can contain usage data in various formats:
|
||||||
|
- Direct usage field: {"usage": {...}}
|
||||||
|
- Nested in metadata: {"metadata": {"usage": {...}}}
|
||||||
|
- Or nested within other fields
|
||||||
|
|
||||||
|
Args:
|
||||||
|
payload: The payload to search for usage data
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
The usage dictionary if found, None otherwise
|
||||||
|
"""
|
||||||
|
# Check for direct usage field
|
||||||
|
usage_candidate = payload.get("usage")
|
||||||
|
if isinstance(usage_candidate, Mapping):
|
||||||
|
return usage_candidate
|
||||||
|
|
||||||
|
# Check for metadata nested usage
|
||||||
|
metadata_candidate = payload.get("metadata")
|
||||||
|
if isinstance(metadata_candidate, Mapping):
|
||||||
|
usage_candidate = metadata_candidate.get("usage")
|
||||||
|
if isinstance(usage_candidate, Mapping):
|
||||||
|
return usage_candidate
|
||||||
|
|
||||||
|
# Check for common token counting fields directly in payload
|
||||||
|
# Some MCP servers may include token counts directly
|
||||||
|
if "total_tokens" in payload or "prompt_tokens" in payload or "completion_tokens" in payload:
|
||||||
|
usage_dict: dict[str, Any] = {}
|
||||||
|
for key in (
|
||||||
|
"prompt_tokens",
|
||||||
|
"completion_tokens",
|
||||||
|
"total_tokens",
|
||||||
|
"prompt_unit_price",
|
||||||
|
"completion_unit_price",
|
||||||
|
"total_price",
|
||||||
|
"currency",
|
||||||
|
"prompt_price_unit",
|
||||||
|
"completion_price_unit",
|
||||||
|
"prompt_price",
|
||||||
|
"completion_price",
|
||||||
|
"latency",
|
||||||
|
"time_to_first_token",
|
||||||
|
"time_to_generate",
|
||||||
|
):
|
||||||
|
if key in payload:
|
||||||
|
usage_dict[key] = payload[key]
|
||||||
|
if usage_dict:
|
||||||
|
return usage_dict
|
||||||
|
|
||||||
|
# Recursively search through nested structures
|
||||||
|
for value in payload.values():
|
||||||
|
if isinstance(value, Mapping):
|
||||||
|
found = cls._extract_usage_dict(value)
|
||||||
|
if found is not None:
|
||||||
|
return found
|
||||||
|
elif isinstance(value, list) and not isinstance(value, (str, bytes, bytearray)):
|
||||||
|
for item in value:
|
||||||
|
if isinstance(item, Mapping):
|
||||||
|
found = cls._extract_usage_dict(item)
|
||||||
|
if found is not None:
|
||||||
|
return found
|
||||||
|
return None
|
||||||
|
|
||||||
def fork_tool_runtime(self, runtime: ToolRuntime) -> MCPTool:
|
def fork_tool_runtime(self, runtime: ToolRuntime) -> MCPTool:
|
||||||
return MCPTool(
|
return MCPTool(
|
||||||
entity=self.entity,
|
entity=self.entity,
|
||||||
|
|||||||
@@ -3,6 +3,8 @@ from typing import Any
|
|||||||
|
|
||||||
from core.app.app_config.entities import VariableEntity
|
from core.app.app_config.entities import VariableEntity
|
||||||
from core.tools.entities.tool_entities import WorkflowToolParameterConfiguration
|
from core.tools.entities.tool_entities import WorkflowToolParameterConfiguration
|
||||||
|
from core.tools.errors import WorkflowToolHumanInputNotSupportedError
|
||||||
|
from core.workflow.enums import NodeType
|
||||||
from core.workflow.nodes.base.entities import OutputVariableEntity
|
from core.workflow.nodes.base.entities import OutputVariableEntity
|
||||||
|
|
||||||
|
|
||||||
@@ -45,6 +47,13 @@ class WorkflowToolConfigurationUtils:
|
|||||||
|
|
||||||
return [outputs_by_variable[variable] for variable in variable_order]
|
return [outputs_by_variable[variable] for variable in variable_order]
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def ensure_no_human_input_nodes(cls, graph: Mapping[str, Any]) -> None:
|
||||||
|
nodes = graph.get("nodes", [])
|
||||||
|
for node in nodes:
|
||||||
|
if node.get("data", {}).get("type") == NodeType.HUMAN_INPUT:
|
||||||
|
raise WorkflowToolHumanInputNotSupportedError()
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def check_is_synced(
|
def check_is_synced(
|
||||||
cls, variables: list[VariableEntity], tool_configurations: list[WorkflowToolParameterConfiguration]
|
cls, variables: list[VariableEntity], tool_configurations: list[WorkflowToolParameterConfiguration]
|
||||||
|
|||||||
@@ -98,6 +98,10 @@ class WorkflowTool(Tool):
|
|||||||
invoke_from=self.runtime.invoke_from,
|
invoke_from=self.runtime.invoke_from,
|
||||||
streaming=False,
|
streaming=False,
|
||||||
call_depth=self.workflow_call_depth + 1,
|
call_depth=self.workflow_call_depth + 1,
|
||||||
|
# NOTE(QuantumGhost): We explicitly set `pause_state_config` to `None`
|
||||||
|
# because workflow pausing mechanisms (such as HumanInput) are not
|
||||||
|
# supported within WorkflowTool execution context.
|
||||||
|
pause_state_config=None,
|
||||||
)
|
)
|
||||||
assert isinstance(result, dict)
|
assert isinstance(result, dict)
|
||||||
data = result.get("data", {})
|
data = result.get("data", {})
|
||||||
|
|||||||
@@ -112,7 +112,7 @@ class ArrayBooleanVariable(ArrayBooleanSegment, ArrayVariable):
|
|||||||
|
|
||||||
class RAGPipelineVariable(BaseModel):
|
class RAGPipelineVariable(BaseModel):
|
||||||
belong_to_node_id: str = Field(description="belong to which node id, shared means public")
|
belong_to_node_id: str = Field(description="belong to which node id, shared means public")
|
||||||
type: str = Field(description="variable type, text-input, paragraph, select, number, file, file-list")
|
type: str = Field(description="variable type, text-input, paragraph, select, number, file, file-list")
|
||||||
label: str = Field(description="label")
|
label: str = Field(description="label")
|
||||||
description: str | None = Field(description="description", default="")
|
description: str | None = Field(description="description", default="")
|
||||||
variable: str = Field(description="variable key", default="")
|
variable: str = Field(description="variable key", default="")
|
||||||
|
|||||||
@@ -2,10 +2,12 @@ from .agent import AgentNodeStrategyInit
|
|||||||
from .graph_init_params import GraphInitParams
|
from .graph_init_params import GraphInitParams
|
||||||
from .workflow_execution import WorkflowExecution
|
from .workflow_execution import WorkflowExecution
|
||||||
from .workflow_node_execution import WorkflowNodeExecution
|
from .workflow_node_execution import WorkflowNodeExecution
|
||||||
|
from .workflow_start_reason import WorkflowStartReason
|
||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
"AgentNodeStrategyInit",
|
"AgentNodeStrategyInit",
|
||||||
"GraphInitParams",
|
"GraphInitParams",
|
||||||
"WorkflowExecution",
|
"WorkflowExecution",
|
||||||
"WorkflowNodeExecution",
|
"WorkflowNodeExecution",
|
||||||
|
"WorkflowStartReason",
|
||||||
]
|
]
|
||||||
|
|||||||
@@ -5,6 +5,16 @@ from pydantic import BaseModel, Field
|
|||||||
|
|
||||||
|
|
||||||
class GraphInitParams(BaseModel):
|
class GraphInitParams(BaseModel):
|
||||||
|
"""GraphInitParams encapsulates the configurations and contextual information
|
||||||
|
that remain constant throughout a single execution of the graph engine.
|
||||||
|
|
||||||
|
A single execution is defined as follows: as long as the execution has not reached
|
||||||
|
its conclusion, it is considered one execution. For instance, if a workflow is suspended
|
||||||
|
and later resumed, it is still regarded as a single execution, not two.
|
||||||
|
|
||||||
|
For the state diagram of workflow execution, refer to `WorkflowExecutionStatus`.
|
||||||
|
"""
|
||||||
|
|
||||||
# init params
|
# init params
|
||||||
tenant_id: str = Field(..., description="tenant / workspace id")
|
tenant_id: str = Field(..., description="tenant / workspace id")
|
||||||
app_id: str = Field(..., description="app id")
|
app_id: str = Field(..., description="app id")
|
||||||
|
|||||||
@@ -1,8 +1,11 @@
|
|||||||
|
from collections.abc import Mapping
|
||||||
from enum import StrEnum, auto
|
from enum import StrEnum, auto
|
||||||
from typing import Annotated, Literal, TypeAlias
|
from typing import Annotated, Any, Literal, TypeAlias
|
||||||
|
|
||||||
from pydantic import BaseModel, Field
|
from pydantic import BaseModel, Field
|
||||||
|
|
||||||
|
from core.workflow.nodes.human_input.entities import FormInput, UserAction
|
||||||
|
|
||||||
|
|
||||||
class PauseReasonType(StrEnum):
|
class PauseReasonType(StrEnum):
|
||||||
HUMAN_INPUT_REQUIRED = auto()
|
HUMAN_INPUT_REQUIRED = auto()
|
||||||
@@ -11,10 +14,31 @@ class PauseReasonType(StrEnum):
|
|||||||
|
|
||||||
class HumanInputRequired(BaseModel):
|
class HumanInputRequired(BaseModel):
|
||||||
TYPE: Literal[PauseReasonType.HUMAN_INPUT_REQUIRED] = PauseReasonType.HUMAN_INPUT_REQUIRED
|
TYPE: Literal[PauseReasonType.HUMAN_INPUT_REQUIRED] = PauseReasonType.HUMAN_INPUT_REQUIRED
|
||||||
|
|
||||||
form_id: str
|
form_id: str
|
||||||
# The identifier of the human input node causing the pause.
|
form_content: str
|
||||||
|
inputs: list[FormInput] = Field(default_factory=list)
|
||||||
|
actions: list[UserAction] = Field(default_factory=list)
|
||||||
|
display_in_ui: bool = False
|
||||||
node_id: str
|
node_id: str
|
||||||
|
node_title: str
|
||||||
|
|
||||||
|
# The `resolved_default_values` stores the resolved values of variable defaults. It's a mapping from
|
||||||
|
# `output_variable_name` to their resolved values.
|
||||||
|
#
|
||||||
|
# For example, The form contains a input with output variable name `name` and placeholder type `VARIABLE`, its
|
||||||
|
# selector is ["start", "name"]. While the HumanInputNode is executed, the correspond value of variable
|
||||||
|
# `start.name` in variable pool is `John`. Thus, the resolved value of the output variable `name` is `John`. The
|
||||||
|
# `resolved_default_values` is `{"name": "John"}`.
|
||||||
|
#
|
||||||
|
# Only form inputs with default value type `VARIABLE` will be resolved and stored in `resolved_default_values`.
|
||||||
|
resolved_default_values: Mapping[str, Any] = Field(default_factory=dict)
|
||||||
|
|
||||||
|
# The `form_token` is the token used to submit the form via UI surfaces. It corresponds to
|
||||||
|
# `HumanInputFormRecipient.access_token`.
|
||||||
|
#
|
||||||
|
# This field is `None` if webapp delivery is not set and not
|
||||||
|
# in orchestrating mode.
|
||||||
|
form_token: str | None = None
|
||||||
|
|
||||||
|
|
||||||
class SchedulingPause(BaseModel):
|
class SchedulingPause(BaseModel):
|
||||||
|
|||||||
@@ -0,0 +1,8 @@
|
|||||||
|
from enum import StrEnum
|
||||||
|
|
||||||
|
|
||||||
|
class WorkflowStartReason(StrEnum):
|
||||||
|
"""Reason for workflow start events across graph/queue/SSE layers."""
|
||||||
|
|
||||||
|
INITIAL = "initial" # First start of a workflow run.
|
||||||
|
RESUMPTION = "resumption" # Start triggered after resuming a paused run.
|
||||||
@@ -0,0 +1,15 @@
|
|||||||
|
import time
|
||||||
|
|
||||||
|
|
||||||
|
def get_timestamp() -> float:
|
||||||
|
"""Retrieve a timestamp as a float point numer representing the number of seconds
|
||||||
|
since the Unix epoch.
|
||||||
|
|
||||||
|
This function is primarily used to measure the execution time of the workflow engine.
|
||||||
|
Since workflow execution may be paused and resumed on a different machine,
|
||||||
|
`time.perf_counter` cannot be used as it is inconsistent across machines.
|
||||||
|
|
||||||
|
To address this, the function uses the wall clock as the time source.
|
||||||
|
However, it assumes that the clocks of all servers are properly synchronized.
|
||||||
|
"""
|
||||||
|
return round(time.time())
|
||||||
@@ -2,12 +2,14 @@
|
|||||||
GraphEngine configuration models.
|
GraphEngine configuration models.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from pydantic import BaseModel
|
from pydantic import BaseModel, ConfigDict
|
||||||
|
|
||||||
|
|
||||||
class GraphEngineConfig(BaseModel):
|
class GraphEngineConfig(BaseModel):
|
||||||
"""Configuration for GraphEngine worker pool scaling."""
|
"""Configuration for GraphEngine worker pool scaling."""
|
||||||
|
|
||||||
|
model_config = ConfigDict(frozen=True)
|
||||||
|
|
||||||
min_workers: int = 1
|
min_workers: int = 1
|
||||||
max_workers: int = 5
|
max_workers: int = 5
|
||||||
scale_up_threshold: int = 3
|
scale_up_threshold: int = 3
|
||||||
|
|||||||
@@ -10,6 +10,7 @@ from pydantic import BaseModel, Field
|
|||||||
|
|
||||||
from core.workflow.entities.pause_reason import PauseReason
|
from core.workflow.entities.pause_reason import PauseReason
|
||||||
from core.workflow.enums import NodeState
|
from core.workflow.enums import NodeState
|
||||||
|
from core.workflow.runtime.graph_runtime_state import GraphExecutionProtocol
|
||||||
|
|
||||||
from .node_execution import NodeExecution
|
from .node_execution import NodeExecution
|
||||||
|
|
||||||
@@ -236,3 +237,6 @@ class GraphExecution:
|
|||||||
def record_node_failure(self) -> None:
|
def record_node_failure(self) -> None:
|
||||||
"""Increment the count of node failures encountered during execution."""
|
"""Increment the count of node failures encountered during execution."""
|
||||||
self.exceptions_count += 1
|
self.exceptions_count += 1
|
||||||
|
|
||||||
|
|
||||||
|
_: GraphExecutionProtocol = GraphExecution(workflow_id="")
|
||||||
|
|||||||
@@ -192,9 +192,13 @@ class EventHandler:
|
|||||||
self._event_collector.collect(edge_event)
|
self._event_collector.collect(edge_event)
|
||||||
|
|
||||||
# Enqueue ready nodes
|
# Enqueue ready nodes
|
||||||
for node_id in ready_nodes:
|
if self._graph_execution.is_paused:
|
||||||
self._state_manager.enqueue_node(node_id)
|
for node_id in ready_nodes:
|
||||||
self._state_manager.start_execution(node_id)
|
self._graph_runtime_state.register_deferred_node(node_id)
|
||||||
|
else:
|
||||||
|
for node_id in ready_nodes:
|
||||||
|
self._state_manager.enqueue_node(node_id)
|
||||||
|
self._state_manager.start_execution(node_id)
|
||||||
|
|
||||||
# Update execution tracking
|
# Update execution tracking
|
||||||
self._state_manager.finish_execution(event.node_id)
|
self._state_manager.finish_execution(event.node_id)
|
||||||
|
|||||||
@@ -14,6 +14,7 @@ from collections.abc import Generator
|
|||||||
from typing import TYPE_CHECKING, cast, final
|
from typing import TYPE_CHECKING, cast, final
|
||||||
|
|
||||||
from core.workflow.context import capture_current_context
|
from core.workflow.context import capture_current_context
|
||||||
|
from core.workflow.entities.workflow_start_reason import WorkflowStartReason
|
||||||
from core.workflow.enums import NodeExecutionType
|
from core.workflow.enums import NodeExecutionType
|
||||||
from core.workflow.graph import Graph
|
from core.workflow.graph import Graph
|
||||||
from core.workflow.graph_events import (
|
from core.workflow.graph_events import (
|
||||||
@@ -55,6 +56,9 @@ if TYPE_CHECKING:
|
|||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
_DEFAULT_CONFIG = GraphEngineConfig()
|
||||||
|
|
||||||
|
|
||||||
@final
|
@final
|
||||||
class GraphEngine:
|
class GraphEngine:
|
||||||
"""
|
"""
|
||||||
@@ -70,7 +74,7 @@ class GraphEngine:
|
|||||||
graph: Graph,
|
graph: Graph,
|
||||||
graph_runtime_state: GraphRuntimeState,
|
graph_runtime_state: GraphRuntimeState,
|
||||||
command_channel: CommandChannel,
|
command_channel: CommandChannel,
|
||||||
config: GraphEngineConfig,
|
config: GraphEngineConfig = _DEFAULT_CONFIG,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""Initialize the graph engine with all subsystems and dependencies."""
|
"""Initialize the graph engine with all subsystems and dependencies."""
|
||||||
# stop event
|
# stop event
|
||||||
@@ -234,7 +238,9 @@ class GraphEngine:
|
|||||||
self._graph_execution.paused = False
|
self._graph_execution.paused = False
|
||||||
self._graph_execution.pause_reasons = []
|
self._graph_execution.pause_reasons = []
|
||||||
|
|
||||||
start_event = GraphRunStartedEvent()
|
start_event = GraphRunStartedEvent(
|
||||||
|
reason=WorkflowStartReason.RESUMPTION if is_resume else WorkflowStartReason.INITIAL,
|
||||||
|
)
|
||||||
self._event_manager.notify_layers(start_event)
|
self._event_manager.notify_layers(start_event)
|
||||||
yield start_event
|
yield start_event
|
||||||
|
|
||||||
@@ -303,15 +309,17 @@ class GraphEngine:
|
|||||||
for layer in self._layers:
|
for layer in self._layers:
|
||||||
try:
|
try:
|
||||||
layer.on_graph_start()
|
layer.on_graph_start()
|
||||||
except Exception as e:
|
except Exception:
|
||||||
logger.warning("Layer %s failed on_graph_start: %s", layer.__class__.__name__, e)
|
logger.exception("Layer %s failed on_graph_start", layer.__class__.__name__)
|
||||||
|
|
||||||
def _start_execution(self, *, resume: bool = False) -> None:
|
def _start_execution(self, *, resume: bool = False) -> None:
|
||||||
"""Start execution subsystems."""
|
"""Start execution subsystems."""
|
||||||
self._stop_event.clear()
|
self._stop_event.clear()
|
||||||
paused_nodes: list[str] = []
|
paused_nodes: list[str] = []
|
||||||
|
deferred_nodes: list[str] = []
|
||||||
if resume:
|
if resume:
|
||||||
paused_nodes = self._graph_runtime_state.consume_paused_nodes()
|
paused_nodes = self._graph_runtime_state.consume_paused_nodes()
|
||||||
|
deferred_nodes = self._graph_runtime_state.consume_deferred_nodes()
|
||||||
|
|
||||||
# Start worker pool (it calculates initial workers internally)
|
# Start worker pool (it calculates initial workers internally)
|
||||||
self._worker_pool.start()
|
self._worker_pool.start()
|
||||||
@@ -327,7 +335,11 @@ class GraphEngine:
|
|||||||
self._state_manager.enqueue_node(root_node.id)
|
self._state_manager.enqueue_node(root_node.id)
|
||||||
self._state_manager.start_execution(root_node.id)
|
self._state_manager.start_execution(root_node.id)
|
||||||
else:
|
else:
|
||||||
for node_id in paused_nodes:
|
seen_nodes: set[str] = set()
|
||||||
|
for node_id in paused_nodes + deferred_nodes:
|
||||||
|
if node_id in seen_nodes:
|
||||||
|
continue
|
||||||
|
seen_nodes.add(node_id)
|
||||||
self._state_manager.enqueue_node(node_id)
|
self._state_manager.enqueue_node(node_id)
|
||||||
self._state_manager.start_execution(node_id)
|
self._state_manager.start_execution(node_id)
|
||||||
|
|
||||||
@@ -345,8 +357,8 @@ class GraphEngine:
|
|||||||
for layer in self._layers:
|
for layer in self._layers:
|
||||||
try:
|
try:
|
||||||
layer.on_graph_end(self._graph_execution.error)
|
layer.on_graph_end(self._graph_execution.error)
|
||||||
except Exception as e:
|
except Exception:
|
||||||
logger.warning("Layer %s failed on_graph_end: %s", layer.__class__.__name__, e)
|
logger.exception("Layer %s failed on_graph_end", layer.__class__.__name__)
|
||||||
|
|
||||||
# Public property accessors for attributes that need external access
|
# Public property accessors for attributes that need external access
|
||||||
@property
|
@property
|
||||||
|
|||||||
@@ -224,6 +224,8 @@ class GraphStateManager:
|
|||||||
Returns:
|
Returns:
|
||||||
Number of executing nodes
|
Number of executing nodes
|
||||||
"""
|
"""
|
||||||
|
# This count is a best-effort snapshot and can change concurrently.
|
||||||
|
# Only use it for pause-drain checks where scheduling is already frozen.
|
||||||
with self._lock:
|
with self._lock:
|
||||||
return len(self._executing_nodes)
|
return len(self._executing_nodes)
|
||||||
|
|
||||||
|
|||||||
@@ -83,12 +83,12 @@ class Dispatcher:
|
|||||||
"""Main dispatcher loop."""
|
"""Main dispatcher loop."""
|
||||||
try:
|
try:
|
||||||
self._process_commands()
|
self._process_commands()
|
||||||
|
paused = False
|
||||||
while not self._stop_event.is_set():
|
while not self._stop_event.is_set():
|
||||||
if (
|
if self._execution_coordinator.aborted or self._execution_coordinator.execution_complete:
|
||||||
self._execution_coordinator.aborted
|
break
|
||||||
or self._execution_coordinator.paused
|
if self._execution_coordinator.paused:
|
||||||
or self._execution_coordinator.execution_complete
|
paused = True
|
||||||
):
|
|
||||||
break
|
break
|
||||||
|
|
||||||
self._execution_coordinator.check_scaling()
|
self._execution_coordinator.check_scaling()
|
||||||
@@ -101,13 +101,10 @@ class Dispatcher:
|
|||||||
time.sleep(0.1)
|
time.sleep(0.1)
|
||||||
|
|
||||||
self._process_commands()
|
self._process_commands()
|
||||||
while True:
|
if paused:
|
||||||
try:
|
self._drain_events_until_idle()
|
||||||
event = self._event_queue.get(block=False)
|
else:
|
||||||
self._event_handler.dispatch(event)
|
self._drain_event_queue()
|
||||||
self._event_queue.task_done()
|
|
||||||
except queue.Empty:
|
|
||||||
break
|
|
||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.exception("Dispatcher error")
|
logger.exception("Dispatcher error")
|
||||||
@@ -122,3 +119,24 @@ class Dispatcher:
|
|||||||
def _process_commands(self, event: GraphNodeEventBase | None = None):
|
def _process_commands(self, event: GraphNodeEventBase | None = None):
|
||||||
if event is None or isinstance(event, self._COMMAND_TRIGGER_EVENTS):
|
if event is None or isinstance(event, self._COMMAND_TRIGGER_EVENTS):
|
||||||
self._execution_coordinator.process_commands()
|
self._execution_coordinator.process_commands()
|
||||||
|
|
||||||
|
def _drain_event_queue(self) -> None:
|
||||||
|
while True:
|
||||||
|
try:
|
||||||
|
event = self._event_queue.get(block=False)
|
||||||
|
self._event_handler.dispatch(event)
|
||||||
|
self._event_queue.task_done()
|
||||||
|
except queue.Empty:
|
||||||
|
break
|
||||||
|
|
||||||
|
def _drain_events_until_idle(self) -> None:
|
||||||
|
while not self._stop_event.is_set():
|
||||||
|
try:
|
||||||
|
event = self._event_queue.get(timeout=0.1)
|
||||||
|
self._event_handler.dispatch(event)
|
||||||
|
self._event_queue.task_done()
|
||||||
|
self._process_commands(event)
|
||||||
|
except queue.Empty:
|
||||||
|
if not self._execution_coordinator.has_executing_nodes():
|
||||||
|
break
|
||||||
|
self._drain_event_queue()
|
||||||
|
|||||||
@@ -94,3 +94,11 @@ class ExecutionCoordinator:
|
|||||||
|
|
||||||
self._worker_pool.stop()
|
self._worker_pool.stop()
|
||||||
self._state_manager.clear_executing()
|
self._state_manager.clear_executing()
|
||||||
|
|
||||||
|
def has_executing_nodes(self) -> bool:
|
||||||
|
"""Return True if any nodes are currently marked as executing."""
|
||||||
|
# This check is only safe once execution has already paused.
|
||||||
|
# Before pause, executing state can change concurrently, which makes the result unreliable.
|
||||||
|
if not self._graph_execution.is_paused:
|
||||||
|
raise AssertionError("has_executing_nodes should only be called after execution is paused")
|
||||||
|
return self._state_manager.get_executing_count() > 0
|
||||||
|
|||||||
@@ -38,6 +38,8 @@ from .loop import (
|
|||||||
from .node import (
|
from .node import (
|
||||||
NodeRunExceptionEvent,
|
NodeRunExceptionEvent,
|
||||||
NodeRunFailedEvent,
|
NodeRunFailedEvent,
|
||||||
|
NodeRunHumanInputFormFilledEvent,
|
||||||
|
NodeRunHumanInputFormTimeoutEvent,
|
||||||
NodeRunPauseRequestedEvent,
|
NodeRunPauseRequestedEvent,
|
||||||
NodeRunRetrieverResourceEvent,
|
NodeRunRetrieverResourceEvent,
|
||||||
NodeRunRetryEvent,
|
NodeRunRetryEvent,
|
||||||
@@ -60,6 +62,8 @@ __all__ = [
|
|||||||
"NodeRunAgentLogEvent",
|
"NodeRunAgentLogEvent",
|
||||||
"NodeRunExceptionEvent",
|
"NodeRunExceptionEvent",
|
||||||
"NodeRunFailedEvent",
|
"NodeRunFailedEvent",
|
||||||
|
"NodeRunHumanInputFormFilledEvent",
|
||||||
|
"NodeRunHumanInputFormTimeoutEvent",
|
||||||
"NodeRunIterationFailedEvent",
|
"NodeRunIterationFailedEvent",
|
||||||
"NodeRunIterationNextEvent",
|
"NodeRunIterationNextEvent",
|
||||||
"NodeRunIterationStartedEvent",
|
"NodeRunIterationStartedEvent",
|
||||||
|
|||||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user