Compare commits
11
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
f0c38e7da1 | ||
|
|
5f49fb4e4b | ||
|
|
aabe0258da | ||
|
|
723b60079f | ||
|
|
1e8d72d9ae | ||
|
|
0e546adac0 | ||
|
|
16413caa06 | ||
|
|
0f6a3c79e8 | ||
|
|
1795eedf03 | ||
|
|
9b68dd68b8 | ||
|
|
03cbdbe355 |
+19
-3
@@ -559,6 +559,20 @@ INDEXING_MAX_SEGMENTATION_TOKENS_LENGTH=4000
|
||||
# Workflow runtime configuration
|
||||
WORKFLOW_MAX_EXECUTION_STEPS=500
|
||||
WORKFLOW_MAX_EXECUTION_TIME=1200
|
||||
# Durable handoff for planned worker shutdowns. Deploy migrations and resume consumers first, then enable it.
|
||||
# Every API/worker replica must share the configured object storage (or the same RWX volume).
|
||||
# Celery's prefork pool is intentionally rejected while enabled because drain state is process-local.
|
||||
# EVENT_BUS_REDIS_CHANNEL_TYPE=streams is mandatory so clients can replay events with Last-Event-ID.
|
||||
WORKFLOW_HANDOFF_ENABLED=false
|
||||
# Per-phase timeout for checkpoint preparation/activation and the first resume claim.
|
||||
WORKFLOW_HANDOFF_DRAIN_TIMEOUT_SECONDS=600
|
||||
WORKFLOW_HANDOFF_SCAN_INTERVAL_SECONDS=15
|
||||
WORKFLOW_HANDOFF_LEASE_SECONDS=120
|
||||
WORKFLOW_HANDOFF_MAX_ATTEMPTS=20
|
||||
# Retain terminal handoff audit rows and completed snapshot-GC records for at least this many days.
|
||||
WORKFLOW_HANDOFF_RETENTION_DAYS=7
|
||||
# Capability queue: only upgraded workers should consume it during an adjacent-version rollout.
|
||||
WORKFLOW_HANDOFF_QUEUE=workflow_handoff
|
||||
WORKFLOW_CALL_MAX_DEPTH=5
|
||||
MAX_VARIABLE_SIZE=204800
|
||||
# Maximum concurrent node-builder LLM calls per workflow generation request
|
||||
@@ -792,10 +806,12 @@ EVENT_BUS_REDIS_URL=
|
||||
# - streams: Redis Streams (at-least-once, recommended to avoid subscriber races)
|
||||
#
|
||||
# Note: Before enabling 'streams' in production, estimate your expected event volume and retention needs.
|
||||
# Configure Redis memory limits and stream trimming appropriately (e.g., MAXLEN and key expiry) to reduce
|
||||
# the risk of data loss from Redis auto-eviction under memory pressure.
|
||||
# Also accepts ENV: EVENT_BUS_REDIS_CHANNEL_TYPE.
|
||||
# Each per-run stream retains its complete cursor window until key expiry; size Redis to avoid eviction,
|
||||
# because trimming a live stream would invalidate replay guarantees.
|
||||
# Also accepts ENV: EVENT_BUS_REDIS_CHANNEL_TYPE. Must be streams when WORKFLOW_HANDOFF_ENABLED=true.
|
||||
EVENT_BUS_REDIS_CHANNEL_TYPE=pubsub
|
||||
# When handoff is enabled, must be at least WORKFLOW_HANDOFF_DRAIN_TIMEOUT_SECONDS + 60 seconds.
|
||||
EVENT_BUS_STREAMS_RETENTION_SECONDS=900
|
||||
# Whether to use Redis cluster mode while use redis as event bus.
|
||||
# It's highly recommended to enable this for large deployments.
|
||||
EVENT_BUS_REDIS_USE_CLUSTERS=false
|
||||
|
||||
@@ -1,7 +1,8 @@
|
||||
import logging
|
||||
from pathlib import Path
|
||||
from typing import Any, override
|
||||
from typing import Any, Self, override
|
||||
|
||||
from pydantic import model_validator
|
||||
from pydantic.fields import FieldInfo
|
||||
from pydantic_settings import BaseSettings, PydanticBaseSettingsSource, SettingsConfigDict, TomlConfigSettingsSource
|
||||
|
||||
@@ -92,6 +93,22 @@ class DifyConfig(
|
||||
# for better readability and maintainability.
|
||||
# Thanks for your concentration and consideration.
|
||||
|
||||
@model_validator(mode="after")
|
||||
def validate_workflow_handoff_event_transport(self) -> Self:
|
||||
if self.WORKFLOW_HANDOFF_ENABLED and self.PUBSUB_REDIS_CHANNEL_TYPE != "streams":
|
||||
raise ValueError(
|
||||
"WORKFLOW_HANDOFF_ENABLED requires EVENT_BUS_REDIS_CHANNEL_TYPE=streams "
|
||||
"so workflow events remain replayable across rolling updates"
|
||||
)
|
||||
minimum_stream_retention = self.WORKFLOW_HANDOFF_DRAIN_TIMEOUT_SECONDS + 60
|
||||
if self.WORKFLOW_HANDOFF_ENABLED and minimum_stream_retention > self.PUBSUB_STREAMS_RETENTION_SECONDS:
|
||||
raise ValueError(
|
||||
"WORKFLOW_HANDOFF_ENABLED requires EVENT_BUS_STREAMS_RETENTION_SECONDS to be at least "
|
||||
"WORKFLOW_HANDOFF_DRAIN_TIMEOUT_SECONDS + 60 seconds "
|
||||
f"({minimum_stream_retention} seconds for the current drain timeout)"
|
||||
)
|
||||
return self
|
||||
|
||||
@classmethod
|
||||
@override
|
||||
def settings_customise_sources(
|
||||
|
||||
@@ -893,6 +893,53 @@ class WorkflowConfig(BaseSettings):
|
||||
default=1200,
|
||||
)
|
||||
|
||||
WORKFLOW_HANDOFF_ENABLED: bool = Field(
|
||||
description=(
|
||||
"Enable durable workflow handoff during planned worker shutdowns; all replicas must share object storage "
|
||||
"and EVENT_BUS_REDIS_CHANNEL_TYPE must be streams"
|
||||
),
|
||||
default=False,
|
||||
)
|
||||
|
||||
WORKFLOW_HANDOFF_DRAIN_TIMEOUT_SECONDS: PositiveInt = Field(
|
||||
description=(
|
||||
"Maximum time in seconds for checkpoint preparation/activation and for a READY handoff's first resume claim"
|
||||
),
|
||||
default=600,
|
||||
)
|
||||
|
||||
WORKFLOW_HANDOFF_SCAN_INTERVAL_SECONDS: PositiveInt = Field(
|
||||
description="Interval in seconds between scans for workflow handoffs that need to be resumed",
|
||||
default=15,
|
||||
)
|
||||
|
||||
WORKFLOW_HANDOFF_LEASE_SECONDS: PositiveInt = Field(
|
||||
description="Duration in seconds of an exclusive workflow handoff resume lease",
|
||||
default=120,
|
||||
)
|
||||
|
||||
WORKFLOW_HANDOFF_MAX_ATTEMPTS: PositiveInt = Field(
|
||||
description="Maximum number of attempts to resume a workflow handoff before failing closed",
|
||||
default=20,
|
||||
)
|
||||
|
||||
WORKFLOW_HANDOFF_RETENTION_DAYS: PositiveInt = Field(
|
||||
description=(
|
||||
"Days to retain terminal workflow handoff audit rows and completed snapshot-GC records after their "
|
||||
"durable side effects finish"
|
||||
),
|
||||
default=7,
|
||||
)
|
||||
|
||||
WORKFLOW_HANDOFF_QUEUE: str = Field(
|
||||
description=(
|
||||
"Capability-isolated Celery queue used for workflow handoff scan and resume tasks; only upgraded workers "
|
||||
"must consume it during an adjacent-version rollout"
|
||||
),
|
||||
default="workflow_handoff",
|
||||
min_length=1,
|
||||
)
|
||||
|
||||
WORKFLOW_CALL_MAX_DEPTH: PositiveInt = Field(
|
||||
description="Maximum allowed depth for nested workflow calls",
|
||||
default=5,
|
||||
|
||||
+4
-3
@@ -54,8 +54,8 @@ class RedisPubSubConfig(BaseSettings):
|
||||
" - sharded: sharded Pub/Sub (at-most-once)\n"
|
||||
" - streams: Redis Streams (at-least-once, recommended to avoid subscriber races)\n\n"
|
||||
"Note: Before enabling 'streams' in production, estimate your expected event volume and retention needs.\n"
|
||||
"Configure Redis memory limits and stream trimming appropriately (e.g., MAXLEN and key expiry) to reduce\n"
|
||||
"the risk of data loss from Redis auto-eviction under memory pressure.\n"
|
||||
"Each per-run stream retains its complete cursor window until key expiry; size Redis to avoid eviction,\n"
|
||||
"because trimming a live stream would invalidate replay guarantees.\n"
|
||||
"Also accepts ENV: EVENT_BUS_REDIS_CHANNEL_TYPE."
|
||||
),
|
||||
default="pubsub",
|
||||
@@ -65,9 +65,10 @@ class RedisPubSubConfig(BaseSettings):
|
||||
validation_alias=AliasChoices("EVENT_BUS_STREAMS_RETENTION_SECONDS", "PUBSUB_STREAMS_RETENTION_SECONDS"),
|
||||
description=(
|
||||
"When using 'streams', expire each stream key this many seconds after the last event is published. "
|
||||
"Set this longer than the maximum client reconnect and rolling-update handoff window. "
|
||||
"Also accepts ENV: EVENT_BUS_STREAMS_RETENTION_SECONDS."
|
||||
),
|
||||
default=600,
|
||||
default=900,
|
||||
)
|
||||
|
||||
def _build_default_pubsub_url(self) -> str:
|
||||
|
||||
@@ -0,0 +1,20 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from flask import Request
|
||||
from werkzeug.exceptions import BadRequest
|
||||
|
||||
from libs.broadcast_channel.cursor import normalize_stream_cursor
|
||||
|
||||
|
||||
def get_workflow_event_replay_cursor(request: Request) -> str | None:
|
||||
"""Read the standard SSE cursor, with a query fallback for non-EventSource clients."""
|
||||
|
||||
raw_cursor = (request.headers.get("Last-Event-ID") or "").strip()
|
||||
if not raw_cursor:
|
||||
raw_cursor = (request.args.get("cursor") or "").strip()
|
||||
if not raw_cursor:
|
||||
return None
|
||||
try:
|
||||
return normalize_stream_cursor(raw_cursor)
|
||||
except ValueError as exc:
|
||||
raise BadRequest(str(exc)) from exc
|
||||
@@ -219,6 +219,8 @@ class CompletionMessageStopApi(Resource):
|
||||
invoke_from=InvokeFrom.DEBUGGER,
|
||||
user_id=current_user_id,
|
||||
app_mode=AppMode.value_of(app_model.mode),
|
||||
tenant_id=app_model.tenant_id,
|
||||
app_id=app_model.id,
|
||||
)
|
||||
|
||||
return SimpleResultResponse(result="success").model_dump(mode="json"), 200
|
||||
@@ -582,6 +584,8 @@ def _stop_chat_message(*, current_user_id: str, app_model: App, task_id: str):
|
||||
invoke_from=InvokeFrom.DEBUGGER,
|
||||
user_id=current_user_id,
|
||||
app_mode=AppMode.value_of(app_model.mode),
|
||||
tenant_id=app_model.tenant_id,
|
||||
app_id=app_model.id,
|
||||
)
|
||||
|
||||
return SimpleResultResponse(result="success").model_dump(mode="json"), 200
|
||||
|
||||
@@ -40,7 +40,6 @@ from controllers.console.wraps import (
|
||||
)
|
||||
from controllers.web.error import InvokeRateLimitError as InvokeRateLimitHttpError
|
||||
from core.app.app_config.features.file_upload.manager import FileUploadConfigManager
|
||||
from core.app.apps.base_app_queue_manager import AppQueueManager
|
||||
from core.app.apps.workflow.app_generator import SKIP_PREPARE_USER_INPUTS_KEY
|
||||
from core.app.entities.app_invoke_entities import InvokeFrom
|
||||
from core.app.file_access import DatabaseFileAccessController
|
||||
@@ -64,7 +63,6 @@ from fields.workflow_run_fields import WorkflowRunNodeExecutionResponse
|
||||
from graphon.enums import NodeType
|
||||
from graphon.file import File
|
||||
from graphon.file import helpers as file_helpers
|
||||
from graphon.graph_engine.manager import GraphEngineManager
|
||||
from graphon.model_runtime.utils.encoders import jsonable_encoder
|
||||
from graphon.variables import SecretVariable, SegmentType, VariableBase
|
||||
from graphon.variables.exc import VariableError
|
||||
@@ -78,6 +76,7 @@ from models.workflow import Workflow
|
||||
from repositories.workflow_collaboration_repository import WORKFLOW_ONLINE_USERS_PREFIX
|
||||
from services.agent.retirement_service import WorkflowAgentRetirementService
|
||||
from services.app_generate_service import AppGenerateService
|
||||
from services.app_task_service import AppTaskService
|
||||
from services.errors.app import IsDraftWorkflowError, WorkflowHashNotEqualError, WorkflowNotFoundError
|
||||
from services.errors.llm import InvokeRateLimitError
|
||||
from services.workflow_ref_service import WorkflowRefService
|
||||
@@ -1131,16 +1130,19 @@ class WorkflowTaskStopApi(Resource):
|
||||
@edit_permission_required
|
||||
@rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_TEST_AND_RUN)
|
||||
@get_app_model(mode=[AppMode.ADVANCED_CHAT, AppMode.WORKFLOW])
|
||||
def post(self, app_model: App, task_id: str):
|
||||
@with_current_user
|
||||
def post(self, current_user: Account, app_model: App, task_id: str):
|
||||
"""
|
||||
Stop workflow task
|
||||
"""
|
||||
# Stop using both mechanisms for backward compatibility
|
||||
# Legacy stop flag mechanism (without user check)
|
||||
AppQueueManager.set_stop_flag_no_user_check(task_id)
|
||||
|
||||
# New graph engine command channel mechanism
|
||||
GraphEngineManager(redis_client).send_stop_command(task_id)
|
||||
AppTaskService.stop_task(
|
||||
task_id,
|
||||
InvokeFrom.DEBUGGER,
|
||||
current_user.id,
|
||||
AppMode.value_of(app_model.mode),
|
||||
tenant_id=app_model.tenant_id,
|
||||
app_id=app_model.id,
|
||||
)
|
||||
|
||||
return {"result": "success"}
|
||||
|
||||
|
||||
@@ -71,6 +71,7 @@ class WorkflowRunForLogResponse(ResponseModel):
|
||||
triggered_from: str | None = None
|
||||
error: str | None = None
|
||||
elapsed_time: float | None = None
|
||||
handoff_duration: float = 0.0
|
||||
total_tokens: int | None = None
|
||||
total_steps: int | None = None
|
||||
created_at: int | None = None
|
||||
@@ -97,6 +98,7 @@ class WorkflowRunForArchivedLogResponse(ResponseModel):
|
||||
status: str | None = None
|
||||
triggered_from: str | None = None
|
||||
elapsed_time: float | None = None
|
||||
handoff_duration: float = 0.0
|
||||
total_tokens: int | None = None
|
||||
|
||||
@field_validator("status", mode="before")
|
||||
|
||||
@@ -3,7 +3,7 @@ import logging
|
||||
from typing import Any, Literal, cast
|
||||
from uuid import UUID
|
||||
|
||||
from flask import abort, request
|
||||
from flask import Response, abort, request
|
||||
from flask_restx import Resource
|
||||
from pydantic import BaseModel, Field, RootModel, ValidationError
|
||||
from sqlalchemy.orm import Session, sessionmaker
|
||||
@@ -13,6 +13,7 @@ import services
|
||||
from controllers.common.controller_schemas import DefaultBlockConfigQuery, WorkflowListQuery, WorkflowUpdatePayload
|
||||
from controllers.common.fields import SimpleResultResponse
|
||||
from controllers.common.schema import query_params_from_model, register_response_schema_models, register_schema_models
|
||||
from controllers.common.workflow_event_cursor import get_workflow_event_replay_cursor
|
||||
from controllers.console import console_ns
|
||||
from controllers.console.app.error import (
|
||||
ConversationCompletedError,
|
||||
@@ -40,9 +41,14 @@ from controllers.console.wraps import (
|
||||
)
|
||||
from controllers.web.error import InvokeRateLimitError as InvokeRateLimitHttpError
|
||||
from core.app.apps.base_app_queue_manager import AppQueueManager
|
||||
from core.app.apps.common.workflow_response_converter import WorkflowResponseConverter
|
||||
from core.app.apps.message_generator import MessageGenerator
|
||||
from core.app.apps.pipeline.pipeline_generator import PipelineGenerator
|
||||
from core.app.entities.app_invoke_entities import InvokeFrom
|
||||
from core.app.entities.task_entities import StreamEvent
|
||||
from core.workflow.human_input_policy import HumanInputSurface
|
||||
from extensions.ext_database import db
|
||||
from extensions.ext_redis import redis_client
|
||||
from factories import variable_factory
|
||||
from fields.base import ResponseModel
|
||||
from fields.workflow_run_fields import (
|
||||
@@ -51,13 +57,15 @@ from fields.workflow_run_fields import (
|
||||
WorkflowRunNodeExecutionResponse,
|
||||
WorkflowRunPaginationResponse,
|
||||
)
|
||||
from graphon.graph_engine.manager import GraphEngineManager
|
||||
from graphon.model_runtime.utils.encoders import jsonable_encoder
|
||||
from libs import helper
|
||||
from libs.helper import TimestampField, UUIDStrOrEmpty, dump_response
|
||||
from libs.login import login_required
|
||||
from models import Account
|
||||
from models.dataset import Pipeline
|
||||
from models.model import EndUser
|
||||
from models.enums import CreatorUserRole
|
||||
from models.model import AppMode, EndUser
|
||||
from models.workflow import Workflow
|
||||
from services.errors.app import IsDraftWorkflowError, WorkflowHashNotEqualError, WorkflowNotFoundError
|
||||
from services.errors.llm import InvokeRateLimitError
|
||||
@@ -65,6 +73,8 @@ from services.rag_pipeline.pipeline_generate_service import PipelineGenerateServ
|
||||
from services.rag_pipeline.rag_pipeline import RagPipelineService
|
||||
from services.rag_pipeline.rag_pipeline_manage_service import RagPipelineManageService
|
||||
from services.rag_pipeline.rag_pipeline_transform_service import RagPipelineTransformService
|
||||
from services.workflow_event_snapshot_service import build_workflow_event_stream, resolve_workflow_event_task_id
|
||||
from services.workflow_handoff_cancellation_service import request_workflow_handoff_cancel_for_app
|
||||
from services.workflow_ref_service import WorkflowRefService
|
||||
from services.workflow_service import DraftWorkflowDeletionError, WorkflowInUseError, WorkflowService
|
||||
|
||||
@@ -517,7 +527,22 @@ class RagPipelineTaskStopApi(Resource):
|
||||
"""
|
||||
Stop workflow task
|
||||
"""
|
||||
AppQueueManager.set_stop_flag(task_id, InvokeFrom.DEBUGGER, current_user.id)
|
||||
# Preserve the legacy account ownership check. The task id is
|
||||
# caller-controlled, so an unscoped Redis flag or Graph command could
|
||||
# otherwise stop a run owned by another account or tenant.
|
||||
live_task_owned_by_user = AppQueueManager.set_stop_flag(task_id, InvokeFrom.DEBUGGER, current_user.id)
|
||||
cancelled_handoffs = request_workflow_handoff_cancel_for_app(
|
||||
task_id,
|
||||
tenant_id=pipeline.tenant_id,
|
||||
app_id=pipeline.id,
|
||||
created_by_role=CreatorUserRole.ACCOUNT,
|
||||
created_by=current_user.id,
|
||||
)
|
||||
# A live non-handoff run observes the scoped Redis flag and propagates
|
||||
# its own graph abort. Send directly only when the owner-scoped
|
||||
# durable update proves this task belongs to the requested pipeline.
|
||||
if live_task_owned_by_user or cancelled_handoffs > 0:
|
||||
GraphEngineManager(redis_client).send_stop_command(task_id)
|
||||
|
||||
return {"result": "success"}
|
||||
|
||||
@@ -944,6 +969,82 @@ class RagPipelineWorkflowRunDetailApi(Resource):
|
||||
return WorkflowRunDetailResponse.model_validate(workflow_run, from_attributes=True).model_dump(mode="json")
|
||||
|
||||
|
||||
@console_ns.route("/rag/pipelines/<uuid:pipeline_id>/workflow-runs/<uuid:run_id>/events")
|
||||
class RagPipelineWorkflowRunEventsApi(Resource):
|
||||
"""Reconnect to the durable event stream of a RAG pipeline run."""
|
||||
|
||||
@setup_required
|
||||
@login_required
|
||||
@account_initialization_required
|
||||
@with_current_user
|
||||
@get_rag_pipeline
|
||||
def get(self, current_user: Account, pipeline: Pipeline, run_id: UUID):
|
||||
rag_pipeline_service = RagPipelineService(db.session())
|
||||
session_maker = rag_pipeline_service.session_maker
|
||||
workflow_run = rag_pipeline_service.get_rag_pipeline_workflow_run(pipeline=pipeline, run_id=str(run_id))
|
||||
if (
|
||||
workflow_run is None
|
||||
or workflow_run.app_id != pipeline.id
|
||||
or workflow_run.created_by_role != CreatorUserRole.ACCOUNT
|
||||
or workflow_run.created_by != current_user.id
|
||||
):
|
||||
raise NotFound("Workflow run not found")
|
||||
|
||||
cursor = get_workflow_event_replay_cursor(request)
|
||||
include_state_snapshot = request.args.get("include_state_snapshot", "false").lower() == "true"
|
||||
continue_on_pause = request.args.get("continue_on_pause", "false").lower() == "true"
|
||||
|
||||
if workflow_run.finished_at is not None and cursor is None:
|
||||
finished = WorkflowResponseConverter.workflow_run_result_to_finish_response(
|
||||
task_id=resolve_workflow_event_task_id(
|
||||
workflow_run=workflow_run,
|
||||
session_maker=session_maker,
|
||||
),
|
||||
workflow_run=workflow_run,
|
||||
creator_user=current_user,
|
||||
)
|
||||
payload = finished.model_dump(mode="json")
|
||||
payload["event"] = finished.event.value
|
||||
|
||||
def event_generator():
|
||||
yield f"data: {json.dumps(payload)}\n\n"
|
||||
|
||||
else:
|
||||
message_generator = MessageGenerator()
|
||||
|
||||
def event_generator():
|
||||
if include_state_snapshot or cursor is not None:
|
||||
source = build_workflow_event_stream(
|
||||
app_mode=AppMode.RAG_PIPELINE,
|
||||
workflow_run=workflow_run,
|
||||
tenant_id=pipeline.tenant_id,
|
||||
app_id=pipeline.id,
|
||||
session_maker=session_maker,
|
||||
human_input_surface=HumanInputSurface.CONSOLE,
|
||||
close_on_pause=not continue_on_pause,
|
||||
cursor=cursor,
|
||||
)
|
||||
else:
|
||||
terminal_events = [StreamEvent.WORKFLOW_FINISHED] if continue_on_pause else None
|
||||
retrieve_kwargs: dict[str, Any] = {"terminal_events": terminal_events}
|
||||
if cursor is not None:
|
||||
retrieve_kwargs["cursor"] = cursor
|
||||
if workflow_run.finished_at is not None:
|
||||
retrieve_kwargs["idle_timeout"] = 0
|
||||
source = message_generator.retrieve_events(
|
||||
AppMode.RAG_PIPELINE,
|
||||
workflow_run.id,
|
||||
**retrieve_kwargs,
|
||||
)
|
||||
yield from PipelineGenerator.convert_to_event_stream(source)
|
||||
|
||||
return Response(
|
||||
event_generator(),
|
||||
mimetype="text/event-stream",
|
||||
headers={"Cache-Control": "no-cache", "Connection": "keep-alive"},
|
||||
)
|
||||
|
||||
|
||||
@console_ns.route("/rag/pipelines/<uuid:pipeline_id>/workflow-runs/<uuid:run_id>/node-executions")
|
||||
class RagPipelineWorkflowRunNodeExecutionListApi(Resource):
|
||||
@console_ns.response(
|
||||
|
||||
@@ -159,6 +159,8 @@ class CompletionStopApi(InstalledAppResource):
|
||||
invoke_from=InvokeFrom.EXPLORE,
|
||||
user_id=current_user_id,
|
||||
app_mode=AppMode.value_of(app_model.mode),
|
||||
tenant_id=app_model.tenant_id,
|
||||
app_id=app_model.id,
|
||||
)
|
||||
|
||||
return SimpleResultResponse(result="success").model_dump(mode="json"), 200
|
||||
@@ -255,6 +257,8 @@ class ChatStopApi(InstalledAppResource):
|
||||
invoke_from=InvokeFrom.EXPLORE,
|
||||
user_id=current_user_id,
|
||||
app_mode=app_mode,
|
||||
tenant_id=app_model.tenant_id,
|
||||
app_id=app_model.id,
|
||||
)
|
||||
|
||||
return SimpleResultResponse(result="success").model_dump(mode="json"), 200
|
||||
|
||||
@@ -52,7 +52,6 @@ from controllers.console.remote_files import RemoteFileUploadPayload, upload_rem
|
||||
from controllers.console.wraps import cloud_edition_billing_resource_check, with_current_user
|
||||
from controllers.web.error import InvokeRateLimitError as InvokeRateLimitHttpError
|
||||
from core.app.app_config.common.parameters_mapping import get_parameters_from_feature_dict
|
||||
from core.app.apps.base_app_queue_manager import AppQueueManager
|
||||
from core.app.entities.app_invoke_entities import InvokeFrom
|
||||
from core.errors.error import (
|
||||
ModelCurrentlyNotSupportError,
|
||||
@@ -60,12 +59,10 @@ from core.errors.error import (
|
||||
QuotaExceededError,
|
||||
)
|
||||
from extensions.ext_database import db
|
||||
from extensions.ext_redis import redis_client
|
||||
from fields.base import ResponseModel
|
||||
from fields.conversation_variable_fields import WorkflowConversationVariableResponse
|
||||
from fields.file_fields import FileResponse, FileWithSignedUrl
|
||||
from fields.message_fields import SuggestedQuestionsResponse
|
||||
from graphon.graph_engine.manager import GraphEngineManager
|
||||
from graphon.model_runtime.errors.invoke import InvokeError
|
||||
from libs import helper
|
||||
from libs.helper import dump_response, to_timestamp, uuid_value
|
||||
@@ -77,6 +74,7 @@ from services.account_service import TenantService
|
||||
from services.app_generate_service import AppGenerateService
|
||||
from services.app_ref_service import AppRefService
|
||||
from services.app_service import AppResponseView, AppService
|
||||
from services.app_task_service import AppTaskService
|
||||
from services.audio_service import AudioService
|
||||
from services.dataset_service import DatasetService
|
||||
from services.errors.audio import (
|
||||
@@ -510,7 +508,8 @@ class TrialAppWorkflowRunApi(TrialAppResource):
|
||||
|
||||
class TrialAppWorkflowTaskStopApi(TrialAppResource):
|
||||
@console_ns.response(200, "Success", console_ns.models[SimpleResultResponse.__name__])
|
||||
def post(self, trial_app, task_id: str):
|
||||
@with_current_user
|
||||
def post(self, current_user: Account, trial_app: App, task_id: str):
|
||||
"""
|
||||
Stop workflow task
|
||||
"""
|
||||
@@ -521,12 +520,14 @@ class TrialAppWorkflowTaskStopApi(TrialAppResource):
|
||||
if app_mode != AppMode.WORKFLOW:
|
||||
raise NotWorkflowAppError()
|
||||
|
||||
# Stop using both mechanisms for backward compatibility
|
||||
# Legacy stop flag mechanism (without user check)
|
||||
AppQueueManager.set_stop_flag_no_user_check(task_id)
|
||||
|
||||
# New graph engine command channel mechanism
|
||||
GraphEngineManager(redis_client).send_stop_command(task_id)
|
||||
AppTaskService.stop_task(
|
||||
task_id,
|
||||
InvokeFrom.EXPLORE,
|
||||
current_user.id,
|
||||
app_mode,
|
||||
tenant_id=app_model.tenant_id,
|
||||
app_id=app_model.id,
|
||||
)
|
||||
|
||||
return {"result": "success"}
|
||||
|
||||
|
||||
@@ -17,20 +17,18 @@ from controllers.console.explore.error import NotWorkflowAppError
|
||||
from controllers.console.explore.wraps import InstalledAppResource
|
||||
from controllers.console.wraps import with_current_user
|
||||
from controllers.web.error import InvokeRateLimitError as InvokeRateLimitHttpError
|
||||
from core.app.apps.base_app_queue_manager import AppQueueManager
|
||||
from core.app.entities.app_invoke_entities import InvokeFrom
|
||||
from core.errors.error import (
|
||||
ModelCurrentlyNotSupportError,
|
||||
ProviderTokenNotInitError,
|
||||
QuotaExceededError,
|
||||
)
|
||||
from extensions.ext_redis import redis_client
|
||||
from graphon.graph_engine.manager import GraphEngineManager
|
||||
from graphon.model_runtime.errors.invoke import InvokeError
|
||||
from libs import helper
|
||||
from models import Account
|
||||
from models.model import AppMode, InstalledApp
|
||||
from services.app_generate_service import AppGenerateService
|
||||
from services.app_task_service import AppTaskService
|
||||
from services.errors.llm import InvokeRateLimitError
|
||||
|
||||
from .. import console_ns
|
||||
@@ -92,8 +90,9 @@ class InstalledAppWorkflowRunApi(InstalledAppResource):
|
||||
@console_ns.route("/installed-apps/<uuid:installed_app_id>/workflows/tasks/<string:task_id>/stop")
|
||||
class InstalledAppWorkflowTaskStopApi(InstalledAppResource):
|
||||
@console_ns.response(200, "Success", console_ns.models[SimpleResultResponse.__name__])
|
||||
@with_current_user
|
||||
@with_session(write=False)
|
||||
def post(self, session: Session, installed_app: InstalledApp, task_id: str):
|
||||
def post(self, session: Session, current_user: Account, installed_app: InstalledApp, task_id: str):
|
||||
"""
|
||||
Stop workflow task
|
||||
"""
|
||||
@@ -104,11 +103,13 @@ class InstalledAppWorkflowTaskStopApi(InstalledAppResource):
|
||||
if app_mode != AppMode.WORKFLOW:
|
||||
raise NotWorkflowAppError()
|
||||
|
||||
# Stop using both mechanisms for backward compatibility
|
||||
# Legacy stop flag mechanism (without user check)
|
||||
AppQueueManager.set_stop_flag_no_user_check(task_id)
|
||||
|
||||
# New graph engine command channel mechanism
|
||||
GraphEngineManager(redis_client).send_stop_command(task_id)
|
||||
AppTaskService.stop_task(
|
||||
task_id,
|
||||
InvokeFrom.EXPLORE,
|
||||
current_user.id,
|
||||
app_mode,
|
||||
tenant_id=app_model.tenant_id,
|
||||
app_id=app_model.id,
|
||||
)
|
||||
|
||||
return SimpleResultResponse(result="success").model_dump(mode="json")
|
||||
|
||||
@@ -17,6 +17,7 @@ from controllers.common.errors import InvalidArgumentError, NotFoundError
|
||||
from controllers.common.fields import EventStreamResponse
|
||||
from controllers.common.human_input import HumanInputFormSubmitPayload
|
||||
from controllers.common.schema import register_response_schema_models, register_schema_models
|
||||
from controllers.common.workflow_event_cursor import get_workflow_event_replay_cursor
|
||||
from controllers.console import console_ns
|
||||
from controllers.console.wraps import (
|
||||
account_initialization_required,
|
||||
@@ -30,6 +31,7 @@ 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 core.app.entities.task_entities import StreamEvent
|
||||
from core.workflow.human_input_policy import HumanInputSurface, is_recipient_type_allowed_for_surface
|
||||
from extensions.ext_database import db
|
||||
from libs.login import login_required
|
||||
@@ -39,7 +41,7 @@ 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
|
||||
from services.workflow_event_snapshot_service import build_workflow_event_stream, resolve_workflow_event_task_id
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -187,10 +189,18 @@ class ConsoleWorkflowEventsApi(Resource):
|
||||
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:
|
||||
app_mode = AppMode.value_of(app.mode)
|
||||
if app_mode not in {AppMode.ADVANCED_CHAT, AppMode.WORKFLOW}:
|
||||
raise InvalidArgumentError(f"cannot subscribe to workflow run, workflow_run_id={workflow_run.id}")
|
||||
cursor = get_workflow_event_replay_cursor(request)
|
||||
|
||||
if workflow_run.finished_at is not None and cursor is 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,
|
||||
task_id=resolve_workflow_event_task_id(
|
||||
workflow_run=workflow_run,
|
||||
session_maker=session_maker,
|
||||
),
|
||||
workflow_run=workflow_run,
|
||||
creator_user=user,
|
||||
)
|
||||
@@ -206,29 +216,37 @@ class ConsoleWorkflowEventsApi(Resource):
|
||||
else:
|
||||
msg_generator = MessageGenerator()
|
||||
generator: BaseAppGenerator
|
||||
match app.mode:
|
||||
match app_mode:
|
||||
case AppMode.ADVANCED_CHAT:
|
||||
generator = AdvancedChatAppGenerator()
|
||||
case AppMode.WORKFLOW:
|
||||
generator = WorkflowAppGenerator()
|
||||
case _:
|
||||
raise InvalidArgumentError(f"cannot subscribe to workflow run, workflow_run_id={workflow_run.id}")
|
||||
raise AssertionError("app mode was validated before event stream construction")
|
||||
|
||||
include_state_snapshot = request.args.get("include_state_snapshot", "false").lower() == "true"
|
||||
continue_on_pause = request.args.get("continue_on_pause", "false").lower() == "true"
|
||||
|
||||
def _generate_stream_events():
|
||||
if include_state_snapshot:
|
||||
if include_state_snapshot or cursor is not None:
|
||||
return generator.convert_to_event_stream(
|
||||
build_workflow_event_stream(
|
||||
app_mode=AppMode(app.mode),
|
||||
app_mode=app_mode,
|
||||
workflow_run=workflow_run,
|
||||
tenant_id=workflow_run.tenant_id,
|
||||
app_id=workflow_run.app_id,
|
||||
session_maker=session_maker,
|
||||
close_on_pause=not continue_on_pause,
|
||||
cursor=cursor,
|
||||
)
|
||||
)
|
||||
terminal_events = [StreamEvent.WORKFLOW_FINISHED] if continue_on_pause else None
|
||||
return generator.convert_to_event_stream(
|
||||
msg_generator.retrieve_events(AppMode(app.mode), workflow_run.id),
|
||||
msg_generator.retrieve_events(
|
||||
app_mode,
|
||||
workflow_run.id,
|
||||
terminal_events=terminal_events,
|
||||
)
|
||||
)
|
||||
|
||||
event_generator = _generate_stream_events
|
||||
|
||||
@@ -1,8 +1,11 @@
|
||||
import json
|
||||
import logging
|
||||
from collections.abc import Callable
|
||||
from collections.abc import Callable, Generator
|
||||
from functools import wraps
|
||||
from typing import Any
|
||||
from uuid import UUID
|
||||
|
||||
from flask import request
|
||||
from flask import Response, request
|
||||
from flask_restx import Resource
|
||||
from pydantic import BaseModel, Field
|
||||
from sqlalchemy.orm import Session, sessionmaker
|
||||
@@ -11,6 +14,7 @@ from werkzeug.exceptions import BadRequest, InternalServerError, NotFound
|
||||
from controllers.common.controller_schemas import WorkflowUpdatePayload
|
||||
from controllers.common.fields import GeneratedAppResponse, SimpleResultResponse
|
||||
from controllers.common.schema import query_params_from_model, register_response_schema_models, register_schema_models
|
||||
from controllers.common.workflow_event_cursor import get_workflow_event_replay_cursor
|
||||
from controllers.console import console_ns
|
||||
from controllers.console.app.error import DraftWorkflowNotExist, DraftWorkflowNotSync
|
||||
from controllers.console.app.workflow import (
|
||||
@@ -40,27 +44,34 @@ from controllers.console.wraps import (
|
||||
setup_required,
|
||||
with_current_user,
|
||||
)
|
||||
from core.app.apps.base_app_queue_manager import AppQueueManager
|
||||
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 core.app.entities.app_invoke_entities import InvokeFrom
|
||||
from core.app.entities.task_entities import StreamEvent
|
||||
from core.workflow.human_input_policy import HumanInputSurface
|
||||
from extensions.ext_database import db
|
||||
from extensions.ext_redis import redis_client
|
||||
from fields.workflow_run_fields import (
|
||||
WorkflowRunDetailResponse,
|
||||
WorkflowRunNodeExecutionListResponse,
|
||||
WorkflowRunNodeExecutionResponse,
|
||||
WorkflowRunPaginationResponse,
|
||||
)
|
||||
from graphon.graph_engine.manager import GraphEngineManager
|
||||
from libs import helper
|
||||
from libs.helper import TimestampField
|
||||
from libs.login import current_account_with_tenant, login_required
|
||||
from models import Account
|
||||
from models.enums import CreatorUserRole
|
||||
from models.model import AppMode
|
||||
from models.snippet import CustomizedSnippet
|
||||
from repositories.factory import DifyAPIRepositoryFactory
|
||||
from services.agent.retirement_service import WorkflowAgentRetirementService
|
||||
from services.agent.workflow_publish_service import WorkflowAgentPublishService
|
||||
from services.app_task_service import AppTaskService
|
||||
from services.errors.app import IsDraftWorkflowError, WorkflowHashNotEqualError, WorkflowNotFoundError
|
||||
from services.snippet_generate_service import SnippetGenerateService
|
||||
from services.snippet_service import SnippetService
|
||||
from services.workflow_event_snapshot_service import build_workflow_event_stream, resolve_workflow_event_task_id
|
||||
from tasks.collect_agent_resources_task import enqueue_agent_resource_collection
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -532,6 +543,87 @@ class SnippetWorkflowRunDetailApi(Resource):
|
||||
return WorkflowRunDetailResponse.model_validate(workflow_run, from_attributes=True).model_dump(mode="json")
|
||||
|
||||
|
||||
@console_ns.route("/snippets/<uuid:snippet_id>/workflow-runs/<uuid:run_id>/events")
|
||||
class SnippetWorkflowRunEventsApi(Resource):
|
||||
"""Reconnect to one durable Snippet workflow event stream."""
|
||||
|
||||
@setup_required
|
||||
@login_required
|
||||
@account_initialization_required
|
||||
@with_current_user
|
||||
@get_snippet
|
||||
def get(self, current_user: Account, snippet: CustomizedSnippet, run_id: UUID):
|
||||
session_maker = _snippet_session_maker()
|
||||
repository = DifyAPIRepositoryFactory.create_api_workflow_run_repository(session_maker)
|
||||
workflow_run = repository.get_workflow_run_by_id_and_tenant_id(
|
||||
tenant_id=snippet.tenant_id,
|
||||
run_id=str(run_id),
|
||||
)
|
||||
if (
|
||||
workflow_run is None
|
||||
or workflow_run.app_id != snippet.id
|
||||
or workflow_run.created_by_role != CreatorUserRole.ACCOUNT
|
||||
or workflow_run.created_by != current_user.id
|
||||
):
|
||||
raise NotFound("Workflow run not found")
|
||||
|
||||
cursor = get_workflow_event_replay_cursor(request)
|
||||
include_state_snapshot = request.args.get("include_state_snapshot", "false").lower() == "true"
|
||||
continue_on_pause = request.args.get("continue_on_pause", "false").lower() == "true"
|
||||
|
||||
if workflow_run.finished_at is not None and cursor is None:
|
||||
finished = WorkflowResponseConverter.workflow_run_result_to_finish_response(
|
||||
task_id=resolve_workflow_event_task_id(
|
||||
workflow_run=workflow_run,
|
||||
session_maker=session_maker,
|
||||
),
|
||||
workflow_run=workflow_run,
|
||||
creator_user=current_user,
|
||||
)
|
||||
payload = finished.model_dump(mode="json")
|
||||
payload["event"] = finished.event.value
|
||||
|
||||
def event_generator() -> Generator[str, None, None]:
|
||||
yield f"data: {json.dumps(payload)}\n\n"
|
||||
|
||||
else:
|
||||
message_generator = MessageGenerator()
|
||||
|
||||
def event_generator() -> Generator[str, None, None]:
|
||||
if include_state_snapshot or cursor is not None:
|
||||
source = build_workflow_event_stream(
|
||||
app_mode=AppMode.WORKFLOW,
|
||||
workflow_run=workflow_run,
|
||||
tenant_id=snippet.tenant_id,
|
||||
app_id=snippet.id,
|
||||
session_maker=session_maker,
|
||||
human_input_surface=HumanInputSurface.CONSOLE,
|
||||
close_on_pause=not continue_on_pause,
|
||||
cursor=cursor,
|
||||
)
|
||||
else:
|
||||
terminal_events = [StreamEvent.WORKFLOW_FINISHED] if continue_on_pause else None
|
||||
retrieve_kwargs: dict[str, Any] = {"terminal_events": terminal_events}
|
||||
if cursor is not None:
|
||||
retrieve_kwargs["cursor"] = cursor
|
||||
if workflow_run.finished_at is not None:
|
||||
retrieve_kwargs["idle_timeout"] = 0
|
||||
source = message_generator.retrieve_events(
|
||||
AppMode.WORKFLOW,
|
||||
workflow_run.id,
|
||||
**retrieve_kwargs,
|
||||
)
|
||||
yield from WorkflowAppGenerator.convert_to_event_stream(
|
||||
SnippetGenerateService.filter_virtual_start_events(source)
|
||||
)
|
||||
|
||||
return Response(
|
||||
event_generator(),
|
||||
mimetype="text/event-stream",
|
||||
headers={"Cache-Control": "no-cache", "Connection": "keep-alive"},
|
||||
)
|
||||
|
||||
|
||||
@console_ns.route("/snippets/<uuid:snippet_id>/workflow-runs/<uuid:run_id>/node-executions")
|
||||
class SnippetWorkflowRunNodeExecutionsApi(Resource):
|
||||
@console_ns.doc("list_snippet_workflow_run_node_executions")
|
||||
@@ -787,20 +879,23 @@ class SnippetWorkflowTaskStopApi(Resource):
|
||||
@setup_required
|
||||
@login_required
|
||||
@account_initialization_required
|
||||
@with_current_user
|
||||
@get_snippet
|
||||
@edit_permission_required
|
||||
def post(self, snippet: CustomizedSnippet, task_id: str):
|
||||
def post(self, current_user: Account, snippet: CustomizedSnippet, task_id: str):
|
||||
"""
|
||||
Stop a running snippet workflow task.
|
||||
|
||||
Uses both the legacy stop flag mechanism and the graph engine
|
||||
command channel for backward compatibility.
|
||||
"""
|
||||
# Stop using both mechanisms for backward compatibility
|
||||
# Legacy stop flag mechanism (without user check)
|
||||
AppQueueManager.set_stop_flag_no_user_check(task_id)
|
||||
|
||||
# New graph engine command channel mechanism
|
||||
GraphEngineManager(redis_client).send_stop_command(task_id)
|
||||
AppTaskService.stop_task(
|
||||
task_id,
|
||||
InvokeFrom.DEBUGGER,
|
||||
current_user.id,
|
||||
AppMode.WORKFLOW,
|
||||
tenant_id=snippet.tenant_id,
|
||||
app_id=snippet.id,
|
||||
)
|
||||
|
||||
return {"result": "success"}
|
||||
|
||||
@@ -138,6 +138,7 @@ class WorkflowRunData(BaseModel):
|
||||
outputs: dict[str, Any] = Field(default_factory=dict)
|
||||
error: str | None = None
|
||||
elapsed_time: float | None = None
|
||||
handoff_duration: float = 0.0
|
||||
total_tokens: int | None = None
|
||||
total_steps: int | None = None
|
||||
created_at: int | None = None
|
||||
|
||||
@@ -27,7 +27,7 @@ from controllers.openapi._audit import emit_app_run
|
||||
from controllers.openapi._contract import accepts, returns
|
||||
from controllers.openapi._models import AppRunRequest, TaskStopResponse
|
||||
from controllers.openapi.auth.composition import auth_router
|
||||
from controllers.openapi.auth.data import AuthData, RBACRequirement
|
||||
from controllers.openapi.auth.data import AuthData, CallerKind, RBACRequirement
|
||||
from controllers.service_api.app.error import (
|
||||
AppUnavailableError,
|
||||
CompletionRequestError,
|
||||
@@ -37,7 +37,6 @@ from controllers.service_api.app.error import (
|
||||
ProviderQuotaExceededError,
|
||||
)
|
||||
from controllers.web.error import InvokeRateLimitError as InvokeRateLimitHttpError
|
||||
from core.app.apps.base_app_queue_manager import AppQueueManager
|
||||
from core.app.entities.app_invoke_entities import InvokeFrom
|
||||
from core.errors.error import (
|
||||
AppInvokeQuotaExceededError,
|
||||
@@ -45,13 +44,13 @@ from core.errors.error import (
|
||||
ProviderTokenNotInitError,
|
||||
QuotaExceededError,
|
||||
)
|
||||
from extensions.ext_redis import redis_client
|
||||
from graphon.graph_engine.manager import GraphEngineManager
|
||||
from graphon.model_runtime.errors.invoke import InvokeError
|
||||
from libs import helper
|
||||
from libs.oauth_bearer import Scope
|
||||
from models.enums import CreatorUserRole
|
||||
from models.model import App, AppMode
|
||||
from services.app_generate_service import AppGenerateService
|
||||
from services.app_task_service import AppTaskService
|
||||
from services.errors.app import (
|
||||
IsDraftWorkflowError,
|
||||
WorkflowIdFormatError,
|
||||
@@ -183,6 +182,15 @@ class AppRunTaskStopApi(Resource):
|
||||
@returns(200, TaskStopResponse, description="Task stopped")
|
||||
def post(self, app_id: str, task_id: str, *, auth_data: AuthData):
|
||||
app_model, caller, caller_kind = auth_data.require_app_context()
|
||||
AppQueueManager.set_stop_flag_no_user_check(task_id)
|
||||
GraphEngineManager(redis_client).send_stop_command(task_id)
|
||||
AppTaskService.stop_task(
|
||||
task_id,
|
||||
InvokeFrom.OPENAPI,
|
||||
caller.id,
|
||||
AppMode.value_of(app_model.mode),
|
||||
tenant_id=app_model.tenant_id,
|
||||
app_id=app_model.id,
|
||||
created_by_role=(
|
||||
CreatorUserRole.ACCOUNT if caller_kind == CallerKind.ACCOUNT else CreatorUserRole.END_USER
|
||||
),
|
||||
)
|
||||
return TaskStopResponse(result="success")
|
||||
|
||||
@@ -10,6 +10,7 @@ from __future__ import annotations
|
||||
|
||||
import json
|
||||
from collections.abc import Generator
|
||||
from typing import Any
|
||||
|
||||
from flask import Response, request
|
||||
from flask_restx import Resource
|
||||
@@ -19,6 +20,7 @@ from werkzeug.exceptions import NotFound, UnprocessableEntity
|
||||
|
||||
from controllers.common.fields import EventStreamResponse
|
||||
from controllers.common.schema import query_params_from_model
|
||||
from controllers.common.workflow_event_cursor import get_workflow_event_replay_cursor
|
||||
from controllers.common.wraps import RBACPermission, RBACResourceScope
|
||||
from controllers.openapi import openapi_ns
|
||||
from controllers.openapi.auth.composition import auth_router
|
||||
@@ -35,12 +37,16 @@ from libs.oauth_bearer import Scope
|
||||
from models.enums import CreatorUserRole
|
||||
from models.model import AppMode
|
||||
from repositories.factory import DifyAPIRepositoryFactory
|
||||
from services.workflow_event_snapshot_service import build_workflow_event_stream
|
||||
from services.workflow_event_snapshot_service import build_workflow_event_stream, resolve_workflow_event_task_id
|
||||
|
||||
|
||||
class WorkflowEventsQuery(BaseModel):
|
||||
include_state_snapshot: bool = Field(default=False, description="Whether to include workflow state snapshots")
|
||||
continue_on_pause: bool = Field(default=False, description="Whether to keep the event stream open on pause")
|
||||
cursor: str | None = Field(
|
||||
default=None,
|
||||
description="Replay events strictly after this SSE event ID. Last-Event-ID header takes precedence.",
|
||||
)
|
||||
|
||||
|
||||
@openapi_ns.route("/apps/<string:app_id>/tasks/<string:task_id>/events")
|
||||
@@ -78,10 +84,14 @@ class OpenApiWorkflowEventsApi(Resource):
|
||||
raise NotFound("Workflow run not found")
|
||||
|
||||
workflow_run_entity = workflow_run
|
||||
cursor = get_workflow_event_replay_cursor(request)
|
||||
|
||||
if workflow_run_entity.finished_at is not None:
|
||||
if workflow_run_entity.finished_at is not None and cursor is None:
|
||||
response = WorkflowResponseConverter.workflow_run_result_to_finish_response(
|
||||
task_id=workflow_run_entity.id,
|
||||
task_id=resolve_workflow_event_task_id(
|
||||
workflow_run=workflow_run_entity,
|
||||
session_maker=session_maker,
|
||||
),
|
||||
workflow_run=workflow_run_entity,
|
||||
creator_user=caller,
|
||||
)
|
||||
@@ -102,10 +112,10 @@ class OpenApiWorkflowEventsApi(Resource):
|
||||
|
||||
include_state_snapshot = request.args.get("include_state_snapshot", "false").lower() == "true"
|
||||
continue_on_pause = request.args.get("continue_on_pause", "false").lower() == "true"
|
||||
terminal_events: list[StreamEvent] | None = [] if continue_on_pause else None
|
||||
terminal_events: list[StreamEvent] | None = [StreamEvent.WORKFLOW_FINISHED] if continue_on_pause else None
|
||||
|
||||
def _generate_stream_events():
|
||||
if include_state_snapshot:
|
||||
if include_state_snapshot or cursor is not None:
|
||||
return generator.convert_to_event_stream(
|
||||
build_workflow_event_stream(
|
||||
app_mode=app_mode,
|
||||
@@ -115,14 +125,16 @@ class OpenApiWorkflowEventsApi(Resource):
|
||||
session_maker=session_maker,
|
||||
human_input_surface=HumanInputSurface.OPENAPI,
|
||||
close_on_pause=not continue_on_pause,
|
||||
cursor=cursor,
|
||||
)
|
||||
)
|
||||
retrieve_kwargs: dict[str, Any] = {"terminal_events": terminal_events}
|
||||
if cursor is not None:
|
||||
retrieve_kwargs["cursor"] = cursor
|
||||
if workflow_run_entity.finished_at is not None:
|
||||
retrieve_kwargs["idle_timeout"] = 0
|
||||
return generator.convert_to_event_stream(
|
||||
msg_generator.retrieve_events(
|
||||
app_mode,
|
||||
workflow_run_entity.id,
|
||||
terminal_events=terminal_events,
|
||||
),
|
||||
msg_generator.retrieve_events(app_mode, workflow_run_entity.id, **retrieve_kwargs),
|
||||
)
|
||||
|
||||
event_generator = _generate_stream_events
|
||||
|
||||
@@ -304,6 +304,8 @@ class CompletionStopApi(Resource):
|
||||
invoke_from=InvokeFrom.SERVICE_API,
|
||||
user_id=end_user.id,
|
||||
app_mode=AppMode.value_of(app_model.mode),
|
||||
tenant_id=app_model.tenant_id,
|
||||
app_id=app_model.id,
|
||||
)
|
||||
|
||||
return SimpleResultResponse(result="success").model_dump(mode="json"), 200
|
||||
@@ -480,6 +482,8 @@ class ChatStopApi(Resource):
|
||||
invoke_from=InvokeFrom.SERVICE_API,
|
||||
user_id=end_user.id,
|
||||
app_mode=app_mode,
|
||||
tenant_id=app_model.tenant_id,
|
||||
app_id=app_model.id,
|
||||
)
|
||||
|
||||
return SimpleResultResponse(result="success").model_dump(mode="json"), 200
|
||||
|
||||
@@ -37,7 +37,6 @@ from controllers.service_api.schema import (
|
||||
)
|
||||
from controllers.service_api.wraps import FetchUserArg, WhereisUserArg, validate_app_token
|
||||
from controllers.web.error import InvokeRateLimitError as InvokeRateLimitHttpError
|
||||
from core.app.apps.base_app_queue_manager import AppQueueManager
|
||||
from core.app.entities.app_invoke_entities import InvokeFrom
|
||||
from core.errors.error import (
|
||||
ModelCurrentlyNotSupportError,
|
||||
@@ -47,18 +46,17 @@ from core.errors.error import (
|
||||
from core.helper.trace_id_helper import get_external_trace_id, get_trace_session_id, omit_trace_session_id_from_payload
|
||||
from enums.cloud_plan import CloudPlan
|
||||
from extensions.ext_database import db
|
||||
from extensions.ext_redis import redis_client
|
||||
from fields.base import ResponseModel
|
||||
from fields.end_user_fields import SimpleEndUser
|
||||
from fields.member_fields import SimpleAccountResponse
|
||||
from graphon.enums import WorkflowExecutionStatus
|
||||
from graphon.graph_engine.manager import GraphEngineManager
|
||||
from graphon.model_runtime.errors.invoke import InvokeError
|
||||
from libs import helper
|
||||
from libs.helper import dump_response, to_timestamp
|
||||
from models.model import App, AppMode, EndUser
|
||||
from repositories.factory import DifyAPIRepositoryFactory
|
||||
from services.app_generate_service import AppGenerateService
|
||||
from services.app_task_service import AppTaskService
|
||||
from services.billing_service import BillingService
|
||||
from services.errors.app import IsDraftWorkflowError, WorkflowIdFormatError, WorkflowNotFoundError
|
||||
from services.errors.llm import InvokeRateLimitError
|
||||
@@ -125,6 +123,7 @@ class WorkflowRunResponse(ResponseModel):
|
||||
created_at: int | None = None
|
||||
finished_at: int | None = None
|
||||
elapsed_time: float | int | None = None
|
||||
handoff_duration: float = 0.0
|
||||
|
||||
@field_validator("status", mode="before")
|
||||
@classmethod
|
||||
@@ -161,6 +160,7 @@ class WorkflowRunForLogResponse(ResponseModel):
|
||||
triggered_from: str | None = None
|
||||
error: str | None = None
|
||||
elapsed_time: float | int | None = None
|
||||
handoff_duration: float = 0.0
|
||||
total_tokens: int | None = None
|
||||
total_steps: int | None = None
|
||||
created_at: int | None = None
|
||||
@@ -539,12 +539,14 @@ class WorkflowTaskStopApi(Resource):
|
||||
if app_mode != AppMode.WORKFLOW:
|
||||
raise NotWorkflowAppError()
|
||||
|
||||
# Stop using both mechanisms for backward compatibility
|
||||
# Legacy stop flag mechanism (without user check)
|
||||
AppQueueManager.set_stop_flag_no_user_check(task_id)
|
||||
|
||||
# New graph engine command channel mechanism
|
||||
GraphEngineManager(redis_client).send_stop_command(task_id)
|
||||
AppTaskService.stop_task(
|
||||
task_id,
|
||||
InvokeFrom.SERVICE_API,
|
||||
end_user.id,
|
||||
app_mode,
|
||||
tenant_id=app_model.tenant_id,
|
||||
app_id=app_model.id,
|
||||
)
|
||||
|
||||
return SimpleResultResponse(result="success").model_dump()
|
||||
|
||||
|
||||
@@ -4,6 +4,7 @@ Service API workflow resume event stream endpoints.
|
||||
|
||||
import json
|
||||
from collections.abc import Generator
|
||||
from typing import Any
|
||||
|
||||
from flask import Response, request
|
||||
from flask_restx import Resource
|
||||
@@ -13,6 +14,7 @@ from werkzeug.exceptions import NotFound
|
||||
|
||||
from controllers.common.fields import EventStreamResponse
|
||||
from controllers.common.schema import query_params_from_model, register_response_schema_model, register_schema_models
|
||||
from controllers.common.workflow_event_cursor import get_workflow_event_replay_cursor
|
||||
from controllers.service_api import service_api_ns
|
||||
from controllers.service_api.app.error import NotWorkflowAppError
|
||||
from controllers.service_api.schema import event_stream_response
|
||||
@@ -28,7 +30,7 @@ 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
|
||||
from services.workflow_event_snapshot_service import build_workflow_event_stream, resolve_workflow_event_task_id
|
||||
|
||||
|
||||
class WorkflowEventsQuery(BaseModel):
|
||||
@@ -51,6 +53,10 @@ class WorkflowEventsQuery(BaseModel):
|
||||
"first pause."
|
||||
),
|
||||
)
|
||||
cursor: str | None = Field(
|
||||
default=None,
|
||||
description="Replay events strictly after this SSE event ID. Last-Event-ID header takes precedence.",
|
||||
)
|
||||
|
||||
|
||||
register_schema_models(service_api_ns, WorkflowEventsQuery)
|
||||
@@ -71,8 +77,9 @@ class WorkflowEventsApi(Resource):
|
||||
tags=["Chatflows", "Workflows"],
|
||||
responses={
|
||||
200: (
|
||||
"Server-Sent Events stream. Each event is delivered as `data: {JSON}\\n\\n`. Event payloads "
|
||||
"follow the same schemas as the original streaming response."
|
||||
"Server-Sent Events stream. Durable events are delivered as "
|
||||
"`id: {cursor}\\ndata: {JSON}\\n\\n`; reconnect with Last-Event-ID. Event payloads follow the "
|
||||
"same schemas as the original streaming response."
|
||||
),
|
||||
400: "`not_workflow_app` : Please check if your app mode matches the right API route.",
|
||||
404: "`not_found` : Workflow run not found.",
|
||||
@@ -117,10 +124,14 @@ class WorkflowEventsApi(Resource):
|
||||
raise NotFound("Workflow run not found")
|
||||
|
||||
workflow_run_entity = workflow_run
|
||||
cursor = get_workflow_event_replay_cursor(request)
|
||||
|
||||
if workflow_run_entity.finished_at is not None:
|
||||
if workflow_run_entity.finished_at is not None and cursor is None:
|
||||
response = WorkflowResponseConverter.workflow_run_result_to_finish_response(
|
||||
task_id=workflow_run_entity.id,
|
||||
task_id=resolve_workflow_event_task_id(
|
||||
workflow_run=workflow_run_entity,
|
||||
session_maker=session_maker,
|
||||
),
|
||||
workflow_run=workflow_run_entity,
|
||||
creator_user=end_user,
|
||||
)
|
||||
@@ -144,10 +155,10 @@ class WorkflowEventsApi(Resource):
|
||||
|
||||
include_state_snapshot = request.args.get("include_state_snapshot", "false").lower() == "true"
|
||||
continue_on_pause = request.args.get("continue_on_pause", "false").lower() == "true"
|
||||
terminal_events: list[StreamEvent] | None = [] if continue_on_pause else None
|
||||
terminal_events: list[StreamEvent] | None = [StreamEvent.WORKFLOW_FINISHED] if continue_on_pause else None
|
||||
|
||||
def _generate_stream_events():
|
||||
if include_state_snapshot:
|
||||
if include_state_snapshot or cursor is not None:
|
||||
return generator.convert_to_event_stream(
|
||||
build_workflow_event_stream(
|
||||
app_mode=app_mode,
|
||||
@@ -157,14 +168,16 @@ class WorkflowEventsApi(Resource):
|
||||
session_maker=session_maker,
|
||||
human_input_surface=HumanInputSurface.SERVICE_API,
|
||||
close_on_pause=not continue_on_pause,
|
||||
cursor=cursor,
|
||||
)
|
||||
)
|
||||
retrieve_kwargs: dict[str, Any] = {"terminal_events": terminal_events}
|
||||
if cursor is not None:
|
||||
retrieve_kwargs["cursor"] = cursor
|
||||
if workflow_run_entity.finished_at is not None:
|
||||
retrieve_kwargs["idle_timeout"] = 0
|
||||
return generator.convert_to_event_stream(
|
||||
msg_generator.retrieve_events(
|
||||
app_mode,
|
||||
workflow_run_entity.id,
|
||||
terminal_events=terminal_events,
|
||||
),
|
||||
msg_generator.retrieve_events(app_mode, workflow_run_entity.id, **retrieve_kwargs),
|
||||
)
|
||||
|
||||
event_generator = _generate_stream_events
|
||||
|
||||
@@ -1,9 +1,10 @@
|
||||
from collections.abc import Generator
|
||||
import json
|
||||
from collections.abc import Iterator
|
||||
from datetime import datetime
|
||||
from typing import Any
|
||||
from uuid import UUID
|
||||
|
||||
from flask import request
|
||||
from flask import Response, request
|
||||
from pydantic import BaseModel, Field, RootModel, field_validator
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.orm import Session
|
||||
@@ -18,13 +19,18 @@ from controllers.common.schema import (
|
||||
register_response_schema_models,
|
||||
register_schema_model,
|
||||
)
|
||||
from controllers.common.workflow_event_cursor import get_workflow_event_replay_cursor
|
||||
from controllers.console.app.wraps import with_session
|
||||
from controllers.service_api import service_api_ns
|
||||
from controllers.service_api.dataset.error import PipelineRunError
|
||||
from controllers.service_api.schema import event_stream_response, json_or_event_stream_response, multipart_file_params
|
||||
from controllers.service_api.wraps import DatasetApiResource
|
||||
from core.app.apps.common.workflow_response_converter import WorkflowResponseConverter
|
||||
from core.app.apps.message_generator import MessageGenerator
|
||||
from core.app.apps.pipeline.pipeline_generator import PipelineGenerator
|
||||
from core.app.entities.app_invoke_entities import InvokeFrom
|
||||
from core.app.entities.task_entities import StreamEvent
|
||||
from core.workflow.human_input_policy import HumanInputSurface
|
||||
from fields.base import ResponseModel
|
||||
from libs import helper
|
||||
from libs.helper import dump_response
|
||||
@@ -32,6 +38,8 @@ from libs.login import current_user
|
||||
from models import Account
|
||||
from models.dataset import Dataset, Pipeline
|
||||
from models.engine import db
|
||||
from models.enums import CreatorUserRole
|
||||
from models.model import AppMode
|
||||
from services.errors.file import FileTooLargeError, UnsupportedFileTypeError
|
||||
from services.file_service import FileService
|
||||
from services.rag_pipeline.entity.pipeline_service_api_entities import (
|
||||
@@ -41,6 +49,7 @@ from services.rag_pipeline.entity.pipeline_service_api_entities import (
|
||||
)
|
||||
from services.rag_pipeline.pipeline_generate_service import PipelineGenerateService
|
||||
from services.rag_pipeline.rag_pipeline import RagPipelineService
|
||||
from services.workflow_event_snapshot_service import build_workflow_event_stream, resolve_workflow_event_task_id
|
||||
|
||||
|
||||
class DatasourceNodeRunPayload(BaseModel):
|
||||
@@ -284,7 +293,7 @@ class PipelineRunApi(DatasetApiResource):
|
||||
rag_pipeline_service = RagPipelineService(session)
|
||||
pipeline = rag_pipeline_service.get_pipeline(tenant_id=tenant_id, dataset_id=dataset_id_str)
|
||||
try:
|
||||
response: dict[Any, Any] | Generator[str, Any, None] = PipelineGenerateService.generate(
|
||||
response: dict[Any, Any] | Iterator[str] = PipelineGenerateService.generate(
|
||||
session=session,
|
||||
pipeline=pipeline,
|
||||
user=current_user,
|
||||
@@ -299,6 +308,88 @@ class PipelineRunApi(DatasetApiResource):
|
||||
raise PipelineRunError(description=str(ex))
|
||||
|
||||
|
||||
@service_api_ns.route("/datasets/<uuid:dataset_id>/pipeline/workflow-runs/<uuid:run_id>/events")
|
||||
class PipelineWorkflowRunEventsApi(DatasetApiResource):
|
||||
"""Reconnect to the durable event stream of a service-API pipeline run."""
|
||||
|
||||
@event_stream_response(service_api_ns)
|
||||
def get(self, tenant_id: str, dataset_id: UUID, run_id: UUID):
|
||||
if not isinstance(current_user, Account):
|
||||
raise Forbidden()
|
||||
|
||||
rag_pipeline_service = RagPipelineService(db.session())
|
||||
try:
|
||||
pipeline = rag_pipeline_service.get_pipeline(
|
||||
tenant_id=tenant_id,
|
||||
dataset_id=str(dataset_id),
|
||||
)
|
||||
except ValueError as error:
|
||||
raise NotFound(str(error)) from error
|
||||
session_maker = rag_pipeline_service.session_maker
|
||||
workflow_run = rag_pipeline_service.get_rag_pipeline_workflow_run(pipeline=pipeline, run_id=str(run_id))
|
||||
if (
|
||||
workflow_run is None
|
||||
or workflow_run.app_id != pipeline.id
|
||||
or workflow_run.created_by_role != CreatorUserRole.ACCOUNT
|
||||
or workflow_run.created_by != current_user.id
|
||||
):
|
||||
raise NotFound("Workflow run not found")
|
||||
|
||||
cursor = get_workflow_event_replay_cursor(request)
|
||||
include_state_snapshot = request.args.get("include_state_snapshot", "false").lower() == "true"
|
||||
continue_on_pause = request.args.get("continue_on_pause", "false").lower() == "true"
|
||||
|
||||
if workflow_run.finished_at is not None and cursor is None:
|
||||
finished = WorkflowResponseConverter.workflow_run_result_to_finish_response(
|
||||
task_id=resolve_workflow_event_task_id(
|
||||
workflow_run=workflow_run,
|
||||
session_maker=session_maker,
|
||||
),
|
||||
workflow_run=workflow_run,
|
||||
creator_user=current_user,
|
||||
)
|
||||
payload = finished.model_dump(mode="json")
|
||||
payload["event"] = finished.event.value
|
||||
|
||||
def event_generator():
|
||||
yield f"data: {json.dumps(payload)}\n\n"
|
||||
|
||||
else:
|
||||
message_generator = MessageGenerator()
|
||||
|
||||
def event_generator():
|
||||
if include_state_snapshot or cursor is not None:
|
||||
source = build_workflow_event_stream(
|
||||
app_mode=AppMode.RAG_PIPELINE,
|
||||
workflow_run=workflow_run,
|
||||
tenant_id=tenant_id,
|
||||
app_id=pipeline.id,
|
||||
session_maker=session_maker,
|
||||
human_input_surface=HumanInputSurface.SERVICE_API,
|
||||
close_on_pause=not continue_on_pause,
|
||||
cursor=cursor,
|
||||
)
|
||||
else:
|
||||
terminal_events = [StreamEvent.WORKFLOW_FINISHED] if continue_on_pause else None
|
||||
retrieve_kwargs: dict[str, Any] = {"terminal_events": terminal_events}
|
||||
if cursor is not None:
|
||||
retrieve_kwargs["cursor"] = cursor
|
||||
if workflow_run.finished_at is not None:
|
||||
retrieve_kwargs["idle_timeout"] = 0
|
||||
source = message_generator.retrieve_events(
|
||||
AppMode.RAG_PIPELINE,
|
||||
workflow_run.id,
|
||||
**retrieve_kwargs,
|
||||
)
|
||||
yield from PipelineGenerator.convert_to_event_stream(source)
|
||||
|
||||
return Response(
|
||||
event_generator(),
|
||||
mimetype="text/event-stream",
|
||||
headers={"Cache-Control": "no-cache", "Connection": "keep-alive"},
|
||||
)
|
||||
|
||||
|
||||
@service_api_ns.route("/datasets/pipeline/file-upload")
|
||||
class KnowledgebasePipelineFileUploadApi(DatasetApiResource):
|
||||
"""Resource for uploading a file to a knowledgebase pipeline."""
|
||||
|
||||
@@ -183,6 +183,8 @@ class CompletionStopApi(WebApiResource):
|
||||
invoke_from=InvokeFrom.WEB_APP,
|
||||
user_id=end_user.id,
|
||||
app_mode=AppMode.value_of(app_model.mode),
|
||||
tenant_id=app_model.tenant_id,
|
||||
app_id=app_model.id,
|
||||
)
|
||||
|
||||
return SimpleResultResponse(result="success").model_dump(mode="json"), 200
|
||||
@@ -289,6 +291,8 @@ class ChatStopApi(WebApiResource):
|
||||
invoke_from=InvokeFrom.WEB_APP,
|
||||
user_id=end_user.id,
|
||||
app_mode=app_mode,
|
||||
tenant_id=app_model.tenant_id,
|
||||
app_id=app_model.id,
|
||||
)
|
||||
|
||||
return SimpleResultResponse(result="success").model_dump(mode="json"), 200
|
||||
|
||||
@@ -17,19 +17,17 @@ from controllers.web.error import (
|
||||
)
|
||||
from controllers.web.error import InvokeRateLimitError as InvokeRateLimitHttpError
|
||||
from controllers.web.wraps import WebApiResource
|
||||
from core.app.apps.base_app_queue_manager import AppQueueManager
|
||||
from core.app.entities.app_invoke_entities import InvokeFrom
|
||||
from core.errors.error import (
|
||||
ModelCurrentlyNotSupportError,
|
||||
ProviderTokenNotInitError,
|
||||
QuotaExceededError,
|
||||
)
|
||||
from extensions.ext_redis import redis_client
|
||||
from graphon.graph_engine.manager import GraphEngineManager
|
||||
from graphon.model_runtime.errors.invoke import InvokeError
|
||||
from libs import helper
|
||||
from models.model import App, AppMode, EndUser
|
||||
from services.app_generate_service import AppGenerateService
|
||||
from services.app_task_service import AppTaskService
|
||||
from services.errors.llm import InvokeRateLimitError
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -123,11 +121,13 @@ class WorkflowTaskStopApi(WebApiResource):
|
||||
if app_mode != AppMode.WORKFLOW:
|
||||
raise NotWorkflowAppError()
|
||||
|
||||
# Stop using both mechanisms for backward compatibility
|
||||
# Legacy stop flag mechanism (without user check)
|
||||
AppQueueManager.set_stop_flag_no_user_check(task_id)
|
||||
|
||||
# New graph engine command channel mechanism
|
||||
GraphEngineManager(redis_client).send_stop_command(task_id)
|
||||
AppTaskService.stop_task(
|
||||
task_id,
|
||||
InvokeFrom.WEB_APP,
|
||||
end_user.id,
|
||||
app_mode,
|
||||
tenant_id=app_model.tenant_id,
|
||||
app_id=app_model.id,
|
||||
)
|
||||
|
||||
return SimpleResultResponse(result="success").model_dump(mode="json")
|
||||
|
||||
@@ -11,6 +11,7 @@ from sqlalchemy.orm import sessionmaker
|
||||
from controllers.common.errors import InvalidArgumentError, NotFoundError
|
||||
from controllers.common.fields import EventStreamResponse
|
||||
from controllers.common.schema import register_response_schema_model
|
||||
from controllers.common.workflow_event_cursor import get_workflow_event_replay_cursor
|
||||
from controllers.web import api, web_ns
|
||||
from controllers.web.wraps import WebApiResource
|
||||
from core.app.apps.advanced_chat.app_generator import AdvancedChatAppGenerator
|
||||
@@ -18,11 +19,12 @@ 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 core.app.entities.task_entities import StreamEvent
|
||||
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
|
||||
from services.workflow_event_snapshot_service import build_workflow_event_stream, resolve_workflow_event_task_id
|
||||
|
||||
register_response_schema_model(web_ns, EventStreamResponse)
|
||||
|
||||
@@ -59,9 +61,17 @@ class WorkflowEventsApi(WebApiResource):
|
||||
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:
|
||||
app_mode = AppMode.value_of(app_model.mode)
|
||||
if app_mode not in {AppMode.ADVANCED_CHAT, AppMode.WORKFLOW}:
|
||||
raise InvalidArgumentError(f"cannot subscribe to workflow run, workflow_run_id={workflow_run.id}")
|
||||
cursor = get_workflow_event_replay_cursor(request)
|
||||
|
||||
if workflow_run.finished_at is not None and cursor is None:
|
||||
response = WorkflowResponseConverter.workflow_run_result_to_finish_response(
|
||||
task_id=workflow_run.id,
|
||||
task_id=resolve_workflow_event_task_id(
|
||||
workflow_run=workflow_run,
|
||||
session_maker=session_maker,
|
||||
),
|
||||
workflow_run=workflow_run,
|
||||
creator_user=end_user,
|
||||
)
|
||||
@@ -74,7 +84,6 @@ class WorkflowEventsApi(WebApiResource):
|
||||
|
||||
event_generator = _generate_finished_events
|
||||
else:
|
||||
app_mode = AppMode.value_of(app_model.mode)
|
||||
msg_generator = MessageGenerator()
|
||||
generator: BaseAppGenerator
|
||||
match app_mode:
|
||||
@@ -83,12 +92,13 @@ class WorkflowEventsApi(WebApiResource):
|
||||
case AppMode.WORKFLOW:
|
||||
generator = WorkflowAppGenerator()
|
||||
case _:
|
||||
raise InvalidArgumentError(f"cannot subscribe to workflow run, workflow_run_id={workflow_run.id}")
|
||||
raise AssertionError("app mode was validated before event stream construction")
|
||||
|
||||
include_state_snapshot = request.args.get("include_state_snapshot", "false").lower() == "true"
|
||||
continue_on_pause = request.args.get("continue_on_pause", "false").lower() == "true"
|
||||
|
||||
def _generate_stream_events():
|
||||
if include_state_snapshot:
|
||||
if include_state_snapshot or cursor is not None:
|
||||
return generator.convert_to_event_stream(
|
||||
build_workflow_event_stream(
|
||||
app_mode=app_mode,
|
||||
@@ -96,10 +106,17 @@ class WorkflowEventsApi(WebApiResource):
|
||||
tenant_id=app_model.tenant_id,
|
||||
app_id=app_model.id,
|
||||
session_maker=session_maker,
|
||||
close_on_pause=not continue_on_pause,
|
||||
cursor=cursor,
|
||||
)
|
||||
)
|
||||
terminal_events = [StreamEvent.WORKFLOW_FINISHED] if continue_on_pause else None
|
||||
return generator.convert_to_event_stream(
|
||||
msg_generator.retrieve_events(app_mode, workflow_run.id),
|
||||
msg_generator.retrieve_events(
|
||||
app_mode,
|
||||
workflow_run.id,
|
||||
terminal_events=terminal_events,
|
||||
)
|
||||
)
|
||||
|
||||
event_generator = _generate_stream_events
|
||||
|
||||
@@ -39,6 +39,7 @@ from core.app.entities.task_entities import (
|
||||
AdvancedChatPausedBlockingResponse,
|
||||
ChatbotAppBlockingResponse,
|
||||
ChatbotAppStreamResponse,
|
||||
WorkflowMaintenancePausedBlockingResponse,
|
||||
)
|
||||
from core.app.layers.pause_state_persist_layer import PauseStateLayerConfig, PauseStatePersistenceLayer
|
||||
from core.helper.trace_id_helper import extract_external_trace_id_from_args, extract_trace_session_id_from_args
|
||||
@@ -56,12 +57,14 @@ from graphon.variable_loader import DUMMY_VARIABLE_LOADER, VariableLoader
|
||||
from libs.flask_utils import preserve_flask_contexts
|
||||
from models import Account, App, Conversation, EndUser, Message, Workflow, WorkflowNodeExecutionTriggeredFrom
|
||||
from models.enums import WorkflowRunTriggeredFrom
|
||||
from models.workflow_handoff import WorkflowHandoffResumeRoute
|
||||
from services.conversation_service import ConversationService
|
||||
from services.errors.conversation import ConversationNotExistsError
|
||||
from services.workflow_draft_variable_service import (
|
||||
DraftVarLoader,
|
||||
WorkflowDraftVariableService,
|
||||
)
|
||||
from services.workflow_handoff_runtime_service import build_workflow_handoff_persistence_layer
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -284,6 +287,11 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator):
|
||||
graph_runtime_state: GraphRuntimeState,
|
||||
pause_state_config: PauseStateLayerConfig | None = None,
|
||||
response_stream_filter: ResponseStreamFilter | None = None,
|
||||
graph_engine_layers: Sequence[GraphEngineLayer] = (),
|
||||
handoff_resume_route: WorkflowHandoffResumeRoute | None = None,
|
||||
graph_config: Mapping[str, Any] | None = None,
|
||||
workflow_version: str | None = None,
|
||||
root_node_id: str | None = None,
|
||||
):
|
||||
"""
|
||||
Resume a paused advanced chat execution.
|
||||
@@ -314,6 +322,11 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator):
|
||||
pause_state_config=pause_state_config,
|
||||
graph_runtime_state=graph_runtime_state,
|
||||
response_stream_filter=response_stream_filter,
|
||||
graph_engine_layers=graph_engine_layers,
|
||||
handoff_resume_route=handoff_resume_route,
|
||||
graph_config=graph_config,
|
||||
workflow_version=workflow_version,
|
||||
root_node_id=root_node_id,
|
||||
session=session,
|
||||
)
|
||||
|
||||
@@ -366,6 +379,7 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator):
|
||||
single_iteration_run=AdvancedChatAppGenerateEntity.SingleIterationRunEntity(
|
||||
node_id=node_id, inputs=args["inputs"]
|
||||
),
|
||||
workflow_run_id=str(uuid.uuid4()),
|
||||
)
|
||||
contexts.plugin_tool_providers.set({})
|
||||
contexts.plugin_tool_providers_lock.set(threading.Lock())
|
||||
@@ -459,6 +473,7 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator):
|
||||
**_extract_trace_session_id_from_debug_args(args),
|
||||
},
|
||||
single_loop_run=AdvancedChatAppGenerateEntity.SingleLoopRunEntity(node_id=node_id, inputs=args.inputs),
|
||||
workflow_run_id=str(uuid.uuid4()),
|
||||
)
|
||||
contexts.plugin_tool_providers.set({})
|
||||
contexts.plugin_tool_providers_lock.set(threading.Lock())
|
||||
@@ -523,6 +538,10 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator):
|
||||
graph_runtime_state: GraphRuntimeState | None = None,
|
||||
graph_engine_layers: Sequence[GraphEngineLayer] = (),
|
||||
response_stream_filter: ResponseStreamFilter | None = None,
|
||||
handoff_resume_route: WorkflowHandoffResumeRoute | None = None,
|
||||
graph_config: Mapping[str, Any] | None = None,
|
||||
workflow_version: str | None = None,
|
||||
root_node_id: str | None = None,
|
||||
) -> Mapping[str, Any] | Generator[str | Mapping[str, Any], None, None]:
|
||||
"""
|
||||
Generate App response.
|
||||
@@ -576,6 +595,13 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator):
|
||||
|
||||
graph_layers: list[GraphEngineLayer] = list(graph_engine_layers)
|
||||
resolved_response_stream_filter = response_stream_filter or ResponseStreamFilter()
|
||||
handoff_layer = build_workflow_handoff_persistence_layer(
|
||||
generate_entity=application_generate_entity,
|
||||
response_stream_filter=resolved_response_stream_filter,
|
||||
resume_route=handoff_resume_route,
|
||||
)
|
||||
if handoff_layer is not None:
|
||||
graph_layers.append(handoff_layer)
|
||||
if pause_state_config is not None:
|
||||
graph_layers.append(
|
||||
PauseStatePersistenceLayer(
|
||||
@@ -604,6 +630,9 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator):
|
||||
"graph_engine_layers": tuple(graph_layers),
|
||||
"graph_runtime_state": graph_runtime_state,
|
||||
"response_stream_filter": resolved_response_stream_filter,
|
||||
"graph_config": graph_config,
|
||||
"workflow_version": workflow_version,
|
||||
"root_node_id": root_node_id,
|
||||
},
|
||||
)
|
||||
|
||||
@@ -613,7 +642,7 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator):
|
||||
# releasing the request-scoped SQLAlchemy session.
|
||||
workflow_snapshot = WorkflowSnapshot.from_workflow(workflow)
|
||||
conversation_snapshot = ConversationSnapshot.from_conversation(conversation)
|
||||
message_snapshot = MessageSnapshot.from_message(message)
|
||||
message_snapshot = MessageSnapshot.from_message(message, session=session)
|
||||
session.close()
|
||||
|
||||
try:
|
||||
@@ -659,6 +688,9 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator):
|
||||
graph_engine_layers: Sequence[GraphEngineLayer] = (),
|
||||
graph_runtime_state: GraphRuntimeState | None = None,
|
||||
response_stream_filter: ResponseStreamFilter | None = None,
|
||||
graph_config: Mapping[str, Any] | None = None,
|
||||
workflow_version: str | None = None,
|
||||
root_node_id: str | None = None,
|
||||
):
|
||||
"""
|
||||
Generate worker in a new thread.
|
||||
@@ -719,10 +751,16 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator):
|
||||
graph_engine_layers=graph_engine_layers,
|
||||
graph_runtime_state=graph_runtime_state,
|
||||
response_stream_filter=response_stream_filter,
|
||||
graph_config=graph_config,
|
||||
workflow_version=workflow_version,
|
||||
root_node_id=root_node_id,
|
||||
)
|
||||
|
||||
try:
|
||||
with active_workflow_task(application_generate_entity.task_id):
|
||||
with active_workflow_task(
|
||||
application_generate_entity.task_id,
|
||||
workflow_run_id=application_generate_entity.workflow_run_id,
|
||||
):
|
||||
runner.run()
|
||||
except GenerateTaskStoppedError:
|
||||
pass
|
||||
@@ -757,6 +795,7 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator):
|
||||
) -> (
|
||||
ChatbotAppBlockingResponse
|
||||
| AdvancedChatPausedBlockingResponse
|
||||
| WorkflowMaintenancePausedBlockingResponse
|
||||
| Generator[ChatbotAppStreamResponse, None, None]
|
||||
):
|
||||
"""
|
||||
|
||||
@@ -6,6 +6,7 @@ from typing import Any, cast
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from configs import dify_config
|
||||
from core.app.apps.advanced_chat.app_config_manager import AdvancedChatAppConfig
|
||||
from core.app.apps.base_app_queue_manager import AppQueueManager
|
||||
from core.app.apps.workflow.command_channels import (
|
||||
@@ -26,6 +27,7 @@ from core.app.entities.queue_entities import (
|
||||
)
|
||||
from core.app.features.annotation_reply.annotation_reply import AnnotationReplyFeature
|
||||
from core.app.layers.conversation_variable_persist_layer import ConversationVariablePersistenceLayer
|
||||
from core.app.layers.pause_state_persist_layer import get_workflow_handoff_active_execution_seconds
|
||||
from core.app.workflow.layers.persistence import PersistenceWorkflowInfo, WorkflowPersistenceLayer
|
||||
from core.db.session_factory import create_session, session_factory
|
||||
from core.moderation.base import ModerationError
|
||||
@@ -42,7 +44,7 @@ from core.workflow.variable_pool_initializer import add_node_inputs_to_pool, add
|
||||
from core.workflow.workflow_entry import WorkflowEntry
|
||||
from extensions.ext_redis import redis_client
|
||||
from extensions.otel import WorkflowAppRunnerHandler, trace_span
|
||||
from extensions.workflow_warm_shutdown import WORKFLOW_WARM_SHUTDOWN_ABORT_REASON, celery_warm_shutdown_started
|
||||
from extensions.workflow_warm_shutdown import celery_warm_shutdown_started
|
||||
from graphon.enums import WorkflowType
|
||||
from graphon.filters import ResponseStreamFilter
|
||||
from graphon.graph_engine.command_channels import RedisChannel
|
||||
@@ -80,6 +82,9 @@ class AdvancedChatAppRunner(WorkflowBasedAppRunner):
|
||||
graph_engine_layers: Sequence[GraphEngineLayer] = (),
|
||||
graph_runtime_state: GraphRuntimeState | None = None,
|
||||
response_stream_filter: ResponseStreamFilter | None = None,
|
||||
graph_config: Mapping[str, Any] | None = None,
|
||||
workflow_version: str | None = None,
|
||||
root_node_id: str | None = None,
|
||||
):
|
||||
super().__init__(
|
||||
queue_manager=queue_manager,
|
||||
@@ -98,6 +103,9 @@ class AdvancedChatAppRunner(WorkflowBasedAppRunner):
|
||||
self._workflow_node_execution_repository = workflow_node_execution_repository
|
||||
self._resume_graph_runtime_state = graph_runtime_state
|
||||
self._response_stream_filter = response_stream_filter
|
||||
self._execution_graph_config = graph_config
|
||||
self._workflow_version = workflow_version
|
||||
self._root_node_id = root_node_id
|
||||
|
||||
@trace_span(WorkflowAppRunnerHandler)
|
||||
def run(self):
|
||||
@@ -125,6 +133,10 @@ class AdvancedChatAppRunner(WorkflowBasedAppRunner):
|
||||
if self.application_generate_entity.single_iteration_run or self.application_generate_entity.single_loop_run:
|
||||
invoke_from = InvokeFrom.DEBUGGER
|
||||
user_from = self._resolve_user_from(invoke_from)
|
||||
graph_config = (
|
||||
self._execution_graph_config if self._execution_graph_config is not None else self._workflow.graph_dict
|
||||
)
|
||||
workflow_version = self._workflow_version if self._workflow_version is not None else self._workflow.version
|
||||
|
||||
resume_state = self._resume_graph_runtime_state
|
||||
|
||||
@@ -132,22 +144,30 @@ class AdvancedChatAppRunner(WorkflowBasedAppRunner):
|
||||
graph_runtime_state = resume_state
|
||||
variable_pool = graph_runtime_state.variable_pool
|
||||
graph = self._init_graph(
|
||||
graph_config=self._workflow.graph_dict,
|
||||
graph_config=graph_config,
|
||||
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,
|
||||
root_node_id=self._root_node_id,
|
||||
trace_session_id=self.application_generate_entity.extras.get("trace_session_id"),
|
||||
skip_validation=bool(
|
||||
self.application_generate_entity.single_iteration_run
|
||||
or self.application_generate_entity.single_loop_run
|
||||
),
|
||||
)
|
||||
elif self.application_generate_entity.single_iteration_run or self.application_generate_entity.single_loop_run:
|
||||
# Handle single iteration or single loop run
|
||||
graph, variable_pool, graph_runtime_state = self._prepare_single_node_execution(
|
||||
if self.application_generate_entity.workflow_run_id is None:
|
||||
raise ValueError("Workflow execution id is required for a single-node chatflow run")
|
||||
graph, graph_config, variable_pool, graph_runtime_state = self._prepare_single_node_execution(
|
||||
workflow=self._workflow,
|
||||
single_iteration_run=self.application_generate_entity.single_iteration_run,
|
||||
single_loop_run=self.application_generate_entity.single_loop_run,
|
||||
user_id=self.application_generate_entity.user_id,
|
||||
workflow_execution_id=self.application_generate_entity.workflow_run_id,
|
||||
trace_session_id=self.application_generate_entity.extras.get("trace_session_id"),
|
||||
)
|
||||
else:
|
||||
@@ -204,13 +224,13 @@ class AdvancedChatAppRunner(WorkflowBasedAppRunner):
|
||||
conversation_variables=conversation_variables,
|
||||
),
|
||||
)
|
||||
root_node_id = get_default_root_node_id(self._workflow.graph_dict)
|
||||
root_node_id = get_default_root_node_id(graph_config)
|
||||
add_node_inputs_to_pool(variable_pool, node_id=root_node_id, inputs=new_inputs)
|
||||
|
||||
# init graph
|
||||
graph_runtime_state = GraphRuntimeState(variable_pool=variable_pool, start_at=time.time())
|
||||
graph = self._init_graph(
|
||||
graph_config=self._workflow.graph_dict,
|
||||
graph_config=graph_config,
|
||||
graph_runtime_state=graph_runtime_state,
|
||||
workflow_id=self._workflow.id,
|
||||
tenant_id=self._workflow.tenant_id,
|
||||
@@ -222,12 +242,16 @@ class AdvancedChatAppRunner(WorkflowBasedAppRunner):
|
||||
)
|
||||
|
||||
# RUN WORKFLOW
|
||||
self._configure_execution_root(graph.root_node.id)
|
||||
# Create Redis command channel for this workflow execution
|
||||
task_id = self.application_generate_entity.task_id
|
||||
channel_key = f"workflow:{task_id}:commands"
|
||||
celery_signal_channel = CelerySignalCommandChannel(
|
||||
shutdown_state_getter=celery_warm_shutdown_started,
|
||||
abort_reason=WORKFLOW_WARM_SHUTDOWN_ABORT_REASON,
|
||||
pause_on_shutdown=(
|
||||
dify_config.WORKFLOW_HANDOFF_ENABLED and self.application_generate_entity.call_depth == 0
|
||||
),
|
||||
ignore_shutdown=(dify_config.WORKFLOW_HANDOFF_ENABLED and self.application_generate_entity.call_depth > 0),
|
||||
)
|
||||
command_channel = CombinedCommandChannel(
|
||||
(
|
||||
@@ -241,7 +265,7 @@ class AdvancedChatAppRunner(WorkflowBasedAppRunner):
|
||||
app_id=self._workflow.app_id,
|
||||
workflow_id=self._workflow.id,
|
||||
graph=graph,
|
||||
graph_config=self._workflow.graph_dict,
|
||||
graph_config=graph_config,
|
||||
user_id=self.application_generate_entity.user_id,
|
||||
user_from=user_from,
|
||||
invoke_from=invoke_from,
|
||||
@@ -250,6 +274,9 @@ class AdvancedChatAppRunner(WorkflowBasedAppRunner):
|
||||
graph_runtime_state=graph_runtime_state,
|
||||
command_channel=command_channel,
|
||||
response_stream_filter=self._response_stream_filter,
|
||||
prior_active_execution_seconds=get_workflow_handoff_active_execution_seconds(
|
||||
self.application_generate_entity
|
||||
),
|
||||
)
|
||||
|
||||
self._queue_manager.graph_runtime_state = graph_runtime_state
|
||||
@@ -259,8 +286,8 @@ class AdvancedChatAppRunner(WorkflowBasedAppRunner):
|
||||
workflow_info=PersistenceWorkflowInfo(
|
||||
workflow_id=self._workflow.id,
|
||||
workflow_type=WorkflowType(self._workflow.type),
|
||||
version=self._workflow.version,
|
||||
graph_data=self._workflow.graph_dict,
|
||||
version=workflow_version,
|
||||
graph_data=graph_config,
|
||||
),
|
||||
workflow_execution_repository=self._workflow_execution_repository,
|
||||
workflow_node_execution_repository=self._workflow_node_execution_repository,
|
||||
@@ -289,7 +316,7 @@ class AdvancedChatAppRunner(WorkflowBasedAppRunner):
|
||||
|
||||
generator = workflow_entry.run()
|
||||
for event in generator:
|
||||
self._handle_event(workflow_entry, event)
|
||||
self._handle_event_with_handoff_contracts(workflow_entry, event)
|
||||
|
||||
def handle_input_moderation(
|
||||
self,
|
||||
|
||||
@@ -2,6 +2,7 @@ from collections.abc import Generator
|
||||
from typing import Any, cast, override
|
||||
|
||||
from core.app.apps.base_app_generate_response_converter import AppGenerateResponseConverter
|
||||
from core.app.apps.streaming_utils import close_stream
|
||||
from core.app.entities.task_entities import (
|
||||
AdvancedChatPausedBlockingResponse,
|
||||
AppStreamResponse,
|
||||
@@ -13,22 +14,31 @@ from core.app.entities.task_entities import (
|
||||
NodeStartStreamResponse,
|
||||
PingStreamResponse,
|
||||
StreamEvent,
|
||||
WorkflowMaintenancePausedBlockingResponse,
|
||||
)
|
||||
|
||||
|
||||
class AdvancedChatAppGenerateResponseConverter(
|
||||
AppGenerateResponseConverter[ChatbotAppBlockingResponse | AdvancedChatPausedBlockingResponse]
|
||||
AppGenerateResponseConverter[
|
||||
ChatbotAppBlockingResponse | AdvancedChatPausedBlockingResponse | WorkflowMaintenancePausedBlockingResponse
|
||||
]
|
||||
):
|
||||
@classmethod
|
||||
@override
|
||||
def convert_blocking_full_response(
|
||||
cls, blocking_response: ChatbotAppBlockingResponse | AdvancedChatPausedBlockingResponse
|
||||
cls,
|
||||
blocking_response: (
|
||||
ChatbotAppBlockingResponse | AdvancedChatPausedBlockingResponse | WorkflowMaintenancePausedBlockingResponse
|
||||
),
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
Convert blocking full response.
|
||||
:param blocking_response: blocking response
|
||||
:return:
|
||||
"""
|
||||
if isinstance(blocking_response, WorkflowMaintenancePausedBlockingResponse):
|
||||
return blocking_response.model_dump(mode="json")
|
||||
|
||||
if isinstance(blocking_response, AdvancedChatPausedBlockingResponse):
|
||||
paused_data = blocking_response.data.model_dump(mode="json")
|
||||
return {
|
||||
@@ -62,7 +72,10 @@ class AdvancedChatAppGenerateResponseConverter(
|
||||
@classmethod
|
||||
@override
|
||||
def convert_blocking_simple_response(
|
||||
cls, blocking_response: ChatbotAppBlockingResponse | AdvancedChatPausedBlockingResponse
|
||||
cls,
|
||||
blocking_response: (
|
||||
ChatbotAppBlockingResponse | AdvancedChatPausedBlockingResponse | WorkflowMaintenancePausedBlockingResponse
|
||||
),
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
Convert blocking simple response.
|
||||
@@ -87,27 +100,30 @@ class AdvancedChatAppGenerateResponseConverter(
|
||||
:param stream_response: stream response
|
||||
:return:
|
||||
"""
|
||||
for chunk in stream_response:
|
||||
chunk = cast(ChatbotAppStreamResponse, chunk)
|
||||
sub_stream_response = chunk.stream_response
|
||||
try:
|
||||
for chunk in stream_response:
|
||||
chunk = cast(ChatbotAppStreamResponse, chunk)
|
||||
sub_stream_response = chunk.stream_response
|
||||
|
||||
if isinstance(sub_stream_response, PingStreamResponse):
|
||||
yield "ping"
|
||||
continue
|
||||
if isinstance(sub_stream_response, PingStreamResponse):
|
||||
yield "ping"
|
||||
continue
|
||||
|
||||
response_chunk: dict[str, Any] = {
|
||||
"event": sub_stream_response.event.value,
|
||||
"conversation_id": chunk.conversation_id,
|
||||
"message_id": chunk.message_id,
|
||||
"created_at": chunk.created_at,
|
||||
}
|
||||
response_chunk: dict[str, Any] = {
|
||||
"event": sub_stream_response.event.value,
|
||||
"conversation_id": chunk.conversation_id,
|
||||
"message_id": chunk.message_id,
|
||||
"created_at": chunk.created_at,
|
||||
}
|
||||
|
||||
if isinstance(sub_stream_response, ErrorStreamResponse):
|
||||
data = cls._error_to_stream_response(sub_stream_response.err)
|
||||
response_chunk.update(data)
|
||||
else:
|
||||
response_chunk.update(sub_stream_response.model_dump(mode="json"))
|
||||
yield response_chunk
|
||||
if isinstance(sub_stream_response, ErrorStreamResponse):
|
||||
data = cls._error_to_stream_response(sub_stream_response.err)
|
||||
response_chunk.update(data)
|
||||
else:
|
||||
response_chunk.update(sub_stream_response.model_dump(mode="json"))
|
||||
yield response_chunk
|
||||
finally:
|
||||
close_stream(stream_response)
|
||||
|
||||
@classmethod
|
||||
@override
|
||||
@@ -119,33 +135,36 @@ class AdvancedChatAppGenerateResponseConverter(
|
||||
:param stream_response: stream response
|
||||
:return:
|
||||
"""
|
||||
for chunk in stream_response:
|
||||
chunk = cast(ChatbotAppStreamResponse, chunk)
|
||||
sub_stream_response = chunk.stream_response
|
||||
try:
|
||||
for chunk in stream_response:
|
||||
chunk = cast(ChatbotAppStreamResponse, chunk)
|
||||
sub_stream_response = chunk.stream_response
|
||||
|
||||
if isinstance(sub_stream_response, PingStreamResponse):
|
||||
yield "ping"
|
||||
continue
|
||||
if isinstance(sub_stream_response, PingStreamResponse):
|
||||
yield "ping"
|
||||
continue
|
||||
|
||||
response_chunk: dict[str, Any] = {
|
||||
"event": sub_stream_response.event.value,
|
||||
"conversation_id": chunk.conversation_id,
|
||||
"message_id": chunk.message_id,
|
||||
"created_at": chunk.created_at,
|
||||
}
|
||||
response_chunk: dict[str, Any] = {
|
||||
"event": sub_stream_response.event.value,
|
||||
"conversation_id": chunk.conversation_id,
|
||||
"message_id": chunk.message_id,
|
||||
"created_at": chunk.created_at,
|
||||
}
|
||||
|
||||
match sub_stream_response:
|
||||
case MessageEndStreamResponse():
|
||||
sub_stream_response_dict = sub_stream_response.model_dump(mode="json")
|
||||
metadata = sub_stream_response_dict.get("metadata", {})
|
||||
sub_stream_response_dict["metadata"] = cls._get_simple_metadata(metadata)
|
||||
response_chunk.update(sub_stream_response_dict)
|
||||
case ErrorStreamResponse():
|
||||
data = cls._error_to_stream_response(sub_stream_response.err)
|
||||
response_chunk.update(data)
|
||||
case NodeStartStreamResponse() | NodeFinishStreamResponse():
|
||||
response_chunk.update(sub_stream_response.to_ignore_detail_dict())
|
||||
case _:
|
||||
response_chunk.update(sub_stream_response.model_dump(mode="json"))
|
||||
match sub_stream_response:
|
||||
case MessageEndStreamResponse():
|
||||
sub_stream_response_dict = sub_stream_response.model_dump(mode="json")
|
||||
metadata = sub_stream_response_dict.get("metadata", {})
|
||||
sub_stream_response_dict["metadata"] = cls._get_simple_metadata(metadata)
|
||||
response_chunk.update(sub_stream_response_dict)
|
||||
case ErrorStreamResponse():
|
||||
data = cls._error_to_stream_response(sub_stream_response.err)
|
||||
response_chunk.update(data)
|
||||
case NodeStartStreamResponse() | NodeFinishStreamResponse():
|
||||
response_chunk.update(sub_stream_response.to_ignore_detail_dict())
|
||||
case _:
|
||||
response_chunk.update(sub_stream_response.model_dump(mode="json"))
|
||||
|
||||
yield response_chunk
|
||||
yield response_chunk
|
||||
finally:
|
||||
close_stream(stream_response)
|
||||
|
||||
@@ -1,9 +1,9 @@
|
||||
import json
|
||||
import logging
|
||||
import time
|
||||
from collections.abc import Callable, Generator, Mapping
|
||||
from collections.abc import Callable, Generator, Mapping, Sequence
|
||||
from contextlib import contextmanager
|
||||
from dataclasses import dataclass
|
||||
from dataclasses import dataclass, field
|
||||
from datetime import datetime
|
||||
from threading import Thread
|
||||
from typing import Any, Union
|
||||
@@ -16,6 +16,7 @@ from core.app.apps.base_app_queue_manager import AppQueueManager, PublishFrom
|
||||
from core.app.apps.common.graph_runtime_state_support import GraphRuntimeStateSupport
|
||||
from core.app.apps.common.workflow_response_converter import WorkflowResponseConverter
|
||||
from core.app.apps.draft_variable_saver import DraftVariableSaverFactory
|
||||
from core.app.apps.streaming_utils import close_stream
|
||||
from core.app.entities.app_invoke_entities import (
|
||||
AdvancedChatAppGenerateEntity,
|
||||
InvokeFrom,
|
||||
@@ -46,6 +47,7 @@ from core.app.entities.queue_entities import (
|
||||
QueueStopEvent,
|
||||
QueueTextChunkEvent,
|
||||
QueueWorkflowFailedEvent,
|
||||
QueueWorkflowMaintenancePausedEvent,
|
||||
QueueWorkflowPartialSuccessEvent,
|
||||
QueueWorkflowPausedEvent,
|
||||
QueueWorkflowStartedEvent,
|
||||
@@ -62,14 +64,19 @@ from core.app.entities.task_entities import (
|
||||
MessageAudioEndStreamResponse,
|
||||
MessageAudioStreamResponse,
|
||||
MessageEndStreamResponse,
|
||||
MessageReplaceStreamResponse,
|
||||
PingStreamResponse,
|
||||
ReasoningChunkStreamResponse,
|
||||
StreamResponse,
|
||||
TaskStateMetadata,
|
||||
WorkflowMaintenancePausedBlockingResponse,
|
||||
WorkflowMaintenancePausedStreamResponse,
|
||||
WorkflowPauseStreamResponse,
|
||||
WorkflowTaskState,
|
||||
)
|
||||
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_file_utils import prepare_file_dict
|
||||
from core.base.tts import AppGeneratorTTSPublisher, AudioTrunk
|
||||
from core.db.session_factory import session_factory
|
||||
from core.ops.ops_trace_manager import TraceQueueManager
|
||||
@@ -77,17 +84,21 @@ from core.repositories.human_input_repository import HumanInputFormRepositoryImp
|
||||
from core.workflow.file_reference import resolve_file_record_id
|
||||
from core.workflow.nodes.human_input.pause_reason import HumanInputRequired
|
||||
from core.workflow.system_variables import build_system_variables
|
||||
from graphon.entities import WorkflowStartReason
|
||||
from graphon.enums import WorkflowExecutionStatus
|
||||
from graphon.file import FileTransferMethod
|
||||
from graphon.model_runtime.entities.llm_entities import LLMUsage
|
||||
from graphon.model_runtime.utils.encoders import jsonable_encoder
|
||||
from graphon.nodes import BuiltinNodeTypes
|
||||
from graphon.runtime import GraphRuntimeState
|
||||
from libs.datetime_utils import naive_utc_now
|
||||
from models import Account, Conversation, EndUser, Message, MessageFile
|
||||
from models import Account, Conversation, EndUser, Message, MessageFile, UploadFile
|
||||
from models.enums import CreatorUserRole, MessageFileBelongsTo, MessageStatus
|
||||
from models.execution_extra_content import HumanInputContent
|
||||
from models.model import AppMode
|
||||
from models.workflow import Workflow
|
||||
from services.workflow_handoff_activation_service import activate_workflow_handoff_by_task_id
|
||||
from services.workflow_run_timing_service import get_workflow_run_public_timing
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -127,15 +138,33 @@ class MessageSnapshot:
|
||||
created_at: datetime
|
||||
status: MessageStatus
|
||||
answer: str
|
||||
message_metadata: Mapping[str, Any] = field(default_factory=dict)
|
||||
recorded_files: Sequence[Mapping[str, Any]] = field(default_factory=tuple)
|
||||
provider_response_latency: float = 0.0
|
||||
|
||||
@classmethod
|
||||
def from_message(cls, message: Message) -> "MessageSnapshot":
|
||||
def from_message(cls, message: Message, *, session: Session) -> "MessageSnapshot":
|
||||
message_files = list(session.scalars(select(MessageFile).where(MessageFile.message_id == message.id)))
|
||||
upload_file_ids = list(
|
||||
dict.fromkeys(
|
||||
item.upload_file_id
|
||||
for item in message_files
|
||||
if item.transfer_method == FileTransferMethod.LOCAL_FILE and item.upload_file_id
|
||||
)
|
||||
)
|
||||
upload_files_map: dict[str, UploadFile] = {}
|
||||
if upload_file_ids:
|
||||
upload_files = session.scalars(select(UploadFile).where(UploadFile.id.in_(upload_file_ids))).all()
|
||||
upload_files_map = {item.id: item for item in upload_files}
|
||||
return cls(
|
||||
id=message.id,
|
||||
query=message.query,
|
||||
created_at=message.created_at,
|
||||
status=message.status,
|
||||
answer=message.answer,
|
||||
message_metadata=message.message_metadata_dict,
|
||||
recorded_files=tuple(prepare_file_dict(item, upload_files_map) for item in message_files),
|
||||
provider_response_latency=message.provider_response_latency or 0.0,
|
||||
)
|
||||
|
||||
|
||||
@@ -191,6 +220,8 @@ class AdvancedChatAppGenerateTaskPipeline(GraphRuntimeStateSupport):
|
||||
)
|
||||
|
||||
self._task_state = WorkflowTaskState()
|
||||
self._recorded_files: list[Mapping[str, Any]] = list(message.recorded_files)
|
||||
self._prior_provider_response_latency = message.provider_response_latency
|
||||
self._seed_task_state_from_message(message)
|
||||
self._message_cycle_manager = MessageCycleManager(
|
||||
application_generate_entity=application_generate_entity, task_state=self._task_state
|
||||
@@ -205,7 +236,6 @@ class AdvancedChatAppGenerateTaskPipeline(GraphRuntimeStateSupport):
|
||||
self._message_id = message.id
|
||||
self._message_created_at = int(message.created_at.timestamp())
|
||||
self._conversation_name_generate_thread: Thread | None = None
|
||||
self._recorded_files: list[Mapping[str, Any]] = []
|
||||
self._workflow_run_id: str = ""
|
||||
self._draft_var_saver_factory = draft_var_saver_factory
|
||||
self._graph_runtime_state: GraphRuntimeState | None = None
|
||||
@@ -213,14 +243,21 @@ class AdvancedChatAppGenerateTaskPipeline(GraphRuntimeStateSupport):
|
||||
self._seed_graph_runtime_state_from_queue_manager()
|
||||
|
||||
def _seed_task_state_from_message(self, message: MessageSnapshot) -> None:
|
||||
if message.status == MessageStatus.PAUSED and message.answer:
|
||||
# Resumed executions reuse the original Message. Human-input pauses mark
|
||||
# it as PAUSED, while maintenance handoffs deliberately preserve its
|
||||
# public status, so the persisted answer itself is the common resume
|
||||
# checkpoint for both paths. Newly created messages always start empty.
|
||||
if message.answer:
|
||||
self._task_state.answer = message.answer
|
||||
if message.message_metadata:
|
||||
self._task_state.metadata = TaskStateMetadata.model_validate(message.message_metadata)
|
||||
|
||||
def process(
|
||||
self,
|
||||
) -> Union[
|
||||
ChatbotAppBlockingResponse,
|
||||
AdvancedChatPausedBlockingResponse,
|
||||
WorkflowMaintenancePausedBlockingResponse,
|
||||
Generator[ChatbotAppStreamResponse, None, None],
|
||||
]:
|
||||
"""
|
||||
@@ -240,57 +277,70 @@ class AdvancedChatAppGenerateTaskPipeline(GraphRuntimeStateSupport):
|
||||
|
||||
def _to_blocking_response(
|
||||
self, generator: Generator[StreamResponse, None, None]
|
||||
) -> Union[ChatbotAppBlockingResponse, AdvancedChatPausedBlockingResponse]:
|
||||
) -> Union[
|
||||
ChatbotAppBlockingResponse,
|
||||
AdvancedChatPausedBlockingResponse,
|
||||
WorkflowMaintenancePausedBlockingResponse,
|
||||
]:
|
||||
"""
|
||||
Process blocking response.
|
||||
:return:
|
||||
"""
|
||||
human_input_responses: list[HumanInputRequiredResponse] = []
|
||||
for stream_response in generator:
|
||||
match stream_response:
|
||||
case ErrorStreamResponse():
|
||||
raise stream_response.err
|
||||
case HumanInputRequiredResponse():
|
||||
human_input_responses.append(stream_response)
|
||||
case WorkflowPauseStreamResponse():
|
||||
return AdvancedChatPausedBlockingResponse(
|
||||
task_id=stream_response.task_id,
|
||||
data=AdvancedChatPausedBlockingResponse.Data(
|
||||
id=self._message_id,
|
||||
mode=self._conversation_mode,
|
||||
conversation_id=self._conversation_id,
|
||||
message_id=self._message_id,
|
||||
workflow_run_id=stream_response.data.workflow_run_id,
|
||||
answer=self._task_state.answer,
|
||||
metadata=self._message_end_to_stream_response().metadata,
|
||||
created_at=self._message_created_at,
|
||||
paused_nodes=stream_response.data.paused_nodes,
|
||||
reasons=stream_response.data.reasons,
|
||||
status=stream_response.data.status,
|
||||
elapsed_time=stream_response.data.elapsed_time,
|
||||
total_tokens=stream_response.data.total_tokens,
|
||||
total_steps=stream_response.data.total_steps,
|
||||
),
|
||||
)
|
||||
case MessageEndStreamResponse():
|
||||
extras = {}
|
||||
if stream_response.metadata:
|
||||
extras["metadata"] = stream_response.metadata
|
||||
try:
|
||||
for stream_response in generator:
|
||||
match stream_response:
|
||||
case ErrorStreamResponse():
|
||||
raise stream_response.err
|
||||
case HumanInputRequiredResponse():
|
||||
human_input_responses.append(stream_response)
|
||||
case WorkflowPauseStreamResponse():
|
||||
return AdvancedChatPausedBlockingResponse(
|
||||
task_id=stream_response.task_id,
|
||||
data=AdvancedChatPausedBlockingResponse.Data(
|
||||
id=self._message_id,
|
||||
mode=self._conversation_mode,
|
||||
conversation_id=self._conversation_id,
|
||||
message_id=self._message_id,
|
||||
workflow_run_id=stream_response.data.workflow_run_id,
|
||||
answer=self._task_state.answer,
|
||||
metadata=self._message_end_to_stream_response().metadata,
|
||||
created_at=self._message_created_at,
|
||||
paused_nodes=stream_response.data.paused_nodes,
|
||||
reasons=stream_response.data.reasons,
|
||||
status=stream_response.data.status,
|
||||
elapsed_time=stream_response.data.elapsed_time,
|
||||
total_tokens=stream_response.data.total_tokens,
|
||||
total_steps=stream_response.data.total_steps,
|
||||
handoff_duration=stream_response.data.handoff_duration,
|
||||
),
|
||||
)
|
||||
case WorkflowMaintenancePausedStreamResponse():
|
||||
return WorkflowMaintenancePausedBlockingResponse(
|
||||
task_id=stream_response.task_id,
|
||||
workflow_run_id=stream_response.workflow_run_id,
|
||||
)
|
||||
case MessageEndStreamResponse():
|
||||
extras = {}
|
||||
if stream_response.metadata:
|
||||
extras["metadata"] = stream_response.metadata
|
||||
|
||||
return ChatbotAppBlockingResponse(
|
||||
task_id=stream_response.task_id,
|
||||
data=ChatbotAppBlockingResponse.Data(
|
||||
id=self._message_id,
|
||||
mode=self._conversation_mode,
|
||||
conversation_id=self._conversation_id,
|
||||
message_id=self._message_id,
|
||||
answer=self._task_state.answer,
|
||||
created_at=self._message_created_at,
|
||||
**extras,
|
||||
),
|
||||
)
|
||||
case _:
|
||||
continue
|
||||
return ChatbotAppBlockingResponse(
|
||||
task_id=stream_response.task_id,
|
||||
data=ChatbotAppBlockingResponse.Data(
|
||||
id=self._message_id,
|
||||
mode=self._conversation_mode,
|
||||
conversation_id=self._conversation_id,
|
||||
message_id=self._message_id,
|
||||
answer=self._task_state.answer,
|
||||
created_at=self._message_created_at,
|
||||
**extras,
|
||||
),
|
||||
)
|
||||
case _:
|
||||
continue
|
||||
finally:
|
||||
close_stream(generator)
|
||||
|
||||
if human_input_responses:
|
||||
return self._build_paused_blocking_response_from_human_input(human_input_responses)
|
||||
@@ -334,13 +384,16 @@ class AdvancedChatAppGenerateTaskPipeline(GraphRuntimeStateSupport):
|
||||
To stream response.
|
||||
:return:
|
||||
"""
|
||||
for stream_response in generator:
|
||||
yield ChatbotAppStreamResponse(
|
||||
conversation_id=self._conversation_id,
|
||||
message_id=self._message_id,
|
||||
created_at=self._message_created_at,
|
||||
stream_response=stream_response,
|
||||
)
|
||||
try:
|
||||
for stream_response in generator:
|
||||
yield ChatbotAppStreamResponse(
|
||||
conversation_id=self._conversation_id,
|
||||
message_id=self._message_id,
|
||||
created_at=self._message_created_at,
|
||||
stream_response=stream_response,
|
||||
)
|
||||
finally:
|
||||
close_stream(generator)
|
||||
|
||||
def _listen_audio_msg(self, publisher: AppGeneratorTTSPublisher | None, task_id: str):
|
||||
if not publisher:
|
||||
@@ -367,14 +420,18 @@ class AdvancedChatAppGenerateTaskPipeline(GraphRuntimeStateSupport):
|
||||
tenant_id, features_dict["text_to_speech"].get("voice"), features_dict["text_to_speech"].get("language")
|
||||
)
|
||||
|
||||
for response in self._process_stream_response(tts_publisher=tts_publisher, trace_manager=trace_manager):
|
||||
while True:
|
||||
audio_response = self._listen_audio_msg(publisher=tts_publisher, task_id=task_id)
|
||||
if audio_response:
|
||||
yield audio_response
|
||||
else:
|
||||
break
|
||||
yield response
|
||||
response_stream = self._process_stream_response(tts_publisher=tts_publisher, trace_manager=trace_manager)
|
||||
try:
|
||||
for response in response_stream:
|
||||
while True:
|
||||
audio_response = self._listen_audio_msg(publisher=tts_publisher, task_id=task_id)
|
||||
if audio_response:
|
||||
yield audio_response
|
||||
else:
|
||||
break
|
||||
yield response
|
||||
finally:
|
||||
close_stream(response_stream)
|
||||
|
||||
start_listener_time = time.time()
|
||||
while (time.time() - start_listener_time) < TTS_AUTO_PLAY_TIMEOUT:
|
||||
@@ -435,11 +492,24 @@ class AdvancedChatAppGenerateTaskPipeline(GraphRuntimeStateSupport):
|
||||
with self._database_session() as session:
|
||||
session.execute(update(Message).where(Message.id == self._message_id).values(workflow_run_id=run_id))
|
||||
|
||||
logical_timing = None
|
||||
if event.reason == WorkflowStartReason.RESUMPTION:
|
||||
with self._database_session() as session:
|
||||
logical_timing = get_workflow_run_public_timing(
|
||||
session=session,
|
||||
workflow_run_id=run_id,
|
||||
tenant_id=self._workflow_tenant_id,
|
||||
app_id=self._application_generate_entity.app_config.app_id,
|
||||
workflow_id=self._workflow_id,
|
||||
)
|
||||
|
||||
workflow_start_resp = self._workflow_response_converter.workflow_start_to_stream_response(
|
||||
task_id=self._application_generate_entity.task_id,
|
||||
workflow_run_id=run_id,
|
||||
workflow_id=self._workflow_id,
|
||||
reason=event.reason,
|
||||
logical_started_at=logical_timing.started_at if logical_timing is not None else None,
|
||||
handoff_duration=logical_timing.handoff_duration if logical_timing is not None else 0.0,
|
||||
)
|
||||
|
||||
yield workflow_start_resp
|
||||
@@ -710,7 +780,6 @@ class AdvancedChatAppGenerateTaskPipeline(GraphRuntimeStateSupport):
|
||||
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()
|
||||
@@ -723,7 +792,32 @@ class AdvancedChatAppGenerateTaskPipeline(GraphRuntimeStateSupport):
|
||||
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)
|
||||
# Publish a user-visible pause only after both the graph checkpoint and
|
||||
# accumulated chat state are committed. A reconnect can therefore
|
||||
# always reconstruct the exact state advertised by this event.
|
||||
yield from responses
|
||||
|
||||
def _handle_workflow_maintenance_paused_event(
|
||||
self,
|
||||
event: QueueWorkflowMaintenancePausedEvent,
|
||||
**kwargs,
|
||||
) -> Generator[StreamResponse, None, None]:
|
||||
"""End this worker's segment without pausing the public chat message."""
|
||||
_ = event, kwargs
|
||||
self._ensure_workflow_initialized()
|
||||
validated_state = self._ensure_graph_runtime_initialized()
|
||||
with self._database_session() as session:
|
||||
self._save_message(
|
||||
session=session,
|
||||
graph_runtime_state=validated_state,
|
||||
preserve_status=True,
|
||||
)
|
||||
activate_workflow_handoff_by_task_id(self._application_generate_entity.task_id)
|
||||
self._base_task_pipeline.queue_manager.mark_execution_terminal()
|
||||
yield WorkflowMaintenancePausedStreamResponse(
|
||||
task_id=self._application_generate_entity.task_id,
|
||||
workflow_run_id=self._workflow_run_id,
|
||||
)
|
||||
|
||||
def _handle_workflow_failed_event(
|
||||
self,
|
||||
@@ -747,11 +841,17 @@ class AdvancedChatAppGenerateTaskPipeline(GraphRuntimeStateSupport):
|
||||
)
|
||||
|
||||
with self._database_session() as session:
|
||||
self._save_message(session=session, graph_runtime_state=validated_state)
|
||||
err_event = QueueErrorEvent(error=ValueError(f"Run failed: {event.error}"))
|
||||
err = self._base_task_pipeline.handle_error(event=err_event, session=session, message_id=self._message_id)
|
||||
self._base_task_pipeline.handle_error(event=err_event, session=session, message_id=self._message_id)
|
||||
|
||||
yield MessageReplaceStreamResponse(
|
||||
task_id=self._application_generate_entity.task_id,
|
||||
answer=self._task_state.answer,
|
||||
reason="workflow_resumption_terminal",
|
||||
)
|
||||
yield self._message_end_to_stream_response()
|
||||
yield workflow_finish_resp
|
||||
yield self._base_task_pipeline.error_to_stream_response(err)
|
||||
|
||||
def _handle_stop_event(
|
||||
self,
|
||||
@@ -780,6 +880,7 @@ class AdvancedChatAppGenerateTaskPipeline(GraphRuntimeStateSupport):
|
||||
# Save message
|
||||
self._save_message(session=session, graph_runtime_state=resolved_state)
|
||||
|
||||
yield self._message_end_to_stream_response()
|
||||
yield workflow_finish_resp
|
||||
elif event.stopped_by in (
|
||||
QueueStopEvent.StopBy.INPUT_MODERATION,
|
||||
@@ -789,8 +890,7 @@ class AdvancedChatAppGenerateTaskPipeline(GraphRuntimeStateSupport):
|
||||
with self._database_session() as session:
|
||||
# Save message
|
||||
self._save_message(session=session)
|
||||
|
||||
yield self._message_end_to_stream_response()
|
||||
yield self._message_end_to_stream_response()
|
||||
|
||||
def _handle_advanced_chat_message_end_event(
|
||||
self,
|
||||
@@ -918,6 +1018,7 @@ class AdvancedChatAppGenerateTaskPipeline(GraphRuntimeStateSupport):
|
||||
QueueWorkflowSucceededEvent: self._handle_workflow_succeeded_event,
|
||||
QueueWorkflowPartialSuccessEvent: self._handle_workflow_partial_success_event,
|
||||
QueueWorkflowPausedEvent: self._handle_workflow_paused_event,
|
||||
QueueWorkflowMaintenancePausedEvent: self._handle_workflow_maintenance_paused_event,
|
||||
QueueWorkflowFailedEvent: self._handle_workflow_failed_event,
|
||||
# Node events
|
||||
QueueNodeRetryEvent: self._handle_node_retry_event,
|
||||
@@ -993,48 +1094,55 @@ class AdvancedChatAppGenerateTaskPipeline(GraphRuntimeStateSupport):
|
||||
Process stream response using elegant Fluent Python patterns.
|
||||
Maintains exact same functionality as original 57-if-statement version.
|
||||
"""
|
||||
for queue_message in self._base_task_pipeline.queue_manager.listen():
|
||||
event = queue_message.event
|
||||
queue_stream = self._base_task_pipeline.queue_manager.listen()
|
||||
try:
|
||||
for queue_message in queue_stream:
|
||||
event = queue_message.event
|
||||
|
||||
match event:
|
||||
case QueueWorkflowStartedEvent():
|
||||
self._resolve_graph_runtime_state()
|
||||
yield from self._handle_workflow_started_event(event)
|
||||
match event:
|
||||
case QueueWorkflowStartedEvent():
|
||||
self._resolve_graph_runtime_state()
|
||||
yield from self._handle_workflow_started_event(event)
|
||||
|
||||
case QueueErrorEvent():
|
||||
yield from self._handle_error_event(event)
|
||||
break
|
||||
case QueueErrorEvent():
|
||||
yield from self._handle_error_event(event)
|
||||
break
|
||||
|
||||
case QueueWorkflowFailedEvent():
|
||||
yield from self._handle_workflow_failed_event(event, trace_manager=trace_manager)
|
||||
break
|
||||
case QueueWorkflowPausedEvent():
|
||||
yield from self._handle_workflow_paused_event(event)
|
||||
break
|
||||
case QueueWorkflowFailedEvent():
|
||||
yield from self._handle_workflow_failed_event(event, trace_manager=trace_manager)
|
||||
break
|
||||
case QueueWorkflowPausedEvent():
|
||||
yield from self._handle_workflow_paused_event(event)
|
||||
break
|
||||
case QueueWorkflowMaintenancePausedEvent():
|
||||
yield from self._handle_workflow_maintenance_paused_event(event)
|
||||
break
|
||||
|
||||
case QueueWorkflowSucceededEvent():
|
||||
yield from self._handle_workflow_succeeded_event(event, trace_manager=trace_manager)
|
||||
break
|
||||
case QueueWorkflowSucceededEvent():
|
||||
yield from self._handle_workflow_succeeded_event(event, trace_manager=trace_manager)
|
||||
break
|
||||
|
||||
case QueueWorkflowPartialSuccessEvent():
|
||||
yield from self._handle_workflow_partial_success_event(event, trace_manager=trace_manager)
|
||||
break
|
||||
case QueueWorkflowPartialSuccessEvent():
|
||||
yield from self._handle_workflow_partial_success_event(event, trace_manager=trace_manager)
|
||||
break
|
||||
|
||||
case QueueStopEvent():
|
||||
yield from self._handle_stop_event(event, graph_runtime_state=None, trace_manager=trace_manager)
|
||||
break
|
||||
case QueueStopEvent():
|
||||
yield from self._handle_stop_event(event, graph_runtime_state=None, trace_manager=trace_manager)
|
||||
break
|
||||
|
||||
# Handle all other events through elegant dispatch
|
||||
case _:
|
||||
if responses := list(
|
||||
self._dispatch_event(
|
||||
event,
|
||||
tts_publisher=tts_publisher,
|
||||
trace_manager=trace_manager,
|
||||
queue_message=queue_message,
|
||||
)
|
||||
):
|
||||
yield from responses
|
||||
# Handle all other events through elegant dispatch
|
||||
case _:
|
||||
if responses := list(
|
||||
self._dispatch_event(
|
||||
event,
|
||||
tts_publisher=tts_publisher,
|
||||
trace_manager=trace_manager,
|
||||
queue_message=queue_message,
|
||||
)
|
||||
):
|
||||
yield from responses
|
||||
finally:
|
||||
close_stream(queue_stream)
|
||||
|
||||
if tts_publisher:
|
||||
tts_publisher.publish(None)
|
||||
@@ -1042,19 +1150,29 @@ class AdvancedChatAppGenerateTaskPipeline(GraphRuntimeStateSupport):
|
||||
if self._conversation_name_generate_thread:
|
||||
logger.debug("Conversation name generation running as daemon thread")
|
||||
|
||||
def _save_message(self, *, session: Session, graph_runtime_state: GraphRuntimeState | None = None):
|
||||
def _save_message(
|
||||
self,
|
||||
*,
|
||||
session: Session,
|
||||
graph_runtime_state: GraphRuntimeState | None = None,
|
||||
preserve_status: bool = False,
|
||||
):
|
||||
message = self._get_message(session=session)
|
||||
if message is None:
|
||||
return
|
||||
|
||||
if message.status == MessageStatus.PAUSED:
|
||||
if not preserve_status and message.status == MessageStatus.PAUSED:
|
||||
message.status = MessageStatus.NORMAL
|
||||
|
||||
answer_text = self._task_state.answer
|
||||
|
||||
message.answer = answer_text
|
||||
message.updated_at = naive_utc_now()
|
||||
message.provider_response_latency = time.perf_counter() - self._base_task_pipeline.start_at
|
||||
segment_latency = max(time.perf_counter() - self._base_task_pipeline.start_at, 0.0)
|
||||
message.provider_response_latency = max(
|
||||
message.provider_response_latency or 0.0,
|
||||
self._prior_provider_response_latency + segment_latency,
|
||||
)
|
||||
|
||||
# Set usage first before dumping metadata
|
||||
if graph_runtime_state and graph_runtime_state.llm_usage:
|
||||
@@ -1069,8 +1187,10 @@ class AdvancedChatAppGenerateTaskPipeline(GraphRuntimeStateSupport):
|
||||
message.currency = usage.currency
|
||||
self._task_state.metadata.usage = usage
|
||||
else:
|
||||
usage = LLMUsage.empty_usage()
|
||||
self._task_state.metadata.usage = usage
|
||||
# A resumed segment may not invoke an LLM. Keep the metadata and
|
||||
# accounting persisted by the previous segment instead of
|
||||
# replacing it with an empty usage object.
|
||||
usage = self._task_state.metadata.usage or LLMUsage.empty_usage()
|
||||
|
||||
# Add streaming metrics to usage if available
|
||||
if self._task_state.is_streaming_response and self._task_state.first_token_time:
|
||||
@@ -1082,23 +1202,48 @@ class AdvancedChatAppGenerateTaskPipeline(GraphRuntimeStateSupport):
|
||||
|
||||
metadata = self._task_state.metadata.model_dump()
|
||||
message.message_metadata = json.dumps(jsonable_encoder(metadata))
|
||||
existing_message_files = list(session.scalars(select(MessageFile).where(MessageFile.message_id == message.id)))
|
||||
existing_related_ids = {item.id for item in existing_message_files}
|
||||
existing_file_keys = {
|
||||
(
|
||||
str(item.type),
|
||||
str(item.transfer_method),
|
||||
item.url or "",
|
||||
item.upload_file_id or "",
|
||||
)
|
||||
for item in existing_message_files
|
||||
}
|
||||
message_files: list[MessageFile] = []
|
||||
for file in self._recorded_files:
|
||||
reference = file.get("reference") or file.get("related_id")
|
||||
if isinstance(reference, str) and reference in existing_related_ids:
|
||||
continue
|
||||
upload_file_id = resolve_file_record_id(reference if isinstance(reference, str) else None)
|
||||
remote_url = file.get("remote_url")
|
||||
normalized_url = remote_url if isinstance(remote_url, str) else ""
|
||||
file_key = (
|
||||
str(file["type"]),
|
||||
str(file["transfer_method"]),
|
||||
normalized_url,
|
||||
upload_file_id or "",
|
||||
)
|
||||
if file_key in existing_file_keys:
|
||||
continue
|
||||
message_files.append(
|
||||
MessageFile(
|
||||
message_id=message.id,
|
||||
type=file["type"],
|
||||
transfer_method=file["transfer_method"],
|
||||
url=file["remote_url"],
|
||||
url=normalized_url,
|
||||
belongs_to=MessageFileBelongsTo.ASSISTANT,
|
||||
upload_file_id=resolve_file_record_id(reference if isinstance(reference, str) else None),
|
||||
upload_file_id=upload_file_id,
|
||||
created_by_role=CreatorUserRole.ACCOUNT
|
||||
if message.invoke_from in {InvokeFrom.EXPLORE, InvokeFrom.DEBUGGER}
|
||||
else CreatorUserRole.END_USER,
|
||||
created_by=message.from_account_id or message.from_end_user_id or "",
|
||||
)
|
||||
)
|
||||
existing_file_keys.add(file_key)
|
||||
session.add_all(message_files)
|
||||
|
||||
def _seed_graph_runtime_state_from_queue_manager(self) -> None:
|
||||
|
||||
@@ -11,6 +11,7 @@ from core.app.apps.draft_variable_saver import (
|
||||
DraftVariableSaverFactory,
|
||||
NoopDraftVariableSaver,
|
||||
)
|
||||
from core.app.apps.streaming_utils import StreamEventWithCursor, close_stream
|
||||
from core.app.entities.app_invoke_entities import InvokeFrom, UserFrom
|
||||
from core.app.file_access import DatabaseFileAccessController, FileAccessScope, bind_file_access_scope
|
||||
from extensions.ext_database import db
|
||||
@@ -91,6 +92,7 @@ class BaseAppGenerator:
|
||||
try:
|
||||
yield from response_stream
|
||||
finally:
|
||||
close_stream(response_stream)
|
||||
BaseAppGenerator._join_worker_thread(worker_thread)
|
||||
|
||||
@staticmethod
|
||||
@@ -300,7 +302,10 @@ class BaseAppGenerator:
|
||||
return value
|
||||
|
||||
@classmethod
|
||||
def convert_to_event_stream(cls, generator: Union[Mapping, Generator[Mapping | str, None, None]]):
|
||||
def convert_to_event_stream(
|
||||
cls,
|
||||
generator: Union[Mapping, Generator[Mapping | StreamEventWithCursor | str, None, None]],
|
||||
):
|
||||
"""
|
||||
Convert messages into event stream
|
||||
"""
|
||||
@@ -309,11 +314,16 @@ class BaseAppGenerator:
|
||||
else:
|
||||
|
||||
def gen():
|
||||
for message in generator:
|
||||
if isinstance(message, Mapping | dict):
|
||||
yield f"data: {orjson_dumps(message)}\n\n"
|
||||
else:
|
||||
yield f"event: {message}\n\n"
|
||||
try:
|
||||
for message in generator:
|
||||
if isinstance(message, StreamEventWithCursor):
|
||||
yield f"id: {message.cursor}\ndata: {orjson_dumps(message.event)}\n\n"
|
||||
elif isinstance(message, Mapping | dict):
|
||||
yield f"data: {orjson_dumps(message)}\n\n"
|
||||
else:
|
||||
yield f"event: {message}\n\n"
|
||||
finally:
|
||||
close_stream(generator)
|
||||
|
||||
return gen()
|
||||
|
||||
|
||||
@@ -108,6 +108,10 @@ class AppQueueManager(ABC):
|
||||
self._clear_task_belong_cache()
|
||||
self._q.put(None)
|
||||
|
||||
def mark_execution_terminal(self) -> None:
|
||||
"""Confirm that the consumer durably finalized this local segment."""
|
||||
self._execution_terminal.set()
|
||||
|
||||
def _abort_execution(self, reason: str) -> None:
|
||||
"""Propagate response timeout/disconnect to legacy and GraphEngine runners."""
|
||||
with self._lifecycle_lock:
|
||||
@@ -117,9 +121,20 @@ class AppQueueManager(ABC):
|
||||
|
||||
try:
|
||||
self.set_stop_flag_no_user_check(self._task_id)
|
||||
except Exception:
|
||||
logger.exception("Failed to set the stop flag for task %s", self._task_id)
|
||||
try:
|
||||
from services.workflow_handoff_cancellation_service import (
|
||||
request_workflow_handoff_cancel_by_task_id,
|
||||
)
|
||||
|
||||
request_workflow_handoff_cancel_by_task_id(self._task_id, reason=reason)
|
||||
except Exception:
|
||||
logger.exception("Failed to cancel a durable handoff for task %s", self._task_id)
|
||||
try:
|
||||
GraphEngineManager(redis_client).send_stop_command(self._task_id, reason=reason)
|
||||
except Exception:
|
||||
logger.exception("Failed to abort app execution for task %s", self._task_id)
|
||||
logger.exception("Failed to send the graph abort command for task %s", self._task_id)
|
||||
|
||||
def _clear_task_belong_cache(self) -> None:
|
||||
"""
|
||||
@@ -176,21 +191,22 @@ class AppQueueManager(ABC):
|
||||
raise NotImplementedError
|
||||
|
||||
@classmethod
|
||||
def set_stop_flag(cls, task_id: str, invoke_from: InvokeFrom, user_id: str):
|
||||
def set_stop_flag(cls, task_id: str, invoke_from: InvokeFrom, user_id: str) -> bool:
|
||||
"""
|
||||
Set task stop flag
|
||||
:return:
|
||||
"""
|
||||
result: Any | None = redis_client.get(cls._generate_task_belong_cache_key(task_id))
|
||||
if result is None:
|
||||
return
|
||||
return False
|
||||
|
||||
user_prefix = "account" if invoke_from in {InvokeFrom.EXPLORE, InvokeFrom.DEBUGGER} else "end-user"
|
||||
if result.decode("utf-8") != f"{user_prefix}-{user_id}":
|
||||
return
|
||||
return False
|
||||
|
||||
stopped_cache_key = cls._generate_stopped_cache_key(task_id)
|
||||
redis_client.setex(stopped_cache_key, 600, 1)
|
||||
return True
|
||||
|
||||
@classmethod
|
||||
def set_stop_flag_no_user_check(cls, task_id: str) -> None:
|
||||
|
||||
@@ -145,6 +145,7 @@ class WorkflowResponseConverter:
|
||||
self._node_snapshots: dict[NodeExecutionId, _NodeSnapshot] = {}
|
||||
self._workflow_execution_id: str | None = None
|
||||
self._workflow_started_at: datetime | None = None
|
||||
self._workflow_handoff_duration = 0.0
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Workflow lifecycle helpers
|
||||
@@ -241,10 +242,13 @@ class WorkflowResponseConverter:
|
||||
workflow_run_id: str,
|
||||
workflow_id: str,
|
||||
reason: WorkflowStartReason,
|
||||
logical_started_at: datetime | None = None,
|
||||
handoff_duration: float = 0.0,
|
||||
) -> WorkflowStartStreamResponse:
|
||||
run_id = self._ensure_workflow_run_id(workflow_run_id)
|
||||
started_at = naive_utc_now()
|
||||
started_at = logical_started_at or naive_utc_now()
|
||||
self._workflow_started_at = started_at
|
||||
self._workflow_handoff_duration = max(handoff_duration, 0.0)
|
||||
|
||||
return WorkflowStartStreamResponse(
|
||||
task_id=task_id,
|
||||
@@ -276,7 +280,7 @@ class WorkflowResponseConverter:
|
||||
)
|
||||
|
||||
finished_at = naive_utc_now()
|
||||
elapsed_time = (finished_at - started_at).total_seconds()
|
||||
elapsed_time = max((finished_at - started_at).total_seconds(), 0.0)
|
||||
|
||||
outputs_mapping = graph_runtime_state.outputs or {}
|
||||
encoded_outputs = WorkflowRuntimeTypeConverter().to_json_encodable(outputs_mapping)
|
||||
@@ -313,6 +317,7 @@ class WorkflowResponseConverter:
|
||||
finished_at=int(finished_at.timestamp()),
|
||||
files=self.fetch_files_from_node_outputs(outputs_mapping),
|
||||
exceptions_count=exceptions_count,
|
||||
handoff_duration=self._workflow_handoff_duration,
|
||||
),
|
||||
)
|
||||
|
||||
@@ -330,7 +335,7 @@ class WorkflowResponseConverter:
|
||||
"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()
|
||||
elapsed_time = max((paused_at - started_at).total_seconds(), 0.0)
|
||||
encoded_outputs = self._encode_outputs(event.outputs) or {}
|
||||
if self._application_generate_entity.invoke_from == InvokeFrom.SERVICE_API:
|
||||
encoded_outputs = {}
|
||||
@@ -418,6 +423,7 @@ class WorkflowResponseConverter:
|
||||
elapsed_time=elapsed_time,
|
||||
total_tokens=graph_runtime_state.total_tokens,
|
||||
total_steps=graph_runtime_state.node_run_steps,
|
||||
handoff_duration=self._workflow_handoff_duration,
|
||||
),
|
||||
)
|
||||
)
|
||||
@@ -503,6 +509,7 @@ class WorkflowResponseConverter:
|
||||
finished_at=int(finished_at.timestamp()),
|
||||
files=cls.fetch_files_from_node_outputs(encoded_outputs),
|
||||
exceptions_count=workflow_run.exceptions_count,
|
||||
handoff_duration=workflow_run.handoff_duration,
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
@@ -10,7 +10,7 @@ from core.app.app_config.entities import EasyUIBasedAppConfig, EasyUIBasedAppMod
|
||||
from core.app.apps.base_app_generator import BaseAppGenerator
|
||||
from core.app.apps.base_app_queue_manager import AppQueueManager
|
||||
from core.app.apps.exc import GenerateTaskStoppedError
|
||||
from core.app.apps.streaming_utils import stream_topic_events
|
||||
from core.app.apps.streaming_utils import StreamEventWithCursor, stream_topic_events
|
||||
from core.app.entities.app_invoke_entities import (
|
||||
AdvancedChatAppGenerateEntity,
|
||||
AgentChatAppGenerateEntity,
|
||||
@@ -321,10 +321,18 @@ class MessageBasedAppGenerator(BaseAppGenerator):
|
||||
workflow_run_id: str,
|
||||
idle_timeout=300,
|
||||
on_subscribe: Callable[[], None] | None = None,
|
||||
) -> Generator[Mapping | str, None, None]:
|
||||
cursor: str | None = None,
|
||||
) -> Generator[Mapping | StreamEventWithCursor | str, None, None]:
|
||||
topic = cls.get_response_topic(app_mode, workflow_run_id)
|
||||
if cursor is None:
|
||||
return stream_topic_events(
|
||||
topic=topic,
|
||||
idle_timeout=idle_timeout,
|
||||
on_subscribe=on_subscribe,
|
||||
)
|
||||
return stream_topic_events(
|
||||
topic=topic,
|
||||
idle_timeout=idle_timeout,
|
||||
on_subscribe=on_subscribe,
|
||||
cursor=cursor,
|
||||
)
|
||||
|
||||
@@ -10,6 +10,8 @@ from core.app.entities.queue_entities import (
|
||||
QueueErrorEvent,
|
||||
QueueMessageEndEvent,
|
||||
QueueStopEvent,
|
||||
QueueWorkflowMaintenancePausedEvent,
|
||||
QueueWorkflowPausedEvent,
|
||||
)
|
||||
from models.model import AppMode
|
||||
|
||||
@@ -43,9 +45,16 @@ class MessageBasedAppQueueManager(AppQueueManager):
|
||||
self._q.put(message)
|
||||
|
||||
if isinstance(
|
||||
event, QueueStopEvent | QueueErrorEvent | QueueMessageEndEvent | QueueAdvancedChatMessageEndEvent
|
||||
event,
|
||||
QueueStopEvent
|
||||
| QueueErrorEvent
|
||||
| QueueMessageEndEvent
|
||||
| QueueAdvancedChatMessageEndEvent
|
||||
| QueueWorkflowPausedEvent,
|
||||
):
|
||||
self.stop_listen(execution_terminal=True)
|
||||
elif isinstance(event, QueueWorkflowMaintenancePausedEvent):
|
||||
self.stop_listen(execution_terminal=False)
|
||||
|
||||
if pub_from == PublishFrom.APPLICATION_MANAGER and self._is_stopped():
|
||||
if self._app_mode == AppMode.ADVANCED_CHAT.value:
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
from collections.abc import Callable, Generator, Iterable, Mapping
|
||||
|
||||
from core.app.apps.streaming_utils import stream_topic_events
|
||||
from core.app.apps.streaming_utils import StreamEventWithCursor, stream_topic_events
|
||||
from core.app.entities.task_entities import StreamEvent
|
||||
from extensions.ext_redis import get_pubsub_broadcast_channel
|
||||
from libs.broadcast_channel.channel import Topic
|
||||
@@ -28,12 +28,22 @@ class MessageGenerator:
|
||||
ping_interval: float = 10.0,
|
||||
on_subscribe: Callable[[], None] | None = None,
|
||||
terminal_events: Iterable[str | StreamEvent] | None = None,
|
||||
) -> Generator[Mapping | str, None, None]:
|
||||
cursor: str | None = None,
|
||||
) -> Generator[Mapping | StreamEventWithCursor | str, None, None]:
|
||||
topic = cls.get_response_topic(app_mode, workflow_run_id)
|
||||
if cursor is None:
|
||||
return stream_topic_events(
|
||||
topic=topic,
|
||||
idle_timeout=idle_timeout,
|
||||
ping_interval=ping_interval,
|
||||
on_subscribe=on_subscribe,
|
||||
terminal_events=terminal_events,
|
||||
)
|
||||
return stream_topic_events(
|
||||
topic=topic,
|
||||
idle_timeout=idle_timeout,
|
||||
ping_interval=ping_interval,
|
||||
on_subscribe=on_subscribe,
|
||||
terminal_events=terminal_events,
|
||||
cursor=cursor,
|
||||
)
|
||||
|
||||
@@ -2,6 +2,7 @@ from collections.abc import Generator
|
||||
from typing import Any, cast, override
|
||||
|
||||
from core.app.apps.base_app_generate_response_converter import AppGenerateResponseConverter
|
||||
from core.app.apps.streaming_utils import close_stream
|
||||
from core.app.entities.task_entities import (
|
||||
AppStreamResponse,
|
||||
ErrorStreamResponse,
|
||||
@@ -44,29 +45,32 @@ class WorkflowAppGenerateResponseConverter(AppGenerateResponseConverter[Workflow
|
||||
:param stream_response: stream response
|
||||
:return:
|
||||
"""
|
||||
for chunk in stream_response:
|
||||
chunk = cast(WorkflowAppStreamResponse, chunk)
|
||||
sub_stream_response = chunk.stream_response
|
||||
try:
|
||||
for chunk in stream_response:
|
||||
chunk = cast(WorkflowAppStreamResponse, chunk)
|
||||
sub_stream_response = chunk.stream_response
|
||||
|
||||
match sub_stream_response:
|
||||
case PingStreamResponse():
|
||||
yield "ping"
|
||||
continue
|
||||
case ErrorStreamResponse():
|
||||
response_chunk = {
|
||||
"event": sub_stream_response.event.value,
|
||||
"workflow_run_id": chunk.workflow_run_id,
|
||||
}
|
||||
data = cls._error_to_stream_response(sub_stream_response.err)
|
||||
response_chunk.update(cast(dict, data))
|
||||
case _:
|
||||
response_chunk = {
|
||||
"event": sub_stream_response.event.value,
|
||||
"workflow_run_id": chunk.workflow_run_id,
|
||||
}
|
||||
response_chunk.update(sub_stream_response.model_dump())
|
||||
match sub_stream_response:
|
||||
case PingStreamResponse():
|
||||
yield "ping"
|
||||
continue
|
||||
case ErrorStreamResponse():
|
||||
response_chunk = {
|
||||
"event": sub_stream_response.event.value,
|
||||
"workflow_run_id": chunk.workflow_run_id,
|
||||
}
|
||||
data = cls._error_to_stream_response(sub_stream_response.err)
|
||||
response_chunk.update(cast(dict, data))
|
||||
case _:
|
||||
response_chunk = {
|
||||
"event": sub_stream_response.event.value,
|
||||
"workflow_run_id": chunk.workflow_run_id,
|
||||
}
|
||||
response_chunk.update(sub_stream_response.model_dump())
|
||||
|
||||
yield response_chunk
|
||||
yield response_chunk
|
||||
finally:
|
||||
close_stream(stream_response)
|
||||
|
||||
@classmethod
|
||||
@override
|
||||
@@ -78,32 +82,35 @@ class WorkflowAppGenerateResponseConverter(AppGenerateResponseConverter[Workflow
|
||||
:param stream_response: stream response
|
||||
:return:
|
||||
"""
|
||||
for chunk in stream_response:
|
||||
chunk = cast(WorkflowAppStreamResponse, chunk)
|
||||
sub_stream_response = chunk.stream_response
|
||||
try:
|
||||
for chunk in stream_response:
|
||||
chunk = cast(WorkflowAppStreamResponse, chunk)
|
||||
sub_stream_response = chunk.stream_response
|
||||
|
||||
match sub_stream_response:
|
||||
case PingStreamResponse():
|
||||
yield "ping"
|
||||
continue
|
||||
case ErrorStreamResponse():
|
||||
response_chunk = {
|
||||
"event": sub_stream_response.event.value,
|
||||
"workflow_run_id": chunk.workflow_run_id,
|
||||
}
|
||||
data = cls._error_to_stream_response(sub_stream_response.err)
|
||||
response_chunk.update(cast(dict, data))
|
||||
case NodeStartStreamResponse() | NodeFinishStreamResponse():
|
||||
response_chunk = {
|
||||
"event": sub_stream_response.event.value,
|
||||
"workflow_run_id": chunk.workflow_run_id,
|
||||
}
|
||||
response_chunk.update(cast(dict, sub_stream_response.to_ignore_detail_dict()))
|
||||
case _:
|
||||
response_chunk = {
|
||||
"event": sub_stream_response.event.value,
|
||||
"workflow_run_id": chunk.workflow_run_id,
|
||||
}
|
||||
response_chunk.update(sub_stream_response.model_dump())
|
||||
match sub_stream_response:
|
||||
case PingStreamResponse():
|
||||
yield "ping"
|
||||
continue
|
||||
case ErrorStreamResponse():
|
||||
response_chunk = {
|
||||
"event": sub_stream_response.event.value,
|
||||
"workflow_run_id": chunk.workflow_run_id,
|
||||
}
|
||||
data = cls._error_to_stream_response(sub_stream_response.err)
|
||||
response_chunk.update(cast(dict, data))
|
||||
case NodeStartStreamResponse() | NodeFinishStreamResponse():
|
||||
response_chunk = {
|
||||
"event": sub_stream_response.event.value,
|
||||
"workflow_run_id": chunk.workflow_run_id,
|
||||
}
|
||||
response_chunk.update(cast(dict, sub_stream_response.to_ignore_detail_dict()))
|
||||
case _:
|
||||
response_chunk = {
|
||||
"event": sub_stream_response.event.value,
|
||||
"workflow_run_id": chunk.workflow_run_id,
|
||||
}
|
||||
response_chunk.update(sub_stream_response.model_dump())
|
||||
|
||||
yield response_chunk
|
||||
yield response_chunk
|
||||
finally:
|
||||
close_stream(stream_response)
|
||||
|
||||
@@ -6,7 +6,7 @@ import secrets
|
||||
import threading
|
||||
import time
|
||||
import uuid
|
||||
from collections.abc import Generator, Mapping
|
||||
from collections.abc import Generator, Mapping, Sequence
|
||||
from typing import Any, Literal, cast, overload
|
||||
|
||||
from flask import Flask, current_app
|
||||
@@ -23,6 +23,7 @@ from core.app.apps.exc import GenerateTaskStoppedError
|
||||
from core.app.apps.pipeline.pipeline_config_manager import PipelineConfigManager
|
||||
from core.app.apps.pipeline.pipeline_queue_manager import PipelineQueueManager
|
||||
from core.app.apps.pipeline.pipeline_runner import PipelineRunner
|
||||
from core.app.apps.workflow.active_workflow_tasks import active_workflow_task
|
||||
from core.app.apps.workflow.generate_response_converter import WorkflowAppGenerateResponseConverter
|
||||
from core.app.apps.workflow.generate_task_pipeline import WorkflowAppGenerateTaskPipeline
|
||||
from core.app.entities.app_invoke_entities import InvokeFrom, RagPipelineGenerateEntity
|
||||
@@ -31,7 +32,9 @@ from core.app.entities.task_entities import (
|
||||
WorkflowAppBlockingResponse,
|
||||
WorkflowAppPausedBlockingResponse,
|
||||
WorkflowAppStreamResponse,
|
||||
WorkflowMaintenancePausedBlockingResponse,
|
||||
)
|
||||
from core.app.layers.pause_state_persist_layer import PauseStateLayerConfig, PauseStatePersistenceLayer
|
||||
from core.datasource.entities.datasource_entities import (
|
||||
DatasourceProviderType,
|
||||
OnlineDriveBrowseFilesRequest,
|
||||
@@ -45,16 +48,21 @@ from core.repositories.factory import (
|
||||
WorkflowNodeExecutionRepository,
|
||||
)
|
||||
from extensions.ext_database import db
|
||||
from graphon.filters import ResponseStreamFilter
|
||||
from graphon.graph_engine.layers import GraphEngineLayer
|
||||
from graphon.model_runtime.errors.invoke import InvokeAuthorizationError
|
||||
from graphon.runtime import GraphRuntimeState
|
||||
from graphon.variable_loader import DUMMY_VARIABLE_LOADER, VariableLoader
|
||||
from libs.flask_utils import preserve_flask_contexts
|
||||
from models import Account, EndUser, Workflow, WorkflowNodeExecutionTriggeredFrom
|
||||
from models.dataset import Document, DocumentPipelineExecutionLog, Pipeline
|
||||
from models.enums import WorkflowRunTriggeredFrom
|
||||
from models.model import AppMode
|
||||
from models.workflow_handoff import WorkflowHandoffResumeRoute
|
||||
from services.datasource_provider_service import DatasourceProviderService
|
||||
from services.rag_pipeline.rag_pipeline_task_proxy import RagPipelineTaskProxy
|
||||
from services.workflow_draft_variable_service import DraftVarLoader, WorkflowDraftVariableService
|
||||
from services.workflow_handoff_runtime_service import build_workflow_handoff_persistence_layer
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -73,6 +81,7 @@ class PipelineGenerator(BaseAppGenerator):
|
||||
streaming: Literal[True],
|
||||
call_depth: int,
|
||||
workflow_thread_pool_id: str | None,
|
||||
workflow_run_id: str | uuid.UUID | None = None,
|
||||
is_retry: bool = False,
|
||||
) -> Generator[Mapping | str, None, None]: ...
|
||||
|
||||
@@ -89,6 +98,7 @@ class PipelineGenerator(BaseAppGenerator):
|
||||
streaming: Literal[False],
|
||||
call_depth: int,
|
||||
workflow_thread_pool_id: str | None,
|
||||
workflow_run_id: str | uuid.UUID | None = None,
|
||||
is_retry: bool = False,
|
||||
) -> Mapping[str, Any]: ...
|
||||
|
||||
@@ -105,6 +115,7 @@ class PipelineGenerator(BaseAppGenerator):
|
||||
streaming: bool,
|
||||
call_depth: int,
|
||||
workflow_thread_pool_id: str | None,
|
||||
workflow_run_id: str | uuid.UUID | None = None,
|
||||
is_retry: bool = False,
|
||||
) -> Mapping[str, Any] | Generator[Mapping | str, None, None]: ...
|
||||
|
||||
@@ -120,6 +131,7 @@ class PipelineGenerator(BaseAppGenerator):
|
||||
streaming: bool = True,
|
||||
call_depth: int = 0,
|
||||
workflow_thread_pool_id: str | None = None,
|
||||
workflow_run_id: str | uuid.UUID | None = None,
|
||||
is_retry: bool = False,
|
||||
) -> Mapping[str, Any] | Generator[Mapping | str, None, None] | None:
|
||||
# Add null check for dataset
|
||||
@@ -167,7 +179,7 @@ class PipelineGenerator(BaseAppGenerator):
|
||||
# run in child thread
|
||||
rag_pipeline_invoke_entities = []
|
||||
for i, datasource_info in enumerate(datasource_info_list):
|
||||
workflow_run_id = str(uuid.uuid4())
|
||||
current_workflow_run_id = str(workflow_run_id if i == 0 and workflow_run_id else uuid.uuid4())
|
||||
document_id = args.get("original_document_id") or None
|
||||
if invoke_from == InvokeFrom.PUBLISHED_PIPELINE and not is_retry:
|
||||
document_id = document_id or documents[i].id
|
||||
@@ -203,7 +215,7 @@ class PipelineGenerator(BaseAppGenerator):
|
||||
stream=streaming,
|
||||
invoke_from=invoke_from,
|
||||
call_depth=call_depth,
|
||||
workflow_execution_id=workflow_run_id,
|
||||
workflow_execution_id=current_workflow_run_id,
|
||||
)
|
||||
|
||||
contexts.plugin_tool_providers.set({})
|
||||
@@ -252,7 +264,7 @@ class PipelineGenerator(BaseAppGenerator):
|
||||
tenant_id=pipeline.tenant_id,
|
||||
workflow_id=workflow.id,
|
||||
streaming=streaming,
|
||||
workflow_execution_id=workflow_run_id,
|
||||
workflow_execution_id=current_workflow_run_id,
|
||||
workflow_thread_pool_id=workflow_thread_pool_id,
|
||||
application_generate_entity=application_generate_entity.model_dump(),
|
||||
)
|
||||
@@ -286,6 +298,58 @@ class PipelineGenerator(BaseAppGenerator):
|
||||
],
|
||||
}
|
||||
|
||||
def resume(
|
||||
self,
|
||||
*,
|
||||
session: Session,
|
||||
pipeline: Pipeline,
|
||||
workflow: Workflow,
|
||||
user: Account | EndUser,
|
||||
application_generate_entity: RagPipelineGenerateEntity,
|
||||
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,
|
||||
response_stream_filter: ResponseStreamFilter | None = None,
|
||||
workflow_thread_pool_id: str | None = None,
|
||||
handoff_resume_route: WorkflowHandoffResumeRoute | None = None,
|
||||
graph_config: Mapping[str, Any] | None = None,
|
||||
workflow_version: str | None = None,
|
||||
root_node_id: str | None = None,
|
||||
) -> Mapping[str, Any] | Generator[str | Mapping[str, Any], None, None]:
|
||||
"""Resume a RAG pipeline execution from a persisted graph checkpoint.
|
||||
|
||||
Reusing ``_generate`` keeps document ownership checks, queue semantics,
|
||||
and graph layers identical to a newly dispatched pipeline execution.
|
||||
The caller owns the durable handoff claim and must exhaust a streaming
|
||||
result before marking the handoff resumed.
|
||||
"""
|
||||
return self._generate(
|
||||
session=session,
|
||||
flask_app=current_app._get_current_object(), # type: ignore
|
||||
context=contextvars.copy_context(),
|
||||
pipeline=pipeline,
|
||||
workflow_id=workflow.id,
|
||||
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,
|
||||
workflow_thread_pool_id=workflow_thread_pool_id,
|
||||
graph_engine_layers=graph_engine_layers,
|
||||
graph_runtime_state=graph_runtime_state,
|
||||
pause_state_config=pause_state_config,
|
||||
response_stream_filter=response_stream_filter,
|
||||
handoff_resume_route=handoff_resume_route,
|
||||
graph_config=graph_config,
|
||||
workflow_version=workflow_version,
|
||||
root_node_id=root_node_id,
|
||||
)
|
||||
|
||||
def _generate(
|
||||
self,
|
||||
*,
|
||||
@@ -302,6 +366,14 @@ class PipelineGenerator(BaseAppGenerator):
|
||||
streaming: bool = True,
|
||||
variable_loader: VariableLoader = DUMMY_VARIABLE_LOADER,
|
||||
workflow_thread_pool_id: str | None = None,
|
||||
graph_engine_layers: Sequence[GraphEngineLayer] = (),
|
||||
graph_runtime_state: GraphRuntimeState | None = None,
|
||||
pause_state_config: PauseStateLayerConfig | None = None,
|
||||
response_stream_filter: ResponseStreamFilter | None = None,
|
||||
handoff_resume_route: WorkflowHandoffResumeRoute | None = None,
|
||||
graph_config: Mapping[str, Any] | None = None,
|
||||
workflow_version: str | None = None,
|
||||
root_node_id: str | None = None,
|
||||
) -> Mapping[str, Any] | Generator[str | Mapping[str, Any], None, None]:
|
||||
"""
|
||||
Generate App response.
|
||||
@@ -321,12 +393,39 @@ class PipelineGenerator(BaseAppGenerator):
|
||||
workflow = session.get(Workflow, workflow_id)
|
||||
if not workflow:
|
||||
raise ValueError(f"Workflow not found: {workflow_id}")
|
||||
# The graph runs in a worker thread and uses repository-scoped
|
||||
# sessions. Detach the immutable configuration snapshots and
|
||||
# release the caller's connection before waiting on a potentially
|
||||
# long pipeline execution.
|
||||
if workflow in session:
|
||||
session.expunge(workflow)
|
||||
if pipeline in session:
|
||||
session.expunge(pipeline)
|
||||
session.close()
|
||||
queue_manager = PipelineQueueManager(
|
||||
task_id=application_generate_entity.task_id,
|
||||
user_id=application_generate_entity.user_id,
|
||||
invoke_from=application_generate_entity.invoke_from,
|
||||
app_mode=AppMode.RAG_PIPELINE,
|
||||
)
|
||||
resolved_response_stream_filter = response_stream_filter or ResponseStreamFilter()
|
||||
resolved_graph_engine_layers = list(graph_engine_layers)
|
||||
handoff_layer = build_workflow_handoff_persistence_layer(
|
||||
generate_entity=application_generate_entity,
|
||||
response_stream_filter=resolved_response_stream_filter,
|
||||
resume_route=handoff_resume_route,
|
||||
)
|
||||
if handoff_layer is not None:
|
||||
resolved_graph_engine_layers.append(handoff_layer)
|
||||
if pause_state_config is not None:
|
||||
resolved_graph_engine_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,
|
||||
response_stream_filter=resolved_response_stream_filter,
|
||||
)
|
||||
)
|
||||
context = contextvars.copy_context()
|
||||
|
||||
# new thread
|
||||
@@ -341,6 +440,12 @@ class PipelineGenerator(BaseAppGenerator):
|
||||
"variable_loader": variable_loader,
|
||||
"workflow_execution_repository": workflow_execution_repository,
|
||||
"workflow_node_execution_repository": workflow_node_execution_repository,
|
||||
"graph_engine_layers": tuple(resolved_graph_engine_layers),
|
||||
"graph_runtime_state": graph_runtime_state,
|
||||
"response_stream_filter": resolved_response_stream_filter,
|
||||
"graph_config": graph_config,
|
||||
"workflow_version": workflow_version,
|
||||
"root_node_id": root_node_id,
|
||||
},
|
||||
)
|
||||
|
||||
@@ -384,6 +489,7 @@ class PipelineGenerator(BaseAppGenerator):
|
||||
streaming: bool = True,
|
||||
*,
|
||||
session: Session,
|
||||
workflow_run_id: str | uuid.UUID | None = None,
|
||||
) -> Mapping[str, Any] | Generator[str | Mapping[str, Any], None, None]:
|
||||
"""
|
||||
Generate App response.
|
||||
@@ -427,7 +533,7 @@ class PipelineGenerator(BaseAppGenerator):
|
||||
stream=streaming,
|
||||
invoke_from=InvokeFrom.DEBUGGER,
|
||||
call_depth=0,
|
||||
workflow_execution_id=str(uuid.uuid4()),
|
||||
workflow_execution_id=str(workflow_run_id or uuid.uuid4()),
|
||||
single_iteration_run=RagPipelineGenerateEntity.SingleIterationRunEntity(
|
||||
node_id=node_id, inputs=args["inputs"]
|
||||
),
|
||||
@@ -486,6 +592,7 @@ class PipelineGenerator(BaseAppGenerator):
|
||||
streaming: bool = True,
|
||||
*,
|
||||
session: Session,
|
||||
workflow_run_id: str | uuid.UUID | None = None,
|
||||
) -> Mapping[str, Any] | Generator[str | Mapping[str, Any], None, None]:
|
||||
"""
|
||||
Generate App response.
|
||||
@@ -530,7 +637,7 @@ class PipelineGenerator(BaseAppGenerator):
|
||||
invoke_from=InvokeFrom.DEBUGGER,
|
||||
extras={"auto_generate_conversation_name": False},
|
||||
single_loop_run=RagPipelineGenerateEntity.SingleLoopRunEntity(node_id=node_id, inputs=args["inputs"]),
|
||||
workflow_execution_id=str(uuid.uuid4()),
|
||||
workflow_execution_id=str(workflow_run_id or uuid.uuid4()),
|
||||
)
|
||||
contexts.plugin_tool_providers.set({})
|
||||
contexts.plugin_tool_providers_lock.set(threading.Lock())
|
||||
@@ -587,6 +694,12 @@ class PipelineGenerator(BaseAppGenerator):
|
||||
workflow_execution_repository: WorkflowExecutionRepository,
|
||||
workflow_node_execution_repository: WorkflowNodeExecutionRepository,
|
||||
workflow_thread_pool_id: str | None = None,
|
||||
graph_engine_layers: Sequence[GraphEngineLayer] = (),
|
||||
graph_runtime_state: GraphRuntimeState | None = None,
|
||||
response_stream_filter: ResponseStreamFilter | None = None,
|
||||
graph_config: Mapping[str, Any] | None = None,
|
||||
workflow_version: str | None = None,
|
||||
root_node_id: str | None = None,
|
||||
) -> None:
|
||||
"""
|
||||
Generate worker in a new thread.
|
||||
@@ -635,9 +748,19 @@ class PipelineGenerator(BaseAppGenerator):
|
||||
system_user_id=system_user_id,
|
||||
workflow_execution_repository=workflow_execution_repository,
|
||||
workflow_node_execution_repository=workflow_node_execution_repository,
|
||||
graph_engine_layers=graph_engine_layers,
|
||||
graph_runtime_state=graph_runtime_state,
|
||||
response_stream_filter=response_stream_filter,
|
||||
graph_config=graph_config,
|
||||
workflow_version=workflow_version,
|
||||
root_node_id=root_node_id,
|
||||
)
|
||||
|
||||
runner.run()
|
||||
with active_workflow_task(
|
||||
application_generate_entity.task_id,
|
||||
workflow_run_id=application_generate_entity.workflow_execution_id,
|
||||
):
|
||||
runner.run()
|
||||
except GenerateTaskStoppedError:
|
||||
pass
|
||||
except InvokeAuthorizationError:
|
||||
@@ -668,6 +791,7 @@ class PipelineGenerator(BaseAppGenerator):
|
||||
) -> (
|
||||
WorkflowAppBlockingResponse
|
||||
| WorkflowAppPausedBlockingResponse
|
||||
| WorkflowMaintenancePausedBlockingResponse
|
||||
| Generator[WorkflowAppStreamResponse, None, None]
|
||||
):
|
||||
"""
|
||||
|
||||
@@ -9,7 +9,9 @@ from core.app.entities.queue_entities import (
|
||||
QueueMessageEndEvent,
|
||||
QueueStopEvent,
|
||||
QueueWorkflowFailedEvent,
|
||||
QueueWorkflowMaintenancePausedEvent,
|
||||
QueueWorkflowPartialSuccessEvent,
|
||||
QueueWorkflowPausedEvent,
|
||||
QueueWorkflowSucceededEvent,
|
||||
WorkflowQueueMessage,
|
||||
)
|
||||
@@ -40,9 +42,12 @@ class PipelineQueueManager(AppQueueManager):
|
||||
| QueueMessageEndEvent
|
||||
| QueueWorkflowSucceededEvent
|
||||
| QueueWorkflowFailedEvent
|
||||
| QueueWorkflowPartialSuccessEvent,
|
||||
| QueueWorkflowPartialSuccessEvent
|
||||
| QueueWorkflowPausedEvent,
|
||||
):
|
||||
self.stop_listen(execution_terminal=True)
|
||||
elif isinstance(event, QueueWorkflowMaintenancePausedEvent):
|
||||
self.stop_listen(execution_terminal=False)
|
||||
|
||||
if pub_from == PublishFrom.APPLICATION_MANAGER and self._is_stopped():
|
||||
raise GenerateTaskStoppedError()
|
||||
|
||||
@@ -1,12 +1,15 @@
|
||||
import logging
|
||||
import time
|
||||
from typing import cast
|
||||
from collections.abc import Mapping, Sequence
|
||||
from typing import Any, cast
|
||||
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from configs import dify_config
|
||||
from core.app.apps.base_app_queue_manager import AppQueueManager
|
||||
from core.app.apps.pipeline.pipeline_config_manager import PipelineConfig
|
||||
from core.app.apps.workflow.command_channels import CelerySignalCommandChannel, CombinedCommandChannel
|
||||
from core.app.apps.workflow_app_runner import WorkflowBasedAppRunner
|
||||
from core.app.entities.app_invoke_entities import (
|
||||
InvokeFrom,
|
||||
@@ -14,6 +17,7 @@ from core.app.entities.app_invoke_entities import (
|
||||
UserFrom,
|
||||
build_dify_run_context,
|
||||
)
|
||||
from core.app.layers.pause_state_persist_layer import get_workflow_handoff_active_execution_seconds
|
||||
from core.app.workflow.layers.persistence import PersistenceWorkflowInfo, WorkflowPersistenceLayer
|
||||
from core.db.session_factory import create_session
|
||||
from core.repositories.factory import WorkflowExecutionRepository, WorkflowNodeExecutionRepository
|
||||
@@ -21,8 +25,13 @@ from core.workflow.node_factory import DifyGraphInitContext, DifyNodeFactory, ge
|
||||
from core.workflow.system_variables import build_bootstrap_variables, build_system_variables
|
||||
from core.workflow.variable_pool_initializer import add_node_inputs_to_pool, add_variables_to_pool
|
||||
from core.workflow.workflow_entry import WorkflowEntry
|
||||
from extensions.ext_redis import redis_client
|
||||
from extensions.workflow_warm_shutdown import celery_warm_shutdown_started
|
||||
from graphon.enums import WorkflowType
|
||||
from graphon.filters import ResponseStreamFilter
|
||||
from graphon.graph import Graph
|
||||
from graphon.graph_engine.command_channels import RedisChannel
|
||||
from graphon.graph_engine.layers import GraphEngineLayer
|
||||
from graphon.graph_events import GraphEngineEvent, GraphRunFailedEvent
|
||||
from graphon.runtime import GraphRuntimeState, VariablePool
|
||||
from graphon.variable_loader import VariableLoader
|
||||
@@ -50,6 +59,12 @@ class PipelineRunner(WorkflowBasedAppRunner):
|
||||
workflow_execution_repository: WorkflowExecutionRepository,
|
||||
workflow_node_execution_repository: WorkflowNodeExecutionRepository,
|
||||
workflow_thread_pool_id: str | None = None,
|
||||
graph_engine_layers: Sequence[GraphEngineLayer] = (),
|
||||
graph_runtime_state: GraphRuntimeState | None = None,
|
||||
response_stream_filter: ResponseStreamFilter | None = None,
|
||||
graph_config: Mapping[str, Any] | None = None,
|
||||
workflow_version: str | None = None,
|
||||
root_node_id: str | None = None,
|
||||
) -> None:
|
||||
"""
|
||||
:param application_generate_entity: application generate entity
|
||||
@@ -60,6 +75,7 @@ class PipelineRunner(WorkflowBasedAppRunner):
|
||||
queue_manager=queue_manager,
|
||||
variable_loader=variable_loader,
|
||||
app_id=application_generate_entity.app_config.app_id,
|
||||
graph_engine_layers=graph_engine_layers,
|
||||
)
|
||||
self.application_generate_entity = application_generate_entity
|
||||
self.workflow_thread_pool_id = workflow_thread_pool_id
|
||||
@@ -67,6 +83,11 @@ class PipelineRunner(WorkflowBasedAppRunner):
|
||||
self._sys_user_id = system_user_id
|
||||
self._workflow_execution_repository = workflow_execution_repository
|
||||
self._workflow_node_execution_repository = workflow_node_execution_repository
|
||||
self._resume_graph_runtime_state = graph_runtime_state
|
||||
self._response_stream_filter = response_stream_filter
|
||||
self._execution_graph_config = graph_config
|
||||
self._workflow_version = workflow_version
|
||||
self._root_node_id = root_node_id
|
||||
|
||||
def _get_app_id(self) -> str:
|
||||
return self.application_generate_entity.app_config.app_id
|
||||
@@ -133,14 +154,33 @@ class PipelineRunner(WorkflowBasedAppRunner):
|
||||
session.expunge(pipeline)
|
||||
session.expunge(workflow)
|
||||
|
||||
resume_state = self._resume_graph_runtime_state
|
||||
graph_config = self._execution_graph_config if self._execution_graph_config is not None else workflow.graph_dict
|
||||
workflow_version = self._workflow_version if self._workflow_version is not None else workflow.version
|
||||
if resume_state is not None:
|
||||
graph_runtime_state = resume_state
|
||||
variable_pool = graph_runtime_state.variable_pool
|
||||
graph = self._init_rag_pipeline_graph(
|
||||
graph_runtime_state=graph_runtime_state,
|
||||
start_node_id=self._root_node_id or self.application_generate_entity.start_node_id,
|
||||
workflow=workflow,
|
||||
graph_config=graph_config,
|
||||
user_from=user_from,
|
||||
invoke_from=invoke_from,
|
||||
skip_validation=bool(
|
||||
self.application_generate_entity.single_iteration_run
|
||||
or self.application_generate_entity.single_loop_run
|
||||
),
|
||||
)
|
||||
# if only single iteration run is requested
|
||||
if self.application_generate_entity.single_iteration_run or self.application_generate_entity.single_loop_run:
|
||||
elif self.application_generate_entity.single_iteration_run or self.application_generate_entity.single_loop_run:
|
||||
# Handle single iteration or single loop run
|
||||
graph, variable_pool, graph_runtime_state = self._prepare_single_node_execution(
|
||||
graph, graph_config, variable_pool, graph_runtime_state = self._prepare_single_node_execution(
|
||||
workflow=workflow,
|
||||
single_iteration_run=self.application_generate_entity.single_iteration_run,
|
||||
single_loop_run=self.application_generate_entity.single_loop_run,
|
||||
user_id=self.application_generate_entity.user_id,
|
||||
workflow_execution_id=self.application_generate_entity.workflow_execution_id,
|
||||
)
|
||||
else:
|
||||
inputs = self.application_generate_entity.inputs
|
||||
@@ -197,23 +237,46 @@ class PipelineRunner(WorkflowBasedAppRunner):
|
||||
graph_runtime_state=graph_runtime_state,
|
||||
start_node_id=root_node_id,
|
||||
workflow=workflow,
|
||||
graph_config=graph_config,
|
||||
user_from=user_from,
|
||||
invoke_from=invoke_from,
|
||||
)
|
||||
|
||||
# RUN WORKFLOW
|
||||
self._configure_execution_root(graph.root_node.id)
|
||||
task_id = self.application_generate_entity.task_id
|
||||
channel_key = f"workflow:{task_id}:commands"
|
||||
command_channel = CombinedCommandChannel(
|
||||
(
|
||||
RedisChannel(redis_client, channel_key),
|
||||
CelerySignalCommandChannel(
|
||||
shutdown_state_getter=celery_warm_shutdown_started,
|
||||
pause_on_shutdown=(
|
||||
dify_config.WORKFLOW_HANDOFF_ENABLED and self.application_generate_entity.call_depth == 0
|
||||
),
|
||||
ignore_shutdown=(
|
||||
dify_config.WORKFLOW_HANDOFF_ENABLED and self.application_generate_entity.call_depth > 0
|
||||
),
|
||||
),
|
||||
)
|
||||
)
|
||||
workflow_entry = WorkflowEntry(
|
||||
tenant_id=workflow.tenant_id,
|
||||
app_id=workflow.app_id,
|
||||
workflow_id=workflow.id,
|
||||
graph=graph,
|
||||
graph_config=workflow.graph_dict,
|
||||
graph_config=graph_config,
|
||||
user_id=self.application_generate_entity.user_id,
|
||||
user_from=user_from,
|
||||
invoke_from=invoke_from,
|
||||
call_depth=self.application_generate_entity.call_depth,
|
||||
graph_runtime_state=graph_runtime_state,
|
||||
variable_pool=variable_pool,
|
||||
command_channel=command_channel,
|
||||
response_stream_filter=self._response_stream_filter,
|
||||
prior_active_execution_seconds=get_workflow_handoff_active_execution_seconds(
|
||||
self.application_generate_entity
|
||||
),
|
||||
)
|
||||
|
||||
self._queue_manager.graph_runtime_state = graph_runtime_state
|
||||
@@ -223,8 +286,8 @@ class PipelineRunner(WorkflowBasedAppRunner):
|
||||
workflow_info=PersistenceWorkflowInfo(
|
||||
workflow_id=workflow.id,
|
||||
workflow_type=WorkflowType(workflow.type),
|
||||
version=workflow.version,
|
||||
graph_data=workflow.graph_dict,
|
||||
version=workflow_version,
|
||||
graph_data=graph_config,
|
||||
),
|
||||
workflow_execution_repository=self._workflow_execution_repository,
|
||||
workflow_node_execution_repository=self._workflow_node_execution_repository,
|
||||
@@ -232,12 +295,14 @@ class PipelineRunner(WorkflowBasedAppRunner):
|
||||
)
|
||||
|
||||
workflow_entry.graph_engine.layer(persistence_layer)
|
||||
for layer in self._graph_engine_layers:
|
||||
workflow_entry.graph_engine.layer(layer)
|
||||
|
||||
generator = workflow_entry.run()
|
||||
|
||||
for event in generator:
|
||||
self._update_document_status(event, document_ref)
|
||||
self._handle_event(workflow_entry, event)
|
||||
self._handle_event_with_handoff_contracts(workflow_entry, event)
|
||||
|
||||
def get_workflow(self, session: Session, pipeline: Pipeline, workflow_id: str) -> Workflow | None:
|
||||
"""
|
||||
@@ -258,13 +323,15 @@ class PipelineRunner(WorkflowBasedAppRunner):
|
||||
workflow: Workflow,
|
||||
graph_runtime_state: GraphRuntimeState,
|
||||
start_node_id: str | None = None,
|
||||
graph_config: Mapping[str, Any] | None = None,
|
||||
user_from: UserFrom = UserFrom.ACCOUNT,
|
||||
invoke_from: InvokeFrom = InvokeFrom.SERVICE_API,
|
||||
skip_validation: bool = False,
|
||||
) -> Graph:
|
||||
"""
|
||||
Init pipeline graph
|
||||
"""
|
||||
graph_config = workflow.graph_dict
|
||||
graph_config = graph_config if graph_config is not None else workflow.graph_dict
|
||||
if "nodes" not in graph_config or "edges" not in graph_config:
|
||||
raise ValueError("nodes or edges not found in workflow graph")
|
||||
|
||||
@@ -316,7 +383,12 @@ class PipelineRunner(WorkflowBasedAppRunner):
|
||||
)
|
||||
if start_node_id is None:
|
||||
start_node_id = get_default_root_node_id(graph_config)
|
||||
graph = Graph.init(graph_config=graph_config, node_factory=node_factory, root_node_id=start_node_id)
|
||||
graph = Graph.init(
|
||||
graph_config=graph_config,
|
||||
node_factory=node_factory,
|
||||
root_node_id=start_node_id,
|
||||
skip_validation=skip_validation,
|
||||
)
|
||||
|
||||
if not graph:
|
||||
raise ValueError("graph not found in workflow")
|
||||
|
||||
@@ -2,14 +2,59 @@ from __future__ import annotations
|
||||
|
||||
import json
|
||||
import time
|
||||
from collections.abc import Callable, Generator, Iterable, Mapping
|
||||
from typing import Any
|
||||
from collections.abc import Callable, Generator, Iterable, Iterator, Mapping
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Protocol, override, runtime_checkable
|
||||
|
||||
from core.app.entities.task_entities import StreamEvent
|
||||
from libs.broadcast_channel.channel import Topic
|
||||
from libs.broadcast_channel.channel import CursorSubscription, Topic
|
||||
from libs.broadcast_channel.exc import SubscriptionClosedError
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class StreamEventWithCursor:
|
||||
"""A decoded application event and its durable replay cursor."""
|
||||
|
||||
event: Mapping[str, Any]
|
||||
cursor: str
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class _Closable(Protocol):
|
||||
def close(self) -> None: ...
|
||||
|
||||
|
||||
def close_stream(stream: object) -> None:
|
||||
"""Close a stream when its concrete iterator exposes the close protocol."""
|
||||
if isinstance(stream, _Closable):
|
||||
stream.close()
|
||||
|
||||
|
||||
class WorkflowRunIdentifiedStream(Iterator[str]):
|
||||
"""Streaming response carrying its stable logical workflow-run identifier.
|
||||
|
||||
The workflow run is allocated before the first SSE event. Keeping it on
|
||||
the iterable lets the final HTTP boundary expose ``X-Workflow-Run-ID`` even
|
||||
if the socket drops before ``workflow_started`` is delivered.
|
||||
"""
|
||||
|
||||
def __init__(self, stream: Iterable[str], *, workflow_run_id: str) -> None:
|
||||
self._stream = stream
|
||||
self._iterator = iter(stream)
|
||||
self.workflow_run_id = workflow_run_id
|
||||
|
||||
@override
|
||||
def __iter__(self) -> WorkflowRunIdentifiedStream:
|
||||
return self
|
||||
|
||||
@override
|
||||
def __next__(self) -> str:
|
||||
return next(self._iterator)
|
||||
|
||||
def close(self) -> None:
|
||||
close_stream(self._stream)
|
||||
|
||||
|
||||
def stream_topic_events(
|
||||
*,
|
||||
topic: Topic,
|
||||
@@ -17,24 +62,31 @@ def stream_topic_events(
|
||||
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
|
||||
|
||||
cursor: str | None = None,
|
||||
) -> Generator[Mapping[str, Any] | StreamEventWithCursor | str, None, None]:
|
||||
terminal_values = _normalize_terminal_events(terminal_events)
|
||||
last_msg_time = time.time()
|
||||
last_ping_time = last_msg_time
|
||||
with topic.subscribe() as sub:
|
||||
subscription = topic.subscribe(cursor=cursor) if cursor is not None else topic.subscribe()
|
||||
with subscription 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()
|
||||
|
||||
# Do not expose the first response byte until the subscription is live
|
||||
# and task dispatch has succeeded. Otherwise a process can disappear
|
||||
# after returning the stable run ID but before creating any recoverable
|
||||
# execution for it.
|
||||
yield StreamEvent.PING.value
|
||||
while True:
|
||||
try:
|
||||
msg = sub.receive(timeout=1)
|
||||
if isinstance(sub, CursorSubscription):
|
||||
cursor_message = sub.receive_with_cursor(timeout=1)
|
||||
msg = None if cursor_message is None else cursor_message.payload
|
||||
else:
|
||||
cursor_message = None
|
||||
msg = sub.receive(timeout=1)
|
||||
except SubscriptionClosedError:
|
||||
return
|
||||
if msg is None:
|
||||
@@ -49,7 +101,10 @@ def stream_topic_events(
|
||||
last_msg_time = time.time()
|
||||
last_ping_time = last_msg_time
|
||||
event = json.loads(msg)
|
||||
yield event
|
||||
if cursor_message is not None and isinstance(event, Mapping):
|
||||
yield StreamEventWithCursor(event=event, cursor=cursor_message.cursor)
|
||||
else:
|
||||
yield event
|
||||
if not isinstance(event, dict):
|
||||
continue
|
||||
|
||||
|
||||
@@ -1,38 +1,99 @@
|
||||
"""In-process registry for workflow application task IDs."""
|
||||
"""In-process registry for workflow execution segments owned by this worker."""
|
||||
|
||||
import threading
|
||||
import time
|
||||
from collections.abc import Iterator
|
||||
from contextlib import contextmanager
|
||||
from dataclasses import dataclass
|
||||
from uuid import uuid4
|
||||
|
||||
_active_task_ids: set[str] = set()
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ActiveWorkflowTask:
|
||||
"""One workflow segment currently executing in this worker process.
|
||||
|
||||
``registration_id`` distinguishes a resumed segment that reuses the same
|
||||
task and run IDs from an older segment observed by a shutdown watchdog.
|
||||
"""
|
||||
|
||||
task_id: str
|
||||
workflow_run_id: str | None
|
||||
registration_id: str
|
||||
|
||||
|
||||
_active_tasks: dict[str, ActiveWorkflowTask] = {}
|
||||
_active_task_ids_lock = threading.RLock()
|
||||
_active_task_ids_changed = threading.Condition(_active_task_ids_lock)
|
||||
|
||||
|
||||
@contextmanager
|
||||
def active_workflow_task(task_id: str) -> Iterator[None]:
|
||||
"""Register a workflow application task ID for the duration of a workflow run."""
|
||||
def active_workflow_task(task_id: str, *, workflow_run_id: str | None = None) -> Iterator[None]:
|
||||
"""Register an execution segment for the duration of a workflow run."""
|
||||
if not task_id:
|
||||
raise ValueError("task_id must not be empty")
|
||||
if workflow_run_id is not None and not workflow_run_id:
|
||||
raise ValueError("workflow_run_id must not be empty")
|
||||
|
||||
with _active_task_ids_lock:
|
||||
if task_id in _active_task_ids:
|
||||
registration = ActiveWorkflowTask(
|
||||
task_id=task_id,
|
||||
workflow_run_id=workflow_run_id,
|
||||
registration_id=str(uuid4()),
|
||||
)
|
||||
|
||||
with _active_task_ids_changed:
|
||||
if task_id in _active_tasks:
|
||||
raise ValueError(f"Workflow task already active for task_id={task_id}")
|
||||
_active_task_ids.add(task_id)
|
||||
_active_tasks[task_id] = registration
|
||||
_active_task_ids_changed.notify_all()
|
||||
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
with _active_task_ids_lock:
|
||||
_active_task_ids.discard(task_id)
|
||||
with _active_task_ids_changed:
|
||||
if _active_tasks.get(task_id) == registration:
|
||||
del _active_tasks[task_id]
|
||||
_active_task_ids_changed.notify_all()
|
||||
|
||||
|
||||
def get_active_workflow_task_count() -> int:
|
||||
"""Return the number of active workflow application task IDs in this process."""
|
||||
with _active_task_ids_lock:
|
||||
return len(_active_task_ids)
|
||||
return len(_active_tasks)
|
||||
|
||||
|
||||
def get_active_workflow_tasks() -> tuple[ActiveWorkflowTask, ...]:
|
||||
"""Return a stable snapshot of workflow segments owned by this worker."""
|
||||
with _active_task_ids_lock:
|
||||
return tuple(_active_tasks.values())
|
||||
|
||||
|
||||
def retain_active_workflow_tasks(
|
||||
registrations: tuple[ActiveWorkflowTask, ...],
|
||||
) -> tuple[ActiveWorkflowTask, ...]:
|
||||
"""Keep only registrations that still refer to the same active segment."""
|
||||
with _active_task_ids_lock:
|
||||
return tuple(
|
||||
registration for registration in registrations if _active_tasks.get(registration.task_id) == registration
|
||||
)
|
||||
|
||||
|
||||
def wait_for_active_workflow_tasks(timeout: float) -> tuple[ActiveWorkflowTask, ...]:
|
||||
"""Wait for the registry to empty and return registrations left at timeout."""
|
||||
if timeout < 0:
|
||||
raise ValueError("timeout must be non-negative")
|
||||
|
||||
deadline = time.monotonic() + timeout
|
||||
with _active_task_ids_changed:
|
||||
while _active_tasks:
|
||||
remaining = deadline - time.monotonic()
|
||||
if remaining <= 0:
|
||||
return tuple(_active_tasks.values())
|
||||
_active_task_ids_changed.wait(timeout=remaining)
|
||||
return ()
|
||||
|
||||
|
||||
def reset_active_workflow_tasks() -> None:
|
||||
"""Clear active workflow application task IDs for worker initialization and tests."""
|
||||
with _active_task_ids_lock:
|
||||
_active_task_ids.clear()
|
||||
with _active_task_ids_changed:
|
||||
_active_tasks.clear()
|
||||
_active_task_ids_changed.notify_all()
|
||||
|
||||
@@ -30,6 +30,7 @@ from core.app.entities.task_entities import (
|
||||
WorkflowAppBlockingResponse,
|
||||
WorkflowAppPausedBlockingResponse,
|
||||
WorkflowAppStreamResponse,
|
||||
WorkflowMaintenancePausedBlockingResponse,
|
||||
)
|
||||
from core.app.layers.pause_state_persist_layer import PauseStateLayerConfig, PauseStatePersistenceLayer
|
||||
from core.db.session_factory import session_factory
|
||||
@@ -53,7 +54,12 @@ from models.account import Account
|
||||
from models.enums import WorkflowRunTriggeredFrom
|
||||
from models.model import App, EndUser
|
||||
from models.workflow import Workflow, WorkflowNodeExecutionTriggeredFrom
|
||||
from models.workflow_handoff import WorkflowHandoffResumeRoute
|
||||
from services.workflow_draft_variable_service import DraftVarLoader, WorkflowDraftVariableService
|
||||
from services.workflow_handoff_runtime_service import (
|
||||
build_workflow_handoff_persistence_layer,
|
||||
infer_initial_handoff_resume_route,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from controllers.console.app.workflow import LoopNodeRunPayload
|
||||
@@ -109,6 +115,7 @@ class WorkflowAppGenerator(BaseAppGenerator):
|
||||
root_node_id: str | None = None,
|
||||
graph_engine_layers: Sequence[GraphEngineLayer] = (),
|
||||
pause_state_config: PauseStateLayerConfig | None = None,
|
||||
handoff_resume_route: WorkflowHandoffResumeRoute | None = None,
|
||||
) -> Generator[Mapping[str, Any] | str, None, None]: ...
|
||||
|
||||
@overload
|
||||
@@ -127,6 +134,7 @@ class WorkflowAppGenerator(BaseAppGenerator):
|
||||
root_node_id: str | None = None,
|
||||
graph_engine_layers: Sequence[GraphEngineLayer] = (),
|
||||
pause_state_config: PauseStateLayerConfig | None = None,
|
||||
handoff_resume_route: WorkflowHandoffResumeRoute | None = None,
|
||||
) -> Mapping[str, Any]: ...
|
||||
|
||||
@overload
|
||||
@@ -145,6 +153,7 @@ class WorkflowAppGenerator(BaseAppGenerator):
|
||||
root_node_id: str | None = None,
|
||||
graph_engine_layers: Sequence[GraphEngineLayer] = (),
|
||||
pause_state_config: PauseStateLayerConfig | None = None,
|
||||
handoff_resume_route: WorkflowHandoffResumeRoute | None = None,
|
||||
) -> Mapping[str, Any] | Generator[Mapping[str, Any] | str, None, None]: ...
|
||||
|
||||
def generate(
|
||||
@@ -162,6 +171,7 @@ class WorkflowAppGenerator(BaseAppGenerator):
|
||||
root_node_id: str | None = None,
|
||||
graph_engine_layers: Sequence[GraphEngineLayer] = (),
|
||||
pause_state_config: PauseStateLayerConfig | None = None,
|
||||
handoff_resume_route: WorkflowHandoffResumeRoute | None = None,
|
||||
) -> Mapping[str, Any] | Generator[Mapping[str, Any] | str, None, None]:
|
||||
with self._bind_file_access_scope(tenant_id=app_model.tenant_id, user=user, invoke_from=invoke_from):
|
||||
files: Sequence[Mapping[str, Any]] = args.get("files") or []
|
||||
@@ -269,6 +279,18 @@ class WorkflowAppGenerator(BaseAppGenerator):
|
||||
root_node_id=root_node_id,
|
||||
graph_engine_layers=graph_engine_layers,
|
||||
pause_state_config=pause_state_config,
|
||||
handoff_resume_route=(
|
||||
handoff_resume_route
|
||||
or infer_initial_handoff_resume_route(
|
||||
application_generate_entity,
|
||||
triggered=workflow_triggered_from
|
||||
in {
|
||||
WorkflowRunTriggeredFrom.WEBHOOK,
|
||||
WorkflowRunTriggeredFrom.SCHEDULE,
|
||||
WorkflowRunTriggeredFrom.PLUGIN,
|
||||
},
|
||||
)
|
||||
),
|
||||
)
|
||||
|
||||
def resume(
|
||||
@@ -285,6 +307,10 @@ class WorkflowAppGenerator(BaseAppGenerator):
|
||||
pause_state_config: PauseStateLayerConfig | None = None,
|
||||
variable_loader: VariableLoader = DUMMY_VARIABLE_LOADER,
|
||||
response_stream_filter: ResponseStreamFilter | None = None,
|
||||
handoff_resume_route: WorkflowHandoffResumeRoute | None = None,
|
||||
graph_config: Mapping[str, Any] | None = None,
|
||||
workflow_version: str | None = None,
|
||||
root_node_id: str | None = None,
|
||||
) -> Mapping[str, Any] | Generator[str | Mapping[str, Any], None, None]:
|
||||
"""
|
||||
Resume a paused workflow execution using the persisted runtime state.
|
||||
@@ -316,6 +342,10 @@ class WorkflowAppGenerator(BaseAppGenerator):
|
||||
graph_runtime_state=graph_runtime_state,
|
||||
pause_state_config=pause_state_config,
|
||||
response_stream_filter=response_stream_filter,
|
||||
handoff_resume_route=handoff_resume_route,
|
||||
graph_config=graph_config,
|
||||
workflow_version=workflow_version,
|
||||
root_node_id=root_node_id,
|
||||
)
|
||||
|
||||
def _generate(
|
||||
@@ -335,6 +365,9 @@ class WorkflowAppGenerator(BaseAppGenerator):
|
||||
graph_runtime_state: GraphRuntimeState | None = None,
|
||||
pause_state_config: PauseStateLayerConfig | None = None,
|
||||
response_stream_filter: ResponseStreamFilter | None = None,
|
||||
handoff_resume_route: WorkflowHandoffResumeRoute | None = None,
|
||||
graph_config: Mapping[str, Any] | None = None,
|
||||
workflow_version: str | None = None,
|
||||
) -> Mapping[str, Any] | Generator[str | Mapping[str, Any], None, None]:
|
||||
"""
|
||||
Generate App response.
|
||||
@@ -364,6 +397,13 @@ class WorkflowAppGenerator(BaseAppGenerator):
|
||||
)
|
||||
|
||||
resolved_response_stream_filter = response_stream_filter or ResponseStreamFilter()
|
||||
handoff_layer = build_workflow_handoff_persistence_layer(
|
||||
generate_entity=application_generate_entity,
|
||||
response_stream_filter=resolved_response_stream_filter,
|
||||
resume_route=handoff_resume_route,
|
||||
)
|
||||
if handoff_layer is not None:
|
||||
graph_layers.append(handoff_layer)
|
||||
if pause_state_config is not None:
|
||||
graph_layers.append(
|
||||
PauseStatePersistenceLayer(
|
||||
@@ -394,6 +434,8 @@ class WorkflowAppGenerator(BaseAppGenerator):
|
||||
"graph_engine_layers": tuple(graph_layers),
|
||||
"graph_runtime_state": graph_runtime_state,
|
||||
"response_stream_filter": resolved_response_stream_filter,
|
||||
"graph_config": graph_config,
|
||||
"workflow_version": workflow_version,
|
||||
},
|
||||
)
|
||||
|
||||
@@ -438,6 +480,8 @@ class WorkflowAppGenerator(BaseAppGenerator):
|
||||
streaming: bool = True,
|
||||
*,
|
||||
session: Session,
|
||||
workflow_run_id: str | uuid.UUID | None = None,
|
||||
handoff_resume_route: WorkflowHandoffResumeRoute | None = None,
|
||||
) -> Mapping[str, Any] | Generator[str | Mapping[str, Any], None, None]:
|
||||
"""
|
||||
Generate App response.
|
||||
@@ -475,7 +519,7 @@ class WorkflowAppGenerator(BaseAppGenerator):
|
||||
single_iteration_run=WorkflowAppGenerateEntity.SingleIterationRunEntity(
|
||||
node_id=node_id, inputs=args["inputs"]
|
||||
),
|
||||
workflow_execution_id=str(uuid.uuid4()),
|
||||
workflow_execution_id=str(workflow_run_id or uuid.uuid4()),
|
||||
)
|
||||
contexts.plugin_tool_providers.set({})
|
||||
contexts.plugin_tool_providers_lock.set(threading.Lock())
|
||||
@@ -520,6 +564,7 @@ class WorkflowAppGenerator(BaseAppGenerator):
|
||||
streaming=streaming,
|
||||
variable_loader=var_loader,
|
||||
pause_state_config=None,
|
||||
handoff_resume_route=handoff_resume_route,
|
||||
)
|
||||
|
||||
def single_loop_generate(
|
||||
@@ -532,6 +577,8 @@ class WorkflowAppGenerator(BaseAppGenerator):
|
||||
streaming: bool = True,
|
||||
*,
|
||||
session: Session,
|
||||
workflow_run_id: str | uuid.UUID | None = None,
|
||||
handoff_resume_route: WorkflowHandoffResumeRoute | None = None,
|
||||
) -> Mapping[str, Any] | Generator[str | Mapping[str, Any], None, None]:
|
||||
"""
|
||||
Generate App response.
|
||||
@@ -567,7 +614,7 @@ class WorkflowAppGenerator(BaseAppGenerator):
|
||||
**_extract_trace_session_id_from_debug_args(args),
|
||||
},
|
||||
single_loop_run=WorkflowAppGenerateEntity.SingleLoopRunEntity(node_id=node_id, inputs=args.inputs or {}),
|
||||
workflow_execution_id=str(uuid.uuid4()),
|
||||
workflow_execution_id=str(workflow_run_id or uuid.uuid4()),
|
||||
)
|
||||
contexts.plugin_tool_providers.set({})
|
||||
contexts.plugin_tool_providers_lock.set(threading.Lock())
|
||||
@@ -611,6 +658,7 @@ class WorkflowAppGenerator(BaseAppGenerator):
|
||||
streaming=streaming,
|
||||
variable_loader=var_loader,
|
||||
pause_state_config=None,
|
||||
handoff_resume_route=handoff_resume_route,
|
||||
)
|
||||
|
||||
def _generate_worker(
|
||||
@@ -626,6 +674,8 @@ class WorkflowAppGenerator(BaseAppGenerator):
|
||||
graph_engine_layers: Sequence[GraphEngineLayer] = (),
|
||||
graph_runtime_state: GraphRuntimeState | None = None,
|
||||
response_stream_filter: ResponseStreamFilter | None = None,
|
||||
graph_config: Mapping[str, Any] | None = None,
|
||||
workflow_version: str | None = None,
|
||||
) -> None:
|
||||
"""
|
||||
Generate worker in a new thread.
|
||||
@@ -675,10 +725,15 @@ class WorkflowAppGenerator(BaseAppGenerator):
|
||||
graph_engine_layers=graph_engine_layers,
|
||||
graph_runtime_state=graph_runtime_state,
|
||||
response_stream_filter=response_stream_filter,
|
||||
graph_config=graph_config,
|
||||
workflow_version=workflow_version,
|
||||
)
|
||||
|
||||
try:
|
||||
with active_workflow_task(application_generate_entity.task_id):
|
||||
with active_workflow_task(
|
||||
application_generate_entity.task_id,
|
||||
workflow_run_id=application_generate_entity.workflow_execution_id,
|
||||
):
|
||||
runner.run()
|
||||
except GenerateTaskStoppedError as e:
|
||||
logger.warning("Task stopped: %s", str(e))
|
||||
@@ -709,6 +764,7 @@ class WorkflowAppGenerator(BaseAppGenerator):
|
||||
) -> (
|
||||
WorkflowAppBlockingResponse
|
||||
| WorkflowAppPausedBlockingResponse
|
||||
| WorkflowMaintenancePausedBlockingResponse
|
||||
| Generator[WorkflowAppStreamResponse, None, None]
|
||||
):
|
||||
"""
|
||||
|
||||
@@ -8,7 +8,9 @@ from core.app.entities.queue_entities import (
|
||||
QueueMessageEndEvent,
|
||||
QueueStopEvent,
|
||||
QueueWorkflowFailedEvent,
|
||||
QueueWorkflowMaintenancePausedEvent,
|
||||
QueueWorkflowPartialSuccessEvent,
|
||||
QueueWorkflowPausedEvent,
|
||||
QueueWorkflowSucceededEvent,
|
||||
WorkflowQueueMessage,
|
||||
)
|
||||
@@ -39,6 +41,10 @@ class WorkflowAppQueueManager(AppQueueManager):
|
||||
| QueueMessageEndEvent
|
||||
| QueueWorkflowSucceededEvent
|
||||
| QueueWorkflowFailedEvent
|
||||
| QueueWorkflowPartialSuccessEvent,
|
||||
| QueueWorkflowPartialSuccessEvent
|
||||
| QueueWorkflowPausedEvent,
|
||||
):
|
||||
self.stop_listen(execution_terminal=True)
|
||||
elif isinstance(event, QueueWorkflowMaintenancePausedEvent):
|
||||
# The response pipeline must first cross the durable flush barrier.
|
||||
self.stop_listen(execution_terminal=False)
|
||||
|
||||
@@ -1,8 +1,9 @@
|
||||
import logging
|
||||
import time
|
||||
from collections.abc import Sequence
|
||||
from typing import cast
|
||||
from collections.abc import Mapping, Sequence
|
||||
from typing import Any, cast
|
||||
|
||||
from configs import dify_config
|
||||
from core.app.apps.base_app_queue_manager import AppQueueManager
|
||||
from core.app.apps.workflow.app_config_manager import WorkflowAppConfig
|
||||
from core.app.apps.workflow.command_channels import (
|
||||
@@ -11,6 +12,7 @@ from core.app.apps.workflow.command_channels import (
|
||||
)
|
||||
from core.app.apps.workflow_app_runner import WorkflowBasedAppRunner
|
||||
from core.app.entities.app_invoke_entities import DifyRunContext, InvokeFrom, WorkflowAppGenerateEntity
|
||||
from core.app.layers.pause_state_persist_layer import get_workflow_handoff_active_execution_seconds
|
||||
from core.app.workflow.layers.persistence import PersistenceWorkflowInfo, WorkflowPersistenceLayer
|
||||
from core.repositories.factory import WorkflowExecutionRepository, WorkflowNodeExecutionRepository
|
||||
from core.workflow.node_factory import get_default_root_node_id
|
||||
@@ -21,7 +23,7 @@ from core.workflow.variable_pool_initializer import add_node_inputs_to_pool, add
|
||||
from core.workflow.workflow_entry import WorkflowEntry
|
||||
from extensions.ext_redis import redis_client
|
||||
from extensions.otel import WorkflowAppRunnerHandler, trace_span
|
||||
from extensions.workflow_warm_shutdown import WORKFLOW_WARM_SHUTDOWN_ABORT_REASON, celery_warm_shutdown_started
|
||||
from extensions.workflow_warm_shutdown import celery_warm_shutdown_started
|
||||
from graphon.enums import WorkflowType
|
||||
from graphon.filters import ResponseStreamFilter
|
||||
from graphon.graph_engine.command_channels import RedisChannel
|
||||
@@ -53,6 +55,8 @@ class WorkflowAppRunner(WorkflowBasedAppRunner):
|
||||
graph_engine_layers: Sequence[GraphEngineLayer] = (),
|
||||
graph_runtime_state: GraphRuntimeState | None = None,
|
||||
response_stream_filter: ResponseStreamFilter | None = None,
|
||||
graph_config: Mapping[str, Any] | None = None,
|
||||
workflow_version: str | None = None,
|
||||
):
|
||||
super().__init__(
|
||||
queue_manager=queue_manager,
|
||||
@@ -68,6 +72,8 @@ class WorkflowAppRunner(WorkflowBasedAppRunner):
|
||||
self._workflow_node_execution_repository = workflow_node_execution_repository
|
||||
self._resume_graph_runtime_state = graph_runtime_state
|
||||
self._response_stream_filter = response_stream_filter
|
||||
self._execution_graph_config = graph_config
|
||||
self._workflow_version = workflow_version
|
||||
|
||||
@trace_span(WorkflowAppRunnerHandler)
|
||||
def run(self):
|
||||
@@ -81,6 +87,10 @@ class WorkflowAppRunner(WorkflowBasedAppRunner):
|
||||
if self.application_generate_entity.single_iteration_run or self.application_generate_entity.single_loop_run:
|
||||
invoke_from = InvokeFrom.DEBUGGER
|
||||
user_from = self._resolve_user_from(invoke_from)
|
||||
graph_config = (
|
||||
self._execution_graph_config if self._execution_graph_config is not None else self._workflow.graph_dict
|
||||
)
|
||||
workflow_version = self._workflow_version if self._workflow_version is not None else self._workflow.version
|
||||
|
||||
resume_state = self._resume_graph_runtime_state
|
||||
|
||||
@@ -88,7 +98,7 @@ class WorkflowAppRunner(WorkflowBasedAppRunner):
|
||||
graph_runtime_state = resume_state
|
||||
variable_pool = graph_runtime_state.variable_pool
|
||||
graph = self._init_graph(
|
||||
graph_config=self._workflow.graph_dict,
|
||||
graph_config=graph_config,
|
||||
graph_runtime_state=graph_runtime_state,
|
||||
workflow_id=self._workflow.id,
|
||||
tenant_id=self._workflow.tenant_id,
|
||||
@@ -97,13 +107,18 @@ class WorkflowAppRunner(WorkflowBasedAppRunner):
|
||||
invoke_from=invoke_from,
|
||||
root_node_id=self._root_node_id,
|
||||
trace_session_id=self.application_generate_entity.extras.get("trace_session_id"),
|
||||
skip_validation=bool(
|
||||
self.application_generate_entity.single_iteration_run
|
||||
or self.application_generate_entity.single_loop_run
|
||||
),
|
||||
)
|
||||
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, graph_config, variable_pool, graph_runtime_state = self._prepare_single_node_execution(
|
||||
workflow=self._workflow,
|
||||
single_iteration_run=self.application_generate_entity.single_iteration_run,
|
||||
single_loop_run=self.application_generate_entity.single_loop_run,
|
||||
user_id=self.application_generate_entity.user_id,
|
||||
workflow_execution_id=self.application_generate_entity.workflow_execution_id,
|
||||
trace_session_id=self.application_generate_entity.extras.get("trace_session_id"),
|
||||
)
|
||||
else:
|
||||
@@ -126,7 +141,7 @@ class WorkflowAppRunner(WorkflowBasedAppRunner):
|
||||
environment_variables=self._workflow.environment_variables,
|
||||
),
|
||||
)
|
||||
root_node_id = self._root_node_id or get_default_root_node_id(self._workflow.graph_dict)
|
||||
root_node_id = self._root_node_id or get_default_root_node_id(graph_config)
|
||||
add_node_inputs_to_pool(
|
||||
variable_pool,
|
||||
node_id=root_node_id,
|
||||
@@ -139,7 +154,7 @@ class WorkflowAppRunner(WorkflowBasedAppRunner):
|
||||
|
||||
graph_runtime_state = GraphRuntimeState(variable_pool=variable_pool, start_at=time.perf_counter())
|
||||
graph = self._init_graph(
|
||||
graph_config=self._workflow.graph_dict,
|
||||
graph_config=graph_config,
|
||||
graph_runtime_state=graph_runtime_state,
|
||||
workflow_id=self._workflow.id,
|
||||
tenant_id=self._workflow.tenant_id,
|
||||
@@ -151,12 +166,16 @@ class WorkflowAppRunner(WorkflowBasedAppRunner):
|
||||
)
|
||||
|
||||
# RUN WORKFLOW
|
||||
self._configure_execution_root(graph.root_node.id)
|
||||
# Create Redis command channel for this workflow execution
|
||||
task_id = self.application_generate_entity.task_id
|
||||
channel_key = f"workflow:{task_id}:commands"
|
||||
celery_signal_channel = CelerySignalCommandChannel(
|
||||
shutdown_state_getter=celery_warm_shutdown_started,
|
||||
abort_reason=WORKFLOW_WARM_SHUTDOWN_ABORT_REASON,
|
||||
pause_on_shutdown=(
|
||||
dify_config.WORKFLOW_HANDOFF_ENABLED and self.application_generate_entity.call_depth == 0
|
||||
),
|
||||
ignore_shutdown=(dify_config.WORKFLOW_HANDOFF_ENABLED and self.application_generate_entity.call_depth > 0),
|
||||
)
|
||||
command_channel = CombinedCommandChannel(
|
||||
(
|
||||
@@ -172,7 +191,7 @@ class WorkflowAppRunner(WorkflowBasedAppRunner):
|
||||
app_id=self._workflow.app_id,
|
||||
workflow_id=self._workflow.id,
|
||||
graph=graph,
|
||||
graph_config=self._workflow.graph_dict,
|
||||
graph_config=graph_config,
|
||||
user_id=self.application_generate_entity.user_id,
|
||||
user_from=user_from,
|
||||
invoke_from=invoke_from,
|
||||
@@ -181,6 +200,9 @@ class WorkflowAppRunner(WorkflowBasedAppRunner):
|
||||
graph_runtime_state=graph_runtime_state,
|
||||
command_channel=command_channel,
|
||||
response_stream_filter=self._response_stream_filter,
|
||||
prior_active_execution_seconds=get_workflow_handoff_active_execution_seconds(
|
||||
self.application_generate_entity
|
||||
),
|
||||
)
|
||||
|
||||
persistence_layer = WorkflowPersistenceLayer(
|
||||
@@ -188,8 +210,8 @@ class WorkflowAppRunner(WorkflowBasedAppRunner):
|
||||
workflow_info=PersistenceWorkflowInfo(
|
||||
workflow_id=self._workflow.id,
|
||||
workflow_type=WorkflowType(self._workflow.type),
|
||||
version=self._workflow.version,
|
||||
graph_data=self._workflow.graph_dict,
|
||||
version=workflow_version,
|
||||
graph_data=graph_config,
|
||||
),
|
||||
workflow_execution_repository=self._workflow_execution_repository,
|
||||
workflow_node_execution_repository=self._workflow_node_execution_repository,
|
||||
@@ -215,4 +237,4 @@ class WorkflowAppRunner(WorkflowBasedAppRunner):
|
||||
generator = workflow_entry.run()
|
||||
|
||||
for event in generator:
|
||||
self._handle_event(workflow_entry, event)
|
||||
self._handle_event_with_handoff_contracts(workflow_entry, event)
|
||||
|
||||
@@ -4,12 +4,29 @@ import logging
|
||||
from collections.abc import Callable, Sequence
|
||||
from typing import final, override
|
||||
|
||||
from graphon.entities.pause_reason import PauseReason, SchedulingPause
|
||||
from graphon.graph_engine.command_channels import CommandChannel
|
||||
from graphon.graph_engine.entities.commands import AbortCommand, GraphEngineCommand
|
||||
from graphon.graph_engine.entities.commands import AbortCommand, GraphEngineCommand, PauseCommand
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
ShutdownStateGetter = Callable[[], bool]
|
||||
|
||||
WORKFLOW_WARM_SHUTDOWN_ABORT_REASON = "Workflow stopped because the worker is shutting down."
|
||||
WORKFLOW_WARM_SHUTDOWN_PAUSE_REASON = "Workflow paused because the worker is shutting down."
|
||||
|
||||
|
||||
def is_workflow_warm_shutdown_pause(reasons: Sequence[PauseReason]) -> bool:
|
||||
"""Return whether a graph pause was requested by Celery worker drain.
|
||||
|
||||
Graphon 0.6 represents a scheduling pause with only a free-form message.
|
||||
Keep the comparison behind this adapter so callers do not duplicate that
|
||||
compatibility detail while Dify moves toward a structured pause origin.
|
||||
"""
|
||||
return bool(reasons) and all(
|
||||
isinstance(reason, SchedulingPause) and reason.message == WORKFLOW_WARM_SHUTDOWN_PAUSE_REASON
|
||||
for reason in reasons
|
||||
)
|
||||
|
||||
|
||||
@final
|
||||
class CombinedCommandChannel:
|
||||
@@ -29,7 +46,27 @@ class CombinedCommandChannel:
|
||||
commands.extend(channel.fetch_commands())
|
||||
except Exception:
|
||||
logger.exception("Failed to fetch GraphEngine commands from %s", channel.__class__.__name__)
|
||||
return commands
|
||||
|
||||
abort_commands = [command for command in commands if isinstance(command, AbortCommand)]
|
||||
if not abort_commands:
|
||||
return commands
|
||||
|
||||
# Abort is terminal and must win when a user Stop races a worker-drain
|
||||
# pause. Also keep a user Abort reason ahead of the shutdown fallback
|
||||
# when handoff is disabled.
|
||||
has_non_shutdown_abort = any(
|
||||
command.reason != WORKFLOW_WARM_SHUTDOWN_ABORT_REASON for command in abort_commands
|
||||
)
|
||||
return [
|
||||
command
|
||||
for command in commands
|
||||
if not isinstance(command, PauseCommand)
|
||||
and not (
|
||||
has_non_shutdown_abort
|
||||
and isinstance(command, AbortCommand)
|
||||
and command.reason == WORKFLOW_WARM_SHUTDOWN_ABORT_REASON
|
||||
)
|
||||
]
|
||||
|
||||
def send_command(self, command: GraphEngineCommand) -> None:
|
||||
"""Send commands through the first channel, which is the runner's primary command sink."""
|
||||
@@ -38,29 +75,41 @@ class CombinedCommandChannel:
|
||||
|
||||
@final
|
||||
class CelerySignalCommandChannel(CommandChannel):
|
||||
"""Translate process-local Celery shutdown state into one GraphEngine abort command."""
|
||||
"""Translate process-local Celery shutdown state into one GraphEngine control command."""
|
||||
|
||||
_shutdown_state_getter: ShutdownStateGetter
|
||||
_abort_reason: str
|
||||
_abort_emitted: bool
|
||||
_pause_on_shutdown: bool
|
||||
_ignore_shutdown: bool
|
||||
_command_emitted: bool
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
shutdown_state_getter: ShutdownStateGetter,
|
||||
abort_reason: str,
|
||||
pause_on_shutdown: bool,
|
||||
ignore_shutdown: bool = False,
|
||||
) -> None:
|
||||
if pause_on_shutdown and ignore_shutdown:
|
||||
raise ValueError("pause_on_shutdown and ignore_shutdown are mutually exclusive")
|
||||
self._shutdown_state_getter = shutdown_state_getter
|
||||
self._abort_reason = abort_reason
|
||||
self._abort_emitted = False
|
||||
self._pause_on_shutdown = pause_on_shutdown
|
||||
self._ignore_shutdown = ignore_shutdown
|
||||
self._command_emitted = False
|
||||
|
||||
@override
|
||||
def fetch_commands(self) -> list[GraphEngineCommand]:
|
||||
if self._abort_emitted or not self._shutdown_state_getter():
|
||||
if self._command_emitted or not self._shutdown_state_getter():
|
||||
return []
|
||||
|
||||
self._abort_emitted = True
|
||||
return [AbortCommand(reason=self._abort_reason)]
|
||||
self._command_emitted = True
|
||||
if self._ignore_shutdown:
|
||||
# Nested Workflow-as-Tool execution must drain back into its parent
|
||||
# node. The top-level graph then checkpoints the complete result at
|
||||
# the next safe scheduling boundary.
|
||||
return []
|
||||
if self._pause_on_shutdown:
|
||||
return [PauseCommand(reason=WORKFLOW_WARM_SHUTDOWN_PAUSE_REASON)]
|
||||
return [AbortCommand(reason=WORKFLOW_WARM_SHUTDOWN_ABORT_REASON)]
|
||||
|
||||
@override
|
||||
def send_command(self, command: GraphEngineCommand) -> None:
|
||||
|
||||
@@ -2,6 +2,7 @@ from collections.abc import Generator
|
||||
from typing import Any, cast, override
|
||||
|
||||
from core.app.apps.base_app_generate_response_converter import AppGenerateResponseConverter
|
||||
from core.app.apps.streaming_utils import close_stream
|
||||
from core.app.entities.task_entities import (
|
||||
AppStreamResponse,
|
||||
ErrorStreamResponse,
|
||||
@@ -11,16 +12,22 @@ from core.app.entities.task_entities import (
|
||||
WorkflowAppBlockingResponse,
|
||||
WorkflowAppPausedBlockingResponse,
|
||||
WorkflowAppStreamResponse,
|
||||
WorkflowMaintenancePausedBlockingResponse,
|
||||
)
|
||||
|
||||
|
||||
class WorkflowAppGenerateResponseConverter(
|
||||
AppGenerateResponseConverter[WorkflowAppBlockingResponse | WorkflowAppPausedBlockingResponse]
|
||||
AppGenerateResponseConverter[
|
||||
WorkflowAppBlockingResponse | WorkflowAppPausedBlockingResponse | WorkflowMaintenancePausedBlockingResponse
|
||||
]
|
||||
):
|
||||
@classmethod
|
||||
@override
|
||||
def convert_blocking_full_response(
|
||||
cls, blocking_response: WorkflowAppBlockingResponse | WorkflowAppPausedBlockingResponse
|
||||
cls,
|
||||
blocking_response: (
|
||||
WorkflowAppBlockingResponse | WorkflowAppPausedBlockingResponse | WorkflowMaintenancePausedBlockingResponse
|
||||
),
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
Convert blocking full response.
|
||||
@@ -32,7 +39,10 @@ class WorkflowAppGenerateResponseConverter(
|
||||
@classmethod
|
||||
@override
|
||||
def convert_blocking_simple_response(
|
||||
cls, blocking_response: WorkflowAppBlockingResponse | WorkflowAppPausedBlockingResponse
|
||||
cls,
|
||||
blocking_response: (
|
||||
WorkflowAppBlockingResponse | WorkflowAppPausedBlockingResponse | WorkflowMaintenancePausedBlockingResponse
|
||||
),
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
Convert blocking simple response.
|
||||
@@ -51,29 +61,32 @@ class WorkflowAppGenerateResponseConverter(
|
||||
:param stream_response: stream response
|
||||
:return:
|
||||
"""
|
||||
for chunk in stream_response:
|
||||
chunk = cast(WorkflowAppStreamResponse, chunk)
|
||||
sub_stream_response = chunk.stream_response
|
||||
try:
|
||||
for chunk in stream_response:
|
||||
chunk = cast(WorkflowAppStreamResponse, chunk)
|
||||
sub_stream_response = chunk.stream_response
|
||||
|
||||
match sub_stream_response:
|
||||
case PingStreamResponse():
|
||||
yield "ping"
|
||||
continue
|
||||
case ErrorStreamResponse():
|
||||
response_chunk: dict[str, object] = {
|
||||
"event": sub_stream_response.event.value,
|
||||
"workflow_run_id": chunk.workflow_run_id,
|
||||
}
|
||||
data = cls._error_to_stream_response(sub_stream_response.err)
|
||||
response_chunk.update(data)
|
||||
case _:
|
||||
response_chunk = {
|
||||
"event": sub_stream_response.event.value,
|
||||
"workflow_run_id": chunk.workflow_run_id,
|
||||
}
|
||||
response_chunk.update(sub_stream_response.model_dump(mode="json"))
|
||||
match sub_stream_response:
|
||||
case PingStreamResponse():
|
||||
yield "ping"
|
||||
continue
|
||||
case ErrorStreamResponse():
|
||||
response_chunk: dict[str, object] = {
|
||||
"event": sub_stream_response.event.value,
|
||||
"workflow_run_id": chunk.workflow_run_id,
|
||||
}
|
||||
data = cls._error_to_stream_response(sub_stream_response.err)
|
||||
response_chunk.update(data)
|
||||
case _:
|
||||
response_chunk = {
|
||||
"event": sub_stream_response.event.value,
|
||||
"workflow_run_id": chunk.workflow_run_id,
|
||||
}
|
||||
response_chunk.update(sub_stream_response.model_dump(mode="json"))
|
||||
|
||||
yield response_chunk
|
||||
yield response_chunk
|
||||
finally:
|
||||
close_stream(stream_response)
|
||||
|
||||
@classmethod
|
||||
@override
|
||||
@@ -85,32 +98,35 @@ class WorkflowAppGenerateResponseConverter(
|
||||
:param stream_response: stream response
|
||||
:return:
|
||||
"""
|
||||
for chunk in stream_response:
|
||||
chunk = cast(WorkflowAppStreamResponse, chunk)
|
||||
sub_stream_response = chunk.stream_response
|
||||
try:
|
||||
for chunk in stream_response:
|
||||
chunk = cast(WorkflowAppStreamResponse, chunk)
|
||||
sub_stream_response = chunk.stream_response
|
||||
|
||||
match sub_stream_response:
|
||||
case PingStreamResponse():
|
||||
yield "ping"
|
||||
continue
|
||||
case ErrorStreamResponse():
|
||||
response_chunk: dict[str, object] = {
|
||||
"event": sub_stream_response.event.value,
|
||||
"workflow_run_id": chunk.workflow_run_id,
|
||||
}
|
||||
data = cls._error_to_stream_response(sub_stream_response.err)
|
||||
response_chunk.update(data)
|
||||
case NodeStartStreamResponse() | NodeFinishStreamResponse():
|
||||
response_chunk = {
|
||||
"event": sub_stream_response.event.value,
|
||||
"workflow_run_id": chunk.workflow_run_id,
|
||||
}
|
||||
response_chunk.update(sub_stream_response.to_ignore_detail_dict())
|
||||
case _:
|
||||
response_chunk = {
|
||||
"event": sub_stream_response.event.value,
|
||||
"workflow_run_id": chunk.workflow_run_id,
|
||||
}
|
||||
response_chunk.update(sub_stream_response.model_dump(mode="json"))
|
||||
match sub_stream_response:
|
||||
case PingStreamResponse():
|
||||
yield "ping"
|
||||
continue
|
||||
case ErrorStreamResponse():
|
||||
response_chunk: dict[str, object] = {
|
||||
"event": sub_stream_response.event.value,
|
||||
"workflow_run_id": chunk.workflow_run_id,
|
||||
}
|
||||
data = cls._error_to_stream_response(sub_stream_response.err)
|
||||
response_chunk.update(data)
|
||||
case NodeStartStreamResponse() | NodeFinishStreamResponse():
|
||||
response_chunk = {
|
||||
"event": sub_stream_response.event.value,
|
||||
"workflow_run_id": chunk.workflow_run_id,
|
||||
}
|
||||
response_chunk.update(sub_stream_response.to_ignore_detail_dict())
|
||||
case _:
|
||||
response_chunk = {
|
||||
"event": sub_stream_response.event.value,
|
||||
"workflow_run_id": chunk.workflow_run_id,
|
||||
}
|
||||
response_chunk.update(sub_stream_response.model_dump(mode="json"))
|
||||
|
||||
yield response_chunk
|
||||
yield response_chunk
|
||||
finally:
|
||||
close_stream(stream_response)
|
||||
|
||||
@@ -11,6 +11,7 @@ from core.app.apps.base_app_queue_manager import AppQueueManager
|
||||
from core.app.apps.common.graph_runtime_state_support import GraphRuntimeStateSupport
|
||||
from core.app.apps.common.workflow_response_converter import WorkflowResponseConverter
|
||||
from core.app.apps.draft_variable_saver import DraftVariableSaverFactory
|
||||
from core.app.apps.streaming_utils import close_stream
|
||||
from core.app.entities.app_invoke_entities import InvokeFrom, WorkflowAppGenerateEntity
|
||||
from core.app.entities.queue_entities import (
|
||||
AppQueueEvent,
|
||||
@@ -35,6 +36,7 @@ from core.app.entities.queue_entities import (
|
||||
QueueStopEvent,
|
||||
QueueTextChunkEvent,
|
||||
QueueWorkflowFailedEvent,
|
||||
QueueWorkflowMaintenancePausedEvent,
|
||||
QueueWorkflowPartialSuccessEvent,
|
||||
QueueWorkflowPausedEvent,
|
||||
QueueWorkflowStartedEvent,
|
||||
@@ -55,6 +57,8 @@ from core.app.entities.task_entities import (
|
||||
WorkflowAppPausedBlockingResponse,
|
||||
WorkflowAppStreamResponse,
|
||||
WorkflowFinishStreamResponse,
|
||||
WorkflowMaintenancePausedBlockingResponse,
|
||||
WorkflowMaintenancePausedStreamResponse,
|
||||
WorkflowPauseStreamResponse,
|
||||
WorkflowStartStreamResponse,
|
||||
)
|
||||
@@ -70,6 +74,8 @@ from models import Account
|
||||
from models.enums import CreatorUserRole
|
||||
from models.model import EndUser
|
||||
from models.workflow import Workflow, WorkflowAppLog, WorkflowAppLogCreatedFrom
|
||||
from services.workflow_handoff_activation_service import activate_workflow_handoff_by_task_id
|
||||
from services.workflow_run_timing_service import get_workflow_run_public_timing
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -126,7 +132,10 @@ class WorkflowAppGenerateTaskPipeline(GraphRuntimeStateSupport):
|
||||
def process(
|
||||
self,
|
||||
) -> Union[
|
||||
WorkflowAppBlockingResponse, WorkflowAppPausedBlockingResponse, Generator[WorkflowAppStreamResponse, None, None]
|
||||
WorkflowAppBlockingResponse,
|
||||
WorkflowAppPausedBlockingResponse,
|
||||
WorkflowMaintenancePausedBlockingResponse,
|
||||
Generator[WorkflowAppStreamResponse, None, None],
|
||||
]:
|
||||
"""
|
||||
Process generate task pipeline.
|
||||
@@ -140,58 +149,72 @@ class WorkflowAppGenerateTaskPipeline(GraphRuntimeStateSupport):
|
||||
|
||||
def _to_blocking_response(
|
||||
self, generator: Generator[StreamResponse, None, None]
|
||||
) -> Union[WorkflowAppBlockingResponse, WorkflowAppPausedBlockingResponse]:
|
||||
) -> Union[
|
||||
WorkflowAppBlockingResponse,
|
||||
WorkflowAppPausedBlockingResponse,
|
||||
WorkflowMaintenancePausedBlockingResponse,
|
||||
]:
|
||||
"""
|
||||
To blocking response.
|
||||
:return:
|
||||
"""
|
||||
human_input_responses: list[HumanInputRequiredResponse] = []
|
||||
for stream_response in generator:
|
||||
match stream_response:
|
||||
case ErrorStreamResponse():
|
||||
raise stream_response.err
|
||||
case HumanInputRequiredResponse():
|
||||
human_input_responses.append(stream_response)
|
||||
case WorkflowPauseStreamResponse():
|
||||
return WorkflowAppPausedBlockingResponse(
|
||||
task_id=self._application_generate_entity.task_id,
|
||||
workflow_run_id=stream_response.data.workflow_run_id,
|
||||
data=WorkflowAppPausedBlockingResponse.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,
|
||||
paused_nodes=stream_response.data.paused_nodes,
|
||||
reasons=stream_response.data.reasons,
|
||||
),
|
||||
)
|
||||
case WorkflowFinishStreamResponse():
|
||||
return WorkflowAppBlockingResponse(
|
||||
task_id=self._application_generate_entity.task_id,
|
||||
workflow_run_id=stream_response.data.id,
|
||||
data=WorkflowAppBlockingResponse.Data(
|
||||
id=stream_response.data.id,
|
||||
workflow_id=stream_response.data.workflow_id,
|
||||
status=stream_response.data.status,
|
||||
outputs=stream_response.data.outputs,
|
||||
error=stream_response.data.error,
|
||||
elapsed_time=stream_response.data.elapsed_time,
|
||||
total_tokens=stream_response.data.total_tokens,
|
||||
total_steps=stream_response.data.total_steps,
|
||||
created_at=int(stream_response.data.created_at),
|
||||
finished_at=int(stream_response.data.finished_at)
|
||||
if stream_response.data.finished_at
|
||||
else None,
|
||||
),
|
||||
)
|
||||
case _:
|
||||
continue
|
||||
try:
|
||||
for stream_response in generator:
|
||||
match stream_response:
|
||||
case ErrorStreamResponse():
|
||||
raise stream_response.err
|
||||
case HumanInputRequiredResponse():
|
||||
human_input_responses.append(stream_response)
|
||||
case WorkflowPauseStreamResponse():
|
||||
return WorkflowAppPausedBlockingResponse(
|
||||
task_id=self._application_generate_entity.task_id,
|
||||
workflow_run_id=stream_response.data.workflow_run_id,
|
||||
data=WorkflowAppPausedBlockingResponse.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,
|
||||
paused_nodes=stream_response.data.paused_nodes,
|
||||
reasons=stream_response.data.reasons,
|
||||
handoff_duration=stream_response.data.handoff_duration,
|
||||
),
|
||||
)
|
||||
case WorkflowMaintenancePausedStreamResponse():
|
||||
return WorkflowMaintenancePausedBlockingResponse(
|
||||
task_id=stream_response.task_id,
|
||||
workflow_run_id=stream_response.workflow_run_id,
|
||||
)
|
||||
case WorkflowFinishStreamResponse():
|
||||
return WorkflowAppBlockingResponse(
|
||||
task_id=self._application_generate_entity.task_id,
|
||||
workflow_run_id=stream_response.data.id,
|
||||
data=WorkflowAppBlockingResponse.Data(
|
||||
id=stream_response.data.id,
|
||||
workflow_id=stream_response.data.workflow_id,
|
||||
status=stream_response.data.status,
|
||||
outputs=stream_response.data.outputs,
|
||||
error=stream_response.data.error,
|
||||
elapsed_time=stream_response.data.elapsed_time,
|
||||
total_tokens=stream_response.data.total_tokens,
|
||||
total_steps=stream_response.data.total_steps,
|
||||
created_at=int(stream_response.data.created_at),
|
||||
finished_at=int(stream_response.data.finished_at)
|
||||
if stream_response.data.finished_at
|
||||
else None,
|
||||
handoff_duration=stream_response.data.handoff_duration,
|
||||
),
|
||||
)
|
||||
case _:
|
||||
continue
|
||||
finally:
|
||||
close_stream(generator)
|
||||
|
||||
if human_input_responses:
|
||||
return self._build_paused_blocking_response_from_human_input(human_input_responses)
|
||||
@@ -236,11 +259,14 @@ class WorkflowAppGenerateTaskPipeline(GraphRuntimeStateSupport):
|
||||
:return:
|
||||
"""
|
||||
workflow_run_id = None
|
||||
for stream_response in generator:
|
||||
if isinstance(stream_response, WorkflowStartStreamResponse):
|
||||
workflow_run_id = stream_response.workflow_run_id
|
||||
try:
|
||||
for stream_response in generator:
|
||||
if isinstance(stream_response, WorkflowStartStreamResponse):
|
||||
workflow_run_id = stream_response.workflow_run_id
|
||||
|
||||
yield WorkflowAppStreamResponse(workflow_run_id=workflow_run_id, stream_response=stream_response)
|
||||
yield WorkflowAppStreamResponse(workflow_run_id=workflow_run_id, stream_response=stream_response)
|
||||
finally:
|
||||
close_stream(generator)
|
||||
|
||||
def _listen_audio_msg(self, publisher: AppGeneratorTTSPublisher | None, task_id: str):
|
||||
if not publisher:
|
||||
@@ -267,14 +293,18 @@ class WorkflowAppGenerateTaskPipeline(GraphRuntimeStateSupport):
|
||||
tenant_id, features_dict["text_to_speech"].get("voice"), features_dict["text_to_speech"].get("language")
|
||||
)
|
||||
|
||||
for response in self._process_stream_response(tts_publisher=tts_publisher, trace_manager=trace_manager):
|
||||
while True:
|
||||
audio_response = self._listen_audio_msg(publisher=tts_publisher, task_id=task_id)
|
||||
if audio_response:
|
||||
yield audio_response
|
||||
else:
|
||||
break
|
||||
yield response
|
||||
response_stream = self._process_stream_response(tts_publisher=tts_publisher, trace_manager=trace_manager)
|
||||
try:
|
||||
for response in response_stream:
|
||||
while True:
|
||||
audio_response = self._listen_audio_msg(publisher=tts_publisher, task_id=task_id)
|
||||
if audio_response:
|
||||
yield audio_response
|
||||
else:
|
||||
break
|
||||
yield response
|
||||
finally:
|
||||
close_stream(response_stream)
|
||||
|
||||
start_listener_time = time.time()
|
||||
while (time.time() - start_listener_time) < TTS_AUTO_PLAY_TIMEOUT:
|
||||
@@ -330,11 +360,24 @@ class WorkflowAppGenerateTaskPipeline(GraphRuntimeStateSupport):
|
||||
with self._database_session() as session:
|
||||
self._save_workflow_app_log(session=session, workflow_run_id=self._workflow_execution_id)
|
||||
|
||||
logical_timing = None
|
||||
if event.reason == WorkflowStartReason.RESUMPTION:
|
||||
with self._database_session() as session:
|
||||
logical_timing = get_workflow_run_public_timing(
|
||||
session=session,
|
||||
workflow_run_id=run_id,
|
||||
tenant_id=self._application_generate_entity.app_config.tenant_id,
|
||||
app_id=self._application_generate_entity.app_config.app_id,
|
||||
workflow_id=self._workflow.id,
|
||||
)
|
||||
|
||||
start_resp = self._workflow_response_converter.workflow_start_to_stream_response(
|
||||
task_id=self._application_generate_entity.task_id,
|
||||
workflow_run_id=run_id,
|
||||
workflow_id=self._workflow.id,
|
||||
reason=event.reason,
|
||||
logical_started_at=logical_timing.started_at if logical_timing is not None else None,
|
||||
handoff_duration=logical_timing.handoff_duration if logical_timing is not None else 0.0,
|
||||
)
|
||||
yield start_resp
|
||||
|
||||
@@ -524,6 +567,21 @@ class WorkflowAppGenerateTaskPipeline(GraphRuntimeStateSupport):
|
||||
)
|
||||
yield from responses
|
||||
|
||||
def _handle_workflow_maintenance_paused_event(
|
||||
self,
|
||||
event: QueueWorkflowMaintenancePausedEvent,
|
||||
**kwargs,
|
||||
) -> Generator[StreamResponse, None, None]:
|
||||
"""End this worker's segment without exposing a public pause event."""
|
||||
_ = event, kwargs
|
||||
self._ensure_workflow_initialized()
|
||||
activate_workflow_handoff_by_task_id(self._application_generate_entity.task_id)
|
||||
self._base_task_pipeline.queue_manager.mark_execution_terminal()
|
||||
yield WorkflowMaintenancePausedStreamResponse(
|
||||
task_id=self._application_generate_entity.task_id,
|
||||
workflow_run_id=self._workflow_execution_id,
|
||||
)
|
||||
|
||||
def _handle_workflow_failed_and_stop_events(
|
||||
self,
|
||||
event: Union[QueueWorkflowFailedEvent, QueueStopEvent],
|
||||
@@ -624,6 +682,7 @@ class WorkflowAppGenerateTaskPipeline(GraphRuntimeStateSupport):
|
||||
QueueWorkflowSucceededEvent: self._handle_workflow_succeeded_event,
|
||||
QueueWorkflowPartialSuccessEvent: self._handle_workflow_partial_success_event,
|
||||
QueueWorkflowPausedEvent: self._handle_workflow_paused_event,
|
||||
QueueWorkflowMaintenancePausedEvent: self._handle_workflow_maintenance_paused_event,
|
||||
# Node events
|
||||
QueueNodeRetryEvent: self._handle_node_retry_event,
|
||||
QueueNodeStartedEvent: self._handle_node_started_event,
|
||||
@@ -702,45 +761,52 @@ class WorkflowAppGenerateTaskPipeline(GraphRuntimeStateSupport):
|
||||
Process stream response using elegant Fluent Python patterns.
|
||||
Maintains exact same functionality as original 44-if-statement version.
|
||||
"""
|
||||
for queue_message in self._base_task_pipeline.queue_manager.listen():
|
||||
event = queue_message.event
|
||||
queue_stream = self._base_task_pipeline.queue_manager.listen()
|
||||
try:
|
||||
for queue_message in queue_stream:
|
||||
event = queue_message.event
|
||||
|
||||
match event:
|
||||
case QueueWorkflowStartedEvent():
|
||||
self._resolve_graph_runtime_state()
|
||||
yield from self._handle_workflow_started_event(event)
|
||||
match event:
|
||||
case QueueWorkflowStartedEvent():
|
||||
self._resolve_graph_runtime_state()
|
||||
yield from self._handle_workflow_started_event(event)
|
||||
|
||||
case QueueTextChunkEvent():
|
||||
yield from self._handle_text_chunk_event(
|
||||
event, tts_publisher=tts_publisher, queue_message=queue_message
|
||||
)
|
||||
|
||||
case QueueErrorEvent():
|
||||
yield from self._handle_error_event(event)
|
||||
break
|
||||
|
||||
case QueueWorkflowFailedEvent():
|
||||
yield from self._handle_workflow_failed_and_stop_events(event)
|
||||
break
|
||||
case QueueWorkflowPausedEvent():
|
||||
yield from self._handle_workflow_paused_event(event)
|
||||
break
|
||||
|
||||
case QueueStopEvent():
|
||||
yield from self._handle_workflow_failed_and_stop_events(event)
|
||||
break
|
||||
|
||||
# Handle all other events through elegant dispatch
|
||||
case _:
|
||||
if responses := list(
|
||||
self._dispatch_event(
|
||||
event,
|
||||
tts_publisher=tts_publisher,
|
||||
trace_manager=trace_manager,
|
||||
queue_message=queue_message,
|
||||
case QueueTextChunkEvent():
|
||||
yield from self._handle_text_chunk_event(
|
||||
event, tts_publisher=tts_publisher, queue_message=queue_message
|
||||
)
|
||||
):
|
||||
yield from responses
|
||||
|
||||
case QueueErrorEvent():
|
||||
yield from self._handle_error_event(event)
|
||||
break
|
||||
|
||||
case QueueWorkflowFailedEvent():
|
||||
yield from self._handle_workflow_failed_and_stop_events(event)
|
||||
break
|
||||
case QueueWorkflowPausedEvent():
|
||||
yield from self._handle_workflow_paused_event(event)
|
||||
break
|
||||
case QueueWorkflowMaintenancePausedEvent():
|
||||
yield from self._handle_workflow_maintenance_paused_event(event)
|
||||
break
|
||||
|
||||
case QueueStopEvent():
|
||||
yield from self._handle_workflow_failed_and_stop_events(event)
|
||||
break
|
||||
|
||||
# Handle all other events through elegant dispatch
|
||||
case _:
|
||||
if responses := list(
|
||||
self._dispatch_event(
|
||||
event,
|
||||
tts_publisher=tts_publisher,
|
||||
trace_manager=trace_manager,
|
||||
queue_message=queue_message,
|
||||
)
|
||||
):
|
||||
yield from responses
|
||||
finally:
|
||||
close_stream(queue_stream)
|
||||
|
||||
if tts_publisher:
|
||||
tts_publisher.publish(None)
|
||||
|
||||
@@ -6,6 +6,7 @@ from typing import Any, cast
|
||||
from pydantic import ValidationError
|
||||
|
||||
from core.app.apps.base_app_queue_manager import AppQueueManager, PublishFrom
|
||||
from core.app.apps.workflow.command_channels import is_workflow_warm_shutdown_pause
|
||||
from core.app.entities.agent_strategy import AgentStrategyInfo
|
||||
from core.app.entities.app_invoke_entities import InvokeFrom, UserFrom, build_dify_run_context
|
||||
from core.app.entities.queue_entities import (
|
||||
@@ -29,11 +30,18 @@ from core.app.entities.queue_entities import (
|
||||
QueueStopEvent,
|
||||
QueueTextChunkEvent,
|
||||
QueueWorkflowFailedEvent,
|
||||
QueueWorkflowMaintenancePausedEvent,
|
||||
QueueWorkflowPartialSuccessEvent,
|
||||
QueueWorkflowPausedEvent,
|
||||
QueueWorkflowStartedEvent,
|
||||
QueueWorkflowSucceededEvent,
|
||||
)
|
||||
from core.app.layers.pause_state_persist_layer import PauseStatePersistenceLayer
|
||||
from core.app.layers.workflow_handoff_persist_layer import (
|
||||
WorkflowHandoffPersistenceError,
|
||||
WorkflowHandoffPersistenceLayer,
|
||||
)
|
||||
from core.app.layers.workflow_handoff_resume_layer import WorkflowHandoffResumeAcknowledgementLayer
|
||||
from core.rag.entities import RetrievalSourceMetadata
|
||||
from core.repositories.human_input_repository import HumanInputFormSubmissionRepository
|
||||
from core.workflow.node_factory import (
|
||||
@@ -45,16 +53,21 @@ from core.workflow.node_factory import (
|
||||
from core.workflow.nodes.human_input.boundary import enrich_graph_pause_reasons
|
||||
from core.workflow.nodes.human_input.pause_reason import HumanInputRequired
|
||||
from core.workflow.system_variables import (
|
||||
SystemVariableKey,
|
||||
build_bootstrap_variables,
|
||||
default_system_variables,
|
||||
build_system_variables,
|
||||
get_node_creation_preload_selectors,
|
||||
get_system_text,
|
||||
inject_default_system_variable_mappings,
|
||||
preload_node_creation_variables,
|
||||
)
|
||||
from core.workflow.variable_pool_initializer import add_variables_to_pool
|
||||
from core.workflow.workflow_entry import WorkflowEntry
|
||||
from core.workflow.workflow_run_outputs import project_node_outputs_for_workflow_run
|
||||
from extensions.workflow_warm_shutdown import mark_workflow_runs_stopped_if_running_without_active_handoff
|
||||
from graphon.entities import WorkflowStartReason
|
||||
from graphon.entities.graph_config import NodeConfigDictAdapter
|
||||
from graphon.entities.pause_reason import SchedulingPause
|
||||
from graphon.graph import Graph
|
||||
from graphon.graph_engine.layers import GraphEngineLayer
|
||||
from graphon.graph_events import (
|
||||
@@ -91,6 +104,9 @@ from models.workflow import Workflow
|
||||
from tasks.mail_human_input_delivery_task import dispatch_human_input_email_task
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
WORKFLOW_HANDOFF_PERSISTENCE_FAILURE_STOP_REASON = (
|
||||
"Workflow stopped because its worker-drain checkpoint could not be persisted."
|
||||
)
|
||||
|
||||
|
||||
class WorkflowBasedAppRunner:
|
||||
@@ -124,6 +140,7 @@ class WorkflowBasedAppRunner:
|
||||
user_id: str = "",
|
||||
root_node_id: str | None = None,
|
||||
trace_session_id: str | None = None,
|
||||
skip_validation: bool = False,
|
||||
) -> Graph:
|
||||
"""
|
||||
Init graph
|
||||
@@ -164,7 +181,12 @@ class WorkflowBasedAppRunner:
|
||||
root_node_id = get_default_root_node_id(graph_config)
|
||||
|
||||
# init graph
|
||||
graph = Graph.init(graph_config=graph_config, node_factory=node_factory, root_node_id=root_node_id)
|
||||
graph = Graph.init(
|
||||
graph_config=graph_config,
|
||||
node_factory=node_factory,
|
||||
root_node_id=root_node_id,
|
||||
skip_validation=skip_validation,
|
||||
)
|
||||
|
||||
if not graph:
|
||||
raise ValueError("graph not found in workflow")
|
||||
@@ -178,8 +200,9 @@ class WorkflowBasedAppRunner:
|
||||
single_loop_run: Any | None = None,
|
||||
*,
|
||||
user_id: str,
|
||||
workflow_execution_id: str,
|
||||
trace_session_id: str | None = None,
|
||||
) -> tuple[Graph, VariablePool, GraphRuntimeState]:
|
||||
) -> tuple[Graph, Mapping[str, Any], VariablePool, GraphRuntimeState]:
|
||||
"""
|
||||
Prepare graph, variable pool, and runtime state for single node execution
|
||||
(either single iteration or single loop).
|
||||
@@ -190,7 +213,7 @@ class WorkflowBasedAppRunner:
|
||||
single_loop_run: SingleLoopRunEntity if running single loop, None otherwise
|
||||
|
||||
Returns:
|
||||
A tuple containing (graph, variable_pool, graph_runtime_state)
|
||||
A tuple containing (graph, effective graph config, variable pool, runtime state)
|
||||
|
||||
Raises:
|
||||
ValueError: If neither single_iteration_run nor single_loop_run is specified
|
||||
@@ -200,7 +223,7 @@ class WorkflowBasedAppRunner:
|
||||
add_variables_to_pool(
|
||||
variable_pool,
|
||||
build_bootstrap_variables(
|
||||
system_variables=default_system_variables(),
|
||||
system_variables=build_system_variables(workflow_execution_id=workflow_execution_id),
|
||||
environment_variables=workflow.environment_variables,
|
||||
),
|
||||
)
|
||||
@@ -208,7 +231,7 @@ class WorkflowBasedAppRunner:
|
||||
|
||||
# Determine which type of single node execution and get graph/variable_pool
|
||||
if single_iteration_run:
|
||||
graph, variable_pool = self._get_graph_and_variable_pool_for_single_node_run(
|
||||
graph, graph_config, variable_pool = self._get_graph_and_variable_pool_for_single_node_run(
|
||||
workflow=workflow,
|
||||
node_id=single_iteration_run.node_id,
|
||||
user_inputs=dict(single_iteration_run.inputs),
|
||||
@@ -219,7 +242,7 @@ class WorkflowBasedAppRunner:
|
||||
trace_session_id=trace_session_id,
|
||||
)
|
||||
elif single_loop_run:
|
||||
graph, variable_pool = self._get_graph_and_variable_pool_for_single_node_run(
|
||||
graph, graph_config, variable_pool = self._get_graph_and_variable_pool_for_single_node_run(
|
||||
workflow=workflow,
|
||||
node_id=single_loop_run.node_id,
|
||||
user_inputs=dict(single_loop_run.inputs),
|
||||
@@ -234,7 +257,7 @@ class WorkflowBasedAppRunner:
|
||||
|
||||
# Return the graph, variable_pool, and the same graph_runtime_state used during graph creation
|
||||
# This ensures all nodes in the graph reference the same GraphRuntimeState instance
|
||||
return graph, variable_pool, graph_runtime_state
|
||||
return graph, graph_config, variable_pool, graph_runtime_state
|
||||
|
||||
def _get_graph_and_variable_pool_for_single_node_run(
|
||||
self,
|
||||
@@ -247,7 +270,7 @@ class WorkflowBasedAppRunner:
|
||||
*,
|
||||
user_id: str = "",
|
||||
trace_session_id: str | None = None,
|
||||
) -> tuple[Graph, VariablePool]:
|
||||
) -> tuple[Graph, Mapping[str, Any], VariablePool]:
|
||||
"""
|
||||
Get graph and variable pool for single node execution (iteration or loop).
|
||||
|
||||
@@ -260,7 +283,7 @@ class WorkflowBasedAppRunner:
|
||||
node_type_label: Label for error messages ('iteration' or 'loop')
|
||||
|
||||
Returns:
|
||||
A tuple containing (graph, variable_pool)
|
||||
A tuple containing (graph, effective graph config, variable pool)
|
||||
"""
|
||||
# fetch workflow graph
|
||||
graph_config = workflow.graph_dict
|
||||
@@ -391,7 +414,13 @@ class WorkflowBasedAppRunner:
|
||||
if not graph:
|
||||
raise ValueError("graph not found in workflow")
|
||||
|
||||
return graph, variable_pool
|
||||
return graph, graph_config, variable_pool
|
||||
|
||||
def _configure_execution_root(self, root_node_id: str) -> None:
|
||||
"""Give pause persistence layers the exact root used by Graph.init."""
|
||||
for layer in self._graph_engine_layers:
|
||||
if isinstance(layer, (WorkflowHandoffPersistenceLayer, PauseStatePersistenceLayer)):
|
||||
layer.set_execution_root_node_id(root_node_id)
|
||||
|
||||
@staticmethod
|
||||
def _build_agent_strategy_info(event: NodeRunStartedEvent) -> AgentStrategyInfo | None:
|
||||
@@ -405,6 +434,68 @@ class WorkflowBasedAppRunner:
|
||||
logger.warning("Invalid agent strategy payload for node %s", event.node_id, exc_info=True)
|
||||
return None
|
||||
|
||||
def _handle_event_with_handoff_contracts(self, workflow_entry: WorkflowEntry, event: GraphEngineEvent) -> None:
|
||||
"""Enforce durable handoff contracts before publishing an event."""
|
||||
self._require_handoff_layer_contracts(workflow_entry, event)
|
||||
self._handle_event(workflow_entry, event)
|
||||
|
||||
def _require_handoff_layer_contracts(self, workflow_entry: WorkflowEntry, event: GraphEngineEvent) -> None:
|
||||
"""Surface layer failures that Graphon intentionally logs and swallows."""
|
||||
if isinstance(event, GraphRunStartedEvent) and event.reason == WorkflowStartReason.RESUMPTION:
|
||||
acknowledgement_error: Exception | None = None
|
||||
for layer in self._graph_engine_layers:
|
||||
if isinstance(layer, WorkflowHandoffResumeAcknowledgementLayer):
|
||||
try:
|
||||
layer.require_acknowledged()
|
||||
except Exception as error:
|
||||
acknowledgement_error = acknowledgement_error or error
|
||||
if acknowledgement_error is not None:
|
||||
raise acknowledgement_error
|
||||
return
|
||||
|
||||
if not isinstance(event, GraphRunPausedEvent) or not is_workflow_warm_shutdown_pause(event.reasons):
|
||||
return
|
||||
|
||||
persistence_layers = [
|
||||
layer for layer in self._graph_engine_layers if isinstance(layer, WorkflowHandoffPersistenceLayer)
|
||||
]
|
||||
persistence_error: Exception | None = None
|
||||
if not persistence_layers:
|
||||
persistence_error = WorkflowHandoffPersistenceError(
|
||||
"Worker-drain pause was emitted without a workflow handoff persistence layer"
|
||||
)
|
||||
for layer in persistence_layers:
|
||||
try:
|
||||
layer.require_persisted_handoff()
|
||||
except Exception as error:
|
||||
persistence_error = persistence_error or error
|
||||
if persistence_error is not None:
|
||||
self._fail_closed_after_handoff_persistence_error(workflow_entry)
|
||||
raise persistence_error
|
||||
|
||||
@staticmethod
|
||||
def _fail_closed_after_handoff_persistence_error(workflow_entry: WorkflowEntry) -> None:
|
||||
"""Best-effort terminal update without masking the layer contract error."""
|
||||
try:
|
||||
workflow_run_id = get_system_text(
|
||||
workflow_entry.graph_engine.graph_runtime_state.variable_pool,
|
||||
SystemVariableKey.WORKFLOW_EXECUTION_ID,
|
||||
)
|
||||
if workflow_run_id is None:
|
||||
logger.error("Cannot fail closed after handoff persistence error: workflow run id is missing")
|
||||
return
|
||||
updated_count = mark_workflow_runs_stopped_if_running_without_active_handoff(
|
||||
(workflow_run_id,),
|
||||
reason=WORKFLOW_HANDOFF_PERSISTENCE_FAILURE_STOP_REASON,
|
||||
)
|
||||
if updated_count == 0:
|
||||
logger.warning(
|
||||
"Handoff persistence failed for workflow run %s, but its state was already terminal or recoverable",
|
||||
workflow_run_id,
|
||||
)
|
||||
except Exception:
|
||||
logger.exception("Failed to mark workflow run STOPPED after handoff persistence error")
|
||||
|
||||
def _handle_event(self, workflow_entry: WorkflowEntry, event: GraphEngineEvent):
|
||||
"""
|
||||
Handle event
|
||||
@@ -434,6 +525,16 @@ class WorkflowBasedAppRunner:
|
||||
case GraphRunPausedEvent():
|
||||
runtime_state = workflow_entry.graph_engine.graph_runtime_state
|
||||
paused_nodes = runtime_state.get_paused_nodes()
|
||||
if is_workflow_warm_shutdown_pause(event.reasons):
|
||||
self._publish_event(
|
||||
QueueWorkflowMaintenancePausedEvent(
|
||||
reasons=[reason for reason in event.reasons if isinstance(reason, SchedulingPause)],
|
||||
outputs=event.outputs,
|
||||
paused_nodes=paused_nodes,
|
||||
)
|
||||
)
|
||||
return
|
||||
|
||||
enriched_reasons = enrich_graph_pause_reasons(
|
||||
reasons=event.reasons,
|
||||
form_repository=HumanInputFormSubmissionRepository(),
|
||||
|
||||
@@ -50,6 +50,7 @@ class QueueEvent(StrEnum):
|
||||
STOP = "stop"
|
||||
RETRY = "retry"
|
||||
PAUSE = "pause"
|
||||
WORKFLOW_MAINTENANCE_PAUSED = "workflow_maintenance_paused"
|
||||
HUMAN_INPUT_FORM_FILLED = "human_input_form_filled"
|
||||
HUMAN_INPUT_FORM_TIMEOUT = "human_input_form_timeout"
|
||||
|
||||
@@ -588,3 +589,17 @@ class QueueWorkflowPausedEvent(AppQueueEvent):
|
||||
reasons: Sequence[PauseReason] = Field(default_factory=list)
|
||||
outputs: Mapping[str, object] = Field(default_factory=dict)
|
||||
paused_nodes: Sequence[str] = Field(default_factory=list)
|
||||
|
||||
|
||||
class QueueWorkflowMaintenancePausedEvent(AppQueueEvent):
|
||||
"""Internal event marking a durable worker-drain checkpoint.
|
||||
|
||||
Unlike ``QueueWorkflowPausedEvent``, this event must never be exposed as a
|
||||
public workflow-paused lifecycle event. It terminates only the execution
|
||||
segment owned by the draining worker; another worker resumes the same run.
|
||||
"""
|
||||
|
||||
event: QueueEvent = QueueEvent.WORKFLOW_MAINTENANCE_PAUSED
|
||||
reasons: Sequence[PauseReason] = Field(default_factory=list)
|
||||
outputs: Mapping[str, object] = Field(default_factory=dict)
|
||||
paused_nodes: Sequence[str] = Field(default_factory=list)
|
||||
|
||||
@@ -12,3 +12,6 @@ class RagPipelineInvokeEntity(BaseModel):
|
||||
streaming: bool
|
||||
workflow_execution_id: str | None = None
|
||||
workflow_thread_pool_id: str | None = None
|
||||
# Stored inside the uploaded batch instead of the Celery task signature so
|
||||
# new producers remain compatible with old workers during a rolling update.
|
||||
tenant_isolated: bool | None = None
|
||||
|
||||
@@ -76,6 +76,7 @@ class StreamEvent(StrEnum):
|
||||
AGENT_MESSAGE = "agent_message"
|
||||
WORKFLOW_STARTED = "workflow_started"
|
||||
WORKFLOW_PAUSED = "workflow_paused"
|
||||
WORKFLOW_MAINTENANCE_PAUSED = "workflow_maintenance_paused"
|
||||
WORKFLOW_FINISHED = "workflow_finished"
|
||||
NODE_STARTED = "node_started"
|
||||
NODE_FINISHED = "node_finished"
|
||||
@@ -247,6 +248,7 @@ class WorkflowFinishStreamResponse(StreamResponse):
|
||||
finished_at: int | None
|
||||
exceptions_count: int = 0
|
||||
files: Sequence[Mapping[str, Any]] | None = []
|
||||
handoff_duration: float = 0.0
|
||||
|
||||
event: StreamEvent = StreamEvent.WORKFLOW_FINISHED
|
||||
workflow_run_id: str
|
||||
@@ -272,12 +274,25 @@ class WorkflowPauseStreamResponse(StreamResponse):
|
||||
elapsed_time: float
|
||||
total_tokens: int
|
||||
total_steps: int
|
||||
handoff_duration: float = 0.0
|
||||
|
||||
event: StreamEvent = StreamEvent.WORKFLOW_PAUSED
|
||||
workflow_run_id: str
|
||||
data: Data
|
||||
|
||||
|
||||
class WorkflowMaintenancePausedStreamResponse(StreamResponse):
|
||||
"""Internal control event for a completed worker-drain execution segment.
|
||||
|
||||
Publishers must consume this sentinel without forwarding it to public
|
||||
streams. It is intentionally distinct from ``WorkflowPauseStreamResponse``
|
||||
because maintenance handoff is transparent to clients.
|
||||
"""
|
||||
|
||||
event: StreamEvent = StreamEvent.WORKFLOW_MAINTENANCE_PAUSED
|
||||
workflow_run_id: str
|
||||
|
||||
|
||||
class HumanInputRequiredResponse(StreamResponse):
|
||||
class Data(BaseModel):
|
||||
"""
|
||||
@@ -820,6 +835,13 @@ class AppBlockingResponse(BaseModel):
|
||||
task_id: str
|
||||
|
||||
|
||||
class WorkflowMaintenancePausedBlockingResponse(AppBlockingResponse):
|
||||
"""Internal blocking-mode result for a completed worker-drain segment."""
|
||||
|
||||
event: StreamEvent = StreamEvent.WORKFLOW_MAINTENANCE_PAUSED
|
||||
workflow_run_id: str
|
||||
|
||||
|
||||
class ChatbotAppBlockingResponse(AppBlockingResponse):
|
||||
"""
|
||||
ChatbotAppBlockingResponse entity
|
||||
@@ -865,6 +887,7 @@ class AdvancedChatPausedBlockingResponse(AppBlockingResponse):
|
||||
elapsed_time: float
|
||||
total_tokens: int
|
||||
total_steps: int
|
||||
handoff_duration: float = 0.0
|
||||
|
||||
data: Data
|
||||
|
||||
@@ -909,6 +932,7 @@ class WorkflowAppBlockingResponse(AppBlockingResponse):
|
||||
total_steps: int
|
||||
created_at: int
|
||||
finished_at: int | None
|
||||
handoff_duration: float = 0.0
|
||||
|
||||
workflow_run_id: str
|
||||
data: Data
|
||||
@@ -936,6 +960,7 @@ class WorkflowAppPausedBlockingResponse(AppBlockingResponse):
|
||||
finished_at: int | None
|
||||
paused_nodes: Sequence[str] = Field(default_factory=list)
|
||||
reasons: Sequence[Mapping[str, Any]] = Field(default_factory=list)
|
||||
handoff_duration: float = 0.0
|
||||
|
||||
workflow_run_id: str
|
||||
data: Data
|
||||
|
||||
@@ -124,6 +124,7 @@ class RateLimitGenerator:
|
||||
self.generator = generator
|
||||
self.request_id = request_id
|
||||
self.closed = False
|
||||
self.workflow_run_id: str | None = None
|
||||
|
||||
def __iter__(self):
|
||||
return self
|
||||
|
||||
@@ -1,11 +1,18 @@
|
||||
import time
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass
|
||||
from typing import Annotated, Literal, Self, override
|
||||
from typing import Annotated, Literal, Protocol, Self, override, runtime_checkable
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
from sqlalchemy import Engine
|
||||
from sqlalchemy.orm import Session, sessionmaker
|
||||
|
||||
from core.app.entities.app_invoke_entities import AdvancedChatAppGenerateEntity, WorkflowAppGenerateEntity
|
||||
from core.app.apps.workflow.command_channels import is_workflow_warm_shutdown_pause
|
||||
from core.app.entities.app_invoke_entities import (
|
||||
AdvancedChatAppGenerateEntity,
|
||||
RagPipelineGenerateEntity,
|
||||
WorkflowAppGenerateEntity,
|
||||
)
|
||||
from core.repositories.human_input_repository import HumanInputFormSubmissionRepository
|
||||
from core.workflow.nodes.human_input.boundary import enrich_graph_pause_reasons
|
||||
from core.workflow.system_variables import SystemVariableKey, get_system_text
|
||||
@@ -16,6 +23,32 @@ from models.model import AppMode
|
||||
from repositories.api_workflow_run_repository import APIWorkflowRunRepository
|
||||
from repositories.factory import DifyAPIRepositoryFactory
|
||||
|
||||
WORKFLOW_HANDOFF_ACTIVE_EXECUTION_SECONDS_EXTRA_KEY = "workflow_handoff_active_execution_seconds"
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class _GenerateEntityExtras(Protocol):
|
||||
@property
|
||||
def extras(self) -> object: ...
|
||||
|
||||
|
||||
def get_workflow_handoff_active_execution_seconds(
|
||||
generate_entity: WorkflowAppGenerateEntity | AdvancedChatAppGenerateEntity | RagPipelineGenerateEntity,
|
||||
) -> float:
|
||||
if not isinstance(generate_entity, _GenerateEntityExtras):
|
||||
return 0.0
|
||||
extras = generate_entity.extras
|
||||
if extras is None:
|
||||
return 0.0
|
||||
if not isinstance(extras, dict):
|
||||
raise ValueError("Workflow generate entity extras are invalid")
|
||||
value = extras.get(WORKFLOW_HANDOFF_ACTIVE_EXECUTION_SECONDS_EXTRA_KEY, 0.0)
|
||||
if isinstance(value, bool) or not isinstance(value, (int, float)):
|
||||
raise ValueError("Workflow handoff active execution time is invalid")
|
||||
if value < 0:
|
||||
raise ValueError("Workflow handoff active execution time must be non-negative")
|
||||
return float(value)
|
||||
|
||||
|
||||
# Wrapper types for `WorkflowAppGenerateEntity` and
|
||||
# `AdvancedChatAppGenerateEntity`. These wrappers enable type discrimination
|
||||
@@ -30,8 +63,13 @@ class _AdvancedChatAppGenerateEntityWrapper(BaseModel):
|
||||
entity: AdvancedChatAppGenerateEntity
|
||||
|
||||
|
||||
class _RagPipelineGenerateEntityWrapper(BaseModel):
|
||||
type: Literal[AppMode.RAG_PIPELINE] = AppMode.RAG_PIPELINE
|
||||
entity: RagPipelineGenerateEntity
|
||||
|
||||
|
||||
type _GenerateEntityUnion = Annotated[
|
||||
_WorkflowGenerateEntityWrapper | _AdvancedChatAppGenerateEntityWrapper,
|
||||
_WorkflowGenerateEntityWrapper | _AdvancedChatAppGenerateEntityWrapper | _RagPipelineGenerateEntityWrapper,
|
||||
Field(discriminator="type"),
|
||||
]
|
||||
|
||||
@@ -41,13 +79,52 @@ class WorkflowResumptionContext(BaseModel):
|
||||
|
||||
version: Literal["1"] = "1"
|
||||
|
||||
# Only workflow / chatflow could be paused.
|
||||
generate_entity: _GenerateEntityUnion
|
||||
serialized_graph_runtime_state: str
|
||||
# The graph is reconstructed from the immutable WorkflowRun.graph snapshot,
|
||||
# but the active root cannot always be inferred from that graph. Triggered
|
||||
# executions may select a non-default root, and single iteration/loop runs
|
||||
# intentionally root a filtered graph at the container node.
|
||||
#
|
||||
# Optional for adjacent-version snapshots created before exact-root handoff
|
||||
# support. New maintenance handoffs require this value when they are written.
|
||||
root_node_id: str | None = None
|
||||
# Optional so that a workflow run paused before this field existed still
|
||||
# loads: it just degrades to fresh-filter behavior on resume for that one
|
||||
# stale run.
|
||||
serialized_response_stream_filter_state: str | None = None
|
||||
# Cumulative time spent actively executing before this checkpoint. This is
|
||||
# optional for adjacent-version snapshots created before handoff timing was
|
||||
# introduced; maintenance wait itself is deliberately not included.
|
||||
active_execution_seconds: float = Field(default=0.0, ge=0.0)
|
||||
|
||||
@classmethod
|
||||
def from_runtime_snapshot(
|
||||
cls,
|
||||
*,
|
||||
generate_entity: WorkflowAppGenerateEntity | AdvancedChatAppGenerateEntity | RagPipelineGenerateEntity,
|
||||
serialized_graph_runtime_state: str,
|
||||
serialized_response_stream_filter_state: str | None,
|
||||
active_execution_seconds: float = 0.0,
|
||||
root_node_id: str | None = None,
|
||||
) -> Self:
|
||||
"""Build a versioned context shared by user pauses and worker handoffs."""
|
||||
entity_wrapper: _GenerateEntityUnion
|
||||
# RagPipelineGenerateEntity subclasses WorkflowAppGenerateEntity, so it
|
||||
# must be checked first or its dataset/document resume fields are lost.
|
||||
if isinstance(generate_entity, RagPipelineGenerateEntity):
|
||||
entity_wrapper = _RagPipelineGenerateEntityWrapper(entity=generate_entity)
|
||||
elif isinstance(generate_entity, WorkflowAppGenerateEntity):
|
||||
entity_wrapper = _WorkflowGenerateEntityWrapper(entity=generate_entity)
|
||||
else:
|
||||
entity_wrapper = _AdvancedChatAppGenerateEntityWrapper(entity=generate_entity)
|
||||
return cls(
|
||||
serialized_graph_runtime_state=serialized_graph_runtime_state,
|
||||
generate_entity=entity_wrapper,
|
||||
serialized_response_stream_filter_state=serialized_response_stream_filter_state,
|
||||
active_execution_seconds=active_execution_seconds,
|
||||
root_node_id=root_node_id,
|
||||
)
|
||||
|
||||
def dumps(self) -> str:
|
||||
return self.model_dump_json()
|
||||
@@ -56,7 +133,9 @@ class WorkflowResumptionContext(BaseModel):
|
||||
def loads(cls, value: str) -> Self:
|
||||
return cls.model_validate_json(value)
|
||||
|
||||
def get_generate_entity(self) -> WorkflowAppGenerateEntity | AdvancedChatAppGenerateEntity:
|
||||
def get_generate_entity(
|
||||
self,
|
||||
) -> WorkflowAppGenerateEntity | AdvancedChatAppGenerateEntity | RagPipelineGenerateEntity:
|
||||
return self.generate_entity.entity
|
||||
|
||||
def get_response_stream_filter(self) -> ResponseStreamFilter:
|
||||
@@ -65,6 +144,12 @@ class WorkflowResumptionContext(BaseModel):
|
||||
response_stream_filter.loads(self.serialized_response_stream_filter_state)
|
||||
return response_stream_filter
|
||||
|
||||
def apply_handoff_execution_timing(self) -> None:
|
||||
"""Carry cumulative active time into the next segment and checkpoint."""
|
||||
self.get_generate_entity().extras[WORKFLOW_HANDOFF_ACTIVE_EXECUTION_SECONDS_EXTRA_KEY] = (
|
||||
self.active_execution_seconds
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class PauseStateLayerConfig:
|
||||
@@ -78,9 +163,10 @@ class PauseStatePersistenceLayer(GraphEngineLayer):
|
||||
def __init__(
|
||||
self,
|
||||
session_factory: Engine | sessionmaker[Session],
|
||||
generate_entity: WorkflowAppGenerateEntity | AdvancedChatAppGenerateEntity,
|
||||
generate_entity: WorkflowAppGenerateEntity | AdvancedChatAppGenerateEntity | RagPipelineGenerateEntity,
|
||||
state_owner_user_id: str,
|
||||
response_stream_filter: ResponseStreamFilter,
|
||||
monotonic_clock: Callable[[], float] = time.monotonic,
|
||||
):
|
||||
"""Create a PauseStatePersistenceLayer.
|
||||
|
||||
@@ -99,6 +185,19 @@ class PauseStatePersistenceLayer(GraphEngineLayer):
|
||||
self._state_owner_user_id = state_owner_user_id
|
||||
self._generate_entity = generate_entity
|
||||
self._response_stream_filter = response_stream_filter
|
||||
self._prior_active_execution_seconds = get_workflow_handoff_active_execution_seconds(generate_entity)
|
||||
self._monotonic_clock = monotonic_clock
|
||||
self._active_segment_started_at: float | None = None
|
||||
self._root_node_id: str | None = None
|
||||
|
||||
def set_execution_root_node_id(self, root_node_id: str) -> None:
|
||||
if not root_node_id:
|
||||
raise ValueError("Workflow pause root node id must not be empty")
|
||||
if self._active_segment_started_at is not None:
|
||||
raise RuntimeError("Workflow pause root node id cannot change after graph start")
|
||||
if self._root_node_id is not None and self._root_node_id != root_node_id:
|
||||
raise RuntimeError("Workflow pause root node id was configured more than once")
|
||||
self._root_node_id = root_node_id
|
||||
|
||||
def _get_repo(self) -> APIWorkflowRunRepository:
|
||||
return DifyAPIRepositoryFactory.create_api_workflow_run_repository(self._session_maker)
|
||||
@@ -111,7 +210,7 @@ class PauseStatePersistenceLayer(GraphEngineLayer):
|
||||
This is called after the engine has been initialized but before any nodes
|
||||
are executed. Layers can use this to set up resources or log start information.
|
||||
"""
|
||||
pass
|
||||
self._active_segment_started_at = self._monotonic_clock()
|
||||
|
||||
@override
|
||||
def on_event(self, event: GraphEngineEvent) -> None:
|
||||
@@ -129,21 +228,30 @@ class PauseStatePersistenceLayer(GraphEngineLayer):
|
||||
"""
|
||||
if not isinstance(event, GraphRunPausedEvent):
|
||||
return
|
||||
graph_runtime_state = self.graph_runtime_state
|
||||
if is_workflow_warm_shutdown_pause(event.reasons):
|
||||
# Planned worker drains use the durable handoff layer. Persisting
|
||||
# them here would expose the internal checkpoint as user-visible
|
||||
# PAUSED state and make HITL resume semantics race the handoff.
|
||||
return
|
||||
|
||||
entity_wrapper: _GenerateEntityUnion
|
||||
if isinstance(self._generate_entity, WorkflowAppGenerateEntity):
|
||||
entity_wrapper = _WorkflowGenerateEntityWrapper(entity=self._generate_entity)
|
||||
else:
|
||||
entity_wrapper = _AdvancedChatAppGenerateEntityWrapper(entity=self._generate_entity)
|
||||
if self._active_segment_started_at is None:
|
||||
raise RuntimeError("Workflow pause layer did not observe graph start")
|
||||
segment_execution_seconds = max(
|
||||
self._monotonic_clock() - self._active_segment_started_at,
|
||||
0.0,
|
||||
)
|
||||
|
||||
state = WorkflowResumptionContext(
|
||||
serialized_graph_runtime_state=self.graph_runtime_state.dumps(),
|
||||
generate_entity=entity_wrapper,
|
||||
state = WorkflowResumptionContext.from_runtime_snapshot(
|
||||
serialized_graph_runtime_state=graph_runtime_state.dumps(),
|
||||
generate_entity=self._generate_entity,
|
||||
serialized_response_stream_filter_state=self._response_stream_filter.dumps(),
|
||||
active_execution_seconds=self._prior_active_execution_seconds + segment_execution_seconds,
|
||||
root_node_id=self._root_node_id,
|
||||
)
|
||||
|
||||
workflow_run_id = get_system_text(
|
||||
self.graph_runtime_state.variable_pool,
|
||||
graph_runtime_state.variable_pool,
|
||||
SystemVariableKey.WORKFLOW_EXECUTION_ID,
|
||||
)
|
||||
assert workflow_run_id is not None
|
||||
@@ -153,7 +261,7 @@ class PauseStatePersistenceLayer(GraphEngineLayer):
|
||||
pause_reasons = enrich_graph_pause_reasons(
|
||||
reasons=event.reasons,
|
||||
form_repository=HumanInputFormSubmissionRepository(),
|
||||
variable_pool=self.graph_runtime_state.variable_pool,
|
||||
variable_pool=graph_runtime_state.variable_pool,
|
||||
)
|
||||
repo = self._get_repo()
|
||||
repo.create_workflow_pause(
|
||||
|
||||
@@ -4,6 +4,7 @@ from typing import Any, ClassVar, override
|
||||
|
||||
from pydantic import TypeAdapter
|
||||
|
||||
from core.app.apps.workflow.command_channels import is_workflow_warm_shutdown_pause
|
||||
from core.db.session_factory import session_factory
|
||||
from core.workflow.system_variables import SystemVariableKey, get_system_text
|
||||
from graphon.graph_engine.layers import GraphEngineLayer
|
||||
@@ -14,7 +15,9 @@ from graphon.graph_events import (
|
||||
GraphRunPausedEvent,
|
||||
GraphRunSucceededEvent,
|
||||
)
|
||||
from libs.datetime_utils import ensure_naive_utc
|
||||
from models.enums import WorkflowTriggerStatus
|
||||
from models.workflow import WorkflowRun
|
||||
from repositories.sqlalchemy_workflow_trigger_log_repository import SQLAlchemyWorkflowTriggerLogRepository
|
||||
from tasks.workflow_cfs_scheduler.cfs_scheduler import AsyncWorkflowCFSPlanEntity
|
||||
|
||||
@@ -46,13 +49,41 @@ class TriggerPostLayer(GraphEngineLayer):
|
||||
|
||||
@override
|
||||
def on_graph_start(self):
|
||||
pass
|
||||
# Persist the association before the graph can enter a maintenance
|
||||
# handoff. The pause event is intentionally transparent, so waiting
|
||||
# until a terminal event would leave the compensation scanner unable
|
||||
# to find and terminalize the trigger log after resume exhaustion.
|
||||
workflow_run_id = get_system_text(
|
||||
self.graph_runtime_state.variable_pool,
|
||||
SystemVariableKey.WORKFLOW_EXECUTION_ID,
|
||||
)
|
||||
if not workflow_run_id:
|
||||
logger.warning("Workflow run id is not set when trigger graph starts: %s", self.trigger_log_id)
|
||||
return
|
||||
|
||||
with session_factory.create_session() as session:
|
||||
repo = SQLAlchemyWorkflowTriggerLogRepository(session)
|
||||
trigger_log = repo.get_by_id(self.trigger_log_id)
|
||||
if not trigger_log:
|
||||
logger.error("Trigger log not found: %s", self.trigger_log_id)
|
||||
return
|
||||
if trigger_log.workflow_run_id == workflow_run_id:
|
||||
return
|
||||
trigger_log.workflow_run_id = workflow_run_id
|
||||
repo.update(trigger_log)
|
||||
session.commit()
|
||||
|
||||
@override
|
||||
def on_event(self, event: GraphEngineEvent):
|
||||
"""
|
||||
Update trigger log with success or failure.
|
||||
"""
|
||||
if isinstance(event, GraphRunPausedEvent) and is_workflow_warm_shutdown_pause(event.reasons):
|
||||
# Maintenance handoff is transparent to the logical trigger run.
|
||||
# The resumed worker will eventually persist its actual terminal
|
||||
# status and cumulative execution statistics.
|
||||
return
|
||||
|
||||
if isinstance(event, tuple(self._STATUS_MAP.keys())):
|
||||
with session_factory.create_session() as session:
|
||||
repo = SQLAlchemyWorkflowTriggerLogRepository(session)
|
||||
@@ -61,8 +92,8 @@ class TriggerPostLayer(GraphEngineLayer):
|
||||
logger.exception("Trigger log not found: %s", self.trigger_log_id)
|
||||
return
|
||||
|
||||
# Calculate elapsed time
|
||||
elapsed_time = (datetime.now(UTC) - self.start_time).total_seconds()
|
||||
now = datetime.now(UTC)
|
||||
segment_elapsed_time = (now - self.start_time).total_seconds()
|
||||
|
||||
# Extract relevant data from result
|
||||
outputs = self.graph_runtime_state.outputs
|
||||
@@ -83,13 +114,23 @@ class TriggerPostLayer(GraphEngineLayer):
|
||||
if isinstance(event, GraphRunAbortedEvent):
|
||||
trigger_log.error = event.reason or "Workflow execution aborted"
|
||||
|
||||
if trigger_log.elapsed_time is None:
|
||||
trigger_log.elapsed_time = elapsed_time
|
||||
workflow_run = session.get(WorkflowRun, workflow_run_id)
|
||||
if workflow_run is not None:
|
||||
# WorkflowRun is the source of user-visible wall-clock
|
||||
# timing across maintenance handoffs. Execution limits and
|
||||
# quota accounting continue to use Graphon's active segment
|
||||
# clocks and token counters instead.
|
||||
trigger_log.elapsed_time = max(
|
||||
(ensure_naive_utc(now) - ensure_naive_utc(workflow_run.created_at)).total_seconds(),
|
||||
0.0,
|
||||
)
|
||||
elif trigger_log.elapsed_time is None:
|
||||
trigger_log.elapsed_time = segment_elapsed_time
|
||||
else:
|
||||
trigger_log.elapsed_time += elapsed_time
|
||||
trigger_log.elapsed_time += segment_elapsed_time
|
||||
|
||||
trigger_log.total_tokens = total_tokens
|
||||
trigger_log.finished_at = datetime.now(UTC)
|
||||
trigger_log.finished_at = now
|
||||
repo.update(trigger_log)
|
||||
session.commit()
|
||||
|
||||
|
||||
@@ -0,0 +1,257 @@
|
||||
import logging
|
||||
import time
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass
|
||||
from typing import override
|
||||
|
||||
from core.app.apps.workflow.command_channels import is_workflow_warm_shutdown_pause
|
||||
from core.app.entities.app_invoke_entities import (
|
||||
AdvancedChatAppGenerateEntity,
|
||||
RagPipelineGenerateEntity,
|
||||
WorkflowAppGenerateEntity,
|
||||
)
|
||||
from core.app.layers.pause_state_persist_layer import (
|
||||
WorkflowResumptionContext,
|
||||
get_workflow_handoff_active_execution_seconds,
|
||||
)
|
||||
from core.workflow.system_variables import SystemVariableKey, get_system_text
|
||||
from graphon.filters import ResponseStreamFilter
|
||||
from graphon.graph_engine.layers import GraphEngineLayer
|
||||
from graphon.graph_events import GraphEngineEvent, GraphRunPausedEvent
|
||||
from models.workflow_handoff import (
|
||||
RAG_PIPELINE_QUEUE_KIND_EXTRA_KEY,
|
||||
RAG_PIPELINE_SOURCE_BATCH_ID_EXTRA_KEY,
|
||||
RAG_PIPELINE_TENANT_ID_EXTRA_KEY,
|
||||
RAG_PIPELINE_TENANT_ISOLATED_EXTRA_KEY,
|
||||
RagPipelineHandoffGroupMetadata,
|
||||
RagPipelineQueueKind,
|
||||
WorkflowHandoffResumeRoute,
|
||||
WorkflowHandoffState,
|
||||
WorkflowRunHandoff,
|
||||
)
|
||||
from services.workflow_handoff_service import WorkflowHandoffService
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
type ResumableWorkflowGenerateEntity = (
|
||||
WorkflowAppGenerateEntity | AdvancedChatAppGenerateEntity | RagPipelineGenerateEntity
|
||||
)
|
||||
|
||||
|
||||
class WorkflowHandoffPersistenceError(RuntimeError):
|
||||
"""Raised by the explicit post-pause durability check."""
|
||||
|
||||
|
||||
class WorkflowHandoffNotObservedError(RuntimeError):
|
||||
"""Raised when a caller checks a layer that did not observe maintenance pause."""
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class WorkflowHandoffLayerConfig:
|
||||
"""Dependencies used to inject durable handoff persistence into a runner."""
|
||||
|
||||
handoff_service: WorkflowHandoffService
|
||||
source_worker_id: str
|
||||
resume_route: WorkflowHandoffResumeRoute | None = None
|
||||
|
||||
|
||||
class WorkflowHandoffPersistenceLayer(GraphEngineLayer):
|
||||
"""Persist only planned worker-drain pauses as durable handoff checkpoints.
|
||||
|
||||
Graphon 0.6 logs and swallows exceptions raised by ``GraphEngineLayer.on_event``.
|
||||
This layer therefore records any persistence error and exposes
|
||||
``require_persisted_handoff`` as a mandatory post-pause contract for the runner.
|
||||
The terminal pause event is notified synchronously before GraphEngine yields it,
|
||||
so the result is available when the runner receives that event.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
handoff_service: WorkflowHandoffService,
|
||||
generate_entity: ResumableWorkflowGenerateEntity,
|
||||
source_worker_id: str,
|
||||
response_stream_filter: ResponseStreamFilter,
|
||||
resume_route: WorkflowHandoffResumeRoute | None = None,
|
||||
monotonic_clock: Callable[[], float] = time.monotonic,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self._handoff_service = handoff_service
|
||||
self._generate_entity = generate_entity
|
||||
self._source_worker_id = source_worker_id
|
||||
self._response_stream_filter = response_stream_filter
|
||||
self._resume_route = resume_route or infer_workflow_handoff_resume_route(generate_entity)
|
||||
self._prior_active_execution_seconds = get_workflow_handoff_active_execution_seconds(generate_entity)
|
||||
self._monotonic_clock = monotonic_clock
|
||||
self._active_segment_started_at: float | None = None
|
||||
self._root_node_id: str | None = None
|
||||
self._maintenance_pause_observed = False
|
||||
self._persisted_handoff: WorkflowRunHandoff | None = None
|
||||
self._persistence_error: Exception | None = None
|
||||
|
||||
@property
|
||||
def maintenance_pause_observed(self) -> bool:
|
||||
return self._maintenance_pause_observed
|
||||
|
||||
@property
|
||||
def persisted_handoff(self) -> WorkflowRunHandoff | None:
|
||||
return self._persisted_handoff
|
||||
|
||||
@property
|
||||
def persistence_error(self) -> Exception | None:
|
||||
return self._persistence_error
|
||||
|
||||
def set_execution_root_node_id(self, root_node_id: str) -> None:
|
||||
"""Record the exact Graph root before execution starts.
|
||||
|
||||
WorkflowRun.graph preserves the effective graph and version. The root
|
||||
is persisted alongside the runtime state because it cannot be inferred
|
||||
for custom trigger roots or single iteration/loop debugger runs.
|
||||
"""
|
||||
if not root_node_id:
|
||||
raise ValueError("Workflow handoff root node id must not be empty")
|
||||
if self._active_segment_started_at is not None:
|
||||
raise RuntimeError("Workflow handoff root node id cannot change after graph start")
|
||||
if self._root_node_id is not None and self._root_node_id != root_node_id:
|
||||
raise RuntimeError("Workflow handoff root node id was configured more than once")
|
||||
self._root_node_id = root_node_id
|
||||
|
||||
@override
|
||||
def on_graph_start(self) -> None:
|
||||
self._maintenance_pause_observed = False
|
||||
self._persisted_handoff = None
|
||||
self._persistence_error = None
|
||||
self._active_segment_started_at = self._monotonic_clock()
|
||||
|
||||
@override
|
||||
def on_event(self, event: GraphEngineEvent) -> None:
|
||||
if not isinstance(event, GraphRunPausedEvent) or not is_workflow_warm_shutdown_pause(event.reasons):
|
||||
return
|
||||
|
||||
self._maintenance_pause_observed = True
|
||||
if self._persisted_handoff is not None or self._persistence_error is not None:
|
||||
return
|
||||
|
||||
try:
|
||||
workflow_run_id = get_system_text(
|
||||
self.graph_runtime_state.variable_pool,
|
||||
SystemVariableKey.WORKFLOW_EXECUTION_ID,
|
||||
)
|
||||
if workflow_run_id is None:
|
||||
raise ValueError("Workflow execution id is missing from graph runtime state")
|
||||
if self._active_segment_started_at is None:
|
||||
raise RuntimeError("Workflow handoff layer did not observe graph start")
|
||||
if self._root_node_id is None:
|
||||
raise RuntimeError("Workflow handoff layer has no execution root node id")
|
||||
segment_execution_seconds = max(
|
||||
self._monotonic_clock() - self._active_segment_started_at,
|
||||
0.0,
|
||||
)
|
||||
|
||||
context = WorkflowResumptionContext.from_runtime_snapshot(
|
||||
generate_entity=self._generate_entity,
|
||||
serialized_graph_runtime_state=self.graph_runtime_state.dumps(),
|
||||
serialized_response_stream_filter_state=self._response_stream_filter.dumps(),
|
||||
active_execution_seconds=(self._prior_active_execution_seconds + segment_execution_seconds),
|
||||
root_node_id=self._root_node_id,
|
||||
)
|
||||
handoff = self._handoff_service.create_prepared_from_state(
|
||||
workflow_run_id=workflow_run_id,
|
||||
task_id=self._generate_entity.task_id,
|
||||
serialized_state=context.dumps(),
|
||||
resume_route=self._resume_route,
|
||||
source_worker_id=self._source_worker_id,
|
||||
rag_group_metadata=self._rag_group_metadata(),
|
||||
)
|
||||
if handoff.state != WorkflowHandoffState.PREPARED:
|
||||
raise RuntimeError(
|
||||
f"Workflow handoff was persisted in non-resumable state: "
|
||||
f"handoff_id={handoff.id}, state={handoff.state}"
|
||||
)
|
||||
self._persisted_handoff = handoff
|
||||
except Exception as error:
|
||||
# Do not rely on raising here: Graphon 0.6 catches layer exceptions.
|
||||
# The runner must call require_persisted_handoff after receiving the
|
||||
# maintenance pause event and fail closed when this error is present.
|
||||
self._persistence_error = error
|
||||
logger.exception("Failed to persist workflow handoff checkpoint")
|
||||
|
||||
def _rag_group_metadata(self) -> RagPipelineHandoffGroupMetadata | None:
|
||||
if not isinstance(self._generate_entity, RagPipelineGenerateEntity):
|
||||
return None
|
||||
extras = self._generate_entity.extras
|
||||
source_batch_id = extras.get(RAG_PIPELINE_SOURCE_BATCH_ID_EXTRA_KEY)
|
||||
tenant_id = extras.get(RAG_PIPELINE_TENANT_ID_EXTRA_KEY)
|
||||
queue_kind = extras.get(RAG_PIPELINE_QUEUE_KIND_EXTRA_KEY)
|
||||
tenant_isolated = extras.get(RAG_PIPELINE_TENANT_ISOLATED_EXTRA_KEY)
|
||||
if source_batch_id is None and tenant_id is None and queue_kind is None and tenant_isolated is None:
|
||||
return None
|
||||
if not all(
|
||||
isinstance(value, str) and value for value in (source_batch_id, tenant_id, queue_kind)
|
||||
) or not isinstance(tenant_isolated, bool):
|
||||
raise ValueError("RAG pipeline handoff group metadata is incomplete")
|
||||
assert isinstance(source_batch_id, str)
|
||||
assert isinstance(tenant_id, str)
|
||||
assert isinstance(queue_kind, str)
|
||||
return RagPipelineHandoffGroupMetadata(
|
||||
source_batch_id=source_batch_id,
|
||||
tenant_id=tenant_id,
|
||||
queue_kind=RagPipelineQueueKind(queue_kind),
|
||||
document_id=self._generate_entity.document_id,
|
||||
dataset_id=self._generate_entity.dataset_id,
|
||||
tenant_isolated=tenant_isolated,
|
||||
)
|
||||
|
||||
@override
|
||||
def on_graph_end(self, error: Exception | None) -> None:
|
||||
_ = error
|
||||
|
||||
def require_persisted_handoff(self) -> WorkflowRunHandoff:
|
||||
"""Return the durable row or raise so the caller can fail closed."""
|
||||
if not self._maintenance_pause_observed:
|
||||
raise WorkflowHandoffNotObservedError("Worker-drain pause was not observed")
|
||||
if self._persistence_error is not None:
|
||||
raise WorkflowHandoffPersistenceError(
|
||||
"Workflow handoff checkpoint was not persisted"
|
||||
) from self._persistence_error
|
||||
if self._persisted_handoff is None:
|
||||
raise WorkflowHandoffPersistenceError("Workflow handoff checkpoint result is missing")
|
||||
return self._persisted_handoff
|
||||
|
||||
|
||||
def infer_workflow_handoff_resume_route(
|
||||
generate_entity: ResumableWorkflowGenerateEntity,
|
||||
) -> WorkflowHandoffResumeRoute:
|
||||
"""Infer the standard route; triggered workflows explicitly override it."""
|
||||
if isinstance(generate_entity, RagPipelineGenerateEntity):
|
||||
return WorkflowHandoffResumeRoute.RAG_PIPELINE
|
||||
if isinstance(generate_entity, AdvancedChatAppGenerateEntity):
|
||||
return WorkflowHandoffResumeRoute.ADVANCED_CHAT
|
||||
return WorkflowHandoffResumeRoute.WORKFLOW
|
||||
|
||||
|
||||
def create_workflow_handoff_persistence_layer(
|
||||
*,
|
||||
config: WorkflowHandoffLayerConfig,
|
||||
generate_entity: ResumableWorkflowGenerateEntity,
|
||||
response_stream_filter: ResponseStreamFilter,
|
||||
) -> WorkflowHandoffPersistenceLayer:
|
||||
"""Construct an injectable handoff layer while keeping generator wiring thin."""
|
||||
return WorkflowHandoffPersistenceLayer(
|
||||
handoff_service=config.handoff_service,
|
||||
generate_entity=generate_entity,
|
||||
source_worker_id=config.source_worker_id,
|
||||
response_stream_filter=response_stream_filter,
|
||||
resume_route=config.resume_route,
|
||||
)
|
||||
|
||||
|
||||
__all__ = [
|
||||
"ResumableWorkflowGenerateEntity",
|
||||
"WorkflowHandoffLayerConfig",
|
||||
"WorkflowHandoffNotObservedError",
|
||||
"WorkflowHandoffPersistenceError",
|
||||
"WorkflowHandoffPersistenceLayer",
|
||||
"create_workflow_handoff_persistence_layer",
|
||||
"infer_workflow_handoff_resume_route",
|
||||
]
|
||||
@@ -0,0 +1,138 @@
|
||||
import logging
|
||||
from collections.abc import Callable
|
||||
from datetime import datetime
|
||||
from typing import override
|
||||
|
||||
from graphon.entities import WorkflowStartReason
|
||||
from graphon.graph_engine.entities.commands import AbortCommand
|
||||
from graphon.graph_engine.layers import GraphEngineLayer
|
||||
from graphon.graph_events import GraphEngineEvent, GraphRunStartedEvent
|
||||
from libs.datetime_utils import naive_utc_now
|
||||
from models.workflow_handoff import WorkflowHandoffState, WorkflowRunHandoff
|
||||
from repositories.workflow_handoff_repository import WorkflowRunHandoffRepository
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
WORKFLOW_HANDOFF_ACKNOWLEDGEMENT_ABORT_REASON = "Workflow handoff acknowledgement failed."
|
||||
|
||||
|
||||
class WorkflowHandoffAcknowledgementError(RuntimeError):
|
||||
"""Raised when a resumed graph did not durably acknowledge its handoff."""
|
||||
|
||||
|
||||
class WorkflowHandoffAcknowledgementNotObservedError(RuntimeError):
|
||||
"""Raised when the graph did not emit the expected resumption start event."""
|
||||
|
||||
|
||||
class WorkflowHandoffResumeAcknowledgementLayer(GraphEngineLayer):
|
||||
"""Complete a claimed handoff immediately before resumed nodes can run.
|
||||
|
||||
Graphon yields ``GraphRunStartedEvent`` before it starts the worker pool. The
|
||||
layer marks the old generation ``RESUMED`` while handling that event, which
|
||||
permits another planned drain to create the next generation. The runner must
|
||||
call :meth:`require_acknowledged` when it receives the start event; Graphon
|
||||
logs and swallows layer exceptions by design.
|
||||
|
||||
If acknowledgement fails, an Abort command is also queued as a second safety
|
||||
net so execution cannot continue if a caller accidentally misses the explicit
|
||||
check.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
repository: WorkflowRunHandoffRepository,
|
||||
claimed_handoff: WorkflowRunHandoff,
|
||||
clock: Callable[[], datetime] = naive_utc_now,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
if claimed_handoff.state != WorkflowHandoffState.CLAIMED:
|
||||
raise ValueError(f"Workflow handoff is not claimed: {claimed_handoff.id}")
|
||||
if not claimed_handoff.lease_owner or not claimed_handoff.lease_token:
|
||||
raise ValueError(f"Workflow handoff claim identity is incomplete: {claimed_handoff.id}")
|
||||
|
||||
self._repository = repository
|
||||
self._handoff_id = claimed_handoff.id
|
||||
self._generation = claimed_handoff.generation
|
||||
self._lease_owner = claimed_handoff.lease_owner
|
||||
self._lease_token = claimed_handoff.lease_token
|
||||
self._clock = clock
|
||||
self._resumption_start_observed = False
|
||||
self._acknowledged = False
|
||||
self._acknowledgement_error: Exception | None = None
|
||||
|
||||
@property
|
||||
def resumption_start_observed(self) -> bool:
|
||||
return self._resumption_start_observed
|
||||
|
||||
@property
|
||||
def acknowledged(self) -> bool:
|
||||
return self._acknowledged
|
||||
|
||||
@property
|
||||
def acknowledgement_error(self) -> Exception | None:
|
||||
return self._acknowledgement_error
|
||||
|
||||
@override
|
||||
def on_graph_start(self) -> None:
|
||||
self._resumption_start_observed = False
|
||||
self._acknowledged = False
|
||||
self._acknowledgement_error = None
|
||||
|
||||
@override
|
||||
def on_event(self, event: GraphEngineEvent) -> None:
|
||||
if not isinstance(event, GraphRunStartedEvent) or event.reason != WorkflowStartReason.RESUMPTION:
|
||||
return
|
||||
if self._resumption_start_observed:
|
||||
return
|
||||
|
||||
self._resumption_start_observed = True
|
||||
try:
|
||||
resumed_at = self._clock()
|
||||
if not self._repository.mark_resumed(
|
||||
handoff_id=self._handoff_id,
|
||||
generation=self._generation,
|
||||
lease_owner=self._lease_owner,
|
||||
lease_token=self._lease_token,
|
||||
resumed_at=resumed_at,
|
||||
):
|
||||
raise RuntimeError(
|
||||
f"Workflow handoff claim is no longer current: "
|
||||
f"handoff_id={self._handoff_id}, generation={self._generation}"
|
||||
)
|
||||
self._acknowledged = True
|
||||
except Exception as error:
|
||||
self._acknowledgement_error = error
|
||||
logger.exception("Failed to acknowledge resumed workflow handoff")
|
||||
self._abort_resumed_graph()
|
||||
|
||||
@override
|
||||
def on_graph_end(self, error: Exception | None) -> None:
|
||||
_ = error
|
||||
|
||||
def require_acknowledged(self) -> None:
|
||||
"""Fail closed before Graphon starts any resumed nodes."""
|
||||
if not self._resumption_start_observed:
|
||||
raise WorkflowHandoffAcknowledgementNotObservedError("Workflow handoff resumption start was not observed")
|
||||
if self._acknowledgement_error is not None:
|
||||
raise WorkflowHandoffAcknowledgementError(
|
||||
"Workflow handoff resumption was not acknowledged"
|
||||
) from self._acknowledgement_error
|
||||
if not self._acknowledged:
|
||||
raise WorkflowHandoffAcknowledgementError("Workflow handoff acknowledgement result is missing")
|
||||
|
||||
def _abort_resumed_graph(self) -> None:
|
||||
if self.command_channel is None:
|
||||
return
|
||||
try:
|
||||
self.command_channel.send_command(AbortCommand(reason=WORKFLOW_HANDOFF_ACKNOWLEDGEMENT_ABORT_REASON))
|
||||
except Exception:
|
||||
logger.exception("Failed to abort graph after workflow handoff acknowledgement failure")
|
||||
|
||||
|
||||
__all__ = [
|
||||
"WORKFLOW_HANDOFF_ACKNOWLEDGEMENT_ABORT_REASON",
|
||||
"WorkflowHandoffAcknowledgementError",
|
||||
"WorkflowHandoffAcknowledgementNotObservedError",
|
||||
"WorkflowHandoffResumeAcknowledgementLayer",
|
||||
]
|
||||
@@ -14,6 +14,7 @@ from dataclasses import dataclass
|
||||
from datetime import datetime
|
||||
from typing import Any, Union, override
|
||||
|
||||
from core.app.apps.workflow.command_channels import is_workflow_warm_shutdown_pause
|
||||
from core.app.entities.app_invoke_entities import AdvancedChatAppGenerateEntity, WorkflowAppGenerateEntity
|
||||
from core.app.workflow.retry_history import RETRY_HISTORY_PROCESS_DATA_KEY, WorkflowNodeRetryAttempt
|
||||
from core.helper.trace_id_helper import ParentTraceContext
|
||||
@@ -24,7 +25,7 @@ from core.workflow.node_execution_process_data import preserve_workflow_agent_bi
|
||||
from core.workflow.system_variables import SystemVariableKey
|
||||
from core.workflow.variable_prefixes import SYSTEM_VARIABLE_NODE_ID
|
||||
from core.workflow.workflow_run_outputs import project_node_outputs_for_workflow_run
|
||||
from graphon.entities import WorkflowExecution, WorkflowNodeExecution
|
||||
from graphon.entities import WorkflowExecution, WorkflowNodeExecution, WorkflowStartReason
|
||||
from graphon.enums import (
|
||||
BuiltinNodeTypes,
|
||||
WorkflowExecutionStatus,
|
||||
@@ -118,7 +119,7 @@ class WorkflowPersistenceLayer(GraphEngineLayer):
|
||||
def on_event(self, event: GraphEngineEvent) -> None:
|
||||
match event:
|
||||
case GraphRunStartedEvent():
|
||||
self._handle_graph_run_started()
|
||||
self._handle_graph_run_started(event)
|
||||
case GraphRunSucceededEvent():
|
||||
self._handle_graph_run_succeeded(event)
|
||||
case GraphRunPartialSucceededEvent():
|
||||
@@ -149,8 +150,12 @@ class WorkflowPersistenceLayer(GraphEngineLayer):
|
||||
# ------------------------------------------------------------------
|
||||
# Graph-level handlers
|
||||
# ------------------------------------------------------------------
|
||||
def _handle_graph_run_started(self) -> None:
|
||||
def _handle_graph_run_started(self, event: GraphRunStartedEvent | None = None) -> None:
|
||||
execution_id = self._get_execution_id()
|
||||
# Every handoff creates a fresh layer instance. Continue the logical
|
||||
# run's persisted node sequence instead of reusing indices from zero.
|
||||
if event is not None and event.reason == WorkflowStartReason.RESUMPTION:
|
||||
self._node_sequence = self._workflow_node_execution_repository.get_max_index(execution_id)
|
||||
workflow_execution = WorkflowExecution.new(
|
||||
id_=execution_id,
|
||||
workflow_id=self._workflow_info.workflow_id,
|
||||
@@ -209,6 +214,12 @@ class WorkflowPersistenceLayer(GraphEngineLayer):
|
||||
_inspector_publish_workflow_completed(workflow_run_id=execution.id_, status=str(execution.status.value))
|
||||
|
||||
def _handle_graph_run_paused(self, event: GraphRunPausedEvent) -> None:
|
||||
if is_workflow_warm_shutdown_pause(event.reasons):
|
||||
# A maintenance pause ends only this worker's execution segment.
|
||||
# The durable handoff resumes the same logical run on another
|
||||
# worker, so keep its externally visible state as RUNNING.
|
||||
return
|
||||
|
||||
execution = self._get_workflow_execution()
|
||||
execution.status = WorkflowExecutionStatus.PAUSED
|
||||
execution.outputs = event.outputs
|
||||
|
||||
+219
-16
@@ -1,7 +1,9 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
import json
|
||||
from collections.abc import Sequence
|
||||
from enum import StrEnum
|
||||
from typing import Any
|
||||
|
||||
from pydantic import BaseModel, ValidationError
|
||||
@@ -9,6 +11,79 @@ from pydantic import BaseModel, ValidationError
|
||||
from extensions.ext_redis import redis_client
|
||||
|
||||
_DEFAULT_TASK_TTL = 60 * 60 # 1 hour
|
||||
_CLAIM_EMPTY_VALUE = "empty:"
|
||||
_CLAIM_TASK_PREFIX = "task:"
|
||||
_DISPATCH_LEASE_PREFIX = "lease:"
|
||||
_DISPATCH_DONE_VALUE = "done"
|
||||
|
||||
_ENQUEUE_OR_ACQUIRE_SCRIPT = """
|
||||
if redis.call('EXISTS', KEYS[2]) == 1 then
|
||||
redis.call('LPUSH', KEYS[1], ARGV[2])
|
||||
return 0
|
||||
end
|
||||
redis.call('SETEX', KEYS[2], ARGV[1], '1')
|
||||
return 1
|
||||
"""
|
||||
|
||||
_CLAIM_TASK_ONCE_SCRIPT = """
|
||||
local existing = redis.call('GET', KEYS[3])
|
||||
if existing then
|
||||
return existing
|
||||
end
|
||||
local task = redis.call('RPOP', KEYS[1])
|
||||
if task then
|
||||
local claimed = ARGV[2] .. task
|
||||
redis.call('SETEX', KEYS[3], ARGV[1], claimed)
|
||||
redis.call('SETEX', KEYS[2], ARGV[1], '1')
|
||||
return claimed
|
||||
end
|
||||
redis.call('SETEX', KEYS[3], ARGV[1], ARGV[3])
|
||||
redis.call('DEL', KEYS[2])
|
||||
return ARGV[3]
|
||||
"""
|
||||
|
||||
_CLAIM_DISPATCH_SCRIPT = """
|
||||
local current = redis.call('GET', KEYS[1])
|
||||
if not current then
|
||||
redis.call('SETEX', KEYS[1], ARGV[1], ARGV[2])
|
||||
return 1
|
||||
end
|
||||
if current == ARGV[3] then
|
||||
return 2
|
||||
end
|
||||
if current == ARGV[2] then
|
||||
redis.call('EXPIRE', KEYS[1], ARGV[1])
|
||||
return 1
|
||||
end
|
||||
return 0
|
||||
"""
|
||||
|
||||
_RENEW_DISPATCH_SCRIPT = """
|
||||
if redis.call('GET', KEYS[1]) == ARGV[2] then
|
||||
redis.call('EXPIRE', KEYS[1], ARGV[1])
|
||||
return 1
|
||||
end
|
||||
return 0
|
||||
"""
|
||||
|
||||
_COMPLETE_DISPATCH_SCRIPT = """
|
||||
local current = redis.call('GET', KEYS[1])
|
||||
if current == ARGV[2] then
|
||||
redis.call('SETEX', KEYS[1], ARGV[1], ARGV[3])
|
||||
return 1
|
||||
end
|
||||
if current == ARGV[3] then
|
||||
redis.call('EXPIRE', KEYS[1], ARGV[1])
|
||||
return 1
|
||||
end
|
||||
return 0
|
||||
"""
|
||||
|
||||
|
||||
class TenantTaskDispatchClaimOutcome(StrEnum):
|
||||
BUSY = "busy"
|
||||
ACQUIRED = "acquired"
|
||||
DONE = "done"
|
||||
|
||||
|
||||
class TaskWrapper(BaseModel):
|
||||
@@ -35,6 +110,12 @@ class TenantIsolatedTaskQueue:
|
||||
self._queue = f"tenant_self_{unique_key}_task_queue:{tenant_id}"
|
||||
self._task_key = f"tenant_{unique_key}_task:{tenant_id}"
|
||||
|
||||
def _dispatch_key(self, dispatch_token: str) -> str:
|
||||
if not dispatch_token:
|
||||
raise ValueError("dispatch_token must not be empty")
|
||||
token_digest = hashlib.sha256(dispatch_token.encode()).hexdigest()
|
||||
return f"tenant_{self._unique_key}_task_dispatch:{self._tenant_id}:{token_digest}"
|
||||
|
||||
def get_task_key(self):
|
||||
return redis_client.get(self._task_key)
|
||||
|
||||
@@ -45,22 +126,134 @@ class TenantIsolatedTaskQueue:
|
||||
redis_client.delete(self._task_key)
|
||||
|
||||
def push_tasks(self, tasks: Sequence[Any]):
|
||||
serialized_tasks = []
|
||||
for task in tasks:
|
||||
# Store str list directly, maintaining full compatibility for pipeline scenarios
|
||||
if isinstance(task, str):
|
||||
serialized_tasks.append(task)
|
||||
else:
|
||||
# Use TaskWrapper to do JSON serialization for non-string tasks
|
||||
wrapper = TaskWrapper(data=task)
|
||||
serialized_data = wrapper.serialize()
|
||||
serialized_tasks.append(serialized_data)
|
||||
serialized_tasks = [self._serialize_task(task) for task in tasks]
|
||||
|
||||
if not serialized_tasks:
|
||||
return
|
||||
|
||||
redis_client.lpush(self._queue, *serialized_tasks)
|
||||
|
||||
def enqueue_or_acquire(self, task: Any, ttl: int = _DEFAULT_TASK_TTL) -> bool:
|
||||
"""Atomically enqueue behind an owner or acquire the idle tenant slot.
|
||||
|
||||
Returns True when the caller acquired the slot and must dispatch the
|
||||
task itself. Returns False when the task was appended to the wait queue.
|
||||
"""
|
||||
if ttl <= 0:
|
||||
raise ValueError("ttl must be positive")
|
||||
acquired = redis_client.eval(
|
||||
_ENQUEUE_OR_ACQUIRE_SCRIPT,
|
||||
2,
|
||||
self._queue,
|
||||
self._task_key,
|
||||
ttl,
|
||||
self._serialize_task(task),
|
||||
)
|
||||
return bool(acquired)
|
||||
|
||||
def claim_task_once(self, *, claim_key: str, ttl: int = _DEFAULT_TASK_TTL) -> tuple[bool, Any | None]:
|
||||
"""Atomically claim one queued task once for a retryable release owner.
|
||||
|
||||
The claim is stored before the caller dispatches it. A retry receives
|
||||
the same task instead of consuming another tenant slot.
|
||||
"""
|
||||
if not claim_key:
|
||||
raise ValueError("claim_key must not be empty")
|
||||
if ttl <= 0:
|
||||
raise ValueError("ttl must be positive")
|
||||
claimed = redis_client.eval(
|
||||
_CLAIM_TASK_ONCE_SCRIPT,
|
||||
3,
|
||||
self._queue,
|
||||
self._task_key,
|
||||
claim_key,
|
||||
ttl,
|
||||
_CLAIM_TASK_PREFIX,
|
||||
_CLAIM_EMPTY_VALUE,
|
||||
)
|
||||
if isinstance(claimed, bytes):
|
||||
claimed = claimed.decode("utf-8")
|
||||
if claimed == _CLAIM_EMPTY_VALUE:
|
||||
return False, None
|
||||
if not isinstance(claimed, str) or not claimed.startswith(_CLAIM_TASK_PREFIX):
|
||||
raise ValueError("invalid tenant queue claim payload")
|
||||
return True, self._deserialize_task(claimed.removeprefix(_CLAIM_TASK_PREFIX))
|
||||
|
||||
def claim_dispatch(
|
||||
self,
|
||||
*,
|
||||
dispatch_token: str,
|
||||
owner: str,
|
||||
lease_ttl: int = _DEFAULT_TASK_TTL,
|
||||
) -> TenantTaskDispatchClaimOutcome:
|
||||
"""Claim a dispatched queue item without executing duplicate messages.
|
||||
|
||||
The dispatch token identifies the logical queue item while ``owner``
|
||||
identifies one Celery delivery. The lease is renewable so a worker
|
||||
loss can be recovered without allowing a concurrently delivered copy
|
||||
to execute the same source batch.
|
||||
"""
|
||||
if not owner:
|
||||
raise ValueError("owner must not be empty")
|
||||
if lease_ttl <= 0:
|
||||
raise ValueError("lease_ttl must be positive")
|
||||
result = redis_client.eval(
|
||||
_CLAIM_DISPATCH_SCRIPT,
|
||||
1,
|
||||
self._dispatch_key(dispatch_token),
|
||||
lease_ttl,
|
||||
f"{_DISPATCH_LEASE_PREFIX}{owner}",
|
||||
_DISPATCH_DONE_VALUE,
|
||||
)
|
||||
if result == 1:
|
||||
return TenantTaskDispatchClaimOutcome.ACQUIRED
|
||||
if result == 2:
|
||||
return TenantTaskDispatchClaimOutcome.DONE
|
||||
return TenantTaskDispatchClaimOutcome.BUSY
|
||||
|
||||
def renew_dispatch_claim(
|
||||
self,
|
||||
*,
|
||||
dispatch_token: str,
|
||||
owner: str,
|
||||
lease_ttl: int = _DEFAULT_TASK_TTL,
|
||||
) -> bool:
|
||||
if not owner:
|
||||
raise ValueError("owner must not be empty")
|
||||
if lease_ttl <= 0:
|
||||
raise ValueError("lease_ttl must be positive")
|
||||
return bool(
|
||||
redis_client.eval(
|
||||
_RENEW_DISPATCH_SCRIPT,
|
||||
1,
|
||||
self._dispatch_key(dispatch_token),
|
||||
lease_ttl,
|
||||
f"{_DISPATCH_LEASE_PREFIX}{owner}",
|
||||
)
|
||||
)
|
||||
|
||||
def complete_dispatch_claim(
|
||||
self,
|
||||
*,
|
||||
dispatch_token: str,
|
||||
owner: str,
|
||||
done_ttl: int,
|
||||
) -> bool:
|
||||
if not owner:
|
||||
raise ValueError("owner must not be empty")
|
||||
if done_ttl <= 0:
|
||||
raise ValueError("done_ttl must be positive")
|
||||
return bool(
|
||||
redis_client.eval(
|
||||
_COMPLETE_DISPATCH_SCRIPT,
|
||||
1,
|
||||
self._dispatch_key(dispatch_token),
|
||||
done_ttl,
|
||||
f"{_DISPATCH_LEASE_PREFIX}{owner}",
|
||||
_DISPATCH_DONE_VALUE,
|
||||
)
|
||||
)
|
||||
|
||||
def pull_tasks(self, count: int = 1) -> Sequence[Any]:
|
||||
if count <= 0:
|
||||
return []
|
||||
@@ -74,11 +267,21 @@ class TenantIsolatedTaskQueue:
|
||||
if isinstance(serialized_task, bytes):
|
||||
serialized_task = serialized_task.decode("utf-8")
|
||||
|
||||
try:
|
||||
wrapper = TaskWrapper.deserialize(serialized_task)
|
||||
tasks.append(wrapper.data)
|
||||
except (json.JSONDecodeError, ValidationError, TypeError, ValueError):
|
||||
# Fall back to raw string for legacy format or invalid JSON
|
||||
tasks.append(serialized_task)
|
||||
tasks.append(self._deserialize_task(serialized_task))
|
||||
|
||||
return tasks
|
||||
|
||||
@staticmethod
|
||||
def _serialize_task(task: Any) -> str:
|
||||
# Store strings directly, maintaining full compatibility for pipeline scenarios.
|
||||
if isinstance(task, str):
|
||||
return task
|
||||
return TaskWrapper(data=task).serialize()
|
||||
|
||||
@staticmethod
|
||||
def _deserialize_task(serialized_task: str) -> Any:
|
||||
try:
|
||||
return TaskWrapper.deserialize(serialized_task).data
|
||||
except (json.JSONDecodeError, ValidationError, TypeError, ValueError):
|
||||
# Fall back to raw string for legacy format or invalid JSON.
|
||||
return serialized_task
|
||||
|
||||
@@ -219,3 +219,10 @@ class CeleryWorkflowNodeExecutionRepository(WorkflowNodeExecutionRepository):
|
||||
workflow_execution_id,
|
||||
)
|
||||
return []
|
||||
|
||||
@override
|
||||
def get_max_index(self, workflow_execution_id: str) -> int:
|
||||
# A handoff resumes in a new worker with a fresh in-memory cache. Read
|
||||
# the durable SQL state so the next segment continues the logical run's
|
||||
# node sequence instead of starting again at one.
|
||||
return self._sql_repository.get_max_index(workflow_execution_id)
|
||||
|
||||
@@ -45,6 +45,10 @@ class WorkflowNodeExecutionRepository(Protocol):
|
||||
order_config: OrderConfig | None = None,
|
||||
) -> Sequence[WorkflowNodeExecution]: ...
|
||||
|
||||
def get_max_index(self, workflow_execution_id: str) -> int:
|
||||
"""Return the greatest persisted node sequence for a logical run."""
|
||||
...
|
||||
|
||||
|
||||
class RepositoryImportError(Exception):
|
||||
"""Raised when a repository implementation cannot be imported or instantiated."""
|
||||
|
||||
@@ -13,6 +13,7 @@ from core.repositories.factory import WorkflowExecutionRepository
|
||||
from graphon.entities import WorkflowExecution
|
||||
from graphon.enums import WorkflowExecutionStatus, WorkflowType
|
||||
from graphon.workflow_type_encoder import WorkflowRuntimeTypeConverter
|
||||
from libs.datetime_utils import naive_utc_now
|
||||
from models import (
|
||||
Account,
|
||||
CreatorUserRole,
|
||||
@@ -23,6 +24,15 @@ from models.enums import WorkflowRunTriggeredFrom
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_TERMINAL_WORKFLOW_EXECUTION_STATUSES = frozenset(
|
||||
{
|
||||
WorkflowExecutionStatus.SUCCEEDED,
|
||||
WorkflowExecutionStatus.FAILED,
|
||||
WorkflowExecutionStatus.STOPPED,
|
||||
WorkflowExecutionStatus.PARTIAL_SUCCEEDED,
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
class SQLAlchemyWorkflowExecutionRepository(WorkflowExecutionRepository):
|
||||
"""
|
||||
@@ -201,8 +211,41 @@ class SQLAlchemyWorkflowExecutionRepository(WorkflowExecutionRepository):
|
||||
if existing_model:
|
||||
if existing_model.tenant_id != self._tenant_id:
|
||||
raise ValueError("Unauthorized access to workflow run")
|
||||
# Preserve the original start time for pause/resume flows.
|
||||
if existing_model.status in _TERMINAL_WORKFLOW_EXECUTION_STATUSES:
|
||||
# Terminal status is monotonic. In particular, a resumed
|
||||
# segment can emit GraphRunStarted after Stop has already
|
||||
# atomically failed its CLAIMED handoff. Never let that
|
||||
# late RUNNING write resurrect the logical workflow run.
|
||||
logger.info(
|
||||
"Ignoring late workflow execution update for terminal run: workflow_run_id=%s, "
|
||||
"existing_status=%s, requested_status=%s",
|
||||
existing_model.id,
|
||||
existing_model.status,
|
||||
db_model.status,
|
||||
)
|
||||
self._execution_cache[existing_model.id] = existing_model
|
||||
return
|
||||
# A resumed graph segment has its own domain ``started_at``, but
|
||||
# ``WorkflowRun`` represents the complete logical run. Preserve
|
||||
# the original start and report wall-clock elapsed time so the
|
||||
# user-visible duration includes maintenance handoff waits.
|
||||
# GraphRuntimeState remains unchanged, so execution timeouts and
|
||||
# token billing continue to exclude time spent between workers.
|
||||
db_model.created_at = existing_model.created_at
|
||||
if db_model.finished_at is not None:
|
||||
db_model.elapsed_time = max(
|
||||
(db_model.finished_at - existing_model.created_at).total_seconds(),
|
||||
0.0,
|
||||
)
|
||||
elif db_model.status == WorkflowExecutionStatus.PAUSED:
|
||||
db_model.elapsed_time = max(
|
||||
(naive_utc_now() - existing_model.created_at).total_seconds(),
|
||||
0.0,
|
||||
)
|
||||
else:
|
||||
# Starting another segment must not erase timing already
|
||||
# captured by a previous user-visible pause.
|
||||
db_model.elapsed_time = existing_model.elapsed_time
|
||||
|
||||
# SQLAlchemy merge intelligently handles both insert and update operations
|
||||
# based on the presence of the primary key
|
||||
|
||||
@@ -10,7 +10,7 @@ from concurrent.futures import ThreadPoolExecutor
|
||||
from typing import Any, override
|
||||
|
||||
import psycopg2.errors
|
||||
from sqlalchemy import UnaryExpression, asc, desc, select
|
||||
from sqlalchemy import UnaryExpression, asc, desc, func, select
|
||||
from sqlalchemy.engine import Engine
|
||||
from sqlalchemy.exc import IntegrityError
|
||||
from sqlalchemy.orm import sessionmaker
|
||||
@@ -571,6 +571,18 @@ class SQLAlchemyWorkflowNodeExecutionRepository(WorkflowNodeExecutionRepository)
|
||||
|
||||
return list(domain_models)
|
||||
|
||||
@override
|
||||
def get_max_index(self, workflow_execution_id: str) -> int:
|
||||
"""Seed a resumed segment after every node already persisted for the run."""
|
||||
with self._session_factory() as session:
|
||||
value = session.scalar(
|
||||
select(func.max(WorkflowNodeExecutionModel.index)).where(
|
||||
WorkflowNodeExecutionModel.workflow_run_id == workflow_execution_id,
|
||||
WorkflowNodeExecutionModel.tenant_id == self._tenant_id,
|
||||
)
|
||||
)
|
||||
return int(value or 0)
|
||||
|
||||
|
||||
def _deterministic_json_dump(value: Mapping[str, Any]) -> str:
|
||||
return json.dumps(value, sort_keys=True)
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
import logging
|
||||
import time
|
||||
from collections.abc import Generator, Mapping, Sequence
|
||||
from typing import Any, TypedDict
|
||||
from typing import Any, Protocol, TypedDict, runtime_checkable
|
||||
|
||||
from configs import dify_config
|
||||
from context import capture_current_context
|
||||
@@ -46,6 +46,12 @@ logger = logging.getLogger(__name__)
|
||||
_file_access_controller = DatabaseFileAccessController()
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class _NodeRunStepsState(Protocol):
|
||||
@property
|
||||
def node_run_steps(self) -> object: ...
|
||||
|
||||
|
||||
def iter_dify_graph_engine_events(
|
||||
engine: GraphEngine,
|
||||
response_stream_filter: ResponseStreamFilter | None = None,
|
||||
@@ -176,6 +182,7 @@ class WorkflowEntry:
|
||||
graph_runtime_state: GraphRuntimeState,
|
||||
command_channel: CommandChannel | None = None,
|
||||
response_stream_filter: ResponseStreamFilter | None = None,
|
||||
prior_active_execution_seconds: float = 0.0,
|
||||
) -> None:
|
||||
"""
|
||||
Init workflow entry
|
||||
@@ -200,6 +207,8 @@ class WorkflowEntry:
|
||||
workflow_call_max_depth = dify_config.WORKFLOW_CALL_MAX_DEPTH
|
||||
if call_depth > workflow_call_max_depth:
|
||||
raise ValueError(f"Max workflow call depth {workflow_call_max_depth} reached.")
|
||||
if prior_active_execution_seconds < 0:
|
||||
raise ValueError("prior_active_execution_seconds must be non-negative")
|
||||
|
||||
# Use provided command channel or default to InMemoryChannel
|
||||
if command_channel is None:
|
||||
@@ -237,9 +246,18 @@ class WorkflowEntry:
|
||||
self.graph_engine.layer(debug_layer)
|
||||
|
||||
# Add execution limits layer
|
||||
limits_layer = ExecutionLimitsLayer(
|
||||
max_steps=dify_config.WORKFLOW_MAX_EXECUTION_STEPS, max_time=dify_config.WORKFLOW_MAX_EXECUTION_TIME
|
||||
node_run_steps = 0
|
||||
if isinstance(graph_runtime_state, _NodeRunStepsState) and isinstance(graph_runtime_state.node_run_steps, int):
|
||||
node_run_steps = graph_runtime_state.node_run_steps
|
||||
remaining_steps = max(
|
||||
dify_config.WORKFLOW_MAX_EXECUTION_STEPS - node_run_steps,
|
||||
0,
|
||||
)
|
||||
remaining_execution_seconds = max(
|
||||
int(dify_config.WORKFLOW_MAX_EXECUTION_TIME - prior_active_execution_seconds),
|
||||
0,
|
||||
)
|
||||
limits_layer = ExecutionLimitsLayer(max_steps=remaining_steps, max_time=remaining_execution_seconds)
|
||||
self.graph_engine.layer(limits_layer)
|
||||
self.graph_engine.layer(LLMQuotaLayer(tenant_id=tenant_id))
|
||||
|
||||
|
||||
@@ -55,12 +55,27 @@ if [[ "${MODE}" == "worker" ]]; then
|
||||
echo "Using CELERY_WORKER_QUEUES: ${DEFAULT_QUEUES}"
|
||||
fi
|
||||
|
||||
# Recovery remains active after creation is disabled, so every upgraded
|
||||
# worker must keep consuming the capability-isolated handoff queue. Old
|
||||
# image revisions do not append this queue and therefore cannot steal N+1
|
||||
# scan/resume messages during a rolling update.
|
||||
WORKFLOW_HANDOFF_CAPABILITY_QUEUE="${WORKFLOW_HANDOFF_QUEUE:-workflow_handoff}"
|
||||
case ",${DEFAULT_QUEUES}," in
|
||||
*",${WORKFLOW_HANDOFF_CAPABILITY_QUEUE},"*) ;;
|
||||
*) DEFAULT_QUEUES="${DEFAULT_QUEUES},${WORKFLOW_HANDOFF_CAPABILITY_QUEUE}" ;;
|
||||
esac
|
||||
|
||||
if [[ -n "${CELERY_WORKER_CONCURRENCY}" ]]; then
|
||||
CONCURRENCY_OPTION="-c ${CELERY_WORKER_CONCURRENCY}"
|
||||
echo "Using CELERY_WORKER_CONCURRENCY: ${CELERY_WORKER_CONCURRENCY}"
|
||||
fi
|
||||
|
||||
WORKER_POOL="${CELERY_WORKER_POOL:-${CELERY_WORKER_CLASS:-gevent}}"
|
||||
if [[ "${WORKFLOW_HANDOFF_ENABLED,,}" = "true" && "${WORKER_POOL,,}" = "prefork" ]]; then
|
||||
echo "Error: WORKFLOW_HANDOFF_ENABLED requires a process-shared Celery pool; prefork is unsupported."
|
||||
echo "Use gevent, eventlet, threads, or solo so shutdown state and active-run tracking remain coherent."
|
||||
exit 1
|
||||
fi
|
||||
echo "Starting Celery worker with queues: ${DEFAULT_QUEUES}"
|
||||
|
||||
exec celery -A celery_entrypoint.celery worker -P ${WORKER_POOL} $CONCURRENCY_OPTION \
|
||||
@@ -130,6 +145,7 @@ else
|
||||
--worker-class ${SERVER_WORKER_CLASS:-geventwebsocket.gunicorn.workers.GeventWebSocketWorker} \
|
||||
--worker-connections ${SERVER_WORKER_CONNECTIONS:-10} \
|
||||
--timeout ${GUNICORN_TIMEOUT:-200} \
|
||||
--graceful-timeout ${GUNICORN_GRACEFUL_TIMEOUT:-660} \
|
||||
app:socketio_app
|
||||
fi
|
||||
fi
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
import ssl
|
||||
from datetime import timedelta
|
||||
from typing import Any
|
||||
from typing import Any, NotRequired
|
||||
|
||||
import pytz # type: ignore[import-untyped]
|
||||
from celery import Celery, Task
|
||||
@@ -35,6 +35,26 @@ class CelerySSLOptionsDict(TypedDict):
|
||||
class CeleryBeatScheduleEntry(TypedDict):
|
||||
task: str
|
||||
schedule: crontab | timedelta
|
||||
options: NotRequired[dict[str, Any]]
|
||||
|
||||
|
||||
def _register_workflow_handoff_schedule(
|
||||
*,
|
||||
imports: list[str],
|
||||
beat_schedule: dict[str, CeleryBeatScheduleEntry],
|
||||
) -> None:
|
||||
# Register recovery code even while creation is disabled. This is required
|
||||
# for the first, dormant rollout and ensures toggling the flag off never
|
||||
# strands durable rows created before the configuration change.
|
||||
imports.append("tasks.workflow_handoff_tasks")
|
||||
beat_schedule["workflow_handoff_scan"] = {
|
||||
"task": "workflow_handoff.scan",
|
||||
"schedule": timedelta(seconds=dify_config.WORKFLOW_HANDOFF_SCAN_INTERVAL_SECONDS),
|
||||
# This capability queue is consumed only by upgraded workers during
|
||||
# the dormant N -> N+1 rollout, so an older worker cannot discard the
|
||||
# unknown scanner task from the shared schedule_poller queue.
|
||||
"options": {"queue": dify_config.WORKFLOW_HANDOFF_QUEUE},
|
||||
}
|
||||
|
||||
|
||||
def _enqueue_initial_community_telemetry_heartbeat(sender: Any, **_: Any) -> None:
|
||||
@@ -180,6 +200,7 @@ def init_app(app: DifyApp) -> Celery:
|
||||
|
||||
# if you add a new task, please add the switch to CeleryScheduleTasksConfig
|
||||
beat_schedule: dict[str, CeleryBeatScheduleEntry] = {}
|
||||
_register_workflow_handoff_schedule(imports=imports, beat_schedule=beat_schedule)
|
||||
if dify_config.ENABLE_CLEAN_EMBEDDING_CACHE_TASK:
|
||||
imports.append("schedule.clean_embedding_cache_task")
|
||||
beat_schedule["clean_embedding_cache_task"] = {
|
||||
|
||||
@@ -415,3 +415,8 @@ class LogstoreWorkflowNodeExecutionRepository(WorkflowNodeExecutionRepository):
|
||||
workflow_execution_id,
|
||||
)
|
||||
raise
|
||||
|
||||
@override
|
||||
def get_max_index(self, workflow_execution_id: str) -> int:
|
||||
executions = self.get_by_workflow_execution(workflow_execution_id)
|
||||
return max((int(execution.index or 0) for execution in executions), default=0)
|
||||
|
||||
@@ -1,54 +1,248 @@
|
||||
"""Abort active workflow runs during Celery warm shutdown."""
|
||||
"""Coordinate active workflow runs during planned process shutdown."""
|
||||
|
||||
import logging
|
||||
import os
|
||||
import threading
|
||||
from typing import Any
|
||||
from collections.abc import Callable, Sequence
|
||||
from datetime import datetime
|
||||
from typing import Any, cast
|
||||
|
||||
from celery.signals import worker_shutdown, worker_shutting_down
|
||||
from sqlalchemy.engine import CursorResult
|
||||
from sqlalchemy.orm import Session, sessionmaker
|
||||
|
||||
from configs import dify_config
|
||||
from core.app.apps.workflow import command_channels as workflow_command_channels
|
||||
from core.app.apps.workflow.active_workflow_tasks import (
|
||||
ActiveWorkflowTask,
|
||||
get_active_workflow_task_count,
|
||||
get_active_workflow_tasks,
|
||||
reset_active_workflow_tasks,
|
||||
retain_active_workflow_tasks,
|
||||
)
|
||||
from libs.datetime_utils import naive_utc_now
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
WORKFLOW_WARM_SHUTDOWN_ABORT_REASON = "Workflow stopped because the worker is shutting down."
|
||||
WORKFLOW_WARM_SHUTDOWN_ABORT_REASON = workflow_command_channels.WORKFLOW_WARM_SHUTDOWN_ABORT_REASON
|
||||
WORKFLOW_WARM_SHUTDOWN_PAUSE_REASON = workflow_command_channels.WORKFLOW_WARM_SHUTDOWN_PAUSE_REASON
|
||||
WORKFLOW_WARM_SHUTDOWN_TIMEOUT_REASON = "Workflow stopped because the worker drain deadline expired."
|
||||
_WORKER_SHUTTING_DOWN_DISPATCH_UID = "dify.workflow_warm_shutdown.shutting_down"
|
||||
_WORKER_SHUTDOWN_DISPATCH_UID = "dify.workflow_warm_shutdown.shutdown"
|
||||
_celery_warm_shutdown_started = threading.Event()
|
||||
_workflow_warm_shutdown_started = threading.Event()
|
||||
# Keep the private alias temporarily because existing tests and out-of-tree
|
||||
# extensions may still clear it between worker initializations.
|
||||
_celery_warm_shutdown_started = _workflow_warm_shutdown_started
|
||||
_workflow_drain_watchdog_started = threading.Event()
|
||||
_workflow_drain_watchdog_lock = threading.Lock()
|
||||
|
||||
|
||||
def _is_warm_shutdown(how: Any) -> bool:
|
||||
return str(how).strip().lower() == "warm"
|
||||
|
||||
|
||||
def workflow_warm_shutdown_started() -> bool:
|
||||
"""Return whether this process has started a planned warm shutdown."""
|
||||
return _workflow_warm_shutdown_started.is_set()
|
||||
|
||||
|
||||
def celery_warm_shutdown_started() -> bool:
|
||||
"""Return whether the current worker process started Celery warm shutdown."""
|
||||
return _celery_warm_shutdown_started.is_set()
|
||||
"""Backward-compatible alias for workflow command-channel callers."""
|
||||
return workflow_warm_shutdown_started()
|
||||
|
||||
|
||||
def mark_workflow_warm_shutdown_started() -> None:
|
||||
"""Mark this process as draining workflow executions."""
|
||||
_workflow_warm_shutdown_started.set()
|
||||
|
||||
|
||||
def mark_celery_warm_shutdown_started() -> None:
|
||||
"""Mark the current worker process as being in Celery warm shutdown."""
|
||||
_celery_warm_shutdown_started.set()
|
||||
"""Backward-compatible alias for Celery integrations."""
|
||||
mark_workflow_warm_shutdown_started()
|
||||
|
||||
|
||||
def mark_workflow_runs_stopped_if_running_without_active_handoff(
|
||||
workflow_run_ids: Sequence[str],
|
||||
*,
|
||||
reason: str,
|
||||
now: datetime | None = None,
|
||||
session_maker: sessionmaker[Session] | None = None,
|
||||
) -> int:
|
||||
"""Conditionally stop runs that have no recoverable durable handoff.
|
||||
|
||||
A durable READY/CLAIMED handoff is excluded because another worker can
|
||||
recover it. The RUNNING predicate preserves any terminal state already
|
||||
written by a user Stop or by normal workflow completion.
|
||||
"""
|
||||
normalized_run_ids = sorted({workflow_run_id for workflow_run_id in workflow_run_ids if workflow_run_id})
|
||||
if not normalized_run_ids:
|
||||
return 0
|
||||
if not reason:
|
||||
raise ValueError("reason must not be empty")
|
||||
|
||||
# Resolve the process-global session maker lazily. It is configured after
|
||||
# Celery initialization but before a worker can receive a shutdown signal,
|
||||
# and unlike ``db.engine`` it does not require a Flask app context here.
|
||||
from sqlalchemy import exists, select, update
|
||||
|
||||
from graphon.enums import WorkflowExecutionStatus
|
||||
from models.workflow import WorkflowRun
|
||||
from models.workflow_handoff import WorkflowHandoffState, WorkflowRunHandoff
|
||||
|
||||
active_handoff_exists = exists(
|
||||
select(WorkflowRunHandoff.id).where(
|
||||
WorkflowRunHandoff.workflow_run_id == WorkflowRun.id,
|
||||
WorkflowRunHandoff.state.in_((WorkflowHandoffState.READY, WorkflowHandoffState.CLAIMED)),
|
||||
)
|
||||
)
|
||||
statement = (
|
||||
update(WorkflowRun)
|
||||
.where(
|
||||
WorkflowRun.id.in_(normalized_run_ids),
|
||||
WorkflowRun.status == WorkflowExecutionStatus.RUNNING,
|
||||
~active_handoff_exists,
|
||||
)
|
||||
.values(
|
||||
status=WorkflowExecutionStatus.STOPPED,
|
||||
error=reason,
|
||||
finished_at=now or naive_utc_now(),
|
||||
)
|
||||
)
|
||||
|
||||
if session_maker is None:
|
||||
from core.db.session_factory import session_factory
|
||||
|
||||
session_maker = session_factory.get_session_maker()
|
||||
|
||||
with session_maker.begin() as session:
|
||||
result = session.execute(statement)
|
||||
return max(cast(CursorResult, result).rowcount or 0, 0)
|
||||
|
||||
|
||||
def _mark_timed_out_workflow_runs_stopped(
|
||||
registrations: Sequence[ActiveWorkflowTask],
|
||||
*,
|
||||
now: datetime | None = None,
|
||||
session_maker: sessionmaker[Session] | None = None,
|
||||
) -> int:
|
||||
"""Conditionally stop timed-out runs that still belong to this process."""
|
||||
active_registrations = retain_active_workflow_tasks(tuple(registrations))
|
||||
workflow_run_ids = [
|
||||
registration.workflow_run_id
|
||||
for registration in active_registrations
|
||||
if registration.workflow_run_id is not None
|
||||
]
|
||||
return mark_workflow_runs_stopped_if_running_without_active_handoff(
|
||||
workflow_run_ids,
|
||||
reason=WORKFLOW_WARM_SHUTDOWN_TIMEOUT_REASON,
|
||||
now=now,
|
||||
session_maker=session_maker,
|
||||
)
|
||||
|
||||
|
||||
def _wait_for_workflow_drain_deadline(timeout_seconds: float) -> None:
|
||||
threading.Event().wait(timeout_seconds)
|
||||
|
||||
|
||||
def _run_workflow_drain_watchdog(
|
||||
*,
|
||||
timeout_seconds: float,
|
||||
fail_closed: Callable[[Sequence[ActiveWorkflowTask]], int] = _mark_timed_out_workflow_runs_stopped,
|
||||
hard_exit: Callable[[int], object] = os._exit,
|
||||
wait_for_deadline: Callable[[float], object] = _wait_for_workflow_drain_deadline,
|
||||
) -> None:
|
||||
"""Enforce the workflow drain deadline without depending on Graphon polls."""
|
||||
if timeout_seconds < 0:
|
||||
raise ValueError("timeout_seconds must be non-negative")
|
||||
|
||||
# Do not return merely because the registry is empty immediately after the
|
||||
# signal. An already-admitted API request can still be between request
|
||||
# dispatch and workflow registration. Keeping the deadline alive closes
|
||||
# that race; a naturally exiting process discards this daemon thread.
|
||||
wait_for_deadline(timeout_seconds)
|
||||
registrations = get_active_workflow_tasks()
|
||||
if not registrations:
|
||||
logger.info("All tracked workflow runs ended before the worker drain deadline")
|
||||
return
|
||||
|
||||
try:
|
||||
stopped_count = fail_closed(registrations)
|
||||
logger.error(
|
||||
"Workflow worker drain deadline expired with %s active run(s); marked %s RUNNING run(s) STOPPED",
|
||||
len(registrations),
|
||||
stopped_count,
|
||||
)
|
||||
except Exception:
|
||||
# A failed status write must not turn the deadline into an unbounded
|
||||
# Celery shutdown. The process exits non-zero so orchestration can
|
||||
# surface and replace the unhealthy worker.
|
||||
logger.exception("Failed to mark workflow runs STOPPED after the worker drain deadline; forcing worker exit")
|
||||
finally:
|
||||
hard_exit(1)
|
||||
|
||||
|
||||
def _start_workflow_drain_watchdog() -> None:
|
||||
"""Start at most one process-local drain deadline watchdog."""
|
||||
# Gunicorn's gevent worker invokes Python signal handlers from the hub
|
||||
# callback. Once ``threading`` is monkey-patched, ``Thread.start()`` waits
|
||||
# on a gevent Event and therefore cannot be called from that callback
|
||||
# (gevent raises ``BlockingSwitchOutError``). Keep the shutdown flag and
|
||||
# active-run inspection synchronous, but defer only the watchdog launch to
|
||||
# a normal greenlet where starting the patched thread is safe. Celery
|
||||
# workers without gevent monkey-patching continue down the native path.
|
||||
try:
|
||||
import gevent
|
||||
from gevent import monkey as gevent_monkey
|
||||
|
||||
if gevent_monkey.is_module_patched("threading") and gevent.getcurrent() is gevent.get_hub():
|
||||
gevent.spawn(_start_workflow_drain_watchdog)
|
||||
return
|
||||
except ImportError:
|
||||
# gevent is optional for non-Gunicorn process roles.
|
||||
pass
|
||||
|
||||
with _workflow_drain_watchdog_lock:
|
||||
if _workflow_drain_watchdog_started.is_set():
|
||||
return
|
||||
_workflow_drain_watchdog_started.set()
|
||||
watchdog = threading.Thread(
|
||||
target=_run_workflow_drain_watchdog,
|
||||
kwargs={"timeout_seconds": dify_config.WORKFLOW_HANDOFF_DRAIN_TIMEOUT_SECONDS},
|
||||
name="WorkflowDrainDeadline",
|
||||
daemon=True,
|
||||
)
|
||||
watchdog.start()
|
||||
|
||||
|
||||
def begin_workflow_warm_shutdown(*, source: str = "process") -> None:
|
||||
"""Start process-wide workflow draining for Celery or API workers.
|
||||
|
||||
The shutdown state is set before inspecting the registry so execution
|
||||
command channels created concurrently observe it. When durable handoff is
|
||||
enabled, the deadline watchdog starts even if the registry is momentarily
|
||||
empty, covering already-admitted API requests that register slightly later.
|
||||
"""
|
||||
if not source:
|
||||
raise ValueError("source must not be empty")
|
||||
|
||||
mark_workflow_warm_shutdown_started()
|
||||
active_count = get_active_workflow_task_count()
|
||||
if active_count:
|
||||
logger.info("Marked %s warm shutdown for %s active workflow run(s)", source, active_count)
|
||||
else:
|
||||
logger.info("No active workflow runs found when %s warm shutdown started", source)
|
||||
|
||||
if dify_config.WORKFLOW_HANDOFF_ENABLED:
|
||||
_start_workflow_drain_watchdog()
|
||||
|
||||
|
||||
def _on_worker_shutting_down(*args: object, **kwargs: object) -> None:
|
||||
"""Mark warm shutdown and log the active workflow run count."""
|
||||
how = kwargs.get("how")
|
||||
if not _is_warm_shutdown(how):
|
||||
logger.debug("Skip workflow abort during non-warm Celery shutdown: how=%s", how)
|
||||
logger.debug("Skip workflow handoff during non-warm Celery shutdown: how=%s", how)
|
||||
return
|
||||
|
||||
mark_celery_warm_shutdown_started()
|
||||
abort_count = get_active_workflow_task_count()
|
||||
if abort_count == 0:
|
||||
logger.info("No active workflow runs found during Celery warm shutdown")
|
||||
return
|
||||
|
||||
logger.info(
|
||||
"Marked Celery warm shutdown for %s active workflow run(s)",
|
||||
abort_count,
|
||||
)
|
||||
begin_workflow_warm_shutdown(source="Celery worker")
|
||||
|
||||
|
||||
def _on_worker_shutdown(*args: object, **kwargs: object) -> None:
|
||||
@@ -65,7 +259,7 @@ def _on_worker_shutdown(*args: object, **kwargs: object) -> None:
|
||||
|
||||
|
||||
def setup_workflow_warm_shutdown_handler() -> None:
|
||||
"""Connect Celery worker shutdown handlers for workflow abort and logging."""
|
||||
"""Connect Celery worker shutdown handlers for workflow drain and logging."""
|
||||
reset_active_workflow_tasks()
|
||||
worker_shutting_down.connect(
|
||||
_on_worker_shutting_down,
|
||||
|
||||
@@ -25,6 +25,7 @@ workflow_run_for_log_fields = {
|
||||
"triggered_from": fields.String,
|
||||
"error": fields.String,
|
||||
"elapsed_time": fields.Float,
|
||||
"handoff_duration": fields.Float,
|
||||
"total_tokens": fields.Integer,
|
||||
"total_steps": fields.Integer,
|
||||
"created_at": TimestampField,
|
||||
@@ -42,6 +43,7 @@ workflow_run_for_archived_log_fields = {
|
||||
"status": fields.String,
|
||||
"triggered_from": fields.String,
|
||||
"elapsed_time": fields.Float,
|
||||
"handoff_duration": fields.Float,
|
||||
"total_tokens": fields.Integer,
|
||||
}
|
||||
|
||||
@@ -57,6 +59,7 @@ class WorkflowRunForLogResponse(ResponseModel):
|
||||
triggered_from: str | None = None
|
||||
error: str | None = None
|
||||
elapsed_time: float | None = None
|
||||
handoff_duration: float = 0.0
|
||||
total_tokens: int | None = None
|
||||
total_steps: int | None = None
|
||||
created_at: int | None = None
|
||||
@@ -81,6 +84,7 @@ class WorkflowRunForArchivedLogResponse(ResponseModel):
|
||||
status: str | None = None
|
||||
triggered_from: str | None = None
|
||||
elapsed_time: float | None = None
|
||||
handoff_duration: float = 0.0
|
||||
total_tokens: int | None = None
|
||||
|
||||
@field_validator("status", mode="before")
|
||||
@@ -96,6 +100,7 @@ class WorkflowRunForListResponse(ResponseModel):
|
||||
version: str | None = None
|
||||
status: str | None = None
|
||||
elapsed_time: float | None = None
|
||||
handoff_duration: float = 0.0
|
||||
total_tokens: int | None = None
|
||||
total_steps: int | None = None
|
||||
created_by_account: SimpleAccount | None = None
|
||||
@@ -155,6 +160,7 @@ class WorkflowRunDetailResponse(ResponseModel):
|
||||
outputs: Any = Field(validation_alias="outputs_dict")
|
||||
error: str | None = None
|
||||
elapsed_time: float | None = None
|
||||
handoff_duration: float = 0.0
|
||||
total_tokens: int | None = None
|
||||
total_steps: int | None = None
|
||||
created_by_role: str | None = None
|
||||
|
||||
@@ -43,3 +43,38 @@ def post_patch(event):
|
||||
|
||||
|
||||
gevent_events.subscribers.append(post_patch)
|
||||
|
||||
|
||||
def post_worker_init(worker):
|
||||
"""Install workflow draining before Gunicorn's normal SIGTERM handling."""
|
||||
# Import lazily: this hook runs after the gevent worker has applied stdlib
|
||||
# monkey-patching, while this configuration module itself is imported much
|
||||
# earlier in the worker lifecycle.
|
||||
import signal
|
||||
|
||||
from configs import dify_config
|
||||
from extensions.workflow_warm_shutdown import begin_workflow_warm_shutdown
|
||||
|
||||
# Keep a dormant rollout behaviorally identical to the previous release.
|
||||
# Celery retains its historical warm-shutdown abort fallback, but Gunicorn
|
||||
# already waits for in-flight requests and must not inject an AbortCommand
|
||||
# unless durable handoff is explicitly enabled.
|
||||
if not dify_config.WORKFLOW_HANDOFF_ENABLED:
|
||||
return
|
||||
|
||||
original_sigterm_handler = signal.getsignal(signal.SIGTERM)
|
||||
if not callable(original_sigterm_handler):
|
||||
raise RuntimeError("Gunicorn SIGTERM handler is not installed before post_worker_init")
|
||||
|
||||
def handle_workflow_warm_shutdown(signum, frame):
|
||||
try:
|
||||
begin_workflow_warm_shutdown(source="Gunicorn API worker")
|
||||
except Exception:
|
||||
worker.log.exception("Failed to start workflow draining during Gunicorn SIGTERM")
|
||||
finally:
|
||||
original_sigterm_handler(signum, frame)
|
||||
|
||||
signal.signal(signal.SIGTERM, handle_workflow_warm_shutdown)
|
||||
# ``signal.signal`` can restore interruptible syscalls. Preserve Gunicorn's
|
||||
# behavior so SIGTERM does not interrupt an in-flight API request mid-I/O.
|
||||
signal.siginterrupt(signal.SIGTERM, False)
|
||||
|
||||
@@ -8,7 +8,22 @@ import types
|
||||
from abc import abstractmethod
|
||||
from collections.abc import Iterator
|
||||
from contextlib import AbstractContextManager
|
||||
from typing import Protocol, Self, override
|
||||
from dataclasses import dataclass
|
||||
from typing import Protocol, Self, override, runtime_checkable
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class CursorMessage:
|
||||
"""A broadcast payload paired with a durable transport cursor.
|
||||
|
||||
Cursor-aware consumers use this value to expose Redis Stream entry IDs as
|
||||
standard SSE ``id`` fields. The existing byte-only ``receive`` API stays
|
||||
available so non-HTTP broadcast consumers do not need to know about
|
||||
transport metadata.
|
||||
"""
|
||||
|
||||
payload: bytes
|
||||
cursor: str
|
||||
|
||||
|
||||
class Subscription(AbstractContextManager["Subscription"], Protocol):
|
||||
@@ -75,6 +90,16 @@ class Subscription(AbstractContextManager["Subscription"], Protocol):
|
||||
...
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class CursorSubscription(Subscription, Protocol):
|
||||
"""A subscription that can return a durable cursor with each message."""
|
||||
|
||||
@abstractmethod
|
||||
def receive_with_cursor(self, timeout: float | None = 0.1) -> CursorMessage | None:
|
||||
"""Receive the next payload and its transport cursor."""
|
||||
...
|
||||
|
||||
|
||||
class Producer(Protocol):
|
||||
"""Producer is an interface for message publishing. It is already bound to a specific topic.
|
||||
|
||||
@@ -94,7 +119,12 @@ class Subscriber(Protocol):
|
||||
"""
|
||||
|
||||
@abstractmethod
|
||||
def subscribe(self) -> Subscription:
|
||||
def subscribe(self, *, cursor: str | None = None) -> Subscription:
|
||||
"""Create a subscription after ``cursor`` when the transport supports replay.
|
||||
|
||||
Non-replayable transports reject a non-null cursor. Omitting the
|
||||
cursor preserves their normal live-subscription behavior.
|
||||
"""
|
||||
pass
|
||||
|
||||
|
||||
@@ -118,6 +148,21 @@ class Topic(Producer, Subscriber, Protocol):
|
||||
...
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class ReplayableTopic(Topic, Protocol):
|
||||
"""A topic backed by an ordered, durable event log."""
|
||||
|
||||
@abstractmethod
|
||||
def earliest_cursor(self) -> str | None:
|
||||
"""Return the earliest retained cursor, or ``None`` for an empty log."""
|
||||
...
|
||||
|
||||
@abstractmethod
|
||||
def latest_cursor(self) -> str | None:
|
||||
"""Return the latest retained cursor, or ``None`` for an empty log."""
|
||||
...
|
||||
|
||||
|
||||
class BroadcastChannel(Protocol):
|
||||
"""A broadcasting channel is a channel supporting broadcasting semantics.
|
||||
|
||||
|
||||
@@ -0,0 +1,27 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
|
||||
_REDIS_STREAM_ID_PATTERN = re.compile(r"^[0-9]+-[0-9]+$")
|
||||
_REDIS_STREAM_ID_COMPONENT_MAX = (1 << 64) - 1
|
||||
_REDIS_STREAM_ID_COMPONENT_MAX_DIGITS = len(str(_REDIS_STREAM_ID_COMPONENT_MAX))
|
||||
|
||||
|
||||
def normalize_stream_cursor(cursor: str | bytes) -> str:
|
||||
"""Normalize and validate a public Redis Streams replay cursor."""
|
||||
|
||||
if isinstance(cursor, bytes):
|
||||
try:
|
||||
cursor = cursor.decode("ascii")
|
||||
except UnicodeDecodeError as error:
|
||||
raise ValueError("event cursor must be a Redis Stream ID such as '1712345678901-0'") from error
|
||||
normalized = cursor.strip()
|
||||
if not _REDIS_STREAM_ID_PATTERN.fullmatch(normalized):
|
||||
raise ValueError("event cursor must be a Redis Stream ID such as '1712345678901-0'")
|
||||
milliseconds, sequence = normalized.split("-", maxsplit=1)
|
||||
for component in (milliseconds, sequence):
|
||||
# Bound before int() so adversarial headers cannot hit Python's
|
||||
# max_str_digits exception and turn a client validation error into 500.
|
||||
if len(component) > _REDIS_STREAM_ID_COMPONENT_MAX_DIGITS or int(component) > _REDIS_STREAM_ID_COMPONENT_MAX:
|
||||
raise ValueError("event cursor components must be unsigned 64-bit integers")
|
||||
return normalized
|
||||
@@ -1,4 +1,5 @@
|
||||
from .pubsub_channel import BroadcastChannel
|
||||
from .sharded_channel import ShardedRedisBroadcastChannel
|
||||
from .streams_channel import StreamsBroadcastChannel
|
||||
|
||||
__all__ = ["BroadcastChannel", "ShardedRedisBroadcastChannel"]
|
||||
__all__ = ["BroadcastChannel", "ShardedRedisBroadcastChannel", "StreamsBroadcastChannel"]
|
||||
|
||||
@@ -84,8 +84,8 @@ class RedisSubscriptionBase(Subscription):
|
||||
if raw_message is None:
|
||||
continue
|
||||
|
||||
# If close() sent a control event to unblock us, exit immediately
|
||||
# without processing any message — the subscription is shutting down.
|
||||
# close() flips only local state; exit without processing a message
|
||||
# that happened to arrive while this subscription was shutting down.
|
||||
if self._closed.is_set():
|
||||
break
|
||||
|
||||
@@ -117,9 +117,13 @@ class RedisSubscriptionBase(Subscription):
|
||||
)
|
||||
continue
|
||||
|
||||
self._enqueue_message(payload_bytes)
|
||||
# Older replicas may still publish the legacy shared close marker
|
||||
# during an adjacent-version rollout. It is subscription-local
|
||||
# control data, never an application event, so ignore it rather
|
||||
# than terminating every subscriber on the topic.
|
||||
if payload_bytes == SIG_CLOSE:
|
||||
break
|
||||
continue
|
||||
self._enqueue_message(payload_bytes)
|
||||
|
||||
_logger.debug("%s listener thread stopped for channel %s", self._get_subscription_type().title(), self._topic)
|
||||
try:
|
||||
@@ -193,6 +197,9 @@ class RedisSubscriptionBase(Subscription):
|
||||
except queue.Empty:
|
||||
return None
|
||||
|
||||
if self._closed.is_set():
|
||||
raise SubscriptionClosedError(f"The Redis {self._get_subscription_type()} subscription is closed")
|
||||
|
||||
return item
|
||||
|
||||
@override
|
||||
@@ -227,9 +234,6 @@ class RedisSubscriptionBase(Subscription):
|
||||
if started:
|
||||
self._unblock_message_iterator()
|
||||
|
||||
# Send a control event on the same Redis channel to unblock the
|
||||
self._publish_close_event()
|
||||
|
||||
# NOTE: PubSub is not thread-safe. More specifically, the `PubSub.close` method and the
|
||||
# message retrieval method should NOT be called concurrently.
|
||||
#
|
||||
@@ -255,15 +259,6 @@ class RedisSubscriptionBase(Subscription):
|
||||
"""Return the subscription type (e.g., 'regular' or 'sharded')."""
|
||||
raise NotImplementedError
|
||||
|
||||
def _publish_close_event(self) -> None:
|
||||
"""Publish a control event on the Redis channel to unblock the listener.
|
||||
|
||||
This is called by close() after setting _closed. The subclass should
|
||||
publish an empty message on the same topic so that a blocking
|
||||
get_message() call in the listener thread returns promptly.
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
def _subscribe(self) -> None:
|
||||
"""Subscribe to the Redis topic using the appropriate command."""
|
||||
raise NotImplementedError
|
||||
|
||||
@@ -1,17 +1,13 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from typing import Any, override
|
||||
|
||||
from extensions.redis_names import serialize_redis_name
|
||||
from libs.broadcast_channel.channel import Producer, Subscriber, Subscription
|
||||
from libs.broadcast_channel.signals import SIG_CLOSE
|
||||
from redis import Redis, RedisCluster
|
||||
|
||||
from ._subscription import RedisSubscriptionBase
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class BroadcastChannel:
|
||||
"""
|
||||
@@ -52,7 +48,9 @@ class Topic:
|
||||
def as_subscriber(self) -> Subscriber:
|
||||
return self
|
||||
|
||||
def subscribe(self) -> Subscription:
|
||||
def subscribe(self, *, cursor: str | None = None) -> Subscription:
|
||||
if cursor is not None:
|
||||
raise ValueError("Redis Pub/Sub does not support replay cursors")
|
||||
return _RedisSubscription(
|
||||
client=self._client,
|
||||
pubsub=self._client.pubsub(),
|
||||
@@ -67,13 +65,6 @@ class _RedisSubscription(RedisSubscriptionBase):
|
||||
def _get_subscription_type(self) -> str:
|
||||
return "regular"
|
||||
|
||||
@override
|
||||
def _publish_close_event(self) -> None:
|
||||
try:
|
||||
self._client.publish(self._topic, SIG_CLOSE)
|
||||
except Exception:
|
||||
logger.exception("failed to publish close event")
|
||||
|
||||
@override
|
||||
def _subscribe(self) -> None:
|
||||
assert self._pubsub is not None
|
||||
|
||||
@@ -1,17 +1,13 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from typing import Any, override
|
||||
|
||||
from extensions.redis_names import serialize_redis_name
|
||||
from libs.broadcast_channel.channel import Producer, Subscriber, Subscription
|
||||
from libs.broadcast_channel.signals import SIG_CLOSE
|
||||
from redis import Redis, RedisCluster
|
||||
|
||||
from ._subscription import RedisSubscriptionBase
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class ShardedRedisBroadcastChannel:
|
||||
"""
|
||||
@@ -50,7 +46,9 @@ class ShardedTopic:
|
||||
def as_subscriber(self) -> Subscriber:
|
||||
return self
|
||||
|
||||
def subscribe(self) -> Subscription:
|
||||
def subscribe(self, *, cursor: str | None = None) -> Subscription:
|
||||
if cursor is not None:
|
||||
raise ValueError("Redis Sharded Pub/Sub does not support replay cursors")
|
||||
return _RedisShardedSubscription(
|
||||
client=self._client,
|
||||
pubsub=self._client.pubsub(),
|
||||
@@ -65,13 +63,6 @@ class _RedisShardedSubscription(RedisSubscriptionBase):
|
||||
def _get_subscription_type(self) -> str:
|
||||
return "sharded"
|
||||
|
||||
@override
|
||||
def _publish_close_event(self) -> None:
|
||||
try:
|
||||
self._client.spublish(self._topic, SIG_CLOSE) # type: ignore[attr-defined,union-attr]
|
||||
except Exception:
|
||||
logger.exception("failed to publish close event")
|
||||
|
||||
@override
|
||||
def _subscribe(self) -> None:
|
||||
assert self._pubsub is not None
|
||||
|
||||
@@ -4,16 +4,27 @@ import logging
|
||||
import queue
|
||||
import threading
|
||||
from collections.abc import Iterator
|
||||
from typing import Self, override
|
||||
from typing import Protocol, Self, cast, override
|
||||
|
||||
from extensions.redis_names import serialize_redis_name
|
||||
from libs.broadcast_channel.channel import Producer, Subscriber, Subscription
|
||||
from libs.broadcast_channel.channel import CursorMessage, CursorSubscription, Producer, Subscriber
|
||||
from libs.broadcast_channel.cursor import normalize_stream_cursor
|
||||
from libs.broadcast_channel.exc import SubscriptionClosedError
|
||||
from libs.broadcast_channel.signals import SIG_CLOSE
|
||||
from redis import Redis, RedisCluster
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_XADD_WITH_EXPIRE_LUA = """
|
||||
local entry_id = redis.call('XADD', KEYS[1], '*', 'data', ARGV[1])
|
||||
redis.call('EXPIRE', KEYS[1], ARGV[2])
|
||||
return entry_id
|
||||
"""
|
||||
|
||||
|
||||
class _RedisLuaClient(Protocol):
|
||||
def eval(self, script: str, numkeys: int, *keys_and_args: str | bytes | int) -> object: ...
|
||||
|
||||
|
||||
class StreamsBroadcastChannel:
|
||||
"""
|
||||
@@ -32,7 +43,7 @@ class StreamsBroadcastChannel:
|
||||
retention_seconds: int = 600,
|
||||
):
|
||||
self._client = redis_client
|
||||
self._retention_seconds = max(int(retention_seconds or 0), 0)
|
||||
self._retention_seconds = max(retention_seconds, 0)
|
||||
|
||||
def topic(self, topic: str) -> StreamsTopic:
|
||||
return StreamsTopic(
|
||||
@@ -54,34 +65,62 @@ class StreamsTopic:
|
||||
self._topic = topic
|
||||
self._key = serialize_redis_name(f"stream:{topic}")
|
||||
self._retention_seconds = retention_seconds
|
||||
self.max_length = 5000
|
||||
|
||||
def as_producer(self) -> Producer:
|
||||
return self
|
||||
|
||||
def publish(self, payload: bytes) -> None:
|
||||
self._client.xadd(self._key, {b"data": payload}, maxlen=self.max_length)
|
||||
# Retention is bounded by key expiry rather than MAXLEN. Trimming a
|
||||
# live per-run stream could silently invalidate a reconnect cursor.
|
||||
if self._retention_seconds > 0:
|
||||
try:
|
||||
self._client.expire(self._key, self._retention_seconds)
|
||||
except Exception as e:
|
||||
logger.warning("Failed to set expire for stream key %s: %s", self._key, e, exc_info=True)
|
||||
# A single-key Lua command is atomic on both standalone Redis and
|
||||
# Redis Cluster. A process exit cannot leave a newly appended
|
||||
# stream without its TTL, and failures propagate to the publisher
|
||||
# instead of silently creating an immortal event log.
|
||||
cast(_RedisLuaClient, self._client).eval(
|
||||
_XADD_WITH_EXPIRE_LUA,
|
||||
1,
|
||||
self._key,
|
||||
payload,
|
||||
self._retention_seconds,
|
||||
)
|
||||
return
|
||||
self._client.xadd(self._key, {b"data": payload})
|
||||
|
||||
def as_subscriber(self) -> Subscriber:
|
||||
return self
|
||||
|
||||
def subscribe(self) -> Subscription:
|
||||
return _StreamsSubscription(self._client, self._key)
|
||||
def subscribe(self, *, cursor: str | None = None) -> CursorSubscription:
|
||||
# A new subscriber replays the retained per-run log. XREAD is strictly
|
||||
# greater-than, so a Last-Event-ID can be passed through without
|
||||
# duplicating the event the client already acknowledged.
|
||||
start_cursor = "0-0" if cursor is None else normalize_stream_cursor(cursor)
|
||||
return _StreamsSubscription(self._client, self._key, cursor=start_cursor)
|
||||
|
||||
def earliest_cursor(self) -> str | None:
|
||||
entries = self._client.xrange(self._key, count=1)
|
||||
if not entries:
|
||||
return None
|
||||
entry_id, _ = entries[0]
|
||||
return normalize_stream_cursor(entry_id)
|
||||
|
||||
def latest_cursor(self) -> str | None:
|
||||
entries = self._client.xrevrange(self._key, count=1)
|
||||
if not entries:
|
||||
return None
|
||||
entry_id, _ = entries[0]
|
||||
return normalize_stream_cursor(entry_id)
|
||||
|
||||
|
||||
class _StreamsSubscription(Subscription):
|
||||
class _StreamsSubscription(CursorSubscription):
|
||||
_SENTINEL = object()
|
||||
|
||||
def __init__(self, client: Redis | RedisCluster, key: str):
|
||||
def __init__(self, client: Redis | RedisCluster, key: str, *, cursor: str = "0-0"):
|
||||
self._client = client
|
||||
self._key = key
|
||||
self._cursor = normalize_stream_cursor(cursor)
|
||||
|
||||
self._queue: queue.Queue[object] = queue.Queue()
|
||||
self._queue: queue.Queue[CursorMessage | object] = queue.Queue()
|
||||
|
||||
# The `_lock` lock is used to
|
||||
#
|
||||
@@ -104,21 +143,24 @@ class _StreamsSubscription(Subscription):
|
||||
# since this method runs in a dedicated thread, acquiring `_lock` inside this method won't cause
|
||||
# deadlock.
|
||||
|
||||
# Setting initial last id to `$` to signal redis that we only want new messages.
|
||||
#
|
||||
# ref: https://redis.io/docs/latest/commands/xread/#the-special--id
|
||||
last_id = "$"
|
||||
last_id = self._cursor
|
||||
try:
|
||||
while True:
|
||||
with self._lock:
|
||||
if self._closed:
|
||||
break
|
||||
streams = self._client.xread({self._key: last_id}, block=1000, count=100)
|
||||
# A short bounded block lets close() remain purely local while
|
||||
# still releasing the listener promptly.
|
||||
streams = self._client.xread({self._key: last_id}, block=100, count=100)
|
||||
if not streams:
|
||||
continue
|
||||
|
||||
for _, entries in streams:
|
||||
for entry_id, fields in entries:
|
||||
cursor = normalize_stream_cursor(entry_id)
|
||||
# Advance over malformed/legacy control entries as well;
|
||||
# otherwise the next XREAD would return them forever.
|
||||
last_id = cursor
|
||||
data = None
|
||||
if isinstance(fields, dict):
|
||||
data = fields.get(b"data")
|
||||
@@ -130,9 +172,8 @@ class _StreamsSubscription(Subscription):
|
||||
data_bytes = bytes(data)
|
||||
if data_bytes is not None:
|
||||
if data_bytes == SIG_CLOSE:
|
||||
break
|
||||
self._queue.put_nowait(data_bytes)
|
||||
last_id = entry_id
|
||||
continue
|
||||
self._queue.put_nowait(CursorMessage(payload=data_bytes, cursor=cursor))
|
||||
finally:
|
||||
self._queue.put_nowait(self._SENTINEL)
|
||||
with self._lock:
|
||||
@@ -172,6 +213,11 @@ class _StreamsSubscription(Subscription):
|
||||
|
||||
@override
|
||||
def receive(self, timeout: float | None = 0.1) -> bytes | None:
|
||||
message = self.receive_with_cursor(timeout=timeout)
|
||||
return None if message is None else message.payload
|
||||
|
||||
@override
|
||||
def receive_with_cursor(self, timeout: float | None = 0.1) -> CursorMessage | None:
|
||||
with self._lock:
|
||||
if self._closed:
|
||||
raise SubscriptionClosedError("The Redis streams subscription is closed")
|
||||
@@ -187,15 +233,8 @@ class _StreamsSubscription(Subscription):
|
||||
|
||||
if item is self._SENTINEL:
|
||||
raise SubscriptionClosedError("The Redis streams subscription is closed")
|
||||
assert isinstance(item, (bytes, bytearray)), "Unexpected item type in stream queue"
|
||||
return bytes(item)
|
||||
|
||||
def _publish_close_event(self) -> None:
|
||||
"""Publish an empty message to the stream to unblock the listener's xread."""
|
||||
try:
|
||||
self._client.xadd(self._key, {b"data": SIG_CLOSE})
|
||||
except Exception:
|
||||
logger.exception("failed to publish close event")
|
||||
assert isinstance(item, CursorMessage), "Unexpected item type in stream queue"
|
||||
return item
|
||||
|
||||
@override
|
||||
def close(self) -> None:
|
||||
@@ -207,8 +246,10 @@ class _StreamsSubscription(Subscription):
|
||||
if listener is not None:
|
||||
self._listener = None
|
||||
|
||||
if listener is not None:
|
||||
self._publish_close_event()
|
||||
# Wake local consumers immediately. The Redis XREAD call uses a
|
||||
# bounded short poll and observes _closed without writing a
|
||||
# shared marker into the stream.
|
||||
self._queue.put_nowait(self._SENTINEL)
|
||||
|
||||
if listener is not None and listener.is_alive():
|
||||
listener.join(timeout=2)
|
||||
|
||||
+91
-11
@@ -7,10 +7,10 @@ import struct
|
||||
import subprocess
|
||||
import time
|
||||
import uuid
|
||||
from collections.abc import Callable, Generator, Mapping
|
||||
from collections.abc import Callable, Generator, Iterator, Mapping
|
||||
from datetime import datetime
|
||||
from hashlib import sha256
|
||||
from typing import TYPE_CHECKING, Annotated, Any, Protocol, cast, overload, override
|
||||
from typing import TYPE_CHECKING, Annotated, Any, Protocol, cast, overload, override, runtime_checkable
|
||||
from uuid import UUID
|
||||
from zoneinfo import available_timezones
|
||||
|
||||
@@ -21,6 +21,7 @@ from pydantic.functional_validators import AfterValidator
|
||||
from typing_extensions import TypedDict
|
||||
|
||||
from configs import dify_config
|
||||
from core.app.apps.streaming_utils import WorkflowRunIdentifiedStream
|
||||
from core.app.features.rate_limiting.rate_limit import RateLimitGenerator
|
||||
from extensions.ext_redis import redis_client
|
||||
from graphon.file import helpers as file_helpers
|
||||
@@ -33,6 +34,22 @@ if TYPE_CHECKING:
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class _Closable(Protocol):
|
||||
def close(self) -> object: ...
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class _WorkflowRunIdentified(Protocol):
|
||||
workflow_run_id: str | None
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class _ValueBearing(Protocol):
|
||||
@property
|
||||
def value(self) -> object: ...
|
||||
|
||||
|
||||
@with_config(ConfigDict(extra="allow"))
|
||||
class _TokenData(TypedDict, total=False):
|
||||
"""Shared baseline token payload.
|
||||
@@ -405,31 +422,90 @@ def generate_text_hash(text: str) -> str:
|
||||
return sha256(hash_text.encode()).hexdigest()
|
||||
|
||||
|
||||
def _is_workflow_maintenance_sse_chunk(chunk: object) -> bool:
|
||||
if not isinstance(chunk, str) or "workflow_maintenance_paused" not in chunk:
|
||||
return False
|
||||
for line in chunk.splitlines():
|
||||
if not line.startswith("data: "):
|
||||
continue
|
||||
try:
|
||||
payload = json.loads(line[6:])
|
||||
except (TypeError, json.JSONDecodeError):
|
||||
continue
|
||||
if isinstance(payload, dict) and payload.get("event") == "workflow_maintenance_paused":
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
def compact_generate_response(
|
||||
response: Mapping[str, Any] | Generator[str, None, None] | RateLimitGenerator,
|
||||
response: Mapping[str, Any] | Iterator[str] | RateLimitGenerator | WorkflowRunIdentifiedStream,
|
||||
) -> Response:
|
||||
if isinstance(response, Mapping):
|
||||
event = response.get("event")
|
||||
event_value = event.value if isinstance(event, _ValueBearing) else event
|
||||
is_maintenance_handoff = event_value == "workflow_maintenance_paused"
|
||||
workflow_run_id = response.get("workflow_run_id")
|
||||
headers: dict[str, str] = {}
|
||||
if is_maintenance_handoff:
|
||||
headers["Retry-After"] = "1"
|
||||
if isinstance(workflow_run_id, str) and workflow_run_id:
|
||||
headers.update(
|
||||
{
|
||||
"X-Workflow-Run-ID": workflow_run_id,
|
||||
"Access-Control-Expose-Headers": "X-Workflow-Run-ID, Retry-After",
|
||||
}
|
||||
)
|
||||
return Response(
|
||||
response=json.dumps(jsonable_encoder(response)),
|
||||
status=200,
|
||||
# A blocking request cannot remain attached to a worker that is
|
||||
# intentionally exiting. Keep the logical run alive and return an
|
||||
# explicit continuation response instead of presenting an internal
|
||||
# maintenance sentinel as a successful terminal result.
|
||||
status=202 if is_maintenance_handoff else 200,
|
||||
content_type="application/json; charset=utf-8",
|
||||
headers=headers,
|
||||
)
|
||||
else:
|
||||
stream_response = response
|
||||
workflow_run_id = (
|
||||
stream_response.workflow_run_id if isinstance(stream_response, _WorkflowRunIdentified) else None
|
||||
)
|
||||
|
||||
def generate() -> Generator[str, None, None]:
|
||||
yield from stream_response
|
||||
try:
|
||||
for chunk in stream_response:
|
||||
# Maintenance handoff is an internal segment boundary. Direct
|
||||
# console/debug generators do not pass through the Celery
|
||||
# publisher (which already consumes this sentinel), so filter
|
||||
# it at the final HTTP boundary and let the cursor-aware client
|
||||
# reconnect to the durable per-run stream.
|
||||
if _is_workflow_maintenance_sse_chunk(chunk):
|
||||
continue
|
||||
yield chunk
|
||||
finally:
|
||||
if isinstance(stream_response, _Closable):
|
||||
stream_response.close()
|
||||
|
||||
response_headers: dict[str, str] = {}
|
||||
if isinstance(workflow_run_id, str) and workflow_run_id:
|
||||
response_headers = {
|
||||
"X-Workflow-Run-ID": workflow_run_id,
|
||||
"Access-Control-Expose-Headers": "X-Workflow-Run-ID",
|
||||
}
|
||||
|
||||
return Response(
|
||||
_stream_with_request_context(generate()),
|
||||
status=200,
|
||||
mimetype="text/event-stream",
|
||||
headers=response_headers,
|
||||
)
|
||||
|
||||
|
||||
def length_prefixed_response(
|
||||
magic_number: int,
|
||||
response: Mapping[str, Any] | BaseModel | Generator[str | bytes, None, None] | RateLimitGenerator,
|
||||
response: (
|
||||
Mapping[str, Any] | BaseModel | Iterator[str | bytes] | RateLimitGenerator | WorkflowRunIdentifiedStream
|
||||
),
|
||||
) -> Response:
|
||||
"""
|
||||
This function is used to return a response with a length prefix.
|
||||
@@ -478,11 +554,15 @@ def length_prefixed_response(
|
||||
stream_response = response
|
||||
|
||||
def generate() -> Generator[bytes, None, None]:
|
||||
for chunk in stream_response:
|
||||
if isinstance(chunk, str):
|
||||
yield pack_response_with_length_prefix(chunk.encode("utf-8"))
|
||||
else:
|
||||
yield pack_response_with_length_prefix(chunk)
|
||||
try:
|
||||
for chunk in stream_response:
|
||||
if isinstance(chunk, str):
|
||||
yield pack_response_with_length_prefix(chunk.encode("utf-8"))
|
||||
else:
|
||||
yield pack_response_with_length_prefix(chunk)
|
||||
finally:
|
||||
if isinstance(stream_response, _Closable):
|
||||
stream_response.close()
|
||||
|
||||
return Response(
|
||||
_stream_with_request_context(generate()),
|
||||
|
||||
@@ -0,0 +1,269 @@
|
||||
"""add workflow run handoffs
|
||||
|
||||
Revision ID: 4f3a2b1c9d8e
|
||||
Revises: f6e4c5686857
|
||||
Create Date: 2026-07-28 15:00:00.000000
|
||||
|
||||
"""
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
|
||||
import models
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision = "4f3a2b1c9d8e"
|
||||
down_revision = "f6e4c5686857"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade():
|
||||
with op.batch_alter_table("workflow_runs", schema=None) as batch_op:
|
||||
batch_op.add_column(sa.Column("handoff_duration", sa.Float(), server_default=sa.text("0"), nullable=False))
|
||||
|
||||
with op.batch_alter_table("workflow_archive_logs", schema=None) as batch_op:
|
||||
batch_op.add_column(sa.Column("run_handoff_duration", sa.Float(), server_default=sa.text("0"), nullable=False))
|
||||
|
||||
op.create_table(
|
||||
"workflow_run_handoffs",
|
||||
sa.Column("workflow_run_id", models.types.StringUUID(), nullable=False),
|
||||
sa.Column("generation", sa.Integer(), nullable=False),
|
||||
sa.Column("task_id", sa.String(length=255), nullable=False),
|
||||
sa.Column("snapshot_object_key", sa.String(length=255), nullable=False),
|
||||
sa.Column("snapshot_schema_version", sa.String(length=64), nullable=False),
|
||||
sa.Column("snapshot_checksum", sa.String(length=128), nullable=False),
|
||||
sa.Column("snapshot_size_bytes", sa.BigInteger(), nullable=False),
|
||||
sa.Column("resume_route", sa.String(length=64), nullable=False),
|
||||
sa.Column("source_worker_id", sa.String(length=255), nullable=False),
|
||||
sa.Column("rag_source_batch_id", sa.String(length=255), nullable=True),
|
||||
sa.Column("rag_tenant_id", models.types.StringUUID(), nullable=True),
|
||||
sa.Column("rag_queue_kind", sa.String(length=32), nullable=True),
|
||||
sa.Column("rag_dataset_id", models.types.StringUUID(), nullable=True),
|
||||
sa.Column("rag_document_id", models.types.StringUUID(), nullable=True),
|
||||
sa.Column("rag_tenant_isolated", sa.Boolean(), nullable=True),
|
||||
sa.Column("rag_group_sealed_at", sa.DateTime(), nullable=True),
|
||||
sa.Column("rag_tenant_slot_released_at", sa.DateTime(), nullable=True),
|
||||
sa.Column("rag_document_error_marked_at", sa.DateTime(), nullable=True),
|
||||
sa.Column("state", sa.String(length=32), server_default=sa.text("'prepared'"), nullable=False),
|
||||
sa.Column("target_worker_id", sa.String(length=255), nullable=True),
|
||||
sa.Column("lease_owner", sa.String(length=255), nullable=True),
|
||||
sa.Column("lease_token", models.types.StringUUID(), nullable=True),
|
||||
sa.Column("lease_expires_at", sa.DateTime(), nullable=True),
|
||||
sa.Column("attempts", sa.Integer(), server_default=sa.text("0"), nullable=False),
|
||||
sa.Column("next_retry_at", sa.DateTime(), nullable=True),
|
||||
sa.Column("dispatched_at", sa.DateTime(), nullable=True),
|
||||
sa.Column("last_error", models.types.LongText(), nullable=True),
|
||||
sa.Column("cancel_requested_at", sa.DateTime(), nullable=True),
|
||||
sa.Column("claimed_at", sa.DateTime(), nullable=True),
|
||||
sa.Column("resumed_at", sa.DateTime(), nullable=True),
|
||||
sa.Column("failed_at", sa.DateTime(), nullable=True),
|
||||
sa.Column("terminal_compensated_at", sa.DateTime(), nullable=True),
|
||||
sa.Column("terminal_event_published_at", sa.DateTime(), nullable=True),
|
||||
sa.Column("terminal_attempts", sa.Integer(), server_default=sa.text("0"), nullable=False),
|
||||
sa.Column("terminal_last_error", models.types.LongText(), nullable=True),
|
||||
sa.Column("id", models.types.StringUUID(), nullable=False),
|
||||
sa.Column("created_at", sa.DateTime(), server_default=sa.text("CURRENT_TIMESTAMP"), nullable=False),
|
||||
sa.Column("updated_at", sa.DateTime(), server_default=sa.text("CURRENT_TIMESTAMP"), nullable=False),
|
||||
sa.CheckConstraint(
|
||||
"attempts >= 0",
|
||||
name=op.f("workflow_run_handoffs_attempts_nonnegative_check"),
|
||||
),
|
||||
sa.CheckConstraint(
|
||||
"terminal_attempts >= 0",
|
||||
name=op.f("workflow_run_handoffs_terminal_attempts_nonnegative_check"),
|
||||
),
|
||||
sa.CheckConstraint(
|
||||
"state <> 'claimed' OR "
|
||||
"(lease_owner IS NOT NULL AND lease_token IS NOT NULL AND lease_expires_at IS NOT NULL)",
|
||||
name=op.f("workflow_run_handoffs_claim_lease_present_check"),
|
||||
),
|
||||
sa.CheckConstraint(
|
||||
"state <> 'failed' OR failed_at IS NOT NULL",
|
||||
name=op.f("workflow_run_handoffs_failed_at_present_check"),
|
||||
),
|
||||
sa.CheckConstraint(
|
||||
"generation > 0",
|
||||
name=op.f("workflow_run_handoffs_generation_positive_check"),
|
||||
),
|
||||
sa.CheckConstraint(
|
||||
"(rag_source_batch_id IS NULL AND rag_tenant_id IS NULL AND rag_queue_kind IS NULL "
|
||||
"AND rag_dataset_id IS NULL AND rag_tenant_isolated IS NULL) OR "
|
||||
"(rag_source_batch_id IS NOT NULL AND rag_tenant_id IS NOT NULL AND rag_queue_kind IS NOT NULL "
|
||||
"AND rag_dataset_id IS NOT NULL AND rag_tenant_isolated IS NOT NULL)",
|
||||
name=op.f("workflow_run_handoffs_rag_group_metadata_complete_check"),
|
||||
),
|
||||
sa.CheckConstraint(
|
||||
"rag_queue_kind IS NULL OR rag_queue_kind IN ('regular', 'priority')",
|
||||
name=op.f("workflow_run_handoffs_rag_queue_kind_valid_check"),
|
||||
),
|
||||
sa.CheckConstraint(
|
||||
"rag_tenant_slot_released_at IS NULL OR rag_group_sealed_at IS NOT NULL",
|
||||
name=op.f("workflow_run_handoffs_rag_release_requires_seal_check"),
|
||||
),
|
||||
sa.CheckConstraint(
|
||||
"resume_route IN ('workflow', 'snippet', 'advanced_chat', 'triggered_workflow', 'rag_pipeline')",
|
||||
name=op.f("workflow_run_handoffs_resume_route_valid_check"),
|
||||
),
|
||||
sa.CheckConstraint(
|
||||
"state <> 'resumed' OR resumed_at IS NOT NULL",
|
||||
name=op.f("workflow_run_handoffs_resumed_at_present_check"),
|
||||
),
|
||||
sa.CheckConstraint(
|
||||
"snapshot_size_bytes >= 0",
|
||||
name=op.f("workflow_run_handoffs_snapshot_size_nonnegative_check"),
|
||||
),
|
||||
sa.CheckConstraint(
|
||||
"state IN ('preparing', 'prepared', 'ready', 'claimed', 'resumed', 'failed')",
|
||||
name=op.f("workflow_run_handoffs_state_valid_check"),
|
||||
),
|
||||
sa.PrimaryKeyConstraint("id", name=op.f("workflow_run_handoffs_pkey")),
|
||||
sa.UniqueConstraint(
|
||||
"workflow_run_id",
|
||||
"generation",
|
||||
name="workflow_run_handoffs_run_generation_key",
|
||||
),
|
||||
)
|
||||
with op.batch_alter_table("workflow_run_handoffs", schema=None) as batch_op:
|
||||
batch_op.create_index(
|
||||
"workflow_run_handoffs_rag_reconcile_idx",
|
||||
["rag_group_sealed_at", "rag_tenant_slot_released_at", "rag_source_batch_id"],
|
||||
unique=False,
|
||||
)
|
||||
batch_op.create_index(
|
||||
"workflow_run_handoffs_rag_group_release_idx",
|
||||
[
|
||||
"rag_source_batch_id",
|
||||
"rag_tenant_id",
|
||||
"rag_queue_kind",
|
||||
"rag_group_sealed_at",
|
||||
"rag_tenant_slot_released_at",
|
||||
],
|
||||
unique=False,
|
||||
)
|
||||
batch_op.create_index(
|
||||
"workflow_run_handoffs_dispatch_idx",
|
||||
["state", "next_retry_at", "lease_expires_at", "dispatched_at", "created_at"],
|
||||
unique=False,
|
||||
)
|
||||
batch_op.create_index(
|
||||
"workflow_run_handoffs_run_state_idx",
|
||||
["workflow_run_id", "state"],
|
||||
unique=False,
|
||||
)
|
||||
batch_op.create_index(
|
||||
"workflow_run_handoffs_snapshot_object_key_idx",
|
||||
["snapshot_object_key"],
|
||||
unique=False,
|
||||
)
|
||||
batch_op.create_index(
|
||||
"workflow_run_handoffs_state_created_idx",
|
||||
["state", "created_at", "id"],
|
||||
unique=False,
|
||||
)
|
||||
batch_op.create_index(
|
||||
"workflow_run_handoffs_resumed_retention_idx",
|
||||
["state", "resumed_at", "id"],
|
||||
unique=False,
|
||||
)
|
||||
batch_op.create_index(
|
||||
"workflow_run_handoffs_failed_retention_idx",
|
||||
["state", "failed_at", "id"],
|
||||
unique=False,
|
||||
)
|
||||
batch_op.create_index(
|
||||
"workflow_run_handoffs_task_state_idx",
|
||||
["task_id", "state"],
|
||||
unique=False,
|
||||
)
|
||||
batch_op.create_index(
|
||||
"workflow_run_handoffs_terminal_idx",
|
||||
["state", "terminal_compensated_at", "terminal_event_published_at", "created_at"],
|
||||
unique=False,
|
||||
)
|
||||
|
||||
op.create_table(
|
||||
"workflow_handoff_cancellations",
|
||||
sa.Column("task_id", sa.String(length=255), nullable=False),
|
||||
sa.Column("scope_tenant_id", models.types.StringUUID(), nullable=True),
|
||||
sa.Column("scope_app_id", models.types.StringUUID(), nullable=True),
|
||||
sa.Column("scope_created_by_role", sa.String(length=255), nullable=True),
|
||||
sa.Column("scope_created_by", models.types.StringUUID(), nullable=True),
|
||||
sa.Column("requested_at", sa.DateTime(), nullable=False),
|
||||
sa.Column("expires_at", sa.DateTime(), nullable=False),
|
||||
sa.Column("reason", models.types.LongText(), nullable=False),
|
||||
sa.Column("id", models.types.StringUUID(), nullable=False),
|
||||
sa.Column("created_at", sa.DateTime(), server_default=sa.text("CURRENT_TIMESTAMP"), nullable=False),
|
||||
sa.Column("updated_at", sa.DateTime(), server_default=sa.text("CURRENT_TIMESTAMP"), nullable=False),
|
||||
sa.CheckConstraint(
|
||||
"(scope_tenant_id IS NULL AND scope_app_id IS NULL) OR "
|
||||
"(scope_tenant_id IS NOT NULL AND scope_app_id IS NOT NULL)",
|
||||
name=op.f("workflow_handoff_cancellations_owner_scope_pair_check"),
|
||||
),
|
||||
sa.CheckConstraint(
|
||||
"(scope_created_by_role IS NULL AND scope_created_by IS NULL) OR "
|
||||
"(scope_created_by_role IS NOT NULL AND scope_created_by IS NOT NULL)",
|
||||
name=op.f("workflow_handoff_cancellations_creator_scope_pair_check"),
|
||||
),
|
||||
sa.CheckConstraint(
|
||||
"scope_created_by IS NULL OR (scope_tenant_id IS NOT NULL AND scope_app_id IS NOT NULL)",
|
||||
name=op.f("workflow_handoff_cancellations_creator_scope_app_check"),
|
||||
),
|
||||
sa.PrimaryKeyConstraint("id", name=op.f("workflow_handoff_cancellations_pkey")),
|
||||
)
|
||||
with op.batch_alter_table("workflow_handoff_cancellations", schema=None) as batch_op:
|
||||
batch_op.create_index(
|
||||
"workflow_handoff_cancellations_expires_idx",
|
||||
["expires_at"],
|
||||
unique=False,
|
||||
)
|
||||
batch_op.create_index(
|
||||
"workflow_handoff_cancellations_task_scope_idx",
|
||||
[
|
||||
"task_id",
|
||||
"scope_tenant_id",
|
||||
"scope_app_id",
|
||||
"scope_created_by_role",
|
||||
"scope_created_by",
|
||||
"expires_at",
|
||||
],
|
||||
unique=False,
|
||||
)
|
||||
|
||||
op.create_table(
|
||||
"workflow_handoff_snapshot_gc",
|
||||
sa.Column("snapshot_object_key", sa.String(length=255), nullable=False),
|
||||
sa.Column("upload_completed_at", sa.DateTime(), nullable=True),
|
||||
sa.Column("deleted_at", sa.DateTime(), nullable=True),
|
||||
sa.Column("attempts", sa.Integer(), server_default=sa.text("0"), nullable=False),
|
||||
sa.Column("next_retry_at", sa.DateTime(), nullable=True),
|
||||
sa.Column("last_error", models.types.LongText(), nullable=True),
|
||||
sa.Column("id", models.types.StringUUID(), nullable=False),
|
||||
sa.Column("created_at", sa.DateTime(), server_default=sa.text("CURRENT_TIMESTAMP"), nullable=False),
|
||||
sa.Column("updated_at", sa.DateTime(), server_default=sa.text("CURRENT_TIMESTAMP"), nullable=False),
|
||||
sa.CheckConstraint(
|
||||
"attempts >= 0",
|
||||
name=op.f("workflow_handoff_snapshot_gc_attempts_nonnegative_check"),
|
||||
),
|
||||
sa.PrimaryKeyConstraint("id", name=op.f("workflow_handoff_snapshot_gc_pkey")),
|
||||
sa.UniqueConstraint(
|
||||
"snapshot_object_key",
|
||||
name="workflow_handoff_snapshot_gc_object_key_key",
|
||||
),
|
||||
)
|
||||
with op.batch_alter_table("workflow_handoff_snapshot_gc", schema=None) as batch_op:
|
||||
batch_op.create_index(
|
||||
"workflow_handoff_snapshot_gc_pending_idx",
|
||||
["deleted_at", "next_retry_at", "created_at"],
|
||||
unique=False,
|
||||
)
|
||||
|
||||
|
||||
def downgrade():
|
||||
op.drop_table("workflow_handoff_snapshot_gc")
|
||||
op.drop_table("workflow_handoff_cancellations")
|
||||
op.drop_table("workflow_run_handoffs")
|
||||
with op.batch_alter_table("workflow_archive_logs", schema=None) as batch_op:
|
||||
batch_op.drop_column("run_handoff_duration")
|
||||
with op.batch_alter_table("workflow_runs", schema=None) as batch_op:
|
||||
batch_op.drop_column("handoff_duration")
|
||||
@@ -150,6 +150,13 @@ from .workflow import (
|
||||
WorkflowType,
|
||||
resolve_workflow_kind,
|
||||
)
|
||||
from .workflow_handoff import (
|
||||
WorkflowHandoffCancellation,
|
||||
WorkflowHandoffResumeRoute,
|
||||
WorkflowHandoffSnapshotGC,
|
||||
WorkflowHandoffState,
|
||||
WorkflowRunHandoff,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"APIBasedExtension",
|
||||
@@ -281,6 +288,10 @@ __all__ = [
|
||||
"WorkflowComment",
|
||||
"WorkflowCommentMention",
|
||||
"WorkflowCommentReply",
|
||||
"WorkflowHandoffCancellation",
|
||||
"WorkflowHandoffResumeRoute",
|
||||
"WorkflowHandoffSnapshotGC",
|
||||
"WorkflowHandoffState",
|
||||
"WorkflowKind",
|
||||
"WorkflowNodeExecutionModel",
|
||||
"WorkflowNodeExecutionOffload",
|
||||
@@ -288,6 +299,7 @@ __all__ = [
|
||||
"WorkflowPause",
|
||||
"WorkflowRun",
|
||||
"WorkflowRunArchiveBundle",
|
||||
"WorkflowRunHandoff",
|
||||
"WorkflowRunTriggeredFrom",
|
||||
"WorkflowSchedulePlan",
|
||||
"WorkflowToolProvider",
|
||||
|
||||
@@ -109,6 +109,7 @@ class WorkflowRunSummaryDict(TypedDict):
|
||||
status: str
|
||||
triggered_from: str
|
||||
elapsed_time: float
|
||||
handoff_duration: float
|
||||
total_tokens: int
|
||||
|
||||
|
||||
@@ -748,6 +749,7 @@ class WorkflowRunDict(TypedDict):
|
||||
outputs: Mapping[str, Any]
|
||||
error: str | None
|
||||
elapsed_time: float
|
||||
handoff_duration: float
|
||||
total_tokens: int
|
||||
total_steps: int
|
||||
created_by_role: CreatorUserRole
|
||||
@@ -782,6 +784,7 @@ class WorkflowRun(Base):
|
||||
- outputs (text) `optional` Output content
|
||||
- error (string) `optional` Error reason
|
||||
- elapsed_time (float) `optional` Time consumption (s)
|
||||
- handoff_duration (float) Planned worker-handoff wait included in elapsed_time (s)
|
||||
- total_tokens (int) `optional` Total tokens used
|
||||
- total_steps (int) Total steps (redundant), default 0
|
||||
- created_by_role (string) Creator role
|
||||
@@ -819,6 +822,7 @@ class WorkflowRun(Base):
|
||||
outputs: Mapped[str | None] = mapped_column(LongText, default="{}")
|
||||
error: Mapped[str | None] = mapped_column(LongText)
|
||||
elapsed_time: Mapped[float] = mapped_column(sa.Float, nullable=False, server_default=sa.text("0"))
|
||||
handoff_duration: Mapped[float] = mapped_column(sa.Float, nullable=False, default=0.0, server_default=sa.text("0"))
|
||||
total_tokens: Mapped[int] = mapped_column(sa.BigInteger, server_default=sa.text("0"))
|
||||
total_steps: Mapped[int] = mapped_column(sa.Integer, server_default=sa.text("0"), nullable=True)
|
||||
created_by_role: Mapped[CreatorUserRole] = mapped_column(EnumText(CreatorUserRole, length=255)) # account, end_user
|
||||
@@ -891,6 +895,7 @@ class WorkflowRun(Base):
|
||||
outputs=self.outputs_dict,
|
||||
error=self.error,
|
||||
elapsed_time=self.elapsed_time,
|
||||
handoff_duration=self.handoff_duration,
|
||||
total_tokens=self.total_tokens,
|
||||
total_steps=self.total_steps,
|
||||
created_by_role=self.created_by_role,
|
||||
@@ -916,6 +921,7 @@ class WorkflowRun(Base):
|
||||
outputs=json.dumps(data.get("outputs")),
|
||||
error=data.get("error"),
|
||||
elapsed_time=data.get("elapsed_time"),
|
||||
handoff_duration=data.get("handoff_duration", 0),
|
||||
total_tokens=data.get("total_tokens"),
|
||||
total_steps=data.get("total_steps"),
|
||||
created_by_role=data.get("created_by_role"),
|
||||
@@ -1441,6 +1447,12 @@ class WorkflowArchiveLog(TypeBase):
|
||||
run_exceptions_count: Mapped[int] = mapped_column(sa.Integer, server_default=sa.text("0"), nullable=True)
|
||||
|
||||
trigger_metadata: Mapped[str | None] = mapped_column(LongText, nullable=True)
|
||||
run_handoff_duration: Mapped[float] = mapped_column(
|
||||
sa.Float,
|
||||
nullable=False,
|
||||
default=0.0,
|
||||
server_default=sa.text("0"),
|
||||
)
|
||||
archived_at: Mapped[datetime] = mapped_column(
|
||||
DateTime, nullable=False, server_default=func.current_timestamp(), init=False
|
||||
)
|
||||
@@ -1452,6 +1464,7 @@ class WorkflowArchiveLog(TypeBase):
|
||||
"status": self.run_status,
|
||||
"triggered_from": self.run_triggered_from,
|
||||
"elapsed_time": self.run_elapsed_time,
|
||||
"handoff_duration": self.run_handoff_duration,
|
||||
"total_tokens": self.run_total_tokens,
|
||||
}
|
||||
|
||||
|
||||
@@ -0,0 +1,331 @@
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime
|
||||
from enum import StrEnum
|
||||
|
||||
import sqlalchemy as sa
|
||||
from sqlalchemy import Boolean, DateTime, Index, Integer, String, UniqueConstraint
|
||||
from sqlalchemy.orm import Mapped, mapped_column
|
||||
|
||||
from models.enums import CreatorUserRole
|
||||
|
||||
from .base import DefaultFieldsDCMixin, TypeBase
|
||||
from .types import EnumText, LongText, StringUUID
|
||||
|
||||
RAG_PIPELINE_SOURCE_BATCH_ID_EXTRA_KEY = "source_batch_id"
|
||||
RAG_PIPELINE_TENANT_ID_EXTRA_KEY = "tenant_id"
|
||||
RAG_PIPELINE_QUEUE_KIND_EXTRA_KEY = "queue_kind"
|
||||
RAG_PIPELINE_TENANT_ISOLATED_EXTRA_KEY = "tenant_isolated"
|
||||
|
||||
|
||||
class WorkflowHandoffState(StrEnum):
|
||||
"""Durable states for transferring a workflow execution between workers."""
|
||||
|
||||
PREPARING = "preparing"
|
||||
PREPARED = "prepared"
|
||||
READY = "ready"
|
||||
CLAIMED = "claimed"
|
||||
RESUMED = "resumed"
|
||||
FAILED = "failed"
|
||||
|
||||
|
||||
class WorkflowHandoffResumeRoute(StrEnum):
|
||||
"""Execution entry point that must rebuild the workflow from a checkpoint."""
|
||||
|
||||
WORKFLOW = "workflow"
|
||||
SNIPPET = "snippet"
|
||||
ADVANCED_CHAT = "advanced_chat"
|
||||
TRIGGERED_WORKFLOW = "triggered_workflow"
|
||||
RAG_PIPELINE = "rag_pipeline"
|
||||
|
||||
|
||||
class RagPipelineQueueKind(StrEnum):
|
||||
"""Celery lane owning a RAG pipeline source batch."""
|
||||
|
||||
REGULAR = "regular"
|
||||
PRIORITY = "priority"
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class RagPipelineHandoffGroupIdentity:
|
||||
source_batch_id: str
|
||||
tenant_id: str
|
||||
queue_kind: RagPipelineQueueKind
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class RagPipelineHandoffGroupMetadata:
|
||||
"""Durable tenant-slot ownership carried by a RAG handoff checkpoint."""
|
||||
|
||||
source_batch_id: str
|
||||
tenant_id: str
|
||||
queue_kind: RagPipelineQueueKind
|
||||
dataset_id: str
|
||||
document_id: str | None
|
||||
tenant_isolated: bool
|
||||
|
||||
@property
|
||||
def identity(self) -> RagPipelineHandoffGroupIdentity:
|
||||
return RagPipelineHandoffGroupIdentity(
|
||||
source_batch_id=self.source_batch_id,
|
||||
tenant_id=self.tenant_id,
|
||||
queue_kind=self.queue_kind,
|
||||
)
|
||||
|
||||
|
||||
class WorkflowRunHandoff(DefaultFieldsDCMixin, TypeBase):
|
||||
"""A durable checkpoint handoff created during a planned worker drain.
|
||||
|
||||
A workflow run may be handed off multiple times. ``generation`` fences commands
|
||||
and resume attempts from older handoffs, while ``lease_token`` fences retries of
|
||||
the same generation after a lease expires. The checkpoint object is written to
|
||||
shared storage before this row transitions into ``READY``.
|
||||
"""
|
||||
|
||||
__tablename__ = "workflow_run_handoffs"
|
||||
__table_args__ = (
|
||||
UniqueConstraint(
|
||||
"workflow_run_id",
|
||||
"generation",
|
||||
name="workflow_run_handoffs_run_generation_key",
|
||||
),
|
||||
sa.CheckConstraint("generation > 0", name="generation_positive"),
|
||||
sa.CheckConstraint("attempts >= 0", name="attempts_nonnegative"),
|
||||
sa.CheckConstraint("terminal_attempts >= 0", name="terminal_attempts_nonnegative"),
|
||||
sa.CheckConstraint("snapshot_size_bytes >= 0", name="snapshot_size_nonnegative"),
|
||||
sa.CheckConstraint(
|
||||
"state IN ('preparing', 'prepared', 'ready', 'claimed', 'resumed', 'failed')",
|
||||
name="state_valid",
|
||||
),
|
||||
sa.CheckConstraint(
|
||||
"resume_route IN ('workflow', 'snippet', 'advanced_chat', 'triggered_workflow', 'rag_pipeline')",
|
||||
name="resume_route_valid",
|
||||
),
|
||||
sa.CheckConstraint(
|
||||
"state <> 'claimed' OR "
|
||||
"(lease_owner IS NOT NULL AND lease_token IS NOT NULL AND lease_expires_at IS NOT NULL)",
|
||||
name="claim_lease_present",
|
||||
),
|
||||
sa.CheckConstraint(
|
||||
"state <> 'resumed' OR resumed_at IS NOT NULL",
|
||||
name="resumed_at_present",
|
||||
),
|
||||
sa.CheckConstraint(
|
||||
"state <> 'failed' OR failed_at IS NOT NULL",
|
||||
name="failed_at_present",
|
||||
),
|
||||
sa.CheckConstraint(
|
||||
"(rag_source_batch_id IS NULL AND rag_tenant_id IS NULL AND rag_queue_kind IS NULL "
|
||||
"AND rag_dataset_id IS NULL AND rag_tenant_isolated IS NULL) OR "
|
||||
"(rag_source_batch_id IS NOT NULL AND rag_tenant_id IS NOT NULL AND rag_queue_kind IS NOT NULL "
|
||||
"AND rag_dataset_id IS NOT NULL AND rag_tenant_isolated IS NOT NULL)",
|
||||
name="rag_group_metadata_complete",
|
||||
),
|
||||
sa.CheckConstraint(
|
||||
"rag_queue_kind IS NULL OR rag_queue_kind IN ('regular', 'priority')",
|
||||
name="rag_queue_kind_valid",
|
||||
),
|
||||
sa.CheckConstraint(
|
||||
"rag_tenant_slot_released_at IS NULL OR rag_group_sealed_at IS NOT NULL",
|
||||
name="rag_release_requires_seal",
|
||||
),
|
||||
Index("workflow_run_handoffs_run_state_idx", "workflow_run_id", "state"),
|
||||
Index("workflow_run_handoffs_snapshot_object_key_idx", "snapshot_object_key"),
|
||||
Index("workflow_run_handoffs_task_state_idx", "task_id", "state"),
|
||||
Index("workflow_run_handoffs_state_created_idx", "state", "created_at", "id"),
|
||||
Index("workflow_run_handoffs_resumed_retention_idx", "state", "resumed_at", "id"),
|
||||
Index("workflow_run_handoffs_failed_retention_idx", "state", "failed_at", "id"),
|
||||
Index(
|
||||
"workflow_run_handoffs_terminal_idx",
|
||||
"state",
|
||||
"terminal_compensated_at",
|
||||
"terminal_event_published_at",
|
||||
"created_at",
|
||||
),
|
||||
Index(
|
||||
"workflow_run_handoffs_rag_group_release_idx",
|
||||
"rag_source_batch_id",
|
||||
"rag_tenant_id",
|
||||
"rag_queue_kind",
|
||||
"rag_group_sealed_at",
|
||||
"rag_tenant_slot_released_at",
|
||||
),
|
||||
Index(
|
||||
"workflow_run_handoffs_rag_reconcile_idx",
|
||||
"rag_group_sealed_at",
|
||||
"rag_tenant_slot_released_at",
|
||||
"rag_source_batch_id",
|
||||
),
|
||||
Index(
|
||||
"workflow_run_handoffs_dispatch_idx",
|
||||
"state",
|
||||
"next_retry_at",
|
||||
"lease_expires_at",
|
||||
"dispatched_at",
|
||||
"created_at",
|
||||
),
|
||||
)
|
||||
|
||||
workflow_run_id: Mapped[str] = mapped_column(StringUUID, nullable=False)
|
||||
generation: Mapped[int] = mapped_column(Integer, nullable=False)
|
||||
task_id: Mapped[str] = mapped_column(String(255), nullable=False)
|
||||
|
||||
snapshot_object_key: Mapped[str] = mapped_column(String(255), nullable=False)
|
||||
snapshot_schema_version: Mapped[str] = mapped_column(String(64), nullable=False)
|
||||
snapshot_checksum: Mapped[str] = mapped_column(String(128), nullable=False)
|
||||
snapshot_size_bytes: Mapped[int] = mapped_column(sa.BigInteger, nullable=False)
|
||||
|
||||
resume_route: Mapped[WorkflowHandoffResumeRoute] = mapped_column(
|
||||
EnumText(WorkflowHandoffResumeRoute, length=64),
|
||||
nullable=False,
|
||||
)
|
||||
source_worker_id: Mapped[str] = mapped_column(String(255), nullable=False)
|
||||
|
||||
rag_source_batch_id: Mapped[str | None] = mapped_column(String(255), nullable=True, default=None)
|
||||
rag_tenant_id: Mapped[str | None] = mapped_column(StringUUID, nullable=True, default=None)
|
||||
rag_queue_kind: Mapped[RagPipelineQueueKind | None] = mapped_column(
|
||||
EnumText(RagPipelineQueueKind, length=32), nullable=True, default=None
|
||||
)
|
||||
rag_dataset_id: Mapped[str | None] = mapped_column(StringUUID, nullable=True, default=None)
|
||||
rag_document_id: Mapped[str | None] = mapped_column(StringUUID, nullable=True, default=None)
|
||||
rag_tenant_isolated: Mapped[bool | None] = mapped_column(Boolean, nullable=True, default=None)
|
||||
rag_group_sealed_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True, default=None)
|
||||
rag_tenant_slot_released_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True, default=None)
|
||||
rag_document_error_marked_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True, default=None)
|
||||
|
||||
state: Mapped[WorkflowHandoffState] = mapped_column(
|
||||
EnumText(WorkflowHandoffState, length=32),
|
||||
nullable=False,
|
||||
default=WorkflowHandoffState.PREPARED,
|
||||
server_default=sa.text("'prepared'"),
|
||||
)
|
||||
target_worker_id: Mapped[str | None] = mapped_column(String(255), nullable=True, default=None)
|
||||
lease_owner: Mapped[str | None] = mapped_column(String(255), nullable=True, default=None)
|
||||
lease_token: Mapped[str | None] = mapped_column(StringUUID, nullable=True, default=None)
|
||||
lease_expires_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True, default=None)
|
||||
|
||||
attempts: Mapped[int] = mapped_column(
|
||||
Integer,
|
||||
nullable=False,
|
||||
default=0,
|
||||
server_default=sa.text("0"),
|
||||
)
|
||||
next_retry_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True, default=None)
|
||||
dispatched_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True, default=None)
|
||||
last_error: Mapped[str | None] = mapped_column(LongText, nullable=True, default=None)
|
||||
cancel_requested_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True, default=None)
|
||||
claimed_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True, default=None)
|
||||
resumed_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True, default=None)
|
||||
failed_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True, default=None)
|
||||
terminal_compensated_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True, default=None)
|
||||
terminal_event_published_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True, default=None)
|
||||
terminal_attempts: Mapped[int] = mapped_column(
|
||||
Integer,
|
||||
nullable=False,
|
||||
default=0,
|
||||
server_default=sa.text("0"),
|
||||
)
|
||||
terminal_last_error: Mapped[str | None] = mapped_column(LongText, nullable=True, default=None)
|
||||
|
||||
|
||||
class WorkflowHandoffSnapshotGC(DefaultFieldsDCMixin, TypeBase):
|
||||
"""Durable object-GC reference independent of WorkflowRun retention.
|
||||
|
||||
The row is created in the PREPARING intent transaction, before object
|
||||
storage is touched. It therefore remains the last durable reference even
|
||||
when a workflow run and its operational handoff rows are later removed.
|
||||
"""
|
||||
|
||||
__tablename__ = "workflow_handoff_snapshot_gc"
|
||||
__table_args__ = (
|
||||
UniqueConstraint(
|
||||
"snapshot_object_key",
|
||||
name="workflow_handoff_snapshot_gc_object_key_key",
|
||||
),
|
||||
sa.CheckConstraint("attempts >= 0", name="attempts_nonnegative"),
|
||||
Index(
|
||||
"workflow_handoff_snapshot_gc_pending_idx",
|
||||
"deleted_at",
|
||||
"next_retry_at",
|
||||
"created_at",
|
||||
),
|
||||
)
|
||||
|
||||
snapshot_object_key: Mapped[str] = mapped_column(String(255), nullable=False)
|
||||
upload_completed_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True, default=None)
|
||||
deleted_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True, default=None)
|
||||
attempts: Mapped[int] = mapped_column(
|
||||
Integer,
|
||||
nullable=False,
|
||||
default=0,
|
||||
server_default=sa.text("0"),
|
||||
)
|
||||
next_retry_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True, default=None)
|
||||
last_error: Mapped[str | None] = mapped_column(LongText, nullable=True, default=None)
|
||||
|
||||
|
||||
class WorkflowHandoffCancellation(DefaultFieldsDCMixin, TypeBase):
|
||||
"""Durable Stop tombstone fencing handoff preparation by task and owner scope.
|
||||
|
||||
A Stop request can arrive after GraphEngine accepts a maintenance pause but
|
||||
before the checkpoint row exists. This tombstone is committed independently
|
||||
of object storage so a later preparation transaction can observe the Stop
|
||||
and terminalize the parent run instead of reviving it.
|
||||
|
||||
``scope_tenant_id``/``scope_app_id`` are nullable only for trusted internal
|
||||
cancellation paths. Public paths should persist the authorized owner scope.
|
||||
"""
|
||||
|
||||
__tablename__ = "workflow_handoff_cancellations"
|
||||
__table_args__ = (
|
||||
sa.CheckConstraint(
|
||||
"(scope_tenant_id IS NULL AND scope_app_id IS NULL) OR "
|
||||
"(scope_tenant_id IS NOT NULL AND scope_app_id IS NOT NULL)",
|
||||
name="owner_scope_pair",
|
||||
),
|
||||
sa.CheckConstraint(
|
||||
"(scope_created_by_role IS NULL AND scope_created_by IS NULL) OR "
|
||||
"(scope_created_by_role IS NOT NULL AND scope_created_by IS NOT NULL)",
|
||||
name="creator_scope_pair",
|
||||
),
|
||||
sa.CheckConstraint(
|
||||
"scope_created_by IS NULL OR (scope_tenant_id IS NOT NULL AND scope_app_id IS NOT NULL)",
|
||||
name="creator_scope_app",
|
||||
),
|
||||
Index(
|
||||
"workflow_handoff_cancellations_task_scope_idx",
|
||||
"task_id",
|
||||
"scope_tenant_id",
|
||||
"scope_app_id",
|
||||
"scope_created_by_role",
|
||||
"scope_created_by",
|
||||
"expires_at",
|
||||
),
|
||||
Index("workflow_handoff_cancellations_expires_idx", "expires_at"),
|
||||
)
|
||||
|
||||
task_id: Mapped[str] = mapped_column(String(255), nullable=False)
|
||||
requested_at: Mapped[datetime] = mapped_column(DateTime, nullable=False)
|
||||
expires_at: Mapped[datetime] = mapped_column(DateTime, nullable=False)
|
||||
reason: Mapped[str] = mapped_column(LongText, nullable=False)
|
||||
scope_tenant_id: Mapped[str | None] = mapped_column(StringUUID, nullable=True, default=None)
|
||||
scope_app_id: Mapped[str | None] = mapped_column(StringUUID, nullable=True, default=None)
|
||||
scope_created_by_role: Mapped[CreatorUserRole | None] = mapped_column(
|
||||
EnumText(CreatorUserRole, length=255), nullable=True, default=None
|
||||
)
|
||||
scope_created_by: Mapped[str | None] = mapped_column(StringUUID, nullable=True, default=None)
|
||||
|
||||
|
||||
__all__ = [
|
||||
"RAG_PIPELINE_QUEUE_KIND_EXTRA_KEY",
|
||||
"RAG_PIPELINE_SOURCE_BATCH_ID_EXTRA_KEY",
|
||||
"RAG_PIPELINE_TENANT_ID_EXTRA_KEY",
|
||||
"RAG_PIPELINE_TENANT_ISOLATED_EXTRA_KEY",
|
||||
"RagPipelineHandoffGroupIdentity",
|
||||
"RagPipelineHandoffGroupMetadata",
|
||||
"RagPipelineQueueKind",
|
||||
"WorkflowHandoffCancellation",
|
||||
"WorkflowHandoffResumeRoute",
|
||||
"WorkflowHandoffSnapshotGC",
|
||||
"WorkflowHandoffState",
|
||||
"WorkflowRunHandoff",
|
||||
]
|
||||
@@ -8087,6 +8087,20 @@ Update account-level Step-by-step Tour state
|
||||
| ---- | ----------- | ------ |
|
||||
| 200 | Workflow run detail retrieved successfully | **application/json**: [WorkflowRunDetailResponse](#workflowrundetailresponse)<br> |
|
||||
|
||||
### [GET] /rag/pipelines/{pipeline_id}/workflow-runs/{run_id}/events
|
||||
#### Parameters
|
||||
|
||||
| Name | Located in | Description | Required | Schema |
|
||||
| ---- | ---------- | ----------- | -------- | ------ |
|
||||
| pipeline_id | path | | Yes | string (uuid) |
|
||||
| run_id | path | | Yes | string (uuid) |
|
||||
|
||||
#### Responses
|
||||
|
||||
| Code | Description |
|
||||
| ---- | ----------- |
|
||||
| 200 | Success |
|
||||
|
||||
### [GET] /rag/pipelines/{pipeline_id}/workflow-runs/{run_id}/node-executions
|
||||
**Get workflow run node execution list**
|
||||
|
||||
@@ -8851,6 +8865,20 @@ command channel for backward compatibility.
|
||||
| 200 | Workflow run detail retrieved successfully | **application/json**: [WorkflowRunDetailResponse](#workflowrundetailresponse)<br> |
|
||||
| 404 | Workflow run not found | |
|
||||
|
||||
### [GET] /snippets/{snippet_id}/workflow-runs/{run_id}/events
|
||||
#### Parameters
|
||||
|
||||
| Name | Located in | Description | Required | Schema |
|
||||
| ---- | ---------- | ----------- | -------- | ------ |
|
||||
| run_id | path | | Yes | string (uuid) |
|
||||
| snippet_id | path | | Yes | string (uuid) |
|
||||
|
||||
#### Responses
|
||||
|
||||
| Code | Description |
|
||||
| ---- | ----------- |
|
||||
| 200 | Success |
|
||||
|
||||
### [GET] /snippets/{snippet_id}/workflow-runs/{run_id}/node-executions
|
||||
**List node executions for a workflow run**
|
||||
|
||||
@@ -13186,6 +13214,7 @@ Model class for AI model.
|
||||
| elapsed_time | number | | No |
|
||||
| exceptions_count | integer | | No |
|
||||
| finished_at | integer | | No |
|
||||
| handoff_duration | number | | No |
|
||||
| id | string | | Yes |
|
||||
| message_id | string | | No |
|
||||
| retry_index | integer | | No |
|
||||
@@ -23987,6 +24016,7 @@ Lifecycle state for an asynchronous archive download request.
|
||||
| exceptions_count | integer | | No |
|
||||
| finished_at | integer | | No |
|
||||
| graph | | | Yes |
|
||||
| handoff_duration | number | | No |
|
||||
| id | string | | Yes |
|
||||
| inputs | | | Yes |
|
||||
| outputs | | | Yes |
|
||||
@@ -24008,6 +24038,7 @@ Lifecycle state for an asynchronous archive download request.
|
||||
| Name | Type | Description | Required |
|
||||
| ---- | ---- | ----------- | -------- |
|
||||
| elapsed_time | number | | No |
|
||||
| handoff_duration | number | | No |
|
||||
| id | string | | Yes |
|
||||
| status | string | | No |
|
||||
| total_tokens | integer | | No |
|
||||
@@ -24022,6 +24053,7 @@ Lifecycle state for an asynchronous archive download request.
|
||||
| elapsed_time | number | | No |
|
||||
| exceptions_count | integer | | No |
|
||||
| finished_at | integer | | No |
|
||||
| handoff_duration | number | | No |
|
||||
| id | string | | Yes |
|
||||
| retry_index | integer | | No |
|
||||
| status | string | | No |
|
||||
@@ -24038,6 +24070,7 @@ Lifecycle state for an asynchronous archive download request.
|
||||
| error | string | | No |
|
||||
| exceptions_count | integer | | No |
|
||||
| finished_at | integer | | No |
|
||||
| handoff_duration | number | | No |
|
||||
| id | string | | Yes |
|
||||
| status | string | | No |
|
||||
| total_steps | integer | | No |
|
||||
|
||||
@@ -202,6 +202,7 @@ Upload a file to use as an input variable when running the app
|
||||
| Name | Located in | Description | Required | Schema |
|
||||
| ---- | ---------- | ----------- | -------- | ------ |
|
||||
| continue_on_pause | query | Whether to keep the event stream open on pause | No | boolean |
|
||||
| cursor | query | Replay events strictly after this SSE event ID. Last-Event-ID header takes precedence. | No | string |
|
||||
| include_state_snapshot | query | Whether to include workflow state snapshots | No | boolean |
|
||||
| app_id | path | | Yes | string |
|
||||
| task_id | path | | Yes | string |
|
||||
@@ -1058,6 +1059,7 @@ types it as a required `'success'` rather than an optional field.
|
||||
| elapsed_time | number | | No |
|
||||
| error | string | | No |
|
||||
| finished_at | integer | | No |
|
||||
| handoff_duration | number | | No |
|
||||
| id | string | | Yes |
|
||||
| outputs | object | | No |
|
||||
| status | string | | Yes |
|
||||
|
||||
@@ -74,6 +74,21 @@ Deprecated legacy alias for updating an existing document by providing text cont
|
||||
| 403 | Forbidden - dataset API access or workspace access denied | |
|
||||
| 404 | Document not found | |
|
||||
|
||||
### [GET] /datasets/{dataset_id}/pipeline/workflow-runs/{run_id}/events
|
||||
#### Parameters
|
||||
|
||||
| Name | Located in | Description | Required | Schema |
|
||||
| ---- | ---------- | ----------- | -------- | ------ |
|
||||
| dataset_id | path | | Yes | string (uuid) |
|
||||
| run_id | path | | Yes | string (uuid) |
|
||||
|
||||
#### Responses
|
||||
|
||||
| Code | Description |
|
||||
| ---- | ----------- |
|
||||
| 401 | Unauthorized - invalid API token |
|
||||
| 403 | Forbidden - dataset API access or workspace access denied |
|
||||
|
||||
---
|
||||
## default
|
||||
|
||||
@@ -393,6 +408,7 @@ Resume the Server-Sent Events stream for a workflow run after a pause or a dropp
|
||||
| ---- | ---------- | ----------- | -------- | ------ |
|
||||
| task_id | path | Workflow run ID returned by the original workflow run request. | Yes | string |
|
||||
| continue_on_pause | query | Set to `true` to keep the stream open across multiple `workflow_paused` events, which is useful when the workflow has more than one Human Input node in sequence. By default, the stream closes after the first pause. | No | boolean |
|
||||
| cursor | query | Replay events strictly after this SSE event ID. Last-Event-ID header takes precedence. | No | string |
|
||||
| include_state_snapshot | query | When `true`, replay from the persisted state snapshot to include a status summary of already-executed nodes before streaming new events. | No | boolean |
|
||||
| user | query | End-user identifier that originally triggered the run. Must match the creator of the run. | Yes | string |
|
||||
|
||||
@@ -400,7 +416,7 @@ Resume the Server-Sent Events stream for a workflow run after a pause or a dropp
|
||||
|
||||
| Code | Description | Schema |
|
||||
| ---- | ----------- | ------ |
|
||||
| 200 | Server-Sent Events stream. Each event is delivered as `data: {JSON}\\n\\n`. Event payloads follow the same schemas as the original streaming response. | **text/event-stream**: [EventStreamResponse](#eventstreamresponse)<br> |
|
||||
| 200 | Server-Sent Events stream. Durable events are delivered as `id: {cursor}\\ndata: {JSON}\\n\\n`; reconnect with Last-Event-ID. Event payloads follow the same schemas as the original streaming response. | **text/event-stream**: [EventStreamResponse](#eventstreamresponse)<br> |
|
||||
| 400 | `not_workflow_app` : Please check if your app mode matches the right API route. | |
|
||||
| 401 | Unauthorized - invalid API token | |
|
||||
| 403 | Forbidden - token scope, app, dataset, or workspace access denied | |
|
||||
@@ -2101,6 +2117,7 @@ Resume the Server-Sent Events stream for a workflow run after a pause or a dropp
|
||||
| ---- | ---------- | ----------- | -------- | ------ |
|
||||
| task_id | path | Workflow run ID returned by the original workflow run request. | Yes | string |
|
||||
| continue_on_pause | query | Set to `true` to keep the stream open across multiple `workflow_paused` events, which is useful when the workflow has more than one Human Input node in sequence. By default, the stream closes after the first pause. | No | boolean |
|
||||
| cursor | query | Replay events strictly after this SSE event ID. Last-Event-ID header takes precedence. | No | string |
|
||||
| include_state_snapshot | query | When `true`, replay from the persisted state snapshot to include a status summary of already-executed nodes before streaming new events. | No | boolean |
|
||||
| user | query | End-user identifier that originally triggered the run. Must match the creator of the run. | Yes | string |
|
||||
|
||||
@@ -2108,7 +2125,7 @@ Resume the Server-Sent Events stream for a workflow run after a pause or a dropp
|
||||
|
||||
| Code | Description | Schema |
|
||||
| ---- | ----------- | ------ |
|
||||
| 200 | Server-Sent Events stream. Each event is delivered as `data: {JSON}\\n\\n`. Event payloads follow the same schemas as the original streaming response. | **text/event-stream**: [EventStreamResponse](#eventstreamresponse)<br> |
|
||||
| 200 | Server-Sent Events stream. Durable events are delivered as `id: {cursor}\\ndata: {JSON}\\n\\n`; reconnect with Last-Event-ID. Event payloads follow the same schemas as the original streaming response. | **text/event-stream**: [EventStreamResponse](#eventstreamresponse)<br> |
|
||||
| 400 | `not_workflow_app` : Please check if your app mode matches the right API route. | |
|
||||
| 401 | Unauthorized - invalid API token | |
|
||||
| 403 | Forbidden - token scope, app, dataset, or workspace access denied | |
|
||||
@@ -4108,6 +4125,7 @@ in form definition, or a variable while the workflow is running.
|
||||
| Name | Type | Description | Required |
|
||||
| ---- | ---- | ----------- | -------- |
|
||||
| continue_on_pause | boolean | Set to `true` to keep the stream open across multiple `workflow_paused` events, which is useful when the workflow has more than one Human Input node in sequence. By default, the stream closes after the first pause. | No |
|
||||
| cursor | string | Replay events strictly after this SSE event ID. Last-Event-ID header takes precedence. | No |
|
||||
| include_state_snapshot | boolean | When `true`, replay from the persisted state snapshot to include a status summary of already-executed nodes before streaming new events. | No |
|
||||
| user | string | End-user identifier that originally triggered the run. Must match the creator of the run. | Yes |
|
||||
|
||||
@@ -4133,6 +4151,7 @@ in form definition, or a variable while the workflow is running.
|
||||
| error | string | | No |
|
||||
| exceptions_count | integer | | No |
|
||||
| finished_at | integer | | No |
|
||||
| handoff_duration | number | | No |
|
||||
| id | string | | Yes |
|
||||
| status | string | | No |
|
||||
| total_steps | integer | | No |
|
||||
@@ -4165,6 +4184,7 @@ in form definition, or a variable while the workflow is running.
|
||||
| elapsed_time | number<br>integer | | No |
|
||||
| error | string | | No |
|
||||
| finished_at | integer | | No |
|
||||
| handoff_duration | number | | No |
|
||||
| id | string | | Yes |
|
||||
| inputs | object<br>[ object ]<br>string<br>integer<br>number<br>boolean | | No |
|
||||
| outputs | object | | No |
|
||||
|
||||
@@ -0,0 +1,276 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Sequence
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime
|
||||
from typing import Protocol, cast, override
|
||||
|
||||
from sqlalchemy import and_, exists, or_, select
|
||||
from sqlalchemy.orm import Session, sessionmaker
|
||||
|
||||
from graphon.enums import WorkflowExecutionStatus
|
||||
from models.dataset import Document
|
||||
from models.enums import IndexingStatus
|
||||
from models.workflow import WorkflowRun
|
||||
from models.workflow_handoff import (
|
||||
RagPipelineHandoffGroupIdentity,
|
||||
WorkflowHandoffState,
|
||||
WorkflowRunHandoff,
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class RagPipelineHandoffGroupSnapshot:
|
||||
identity: RagPipelineHandoffGroupIdentity
|
||||
sealed_at: datetime | None
|
||||
released_at: datetime | None
|
||||
tenant_isolated: bool
|
||||
has_running_workflow_runs: bool
|
||||
|
||||
|
||||
class RagPipelineHandoffGroupRepository(Protocol):
|
||||
def seal_group(self, *, identity: RagPipelineHandoffGroupIdentity, sealed_at: datetime) -> int: ...
|
||||
|
||||
def get_group(self, identity: RagPipelineHandoffGroupIdentity) -> RagPipelineHandoffGroupSnapshot | None: ...
|
||||
|
||||
def list_reconcilable_groups(self, *, limit: int) -> Sequence[RagPipelineHandoffGroupIdentity]: ...
|
||||
|
||||
def mark_failed_documents(self, *, identity: RagPipelineHandoffGroupIdentity, marked_at: datetime) -> int: ...
|
||||
|
||||
def mark_released_once(self, *, identity: RagPipelineHandoffGroupIdentity, released_at: datetime) -> bool: ...
|
||||
|
||||
|
||||
class SQLAlchemyRagPipelineHandoffGroupRepository(RagPipelineHandoffGroupRepository):
|
||||
"""Database half of RAG tenant-slot handoff ownership.
|
||||
|
||||
Redis/Celery side effects stay in the service. Group lifecycle markers live
|
||||
beside every handoff generation so the periodic scanner can recover after a
|
||||
worker exits between the handoff and tenant-slot release.
|
||||
"""
|
||||
|
||||
def __init__(self, session_factory: sessionmaker[Session]) -> None:
|
||||
self._session_factory = session_factory
|
||||
|
||||
@override
|
||||
def seal_group(self, *, identity: RagPipelineHandoffGroupIdentity, sealed_at: datetime) -> int:
|
||||
with self._session_factory.begin() as session:
|
||||
rows = self._locked_group_rows(session, identity)
|
||||
for row in rows:
|
||||
if row.rag_group_sealed_at is None:
|
||||
row.rag_group_sealed_at = sealed_at
|
||||
session.flush()
|
||||
return len(rows)
|
||||
|
||||
@override
|
||||
def get_group(self, identity: RagPipelineHandoffGroupIdentity) -> RagPipelineHandoffGroupSnapshot | None:
|
||||
with self._session_factory() as session:
|
||||
rows = list(session.scalars(select(WorkflowRunHandoff).where(*self._group_filters(identity))))
|
||||
if not rows:
|
||||
return None
|
||||
self._validate_group_rows(rows)
|
||||
run_ids = sorted({row.workflow_run_id for row in rows})
|
||||
has_running = session.scalar(
|
||||
select(
|
||||
exists().where(
|
||||
WorkflowRun.id.in_(run_ids),
|
||||
WorkflowRun.status == WorkflowExecutionStatus.RUNNING,
|
||||
)
|
||||
)
|
||||
)
|
||||
sealed_values = [row.rag_group_sealed_at for row in rows if row.rag_group_sealed_at is not None]
|
||||
released_values = [
|
||||
row.rag_tenant_slot_released_at for row in rows if row.rag_tenant_slot_released_at is not None
|
||||
]
|
||||
return RagPipelineHandoffGroupSnapshot(
|
||||
identity=identity,
|
||||
sealed_at=max(sealed_values) if len(sealed_values) == len(rows) else None,
|
||||
released_at=max(released_values) if len(released_values) == len(rows) else None,
|
||||
tenant_isolated=bool(rows[0].rag_tenant_isolated),
|
||||
has_running_workflow_runs=bool(has_running),
|
||||
)
|
||||
|
||||
@override
|
||||
def list_reconcilable_groups(self, *, limit: int) -> Sequence[RagPipelineHandoffGroupIdentity]:
|
||||
if limit <= 0:
|
||||
raise ValueError("limit must be positive")
|
||||
latest = WorkflowRunHandoff.__table__.alias("latest_rag_handoff")
|
||||
failed_document_pending = and_(
|
||||
WorkflowRunHandoff.state == WorkflowHandoffState.FAILED,
|
||||
WorkflowRunHandoff.rag_document_id.is_not(None),
|
||||
WorkflowRunHandoff.rag_document_error_marked_at.is_(None),
|
||||
~exists(
|
||||
select(1).where(
|
||||
latest.c.workflow_run_id == WorkflowRunHandoff.workflow_run_id,
|
||||
latest.c.generation > WorkflowRunHandoff.generation,
|
||||
)
|
||||
),
|
||||
)
|
||||
candidates = (
|
||||
select(
|
||||
WorkflowRunHandoff.rag_source_batch_id,
|
||||
WorkflowRunHandoff.rag_tenant_id,
|
||||
WorkflowRunHandoff.rag_queue_kind,
|
||||
)
|
||||
.where(
|
||||
WorkflowRunHandoff.rag_source_batch_id.is_not(None),
|
||||
WorkflowRunHandoff.rag_tenant_id.is_not(None),
|
||||
WorkflowRunHandoff.rag_queue_kind.is_not(None),
|
||||
or_(
|
||||
and_(
|
||||
WorkflowRunHandoff.rag_group_sealed_at.is_not(None),
|
||||
or_(WorkflowRunHandoff.rag_tenant_slot_released_at.is_(None), failed_document_pending),
|
||||
),
|
||||
WorkflowRunHandoff.rag_group_sealed_at.is_(None),
|
||||
),
|
||||
)
|
||||
.distinct()
|
||||
.subquery("rag_handoff_group_candidates")
|
||||
)
|
||||
group_member = WorkflowRunHandoff.__table__.alias("rag_handoff_group_member")
|
||||
group_run = WorkflowRun.__table__.alias("rag_handoff_group_run")
|
||||
group_has_running_run = exists(
|
||||
select(1)
|
||||
.select_from(group_member.join(group_run, group_run.c.id == group_member.c.workflow_run_id))
|
||||
.where(
|
||||
group_member.c.rag_source_batch_id == candidates.c.rag_source_batch_id,
|
||||
group_member.c.rag_tenant_id == candidates.c.rag_tenant_id,
|
||||
group_member.c.rag_queue_kind == candidates.c.rag_queue_kind,
|
||||
group_run.c.status == WorkflowExecutionStatus.RUNNING,
|
||||
)
|
||||
)
|
||||
stmt = (
|
||||
select(
|
||||
candidates.c.rag_source_batch_id,
|
||||
candidates.c.rag_tenant_id,
|
||||
candidates.c.rag_queue_kind,
|
||||
)
|
||||
.order_by(
|
||||
group_has_running_run,
|
||||
candidates.c.rag_source_batch_id,
|
||||
candidates.c.rag_tenant_id,
|
||||
candidates.c.rag_queue_kind,
|
||||
)
|
||||
.limit(limit)
|
||||
)
|
||||
with self._session_factory() as session:
|
||||
return [
|
||||
RagPipelineHandoffGroupIdentity(
|
||||
source_batch_id=cast(str, source_batch_id),
|
||||
tenant_id=cast(str, tenant_id),
|
||||
queue_kind=queue_kind,
|
||||
)
|
||||
for source_batch_id, tenant_id, queue_kind in session.execute(stmt).tuples()
|
||||
]
|
||||
|
||||
@override
|
||||
def mark_failed_documents(self, *, identity: RagPipelineHandoffGroupIdentity, marked_at: datetime) -> int:
|
||||
newer = WorkflowRunHandoff.__table__.alias("newer_rag_handoff")
|
||||
latest_generation = ~exists(
|
||||
select(1).where(
|
||||
newer.c.workflow_run_id == WorkflowRunHandoff.workflow_run_id,
|
||||
newer.c.generation > WorkflowRunHandoff.generation,
|
||||
)
|
||||
)
|
||||
with self._session_factory.begin() as session:
|
||||
failed_rows = list(
|
||||
session.scalars(
|
||||
select(WorkflowRunHandoff)
|
||||
.where(
|
||||
*self._group_filters(identity),
|
||||
latest_generation,
|
||||
WorkflowRunHandoff.state == WorkflowHandoffState.FAILED,
|
||||
WorkflowRunHandoff.rag_document_id.is_not(None),
|
||||
WorkflowRunHandoff.rag_document_error_marked_at.is_(None),
|
||||
)
|
||||
.order_by(WorkflowRunHandoff.workflow_run_id, WorkflowRunHandoff.generation)
|
||||
.with_for_update()
|
||||
)
|
||||
)
|
||||
if not failed_rows:
|
||||
return 0
|
||||
document_ids = sorted({cast(str, row.rag_document_id) for row in failed_rows})
|
||||
dataset_ids = sorted({cast(str, row.rag_dataset_id) for row in failed_rows})
|
||||
documents = {
|
||||
(document.dataset_id, document.id): document
|
||||
for document in session.scalars(
|
||||
select(Document)
|
||||
.where(
|
||||
Document.id.in_(document_ids),
|
||||
Document.dataset_id.in_(dataset_ids),
|
||||
Document.tenant_id == identity.tenant_id,
|
||||
)
|
||||
.order_by(Document.id)
|
||||
.with_for_update()
|
||||
)
|
||||
}
|
||||
for row in failed_rows:
|
||||
document = documents.get((cast(str, row.rag_dataset_id), cast(str, row.rag_document_id)))
|
||||
if document is not None and document.indexing_status != IndexingStatus.COMPLETED:
|
||||
document.indexing_status = IndexingStatus.ERROR
|
||||
document.error = row.last_error or "RAG pipeline handoff permanently failed"
|
||||
document.stopped_at = marked_at
|
||||
row.rag_document_error_marked_at = marked_at
|
||||
session.flush()
|
||||
return len(failed_rows)
|
||||
|
||||
@override
|
||||
def mark_released_once(self, *, identity: RagPipelineHandoffGroupIdentity, released_at: datetime) -> bool:
|
||||
with self._session_factory.begin() as session:
|
||||
rows = self._locked_group_rows(session, identity)
|
||||
if not rows or any(row.rag_tenant_slot_released_at is not None for row in rows):
|
||||
return False
|
||||
if any(row.rag_group_sealed_at is None for row in rows):
|
||||
return False
|
||||
run_ids = sorted({row.workflow_run_id for row in rows})
|
||||
has_running = session.scalar(
|
||||
select(
|
||||
exists().where(
|
||||
WorkflowRun.id.in_(run_ids),
|
||||
WorkflowRun.status == WorkflowExecutionStatus.RUNNING,
|
||||
)
|
||||
)
|
||||
)
|
||||
if has_running:
|
||||
return False
|
||||
for row in rows:
|
||||
row.rag_tenant_slot_released_at = released_at
|
||||
session.flush()
|
||||
return True
|
||||
|
||||
@staticmethod
|
||||
def _group_filters(identity: RagPipelineHandoffGroupIdentity):
|
||||
return (
|
||||
WorkflowRunHandoff.rag_source_batch_id == identity.source_batch_id,
|
||||
WorkflowRunHandoff.rag_tenant_id == identity.tenant_id,
|
||||
WorkflowRunHandoff.rag_queue_kind == identity.queue_kind,
|
||||
)
|
||||
|
||||
def _locked_group_rows(
|
||||
self, session: Session, identity: RagPipelineHandoffGroupIdentity
|
||||
) -> list[WorkflowRunHandoff]:
|
||||
rows = list(
|
||||
session.scalars(
|
||||
select(WorkflowRunHandoff)
|
||||
.where(*self._group_filters(identity))
|
||||
.order_by(WorkflowRunHandoff.workflow_run_id, WorkflowRunHandoff.generation)
|
||||
.with_for_update()
|
||||
)
|
||||
)
|
||||
self._validate_group_rows(rows)
|
||||
return rows
|
||||
|
||||
@staticmethod
|
||||
def _validate_group_rows(rows: Sequence[WorkflowRunHandoff]) -> None:
|
||||
isolation_values = {row.rag_tenant_isolated for row in rows}
|
||||
if None in isolation_values or len(isolation_values) > 1:
|
||||
raise RuntimeError("RAG handoff group has inconsistent tenant isolation metadata")
|
||||
dataset_values = {row.rag_dataset_id for row in rows}
|
||||
if None in dataset_values or len(dataset_values) > 1:
|
||||
raise RuntimeError("RAG handoff group has inconsistent dataset ownership metadata")
|
||||
|
||||
|
||||
__all__ = [
|
||||
"RagPipelineHandoffGroupRepository",
|
||||
"RagPipelineHandoffGroupSnapshot",
|
||||
"SQLAlchemyRagPipelineHandoffGroupRepository",
|
||||
]
|
||||
@@ -651,6 +651,9 @@ class DifyAPISQLAlchemyWorkflowRunRepository(APIWorkflowRunRepository):
|
||||
|
||||
trigger_logs_deleted = delete_trigger_logs(session, run_ids) if delete_trigger_logs else 0
|
||||
|
||||
# Keep handoff rows after run retention. They fence snapshot GC
|
||||
# until terminal compensation has completed, while the independent
|
||||
# snapshot-GC outbox remains the final durable object reference.
|
||||
runs_result = session.execute(delete(WorkflowRun).where(WorkflowRun.id.in_(run_ids)))
|
||||
runs_deleted = cast(CursorResult, runs_result).rowcount or 0
|
||||
|
||||
@@ -701,6 +704,8 @@ class DifyAPISQLAlchemyWorkflowRunRepository(APIWorkflowRunRepository):
|
||||
|
||||
trigger_logs_deleted = delete_trigger_logs(session, run_ids) if delete_trigger_logs else 0
|
||||
|
||||
# Keep handoff rows until durable terminal compensation and guarded
|
||||
# snapshot GC have both observed them.
|
||||
runs_result = session.execute(delete(WorkflowRun).where(WorkflowRun.id.in_(run_ids)))
|
||||
runs_deleted = cast(CursorResult, runs_result).rowcount or 0
|
||||
|
||||
@@ -755,6 +760,7 @@ class DifyAPISQLAlchemyWorkflowRunRepository(APIWorkflowRunRepository):
|
||||
run_finished_at=run.finished_at,
|
||||
run_exceptions_count=run.exceptions_count,
|
||||
trigger_metadata=trigger_metadata,
|
||||
run_handoff_duration=run.handoff_duration,
|
||||
)
|
||||
session.add(archive_log)
|
||||
return 1
|
||||
@@ -781,6 +787,7 @@ class DifyAPISQLAlchemyWorkflowRunRepository(APIWorkflowRunRepository):
|
||||
run_finished_at=run.finished_at,
|
||||
run_exceptions_count=run.exceptions_count,
|
||||
trigger_metadata=trigger_metadata,
|
||||
run_handoff_duration=run.handoff_duration,
|
||||
)
|
||||
for app_log in app_logs
|
||||
]
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,310 @@
|
||||
from collections.abc import Callable, Mapping, Sequence
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timedelta
|
||||
from enum import StrEnum
|
||||
from typing import Any, Protocol
|
||||
|
||||
from graphon.enums import WorkflowExecutionStatus
|
||||
from models.enums import CreatorUserRole
|
||||
from models.workflow_handoff import (
|
||||
RagPipelineHandoffGroupMetadata,
|
||||
WorkflowHandoffResumeRoute,
|
||||
WorkflowHandoffSnapshotGC,
|
||||
WorkflowRunHandoff,
|
||||
)
|
||||
|
||||
|
||||
class WorkflowHandoffPreparationCancelledError(RuntimeError):
|
||||
"""Raised when a durable Stop fences checkpoint preparation."""
|
||||
|
||||
|
||||
class WorkflowHandoffTerminalOwnershipError(RuntimeError):
|
||||
"""Raised when a runtime terminal write no longer owns the resumed run."""
|
||||
|
||||
|
||||
class WorkflowHandoffSnapshotDeleteOutcome(StrEnum):
|
||||
DELETED = "deleted"
|
||||
MISSING = "missing"
|
||||
ALREADY_DELETED = "already_deleted"
|
||||
BLOCKED = "blocked"
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class WorkflowHandoffTerminalEvent:
|
||||
handoff_id: str
|
||||
generation: int
|
||||
task_id: str
|
||||
resume_route: WorkflowHandoffResumeRoute
|
||||
workflow_run_id: str
|
||||
workflow_id: str
|
||||
status: WorkflowExecutionStatus
|
||||
outputs: Mapping[str, Any]
|
||||
error: str | None
|
||||
elapsed_time: float
|
||||
total_tokens: int
|
||||
total_steps: int
|
||||
created_at: datetime
|
||||
finished_at: datetime | None
|
||||
exceptions_count: int
|
||||
handoff_duration: float
|
||||
message_id: str | None = None
|
||||
message_answer: str | None = None
|
||||
message_metadata: Mapping[str, Any] | None = None
|
||||
message_files: Sequence[Mapping[str, Any]] = ()
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class WorkflowHandoffTerminalScope:
|
||||
"""Immutable ownership fence for terminalizing one resumed generation."""
|
||||
|
||||
workflow_run_id: str
|
||||
task_id: str
|
||||
tenant_id: str
|
||||
app_id: str
|
||||
workflow_id: str
|
||||
resume_route: WorkflowHandoffResumeRoute
|
||||
|
||||
|
||||
class WorkflowRunHandoffRepository(Protocol):
|
||||
"""Persistence boundary for durable workflow checkpoint handoffs."""
|
||||
|
||||
def create_preparing(
|
||||
self,
|
||||
*,
|
||||
workflow_run_id: str,
|
||||
task_id: str,
|
||||
snapshot_object_key: str,
|
||||
snapshot_schema_version: str,
|
||||
snapshot_checksum: str,
|
||||
snapshot_size_bytes: int,
|
||||
resume_route: WorkflowHandoffResumeRoute,
|
||||
source_worker_id: str,
|
||||
rag_group_metadata: RagPipelineHandoffGroupMetadata | None = None,
|
||||
) -> WorkflowRunHandoff:
|
||||
"""Commit an upload intent before object storage is touched."""
|
||||
...
|
||||
|
||||
def finish_preparing(
|
||||
self,
|
||||
*,
|
||||
handoff_id: str,
|
||||
generation: int,
|
||||
) -> WorkflowRunHandoff | None:
|
||||
"""Transition PREPARING to PREPARED unless a durable Stop won."""
|
||||
...
|
||||
|
||||
def create_prepared(
|
||||
self,
|
||||
*,
|
||||
workflow_run_id: str,
|
||||
task_id: str,
|
||||
snapshot_object_key: str,
|
||||
snapshot_schema_version: str,
|
||||
snapshot_checksum: str,
|
||||
snapshot_size_bytes: int,
|
||||
resume_route: WorkflowHandoffResumeRoute,
|
||||
source_worker_id: str,
|
||||
rag_group_metadata: RagPipelineHandoffGroupMetadata | None = None,
|
||||
) -> WorkflowRunHandoff:
|
||||
"""Create a checkpoint that remains undispatchable until explicitly activated."""
|
||||
...
|
||||
|
||||
def activate_latest_prepared_by_task_id(
|
||||
self,
|
||||
*,
|
||||
task_id: str,
|
||||
activated_at: datetime,
|
||||
) -> WorkflowRunHandoff | None:
|
||||
"""Atomically activate the latest PREPARED checkpoint for a task."""
|
||||
...
|
||||
|
||||
def get(self, handoff_id: str, generation: int | None = None) -> WorkflowRunHandoff | None: ...
|
||||
|
||||
def get_latest_by_run(self, workflow_run_id: str) -> WorkflowRunHandoff | None: ...
|
||||
|
||||
def list_due(
|
||||
self,
|
||||
*,
|
||||
now: datetime,
|
||||
redispatch_interval: timedelta,
|
||||
max_attempts: int,
|
||||
limit: int,
|
||||
) -> Sequence[WorkflowRunHandoff]:
|
||||
"""List durable outbox rows that should be dispatched or redispatched."""
|
||||
...
|
||||
|
||||
def mark_dispatched(self, *, handoff_id: str, generation: int, dispatched_at: datetime) -> bool: ...
|
||||
|
||||
def claim(
|
||||
self,
|
||||
*,
|
||||
handoff_id: str,
|
||||
generation: int,
|
||||
lease_owner: str,
|
||||
lease_duration: timedelta,
|
||||
max_attempts: int,
|
||||
now: datetime,
|
||||
) -> WorkflowRunHandoff | None:
|
||||
"""Atomically claim a due handoff or reclaim an expired lease."""
|
||||
...
|
||||
|
||||
def renew_lease(
|
||||
self,
|
||||
*,
|
||||
handoff_id: str,
|
||||
generation: int,
|
||||
lease_owner: str,
|
||||
lease_token: str,
|
||||
lease_duration: timedelta,
|
||||
now: datetime,
|
||||
) -> bool: ...
|
||||
|
||||
def record_failure(
|
||||
self,
|
||||
*,
|
||||
handoff_id: str,
|
||||
generation: int,
|
||||
lease_owner: str,
|
||||
lease_token: str,
|
||||
error: str,
|
||||
retry_at: datetime,
|
||||
max_attempts: int,
|
||||
now: datetime,
|
||||
) -> WorkflowRunHandoff | None:
|
||||
"""Release a claim for retry, or fail it when the retry budget is exhausted."""
|
||||
...
|
||||
|
||||
def mark_resumed(
|
||||
self,
|
||||
*,
|
||||
handoff_id: str,
|
||||
generation: int,
|
||||
lease_owner: str,
|
||||
lease_token: str,
|
||||
resumed_at: datetime,
|
||||
) -> bool: ...
|
||||
|
||||
def mark_failed(
|
||||
self,
|
||||
*,
|
||||
handoff_id: str,
|
||||
generation: int,
|
||||
error: str,
|
||||
failed_at: datetime,
|
||||
lease_owner: str | None = None,
|
||||
lease_token: str | None = None,
|
||||
) -> bool: ...
|
||||
|
||||
def request_cancel(
|
||||
self,
|
||||
*,
|
||||
workflow_run_id: str,
|
||||
requested_at: datetime,
|
||||
reason: str = "workflow run cancellation requested",
|
||||
) -> int:
|
||||
"""Cancel active durable handoffs and stop the parent run in the same transaction."""
|
||||
...
|
||||
|
||||
def request_cancel_by_task_id(
|
||||
self,
|
||||
*,
|
||||
task_id: str,
|
||||
requested_at: datetime,
|
||||
reason: str = "workflow task cancellation requested",
|
||||
scope_tenant_id: str | None = None,
|
||||
scope_app_id: str | None = None,
|
||||
scope_created_by_role: CreatorUserRole | None = None,
|
||||
scope_created_by: str | None = None,
|
||||
expires_at: datetime | None = None,
|
||||
) -> int:
|
||||
"""Record Stop and cancel scoped PREPARING/PREPARED/READY/CLAIMED rows.
|
||||
|
||||
Omitting owner scope is reserved for trusted internal call sites. A
|
||||
latest RESUMED row is left to the live graph's Abort command.
|
||||
"""
|
||||
...
|
||||
|
||||
def fail_exhausted(self, *, now: datetime, max_attempts: int, error: str) -> int:
|
||||
"""Fail exhausted READY rows and exhausted CLAIMED rows whose lease expired."""
|
||||
...
|
||||
|
||||
def fail_stale_prepared(self, *, now: datetime, stale_before: datetime, error: str, limit: int) -> int:
|
||||
"""Fail stale PREPARED rows and conditionally stop their still-running runs."""
|
||||
...
|
||||
|
||||
def fail_stale_ready(self, *, now: datetime, stale_before: datetime, error: str, limit: int) -> int:
|
||||
"""Fail latest READY rows that never reached a first resume claim."""
|
||||
...
|
||||
|
||||
def list_failed_pending_terminal_compensation(self, *, limit: int) -> Sequence[WorkflowRunHandoff]:
|
||||
"""List FAILED handoffs whose route-specific terminal state is not reconciled."""
|
||||
...
|
||||
|
||||
def compensate_failed_terminal(self, *, handoff_id: str, generation: int, compensated_at: datetime) -> bool:
|
||||
"""Idempotently reconcile WorkflowRun and route-specific terminal records."""
|
||||
...
|
||||
|
||||
def reconcile_resumed_terminal_failure(
|
||||
self,
|
||||
*,
|
||||
handoff_id: str,
|
||||
generation: int,
|
||||
scope: WorkflowHandoffTerminalScope,
|
||||
error: str,
|
||||
failed_at: datetime,
|
||||
message_answer_delta: str = "",
|
||||
message_answer_replacement: str | None = None,
|
||||
) -> WorkflowHandoffTerminalEvent | None:
|
||||
"""Atomically terminalize an owned, latest RESUMED generation.
|
||||
|
||||
A terminal event is returned only after the durable business records
|
||||
and terminal outbox marker have committed. ``None`` means that an
|
||||
already-published event or the durable PAUSED reconnect snapshot owns
|
||||
delivery. Ownership mismatches are raised instead of falling back to an
|
||||
unscoped write.
|
||||
"""
|
||||
...
|
||||
|
||||
def list_pending_terminal_events(self, *, limit: int) -> Sequence[WorkflowHandoffTerminalEvent]:
|
||||
"""Build durable terminal events that still need at-least-once publication."""
|
||||
...
|
||||
|
||||
def mark_terminal_event_published(self, *, handoff_id: str, generation: int, published_at: datetime) -> bool: ...
|
||||
|
||||
def record_terminal_processing_failure(self, *, handoff_id: str, generation: int, error: str) -> bool: ...
|
||||
|
||||
def list_snapshot_gc_candidates(self, *, now: datetime, limit: int) -> Sequence[WorkflowHandoffSnapshotGC]:
|
||||
"""List snapshot keys that may be checked for guarded deletion."""
|
||||
...
|
||||
|
||||
def delete_snapshot_if_unreferenced(
|
||||
self,
|
||||
*,
|
||||
snapshot_object_key: str,
|
||||
deleted_at: datetime,
|
||||
delete_object: Callable[[str], bool],
|
||||
) -> WorkflowHandoffSnapshotDeleteOutcome:
|
||||
"""Delete while serializing against new active references to the same key."""
|
||||
...
|
||||
|
||||
def record_snapshot_gc_failure(self, *, snapshot_object_key: str, error: str, retry_at: datetime) -> bool: ...
|
||||
|
||||
def cleanup_expired_cancellations(self, *, now: datetime, limit: int) -> int: ...
|
||||
|
||||
def cleanup_terminal_handoffs(self, *, terminal_before: datetime, limit: int) -> int:
|
||||
"""Delete bounded, fully reconciled terminal audit rows older than the cutoff."""
|
||||
...
|
||||
|
||||
def cleanup_completed_snapshot_gc(self, *, deleted_before: datetime, limit: int) -> int:
|
||||
"""Delete bounded, aged GC outbox rows after every handoff reference is gone."""
|
||||
...
|
||||
|
||||
|
||||
__all__ = [
|
||||
"WorkflowHandoffPreparationCancelledError",
|
||||
"WorkflowHandoffSnapshotDeleteOutcome",
|
||||
"WorkflowHandoffTerminalEvent",
|
||||
"WorkflowHandoffTerminalOwnershipError",
|
||||
"WorkflowHandoffTerminalScope",
|
||||
"WorkflowRunHandoffRepository",
|
||||
]
|
||||
@@ -18,7 +18,7 @@ from core.app.apps.message_based_app_generator import MessageBasedAppGenerator
|
||||
from core.app.apps.workflow.app_generator import WorkflowAppGenerator
|
||||
from core.app.entities.app_invoke_entities import InvokeFrom
|
||||
from core.app.features.rate_limiting import RateLimit
|
||||
from core.app.features.rate_limiting.rate_limit import rate_limit_context
|
||||
from core.app.features.rate_limiting.rate_limit import RateLimitGenerator, rate_limit_context
|
||||
from core.app.layers.pause_state_persist_layer import PauseStateLayerConfig
|
||||
from core.db import session_factory
|
||||
from enums.quota_type import QuotaType
|
||||
@@ -72,7 +72,8 @@ class AppGenerateService:
|
||||
|
||||
# Keep return type Callable[[], None] consistent while allowing an extra (no-op) call.
|
||||
def _on_subscribe_streams() -> None:
|
||||
_try_start()
|
||||
if not _try_start():
|
||||
raise RuntimeError("Failed to enqueue streaming workflow task")
|
||||
|
||||
return _on_subscribe_streams
|
||||
|
||||
@@ -84,6 +85,8 @@ class AppGenerateService:
|
||||
def _on_subscribe() -> None:
|
||||
if _try_start():
|
||||
timer.cancel()
|
||||
return
|
||||
raise RuntimeError("Failed to enqueue streaming workflow task")
|
||||
|
||||
return _on_subscribe
|
||||
|
||||
@@ -257,7 +260,7 @@ class AppGenerateService:
|
||||
|
||||
on_subscribe = cls._build_streaming_task_on_subscribe(on_subscribe)
|
||||
generator = AdvancedChatAppGenerator()
|
||||
return rate_limit.generate(
|
||||
stream_response = rate_limit.generate(
|
||||
generator.convert_to_event_stream(
|
||||
generator.retrieve_events(
|
||||
AppMode.ADVANCED_CHAT,
|
||||
@@ -267,6 +270,14 @@ class AppGenerateService:
|
||||
),
|
||||
request_id=request_id,
|
||||
)
|
||||
# The API owns this stable run identifier before the Celery
|
||||
# worker has persisted WorkflowRun. Expose it on the
|
||||
# streaming wrapper so the HTTP response can carry a
|
||||
# reconnect token even when the connection drops before
|
||||
# the first workflow_started event.
|
||||
if isinstance(stream_response, RateLimitGenerator):
|
||||
stream_response.workflow_run_id = payload.workflow_run_id
|
||||
return stream_response
|
||||
|
||||
# Blocking mode: run synchronously and return JSON instead of SSE
|
||||
# Keep behaviour consistent with WORKFLOW blocking branch.
|
||||
@@ -313,7 +324,7 @@ class AppGenerateService:
|
||||
workflow_based_app_execution_task.delay(payload_json)
|
||||
|
||||
on_subscribe = cls._build_streaming_task_on_subscribe(on_subscribe)
|
||||
return rate_limit.generate(
|
||||
stream_response = rate_limit.generate(
|
||||
WorkflowAppGenerator.convert_to_event_stream(
|
||||
MessageBasedAppGenerator.retrieve_events(
|
||||
AppMode.WORKFLOW,
|
||||
@@ -323,6 +334,9 @@ class AppGenerateService:
|
||||
),
|
||||
request_id,
|
||||
)
|
||||
if isinstance(stream_response, RateLimitGenerator):
|
||||
stream_response.workflow_run_id = payload.workflow_run_id
|
||||
return stream_response
|
||||
|
||||
pause_config = PauseStateLayerConfig(
|
||||
session_factory=session_factory.get_session_maker(),
|
||||
|
||||
@@ -9,7 +9,9 @@ from core.app.apps.base_app_queue_manager import AppQueueManager
|
||||
from core.app.entities.app_invoke_entities import InvokeFrom
|
||||
from extensions.ext_redis import redis_client
|
||||
from graphon.graph_engine.manager import GraphEngineManager
|
||||
from models.enums import CreatorUserRole
|
||||
from models.model import AppMode
|
||||
from services.workflow_handoff_cancellation_service import request_workflow_handoff_cancel_for_app
|
||||
|
||||
|
||||
class AppTaskService:
|
||||
@@ -21,26 +23,47 @@ class AppTaskService:
|
||||
invoke_from: InvokeFrom,
|
||||
user_id: str,
|
||||
app_mode: AppMode,
|
||||
*,
|
||||
tenant_id: str,
|
||||
app_id: str,
|
||||
created_by_role: CreatorUserRole | None = None,
|
||||
) -> None:
|
||||
"""Stop a running task.
|
||||
|
||||
This method handles stopping tasks using both mechanisms:
|
||||
This method handles stopping tasks using all applicable mechanisms:
|
||||
1. Legacy Redis flag mechanism (for backward compatibility)
|
||||
2. New GraphEngine command channel (for workflow-based apps)
|
||||
3. Durable handoff cancellation (when workflow handoff is enabled)
|
||||
|
||||
Args:
|
||||
task_id: The task ID to stop
|
||||
invoke_from: The source of the invoke (e.g., DEBUGGER, WEB_APP, SERVICE_API)
|
||||
user_id: The user ID requesting the stop
|
||||
app_mode: The application mode (CHAT, AGENT_CHAT, ADVANCED_CHAT, WORKFLOW, etc.)
|
||||
tenant_id: The owning tenant used to scope durable handoff cancellation
|
||||
app_id: The owning app used to scope durable handoff cancellation
|
||||
created_by_role: Optional creator-role override for entry points,
|
||||
such as OpenAPI, whose caller type is independent of invoke_from.
|
||||
|
||||
Returns:
|
||||
None
|
||||
"""
|
||||
# Legacy mechanism: Set stop flag in Redis
|
||||
AppQueueManager.set_stop_flag(task_id, invoke_from, user_id)
|
||||
live_task_owned_by_user = AppQueueManager.set_stop_flag(task_id, invoke_from, user_id)
|
||||
|
||||
# New mechanism: Send stop command via GraphEngine for workflow-based apps
|
||||
# This ensures proper workflow status recording in the persistence layer
|
||||
if app_mode in (AppMode.ADVANCED_CHAT, AppMode.WORKFLOW):
|
||||
GraphEngineManager(redis_client).send_stop_command(task_id)
|
||||
if app_mode in (AppMode.ADVANCED_CHAT, AppMode.WORKFLOW, AppMode.RAG_PIPELINE):
|
||||
cancelled_handoffs = request_workflow_handoff_cancel_for_app(
|
||||
task_id,
|
||||
tenant_id=tenant_id,
|
||||
app_id=app_id,
|
||||
created_by_role=created_by_role
|
||||
or (CreatorUserRole.ACCOUNT if invoke_from.runs_as_account() else CreatorUserRole.END_USER),
|
||||
created_by=user_id,
|
||||
)
|
||||
# The graph command channel is keyed only by caller-provided task
|
||||
# id. Send it only after either the legacy Redis owner record or
|
||||
# the durable creator-scoped handoff row proves ownership.
|
||||
if live_task_owned_by_user or cancelled_handoffs > 0:
|
||||
GraphEngineManager(redis_client).send_stop_command(task_id)
|
||||
|
||||
@@ -1,10 +1,12 @@
|
||||
from collections.abc import Mapping
|
||||
import uuid
|
||||
from collections.abc import Generator, Mapping
|
||||
from typing import Any
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from configs import dify_config
|
||||
from core.app.apps.pipeline.pipeline_generator import PipelineGenerator
|
||||
from core.app.apps.streaming_utils import WorkflowRunIdentifiedStream
|
||||
from core.app.entities.app_invoke_entities import InvokeFrom
|
||||
from models.dataset import Pipeline
|
||||
from models.enums import IndexingStatus
|
||||
@@ -44,7 +46,8 @@ class PipelineGenerateService:
|
||||
dataset_ref = DatasetRefService.create_dataset_ref(dataset)
|
||||
document_ref = DatasetRefService.create_document_ref_from_id(dataset_ref, original_document_id)
|
||||
cls.update_document_status(document_ref, session=session)
|
||||
return PipelineGenerator.convert_to_event_stream(
|
||||
workflow_run_id = str(uuid.uuid4()) if streaming and invoke_from == InvokeFrom.DEBUGGER else None
|
||||
converted = PipelineGenerator.convert_to_event_stream(
|
||||
PipelineGenerator().generate(
|
||||
session=session,
|
||||
pipeline=pipeline,
|
||||
@@ -55,8 +58,12 @@ class PipelineGenerateService:
|
||||
streaming=streaming,
|
||||
call_depth=0,
|
||||
workflow_thread_pool_id=None,
|
||||
workflow_run_id=workflow_run_id,
|
||||
),
|
||||
)
|
||||
if workflow_run_id is not None and isinstance(converted, Generator):
|
||||
return WorkflowRunIdentifiedStream(converted, workflow_run_id=workflow_run_id)
|
||||
return converted
|
||||
|
||||
except Exception:
|
||||
raise
|
||||
@@ -74,7 +81,8 @@ class PipelineGenerateService:
|
||||
cls, pipeline: Pipeline, user: Account, node_id: str, args: Any, session: Session, streaming: bool = True
|
||||
):
|
||||
workflow = cls._get_workflow(pipeline, InvokeFrom.DEBUGGER, session)
|
||||
return PipelineGenerator.convert_to_event_stream(
|
||||
workflow_run_id = str(uuid.uuid4())
|
||||
converted = PipelineGenerator.convert_to_event_stream(
|
||||
PipelineGenerator().single_iteration_generate(
|
||||
pipeline=pipeline,
|
||||
workflow=workflow,
|
||||
@@ -83,15 +91,20 @@ class PipelineGenerateService:
|
||||
args=args,
|
||||
streaming=streaming,
|
||||
session=session,
|
||||
workflow_run_id=workflow_run_id,
|
||||
)
|
||||
)
|
||||
if streaming and isinstance(converted, Generator):
|
||||
return WorkflowRunIdentifiedStream(converted, workflow_run_id=workflow_run_id)
|
||||
return converted
|
||||
|
||||
@classmethod
|
||||
def generate_single_loop(
|
||||
cls, pipeline: Pipeline, user: Account, node_id: str, args: Any, session: Session, streaming: bool = True
|
||||
):
|
||||
workflow = cls._get_workflow(pipeline, InvokeFrom.DEBUGGER, session)
|
||||
return PipelineGenerator.convert_to_event_stream(
|
||||
workflow_run_id = str(uuid.uuid4())
|
||||
converted = PipelineGenerator.convert_to_event_stream(
|
||||
PipelineGenerator().single_loop_generate(
|
||||
pipeline=pipeline,
|
||||
workflow=workflow,
|
||||
@@ -100,8 +113,12 @@ class PipelineGenerateService:
|
||||
args=args,
|
||||
streaming=streaming,
|
||||
session=session,
|
||||
workflow_run_id=workflow_run_id,
|
||||
)
|
||||
)
|
||||
if streaming and isinstance(converted, Generator):
|
||||
return WorkflowRunIdentifiedStream(converted, workflow_run_id=workflow_run_id)
|
||||
return converted
|
||||
|
||||
@classmethod
|
||||
def _get_workflow(cls, pipeline: Pipeline, invoke_from: InvokeFrom, session: Session) -> Workflow:
|
||||
|
||||
@@ -101,17 +101,24 @@ def _build_seeded_variable_pool(variables: Sequence[Variable]) -> VariablePool:
|
||||
|
||||
class RagPipelineService:
|
||||
_session: Session
|
||||
_session_maker: sessionmaker[Session]
|
||||
|
||||
def __init__(self, session: Session, session_maker: sessionmaker | None = None):
|
||||
def __init__(self, session: Session, session_maker: sessionmaker[Session] | None = None):
|
||||
"""Initialize RagPipelineService with repository dependencies."""
|
||||
self._session = session
|
||||
if session_maker is None:
|
||||
session_maker = sessionmaker(bind=db.engine, expire_on_commit=False)
|
||||
self._session_maker = session_maker
|
||||
self._node_execution_service_repo = DifyAPIRepositoryFactory.create_api_workflow_node_execution_repository(
|
||||
session_maker
|
||||
)
|
||||
self._workflow_run_repo = DifyAPIRepositoryFactory.create_api_workflow_run_repository(session_maker)
|
||||
|
||||
@property
|
||||
def session_maker(self) -> sessionmaker[Session]:
|
||||
"""Return the service-owned factory used by workflow event repositories."""
|
||||
return self._session_maker
|
||||
|
||||
@staticmethod
|
||||
def get_pipeline_by_id(pipeline_id: str, tenant_id: str, *, session: Session) -> Pipeline | None:
|
||||
return session.scalar(
|
||||
|
||||
@@ -0,0 +1,245 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
import logging
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime
|
||||
from enum import StrEnum
|
||||
from typing import Any, Protocol
|
||||
|
||||
from core.rag.pipeline.queue import TenantIsolatedTaskQueue
|
||||
from extensions.ext_redis import redis_client
|
||||
from models.workflow_handoff import RagPipelineHandoffGroupIdentity, RagPipelineQueueKind
|
||||
from repositories.rag_pipeline_handoff_group_repository import RagPipelineHandoffGroupRepository
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_RELEASE_LOCK_SECONDS = 60
|
||||
_RELEASE_MARKER_SECONDS = 7 * 24 * 60 * 60
|
||||
|
||||
|
||||
class _RedisLock(Protocol):
|
||||
def acquire(self, blocking: bool = True) -> bool: ...
|
||||
|
||||
def release(self) -> None: ...
|
||||
|
||||
|
||||
class _RedisClient(Protocol):
|
||||
def get(self, name: str | bytes) -> Any: ...
|
||||
|
||||
def set(
|
||||
self,
|
||||
name: str | bytes,
|
||||
value: Any,
|
||||
ex: int | None = None,
|
||||
px: int | None = None,
|
||||
nx: bool = False,
|
||||
xx: bool = False,
|
||||
keepttl: bool = False,
|
||||
get: bool = False,
|
||||
exat: int | None = None,
|
||||
pxat: int | None = None,
|
||||
) -> Any: ...
|
||||
|
||||
def lock(
|
||||
self,
|
||||
name: str,
|
||||
timeout: float | None = None,
|
||||
sleep: float = 0.1,
|
||||
blocking: bool = True,
|
||||
blocking_timeout: float | None = None,
|
||||
thread_local: bool = True,
|
||||
) -> _RedisLock: ...
|
||||
|
||||
|
||||
type RagPipelineBatchEnqueue = Callable[[str, str, str], None]
|
||||
|
||||
|
||||
class RagPipelineHandoffGroupOutcome(StrEnum):
|
||||
MISSING = "missing"
|
||||
NOT_READY = "not_ready"
|
||||
LOCK_BUSY = "lock_busy"
|
||||
RELEASED = "released"
|
||||
ALREADY_RELEASED = "already_released"
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class RagPipelineHandoffGroupScanResult:
|
||||
scanned: int
|
||||
released: int
|
||||
not_ready: int
|
||||
lock_busy: int
|
||||
errors: int
|
||||
|
||||
|
||||
class RagPipelineHandoffGroupService:
|
||||
"""Release one tenant slot after every handed-off run in a sealed batch ends."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
repository: RagPipelineHandoffGroupRepository,
|
||||
regular_enqueue: RagPipelineBatchEnqueue,
|
||||
priority_enqueue: RagPipelineBatchEnqueue,
|
||||
redis: _RedisClient = redis_client,
|
||||
) -> None:
|
||||
self._repository = repository
|
||||
self._regular_enqueue = regular_enqueue
|
||||
self._priority_enqueue = priority_enqueue
|
||||
self._redis = redis
|
||||
|
||||
def seal_group(self, *, identity: RagPipelineHandoffGroupIdentity, sealed_at: datetime) -> int:
|
||||
return self._repository.seal_group(identity=identity, sealed_at=sealed_at)
|
||||
|
||||
def reconcile_group(
|
||||
self, *, identity: RagPipelineHandoffGroupIdentity, now: datetime
|
||||
) -> RagPipelineHandoffGroupOutcome:
|
||||
# Document repair is database-only and remains retryable even when the
|
||||
# tenant-slot Redis lock is contended or unavailable.
|
||||
self._repository.mark_failed_documents(identity=identity, marked_at=now)
|
||||
snapshot = self._repository.get_group(identity)
|
||||
if snapshot is None:
|
||||
return RagPipelineHandoffGroupOutcome.MISSING
|
||||
if snapshot.released_at is not None:
|
||||
return RagPipelineHandoffGroupOutcome.ALREADY_RELEASED
|
||||
if snapshot.sealed_at is None:
|
||||
return RagPipelineHandoffGroupOutcome.NOT_READY
|
||||
if snapshot.has_running_workflow_runs:
|
||||
if snapshot.tenant_isolated:
|
||||
TenantIsolatedTaskQueue(identity.tenant_id, "pipeline").set_task_waiting_time(
|
||||
ttl=_RELEASE_MARKER_SECONDS
|
||||
)
|
||||
return RagPipelineHandoffGroupOutcome.NOT_READY
|
||||
|
||||
if not snapshot.tenant_isolated:
|
||||
marked = self._repository.mark_released_once(identity=identity, released_at=now)
|
||||
return (
|
||||
RagPipelineHandoffGroupOutcome.RELEASED if marked else RagPipelineHandoffGroupOutcome.ALREADY_RELEASED
|
||||
)
|
||||
|
||||
lock = self._redis.lock(
|
||||
self._release_lock_key(identity),
|
||||
timeout=_RELEASE_LOCK_SECONDS,
|
||||
blocking_timeout=0,
|
||||
)
|
||||
if not lock.acquire(blocking=False):
|
||||
return RagPipelineHandoffGroupOutcome.LOCK_BUSY
|
||||
try:
|
||||
snapshot = self._repository.get_group(identity)
|
||||
if snapshot is None:
|
||||
return RagPipelineHandoffGroupOutcome.MISSING
|
||||
if snapshot.released_at is not None:
|
||||
return RagPipelineHandoffGroupOutcome.ALREADY_RELEASED
|
||||
if snapshot.sealed_at is None or snapshot.has_running_workflow_runs:
|
||||
return RagPipelineHandoffGroupOutcome.NOT_READY
|
||||
|
||||
marker_key = self._release_marker_key(identity)
|
||||
if not self._redis.get(marker_key):
|
||||
self._release_tenant_slot(identity, claim_key=self._release_claim_key(identity))
|
||||
# If the database CAS transiently fails, the next scanner pass
|
||||
# observes this marker and skips the non-transactional queue
|
||||
# side effect before retrying the durable success marker.
|
||||
self._redis.set(marker_key, "1", ex=_RELEASE_MARKER_SECONDS)
|
||||
|
||||
marked = self._repository.mark_released_once(identity=identity, released_at=now)
|
||||
if marked:
|
||||
return RagPipelineHandoffGroupOutcome.RELEASED
|
||||
refreshed = self._repository.get_group(identity)
|
||||
if refreshed is not None and refreshed.released_at is not None:
|
||||
return RagPipelineHandoffGroupOutcome.ALREADY_RELEASED
|
||||
raise RuntimeError(f"Failed to persist RAG tenant-slot release marker: {identity}")
|
||||
finally:
|
||||
try:
|
||||
lock.release()
|
||||
except Exception:
|
||||
logger.warning("Failed to release RAG handoff group Redis lock", exc_info=True)
|
||||
|
||||
def scan(self, *, now: datetime, limit: int) -> RagPipelineHandoffGroupScanResult:
|
||||
identities = self._repository.list_reconcilable_groups(limit=limit)
|
||||
released = 0
|
||||
not_ready = 0
|
||||
lock_busy = 0
|
||||
errors = 0
|
||||
for identity in identities:
|
||||
try:
|
||||
if self._redis.get(self.batch_heartbeat_key(identity)):
|
||||
not_ready += 1
|
||||
continue
|
||||
# Unsealed rows are returned only after the uploaded source
|
||||
# batch owner heartbeat expires (or normal finalization clears
|
||||
# it). This repairs both worker loss and a transient seal loss.
|
||||
self._repository.seal_group(identity=identity, sealed_at=now)
|
||||
outcome = self.reconcile_group(identity=identity, now=now)
|
||||
except Exception:
|
||||
errors += 1
|
||||
logger.exception("Failed to reconcile RAG handoff group: %s", identity)
|
||||
continue
|
||||
if outcome == RagPipelineHandoffGroupOutcome.RELEASED:
|
||||
released += 1
|
||||
elif outcome == RagPipelineHandoffGroupOutcome.NOT_READY:
|
||||
not_ready += 1
|
||||
elif outcome == RagPipelineHandoffGroupOutcome.LOCK_BUSY:
|
||||
lock_busy += 1
|
||||
return RagPipelineHandoffGroupScanResult(
|
||||
scanned=len(identities),
|
||||
released=released,
|
||||
not_ready=not_ready,
|
||||
lock_busy=lock_busy,
|
||||
errors=errors,
|
||||
)
|
||||
|
||||
def _release_tenant_slot(self, identity: RagPipelineHandoffGroupIdentity, *, claim_key: str) -> None:
|
||||
queue = TenantIsolatedTaskQueue(identity.tenant_id, "pipeline")
|
||||
# The Redis claim is durable before Celery dispatch. A scanner retry
|
||||
# receives this same item instead of consuming an additional slot.
|
||||
has_next, raw_file_id = queue.claim_task_once(claim_key=claim_key, ttl=_RELEASE_MARKER_SECONDS)
|
||||
if not has_next:
|
||||
return
|
||||
|
||||
enqueue = (
|
||||
self._regular_enqueue if identity.queue_kind == RagPipelineQueueKind.REGULAR else self._priority_enqueue
|
||||
)
|
||||
if isinstance(raw_file_id, dict):
|
||||
file_id = raw_file_id.get("file_id")
|
||||
else:
|
||||
file_id = raw_file_id.decode("utf-8") if isinstance(raw_file_id, bytes) else raw_file_id
|
||||
if not isinstance(file_id, str) or not file_id:
|
||||
raise ValueError(f"Invalid queued RAG pipeline source batch: {raw_file_id!r}")
|
||||
enqueue(file_id, identity.tenant_id, self._dispatch_token(identity, file_id=file_id))
|
||||
|
||||
@staticmethod
|
||||
def _identity_digest(identity: RagPipelineHandoffGroupIdentity) -> str:
|
||||
value = f"{identity.source_batch_id}:{identity.tenant_id}:{identity.queue_kind.value}"
|
||||
return hashlib.sha256(value.encode()).hexdigest()
|
||||
|
||||
@classmethod
|
||||
def _release_lock_key(cls, identity: RagPipelineHandoffGroupIdentity) -> str:
|
||||
# Slot ownership is tenant-wide across regular and priority lanes.
|
||||
tenant_digest = hashlib.sha256(identity.tenant_id.encode()).hexdigest()
|
||||
return f"rag_pipeline_handoff_release_lock:{tenant_digest}"
|
||||
|
||||
@classmethod
|
||||
def _release_marker_key(cls, identity: RagPipelineHandoffGroupIdentity) -> str:
|
||||
return f"rag_pipeline_handoff_released:{cls._identity_digest(identity)}"
|
||||
|
||||
@classmethod
|
||||
def _release_claim_key(cls, identity: RagPipelineHandoffGroupIdentity) -> str:
|
||||
return f"rag_pipeline_handoff_release_claim:{cls._identity_digest(identity)}"
|
||||
|
||||
@classmethod
|
||||
def _dispatch_token(cls, identity: RagPipelineHandoffGroupIdentity, *, file_id: str) -> str:
|
||||
value = f"{cls._identity_digest(identity)}:{file_id}"
|
||||
return f"rag-pipeline-handoff:{hashlib.sha256(value.encode()).hexdigest()}"
|
||||
|
||||
@classmethod
|
||||
def batch_heartbeat_key(cls, identity: RagPipelineHandoffGroupIdentity) -> str:
|
||||
return f"rag_pipeline_handoff_batch_heartbeat:{cls._identity_digest(identity)}"
|
||||
|
||||
|
||||
__all__ = [
|
||||
"RagPipelineBatchEnqueue",
|
||||
"RagPipelineHandoffGroupOutcome",
|
||||
"RagPipelineHandoffGroupScanResult",
|
||||
"RagPipelineHandoffGroupService",
|
||||
]
|
||||
@@ -1,7 +1,8 @@
|
||||
import json
|
||||
import logging
|
||||
from collections.abc import Callable, Sequence
|
||||
from collections.abc import Sequence
|
||||
from functools import cached_property
|
||||
from typing import Any, Protocol
|
||||
|
||||
from core.app.entities.rag_pipeline_invoke_entities import RagPipelineInvokeEntity
|
||||
from core.rag.pipeline.queue import TenantIsolatedTaskQueue
|
||||
@@ -15,6 +16,10 @@ from tasks.rag_pipeline.rag_pipeline_run_task import rag_pipeline_run_task
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class _CeleryTask(Protocol):
|
||||
def delay(self, *args: Any, **kwargs: Any) -> Any: ...
|
||||
|
||||
|
||||
class RagPipelineTaskProxy:
|
||||
# Default uploaded file name for rag pipeline invoke entities
|
||||
_RAG_PIPELINE_INVOKE_ENTITIES_FILE_NAME = "rag_pipeline_invoke_entities.json"
|
||||
@@ -31,8 +36,13 @@ class RagPipelineTaskProxy:
|
||||
def features(self):
|
||||
return FeatureService.get_features(self._dataset_tenant_id, exclude_vector_space=True)
|
||||
|
||||
def _upload_invoke_entities(self) -> str:
|
||||
text = [item.model_dump() for item in self._rag_pipeline_invoke_entities]
|
||||
def _upload_invoke_entities(self, *, tenant_isolated: bool | None = None) -> str:
|
||||
text = []
|
||||
for item in self._rag_pipeline_invoke_entities:
|
||||
payload = item.model_dump()
|
||||
if tenant_isolated is not None:
|
||||
payload["tenant_isolated"] = tenant_isolated
|
||||
text.append(payload)
|
||||
# Convert list to proper JSON string
|
||||
json_text = json.dumps(text)
|
||||
upload_file = FileService(db.engine).upload_text(
|
||||
@@ -43,23 +53,19 @@ class RagPipelineTaskProxy:
|
||||
)
|
||||
return upload_file.id
|
||||
|
||||
def _send_to_direct_queue(self, upload_file_id: str, task_func: Callable[[str, str], None]):
|
||||
def _send_to_direct_queue(self, upload_file_id: str, task_func: _CeleryTask):
|
||||
logger.info("tenant %s send file %s to direct queue", self._dataset_tenant_id, upload_file_id)
|
||||
task_func.delay( # type: ignore
|
||||
task_func.delay(
|
||||
rag_pipeline_invoke_entities_file_id=upload_file_id,
|
||||
tenant_id=self._dataset_tenant_id,
|
||||
)
|
||||
|
||||
def _send_to_tenant_queue(self, upload_file_id: str, task_func: Callable[[str, str], None]):
|
||||
def _send_to_tenant_queue(self, upload_file_id: str, task_func: _CeleryTask):
|
||||
logger.info("tenant %s send file %s to tenant queue", self._dataset_tenant_id, upload_file_id)
|
||||
if self._tenant_isolated_task_queue.get_task_key():
|
||||
# Add to waiting queue using List operations (lpush)
|
||||
self._tenant_isolated_task_queue.push_tasks([upload_file_id])
|
||||
if not self._tenant_isolated_task_queue.enqueue_or_acquire(upload_file_id):
|
||||
logger.info("tenant %s push tasks: %s", self._dataset_tenant_id, upload_file_id)
|
||||
else:
|
||||
# Set flag and execute task
|
||||
self._tenant_isolated_task_queue.set_task_waiting_time()
|
||||
task_func.delay( # type: ignore
|
||||
task_func.delay(
|
||||
rag_pipeline_invoke_entities_file_id=upload_file_id,
|
||||
tenant_id=self._dataset_tenant_id,
|
||||
)
|
||||
@@ -75,7 +81,18 @@ class RagPipelineTaskProxy:
|
||||
self._send_to_direct_queue(upload_file_id, priority_rag_pipeline_run_task)
|
||||
|
||||
def _dispatch(self):
|
||||
upload_file_id = self._upload_invoke_entities()
|
||||
if self.features.billing.enabled:
|
||||
tenant_isolated = True
|
||||
send = (
|
||||
self._send_to_default_tenant_queue
|
||||
if self.features.billing.subscription.plan == CloudPlan.SANDBOX
|
||||
else self._send_to_priority_tenant_queue
|
||||
)
|
||||
else:
|
||||
tenant_isolated = False
|
||||
send = self._send_to_priority_direct_queue
|
||||
|
||||
upload_file_id = self._upload_invoke_entities(tenant_isolated=tenant_isolated)
|
||||
if not upload_file_id:
|
||||
raise ValueError("upload_file_id is empty")
|
||||
|
||||
@@ -86,17 +103,7 @@ class RagPipelineTaskProxy:
|
||||
self.features.billing.subscription.plan,
|
||||
)
|
||||
|
||||
# dispatch to different pipeline queue with tenant isolation when billing enabled
|
||||
if self.features.billing.enabled:
|
||||
if self.features.billing.subscription.plan == CloudPlan.SANDBOX:
|
||||
# dispatch to normal pipeline queue with tenant isolation for sandbox plan
|
||||
self._send_to_default_tenant_queue(upload_file_id)
|
||||
else:
|
||||
# dispatch to priority pipeline queue with tenant isolation for other plans
|
||||
self._send_to_priority_tenant_queue(upload_file_id)
|
||||
else:
|
||||
# dispatch to priority pipeline queue without tenant isolation for others, e.g.: self-hosted or enterprise
|
||||
self._send_to_priority_direct_queue(upload_file_id)
|
||||
send(upload_file_id)
|
||||
|
||||
def delay(self):
|
||||
if not self._rag_pipeline_invoke_entities:
|
||||
|
||||
@@ -20,12 +20,14 @@ Supported execution modes:
|
||||
|
||||
import json
|
||||
import logging
|
||||
from collections.abc import Generator, Mapping, Sequence
|
||||
from typing import Any, Union, cast
|
||||
import uuid
|
||||
from collections.abc import Generator, Iterator, Mapping, Sequence
|
||||
from typing import Any, Protocol, Union, cast, runtime_checkable
|
||||
|
||||
from sqlalchemy.orm import Session, make_transient, sessionmaker
|
||||
|
||||
from core.app.app_config.features.file_upload.manager import FileUploadConfigManager
|
||||
from core.app.apps.streaming_utils import StreamEventWithCursor, WorkflowRunIdentifiedStream
|
||||
from core.app.apps.workflow.app_generator import WorkflowAppGenerator
|
||||
from core.app.entities.app_invoke_entities import InvokeFrom
|
||||
from core.app.file_access import DatabaseFileAccessController
|
||||
@@ -36,13 +38,21 @@ from models import Account
|
||||
from models.model import App, AppMode, EndUser
|
||||
from models.snippet import CustomizedSnippet
|
||||
from models.workflow import Workflow, WorkflowNodeExecutionModel
|
||||
from models.workflow_handoff import WorkflowHandoffResumeRoute
|
||||
from services.snippet_service import SnippetService
|
||||
from services.workflow_service import WorkflowService
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
type SnippetGenerateResponse = Mapping[str, Any] | Iterator[str]
|
||||
_file_access_controller = DatabaseFileAccessController()
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class _Closable(Protocol):
|
||||
def close(self) -> object: ...
|
||||
|
||||
|
||||
class _SnippetAsApp:
|
||||
"""
|
||||
Minimal adapter that wraps a CustomizedSnippet to satisfy the App-like
|
||||
@@ -84,7 +94,7 @@ class SnippetGenerateService:
|
||||
_VIRTUAL_START_NODE_ID = SNIPPET_VIRTUAL_START_NODE_ID
|
||||
|
||||
@classmethod
|
||||
def _is_virtual_start_event(cls, message: Mapping[str, Any] | str) -> bool:
|
||||
def _is_virtual_start_event(cls, message: Mapping[str, Any] | StreamEventWithCursor | str) -> bool:
|
||||
"""
|
||||
Return True when *message* is a snippet-only virtual Start node event.
|
||||
|
||||
@@ -93,13 +103,14 @@ class SnippetGenerateService:
|
||||
out of the SSE stream so the frontend only receives nodes that exist on
|
||||
the canvas.
|
||||
"""
|
||||
if not isinstance(message, Mapping):
|
||||
payload = message.event if isinstance(message, StreamEventWithCursor) else message
|
||||
if not isinstance(payload, Mapping):
|
||||
return False
|
||||
|
||||
if message.get("event") not in {"node_started", "node_finished"}:
|
||||
if payload.get("event") not in {"node_started", "node_finished"}:
|
||||
return False
|
||||
|
||||
data = message.get("data")
|
||||
data = payload.get("data")
|
||||
if not isinstance(data, Mapping):
|
||||
return False
|
||||
|
||||
@@ -108,8 +119,8 @@ class SnippetGenerateService:
|
||||
@classmethod
|
||||
def _filter_virtual_start_events(
|
||||
cls,
|
||||
response: Mapping[str, Any] | Generator[Mapping[str, Any] | str, None, None],
|
||||
) -> Mapping[str, Any] | Generator[Mapping[str, Any] | str, None, None]:
|
||||
response: (Mapping[str, Any] | Generator[Mapping[str, Any] | StreamEventWithCursor | str, None, None]),
|
||||
) -> Mapping[str, Any] | Generator[Mapping[str, Any] | StreamEventWithCursor | str, None, None]:
|
||||
"""
|
||||
Drop snippet virtual Start node lifecycle events from stream responses.
|
||||
|
||||
@@ -119,14 +130,33 @@ class SnippetGenerateService:
|
||||
if isinstance(response, Mapping):
|
||||
return response
|
||||
|
||||
def _stream() -> Generator[Mapping[str, Any] | str, None, None]:
|
||||
for message in response:
|
||||
if cls._is_virtual_start_event(message):
|
||||
continue
|
||||
yield message
|
||||
def _stream() -> Generator[Mapping[str, Any] | StreamEventWithCursor | str, None, None]:
|
||||
try:
|
||||
for message in response:
|
||||
if cls._is_virtual_start_event(message):
|
||||
continue
|
||||
yield message
|
||||
finally:
|
||||
if isinstance(response, _Closable):
|
||||
response.close()
|
||||
|
||||
return _stream()
|
||||
|
||||
@classmethod
|
||||
def filter_virtual_start_events(
|
||||
cls,
|
||||
response: Generator[Mapping[str, Any] | StreamEventWithCursor | str, None, None],
|
||||
) -> Generator[Mapping[str, Any] | StreamEventWithCursor | str, None, None]:
|
||||
"""Apply the Snippet public-event filter to a worker or reconnect stream."""
|
||||
filtered = cls._filter_virtual_start_events(response)
|
||||
assert isinstance(filtered, Generator)
|
||||
return filtered
|
||||
|
||||
@staticmethod
|
||||
def build_app_model(snippet: CustomizedSnippet) -> App:
|
||||
"""Rebuild the App-shaped adapter used by both initial and resumed runs."""
|
||||
return cast(App, _SnippetAsApp(snippet))
|
||||
|
||||
@classmethod
|
||||
def generate(
|
||||
cls,
|
||||
@@ -136,7 +166,7 @@ class SnippetGenerateService:
|
||||
invoke_from: InvokeFrom,
|
||||
streaming: bool = True,
|
||||
session_maker: sessionmaker[Session] | None = None,
|
||||
) -> Mapping[str, Any] | Generator[str, None, None]:
|
||||
) -> SnippetGenerateResponse:
|
||||
"""
|
||||
Run a snippet's draft workflow.
|
||||
|
||||
@@ -165,8 +195,9 @@ class SnippetGenerateService:
|
||||
workflow = cls._ensure_start_node(workflow, snippet)
|
||||
|
||||
# Adapt snippet to App-like interface for WorkflowAppGenerator
|
||||
app_proxy = cast(App, _SnippetAsApp(snippet))
|
||||
app_proxy = cls.build_app_model(snippet)
|
||||
|
||||
workflow_run_id = str(uuid.uuid4())
|
||||
response = WorkflowAppGenerator().generate(
|
||||
app_model=app_proxy,
|
||||
workflow=workflow,
|
||||
@@ -175,9 +206,14 @@ class SnippetGenerateService:
|
||||
invoke_from=invoke_from,
|
||||
streaming=streaming,
|
||||
call_depth=0,
|
||||
workflow_run_id=workflow_run_id,
|
||||
handoff_resume_route=WorkflowHandoffResumeRoute.SNIPPET,
|
||||
)
|
||||
|
||||
return WorkflowAppGenerator.convert_to_event_stream(cls._filter_virtual_start_events(response))
|
||||
converted = WorkflowAppGenerator.convert_to_event_stream(cls._filter_virtual_start_events(response))
|
||||
if streaming and isinstance(converted, Generator):
|
||||
return WorkflowRunIdentifiedStream(converted, workflow_run_id=workflow_run_id)
|
||||
return converted
|
||||
|
||||
@classmethod
|
||||
def run_published(
|
||||
@@ -210,7 +246,7 @@ class SnippetGenerateService:
|
||||
# Inject a virtual Start node when the graph doesn't have one.
|
||||
workflow = cls._ensure_start_node(workflow, snippet)
|
||||
|
||||
app_proxy = cast(App, _SnippetAsApp(snippet))
|
||||
app_proxy = cls.build_app_model(snippet)
|
||||
|
||||
response = WorkflowAppGenerator().generate(
|
||||
app_model=app_proxy,
|
||||
@@ -220,6 +256,7 @@ class SnippetGenerateService:
|
||||
invoke_from=invoke_from,
|
||||
streaming=False,
|
||||
call_depth=0,
|
||||
handoff_resume_route=WorkflowHandoffResumeRoute.SNIPPET,
|
||||
)
|
||||
return response
|
||||
|
||||
@@ -355,7 +392,7 @@ class SnippetGenerateService:
|
||||
if not draft_workflow:
|
||||
raise ValueError("Workflow not initialized")
|
||||
|
||||
app_proxy = cast(App, _SnippetAsApp(snippet))
|
||||
app_proxy = cls.build_app_model(snippet)
|
||||
|
||||
workflow_service = WorkflowService()
|
||||
return workflow_service.run_draft_workflow_node(
|
||||
@@ -378,7 +415,7 @@ class SnippetGenerateService:
|
||||
streaming: bool = True,
|
||||
*,
|
||||
session_maker: sessionmaker[Session],
|
||||
) -> Mapping[str, Any] | Generator[str, None, None]:
|
||||
) -> SnippetGenerateResponse:
|
||||
"""
|
||||
Run a single iteration node in a snippet's draft workflow.
|
||||
|
||||
@@ -400,10 +437,11 @@ class SnippetGenerateService:
|
||||
if not workflow:
|
||||
raise ValueError("Workflow not initialized")
|
||||
|
||||
app_proxy = cast(App, _SnippetAsApp(snippet))
|
||||
app_proxy = cls.build_app_model(snippet)
|
||||
workflow_run_id = str(uuid.uuid4())
|
||||
|
||||
with session_maker() as session:
|
||||
return WorkflowAppGenerator.convert_to_event_stream(
|
||||
converted = WorkflowAppGenerator.convert_to_event_stream(
|
||||
WorkflowAppGenerator().single_iteration_generate(
|
||||
app_model=app_proxy,
|
||||
workflow=workflow,
|
||||
@@ -412,8 +450,13 @@ class SnippetGenerateService:
|
||||
args=args,
|
||||
streaming=streaming,
|
||||
session=session,
|
||||
workflow_run_id=workflow_run_id,
|
||||
handoff_resume_route=WorkflowHandoffResumeRoute.SNIPPET,
|
||||
)
|
||||
)
|
||||
if streaming and isinstance(converted, Generator):
|
||||
return WorkflowRunIdentifiedStream(converted, workflow_run_id=workflow_run_id)
|
||||
return converted
|
||||
|
||||
@classmethod
|
||||
def generate_single_loop(
|
||||
@@ -425,7 +468,7 @@ class SnippetGenerateService:
|
||||
streaming: bool = True,
|
||||
*,
|
||||
session_maker: sessionmaker[Session],
|
||||
) -> Mapping[str, Any] | Generator[str, None, None]:
|
||||
) -> SnippetGenerateResponse:
|
||||
"""
|
||||
Run a single loop node in a snippet's draft workflow.
|
||||
|
||||
@@ -447,10 +490,11 @@ class SnippetGenerateService:
|
||||
if not workflow:
|
||||
raise ValueError("Workflow not initialized")
|
||||
|
||||
app_proxy = cast(App, _SnippetAsApp(snippet))
|
||||
app_proxy = cls.build_app_model(snippet)
|
||||
workflow_run_id = str(uuid.uuid4())
|
||||
|
||||
with session_maker() as session:
|
||||
return WorkflowAppGenerator.convert_to_event_stream(
|
||||
converted = WorkflowAppGenerator.convert_to_event_stream(
|
||||
WorkflowAppGenerator().single_loop_generate(
|
||||
app_model=app_proxy,
|
||||
workflow=workflow,
|
||||
@@ -459,8 +503,13 @@ class SnippetGenerateService:
|
||||
args=args, # type: ignore[arg-type]
|
||||
streaming=streaming,
|
||||
session=session,
|
||||
workflow_run_id=workflow_run_id,
|
||||
handoff_resume_route=WorkflowHandoffResumeRoute.SNIPPET,
|
||||
)
|
||||
)
|
||||
if streaming and isinstance(converted, Generator):
|
||||
return WorkflowRunIdentifiedStream(converted, workflow_run_id=workflow_run_id)
|
||||
return converted
|
||||
|
||||
@staticmethod
|
||||
def parse_files(workflow: Workflow, files: list[dict] | None = None) -> Sequence[File]:
|
||||
|
||||
@@ -7,19 +7,23 @@ import threading
|
||||
import time
|
||||
from collections.abc import Generator, Mapping, Sequence
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
from typing import Any, Protocol, runtime_checkable
|
||||
|
||||
from sqlalchemy import desc, select
|
||||
from sqlalchemy.orm import Session, sessionmaker
|
||||
|
||||
from core.app.apps.common.workflow_response_converter import WorkflowResponseConverter
|
||||
from core.app.apps.message_generator import MessageGenerator
|
||||
from core.app.apps.streaming_utils import StreamEventWithCursor, stream_topic_events
|
||||
from core.app.entities.app_invoke_entities import AdvancedChatAppGenerateEntity
|
||||
from core.app.entities.task_entities import (
|
||||
HumanInputRequiredResponse,
|
||||
MessageEndStreamResponse,
|
||||
MessageReplaceStreamResponse,
|
||||
NodeFinishStreamResponse,
|
||||
NodeStartStreamResponse,
|
||||
StreamEvent,
|
||||
WorkflowFinishStreamResponse,
|
||||
WorkflowPauseStreamResponse,
|
||||
WorkflowStartStreamResponse,
|
||||
)
|
||||
@@ -38,20 +42,37 @@ from core.workflow.nodes.human_input.pause_reason import (
|
||||
DifyHITLEventType,
|
||||
HumanInputRequired,
|
||||
)
|
||||
from extensions.ext_storage import storage
|
||||
from graphon.entities import WorkflowStartReason
|
||||
from graphon.enums import WorkflowExecutionStatus, WorkflowNodeExecutionStatus
|
||||
from graphon.runtime import GraphRuntimeState
|
||||
from graphon.runtime.graph_runtime_state_protocol import ReadOnlyVariablePool
|
||||
from graphon.workflow_type_encoder import WorkflowRuntimeTypeConverter
|
||||
from libs.broadcast_channel.channel import CursorMessage, Topic
|
||||
from libs.broadcast_channel.cursor import normalize_stream_cursor
|
||||
from libs.broadcast_channel.exc import SubscriptionClosedError
|
||||
from models.enums import WorkflowRunTriggeredFrom
|
||||
from models.human_input import HumanInputForm
|
||||
from models.model import AppMode, Message
|
||||
from models.workflow import WorkflowNodeExecutionTriggeredFrom, WorkflowRun
|
||||
from models.workflow_handoff import WorkflowRunHandoff
|
||||
from repositories.api_workflow_node_execution_repository import WorkflowNodeExecutionSnapshot
|
||||
from repositories.entities.workflow_pause import WorkflowPauseEntity
|
||||
from repositories.factory import DifyAPIRepositoryFactory
|
||||
from repositories.sqlalchemy_workflow_handoff_repository import SQLAlchemyWorkflowRunHandoffRepository
|
||||
from services.workflow_handoff_service import WorkflowHandoffService
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_TERMINAL_WORKFLOW_STATUSES = frozenset(
|
||||
{
|
||||
WorkflowExecutionStatus.SUCCEEDED,
|
||||
WorkflowExecutionStatus.FAILED,
|
||||
WorkflowExecutionStatus.STOPPED,
|
||||
WorkflowExecutionStatus.PARTIAL_SUCCEEDED,
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class MessageContext:
|
||||
@@ -63,13 +84,25 @@ class MessageContext:
|
||||
|
||||
@dataclass
|
||||
class BufferState:
|
||||
queue: queue.Queue[Mapping[str, Any]]
|
||||
queue: queue.Queue[Mapping[str, Any] | StreamEventWithCursor]
|
||||
stop_event: threading.Event
|
||||
done_event: threading.Event
|
||||
task_id_ready: threading.Event
|
||||
task_id_hint: str | None = None
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class _RetainedCursorTopic(Protocol):
|
||||
def earliest_cursor(self) -> str | None: ...
|
||||
|
||||
def latest_cursor(self) -> str | None: ...
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class _CursorReceiver(Protocol):
|
||||
def receive_with_cursor(self, timeout: float | None = 0.1) -> CursorMessage | None: ...
|
||||
|
||||
|
||||
def build_workflow_event_stream(
|
||||
*,
|
||||
app_mode: AppMode,
|
||||
@@ -81,8 +114,68 @@ def build_workflow_event_stream(
|
||||
idle_timeout: float = 300,
|
||||
ping_interval: float = 10.0,
|
||||
close_on_pause: bool = True,
|
||||
) -> Generator[Mapping[str, Any] | str, None, None]:
|
||||
cursor: str | None = None,
|
||||
node_execution_triggered_from: WorkflowNodeExecutionTriggeredFrom | None = None,
|
||||
) -> Generator[Mapping[str, Any] | StreamEventWithCursor | str, None, None]:
|
||||
topic = MessageGenerator.get_response_topic(app_mode, workflow_run.id)
|
||||
|
||||
terminal_events = None if close_on_pause else [StreamEvent.WORKFLOW_FINISHED]
|
||||
has_retained_events = _topic_has_retained_events(topic)
|
||||
force_cursor_snapshot = False
|
||||
if cursor is not None:
|
||||
normalized_cursor = normalize_stream_cursor(cursor)
|
||||
retained_window = _topic_retained_cursor_window(topic)
|
||||
cursor_key = _stream_cursor_key(normalized_cursor)
|
||||
|
||||
# A terminal database row is the authoritative full-state replacement.
|
||||
# Replaying strictly after a tail cursor would otherwise wait for the
|
||||
# normal 300-second idle timeout even though no future event can arrive.
|
||||
if workflow_run.status in _TERMINAL_WORKFLOW_STATUSES:
|
||||
force_cursor_snapshot = True
|
||||
elif retained_window is not None:
|
||||
earliest_cursor, latest_cursor = retained_window
|
||||
earliest_key = _stream_cursor_key(earliest_cursor)
|
||||
latest_key = _stream_cursor_key(latest_cursor)
|
||||
cursor_is_replayable = normalized_cursor == "0-0" or earliest_key <= cursor_key <= latest_key
|
||||
|
||||
# A closed pause stream whose cursor already addresses the retained
|
||||
# tail also needs a persisted pause event rather than a long wait.
|
||||
cursor_is_closed_pause_tail = (
|
||||
close_on_pause and workflow_run.status == WorkflowExecutionStatus.PAUSED and cursor_key == latest_key
|
||||
)
|
||||
if cursor_is_replayable and not cursor_is_closed_pause_tail:
|
||||
return stream_topic_events(
|
||||
topic=topic,
|
||||
idle_timeout=idle_timeout,
|
||||
ping_interval=ping_interval,
|
||||
terminal_events=terminal_events,
|
||||
cursor=normalized_cursor,
|
||||
)
|
||||
force_cursor_snapshot = True
|
||||
else:
|
||||
# The key expired (or the transport cannot prove that this cursor
|
||||
# is still retained). Reconstruct current state from the database
|
||||
# and buffer live events so RUNNING runs do not lose the gap.
|
||||
force_cursor_snapshot = True
|
||||
|
||||
# A paused or terminal run can outlive its Redis Streams retention window.
|
||||
# In that case a Last-Event-ID no longer has a log to address; fall through
|
||||
# to the persisted snapshot so reconnect emits a full-state pause/terminal
|
||||
# event instead of a lone ping followed by EOF.
|
||||
|
||||
if has_retained_events and not force_cursor_snapshot:
|
||||
# The durable event log is the primary source of continuation truth.
|
||||
# Replaying it is both ordered and cursor-addressable; the DB snapshot
|
||||
# remains a compatibility fallback for runs whose event log predates
|
||||
# Streams or has expired.
|
||||
return stream_topic_events(
|
||||
topic=topic,
|
||||
idle_timeout=idle_timeout,
|
||||
ping_interval=ping_interval,
|
||||
terminal_events=terminal_events,
|
||||
cursor="0-0",
|
||||
)
|
||||
|
||||
workflow_run_repo = DifyAPIRepositoryFactory.create_api_workflow_run_repository(session_maker)
|
||||
node_execution_repo = DifyAPIRepositoryFactory.create_api_workflow_node_execution_repository(session_maker)
|
||||
|
||||
@@ -95,6 +188,18 @@ def build_workflow_event_stream(
|
||||
pause_entity = None
|
||||
|
||||
resumption_context = _load_resumption_context(pause_entity)
|
||||
latest_handoff = _get_latest_workflow_handoff(session_maker, workflow_run.id)
|
||||
handoff_resumption_context = None
|
||||
if resumption_context is None and latest_handoff is not None:
|
||||
handoff_resumption_context = _load_handoff_resumption_context(
|
||||
session_maker=session_maker,
|
||||
handoff=latest_handoff,
|
||||
)
|
||||
resolved_node_triggered_from = _resolve_node_execution_triggered_from(
|
||||
workflow_run=workflow_run,
|
||||
resumption_context=resumption_context or handoff_resumption_context,
|
||||
override=node_execution_triggered_from,
|
||||
)
|
||||
message_context: MessageContext | None = None
|
||||
if app_mode == AppMode.ADVANCED_CHAT:
|
||||
if workflow_run.status == WorkflowExecutionStatus.PAUSED:
|
||||
@@ -131,18 +236,28 @@ def build_workflow_event_stream(
|
||||
tenant_id=tenant_id,
|
||||
app_id=app_id,
|
||||
workflow_id=workflow_run.workflow_id,
|
||||
# NOTE(QuantumGhost): for events resumption, we only care about
|
||||
# the execution records from `WORKFLOW_RUN`.
|
||||
#
|
||||
# Ideally filtering with `workflow_run_id` is enough. However,
|
||||
# due to the index of `WorkflowNodeExecution` table, we have to
|
||||
# add a filter condition of `triggered_from` to
|
||||
# ensure that we can utilize the index.
|
||||
triggered_from=WorkflowNodeExecutionTriggeredFrom.WORKFLOW_RUN,
|
||||
# ``triggered_from`` is part of the node-execution lookup index. It must
|
||||
# match the repository used by the original or resumed execution: one-
|
||||
# step runs write SINGLE_STEP rows, full RAG runs write
|
||||
# RAG_PIPELINE_RUN rows, and ordinary runs write WORKFLOW_RUN rows.
|
||||
triggered_from=resolved_node_triggered_from,
|
||||
workflow_run_id=workflow_run.id,
|
||||
)
|
||||
|
||||
def _generate() -> Generator[Mapping[str, Any] | str, None, None]:
|
||||
def _generate() -> Generator[Mapping[str, Any] | StreamEventWithCursor | str, None, None]:
|
||||
# Close the small check/query race without combining a DB snapshot with
|
||||
# an already-retained event log. If events arrived while the fallback
|
||||
# snapshot was being read, replay the log instead.
|
||||
if not force_cursor_snapshot and _topic_has_retained_events(topic):
|
||||
yield from stream_topic_events(
|
||||
topic=topic,
|
||||
idle_timeout=idle_timeout,
|
||||
ping_interval=ping_interval,
|
||||
terminal_events=terminal_events,
|
||||
cursor="0-0",
|
||||
)
|
||||
return
|
||||
|
||||
# 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
|
||||
@@ -155,7 +270,12 @@ def build_workflow_event_stream(
|
||||
with topic.subscribe() as sub:
|
||||
buffer_state = _start_buffering(sub)
|
||||
try:
|
||||
task_id = _resolve_task_id(resumption_context, buffer_state, workflow_run.id)
|
||||
task_id = _resolve_task_id(
|
||||
resumption_context,
|
||||
buffer_state,
|
||||
workflow_run.id,
|
||||
latest_handoff_task_id=latest_handoff.task_id if latest_handoff is not None else None,
|
||||
)
|
||||
|
||||
snapshot_events = _build_snapshot_events(
|
||||
workflow_run=workflow_run,
|
||||
@@ -168,11 +288,11 @@ def build_workflow_event_stream(
|
||||
human_input_surface=human_input_surface,
|
||||
)
|
||||
|
||||
for event in snapshot_events:
|
||||
for snapshot_event in snapshot_events:
|
||||
last_msg_time = time.time()
|
||||
last_ping_time = last_msg_time
|
||||
yield event
|
||||
if _is_terminal_event(event, close_on_pause=close_on_pause):
|
||||
yield snapshot_event
|
||||
if _is_terminal_event(snapshot_event, close_on_pause=close_on_pause):
|
||||
return
|
||||
|
||||
while True:
|
||||
@@ -205,6 +325,35 @@ def build_workflow_event_stream(
|
||||
return _generate()
|
||||
|
||||
|
||||
def _topic_has_retained_events(topic: Topic) -> bool:
|
||||
if not isinstance(topic, _RetainedCursorTopic):
|
||||
return False
|
||||
try:
|
||||
return topic.latest_cursor() is not None
|
||||
except Exception:
|
||||
logger.exception("Failed to inspect retained workflow events")
|
||||
return False
|
||||
|
||||
|
||||
def _topic_retained_cursor_window(topic: Topic) -> tuple[str, str] | None:
|
||||
if not isinstance(topic, _RetainedCursorTopic):
|
||||
return None
|
||||
try:
|
||||
earliest = topic.earliest_cursor()
|
||||
latest = topic.latest_cursor()
|
||||
except Exception:
|
||||
logger.exception("Failed to inspect retained workflow event cursor window")
|
||||
return None
|
||||
if earliest is None or latest is None:
|
||||
return None
|
||||
return normalize_stream_cursor(earliest), normalize_stream_cursor(latest)
|
||||
|
||||
|
||||
def _stream_cursor_key(cursor: str) -> tuple[int, int]:
|
||||
milliseconds, sequence = normalize_stream_cursor(cursor).split("-", maxsplit=1)
|
||||
return int(milliseconds), int(sequence)
|
||||
|
||||
|
||||
def _get_message_context_by_conversation(
|
||||
session_maker: sessionmaker[Session],
|
||||
*,
|
||||
@@ -280,12 +429,115 @@ def _load_resumption_context(pause_entity: WorkflowPauseEntity | None) -> Workfl
|
||||
return None
|
||||
|
||||
|
||||
def _get_latest_workflow_handoff(
|
||||
session_maker: sessionmaker[Session],
|
||||
workflow_run_id: str,
|
||||
) -> WorkflowRunHandoff | None:
|
||||
"""Return the newest durable execution segment, if the run was handed off."""
|
||||
try:
|
||||
handoff = SQLAlchemyWorkflowRunHandoffRepository(session_maker).get_latest_by_run(workflow_run_id)
|
||||
except Exception:
|
||||
# Snapshot reconnect must remain compatible with runs created before
|
||||
# handoff support and with a rolling migration where this best-effort
|
||||
# lookup is temporarily unavailable.
|
||||
logger.warning(
|
||||
"Failed to load latest workflow handoff for event snapshot, workflow_run_id=%s",
|
||||
workflow_run_id,
|
||||
exc_info=True,
|
||||
)
|
||||
return None
|
||||
# Test doubles and compatibility session factories can return an untyped
|
||||
# sentinel. Do not let it leak into the public event contract.
|
||||
return handoff if isinstance(handoff, WorkflowRunHandoff) else None
|
||||
|
||||
|
||||
def _load_handoff_resumption_context(
|
||||
*,
|
||||
session_maker: sessionmaker[Session],
|
||||
handoff: WorkflowRunHandoff,
|
||||
) -> WorkflowResumptionContext | None:
|
||||
"""Best-effort load of the generate entity that produced a handoff segment."""
|
||||
try:
|
||||
handoff_service = WorkflowHandoffService(
|
||||
repository=SQLAlchemyWorkflowRunHandoffRepository(session_maker),
|
||||
storage=storage,
|
||||
)
|
||||
serialized_state = handoff_service.load_and_verify_state(handoff)
|
||||
return WorkflowResumptionContext.loads(serialized_state.decode())
|
||||
except Exception:
|
||||
# A terminal handoff snapshot may already have been garbage-collected.
|
||||
# The run-level trigger remains a safe fallback for full executions.
|
||||
logger.warning(
|
||||
"Failed to load workflow handoff context for event snapshot, "
|
||||
"workflow_run_id=%s, handoff_id=%s, generation=%s",
|
||||
handoff.workflow_run_id,
|
||||
handoff.id,
|
||||
handoff.generation,
|
||||
exc_info=True,
|
||||
)
|
||||
return None
|
||||
|
||||
|
||||
def _resolve_node_execution_triggered_from(
|
||||
*,
|
||||
workflow_run: WorkflowRun,
|
||||
resumption_context: WorkflowResumptionContext | None,
|
||||
override: WorkflowNodeExecutionTriggeredFrom | None = None,
|
||||
) -> WorkflowNodeExecutionTriggeredFrom:
|
||||
"""Resolve the indexed node-execution source used by this run segment."""
|
||||
if override is not None:
|
||||
return override
|
||||
|
||||
if resumption_context is not None:
|
||||
try:
|
||||
generate_entity = resumption_context.get_generate_entity()
|
||||
if generate_entity.single_iteration_run is not None or generate_entity.single_loop_run is not None:
|
||||
return WorkflowNodeExecutionTriggeredFrom.SINGLE_STEP
|
||||
except Exception:
|
||||
logger.warning(
|
||||
"Failed to inspect workflow resumption context for event snapshot, workflow_run_id=%s",
|
||||
workflow_run.id,
|
||||
exc_info=True,
|
||||
)
|
||||
|
||||
try:
|
||||
run_triggered_from = WorkflowRunTriggeredFrom(workflow_run.triggered_from)
|
||||
except ValueError:
|
||||
logger.warning(
|
||||
"Unknown workflow run trigger source for event snapshot, workflow_run_id=%s, triggered_from=%s",
|
||||
workflow_run.id,
|
||||
workflow_run.triggered_from,
|
||||
)
|
||||
return WorkflowNodeExecutionTriggeredFrom.WORKFLOW_RUN
|
||||
|
||||
if run_triggered_from in {
|
||||
WorkflowRunTriggeredFrom.RAG_PIPELINE_RUN,
|
||||
WorkflowRunTriggeredFrom.RAG_PIPELINE_DEBUGGING,
|
||||
}:
|
||||
return WorkflowNodeExecutionTriggeredFrom.RAG_PIPELINE_RUN
|
||||
return WorkflowNodeExecutionTriggeredFrom.WORKFLOW_RUN
|
||||
|
||||
|
||||
def resolve_workflow_event_task_id(
|
||||
*,
|
||||
workflow_run: WorkflowRun,
|
||||
session_maker: sessionmaker[Session],
|
||||
) -> str:
|
||||
"""Use the task identity of the newest execution segment when available."""
|
||||
latest_handoff = _get_latest_workflow_handoff(session_maker, workflow_run.id)
|
||||
return latest_handoff.task_id if latest_handoff is not None else workflow_run.id
|
||||
|
||||
|
||||
def _resolve_task_id(
|
||||
resumption_context: WorkflowResumptionContext | None,
|
||||
buffer_state: BufferState | None,
|
||||
workflow_run_id: str,
|
||||
wait_timeout: float = 0.2,
|
||||
*,
|
||||
latest_handoff_task_id: str | None = None,
|
||||
) -> str:
|
||||
if latest_handoff_task_id:
|
||||
return latest_handoff_task_id
|
||||
if resumption_context is not None:
|
||||
generate_entity = resumption_context.get_generate_entity()
|
||||
if generate_entity.task_id:
|
||||
@@ -368,6 +620,27 @@ def _build_snapshot_events(
|
||||
_apply_message_context(pause_event, message_context)
|
||||
events.append(pause_event)
|
||||
|
||||
if workflow_run.status in _TERMINAL_WORKFLOW_STATUSES:
|
||||
# Advanced Chat live streams always emit ``message_end`` before the
|
||||
# workflow terminal event. Preserve that contract when Redis history
|
||||
# has expired and the stream is reconstructed from the database. A
|
||||
# message context is only loaded for Advanced Chat, so workflow-only
|
||||
# and RAG snapshots remain unchanged.
|
||||
if message_context is not None:
|
||||
message_end = _build_message_end_event(
|
||||
task_id=task_id,
|
||||
message_id=message_context.message_id,
|
||||
)
|
||||
_apply_message_context(message_end, message_context)
|
||||
events.append(message_end)
|
||||
|
||||
workflow_finished = _build_workflow_finished_event(
|
||||
workflow_run=workflow_run,
|
||||
task_id=task_id,
|
||||
)
|
||||
_apply_message_context(workflow_finished, message_context)
|
||||
events.append(workflow_finished)
|
||||
|
||||
return events
|
||||
|
||||
|
||||
@@ -392,6 +665,38 @@ def _build_workflow_started_event(
|
||||
return payload
|
||||
|
||||
|
||||
def _build_workflow_finished_event(
|
||||
*,
|
||||
workflow_run: WorkflowRun,
|
||||
task_id: str,
|
||||
) -> dict[str, Any]:
|
||||
outputs = workflow_run.outputs_dict
|
||||
finished_at = workflow_run.finished_at
|
||||
response = WorkflowFinishStreamResponse(
|
||||
task_id=task_id,
|
||||
workflow_run_id=workflow_run.id,
|
||||
data=WorkflowFinishStreamResponse.Data(
|
||||
id=workflow_run.id,
|
||||
workflow_id=workflow_run.workflow_id,
|
||||
status=workflow_run.status,
|
||||
outputs=outputs,
|
||||
error=workflow_run.error,
|
||||
elapsed_time=float(workflow_run.elapsed_time or 0.0),
|
||||
total_tokens=int(workflow_run.total_tokens or 0),
|
||||
total_steps=int(workflow_run.total_steps or 0),
|
||||
created_by={},
|
||||
created_at=int(workflow_run.created_at.timestamp()),
|
||||
finished_at=int(finished_at.timestamp()) if finished_at is not None else None,
|
||||
files=WorkflowResponseConverter.fetch_files_from_node_outputs(outputs),
|
||||
exceptions_count=int(workflow_run.exceptions_count or 0),
|
||||
handoff_duration=float(workflow_run.handoff_duration or 0.0),
|
||||
),
|
||||
)
|
||||
payload = response.model_dump(mode="json")
|
||||
payload["event"] = response.event.value
|
||||
return payload
|
||||
|
||||
|
||||
def _build_message_replace_event(*, task_id: str, answer: str) -> dict[str, Any]:
|
||||
response = MessageReplaceStreamResponse(
|
||||
task_id=task_id,
|
||||
@@ -403,6 +708,16 @@ def _build_message_replace_event(*, task_id: str, answer: str) -> dict[str, Any]
|
||||
return payload
|
||||
|
||||
|
||||
def _build_message_end_event(*, task_id: str, message_id: str) -> dict[str, Any]:
|
||||
response = MessageEndStreamResponse(
|
||||
task_id=task_id,
|
||||
id=message_id,
|
||||
)
|
||||
payload = response.model_dump(mode="json")
|
||||
payload["event"] = response.event.value
|
||||
return payload
|
||||
|
||||
|
||||
def _build_node_started_event(
|
||||
*,
|
||||
workflow_run_id: str,
|
||||
@@ -616,6 +931,7 @@ def _build_pause_event(
|
||||
elapsed_time=float(workflow_run.elapsed_time or 0.0),
|
||||
total_tokens=int(workflow_run.total_tokens or 0),
|
||||
total_steps=int(workflow_run.total_steps or 0),
|
||||
handoff_duration=float(workflow_run.handoff_duration or 0.0),
|
||||
),
|
||||
)
|
||||
payload = response.model_dump(mode="json")
|
||||
@@ -640,10 +956,14 @@ def _start_buffering(subscription) -> BufferState:
|
||||
)
|
||||
|
||||
def _worker() -> None:
|
||||
dropped_count = 0
|
||||
try:
|
||||
while not buffer_state.stop_event.is_set():
|
||||
msg = subscription.receive(timeout=1)
|
||||
if isinstance(subscription, _CursorReceiver):
|
||||
cursor_message = subscription.receive_with_cursor(timeout=1)
|
||||
msg = None if cursor_message is None else cursor_message.payload
|
||||
else:
|
||||
cursor_message = None
|
||||
msg = subscription.receive(timeout=1)
|
||||
if msg is None:
|
||||
continue
|
||||
event = _parse_event_message(msg)
|
||||
@@ -653,19 +973,22 @@ def _start_buffering(subscription) -> BufferState:
|
||||
if task_id and buffer_state.task_id_hint is None:
|
||||
buffer_state.task_id_hint = str(task_id)
|
||||
buffer_state.task_id_ready.set()
|
||||
try:
|
||||
buffer_state.queue.put_nowait(event)
|
||||
except queue.Full:
|
||||
dropped_count += 1
|
||||
buffered_event: Mapping[str, Any] | StreamEventWithCursor = event
|
||||
if cursor_message is not None:
|
||||
buffered_event = StreamEventWithCursor(event=event, cursor=cursor_message.cursor)
|
||||
# Apply lossless backpressure while the database snapshot is
|
||||
# being built. Advancing past a dropped Redis cursor would let
|
||||
# the client acknowledge a gap that Last-Event-ID can never
|
||||
# repair. The Streams subscription itself remains replayable
|
||||
# while this bounded queue waits for the HTTP consumer.
|
||||
while not buffer_state.stop_event.is_set():
|
||||
try:
|
||||
buffer_state.queue.get_nowait()
|
||||
except queue.Empty:
|
||||
pass
|
||||
try:
|
||||
buffer_state.queue.put_nowait(event)
|
||||
buffer_state.queue.put(buffered_event, timeout=1)
|
||||
break
|
||||
except queue.Full:
|
||||
continue
|
||||
logger.warning("Dropped buffered workflow event, total_dropped=%s", dropped_count)
|
||||
except SubscriptionClosedError:
|
||||
pass
|
||||
except Exception:
|
||||
logger.exception("Failed while buffering workflow events")
|
||||
finally:
|
||||
@@ -688,13 +1011,15 @@ def _parse_event_message(message: bytes) -> Mapping[str, Any] | None:
|
||||
|
||||
|
||||
def _is_terminal_event(
|
||||
event: Mapping[str, Any] | str,
|
||||
event: Mapping[str, Any] | StreamEventWithCursor | str,
|
||||
close_on_pause: bool = True,
|
||||
*,
|
||||
include_paused: bool | None = None,
|
||||
) -> bool:
|
||||
if include_paused is not None:
|
||||
close_on_pause = include_paused
|
||||
if isinstance(event, StreamEventWithCursor):
|
||||
event = event.event
|
||||
if not isinstance(event, Mapping):
|
||||
return False
|
||||
event_type = event.get("event")
|
||||
|
||||
@@ -0,0 +1,115 @@
|
||||
import logging
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime
|
||||
|
||||
from celery import current_app as current_celery_app
|
||||
from sqlalchemy.orm import sessionmaker
|
||||
|
||||
from configs import dify_config
|
||||
from extensions.ext_database import db
|
||||
from libs.datetime_utils import naive_utc_now
|
||||
from models.workflow_handoff import WorkflowRunHandoff
|
||||
from repositories.sqlalchemy_workflow_handoff_repository import SQLAlchemyWorkflowRunHandoffRepository
|
||||
from repositories.workflow_handoff_repository import WorkflowRunHandoffRepository
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
WORKFLOW_HANDOFF_RESUME_TASK_NAME = "workflow_handoff.resume"
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class WorkflowHandoffActivationResult:
|
||||
handoff: WorkflowRunHandoff | None
|
||||
enqueued: bool = False
|
||||
dispatch_marked: bool = False
|
||||
errors: int = 0
|
||||
|
||||
@property
|
||||
def activated(self) -> bool:
|
||||
return self.handoff is not None
|
||||
|
||||
|
||||
class WorkflowHandoffActivationService:
|
||||
"""Cross the drain barrier, then dispatch through the durable outbox.
|
||||
|
||||
The repository commits PREPARED -> READY before this class touches the
|
||||
broker. If broker publication or ``mark_dispatched`` fails, the READY row is
|
||||
deliberately left for the periodic scanner to repair.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
repository: WorkflowRunHandoffRepository,
|
||||
enqueue: Callable[[str, int], None],
|
||||
) -> None:
|
||||
self._repository = repository
|
||||
self._enqueue = enqueue
|
||||
|
||||
def activate(self, *, task_id: str, now: datetime) -> WorkflowHandoffActivationResult:
|
||||
handoff = self._repository.activate_latest_prepared_by_task_id(
|
||||
task_id=task_id,
|
||||
activated_at=now,
|
||||
)
|
||||
if handoff is None:
|
||||
return WorkflowHandoffActivationResult(handoff=None)
|
||||
|
||||
try:
|
||||
self._enqueue(handoff.id, handoff.generation)
|
||||
except Exception:
|
||||
logger.exception(
|
||||
"Failed to enqueue activated workflow handoff; scanner will retry: handoff_id=%s, generation=%s",
|
||||
handoff.id,
|
||||
handoff.generation,
|
||||
)
|
||||
return WorkflowHandoffActivationResult(handoff=handoff, errors=1)
|
||||
|
||||
try:
|
||||
dispatch_marked = self._repository.mark_dispatched(
|
||||
handoff_id=handoff.id,
|
||||
generation=handoff.generation,
|
||||
dispatched_at=now,
|
||||
)
|
||||
except Exception:
|
||||
logger.exception(
|
||||
"Failed to mark activated workflow handoff dispatched; scanner will retry: "
|
||||
"handoff_id=%s, generation=%s",
|
||||
handoff.id,
|
||||
handoff.generation,
|
||||
)
|
||||
return WorkflowHandoffActivationResult(handoff=handoff, enqueued=True, errors=1)
|
||||
|
||||
return WorkflowHandoffActivationResult(
|
||||
handoff=handoff,
|
||||
enqueued=True,
|
||||
dispatch_marked=dispatch_marked,
|
||||
)
|
||||
|
||||
|
||||
def activate_workflow_handoff_by_task_id(task_id: str) -> WorkflowHandoffActivationResult:
|
||||
"""Runtime helper for response pipelines after the old segment fully drains."""
|
||||
|
||||
repository = SQLAlchemyWorkflowRunHandoffRepository(
|
||||
sessionmaker(bind=db.engine, expire_on_commit=False),
|
||||
)
|
||||
|
||||
def _enqueue(handoff_id: str, generation: int) -> None:
|
||||
current_celery_app.send_task(
|
||||
WORKFLOW_HANDOFF_RESUME_TASK_NAME,
|
||||
kwargs={"handoff_id": handoff_id, "generation": generation},
|
||||
queue=dify_config.WORKFLOW_HANDOFF_QUEUE,
|
||||
)
|
||||
|
||||
return WorkflowHandoffActivationService(
|
||||
repository=repository,
|
||||
enqueue=_enqueue,
|
||||
).activate(task_id=task_id, now=naive_utc_now())
|
||||
|
||||
|
||||
__all__ = [
|
||||
"WORKFLOW_HANDOFF_RESUME_TASK_NAME",
|
||||
"WorkflowHandoffActivationResult",
|
||||
"WorkflowHandoffActivationService",
|
||||
"activate_workflow_handoff_by_task_id",
|
||||
]
|
||||
@@ -0,0 +1,118 @@
|
||||
from datetime import datetime, timedelta
|
||||
|
||||
from sqlalchemy.orm import sessionmaker
|
||||
|
||||
from extensions.ext_database import db
|
||||
from libs.datetime_utils import naive_utc_now
|
||||
from models.enums import CreatorUserRole
|
||||
from repositories.sqlalchemy_workflow_handoff_repository import SQLAlchemyWorkflowRunHandoffRepository
|
||||
from repositories.workflow_handoff_repository import WorkflowRunHandoffRepository
|
||||
|
||||
WORKFLOW_HANDOFF_CANCELLATION_RETENTION = timedelta(days=7)
|
||||
|
||||
|
||||
class WorkflowHandoffCancellationService:
|
||||
"""Persist a Stop request so a queued handoff cannot start later."""
|
||||
|
||||
def __init__(self, repository: WorkflowRunHandoffRepository) -> None:
|
||||
self._repository = repository
|
||||
|
||||
def request_by_task_id(
|
||||
self,
|
||||
*,
|
||||
task_id: str,
|
||||
requested_at: datetime,
|
||||
reason: str = "workflow task cancellation requested",
|
||||
scope_tenant_id: str | None = None,
|
||||
scope_app_id: str | None = None,
|
||||
scope_created_by_role: CreatorUserRole | None = None,
|
||||
scope_created_by: str | None = None,
|
||||
) -> int:
|
||||
if not task_id:
|
||||
raise ValueError("task_id must not be empty")
|
||||
if not reason:
|
||||
raise ValueError("reason must not be empty")
|
||||
if (scope_tenant_id is None) != (scope_app_id is None):
|
||||
raise ValueError("scope_tenant_id and scope_app_id must be provided together")
|
||||
if (scope_created_by_role is None) != (scope_created_by is None):
|
||||
raise ValueError("scope_created_by_role and scope_created_by must be provided together")
|
||||
if scope_created_by is not None and scope_app_id is None:
|
||||
raise ValueError("creator scope requires tenant and app scope")
|
||||
return self._repository.request_cancel_by_task_id(
|
||||
task_id=task_id,
|
||||
requested_at=requested_at,
|
||||
reason=reason,
|
||||
scope_tenant_id=scope_tenant_id,
|
||||
scope_app_id=scope_app_id,
|
||||
scope_created_by_role=scope_created_by_role,
|
||||
scope_created_by=scope_created_by,
|
||||
expires_at=requested_at + WORKFLOW_HANDOFF_CANCELLATION_RETENTION,
|
||||
)
|
||||
|
||||
|
||||
def request_workflow_handoff_cancel_by_task_id(
|
||||
task_id: str,
|
||||
*,
|
||||
reason: str = "workflow task cancellation requested",
|
||||
requested_at: datetime | None = None,
|
||||
scope_tenant_id: str | None = None,
|
||||
scope_app_id: str | None = None,
|
||||
scope_created_by_role: CreatorUserRole | None = None,
|
||||
scope_created_by: str | None = None,
|
||||
) -> int:
|
||||
"""Controller-facing helper for the durable half of user Stop.
|
||||
|
||||
Callers should still publish the existing Redis stop flag and GraphEngine
|
||||
Abort command for a live segment. This helper closes the
|
||||
PREPARED/READY/CLAIMED gap;
|
||||
it intentionally lets database errors propagate so an API must not report a
|
||||
durable Stop that was never recorded. Public/authenticated callers should
|
||||
pass both owner IDs; omitting them is reserved for trusted in-process aborts
|
||||
that already own the task stream.
|
||||
"""
|
||||
repository = SQLAlchemyWorkflowRunHandoffRepository(
|
||||
sessionmaker(bind=db.engine, expire_on_commit=False),
|
||||
)
|
||||
return WorkflowHandoffCancellationService(repository).request_by_task_id(
|
||||
task_id=task_id,
|
||||
requested_at=requested_at or naive_utc_now(),
|
||||
reason=reason,
|
||||
scope_tenant_id=scope_tenant_id,
|
||||
scope_app_id=scope_app_id,
|
||||
scope_created_by_role=scope_created_by_role,
|
||||
scope_created_by=scope_created_by,
|
||||
)
|
||||
|
||||
|
||||
def request_workflow_handoff_cancel_for_app(
|
||||
task_id: str,
|
||||
*,
|
||||
tenant_id: str,
|
||||
app_id: str,
|
||||
created_by_role: CreatorUserRole,
|
||||
created_by: str,
|
||||
reason: str = "workflow task cancellation requested",
|
||||
requested_at: datetime | None = None,
|
||||
) -> int:
|
||||
"""Owner-scoped helper for authenticated/public Stop endpoints."""
|
||||
if not tenant_id or not app_id:
|
||||
raise ValueError("tenant_id and app_id must not be empty")
|
||||
if not created_by:
|
||||
raise ValueError("created_by must not be empty")
|
||||
return request_workflow_handoff_cancel_by_task_id(
|
||||
task_id,
|
||||
reason=reason,
|
||||
requested_at=requested_at,
|
||||
scope_tenant_id=tenant_id,
|
||||
scope_app_id=app_id,
|
||||
scope_created_by_role=created_by_role,
|
||||
scope_created_by=created_by,
|
||||
)
|
||||
|
||||
|
||||
__all__ = [
|
||||
"WORKFLOW_HANDOFF_CANCELLATION_RETENTION",
|
||||
"WorkflowHandoffCancellationService",
|
||||
"request_workflow_handoff_cancel_by_task_id",
|
||||
"request_workflow_handoff_cancel_for_app",
|
||||
]
|
||||
@@ -0,0 +1,120 @@
|
||||
import logging
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timedelta
|
||||
|
||||
from repositories.workflow_handoff_repository import WorkflowRunHandoffRepository
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class WorkflowHandoffDispatchResult:
|
||||
exhausted_failed: int
|
||||
due: int
|
||||
enqueued: int
|
||||
dispatch_marked: int
|
||||
errors: int
|
||||
stale_prepared_failed: int = 0
|
||||
stale_ready_failed: int = 0
|
||||
|
||||
|
||||
class WorkflowHandoffDispatcher:
|
||||
"""Scan the durable handoff outbox and enqueue fenced resume attempts.
|
||||
|
||||
Enqueueing intentionally happens before ``mark_dispatched``. A process crash
|
||||
can therefore create a duplicate message, but the lease-token claim makes it
|
||||
harmless. Reversing the order could lose a handoff until the redispatch
|
||||
deadline when the broker write fails.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
repository: WorkflowRunHandoffRepository,
|
||||
enqueue: Callable[[str, int], None],
|
||||
) -> None:
|
||||
self._repository = repository
|
||||
self._enqueue = enqueue
|
||||
|
||||
def scan(
|
||||
self,
|
||||
*,
|
||||
now: datetime,
|
||||
redispatch_interval: timedelta,
|
||||
prepared_timeout: timedelta,
|
||||
max_attempts: int,
|
||||
limit: int,
|
||||
) -> WorkflowHandoffDispatchResult:
|
||||
if prepared_timeout.total_seconds() <= 0:
|
||||
raise ValueError("prepared_timeout must be positive")
|
||||
stale_prepared_failed = self._repository.fail_stale_prepared(
|
||||
now=now,
|
||||
stale_before=now - prepared_timeout,
|
||||
error="workflow handoff drain barrier timed out before activation",
|
||||
limit=limit,
|
||||
)
|
||||
stale_ready_failed = self._repository.fail_stale_ready(
|
||||
now=now,
|
||||
stale_before=now - prepared_timeout,
|
||||
error="workflow handoff timed out before the first resume attempt",
|
||||
limit=limit,
|
||||
)
|
||||
exhausted_failed = self._repository.fail_exhausted(
|
||||
now=now,
|
||||
max_attempts=max_attempts,
|
||||
error=f"workflow handoff exhausted {max_attempts} resume attempts",
|
||||
)
|
||||
due_handoffs = self._repository.list_due(
|
||||
now=now,
|
||||
redispatch_interval=redispatch_interval,
|
||||
max_attempts=max_attempts,
|
||||
limit=limit,
|
||||
)
|
||||
|
||||
enqueued = 0
|
||||
dispatch_marked = 0
|
||||
errors = 0
|
||||
for handoff in due_handoffs:
|
||||
try:
|
||||
self._enqueue(handoff.id, handoff.generation)
|
||||
enqueued += 1
|
||||
except Exception:
|
||||
errors += 1
|
||||
logger.exception(
|
||||
"Failed to enqueue workflow handoff resume: handoff_id=%s, generation=%s",
|
||||
handoff.id,
|
||||
handoff.generation,
|
||||
)
|
||||
continue
|
||||
|
||||
try:
|
||||
if self._repository.mark_dispatched(
|
||||
handoff_id=handoff.id,
|
||||
generation=handoff.generation,
|
||||
dispatched_at=now,
|
||||
):
|
||||
dispatch_marked += 1
|
||||
except Exception:
|
||||
# The broker already owns a message. Leaving the row unmarked is
|
||||
# safe and makes the next scan redispatch it; claim fencing keeps
|
||||
# that duplicate from starting another graph.
|
||||
errors += 1
|
||||
logger.exception(
|
||||
"Failed to mark workflow handoff dispatched: handoff_id=%s, generation=%s",
|
||||
handoff.id,
|
||||
handoff.generation,
|
||||
)
|
||||
|
||||
return WorkflowHandoffDispatchResult(
|
||||
exhausted_failed=exhausted_failed,
|
||||
due=len(due_handoffs),
|
||||
enqueued=enqueued,
|
||||
dispatch_marked=dispatch_marked,
|
||||
errors=errors,
|
||||
stale_prepared_failed=stale_prepared_failed,
|
||||
stale_ready_failed=stale_ready_failed,
|
||||
)
|
||||
|
||||
|
||||
__all__ = ["WorkflowHandoffDispatchResult", "WorkflowHandoffDispatcher"]
|
||||
@@ -0,0 +1,408 @@
|
||||
import logging
|
||||
from collections.abc import Callable, Mapping
|
||||
from contextlib import AbstractContextManager
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timedelta
|
||||
from enum import StrEnum
|
||||
from threading import Event, Thread
|
||||
from types import TracebackType
|
||||
from typing import Protocol, override
|
||||
|
||||
from libs.datetime_utils import naive_utc_now
|
||||
from models.workflow_handoff import (
|
||||
WorkflowHandoffResumeRoute,
|
||||
WorkflowHandoffState,
|
||||
WorkflowRunHandoff,
|
||||
)
|
||||
from repositories.workflow_handoff_repository import WorkflowRunHandoffRepository
|
||||
from services.workflow_handoff_service import (
|
||||
UnsupportedWorkflowHandoffSnapshotVersionError,
|
||||
WorkflowHandoffService,
|
||||
WorkflowHandoffSnapshotIntegrityError,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class PermanentWorkflowHandoffResumeError(RuntimeError):
|
||||
"""A checkpoint cannot become resumable without operator intervention."""
|
||||
|
||||
|
||||
class UnsupportedWorkflowHandoffResumeRouteError(PermanentWorkflowHandoffResumeError):
|
||||
pass
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class WorkflowHandoffLease:
|
||||
"""Fenced lease handle available to route-specific resume setup."""
|
||||
|
||||
repository: WorkflowRunHandoffRepository
|
||||
handoff_id: str
|
||||
generation: int
|
||||
lease_owner: str
|
||||
lease_token: str
|
||||
lease_duration: timedelta
|
||||
|
||||
def renew(self, *, now: datetime) -> bool:
|
||||
return self.repository.renew_lease(
|
||||
handoff_id=self.handoff_id,
|
||||
generation=self.generation,
|
||||
lease_owner=self.lease_owner,
|
||||
lease_token=self.lease_token,
|
||||
lease_duration=self.lease_duration,
|
||||
now=now,
|
||||
)
|
||||
|
||||
|
||||
class _WorkflowHandoffLeaseHeartbeat(AbstractContextManager[None]):
|
||||
"""Keep a claimed checkpoint fenced while route setup reaches graph ACK.
|
||||
|
||||
A resume handler may spend longer than one lease loading plugins or building
|
||||
a large graph. The acknowledgement layer stops accepting renewals once the
|
||||
row becomes RESUMED, so this daemon naturally exits after graph acceptance.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
lease: WorkflowHandoffLease,
|
||||
clock: Callable[[], datetime],
|
||||
interval: timedelta,
|
||||
) -> None:
|
||||
if interval.total_seconds() <= 0:
|
||||
raise ValueError("heartbeat interval must be positive")
|
||||
self._lease = lease
|
||||
self._clock = clock
|
||||
self._interval_seconds = interval.total_seconds()
|
||||
self._stopped = Event()
|
||||
self._thread = Thread(
|
||||
target=self._renew_until_stopped,
|
||||
name=f"workflow-handoff-lease-{lease.handoff_id}",
|
||||
daemon=True,
|
||||
)
|
||||
|
||||
@override
|
||||
def __enter__(self) -> None:
|
||||
self._thread.start()
|
||||
|
||||
@override
|
||||
def __exit__(
|
||||
self,
|
||||
exc_type: type[BaseException] | None,
|
||||
exc_value: BaseException | None,
|
||||
traceback: TracebackType | None,
|
||||
) -> None:
|
||||
self._stopped.set()
|
||||
self._thread.join()
|
||||
|
||||
def _renew_until_stopped(self) -> None:
|
||||
while not self._stopped.wait(self._interval_seconds):
|
||||
try:
|
||||
if not self._lease.renew(now=self._clock()):
|
||||
return
|
||||
except Exception:
|
||||
# A transient database failure must not crash the process. The
|
||||
# next heartbeat can recover before expiry; the fenced ACK still
|
||||
# prevents this worker from accepting a lease it ultimately lost.
|
||||
logger.exception(
|
||||
"Failed to renew workflow handoff lease: handoff_id=%s, generation=%s",
|
||||
self._lease.handoff_id,
|
||||
self._lease.generation,
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class WorkflowHandoffResumeRequest:
|
||||
"""Verified checkpoint and claim identity passed to a business resumer."""
|
||||
|
||||
handoff: WorkflowRunHandoff
|
||||
serialized_state: bytes
|
||||
lease: WorkflowHandoffLease
|
||||
|
||||
|
||||
class WorkflowHandoffResumeDispatcher(Protocol):
|
||||
"""Route a claimed checkpoint to Workflow, Chatflow, Trigger, or RAG resume code.
|
||||
|
||||
The handler must install ``WorkflowHandoffResumeAcknowledgementLayer`` and
|
||||
call its explicit check on the resumption start event. Returning before the
|
||||
handoff reaches ``RESUMED`` is treated as a retryable setup failure.
|
||||
"""
|
||||
|
||||
def dispatch(self, request: WorkflowHandoffResumeRequest) -> None: ...
|
||||
|
||||
|
||||
type WorkflowHandoffResumeHandler = Callable[[WorkflowHandoffResumeRequest], None]
|
||||
|
||||
|
||||
class MappingWorkflowHandoffResumeDispatcher:
|
||||
"""Small dependency-injection adapter for route-specific callback functions."""
|
||||
|
||||
def __init__(self, handlers: Mapping[WorkflowHandoffResumeRoute, WorkflowHandoffResumeHandler]):
|
||||
self._handlers = dict(handlers)
|
||||
|
||||
def dispatch(self, request: WorkflowHandoffResumeRequest) -> None:
|
||||
handler = self._handlers.get(request.handoff.resume_route)
|
||||
if handler is None:
|
||||
raise UnsupportedWorkflowHandoffResumeRouteError(
|
||||
f"No workflow handoff resume handler for route: {request.handoff.resume_route}"
|
||||
)
|
||||
handler(request)
|
||||
|
||||
|
||||
class WorkflowHandoffResumeOutcome(StrEnum):
|
||||
CLAIM_NOT_ACQUIRED = "claim_not_acquired"
|
||||
RESUMED = "resumed"
|
||||
RETRY_SCHEDULED = "retry_scheduled"
|
||||
FAILED = "failed"
|
||||
LEASE_LOST = "lease_lost"
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class WorkflowHandoffResumeResult:
|
||||
outcome: WorkflowHandoffResumeOutcome
|
||||
handoff_id: str
|
||||
generation: int
|
||||
error: str | None = None
|
||||
|
||||
|
||||
class WorkflowHandoffResumeCoordinator:
|
||||
"""Claim, verify, and dispatch one idempotent workflow handoff message.
|
||||
|
||||
Celery delivery is at-least-once. ``claim`` provides the exclusive fence;
|
||||
only that claim's lease token may acknowledge or release the generation.
|
||||
Retry state is persisted in the outbox row rather than delegated to Celery,
|
||||
so a periodic scan recovers broker loss and worker interruption uniformly.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
repository: WorkflowRunHandoffRepository,
|
||||
handoff_service: WorkflowHandoffService,
|
||||
lease_duration: timedelta,
|
||||
retry_delay: timedelta,
|
||||
max_attempts: int,
|
||||
clock: Callable[[], datetime] = naive_utc_now,
|
||||
lease_heartbeat_factory: Callable[[WorkflowHandoffLease], AbstractContextManager[None]] | None = None,
|
||||
) -> None:
|
||||
if lease_duration.total_seconds() <= 0:
|
||||
raise ValueError("lease_duration must be positive")
|
||||
if retry_delay.total_seconds() < 0:
|
||||
raise ValueError("retry_delay must be non-negative")
|
||||
if max_attempts <= 0:
|
||||
raise ValueError("max_attempts must be positive")
|
||||
self._repository = repository
|
||||
self._handoff_service = handoff_service
|
||||
self._lease_duration = lease_duration
|
||||
self._retry_delay = retry_delay
|
||||
self._max_attempts = max_attempts
|
||||
self._clock = clock
|
||||
heartbeat_interval = timedelta(seconds=max(1.0, lease_duration.total_seconds() / 3))
|
||||
self._lease_heartbeat_factory = lease_heartbeat_factory or (
|
||||
lambda lease: _WorkflowHandoffLeaseHeartbeat(
|
||||
lease=lease,
|
||||
clock=self._clock,
|
||||
interval=heartbeat_interval,
|
||||
)
|
||||
)
|
||||
|
||||
def resume(
|
||||
self,
|
||||
*,
|
||||
handoff_id: str,
|
||||
generation: int,
|
||||
lease_owner: str,
|
||||
now: datetime,
|
||||
dispatcher: WorkflowHandoffResumeDispatcher,
|
||||
) -> WorkflowHandoffResumeResult:
|
||||
claimed = self._repository.claim(
|
||||
handoff_id=handoff_id,
|
||||
generation=generation,
|
||||
lease_owner=lease_owner,
|
||||
lease_duration=self._lease_duration,
|
||||
max_attempts=self._max_attempts,
|
||||
now=now,
|
||||
)
|
||||
if claimed is None:
|
||||
return WorkflowHandoffResumeResult(
|
||||
outcome=WorkflowHandoffResumeOutcome.CLAIM_NOT_ACQUIRED,
|
||||
handoff_id=handoff_id,
|
||||
generation=generation,
|
||||
)
|
||||
|
||||
lease = self._lease_from_claim(claimed)
|
||||
try:
|
||||
serialized_state = self._handoff_service.load_and_verify_state(claimed)
|
||||
except (UnsupportedWorkflowHandoffSnapshotVersionError, WorkflowHandoffSnapshotIntegrityError) as error:
|
||||
return self._fail_permanently(claimed=claimed, error=error, now=self._clock())
|
||||
except Exception as error:
|
||||
logger.exception("Failed to load workflow handoff checkpoint: handoff_id=%s", claimed.id)
|
||||
return self._schedule_retry(claimed=claimed, error=error, now=self._clock())
|
||||
|
||||
# Loading from remote object storage can consume a meaningful portion of
|
||||
# a short lease. Refresh it once before route-specific setup. The handler
|
||||
# also receives this fenced lease handle if additional setup is lengthy.
|
||||
if not lease.renew(now=self._clock()):
|
||||
return WorkflowHandoffResumeResult(
|
||||
outcome=WorkflowHandoffResumeOutcome.LEASE_LOST,
|
||||
handoff_id=claimed.id,
|
||||
generation=claimed.generation,
|
||||
error="workflow handoff lease was lost before dispatch",
|
||||
)
|
||||
|
||||
request = WorkflowHandoffResumeRequest(
|
||||
handoff=claimed,
|
||||
serialized_state=serialized_state,
|
||||
lease=lease,
|
||||
)
|
||||
try:
|
||||
with self._lease_heartbeat_factory(lease):
|
||||
dispatcher.dispatch(request)
|
||||
except PermanentWorkflowHandoffResumeError as error:
|
||||
return self._fail_permanently(claimed=claimed, error=error, now=self._clock())
|
||||
except Exception as error:
|
||||
logger.exception(
|
||||
"Workflow handoff resume handler failed: handoff_id=%s, route=%s",
|
||||
claimed.id,
|
||||
claimed.resume_route,
|
||||
)
|
||||
current = self._repository.get(claimed.id, claimed.generation)
|
||||
if current is not None and current.state == WorkflowHandoffState.RESUMED:
|
||||
# The graph already accepted the checkpoint. Retrying the handoff
|
||||
# would duplicate execution; ordinary workflow failure handling
|
||||
# now owns this runtime error.
|
||||
return WorkflowHandoffResumeResult(
|
||||
outcome=WorkflowHandoffResumeOutcome.RESUMED,
|
||||
handoff_id=claimed.id,
|
||||
generation=claimed.generation,
|
||||
error=self._error_text(error),
|
||||
)
|
||||
if current is not None and current.state == WorkflowHandoffState.FAILED:
|
||||
# The post-ACK stream drain reconciles the run and route-owned
|
||||
# records before raising. Report that durable terminal outcome
|
||||
# directly; retrying the released claim would be both stale and
|
||||
# capable of duplicating resumed node execution.
|
||||
return WorkflowHandoffResumeResult(
|
||||
outcome=WorkflowHandoffResumeOutcome.FAILED,
|
||||
handoff_id=claimed.id,
|
||||
generation=claimed.generation,
|
||||
error=current.last_error or self._error_text(error),
|
||||
)
|
||||
return self._schedule_retry(claimed=claimed, error=error, now=self._clock())
|
||||
|
||||
current = self._repository.get(claimed.id, claimed.generation)
|
||||
if current is None:
|
||||
return WorkflowHandoffResumeResult(
|
||||
outcome=WorkflowHandoffResumeOutcome.LEASE_LOST,
|
||||
handoff_id=claimed.id,
|
||||
generation=claimed.generation,
|
||||
error="workflow handoff disappeared after dispatch",
|
||||
)
|
||||
if current.state == WorkflowHandoffState.RESUMED:
|
||||
return WorkflowHandoffResumeResult(
|
||||
outcome=WorkflowHandoffResumeOutcome.RESUMED,
|
||||
handoff_id=claimed.id,
|
||||
generation=claimed.generation,
|
||||
)
|
||||
if current.state == WorkflowHandoffState.FAILED:
|
||||
return WorkflowHandoffResumeResult(
|
||||
outcome=WorkflowHandoffResumeOutcome.FAILED,
|
||||
handoff_id=claimed.id,
|
||||
generation=claimed.generation,
|
||||
error=current.last_error,
|
||||
)
|
||||
return self._schedule_retry(
|
||||
claimed=claimed,
|
||||
error=RuntimeError("workflow handoff handler returned before acknowledgement"),
|
||||
now=self._clock(),
|
||||
)
|
||||
|
||||
def _lease_from_claim(self, claimed: WorkflowRunHandoff) -> WorkflowHandoffLease:
|
||||
if not claimed.lease_owner or not claimed.lease_token:
|
||||
raise RuntimeError(f"Claimed workflow handoff has incomplete lease identity: {claimed.id}")
|
||||
return WorkflowHandoffLease(
|
||||
repository=self._repository,
|
||||
handoff_id=claimed.id,
|
||||
generation=claimed.generation,
|
||||
lease_owner=claimed.lease_owner,
|
||||
lease_token=claimed.lease_token,
|
||||
lease_duration=self._lease_duration,
|
||||
)
|
||||
|
||||
def _schedule_retry(
|
||||
self,
|
||||
*,
|
||||
claimed: WorkflowRunHandoff,
|
||||
error: Exception,
|
||||
now: datetime,
|
||||
) -> WorkflowHandoffResumeResult:
|
||||
if not claimed.lease_owner or not claimed.lease_token:
|
||||
return WorkflowHandoffResumeResult(
|
||||
outcome=WorkflowHandoffResumeOutcome.LEASE_LOST,
|
||||
handoff_id=claimed.id,
|
||||
generation=claimed.generation,
|
||||
error=self._error_text(error),
|
||||
)
|
||||
updated = self._repository.record_failure(
|
||||
handoff_id=claimed.id,
|
||||
generation=claimed.generation,
|
||||
lease_owner=claimed.lease_owner,
|
||||
lease_token=claimed.lease_token,
|
||||
error=self._error_text(error),
|
||||
retry_at=now + self._retry_delay,
|
||||
max_attempts=self._max_attempts,
|
||||
now=now,
|
||||
)
|
||||
if updated is None:
|
||||
outcome = WorkflowHandoffResumeOutcome.LEASE_LOST
|
||||
elif updated.state == WorkflowHandoffState.FAILED:
|
||||
outcome = WorkflowHandoffResumeOutcome.FAILED
|
||||
else:
|
||||
outcome = WorkflowHandoffResumeOutcome.RETRY_SCHEDULED
|
||||
return WorkflowHandoffResumeResult(
|
||||
outcome=outcome,
|
||||
handoff_id=claimed.id,
|
||||
generation=claimed.generation,
|
||||
error=self._error_text(error),
|
||||
)
|
||||
|
||||
def _fail_permanently(
|
||||
self,
|
||||
*,
|
||||
claimed: WorkflowRunHandoff,
|
||||
error: Exception,
|
||||
now: datetime,
|
||||
) -> WorkflowHandoffResumeResult:
|
||||
marked = self._repository.mark_failed(
|
||||
handoff_id=claimed.id,
|
||||
generation=claimed.generation,
|
||||
error=self._error_text(error),
|
||||
failed_at=now,
|
||||
lease_owner=claimed.lease_owner,
|
||||
lease_token=claimed.lease_token,
|
||||
)
|
||||
return WorkflowHandoffResumeResult(
|
||||
outcome=(WorkflowHandoffResumeOutcome.FAILED if marked else WorkflowHandoffResumeOutcome.LEASE_LOST),
|
||||
handoff_id=claimed.id,
|
||||
generation=claimed.generation,
|
||||
error=self._error_text(error),
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _error_text(error: Exception) -> str:
|
||||
text = str(error) or error.__class__.__name__
|
||||
return text[:4000]
|
||||
|
||||
|
||||
__all__ = [
|
||||
"MappingWorkflowHandoffResumeDispatcher",
|
||||
"PermanentWorkflowHandoffResumeError",
|
||||
"UnsupportedWorkflowHandoffResumeRouteError",
|
||||
"WorkflowHandoffLease",
|
||||
"WorkflowHandoffResumeCoordinator",
|
||||
"WorkflowHandoffResumeDispatcher",
|
||||
"WorkflowHandoffResumeHandler",
|
||||
"WorkflowHandoffResumeOutcome",
|
||||
"WorkflowHandoffResumeRequest",
|
||||
"WorkflowHandoffResumeResult",
|
||||
]
|
||||
@@ -0,0 +1,620 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Generator
|
||||
from datetime import UTC, datetime
|
||||
|
||||
from flask import g
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.orm import Session, sessionmaker
|
||||
|
||||
from configs import dify_config
|
||||
from core.app.apps.advanced_chat.app_generator import AdvancedChatAppGenerator
|
||||
from core.app.apps.pipeline.pipeline_generator import PipelineGenerator
|
||||
from core.app.apps.workflow.app_generator import WorkflowAppGenerator
|
||||
from core.app.entities.app_invoke_entities import (
|
||||
AdvancedChatAppGenerateEntity,
|
||||
RagPipelineGenerateEntity,
|
||||
WorkflowAppGenerateEntity,
|
||||
)
|
||||
from core.app.layers.pause_state_persist_layer import PauseStateLayerConfig, WorkflowResumptionContext
|
||||
from core.app.layers.trigger_post_layer import TriggerPostLayer
|
||||
from core.app.layers.workflow_handoff_resume_layer import WorkflowHandoffResumeAcknowledgementLayer
|
||||
from core.repositories import DifyCoreRepositoryFactory
|
||||
from extensions.ext_database import db
|
||||
from extensions.ext_storage import storage
|
||||
from graphon.enums import WorkflowExecutionStatus
|
||||
from graphon.runtime import GraphRuntimeState
|
||||
from graphon.variable_loader import DUMMY_VARIABLE_LOADER, VariableLoader
|
||||
from libs.datetime_utils import naive_utc_now
|
||||
from libs.flask_utils import set_login_user
|
||||
from models.account import Account
|
||||
from models.dataset import Document, Pipeline
|
||||
from models.enums import CreatorUserRole, WorkflowRunTriggeredFrom
|
||||
from models.model import App, AppMode, Conversation, EndUser, Message, Tenant
|
||||
from models.snippet import CustomizedSnippet
|
||||
from models.trigger import WorkflowTriggerLog
|
||||
from models.workflow import Workflow, WorkflowKind, WorkflowNodeExecutionTriggeredFrom, WorkflowRun
|
||||
from models.workflow_handoff import WorkflowHandoffResumeRoute
|
||||
from repositories.workflow_handoff_repository import WorkflowHandoffTerminalScope
|
||||
from services.snippet_generate_service import SnippetGenerateService
|
||||
from services.workflow_draft_variable_service import DraftVarLoader
|
||||
from services.workflow_handoff_resume_coordinator import (
|
||||
MappingWorkflowHandoffResumeDispatcher,
|
||||
PermanentWorkflowHandoffResumeError,
|
||||
WorkflowHandoffResumeRequest,
|
||||
)
|
||||
from services.workflow_handoff_terminal_service import WorkflowHandoffTerminalService
|
||||
from tasks.app_generate.workflow_execute_task import (
|
||||
WorkflowStreamTerminalFailure,
|
||||
WorkflowStreamTerminalFailureHandler,
|
||||
_publish_streaming_response,
|
||||
)
|
||||
from tasks.workflow_cfs_scheduler.cfs_scheduler import AsyncWorkflowCFSPlanEntity
|
||||
from tasks.workflow_cfs_scheduler.entities import AsyncWorkflowQueue, AsyncWorkflowSystemStrategy
|
||||
|
||||
|
||||
def create_workflow_handoff_resume_dispatcher() -> MappingWorkflowHandoffResumeDispatcher:
|
||||
"""Build the business resumers behind the generic fenced claim task."""
|
||||
return MappingWorkflowHandoffResumeDispatcher(
|
||||
{
|
||||
WorkflowHandoffResumeRoute.WORKFLOW: _resume_workflow_handoff,
|
||||
WorkflowHandoffResumeRoute.SNIPPET: _resume_snippet_handoff,
|
||||
WorkflowHandoffResumeRoute.ADVANCED_CHAT: _resume_advanced_chat_handoff,
|
||||
WorkflowHandoffResumeRoute.TRIGGERED_WORKFLOW: _resume_triggered_workflow_handoff,
|
||||
WorkflowHandoffResumeRoute.RAG_PIPELINE: _resume_rag_pipeline_handoff,
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def _load_context(
|
||||
request: WorkflowHandoffResumeRequest,
|
||||
) -> tuple[
|
||||
WorkflowResumptionContext,
|
||||
WorkflowAppGenerateEntity | AdvancedChatAppGenerateEntity | RagPipelineGenerateEntity,
|
||||
GraphRuntimeState,
|
||||
]:
|
||||
# Creation can be disabled independently during a rollout, but an already
|
||||
# durable handoff must never resume onto a lossy event transport. Without
|
||||
# Redis Streams, reconnecting clients could silently miss the resumed
|
||||
# segment, so treat this as an incompatible runtime and fail closed.
|
||||
if dify_config.PUBSUB_REDIS_CHANNEL_TYPE != "streams":
|
||||
raise PermanentWorkflowHandoffResumeError(
|
||||
"Workflow handoff resumption requires EVENT_BUS_REDIS_CHANNEL_TYPE=streams"
|
||||
)
|
||||
try:
|
||||
context = WorkflowResumptionContext.loads(request.serialized_state.decode())
|
||||
generate_entity = context.get_generate_entity()
|
||||
graph_runtime_state = GraphRuntimeState.from_snapshot(context.serialized_graph_runtime_state)
|
||||
except Exception as error:
|
||||
raise PermanentWorkflowHandoffResumeError("Invalid workflow handoff resumption context") from error
|
||||
|
||||
if generate_entity.task_id != request.handoff.task_id:
|
||||
raise PermanentWorkflowHandoffResumeError("Workflow handoff task identity does not match its checkpoint")
|
||||
|
||||
entity_run_id = (
|
||||
generate_entity.workflow_run_id
|
||||
if isinstance(generate_entity, AdvancedChatAppGenerateEntity)
|
||||
else generate_entity.workflow_execution_id
|
||||
)
|
||||
if entity_run_id != request.handoff.workflow_run_id:
|
||||
raise PermanentWorkflowHandoffResumeError("Workflow handoff run identity does not match its checkpoint")
|
||||
# A resumed segment has no original blocking HTTP request to return to.
|
||||
# Always expose its events through the durable per-run continuation stream;
|
||||
# the original response mode remains relevant only to the source segment's
|
||||
# public 202/streaming response contract.
|
||||
generate_entity = generate_entity.model_copy(update={"stream": True})
|
||||
context.apply_handoff_execution_timing()
|
||||
return context, generate_entity, graph_runtime_state
|
||||
|
||||
|
||||
def _load_run_dependencies(
|
||||
session: Session,
|
||||
request: WorkflowHandoffResumeRequest,
|
||||
) -> tuple[WorkflowRun, Workflow, Account | EndUser]:
|
||||
workflow_run = session.get(WorkflowRun, request.handoff.workflow_run_id)
|
||||
if workflow_run is None:
|
||||
raise PermanentWorkflowHandoffResumeError("Workflow run no longer exists")
|
||||
if workflow_run.status != WorkflowExecutionStatus.RUNNING:
|
||||
raise PermanentWorkflowHandoffResumeError(f"Workflow run is no longer resumable: status={workflow_run.status}")
|
||||
|
||||
workflow = session.get(Workflow, workflow_run.workflow_id)
|
||||
if workflow is None:
|
||||
raise PermanentWorkflowHandoffResumeError("Workflow definition no longer exists")
|
||||
_validate_owned_resource(
|
||||
"Workflow definition",
|
||||
{
|
||||
"tenant": (workflow.tenant_id, workflow_run.tenant_id),
|
||||
"app": (workflow.app_id, workflow_run.app_id),
|
||||
"workflow": (workflow.id, workflow_run.workflow_id),
|
||||
},
|
||||
)
|
||||
user = _resolve_user(session, workflow_run)
|
||||
return workflow_run, workflow, user
|
||||
|
||||
|
||||
def _validate_entity_identity(
|
||||
*,
|
||||
generate_entity: WorkflowAppGenerateEntity | AdvancedChatAppGenerateEntity | RagPipelineGenerateEntity,
|
||||
workflow_run: WorkflowRun,
|
||||
workflow: Workflow,
|
||||
) -> None:
|
||||
app_config = generate_entity.app_config
|
||||
mismatches = {
|
||||
"tenant": (app_config.tenant_id, workflow_run.tenant_id),
|
||||
"app": (app_config.app_id, workflow_run.app_id),
|
||||
"workflow": (app_config.workflow_id, workflow_run.workflow_id),
|
||||
"user": (generate_entity.user_id, workflow_run.created_by),
|
||||
"workflow_definition": (workflow.id, workflow_run.workflow_id),
|
||||
}
|
||||
invalid = [name for name, (snapshot_value, run_value) in mismatches.items() if snapshot_value != run_value]
|
||||
if invalid:
|
||||
raise PermanentWorkflowHandoffResumeError(f"Workflow handoff identity mismatch: {', '.join(sorted(invalid))}")
|
||||
|
||||
|
||||
def _validate_owned_resource(resource: str, identities: dict[str, tuple[object, object]]) -> None:
|
||||
invalid = [name for name, (resource_value, run_value) in identities.items() if resource_value != run_value]
|
||||
if invalid:
|
||||
raise PermanentWorkflowHandoffResumeError(f"{resource} ownership mismatch: {', '.join(sorted(invalid))}")
|
||||
|
||||
|
||||
def _validate_app_ownership(app: App, workflow_run: WorkflowRun) -> None:
|
||||
_validate_owned_resource(
|
||||
"Workflow app",
|
||||
{
|
||||
"app": (app.id, workflow_run.app_id),
|
||||
"tenant": (app.tenant_id, workflow_run.tenant_id),
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
def _validate_chat_records(
|
||||
*,
|
||||
conversation: Conversation,
|
||||
message: Message,
|
||||
workflow_run: WorkflowRun,
|
||||
) -> None:
|
||||
conversation_identities: dict[str, tuple[object, object]] = {
|
||||
"app": (conversation.app_id, workflow_run.app_id),
|
||||
}
|
||||
# Conversation has no tenant column in current schemas. Keep this check
|
||||
# forward-compatible for deployments that expose one without inferring
|
||||
# tenant ownership from unrelated fields.
|
||||
conversation_tenant_id = getattr( # guard-ignore: no-new-getattr -- optional forward-schema tenant column
|
||||
conversation, "tenant_id", None
|
||||
)
|
||||
if conversation_tenant_id is not None:
|
||||
conversation_identities["tenant"] = (conversation_tenant_id, workflow_run.tenant_id)
|
||||
_validate_owned_resource("Chatflow conversation", conversation_identities)
|
||||
_validate_owned_resource(
|
||||
"Chatflow message",
|
||||
{
|
||||
"app": (message.app_id, workflow_run.app_id),
|
||||
"conversation": (message.conversation_id, conversation.id),
|
||||
"workflow_run": (message.workflow_run_id, workflow_run.id),
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
def _resolve_user(session: Session, workflow_run: WorkflowRun) -> Account | EndUser:
|
||||
tenant = session.get(Tenant, workflow_run.tenant_id)
|
||||
if tenant is None:
|
||||
raise PermanentWorkflowHandoffResumeError("Workflow tenant no longer exists")
|
||||
|
||||
if workflow_run.created_by_role == CreatorUserRole.ACCOUNT:
|
||||
account = session.get(Account, workflow_run.created_by)
|
||||
if account is None:
|
||||
raise PermanentWorkflowHandoffResumeError("Workflow account no longer exists")
|
||||
account.set_current_tenant_with_session(tenant, session=session)
|
||||
return account
|
||||
|
||||
end_user = session.get(EndUser, workflow_run.created_by)
|
||||
if end_user is None:
|
||||
raise PermanentWorkflowHandoffResumeError("Workflow end user no longer exists")
|
||||
return end_user
|
||||
|
||||
|
||||
def _acknowledgement_layer(request: WorkflowHandoffResumeRequest) -> WorkflowHandoffResumeAcknowledgementLayer:
|
||||
return WorkflowHandoffResumeAcknowledgementLayer(
|
||||
repository=request.lease.repository,
|
||||
claimed_handoff=request.handoff,
|
||||
)
|
||||
|
||||
|
||||
def _resumed_terminal_failure_handler(
|
||||
request: WorkflowHandoffResumeRequest,
|
||||
workflow_run: WorkflowRun,
|
||||
) -> WorkflowStreamTerminalFailureHandler:
|
||||
service = WorkflowHandoffTerminalService(repository=request.lease.repository, storage=storage)
|
||||
scope = WorkflowHandoffTerminalScope(
|
||||
workflow_run_id=workflow_run.id,
|
||||
task_id=request.handoff.task_id,
|
||||
tenant_id=workflow_run.tenant_id,
|
||||
app_id=workflow_run.app_id,
|
||||
workflow_id=workflow_run.workflow_id,
|
||||
resume_route=request.handoff.resume_route,
|
||||
)
|
||||
|
||||
def reconcile(failure: WorkflowStreamTerminalFailure) -> None:
|
||||
service.reconcile_resumed_failure(
|
||||
handoff_id=request.handoff.id,
|
||||
generation=request.handoff.generation,
|
||||
scope=scope,
|
||||
error=failure.error,
|
||||
failed_at=naive_utc_now(),
|
||||
message_answer_delta=failure.message_answer_delta,
|
||||
message_answer_replacement=failure.message_answer_replacement,
|
||||
)
|
||||
|
||||
return reconcile
|
||||
|
||||
|
||||
def _build_repositories(
|
||||
*,
|
||||
session_factory: sessionmaker[Session],
|
||||
workflow_run: WorkflowRun,
|
||||
user: Account | EndUser,
|
||||
generate_entity: WorkflowAppGenerateEntity | AdvancedChatAppGenerateEntity | RagPipelineGenerateEntity,
|
||||
):
|
||||
triggered_from = WorkflowRunTriggeredFrom(workflow_run.triggered_from)
|
||||
workflow_execution_repository = DifyCoreRepositoryFactory.create_workflow_execution_repository(
|
||||
session_factory=session_factory,
|
||||
tenant_id=workflow_run.tenant_id,
|
||||
user=user,
|
||||
app_id=workflow_run.app_id,
|
||||
triggered_from=triggered_from,
|
||||
)
|
||||
if generate_entity.single_iteration_run is not None or generate_entity.single_loop_run is not None:
|
||||
node_triggered_from = WorkflowNodeExecutionTriggeredFrom.SINGLE_STEP
|
||||
elif triggered_from in {
|
||||
WorkflowRunTriggeredFrom.RAG_PIPELINE_RUN,
|
||||
WorkflowRunTriggeredFrom.RAG_PIPELINE_DEBUGGING,
|
||||
}:
|
||||
node_triggered_from = WorkflowNodeExecutionTriggeredFrom.RAG_PIPELINE_RUN
|
||||
else:
|
||||
node_triggered_from = WorkflowNodeExecutionTriggeredFrom.WORKFLOW_RUN
|
||||
workflow_node_execution_repository = DifyCoreRepositoryFactory.create_workflow_node_execution_repository(
|
||||
session_factory=session_factory,
|
||||
tenant_id=workflow_run.tenant_id,
|
||||
user=user,
|
||||
app_id=workflow_run.app_id,
|
||||
triggered_from=node_triggered_from,
|
||||
)
|
||||
return workflow_execution_repository, workflow_node_execution_repository
|
||||
|
||||
|
||||
def _trigger_layers(
|
||||
session: Session,
|
||||
workflow_run: WorkflowRun,
|
||||
) -> list[TriggerPostLayer]:
|
||||
trigger_log = session.scalar(
|
||||
select(WorkflowTriggerLog).where(WorkflowTriggerLog.workflow_run_id == workflow_run.id).limit(1)
|
||||
)
|
||||
if trigger_log is None:
|
||||
return []
|
||||
|
||||
_validate_owned_resource(
|
||||
"Workflow trigger log",
|
||||
{
|
||||
"tenant": (trigger_log.tenant_id, workflow_run.tenant_id),
|
||||
"app": (trigger_log.app_id, workflow_run.app_id),
|
||||
"workflow": (trigger_log.workflow_id, workflow_run.workflow_id),
|
||||
"workflow_run": (trigger_log.workflow_run_id, workflow_run.id),
|
||||
},
|
||||
)
|
||||
|
||||
scheduler_entity = AsyncWorkflowCFSPlanEntity(
|
||||
queue=AsyncWorkflowQueue(trigger_log.queue_name),
|
||||
schedule_strategy=AsyncWorkflowSystemStrategy,
|
||||
granularity=dify_config.ASYNC_WORKFLOW_SCHEDULER_GRANULARITY,
|
||||
)
|
||||
# Match initial triggered execution exactly: time slicing remains disabled
|
||||
# there, so a maintenance handoff must not silently enable it on resume.
|
||||
return [TriggerPostLayer(scheduler_entity, datetime.now(UTC), trigger_log.id)]
|
||||
|
||||
|
||||
def _resume_workflow_handoff(request: WorkflowHandoffResumeRequest) -> None:
|
||||
_resume_workflow_route(request, require_trigger_log=False)
|
||||
|
||||
|
||||
def _resume_triggered_workflow_handoff(request: WorkflowHandoffResumeRequest) -> None:
|
||||
_resume_workflow_route(request, require_trigger_log=True)
|
||||
|
||||
|
||||
def _resume_snippet_handoff(request: WorkflowHandoffResumeRequest) -> None:
|
||||
context, generate_entity, graph_runtime_state = _load_context(request)
|
||||
if not isinstance(generate_entity, WorkflowAppGenerateEntity) or isinstance(
|
||||
generate_entity, RagPipelineGenerateEntity
|
||||
):
|
||||
raise PermanentWorkflowHandoffResumeError("Snippet handoff contains an incompatible generate entity")
|
||||
|
||||
session_factory = sessionmaker(bind=db.engine, expire_on_commit=False)
|
||||
with session_factory() as session:
|
||||
workflow_run, workflow, user = _load_run_dependencies(session, request)
|
||||
_validate_entity_identity(
|
||||
generate_entity=generate_entity,
|
||||
workflow_run=workflow_run,
|
||||
workflow=workflow,
|
||||
)
|
||||
snippet = session.get(CustomizedSnippet, workflow_run.app_id)
|
||||
if snippet is None:
|
||||
raise PermanentWorkflowHandoffResumeError("Snippet no longer exists")
|
||||
_validate_owned_resource(
|
||||
"Snippet",
|
||||
{
|
||||
"snippet": (snippet.id, workflow_run.app_id),
|
||||
"tenant": (snippet.tenant_id, workflow_run.tenant_id),
|
||||
},
|
||||
)
|
||||
if workflow.kind_or_standard != WorkflowKind.SNIPPET.value:
|
||||
raise PermanentWorkflowHandoffResumeError("Workflow definition is not owned by a Snippet")
|
||||
|
||||
app_model = SnippetGenerateService.build_app_model(snippet)
|
||||
workflow_execution_repository, workflow_node_execution_repository = _build_repositories(
|
||||
session_factory=session_factory,
|
||||
workflow_run=workflow_run,
|
||||
user=user,
|
||||
generate_entity=generate_entity,
|
||||
)
|
||||
variable_loader: VariableLoader = DUMMY_VARIABLE_LOADER
|
||||
if generate_entity.single_iteration_run is not None or generate_entity.single_loop_run is not None:
|
||||
variable_loader = DraftVarLoader(
|
||||
engine=db.engine,
|
||||
app_id=workflow_run.app_id,
|
||||
tenant_id=workflow_run.tenant_id,
|
||||
user_id=user.id,
|
||||
)
|
||||
|
||||
set_login_user(user)
|
||||
response = WorkflowAppGenerator().resume(
|
||||
app_model=app_model,
|
||||
workflow=workflow,
|
||||
user=user,
|
||||
application_generate_entity=generate_entity,
|
||||
graph_runtime_state=graph_runtime_state,
|
||||
workflow_execution_repository=workflow_execution_repository,
|
||||
workflow_node_execution_repository=workflow_node_execution_repository,
|
||||
graph_engine_layers=[_acknowledgement_layer(request)],
|
||||
pause_state_config=PauseStateLayerConfig(
|
||||
session_factory=session_factory,
|
||||
state_owner_user_id=workflow.created_by,
|
||||
),
|
||||
variable_loader=variable_loader,
|
||||
response_stream_filter=context.get_response_stream_filter(),
|
||||
handoff_resume_route=request.handoff.resume_route,
|
||||
graph_config=workflow_run.graph_dict,
|
||||
workflow_version=workflow_run.version,
|
||||
root_node_id=context.root_node_id,
|
||||
)
|
||||
if isinstance(response, Generator):
|
||||
_publish_streaming_response(
|
||||
SnippetGenerateService.filter_virtual_start_events(response),
|
||||
workflow_run.id,
|
||||
AppMode.WORKFLOW,
|
||||
workflow.id,
|
||||
generate_entity.inputs,
|
||||
started_reason=_resumption_reason(),
|
||||
terminal_failure_handler=_resumed_terminal_failure_handler(request, workflow_run),
|
||||
)
|
||||
|
||||
|
||||
def _resume_workflow_route(request: WorkflowHandoffResumeRequest, *, require_trigger_log: bool) -> None:
|
||||
context, generate_entity, graph_runtime_state = _load_context(request)
|
||||
if not isinstance(generate_entity, WorkflowAppGenerateEntity) or isinstance(
|
||||
generate_entity, RagPipelineGenerateEntity
|
||||
):
|
||||
raise PermanentWorkflowHandoffResumeError("Workflow handoff contains an incompatible generate entity")
|
||||
|
||||
session_factory = sessionmaker(bind=db.engine, expire_on_commit=False)
|
||||
with session_factory() as session:
|
||||
workflow_run, workflow, user = _load_run_dependencies(session, request)
|
||||
_validate_entity_identity(
|
||||
generate_entity=generate_entity,
|
||||
workflow_run=workflow_run,
|
||||
workflow=workflow,
|
||||
)
|
||||
triggered_sources = {
|
||||
WorkflowRunTriggeredFrom.WEBHOOK,
|
||||
WorkflowRunTriggeredFrom.SCHEDULE,
|
||||
WorkflowRunTriggeredFrom.PLUGIN,
|
||||
}
|
||||
if (workflow_run.triggered_from in triggered_sources) != require_trigger_log:
|
||||
raise PermanentWorkflowHandoffResumeError("Workflow handoff trigger route does not match the run source")
|
||||
app_model = session.get(App, workflow_run.app_id)
|
||||
if app_model is None:
|
||||
raise PermanentWorkflowHandoffResumeError("Workflow app no longer exists")
|
||||
_validate_app_ownership(app_model, workflow_run)
|
||||
trigger_layers = _trigger_layers(session, workflow_run) if require_trigger_log else []
|
||||
if require_trigger_log and not trigger_layers:
|
||||
raise PermanentWorkflowHandoffResumeError("Triggered workflow handoff has no trigger log")
|
||||
|
||||
set_login_user(user)
|
||||
workflow_execution_repository, workflow_node_execution_repository = _build_repositories(
|
||||
session_factory=session_factory,
|
||||
workflow_run=workflow_run,
|
||||
user=user,
|
||||
generate_entity=generate_entity,
|
||||
)
|
||||
layers = [_acknowledgement_layer(request), *trigger_layers]
|
||||
response = WorkflowAppGenerator().resume(
|
||||
app_model=app_model,
|
||||
workflow=workflow,
|
||||
user=user,
|
||||
application_generate_entity=generate_entity,
|
||||
graph_runtime_state=graph_runtime_state,
|
||||
workflow_execution_repository=workflow_execution_repository,
|
||||
workflow_node_execution_repository=workflow_node_execution_repository,
|
||||
graph_engine_layers=layers,
|
||||
pause_state_config=PauseStateLayerConfig(
|
||||
session_factory=session_factory,
|
||||
state_owner_user_id=workflow.created_by,
|
||||
),
|
||||
response_stream_filter=context.get_response_stream_filter(),
|
||||
handoff_resume_route=request.handoff.resume_route,
|
||||
graph_config=workflow_run.graph_dict,
|
||||
workflow_version=workflow_run.version,
|
||||
root_node_id=context.root_node_id,
|
||||
)
|
||||
if isinstance(response, Generator):
|
||||
_publish_streaming_response(
|
||||
response,
|
||||
workflow_run.id,
|
||||
AppMode.WORKFLOW,
|
||||
workflow.id,
|
||||
generate_entity.inputs,
|
||||
started_reason=_resumption_reason(),
|
||||
terminal_failure_handler=_resumed_terminal_failure_handler(request, workflow_run),
|
||||
)
|
||||
|
||||
|
||||
def _resume_advanced_chat_handoff(request: WorkflowHandoffResumeRequest) -> None:
|
||||
context, generate_entity, graph_runtime_state = _load_context(request)
|
||||
if not isinstance(generate_entity, AdvancedChatAppGenerateEntity):
|
||||
raise PermanentWorkflowHandoffResumeError("Chatflow handoff contains an incompatible generate entity")
|
||||
if generate_entity.conversation_id is None:
|
||||
raise PermanentWorkflowHandoffResumeError("Chatflow handoff has no conversation identity")
|
||||
|
||||
session_factory = sessionmaker(bind=db.engine, expire_on_commit=False)
|
||||
with session_factory() as session:
|
||||
workflow_run, workflow, user = _load_run_dependencies(session, request)
|
||||
_validate_entity_identity(
|
||||
generate_entity=generate_entity,
|
||||
workflow_run=workflow_run,
|
||||
workflow=workflow,
|
||||
)
|
||||
app_model = session.get(App, workflow_run.app_id)
|
||||
conversation = session.get(Conversation, generate_entity.conversation_id)
|
||||
message = session.scalar(
|
||||
select(Message)
|
||||
.where(
|
||||
Message.conversation_id == generate_entity.conversation_id,
|
||||
Message.workflow_run_id == workflow_run.id,
|
||||
)
|
||||
.order_by(Message.created_at.desc())
|
||||
.limit(1)
|
||||
)
|
||||
if app_model is None or conversation is None or message is None:
|
||||
raise PermanentWorkflowHandoffResumeError("Chatflow records required for resumption no longer exist")
|
||||
_validate_app_ownership(app_model, workflow_run)
|
||||
_validate_chat_records(conversation=conversation, message=message, workflow_run=workflow_run)
|
||||
|
||||
set_login_user(user)
|
||||
workflow_execution_repository, workflow_node_execution_repository = _build_repositories(
|
||||
session_factory=session_factory,
|
||||
workflow_run=workflow_run,
|
||||
user=user,
|
||||
generate_entity=generate_entity,
|
||||
)
|
||||
response = AdvancedChatAppGenerator().resume(
|
||||
app_model=app_model,
|
||||
workflow=workflow,
|
||||
user=user,
|
||||
conversation=conversation,
|
||||
message=message,
|
||||
session=session,
|
||||
application_generate_entity=generate_entity,
|
||||
workflow_execution_repository=workflow_execution_repository,
|
||||
workflow_node_execution_repository=workflow_node_execution_repository,
|
||||
graph_runtime_state=graph_runtime_state,
|
||||
graph_engine_layers=[_acknowledgement_layer(request)],
|
||||
pause_state_config=PauseStateLayerConfig(
|
||||
session_factory=session_factory,
|
||||
state_owner_user_id=workflow.created_by,
|
||||
),
|
||||
response_stream_filter=context.get_response_stream_filter(),
|
||||
handoff_resume_route=request.handoff.resume_route,
|
||||
graph_config=workflow_run.graph_dict,
|
||||
workflow_version=workflow_run.version,
|
||||
root_node_id=context.root_node_id,
|
||||
)
|
||||
|
||||
if isinstance(response, Generator):
|
||||
_publish_streaming_response(
|
||||
response,
|
||||
workflow_run.id,
|
||||
AppMode.ADVANCED_CHAT,
|
||||
workflow.id,
|
||||
generate_entity.inputs,
|
||||
started_reason=_resumption_reason(),
|
||||
terminal_failure_handler=_resumed_terminal_failure_handler(request, workflow_run),
|
||||
)
|
||||
|
||||
|
||||
def _resume_rag_pipeline_handoff(request: WorkflowHandoffResumeRequest) -> None:
|
||||
context, generate_entity, graph_runtime_state = _load_context(request)
|
||||
if not isinstance(generate_entity, RagPipelineGenerateEntity):
|
||||
raise PermanentWorkflowHandoffResumeError("RAG pipeline handoff contains an incompatible generate entity")
|
||||
|
||||
session_factory = sessionmaker(bind=db.engine, expire_on_commit=False)
|
||||
with session_factory() as session:
|
||||
workflow_run, workflow, user = _load_run_dependencies(session, request)
|
||||
_validate_entity_identity(
|
||||
generate_entity=generate_entity,
|
||||
workflow_run=workflow_run,
|
||||
workflow=workflow,
|
||||
)
|
||||
pipeline = session.get(Pipeline, workflow_run.app_id)
|
||||
if pipeline is None or pipeline.tenant_id != workflow_run.tenant_id:
|
||||
raise PermanentWorkflowHandoffResumeError("RAG pipeline no longer exists")
|
||||
dataset = pipeline.retrieve_dataset(session)
|
||||
if dataset is None or dataset.id != generate_entity.dataset_id or dataset.tenant_id != workflow_run.tenant_id:
|
||||
raise PermanentWorkflowHandoffResumeError("RAG pipeline dataset identity does not match the checkpoint")
|
||||
if generate_entity.document_id is not None:
|
||||
document = session.get(Document, generate_entity.document_id)
|
||||
if (
|
||||
document is None
|
||||
or document.dataset_id != generate_entity.dataset_id
|
||||
or document.tenant_id != workflow_run.tenant_id
|
||||
):
|
||||
raise PermanentWorkflowHandoffResumeError(
|
||||
"RAG pipeline document identity does not match the checkpoint"
|
||||
)
|
||||
|
||||
set_login_user(user)
|
||||
g._login_user = user
|
||||
workflow_execution_repository, workflow_node_execution_repository = _build_repositories(
|
||||
session_factory=session_factory,
|
||||
workflow_run=workflow_run,
|
||||
user=user,
|
||||
generate_entity=generate_entity,
|
||||
)
|
||||
response = PipelineGenerator().resume(
|
||||
session=session,
|
||||
pipeline=pipeline,
|
||||
workflow=workflow,
|
||||
user=user,
|
||||
application_generate_entity=generate_entity,
|
||||
graph_runtime_state=graph_runtime_state,
|
||||
workflow_execution_repository=workflow_execution_repository,
|
||||
workflow_node_execution_repository=workflow_node_execution_repository,
|
||||
graph_engine_layers=[_acknowledgement_layer(request)],
|
||||
pause_state_config=PauseStateLayerConfig(
|
||||
session_factory=session_factory,
|
||||
state_owner_user_id=workflow.created_by,
|
||||
),
|
||||
response_stream_filter=context.get_response_stream_filter(),
|
||||
handoff_resume_route=request.handoff.resume_route,
|
||||
graph_config=workflow_run.graph_dict,
|
||||
workflow_version=workflow_run.version,
|
||||
root_node_id=context.root_node_id,
|
||||
)
|
||||
if isinstance(response, Generator):
|
||||
_publish_streaming_response(
|
||||
response,
|
||||
workflow_run.id,
|
||||
AppMode.RAG_PIPELINE,
|
||||
workflow.id,
|
||||
generate_entity.inputs,
|
||||
started_reason=_resumption_reason(),
|
||||
terminal_failure_handler=_resumed_terminal_failure_handler(request, workflow_run),
|
||||
)
|
||||
|
||||
|
||||
def _resumption_reason():
|
||||
# Lazy import keeps this module's public surface focused on Dify route
|
||||
# reconstruction while preserving Graphon's strongly typed event reason.
|
||||
from graphon.entities import WorkflowStartReason
|
||||
|
||||
return WorkflowStartReason.RESUMPTION
|
||||
|
||||
|
||||
__all__ = ["create_workflow_handoff_resume_dispatcher"]
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user