feat(api): abort active workflow runs during Celery warm shutdown (#38220)

This commit is contained in:
林玮 (Jade Lin) 2026-07-03 12:51:53 +08:00 committed by GitHub
parent c080e2c3b8
commit 5622e8f7ea
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
14 changed files with 537 additions and 10 deletions

View File

@ -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:

View File

@ -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)

View File

@ -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()

View File

@ -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

View File

@ -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

View File

@ -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

View File

@ -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,
},
)

View File

@ -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

View File

@ -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."""

View File

@ -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,
)

View File

@ -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

View File

@ -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() == []

View File

@ -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."""

View File

@ -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