diff --git a/python/packages/ag-ui/agent_framework_ag_ui/_message_adapters.py b/python/packages/ag-ui/agent_framework_ag_ui/_message_adapters.py index 48e2f45e35..24624517eb 100644 --- a/python/packages/ag-ui/agent_framework_ag_ui/_message_adapters.py +++ b/python/packages/ag-ui/agent_framework_ag_ui/_message_adapters.py @@ -958,6 +958,162 @@ def _filter_modified_args( return result +def _encode_agui_segment(contents: list[Content]) -> tuple[str, list[dict[str, Any]]]: + """Encode assistant contents into an AG-UI ``(content, tool_calls)`` pair. + + Shared by both the single-message path (``agent_framework_messages_to_agui``) and the + split path (``_split_mixed_message_to_agui``) so the text / function_call + serialization lives in one place. A future argument-format or supported-content + change then updates both paths at once instead of drifting between them. + """ + text = "" + tool_calls: list[dict[str, Any]] = [] + for content in contents: + if content.type == "text": + text += content.text or "" + elif content.type == "function_call": + tool_calls.append( + { + "id": content.call_id, + "type": "function", + "function": { + "name": content.name, + "arguments": content.arguments, + }, + } + ) + return text, tool_calls + + +def _split_mixed_message_to_agui(msg: Message, role: str, unresolved_call_ids: set[str]) -> list[dict[str, Any]]: + """Convert a Message that carries function_result content into ordered AG-UI messages. + + A single Agent Framework message can interleave assistant content (text, + function_call) with one or more function_result (tool) contents -- for example a + parallel tool-call batch or a finalized turn. AG-UI needs each function_result as + its own ``tool`` message. ``_sanitize_tool_history`` resets its pending-call set on + every non-tool message, so a result is dropped as orphaned whenever an assistant + message lands between a call and that call's result. Four ordering rules keep the + transcript provider-valid: + + * A ``function_call`` must precede its matching result, so a pending assistant + segment that carries tool calls is flushed (together with any buffered text) right + before the result. + * A text-only segment is NOT flushed before a result. Such text is deferred and + emitted after the results instead. + * A new assistant segment is NOT flushed while earlier emitted calls are still + awaiting their results. For ``[call A, call B, result A, call C, result B, + result C]``, flushing ``assistant(C)`` before ``result B`` would separate the + still-open call B from its result; deferring it yields ``[assistant(A,B), tool(A), + tool(B), assistant(C), tool(C)]``. + * A result whose own call is still buffered (not yet emitted) is QUEUED rather than + emitted, because emitting it would leave it ahead of its call. For + ``[call A, call B, result A, call C, result C, result B]`` the queued ``result C`` + is held until ``assistant(C)`` is flushed, giving ``[assistant(A,B), tool(A), + tool(B), assistant(C), tool(C)]``. + + ``unresolved_call_ids`` is owned by ``agent_framework_messages_to_agui`` and carried + across the whole conversion, not rebuilt per message: a call emitted by an earlier + message stays open until its result is emitted, so a later mixed message cannot slip + a new assistant segment between that call and its result (e.g. a prior + ``assistant(call A)`` followed by ``[call C, result A, result C]``). + + The source message id is kept on the first emitted message; every additional + message gets an independent generated id. Deriving suffixes from the source id + (e.g. ``f"{base_id}-1"``) risks colliding with a legitimate id elsewhere in the + history, which would let id-keyed clients re-collapse the split messages. + """ + from ._utils import generate_event_id + + messages: list[dict[str, Any]] = [] + seg_contents: list[Content] = [] + seg_has_call = False + seg_call_ids: set[str] = set() + # Results whose own call is still buffered; emitted once that segment is flushed. + queued_results: list[Content] = [] + source_id_available = bool(msg.message_id) + + def next_id() -> str: + nonlocal source_id_available + if source_id_available and msg.message_id: + source_id_available = False + return msg.message_id + source_id_available = False + return generate_event_id() + + def flush_segment() -> None: + nonlocal seg_contents, seg_has_call + if not seg_contents: + return + seg_text, seg_tool_calls = _encode_agui_segment(seg_contents) + seg_contents = [] + seg_has_call = False + seg_call_ids.clear() + if not seg_text and not seg_tool_calls: + return + assistant_msg: dict[str, Any] = {"id": next_id(), "role": role, "content": seg_text} + if seg_tool_calls: + assistant_msg["tool_calls"] = seg_tool_calls + unresolved_call_ids.update(str(tc["id"]) for tc in seg_tool_calls if tc["id"] is not None) + messages.append(assistant_msg) + + def emit_result(content: Content) -> None: + messages.append( + { + "id": next_id(), + "role": "tool", + "content": content.result if content.result is not None else "", + "toolCallId": content.call_id, + } + ) + if content.call_id is not None: + unresolved_call_ids.discard(str(content.call_id)) + + def drain_queued() -> None: + """Flush the buffered segment, then release the results waiting on its calls.""" + if not queued_results: + return + flush_segment() + for queued in queued_results: + emit_result(queued) + queued_results.clear() + + for content in msg.contents: + if content.type in ("text", "function_call"): + seg_contents.append(content) + if content.type == "function_call": + seg_has_call = True + if content.call_id is not None: + seg_call_ids.add(str(content.call_id)) + elif content.type == "function_result": + call_id = str(content.call_id) if content.call_id is not None else None + if call_id is not None and call_id in unresolved_call_ids: + # Its call is already emitted and still open, so the result can go now. + emit_result(content) + if not unresolved_call_ids: + drain_queued() + elif seg_has_call and not unresolved_call_ids: + # No older batch is open: flush the buffered segment so its calls precede + # this result, then emit it. + flush_segment() + emit_result(content) + drain_queued() + elif call_id is not None and call_id in seg_call_ids: + # This result's call is still buffered behind an open older batch. + # Emitting now would put the result ahead of its call, so hold it. + queued_results.append(content) + else: + # Its call came from an already-emitted message: emit in place. + emit_result(content) + + # Emit any deferred / trailing segment (buffered text, or a new-call segment whose + # results arrive in a later message), then release anything still queued behind it. + flush_segment() + for queued in queued_results: + emit_result(queued) + return messages + + def agent_framework_messages_to_agui(messages: list[Message] | list[dict[str, Any]]) -> list[dict[str, Any]]: """Convert Agent Framework messages to AG-UI format. @@ -970,6 +1126,25 @@ def agent_framework_messages_to_agui(messages: list[Message] | list[dict[str, An from ._utils import generate_event_id result: list[dict[str, Any]] = [] + # Calls emitted so far whose results have not been emitted yet. Carried across every + # message (mirroring _sanitize_tool_history's pending set) so a mixed message never + # slips a new assistant segment between an earlier call and its result. + unresolved_call_ids: set[str] = set() + + def track_emitted( + role_value: str | None, tool_calls: list[dict[str, Any]] | None, tool_call_id: Any = None + ) -> None: + """Mirror _sanitize_tool_history's pending-call bookkeeping for an emitted message.""" + if role_value == "tool": + if tool_call_id: + unresolved_call_ids.discard(str(tool_call_id)) + return + # Any non-tool message resets the pending set to the calls it introduces. + unresolved_call_ids.clear() + for tool_call in tool_calls or []: + if isinstance(tool_call, dict) and tool_call.get("id") is not None: + unresolved_call_ids.add(str(tool_call["id"])) + for msg in messages: # If already a dict (AG-UI format), ensure it has an ID and normalize keys for Pydantic if isinstance(msg, dict): @@ -989,64 +1164,28 @@ def agent_framework_messages_to_agui(messages: list[Message] | list[dict[str, An normalized_msg["toolCallId"] = "" # Always append the normalized copy, not the original result.append(normalized_msg) + track_emitted( + normalized_msg.get("role"), + normalized_msg.get("tool_calls"), + normalized_msg.get("toolCallId"), + ) continue # Convert Message to AG-UI format role_value: str = msg.role if hasattr(msg.role, "value") else msg.role role = FRAMEWORK_TO_AGUI_ROLE.get(role_value, "user") - content_text = "" - tool_calls: list[dict[str, Any]] = [] - function_results: list[Any] = [] - - for content in msg.contents: - if content.type == "text": - content_text += content.text # type: ignore[operator] - elif content.type == "function_call": - tool_calls.append( - { - "id": content.call_id, - "type": "function", - "function": { - "name": content.name, - "arguments": content.arguments, - }, - } - ) - elif content.type == "function_result": - function_results.append(content) - - # A single Agent Framework message can carry several function_result - # contents (parallel tool calls). Emit one AG-UI tool message per result so - # none are dropped and each keeps its own toolCallId. - if function_results: - # Preserve the source id for the first result; give every additional - # message an independent generated id. Deriving suffixes from the source - # id (e.g. f"{base_id}-1") risks colliding with a legitimate id elsewhere - # in the history, which would let id-keyed clients re-collapse results. - for idx, fr in enumerate(function_results): - result.append( - { - "id": msg.message_id if (idx == 0 and msg.message_id) else generate_event_id(), - "role": "tool", - "content": fr.result if fr.result is not None else "", - "toolCallId": fr.call_id, - } - ) - # A mixed message may also carry text / function_call contents alongside - # the tool results (e.g. a finalized assistant turn). Emit those as a - # separate, distinctly-identified message so they are not lost. - if content_text or tool_calls: - extra_msg: dict[str, Any] = { - "id": generate_event_id(), - "role": role, - "content": content_text, - } - if tool_calls: - extra_msg["tool_calls"] = tool_calls - result.append(extra_msg) + # A message carrying function_result content may interleave assistant + # (text/function_call) and tool (function_result) segments -- e.g. parallel + # tool calls or a finalized turn. Split it into ordered AG-UI messages so no + # result is dropped and each result stays after its matching call. Messages + # with no result use the simple single-message form below. + if any(content.type == "function_result" for content in msg.contents): + result.extend(_split_mixed_message_to_agui(msg, role, unresolved_call_ids)) continue + content_text, tool_calls = _encode_agui_segment(msg.contents) + agui_msg: dict[str, Any] = { "id": msg.message_id if msg.message_id else generate_event_id(), # Always include id "role": role, @@ -1057,6 +1196,7 @@ def agent_framework_messages_to_agui(messages: list[Message] | list[dict[str, An agui_msg["tool_calls"] = tool_calls result.append(agui_msg) + track_emitted(role, tool_calls) return result diff --git a/python/packages/ag-ui/tests/ag_ui/test_message_adapters.py b/python/packages/ag-ui/tests/ag_ui/test_message_adapters.py index 1749cafe72..a756009df6 100644 --- a/python/packages/ag-ui/tests/ag_ui/test_message_adapters.py +++ b/python/packages/ag-ui/tests/ag_ui/test_message_adapters.py @@ -5,6 +5,7 @@ import base64 import json import logging +from itertools import permutations from typing import Any import pytest @@ -15,6 +16,7 @@ agui_messages_to_agent_framework, agui_messages_to_snapshot_format, extract_text_from_contents, + normalize_agui_input_messages, ) @@ -963,6 +965,344 @@ def test_agent_framework_to_agui_function_result_with_text_preserves_both(): assert text_msg["id"] != tool_msg["id"] +def test_agent_framework_to_agui_function_call_precedes_its_result(): + """A [function_call, function_result] message keeps the call before the result (no orphan).""" + msg = Message( + role="assistant", + contents=[ + Content.from_function_call(call_id="call_a", name="get_weather", arguments={"city": "Seattle"}), + Content.from_function_result(call_id="call_a", result="Sunny"), + ], + message_id="mixed-order-1", + ) + + messages = agent_framework_messages_to_agui([msg]) + + assert len(messages) == 2 + assistant_msg, tool_msg = messages + # The assistant call must be emitted before its result; a result ahead of its + # matching call would be an orphan that providers reject. + assert assistant_msg["role"] == "assistant" + assert [tc["id"] for tc in assistant_msg["tool_calls"]] == ["call_a"] + assert tool_msg["role"] == "tool" + assert tool_msg["toolCallId"] == "call_a" + assert tool_msg["content"] == "Sunny" + # First emitted message keeps the source id; the split-off message gets its own. + assert assistant_msg["id"] == "mixed-order-1" + assert tool_msg["id"] != assistant_msg["id"] + + +def test_agent_framework_to_agui_call_result_text_order_preserved(): + """[function_call, function_result, text] round-trips in order: call, result, then summary text.""" + msg = Message( + role="assistant", + contents=[ + Content.from_function_call(call_id="call_a", name="get_weather", arguments="{}"), + Content.from_function_result(call_id="call_a", result="Sunny"), + Content.from_text("It is sunny."), + ], + message_id="mixed-order-2", + ) + + messages = agent_framework_messages_to_agui([msg]) + + assert [m["role"] for m in messages] == ["assistant", "tool", "assistant"] + assert messages[0]["tool_calls"][0]["id"] == "call_a" + assert messages[1]["toolCallId"] == "call_a" + assert messages[2]["content"] == "It is sunny." + # Only the first emitted message reuses the source id, and all ids are distinct. + assert messages[0]["id"] == "mixed-order-2" + assert len({m["id"] for m in messages}) == 3 + + +def test_agent_framework_to_agui_text_before_result_deferred_after_result(): + """Text preceding a result is emitted AFTER the result, never as an assistant-only message before it.""" + msg = Message( + role="assistant", + contents=[ + Content.from_text("Here is the weather."), + Content.from_function_result(call_id="weather-call", result="Sunny"), + ], + message_id="mixed-text-first", + ) + + messages = agent_framework_messages_to_agui([msg]) + + assert len(messages) == 2 + tool_msg, text_msg = messages + # The tool result comes first; the text-only assistant message follows it, so it can + # never separate a prior outstanding call from this result. + assert tool_msg["role"] == "tool" + assert tool_msg["toolCallId"] == "weather-call" + assert text_msg["role"] == "assistant" + assert "tool_calls" not in text_msg + assert text_msg["content"] == "Here is the weather." + + +def test_agent_framework_to_agui_text_before_result_round_trips_without_dropping(): + """A prior call + a [text, result] message must not drop the result through sanitize_tool_history. + + Regression for the MAF review finding: emitting a text-only assistant message between an + outstanding call and its result made ``_sanitize_tool_history`` clear the pending call and + drop the real result, leaving the provider with an unanswered tool call. + """ + framework_messages = [ + Message( + role="assistant", + contents=[Content.from_function_call(call_id="call_a", name="get_weather", arguments="{}")], + message_id="m1", + ), + Message( + role="assistant", + contents=[ + Content.from_text("Let me check the weather."), + Content.from_function_result(call_id="call_a", result="Sunny"), + ], + message_id="m2", + ), + ] + + agui_messages = agent_framework_messages_to_agui(framework_messages) + + # The call is immediately followed by its result (no assistant message in between). + assert [m["role"] for m in agui_messages] == ["assistant", "tool", "assistant"] + assert agui_messages[0]["tool_calls"][0]["id"] == "call_a" + assert agui_messages[1]["toolCallId"] == "call_a" + + # Round-trip through provider normalization: the real result must survive. + provider_messages, _ = normalize_agui_input_messages(agui_messages, sanitize_tool_history=True) + surviving_result_ids = { + content.call_id + for message in provider_messages + for content in (message.contents or []) + if content.type == "function_result" + } + assert "call_a" in surviving_result_ids + + +def test_agent_framework_to_agui_interleaved_parallel_batch_order_preserved(): + """An interleaved parallel batch keeps every call adjacent to its results. + + A new call (C) appearing before the preceding batch's results are all emitted must not + start a new assistant segment ahead of the still-open results (A, B). The split defers + ``assistant(C)`` until B's result has been emitted. + """ + msg = Message( + role="assistant", + contents=[ + Content.from_function_call(call_id="call_a", name="fa", arguments="{}"), + Content.from_function_call(call_id="call_b", name="fb", arguments="{}"), + Content.from_function_result(call_id="call_a", result="ra"), + Content.from_function_call(call_id="call_c", name="fc", arguments="{}"), + Content.from_function_result(call_id="call_b", result="rb"), + Content.from_function_result(call_id="call_c", result="rc"), + ], + message_id="interleaved-1", + ) + + messages = agent_framework_messages_to_agui([msg]) + + # assistant(A,B) -> tool(A) -> tool(B) -> assistant(C) -> tool(C): the new call C is + # deferred until the open batch {A, B} is fully resolved, so no assistant message ever + # separates B's call from B's result. + assert [m["role"] for m in messages] == ["assistant", "tool", "tool", "assistant", "tool"] + assert [tc["id"] for tc in messages[0]["tool_calls"]] == ["call_a", "call_b"] + assert messages[1]["toolCallId"] == "call_a" + assert messages[2]["toolCallId"] == "call_b" + assert [tc["id"] for tc in messages[3]["tool_calls"]] == ["call_c"] + assert messages[4]["toolCallId"] == "call_c" + # First emitted message keeps the source id; every other id is independent. + assert messages[0]["id"] == "interleaved-1" + assert len({m["id"] for m in messages}) == len(messages) + + +def test_agent_framework_to_agui_interleaved_batch_round_trips_without_dropping(): + """An interleaved parallel batch must not drop any result through sanitize_tool_history. + + Regression for the review finding: with the naive split, ``[call A, call B, result A, + call C, result B, result C]`` became ``[assistant(A,B), tool(A), assistant(C), tool(B), + tool(C)]``; the intervening ``assistant(C)`` cleared the pending call B, so + ``_sanitize_tool_history`` dropped B's real result. + """ + msg = Message( + role="assistant", + contents=[ + Content.from_function_call(call_id="call_a", name="fa", arguments="{}"), + Content.from_function_call(call_id="call_b", name="fb", arguments="{}"), + Content.from_function_result(call_id="call_a", result="ra"), + Content.from_function_call(call_id="call_c", name="fc", arguments="{}"), + Content.from_function_result(call_id="call_b", result="rb"), + Content.from_function_result(call_id="call_c", result="rc"), + ], + message_id="interleaved-2", + ) + + agui_messages = agent_framework_messages_to_agui([msg]) + + # Round-trip through provider normalization: every result must survive. + provider_messages, _ = normalize_agui_input_messages(agui_messages, sanitize_tool_history=True) + surviving_result_ids = { + content.call_id + for message in provider_messages + for content in (message.contents or []) + if content.type == "function_result" + } + assert surviving_result_ids == {"call_a", "call_b", "call_c"} + + +def test_agent_framework_to_agui_pending_call_carries_across_messages(): + """A call opened by an earlier message stays open across the conversion. + + Regression: the unresolved-call set used to be rebuilt per mixed message, so a prior + ``assistant(call A)`` followed by ``[call C, result A, result C]`` emitted + ``assistant(A), assistant(C), tool(A), tool(C)``; the intervening ``assistant(C)`` + cleared pending call A and ``_sanitize_tool_history`` dropped A's real result. + """ + framework_messages = [ + Message( + role="assistant", + contents=[Content.from_function_call(call_id="call_a", name="fa", arguments="{}")], + message_id="m1", + ), + Message( + role="assistant", + contents=[ + Content.from_function_call(call_id="call_c", name="fc", arguments="{}"), + Content.from_function_result(call_id="call_a", result="ra"), + Content.from_function_result(call_id="call_c", result="rc"), + ], + message_id="m2", + ), + ] + + agui_messages = agent_framework_messages_to_agui(framework_messages) + + # assistant(A) -> tool(A) -> assistant(C) -> tool(C): the new call C is deferred until + # the earlier call A has received its result. + assert [m["role"] for m in agui_messages] == ["assistant", "tool", "assistant", "tool"] + assert [tc["id"] for tc in agui_messages[0]["tool_calls"]] == ["call_a"] + assert agui_messages[1]["toolCallId"] == "call_a" + assert [tc["id"] for tc in agui_messages[2]["tool_calls"]] == ["call_c"] + assert agui_messages[3]["toolCallId"] == "call_c" + + provider_messages, _ = normalize_agui_input_messages(agui_messages, sanitize_tool_history=True) + surviving_result_ids = { + content.call_id + for message in provider_messages + for content in (message.contents or []) + if content.type == "function_result" + } + assert surviving_result_ids == {"call_a", "call_c"} + + +def test_agent_framework_to_agui_result_for_buffered_call_waits_for_its_call(): + """A result whose own call is still buffered is held until that call is emitted. + + Regression: for ``[call A, call B, result A, call C, result C, result B]`` the split + correctly deferred ``assistant(C)`` but still emitted ``tool(C)`` immediately, leaving + C's result ahead of its call so ``_sanitize_tool_history`` dropped it. + """ + msg = Message( + role="assistant", + contents=[ + Content.from_function_call(call_id="call_a", name="fa", arguments="{}"), + Content.from_function_call(call_id="call_b", name="fb", arguments="{}"), + Content.from_function_result(call_id="call_a", result="ra"), + Content.from_function_call(call_id="call_c", name="fc", arguments="{}"), + Content.from_function_result(call_id="call_c", result="rc"), + Content.from_function_result(call_id="call_b", result="rb"), + ], + message_id="out-of-order-1", + ) + + agui_messages = agent_framework_messages_to_agui([msg]) + + # Every tool message must come after an assistant message that declared its call. + declared: set[str] = set() + for agui_msg in agui_messages: + if agui_msg["role"] == "assistant": + declared.update(tc["id"] for tc in agui_msg.get("tool_calls") or []) + else: + assert agui_msg["toolCallId"] in declared, f"{agui_msg['toolCallId']} emitted before its call" + + provider_messages, _ = normalize_agui_input_messages(agui_messages, sanitize_tool_history=True) + surviving_result_ids = { + content.call_id + for message in provider_messages + for content in (message.contents or []) + if content.type == "function_result" + } + assert surviving_result_ids == {"call_a", "call_b", "call_c"} + + +def _call_result_interleavings(call_ids: list[str]) -> list[tuple[tuple[str, str], ...]]: + """Every ordering of the given calls and their results where each call precedes its result.""" + events = [("call", call_id) for call_id in call_ids] + [("result", call_id) for call_id in call_ids] + orderings: list[tuple[tuple[str, str], ...]] = [] + for ordering in dict.fromkeys(permutations(events)): + emitted: set[str] = set() + for kind, call_id in ordering: + if kind == "call": + emitted.add(call_id) + elif call_id not in emitted: + break + else: + orderings.append(ordering) + return orderings + + +def _contents_for(ordering: tuple[tuple[str, str], ...]) -> list[Content]: + return [ + Content.from_function_call(call_id=call_id, name="f", arguments="{}") + if kind == "call" + else Content.from_function_result(call_id=call_id, result=f"r-{call_id}") + for kind, call_id in ordering + ] + + +def _assert_no_result_is_orphaned(agui_messages: list[dict[str, Any]], call_ids: list[str]) -> None: + """Every tool message follows an assistant message declaring its call, and no result is lost.""" + declared: set[str] = set() + for agui_msg in agui_messages: + if agui_msg["role"] == "assistant": + declared.update(tool_call["id"] for tool_call in agui_msg.get("tool_calls") or []) + elif agui_msg["role"] == "tool": + assert agui_msg["toolCallId"] in declared, f"{agui_msg['toolCallId']} emitted before its call" + + provider_messages, _ = normalize_agui_input_messages(agui_messages, sanitize_tool_history=True) + surviving = { + content.call_id + for message in provider_messages + for content in (message.contents or []) + if content.type == "function_result" + } + assert surviving == set(call_ids) + + +@pytest.mark.parametrize("call_ids", [["call_a", "call_b"], ["call_a", "call_b", "call_c"]]) +def test_agent_framework_to_agui_no_interleaving_drops_a_result(call_ids: list[str]): + """Exhaustive: no ordering of calls/results may lose a result or emit one ahead of its call. + + ``_sanitize_tool_history`` resets its pending-call set on every non-tool message, so any + assistant message emitted between a call and that call's result causes the result to be + dropped. This sweeps every valid interleaving -- in a single mixed message and split across + two messages -- because that whole class of ordering bug has recurred here repeatedly. + """ + for ordering in _call_result_interleavings(call_ids): + contents = _contents_for(ordering) + + single = [Message(role="assistant", contents=contents, message_id="m1")] + _assert_no_result_is_orphaned(agent_framework_messages_to_agui(single), call_ids) + + # Splitting the leading content into its own message exercises the cross-message + # pending-call state that a per-message set would lose. + split = [ + Message(role="assistant", contents=contents[:1], message_id="m1"), + Message(role="assistant", contents=contents[1:], message_id="m2"), + ] + _assert_no_result_is_orphaned(agent_framework_messages_to_agui(split), call_ids) + + # Additional tests for better coverage