mirror of
https://github.com/langgenius/dify.git
synced 2026-09-08 11:04:27 +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.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:
|
||||||
|
|||||||
@ -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)
|
||||||
|
|
||||||
|
|||||||
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.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
|
||||||
|
|||||||
@ -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
|
||||||
|
|
||||||
|
|||||||
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 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,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|||||||
@ -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
|
||||||
|
|||||||
@ -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."""
|
||||||
|
|
||||||
|
|||||||
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 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."""
|
||||||
|
|||||||
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