mirror of
https://github.com/langgenius/dify.git
synced 2026-09-24 23:22:26 +08:00
fix: preserve ResponseStreamFilter state across workflow pause/resume (#38540)
Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
This commit is contained in:
co-authored by
autofix-ci[bot]
parent
2b35f48d77
commit
d72ee32ba1
+216
@@ -0,0 +1,216 @@
|
||||
"""Regression test: if-else branch + human_input pause + downstream answer nodes.
|
||||
|
||||
Reproduces https://github.com/langgenius/dify/issues/38525 at the
|
||||
iter_dify_graph_engine_events layer: without a restored ResponseStreamFilter,
|
||||
answer nodes downstream of a pre-pause branch never unlock for streaming on
|
||||
resume, even though the graph executes correctly.
|
||||
"""
|
||||
|
||||
from datetime import timedelta
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
from core.app.entities.app_invoke_entities import InvokeFrom, UserFrom
|
||||
from core.repositories.human_input_repository import HumanInputFormEntity, HumanInputFormRepository
|
||||
from core.workflow.nodes.human_input.callback import DifyHITLCallback
|
||||
from core.workflow.nodes.human_input.entities import HumanInputNodeData, UserActionConfig
|
||||
from core.workflow.nodes.human_input.enums import HumanInputFormStatus
|
||||
from core.workflow.system_variables import build_system_variables
|
||||
from core.workflow.workflow_entry import iter_dify_graph_engine_events
|
||||
from graphon.filters import GraphEventFilterContext, ResponseStreamFilter, filter_graph_events
|
||||
from graphon.graph import Graph
|
||||
from graphon.graph_engine import GraphEngine, GraphEngineConfig
|
||||
from graphon.graph_engine.command_channels import InMemoryChannel
|
||||
from graphon.graph_events import GraphRunPausedEvent, GraphRunSucceededEvent, NodeRunStreamChunkEvent
|
||||
from graphon.nodes.answer.answer_node import AnswerNode
|
||||
from graphon.nodes.answer.entities import AnswerNodeData
|
||||
from graphon.nodes.human_input.human_input_node import HumanInputNode
|
||||
from graphon.nodes.if_else.entities import IfElseNodeData
|
||||
from graphon.nodes.if_else.if_else_node import IfElseNode
|
||||
from graphon.nodes.start.entities import StartNodeData
|
||||
from graphon.nodes.start.start_node import StartNode
|
||||
from graphon.runtime import GraphRuntimeState, VariablePool
|
||||
from graphon.utils.condition.entities import Condition
|
||||
from libs.datetime_utils import naive_utc_now
|
||||
from tests.workflow_test_utils import build_test_graph_init_params
|
||||
|
||||
WORKFLOW_EXECUTION_ID = "wf-exec-38525"
|
||||
|
||||
|
||||
def _mock_repo_paused() -> HumanInputFormRepository:
|
||||
repo = MagicMock(spec=HumanInputFormRepository)
|
||||
form = MagicMock(spec=HumanInputFormEntity)
|
||||
form.id = "form-1"
|
||||
form.submission_token = "token-1"
|
||||
form.recipients = []
|
||||
form.rendered_content = "rendered"
|
||||
form.submitted = False
|
||||
repo.create_form.return_value = form
|
||||
repo.get_form.return_value = None
|
||||
return repo
|
||||
|
||||
|
||||
def _mock_repo_resumed(action_id: str = "continue") -> HumanInputFormRepository:
|
||||
repo = MagicMock(spec=HumanInputFormRepository)
|
||||
form = MagicMock(spec=HumanInputFormEntity)
|
||||
form.id = "form-1"
|
||||
form.submission_token = "token-1"
|
||||
form.recipients = []
|
||||
form.rendered_content = "rendered"
|
||||
form.submitted = True
|
||||
form.selected_action_id = action_id
|
||||
form.submitted_data = {}
|
||||
form.status = HumanInputFormStatus.WAITING
|
||||
form.expiration_time = naive_utc_now() + timedelta(hours=1)
|
||||
repo.get_form.return_value = form
|
||||
return repo
|
||||
|
||||
|
||||
def _build_graph(runtime_state: GraphRuntimeState, form_repository: HumanInputFormRepository) -> Graph:
|
||||
params = build_test_graph_init_params(
|
||||
workflow_id="wf",
|
||||
graph_config={"nodes": [], "edges": []},
|
||||
user_from=UserFrom.ACCOUNT,
|
||||
invoke_from=InvokeFrom.DEBUGGER,
|
||||
)
|
||||
|
||||
start_node = StartNode(
|
||||
node_id="start",
|
||||
data=StartNodeData(title="start", variables=[]),
|
||||
graph_init_params=params,
|
||||
graph_runtime_state=runtime_state,
|
||||
)
|
||||
|
||||
if_else_node = IfElseNode(
|
||||
node_id="if_else",
|
||||
data=IfElseNodeData(
|
||||
title="if-else",
|
||||
cases=[
|
||||
IfElseNodeData.Case(
|
||||
case_id="true",
|
||||
logical_operator="and",
|
||||
conditions=[
|
||||
Condition(
|
||||
variable_selector=["start", "category"],
|
||||
comparison_operator="is",
|
||||
value="fruit",
|
||||
)
|
||||
],
|
||||
)
|
||||
],
|
||||
),
|
||||
graph_init_params=params,
|
||||
graph_runtime_state=runtime_state,
|
||||
)
|
||||
|
||||
human_data = HumanInputNodeData(
|
||||
title="human",
|
||||
form_content="Awaiting human input",
|
||||
inputs=[],
|
||||
user_actions=[UserActionConfig(id="continue", title="Continue")],
|
||||
)
|
||||
human_node = HumanInputNode(
|
||||
node_id="human_input",
|
||||
data=human_data,
|
||||
graph_init_params=params,
|
||||
graph_runtime_state=runtime_state,
|
||||
hitl_callback=DifyHITLCallback(form_repository=form_repository, node_data=human_data),
|
||||
)
|
||||
|
||||
answer_false_node = AnswerNode(
|
||||
node_id="answer_false",
|
||||
data=AnswerNodeData(title="answer_false", answer="unreachable branch"),
|
||||
graph_init_params=params,
|
||||
graph_runtime_state=runtime_state,
|
||||
)
|
||||
|
||||
answer_after_pause = AnswerNode(
|
||||
node_id="answer_after_pause",
|
||||
data=AnswerNodeData(title="answer_after_pause", answer="Post-branch answer chunk 1"),
|
||||
graph_init_params=params,
|
||||
graph_runtime_state=runtime_state,
|
||||
)
|
||||
|
||||
answer_after_pause_2 = AnswerNode(
|
||||
node_id="answer_after_pause_2",
|
||||
data=AnswerNodeData(title="answer_after_pause_2", answer="Post-branch answer chunk 2"),
|
||||
graph_init_params=params,
|
||||
graph_runtime_state=runtime_state,
|
||||
)
|
||||
|
||||
return (
|
||||
Graph.new()
|
||||
.add_root(start_node)
|
||||
.add_node(if_else_node, from_node_id="start")
|
||||
.add_node(human_node, from_node_id="if_else", source_handle="true")
|
||||
.add_node(answer_false_node, from_node_id="if_else", source_handle="false")
|
||||
.add_node(answer_after_pause, from_node_id="human_input", source_handle="continue")
|
||||
.add_node(answer_after_pause_2, from_node_id="answer_after_pause")
|
||||
.build()
|
||||
)
|
||||
|
||||
|
||||
def _build_runtime_state() -> GraphRuntimeState:
|
||||
variable_pool = VariablePool.from_bootstrap(
|
||||
system_variables=build_system_variables(
|
||||
workflow_execution_id=WORKFLOW_EXECUTION_ID,
|
||||
app_id="app",
|
||||
workflow_id="wf",
|
||||
user_id="user",
|
||||
),
|
||||
user_inputs={},
|
||||
conversation_variables=[],
|
||||
)
|
||||
variable_pool.add(("start", "category"), "fruit") # drives the if-else "true" branch
|
||||
return GraphRuntimeState(variable_pool=variable_pool, start_at=0.0)
|
||||
|
||||
|
||||
def test_if_else_human_input_pause_resume_answer_chunks_survive_resume() -> None:
|
||||
# ---- Phase 1: run to GraphRunPausedEvent ----
|
||||
runtime_state_1 = _build_runtime_state()
|
||||
graph_1 = _build_graph(runtime_state_1, _mock_repo_paused())
|
||||
engine_1 = GraphEngine(
|
||||
workflow_id="wf",
|
||||
graph=graph_1,
|
||||
graph_runtime_state=runtime_state_1,
|
||||
command_channel=InMemoryChannel(),
|
||||
config=GraphEngineConfig(),
|
||||
)
|
||||
filter_1 = ResponseStreamFilter()
|
||||
phase1_events = list(
|
||||
filter_graph_events(
|
||||
engine_1.run(),
|
||||
context=GraphEventFilterContext.from_engine(engine_1),
|
||||
filters=[filter_1],
|
||||
)
|
||||
)
|
||||
|
||||
assert any(isinstance(e, GraphRunPausedEvent) for e in phase1_events)
|
||||
phase1_chunks = [e for e in phase1_events if isinstance(e, NodeRunStreamChunkEvent)]
|
||||
assert not any(e.node_id in ("answer_after_pause", "answer_after_pause_2") for e in phase1_chunks)
|
||||
|
||||
response_filter_snapshot = filter_1.dumps()
|
||||
runtime_snapshot = runtime_state_1.dumps()
|
||||
|
||||
# ---- Phase 2: rebuild engine + filter from snapshots, resume to completion ----
|
||||
runtime_state_2 = GraphRuntimeState.from_snapshot(runtime_snapshot)
|
||||
graph_2 = _build_graph(runtime_state_2, _mock_repo_resumed(action_id="continue"))
|
||||
engine_2 = GraphEngine(
|
||||
workflow_id="wf",
|
||||
graph=graph_2,
|
||||
graph_runtime_state=runtime_state_2,
|
||||
command_channel=InMemoryChannel(),
|
||||
config=GraphEngineConfig(),
|
||||
)
|
||||
filter_2 = ResponseStreamFilter()
|
||||
filter_2.loads(response_filter_snapshot)
|
||||
|
||||
phase2_events = list(iter_dify_graph_engine_events(engine_2, filter_2))
|
||||
|
||||
assert any(isinstance(e, GraphRunSucceededEvent) for e in phase2_events)
|
||||
|
||||
phase2_chunks = [e for e in phase2_events if isinstance(e, NodeRunStreamChunkEvent)]
|
||||
answer_1_chunks = [e for e in phase2_chunks if e.node_id == "answer_after_pause"]
|
||||
answer_2_chunks = [e for e in phase2_chunks if e.node_id == "answer_after_pause_2"]
|
||||
|
||||
assert answer_1_chunks, "answer_after_pause produced no stream chunks after resume"
|
||||
assert answer_2_chunks, "answer_after_pause_2 produced no stream chunks after resume"
|
||||
+19
@@ -20,6 +20,7 @@ providing more reliable and realistic test scenarios than mocks.
|
||||
import json
|
||||
import uuid
|
||||
from time import time
|
||||
from unittest.mock import Mock
|
||||
|
||||
import pytest
|
||||
from sqlalchemy import Engine, delete, select
|
||||
@@ -35,6 +36,7 @@ from core.workflow.system_variables import build_system_variables
|
||||
from extensions.ext_storage import storage
|
||||
from graphon.entities.pause_reason import SchedulingPause
|
||||
from graphon.enums import WorkflowExecutionStatus
|
||||
from graphon.filters import GraphEventFilterContext, ResponseStreamFilter
|
||||
from graphon.graph_engine.entities.commands import GraphEngineCommand
|
||||
from graphon.graph_engine.layers.base import GraphEngineLayerNotInitializedError
|
||||
from graphon.graph_events import GraphRunPausedEvent
|
||||
@@ -49,6 +51,22 @@ from services.file_service import FileService
|
||||
from services.workflow_run_service import WorkflowRunService
|
||||
|
||||
|
||||
def _create_initialized_response_stream_filter() -> ResponseStreamFilter:
|
||||
"""Build a `ResponseStreamFilter` that has already run `initialize()`.
|
||||
|
||||
`ResponseStreamFilter.dumps()` raises `RuntimeError` unless the filter has
|
||||
processed a `GraphEventFilterContext` first. In production this always
|
||||
happens before any event (including `GraphRunPausedEvent`) reaches
|
||||
`PauseStatePersistenceLayer.on_event`, so tests that exercise `on_event`
|
||||
or a subsequent `dumps()` call need a filter in that same state. A
|
||||
nodeless graph is enough to satisfy the precondition.
|
||||
"""
|
||||
response_stream_filter = ResponseStreamFilter()
|
||||
context = GraphEventFilterContext(graph=Mock(nodes={}), runtime_state=Mock())
|
||||
response_stream_filter.initialize(context)
|
||||
return response_stream_filter
|
||||
|
||||
|
||||
class _TestCommandChannelImpl:
|
||||
"""Real implementation of CommandChannel for testing."""
|
||||
|
||||
@@ -295,6 +313,7 @@ class TestPauseStatePersistenceLayerTestContainers:
|
||||
session_factory=self.session.get_bind(),
|
||||
state_owner_user_id=owner_id,
|
||||
generate_entity=entity,
|
||||
response_stream_filter=_create_initialized_response_stream_filter(),
|
||||
)
|
||||
|
||||
def test_complete_pause_flow_with_real_dependencies(self, db_session_with_containers: Session):
|
||||
|
||||
@@ -17,6 +17,7 @@ from core.app.layers.pause_state_persist_layer import (
|
||||
from core.workflow.nodes.human_input.pause_reason import HumanInputRequired
|
||||
from core.workflow.system_variables import SystemVariableKey
|
||||
from graphon.entities.pause_reason import HitlRequired, SchedulingPause
|
||||
from graphon.filters import GraphEventFilterContext, ResponseStreamFilter
|
||||
from graphon.graph_engine.entities.commands import GraphEngineCommand
|
||||
from graphon.graph_engine.layers.base import GraphEngineLayerNotInitializedError
|
||||
from graphon.graph_events import (
|
||||
@@ -31,6 +32,22 @@ from models.model import AppMode
|
||||
from repositories.factory import DifyAPIRepositoryFactory
|
||||
|
||||
|
||||
def _create_initialized_response_stream_filter() -> ResponseStreamFilter:
|
||||
"""Build a `ResponseStreamFilter` that has already run `initialize()`.
|
||||
|
||||
`ResponseStreamFilter.dumps()` raises `RuntimeError` unless the filter has
|
||||
processed a `GraphEventFilterContext` first. In production this always
|
||||
happens before any event (including `GraphRunPausedEvent`) reaches
|
||||
`PauseStatePersistenceLayer.on_event`, so tests that exercise `on_event`
|
||||
or a subsequent `dumps()` call need a filter in that same state. A
|
||||
nodeless graph is enough to satisfy the precondition.
|
||||
"""
|
||||
response_stream_filter = ResponseStreamFilter()
|
||||
context = GraphEventFilterContext(graph=Mock(nodes={}), runtime_state=Mock())
|
||||
response_stream_filter.initialize(context)
|
||||
return response_stream_filter
|
||||
|
||||
|
||||
class TestDataFactory:
|
||||
"""Factory helpers for constructing graph events used in tests."""
|
||||
|
||||
@@ -202,6 +219,7 @@ class TestPauseStatePersistenceLayer:
|
||||
session_factory=session_factory,
|
||||
state_owner_user_id=state_owner_user_id,
|
||||
generate_entity=self._create_generate_entity(),
|
||||
response_stream_filter=ResponseStreamFilter(),
|
||||
)
|
||||
|
||||
assert layer._session_maker is session_factory
|
||||
@@ -216,6 +234,7 @@ class TestPauseStatePersistenceLayer:
|
||||
session_factory=session_factory,
|
||||
state_owner_user_id="owner",
|
||||
generate_entity=self._create_generate_entity(),
|
||||
response_stream_filter=ResponseStreamFilter(),
|
||||
)
|
||||
|
||||
graph_runtime_state = MockReadOnlyGraphRuntimeState()
|
||||
@@ -233,6 +252,7 @@ class TestPauseStatePersistenceLayer:
|
||||
session_factory=session_factory,
|
||||
state_owner_user_id="owner-123",
|
||||
generate_entity=generate_entity,
|
||||
response_stream_filter=_create_initialized_response_stream_filter(),
|
||||
)
|
||||
|
||||
mock_repo = Mock()
|
||||
@@ -272,6 +292,7 @@ class TestPauseStatePersistenceLayer:
|
||||
session_factory=session_factory,
|
||||
state_owner_user_id="owner-123",
|
||||
generate_entity=generate_entity,
|
||||
response_stream_filter=_create_initialized_response_stream_filter(),
|
||||
)
|
||||
|
||||
mock_repo = Mock()
|
||||
@@ -328,6 +349,7 @@ class TestPauseStatePersistenceLayer:
|
||||
session_factory=session_factory,
|
||||
state_owner_user_id="owner-123",
|
||||
generate_entity=self._create_generate_entity(),
|
||||
response_stream_filter=ResponseStreamFilter(),
|
||||
)
|
||||
|
||||
mock_repo = Mock()
|
||||
@@ -356,6 +378,7 @@ class TestPauseStatePersistenceLayer:
|
||||
session_factory=session_factory,
|
||||
state_owner_user_id="owner-123",
|
||||
generate_entity=self._create_generate_entity(),
|
||||
response_stream_filter=ResponseStreamFilter(),
|
||||
)
|
||||
|
||||
event = TestDataFactory.create_graph_run_paused_event()
|
||||
@@ -369,6 +392,7 @@ class TestPauseStatePersistenceLayer:
|
||||
session_factory=session_factory,
|
||||
state_owner_user_id="owner-123",
|
||||
generate_entity=self._create_generate_entity(),
|
||||
response_stream_filter=_create_initialized_response_stream_filter(),
|
||||
)
|
||||
|
||||
mock_repo = Mock()
|
||||
@@ -468,3 +492,53 @@ def test_workflow_resumption_context_dumps_loads_roundtrip(state: WorkflowResump
|
||||
restored_entity = loaded.get_generate_entity()
|
||||
assert isinstance(restored_entity, type(state.generate_entity.entity))
|
||||
assert restored_entity.extras["trace_session_id"] == "session-1"
|
||||
|
||||
|
||||
def test_on_event_persists_response_stream_filter_dump(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
session_factory = Mock(name="session_factory")
|
||||
generate_entity = TestPauseStatePersistenceLayer._create_generate_entity(workflow_execution_id="run-123")
|
||||
response_stream_filter = _create_initialized_response_stream_filter()
|
||||
layer = PauseStatePersistenceLayer(
|
||||
session_factory=session_factory,
|
||||
state_owner_user_id="owner-123",
|
||||
generate_entity=generate_entity,
|
||||
response_stream_filter=response_stream_filter,
|
||||
)
|
||||
|
||||
mock_repo = Mock()
|
||||
mock_factory = Mock(return_value=mock_repo)
|
||||
monkeypatch.setattr(DifyAPIRepositoryFactory, "create_api_workflow_run_repository", mock_factory)
|
||||
|
||||
graph_runtime_state = MockReadOnlyGraphRuntimeState(workflow_execution_id="run-123")
|
||||
layer.initialize(graph_runtime_state, MockCommandChannel())
|
||||
|
||||
event = TestDataFactory.create_graph_run_paused_event()
|
||||
layer.on_event(event)
|
||||
|
||||
serialized_state = mock_repo.create_workflow_pause.call_args.kwargs["state"]
|
||||
resumption_context = WorkflowResumptionContext.loads(serialized_state)
|
||||
assert resumption_context.serialized_response_stream_filter_state == response_stream_filter.dumps()
|
||||
|
||||
|
||||
def test_get_response_stream_filter_restores_dumped_state() -> None:
|
||||
original = _create_initialized_response_stream_filter()
|
||||
context = WorkflowResumptionContext(
|
||||
serialized_graph_runtime_state=json.dumps({"state": "workflow"}),
|
||||
generate_entity=_WorkflowGenerateEntityWrapper(entity=TestPauseStatePersistenceLayer._create_generate_entity()),
|
||||
serialized_response_stream_filter_state=original.dumps(),
|
||||
)
|
||||
|
||||
restored = context.get_response_stream_filter()
|
||||
|
||||
assert restored.dumps() == original.dumps()
|
||||
|
||||
|
||||
def test_get_response_stream_filter_defaults_when_state_missing() -> None:
|
||||
context = WorkflowResumptionContext(
|
||||
serialized_graph_runtime_state=json.dumps({"state": "workflow"}),
|
||||
generate_entity=_WorkflowGenerateEntityWrapper(entity=TestPauseStatePersistenceLayer._create_generate_entity()),
|
||||
)
|
||||
|
||||
restored = context.get_response_stream_filter()
|
||||
|
||||
assert isinstance(restored, ResponseStreamFilter)
|
||||
|
||||
@@ -13,6 +13,7 @@ from graphon.entities.base_node_data import BaseNodeData
|
||||
from graphon.enums import NodeType, WorkflowNodeExecutionStatus
|
||||
from graphon.errors import WorkflowNodeRunFailedError
|
||||
from graphon.file import File, FileTransferMethod, FileType
|
||||
from graphon.filters import ResponseStreamFilter
|
||||
from graphon.graph import Graph
|
||||
from graphon.graph_events import GraphRunFailedEvent
|
||||
from graphon.model_runtime.entities.llm_entities import LLMMode, LLMUsage
|
||||
@@ -241,6 +242,37 @@ class TestWorkflowChildEngineBuilder:
|
||||
)
|
||||
|
||||
|
||||
def _build_minimal_workflow_entry(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
*,
|
||||
response_stream_filter: ResponseStreamFilter | None = None,
|
||||
) -> workflow_entry.WorkflowEntry:
|
||||
"""Construct a minimal WorkflowEntry with GraphEngine construction mocked out."""
|
||||
graph_engine = MagicMock()
|
||||
graph_runtime_state = SimpleNamespace(execution_context=None)
|
||||
|
||||
monkeypatch.setattr(workflow_entry, "capture_current_context", lambda: sentinel.execution_context)
|
||||
monkeypatch.setattr(workflow_entry, "GraphEngine", MagicMock(return_value=graph_engine))
|
||||
monkeypatch.setattr(workflow_entry, "GraphEngineConfig", MagicMock(return_value=sentinel.graph_engine_config))
|
||||
monkeypatch.setattr(workflow_entry, "InMemoryChannel", MagicMock(return_value=sentinel.command_channel))
|
||||
monkeypatch.setattr(workflow_entry, "LLMQuotaLayer", MagicMock(return_value=sentinel.llm_quota_layer))
|
||||
|
||||
return workflow_entry.WorkflowEntry(
|
||||
tenant_id="tenant-id",
|
||||
app_id="app-id",
|
||||
workflow_id="workflow-id",
|
||||
graph_config={"nodes": [], "edges": []},
|
||||
graph=sentinel.graph,
|
||||
user_id="user-id",
|
||||
user_from=UserFrom.ACCOUNT,
|
||||
invoke_from=InvokeFrom.DEBUGGER,
|
||||
call_depth=0,
|
||||
variable_pool=sentinel.variable_pool,
|
||||
graph_runtime_state=graph_runtime_state,
|
||||
response_stream_filter=response_stream_filter,
|
||||
)
|
||||
|
||||
|
||||
class TestWorkflowEntryInit:
|
||||
def test_rejects_call_depth_above_limit(self):
|
||||
call_depth = workflow_entry.dify_config.WORKFLOW_CALL_MAX_DEPTH + 1
|
||||
@@ -329,12 +361,24 @@ class TestWorkflowEntryInit:
|
||||
((observability_layer,), {}),
|
||||
]
|
||||
|
||||
def test_workflow_entry_stores_supplied_response_stream_filter(self, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
supplied_filter = ResponseStreamFilter()
|
||||
entry = _build_minimal_workflow_entry(monkeypatch, response_stream_filter=supplied_filter)
|
||||
|
||||
assert entry._response_stream_filter is supplied_filter
|
||||
|
||||
def test_workflow_entry_defaults_to_fresh_response_stream_filter(self, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
entry = _build_minimal_workflow_entry(monkeypatch, response_stream_filter=None)
|
||||
|
||||
assert isinstance(entry._response_stream_filter, ResponseStreamFilter)
|
||||
|
||||
|
||||
class TestWorkflowEntryRun:
|
||||
def test_run_swallows_generate_task_stopped_errors(self):
|
||||
entry = object.__new__(workflow_entry.WorkflowEntry)
|
||||
entry.graph_engine = MagicMock()
|
||||
entry.graph_engine.run.side_effect = GenerateTaskStoppedError()
|
||||
entry._response_stream_filter = ResponseStreamFilter()
|
||||
|
||||
assert list(entry.run()) == []
|
||||
|
||||
@@ -373,6 +417,7 @@ class TestWorkflowEntryRun:
|
||||
def test_run_delegates_to_dify_event_iterator(self):
|
||||
entry = object.__new__(workflow_entry.WorkflowEntry)
|
||||
entry.graph_engine = sentinel.graph_engine
|
||||
entry._response_stream_filter = sentinel.response_stream_filter
|
||||
|
||||
with patch.object(
|
||||
workflow_entry,
|
||||
@@ -382,12 +427,13 @@ class TestWorkflowEntryRun:
|
||||
events = list(entry.run())
|
||||
|
||||
assert events == [sentinel.filtered_event]
|
||||
iter_dify_graph_engine_events.assert_called_once_with(sentinel.graph_engine)
|
||||
iter_dify_graph_engine_events.assert_called_once_with(sentinel.graph_engine, sentinel.response_stream_filter)
|
||||
|
||||
def test_run_emits_failed_event_for_unexpected_errors(self):
|
||||
entry = object.__new__(workflow_entry.WorkflowEntry)
|
||||
entry.graph_engine = MagicMock()
|
||||
entry.graph_engine.run.side_effect = RuntimeError("boom")
|
||||
entry._response_stream_filter = ResponseStreamFilter()
|
||||
|
||||
events = list(entry.run())
|
||||
|
||||
|
||||
@@ -723,6 +723,7 @@ def test_resume_advanced_chat_publishes_events_for_originally_blocking_runs(monk
|
||||
message=MagicMock(),
|
||||
generate_entity=generate_entity,
|
||||
graph_runtime_state=MagicMock(),
|
||||
response_stream_filter=MagicMock(),
|
||||
session_factory=MagicMock(),
|
||||
pause_state_config=MagicMock(),
|
||||
workflow_run_id="workflow-run-id",
|
||||
@@ -774,6 +775,7 @@ def test_resume_workflow_publishes_events_for_originally_blocking_runs(monkeypat
|
||||
user=MagicMock(),
|
||||
generate_entity=generate_entity,
|
||||
graph_runtime_state=MagicMock(),
|
||||
response_stream_filter=MagicMock(),
|
||||
session_factory=MagicMock(),
|
||||
pause_state_config=MagicMock(),
|
||||
workflow_run_id="workflow-run-id",
|
||||
@@ -829,6 +831,7 @@ def test_resume_workflow_ignores_missing_old_pause_after_repause(monkeypatch: py
|
||||
user=MagicMock(),
|
||||
generate_entity=generate_entity,
|
||||
graph_runtime_state=MagicMock(),
|
||||
response_stream_filter=MagicMock(),
|
||||
session_factory=MagicMock(),
|
||||
pause_state_config=MagicMock(),
|
||||
workflow_run_id="workflow-run-id",
|
||||
|
||||
Reference in New Issue
Block a user