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:
Xiyuan Chen
2026-07-09 01:05:47 +00:00
committed by GitHub
co-authored by autofix-ci[bot]
parent 2b35f48d77
commit d72ee32ba1
13 changed files with 429 additions and 4 deletions
@@ -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"
@@ -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",