mirror of
https://github.com/microsoft/agent-framework.git
synced 2026-06-16 21:04:09 +08:00
Python: Add Handoff orchestration pattern support (#1469)
* Add Handoff orchestration pattern support * PR feedback * Use AOAI client in samples * Adjust to tool * Handoff to sub-agent via ai function * PR feedback * More cleanup * Improvements * PR feedback cleanup * Add handoff migration sample. * Remove type ignore * fix markdown link formatting * Remove readme link for non-existent sample
This commit is contained in:
committed by
GitHub
Unverified
parent
4554de00ab
commit
b66619a544
@@ -52,6 +52,7 @@ from ._executor import (
|
||||
handler,
|
||||
)
|
||||
from ._function_executor import FunctionExecutor, executor
|
||||
from ._handoff import HandoffBuilder, HandoffUserInputRequest
|
||||
from ._magentic import (
|
||||
MagenticAgentDeltaEvent,
|
||||
MagenticAgentExecutor,
|
||||
@@ -127,6 +128,8 @@ __all__ = [
|
||||
"FileCheckpointStorage",
|
||||
"FunctionExecutor",
|
||||
"GraphConnectivityError",
|
||||
"HandoffBuilder",
|
||||
"HandoffUserInputRequest",
|
||||
"InMemoryCheckpointStorage",
|
||||
"InProcRunnerContext",
|
||||
"MagenticAgentDeltaEvent",
|
||||
|
||||
@@ -50,6 +50,7 @@ from ._executor import (
|
||||
handler,
|
||||
)
|
||||
from ._function_executor import FunctionExecutor, executor
|
||||
from ._handoff import HandoffBuilder, HandoffUserInputRequest
|
||||
from ._magentic import (
|
||||
MagenticAgentDeltaEvent,
|
||||
MagenticAgentExecutor,
|
||||
@@ -125,6 +126,8 @@ __all__ = [
|
||||
"FileCheckpointStorage",
|
||||
"FunctionExecutor",
|
||||
"GraphConnectivityError",
|
||||
"HandoffBuilder",
|
||||
"HandoffUserInputRequest",
|
||||
"InMemoryCheckpointStorage",
|
||||
"InProcRunnerContext",
|
||||
"MagenticAgentDeltaEvent",
|
||||
|
||||
@@ -14,6 +14,8 @@ from agent_framework import (
|
||||
AgentThread,
|
||||
BaseAgent,
|
||||
ChatMessage,
|
||||
FunctionApprovalRequestContent,
|
||||
FunctionApprovalResponseContent,
|
||||
FunctionCallContent,
|
||||
FunctionResultContent,
|
||||
Role,
|
||||
@@ -266,16 +268,20 @@ class WorkflowAgent(BaseAgent):
|
||||
# Store the pending request for later correlation
|
||||
self.pending_requests[request_id] = event
|
||||
|
||||
# Convert to function call content
|
||||
# TODO(ekzhu): update this to FunctionApprovalRequestContent
|
||||
# monitor: https://github.com/microsoft/agent-framework/issues/285
|
||||
args = self.RequestInfoFunctionArgs(request_id=request_id, data=event.data).to_dict()
|
||||
|
||||
function_call = FunctionCallContent(
|
||||
call_id=request_id,
|
||||
name=self.REQUEST_INFO_FUNCTION_NAME,
|
||||
arguments=self.RequestInfoFunctionArgs(request_id=request_id, data=event.data).to_dict(),
|
||||
arguments=args,
|
||||
)
|
||||
approval_request = FunctionApprovalRequestContent(
|
||||
id=request_id,
|
||||
function_call=function_call,
|
||||
additional_properties={"request_id": request_id},
|
||||
)
|
||||
return AgentRunResponseUpdate(
|
||||
contents=[function_call],
|
||||
contents=[function_call, approval_request],
|
||||
role=Role.ASSISTANT,
|
||||
author_name=self.name,
|
||||
response_id=response_id,
|
||||
@@ -293,26 +299,45 @@ class WorkflowAgent(BaseAgent):
|
||||
function_responses: dict[str, Any] = {}
|
||||
for message in input_messages:
|
||||
for content in message.contents:
|
||||
# TODO(ekzhu): update this to FunctionApprovalResponseContent
|
||||
# monitor: https://github.com/microsoft/agent-framework/issues/285
|
||||
if isinstance(content, FunctionResultContent):
|
||||
if isinstance(content, FunctionApprovalResponseContent):
|
||||
# Parse the function arguments to recover request payload
|
||||
arguments_payload = content.function_call.arguments
|
||||
if isinstance(arguments_payload, str):
|
||||
try:
|
||||
parsed_args = self.RequestInfoFunctionArgs.from_json(arguments_payload)
|
||||
except ValueError as exc:
|
||||
raise AgentExecutionException(
|
||||
"FunctionApprovalResponseContent arguments must decode to a mapping."
|
||||
) from exc
|
||||
elif isinstance(arguments_payload, dict):
|
||||
parsed_args = self.RequestInfoFunctionArgs.from_dict(arguments_payload)
|
||||
else:
|
||||
raise AgentExecutionException(
|
||||
"FunctionApprovalResponseContent arguments must be a mapping or JSON string."
|
||||
)
|
||||
|
||||
request_id = parsed_args.request_id or content.id
|
||||
if not content.approved:
|
||||
raise AgentExecutionException(f"Request '{request_id}' was not approved by the caller.")
|
||||
|
||||
if request_id in self.pending_requests:
|
||||
function_responses[request_id] = parsed_args.data
|
||||
elif bool(self.pending_requests):
|
||||
raise AgentExecutionException(
|
||||
"Only responses for pending requests are allowed when there are outstanding approvals."
|
||||
)
|
||||
elif isinstance(content, FunctionResultContent):
|
||||
request_id = content.call_id
|
||||
# Check if we have a pending request for this call_id
|
||||
if request_id in self.pending_requests:
|
||||
response_data = content.result if hasattr(content, "result") else str(content)
|
||||
function_responses[request_id] = response_data
|
||||
elif bool(self.pending_requests):
|
||||
# Function result for unknown request when we have pending requests - this is an error
|
||||
raise AgentExecutionException(
|
||||
"Only FunctionResultContent for pending requests is allowed in input messages "
|
||||
"when there are pending requests."
|
||||
"Only function responses for pending requests are allowed while requests are outstanding."
|
||||
)
|
||||
else:
|
||||
if bool(self.pending_requests):
|
||||
# Non-function content when we have pending requests - this is an error
|
||||
raise AgentExecutionException(
|
||||
"Only FunctionResultContent is allowed in input messages when there are pending requests."
|
||||
)
|
||||
raise AgentExecutionException("Unexpected content type while awaiting request info responses.")
|
||||
return function_responses
|
||||
|
||||
class _ResponseState(TypedDict):
|
||||
|
||||
@@ -0,0 +1,77 @@
|
||||
# Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
from collections.abc import Iterable
|
||||
from typing import Any, cast
|
||||
|
||||
from agent_framework import ChatMessage, Role
|
||||
|
||||
from ._checkpoint_encoding import decode_checkpoint_value, encode_checkpoint_value
|
||||
|
||||
"""Utilities for serializing and deserializing chat conversations for persistence.
|
||||
|
||||
These helpers convert rich `ChatMessage` instances to checkpoint-friendly payloads
|
||||
using the same encoding primitives as the workflow runner. This preserves
|
||||
`additional_properties` and other metadata without relying on unsafe mechanisms
|
||||
such as pickling.
|
||||
"""
|
||||
|
||||
|
||||
def encode_chat_messages(messages: Iterable[ChatMessage]) -> list[dict[str, Any]]:
|
||||
"""Serialize chat messages into checkpoint-safe payloads."""
|
||||
encoded: list[dict[str, Any]] = []
|
||||
for message in messages:
|
||||
encoded.append({
|
||||
"role": encode_checkpoint_value(message.role),
|
||||
"contents": [encode_checkpoint_value(content) for content in message.contents],
|
||||
"author_name": message.author_name,
|
||||
"message_id": message.message_id,
|
||||
"additional_properties": {
|
||||
key: encode_checkpoint_value(value) for key, value in message.additional_properties.items()
|
||||
},
|
||||
})
|
||||
return encoded
|
||||
|
||||
|
||||
def decode_chat_messages(payload: Iterable[dict[str, Any]]) -> list[ChatMessage]:
|
||||
"""Restore chat messages from checkpoint-safe payloads."""
|
||||
restored: list[ChatMessage] = []
|
||||
for item in payload:
|
||||
if not isinstance(item, dict):
|
||||
continue
|
||||
|
||||
role_value = decode_checkpoint_value(item.get("role"))
|
||||
if isinstance(role_value, Role):
|
||||
role = role_value
|
||||
elif isinstance(role_value, dict):
|
||||
role_dict = cast(dict[str, Any], role_value)
|
||||
role = Role.from_dict(role_dict)
|
||||
elif isinstance(role_value, str):
|
||||
role = Role(value=role_value)
|
||||
else:
|
||||
role = Role.ASSISTANT
|
||||
|
||||
contents_field = item.get("contents", [])
|
||||
contents: list[Any] = []
|
||||
if isinstance(contents_field, list):
|
||||
contents_iter: list[Any] = contents_field # type: ignore[assignment]
|
||||
for entry in contents_iter:
|
||||
decoded_entry: Any = decode_checkpoint_value(entry)
|
||||
contents.append(decoded_entry)
|
||||
|
||||
additional_field = item.get("additional_properties", {})
|
||||
additional: dict[str, Any] = {}
|
||||
if isinstance(additional_field, dict):
|
||||
additional_dict = cast(dict[str, Any], additional_field)
|
||||
for key, value in additional_dict.items():
|
||||
additional[key] = decode_checkpoint_value(value)
|
||||
|
||||
restored.append(
|
||||
ChatMessage(
|
||||
role=role,
|
||||
contents=contents,
|
||||
author_name=item.get("author_name"),
|
||||
message_id=item.get("message_id"),
|
||||
additional_properties=additional,
|
||||
)
|
||||
)
|
||||
return restored
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,359 @@
|
||||
# Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
from collections.abc import AsyncIterable, AsyncIterator
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, cast
|
||||
|
||||
import pytest
|
||||
|
||||
from agent_framework import (
|
||||
AgentRunResponse,
|
||||
AgentRunResponseUpdate,
|
||||
BaseAgent,
|
||||
ChatMessage,
|
||||
FunctionCallContent,
|
||||
HandoffBuilder,
|
||||
HandoffUserInputRequest,
|
||||
RequestInfoEvent,
|
||||
Role,
|
||||
TextContent,
|
||||
WorkflowEvent,
|
||||
WorkflowOutputEvent,
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class _ComplexMetadata:
|
||||
reason: str
|
||||
payload: dict[str, str]
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def complex_metadata() -> _ComplexMetadata:
|
||||
return _ComplexMetadata(reason="route", payload={"code": "X1"})
|
||||
|
||||
|
||||
def _metadata_from_conversation(conversation: list[ChatMessage], key: str) -> list[object]:
|
||||
return [msg.additional_properties[key] for msg in conversation if key in msg.additional_properties]
|
||||
|
||||
|
||||
def _conversation_debug(conversation: list[ChatMessage]) -> list[tuple[str, str | None, str]]:
|
||||
return [
|
||||
(msg.role.value if hasattr(msg.role, "value") else str(msg.role), msg.author_name, msg.text)
|
||||
for msg in conversation
|
||||
]
|
||||
|
||||
|
||||
class _RecordingAgent(BaseAgent):
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
name: str,
|
||||
handoff_to: str | None = None,
|
||||
text_handoff: bool = False,
|
||||
extra_properties: dict[str, object] | None = None,
|
||||
) -> None:
|
||||
super().__init__(id=name, name=name, display_name=name)
|
||||
self.handoff_to = handoff_to
|
||||
self.calls: list[list[ChatMessage]] = []
|
||||
self._text_handoff = text_handoff
|
||||
self._extra_properties = dict(extra_properties or {})
|
||||
self._call_index = 0
|
||||
|
||||
async def run( # type: ignore[override]
|
||||
self,
|
||||
messages: str | ChatMessage | list[str] | list[ChatMessage] | None = None,
|
||||
*,
|
||||
thread: Any = None,
|
||||
**kwargs: Any,
|
||||
) -> AgentRunResponse:
|
||||
conversation = _normalise(messages)
|
||||
self.calls.append(conversation)
|
||||
additional_properties = _merge_additional_properties(
|
||||
self.handoff_to, self._text_handoff, self._extra_properties
|
||||
)
|
||||
contents = _build_reply_contents(self.name, self.handoff_to, self._text_handoff, self._next_call_id())
|
||||
reply = ChatMessage(
|
||||
role=Role.ASSISTANT,
|
||||
contents=contents,
|
||||
author_name=self.display_name,
|
||||
additional_properties=additional_properties,
|
||||
)
|
||||
return AgentRunResponse(messages=[reply])
|
||||
|
||||
async def run_stream( # type: ignore[override]
|
||||
self,
|
||||
messages: str | ChatMessage | list[str] | list[ChatMessage] | None = None,
|
||||
*,
|
||||
thread: Any = None,
|
||||
**kwargs: Any,
|
||||
) -> AsyncIterator[AgentRunResponseUpdate]:
|
||||
conversation = _normalise(messages)
|
||||
self.calls.append(conversation)
|
||||
additional_props = _merge_additional_properties(self.handoff_to, self._text_handoff, self._extra_properties)
|
||||
contents = _build_reply_contents(self.name, self.handoff_to, self._text_handoff, self._next_call_id())
|
||||
yield AgentRunResponseUpdate(
|
||||
contents=contents,
|
||||
role=Role.ASSISTANT,
|
||||
additional_properties=additional_props,
|
||||
)
|
||||
|
||||
def _next_call_id(self) -> str | None:
|
||||
if not self.handoff_to:
|
||||
return None
|
||||
call_id = f"{self.id}-handoff-{self._call_index}"
|
||||
self._call_index += 1
|
||||
return call_id
|
||||
|
||||
|
||||
def _merge_additional_properties(
|
||||
handoff_to: str | None, use_text_hint: bool, extras: dict[str, object]
|
||||
) -> dict[str, object]:
|
||||
additional_properties: dict[str, object] = {}
|
||||
if handoff_to and not use_text_hint:
|
||||
additional_properties["handoff_to"] = handoff_to
|
||||
additional_properties.update(extras)
|
||||
return additional_properties
|
||||
|
||||
|
||||
def _build_reply_contents(
|
||||
agent_name: str,
|
||||
handoff_to: str | None,
|
||||
use_text_hint: bool,
|
||||
call_id: str | None,
|
||||
) -> list[TextContent | FunctionCallContent]:
|
||||
contents: list[TextContent | FunctionCallContent] = []
|
||||
if handoff_to and call_id:
|
||||
contents.append(
|
||||
FunctionCallContent(call_id=call_id, name=f"handoff_to_{handoff_to}", arguments={"handoff_to": handoff_to})
|
||||
)
|
||||
text = f"{agent_name} reply"
|
||||
if use_text_hint and handoff_to:
|
||||
text += f"\nHANDOFF_TO: {handoff_to}"
|
||||
contents.append(TextContent(text=text))
|
||||
return contents
|
||||
|
||||
|
||||
def _normalise(messages: str | ChatMessage | list[str] | list[ChatMessage] | None) -> list[ChatMessage]:
|
||||
if isinstance(messages, list):
|
||||
result: list[ChatMessage] = []
|
||||
for msg in messages:
|
||||
if isinstance(msg, ChatMessage):
|
||||
result.append(msg)
|
||||
elif isinstance(msg, str):
|
||||
result.append(ChatMessage(Role.USER, text=msg))
|
||||
return result
|
||||
if isinstance(messages, ChatMessage):
|
||||
return [messages]
|
||||
if isinstance(messages, str):
|
||||
return [ChatMessage(Role.USER, text=messages)]
|
||||
return []
|
||||
|
||||
|
||||
async def _drain(stream: AsyncIterable[WorkflowEvent]) -> list[WorkflowEvent]:
|
||||
return [event async for event in stream]
|
||||
|
||||
|
||||
async def test_handoff_routes_to_specialist_and_requests_user_input():
|
||||
triage = _RecordingAgent(name="triage", handoff_to="specialist")
|
||||
specialist = _RecordingAgent(name="specialist")
|
||||
|
||||
workflow = HandoffBuilder(participants=[triage, specialist]).set_coordinator("triage").build()
|
||||
|
||||
events = await _drain(workflow.run_stream("Need help with a refund"))
|
||||
|
||||
assert triage.calls, "Starting agent should receive initial conversation"
|
||||
assert specialist.calls, "Specialist should be invoked after handoff"
|
||||
assert len(specialist.calls[0]) == 2 # user + triage reply
|
||||
|
||||
requests = [ev for ev in events if isinstance(ev, RequestInfoEvent)]
|
||||
assert requests, "Workflow should request additional user input"
|
||||
request_payload = requests[-1].data
|
||||
assert isinstance(request_payload, HandoffUserInputRequest)
|
||||
assert len(request_payload.conversation) == 4 # user, triage tool call, tool ack, specialist
|
||||
assert request_payload.conversation[2].role == Role.TOOL
|
||||
assert request_payload.conversation[3].role == Role.ASSISTANT
|
||||
assert "specialist reply" in request_payload.conversation[3].text
|
||||
|
||||
follow_up = await _drain(workflow.send_responses_streaming({requests[-1].request_id: "Thanks"}))
|
||||
assert any(isinstance(ev, RequestInfoEvent) for ev in follow_up)
|
||||
|
||||
|
||||
async def test_specialist_to_specialist_handoff():
|
||||
"""Test that specialists can hand off to other specialists via .add_handoff() configuration."""
|
||||
triage = _RecordingAgent(name="triage", handoff_to="specialist")
|
||||
specialist = _RecordingAgent(name="specialist", handoff_to="escalation")
|
||||
escalation = _RecordingAgent(name="escalation")
|
||||
|
||||
workflow = (
|
||||
HandoffBuilder(participants=[triage, specialist, escalation])
|
||||
.set_coordinator(triage)
|
||||
.add_handoff(triage, [specialist, escalation])
|
||||
.add_handoff(specialist, escalation)
|
||||
.with_termination_condition(lambda conv: sum(1 for m in conv if m.role == Role.USER) >= 2)
|
||||
.build()
|
||||
)
|
||||
|
||||
# Start conversation - triage hands off to specialist
|
||||
events = await _drain(workflow.run_stream("Need technical support"))
|
||||
requests = [ev for ev in events if isinstance(ev, RequestInfoEvent)]
|
||||
assert requests
|
||||
|
||||
# Specialist should have been called
|
||||
assert len(specialist.calls) > 0
|
||||
|
||||
# Second user message - specialist hands off to escalation
|
||||
events = await _drain(workflow.send_responses_streaming({requests[-1].request_id: "This is complex"}))
|
||||
outputs = [ev for ev in events if isinstance(ev, WorkflowOutputEvent)]
|
||||
assert outputs
|
||||
|
||||
# Escalation should have been called
|
||||
assert len(escalation.calls) > 0
|
||||
|
||||
|
||||
async def test_handoff_preserves_complex_additional_properties(complex_metadata: _ComplexMetadata):
|
||||
triage = _RecordingAgent(name="triage", handoff_to="specialist", extra_properties={"complex": complex_metadata})
|
||||
specialist = _RecordingAgent(name="specialist")
|
||||
|
||||
# Sanity check: agent response contains complex metadata before entering workflow
|
||||
triage_response = await triage.run([ChatMessage(role=Role.USER, text="Need help with a return")])
|
||||
assert triage_response.messages
|
||||
assert "complex" in triage_response.messages[0].additional_properties
|
||||
|
||||
workflow = (
|
||||
HandoffBuilder(participants=[triage, specialist])
|
||||
.set_coordinator("triage")
|
||||
.with_termination_condition(lambda conv: sum(1 for msg in conv if msg.role == Role.USER) >= 2)
|
||||
.build()
|
||||
)
|
||||
|
||||
# Initial run should preserve complex metadata in the triage response
|
||||
events = await _drain(workflow.run_stream("Need help with a return"))
|
||||
agent_events = [ev for ev in events if hasattr(ev, "data") and hasattr(ev.data, "messages")]
|
||||
if agent_events:
|
||||
first_agent_event = agent_events[0]
|
||||
first_agent_event_data = first_agent_event.data
|
||||
if first_agent_event_data and hasattr(first_agent_event_data, "messages"):
|
||||
first_agent_message = first_agent_event_data.messages[0] # type: ignore[attr-defined]
|
||||
assert "complex" in first_agent_message.additional_properties, "Agent event lost complex metadata"
|
||||
requests = [ev for ev in events if isinstance(ev, RequestInfoEvent)]
|
||||
assert requests, "Workflow should request additional user input"
|
||||
|
||||
request_data = requests[-1].data
|
||||
assert isinstance(request_data, HandoffUserInputRequest)
|
||||
conversation_snapshot = request_data.conversation
|
||||
metadata_values = _metadata_from_conversation(conversation_snapshot, "complex")
|
||||
assert metadata_values, (
|
||||
"Expected triage message in conversation, found "
|
||||
f"additional_properties={[msg.additional_properties for msg in conversation_snapshot]},"
|
||||
f" messages={_conversation_debug(conversation_snapshot)}"
|
||||
)
|
||||
assert any(isinstance(value, _ComplexMetadata) for value in metadata_values), (
|
||||
"Complex metadata lost after first hop"
|
||||
)
|
||||
restored_meta = next(value for value in metadata_values if isinstance(value, _ComplexMetadata))
|
||||
assert restored_meta.payload["code"] == "X1"
|
||||
|
||||
# Respond and ensure metadata survives subsequent cycles
|
||||
follow_up_events = await _drain(
|
||||
workflow.send_responses_streaming({requests[-1].request_id: "Here are more details"})
|
||||
)
|
||||
follow_up_requests = [ev for ev in follow_up_events if isinstance(ev, RequestInfoEvent)]
|
||||
outputs = [ev for ev in follow_up_events if isinstance(ev, WorkflowOutputEvent)]
|
||||
|
||||
follow_up_conversation: list[ChatMessage]
|
||||
if follow_up_requests:
|
||||
follow_up_request_data = follow_up_requests[-1].data
|
||||
assert isinstance(follow_up_request_data, HandoffUserInputRequest)
|
||||
follow_up_conversation = follow_up_request_data.conversation
|
||||
else:
|
||||
assert outputs, "Workflow produced neither follow-up request nor output"
|
||||
output_data = outputs[-1].data
|
||||
follow_up_conversation = cast(list[ChatMessage], output_data) if isinstance(output_data, list) else []
|
||||
|
||||
metadata_values_after = _metadata_from_conversation(follow_up_conversation, "complex")
|
||||
assert metadata_values_after, "Expected triage message after follow-up"
|
||||
assert any(isinstance(value, _ComplexMetadata) for value in metadata_values_after), (
|
||||
"Complex metadata lost after restore"
|
||||
)
|
||||
|
||||
restored_meta_after = next(value for value in metadata_values_after if isinstance(value, _ComplexMetadata))
|
||||
assert restored_meta_after.payload["code"] == "X1"
|
||||
|
||||
|
||||
async def test_tool_call_handoff_detection_with_text_hint():
|
||||
triage = _RecordingAgent(name="triage", handoff_to="specialist", text_handoff=True)
|
||||
specialist = _RecordingAgent(name="specialist")
|
||||
|
||||
workflow = HandoffBuilder(participants=[triage, specialist]).set_coordinator("triage").build()
|
||||
|
||||
await _drain(workflow.run_stream("Package arrived broken"))
|
||||
|
||||
assert specialist.calls, "Specialist should be invoked using handoff tool call"
|
||||
assert len(specialist.calls[0]) >= 2
|
||||
|
||||
|
||||
def test_build_fails_without_coordinator():
|
||||
"""Verify that build() raises ValueError when set_coordinator() was not called."""
|
||||
triage = _RecordingAgent(name="triage")
|
||||
specialist = _RecordingAgent(name="specialist")
|
||||
|
||||
with pytest.raises(ValueError, match="coordinator must be defined before build"):
|
||||
HandoffBuilder(participants=[triage, specialist]).build()
|
||||
|
||||
|
||||
def test_build_fails_without_participants():
|
||||
"""Verify that build() raises ValueError when no participants are provided."""
|
||||
with pytest.raises(ValueError, match="No participants provided"):
|
||||
HandoffBuilder().build()
|
||||
|
||||
|
||||
async def test_multiple_runs_dont_leak_conversation():
|
||||
"""Verify that running the same workflow multiple times doesn't leak conversation history."""
|
||||
triage = _RecordingAgent(name="triage", handoff_to="specialist")
|
||||
specialist = _RecordingAgent(name="specialist")
|
||||
|
||||
workflow = (
|
||||
HandoffBuilder(participants=[triage, specialist])
|
||||
.set_coordinator("triage")
|
||||
.with_termination_condition(lambda conv: sum(1 for m in conv if m.role == Role.USER) >= 2)
|
||||
.build()
|
||||
)
|
||||
|
||||
# First run
|
||||
events = await _drain(workflow.run_stream("First run message"))
|
||||
requests = [ev for ev in events if isinstance(ev, RequestInfoEvent)]
|
||||
assert requests
|
||||
events = await _drain(workflow.send_responses_streaming({requests[-1].request_id: "Second message"}))
|
||||
outputs = [ev for ev in events if isinstance(ev, WorkflowOutputEvent)]
|
||||
assert outputs, "First run should emit output"
|
||||
|
||||
first_run_conversation = outputs[-1].data
|
||||
assert isinstance(first_run_conversation, list)
|
||||
first_run_conv_list = cast(list[ChatMessage], first_run_conversation)
|
||||
first_run_user_messages = [msg for msg in first_run_conv_list if msg.role == Role.USER]
|
||||
assert len(first_run_user_messages) == 2
|
||||
assert any("First run message" in msg.text for msg in first_run_user_messages if msg.text)
|
||||
|
||||
# Second run - should start fresh, not include first run's messages
|
||||
triage.calls.clear()
|
||||
specialist.calls.clear()
|
||||
|
||||
events = await _drain(workflow.run_stream("Second run different message"))
|
||||
requests = [ev for ev in events if isinstance(ev, RequestInfoEvent)]
|
||||
assert requests
|
||||
events = await _drain(workflow.send_responses_streaming({requests[-1].request_id: "Another message"}))
|
||||
outputs = [ev for ev in events if isinstance(ev, WorkflowOutputEvent)]
|
||||
assert outputs, "Second run should emit output"
|
||||
|
||||
second_run_conversation = outputs[-1].data
|
||||
assert isinstance(second_run_conversation, list)
|
||||
second_run_conv_list = cast(list[ChatMessage], second_run_conversation)
|
||||
second_run_user_messages = [msg for msg in second_run_conv_list if msg.role == Role.USER]
|
||||
assert len(second_run_user_messages) == 2, (
|
||||
"Second run should have exactly 2 user messages, not accumulate first run"
|
||||
)
|
||||
assert any("Second run different message" in msg.text for msg in second_run_user_messages if msg.text)
|
||||
assert not any("First run message" in msg.text for msg in second_run_user_messages if msg.text), (
|
||||
"Second run should NOT contain first run's messages"
|
||||
)
|
||||
@@ -11,8 +11,9 @@ from agent_framework import (
|
||||
AgentRunUpdateEvent,
|
||||
ChatMessage,
|
||||
Executor,
|
||||
FunctionApprovalRequestContent,
|
||||
FunctionApprovalResponseContent,
|
||||
FunctionCallContent,
|
||||
FunctionResultContent,
|
||||
RequestInfoExecutor,
|
||||
RequestInfoMessage,
|
||||
Role,
|
||||
@@ -163,35 +164,56 @@ class TestWorkflowAgent:
|
||||
updates: list[AgentRunResponseUpdate] = []
|
||||
async for update in agent.run_stream("Start request"):
|
||||
updates.append(update)
|
||||
# Should have received a function call for the request info
|
||||
# Should have received an approval request for the request info
|
||||
assert len(updates) > 0
|
||||
|
||||
# Find the function call update (RequestInfoEvent converted to function call)
|
||||
function_call_update: AgentRunResponseUpdate | None = None
|
||||
approval_update: AgentRunResponseUpdate | None = None
|
||||
for update in updates:
|
||||
if update.contents and hasattr(update.contents[0], "name") and update.contents[0].name == "request_info": # type: ignore[attr-defined]
|
||||
function_call_update = update
|
||||
if any(isinstance(content, FunctionApprovalRequestContent) for content in update.contents):
|
||||
approval_update = update
|
||||
break
|
||||
|
||||
assert function_call_update is not None, "Should have received a request_info function call"
|
||||
function_call: FunctionCallContent = function_call_update.contents[0] # type: ignore[assignment]
|
||||
assert approval_update is not None, "Should have received a request_info approval request"
|
||||
|
||||
function_call = next(
|
||||
content for content in approval_update.contents if isinstance(content, FunctionCallContent)
|
||||
)
|
||||
approval_request = next(
|
||||
content for content in approval_update.contents if isinstance(content, FunctionApprovalRequestContent)
|
||||
)
|
||||
|
||||
# Verify the function call has expected structure
|
||||
assert function_call.call_id is not None
|
||||
assert function_call.name == "request_info"
|
||||
assert isinstance(function_call.arguments, dict)
|
||||
assert "request_id" in function_call.arguments
|
||||
assert function_call.arguments.get("request_id") == approval_request.id
|
||||
|
||||
# Approval request should reference the same function call
|
||||
assert approval_request.function_call.call_id == function_call.call_id
|
||||
assert approval_request.function_call.name == function_call.name
|
||||
|
||||
# Verify the request is tracked in pending_requests
|
||||
assert len(agent.pending_requests) == 1
|
||||
assert function_call.call_id in agent.pending_requests
|
||||
|
||||
# Now provide a function result response to test continuation
|
||||
response_message = ChatMessage(
|
||||
role=Role.USER,
|
||||
contents=[FunctionResultContent(call_id=function_call.call_id, result="User provided answer")],
|
||||
# Now provide an approval response with updated arguments to test continuation
|
||||
response_args = WorkflowAgent.RequestInfoFunctionArgs(
|
||||
request_id=approval_request.id,
|
||||
data="User provided answer",
|
||||
).to_dict()
|
||||
|
||||
approval_response = FunctionApprovalResponseContent(
|
||||
approved=True,
|
||||
id=approval_request.id,
|
||||
function_call=FunctionCallContent(
|
||||
call_id=function_call.call_id,
|
||||
name=function_call.name,
|
||||
arguments=response_args,
|
||||
),
|
||||
)
|
||||
|
||||
response_message = ChatMessage(role=Role.USER, contents=[approval_response])
|
||||
|
||||
# Continue the workflow with the response
|
||||
continuation_result = await agent.run(response_message)
|
||||
|
||||
|
||||
@@ -62,16 +62,16 @@ If a policy violation is detected on the prompt, the middleware terminates the r
|
||||
`PurviewClient` uses the `azure-identity` library for token acquisition. You can use any `TokenCredential` or `AsyncTokenCredential` implementation.
|
||||
|
||||
The APIs require the following Graph Permissions:
|
||||
- ProtectionScopes.Compute.All : (userProtectionScopeContainer)[https://learn.microsoft.com/en-us/graph/api/userprotectionscopecontainer-compute]
|
||||
- Content.Process.All : (processContent)[https://learn.microsoft.com/en-us/graph/api/userdatasecurityandgovernance-processcontent]
|
||||
- ContentActivity.Write : (contentActivity)[https://learn.microsoft.com/en-us/graph/api/activitiescontainer-post-contentactivities]
|
||||
- ProtectionScopes.Compute.All : [userProtectionScopeContainer](https://learn.microsoft.com/en-us/graph/api/userprotectionscopecontainer-compute)
|
||||
- Content.Process.All : [processContent](https://learn.microsoft.com/en-us/graph/api/userdatasecurityandgovernance-processcontent)
|
||||
- ContentActivity.Write : [contentActivity](https://learn.microsoft.com/en-us/graph/api/activitiescontainer-post-contentactivities)
|
||||
|
||||
### Scopes
|
||||
`PurviewSettings.get_scopes()` derives the Graph scope list (currently `https://graph.microsoft.com/.default` style).
|
||||
|
||||
### Tenant Enablement for Purview
|
||||
- The tenant requires an e5 license and consumptive billing setup.
|
||||
- There need to be (Data Loss Prevention)[https://learn.microsoft.com/en-us/purview/dlp-create-deploy-policy] or (Data Collection Policies)[https://learn.microsoft.com/en-us/purview/collection-policies-policy-reference] that apply to the user to call Process Content API else it calls Content Activities API for auditing the message.
|
||||
- There need to be [Data Loss Prevention](https://learn.microsoft.com/en-us/purview/dlp-create-deploy-policy) or [Data Collection Policies](https://learn.microsoft.com/en-us/purview/collection-policies-policy-reference) that apply to the user to call Process Content API else it calls Content Activities API for auditing the message.
|
||||
|
||||
---
|
||||
|
||||
|
||||
Reference in New Issue
Block a user