From 5a792945f5b29b10f4ba25b9231c40b8322261ef Mon Sep 17 00:00:00 2001 From: zyssyz123 <916125788@qq.com> Date: Fri, 17 Jul 2026 15:46:42 +0800 Subject: [PATCH] fix(agent): bound streams and cancel zombie workflow runs (#39186) --- api/.env.example | 3 + api/clients/agent_backend/client.py | 46 +++++- api/clients/agent_backend/factory.py | 9 +- api/clients/agent_backend/fake_client.py | 12 +- api/configs/extra/agent_backend_config.py | 17 ++- api/core/app/apps/agent_app/app_generator.py | 3 + api/core/app/apps/agent_app/app_runner.py | 86 ++++++----- api/core/app/apps/base_app_queue_manager.py | 35 ++++- .../apps/message_based_app_queue_manager.py | 2 +- .../apps/pipeline/pipeline_queue_manager.py | 2 +- .../app/apps/workflow/app_queue_manager.py | 2 +- api/core/app/apps/workflow_app_runner.py | 8 +- api/core/app/entities/queue_entities.py | 4 + api/core/workflow/node_factory.py | 3 + .../workflow/nodes/agent_v2/agent_node.py | 29 +++- .../agent_backend_session_cleanup_task.py | 3 + .../clients/agent_backend/test_client.py | 35 ++++- .../app/apps/agent_app/test_app_runner.py | 82 ++++++++--- .../app/apps/test_base_app_queue_manager.py | 35 ++++- .../app/apps/test_workflow_app_runner_core.py | 19 +++ .../apps/workflow/test_app_queue_manager.py | 66 ++++++++- .../core/app/entities/test_queue_entities.py | 7 + .../nodes/agent_v2/test_agent_node.py | 138 +++++++++++++++++- dify-agent/src/dify_agent/client/_client.py | 96 ++++++++++-- .../src/dify_agent/runtime/event_sink.py | 21 +++ .../src/dify_agent/runtime/run_scheduler.py | 63 +++++++- dify-agent/src/dify_agent/runtime/runner.py | 14 ++ .../src/dify_agent/server/routes/runs.py | 24 +-- dify-agent/src/dify_agent/server/sse.py | 32 +++- .../local/dify_agent/client/test_client.py | 57 ++++++++ .../dify_agent/runtime/test_run_scheduler.py | 104 ++++++++++++- .../local/dify_agent/runtime/test_runner.py | 39 ++++- .../dify_agent/server/test_runs_routes.py | 34 ++++- .../tests/local/dify_agent/server/test_sse.py | 23 ++- docker/.env.example | 3 + docker/docker-compose-template.yaml | 6 + docker/docker-compose.yaml | 6 + .../envs/core-services/dify-agent.env.example | 3 + 38 files changed, 1042 insertions(+), 129 deletions(-) diff --git a/api/.env.example b/api/.env.example index 804b510ee27..4611bb0d921 100644 --- a/api/.env.example +++ b/api/.env.example @@ -677,6 +677,9 @@ INNER_API_KEY_FOR_PLUGIN=QaHbTe77CtuXmsfyhR7+vRjI/+XbV1AaFy691iy+kGDv2Jvy0/eAh8Y # Dify Agent backend AGENT_BACKEND_BASE_URL=http://localhost:5050 +AGENT_BACKEND_STREAM_READ_TIMEOUT_SECONDS=30 +AGENT_BACKEND_STREAM_MAX_RECONNECTS=3 +AGENT_BACKEND_RUN_TIMEOUT_SECONDS=1200 # Marketplace configuration MARKETPLACE_ENABLED=true diff --git a/api/clients/agent_backend/client.py b/api/clients/agent_backend/client.py index 2b2cfbdf955..270ef5a984c 100644 --- a/api/clients/agent_backend/client.py +++ b/api/clients/agent_backend/client.py @@ -8,7 +8,7 @@ creating another wire contract. from __future__ import annotations -from collections.abc import Iterator +from collections.abc import Callable, Iterator from typing import Protocol from dify_agent.client import ( @@ -45,7 +45,13 @@ class AgentBackendRunClient(Protocol): def cancel_run(self, run_id: str, request: CancelRunRequest | None = None) -> CancelRunResponse: """Request explicit cancellation for one Agent backend run.""" - def stream_events(self, run_id: str, *, after: str | None = None) -> Iterator[RunEvent]: + def stream_events( + self, + run_id: str, + *, + after: str | None = None, + should_stop: Callable[[], bool] | None = None, + ) -> Iterator[RunEvent]: """Yield public ``dify-agent`` run events in stream order.""" def wait_run(self, run_id: str, *, timeout_seconds: float | None = None) -> RunStatusResponse: @@ -61,7 +67,15 @@ class _DifyAgentSyncClient(Protocol): def cancel_run_sync(self, run_id: str, request: CancelRunRequest | None = None) -> CancelRunResponse: """Cancel one run synchronously.""" - def stream_events_sync(self, run_id: str, *, after: str | None = None) -> Iterator[RunEvent]: + def stream_events_sync( + self, + run_id: str, + *, + after: str | None = None, + max_reconnects: int | None = None, + timeout_seconds: float | None = None, + should_stop: Callable[[], bool] | None = None, + ) -> Iterator[RunEvent]: """Stream run events synchronously.""" def wait_run_sync(self, run_id: str, *, timeout_seconds: float | None = None) -> RunStatusResponse: @@ -73,8 +87,16 @@ class DifyAgentBackendRunClient: client: _DifyAgentSyncClient - def __init__(self, client: _DifyAgentSyncClient) -> None: + def __init__( + self, + client: _DifyAgentSyncClient, + *, + stream_max_reconnects: int = 3, + stream_timeout_seconds: float = 1200, + ) -> None: self.client = client + self._stream_max_reconnects = stream_max_reconnects + self._stream_timeout_seconds = stream_timeout_seconds def create_run(self, request: CreateRunRequest) -> CreateRunResponse: """Create one run through ``POST /runs`` and normalize client exceptions.""" @@ -90,10 +112,22 @@ class DifyAgentBackendRunClient: except Exception as exc: raise _normalize_dify_agent_error(exc) from exc - def stream_events(self, run_id: str, *, after: str | None = None) -> Iterator[RunEvent]: + def stream_events( + self, + run_id: str, + *, + after: str | None = None, + should_stop: Callable[[], bool] | None = None, + ) -> Iterator[RunEvent]: """Stream run events from ``/events/sse`` with the wrapped client's reconnect policy.""" try: - yield from self.client.stream_events_sync(run_id, after=after) + yield from self.client.stream_events_sync( + run_id, + after=after, + max_reconnects=self._stream_max_reconnects, + timeout_seconds=self._stream_timeout_seconds, + should_stop=should_stop, + ) except Exception as exc: raise _normalize_dify_agent_error(exc) from exc diff --git a/api/clients/agent_backend/factory.py b/api/clients/agent_backend/factory.py index 133eb42b28d..0fcbf02bf70 100644 --- a/api/clients/agent_backend/factory.py +++ b/api/clients/agent_backend/factory.py @@ -13,10 +13,17 @@ def create_agent_backend_run_client( base_url: str | None = None, use_fake: bool = False, fake_scenario: str | FakeAgentBackendScenario = FakeAgentBackendScenario.SUCCESS, + stream_read_timeout_seconds: float = 30, + stream_max_reconnects: int = 3, + stream_run_timeout_seconds: float = 1200, ) -> AgentBackendRunClient: """Create the API-side run client without hiding the ``dify-agent`` protocol.""" if use_fake: return FakeAgentBackendRunClient(scenario=FakeAgentBackendScenario(fake_scenario)) if base_url is None: raise ValueError("base_url is required when creating a real Agent backend client") - return DifyAgentBackendRunClient(Client(base_url=base_url)) + return DifyAgentBackendRunClient( + Client(base_url=base_url, stream_timeout=stream_read_timeout_seconds), + stream_max_reconnects=stream_max_reconnects, + stream_timeout_seconds=stream_run_timeout_seconds, + ) diff --git a/api/clients/agent_backend/fake_client.py b/api/clients/agent_backend/fake_client.py index 11de90c94b7..2d0881c03ff 100644 --- a/api/clients/agent_backend/fake_client.py +++ b/api/clients/agent_backend/fake_client.py @@ -7,7 +7,7 @@ separate ``agent-backend.v1`` event stream. from __future__ import annotations -from collections.abc import Iterator +from collections.abc import Callable, Iterator from datetime import UTC, datetime from enum import StrEnum @@ -69,9 +69,17 @@ class FakeAgentBackendRunClient: del request return CancelRunResponse(run_id=run_id, status="cancelled") - def stream_events(self, run_id: str, *, after: str | None = None) -> Iterator[RunEvent]: + def stream_events( + self, + run_id: str, + *, + after: str | None = None, + should_stop: Callable[[], bool] | None = None, + ) -> Iterator[RunEvent]: """Yield the deterministic public ``RunEvent`` sequence for ``run_id``.""" for event in self._events(run_id): + if should_stop is not None and should_stop(): + return if after is not None and event.id is not None and event.id <= after: continue yield event diff --git a/api/configs/extra/agent_backend_config.py b/api/configs/extra/agent_backend_config.py index 2e6b9814954..7baad3d0b44 100644 --- a/api/configs/extra/agent_backend_config.py +++ b/api/configs/extra/agent_backend_config.py @@ -1,4 +1,4 @@ -from pydantic import Field, NonNegativeFloat +from pydantic import Field, NonNegativeFloat, NonNegativeInt, PositiveFloat from pydantic_settings import BaseSettings @@ -22,6 +22,21 @@ class AgentBackendConfig(BaseSettings): default="success", ) + AGENT_BACKEND_STREAM_READ_TIMEOUT_SECONDS: PositiveFloat = Field( + description="Read timeout for one Agent backend SSE connection.", + default=30, + ) + + AGENT_BACKEND_STREAM_MAX_RECONNECTS: NonNegativeInt = Field( + description="Maximum Agent backend SSE reconnects before failing the run.", + default=3, + ) + + AGENT_BACKEND_RUN_TIMEOUT_SECONDS: PositiveFloat = Field( + description="Total deadline for one Agent backend run event stream.", + default=1200, + ) + AGENT_SHELL_ENABLED: bool = Field( description=( "Inject the dify.shell layer (sandboxed bash workspace) into Agent runs. " diff --git a/api/core/app/apps/agent_app/app_generator.py b/api/core/app/apps/agent_app/app_generator.py index 68f6edbd550..73a663f92b2 100644 --- a/api/core/app/apps/agent_app/app_generator.py +++ b/api/core/app/apps/agent_app/app_generator.py @@ -540,6 +540,9 @@ class AgentAppGenerator(MessageBasedAppGenerator): base_url=dify_config.AGENT_BACKEND_BASE_URL, use_fake=dify_config.AGENT_BACKEND_USE_FAKE, fake_scenario=dify_config.AGENT_BACKEND_FAKE_SCENARIO, + stream_read_timeout_seconds=dify_config.AGENT_BACKEND_STREAM_READ_TIMEOUT_SECONDS, + stream_max_reconnects=dify_config.AGENT_BACKEND_STREAM_MAX_RECONNECTS, + stream_run_timeout_seconds=dify_config.AGENT_BACKEND_RUN_TIMEOUT_SECONDS, ), event_adapter=AgentBackendRunEventAdapter(), session_store=AgentAppRuntimeSessionStore(), diff --git a/api/core/app/apps/agent_app/app_runner.py b/api/core/app/apps/agent_app/app_runner.py index 3ea2919d631..f844068073d 100644 --- a/api/core/app/apps/agent_app/app_runner.py +++ b/api/core/app/apps/agent_app/app_runner.py @@ -941,48 +941,64 @@ class AgentAppRunner: if pending_text: persist_answer_text(pending_text) - for public_event in self._agent_backend_client.stream_events(run_id): - if queue_manager.is_stopped(): - flush_pending_agent_message_text() - self._cancel_run(run_id) - raise GenerateTaskStoppedError() - for internal_event in self._event_adapter.adapt(public_event): + try: + public_events = self._agent_backend_client.stream_events( + run_id, + should_stop=queue_manager.is_stopped, + ) + for public_event in public_events: if queue_manager.is_stopped(): flush_pending_agent_message_text() self._cancel_run(run_id) raise GenerateTaskStoppedError() - if internal_event.type in ( - AgentBackendInternalEventType.RUN_STARTED, - AgentBackendInternalEventType.STREAM_EVENT, - AgentBackendInternalEventType.AGENT_MESSAGE_DELTA, - ): - if isinstance(internal_event, AgentBackendAgentMessageDeltaInternalEvent): - debounced_delta = text_delta_debouncer.push(internal_event.delta) - if debounced_delta: - persist_answer_text(debounced_delta) - continue - - if isinstance(internal_event, AgentBackendStreamInternalEvent): + for internal_event in self._event_adapter.adapt(public_event): + if queue_manager.is_stopped(): flush_pending_agent_message_text() - try: - process_recorder.handle_stream_event(internal_event) - except Exception: - db.session.rollback() - logger.warning( - "Failed to persist Agent App process event: run_id=%s message_id=%s event_kind=%s", - run_id, - message_id, - internal_event.event_kind, - exc_info=True, - ) + self._cancel_run(run_id) + raise GenerateTaskStoppedError() + if internal_event.type in ( + AgentBackendInternalEventType.RUN_STARTED, + AgentBackendInternalEventType.STREAM_EVENT, + AgentBackendInternalEventType.AGENT_MESSAGE_DELTA, + ): + if isinstance(internal_event, AgentBackendAgentMessageDeltaInternalEvent): + debounced_delta = text_delta_debouncer.push(internal_event.delta) + if debounced_delta: + persist_answer_text(debounced_delta) + continue + + if isinstance(internal_event, AgentBackendStreamInternalEvent): + flush_pending_agent_message_text() + try: + process_recorder.handle_stream_event(internal_event) + except Exception: + db.session.rollback() + logger.warning( + "Failed to persist Agent App process event: run_id=%s message_id=%s event_kind=%s", + run_id, + message_id, + internal_event.event_kind, + exc_info=True, + ) + continue continue - continue - flush_pending_agent_message_text() - terminal = internal_event - break - if terminal is not None: - break + flush_pending_agent_message_text() + terminal = internal_event + break + if terminal is not None: + break + except GenerateTaskStoppedError: + raise + except Exception as error: + flush_pending_agent_message_text() + self._cancel_run(run_id) + if queue_manager.is_stopped(): + raise GenerateTaskStoppedError() from error + raise flush_pending_agent_message_text() + if queue_manager.is_stopped(): + self._cancel_run(run_id) + raise GenerateTaskStoppedError() return terminal, process_recorder def _cancel_run(self, run_id: str) -> None: diff --git a/api/core/app/apps/base_app_queue_manager.py b/api/core/app/apps/base_app_queue_manager.py index 9551a5e38c4..365d66ad0e3 100644 --- a/api/core/app/apps/base_app_queue_manager.py +++ b/api/core/app/apps/base_app_queue_manager.py @@ -21,6 +21,7 @@ from core.app.entities.queue_entities import ( WorkflowQueueMessage, ) from extensions.ext_redis import redis_client +from graphon.graph_engine.manager import GraphEngineManager from graphon.runtime import GraphRuntimeState logger = logging.getLogger(__name__) @@ -51,6 +52,9 @@ class AppQueueManager(ABC): self._graph_runtime_state: GraphRuntimeState | None = None self._stopped_cache: TTLCache[tuple, bool] = TTLCache(maxsize=1, ttl=1) self._cache_lock = threading.Lock() + self._execution_terminal = threading.Event() + self._abort_sent = threading.Event() + self._lifecycle_lock = threading.Lock() def listen(self): """ @@ -59,7 +63,7 @@ class AppQueueManager(ABC): """ # wait for APP_MAX_EXECUTION_TIME seconds to stop listen listen_timeout = dify_config.APP_MAX_EXECUTION_TIME - start_time = time.time() + start_time = time.monotonic() last_ping_time: int | float = 0 try: while True: @@ -72,8 +76,14 @@ class AppQueueManager(ABC): except queue.Empty: continue finally: - elapsed_time = time.time() - start_time - if elapsed_time >= listen_timeout or self._is_stopped(): + elapsed_time = time.monotonic() - start_time + timed_out = elapsed_time >= listen_timeout + manually_stopped = self._is_stopped() + if not self._execution_terminal.is_set() and (timed_out or manually_stopped): + reason = ( + f"App execution exceeded {listen_timeout} seconds" if timed_out else "App task was stopped" + ) + self._abort_execution(reason) # publish two messages to make sure the client can receive the stop signal # and stop listening after the stop signal processed self.publish( @@ -84,16 +94,33 @@ class AppQueueManager(ABC): self.publish(QueuePingEvent(), PublishFrom.TASK_PIPELINE) last_ping_time = elapsed_time // 10 finally: + if not self._execution_terminal.is_set(): + self._abort_execution("Client response stream closed before app execution completed") self._graph_runtime_state = None # Release reference once consumers finish or close the generator. - def stop_listen(self): + def stop_listen(self, *, execution_terminal: bool = False): """ Stop listen to queue :return: """ + if execution_terminal: + self._execution_terminal.set() self._clear_task_belong_cache() self._q.put(None) + def _abort_execution(self, reason: str) -> None: + """Propagate response timeout/disconnect to legacy and GraphEngine runners.""" + with self._lifecycle_lock: + if self._execution_terminal.is_set() or self._abort_sent.is_set(): + return + self._abort_sent.set() + + try: + self.set_stop_flag_no_user_check(self._task_id) + 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) + def _clear_task_belong_cache(self) -> None: """ Remove the task belong cache key once listening is finished. diff --git a/api/core/app/apps/message_based_app_queue_manager.py b/api/core/app/apps/message_based_app_queue_manager.py index 3c7102971f1..b253d93ee52 100644 --- a/api/core/app/apps/message_based_app_queue_manager.py +++ b/api/core/app/apps/message_based_app_queue_manager.py @@ -45,7 +45,7 @@ class MessageBasedAppQueueManager(AppQueueManager): if isinstance( event, QueueStopEvent | QueueErrorEvent | QueueMessageEndEvent | QueueAdvancedChatMessageEndEvent ): - self.stop_listen() + self.stop_listen(execution_terminal=True) if pub_from == PublishFrom.APPLICATION_MANAGER and self._is_stopped(): if self._app_mode == AppMode.ADVANCED_CHAT.value: diff --git a/api/core/app/apps/pipeline/pipeline_queue_manager.py b/api/core/app/apps/pipeline/pipeline_queue_manager.py index c34b51c98c7..d424266454e 100644 --- a/api/core/app/apps/pipeline/pipeline_queue_manager.py +++ b/api/core/app/apps/pipeline/pipeline_queue_manager.py @@ -42,7 +42,7 @@ class PipelineQueueManager(AppQueueManager): | QueueWorkflowFailedEvent | QueueWorkflowPartialSuccessEvent, ): - self.stop_listen() + self.stop_listen(execution_terminal=True) if pub_from == PublishFrom.APPLICATION_MANAGER and self._is_stopped(): raise GenerateTaskStoppedError() diff --git a/api/core/app/apps/workflow/app_queue_manager.py b/api/core/app/apps/workflow/app_queue_manager.py index 7824d33b875..67df3044fb2 100644 --- a/api/core/app/apps/workflow/app_queue_manager.py +++ b/api/core/app/apps/workflow/app_queue_manager.py @@ -41,4 +41,4 @@ class WorkflowAppQueueManager(AppQueueManager): | QueueWorkflowFailedEvent | QueueWorkflowPartialSuccessEvent, ): - self.stop_listen() + self.stop_listen(execution_terminal=True) diff --git a/api/core/app/apps/workflow_app_runner.py b/api/core/app/apps/workflow_app_runner.py index 01dcaef10db..3d2857f130a 100644 --- a/api/core/app/apps/workflow_app_runner.py +++ b/api/core/app/apps/workflow_app_runner.py @@ -26,6 +26,7 @@ from core.app.entities.queue_entities import ( QueueNodeSucceededEvent, QueueReasoningChunkEvent, QueueRetrieverResourcesEvent, + QueueStopEvent, QueueTextChunkEvent, QueueWorkflowFailedEvent, QueueWorkflowPartialSuccessEvent, @@ -424,7 +425,12 @@ class WorkflowBasedAppRunner: QueueWorkflowFailedEvent(error=event.error, exceptions_count=event.exceptions_count) ) case GraphRunAbortedEvent(): - self._publish_event(QueueWorkflowFailedEvent(error=event.reason or "Unknown error", exceptions_count=0)) + self._publish_event( + QueueStopEvent( + stopped_by=QueueStopEvent.StopBy.USER_MANUAL, + reason=event.reason or "Workflow execution aborted", + ) + ) case GraphRunPausedEvent(): runtime_state = workflow_entry.graph_engine.graph_runtime_state paused_nodes = runtime_state.get_paused_nodes() diff --git a/api/core/app/entities/queue_entities.py b/api/core/app/entities/queue_entities.py index 6bbffbbb160..88a3cf6e550 100644 --- a/api/core/app/entities/queue_entities.py +++ b/api/core/app/entities/queue_entities.py @@ -500,11 +500,15 @@ class QueueStopEvent(AppQueueEvent): event: QueueEvent = QueueEvent.STOP stopped_by: StopBy + reason: str | None = None def get_stop_reason(self) -> str: """ To stop reason """ + if self.reason: + return self.reason + reason_mapping = { QueueStopEvent.StopBy.USER_MANUAL: "Stopped by user.", QueueStopEvent.StopBy.ANNOTATION_REPLY: "Stopped by annotation reply.", diff --git a/api/core/workflow/node_factory.py b/api/core/workflow/node_factory.py index 5b8db068d0d..3b47e32adf1 100644 --- a/api/core/workflow/node_factory.py +++ b/api/core/workflow/node_factory.py @@ -499,6 +499,9 @@ class DifyNodeFactory(NodeFactory): base_url=dify_config.AGENT_BACKEND_BASE_URL, use_fake=dify_config.AGENT_BACKEND_USE_FAKE, fake_scenario=dify_config.AGENT_BACKEND_FAKE_SCENARIO, + stream_read_timeout_seconds=dify_config.AGENT_BACKEND_STREAM_READ_TIMEOUT_SECONDS, + stream_max_reconnects=dify_config.AGENT_BACKEND_STREAM_MAX_RECONNECTS, + stream_run_timeout_seconds=dify_config.AGENT_BACKEND_RUN_TIMEOUT_SECONDS, ), "event_adapter": AgentBackendRunEventAdapter(), # Agent Files ยง4.6: reback file outputs from the ToolFile row so diff --git a/api/core/workflow/nodes/agent_v2/agent_node.py b/api/core/workflow/nodes/agent_v2/agent_node.py index 6f27905f9a6..e100b7414b3 100644 --- a/api/core/workflow/nodes/agent_v2/agent_node.py +++ b/api/core/workflow/nodes/agent_v2/agent_node.py @@ -5,6 +5,7 @@ from collections.abc import Generator, Mapping, Sequence from typing import TYPE_CHECKING, Any, override from agenton.compositor import CompositorSessionSnapshot +from dify_agent.protocol import CancelRunRequest from clients.agent_backend import ( AgentBackendAgentMessageDeltaInternalEvent, @@ -473,7 +474,10 @@ class DifyAgentNode(Node[DifyAgentNodeData]): """ stream_event_count = 0 try: - for public_event in self._agent_backend_client.stream_events(run_id): + for public_event in self._agent_backend_client.stream_events( + run_id, + should_stop=self._is_graph_aborted, + ): stream_event_count += 1 for internal_event in self._event_adapter.adapt(public_event): if internal_event.type == AgentBackendInternalEventType.RUN_STARTED: @@ -501,6 +505,7 @@ class DifyAgentNode(Node[DifyAgentNodeData]): | AgentBackendDeferredToolCallInternalEvent, ): return internal_event, None + self._cancel_backend_run(run_id, reason="unexpected_event") return None, self._failure_event( inputs={}, process_data={}, @@ -509,6 +514,7 @@ class DifyAgentNode(Node[DifyAgentNodeData]): error_type="agent_backend_stream_error", ) except AgentBackendError as error: + self._cancel_backend_run(run_id, reason=self._stream_stop_reason()) return None, self._failure_event( inputs={}, process_data={}, @@ -517,6 +523,7 @@ class DifyAgentNode(Node[DifyAgentNodeData]): error_type=self._agent_backend_error_type(error), ) except Exception as error: + self._cancel_backend_run(run_id, reason=self._stream_stop_reason()) return None, self._failure_event( inputs={}, process_data={}, @@ -525,8 +532,28 @@ class DifyAgentNode(Node[DifyAgentNodeData]): error_type="agent_backend_stream_error", ) + self._cancel_backend_run(run_id, reason="stream_ended_without_terminal_event") return None, None + def _is_graph_aborted(self) -> bool: + """Let Agent SSE consumption observe GraphEngine's cooperative abort state.""" + try: + return self.graph_runtime_state.graph_execution.aborted + except (AttributeError, RuntimeError): + return False + + def _stream_stop_reason(self) -> str: + return "workflow_graph_aborted" if self._is_graph_aborted() else "event_stream_failed" + + def _cancel_backend_run(self, run_id: str, *, reason: str) -> None: + try: + self._agent_backend_client.cancel_run( + run_id, + CancelRunRequest(reason=reason, message="Workflow Agent event consumption stopped"), + ) + except Exception: + logger.warning("Failed to cancel Workflow Agent backend run: run_id=%s", run_id, exc_info=True) + @staticmethod def _record_type_check_metadata(metadata: dict[str, Any], outcome: OutputTypeCheckOutcome) -> None: # Surface enough detail in metadata for Inspector / debug logs without diff --git a/api/tasks/agent_backend_session_cleanup_task.py b/api/tasks/agent_backend_session_cleanup_task.py index 2c799bda97c..f1316266db7 100644 --- a/api/tasks/agent_backend_session_cleanup_task.py +++ b/api/tasks/agent_backend_session_cleanup_task.py @@ -24,6 +24,9 @@ def _create_agent_backend_client(): base_url=dify_config.AGENT_BACKEND_BASE_URL, use_fake=dify_config.AGENT_BACKEND_USE_FAKE, fake_scenario=dify_config.AGENT_BACKEND_FAKE_SCENARIO, + stream_read_timeout_seconds=dify_config.AGENT_BACKEND_STREAM_READ_TIMEOUT_SECONDS, + stream_max_reconnects=dify_config.AGENT_BACKEND_STREAM_MAX_RECONNECTS, + stream_run_timeout_seconds=dify_config.AGENT_BACKEND_RUN_TIMEOUT_SECONDS, ) diff --git a/api/tests/unit_tests/clients/agent_backend/test_client.py b/api/tests/unit_tests/clients/agent_backend/test_client.py index 6b61c920a7b..421ffad9092 100644 --- a/api/tests/unit_tests/clients/agent_backend/test_client.py +++ b/api/tests/unit_tests/clients/agent_backend/test_client.py @@ -1,4 +1,4 @@ -from collections.abc import Iterator +from collections.abc import Callable, Iterator from typing import override import pytest @@ -47,6 +47,8 @@ def _request(): class _SuccessfulClient: + stream_options: tuple[int | None, float | None, Callable[[], bool] | None] | None = None + def create_run_sync(self, request: CreateRunRequest) -> CreateRunResponse: assert isinstance(request, CreateRunRequest) return CreateRunResponse(run_id="run-1", status="running") @@ -55,8 +57,17 @@ class _SuccessfulClient: del request return CancelRunResponse(run_id=run_id, status="cancelled") - def stream_events_sync(self, run_id: str, *, after: str | None = None) -> Iterator[RunEvent]: + def stream_events_sync( + self, + run_id: str, + *, + after: str | None = None, + max_reconnects: int | None = None, + timeout_seconds: float | None = None, + should_stop: Callable[[], bool] | None = None, + ) -> Iterator[RunEvent]: del after + self.stream_options = (max_reconnects, timeout_seconds, should_stop) yield RunStartedEvent(id="1-0", run_id=run_id) def wait_run_sync(self, run_id: str, *, timeout_seconds: float | None = None) -> RunStatusResponse: @@ -72,17 +83,22 @@ class _SuccessfulClient: def test_dify_agent_backend_run_client_delegates_sync_methods(): - client = DifyAgentBackendRunClient(_SuccessfulClient()) + wrapped = _SuccessfulClient() + client = DifyAgentBackendRunClient(wrapped, stream_max_reconnects=2, stream_timeout_seconds=45) + + def should_stop() -> bool: + return False created = client.create_run(_request()) cancelled = client.cancel_run(created.run_id) - events = list(client.stream_events(created.run_id)) + events = list(client.stream_events(created.run_id, should_stop=should_stop)) status = client.wait_run(created.run_id) assert created.run_id == "run-1" assert cancelled.status == "cancelled" assert events[0].type == "run_started" assert status.status == "succeeded" + assert wrapped.stream_options == (2, 45, should_stop) def test_dify_agent_backend_run_client_maps_validation_error(): @@ -125,7 +141,16 @@ def test_dify_agent_backend_run_client_maps_timeout_error(): def test_dify_agent_backend_run_client_maps_stream_error(): class StreamClient(_SuccessfulClient): @override - def stream_events_sync(self, run_id: str, *, after: str | None = None) -> Iterator[RunEvent]: + def stream_events_sync( + self, + run_id: str, + *, + after: str | None = None, + max_reconnects: int | None = None, + timeout_seconds: float | None = None, + should_stop: Callable[[], bool] | None = None, + ) -> Iterator[RunEvent]: + del run_id, after, max_reconnects, timeout_seconds, should_stop raise DifyAgentStreamError("bad stream") yield diff --git a/api/tests/unit_tests/core/app/apps/agent_app/test_app_runner.py b/api/tests/unit_tests/core/app/apps/agent_app/test_app_runner.py index 09038920230..6321c8ea48b 100644 --- a/api/tests/unit_tests/core/app/apps/agent_app/test_app_runner.py +++ b/api/tests/unit_tests/core/app/apps/agent_app/test_app_runner.py @@ -4,7 +4,7 @@ saved, using the deterministic fake backend client (no live stack).""" from __future__ import annotations -from collections.abc import Iterator +from collections.abc import Callable, Iterator from datetime import UTC, datetime from types import SimpleNamespace from typing import Any, override @@ -106,8 +106,14 @@ class _RecordingFakeAgentBackendRunClient(FakeAgentBackendRunClient): class _StreamingFakeAgentBackendRunClient(FakeAgentBackendRunClient): @override - def stream_events(self, run_id: str, *, after: str | None = None) -> Iterator[RunEvent]: - del after + def stream_events( + self, + run_id: str, + *, + after: str | None = None, + should_stop: Callable[[], bool] | None = None, + ) -> Iterator[RunEvent]: + del after, should_stop created_at = datetime(2026, 1, 1, tzinfo=UTC) yield RunStartedEvent(id="1-0", run_id=run_id, created_at=created_at) yield PydanticAIStreamRunEvent( @@ -138,8 +144,14 @@ class _StreamingFakeAgentBackendRunClient(FakeAgentBackendRunClient): class _StreamingRecordingFakeAgentBackendRunClient(_RecordingFakeAgentBackendRunClient): @override - def stream_events(self, run_id: str, *, after: str | None = None) -> Iterator[RunEvent]: - del after + def stream_events( + self, + run_id: str, + *, + after: str | None = None, + should_stop: Callable[[], bool] | None = None, + ) -> Iterator[RunEvent]: + del after, should_stop created_at = datetime(2026, 1, 1, tzinfo=UTC) yield RunStartedEvent(id="1-0", run_id=run_id, created_at=created_at) yield PydanticAIStreamRunEvent( @@ -173,8 +185,14 @@ class _StreamingStopAfterFirstDeltaFakeAgentBackendRunClient(_RecordingFakeAgent self._queue_manager = queue_manager @override - def stream_events(self, run_id: str, *, after: str | None = None) -> Iterator[RunEvent]: - del after + def stream_events( + self, + run_id: str, + *, + after: str | None = None, + should_stop: Callable[[], bool] | None = None, + ) -> Iterator[RunEvent]: + del after, should_stop created_at = datetime(2026, 1, 1, tzinfo=UTC) yield RunStartedEvent(id="1-0", run_id=run_id, created_at=created_at) yield PydanticAIStreamRunEvent( @@ -196,8 +214,14 @@ class _StreamingStopAfterFirstDeltaFakeAgentBackendRunClient(_RecordingFakeAgent class _StreamingSingleAgentMessageDeltaFakeAgentBackendRunClient(FakeAgentBackendRunClient): @override - def stream_events(self, run_id: str, *, after: str | None = None) -> Iterator[RunEvent]: - del after + def stream_events( + self, + run_id: str, + *, + after: str | None = None, + should_stop: Callable[[], bool] | None = None, + ) -> Iterator[RunEvent]: + del after, should_stop created_at = datetime(2026, 1, 1, tzinfo=UTC) yield RunStartedEvent(id="1-0", run_id=run_id, created_at=created_at) yield PydanticAIStreamRunEvent( @@ -220,8 +244,14 @@ class _StreamingSingleAgentMessageDeltaFakeAgentBackendRunClient(FakeAgentBacken class _NullOutputFakeAgentBackendRunClient(FakeAgentBackendRunClient): @override - def stream_events(self, run_id: str, *, after: str | None = None) -> Iterator[RunEvent]: - del after + def stream_events( + self, + run_id: str, + *, + after: str | None = None, + should_stop: Callable[[], bool] | None = None, + ) -> Iterator[RunEvent]: + del after, should_stop created_at = datetime(2026, 1, 1, tzinfo=UTC) yield RunStartedEvent(id="1-0", run_id=run_id, created_at=created_at) yield RunSucceededEvent( @@ -237,8 +267,14 @@ class _NullOutputFakeAgentBackendRunClient(FakeAgentBackendRunClient): class _StreamingTextNullOutputFakeAgentBackendRunClient(FakeAgentBackendRunClient): @override - def stream_events(self, run_id: str, *, after: str | None = None) -> Iterator[RunEvent]: - del after + def stream_events( + self, + run_id: str, + *, + after: str | None = None, + should_stop: Callable[[], bool] | None = None, + ) -> Iterator[RunEvent]: + del after, should_stop created_at = datetime(2026, 1, 1, tzinfo=UTC) yield RunStartedEvent(id="1-0", run_id=run_id, created_at=created_at) yield PydanticAIStreamRunEvent( @@ -261,8 +297,14 @@ class _StreamingTextNullOutputFakeAgentBackendRunClient(FakeAgentBackendRunClien class _AgentAnswerStreamingFakeAgentBackendRunClient(FakeAgentBackendRunClient): @override - def stream_events(self, run_id: str, *, after: str | None = None) -> Iterator[RunEvent]: - del after + def stream_events( + self, + run_id: str, + *, + after: str | None = None, + should_stop: Callable[[], bool] | None = None, + ) -> Iterator[RunEvent]: + del after, should_stop created_at = datetime(2026, 1, 1, tzinfo=UTC) yield RunStartedEvent(id="1-0", run_id=run_id, created_at=created_at) yield PydanticAIStreamRunEvent( @@ -292,8 +334,14 @@ class _AgentAnswerStreamingFakeAgentBackendRunClient(FakeAgentBackendRunClient): class _ProcessStreamingFakeAgentBackendRunClient(FakeAgentBackendRunClient): @override - def stream_events(self, run_id: str, *, after: str | None = None) -> Iterator[RunEvent]: - del after + def stream_events( + self, + run_id: str, + *, + after: str | None = None, + should_stop: Callable[[], bool] | None = None, + ) -> Iterator[RunEvent]: + del after, should_stop created_at = datetime(2026, 1, 1, tzinfo=UTC) yield RunStartedEvent(id="1-0", run_id=run_id, created_at=created_at) yield PydanticAIStreamRunEvent( diff --git a/api/tests/unit_tests/core/app/apps/test_base_app_queue_manager.py b/api/tests/unit_tests/core/app/apps/test_base_app_queue_manager.py index 842d14bbd25..3a5591613b1 100644 --- a/api/tests/unit_tests/core/app/apps/test_base_app_queue_manager.py +++ b/api/tests/unit_tests/core/app/apps/test_base_app_queue_manager.py @@ -61,16 +61,37 @@ class TestBaseAppQueueManager: manager._check_for_sqlalchemy_models(bad) def test_stop_listen_defers_graph_runtime_state_cleanup_until_listener_exits(self): - with patch("core.app.apps.base_app_queue_manager.redis_client") as mock_redis: + with ( + patch("core.app.apps.base_app_queue_manager.redis_client") as mock_redis, + patch("core.app.apps.base_app_queue_manager.GraphEngineManager") as graph_engine_manager, + ): mock_redis.setex.return_value = True mock_redis.get.return_value = None manager = DummyQueueManager(task_id="t1", user_id="u1", invoke_from=InvokeFrom.SERVICE_API) + runtime_state = SimpleNamespace(name="runtime-state") + manager.graph_runtime_state = runtime_state - runtime_state = SimpleNamespace(name="runtime-state") - manager.graph_runtime_state = runtime_state + manager.stop_listen() - manager.stop_listen() + assert manager.graph_runtime_state is runtime_state + assert list(manager.listen()) == [] + assert manager.graph_runtime_state is None + graph_engine_manager.return_value.send_stop_command.assert_called_once_with( + "t1", + reason="Client response stream closed before app execution completed", + ) - assert manager.graph_runtime_state is runtime_state - assert list(manager.listen()) == [] - assert manager.graph_runtime_state is None + def test_abort_execution_is_idempotent_when_graph_stop_command_fails(self, caplog): + with ( + patch("core.app.apps.base_app_queue_manager.redis_client") as mock_redis, + patch("core.app.apps.base_app_queue_manager.GraphEngineManager") as graph_engine_manager, + ): + mock_redis.setex.return_value = True + graph_engine_manager.return_value.send_stop_command.side_effect = RuntimeError("redis unavailable") + manager = DummyQueueManager(task_id="t1", user_id="u1", invoke_from=InvokeFrom.SERVICE_API) + + manager._abort_execution("stream closed") + manager._abort_execution("duplicate") + + graph_engine_manager.return_value.send_stop_command.assert_called_once_with("t1", reason="stream closed") + assert "Failed to abort app execution for task t1" in caplog.text diff --git a/api/tests/unit_tests/core/app/apps/test_workflow_app_runner_core.py b/api/tests/unit_tests/core/app/apps/test_workflow_app_runner_core.py index 0c297d50a01..fd643893f69 100644 --- a/api/tests/unit_tests/core/app/apps/test_workflow_app_runner_core.py +++ b/api/tests/unit_tests/core/app/apps/test_workflow_app_runner_core.py @@ -17,6 +17,7 @@ from core.app.entities.queue_entities import ( QueueNodeRetryEvent, QueueNodeSucceededEvent, QueueReasoningChunkEvent, + QueueStopEvent, QueueTextChunkEvent, QueueWorkflowPausedEvent, QueueWorkflowStartedEvent, @@ -27,6 +28,7 @@ from core.workflow.system_variables import default_system_variables from graphon.entities.pause_reason import HitlRequired from graphon.enums import BuiltinNodeTypes from graphon.graph_events import ( + GraphRunAbortedEvent, GraphRunPausedEvent, GraphRunStartedEvent, GraphRunSucceededEvent, @@ -373,6 +375,23 @@ class TestWorkflowBasedAppRunner: assert paused_event.paused_nodes == ["node-1"] assert emails + def test_handle_graph_aborted_publishes_stopped_terminal(self): + published: list[object] = [] + + class _QueueManager: + def publish(self, event, publish_from): + del publish_from + published.append(event) + + runner = WorkflowBasedAppRunner(queue_manager=_QueueManager(), app_id="app") + workflow_entry = SimpleNamespace() + + runner._handle_event(workflow_entry, GraphRunAbortedEvent(reason="User requested stop", outputs={})) + + event = published[-1] + assert isinstance(event, QueueStopEvent) + assert event.get_stop_reason() == "User requested stop" + def test_handle_node_events_publishes_queue_events(self): published: list[object] = [] diff --git a/api/tests/unit_tests/core/app/apps/workflow/test_app_queue_manager.py b/api/tests/unit_tests/core/app/apps/workflow/test_app_queue_manager.py index e3b86530098..5e9a69738c4 100644 --- a/api/tests/unit_tests/core/app/apps/workflow/test_app_queue_manager.py +++ b/api/tests/unit_tests/core/app/apps/workflow/test_app_queue_manager.py @@ -5,7 +5,7 @@ from unittest.mock import patch from core.app.apps.base_app_queue_manager import PublishFrom from core.app.apps.workflow.app_queue_manager import WorkflowAppQueueManager from core.app.entities.app_invoke_entities import InvokeFrom -from core.app.entities.queue_entities import QueueMessageEndEvent, QueuePingEvent +from core.app.entities.queue_entities import QueueMessageEndEvent, QueuePingEvent, QueueStopEvent class TestWorkflowAppQueueManager: @@ -35,3 +35,67 @@ class TestWorkflowAppQueueManager: ) manager._publish(QueuePingEvent(), PublishFrom.TASK_PIPELINE) + + def test_listener_close_aborts_unfinished_execution(self): + with ( + patch("core.app.apps.base_app_queue_manager.redis_client") as redis_client, + patch("core.app.apps.base_app_queue_manager.GraphEngineManager") as graph_engine_manager, + ): + redis_client.get.return_value = None + manager = WorkflowAppQueueManager( + task_id="task", + user_id="user", + invoke_from=InvokeFrom.DEBUGGER, + app_mode="workflow", + ) + manager.publish(QueuePingEvent(), PublishFrom.TASK_PIPELINE) + listener = manager.listen() + + assert isinstance(next(listener).event, QueuePingEvent) + listener.close() + + graph_engine_manager.return_value.send_stop_command.assert_called_once_with( + "task", + reason="Client response stream closed before app execution completed", + ) + + def test_execution_timeout_aborts_graph_before_stop_event(self): + with ( + patch("core.app.apps.base_app_queue_manager.redis_client") as redis_client, + patch("core.app.apps.base_app_queue_manager.GraphEngineManager") as graph_engine_manager, + patch("core.app.apps.base_app_queue_manager.dify_config.APP_MAX_EXECUTION_TIME", 0), + ): + redis_client.get.return_value = None + manager = WorkflowAppQueueManager( + task_id="task", + user_id="user", + invoke_from=InvokeFrom.DEBUGGER, + app_mode="workflow", + ) + manager.publish(QueuePingEvent(), PublishFrom.TASK_PIPELINE) + + messages = list(manager.listen()) + + assert any(isinstance(message.event, QueueStopEvent) for message in messages) + graph_engine_manager.return_value.send_stop_command.assert_called_once_with( + "task", + reason="App execution exceeded 0 seconds", + ) + + def test_terminal_event_does_not_abort_completed_execution(self): + with ( + patch("core.app.apps.base_app_queue_manager.redis_client") as redis_client, + patch("core.app.apps.base_app_queue_manager.GraphEngineManager") as graph_engine_manager, + ): + redis_client.get.return_value = None + manager = WorkflowAppQueueManager( + task_id="task", + user_id="user", + invoke_from=InvokeFrom.DEBUGGER, + app_mode="workflow", + ) + manager.publish(QueueMessageEndEvent(llm_result=None), PublishFrom.APPLICATION_MANAGER) + + _ = list(manager.listen()) + + graph_engine_manager.return_value.send_stop_command.assert_not_called() diff --git a/api/tests/unit_tests/core/app/entities/test_queue_entities.py b/api/tests/unit_tests/core/app/entities/test_queue_entities.py index a930d7325d4..43d69023e70 100644 --- a/api/tests/unit_tests/core/app/entities/test_queue_entities.py +++ b/api/tests/unit_tests/core/app/entities/test_queue_entities.py @@ -6,6 +6,13 @@ class TestQueueEntities: event = QueueStopEvent(stopped_by=QueueStopEvent.StopBy.USER_MANUAL) assert event.get_stop_reason() == "Stopped by user." + def test_get_stop_reason_prefers_explicit_reason(self): + event = QueueStopEvent( + stopped_by=QueueStopEvent.StopBy.USER_MANUAL, + reason="Workflow execution timed out", + ) + assert event.get_stop_reason() == "Workflow execution timed out" + def test_get_stop_reason_for_unknown_stop_by(self): event = QueueStopEvent(stopped_by=QueueStopEvent.StopBy.USER_MANUAL) event.stopped_by = "unknown" diff --git a/api/tests/unit_tests/core/workflow/nodes/agent_v2/test_agent_node.py b/api/tests/unit_tests/core/workflow/nodes/agent_v2/test_agent_node.py index b642ea7780d..d142591c3fb 100644 --- a/api/tests/unit_tests/core/workflow/nodes/agent_v2/test_agent_node.py +++ b/api/tests/unit_tests/core/workflow/nodes/agent_v2/test_agent_node.py @@ -1,3 +1,4 @@ +from collections.abc import Callable, Iterator from datetime import UTC, datetime from types import SimpleNamespace from typing import cast @@ -5,11 +6,21 @@ from unittest.mock import MagicMock, patch from agenton.compositor import CompositorSessionSnapshot from dify_agent.layers.ask_human import AskHumanToolResult -from dify_agent.protocol import PydanticAIStreamRunEvent, RunStartedEvent, RunSucceededEvent, RunSucceededEventData +from dify_agent.protocol import ( + CancelRunRequest, + CancelRunResponse, + PydanticAIStreamRunEvent, + RunEvent, + RunStartedEvent, + RunSucceededEvent, + RunSucceededEventData, +) from pydantic_ai.messages import PartDeltaEvent, TextPartDelta from clients.agent_backend import ( + AgentBackendInternalEventType, AgentBackendRunEventAdapter, + AgentBackendStreamError, AgentBackendStreamInternalEvent, FakeAgentBackendRunClient, FakeAgentBackendScenario, @@ -216,6 +227,59 @@ class AgentMessageDeltaBackendClient(FakeAgentBackendRunClient): ) +class FailingStreamBackendClient(FakeAgentBackendRunClient): + def __init__(self) -> None: + super().__init__() + self.cancel_requests: list[CancelRunRequest | None] = [] + + def stream_events( + self, + run_id: str, + *, + after: str | None = None, + should_stop: Callable[[], bool] | None = None, + ) -> Iterator[RunEvent]: + del run_id, after, should_stop + raise AgentBackendStreamError("stream reconnect attempts exhausted") + yield + + def cancel_run(self, run_id: str, request: CancelRunRequest | None = None) -> CancelRunResponse: + self.cancel_requests.append(request) + return CancelRunResponse(run_id=run_id, status="cancelled") + + +class EmptyStreamBackendClient(FailingStreamBackendClient): + def stream_events( + self, + run_id: str, + *, + after: str | None = None, + should_stop: Callable[[], bool] | None = None, + ) -> Iterator[RunEvent]: + del run_id, after, should_stop + return + yield + + +class GenericFailingStreamBackendClient(FailingStreamBackendClient): + def stream_events( + self, + run_id: str, + *, + after: str | None = None, + should_stop: Callable[[], bool] | None = None, + ) -> Iterator[RunEvent]: + del run_id, after, should_stop + raise RuntimeError("unexpected stream failure") + yield + + +class CancelFailingStreamBackendClient(FailingStreamBackendClient): + def cancel_run(self, run_id: str, request: CancelRunRequest | None = None) -> CancelRunResponse: + self.cancel_requests.append(request) + raise RuntimeError(f"failed to cancel {run_id}") + + def _node( *, scenario: FakeAgentBackendScenario = FakeAgentBackendScenario.SUCCESS, @@ -668,6 +732,78 @@ def test_agent_node_repauses_when_resumed_form_still_waiting(monkeypatch): assert client.request is None # no second Agent run was created +def test_agent_node_cancels_backend_run_when_stream_fails(): + client = FailingStreamBackendClient() + node = _node(agent_backend_client=client) + + terminal, failure = node._consume_event_stream("run-1", {"agent_backend": {}}) + + assert terminal is None + assert failure is not None + assert len(client.cancel_requests) == 1 + assert client.cancel_requests[0] is not None + assert client.cancel_requests[0].reason == "event_stream_failed" + + +def test_agent_node_cancels_backend_run_when_stream_ends_without_terminal_event(): + client = EmptyStreamBackendClient() + node = _node(agent_backend_client=client) + + terminal, failure = node._consume_event_stream("run-1", {"agent_backend": {}}) + + assert terminal is None + assert failure is None + assert client.cancel_requests[0] is not None + assert client.cancel_requests[0].reason == "stream_ended_without_terminal_event" + + +def test_agent_node_cancels_backend_run_when_stream_raises_unexpected_error(): + client = GenericFailingStreamBackendClient() + node = _node(agent_backend_client=client) + + terminal, failure = node._consume_event_stream("run-1", {"agent_backend": {}}) + + assert terminal is None + assert failure is not None + assert failure.node_run_result.error == "unexpected stream failure" + assert client.cancel_requests[0] is not None + assert client.cancel_requests[0].reason == "event_stream_failed" + + +def test_agent_node_uses_graph_abort_reason_when_cancel_request_fails(caplog): + client = CancelFailingStreamBackendClient() + node = _node(agent_backend_client=client) + node.graph_runtime_state.graph_execution = SimpleNamespace(aborted=True) + + terminal, failure = node._consume_event_stream("run-1", {"agent_backend": {}}) + + assert terminal is None + assert failure is not None + assert client.cancel_requests[0] is not None + assert client.cancel_requests[0].reason == "workflow_graph_aborted" + assert "Failed to cancel Workflow Agent backend run" in caplog.text + + +def test_agent_node_cancels_backend_run_for_unexpected_internal_event(): + client = FakeAgentBackendRunClient() + node = _node(agent_backend_client=client) + node._agent_backend_client.cancel_run = MagicMock( # type: ignore[method-assign] + return_value=CancelRunResponse(run_id="run-1", status="cancelled") + ) + node._event_adapter.adapt = MagicMock( # type: ignore[method-assign] + return_value=[SimpleNamespace(type=AgentBackendInternalEventType.RUN_FAILED)] + ) + + terminal, failure = node._consume_event_stream("run-1", {"agent_backend": {}}) + + assert terminal is None + assert failure is not None + assert failure.node_run_result.error == ( + "Unexpected internal event type " + ) + node._agent_backend_client.cancel_run.assert_called_once() + + def test_agent_node_records_stream_usage_metadata(): metadata = {"agent_backend": {"run_id": "run-1"}} diff --git a/dify-agent/src/dify_agent/client/_client.py b/dify-agent/src/dify_agent/client/_client.py index c39002bb036..9c15d96cb59 100644 --- a/dify-agent/src/dify_agent/client/_client.py +++ b/dify-agent/src/dify_agent/client/_client.py @@ -15,7 +15,7 @@ import asyncio import inspect import json import time -from collections.abc import AsyncIterator, Iterator +from collections.abc import AsyncIterator, Callable, Iterator from json import JSONDecodeError from types import TracebackType from typing import Any, Self, TypeVar, cast @@ -277,7 +277,7 @@ class Client: *, base_url: str, timeout: float | httpx.Timeout = 30.0, - stream_timeout: float | httpx.Timeout | None = None, + stream_timeout: float | httpx.Timeout | None = 30.0, headers: dict[str, str] | None = None, sync_http_client: httpx.Client | None = None, async_http_client: httpx.AsyncClient | None = None, @@ -531,9 +531,11 @@ class Client: *, after: str | None = None, reconnect: bool = True, - max_reconnects: int | None = None, + max_reconnects: int | None = 3, reconnect_delay_seconds: float = 1.0, until_terminal: bool = True, + timeout_seconds: float | None = None, + should_stop: Callable[[], bool] | None = None, ) -> AsyncIterator[RunEvent]: """Yield typed events from SSE with cursor-based reconnect. @@ -541,14 +543,21 @@ class Client: with an id, reconnects resume from that id using the ``after`` query parameter. HTTP 5xx stream responses are retried, but HTTP 4xx responses, DTO validation failures, and malformed SSE frames are not retried. By - default iteration stops after ``run_succeeded`` or ``run_failed``. + default iteration stops after a succeeded, failed, or cancelled terminal event. """ - _validate_stream_options(max_reconnects, reconnect_delay_seconds) + _validate_stream_options(max_reconnects, reconnect_delay_seconds, timeout_seconds) cursor = after or "0-0" reconnect_attempts = 0 + deadline = time.monotonic() + timeout_seconds if timeout_seconds is not None else None while True: + _raise_if_stream_stopped(run_id, deadline=deadline, should_stop=should_stop) try: - async for event in self._stream_events_once(run_id, after=cursor): + async for event in self._stream_events_once( + run_id, + after=cursor, + deadline=deadline, + should_stop=should_stop, + ): if event.id is not None: cursor = event.id yield event @@ -562,7 +571,8 @@ class Client: max_reconnects=max_reconnects, error=exc.error, ) - await _sleep_async(reconnect_delay_seconds) + _raise_if_stream_stopped(run_id, deadline=deadline, should_stop=should_stop) + await _sleep_async(_bounded_sleep_seconds(reconnect_delay_seconds, deadline)) continue if not reconnect: return @@ -571,7 +581,8 @@ class Client: max_reconnects=max_reconnects, error=DifyAgentStreamError("SSE stream ended before a terminal event"), ) - await _sleep_async(reconnect_delay_seconds) + _raise_if_stream_stopped(run_id, deadline=deadline, should_stop=should_stop) + await _sleep_async(_bounded_sleep_seconds(reconnect_delay_seconds, deadline)) def stream_events_sync( self, @@ -579,17 +590,26 @@ class Client: *, after: str | None = None, reconnect: bool = True, - max_reconnects: int | None = None, + max_reconnects: int | None = 3, reconnect_delay_seconds: float = 1.0, until_terminal: bool = True, + timeout_seconds: float | None = None, + should_stop: Callable[[], bool] | None = None, ) -> Iterator[RunEvent]: """Synchronous variant of ``stream_events`` with the same reconnect rules.""" - _validate_stream_options(max_reconnects, reconnect_delay_seconds) + _validate_stream_options(max_reconnects, reconnect_delay_seconds, timeout_seconds) cursor = after or "0-0" reconnect_attempts = 0 + deadline = time.monotonic() + timeout_seconds if timeout_seconds is not None else None while True: + _raise_if_stream_stopped(run_id, deadline=deadline, should_stop=should_stop) try: - for event in self._stream_events_once_sync(run_id, after=cursor): + for event in self._stream_events_once_sync( + run_id, + after=cursor, + deadline=deadline, + should_stop=should_stop, + ): if event.id is not None: cursor = event.id yield event @@ -603,7 +623,8 @@ class Client: max_reconnects=max_reconnects, error=exc.error, ) - _sleep_sync(reconnect_delay_seconds) + _raise_if_stream_stopped(run_id, deadline=deadline, should_stop=should_stop) + _sleep_sync(_bounded_sleep_seconds(reconnect_delay_seconds, deadline)) continue if not reconnect: return @@ -612,7 +633,8 @@ class Client: max_reconnects=max_reconnects, error=DifyAgentStreamError("SSE stream ended before a terminal event"), ) - _sleep_sync(reconnect_delay_seconds) + _raise_if_stream_stopped(run_id, deadline=deadline, should_stop=should_stop) + _sleep_sync(_bounded_sleep_seconds(reconnect_delay_seconds, deadline)) async def wait_run( self, @@ -652,7 +674,14 @@ class Client: raise DifyAgentTimeoutError(f"run {run_id!r} did not finish before timeout") _sleep_sync(sleep_for) - async def _stream_events_once(self, run_id: str, *, after: str) -> AsyncIterator[RunEvent]: + async def _stream_events_once( + self, + run_id: str, + *, + after: str, + deadline: float | None, + should_stop: Callable[[], bool] | None, + ) -> AsyncIterator[RunEvent]: """Open one SSE connection and yield events until it ends or fails.""" try: async with self._get_async_http_client().stream( @@ -668,6 +697,7 @@ class Client: decoder = _SSEDecoder() line_decoder = _SSELineDecoder() async for text in response.aiter_text(): + _raise_if_stream_stopped(run_id, deadline=deadline, should_stop=should_stop) for line in line_decoder.decode(text): event = decoder.feed_line(line) if event is not None: @@ -687,7 +717,14 @@ class Client: except httpx.StreamError as exc: raise _ReconnectableStreamError(DifyAgentStreamError(f"SSE stream failed: {exc}")) from exc - def _stream_events_once_sync(self, run_id: str, *, after: str) -> Iterator[RunEvent]: + def _stream_events_once_sync( + self, + run_id: str, + *, + after: str, + deadline: float | None, + should_stop: Callable[[], bool] | None, + ) -> Iterator[RunEvent]: """Open one synchronous SSE connection and yield events until it ends or fails.""" try: with self._get_sync_http_client().stream( @@ -703,6 +740,7 @@ class Client: decoder = _SSEDecoder() line_decoder = _SSELineDecoder() for text in response.iter_text(): + _raise_if_stream_stopped(run_id, deadline=deadline, should_stop=should_stop) for line in line_decoder.decode(text): event = decoder.feed_line(line) if event is not None: @@ -852,12 +890,38 @@ def _next_reconnect_attempt( return reconnect_attempts + 1 -def _validate_stream_options(max_reconnects: int | None, reconnect_delay_seconds: float) -> None: +def _validate_stream_options( + max_reconnects: int | None, + reconnect_delay_seconds: float, + timeout_seconds: float | None, +) -> None: """Reject stream options that cannot produce deterministic reconnect behavior.""" if max_reconnects is not None and max_reconnects < 0: raise DifyAgentValidationError(detail="max_reconnects must be non-negative") if reconnect_delay_seconds < 0: raise DifyAgentValidationError(detail="reconnect_delay_seconds must be non-negative") + if timeout_seconds is not None and timeout_seconds < 0: + raise DifyAgentValidationError(detail="timeout_seconds must be non-negative") + + +def _raise_if_stream_stopped( + run_id: str, + *, + deadline: float | None, + should_stop: Callable[[], bool] | None, +) -> None: + """Stop a live stream when its caller cancels or its total deadline expires.""" + if should_stop is not None and should_stop(): + raise DifyAgentStreamError(f"SSE stream for run {run_id!r} was cancelled by the caller") + if deadline is not None and time.monotonic() >= deadline: + raise DifyAgentTimeoutError(f"SSE stream for run {run_id!r} exceeded its timeout") + + +def _bounded_sleep_seconds(seconds: float, deadline: float | None) -> float: + """Keep reconnect backoff inside the total stream deadline.""" + if deadline is None: + return seconds + return max(0.0, min(seconds, deadline - time.monotonic())) def _validate_wait_options(poll_interval_seconds: float, timeout_seconds: float | None) -> None: diff --git a/dify-agent/src/dify_agent/runtime/event_sink.py b/dify-agent/src/dify_agent/runtime/event_sink.py index 71df75b4916..67babdc3754 100644 --- a/dify-agent/src/dify_agent/runtime/event_sink.py +++ b/dify-agent/src/dify_agent/runtime/event_sink.py @@ -22,6 +22,8 @@ from dify_agent.protocol.schemas import ( EmptyRunEventData, PydanticAIStreamRunEvent, RunEvent, + RunCancelledEvent, + RunCancelledEventData, RunFailedEvent, RunFailedEventData, RunStartedEvent, @@ -159,10 +161,29 @@ async def emit_run_failed( ) +async def emit_run_cancelled( + sink: RunEventSink, + *, + run_id: str, + reason: str | None = None, + message: str | None = None, +) -> str: + """Emit the terminal cancellation lifecycle event.""" + return await emit_run_event( + sink, + event=RunCancelledEvent( + run_id=run_id, + data=RunCancelledEventData(reason=reason, message=message), + created_at=utc_now(), + ), + ) + + __all__ = [ "InMemoryRunEventSink", "RunEventSink", "emit_pydantic_ai_event", + "emit_run_cancelled", "emit_run_event", "emit_run_failed", "emit_run_started", diff --git a/dify-agent/src/dify_agent/runtime/run_scheduler.py b/dify-agent/src/dify_agent/runtime/run_scheduler.py index 4186b6afd76..c02799a66ec 100644 --- a/dify-agent/src/dify_agent/runtime/run_scheduler.py +++ b/dify-agent/src/dify_agent/runtime/run_scheduler.py @@ -20,9 +20,9 @@ from typing import Protocol import httpx from agenton.compositor import LayerProviderInput -from dify_agent.protocol.schemas import CreateRunRequest +from dify_agent.protocol.schemas import CancelRunRequest, CancelRunResponse, CreateRunRequest from dify_agent.runtime.compositor_factory import create_default_layer_providers -from dify_agent.runtime.event_sink import RunEventSink, emit_run_failed +from dify_agent.runtime.event_sink import RunEventSink, emit_run_cancelled, emit_run_failed from dify_agent.runtime.runner import AgentRunRunner from dify_agent.server.schemas import RunRecord @@ -33,6 +33,10 @@ class SchedulerStoppingError(RuntimeError): """Raised when a create-run request arrives after shutdown has started.""" +class RunCancellationConflictError(RuntimeError): + """Raised when a run exists but can no longer be cancelled by this scheduler.""" + + class RunStore(RunEventSink, Protocol): """Persistence boundary needed by the scheduler.""" @@ -40,6 +44,10 @@ class RunStore(RunEventSink, Protocol): """Persist a new run record and return it with status ``running``.""" ... + async def get_run(self, run_id: str) -> RunRecord: + """Return the latest persisted run record.""" + ... + class RunnableRun(Protocol): """Executable unit for one scheduled run.""" @@ -65,6 +73,7 @@ class RunScheduler: store: RunStore shutdown_grace_seconds: float active_tasks: dict[str, asyncio.Task[None]] + cancelled_run_ids: set[str] stopping: bool runner_factory: RunRunnerFactory layer_providers: tuple[LayerProviderInput, ...] @@ -85,6 +94,7 @@ class RunScheduler: self.store = store self.shutdown_grace_seconds = shutdown_grace_seconds self.active_tasks = {} + self.cancelled_run_ids = set() self.stopping = False self.plugin_daemon_http_client = plugin_daemon_http_client self.dify_api_http_client = dify_api_http_client @@ -106,9 +116,43 @@ class RunScheduler: record = await self.store.create_run() task = asyncio.create_task(self._run_record(record, request), name=f"dify-agent-run-{record.run_id}") self.active_tasks[record.run_id] = task - task.add_done_callback(lambda _task, run_id=record.run_id: self.active_tasks.pop(run_id, None)) + task.add_done_callback(lambda _task, run_id=record.run_id: self._discard_active_run(run_id)) return record + async def cancel_run(self, run_id: str, request: CancelRunRequest) -> CancelRunResponse: + """Cancel one active task and persist an idempotent cancelled terminal state.""" + async with self._lifecycle_lock: + record = await self.store.get_run(run_id) + if record.status == "cancelled": + return CancelRunResponse(run_id=run_id, status="cancelled") + if record.status != "running": + raise RunCancellationConflictError(f"run already finished with status {record.status!r}") + + task = self.active_tasks.get(run_id) + if task is None: + raise RunCancellationConflictError("run is not active in this scheduler process") + self.cancelled_run_ids.add(run_id) + _ = task.cancel(request.message or request.reason) + _ = await emit_run_cancelled( + self.store, + run_id=run_id, + reason=request.reason, + message=request.message, + ) + await self.store.update_status(run_id, "cancelled", request.message or request.reason) + + # Some model/tool stacks can consume one CancelledError. Re-inject it + # after the terminal state is durable without making the HTTP request + # wait for arbitrary third-party cleanup. + for _attempt in range(2): + if task.done(): + break + _ = task.cancel(request.message or request.reason) + await asyncio.sleep(0) + if task.done(): + self._discard_active_run(run_id) + return CancelRunResponse(run_id=run_id, status="cancelled") + async def shutdown(self) -> None: """Stop accepting runs, wait briefly, then cancel and fail unfinished runs.""" async with self._lifecycle_lock: @@ -121,7 +165,11 @@ class RunScheduler: if not pending: return - pending_run_ids = [run_id for run_id, task in tasks_by_run_id.items() if task in pending] + pending_run_ids = [ + run_id + for run_id, task in tasks_by_run_id.items() + if task in pending and run_id not in self.cancelled_run_ids + ] for task in pending: _ = task.cancel() _ = await asyncio.gather(*pending, return_exceptions=True) @@ -146,8 +194,13 @@ class RunScheduler: plugin_daemon_http_client=self.plugin_daemon_http_client, dify_api_http_client=self.dify_api_http_client, layer_providers=self.layer_providers, + is_cancelled=lambda: record.run_id in self.cancelled_run_ids, ) + def _discard_active_run(self, run_id: str) -> None: + _ = self.active_tasks.pop(run_id, None) + self.cancelled_run_ids.discard(run_id) + async def _mark_cancelled_run_failed(self, run_id: str) -> None: """Best-effort failure event/status for shutdown-cancelled runs.""" message = "run cancelled during server shutdown" @@ -158,4 +211,4 @@ class RunScheduler: logger.exception("failed to mark cancelled run failed", extra={"run_id": run_id}) -__all__ = ["RunScheduler", "SchedulerStoppingError"] +__all__ = ["RunCancellationConflictError", "RunScheduler", "SchedulerStoppingError"] diff --git a/dify-agent/src/dify_agent/runtime/runner.py b/dify-agent/src/dify_agent/runtime/runner.py index 97c1ce7ef1f..315f4484c4b 100644 --- a/dify-agent/src/dify_agent/runtime/runner.py +++ b/dify-agent/src/dify_agent/runtime/runner.py @@ -31,6 +31,7 @@ both the JSON-safe final output or deferred tool call and the session snapshot; there are no separate output or snapshot events to correlate. """ +import asyncio from collections.abc import AsyncIterable, Callable, Mapping from collections import Counter from dataclasses import dataclass @@ -170,6 +171,7 @@ class AgentRunRunner: layer_providers: tuple[LayerProviderInput, ...] plugin_daemon_http_client: httpx.AsyncClient dify_api_http_client: httpx.AsyncClient + is_cancelled: Callable[[], bool] def __init__( self, @@ -180,6 +182,7 @@ class AgentRunRunner: plugin_daemon_http_client: httpx.AsyncClient, dify_api_http_client: httpx.AsyncClient, layer_providers: tuple[LayerProviderInput, ...] | None = None, + is_cancelled: Callable[[], bool] | None = None, ) -> None: self.sink = sink self.request = request @@ -187,20 +190,29 @@ class AgentRunRunner: self.plugin_daemon_http_client = plugin_daemon_http_client self.dify_api_http_client = dify_api_http_client self.layer_providers = layer_providers if layer_providers is not None else create_default_layer_providers() + self.is_cancelled = is_cancelled or (lambda: False) async def run(self) -> None: """Execute the run and emit the documented event sequence.""" + if self.is_cancelled(): + return await self.sink.update_status(self.run_id, "running") + if self.is_cancelled(): + return _ = await emit_run_started(self.sink, run_id=self.run_id) try: outcome = await self._run_agent() except Exception as exc: + if self.is_cancelled(): + return message, reason = _run_failed_error_payload(exc) _ = await emit_run_failed(self.sink, run_id=self.run_id, error=message, reason=reason) await self.sink.update_status(self.run_id, "failed", message) raise + if self.is_cancelled(): + return _ = await emit_run_succeeded( self.sink, run_id=self.run_id, @@ -309,6 +321,8 @@ class AgentRunRunner: async def handle_events(_ctx: object, events: AsyncIterable[AgentStreamEvent]) -> None: async for event in events: + if self.is_cancelled(): + raise asyncio.CancelledError text_delta = _extract_agent_message_delta(event) _ = await emit_pydantic_ai_event( self.sink, diff --git a/dify-agent/src/dify_agent/server/routes/runs.py b/dify-agent/src/dify_agent/server/routes/runs.py index 1cbd9d20945..f41567648e8 100644 --- a/dify-agent/src/dify_agent/server/routes/runs.py +++ b/dify-agent/src/dify_agent/server/routes/runs.py @@ -24,7 +24,7 @@ from dify_agent.protocol.schemas import ( RunEventsResponse, RunStatusResponse, ) -from dify_agent.runtime.run_scheduler import RunScheduler, SchedulerStoppingError +from dify_agent.runtime.run_scheduler import RunCancellationConflictError, RunScheduler, SchedulerStoppingError from dify_agent.server.sse import sse_event_stream from dify_agent.storage.redis_run_store import RedisRunStore, RunNotFoundError @@ -68,16 +68,18 @@ def create_runs_router( ) @router.post("/{run_id}/cancel", response_model=CancelRunResponse) - async def cancel_run(run_id: str, request: CancelRunRequest) -> CancelRunResponse: - """Reserve the cancellation endpoint in the public protocol. - - Runtime cancellation requires scheduler task lookup and persistence - semantics that are outside the current server implementation. Exposing a - typed endpoint now lets clients bind to the final route while receiving - an explicit 501 until execution support lands. - """ - del run_id, request - raise HTTPException(status_code=501, detail="run cancellation is not implemented") + async def cancel_run( + run_id: str, + request: CancelRunRequest, + scheduler: Annotated[RunScheduler, Depends(scheduler_dep)], + ) -> CancelRunResponse: + """Cancel a process-local run and publish its terminal event/status.""" + try: + return await scheduler.cancel_run(run_id, request) + except RunNotFoundError as exc: + raise HTTPException(status_code=404, detail="run not found") from exc + except RunCancellationConflictError as exc: + raise HTTPException(status_code=409, detail=str(exc)) from exc @router.get("/{run_id}/events", response_model=RunEventsResponse) async def get_run_events( diff --git a/dify-agent/src/dify_agent/server/sse.py b/dify-agent/src/dify_agent/server/sse.py index cd9018f2c6a..09c72673ff2 100644 --- a/dify-agent/src/dify_agent/server/sse.py +++ b/dify-agent/src/dify_agent/server/sse.py @@ -5,6 +5,7 @@ browsers can resume with ``Last-Event-ID`` while clients can subscribe by event name. Payload data is the full public ``RunEvent`` JSON object. """ +import asyncio from collections.abc import AsyncIterable, AsyncIterator from dify_agent.protocol.schemas import RUN_EVENT_ADAPTER, RunEvent @@ -29,10 +30,33 @@ def format_sse_event(event: RunEvent) -> str: return "\n".join(lines) + "\n\n" -async def sse_event_stream(events: AsyncIterable[RunEvent]) -> AsyncIterator[str]: - """Yield formatted SSE frames from public run events.""" - async for event in events: - yield format_sse_event(event) +async def sse_event_stream( + events: AsyncIterable[RunEvent], + *, + heartbeat_interval_seconds: float = 15.0, +) -> AsyncIterator[str]: + """Yield events and keep idle SSE connections observable to clients.""" + if heartbeat_interval_seconds <= 0: + raise ValueError("heartbeat_interval_seconds must be positive") + + iterator = events.__aiter__() + next_event = asyncio.ensure_future(anext(iterator)) + try: + while True: + done, _ = await asyncio.wait({next_event}, timeout=heartbeat_interval_seconds) + if not done: + yield ": keepalive\n\n" + continue + try: + event = next_event.result() + except StopAsyncIteration: + return + yield format_sse_event(event) + next_event = asyncio.ensure_future(anext(iterator)) + finally: + if not next_event.done(): + _ = next_event.cancel() + _ = await asyncio.gather(next_event, return_exceptions=True) __all__ = ["format_sse_event", "sse_event_stream"] diff --git a/dify-agent/tests/local/dify_agent/client/test_client.py b/dify-agent/tests/local/dify_agent/client/test_client.py index 9d4c16471b5..37d5c540987 100644 --- a/dify-agent/tests/local/dify_agent/client/test_client.py +++ b/dify-agent/tests/local/dify_agent/client/test_client.py @@ -589,6 +589,63 @@ def test_stream_events_raises_when_reconnects_are_exhausted() -> None: assert calls == 2 +def test_stream_events_default_reconnect_budget_is_finite() -> None: + calls = 0 + + def handler(_request: httpx.Request) -> httpx.Response: + nonlocal calls + calls += 1 + return httpx.Response(200, content="") + + client = Client( + base_url="http://testserver", + sync_http_client=httpx.Client(transport=httpx.MockTransport(handler)), + ) + + with pytest.raises(DifyAgentStreamError, match="reconnect attempts exhausted"): + _ = list(client.stream_events_sync("run-1", reconnect_delay_seconds=0)) + assert calls == 4 + + +def test_stream_events_enforces_total_timeout_before_connecting() -> None: + calls = 0 + + def handler(_request: httpx.Request) -> httpx.Response: + nonlocal calls + calls += 1 + return httpx.Response(200, content="") + + client = Client( + base_url="http://testserver", + sync_http_client=httpx.Client(transport=httpx.MockTransport(handler)), + ) + + with pytest.raises(DifyAgentTimeoutError, match="exceeded its timeout"): + _ = list(client.stream_events_sync("run-1", timeout_seconds=0)) + assert calls == 0 + + +def test_stream_events_observes_caller_stop_on_heartbeat() -> None: + stop_checks = 0 + + def should_stop() -> bool: + nonlocal stop_checks + stop_checks += 1 + return stop_checks >= 2 + + def handler(_request: httpx.Request) -> httpx.Response: + return httpx.Response(200, content=": keepalive\n\n") + + client = Client( + base_url="http://testserver", + sync_http_client=httpx.Client(transport=httpx.MockTransport(handler)), + ) + + with pytest.raises(DifyAgentStreamError, match="cancelled by the caller"): + _ = list(client.stream_events_sync("run-1", should_stop=should_stop)) + assert stop_checks == 2 + + def test_malformed_sse_frame_does_not_reconnect() -> None: calls = 0 diff --git a/dify-agent/tests/local/dify_agent/runtime/test_run_scheduler.py b/dify-agent/tests/local/dify_agent/runtime/test_run_scheduler.py index e1124560ac6..ac72d0502f1 100644 --- a/dify-agent/tests/local/dify_agent/runtime/test_run_scheduler.py +++ b/dify-agent/tests/local/dify_agent/runtime/test_run_scheduler.py @@ -11,13 +11,14 @@ from agenton_collections.layers.plain import PromptLayerConfig from dify_agent.layers.output import DIFY_OUTPUT_LAYER_TYPE_ID, DifyOutputLayerConfig from dify_agent.protocol import DIFY_AGENT_OUTPUT_LAYER_ID from dify_agent.protocol.schemas import ( + CancelRunRequest, CreateRunRequest, RunComposition, RunEvent, RunLayerSpec, RunStatus, ) -from dify_agent.runtime.run_scheduler import RunScheduler, SchedulerStoppingError +from dify_agent.runtime.run_scheduler import RunCancellationConflictError, RunScheduler, SchedulerStoppingError from dify_agent.server.schemas import RunRecord @@ -78,6 +79,11 @@ class FakeStore: self.events[event.run_id].append(event.model_copy(update={"id": event_id})) return event_id + async def get_run(self, run_id: str) -> RunRecord: + return self.records[run_id].model_copy( + update={"status": self.statuses[run_id], "error": self.errors.get(run_id)}, + ) + async def update_status(self, run_id: str, status: RunStatus, error: str | None = None) -> None: self.statuses[run_id] = status self.errors[run_id] = error @@ -111,6 +117,23 @@ class ControlledRunner: await self.release.wait() +class SwallowOneCancellationRunner: + started: asyncio.Event + first_cancellation: asyncio.Event + + def __init__(self, *, started: asyncio.Event, first_cancellation: asyncio.Event) -> None: + self.started = started + self.first_cancellation = first_cancellation + + async def run(self) -> None: + _ = self.started.set() + try: + await asyncio.Event().wait() + except asyncio.CancelledError: + _ = self.first_cancellation.set() + await asyncio.Event().wait() + + def test_create_run_starts_background_task_and_returns_running() -> None: async def scenario() -> None: store = FakeStore() @@ -163,6 +186,85 @@ def test_shutdown_marks_unfinished_runs_failed_and_appends_event() -> None: asyncio.run(scenario()) +def test_cancel_run_stops_task_and_persists_cancelled_terminal() -> None: + async def scenario() -> None: + store = FakeStore() + started = asyncio.Event() + async with httpx.AsyncClient() as client: + scheduler = RunScheduler( + store=store, + plugin_daemon_http_client=client, + dify_api_http_client=client, + runner_factory=lambda _record, _request: ControlledRunner(started=started, release=asyncio.Event()), + ) + record = await scheduler.create_run(_request()) + await asyncio.wait_for(started.wait(), timeout=1) + + response = await scheduler.cancel_run( + record.run_id, + CancelRunRequest(reason="workflow_aborted", message="outer workflow stopped"), + ) + + assert response.status == "cancelled" + assert scheduler.active_tasks == {} + assert store.statuses[record.run_id] == "cancelled" + assert store.errors[record.run_id] == "outer workflow stopped" + assert [event.type for event in store.events[record.run_id]] == ["run_cancelled"] + + repeated = await scheduler.cancel_run(record.run_id, CancelRunRequest(reason="duplicate")) + assert repeated.status == "cancelled" + assert [event.type for event in store.events[record.run_id]] == ["run_cancelled"] + + asyncio.run(scenario()) + + +def test_cancel_run_reinjects_cancellation_without_waiting_for_runner_cleanup() -> None: + async def scenario() -> None: + store = FakeStore() + started = asyncio.Event() + first_cancellation = asyncio.Event() + async with httpx.AsyncClient() as client: + scheduler = RunScheduler( + store=store, + plugin_daemon_http_client=client, + dify_api_http_client=client, + runner_factory=lambda _record, _request: SwallowOneCancellationRunner( + started=started, + first_cancellation=first_cancellation, + ), + ) + record = await scheduler.create_run(_request()) + await asyncio.wait_for(started.wait(), timeout=1) + + response = await asyncio.wait_for( + scheduler.cancel_run(record.run_id, CancelRunRequest(reason="workflow_aborted")), + timeout=1, + ) + + assert response.status == "cancelled" + assert first_cancellation.is_set() + assert store.statuses[record.run_id] == "cancelled" + assert [event.type for event in store.events[record.run_id]] == ["run_cancelled"] + await asyncio.sleep(0) + assert scheduler.active_tasks == {} + + asyncio.run(scenario()) + + +def test_cancel_run_rejects_finished_run() -> None: + async def scenario() -> None: + store = FakeStore() + async with httpx.AsyncClient() as client: + scheduler = RunScheduler(store=store, plugin_daemon_http_client=client, dify_api_http_client=client) + record = await store.create_run() + await store.update_status(record.run_id, "succeeded") + + with pytest.raises(RunCancellationConflictError, match="already finished"): + await scheduler.cancel_run(record.run_id, CancelRunRequest()) + + asyncio.run(scenario()) + + def test_create_run_accepts_blank_prompt_and_runner_fails_asynchronously() -> None: async def scenario() -> None: store = FakeStore() diff --git a/dify-agent/tests/local/dify_agent/runtime/test_runner.py b/dify-agent/tests/local/dify_agent/runtime/test_runner.py index 25814034220..07886a9dfa2 100644 --- a/dify-agent/tests/local/dify_agent/runtime/test_runner.py +++ b/dify-agent/tests/local/dify_agent/runtime/test_runner.py @@ -64,7 +64,12 @@ from dify_agent.protocol.schemas import ( ) from dify_agent.runtime.event_sink import InMemoryRunEventSink from dify_agent.runtime.compositor_factory import create_default_layer_providers -from dify_agent.runtime.runner import AgentRunRunner, AgentRunValidationError, _run_failed_error_payload +from dify_agent.runtime.runner import ( + AgentRunRunner, + AgentRunValidationError, + RunSuccessOutcome, + _run_failed_error_payload, +) from shellctl.shared import DeleteJobResponse, JobResult, JobStatusName, JobStatusView @@ -164,6 +169,38 @@ def test_run_failed_error_payload_preserves_knowledge_error_code() -> None: assert reason == "dataset_not_found" +def test_cancelled_runner_does_not_overwrite_cancelled_status_with_late_failure( + monkeypatch: pytest.MonkeyPatch, +) -> None: + async def scenario() -> None: + sink = InMemoryRunEventSink() + cancelled = False + async with httpx.AsyncClient() as client: + runner = AgentRunRunner( + sink=sink, + request=_request(), + run_id="run-cancelled", + plugin_daemon_http_client=client, + dify_api_http_client=client, + is_cancelled=lambda: cancelled, + ) + + async def fail_after_cancel() -> RunSuccessOutcome: + nonlocal cancelled + cancelled = True + await sink.update_status("run-cancelled", "cancelled", "workflow stopped") + raise RuntimeError("late model failure") + + monkeypatch.setattr(runner, "_run_agent", fail_after_cancel) + await runner.run() + + assert sink.statuses["run-cancelled"] == "cancelled" + assert sink.errors["run-cancelled"] == "workflow stopped" + assert [event.type for event in sink.events["run-cancelled"]] == ["run_started"] + + asyncio.run(scenario()) + + def _request( user: str | list[str] = "hello", *, diff --git a/dify-agent/tests/local/dify_agent/server/test_runs_routes.py b/dify-agent/tests/local/dify_agent/server/test_runs_routes.py index d925580d6e5..d90f82a152a 100644 --- a/dify-agent/tests/local/dify_agent/server/test_runs_routes.py +++ b/dify-agent/tests/local/dify_agent/server/test_runs_routes.py @@ -1,7 +1,7 @@ from fastapi.testclient import TestClient -from dify_agent.protocol import DIFY_AGENT_MODEL_LAYER_ID -from dify_agent.runtime.run_scheduler import SchedulerStoppingError +from dify_agent.protocol import CancelRunResponse, DIFY_AGENT_MODEL_LAYER_ID +from dify_agent.runtime.run_scheduler import RunCancellationConflictError, SchedulerStoppingError from dify_agent.server.routes.runs import create_runs_router from dify_agent.server.schemas import RunRecord @@ -11,6 +11,10 @@ class FakeScheduler: del request return RunRecord(run_id="run-1", status="running") + async def cancel_run(self, run_id: str, request: object) -> CancelRunResponse: + del request + return CancelRunResponse(run_id=run_id, status="cancelled") + class FakeStore: pass @@ -67,7 +71,7 @@ def test_create_run_returns_running_from_scheduler() -> None: assert response.json() == {"run_id": "run-1", "status": "running"} -def test_cancel_run_endpoint_is_reserved_but_not_implemented() -> None: +def test_cancel_run_endpoint_returns_scheduler_result() -> None: from fastapi import FastAPI app = FastAPI() @@ -78,8 +82,28 @@ def test_cancel_run_endpoint_is_reserved_but_not_implemented() -> None: response = client.post("/runs/run-1/cancel", json={"reason": "user_cancelled"}) - assert response.status_code == 501 - assert response.json()["detail"] == "run cancellation is not implemented" + assert response.status_code == 200 + assert response.json() == {"run_id": "run-1", "status": "cancelled"} + + +def test_cancel_run_endpoint_maps_conflict() -> None: + from fastapi import FastAPI + + class ConflictingScheduler(FakeScheduler): + async def cancel_run(self, run_id: str, request: object) -> CancelRunResponse: + del run_id, request + raise RunCancellationConflictError("run already finished with status 'succeeded'") + + app = FastAPI() + app.include_router( + create_runs_router(lambda: FakeStore(), lambda: ConflictingScheduler()) # pyright: ignore[reportArgumentType] + ) + client = TestClient(app) + + response = client.post("/runs/run-1/cancel", json={}) + + assert response.status_code == 409 + assert "already finished" in response.json()["detail"] def test_create_run_accepts_valid_full_plugin_graph() -> None: diff --git a/dify-agent/tests/local/dify_agent/server/test_sse.py b/dify-agent/tests/local/dify_agent/server/test_sse.py index 1c1188e8611..8f146ad396c 100644 --- a/dify-agent/tests/local/dify_agent/server/test_sse.py +++ b/dify-agent/tests/local/dify_agent/server/test_sse.py @@ -1,7 +1,10 @@ +import asyncio import json +from collections.abc import AsyncGenerator +from typing import cast from dify_agent.protocol.schemas import RunFailedEvent, RunFailedEventData, RunStartedEvent -from dify_agent.server.sse import format_sse_event +from dify_agent.server.sse import format_sse_event, sse_event_stream def test_format_sse_event_uses_id_event_and_json_data() -> None: @@ -28,3 +31,21 @@ def test_format_sse_event_escapes_unicode_line_separators() -> None: assert "\\u2028" in frame assert "\\u2029" in frame assert json.loads(data)["data"]["error"] == error + + +def test_sse_event_stream_emits_heartbeats_while_waiting() -> None: + async def scenario() -> None: + release = asyncio.Event() + + async def events(): + await release.wait() + yield RunStartedEvent(id="1-0", run_id="run-1") + + stream = cast(AsyncGenerator[str, None], sse_event_stream(events(), heartbeat_interval_seconds=0.001)) + assert await anext(stream) == ": keepalive\n\n" + + _ = release.set() + assert (await anext(stream)).startswith("id: 1-0\nevent: run_started") + await stream.aclose() + + asyncio.run(scenario()) diff --git a/docker/.env.example b/docker/.env.example index 8d5d81e3f28..2071618fba8 100644 --- a/docker/.env.example +++ b/docker/.env.example @@ -252,6 +252,9 @@ MARKETPLACE_URL= # Dify Agent backend AGENT_BACKEND_BASE_URL=http://agent_backend:5050 +AGENT_BACKEND_STREAM_READ_TIMEOUT_SECONDS=30 +AGENT_BACKEND_STREAM_MAX_RECONNECTS=3 +AGENT_BACKEND_RUN_TIMEOUT_SECONDS=1200 # Leave empty to derive from REDIS_PASSWORD. DIFY_AGENT_REDIS_URL= DIFY_AGENT_REDIS_PREFIX=dify-agent diff --git a/docker/docker-compose-template.yaml b/docker/docker-compose-template.yaml index fc4f2471863..ed42c2c5e8b 100644 --- a/docker/docker-compose-template.yaml +++ b/docker/docker-compose-template.yaml @@ -232,6 +232,9 @@ services: PLUGIN_DAEMON_TIMEOUT: ${PLUGIN_DAEMON_TIMEOUT:-600.0} INNER_API_KEY_FOR_PLUGIN: ${PLUGIN_DIFY_INNER_API_KEY:-QaHbTe77CtuXmsfyhR7+vRjI/+XbV1AaFy691iy+kGDv2Jvy0/eAh8Y1} AGENT_BACKEND_BASE_URL: ${AGENT_BACKEND_BASE_URL:-http://agent_backend:5050} + AGENT_BACKEND_STREAM_READ_TIMEOUT_SECONDS: ${AGENT_BACKEND_STREAM_READ_TIMEOUT_SECONDS:-30} + AGENT_BACKEND_STREAM_MAX_RECONNECTS: ${AGENT_BACKEND_STREAM_MAX_RECONNECTS:-3} + AGENT_BACKEND_RUN_TIMEOUT_SECONDS: ${AGENT_BACKEND_RUN_TIMEOUT_SECONDS:-1200} depends_on: init_permissions: condition: service_completed_successfully @@ -302,6 +305,9 @@ services: PLUGIN_MAX_PACKAGE_SIZE: ${PLUGIN_MAX_PACKAGE_SIZE:-52428800} INNER_API_KEY_FOR_PLUGIN: ${PLUGIN_DIFY_INNER_API_KEY:-QaHbTe77CtuXmsfyhR7+vRjI/+XbV1AaFy691iy+kGDv2Jvy0/eAh8Y1} AGENT_BACKEND_BASE_URL: ${AGENT_BACKEND_BASE_URL:-http://agent_backend:5050} + AGENT_BACKEND_STREAM_READ_TIMEOUT_SECONDS: ${AGENT_BACKEND_STREAM_READ_TIMEOUT_SECONDS:-30} + AGENT_BACKEND_STREAM_MAX_RECONNECTS: ${AGENT_BACKEND_STREAM_MAX_RECONNECTS:-3} + AGENT_BACKEND_RUN_TIMEOUT_SECONDS: ${AGENT_BACKEND_RUN_TIMEOUT_SECONDS:-1200} depends_on: init_permissions: condition: service_completed_successfully diff --git a/docker/docker-compose.yaml b/docker/docker-compose.yaml index 7f29ab24168..9c439cad8af 100644 --- a/docker/docker-compose.yaml +++ b/docker/docker-compose.yaml @@ -238,6 +238,9 @@ services: PLUGIN_DAEMON_TIMEOUT: ${PLUGIN_DAEMON_TIMEOUT:-600.0} INNER_API_KEY_FOR_PLUGIN: ${PLUGIN_DIFY_INNER_API_KEY:-QaHbTe77CtuXmsfyhR7+vRjI/+XbV1AaFy691iy+kGDv2Jvy0/eAh8Y1} AGENT_BACKEND_BASE_URL: ${AGENT_BACKEND_BASE_URL:-http://agent_backend:5050} + AGENT_BACKEND_STREAM_READ_TIMEOUT_SECONDS: ${AGENT_BACKEND_STREAM_READ_TIMEOUT_SECONDS:-30} + AGENT_BACKEND_STREAM_MAX_RECONNECTS: ${AGENT_BACKEND_STREAM_MAX_RECONNECTS:-3} + AGENT_BACKEND_RUN_TIMEOUT_SECONDS: ${AGENT_BACKEND_RUN_TIMEOUT_SECONDS:-1200} depends_on: init_permissions: condition: service_completed_successfully @@ -308,6 +311,9 @@ services: PLUGIN_MAX_PACKAGE_SIZE: ${PLUGIN_MAX_PACKAGE_SIZE:-52428800} INNER_API_KEY_FOR_PLUGIN: ${PLUGIN_DIFY_INNER_API_KEY:-QaHbTe77CtuXmsfyhR7+vRjI/+XbV1AaFy691iy+kGDv2Jvy0/eAh8Y1} AGENT_BACKEND_BASE_URL: ${AGENT_BACKEND_BASE_URL:-http://agent_backend:5050} + AGENT_BACKEND_STREAM_READ_TIMEOUT_SECONDS: ${AGENT_BACKEND_STREAM_READ_TIMEOUT_SECONDS:-30} + AGENT_BACKEND_STREAM_MAX_RECONNECTS: ${AGENT_BACKEND_STREAM_MAX_RECONNECTS:-3} + AGENT_BACKEND_RUN_TIMEOUT_SECONDS: ${AGENT_BACKEND_RUN_TIMEOUT_SECONDS:-1200} depends_on: init_permissions: condition: service_completed_successfully diff --git a/docker/envs/core-services/dify-agent.env.example b/docker/envs/core-services/dify-agent.env.example index be757e456bb..52ac6bc0990 100644 --- a/docker/envs/core-services/dify-agent.env.example +++ b/docker/envs/core-services/dify-agent.env.example @@ -3,6 +3,9 @@ # ------------------------------ AGENT_BACKEND_BASE_URL=http://agent_backend:5050 +AGENT_BACKEND_STREAM_READ_TIMEOUT_SECONDS=30 +AGENT_BACKEND_STREAM_MAX_RECONNECTS=3 +AGENT_BACKEND_RUN_TIMEOUT_SECONDS=1200 # Leave empty to derive from REDIS_PASSWORD in Docker Compose. DIFY_AGENT_REDIS_URL=