mirror of
https://github.com/langgenius/dify.git
synced 2026-09-24 23:22:26 +08:00
feat(api): abort active workflow runs during Celery warm shutdown (#38220)
This commit is contained in:
@@ -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
|
||||
Reference in New Issue
Block a user