mirror of
https://github.com/microsoft/agent-framework.git
synced 2026-06-16 21:04:09 +08:00
Python: name changes executed (#607)
* name changes executed * updated adr to accepted * renamed openai base config * renamed openai config to mixin * added renames in user docs * reverted mcperror * fix tests * remove sse from tests
This commit is contained in:
@@ -13,7 +13,7 @@ from agent_framework import (
|
||||
TextContent,
|
||||
)
|
||||
from agent_framework.exceptions import ServiceInitializationError
|
||||
from agent_framework.openai._chat_client import OpenAIChatClientBase
|
||||
from agent_framework.openai._chat_client import OpenAIBaseChatClient
|
||||
from azure.core.credentials import TokenCredential
|
||||
from openai.lib.azure import AsyncAzureADTokenProvider, AsyncAzureOpenAI
|
||||
from openai.types.chat.chat_completion import Choice
|
||||
@@ -22,7 +22,7 @@ from pydantic import SecretStr, ValidationError
|
||||
from pydantic.networks import AnyUrl
|
||||
|
||||
from ._shared import (
|
||||
AzureOpenAIConfigBase,
|
||||
AzureOpenAIConfigMixin,
|
||||
AzureOpenAISettings,
|
||||
)
|
||||
|
||||
@@ -37,7 +37,7 @@ TChatResponse = TypeVar("TChatResponse", ChatResponse, ChatResponseUpdate)
|
||||
TAzureChatClient = TypeVar("TAzureChatClient", bound="AzureChatClient")
|
||||
|
||||
|
||||
class AzureChatClient(AzureOpenAIConfigBase, OpenAIChatClientBase):
|
||||
class AzureChatClient(AzureOpenAIConfigMixin, OpenAIBaseChatClient):
|
||||
"""Azure Chat completion class."""
|
||||
|
||||
def __init__(
|
||||
@@ -143,7 +143,7 @@ class AzureChatClient(AzureOpenAIConfigBase, OpenAIChatClientBase):
|
||||
def _parse_text_from_choice(self, choice: Choice | ChunkChoice) -> TextContent | None:
|
||||
"""Parse the choice into a TextContent object.
|
||||
|
||||
Overwritten from OpenAIChatClientBase to deal with Azure On Your Data function.
|
||||
Overwritten from OpenAIBaseChatClient to deal with Azure On Your Data function.
|
||||
For docs see:
|
||||
https://learn.microsoft.com/en-us/azure/ai-foundry/openai/references/on-your-data?tabs=python#context
|
||||
"""
|
||||
|
||||
@@ -6,7 +6,7 @@ from urllib.parse import urljoin
|
||||
|
||||
from agent_framework import use_tool_calling
|
||||
from agent_framework.exceptions import ServiceInitializationError
|
||||
from agent_framework.openai._responses_client import OpenAIResponsesClientBase
|
||||
from agent_framework.openai._responses_client import OpenAIBaseResponsesClient
|
||||
from agent_framework.telemetry import use_telemetry
|
||||
from azure.core.credentials import TokenCredential
|
||||
from openai.lib.azure import AsyncAzureADTokenProvider, AsyncAzureOpenAI
|
||||
@@ -14,7 +14,7 @@ from pydantic import SecretStr, ValidationError
|
||||
from pydantic.networks import AnyUrl
|
||||
|
||||
from ._shared import (
|
||||
AzureOpenAIConfigBase,
|
||||
AzureOpenAIConfigMixin,
|
||||
AzureOpenAISettings,
|
||||
)
|
||||
|
||||
@@ -23,7 +23,7 @@ TAzureResponsesClient = TypeVar("TAzureResponsesClient", bound="AzureResponsesCl
|
||||
|
||||
@use_telemetry
|
||||
@use_tool_calling
|
||||
class AzureResponsesClient(AzureOpenAIConfigBase, OpenAIResponsesClientBase):
|
||||
class AzureResponsesClient(AzureOpenAIConfigMixin, OpenAIBaseResponsesClient):
|
||||
"""Azure Responses completion class."""
|
||||
|
||||
def __init__(
|
||||
|
||||
@@ -6,9 +6,9 @@ from collections.abc import Awaitable, Callable, Mapping
|
||||
from copy import copy
|
||||
from typing import Any, ClassVar, Final
|
||||
|
||||
from agent_framework._pydantic import AFBaseSettings, HttpsUrl
|
||||
from agent_framework._pydantic import AFBaseSettings, HTTPsUrl
|
||||
from agent_framework.exceptions import ServiceInitializationError
|
||||
from agent_framework.openai._shared import OpenAIHandler
|
||||
from agent_framework.openai._shared import OpenAIBase
|
||||
from agent_framework.telemetry import USER_AGENT_KEY
|
||||
from azure.core.credentials import TokenCredential
|
||||
from openai.lib.azure import AsyncAzureOpenAI
|
||||
@@ -126,8 +126,8 @@ class AzureOpenAISettings(AFBaseSettings):
|
||||
audio_to_text_deployment_name: str | None = None
|
||||
text_to_audio_deployment_name: str | None = None
|
||||
realtime_deployment_name: str | None = None
|
||||
endpoint: HttpsUrl | None = None
|
||||
base_url: HttpsUrl | None = None
|
||||
endpoint: HTTPsUrl | None = None
|
||||
base_url: HTTPsUrl | None = None
|
||||
api_key: SecretStr | None = None
|
||||
api_version: str | None = None
|
||||
token_endpoint: str | None = None
|
||||
@@ -165,7 +165,7 @@ class AzureOpenAISettings(AFBaseSettings):
|
||||
return self
|
||||
|
||||
|
||||
class AzureOpenAIConfigBase(OpenAIHandler):
|
||||
class AzureOpenAIConfigMixin(OpenAIBase):
|
||||
"""Internal class for configuring a connection to an Azure OpenAI service."""
|
||||
|
||||
MODEL_PROVIDER_NAME: ClassVar[str] = "azure_openai" # type: ignore[reportIncompatibleVariableOverride, misc]
|
||||
@@ -174,8 +174,8 @@ class AzureOpenAIConfigBase(OpenAIHandler):
|
||||
def __init__(
|
||||
self,
|
||||
deployment_name: str,
|
||||
endpoint: HttpsUrl | None = None,
|
||||
base_url: HttpsUrl | None = None,
|
||||
endpoint: HTTPsUrl | None = None,
|
||||
base_url: HTTPsUrl | None = None,
|
||||
api_version: str = DEFAULT_AZURE_API_VERSION,
|
||||
api_key: str | None = None,
|
||||
ad_token: str | None = None,
|
||||
@@ -190,7 +190,7 @@ class AzureOpenAIConfigBase(OpenAIHandler):
|
||||
"""Internal class for configuring a connection to an Azure OpenAI service.
|
||||
|
||||
The `validate_call` decorator is used with a configuration that allows arbitrary types.
|
||||
This is necessary for types like `HttpsUrl` and `OpenAIModelTypes`.
|
||||
This is necessary for types like `HTTPsUrl` and `OpenAIModelTypes`.
|
||||
|
||||
Args:
|
||||
deployment_name: Name of the deployment.
|
||||
|
||||
@@ -9,8 +9,8 @@ from agent_framework import (
|
||||
AgentRunResponse,
|
||||
AgentRunResponseUpdate,
|
||||
AgentThread,
|
||||
ChatClient,
|
||||
ChatClientAgent,
|
||||
ChatAgent,
|
||||
ChatClientProtocol,
|
||||
ChatMessage,
|
||||
ChatResponse,
|
||||
ChatResponseUpdate,
|
||||
@@ -93,7 +93,7 @@ def test_azure_assistants_client_init_with_client(mock_async_azure_openai: Magic
|
||||
assert chat_client.assistant_id == "existing-assistant-id"
|
||||
assert chat_client.thread_id == "test-thread-id"
|
||||
assert not chat_client._should_delete_assistant # type: ignore
|
||||
assert isinstance(chat_client, ChatClient)
|
||||
assert isinstance(chat_client, ChatClientProtocol)
|
||||
|
||||
|
||||
def test_azure_assistants_client_init_auto_create_client(
|
||||
@@ -147,7 +147,7 @@ def test_azure_assistants_client_init_with_default_headers(azure_openai_unit_tes
|
||||
)
|
||||
|
||||
assert chat_client.ai_model_id == "test_chat_deployment"
|
||||
assert isinstance(chat_client, ChatClient)
|
||||
assert isinstance(chat_client, ChatClientProtocol)
|
||||
|
||||
# Assert that the default header we added is present in the client's default headers
|
||||
for key, value in default_headers.items():
|
||||
@@ -267,7 +267,7 @@ def get_weather(
|
||||
async def test_azure_assistants_client_get_response() -> None:
|
||||
"""Test Azure Assistants Client response."""
|
||||
async with AzureAssistantsClient(credential=AzureCliCredential()) as azure_assistants_client:
|
||||
assert isinstance(azure_assistants_client, ChatClient)
|
||||
assert isinstance(azure_assistants_client, ChatClientProtocol)
|
||||
|
||||
messages: list[ChatMessage] = []
|
||||
messages.append(
|
||||
@@ -291,7 +291,7 @@ async def test_azure_assistants_client_get_response() -> None:
|
||||
async def test_azure_assistants_client_get_response_tools() -> None:
|
||||
"""Test Azure Assistants Client response with tools."""
|
||||
async with AzureAssistantsClient(credential=AzureCliCredential()) as azure_assistants_client:
|
||||
assert isinstance(azure_assistants_client, ChatClient)
|
||||
assert isinstance(azure_assistants_client, ChatClientProtocol)
|
||||
|
||||
messages: list[ChatMessage] = []
|
||||
messages.append(ChatMessage(role="user", text="What's the weather like in Seattle?"))
|
||||
@@ -312,7 +312,7 @@ async def test_azure_assistants_client_get_response_tools() -> None:
|
||||
async def test_azure_assistants_client_streaming() -> None:
|
||||
"""Test Azure Assistants Client streaming response."""
|
||||
async with AzureAssistantsClient(credential=AzureCliCredential()) as azure_assistants_client:
|
||||
assert isinstance(azure_assistants_client, ChatClient)
|
||||
assert isinstance(azure_assistants_client, ChatClientProtocol)
|
||||
|
||||
messages: list[ChatMessage] = []
|
||||
messages.append(
|
||||
@@ -342,7 +342,7 @@ async def test_azure_assistants_client_streaming() -> None:
|
||||
async def test_azure_assistants_client_streaming_tools() -> None:
|
||||
"""Test Azure Assistants Client streaming response with tools."""
|
||||
async with AzureAssistantsClient(credential=AzureCliCredential()) as azure_assistants_client:
|
||||
assert isinstance(azure_assistants_client, ChatClient)
|
||||
assert isinstance(azure_assistants_client, ChatClientProtocol)
|
||||
|
||||
messages: list[ChatMessage] = []
|
||||
messages.append(ChatMessage(role="user", text="What's the weather like in Seattle?"))
|
||||
@@ -378,7 +378,7 @@ async def test_azure_assistants_client_with_existing_assistant() -> None:
|
||||
async with AzureAssistantsClient(
|
||||
assistant_id=assistant_id, credential=AzureCliCredential()
|
||||
) as azure_assistants_client:
|
||||
assert isinstance(azure_assistants_client, ChatClient)
|
||||
assert isinstance(azure_assistants_client, ChatClientProtocol)
|
||||
assert azure_assistants_client.assistant_id == assistant_id
|
||||
|
||||
messages = [ChatMessage(role="user", text="What can you do?")]
|
||||
@@ -393,8 +393,8 @@ async def test_azure_assistants_client_with_existing_assistant() -> None:
|
||||
|
||||
@skip_if_azure_integration_tests_disabled
|
||||
async def test_azure_assistants_agent_basic_run():
|
||||
"""Test ChatClientAgent basic run functionality with AzureAssistantsClient."""
|
||||
async with ChatClientAgent(
|
||||
"""Test ChatAgent basic run functionality with AzureAssistantsClient."""
|
||||
async with ChatAgent(
|
||||
chat_client=AzureAssistantsClient(credential=AzureCliCredential()),
|
||||
) as agent:
|
||||
# Run a simple query
|
||||
@@ -409,13 +409,13 @@ async def test_azure_assistants_agent_basic_run():
|
||||
|
||||
@skip_if_azure_integration_tests_disabled
|
||||
async def test_azure_assistants_agent_basic_run_streaming():
|
||||
"""Test ChatClientAgent basic streaming functionality with AzureAssistantsClient."""
|
||||
async with ChatClientAgent(
|
||||
"""Test ChatAgent basic streaming functionality with AzureAssistantsClient."""
|
||||
async with ChatAgent(
|
||||
chat_client=AzureAssistantsClient(credential=AzureCliCredential()),
|
||||
) as agent:
|
||||
# Run streaming query
|
||||
full_message: str = ""
|
||||
async for chunk in agent.run_streaming("Please respond with exactly: 'This is a streaming response test.'"):
|
||||
async for chunk in agent.run_stream("Please respond with exactly: 'This is a streaming response test.'"):
|
||||
assert chunk is not None
|
||||
assert isinstance(chunk, AgentRunResponseUpdate)
|
||||
if chunk.text:
|
||||
@@ -428,8 +428,8 @@ async def test_azure_assistants_agent_basic_run_streaming():
|
||||
|
||||
@skip_if_azure_integration_tests_disabled
|
||||
async def test_azure_assistants_agent_thread_persistence():
|
||||
"""Test ChatClientAgent thread persistence across runs with AzureAssistantsClient."""
|
||||
async with ChatClientAgent(
|
||||
"""Test ChatAgent thread persistence across runs with AzureAssistantsClient."""
|
||||
async with ChatAgent(
|
||||
chat_client=AzureAssistantsClient(credential=AzureCliCredential()),
|
||||
instructions="You are a helpful assistant with good memory.",
|
||||
) as agent:
|
||||
@@ -456,11 +456,11 @@ async def test_azure_assistants_agent_thread_persistence():
|
||||
|
||||
@skip_if_azure_integration_tests_disabled
|
||||
async def test_azure_assistants_agent_existing_thread_id():
|
||||
"""Test ChatClientAgent with existing thread ID to continue conversations across agent instances."""
|
||||
"""Test ChatAgent with existing thread ID to continue conversations across agent instances."""
|
||||
# First, create a conversation and capture the thread ID
|
||||
existing_thread_id = None
|
||||
|
||||
async with ChatClientAgent(
|
||||
async with ChatAgent(
|
||||
chat_client=AzureAssistantsClient(credential=AzureCliCredential()),
|
||||
instructions="You are a helpful weather agent.",
|
||||
tools=[get_weather],
|
||||
@@ -480,7 +480,7 @@ async def test_azure_assistants_agent_existing_thread_id():
|
||||
|
||||
# Now continue with the same thread ID in a new agent instance
|
||||
|
||||
async with ChatClientAgent(
|
||||
async with ChatAgent(
|
||||
chat_client=AzureAssistantsClient(thread_id=existing_thread_id, credential=AzureCliCredential()),
|
||||
instructions="You are a helpful weather agent.",
|
||||
tools=[get_weather],
|
||||
@@ -500,9 +500,9 @@ async def test_azure_assistants_agent_existing_thread_id():
|
||||
|
||||
@skip_if_azure_integration_tests_disabled
|
||||
async def test_azure_assistants_agent_code_interpreter():
|
||||
"""Test ChatClientAgent with code interpreter through AzureAssistantsClient."""
|
||||
"""Test ChatAgent with code interpreter through AzureAssistantsClient."""
|
||||
|
||||
async with ChatClientAgent(
|
||||
async with ChatAgent(
|
||||
chat_client=AzureAssistantsClient(credential=AzureCliCredential()),
|
||||
instructions="You are a helpful assistant that can write and execute Python code.",
|
||||
tools=[HostedCodeInterpreterTool()],
|
||||
@@ -521,7 +521,7 @@ async def test_azure_assistants_agent_code_interpreter():
|
||||
async def test_azure_assistants_client_agent_level_tool_persistence():
|
||||
"""Test that agent-level tools persist across multiple runs with Azure Assistants Client."""
|
||||
|
||||
async with ChatClientAgent(
|
||||
async with ChatAgent(
|
||||
chat_client=AzureAssistantsClient(credential=AzureCliCredential()),
|
||||
instructions="You are a helpful assistant that uses available tools.",
|
||||
tools=[get_weather], # Agent-level tool
|
||||
@@ -556,7 +556,7 @@ async def test_azure_assistants_client_run_level_tool_isolation():
|
||||
call_count += 1
|
||||
return f"The weather in {location} is sunny and 72°F."
|
||||
|
||||
async with ChatClientAgent(
|
||||
async with ChatAgent(
|
||||
chat_client=AzureAssistantsClient(credential=AzureCliCredential()),
|
||||
instructions="You are a helpful assistant.",
|
||||
) as agent:
|
||||
|
||||
@@ -10,9 +10,9 @@ import pytest
|
||||
from agent_framework import (
|
||||
AgentRunResponse,
|
||||
AgentRunResponseUpdate,
|
||||
ChatClient,
|
||||
ChatClientAgent,
|
||||
ChatClientBase,
|
||||
BaseChatClient,
|
||||
ChatAgent,
|
||||
ChatClientProtocol,
|
||||
ChatMessage,
|
||||
ChatResponse,
|
||||
ChatResponseUpdate,
|
||||
@@ -55,7 +55,7 @@ def test_init(azure_openai_unit_test_env: dict[str, str]) -> None:
|
||||
assert azure_chat_client.client is not None
|
||||
assert isinstance(azure_chat_client.client, AsyncAzureOpenAI)
|
||||
assert azure_chat_client.ai_model_id == azure_openai_unit_test_env["AZURE_OPENAI_CHAT_DEPLOYMENT_NAME"]
|
||||
assert isinstance(azure_chat_client, ChatClientBase)
|
||||
assert isinstance(azure_chat_client, BaseChatClient)
|
||||
|
||||
|
||||
def test_init_client(azure_openai_unit_test_env: dict[str, str]) -> None:
|
||||
@@ -78,7 +78,7 @@ def test_init_base_url(azure_openai_unit_test_env: dict[str, str]) -> None:
|
||||
assert azure_chat_client.client is not None
|
||||
assert isinstance(azure_chat_client.client, AsyncAzureOpenAI)
|
||||
assert azure_chat_client.ai_model_id == azure_openai_unit_test_env["AZURE_OPENAI_CHAT_DEPLOYMENT_NAME"]
|
||||
assert isinstance(azure_chat_client, ChatClientBase)
|
||||
assert isinstance(azure_chat_client, BaseChatClient)
|
||||
for key, value in default_headers.items():
|
||||
assert key in azure_chat_client.client.default_headers
|
||||
assert azure_chat_client.client.default_headers[key] == value
|
||||
@@ -91,7 +91,7 @@ def test_init_endpoint(azure_openai_unit_test_env: dict[str, str]) -> None:
|
||||
assert azure_chat_client.client is not None
|
||||
assert isinstance(azure_chat_client.client, AsyncAzureOpenAI)
|
||||
assert azure_chat_client.ai_model_id == azure_openai_unit_test_env["AZURE_OPENAI_CHAT_DEPLOYMENT_NAME"]
|
||||
assert isinstance(azure_chat_client, ChatClientBase)
|
||||
assert isinstance(azure_chat_client, BaseChatClient)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("exclude_list", [["AZURE_OPENAI_CHAT_DEPLOYMENT_NAME"]], indirect=True)
|
||||
@@ -614,7 +614,7 @@ def get_weather(location: str) -> str:
|
||||
async def test_azure_openai_chat_client_response() -> None:
|
||||
"""Test Azure OpenAI chat completion responses."""
|
||||
azure_chat_client = AzureChatClient(credential=AzureCliCredential())
|
||||
assert isinstance(azure_chat_client, ChatClient)
|
||||
assert isinstance(azure_chat_client, ChatClientProtocol)
|
||||
|
||||
messages: list[ChatMessage] = []
|
||||
messages.append(
|
||||
@@ -643,7 +643,7 @@ async def test_azure_openai_chat_client_response() -> None:
|
||||
async def test_azure_openai_chat_client_response_tools() -> None:
|
||||
"""Test AzureOpenAI chat completion responses."""
|
||||
azure_chat_client = AzureChatClient(credential=AzureCliCredential())
|
||||
assert isinstance(azure_chat_client, ChatClient)
|
||||
assert isinstance(azure_chat_client, ChatClientProtocol)
|
||||
|
||||
messages: list[ChatMessage] = []
|
||||
messages.append(ChatMessage(role="user", text="who are Emily and David?"))
|
||||
@@ -664,7 +664,7 @@ async def test_azure_openai_chat_client_response_tools() -> None:
|
||||
async def test_azure_openai_chat_client_streaming() -> None:
|
||||
"""Test Azure OpenAI chat completion responses."""
|
||||
azure_chat_client = AzureChatClient(credential=AzureCliCredential())
|
||||
assert isinstance(azure_chat_client, ChatClient)
|
||||
assert isinstance(azure_chat_client, ChatClientProtocol)
|
||||
|
||||
messages: list[ChatMessage] = []
|
||||
messages.append(
|
||||
@@ -698,7 +698,7 @@ async def test_azure_openai_chat_client_streaming() -> None:
|
||||
async def test_azure_openai_chat_client_streaming_tools() -> None:
|
||||
"""Test AzureOpenAI chat completion responses."""
|
||||
azure_chat_client = AzureChatClient(credential=AzureCliCredential())
|
||||
assert isinstance(azure_chat_client, ChatClient)
|
||||
assert isinstance(azure_chat_client, ChatClientProtocol)
|
||||
|
||||
messages: list[ChatMessage] = []
|
||||
messages.append(ChatMessage(role="user", text="who are Emily and David?"))
|
||||
@@ -723,7 +723,7 @@ async def test_azure_openai_chat_client_streaming_tools() -> None:
|
||||
@skip_if_azure_integration_tests_disabled
|
||||
async def test_azure_openai_chat_client_agent_basic_run():
|
||||
"""Test Azure OpenAI chat client agent basic run functionality with AzureChatClient."""
|
||||
async with ChatClientAgent(
|
||||
async with ChatAgent(
|
||||
chat_client=AzureChatClient(credential=AzureCliCredential()),
|
||||
) as agent:
|
||||
# Test basic run
|
||||
@@ -738,12 +738,12 @@ async def test_azure_openai_chat_client_agent_basic_run():
|
||||
@skip_if_azure_integration_tests_disabled
|
||||
async def test_azure_openai_chat_client_agent_basic_run_streaming():
|
||||
"""Test Azure OpenAI chat client agent basic streaming functionality with AzureChatClient."""
|
||||
async with ChatClientAgent(
|
||||
async with ChatAgent(
|
||||
chat_client=AzureChatClient(credential=AzureCliCredential()),
|
||||
) as agent:
|
||||
# Test streaming run
|
||||
full_text = ""
|
||||
async for chunk in agent.run_streaming("Please respond with exactly: 'This is a streaming response test.'"):
|
||||
async for chunk in agent.run_stream("Please respond with exactly: 'This is a streaming response test.'"):
|
||||
assert isinstance(chunk, AgentRunResponseUpdate)
|
||||
if chunk.text:
|
||||
full_text += chunk.text
|
||||
@@ -755,7 +755,7 @@ async def test_azure_openai_chat_client_agent_basic_run_streaming():
|
||||
@skip_if_azure_integration_tests_disabled
|
||||
async def test_azure_openai_chat_client_agent_thread_persistence():
|
||||
"""Test Azure OpenAI chat client agent thread persistence across runs with AzureChatClient."""
|
||||
async with ChatClientAgent(
|
||||
async with ChatAgent(
|
||||
chat_client=AzureChatClient(credential=AzureCliCredential()),
|
||||
instructions="You are a helpful assistant with good memory.",
|
||||
) as agent:
|
||||
@@ -782,7 +782,7 @@ async def test_azure_openai_chat_client_agent_existing_thread():
|
||||
# First conversation - capture the thread
|
||||
preserved_thread = None
|
||||
|
||||
async with ChatClientAgent(
|
||||
async with ChatAgent(
|
||||
chat_client=AzureChatClient(credential=AzureCliCredential()),
|
||||
instructions="You are a helpful assistant with good memory.",
|
||||
) as first_agent:
|
||||
@@ -798,7 +798,7 @@ async def test_azure_openai_chat_client_agent_existing_thread():
|
||||
|
||||
# Second conversation - reuse the thread in a new agent instance
|
||||
if preserved_thread:
|
||||
async with ChatClientAgent(
|
||||
async with ChatAgent(
|
||||
chat_client=AzureChatClient(credential=AzureCliCredential()),
|
||||
instructions="You are a helpful assistant with good memory.",
|
||||
) as second_agent:
|
||||
@@ -814,7 +814,7 @@ async def test_azure_openai_chat_client_agent_existing_thread():
|
||||
async def test_azure_chat_client_agent_level_tool_persistence():
|
||||
"""Test that agent-level tools persist across multiple runs with Azure Chat Client."""
|
||||
|
||||
async with ChatClientAgent(
|
||||
async with ChatAgent(
|
||||
chat_client=AzureChatClient(credential=AzureCliCredential()),
|
||||
instructions="You are a helpful assistant that uses available tools.",
|
||||
tools=[get_weather], # Agent-level tool
|
||||
@@ -849,7 +849,7 @@ async def test_azure_chat_client_run_level_tool_isolation():
|
||||
call_count += 1
|
||||
return f"The weather in {location} is sunny and 72°F."
|
||||
|
||||
async with ChatClientAgent(
|
||||
async with ChatAgent(
|
||||
chat_client=AzureChatClient(credential=AzureCliCredential()),
|
||||
instructions="You are a helpful assistant.",
|
||||
) as agent:
|
||||
|
||||
@@ -8,8 +8,8 @@ from agent_framework import (
|
||||
AgentRunResponse,
|
||||
AgentRunResponseUpdate,
|
||||
AgentThread,
|
||||
ChatClient,
|
||||
ChatClientAgent,
|
||||
ChatAgent,
|
||||
ChatClientProtocol,
|
||||
ChatMessage,
|
||||
ChatResponse,
|
||||
ChatResponseUpdate,
|
||||
@@ -50,7 +50,7 @@ def test_init(azure_openai_unit_test_env: dict[str, str]) -> None:
|
||||
azure_responses_client = AzureResponsesClient()
|
||||
|
||||
assert azure_responses_client.ai_model_id == azure_openai_unit_test_env["AZURE_OPENAI_RESPONSES_DEPLOYMENT_NAME"]
|
||||
assert isinstance(azure_responses_client, ChatClient)
|
||||
assert isinstance(azure_responses_client, ChatClientProtocol)
|
||||
|
||||
|
||||
def test_init_validation_fail() -> None:
|
||||
@@ -65,7 +65,7 @@ def test_init_ai_model_id_constructor(azure_openai_unit_test_env: dict[str, str]
|
||||
azure_responses_client = AzureResponsesClient(deployment_name=ai_model_id)
|
||||
|
||||
assert azure_responses_client.ai_model_id == ai_model_id
|
||||
assert isinstance(azure_responses_client, ChatClient)
|
||||
assert isinstance(azure_responses_client, ChatClientProtocol)
|
||||
|
||||
|
||||
def test_init_with_default_header(azure_openai_unit_test_env: dict[str, str]) -> None:
|
||||
@@ -77,7 +77,7 @@ def test_init_with_default_header(azure_openai_unit_test_env: dict[str, str]) ->
|
||||
)
|
||||
|
||||
assert azure_responses_client.ai_model_id == azure_openai_unit_test_env["AZURE_OPENAI_RESPONSES_DEPLOYMENT_NAME"]
|
||||
assert isinstance(azure_responses_client, ChatClient)
|
||||
assert isinstance(azure_responses_client, ChatClientProtocol)
|
||||
|
||||
# Assert that the default header we added is present in the client's default headers
|
||||
for key, value in default_headers.items():
|
||||
@@ -119,7 +119,7 @@ async def test_azure_responses_client_response() -> None:
|
||||
"""Test azure responses client responses."""
|
||||
azure_responses_client = AzureResponsesClient(credential=AzureCliCredential())
|
||||
|
||||
assert isinstance(azure_responses_client, ChatClient)
|
||||
assert isinstance(azure_responses_client, ChatClientProtocol)
|
||||
|
||||
messages: list[ChatMessage] = []
|
||||
messages.append(
|
||||
@@ -162,7 +162,7 @@ async def test_azure_responses_client_response_tools() -> None:
|
||||
"""Test azure responses client tools."""
|
||||
azure_responses_client = AzureResponsesClient(credential=AzureCliCredential())
|
||||
|
||||
assert isinstance(azure_responses_client, ChatClient)
|
||||
assert isinstance(azure_responses_client, ChatClientProtocol)
|
||||
|
||||
messages: list[ChatMessage] = []
|
||||
messages.append(ChatMessage(role="user", text="What is the weather in New York?"))
|
||||
@@ -201,7 +201,7 @@ async def test_azure_responses_client_streaming() -> None:
|
||||
"""Test Azure azure responses client streaming responses."""
|
||||
azure_responses_client = AzureResponsesClient(credential=AzureCliCredential())
|
||||
|
||||
assert isinstance(azure_responses_client, ChatClient)
|
||||
assert isinstance(azure_responses_client, ChatClientProtocol)
|
||||
|
||||
messages: list[ChatMessage] = []
|
||||
messages.append(
|
||||
@@ -251,7 +251,7 @@ async def test_azure_responses_client_streaming_tools() -> None:
|
||||
"""Test azure responses client streaming tools."""
|
||||
azure_responses_client = AzureResponsesClient(credential=AzureCliCredential())
|
||||
|
||||
assert isinstance(azure_responses_client, ChatClient)
|
||||
assert isinstance(azure_responses_client, ChatClientProtocol)
|
||||
|
||||
messages: list[ChatMessage] = [ChatMessage(role="user", text="What is the weather in Seattle?")]
|
||||
|
||||
@@ -312,12 +312,12 @@ async def test_azure_responses_client_agent_basic_run():
|
||||
@skip_if_azure_integration_tests_disabled
|
||||
async def test_azure_responses_client_agent_basic_run_streaming():
|
||||
"""Test Azure Responses Client agent basic streaming functionality with AzureResponsesClient."""
|
||||
async with ChatClientAgent(
|
||||
async with ChatAgent(
|
||||
chat_client=AzureResponsesClient(credential=AzureCliCredential()),
|
||||
) as agent:
|
||||
# Test streaming run
|
||||
full_text = ""
|
||||
async for chunk in agent.run_streaming("Please respond with exactly: 'This is a streaming response test.'"):
|
||||
async for chunk in agent.run_stream("Please respond with exactly: 'This is a streaming response test.'"):
|
||||
assert isinstance(chunk, AgentRunResponseUpdate)
|
||||
if chunk.text:
|
||||
full_text += chunk.text
|
||||
@@ -329,7 +329,7 @@ async def test_azure_responses_client_agent_basic_run_streaming():
|
||||
@skip_if_azure_integration_tests_disabled
|
||||
async def test_azure_responses_client_agent_thread_persistence():
|
||||
"""Test Azure Responses Client agent thread persistence across runs with AzureResponsesClient."""
|
||||
async with ChatClientAgent(
|
||||
async with ChatAgent(
|
||||
chat_client=AzureResponsesClient(credential=AzureCliCredential()),
|
||||
instructions="You are a helpful assistant with good memory.",
|
||||
) as agent:
|
||||
@@ -352,7 +352,7 @@ async def test_azure_responses_client_agent_thread_persistence():
|
||||
@skip_if_azure_integration_tests_disabled
|
||||
async def test_azure_responses_client_agent_thread_storage_with_store_true():
|
||||
"""Test Azure Responses Client agent with store=True to verify service_thread_id is returned."""
|
||||
async with ChatClientAgent(
|
||||
async with ChatAgent(
|
||||
chat_client=AzureResponsesClient(credential=AzureCliCredential()),
|
||||
instructions="You are a helpful assistant.",
|
||||
) as agent:
|
||||
@@ -386,7 +386,7 @@ async def test_azure_responses_client_agent_existing_thread():
|
||||
# First conversation - capture the thread
|
||||
preserved_thread = None
|
||||
|
||||
async with ChatClientAgent(
|
||||
async with ChatAgent(
|
||||
chat_client=AzureResponsesClient(credential=AzureCliCredential()),
|
||||
instructions="You are a helpful assistant with good memory.",
|
||||
) as first_agent:
|
||||
@@ -402,7 +402,7 @@ async def test_azure_responses_client_agent_existing_thread():
|
||||
|
||||
# Second conversation - reuse the thread in a new agent instance
|
||||
if preserved_thread:
|
||||
async with ChatClientAgent(
|
||||
async with ChatAgent(
|
||||
chat_client=AzureResponsesClient(credential=AzureCliCredential()),
|
||||
instructions="You are a helpful assistant with good memory.",
|
||||
) as second_agent:
|
||||
@@ -417,7 +417,7 @@ async def test_azure_responses_client_agent_existing_thread():
|
||||
@skip_if_azure_integration_tests_disabled
|
||||
async def test_azure_responses_client_agent_hosted_code_interpreter_tool():
|
||||
"""Test Azure Responses Client agent with HostedCodeInterpreterTool through AzureResponsesClient."""
|
||||
async with ChatClientAgent(
|
||||
async with ChatAgent(
|
||||
chat_client=AzureResponsesClient(credential=AzureCliCredential()),
|
||||
instructions="You are a helpful assistant that can execute Python code.",
|
||||
tools=[HostedCodeInterpreterTool()],
|
||||
@@ -439,7 +439,7 @@ async def test_azure_responses_client_agent_hosted_code_interpreter_tool():
|
||||
async def test_azure_responses_client_agent_level_tool_persistence():
|
||||
"""Test that agent-level tools persist across multiple runs with Azure Responses Client."""
|
||||
|
||||
async with ChatClientAgent(
|
||||
async with ChatAgent(
|
||||
chat_client=AzureResponsesClient(credential=AzureCliCredential()),
|
||||
instructions="You are a helpful assistant that uses available tools.",
|
||||
tools=[get_weather], # Agent-level tool
|
||||
@@ -474,7 +474,7 @@ async def test_azure_responses_client_run_level_tool_isolation():
|
||||
call_count += 1
|
||||
return f"The weather in {location} is sunny and 72°F."
|
||||
|
||||
async with ChatClientAgent(
|
||||
async with ChatAgent(
|
||||
chat_client=AzureResponsesClient(credential=AzureCliCredential()),
|
||||
instructions="You are a helpful assistant.",
|
||||
) as agent:
|
||||
|
||||
@@ -7,19 +7,19 @@ from collections.abc import AsyncIterable, MutableMapping, MutableSequence
|
||||
from typing import Any, ClassVar, TypeVar
|
||||
|
||||
from agent_framework import (
|
||||
AIContents,
|
||||
AIFunction,
|
||||
ChatClientBase,
|
||||
BaseChatClient,
|
||||
ChatMessage,
|
||||
ChatOptions,
|
||||
ChatResponse,
|
||||
ChatResponseUpdate,
|
||||
ChatRole,
|
||||
ChatToolMode,
|
||||
Contents,
|
||||
DataContent,
|
||||
FunctionCallContent,
|
||||
FunctionResultContent,
|
||||
HostedCodeInterpreterTool,
|
||||
Role,
|
||||
TextContent,
|
||||
UriContent,
|
||||
UsageContent,
|
||||
@@ -98,7 +98,7 @@ TFoundryChatClient = TypeVar("TFoundryChatClient", bound="FoundryChatClient")
|
||||
|
||||
@use_telemetry
|
||||
@use_tool_calling
|
||||
class FoundryChatClient(ChatClientBase):
|
||||
class FoundryChatClient(BaseChatClient):
|
||||
"""Azure AI Foundry Chat client."""
|
||||
|
||||
MODEL_PROVIDER_NAME: ClassVar[str] = "azure_ai_foundry" # type: ignore[reportIncompatibleVariableOverride, misc]
|
||||
@@ -385,12 +385,12 @@ class FoundryChatClient(ChatClientBase):
|
||||
message_id=response_id,
|
||||
raw_representation=event_data,
|
||||
response_id=response_id,
|
||||
role=ChatRole.ASSISTANT,
|
||||
role=Role.ASSISTANT,
|
||||
)
|
||||
elif event_type == AgentStreamEvent.THREAD_RUN_STEP_CREATED and isinstance(event_data, RunStep):
|
||||
response_id = event_data.run_id
|
||||
elif event_type == AgentStreamEvent.THREAD_MESSAGE_DELTA and isinstance(event_data, MessageDeltaChunk):
|
||||
role = ChatRole.USER if event_data.delta.role == MessageRole.USER else ChatRole.ASSISTANT
|
||||
role = Role.USER if event_data.delta.role == MessageRole.USER else Role.ASSISTANT
|
||||
yield ChatResponseUpdate(
|
||||
role=role,
|
||||
text=event_data.text,
|
||||
@@ -407,7 +407,7 @@ class FoundryChatClient(ChatClientBase):
|
||||
contents = self._create_function_call_contents(event_data, response_id)
|
||||
if contents:
|
||||
yield ChatResponseUpdate(
|
||||
role=ChatRole.ASSISTANT,
|
||||
role=Role.ASSISTANT,
|
||||
contents=contents,
|
||||
conversation_id=thread_id,
|
||||
message_id=response_id,
|
||||
@@ -427,7 +427,7 @@ class FoundryChatClient(ChatClientBase):
|
||||
)
|
||||
)
|
||||
yield ChatResponseUpdate(
|
||||
role=ChatRole.ASSISTANT,
|
||||
role=Role.ASSISTANT,
|
||||
contents=[usage_content],
|
||||
conversation_id=thread_id,
|
||||
message_id=response_id,
|
||||
@@ -447,12 +447,12 @@ class FoundryChatClient(ChatClientBase):
|
||||
message_id=response_id,
|
||||
raw_representation=event_data, # type: ignore
|
||||
response_id=response_id,
|
||||
role=ChatRole.ASSISTANT,
|
||||
role=Role.ASSISTANT,
|
||||
)
|
||||
|
||||
def _create_function_call_contents(self, event_data: ThreadRun, response_id: str | None) -> list[AIContents]:
|
||||
def _create_function_call_contents(self, event_data: ThreadRun, response_id: str | None) -> list[Contents]:
|
||||
"""Create function call contents from a tool action event."""
|
||||
contents: list[AIContents] = []
|
||||
contents: list[Contents] = []
|
||||
|
||||
if isinstance(event_data.required_action, SubmitToolOutputsAction):
|
||||
for tool_call in event_data.required_action.submit_tool_outputs.tool_calls:
|
||||
@@ -563,7 +563,7 @@ class FoundryChatClient(ChatClientBase):
|
||||
additional_messages = []
|
||||
additional_messages.append(
|
||||
ThreadMessageOptions(
|
||||
role=MessageRole.AGENT if chat_message.role == ChatRole.ASSISTANT else MessageRole.USER,
|
||||
role=MessageRole.AGENT if chat_message.role == Role.ASSISTANT else MessageRole.USER,
|
||||
content=message_contents,
|
||||
)
|
||||
)
|
||||
|
||||
@@ -9,16 +9,16 @@ from agent_framework import (
|
||||
AgentRunResponse,
|
||||
AgentRunResponseUpdate,
|
||||
AgentThread,
|
||||
ChatClient,
|
||||
ChatClientAgent,
|
||||
ChatAgent,
|
||||
ChatClientProtocol,
|
||||
ChatMessage,
|
||||
ChatOptions,
|
||||
ChatResponse,
|
||||
ChatResponseUpdate,
|
||||
ChatRole,
|
||||
FunctionCallContent,
|
||||
FunctionResultContent,
|
||||
HostedCodeInterpreterTool,
|
||||
Role,
|
||||
TextContent,
|
||||
UriContent,
|
||||
ai_function,
|
||||
@@ -98,7 +98,7 @@ def test_foundry_chat_client_init_with_client(mock_ai_project_client: MagicMock)
|
||||
assert chat_client.agent_id == "existing-agent-id"
|
||||
assert chat_client.thread_id == "test-thread-id"
|
||||
assert not chat_client._should_delete_agent # type: ignore
|
||||
assert isinstance(chat_client, ChatClient)
|
||||
assert isinstance(chat_client, ChatClientProtocol)
|
||||
|
||||
|
||||
def test_foundry_chat_client_init_auto_create_client(
|
||||
@@ -252,9 +252,9 @@ async def test_foundry_chat_client_tool_results_without_thread_error_via_public_
|
||||
|
||||
# Create messages with tool results but no thread/conversation ID
|
||||
messages = [
|
||||
ChatMessage(role=ChatRole.USER, text="Hello"),
|
||||
ChatMessage(role=Role.USER, text="Hello"),
|
||||
ChatMessage(
|
||||
role=ChatRole.TOOL, contents=[FunctionResultContent(call_id='["run_123", "call_456"]', result="Result")]
|
||||
role=Role.TOOL, contents=[FunctionResultContent(call_id='["run_123", "call_456"]', result="Result")]
|
||||
),
|
||||
]
|
||||
|
||||
@@ -283,7 +283,7 @@ async def test_foundry_chat_client_thread_management_through_public_api(mock_ai_
|
||||
mock_stream.__aenter__ = AsyncMock(return_value=empty_async_iter())
|
||||
mock_stream.__aexit__ = AsyncMock(return_value=None)
|
||||
|
||||
messages = [ChatMessage(role=ChatRole.USER, text="Hello")]
|
||||
messages = [ChatMessage(role=Role.USER, text="Hello")]
|
||||
|
||||
# Call without existing thread - should create new one
|
||||
response = chat_client.get_streaming_response(messages)
|
||||
@@ -379,7 +379,7 @@ def test_foundry_chat_client_create_run_options_basic(mock_ai_project_client: Ma
|
||||
"""Test _create_run_options with basic ChatOptions."""
|
||||
chat_client = create_test_foundry_chat_client(mock_ai_project_client)
|
||||
|
||||
messages = [ChatMessage(role=ChatRole.USER, text="Hello")]
|
||||
messages = [ChatMessage(role=Role.USER, text="Hello")]
|
||||
chat_options = ChatOptions(max_tokens=100, temperature=0.7)
|
||||
|
||||
run_options, tool_results = chat_client._create_run_options(messages, chat_options) # type: ignore
|
||||
@@ -392,7 +392,7 @@ def test_foundry_chat_client_create_run_options_no_chat_options(mock_ai_project_
|
||||
"""Test _create_run_options with no ChatOptions."""
|
||||
chat_client = create_test_foundry_chat_client(mock_ai_project_client)
|
||||
|
||||
messages = [ChatMessage(role=ChatRole.USER, text="Hello")]
|
||||
messages = [ChatMessage(role=Role.USER, text="Hello")]
|
||||
|
||||
run_options, tool_results = chat_client._create_run_options(messages, None) # type: ignore
|
||||
|
||||
@@ -406,7 +406,7 @@ def test_foundry_chat_client_create_run_options_with_image_content(mock_ai_proje
|
||||
chat_client = create_test_foundry_chat_client(mock_ai_project_client, agent_id="test-agent")
|
||||
|
||||
image_content = UriContent(uri="https://example.com/image.jpg", media_type="image/jpeg")
|
||||
messages = [ChatMessage(role=ChatRole.USER, contents=[image_content])]
|
||||
messages = [ChatMessage(role=Role.USER, contents=[image_content])]
|
||||
|
||||
run_options, _ = chat_client._create_run_options(messages, None) # type: ignore
|
||||
|
||||
@@ -502,8 +502,8 @@ def test_foundry_chat_client_create_run_options_with_messages(mock_ai_project_cl
|
||||
|
||||
# Test with system message (becomes instruction)
|
||||
messages = [
|
||||
ChatMessage(role=ChatRole.SYSTEM, text="You are a helpful assistant"),
|
||||
ChatMessage(role=ChatRole.USER, text="Hello"),
|
||||
ChatMessage(role=Role.SYSTEM, text="You are a helpful assistant"),
|
||||
ChatMessage(role=Role.USER, text="Hello"),
|
||||
]
|
||||
|
||||
run_options, _ = chat_client._create_run_options(messages, None) # type: ignore
|
||||
@@ -517,17 +517,17 @@ def test_foundry_chat_client_create_run_options_with_messages(mock_ai_project_cl
|
||||
async def test_foundry_chat_client_inner_get_response(mock_ai_project_client: MagicMock) -> None:
|
||||
"""Test _inner_get_response method."""
|
||||
chat_client = create_test_foundry_chat_client(mock_ai_project_client, agent_id="test-agent")
|
||||
messages = [ChatMessage(role=ChatRole.USER, text="Hello")]
|
||||
messages = [ChatMessage(role=Role.USER, text="Hello")]
|
||||
chat_options = ChatOptions()
|
||||
|
||||
async def mock_streaming_response():
|
||||
yield ChatResponseUpdate(role=ChatRole.ASSISTANT, text="Hello back")
|
||||
yield ChatResponseUpdate(role=Role.ASSISTANT, text="Hello back")
|
||||
|
||||
with (
|
||||
patch.object(chat_client, "_inner_get_streaming_response", return_value=mock_streaming_response()),
|
||||
patch("agent_framework.ChatResponse.from_chat_response_generator") as mock_from_generator,
|
||||
):
|
||||
mock_response = ChatResponse(role=ChatRole.ASSISTANT, text="Hello back")
|
||||
mock_response = ChatResponse(role=Role.ASSISTANT, text="Hello back")
|
||||
mock_from_generator.return_value = mock_response
|
||||
|
||||
result = await chat_client._inner_get_response(messages=messages, chat_options=chat_options) # type: ignore
|
||||
@@ -667,7 +667,7 @@ def get_weather(
|
||||
async def test_foundry_chat_client_get_response() -> None:
|
||||
"""Test Foundry Chat Client response."""
|
||||
async with FoundryChatClient(async_credential=AzureCliCredential()) as foundry_chat_client:
|
||||
assert isinstance(foundry_chat_client, ChatClient)
|
||||
assert isinstance(foundry_chat_client, ChatClientProtocol)
|
||||
|
||||
messages: list[ChatMessage] = []
|
||||
messages.append(
|
||||
@@ -691,7 +691,7 @@ async def test_foundry_chat_client_get_response() -> None:
|
||||
async def test_foundry_chat_client_get_response_tools() -> None:
|
||||
"""Test Foundry Chat Client response with tools."""
|
||||
async with FoundryChatClient(async_credential=AzureCliCredential()) as foundry_chat_client:
|
||||
assert isinstance(foundry_chat_client, ChatClient)
|
||||
assert isinstance(foundry_chat_client, ChatClientProtocol)
|
||||
|
||||
messages: list[ChatMessage] = []
|
||||
messages.append(ChatMessage(role="user", text="What's the weather like in Seattle?"))
|
||||
@@ -712,7 +712,7 @@ async def test_foundry_chat_client_get_response_tools() -> None:
|
||||
async def test_foundry_chat_client_streaming() -> None:
|
||||
"""Test Foundry Chat Client streaming response."""
|
||||
async with FoundryChatClient(async_credential=AzureCliCredential()) as foundry_chat_client:
|
||||
assert isinstance(foundry_chat_client, ChatClient)
|
||||
assert isinstance(foundry_chat_client, ChatClientProtocol)
|
||||
|
||||
messages: list[ChatMessage] = []
|
||||
messages.append(
|
||||
@@ -742,7 +742,7 @@ async def test_foundry_chat_client_streaming() -> None:
|
||||
async def test_foundry_chat_client_streaming_tools() -> None:
|
||||
"""Test Foundry Chat Client streaming response with tools."""
|
||||
async with FoundryChatClient(async_credential=AzureCliCredential()) as foundry_chat_client:
|
||||
assert isinstance(foundry_chat_client, ChatClient)
|
||||
assert isinstance(foundry_chat_client, ChatClientProtocol)
|
||||
|
||||
messages: list[ChatMessage] = []
|
||||
messages.append(ChatMessage(role="user", text="What's the weather like in Seattle?"))
|
||||
@@ -766,8 +766,8 @@ async def test_foundry_chat_client_streaming_tools() -> None:
|
||||
|
||||
@skip_if_foundry_integration_tests_disabled
|
||||
async def test_foundry_chat_client_agent_basic_run() -> None:
|
||||
"""Test ChatClientAgent basic run functionality with FoundryChatClient."""
|
||||
async with ChatClientAgent(
|
||||
"""Test ChatAgent basic run functionality with FoundryChatClient."""
|
||||
async with ChatAgent(
|
||||
chat_client=FoundryChatClient(async_credential=AzureCliCredential()),
|
||||
) as agent:
|
||||
# Run a simple query
|
||||
@@ -782,13 +782,13 @@ async def test_foundry_chat_client_agent_basic_run() -> None:
|
||||
|
||||
@skip_if_foundry_integration_tests_disabled
|
||||
async def test_foundry_chat_client_agent_basic_run_streaming() -> None:
|
||||
"""Test ChatClientAgent basic streaming functionality with FoundryChatClient."""
|
||||
async with ChatClientAgent(
|
||||
"""Test ChatAgent basic streaming functionality with FoundryChatClient."""
|
||||
async with ChatAgent(
|
||||
chat_client=FoundryChatClient(async_credential=AzureCliCredential()),
|
||||
) as agent:
|
||||
# Run streaming query
|
||||
full_message: str = ""
|
||||
async for chunk in agent.run_streaming("Please respond with exactly: 'This is a streaming response test.'"):
|
||||
async for chunk in agent.run_stream("Please respond with exactly: 'This is a streaming response test.'"):
|
||||
assert chunk is not None
|
||||
assert isinstance(chunk, AgentRunResponseUpdate)
|
||||
if chunk.text:
|
||||
@@ -801,8 +801,8 @@ async def test_foundry_chat_client_agent_basic_run_streaming() -> None:
|
||||
|
||||
@skip_if_foundry_integration_tests_disabled
|
||||
async def test_foundry_chat_client_agent_thread_persistence() -> None:
|
||||
"""Test ChatClientAgent thread persistence across runs with FoundryChatClient."""
|
||||
async with ChatClientAgent(
|
||||
"""Test ChatAgent thread persistence across runs with FoundryChatClient."""
|
||||
async with ChatAgent(
|
||||
chat_client=FoundryChatClient(async_credential=AzureCliCredential()),
|
||||
instructions="You are a helpful assistant with good memory.",
|
||||
) as agent:
|
||||
@@ -826,8 +826,8 @@ async def test_foundry_chat_client_agent_thread_persistence() -> None:
|
||||
|
||||
@skip_if_foundry_integration_tests_disabled
|
||||
async def test_foundry_chat_client_agent_existing_thread_id() -> None:
|
||||
"""Test ChatClientAgent existing thread ID functionality with FoundryChatClient."""
|
||||
async with ChatClientAgent(
|
||||
"""Test ChatAgent existing thread ID functionality with FoundryChatClient."""
|
||||
async with ChatAgent(
|
||||
chat_client=FoundryChatClient(async_credential=AzureCliCredential()),
|
||||
instructions="You are a helpful assistant with good memory.",
|
||||
) as first_agent:
|
||||
@@ -844,7 +844,7 @@ async def test_foundry_chat_client_agent_existing_thread_id() -> None:
|
||||
assert existing_thread_id is not None
|
||||
|
||||
# Now continue with the same thread ID in a new agent instance
|
||||
async with ChatClientAgent(
|
||||
async with ChatAgent(
|
||||
chat_client=FoundryChatClient(thread_id=existing_thread_id, async_credential=AzureCliCredential()),
|
||||
instructions="You are a helpful assistant with good memory.",
|
||||
) as second_agent:
|
||||
@@ -863,9 +863,9 @@ async def test_foundry_chat_client_agent_existing_thread_id() -> None:
|
||||
|
||||
@skip_if_foundry_integration_tests_disabled
|
||||
async def test_foundry_chat_client_agent_code_interpreter():
|
||||
"""Test ChatClientAgent with code interpreter through FoundryChatClient."""
|
||||
"""Test ChatAgent with code interpreter through FoundryChatClient."""
|
||||
|
||||
async with ChatClientAgent(
|
||||
async with ChatAgent(
|
||||
chat_client=FoundryChatClient(async_credential=AzureCliCredential()),
|
||||
instructions="You are a helpful assistant that can write and execute Python code.",
|
||||
tools=[HostedCodeInterpreterTool()],
|
||||
@@ -883,7 +883,7 @@ async def test_foundry_chat_client_agent_code_interpreter():
|
||||
@skip_if_foundry_integration_tests_disabled
|
||||
async def test_foundry_chat_client_agent_level_tool_persistence():
|
||||
"""Test that agent-level tools persist across multiple runs with FoundryChatClient."""
|
||||
async with ChatClientAgent(
|
||||
async with ChatAgent(
|
||||
chat_client=FoundryChatClient(async_credential=AzureCliCredential()),
|
||||
instructions="You are a helpful assistant that uses available tools.",
|
||||
tools=[get_weather],
|
||||
@@ -918,7 +918,7 @@ async def test_foundry_chat_client_run_level_tool_isolation():
|
||||
call_count += 1
|
||||
return f"The weather in {location} is sunny and 25°C."
|
||||
|
||||
async with ChatClientAgent(
|
||||
async with ChatAgent(
|
||||
chat_client=FoundryChatClient(async_credential=AzureCliCredential()),
|
||||
instructions="You are a helpful assistant.",
|
||||
) as agent:
|
||||
|
||||
@@ -65,11 +65,11 @@ Create agents and invoke them directly:
|
||||
|
||||
```python
|
||||
import asyncio
|
||||
from agent_framework import ChatClientAgent
|
||||
from agent_framework import ChatAgent
|
||||
from agent_framework.openai import OpenAIChatClient
|
||||
|
||||
async def main():
|
||||
agent = ChatClientAgent(
|
||||
agent = ChatAgent(
|
||||
chat_client=OpenAIChatClient(),
|
||||
instructions="""
|
||||
1) A robot may not injure a human being...
|
||||
@@ -94,14 +94,14 @@ You can use the chat client classes directly for advanced workflows:
|
||||
```python
|
||||
import asyncio
|
||||
from agent_framework.openai import OpenAIChatClient
|
||||
from agent_framework import ChatMessage, ChatRole
|
||||
from agent_framework import ChatMessage, Role
|
||||
|
||||
async def main():
|
||||
client = OpenAIChatClient()
|
||||
|
||||
messages = [
|
||||
ChatMessage(role=ChatRole.SYSTEM, text="You are a helpful assistant."),
|
||||
ChatMessage(role=ChatRole.USER, text="Write a haiku about Agent Framework.")
|
||||
ChatMessage(role=Role.SYSTEM, text="You are a helpful assistant."),
|
||||
ChatMessage(role=Role.USER, text="Write a haiku about Agent Framework.")
|
||||
]
|
||||
|
||||
response = await client.get_response(messages)
|
||||
@@ -127,7 +127,7 @@ import asyncio
|
||||
from typing import Annotated
|
||||
from random import randint
|
||||
from pydantic import Field
|
||||
from agent_framework import ChatClientAgent
|
||||
from agent_framework import ChatAgent
|
||||
from agent_framework.openai import OpenAIChatClient
|
||||
|
||||
|
||||
@@ -149,7 +149,7 @@ def get_menu_specials() -> str:
|
||||
|
||||
|
||||
async def main():
|
||||
agent = ChatClientAgent(
|
||||
agent = ChatAgent(
|
||||
chat_client=OpenAIChatClient(),
|
||||
instructions="You are a helpful assistant that can provide weather and restaurant information.",
|
||||
tools=[get_weather, get_menu_specials]
|
||||
@@ -173,19 +173,19 @@ Coordinate multiple agents to collaborate on complex tasks using orchestration p
|
||||
|
||||
```python
|
||||
import asyncio
|
||||
from agent_framework import ChatClientAgent
|
||||
from agent_framework import ChatAgent
|
||||
from agent_framework.openai import OpenAIChatClient
|
||||
|
||||
|
||||
async def main():
|
||||
# Create specialized agents
|
||||
writer = ChatClientAgent(
|
||||
writer = ChatAgent(
|
||||
chat_client=OpenAIChatClient(),
|
||||
name="Writer",
|
||||
instructions="You are a creative content writer. Generate and refine slogans based on feedback."
|
||||
)
|
||||
|
||||
reviewer = ChatClientAgent(
|
||||
reviewer = ChatAgent(
|
||||
chat_client=OpenAIChatClient(),
|
||||
name="Reviewer",
|
||||
instructions="You are a critical reviewer. Provide detailed feedback on proposed slogans."
|
||||
|
||||
@@ -9,11 +9,11 @@ from uuid import uuid4
|
||||
|
||||
from pydantic import BaseModel, Field, PrivateAttr
|
||||
|
||||
from ._clients import ChatClient
|
||||
from ._mcp import McpTool
|
||||
from ._clients import ChatClientProtocol
|
||||
from ._mcp import MCPTool
|
||||
from ._pydantic import AFBaseModel
|
||||
from ._threads import AgentThread, ChatMessageStore, deserialize_thread_state, thread_on_new_messages
|
||||
from ._tools import AITool
|
||||
from ._tools import ToolProtocol
|
||||
from ._types import (
|
||||
AgentRunResponse,
|
||||
AgentRunResponseUpdate,
|
||||
@@ -21,8 +21,8 @@ from ._types import (
|
||||
ChatOptions,
|
||||
ChatResponse,
|
||||
ChatResponseUpdate,
|
||||
ChatRole,
|
||||
ChatToolMode,
|
||||
Role,
|
||||
)
|
||||
from .exceptions import AgentExecutionException
|
||||
from .telemetry import use_agent_telemetry
|
||||
@@ -34,14 +34,14 @@ else:
|
||||
|
||||
TThreadType = TypeVar("TThreadType", bound="AgentThread")
|
||||
|
||||
__all__ = ["AIAgent", "AgentBase", "ChatClientAgent"]
|
||||
__all__ = ["AgentProtocol", "BaseAgent", "ChatAgent"]
|
||||
|
||||
|
||||
# region Agent Protocol
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class AIAgent(Protocol):
|
||||
class AgentProtocol(Protocol):
|
||||
"""A protocol for an agent that can be invoked."""
|
||||
|
||||
@property
|
||||
@@ -93,7 +93,7 @@ class AIAgent(Protocol):
|
||||
"""
|
||||
...
|
||||
|
||||
def run_streaming(
|
||||
def run_stream(
|
||||
self,
|
||||
messages: str | ChatMessage | list[str] | list[ChatMessage] | None = None,
|
||||
*,
|
||||
@@ -122,10 +122,10 @@ class AIAgent(Protocol):
|
||||
...
|
||||
|
||||
|
||||
# region AgentBase
|
||||
# region BaseAgent
|
||||
|
||||
|
||||
class AgentBase(AFBaseModel):
|
||||
class BaseAgent(AFBaseModel):
|
||||
"""Base class for all Agent Framework agents.
|
||||
|
||||
Attributes:
|
||||
@@ -167,24 +167,24 @@ class AgentBase(AFBaseModel):
|
||||
return thread
|
||||
|
||||
|
||||
# region ChatClientAgent
|
||||
# region ChatAgent
|
||||
|
||||
|
||||
@use_agent_telemetry
|
||||
class ChatClientAgent(AgentBase):
|
||||
class ChatAgent(BaseAgent):
|
||||
"""A Chat Client Agent."""
|
||||
|
||||
AGENT_SYSTEM_NAME: ClassVar[str] = "microsoft.agent_framework"
|
||||
chat_client: ChatClient
|
||||
chat_client: ChatClientProtocol
|
||||
instructions: str | None = None
|
||||
chat_options: ChatOptions
|
||||
chat_message_store_factory: Callable[[], ChatMessageStore] | None = None
|
||||
_local_mcp_tools: list[McpTool] = PrivateAttr(default_factory=list) # type: ignore[reportUnknownVariableType]
|
||||
_local_mcp_tools: list[MCPTool] = PrivateAttr(default_factory=list) # type: ignore[reportUnknownVariableType]
|
||||
_async_exit_stack: AsyncExitStack = PrivateAttr(default_factory=AsyncExitStack)
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
chat_client: ChatClient,
|
||||
chat_client: ChatClientProtocol,
|
||||
instructions: str | None = None,
|
||||
*,
|
||||
id: str | None = None,
|
||||
@@ -202,10 +202,10 @@ class ChatClientAgent(AgentBase):
|
||||
store: bool | None = None,
|
||||
temperature: float | None = None,
|
||||
tool_choice: ChatToolMode | Literal["auto", "required", "none"] | dict[str, Any] | None = "auto",
|
||||
tools: AITool
|
||||
tools: ToolProtocol
|
||||
| Callable[..., Any]
|
||||
| MutableMapping[str, Any]
|
||||
| list[AITool | Callable[..., Any] | MutableMapping[str, Any]]
|
||||
| list[ToolProtocol | Callable[..., Any] | MutableMapping[str, Any]]
|
||||
| None = None,
|
||||
top_p: float | None = None,
|
||||
user: str | None = None,
|
||||
@@ -213,7 +213,7 @@ class ChatClientAgent(AgentBase):
|
||||
chat_message_store_factory: Callable[[], ChatMessageStore] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> None:
|
||||
"""Create a ChatClientAgent.
|
||||
"""Create a ChatAgent.
|
||||
|
||||
Remarks:
|
||||
The set of attributes from frequency_penalty to additional_properties are used to
|
||||
@@ -253,8 +253,8 @@ class ChatClientAgent(AgentBase):
|
||||
# We ignore the MCP Servers here and store them separately,
|
||||
# we add their functions to the tools list at runtime
|
||||
normalized_tools = [] if tools is None else tools if isinstance(tools, list) else [tools]
|
||||
local_mcp_tools = [tool for tool in normalized_tools if isinstance(tool, McpTool)]
|
||||
final_tools = [tool for tool in normalized_tools if not isinstance(tool, McpTool)]
|
||||
local_mcp_tools = [tool for tool in normalized_tools if isinstance(tool, MCPTool)]
|
||||
final_tools = [tool for tool in normalized_tools if not isinstance(tool, MCPTool)]
|
||||
args: dict[str, Any] = {
|
||||
"chat_client": chat_client,
|
||||
"chat_message_store_factory": chat_message_store_factory,
|
||||
@@ -337,8 +337,8 @@ class ChatClientAgent(AgentBase):
|
||||
store: bool | None = None,
|
||||
temperature: float | None = None,
|
||||
tool_choice: ChatToolMode | Literal["auto", "required", "none"] | dict[str, Any] | None = None,
|
||||
tools: AITool
|
||||
| list[AITool]
|
||||
tools: ToolProtocol
|
||||
| list[ToolProtocol]
|
||||
| Callable[..., Any]
|
||||
| list[Callable[..., Any]]
|
||||
| MutableMapping[str, Any]
|
||||
@@ -384,11 +384,11 @@ class ChatClientAgent(AgentBase):
|
||||
agent_name = self._get_agent_name()
|
||||
|
||||
# Resolve final tool list (runtime provided tools + local MCP server tools)
|
||||
final_tools: list[AITool | Callable[..., Any] | dict[str, Any]] = []
|
||||
final_tools: list[ToolProtocol | Callable[..., Any] | dict[str, Any]] = []
|
||||
# Normalize tools argument to a list without mutating the original parameter
|
||||
normalized_tools = [] if tools is None else tools if isinstance(tools, list) else [tools]
|
||||
for tool in normalized_tools:
|
||||
if isinstance(tool, McpTool):
|
||||
if isinstance(tool, MCPTool):
|
||||
final_tools.extend(tool.functions) # type: ignore
|
||||
else:
|
||||
final_tools.append(tool) # type: ignore
|
||||
@@ -442,7 +442,7 @@ class ChatClientAgent(AgentBase):
|
||||
additional_properties=response.additional_properties,
|
||||
)
|
||||
|
||||
async def run_streaming(
|
||||
async def run_stream(
|
||||
self,
|
||||
messages: str | ChatMessage | list[str] | list[ChatMessage] | None = None,
|
||||
*,
|
||||
@@ -459,10 +459,10 @@ class ChatClientAgent(AgentBase):
|
||||
store: bool | None = None,
|
||||
temperature: float | None = None,
|
||||
tool_choice: ChatToolMode | Literal["auto", "required", "none"] | dict[str, Any] | None = None,
|
||||
tools: AITool
|
||||
tools: ToolProtocol
|
||||
| Callable[..., Any]
|
||||
| MutableMapping[str, Any]
|
||||
| list[AITool | Callable[..., Any] | MutableMapping[str, Any]]
|
||||
| list[ToolProtocol | Callable[..., Any] | MutableMapping[str, Any]]
|
||||
| None = None,
|
||||
top_p: float | None = None,
|
||||
user: str | None = None,
|
||||
@@ -472,7 +472,7 @@ class ChatClientAgent(AgentBase):
|
||||
"""Stream the agent with the given messages and options.
|
||||
|
||||
Remarks:
|
||||
Since you won't always call the agent.run_streaming directly, but it get's called
|
||||
Since you won't always call the agent.run_stream directly, but it get's called
|
||||
through orchestration, it is advised to set your default values for
|
||||
all the chat client parameters in the agent constructor.
|
||||
If both parameters are used, the ones passed to the run methods take precedence.
|
||||
@@ -506,11 +506,11 @@ class ChatClientAgent(AgentBase):
|
||||
response_updates: list[ChatResponseUpdate] = []
|
||||
|
||||
# Resolve final tool list (runtime provided tools + local MCP server tools)
|
||||
final_tools: list[AITool | MutableMapping[str, Any] | Callable[..., Any]] = []
|
||||
final_tools: list[ToolProtocol | MutableMapping[str, Any] | Callable[..., Any]] = []
|
||||
# Normalize tools argument to a list without mutating the original parameter
|
||||
normalized_tools = [] if tools is None else tools if isinstance(tools, list) else [tools]
|
||||
for tool in normalized_tools:
|
||||
if isinstance(tool, McpTool):
|
||||
if isinstance(tool, MCPTool):
|
||||
final_tools.extend(tool.functions) # type: ignore
|
||||
else:
|
||||
final_tools.append(tool)
|
||||
@@ -627,7 +627,7 @@ class ChatClientAgent(AgentBase):
|
||||
|
||||
messages: list[ChatMessage] = []
|
||||
if self.instructions:
|
||||
messages.append(ChatMessage(role=ChatRole.SYSTEM, text=self.instructions))
|
||||
messages.append(ChatMessage(role=Role.SYSTEM, text=self.instructions))
|
||||
if thread.message_store:
|
||||
messages.extend(await thread.message_store.list_messages() or [])
|
||||
messages.extend(input_messages or [])
|
||||
@@ -641,12 +641,12 @@ class ChatClientAgent(AgentBase):
|
||||
return []
|
||||
|
||||
if isinstance(messages, str):
|
||||
return [ChatMessage(role=ChatRole.USER, text=messages)]
|
||||
return [ChatMessage(role=Role.USER, text=messages)]
|
||||
|
||||
if isinstance(messages, ChatMessage):
|
||||
return [messages]
|
||||
|
||||
return [ChatMessage(role=ChatRole.USER, text=msg) if isinstance(msg, str) else msg for msg in messages]
|
||||
return [ChatMessage(role=Role.USER, text=msg) if isinstance(msg, str) else msg for msg in messages]
|
||||
|
||||
def _get_agent_name(self) -> str:
|
||||
return self.name or "UnnamedAgent"
|
||||
|
||||
@@ -11,31 +11,31 @@ from pydantic import BaseModel
|
||||
from ._logging import get_logger
|
||||
from ._pydantic import AFBaseModel
|
||||
from ._threads import ChatMessageStore
|
||||
from ._tools import AIFunction, AITool
|
||||
from ._tools import AIFunction, ToolProtocol
|
||||
from ._types import (
|
||||
AIContents,
|
||||
ChatMessage,
|
||||
ChatOptions,
|
||||
ChatResponse,
|
||||
ChatResponseUpdate,
|
||||
ChatToolMode,
|
||||
Contents,
|
||||
FunctionCallContent,
|
||||
FunctionResultContent,
|
||||
GeneratedEmbeddings,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from ._agents import ChatClientAgent
|
||||
from ._agents import ChatAgent
|
||||
|
||||
TInput = TypeVar("TInput", contravariant=True)
|
||||
TEmbedding = TypeVar("TEmbedding")
|
||||
TChatClientBase = TypeVar("TChatClientBase", bound="ChatClientBase")
|
||||
TBaseChatClient = TypeVar("TBaseChatClient", bound="BaseChatClient")
|
||||
|
||||
logger = get_logger()
|
||||
|
||||
__all__ = [
|
||||
"ChatClient",
|
||||
"ChatClientBase",
|
||||
"BaseChatClient",
|
||||
"ChatClientProtocol",
|
||||
"EmbeddingGenerator",
|
||||
"use_tool_calling",
|
||||
]
|
||||
@@ -50,7 +50,7 @@ async def _auto_invoke_function(
|
||||
tool_map: dict[str, AIFunction[BaseModel, Any]],
|
||||
sequence_index: int | None = None,
|
||||
request_index: int | None = None,
|
||||
) -> AIContents:
|
||||
) -> Contents:
|
||||
"""Invoke a function call requested by the agent, applying filters that are defined in the agent."""
|
||||
tool: AIFunction[BaseModel, Any] | None = tool_map.get(function_call_content.name)
|
||||
if tool is None:
|
||||
@@ -81,7 +81,7 @@ def _tool_call_non_streaming(
|
||||
|
||||
@wraps(func)
|
||||
async def wrapper(
|
||||
self: "ChatClientBase",
|
||||
self: "BaseChatClient",
|
||||
*,
|
||||
messages: MutableSequence[ChatMessage],
|
||||
chat_options: ChatOptions,
|
||||
@@ -157,7 +157,7 @@ def _tool_call_streaming(
|
||||
|
||||
@wraps(func)
|
||||
async def wrapper(
|
||||
self: "ChatClientBase",
|
||||
self: "BaseChatClient",
|
||||
*,
|
||||
messages: MutableSequence[ChatMessage],
|
||||
chat_options: ChatOptions,
|
||||
@@ -217,11 +217,11 @@ def _tool_call_streaming(
|
||||
return wrapper
|
||||
|
||||
|
||||
def use_tool_calling(cls: type[TChatClientBase]) -> type[TChatClientBase]:
|
||||
def use_tool_calling(cls: type[TBaseChatClient]) -> type[TBaseChatClient]:
|
||||
"""Class decorator that enables tool calling for a chat client.
|
||||
|
||||
Remarks:
|
||||
This only works on classes that derive from ChatClientBase
|
||||
This only works on classes that derive from BaseChatClient
|
||||
and the `_inner_get_response`
|
||||
and `_inner_get_streaming_response` methods.
|
||||
It also sets a `__maximum_iterations_per_request` attribute on the class.
|
||||
@@ -247,11 +247,11 @@ def use_tool_calling(cls: type[TChatClientBase]) -> type[TChatClientBase]:
|
||||
return cls
|
||||
|
||||
|
||||
# region ChatClient Protocol
|
||||
# region ChatClientProtocol Protocol
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class ChatClient(Protocol):
|
||||
class ChatClientProtocol(Protocol):
|
||||
"""A protocol for a chat client that can generate responses."""
|
||||
|
||||
async def get_response(
|
||||
@@ -270,10 +270,10 @@ class ChatClient(Protocol):
|
||||
store: bool | None = None,
|
||||
temperature: float | None = None,
|
||||
tool_choice: ChatToolMode | Literal["auto", "required", "none"] | dict[str, Any] | None = "auto",
|
||||
tools: AITool
|
||||
tools: ToolProtocol
|
||||
| Callable[..., Any]
|
||||
| MutableMapping[str, Any]
|
||||
| list[AITool | Callable[..., Any] | MutableMapping[str, Any]]
|
||||
| list[ToolProtocol | Callable[..., Any] | MutableMapping[str, Any]]
|
||||
| None = None,
|
||||
top_p: float | None = None,
|
||||
user: str | None = None,
|
||||
@@ -327,10 +327,10 @@ class ChatClient(Protocol):
|
||||
store: bool | None = None,
|
||||
temperature: float | None = None,
|
||||
tool_choice: ChatToolMode | Literal["auto", "required", "none"] | dict[str, Any] | None = "auto",
|
||||
tools: AITool
|
||||
tools: ToolProtocol
|
||||
| Callable[..., Any]
|
||||
| MutableMapping[str, Any]
|
||||
| list[AITool | Callable[..., Any] | MutableMapping[str, Any]]
|
||||
| list[ToolProtocol | Callable[..., Any] | MutableMapping[str, Any]]
|
||||
| None = None,
|
||||
top_p: float | None = None,
|
||||
user: str | None = None,
|
||||
@@ -370,7 +370,7 @@ class ChatClient(Protocol):
|
||||
...
|
||||
|
||||
|
||||
class ChatClientBase(AFBaseModel, ABC):
|
||||
class BaseChatClient(AFBaseModel, ABC):
|
||||
"""Base class for chat clients."""
|
||||
|
||||
MODEL_PROVIDER_NAME: str = "unknown"
|
||||
@@ -457,10 +457,10 @@ class ChatClientBase(AFBaseModel, ABC):
|
||||
store: bool | None = None,
|
||||
temperature: float | None = None,
|
||||
tool_choice: ChatToolMode | Literal["auto", "required", "none"] | dict[str, Any] | None = "auto",
|
||||
tools: AITool
|
||||
tools: ToolProtocol
|
||||
| Callable[..., Any]
|
||||
| MutableMapping[str, Any]
|
||||
| list[AITool | Callable[..., Any] | MutableMapping[str, Any]]
|
||||
| list[ToolProtocol | Callable[..., Any] | MutableMapping[str, Any]]
|
||||
| None = None,
|
||||
top_p: float | None = None,
|
||||
user: str | None = None,
|
||||
@@ -537,10 +537,10 @@ class ChatClientBase(AFBaseModel, ABC):
|
||||
store: bool | None = None,
|
||||
temperature: float | None = None,
|
||||
tool_choice: ChatToolMode | Literal["auto", "required", "none"] | dict[str, Any] | None = "auto",
|
||||
tools: AITool
|
||||
tools: ToolProtocol
|
||||
| Callable[..., Any]
|
||||
| MutableMapping[str, Any]
|
||||
| list[AITool | Callable[..., Any] | MutableMapping[str, Any]]
|
||||
| list[ToolProtocol | Callable[..., Any] | MutableMapping[str, Any]]
|
||||
| None = None,
|
||||
top_p: float | None = None,
|
||||
user: str | None = None,
|
||||
@@ -633,14 +633,14 @@ class ChatClientBase(AFBaseModel, ABC):
|
||||
*,
|
||||
name: str | None = None,
|
||||
instructions: str | None = None,
|
||||
tools: AITool
|
||||
tools: ToolProtocol
|
||||
| Callable[..., Any]
|
||||
| MutableMapping[str, Any]
|
||||
| list[AITool | Callable[..., Any] | MutableMapping[str, Any]]
|
||||
| list[ToolProtocol | Callable[..., Any] | MutableMapping[str, Any]]
|
||||
| None = None,
|
||||
chat_message_store_factory: Callable[[], ChatMessageStore] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> "ChatClientAgent":
|
||||
) -> "ChatAgent":
|
||||
"""Create an agent with the given name and instructions.
|
||||
|
||||
Args:
|
||||
@@ -650,14 +650,14 @@ class ChatClientBase(AFBaseModel, ABC):
|
||||
chat_message_store_factory: Factory function to create an instance of ChatMessageStore. If not provided,
|
||||
the default in-memory store will be used.
|
||||
**kwargs: Additional keyword arguments to pass to the agent.
|
||||
See ChatClientAgent for all the available options.
|
||||
See ChatAgent for all the available options.
|
||||
|
||||
Returns:
|
||||
An instance of ChatClientAgent.
|
||||
An instance of ChatAgent.
|
||||
"""
|
||||
from ._agents import ChatClientAgent
|
||||
from ._agents import ChatAgent
|
||||
|
||||
return ChatClientAgent(
|
||||
return ChatAgent(
|
||||
chat_client=self,
|
||||
name=name,
|
||||
instructions=instructions,
|
||||
|
||||
@@ -22,7 +22,7 @@ from mcp.shared.session import RequestResponder
|
||||
from pydantic import BaseModel, create_model
|
||||
|
||||
from ._tools import AIFunction
|
||||
from ._types import AIContents, ChatMessage, ChatRole, DataContent, TextContent, UriContent
|
||||
from ._types import ChatMessage, Contents, DataContent, Role, TextContent, UriContent
|
||||
from .exceptions import ToolException, ToolExecutionException
|
||||
|
||||
if sys.version_info >= (3, 11):
|
||||
@@ -31,7 +31,7 @@ else:
|
||||
from typing_extensions import Self # pragma: no cover
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from ._clients import ChatClient
|
||||
from ._clients import ChatClientProtocol
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -49,10 +49,10 @@ LOG_LEVEL_MAPPING: dict[types.LoggingLevel, int] = {
|
||||
}
|
||||
|
||||
__all__ = [
|
||||
"McpSseTools",
|
||||
"McpStdioTool",
|
||||
"McpStreamableHttpTool",
|
||||
"McpWebsocketTool",
|
||||
"MCPSseTools",
|
||||
"MCPStdioTool",
|
||||
"MCPStreamableHTTPTool",
|
||||
"MCPWebsocketTool",
|
||||
]
|
||||
|
||||
|
||||
@@ -61,7 +61,7 @@ def _mcp_prompt_message_to_chat_message(
|
||||
) -> ChatMessage:
|
||||
"""Convert a MCP container type to a Agent Framework type."""
|
||||
return ChatMessage(
|
||||
role=ChatRole(value=mcp_type.role),
|
||||
role=Role(value=mcp_type.role),
|
||||
contents=[_mcp_type_to_ai_content(mcp_type.content)], # type: ignore[call-arg]
|
||||
raw_representation=mcp_type,
|
||||
)
|
||||
@@ -69,14 +69,14 @@ def _mcp_prompt_message_to_chat_message(
|
||||
|
||||
def _mcp_call_tool_result_to_ai_contents(
|
||||
mcp_type: types.CallToolResult,
|
||||
) -> list[AIContents]:
|
||||
) -> list[Contents]:
|
||||
"""Convert a MCP container type to a Agent Framework type."""
|
||||
return [_mcp_type_to_ai_content(item) for item in mcp_type.content]
|
||||
|
||||
|
||||
def _mcp_type_to_ai_content(
|
||||
mcp_type: types.ImageContent | types.TextContent | types.AudioContent | types.EmbeddedResource | types.ResourceLink,
|
||||
) -> AIContents:
|
||||
) -> Contents:
|
||||
"""Convert a MCP type to a Agent Framework type."""
|
||||
match mcp_type:
|
||||
case types.TextContent():
|
||||
@@ -105,9 +105,9 @@ def _mcp_type_to_ai_content(
|
||||
|
||||
|
||||
def _ai_content_to_mcp_types(
|
||||
content: AIContents,
|
||||
content: Contents,
|
||||
) -> types.TextContent | types.ImageContent | types.AudioContent | types.EmbeddedResource | types.ResourceLink | None:
|
||||
"""Convert a AIContent type to a MCP type."""
|
||||
"""Convert a BaseContent type to a MCP type."""
|
||||
match content:
|
||||
case TextContent():
|
||||
return types.TextContent(type="text", text=content.text)
|
||||
@@ -223,7 +223,7 @@ def _normalize_mcp_name(name: str) -> str:
|
||||
# region: MCP Plugin
|
||||
|
||||
|
||||
class McpTool:
|
||||
class MCPTool:
|
||||
"""Base class with the MCP logic."""
|
||||
|
||||
def __init__(
|
||||
@@ -235,7 +235,7 @@ class McpTool:
|
||||
load_prompts: bool = True,
|
||||
session: ClientSession | None = None,
|
||||
request_timeout: int | None = None,
|
||||
chat_client: "ChatClient | None" = None,
|
||||
chat_client: "ChatClientProtocol | None" = None,
|
||||
) -> None:
|
||||
"""Initialize the MCP Plugin Base."""
|
||||
self.name = name
|
||||
@@ -250,7 +250,7 @@ class McpTool:
|
||||
self.functions: list[AIFunction[Any, Any]] = []
|
||||
|
||||
def __str__(self) -> str:
|
||||
return f"McpTool(name={self.name}, description={self.description})"
|
||||
return f"MCPTool(name={self.name}, description={self.description})"
|
||||
|
||||
async def connect(self) -> None:
|
||||
"""Connect to the MCP server."""
|
||||
@@ -424,7 +424,7 @@ class McpTool:
|
||||
local_name = _normalize_mcp_name(tool.name)
|
||||
input_model = _get_input_model_from_mcp_tool(tool)
|
||||
# Create AIFunctions out of each tool
|
||||
func: AIFunction[BaseModel, list[AIContents]] = AIFunction(
|
||||
func: AIFunction[BaseModel, list[Contents]] = AIFunction(
|
||||
func=partial(self.call_tool, tool.name),
|
||||
name=local_name,
|
||||
description=tool.description or "",
|
||||
@@ -442,7 +442,7 @@ class McpTool:
|
||||
"""Get an MCP client."""
|
||||
pass
|
||||
|
||||
async def call_tool(self, tool_name: str, **kwargs: Any) -> list[AIContents]:
|
||||
async def call_tool(self, tool_name: str, **kwargs: Any) -> list[Contents]:
|
||||
"""Call a tool with the given arguments."""
|
||||
if not self.session:
|
||||
raise ToolExecutionException("MCP server not connected, please call connect() before using this method.")
|
||||
@@ -494,7 +494,7 @@ class McpTool:
|
||||
# region: MCP Plugin Implementations
|
||||
|
||||
|
||||
class McpStdioTool(McpTool):
|
||||
class MCPStdioTool(MCPTool):
|
||||
"""MCP stdio server configuration."""
|
||||
|
||||
def __init__(
|
||||
@@ -511,7 +511,7 @@ class McpStdioTool(McpTool):
|
||||
args: list[str] | None = None,
|
||||
env: dict[str, str] | None = None,
|
||||
encoding: str | None = None,
|
||||
chat_client: "ChatClient | None" = None,
|
||||
chat_client: "ChatClientProtocol | None" = None,
|
||||
**kwargs: Any,
|
||||
) -> None:
|
||||
"""Initialize the MCP stdio plugin.
|
||||
@@ -567,7 +567,7 @@ class McpStdioTool(McpTool):
|
||||
return stdio_client(server=StdioServerParameters(**args))
|
||||
|
||||
|
||||
class McpSseTools(McpTool):
|
||||
class MCPSseTools(MCPTool):
|
||||
"""MCP sse server configuration."""
|
||||
|
||||
def __init__(
|
||||
@@ -584,7 +584,7 @@ class McpSseTools(McpTool):
|
||||
headers: dict[str, Any] | None = None,
|
||||
timeout: float | None = None,
|
||||
sse_read_timeout: float | None = None,
|
||||
chat_client: "ChatClient | None" = None,
|
||||
chat_client: "ChatClientProtocol | None" = None,
|
||||
**kwargs: Any,
|
||||
) -> None:
|
||||
"""Initialize the MCP sse plugin.
|
||||
@@ -643,7 +643,7 @@ class McpSseTools(McpTool):
|
||||
return sse_client(**args)
|
||||
|
||||
|
||||
class McpStreamableHttpTool(McpTool):
|
||||
class MCPStreamableHTTPTool(MCPTool):
|
||||
"""MCP streamable http server configuration."""
|
||||
|
||||
def __init__(
|
||||
@@ -661,7 +661,7 @@ class McpStreamableHttpTool(McpTool):
|
||||
timeout: float | None = None,
|
||||
sse_read_timeout: float | None = None,
|
||||
terminate_on_close: bool | None = None,
|
||||
chat_client: "ChatClient | None" = None,
|
||||
chat_client: "ChatClientProtocol | None" = None,
|
||||
**kwargs: Any,
|
||||
) -> None:
|
||||
"""Initialize the MCP streamable http plugin.
|
||||
@@ -723,7 +723,7 @@ class McpStreamableHttpTool(McpTool):
|
||||
return streamablehttp_client(**args)
|
||||
|
||||
|
||||
class McpWebsocketTool(McpTool):
|
||||
class MCPWebsocketTool(MCPTool):
|
||||
"""MCP websocket server configuration."""
|
||||
|
||||
def __init__(
|
||||
@@ -737,7 +737,7 @@ class McpWebsocketTool(McpTool):
|
||||
session: ClientSession | None = None,
|
||||
description: str | None = None,
|
||||
additional_properties: dict[str, Any] | None = None,
|
||||
chat_client: "ChatClient | None" = None,
|
||||
chat_client: "ChatClientProtocol | None" = None,
|
||||
**kwargs: Any,
|
||||
) -> None:
|
||||
"""Initialize the MCP websocket plugin.
|
||||
|
||||
@@ -7,9 +7,9 @@ from pydantic import BaseModel, ConfigDict, Field, UrlConstraints
|
||||
from pydantic.networks import AnyUrl
|
||||
from pydantic_settings import BaseSettings, SettingsConfigDict
|
||||
|
||||
HttpsUrl = Annotated[AnyUrl, UrlConstraints(max_length=2083, allowed_schemes=["https"])]
|
||||
HTTPsUrl = Annotated[AnyUrl, UrlConstraints(max_length=2083, allowed_schemes=["https"])]
|
||||
|
||||
__all__ = ["AFBaseModel", "AFBaseSettings", "HttpsUrl"]
|
||||
__all__ = ["AFBaseModel", "AFBaseSettings", "HTTPsUrl"]
|
||||
|
||||
|
||||
class AFBaseModel(BaseModel):
|
||||
|
||||
@@ -24,7 +24,7 @@ from ._pydantic import AFBaseModel
|
||||
from .telemetry import GenAIAttributes, start_as_current_span
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from ._types import AIContents
|
||||
from ._types import Contents
|
||||
|
||||
tracer: trace.Tracer = trace.get_tracer("agent_framework")
|
||||
meter: metrics.Meter = metrics.get_meter_provider().get_meter("agent_framework")
|
||||
@@ -32,24 +32,24 @@ logger = get_logger()
|
||||
|
||||
__all__ = [
|
||||
"AIFunction",
|
||||
"AITool",
|
||||
"HostedCodeInterpreterTool",
|
||||
"HostedFileSearchTool",
|
||||
"HostedWebSearchTool",
|
||||
"ToolProtocol",
|
||||
"ai_function",
|
||||
]
|
||||
|
||||
|
||||
def _parse_inputs(
|
||||
inputs: "AIContents | dict[str, Any] | str | list[AIContents | dict[str, Any] | str] | None",
|
||||
) -> list["AIContents"]:
|
||||
"""Parse the inputs for a tool, ensuring they are of type AIContents."""
|
||||
inputs: "Contents | dict[str, Any] | str | list[Contents | dict[str, Any] | str] | None",
|
||||
) -> list["Contents"]:
|
||||
"""Parse the inputs for a tool, ensuring they are of type Contents."""
|
||||
if inputs is None:
|
||||
return []
|
||||
|
||||
from ._types import AIContent, DataContent, HostedFileContent, HostedVectorStoreContent, UriContent
|
||||
from ._types import BaseContent, DataContent, HostedFileContent, HostedVectorStoreContent, UriContent
|
||||
|
||||
parsed_inputs: list["AIContents"] = []
|
||||
parsed_inputs: list["Contents"] = []
|
||||
if not isinstance(inputs, list):
|
||||
inputs = [inputs]
|
||||
for input_item in inputs:
|
||||
@@ -75,15 +75,15 @@ def _parse_inputs(
|
||||
parsed_inputs.append(DataContent(**input_item))
|
||||
else:
|
||||
raise ValueError(f"Unsupported input type: {input_item}")
|
||||
elif isinstance(input_item, AIContent):
|
||||
elif isinstance(input_item, BaseContent):
|
||||
parsed_inputs.append(input_item)
|
||||
else:
|
||||
raise TypeError(f"Unsupported input type: {type(input_item).__name__}. Expected AIContents or dict.")
|
||||
raise TypeError(f"Unsupported input type: {type(input_item).__name__}. Expected Contents or dict.")
|
||||
return parsed_inputs
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class AITool(Protocol):
|
||||
class ToolProtocol(Protocol):
|
||||
"""Represents a generic tool that can be specified to an AI service.
|
||||
|
||||
Attributes:
|
||||
@@ -111,7 +111,7 @@ ArgsT = TypeVar("ArgsT", bound=BaseModel)
|
||||
ReturnT = TypeVar("ReturnT")
|
||||
|
||||
|
||||
class AIToolBase(AFBaseModel):
|
||||
class BaseTool(AFBaseModel):
|
||||
"""Base class for AI tools, providing common attributes and methods.
|
||||
|
||||
Args:
|
||||
@@ -131,7 +131,7 @@ class AIToolBase(AFBaseModel):
|
||||
return f"{self.__class__.__name__}(name={self.name})"
|
||||
|
||||
|
||||
class HostedCodeInterpreterTool(AIToolBase):
|
||||
class HostedCodeInterpreterTool(BaseTool):
|
||||
"""Represents a hosted tool that can be specified to an AI service to enable it to execute generated code.
|
||||
|
||||
This tool does not implement code interpretation itself. It serves as a marker to inform a service
|
||||
@@ -143,7 +143,7 @@ class HostedCodeInterpreterTool(AIToolBase):
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
inputs: "AIContents | dict[str, Any] | str | list[AIContents | dict[str, Any] | str] | None" = None,
|
||||
inputs: "Contents | dict[str, Any] | str | list[Contents | dict[str, Any] | str] | None" = None,
|
||||
description: str | None = None,
|
||||
additional_properties: dict[str, Any] | None = None,
|
||||
**kwargs: Any,
|
||||
@@ -155,8 +155,8 @@ class HostedCodeInterpreterTool(AIToolBase):
|
||||
This should mostly be HostedFileContent or HostedVectorStoreContent.
|
||||
Can also be DataContent, depending on the service used.
|
||||
When supplying a list, it can contain:
|
||||
- AIContents instances
|
||||
- dicts with properties for AIContents (e.g., {"uri": "http://example.com", "media_type": "text/html"})
|
||||
- Contents instances
|
||||
- dicts with properties for Contents (e.g., {"uri": "http://example.com", "media_type": "text/html"})
|
||||
- strings (which will be converted to UriContent with media_type "text/plain").
|
||||
If None, defaults to an empty list.
|
||||
description: A description of the tool.
|
||||
@@ -177,7 +177,7 @@ class HostedCodeInterpreterTool(AIToolBase):
|
||||
super().__init__(**args, **kwargs)
|
||||
|
||||
|
||||
class HostedWebSearchTool(AIToolBase):
|
||||
class HostedWebSearchTool(BaseTool):
|
||||
"""Represents a web search tool that can be specified to an AI service to enable it to perform web searches."""
|
||||
|
||||
def __init__(
|
||||
@@ -206,7 +206,7 @@ class HostedWebSearchTool(AIToolBase):
|
||||
super().__init__(**args, **kwargs)
|
||||
|
||||
|
||||
class HostedFileSearchTool(AIToolBase):
|
||||
class HostedFileSearchTool(BaseTool):
|
||||
"""Represents a file search tool that can be specified to an AI service to enable it to perform file searches."""
|
||||
|
||||
inputs: list[Any] | None = None
|
||||
@@ -214,7 +214,7 @@ class HostedFileSearchTool(AIToolBase):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
inputs: "AIContents | dict[str, Any] | str | list[AIContents | dict[str, Any] | str] | None" = None,
|
||||
inputs: "Contents | dict[str, Any] | str | list[Contents | dict[str, Any] | str] | None" = None,
|
||||
max_results: int | None = None,
|
||||
description: str | None = None,
|
||||
additional_properties: dict[str, Any] | None = None,
|
||||
@@ -226,8 +226,8 @@ class HostedFileSearchTool(AIToolBase):
|
||||
inputs: A list of contents that the tool can accept as input. Defaults to None.
|
||||
This should be one or more HostedVectorStoreContents.
|
||||
When supplying a list, it can contain:
|
||||
- AIContents instances
|
||||
- dicts with properties for AIContents (e.g., {"uri": "http://example.com", "media_type": "text/html"})
|
||||
- Contents instances
|
||||
- dicts with properties for Contents (e.g., {"uri": "http://example.com", "media_type": "text/html"})
|
||||
- strings (which will be converted to UriContent with media_type "text/plain").
|
||||
If None, defaults to an empty list.
|
||||
max_results: The maximum number of results to return from the file search.
|
||||
@@ -252,8 +252,8 @@ class HostedFileSearchTool(AIToolBase):
|
||||
super().__init__(**args, **kwargs)
|
||||
|
||||
|
||||
class AIFunction(AIToolBase, Generic[ArgsT, ReturnT]):
|
||||
"""A AITool that is callable as code.
|
||||
class AIFunction(BaseTool, Generic[ArgsT, ReturnT]):
|
||||
"""A ToolProtocol that is callable as code.
|
||||
|
||||
Args:
|
||||
name: The name of the function.
|
||||
|
||||
@@ -28,7 +28,7 @@ from pydantic import (
|
||||
|
||||
from ._logging import get_logger
|
||||
from ._pydantic import AFBaseModel
|
||||
from ._tools import AITool, ai_function
|
||||
from ._tools import ToolProtocol, ai_function
|
||||
from .exceptions import AgentFrameworkException
|
||||
|
||||
if sys.version_info >= (3, 11):
|
||||
@@ -77,29 +77,28 @@ KNOWN_MEDIA_TYPES = [
|
||||
|
||||
|
||||
__all__ = [
|
||||
"AIAnnotation",
|
||||
"AIAnnotations",
|
||||
"AIContent",
|
||||
"AIContents",
|
||||
"AgentRunResponse",
|
||||
"AgentRunResponseUpdate",
|
||||
"AnnotatedRegion",
|
||||
"AnnotatedRegions",
|
||||
"ChatFinishReason",
|
||||
"Annotations",
|
||||
"BaseAnnotation",
|
||||
"BaseContent",
|
||||
"ChatMessage",
|
||||
"ChatOptions",
|
||||
"ChatResponse",
|
||||
"ChatResponseUpdate",
|
||||
"ChatRole",
|
||||
"ChatToolMode",
|
||||
"CitationAnnotation",
|
||||
"Contents",
|
||||
"DataContent",
|
||||
"ErrorContent",
|
||||
"FinishReason",
|
||||
"FunctionCallContent",
|
||||
"FunctionResultContent",
|
||||
"GeneratedEmbeddings",
|
||||
"HostedFileContent",
|
||||
"HostedVectorStoreContent",
|
||||
"Role",
|
||||
"SpeechToTextOptions",
|
||||
"TextContent",
|
||||
"TextReasoningContent",
|
||||
@@ -231,7 +230,7 @@ def _process_update(
|
||||
is_new_message = True
|
||||
|
||||
if is_new_message:
|
||||
message = ChatMessage(role=ChatRole.ASSISTANT, contents=[])
|
||||
message = ChatMessage(role=Role.ASSISTANT, contents=[])
|
||||
response.messages.append(message)
|
||||
else:
|
||||
message = response.messages[-1]
|
||||
@@ -278,12 +277,12 @@ def _process_update(
|
||||
|
||||
|
||||
def _coalesce_text_content(
|
||||
contents: list["AIContents"], type_: type["TextContent"] | type["TextReasoningContent"]
|
||||
contents: list["Contents"], type_: type["TextContent"] | type["TextReasoningContent"]
|
||||
) -> None:
|
||||
"""Take any subsequence Text or TextReasoningContent items and coalesce them into a single item."""
|
||||
if not contents:
|
||||
return
|
||||
coalesced_contents: list["AIContents"] = []
|
||||
coalesced_contents: list["Contents"] = []
|
||||
first_new_content: Any | None = None
|
||||
for content in contents:
|
||||
if isinstance(content, type_):
|
||||
@@ -313,22 +312,10 @@ def _finalize_response(response: "ChatResponse | AgentRunResponse") -> None:
|
||||
_coalesce_text_content(msg.contents, TextReasoningContent)
|
||||
|
||||
|
||||
# region AIAnnotation
|
||||
# region BaseAnnotation
|
||||
|
||||
|
||||
class AnnotatedRegion(AFBaseModel):
|
||||
"""Represents a collection of annotated regions.
|
||||
|
||||
Attributes:
|
||||
regions: A list of regions that have been annotated.
|
||||
additional_properties: Optional additional properties associated with the content.
|
||||
raw_representation: Optional raw representation of the content from an underlying implementation.
|
||||
"""
|
||||
|
||||
type: Literal["annotated_regions"] = "annotated_regions" # type: ignore[assignment]
|
||||
|
||||
|
||||
class TextSpanRegion(AnnotatedRegion):
|
||||
class TextSpanRegion(AFBaseModel):
|
||||
"""Represents a region of text that has been annotated."""
|
||||
|
||||
type: Literal["text_span"] = "text_span" # type: ignore[assignment]
|
||||
@@ -337,28 +324,26 @@ class TextSpanRegion(AnnotatedRegion):
|
||||
|
||||
|
||||
AnnotatedRegions = Annotated[
|
||||
TextSpanRegion | AnnotatedRegion,
|
||||
TextSpanRegion,
|
||||
Field(discriminator="type"),
|
||||
]
|
||||
|
||||
|
||||
class AIAnnotation(AFBaseModel):
|
||||
class BaseAnnotation(AFBaseModel):
|
||||
"""Base class for all AI Annotation types.
|
||||
|
||||
Args:
|
||||
type: The type of content, which is always "ai_annotation" for this class.
|
||||
additional_properties: Optional additional properties associated with the content.
|
||||
raw_representation: Optional raw representation of the content from an underlying implementation.
|
||||
|
||||
"""
|
||||
|
||||
type: Literal["ai_annotation"] = "ai_annotation"
|
||||
annotated_regions: list[AnnotatedRegions] | None = None
|
||||
additional_properties: dict[str, Any] | None = None
|
||||
raw_representation: Any | None = Field(default=None, repr=False)
|
||||
|
||||
|
||||
class CitationAnnotation(AIAnnotation):
|
||||
class CitationAnnotation(BaseAnnotation):
|
||||
"""Represents a citation annotation.
|
||||
|
||||
Attributes:
|
||||
@@ -381,33 +366,31 @@ class CitationAnnotation(AIAnnotation):
|
||||
snippet: str | None = None
|
||||
|
||||
|
||||
AIAnnotations = Annotated[
|
||||
CitationAnnotation | AIAnnotation,
|
||||
Annotations = Annotated[
|
||||
CitationAnnotation,
|
||||
Field(discriminator="type"),
|
||||
]
|
||||
|
||||
|
||||
# region AIContent
|
||||
# region BaseContent
|
||||
|
||||
|
||||
class AIContent(AFBaseModel):
|
||||
class BaseContent(AFBaseModel):
|
||||
"""Represents content used by AI services.
|
||||
|
||||
Attributes:
|
||||
type: The type of content, which is always "ai" for this class.
|
||||
annotations: Optional annotations associated with the content.
|
||||
additional_properties: Optional additional properties associated with the content.
|
||||
raw_representation: Optional raw representation of the content from an underlying implementation.
|
||||
|
||||
"""
|
||||
|
||||
type: Literal["ai"] = "ai"
|
||||
annotations: list[AIAnnotations] | None = None
|
||||
annotations: list[Annotations] | None = None
|
||||
additional_properties: dict[str, Any] | None = None
|
||||
raw_representation: Any | None = Field(default=None, repr=False, exclude=True)
|
||||
|
||||
|
||||
class TextContent(AIContent):
|
||||
class TextContent(BaseContent):
|
||||
"""Represents text content in a chat.
|
||||
|
||||
Attributes:
|
||||
@@ -508,7 +491,7 @@ class TextContent(AIContent):
|
||||
return self
|
||||
|
||||
|
||||
class TextReasoningContent(AIContent):
|
||||
class TextReasoningContent(BaseContent):
|
||||
"""Represents text reasoning content in a chat.
|
||||
|
||||
Remarks:
|
||||
@@ -609,7 +592,7 @@ class TextReasoningContent(AIContent):
|
||||
return self
|
||||
|
||||
|
||||
class DataContent(AIContent):
|
||||
class DataContent(BaseContent):
|
||||
"""Represents binary data content with an associated media type (also known as a MIME type).
|
||||
|
||||
Attributes:
|
||||
@@ -632,7 +615,7 @@ class DataContent(AIContent):
|
||||
self,
|
||||
*,
|
||||
uri: str,
|
||||
annotations: list[AIAnnotations] | None = None,
|
||||
annotations: list[Annotations] | None = None,
|
||||
additional_properties: dict[str, Any] | None = None,
|
||||
raw_representation: Any | None = None,
|
||||
**kwargs: Any,
|
||||
@@ -658,7 +641,7 @@ class DataContent(AIContent):
|
||||
*,
|
||||
data: bytes,
|
||||
media_type: str,
|
||||
annotations: list[AIAnnotations] | None = None,
|
||||
annotations: list[Annotations] | None = None,
|
||||
additional_properties: dict[str, Any] | None = None,
|
||||
raw_representation: Any | None = None,
|
||||
**kwargs: Any,
|
||||
@@ -685,7 +668,7 @@ class DataContent(AIContent):
|
||||
uri: str | None = None,
|
||||
data: bytes | None = None,
|
||||
media_type: str | None = None,
|
||||
annotations: list[AIAnnotations] | None = None,
|
||||
annotations: list[Annotations] | None = None,
|
||||
additional_properties: dict[str, Any] | None = None,
|
||||
raw_representation: Any | None = None,
|
||||
**kwargs: Any,
|
||||
@@ -739,7 +722,7 @@ class DataContent(AIContent):
|
||||
return _has_top_level_media_type(self.media_type, top_level_media_type)
|
||||
|
||||
|
||||
class UriContent(AIContent):
|
||||
class UriContent(BaseContent):
|
||||
"""Represents a URI content.
|
||||
|
||||
Remarks:
|
||||
@@ -765,7 +748,7 @@ class UriContent(AIContent):
|
||||
uri: str,
|
||||
media_type: str,
|
||||
*,
|
||||
annotations: list[AIAnnotations] | None = None,
|
||||
annotations: list[Annotations] | None = None,
|
||||
additional_properties: dict[str, Any] | None = None,
|
||||
raw_representation: Any | None = None,
|
||||
**kwargs: Any,
|
||||
@@ -807,7 +790,7 @@ def _has_top_level_media_type(media_type: str | None, top_level_media_type: str)
|
||||
return span.lower() == top_level_media_type.lower()
|
||||
|
||||
|
||||
class ErrorContent(AIContent):
|
||||
class ErrorContent(BaseContent):
|
||||
"""Represents an error.
|
||||
|
||||
Remarks:
|
||||
@@ -837,7 +820,7 @@ class ErrorContent(AIContent):
|
||||
message: str | None = None,
|
||||
error_code: str | None = None,
|
||||
details: str | None = None,
|
||||
annotations: list[AIAnnotations] | None = None,
|
||||
annotations: list[Annotations] | None = None,
|
||||
additional_properties: dict[str, Any] | None = None,
|
||||
raw_representation: Any | None = None,
|
||||
**kwargs: Any,
|
||||
@@ -868,7 +851,7 @@ class ErrorContent(AIContent):
|
||||
return f"Error {self.error_code}: {self.message}" if self.error_code else self.message or "Unknown error"
|
||||
|
||||
|
||||
class FunctionCallContent(AIContent):
|
||||
class FunctionCallContent(BaseContent):
|
||||
"""Represents a function call request.
|
||||
|
||||
Attributes:
|
||||
@@ -896,7 +879,7 @@ class FunctionCallContent(AIContent):
|
||||
name: str,
|
||||
arguments: str | dict[str, Any | None] | None = None,
|
||||
exception: Exception | None = None,
|
||||
annotations: list[AIAnnotations] | None = None,
|
||||
annotations: list[Annotations] | None = None,
|
||||
additional_properties: dict[str, Any] | None = None,
|
||||
raw_representation: Any | None = None,
|
||||
**kwargs: Any,
|
||||
@@ -962,7 +945,7 @@ class FunctionCallContent(AIContent):
|
||||
)
|
||||
|
||||
|
||||
class FunctionResultContent(AIContent):
|
||||
class FunctionResultContent(BaseContent):
|
||||
"""Represents the result of a function call.
|
||||
|
||||
Attributes:
|
||||
@@ -987,7 +970,7 @@ class FunctionResultContent(AIContent):
|
||||
call_id: str,
|
||||
result: Any | None = None,
|
||||
exception: Exception | None = None,
|
||||
annotations: list[AIAnnotations] | None = None,
|
||||
annotations: list[Annotations] | None = None,
|
||||
additional_properties: dict[str, Any] | None = None,
|
||||
raw_representation: Any | None = None,
|
||||
**kwargs: Any,
|
||||
@@ -1014,7 +997,7 @@ class FunctionResultContent(AIContent):
|
||||
)
|
||||
|
||||
|
||||
class UsageContent(AIContent):
|
||||
class UsageContent(BaseContent):
|
||||
"""Represents usage information associated with a chat request and response.
|
||||
|
||||
Attributes:
|
||||
@@ -1033,7 +1016,7 @@ class UsageContent(AIContent):
|
||||
self,
|
||||
details: UsageDetails,
|
||||
*,
|
||||
annotations: list[AIAnnotations] | None = None,
|
||||
annotations: list[Annotations] | None = None,
|
||||
additional_properties: dict[str, Any] | None = None,
|
||||
raw_representation: Any | None = None,
|
||||
**kwargs: Any,
|
||||
@@ -1048,7 +1031,7 @@ class UsageContent(AIContent):
|
||||
)
|
||||
|
||||
|
||||
class HostedFileContent(AIContent):
|
||||
class HostedFileContent(BaseContent):
|
||||
"""Represents a hosted file content.
|
||||
|
||||
Attributes:
|
||||
@@ -1079,7 +1062,7 @@ class HostedFileContent(AIContent):
|
||||
)
|
||||
|
||||
|
||||
class HostedVectorStoreContent(AIContent):
|
||||
class HostedVectorStoreContent(BaseContent):
|
||||
"""Represents a hosted vector store content.
|
||||
|
||||
Attributes:
|
||||
@@ -1110,7 +1093,7 @@ class HostedVectorStoreContent(AIContent):
|
||||
)
|
||||
|
||||
|
||||
AIContents = Annotated[
|
||||
Contents = Annotated[
|
||||
TextContent
|
||||
| DataContent
|
||||
| TextReasoningContent
|
||||
@@ -1127,7 +1110,7 @@ AIContents = Annotated[
|
||||
# region Chat Response constants
|
||||
|
||||
|
||||
class ChatRole(AFBaseModel):
|
||||
class Role(AFBaseModel):
|
||||
"""Describes the intended purpose of a message within a chat interaction.
|
||||
|
||||
Attributes:
|
||||
@@ -1157,19 +1140,19 @@ class ChatRole(AFBaseModel):
|
||||
|
||||
def __repr__(self) -> str:
|
||||
"""Returns the string representation of the role."""
|
||||
return f"ChatRole(value={self.value!r})"
|
||||
return f"Role(value={self.value!r})"
|
||||
|
||||
|
||||
# Note: ClassVar is used to indicate that these are class-level constants, not instance attributes.
|
||||
# The type: ignore[assignment] is used to suppress the type checker warning about assigning to a ClassVar,
|
||||
# it gets assigned immediately after the class definition.
|
||||
ChatRole.SYSTEM = ChatRole(value="system") # type: ignore[assignment]
|
||||
ChatRole.USER = ChatRole(value="user") # type: ignore[assignment]
|
||||
ChatRole.ASSISTANT = ChatRole(value="assistant") # type: ignore[assignment]
|
||||
ChatRole.TOOL = ChatRole(value="tool") # type: ignore[assignment]
|
||||
Role.SYSTEM = Role(value="system") # type: ignore[assignment]
|
||||
Role.USER = Role(value="user") # type: ignore[assignment]
|
||||
Role.ASSISTANT = Role(value="assistant") # type: ignore[assignment]
|
||||
Role.TOOL = Role(value="tool") # type: ignore[assignment]
|
||||
|
||||
|
||||
class ChatFinishReason(AFBaseModel):
|
||||
class FinishReason(AFBaseModel):
|
||||
"""Represents the reason a chat response completed.
|
||||
|
||||
Attributes:
|
||||
@@ -1179,21 +1162,21 @@ class ChatFinishReason(AFBaseModel):
|
||||
value: str
|
||||
|
||||
CONTENT_FILTER: ClassVar[Self] # type: ignore[assignment]
|
||||
"""A ChatFinishReason representing the model filtering content, whether for safety, prohibited content,
|
||||
"""A FinishReason representing the model filtering content, whether for safety, prohibited content,
|
||||
sensitive content, or other such issues."""
|
||||
LENGTH: ClassVar[Self] # type: ignore[assignment]
|
||||
"""A ChatFinishReason representing the model reaching the maximum length allowed for the request and/or
|
||||
"""A FinishReason representing the model reaching the maximum length allowed for the request and/or
|
||||
response (typically in terms of tokens)."""
|
||||
STOP: ClassVar[Self] # type: ignore[assignment]
|
||||
"""A ChatFinishReason representing the model encountering a natural stop point or provided stop sequence."""
|
||||
"""A FinishReason representing the model encountering a natural stop point or provided stop sequence."""
|
||||
TOOL_CALLS: ClassVar[Self] # type: ignore[assignment]
|
||||
"""A ChatFinishReason representing the model requesting the use of a tool that was defined in the request."""
|
||||
"""A FinishReason representing the model requesting the use of a tool that was defined in the request."""
|
||||
|
||||
|
||||
ChatFinishReason.CONTENT_FILTER = ChatFinishReason(value="content_filter") # type: ignore[assignment]
|
||||
ChatFinishReason.LENGTH = ChatFinishReason(value="length") # type: ignore[assignment]
|
||||
ChatFinishReason.STOP = ChatFinishReason(value="stop") # type: ignore[assignment]
|
||||
ChatFinishReason.TOOL_CALLS = ChatFinishReason(value="tool_calls") # type: ignore[assignment]
|
||||
FinishReason.CONTENT_FILTER = FinishReason(value="content_filter") # type: ignore[assignment]
|
||||
FinishReason.LENGTH = FinishReason(value="length") # type: ignore[assignment]
|
||||
FinishReason.STOP = FinishReason(value="stop") # type: ignore[assignment]
|
||||
FinishReason.TOOL_CALLS = FinishReason(value="tool_calls") # type: ignore[assignment]
|
||||
|
||||
# region ChatMessage
|
||||
|
||||
@@ -1211,9 +1194,9 @@ class ChatMessage(AFBaseModel):
|
||||
|
||||
"""
|
||||
|
||||
role: ChatRole
|
||||
role: Role
|
||||
"""The role of the author of the message."""
|
||||
contents: list[AIContents]
|
||||
contents: list[Contents]
|
||||
"""The chat message content items."""
|
||||
author_name: str | None
|
||||
"""The name of the author of the message."""
|
||||
@@ -1227,7 +1210,7 @@ class ChatMessage(AFBaseModel):
|
||||
@overload
|
||||
def __init__(
|
||||
self,
|
||||
role: ChatRole | Literal["system", "user", "assistant", "tool"],
|
||||
role: Role | Literal["system", "user", "assistant", "tool"],
|
||||
*,
|
||||
text: str,
|
||||
author_name: str | None = None,
|
||||
@@ -1249,9 +1232,9 @@ class ChatMessage(AFBaseModel):
|
||||
@overload
|
||||
def __init__(
|
||||
self,
|
||||
role: ChatRole | Literal["system", "user", "assistant", "tool"],
|
||||
role: Role | Literal["system", "user", "assistant", "tool"],
|
||||
*,
|
||||
contents: MutableSequence[AIContents],
|
||||
contents: MutableSequence[Contents],
|
||||
author_name: str | None = None,
|
||||
message_id: str | None = None,
|
||||
additional_properties: dict[str, Any] | None = None,
|
||||
@@ -1261,7 +1244,7 @@ class ChatMessage(AFBaseModel):
|
||||
|
||||
Args:
|
||||
role: The role of the author of the message.
|
||||
contents: Optional list of AIContent items to include in the message.
|
||||
contents: Optional list of BaseContent items to include in the message.
|
||||
author_name: Optional name of the author of the message.
|
||||
message_id: Optional ID of the chat message.
|
||||
additional_properties: Optional additional properties associated with the chat message.
|
||||
@@ -1270,10 +1253,10 @@ class ChatMessage(AFBaseModel):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
role: ChatRole | Literal["system", "user", "assistant", "tool"],
|
||||
role: Role | Literal["system", "user", "assistant", "tool"],
|
||||
*,
|
||||
text: str | None = None,
|
||||
contents: MutableSequence[AIContents] | None = None,
|
||||
contents: MutableSequence[Contents] | None = None,
|
||||
author_name: str | None = None,
|
||||
message_id: str | None = None,
|
||||
additional_properties: dict[str, Any] | None = None,
|
||||
@@ -1284,7 +1267,7 @@ class ChatMessage(AFBaseModel):
|
||||
if text is not None:
|
||||
contents.append(TextContent(text=text))
|
||||
if isinstance(role, str):
|
||||
role = ChatRole(value=role)
|
||||
role = Role(value=role)
|
||||
super().__init__(
|
||||
role=role, # type: ignore[reportCallIssue]
|
||||
contents=contents, # type: ignore[reportCallIssue]
|
||||
@@ -1334,7 +1317,7 @@ class ChatResponse(AFBaseModel):
|
||||
"""The model ID used in the creation of the chat response."""
|
||||
created_at: CreatedAtT | None = None # use a datetimeoffset type?
|
||||
"""A timestamp for the chat response."""
|
||||
finish_reason: ChatFinishReason | None = None
|
||||
finish_reason: FinishReason | None = None
|
||||
"""The reason for the chat response."""
|
||||
usage_details: UsageDetails | None = None
|
||||
"""The usage details for the chat response."""
|
||||
@@ -1354,7 +1337,7 @@ class ChatResponse(AFBaseModel):
|
||||
conversation_id: str | None = None,
|
||||
model_id: str | None = None,
|
||||
created_at: CreatedAtT | None = None,
|
||||
finish_reason: ChatFinishReason | None = None,
|
||||
finish_reason: FinishReason | None = None,
|
||||
usage_details: UsageDetails | None = None,
|
||||
value: Any | None = None,
|
||||
response_format: type[BaseModel] | None = None,
|
||||
@@ -1389,7 +1372,7 @@ class ChatResponse(AFBaseModel):
|
||||
conversation_id: str | None = None,
|
||||
model_id: str | None = None,
|
||||
created_at: CreatedAtT | None = None,
|
||||
finish_reason: ChatFinishReason | None = None,
|
||||
finish_reason: FinishReason | None = None,
|
||||
usage_details: UsageDetails | None = None,
|
||||
value: Any | None = None,
|
||||
response_format: type[BaseModel] | None = None,
|
||||
@@ -1424,7 +1407,7 @@ class ChatResponse(AFBaseModel):
|
||||
conversation_id: str | None = None,
|
||||
model_id: str | None = None,
|
||||
created_at: CreatedAtT | None = None,
|
||||
finish_reason: ChatFinishReason | None = None,
|
||||
finish_reason: FinishReason | None = None,
|
||||
usage_details: UsageDetails | None = None,
|
||||
value: Any | None = None,
|
||||
response_format: type[BaseModel] | None = None,
|
||||
@@ -1440,7 +1423,7 @@ class ChatResponse(AFBaseModel):
|
||||
if text is not None:
|
||||
if isinstance(text, str):
|
||||
text = TextContent(text=text)
|
||||
messages.append(ChatMessage(role=ChatRole.ASSISTANT, contents=[text]))
|
||||
messages.append(ChatMessage(role=Role.ASSISTANT, contents=[text]))
|
||||
|
||||
super().__init__(
|
||||
messages=messages, # type: ignore[reportCallIssue]
|
||||
@@ -1528,10 +1511,10 @@ class ChatResponseUpdate(AFBaseModel):
|
||||
|
||||
"""
|
||||
|
||||
contents: list[AIContents]
|
||||
contents: list[Contents]
|
||||
"""The chat response update content items."""
|
||||
|
||||
role: ChatRole | None = None
|
||||
role: Role | None = None
|
||||
"""The role of the author of the response update."""
|
||||
author_name: str | None = None
|
||||
"""The name of the author of the response update."""
|
||||
@@ -1546,7 +1529,7 @@ class ChatResponseUpdate(AFBaseModel):
|
||||
"""The model ID associated with this response update."""
|
||||
created_at: CreatedAtT | None = None # use a datetimeoffset type?
|
||||
"""A timestamp for the chat response update."""
|
||||
finish_reason: ChatFinishReason | None = None
|
||||
finish_reason: FinishReason | None = None
|
||||
"""The finish reason for the operation."""
|
||||
|
||||
additional_properties: dict[str, Any] | None = None
|
||||
@@ -1558,15 +1541,15 @@ class ChatResponseUpdate(AFBaseModel):
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
contents: list[AIContents],
|
||||
role: ChatRole | Literal["system", "user", "assistant", "tool"] | None = None,
|
||||
contents: list[Contents],
|
||||
role: Role | Literal["system", "user", "assistant", "tool"] | None = None,
|
||||
author_name: str | None = None,
|
||||
response_id: str | None = None,
|
||||
message_id: str | None = None,
|
||||
conversation_id: str | None = None,
|
||||
ai_model_id: str | None = None,
|
||||
created_at: CreatedAtT | None = None,
|
||||
finish_reason: ChatFinishReason | None = None,
|
||||
finish_reason: FinishReason | None = None,
|
||||
additional_properties: dict[str, Any] | None = None,
|
||||
raw_representation: Any | None = None,
|
||||
) -> None:
|
||||
@@ -1577,14 +1560,14 @@ class ChatResponseUpdate(AFBaseModel):
|
||||
self,
|
||||
*,
|
||||
text: TextContent | str,
|
||||
role: ChatRole | Literal["system", "user", "assistant", "tool"] | None = None,
|
||||
role: Role | Literal["system", "user", "assistant", "tool"] | None = None,
|
||||
author_name: str | None = None,
|
||||
response_id: str | None = None,
|
||||
message_id: str | None = None,
|
||||
conversation_id: str | None = None,
|
||||
ai_model_id: str | None = None,
|
||||
created_at: CreatedAtT | None = None,
|
||||
finish_reason: ChatFinishReason | None = None,
|
||||
finish_reason: FinishReason | None = None,
|
||||
additional_properties: dict[str, Any] | None = None,
|
||||
raw_representation: Any | None = None,
|
||||
) -> None:
|
||||
@@ -1593,16 +1576,16 @@ class ChatResponseUpdate(AFBaseModel):
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
contents: list[AIContents] | None = None,
|
||||
contents: list[Contents] | None = None,
|
||||
text: TextContent | str | None = None,
|
||||
role: ChatRole | Literal["system", "user", "assistant", "tool"] | None = None,
|
||||
role: Role | Literal["system", "user", "assistant", "tool"] | None = None,
|
||||
author_name: str | None = None,
|
||||
response_id: str | None = None,
|
||||
message_id: str | None = None,
|
||||
conversation_id: str | None = None,
|
||||
ai_model_id: str | None = None,
|
||||
created_at: CreatedAtT | None = None,
|
||||
finish_reason: ChatFinishReason | None = None,
|
||||
finish_reason: FinishReason | None = None,
|
||||
additional_properties: dict[str, Any] | None = None,
|
||||
raw_representation: Any | None = None,
|
||||
) -> None:
|
||||
@@ -1614,7 +1597,7 @@ class ChatResponseUpdate(AFBaseModel):
|
||||
text = TextContent(text=text)
|
||||
contents.append(text)
|
||||
if role and isinstance(role, str):
|
||||
role = ChatRole(value=role)
|
||||
role = Role(value=role)
|
||||
super().__init__(
|
||||
contents=contents, # type: ignore[reportCallIssue]
|
||||
additional_properties=additional_properties, # type: ignore[reportCallIssue]
|
||||
@@ -1637,7 +1620,7 @@ class ChatResponseUpdate(AFBaseModel):
|
||||
def __str__(self) -> str:
|
||||
return self.text
|
||||
|
||||
def with_(self, contents: list[AIContent] | None = None, message_id: str | None = None) -> Self:
|
||||
def with_(self, contents: list[BaseContent] | None = None, message_id: str | None = None) -> Self:
|
||||
"""Returns a new instance with the specified contents and message_id."""
|
||||
if contents is None:
|
||||
contents = []
|
||||
@@ -1709,7 +1692,7 @@ class ChatOptions(AFBaseModel):
|
||||
store: bool | None = None
|
||||
temperature: Annotated[float | None, Field(ge=0.0, le=2.0)] = None
|
||||
tool_choice: ChatToolMode | Literal["auto", "required", "none"] | Mapping[str, Any] | None = None
|
||||
tools: list[AITool | MutableMapping[str, Any]] | None = None
|
||||
tools: list[ToolProtocol | MutableMapping[str, Any]] | None = None
|
||||
top_p: Annotated[float | None, Field(ge=0.0, le=1.0)] = None
|
||||
user: str | None = None
|
||||
|
||||
@@ -1718,21 +1701,21 @@ class ChatOptions(AFBaseModel):
|
||||
def _validate_tools(
|
||||
cls,
|
||||
tools: (
|
||||
AITool
|
||||
ToolProtocol
|
||||
| Callable[..., Any]
|
||||
| MutableMapping[str, Any]
|
||||
| list[AITool | Callable[..., Any] | MutableMapping[str, Any]]
|
||||
| list[ToolProtocol | Callable[..., Any] | MutableMapping[str, Any]]
|
||||
| None
|
||||
),
|
||||
) -> list[AITool | MutableMapping[str, Any]] | None:
|
||||
) -> list[ToolProtocol | MutableMapping[str, Any]] | None:
|
||||
"""Parse the tools field."""
|
||||
if not tools:
|
||||
return None
|
||||
if not isinstance(tools, list):
|
||||
tools = [tools] # type: ignore[reportAssignmentType, assignment]
|
||||
for idx, tool in enumerate(tools): # type: ignore[reportArgumentType, arg-type]
|
||||
if not isinstance(tool, (AITool, MutableMapping)):
|
||||
# Convert to AITool if it's a function or callable
|
||||
if not isinstance(tool, (ToolProtocol, MutableMapping)):
|
||||
# Convert to ToolProtocol if it's a function or callable
|
||||
tools[idx] = ai_function(tool) # type: ignore[reportIndexIssues, reportCallIssue, reportArgumentType, index, call-overload, arg-type]
|
||||
return tools # type: ignore[reportReturnType, return-value]
|
||||
|
||||
@@ -2006,8 +1989,8 @@ class AgentRunResponse(AFBaseModel):
|
||||
class AgentRunResponseUpdate(AFBaseModel):
|
||||
"""Represents a single streaming response chunk from an Agent."""
|
||||
|
||||
contents: list[AIContents] = Field(default_factory=list[AIContents])
|
||||
role: ChatRole | None = None
|
||||
contents: list[Contents] = Field(default_factory=list[Contents])
|
||||
role: Role | None = None
|
||||
author_name: str | None = None
|
||||
response_id: str | None = None
|
||||
message_id: str | None = None
|
||||
|
||||
@@ -20,18 +20,18 @@ from openai.types.beta.threads.run_submit_tool_outputs_params import ToolOutput
|
||||
from openai.types.beta.threads.runs import RunStep
|
||||
from pydantic import Field, PrivateAttr, SecretStr, ValidationError
|
||||
|
||||
from .._clients import ChatClientBase, use_tool_calling
|
||||
from .._clients import BaseChatClient, use_tool_calling
|
||||
from .._tools import AIFunction, HostedCodeInterpreterTool, HostedFileSearchTool
|
||||
from .._types import (
|
||||
AIContents,
|
||||
ChatMessage,
|
||||
ChatOptions,
|
||||
ChatResponse,
|
||||
ChatResponseUpdate,
|
||||
ChatRole,
|
||||
ChatToolMode,
|
||||
Contents,
|
||||
FunctionCallContent,
|
||||
FunctionResultContent,
|
||||
Role,
|
||||
TextContent,
|
||||
UriContent,
|
||||
UsageContent,
|
||||
@@ -39,7 +39,7 @@ from .._types import (
|
||||
)
|
||||
from ..exceptions import ServiceInitializationError
|
||||
from ..telemetry import use_telemetry
|
||||
from ._shared import OpenAIConfigBase, OpenAISettings
|
||||
from ._shared import OpenAIConfigMixin, OpenAISettings
|
||||
|
||||
if sys.version_info >= (3, 11):
|
||||
from typing import Self # pragma: no cover
|
||||
@@ -52,7 +52,7 @@ __all__ = ["OpenAIAssistantsClient"]
|
||||
|
||||
@use_telemetry
|
||||
@use_tool_calling
|
||||
class OpenAIAssistantsClient(OpenAIConfigBase, ChatClientBase):
|
||||
class OpenAIAssistantsClient(OpenAIConfigMixin, BaseChatClient):
|
||||
"""OpenAI Assistants client."""
|
||||
|
||||
assistant_id: str | None = Field(default=None)
|
||||
@@ -274,13 +274,13 @@ class OpenAIAssistantsClient(OpenAIConfigBase, ChatClientBase):
|
||||
message_id=response_id,
|
||||
raw_representation=response.data,
|
||||
response_id=response_id,
|
||||
role=ChatRole.ASSISTANT,
|
||||
role=Role.ASSISTANT,
|
||||
)
|
||||
elif response.event == "thread.run.step.created" and isinstance(response.data, RunStep):
|
||||
response_id = response.data.run_id
|
||||
elif response.event == "thread.message.delta" and isinstance(response.data, MessageDeltaEvent):
|
||||
delta = response.data.delta
|
||||
role = ChatRole.USER if delta.role == "user" else ChatRole.ASSISTANT
|
||||
role = Role.USER if delta.role == "user" else Role.ASSISTANT
|
||||
|
||||
for delta_block in delta.content or []:
|
||||
if isinstance(delta_block, TextDeltaBlock) and delta_block.text and delta_block.text.value:
|
||||
@@ -296,7 +296,7 @@ class OpenAIAssistantsClient(OpenAIConfigBase, ChatClientBase):
|
||||
contents = self._create_function_call_contents(response.data, response_id)
|
||||
if contents:
|
||||
yield ChatResponseUpdate(
|
||||
role=ChatRole.ASSISTANT,
|
||||
role=Role.ASSISTANT,
|
||||
contents=contents,
|
||||
conversation_id=thread_id,
|
||||
message_id=response_id,
|
||||
@@ -317,7 +317,7 @@ class OpenAIAssistantsClient(OpenAIConfigBase, ChatClientBase):
|
||||
)
|
||||
)
|
||||
yield ChatResponseUpdate(
|
||||
role=ChatRole.ASSISTANT,
|
||||
role=Role.ASSISTANT,
|
||||
contents=[usage_content],
|
||||
conversation_id=thread_id,
|
||||
message_id=response_id,
|
||||
@@ -331,12 +331,12 @@ class OpenAIAssistantsClient(OpenAIConfigBase, ChatClientBase):
|
||||
message_id=response_id,
|
||||
raw_representation=response.data,
|
||||
response_id=response_id,
|
||||
role=ChatRole.ASSISTANT,
|
||||
role=Role.ASSISTANT,
|
||||
)
|
||||
|
||||
def _create_function_call_contents(self, event_data: Run, response_id: str | None) -> list[AIContents]:
|
||||
def _create_function_call_contents(self, event_data: Run, response_id: str | None) -> list[Contents]:
|
||||
"""Create function call contents from a tool action event."""
|
||||
contents: list[AIContents] = []
|
||||
contents: list[Contents] = []
|
||||
|
||||
if event_data.required_action is not None:
|
||||
for tool_call in event_data.required_action.submit_tool_outputs.tool_calls:
|
||||
@@ -437,7 +437,7 @@ class OpenAIAssistantsClient(OpenAIConfigBase, ChatClientBase):
|
||||
additional_messages = []
|
||||
additional_messages.append(
|
||||
AdditionalMessage(
|
||||
role="assistant" if chat_message.role == ChatRole.ASSISTANT else "user",
|
||||
role="assistant" if chat_message.role == Role.ASSISTANT else "user",
|
||||
content=message_contents,
|
||||
)
|
||||
)
|
||||
|
||||
@@ -15,19 +15,19 @@ from openai.types.chat.chat_completion_chunk import Choice as ChunkChoice
|
||||
from openai.types.chat.chat_completion_message_custom_tool_call import ChatCompletionMessageCustomToolCall
|
||||
from pydantic import BaseModel, SecretStr, ValidationError
|
||||
|
||||
from .._clients import ChatClientBase, use_tool_calling
|
||||
from .._clients import BaseChatClient, use_tool_calling
|
||||
from .._logging import get_logger
|
||||
from .._tools import AIFunction, AITool, HostedWebSearchTool
|
||||
from .._tools import AIFunction, HostedWebSearchTool, ToolProtocol
|
||||
from .._types import (
|
||||
AIContents,
|
||||
ChatFinishReason,
|
||||
ChatMessage,
|
||||
ChatOptions,
|
||||
ChatResponse,
|
||||
ChatResponseUpdate,
|
||||
ChatRole,
|
||||
Contents,
|
||||
FinishReason,
|
||||
FunctionCallContent,
|
||||
FunctionResultContent,
|
||||
Role,
|
||||
TextContent,
|
||||
UsageContent,
|
||||
UsageDetails,
|
||||
@@ -39,7 +39,7 @@ from ..exceptions import (
|
||||
)
|
||||
from ..telemetry import use_telemetry
|
||||
from ._exceptions import OpenAIContentFilterException
|
||||
from ._shared import OpenAIConfigBase, OpenAIHandler, OpenAISettings, prepare_function_call_results
|
||||
from ._shared import OpenAIBase, OpenAIConfigMixin, OpenAISettings, prepare_function_call_results
|
||||
|
||||
__all__ = ["OpenAIChatClient"]
|
||||
|
||||
@@ -49,7 +49,7 @@ logger = get_logger("agent_framework.openai")
|
||||
# region Base Client
|
||||
@use_telemetry
|
||||
@use_tool_calling
|
||||
class OpenAIChatClientBase(OpenAIHandler, ChatClientBase):
|
||||
class OpenAIBaseChatClient(OpenAIBase, BaseChatClient):
|
||||
"""OpenAI Chat completion class."""
|
||||
|
||||
async def _inner_get_response(
|
||||
@@ -112,10 +112,10 @@ class OpenAIChatClientBase(OpenAIHandler, ChatClientBase):
|
||||
|
||||
# region content creation
|
||||
|
||||
def _chat_to_tool_spec(self, tools: list[AITool | MutableMapping[str, Any]]) -> list[dict[str, Any]]:
|
||||
def _chat_to_tool_spec(self, tools: list[ToolProtocol | MutableMapping[str, Any]]) -> list[dict[str, Any]]:
|
||||
chat_tools: list[dict[str, Any]] = []
|
||||
for tool in tools:
|
||||
if isinstance(tool, AITool):
|
||||
if isinstance(tool, ToolProtocol):
|
||||
match tool:
|
||||
case AIFunction():
|
||||
chat_tools.append(tool.to_json_schema_spec())
|
||||
@@ -125,7 +125,7 @@ class OpenAIChatClientBase(OpenAIHandler, ChatClientBase):
|
||||
chat_tools.append(tool if isinstance(tool, dict) else dict(tool))
|
||||
return chat_tools
|
||||
|
||||
def _process_web_search_tool(self, tools: list[AITool | MutableMapping[str, Any]]) -> dict[str, Any] | None:
|
||||
def _process_web_search_tool(self, tools: list[ToolProtocol | MutableMapping[str, Any]]) -> dict[str, Any] | None:
|
||||
for tool in tools:
|
||||
if isinstance(tool, HostedWebSearchTool):
|
||||
# Web search tool requires special handling
|
||||
@@ -173,12 +173,12 @@ class OpenAIChatClientBase(OpenAIHandler, ChatClientBase):
|
||||
"""Create a chat message content object from a choice."""
|
||||
response_metadata = self._get_metadata_from_chat_response(response)
|
||||
messages: list[ChatMessage] = []
|
||||
finish_reason: ChatFinishReason | None = None
|
||||
finish_reason: FinishReason | None = None
|
||||
for choice in response.choices:
|
||||
response_metadata.update(self._get_metadata_from_chat_choice(choice))
|
||||
if choice.finish_reason:
|
||||
finish_reason = ChatFinishReason(value=choice.finish_reason)
|
||||
contents: list[AIContents] = []
|
||||
finish_reason = FinishReason(value=choice.finish_reason)
|
||||
contents: list[Contents] = []
|
||||
if parsed_tool_calls := [tool for tool in self._get_tool_calls_from_chat_choice(choice)]:
|
||||
contents.extend(parsed_tool_calls)
|
||||
if text_content := self._parse_text_from_choice(choice):
|
||||
@@ -203,27 +203,27 @@ class OpenAIChatClientBase(OpenAIHandler, ChatClientBase):
|
||||
chunk_metadata = self._get_metadata_from_streaming_chat_response(chunk)
|
||||
if chunk.usage:
|
||||
return ChatResponseUpdate(
|
||||
role=ChatRole.ASSISTANT,
|
||||
role=Role.ASSISTANT,
|
||||
contents=[UsageContent(details=self._usage_details_from_openai(chunk.usage), raw_representation=chunk)],
|
||||
ai_model_id=chunk.model,
|
||||
additional_properties=chunk_metadata,
|
||||
response_id=chunk.id,
|
||||
message_id=chunk.id,
|
||||
)
|
||||
contents: list[AIContents] = []
|
||||
finish_reason: ChatFinishReason | None = None
|
||||
contents: list[Contents] = []
|
||||
finish_reason: FinishReason | None = None
|
||||
for choice in chunk.choices:
|
||||
chunk_metadata.update(self._get_metadata_from_chat_choice(choice))
|
||||
contents.extend(self._get_tool_calls_from_chat_choice(choice))
|
||||
if choice.finish_reason:
|
||||
finish_reason = ChatFinishReason(value=choice.finish_reason)
|
||||
finish_reason = FinishReason(value=choice.finish_reason)
|
||||
|
||||
if text_content := self._parse_text_from_choice(choice):
|
||||
contents.append(text_content)
|
||||
return ChatResponseUpdate(
|
||||
created_at=datetime.fromtimestamp(chunk.created).strftime("%Y-%m-%dT%H:%M:%S.%fZ"),
|
||||
contents=contents,
|
||||
role=ChatRole.ASSISTANT,
|
||||
role=Role.ASSISTANT,
|
||||
ai_model_id=chunk.model,
|
||||
additional_properties=chunk_metadata,
|
||||
finish_reason=finish_reason,
|
||||
@@ -266,9 +266,9 @@ class OpenAIChatClientBase(OpenAIHandler, ChatClientBase):
|
||||
"logprobs": getattr(choice, "logprobs", None),
|
||||
}
|
||||
|
||||
def _get_tool_calls_from_chat_choice(self, choice: Choice | ChunkChoice) -> list[AIContents]:
|
||||
def _get_tool_calls_from_chat_choice(self, choice: Choice | ChunkChoice) -> list[Contents]:
|
||||
"""Get tool calls from a chat choice."""
|
||||
resp: list[AIContents] = []
|
||||
resp: list[Contents] = []
|
||||
content = choice.message if isinstance(choice, Choice) else choice.delta
|
||||
if content and content.tool_calls:
|
||||
for tool in content.tool_calls:
|
||||
@@ -295,7 +295,7 @@ class OpenAIChatClientBase(OpenAIHandler, ChatClientBase):
|
||||
|
||||
Allowing customization of the key names for role/author, and optionally overriding the role.
|
||||
|
||||
ChatRole.TOOL messages need to be formatted different than system/user/assistant messages:
|
||||
Role.TOOL messages need to be formatted different than system/user/assistant messages:
|
||||
They require a "tool_call_id" and (function) "name" key, and the "metadata" key should
|
||||
be removed. The "encoding" key should also be removed.
|
||||
|
||||
@@ -320,7 +320,7 @@ class OpenAIChatClientBase(OpenAIHandler, ChatClientBase):
|
||||
all_messages: list[dict[str, Any]] = []
|
||||
for content in message.contents:
|
||||
args: dict[str, Any] = {
|
||||
"role": message.role.value if isinstance(message.role, ChatRole) else message.role,
|
||||
"role": message.role.value if isinstance(message.role, Role) else message.role,
|
||||
}
|
||||
if message.additional_properties:
|
||||
args["metadata"] = message.additional_properties
|
||||
@@ -344,7 +344,7 @@ class OpenAIChatClientBase(OpenAIHandler, ChatClientBase):
|
||||
all_messages.append(args)
|
||||
return all_messages
|
||||
|
||||
def _openai_content_parser(self, content: AIContents) -> dict[str, Any]:
|
||||
def _openai_content_parser(self, content: Contents) -> dict[str, Any]:
|
||||
"""Parse contents into the openai format."""
|
||||
match content:
|
||||
case FunctionCallContent():
|
||||
@@ -376,7 +376,7 @@ class OpenAIChatClientBase(OpenAIHandler, ChatClientBase):
|
||||
TOpenAIChatClient = TypeVar("TOpenAIChatClient", bound="OpenAIChatClient")
|
||||
|
||||
|
||||
class OpenAIChatClient(OpenAIConfigBase, OpenAIChatClientBase):
|
||||
class OpenAIChatClient(OpenAIConfigMixin, OpenAIBaseChatClient):
|
||||
"""OpenAI Chat completion class."""
|
||||
|
||||
def __init__(
|
||||
|
||||
@@ -31,22 +31,22 @@ from openai.types.responses.web_search_tool_param import UserLocation as WebSear
|
||||
from openai.types.responses.web_search_tool_param import WebSearchToolParam
|
||||
from pydantic import BaseModel, SecretStr, ValidationError
|
||||
|
||||
from .._clients import ChatClientBase, use_tool_calling
|
||||
from .._clients import BaseChatClient, use_tool_calling
|
||||
from .._logging import get_logger
|
||||
from .._tools import AIFunction, AITool, HostedCodeInterpreterTool, HostedFileSearchTool, HostedWebSearchTool
|
||||
from .._tools import AIFunction, HostedCodeInterpreterTool, HostedFileSearchTool, HostedWebSearchTool, ToolProtocol
|
||||
from .._types import (
|
||||
AIContents,
|
||||
ChatMessage,
|
||||
ChatOptions,
|
||||
ChatResponse,
|
||||
ChatResponseUpdate,
|
||||
ChatRole,
|
||||
CitationAnnotation,
|
||||
Contents,
|
||||
DataContent,
|
||||
FunctionCallContent,
|
||||
FunctionResultContent,
|
||||
HostedFileContent,
|
||||
HostedVectorStoreContent,
|
||||
Role,
|
||||
TextContent,
|
||||
TextReasoningContent,
|
||||
TextSpanRegion,
|
||||
@@ -61,7 +61,7 @@ from ..exceptions import (
|
||||
)
|
||||
from ..telemetry import use_telemetry
|
||||
from ._exceptions import OpenAIContentFilterException
|
||||
from ._shared import OpenAIConfigBase, OpenAIHandler, OpenAISettings, prepare_function_call_results
|
||||
from ._shared import OpenAIBase, OpenAIConfigMixin, OpenAISettings, prepare_function_call_results
|
||||
|
||||
if sys.version_info >= (3, 12):
|
||||
from typing import override # type: ignore # pragma: no cover
|
||||
@@ -81,7 +81,7 @@ __all__ = ["OpenAIResponsesClient"]
|
||||
# region ResponsesClient
|
||||
|
||||
|
||||
class OpenAIResponsesClientBase(OpenAIHandler, ChatClientBase):
|
||||
class OpenAIBaseResponsesClient(OpenAIBase, BaseChatClient):
|
||||
"""Base class for all OpenAI Responses based API's."""
|
||||
|
||||
FILE_SEARCH_MAX_RESULTS: int = 50
|
||||
@@ -110,10 +110,10 @@ class OpenAIResponsesClientBase(OpenAIHandler, ChatClientBase):
|
||||
store: bool | None = None,
|
||||
temperature: float | None = None,
|
||||
tool_choice: "ChatToolMode" | Literal["auto", "required", "none"] | dict[str, Any] | None = "auto",
|
||||
tools: AITool
|
||||
tools: ToolProtocol
|
||||
| Callable[..., Any]
|
||||
| MutableMapping[str, Any]
|
||||
| list[AITool | Callable[..., Any] | MutableMapping[str, Any]]
|
||||
| list[ToolProtocol | Callable[..., Any] | MutableMapping[str, Any]]
|
||||
| None = None,
|
||||
top_p: float | None = None,
|
||||
user: str | None = None,
|
||||
@@ -200,10 +200,10 @@ class OpenAIResponsesClientBase(OpenAIHandler, ChatClientBase):
|
||||
store: bool | None = None,
|
||||
temperature: float | None = None,
|
||||
tool_choice: "ChatToolMode" | Literal["auto", "required", "none"] | dict[str, Any] | None = "auto",
|
||||
tools: AITool
|
||||
tools: ToolProtocol
|
||||
| Callable[..., Any]
|
||||
| MutableMapping[str, Any]
|
||||
| list[AITool | Callable[..., Any] | MutableMapping[str, Any]]
|
||||
| list[ToolProtocol | Callable[..., Any] | MutableMapping[str, Any]]
|
||||
| None = None,
|
||||
top_p: float | None = None,
|
||||
user: str | None = None,
|
||||
@@ -365,11 +365,11 @@ class OpenAIResponsesClientBase(OpenAIHandler, ChatClientBase):
|
||||
# region Prep methods
|
||||
|
||||
def _chat_to_response_tool_spec(
|
||||
self, tools: list[AITool | MutableMapping[str, Any]]
|
||||
self, tools: list[ToolProtocol | MutableMapping[str, Any]]
|
||||
) -> list[ToolParam | dict[str, Any]]:
|
||||
response_tools: list[ToolParam | dict[str, Any]] = []
|
||||
for tool in tools:
|
||||
if isinstance(tool, AITool):
|
||||
if isinstance(tool, ToolProtocol):
|
||||
match tool:
|
||||
case HostedCodeInterpreterTool():
|
||||
tool_args: dict[str, Any] = {"type": "auto"}
|
||||
@@ -471,7 +471,7 @@ class OpenAIResponsesClientBase(OpenAIHandler, ChatClientBase):
|
||||
|
||||
Allowing customization of the key names for role/author, and optionally overriding the role.
|
||||
|
||||
ChatRole.TOOL messages need to be formatted different than system/user/assistant messages:
|
||||
Role.TOOL messages need to be formatted different than system/user/assistant messages:
|
||||
They require a "tool_call_id" and (function) "name" key, and the "metadata" key should
|
||||
be removed. The "encoding" key should also be removed.
|
||||
|
||||
@@ -507,7 +507,7 @@ class OpenAIResponsesClientBase(OpenAIHandler, ChatClientBase):
|
||||
structured_response: BaseModel | None = response.output_parsed if isinstance(response, ParsedResponse) else None # type: ignore[reportUnknownMemberType]
|
||||
|
||||
metadata: dict[str, Any] = response.metadata or {}
|
||||
contents: list[AIContents] = []
|
||||
contents: list[Contents] = []
|
||||
for item in response.output: # type: ignore[reportUnknownMemberType]
|
||||
match item.type:
|
||||
# types:
|
||||
@@ -517,12 +517,12 @@ class OpenAIResponsesClientBase(OpenAIHandler, ChatClientBase):
|
||||
# ResponseFunctionWebSearch |
|
||||
# ResponseComputerToolCall |
|
||||
# ResponseReasoningItem |
|
||||
# McpCall |
|
||||
# McpApprovalRequest |
|
||||
# MCPCall |
|
||||
# MCPApprovalRequest |
|
||||
# ImageGenerationCall |
|
||||
# LocalShellCall |
|
||||
# LocalShellCallAction |
|
||||
# McpListTools |
|
||||
# MCPListTools |
|
||||
# ResponseCodeInterpreterToolCall |
|
||||
# ResponseCustomToolCall |
|
||||
# ParsedResponseOutputMessage[BaseModel] |
|
||||
@@ -677,7 +677,7 @@ class OpenAIResponsesClientBase(OpenAIHandler, ChatClientBase):
|
||||
) -> ChatResponseUpdate:
|
||||
"""Create a streaming chat message content object from a choice."""
|
||||
metadata: dict[str, Any] = {}
|
||||
items: list[AIContents] = []
|
||||
items: list[Contents] = []
|
||||
conversation_id: str | None = None
|
||||
model = self.ai_model_id
|
||||
# TODO(peterychang): Add support for other content types
|
||||
@@ -720,7 +720,7 @@ class OpenAIResponsesClientBase(OpenAIHandler, ChatClientBase):
|
||||
return ChatResponseUpdate(
|
||||
contents=items,
|
||||
conversation_id=conversation_id,
|
||||
role=ChatRole.ASSISTANT,
|
||||
role=Role.ASSISTANT,
|
||||
ai_model_id=model,
|
||||
additional_properties=metadata,
|
||||
raw_representation=event,
|
||||
@@ -746,7 +746,7 @@ class OpenAIResponsesClientBase(OpenAIHandler, ChatClientBase):
|
||||
"""Parse a chat message into the openai format."""
|
||||
all_messages: list[dict[str, Any]] = []
|
||||
args: dict[str, Any] = {
|
||||
"role": message.role.value if isinstance(message.role, ChatRole) else message.role,
|
||||
"role": message.role.value if isinstance(message.role, Role) else message.role,
|
||||
}
|
||||
if message.additional_properties:
|
||||
args["metadata"] = message.additional_properties
|
||||
@@ -769,8 +769,8 @@ class OpenAIResponsesClientBase(OpenAIHandler, ChatClientBase):
|
||||
|
||||
def _openai_content_parser(
|
||||
self,
|
||||
role: ChatRole,
|
||||
content: AIContents,
|
||||
role: Role,
|
||||
content: Contents,
|
||||
call_id_to_id: dict[str, str],
|
||||
) -> dict[str, Any]:
|
||||
"""Parse contents into the openai format."""
|
||||
@@ -794,7 +794,7 @@ class OpenAIResponsesClientBase(OpenAIHandler, ChatClientBase):
|
||||
return args
|
||||
case TextContent():
|
||||
return {
|
||||
"type": "output_text" if role == ChatRole.ASSISTANT else "input_text",
|
||||
"type": "output_text" if role == Role.ASSISTANT else "input_text",
|
||||
"text": content.text,
|
||||
}
|
||||
# TODO(peterychang): We'll probably need to specialize the other content types as well
|
||||
@@ -815,7 +815,7 @@ TOpenAIResponsesClient = TypeVar("TOpenAIResponsesClient", bound="OpenAIResponse
|
||||
|
||||
@use_telemetry
|
||||
@use_tool_calling
|
||||
class OpenAIResponsesClient(OpenAIConfigBase, OpenAIResponsesClientBase):
|
||||
class OpenAIResponsesClient(OpenAIConfigMixin, OpenAIBaseResponsesClient):
|
||||
"""OpenAI Responses client class."""
|
||||
|
||||
def __init__(
|
||||
|
||||
@@ -22,7 +22,7 @@ from pydantic.types import StringConstraints
|
||||
|
||||
from .._logging import get_logger
|
||||
from .._pydantic import AFBaseModel, AFBaseSettings
|
||||
from .._types import AIContents, ChatOptions, SpeechToTextOptions, TextToSpeechOptions
|
||||
from .._types import ChatOptions, Contents, SpeechToTextOptions, TextToSpeechOptions
|
||||
from ..exceptions import ServiceInitializationError
|
||||
from ..telemetry import APP_INFO, USER_AGENT_KEY, prepend_agent_framework_to_user_agent
|
||||
|
||||
@@ -50,7 +50,7 @@ __all__ = [
|
||||
]
|
||||
|
||||
|
||||
def prepare_function_call_results(content: AIContents | Any | list[AIContents | Any]) -> str | list[str]:
|
||||
def prepare_function_call_results(content: Contents | Any | list[Contents | Any]) -> str | list[str]:
|
||||
"""Prepare the values of the function call results."""
|
||||
if isinstance(content, list):
|
||||
results: list[str] = []
|
||||
@@ -117,14 +117,14 @@ class OpenAISettings(AFBaseSettings):
|
||||
realtime_model_id: str | None = None
|
||||
|
||||
|
||||
class OpenAIHandler(AFBaseModel):
|
||||
class OpenAIBase(AFBaseModel):
|
||||
"""Base class for OpenAI Clients."""
|
||||
|
||||
client: AsyncOpenAI
|
||||
ai_model_id: Annotated[str, StringConstraints(strip_whitespace=True, min_length=1)]
|
||||
|
||||
|
||||
class OpenAIConfigBase(OpenAIHandler):
|
||||
class OpenAIConfigMixin(OpenAIBase):
|
||||
"""Internal class for configuring a connection to an OpenAI service."""
|
||||
|
||||
MODEL_PROVIDER_NAME: ClassVar[str] = "openai" # type: ignore[reportIncompatibleVariableOverride, misc]
|
||||
|
||||
@@ -18,8 +18,8 @@ from ._pydantic import AFBaseSettings
|
||||
if TYPE_CHECKING: # pragma: no cover
|
||||
from opentelemetry.util._decorator import _AgnosticContextManager # type: ignore[reportPrivateUsage]
|
||||
|
||||
from ._agents import AIAgent, ChatClientAgent
|
||||
from ._clients import ChatClientBase
|
||||
from ._agents import AgentProtocol, ChatAgent
|
||||
from ._clients import BaseChatClient
|
||||
from ._threads import AgentThread
|
||||
from ._tools import AIFunction
|
||||
from ._types import (
|
||||
@@ -31,8 +31,8 @@ if TYPE_CHECKING: # pragma: no cover
|
||||
ChatResponseUpdate,
|
||||
)
|
||||
|
||||
TChatClientBase = TypeVar("TChatClientBase", bound="ChatClientBase")
|
||||
TChatClientAgent = TypeVar("TChatClientAgent", bound="ChatClientAgent")
|
||||
TBaseChatClient = TypeVar("TBaseChatClient", bound="BaseChatClient")
|
||||
TChatClientAgent = TypeVar("TChatClientAgent", bound="ChatAgent")
|
||||
|
||||
tracer = get_tracer("agent_framework")
|
||||
logger = get_logger()
|
||||
@@ -269,7 +269,7 @@ def _set_error(span: Span, error: Exception) -> None:
|
||||
span.set_status(StatusCode.ERROR, repr(error))
|
||||
|
||||
|
||||
# region ChatClient
|
||||
# region ChatClientProtocol
|
||||
|
||||
|
||||
def _trace_chat_get_response(
|
||||
@@ -283,7 +283,7 @@ def _trace_chat_get_response(
|
||||
|
||||
@functools.wraps(completion_func)
|
||||
async def wrap_inner_get_response(
|
||||
self: "ChatClientBase",
|
||||
self: "BaseChatClient",
|
||||
*,
|
||||
messages: MutableSequence["ChatMessage"],
|
||||
chat_options: "ChatOptions",
|
||||
@@ -334,7 +334,7 @@ def _trace_chat_get_streaming_response(
|
||||
|
||||
@functools.wraps(completion_func)
|
||||
async def wrap_inner_get_streaming_response(
|
||||
self: "ChatClientBase", *, messages: MutableSequence["ChatMessage"], chat_options: "ChatOptions", **kwargs: Any
|
||||
self: "BaseChatClient", *, messages: MutableSequence["ChatMessage"], chat_options: "ChatOptions", **kwargs: Any
|
||||
) -> AsyncIterable["ChatResponseUpdate"]:
|
||||
if not MODEL_DIAGNOSTICS_SETTINGS.ENABLED:
|
||||
# If model diagnostics are not enabled, just return the completion
|
||||
@@ -375,11 +375,11 @@ def _trace_chat_get_streaming_response(
|
||||
return wrap_inner_get_streaming_response
|
||||
|
||||
|
||||
def use_telemetry(cls: type[TChatClientBase]) -> type[TChatClientBase]:
|
||||
def use_telemetry(cls: type[TBaseChatClient]) -> type[TBaseChatClient]:
|
||||
"""Class decorator that enables telemetry for a chat client.
|
||||
|
||||
Remarks:
|
||||
This only works on classes that derive from ChatClientBase
|
||||
This only works on classes that derive from BaseChatClient
|
||||
and the _inner_get_response
|
||||
and _inner_get_streaming_response methods.
|
||||
It also relies on the presence of the MODEL_PROVIDER_NAME class variable.
|
||||
@@ -520,7 +520,7 @@ def _trace_agent_run(
|
||||
|
||||
@functools.wraps(run_func)
|
||||
async def wrap_run(
|
||||
self: "ChatClientAgent",
|
||||
self: "ChatAgent",
|
||||
messages: "str | ChatMessage | list[str] | list[ChatMessage] | None" = None,
|
||||
*,
|
||||
thread: "AgentThread | None" = None,
|
||||
@@ -560,7 +560,7 @@ def _trace_agent_run(
|
||||
return wrap_run
|
||||
|
||||
|
||||
def _trace_agent_run_streaming(
|
||||
def _trace_agent_run_stream(
|
||||
run_func: Callable[..., AsyncIterable["AgentRunResponseUpdate"]],
|
||||
) -> Callable[..., AsyncIterable["AgentRunResponseUpdate"]]:
|
||||
"""Decorator to trace streaming agent run activities.
|
||||
@@ -570,8 +570,8 @@ def _trace_agent_run_streaming(
|
||||
"""
|
||||
|
||||
@functools.wraps(run_func)
|
||||
async def wrap_run_streaming(
|
||||
self: "ChatClientAgent",
|
||||
async def wrap_run_stream(
|
||||
self: "ChatAgent",
|
||||
messages: "str | ChatMessage | list[str] | list[ChatMessage] | None" = None,
|
||||
*,
|
||||
thread: "AgentThread | None" = None,
|
||||
@@ -610,23 +610,23 @@ def _trace_agent_run_streaming(
|
||||
raise
|
||||
|
||||
# Mark the wrapper decorator as a streaming agent run decorator
|
||||
wrap_run_streaming.__model_diagnostics_streaming_agent_run__ = True # type: ignore
|
||||
return wrap_run_streaming
|
||||
wrap_run_stream.__model_diagnostics_streaming_agent_run__ = True # type: ignore
|
||||
return wrap_run_stream
|
||||
|
||||
|
||||
def use_agent_telemetry(cls: type[TChatClientAgent]) -> type[TChatClientAgent]:
|
||||
"""Class decorator that enables telemetry for an agent."""
|
||||
if run := getattr(cls, "run", None):
|
||||
cls.run = _trace_agent_run(run) # type: ignore
|
||||
if run_streaming := getattr(cls, "run_streaming", None):
|
||||
cls.run_streaming = _trace_agent_run_streaming(run_streaming) # type: ignore
|
||||
if run_stream := getattr(cls, "run_stream", None):
|
||||
cls.run_stream = _trace_agent_run_stream(run_stream) # type: ignore
|
||||
return cls
|
||||
|
||||
|
||||
def _get_agent_run_span(
|
||||
*,
|
||||
operation_name: str,
|
||||
agent: "AIAgent",
|
||||
agent: "AgentProtocol",
|
||||
system: str,
|
||||
thread: "AgentThread | None",
|
||||
**kwargs: Any,
|
||||
|
||||
@@ -4,7 +4,7 @@ from typing import Any
|
||||
from pydantic import BaseModel
|
||||
from pytest import fixture
|
||||
|
||||
from agent_framework import AITool, ChatMessage, ai_function
|
||||
from agent_framework import ChatMessage, ToolProtocol, ai_function
|
||||
from agent_framework.telemetry import ModelDiagnosticSettings
|
||||
|
||||
|
||||
@@ -14,8 +14,8 @@ def chat_history() -> list[ChatMessage]:
|
||||
|
||||
|
||||
@fixture
|
||||
def ai_tool() -> AITool:
|
||||
"""Returns a generic AITool."""
|
||||
def ai_tool() -> ToolProtocol:
|
||||
"""Returns a generic ToolProtocol."""
|
||||
|
||||
class GenericTool(BaseModel):
|
||||
name: str
|
||||
@@ -32,8 +32,8 @@ def ai_tool() -> AITool:
|
||||
|
||||
|
||||
@fixture
|
||||
def ai_function_tool() -> AITool:
|
||||
"""Returns a executable AITool."""
|
||||
def ai_function_tool() -> ToolProtocol:
|
||||
"""Returns a executable ToolProtocol."""
|
||||
|
||||
@ai_function
|
||||
def simple_function(x: int, y: int) -> int:
|
||||
|
||||
@@ -7,19 +7,19 @@ from uuid import uuid4
|
||||
from pytest import fixture, raises
|
||||
|
||||
from agent_framework import (
|
||||
AgentProtocol,
|
||||
AgentRunResponse,
|
||||
AgentRunResponseUpdate,
|
||||
AgentThread,
|
||||
AIAgent,
|
||||
ChatClient,
|
||||
ChatClientAgent,
|
||||
ChatClientBase,
|
||||
BaseChatClient,
|
||||
ChatAgent,
|
||||
ChatClientProtocol,
|
||||
ChatMessage,
|
||||
ChatMessageList,
|
||||
ChatOptions,
|
||||
ChatResponse,
|
||||
ChatResponseUpdate,
|
||||
ChatRole,
|
||||
Role,
|
||||
TextContent,
|
||||
)
|
||||
from agent_framework.exceptions import AgentExecutionException
|
||||
@@ -31,7 +31,7 @@ class MockAgentThread(AgentThread):
|
||||
|
||||
|
||||
# Mock Agent implementation for testing
|
||||
class MockAgent(AIAgent):
|
||||
class MockAgent(AgentProtocol):
|
||||
@property
|
||||
def id(self) -> str:
|
||||
return str(uuid4())
|
||||
@@ -57,9 +57,9 @@ class MockAgent(AIAgent):
|
||||
thread: AgentThread | None = None,
|
||||
**kwargs: Any,
|
||||
) -> AgentRunResponse:
|
||||
return AgentRunResponse(messages=[ChatMessage(role=ChatRole.ASSISTANT, contents=[TextContent("Response")])])
|
||||
return AgentRunResponse(messages=[ChatMessage(role=Role.ASSISTANT, contents=[TextContent("Response")])])
|
||||
|
||||
async def run_streaming(
|
||||
async def run_stream(
|
||||
self,
|
||||
messages: str | ChatMessage | list[str] | list[ChatMessage] | None = None,
|
||||
*,
|
||||
@@ -72,8 +72,8 @@ class MockAgent(AIAgent):
|
||||
return MockAgentThread()
|
||||
|
||||
|
||||
# Mock ChatClient implementation for testing
|
||||
class MockChatClient(ChatClientBase):
|
||||
# Mock ChatClientProtocol implementation for testing
|
||||
class MockChatClient(BaseChatClient):
|
||||
_mock_response: ChatResponse | None = None
|
||||
|
||||
def __init__(self, mock_response: ChatResponse | None = None) -> None:
|
||||
@@ -89,7 +89,7 @@ class MockChatClient(ChatClientBase):
|
||||
return (
|
||||
self._mock_response
|
||||
if self._mock_response
|
||||
else ChatResponse(messages=ChatMessage(role=ChatRole.ASSISTANT, text="test response"))
|
||||
else ChatResponse(messages=ChatMessage(role=Role.ASSISTANT, text="test response"))
|
||||
)
|
||||
|
||||
async def _inner_get_streaming_response(
|
||||
@@ -99,7 +99,7 @@ class MockChatClient(ChatClientBase):
|
||||
chat_options: ChatOptions,
|
||||
**kwargs: Any,
|
||||
) -> AsyncIterable[ChatResponseUpdate]:
|
||||
yield ChatResponseUpdate(role=ChatRole.ASSISTANT, text=TextContent(text="test streaming response"))
|
||||
yield ChatResponseUpdate(role=Role.ASSISTANT, text=TextContent(text="test streaming response"))
|
||||
|
||||
|
||||
@fixture
|
||||
@@ -108,12 +108,12 @@ def agent_thread() -> AgentThread:
|
||||
|
||||
|
||||
@fixture
|
||||
def agent() -> AIAgent:
|
||||
def agent() -> AgentProtocol:
|
||||
return MockAgent()
|
||||
|
||||
|
||||
@fixture
|
||||
def chat_client() -> ChatClientBase:
|
||||
def chat_client() -> BaseChatClient:
|
||||
return MockChatClient()
|
||||
|
||||
|
||||
@@ -121,33 +121,33 @@ def test_agent_thread_type(agent_thread: AgentThread) -> None:
|
||||
assert isinstance(agent_thread, AgentThread)
|
||||
|
||||
|
||||
def test_agent_type(agent: AIAgent) -> None:
|
||||
assert isinstance(agent, AIAgent)
|
||||
def test_agent_type(agent: AgentProtocol) -> None:
|
||||
assert isinstance(agent, AgentProtocol)
|
||||
|
||||
|
||||
async def test_agent_run(agent: AIAgent) -> None:
|
||||
async def test_agent_run(agent: AgentProtocol) -> None:
|
||||
response = await agent.run("test")
|
||||
assert response.messages[0].role == ChatRole.ASSISTANT
|
||||
assert response.messages[0].role == Role.ASSISTANT
|
||||
assert response.messages[0].text == "Response"
|
||||
|
||||
|
||||
async def test_agent_run_streaming(agent: AIAgent) -> None:
|
||||
async def test_agent_run_streaming(agent: AgentProtocol) -> None:
|
||||
async def collect_updates(updates: AsyncIterable[AgentRunResponseUpdate]) -> list[AgentRunResponseUpdate]:
|
||||
return [u async for u in updates]
|
||||
|
||||
updates = await collect_updates(agent.run_streaming(messages="test"))
|
||||
updates = await collect_updates(agent.run_stream(messages="test"))
|
||||
assert len(updates) == 1
|
||||
assert updates[0].text == "Response"
|
||||
|
||||
|
||||
def test_chat_client_agent_type(chat_client: ChatClient) -> None:
|
||||
chat_client_agent = ChatClientAgent(chat_client=chat_client)
|
||||
assert isinstance(chat_client_agent, AIAgent)
|
||||
def test_chat_client_agent_type(chat_client: ChatClientProtocol) -> None:
|
||||
chat_client_agent = ChatAgent(chat_client=chat_client)
|
||||
assert isinstance(chat_client_agent, AgentProtocol)
|
||||
|
||||
|
||||
async def test_chat_client_agent_init(chat_client: ChatClient) -> None:
|
||||
async def test_chat_client_agent_init(chat_client: ChatClientProtocol) -> None:
|
||||
agent_id = str(uuid4())
|
||||
agent = ChatClientAgent(chat_client=chat_client, id=agent_id, description="Test")
|
||||
agent = ChatAgent(chat_client=chat_client, id=agent_id, description="Test")
|
||||
|
||||
assert agent.id == agent_id
|
||||
assert agent.name is None
|
||||
@@ -155,9 +155,9 @@ async def test_chat_client_agent_init(chat_client: ChatClient) -> None:
|
||||
assert agent.display_name == agent_id # Display name defaults to id if name is None
|
||||
|
||||
|
||||
async def test_chat_client_agent_init_with_name(chat_client: ChatClient) -> None:
|
||||
async def test_chat_client_agent_init_with_name(chat_client: ChatClientProtocol) -> None:
|
||||
agent_id = str(uuid4())
|
||||
agent = ChatClientAgent(chat_client=chat_client, id=agent_id, name="Test Agent", description="Test")
|
||||
agent = ChatAgent(chat_client=chat_client, id=agent_id, name="Test Agent", description="Test")
|
||||
|
||||
assert agent.id == agent_id
|
||||
assert agent.name == "Test Agent"
|
||||
@@ -165,37 +165,37 @@ async def test_chat_client_agent_init_with_name(chat_client: ChatClient) -> None
|
||||
assert agent.display_name == "Test Agent" # Display name is the name if present
|
||||
|
||||
|
||||
async def test_chat_client_agent_run(chat_client: ChatClient) -> None:
|
||||
agent = ChatClientAgent(chat_client=chat_client)
|
||||
async def test_chat_client_agent_run(chat_client: ChatClientProtocol) -> None:
|
||||
agent = ChatAgent(chat_client=chat_client)
|
||||
|
||||
result = await agent.run("Hello")
|
||||
|
||||
assert result.text == "test response"
|
||||
|
||||
|
||||
async def test_chat_client_agent_run_streaming(chat_client: ChatClient) -> None:
|
||||
agent = ChatClientAgent(chat_client=chat_client)
|
||||
async def test_chat_client_agent_run_streaming(chat_client: ChatClientProtocol) -> None:
|
||||
agent = ChatAgent(chat_client=chat_client)
|
||||
|
||||
result = await AgentRunResponse.from_agent_response_generator(agent.run_streaming("Hello"))
|
||||
result = await AgentRunResponse.from_agent_response_generator(agent.run_stream("Hello"))
|
||||
|
||||
assert result.text == "test streaming response"
|
||||
|
||||
|
||||
async def test_chat_client_agent_get_new_thread(chat_client: ChatClient) -> None:
|
||||
agent = ChatClientAgent(chat_client=chat_client)
|
||||
async def test_chat_client_agent_get_new_thread(chat_client: ChatClientProtocol) -> None:
|
||||
agent = ChatAgent(chat_client=chat_client)
|
||||
thread = agent.get_new_thread()
|
||||
|
||||
assert isinstance(thread, AgentThread)
|
||||
|
||||
|
||||
async def test_chat_client_agent_prepare_thread_and_messages(chat_client: ChatClient) -> None:
|
||||
agent = ChatClientAgent(chat_client=chat_client)
|
||||
message = ChatMessage(role=ChatRole.USER, text="Hello")
|
||||
async def test_chat_client_agent_prepare_thread_and_messages(chat_client: ChatClientProtocol) -> None:
|
||||
agent = ChatAgent(chat_client=chat_client)
|
||||
message = ChatMessage(role=Role.USER, text="Hello")
|
||||
thread = AgentThread(message_store=ChatMessageList(messages=[message]))
|
||||
|
||||
_, result_messages = await agent._prepare_thread_and_messages( # type: ignore[reportPrivateUsage]
|
||||
thread=thread,
|
||||
input_messages=[ChatMessage(role=ChatRole.USER, text="Test")],
|
||||
input_messages=[ChatMessage(role=Role.USER, text="Test")],
|
||||
)
|
||||
|
||||
assert len(result_messages) == 2
|
||||
@@ -206,11 +206,11 @@ async def test_chat_client_agent_prepare_thread_and_messages(chat_client: ChatCl
|
||||
async def test_chat_client_agent_update_thread_id() -> None:
|
||||
chat_client = MockChatClient(
|
||||
mock_response=ChatResponse(
|
||||
messages=[ChatMessage(role=ChatRole.ASSISTANT, contents=[TextContent("test response")])],
|
||||
messages=[ChatMessage(role=Role.ASSISTANT, contents=[TextContent("test response")])],
|
||||
conversation_id="123",
|
||||
)
|
||||
)
|
||||
agent = ChatClientAgent(chat_client=chat_client)
|
||||
agent = ChatAgent(chat_client=chat_client)
|
||||
thread = agent.get_new_thread()
|
||||
|
||||
result = await agent.run("Hello", thread=thread)
|
||||
@@ -219,8 +219,8 @@ async def test_chat_client_agent_update_thread_id() -> None:
|
||||
assert thread.service_thread_id == "123"
|
||||
|
||||
|
||||
async def test_chat_client_agent_update_thread_messages(chat_client: ChatClient) -> None:
|
||||
agent = ChatClientAgent(chat_client=chat_client)
|
||||
async def test_chat_client_agent_update_thread_messages(chat_client: ChatClientProtocol) -> None:
|
||||
agent = ChatAgent(chat_client=chat_client)
|
||||
thread = agent.get_new_thread()
|
||||
|
||||
result = await agent.run("Hello", thread=thread)
|
||||
@@ -237,26 +237,26 @@ async def test_chat_client_agent_update_thread_messages(chat_client: ChatClient)
|
||||
assert chat_messages[1].text == "test response"
|
||||
|
||||
|
||||
async def test_chat_client_agent_update_thread_conversation_id_missing(chat_client: ChatClient) -> None:
|
||||
agent = ChatClientAgent(chat_client=chat_client)
|
||||
async def test_chat_client_agent_update_thread_conversation_id_missing(chat_client: ChatClientProtocol) -> None:
|
||||
agent = ChatAgent(chat_client=chat_client)
|
||||
thread = AgentThread(service_thread_id="123")
|
||||
|
||||
with raises(AgentExecutionException, match="Service did not return a valid conversation id"):
|
||||
agent._update_thread_with_type_and_conversation_id(thread, None) # type: ignore[reportPrivateUsage]
|
||||
|
||||
|
||||
async def test_chat_client_agent_default_author_name(chat_client: ChatClient) -> None:
|
||||
async def test_chat_client_agent_default_author_name(chat_client: ChatClientProtocol) -> None:
|
||||
# Name is not specified here, so default name should be used
|
||||
agent = ChatClientAgent(chat_client=chat_client)
|
||||
agent = ChatAgent(chat_client=chat_client)
|
||||
|
||||
result = await agent.run("Hello")
|
||||
assert result.text == "test response"
|
||||
assert result.messages[0].author_name == "UnnamedAgent"
|
||||
|
||||
|
||||
async def test_chat_client_agent_author_name_as_agent_name(chat_client: ChatClient) -> None:
|
||||
async def test_chat_client_agent_author_name_as_agent_name(chat_client: ChatClientProtocol) -> None:
|
||||
# Name is specified here, so it should be used as author name
|
||||
agent = ChatClientAgent(chat_client=chat_client, name="TestAgent")
|
||||
agent = ChatAgent(chat_client=chat_client, name="TestAgent")
|
||||
|
||||
result = await agent.run("Hello")
|
||||
assert result.text == "test response"
|
||||
@@ -267,11 +267,11 @@ async def test_chat_client_agent_author_name_is_used_from_response() -> None:
|
||||
chat_client = MockChatClient(
|
||||
mock_response=ChatResponse(
|
||||
messages=[
|
||||
ChatMessage(role=ChatRole.ASSISTANT, contents=[TextContent("test response")], author_name="TestAuthor")
|
||||
ChatMessage(role=Role.ASSISTANT, contents=[TextContent("test response")], author_name="TestAuthor")
|
||||
]
|
||||
)
|
||||
)
|
||||
agent = ChatClientAgent(chat_client=chat_client)
|
||||
agent = ChatAgent(chat_client=chat_client)
|
||||
|
||||
result = await agent.run("Hello")
|
||||
assert result.text == "test response"
|
||||
|
||||
@@ -9,17 +9,17 @@ from pydantic import Field
|
||||
from pytest import fixture
|
||||
|
||||
from agent_framework import (
|
||||
ChatClient,
|
||||
ChatClientBase,
|
||||
BaseChatClient,
|
||||
ChatClientProtocol,
|
||||
ChatMessage,
|
||||
ChatOptions,
|
||||
ChatResponse,
|
||||
ChatResponseUpdate,
|
||||
ChatRole,
|
||||
EmbeddingGenerator,
|
||||
FunctionCallContent,
|
||||
FunctionResultContent,
|
||||
GeneratedEmbeddings,
|
||||
Role,
|
||||
TextContent,
|
||||
ai_function,
|
||||
use_tool_calling,
|
||||
@@ -54,8 +54,8 @@ class MockChatClient:
|
||||
|
||||
|
||||
@use_tool_calling
|
||||
class MockChatClientBase(ChatClientBase):
|
||||
"""Mock implementation of the ChatClientBase."""
|
||||
class MockBaseChatClient(BaseChatClient):
|
||||
"""Mock implementation of the BaseChatClient."""
|
||||
|
||||
run_responses: list[ChatResponse] = Field(default_factory=list)
|
||||
streaming_responses: list[list[ChatResponseUpdate]] = Field(default_factory=list)
|
||||
@@ -120,8 +120,8 @@ def chat_client() -> MockChatClient:
|
||||
|
||||
|
||||
@fixture
|
||||
def chat_client_base() -> MockChatClientBase:
|
||||
return MockChatClientBase()
|
||||
def chat_client_base() -> MockBaseChatClient:
|
||||
return MockBaseChatClient()
|
||||
|
||||
|
||||
@fixture
|
||||
@@ -131,19 +131,19 @@ def embedding_generator() -> MockEmbeddingGenerator:
|
||||
|
||||
|
||||
def test_chat_client_type(chat_client: MockChatClient):
|
||||
assert isinstance(chat_client, ChatClient)
|
||||
assert isinstance(chat_client, ChatClientProtocol)
|
||||
|
||||
|
||||
async def test_chat_client_get_response(chat_client: MockChatClient):
|
||||
response = await chat_client.get_response(ChatMessage(role="user", text="Hello"))
|
||||
assert response.text == "test response"
|
||||
assert response.messages[0].role == ChatRole.ASSISTANT
|
||||
assert response.messages[0].role == Role.ASSISTANT
|
||||
|
||||
|
||||
async def test_chat_client_get_streaming_response(chat_client: MockChatClient):
|
||||
async for update in chat_client.get_streaming_response(ChatMessage(role="user", text="Hello")):
|
||||
assert update.text == "test streaming response" or update.text == "another update"
|
||||
assert update.role == ChatRole.ASSISTANT
|
||||
assert update.role == Role.ASSISTANT
|
||||
|
||||
|
||||
def test_embedding_generator_type(embedding_generator: MockEmbeddingGenerator):
|
||||
@@ -158,23 +158,23 @@ async def test_embedding_generator_generate(embedding_generator: MockEmbeddingGe
|
||||
assert len(emb) == 5
|
||||
|
||||
|
||||
def test_base_client(chat_client_base: MockChatClientBase):
|
||||
assert isinstance(chat_client_base, ChatClientBase)
|
||||
assert isinstance(chat_client_base, ChatClient)
|
||||
def test_base_client(chat_client_base: MockBaseChatClient):
|
||||
assert isinstance(chat_client_base, BaseChatClient)
|
||||
assert isinstance(chat_client_base, ChatClientProtocol)
|
||||
|
||||
|
||||
async def test_base_client_get_response(chat_client_base: MockChatClientBase):
|
||||
async def test_base_client_get_response(chat_client_base: MockBaseChatClient):
|
||||
response = await chat_client_base.get_response(ChatMessage(role="user", text="Hello"))
|
||||
assert response.messages[0].role == ChatRole.ASSISTANT
|
||||
assert response.messages[0].role == Role.ASSISTANT
|
||||
assert response.messages[0].text == "test response - Hello"
|
||||
|
||||
|
||||
async def test_base_client_get_streaming_response(chat_client_base: MockChatClientBase):
|
||||
async def test_base_client_get_streaming_response(chat_client_base: MockBaseChatClient):
|
||||
async for update in chat_client_base.get_streaming_response(ChatMessage(role="user", text="Hello")):
|
||||
assert update.text == "update - Hello" or update.text == "another update"
|
||||
|
||||
|
||||
async def test_base_client_with_function_calling(chat_client_base: MockChatClientBase):
|
||||
async def test_base_client_with_function_calling(chat_client_base: MockBaseChatClient):
|
||||
exec_counter = 0
|
||||
|
||||
@ai_function(name="test_function")
|
||||
@@ -195,20 +195,20 @@ async def test_base_client_with_function_calling(chat_client_base: MockChatClien
|
||||
response = await chat_client_base.get_response("hello", tool_choice="auto", tools=[ai_func])
|
||||
assert exec_counter == 1
|
||||
assert len(response.messages) == 3
|
||||
assert response.messages[0].role == ChatRole.ASSISTANT
|
||||
assert response.messages[0].role == Role.ASSISTANT
|
||||
assert isinstance(response.messages[0].contents[0], FunctionCallContent)
|
||||
assert response.messages[0].contents[0].name == "test_function"
|
||||
assert response.messages[0].contents[0].arguments == '{"arg1": "value1"}'
|
||||
assert response.messages[0].contents[0].call_id == "1"
|
||||
assert response.messages[1].role == ChatRole.TOOL
|
||||
assert response.messages[1].role == Role.TOOL
|
||||
assert isinstance(response.messages[1].contents[0], FunctionResultContent)
|
||||
assert response.messages[1].contents[0].call_id == "1"
|
||||
assert response.messages[1].contents[0].result == "Processed value1"
|
||||
assert response.messages[2].role == ChatRole.ASSISTANT
|
||||
assert response.messages[2].role == Role.ASSISTANT
|
||||
assert response.messages[2].text == "done"
|
||||
|
||||
|
||||
async def test_base_client_with_function_calling_disabled(chat_client_base: MockChatClientBase):
|
||||
async def test_base_client_with_function_calling_disabled(chat_client_base: MockBaseChatClient):
|
||||
chat_client_base.__maximum_iterations_per_request = 0
|
||||
exec_counter = 0
|
||||
|
||||
@@ -230,11 +230,11 @@ async def test_base_client_with_function_calling_disabled(chat_client_base: Mock
|
||||
response = await chat_client_base.get_response("hello", tool_choice="auto", tools=[ai_func])
|
||||
assert exec_counter == 0
|
||||
assert len(response.messages) == 1
|
||||
assert response.messages[0].role == ChatRole.ASSISTANT
|
||||
assert response.messages[0].role == Role.ASSISTANT
|
||||
assert response.messages[0].text == "test response - hello"
|
||||
|
||||
|
||||
async def test_base_client_with_streaming_function_calling(chat_client_base: MockChatClientBase):
|
||||
async def test_base_client_with_streaming_function_calling(chat_client_base: MockBaseChatClient):
|
||||
exec_counter = 0
|
||||
|
||||
@ai_function(name="test_function")
|
||||
@@ -272,7 +272,7 @@ async def test_base_client_with_streaming_function_calling(chat_client_base: Moc
|
||||
assert exec_counter == 1
|
||||
|
||||
|
||||
async def test_base_client_with_streaming_function_calling_disabled(chat_client_base: MockChatClientBase):
|
||||
async def test_base_client_with_streaming_function_calling_disabled(chat_client_base: MockBaseChatClient):
|
||||
chat_client_base.__maximum_iterations_per_request = 0
|
||||
exec_counter = 0
|
||||
|
||||
|
||||
@@ -12,19 +12,18 @@ from mcp.shared.exceptions import McpError
|
||||
from pydantic import AnyUrl, ValidationError
|
||||
|
||||
from agent_framework import (
|
||||
AITool,
|
||||
ChatMessage,
|
||||
ChatRole,
|
||||
DataContent,
|
||||
McpSseTools,
|
||||
McpStdioTool,
|
||||
McpStreamableHttpTool,
|
||||
McpWebsocketTool,
|
||||
MCPStdioTool,
|
||||
MCPStreamableHTTPTool,
|
||||
MCPWebsocketTool,
|
||||
Role,
|
||||
TextContent,
|
||||
ToolProtocol,
|
||||
UriContent,
|
||||
)
|
||||
from agent_framework._mcp import (
|
||||
McpTool,
|
||||
MCPTool,
|
||||
_ai_content_to_mcp_types,
|
||||
_chat_message_to_mcp_types,
|
||||
_get_input_model_from_mcp_prompt,
|
||||
@@ -275,20 +274,20 @@ def test_get_input_model_from_mcp_prompt():
|
||||
model(arg2="optional")
|
||||
|
||||
|
||||
# McpTool tests
|
||||
# MCPTool tests
|
||||
async def test_local_mcp_server_initialization():
|
||||
"""Test McpTool initialization."""
|
||||
server = McpTool(name="test_server")
|
||||
assert isinstance(server, AITool)
|
||||
"""Test MCPTool initialization."""
|
||||
server = MCPTool(name="test_server")
|
||||
assert isinstance(server, ToolProtocol)
|
||||
assert server.name == "test_server"
|
||||
assert server.session is None
|
||||
assert server.functions == []
|
||||
|
||||
|
||||
async def test_local_mcp_server_context_manager():
|
||||
"""Test McpTool as context manager."""
|
||||
"""Test MCPTool as context manager."""
|
||||
|
||||
class TestServer(McpTool):
|
||||
class TestServer(MCPTool):
|
||||
async def connect(self):
|
||||
# Mock connection
|
||||
self.session = Mock(spec=ClientSession)
|
||||
@@ -306,7 +305,7 @@ async def test_local_mcp_server_context_manager():
|
||||
async def test_local_mcp_server_load_functions():
|
||||
"""Test loading functions from MCP server."""
|
||||
|
||||
class TestServer(McpTool):
|
||||
class TestServer(MCPTool):
|
||||
async def connect(self):
|
||||
self.session = Mock(spec=ClientSession)
|
||||
# Mock tools list response
|
||||
@@ -330,7 +329,7 @@ async def test_local_mcp_server_load_functions():
|
||||
return None
|
||||
|
||||
server = TestServer(name="test_server")
|
||||
assert isinstance(server, AITool)
|
||||
assert isinstance(server, ToolProtocol)
|
||||
async with server:
|
||||
await server.load_tools()
|
||||
assert len(server.functions) == 1
|
||||
@@ -340,7 +339,7 @@ async def test_local_mcp_server_load_functions():
|
||||
async def test_local_mcp_server_load_prompts():
|
||||
"""Test loading prompts from MCP server."""
|
||||
|
||||
class TestServer(McpTool):
|
||||
class TestServer(MCPTool):
|
||||
async def connect(self):
|
||||
self.session = Mock(spec=ClientSession)
|
||||
# Mock prompts list response
|
||||
@@ -369,7 +368,7 @@ async def test_local_mcp_server_load_prompts():
|
||||
async def test_local_mcp_server_function_execution():
|
||||
"""Test function execution through MCP server."""
|
||||
|
||||
class TestServer(McpTool):
|
||||
class TestServer(MCPTool):
|
||||
async def connect(self):
|
||||
self.session = Mock(spec=ClientSession)
|
||||
self.session.list_tools = AsyncMock(
|
||||
@@ -410,7 +409,7 @@ async def test_local_mcp_server_function_execution():
|
||||
async def test_local_mcp_server_function_execution_error():
|
||||
"""Test function execution error handling."""
|
||||
|
||||
class TestServer(McpTool):
|
||||
class TestServer(MCPTool):
|
||||
async def connect(self):
|
||||
self.session = Mock(spec=ClientSession)
|
||||
self.session.list_tools = AsyncMock(
|
||||
@@ -448,7 +447,7 @@ async def test_local_mcp_server_function_execution_error():
|
||||
async def test_local_mcp_server_prompt_execution():
|
||||
"""Test prompt execution through MCP server."""
|
||||
|
||||
class TestMcpTool(McpTool):
|
||||
class TestMCPTool(MCPTool):
|
||||
async def connect(self):
|
||||
self.session = Mock(spec=ClientSession)
|
||||
self.session.list_prompts = AsyncMock(
|
||||
@@ -474,7 +473,7 @@ async def test_local_mcp_server_prompt_execution():
|
||||
def get_mcp_client(self) -> _AsyncGeneratorContextManager[Any, None]:
|
||||
return None
|
||||
|
||||
server = TestMcpTool(name="test_server")
|
||||
server = TestMCPTool(name="test_server")
|
||||
async with server:
|
||||
await server.load_prompts()
|
||||
prompt = server.functions[0]
|
||||
@@ -482,37 +481,30 @@ async def test_local_mcp_server_prompt_execution():
|
||||
|
||||
assert len(result) == 1
|
||||
assert isinstance(result[0], ChatMessage)
|
||||
assert result[0].role == ChatRole.USER
|
||||
assert result[0].role == Role.USER
|
||||
assert len(result[0].contents) == 1
|
||||
assert result[0].contents[0].text == "Test message"
|
||||
|
||||
|
||||
# Server implementation tests
|
||||
def test_local_mcp_stdio_tool_init():
|
||||
"""Test McpStdioTool initialization."""
|
||||
tool = McpStdioTool(name="test", command="echo", args=["hello"])
|
||||
"""Test MCPStdioTool initialization."""
|
||||
tool = MCPStdioTool(name="test", command="echo", args=["hello"])
|
||||
assert tool.name == "test"
|
||||
assert tool.command == "echo"
|
||||
assert tool.args == ["hello"]
|
||||
|
||||
|
||||
def test_local_mcp_sse_tools_init():
|
||||
"""Test McpSseTools initialization."""
|
||||
tool = McpSseTools(name="test", url="http://localhost:8080")
|
||||
assert tool.name == "test"
|
||||
assert tool.url == "http://localhost:8080"
|
||||
|
||||
|
||||
def test_local_mcp_websocket_tool_init():
|
||||
"""Test McpWebsocketTool initialization."""
|
||||
tool = McpWebsocketTool(name="test", url="ws://localhost:8080")
|
||||
"""Test MCPWebsocketTool initialization."""
|
||||
tool = MCPWebsocketTool(name="test", url="ws://localhost:8080")
|
||||
assert tool.name == "test"
|
||||
assert tool.url == "ws://localhost:8080"
|
||||
|
||||
|
||||
def test_local_mcp_streamable_http_tool_init():
|
||||
"""Test McpStreamableHttpTool initialization."""
|
||||
tool = McpStreamableHttpTool(name="test", url="http://localhost:8080")
|
||||
"""Test MCPStreamableHTTPTool initialization."""
|
||||
tool = MCPStreamableHTTPTool(name="test", url="http://localhost:8080")
|
||||
assert tool.name == "test"
|
||||
assert tool.url == "http://localhost:8080"
|
||||
|
||||
@@ -525,7 +517,7 @@ async def test_streamable_http_integration():
|
||||
if not url.startswith("http"):
|
||||
pytest.skip("LOCAL_MCP_URL is not an HTTP URL")
|
||||
|
||||
tool = McpStreamableHttpTool(name="integration_test", url=url)
|
||||
tool = MCPStreamableHTTPTool(name="integration_test", url=url)
|
||||
|
||||
async with tool:
|
||||
# Test that we can connect and load tools
|
||||
|
||||
@@ -13,7 +13,7 @@ from agent_framework import (
|
||||
ChatOptions,
|
||||
ChatResponse,
|
||||
ChatResponseUpdate,
|
||||
ChatRole,
|
||||
Role,
|
||||
UsageDetails,
|
||||
)
|
||||
from agent_framework.telemetry import (
|
||||
@@ -301,7 +301,7 @@ def test_start_span_empty_metadata():
|
||||
|
||||
|
||||
def test_decorator_with_valid_class():
|
||||
"""Test that decorator works with a valid ChatClientBase-like class."""
|
||||
"""Test that decorator works with a valid BaseChatClient-like class."""
|
||||
|
||||
# Create a mock class with the required methods
|
||||
class MockChatClient:
|
||||
@@ -373,7 +373,7 @@ def mock_chat_client():
|
||||
self, *, messages: MutableSequence[ChatMessage], chat_options: ChatOptions, **kwargs: Any
|
||||
):
|
||||
return ChatResponse(
|
||||
messages=[ChatMessage(role=ChatRole.ASSISTANT, text="Test response")],
|
||||
messages=[ChatMessage(role=Role.ASSISTANT, text="Test response")],
|
||||
usage_details=UsageDetails(input_token_count=10, output_token_count=20),
|
||||
finish_reason=None,
|
||||
)
|
||||
@@ -381,8 +381,8 @@ def mock_chat_client():
|
||||
async def _inner_get_streaming_response(
|
||||
self, *, messages: MutableSequence[ChatMessage], chat_options: ChatOptions, **kwargs: Any
|
||||
):
|
||||
yield ChatResponseUpdate(text="Hello", role=ChatRole.ASSISTANT)
|
||||
yield ChatResponseUpdate(text=" world", role=ChatRole.ASSISTANT)
|
||||
yield ChatResponseUpdate(text="Hello", role=Role.ASSISTANT)
|
||||
yield ChatResponseUpdate(text=" world", role=Role.ASSISTANT)
|
||||
|
||||
return MockChatClient()
|
||||
|
||||
@@ -393,7 +393,7 @@ async def test_telemetry_disabled_bypasses_instrumentation(mock_chat_client, mod
|
||||
decorated_class = use_telemetry(type(mock_chat_client))
|
||||
client = decorated_class()
|
||||
|
||||
messages = [ChatMessage(role=ChatRole.USER, text="Test message")]
|
||||
messages = [ChatMessage(role=Role.USER, text="Test message")]
|
||||
chat_options = ChatOptions()
|
||||
|
||||
with (
|
||||
@@ -412,7 +412,7 @@ async def test_instrumentation_enabled(mock_chat_client, model_diagnostic_settin
|
||||
decorated_class = use_telemetry(type(mock_chat_client))
|
||||
client = decorated_class()
|
||||
|
||||
messages = [ChatMessage(role=ChatRole.USER, text="Test message")]
|
||||
messages = [ChatMessage(role=Role.USER, text="Test message")]
|
||||
chat_options = ChatOptions()
|
||||
|
||||
with (
|
||||
@@ -432,7 +432,7 @@ async def test_streaming_response_with_diagnostics_enabled_via_decorator(mock_ch
|
||||
"""Test streaming telemetry through the use_telemetry decorator."""
|
||||
decorated_class = use_telemetry(type(mock_chat_client))
|
||||
client = decorated_class()
|
||||
messages = [ChatMessage(role=ChatRole.USER, text="Test")]
|
||||
messages = [ChatMessage(role=Role.USER, text="Test")]
|
||||
chat_options = ChatOptions()
|
||||
|
||||
with (
|
||||
@@ -470,7 +470,7 @@ async def test_streaming_response_with_exception_via_decorator(mock_chat_client,
|
||||
async def _inner_get_streaming_response(
|
||||
self, *, messages: MutableSequence[ChatMessage], chat_options: ChatOptions, **kwargs: Any
|
||||
) -> AsyncIterable[ChatResponseUpdate]:
|
||||
yield ChatResponseUpdate(text="Partial", role=ChatRole.ASSISTANT)
|
||||
yield ChatResponseUpdate(text="Partial", role=Role.ASSISTANT)
|
||||
raise ValueError("Test streaming error")
|
||||
|
||||
type(mock_chat_client)._inner_get_streaming_response = _inner_get_streaming_response
|
||||
@@ -478,7 +478,7 @@ async def test_streaming_response_with_exception_via_decorator(mock_chat_client,
|
||||
decorated_class = use_telemetry(type(mock_chat_client))
|
||||
client = decorated_class()
|
||||
|
||||
messages = [ChatMessage(role=ChatRole.USER, text="Test")]
|
||||
messages = [ChatMessage(role=Role.USER, text="Test")]
|
||||
chat_options = ChatOptions()
|
||||
|
||||
with (
|
||||
@@ -513,12 +513,12 @@ async def test_streaming_response_diagnostics_disabled_via_decorator(model_diagn
|
||||
async def _inner_get_streaming_response(
|
||||
self, *, messages: MutableSequence[ChatMessage], chat_options: ChatOptions, **kwargs: Any
|
||||
) -> AsyncIterable[ChatResponseUpdate]:
|
||||
yield ChatResponseUpdate(text="Test", role=ChatRole.ASSISTANT)
|
||||
yield ChatResponseUpdate(text="Test", role=Role.ASSISTANT)
|
||||
|
||||
decorated_class = use_telemetry(MockStreamingClientNoDiagnostics)
|
||||
client = decorated_class()
|
||||
|
||||
messages = [ChatMessage(role=ChatRole.USER, text="Test")]
|
||||
messages = [ChatMessage(role=Role.USER, text="Test")]
|
||||
chat_options = ChatOptions()
|
||||
|
||||
with (
|
||||
@@ -561,7 +561,7 @@ async def test_empty_streaming_response_via_decorator(model_diagnostic_settings)
|
||||
decorated_class = use_telemetry(MockEmptyStreamingClient)
|
||||
client = decorated_class()
|
||||
|
||||
messages = [ChatMessage(role=ChatRole.USER, text="Test")]
|
||||
messages = [ChatMessage(role=Role.USER, text="Test")]
|
||||
chat_options = ChatOptions()
|
||||
|
||||
with (
|
||||
@@ -617,7 +617,7 @@ def test_prepend_user_agent_with_none_value():
|
||||
|
||||
|
||||
def test_agent_decorator_with_valid_class():
|
||||
"""Test that agent decorator works with a valid ChatClientAgent-like class."""
|
||||
"""Test that agent decorator works with a valid ChatAgent-like class."""
|
||||
from agent_framework.telemetry import use_agent_telemetry
|
||||
|
||||
# Create a mock class with the required methods
|
||||
@@ -633,7 +633,7 @@ def test_agent_decorator_with_valid_class():
|
||||
async def run(self, messages=None, *, thread=None, **kwargs):
|
||||
return Mock()
|
||||
|
||||
async def run_streaming(self, messages=None, *, thread=None, **kwargs):
|
||||
async def run_stream(self, messages=None, *, thread=None, **kwargs):
|
||||
async def gen():
|
||||
yield Mock()
|
||||
|
||||
@@ -644,7 +644,7 @@ def test_agent_decorator_with_valid_class():
|
||||
|
||||
# Check that the methods were wrapped
|
||||
assert hasattr(decorated_class.run, "__model_diagnostics_agent_run__")
|
||||
assert hasattr(decorated_class.run_streaming, "__model_diagnostics_streaming_agent_run__")
|
||||
assert hasattr(decorated_class.run_stream, "__model_diagnostics_streaming_agent_run__")
|
||||
|
||||
|
||||
def test_agent_decorator_with_missing_methods():
|
||||
@@ -680,7 +680,7 @@ def test_agent_decorator_with_partial_methods():
|
||||
|
||||
# Only the present method should be wrapped
|
||||
assert hasattr(decorated_class.run, "__model_diagnostics_agent_run__")
|
||||
assert not hasattr(decorated_class, "run_streaming")
|
||||
assert not hasattr(decorated_class, "run_stream")
|
||||
|
||||
|
||||
# region Test agent telemetry decorator with mock agent
|
||||
@@ -689,7 +689,7 @@ def test_agent_decorator_with_partial_methods():
|
||||
@pytest.fixture
|
||||
def mock_chat_client_agent():
|
||||
"""Create a mock chat client agent for testing."""
|
||||
from agent_framework import AgentRunResponse, ChatMessage, ChatRole, UsageDetails
|
||||
from agent_framework import AgentRunResponse, ChatMessage, Role, UsageDetails
|
||||
|
||||
class MockChatClientAgent:
|
||||
AGENT_SYSTEM_NAME = "test_agent_system"
|
||||
@@ -702,17 +702,17 @@ def mock_chat_client_agent():
|
||||
|
||||
async def run(self, messages=None, *, thread=None, **kwargs):
|
||||
return AgentRunResponse(
|
||||
messages=[ChatMessage(role=ChatRole.ASSISTANT, text="Agent response")],
|
||||
messages=[ChatMessage(role=Role.ASSISTANT, text="Agent response")],
|
||||
usage_details=UsageDetails(input_token_count=15, output_token_count=25),
|
||||
response_id="test_response_id",
|
||||
raw_representation=Mock(finish_reason=Mock(value="stop")),
|
||||
)
|
||||
|
||||
async def run_streaming(self, messages=None, *, thread=None, **kwargs):
|
||||
async def run_stream(self, messages=None, *, thread=None, **kwargs):
|
||||
from agent_framework import AgentRunResponseUpdate
|
||||
|
||||
yield AgentRunResponseUpdate(text="Hello", role=ChatRole.ASSISTANT)
|
||||
yield AgentRunResponseUpdate(text=" from agent", role=ChatRole.ASSISTANT)
|
||||
yield AgentRunResponseUpdate(text="Hello", role=Role.ASSISTANT)
|
||||
yield AgentRunResponseUpdate(text=" from agent", role=Role.ASSISTANT)
|
||||
|
||||
return MockChatClientAgent()
|
||||
|
||||
@@ -778,7 +778,7 @@ async def test_agent_streaming_response_with_diagnostics_enabled_via_decorator(
|
||||
|
||||
# Collect all yielded updates
|
||||
updates = []
|
||||
async for update in agent.run_streaming("Test message"):
|
||||
async for update in agent.run_stream("Test message"):
|
||||
updates.append(update)
|
||||
|
||||
# Verify we got the expected updates
|
||||
@@ -795,13 +795,13 @@ async def test_agent_streaming_response_with_exception_via_decorator(mock_chat_c
|
||||
"""Test agent streaming telemetry exception handling through decorator."""
|
||||
from agent_framework.telemetry import use_agent_telemetry
|
||||
|
||||
async def run_streaming(self, messages=None, *, thread=None, **kwargs):
|
||||
from agent_framework import AgentRunResponseUpdate, ChatRole
|
||||
async def run_stream(self, messages=None, *, thread=None, **kwargs):
|
||||
from agent_framework import AgentRunResponseUpdate, Role
|
||||
|
||||
yield AgentRunResponseUpdate(text="Partial", role=ChatRole.ASSISTANT)
|
||||
yield AgentRunResponseUpdate(text="Partial", role=Role.ASSISTANT)
|
||||
raise ValueError("Test agent streaming error")
|
||||
|
||||
type(mock_chat_client_agent).run_streaming = run_streaming
|
||||
type(mock_chat_client_agent).run_stream = run_stream
|
||||
|
||||
decorated_class = use_agent_telemetry(type(mock_chat_client_agent))
|
||||
agent = decorated_class()
|
||||
@@ -819,7 +819,7 @@ async def test_agent_streaming_response_with_exception_via_decorator(mock_chat_c
|
||||
|
||||
# Should raise the exception and call error handler
|
||||
with pytest.raises(ValueError, match="Test agent streaming error"):
|
||||
async for _ in agent.run_streaming("Test message"):
|
||||
async for _ in agent.run_stream("Test message"):
|
||||
pass
|
||||
|
||||
# Verify error was recorded
|
||||
@@ -830,7 +830,7 @@ async def test_agent_streaming_response_with_exception_via_decorator(mock_chat_c
|
||||
@pytest.mark.parametrize("model_diagnostic_settings", [(False, False)], indirect=True)
|
||||
async def test_agent_streaming_response_diagnostics_disabled_via_decorator(model_diagnostic_settings):
|
||||
"""Test agent streaming response when diagnostics are disabled."""
|
||||
from agent_framework import AgentRunResponseUpdate, ChatRole
|
||||
from agent_framework import AgentRunResponseUpdate, Role
|
||||
from agent_framework.telemetry import use_agent_telemetry
|
||||
|
||||
class MockStreamingAgentNoDiagnostics:
|
||||
@@ -841,8 +841,8 @@ async def test_agent_streaming_response_diagnostics_disabled_via_decorator(model
|
||||
self.name = "test_agent"
|
||||
self.display_name = "Test Agent"
|
||||
|
||||
async def run_streaming(self, messages=None, *, thread=None, **kwargs):
|
||||
yield AgentRunResponseUpdate(text="Test", role=ChatRole.ASSISTANT)
|
||||
async def run_stream(self, messages=None, *, thread=None, **kwargs):
|
||||
yield AgentRunResponseUpdate(text="Test", role=Role.ASSISTANT)
|
||||
|
||||
decorated_class = use_agent_telemetry(MockStreamingAgentNoDiagnostics)
|
||||
agent = decorated_class()
|
||||
@@ -853,7 +853,7 @@ async def test_agent_streaming_response_diagnostics_disabled_via_decorator(model
|
||||
):
|
||||
# Should not create spans when diagnostics are disabled
|
||||
updates = []
|
||||
async for update in agent.run_streaming("Test message"):
|
||||
async for update in agent.run_stream("Test message"):
|
||||
updates.append(update)
|
||||
|
||||
assert len(updates) == 1
|
||||
@@ -874,7 +874,7 @@ async def test_agent_empty_streaming_response_via_decorator(model_diagnostic_set
|
||||
self.name = "test_agent"
|
||||
self.display_name = "Test Agent"
|
||||
|
||||
async def run_streaming(self, messages=None, *, thread=None, **kwargs):
|
||||
async def run_stream(self, messages=None, *, thread=None, **kwargs):
|
||||
# Return empty stream
|
||||
return
|
||||
yield # This will never be reached
|
||||
@@ -895,7 +895,7 @@ async def test_agent_empty_streaming_response_via_decorator(model_diagnostic_set
|
||||
|
||||
# Should handle empty stream gracefully
|
||||
updates = []
|
||||
async for update in agent.run_streaming("Test message"):
|
||||
async for update in agent.run_stream("Test message"):
|
||||
updates.append(update)
|
||||
|
||||
assert len(updates) == 0
|
||||
@@ -943,16 +943,16 @@ async def test_agent_run_with_thread_and_kwargs(mock_chat_client_agent, model_di
|
||||
@pytest.mark.parametrize("model_diagnostic_settings", [(True, False)], indirect=True)
|
||||
async def test_agent_run_with_list_messages(mock_chat_client_agent, model_diagnostic_settings):
|
||||
"""Test agent run with list of messages."""
|
||||
from agent_framework import ChatMessage, ChatRole
|
||||
from agent_framework import ChatMessage, Role
|
||||
from agent_framework.telemetry import use_agent_telemetry
|
||||
|
||||
decorated_class = use_agent_telemetry(type(mock_chat_client_agent))
|
||||
agent = decorated_class()
|
||||
|
||||
messages = [
|
||||
ChatMessage(role=ChatRole.USER, text="First message"),
|
||||
ChatMessage(role=ChatRole.ASSISTANT, text="Response"),
|
||||
ChatMessage(role=ChatRole.USER, text="Second message"),
|
||||
ChatMessage(role=Role.USER, text="First message"),
|
||||
ChatMessage(role=Role.ASSISTANT, text="Response"),
|
||||
ChatMessage(role=Role.USER, text="Second message"),
|
||||
]
|
||||
|
||||
with (
|
||||
|
||||
@@ -5,7 +5,7 @@ from typing import Any
|
||||
|
||||
import pytest
|
||||
|
||||
from agent_framework import AgentThread, ChatMessage, ChatMessageList, ChatRole
|
||||
from agent_framework import AgentThread, ChatMessage, ChatMessageList, Role
|
||||
from agent_framework._threads import StoreState, ThreadState, deserialize_thread_state, thread_on_new_messages
|
||||
|
||||
|
||||
@@ -37,16 +37,16 @@ class MockChatMessageStore:
|
||||
def sample_messages() -> list[ChatMessage]:
|
||||
"""Fixture providing sample chat messages for testing."""
|
||||
return [
|
||||
ChatMessage(role=ChatRole.USER, text="Hello", message_id="msg1"),
|
||||
ChatMessage(role=ChatRole.ASSISTANT, text="Hi there!", message_id="msg2"),
|
||||
ChatMessage(role=ChatRole.USER, text="How are you?", message_id="msg3"),
|
||||
ChatMessage(role=Role.USER, text="Hello", message_id="msg1"),
|
||||
ChatMessage(role=Role.ASSISTANT, text="Hi there!", message_id="msg2"),
|
||||
ChatMessage(role=Role.USER, text="How are you?", message_id="msg3"),
|
||||
]
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def sample_message() -> ChatMessage:
|
||||
"""Fixture providing a single sample chat message for testing."""
|
||||
return ChatMessage(role=ChatRole.USER, text="Test message", message_id="test1")
|
||||
return ChatMessage(role=Role.USER, text="Test message", message_id="test1")
|
||||
|
||||
|
||||
class TestAgentThread:
|
||||
@@ -171,7 +171,7 @@ class TestAgentThread:
|
||||
|
||||
async def test_on_new_messages_with_existing_store(self, sample_message: ChatMessage) -> None:
|
||||
"""Test _on_new_messages adds to existing message store."""
|
||||
initial_messages = [ChatMessage(role=ChatRole.USER, text="Initial", message_id="init1")]
|
||||
initial_messages = [ChatMessage(role=Role.USER, text="Initial", message_id="init1")]
|
||||
store = ChatMessageList(initial_messages)
|
||||
thread = AgentThread(message_store=store)
|
||||
|
||||
|
||||
@@ -5,7 +5,7 @@ from unittest.mock import Mock, patch
|
||||
import pytest
|
||||
from pydantic import BaseModel
|
||||
|
||||
from agent_framework import AIFunction, AITool, HostedCodeInterpreterTool, ai_function
|
||||
from agent_framework import AIFunction, HostedCodeInterpreterTool, ToolProtocol, ai_function
|
||||
from agent_framework._tools import _parse_inputs
|
||||
from agent_framework.telemetry import GenAIAttributes
|
||||
|
||||
@@ -18,7 +18,7 @@ def test_ai_function_decorator():
|
||||
"""A simple function that adds two numbers."""
|
||||
return x + y
|
||||
|
||||
assert isinstance(test_tool, AITool)
|
||||
assert isinstance(test_tool, ToolProtocol)
|
||||
assert isinstance(test_tool, AIFunction)
|
||||
assert test_tool.name == "test_tool"
|
||||
assert test_tool.description == "A test tool"
|
||||
@@ -39,7 +39,7 @@ def test_ai_function_decorator_without_args():
|
||||
"""A simple function that adds two numbers."""
|
||||
return x + y
|
||||
|
||||
assert isinstance(test_tool, AITool)
|
||||
assert isinstance(test_tool, ToolProtocol)
|
||||
assert isinstance(test_tool, AIFunction)
|
||||
assert test_tool.name == "test_tool"
|
||||
assert test_tool.description == "A simple function that adds two numbers."
|
||||
@@ -60,7 +60,7 @@ async def test_ai_function_decorator_with_async():
|
||||
"""An async function that adds two numbers."""
|
||||
return x + y
|
||||
|
||||
assert isinstance(async_test_tool, AITool)
|
||||
assert isinstance(async_test_tool, ToolProtocol)
|
||||
assert isinstance(async_test_tool, AIFunction)
|
||||
assert async_test_tool.name == "async_test_tool"
|
||||
assert async_test_tool.description == "An async test tool"
|
||||
@@ -399,7 +399,7 @@ def test_parse_inputs_data_dict():
|
||||
|
||||
|
||||
def test_parse_inputs_ai_contents_instance():
|
||||
"""Test _parse_inputs with AIContents instance."""
|
||||
"""Test _parse_inputs with Contents instance."""
|
||||
from agent_framework import TextContent
|
||||
|
||||
text_content = TextContent(text="Hello, world!")
|
||||
@@ -418,7 +418,7 @@ def test_parse_inputs_mixed_list():
|
||||
"http://example.com", # string
|
||||
{"uri": "https://test.org", "media_type": "text/html"}, # URI dict
|
||||
{"file_id": "file-456"}, # hosted file dict
|
||||
TextContent(text="Hello"), # AIContents instance
|
||||
TextContent(text="Hello"), # Contents instance
|
||||
]
|
||||
|
||||
result = _parse_inputs(inputs)
|
||||
@@ -477,7 +477,7 @@ def test_hosted_code_interpreter_tool_with_dict_inputs():
|
||||
|
||||
|
||||
def test_hosted_code_interpreter_tool_with_ai_contents():
|
||||
"""Test HostedCodeInterpreterTool with AIContents instances."""
|
||||
"""Test HostedCodeInterpreterTool with Contents instances."""
|
||||
from agent_framework import DataContent, TextContent
|
||||
|
||||
inputs = [TextContent(text="Hello, world!"), DataContent(data=b"test", media_type="text/plain")]
|
||||
|
||||
@@ -9,32 +9,30 @@ from pytest import fixture, mark, raises
|
||||
from agent_framework import (
|
||||
AgentRunResponse,
|
||||
AgentRunResponseUpdate,
|
||||
AIAnnotation,
|
||||
AIContent,
|
||||
AIContents,
|
||||
AIFunction,
|
||||
AITool,
|
||||
AnnotatedRegion,
|
||||
ChatFinishReason,
|
||||
BaseContent,
|
||||
ChatMessage,
|
||||
ChatOptions,
|
||||
ChatResponse,
|
||||
ChatResponseUpdate,
|
||||
ChatRole,
|
||||
ChatToolMode,
|
||||
CitationAnnotation,
|
||||
Contents,
|
||||
DataContent,
|
||||
ErrorContent,
|
||||
FinishReason,
|
||||
FunctionCallContent,
|
||||
FunctionResultContent,
|
||||
GeneratedEmbeddings,
|
||||
HostedFileContent,
|
||||
HostedVectorStoreContent,
|
||||
Role,
|
||||
SpeechToTextOptions,
|
||||
TextContent,
|
||||
TextReasoningContent,
|
||||
TextSpanRegion,
|
||||
TextToSpeechOptions,
|
||||
ToolProtocol,
|
||||
UriContent,
|
||||
UsageContent,
|
||||
UsageDetails,
|
||||
@@ -43,8 +41,8 @@ from agent_framework import (
|
||||
|
||||
|
||||
@fixture
|
||||
def ai_tool() -> AITool:
|
||||
"""Returns a generic AITool."""
|
||||
def ai_tool() -> ToolProtocol:
|
||||
"""Returns a generic ToolProtocol."""
|
||||
|
||||
class GenericTool(BaseModel):
|
||||
name: str
|
||||
@@ -61,8 +59,8 @@ def ai_tool() -> AITool:
|
||||
|
||||
|
||||
@fixture
|
||||
def ai_function_tool() -> AITool:
|
||||
"""Returns a executable AITool."""
|
||||
def ai_function_tool() -> ToolProtocol:
|
||||
"""Returns a executable ToolProtocol."""
|
||||
|
||||
@ai_function
|
||||
def simple_function(x: int, y: int) -> int:
|
||||
@@ -76,7 +74,7 @@ def ai_function_tool() -> AITool:
|
||||
|
||||
|
||||
def test_text_content_positional():
|
||||
"""Test the TextContent class to ensure it initializes correctly and inherits from AIContent."""
|
||||
"""Test the TextContent class to ensure it initializes correctly and inherits from BaseContent."""
|
||||
# Create an instance of TextContent
|
||||
content = TextContent("Hello, world!", raw_representation="Hello, world!", additional_properties={"version": 1})
|
||||
|
||||
@@ -85,14 +83,14 @@ def test_text_content_positional():
|
||||
assert content.text == "Hello, world!"
|
||||
assert content.raw_representation == "Hello, world!"
|
||||
assert content.additional_properties["version"] == 1
|
||||
# Ensure the instance is of type AIContent
|
||||
assert isinstance(content, AIContent)
|
||||
# Ensure the instance is of type BaseContent
|
||||
assert isinstance(content, BaseContent)
|
||||
with raises(ValidationError):
|
||||
content.type = "ai"
|
||||
|
||||
|
||||
def test_text_content_keyword():
|
||||
"""Test the TextContent class to ensure it initializes correctly and inherits from AIContent."""
|
||||
"""Test the TextContent class to ensure it initializes correctly and inherits from BaseContent."""
|
||||
# Create an instance of TextContent
|
||||
content = TextContent(
|
||||
text="Hello, world!", raw_representation="Hello, world!", additional_properties={"version": 1}
|
||||
@@ -103,8 +101,8 @@ def test_text_content_keyword():
|
||||
assert content.text == "Hello, world!"
|
||||
assert content.raw_representation == "Hello, world!"
|
||||
assert content.additional_properties["version"] == 1
|
||||
# Ensure the instance is of type AIContent
|
||||
assert isinstance(content, AIContent)
|
||||
# Ensure the instance is of type BaseContent
|
||||
assert isinstance(content, BaseContent)
|
||||
with raises(ValidationError):
|
||||
content.type = "ai"
|
||||
|
||||
@@ -124,8 +122,8 @@ def test_data_content_bytes():
|
||||
assert content.has_top_level_media_type("image") is False
|
||||
assert content.additional_properties["version"] == 1
|
||||
|
||||
# Ensure the instance is of type AIContent
|
||||
assert isinstance(content, AIContent)
|
||||
# Ensure the instance is of type BaseContent
|
||||
assert isinstance(content, BaseContent)
|
||||
|
||||
|
||||
def test_data_content_uri():
|
||||
@@ -140,8 +138,8 @@ def test_data_content_uri():
|
||||
assert content.has_top_level_media_type("application") is False
|
||||
assert content.additional_properties["version"] == 1
|
||||
|
||||
# Ensure the instance is of type AIContent
|
||||
assert isinstance(content, AIContent)
|
||||
# Ensure the instance is of type BaseContent
|
||||
assert isinstance(content, BaseContent)
|
||||
|
||||
|
||||
def test_data_content_invalid():
|
||||
@@ -185,8 +183,8 @@ def test_uri_content():
|
||||
assert content.has_top_level_media_type("application") is False
|
||||
assert content.additional_properties["version"] == 1
|
||||
|
||||
# Ensure the instance is of type AIContent
|
||||
assert isinstance(content, AIContent)
|
||||
# Ensure the instance is of type BaseContent
|
||||
assert isinstance(content, BaseContent)
|
||||
|
||||
|
||||
# region: HostedFileContent
|
||||
@@ -201,8 +199,8 @@ def test_hosted_file_content():
|
||||
assert content.file_id == "file-123"
|
||||
assert content.additional_properties["version"] == 1
|
||||
|
||||
# Ensure the instance is of type AIContent
|
||||
assert isinstance(content, AIContent)
|
||||
# Ensure the instance is of type BaseContent
|
||||
assert isinstance(content, BaseContent)
|
||||
|
||||
|
||||
def test_hosted_file_content_minimal():
|
||||
@@ -215,8 +213,8 @@ def test_hosted_file_content_minimal():
|
||||
assert content.additional_properties is None
|
||||
assert content.raw_representation is None
|
||||
|
||||
# Ensure the instance is of type AIContent
|
||||
assert isinstance(content, AIContent)
|
||||
# Ensure the instance is of type BaseContent
|
||||
assert isinstance(content, BaseContent)
|
||||
|
||||
|
||||
# region: HostedVectorStoreContent
|
||||
@@ -231,9 +229,9 @@ def test_hosted_vector_store_content():
|
||||
assert content.vector_store_id == "vs-789"
|
||||
assert content.additional_properties["version"] == 1
|
||||
|
||||
# Ensure the instance is of type AIContent
|
||||
# Ensure the instance is of type BaseContent
|
||||
assert isinstance(content, HostedVectorStoreContent)
|
||||
assert isinstance(content, AIContent)
|
||||
assert isinstance(content, BaseContent)
|
||||
|
||||
|
||||
def test_hosted_vector_store_content_minimal():
|
||||
@@ -246,9 +244,9 @@ def test_hosted_vector_store_content_minimal():
|
||||
assert content.additional_properties is None
|
||||
assert content.raw_representation is None
|
||||
|
||||
# Ensure the instance is of type AIContent
|
||||
# Ensure the instance is of type BaseContent
|
||||
assert isinstance(content, HostedVectorStoreContent)
|
||||
assert isinstance(content, AIContent)
|
||||
assert isinstance(content, BaseContent)
|
||||
|
||||
|
||||
# region FunctionCallContent
|
||||
@@ -263,8 +261,8 @@ def test_function_call_content():
|
||||
assert content.name == "example_function"
|
||||
assert content.arguments == {"param1": "value1"}
|
||||
|
||||
# Ensure the instance is of type AIContent
|
||||
assert isinstance(content, AIContent)
|
||||
# Ensure the instance is of type BaseContent
|
||||
assert isinstance(content, BaseContent)
|
||||
|
||||
|
||||
def test_function_call_content_parse_arguments():
|
||||
@@ -315,8 +313,8 @@ def test_function_result_content():
|
||||
assert content.type == "function_result"
|
||||
assert content.result == {"param1": "value1"}
|
||||
|
||||
# Ensure the instance is of type AIContent
|
||||
assert isinstance(content, AIContent)
|
||||
# Ensure the instance is of type BaseContent
|
||||
assert isinstance(content, BaseContent)
|
||||
|
||||
|
||||
# region UsageDetails
|
||||
@@ -381,7 +379,7 @@ def test_usage_details_add_with_none_and_type_errors():
|
||||
u += 42 # type: ignore[arg-type]
|
||||
|
||||
|
||||
# region AIContent Serialization
|
||||
# region BaseContent Serialization
|
||||
|
||||
|
||||
@mark.parametrize(
|
||||
@@ -396,14 +394,14 @@ def test_usage_details_add_with_none_and_type_errors():
|
||||
(HostedVectorStoreContent, {"vector_store_id": "vs-789"}),
|
||||
],
|
||||
)
|
||||
def test_ai_content_serialization(content_type: type[AIContent], args: dict):
|
||||
def test_ai_content_serialization(content_type: type[BaseContent], args: dict):
|
||||
content = content_type(**args)
|
||||
serialized = content.model_dump()
|
||||
deserialized = content_type.model_validate(serialized)
|
||||
assert deserialized == content
|
||||
|
||||
class TestModel(BaseModel):
|
||||
content: AIContents
|
||||
content: Contents
|
||||
|
||||
test_item = TestModel.model_validate({"content": serialized})
|
||||
|
||||
@@ -419,14 +417,14 @@ def test_chat_message_text():
|
||||
message = ChatMessage(role="user", text="Hello, how are you?")
|
||||
|
||||
# Check the type and content
|
||||
assert message.role == ChatRole.USER
|
||||
assert message.role == Role.USER
|
||||
assert len(message.contents) == 1
|
||||
assert isinstance(message.contents[0], TextContent)
|
||||
assert message.contents[0].text == "Hello, how are you?"
|
||||
assert message.text == "Hello, how are you?"
|
||||
|
||||
# Ensure the instance is of type AIContent
|
||||
assert isinstance(message.contents[0], AIContent)
|
||||
# Ensure the instance is of type BaseContent
|
||||
assert isinstance(message.contents[0], BaseContent)
|
||||
|
||||
|
||||
def test_chat_message_contents():
|
||||
@@ -437,7 +435,7 @@ def test_chat_message_contents():
|
||||
message = ChatMessage(role="user", contents=[content1, content2])
|
||||
|
||||
# Check the type and content
|
||||
assert message.role == ChatRole.USER
|
||||
assert message.role == Role.USER
|
||||
assert len(message.contents) == 2
|
||||
assert isinstance(message.contents[0], TextContent)
|
||||
assert isinstance(message.contents[1], TextContent)
|
||||
@@ -447,8 +445,8 @@ def test_chat_message_contents():
|
||||
|
||||
|
||||
def test_chat_message_with_chatrole_instance():
|
||||
m = ChatMessage(role=ChatRole.USER, text="hi")
|
||||
assert m.role == ChatRole.USER
|
||||
m = ChatMessage(role=Role.USER, text="hi")
|
||||
assert m.role == Role.USER
|
||||
assert m.text == "hi"
|
||||
|
||||
|
||||
@@ -464,7 +462,7 @@ def test_chat_response():
|
||||
response = ChatResponse(messages=message)
|
||||
|
||||
# Check the type and content
|
||||
assert response.messages[0].role == ChatRole.ASSISTANT
|
||||
assert response.messages[0].role == Role.ASSISTANT
|
||||
assert response.messages[0].text == "I'm doing well, thank you!"
|
||||
assert isinstance(response.messages[0], ChatMessage)
|
||||
# __str__ returns text
|
||||
@@ -484,7 +482,7 @@ def test_chat_response_with_format():
|
||||
response = ChatResponse(messages=message)
|
||||
|
||||
# Check the type and content
|
||||
assert response.messages[0].role == ChatRole.ASSISTANT
|
||||
assert response.messages[0].role == Role.ASSISTANT
|
||||
assert response.messages[0].text == '{"response": "Hello"}'
|
||||
assert isinstance(response.messages[0], ChatMessage)
|
||||
assert response.text == '{"response": "Hello"}'
|
||||
@@ -503,7 +501,7 @@ def test_chat_response_with_format_init():
|
||||
response = ChatResponse(messages=message, response_format=OutputModel)
|
||||
|
||||
# Check the type and content
|
||||
assert response.messages[0].role == ChatRole.ASSISTANT
|
||||
assert response.messages[0].role == Role.ASSISTANT
|
||||
assert response.messages[0].text == '{"response": "Hello"}'
|
||||
assert isinstance(response.messages[0], ChatMessage)
|
||||
assert response.text == '{"response": "Hello"}'
|
||||
@@ -767,7 +765,7 @@ def test_chat_options_init_with_args(ai_function_tool, ai_tool) -> None:
|
||||
assert options.frequency_penalty == 0.0
|
||||
assert options.user == "user-123"
|
||||
for tool in options.tools:
|
||||
assert isinstance(tool, AITool)
|
||||
assert isinstance(tool, ToolProtocol)
|
||||
assert tool.name is not None
|
||||
assert tool.description is not None
|
||||
if isinstance(tool, AIFunction):
|
||||
@@ -809,7 +807,7 @@ def test_chat_options_and(ai_function_tool, ai_tool) -> None:
|
||||
|
||||
@fixture
|
||||
def chat_message() -> ChatMessage:
|
||||
return ChatMessage(role=ChatRole.USER, text="Hello")
|
||||
return ChatMessage(role=Role.USER, text="Hello")
|
||||
|
||||
|
||||
@fixture
|
||||
@@ -824,7 +822,7 @@ def agent_run_response(chat_message: ChatMessage) -> AgentRunResponse:
|
||||
|
||||
@fixture
|
||||
def agent_run_response_update(text_content: TextContent) -> AgentRunResponseUpdate:
|
||||
return AgentRunResponseUpdate(role=ChatRole.ASSISTANT, contents=[text_content])
|
||||
return AgentRunResponseUpdate(role=Role.ASSISTANT, contents=[text_content])
|
||||
|
||||
|
||||
# region AgentRunResponse
|
||||
@@ -914,18 +912,16 @@ def test_error_content_str():
|
||||
|
||||
def test_annotations_models_and_roundtrip():
|
||||
span = TextSpanRegion(start_index=0, end_index=5)
|
||||
base_region = AnnotatedRegion()
|
||||
ann: AIAnnotation = AIAnnotation(annotated_regions=[span, base_region])
|
||||
cit = CitationAnnotation(title="Doc", url="http://example.com", snippet="Snippet", annotated_regions=[span])
|
||||
|
||||
# Attach to content
|
||||
content = TextContent(text="hello", additional_properties={"v": 1})
|
||||
content.annotations = [ann, cit]
|
||||
content.annotations = [cit]
|
||||
|
||||
dumped = content.model_dump()
|
||||
loaded = TextContent.model_validate(dumped)
|
||||
assert isinstance(loaded.annotations, list)
|
||||
assert len(loaded.annotations) == 2
|
||||
assert len(loaded.annotations) == 1
|
||||
assert isinstance(loaded.annotations[0], dict) is False # pydantic parsed into models
|
||||
# discriminators preserved
|
||||
assert any(getattr(a, "type", None) == "citation" for a in loaded.annotations)
|
||||
@@ -1032,16 +1028,16 @@ def test_generated_embeddings_operations():
|
||||
assert g.additional_properties == {}
|
||||
|
||||
|
||||
# region ChatRole & ChatFinishReason basics
|
||||
# region Role & FinishReason basics
|
||||
|
||||
|
||||
def test_chat_role_str_and_repr():
|
||||
assert str(ChatRole.USER) == "user"
|
||||
assert "ChatRole(value=" in repr(ChatRole.USER)
|
||||
assert str(Role.USER) == "user"
|
||||
assert "Role(value=" in repr(Role.USER)
|
||||
|
||||
|
||||
def test_chat_finish_reason_constants():
|
||||
assert ChatFinishReason.STOP.value == "stop"
|
||||
assert FinishReason.STOP.value == "stop"
|
||||
|
||||
|
||||
def test_response_update_propagates_fields_and_metadata():
|
||||
@@ -1054,7 +1050,7 @@ def test_response_update_propagates_fields_and_metadata():
|
||||
conversation_id="cid",
|
||||
ai_model_id="model-x",
|
||||
created_at="t0",
|
||||
finish_reason=ChatFinishReason.STOP,
|
||||
finish_reason=FinishReason.STOP,
|
||||
additional_properties={"k": "v"},
|
||||
)
|
||||
resp = ChatResponse.from_chat_response_updates([upd])
|
||||
@@ -1062,9 +1058,9 @@ def test_response_update_propagates_fields_and_metadata():
|
||||
assert resp.created_at == "t0"
|
||||
assert resp.conversation_id == "cid"
|
||||
assert resp.ai_model_id == "model-x"
|
||||
assert resp.finish_reason == ChatFinishReason.STOP
|
||||
assert resp.finish_reason == FinishReason.STOP
|
||||
assert resp.additional_properties and resp.additional_properties["k"] == "v"
|
||||
assert resp.messages[0].role == ChatRole.ASSISTANT
|
||||
assert resp.messages[0].role == Role.ASSISTANT
|
||||
assert resp.messages[0].author_name == "bot"
|
||||
assert resp.messages[0].message_id == "mid"
|
||||
|
||||
|
||||
@@ -14,19 +14,19 @@ from agent_framework import (
|
||||
AgentRunResponse,
|
||||
AgentRunResponseUpdate,
|
||||
AgentThread,
|
||||
ChatClient,
|
||||
ChatClientAgent,
|
||||
ChatAgent,
|
||||
ChatClientProtocol,
|
||||
ChatMessage,
|
||||
ChatOptions,
|
||||
ChatResponse,
|
||||
ChatResponseUpdate,
|
||||
ChatRole,
|
||||
ChatToolMode,
|
||||
FunctionCallContent,
|
||||
FunctionResultContent,
|
||||
HostedCodeInterpreterTool,
|
||||
HostedFileSearchTool,
|
||||
HostedVectorStoreContent,
|
||||
Role,
|
||||
TextContent,
|
||||
UriContent,
|
||||
UsageContent,
|
||||
@@ -125,7 +125,7 @@ def test_openai_assistants_client_init_with_client(mock_async_openai: MagicMock)
|
||||
assert chat_client.assistant_id == "existing-assistant-id"
|
||||
assert chat_client.thread_id == "test-thread-id"
|
||||
assert not chat_client._should_delete_assistant # type: ignore
|
||||
assert isinstance(chat_client, ChatClient)
|
||||
assert isinstance(chat_client, ChatClientProtocol)
|
||||
|
||||
|
||||
def test_openai_assistants_client_init_auto_create_client(
|
||||
@@ -185,7 +185,7 @@ def test_openai_assistants_client_init_with_default_headers(openai_unit_test_env
|
||||
)
|
||||
|
||||
assert chat_client.ai_model_id == "gpt-4"
|
||||
assert isinstance(chat_client, ChatClient)
|
||||
assert isinstance(chat_client, ChatClientProtocol)
|
||||
|
||||
# Assert that the default header we added is present in the client's default headers
|
||||
for key, value in default_headers.items():
|
||||
@@ -412,7 +412,7 @@ async def test_openai_assistants_client_process_stream_events_thread_run_created
|
||||
update = updates[0]
|
||||
assert isinstance(update, ChatResponseUpdate)
|
||||
assert update.conversation_id == thread_id
|
||||
assert update.role == ChatRole.ASSISTANT
|
||||
assert update.role == Role.ASSISTANT
|
||||
assert update.contents == []
|
||||
assert update.raw_representation == mock_response.data
|
||||
|
||||
@@ -457,7 +457,7 @@ async def test_openai_assistants_client_process_stream_events_message_delta_text
|
||||
update = updates[0]
|
||||
assert isinstance(update, ChatResponseUpdate)
|
||||
assert update.conversation_id == thread_id
|
||||
assert update.role == ChatRole.ASSISTANT
|
||||
assert update.role == Role.ASSISTANT
|
||||
assert update.text == "Hello from assistant"
|
||||
assert update.raw_representation == mock_message_delta
|
||||
|
||||
@@ -497,7 +497,7 @@ async def test_openai_assistants_client_process_stream_events_requires_action(mo
|
||||
update = updates[0]
|
||||
assert isinstance(update, ChatResponseUpdate)
|
||||
assert update.conversation_id == thread_id
|
||||
assert update.role == ChatRole.ASSISTANT
|
||||
assert update.role == Role.ASSISTANT
|
||||
assert len(update.contents) == 1
|
||||
assert update.contents[0] == test_function_content
|
||||
assert update.raw_representation == mock_run
|
||||
@@ -579,7 +579,7 @@ async def test_openai_assistants_client_process_stream_events_run_completed_with
|
||||
update = updates[0]
|
||||
assert isinstance(update, ChatResponseUpdate)
|
||||
assert update.conversation_id == thread_id
|
||||
assert update.role == ChatRole.ASSISTANT
|
||||
assert update.role == Role.ASSISTANT
|
||||
assert len(update.contents) == 1
|
||||
|
||||
# Check the usage content
|
||||
@@ -632,7 +632,7 @@ def test_openai_assistants_client_create_run_options_basic(mock_async_openai: Ma
|
||||
top_p=0.9,
|
||||
)
|
||||
|
||||
messages = [ChatMessage(role=ChatRole.USER, text="Hello")]
|
||||
messages = [ChatMessage(role=Role.USER, text="Hello")]
|
||||
|
||||
# Call the method
|
||||
run_options, tool_results = chat_client._create_run_options(messages, chat_options) # type: ignore
|
||||
@@ -661,7 +661,7 @@ def test_openai_assistants_client_create_run_options_with_ai_function_tool(mock_
|
||||
tool_choice="auto",
|
||||
)
|
||||
|
||||
messages = [ChatMessage(role=ChatRole.USER, text="Hello")]
|
||||
messages = [ChatMessage(role=Role.USER, text="Hello")]
|
||||
|
||||
# Call the method
|
||||
run_options, tool_results = chat_client._create_run_options(messages, chat_options) # type: ignore
|
||||
@@ -686,7 +686,7 @@ def test_openai_assistants_client_create_run_options_with_code_interpreter(mock_
|
||||
tool_choice="auto",
|
||||
)
|
||||
|
||||
messages = [ChatMessage(role=ChatRole.USER, text="Calculate something")]
|
||||
messages = [ChatMessage(role=Role.USER, text="Calculate something")]
|
||||
|
||||
# Call the method
|
||||
run_options, tool_results = chat_client._create_run_options(messages, chat_options) # type: ignore
|
||||
@@ -706,7 +706,7 @@ def test_openai_assistants_client_create_run_options_tool_choice_none(mock_async
|
||||
tool_choice="none",
|
||||
)
|
||||
|
||||
messages = [ChatMessage(role=ChatRole.USER, text="Hello")]
|
||||
messages = [ChatMessage(role=Role.USER, text="Hello")]
|
||||
|
||||
# Call the method
|
||||
run_options, tool_results = chat_client._create_run_options(messages, chat_options) # type: ignore
|
||||
@@ -727,7 +727,7 @@ def test_openai_assistants_client_create_run_options_required_function(mock_asyn
|
||||
tool_choice=tool_choice,
|
||||
)
|
||||
|
||||
messages = [ChatMessage(role=ChatRole.USER, text="Hello")]
|
||||
messages = [ChatMessage(role=Role.USER, text="Hello")]
|
||||
|
||||
# Call the method
|
||||
run_options, tool_results = chat_client._create_run_options(messages, chat_options) # type: ignore
|
||||
@@ -753,7 +753,7 @@ def test_openai_assistants_client_create_run_options_with_file_search_tool(mock_
|
||||
tool_choice="auto",
|
||||
)
|
||||
|
||||
messages = [ChatMessage(role=ChatRole.USER, text="Search for information")]
|
||||
messages = [ChatMessage(role=Role.USER, text="Search for information")]
|
||||
|
||||
# Call the method
|
||||
run_options, tool_results = chat_client._create_run_options(messages, chat_options) # type: ignore
|
||||
@@ -778,7 +778,7 @@ def test_openai_assistants_client_create_run_options_with_mapping_tool(mock_asyn
|
||||
tool_choice="auto",
|
||||
)
|
||||
|
||||
messages = [ChatMessage(role=ChatRole.USER, text="Use custom tool")]
|
||||
messages = [ChatMessage(role=Role.USER, text="Use custom tool")]
|
||||
|
||||
# Call the method
|
||||
run_options, tool_results = chat_client._create_run_options(messages, chat_options) # type: ignore
|
||||
@@ -795,8 +795,8 @@ def test_openai_assistants_client_create_run_options_with_system_message(mock_as
|
||||
chat_client = create_test_openai_assistants_client(mock_async_openai)
|
||||
|
||||
messages = [
|
||||
ChatMessage(role=ChatRole.SYSTEM, text="You are a helpful assistant."),
|
||||
ChatMessage(role=ChatRole.USER, text="Hello"),
|
||||
ChatMessage(role=Role.SYSTEM, text="You are a helpful assistant."),
|
||||
ChatMessage(role=Role.USER, text="Hello"),
|
||||
]
|
||||
|
||||
# Call the method
|
||||
@@ -816,7 +816,7 @@ def test_openai_assistants_client_create_run_options_with_image_content(mock_asy
|
||||
|
||||
# Create message with image content
|
||||
image_content = UriContent(uri="https://example.com/image.jpg", media_type="image/jpeg")
|
||||
messages = [ChatMessage(role=ChatRole.USER, contents=[image_content])]
|
||||
messages = [ChatMessage(role=Role.USER, contents=[image_content])]
|
||||
|
||||
# Call the method
|
||||
run_options, tool_results = chat_client._create_run_options(messages, None) # type: ignore
|
||||
@@ -924,7 +924,7 @@ def get_weather(
|
||||
async def test_openai_assistants_client_get_response() -> None:
|
||||
"""Test OpenAI Assistants Client response."""
|
||||
async with OpenAIAssistantsClient() as openai_assistants_client:
|
||||
assert isinstance(openai_assistants_client, ChatClient)
|
||||
assert isinstance(openai_assistants_client, ChatClientProtocol)
|
||||
|
||||
messages: list[ChatMessage] = []
|
||||
messages.append(
|
||||
@@ -948,7 +948,7 @@ async def test_openai_assistants_client_get_response() -> None:
|
||||
async def test_openai_assistants_client_get_response_tools() -> None:
|
||||
"""Test OpenAI Assistants Client response with tools."""
|
||||
async with OpenAIAssistantsClient() as openai_assistants_client:
|
||||
assert isinstance(openai_assistants_client, ChatClient)
|
||||
assert isinstance(openai_assistants_client, ChatClientProtocol)
|
||||
|
||||
messages: list[ChatMessage] = []
|
||||
messages.append(ChatMessage(role="user", text="What's the weather like in Seattle?"))
|
||||
@@ -969,7 +969,7 @@ async def test_openai_assistants_client_get_response_tools() -> None:
|
||||
async def test_openai_assistants_client_streaming() -> None:
|
||||
"""Test OpenAI Assistants Client streaming response."""
|
||||
async with OpenAIAssistantsClient() as openai_assistants_client:
|
||||
assert isinstance(openai_assistants_client, ChatClient)
|
||||
assert isinstance(openai_assistants_client, ChatClientProtocol)
|
||||
|
||||
messages: list[ChatMessage] = []
|
||||
messages.append(
|
||||
@@ -999,7 +999,7 @@ async def test_openai_assistants_client_streaming() -> None:
|
||||
async def test_openai_assistants_client_streaming_tools() -> None:
|
||||
"""Test OpenAI Assistants Client streaming response with tools."""
|
||||
async with OpenAIAssistantsClient() as openai_assistants_client:
|
||||
assert isinstance(openai_assistants_client, ChatClient)
|
||||
assert isinstance(openai_assistants_client, ChatClientProtocol)
|
||||
|
||||
messages: list[ChatMessage] = []
|
||||
messages.append(ChatMessage(role="user", text="What's the weather like in Seattle?"))
|
||||
@@ -1035,7 +1035,7 @@ async def test_openai_assistants_client_with_existing_assistant() -> None:
|
||||
async with OpenAIAssistantsClient(
|
||||
ai_model_id="gpt-4o-mini", assistant_id=assistant_id
|
||||
) as openai_assistants_client:
|
||||
assert isinstance(openai_assistants_client, ChatClient)
|
||||
assert isinstance(openai_assistants_client, ChatClientProtocol)
|
||||
assert openai_assistants_client.assistant_id == assistant_id
|
||||
|
||||
messages = [ChatMessage(role="user", text="What can you do?")]
|
||||
@@ -1052,7 +1052,7 @@ async def test_openai_assistants_client_with_existing_assistant() -> None:
|
||||
async def test_openai_assistants_client_file_search() -> None:
|
||||
"""Test OpenAI Assistants Client response."""
|
||||
async with OpenAIAssistantsClient() as openai_assistants_client:
|
||||
assert isinstance(openai_assistants_client, ChatClient)
|
||||
assert isinstance(openai_assistants_client, ChatClientProtocol)
|
||||
|
||||
messages: list[ChatMessage] = []
|
||||
messages.append(ChatMessage(role="user", text="What's the weather like today?"))
|
||||
@@ -1074,7 +1074,7 @@ async def test_openai_assistants_client_file_search() -> None:
|
||||
async def test_openai_assistants_client_file_search_streaming() -> None:
|
||||
"""Test OpenAI Assistants Client response."""
|
||||
async with OpenAIAssistantsClient() as openai_assistants_client:
|
||||
assert isinstance(openai_assistants_client, ChatClient)
|
||||
assert isinstance(openai_assistants_client, ChatClientProtocol)
|
||||
|
||||
messages: list[ChatMessage] = []
|
||||
messages.append(ChatMessage(role="user", text="What's the weather like today?"))
|
||||
@@ -1101,8 +1101,8 @@ async def test_openai_assistants_client_file_search_streaming() -> None:
|
||||
|
||||
@skip_if_openai_integration_tests_disabled
|
||||
async def test_openai_assistants_agent_basic_run():
|
||||
"""Test ChatClientAgent basic run functionality with OpenAIAssistantsClient."""
|
||||
async with ChatClientAgent(
|
||||
"""Test ChatAgent basic run functionality with OpenAIAssistantsClient."""
|
||||
async with ChatAgent(
|
||||
chat_client=OpenAIAssistantsClient(),
|
||||
) as agent:
|
||||
# Run a simple query
|
||||
@@ -1117,13 +1117,13 @@ async def test_openai_assistants_agent_basic_run():
|
||||
|
||||
@skip_if_openai_integration_tests_disabled
|
||||
async def test_openai_assistants_agent_basic_run_streaming():
|
||||
"""Test ChatClientAgent basic streaming functionality with OpenAIAssistantsClient."""
|
||||
async with ChatClientAgent(
|
||||
"""Test ChatAgent basic streaming functionality with OpenAIAssistantsClient."""
|
||||
async with ChatAgent(
|
||||
chat_client=OpenAIAssistantsClient(),
|
||||
) as agent:
|
||||
# Run streaming query
|
||||
full_message: str = ""
|
||||
async for chunk in agent.run_streaming("Please respond with exactly: 'This is a streaming response test.'"):
|
||||
async for chunk in agent.run_stream("Please respond with exactly: 'This is a streaming response test.'"):
|
||||
assert chunk is not None
|
||||
assert isinstance(chunk, AgentRunResponseUpdate)
|
||||
if chunk.text:
|
||||
@@ -1136,8 +1136,8 @@ async def test_openai_assistants_agent_basic_run_streaming():
|
||||
|
||||
@skip_if_openai_integration_tests_disabled
|
||||
async def test_openai_assistants_agent_thread_persistence():
|
||||
"""Test ChatClientAgent thread persistence across runs with OpenAIAssistantsClient."""
|
||||
async with ChatClientAgent(
|
||||
"""Test ChatAgent thread persistence across runs with OpenAIAssistantsClient."""
|
||||
async with ChatAgent(
|
||||
chat_client=OpenAIAssistantsClient(),
|
||||
instructions="You are a helpful assistant with good memory.",
|
||||
) as agent:
|
||||
@@ -1164,11 +1164,11 @@ async def test_openai_assistants_agent_thread_persistence():
|
||||
|
||||
@skip_if_openai_integration_tests_disabled
|
||||
async def test_openai_assistants_agent_existing_thread_id():
|
||||
"""Test ChatClientAgent with existing thread ID to continue conversations across agent instances."""
|
||||
"""Test ChatAgent with existing thread ID to continue conversations across agent instances."""
|
||||
# First, create a conversation and capture the thread ID
|
||||
existing_thread_id = None
|
||||
|
||||
async with ChatClientAgent(
|
||||
async with ChatAgent(
|
||||
chat_client=OpenAIAssistantsClient(),
|
||||
instructions="You are a helpful weather agent.",
|
||||
tools=[get_weather],
|
||||
@@ -1188,7 +1188,7 @@ async def test_openai_assistants_agent_existing_thread_id():
|
||||
|
||||
# Now continue with the same thread ID in a new agent instance
|
||||
|
||||
async with ChatClientAgent(
|
||||
async with ChatAgent(
|
||||
chat_client=OpenAIAssistantsClient(thread_id=existing_thread_id),
|
||||
instructions="You are a helpful weather agent.",
|
||||
tools=[get_weather],
|
||||
@@ -1208,9 +1208,9 @@ async def test_openai_assistants_agent_existing_thread_id():
|
||||
|
||||
@skip_if_openai_integration_tests_disabled
|
||||
async def test_openai_assistants_agent_code_interpreter():
|
||||
"""Test ChatClientAgent with code interpreter through OpenAIAssistantsClient."""
|
||||
"""Test ChatAgent with code interpreter through OpenAIAssistantsClient."""
|
||||
|
||||
async with ChatClientAgent(
|
||||
async with ChatAgent(
|
||||
chat_client=OpenAIAssistantsClient(),
|
||||
instructions="You are a helpful assistant that can write and execute Python code.",
|
||||
tools=[HostedCodeInterpreterTool()],
|
||||
@@ -1229,7 +1229,7 @@ async def test_openai_assistants_agent_code_interpreter():
|
||||
async def test_openai_assistants_client_agent_level_tool_persistence():
|
||||
"""Test that agent-level tools persist across multiple runs with OpenAI Assistants Client."""
|
||||
|
||||
async with ChatClientAgent(
|
||||
async with ChatAgent(
|
||||
chat_client=OpenAIAssistantsClient(),
|
||||
instructions="You are a helpful assistant that uses available tools.",
|
||||
tools=[get_weather], # Agent-level tool
|
||||
@@ -1264,7 +1264,7 @@ async def test_openai_assistants_client_run_level_tool_isolation():
|
||||
call_count += 1
|
||||
return f"The weather in {location} is sunny and 72°F."
|
||||
|
||||
async with ChatClientAgent(
|
||||
async with ChatAgent(
|
||||
chat_client=OpenAIAssistantsClient(),
|
||||
instructions="You are a helpful assistant.",
|
||||
) as agent:
|
||||
|
||||
@@ -10,15 +10,15 @@ from openai import BadRequestError
|
||||
from agent_framework import (
|
||||
AgentRunResponse,
|
||||
AgentRunResponseUpdate,
|
||||
AITool,
|
||||
ChatClient,
|
||||
ChatClientAgent,
|
||||
ChatAgent,
|
||||
ChatClientProtocol,
|
||||
ChatMessage,
|
||||
ChatOptions,
|
||||
ChatResponse,
|
||||
ChatResponseUpdate,
|
||||
HostedWebSearchTool,
|
||||
TextContent,
|
||||
ToolProtocol,
|
||||
ai_function,
|
||||
)
|
||||
from agent_framework.exceptions import ServiceInitializationError, ServiceResponseException
|
||||
@@ -39,7 +39,7 @@ def test_init(openai_unit_test_env: dict[str, str]) -> None:
|
||||
open_ai_chat_completion = OpenAIChatClient()
|
||||
|
||||
assert open_ai_chat_completion.ai_model_id == openai_unit_test_env["OPENAI_CHAT_MODEL_ID"]
|
||||
assert isinstance(open_ai_chat_completion, ChatClient)
|
||||
assert isinstance(open_ai_chat_completion, ChatClientProtocol)
|
||||
|
||||
|
||||
def test_init_validation_fail() -> None:
|
||||
@@ -54,7 +54,7 @@ def test_init_ai_model_id_constructor(openai_unit_test_env: dict[str, str]) -> N
|
||||
open_ai_chat_completion = OpenAIChatClient(ai_model_id=ai_model_id)
|
||||
|
||||
assert open_ai_chat_completion.ai_model_id == ai_model_id
|
||||
assert isinstance(open_ai_chat_completion, ChatClient)
|
||||
assert isinstance(open_ai_chat_completion, ChatClientProtocol)
|
||||
|
||||
|
||||
def test_init_with_default_header(openai_unit_test_env: dict[str, str]) -> None:
|
||||
@@ -66,7 +66,7 @@ def test_init_with_default_header(openai_unit_test_env: dict[str, str]) -> None:
|
||||
)
|
||||
|
||||
assert open_ai_chat_completion.ai_model_id == openai_unit_test_env["OPENAI_CHAT_MODEL_ID"]
|
||||
assert isinstance(open_ai_chat_completion, ChatClient)
|
||||
assert isinstance(open_ai_chat_completion, ChatClientProtocol)
|
||||
|
||||
# Assert that the default header we added is present in the client's default headers
|
||||
for key, value in default_headers.items():
|
||||
@@ -154,15 +154,15 @@ def test_unsupported_tool_handling(openai_unit_test_env: dict[str, str]) -> None
|
||||
"""Test that unsupported tool types are handled correctly."""
|
||||
client = OpenAIChatClient()
|
||||
|
||||
# Create a mock AITool that's not an AIFunction
|
||||
unsupported_tool = MagicMock(spec=AITool)
|
||||
# Create a mock ToolProtocol that's not an AIFunction
|
||||
unsupported_tool = MagicMock(spec=ToolProtocol)
|
||||
unsupported_tool.__class__.__name__ = "UnsupportedAITool"
|
||||
|
||||
# This should ignore the unsupported AITool and return empty list
|
||||
# This should ignore the unsupported ToolProtocol and return empty list
|
||||
result = client._chat_to_tool_spec([unsupported_tool]) # type: ignore
|
||||
assert result == []
|
||||
|
||||
# Also test with a non-AITool that should be converted to dict
|
||||
# Also test with a non-ToolProtocol that should be converted to dict
|
||||
dict_tool = {"type": "function", "name": "test"}
|
||||
result = client._chat_to_tool_spec([dict_tool]) # type: ignore
|
||||
assert result == [dict_tool]
|
||||
@@ -190,7 +190,7 @@ async def test_openai_chat_completion_response() -> None:
|
||||
"""Test OpenAI chat completion responses."""
|
||||
openai_chat_client = OpenAIChatClient()
|
||||
|
||||
assert isinstance(openai_chat_client, ChatClient)
|
||||
assert isinstance(openai_chat_client, ChatClientProtocol)
|
||||
|
||||
messages: list[ChatMessage] = []
|
||||
messages.append(
|
||||
@@ -217,7 +217,7 @@ async def test_openai_chat_completion_response_tools() -> None:
|
||||
"""Test OpenAI chat completion responses."""
|
||||
openai_chat_client = OpenAIChatClient()
|
||||
|
||||
assert isinstance(openai_chat_client, ChatClient)
|
||||
assert isinstance(openai_chat_client, ChatClientProtocol)
|
||||
|
||||
messages: list[ChatMessage] = []
|
||||
messages.append(ChatMessage(role="user", text="who are Emily and David?"))
|
||||
@@ -239,7 +239,7 @@ async def test_openai_chat_client_streaming() -> None:
|
||||
"""Test Azure OpenAI chat completion responses."""
|
||||
openai_chat_client = OpenAIChatClient()
|
||||
|
||||
assert isinstance(openai_chat_client, ChatClient)
|
||||
assert isinstance(openai_chat_client, ChatClientProtocol)
|
||||
|
||||
messages: list[ChatMessage] = []
|
||||
messages.append(
|
||||
@@ -274,7 +274,7 @@ async def test_openai_chat_client_streaming_tools() -> None:
|
||||
"""Test AzureOpenAI chat completion responses."""
|
||||
openai_chat_client = OpenAIChatClient()
|
||||
|
||||
assert isinstance(openai_chat_client, ChatClient)
|
||||
assert isinstance(openai_chat_client, ChatClientProtocol)
|
||||
|
||||
messages: list[ChatMessage] = []
|
||||
messages.append(ChatMessage(role="user", text="who are Emily and David?"))
|
||||
@@ -301,7 +301,7 @@ async def test_openai_chat_client_web_search() -> None:
|
||||
# Currently only a select few models support web search tool calls
|
||||
openai_chat_client = OpenAIChatClient(ai_model_id="gpt-4o-search-preview")
|
||||
|
||||
assert isinstance(openai_chat_client, ChatClient)
|
||||
assert isinstance(openai_chat_client, ChatClientProtocol)
|
||||
|
||||
# Test that the client will use the web search tool
|
||||
response = await openai_chat_client.get_response(
|
||||
@@ -340,7 +340,7 @@ async def test_openai_chat_client_web_search() -> None:
|
||||
async def test_openai_chat_client_web_search_streaming() -> None:
|
||||
openai_chat_client = OpenAIChatClient(ai_model_id="gpt-4o-search-preview")
|
||||
|
||||
assert isinstance(openai_chat_client, ChatClient)
|
||||
assert isinstance(openai_chat_client, ChatClientProtocol)
|
||||
|
||||
# Test that the client will use the web search tool
|
||||
response = openai_chat_client.get_streaming_response(
|
||||
@@ -392,7 +392,7 @@ async def test_openai_chat_client_web_search_streaming() -> None:
|
||||
@skip_if_openai_integration_tests_disabled
|
||||
async def test_openai_chat_client_agent_basic_run():
|
||||
"""Test OpenAI chat client agent basic run functionality with OpenAIChatClient."""
|
||||
async with ChatClientAgent(
|
||||
async with ChatAgent(
|
||||
chat_client=OpenAIChatClient(ai_model_id="gpt-4o-search-preview"),
|
||||
) as agent:
|
||||
# Test basic run
|
||||
@@ -407,12 +407,12 @@ async def test_openai_chat_client_agent_basic_run():
|
||||
@skip_if_openai_integration_tests_disabled
|
||||
async def test_openai_chat_client_agent_basic_run_streaming():
|
||||
"""Test OpenAI chat client agent basic streaming functionality with OpenAIChatClient."""
|
||||
async with ChatClientAgent(
|
||||
async with ChatAgent(
|
||||
chat_client=OpenAIChatClient(ai_model_id="gpt-4o-search-preview"),
|
||||
) as agent:
|
||||
# Test streaming run
|
||||
full_text = ""
|
||||
async for chunk in agent.run_streaming("Please respond with exactly: 'This is a streaming response test.'"):
|
||||
async for chunk in agent.run_stream("Please respond with exactly: 'This is a streaming response test.'"):
|
||||
assert isinstance(chunk, AgentRunResponseUpdate)
|
||||
if chunk.text:
|
||||
full_text += chunk.text
|
||||
@@ -424,7 +424,7 @@ async def test_openai_chat_client_agent_basic_run_streaming():
|
||||
@skip_if_openai_integration_tests_disabled
|
||||
async def test_openai_chat_client_agent_thread_persistence():
|
||||
"""Test OpenAI chat client agent thread persistence across runs with OpenAIChatClient."""
|
||||
async with ChatClientAgent(
|
||||
async with ChatAgent(
|
||||
chat_client=OpenAIChatClient(ai_model_id="gpt-4o-search-preview"),
|
||||
instructions="You are a helpful assistant with good memory.",
|
||||
) as agent:
|
||||
@@ -451,7 +451,7 @@ async def test_openai_chat_client_agent_existing_thread():
|
||||
# First conversation - capture the thread
|
||||
preserved_thread = None
|
||||
|
||||
async with ChatClientAgent(
|
||||
async with ChatAgent(
|
||||
chat_client=OpenAIChatClient(ai_model_id="gpt-4o-search-preview"),
|
||||
instructions="You are a helpful assistant with good memory.",
|
||||
) as first_agent:
|
||||
@@ -467,7 +467,7 @@ async def test_openai_chat_client_agent_existing_thread():
|
||||
|
||||
# Second conversation - reuse the thread in a new agent instance
|
||||
if preserved_thread:
|
||||
async with ChatClientAgent(
|
||||
async with ChatAgent(
|
||||
chat_client=OpenAIChatClient(ai_model_id="gpt-4o-search-preview"),
|
||||
instructions="You are a helpful assistant with good memory.",
|
||||
) as second_agent:
|
||||
@@ -483,7 +483,7 @@ async def test_openai_chat_client_agent_existing_thread():
|
||||
async def test_openai_chat_client_agent_level_tool_persistence():
|
||||
"""Test that agent-level tools persist across multiple runs with OpenAI Chat Client."""
|
||||
|
||||
async with ChatClientAgent(
|
||||
async with ChatAgent(
|
||||
chat_client=OpenAIChatClient(ai_model_id="gpt-4.1"),
|
||||
instructions="You are a helpful assistant that uses available tools.",
|
||||
tools=[get_weather], # Agent-level tool
|
||||
@@ -518,7 +518,7 @@ async def test_openai_chat_client_run_level_tool_isolation():
|
||||
call_count += 1
|
||||
return f"The weather in {location} is sunny and 72°F."
|
||||
|
||||
async with ChatClientAgent(
|
||||
async with ChatAgent(
|
||||
chat_client=OpenAIChatClient(ai_model_id="gpt-4.1"),
|
||||
instructions="You are a helpful assistant.",
|
||||
) as agent:
|
||||
|
||||
@@ -13,12 +13,11 @@ from agent_framework import (
|
||||
AgentRunResponse,
|
||||
AgentRunResponseUpdate,
|
||||
AgentThread,
|
||||
ChatClient,
|
||||
ChatClientAgent,
|
||||
ChatAgent,
|
||||
ChatClientProtocol,
|
||||
ChatMessage,
|
||||
ChatResponse,
|
||||
ChatResponseUpdate,
|
||||
ChatRole,
|
||||
FunctionCallContent,
|
||||
FunctionResultContent,
|
||||
HostedCodeInterpreterTool,
|
||||
@@ -26,6 +25,7 @@ from agent_framework import (
|
||||
HostedFileSearchTool,
|
||||
HostedVectorStoreContent,
|
||||
HostedWebSearchTool,
|
||||
Role,
|
||||
TextContent,
|
||||
TextReasoningContent,
|
||||
UriContent,
|
||||
@@ -87,7 +87,7 @@ def test_init(openai_unit_test_env: dict[str, str]) -> None:
|
||||
openai_responses_client = OpenAIResponsesClient()
|
||||
|
||||
assert openai_responses_client.ai_model_id == openai_unit_test_env["OPENAI_RESPONSES_MODEL_ID"]
|
||||
assert isinstance(openai_responses_client, ChatClient)
|
||||
assert isinstance(openai_responses_client, ChatClientProtocol)
|
||||
|
||||
|
||||
def test_init_validation_fail() -> None:
|
||||
@@ -102,7 +102,7 @@ def test_init_ai_model_id_constructor(openai_unit_test_env: dict[str, str]) -> N
|
||||
openai_responses_client = OpenAIResponsesClient(ai_model_id=ai_model_id)
|
||||
|
||||
assert openai_responses_client.ai_model_id == ai_model_id
|
||||
assert isinstance(openai_responses_client, ChatClient)
|
||||
assert isinstance(openai_responses_client, ChatClientProtocol)
|
||||
|
||||
|
||||
def test_init_with_default_header(openai_unit_test_env: dict[str, str]) -> None:
|
||||
@@ -114,7 +114,7 @@ def test_init_with_default_header(openai_unit_test_env: dict[str, str]) -> None:
|
||||
)
|
||||
|
||||
assert openai_responses_client.ai_model_id == openai_unit_test_env["OPENAI_RESPONSES_MODEL_ID"]
|
||||
assert isinstance(openai_responses_client, ChatClient)
|
||||
assert isinstance(openai_responses_client, ChatClientProtocol)
|
||||
|
||||
# Assert that the default header we added is present in the client's default headers
|
||||
for key, value in default_headers.items():
|
||||
@@ -731,7 +731,7 @@ def test_streaming_response_basic_structure() -> None:
|
||||
|
||||
# Should get a valid ChatResponseUpdate structure
|
||||
assert isinstance(response, ChatResponseUpdate)
|
||||
assert response.role == ChatRole.ASSISTANT
|
||||
assert response.role == Role.ASSISTANT
|
||||
assert response.ai_model_id == "test-model"
|
||||
assert isinstance(response.contents, list)
|
||||
assert response.raw_representation is mock_event
|
||||
@@ -742,7 +742,7 @@ async def test_openai_responses_client_response() -> None:
|
||||
"""Test OpenAI chat completion responses."""
|
||||
openai_responses_client = OpenAIResponsesClient()
|
||||
|
||||
assert isinstance(openai_responses_client, ChatClient)
|
||||
assert isinstance(openai_responses_client, ChatClientProtocol)
|
||||
|
||||
messages: list[ChatMessage] = []
|
||||
messages.append(
|
||||
@@ -785,7 +785,7 @@ async def test_openai_responses_client_response_tools() -> None:
|
||||
"""Test OpenAI chat completion responses."""
|
||||
openai_responses_client = OpenAIResponsesClient()
|
||||
|
||||
assert isinstance(openai_responses_client, ChatClient)
|
||||
assert isinstance(openai_responses_client, ChatClientProtocol)
|
||||
|
||||
messages: list[ChatMessage] = []
|
||||
messages.append(ChatMessage(role="user", text="What is the weather in New York?"))
|
||||
@@ -824,7 +824,7 @@ async def test_openai_responses_client_streaming() -> None:
|
||||
"""Test Azure OpenAI chat completion responses."""
|
||||
openai_responses_client = OpenAIResponsesClient()
|
||||
|
||||
assert isinstance(openai_responses_client, ChatClient)
|
||||
assert isinstance(openai_responses_client, ChatClientProtocol)
|
||||
|
||||
messages: list[ChatMessage] = []
|
||||
messages.append(
|
||||
@@ -877,7 +877,7 @@ async def test_openai_responses_client_streaming_tools() -> None:
|
||||
"""Test OpenAI chat completion responses."""
|
||||
openai_responses_client = OpenAIResponsesClient()
|
||||
|
||||
assert isinstance(openai_responses_client, ChatClient)
|
||||
assert isinstance(openai_responses_client, ChatClientProtocol)
|
||||
|
||||
messages: list[ChatMessage] = [ChatMessage(role="user", text="What is the weather in Seattle?")]
|
||||
|
||||
@@ -923,7 +923,7 @@ async def test_openai_responses_client_streaming_tools() -> None:
|
||||
async def test_openai_responses_client_web_search() -> None:
|
||||
openai_responses_client = OpenAIResponsesClient()
|
||||
|
||||
assert isinstance(openai_responses_client, ChatClient)
|
||||
assert isinstance(openai_responses_client, ChatClientProtocol)
|
||||
|
||||
# Test that the client will use the web search tool
|
||||
response = await openai_responses_client.get_response(
|
||||
@@ -962,7 +962,7 @@ async def test_openai_responses_client_web_search() -> None:
|
||||
async def test_openai_responses_client_web_search_streaming() -> None:
|
||||
openai_responses_client = OpenAIResponsesClient()
|
||||
|
||||
assert isinstance(openai_responses_client, ChatClient)
|
||||
assert isinstance(openai_responses_client, ChatClientProtocol)
|
||||
|
||||
# Test that the client will use the web search tool
|
||||
response = openai_responses_client.get_streaming_response(
|
||||
@@ -1015,7 +1015,7 @@ async def test_openai_responses_client_web_search_streaming() -> None:
|
||||
async def test_openai_responses_client_file_search() -> None:
|
||||
openai_responses_client = OpenAIResponsesClient()
|
||||
|
||||
assert isinstance(openai_responses_client, ChatClient)
|
||||
assert isinstance(openai_responses_client, ChatClientProtocol)
|
||||
|
||||
file_id, vector_store = await create_vector_store(openai_responses_client)
|
||||
# Test that the client will use the web search tool
|
||||
@@ -1039,7 +1039,7 @@ async def test_openai_responses_client_file_search() -> None:
|
||||
async def test_openai_responses_client_streaming_file_search() -> None:
|
||||
openai_responses_client = OpenAIResponsesClient()
|
||||
|
||||
assert isinstance(openai_responses_client, ChatClient)
|
||||
assert isinstance(openai_responses_client, ChatClientProtocol)
|
||||
|
||||
file_id, vector_store = await create_vector_store(openai_responses_client)
|
||||
# Test that the client will use the web search tool
|
||||
@@ -1088,12 +1088,12 @@ async def test_openai_responses_client_agent_basic_run():
|
||||
@skip_if_openai_integration_tests_disabled
|
||||
async def test_openai_responses_client_agent_basic_run_streaming():
|
||||
"""Test OpenAI Responses Client agent basic streaming functionality with OpenAIResponsesClient."""
|
||||
async with ChatClientAgent(
|
||||
async with ChatAgent(
|
||||
chat_client=OpenAIResponsesClient(),
|
||||
) as agent:
|
||||
# Test streaming run
|
||||
full_text = ""
|
||||
async for chunk in agent.run_streaming("Please respond with exactly: 'This is a streaming response test.'"):
|
||||
async for chunk in agent.run_stream("Please respond with exactly: 'This is a streaming response test.'"):
|
||||
assert isinstance(chunk, AgentRunResponseUpdate)
|
||||
if chunk.text:
|
||||
full_text += chunk.text
|
||||
@@ -1105,7 +1105,7 @@ async def test_openai_responses_client_agent_basic_run_streaming():
|
||||
@skip_if_openai_integration_tests_disabled
|
||||
async def test_openai_responses_client_agent_thread_persistence():
|
||||
"""Test OpenAI Responses Client agent thread persistence across runs with OpenAIResponsesClient."""
|
||||
async with ChatClientAgent(
|
||||
async with ChatAgent(
|
||||
chat_client=OpenAIResponsesClient(),
|
||||
instructions="You are a helpful assistant with good memory.",
|
||||
) as agent:
|
||||
@@ -1128,7 +1128,7 @@ async def test_openai_responses_client_agent_thread_persistence():
|
||||
@skip_if_openai_integration_tests_disabled
|
||||
async def test_openai_responses_client_agent_thread_storage_with_store_true():
|
||||
"""Test OpenAI Responses Client agent with store=True to verify service_thread_id is returned."""
|
||||
async with ChatClientAgent(
|
||||
async with ChatAgent(
|
||||
chat_client=OpenAIResponsesClient(),
|
||||
instructions="You are a helpful assistant.",
|
||||
) as agent:
|
||||
@@ -1162,7 +1162,7 @@ async def test_openai_responses_client_agent_existing_thread():
|
||||
# First conversation - capture the thread
|
||||
preserved_thread = None
|
||||
|
||||
async with ChatClientAgent(
|
||||
async with ChatAgent(
|
||||
chat_client=OpenAIResponsesClient(),
|
||||
instructions="You are a helpful assistant with good memory.",
|
||||
) as first_agent:
|
||||
@@ -1178,7 +1178,7 @@ async def test_openai_responses_client_agent_existing_thread():
|
||||
|
||||
# Second conversation - reuse the thread in a new agent instance
|
||||
if preserved_thread:
|
||||
async with ChatClientAgent(
|
||||
async with ChatAgent(
|
||||
chat_client=OpenAIResponsesClient(),
|
||||
instructions="You are a helpful assistant with good memory.",
|
||||
) as second_agent:
|
||||
@@ -1193,7 +1193,7 @@ async def test_openai_responses_client_agent_existing_thread():
|
||||
@skip_if_openai_integration_tests_disabled
|
||||
async def test_openai_responses_client_agent_hosted_code_interpreter_tool():
|
||||
"""Test OpenAI Responses Client agent with HostedCodeInterpreterTool through OpenAIResponsesClient."""
|
||||
async with ChatClientAgent(
|
||||
async with ChatAgent(
|
||||
chat_client=OpenAIResponsesClient(),
|
||||
instructions="You are a helpful assistant that can execute Python code.",
|
||||
tools=[HostedCodeInterpreterTool()],
|
||||
@@ -1215,7 +1215,7 @@ async def test_openai_responses_client_agent_hosted_code_interpreter_tool():
|
||||
async def test_openai_responses_client_agent_level_tool_persistence():
|
||||
"""Test that agent-level tools persist across multiple runs with OpenAI Responses Client."""
|
||||
|
||||
async with ChatClientAgent(
|
||||
async with ChatAgent(
|
||||
chat_client=OpenAIResponsesClient(),
|
||||
instructions="You are a helpful assistant that uses available tools.",
|
||||
tools=[get_weather], # Agent-level tool
|
||||
@@ -1250,7 +1250,7 @@ async def test_openai_responses_client_run_level_tool_isolation():
|
||||
call_count += 1
|
||||
return f"The weather in {location} is sunny and 72°F."
|
||||
|
||||
async with ChatClientAgent(
|
||||
async with ChatAgent(
|
||||
chat_client=OpenAIResponsesClient(),
|
||||
instructions="You are a helpful assistant.",
|
||||
) as agent:
|
||||
|
||||
@@ -7,14 +7,14 @@ from datetime import datetime
|
||||
from typing import TYPE_CHECKING, Any, ClassVar, TypedDict, cast
|
||||
|
||||
from agent_framework import (
|
||||
AgentBase,
|
||||
AgentRunResponse,
|
||||
AgentRunResponseUpdate,
|
||||
AgentThread,
|
||||
BaseAgent,
|
||||
ChatMessage,
|
||||
ChatRole,
|
||||
FunctionCallContent,
|
||||
FunctionResultContent,
|
||||
Role,
|
||||
TextContent,
|
||||
UsageDetails,
|
||||
)
|
||||
@@ -34,8 +34,8 @@ if TYPE_CHECKING:
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class WorkflowAgent(AgentBase):
|
||||
"""An `AIAgent` subclass that wraps a workflow and exposes it as an agent."""
|
||||
class WorkflowAgent(BaseAgent):
|
||||
"""An `Agent` subclass that wraps a workflow and exposes it as an agent."""
|
||||
|
||||
# Class variable for the request info function name
|
||||
REQUEST_INFO_FUNCTION_NAME: ClassVar[str] = "request_info"
|
||||
@@ -65,11 +65,11 @@ class WorkflowAgent(AgentBase):
|
||||
id: Unique identifier for the agent. If None, will be generated.
|
||||
name: Optional name for the agent.
|
||||
description: Optional description of the agent.
|
||||
**kwargs: Additional keyword arguments passed to AgentBase.
|
||||
**kwargs: Additional keyword arguments passed to BaseAgent.
|
||||
"""
|
||||
if id is None:
|
||||
id = f"WorkflowAgent_{uuid.uuid4().hex[:8]}"
|
||||
# Initialize with standard AgentBase parameters first
|
||||
# Initialize with standard BaseAgent parameters first
|
||||
kwargs["workflow"] = workflow
|
||||
|
||||
# Validate the workflow's start executor can handle agent-facing message inputs
|
||||
@@ -107,7 +107,7 @@ class WorkflowAgent(AgentBase):
|
||||
thread = thread or self.get_new_thread()
|
||||
response_id = str(uuid.uuid4())
|
||||
|
||||
async for update in self._run_streaming_impl(input_messages, response_id):
|
||||
async for update in self._run_stream_impl(input_messages, response_id):
|
||||
response_updates.append(update)
|
||||
|
||||
# Convert updates to final response.
|
||||
@@ -119,7 +119,7 @@ class WorkflowAgent(AgentBase):
|
||||
|
||||
return response
|
||||
|
||||
async def run_streaming(
|
||||
async def run_stream(
|
||||
self,
|
||||
messages: str | ChatMessage | list[str] | list[ChatMessage] | None = None,
|
||||
*,
|
||||
@@ -141,7 +141,7 @@ class WorkflowAgent(AgentBase):
|
||||
response_updates: list[AgentRunResponseUpdate] = []
|
||||
response_id = str(uuid.uuid4())
|
||||
|
||||
async for update in self._run_streaming_impl(input_messages, response_id):
|
||||
async for update in self._run_stream_impl(input_messages, response_id):
|
||||
response_updates.append(update)
|
||||
yield update
|
||||
|
||||
@@ -152,7 +152,7 @@ class WorkflowAgent(AgentBase):
|
||||
await self._notify_thread_of_new_messages(thread, input_messages)
|
||||
await self._notify_thread_of_new_messages(thread, response.messages)
|
||||
|
||||
async def _run_streaming_impl(
|
||||
async def _run_stream_impl(
|
||||
self,
|
||||
input_messages: list[ChatMessage],
|
||||
response_id: str,
|
||||
@@ -188,7 +188,7 @@ class WorkflowAgent(AgentBase):
|
||||
else:
|
||||
# Execute workflow with streaming (initial run or no function responses)
|
||||
# Pass the new input messages directly to the workflow
|
||||
event_stream = self.workflow.run_streaming(input_messages)
|
||||
event_stream = self.workflow.run_stream(input_messages)
|
||||
|
||||
# Process events from the stream
|
||||
async for event in event_stream:
|
||||
@@ -206,7 +206,7 @@ class WorkflowAgent(AgentBase):
|
||||
return []
|
||||
|
||||
if isinstance(messages, str):
|
||||
return [ChatMessage(role=ChatRole.USER, contents=[TextContent(text=messages)])]
|
||||
return [ChatMessage(role=Role.USER, contents=[TextContent(text=messages)])]
|
||||
|
||||
if isinstance(messages, ChatMessage):
|
||||
return [messages]
|
||||
@@ -214,7 +214,7 @@ class WorkflowAgent(AgentBase):
|
||||
normalized = []
|
||||
for msg in messages:
|
||||
if isinstance(msg, str):
|
||||
normalized.append(ChatMessage(role=ChatRole.USER, contents=[TextContent(text=msg)]))
|
||||
normalized.append(ChatMessage(role=Role.USER, contents=[TextContent(text=msg)]))
|
||||
elif isinstance(msg, ChatMessage):
|
||||
normalized.append(msg)
|
||||
return normalized
|
||||
@@ -250,7 +250,7 @@ class WorkflowAgent(AgentBase):
|
||||
)
|
||||
return AgentRunResponseUpdate(
|
||||
contents=[function_call],
|
||||
role=ChatRole.ASSISTANT,
|
||||
role=Role.ASSISTANT,
|
||||
author_name=self.name,
|
||||
response_id=response_id,
|
||||
message_id=str(uuid.uuid4()),
|
||||
|
||||
@@ -12,7 +12,7 @@ from typing import TYPE_CHECKING, Any, Generic, TypeVar, Union, get_args, get_or
|
||||
if TYPE_CHECKING:
|
||||
from ._workflow import Workflow
|
||||
|
||||
from agent_framework import AgentRunResponse, AgentRunResponseUpdate, AgentThread, AIAgent, ChatMessage
|
||||
from agent_framework import AgentProtocol, AgentRunResponse, AgentRunResponseUpdate, AgentThread, ChatMessage
|
||||
from agent_framework._pydantic import AFBaseModel
|
||||
from pydantic import Field
|
||||
|
||||
@@ -789,7 +789,7 @@ class AgentExecutor(Executor):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
agent: AIAgent,
|
||||
agent: AgentProtocol,
|
||||
*,
|
||||
agent_thread: AgentThread | None = None,
|
||||
streaming: bool = False,
|
||||
@@ -818,7 +818,7 @@ class AgentExecutor(Executor):
|
||||
if request.should_respond:
|
||||
if self._streaming:
|
||||
updates: list[AgentRunResponseUpdate] = []
|
||||
async for update in self._agent.run_streaming(
|
||||
async for update in self._agent.run_stream(
|
||||
self._cache,
|
||||
thread=self._agent_thread,
|
||||
):
|
||||
@@ -894,7 +894,7 @@ class WorkflowExecutor(Executor):
|
||||
|
||||
try:
|
||||
# Run the sub-workflow and collect all events
|
||||
events = [event async for event in self.workflow.run_streaming(input_data)]
|
||||
events = [event async for event in self.workflow.run_stream(input_data)]
|
||||
|
||||
# Count requests and initialize response tracking
|
||||
request_count = 0
|
||||
|
||||
@@ -14,16 +14,16 @@ from typing import Annotated, Any, Literal, Protocol, TypeVar, Union, cast
|
||||
from uuid import uuid4
|
||||
|
||||
from agent_framework import (
|
||||
AgentProtocol,
|
||||
AgentRunResponse,
|
||||
AgentRunResponseUpdate,
|
||||
AIAgent,
|
||||
ChatClient,
|
||||
ChatClientProtocol,
|
||||
ChatMessage,
|
||||
ChatRole,
|
||||
FunctionCallContent,
|
||||
FunctionResultContent,
|
||||
Role,
|
||||
)
|
||||
from agent_framework._agents import AgentBase
|
||||
from agent_framework._agents import BaseAgent
|
||||
from agent_framework._pydantic import AFBaseModel
|
||||
from pydantic import BaseModel, ConfigDict, Field
|
||||
|
||||
@@ -83,7 +83,7 @@ class MagenticAgentDeltaEvent:
|
||||
function_call_arguments: Any | None = None
|
||||
function_result_id: str | None = None
|
||||
function_result: Any | None = None
|
||||
role: ChatRole | None = None
|
||||
role: Role | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -289,7 +289,7 @@ class MagenticStartMessage:
|
||||
Returns:
|
||||
A MagenticStartMessage with the string converted to a ChatMessage.
|
||||
"""
|
||||
return cls(task=ChatMessage(role=ChatRole.USER, text=task_text))
|
||||
return cls(task=ChatMessage(role=Role.USER, text=task_text))
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -401,7 +401,7 @@ def _team_block(participants: dict[str, str]) -> str:
|
||||
|
||||
def _first_assistant(messages: list[ChatMessage]) -> ChatMessage | None:
|
||||
for msg in reversed(messages):
|
||||
if msg.role == ChatRole.ASSISTANT:
|
||||
if msg.role == Role.ASSISTANT:
|
||||
return msg
|
||||
return None
|
||||
|
||||
@@ -409,7 +409,7 @@ def _first_assistant(messages: list[ChatMessage]) -> ChatMessage | None:
|
||||
def _extract_json(text: str) -> dict[str, Any]:
|
||||
"""Potentially temp helper method.
|
||||
|
||||
Note: this method is required right now because the ChatClient, when calling
|
||||
Note: this method is required right now because the ChatClientProtocol, when calling
|
||||
response.text, returns duplicate JSON payloads - need to figure out why.
|
||||
|
||||
The `text` method is concatenating multiple text contents from diff msgs into a single string.
|
||||
@@ -497,7 +497,7 @@ class MagenticManagerBase(AFBaseModel, ABC):
|
||||
|
||||
|
||||
class StandardMagenticManager(MagenticManagerBase):
|
||||
"""Standard Magentic manager that performs real LLM calls via a ChatClientAgent.
|
||||
"""Standard Magentic manager that performs real LLM calls via a ChatAgent.
|
||||
|
||||
The manager constructs prompts that mirror the original Magentic One orchestration:
|
||||
- Facts gathering
|
||||
@@ -509,7 +509,7 @@ class StandardMagenticManager(MagenticManagerBase):
|
||||
|
||||
model_config = ConfigDict(arbitrary_types_allowed=True)
|
||||
|
||||
chat_client: ChatClient
|
||||
chat_client: ChatClientProtocol
|
||||
task_ledger: MagenticTaskLedger | None = None
|
||||
instructions: str | None = None
|
||||
|
||||
@@ -526,7 +526,7 @@ class StandardMagenticManager(MagenticManagerBase):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
chat_client: ChatClient,
|
||||
chat_client: ChatClientProtocol,
|
||||
task_ledger: MagenticTaskLedger | None = None,
|
||||
*,
|
||||
instructions: str | None = None,
|
||||
@@ -597,7 +597,7 @@ class StandardMagenticManager(MagenticManagerBase):
|
||||
*,
|
||||
response_format: type[BaseModel] | None = None,
|
||||
) -> ChatMessage:
|
||||
"""Call the underlying ChatClient directly and return the last assistant message.
|
||||
"""Call the underlying ChatClientProtocol directly and return the last assistant message.
|
||||
|
||||
If manager instructions are provided, they are injected as a SYSTEM message
|
||||
at the start of the request to guide the model consistently without needing
|
||||
@@ -606,7 +606,7 @@ class StandardMagenticManager(MagenticManagerBase):
|
||||
# Prepend system instructions if present
|
||||
request_messages: list[ChatMessage] = []
|
||||
if self.instructions:
|
||||
request_messages.append(ChatMessage(role=ChatRole.SYSTEM, text=self.instructions))
|
||||
request_messages.append(ChatMessage(role=Role.SYSTEM, text=self.instructions))
|
||||
request_messages.extend(messages)
|
||||
|
||||
# Invoke the chat client non-streaming API
|
||||
@@ -619,13 +619,13 @@ class StandardMagenticManager(MagenticManagerBase):
|
||||
if out_messages:
|
||||
last = out_messages[-1]
|
||||
return ChatMessage(
|
||||
role=last.role or ChatRole.ASSISTANT,
|
||||
role=last.role or Role.ASSISTANT,
|
||||
text=last.text or "",
|
||||
author_name=last.author_name or MAGENTIC_MANAGER_NAME,
|
||||
)
|
||||
|
||||
# Fallback if no messages
|
||||
return ChatMessage(role=ChatRole.ASSISTANT, text="No output produced.", author_name=MAGENTIC_MANAGER_NAME)
|
||||
return ChatMessage(role=Role.ASSISTANT, text="No output produced.", author_name=MAGENTIC_MANAGER_NAME)
|
||||
|
||||
async def plan(self, magentic_context: MagenticContext) -> ChatMessage:
|
||||
"""Create facts and plan using the model, then render a combined task ledger as a single assistant message."""
|
||||
@@ -634,14 +634,14 @@ class StandardMagenticManager(MagenticManagerBase):
|
||||
|
||||
# Gather facts
|
||||
facts_user = ChatMessage(
|
||||
role=ChatRole.USER,
|
||||
role=Role.USER,
|
||||
text=self.task_ledger_facts_prompt.format(task=task_text),
|
||||
)
|
||||
facts_msg = await self._complete([*magentic_context.chat_history, facts_user])
|
||||
|
||||
# Create plan
|
||||
plan_user = ChatMessage(
|
||||
role=ChatRole.USER,
|
||||
role=Role.USER,
|
||||
text=self.task_ledger_plan_prompt.format(team=team_text),
|
||||
)
|
||||
plan_msg = await self._complete([*magentic_context.chat_history, facts_user, facts_msg, plan_user])
|
||||
@@ -659,7 +659,7 @@ class StandardMagenticManager(MagenticManagerBase):
|
||||
facts=facts_msg.text,
|
||||
plan=plan_msg.text,
|
||||
)
|
||||
return ChatMessage(role=ChatRole.ASSISTANT, text=combined, author_name=MAGENTIC_MANAGER_NAME)
|
||||
return ChatMessage(role=Role.ASSISTANT, text=combined, author_name=MAGENTIC_MANAGER_NAME)
|
||||
|
||||
async def replan(self, magentic_context: MagenticContext) -> ChatMessage:
|
||||
"""Update facts and plan when stalling or looping has been detected."""
|
||||
@@ -671,14 +671,14 @@ class StandardMagenticManager(MagenticManagerBase):
|
||||
|
||||
# Update facts
|
||||
facts_update_user = ChatMessage(
|
||||
role=ChatRole.USER,
|
||||
role=Role.USER,
|
||||
text=self.task_ledger_facts_update_prompt.format(task=task_text, old_facts=self.task_ledger.facts.text),
|
||||
)
|
||||
updated_facts = await self._complete([*magentic_context.chat_history, facts_update_user])
|
||||
|
||||
# Update plan
|
||||
plan_update_user = ChatMessage(
|
||||
role=ChatRole.USER,
|
||||
role=Role.USER,
|
||||
text=self.task_ledger_plan_update_prompt.format(team=team_text),
|
||||
)
|
||||
updated_plan = await self._complete([
|
||||
@@ -701,7 +701,7 @@ class StandardMagenticManager(MagenticManagerBase):
|
||||
facts=updated_facts.text,
|
||||
plan=updated_plan.text,
|
||||
)
|
||||
return ChatMessage(role=ChatRole.ASSISTANT, text=combined, author_name=MAGENTIC_MANAGER_NAME)
|
||||
return ChatMessage(role=Role.ASSISTANT, text=combined, author_name=MAGENTIC_MANAGER_NAME)
|
||||
|
||||
async def create_progress_ledger(self, magentic_context: MagenticContext) -> MagenticProgressLedger:
|
||||
"""Use the model to produce a JSON progress ledger based on the conversation so far.
|
||||
@@ -721,7 +721,7 @@ class StandardMagenticManager(MagenticManagerBase):
|
||||
team=team_text,
|
||||
names=names_csv,
|
||||
)
|
||||
user_message = ChatMessage(role=ChatRole.USER, text=prompt)
|
||||
user_message = ChatMessage(role=Role.USER, text=prompt)
|
||||
|
||||
# Include full context to help the model decide current stage, with small retry loop
|
||||
attempts = 0
|
||||
@@ -751,11 +751,11 @@ class StandardMagenticManager(MagenticManagerBase):
|
||||
async def prepare_final_answer(self, magentic_context: MagenticContext) -> ChatMessage:
|
||||
"""Ask the model to produce the final answer addressed to the user."""
|
||||
prompt = self.final_answer_prompt.format(task=magentic_context.task.text)
|
||||
user_message = ChatMessage(role=ChatRole.USER, text=prompt)
|
||||
user_message = ChatMessage(role=Role.USER, text=prompt)
|
||||
response = await self._complete([*magentic_context.chat_history, user_message])
|
||||
# Ensure role is assistant
|
||||
return ChatMessage(
|
||||
role=ChatRole.ASSISTANT,
|
||||
role=Role.ASSISTANT,
|
||||
text=response.text,
|
||||
author_name=response.author_name or MAGENTIC_MANAGER_NAME,
|
||||
)
|
||||
@@ -896,9 +896,9 @@ class MagenticOrchestratorExecutor(Executor):
|
||||
logger.debug("Magentic Orchestrator: Received response from agent")
|
||||
|
||||
# Add transfer message if needed
|
||||
if message.body.role != ChatRole.USER:
|
||||
if message.body.role != Role.USER:
|
||||
transfer_msg = ChatMessage(
|
||||
role=ChatRole.USER,
|
||||
role=Role.USER,
|
||||
text=f"Transferred to {getattr(message.body, 'author_name', 'agent')}",
|
||||
)
|
||||
self._context.chat_history.append(transfer_msg)
|
||||
@@ -945,7 +945,7 @@ class MagenticOrchestratorExecutor(Executor):
|
||||
plan=human.edited_plan_text,
|
||||
)
|
||||
self._task_ledger = ChatMessage(
|
||||
role=ChatRole.ASSISTANT,
|
||||
role=Role.ASSISTANT,
|
||||
text=combined,
|
||||
author_name=MAGENTIC_MANAGER_NAME,
|
||||
)
|
||||
@@ -953,7 +953,7 @@ class MagenticOrchestratorExecutor(Executor):
|
||||
elif human.comments:
|
||||
# Record the human feedback for grounding
|
||||
self._context.chat_history.append(
|
||||
ChatMessage(role=ChatRole.USER, text=f"Human plan feedback: {human.comments}")
|
||||
ChatMessage(role=Role.USER, text=f"Human plan feedback: {human.comments}")
|
||||
)
|
||||
# Ask the manager to replan based on comments; proceed immediately
|
||||
self._task_ledger = await self._manager.replan(self._context.model_copy(deep=True))
|
||||
@@ -981,7 +981,7 @@ class MagenticOrchestratorExecutor(Executor):
|
||||
self._require_plan_signoff = False
|
||||
# Add a clear note to the conversation so users know review is closed
|
||||
notice = ChatMessage(
|
||||
role=ChatRole.ASSISTANT,
|
||||
role=Role.ASSISTANT,
|
||||
text=(
|
||||
"Plan review closed after max rounds. Proceeding with the current plan and will no longer "
|
||||
"prompt for plan approval."
|
||||
@@ -1015,14 +1015,14 @@ class MagenticOrchestratorExecutor(Executor):
|
||||
facts=(mgr_ledger2.facts.text if mgr_ledger2 else ""),
|
||||
plan=human.edited_plan_text,
|
||||
)
|
||||
self._task_ledger = ChatMessage(role=ChatRole.ASSISTANT, text=combined, author_name=MAGENTIC_MANAGER_NAME)
|
||||
self._task_ledger = ChatMessage(role=Role.ASSISTANT, text=combined, author_name=MAGENTIC_MANAGER_NAME)
|
||||
await self._send_plan_review_request(context)
|
||||
return
|
||||
|
||||
# Else pass comments into the chat history and replan with the manager
|
||||
if human.comments:
|
||||
self._context.chat_history.append(
|
||||
ChatMessage(role=ChatRole.USER, text=f"Human plan feedback: {human.comments}")
|
||||
ChatMessage(role=Role.USER, text=f"Human plan feedback: {human.comments}")
|
||||
)
|
||||
|
||||
# Ask the manager to replan; this only adjusts the plan stage, not a full reset
|
||||
@@ -1127,7 +1127,7 @@ class MagenticOrchestratorExecutor(Executor):
|
||||
|
||||
# Add instruction to conversation (assistant guidance)
|
||||
instruction_msg = ChatMessage(
|
||||
role=ChatRole.ASSISTANT,
|
||||
role=Role.ASSISTANT,
|
||||
text=str(instruction),
|
||||
author_name=MAGENTIC_MANAGER_NAME,
|
||||
)
|
||||
@@ -1215,7 +1215,7 @@ class MagenticOrchestratorExecutor(Executor):
|
||||
partial_result = _first_assistant(ctx.chat_history)
|
||||
if partial_result is None:
|
||||
partial_result = ChatMessage(
|
||||
role=ChatRole.ASSISTANT,
|
||||
role=Role.ASSISTANT,
|
||||
text=f"Stopped due to {limit_type} limit. No partial result available.",
|
||||
author_name=MAGENTIC_MANAGER_NAME,
|
||||
)
|
||||
@@ -1262,7 +1262,7 @@ class MagenticAgentExecutor(Executor):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
agent: AIAgent | Executor,
|
||||
agent: AgentProtocol | Executor,
|
||||
agent_id: str,
|
||||
agent_response_callback: Callable[[str, ChatMessage], Awaitable[None]] | None = None,
|
||||
streaming_agent_response_callback: Callable[[str, AgentRunResponseUpdate, bool], Awaitable[None]] | None = None,
|
||||
@@ -1288,9 +1288,9 @@ class MagenticAgentExecutor(Executor):
|
||||
return
|
||||
|
||||
# Add transfer message if needed
|
||||
if message.body.role != ChatRole.USER:
|
||||
if message.body.role != Role.USER:
|
||||
transfer_msg = ChatMessage(
|
||||
role=ChatRole.USER,
|
||||
role=Role.USER,
|
||||
text=f"Transferred to {getattr(message.body, 'author_name', 'agent')}",
|
||||
)
|
||||
self._chat_history.append(transfer_msg)
|
||||
@@ -1298,18 +1298,18 @@ class MagenticAgentExecutor(Executor):
|
||||
# Add message to agent's history
|
||||
self._chat_history.append(message.body)
|
||||
|
||||
def _get_persona_adoption_role(self) -> ChatRole:
|
||||
def _get_persona_adoption_role(self) -> Role:
|
||||
"""Determine the best role for persona adoption messages.
|
||||
|
||||
Uses SYSTEM role if the agent supports it, otherwise falls back to USER.
|
||||
"""
|
||||
# Only AgentBase-derived agents are assumed to support SYSTEM messages reliably.
|
||||
from agent_framework import AgentBase as _AF_AgentBase # local import to avoid cycles
|
||||
# Only BaseAgent-derived agents are assumed to support SYSTEM messages reliably.
|
||||
from agent_framework import BaseAgent as _AF_AgentBase # local import to avoid cycles
|
||||
|
||||
if isinstance(self._agent, _AF_AgentBase) and hasattr(self._agent, "chat_client"):
|
||||
return ChatRole.SYSTEM
|
||||
return Role.SYSTEM
|
||||
# For other agent types or when we can't determine support, use USER
|
||||
return ChatRole.USER
|
||||
return Role.USER
|
||||
|
||||
@handler
|
||||
async def handle_request_message(
|
||||
@@ -1331,14 +1331,14 @@ class MagenticAgentExecutor(Executor):
|
||||
|
||||
# Add the orchestrator's instruction as a USER message so the agent treats it as the prompt
|
||||
if message.instruction:
|
||||
self._chat_history.append(ChatMessage(role=ChatRole.USER, text=message.instruction))
|
||||
self._chat_history.append(ChatMessage(role=Role.USER, text=message.instruction))
|
||||
try:
|
||||
# If the participant is not an invokable AgentBase, return a no-op response.
|
||||
from agent_framework import AgentBase as _AF_AgentBase # local import to avoid cycles
|
||||
# If the participant is not an invokable BaseAgent, return a no-op response.
|
||||
from agent_framework import BaseAgent as _AF_AgentBase # local import to avoid cycles
|
||||
|
||||
if not isinstance(self._agent, _AF_AgentBase):
|
||||
response = ChatMessage(
|
||||
role=ChatRole.ASSISTANT,
|
||||
role=Role.ASSISTANT,
|
||||
text=f"{self._agent_id} is a workflow executor and cannot be invoked directly.",
|
||||
author_name=self._agent_id,
|
||||
)
|
||||
@@ -1354,7 +1354,7 @@ class MagenticAgentExecutor(Executor):
|
||||
logger.warning("Agent %s invoke failed: %s", self._agent_id, e)
|
||||
# Fallback response
|
||||
response = ChatMessage(
|
||||
role=ChatRole.ASSISTANT,
|
||||
role=Role.ASSISTANT,
|
||||
text=f"Agent {self._agent_id}: Error processing request - {str(e)[:100]}",
|
||||
)
|
||||
self._chat_history.append(response)
|
||||
@@ -1370,9 +1370,9 @@ class MagenticAgentExecutor(Executor):
|
||||
logger.debug(f"Agent {self._agent_id}: Running with {len(self._chat_history)} messages")
|
||||
|
||||
updates: list[AgentRunResponseUpdate] = []
|
||||
# The wrapped participant is guaranteed to be an AgentBase when this is called.
|
||||
agent = cast("AIAgent", self._agent)
|
||||
async for update in agent.run_streaming(messages=self._chat_history): # type: ignore[attr-defined]
|
||||
# The wrapped participant is guaranteed to be an BaseAgent when this is called.
|
||||
agent = cast("AgentProtocol", self._agent)
|
||||
async for update in agent.run_stream(messages=self._chat_history): # type: ignore[attr-defined]
|
||||
updates.append(update)
|
||||
if self._streaming_agent_response_callback is not None:
|
||||
with contextlib.suppress(Exception):
|
||||
@@ -1394,7 +1394,7 @@ class MagenticAgentExecutor(Executor):
|
||||
if messages and len(messages) > 0:
|
||||
last: ChatMessage = messages[-1]
|
||||
author = last.author_name or self._agent_id
|
||||
role: ChatRole = last.role if last.role else ChatRole.ASSISTANT
|
||||
role: Role = last.role if last.role else Role.ASSISTANT
|
||||
text = last.text or str(last)
|
||||
msg = ChatMessage(role=role, text=text, author_name=author)
|
||||
if self._agent_response_callback is not None:
|
||||
@@ -1403,7 +1403,7 @@ class MagenticAgentExecutor(Executor):
|
||||
return msg
|
||||
|
||||
msg = ChatMessage(
|
||||
role=ChatRole.ASSISTANT,
|
||||
role=Role.ASSISTANT,
|
||||
text=f"Agent {self._agent_id}: No output produced",
|
||||
author_name=self._agent_id,
|
||||
)
|
||||
@@ -1422,7 +1422,7 @@ class MagenticBuilder:
|
||||
"""High-level builder for creating Magentic One workflows."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._participants: dict[str, AIAgent | Executor] = {}
|
||||
self._participants: dict[str, AgentProtocol | Executor] = {}
|
||||
self._manager: MagenticManagerBase | None = None
|
||||
self._exception_callback: Callable[[Exception], None] | None = None
|
||||
self._result_callback: Callable[[ChatMessage], Awaitable[None]] | None = None
|
||||
@@ -1435,7 +1435,7 @@ class MagenticBuilder:
|
||||
self._unified_callback: CallbackSink | None = None
|
||||
self._callback_mode: MagenticCallbackMode | None = None
|
||||
|
||||
def participants(self, **participants: AIAgent | Executor) -> Self:
|
||||
def participants(self, **participants: AgentProtocol | Executor) -> Self:
|
||||
"""Add participants (agents) to the workflow."""
|
||||
self._participants.update(participants)
|
||||
return self
|
||||
@@ -1450,7 +1450,7 @@ class MagenticBuilder:
|
||||
manager: MagenticManagerBase | None = None,
|
||||
*,
|
||||
# Constructor args for StandardMagenticManager when manager is not provided
|
||||
chat_client: ChatClient | None = None,
|
||||
chat_client: ChatClientProtocol | None = None,
|
||||
task_ledger: MagenticTaskLedger | None = None,
|
||||
instructions: str | None = None,
|
||||
# Prompt overrides
|
||||
@@ -1540,7 +1540,7 @@ class MagenticBuilder:
|
||||
# Create participant descriptions
|
||||
participant_descriptions: dict[str, str] = {}
|
||||
for name, participant in self._participants.items():
|
||||
if isinstance(participant, AgentBase):
|
||||
if isinstance(participant, BaseAgent):
|
||||
description = getattr(participant, "description", None) or f"Agent {name}"
|
||||
else:
|
||||
description = f"Executor {name}"
|
||||
@@ -1745,7 +1745,7 @@ class MagenticWorkflow:
|
||||
WorkflowEvent: The events generated during the workflow execution.
|
||||
"""
|
||||
start_message = MagenticStartMessage.from_string(task_text)
|
||||
async for event in self._workflow.run_streaming(start_message):
|
||||
async for event in self._workflow.run_stream(start_message):
|
||||
yield event
|
||||
|
||||
async def run_streaming_with_message(self, task_message: ChatMessage) -> AsyncIterable[WorkflowEvent]:
|
||||
@@ -1758,10 +1758,10 @@ class MagenticWorkflow:
|
||||
WorkflowEvent: The events generated during the workflow execution.
|
||||
"""
|
||||
start_message = MagenticStartMessage(task=task_message)
|
||||
async for event in self._workflow.run_streaming(start_message):
|
||||
async for event in self._workflow.run_stream(start_message):
|
||||
yield event
|
||||
|
||||
async def run_streaming(self, message: Any | None = None) -> AsyncIterable[WorkflowEvent]:
|
||||
async def run_stream(self, message: Any | None = None) -> AsyncIterable[WorkflowEvent]:
|
||||
"""Run the workflow with either a message object or the preset task string.
|
||||
|
||||
Args:
|
||||
@@ -1780,7 +1780,7 @@ class MagenticWorkflow:
|
||||
elif isinstance(message, ChatMessage):
|
||||
message = MagenticStartMessage(task=message)
|
||||
|
||||
async for event in self._workflow.run_streaming(message):
|
||||
async for event in self._workflow.run_stream(message):
|
||||
yield event
|
||||
|
||||
async def run_with_string(self, task_text: str) -> WorkflowRunResult:
|
||||
@@ -1822,7 +1822,7 @@ class MagenticWorkflow:
|
||||
WorkflowRunResult: All events generated during the workflow execution.
|
||||
"""
|
||||
events: list[WorkflowEvent] = []
|
||||
async for event in self.run_streaming(message):
|
||||
async for event in self.run_stream(message):
|
||||
events.append(event)
|
||||
return WorkflowRunResult(events)
|
||||
|
||||
|
||||
@@ -225,7 +225,7 @@ class Workflow(AFBaseModel):
|
||||
workflow_tracer.add_workflow_error_event(e)
|
||||
raise
|
||||
|
||||
async def run_streaming(self, message: Any) -> AsyncIterable[WorkflowEvent]:
|
||||
async def run_stream(self, message: Any) -> AsyncIterable[WorkflowEvent]:
|
||||
"""Run the workflow with a starting message and stream events.
|
||||
|
||||
Args:
|
||||
@@ -252,7 +252,7 @@ class Workflow(AFBaseModel):
|
||||
async for event in self._run_workflow_with_tracing(initial_executor_fn=initial_execution, reset_context=True):
|
||||
yield event
|
||||
|
||||
async def run_streaming_from_checkpoint(
|
||||
async def run_stream_from_checkpoint(
|
||||
self,
|
||||
checkpoint_id: str,
|
||||
checkpoint_storage: CheckpointStorage | None = None,
|
||||
@@ -368,7 +368,7 @@ class Workflow(AFBaseModel):
|
||||
Returns:
|
||||
A WorkflowRunResult instance containing a list of events generated during the workflow execution.
|
||||
"""
|
||||
events = [event async for event in self.run_streaming(message)]
|
||||
events = [event async for event in self.run_stream(message)]
|
||||
return WorkflowRunResult(events)
|
||||
|
||||
async def run_from_checkpoint(
|
||||
@@ -394,7 +394,7 @@ class Workflow(AFBaseModel):
|
||||
RuntimeError: If checkpoint restoration fails.
|
||||
"""
|
||||
events = [
|
||||
event async for event in self.run_streaming_from_checkpoint(checkpoint_id, checkpoint_storage, responses)
|
||||
event async for event in self.run_stream_from_checkpoint(checkpoint_id, checkpoint_storage, responses)
|
||||
]
|
||||
return WorkflowRunResult(events)
|
||||
|
||||
|
||||
@@ -11,11 +11,11 @@ from agent_framework import (
|
||||
ChatMessage,
|
||||
ChatResponse,
|
||||
ChatResponseUpdate,
|
||||
ChatRole,
|
||||
Role,
|
||||
TextContent,
|
||||
)
|
||||
from agent_framework._agents import AgentBase
|
||||
from agent_framework._clients import ChatClient as AFChatClient
|
||||
from agent_framework._agents import BaseAgent
|
||||
from agent_framework._clients import ChatClientProtocol as AFChatClient
|
||||
|
||||
from agent_framework_workflow import (
|
||||
Executor,
|
||||
@@ -42,7 +42,7 @@ def test_magentic_start_message_from_string():
|
||||
msg = MagenticStartMessage.from_string("Do the thing")
|
||||
assert isinstance(msg, MagenticStartMessage)
|
||||
assert isinstance(msg.task, ChatMessage)
|
||||
assert msg.task.role == ChatRole.USER
|
||||
assert msg.task.role == Role.USER
|
||||
assert msg.task.text == "Do the thing"
|
||||
|
||||
|
||||
@@ -67,11 +67,11 @@ def test_plan_review_request_defaults_and_reply_variants():
|
||||
|
||||
def test_magentic_context_reset_behavior():
|
||||
ctx = MagenticContext(
|
||||
task=ChatMessage(role=ChatRole.USER, text="task"),
|
||||
task=ChatMessage(role=Role.USER, text="task"),
|
||||
participant_descriptions={"Alice": "Researcher"},
|
||||
)
|
||||
# seed context state
|
||||
ctx.chat_history.append(ChatMessage(role=ChatRole.ASSISTANT, text="draft"))
|
||||
ctx.chat_history.append(ChatMessage(role=Role.ASSISTANT, text="draft"))
|
||||
ctx.stall_count = 2
|
||||
prev_reset = ctx.reset_count
|
||||
|
||||
@@ -97,18 +97,18 @@ class FakeManager(MagenticManagerBase):
|
||||
instruction_text: str = "Proceed with step 1"
|
||||
|
||||
async def plan(self, magentic_context: MagenticContext) -> ChatMessage:
|
||||
facts = ChatMessage(role=ChatRole.ASSISTANT, text="GIVEN OR VERIFIED FACTS\n- A\n")
|
||||
plan = ChatMessage(role=ChatRole.ASSISTANT, text="- Do X\n- Do Y\n")
|
||||
facts = ChatMessage(role=Role.ASSISTANT, text="GIVEN OR VERIFIED FACTS\n- A\n")
|
||||
plan = ChatMessage(role=Role.ASSISTANT, text="- Do X\n- Do Y\n")
|
||||
self.task_ledger = _SimpleLedger(facts=facts, plan=plan)
|
||||
combined = f"Task: {magentic_context.task.text}\n\nFacts:\n{facts.text}\n\nPlan:\n{plan.text}"
|
||||
return ChatMessage(role=ChatRole.ASSISTANT, text=combined, author_name="magentic_manager")
|
||||
return ChatMessage(role=Role.ASSISTANT, text=combined, author_name="magentic_manager")
|
||||
|
||||
async def replan(self, magentic_context: MagenticContext) -> ChatMessage:
|
||||
facts = ChatMessage(role=ChatRole.ASSISTANT, text="GIVEN OR VERIFIED FACTS\n- A2\n")
|
||||
plan = ChatMessage(role=ChatRole.ASSISTANT, text="- Do Z\n")
|
||||
facts = ChatMessage(role=Role.ASSISTANT, text="GIVEN OR VERIFIED FACTS\n- A2\n")
|
||||
plan = ChatMessage(role=Role.ASSISTANT, text="- Do Z\n")
|
||||
self.task_ledger = _SimpleLedger(facts=facts, plan=plan)
|
||||
combined = f"Task: {magentic_context.task.text}\n\nFacts:\n{facts.text}\n\nPlan:\n{plan.text}"
|
||||
return ChatMessage(role=ChatRole.ASSISTANT, text=combined, author_name="magentic_manager")
|
||||
return ChatMessage(role=Role.ASSISTANT, text=combined, author_name="magentic_manager")
|
||||
|
||||
async def create_progress_ledger(self, magentic_context: MagenticContext) -> MagenticProgressLedger:
|
||||
is_satisfied = self.satisfied_after_signoff and len(magentic_context.chat_history) > 0
|
||||
@@ -121,18 +121,18 @@ class FakeManager(MagenticManagerBase):
|
||||
)
|
||||
|
||||
async def prepare_final_answer(self, magentic_context: MagenticContext) -> ChatMessage:
|
||||
return ChatMessage(role=ChatRole.ASSISTANT, text="FINAL", author_name="magentic_manager")
|
||||
return ChatMessage(role=Role.ASSISTANT, text="FINAL", author_name="magentic_manager")
|
||||
|
||||
|
||||
async def test_standard_manager_plan_and_replan_combined_ledger():
|
||||
manager = FakeManager(max_round_count=10, max_stall_count=3, max_reset_count=2)
|
||||
ctx = MagenticContext(
|
||||
task=ChatMessage(role=ChatRole.USER, text="demo task"),
|
||||
task=ChatMessage(role=Role.USER, text="demo task"),
|
||||
participant_descriptions={"agentA": "Agent A"},
|
||||
)
|
||||
|
||||
first = await manager.plan(ctx.model_copy(deep=True))
|
||||
assert first.role == ChatRole.ASSISTANT and "Facts:" in first.text and "Plan:" in first.text
|
||||
assert first.role == Role.ASSISTANT and "Facts:" in first.text and "Plan:" in first.text
|
||||
assert manager.task_ledger is not None
|
||||
|
||||
replanned = await manager.replan(ctx.model_copy(deep=True))
|
||||
@@ -142,7 +142,7 @@ async def test_standard_manager_plan_and_replan_combined_ledger():
|
||||
async def test_standard_manager_progress_ledger_and_fallback():
|
||||
manager = FakeManager(max_round_count=10)
|
||||
ctx = MagenticContext(
|
||||
task=ChatMessage(role=ChatRole.USER, text="demo"),
|
||||
task=ChatMessage(role=Role.USER, text="demo"),
|
||||
participant_descriptions={"agentA": "Agent A"},
|
||||
)
|
||||
|
||||
@@ -166,7 +166,7 @@ async def test_magentic_workflow_plan_review_approval_to_completion():
|
||||
)
|
||||
|
||||
req_event: RequestInfoEvent | None = None
|
||||
async for ev in wf.run_streaming("do work"):
|
||||
async for ev in wf.run_stream("do work"):
|
||||
if isinstance(ev, RequestInfoEvent) and ev.request_type is MagenticPlanReviewRequest:
|
||||
req_event = ev
|
||||
assert req_event is not None
|
||||
@@ -205,7 +205,7 @@ async def test_magentic_plan_review_approve_with_comments_replans_and_proceeds()
|
||||
|
||||
# Wait for the initial plan review request
|
||||
req_event: RequestInfoEvent | None = None
|
||||
async for ev in wf.run_streaming("do work"):
|
||||
async for ev in wf.run_stream("do work"):
|
||||
if isinstance(ev, RequestInfoEvent) and ev.request_type is MagenticPlanReviewRequest:
|
||||
req_event = ev
|
||||
assert req_event is not None
|
||||
@@ -242,7 +242,7 @@ async def test_magentic_orchestrator_round_limit_produces_partial_result():
|
||||
from agent_framework_workflow import WorkflowEvent # type: ignore
|
||||
|
||||
events: list[WorkflowEvent] = []
|
||||
async for ev in wf.run_streaming("round limit test"):
|
||||
async for ev in wf.run_stream("round limit test"):
|
||||
events.append(ev)
|
||||
if len(events) > 50:
|
||||
break
|
||||
@@ -251,7 +251,7 @@ async def test_magentic_orchestrator_round_limit_produces_partial_result():
|
||||
assert completed is not None
|
||||
data = getattr(completed, "data", None)
|
||||
assert isinstance(data, ChatMessage)
|
||||
assert data.role == ChatRole.ASSISTANT
|
||||
assert data.role == Role.ASSISTANT
|
||||
|
||||
|
||||
class _DummyExec(Executor):
|
||||
@@ -268,7 +268,7 @@ from agent_framework_workflow import StandardMagenticManager # noqa: E402
|
||||
|
||||
class _StubChatClient(AFChatClient):
|
||||
async def get_response(self, messages, **kwargs): # type: ignore[override]
|
||||
return ChatResponse(messages=[ChatMessage(role=ChatRole.ASSISTANT, text="ok")])
|
||||
return ChatResponse(messages=[ChatMessage(role=Role.ASSISTANT, text="ok")])
|
||||
|
||||
def get_streaming_response(self, messages, **kwargs) -> AsyncIterable[ChatResponseUpdate]: # type: ignore[override]
|
||||
async def _gen():
|
||||
@@ -284,14 +284,14 @@ async def test_standard_manager_plan_and_replan_via_complete_monkeypatch():
|
||||
async def fake_complete_plan(messages: list[ChatMessage], **kwargs: Any) -> ChatMessage:
|
||||
# Return a different response depending on call order length
|
||||
if any("FACTS" in (m.text or "") for m in messages):
|
||||
return ChatMessage(role=ChatRole.ASSISTANT, text="- step A\n- step B")
|
||||
return ChatMessage(role=ChatRole.ASSISTANT, text="GIVEN OR VERIFIED FACTS\n- fact1")
|
||||
return ChatMessage(role=Role.ASSISTANT, text="- step A\n- step B")
|
||||
return ChatMessage(role=Role.ASSISTANT, text="GIVEN OR VERIFIED FACTS\n- fact1")
|
||||
|
||||
# First, patch to produce facts then plan
|
||||
mgr._complete = fake_complete_plan # type: ignore[attr-defined]
|
||||
|
||||
ctx = MagenticContext(
|
||||
task=ChatMessage(role=ChatRole.USER, text="T"),
|
||||
task=ChatMessage(role=Role.USER, text="T"),
|
||||
participant_descriptions={"A": "desc"},
|
||||
)
|
||||
combined = await mgr.plan(ctx.model_copy(deep=True))
|
||||
@@ -303,8 +303,8 @@ async def test_standard_manager_plan_and_replan_via_complete_monkeypatch():
|
||||
# Now replan with new outputs
|
||||
async def fake_complete_replan(messages: list[ChatMessage], **kwargs: Any) -> ChatMessage:
|
||||
if any("Please briefly explain" in (m.text or "") for m in messages):
|
||||
return ChatMessage(role=ChatRole.ASSISTANT, text="- new step")
|
||||
return ChatMessage(role=ChatRole.ASSISTANT, text="GIVEN OR VERIFIED FACTS\n- updated")
|
||||
return ChatMessage(role=Role.ASSISTANT, text="- new step")
|
||||
return ChatMessage(role=Role.ASSISTANT, text="GIVEN OR VERIFIED FACTS\n- updated")
|
||||
|
||||
mgr._complete = fake_complete_replan # type: ignore[attr-defined]
|
||||
combined2 = await mgr.replan(ctx.model_copy(deep=True))
|
||||
@@ -314,7 +314,7 @@ async def test_standard_manager_plan_and_replan_via_complete_monkeypatch():
|
||||
async def test_standard_manager_progress_ledger_success_and_error():
|
||||
mgr = StandardMagenticManager(chat_client=_StubChatClient())
|
||||
ctx = MagenticContext(
|
||||
task=ChatMessage(role=ChatRole.USER, text="task"),
|
||||
task=ChatMessage(role=Role.USER, text="task"),
|
||||
participant_descriptions={"alice": "desc"},
|
||||
)
|
||||
|
||||
@@ -327,7 +327,7 @@ async def test_standard_manager_progress_ledger_success_and_error():
|
||||
'"next_speaker": {"reason": "r", "answer": "alice"}, '
|
||||
'"instruction_or_question": {"reason": "r", "answer": "do"}}'
|
||||
)
|
||||
return ChatMessage(role=ChatRole.ASSISTANT, text=json_text)
|
||||
return ChatMessage(role=Role.ASSISTANT, text=json_text)
|
||||
|
||||
mgr._complete = fake_complete_ok # type: ignore[attr-defined]
|
||||
ledger = await mgr.create_progress_ledger(ctx.model_copy(deep=True))
|
||||
@@ -335,7 +335,7 @@ async def test_standard_manager_progress_ledger_success_and_error():
|
||||
|
||||
# Error path: invalid JSON now raises to avoid emitting planner-oriented instructions to agents
|
||||
async def fake_complete_bad(messages: list[ChatMessage], **kwargs: Any) -> ChatMessage:
|
||||
return ChatMessage(role=ChatRole.ASSISTANT, text="not-json")
|
||||
return ChatMessage(role=Role.ASSISTANT, text="not-json")
|
||||
|
||||
mgr._complete = fake_complete_bad # type: ignore[attr-defined]
|
||||
with pytest.raises(RuntimeError):
|
||||
@@ -348,10 +348,10 @@ class InvokeOnceManager(MagenticManagerBase):
|
||||
self._invoked = False
|
||||
|
||||
async def plan(self, magentic_context: MagenticContext) -> ChatMessage:
|
||||
return ChatMessage(role=ChatRole.ASSISTANT, text="ledger")
|
||||
return ChatMessage(role=Role.ASSISTANT, text="ledger")
|
||||
|
||||
async def replan(self, magentic_context: MagenticContext) -> ChatMessage:
|
||||
return ChatMessage(role=ChatRole.ASSISTANT, text="re-ledger")
|
||||
return ChatMessage(role=Role.ASSISTANT, text="re-ledger")
|
||||
|
||||
async def create_progress_ledger(self, magentic_context: MagenticContext) -> MagenticProgressLedger:
|
||||
if not self._invoked:
|
||||
@@ -374,43 +374,41 @@ class InvokeOnceManager(MagenticManagerBase):
|
||||
)
|
||||
|
||||
async def prepare_final_answer(self, magentic_context: MagenticContext) -> ChatMessage:
|
||||
return ChatMessage(role=ChatRole.ASSISTANT, text="final")
|
||||
return ChatMessage(role=Role.ASSISTANT, text="final")
|
||||
|
||||
|
||||
class StubThreadAgent(AgentBase):
|
||||
async def run_streaming(self, messages=None, *, thread=None, **kwargs): # type: ignore[override]
|
||||
class StubThreadAgent(BaseAgent):
|
||||
async def run_stream(self, messages=None, *, thread=None, **kwargs): # type: ignore[override]
|
||||
yield AgentRunResponseUpdate(
|
||||
contents=[TextContent(text="thread-ok")],
|
||||
author_name="agentA",
|
||||
role=ChatRole.ASSISTANT,
|
||||
role=Role.ASSISTANT,
|
||||
)
|
||||
|
||||
async def run(self, messages=None, *, thread=None, **kwargs): # type: ignore[override]
|
||||
return AgentRunResponse(messages=[ChatMessage(role=ChatRole.ASSISTANT, text="thread-ok", author_name="agentA")])
|
||||
return AgentRunResponse(messages=[ChatMessage(role=Role.ASSISTANT, text="thread-ok", author_name="agentA")])
|
||||
|
||||
|
||||
class StubAssistantsClient:
|
||||
pass # class name used for branch detection
|
||||
|
||||
|
||||
class StubAssistantsAgent(AgentBase):
|
||||
class StubAssistantsAgent(BaseAgent):
|
||||
chat_client: object | None = None # allow assignment via Pydantic field
|
||||
|
||||
def __init__(self) -> None:
|
||||
super().__init__()
|
||||
self.chat_client = StubAssistantsClient() # type name contains 'AssistantsClient'
|
||||
|
||||
async def run_streaming(self, messages=None, *, thread=None, **kwargs): # type: ignore[override]
|
||||
async def run_stream(self, messages=None, *, thread=None, **kwargs): # type: ignore[override]
|
||||
yield AgentRunResponseUpdate(
|
||||
contents=[TextContent(text="assistants-ok")],
|
||||
author_name="agentA",
|
||||
role=ChatRole.ASSISTANT,
|
||||
role=Role.ASSISTANT,
|
||||
)
|
||||
|
||||
async def run(self, messages=None, *, thread=None, **kwargs): # type: ignore[override]
|
||||
return AgentRunResponse(
|
||||
messages=[ChatMessage(role=ChatRole.ASSISTANT, text="assistants-ok", author_name="agentA")]
|
||||
)
|
||||
return AgentRunResponse(messages=[ChatMessage(role=Role.ASSISTANT, text="assistants-ok", author_name="agentA")])
|
||||
|
||||
|
||||
async def _collect_agent_responses_setup(participant_obj: object):
|
||||
@@ -432,7 +430,7 @@ async def _collect_agent_responses_setup(participant_obj: object):
|
||||
|
||||
# Run a bounded stream to allow one invoke and then completion
|
||||
events: list[WorkflowEvent] = []
|
||||
async for ev in wf.run_streaming("task"): # plan review disabled
|
||||
async for ev in wf.run_stream("task"): # plan review disabled
|
||||
events.append(ev)
|
||||
if len(events) > 50:
|
||||
break
|
||||
|
||||
@@ -343,7 +343,7 @@ async def test_end_to_end_workflow_tracing(tracing_enabled: Any, span_exporter:
|
||||
|
||||
# Run workflow (this should create run spans)
|
||||
events = []
|
||||
async for event in workflow.run_streaming("test input"):
|
||||
async for event in workflow.run_stream("test input"):
|
||||
events.append(event)
|
||||
|
||||
# Verify workflow executed correctly
|
||||
@@ -444,7 +444,7 @@ async def test_workflow_error_handling_in_tracing(tracing_enabled: Any, span_exp
|
||||
|
||||
# Run workflow and expect error
|
||||
with pytest.raises(ValueError, match="Test error"):
|
||||
async for _ in workflow.run_streaming("test input"):
|
||||
async for _ in workflow.run_stream("test input"):
|
||||
pass
|
||||
|
||||
spans = span_exporter.get_finished_spans()
|
||||
|
||||
@@ -95,7 +95,7 @@ async def test_workflow_run_streaming():
|
||||
)
|
||||
|
||||
result: int | None = None
|
||||
async for event in workflow.run_streaming(NumberMessage(data=0)):
|
||||
async for event in workflow.run_stream(NumberMessage(data=0)):
|
||||
assert isinstance(event, WorkflowEvent)
|
||||
if isinstance(event, WorkflowCompletedEvent):
|
||||
result = event.data
|
||||
@@ -118,7 +118,7 @@ async def test_workflow_run_stream_not_completed():
|
||||
)
|
||||
|
||||
with pytest.raises(RuntimeError):
|
||||
async for _ in workflow.run_streaming(NumberMessage(data=0)):
|
||||
async for _ in workflow.run_stream(NumberMessage(data=0)):
|
||||
pass
|
||||
|
||||
|
||||
@@ -176,7 +176,7 @@ async def test_workflow_send_responses_streaming():
|
||||
)
|
||||
|
||||
request_info_event: RequestInfoEvent | None = None
|
||||
async for event in workflow.run_streaming(NumberMessage(data=0)):
|
||||
async for event in workflow.run_stream(NumberMessage(data=0)):
|
||||
if isinstance(event, RequestInfoEvent):
|
||||
request_info_event = event
|
||||
|
||||
@@ -326,7 +326,7 @@ async def test_workflow_checkpointing_not_enabled_for_external_restore(simple_ex
|
||||
|
||||
# Attempt to restore from checkpoint without providing external storage should fail
|
||||
try:
|
||||
[event async for event in workflow.run_streaming_from_checkpoint("fake-checkpoint-id")]
|
||||
[event async for event in workflow.run_stream_from_checkpoint("fake-checkpoint-id")]
|
||||
raise AssertionError("Expected ValueError to be raised")
|
||||
except ValueError as e:
|
||||
assert "Cannot restore from checkpoint" in str(e)
|
||||
@@ -344,7 +344,7 @@ async def test_workflow_run_stream_from_checkpoint_no_checkpointing_enabled(simp
|
||||
|
||||
# Attempt to run from checkpoint should fail
|
||||
try:
|
||||
async for _ in workflow.run_streaming_from_checkpoint("fake_checkpoint_id"):
|
||||
async for _ in workflow.run_stream_from_checkpoint("fake_checkpoint_id"):
|
||||
pass
|
||||
raise AssertionError("Expected ValueError to be raised")
|
||||
except ValueError as e:
|
||||
@@ -368,7 +368,7 @@ async def test_workflow_run_stream_from_checkpoint_invalid_checkpoint(simple_exe
|
||||
|
||||
# Attempt to run from non-existent checkpoint should fail
|
||||
try:
|
||||
async for _ in workflow.run_streaming_from_checkpoint("nonexistent_checkpoint_id"):
|
||||
async for _ in workflow.run_stream_from_checkpoint("nonexistent_checkpoint_id"):
|
||||
pass
|
||||
raise AssertionError("Expected RuntimeError to be raised")
|
||||
except RuntimeError as e:
|
||||
@@ -401,7 +401,7 @@ async def test_workflow_run_stream_from_checkpoint_with_external_storage(simple_
|
||||
# Resume from checkpoint using external storage parameter
|
||||
try:
|
||||
events: list[WorkflowEvent] = []
|
||||
async for event in workflow_without_checkpointing.run_streaming_from_checkpoint(
|
||||
async for event in workflow_without_checkpointing.run_stream_from_checkpoint(
|
||||
checkpoint_id, checkpoint_storage=storage
|
||||
):
|
||||
events.append(event)
|
||||
@@ -446,7 +446,7 @@ async def test_workflow_run_from_checkpoint_non_streaming(simple_executor: Execu
|
||||
|
||||
|
||||
async def test_workflow_run_stream_from_checkpoint_with_responses(simple_executor: Executor):
|
||||
"""Test that run_streaming_from_checkpoint accepts responses parameter."""
|
||||
"""Test that run_stream_from_checkpoint accepts responses parameter."""
|
||||
with tempfile.TemporaryDirectory() as temp_dir:
|
||||
storage = FileCheckpointStorage(temp_dir)
|
||||
|
||||
@@ -477,7 +477,7 @@ async def test_workflow_run_stream_from_checkpoint_with_responses(simple_executo
|
||||
|
||||
try:
|
||||
events: list[WorkflowEvent] = []
|
||||
async for event in workflow.run_streaming_from_checkpoint(checkpoint_id, responses=responses):
|
||||
async for event in workflow.run_stream_from_checkpoint(checkpoint_id, responses=responses):
|
||||
events.append(event)
|
||||
if len(events) >= 2: # Limit to avoid infinite loops
|
||||
break
|
||||
|
||||
@@ -8,8 +8,8 @@ from agent_framework import (
|
||||
AgentRunResponse,
|
||||
AgentRunResponseUpdate,
|
||||
ChatMessage,
|
||||
ChatRole,
|
||||
FunctionResultContent,
|
||||
Role,
|
||||
TextContent,
|
||||
UsageContent,
|
||||
UsageDetails,
|
||||
@@ -43,11 +43,11 @@ class SimpleExecutor(Executor):
|
||||
response_text = f"{self.response_text}: {input_text}"
|
||||
|
||||
# Create response message for both streaming and non-streaming cases
|
||||
response_message = ChatMessage(role=ChatRole.ASSISTANT, contents=[TextContent(text=response_text)])
|
||||
response_message = ChatMessage(role=Role.ASSISTANT, contents=[TextContent(text=response_text)])
|
||||
|
||||
# Emit update event.
|
||||
streaming_update = AgentRunResponseUpdate(
|
||||
contents=[TextContent(text=response_text)], role=ChatRole.ASSISTANT, message_id=str(uuid.uuid4())
|
||||
contents=[TextContent(text=response_text)], role=Role.ASSISTANT, message_id=str(uuid.uuid4())
|
||||
)
|
||||
await ctx.add_event(AgentRunUpdateEvent(executor_id=self.id, data=streaming_update))
|
||||
|
||||
@@ -68,7 +68,7 @@ class RequestingExecutor(Executor):
|
||||
# Handle the response and emit completion response
|
||||
update = AgentRunResponseUpdate(
|
||||
contents=[TextContent(text="Request completed successfully")],
|
||||
role=ChatRole.ASSISTANT,
|
||||
role=Role.ASSISTANT,
|
||||
message_id=str(uuid.uuid4()),
|
||||
)
|
||||
await ctx.add_event(AgentRunUpdateEvent(executor_id=self.id, data=update))
|
||||
@@ -132,7 +132,7 @@ class TestWorkflowAgent:
|
||||
|
||||
# Execute workflow streaming to capture streaming events
|
||||
updates = []
|
||||
async for update in agent.run_streaming("Test input"):
|
||||
async for update in agent.run_stream("Test input"):
|
||||
updates.append(update)
|
||||
|
||||
# Should have received at least one streaming update
|
||||
@@ -165,7 +165,7 @@ class TestWorkflowAgent:
|
||||
|
||||
# Execute workflow streaming to get request info event
|
||||
updates = []
|
||||
async for update in agent.run_streaming("Start request"):
|
||||
async for update in agent.run_stream("Start request"):
|
||||
updates.append(update)
|
||||
# Should have received a function call for the request info
|
||||
assert len(updates) > 0
|
||||
@@ -192,7 +192,7 @@ class TestWorkflowAgent:
|
||||
|
||||
# Now provide a function result response to test continuation
|
||||
response_message = ChatMessage(
|
||||
role=ChatRole.USER,
|
||||
role=Role.USER,
|
||||
contents=[FunctionResultContent(call_id=function_call.call_id, result="User provided answer")],
|
||||
)
|
||||
|
||||
@@ -252,7 +252,7 @@ class TestWorkflowAgentMergeUpdates:
|
||||
# Response B, Message 2 (latest in resp B)
|
||||
AgentRunResponseUpdate(
|
||||
contents=[TextContent(text="RespB-Msg2")],
|
||||
role=ChatRole.ASSISTANT,
|
||||
role=Role.ASSISTANT,
|
||||
response_id="resp-b",
|
||||
message_id="msg-2",
|
||||
created_at="2024-01-01T12:02:00Z",
|
||||
@@ -260,7 +260,7 @@ class TestWorkflowAgentMergeUpdates:
|
||||
# Response A, Message 1 (earliest overall)
|
||||
AgentRunResponseUpdate(
|
||||
contents=[TextContent(text="RespA-Msg1")],
|
||||
role=ChatRole.ASSISTANT,
|
||||
role=Role.ASSISTANT,
|
||||
response_id="resp-a",
|
||||
message_id="msg-1",
|
||||
created_at="2024-01-01T12:00:00Z",
|
||||
@@ -268,7 +268,7 @@ class TestWorkflowAgentMergeUpdates:
|
||||
# Response B, Message 1 (earlier in resp B)
|
||||
AgentRunResponseUpdate(
|
||||
contents=[TextContent(text="RespB-Msg1")],
|
||||
role=ChatRole.ASSISTANT,
|
||||
role=Role.ASSISTANT,
|
||||
response_id="resp-b",
|
||||
message_id="msg-1",
|
||||
created_at="2024-01-01T12:01:00Z",
|
||||
@@ -276,7 +276,7 @@ class TestWorkflowAgentMergeUpdates:
|
||||
# Response A, Message 2 (later in resp A)
|
||||
AgentRunResponseUpdate(
|
||||
contents=[TextContent(text="RespA-Msg2")],
|
||||
role=ChatRole.ASSISTANT,
|
||||
role=Role.ASSISTANT,
|
||||
response_id="resp-a",
|
||||
message_id="msg-2",
|
||||
created_at="2024-01-01T12:00:30Z",
|
||||
@@ -284,7 +284,7 @@ class TestWorkflowAgentMergeUpdates:
|
||||
# Global dangling update (no response_id) - should go at end
|
||||
AgentRunResponseUpdate(
|
||||
contents=[TextContent(text="Global-Dangling")],
|
||||
role=ChatRole.ASSISTANT,
|
||||
role=Role.ASSISTANT,
|
||||
response_id=None,
|
||||
message_id="msg-global",
|
||||
created_at="2024-01-01T11:59:00Z", # Earliest timestamp but should be last
|
||||
@@ -360,7 +360,7 @@ class TestWorkflowAgentMergeUpdates:
|
||||
details=UsageDetails(input_token_count=10, output_token_count=5, total_token_count=15)
|
||||
),
|
||||
],
|
||||
role=ChatRole.ASSISTANT,
|
||||
role=Role.ASSISTANT,
|
||||
response_id="resp-1",
|
||||
message_id="msg-1",
|
||||
created_at="2024-01-01T12:00:00Z",
|
||||
@@ -373,7 +373,7 @@ class TestWorkflowAgentMergeUpdates:
|
||||
details=UsageDetails(input_token_count=20, output_token_count=8, total_token_count=28)
|
||||
),
|
||||
],
|
||||
role=ChatRole.ASSISTANT,
|
||||
role=Role.ASSISTANT,
|
||||
response_id="resp-2",
|
||||
message_id="msg-2",
|
||||
created_at="2024-01-01T12:01:00Z", # Later timestamp
|
||||
@@ -384,7 +384,7 @@ class TestWorkflowAgentMergeUpdates:
|
||||
TextContent(text="Third"),
|
||||
UsageContent(details=UsageDetails(input_token_count=5, output_token_count=3, total_token_count=8)),
|
||||
],
|
||||
role=ChatRole.ASSISTANT,
|
||||
role=Role.ASSISTANT,
|
||||
response_id="resp-1", # Same response_id as first
|
||||
message_id="msg-3",
|
||||
created_at="2024-01-01T11:59:00Z", # Earlier timestamp
|
||||
|
||||
Reference in New Issue
Block a user