diff --git a/api/core/app/apps/advanced_chat/app_generator.py b/api/core/app/apps/advanced_chat/app_generator.py index ee7ade9e45e..f52fd1046f8 100644 --- a/api/core/app/apps/advanced_chat/app_generator.py +++ b/api/core/app/apps/advanced_chat/app_generator.py @@ -33,6 +33,7 @@ from core.app.apps.draft_variable_saver import DraftVariableSaverFactory from core.app.apps.exc import GenerateTaskStoppedError from core.app.apps.message_based_app_generator import MessageBasedAppGenerator from core.app.apps.message_based_app_queue_manager import MessageBasedAppQueueManager +from core.app.apps.workflow.active_workflow_tasks import active_workflow_task from core.app.entities.app_invoke_entities import AdvancedChatAppGenerateEntity, InvokeFrom from core.app.entities.task_entities import ( AdvancedChatPausedBlockingResponse, @@ -665,7 +666,8 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator): ) try: - runner.run() + with active_workflow_task(application_generate_entity.task_id): + runner.run() except GenerateTaskStoppedError: pass except InvokeAuthorizationError: diff --git a/api/core/app/apps/advanced_chat/app_runner.py b/api/core/app/apps/advanced_chat/app_runner.py index 67397965384..b78a3b5b3dc 100644 --- a/api/core/app/apps/advanced_chat/app_runner.py +++ b/api/core/app/apps/advanced_chat/app_runner.py @@ -8,6 +8,10 @@ from sqlalchemy.orm import Session 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 ( + CelerySignalCommandChannel, + CombinedCommandChannel, +) from core.app.apps.workflow_app_runner import WorkflowBasedAppRunner from core.app.entities.app_invoke_entities import ( AdvancedChatAppGenerateEntity, @@ -38,6 +42,7 @@ from core.workflow.workflow_entry import WorkflowEntry from extensions.ext_database import db 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 graphon.enums import WorkflowType from graphon.graph_engine.command_channels import RedisChannel from graphon.graph_engine.layers import GraphEngineLayer @@ -212,7 +217,16 @@ class AdvancedChatAppRunner(WorkflowBasedAppRunner): # Create Redis command channel for this workflow execution task_id = self.application_generate_entity.task_id channel_key = f"workflow:{task_id}:commands" - command_channel = RedisChannel(redis_client, channel_key) + celery_signal_channel = CelerySignalCommandChannel( + shutdown_state_getter=celery_warm_shutdown_started, + abort_reason=WORKFLOW_WARM_SHUTDOWN_ABORT_REASON, + ) + command_channel = CombinedCommandChannel( + ( + RedisChannel(redis_client, channel_key), + celery_signal_channel, + ) + ) workflow_entry = WorkflowEntry( tenant_id=self._workflow.tenant_id, @@ -254,7 +268,6 @@ class AdvancedChatAppRunner(WorkflowBasedAppRunner): workflow_entry.graph_engine.layer(layer) generator = workflow_entry.run() - for event in generator: self._handle_event(workflow_entry, event) diff --git a/api/core/app/apps/workflow/active_workflow_tasks.py b/api/core/app/apps/workflow/active_workflow_tasks.py new file mode 100644 index 00000000000..4aa23ad1ccf --- /dev/null +++ b/api/core/app/apps/workflow/active_workflow_tasks.py @@ -0,0 +1,38 @@ +"""In-process registry for workflow application task IDs.""" + +import threading +from collections.abc import Iterator +from contextlib import contextmanager + +_active_task_ids: set[str] = set() +_active_task_ids_lock = threading.RLock() + + +@contextmanager +def active_workflow_task(task_id: str) -> Iterator[None]: + """Register a workflow application task ID for the duration of a workflow run.""" + if not task_id: + raise ValueError("task_id must not be empty") + + with _active_task_ids_lock: + if task_id in _active_task_ids: + raise ValueError(f"Workflow task already active for task_id={task_id}") + _active_task_ids.add(task_id) + + try: + yield + finally: + with _active_task_ids_lock: + _active_task_ids.discard(task_id) + + +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) + + +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() diff --git a/api/core/app/apps/workflow/app_generator.py b/api/core/app/apps/workflow/app_generator.py index cdc65c6e415..ab07454ff5b 100644 --- a/api/core/app/apps/workflow/app_generator.py +++ b/api/core/app/apps/workflow/app_generator.py @@ -19,6 +19,7 @@ from core.app.apps.base_app_generator import BaseAppGenerator from core.app.apps.base_app_queue_manager import AppQueueManager, PublishFrom from core.app.apps.draft_variable_saver import DraftVariableSaverFactory from core.app.apps.exc import GenerateTaskStoppedError +from core.app.apps.workflow.active_workflow_tasks import active_workflow_task from core.app.apps.workflow.app_config_manager import WorkflowAppConfigManager from core.app.apps.workflow.app_queue_manager import WorkflowAppQueueManager from core.app.apps.workflow.app_runner import WorkflowAppRunner @@ -641,7 +642,8 @@ class WorkflowAppGenerator(BaseAppGenerator): ) try: - runner.run() + with active_workflow_task(application_generate_entity.task_id): + runner.run() except GenerateTaskStoppedError as e: logger.warning("Task stopped: %s", str(e)) pass diff --git a/api/core/app/apps/workflow/app_runner.py b/api/core/app/apps/workflow/app_runner.py index e35735038c3..6682a395a8c 100644 --- a/api/core/app/apps/workflow/app_runner.py +++ b/api/core/app/apps/workflow/app_runner.py @@ -5,6 +5,10 @@ from typing import cast 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 ( + CelerySignalCommandChannel, + CombinedCommandChannel, +) from core.app.apps.workflow_app_runner import WorkflowBasedAppRunner from core.app.entities.app_invoke_entities import InvokeFrom, WorkflowAppGenerateEntity from core.app.workflow.layers.persistence import PersistenceWorkflowInfo, WorkflowPersistenceLayer @@ -17,6 +21,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 graphon.enums import WorkflowType from graphon.graph_engine.command_channels import RedisChannel from graphon.graph_engine.layers import GraphEngineLayer @@ -146,7 +151,16 @@ class WorkflowAppRunner(WorkflowBasedAppRunner): # Create Redis command channel for this workflow execution task_id = self.application_generate_entity.task_id channel_key = f"workflow:{task_id}:commands" - command_channel = RedisChannel(redis_client, channel_key) + celery_signal_channel = CelerySignalCommandChannel( + shutdown_state_getter=celery_warm_shutdown_started, + abort_reason=WORKFLOW_WARM_SHUTDOWN_ABORT_REASON, + ) + command_channel = CombinedCommandChannel( + ( + RedisChannel(redis_client, channel_key), + celery_signal_channel, + ) + ) self._queue_manager.graph_runtime_state = graph_runtime_state diff --git a/api/core/app/apps/workflow/command_channels.py b/api/core/app/apps/workflow/command_channels.py new file mode 100644 index 00000000000..526476f8c24 --- /dev/null +++ b/api/core/app/apps/workflow/command_channels.py @@ -0,0 +1,67 @@ +"""Command channels used by Dify workflow runners.""" + +import logging +from collections.abc import Callable, Sequence +from typing import final, override + +from graphon.graph_engine.command_channels import CommandChannel +from graphon.graph_engine.entities.commands import AbortCommand, GraphEngineCommand + +logger = logging.getLogger(__name__) +ShutdownStateGetter = Callable[[], bool] + + +@final +class CombinedCommandChannel: + """Fetch commands from all sources and send outbound commands through the primary source.""" + + _command_channels: tuple[CommandChannel, ...] + + def __init__(self, command_channels: Sequence[CommandChannel]) -> None: + if not command_channels: + raise ValueError("command_channels must not be empty") + self._command_channels = tuple(command_channels) + + def fetch_commands(self) -> list[GraphEngineCommand]: + commands: list[GraphEngineCommand] = [] + for channel in self._command_channels: + try: + commands.extend(channel.fetch_commands()) + except Exception: + logger.exception("Failed to fetch GraphEngine commands from %s", channel.__class__.__name__) + return commands + + def send_command(self, command: GraphEngineCommand) -> None: + """Send commands through the first channel, which is the runner's primary command sink.""" + self._command_channels[0].send_command(command) + + +@final +class CelerySignalCommandChannel(CommandChannel): + """Translate process-local Celery shutdown state into one GraphEngine abort command.""" + + _shutdown_state_getter: ShutdownStateGetter + _abort_reason: str + _abort_emitted: bool + + def __init__( + self, + *, + shutdown_state_getter: ShutdownStateGetter, + abort_reason: str, + ) -> None: + self._shutdown_state_getter = shutdown_state_getter + self._abort_reason = abort_reason + self._abort_emitted = False + + @override + def fetch_commands(self) -> list[GraphEngineCommand]: + if self._abort_emitted or not self._shutdown_state_getter(): + return [] + + self._abort_emitted = True + return [AbortCommand(reason=self._abort_reason)] + + @override + def send_command(self, command: GraphEngineCommand) -> None: + _ = command diff --git a/api/core/app/workflow/layers/observability.py b/api/core/app/workflow/layers/observability.py index 8b5a5b9d7ff..317dab91ad3 100644 --- a/api/core/app/workflow/layers/observability.py +++ b/api/core/app/workflow/layers/observability.py @@ -13,7 +13,7 @@ from dataclasses import dataclass from typing import cast, final, override from opentelemetry import context as context_api -from opentelemetry.trace import Span, SpanKind, Tracer, get_tracer, set_span_in_context +from opentelemetry.trace import Span, SpanKind, Tracer, get_current_span, get_tracer, set_span_in_context from configs import dify_config from extensions.otel.parser import ( @@ -24,9 +24,10 @@ from extensions.otel.parser import ( ToolNodeOTelParser, ) from extensions.otel.runtime import is_instrument_flag_enabled +from extensions.otel.semconv import DifySpanAttributes from graphon.enums import BuiltinNodeTypes, NodeType from graphon.graph_engine.layers import GraphEngineLayer -from graphon.graph_events import GraphNodeEventBase +from graphon.graph_events import GraphEngineEvent, GraphNodeEventBase, GraphRunAbortedEvent from graphon.nodes.base.node import Node logger = logging.getLogger(__name__) @@ -158,9 +159,13 @@ class ObservabilityLayer(GraphEngineLayer): logger.warning("Failed to end OpenTelemetry span for node %s: %s", node.id, e) @override - def on_event(self, event) -> None: - """Not used in this layer.""" - pass + def on_event(self, event: GraphEngineEvent) -> None: + """Record graph-level observability events.""" + if self._is_disabled: + return + + if isinstance(event, GraphRunAbortedEvent): + self._record_abort_reason(reason=event.reason or "Workflow execution aborted") @override def on_graph_end(self, error: Exception | None) -> None: @@ -171,3 +176,16 @@ class ObservabilityLayer(GraphEngineLayer): len(self._node_contexts), ) self._node_contexts.clear() + + def _record_abort_reason(self, *, reason: str) -> None: + span = get_current_span() + if not span.is_recording(): + return + + span.set_attribute(DifySpanAttributes.WORKFLOW_ABORT_REASON, reason) + span.add_event( + "dify.workflow.aborted", + attributes={ + DifySpanAttributes.WORKFLOW_ABORT_REASON: reason, + }, + ) diff --git a/api/extensions/ext_celery.py b/api/extensions/ext_celery.py index 42c83b30f2b..810f5f17bc7 100644 --- a/api/extensions/ext_celery.py +++ b/api/extensions/ext_celery.py @@ -10,6 +10,7 @@ from typing_extensions import TypedDict from configs import dify_config from dify_app import DifyApp from extensions.redis_names import normalize_redis_key_prefix +from extensions.workflow_warm_shutdown import setup_workflow_warm_shutdown_handler class _CelerySentinelKwargsDict(TypedDict): @@ -147,6 +148,7 @@ def init_app(app: DifyApp) -> Celery: celery_app.set_default() app.extensions["celery"] = celery_app + setup_workflow_warm_shutdown_handler() imports = [ "tasks.async_workflow_tasks", # trigger workers diff --git a/api/extensions/otel/semconv/dify.py b/api/extensions/otel/semconv/dify.py index 301ddd11aaa..80dbc97387f 100644 --- a/api/extensions/otel/semconv/dify.py +++ b/api/extensions/otel/semconv/dify.py @@ -19,6 +19,9 @@ class DifySpanAttributes: WORKFLOW_ID = "dify.workflow_id" """Workflow identifier.""" + WORKFLOW_ABORT_REASON = "dify.workflow.abort.reason" + """Reason recorded when a workflow run is aborted.""" + INVOKE_FROM = "dify.invoke_from" """Invocation source, e.g. SERVICE_API, WEB_APP, DEBUGGER.""" diff --git a/api/extensions/workflow_warm_shutdown.py b/api/extensions/workflow_warm_shutdown.py new file mode 100644 index 00000000000..e6579ccdb5d --- /dev/null +++ b/api/extensions/workflow_warm_shutdown.py @@ -0,0 +1,79 @@ +"""Abort active workflow runs during Celery warm shutdown.""" + +import logging +import threading +from typing import Any + +from celery.signals import worker_shutdown, worker_shutting_down + +from core.app.apps.workflow.active_workflow_tasks import ( + get_active_workflow_task_count, + reset_active_workflow_tasks, +) + +logger = logging.getLogger(__name__) +WORKFLOW_WARM_SHUTDOWN_ABORT_REASON = "Workflow stopped because the worker is shutting down." +_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() + + +def _is_warm_shutdown(how: Any) -> bool: + return str(how).strip().lower() == "warm" + + +def celery_warm_shutdown_started() -> bool: + """Return whether the current worker process started Celery warm shutdown.""" + return _celery_warm_shutdown_started.is_set() + + +def mark_celery_warm_shutdown_started() -> None: + """Mark the current worker process as being in Celery warm shutdown.""" + _celery_warm_shutdown_started.set() + + +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) + 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, + ) + + +def _on_worker_shutdown(*args: object, **kwargs: object) -> None: + """Log whether tracked workflow tasks ended before Celery worker shutdown.""" + remaining_run_count = get_active_workflow_task_count() + if remaining_run_count: + logger.warning( + "Celery worker is shutting down with %s workflow run(s) still active after warm shutdown wait", + remaining_run_count, + ) + return + + logger.info("Celery worker shutdown reached after all tracked workflow runs ended") + + +def setup_workflow_warm_shutdown_handler() -> None: + """Connect Celery worker shutdown handlers for workflow abort and logging.""" + reset_active_workflow_tasks() + worker_shutting_down.connect( + _on_worker_shutting_down, + weak=False, + dispatch_uid=_WORKER_SHUTTING_DOWN_DISPATCH_UID, + ) + worker_shutdown.connect( + _on_worker_shutdown, + weak=False, + dispatch_uid=_WORKER_SHUTDOWN_DISPATCH_UID, + ) diff --git a/api/tests/unit_tests/core/app/apps/workflow/test_active_workflow_tasks.py b/api/tests/unit_tests/core/app/apps/workflow/test_active_workflow_tasks.py new file mode 100644 index 00000000000..c50b16533ff --- /dev/null +++ b/api/tests/unit_tests/core/app/apps/workflow/test_active_workflow_tasks.py @@ -0,0 +1,30 @@ +import pytest + +from core.app.apps.workflow.active_workflow_tasks import ( + active_workflow_task, + get_active_workflow_task_count, + reset_active_workflow_tasks, +) + + +@pytest.fixture(autouse=True) +def reset_active_tasks() -> None: + reset_active_workflow_tasks() + yield + reset_active_workflow_tasks() + + +def test_active_workflow_task_tracks_count_during_context() -> None: + assert get_active_workflow_task_count() == 0 + + with active_workflow_task("task-a"): + assert get_active_workflow_task_count() == 1 + + assert get_active_workflow_task_count() == 0 + + +def test_active_workflow_task_rejects_duplicate_task_id() -> None: + with active_workflow_task("task-a"): + with pytest.raises(ValueError, match="already active"): + with active_workflow_task("task-a"): + pass diff --git a/api/tests/unit_tests/core/app/apps/workflow/test_command_channels.py b/api/tests/unit_tests/core/app/apps/workflow/test_command_channels.py new file mode 100644 index 00000000000..1a7e7700358 --- /dev/null +++ b/api/tests/unit_tests/core/app/apps/workflow/test_command_channels.py @@ -0,0 +1,107 @@ +from __future__ import annotations + +from types import SimpleNamespace + +import pytest + +from core.app.apps.workflow.command_channels import ( + CelerySignalCommandChannel, + CombinedCommandChannel, +) +from graphon.graph_engine.entities.commands import AbortCommand, PauseCommand + + +class _CommandChannelStub: + def __init__(self, commands=None) -> None: + self.commands = list(commands or []) + self.sent = [] + + def fetch_commands(self): + commands = self.commands + self.commands = [] + return commands + + def send_command(self, command) -> None: + self.sent.append(command) + + +def test_combined_command_channel_fetches_from_all_sources() -> None: + abort = AbortCommand(reason="stop") + pause = PauseCommand(reason="pause") + combined = CombinedCommandChannel( + ( + _CommandChannelStub([abort]), + _CommandChannelStub([pause]), + ) + ) + + assert combined.fetch_commands() == [abort, pause] + + +def test_combined_command_channel_sends_to_primary_source() -> None: + primary = _CommandChannelStub() + secondary = _CommandChannelStub() + combined = CombinedCommandChannel((primary, secondary)) + command = AbortCommand(reason="stop") + + combined.send_command(command) + + assert primary.sent == [command] + assert secondary.sent == [] + + +def test_combined_command_channel_requires_at_least_one_source() -> None: + with pytest.raises(ValueError, match="command_channels must not be empty"): + CombinedCommandChannel(()) + + +def test_combined_command_channel_continues_after_source_failure(caplog: pytest.LogCaptureFixture) -> None: + abort = AbortCommand(reason="stop") + failing = SimpleNamespace( + fetch_commands=lambda: (_ for _ in ()).throw(RuntimeError("boom")), + send_command=lambda _command: None, + ) + combined = CombinedCommandChannel((failing, _CommandChannelStub([abort]))) + + assert combined.fetch_commands() == [abort] + assert "Failed to fetch GraphEngine commands" in caplog.text + + +def test_celery_signal_command_channel_emits_abort_when_shutdown_starts() -> None: + shutdown_started = False + channel = CelerySignalCommandChannel( + shutdown_state_getter=lambda: shutdown_started, + abort_reason="worker shutdown", + ) + + assert channel.fetch_commands() == [] + + shutdown_started = True + + commands = channel.fetch_commands() + + assert len(commands) == 1 + assert isinstance(commands[0], AbortCommand) + assert commands[0].reason == "worker shutdown" + + +def test_celery_signal_command_channel_emits_abort_once_per_instance() -> None: + channel = CelerySignalCommandChannel( + shutdown_state_getter=lambda: True, + abort_reason="worker shutdown", + ) + + assert len(channel.fetch_commands()) == 1 + assert channel.fetch_commands() == [] + + +def test_celery_signal_command_channel_send_command_is_noop() -> None: + channel = CelerySignalCommandChannel( + shutdown_state_getter=lambda: False, + abort_reason="worker shutdown", + ) + command = PauseCommand(reason="pause") + + channel.send_command(command) + + assert channel.fetch_commands() == [] diff --git a/api/tests/unit_tests/core/workflow/graph_engine/layers/test_observability.py b/api/tests/unit_tests/core/workflow/graph_engine/layers/test_observability.py index 919f15efd09..f3903e2e438 100644 --- a/api/tests/unit_tests/core/workflow/graph_engine/layers/test_observability.py +++ b/api/tests/unit_tests/core/workflow/graph_engine/layers/test_observability.py @@ -16,7 +16,9 @@ import pytest from opentelemetry.trace import StatusCode from core.app.workflow.layers.observability import ObservabilityLayer +from extensions.otel.semconv import DifySpanAttributes from graphon.enums import BuiltinNodeTypes +from graphon.graph_events import GraphRunAbortedEvent class TestObservabilityLayerInitialization: @@ -281,6 +283,27 @@ class TestObservabilityLayerGraphLifecycle: assert len(layer._node_contexts) == 0 assert "node spans were not properly ended" in caplog.text + @patch("core.app.workflow.layers.observability.dify_config.ENABLE_OTEL", True) + @pytest.mark.usefixtures("mock_is_instrument_flag_enabled_false") + def test_graph_aborted_event_records_reason_on_current_span( + self, tracer_provider_with_memory_exporter, memory_span_exporter, mock_start_node + ): + layer = ObservabilityLayer() + layer.on_graph_start() + layer.on_node_run_start(mock_start_node) + + layer.on_event(GraphRunAbortedEvent(reason="worker shutdown", outputs={})) + layer.on_node_run_end(mock_start_node, None) + + spans = memory_span_exporter.get_finished_spans() + assert len(spans) == 1 + assert spans[0].attributes[DifySpanAttributes.WORKFLOW_ABORT_REASON] == "worker shutdown" + assert any( + event.name == "dify.workflow.aborted" + and event.attributes[DifySpanAttributes.WORKFLOW_ABORT_REASON] == "worker shutdown" + for event in spans[0].events + ) + class TestObservabilityLayerDisabledMode: """Test behavior when layer is disabled.""" diff --git a/api/tests/unit_tests/extensions/test_workflow_warm_shutdown.py b/api/tests/unit_tests/extensions/test_workflow_warm_shutdown.py new file mode 100644 index 00000000000..ce774e4d160 --- /dev/null +++ b/api/tests/unit_tests/extensions/test_workflow_warm_shutdown.py @@ -0,0 +1,129 @@ +import logging +from unittest.mock import MagicMock + +import pytest + +from core.app.apps.workflow.active_workflow_tasks import reset_active_workflow_tasks +from core.app.apps.workflow.command_channels import CelerySignalCommandChannel +from extensions import workflow_warm_shutdown +from graphon.graph_engine.entities.commands import AbortCommand + + +@pytest.fixture(autouse=True) +def reset_warm_shutdown_state() -> None: + reset_active_workflow_tasks() + workflow_warm_shutdown._celery_warm_shutdown_started.clear() + yield + reset_active_workflow_tasks() + workflow_warm_shutdown._celery_warm_shutdown_started.clear() + + +def _create_warm_shutdown_command_channel() -> CelerySignalCommandChannel: + return CelerySignalCommandChannel( + shutdown_state_getter=workflow_warm_shutdown.celery_warm_shutdown_started, + abort_reason=workflow_warm_shutdown.WORKFLOW_WARM_SHUTDOWN_ABORT_REASON, + ) + + +def test_worker_shutting_down_skips_non_warm_shutdown(monkeypatch: pytest.MonkeyPatch) -> None: + mark_shutdown = MagicMock() + monkeypatch.setattr(workflow_warm_shutdown, "mark_celery_warm_shutdown_started", mark_shutdown) + + workflow_warm_shutdown._on_worker_shutting_down(how="cold") + + mark_shutdown.assert_not_called() + + +def test_worker_shutting_down_marks_warm_shutdown(monkeypatch: pytest.MonkeyPatch) -> None: + mark_shutdown = MagicMock() + monkeypatch.setattr(workflow_warm_shutdown, "mark_celery_warm_shutdown_started", mark_shutdown) + monkeypatch.setattr(workflow_warm_shutdown, "get_active_workflow_task_count", lambda: 2) + + workflow_warm_shutdown._on_worker_shutting_down(how="warm") + + mark_shutdown.assert_called_once_with() + + +def test_warm_shutdown_state_tracks_started_flag() -> None: + assert workflow_warm_shutdown.celery_warm_shutdown_started() is False + + workflow_warm_shutdown.mark_celery_warm_shutdown_started() + + assert workflow_warm_shutdown.celery_warm_shutdown_started() is True + + +def test_setup_configures_warm_shutdown_command_channel(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(workflow_warm_shutdown.worker_shutting_down, "connect", MagicMock()) + monkeypatch.setattr(workflow_warm_shutdown.worker_shutdown, "connect", MagicMock()) + + workflow_warm_shutdown.setup_workflow_warm_shutdown_handler() + workflow_warm_shutdown.mark_celery_warm_shutdown_started() + + commands = _create_warm_shutdown_command_channel().fetch_commands() + + assert len(commands) == 1 + assert isinstance(commands[0], AbortCommand) + assert commands[0].reason == workflow_warm_shutdown.WORKFLOW_WARM_SHUTDOWN_ABORT_REASON + + +def test_warm_shutdown_command_stays_available_for_late_channels(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(workflow_warm_shutdown.worker_shutting_down, "connect", MagicMock()) + monkeypatch.setattr(workflow_warm_shutdown.worker_shutdown, "connect", MagicMock()) + + workflow_warm_shutdown.setup_workflow_warm_shutdown_handler() + workflow_warm_shutdown.mark_celery_warm_shutdown_started() + + first_channel = _create_warm_shutdown_command_channel() + late_channel = _create_warm_shutdown_command_channel() + + assert len(first_channel.fetch_commands()) == 1 + assert len(late_channel.fetch_commands()) == 1 + + +def test_worker_shutdown_logs_when_all_workflow_runs_ended( + monkeypatch: pytest.MonkeyPatch, + caplog: pytest.LogCaptureFixture, +) -> None: + caplog.set_level(logging.INFO, logger=workflow_warm_shutdown.logger.name) + monkeypatch.setattr(workflow_warm_shutdown, "get_active_workflow_task_count", lambda: 0) + + workflow_warm_shutdown._on_worker_shutdown() + + assert "after all tracked workflow runs ended" in caplog.text + + +def test_worker_shutdown_logs_remaining_workflow_runs( + monkeypatch: pytest.MonkeyPatch, + caplog: pytest.LogCaptureFixture, +) -> None: + caplog.set_level(logging.INFO, logger=workflow_warm_shutdown.logger.name) + monkeypatch.setattr(workflow_warm_shutdown, "get_active_workflow_task_count", lambda: 2) + + workflow_warm_shutdown._on_worker_shutdown() + + assert "with 2 workflow run(s) still active after warm shutdown wait" in caplog.text + + +def test_setup_connects_shutdown_handlers(monkeypatch: pytest.MonkeyPatch) -> None: + connect_shutting_down = MagicMock() + connect_shutdown = MagicMock() + monkeypatch.setattr(workflow_warm_shutdown.worker_shutting_down, "connect", connect_shutting_down) + monkeypatch.setattr(workflow_warm_shutdown.worker_shutdown, "connect", connect_shutdown) + + workflow_warm_shutdown.setup_workflow_warm_shutdown_handler() + + connect_shutting_down.assert_called_once() + connect_shutdown.assert_called_once() + + +def test_setup_preserves_warm_shutdown_state(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(workflow_warm_shutdown.worker_shutting_down, "connect", MagicMock()) + monkeypatch.setattr(workflow_warm_shutdown.worker_shutdown, "connect", MagicMock()) + + workflow_warm_shutdown.mark_celery_warm_shutdown_started() + workflow_warm_shutdown.setup_workflow_warm_shutdown_handler() + + commands = _create_warm_shutdown_command_channel().fetch_commands() + + assert workflow_warm_shutdown.celery_warm_shutdown_started() is True + assert len(commands) == 1