diff --git a/api/core/app/apps/advanced_chat/app_generator.py b/api/core/app/apps/advanced_chat/app_generator.py index 75ada4fe888..23acfb7ea62 100644 --- a/api/core/app/apps/advanced_chat/app_generator.py +++ b/api/core/app/apps/advanced_chat/app_generator.py @@ -48,6 +48,7 @@ from core.repositories import DifyCoreRepositoryFactory from core.repositories.factory import WorkflowExecutionRepository, WorkflowNodeExecutionRepository from extensions.ext_database import db from factories import file_factory +from graphon.filters import ResponseStreamFilter from graphon.graph_engine.layers import GraphEngineLayer from graphon.model_runtime.errors.invoke import InvokeAuthorizationError from graphon.runtime import GraphRuntimeState @@ -269,6 +270,7 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator): workflow_node_execution_repository: WorkflowNodeExecutionRepository, graph_runtime_state: GraphRuntimeState, pause_state_config: PauseStateLayerConfig | None = None, + response_stream_filter: ResponseStreamFilter | None = None, ): """ Resume a paused advanced chat execution. @@ -298,6 +300,7 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator): stream=application_generate_entity.stream, pause_state_config=pause_state_config, graph_runtime_state=graph_runtime_state, + response_stream_filter=response_stream_filter, ) def single_iteration_generate( @@ -492,6 +495,7 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator): pause_state_config: PauseStateLayerConfig | None = None, graph_runtime_state: GraphRuntimeState | None = None, graph_engine_layers: Sequence[GraphEngineLayer] = (), + response_stream_filter: ResponseStreamFilter | None = None, ) -> Mapping[str, Any] | Generator[str | Mapping[str, Any], None, None]: """ Generate App response. @@ -539,12 +543,14 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator): ) graph_layers: list[GraphEngineLayer] = list(graph_engine_layers) + resolved_response_stream_filter = response_stream_filter or ResponseStreamFilter() if pause_state_config is not None: graph_layers.append( PauseStatePersistenceLayer( session_factory=pause_state_config.session_factory, generate_entity=application_generate_entity, state_owner_user_id=pause_state_config.state_owner_user_id, + response_stream_filter=resolved_response_stream_filter, ) ) @@ -565,6 +571,7 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator): "workflow_node_execution_repository": workflow_node_execution_repository, "graph_engine_layers": tuple(graph_layers), "graph_runtime_state": graph_runtime_state, + "response_stream_filter": resolved_response_stream_filter, }, ) @@ -604,6 +611,7 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator): workflow_node_execution_repository: WorkflowNodeExecutionRepository, graph_engine_layers: Sequence[GraphEngineLayer] = (), graph_runtime_state: GraphRuntimeState | None = None, + response_stream_filter: ResponseStreamFilter | None = None, ): """ Generate worker in a new thread. @@ -663,6 +671,7 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator): workflow_node_execution_repository=workflow_node_execution_repository, graph_engine_layers=graph_engine_layers, graph_runtime_state=graph_runtime_state, + response_stream_filter=response_stream_filter, ) try: diff --git a/api/core/app/apps/advanced_chat/app_runner.py b/api/core/app/apps/advanced_chat/app_runner.py index b78a3b5b3dc..249cb33a98c 100644 --- a/api/core/app/apps/advanced_chat/app_runner.py +++ b/api/core/app/apps/advanced_chat/app_runner.py @@ -44,6 +44,7 @@ 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.filters import ResponseStreamFilter from graphon.graph_engine.command_channels import RedisChannel from graphon.graph_engine.layers import GraphEngineLayer from graphon.runtime import GraphRuntimeState, VariablePool @@ -78,6 +79,7 @@ class AdvancedChatAppRunner(WorkflowBasedAppRunner): workflow_node_execution_repository: WorkflowNodeExecutionRepository, graph_engine_layers: Sequence[GraphEngineLayer] = (), graph_runtime_state: GraphRuntimeState | None = None, + response_stream_filter: ResponseStreamFilter | None = None, ): super().__init__( queue_manager=queue_manager, @@ -95,6 +97,7 @@ class AdvancedChatAppRunner(WorkflowBasedAppRunner): self._workflow_execution_repository = workflow_execution_repository self._workflow_node_execution_repository = workflow_node_execution_repository self._resume_graph_runtime_state = graph_runtime_state + self._response_stream_filter = response_stream_filter @trace_span(WorkflowAppRunnerHandler) def run(self): @@ -241,6 +244,7 @@ class AdvancedChatAppRunner(WorkflowBasedAppRunner): variable_pool=variable_pool, graph_runtime_state=graph_runtime_state, command_channel=command_channel, + response_stream_filter=self._response_stream_filter, ) self._queue_manager.graph_runtime_state = graph_runtime_state diff --git a/api/core/app/apps/workflow/app_generator.py b/api/core/app/apps/workflow/app_generator.py index ab07454ff5b..168b0e525d6 100644 --- a/api/core/app/apps/workflow/app_generator.py +++ b/api/core/app/apps/workflow/app_generator.py @@ -43,6 +43,7 @@ from core.repositories import DifyCoreRepositoryFactory from core.repositories.factory import WorkflowExecutionRepository, WorkflowNodeExecutionRepository from extensions.ext_database import db from factories import file_factory +from graphon.filters import ResponseStreamFilter from graphon.graph_engine.layers import GraphEngineLayer from graphon.model_runtime.errors.invoke import InvokeAuthorizationError from graphon.runtime import GraphRuntimeState @@ -281,6 +282,7 @@ class WorkflowAppGenerator(BaseAppGenerator): graph_engine_layers: Sequence[GraphEngineLayer] = (), pause_state_config: PauseStateLayerConfig | None = None, variable_loader: VariableLoader = DUMMY_VARIABLE_LOADER, + response_stream_filter: ResponseStreamFilter | None = None, ) -> Mapping[str, Any] | Generator[str | Mapping[str, Any], None, None]: """ Resume a paused workflow execution using the persisted runtime state. @@ -311,6 +313,7 @@ class WorkflowAppGenerator(BaseAppGenerator): graph_engine_layers=graph_engine_layers, graph_runtime_state=graph_runtime_state, pause_state_config=pause_state_config, + response_stream_filter=response_stream_filter, ) def _generate( @@ -329,6 +332,7 @@ class WorkflowAppGenerator(BaseAppGenerator): graph_engine_layers: Sequence[GraphEngineLayer] = (), graph_runtime_state: GraphRuntimeState | None = None, pause_state_config: PauseStateLayerConfig | None = None, + response_stream_filter: ResponseStreamFilter | None = None, ) -> Mapping[str, Any] | Generator[str | Mapping[str, Any], None, None]: """ Generate App response. @@ -357,12 +361,14 @@ class WorkflowAppGenerator(BaseAppGenerator): app_mode=app_model.mode, ) + resolved_response_stream_filter = response_stream_filter or ResponseStreamFilter() if pause_state_config is not None: graph_layers.append( PauseStatePersistenceLayer( session_factory=pause_state_config.session_factory, generate_entity=application_generate_entity, state_owner_user_id=pause_state_config.state_owner_user_id, + response_stream_filter=resolved_response_stream_filter, ) ) @@ -385,6 +391,7 @@ class WorkflowAppGenerator(BaseAppGenerator): "workflow_node_execution_repository": workflow_node_execution_repository, "graph_engine_layers": tuple(graph_layers), "graph_runtime_state": graph_runtime_state, + "response_stream_filter": resolved_response_stream_filter, }, ) @@ -591,6 +598,7 @@ class WorkflowAppGenerator(BaseAppGenerator): root_node_id: str | None = None, graph_engine_layers: Sequence[GraphEngineLayer] = (), graph_runtime_state: GraphRuntimeState | None = None, + response_stream_filter: ResponseStreamFilter | None = None, ) -> None: """ Generate worker in a new thread. @@ -639,6 +647,7 @@ class WorkflowAppGenerator(BaseAppGenerator): root_node_id=root_node_id, graph_engine_layers=graph_engine_layers, graph_runtime_state=graph_runtime_state, + response_stream_filter=response_stream_filter, ) try: diff --git a/api/core/app/apps/workflow/app_runner.py b/api/core/app/apps/workflow/app_runner.py index 6682a395a8c..95c9d777ebd 100644 --- a/api/core/app/apps/workflow/app_runner.py +++ b/api/core/app/apps/workflow/app_runner.py @@ -23,6 +23,7 @@ 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.filters import ResponseStreamFilter from graphon.graph_engine.command_channels import RedisChannel from graphon.graph_engine.layers import GraphEngineLayer from graphon.runtime import GraphRuntimeState, VariablePool @@ -51,6 +52,7 @@ class WorkflowAppRunner(WorkflowBasedAppRunner): workflow_node_execution_repository: WorkflowNodeExecutionRepository, graph_engine_layers: Sequence[GraphEngineLayer] = (), graph_runtime_state: GraphRuntimeState | None = None, + response_stream_filter: ResponseStreamFilter | None = None, ): super().__init__( queue_manager=queue_manager, @@ -65,6 +67,7 @@ class WorkflowAppRunner(WorkflowBasedAppRunner): self._workflow_execution_repository = workflow_execution_repository self._workflow_node_execution_repository = workflow_node_execution_repository self._resume_graph_runtime_state = graph_runtime_state + self._response_stream_filter = response_stream_filter @trace_span(WorkflowAppRunnerHandler) def run(self): @@ -177,6 +180,7 @@ class WorkflowAppRunner(WorkflowBasedAppRunner): variable_pool=variable_pool, graph_runtime_state=graph_runtime_state, command_channel=command_channel, + response_stream_filter=self._response_stream_filter, ) persistence_layer = WorkflowPersistenceLayer( diff --git a/api/core/app/layers/pause_state_persist_layer.py b/api/core/app/layers/pause_state_persist_layer.py index 2a13c73eccf..5d327adb065 100644 --- a/api/core/app/layers/pause_state_persist_layer.py +++ b/api/core/app/layers/pause_state_persist_layer.py @@ -9,6 +9,7 @@ from core.app.entities.app_invoke_entities import AdvancedChatAppGenerateEntity, from core.repositories.human_input_repository import HumanInputFormSubmissionRepository from core.workflow.nodes.human_input.boundary import enrich_graph_pause_reasons from core.workflow.system_variables import SystemVariableKey, get_system_text +from graphon.filters import ResponseStreamFilter from graphon.graph_engine.layers import GraphEngineLayer from graphon.graph_events import GraphEngineEvent, GraphRunPausedEvent from models.model import AppMode @@ -43,6 +44,10 @@ class WorkflowResumptionContext(BaseModel): # Only workflow / chatflow could be paused. generate_entity: _GenerateEntityUnion serialized_graph_runtime_state: str + # Optional so that a workflow run paused before this field existed still + # loads: it just degrades to fresh-filter behavior on resume for that one + # stale run. + serialized_response_stream_filter_state: str | None = None def dumps(self) -> str: return self.model_dump_json() @@ -54,6 +59,12 @@ class WorkflowResumptionContext(BaseModel): def get_generate_entity(self) -> WorkflowAppGenerateEntity | AdvancedChatAppGenerateEntity: return self.generate_entity.entity + def get_response_stream_filter(self) -> ResponseStreamFilter: + response_stream_filter = ResponseStreamFilter() + if self.serialized_response_stream_filter_state is not None: + response_stream_filter.loads(self.serialized_response_stream_filter_state) + return response_stream_filter + @dataclass(frozen=True) class PauseStateLayerConfig: @@ -69,11 +80,17 @@ class PauseStatePersistenceLayer(GraphEngineLayer): session_factory: Engine | sessionmaker[Session], generate_entity: WorkflowAppGenerateEntity | AdvancedChatAppGenerateEntity, state_owner_user_id: str, + response_stream_filter: ResponseStreamFilter, ): """Create a PauseStatePersistenceLayer. The `state_owner_user_id` is used when creating state file for pause. It generally should id of the creator of workflow. + + `response_stream_filter` must be the exact same instance that + `WorkflowEntry` is using to stream this run's events — this layer + dumps its state on pause, and a different instance would silently + persist the wrong (empty) filter state. """ if isinstance(session_factory, Engine): session_factory = sessionmaker(session_factory) @@ -81,6 +98,7 @@ class PauseStatePersistenceLayer(GraphEngineLayer): self._session_maker = session_factory self._state_owner_user_id = state_owner_user_id self._generate_entity = generate_entity + self._response_stream_filter = response_stream_filter def _get_repo(self) -> APIWorkflowRunRepository: return DifyAPIRepositoryFactory.create_api_workflow_run_repository(self._session_maker) @@ -121,6 +139,7 @@ class PauseStatePersistenceLayer(GraphEngineLayer): state = WorkflowResumptionContext( serialized_graph_runtime_state=self.graph_runtime_state.dumps(), generate_entity=entity_wrapper, + serialized_response_stream_filter_state=self._response_stream_filter.dumps(), ) workflow_run_id = get_system_text( diff --git a/api/core/workflow/workflow_entry.py b/api/core/workflow/workflow_entry.py index 9de26b8214b..fb12922ed7f 100644 --- a/api/core/workflow/workflow_entry.py +++ b/api/core/workflow/workflow_entry.py @@ -46,18 +46,26 @@ logger = logging.getLogger(__name__) _file_access_controller = DatabaseFileAccessController() -def iter_dify_graph_engine_events(engine: GraphEngine) -> Generator[GraphEngineEvent, None, None]: +def iter_dify_graph_engine_events( + engine: GraphEngine, + response_stream_filter: ResponseStreamFilter | None = None, +) -> Generator[GraphEngineEvent, None, None]: """ Apply Dify's response streaming compatibility filter to GraphEngine events. Graphon v0.5.0 emits raw variable stream chunks and requires callers to opt into the legacy response-ordered stream behavior that Dify exposes to its workflow runners and tests. + + ``response_stream_filter``, when supplied, must be the same instance a + caller intends to persist on pause (see ``PauseStatePersistenceLayer``) so + the filter's ``paths_map`` reflects everything the engine has actually + streamed for this run. """ yield from filter_graph_events( engine.run(), context=GraphEventFilterContext.from_engine(engine), - filters=[ResponseStreamFilter()], + filters=[response_stream_filter or ResponseStreamFilter()], ) @@ -167,6 +175,7 @@ class WorkflowEntry: variable_pool: VariablePool, graph_runtime_state: GraphRuntimeState, command_channel: CommandChannel | None = None, + response_stream_filter: ResponseStreamFilter | None = None, ) -> None: """ Init workflow entry @@ -183,6 +192,8 @@ class WorkflowEntry: :param variable_pool: variable pool :param graph_runtime_state: pre-created graph runtime state :param command_channel: command channel for external control (optional, defaults to InMemoryChannel) + :param response_stream_filter: pre-restored filter for resumed runs (optional, defaults to a fresh + ResponseStreamFilter for runs with no prior pause) :param thread_pool_id: thread pool id """ # check call depth @@ -195,6 +206,7 @@ class WorkflowEntry: command_channel = InMemoryChannel() self.command_channel = command_channel + self._response_stream_filter = response_stream_filter or ResponseStreamFilter() execution_context = capture_current_context() graph_runtime_state.execution_context = execution_context self._child_engine_builder = _WorkflowChildEngineBuilder(tenant_id=tenant_id) @@ -240,7 +252,7 @@ class WorkflowEntry: try: # Preserve Dify's response-stream semantics on top of Graphon 0.5.0. - generator = iter_dify_graph_engine_events(graph_engine) + generator = iter_dify_graph_engine_events(graph_engine, self._response_stream_filter) yield from generator except GenerateTaskStoppedError: pass diff --git a/api/tasks/app_generate/workflow_execute_task.py b/api/tasks/app_generate/workflow_execute_task.py index a383839bd05..36bd21e16c1 100644 --- a/api/tasks/app_generate/workflow_execute_task.py +++ b/api/tasks/app_generate/workflow_execute_task.py @@ -25,6 +25,7 @@ from core.repositories import DifyCoreRepositoryFactory from extensions.ext_database import db from graphon.entities import WorkflowStartReason from graphon.enums import WorkflowExecutionStatus +from graphon.filters import ResponseStreamFilter from graphon.runtime import GraphRuntimeState from libs.datetime_utils import naive_utc_now from libs.flask_utils import set_login_user @@ -486,6 +487,7 @@ def _resume_app_execution(payload: dict[str, Any]) -> None: generate_entity = resumption_context.get_generate_entity() graph_runtime_state = GraphRuntimeState.from_snapshot(resumption_context.serialized_graph_runtime_state) + response_stream_filter = resumption_context.get_response_stream_filter() conversation = None message = None @@ -562,6 +564,7 @@ def _resume_app_execution(payload: dict[str, Any]) -> None: message=message, generate_entity=generate_entity, graph_runtime_state=graph_runtime_state, + response_stream_filter=response_stream_filter, session_factory=session_factory, pause_state_config=pause_config, workflow_run_id=workflow_run_id, @@ -574,6 +577,7 @@ def _resume_app_execution(payload: dict[str, Any]) -> None: user=user, generate_entity=generate_entity, graph_runtime_state=graph_runtime_state, + response_stream_filter=response_stream_filter, session_factory=session_factory, pause_state_config=pause_config, workflow_run_id=workflow_run_id, @@ -592,6 +596,7 @@ def _resume_advanced_chat( message: Message, generate_entity: AdvancedChatAppGenerateEntity, graph_runtime_state: GraphRuntimeState, + response_stream_filter: ResponseStreamFilter, session_factory: sessionmaker, pause_state_config: PauseStateLayerConfig, workflow_run_id: str, @@ -631,6 +636,7 @@ def _resume_advanced_chat( workflow_node_execution_repository=workflow_node_execution_repository, graph_runtime_state=graph_runtime_state, pause_state_config=pause_state_config, + response_stream_filter=response_stream_filter, ) except Exception: logger.exception("Failed to resume chatflow execution for workflow run %s", workflow_run_id) @@ -654,6 +660,7 @@ def _resume_workflow( user: Account | EndUser, generate_entity: WorkflowAppGenerateEntity, graph_runtime_state: GraphRuntimeState, + response_stream_filter: ResponseStreamFilter, session_factory: sessionmaker, pause_state_config: PauseStateLayerConfig, workflow_run_id: str, @@ -693,6 +700,7 @@ def _resume_workflow( workflow_execution_repository=workflow_execution_repository, workflow_node_execution_repository=workflow_node_execution_repository, pause_state_config=pause_state_config, + response_stream_filter=response_stream_filter, ) except Exception: logger.exception("Failed to resume workflow execution for workflow run %s", workflow_run_id) diff --git a/api/tasks/async_workflow_tasks.py b/api/tasks/async_workflow_tasks.py index 9f6dfc93f4c..a6cdb0a1f96 100644 --- a/api/tasks/async_workflow_tasks.py +++ b/api/tasks/async_workflow_tasks.py @@ -232,6 +232,7 @@ def resume_workflow_execution(task_data_dict: dict[str, Any]) -> None: return graph_runtime_state = GraphRuntimeState.from_snapshot(resumption_context.serialized_graph_runtime_state) + response_stream_filter = resumption_context.get_response_stream_filter() with session_factory() as session: workflow = session.scalar(select(Workflow).where(Workflow.id == workflow_run.workflow_id)) @@ -294,6 +295,7 @@ def resume_workflow_execution(task_data_dict: dict[str, Any]) -> None: workflow_node_execution_repository=workflow_node_execution_repository, graph_engine_layers=graph_engine_layers, pause_state_config=pause_config, + response_stream_filter=response_stream_filter, ) workflow_run_repo.delete_workflow_pause(pause_entity) diff --git a/api/tests/integration_tests/workflow/test_response_stream_filter_pause_resume_integration.py b/api/tests/integration_tests/workflow/test_response_stream_filter_pause_resume_integration.py new file mode 100644 index 00000000000..08ebbfb31a1 --- /dev/null +++ b/api/tests/integration_tests/workflow/test_response_stream_filter_pause_resume_integration.py @@ -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" diff --git a/api/tests/test_containers_integration_tests/core/app/layers/test_pause_state_persist_layer.py b/api/tests/test_containers_integration_tests/core/app/layers/test_pause_state_persist_layer.py index 66b3392a4b4..84f01ea52ee 100644 --- a/api/tests/test_containers_integration_tests/core/app/layers/test_pause_state_persist_layer.py +++ b/api/tests/test_containers_integration_tests/core/app/layers/test_pause_state_persist_layer.py @@ -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): diff --git a/api/tests/unit_tests/core/app/layers/test_pause_state_persist_layer.py b/api/tests/unit_tests/core/app/layers/test_pause_state_persist_layer.py index 18e724ec48b..ff7f27f5efa 100644 --- a/api/tests/unit_tests/core/app/layers/test_pause_state_persist_layer.py +++ b/api/tests/unit_tests/core/app/layers/test_pause_state_persist_layer.py @@ -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) diff --git a/api/tests/unit_tests/core/workflow/test_workflow_entry_helpers.py b/api/tests/unit_tests/core/workflow/test_workflow_entry_helpers.py index 3ccfdf76f5a..41037233b8c 100644 --- a/api/tests/unit_tests/core/workflow/test_workflow_entry_helpers.py +++ b/api/tests/unit_tests/core/workflow/test_workflow_entry_helpers.py @@ -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()) diff --git a/api/tests/unit_tests/tasks/test_workflow_execute_task.py b/api/tests/unit_tests/tasks/test_workflow_execute_task.py index 40965096b39..a3fd70f205f 100644 --- a/api/tests/unit_tests/tasks/test_workflow_execute_task.py +++ b/api/tests/unit_tests/tasks/test_workflow_execute_task.py @@ -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",