diff --git a/python/packages/ag-ui/agent_framework_ag_ui/_workflow.py b/python/packages/ag-ui/agent_framework_ag_ui/_workflow.py index 73b2a1d7ec..5b35f8d00c 100644 --- a/python/packages/ag-ui/agent_framework_ag_ui/_workflow.py +++ b/python/packages/ag-ui/agent_framework_ag_ui/_workflow.py @@ -12,6 +12,12 @@ from ag_ui.core import ( BaseEvent, MessagesSnapshotEvent, + ReasoningEncryptedValueEvent, + ReasoningEndEvent, + ReasoningMessageContentEvent, + ReasoningMessageEndEvent, + ReasoningMessageStartEvent, + ReasoningStartEvent, RunErrorEvent, RunFinishedEvent, RunStartedEvent, @@ -122,6 +128,7 @@ def __init__(self, raw_messages: list[dict[str, Any]]) -> None: self._synthesized_messages = agui_messages_to_snapshot_format(raw_messages) self._emitted_messages: list[dict[str, Any]] | None = None self._open_text_message: dict[str, Any] | None = None + self._open_reasoning_message: dict[str, Any] | None = None self._tool_call_message: dict[str, Any] | None = None self._tool_calls_by_id: dict[str, dict[str, Any]] = {} self.state: dict[str, Any] | None = None @@ -165,14 +172,27 @@ def observe(self, event: BaseEvent) -> None: self._observe_tool_call_args(event) elif isinstance(event, ToolCallResultEvent): self._observe_tool_call_result(event) + elif isinstance(event, ReasoningStartEvent): + # A new reasoning block supersedes anything still open from the last one. + self._flush_open_reasoning_message() + elif isinstance(event, ReasoningMessageStartEvent): + self._observe_reasoning_start(event) + elif isinstance(event, ReasoningMessageContentEvent): + self._observe_reasoning_content(event) + elif isinstance(event, (ReasoningMessageEndEvent, ReasoningEndEvent)): + self._observe_reasoning_end(event) + elif isinstance(event, ReasoningEncryptedValueEvent): + self._observe_reasoning_encrypted_value(event) def build(self) -> AGUIThreadSnapshot: """Return the replayable thread snapshot.""" + self._flush_open_reasoning_message() self._flush_open_text_message() messages = self._emitted_messages if self._emitted_messages is not None else self._synthesized_messages return AGUIThreadSnapshot(messages=messages, state=self.state, interrupt=self.interrupt) def _observe_text_start(self, event: TextMessageStartEvent) -> None: + self._flush_open_reasoning_message() if self._open_text_message is not None and self._open_text_message.get("id") != event.message_id: self._flush_open_text_message() self._open_text_message = {"id": event.message_id, "role": event.role, "content": ""} @@ -188,6 +208,7 @@ def _observe_text_end(self, event: TextMessageEndEvent) -> None: self._flush_open_text_message() def _observe_tool_call_start(self, event: ToolCallStartEvent) -> None: + self._flush_open_reasoning_message() parent_message_id = event.parent_message_id if ( self._open_text_message is not None @@ -245,6 +266,53 @@ def _flush_open_text_message(self) -> None: self._tool_call_message = None self._open_text_message = None + def _observe_reasoning_start(self, event: ReasoningMessageStartEvent) -> None: + if self._open_reasoning_message is not None and self._open_reasoning_message.get("id") != event.message_id: + self._flush_open_reasoning_message() + self._open_reasoning_message = {"id": event.message_id, "role": "reasoning", "content": ""} + + def _observe_reasoning_content(self, event: ReasoningMessageContentEvent) -> None: + if self._open_reasoning_message is None or self._open_reasoning_message.get("id") != event.message_id: + self._flush_open_reasoning_message() + self._open_reasoning_message = {"id": event.message_id, "role": "reasoning", "content": ""} + self._open_reasoning_message["content"] = f"{self._open_reasoning_message.get('content', '')}{event.delta}" + + def _observe_reasoning_end(self, event: ReasoningMessageEndEvent | ReasoningEndEvent) -> None: + if self._open_reasoning_message is None or self._open_reasoning_message.get("id") != event.message_id: + return + if isinstance(event, ReasoningMessageEndEvent): + # REASONING_MESSAGE_END closes the message, not the block, and an + # encrypted value is block-scoped so it legitimately trails it -- that is + # the order `_emit_text_reasoning` produces without a flow. Keep the + # message open so protected-data-only reasoning still has somewhere to + # land; REASONING_END, the next block, or build() finalizes it. + return + self._flush_open_reasoning_message() + + def _observe_reasoning_encrypted_value(self, event: ReasoningEncryptedValueEvent) -> None: + # Only message-scoped encrypted values belong on a reasoning message; the + # protocol also uses this event for other subtypes. + if event.subtype != "message": + return + if self._open_reasoning_message is not None and self._open_reasoning_message.get("id") == event.entity_id: + self._open_reasoning_message["encryptedValue"] = event.encrypted_value + return + # Intervening text or tool output can flush the message before its encrypted + # value arrives; attach it to the message we already synthesized. + for message in reversed(self._synthesized_messages): + if message.get("role") == "reasoning" and message.get("id") == event.entity_id: + message["encryptedValue"] = event.encrypted_value + return + + def _flush_open_reasoning_message(self) -> None: + if self._open_reasoning_message is None: + return + # An encrypted-value-only block carries no display text but still has to + # survive hydration, so it counts as content worth keeping. + if self._open_reasoning_message.get("content") or self._open_reasoning_message.get("encryptedValue"): + self._synthesized_messages.append(self._open_reasoning_message) + self._open_reasoning_message = None + class AgentFrameworkWorkflow: """Base AG-UI workflow wrapper. diff --git a/python/packages/ag-ui/tests/ag_ui/test_snapshots.py b/python/packages/ag-ui/tests/ag_ui/test_snapshots.py index e2ea85e37d..5b4e8e694f 100644 --- a/python/packages/ag-ui/tests/ag_ui/test_snapshots.py +++ b/python/packages/ag-ui/tests/ag_ui/test_snapshots.py @@ -185,6 +185,290 @@ def test_workflow_snapshot_builder_splits_tool_call_groups() -> None: ] +def test_workflow_snapshot_builder_folds_reasoning_into_snapshot() -> None: + """Streamed reasoning deltas accumulate into a replayable reasoning message.""" + from ag_ui.core import ( + ReasoningMessageContentEvent, + ReasoningMessageEndEvent, + ReasoningMessageStartEvent, + ) + + from agent_framework_ag_ui._workflow import _WorkflowSnapshotBuilder + + builder = _WorkflowSnapshotBuilder([]) + builder.observe(ReasoningMessageStartEvent(message_id="reason-1", role="reasoning")) + builder.observe(ReasoningMessageContentEvent(message_id="reason-1", delta="step one ")) + builder.observe(ReasoningMessageContentEvent(message_id="reason-1", delta="step two")) + builder.observe(ReasoningMessageEndEvent(message_id="reason-1")) + + assert builder.build().messages == [{"id": "reason-1", "role": "reasoning", "content": "step one step two"}] + + +def test_workflow_snapshot_builder_keeps_reasoning_in_emission_order() -> None: + """Reasoning is replayed where it streamed, not appended after the visible output.""" + from ag_ui.core import ( + ReasoningMessageContentEvent, + ReasoningMessageEndEvent, + ReasoningMessageStartEvent, + TextMessageContentEvent, + TextMessageEndEvent, + TextMessageStartEvent, + ToolCallResultEvent, + ToolCallStartEvent, + ) + + from agent_framework_ag_ui._workflow import _WorkflowSnapshotBuilder + + builder = _WorkflowSnapshotBuilder([]) + builder.observe(ReasoningMessageStartEvent(message_id="reason-1", role="reasoning")) + builder.observe(ReasoningMessageContentEvent(message_id="reason-1", delta="planning")) + builder.observe(ReasoningMessageEndEvent(message_id="reason-1")) + builder.observe(ToolCallStartEvent(tool_call_id="call-a", tool_call_name="toolA")) + builder.observe(ToolCallResultEvent(message_id="result-a", tool_call_id="call-a", content="resA")) + builder.observe(TextMessageStartEvent(message_id="text-1", role="assistant")) + builder.observe(TextMessageContentEvent(message_id="text-1", delta="done")) + builder.observe(TextMessageEndEvent(message_id="text-1")) + + messages = builder.build().messages + assert [ + ( + message.get("role"), + [tool_call["id"] for tool_call in message.get("tool_calls", [])] + or message.get("toolCallId") + or message.get("content"), + ) + for message in messages + ] == [ + ("reasoning", "planning"), + ("assistant", ["call-a"]), + ("tool", "call-a"), + ("assistant", "done"), + ] + + +def test_workflow_snapshot_builder_captures_reasoning_encrypted_value() -> None: + """Encrypted reasoning payloads survive hydration under the protocol's camelCase key.""" + from ag_ui.core import ( + ReasoningEncryptedValueEvent, + ReasoningMessageContentEvent, + ReasoningMessageEndEvent, + ReasoningMessageStartEvent, + ) + + from agent_framework_ag_ui._workflow import _WorkflowSnapshotBuilder + + builder = _WorkflowSnapshotBuilder([]) + builder.observe(ReasoningMessageStartEvent(message_id="reason-1", role="reasoning")) + builder.observe(ReasoningMessageContentEvent(message_id="reason-1", delta="hidden")) + builder.observe(ReasoningEncryptedValueEvent(subtype="message", entity_id="reason-1", encrypted_value="cipher")) + builder.observe(ReasoningMessageEndEvent(message_id="reason-1")) + + assert builder.build().messages == [ + {"id": "reason-1", "role": "reasoning", "content": "hidden", "encryptedValue": "cipher"} + ] + + +def test_workflow_snapshot_builder_flushes_reasoning_left_open_at_build() -> None: + """A run that ends without REASONING_MESSAGE_END still snapshots what streamed.""" + from ag_ui.core import ReasoningMessageContentEvent, ReasoningMessageStartEvent + + from agent_framework_ag_ui._workflow import _WorkflowSnapshotBuilder + + builder = _WorkflowSnapshotBuilder([]) + builder.observe(ReasoningMessageStartEvent(message_id="reason-1", role="reasoning")) + builder.observe(ReasoningMessageContentEvent(message_id="reason-1", delta="unterminated")) + + assert builder.build().messages == [{"id": "reason-1", "role": "reasoning", "content": "unterminated"}] + + +def test_workflow_snapshot_builder_separates_consecutive_reasoning_blocks() -> None: + """A new reasoning message id closes the previous block instead of merging into it.""" + from ag_ui.core import ReasoningMessageContentEvent, ReasoningMessageStartEvent + + from agent_framework_ag_ui._workflow import _WorkflowSnapshotBuilder + + builder = _WorkflowSnapshotBuilder([]) + builder.observe(ReasoningMessageStartEvent(message_id="reason-1", role="reasoning")) + builder.observe(ReasoningMessageContentEvent(message_id="reason-1", delta="first")) + # Second block opens without the first ever being closed. + builder.observe(ReasoningMessageStartEvent(message_id="reason-2", role="reasoning")) + builder.observe(ReasoningMessageContentEvent(message_id="reason-2", delta="second")) + + assert builder.build().messages == [ + {"id": "reason-1", "role": "reasoning", "content": "first"}, + {"id": "reason-2", "role": "reasoning", "content": "second"}, + ] + + +def test_workflow_snapshot_builder_folds_reasoning_content_without_a_start_event() -> None: + """Reasoning deltas are kept even if the opening REASONING_MESSAGE_START was missed.""" + from ag_ui.core import ReasoningMessageContentEvent + + from agent_framework_ag_ui._workflow import _WorkflowSnapshotBuilder + + builder = _WorkflowSnapshotBuilder([]) + builder.observe(ReasoningMessageContentEvent(message_id="reason-1", delta="orphaned")) + + assert builder.build().messages == [{"id": "reason-1", "role": "reasoning", "content": "orphaned"}] + + +def test_workflow_snapshot_builder_attaches_encrypted_value_after_block_closed() -> None: + """An encrypted value trailing REASONING_MESSAGE_END still lands on its message.""" + from ag_ui.core import ( + ReasoningEncryptedValueEvent, + ReasoningMessageContentEvent, + ReasoningMessageEndEvent, + ReasoningMessageStartEvent, + ) + + from agent_framework_ag_ui._workflow import _WorkflowSnapshotBuilder + + builder = _WorkflowSnapshotBuilder([]) + builder.observe(ReasoningMessageStartEvent(message_id="reason-1", role="reasoning")) + builder.observe(ReasoningMessageContentEvent(message_id="reason-1", delta="hidden")) + builder.observe(ReasoningMessageEndEvent(message_id="reason-1")) + builder.observe( + ReasoningEncryptedValueEvent(subtype="message", entity_id="reason-1", encrypted_value="late-cipher") + ) + + assert builder.build().messages == [ + {"id": "reason-1", "role": "reasoning", "content": "hidden", "encryptedValue": "late-cipher"} + ] + + +def test_workflow_snapshot_builder_keeps_encrypted_only_reasoning_ended_before_its_value() -> None: + """Protected-data-only reasoning survives the no-flow order, where END precedes the value. + + `_emit_text_reasoning` without a flow emits REASONING_MESSAGE_END before + REASONING_ENCRYPTED_VALUE. The message carries no display text, so closing it at + REASONING_MESSAGE_END would discard it and leave the encrypted value nothing to + attach to. + """ + from ag_ui.core import ( + ReasoningEncryptedValueEvent, + ReasoningEndEvent, + ReasoningMessageEndEvent, + ReasoningMessageStartEvent, + ReasoningStartEvent, + ) + + from agent_framework_ag_ui._workflow import _WorkflowSnapshotBuilder + + builder = _WorkflowSnapshotBuilder([]) + builder.observe(ReasoningStartEvent(message_id="reason-1")) + builder.observe(ReasoningMessageStartEvent(message_id="reason-1", role="reasoning")) + builder.observe(ReasoningMessageEndEvent(message_id="reason-1")) + builder.observe(ReasoningEncryptedValueEvent(subtype="message", entity_id="reason-1", encrypted_value="cipher")) + builder.observe(ReasoningEndEvent(message_id="reason-1")) + + assert builder.build().messages == [ + {"id": "reason-1", "role": "reasoning", "content": "", "encryptedValue": "cipher"} + ] + + +def test_workflow_snapshot_builder_drops_reasoning_with_neither_text_nor_encrypted_value() -> None: + """An empty reasoning block with nothing to replay is not synthesized into a message.""" + from ag_ui.core import ( + ReasoningEndEvent, + ReasoningMessageEndEvent, + ReasoningMessageStartEvent, + ReasoningStartEvent, + ) + + from agent_framework_ag_ui._workflow import _WorkflowSnapshotBuilder + + builder = _WorkflowSnapshotBuilder([]) + builder.observe(ReasoningStartEvent(message_id="reason-1")) + builder.observe(ReasoningMessageStartEvent(message_id="reason-1", role="reasoning")) + builder.observe(ReasoningMessageEndEvent(message_id="reason-1")) + builder.observe(ReasoningEndEvent(message_id="reason-1")) + + assert builder.build().messages == [] + + +def test_workflow_snapshot_builder_closes_previous_block_on_new_reasoning_start() -> None: + """A new REASONING_START finalizes a message the previous block left open.""" + from ag_ui.core import ( + ReasoningMessageContentEvent, + ReasoningMessageEndEvent, + ReasoningMessageStartEvent, + ReasoningStartEvent, + ) + + from agent_framework_ag_ui._workflow import _WorkflowSnapshotBuilder + + builder = _WorkflowSnapshotBuilder([]) + builder.observe(ReasoningStartEvent(message_id="reason-1")) + builder.observe(ReasoningMessageStartEvent(message_id="reason-1", role="reasoning")) + builder.observe(ReasoningMessageContentEvent(message_id="reason-1", delta="first")) + # No REASONING_END: the next block's start has to close this one. + builder.observe(ReasoningStartEvent(message_id="reason-2")) + builder.observe(ReasoningMessageStartEvent(message_id="reason-2", role="reasoning")) + builder.observe(ReasoningMessageContentEvent(message_id="reason-2", delta="second")) + builder.observe(ReasoningMessageEndEvent(message_id="reason-2")) + + assert builder.build().messages == [ + {"id": "reason-1", "role": "reasoning", "content": "first"}, + {"id": "reason-2", "role": "reasoning", "content": "second"}, + ] + + +def test_workflow_snapshot_builder_attaches_encrypted_value_after_intervening_output() -> None: + """Text arriving between a reasoning message and its encrypted value does not lose the value.""" + from ag_ui.core import ( + ReasoningEncryptedValueEvent, + ReasoningMessageContentEvent, + ReasoningMessageEndEvent, + ReasoningMessageStartEvent, + TextMessageContentEvent, + TextMessageStartEvent, + ) + + from agent_framework_ag_ui._workflow import _WorkflowSnapshotBuilder + + builder = _WorkflowSnapshotBuilder([]) + builder.observe(ReasoningMessageStartEvent(message_id="reason-1", role="reasoning")) + builder.observe(ReasoningMessageContentEvent(message_id="reason-1", delta="thinking")) + builder.observe(ReasoningMessageEndEvent(message_id="reason-1")) + # Text output flushes the reasoning message before the encrypted value shows up. + builder.observe(TextMessageStartEvent(message_id="text-1", role="assistant")) + builder.observe(TextMessageContentEvent(message_id="text-1", delta="answer")) + builder.observe(ReasoningEncryptedValueEvent(subtype="message", entity_id="reason-1", encrypted_value="cipher")) + + assert builder.build().messages == [ + {"id": "reason-1", "role": "reasoning", "content": "thinking", "encryptedValue": "cipher"}, + {"id": "text-1", "role": "assistant", "content": "answer"}, + ] + + +def test_workflow_snapshot_builder_ignores_reasoning_end_for_unopened_message() -> None: + """A reasoning end event for a message that was never opened is a no-op.""" + from ag_ui.core import ReasoningEndEvent, ReasoningMessageContentEvent, ReasoningMessageStartEvent + + from agent_framework_ag_ui._workflow import _WorkflowSnapshotBuilder + + builder = _WorkflowSnapshotBuilder([]) + builder.observe(ReasoningEndEvent(message_id="never-opened")) + builder.observe(ReasoningMessageStartEvent(message_id="reason-1", role="reasoning")) + builder.observe(ReasoningMessageContentEvent(message_id="reason-1", delta="kept")) + # An end event for a different id must not close the open message. + builder.observe(ReasoningEndEvent(message_id="other")) + + assert builder.build().messages == [{"id": "reason-1", "role": "reasoning", "content": "kept"}] + + +def test_workflow_snapshot_builder_ignores_tool_call_encrypted_values() -> None: + """A tool-call-scoped encrypted value must not be folded in as reasoning content.""" + from ag_ui.core import ReasoningEncryptedValueEvent + + from agent_framework_ag_ui._workflow import _WorkflowSnapshotBuilder + + builder = _WorkflowSnapshotBuilder([]) + builder.observe(ReasoningEncryptedValueEvent(subtype="tool-call", entity_id="call-a", encrypted_value="cipher")) + + assert builder.build().messages == [] + + async def test_in_memory_snapshot_store_rejects_invalid_keys() -> None: """Key parts must be non-empty strings for every store operation.""" import pytest diff --git a/python/packages/ag-ui/tests/ag_ui/test_workflow_agent.py b/python/packages/ag-ui/tests/ag_ui/test_workflow_agent.py index 4a4ec41160..8702adcd04 100644 --- a/python/packages/ag-ui/tests/ag_ui/test_workflow_agent.py +++ b/python/packages/ag-ui/tests/ag_ui/test_workflow_agent.py @@ -359,3 +359,64 @@ async def test_workflow_checkpoint_only_resume_preserves_thread_snapshot() -> No assert "Earlier reply" in contents # ...plus the newly produced output from the resumed run. assert any(isinstance(content, str) and "done" in content for content in contents) + + +async def test_workflow_snapshot_preserves_streamed_reasoning() -> None: + """Reasoning that streamed during a workflow run survives thread hydration. + + Regression test for workflow reasoning rendering live and then vanishing from the + stored AG-UI Thread Snapshot, so a hydrated thread lost the intermediate output + that the agent path retains. + """ + from agent_framework_ag_ui import InMemoryAGUIThreadSnapshotStore + from agent_framework_ag_ui._snapshots import _SNAPSHOT_SCOPE_INPUT_KEY + + @executor(id="thinker") + async def thinker(message: Any, ctx: WorkflowContext[str, str]) -> None: + await ctx.yield_output("Weighing the options...") + await ctx.send_message("go") + + @executor(id="finalizer") + async def finalizer(message: str, ctx: WorkflowContext[None, str]) -> None: + await ctx.yield_output("Final answer.") + + workflow = ( + WorkflowBuilder( + start_executor=thinker, + output_from=[finalizer], + intermediate_output_from=[thinker], + ) + .add_edge(thinker, finalizer) + .build() + ) + + store = InMemoryAGUIThreadSnapshotStore() + agent = AgentFrameworkWorkflow(workflow=workflow, snapshot_store=store) + + events = await _run( + agent, + { + "thread_id": "thread-reasoning", + "run_id": "run-1", + "messages": [{"id": "user-1", "role": "user", "content": "Decide"}], + _SNAPSHOT_SCOPE_INPUT_KEY: "tenant-a", + }, + ) + assert "RUN_ERROR" not in [event.type for event in events] + + # The reasoning really did stream ... + reasoning_deltas = [ + event.delta # type: ignore[attr-defined] # ty: ignore[unresolved-attribute] + for event in events + if event.type == "REASONING_MESSAGE_CONTENT" + ] + assert "Weighing the options..." in reasoning_deltas + + # ... and it is still there once the thread is hydrated from the snapshot. + snapshot = await store.get(scope="tenant-a", thread_id="thread-reasoning") + assert snapshot is not None + reasoning_messages = [message for message in snapshot.messages if message.get("role") == "reasoning"] + assert [message.get("content") for message in reasoning_messages] == ["Weighing the options..."] + + # The final assistant text is still preserved alongside it. + assert any(message.get("content") == "Final answer." for message in snapshot.messages)