Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
85 changes: 69 additions & 16 deletions python/packages/core/agent_framework/_workflows/_agent.py
Original file line number Diff line number Diff line change
Expand Up @@ -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."""

Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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 [
Expand All @@ -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),
Expand All @@ -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),
Expand All @@ -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),
Expand Down
248 changes: 247 additions & 1 deletion python/packages/core/tests/workflow/test_workflow_agent.py
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand Down Expand Up @@ -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"
Loading