Revert "Merge from main"

This reverts commit b8206a85d7.
This commit is contained in:
Dmytro Struk
2025-11-11 18:44:25 -08:00
Unverified
parent b8206a85d7
commit 85fcd230bf
231 changed files with 4138 additions and 19654 deletions
@@ -279,45 +279,6 @@ async def test_chat_client_streaming_observability(
assert span.attributes[OtelAttr.OUTPUT_MESSAGES] is not None
async def test_chat_client_without_model_id_observability(mock_chat_client, span_exporter: InMemorySpanExporter):
"""Test telemetry shouldn't fail when the model_id is not provided for unknown reason."""
client = use_observability(mock_chat_client)()
messages = [ChatMessage(role=Role.USER, text="Test")]
span_exporter.clear()
response = await client.get_response(messages=messages)
assert response is not None
spans = span_exporter.get_finished_spans()
assert len(spans) == 1
span = spans[0]
assert span.name == "chat unknown"
assert span.attributes[OtelAttr.OPERATION.value] == OtelAttr.CHAT_COMPLETION_OPERATION
assert span.attributes[SpanAttributes.LLM_REQUEST_MODEL] == "unknown"
async def test_chat_client_streaming_without_model_id_observability(
mock_chat_client, span_exporter: InMemorySpanExporter
):
"""Test streaming telemetry shouldn't fail when the model_id is not provided for unknown reason."""
client = use_observability(mock_chat_client)()
messages = [ChatMessage(role=Role.USER, text="Test")]
span_exporter.clear()
# Collect all yielded updates
updates = []
async for update in client.get_streaming_response(messages=messages):
updates.append(update)
# Verify we got the expected updates, this shouldn't be dependent on otel
assert len(updates) == 2
spans = span_exporter.get_finished_spans()
assert len(spans) == 1
span = spans[0]
assert span.name == "chat unknown"
assert span.attributes[OtelAttr.OPERATION.value] == OtelAttr.CHAT_COMPLETION_OPERATION
assert span.attributes[SpanAttributes.LLM_REQUEST_MODEL] == "unknown"
def test_prepend_user_agent_with_none_value():
"""Test prepend user agent with None value in headers."""
headers = {"User-Agent": None}
@@ -407,7 +368,6 @@ def mock_chat_agent():
self.name = "test_agent"
self.display_name = "Test Agent"
self.description = "Test agent description"
self.chat_options = ChatOptions(model_id="TestModel")
async def run(self, messages=None, *, thread=None, **kwargs):
return AgentRunResponse(
@@ -445,7 +405,7 @@ async def test_agent_instrumentation_enabled(
assert span.attributes[OtelAttr.AGENT_ID] == "test_agent_id"
assert span.attributes[OtelAttr.AGENT_NAME] == "Test Agent"
assert span.attributes[OtelAttr.AGENT_DESCRIPTION] == "Test agent description"
assert span.attributes[SpanAttributes.LLM_REQUEST_MODEL] == "TestModel"
assert span.attributes[SpanAttributes.LLM_REQUEST_MODEL] == "unknown"
assert span.attributes[OtelAttr.INPUT_TOKENS] == 15
assert span.attributes[OtelAttr.OUTPUT_TOKENS] == 25
if enable_sensitive_data:
@@ -473,7 +433,7 @@ async def test_agent_streaming_response_with_diagnostics_enabled_via_decorator(
assert span.attributes[OtelAttr.AGENT_ID] == "test_agent_id"
assert span.attributes[OtelAttr.AGENT_NAME] == "Test Agent"
assert span.attributes[OtelAttr.AGENT_DESCRIPTION] == "Test agent description"
assert span.attributes[SpanAttributes.LLM_REQUEST_MODEL] == "TestModel"
assert span.attributes[SpanAttributes.LLM_REQUEST_MODEL] == "unknown"
if enable_sensitive_data:
assert span.attributes.get(OtelAttr.OUTPUT_MESSAGES) is not None # Streaming, so no usage yet
@@ -1,6 +1,5 @@
# Copyright (c) Microsoft. All rights reserved.
import base64
from collections.abc import AsyncIterable
from typing import Any
@@ -167,57 +166,6 @@ def test_data_content_empty():
DataContent(uri="")
def test_data_content_detect_image_format_from_base64():
"""Test the detect_image_format_from_base64 static method."""
# Test each supported format
png_data = b"\x89PNG\r\n\x1a\n" + b"fake_data"
assert DataContent.detect_image_format_from_base64(base64.b64encode(png_data).decode()) == "png"
jpeg_data = b"\xff\xd8\xff\xe0" + b"fake_data"
assert DataContent.detect_image_format_from_base64(base64.b64encode(jpeg_data).decode()) == "jpeg"
webp_data = b"RIFF" + b"1234" + b"WEBP" + b"fake_data"
assert DataContent.detect_image_format_from_base64(base64.b64encode(webp_data).decode()) == "webp"
gif_data = b"GIF89a" + b"fake_data"
assert DataContent.detect_image_format_from_base64(base64.b64encode(gif_data).decode()) == "gif"
# Test fallback behavior
unknown_data = b"UNKNOWN_FORMAT"
assert DataContent.detect_image_format_from_base64(base64.b64encode(unknown_data).decode()) == "png"
# Test error handling
assert DataContent.detect_image_format_from_base64("invalid_base64!") == "png"
assert DataContent.detect_image_format_from_base64("") == "png"
def test_data_content_create_data_uri_from_base64():
"""Test the create_data_uri_from_base64 class method."""
# Test with PNG data
png_data = b"\x89PNG\r\n\x1a\n" + b"fake_data"
png_base64 = base64.b64encode(png_data).decode()
uri, media_type = DataContent.create_data_uri_from_base64(png_base64)
assert uri == f"data:image/png;base64,{png_base64}"
assert media_type == "image/png"
# Test with different format
jpeg_data = b"\xff\xd8\xff\xe0" + b"fake_data"
jpeg_base64 = base64.b64encode(jpeg_data).decode()
uri, media_type = DataContent.create_data_uri_from_base64(jpeg_base64)
assert uri == f"data:image/jpeg;base64,{jpeg_base64}"
assert media_type == "image/jpeg"
# Test fallback for unknown format
unknown_data = b"UNKNOWN_FORMAT"
unknown_base64 = base64.b64encode(unknown_data).decode()
uri, media_type = DataContent.create_data_uri_from_base64(unknown_base64)
assert uri == f"data:image/png;base64,{unknown_base64}"
assert media_type == "image/png"
# region UriContent
File diff suppressed because it is too large Load Diff
@@ -111,10 +111,6 @@ async def test_agent_executor_checkpoint_stores_and_restores_state() -> None:
chat_store_state = thread_state["chat_message_store_state"] # type: ignore[index]
assert "messages" in chat_store_state, "Message store state should include messages"
# Verify checkpoint contains pending requests from agents and responses to be sent
assert "pending_agent_requests" in executor_state
assert "pending_responses_to_agent" in executor_state
# Create a new agent and executor for restoration
# This simulates starting from a fresh state and restoring from checkpoint
restored_agent = _CountingAgent(id="test_agent", name="TestAgent")
@@ -5,32 +5,19 @@
from collections.abc import AsyncIterable
from typing import Any
from typing_extensions import Never
from agent_framework import (
AgentExecutor,
AgentExecutorResponse,
AgentRunResponse,
AgentRunResponseUpdate,
AgentRunUpdateEvent,
AgentThread,
BaseAgent,
ChatAgent,
ChatMessage,
ChatResponse,
ChatResponseUpdate,
FunctionApprovalRequestContent,
FunctionCallContent,
FunctionResultContent,
RequestInfoEvent,
Role,
TextContent,
WorkflowBuilder,
WorkflowContext,
WorkflowOutputEvent,
ai_function,
executor,
use_function_invocation,
)
@@ -133,235 +120,3 @@ async def test_agent_executor_emits_tool_calls_in_streaming_mode() -> None:
assert events[3].data is not None
assert isinstance(events[3].data.contents[0], TextContent)
assert "sunny" in events[3].data.contents[0].text
@ai_function(approval_mode="always_require")
def mock_tool_requiring_approval(query: str) -> str:
"""Mock tool that requires approval before execution."""
return f"Executed tool with query: {query}"
@use_function_invocation
class MockChatClient:
"""Simple implementation of a chat client."""
def __init__(self, parallel_request: bool = False) -> None:
self.additional_properties: dict[str, Any] = {}
self._iteration: int = 0
self._parallel_request: bool = parallel_request
async def get_response(
self,
messages: str | ChatMessage | list[str] | list[ChatMessage],
**kwargs: Any,
) -> ChatResponse:
if self._iteration == 0:
if self._parallel_request:
response = ChatResponse(
messages=ChatMessage(
role="assistant",
contents=[
FunctionCallContent(
call_id="1", name="mock_tool_requiring_approval", arguments='{"query": "test"}'
),
FunctionCallContent(
call_id="2", name="mock_tool_requiring_approval", arguments='{"query": "test"}'
),
],
)
)
else:
response = ChatResponse(
messages=ChatMessage(
role="assistant",
contents=[
FunctionCallContent(
call_id="1", name="mock_tool_requiring_approval", arguments='{"query": "test"}'
)
],
)
)
else:
response = ChatResponse(messages=ChatMessage(role="assistant", text="Tool executed successfully."))
self._iteration += 1
return response
async def get_streaming_response(
self,
messages: str | ChatMessage | list[str] | list[ChatMessage],
**kwargs: Any,
) -> AsyncIterable[ChatResponseUpdate]:
if self._iteration == 0:
if self._parallel_request:
yield ChatResponseUpdate(
contents=[
FunctionCallContent(
call_id="1", name="mock_tool_requiring_approval", arguments='{"query": "test"}'
),
FunctionCallContent(
call_id="2", name="mock_tool_requiring_approval", arguments='{"query": "test"}'
),
],
role="assistant",
)
else:
yield ChatResponseUpdate(
contents=[
FunctionCallContent(
call_id="1", name="mock_tool_requiring_approval", arguments='{"query": "test"}'
)
],
role="assistant",
)
else:
yield ChatResponseUpdate(text=TextContent(text="Tool executed "), role="assistant")
yield ChatResponseUpdate(contents=[TextContent(text="successfully.")], role="assistant")
self._iteration += 1
@executor(id="test_executor")
async def test_executor(agent_executor_response: AgentExecutorResponse, ctx: WorkflowContext[Never, str]) -> None:
await ctx.yield_output(agent_executor_response.agent_run_response.text)
async def test_agent_executor_tool_call_with_approval() -> None:
"""Test that AgentExecutor handles tool calls requiring approval."""
# Arrange
agent = ChatAgent(
chat_client=MockChatClient(),
name="ApprovalAgent",
tools=[mock_tool_requiring_approval],
)
workflow = WorkflowBuilder().set_start_executor(agent).add_edge(agent, test_executor).build()
# Act
events = await workflow.run("Invoke tool requiring approval")
# Assert
assert len(events.get_request_info_events()) == 1
approval_request = events.get_request_info_events()[0]
assert isinstance(approval_request.data, FunctionApprovalRequestContent)
assert approval_request.data.function_call.name == "mock_tool_requiring_approval"
assert approval_request.data.function_call.arguments == '{"query": "test"}'
# Act
events = await workflow.send_responses({approval_request.request_id: approval_request.data.create_response(True)})
# Assert
final_response = events.get_outputs()
assert len(final_response) == 1
assert final_response[0] == "Tool executed successfully."
async def test_agent_executor_tool_call_with_approval_streaming() -> None:
"""Test that AgentExecutor handles tool calls requiring approval in streaming mode."""
# Arrange
agent = ChatAgent(
chat_client=MockChatClient(),
name="ApprovalAgent",
tools=[mock_tool_requiring_approval],
)
workflow = WorkflowBuilder().set_start_executor(agent).add_edge(agent, test_executor).build()
# Act
request_info_events: list[RequestInfoEvent] = []
async for event in workflow.run_stream("Invoke tool requiring approval"):
if isinstance(event, RequestInfoEvent):
request_info_events.append(event)
# Assert
assert len(request_info_events) == 1
approval_request = request_info_events[0]
assert isinstance(approval_request.data, FunctionApprovalRequestContent)
assert approval_request.data.function_call.name == "mock_tool_requiring_approval"
assert approval_request.data.function_call.arguments == '{"query": "test"}'
# Act
output: str | None = None
async for event in workflow.send_responses_streaming({
approval_request.request_id: approval_request.data.create_response(True)
}):
if isinstance(event, WorkflowOutputEvent):
output = event.data
# Assert
assert output is not None
assert output == "Tool executed successfully."
async def test_agent_executor_parallel_tool_call_with_approval() -> None:
"""Test that AgentExecutor handles parallel tool calls requiring approval."""
# Arrange
agent = ChatAgent(
chat_client=MockChatClient(parallel_request=True),
name="ApprovalAgent",
tools=[mock_tool_requiring_approval],
)
workflow = WorkflowBuilder().set_start_executor(agent).add_edge(agent, test_executor).build()
# Act
events = await workflow.run("Invoke tool requiring approval")
# Assert
assert len(events.get_request_info_events()) == 2
for approval_request in events.get_request_info_events():
assert isinstance(approval_request.data, FunctionApprovalRequestContent)
assert approval_request.data.function_call.name == "mock_tool_requiring_approval"
assert approval_request.data.function_call.arguments == '{"query": "test"}'
# Act
responses = {
approval_request.request_id: approval_request.data.create_response(True) # type: ignore
for approval_request in events.get_request_info_events()
}
events = await workflow.send_responses(responses)
# Assert
final_response = events.get_outputs()
assert len(final_response) == 1
assert final_response[0] == "Tool executed successfully."
async def test_agent_executor_parallel_tool_call_with_approval_streaming() -> None:
"""Test that AgentExecutor handles parallel tool calls requiring approval in streaming mode."""
# Arrange
agent = ChatAgent(
chat_client=MockChatClient(parallel_request=True),
name="ApprovalAgent",
tools=[mock_tool_requiring_approval],
)
workflow = WorkflowBuilder().set_start_executor(agent).add_edge(agent, test_executor).build()
# Act
request_info_events: list[RequestInfoEvent] = []
async for event in workflow.run_stream("Invoke tool requiring approval"):
if isinstance(event, RequestInfoEvent):
request_info_events.append(event)
# Assert
assert len(request_info_events) == 2
for approval_request in request_info_events:
assert isinstance(approval_request.data, FunctionApprovalRequestContent)
assert approval_request.data.function_call.name == "mock_tool_requiring_approval"
assert approval_request.data.function_call.arguments == '{"query": "test"}'
# Act
responses = {
approval_request.request_id: approval_request.data.create_response(True) # type: ignore
for approval_request in request_info_events
}
output: str | None = None
async for event in workflow.send_responses_streaming(responses):
if isinstance(event, WorkflowOutputEvent):
output = event.data
# Assert
assert output is not None
assert output == "Tool executed successfully."
@@ -23,7 +23,7 @@ from agent_framework import (
WorkflowOutputEvent,
)
from agent_framework._mcp import MCPTool
from agent_framework._workflows._handoff import _clone_chat_agent # type: ignore[reportPrivateUsage]
from agent_framework._workflows._handoff import _clone_chat_agent
@dataclass
@@ -392,218 +392,12 @@ async def test_clone_chat_agent_preserves_mcp_tools() -> None:
)
assert hasattr(original_agent, "_local_mcp_tools")
assert len(original_agent._local_mcp_tools) == 1 # type: ignore[reportPrivateUsage]
assert original_agent._local_mcp_tools[0] == mock_mcp_tool # type: ignore[reportPrivateUsage]
assert len(original_agent._local_mcp_tools) == 1
assert original_agent._local_mcp_tools[0] == mock_mcp_tool
cloned_agent = _clone_chat_agent(original_agent)
assert hasattr(cloned_agent, "_local_mcp_tools")
assert len(cloned_agent._local_mcp_tools) == 1 # type: ignore[reportPrivateUsage]
assert cloned_agent._local_mcp_tools[0] == mock_mcp_tool # type: ignore[reportPrivateUsage]
assert cloned_agent.chat_options.tools is not None
assert len(cloned_agent._local_mcp_tools) == 1
assert cloned_agent._local_mcp_tools[0] == mock_mcp_tool
assert len(cloned_agent.chat_options.tools) == 1
async def test_return_to_previous_routing():
"""Test that return-to-previous routes back to the current specialist handling the conversation."""
triage = _RecordingAgent(name="triage", handoff_to="specialist_a")
specialist_a = _RecordingAgent(name="specialist_a", handoff_to="specialist_b")
specialist_b = _RecordingAgent(name="specialist_b")
workflow = (
HandoffBuilder(participants=[triage, specialist_a, specialist_b])
.set_coordinator(triage)
.add_handoff(triage, [specialist_a, specialist_b])
.add_handoff(specialist_a, specialist_b)
.enable_return_to_previous(True)
.with_termination_condition(lambda conv: sum(1 for m in conv if m.role == Role.USER) >= 4)
.build()
)
# Start conversation - triage hands off to specialist_a
events = await _drain(workflow.run_stream("Initial request"))
requests = [ev for ev in events if isinstance(ev, RequestInfoEvent)]
assert requests
assert len(specialist_a.calls) > 0
# Specialist_a should have been called with initial request
initial_specialist_a_calls = len(specialist_a.calls)
# Second user message - specialist_a hands off to specialist_b
events = await _drain(workflow.send_responses_streaming({requests[-1].request_id: "Need more help"}))
requests = [ev for ev in events if isinstance(ev, RequestInfoEvent)]
assert requests
# Specialist_b should have been called
assert len(specialist_b.calls) > 0
initial_specialist_b_calls = len(specialist_b.calls)
# Third user message - with return_to_previous, should route back to specialist_b (current agent)
events = await _drain(workflow.send_responses_streaming({requests[-1].request_id: "Follow up question"}))
third_requests = [ev for ev in events if isinstance(ev, RequestInfoEvent)]
# Specialist_b should have been called again (return-to-previous routes to current agent)
assert len(specialist_b.calls) > initial_specialist_b_calls, (
"Specialist B should be called again due to return-to-previous routing to current agent"
)
# Specialist_a should NOT be called again (it's no longer the current agent)
assert len(specialist_a.calls) == initial_specialist_a_calls, (
"Specialist A should not be called again - specialist_b is the current agent"
)
# Triage should only have been called once at the start
assert len(triage.calls) == 1, "Triage should only be called once (initial routing)"
# Verify awaiting_agent_id is set to specialist_b (the agent that just responded)
if third_requests:
user_input_req = third_requests[-1].data
assert isinstance(user_input_req, HandoffUserInputRequest)
assert user_input_req.awaiting_agent_id == "specialist_b", (
f"Expected awaiting_agent_id 'specialist_b' but got '{user_input_req.awaiting_agent_id}'"
)
async def test_return_to_previous_disabled_routes_to_coordinator():
"""Test that with return-to-previous disabled, routing goes back to coordinator."""
triage = _RecordingAgent(name="triage", handoff_to="specialist_a")
specialist_a = _RecordingAgent(name="specialist_a", handoff_to="specialist_b")
specialist_b = _RecordingAgent(name="specialist_b")
workflow = (
HandoffBuilder(participants=[triage, specialist_a, specialist_b])
.set_coordinator(triage)
.add_handoff(triage, [specialist_a, specialist_b])
.add_handoff(specialist_a, specialist_b)
.enable_return_to_previous(False)
.with_termination_condition(lambda conv: sum(1 for m in conv if m.role == Role.USER) >= 3)
.build()
)
# Start conversation - triage hands off to specialist_a
events = await _drain(workflow.run_stream("Initial request"))
requests = [ev for ev in events if isinstance(ev, RequestInfoEvent)]
assert requests
assert len(triage.calls) == 1
# Second user message - specialist_a hands off to specialist_b
events = await _drain(workflow.send_responses_streaming({requests[-1].request_id: "Need more help"}))
requests = [ev for ev in events if isinstance(ev, RequestInfoEvent)]
assert requests
# Third user message - without return_to_previous, should route back to triage
await _drain(workflow.send_responses_streaming({requests[-1].request_id: "Follow up question"}))
# Triage should have been called twice total: initial + after specialist_b responds
assert len(triage.calls) == 2, "Triage should be called twice (initial + default routing to coordinator)"
async def test_return_to_previous_enabled():
"""Verify that enable_return_to_previous() keeps control with the current specialist."""
triage = _RecordingAgent(name="triage", handoff_to="specialist_a")
specialist_a = _RecordingAgent(name="specialist_a")
specialist_b = _RecordingAgent(name="specialist_b")
workflow = (
HandoffBuilder(participants=[triage, specialist_a, specialist_b])
.set_coordinator("triage")
.enable_return_to_previous(True)
.with_termination_condition(lambda conv: sum(1 for m in conv if m.role == Role.USER) >= 3)
.build()
)
# Start conversation - triage hands off to specialist_a
events = await _drain(workflow.run_stream("Initial request"))
requests = [ev for ev in events if isinstance(ev, RequestInfoEvent)]
assert requests
assert len(triage.calls) == 1
assert len(specialist_a.calls) == 1
# Second user message - with return_to_previous, should route to specialist_a (not triage)
events = await _drain(workflow.send_responses_streaming({requests[-1].request_id: "Follow up question"}))
requests = [ev for ev in events if isinstance(ev, RequestInfoEvent)]
assert requests
# Triage should only have been called once (initial) - specialist_a handles follow-up
assert len(triage.calls) == 1, "Triage should only be called once (initial)"
assert len(specialist_a.calls) == 2, "Specialist A should handle follow-up with return_to_previous enabled"
async def test_tool_choice_preserved_from_agent_config():
"""Verify that agent-level tool_choice configuration is preserved and not overridden."""
from unittest.mock import AsyncMock
from agent_framework import ChatResponse, ToolMode
# Create a mock chat client that records the tool_choice used
recorded_tool_choices: list[Any] = []
async def mock_get_response(messages: Any, **kwargs: Any) -> ChatResponse:
chat_options = kwargs.get("chat_options")
if chat_options:
recorded_tool_choices.append(chat_options.tool_choice)
return ChatResponse(
messages=[ChatMessage(role=Role.ASSISTANT, text="Response")],
response_id="test_response",
)
mock_client = MagicMock()
mock_client.get_response = AsyncMock(side_effect=mock_get_response)
# Create agent with specific tool_choice configuration
agent = ChatAgent(
chat_client=mock_client,
name="test_agent",
tool_choice=ToolMode(mode="required"), # type: ignore[arg-type]
)
# Run the agent
await agent.run("Test message")
# Verify tool_choice was preserved
assert len(recorded_tool_choices) > 0, "No tool_choice recorded"
last_tool_choice = recorded_tool_choices[-1]
assert last_tool_choice is not None, "tool_choice should not be None"
assert str(last_tool_choice) == "required", f"Expected 'required', got {last_tool_choice}"
async def test_return_to_previous_state_serialization():
"""Test that return_to_previous state is properly serialized/deserialized for checkpointing."""
from agent_framework._workflows._handoff import _HandoffCoordinator # type: ignore[reportPrivateUsage]
# Create a coordinator with return_to_previous enabled
coordinator = _HandoffCoordinator(
starting_agent_id="triage",
specialist_ids={"specialist_a": "specialist_a", "specialist_b": "specialist_b"},
input_gateway_id="gateway",
termination_condition=lambda conv: False,
id="test-coordinator",
return_to_previous=True,
)
# Set the current agent (simulating a handoff scenario)
coordinator._current_agent_id = "specialist_a" # type: ignore[reportPrivateUsage]
# Snapshot the state
state = coordinator.snapshot_state()
# Verify pattern metadata includes current_agent_id
assert "metadata" in state
assert "current_agent_id" in state["metadata"]
assert state["metadata"]["current_agent_id"] == "specialist_a"
# Create a new coordinator and restore state
coordinator2 = _HandoffCoordinator(
starting_agent_id="triage",
specialist_ids={"specialist_a": "specialist_a", "specialist_b": "specialist_b"},
input_gateway_id="gateway",
termination_condition=lambda conv: False,
id="test-coordinator",
return_to_previous=True,
)
# Restore state
coordinator2.restore_state(state)
# Verify current_agent_id was restored
assert coordinator2._current_agent_id == "specialist_a", "Current agent should be restored from checkpoint" # type: ignore[reportPrivateUsage]