Compare commits

...
245 changed files with 23937 additions and 2131 deletions
+19 -3
View File
@@ -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
+18 -1
View File
@@ -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(
+47
View File
@@ -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
View File
@@ -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
+11 -9
View File
@@ -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
+11 -10
View File
@@ -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"}
+11 -10
View File
@@ -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")
+26 -8
View File
@@ -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"}
+1
View File
@@ -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
+14 -6
View File
@@ -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")
+22 -10
View File
@@ -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
+11 -9
View File
@@ -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."""
+4
View File
@@ -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
+9 -9
View File
@@ -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")
+24 -7
View File
@@ -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]
):
"""
+37 -10
View File
@@ -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:
+16 -6
View File
@@ -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()
+20 -4
View File
@@ -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:
+12 -2
View File
@@ -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()
+81 -9
View File
@@ -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")
+68 -13
View File
@@ -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()
+59 -3
View File
@@ -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)
+34 -12
View File
@@ -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)
+60 -11
View File
@@ -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)
+112 -11
View File
@@ -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(),
+15
View File
@@ -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
+25
View File
@@ -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
+125 -17
View File
@@ -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(
+48 -7
View File
@@ -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 -3
View File
@@ -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
View File
@@ -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)
+4
View File
@@ -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)
+21 -3
View File
@@ -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))
+16
View File
@@ -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
+22 -1
View File
@@ -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)
+214 -20
View File
@@ -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,
+6
View File
@@ -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
+35
View File
@@ -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)
+47 -2
View File
@@ -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.
+27
View File
@@ -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
+2 -1
View File
@@ -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
View File
@@ -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")
+12
View File
@@ -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",
+13
View File
@@ -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,
}
+331
View File
@@ -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",
]
+33
View File
@@ -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 |
+2
View File
@@ -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 |
+22 -2
View File
@@ -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 -4
View File
@@ -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(),
+27 -4
View File
@@ -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:
+8 -1
View File
@@ -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:
+73 -24
View File
@@ -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]:
+354 -29
View 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",
]
+120
View File
@@ -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