mirror of
https://github.com/microsoft/agent-framework.git
synced 2026-06-16 21:04:09 +08:00
Python: [BREAKING] Moved to a single get_response and run API (#3379)
* WIP * big update to new ResponseStream model * fixed tests and typing * fixed tests and typing * fixed tools typevar import * fix * mypy fix * mypy fixes and some cleanup * fix missing quoted names * and client * fix imports agui * fix anthropic override * fix agui * fix ag ui * fix import * fix anthropic types * fix mypy * refactoring * updated typing * fix 3.11 * fixes * redid layering of chat clients and agents * redid layering of chat clients and agents * Fix lint, type, and test issues after rebase - Add @overload decorators to AgentProtocol.run() for type compatibility - Add missing docstring params (middleware, function_invocation_configuration) - Fix TODO format (TD002) by adding author tags - Fix broken observability tests from upstream: - Replace non-existent use_instrumentation with direct instantiation - Replace non-existent use_agent_instrumentation with AgentTelemetryLayer mixin - Fix get_streaming_response to use get_response(stream=True) - Add AgentInitializationError import - Update streaming exception tests to match actual behavior * Fix AgentExecutionException import error in test_agents.py - Replace non-existent AgentExecutionException with AgentRunException * Fix test import and asyncio deprecation issues - Add 'tests' to pythonpath in ag-ui pyproject.toml for utils_test_ag_ui import - Replace deprecated asyncio.get_event_loop().run_until_complete with asyncio.run * Fix azure-ai test failures - Update _prepare_options patching to use correct class path - Fix test_to_azure_ai_agent_tools_web_search_missing_connection to clear env vars * Convert ag-ui utils_test_ag_ui.py to conftest.py - Move test utilities to conftest.py for proper pytest discovery - Update all test imports to use conftest instead of utils_test_ag_ui - Remove old utils_test_ag_ui.py file - Revert pythonpath change in pyproject.toml * fix: use relative imports for ag-ui test utilities * fix agui * Rename Bare*Client to Raw*Client and BaseChatClient - Renamed BareChatClient to BaseChatClient (abstract base class) - Renamed BareOpenAIChatClient to RawOpenAIChatClient - Renamed BareOpenAIResponsesClient to RawOpenAIResponsesClient - Renamed BareAzureAIClient to RawAzureAIClient - Added warning docstrings to Raw* classes about layer ordering - Updated README in samples/getting_started/agents/custom with layer docs - Added test for span ordering with function calling * Fix layer ordering: FunctionInvocationLayer before ChatTelemetryLayer This ensures each inner LLM call gets its own telemetry span, resulting in the correct span sequence: chat -> execute_tool -> chat Updated all production clients and test mocks to use correct ordering: - ChatMiddlewareLayer (first) - FunctionInvocationLayer (second) - ChatTelemetryLayer (third) - BaseChatClient/Raw...Client (fourth) * Remove run_stream usage * Fix conversation_id propagation * Python: Add BaseAgent implementation for Claude Agent SDK (#3509) * Added ClaudeAgent implementation * Updated streaming logic * Small updates * Small update * Fixes * Small fix * Naming improvements * Updated imports * Addressed comments * Updated package versions * Update Claude agent connector layering * fix test and plugin * Store function middleware in invocation layer * Fix telemetry streaming and ag-ui tests * Remove legacy ag-ui tests folder * updates * Remove terminate flag from FunctionInvocationContext, use MiddlewareTermination instead - Remove terminate attribute from FunctionInvocationContext - Add result attribute to MiddlewareTermination to carry function results - FunctionMiddlewarePipeline.execute() now lets MiddlewareTermination propagate - _auto_invoke_function captures context.result in exception before re-raising - _try_execute_function_calls catches MiddlewareTermination and sets should_terminate - Fix handoff middleware to append to chat_client.function_middleware directly - Update tests to use raise MiddlewareTermination instead of context.terminate - Add middleware flow documentation in samples/concepts/tools/README.md - Fix ag-ui to use FunctionMiddlewarePipeline instead of removed create_function_middleware_pipeline * fix: remove references to removed terminate flag in purview tests, add type ignore * fix: move _test_utils.py from package to test folder * fix: call get_final_response() to trigger context provider notification in streaming test * fix: correct broken links in tools README * docs: clarify default middleware behavior in summary table * fix: ensure inner stream result hooks are called when using map()/from_awaitable() * Fix mypy type errors * Address PR review comments on observability.py - Remove TODO comment about unconsumed streams, add explanatory note instead - Remove redundant _close_span cleanup hook (already called in _finalize_stream) - Clarify behavior: cleanup hooks run after stream iteration, if stream is not consumed the span remains open until garbage collected * Remove gen_ai.client.operation.duration from span attributes Duration is a metrics-only attribute per OpenTelemetry semantic conventions. It should be recorded to the histogram but not set as a span attribute. * Remove duration from _get_response_attributes, pass directly to _capture_response Duration is a metrics-only attribute. It's now passed directly to _capture_response instead of being included in the attributes dict that gets set on the span. * Remove redundant _close_span cleanup hook in AgentTelemetryLayer _finalize_stream already calls _close_span() in its finally block, so adding it as a separate cleanup hook is redundant. * Use weakref.finalize to close span when stream is garbage collected If a user creates a streaming response but never consumes it, the cleanup hooks won't run. Now we register a weak reference finalizer that will close the span when the stream object is garbage collected, ensuring spans don't leak in this scenario. * Fix _get_finalizers_from_stream to use _result_hooks attribute Renamed function to _get_result_hooks_from_stream and fixed it to look for the _result_hooks attribute which is the correct name in ResponseStream class. * Add missing asyncio import in test_request_info_mixin.py * Fix leftover merge conflict marker in image_generation sample * Update integration tests * Fix integration tests: increase max_iterations from 1 to 2 Tests with tool_choice options require at least 2 iterations: 1. First iteration to get function call and execute the tool 2. Second iteration to get the final text response With max_iterations=1, streaming tests would return early with only the function call/result but no final text content. * Fix duplicate function call error in conversation-based APIs When using conversation_id (for Responses/Assistants APIs), the server already has the function call message from the previous response. We should only send the new function result message, not all messages including the function call which would cause a duplicate ID error. Fix: When conversation_id is set, only send the last message (the tool result) instead of all response.messages. * Add regression test for conversation_id propagation between tool iterations Port test from PR #3664 with updates for new streaming API pattern. Tests that conversation_id is properly updated in options dict during function invocation loop iterations. * Fix tool_choice=required to return after tool execution When tool_choice is 'required', the user's intent is to force exactly one tool call. After the tool executes, return immediately with the function call and result - don't continue to call the model again. This fixes integration tests that were failing with empty text responses because with tool_choice=required, the model would keep returning function calls instead of text. Also adds regression tests for: - conversation_id propagation between tool iterations (from PR #3664) - tool_choice=required returns after tool execution * Document tool_choice behavior in tools README - Add table explaining tool_choice values (auto, none, required) - Explain why tool_choice=required returns immediately after tool execution - Add code example showing the difference between required and auto - Update flow diagram to show the early return path for tool_choice=required * Fix tool_choice=None behavior - don't default to 'auto' Remove the hardcoded default of 'auto' for tool_choice in ChatAgent init. When tool_choice is not specified (None), it will now not be sent to the API, allowing the API's default behavior to be used. Users who want tool_choice='auto' can still explicitly set it either in default_options or at runtime. Fixes #3585 * Fix tool_choice=none should not remove tools In OpenAI Assistants client, tools were not being sent when tool_choice='none'. This was incorrect - tool_choice='none' means the model won't call tools, but tools should still be available in the request (they may be used later in the conversation). Fixes #3585 * Add test for tool_choice=none preserving tools Adds a regression test to ensure that when tool_choice='none' is set but tools are provided, the tools are still sent to the API. This verifies the fix for #3585. * Fix tool_choice=none should not remove tools in all clients Apply the same fix to OpenAI Responses client and Azure AI client: - OpenAI Responses: Remove else block that popped tool_choice/parallel_tool_calls - Azure AI: Remove tool_choice != 'none' check when adding tools When tool_choice='none', the model won't call tools, but tools should still be sent to the API so they're available for future turns. Also update README to clarify tool_choice=required supports multiple tools. Fixes #3585 * Keep tool_choice even when tools is None Move tool_choice processing outside of the 'if tools' block in OpenAI Responses client so tool_choice is sent to the API even when no tools are provided. * Update test to match new parallel_tool_calls behavior Changed test_prepare_options_removes_parallel_tool_calls_when_no_tools to test_prepare_options_preserves_parallel_tool_calls_when_no_tools to reflect that parallel_tool_calls is now preserved even when no tools are present, consistent with the tool_choice behavior. * Fix ChatMessage API and Role enum usage after rebase - Update ChatMessage instantiation to use keyword args (role=, text=, contents=) - Fix Role enum comparisons to use .value for string comparison - Add created_at to AgentResponse in error handling - Fix AgentResponse.from_updates -> from_agent_run_response_updates - Fix DurableAgentStateMessage.from_chat_message to convert Role enum to string - Add Role import where needed * Fix additional ChatMessage API and method name changes - Fix ChatMessage usage in workflow files (use text= instead of contents= for strings) - Fix AgentResponse.from_updates -> from_agent_run_response_updates in workflow files - Fix test files for ChatMessage and Role enum usage * Fix remaining ChatMessage API usage in test files * Fix more ChatMessage and Role API changes in source and test files - Fix ChatMessage in _magentic.py replan method - Fix Role enum comparison in test assertions - Fix remaining test files with old ChatMessage syntax * Fix ChatMessage and Role API changes across packages - Add Role import where missing - Fix ChatMessage signature: positional args to keyword args (role=, text=, contents=) - Fix Role enum comparisons: .role.value instead of .role string - Fix FinishReason enum usage in ag-ui event converters - Rename AgentResponse.from_updates to from_agent_run_response_updates in ag-ui Fixes API compatibility after Types API Review improvements merge * Fix ChatMessage and Role API changes in github_copilot tests * Fix ChatMessage and Role API changes in redis and github_copilot packages - Fix redis provider: Role enum comparison using .value - Fix redis tests: ChatMessage signature and Role comparisons - Fix github_copilot tests: ChatMessage signature and Role comparisons - Update docstring examples in redis chat message store * Fix ChatMessage and Role API changes in devui package - Fix executor: ChatMessage signature change - Fix conversations: Role enum to string conversion in two places - Fix tests: ChatMessage signatures and Role comparisons * Fix ChatMessage and Role API changes in a2a and lab packages - Fix a2a tests: Role comparisons and ChatMessage signatures - Fix lab tau2 source: Role enum comparison in flip_messages, log_messages, sliding_window - Fix lab tau2 tests: ChatMessage signatures and Role comparisons * Remove duplicate test files from ag-ui/tests (tests are in ag_ui_tests) * Fix ChatMessage and Role API changes across packages After rebasing on upstream/main which merged PR #3647 (Types API Review improvements), fix all packages to use the new API: - ChatMessage: Use keyword args (role=, text=, contents=) instead of positional args - Role: Compare using .value attribute since it's now an enum Packages fixed: - ag-ui: Fixed Role value extraction bugs in _message_adapters.py - anthropic: Fixed ChatMessage and Role comparisons in tests - azure-ai: Fixed Role comparison in _client.py - azure-ai-search: Fixed ChatMessage and Role in source/tests - bedrock: Fixed ChatMessage signatures in tests - chatkit: Fixed ChatMessage and Role in source/tests - copilotstudio: Fixed ChatMessage and Role in tests - declarative: Fixed ChatMessage in _executors_agents.py - mem0: Fixed ChatMessage and Role in source/tests - purview: Fixed ChatMessage in source/tests * Fix mypy errors for ChatMessage and Role API changes - durabletask: Use str() fallback in role value extraction - core: Fix ChatMessage in _orchestrator_helpers.py to use keyword args - core: Add type ignore for _conversation_state.py contents deserialization - ag-ui: Fix type ignore comments (call-overload instead of arg-type) - azure-ai-search: Fix get_role_value type hint to accept Any - lab: Move get_role_value to module level with Any type hint * Improve CI test timeout configuration - Increase job timeout from 10 to 15 minutes - Reduce per-test timeout to 60s (was 900s/300s) - Add --timeout_method thread for better timeout handling - Add --timeout-verbose to see which tests are slow - Reduce retries from 3 to 2 and delay from 10s to 5s This ensures individual test timeouts are shorter than the job timeout, providing better visibility when tests hang. With 60s timeout and 2 retries, worst case per test is ~180s. * Fix ChatMessage API usage in docstrings and source - Fix ChatMessage positional args in docstrings: _serialization.py, _threads.py, _middleware.py - Fix ChatMessage in tau2 runner.py - Fix role comparison in _orchestrator_helpers.py to use .value - Fix role comparison in _group_chat.py docstring example - Fix role assertions in test_durable_entities.py to use .value * Revert tool_choice/parallel_tool_calls changes - must be removed when no tools OpenAI API requires tool_choice and parallel_tool_calls to only be present when tools are specified. Restored the logic that removes these options when there are no tools. - Restored check in _chat_client.py to remove tool_choice and parallel_tool_calls when no tools present - Restored same logic in _responses_client.py - Reverted test to expect the correct behavior * fixed issue in tests * fix: resolve merge conflict markers in ag-ui tests * fix: restructure ag-ui tests and fix Role/FinishReason to use string types * fix: streaming function invocation and middleware termination - Refactor streaming function invocation to use get_final_response() on inner streams - Fix MiddlewareTermination to accept result parameter for passing results - Fix _AutoHandoffMiddleware to use MiddlewareTermination instead of context.terminate - Fix AgentMiddlewareLayer.run() to properly forward function/chat middleware - Remove duplicate middleware registration in AgentMiddlewareLayer.__init__ - Fix exception handling in _auto_invoke_function to properly capture termination - Fix mypy errors in core package - Update tests to use stream=True parameter for unified run API * fix all tests command * Refactor integration tests to use pytest fixtures - Merge testutils.py into conftest.py for azurefunctions integration tests - Merge dt_testutils.py into conftest.py for durabletask integration tests - Convert all integration tests to use fixtures instead of direct imports (fixes ModuleNotFoundError with --import-mode=importlib) - Add sample_helper fixture for azurefunctions tests - Add agent_client_factory and orchestration_helper fixtures for durabletask - Integration tests now skip with descriptive messages when services unavailable - Restructure devui tests into tests/devui/ with proper conftest.py - Add test organization guidelines to CODING_STANDARD.md - Remove __init__.py from test directories per pytest best practices * Fix pytest_collection_modifyitems to only skip integration tests The hook was skipping all tests in the test session, not just integration tests. Now it only skips items in the integration_tests directory. * Fix mem0 tests failing on Python 3.13 Use patch.object on the imported module instead of @patch with string path to ensure the mock takes effect regardless of import timing. * fix mem0 * another attempt for mem0 * fix for mem0 * fix mem0 * Increase worker initialization wait time in durabletask tests Increase from 2 to 8 seconds to allow time for: - Python startup and module imports - Azure OpenAI client creation - Agent registration with DTS worker - Worker connection to DTS This helps prevent test failures in CI where the first tests may run before the worker is fully ready to process requests. * Fix streaming test to use ResponseStream with finalizer The _consume_stream method now expects a ResponseStream that can provide a final AgentResponse via get_final_response(). Update the test to use ResponseStream with AgentResponse.from_updates as the finalizer. * Fix MockToolCallingAgent to use new ResponseStream API and update samples * small updates to run_stream to run * fix sub workflow * temp fix for az func test --------- Co-authored-by: Dmytro Struk <13853051+dmytrostruk@users.noreply.github.com>
This commit is contained in:
co-authored by
Dmytro Struk
parent
d1205896a1
commit
3dc59c83b5
@@ -46,7 +46,7 @@ from agent_framework.ag_ui import AGUIChatClient
|
||||
async def main():
|
||||
async with AGUIChatClient(endpoint="http://localhost:8000/") as client:
|
||||
# Stream responses
|
||||
async for update in client.get_streaming_response("Hello!"):
|
||||
async for update in client.get_response("Hello!", stream=True):
|
||||
for content in update.contents:
|
||||
if isinstance(content, TextContent):
|
||||
print(content.text, end="", flush=True)
|
||||
|
||||
@@ -6,9 +6,9 @@ import json
|
||||
import logging
|
||||
import sys
|
||||
import uuid
|
||||
from collections.abc import AsyncIterable, MutableSequence
|
||||
from collections.abc import AsyncIterable, Awaitable, Mapping, MutableSequence, Sequence
|
||||
from functools import wraps
|
||||
from typing import TYPE_CHECKING, Any, Generic, cast
|
||||
from typing import TYPE_CHECKING, Any, Generic, TypedDict, cast
|
||||
|
||||
import httpx
|
||||
from agent_framework import (
|
||||
@@ -18,10 +18,11 @@ from agent_framework import (
|
||||
ChatResponseUpdate,
|
||||
Content,
|
||||
FunctionTool,
|
||||
use_chat_middleware,
|
||||
use_function_invocation,
|
||||
ResponseStream,
|
||||
)
|
||||
from agent_framework.observability import use_instrumentation
|
||||
from agent_framework._middleware import ChatMiddlewareLayer
|
||||
from agent_framework._tools import FunctionInvocationConfiguration, FunctionInvocationLayer
|
||||
from agent_framework.observability import ChatTelemetryLayer
|
||||
|
||||
from ._event_converters import AGUIEventConverter
|
||||
from ._http_service import AGUIHttpService
|
||||
@@ -42,6 +43,8 @@ else:
|
||||
from typing_extensions import Self, TypedDict # pragma: no cover
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from agent_framework._middleware import ChatAndFunctionMiddlewareTypes
|
||||
|
||||
from ._types import AGUIChatOptions
|
||||
|
||||
logger: logging.Logger = logging.getLogger(__name__)
|
||||
@@ -67,35 +70,51 @@ TAGUIChatOptions = TypeVar(
|
||||
def _apply_server_function_call_unwrap(chat_client: TBaseChatClient) -> TBaseChatClient:
|
||||
"""Class decorator that unwraps server-side function calls after tool handling."""
|
||||
|
||||
original_get_streaming_response = chat_client.get_streaming_response
|
||||
|
||||
@wraps(original_get_streaming_response)
|
||||
async def streaming_wrapper(self: Any, *args: Any, **kwargs: Any) -> AsyncIterable[ChatResponseUpdate]:
|
||||
async for update in original_get_streaming_response(self, *args, **kwargs):
|
||||
_unwrap_server_function_call_contents(cast(MutableSequence[Content | dict[str, Any]], update.contents))
|
||||
yield update
|
||||
|
||||
chat_client.get_streaming_response = streaming_wrapper # type: ignore[assignment]
|
||||
|
||||
original_get_response = chat_client.get_response
|
||||
|
||||
@wraps(original_get_response)
|
||||
async def response_wrapper(self: Any, *args: Any, **kwargs: Any) -> ChatResponse:
|
||||
response: ChatResponse[Any] = await original_get_response(self, *args, **kwargs) # type: ignore[var-annotated]
|
||||
def response_wrapper(
|
||||
self, *args: Any, stream: bool = False, **kwargs: Any
|
||||
) -> Awaitable[ChatResponse] | ResponseStream[ChatResponseUpdate, ChatResponse]:
|
||||
if stream:
|
||||
stream_response = original_get_response(self, *args, stream=True, **kwargs)
|
||||
if isinstance(stream_response, ResponseStream):
|
||||
return stream_response.with_transform_hook(_map_update)
|
||||
return ResponseStream(_stream_wrapper_impl(stream_response))
|
||||
return _response_wrapper_impl(self, original_get_response, *args, **kwargs)
|
||||
|
||||
async def _response_wrapper_impl(self, original_func: Any, *args: Any, **kwargs: Any) -> ChatResponse:
|
||||
"""Non-streaming wrapper implementation."""
|
||||
response = await original_func(self, *args, stream=False, **kwargs)
|
||||
if response.messages:
|
||||
for message in response.messages:
|
||||
_unwrap_server_function_call_contents(cast(MutableSequence[Content | dict[str, Any]], message.contents))
|
||||
return response
|
||||
return response # type: ignore[no-any-return]
|
||||
|
||||
async def _stream_wrapper_impl(stream: Any) -> AsyncIterable[ChatResponseUpdate]:
|
||||
"""Streaming wrapper implementation."""
|
||||
if isinstance(stream, Awaitable):
|
||||
stream = await stream
|
||||
async for update in stream:
|
||||
_unwrap_server_function_call_contents(cast(MutableSequence[Content | dict[str, Any]], update.contents))
|
||||
yield update
|
||||
|
||||
def _map_update(update: ChatResponseUpdate) -> ChatResponseUpdate:
|
||||
_unwrap_server_function_call_contents(cast(MutableSequence[Content | dict[str, Any]], update.contents))
|
||||
return update
|
||||
|
||||
chat_client.get_response = response_wrapper # type: ignore[assignment]
|
||||
return chat_client
|
||||
|
||||
|
||||
@_apply_server_function_call_unwrap
|
||||
@use_function_invocation
|
||||
@use_instrumentation
|
||||
@use_chat_middleware
|
||||
class AGUIChatClient(BaseChatClient[TAGUIChatOptions], Generic[TAGUIChatOptions]):
|
||||
class AGUIChatClient(
|
||||
ChatMiddlewareLayer[TAGUIChatOptions],
|
||||
FunctionInvocationLayer[TAGUIChatOptions],
|
||||
ChatTelemetryLayer[TAGUIChatOptions],
|
||||
BaseChatClient[TAGUIChatOptions],
|
||||
Generic[TAGUIChatOptions],
|
||||
):
|
||||
"""Chat client for communicating with AG-UI compliant servers.
|
||||
|
||||
This client implements the BaseChatClient interface and automatically handles:
|
||||
@@ -103,6 +122,7 @@ class AGUIChatClient(BaseChatClient[TAGUIChatOptions], Generic[TAGUIChatOptions]
|
||||
- State synchronization between client and server
|
||||
- Server-Sent Events (SSE) streaming
|
||||
- Event conversion to Agent Framework types
|
||||
- MiddlewareTypes, telemetry, and function invocation support
|
||||
|
||||
Important: Message History Management
|
||||
This client sends exactly the messages it receives to the server. It does NOT
|
||||
@@ -115,10 +135,10 @@ class AGUIChatClient(BaseChatClient[TAGUIChatOptions], Generic[TAGUIChatOptions]
|
||||
Important: Tool Handling (Hybrid Execution - matches .NET)
|
||||
1. Client tool metadata sent to server - LLM knows about both client and server tools
|
||||
2. Server has its own tools that execute server-side
|
||||
3. When LLM calls a client tool, @use_function_invocation executes it locally
|
||||
3. When LLM calls a client tool, function invocation executes it locally
|
||||
4. Both client and server tools work together (hybrid pattern)
|
||||
|
||||
The wrapping ChatAgent's @use_function_invocation handles client tool execution
|
||||
The wrapping ChatAgent's function invocation handles client tool execution
|
||||
automatically when the server's LLM decides to call them.
|
||||
|
||||
Examples:
|
||||
@@ -159,7 +179,7 @@ class AGUIChatClient(BaseChatClient[TAGUIChatOptions], Generic[TAGUIChatOptions]
|
||||
|
||||
.. code-block:: python
|
||||
|
||||
async for update in client.get_streaming_response("Tell me a story"):
|
||||
async for update in client.get_response("Tell me a story", stream=True):
|
||||
if update.contents:
|
||||
for content in update.contents:
|
||||
if hasattr(content, "text"):
|
||||
@@ -196,6 +216,8 @@ class AGUIChatClient(BaseChatClient[TAGUIChatOptions], Generic[TAGUIChatOptions]
|
||||
http_client: httpx.AsyncClient | None = None,
|
||||
timeout: float = 60.0,
|
||||
additional_properties: dict[str, Any] | None = None,
|
||||
middleware: Sequence["ChatAndFunctionMiddlewareTypes"] | None = None,
|
||||
function_invocation_configuration: FunctionInvocationConfiguration | None = None,
|
||||
**kwargs: Any,
|
||||
) -> None:
|
||||
"""Initialize the AG-UI chat client.
|
||||
@@ -205,9 +227,16 @@ class AGUIChatClient(BaseChatClient[TAGUIChatOptions], Generic[TAGUIChatOptions]
|
||||
http_client: Optional httpx.AsyncClient instance. If None, one will be created.
|
||||
timeout: Request timeout in seconds (default: 60.0)
|
||||
additional_properties: Additional properties to store
|
||||
middleware: Optional middleware to apply to the client.
|
||||
function_invocation_configuration: Optional function invocation configuration override.
|
||||
**kwargs: Additional arguments passed to BaseChatClient
|
||||
"""
|
||||
super().__init__(additional_properties=additional_properties, **kwargs)
|
||||
super().__init__(
|
||||
additional_properties=additional_properties,
|
||||
middleware=middleware,
|
||||
function_invocation_configuration=function_invocation_configuration,
|
||||
**kwargs,
|
||||
)
|
||||
self._http_service = AGUIHttpService(
|
||||
endpoint=endpoint,
|
||||
http_client=http_client,
|
||||
@@ -230,9 +259,10 @@ class AGUIChatClient(BaseChatClient[TAGUIChatOptions], Generic[TAGUIChatOptions]
|
||||
"""Register a declaration-only placeholder so function invocation skips execution."""
|
||||
|
||||
config = getattr(self, "function_invocation_configuration", None)
|
||||
if not config:
|
||||
if not isinstance(config, dict):
|
||||
return
|
||||
if any(getattr(tool, "name", None) == tool_name for tool in config.additional_tools):
|
||||
additional_tools = list(config.get("additional_tools", []))
|
||||
if any(getattr(tool, "name", None) == tool_name for tool in additional_tools):
|
||||
return
|
||||
|
||||
placeholder: FunctionTool[Any, Any] = FunctionTool(
|
||||
@@ -240,7 +270,8 @@ class AGUIChatClient(BaseChatClient[TAGUIChatOptions], Generic[TAGUIChatOptions]
|
||||
description="Server-managed tool placeholder (AG-UI)",
|
||||
func=None,
|
||||
)
|
||||
config.additional_tools = list(config.additional_tools) + [placeholder]
|
||||
additional_tools.append(placeholder)
|
||||
config["additional_tools"] = additional_tools
|
||||
registered: set[str] = getattr(self, "_registered_server_tools", set())
|
||||
registered.add(tool_name)
|
||||
self._registered_server_tools = registered # type: ignore[attr-defined]
|
||||
@@ -250,7 +281,7 @@ class AGUIChatClient(BaseChatClient[TAGUIChatOptions], Generic[TAGUIChatOptions]
|
||||
logger.debug(f"[AGUIChatClient] Registered server placeholder: {tool_name}")
|
||||
|
||||
def _extract_state_from_messages(
|
||||
self, messages: MutableSequence[ChatMessage]
|
||||
self, messages: Sequence[ChatMessage]
|
||||
) -> tuple[list[ChatMessage], dict[str, Any] | None]:
|
||||
"""Extract state from last message if present.
|
||||
|
||||
@@ -297,7 +328,7 @@ class AGUIChatClient(BaseChatClient[TAGUIChatOptions], Generic[TAGUIChatOptions]
|
||||
"""
|
||||
return agent_framework_messages_to_agui(messages)
|
||||
|
||||
def _get_thread_id(self, options: dict[str, Any]) -> str:
|
||||
def _get_thread_id(self, options: Mapping[str, Any]) -> str:
|
||||
"""Get or generate thread ID from chat options.
|
||||
|
||||
Args:
|
||||
@@ -317,43 +348,57 @@ class AGUIChatClient(BaseChatClient[TAGUIChatOptions], Generic[TAGUIChatOptions]
|
||||
return thread_id
|
||||
|
||||
@override
|
||||
async def _inner_get_response(
|
||||
def _inner_get_response(
|
||||
self,
|
||||
*,
|
||||
messages: MutableSequence[ChatMessage],
|
||||
options: dict[str, Any],
|
||||
messages: Sequence[ChatMessage],
|
||||
stream: bool = False,
|
||||
options: Mapping[str, Any],
|
||||
**kwargs: Any,
|
||||
) -> ChatResponse:
|
||||
) -> Awaitable[ChatResponse] | ResponseStream[ChatResponseUpdate, ChatResponse]:
|
||||
"""Internal method to get non-streaming response.
|
||||
|
||||
Keyword Args:
|
||||
messages: List of chat messages
|
||||
stream: Whether to stream the response.
|
||||
options: Chat options for the request
|
||||
**kwargs: Additional keyword arguments
|
||||
|
||||
Returns:
|
||||
ChatResponse object
|
||||
"""
|
||||
return await ChatResponse.from_update_generator(
|
||||
self._inner_get_streaming_response(
|
||||
messages=messages,
|
||||
options=options,
|
||||
**kwargs,
|
||||
if stream:
|
||||
return ResponseStream(
|
||||
self._streaming_impl(
|
||||
messages=messages,
|
||||
options=options,
|
||||
**kwargs,
|
||||
),
|
||||
finalizer=ChatResponse.from_updates,
|
||||
)
|
||||
)
|
||||
|
||||
@override
|
||||
async def _inner_get_streaming_response(
|
||||
async def _get_response() -> ChatResponse:
|
||||
return await ChatResponse.from_update_generator(
|
||||
self._streaming_impl(
|
||||
messages=messages,
|
||||
options=options,
|
||||
**kwargs,
|
||||
)
|
||||
)
|
||||
|
||||
return _get_response()
|
||||
|
||||
async def _streaming_impl(
|
||||
self,
|
||||
*,
|
||||
messages: MutableSequence[ChatMessage],
|
||||
options: dict[str, Any],
|
||||
messages: Sequence[ChatMessage],
|
||||
options: Mapping[str, Any],
|
||||
**kwargs: Any,
|
||||
) -> AsyncIterable[ChatResponseUpdate]:
|
||||
"""Internal method to get streaming response.
|
||||
|
||||
Keyword Args:
|
||||
messages: List of chat messages
|
||||
messages: Sequence of chat messages
|
||||
options: Chat options for the request
|
||||
**kwargs: Additional keyword arguments
|
||||
|
||||
@@ -368,7 +413,7 @@ class AGUIChatClient(BaseChatClient[TAGUIChatOptions], Generic[TAGUIChatOptions]
|
||||
agui_messages = self._convert_messages_to_agui_format(messages_to_send)
|
||||
|
||||
# Send client tools to server so LLM knows about them
|
||||
# Client tools execute via ChatAgent's @use_function_invocation wrapper
|
||||
# Client tools execute via ChatAgent's function invocation wrapper
|
||||
agui_tools = convert_tools_to_agui_format(options.get("tools"))
|
||||
|
||||
# Build set of client tool names (matches .NET clientToolSet)
|
||||
@@ -415,12 +460,12 @@ class AGUIChatClient(BaseChatClient[TAGUIChatOptions], Generic[TAGUIChatOptions]
|
||||
f"[AGUIChatClient] Function call: {content.name}, in client_tool_set: {content.name in client_tool_set}" # type: ignore[attr-defined]
|
||||
)
|
||||
if content.name in client_tool_set: # type: ignore[attr-defined]
|
||||
# Client tool - let @use_function_invocation execute it
|
||||
# Client tool - let function invocation execute it
|
||||
if not content.additional_properties: # type: ignore[attr-defined]
|
||||
content.additional_properties = {} # type: ignore[attr-defined]
|
||||
content.additional_properties["agui_thread_id"] = thread_id # type: ignore[attr-defined]
|
||||
else:
|
||||
# Server tool - wrap so @use_function_invocation ignores it
|
||||
# Server tool - wrap so function invocation ignores it
|
||||
logger.debug(f"[AGUIChatClient] Wrapping server tool: {content.name}") # type: ignore[union-attr]
|
||||
self._register_server_tool_placeholder(content.name) # type: ignore[arg-type]
|
||||
update.contents[i] = Content(type="server_function_call", function_call=content) # type: ignore
|
||||
|
||||
@@ -590,7 +590,7 @@ def agui_messages_to_agent_framework(messages: list[dict[str, Any]]) -> list[Cha
|
||||
arguments=arguments,
|
||||
)
|
||||
)
|
||||
chat_msg = ChatMessage("assistant", contents)
|
||||
chat_msg = ChatMessage(role="assistant", contents=contents)
|
||||
if "id" in msg:
|
||||
chat_msg.message_id = msg["id"]
|
||||
result.append(chat_msg)
|
||||
@@ -620,14 +620,14 @@ def agui_messages_to_agent_framework(messages: list[dict[str, Any]]) -> list[Cha
|
||||
)
|
||||
approval_contents.append(approval_response)
|
||||
|
||||
chat_msg = ChatMessage(role, approval_contents) # type: ignore[arg-type]
|
||||
chat_msg = ChatMessage(role=role, contents=approval_contents) # type: ignore[call-overload]
|
||||
else:
|
||||
# Regular text message
|
||||
content = msg.get("content", "")
|
||||
if isinstance(content, str):
|
||||
chat_msg = ChatMessage(role, [Content.from_text(text=content)])
|
||||
chat_msg = ChatMessage(role=role, contents=[Content.from_text(text=content)]) # type: ignore[call-overload]
|
||||
else:
|
||||
chat_msg = ChatMessage(role, [Content.from_text(text=str(content))])
|
||||
chat_msg = ChatMessage(role=role, contents=[Content.from_text(text=str(content))]) # type: ignore[call-overload]
|
||||
|
||||
if "id" in msg:
|
||||
chat_msg.message_id = msg["id"]
|
||||
@@ -671,7 +671,8 @@ def agent_framework_messages_to_agui(messages: list[ChatMessage] | list[dict[str
|
||||
continue
|
||||
|
||||
# Convert ChatMessage to AG-UI format
|
||||
role = FRAMEWORK_TO_AGUI_ROLE.get(msg.role, "user")
|
||||
role_value: str = msg.role if hasattr(msg.role, "value") else msg.role # type: ignore[assignment]
|
||||
role = FRAMEWORK_TO_AGUI_ROLE.get(role_value, "user")
|
||||
|
||||
content_text = ""
|
||||
tool_calls: list[dict[str, Any]] = []
|
||||
|
||||
@@ -79,8 +79,8 @@ def register_additional_client_tools(agent: "AgentProtocol", client_tools: list[
|
||||
if chat_client is None:
|
||||
return
|
||||
|
||||
if isinstance(chat_client, BaseChatClient) and chat_client.function_invocation_configuration is not None:
|
||||
chat_client.function_invocation_configuration.additional_tools = client_tools
|
||||
if isinstance(chat_client, BaseChatClient) and chat_client.function_invocation_configuration is not None: # type: ignore[attr-defined]
|
||||
chat_client.function_invocation_configuration["additional_tools"] = client_tools # type: ignore[attr-defined]
|
||||
logger.debug(f"[TOOLS] Registered {len(client_tools)} client tools as additional_tools (declaration-only)")
|
||||
|
||||
|
||||
|
||||
@@ -5,8 +5,9 @@
|
||||
import json
|
||||
import logging
|
||||
import uuid
|
||||
from collections.abc import Awaitable
|
||||
from dataclasses import dataclass, field
|
||||
from typing import TYPE_CHECKING, Any
|
||||
from typing import TYPE_CHECKING, Any, cast
|
||||
|
||||
from ag_ui.core import (
|
||||
BaseEvent,
|
||||
@@ -30,13 +31,15 @@ from agent_framework import (
|
||||
Content,
|
||||
prepare_function_call_results,
|
||||
)
|
||||
from agent_framework._middleware import extract_and_merge_function_middleware
|
||||
from agent_framework._middleware import FunctionMiddlewarePipeline
|
||||
from agent_framework._tools import (
|
||||
FunctionInvocationConfiguration,
|
||||
_collect_approval_responses, # type: ignore
|
||||
_replace_approval_contents_with_results, # type: ignore
|
||||
_try_execute_function_calls, # type: ignore
|
||||
normalize_function_invocation_configuration,
|
||||
)
|
||||
from agent_framework._types import ResponseStream
|
||||
from agent_framework.exceptions import AgentExecutionException
|
||||
|
||||
from ._message_adapters import normalize_agui_input_messages
|
||||
from ._orchestration._predictive_state import PredictiveStateHandler
|
||||
@@ -601,8 +604,13 @@ async def _resolve_approval_responses(
|
||||
# Execute approved tool calls
|
||||
if approved_responses and tools:
|
||||
chat_client = getattr(agent, "chat_client", None)
|
||||
config = getattr(chat_client, "function_invocation_configuration", None) or FunctionInvocationConfiguration()
|
||||
middleware_pipeline = extract_and_merge_function_middleware(chat_client, run_kwargs)
|
||||
config = normalize_function_invocation_configuration(
|
||||
getattr(chat_client, "function_invocation_configuration", None)
|
||||
)
|
||||
middleware_pipeline = FunctionMiddlewarePipeline(
|
||||
*getattr(chat_client, "function_middleware", ()),
|
||||
*run_kwargs.get("middleware", ()),
|
||||
)
|
||||
# Filter out AG-UI-specific kwargs that should not be passed to tool execution
|
||||
tool_kwargs = {k: v for k, v in run_kwargs.items() if k != "options"}
|
||||
try:
|
||||
@@ -862,7 +870,14 @@ async def run_agent_stream(
|
||||
# Stream from agent - emit RunStarted after first update to get service IDs
|
||||
run_started_emitted = False
|
||||
all_updates: list[Any] = [] # Collect for structured output processing
|
||||
async for update in agent.run_stream(messages, **run_kwargs):
|
||||
response_stream = agent.run(messages, stream=True, **run_kwargs)
|
||||
if isinstance(response_stream, ResponseStream):
|
||||
stream = response_stream
|
||||
else:
|
||||
stream = await cast(Awaitable[ResponseStream[Any, Any]], response_stream)
|
||||
if not isinstance(stream, ResponseStream):
|
||||
raise AgentExecutionException("Chat client did not return a ResponseStream.")
|
||||
async for update in stream:
|
||||
# Collect updates for structured output processing
|
||||
if response_format is not None:
|
||||
all_updates.append(update)
|
||||
|
||||
@@ -102,7 +102,7 @@ class AGUIChatOptions(ChatOptions[TResponseModel], Generic[TResponseModel], tota
|
||||
stop: Stop sequences.
|
||||
tools: List of tools - sent to server so LLM knows about client tools.
|
||||
Server executes its own tools; client tools execute locally via
|
||||
@use_function_invocation middleware.
|
||||
function invocation middleware.
|
||||
tool_choice: How the model should use tools.
|
||||
metadata: Metadata dict containing thread_id for conversation continuity.
|
||||
|
||||
|
||||
@@ -165,7 +165,7 @@ def convert_agui_tools_to_agent_framework(
|
||||
|
||||
Creates declaration-only FunctionTool instances (no executable implementation).
|
||||
These are used to tell the LLM about available tools. The actual execution
|
||||
happens on the client side via @use_function_invocation.
|
||||
happens on the client side via function invocation mixin.
|
||||
|
||||
CRITICAL: These tools MUST have func=None so that declaration_only returns True.
|
||||
This prevents the server from trying to execute client-side tools.
|
||||
@@ -183,7 +183,7 @@ def convert_agui_tools_to_agent_framework(
|
||||
for tool_def in agui_tools:
|
||||
# Create declaration-only FunctionTool (func=None means no implementation)
|
||||
# When func=None, the declaration_only property returns True,
|
||||
# which tells @use_function_invocation to return the function call
|
||||
# which tells the function invocation mixin to return the function call
|
||||
# without executing it (so it can be sent back to the client)
|
||||
func: FunctionTool[Any, Any] = FunctionTool(
|
||||
name=tool_def.get("name", ""),
|
||||
@@ -209,7 +209,7 @@ def convert_tools_to_agui_format(
|
||||
|
||||
This sends only the metadata (name, description, JSON schema) to the server.
|
||||
The actual executable implementation stays on the client side.
|
||||
The @use_function_invocation decorator handles client-side execution when
|
||||
The function invocation mixin handles client-side execution when
|
||||
the server requests a function.
|
||||
|
||||
Args:
|
||||
|
||||
@@ -268,7 +268,7 @@ class TaskStepsAgentWithExecution:
|
||||
|
||||
# Stream completion
|
||||
accumulated_text = ""
|
||||
async for chunk in chat_client.get_streaming_response(messages=messages):
|
||||
async for chunk in chat_client.get_response(messages=messages, stream=True):
|
||||
# chunk is ChatResponseUpdate
|
||||
if hasattr(chunk, "text") and chunk.text:
|
||||
accumulated_text += chunk.text
|
||||
|
||||
+4
-1
@@ -2,6 +2,9 @@
|
||||
|
||||
"""Backend tool rendering endpoint."""
|
||||
|
||||
from typing import Any, cast
|
||||
|
||||
from agent_framework._clients import ChatClientProtocol
|
||||
from agent_framework.ag_ui import add_agent_framework_fastapi_endpoint
|
||||
from agent_framework.azure import AzureOpenAIChatClient
|
||||
from fastapi import FastAPI
|
||||
@@ -16,7 +19,7 @@ def register_backend_tool_rendering(app: FastAPI) -> None:
|
||||
app: The FastAPI application.
|
||||
"""
|
||||
# Create a chat client and call the factory function
|
||||
chat_client = AzureOpenAIChatClient()
|
||||
chat_client = cast(ChatClientProtocol[Any], AzureOpenAIChatClient())
|
||||
|
||||
add_agent_framework_fastapi_endpoint(
|
||||
app,
|
||||
|
||||
@@ -4,10 +4,11 @@
|
||||
|
||||
import logging
|
||||
import os
|
||||
from typing import cast
|
||||
|
||||
import uvicorn
|
||||
from agent_framework import ChatOptions
|
||||
from agent_framework._clients import BaseChatClient
|
||||
from agent_framework._clients import ChatClientProtocol
|
||||
from agent_framework.ag_ui import add_agent_framework_fastapi_endpoint
|
||||
from agent_framework.anthropic import AnthropicClient
|
||||
from agent_framework.azure import AzureOpenAIChatClient
|
||||
@@ -64,8 +65,9 @@ app.add_middleware(
|
||||
# Create a shared chat client for all agents
|
||||
# You can use different chat clients for different agents if needed
|
||||
# Set CHAT_CLIENT=anthropic to use Anthropic, defaults to Azure OpenAI
|
||||
chat_client: BaseChatClient[ChatOptions] = (
|
||||
AnthropicClient() if os.getenv("CHAT_CLIENT", "").lower() == "anthropic" else AzureOpenAIChatClient()
|
||||
chat_client: ChatClientProtocol[ChatOptions] = cast(
|
||||
ChatClientProtocol[ChatOptions],
|
||||
AnthropicClient() if os.getenv("CHAT_CLIENT", "").lower() == "anthropic" else AzureOpenAIChatClient(),
|
||||
)
|
||||
|
||||
# Agentic Chat - basic chat agent
|
||||
|
||||
@@ -323,7 +323,7 @@ async def main():
|
||||
# Use metadata to maintain conversation continuity
|
||||
metadata = {"thread_id": thread_id} if thread_id else None
|
||||
|
||||
async for update in client.get_streaming_response(message, metadata=metadata):
|
||||
async for update in client.get_response(message, metadata=metadata, stream=True):
|
||||
# Extract thread ID from first update
|
||||
if not thread_id and update.additional_properties:
|
||||
thread_id = update.additional_properties.get("thread_id")
|
||||
@@ -353,7 +353,7 @@ if __name__ == "__main__":
|
||||
- **`AGUIChatClient`**: Built-in client that implements the Agent Framework's `BaseChatClient` interface
|
||||
- **Automatic Event Handling**: The client automatically converts AG-UI events to Agent Framework types
|
||||
- **Thread Management**: Pass `thread_id` in metadata to maintain conversation context across requests
|
||||
- **Streaming Responses**: Use `get_streaming_response()` for real-time streaming or `get_response()` for non-streaming
|
||||
- **Streaming Responses**: Use `get_response(..., stream=True)` for real-time streaming or `get_response(..., stream=False)` for non-streaming
|
||||
- **Context Manager**: Use `async with` for automatic cleanup of HTTP connections
|
||||
- **Standard Interface**: Works with all Agent Framework patterns (ChatAgent, tools, etc.)
|
||||
- **Hybrid Tool Execution**: Supports both client-side and server-side tools executing together in the same conversation
|
||||
|
||||
@@ -9,7 +9,9 @@ standard chat interface.
|
||||
|
||||
import asyncio
|
||||
import os
|
||||
from typing import cast
|
||||
|
||||
from agent_framework import ChatResponse, ChatResponseUpdate, ResponseStream
|
||||
from agent_framework.ag_ui import AGUIChatClient
|
||||
|
||||
|
||||
@@ -41,7 +43,13 @@ async def main():
|
||||
# Use metadata to maintain conversation continuity
|
||||
metadata = {"thread_id": thread_id} if thread_id else None
|
||||
|
||||
async for update in client.get_streaming_response(message, metadata=metadata):
|
||||
stream = client.get_response(
|
||||
message,
|
||||
stream=True,
|
||||
options={"metadata": metadata} if metadata else None,
|
||||
)
|
||||
stream = cast(ResponseStream[ChatResponseUpdate, ChatResponse], stream)
|
||||
async for update in stream:
|
||||
# Extract and display thread ID from first update
|
||||
if not thread_id and update.additional_properties:
|
||||
thread_id = update.additional_properties.get("thread_id")
|
||||
@@ -51,8 +59,8 @@ async def main():
|
||||
|
||||
# Display text content as it streams
|
||||
for content in update.contents:
|
||||
if hasattr(content, "text") and content.text: # type: ignore[attr-defined]
|
||||
print(f"\033[96m{content.text}\033[0m", end="", flush=True) # type: ignore[attr-defined]
|
||||
if content.type == "text" and content.text:
|
||||
print(f"\033[96m{content.text}\033[0m", end="", flush=True)
|
||||
|
||||
# Display finish reason if present
|
||||
if update.finish_reason:
|
||||
|
||||
@@ -11,8 +11,9 @@ This example demonstrates advanced AGUIChatClient features including:
|
||||
|
||||
import asyncio
|
||||
import os
|
||||
from typing import cast
|
||||
|
||||
from agent_framework import tool
|
||||
from agent_framework import ChatResponse, ChatResponseUpdate, ResponseStream, tool
|
||||
from agent_framework.ag_ui import AGUIChatClient
|
||||
|
||||
|
||||
@@ -69,7 +70,13 @@ async def streaming_example(client: AGUIChatClient, thread_id: str | None = None
|
||||
print("\nUser: Tell me a short joke\n")
|
||||
print("Assistant: ", end="", flush=True)
|
||||
|
||||
async for update in client.get_streaming_response("Tell me a short joke", metadata=metadata):
|
||||
stream = client.get_response(
|
||||
"Tell me a short joke",
|
||||
stream=True,
|
||||
options={"metadata": metadata} if metadata else None,
|
||||
)
|
||||
stream = cast(ResponseStream[ChatResponseUpdate, ChatResponse], stream)
|
||||
async for update in stream:
|
||||
if not thread_id and update.additional_properties:
|
||||
thread_id = update.additional_properties.get("thread_id")
|
||||
|
||||
|
||||
@@ -6,11 +6,11 @@ This demonstrates the HYBRID pattern matching .NET AGUIClient implementation:
|
||||
|
||||
1. AgentThread Pattern (like .NET):
|
||||
- Create thread with agent.get_new_thread()
|
||||
- Pass thread to agent.run_stream() on each turn
|
||||
- Pass thread to agent.run(stream=True) on each turn
|
||||
- Thread automatically maintains conversation history via message_store
|
||||
|
||||
2. Hybrid Tool Execution:
|
||||
- AGUIChatClient has @use_function_invocation decorator
|
||||
- AGUIChatClient uses function invocation mixin
|
||||
- Client-side tools (get_weather) can execute locally when server requests them
|
||||
- Server may also have its own tools that execute server-side
|
||||
- Both work together: server LLM decides which tool to call, decorator handles client execution
|
||||
@@ -63,7 +63,7 @@ async def main():
|
||||
Python equivalent:
|
||||
- agent = ChatAgent(chat_client=AGUIChatClient(...), tools=[...])
|
||||
- thread = agent.get_new_thread() # Creates thread with message_store
|
||||
- agent.run_stream(message, thread=thread) # Thread accumulates history
|
||||
- agent.run(message, stream=True, thread=thread) # Thread accumulates history
|
||||
"""
|
||||
server_url = os.environ.get("AGUI_SERVER_URL", "http://127.0.0.1:5100/")
|
||||
|
||||
@@ -73,7 +73,7 @@ async def main():
|
||||
print(f"\nServer: {server_url}")
|
||||
print("\nThis example demonstrates:")
|
||||
print(" 1. AgentThread maintains conversation state (like .NET)")
|
||||
print(" 2. Client-side tools execute locally via @use_function_invocation")
|
||||
print(" 2. Client-side tools execute locally via function invocation mixin")
|
||||
print(" 3. Server may have additional tools that execute server-side")
|
||||
print(" 4. HYBRID: Client and server tools work together simultaneously\n")
|
||||
|
||||
@@ -97,35 +97,39 @@ async def main():
|
||||
|
||||
# Turn 1: Introduce
|
||||
print("\nUser: My name is Alice and I live in Seattle\n")
|
||||
async for chunk in agent.run_stream("My name is Alice and I live in Seattle", thread=thread):
|
||||
async for chunk in agent.run("My name is Alice and I live in Seattle", stream=True, thread=thread):
|
||||
if chunk.text:
|
||||
print(chunk.text, end="", flush=True)
|
||||
print("\n")
|
||||
|
||||
# Turn 2: Ask about name (tests history)
|
||||
print("User: What's my name?\n")
|
||||
async for chunk in agent.run_stream("What's my name?", thread=thread):
|
||||
async for chunk in agent.run("What's my name?", stream=True, thread=thread):
|
||||
if chunk.text:
|
||||
print(chunk.text, end="", flush=True)
|
||||
print("\n")
|
||||
|
||||
# Turn 3: Ask about location (tests history)
|
||||
print("User: Where do I live?\n")
|
||||
async for chunk in agent.run_stream("Where do I live?", thread=thread):
|
||||
async for chunk in agent.run("Where do I live?", stream=True, thread=thread):
|
||||
if chunk.text:
|
||||
print(chunk.text, end="", flush=True)
|
||||
print("\n")
|
||||
|
||||
# Turn 4: Test client-side tool (get_weather is client-side)
|
||||
print("User: What's the weather forecast for today in Seattle?\n")
|
||||
async for chunk in agent.run_stream("What's the weather forecast for today in Seattle?", thread=thread):
|
||||
async for chunk in agent.run(
|
||||
"What's the weather forecast for today in Seattle?",
|
||||
stream=True,
|
||||
thread=thread,
|
||||
):
|
||||
if chunk.text:
|
||||
print(chunk.text, end="", flush=True)
|
||||
print("\n")
|
||||
|
||||
# Turn 5: Test server-side tool (get_time_zone is server-side only)
|
||||
print("User: What time zone is Seattle in?\n")
|
||||
async for chunk in agent.run_stream("What time zone is Seattle in?", thread=thread):
|
||||
async for chunk in agent.run("What time zone is Seattle in?", stream=True, thread=thread):
|
||||
if chunk.text:
|
||||
print(chunk.text, end="", flush=True)
|
||||
print("\n")
|
||||
|
||||
@@ -112,7 +112,7 @@ def get_time_zone(location: str) -> str:
|
||||
# - get_time_zone: SERVER-ONLY tool (only server has this)
|
||||
# - get_weather: CLIENT-ONLY tool (client provides this, server should NOT include it)
|
||||
# The client will send get_weather tool metadata so the LLM knows about it,
|
||||
# and @use_function_invocation on AGUIChatClient will execute it client-side.
|
||||
# and the function invocation mixin on AGUIChatClient will execute it client-side.
|
||||
# This matches the .NET AG-UI hybrid execution pattern.
|
||||
agent = ChatAgent(
|
||||
name="AGUIAssistant",
|
||||
|
||||
@@ -31,7 +31,6 @@ dependencies = [
|
||||
[project.optional-dependencies]
|
||||
dev = [
|
||||
"pytest>=8.0.0",
|
||||
"pytest-asyncio>=0.24.0",
|
||||
"httpx>=0.27.0",
|
||||
]
|
||||
|
||||
@@ -44,7 +43,7 @@ packages = ["agent_framework_ag_ui", "agent_framework_ag_ui_examples"]
|
||||
|
||||
[tool.pytest.ini_options]
|
||||
asyncio_mode = "auto"
|
||||
testpaths = ["tests"]
|
||||
testpaths = ["tests/ag_ui"]
|
||||
pythonpath = ["."]
|
||||
|
||||
[tool.ruff]
|
||||
@@ -62,7 +61,7 @@ warn_unused_configs = true
|
||||
disallow_untyped_defs = false
|
||||
|
||||
[tool.pyright]
|
||||
exclude = ["tests", "examples"]
|
||||
exclude = ["tests", "tests/ag_ui", "examples"]
|
||||
typeCheckingMode = "basic"
|
||||
|
||||
[tool.poe]
|
||||
@@ -71,4 +70,4 @@ include = "../../shared_tasks.toml"
|
||||
|
||||
[tool.poe.tasks]
|
||||
mypy = "mypy --config-file $POE_ROOT/pyproject.toml agent_framework_ag_ui"
|
||||
test = "pytest --cov=agent_framework_ag_ui --cov-report=term-missing:skip-covered tests"
|
||||
test = "pytest --cov=agent_framework_ag_ui --cov-report=term-missing:skip-covered tests/ag_ui"
|
||||
|
||||
@@ -0,0 +1,243 @@
|
||||
# Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
"""Shared test fixtures and stubs for AG-UI tests."""
|
||||
|
||||
import sys
|
||||
from collections.abc import AsyncIterable, AsyncIterator, Awaitable, Callable, Mapping, MutableSequence, Sequence
|
||||
from types import SimpleNamespace
|
||||
from typing import Any, Generic, Literal, cast, overload
|
||||
|
||||
import pytest
|
||||
from agent_framework import (
|
||||
AgentProtocol,
|
||||
AgentResponse,
|
||||
AgentResponseUpdate,
|
||||
AgentThread,
|
||||
BaseChatClient,
|
||||
ChatClientProtocol,
|
||||
ChatMessage,
|
||||
ChatOptions,
|
||||
ChatResponse,
|
||||
ChatResponseUpdate,
|
||||
Content,
|
||||
)
|
||||
from agent_framework._clients import TOptions_co
|
||||
from agent_framework._middleware import ChatMiddlewareLayer
|
||||
from agent_framework._tools import FunctionInvocationLayer
|
||||
from agent_framework._types import ResponseStream
|
||||
from agent_framework.observability import ChatTelemetryLayer
|
||||
|
||||
if sys.version_info >= (3, 12):
|
||||
from typing import override # type: ignore # pragma: no cover
|
||||
else:
|
||||
from typing_extensions import override # type: ignore[import] # pragma: no cover
|
||||
|
||||
StreamFn = Callable[..., AsyncIterable[ChatResponseUpdate]]
|
||||
ResponseFn = Callable[..., Awaitable[ChatResponse]]
|
||||
|
||||
|
||||
class StreamingChatClientStub(
|
||||
ChatMiddlewareLayer[TOptions_co],
|
||||
FunctionInvocationLayer[TOptions_co],
|
||||
ChatTelemetryLayer[TOptions_co],
|
||||
BaseChatClient[TOptions_co],
|
||||
Generic[TOptions_co],
|
||||
):
|
||||
"""Typed streaming stub that satisfies ChatClientProtocol."""
|
||||
|
||||
def __init__(self, stream_fn: StreamFn, response_fn: ResponseFn | None = None) -> None:
|
||||
super().__init__(function_middleware=[])
|
||||
self._stream_fn = stream_fn
|
||||
self._response_fn = response_fn
|
||||
self.last_thread: AgentThread | None = None
|
||||
self.last_service_thread_id: str | None = None
|
||||
|
||||
@overload
|
||||
def get_response(
|
||||
self,
|
||||
messages: str | ChatMessage | Sequence[str | ChatMessage],
|
||||
*,
|
||||
stream: Literal[False] = ...,
|
||||
options: ChatOptions[Any],
|
||||
**kwargs: Any,
|
||||
) -> Awaitable[ChatResponse[Any]]: ...
|
||||
|
||||
@overload
|
||||
def get_response(
|
||||
self,
|
||||
messages: str | ChatMessage | Sequence[str | ChatMessage],
|
||||
*,
|
||||
stream: Literal[False] = ...,
|
||||
options: TOptions_co | ChatOptions[None] | None = ...,
|
||||
**kwargs: Any,
|
||||
) -> Awaitable[ChatResponse[Any]]: ...
|
||||
|
||||
@overload
|
||||
def get_response(
|
||||
self,
|
||||
messages: str | ChatMessage | Sequence[str | ChatMessage],
|
||||
*,
|
||||
stream: Literal[True],
|
||||
options: TOptions_co | ChatOptions[Any] | None = ...,
|
||||
**kwargs: Any,
|
||||
) -> ResponseStream[ChatResponseUpdate, ChatResponse[Any]]: ...
|
||||
|
||||
def get_response(
|
||||
self,
|
||||
messages: str | ChatMessage | Sequence[str | ChatMessage],
|
||||
*,
|
||||
stream: bool = False,
|
||||
options: TOptions_co | ChatOptions[Any] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> Awaitable[ChatResponse[Any]] | ResponseStream[ChatResponseUpdate, ChatResponse[Any]]:
|
||||
self.last_thread = kwargs.get("thread")
|
||||
self.last_service_thread_id = self.last_thread.service_thread_id if self.last_thread else None
|
||||
return cast(
|
||||
Awaitable[ChatResponse[Any]] | ResponseStream[ChatResponseUpdate, ChatResponse[Any]],
|
||||
super().get_response(
|
||||
messages=messages,
|
||||
stream=cast(Literal[True, False], stream),
|
||||
options=options,
|
||||
**kwargs,
|
||||
),
|
||||
)
|
||||
|
||||
@override
|
||||
def _inner_get_response(
|
||||
self,
|
||||
*,
|
||||
messages: Sequence[ChatMessage],
|
||||
stream: bool = False,
|
||||
options: Mapping[str, Any],
|
||||
**kwargs: Any,
|
||||
) -> Awaitable[ChatResponse] | ResponseStream[ChatResponseUpdate, ChatResponse]:
|
||||
if stream:
|
||||
|
||||
def _finalize(updates: Sequence[ChatResponseUpdate]) -> ChatResponse:
|
||||
return ChatResponse.from_updates(updates)
|
||||
|
||||
return ResponseStream(self._stream_fn(messages, options, **kwargs), finalizer=_finalize)
|
||||
|
||||
return self._get_response_impl(messages, options, **kwargs)
|
||||
|
||||
async def _get_response_impl(
|
||||
self, messages: Sequence[ChatMessage], options: Mapping[str, Any], **kwargs: Any
|
||||
) -> ChatResponse:
|
||||
"""Non-streaming implementation."""
|
||||
if self._response_fn is not None:
|
||||
return await self._response_fn(messages, options, **kwargs)
|
||||
|
||||
contents: list[Any] = []
|
||||
async for update in self._stream_fn(list(messages), dict(options), **kwargs):
|
||||
contents.extend(update.contents)
|
||||
|
||||
return ChatResponse(
|
||||
messages=[ChatMessage(role="assistant", contents=contents)],
|
||||
response_id="stub-response",
|
||||
)
|
||||
|
||||
|
||||
def stream_from_updates(updates: list[ChatResponseUpdate]) -> StreamFn:
|
||||
"""Create a stream function that yields from a static list of updates."""
|
||||
|
||||
async def _stream(
|
||||
messages: MutableSequence[ChatMessage], options: dict[str, Any], **kwargs: Any
|
||||
) -> AsyncIterator[ChatResponseUpdate]:
|
||||
for update in updates:
|
||||
yield update
|
||||
|
||||
return _stream
|
||||
|
||||
|
||||
class StubAgent(AgentProtocol):
|
||||
"""Minimal AgentProtocol stub for orchestrator tests."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
updates: list[AgentResponseUpdate] | None = None,
|
||||
*,
|
||||
agent_id: str = "stub-agent",
|
||||
agent_name: str | None = "stub-agent",
|
||||
default_options: Any | None = None,
|
||||
chat_client: Any | None = None,
|
||||
) -> None:
|
||||
self.id = agent_id
|
||||
self.name = agent_name
|
||||
self.description = "stub agent"
|
||||
self.updates = updates or [AgentResponseUpdate(contents=[Content.from_text(text="response")], role="assistant")]
|
||||
self.default_options: dict[str, Any] = (
|
||||
default_options if isinstance(default_options, dict) else {"tools": None, "response_format": None}
|
||||
)
|
||||
self.chat_client = chat_client or SimpleNamespace(function_invocation_configuration=None)
|
||||
self.messages_received: list[Any] = []
|
||||
self.tools_received: list[Any] | None = None
|
||||
|
||||
@overload
|
||||
def run(
|
||||
self,
|
||||
messages: str | ChatMessage | Sequence[str | ChatMessage] | None = None,
|
||||
*,
|
||||
stream: Literal[False] = ...,
|
||||
thread: AgentThread | None = None,
|
||||
**kwargs: Any,
|
||||
) -> Awaitable[AgentResponse[Any]]: ...
|
||||
|
||||
@overload
|
||||
def run(
|
||||
self,
|
||||
messages: str | ChatMessage | Sequence[str | ChatMessage] | None = None,
|
||||
*,
|
||||
stream: Literal[True],
|
||||
thread: AgentThread | None = None,
|
||||
**kwargs: Any,
|
||||
) -> ResponseStream[AgentResponseUpdate, AgentResponse[Any]]: ...
|
||||
|
||||
def run(
|
||||
self,
|
||||
messages: str | ChatMessage | Sequence[str | ChatMessage] | None = None,
|
||||
*,
|
||||
stream: bool = False,
|
||||
thread: AgentThread | None = None,
|
||||
**kwargs: Any,
|
||||
) -> Awaitable[AgentResponse[Any]] | ResponseStream[AgentResponseUpdate, AgentResponse[Any]]:
|
||||
if stream:
|
||||
|
||||
async def _stream() -> AsyncIterator[AgentResponseUpdate]:
|
||||
self.messages_received = [] if messages is None else list(messages) # type: ignore[arg-type]
|
||||
self.tools_received = kwargs.get("tools")
|
||||
for update in self.updates:
|
||||
yield update
|
||||
|
||||
def _finalize(updates: Sequence[AgentResponseUpdate]) -> AgentResponse:
|
||||
return AgentResponse.from_updates(updates)
|
||||
|
||||
return ResponseStream(_stream(), finalizer=_finalize)
|
||||
|
||||
async def _get_response() -> AgentResponse[Any]:
|
||||
return AgentResponse(messages=[], response_id="stub-response")
|
||||
|
||||
return _get_response()
|
||||
|
||||
def get_new_thread(self, **kwargs: Any) -> AgentThread:
|
||||
return AgentThread()
|
||||
|
||||
|
||||
# Fixtures
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def streaming_chat_client_stub() -> type[ChatClientProtocol]:
|
||||
"""Return the StreamingChatClientStub class for creating test instances."""
|
||||
return StreamingChatClientStub # type: ignore[return-value]
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def stream_from_updates_fixture() -> Callable[[list[ChatResponseUpdate]], StreamFn]:
|
||||
"""Return the stream_from_updates helper function."""
|
||||
return stream_from_updates
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def stub_agent() -> type[AgentProtocol]:
|
||||
"""Return the StubAgent class for creating test instances."""
|
||||
return StubAgent # type: ignore[return-value]
|
||||
+24
-28
@@ -3,7 +3,7 @@
|
||||
"""Tests for AGUIChatClient."""
|
||||
|
||||
import json
|
||||
from collections.abc import AsyncGenerator, AsyncIterable, MutableSequence
|
||||
from collections.abc import AsyncGenerator, Awaitable, MutableSequence
|
||||
from typing import Any
|
||||
|
||||
from agent_framework import (
|
||||
@@ -12,6 +12,7 @@ from agent_framework import (
|
||||
ChatResponse,
|
||||
ChatResponseUpdate,
|
||||
Content,
|
||||
ResponseStream,
|
||||
tool,
|
||||
)
|
||||
from pytest import MonkeyPatch
|
||||
@@ -42,18 +43,11 @@ class TestableAGUIChatClient(AGUIChatClient):
|
||||
"""Expose thread id helper."""
|
||||
return self._get_thread_id(options)
|
||||
|
||||
async def inner_get_streaming_response(
|
||||
self, *, messages: MutableSequence[ChatMessage], options: dict[str, Any]
|
||||
) -> AsyncIterable[ChatResponseUpdate]:
|
||||
"""Proxy to protected streaming call."""
|
||||
async for update in self._inner_get_streaming_response(messages=messages, options=options):
|
||||
yield update
|
||||
|
||||
async def inner_get_response(
|
||||
self, *, messages: MutableSequence[ChatMessage], options: dict[str, Any]
|
||||
) -> ChatResponse:
|
||||
def inner_get_response(
|
||||
self, *, messages: MutableSequence[ChatMessage], options: dict[str, Any], stream: bool = False
|
||||
) -> Awaitable[ChatResponse] | ResponseStream[ChatResponseUpdate, ChatResponse]:
|
||||
"""Proxy to protected response call."""
|
||||
return await self._inner_get_response(messages=messages, options=options)
|
||||
return self._inner_get_response(messages=messages, options=options, stream=stream)
|
||||
|
||||
|
||||
class TestAGUIChatClient:
|
||||
@@ -75,8 +69,8 @@ class TestAGUIChatClient:
|
||||
"""Test state extraction when no state is present."""
|
||||
client = TestableAGUIChatClient(endpoint="http://localhost:8888/")
|
||||
messages = [
|
||||
ChatMessage("user", ["Hello"]),
|
||||
ChatMessage("assistant", ["Hi there"]),
|
||||
ChatMessage(role="user", text="Hello"),
|
||||
ChatMessage(role="assistant", text="Hi there"),
|
||||
]
|
||||
|
||||
result_messages, state = client.extract_state_from_messages(messages)
|
||||
@@ -95,7 +89,7 @@ class TestAGUIChatClient:
|
||||
state_b64 = base64.b64encode(state_json.encode("utf-8")).decode("utf-8")
|
||||
|
||||
messages = [
|
||||
ChatMessage("user", ["Hello"]),
|
||||
ChatMessage(role="user", text="Hello"),
|
||||
ChatMessage(
|
||||
role="user",
|
||||
contents=[Content.from_uri(uri=f"data:application/json;base64,{state_b64}")],
|
||||
@@ -133,8 +127,8 @@ class TestAGUIChatClient:
|
||||
"""Test message conversion to AG-UI format."""
|
||||
client = TestableAGUIChatClient(endpoint="http://localhost:8888/")
|
||||
messages = [
|
||||
ChatMessage("user", ["What is the weather?"]),
|
||||
ChatMessage("assistant", ["Let me check."], message_id="msg_123"),
|
||||
ChatMessage(role="user", text="What is the weather?"),
|
||||
ChatMessage(role="assistant", text="Let me check.", message_id="msg_123"),
|
||||
]
|
||||
|
||||
agui_messages = client.convert_messages_to_agui_format(messages)
|
||||
@@ -165,7 +159,7 @@ class TestAGUIChatClient:
|
||||
assert thread_id.startswith("thread_")
|
||||
assert len(thread_id) > 7
|
||||
|
||||
async def test_get_streaming_response(self, monkeypatch: MonkeyPatch) -> None:
|
||||
async def test_get_response_streaming(self, monkeypatch: MonkeyPatch) -> None:
|
||||
"""Test streaming response method."""
|
||||
mock_events = [
|
||||
{"type": "RUN_STARTED", "threadId": "thread_1", "runId": "run_1"},
|
||||
@@ -181,11 +175,11 @@ class TestAGUIChatClient:
|
||||
client = TestableAGUIChatClient(endpoint="http://localhost:8888/")
|
||||
monkeypatch.setattr(client.http_service, "post_run", mock_post_run)
|
||||
|
||||
messages = [ChatMessage("user", ["Test message"])]
|
||||
messages = [ChatMessage(role="user", text="Test message")]
|
||||
chat_options = ChatOptions()
|
||||
|
||||
updates: list[ChatResponseUpdate] = []
|
||||
async for update in client.inner_get_streaming_response(messages=messages, options=chat_options):
|
||||
async for update in client._inner_get_response(messages=messages, stream=True, options=chat_options):
|
||||
updates.append(update)
|
||||
|
||||
assert len(updates) == 4
|
||||
@@ -214,7 +208,7 @@ class TestAGUIChatClient:
|
||||
client = TestableAGUIChatClient(endpoint="http://localhost:8888/")
|
||||
monkeypatch.setattr(client.http_service, "post_run", mock_post_run)
|
||||
|
||||
messages = [ChatMessage("user", ["Test message"])]
|
||||
messages = [ChatMessage(role="user", text="Test message")]
|
||||
chat_options = {}
|
||||
|
||||
response = await client.inner_get_response(messages=messages, options=chat_options)
|
||||
@@ -227,7 +221,7 @@ class TestAGUIChatClient:
|
||||
"""Test that client tool metadata is sent to server.
|
||||
|
||||
Client tool metadata (name, description, schema) is sent to server for planning.
|
||||
When server requests a client function, @use_function_invocation decorator
|
||||
When server requests a client function, function invocation mixin
|
||||
intercepts and executes it locally. This matches .NET AG-UI implementation.
|
||||
"""
|
||||
from agent_framework import tool
|
||||
@@ -257,7 +251,7 @@ class TestAGUIChatClient:
|
||||
client = TestableAGUIChatClient(endpoint="http://localhost:8888/")
|
||||
monkeypatch.setattr(client.http_service, "post_run", mock_post_run)
|
||||
|
||||
messages = [ChatMessage("user", ["Test with tools"])]
|
||||
messages = [ChatMessage(role="user", text="Test with tools")]
|
||||
chat_options = ChatOptions(tools=[test_tool])
|
||||
|
||||
response = await client.inner_get_response(messages=messages, options=chat_options)
|
||||
@@ -281,10 +275,10 @@ class TestAGUIChatClient:
|
||||
client = TestableAGUIChatClient(endpoint="http://localhost:8888/")
|
||||
monkeypatch.setattr(client.http_service, "post_run", mock_post_run)
|
||||
|
||||
messages = [ChatMessage("user", ["Test server tool execution"])]
|
||||
messages = [ChatMessage(role="user", text="Test server tool execution")]
|
||||
|
||||
updates: list[ChatResponseUpdate] = []
|
||||
async for update in client.get_streaming_response(messages):
|
||||
async for update in client.get_response(messages, stream=True):
|
||||
updates.append(update)
|
||||
|
||||
function_calls = [
|
||||
@@ -323,9 +317,11 @@ class TestAGUIChatClient:
|
||||
client = TestableAGUIChatClient(endpoint="http://localhost:8888/")
|
||||
monkeypatch.setattr(client.http_service, "post_run", mock_post_run)
|
||||
|
||||
messages = [ChatMessage("user", ["Test server tool execution"])]
|
||||
messages = [ChatMessage(role="user", text="Test server tool execution")]
|
||||
|
||||
async for _ in client.get_streaming_response(messages, options={"tool_choice": "auto", "tools": [client_tool]}):
|
||||
async for _ in client.get_response(
|
||||
messages, stream=True, options={"tool_choice": "auto", "tools": [client_tool]}
|
||||
):
|
||||
pass
|
||||
|
||||
async def test_state_transmission(self, monkeypatch: MonkeyPatch) -> None:
|
||||
@@ -337,7 +333,7 @@ class TestAGUIChatClient:
|
||||
state_b64 = base64.b64encode(state_json.encode("utf-8")).decode("utf-8")
|
||||
|
||||
messages = [
|
||||
ChatMessage("user", ["Hello"]),
|
||||
ChatMessage(role="user", text="Hello"),
|
||||
ChatMessage(
|
||||
role="user",
|
||||
contents=[Content.from_uri(uri=f"data:application/json;base64,{state_b64}")],
|
||||
+53
-68
@@ -3,20 +3,15 @@
|
||||
"""Comprehensive tests for AgentFrameworkAgent (_agent.py)."""
|
||||
|
||||
import json
|
||||
import sys
|
||||
from collections.abc import AsyncIterator, MutableSequence
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
from agent_framework import ChatAgent, ChatMessage, ChatOptions, ChatResponseUpdate, Content
|
||||
from pydantic import BaseModel
|
||||
|
||||
sys.path.insert(0, str(Path(__file__).parent))
|
||||
from utils_test_ag_ui import StreamingChatClientStub
|
||||
|
||||
|
||||
async def test_agent_initialization_basic():
|
||||
async def test_agent_initialization_basic(streaming_chat_client_stub):
|
||||
"""Test basic agent initialization without state schema."""
|
||||
from agent_framework.ag_ui import AgentFrameworkAgent
|
||||
|
||||
@@ -26,7 +21,7 @@ async def test_agent_initialization_basic():
|
||||
yield ChatResponseUpdate(contents=[Content.from_text(text="Hello")])
|
||||
|
||||
agent = ChatAgent[ChatOptions](
|
||||
chat_client=StreamingChatClientStub(stream_fn),
|
||||
chat_client=streaming_chat_client_stub(stream_fn),
|
||||
name="test_agent",
|
||||
instructions="Test",
|
||||
)
|
||||
@@ -38,7 +33,7 @@ async def test_agent_initialization_basic():
|
||||
assert wrapper.config.predict_state_config == {}
|
||||
|
||||
|
||||
async def test_agent_initialization_with_state_schema():
|
||||
async def test_agent_initialization_with_state_schema(streaming_chat_client_stub):
|
||||
"""Test agent initialization with state_schema."""
|
||||
from agent_framework.ag_ui import AgentFrameworkAgent
|
||||
|
||||
@@ -47,14 +42,14 @@ async def test_agent_initialization_with_state_schema():
|
||||
) -> AsyncIterator[ChatResponseUpdate]:
|
||||
yield ChatResponseUpdate(contents=[Content.from_text(text="Hello")])
|
||||
|
||||
agent = ChatAgent(name="test_agent", instructions="Test", chat_client=StreamingChatClientStub(stream_fn))
|
||||
agent = ChatAgent(name="test_agent", instructions="Test", chat_client=streaming_chat_client_stub(stream_fn))
|
||||
state_schema: dict[str, dict[str, Any]] = {"document": {"type": "string"}}
|
||||
wrapper = AgentFrameworkAgent(agent=agent, state_schema=state_schema)
|
||||
|
||||
assert wrapper.config.state_schema == state_schema
|
||||
|
||||
|
||||
async def test_agent_initialization_with_predict_state_config():
|
||||
async def test_agent_initialization_with_predict_state_config(streaming_chat_client_stub):
|
||||
"""Test agent initialization with predict_state_config."""
|
||||
from agent_framework.ag_ui import AgentFrameworkAgent
|
||||
|
||||
@@ -63,14 +58,14 @@ async def test_agent_initialization_with_predict_state_config():
|
||||
) -> AsyncIterator[ChatResponseUpdate]:
|
||||
yield ChatResponseUpdate(contents=[Content.from_text(text="Hello")])
|
||||
|
||||
agent = ChatAgent(name="test_agent", instructions="Test", chat_client=StreamingChatClientStub(stream_fn))
|
||||
agent = ChatAgent(name="test_agent", instructions="Test", chat_client=streaming_chat_client_stub(stream_fn))
|
||||
predict_config = {"document": {"tool": "write_doc", "tool_argument": "content"}}
|
||||
wrapper = AgentFrameworkAgent(agent=agent, predict_state_config=predict_config)
|
||||
|
||||
assert wrapper.config.predict_state_config == predict_config
|
||||
|
||||
|
||||
async def test_agent_initialization_with_pydantic_state_schema():
|
||||
async def test_agent_initialization_with_pydantic_state_schema(streaming_chat_client_stub):
|
||||
"""Test agent initialization when state_schema is provided as Pydantic model/class."""
|
||||
from agent_framework.ag_ui import AgentFrameworkAgent
|
||||
|
||||
@@ -83,7 +78,7 @@ async def test_agent_initialization_with_pydantic_state_schema():
|
||||
document: str
|
||||
tags: list[str] = []
|
||||
|
||||
agent = ChatAgent(name="test_agent", instructions="Test", chat_client=StreamingChatClientStub(stream_fn))
|
||||
agent = ChatAgent(name="test_agent", instructions="Test", chat_client=streaming_chat_client_stub(stream_fn))
|
||||
|
||||
wrapper_class_schema = AgentFrameworkAgent(agent=agent, state_schema=MyState)
|
||||
wrapper_instance_schema = AgentFrameworkAgent(agent=agent, state_schema=MyState(document="hi"))
|
||||
@@ -93,7 +88,7 @@ async def test_agent_initialization_with_pydantic_state_schema():
|
||||
assert wrapper_instance_schema.config.state_schema == expected_properties
|
||||
|
||||
|
||||
async def test_run_started_event_emission():
|
||||
async def test_run_started_event_emission(streaming_chat_client_stub):
|
||||
"""Test RunStartedEvent is emitted at start of run."""
|
||||
from agent_framework.ag_ui import AgentFrameworkAgent
|
||||
|
||||
@@ -102,7 +97,7 @@ async def test_run_started_event_emission():
|
||||
) -> AsyncIterator[ChatResponseUpdate]:
|
||||
yield ChatResponseUpdate(contents=[Content.from_text(text="Hello")])
|
||||
|
||||
agent = ChatAgent(name="test_agent", instructions="Test", chat_client=StreamingChatClientStub(stream_fn))
|
||||
agent = ChatAgent(name="test_agent", instructions="Test", chat_client=streaming_chat_client_stub(stream_fn))
|
||||
wrapper = AgentFrameworkAgent(agent=agent)
|
||||
|
||||
input_data = {"messages": [{"role": "user", "content": "Hi"}]}
|
||||
@@ -117,7 +112,7 @@ async def test_run_started_event_emission():
|
||||
assert events[0].thread_id is not None
|
||||
|
||||
|
||||
async def test_predict_state_custom_event_emission():
|
||||
async def test_predict_state_custom_event_emission(streaming_chat_client_stub):
|
||||
"""Test PredictState CustomEvent is emitted when predict_state_config is present."""
|
||||
from agent_framework.ag_ui import AgentFrameworkAgent
|
||||
|
||||
@@ -126,7 +121,7 @@ async def test_predict_state_custom_event_emission():
|
||||
) -> AsyncIterator[ChatResponseUpdate]:
|
||||
yield ChatResponseUpdate(contents=[Content.from_text(text="Hello")])
|
||||
|
||||
agent = ChatAgent(name="test_agent", instructions="Test", chat_client=StreamingChatClientStub(stream_fn))
|
||||
agent = ChatAgent(name="test_agent", instructions="Test", chat_client=streaming_chat_client_stub(stream_fn))
|
||||
predict_config = {
|
||||
"document": {"tool": "write_doc", "tool_argument": "content"},
|
||||
"summary": {"tool": "summarize", "tool_argument": "text"},
|
||||
@@ -149,7 +144,7 @@ async def test_predict_state_custom_event_emission():
|
||||
assert {"state_key": "summary", "tool": "summarize", "tool_argument": "text"} in predict_value
|
||||
|
||||
|
||||
async def test_initial_state_snapshot_with_schema():
|
||||
async def test_initial_state_snapshot_with_schema(streaming_chat_client_stub):
|
||||
"""Test initial StateSnapshotEvent emission when state_schema present."""
|
||||
from agent_framework.ag_ui import AgentFrameworkAgent
|
||||
|
||||
@@ -158,7 +153,7 @@ async def test_initial_state_snapshot_with_schema():
|
||||
) -> AsyncIterator[ChatResponseUpdate]:
|
||||
yield ChatResponseUpdate(contents=[Content.from_text(text="Hello")])
|
||||
|
||||
agent = ChatAgent(name="test_agent", instructions="Test", chat_client=StreamingChatClientStub(stream_fn))
|
||||
agent = ChatAgent(name="test_agent", instructions="Test", chat_client=streaming_chat_client_stub(stream_fn))
|
||||
state_schema = {"document": {"type": "string"}}
|
||||
wrapper = AgentFrameworkAgent(agent=agent, state_schema=state_schema)
|
||||
|
||||
@@ -179,7 +174,7 @@ async def test_initial_state_snapshot_with_schema():
|
||||
assert snapshot_events[0].snapshot == {"document": "Initial content"}
|
||||
|
||||
|
||||
async def test_state_initialization_object_type():
|
||||
async def test_state_initialization_object_type(streaming_chat_client_stub):
|
||||
"""Test state initialization with object type in schema."""
|
||||
from agent_framework.ag_ui import AgentFrameworkAgent
|
||||
|
||||
@@ -188,7 +183,7 @@ async def test_state_initialization_object_type():
|
||||
) -> AsyncIterator[ChatResponseUpdate]:
|
||||
yield ChatResponseUpdate(contents=[Content.from_text(text="Hello")])
|
||||
|
||||
agent = ChatAgent(name="test_agent", instructions="Test", chat_client=StreamingChatClientStub(stream_fn))
|
||||
agent = ChatAgent(name="test_agent", instructions="Test", chat_client=streaming_chat_client_stub(stream_fn))
|
||||
state_schema: dict[str, dict[str, Any]] = {"recipe": {"type": "object", "properties": {}}}
|
||||
wrapper = AgentFrameworkAgent(agent=agent, state_schema=state_schema)
|
||||
|
||||
@@ -206,7 +201,7 @@ async def test_state_initialization_object_type():
|
||||
assert snapshot_events[0].snapshot == {"recipe": {}}
|
||||
|
||||
|
||||
async def test_state_initialization_array_type():
|
||||
async def test_state_initialization_array_type(streaming_chat_client_stub):
|
||||
"""Test state initialization with array type in schema."""
|
||||
from agent_framework.ag_ui import AgentFrameworkAgent
|
||||
|
||||
@@ -215,7 +210,7 @@ async def test_state_initialization_array_type():
|
||||
) -> AsyncIterator[ChatResponseUpdate]:
|
||||
yield ChatResponseUpdate(contents=[Content.from_text(text="Hello")])
|
||||
|
||||
agent = ChatAgent(name="test_agent", instructions="Test", chat_client=StreamingChatClientStub(stream_fn))
|
||||
agent = ChatAgent(name="test_agent", instructions="Test", chat_client=streaming_chat_client_stub(stream_fn))
|
||||
state_schema: dict[str, dict[str, Any]] = {"steps": {"type": "array", "items": {}}}
|
||||
wrapper = AgentFrameworkAgent(agent=agent, state_schema=state_schema)
|
||||
|
||||
@@ -233,7 +228,7 @@ async def test_state_initialization_array_type():
|
||||
assert snapshot_events[0].snapshot == {"steps": []}
|
||||
|
||||
|
||||
async def test_run_finished_event_emission():
|
||||
async def test_run_finished_event_emission(streaming_chat_client_stub):
|
||||
"""Test RunFinishedEvent is emitted at end of run."""
|
||||
from agent_framework.ag_ui import AgentFrameworkAgent
|
||||
|
||||
@@ -242,7 +237,7 @@ async def test_run_finished_event_emission():
|
||||
) -> AsyncIterator[ChatResponseUpdate]:
|
||||
yield ChatResponseUpdate(contents=[Content.from_text(text="Hello")])
|
||||
|
||||
agent = ChatAgent(name="test_agent", instructions="Test", chat_client=StreamingChatClientStub(stream_fn))
|
||||
agent = ChatAgent(name="test_agent", instructions="Test", chat_client=streaming_chat_client_stub(stream_fn))
|
||||
wrapper = AgentFrameworkAgent(agent=agent)
|
||||
|
||||
input_data = {"messages": [{"role": "user", "content": "Hi"}]}
|
||||
@@ -255,7 +250,7 @@ async def test_run_finished_event_emission():
|
||||
assert events[-1].type == "RUN_FINISHED"
|
||||
|
||||
|
||||
async def test_tool_result_confirm_changes_accepted():
|
||||
async def test_tool_result_confirm_changes_accepted(streaming_chat_client_stub):
|
||||
"""Test confirm_changes tool result handling when accepted."""
|
||||
from agent_framework.ag_ui import AgentFrameworkAgent
|
||||
|
||||
@@ -264,7 +259,7 @@ async def test_tool_result_confirm_changes_accepted():
|
||||
) -> AsyncIterator[ChatResponseUpdate]:
|
||||
yield ChatResponseUpdate(contents=[Content.from_text(text="Document updated")])
|
||||
|
||||
agent = ChatAgent(name="test_agent", instructions="Test", chat_client=StreamingChatClientStub(stream_fn))
|
||||
agent = ChatAgent(name="test_agent", instructions="Test", chat_client=streaming_chat_client_stub(stream_fn))
|
||||
wrapper = AgentFrameworkAgent(
|
||||
agent=agent,
|
||||
state_schema={"document": {"type": "string"}},
|
||||
@@ -302,7 +297,7 @@ async def test_tool_result_confirm_changes_accepted():
|
||||
assert confirmation_found, f"No confirmation in deltas: {[e.delta for e in text_content_events]}"
|
||||
|
||||
|
||||
async def test_tool_result_confirm_changes_rejected():
|
||||
async def test_tool_result_confirm_changes_rejected(streaming_chat_client_stub):
|
||||
"""Test confirm_changes tool result handling when rejected."""
|
||||
from agent_framework.ag_ui import AgentFrameworkAgent
|
||||
|
||||
@@ -311,7 +306,7 @@ async def test_tool_result_confirm_changes_rejected():
|
||||
) -> AsyncIterator[ChatResponseUpdate]:
|
||||
yield ChatResponseUpdate(contents=[Content.from_text(text="OK")])
|
||||
|
||||
agent = ChatAgent(name="test_agent", instructions="Test", chat_client=StreamingChatClientStub(stream_fn))
|
||||
agent = ChatAgent(name="test_agent", instructions="Test", chat_client=streaming_chat_client_stub(stream_fn))
|
||||
wrapper = AgentFrameworkAgent(agent=agent)
|
||||
|
||||
# Simulate tool result message with rejection
|
||||
@@ -336,7 +331,7 @@ async def test_tool_result_confirm_changes_rejected():
|
||||
assert any("what would you like me to change" in e.delta.lower() for e in text_content_events)
|
||||
|
||||
|
||||
async def test_tool_result_function_approval_accepted():
|
||||
async def test_tool_result_function_approval_accepted(streaming_chat_client_stub):
|
||||
"""Test function approval tool result when steps are accepted."""
|
||||
from agent_framework.ag_ui import AgentFrameworkAgent
|
||||
|
||||
@@ -345,7 +340,7 @@ async def test_tool_result_function_approval_accepted():
|
||||
) -> AsyncIterator[ChatResponseUpdate]:
|
||||
yield ChatResponseUpdate(contents=[Content.from_text(text="OK")])
|
||||
|
||||
agent = ChatAgent(name="test_agent", instructions="Test", chat_client=StreamingChatClientStub(stream_fn))
|
||||
agent = ChatAgent(name="test_agent", instructions="Test", chat_client=streaming_chat_client_stub(stream_fn))
|
||||
wrapper = AgentFrameworkAgent(agent=agent)
|
||||
|
||||
# Simulate tool result with multiple steps
|
||||
@@ -382,7 +377,7 @@ async def test_tool_result_function_approval_accepted():
|
||||
assert "create calendar event" in full_text.lower()
|
||||
|
||||
|
||||
async def test_tool_result_function_approval_rejected():
|
||||
async def test_tool_result_function_approval_rejected(streaming_chat_client_stub):
|
||||
"""Test function approval tool result when rejected."""
|
||||
from agent_framework.ag_ui import AgentFrameworkAgent
|
||||
|
||||
@@ -391,7 +386,7 @@ async def test_tool_result_function_approval_rejected():
|
||||
) -> AsyncIterator[ChatResponseUpdate]:
|
||||
yield ChatResponseUpdate(contents=[Content.from_text(text="OK")])
|
||||
|
||||
agent = ChatAgent(name="test_agent", instructions="Test", chat_client=StreamingChatClientStub(stream_fn))
|
||||
agent = ChatAgent(name="test_agent", instructions="Test", chat_client=streaming_chat_client_stub(stream_fn))
|
||||
wrapper = AgentFrameworkAgent(agent=agent)
|
||||
|
||||
# Simulate tool result rejection with steps
|
||||
@@ -419,7 +414,7 @@ async def test_tool_result_function_approval_rejected():
|
||||
assert any("what would you like me to change about the plan" in e.delta.lower() for e in text_content_events)
|
||||
|
||||
|
||||
async def test_thread_metadata_tracking():
|
||||
async def test_thread_metadata_tracking(streaming_chat_client_stub):
|
||||
"""Test that thread metadata includes ag_ui_thread_id and ag_ui_run_id.
|
||||
|
||||
AG-UI internal metadata is stored in thread.metadata for orchestration,
|
||||
@@ -427,21 +422,16 @@ async def test_thread_metadata_tracking():
|
||||
"""
|
||||
from agent_framework.ag_ui import AgentFrameworkAgent
|
||||
|
||||
captured_thread: dict[str, Any] = {}
|
||||
captured_options: dict[str, Any] = {}
|
||||
|
||||
async def stream_fn(
|
||||
messages: MutableSequence[ChatMessage], options: dict[str, Any], **kwargs: Any
|
||||
) -> AsyncIterator[ChatResponseUpdate]:
|
||||
# Capture the thread object from kwargs
|
||||
thread = kwargs.get("thread")
|
||||
if thread and hasattr(thread, "metadata"):
|
||||
captured_thread["metadata"] = thread.metadata
|
||||
# Capture options to verify internal keys are NOT passed to chat client
|
||||
captured_options.update(options)
|
||||
yield ChatResponseUpdate(contents=[Content.from_text(text="Hello")])
|
||||
|
||||
agent = ChatAgent(name="test_agent", instructions="Test", chat_client=StreamingChatClientStub(stream_fn))
|
||||
agent = ChatAgent(name="test_agent", instructions="Test", chat_client=streaming_chat_client_stub(stream_fn))
|
||||
wrapper = AgentFrameworkAgent(agent=agent)
|
||||
|
||||
input_data = {
|
||||
@@ -455,7 +445,8 @@ async def test_thread_metadata_tracking():
|
||||
events.append(event)
|
||||
|
||||
# AG-UI internal metadata should be stored in thread.metadata
|
||||
thread_metadata = captured_thread.get("metadata", {})
|
||||
thread = agent.chat_client.last_thread
|
||||
thread_metadata = thread.metadata if thread and hasattr(thread, "metadata") else {}
|
||||
assert thread_metadata.get("ag_ui_thread_id") == "test_thread_123"
|
||||
assert thread_metadata.get("ag_ui_run_id") == "test_run_456"
|
||||
|
||||
@@ -465,7 +456,7 @@ async def test_thread_metadata_tracking():
|
||||
assert "ag_ui_run_id" not in options_metadata
|
||||
|
||||
|
||||
async def test_state_context_injection():
|
||||
async def test_state_context_injection(streaming_chat_client_stub):
|
||||
"""Test that current state is injected into thread metadata.
|
||||
|
||||
AG-UI internal metadata (including current_state) is stored in thread.metadata
|
||||
@@ -473,21 +464,16 @@ async def test_state_context_injection():
|
||||
"""
|
||||
from agent_framework_ag_ui import AgentFrameworkAgent
|
||||
|
||||
captured_thread: dict[str, Any] = {}
|
||||
captured_options: dict[str, Any] = {}
|
||||
|
||||
async def stream_fn(
|
||||
messages: MutableSequence[ChatMessage], options: dict[str, Any], **kwargs: Any
|
||||
) -> AsyncIterator[ChatResponseUpdate]:
|
||||
# Capture the thread object from kwargs
|
||||
thread = kwargs.get("thread")
|
||||
if thread and hasattr(thread, "metadata"):
|
||||
captured_thread["metadata"] = thread.metadata
|
||||
# Capture options to verify internal keys are NOT passed to chat client
|
||||
captured_options.update(options)
|
||||
yield ChatResponseUpdate(contents=[Content.from_text(text="Hello")])
|
||||
|
||||
agent = ChatAgent(name="test_agent", instructions="Test", chat_client=StreamingChatClientStub(stream_fn))
|
||||
agent = ChatAgent(name="test_agent", instructions="Test", chat_client=streaming_chat_client_stub(stream_fn))
|
||||
wrapper = AgentFrameworkAgent(
|
||||
agent=agent,
|
||||
state_schema={"document": {"type": "string"}},
|
||||
@@ -503,7 +489,8 @@ async def test_state_context_injection():
|
||||
events.append(event)
|
||||
|
||||
# Current state should be stored in thread.metadata
|
||||
thread_metadata = captured_thread.get("metadata", {})
|
||||
thread = agent.chat_client.last_thread
|
||||
thread_metadata = thread.metadata if thread and hasattr(thread, "metadata") else {}
|
||||
current_state = thread_metadata.get("current_state")
|
||||
if isinstance(current_state, str):
|
||||
current_state = json.loads(current_state)
|
||||
@@ -514,7 +501,7 @@ async def test_state_context_injection():
|
||||
assert "current_state" not in options_metadata
|
||||
|
||||
|
||||
async def test_no_messages_provided():
|
||||
async def test_no_messages_provided(streaming_chat_client_stub):
|
||||
"""Test handling when no messages are provided."""
|
||||
from agent_framework.ag_ui import AgentFrameworkAgent
|
||||
|
||||
@@ -523,7 +510,7 @@ async def test_no_messages_provided():
|
||||
) -> AsyncIterator[ChatResponseUpdate]:
|
||||
yield ChatResponseUpdate(contents=[Content.from_text(text="Hello")])
|
||||
|
||||
agent = ChatAgent(name="test_agent", instructions="Test", chat_client=StreamingChatClientStub(stream_fn))
|
||||
agent = ChatAgent(name="test_agent", instructions="Test", chat_client=streaming_chat_client_stub(stream_fn))
|
||||
wrapper = AgentFrameworkAgent(agent=agent)
|
||||
|
||||
input_data: dict[str, Any] = {"messages": []}
|
||||
@@ -538,7 +525,7 @@ async def test_no_messages_provided():
|
||||
assert events[-1].type == "RUN_FINISHED"
|
||||
|
||||
|
||||
async def test_message_end_event_emission():
|
||||
async def test_message_end_event_emission(streaming_chat_client_stub):
|
||||
"""Test TextMessageEndEvent is emitted for assistant messages."""
|
||||
from agent_framework.ag_ui import AgentFrameworkAgent
|
||||
|
||||
@@ -547,7 +534,7 @@ async def test_message_end_event_emission():
|
||||
) -> AsyncIterator[ChatResponseUpdate]:
|
||||
yield ChatResponseUpdate(contents=[Content.from_text(text="Hello world")])
|
||||
|
||||
agent = ChatAgent(name="test_agent", instructions="Test", chat_client=StreamingChatClientStub(stream_fn))
|
||||
agent = ChatAgent(name="test_agent", instructions="Test", chat_client=streaming_chat_client_stub(stream_fn))
|
||||
wrapper = AgentFrameworkAgent(agent=agent)
|
||||
|
||||
input_data: dict[str, Any] = {"messages": [{"role": "user", "content": "Hi"}]}
|
||||
@@ -566,7 +553,7 @@ async def test_message_end_event_emission():
|
||||
assert end_index < finished_index
|
||||
|
||||
|
||||
async def test_error_handling_with_exception():
|
||||
async def test_error_handling_with_exception(streaming_chat_client_stub):
|
||||
"""Test that exceptions during agent execution are re-raised."""
|
||||
from agent_framework.ag_ui import AgentFrameworkAgent
|
||||
|
||||
@@ -577,7 +564,7 @@ async def test_error_handling_with_exception():
|
||||
yield ChatResponseUpdate(contents=[])
|
||||
raise RuntimeError("Simulated failure")
|
||||
|
||||
agent = ChatAgent(name="test_agent", instructions="Test", chat_client=StreamingChatClientStub(stream_fn))
|
||||
agent = ChatAgent(name="test_agent", instructions="Test", chat_client=streaming_chat_client_stub(stream_fn))
|
||||
wrapper = AgentFrameworkAgent(agent=agent)
|
||||
|
||||
input_data: dict[str, Any] = {"messages": [{"role": "user", "content": "Hi"}]}
|
||||
@@ -587,7 +574,7 @@ async def test_error_handling_with_exception():
|
||||
pass
|
||||
|
||||
|
||||
async def test_json_decode_error_in_tool_result():
|
||||
async def test_json_decode_error_in_tool_result(streaming_chat_client_stub):
|
||||
"""Test handling of orphaned tool result - should be sanitized out."""
|
||||
from agent_framework.ag_ui import AgentFrameworkAgent
|
||||
|
||||
@@ -598,7 +585,7 @@ async def test_json_decode_error_in_tool_result():
|
||||
yield ChatResponseUpdate(contents=[])
|
||||
raise AssertionError("ChatClient should not be called with orphaned tool result")
|
||||
|
||||
agent = ChatAgent(name="test_agent", instructions="Test", chat_client=StreamingChatClientStub(stream_fn))
|
||||
agent = ChatAgent(name="test_agent", instructions="Test", chat_client=streaming_chat_client_stub(stream_fn))
|
||||
wrapper = AgentFrameworkAgent(agent=agent)
|
||||
|
||||
# Send invalid JSON as tool result without preceding tool call
|
||||
@@ -624,7 +611,7 @@ async def test_json_decode_error_in_tool_result():
|
||||
assert len(tool_events) == 0
|
||||
|
||||
|
||||
async def test_agent_with_use_service_thread_is_false():
|
||||
async def test_agent_with_use_service_thread_is_false(streaming_chat_client_stub):
|
||||
"""Test that when use_service_thread is False, the AgentThread used to run the agent is NOT set to the service thread ID."""
|
||||
from agent_framework.ag_ui import AgentFrameworkAgent
|
||||
|
||||
@@ -633,14 +620,11 @@ async def test_agent_with_use_service_thread_is_false():
|
||||
async def stream_fn(
|
||||
messages: MutableSequence[ChatMessage], chat_options: ChatOptions, **kwargs: Any
|
||||
) -> AsyncIterator[ChatResponseUpdate]:
|
||||
nonlocal request_service_thread_id
|
||||
thread = kwargs.get("thread")
|
||||
request_service_thread_id = thread.service_thread_id if thread else None
|
||||
yield ChatResponseUpdate(
|
||||
contents=[Content.from_text(text="Response")], response_id="resp_67890", conversation_id="conv_12345"
|
||||
)
|
||||
|
||||
agent = ChatAgent(chat_client=StreamingChatClientStub(stream_fn))
|
||||
agent = ChatAgent(chat_client=streaming_chat_client_stub(stream_fn))
|
||||
wrapper = AgentFrameworkAgent(agent=agent, use_service_thread=False)
|
||||
|
||||
input_data = {"messages": [{"role": "user", "content": "Hi"}], "thread_id": "conv_123456"}
|
||||
@@ -651,7 +635,7 @@ async def test_agent_with_use_service_thread_is_false():
|
||||
assert request_service_thread_id is None # type: ignore[attr-defined] (service_thread_id should be set)
|
||||
|
||||
|
||||
async def test_agent_with_use_service_thread_is_true():
|
||||
async def test_agent_with_use_service_thread_is_true(streaming_chat_client_stub):
|
||||
"""Test that when use_service_thread is True, the AgentThread used to run the agent is set to the service thread ID."""
|
||||
from agent_framework.ag_ui import AgentFrameworkAgent
|
||||
|
||||
@@ -667,7 +651,7 @@ async def test_agent_with_use_service_thread_is_true():
|
||||
contents=[Content.from_text(text="Response")], response_id="resp_67890", conversation_id="conv_12345"
|
||||
)
|
||||
|
||||
agent = ChatAgent(chat_client=StreamingChatClientStub(stream_fn))
|
||||
agent = ChatAgent(chat_client=streaming_chat_client_stub(stream_fn))
|
||||
wrapper = AgentFrameworkAgent(agent=agent, use_service_thread=True)
|
||||
|
||||
input_data = {"messages": [{"role": "user", "content": "Hi"}], "thread_id": "conv_123456"}
|
||||
@@ -675,10 +659,11 @@ async def test_agent_with_use_service_thread_is_true():
|
||||
events: list[Any] = []
|
||||
async for event in wrapper.run_agent(input_data):
|
||||
events.append(event)
|
||||
request_service_thread_id = agent.chat_client.last_service_thread_id
|
||||
assert request_service_thread_id == "conv_123456" # type: ignore[attr-defined] (service_thread_id should be set)
|
||||
|
||||
|
||||
async def test_function_approval_mode_executes_tool():
|
||||
async def test_function_approval_mode_executes_tool(streaming_chat_client_stub):
|
||||
"""Test that function approval with approval_mode='always_require' sends the correct messages."""
|
||||
from agent_framework import tool
|
||||
from agent_framework.ag_ui import AgentFrameworkAgent
|
||||
@@ -702,7 +687,7 @@ async def test_function_approval_mode_executes_tool():
|
||||
yield ChatResponseUpdate(contents=[Content.from_text(text="Processing completed")])
|
||||
|
||||
agent = ChatAgent(
|
||||
chat_client=StreamingChatClientStub(stream_fn),
|
||||
chat_client=streaming_chat_client_stub(stream_fn),
|
||||
name="test_agent",
|
||||
instructions="Test",
|
||||
tools=[get_datetime],
|
||||
@@ -769,7 +754,7 @@ async def test_function_approval_mode_executes_tool():
|
||||
)
|
||||
|
||||
|
||||
async def test_function_approval_mode_rejection():
|
||||
async def test_function_approval_mode_rejection(streaming_chat_client_stub):
|
||||
"""Test that function approval rejection creates a rejection response."""
|
||||
from agent_framework import tool
|
||||
from agent_framework.ag_ui import AgentFrameworkAgent
|
||||
@@ -795,7 +780,7 @@ async def test_function_approval_mode_rejection():
|
||||
agent = ChatAgent(
|
||||
name="test_agent",
|
||||
instructions="Test",
|
||||
chat_client=StreamingChatClientStub(stream_fn),
|
||||
chat_client=streaming_chat_client_stub(stream_fn),
|
||||
tools=[delete_all_data],
|
||||
)
|
||||
wrapper = AgentFrameworkAgent(agent=agent)
|
||||
+31
-30
@@ -3,9 +3,8 @@
|
||||
"""Tests for FastAPI endpoint creation (_endpoint.py)."""
|
||||
|
||||
import json
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
from agent_framework import ChatAgent, ChatResponseUpdate, Content
|
||||
from fastapi import FastAPI, Header, HTTPException
|
||||
from fastapi.params import Depends
|
||||
@@ -14,17 +13,19 @@ from fastapi.testclient import TestClient
|
||||
from agent_framework_ag_ui import add_agent_framework_fastapi_endpoint
|
||||
from agent_framework_ag_ui._agent import AgentFrameworkAgent
|
||||
|
||||
sys.path.insert(0, str(Path(__file__).parent))
|
||||
from utils_test_ag_ui import StreamingChatClientStub, stream_from_updates
|
||||
|
||||
|
||||
def build_chat_client(response_text: str = "Test response") -> StreamingChatClientStub:
|
||||
@pytest.fixture
|
||||
def build_chat_client(streaming_chat_client_stub, stream_from_updates_fixture):
|
||||
"""Create a typed chat client stub for endpoint tests."""
|
||||
updates = [ChatResponseUpdate(contents=[Content.from_text(text=response_text)])]
|
||||
return StreamingChatClientStub(stream_from_updates(updates))
|
||||
|
||||
def _build(response_text: str = "Test response"):
|
||||
updates = [ChatResponseUpdate(contents=[Content.from_text(text=response_text)])]
|
||||
return streaming_chat_client_stub(stream_from_updates_fixture(updates))
|
||||
|
||||
return _build
|
||||
|
||||
|
||||
async def test_add_endpoint_with_agent_protocol():
|
||||
async def test_add_endpoint_with_agent_protocol(build_chat_client):
|
||||
"""Test adding endpoint with raw AgentProtocol."""
|
||||
app = FastAPI()
|
||||
agent = ChatAgent(name="test", instructions="Test agent", chat_client=build_chat_client())
|
||||
@@ -38,7 +39,7 @@ async def test_add_endpoint_with_agent_protocol():
|
||||
assert response.headers["content-type"] == "text/event-stream; charset=utf-8"
|
||||
|
||||
|
||||
async def test_add_endpoint_with_wrapped_agent():
|
||||
async def test_add_endpoint_with_wrapped_agent(build_chat_client):
|
||||
"""Test adding endpoint with pre-wrapped AgentFrameworkAgent."""
|
||||
app = FastAPI()
|
||||
agent = ChatAgent(name="test", instructions="Test agent", chat_client=build_chat_client())
|
||||
@@ -53,7 +54,7 @@ async def test_add_endpoint_with_wrapped_agent():
|
||||
assert response.headers["content-type"] == "text/event-stream; charset=utf-8"
|
||||
|
||||
|
||||
async def test_endpoint_with_state_schema():
|
||||
async def test_endpoint_with_state_schema(build_chat_client):
|
||||
"""Test endpoint with state_schema parameter."""
|
||||
app = FastAPI()
|
||||
agent = ChatAgent(name="test", instructions="Test agent", chat_client=build_chat_client())
|
||||
@@ -69,7 +70,7 @@ async def test_endpoint_with_state_schema():
|
||||
assert response.status_code == 200
|
||||
|
||||
|
||||
async def test_endpoint_with_default_state_seed():
|
||||
async def test_endpoint_with_default_state_seed(build_chat_client):
|
||||
"""Test endpoint seeds default state when client omits it."""
|
||||
app = FastAPI()
|
||||
agent = ChatAgent(name="test", instructions="Test agent", chat_client=build_chat_client())
|
||||
@@ -96,7 +97,7 @@ async def test_endpoint_with_default_state_seed():
|
||||
assert snapshots[0]["snapshot"]["proverbs"] == default_state["proverbs"]
|
||||
|
||||
|
||||
async def test_endpoint_with_predict_state_config():
|
||||
async def test_endpoint_with_predict_state_config(build_chat_client):
|
||||
"""Test endpoint with predict_state_config parameter."""
|
||||
app = FastAPI()
|
||||
agent = ChatAgent(name="test", instructions="Test agent", chat_client=build_chat_client())
|
||||
@@ -110,7 +111,7 @@ async def test_endpoint_with_predict_state_config():
|
||||
assert response.status_code == 200
|
||||
|
||||
|
||||
async def test_endpoint_request_logging():
|
||||
async def test_endpoint_request_logging(build_chat_client):
|
||||
"""Test that endpoint logs request details."""
|
||||
app = FastAPI()
|
||||
agent = ChatAgent(name="test", instructions="Test agent", chat_client=build_chat_client())
|
||||
@@ -130,7 +131,7 @@ async def test_endpoint_request_logging():
|
||||
assert response.status_code == 200
|
||||
|
||||
|
||||
async def test_endpoint_event_streaming():
|
||||
async def test_endpoint_event_streaming(build_chat_client):
|
||||
"""Test that endpoint streams events correctly."""
|
||||
app = FastAPI()
|
||||
agent = ChatAgent(name="test", instructions="Test agent", chat_client=build_chat_client("Streamed response"))
|
||||
@@ -164,7 +165,7 @@ async def test_endpoint_event_streaming():
|
||||
assert found_run_finished
|
||||
|
||||
|
||||
async def test_endpoint_error_handling():
|
||||
async def test_endpoint_error_handling(build_chat_client):
|
||||
"""Test endpoint error handling during request parsing."""
|
||||
app = FastAPI()
|
||||
agent = ChatAgent(name="test", instructions="Test agent", chat_client=build_chat_client())
|
||||
@@ -180,7 +181,7 @@ async def test_endpoint_error_handling():
|
||||
assert response.status_code == 422
|
||||
|
||||
|
||||
async def test_endpoint_multiple_paths():
|
||||
async def test_endpoint_multiple_paths(build_chat_client):
|
||||
"""Test adding multiple endpoints with different paths."""
|
||||
app = FastAPI()
|
||||
agent1 = ChatAgent(name="agent1", instructions="First agent", chat_client=build_chat_client("Response 1"))
|
||||
@@ -198,7 +199,7 @@ async def test_endpoint_multiple_paths():
|
||||
assert response2.status_code == 200
|
||||
|
||||
|
||||
async def test_endpoint_default_path():
|
||||
async def test_endpoint_default_path(build_chat_client):
|
||||
"""Test endpoint with default path."""
|
||||
app = FastAPI()
|
||||
agent = ChatAgent(name="test", instructions="Test agent", chat_client=build_chat_client())
|
||||
@@ -211,7 +212,7 @@ async def test_endpoint_default_path():
|
||||
assert response.status_code == 200
|
||||
|
||||
|
||||
async def test_endpoint_response_headers():
|
||||
async def test_endpoint_response_headers(build_chat_client):
|
||||
"""Test that endpoint sets correct response headers."""
|
||||
app = FastAPI()
|
||||
agent = ChatAgent(name="test", instructions="Test agent", chat_client=build_chat_client())
|
||||
@@ -227,7 +228,7 @@ async def test_endpoint_response_headers():
|
||||
assert response.headers["cache-control"] == "no-cache"
|
||||
|
||||
|
||||
async def test_endpoint_empty_messages():
|
||||
async def test_endpoint_empty_messages(build_chat_client):
|
||||
"""Test endpoint with empty messages list."""
|
||||
app = FastAPI()
|
||||
agent = ChatAgent(name="test", instructions="Test agent", chat_client=build_chat_client())
|
||||
@@ -240,7 +241,7 @@ async def test_endpoint_empty_messages():
|
||||
assert response.status_code == 200
|
||||
|
||||
|
||||
async def test_endpoint_complex_input():
|
||||
async def test_endpoint_complex_input(build_chat_client):
|
||||
"""Test endpoint with complex input data."""
|
||||
app = FastAPI()
|
||||
agent = ChatAgent(name="test", instructions="Test agent", chat_client=build_chat_client())
|
||||
@@ -265,7 +266,7 @@ async def test_endpoint_complex_input():
|
||||
assert response.status_code == 200
|
||||
|
||||
|
||||
async def test_endpoint_openapi_schema():
|
||||
async def test_endpoint_openapi_schema(build_chat_client):
|
||||
"""Test that endpoint generates proper OpenAPI schema with request model."""
|
||||
app = FastAPI()
|
||||
agent = ChatAgent(name="test", instructions="Test agent", chat_client=build_chat_client())
|
||||
@@ -309,7 +310,7 @@ async def test_endpoint_openapi_schema():
|
||||
assert "messages" in agui_request_schema["required"]
|
||||
|
||||
|
||||
async def test_endpoint_default_tags():
|
||||
async def test_endpoint_default_tags(build_chat_client):
|
||||
"""Test that endpoint uses default 'AG-UI' tag."""
|
||||
app = FastAPI()
|
||||
agent = ChatAgent(name="test", instructions="Test agent", chat_client=build_chat_client())
|
||||
@@ -327,7 +328,7 @@ async def test_endpoint_default_tags():
|
||||
assert endpoint_spec["tags"] == ["AG-UI"]
|
||||
|
||||
|
||||
async def test_endpoint_custom_tags():
|
||||
async def test_endpoint_custom_tags(build_chat_client):
|
||||
"""Test that endpoint accepts custom tags."""
|
||||
app = FastAPI()
|
||||
agent = ChatAgent(name="test", instructions="Test agent", chat_client=build_chat_client())
|
||||
@@ -345,7 +346,7 @@ async def test_endpoint_custom_tags():
|
||||
assert endpoint_spec["tags"] == ["Custom", "Agent"]
|
||||
|
||||
|
||||
async def test_endpoint_missing_required_field():
|
||||
async def test_endpoint_missing_required_field(build_chat_client):
|
||||
"""Test that endpoint validates required fields with Pydantic."""
|
||||
app = FastAPI()
|
||||
agent = ChatAgent(name="test", instructions="Test agent", chat_client=build_chat_client())
|
||||
@@ -362,7 +363,7 @@ async def test_endpoint_missing_required_field():
|
||||
assert "detail" in error_detail
|
||||
|
||||
|
||||
async def test_endpoint_internal_error_handling():
|
||||
async def test_endpoint_internal_error_handling(build_chat_client):
|
||||
"""Test endpoint error handling when an exception occurs before streaming starts."""
|
||||
from unittest.mock import patch
|
||||
|
||||
@@ -383,7 +384,7 @@ async def test_endpoint_internal_error_handling():
|
||||
assert response.json() == {"error": "An internal error has occurred."}
|
||||
|
||||
|
||||
async def test_endpoint_with_dependencies_blocks_unauthorized():
|
||||
async def test_endpoint_with_dependencies_blocks_unauthorized(build_chat_client):
|
||||
"""Test that endpoint blocks requests when authentication dependency fails."""
|
||||
app = FastAPI()
|
||||
agent = ChatAgent(name="test", instructions="Test agent", chat_client=build_chat_client())
|
||||
@@ -402,7 +403,7 @@ async def test_endpoint_with_dependencies_blocks_unauthorized():
|
||||
assert response.json()["detail"] == "Unauthorized"
|
||||
|
||||
|
||||
async def test_endpoint_with_dependencies_allows_authorized():
|
||||
async def test_endpoint_with_dependencies_allows_authorized(build_chat_client):
|
||||
"""Test that endpoint allows requests when authentication dependency passes."""
|
||||
app = FastAPI()
|
||||
agent = ChatAgent(name="test", instructions="Test agent", chat_client=build_chat_client())
|
||||
@@ -425,7 +426,7 @@ async def test_endpoint_with_dependencies_allows_authorized():
|
||||
assert response.headers["content-type"] == "text/event-stream; charset=utf-8"
|
||||
|
||||
|
||||
async def test_endpoint_with_multiple_dependencies():
|
||||
async def test_endpoint_with_multiple_dependencies(build_chat_client):
|
||||
"""Test that endpoint supports multiple dependencies."""
|
||||
app = FastAPI()
|
||||
agent = ChatAgent(name="test", instructions="Test agent", chat_client=build_chat_client())
|
||||
@@ -453,7 +454,7 @@ async def test_endpoint_with_multiple_dependencies():
|
||||
assert "second" in execution_order
|
||||
|
||||
|
||||
async def test_endpoint_without_dependencies_is_accessible():
|
||||
async def test_endpoint_without_dependencies_is_accessible(build_chat_client):
|
||||
"""Test that endpoint without dependencies remains accessible (backward compatibility)."""
|
||||
app = FastAPI()
|
||||
agent = ChatAgent(name="test", instructions="Test agent", chat_client=build_chat_client())
|
||||
+5
-5
@@ -29,8 +29,8 @@ class TestPendingToolCallIds:
|
||||
def test_no_tool_calls(self):
|
||||
"""Returns empty set when no tool calls in messages."""
|
||||
messages = [
|
||||
ChatMessage("user", [Content.from_text("Hello")]),
|
||||
ChatMessage("assistant", [Content.from_text("Hi there")]),
|
||||
ChatMessage(role="user", contents=[Content.from_text("Hello")]),
|
||||
ChatMessage(role="assistant", contents=[Content.from_text("Hi there")]),
|
||||
]
|
||||
result = pending_tool_call_ids(messages)
|
||||
assert result == set()
|
||||
@@ -114,7 +114,7 @@ class TestIsStateContextMessage:
|
||||
|
||||
def test_empty_contents(self):
|
||||
"""Returns False for message with empty contents."""
|
||||
message = ChatMessage("system", [])
|
||||
message = ChatMessage(role="system", contents=[])
|
||||
assert is_state_context_message(message) is False
|
||||
|
||||
|
||||
@@ -342,7 +342,7 @@ class TestLatestApprovalResponse:
|
||||
def test_no_approval_response(self):
|
||||
"""Returns None when no approval response in last message."""
|
||||
messages = [
|
||||
ChatMessage("assistant", [Content.from_text("Hello")]),
|
||||
ChatMessage(role="assistant", contents=[Content.from_text("Hello")]),
|
||||
]
|
||||
result = latest_approval_response(messages)
|
||||
assert result is None
|
||||
@@ -357,7 +357,7 @@ class TestLatestApprovalResponse:
|
||||
function_call=fc,
|
||||
)
|
||||
messages = [
|
||||
ChatMessage("user", [approval_content]),
|
||||
ChatMessage(role="user", contents=[approval_content]),
|
||||
]
|
||||
result = latest_approval_response(messages)
|
||||
assert result is approval_content
|
||||
+3
-3
@@ -24,7 +24,7 @@ def sample_agui_message():
|
||||
@pytest.fixture
|
||||
def sample_agent_framework_message():
|
||||
"""Create a sample Agent Framework message."""
|
||||
return ChatMessage("user", [Content.from_text(text="Hello")], message_id="msg-123")
|
||||
return ChatMessage(role="user", contents=[Content.from_text(text="Hello")], message_id="msg-123")
|
||||
|
||||
|
||||
def test_agui_to_agent_framework_basic(sample_agui_message):
|
||||
@@ -484,7 +484,7 @@ def test_agent_framework_to_agui_multiple_text_contents():
|
||||
|
||||
def test_agent_framework_to_agui_no_message_id():
|
||||
"""Test message without message_id - should auto-generate ID."""
|
||||
msg = ChatMessage("user", [Content.from_text(text="Hello")])
|
||||
msg = ChatMessage(role="user", contents=[Content.from_text(text="Hello")])
|
||||
|
||||
messages = agent_framework_messages_to_agui([msg])
|
||||
|
||||
@@ -496,7 +496,7 @@ def test_agent_framework_to_agui_no_message_id():
|
||||
|
||||
def test_agent_framework_to_agui_system_role():
|
||||
"""Test system role conversion."""
|
||||
msg = ChatMessage("system", [Content.from_text(text="System")])
|
||||
msg = ChatMessage(role="system", contents=[Content.from_text(text="System")])
|
||||
|
||||
messages = agent_framework_messages_to_agui([msg])
|
||||
|
||||
+6
-12
@@ -33,14 +33,12 @@ def test_sanitize_tool_history_filters_out_confirm_changes_only_message() -> Non
|
||||
|
||||
# Assistant message with only confirm_changes should be filtered out
|
||||
assistant_messages = [
|
||||
msg for msg in sanitized if (msg.role.value if hasattr(msg.role, "value") else str(msg.role)) == "assistant"
|
||||
msg for msg in sanitized if (msg.role if hasattr(msg.role, "value") else str(msg.role)) == "assistant"
|
||||
]
|
||||
assert len(assistant_messages) == 0
|
||||
|
||||
# No synthetic tool result should be injected since confirm_changes was filtered out
|
||||
tool_messages = [
|
||||
msg for msg in sanitized if (msg.role.value if hasattr(msg.role, "value") else str(msg.role)) == "tool"
|
||||
]
|
||||
tool_messages = [msg for msg in sanitized if (msg.role if hasattr(msg.role, "value") else str(msg.role)) == "tool"]
|
||||
assert len(tool_messages) == 0
|
||||
|
||||
|
||||
@@ -182,7 +180,7 @@ def test_sanitize_tool_history_filters_confirm_changes_keeps_other_tools() -> No
|
||||
|
||||
# Find the assistant message
|
||||
assistant_messages = [
|
||||
msg for msg in sanitized if (msg.role.value if hasattr(msg.role, "value") else str(msg.role)) == "assistant"
|
||||
msg for msg in sanitized if (msg.role if hasattr(msg.role, "value") else str(msg.role)) == "assistant"
|
||||
]
|
||||
assert len(assistant_messages) == 1
|
||||
|
||||
@@ -192,9 +190,7 @@ def test_sanitize_tool_history_filters_confirm_changes_keeps_other_tools() -> No
|
||||
assert "confirm_changes" not in function_call_names
|
||||
|
||||
# Only one tool message (for call_1), no synthetic for confirm_changes
|
||||
tool_messages = [
|
||||
msg for msg in sanitized if (msg.role.value if hasattr(msg.role, "value") else str(msg.role)) == "tool"
|
||||
]
|
||||
tool_messages = [msg for msg in sanitized if (msg.role if hasattr(msg.role, "value") else str(msg.role)) == "tool"]
|
||||
assert len(tool_messages) == 1
|
||||
assert str(tool_messages[0].contents[0].call_id) == "call_1"
|
||||
|
||||
@@ -249,7 +245,7 @@ def test_sanitize_tool_history_filters_confirm_changes_from_assistant_messages()
|
||||
|
||||
# Find the assistant message in sanitized output
|
||||
assistant_messages = [
|
||||
msg for msg in sanitized if (msg.role.value if hasattr(msg.role, "value") else str(msg.role)) == "assistant"
|
||||
msg for msg in sanitized if (msg.role if hasattr(msg.role, "value") else str(msg.role)) == "assistant"
|
||||
]
|
||||
|
||||
assert len(assistant_messages) == 1
|
||||
@@ -261,9 +257,7 @@ def test_sanitize_tool_history_filters_confirm_changes_from_assistant_messages()
|
||||
assert "confirm_changes" not in function_call_names
|
||||
|
||||
# No synthetic tool result for confirm_changes (it was filtered from the message)
|
||||
tool_messages = [
|
||||
msg for msg in sanitized if (msg.role.value if hasattr(msg.role, "value") else str(msg.role)) == "tool"
|
||||
]
|
||||
tool_messages = [msg for msg in sanitized if (msg.role if hasattr(msg.role, "value") else str(msg.role)) == "tool"]
|
||||
# No tool results expected since there are no completed tool calls
|
||||
# (the approval response is handled separately by the framework)
|
||||
tool_call_ids = {str(msg.contents[0].call_id) for msg in tool_messages}
|
||||
+7
-7
@@ -212,7 +212,7 @@ class TestInjectStateContext:
|
||||
|
||||
def test_no_state_message(self):
|
||||
"""Returns original messages when no state context needed."""
|
||||
messages = [ChatMessage("user", [Content.from_text("Hello")])]
|
||||
messages = [ChatMessage(role="user", contents=[Content.from_text("Hello")])]
|
||||
result = _inject_state_context(messages, {}, {})
|
||||
assert result == messages
|
||||
|
||||
@@ -224,8 +224,8 @@ class TestInjectStateContext:
|
||||
def test_last_message_not_user(self):
|
||||
"""Returns original messages when last message is not from user."""
|
||||
messages = [
|
||||
ChatMessage("user", [Content.from_text("Hello")]),
|
||||
ChatMessage("assistant", [Content.from_text("Hi")]),
|
||||
ChatMessage(role="user", contents=[Content.from_text("Hello")]),
|
||||
ChatMessage(role="assistant", contents=[Content.from_text("Hi")]),
|
||||
]
|
||||
state = {"key": "value"}
|
||||
schema = {"properties": {"key": {"type": "string"}}}
|
||||
@@ -237,8 +237,8 @@ class TestInjectStateContext:
|
||||
"""Injects state context before last user message."""
|
||||
|
||||
messages = [
|
||||
ChatMessage("system", [Content.from_text("You are helpful")]),
|
||||
ChatMessage("user", [Content.from_text("Hello")]),
|
||||
ChatMessage(role="system", contents=[Content.from_text("You are helpful")]),
|
||||
ChatMessage(role="user", contents=[Content.from_text("Hello")]),
|
||||
]
|
||||
state = {"document": "content"}
|
||||
schema = {"properties": {"document": {"type": "string"}}}
|
||||
@@ -405,7 +405,7 @@ def test_extract_approved_state_updates_no_handler():
|
||||
"""Test _extract_approved_state_updates returns empty with no handler."""
|
||||
from agent_framework_ag_ui._run import _extract_approved_state_updates
|
||||
|
||||
messages = [ChatMessage("user", [Content.from_text("Hello")])]
|
||||
messages = [ChatMessage(role="user", contents=[Content.from_text("Hello")])]
|
||||
result = _extract_approved_state_updates(messages, None)
|
||||
assert result == {}
|
||||
|
||||
@@ -416,7 +416,7 @@ def test_extract_approved_state_updates_no_approval():
|
||||
from agent_framework_ag_ui._run import _extract_approved_state_updates
|
||||
|
||||
handler = PredictiveStateHandler(predict_state_config={"doc": {"tool": "write", "tool_argument": "content"}})
|
||||
messages = [ChatMessage("user", [Content.from_text("Hello")])]
|
||||
messages = [ChatMessage(role="user", contents=[Content.from_text("Hello")])]
|
||||
result = _extract_approved_state_updates(messages, handler)
|
||||
assert result == {}
|
||||
|
||||
+6
-11
@@ -2,19 +2,14 @@
|
||||
|
||||
"""Tests for service-managed thread IDs, and service-generated response ids."""
|
||||
|
||||
import sys
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
from ag_ui.core import RunFinishedEvent, RunStartedEvent
|
||||
from agent_framework import Content
|
||||
from agent_framework._types import AgentResponseUpdate, ChatResponseUpdate
|
||||
|
||||
sys.path.insert(0, str(Path(__file__).parent))
|
||||
from utils_test_ag_ui import StubAgent
|
||||
|
||||
|
||||
async def test_service_thread_id_when_there_are_updates():
|
||||
async def test_service_thread_id_when_there_are_updates(stub_agent):
|
||||
"""Test that service-managed thread IDs (conversation_id) are correctly set as the thread_id in events."""
|
||||
from agent_framework.ag_ui import AgentFrameworkAgent
|
||||
|
||||
@@ -29,7 +24,7 @@ async def test_service_thread_id_when_there_are_updates():
|
||||
),
|
||||
)
|
||||
]
|
||||
agent = StubAgent(updates=updates)
|
||||
agent = stub_agent(updates=updates)
|
||||
wrapper = AgentFrameworkAgent(agent=agent)
|
||||
|
||||
input_data = {
|
||||
@@ -46,12 +41,12 @@ async def test_service_thread_id_when_there_are_updates():
|
||||
assert isinstance(events[-1], RunFinishedEvent)
|
||||
|
||||
|
||||
async def test_service_thread_id_when_no_user_message():
|
||||
async def test_service_thread_id_when_no_user_message(stub_agent):
|
||||
"""Test when user submits no messages, emitted events still have with a thread_id"""
|
||||
from agent_framework.ag_ui import AgentFrameworkAgent
|
||||
|
||||
updates: list[AgentResponseUpdate] = []
|
||||
agent = StubAgent(updates=updates)
|
||||
agent = stub_agent(updates=updates)
|
||||
wrapper = AgentFrameworkAgent(agent=agent)
|
||||
|
||||
input_data: dict[str, list[dict[str, str]]] = {
|
||||
@@ -68,12 +63,12 @@ async def test_service_thread_id_when_no_user_message():
|
||||
assert isinstance(events[-1], RunFinishedEvent)
|
||||
|
||||
|
||||
async def test_service_thread_id_when_user_supplied_thread_id():
|
||||
async def test_service_thread_id_when_user_supplied_thread_id(stub_agent):
|
||||
"""Test that user-supplied thread IDs are preserved in emitted events."""
|
||||
from agent_framework.ag_ui import AgentFrameworkAgent
|
||||
|
||||
updates: list[AgentResponseUpdate] = []
|
||||
agent = StubAgent(updates=updates)
|
||||
agent = stub_agent(updates=updates)
|
||||
wrapper = AgentFrameworkAgent(agent=agent)
|
||||
|
||||
input_data: dict[str, Any] = {"messages": [{"role": "user", "content": "Hi"}], "threadId": "conv_12345"}
|
||||
+14
-19
@@ -3,17 +3,12 @@
|
||||
"""Tests for structured output handling in _agent.py."""
|
||||
|
||||
import json
|
||||
import sys
|
||||
from collections.abc import AsyncIterator, MutableSequence
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
from agent_framework import ChatAgent, ChatMessage, ChatOptions, ChatResponseUpdate, Content
|
||||
from pydantic import BaseModel
|
||||
|
||||
sys.path.insert(0, str(Path(__file__).parent))
|
||||
from utils_test_ag_ui import StreamingChatClientStub, stream_from_updates
|
||||
|
||||
|
||||
class RecipeOutput(BaseModel):
|
||||
"""Test Pydantic model for recipe output."""
|
||||
@@ -35,7 +30,7 @@ class GenericOutput(BaseModel):
|
||||
data: dict[str, Any]
|
||||
|
||||
|
||||
async def test_structured_output_with_recipe():
|
||||
async def test_structured_output_with_recipe(streaming_chat_client_stub, stream_from_updates_fixture):
|
||||
"""Test structured output processing with recipe state."""
|
||||
from agent_framework.ag_ui import AgentFrameworkAgent
|
||||
|
||||
@@ -46,7 +41,7 @@ async def test_structured_output_with_recipe():
|
||||
contents=[Content.from_text(text='{"recipe": {"name": "Pasta"}, "message": "Here is your recipe"}')]
|
||||
)
|
||||
|
||||
agent = ChatAgent(name="test", instructions="Test", chat_client=StreamingChatClientStub(stream_fn))
|
||||
agent = ChatAgent(name="test", instructions="Test", chat_client=streaming_chat_client_stub(stream_fn))
|
||||
agent.default_options = ChatOptions(response_format=RecipeOutput)
|
||||
|
||||
wrapper = AgentFrameworkAgent(
|
||||
@@ -73,7 +68,7 @@ async def test_structured_output_with_recipe():
|
||||
assert any("Here is your recipe" in e.delta for e in text_events)
|
||||
|
||||
|
||||
async def test_structured_output_with_steps():
|
||||
async def test_structured_output_with_steps(streaming_chat_client_stub, stream_from_updates_fixture):
|
||||
"""Test structured output processing with steps state."""
|
||||
from agent_framework.ag_ui import AgentFrameworkAgent
|
||||
|
||||
@@ -88,7 +83,7 @@ async def test_structured_output_with_steps():
|
||||
}
|
||||
yield ChatResponseUpdate(contents=[Content.from_text(text=json.dumps(steps_data))])
|
||||
|
||||
agent = ChatAgent(name="test", instructions="Test", chat_client=StreamingChatClientStub(stream_fn))
|
||||
agent = ChatAgent(name="test", instructions="Test", chat_client=streaming_chat_client_stub(stream_fn))
|
||||
agent.default_options = ChatOptions(response_format=StepsOutput)
|
||||
|
||||
wrapper = AgentFrameworkAgent(
|
||||
@@ -113,7 +108,7 @@ async def test_structured_output_with_steps():
|
||||
assert steps_snapshots[0].snapshot["steps"][0]["id"] == "1"
|
||||
|
||||
|
||||
async def test_structured_output_with_no_schema_match():
|
||||
async def test_structured_output_with_no_schema_match(streaming_chat_client_stub, stream_from_updates_fixture):
|
||||
"""Test structured output when response fields don't match state_schema keys."""
|
||||
from agent_framework.ag_ui import AgentFrameworkAgent
|
||||
|
||||
@@ -122,7 +117,7 @@ async def test_structured_output_with_no_schema_match():
|
||||
]
|
||||
|
||||
agent = ChatAgent(
|
||||
name="test", instructions="Test", chat_client=StreamingChatClientStub(stream_from_updates(updates))
|
||||
name="test", instructions="Test", chat_client=streaming_chat_client_stub(stream_from_updates_fixture(updates))
|
||||
)
|
||||
agent.default_options = ChatOptions(response_format=GenericOutput)
|
||||
|
||||
@@ -143,7 +138,7 @@ async def test_structured_output_with_no_schema_match():
|
||||
assert len(snapshot_events) >= 1
|
||||
|
||||
|
||||
async def test_structured_output_without_schema():
|
||||
async def test_structured_output_without_schema(streaming_chat_client_stub, stream_from_updates_fixture):
|
||||
"""Test structured output without state_schema treats all fields as state."""
|
||||
from agent_framework.ag_ui import AgentFrameworkAgent
|
||||
|
||||
@@ -158,7 +153,7 @@ async def test_structured_output_without_schema():
|
||||
) -> AsyncIterator[ChatResponseUpdate]:
|
||||
yield ChatResponseUpdate(contents=[Content.from_text(text='{"data": {"key": "value"}, "info": "processed"}')])
|
||||
|
||||
agent = ChatAgent(name="test", instructions="Test", chat_client=StreamingChatClientStub(stream_fn))
|
||||
agent = ChatAgent(name="test", instructions="Test", chat_client=streaming_chat_client_stub(stream_fn))
|
||||
agent.default_options = ChatOptions(response_format=DataOutput)
|
||||
|
||||
wrapper = AgentFrameworkAgent(
|
||||
@@ -181,7 +176,7 @@ async def test_structured_output_without_schema():
|
||||
assert snapshot_events[0].snapshot["info"] == "processed"
|
||||
|
||||
|
||||
async def test_no_structured_output_when_no_response_format():
|
||||
async def test_no_structured_output_when_no_response_format(streaming_chat_client_stub, stream_from_updates_fixture):
|
||||
"""Test that structured output path is skipped when no response_format."""
|
||||
from agent_framework.ag_ui import AgentFrameworkAgent
|
||||
|
||||
@@ -190,7 +185,7 @@ async def test_no_structured_output_when_no_response_format():
|
||||
agent = ChatAgent(
|
||||
name="test",
|
||||
instructions="Test",
|
||||
chat_client=StreamingChatClientStub(stream_from_updates(updates)),
|
||||
chat_client=streaming_chat_client_stub(stream_from_updates_fixture(updates)),
|
||||
)
|
||||
# No response_format set
|
||||
|
||||
@@ -208,7 +203,7 @@ async def test_no_structured_output_when_no_response_format():
|
||||
assert text_events[0].delta == "Regular text"
|
||||
|
||||
|
||||
async def test_structured_output_with_message_field():
|
||||
async def test_structured_output_with_message_field(streaming_chat_client_stub, stream_from_updates_fixture):
|
||||
"""Test structured output that includes a message field."""
|
||||
from agent_framework.ag_ui import AgentFrameworkAgent
|
||||
|
||||
@@ -218,7 +213,7 @@ async def test_structured_output_with_message_field():
|
||||
output_data = {"recipe": {"name": "Salad"}, "message": "Fresh salad recipe ready"}
|
||||
yield ChatResponseUpdate(contents=[Content.from_text(text=json.dumps(output_data))])
|
||||
|
||||
agent = ChatAgent(name="test", instructions="Test", chat_client=StreamingChatClientStub(stream_fn))
|
||||
agent = ChatAgent(name="test", instructions="Test", chat_client=streaming_chat_client_stub(stream_fn))
|
||||
agent.default_options = ChatOptions(response_format=RecipeOutput)
|
||||
|
||||
wrapper = AgentFrameworkAgent(
|
||||
@@ -243,7 +238,7 @@ async def test_structured_output_with_message_field():
|
||||
assert len(end_events) >= 1
|
||||
|
||||
|
||||
async def test_empty_updates_no_structured_processing():
|
||||
async def test_empty_updates_no_structured_processing(streaming_chat_client_stub, stream_from_updates_fixture):
|
||||
"""Test that empty updates don't trigger structured output processing."""
|
||||
from agent_framework.ag_ui import AgentFrameworkAgent
|
||||
|
||||
@@ -253,7 +248,7 @@ async def test_empty_updates_no_structured_processing():
|
||||
if False:
|
||||
yield ChatResponseUpdate(contents=[])
|
||||
|
||||
agent = ChatAgent(name="test", instructions="Test", chat_client=StreamingChatClientStub(stream_fn))
|
||||
agent = ChatAgent(name="test", instructions="Test", chat_client=streaming_chat_client_stub(stream_fn))
|
||||
agent.default_options = ChatOptions(response_format=RecipeOutput)
|
||||
|
||||
wrapper = AgentFrameworkAgent(agent=agent)
|
||||
+3
-3
@@ -54,17 +54,17 @@ def test_merge_tools_filters_duplicates() -> None:
|
||||
|
||||
def test_register_additional_client_tools_assigns_when_configured() -> None:
|
||||
"""register_additional_client_tools should set additional_tools on the chat client."""
|
||||
from agent_framework import BaseChatClient, FunctionInvocationConfiguration
|
||||
from agent_framework import BaseChatClient, normalize_function_invocation_configuration
|
||||
|
||||
mock_chat_client = MagicMock(spec=BaseChatClient)
|
||||
mock_chat_client.function_invocation_configuration = FunctionInvocationConfiguration()
|
||||
mock_chat_client.function_invocation_configuration = normalize_function_invocation_configuration(None)
|
||||
|
||||
agent = ChatAgent(chat_client=mock_chat_client)
|
||||
|
||||
tools = [DummyTool("x")]
|
||||
register_additional_client_tools(agent, tools)
|
||||
|
||||
assert mock_chat_client.function_invocation_configuration.additional_tools == tools
|
||||
assert mock_chat_client.function_invocation_configuration["additional_tools"] == tools
|
||||
|
||||
|
||||
def test_collect_server_tools_includes_mcp_tools_when_connected() -> None:
|
||||
+1
-1
@@ -408,7 +408,7 @@ def test_get_role_value_with_enum():
|
||||
|
||||
from agent_framework_ag_ui._utils import get_role_value
|
||||
|
||||
message = ChatMessage("user", [Content.from_text("test")])
|
||||
message = ChatMessage(role="user", contents=[Content.from_text("test")])
|
||||
result = get_role_value(message)
|
||||
assert result == "user"
|
||||
|
||||
@@ -1,124 +0,0 @@
|
||||
# Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
"""Shared test stubs for AG-UI tests."""
|
||||
|
||||
import sys
|
||||
from collections.abc import AsyncIterable, AsyncIterator, Awaitable, Callable, MutableSequence
|
||||
from types import SimpleNamespace
|
||||
from typing import Any, Generic
|
||||
|
||||
from agent_framework import (
|
||||
AgentProtocol,
|
||||
AgentResponse,
|
||||
AgentResponseUpdate,
|
||||
AgentThread,
|
||||
BaseChatClient,
|
||||
ChatMessage,
|
||||
ChatResponse,
|
||||
ChatResponseUpdate,
|
||||
Content,
|
||||
)
|
||||
from agent_framework._clients import TOptions_co
|
||||
|
||||
if sys.version_info >= (3, 12):
|
||||
from typing import override # type: ignore # pragma: no cover
|
||||
else:
|
||||
from typing_extensions import override # type: ignore[import] # pragma: no cover
|
||||
|
||||
StreamFn = Callable[..., AsyncIterator[ChatResponseUpdate]]
|
||||
ResponseFn = Callable[..., Awaitable[ChatResponse]]
|
||||
|
||||
|
||||
class StreamingChatClientStub(BaseChatClient[TOptions_co], Generic[TOptions_co]):
|
||||
"""Typed streaming stub that satisfies ChatClientProtocol."""
|
||||
|
||||
def __init__(self, stream_fn: StreamFn, response_fn: ResponseFn | None = None) -> None:
|
||||
super().__init__()
|
||||
self._stream_fn = stream_fn
|
||||
self._response_fn = response_fn
|
||||
|
||||
@override
|
||||
async def _inner_get_streaming_response(
|
||||
self, *, messages: MutableSequence[ChatMessage], options: dict[str, Any], **kwargs: Any
|
||||
) -> AsyncIterator[ChatResponseUpdate]:
|
||||
async for update in self._stream_fn(messages, options, **kwargs):
|
||||
yield update
|
||||
|
||||
@override
|
||||
async def _inner_get_response(
|
||||
self, *, messages: MutableSequence[ChatMessage], options: dict[str, Any], **kwargs: Any
|
||||
) -> ChatResponse:
|
||||
if self._response_fn is not None:
|
||||
return await self._response_fn(messages, options, **kwargs)
|
||||
|
||||
contents: list[Any] = []
|
||||
async for update in self._stream_fn(messages, options, **kwargs):
|
||||
contents.extend(update.contents)
|
||||
|
||||
return ChatResponse(
|
||||
messages=[ChatMessage("assistant", contents)],
|
||||
response_id="stub-response",
|
||||
)
|
||||
|
||||
|
||||
def stream_from_updates(updates: list[ChatResponseUpdate]) -> StreamFn:
|
||||
"""Create a stream function that yields from a static list of updates."""
|
||||
|
||||
async def _stream(
|
||||
messages: MutableSequence[ChatMessage], options: dict[str, Any], **kwargs: Any
|
||||
) -> AsyncIterator[ChatResponseUpdate]:
|
||||
for update in updates:
|
||||
yield update
|
||||
|
||||
return _stream
|
||||
|
||||
|
||||
class StubAgent(AgentProtocol):
|
||||
"""Minimal AgentProtocol stub for orchestrator tests."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
updates: list[AgentResponseUpdate] | None = None,
|
||||
*,
|
||||
agent_id: str = "stub-agent",
|
||||
agent_name: str | None = "stub-agent",
|
||||
default_options: Any | None = None,
|
||||
chat_client: Any | None = None,
|
||||
) -> None:
|
||||
self.id = agent_id
|
||||
self.name = agent_name
|
||||
self.description = "stub agent"
|
||||
self.updates = updates or [AgentResponseUpdate(contents=[Content.from_text(text="response")], role="assistant")]
|
||||
self.default_options: dict[str, Any] = (
|
||||
default_options if isinstance(default_options, dict) else {"tools": None, "response_format": None}
|
||||
)
|
||||
self.chat_client = chat_client or SimpleNamespace(function_invocation_configuration=None)
|
||||
self.messages_received: list[Any] = []
|
||||
self.tools_received: list[Any] | None = None
|
||||
|
||||
async def run(
|
||||
self,
|
||||
messages: str | ChatMessage | list[str] | list[ChatMessage] | None = None,
|
||||
*,
|
||||
thread: AgentThread | None = None,
|
||||
**kwargs: Any,
|
||||
) -> AgentResponse:
|
||||
return AgentResponse(messages=[], response_id="stub-response")
|
||||
|
||||
def run_stream(
|
||||
self,
|
||||
messages: str | ChatMessage | list[str] | list[ChatMessage] | None = None,
|
||||
*,
|
||||
thread: AgentThread | None = None,
|
||||
**kwargs: Any,
|
||||
) -> AsyncIterable[AgentResponseUpdate]:
|
||||
async def _stream() -> AsyncIterator[AgentResponseUpdate]:
|
||||
self.messages_received = [] if messages is None else list(messages) # type: ignore[arg-type]
|
||||
self.tools_received = kwargs.get("tools")
|
||||
for update in self.updates:
|
||||
yield update
|
||||
|
||||
return _stream()
|
||||
|
||||
def get_new_thread(self, **kwargs: Any) -> AgentThread:
|
||||
return AgentThread()
|
||||
Reference in New Issue
Block a user