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

This commit is contained in:
林玮 (Jade Lin)
2026-07-03 04:51:53 +00:00
committed by GitHub
parent c080e2c3b8
commit 5622e8f7ea
14 changed files with 537 additions and 10 deletions
@@ -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
@@ -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() == []
@@ -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."""
@@ -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