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.exc import GenerateTaskStoppedError
from core.app.apps.message_based_app_generator import MessageBasedAppGenerator 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.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.app_invoke_entities import AdvancedChatAppGenerateEntity, InvokeFrom
from core.app.entities.task_entities import ( from core.app.entities.task_entities import (
AdvancedChatPausedBlockingResponse, AdvancedChatPausedBlockingResponse,
@ -665,7 +666,8 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator):
) )
try: try:
runner.run() with active_workflow_task(application_generate_entity.task_id):
runner.run()
except GenerateTaskStoppedError: except GenerateTaskStoppedError:
pass pass
except InvokeAuthorizationError: 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.advanced_chat.app_config_manager import AdvancedChatAppConfig
from core.app.apps.base_app_queue_manager import AppQueueManager 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.apps.workflow_app_runner import WorkflowBasedAppRunner
from core.app.entities.app_invoke_entities import ( from core.app.entities.app_invoke_entities import (
AdvancedChatAppGenerateEntity, AdvancedChatAppGenerateEntity,
@ -38,6 +42,7 @@ from core.workflow.workflow_entry import WorkflowEntry
from extensions.ext_database import db from extensions.ext_database import db
from extensions.ext_redis import redis_client from extensions.ext_redis import redis_client
from extensions.otel import WorkflowAppRunnerHandler, trace_span 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.enums import WorkflowType
from graphon.graph_engine.command_channels import RedisChannel from graphon.graph_engine.command_channels import RedisChannel
from graphon.graph_engine.layers import GraphEngineLayer from graphon.graph_engine.layers import GraphEngineLayer
@ -212,7 +217,16 @@ class AdvancedChatAppRunner(WorkflowBasedAppRunner):
# Create Redis command channel for this workflow execution # Create Redis command channel for this workflow execution
task_id = self.application_generate_entity.task_id task_id = self.application_generate_entity.task_id
channel_key = f"workflow:{task_id}:commands" 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( workflow_entry = WorkflowEntry(
tenant_id=self._workflow.tenant_id, tenant_id=self._workflow.tenant_id,
@ -254,7 +268,6 @@ class AdvancedChatAppRunner(WorkflowBasedAppRunner):
workflow_entry.graph_engine.layer(layer) workflow_entry.graph_engine.layer(layer)
generator = workflow_entry.run() generator = workflow_entry.run()
for event in generator: for event in generator:
self._handle_event(workflow_entry, event) 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.base_app_queue_manager import AppQueueManager, PublishFrom
from core.app.apps.draft_variable_saver import DraftVariableSaverFactory from core.app.apps.draft_variable_saver import DraftVariableSaverFactory
from core.app.apps.exc import GenerateTaskStoppedError 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_config_manager import WorkflowAppConfigManager
from core.app.apps.workflow.app_queue_manager import WorkflowAppQueueManager from core.app.apps.workflow.app_queue_manager import WorkflowAppQueueManager
from core.app.apps.workflow.app_runner import WorkflowAppRunner from core.app.apps.workflow.app_runner import WorkflowAppRunner
@ -641,7 +642,8 @@ class WorkflowAppGenerator(BaseAppGenerator):
) )
try: try:
runner.run() with active_workflow_task(application_generate_entity.task_id):
runner.run()
except GenerateTaskStoppedError as e: except GenerateTaskStoppedError as e:
logger.warning("Task stopped: %s", str(e)) logger.warning("Task stopped: %s", str(e))
pass 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.base_app_queue_manager import AppQueueManager
from core.app.apps.workflow.app_config_manager import WorkflowAppConfig 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.apps.workflow_app_runner import WorkflowBasedAppRunner
from core.app.entities.app_invoke_entities import InvokeFrom, WorkflowAppGenerateEntity from core.app.entities.app_invoke_entities import InvokeFrom, WorkflowAppGenerateEntity
from core.app.workflow.layers.persistence import PersistenceWorkflowInfo, WorkflowPersistenceLayer 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 core.workflow.workflow_entry import WorkflowEntry
from extensions.ext_redis import redis_client from extensions.ext_redis import redis_client
from extensions.otel import WorkflowAppRunnerHandler, trace_span 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.enums import WorkflowType
from graphon.graph_engine.command_channels import RedisChannel from graphon.graph_engine.command_channels import RedisChannel
from graphon.graph_engine.layers import GraphEngineLayer from graphon.graph_engine.layers import GraphEngineLayer
@ -146,7 +151,16 @@ class WorkflowAppRunner(WorkflowBasedAppRunner):
# Create Redis command channel for this workflow execution # Create Redis command channel for this workflow execution
task_id = self.application_generate_entity.task_id task_id = self.application_generate_entity.task_id
channel_key = f"workflow:{task_id}:commands" 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 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 typing import cast, final, override
from opentelemetry import context as context_api 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 configs import dify_config
from extensions.otel.parser import ( from extensions.otel.parser import (
@ -24,9 +24,10 @@ from extensions.otel.parser import (
ToolNodeOTelParser, ToolNodeOTelParser,
) )
from extensions.otel.runtime import is_instrument_flag_enabled from extensions.otel.runtime import is_instrument_flag_enabled
from extensions.otel.semconv import DifySpanAttributes
from graphon.enums import BuiltinNodeTypes, NodeType from graphon.enums import BuiltinNodeTypes, NodeType
from graphon.graph_engine.layers import GraphEngineLayer 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 from graphon.nodes.base.node import Node
logger = logging.getLogger(__name__) 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) logger.warning("Failed to end OpenTelemetry span for node %s: %s", node.id, e)
@override @override
def on_event(self, event) -> None: def on_event(self, event: GraphEngineEvent) -> None:
"""Not used in this layer.""" """Record graph-level observability events."""
pass if self._is_disabled:
return
if isinstance(event, GraphRunAbortedEvent):
self._record_abort_reason(reason=event.reason or "Workflow execution aborted")
@override @override
def on_graph_end(self, error: Exception | None) -> None: def on_graph_end(self, error: Exception | None) -> None:
@ -171,3 +176,16 @@ class ObservabilityLayer(GraphEngineLayer):
len(self._node_contexts), len(self._node_contexts),
) )
self._node_contexts.clear() 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 configs import dify_config
from dify_app import DifyApp from dify_app import DifyApp
from extensions.redis_names import normalize_redis_key_prefix from extensions.redis_names import normalize_redis_key_prefix
from extensions.workflow_warm_shutdown import setup_workflow_warm_shutdown_handler
class _CelerySentinelKwargsDict(TypedDict): class _CelerySentinelKwargsDict(TypedDict):
@ -147,6 +148,7 @@ def init_app(app: DifyApp) -> Celery:
celery_app.set_default() celery_app.set_default()
app.extensions["celery"] = celery_app app.extensions["celery"] = celery_app
setup_workflow_warm_shutdown_handler()
imports = [ imports = [
"tasks.async_workflow_tasks", # trigger workers "tasks.async_workflow_tasks", # trigger workers

View File

@ -19,6 +19,9 @@ class DifySpanAttributes:
WORKFLOW_ID = "dify.workflow_id" WORKFLOW_ID = "dify.workflow_id"
"""Workflow identifier.""" """Workflow identifier."""
WORKFLOW_ABORT_REASON = "dify.workflow.abort.reason"
"""Reason recorded when a workflow run is aborted."""
INVOKE_FROM = "dify.invoke_from" INVOKE_FROM = "dify.invoke_from"
"""Invocation source, e.g. SERVICE_API, WEB_APP, DEBUGGER.""" """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 opentelemetry.trace import StatusCode
from core.app.workflow.layers.observability import ObservabilityLayer from core.app.workflow.layers.observability import ObservabilityLayer
from extensions.otel.semconv import DifySpanAttributes
from graphon.enums import BuiltinNodeTypes from graphon.enums import BuiltinNodeTypes
from graphon.graph_events import GraphRunAbortedEvent
class TestObservabilityLayerInitialization: class TestObservabilityLayerInitialization:
@ -281,6 +283,27 @@ class TestObservabilityLayerGraphLifecycle:
assert len(layer._node_contexts) == 0 assert len(layer._node_contexts) == 0
assert "node spans were not properly ended" in caplog.text 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: class TestObservabilityLayerDisabledMode:
"""Test behavior when layer is disabled.""" """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