mirror of
https://github.com/microsoft/agent-framework.git
synced 2026-06-16 21:04:09 +08:00
Python: [BREAKING] Python: Provider-leading client design & OpenAI package extraction (#4818)
* Python: Provider-leading client design & OpenAI package extraction Major refactoring of the Python Agent Framework client architecture: - Extract OpenAI clients into new `agent-framework-openai` package - Core package no longer depends on openai, azure-identity, azure-ai-projects - Rename clients for discoverability: OpenAIResponsesClient → OpenAIChatClient, OpenAIChatClient → OpenAIChatCompletionClient - Unify `model_id`/`deployment_name`/`model_deployment_name` → `model` param - New FoundryChatClient for Azure AI Foundry Responses API - New FoundryAgent/FoundryAgentClient for connecting to pre-configured Foundry agents - Remove OpenAIBase/OpenAIConfigMixin from non-deprecated client MRO - Deprecate AzureOpenAI* clients, AzureAIClient, OpenAIAssistantsClient - Reorganize samples: azure_openai+azure_ai+azure_ai_agent → azure/ - ADR-0020: Provider-Leading Client Design Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> * fix: missing Agent imports in samples, .model_id → .model in foundry_local sample Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> * fix: CI failures — mypy errors, coverage targets, sample imports - azure-ai mypy: add type ignores for TypedDict total=, model arg, forward ref - Coverage: replace core.azure/openai targets with openai package target - project_provider: add type annotation for opts dict Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> * fix: populate openai .pyi stub, fix broken README links, coverage targets Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> * fixes * updated observabilitty * reset azure init.pyi * fix errors * updated adr number * fix foundry local * fixed not renamed docstrings and comments, and added deprecated markers to old classes * fix tests and pyprojects * fix test vars * updated function tests * update durable * updated test setup for functions * Fix Foundry auth in workflow samples Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> * Stabilize Python integration workflows Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> * Update hosting samples for Foundry Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> * Trigger full CI rerun Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> * Trigger CI rerun again Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> * trigger rerun * trigger rerun * fix for litellm * undo durabletask changes * Move Foundry APIs into foundry namespace Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> * Fix Foundry pyproject formatting Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> * Split provider samples by Foundry surface Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> * Restore hosting sample requirements Also fix the Foundry Local sample link after the provider sample move. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> * updated tests * udpated foundry integration tests * removed dist from azurefunctions tests * Use separate Foundry clients for concurrent agents Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> * fix client setup in azfunc and durable * disabled two tests * updated setup for some function and durable tests * improved azure openai setup with new clients * ignore deprecated * fixes * skip 11 * remove openai assistants int tests --------- Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
This commit is contained in:
committed by
GitHub
Unverified
parent
4b533608b6
commit
5e056b672e
@@ -0,0 +1,34 @@
|
||||
# AGENTS.md — agent-framework-openai
|
||||
|
||||
OpenAI integration package for Agent Framework. Contains OpenAI Responses API and Chat Completions API clients.
|
||||
|
||||
## Package Structure
|
||||
|
||||
```
|
||||
agent_framework_openai/
|
||||
├── __init__.py # Public API exports
|
||||
├── _chat_client.py # OpenAIChatClient (Responses API) + RawOpenAIChatClient
|
||||
├── _chat_completion_client.py # OpenAIChatCompletionClient (Chat Completions API) + RawOpenAIChatCompletionClient
|
||||
├── _embedding_client.py # OpenAIEmbeddingClient
|
||||
├── _exceptions.py # OpenAI-specific exceptions
|
||||
├── _shared.py # OpenAIBase, OpenAIConfigMixin, OpenAISettings
|
||||
├── _assistants_client.py # OpenAIAssistantsClient (DEPRECATED)
|
||||
└── _assistant_provider.py # OpenAIAssistantProvider (DEPRECATED)
|
||||
```
|
||||
|
||||
## Key Classes
|
||||
|
||||
| Class | API | Status |
|
||||
|---|---|---|
|
||||
| `OpenAIChatClient` | Responses API | Primary |
|
||||
| `OpenAIChatCompletionClient` | Chat Completions API | Primary |
|
||||
| `OpenAIEmbeddingClient` | Embeddings API | Primary |
|
||||
| `OpenAIAssistantsClient` | Assistants API | Deprecated |
|
||||
|
||||
All clients follow the Raw + Full-Featured pattern (e.g., `RawOpenAIChatClient` + `OpenAIChatClient`).
|
||||
|
||||
## Dependencies
|
||||
|
||||
- `agent-framework-core` — core abstractions
|
||||
- `openai` — OpenAI Python SDK
|
||||
- `packaging` — version checking
|
||||
@@ -0,0 +1,21 @@
|
||||
MIT License
|
||||
|
||||
Copyright (c) Microsoft Corporation.
|
||||
|
||||
Permission is hereby granted, free of charge, to any person obtaining a copy
|
||||
of this software and associated documentation files (the "Software"), to deal
|
||||
in the Software without restriction, including without limitation the rights
|
||||
to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
|
||||
copies of the Software, and to permit persons to whom the Software is
|
||||
furnished to do so, subject to the following conditions:
|
||||
|
||||
The above copyright notice and this permission notice shall be included in all
|
||||
copies or substantial portions of the Software.
|
||||
|
||||
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
|
||||
IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
|
||||
FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
|
||||
AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
|
||||
LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
|
||||
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
|
||||
SOFTWARE
|
||||
@@ -0,0 +1,17 @@
|
||||
# agent-framework-openai
|
||||
|
||||
OpenAI integration for Microsoft Agent Framework. Provides chat clients for the OpenAI Responses API and Chat Completions API.
|
||||
|
||||
## Installation
|
||||
|
||||
```bash
|
||||
pip install agent-framework-openai
|
||||
```
|
||||
|
||||
## Usage
|
||||
|
||||
```python
|
||||
from agent_framework.openai import OpenAIChatClient
|
||||
|
||||
client = OpenAIChatClient(model_id="gpt-4o")
|
||||
```
|
||||
@@ -0,0 +1,87 @@
|
||||
# Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
"""OpenAI integration for Microsoft Agent Framework.
|
||||
|
||||
This package provides OpenAI client implementations for the Agent Framework,
|
||||
including clients for the Responses API and Chat Completions API.
|
||||
"""
|
||||
|
||||
import importlib.metadata
|
||||
import sys
|
||||
|
||||
if sys.version_info >= (3, 13):
|
||||
from warnings import deprecated # type: ignore # pragma: no cover
|
||||
else:
|
||||
from typing_extensions import deprecated # type: ignore # pragma: no cover
|
||||
|
||||
from ._assistant_provider import OpenAIAssistantProvider
|
||||
from ._assistants_client import (
|
||||
AssistantToolResources,
|
||||
OpenAIAssistantsClient,
|
||||
OpenAIAssistantsOptions,
|
||||
)
|
||||
from ._chat_client import (
|
||||
OpenAIChatClient,
|
||||
OpenAIChatOptions,
|
||||
OpenAIContinuationToken,
|
||||
RawOpenAIChatClient,
|
||||
)
|
||||
from ._chat_completion_client import (
|
||||
OpenAIChatCompletionClient,
|
||||
OpenAIChatCompletionOptions,
|
||||
RawOpenAIChatCompletionClient,
|
||||
)
|
||||
from ._embedding_client import OpenAIEmbeddingClient, OpenAIEmbeddingOptions
|
||||
from ._exceptions import ContentFilterResultSeverity, OpenAIContentFilterException
|
||||
from ._shared import OpenAISettings
|
||||
|
||||
try:
|
||||
__version__ = importlib.metadata.version("agent-framework-openai")
|
||||
except importlib.metadata.PackageNotFoundError:
|
||||
__version__ = "0.0.0" # Fallback for development mode
|
||||
|
||||
# Deprecated aliases for old names — use subclasses so the warning only fires for the alias
|
||||
|
||||
|
||||
@deprecated(
|
||||
"OpenAIResponsesClient is deprecated, use OpenAIChatClient instead.",
|
||||
category=DeprecationWarning,
|
||||
)
|
||||
class OpenAIResponsesClient(OpenAIChatClient): # type: ignore[misc]
|
||||
"""Deprecated alias for :class:`OpenAIChatClient`."""
|
||||
|
||||
|
||||
@deprecated(
|
||||
"RawOpenAIResponsesClient is deprecated, use RawOpenAIChatClient instead.",
|
||||
category=DeprecationWarning,
|
||||
)
|
||||
class RawOpenAIResponsesClient(RawOpenAIChatClient): # type: ignore[misc]
|
||||
"""Deprecated alias for :class:`RawOpenAIChatClient`."""
|
||||
|
||||
|
||||
OpenAIResponsesOptions = OpenAIChatOptions
|
||||
"""Deprecated alias for :class:`OpenAIChatOptions`."""
|
||||
|
||||
|
||||
__all__ = [
|
||||
"AssistantToolResources",
|
||||
"ContentFilterResultSeverity",
|
||||
"OpenAIAssistantProvider",
|
||||
"OpenAIAssistantsClient",
|
||||
"OpenAIAssistantsOptions",
|
||||
"OpenAIChatClient",
|
||||
"OpenAIChatCompletionClient",
|
||||
"OpenAIChatCompletionOptions",
|
||||
"OpenAIChatOptions",
|
||||
"OpenAIContentFilterException",
|
||||
"OpenAIContinuationToken",
|
||||
"OpenAIEmbeddingClient",
|
||||
"OpenAIEmbeddingOptions",
|
||||
"OpenAIResponsesClient",
|
||||
"OpenAIResponsesOptions",
|
||||
"OpenAISettings",
|
||||
"RawOpenAIChatClient",
|
||||
"RawOpenAIChatCompletionClient",
|
||||
"RawOpenAIResponsesClient",
|
||||
"__version__",
|
||||
]
|
||||
@@ -0,0 +1,564 @@
|
||||
# Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import sys
|
||||
from collections.abc import Awaitable, Callable, Mapping, MutableMapping, Sequence
|
||||
from typing import TYPE_CHECKING, Any, Generic, cast
|
||||
|
||||
from agent_framework._agents import Agent
|
||||
from agent_framework._middleware import MiddlewareTypes
|
||||
from agent_framework._sessions import BaseContextProvider
|
||||
from agent_framework._settings import SecretString, load_settings
|
||||
from agent_framework._tools import FunctionTool, ToolTypes, normalize_tools
|
||||
from openai import AsyncOpenAI
|
||||
from openai.types.beta.assistant import Assistant
|
||||
from pydantic import BaseModel
|
||||
|
||||
from ._assistants_client import OpenAIAssistantsClient
|
||||
from ._shared import OpenAISettings, from_assistant_tools, to_assistant_tools
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from ._assistants_client import OpenAIAssistantsOptions
|
||||
|
||||
if sys.version_info >= (3, 13):
|
||||
from typing import TypeVar # type:ignore # pragma: no cover
|
||||
else:
|
||||
from typing_extensions import TypeVar # type:ignore # pragma: no cover
|
||||
if sys.version_info >= (3, 11):
|
||||
from typing import Self, TypedDict # type:ignore # pragma: no cover
|
||||
else:
|
||||
from typing_extensions import Self, TypedDict # type:ignore # pragma: no cover
|
||||
|
||||
|
||||
# Type variable for options - allows typed OpenAIAssistantProvider[OptionsCoT] returns
|
||||
# Default matches OpenAIAssistantsClient's default options type
|
||||
OptionsCoT = TypeVar(
|
||||
"OptionsCoT",
|
||||
bound=TypedDict, # type: ignore[valid-type]
|
||||
default="OpenAIAssistantsOptions",
|
||||
covariant=True,
|
||||
)
|
||||
|
||||
|
||||
class OpenAIAssistantProvider(Generic[OptionsCoT]):
|
||||
"""Provider for creating Agent instances from OpenAI Assistants API.
|
||||
|
||||
This provider allows you to create, retrieve, and wrap OpenAI Assistants
|
||||
as Agent instances for use in the agent framework.
|
||||
|
||||
Examples:
|
||||
Basic usage with automatic client creation:
|
||||
|
||||
.. code-block:: python
|
||||
|
||||
from agent_framework.openai import OpenAIAssistantProvider
|
||||
|
||||
# Uses OPENAI_API_KEY environment variable
|
||||
provider = OpenAIAssistantProvider()
|
||||
|
||||
# Create a new assistant
|
||||
agent = await provider.create_agent(
|
||||
name="MyAssistant",
|
||||
model="gpt-4",
|
||||
instructions="You are a helpful assistant.",
|
||||
tools=[my_function],
|
||||
)
|
||||
|
||||
result = await agent.run("Hello!")
|
||||
|
||||
Using an existing client:
|
||||
|
||||
.. code-block:: python
|
||||
|
||||
from openai import AsyncOpenAI
|
||||
from agent_framework.openai import OpenAIAssistantProvider
|
||||
|
||||
client = AsyncOpenAI()
|
||||
provider = OpenAIAssistantProvider(client)
|
||||
|
||||
# Get an existing assistant by ID
|
||||
agent = await provider.get_agent(
|
||||
assistant_id="asst_123",
|
||||
tools=[my_function], # Provide implementations for function tools
|
||||
)
|
||||
|
||||
Wrapping an SDK Assistant object:
|
||||
|
||||
.. code-block:: python
|
||||
|
||||
# Fetch assistant directly via SDK
|
||||
assistant = await client.beta.assistants.retrieve("asst_123")
|
||||
|
||||
# Wrap without additional HTTP call
|
||||
agent = provider.as_agent(assistant, tools=[my_function])
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
client: AsyncOpenAI | None = None,
|
||||
*,
|
||||
api_key: str | SecretString | Callable[[], str | Awaitable[str]] | None = None,
|
||||
org_id: str | None = None,
|
||||
base_url: str | None = None,
|
||||
env_file_path: str | None = None,
|
||||
env_file_encoding: str | None = None,
|
||||
) -> None:
|
||||
"""Initialize the OpenAI Assistant Provider.
|
||||
|
||||
Args:
|
||||
client: An existing AsyncOpenAI client to use. If not provided,
|
||||
a new client will be created using the other parameters.
|
||||
|
||||
Keyword Args:
|
||||
api_key: OpenAI API key. Can also be set via OPENAI_API_KEY env var.
|
||||
org_id: OpenAI organization ID. Can also be set via OPENAI_ORG_ID env var.
|
||||
base_url: Base URL for the OpenAI API. Can also be set via OPENAI_BASE_URL env var.
|
||||
env_file_path: Path to .env file for configuration.
|
||||
env_file_encoding: Encoding of the .env file.
|
||||
|
||||
Raises:
|
||||
ValueError: If no client is provided and API key is missing.
|
||||
|
||||
Examples:
|
||||
.. code-block:: python
|
||||
|
||||
# Using environment variables
|
||||
provider = OpenAIAssistantProvider()
|
||||
|
||||
# Using explicit API key
|
||||
provider = OpenAIAssistantProvider(api_key="sk-...")
|
||||
|
||||
# Using existing client
|
||||
client = AsyncOpenAI()
|
||||
provider = OpenAIAssistantProvider(client)
|
||||
"""
|
||||
self._client: AsyncOpenAI | None = client
|
||||
self._should_close_client: bool = client is None
|
||||
|
||||
if client is None:
|
||||
# Load settings and create client
|
||||
settings = load_settings(
|
||||
OpenAISettings,
|
||||
env_prefix="OPENAI_",
|
||||
api_key=api_key,
|
||||
org_id=org_id,
|
||||
base_url=base_url,
|
||||
env_file_path=env_file_path,
|
||||
env_file_encoding=env_file_encoding,
|
||||
)
|
||||
|
||||
api_key_setting = settings.get("api_key")
|
||||
if not api_key_setting:
|
||||
raise ValueError(
|
||||
"OpenAI API key is required. Set via 'api_key' parameter or 'OPENAI_API_KEY' environment variable."
|
||||
)
|
||||
|
||||
# Get API key value
|
||||
api_key_value: str | Callable[[], str | Awaitable[str]]
|
||||
if isinstance(api_key_setting, SecretString):
|
||||
api_key_value = api_key_setting.get_secret_value()
|
||||
else:
|
||||
api_key_value = api_key_setting
|
||||
|
||||
# Create client
|
||||
client_args: dict[str, Any] = {"api_key": api_key_value}
|
||||
if org_id_value := settings.get("org_id"):
|
||||
client_args["organization"] = org_id_value
|
||||
if base_url_value := settings.get("base_url"):
|
||||
client_args["base_url"] = base_url_value
|
||||
|
||||
self._client = AsyncOpenAI(**client_args)
|
||||
|
||||
async def __aenter__(self) -> Self:
|
||||
"""Async context manager entry."""
|
||||
return self
|
||||
|
||||
async def __aexit__(self, exc_type: type[BaseException] | None, exc_val: BaseException | None, exc_tb: Any) -> None:
|
||||
"""Async context manager exit."""
|
||||
await self.close()
|
||||
|
||||
async def close(self) -> None:
|
||||
"""Close the provider and clean up resources.
|
||||
|
||||
If the provider created its own client, it will be closed.
|
||||
If an external client was provided, it will not be closed.
|
||||
"""
|
||||
if self._should_close_client and self._client is not None:
|
||||
await self._client.close()
|
||||
|
||||
async def create_agent(
|
||||
self,
|
||||
*,
|
||||
name: str,
|
||||
model: str,
|
||||
instructions: str | None = None,
|
||||
description: str | None = None,
|
||||
tools: ToolTypes | Callable[..., Any] | Sequence[ToolTypes | Callable[..., Any]] | None = None,
|
||||
metadata: dict[str, str] | None = None,
|
||||
default_options: OptionsCoT | None = None,
|
||||
middleware: Sequence[MiddlewareTypes] | None = None,
|
||||
context_providers: Sequence[BaseContextProvider] | None = None,
|
||||
) -> Agent[OptionsCoT]:
|
||||
"""Create a new assistant on OpenAI and return a Agent.
|
||||
|
||||
This method creates a new assistant on the OpenAI service and wraps it
|
||||
in a Agent instance. The assistant will persist on OpenAI until deleted.
|
||||
|
||||
Keyword Args:
|
||||
name: The name of the assistant (required).
|
||||
model: The model ID to use, e.g., "gpt-4", "gpt-4o" (required).
|
||||
instructions: System instructions for the assistant.
|
||||
description: A description of the assistant.
|
||||
tools: Tools available to the assistant. Can include:
|
||||
- FunctionTool instances or callables decorated with @tool
|
||||
- Dict-based tools from OpenAIAssistantsClient.get_code_interpreter_tool()
|
||||
- Dict-based tools from OpenAIAssistantsClient.get_file_search_tool()
|
||||
- Raw tool dictionaries
|
||||
metadata: Metadata to attach to the assistant (max 16 key-value pairs).
|
||||
default_options: A TypedDict containing default chat options for the agent.
|
||||
These options are applied to every run unless overridden.
|
||||
Include ``response_format`` here for structured output responses.
|
||||
middleware: MiddlewareTypes for the Agent.
|
||||
context_providers: Context providers for the Agent.
|
||||
|
||||
Returns:
|
||||
A Agent instance wrapping the created assistant.
|
||||
|
||||
Raises:
|
||||
ValueError: If assistant creation fails.
|
||||
|
||||
Examples:
|
||||
.. code-block:: python
|
||||
|
||||
provider = OpenAIAssistantProvider()
|
||||
|
||||
# Create with function tools
|
||||
agent = await provider.create_agent(
|
||||
name="WeatherBot",
|
||||
model="gpt-4",
|
||||
instructions="You are a helpful weather assistant.",
|
||||
tools=[get_weather],
|
||||
)
|
||||
|
||||
# Create with structured output
|
||||
agent = await provider.create_agent(
|
||||
name="StructuredBot",
|
||||
model="gpt-4",
|
||||
default_options={"response_format": MyPydanticModel},
|
||||
)
|
||||
"""
|
||||
# Normalize tools
|
||||
normalized_tools = normalize_tools(tools)
|
||||
assistant_tools: list[FunctionTool | MutableMapping[str, Any]] = [
|
||||
tool for tool in normalized_tools if isinstance(tool, (FunctionTool, MutableMapping))
|
||||
]
|
||||
api_tools = to_assistant_tools(assistant_tools) if assistant_tools else []
|
||||
|
||||
# Extract response_format from default_options if present
|
||||
opts = dict(default_options) if default_options else {}
|
||||
response_format = opts.get("response_format")
|
||||
|
||||
# Build assistant creation parameters
|
||||
create_params: dict[str, Any] = {
|
||||
"model": model,
|
||||
"name": name,
|
||||
}
|
||||
|
||||
if instructions is not None:
|
||||
create_params["instructions"] = instructions
|
||||
if description is not None:
|
||||
create_params["description"] = description
|
||||
if api_tools:
|
||||
create_params["tools"] = api_tools
|
||||
if metadata is not None:
|
||||
create_params["metadata"] = metadata
|
||||
|
||||
# Handle response format for OpenAI API
|
||||
if response_format is not None and isinstance(response_format, type) and issubclass(response_format, BaseModel):
|
||||
create_params["response_format"] = {
|
||||
"type": "json_schema",
|
||||
"json_schema": {
|
||||
"name": response_format.__name__,
|
||||
"schema": response_format.model_json_schema(),
|
||||
"strict": True,
|
||||
},
|
||||
}
|
||||
|
||||
# Create the assistant
|
||||
if not self._client:
|
||||
raise RuntimeError("OpenAI client is not initialized.")
|
||||
|
||||
assistant = await self._client.beta.assistants.create(**create_params) # type: ignore[reportDeprecated]
|
||||
|
||||
# Create Agent - pass default_options which contains response_format
|
||||
return self._create_chat_agent_from_assistant(
|
||||
assistant=assistant,
|
||||
tools=normalized_tools,
|
||||
instructions=instructions,
|
||||
middleware=middleware,
|
||||
context_providers=context_providers,
|
||||
default_options=default_options,
|
||||
)
|
||||
|
||||
async def get_agent(
|
||||
self,
|
||||
assistant_id: str,
|
||||
*,
|
||||
tools: ToolTypes | Callable[..., Any] | Sequence[ToolTypes | Callable[..., Any]] | None = None,
|
||||
instructions: str | None = None,
|
||||
default_options: OptionsCoT | None = None,
|
||||
middleware: Sequence[MiddlewareTypes] | None = None,
|
||||
context_providers: Sequence[BaseContextProvider] | None = None,
|
||||
) -> Agent[OptionsCoT]:
|
||||
"""Retrieve an existing assistant by ID and return a Agent.
|
||||
|
||||
This method fetches an existing assistant from OpenAI by its ID
|
||||
and wraps it in a Agent instance.
|
||||
|
||||
Args:
|
||||
assistant_id: The ID of the assistant to retrieve (e.g., "asst_123").
|
||||
|
||||
Keyword Args:
|
||||
tools: Function tools to make available. IMPORTANT: If the assistant
|
||||
was created with function tools, you MUST provide matching
|
||||
implementations here. Hosted tools (code_interpreter, file_search)
|
||||
are automatically included.
|
||||
instructions: Override the assistant's instructions (optional).
|
||||
default_options: A TypedDict containing default chat options for the agent.
|
||||
These options are applied to every run unless overridden.
|
||||
middleware: MiddlewareTypes for the Agent.
|
||||
context_providers: Context providers for the Agent.
|
||||
|
||||
Returns:
|
||||
A Agent instance wrapping the retrieved assistant.
|
||||
|
||||
Raises:
|
||||
RuntimeError: If the assistant cannot be retrieved.
|
||||
ValueError: If required function tools are missing.
|
||||
|
||||
Examples:
|
||||
.. code-block:: python
|
||||
|
||||
provider = OpenAIAssistantProvider()
|
||||
|
||||
# Get assistant without function tools
|
||||
agent = await provider.get_agent(assistant_id="asst_123")
|
||||
|
||||
# Get assistant with function tools
|
||||
agent = await provider.get_agent(
|
||||
assistant_id="asst_456",
|
||||
tools=[get_weather, search_database], # Implementations required!
|
||||
)
|
||||
"""
|
||||
# Fetch the assistant
|
||||
if not self._client:
|
||||
raise RuntimeError("OpenAI client is not initialized.")
|
||||
|
||||
assistant = await self._client.beta.assistants.retrieve(assistant_id) # type: ignore[reportDeprecated]
|
||||
|
||||
# Use as_agent to wrap it
|
||||
return self.as_agent(
|
||||
assistant=assistant,
|
||||
tools=tools,
|
||||
instructions=instructions,
|
||||
default_options=default_options,
|
||||
middleware=middleware,
|
||||
context_providers=context_providers,
|
||||
)
|
||||
|
||||
def as_agent(
|
||||
self,
|
||||
assistant: Assistant,
|
||||
*,
|
||||
tools: ToolTypes | Callable[..., Any] | Sequence[ToolTypes | Callable[..., Any]] | None = None,
|
||||
instructions: str | None = None,
|
||||
default_options: OptionsCoT | None = None,
|
||||
middleware: Sequence[MiddlewareTypes] | None = None,
|
||||
context_providers: Sequence[BaseContextProvider] | None = None,
|
||||
) -> Agent[OptionsCoT]:
|
||||
"""Wrap an existing SDK Assistant object as a Agent.
|
||||
|
||||
This method does NOT make any HTTP calls. It simply wraps an already-
|
||||
fetched Assistant object in a Agent.
|
||||
|
||||
Args:
|
||||
assistant: The OpenAI Assistant SDK object to wrap.
|
||||
|
||||
Keyword Args:
|
||||
tools: Function tools to make available. If the assistant has
|
||||
function tools defined, you MUST provide matching implementations.
|
||||
Hosted tools (code_interpreter, file_search) are automatically included.
|
||||
instructions: Override the assistant's instructions (optional).
|
||||
default_options: A TypedDict containing default chat options for the agent.
|
||||
These options are applied to every run unless overridden.
|
||||
middleware: MiddlewareTypes for the Agent.
|
||||
context_providers: Context providers for the Agent.
|
||||
|
||||
Returns:
|
||||
A Agent instance wrapping the assistant.
|
||||
|
||||
Raises:
|
||||
ValueError: If required function tools are missing.
|
||||
|
||||
Examples:
|
||||
.. code-block:: python
|
||||
|
||||
client = AsyncOpenAI()
|
||||
provider = OpenAIAssistantProvider(client)
|
||||
|
||||
# Fetch assistant via SDK
|
||||
assistant = await client.beta.assistants.retrieve("asst_123")
|
||||
|
||||
# Wrap without additional HTTP call
|
||||
agent = provider.as_agent(
|
||||
assistant,
|
||||
tools=[my_function],
|
||||
instructions="Custom instructions override",
|
||||
)
|
||||
"""
|
||||
# Validate that required function tools are provided
|
||||
self._validate_function_tools(assistant.tools or [], tools)
|
||||
|
||||
# Merge hosted tools with user-provided function tools
|
||||
merged_tools = self._merge_tools(assistant.tools or [], tools)
|
||||
|
||||
# Create Agent
|
||||
return self._create_chat_agent_from_assistant(
|
||||
assistant=assistant,
|
||||
tools=merged_tools,
|
||||
instructions=instructions,
|
||||
default_options=default_options,
|
||||
middleware=middleware,
|
||||
context_providers=context_providers,
|
||||
)
|
||||
|
||||
def _validate_function_tools(
|
||||
self,
|
||||
assistant_tools: list[Any],
|
||||
provided_tools: ToolTypes | Callable[..., Any] | Sequence[ToolTypes | Callable[..., Any]] | None,
|
||||
) -> None:
|
||||
"""Validate that required function tools are provided.
|
||||
|
||||
Args:
|
||||
assistant_tools: Tools defined on the assistant.
|
||||
provided_tools: Tools provided by the user.
|
||||
|
||||
Raises:
|
||||
ValueError: If a required function tool is missing.
|
||||
"""
|
||||
# Get function tool names from assistant
|
||||
required_functions: set[str] = set()
|
||||
for tool in assistant_tools:
|
||||
if (
|
||||
hasattr(tool, "type")
|
||||
and tool.type == "function"
|
||||
and hasattr(tool, "function")
|
||||
and hasattr(tool.function, "name")
|
||||
):
|
||||
required_functions.add(tool.function.name)
|
||||
|
||||
if not required_functions:
|
||||
return # No function tools required
|
||||
|
||||
# Get provided function names using normalize_tools
|
||||
provided_functions: set[str] = set()
|
||||
if provided_tools is not None:
|
||||
normalized = normalize_tools(provided_tools)
|
||||
for tool in normalized:
|
||||
if isinstance(tool, FunctionTool):
|
||||
provided_functions.add(tool.name)
|
||||
elif isinstance(tool, Mapping):
|
||||
typed_tool = cast(Mapping[str, Any], tool)
|
||||
raw_func_spec = typed_tool.get("function")
|
||||
if isinstance(raw_func_spec, Mapping):
|
||||
typed_func_spec = cast(Mapping[str, Any], raw_func_spec)
|
||||
raw_name = typed_func_spec.get("name")
|
||||
if isinstance(raw_name, str) and raw_name:
|
||||
provided_functions.add(raw_name)
|
||||
|
||||
# Check for missing functions
|
||||
missing = required_functions - provided_functions
|
||||
if missing:
|
||||
missing_list = ", ".join(sorted(missing))
|
||||
raise ValueError(
|
||||
f"Assistant requires function tool(s) '{missing_list}' but no implementation was provided. "
|
||||
f"Please pass the function implementation(s) in the 'tools' parameter."
|
||||
)
|
||||
|
||||
def _merge_tools(
|
||||
self,
|
||||
assistant_tools: list[Any],
|
||||
user_tools: ToolTypes | Callable[..., Any] | Sequence[ToolTypes | Callable[..., Any]] | None,
|
||||
) -> list[FunctionTool | MutableMapping[str, Any] | Any]:
|
||||
"""Merge hosted tools from assistant with user-provided function tools.
|
||||
|
||||
Args:
|
||||
assistant_tools: Tools defined on the assistant.
|
||||
user_tools: Tools provided by the user.
|
||||
|
||||
Returns:
|
||||
A list of all tools (hosted tools + user function implementations).
|
||||
"""
|
||||
merged: list[FunctionTool | MutableMapping[str, Any] | Any] = []
|
||||
|
||||
# Add hosted tools from assistant using shared conversion
|
||||
hosted_tools = from_assistant_tools(assistant_tools)
|
||||
merged.extend(hosted_tools)
|
||||
|
||||
# Add user-provided tools (normalized)
|
||||
if user_tools is not None:
|
||||
normalized_user_tools = normalize_tools(user_tools)
|
||||
merged.extend(normalized_user_tools)
|
||||
|
||||
return merged
|
||||
|
||||
def _create_chat_agent_from_assistant(
|
||||
self,
|
||||
assistant: Assistant,
|
||||
tools: list[FunctionTool | MutableMapping[str, Any] | Any] | None,
|
||||
instructions: str | None,
|
||||
middleware: Sequence[MiddlewareTypes] | None,
|
||||
context_providers: Sequence[BaseContextProvider] | None,
|
||||
default_options: OptionsCoT | None = None,
|
||||
**kwargs: Any,
|
||||
) -> Agent[OptionsCoT]:
|
||||
"""Create a Agent from an Assistant.
|
||||
|
||||
Args:
|
||||
assistant: The OpenAI Assistant object.
|
||||
tools: Tools for the agent.
|
||||
instructions: Instructions override.
|
||||
middleware: MiddlewareTypes for the agent.
|
||||
context_providers: Context providers for the agent.
|
||||
default_options: Default chat options for the agent (may include response_format).
|
||||
**kwargs: Additional arguments passed to Agent.
|
||||
|
||||
Returns:
|
||||
A configured Agent instance.
|
||||
"""
|
||||
# Create the chat client with the assistant
|
||||
client = OpenAIAssistantsClient(
|
||||
model=assistant.model,
|
||||
assistant_id=assistant.id,
|
||||
assistant_name=assistant.name,
|
||||
assistant_description=assistant.description,
|
||||
async_client=self._client,
|
||||
)
|
||||
|
||||
# Use instructions from assistant if not overridden
|
||||
final_instructions = instructions if instructions is not None else assistant.instructions
|
||||
|
||||
# Create and return Agent
|
||||
return Agent(
|
||||
client=client,
|
||||
id=assistant.id,
|
||||
name=assistant.name,
|
||||
description=assistant.description,
|
||||
instructions=final_instructions,
|
||||
tools=tools if tools else None,
|
||||
middleware=middleware,
|
||||
context_providers=context_providers,
|
||||
default_options=default_options, # type: ignore[arg-type]
|
||||
**kwargs,
|
||||
)
|
||||
@@ -0,0 +1,958 @@
|
||||
# Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
import sys
|
||||
from collections.abc import (
|
||||
AsyncIterable,
|
||||
Awaitable,
|
||||
Callable,
|
||||
Mapping,
|
||||
MutableMapping,
|
||||
Sequence,
|
||||
)
|
||||
from typing import TYPE_CHECKING, Any, Generic, Literal, TypedDict, cast
|
||||
|
||||
from agent_framework._clients import BaseChatClient
|
||||
from agent_framework._middleware import ChatMiddlewareLayer
|
||||
from agent_framework._settings import load_settings
|
||||
from agent_framework._tools import (
|
||||
FunctionInvocationConfiguration,
|
||||
FunctionInvocationLayer,
|
||||
FunctionTool,
|
||||
normalize_tools,
|
||||
)
|
||||
from agent_framework._types import (
|
||||
Annotation,
|
||||
ChatOptions,
|
||||
ChatResponse,
|
||||
ChatResponseUpdate,
|
||||
Content,
|
||||
Message,
|
||||
ResponseStream,
|
||||
TextSpanRegion,
|
||||
UsageDetails,
|
||||
)
|
||||
from agent_framework.observability import ChatTelemetryLayer
|
||||
from openai import AsyncOpenAI
|
||||
from openai.types.beta.threads import (
|
||||
FileCitationAnnotation,
|
||||
FileCitationDeltaAnnotation,
|
||||
FilePathAnnotation,
|
||||
FilePathDeltaAnnotation,
|
||||
ImageURLContentBlockParam,
|
||||
ImageURLParam,
|
||||
MessageContentPartParam,
|
||||
MessageDeltaEvent,
|
||||
Run,
|
||||
TextContentBlockParam,
|
||||
TextDeltaBlock,
|
||||
)
|
||||
from openai.types.beta.threads import (
|
||||
Message as ThreadMessage,
|
||||
)
|
||||
from openai.types.beta.threads.run_create_params import AdditionalMessage
|
||||
from openai.types.beta.threads.run_submit_tool_outputs_params import ToolOutput
|
||||
from openai.types.beta.threads.runs import RunStep
|
||||
from pydantic import BaseModel
|
||||
|
||||
from ._shared import OpenAIConfigMixin, OpenAISettings
|
||||
|
||||
if sys.version_info >= (3, 13):
|
||||
from typing import TypeVar # type: ignore # pragma: no cover
|
||||
else:
|
||||
from typing_extensions import TypeVar # type: ignore # pragma: no cover
|
||||
|
||||
if sys.version_info >= (3, 12):
|
||||
from typing import override # type: ignore # pragma: no cover
|
||||
else:
|
||||
from typing_extensions import override # type: ignore # pragma: no cover
|
||||
|
||||
if sys.version_info >= (3, 11):
|
||||
from typing import Self, TypedDict # type: ignore # pragma: no cover
|
||||
else:
|
||||
from typing_extensions import Self, TypedDict # type: ignore # pragma: no cover
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from agent_framework._middleware import MiddlewareTypes
|
||||
|
||||
logger = logging.getLogger("agent_framework.openai")
|
||||
|
||||
|
||||
# region OpenAI Assistants Options TypedDict
|
||||
|
||||
ResponseModelT = TypeVar("ResponseModelT", bound=BaseModel | None, default=None)
|
||||
|
||||
|
||||
class VectorStoreToolResource(TypedDict, total=False):
|
||||
"""Vector store configuration for file search tool resources."""
|
||||
|
||||
vector_store_ids: list[str]
|
||||
"""IDs of vector stores attached to this assistant."""
|
||||
|
||||
|
||||
class CodeInterpreterToolResource(TypedDict, total=False):
|
||||
"""Code interpreter tool resource configuration."""
|
||||
|
||||
file_ids: list[str]
|
||||
"""File IDs accessible by the code interpreter tool. Max 20 files per assistant."""
|
||||
|
||||
|
||||
class AssistantToolResources(TypedDict, total=False):
|
||||
"""Tool resources attached to the assistant.
|
||||
|
||||
See: https://platform.openai.com/docs/api-reference/assistants/createAssistant#assistants-createassistant-tool_resources
|
||||
"""
|
||||
|
||||
code_interpreter: CodeInterpreterToolResource
|
||||
"""Resources for code interpreter tool, including file IDs."""
|
||||
|
||||
file_search: VectorStoreToolResource
|
||||
"""Resources for file search tool, including vector store IDs."""
|
||||
|
||||
|
||||
class OpenAIAssistantsOptions(ChatOptions[ResponseModelT], Generic[ResponseModelT], total=False):
|
||||
"""OpenAI Assistants API-specific options dict.
|
||||
|
||||
Extends base ChatOptions with Assistants API-specific parameters
|
||||
for creating and running assistants.
|
||||
|
||||
See: https://platform.openai.com/docs/api-reference/assistants
|
||||
|
||||
Keys:
|
||||
# Inherited from ChatOptions:
|
||||
model_id: Deprecated. The model to use for the assistant,
|
||||
translates to ``model`` in OpenAI API.
|
||||
temperature: Sampling temperature between 0 and 2.
|
||||
top_p: Nucleus sampling parameter.
|
||||
max_tokens: Maximum number of tokens to generate,
|
||||
translates to ``max_completion_tokens`` in OpenAI API.
|
||||
tools: List of tools (functions, code_interpreter, file_search).
|
||||
tool_choice: How the model should use tools.
|
||||
allow_multiple_tool_calls: Whether to allow parallel tool calls,
|
||||
translates to ``parallel_tool_calls`` in OpenAI API.
|
||||
response_format: Structured output schema.
|
||||
metadata: Request metadata for tracking.
|
||||
|
||||
# Options not supported in Assistants API (inherited but unused):
|
||||
stop: Not supported.
|
||||
seed: Not supported (use assistant-level configuration instead).
|
||||
frequency_penalty: Not supported.
|
||||
presence_penalty: Not supported.
|
||||
user: Not supported.
|
||||
store: Not supported.
|
||||
|
||||
# Assistants-specific options:
|
||||
name: Name of the assistant.
|
||||
description: Description of the assistant.
|
||||
instructions: System instructions for the assistant.
|
||||
tool_resources: Resources for tools (file IDs, vector stores).
|
||||
reasoning_effort: Effort level for o-series reasoning models.
|
||||
conversation_id: Thread ID to continue conversation in.
|
||||
"""
|
||||
|
||||
# Assistants-specific options
|
||||
name: str
|
||||
"""Name of the assistant (max 256 characters)."""
|
||||
|
||||
description: str
|
||||
"""Description of the assistant (max 512 characters)."""
|
||||
|
||||
tool_resources: AssistantToolResources
|
||||
"""Tool-specific resources like file IDs and vector stores."""
|
||||
|
||||
reasoning_effort: Literal["low", "medium", "high"]
|
||||
"""Effort level for o-series reasoning models (o1, o3-mini).
|
||||
Higher effort = more reasoning time and potentially better results."""
|
||||
|
||||
conversation_id: str # type: ignore[misc]
|
||||
"""Thread ID to continue a conversation in an existing thread."""
|
||||
|
||||
# OpenAI/ChatOptions fields not supported in Assistants API
|
||||
stop: None # type: ignore[misc]
|
||||
"""Not supported in Assistants API."""
|
||||
|
||||
seed: None # type: ignore[misc]
|
||||
"""Not supported in Assistants API (use assistant-level configuration)."""
|
||||
|
||||
frequency_penalty: None # type: ignore[misc]
|
||||
"""Not supported in Assistants API."""
|
||||
|
||||
presence_penalty: None # type: ignore[misc]
|
||||
"""Not supported in Assistants API."""
|
||||
|
||||
user: None # type: ignore[misc]
|
||||
"""Not supported in Assistants API."""
|
||||
|
||||
store: None # type: ignore[misc]
|
||||
"""Not supported in Assistants API."""
|
||||
|
||||
|
||||
ASSISTANTS_OPTION_TRANSLATIONS: dict[str, str] = {
|
||||
"model_id": "model", # backward compat: accept model_id in options
|
||||
"max_tokens": "max_completion_tokens",
|
||||
"allow_multiple_tool_calls": "parallel_tool_calls",
|
||||
}
|
||||
"""Maps ChatOptions keys to OpenAI Assistants API parameter names."""
|
||||
|
||||
OpenAIAssistantsOptionsT = TypeVar(
|
||||
"OpenAIAssistantsOptionsT",
|
||||
bound=TypedDict, # type: ignore[valid-type]
|
||||
default="OpenAIAssistantsOptions",
|
||||
covariant=True,
|
||||
)
|
||||
|
||||
|
||||
# endregion
|
||||
|
||||
|
||||
class OpenAIAssistantsClient( # type: ignore[misc]
|
||||
OpenAIConfigMixin,
|
||||
FunctionInvocationLayer[OpenAIAssistantsOptionsT],
|
||||
ChatMiddlewareLayer[OpenAIAssistantsOptionsT],
|
||||
ChatTelemetryLayer[OpenAIAssistantsOptionsT],
|
||||
BaseChatClient[OpenAIAssistantsOptionsT],
|
||||
Generic[OpenAIAssistantsOptionsT],
|
||||
):
|
||||
"""OpenAI Assistants client with middleware, telemetry, and function invocation support."""
|
||||
|
||||
# region Hosted Tool Factory Methods
|
||||
|
||||
@staticmethod
|
||||
def get_code_interpreter_tool() -> dict[str, Any]:
|
||||
"""Create a code interpreter tool configuration for the Assistants API.
|
||||
|
||||
Returns:
|
||||
A dict tool configuration ready to pass to ChatAgent.
|
||||
|
||||
Examples:
|
||||
.. code-block:: python
|
||||
|
||||
from agent_framework.openai import OpenAIAssistantsClient
|
||||
|
||||
# Enable code interpreter
|
||||
tool = OpenAIAssistantsClient.get_code_interpreter_tool()
|
||||
|
||||
agent = ChatAgent(client, tools=[tool])
|
||||
"""
|
||||
return {"type": "code_interpreter"}
|
||||
|
||||
@staticmethod
|
||||
def get_file_search_tool(
|
||||
*,
|
||||
max_num_results: int | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""Create a file search tool configuration for the Assistants API.
|
||||
|
||||
Keyword Args:
|
||||
max_num_results: Maximum number of results to return from file search.
|
||||
|
||||
Returns:
|
||||
A dict tool configuration ready to pass to ChatAgent.
|
||||
|
||||
Examples:
|
||||
.. code-block:: python
|
||||
|
||||
from agent_framework.openai import OpenAIAssistantsClient
|
||||
|
||||
# Basic file search
|
||||
tool = OpenAIAssistantsClient.get_file_search_tool()
|
||||
|
||||
# With result limit
|
||||
tool = OpenAIAssistantsClient.get_file_search_tool(max_num_results=10)
|
||||
|
||||
agent = ChatAgent(client, tools=[tool])
|
||||
"""
|
||||
tool: dict[str, Any] = {"type": "file_search"}
|
||||
|
||||
if max_num_results is not None:
|
||||
tool["file_search"] = {"max_num_results": max_num_results}
|
||||
|
||||
return tool
|
||||
|
||||
# endregion
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
model: str | None = None,
|
||||
model_id: str | None = None,
|
||||
assistant_id: str | None = None,
|
||||
assistant_name: str | None = None,
|
||||
assistant_description: str | None = None,
|
||||
thread_id: str | None = None,
|
||||
api_key: str | Callable[[], str | Awaitable[str]] | None = None,
|
||||
org_id: str | None = None,
|
||||
base_url: str | None = None,
|
||||
default_headers: Mapping[str, str] | None = None,
|
||||
async_client: AsyncOpenAI | None = None,
|
||||
env_file_path: str | None = None,
|
||||
env_file_encoding: str | None = None,
|
||||
middleware: Sequence[MiddlewareTypes] | None = None,
|
||||
function_invocation_configuration: FunctionInvocationConfiguration | None = None,
|
||||
**kwargs: Any,
|
||||
) -> None:
|
||||
"""Initialize an OpenAI Assistants client.
|
||||
|
||||
Keyword Args:
|
||||
model: OpenAI model name, see https://platform.openai.com/docs/models.
|
||||
Can also be set via environment variable OPENAI_MODEL.
|
||||
model_id: Deprecated alias for ``model``.
|
||||
assistant_id: The ID of an OpenAI assistant to use.
|
||||
If not provided, a new assistant will be created (and deleted after the request).
|
||||
assistant_name: The name to use when creating new assistants.
|
||||
assistant_description: The description to use when creating new assistants.
|
||||
thread_id: Default thread ID to use for conversations. Can be overridden by
|
||||
conversation_id property when making a request.
|
||||
If not provided, a new thread will be created (and deleted after the request).
|
||||
api_key: The API key to use. If provided will override the env vars or .env file value.
|
||||
Can also be set via environment variable OPENAI_API_KEY.
|
||||
org_id: The org ID to use. If provided will override the env vars or .env file value.
|
||||
Can also be set via environment variable OPENAI_ORG_ID.
|
||||
base_url: The base URL to use. If provided will override the standard value.
|
||||
Can also be set via environment variable OPENAI_BASE_URL.
|
||||
default_headers: The default headers mapping of string keys to
|
||||
string values for HTTP requests.
|
||||
async_client: An existing client to use.
|
||||
env_file_path: Use the environment settings file as a fallback
|
||||
to environment variables.
|
||||
env_file_encoding: The encoding of the environment settings file.
|
||||
middleware: Optional sequence of middleware to apply to requests.
|
||||
function_invocation_configuration: Optional configuration for function invocation behavior.
|
||||
kwargs: Other keyword parameters.
|
||||
|
||||
Examples:
|
||||
.. code-block:: python
|
||||
|
||||
from agent_framework.openai import OpenAIAssistantsClient
|
||||
|
||||
# Using environment variables
|
||||
# Set OPENAI_API_KEY=sk-...
|
||||
# Set OPENAI_MODEL=gpt-4
|
||||
client = OpenAIAssistantsClient()
|
||||
|
||||
# Or passing parameters directly
|
||||
client = OpenAIAssistantsClient(model="gpt-4", api_key="sk-...")
|
||||
|
||||
# Or loading from a .env file
|
||||
client = OpenAIAssistantsClient(env_file_path="path/to/.env")
|
||||
|
||||
# Using custom ChatOptions with type safety:
|
||||
from typing import TypedDict
|
||||
from agent_framework.openai import OpenAIAssistantsOptions
|
||||
|
||||
|
||||
class MyOptions(OpenAIAssistantsOptions, total=False):
|
||||
my_custom_option: str
|
||||
|
||||
|
||||
client: OpenAIAssistantsClient[MyOptions] = OpenAIAssistantsClient(model="gpt-4")
|
||||
response = await client.get_response("Hello", options={"my_custom_option": "value"})
|
||||
"""
|
||||
if model_id is not None and model is None:
|
||||
import warnings
|
||||
|
||||
warnings.warn("model_id is deprecated, use model instead", DeprecationWarning, stacklevel=2)
|
||||
model = model_id
|
||||
openai_settings = load_settings(
|
||||
OpenAISettings,
|
||||
env_prefix="OPENAI_",
|
||||
api_key=api_key,
|
||||
base_url=base_url,
|
||||
org_id=org_id,
|
||||
model=model,
|
||||
env_file_path=env_file_path,
|
||||
env_file_encoding=env_file_encoding,
|
||||
)
|
||||
|
||||
api_key_value = openai_settings.get("api_key")
|
||||
if not async_client and not api_key_value:
|
||||
raise ValueError(
|
||||
"OpenAI API key is required. Set via 'api_key' parameter or 'OPENAI_API_KEY' environment variable."
|
||||
)
|
||||
|
||||
resolved_model = openai_settings.get("model")
|
||||
if not resolved_model:
|
||||
raise ValueError(
|
||||
"OpenAI model is required. Set via 'model' parameter or 'OPENAI_MODEL' environment variable."
|
||||
)
|
||||
|
||||
super().__init__(
|
||||
model=resolved_model,
|
||||
api_key=self._get_api_key(api_key_value),
|
||||
org_id=openai_settings.get("org_id"),
|
||||
default_headers=default_headers,
|
||||
client=async_client,
|
||||
base_url=openai_settings.get("base_url"),
|
||||
middleware=middleware,
|
||||
function_invocation_configuration=function_invocation_configuration,
|
||||
)
|
||||
self.assistant_id: str | None = assistant_id
|
||||
self.assistant_name: str | None = assistant_name
|
||||
self.assistant_description: str | None = assistant_description
|
||||
self.thread_id: str | None = thread_id
|
||||
self._should_delete_assistant: bool = False
|
||||
|
||||
async def __aenter__(self) -> Self:
|
||||
"""Async context manager entry."""
|
||||
return self
|
||||
|
||||
async def __aexit__(
|
||||
self,
|
||||
exc_type: type[BaseException] | None,
|
||||
exc_val: BaseException | None,
|
||||
exc_tb: Any,
|
||||
) -> None:
|
||||
"""Async context manager exit - clean up any assistants we created."""
|
||||
await self.close()
|
||||
|
||||
async def close(self) -> None:
|
||||
"""Clean up any assistants we created."""
|
||||
if self._should_delete_assistant and self.assistant_id is not None:
|
||||
client = await self._ensure_client()
|
||||
await client.beta.assistants.delete(self.assistant_id) # type: ignore[reportDeprecated]
|
||||
object.__setattr__(self, "assistant_id", None)
|
||||
object.__setattr__(self, "_should_delete_assistant", False)
|
||||
|
||||
@override
|
||||
def _inner_get_response(
|
||||
self,
|
||||
*,
|
||||
messages: Sequence[Message],
|
||||
options: Mapping[str, Any],
|
||||
stream: bool = False,
|
||||
**kwargs: Any,
|
||||
) -> Awaitable[ChatResponse] | ResponseStream[ChatResponseUpdate, ChatResponse]:
|
||||
if stream:
|
||||
# Streaming mode - return the async generator directly
|
||||
async def _stream() -> AsyncIterable[ChatResponseUpdate]:
|
||||
# prepare
|
||||
run_options, tool_results = self._prepare_options(messages, options, **kwargs)
|
||||
|
||||
# Get the thread ID
|
||||
thread_id: str | None = options.get(
|
||||
"conversation_id", run_options.get("conversation_id", self.thread_id)
|
||||
)
|
||||
|
||||
if thread_id is None and tool_results is not None:
|
||||
raise ValueError("No thread ID was provided, but chat messages includes tool results.")
|
||||
|
||||
# Determine which assistant to use and create if needed
|
||||
assistant_id = await self._get_assistant_id_or_create()
|
||||
|
||||
# execute
|
||||
stream_obj, thread_id = await self._create_assistant_stream(
|
||||
thread_id, assistant_id, run_options, tool_results
|
||||
)
|
||||
|
||||
# process
|
||||
async for update in self._process_stream_events(stream_obj, thread_id):
|
||||
yield update
|
||||
|
||||
return self._build_response_stream(_stream(), response_format=options.get("response_format"))
|
||||
|
||||
# Non-streaming mode - collect updates and convert to response
|
||||
async def _get_response() -> ChatResponse:
|
||||
stream_result = self._inner_get_response(messages=messages, options=options, stream=True, **kwargs)
|
||||
return await ChatResponse.from_update_generator(
|
||||
updates=stream_result, # type: ignore[arg-type]
|
||||
output_format_type=options.get("response_format"), # type: ignore[arg-type]
|
||||
)
|
||||
|
||||
return _get_response()
|
||||
|
||||
async def _get_assistant_id_or_create(self) -> str:
|
||||
"""Determine which assistant to use and create if needed.
|
||||
|
||||
Returns:
|
||||
str: The assistant_id to use.
|
||||
"""
|
||||
# If no assistant is provided, create a temporary assistant
|
||||
if self.assistant_id is None:
|
||||
if not self.model:
|
||||
raise ValueError("Parameter 'model' is required for assistant creation.")
|
||||
|
||||
client = await self._ensure_client()
|
||||
created_assistant = await client.beta.assistants.create( # type: ignore[reportDeprecated]
|
||||
model=self.model,
|
||||
description=self.assistant_description,
|
||||
name=self.assistant_name,
|
||||
)
|
||||
self.assistant_id = created_assistant.id
|
||||
self._should_delete_assistant = True
|
||||
|
||||
return self.assistant_id
|
||||
|
||||
async def _create_assistant_stream(
|
||||
self,
|
||||
thread_id: str | None,
|
||||
assistant_id: str,
|
||||
run_options: dict[str, Any],
|
||||
tool_results: list[Content] | None,
|
||||
) -> tuple[Any, str]:
|
||||
"""Create the assistant stream for processing.
|
||||
|
||||
Returns:
|
||||
tuple: (stream, final_thread_id)
|
||||
"""
|
||||
client = await self._ensure_client()
|
||||
# Get any active run for this thread
|
||||
thread_run = await self._get_active_thread_run(thread_id)
|
||||
|
||||
tool_run_id, tool_outputs = self._prepare_tool_outputs_for_assistants(tool_results)
|
||||
|
||||
if thread_run is not None and tool_run_id is not None and tool_run_id == thread_run.id and tool_outputs:
|
||||
# There's an active run and we have tool results to submit, so submit the results.
|
||||
stream = client.beta.threads.runs.submit_tool_outputs_stream( # type: ignore[reportDeprecated]
|
||||
run_id=tool_run_id,
|
||||
thread_id=thread_run.thread_id,
|
||||
tool_outputs=tool_outputs,
|
||||
)
|
||||
final_thread_id = thread_run.thread_id
|
||||
else:
|
||||
# Handle thread creation or cancellation
|
||||
final_thread_id = await self._prepare_thread(thread_id, thread_run, run_options)
|
||||
|
||||
# Now create a new run and stream the results.
|
||||
stream = client.beta.threads.runs.stream( # type: ignore[reportDeprecated]
|
||||
assistant_id=assistant_id, thread_id=final_thread_id, **run_options
|
||||
)
|
||||
|
||||
return stream, final_thread_id
|
||||
|
||||
async def _get_active_thread_run(self, thread_id: str | None) -> Run | None:
|
||||
"""Get any active run for the given thread."""
|
||||
client = await self._ensure_client()
|
||||
if thread_id is None:
|
||||
return None
|
||||
|
||||
async for run in client.beta.threads.runs.list(thread_id=thread_id, limit=1, order="desc"): # type: ignore[reportDeprecated]
|
||||
if run.status not in ["completed", "cancelled", "failed", "expired"]:
|
||||
return run
|
||||
return None
|
||||
|
||||
async def _prepare_thread(self, thread_id: str | None, thread_run: Run | None, run_options: dict[str, Any]) -> str:
|
||||
"""Prepare the thread for a new run, creating or cleaning up as needed."""
|
||||
client = await self._ensure_client()
|
||||
if thread_id is None:
|
||||
# No thread ID was provided, so create a new thread.
|
||||
thread = await client.beta.threads.create( # type: ignore[reportDeprecated]
|
||||
messages=run_options["additional_messages"],
|
||||
tool_resources=run_options.get("tool_resources"),
|
||||
metadata=run_options.get("metadata"),
|
||||
)
|
||||
run_options["additional_messages"] = []
|
||||
run_options.pop("tool_resources", None)
|
||||
return thread.id
|
||||
|
||||
if thread_run is not None:
|
||||
# There was an active run; we need to cancel it before starting a new run.
|
||||
await client.beta.threads.runs.cancel(run_id=thread_run.id, thread_id=thread_id) # type: ignore[reportDeprecated]
|
||||
|
||||
return thread_id
|
||||
|
||||
async def _process_stream_events(self, stream: Any, thread_id: str) -> AsyncIterable[ChatResponseUpdate]:
|
||||
response_id: str | None = None
|
||||
|
||||
async with stream as response_stream:
|
||||
async for response in response_stream:
|
||||
if response.event == "thread.run.created":
|
||||
yield ChatResponseUpdate(
|
||||
contents=[],
|
||||
conversation_id=thread_id,
|
||||
message_id=response_id,
|
||||
raw_representation=response.data,
|
||||
response_id=response_id,
|
||||
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 = "user" if delta.role == "user" else "assistant"
|
||||
|
||||
for delta_block in delta.content or []:
|
||||
if isinstance(delta_block, TextDeltaBlock) and delta_block.text and delta_block.text.value:
|
||||
text_content = Content.from_text(delta_block.text.value)
|
||||
if delta_block.text.annotations:
|
||||
annotations: list[Annotation] = []
|
||||
text_content.annotations = annotations
|
||||
for annotation in delta_block.text.annotations:
|
||||
if isinstance(annotation, FileCitationDeltaAnnotation):
|
||||
ann: Annotation = Annotation(
|
||||
type="citation",
|
||||
additional_properties={
|
||||
"text": annotation.text,
|
||||
"index": annotation.index,
|
||||
},
|
||||
raw_representation=annotation,
|
||||
)
|
||||
if annotation.file_citation and annotation.file_citation.file_id:
|
||||
ann["file_id"] = annotation.file_citation.file_id
|
||||
if annotation.start_index is not None and annotation.end_index is not None:
|
||||
ann["annotated_regions"] = [
|
||||
TextSpanRegion(
|
||||
type="text_span",
|
||||
start_index=annotation.start_index,
|
||||
end_index=annotation.end_index,
|
||||
)
|
||||
]
|
||||
annotations.append(ann)
|
||||
elif isinstance(annotation, FilePathDeltaAnnotation):
|
||||
ann = Annotation(
|
||||
type="citation",
|
||||
additional_properties={
|
||||
"text": annotation.text,
|
||||
"index": annotation.index,
|
||||
},
|
||||
raw_representation=annotation,
|
||||
)
|
||||
if annotation.file_path and annotation.file_path.file_id:
|
||||
ann["file_id"] = annotation.file_path.file_id
|
||||
if annotation.start_index is not None and annotation.end_index is not None:
|
||||
ann["annotated_regions"] = [
|
||||
TextSpanRegion(
|
||||
type="text_span",
|
||||
start_index=annotation.start_index,
|
||||
end_index=annotation.end_index,
|
||||
)
|
||||
]
|
||||
annotations.append(ann)
|
||||
yield ChatResponseUpdate(
|
||||
role=role, # type: ignore[arg-type]
|
||||
contents=[text_content],
|
||||
conversation_id=thread_id,
|
||||
message_id=response_id,
|
||||
raw_representation=response.data,
|
||||
response_id=response_id,
|
||||
)
|
||||
elif response.event == "thread.message.completed" and isinstance(response.data, ThreadMessage):
|
||||
# Process completed message to extract fully resolved annotations.
|
||||
# Delta events may carry partial/empty annotation data; the completed
|
||||
# message contains the final text with all citation details populated.
|
||||
completed_contents: list[Content] = []
|
||||
for block in response.data.content:
|
||||
if block.type != "text":
|
||||
continue
|
||||
text_content = Content.from_text(block.text.value)
|
||||
if block.text.annotations:
|
||||
completed_annotations: list[Annotation] = []
|
||||
text_content.annotations = completed_annotations
|
||||
for completed_annotation in block.text.annotations:
|
||||
if isinstance(completed_annotation, FileCitationAnnotation):
|
||||
props: dict[str, Any] = {
|
||||
"text": completed_annotation.text,
|
||||
}
|
||||
ann = Annotation(
|
||||
type="citation",
|
||||
additional_properties=props,
|
||||
raw_representation=completed_annotation,
|
||||
)
|
||||
if (
|
||||
completed_annotation.file_citation
|
||||
and completed_annotation.file_citation.file_id
|
||||
):
|
||||
ann["file_id"] = completed_annotation.file_citation.file_id
|
||||
ann["annotated_regions"] = [
|
||||
TextSpanRegion(
|
||||
type="text_span",
|
||||
start_index=completed_annotation.start_index,
|
||||
end_index=completed_annotation.end_index,
|
||||
)
|
||||
]
|
||||
text_content.annotations.append(ann)
|
||||
elif isinstance(completed_annotation, FilePathAnnotation):
|
||||
ann = Annotation(
|
||||
type="citation",
|
||||
additional_properties={
|
||||
"text": completed_annotation.text,
|
||||
},
|
||||
raw_representation=completed_annotation,
|
||||
)
|
||||
if completed_annotation.file_path and completed_annotation.file_path.file_id:
|
||||
ann["file_id"] = completed_annotation.file_path.file_id
|
||||
ann["annotated_regions"] = [
|
||||
TextSpanRegion(
|
||||
type="text_span",
|
||||
start_index=completed_annotation.start_index,
|
||||
end_index=completed_annotation.end_index,
|
||||
)
|
||||
]
|
||||
text_content.annotations.append(ann)
|
||||
else:
|
||||
logger.debug("Unparsed annotation type: %s", completed_annotation.type)
|
||||
completed_contents.append(text_content)
|
||||
if completed_contents:
|
||||
yield ChatResponseUpdate(
|
||||
role="assistant",
|
||||
contents=completed_contents,
|
||||
conversation_id=thread_id,
|
||||
message_id=response_id,
|
||||
raw_representation=response.data,
|
||||
response_id=response_id,
|
||||
)
|
||||
elif response.event == "thread.run.requires_action" and isinstance(response.data, Run):
|
||||
contents = self._parse_function_calls_from_assistants(response.data, response_id)
|
||||
if contents:
|
||||
yield ChatResponseUpdate(
|
||||
role="assistant",
|
||||
contents=contents,
|
||||
conversation_id=thread_id,
|
||||
message_id=response_id,
|
||||
raw_representation=response.data,
|
||||
response_id=response_id,
|
||||
)
|
||||
elif (
|
||||
response.event == "thread.run.completed"
|
||||
and isinstance(response.data, Run)
|
||||
and response.data.usage is not None
|
||||
):
|
||||
usage = response.data.usage
|
||||
usage_content = Content.from_usage(
|
||||
UsageDetails(
|
||||
input_token_count=usage.prompt_tokens,
|
||||
output_token_count=usage.completion_tokens,
|
||||
total_token_count=usage.total_tokens,
|
||||
)
|
||||
)
|
||||
yield ChatResponseUpdate(
|
||||
role="assistant",
|
||||
contents=[usage_content],
|
||||
conversation_id=thread_id,
|
||||
message_id=response_id,
|
||||
raw_representation=response.data,
|
||||
response_id=response_id,
|
||||
)
|
||||
else:
|
||||
yield ChatResponseUpdate(
|
||||
contents=[],
|
||||
conversation_id=thread_id,
|
||||
message_id=response_id,
|
||||
raw_representation=response.data,
|
||||
response_id=response_id,
|
||||
role="assistant",
|
||||
)
|
||||
|
||||
def _parse_function_calls_from_assistants(self, event_data: Run, response_id: str | None) -> list[Content]:
|
||||
"""Parse function call contents from an assistants tool action event."""
|
||||
contents: list[Content] = []
|
||||
|
||||
if event_data.required_action is not None:
|
||||
for tool_call in event_data.required_action.submit_tool_outputs.tool_calls:
|
||||
tool_call_any = cast(Any, tool_call)
|
||||
call_id = json.dumps([response_id, tool_call.id])
|
||||
tool_type = getattr(tool_call, "type", None)
|
||||
if tool_type == "code_interpreter" and getattr(tool_call_any, "code_interpreter", None):
|
||||
code_input = getattr(tool_call_any.code_interpreter, "input", None)
|
||||
inputs = (
|
||||
[Content.from_text(text=code_input, raw_representation=tool_call)]
|
||||
if code_input is not None
|
||||
else None
|
||||
)
|
||||
contents.append(
|
||||
Content.from_code_interpreter_tool_call(
|
||||
call_id=call_id,
|
||||
inputs=inputs,
|
||||
raw_representation=tool_call,
|
||||
)
|
||||
)
|
||||
elif tool_type == "mcp":
|
||||
contents.append(
|
||||
Content.from_mcp_server_tool_call(
|
||||
call_id=call_id,
|
||||
tool_name=getattr(tool_call, "name", "") or "",
|
||||
server_name=getattr(tool_call, "server_label", None),
|
||||
arguments=getattr(tool_call, "args", None),
|
||||
raw_representation=tool_call,
|
||||
)
|
||||
)
|
||||
else:
|
||||
function_name = tool_call.function.name
|
||||
function_arguments = json.loads(tool_call.function.arguments)
|
||||
contents.append(
|
||||
Content.from_function_call(
|
||||
call_id=call_id,
|
||||
name=function_name,
|
||||
arguments=function_arguments,
|
||||
)
|
||||
)
|
||||
|
||||
return contents
|
||||
|
||||
def _prepare_options(
|
||||
self,
|
||||
messages: Sequence[Message],
|
||||
options: Mapping[str, Any],
|
||||
**kwargs: Any,
|
||||
) -> tuple[dict[str, Any], list[Content] | None]:
|
||||
from agent_framework._types import validate_tool_mode
|
||||
|
||||
run_options: dict[str, Any] = {**kwargs}
|
||||
|
||||
# Extract options from the dict
|
||||
max_tokens = options.get("max_tokens")
|
||||
model = options.get("model") or options.get("model_id") # backward compat
|
||||
top_p = options.get("top_p")
|
||||
temperature = options.get("temperature")
|
||||
allow_multiple_tool_calls = options.get("allow_multiple_tool_calls")
|
||||
tool_choice = options.get("tool_choice")
|
||||
tools = options.get("tools")
|
||||
response_format = options.get("response_format")
|
||||
tool_resources = options.get("tool_resources")
|
||||
|
||||
if max_tokens is not None:
|
||||
run_options["max_completion_tokens"] = max_tokens
|
||||
if model is not None:
|
||||
run_options["model"] = model
|
||||
if top_p is not None:
|
||||
run_options["top_p"] = top_p
|
||||
if temperature is not None:
|
||||
run_options["temperature"] = temperature
|
||||
|
||||
if allow_multiple_tool_calls is not None:
|
||||
run_options["parallel_tool_calls"] = allow_multiple_tool_calls
|
||||
|
||||
if tool_resources is not None:
|
||||
run_options["tool_resources"] = tool_resources
|
||||
|
||||
tool_mode = validate_tool_mode(tool_choice)
|
||||
tool_definitions: list[MutableMapping[str, Any]] = []
|
||||
# Always include tools if provided, regardless of tool_choice
|
||||
# tool_choice="none" means the model won't call tools, but tools should still be available
|
||||
for tool in normalize_tools(tools):
|
||||
if isinstance(tool, FunctionTool):
|
||||
tool_definitions.append(tool.to_json_schema_spec()) # type: ignore[reportUnknownArgumentType]
|
||||
elif isinstance(tool, MutableMapping):
|
||||
# Pass through dict-based tools directly (from static factory methods)
|
||||
tool_definitions.append(cast(MutableMapping[str, Any], tool))
|
||||
|
||||
if len(tool_definitions) > 0:
|
||||
run_options["tools"] = tool_definitions
|
||||
|
||||
if tool_mode is not None:
|
||||
mode = tool_mode.get("mode")
|
||||
if mode is None:
|
||||
raise ValueError("tool_choice mode is required")
|
||||
if mode == "required" and (func_name := tool_mode.get("required_function_name")) is not None:
|
||||
run_options["tool_choice"] = {
|
||||
"type": "function",
|
||||
"function": {"name": func_name},
|
||||
}
|
||||
else:
|
||||
run_options["tool_choice"] = mode
|
||||
|
||||
if response_format is not None:
|
||||
if isinstance(response_format, dict):
|
||||
run_options["response_format"] = response_format
|
||||
else:
|
||||
run_options["response_format"] = {
|
||||
"type": "json_schema",
|
||||
"json_schema": {
|
||||
"name": response_format.__name__,
|
||||
"schema": response_format.model_json_schema(),
|
||||
"strict": True,
|
||||
},
|
||||
}
|
||||
|
||||
instructions: list[str] = []
|
||||
tool_results: list[Content] | None = None
|
||||
|
||||
additional_messages: list[AdditionalMessage] | None = None
|
||||
|
||||
# System/developer messages are turned into instructions,
|
||||
# since there is no such message roles in OpenAI Assistants.
|
||||
# All other messages are added 1:1.
|
||||
for chat_message in messages:
|
||||
if chat_message.role in ["system", "developer"]:
|
||||
for text_content in [content for content in chat_message.contents if content.type == "text"]:
|
||||
text = getattr(text_content, "text", None)
|
||||
if text:
|
||||
instructions.append(text)
|
||||
|
||||
continue
|
||||
|
||||
message_contents: list[MessageContentPartParam] = []
|
||||
|
||||
for content in chat_message.contents:
|
||||
if content.type == "text":
|
||||
message_contents.append(TextContentBlockParam(type="text", text=content.text)) # type: ignore[attr-defined, typeddict-item]
|
||||
elif content.type == "uri" and content.has_top_level_media_type("image"):
|
||||
message_contents.append(
|
||||
ImageURLContentBlockParam(type="image_url", image_url=ImageURLParam(url=content.uri)) # type: ignore[attr-defined, typeddict-item]
|
||||
)
|
||||
elif content.type == "function_result":
|
||||
if tool_results is None:
|
||||
tool_results = []
|
||||
tool_results.append(content)
|
||||
|
||||
if len(message_contents) > 0:
|
||||
if additional_messages is None:
|
||||
additional_messages = []
|
||||
additional_messages.append(
|
||||
AdditionalMessage(
|
||||
role="assistant" if chat_message.role == "assistant" else "user",
|
||||
content=message_contents,
|
||||
)
|
||||
)
|
||||
|
||||
if additional_messages is not None:
|
||||
run_options["additional_messages"] = additional_messages
|
||||
|
||||
if len(instructions) > 0:
|
||||
run_options["instructions"] = "".join(instructions)
|
||||
|
||||
return run_options, tool_results
|
||||
|
||||
def _prepare_tool_outputs_for_assistants(
|
||||
self,
|
||||
tool_results: list[Content] | None,
|
||||
) -> tuple[str | None, list[ToolOutput] | None]:
|
||||
"""Prepare function results for submission to the assistants API."""
|
||||
run_id: str | None = None
|
||||
tool_outputs: list[ToolOutput] | None = None
|
||||
|
||||
if tool_results:
|
||||
for function_result_content in tool_results:
|
||||
# When creating the FunctionCallContent, we created it with a CallId == [runId, callId].
|
||||
# We need to extract the run ID and ensure that the ToolOutput we send back to Azure
|
||||
# is only the call ID.
|
||||
run_and_call_ids: list[str] = json.loads(function_result_content.call_id) # type: ignore[arg-type]
|
||||
|
||||
if (
|
||||
not run_and_call_ids
|
||||
or len(run_and_call_ids) != 2
|
||||
or not run_and_call_ids[0]
|
||||
or not run_and_call_ids[1]
|
||||
or (run_id is not None and run_id != run_and_call_ids[0])
|
||||
):
|
||||
continue
|
||||
|
||||
run_id = run_and_call_ids[0]
|
||||
call_id = run_and_call_ids[1]
|
||||
|
||||
if tool_outputs is None:
|
||||
tool_outputs = []
|
||||
output = (
|
||||
function_result_content.result
|
||||
if function_result_content.result is not None
|
||||
else "No output received."
|
||||
)
|
||||
tool_outputs.append(ToolOutput(tool_call_id=call_id, output=output))
|
||||
|
||||
return run_id, tool_outputs
|
||||
|
||||
def _update_agent_name_and_description(self, agent_name: str | None, description: str | None = None) -> None:
|
||||
"""Update the agent name in the chat client.
|
||||
|
||||
Args:
|
||||
agent_name: The new name for the agent.
|
||||
description: The new description for the agent.
|
||||
"""
|
||||
# This is a no-op in the base class, but can be overridden by subclasses
|
||||
# to update the agent name in the client.
|
||||
if agent_name and not self.assistant_name:
|
||||
self.assistant_name = agent_name
|
||||
if description and not self.assistant_description:
|
||||
self.assistant_description = description
|
||||
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,303 @@
|
||||
# Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import base64
|
||||
import struct
|
||||
import sys
|
||||
from collections.abc import Awaitable, Callable, Mapping, Sequence
|
||||
from copy import copy
|
||||
from typing import Any, ClassVar, Generic, Literal, TypedDict
|
||||
|
||||
from agent_framework._clients import BaseEmbeddingClient
|
||||
from agent_framework._settings import SecretString, load_settings
|
||||
from agent_framework._telemetry import APP_INFO, USER_AGENT_KEY, prepend_agent_framework_to_user_agent
|
||||
from agent_framework._types import Embedding, EmbeddingGenerationOptions, GeneratedEmbeddings, UsageDetails
|
||||
from agent_framework.observability import EmbeddingTelemetryLayer
|
||||
from openai import AsyncOpenAI
|
||||
|
||||
from ._shared import OpenAISettings, get_api_key
|
||||
|
||||
if sys.version_info >= (3, 13):
|
||||
from typing import TypeVar # type: ignore # pragma: no cover
|
||||
else:
|
||||
from typing_extensions import TypeVar # type: ignore # pragma: no cover
|
||||
|
||||
|
||||
class OpenAIEmbeddingOptions(EmbeddingGenerationOptions, total=False):
|
||||
"""OpenAI-specific embedding options.
|
||||
|
||||
Extends EmbeddingGenerationOptions with OpenAI-specific fields.
|
||||
|
||||
Examples:
|
||||
.. code-block:: python
|
||||
|
||||
from agent_framework.openai import OpenAIEmbeddingOptions
|
||||
|
||||
options: OpenAIEmbeddingOptions = {
|
||||
"model": "text-embedding-3-small",
|
||||
"dimensions": 1536,
|
||||
"encoding_format": "float",
|
||||
}
|
||||
"""
|
||||
|
||||
encoding_format: Literal["float", "base64"]
|
||||
user: str
|
||||
|
||||
|
||||
OpenAIEmbeddingOptionsT = TypeVar(
|
||||
"OpenAIEmbeddingOptionsT",
|
||||
bound=TypedDict, # type: ignore[valid-type]
|
||||
default="OpenAIEmbeddingOptions",
|
||||
covariant=True,
|
||||
)
|
||||
|
||||
|
||||
class RawOpenAIEmbeddingClient(
|
||||
BaseEmbeddingClient[str, list[float], OpenAIEmbeddingOptionsT],
|
||||
Generic[OpenAIEmbeddingOptionsT],
|
||||
):
|
||||
"""Raw OpenAI embedding client without telemetry."""
|
||||
|
||||
INJECTABLE: ClassVar[set[str]] = {"client"}
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
model: str | None = None,
|
||||
model_id: str | None = None,
|
||||
api_key: str | SecretString | Callable[[], str | Awaitable[str]] | None = None,
|
||||
org_id: str | None = None,
|
||||
base_url: str | None = None,
|
||||
default_headers: Mapping[str, str] | None = None,
|
||||
async_client: AsyncOpenAI | None = None,
|
||||
env_file_path: str | None = None,
|
||||
env_file_encoding: str | None = None,
|
||||
**kwargs: Any,
|
||||
) -> None:
|
||||
"""Initialize a raw OpenAI embedding client.
|
||||
|
||||
Keyword Args:
|
||||
model: OpenAI embedding model name.
|
||||
model_id: Deprecated alias for ``model``.
|
||||
api_key: OpenAI API key, SecretString, or callable returning a key.
|
||||
org_id: OpenAI organization ID.
|
||||
base_url: Custom API base URL.
|
||||
default_headers: Additional HTTP headers.
|
||||
async_client: Pre-configured AsyncOpenAI client (skips client creation).
|
||||
env_file_path: Path to .env file for settings.
|
||||
env_file_encoding: Encoding for .env file.
|
||||
kwargs: Additional keyword arguments forwarded to ``BaseEmbeddingClient``.
|
||||
"""
|
||||
if model_id is not None and model is None:
|
||||
import warnings
|
||||
|
||||
warnings.warn("model_id is deprecated, use model instead", DeprecationWarning, stacklevel=2)
|
||||
model = model_id
|
||||
|
||||
if not async_client:
|
||||
openai_settings = load_settings(
|
||||
OpenAISettings,
|
||||
env_prefix="OPENAI_",
|
||||
api_key=api_key,
|
||||
org_id=org_id,
|
||||
base_url=base_url,
|
||||
embedding_model=model,
|
||||
env_file_path=env_file_path,
|
||||
env_file_encoding=env_file_encoding,
|
||||
)
|
||||
|
||||
api_key_value = openai_settings.get("api_key")
|
||||
resolved_model = openai_settings.get("embedding_model") or model
|
||||
|
||||
# Only create a client when we have enough configuration.
|
||||
# Subclasses that manage their own client pass no args here
|
||||
if api_key_value:
|
||||
if not resolved_model:
|
||||
raise ValueError(
|
||||
"OpenAI embedding model is required. "
|
||||
"Set via 'model' parameter or 'OPENAI_EMBEDDING_MODEL' environment variable."
|
||||
)
|
||||
model = resolved_model
|
||||
|
||||
resolved_api_key = get_api_key(api_key_value)
|
||||
|
||||
# Merge APP_INFO into the headers
|
||||
merged_headers = dict(copy(default_headers)) if default_headers else {}
|
||||
if APP_INFO:
|
||||
merged_headers.update(APP_INFO)
|
||||
merged_headers = prepend_agent_framework_to_user_agent(merged_headers)
|
||||
|
||||
client_args: dict[str, Any] = {"api_key": resolved_api_key, "default_headers": merged_headers}
|
||||
if resolved_org_id := openai_settings.get("org_id"):
|
||||
client_args["organization"] = resolved_org_id
|
||||
if resolved_base_url := openai_settings.get("base_url"):
|
||||
client_args["base_url"] = resolved_base_url
|
||||
|
||||
async_client = AsyncOpenAI(**client_args)
|
||||
|
||||
self.client = async_client
|
||||
self.model: str | None = model.strip() if model else None
|
||||
|
||||
# Store configuration for serialization
|
||||
self.org_id = org_id
|
||||
self.base_url = str(base_url) if base_url else None
|
||||
if default_headers:
|
||||
self.default_headers: dict[str, Any] | None = {
|
||||
k: v for k, v in default_headers.items() if k != USER_AGENT_KEY
|
||||
}
|
||||
else:
|
||||
self.default_headers = None
|
||||
|
||||
super().__init__(**kwargs)
|
||||
|
||||
def service_url(self) -> str:
|
||||
"""Get the URL of the service."""
|
||||
return str(self.client.base_url) if self.client else "Unknown"
|
||||
|
||||
async def get_embeddings(
|
||||
self,
|
||||
values: Sequence[str],
|
||||
*,
|
||||
options: OpenAIEmbeddingOptionsT | None = None,
|
||||
) -> GeneratedEmbeddings[list[float], OpenAIEmbeddingOptionsT]:
|
||||
"""Call the OpenAI embeddings API.
|
||||
|
||||
Args:
|
||||
values: The text values to generate embeddings for.
|
||||
options: Optional embedding generation options.
|
||||
|
||||
Returns:
|
||||
Generated embeddings with usage metadata.
|
||||
|
||||
Raises:
|
||||
ValueError: If model is not provided or values is empty.
|
||||
"""
|
||||
if not values:
|
||||
return GeneratedEmbeddings([], options=options) # type: ignore
|
||||
|
||||
opts: dict[str, Any] = options or {} # type: ignore
|
||||
# backward compat: accept model_id in options
|
||||
model = opts.get("model") or opts.get("model_id") or self.model
|
||||
if not model:
|
||||
raise ValueError("model is required")
|
||||
|
||||
kwargs: dict[str, Any] = {"input": list(values), "model": model}
|
||||
if dimensions := opts.get("dimensions"):
|
||||
kwargs["dimensions"] = dimensions
|
||||
if encoding_format := opts.get("encoding_format"):
|
||||
kwargs["encoding_format"] = encoding_format
|
||||
if user := opts.get("user"):
|
||||
kwargs["user"] = user
|
||||
|
||||
response = await self.client.embeddings.create(**kwargs) # type: ignore[union-attr]
|
||||
|
||||
encoding = kwargs.get("encoding_format", "float")
|
||||
embeddings: list[Embedding[list[float]]] = []
|
||||
for item in response.data:
|
||||
vector: list[float]
|
||||
if encoding == "base64" and isinstance(item.embedding, str):
|
||||
# Decode base64-encoded floats (little-endian IEEE 754)
|
||||
raw = base64.b64decode(item.embedding)
|
||||
vector = list(struct.unpack(f"<{len(raw) // 4}f", raw))
|
||||
else:
|
||||
vector = item.embedding # type: ignore[assignment]
|
||||
embeddings.append(
|
||||
Embedding(
|
||||
vector=vector,
|
||||
dimensions=len(vector),
|
||||
model=response.model,
|
||||
)
|
||||
)
|
||||
|
||||
usage_dict: UsageDetails | None = None
|
||||
if response.usage:
|
||||
usage_dict = {
|
||||
"input_token_count": response.usage.prompt_tokens,
|
||||
"total_token_count": response.usage.total_tokens,
|
||||
}
|
||||
|
||||
return GeneratedEmbeddings(embeddings, options=options, usage=usage_dict)
|
||||
|
||||
|
||||
class OpenAIEmbeddingClient(
|
||||
EmbeddingTelemetryLayer[str, list[float], OpenAIEmbeddingOptionsT],
|
||||
RawOpenAIEmbeddingClient[OpenAIEmbeddingOptionsT],
|
||||
Generic[OpenAIEmbeddingOptionsT],
|
||||
):
|
||||
"""OpenAI embedding client with telemetry support.
|
||||
|
||||
Keyword Args:
|
||||
model: The embedding model (e.g. "text-embedding-3-small").
|
||||
Can also be set via environment variable OPENAI_EMBEDDING_MODEL.
|
||||
model_id: Deprecated alias for ``model``.
|
||||
api_key: OpenAI API key.
|
||||
Can also be set via environment variable OPENAI_API_KEY.
|
||||
org_id: OpenAI organization ID.
|
||||
default_headers: Additional HTTP headers.
|
||||
async_client: Pre-configured AsyncOpenAI client.
|
||||
base_url: Custom API base URL.
|
||||
otel_provider_name: Override the OpenTelemetry provider name for telemetry.
|
||||
env_file_path: Path to .env file for settings.
|
||||
env_file_encoding: Encoding for .env file.
|
||||
|
||||
Examples:
|
||||
.. code-block:: python
|
||||
|
||||
from agent_framework.openai import OpenAIEmbeddingClient
|
||||
|
||||
# Using environment variables
|
||||
# Set OPENAI_API_KEY=sk-...
|
||||
# Set OPENAI_EMBEDDING_MODEL=text-embedding-3-small
|
||||
client = OpenAIEmbeddingClient()
|
||||
|
||||
# Or passing parameters directly
|
||||
client = OpenAIEmbeddingClient(
|
||||
model="text-embedding-3-small",
|
||||
api_key="sk-...",
|
||||
)
|
||||
|
||||
# Generate embeddings
|
||||
result = await client.get_embeddings(["Hello, world!"])
|
||||
print(result[0].vector)
|
||||
"""
|
||||
|
||||
OTEL_PROVIDER_NAME: ClassVar[str] = "openai" # type: ignore[reportIncompatibleVariableOverride, misc]
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
model: str | None = None,
|
||||
api_key: str | Callable[[], str | Awaitable[str]] | None = None,
|
||||
org_id: str | None = None,
|
||||
default_headers: Mapping[str, str] | None = None,
|
||||
async_client: AsyncOpenAI | None = None,
|
||||
base_url: str | None = None,
|
||||
otel_provider_name: str | None = None,
|
||||
env_file_path: str | None = None,
|
||||
env_file_encoding: str | None = None,
|
||||
) -> None:
|
||||
"""Initialize an OpenAI embedding client."""
|
||||
super().__init__(
|
||||
model=model,
|
||||
api_key=api_key,
|
||||
org_id=org_id,
|
||||
base_url=base_url,
|
||||
default_headers=default_headers,
|
||||
async_client=async_client,
|
||||
env_file_path=env_file_path,
|
||||
env_file_encoding=env_file_encoding,
|
||||
)
|
||||
if otel_provider_name is not None:
|
||||
self.OTEL_PROVIDER_NAME = otel_provider_name # type: ignore[misc]
|
||||
|
||||
# Validate that the client was created successfully (from explicit args or env vars)
|
||||
if self.client is None:
|
||||
raise ValueError(
|
||||
"OpenAI API key is required. Set via 'api_key' parameter or 'OPENAI_API_KEY' environment variable."
|
||||
)
|
||||
if not self.model:
|
||||
raise ValueError(
|
||||
"OpenAI embedding model is required. "
|
||||
"Set via 'model' parameter or 'OPENAI_EMBEDDING_MODEL' environment variable."
|
||||
)
|
||||
@@ -0,0 +1,90 @@
|
||||
# Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from enum import Enum
|
||||
from typing import Any
|
||||
|
||||
from agent_framework.exceptions import ChatClientContentFilterException
|
||||
from openai import BadRequestError
|
||||
|
||||
|
||||
class ContentFilterResultSeverity(Enum):
|
||||
"""The severity of the content filter result."""
|
||||
|
||||
HIGH = "high"
|
||||
MEDIUM = "medium"
|
||||
SAFE = "safe"
|
||||
LOW = "low"
|
||||
|
||||
|
||||
@dataclass
|
||||
class ContentFilterResult:
|
||||
"""The result of a content filter check."""
|
||||
|
||||
filtered: bool = False
|
||||
detected: bool = False
|
||||
severity: ContentFilterResultSeverity = ContentFilterResultSeverity.SAFE
|
||||
|
||||
@classmethod
|
||||
def from_inner_error_result(cls, inner_error_results: dict[str, Any]) -> ContentFilterResult:
|
||||
"""Creates a ContentFilterResult from the inner error results.
|
||||
|
||||
Args:
|
||||
inner_error_results: The inner error results.
|
||||
|
||||
Returns:
|
||||
ContentFilterResult: The ContentFilterResult.
|
||||
"""
|
||||
return cls(
|
||||
filtered=inner_error_results.get("filtered", False),
|
||||
detected=inner_error_results.get("detected", False),
|
||||
severity=ContentFilterResultSeverity(
|
||||
inner_error_results.get("severity", ContentFilterResultSeverity.SAFE.value)
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
class ContentFilterCodes(Enum):
|
||||
"""Content filter codes."""
|
||||
|
||||
RESPONSIBLE_AI_POLICY_VIOLATION = "ResponsibleAIPolicyViolation"
|
||||
|
||||
|
||||
@dataclass
|
||||
class OpenAIContentFilterException(ChatClientContentFilterException):
|
||||
"""AI exception for an error from Azure OpenAI's content filter."""
|
||||
|
||||
# The parameter that caused the error.
|
||||
param: str | None
|
||||
|
||||
# The error code specific to the content filter.
|
||||
content_filter_code: ContentFilterCodes
|
||||
|
||||
# The results of the different content filter checks.
|
||||
content_filter_result: dict[str, ContentFilterResult]
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
message: str,
|
||||
inner_exception: BadRequestError,
|
||||
) -> None:
|
||||
"""Initializes a new instance of the ContentFilterAIException class.
|
||||
|
||||
Args:
|
||||
message: The error message.
|
||||
inner_exception: The inner exception.
|
||||
"""
|
||||
super().__init__(message)
|
||||
|
||||
self.param = inner_exception.param
|
||||
if inner_exception.body is not None and isinstance(inner_exception.body, dict):
|
||||
inner_error = inner_exception.body.get("innererror", {}) # type: ignore
|
||||
self.content_filter_code = ContentFilterCodes(
|
||||
inner_error.get("code", ContentFilterCodes.RESPONSIBLE_AI_POLICY_VIOLATION.value) # type: ignore
|
||||
)
|
||||
self.content_filter_result = {
|
||||
key: ContentFilterResult.from_inner_error_result(values) # type: ignore
|
||||
for key, values in inner_error.get("content_filter_result", {}).items() # type: ignore
|
||||
}
|
||||
@@ -0,0 +1,508 @@
|
||||
# Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import os
|
||||
import sys
|
||||
from collections.abc import Awaitable, Callable, Mapping, MutableMapping, Sequence
|
||||
from copy import copy
|
||||
from typing import Any, ClassVar, Union, cast
|
||||
|
||||
import openai
|
||||
from agent_framework._serialization import SerializationMixin
|
||||
from agent_framework._settings import SecretString, load_settings
|
||||
from agent_framework._telemetry import APP_INFO, USER_AGENT_KEY, prepend_agent_framework_to_user_agent
|
||||
from agent_framework._tools import FunctionTool
|
||||
from dotenv import dotenv_values
|
||||
from openai import AsyncOpenAI, AsyncStream, _legacy_response # type: ignore
|
||||
from openai.types import Completion
|
||||
from openai.types.audio import Transcription
|
||||
from openai.types.chat import ChatCompletion, ChatCompletionChunk
|
||||
from openai.types.images_response import ImagesResponse
|
||||
from openai.types.responses.response import Response
|
||||
from openai.types.responses.response_stream_event import ResponseStreamEvent
|
||||
from packaging.version import parse
|
||||
|
||||
logger: logging.Logger = logging.getLogger("agent_framework.openai")
|
||||
|
||||
DEFAULT_AZURE_OPENAI_CHAT_COMPLETION_API_VERSION = "2024-10-21"
|
||||
DEFAULT_AZURE_OPENAI_RESPONSES_API_VERSION = "preview"
|
||||
|
||||
|
||||
RESPONSE_TYPE = Union[
|
||||
ChatCompletion,
|
||||
Completion,
|
||||
AsyncStream[ChatCompletionChunk],
|
||||
AsyncStream[Completion],
|
||||
list[Any],
|
||||
ImagesResponse,
|
||||
Response,
|
||||
AsyncStream[ResponseStreamEvent],
|
||||
Transcription,
|
||||
_legacy_response.HttpxBinaryResponseContent,
|
||||
]
|
||||
|
||||
OPTION_TYPE = dict[str, Any]
|
||||
|
||||
if sys.version_info >= (3, 11):
|
||||
from typing import TypedDict # type: ignore # pragma: no cover
|
||||
else:
|
||||
from typing_extensions import TypedDict # type: ignore # pragma: no cover
|
||||
|
||||
|
||||
def _check_openai_version_for_callable_api_key() -> None:
|
||||
"""Check if OpenAI version supports callable API keys.
|
||||
|
||||
Callable API keys require OpenAI >= 1.106.0.
|
||||
If the version is too old, raise a ValueError with helpful message.
|
||||
"""
|
||||
try:
|
||||
current_version = parse(openai.__version__)
|
||||
min_required_version = parse("1.106.0")
|
||||
|
||||
if current_version < min_required_version:
|
||||
raise ValueError(
|
||||
f"Callable API keys require OpenAI SDK >= 1.106.0, but you have {openai.__version__}. "
|
||||
f"Please upgrade with 'pip install openai>=1.106.0' or provide a string API key instead. "
|
||||
f"Note: If you're using mem0ai, you may need to upgrade to mem0ai>=1.0.0 "
|
||||
f"to allow newer OpenAI versions."
|
||||
)
|
||||
except ValueError:
|
||||
raise # Re-raise our own exception
|
||||
except Exception as e:
|
||||
logger.warning(f"Could not check OpenAI version for callable API key support: {e}")
|
||||
|
||||
|
||||
class OpenAISettings(TypedDict, total=False):
|
||||
"""OpenAI environment settings.
|
||||
|
||||
Settings are resolved in this order: explicit keyword arguments, values from an
|
||||
explicitly provided .env file, then environment variables with the prefix
|
||||
'OPENAI_'. If settings are missing after resolution, validation will fail.
|
||||
|
||||
Keyword Args:
|
||||
api_key: OpenAI API key, see https://platform.openai.com/account/api-keys.
|
||||
Can be set via environment variable OPENAI_API_KEY.
|
||||
base_url: The base URL for the OpenAI API.
|
||||
Can be set via environment variable OPENAI_BASE_URL.
|
||||
org_id: This is usually optional unless your account belongs to multiple organizations.
|
||||
Can be set via environment variable OPENAI_ORG_ID.
|
||||
model: The OpenAI model to use, for example, gpt-4o or o1.
|
||||
Can be set via environment variable OPENAI_MODEL.
|
||||
embedding_model: The OpenAI embedding model to use, for example, text-embedding-3-small.
|
||||
Can be set via environment variable OPENAI_EMBEDDING_MODEL.
|
||||
|
||||
Examples:
|
||||
.. code-block:: python
|
||||
|
||||
from agent_framework.openai import OpenAISettings
|
||||
|
||||
# Using environment variables
|
||||
# Set OPENAI_API_KEY=sk-...
|
||||
# Set OPENAI_MODEL=gpt-4o
|
||||
settings = load_settings(OpenAISettings, env_prefix="OPENAI_")
|
||||
|
||||
# Or passing parameters directly
|
||||
settings = load_settings(OpenAISettings, env_prefix="OPENAI_", api_key="sk-...", model="gpt-4o")
|
||||
|
||||
# Or loading from a .env file
|
||||
settings = load_settings(OpenAISettings, env_prefix="OPENAI_", env_file_path="path/to/.env")
|
||||
"""
|
||||
|
||||
api_key: SecretString | Callable[[], str | Awaitable[str]] | None
|
||||
base_url: str | None
|
||||
org_id: str | None
|
||||
model: str | None
|
||||
embedding_model: str | None
|
||||
azure_endpoint: str | None
|
||||
api_version: str | None
|
||||
|
||||
|
||||
def _load_dotenv_values(*, env_file_path: str | None, env_file_encoding: str | None) -> dict[str, str]:
|
||||
"""Load dotenv values for non-standard environment variable aliases."""
|
||||
if env_file_path is None or not os.path.exists(env_file_path):
|
||||
return {}
|
||||
|
||||
raw_dotenv_values = dotenv_values(dotenv_path=env_file_path, encoding=env_file_encoding or "utf-8")
|
||||
return {key: value for key, value in raw_dotenv_values.items() if value is not None}
|
||||
|
||||
|
||||
def _get_setting_from_alias(
|
||||
name: str,
|
||||
*,
|
||||
dotenv_values_by_name: Mapping[str, str],
|
||||
) -> str | None:
|
||||
"""Resolve a setting from an explicit env-var alias."""
|
||||
if dotenv_value := dotenv_values_by_name.get(name):
|
||||
return dotenv_value
|
||||
return os.getenv(name)
|
||||
|
||||
|
||||
def load_openai_service_settings(
|
||||
*,
|
||||
model: str | None,
|
||||
api_key: str | SecretString | Callable[[], str | Awaitable[str]] | None,
|
||||
org_id: str | None,
|
||||
base_url: str | None,
|
||||
azure_endpoint: str | None,
|
||||
api_version: str | None,
|
||||
env_file_path: str | None,
|
||||
env_file_encoding: str | None,
|
||||
azure_model_env_vars: Sequence[str],
|
||||
default_azure_api_version: str,
|
||||
) -> tuple[OpenAISettings, bool]:
|
||||
"""Load OpenAI settings, including Azure OpenAI aliases.
|
||||
|
||||
The generic OpenAI clients primarily read from ``OPENAI_*`` variables. When an
|
||||
``AZURE_OPENAI_ENDPOINT`` (or ``AZURE_OPENAI_BASE_URL``) is available and no
|
||||
explicit OpenAI base URL is configured, this helper switches to Azure-specific
|
||||
environment variables for endpoint, API key, model deployment, and API version.
|
||||
"""
|
||||
openai_settings = load_settings(
|
||||
OpenAISettings,
|
||||
env_prefix="OPENAI_",
|
||||
api_key=api_key,
|
||||
org_id=org_id,
|
||||
base_url=base_url,
|
||||
model=model,
|
||||
azure_endpoint=azure_endpoint,
|
||||
api_version=api_version,
|
||||
env_file_path=env_file_path,
|
||||
env_file_encoding=env_file_encoding,
|
||||
)
|
||||
|
||||
dotenv_values_by_name = _load_dotenv_values(
|
||||
env_file_path=env_file_path,
|
||||
env_file_encoding=env_file_encoding,
|
||||
)
|
||||
|
||||
resolved_azure_endpoint = azure_endpoint
|
||||
resolved_azure_base_url: str | None = None
|
||||
if not openai_settings.get("base_url"):
|
||||
if resolved_azure_endpoint is None:
|
||||
resolved_azure_endpoint = _get_setting_from_alias(
|
||||
"AZURE_OPENAI_ENDPOINT",
|
||||
dotenv_values_by_name=dotenv_values_by_name,
|
||||
)
|
||||
if resolved_azure_endpoint is None:
|
||||
resolved_azure_base_url = _get_setting_from_alias(
|
||||
"AZURE_OPENAI_BASE_URL",
|
||||
dotenv_values_by_name=dotenv_values_by_name,
|
||||
)
|
||||
if resolved_azure_base_url is not None:
|
||||
openai_settings["base_url"] = resolved_azure_base_url
|
||||
|
||||
use_azure_client = resolved_azure_endpoint is not None or resolved_azure_base_url is not None
|
||||
if resolved_azure_endpoint is not None:
|
||||
openai_settings["azure_endpoint"] = resolved_azure_endpoint
|
||||
|
||||
if use_azure_client:
|
||||
if api_key is None:
|
||||
resolved_azure_api_key = _get_setting_from_alias(
|
||||
"AZURE_OPENAI_API_KEY",
|
||||
dotenv_values_by_name=dotenv_values_by_name,
|
||||
)
|
||||
if resolved_azure_api_key is not None:
|
||||
openai_settings["api_key"] = SecretString(resolved_azure_api_key)
|
||||
|
||||
if model is None:
|
||||
for env_var_name in azure_model_env_vars:
|
||||
resolved_model = _get_setting_from_alias(
|
||||
env_var_name,
|
||||
dotenv_values_by_name=dotenv_values_by_name,
|
||||
)
|
||||
if resolved_model is not None:
|
||||
openai_settings["model"] = resolved_model
|
||||
break
|
||||
|
||||
if not openai_settings.get("api_version"):
|
||||
resolved_api_version = _get_setting_from_alias(
|
||||
"AZURE_OPENAI_API_VERSION",
|
||||
dotenv_values_by_name=dotenv_values_by_name,
|
||||
)
|
||||
openai_settings["api_version"] = resolved_api_version or default_azure_api_version
|
||||
|
||||
return openai_settings, use_azure_client
|
||||
|
||||
|
||||
def maybe_append_azure_endpoint_guidance(message: str, *, azure_endpoint: str | None) -> str:
|
||||
"""Append Azure endpoint guidance only when the configured endpoint shape looks suspicious."""
|
||||
if not azure_endpoint or not azure_endpoint.rstrip("/").endswith("/openai/v1"):
|
||||
return message
|
||||
|
||||
return (
|
||||
f"{message} If you are using Azure OpenAI key auth, pass the resource endpoint without "
|
||||
"'/openai/v1' to 'azure_endpoint', or pass the full '/openai/v1' URL via 'base_url' instead."
|
||||
)
|
||||
|
||||
|
||||
def get_api_key(
|
||||
api_key: str | SecretString | Callable[[], str | Awaitable[str]] | None,
|
||||
) -> str | Callable[[], str | Awaitable[str]] | None:
|
||||
"""Get the appropriate API key value for client initialization.
|
||||
|
||||
Args:
|
||||
api_key: The API key parameter which can be a string, SecretString, callable, or None.
|
||||
|
||||
Returns:
|
||||
For callable API keys: returns the callable directly.
|
||||
For SecretString: returns the unwrapped secret value.
|
||||
For string/None API keys: returns as-is.
|
||||
"""
|
||||
if isinstance(api_key, SecretString):
|
||||
return api_key.get_secret_value()
|
||||
|
||||
# Check version compatibility for callable API keys
|
||||
if callable(api_key):
|
||||
_check_openai_version_for_callable_api_key()
|
||||
|
||||
return api_key # Pass callable, string, or None directly to OpenAI SDK
|
||||
|
||||
|
||||
class OpenAIBase(SerializationMixin):
|
||||
"""Base class for OpenAI Clients.
|
||||
|
||||
.. deprecated::
|
||||
``OpenAIBase`` is deprecated and only used by ``OpenAIAssistantsClient``
|
||||
and ``AzureOpenAIAssistantsClient``. New clients should manage ``client``
|
||||
and ``model`` directly in their own ``__init__``.
|
||||
"""
|
||||
|
||||
INJECTABLE: ClassVar[set[str]] = {"client"}
|
||||
|
||||
def __init__(
|
||||
self, *, model: str | None = None, model_id: str | None = None, client: AsyncOpenAI | None = None, **kwargs: Any
|
||||
) -> None:
|
||||
"""Initialize OpenAIBase.
|
||||
|
||||
Keyword Args:
|
||||
client: The AsyncOpenAI client instance.
|
||||
model: The AI model to use.
|
||||
model_id: Deprecated alias for ``model``.
|
||||
**kwargs: Additional keyword arguments.
|
||||
"""
|
||||
if model_id is not None and model is None:
|
||||
import warnings
|
||||
|
||||
warnings.warn("model_id is deprecated, use model instead", DeprecationWarning, stacklevel=2)
|
||||
model = model_id
|
||||
self.client = client
|
||||
self.model: str | None = None
|
||||
if model:
|
||||
self.model = model.strip()
|
||||
|
||||
# Call super().__init__() to continue MRO chain (e.g., RawChatClient)
|
||||
# Extract known kwargs that belong to other base classes
|
||||
additional_properties = kwargs.pop("additional_properties", None)
|
||||
middleware = kwargs.pop("middleware", None)
|
||||
instruction_role = kwargs.pop("instruction_role", None)
|
||||
function_invocation_configuration = kwargs.pop("function_invocation_configuration", None)
|
||||
|
||||
# Build super().__init__() args
|
||||
super_kwargs = {}
|
||||
if additional_properties is not None:
|
||||
super_kwargs["additional_properties"] = additional_properties
|
||||
if middleware is not None:
|
||||
super_kwargs["middleware"] = middleware
|
||||
if function_invocation_configuration is not None:
|
||||
super_kwargs["function_invocation_configuration"] = function_invocation_configuration
|
||||
|
||||
# Call super().__init__() with filtered kwargs
|
||||
super().__init__(**super_kwargs)
|
||||
|
||||
# Store instruction_role and any remaining kwargs as instance attributes
|
||||
if instruction_role is not None:
|
||||
self.instruction_role = instruction_role
|
||||
for key, value in kwargs.items():
|
||||
setattr(self, key, value)
|
||||
|
||||
async def _initialize_client(self) -> None:
|
||||
"""Initialize OpenAI client asynchronously.
|
||||
|
||||
Override in subclasses to initialize the OpenAI client asynchronously.
|
||||
"""
|
||||
pass
|
||||
|
||||
async def _ensure_client(self) -> AsyncOpenAI:
|
||||
"""Ensure OpenAI client is initialized."""
|
||||
await self._initialize_client()
|
||||
if self.client is None:
|
||||
raise RuntimeError("OpenAI client is not initialized")
|
||||
|
||||
return self.client
|
||||
|
||||
def _get_api_key(
|
||||
self, api_key: str | SecretString | Callable[[], str | Awaitable[str]] | None
|
||||
) -> str | Callable[[], str | Awaitable[str]] | None:
|
||||
"""Get the appropriate API key value for client initialization.
|
||||
|
||||
Args:
|
||||
api_key: The API key parameter which can be a string, SecretString, callable, or None.
|
||||
|
||||
Returns:
|
||||
For callable API keys: returns the callable directly.
|
||||
For SecretString/string/None API keys: returns as-is (SecretString is a str subclass).
|
||||
"""
|
||||
if isinstance(api_key, SecretString):
|
||||
return api_key.get_secret_value()
|
||||
|
||||
# Check version compatibility for callable API keys
|
||||
if callable(api_key):
|
||||
_check_openai_version_for_callable_api_key()
|
||||
|
||||
return api_key # Pass callable, string, or None directly to OpenAI SDK
|
||||
|
||||
|
||||
class OpenAIConfigMixin(OpenAIBase):
|
||||
"""Internal class for configuring a connection to an OpenAI service.
|
||||
|
||||
.. deprecated::
|
||||
``OpenAIConfigMixin`` is deprecated and only used by ``OpenAIAssistantsClient``
|
||||
and ``AzureOpenAIAssistantsClient``. New clients handle configuration
|
||||
directly in their own ``__init__``.
|
||||
"""
|
||||
|
||||
OTEL_PROVIDER_NAME: ClassVar[str] = "openai" # type: ignore[reportIncompatibleVariableOverride, misc]
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
model: str,
|
||||
api_key: str | Callable[[], str | Awaitable[str]] | None = None,
|
||||
org_id: str | None = None,
|
||||
default_headers: Mapping[str, str] | None = None,
|
||||
client: AsyncOpenAI | None = None,
|
||||
instruction_role: str | None = None,
|
||||
base_url: str | None = None,
|
||||
**kwargs: Any,
|
||||
) -> None:
|
||||
"""Initialize a client for OpenAI services.
|
||||
|
||||
This constructor sets up a client to interact with OpenAI's API, allowing for
|
||||
different types of AI model interactions, like chat or text completion.
|
||||
|
||||
Args:
|
||||
model: OpenAI model identifier. Must be non-empty.
|
||||
Default to a preset value.
|
||||
api_key: OpenAI API key for authentication, or a callable that returns an API key.
|
||||
Must be non-empty. (Optional)
|
||||
org_id: OpenAI organization ID. This is optional
|
||||
unless the account belongs to multiple organizations.
|
||||
default_headers: Default headers
|
||||
for HTTP requests. (Optional)
|
||||
client: An existing OpenAI client, optional.
|
||||
instruction_role: The role to use for 'instruction'
|
||||
messages, for example, summarization prompts could use `developer` or `system`. (Optional)
|
||||
base_url: The optional base URL to use. If provided will override the standard value for a OpenAI connector.
|
||||
Will not be used when supplying a custom client.
|
||||
kwargs: Additional keyword arguments.
|
||||
|
||||
"""
|
||||
# Merge APP_INFO into the headers if it exists
|
||||
merged_headers = dict(copy(default_headers)) if default_headers else {}
|
||||
if APP_INFO:
|
||||
merged_headers.update(APP_INFO)
|
||||
merged_headers = prepend_agent_framework_to_user_agent(merged_headers)
|
||||
|
||||
# Handle callable API key using base class method
|
||||
api_key_value = self._get_api_key(api_key)
|
||||
|
||||
if not client:
|
||||
if not api_key:
|
||||
raise ValueError("Please provide an api_key")
|
||||
args: dict[str, Any] = {"api_key": api_key_value, "default_headers": merged_headers}
|
||||
if org_id:
|
||||
args["organization"] = org_id
|
||||
if base_url:
|
||||
args["base_url"] = base_url
|
||||
client = AsyncOpenAI(**args)
|
||||
|
||||
# Store configuration as instance attributes for serialization
|
||||
self.org_id = org_id
|
||||
self.base_url = str(base_url)
|
||||
# Store default_headers but filter out USER_AGENT_KEY for serialization
|
||||
if default_headers:
|
||||
self.default_headers: dict[str, Any] | None = {
|
||||
k: v for k, v in default_headers.items() if k != USER_AGENT_KEY
|
||||
}
|
||||
else:
|
||||
self.default_headers = None
|
||||
|
||||
args = {
|
||||
"model": model,
|
||||
"client": client,
|
||||
}
|
||||
if instruction_role:
|
||||
args["instruction_role"] = instruction_role
|
||||
|
||||
# Ensure additional_properties and middleware are passed through kwargs to RawChatClient
|
||||
# These are consumed by RawChatClient.__init__ via kwargs
|
||||
super().__init__(**args, **kwargs)
|
||||
|
||||
|
||||
def to_assistant_tools(
|
||||
tools: Sequence[FunctionTool | MutableMapping[str, Any]] | None,
|
||||
) -> list[dict[str, Any]]:
|
||||
"""Convert Agent Framework tools to OpenAI Assistants API format.
|
||||
|
||||
Handles FunctionTool instances and dict-based tools from static factory methods.
|
||||
|
||||
Args:
|
||||
tools: Sequence of Agent Framework tools.
|
||||
|
||||
Returns:
|
||||
List of tool definitions for OpenAI Assistants API.
|
||||
"""
|
||||
if not tools:
|
||||
return []
|
||||
|
||||
tool_definitions: list[dict[str, Any]] = []
|
||||
|
||||
for tool in tools:
|
||||
if isinstance(tool, FunctionTool):
|
||||
tool_definitions.append(tool.to_json_schema_spec())
|
||||
elif isinstance(tool, MutableMapping):
|
||||
# Pass through dict-based tools directly (from static factory methods)
|
||||
tool_definitions.append(dict(tool))
|
||||
|
||||
return tool_definitions
|
||||
|
||||
|
||||
def from_assistant_tools(
|
||||
assistant_tools: list[Any] | None,
|
||||
) -> list[dict[str, Any]]:
|
||||
"""Convert OpenAI Assistant tools to dict-based format.
|
||||
|
||||
This converts hosted tools (code_interpreter, file_search) from an OpenAI
|
||||
Assistant definition back to dict-based tool definitions.
|
||||
|
||||
Note: Function tools are skipped - user must provide implementations separately.
|
||||
|
||||
Args:
|
||||
assistant_tools: Tools from OpenAI Assistant object (assistant.tools).
|
||||
|
||||
Returns:
|
||||
List of dict-based tool definitions for hosted tools.
|
||||
"""
|
||||
if not assistant_tools:
|
||||
return []
|
||||
|
||||
tools: list[dict[str, Any]] = []
|
||||
|
||||
for tool in assistant_tools:
|
||||
if hasattr(tool, "type"):
|
||||
tool_type = tool.type
|
||||
elif isinstance(tool, Mapping):
|
||||
typed_tool = cast(Mapping[str, Any], tool)
|
||||
tool_type_value: Any = typed_tool.get("type")
|
||||
tool_type = tool_type_value if isinstance(tool_type_value, str) else None
|
||||
else:
|
||||
tool_type = None
|
||||
|
||||
if tool_type == "code_interpreter":
|
||||
tools.append({"type": "code_interpreter"})
|
||||
elif tool_type == "file_search":
|
||||
tools.append({"type": "file_search"})
|
||||
# Skip function tools - user must provide implementations
|
||||
|
||||
return tools
|
||||
@@ -0,0 +1,98 @@
|
||||
[project]
|
||||
name = "agent-framework-openai"
|
||||
description = "OpenAI integration for Microsoft Agent Framework."
|
||||
authors = [{ name = "Microsoft", email = "af-support@microsoft.com"}]
|
||||
readme = "README.md"
|
||||
requires-python = ">=3.10"
|
||||
version = "1.0.0rc5"
|
||||
license-files = ["LICENSE"]
|
||||
urls.homepage = "https://aka.ms/agent-framework"
|
||||
urls.source = "https://github.com/microsoft/agent-framework/tree/main/python"
|
||||
urls.release_notes = "https://github.com/microsoft/agent-framework/releases?q=tag%3Apython-1&expanded=true"
|
||||
urls.issues = "https://github.com/microsoft/agent-framework/issues"
|
||||
classifiers = [
|
||||
"License :: OSI Approved :: MIT License",
|
||||
"Development Status :: 4 - Beta",
|
||||
"Intended Audience :: Developers",
|
||||
"Programming Language :: Python :: 3",
|
||||
"Programming Language :: Python :: 3.10",
|
||||
"Programming Language :: Python :: 3.11",
|
||||
"Programming Language :: Python :: 3.12",
|
||||
"Programming Language :: Python :: 3.13",
|
||||
"Programming Language :: Python :: 3.14",
|
||||
"Typing :: Typed",
|
||||
]
|
||||
dependencies = [
|
||||
"agent-framework-core>=1.0.0rc5",
|
||||
"openai>=1.99.0,<3",
|
||||
"packaging>=24.1,<25",
|
||||
]
|
||||
|
||||
[tool.uv]
|
||||
prerelease = "if-necessary-or-explicit"
|
||||
environments = [
|
||||
"sys_platform == 'darwin'",
|
||||
"sys_platform == 'linux'",
|
||||
"sys_platform == 'win32'"
|
||||
]
|
||||
|
||||
[tool.uv-dynamic-versioning]
|
||||
fallback-version = "0.0.0"
|
||||
|
||||
[tool.pytest.ini_options]
|
||||
testpaths = 'tests'
|
||||
addopts = "-ra -q -r fEX"
|
||||
asyncio_mode = "auto"
|
||||
asyncio_default_fixture_loop_scope = "function"
|
||||
filterwarnings = []
|
||||
timeout = 120
|
||||
markers = [
|
||||
"azure: marks tests as Azure-backed OpenAI specific",
|
||||
"integration: marks tests as integration tests that require external services",
|
||||
]
|
||||
|
||||
[tool.ruff]
|
||||
extend = "../../pyproject.toml"
|
||||
|
||||
[tool.coverage.run]
|
||||
omit = [
|
||||
"**/__init__.py"
|
||||
]
|
||||
|
||||
[tool.pyright]
|
||||
extends = "../../pyproject.toml"
|
||||
exclude = ['tests']
|
||||
|
||||
[tool.mypy]
|
||||
plugins = ['pydantic.mypy']
|
||||
strict = true
|
||||
python_version = "3.10"
|
||||
ignore_missing_imports = true
|
||||
disallow_untyped_defs = true
|
||||
no_implicit_optional = true
|
||||
check_untyped_defs = true
|
||||
warn_return_any = true
|
||||
show_error_codes = true
|
||||
warn_unused_ignores = false
|
||||
disallow_incomplete_defs = true
|
||||
disallow_untyped_decorators = true
|
||||
|
||||
[tool.bandit]
|
||||
targets = ["agent_framework_openai"]
|
||||
exclude_dirs = ["tests"]
|
||||
|
||||
[tool.poe]
|
||||
executor.type = "uv"
|
||||
include = "../../shared_tasks.toml"
|
||||
|
||||
[tool.poe.tasks.mypy]
|
||||
help = "Run MyPy for this package."
|
||||
cmd = "mypy --config-file $POE_ROOT/pyproject.toml agent_framework_openai"
|
||||
|
||||
[tool.poe.tasks.test]
|
||||
help = "Run the default unit test suite for this package."
|
||||
cmd = 'pytest -m "not integration" --cov=agent_framework_openai --cov-report=term-missing:skip-covered -n auto --dist worksteal tests'
|
||||
|
||||
[build-system]
|
||||
requires = ["flit-core >= 3.11,<4.0"]
|
||||
build-backend = "flit_core.buildapi"
|
||||
Binary file not shown.
|
After Width: | Height: | Size: 178 KiB |
@@ -0,0 +1,201 @@
|
||||
# Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
from collections.abc import Generator
|
||||
from typing import Any
|
||||
from unittest.mock import patch
|
||||
|
||||
from opentelemetry.sdk.trace.export import SimpleSpanProcessor, SpanExporter
|
||||
from opentelemetry.sdk.trace.export.in_memory_span_exporter import InMemorySpanExporter
|
||||
from pytest import fixture
|
||||
|
||||
|
||||
def _reset_env(monkeypatch, env_names: list[str]) -> None: # type: ignore
|
||||
for env_name in env_names:
|
||||
monkeypatch.delenv(env_name, raising=False) # type: ignore
|
||||
|
||||
|
||||
# region Connector Settings fixtures
|
||||
@fixture
|
||||
def exclude_list(request: Any) -> list[str]:
|
||||
"""Fixture that returns a list of environment variables to exclude."""
|
||||
return request.param if hasattr(request, "param") else []
|
||||
|
||||
|
||||
@fixture
|
||||
def override_env_param_dict(request: Any) -> dict[str, str]:
|
||||
"""Fixture that returns a dict of environment variables to override."""
|
||||
return request.param if hasattr(request, "param") else {}
|
||||
|
||||
|
||||
@fixture()
|
||||
def openai_unit_test_env(monkeypatch, exclude_list, override_env_param_dict): # type: ignore
|
||||
"""Fixture to set environment variables for OpenAISettings."""
|
||||
if exclude_list is None:
|
||||
exclude_list = []
|
||||
|
||||
if override_env_param_dict is None:
|
||||
override_env_param_dict = {}
|
||||
|
||||
_reset_env(
|
||||
monkeypatch,
|
||||
[
|
||||
"OPENAI_API_KEY",
|
||||
"OPENAI_ORG_ID",
|
||||
"OPENAI_MODEL",
|
||||
"OPENAI_EMBEDDING_MODEL",
|
||||
"OPENAI_TEXT_MODEL_ID",
|
||||
"OPENAI_TEXT_TO_IMAGE_MODEL_ID",
|
||||
"OPENAI_AUDIO_TO_TEXT_MODEL_ID",
|
||||
"OPENAI_TEXT_TO_AUDIO_MODEL_ID",
|
||||
"OPENAI_REALTIME_MODEL_ID",
|
||||
"OPENAI_BASE_URL",
|
||||
"AZURE_OPENAI_ENDPOINT",
|
||||
"AZURE_OPENAI_BASE_URL",
|
||||
"AZURE_OPENAI_API_KEY",
|
||||
"AZURE_OPENAI_DEPLOYMENT_NAME",
|
||||
"AZURE_OPENAI_API_VERSION",
|
||||
],
|
||||
)
|
||||
|
||||
env_vars = {
|
||||
"OPENAI_API_KEY": "test-dummy-key",
|
||||
"OPENAI_ORG_ID": "test_org_id",
|
||||
"OPENAI_MODEL": "test_model_id",
|
||||
"OPENAI_EMBEDDING_MODEL": "test_embedding_model_id",
|
||||
"OPENAI_TEXT_MODEL_ID": "test_text_model_id",
|
||||
"OPENAI_TEXT_TO_IMAGE_MODEL_ID": "test_text_to_image_model_id",
|
||||
"OPENAI_AUDIO_TO_TEXT_MODEL_ID": "test_audio_to_text_model_id",
|
||||
"OPENAI_TEXT_TO_AUDIO_MODEL_ID": "test_text_to_audio_model_id",
|
||||
"OPENAI_REALTIME_MODEL_ID": "test_realtime_model_id",
|
||||
}
|
||||
|
||||
env_vars.update(override_env_param_dict) # type: ignore
|
||||
|
||||
for key, value in env_vars.items():
|
||||
if key in exclude_list:
|
||||
monkeypatch.delenv(key, raising=False) # type: ignore
|
||||
continue
|
||||
monkeypatch.setenv(key, value) # type: ignore
|
||||
|
||||
return env_vars
|
||||
|
||||
|
||||
@fixture()
|
||||
def azure_openai_unit_test_env(monkeypatch, exclude_list, override_env_param_dict): # type: ignore
|
||||
"""Fixture to set environment variables for Azure-backed OpenAI tests."""
|
||||
if exclude_list is None:
|
||||
exclude_list = []
|
||||
|
||||
if override_env_param_dict is None:
|
||||
override_env_param_dict = {}
|
||||
|
||||
_reset_env(
|
||||
monkeypatch,
|
||||
[
|
||||
"OPENAI_API_KEY",
|
||||
"OPENAI_ORG_ID",
|
||||
"OPENAI_MODEL",
|
||||
"OPENAI_EMBEDDING_MODEL",
|
||||
"OPENAI_TEXT_MODEL_ID",
|
||||
"OPENAI_TEXT_TO_IMAGE_MODEL_ID",
|
||||
"OPENAI_AUDIO_TO_TEXT_MODEL_ID",
|
||||
"OPENAI_TEXT_TO_AUDIO_MODEL_ID",
|
||||
"OPENAI_REALTIME_MODEL_ID",
|
||||
"OPENAI_BASE_URL",
|
||||
"AZURE_OPENAI_ENDPOINT",
|
||||
"AZURE_OPENAI_BASE_URL",
|
||||
"AZURE_OPENAI_API_KEY",
|
||||
"AZURE_OPENAI_DEPLOYMENT_NAME",
|
||||
"AZURE_OPENAI_API_VERSION",
|
||||
],
|
||||
)
|
||||
|
||||
env_vars = {
|
||||
"AZURE_OPENAI_ENDPOINT": "https://test-endpoint.openai.azure.com",
|
||||
"AZURE_OPENAI_DEPLOYMENT_NAME": "test_deployment",
|
||||
"AZURE_OPENAI_API_KEY": "test_api_key",
|
||||
"AZURE_OPENAI_API_VERSION": "2024-12-01-preview",
|
||||
}
|
||||
|
||||
env_vars.update(override_env_param_dict) # type: ignore
|
||||
|
||||
for key, value in env_vars.items():
|
||||
if key in exclude_list:
|
||||
monkeypatch.delenv(key, raising=False) # type: ignore
|
||||
continue
|
||||
monkeypatch.setenv(key, value) # type: ignore
|
||||
|
||||
return env_vars
|
||||
|
||||
|
||||
# region Observability fixtures
|
||||
@fixture
|
||||
def enable_instrumentation(request: Any) -> bool:
|
||||
"""Fixture that returns a boolean indicating if Otel is enabled."""
|
||||
return request.param if hasattr(request, "param") else True
|
||||
|
||||
|
||||
@fixture
|
||||
def enable_sensitive_data(request: Any) -> bool:
|
||||
"""Fixture that returns a boolean indicating if sensitive data is enabled."""
|
||||
return request.param if hasattr(request, "param") else True
|
||||
|
||||
|
||||
@fixture
|
||||
def span_exporter(monkeypatch, enable_instrumentation: bool, enable_sensitive_data: bool) -> Generator[SpanExporter]:
|
||||
"""Fixture to remove environment variables for ObservabilitySettings."""
|
||||
env_vars = [
|
||||
"ENABLE_INSTRUMENTATION",
|
||||
"ENABLE_SENSITIVE_DATA",
|
||||
"ENABLE_CONSOLE_EXPORTERS",
|
||||
"OTEL_EXPORTER_OTLP_ENDPOINT",
|
||||
"OTEL_EXPORTER_OTLP_TRACES_ENDPOINT",
|
||||
"OTEL_EXPORTER_OTLP_METRICS_ENDPOINT",
|
||||
"OTEL_EXPORTER_OTLP_LOGS_ENDPOINT",
|
||||
"OTEL_EXPORTER_OTLP_PROTOCOL",
|
||||
"OTEL_EXPORTER_OTLP_HEADERS",
|
||||
"OTEL_EXPORTER_OTLP_TRACES_HEADERS",
|
||||
"OTEL_EXPORTER_OTLP_METRICS_HEADERS",
|
||||
"OTEL_EXPORTER_OTLP_LOGS_HEADERS",
|
||||
"OTEL_SERVICE_NAME",
|
||||
"OTEL_SERVICE_VERSION",
|
||||
"OTEL_RESOURCE_ATTRIBUTES",
|
||||
]
|
||||
|
||||
for key in env_vars:
|
||||
monkeypatch.delenv(key, raising=False) # type: ignore
|
||||
monkeypatch.setenv("ENABLE_INSTRUMENTATION", str(enable_instrumentation)) # type: ignore
|
||||
if not enable_instrumentation:
|
||||
enable_sensitive_data = False
|
||||
monkeypatch.setenv("ENABLE_SENSITIVE_DATA", str(enable_sensitive_data)) # type: ignore
|
||||
import importlib
|
||||
|
||||
import agent_framework.observability as observability
|
||||
from opentelemetry import trace
|
||||
|
||||
importlib.reload(observability)
|
||||
|
||||
observability_settings = observability.ObservabilitySettings()
|
||||
|
||||
if enable_instrumentation or enable_sensitive_data:
|
||||
from opentelemetry.sdk.trace import TracerProvider
|
||||
|
||||
tracer_provider = TracerProvider(resource=observability_settings._resource)
|
||||
trace.set_tracer_provider(tracer_provider)
|
||||
|
||||
monkeypatch.setattr(observability, "OBSERVABILITY_SETTINGS", observability_settings, raising=False) # type: ignore
|
||||
|
||||
with (
|
||||
patch("agent_framework.observability.OBSERVABILITY_SETTINGS", observability_settings),
|
||||
patch("agent_framework.observability.configure_otel_providers"),
|
||||
):
|
||||
exporter = InMemorySpanExporter()
|
||||
if enable_instrumentation or enable_sensitive_data:
|
||||
tracer_provider = trace.get_tracer_provider()
|
||||
if not hasattr(tracer_provider, "add_span_processor"):
|
||||
raise RuntimeError("Tracer provider does not support adding span processors.")
|
||||
|
||||
tracer_provider.add_span_processor(SimpleSpanProcessor(exporter)) # type: ignore
|
||||
|
||||
yield exporter
|
||||
exporter.clear()
|
||||
@@ -0,0 +1,813 @@
|
||||
# Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
import os
|
||||
from typing import Annotated, Any
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import pytest
|
||||
from agent_framework import Agent, normalize_tools, tool
|
||||
from openai.types.beta.assistant import Assistant
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from agent_framework_openai import OpenAIAssistantProvider, OpenAIAssistantsClient
|
||||
from agent_framework_openai._shared import from_assistant_tools, to_assistant_tools
|
||||
|
||||
# region Test Helpers
|
||||
|
||||
|
||||
def create_mock_assistant(
|
||||
assistant_id: str = "asst_test123",
|
||||
name: str = "TestAssistant",
|
||||
model: str = "gpt-4",
|
||||
instructions: str | None = "You are a helpful assistant.",
|
||||
description: str | None = None,
|
||||
tools: list[Any] | None = None,
|
||||
) -> Assistant:
|
||||
"""Create a mock Assistant object."""
|
||||
mock = MagicMock(spec=Assistant)
|
||||
mock.id = assistant_id
|
||||
mock.name = name
|
||||
mock.model = model
|
||||
mock.instructions = instructions
|
||||
mock.description = description
|
||||
mock.tools = tools or []
|
||||
return mock
|
||||
|
||||
|
||||
def create_function_tool(name: str, description: str = "A test function") -> MagicMock:
|
||||
"""Create a mock FunctionTool."""
|
||||
mock = MagicMock()
|
||||
mock.type = "function"
|
||||
mock.function = MagicMock()
|
||||
mock.function.name = name
|
||||
mock.function.description = description
|
||||
return mock
|
||||
|
||||
|
||||
def create_code_interpreter_tool() -> MagicMock:
|
||||
"""Create a mock CodeInterpreterTool."""
|
||||
mock = MagicMock()
|
||||
mock.type = "code_interpreter"
|
||||
return mock
|
||||
|
||||
|
||||
def create_file_search_tool() -> MagicMock:
|
||||
"""Create a mock FileSearchTool."""
|
||||
mock = MagicMock()
|
||||
mock.type = "file_search"
|
||||
return mock
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_async_openai() -> MagicMock:
|
||||
"""Mock AsyncOpenAI client."""
|
||||
mock_client = MagicMock()
|
||||
|
||||
# Mock beta.assistants
|
||||
mock_client.beta.assistants.create = AsyncMock(
|
||||
return_value=create_mock_assistant(assistant_id="asst_created123", name="CreatedAssistant")
|
||||
)
|
||||
mock_client.beta.assistants.retrieve = AsyncMock(
|
||||
return_value=create_mock_assistant(assistant_id="asst_retrieved123", name="RetrievedAssistant")
|
||||
)
|
||||
mock_client.beta.assistants.delete = AsyncMock()
|
||||
|
||||
# Mock close method
|
||||
mock_client.close = AsyncMock()
|
||||
|
||||
return mock_client
|
||||
|
||||
|
||||
# Test function for tool validation
|
||||
def get_weather(location: Annotated[str, Field(description="The location")]) -> str:
|
||||
"""Get the weather for a location."""
|
||||
return f"Weather in {location}: sunny"
|
||||
|
||||
|
||||
def search_database(query: Annotated[str, Field(description="Search query")]) -> str:
|
||||
"""Search the database."""
|
||||
return f"Results for: {query}"
|
||||
|
||||
|
||||
# Pydantic model for structured output tests
|
||||
class WeatherResponse(BaseModel):
|
||||
location: str
|
||||
temperature: float
|
||||
conditions: str
|
||||
|
||||
|
||||
# endregion
|
||||
|
||||
# region Initialization Tests
|
||||
|
||||
|
||||
class TestOpenAIAssistantProviderInit:
|
||||
"""Tests for provider initialization."""
|
||||
|
||||
def test_init_with_client(self, mock_async_openai: MagicMock) -> None:
|
||||
"""Test initialization with existing AsyncOpenAI client."""
|
||||
provider = OpenAIAssistantProvider(mock_async_openai)
|
||||
|
||||
assert provider._client is mock_async_openai # type: ignore[reportPrivateUsage]
|
||||
assert provider._should_close_client is False # type: ignore[reportPrivateUsage]
|
||||
|
||||
def test_init_without_client_creates_one(self, openai_unit_test_env: dict[str, str]) -> None:
|
||||
"""Test initialization creates client from settings."""
|
||||
provider = OpenAIAssistantProvider()
|
||||
|
||||
assert provider._client is not None # type: ignore[reportPrivateUsage]
|
||||
assert provider._should_close_client is True # type: ignore[reportPrivateUsage]
|
||||
|
||||
def test_init_with_api_key(self) -> None:
|
||||
"""Test initialization with explicit API key."""
|
||||
provider = OpenAIAssistantProvider(api_key="sk-test-key")
|
||||
|
||||
assert provider._client is not None # type: ignore[reportPrivateUsage]
|
||||
assert provider._should_close_client is True # type: ignore[reportPrivateUsage]
|
||||
|
||||
def test_init_fails_without_api_key(self) -> None:
|
||||
"""Test initialization fails without API key when settings return None."""
|
||||
from unittest.mock import patch
|
||||
|
||||
# Mock load_settings to return a dict with None for api_key
|
||||
with patch("agent_framework_openai._assistant_provider.load_settings") as mock_load:
|
||||
mock_load.return_value = {
|
||||
"api_key": None,
|
||||
"org_id": None,
|
||||
"base_url": None,
|
||||
"model": None,
|
||||
}
|
||||
|
||||
with pytest.raises(ValueError) as exc_info:
|
||||
OpenAIAssistantProvider()
|
||||
|
||||
assert "API key is required" in str(exc_info.value)
|
||||
|
||||
def test_init_with_org_id_and_base_url(self) -> None:
|
||||
"""Test initialization with organization ID and base URL."""
|
||||
provider = OpenAIAssistantProvider(
|
||||
api_key="sk-test-key",
|
||||
org_id="org-123",
|
||||
base_url="https://custom.openai.com",
|
||||
)
|
||||
|
||||
assert provider._client is not None # type: ignore[reportPrivateUsage]
|
||||
|
||||
|
||||
class TestOpenAIAssistantProviderContextManager:
|
||||
"""Tests for async context manager."""
|
||||
|
||||
async def test_context_manager_enter_exit(self, mock_async_openai: MagicMock) -> None:
|
||||
"""Test async context manager entry and exit."""
|
||||
provider = OpenAIAssistantProvider(mock_async_openai)
|
||||
|
||||
async with provider as p:
|
||||
assert p is provider
|
||||
|
||||
async def test_context_manager_closes_owned_client(self, openai_unit_test_env: dict[str, str]) -> None:
|
||||
"""Test that owned client is closed on exit."""
|
||||
provider = OpenAIAssistantProvider()
|
||||
client = provider._client # type: ignore[reportPrivateUsage]
|
||||
assert client is not None
|
||||
client.close = AsyncMock()
|
||||
|
||||
async with provider:
|
||||
pass
|
||||
|
||||
client.close.assert_called_once()
|
||||
|
||||
async def test_context_manager_does_not_close_external_client(self, mock_async_openai: MagicMock) -> None:
|
||||
"""Test that external client is not closed on exit."""
|
||||
provider = OpenAIAssistantProvider(mock_async_openai)
|
||||
|
||||
async with provider:
|
||||
pass
|
||||
|
||||
mock_async_openai.close.assert_not_called()
|
||||
|
||||
|
||||
# endregion
|
||||
|
||||
# region create_agent Tests
|
||||
|
||||
|
||||
class TestOpenAIAssistantProviderCreateAgent:
|
||||
"""Tests for create_agent method."""
|
||||
|
||||
async def test_create_agent_basic(self, mock_async_openai: MagicMock) -> None:
|
||||
"""Test basic assistant creation."""
|
||||
provider = OpenAIAssistantProvider(mock_async_openai)
|
||||
|
||||
agent = await provider.create_agent(
|
||||
name="TestAgent",
|
||||
model="gpt-4",
|
||||
instructions="You are helpful.",
|
||||
)
|
||||
|
||||
assert isinstance(agent, Agent)
|
||||
assert agent.name == "CreatedAssistant"
|
||||
mock_async_openai.beta.assistants.create.assert_called_once()
|
||||
|
||||
# Verify create was called with correct parameters
|
||||
call_kwargs = mock_async_openai.beta.assistants.create.call_args.kwargs
|
||||
assert call_kwargs["name"] == "TestAgent"
|
||||
assert call_kwargs["model"] == "gpt-4"
|
||||
assert call_kwargs["instructions"] == "You are helpful."
|
||||
|
||||
async def test_create_agent_with_description(self, mock_async_openai: MagicMock) -> None:
|
||||
"""Test assistant creation with description."""
|
||||
provider = OpenAIAssistantProvider(mock_async_openai)
|
||||
|
||||
await provider.create_agent(
|
||||
name="TestAgent",
|
||||
model="gpt-4",
|
||||
description="A test agent description",
|
||||
)
|
||||
|
||||
call_kwargs = mock_async_openai.beta.assistants.create.call_args.kwargs
|
||||
assert call_kwargs["description"] == "A test agent description"
|
||||
|
||||
async def test_create_agent_with_function_tools(self, mock_async_openai: MagicMock) -> None:
|
||||
"""Test assistant creation with function tools."""
|
||||
provider = OpenAIAssistantProvider(mock_async_openai)
|
||||
|
||||
agent = await provider.create_agent(
|
||||
name="WeatherAgent",
|
||||
model="gpt-4",
|
||||
tools=[get_weather],
|
||||
)
|
||||
|
||||
assert isinstance(agent, Agent)
|
||||
|
||||
# Verify tools were passed to create
|
||||
call_kwargs = mock_async_openai.beta.assistants.create.call_args.kwargs
|
||||
assert "tools" in call_kwargs
|
||||
assert len(call_kwargs["tools"]) == 1
|
||||
assert call_kwargs["tools"][0]["type"] == "function"
|
||||
assert call_kwargs["tools"][0]["function"]["name"] == "get_weather"
|
||||
|
||||
async def test_create_agent_with_tool(self, mock_async_openai: MagicMock) -> None:
|
||||
"""Test assistant creation with FunctionTool."""
|
||||
provider = OpenAIAssistantProvider(mock_async_openai)
|
||||
|
||||
@tool
|
||||
def my_function(x: int) -> int:
|
||||
"""Double a number."""
|
||||
return x * 2
|
||||
|
||||
await provider.create_agent(
|
||||
name="TestAgent",
|
||||
model="gpt-4",
|
||||
tools=[my_function],
|
||||
)
|
||||
|
||||
call_kwargs = mock_async_openai.beta.assistants.create.call_args.kwargs
|
||||
assert call_kwargs["tools"][0]["function"]["name"] == "my_function"
|
||||
|
||||
async def test_create_agent_with_code_interpreter(self, mock_async_openai: MagicMock) -> None:
|
||||
"""Test assistant creation with code interpreter."""
|
||||
provider = OpenAIAssistantProvider(mock_async_openai)
|
||||
|
||||
await provider.create_agent(
|
||||
name="CodeAgent",
|
||||
model="gpt-4",
|
||||
tools=[OpenAIAssistantsClient.get_code_interpreter_tool()],
|
||||
)
|
||||
|
||||
call_kwargs = mock_async_openai.beta.assistants.create.call_args.kwargs
|
||||
assert {"type": "code_interpreter"} in call_kwargs["tools"]
|
||||
|
||||
async def test_create_agent_with_file_search(self, mock_async_openai: MagicMock) -> None:
|
||||
"""Test assistant creation with file search."""
|
||||
provider = OpenAIAssistantProvider(mock_async_openai)
|
||||
|
||||
await provider.create_agent(
|
||||
name="SearchAgent",
|
||||
model="gpt-4",
|
||||
tools=[OpenAIAssistantsClient.get_file_search_tool()],
|
||||
)
|
||||
|
||||
call_kwargs = mock_async_openai.beta.assistants.create.call_args.kwargs
|
||||
assert any(t["type"] == "file_search" for t in call_kwargs["tools"])
|
||||
|
||||
async def test_create_agent_with_file_search_max_results(self, mock_async_openai: MagicMock) -> None:
|
||||
"""Test assistant creation with file search and max_results."""
|
||||
provider = OpenAIAssistantProvider(mock_async_openai)
|
||||
|
||||
await provider.create_agent(
|
||||
name="SearchAgent",
|
||||
model="gpt-4",
|
||||
tools=[OpenAIAssistantsClient.get_file_search_tool(max_num_results=10)],
|
||||
)
|
||||
|
||||
call_kwargs = mock_async_openai.beta.assistants.create.call_args.kwargs
|
||||
file_search_tool = next(t for t in call_kwargs["tools"] if t["type"] == "file_search")
|
||||
assert file_search_tool.get("file_search", {}).get("max_num_results") == 10
|
||||
|
||||
async def test_create_agent_with_mixed_tools(self, mock_async_openai: MagicMock) -> None:
|
||||
"""Test assistant creation with multiple tool types."""
|
||||
provider = OpenAIAssistantProvider(mock_async_openai)
|
||||
|
||||
await provider.create_agent(
|
||||
name="MultiToolAgent",
|
||||
model="gpt-4",
|
||||
tools=[
|
||||
get_weather,
|
||||
OpenAIAssistantsClient.get_code_interpreter_tool(),
|
||||
OpenAIAssistantsClient.get_file_search_tool(),
|
||||
],
|
||||
)
|
||||
|
||||
call_kwargs = mock_async_openai.beta.assistants.create.call_args.kwargs
|
||||
assert len(call_kwargs["tools"]) == 3
|
||||
|
||||
async def test_create_agent_with_metadata(self, mock_async_openai: MagicMock) -> None:
|
||||
"""Test assistant creation with metadata."""
|
||||
provider = OpenAIAssistantProvider(mock_async_openai)
|
||||
|
||||
await provider.create_agent(
|
||||
name="TestAgent",
|
||||
model="gpt-4",
|
||||
metadata={"env": "test", "version": "1.0"},
|
||||
)
|
||||
|
||||
call_kwargs = mock_async_openai.beta.assistants.create.call_args.kwargs
|
||||
assert call_kwargs["metadata"] == {"env": "test", "version": "1.0"}
|
||||
|
||||
async def test_create_agent_with_response_format_pydantic(self, mock_async_openai: MagicMock) -> None:
|
||||
"""Test assistant creation with Pydantic response format via default_options."""
|
||||
provider = OpenAIAssistantProvider(mock_async_openai)
|
||||
|
||||
await provider.create_agent(
|
||||
name="StructuredAgent",
|
||||
model="gpt-4",
|
||||
default_options={"response_format": WeatherResponse},
|
||||
)
|
||||
|
||||
call_kwargs = mock_async_openai.beta.assistants.create.call_args.kwargs
|
||||
assert call_kwargs["response_format"]["type"] == "json_schema"
|
||||
assert call_kwargs["response_format"]["json_schema"]["name"] == "WeatherResponse"
|
||||
|
||||
async def test_create_agent_returns_chat_agent(self, mock_async_openai: MagicMock) -> None:
|
||||
"""Test that create_agent returns a Agent instance."""
|
||||
provider = OpenAIAssistantProvider(mock_async_openai)
|
||||
|
||||
agent = await provider.create_agent(
|
||||
name="TestAgent",
|
||||
model="gpt-4",
|
||||
)
|
||||
|
||||
assert isinstance(agent, Agent)
|
||||
|
||||
|
||||
# endregion
|
||||
|
||||
# region get_agent Tests
|
||||
|
||||
|
||||
class TestOpenAIAssistantProviderGetAgent:
|
||||
"""Tests for get_agent method."""
|
||||
|
||||
async def test_get_agent_basic(self, mock_async_openai: MagicMock) -> None:
|
||||
"""Test retrieving an existing assistant."""
|
||||
provider = OpenAIAssistantProvider(mock_async_openai)
|
||||
|
||||
agent = await provider.get_agent(assistant_id="asst_123")
|
||||
|
||||
assert isinstance(agent, Agent)
|
||||
mock_async_openai.beta.assistants.retrieve.assert_called_once_with("asst_123")
|
||||
|
||||
async def test_get_agent_with_instructions_override(self, mock_async_openai: MagicMock) -> None:
|
||||
"""Test retrieving assistant with instruction override."""
|
||||
provider = OpenAIAssistantProvider(mock_async_openai)
|
||||
|
||||
agent = await provider.get_agent(
|
||||
assistant_id="asst_123",
|
||||
instructions="Custom instructions",
|
||||
)
|
||||
|
||||
# Agent should be created successfully with the custom instructions
|
||||
assert isinstance(agent, Agent)
|
||||
assert agent.id == "asst_retrieved123"
|
||||
|
||||
async def test_get_agent_with_function_tools(self, mock_async_openai: MagicMock) -> None:
|
||||
"""Test retrieving assistant with function tools provided."""
|
||||
# Setup assistant with function tool
|
||||
assistant = create_mock_assistant(tools=[create_function_tool("get_weather")])
|
||||
mock_async_openai.beta.assistants.retrieve = AsyncMock(return_value=assistant)
|
||||
|
||||
provider = OpenAIAssistantProvider(mock_async_openai)
|
||||
|
||||
agent = await provider.get_agent(
|
||||
assistant_id="asst_123",
|
||||
tools=[get_weather],
|
||||
)
|
||||
|
||||
assert isinstance(agent, Agent)
|
||||
|
||||
async def test_get_agent_validates_missing_function_tools(self, mock_async_openai: MagicMock) -> None:
|
||||
"""Test that missing function tools raise ValueError."""
|
||||
# Setup assistant with function tool
|
||||
assistant = create_mock_assistant(tools=[create_function_tool("get_weather")])
|
||||
mock_async_openai.beta.assistants.retrieve = AsyncMock(return_value=assistant)
|
||||
|
||||
provider = OpenAIAssistantProvider(mock_async_openai)
|
||||
|
||||
with pytest.raises(ValueError) as exc_info:
|
||||
await provider.get_agent(assistant_id="asst_123")
|
||||
|
||||
assert "get_weather" in str(exc_info.value)
|
||||
assert "no implementation was provided" in str(exc_info.value)
|
||||
|
||||
async def test_get_agent_validates_multiple_missing_function_tools(self, mock_async_openai: MagicMock) -> None:
|
||||
"""Test validation with multiple missing function tools."""
|
||||
assistant = create_mock_assistant(
|
||||
tools=[create_function_tool("get_weather"), create_function_tool("search_database")]
|
||||
)
|
||||
mock_async_openai.beta.assistants.retrieve = AsyncMock(return_value=assistant)
|
||||
|
||||
provider = OpenAIAssistantProvider(mock_async_openai)
|
||||
|
||||
with pytest.raises(ValueError) as exc_info:
|
||||
await provider.get_agent(assistant_id="asst_123")
|
||||
|
||||
error_msg = str(exc_info.value)
|
||||
assert "get_weather" in error_msg or "search_database" in error_msg
|
||||
|
||||
async def test_get_agent_merges_hosted_tools(self, mock_async_openai: MagicMock) -> None:
|
||||
"""Test that hosted tools are automatically included."""
|
||||
assistant = create_mock_assistant(tools=[create_code_interpreter_tool(), create_file_search_tool()])
|
||||
mock_async_openai.beta.assistants.retrieve = AsyncMock(return_value=assistant)
|
||||
|
||||
provider = OpenAIAssistantProvider(mock_async_openai)
|
||||
|
||||
agent = await provider.get_agent(assistant_id="asst_123")
|
||||
|
||||
# Hosted tools should be merged automatically
|
||||
assert isinstance(agent, Agent)
|
||||
|
||||
|
||||
# endregion
|
||||
|
||||
# region as_agent Tests
|
||||
|
||||
|
||||
class TestOpenAIAssistantProviderAsAgent:
|
||||
"""Tests for as_agent method."""
|
||||
|
||||
def test_as_agent_no_http_call(self, mock_async_openai: MagicMock) -> None:
|
||||
"""Test that as_agent doesn't make HTTP calls."""
|
||||
provider = OpenAIAssistantProvider(mock_async_openai)
|
||||
assistant = create_mock_assistant()
|
||||
|
||||
agent = provider.as_agent(assistant)
|
||||
|
||||
assert isinstance(agent, Agent)
|
||||
# Verify no HTTP calls were made
|
||||
mock_async_openai.beta.assistants.create.assert_not_called()
|
||||
mock_async_openai.beta.assistants.retrieve.assert_not_called()
|
||||
|
||||
def test_as_agent_wraps_assistant(self, mock_async_openai: MagicMock) -> None:
|
||||
"""Test wrapping an SDK Assistant object."""
|
||||
provider = OpenAIAssistantProvider(mock_async_openai)
|
||||
assistant = create_mock_assistant(
|
||||
assistant_id="asst_wrap123",
|
||||
name="WrappedAssistant",
|
||||
instructions="Original instructions",
|
||||
)
|
||||
|
||||
agent = provider.as_agent(assistant)
|
||||
|
||||
assert agent.id == "asst_wrap123"
|
||||
assert agent.name == "WrappedAssistant"
|
||||
# Instructions are passed to ChatOptions, not exposed as attribute
|
||||
assert isinstance(agent, Agent)
|
||||
|
||||
def test_as_agent_with_instructions_override(self, mock_async_openai: MagicMock) -> None:
|
||||
"""Test as_agent with instruction override."""
|
||||
provider = OpenAIAssistantProvider(mock_async_openai)
|
||||
assistant = create_mock_assistant(instructions="Original")
|
||||
|
||||
agent = provider.as_agent(assistant, instructions="Override")
|
||||
|
||||
# Agent should be created successfully with override instructions
|
||||
assert isinstance(agent, Agent)
|
||||
|
||||
def test_as_agent_validates_function_tools(self, mock_async_openai: MagicMock) -> None:
|
||||
"""Test that missing function tools raise ValueError."""
|
||||
provider = OpenAIAssistantProvider(mock_async_openai)
|
||||
assistant = create_mock_assistant(tools=[create_function_tool("get_weather")])
|
||||
|
||||
with pytest.raises(ValueError) as exc_info:
|
||||
provider.as_agent(assistant)
|
||||
|
||||
assert "get_weather" in str(exc_info.value)
|
||||
|
||||
def test_as_agent_with_function_tools_provided(self, mock_async_openai: MagicMock) -> None:
|
||||
"""Test as_agent with function tools provided."""
|
||||
provider = OpenAIAssistantProvider(mock_async_openai)
|
||||
assistant = create_mock_assistant(tools=[create_function_tool("get_weather")])
|
||||
|
||||
agent = provider.as_agent(assistant, tools=[get_weather])
|
||||
|
||||
assert isinstance(agent, Agent)
|
||||
|
||||
def test_as_agent_merges_hosted_tools(self, mock_async_openai: MagicMock) -> None:
|
||||
"""Test that hosted tools are merged automatically."""
|
||||
provider = OpenAIAssistantProvider(mock_async_openai)
|
||||
assistant = create_mock_assistant(tools=[create_code_interpreter_tool()])
|
||||
|
||||
agent = provider.as_agent(assistant)
|
||||
|
||||
assert isinstance(agent, Agent)
|
||||
|
||||
def test_as_agent_hosted_tools_not_required(self, mock_async_openai: MagicMock) -> None:
|
||||
"""Test that hosted tools don't require user implementations."""
|
||||
provider = OpenAIAssistantProvider(mock_async_openai)
|
||||
assistant = create_mock_assistant(tools=[create_code_interpreter_tool(), create_file_search_tool()])
|
||||
|
||||
# Should not raise - hosted tools don't need implementations
|
||||
agent = provider.as_agent(assistant)
|
||||
|
||||
assert isinstance(agent, Agent)
|
||||
|
||||
|
||||
# endregion
|
||||
|
||||
# region Tool Conversion Tests
|
||||
|
||||
|
||||
class TestToolConversion:
|
||||
"""Tests for tool conversion utilities (shared functions)."""
|
||||
|
||||
def test_to_assistant_tools_tool(self) -> None:
|
||||
"""Test FunctionTool conversion to API format."""
|
||||
|
||||
@tool
|
||||
def test_func(x: int) -> int:
|
||||
"""Test function."""
|
||||
return x
|
||||
|
||||
# Normalize tools first, then convert
|
||||
normalized = normalize_tools([test_func])
|
||||
api_tools = to_assistant_tools(normalized)
|
||||
|
||||
assert len(api_tools) == 1
|
||||
assert api_tools[0]["type"] == "function"
|
||||
assert api_tools[0]["function"]["name"] == "test_func"
|
||||
|
||||
def test_to_assistant_tools_callable(self) -> None:
|
||||
"""Test raw callable conversion via normalize_tools."""
|
||||
# normalize_tools converts callables to FunctionTool
|
||||
normalized = normalize_tools([get_weather])
|
||||
api_tools = to_assistant_tools(normalized)
|
||||
|
||||
assert len(api_tools) == 1
|
||||
assert api_tools[0]["type"] == "function"
|
||||
assert api_tools[0]["function"]["name"] == "get_weather"
|
||||
|
||||
def test_to_assistant_tools_code_interpreter(self) -> None:
|
||||
"""Test code_interpreter tool dict conversion."""
|
||||
api_tools = to_assistant_tools([OpenAIAssistantsClient.get_code_interpreter_tool()])
|
||||
|
||||
assert len(api_tools) == 1
|
||||
assert api_tools[0] == {"type": "code_interpreter"}
|
||||
|
||||
def test_to_assistant_tools_file_search(self) -> None:
|
||||
"""Test file_search tool dict conversion."""
|
||||
api_tools = to_assistant_tools([OpenAIAssistantsClient.get_file_search_tool()])
|
||||
|
||||
assert len(api_tools) == 1
|
||||
assert api_tools[0]["type"] == "file_search"
|
||||
|
||||
def test_to_assistant_tools_file_search_with_max_results(self) -> None:
|
||||
"""Test file_search tool with max_results conversion."""
|
||||
api_tools = to_assistant_tools([OpenAIAssistantsClient.get_file_search_tool(max_num_results=5)])
|
||||
|
||||
assert api_tools[0]["file_search"]["max_num_results"] == 5
|
||||
|
||||
def test_to_assistant_tools_dict(self) -> None:
|
||||
"""Test raw dict tool passthrough."""
|
||||
raw_tool = {"type": "function", "function": {"name": "custom", "description": "Custom tool"}}
|
||||
|
||||
api_tools = to_assistant_tools([raw_tool])
|
||||
|
||||
assert len(api_tools) == 1
|
||||
assert api_tools[0] == raw_tool
|
||||
|
||||
def test_to_assistant_tools_empty(self) -> None:
|
||||
"""Test conversion with no tools."""
|
||||
api_tools = to_assistant_tools(None)
|
||||
|
||||
assert api_tools == []
|
||||
|
||||
def test_from_assistant_tools_code_interpreter(self) -> None:
|
||||
"""Test converting code_interpreter tool from OpenAI format."""
|
||||
assistant_tools = [create_code_interpreter_tool()]
|
||||
|
||||
tools = from_assistant_tools(assistant_tools)
|
||||
|
||||
assert len(tools) == 1
|
||||
assert tools[0] == {"type": "code_interpreter"}
|
||||
|
||||
def test_from_assistant_tools_file_search(self) -> None:
|
||||
"""Test converting file_search tool from OpenAI format."""
|
||||
assistant_tools = [create_file_search_tool()]
|
||||
|
||||
tools = from_assistant_tools(assistant_tools)
|
||||
|
||||
assert len(tools) == 1
|
||||
assert tools[0] == {"type": "file_search"}
|
||||
|
||||
def test_from_assistant_tools_function_skipped(self) -> None:
|
||||
"""Test that function tools are skipped (no implementations)."""
|
||||
assistant_tools = [create_function_tool("test_func")]
|
||||
|
||||
tools = from_assistant_tools(assistant_tools)
|
||||
|
||||
assert len(tools) == 0 # Function tools are skipped
|
||||
|
||||
def test_from_assistant_tools_empty(self) -> None:
|
||||
"""Test conversion with no tools."""
|
||||
tools = from_assistant_tools(None)
|
||||
|
||||
assert tools == []
|
||||
|
||||
|
||||
# endregion
|
||||
|
||||
# region Tool Validation Tests
|
||||
|
||||
|
||||
class TestToolValidation:
|
||||
"""Tests for tool validation."""
|
||||
|
||||
def test_validate_missing_function_tool_raises(self, mock_async_openai: MagicMock) -> None:
|
||||
"""Test that missing function tools raise ValueError."""
|
||||
provider = OpenAIAssistantProvider(mock_async_openai)
|
||||
assistant_tools = [create_function_tool("my_function")]
|
||||
|
||||
with pytest.raises(ValueError) as exc_info:
|
||||
provider._validate_function_tools(assistant_tools, None) # type: ignore[reportPrivateUsage]
|
||||
|
||||
assert "my_function" in str(exc_info.value)
|
||||
|
||||
def test_validate_all_tools_provided_passes(self, mock_async_openai: MagicMock) -> None:
|
||||
"""Test that validation passes when all tools provided."""
|
||||
provider = OpenAIAssistantProvider(mock_async_openai)
|
||||
assistant_tools = [create_function_tool("get_weather")]
|
||||
|
||||
# Should not raise
|
||||
provider._validate_function_tools(assistant_tools, [get_weather]) # type: ignore[reportPrivateUsage]
|
||||
|
||||
def test_validate_hosted_tools_not_required(self, mock_async_openai: MagicMock) -> None:
|
||||
"""Test that hosted tools don't require implementations."""
|
||||
provider = OpenAIAssistantProvider(mock_async_openai)
|
||||
assistant_tools = [create_code_interpreter_tool(), create_file_search_tool()]
|
||||
|
||||
# Should not raise
|
||||
provider._validate_function_tools(assistant_tools, None) # type: ignore[reportPrivateUsage]
|
||||
|
||||
def test_validate_with_tool(self, mock_async_openai: MagicMock) -> None:
|
||||
"""Test validation with FunctionTool."""
|
||||
provider = OpenAIAssistantProvider(mock_async_openai)
|
||||
assistant_tools = [create_function_tool("get_weather")]
|
||||
|
||||
wrapped = tool(get_weather)
|
||||
|
||||
# Should not raise
|
||||
provider._validate_function_tools(assistant_tools, [wrapped]) # type: ignore[reportPrivateUsage]
|
||||
|
||||
def test_validate_partial_tools_raises(self, mock_async_openai: MagicMock) -> None:
|
||||
"""Test that partial tool provision raises error."""
|
||||
provider = OpenAIAssistantProvider(mock_async_openai)
|
||||
assistant_tools = [
|
||||
create_function_tool("get_weather"),
|
||||
create_function_tool("search_database"),
|
||||
]
|
||||
|
||||
with pytest.raises(ValueError) as exc_info:
|
||||
provider._validate_function_tools(assistant_tools, [get_weather]) # type: ignore[reportPrivateUsage]
|
||||
|
||||
assert "search_database" in str(exc_info.value)
|
||||
|
||||
|
||||
# endregion
|
||||
|
||||
# region Tool Merging Tests
|
||||
|
||||
|
||||
class TestToolMerging:
|
||||
"""Tests for tool merging."""
|
||||
|
||||
def test_merge_code_interpreter(self, mock_async_openai: MagicMock) -> None:
|
||||
"""Test merging code interpreter tool."""
|
||||
provider = OpenAIAssistantProvider(mock_async_openai)
|
||||
assistant_tools = [create_code_interpreter_tool()]
|
||||
|
||||
merged = provider._merge_tools(assistant_tools, None) # type: ignore[reportPrivateUsage]
|
||||
|
||||
assert len(merged) == 1
|
||||
assert merged[0] == {"type": "code_interpreter"}
|
||||
|
||||
def test_merge_file_search(self, mock_async_openai: MagicMock) -> None:
|
||||
"""Test merging file search tool."""
|
||||
provider = OpenAIAssistantProvider(mock_async_openai)
|
||||
assistant_tools = [create_file_search_tool()]
|
||||
|
||||
merged = provider._merge_tools(assistant_tools, None) # type: ignore[reportPrivateUsage]
|
||||
|
||||
assert len(merged) == 1
|
||||
assert merged[0] == {"type": "file_search"}
|
||||
|
||||
def test_merge_with_user_tools(self, mock_async_openai: MagicMock) -> None:
|
||||
"""Test merging hosted and user tools."""
|
||||
provider = OpenAIAssistantProvider(mock_async_openai)
|
||||
assistant_tools = [create_code_interpreter_tool()]
|
||||
|
||||
merged = provider._merge_tools(assistant_tools, [get_weather]) # type: ignore[reportPrivateUsage]
|
||||
|
||||
assert len(merged) == 2
|
||||
assert merged[0] == {"type": "code_interpreter"}
|
||||
|
||||
def test_merge_multiple_hosted_tools(self, mock_async_openai: MagicMock) -> None:
|
||||
"""Test merging multiple hosted tools."""
|
||||
provider = OpenAIAssistantProvider(mock_async_openai)
|
||||
assistant_tools = [create_code_interpreter_tool(), create_file_search_tool()]
|
||||
|
||||
merged = provider._merge_tools(assistant_tools, None) # type: ignore[reportPrivateUsage]
|
||||
|
||||
assert len(merged) == 2
|
||||
|
||||
def test_merge_single_user_tool(self, mock_async_openai: MagicMock) -> None:
|
||||
"""Test merging with single user tool (not list)."""
|
||||
provider = OpenAIAssistantProvider(mock_async_openai)
|
||||
assistant_tools: list[Any] = []
|
||||
|
||||
merged = provider._merge_tools(assistant_tools, get_weather) # type: ignore[reportPrivateUsage]
|
||||
|
||||
assert len(merged) == 1
|
||||
|
||||
|
||||
# endregion
|
||||
|
||||
# region Integration Tests
|
||||
|
||||
skip_if_openai_integration_tests_disabled = pytest.mark.skipif(
|
||||
os.getenv("OPENAI_API_KEY", "") in ("", "test-dummy-key"),
|
||||
reason="No real OPENAI_API_KEY provided; skipping integration tests.",
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.flaky
|
||||
@pytest.mark.integration
|
||||
@skip_if_openai_integration_tests_disabled
|
||||
class TestOpenAIAssistantProviderIntegration:
|
||||
"""Integration tests requiring real OpenAI API."""
|
||||
|
||||
async def test_create_and_run_agent(self) -> None:
|
||||
"""End-to-end test of creating and running an agent."""
|
||||
provider = OpenAIAssistantProvider()
|
||||
|
||||
agent = await provider.create_agent(
|
||||
name="IntegrationTestAgent",
|
||||
model=os.environ.get("OPENAI_MODEL", "gpt-4"),
|
||||
instructions="You are a helpful assistant. Respond briefly.",
|
||||
)
|
||||
|
||||
try:
|
||||
result = await agent.run("Say 'hello' and nothing else.")
|
||||
result_text = str(result)
|
||||
assert "hello" in result_text.lower()
|
||||
finally:
|
||||
# Clean up the assistant
|
||||
await provider._client.beta.assistants.delete(agent.id) # type: ignore[reportPrivateUsage, union-attr]
|
||||
|
||||
async def test_create_agent_with_function_tools_integration(self) -> None:
|
||||
"""Integration test with function tools."""
|
||||
provider = OpenAIAssistantProvider()
|
||||
|
||||
@tool(approval_mode="never_require")
|
||||
def get_current_time() -> str:
|
||||
"""Get the current time."""
|
||||
from datetime import datetime
|
||||
|
||||
return datetime.now().strftime("%H:%M")
|
||||
|
||||
agent = await provider.create_agent(
|
||||
name="TimeAgent",
|
||||
model=os.environ.get("OPENAI_MODEL", "gpt-4"),
|
||||
instructions="You are a helpful assistant.",
|
||||
tools=[get_current_time],
|
||||
)
|
||||
|
||||
try:
|
||||
result = await agent.run("What time is it? Use the get_current_time function.")
|
||||
result_text = str(result)
|
||||
# The response should contain time information
|
||||
assert ":" in result_text or "time" in result_text.lower()
|
||||
finally:
|
||||
await provider._client.beta.assistants.delete(agent.id) # type: ignore[reportPrivateUsage, union-attr]
|
||||
|
||||
|
||||
# endregion
|
||||
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,453 @@
|
||||
# Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import os
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
from agent_framework import Agent, AgentResponse, ChatResponse, Content, Message, SupportsChatGetResponse, tool
|
||||
from azure.identity.aio import AzureCliCredential, get_bearer_token_provider
|
||||
from openai import AsyncAzureOpenAI
|
||||
from pydantic import BaseModel
|
||||
from pytest import param
|
||||
|
||||
from agent_framework_openai import OpenAIChatClient
|
||||
|
||||
pytestmark = pytest.mark.azure
|
||||
|
||||
skip_if_azure_openai_integration_tests_disabled = pytest.mark.skipif(
|
||||
os.getenv("AZURE_OPENAI_ENDPOINT", "") in ("", "https://test-endpoint.openai.azure.com")
|
||||
or os.getenv("AZURE_OPENAI_DEPLOYMENT_NAME", "") == "",
|
||||
reason="No real Azure OpenAI endpoint or responses deployment provided; skipping integration tests.",
|
||||
)
|
||||
|
||||
|
||||
class OutputStruct(BaseModel):
|
||||
"""A structured output for testing purposes."""
|
||||
|
||||
location: str
|
||||
weather: str | None = None
|
||||
|
||||
|
||||
def _create_azure_openai_chat_client(
|
||||
*,
|
||||
api_key: Any = None,
|
||||
) -> OpenAIChatClient:
|
||||
return OpenAIChatClient(
|
||||
model=os.environ["AZURE_OPENAI_DEPLOYMENT_NAME"],
|
||||
api_key=api_key or os.environ["AZURE_OPENAI_API_KEY"],
|
||||
azure_endpoint=os.environ["AZURE_OPENAI_ENDPOINT"],
|
||||
api_version=os.getenv("AZURE_OPENAI_API_VERSION"),
|
||||
)
|
||||
|
||||
|
||||
async def create_vector_store(client: OpenAIChatClient) -> tuple[str, Content]:
|
||||
"""Create a vector store with sample documents for testing."""
|
||||
file = await client.client.files.create(
|
||||
file=("todays_weather.txt", b"The weather today is sunny with a high of 75F."),
|
||||
purpose="assistants",
|
||||
)
|
||||
vector_store = await client.client.vector_stores.create(
|
||||
name="knowledge_base",
|
||||
expires_after={"anchor": "last_active_at", "days": 1},
|
||||
)
|
||||
result = await client.client.vector_stores.files.create_and_poll(
|
||||
vector_store_id=vector_store.id,
|
||||
file_id=file.id,
|
||||
poll_interval_ms=1000,
|
||||
)
|
||||
if result.last_error is not None:
|
||||
raise RuntimeError(f"Vector store file processing failed with status: {result.last_error.message}")
|
||||
|
||||
return file.id, Content.from_hosted_vector_store(vector_store_id=vector_store.id)
|
||||
|
||||
|
||||
async def delete_vector_store(client: OpenAIChatClient, file_id: str, vector_store_id: str) -> None:
|
||||
"""Delete the vector store after tests."""
|
||||
|
||||
await client.client.vector_stores.delete(vector_store_id=vector_store_id)
|
||||
await client.client.files.delete(file_id=file_id)
|
||||
|
||||
|
||||
@tool(approval_mode="never_require")
|
||||
async def get_weather(location: str) -> str:
|
||||
"""Get the current weather in a given location."""
|
||||
return f"The current weather in {location} is sunny."
|
||||
|
||||
|
||||
def test_init_with_azure_endpoint(azure_openai_unit_test_env: dict[str, str]) -> None:
|
||||
client = _create_azure_openai_chat_client()
|
||||
|
||||
assert client.model == azure_openai_unit_test_env["AZURE_OPENAI_DEPLOYMENT_NAME"]
|
||||
assert isinstance(client, SupportsChatGetResponse)
|
||||
assert isinstance(client.client, AsyncAzureOpenAI)
|
||||
assert client.OTEL_PROVIDER_NAME == "azure.ai.openai"
|
||||
assert client.azure_endpoint == azure_openai_unit_test_env["AZURE_OPENAI_ENDPOINT"]
|
||||
assert client.api_version == azure_openai_unit_test_env["AZURE_OPENAI_API_VERSION"]
|
||||
|
||||
|
||||
def test_init_auto_detects_azure_env(azure_openai_unit_test_env: dict[str, str]) -> None:
|
||||
client = OpenAIChatClient()
|
||||
|
||||
assert client.model == azure_openai_unit_test_env["AZURE_OPENAI_DEPLOYMENT_NAME"]
|
||||
assert isinstance(client.client, AsyncAzureOpenAI)
|
||||
assert client.azure_endpoint == azure_openai_unit_test_env["AZURE_OPENAI_ENDPOINT"]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("exclude_list", [["AZURE_OPENAI_API_VERSION"]], indirect=True)
|
||||
def test_init_uses_default_azure_api_version(azure_openai_unit_test_env: dict[str, str]) -> None:
|
||||
client = _create_azure_openai_chat_client()
|
||||
|
||||
assert client.model == azure_openai_unit_test_env["AZURE_OPENAI_DEPLOYMENT_NAME"]
|
||||
assert client.api_version == "preview"
|
||||
|
||||
|
||||
def test_openai_base_url_wins_over_azure_aliases(monkeypatch, azure_openai_unit_test_env: dict[str, str]) -> None:
|
||||
monkeypatch.setenv("OPENAI_API_KEY", "test-dummy-key")
|
||||
monkeypatch.setenv("OPENAI_MODEL", "gpt-5")
|
||||
monkeypatch.setenv("OPENAI_BASE_URL", "https://custom-openai-endpoint.com/v1")
|
||||
|
||||
client = OpenAIChatClient()
|
||||
|
||||
assert client.model == "gpt-5"
|
||||
assert not isinstance(client.client, AsyncAzureOpenAI)
|
||||
assert client.azure_endpoint is None
|
||||
|
||||
|
||||
@pytest.mark.flaky
|
||||
@pytest.mark.integration
|
||||
@skip_if_azure_openai_integration_tests_disabled
|
||||
@pytest.mark.parametrize(
|
||||
"option_name,option_value,needs_validation",
|
||||
[
|
||||
param("temperature", 0.7, False, id="temperature"),
|
||||
param("top_p", 0.9, False, id="top_p"),
|
||||
param("max_tokens", 500, False, id="max_tokens"),
|
||||
param("seed", 123, False, id="seed"),
|
||||
param("user", "test-user-id", False, id="user"),
|
||||
param("metadata", {"test_key": "test_value"}, False, id="metadata"),
|
||||
param("frequency_penalty", 0.5, False, id="frequency_penalty"),
|
||||
param("presence_penalty", 0.3, False, id="presence_penalty"),
|
||||
param("stop", ["END"], False, id="stop"),
|
||||
param("allow_multiple_tool_calls", True, False, id="allow_multiple_tool_calls"),
|
||||
param("tool_choice", "none", True, id="tool_choice_none"),
|
||||
param("safety_identifier", "user-hash-abc123", False, id="safety_identifier"),
|
||||
param("truncation", "auto", False, id="truncation"),
|
||||
param("top_logprobs", 5, False, id="top_logprobs"),
|
||||
param("prompt_cache_key", "test-cache-key", False, id="prompt_cache_key"),
|
||||
param("max_tool_calls", 3, False, id="max_tool_calls"),
|
||||
param("tools", [get_weather], True, id="tools_function"),
|
||||
param("tool_choice", "auto", True, id="tool_choice_auto"),
|
||||
param(
|
||||
"tool_choice",
|
||||
{"mode": "required", "required_function_name": "get_weather"},
|
||||
True,
|
||||
id="tool_choice_required",
|
||||
),
|
||||
param("response_format", OutputStruct, True, id="response_format_pydantic"),
|
||||
param(
|
||||
"response_format",
|
||||
{
|
||||
"type": "json_schema",
|
||||
"json_schema": {
|
||||
"name": "WeatherDigest",
|
||||
"strict": True,
|
||||
"schema": {
|
||||
"title": "WeatherDigest",
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"location": {"type": "string"},
|
||||
"conditions": {"type": "string"},
|
||||
"temperature_c": {"type": "number"},
|
||||
"advisory": {"type": "string"},
|
||||
},
|
||||
"required": ["location", "conditions", "temperature_c", "advisory"],
|
||||
"additionalProperties": False,
|
||||
},
|
||||
},
|
||||
},
|
||||
True,
|
||||
id="response_format_runtime_json_schema",
|
||||
),
|
||||
],
|
||||
)
|
||||
async def test_integration_options(
|
||||
option_name: str,
|
||||
option_value: Any,
|
||||
needs_validation: bool,
|
||||
) -> None:
|
||||
async with AzureCliCredential() as credential:
|
||||
client = _create_azure_openai_chat_client(
|
||||
api_key=get_bearer_token_provider(credential, "https://cognitiveservices.azure.com/.default")
|
||||
)
|
||||
client.function_invocation_configuration["max_iterations"] = 2
|
||||
|
||||
for streaming in [False, True]:
|
||||
if option_name in {"tools", "tool_choice"}:
|
||||
messages = [Message(role="user", text="What is the weather in Seattle?")]
|
||||
elif option_name == "response_format":
|
||||
messages = [
|
||||
Message(role="user", text="The weather in Seattle is sunny"),
|
||||
Message(role="user", text="What is the weather in Seattle?"),
|
||||
]
|
||||
else:
|
||||
messages = [Message(role="user", text="Say 'Hello World' briefly.")]
|
||||
|
||||
options: dict[str, Any] = {option_name: option_value}
|
||||
if option_name == "tool_choice":
|
||||
options["tools"] = [get_weather]
|
||||
|
||||
if streaming:
|
||||
response = await client.get_response(
|
||||
messages=messages,
|
||||
stream=True,
|
||||
options=options,
|
||||
).get_final_response()
|
||||
else:
|
||||
response = await client.get_response(messages=messages, options=options)
|
||||
|
||||
assert isinstance(response, ChatResponse)
|
||||
assert response.text is not None
|
||||
assert len(response.text) > 0
|
||||
|
||||
if needs_validation:
|
||||
if option_name in {"tools", "tool_choice"}:
|
||||
text = response.text.lower()
|
||||
assert "sunny" in text or "seattle" in text
|
||||
elif option_name == "response_format":
|
||||
if option_value == OutputStruct:
|
||||
assert response.value is not None
|
||||
assert isinstance(response.value, OutputStruct)
|
||||
assert "seattle" in response.value.location.lower()
|
||||
else:
|
||||
assert response.value is None
|
||||
response_value = json.loads(response.text)
|
||||
assert isinstance(response_value, dict)
|
||||
assert "location" in response_value
|
||||
assert "seattle" in response_value["location"].lower()
|
||||
|
||||
|
||||
@pytest.mark.flaky
|
||||
@pytest.mark.integration
|
||||
@skip_if_azure_openai_integration_tests_disabled
|
||||
async def test_integration_web_search() -> None:
|
||||
async with AzureCliCredential() as credential:
|
||||
client = _create_azure_openai_chat_client(
|
||||
api_key=get_bearer_token_provider(credential, "https://cognitiveservices.azure.com/.default")
|
||||
)
|
||||
|
||||
for streaming in [False, True]:
|
||||
content = {
|
||||
"messages": [
|
||||
Message(
|
||||
role="user",
|
||||
text="Who are the main characters of Kpop Demon Hunters? Do a web search to find the answer.",
|
||||
)
|
||||
],
|
||||
"options": {
|
||||
"tool_choice": "auto",
|
||||
"tools": [OpenAIChatClient.get_web_search_tool()],
|
||||
},
|
||||
"stream": streaming,
|
||||
}
|
||||
if streaming:
|
||||
response = await client.get_response(**content).get_final_response()
|
||||
else:
|
||||
response = await client.get_response(**content)
|
||||
|
||||
assert isinstance(response, ChatResponse)
|
||||
assert "Rumi" in response.text
|
||||
assert "Mira" in response.text
|
||||
assert "Zoey" in response.text
|
||||
|
||||
content = {
|
||||
"messages": [
|
||||
Message(
|
||||
role="user",
|
||||
text="What is the current weather? Do not ask for my current location.",
|
||||
)
|
||||
],
|
||||
"options": {
|
||||
"tool_choice": "auto",
|
||||
"tools": [OpenAIChatClient.get_web_search_tool(user_location={"country": "US", "city": "Seattle"})],
|
||||
},
|
||||
"stream": streaming,
|
||||
}
|
||||
if streaming:
|
||||
response = await client.get_response(**content).get_final_response()
|
||||
else:
|
||||
response = await client.get_response(**content)
|
||||
assert response.text is not None
|
||||
|
||||
|
||||
@pytest.mark.flaky
|
||||
@pytest.mark.integration
|
||||
@skip_if_azure_openai_integration_tests_disabled
|
||||
async def test_integration_client_file_search() -> None:
|
||||
async with AzureCliCredential() as credential:
|
||||
client = _create_azure_openai_chat_client(
|
||||
api_key=get_bearer_token_provider(credential, "https://cognitiveservices.azure.com/.default")
|
||||
)
|
||||
file_id, vector_store = await create_vector_store(client)
|
||||
try:
|
||||
response = await client.get_response(
|
||||
messages=[Message(role="user", text="What is the weather today? Do a file search to find the answer.")],
|
||||
options={
|
||||
"tools": [OpenAIChatClient.get_file_search_tool(vector_store_ids=[vector_store.vector_store_id])],
|
||||
"tool_choice": "auto",
|
||||
},
|
||||
)
|
||||
|
||||
assert "sunny" in response.text.lower()
|
||||
assert "75" in response.text
|
||||
finally:
|
||||
await delete_vector_store(client, file_id, vector_store.vector_store_id)
|
||||
|
||||
|
||||
@pytest.mark.flaky
|
||||
@pytest.mark.integration
|
||||
@skip_if_azure_openai_integration_tests_disabled
|
||||
async def test_integration_client_file_search_streaming() -> None:
|
||||
async with AzureCliCredential() as credential:
|
||||
client = _create_azure_openai_chat_client(
|
||||
api_key=get_bearer_token_provider(credential, "https://cognitiveservices.azure.com/.default")
|
||||
)
|
||||
file_id, vector_store = await create_vector_store(client)
|
||||
try:
|
||||
response_stream = client.get_response(
|
||||
messages=[Message(role="user", text="What is the weather today? Do a file search to find the answer.")],
|
||||
stream=True,
|
||||
options={
|
||||
"tools": [OpenAIChatClient.get_file_search_tool(vector_store_ids=[vector_store.vector_store_id])],
|
||||
"tool_choice": "auto",
|
||||
},
|
||||
)
|
||||
|
||||
full_response = await response_stream.get_final_response()
|
||||
assert "sunny" in full_response.text.lower()
|
||||
assert "75" in full_response.text
|
||||
finally:
|
||||
await delete_vector_store(client, file_id, vector_store.vector_store_id)
|
||||
|
||||
|
||||
@pytest.mark.flaky
|
||||
@pytest.mark.integration
|
||||
@skip_if_azure_openai_integration_tests_disabled
|
||||
async def test_integration_client_agent_hosted_mcp_tool() -> None:
|
||||
async with AzureCliCredential() as credential:
|
||||
client = _create_azure_openai_chat_client(
|
||||
api_key=get_bearer_token_provider(credential, "https://cognitiveservices.azure.com/.default")
|
||||
)
|
||||
response = await client.get_response(
|
||||
messages=[Message(role="user", text="How to create an Azure storage account using az cli?")],
|
||||
options={
|
||||
"max_tokens": 5000,
|
||||
"tools": OpenAIChatClient.get_mcp_tool(
|
||||
name="Microsoft Learn MCP",
|
||||
url="https://learn.microsoft.com/api/mcp",
|
||||
),
|
||||
},
|
||||
)
|
||||
|
||||
assert isinstance(response, ChatResponse)
|
||||
if not response.text:
|
||||
pytest.skip("MCP server returned empty response - service-side issue")
|
||||
assert any(term in response.text.lower() for term in ["azure", "storage", "account", "cli"])
|
||||
|
||||
|
||||
@pytest.mark.flaky
|
||||
@pytest.mark.integration
|
||||
@skip_if_azure_openai_integration_tests_disabled
|
||||
async def test_integration_client_agent_hosted_code_interpreter_tool() -> None:
|
||||
async with AzureCliCredential() as credential:
|
||||
client = _create_azure_openai_chat_client(
|
||||
api_key=get_bearer_token_provider(credential, "https://cognitiveservices.azure.com/.default")
|
||||
)
|
||||
|
||||
response = await client.get_response(
|
||||
messages=[Message(role="user", text="Calculate the sum of numbers from 1 to 10 using Python code.")],
|
||||
options={"tools": [OpenAIChatClient.get_code_interpreter_tool()]},
|
||||
)
|
||||
|
||||
contains_relevant_content = any(
|
||||
term in response.text.lower() for term in ["55", "sum", "code", "python", "calculate", "10"]
|
||||
)
|
||||
assert contains_relevant_content or len(response.text.strip()) > 10
|
||||
|
||||
|
||||
@pytest.mark.flaky
|
||||
@pytest.mark.integration
|
||||
@skip_if_azure_openai_integration_tests_disabled
|
||||
async def test_integration_client_agent_existing_session() -> None:
|
||||
async with AzureCliCredential() as credential:
|
||||
preserved_session = None
|
||||
|
||||
async with Agent(
|
||||
client=_create_azure_openai_chat_client(
|
||||
api_key=get_bearer_token_provider(credential, "https://cognitiveservices.azure.com/.default")
|
||||
),
|
||||
instructions="You are a helpful assistant with good memory.",
|
||||
) as first_agent:
|
||||
session = first_agent.create_session()
|
||||
first_response = await first_agent.run(
|
||||
"My hobby is photography. Remember this.",
|
||||
session=session,
|
||||
store=True,
|
||||
)
|
||||
|
||||
assert isinstance(first_response, AgentResponse)
|
||||
preserved_session = session
|
||||
|
||||
if preserved_session:
|
||||
async with Agent(
|
||||
client=_create_azure_openai_chat_client(
|
||||
api_key=get_bearer_token_provider(credential, "https://cognitiveservices.azure.com/.default")
|
||||
),
|
||||
instructions="You are a helpful assistant with good memory.",
|
||||
) as second_agent:
|
||||
second_response = await second_agent.run("What is my hobby?", session=preserved_session)
|
||||
|
||||
assert isinstance(second_response, AgentResponse)
|
||||
assert second_response.text is not None
|
||||
assert "photography" in second_response.text.lower()
|
||||
|
||||
|
||||
@pytest.mark.flaky
|
||||
@pytest.mark.integration
|
||||
@skip_if_azure_openai_integration_tests_disabled
|
||||
async def test_azure_openai_chat_client_tool_rich_content_image() -> None:
|
||||
image_path = Path(__file__).parent.parent / "assets" / "sample_image.jpg"
|
||||
image_bytes = image_path.read_bytes()
|
||||
|
||||
@tool(approval_mode="never_require")
|
||||
def get_test_image() -> Content:
|
||||
"""Return a test image for analysis."""
|
||||
return Content.from_data(data=image_bytes, media_type="image/jpeg")
|
||||
|
||||
async with AzureCliCredential() as credential:
|
||||
client = _create_azure_openai_chat_client(
|
||||
api_key=get_bearer_token_provider(credential, "https://cognitiveservices.azure.com/.default")
|
||||
)
|
||||
client.function_invocation_configuration["max_iterations"] = 2
|
||||
|
||||
for streaming in [False, True]:
|
||||
messages = [Message(role="user", text="Call the get_test_image tool and describe what you see.")]
|
||||
options: dict[str, Any] = {"tools": [get_test_image], "tool_choice": "auto"}
|
||||
|
||||
if streaming:
|
||||
response = await client.get_response(
|
||||
messages=messages,
|
||||
stream=True,
|
||||
options=options,
|
||||
).get_final_response()
|
||||
else:
|
||||
response = await client.get_response(messages=messages, options=options)
|
||||
|
||||
assert isinstance(response, ChatResponse)
|
||||
assert response.text is not None
|
||||
assert "house" in response.text.lower(), (
|
||||
f"Model did not describe the house image. Response: {response.text}"
|
||||
)
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,335 @@
|
||||
# Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from collections.abc import Awaitable, Callable
|
||||
|
||||
import pytest
|
||||
from agent_framework import (
|
||||
Agent,
|
||||
AgentResponse,
|
||||
AgentResponseUpdate,
|
||||
ChatResponse,
|
||||
ChatResponseUpdate,
|
||||
Message,
|
||||
SupportsChatGetResponse,
|
||||
tool,
|
||||
)
|
||||
from azure.identity.aio import AzureCliCredential, get_bearer_token_provider
|
||||
from openai import AsyncAzureOpenAI
|
||||
|
||||
from agent_framework_openai import OpenAIChatCompletionClient
|
||||
|
||||
pytestmark = pytest.mark.azure
|
||||
|
||||
skip_if_azure_openai_integration_tests_disabled = pytest.mark.skipif(
|
||||
os.getenv("AZURE_OPENAI_ENDPOINT", "") in ("", "https://test-endpoint.openai.azure.com")
|
||||
or os.getenv("AZURE_OPENAI_DEPLOYMENT_NAME", "") == "",
|
||||
reason="No real Azure OpenAI endpoint or chat deployment provided; skipping integration tests.",
|
||||
)
|
||||
|
||||
|
||||
def _create_azure_chat_completion_client(
|
||||
*,
|
||||
api_key: str | Callable[[], str | Awaitable[str]] | None = None,
|
||||
) -> OpenAIChatCompletionClient:
|
||||
return OpenAIChatCompletionClient(
|
||||
model=os.environ["AZURE_OPENAI_DEPLOYMENT_NAME"],
|
||||
api_key=api_key or os.environ["AZURE_OPENAI_API_KEY"],
|
||||
azure_endpoint=os.environ["AZURE_OPENAI_ENDPOINT"],
|
||||
api_version=os.getenv("AZURE_OPENAI_API_VERSION"),
|
||||
)
|
||||
|
||||
|
||||
@tool(approval_mode="never_require")
|
||||
def get_story_text() -> str:
|
||||
"""Returns a story about Emily and David."""
|
||||
return (
|
||||
"Emily and David, two passionate scientists, met during a research expedition to Antarctica. "
|
||||
"Bonded by their love for the natural world and shared curiosity, they uncovered a "
|
||||
"groundbreaking phenomenon in glaciology that could potentially reshape our understanding "
|
||||
"of climate change."
|
||||
)
|
||||
|
||||
|
||||
@tool(approval_mode="never_require")
|
||||
async def get_weather(location: str) -> str:
|
||||
"""Get the current weather in a given location."""
|
||||
return f"The current weather in {location} is sunny, 72F."
|
||||
|
||||
|
||||
def test_init_with_azure_endpoint(azure_openai_unit_test_env: dict[str, str]) -> None:
|
||||
client = _create_azure_chat_completion_client()
|
||||
|
||||
assert client.model == azure_openai_unit_test_env["AZURE_OPENAI_DEPLOYMENT_NAME"]
|
||||
assert isinstance(client, SupportsChatGetResponse)
|
||||
assert isinstance(client.client, AsyncAzureOpenAI)
|
||||
assert client.OTEL_PROVIDER_NAME == "azure.ai.openai"
|
||||
assert client.azure_endpoint == azure_openai_unit_test_env["AZURE_OPENAI_ENDPOINT"]
|
||||
assert client.api_version == azure_openai_unit_test_env["AZURE_OPENAI_API_VERSION"]
|
||||
|
||||
|
||||
def test_init_auto_detects_azure_env(azure_openai_unit_test_env: dict[str, str]) -> None:
|
||||
client = OpenAIChatCompletionClient()
|
||||
|
||||
assert client.model == azure_openai_unit_test_env["AZURE_OPENAI_DEPLOYMENT_NAME"]
|
||||
assert isinstance(client.client, AsyncAzureOpenAI)
|
||||
assert client.azure_endpoint == azure_openai_unit_test_env["AZURE_OPENAI_ENDPOINT"]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("exclude_list", [["AZURE_OPENAI_API_VERSION"]], indirect=True)
|
||||
def test_init_uses_default_azure_api_version(azure_openai_unit_test_env: dict[str, str]) -> None:
|
||||
client = _create_azure_chat_completion_client()
|
||||
|
||||
assert client.model == azure_openai_unit_test_env["AZURE_OPENAI_DEPLOYMENT_NAME"]
|
||||
assert client.api_version == "2024-10-21"
|
||||
|
||||
|
||||
def test_openai_base_url_wins_over_azure_aliases(monkeypatch, azure_openai_unit_test_env: dict[str, str]) -> None:
|
||||
monkeypatch.setenv("OPENAI_API_KEY", "test-dummy-key")
|
||||
monkeypatch.setenv("OPENAI_MODEL", "gpt-5")
|
||||
monkeypatch.setenv("OPENAI_BASE_URL", "https://custom-openai-endpoint.com/v1")
|
||||
|
||||
client = OpenAIChatCompletionClient()
|
||||
|
||||
assert client.model == "gpt-5"
|
||||
assert not isinstance(client.client, AsyncAzureOpenAI)
|
||||
assert client.azure_endpoint is None
|
||||
|
||||
|
||||
@pytest.mark.flaky
|
||||
@pytest.mark.integration
|
||||
@skip_if_azure_openai_integration_tests_disabled
|
||||
async def test_azure_openai_chat_completion_client_response() -> None:
|
||||
async with AzureCliCredential() as credential:
|
||||
client = _create_azure_chat_completion_client(
|
||||
api_key=get_bearer_token_provider(credential, "https://cognitiveservices.azure.com/.default")
|
||||
)
|
||||
assert isinstance(client, SupportsChatGetResponse)
|
||||
|
||||
messages = [
|
||||
Message(
|
||||
role="user",
|
||||
text=(
|
||||
"Emily and David, two passionate scientists, met during a research expedition to Antarctica. "
|
||||
"Bonded by their love for the natural world and shared curiosity, they uncovered a "
|
||||
"groundbreaking phenomenon in glaciology that could potentially reshape our understanding "
|
||||
"of climate change."
|
||||
),
|
||||
),
|
||||
Message(role="user", text="who are Emily and David?"),
|
||||
]
|
||||
|
||||
response = await client.get_response(messages=messages)
|
||||
|
||||
assert response is not None
|
||||
assert isinstance(response, ChatResponse)
|
||||
assert any(
|
||||
word in response.text.lower() for word in ["scientists", "research", "antarctica", "glaciology", "climate"]
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.flaky
|
||||
@pytest.mark.integration
|
||||
@skip_if_azure_openai_integration_tests_disabled
|
||||
async def test_azure_openai_chat_completion_client_response_tools() -> None:
|
||||
async with AzureCliCredential() as credential:
|
||||
client = _create_azure_chat_completion_client(
|
||||
api_key=get_bearer_token_provider(credential, "https://cognitiveservices.azure.com/.default")
|
||||
)
|
||||
|
||||
response = await client.get_response(
|
||||
messages=[Message(role="user", text="who are Emily and David?")],
|
||||
options={"tools": [get_story_text], "tool_choice": "auto"},
|
||||
)
|
||||
|
||||
assert response is not None
|
||||
assert isinstance(response, ChatResponse)
|
||||
assert "Emily" in response.text or "David" in response.text
|
||||
|
||||
|
||||
@pytest.mark.flaky
|
||||
@pytest.mark.integration
|
||||
@skip_if_azure_openai_integration_tests_disabled
|
||||
async def test_azure_openai_chat_completion_client_streaming() -> None:
|
||||
async with AzureCliCredential() as credential:
|
||||
client = _create_azure_chat_completion_client(
|
||||
api_key=get_bearer_token_provider(credential, "https://cognitiveservices.azure.com/.default")
|
||||
)
|
||||
|
||||
response = client.get_response(
|
||||
messages=[
|
||||
Message(
|
||||
role="user",
|
||||
text=(
|
||||
"Emily and David, two passionate scientists, met during a research expedition to Antarctica. "
|
||||
"Bonded by their love for the natural world and shared curiosity, they uncovered a "
|
||||
"groundbreaking phenomenon in glaciology that could potentially reshape our understanding "
|
||||
"of climate change."
|
||||
),
|
||||
),
|
||||
Message(role="user", text="who are Emily and David?"),
|
||||
],
|
||||
stream=True,
|
||||
)
|
||||
|
||||
full_message = ""
|
||||
async for chunk in response:
|
||||
assert isinstance(chunk, ChatResponseUpdate)
|
||||
assert chunk.message_id is not None
|
||||
assert chunk.response_id is not None
|
||||
for content in chunk.contents:
|
||||
if content.type == "text" and content.text:
|
||||
full_message += content.text
|
||||
|
||||
assert "Emily" in full_message or "David" in full_message
|
||||
|
||||
|
||||
@pytest.mark.flaky
|
||||
@pytest.mark.integration
|
||||
@skip_if_azure_openai_integration_tests_disabled
|
||||
async def test_azure_openai_chat_completion_client_streaming_tools() -> None:
|
||||
async with AzureCliCredential() as credential:
|
||||
client = _create_azure_chat_completion_client(
|
||||
api_key=get_bearer_token_provider(credential, "https://cognitiveservices.azure.com/.default")
|
||||
)
|
||||
|
||||
response = client.get_response(
|
||||
messages=[Message(role="user", text="who are Emily and David?")],
|
||||
stream=True,
|
||||
options={"tools": [get_story_text], "tool_choice": "auto"},
|
||||
)
|
||||
|
||||
full_message = ""
|
||||
async for chunk in response:
|
||||
assert isinstance(chunk, ChatResponseUpdate)
|
||||
for content in chunk.contents:
|
||||
if content.type == "text" and content.text:
|
||||
full_message += content.text
|
||||
|
||||
assert "Emily" in full_message or "David" in full_message
|
||||
|
||||
|
||||
@pytest.mark.flaky
|
||||
@pytest.mark.integration
|
||||
@skip_if_azure_openai_integration_tests_disabled
|
||||
async def test_azure_openai_chat_completion_client_agent_basic_run() -> None:
|
||||
async with (
|
||||
AzureCliCredential() as credential,
|
||||
Agent(
|
||||
client=_create_azure_chat_completion_client(
|
||||
api_key=get_bearer_token_provider(credential, "https://cognitiveservices.azure.com/.default")
|
||||
),
|
||||
) as agent,
|
||||
):
|
||||
response = await agent.run("Please respond with exactly: 'This is a response test.'")
|
||||
|
||||
assert isinstance(response, AgentResponse)
|
||||
assert response.text is not None
|
||||
assert "response test" in response.text.lower()
|
||||
|
||||
|
||||
@pytest.mark.flaky
|
||||
@pytest.mark.integration
|
||||
@skip_if_azure_openai_integration_tests_disabled
|
||||
async def test_azure_openai_chat_completion_client_agent_basic_run_streaming() -> None:
|
||||
async with (
|
||||
AzureCliCredential() as credential,
|
||||
Agent(
|
||||
client=_create_azure_chat_completion_client(
|
||||
api_key=get_bearer_token_provider(credential, "https://cognitiveservices.azure.com/.default")
|
||||
),
|
||||
) as agent,
|
||||
):
|
||||
full_text = ""
|
||||
async for chunk in agent.run(
|
||||
"Please respond with exactly: 'This is a streaming response test.'",
|
||||
stream=True,
|
||||
):
|
||||
assert isinstance(chunk, AgentResponseUpdate)
|
||||
if chunk.text:
|
||||
full_text += chunk.text
|
||||
|
||||
assert "streaming response test" in full_text.lower()
|
||||
|
||||
|
||||
@pytest.mark.flaky
|
||||
@pytest.mark.integration
|
||||
@skip_if_azure_openai_integration_tests_disabled
|
||||
async def test_azure_openai_chat_completion_client_agent_session_persistence() -> None:
|
||||
async with (
|
||||
AzureCliCredential() as credential,
|
||||
Agent(
|
||||
client=_create_azure_chat_completion_client(
|
||||
api_key=get_bearer_token_provider(credential, "https://cognitiveservices.azure.com/.default")
|
||||
),
|
||||
instructions="You are a helpful assistant with good memory.",
|
||||
) as agent,
|
||||
):
|
||||
session = agent.create_session()
|
||||
response1 = await agent.run("My name is Alice. Remember this.", session=session)
|
||||
response2 = await agent.run("What is my name?", session=session)
|
||||
|
||||
assert isinstance(response1, AgentResponse)
|
||||
assert isinstance(response2, AgentResponse)
|
||||
assert response2.text is not None
|
||||
assert "alice" in response2.text.lower()
|
||||
|
||||
|
||||
@pytest.mark.flaky
|
||||
@pytest.mark.integration
|
||||
@skip_if_azure_openai_integration_tests_disabled
|
||||
async def test_azure_openai_chat_completion_client_agent_existing_session() -> None:
|
||||
async with AzureCliCredential() as credential:
|
||||
preserved_session = None
|
||||
|
||||
async with Agent(
|
||||
client=_create_azure_chat_completion_client(
|
||||
api_key=get_bearer_token_provider(credential, "https://cognitiveservices.azure.com/.default")
|
||||
),
|
||||
instructions="You are a helpful assistant with good memory.",
|
||||
) as first_agent:
|
||||
session = first_agent.create_session()
|
||||
first_response = await first_agent.run("My name is Alice. Remember this.", session=session)
|
||||
|
||||
assert isinstance(first_response, AgentResponse)
|
||||
preserved_session = session
|
||||
|
||||
if preserved_session:
|
||||
async with Agent(
|
||||
client=_create_azure_chat_completion_client(
|
||||
api_key=get_bearer_token_provider(credential, "https://cognitiveservices.azure.com/.default")
|
||||
),
|
||||
instructions="You are a helpful assistant with good memory.",
|
||||
) as second_agent:
|
||||
second_response = await second_agent.run("What is my name?", session=preserved_session)
|
||||
|
||||
assert isinstance(second_response, AgentResponse)
|
||||
assert second_response.text is not None
|
||||
assert "alice" in second_response.text.lower()
|
||||
|
||||
|
||||
@pytest.mark.flaky
|
||||
@pytest.mark.integration
|
||||
@skip_if_azure_openai_integration_tests_disabled
|
||||
async def test_azure_chat_completion_client_agent_level_tool_persistence() -> None:
|
||||
async with (
|
||||
AzureCliCredential() as credential,
|
||||
Agent(
|
||||
client=_create_azure_chat_completion_client(
|
||||
api_key=get_bearer_token_provider(credential, "https://cognitiveservices.azure.com/.default")
|
||||
),
|
||||
instructions="You are a helpful assistant that uses available tools.",
|
||||
tools=[get_weather],
|
||||
) as agent,
|
||||
):
|
||||
first_response = await agent.run("What's the weather like in Chicago?")
|
||||
second_response = await agent.run("What's the weather in Miami?")
|
||||
|
||||
assert isinstance(first_response, AgentResponse)
|
||||
assert isinstance(second_response, AgentResponse)
|
||||
assert first_response.text is not None
|
||||
assert second_response.text is not None
|
||||
assert any(term in first_response.text.lower() for term in ["chicago", "sunny", "72"])
|
||||
assert any(term in second_response.text.lower() for term in ["miami", "sunny", "72"])
|
||||
@@ -0,0 +1,425 @@
|
||||
# Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
from copy import deepcopy
|
||||
from datetime import datetime, timezone
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from agent_framework import ChatResponseUpdate, Message
|
||||
from agent_framework.exceptions import ChatClientException
|
||||
from openai import AsyncStream
|
||||
from openai.resources.chat.completions import AsyncCompletions as AsyncChatCompletions
|
||||
from openai.types.chat import ChatCompletion, ChatCompletionChunk
|
||||
from openai.types.chat.chat_completion import Choice
|
||||
from openai.types.chat.chat_completion_chunk import Choice as ChunkChoice
|
||||
from openai.types.chat.chat_completion_chunk import ChoiceDelta as ChunkChoiceDelta
|
||||
from openai.types.chat.chat_completion_message import ChatCompletionMessage
|
||||
from pydantic import BaseModel
|
||||
|
||||
from agent_framework_openai import OpenAIChatCompletionClient
|
||||
|
||||
|
||||
async def mock_async_process_chat_stream_response(_):
|
||||
mock_content = MagicMock(spec=ChatResponseUpdate)
|
||||
yield mock_content, None
|
||||
|
||||
|
||||
@pytest.fixture(scope="function")
|
||||
def chat_history() -> list[Message]:
|
||||
return []
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_chat_completion_response() -> ChatCompletion:
|
||||
return ChatCompletion(
|
||||
id="test_id",
|
||||
choices=[
|
||||
Choice(index=0, message=ChatCompletionMessage(content="test", role="assistant"), finish_reason="stop")
|
||||
],
|
||||
created=0,
|
||||
model="test",
|
||||
object="chat.completion",
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_streaming_chat_completion_response() -> AsyncStream[ChatCompletionChunk]:
|
||||
content = ChatCompletionChunk(
|
||||
id="test_id",
|
||||
choices=[ChunkChoice(index=0, delta=ChunkChoiceDelta(content="test", role="assistant"), finish_reason="stop")],
|
||||
created=0,
|
||||
model="test",
|
||||
object="chat.completion.chunk",
|
||||
)
|
||||
stream = MagicMock(spec=AsyncStream)
|
||||
stream.__aiter__.return_value = [content]
|
||||
return stream
|
||||
|
||||
|
||||
# region Chat Message Content
|
||||
|
||||
|
||||
@patch.object(AsyncChatCompletions, "create", new_callable=AsyncMock)
|
||||
async def test_cmc(
|
||||
mock_create: AsyncMock,
|
||||
chat_history: list[Message],
|
||||
mock_chat_completion_response: ChatCompletion,
|
||||
openai_unit_test_env: dict[str, str],
|
||||
):
|
||||
mock_create.return_value = mock_chat_completion_response
|
||||
chat_history.append(Message(role="user", text="hello world"))
|
||||
|
||||
openai_chat_completion = OpenAIChatCompletionClient()
|
||||
await openai_chat_completion.get_response(messages=chat_history)
|
||||
mock_create.assert_awaited_once_with(
|
||||
model=openai_unit_test_env["OPENAI_MODEL"],
|
||||
stream=False,
|
||||
messages=openai_chat_completion._prepare_messages_for_openai(chat_history), # type: ignore
|
||||
)
|
||||
|
||||
|
||||
@patch.object(AsyncChatCompletions, "create", new_callable=AsyncMock)
|
||||
async def test_cmc_chat_options(
|
||||
mock_create: AsyncMock,
|
||||
chat_history: list[Message],
|
||||
mock_chat_completion_response: ChatCompletion,
|
||||
openai_unit_test_env: dict[str, str],
|
||||
):
|
||||
mock_create.return_value = mock_chat_completion_response
|
||||
chat_history.append(Message(role="user", text="hello world"))
|
||||
|
||||
openai_chat_completion = OpenAIChatCompletionClient()
|
||||
await openai_chat_completion.get_response(
|
||||
messages=chat_history,
|
||||
)
|
||||
mock_create.assert_awaited_once_with(
|
||||
model=openai_unit_test_env["OPENAI_MODEL"],
|
||||
stream=False,
|
||||
messages=openai_chat_completion._prepare_messages_for_openai(chat_history), # type: ignore
|
||||
)
|
||||
|
||||
|
||||
@patch.object(AsyncChatCompletions, "create", new_callable=AsyncMock)
|
||||
async def test_cmc_no_fcc_in_response(
|
||||
mock_create: AsyncMock,
|
||||
chat_history: list[Message],
|
||||
mock_chat_completion_response: ChatCompletion,
|
||||
openai_unit_test_env: dict[str, str],
|
||||
):
|
||||
mock_create.return_value = mock_chat_completion_response
|
||||
chat_history.append(Message(role="user", text="hello world"))
|
||||
orig_chat_history = deepcopy(chat_history)
|
||||
|
||||
openai_chat_completion = OpenAIChatCompletionClient()
|
||||
await openai_chat_completion.get_response(
|
||||
messages=chat_history,
|
||||
)
|
||||
mock_create.assert_awaited_once_with(
|
||||
model=openai_unit_test_env["OPENAI_MODEL"],
|
||||
stream=False,
|
||||
messages=openai_chat_completion._prepare_messages_for_openai(orig_chat_history), # type: ignore
|
||||
)
|
||||
|
||||
|
||||
@patch.object(AsyncChatCompletions, "create", new_callable=AsyncMock)
|
||||
async def test_cmc_structured_output_no_fcc(
|
||||
mock_create: AsyncMock,
|
||||
chat_history: list[Message],
|
||||
mock_chat_completion_response: ChatCompletion,
|
||||
openai_unit_test_env: dict[str, str],
|
||||
):
|
||||
mock_create.return_value = mock_chat_completion_response
|
||||
chat_history.append(Message(role="user", text="hello world"))
|
||||
|
||||
# Define a mock response format
|
||||
class Test(BaseModel):
|
||||
name: str
|
||||
|
||||
openai_chat_completion = OpenAIChatCompletionClient()
|
||||
await openai_chat_completion.get_response(
|
||||
messages=chat_history,
|
||||
response_format=Test,
|
||||
)
|
||||
mock_create.assert_awaited_once()
|
||||
|
||||
|
||||
@patch.object(AsyncChatCompletions, "create", new_callable=AsyncMock)
|
||||
async def test_scmc_chat_options(
|
||||
mock_create: AsyncMock,
|
||||
chat_history: list[Message],
|
||||
mock_streaming_chat_completion_response: AsyncStream[ChatCompletionChunk],
|
||||
openai_unit_test_env: dict[str, str],
|
||||
):
|
||||
mock_create.return_value = mock_streaming_chat_completion_response
|
||||
chat_history.append(Message(role="user", text="hello world"))
|
||||
|
||||
openai_chat_completion = OpenAIChatCompletionClient()
|
||||
async for msg in openai_chat_completion.get_response(
|
||||
stream=True,
|
||||
messages=chat_history,
|
||||
):
|
||||
assert isinstance(msg, ChatResponseUpdate)
|
||||
assert msg.message_id is not None
|
||||
assert msg.response_id is not None
|
||||
mock_create.assert_awaited_once_with(
|
||||
model=openai_unit_test_env["OPENAI_MODEL"],
|
||||
stream=True,
|
||||
stream_options={"include_usage": True},
|
||||
messages=openai_chat_completion._prepare_messages_for_openai(chat_history), # type: ignore
|
||||
)
|
||||
|
||||
|
||||
@patch.object(AsyncChatCompletions, "create", new_callable=AsyncMock, side_effect=Exception)
|
||||
async def test_cmc_general_exception(
|
||||
mock_create: AsyncMock,
|
||||
chat_history: list[Message],
|
||||
mock_chat_completion_response: ChatCompletion,
|
||||
openai_unit_test_env: dict[str, str],
|
||||
):
|
||||
mock_create.return_value = mock_chat_completion_response
|
||||
chat_history.append(Message(role="user", text="hello world"))
|
||||
|
||||
openai_chat_completion = OpenAIChatCompletionClient()
|
||||
with pytest.raises(ChatClientException):
|
||||
await openai_chat_completion.get_response(
|
||||
messages=chat_history,
|
||||
)
|
||||
|
||||
|
||||
@patch.object(AsyncChatCompletions, "create", new_callable=AsyncMock)
|
||||
async def test_cmc_additional_properties(
|
||||
mock_create: AsyncMock,
|
||||
chat_history: list[Message],
|
||||
mock_chat_completion_response: ChatCompletion,
|
||||
openai_unit_test_env: dict[str, str],
|
||||
):
|
||||
mock_create.return_value = mock_chat_completion_response
|
||||
chat_history.append(Message(role="user", text="hello world"))
|
||||
|
||||
openai_chat_completion = OpenAIChatCompletionClient()
|
||||
await openai_chat_completion.get_response(messages=chat_history, options={"reasoning_effort": "low"})
|
||||
mock_create.assert_awaited_once_with(
|
||||
model=openai_unit_test_env["OPENAI_MODEL"],
|
||||
stream=False,
|
||||
messages=openai_chat_completion._prepare_messages_for_openai(chat_history), # type: ignore
|
||||
reasoning_effort="low",
|
||||
)
|
||||
|
||||
|
||||
# region Streaming
|
||||
|
||||
|
||||
@patch.object(AsyncChatCompletions, "create", new_callable=AsyncMock)
|
||||
async def test_get_streaming(
|
||||
mock_create: AsyncMock,
|
||||
chat_history: list[Message],
|
||||
openai_unit_test_env: dict[str, str],
|
||||
):
|
||||
content1 = ChatCompletionChunk(
|
||||
id="test_id",
|
||||
choices=[],
|
||||
created=0,
|
||||
model="test",
|
||||
object="chat.completion.chunk",
|
||||
)
|
||||
content2 = ChatCompletionChunk(
|
||||
id="test_id",
|
||||
choices=[ChunkChoice(index=0, delta=ChunkChoiceDelta(content="test", role="assistant"), finish_reason="stop")],
|
||||
created=0,
|
||||
model="test",
|
||||
object="chat.completion.chunk",
|
||||
)
|
||||
stream = MagicMock(spec=AsyncStream)
|
||||
stream.__aiter__.return_value = [content1, content2]
|
||||
mock_create.return_value = stream
|
||||
chat_history.append(Message(role="user", text="hello world"))
|
||||
orig_chat_history = deepcopy(chat_history)
|
||||
|
||||
openai_chat_completion = OpenAIChatCompletionClient()
|
||||
async for msg in openai_chat_completion.get_response(
|
||||
stream=True,
|
||||
messages=chat_history,
|
||||
):
|
||||
assert isinstance(msg, ChatResponseUpdate)
|
||||
mock_create.assert_awaited_once_with(
|
||||
model=openai_unit_test_env["OPENAI_MODEL"],
|
||||
stream=True,
|
||||
stream_options={"include_usage": True},
|
||||
messages=openai_chat_completion._prepare_messages_for_openai(orig_chat_history), # type: ignore
|
||||
)
|
||||
|
||||
|
||||
@patch.object(AsyncChatCompletions, "create", new_callable=AsyncMock)
|
||||
async def test_get_streaming_singular(
|
||||
mock_create: AsyncMock,
|
||||
chat_history: list[Message],
|
||||
openai_unit_test_env: dict[str, str],
|
||||
):
|
||||
content1 = ChatCompletionChunk(
|
||||
id="test_id",
|
||||
choices=[],
|
||||
created=0,
|
||||
model="test",
|
||||
object="chat.completion.chunk",
|
||||
)
|
||||
content2 = ChatCompletionChunk(
|
||||
id="test_id",
|
||||
choices=[ChunkChoice(index=0, delta=ChunkChoiceDelta(content="test", role="assistant"), finish_reason="stop")],
|
||||
created=0,
|
||||
model="test",
|
||||
object="chat.completion.chunk",
|
||||
)
|
||||
stream = MagicMock(spec=AsyncStream)
|
||||
stream.__aiter__.return_value = [content1, content2]
|
||||
mock_create.return_value = stream
|
||||
chat_history.append(Message(role="user", text="hello world"))
|
||||
orig_chat_history = deepcopy(chat_history)
|
||||
|
||||
openai_chat_completion = OpenAIChatCompletionClient()
|
||||
async for msg in openai_chat_completion.get_response(
|
||||
stream=True,
|
||||
messages=chat_history,
|
||||
):
|
||||
assert isinstance(msg, ChatResponseUpdate)
|
||||
mock_create.assert_awaited_once_with(
|
||||
model=openai_unit_test_env["OPENAI_MODEL"],
|
||||
stream=True,
|
||||
stream_options={"include_usage": True},
|
||||
messages=openai_chat_completion._prepare_messages_for_openai(orig_chat_history), # type: ignore
|
||||
)
|
||||
|
||||
|
||||
@patch.object(AsyncChatCompletions, "create", new_callable=AsyncMock)
|
||||
async def test_get_streaming_structured_output_no_fcc(
|
||||
mock_create: AsyncMock,
|
||||
chat_history: list[Message],
|
||||
openai_unit_test_env: dict[str, str],
|
||||
):
|
||||
content1 = ChatCompletionChunk(
|
||||
id="test_id",
|
||||
choices=[],
|
||||
created=0,
|
||||
model="test",
|
||||
object="chat.completion.chunk",
|
||||
)
|
||||
content2 = ChatCompletionChunk(
|
||||
id="test_id",
|
||||
choices=[ChunkChoice(index=0, delta=ChunkChoiceDelta(content="test", role="assistant"), finish_reason="stop")],
|
||||
created=0,
|
||||
model="test",
|
||||
object="chat.completion.chunk",
|
||||
)
|
||||
stream = MagicMock(spec=AsyncStream)
|
||||
stream.__aiter__.return_value = [content1, content2]
|
||||
mock_create.return_value = stream
|
||||
chat_history.append(Message(role="user", text="hello world"))
|
||||
|
||||
# Define a mock response format
|
||||
class Test(BaseModel):
|
||||
name: str
|
||||
|
||||
openai_chat_completion = OpenAIChatCompletionClient()
|
||||
async for msg in openai_chat_completion.get_response(
|
||||
stream=True,
|
||||
messages=chat_history,
|
||||
response_format=Test,
|
||||
):
|
||||
assert isinstance(msg, ChatResponseUpdate)
|
||||
mock_create.assert_awaited_once()
|
||||
|
||||
|
||||
@patch.object(AsyncChatCompletions, "create", new_callable=AsyncMock)
|
||||
async def test_get_streaming_no_fcc_in_response(
|
||||
mock_create: AsyncMock,
|
||||
chat_history: list[Message],
|
||||
mock_streaming_chat_completion_response: ChatCompletion,
|
||||
openai_unit_test_env: dict[str, str],
|
||||
):
|
||||
mock_create.return_value = mock_streaming_chat_completion_response
|
||||
chat_history.append(Message(role="user", text="hello world"))
|
||||
orig_chat_history = deepcopy(chat_history)
|
||||
|
||||
openai_chat_completion = OpenAIChatCompletionClient()
|
||||
[
|
||||
msg
|
||||
async for msg in openai_chat_completion.get_response(
|
||||
stream=True,
|
||||
messages=chat_history,
|
||||
)
|
||||
]
|
||||
mock_create.assert_awaited_once_with(
|
||||
model=openai_unit_test_env["OPENAI_MODEL"],
|
||||
stream=True,
|
||||
stream_options={"include_usage": True},
|
||||
messages=openai_chat_completion._prepare_messages_for_openai(orig_chat_history), # type: ignore
|
||||
)
|
||||
|
||||
|
||||
# region UTC Timestamp Tests
|
||||
|
||||
|
||||
def test_chat_response_created_at_uses_utc(openai_unit_test_env: dict[str, str]):
|
||||
"""Test that ChatResponse.created_at uses UTC timestamp, not local time.
|
||||
|
||||
This is a regression test for the issue where created_at was using local time
|
||||
but labeling it as UTC (with 'Z' suffix).
|
||||
"""
|
||||
# Use a specific Unix timestamp: 1733011890 = 2024-12-01T00:31:30Z (UTC)
|
||||
# This ensures we test that the timestamp is actually converted to UTC
|
||||
utc_timestamp = 1733011890
|
||||
|
||||
mock_response = ChatCompletion(
|
||||
id="test_id",
|
||||
choices=[
|
||||
Choice(index=0, message=ChatCompletionMessage(content="test", role="assistant"), finish_reason="stop")
|
||||
],
|
||||
created=utc_timestamp,
|
||||
model="test",
|
||||
object="chat.completion",
|
||||
)
|
||||
|
||||
client = OpenAIChatCompletionClient()
|
||||
response = client._parse_response_from_openai(mock_response, {})
|
||||
|
||||
# Verify that created_at is correctly formatted as UTC
|
||||
assert response.created_at is not None
|
||||
assert response.created_at.endswith("Z"), "Timestamp should end with 'Z' for UTC"
|
||||
|
||||
# Parse the timestamp and verify it matches UTC time
|
||||
expected_utc_time = datetime.fromtimestamp(utc_timestamp, tz=timezone.utc)
|
||||
expected_formatted = expected_utc_time.strftime("%Y-%m-%dT%H:%M:%S.%fZ")
|
||||
assert response.created_at == expected_formatted, (
|
||||
f"Expected UTC timestamp {expected_formatted}, got {response.created_at}"
|
||||
)
|
||||
|
||||
|
||||
def test_chat_response_update_created_at_uses_utc(openai_unit_test_env: dict[str, str]):
|
||||
"""Test that ChatResponseUpdate.created_at uses UTC timestamp, not local time.
|
||||
|
||||
This is a regression test for the issue where created_at was using local time
|
||||
but labeling it as UTC (with 'Z' suffix).
|
||||
"""
|
||||
# Use a specific Unix timestamp: 1733011890 = 2024-12-01T00:31:30Z (UTC)
|
||||
utc_timestamp = 1733011890
|
||||
|
||||
mock_chunk = ChatCompletionChunk(
|
||||
id="test_id",
|
||||
choices=[ChunkChoice(index=0, delta=ChunkChoiceDelta(content="test", role="assistant"), finish_reason="stop")],
|
||||
created=utc_timestamp,
|
||||
model="test",
|
||||
object="chat.completion.chunk",
|
||||
)
|
||||
|
||||
client = OpenAIChatCompletionClient()
|
||||
response_update = client._parse_response_update_from_openai(mock_chunk)
|
||||
|
||||
# Verify that created_at is correctly formatted as UTC
|
||||
assert response_update.created_at is not None
|
||||
assert response_update.created_at.endswith("Z"), "Timestamp should end with 'Z' for UTC"
|
||||
|
||||
# Parse the timestamp and verify it matches UTC time
|
||||
expected_utc_time = datetime.fromtimestamp(utc_timestamp, tz=timezone.utc)
|
||||
expected_formatted = expected_utc_time.strftime("%Y-%m-%dT%H:%M:%S.%fZ")
|
||||
assert response_update.created_at == expected_formatted, (
|
||||
f"Expected UTC timestamp {expected_formatted}, got {response_update.created_at}"
|
||||
)
|
||||
@@ -0,0 +1,243 @@
|
||||
# Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import pytest
|
||||
from openai.types import CreateEmbeddingResponse
|
||||
from openai.types import Embedding as OpenAIEmbedding
|
||||
from openai.types.create_embedding_response import Usage
|
||||
|
||||
from agent_framework_openai import (
|
||||
OpenAIEmbeddingClient,
|
||||
OpenAIEmbeddingOptions,
|
||||
)
|
||||
|
||||
|
||||
def _make_openai_response(
|
||||
embeddings: list[list[float]],
|
||||
model: str = "text-embedding-3-small",
|
||||
prompt_tokens: int = 5,
|
||||
total_tokens: int = 5,
|
||||
) -> CreateEmbeddingResponse:
|
||||
"""Helper to create a mock OpenAI embeddings response."""
|
||||
data = [OpenAIEmbedding(embedding=emb, index=i, object="embedding") for i, emb in enumerate(embeddings)]
|
||||
return CreateEmbeddingResponse(
|
||||
data=data,
|
||||
model=model,
|
||||
object="list",
|
||||
usage=Usage(prompt_tokens=prompt_tokens, total_tokens=total_tokens),
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def openai_unit_test_env(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
"""Set up environment variables for OpenAI embedding client."""
|
||||
monkeypatch.setenv("OPENAI_API_KEY", "test-api-key")
|
||||
monkeypatch.setenv("OPENAI_EMBEDDING_MODEL", "text-embedding-3-small")
|
||||
|
||||
|
||||
# --- OpenAI unit tests ---
|
||||
|
||||
|
||||
def test_openai_construction_with_explicit_params() -> None:
|
||||
client = OpenAIEmbeddingClient(
|
||||
model="text-embedding-3-small",
|
||||
api_key="test-key",
|
||||
)
|
||||
assert client.model == "text-embedding-3-small"
|
||||
|
||||
|
||||
def test_openai_construction_from_env(openai_unit_test_env: None) -> None:
|
||||
client = OpenAIEmbeddingClient()
|
||||
assert client.model == "text-embedding-3-small"
|
||||
|
||||
|
||||
def test_openai_construction_missing_api_key_raises(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.delenv("OPENAI_API_KEY", raising=False)
|
||||
with pytest.raises(ValueError, match="API key is required"):
|
||||
OpenAIEmbeddingClient(model="text-embedding-3-small")
|
||||
|
||||
|
||||
def test_openai_construction_missing_model_raises(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.delenv("OPENAI_EMBEDDING_MODEL", raising=False)
|
||||
with pytest.raises(ValueError, match="embedding model is required"):
|
||||
OpenAIEmbeddingClient(api_key="test-key")
|
||||
|
||||
|
||||
async def test_openai_get_embeddings(openai_unit_test_env: None) -> None:
|
||||
mock_response = _make_openai_response(
|
||||
embeddings=[[0.1, 0.2, 0.3], [0.4, 0.5, 0.6]],
|
||||
)
|
||||
client = OpenAIEmbeddingClient()
|
||||
client.client = MagicMock()
|
||||
client.client.embeddings = MagicMock()
|
||||
client.client.embeddings.create = AsyncMock(return_value=mock_response)
|
||||
|
||||
result = await client.get_embeddings(["hello", "world"])
|
||||
|
||||
assert len(result) == 2
|
||||
assert result[0].vector == [0.1, 0.2, 0.3]
|
||||
assert result[1].vector == [0.4, 0.5, 0.6]
|
||||
assert result[0].model == "text-embedding-3-small"
|
||||
assert result[0].dimensions == 3
|
||||
|
||||
|
||||
async def test_openai_get_embeddings_usage(openai_unit_test_env: None) -> None:
|
||||
mock_response = _make_openai_response(
|
||||
embeddings=[[0.1]],
|
||||
prompt_tokens=10,
|
||||
total_tokens=10,
|
||||
)
|
||||
client = OpenAIEmbeddingClient()
|
||||
client.client = MagicMock()
|
||||
client.client.embeddings = MagicMock()
|
||||
client.client.embeddings.create = AsyncMock(return_value=mock_response)
|
||||
|
||||
result = await client.get_embeddings(["test"])
|
||||
|
||||
assert result.usage is not None
|
||||
assert result.usage["input_token_count"] == 10
|
||||
assert result.usage["total_token_count"] == 10
|
||||
|
||||
|
||||
async def test_openai_options_passthrough_dimensions(openai_unit_test_env: None) -> None:
|
||||
mock_response = _make_openai_response(embeddings=[[0.1]])
|
||||
client = OpenAIEmbeddingClient()
|
||||
client.client = MagicMock()
|
||||
client.client.embeddings = MagicMock()
|
||||
client.client.embeddings.create = AsyncMock(return_value=mock_response)
|
||||
|
||||
options: OpenAIEmbeddingOptions = {"dimensions": 256}
|
||||
result = await client.get_embeddings(["test"], options=options)
|
||||
|
||||
call_kwargs = client.client.embeddings.create.call_args[1]
|
||||
assert call_kwargs["dimensions"] == 256
|
||||
assert result.options is options
|
||||
|
||||
|
||||
async def test_openai_options_passthrough_encoding_format(openai_unit_test_env: None) -> None:
|
||||
mock_response = _make_openai_response(embeddings=[[0.1]])
|
||||
client = OpenAIEmbeddingClient()
|
||||
client.client = MagicMock()
|
||||
client.client.embeddings = MagicMock()
|
||||
client.client.embeddings.create = AsyncMock(return_value=mock_response)
|
||||
|
||||
options: OpenAIEmbeddingOptions = {"encoding_format": "base64"}
|
||||
await client.get_embeddings(["test"], options=options)
|
||||
|
||||
call_kwargs = client.client.embeddings.create.call_args[1]
|
||||
assert call_kwargs["encoding_format"] == "base64"
|
||||
|
||||
|
||||
async def test_openai_base64_decoding(openai_unit_test_env: None) -> None:
|
||||
import base64
|
||||
import struct
|
||||
|
||||
# Encode [0.1, 0.2, 0.3] as base64 little-endian floats
|
||||
raw_floats = [0.1, 0.2, 0.3]
|
||||
b64_str = base64.b64encode(struct.pack(f"<{len(raw_floats)}f", *raw_floats)).decode()
|
||||
|
||||
# Mock the embedding item to return a base64 string (as the API does with encoding_format=base64)
|
||||
mock_item = MagicMock()
|
||||
mock_item.embedding = b64_str
|
||||
mock_item.index = 0
|
||||
|
||||
mock_response = MagicMock()
|
||||
mock_response.data = [mock_item]
|
||||
mock_response.model = "text-embedding-3-small"
|
||||
mock_response.usage = MagicMock(prompt_tokens=3, total_tokens=3)
|
||||
|
||||
client = OpenAIEmbeddingClient()
|
||||
client.client = MagicMock()
|
||||
client.client.embeddings = MagicMock()
|
||||
client.client.embeddings.create = AsyncMock(return_value=mock_response)
|
||||
|
||||
options: OpenAIEmbeddingOptions = {"encoding_format": "base64"}
|
||||
result = await client.get_embeddings(["test"], options=options)
|
||||
|
||||
assert len(result) == 1
|
||||
assert len(result[0].vector) == 3
|
||||
assert result[0].dimensions == 3
|
||||
for expected, actual in zip(raw_floats, result[0].vector):
|
||||
assert abs(expected - actual) < 1e-6
|
||||
|
||||
|
||||
async def test_openai_error_when_no_model_id() -> None:
|
||||
client = OpenAIEmbeddingClient.__new__(OpenAIEmbeddingClient)
|
||||
client.model = None
|
||||
client.client = MagicMock()
|
||||
client.additional_properties = {}
|
||||
client.otel_provider_name = "openai"
|
||||
|
||||
with pytest.raises(ValueError, match="model is required"):
|
||||
await client.get_embeddings(["test"])
|
||||
|
||||
|
||||
async def test_openai_empty_values_returns_empty(openai_unit_test_env: None) -> None:
|
||||
client = OpenAIEmbeddingClient()
|
||||
client.client = MagicMock()
|
||||
client.client.embeddings = MagicMock()
|
||||
client.client.embeddings.create = AsyncMock()
|
||||
|
||||
result = await client.get_embeddings([])
|
||||
|
||||
assert len(result) == 0
|
||||
assert result.usage is None
|
||||
client.client.embeddings.create.assert_not_called()
|
||||
|
||||
|
||||
# --- Integration tests ---
|
||||
|
||||
skip_if_openai_integration_tests_disabled = pytest.mark.skipif(
|
||||
os.getenv("OPENAI_API_KEY", "") in ("", "test-dummy-key"),
|
||||
reason="No real OPENAI_API_KEY provided; skipping integration tests.",
|
||||
)
|
||||
|
||||
|
||||
@skip_if_openai_integration_tests_disabled
|
||||
@pytest.mark.flaky
|
||||
@pytest.mark.integration
|
||||
async def test_integration_openai_get_embeddings() -> None:
|
||||
"""End-to-end test of OpenAI embedding generation."""
|
||||
client = OpenAIEmbeddingClient(model="text-embedding-3-small")
|
||||
|
||||
result = await client.get_embeddings(["hello world"])
|
||||
|
||||
assert len(result) == 1
|
||||
assert isinstance(result[0].vector, list)
|
||||
assert len(result[0].vector) > 0
|
||||
assert all(isinstance(v, float) for v in result[0].vector)
|
||||
assert result[0].model is not None
|
||||
assert result.usage is not None
|
||||
assert result.usage["input_token_count"] > 0
|
||||
|
||||
|
||||
@skip_if_openai_integration_tests_disabled
|
||||
@pytest.mark.flaky
|
||||
@pytest.mark.integration
|
||||
async def test_integration_openai_get_embeddings_multiple() -> None:
|
||||
"""Test embedding generation for multiple inputs."""
|
||||
client = OpenAIEmbeddingClient(model="text-embedding-3-small")
|
||||
|
||||
result = await client.get_embeddings(["hello", "world", "test"])
|
||||
|
||||
assert len(result) == 3
|
||||
dims = [len(e.vector) for e in result]
|
||||
assert all(d == dims[0] for d in dims)
|
||||
|
||||
|
||||
@skip_if_openai_integration_tests_disabled
|
||||
@pytest.mark.flaky
|
||||
@pytest.mark.integration
|
||||
async def test_integration_openai_get_embeddings_with_dimensions() -> None:
|
||||
"""Test embedding generation with custom dimensions."""
|
||||
client = OpenAIEmbeddingClient(model="text-embedding-3-small")
|
||||
|
||||
options: OpenAIEmbeddingOptions = {"dimensions": 256}
|
||||
result = await client.get_embeddings(["hello world"], options=options)
|
||||
|
||||
assert len(result) == 1
|
||||
assert len(result[0].vector) == 256
|
||||
Reference in New Issue
Block a user