diff --git a/python/packages/core/agent_framework/_workflows/_agent.py b/python/packages/core/agent_framework/_workflows/_agent.py index f4aa653310d..03166466788 100644 --- a/python/packages/core/agent_framework/_workflows/_agent.py +++ b/python/packages/core/agent_framework/_workflows/_agent.py @@ -52,6 +52,22 @@ logger = logging.getLogger(__name__) +def _is_orphaned_function_call(message: "Message") -> bool: + """True when an assistant message is only a function call whose tool result + is not assistant-role, so keeping the call alone would produce an invalid + (unpaired) transcript for providers that validate call/result history.""" + if message.role != "assistant": + return False + contents = list(getattr(message, "contents", []) or []) + if not contents: + return False + # Orphaned only when every content is a function-call envelope + return all( + getattr(content, "type", None) in ("function_call", "function_approval_response") + for content in contents + ) + + class WorkflowAgent(BaseAgent): """An `Agent` subclass that wraps a workflow and exposes it as an agent.""" @@ -587,23 +603,51 @@ def _convert_workflow_events_to_agent_response( ) if isinstance(data, AgentResponse): - messages.extend(data.messages) - raw_representations.append(data.raw_representation) - merged_usage = add_usage_details(merged_usage, data.usage_details) - latest_created_at = ( - data.created_at - if not latest_created_at - else max(latest_created_at, data.created_at) - if data.created_at - else latest_created_at - ) + # Filter to only assistant messages — system, tool, and user messages + # are intentionally excluded. System prompts and tool results are + # internal workflow artifacts; user messages would be re-emitted + # (e.g., from GroupChat orchestrators that include full conversation history). + # Assistant messages that are bare function calls are also dropped + # when their tool result is not assistant-role: a call without its + # result is an invalid transcript for providers that validate + # call/result pairing on replay. + assistant_messages = [ + msg for msg in data.messages + if msg.role == "assistant" and not _is_orphaned_function_call(msg) + ] + if assistant_messages: + messages.extend(assistant_messages) + raw_representations.append(data.raw_representation) + merged_usage = add_usage_details(merged_usage, data.usage_details) + latest_created_at = ( + data.created_at + if not latest_created_at + else max(latest_created_at, data.created_at) + if data.created_at + else latest_created_at + ) elif isinstance(data, Message): - messages.append(data) - raw_representations.append(data.raw_representation) + if data.role == "assistant": + messages.append(data) + raw_representations.append(data.raw_representation) elif is_instance_of(data, list[Message]): chat_messages = cast(list[Message], data) - messages.extend(chat_messages) - raw_representations.append(data) + # Keep tool results that pair with surviving assistant calls so the + # transcript stays valid; drop user/system and orphaned calls. + assistant_messages = [ + msg for msg in chat_messages + if msg.role == "assistant" and not _is_orphaned_function_call(msg) + ] + if assistant_messages: + messages.extend(assistant_messages) + # raw_representation of a filtered list must not leak the + # non-assistant entries the public messages list dropped. + if len(assistant_messages) == len(chat_messages): + raw_representations.append(data) + else: + raw_representations.extend( + msg.raw_representation for msg in assistant_messages + ) else: contents = self._extract_contents(data) if not contents: @@ -654,6 +698,9 @@ def _convert_workflow_event_to_agent_response_updates( executor_id = event.executor_id if isinstance(data, AgentResponseUpdate): + # Filter out non-assistant updates (e.g. user input echoed back) + if data.role is not None and data.role != "assistant": + return [] # Construct a fresh AgentResponseUpdate so we don't mutate a payload # that AgentExecutor still holds a reference to in its `updates` list. return [ @@ -676,9 +723,11 @@ def _convert_workflow_event_to_agent_response_updates( ) ] if isinstance(data, AgentResponse): - # Convert each message in AgentResponse to an AgentResponseUpdate + # Convert each assistant message in AgentResponse to an AgentResponseUpdate updates: list[AgentResponseUpdate] = [] for msg in data.messages: + if msg.role != "assistant": + continue updates.append( AgentResponseUpdate( contents=list(msg.contents), @@ -698,6 +747,8 @@ def _convert_workflow_event_to_agent_response_updates( updates[-1].additional_properties = dict(data.additional_properties) return updates if isinstance(data, Message): + if data.role != "assistant": + return [] return [ AgentResponseUpdate( contents=list(data.contents), @@ -710,10 +761,12 @@ def _convert_workflow_event_to_agent_response_updates( ) ] if is_instance_of(data, list[Message]): - # Convert each Message to an AgentResponseUpdate + # Convert each assistant Message to an AgentResponseUpdate chat_messages = cast(list[Message], data) updates = [] for msg in chat_messages: + if msg.role != "assistant": + continue updates.append( AgentResponseUpdate( contents=list(msg.contents), diff --git a/python/packages/core/tests/workflow/test_workflow_agent.py b/python/packages/core/tests/workflow/test_workflow_agent.py index 0ef70ca2f6e..14e43502801 100644 --- a/python/packages/core/tests/workflow/test_workflow_agent.py +++ b/python/packages/core/tests/workflow/test_workflow_agent.py @@ -1105,7 +1105,7 @@ async def test_workflow_as_agent_yield_output_with_list_of_chat_messages(self) - async def list_yielding_executor(messages: list[Message], ctx: WorkflowContext[Never, list[Message]]) -> None: # type: ignore[valid-type] # Yield a list of Messages (as SequentialBuilder does) msg_list = [ - Message(role="user", contents=["first message"]), + Message(role="assistant", contents=["first message"]), Message(role="assistant", contents=["second message"]), Message( role="assistant", @@ -2597,3 +2597,249 @@ async def start(messages: list[Message], ctx: WorkflowContext[AgentExecutorReque pending = await workflow._runner_context.get_pending_request_info_events() # The agent's approval id is used as the workflow's pending request id. assert list(pending.keys()) == [approval_id] + + async def test_workflow_as_agent_filters_non_assistant_messages_from_agent_response(self) -> None: + """Verify WorkflowAgent filters user, system, and tool messages from AgentResponse.""" + + @executor + async def mixed_agent_response_executor( + messages: list[Message], + ctx: WorkflowContext[Never, AgentResponse], # type: ignore[valid-type] + ) -> None: + response = AgentResponse( + messages=[ + Message(role="system", contents=["System instructions"]), + Message(role="user", contents=["User input question"]), + Message(role="assistant", contents=["Assistant answer"], author_name="Teacher"), + Message(role="tool", contents=["Tool execution result"]), + ] + ) + await ctx.yield_output(response) + + workflow = WorkflowBuilder(start_executor=mixed_agent_response_executor).build() + agent = workflow.as_agent("mixed-response-agent") + + # Test streaming path + updates: list[AgentResponseUpdate] = [] + async for chunk in agent.run("hello", stream=True): + updates.append(chunk) + + assert len(updates) == 1 + assert updates[0].role == "assistant" + assert updates[0].text == "Assistant answer" + assert updates[0].author_name == "Teacher" + + # Test non-streaming path + result = await agent.run("hello") + assert len(result.messages) == 1 + assert result.messages[0].role == "assistant" + assert result.messages[0].text == "Assistant answer" + assert result.messages[0].author_name == "Teacher" + + async def test_workflow_as_agent_filters_non_assistant_messages_from_list_of_messages(self) -> None: + """Verify WorkflowAgent filters user, system, and tool messages from list[Message].""" + + @executor + async def mixed_list_executor( + messages: list[Message], + ctx: WorkflowContext[Never, list[Message]], # type: ignore[valid-type] + ) -> None: + await ctx.yield_output([ + Message(role="user", contents=["what is 2+2?"]), + Message(role="assistant", contents=["4"], author_name="Maths"), + Message(role="system", contents=["system prompt"]), + Message(role="assistant", contents=["four"], author_name="English"), + ]) + + workflow = WorkflowBuilder(start_executor=mixed_list_executor).build() + agent = workflow.as_agent("mixed-list-agent") + + # Test streaming path + updates: list[AgentResponseUpdate] = [] + async for chunk in agent.run("calc", stream=True): + updates.append(chunk) + + assert len(updates) == 2 + for update in updates: + assert update.role == "assistant" + assert updates[0].author_name == "Maths" + assert updates[0].text == "4" + assert updates[1].author_name == "English" + assert updates[1].text == "four" + + # Test non-streaming path + result = await agent.run("calc") + assert len(result.messages) == 2 + for message in result.messages: + assert message.role == "assistant" + assert result.messages[0].author_name == "Maths" + assert result.messages[0].text == "4" + assert result.messages[1].author_name == "English" + assert result.messages[1].text == "four" + + # raw_representation of the non-streaming result must not leak the + # filtered-out user/system messages through the public payload. + for rep in result.raw_representation or []: + assert rep is None or (isinstance(rep, Message) and rep.role == "assistant") + + async def test_workflow_as_agent_drops_orphaned_function_calls(self) -> None: + """assistant(function_call) whose tool result is tool-role must not survive filtering. + + Keeping the call without its result would produce an invalid transcript for + providers that validate call/result pairing on replay (eavanvalkenburg's review). + """ + + @executor + async def tool_transcript_executor( + messages: list[Message], + ctx: WorkflowContext[Never, list[Message]], # type: ignore[valid-type] + ) -> None: + await ctx.yield_output([ + Message(role="assistant", contents=[ + Content.from_function_call(call_id="call-1", name="get_weather", arguments={"city": "Paris"}), + ]), + Message(role="tool", contents=[ + Content.from_function_result(call_id="call-1", result="18C"), + ]), + Message(role="assistant", contents=[Content.from_text("It is 18C in Paris.")]), + ]) + + workflow = WorkflowBuilder(start_executor=tool_transcript_executor).build() + agent = workflow.as_agent("tool-transcript-agent") + + result = await agent.run("weather") + + # The orphaned call is dropped; the final user-facing answer survives. + assert all(msg.role == "assistant" for msg in result.messages) + assert not any( + getattr(content, "type", None) == "function_call" + for msg in result.messages + for content in msg.contents + ) + assert any( + "18C" in (getattr(content, "text", "") or "") + for msg in result.messages + for content in msg.contents + ) + + async def test_workflow_as_agent_filters_single_non_assistant_message(self) -> None: + """Verify WorkflowAgent filters a single Message when role is not assistant.""" + + @executor + async def user_message_executor( + messages: list[Message], + ctx: WorkflowContext[Never, Message], # type: ignore[valid-type] + ) -> None: + await ctx.yield_output(Message(role="user", contents=["echoed user message"])) + + workflow = WorkflowBuilder(start_executor=user_message_executor).build() + agent = workflow.as_agent("user-msg-agent") + + # Streaming should yield no updates + updates: list[AgentResponseUpdate] = [] + async for chunk in agent.run("test", stream=True): + updates.append(chunk) + assert len(updates) == 0 + + # Non-streaming should produce empty messages list + result = await agent.run("test") + assert len(result.messages) == 0 + + async def test_workflow_as_agent_filters_user_agent_response_update(self) -> None: + """Verify WorkflowAgent drops AgentResponseUpdate when role is user.""" + + @executor + async def update_yielding_executor( + messages: list[Message], + ctx: WorkflowContext[Never, AgentResponseUpdate], # type: ignore[valid-type] + ) -> None: + await ctx.yield_output(AgentResponseUpdate(contents=[Content.from_text(text="echo")], role="user")) + await ctx.yield_output(AgentResponseUpdate(contents=[Content.from_text(text="answer")], role="assistant")) + + workflow = WorkflowBuilder(start_executor=update_yielding_executor).build() + agent = workflow.as_agent("update-agent") + + updates: list[AgentResponseUpdate] = [] + async for chunk in agent.run("test", stream=True): + updates.append(chunk) + + assert len(updates) == 1 + assert updates[0].role == "assistant" + assert updates[0].text == "answer" + + async def test_workflow_as_agent_empty_after_filtering(self) -> None: + """Verify WorkflowAgent handles all non-assistant messages without crashing.""" + + @executor + async def non_assistant_only_executor( + messages: list[Message], + ctx: WorkflowContext[Never, AgentResponse], # type: ignore[valid-type] + ) -> None: + response = AgentResponse( + messages=[ + Message(role="user", contents=["user msg"]), + Message(role="system", contents=["system msg"]), + Message(role="tool", contents=["tool msg"]), + ] + ) + await ctx.yield_output(response) + + workflow = WorkflowBuilder(start_executor=non_assistant_only_executor).build() + agent = workflow.as_agent("all-filtered-agent") + + result = await agent.run("test") + assert len(result.messages) == 0 + assert not result.raw_representation + + async def test_workflow_as_agent_multi_turn_user_input_not_compounded(self) -> None: + """Verify user messages in conversation history do not compound into responses across turns.""" + + class HistoryYieldingExecutor(Executor): + @handler + async def handle_messages( + self, + messages: list[Message], + ctx: WorkflowContext[Never, AgentResponse], # type: ignore[valid-type] + ) -> None: + user_text = messages[-1].text or "" + # Simulates orchestrators that include full conversation history in output + full_history = [ + Message(role="user", contents=[user_text]), + Message(role="assistant", contents=[f"Answer: {user_text}"], author_name="Agent"), + ] + await ctx.yield_output(AgentResponse(messages=full_history)) + + workflow = WorkflowBuilder(start_executor=HistoryYieldingExecutor(id="history-exec")).build() + agent = workflow.as_agent("history-agent") + session = AgentSession() + + # Turn 1 non-streaming + resp1 = await agent.run("first_query", session=session) + assert len(resp1.messages) == 1 + assert resp1.messages[0].role == "assistant" + assert resp1.text == "Answer: first_query" + assert "first_query" not in (resp1.text.replace("Answer: first_query", "")) + + # Turn 2 non-streaming: first_query must not bleed into turn 2 + resp2 = await agent.run("second_query", session=session) + assert len(resp2.messages) == 1 + assert resp2.messages[0].role == "assistant" + assert resp2.text == "Answer: second_query" + + # Streaming check + streaming_agent = workflow.as_agent("streaming-history-agent") + streaming_session = AgentSession() + + chunks1: list[AgentResponseUpdate] = [] + async for chunk in streaming_agent.run("stream_q1", stream=True, session=streaming_session): + chunks1.append(chunk) + assert len(chunks1) == 1 + assert chunks1[0].role == "assistant" + assert chunks1[0].text == "Answer: stream_q1" + + chunks2: list[AgentResponseUpdate] = [] + async for chunk in streaming_agent.run("stream_q2", stream=True, session=streaming_session): + chunks2.append(chunk) + assert len(chunks2) == 1 + assert chunks2[0].role == "assistant" + assert chunks2[0].text == "Answer: stream_q2"