mirror of
https://github.com/langgenius/dify.git
synced 2026-07-21 02:28:30 +08:00
feat(api): abort active workflow runs during Celery warm shutdown (#38220)
This commit is contained in:
parent
c080e2c3b8
commit
5622e8f7ea
@ -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:
|
||||
|
||||
@ -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)
|
||||
|
||||
|
||||
38
api/core/app/apps/workflow/active_workflow_tasks.py
Normal file
38
api/core/app/apps/workflow/active_workflow_tasks.py
Normal 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()
|
||||
@ -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
|
||||
|
||||
@ -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
|
||||
|
||||
|
||||
67
api/core/app/apps/workflow/command_channels.py
Normal file
67
api/core/app/apps/workflow/command_channels.py
Normal 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
|
||||
@ -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,
|
||||
},
|
||||
)
|
||||
|
||||
@ -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
|
||||
|
||||
@ -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."""
|
||||
|
||||
|
||||
79
api/extensions/workflow_warm_shutdown.py
Normal file
79
api/extensions/workflow_warm_shutdown.py
Normal 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,
|
||||
)
|
||||
@ -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."""
|
||||
|
||||
129
api/tests/unit_tests/extensions/test_workflow_warm_shutdown.py
Normal file
129
api/tests/unit_tests/extensions/test_workflow_warm_shutdown.py
Normal 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
|
||||
Loading…
Reference in New Issue
Block a user