Python: [BREAKING] Standardize TypeVar naming convention (TName → NameT) (#3770)

* standardized typevar to use suffix T

* addressed copilot comments
This commit is contained in:
Giles Odigwe
2026-02-10 11:34:10 -08:00
committed by GitHub
Unverified
parent 7dccf3a07b
commit a149aaa926
40 changed files with 443 additions and 437 deletions
+17 -17
View File
@@ -70,15 +70,15 @@ if TYPE_CHECKING:
from ._types import ChatOptions
TResponseModel = TypeVar("TResponseModel", bound=BaseModel | None, default=None, covariant=True)
TResponseModelT = TypeVar("TResponseModelT", bound=BaseModel)
ResponseModelT = TypeVar("ResponseModelT", bound=BaseModel | None, default=None, covariant=True)
ResponseModelBoundT = TypeVar("ResponseModelBoundT", bound=BaseModel)
logger = get_logger("agent_framework")
TThreadType = TypeVar("TThreadType", bound="AgentThread")
TOptions_co = TypeVar(
"TOptions_co",
ThreadTypeT = TypeVar("ThreadTypeT", bound="AgentThread")
OptionsCoT = TypeVar(
"OptionsCoT",
bound=TypedDict, # type: ignore[valid-type]
default="ChatOptions[None]",
covariant=True,
@@ -530,7 +530,7 @@ BareAgent = BaseAgent
# region ChatAgent
class RawChatAgent(BaseAgent, Generic[TOptions_co]): # type: ignore[misc]
class RawChatAgent(BaseAgent, Generic[OptionsCoT]): # type: ignore[misc]
"""A Chat Client Agent without middleware or telemetry layers.
This is the core chat agent implementation. For most use cases,
@@ -613,7 +613,7 @@ class RawChatAgent(BaseAgent, Generic[TOptions_co]): # type: ignore[misc]
def __init__(
self,
chat_client: ChatClientProtocol[TOptions_co],
chat_client: ChatClientProtocol[OptionsCoT],
instructions: str | None = None,
*,
id: str | None = None,
@@ -624,7 +624,7 @@ class RawChatAgent(BaseAgent, Generic[TOptions_co]): # type: ignore[misc]
| MutableMapping[str, Any]
| Sequence[ToolProtocol | Callable[..., Any] | MutableMapping[str, Any]]
| None = None,
default_options: TOptions_co | None = None,
default_options: OptionsCoT | None = None,
chat_message_store_factory: Callable[[], ChatMessageStoreProtocol] | None = None,
context_provider: ContextProvider | None = None,
**kwargs: Any,
@@ -789,9 +789,9 @@ class RawChatAgent(BaseAgent, Generic[TOptions_co]): # type: ignore[misc]
| MutableMapping[str, Any]
| list[ToolProtocol | Callable[..., Any] | MutableMapping[str, Any]]
| None = None,
options: ChatOptions[TResponseModelT],
options: ChatOptions[ResponseModelBoundT],
**kwargs: Any,
) -> Awaitable[AgentResponse[TResponseModelT]]: ...
) -> Awaitable[AgentResponse[ResponseModelBoundT]]: ...
@overload
def run(
@@ -805,7 +805,7 @@ class RawChatAgent(BaseAgent, Generic[TOptions_co]): # type: ignore[misc]
| MutableMapping[str, Any]
| list[ToolProtocol | Callable[..., Any] | MutableMapping[str, Any]]
| None = None,
options: TOptions_co | ChatOptions[None] | None = None,
options: OptionsCoT | ChatOptions[None] | None = None,
**kwargs: Any,
) -> Awaitable[AgentResponse[Any]]: ...
@@ -821,7 +821,7 @@ class RawChatAgent(BaseAgent, Generic[TOptions_co]): # type: ignore[misc]
| MutableMapping[str, Any]
| list[ToolProtocol | Callable[..., Any] | MutableMapping[str, Any]]
| None = None,
options: TOptions_co | ChatOptions[Any] | None = None,
options: OptionsCoT | ChatOptions[Any] | None = None,
**kwargs: Any,
) -> ResponseStream[AgentResponseUpdate, AgentResponse[Any]]: ...
@@ -836,7 +836,7 @@ class RawChatAgent(BaseAgent, Generic[TOptions_co]): # type: ignore[misc]
| MutableMapping[str, Any]
| list[ToolProtocol | Callable[..., Any] | MutableMapping[str, Any]]
| None = None,
options: TOptions_co | ChatOptions[Any] | None = None,
options: OptionsCoT | ChatOptions[Any] | None = None,
**kwargs: Any,
) -> Awaitable[AgentResponse[Any]] | ResponseStream[AgentResponseUpdate, AgentResponse[Any]]:
"""Run the agent with the given messages and options.
@@ -1375,8 +1375,8 @@ class RawChatAgent(BaseAgent, Generic[TOptions_co]): # type: ignore[misc]
class ChatAgent(
AgentTelemetryLayer,
AgentMiddlewareLayer,
RawChatAgent[TOptions_co],
Generic[TOptions_co],
RawChatAgent[OptionsCoT],
Generic[OptionsCoT],
):
"""A Chat Client Agent with middleware, telemetry, and full layer support.
@@ -1389,7 +1389,7 @@ class ChatAgent(
def __init__(
self,
chat_client: ChatClientProtocol[TOptions_co],
chat_client: ChatClientProtocol[OptionsCoT],
instructions: str | None = None,
*,
id: str | None = None,
@@ -1400,7 +1400,7 @@ class ChatAgent(
| MutableMapping[str, Any]
| Sequence[ToolProtocol | Callable[..., Any] | MutableMapping[str, Any]]
| None = None,
default_options: TOptions_co | None = None,
default_options: OptionsCoT | None = None,
chat_message_store_factory: Callable[[], ChatMessageStoreProtocol] | None = None,
context_provider: ContextProvider | None = None,
middleware: Sequence[MiddlewareTypes] | None = None,
@@ -58,10 +58,10 @@ if TYPE_CHECKING:
from ._types import ChatOptions
TInput = TypeVar("TInput", contravariant=True)
InputT = TypeVar("InputT", contravariant=True)
TEmbedding = TypeVar("TEmbedding")
TBaseChatClient = TypeVar("TBaseChatClient", bound="BaseChatClient")
EmbeddingT = TypeVar("EmbeddingT")
BaseChatClientT = TypeVar("BaseChatClientT", bound="BaseChatClient")
logger = get_logger()
@@ -74,19 +74,19 @@ __all__ = [
# region ChatClientProtocol Protocol
# Contravariant for the Protocol
TOptions_contra = TypeVar(
"TOptions_contra",
OptionsContraT = TypeVar(
"OptionsContraT",
bound=TypedDict, # type: ignore[valid-type]
default="ChatOptions[None]",
contravariant=True,
)
# Used for the overloads that capture the response model type from options
TResponseModelT = TypeVar("TResponseModelT", bound=BaseModel)
ResponseModelBoundT = TypeVar("ResponseModelBoundT", bound=BaseModel)
@runtime_checkable
class ChatClientProtocol(Protocol[TOptions_contra]):
class ChatClientProtocol(Protocol[OptionsContraT]):
"""A protocol for a chat client that can generate responses.
This protocol defines the interface that all chat clients must implement,
@@ -139,9 +139,9 @@ class ChatClientProtocol(Protocol[TOptions_contra]):
messages: str | ChatMessage | Sequence[str | ChatMessage],
*,
stream: Literal[False] = ...,
options: ChatOptions[TResponseModelT],
options: ChatOptions[ResponseModelBoundT],
**kwargs: Any,
) -> Awaitable[ChatResponse[TResponseModelT]]: ...
) -> Awaitable[ChatResponse[ResponseModelBoundT]]: ...
@overload
def get_response(
@@ -149,7 +149,7 @@ class ChatClientProtocol(Protocol[TOptions_contra]):
messages: str | ChatMessage | Sequence[str | ChatMessage],
*,
stream: Literal[False] = ...,
options: TOptions_contra | ChatOptions[None] | None = None,
options: OptionsContraT | ChatOptions[None] | None = None,
**kwargs: Any,
) -> Awaitable[ChatResponse[Any]]: ...
@@ -159,7 +159,7 @@ class ChatClientProtocol(Protocol[TOptions_contra]):
messages: str | ChatMessage | Sequence[str | ChatMessage],
*,
stream: Literal[True],
options: TOptions_contra | ChatOptions[Any] | None = None,
options: OptionsContraT | ChatOptions[Any] | None = None,
**kwargs: Any,
) -> ResponseStream[ChatResponseUpdate, ChatResponse[Any]]: ...
@@ -168,7 +168,7 @@ class ChatClientProtocol(Protocol[TOptions_contra]):
messages: str | ChatMessage | Sequence[str | ChatMessage],
*,
stream: bool = False,
options: TOptions_contra | ChatOptions[Any] | None = None,
options: OptionsContraT | ChatOptions[Any] | None = None,
**kwargs: Any,
) -> Awaitable[ChatResponse[Any]] | ResponseStream[ChatResponseUpdate, ChatResponse[Any]]:
"""Send input and return the response.
@@ -195,15 +195,15 @@ class ChatClientProtocol(Protocol[TOptions_contra]):
# region ChatClientBase
# Covariant for the BaseChatClient
TOptions_co = TypeVar(
"TOptions_co",
OptionsCoT = TypeVar(
"OptionsCoT",
bound=TypedDict, # type: ignore[valid-type]
default="ChatOptions[None]",
covariant=True,
)
class BaseChatClient(SerializationMixin, ABC, Generic[TOptions_co]):
class BaseChatClient(SerializationMixin, ABC, Generic[OptionsCoT]):
"""Abstract base class for chat clients without middleware wrapping.
This abstract base class provides core functionality for chat client implementations,
@@ -368,9 +368,9 @@ class BaseChatClient(SerializationMixin, ABC, Generic[TOptions_co]):
messages: str | ChatMessage | Sequence[str | ChatMessage],
*,
stream: Literal[False] = ...,
options: ChatOptions[TResponseModelT],
options: ChatOptions[ResponseModelBoundT],
**kwargs: Any,
) -> Awaitable[ChatResponse[TResponseModelT]]: ...
) -> Awaitable[ChatResponse[ResponseModelBoundT]]: ...
@overload
def get_response(
@@ -378,7 +378,7 @@ class BaseChatClient(SerializationMixin, ABC, Generic[TOptions_co]):
messages: str | ChatMessage | Sequence[str | ChatMessage],
*,
stream: Literal[False] = ...,
options: TOptions_co | ChatOptions[None] | None = None,
options: OptionsCoT | ChatOptions[None] | None = None,
**kwargs: Any,
) -> Awaitable[ChatResponse[Any]]: ...
@@ -388,7 +388,7 @@ class BaseChatClient(SerializationMixin, ABC, Generic[TOptions_co]):
messages: str | ChatMessage | Sequence[str | ChatMessage],
*,
stream: Literal[True],
options: TOptions_co | ChatOptions[Any] | None = None,
options: OptionsCoT | ChatOptions[Any] | None = None,
**kwargs: Any,
) -> ResponseStream[ChatResponseUpdate, ChatResponse[Any]]: ...
@@ -397,7 +397,7 @@ class BaseChatClient(SerializationMixin, ABC, Generic[TOptions_co]):
messages: str | ChatMessage | Sequence[str | ChatMessage],
*,
stream: bool = False,
options: TOptions_co | ChatOptions[Any] | None = None,
options: OptionsCoT | ChatOptions[Any] | None = None,
**kwargs: Any,
) -> Awaitable[ChatResponse[Any]] | ResponseStream[ChatResponseUpdate, ChatResponse[Any]]:
"""Get a response from a chat client.
@@ -442,13 +442,13 @@ class BaseChatClient(SerializationMixin, ABC, Generic[TOptions_co]):
| MutableMapping[str, Any]
| Sequence[ToolProtocol | Callable[..., Any] | MutableMapping[str, Any]]
| None = None,
default_options: TOptions_co | Mapping[str, Any] | None = None,
default_options: OptionsCoT | Mapping[str, Any] | None = None,
chat_message_store_factory: Callable[[], ChatMessageStoreProtocol] | None = None,
context_provider: ContextProvider | None = None,
middleware: Sequence[MiddlewareTypes] | None = None,
function_invocation_configuration: FunctionInvocationConfiguration | None = None,
**kwargs: Any,
) -> ChatAgent[TOptions_co]:
) -> ChatAgent[OptionsCoT]:
"""Create a ChatAgent with this client.
This is a convenience method that creates a ChatAgent instance with this
@@ -40,7 +40,7 @@ if TYPE_CHECKING:
from ._tools import FunctionTool
from ._types import ChatOptions, ChatResponse, ChatResponseUpdate
TResponseModelT = TypeVar("TResponseModelT", bound=BaseModel)
ResponseModelBoundT = TypeVar("ResponseModelBoundT", bound=BaseModel)
__all__ = [
"AgentContext",
@@ -65,21 +65,21 @@ __all__ = [
]
AgentT = TypeVar("AgentT", bound="SupportsAgentRun")
TContext = TypeVar("TContext")
TUpdate = TypeVar("TUpdate")
ContextT = TypeVar("ContextT")
UpdateT = TypeVar("UpdateT")
class _EmptyAsyncIterator(Generic[TUpdate]):
class _EmptyAsyncIterator(Generic[UpdateT]):
"""Empty async iterator that yields nothing.
Used when middleware terminates without setting a result,
and we need to provide an empty stream.
"""
def __aiter__(self) -> _EmptyAsyncIterator[TUpdate]:
def __aiter__(self) -> _EmptyAsyncIterator[UpdateT]:
return self
async def __anext__(self) -> TUpdate:
async def __anext__(self) -> UpdateT:
raise StopAsyncIteration
@@ -656,20 +656,20 @@ def chat_middleware(func: ChatMiddlewareCallable) -> ChatMiddlewareCallable:
return func
class MiddlewareWrapper(Generic[TContext]):
class MiddlewareWrapper(Generic[ContextT]):
"""Generic wrapper to convert pure functions into middleware protocol objects.
This wrapper allows function-based middleware to be used alongside class-based middleware
by providing a unified interface.
Type Parameters:
TContext: The type of context object this middleware operates on.
ContextT: The type of context object this middleware operates on.
"""
def __init__(self, func: Callable[[TContext, Callable[[TContext], Awaitable[None]]], Awaitable[None]]) -> None:
def __init__(self, func: Callable[[ContextT, Callable[[ContextT], Awaitable[None]]], Awaitable[None]]) -> None:
self.func = func
async def process(self, context: TContext, call_next: Callable[[TContext], Awaitable[None]]) -> None:
async def process(self, context: ContextT, call_next: Callable[[ContextT], Awaitable[None]]) -> None:
await self.func(context, call_next)
@@ -953,15 +953,15 @@ class ChatMiddlewarePipeline(BaseMiddlewarePipeline):
# Covariant for chat client options
TOptions_co = TypeVar(
"TOptions_co",
OptionsCoT = TypeVar(
"OptionsCoT",
bound=TypedDict, # type: ignore[valid-type]
default="ChatOptions[None]",
covariant=True,
)
class ChatMiddlewareLayer(Generic[TOptions_co]):
class ChatMiddlewareLayer(Generic[OptionsCoT]):
"""Layer for chat clients to apply chat middleware around response generation."""
def __init__(
@@ -983,9 +983,9 @@ class ChatMiddlewareLayer(Generic[TOptions_co]):
messages: str | ChatMessage | Sequence[str | ChatMessage],
*,
stream: Literal[False] = ...,
options: ChatOptions[TResponseModelT],
options: ChatOptions[ResponseModelBoundT],
**kwargs: Any,
) -> Awaitable[ChatResponse[TResponseModelT]]: ...
) -> Awaitable[ChatResponse[ResponseModelBoundT]]: ...
@overload
def get_response(
@@ -993,7 +993,7 @@ class ChatMiddlewareLayer(Generic[TOptions_co]):
messages: str | ChatMessage | Sequence[str | ChatMessage],
*,
stream: Literal[False] = ...,
options: TOptions_co | ChatOptions[None] | None = None,
options: OptionsCoT | ChatOptions[None] | None = None,
**kwargs: Any,
) -> Awaitable[ChatResponse[Any]]: ...
@@ -1003,7 +1003,7 @@ class ChatMiddlewareLayer(Generic[TOptions_co]):
messages: str | ChatMessage | Sequence[str | ChatMessage],
*,
stream: Literal[True],
options: TOptions_co | ChatOptions[Any] | None = None,
options: OptionsCoT | ChatOptions[Any] | None = None,
**kwargs: Any,
) -> ResponseStream[ChatResponseUpdate, ChatResponse[Any]]: ...
@@ -1012,7 +1012,7 @@ class ChatMiddlewareLayer(Generic[TOptions_co]):
messages: str | ChatMessage | Sequence[str | ChatMessage],
*,
stream: bool = False,
options: TOptions_co | ChatOptions[Any] | None = None,
options: OptionsCoT | ChatOptions[Any] | None = None,
**kwargs: Any,
) -> Awaitable[ChatResponse[Any]] | ResponseStream[ChatResponseUpdate, ChatResponse[Any]]:
"""Execute the chat pipeline if middleware is configured."""
@@ -1102,9 +1102,9 @@ class AgentMiddlewareLayer:
stream: Literal[False] = ...,
thread: AgentThread | None = None,
middleware: Sequence[MiddlewareTypes] | None = None,
options: ChatOptions[TResponseModelT],
options: ChatOptions[ResponseModelBoundT],
**kwargs: Any,
) -> Awaitable[AgentResponse[TResponseModelT]]: ...
) -> Awaitable[AgentResponse[ResponseModelBoundT]]: ...
@overload
def run(
@@ -14,7 +14,7 @@ HTTPsUrl = Annotated[AnyUrl, UrlConstraints(max_length=2083, allowed_schemes=["h
__all__ = ["AFBaseSettings", "HTTPsUrl"]
TSettings = TypeVar("TSettings", bound="AFBaseSettings")
SettingsT = TypeVar("SettingsT", bound="AFBaseSettings")
class AFBaseSettings(BaseSettings):
@@ -50,7 +50,7 @@ class AFBaseSettings(BaseSettings):
kwargs = {k: v for k, v in kwargs.items() if v is not None}
super().__init__(**kwargs)
def __new__(cls: type[TSettings], *args: Any, **kwargs: Any) -> TSettings:
def __new__(cls: type[SettingsT], *args: Any, **kwargs: Any) -> SettingsT:
"""Override the __new__ method to set the env_prefix."""
# for both, if supplied but None, set to default
if "env_file_encoding" in kwargs and kwargs["env_file_encoding"] is not None:
@@ -11,8 +11,8 @@ from ._logging import get_logger
logger = get_logger()
TClass = TypeVar("TClass", bound="SerializationMixin")
TProtocol = TypeVar("TProtocol", bound="SerializationProtocol")
ClassT = TypeVar("ClassT", bound="SerializationMixin")
ProtocolT = TypeVar("ProtocolT", bound="SerializationProtocol")
# Regex pattern for converting CamelCase to snake_case
_CAMEL_TO_SNAKE_PATTERN = re.compile(r"(?<!^)(?=[A-Z])")
@@ -95,7 +95,7 @@ class SerializationProtocol(Protocol):
...
@classmethod
def from_dict(cls: type[TProtocol], value: MutableMapping[str, Any], /, **kwargs: Any) -> TProtocol:
def from_dict(cls: type[ProtocolT], value: MutableMapping[str, Any], /, **kwargs: Any) -> ProtocolT:
"""Create an instance from a dictionary.
Args:
@@ -392,8 +392,8 @@ class SerializationMixin:
@classmethod
def from_dict(
cls: type[TClass], value: MutableMapping[str, Any], /, *, dependencies: MutableMapping[str, Any] | None = None
) -> TClass:
cls: type[ClassT], value: MutableMapping[str, Any], /, *, dependencies: MutableMapping[str, Any] | None = None
) -> ClassT:
"""Create an instance from a dictionary with optional dependency injection.
This method reconstructs an object from its dictionary representation, automatically
@@ -560,7 +560,7 @@ class SerializationMixin:
return cls(**kwargs)
@classmethod
def from_json(cls: type[TClass], value: str, /, *, dependencies: MutableMapping[str, Any] | None = None) -> TClass:
def from_json(cls: type[ClassT], value: str, /, *, dependencies: MutableMapping[str, Any] | None = None) -> ClassT:
"""Create an instance from a JSON string.
This is a convenience method that parses the JSON string using ``json.loads()``
@@ -182,7 +182,7 @@ class AgentThreadState(SerializationMixin):
raise TypeError("Could not parse ChatMessageStoreState.")
TChatMessageStore = TypeVar("TChatMessageStore", bound="ChatMessageStore")
ChatMessageStoreT = TypeVar("ChatMessageStoreT", bound="ChatMessageStore")
class ChatMessageStore:
@@ -243,8 +243,8 @@ class ChatMessageStore:
@classmethod
async def deserialize(
cls: type[TChatMessageStore], serialized_store_state: MutableMapping[str, Any], **kwargs: Any
) -> TChatMessageStore:
cls: type[ChatMessageStoreT], serialized_store_state: MutableMapping[str, Any], **kwargs: Any
) -> ChatMessageStoreT:
"""Create a new ChatMessageStore instance from serialized state data.
Args:
@@ -289,7 +289,7 @@ class ChatMessageStore:
return state.to_dict()
TAgentThread = TypeVar("TAgentThread", bound="AgentThread")
AgentThreadT = TypeVar("AgentThreadT", bound="AgentThread")
class AgentThread:
@@ -437,12 +437,12 @@ class AgentThread:
@classmethod
async def deserialize(
cls: type[TAgentThread],
cls: type[AgentThreadT],
serialized_thread_state: MutableMapping[str, Any],
*,
message_store: ChatMessageStoreProtocol | None = None,
**kwargs: Any,
) -> TAgentThread:
) -> AgentThreadT:
"""Deserializes the state from a dictionary into a new AgentThread instance.
Args:
+11 -11
View File
@@ -76,7 +76,7 @@ if TYPE_CHECKING:
ResponseStream,
)
TResponseModelT = TypeVar("TResponseModelT", bound=BaseModel)
ResponseModelBoundT = TypeVar("ResponseModelBoundT", bound=BaseModel)
logger = get_logger()
@@ -100,7 +100,7 @@ __all__ = [
logger = get_logger()
DEFAULT_MAX_ITERATIONS: Final[int] = 40
DEFAULT_MAX_CONSECUTIVE_ERRORS_PER_REQUEST: Final[int] = 3
TChatClient = TypeVar("TChatClient", bound="ChatClientProtocol[Any]")
ChatClientT = TypeVar("ChatClientT", bound="ChatClientProtocol[Any]")
# region Helpers
ArgsT = TypeVar("ArgsT", bound=BaseModel, default=BaseModel)
@@ -569,7 +569,7 @@ def _default_histogram() -> Histogram:
)
TClass = TypeVar("TClass", bound="SerializationMixin")
ClassT = TypeVar("ClassT", bound="SerializationMixin")
class EmptyInputModel(BaseModel):
@@ -2083,15 +2083,15 @@ async def _process_function_requests(
return result
TOptions_co = TypeVar(
"TOptions_co",
OptionsCoT = TypeVar(
"OptionsCoT",
bound=TypedDict, # type: ignore[valid-type]
default="ChatOptions[None]",
covariant=True,
)
class FunctionInvocationLayer(Generic[TOptions_co]):
class FunctionInvocationLayer(Generic[OptionsCoT]):
"""Layer for chat clients to apply function invocation around get_response."""
def __init__(
@@ -2115,9 +2115,9 @@ class FunctionInvocationLayer(Generic[TOptions_co]):
messages: str | ChatMessage | Sequence[str | ChatMessage],
*,
stream: Literal[False] = ...,
options: ChatOptions[TResponseModelT],
options: ChatOptions[ResponseModelBoundT],
**kwargs: Any,
) -> Awaitable[ChatResponse[TResponseModelT]]: ...
) -> Awaitable[ChatResponse[ResponseModelBoundT]]: ...
@overload
def get_response(
@@ -2125,7 +2125,7 @@ class FunctionInvocationLayer(Generic[TOptions_co]):
messages: str | ChatMessage | Sequence[str | ChatMessage],
*,
stream: Literal[False] = ...,
options: TOptions_co | ChatOptions[None] | None = None,
options: OptionsCoT | ChatOptions[None] | None = None,
**kwargs: Any,
) -> Awaitable[ChatResponse[Any]]: ...
@@ -2135,7 +2135,7 @@ class FunctionInvocationLayer(Generic[TOptions_co]):
messages: str | ChatMessage | Sequence[str | ChatMessage],
*,
stream: Literal[True],
options: TOptions_co | ChatOptions[Any] | None = None,
options: OptionsCoT | ChatOptions[Any] | None = None,
**kwargs: Any,
) -> ResponseStream[ChatResponseUpdate, ChatResponse[Any]]: ...
@@ -2144,7 +2144,7 @@ class FunctionInvocationLayer(Generic[TOptions_co]):
messages: str | ChatMessage | Sequence[str | ChatMessage],
*,
stream: bool = False,
options: TOptions_co | ChatOptions[Any] | None = None,
options: OptionsCoT | ChatOptions[Any] | None = None,
function_middleware: Sequence[FunctionMiddlewareTypes] | None = None,
**kwargs: Any,
) -> Awaitable[ChatResponse[Any]] | ResponseStream[ChatResponseUpdate, ChatResponse[Any]]:
+110 -110
View File
@@ -36,17 +36,17 @@ __all__ = [
"ChatResponse",
"ChatResponseUpdate",
"Content",
"FinalT",
"FinishReason",
"FinishReasonLiteral",
"OuterFinalT",
"OuterUpdateT",
"ResponseStream",
"Role",
"RoleLiteral",
"TFinal",
"TOuterFinal",
"TOuterUpdate",
"TUpdate",
"TextSpanRegion",
"ToolMode",
"UpdateT",
"UsageDetails",
"add_usage_details",
"detect_media_type_from_base64",
@@ -305,12 +305,12 @@ def _serialize_value(value: Any, exclude_none: bool) -> Any:
# region Constants and types
_T = TypeVar("_T")
TEmbedding = TypeVar("TEmbedding")
TChatResponse = TypeVar("TChatResponse", bound="ChatResponse")
TToolMode = TypeVar("TToolMode", bound="ToolMode")
TAgentRunResponse = TypeVar("TAgentRunResponse", bound="AgentResponse")
TResponseModel = TypeVar("TResponseModel", bound=BaseModel | None, default=None, covariant=True)
TResponseModelT = TypeVar("TResponseModelT", bound=BaseModel)
EmbeddingT = TypeVar("EmbeddingT")
ChatResponseT = TypeVar("ChatResponseT", bound="ChatResponse")
ToolModeT = TypeVar("ToolModeT", bound="ToolMode")
AgentResponseT = TypeVar("AgentResponseT", bound="AgentResponse")
ResponseModelT = TypeVar("ResponseModelT", bound=BaseModel | None, default=None, covariant=True)
ResponseModelBoundT = TypeVar("ResponseModelBoundT", bound=BaseModel)
CreatedAtT = str # Use a datetimeoffset type? Or a more specific type like datetime.datetime?
@@ -389,7 +389,7 @@ class Annotation(TypedDict, total=False):
raw_representation: Any
TContent = TypeVar("TContent", bound="Content")
ContentT = TypeVar("ContentT", bound="Content")
# endregion
@@ -544,13 +544,13 @@ class Content:
@classmethod
def from_text(
cls: type[TContent],
cls: type[ContentT],
text: str,
*,
annotations: Sequence[Annotation] | None = None,
additional_properties: MutableMapping[str, Any] | None = None,
raw_representation: Any = None,
) -> TContent:
) -> ContentT:
"""Create text content."""
return cls(
"text",
@@ -562,14 +562,14 @@ class Content:
@classmethod
def from_text_reasoning(
cls: type[TContent],
cls: type[ContentT],
*,
text: str | None = None,
protected_data: str | None = None,
annotations: Sequence[Annotation] | None = None,
additional_properties: MutableMapping[str, Any] | None = None,
raw_representation: Any = None,
) -> TContent:
) -> ContentT:
"""Create text reasoning content."""
return cls(
"text_reasoning",
@@ -582,14 +582,14 @@ class Content:
@classmethod
def from_data(
cls: type[TContent],
cls: type[ContentT],
data: bytes,
media_type: str,
*,
annotations: Sequence[Annotation] | None = None,
additional_properties: MutableMapping[str, Any] | None = None,
raw_representation: Any = None,
) -> TContent:
) -> ContentT:
r"""Create data content from raw binary data.
Use this to create content from binary data (images, audio, documents, etc.).
@@ -658,14 +658,14 @@ class Content:
@classmethod
def from_uri(
cls: type[TContent],
cls: type[ContentT],
uri: str,
*,
media_type: str | None = None,
annotations: Sequence[Annotation] | None = None,
additional_properties: MutableMapping[str, Any] | None = None,
raw_representation: Any = None,
) -> TContent:
) -> ContentT:
"""Create content from a URI, can be both data URI or external URI.
Use this when you already have a properly formed data URI
@@ -720,7 +720,7 @@ class Content:
@classmethod
def from_error(
cls: type[TContent],
cls: type[ContentT],
*,
message: str | None = None,
error_code: str | None = None,
@@ -728,7 +728,7 @@ class Content:
annotations: Sequence[Annotation] | None = None,
additional_properties: MutableMapping[str, Any] | None = None,
raw_representation: Any = None,
) -> TContent:
) -> ContentT:
"""Create error content."""
return cls(
"error",
@@ -742,7 +742,7 @@ class Content:
@classmethod
def from_function_call(
cls: type[TContent],
cls: type[ContentT],
call_id: str,
name: str,
*,
@@ -751,7 +751,7 @@ class Content:
annotations: Sequence[Annotation] | None = None,
additional_properties: MutableMapping[str, Any] | None = None,
raw_representation: Any = None,
) -> TContent:
) -> ContentT:
"""Create function call content."""
return cls(
"function_call",
@@ -766,7 +766,7 @@ class Content:
@classmethod
def from_function_result(
cls: type[TContent],
cls: type[ContentT],
call_id: str,
*,
result: Any = None,
@@ -774,7 +774,7 @@ class Content:
annotations: Sequence[Annotation] | None = None,
additional_properties: MutableMapping[str, Any] | None = None,
raw_representation: Any = None,
) -> TContent:
) -> ContentT:
"""Create function result content."""
return cls(
"function_result",
@@ -788,13 +788,13 @@ class Content:
@classmethod
def from_usage(
cls: type[TContent],
cls: type[ContentT],
usage_details: UsageDetails,
*,
annotations: Sequence[Annotation] | None = None,
additional_properties: MutableMapping[str, Any] | None = None,
raw_representation: Any = None,
) -> TContent:
) -> ContentT:
"""Create usage content."""
return cls(
"usage",
@@ -806,7 +806,7 @@ class Content:
@classmethod
def from_hosted_file(
cls: type[TContent],
cls: type[ContentT],
file_id: str,
*,
media_type: str | None = None,
@@ -814,7 +814,7 @@ class Content:
annotations: Sequence[Annotation] | None = None,
additional_properties: MutableMapping[str, Any] | None = None,
raw_representation: Any = None,
) -> TContent:
) -> ContentT:
"""Create hosted file content."""
return cls(
"hosted_file",
@@ -828,13 +828,13 @@ class Content:
@classmethod
def from_hosted_vector_store(
cls: type[TContent],
cls: type[ContentT],
vector_store_id: str,
*,
annotations: Sequence[Annotation] | None = None,
additional_properties: MutableMapping[str, Any] | None = None,
raw_representation: Any = None,
) -> TContent:
) -> ContentT:
"""Create hosted vector store content."""
return cls(
"hosted_vector_store",
@@ -846,14 +846,14 @@ class Content:
@classmethod
def from_code_interpreter_tool_call(
cls: type[TContent],
cls: type[ContentT],
*,
call_id: str | None = None,
inputs: Sequence[Content] | None = None,
annotations: Sequence[Annotation] | None = None,
additional_properties: MutableMapping[str, Any] | None = None,
raw_representation: Any = None,
) -> TContent:
) -> ContentT:
"""Create code interpreter tool call content."""
return cls(
"code_interpreter_tool_call",
@@ -866,14 +866,14 @@ class Content:
@classmethod
def from_code_interpreter_tool_result(
cls: type[TContent],
cls: type[ContentT],
*,
call_id: str | None = None,
outputs: Sequence[Content] | None = None,
annotations: Sequence[Annotation] | None = None,
additional_properties: MutableMapping[str, Any] | None = None,
raw_representation: Any = None,
) -> TContent:
) -> ContentT:
"""Create code interpreter tool result content."""
return cls(
"code_interpreter_tool_result",
@@ -886,13 +886,13 @@ class Content:
@classmethod
def from_image_generation_tool_call(
cls: type[TContent],
cls: type[ContentT],
*,
image_id: str | None = None,
annotations: Sequence[Annotation] | None = None,
additional_properties: MutableMapping[str, Any] | None = None,
raw_representation: Any = None,
) -> TContent:
) -> ContentT:
"""Create image generation tool call content."""
return cls(
"image_generation_tool_call",
@@ -904,14 +904,14 @@ class Content:
@classmethod
def from_image_generation_tool_result(
cls: type[TContent],
cls: type[ContentT],
*,
image_id: str | None = None,
outputs: Any = None,
annotations: Sequence[Annotation] | None = None,
additional_properties: MutableMapping[str, Any] | None = None,
raw_representation: Any = None,
) -> TContent:
) -> ContentT:
"""Create image generation tool result content."""
return cls(
"image_generation_tool_result",
@@ -924,7 +924,7 @@ class Content:
@classmethod
def from_mcp_server_tool_call(
cls: type[TContent],
cls: type[ContentT],
call_id: str,
tool_name: str,
*,
@@ -933,7 +933,7 @@ class Content:
annotations: Sequence[Annotation] | None = None,
additional_properties: MutableMapping[str, Any] | None = None,
raw_representation: Any = None,
) -> TContent:
) -> ContentT:
"""Create MCP server tool call content."""
return cls(
"mcp_server_tool_call",
@@ -948,14 +948,14 @@ class Content:
@classmethod
def from_mcp_server_tool_result(
cls: type[TContent],
cls: type[ContentT],
call_id: str,
*,
output: Any = None,
annotations: Sequence[Annotation] | None = None,
additional_properties: MutableMapping[str, Any] | None = None,
raw_representation: Any = None,
) -> TContent:
) -> ContentT:
"""Create MCP server tool result content."""
return cls(
"mcp_server_tool_result",
@@ -968,14 +968,14 @@ class Content:
@classmethod
def from_function_approval_request(
cls: type[TContent],
cls: type[ContentT],
id: str,
function_call: Content,
*,
annotations: Sequence[Annotation] | None = None,
additional_properties: MutableMapping[str, Any] | None = None,
raw_representation: Any = None,
) -> TContent:
) -> ContentT:
"""Create function approval request content."""
return cls(
"function_approval_request",
@@ -989,7 +989,7 @@ class Content:
@classmethod
def from_function_approval_response(
cls: type[TContent],
cls: type[ContentT],
approved: bool,
id: str,
function_call: Content,
@@ -997,7 +997,7 @@ class Content:
annotations: Sequence[Annotation] | None = None,
additional_properties: MutableMapping[str, Any] | None = None,
raw_representation: Any = None,
) -> TContent:
) -> ContentT:
"""Create function approval response content."""
return cls(
"function_approval_response",
@@ -1091,7 +1091,7 @@ class Content:
return f"Content(type={self.type})"
@classmethod
def from_dict(cls: type[TContent], data: Mapping[str, Any]) -> TContent:
def from_dict(cls: type[ContentT], data: Mapping[str, Any]) -> ContentT:
"""Create a Content instance from a mapping."""
if not (content_type := data.get("type")):
raise ValueError("Content mapping requires 'type'")
@@ -1796,7 +1796,7 @@ def _finalize_response(response: ChatResponse | AgentResponse) -> None:
_coalesce_text_content(msg.contents, "text_reasoning")
class ChatResponse(SerializationMixin, Generic[TResponseModel]):
class ChatResponse(SerializationMixin, Generic[ResponseModelT]):
"""Represents the response to a chat request.
Attributes:
@@ -1859,7 +1859,7 @@ class ChatResponse(SerializationMixin, Generic[TResponseModel]):
created_at: CreatedAtT | None = None,
finish_reason: FinishReasonLiteral | FinishReason | None = None,
usage_details: UsageDetails | None = None,
value: TResponseModel | None = None,
value: ResponseModelT | None = None,
response_format: type[BaseModel] | None = None,
additional_properties: dict[str, Any] | None = None,
raw_representation: Any | None = None,
@@ -1903,7 +1903,7 @@ class ChatResponse(SerializationMixin, Generic[TResponseModel]):
finish_reason = finish_reason["value"]
self.finish_reason = finish_reason
self.usage_details = usage_details
self._value: TResponseModel | None = value
self._value: ResponseModelT | None = value
self._response_format: type[BaseModel] | None = response_format
self._value_parsed: bool = value is not None
self.additional_properties = additional_properties or {}
@@ -1915,8 +1915,8 @@ class ChatResponse(SerializationMixin, Generic[TResponseModel]):
cls: type[ChatResponse[Any]],
updates: Sequence[ChatResponseUpdate],
*,
output_format_type: type[TResponseModelT],
) -> ChatResponse[TResponseModelT]: ...
output_format_type: type[ResponseModelBoundT],
) -> ChatResponse[ResponseModelBoundT]: ...
@overload
@classmethod
@@ -1929,11 +1929,11 @@ class ChatResponse(SerializationMixin, Generic[TResponseModel]):
@classmethod
def from_updates(
cls: type[TChatResponse],
cls: type[ChatResponseT],
updates: Sequence[ChatResponseUpdate],
*,
output_format_type: type[BaseModel] | None = None,
) -> TChatResponse:
) -> ChatResponseT:
"""Joins multiple updates into a single ChatResponse.
Example:
@@ -1970,8 +1970,8 @@ class ChatResponse(SerializationMixin, Generic[TResponseModel]):
cls: type[ChatResponse[Any]],
updates: AsyncIterable[ChatResponseUpdate],
*,
output_format_type: type[TResponseModelT],
) -> ChatResponse[TResponseModelT]: ...
output_format_type: type[ResponseModelBoundT],
) -> ChatResponse[ResponseModelBoundT]: ...
@overload
@classmethod
@@ -1984,11 +1984,11 @@ class ChatResponse(SerializationMixin, Generic[TResponseModel]):
@classmethod
async def from_update_generator(
cls: type[TChatResponse],
cls: type[ChatResponseT],
updates: AsyncIterable[ChatResponseUpdate],
*,
output_format_type: type[BaseModel] | None = None,
) -> TChatResponse:
) -> ChatResponseT:
"""Joins multiple updates into a single ChatResponse.
Example:
@@ -2021,7 +2021,7 @@ class ChatResponse(SerializationMixin, Generic[TResponseModel]):
return ("\n".join(message.text for message in self.messages if isinstance(message, ChatMessage))).strip()
@property
def value(self) -> TResponseModel | None:
def value(self) -> ResponseModelT | None:
"""Get the parsed structured output value.
If a response_format was provided and parsing hasn't been attempted yet,
@@ -2037,7 +2037,7 @@ class ChatResponse(SerializationMixin, Generic[TResponseModel]):
and isinstance(self._response_format, type)
and issubclass(self._response_format, BaseModel)
):
self._value = cast(TResponseModel, self._response_format.model_validate_json(self.text))
self._value = cast(ResponseModelT, self._response_format.model_validate_json(self.text))
self._value_parsed = True
return self._value
@@ -2166,7 +2166,7 @@ class ChatResponseUpdate(SerializationMixin):
# region AgentResponse
class AgentResponse(SerializationMixin, Generic[TResponseModel]):
class AgentResponse(SerializationMixin, Generic[ResponseModelT]):
"""Represents the response to an Agent run request.
Provides one or more response messages and metadata about the response.
@@ -2220,7 +2220,7 @@ class AgentResponse(SerializationMixin, Generic[TResponseModel]):
agent_id: str | None = None,
created_at: CreatedAtT | None = None,
usage_details: UsageDetails | None = None,
value: TResponseModel | None = None,
value: ResponseModelT | None = None,
response_format: type[BaseModel] | None = None,
raw_representation: Any | None = None,
additional_properties: dict[str, Any] | None = None,
@@ -2258,7 +2258,7 @@ class AgentResponse(SerializationMixin, Generic[TResponseModel]):
self.agent_id = agent_id
self.created_at = created_at
self.usage_details = usage_details
self._value: TResponseModel | None = value
self._value: ResponseModelT | None = value
self._response_format: type[BaseModel] | None = response_format
self._value_parsed: bool = value is not None
self.additional_properties = additional_properties or {}
@@ -2270,7 +2270,7 @@ class AgentResponse(SerializationMixin, Generic[TResponseModel]):
return "".join(msg.text for msg in self.messages) if self.messages else ""
@property
def value(self) -> TResponseModel | None:
def value(self) -> ResponseModelT | None:
"""Get the parsed structured output value.
If a response_format was provided and parsing hasn't been attempted yet,
@@ -2286,7 +2286,7 @@ class AgentResponse(SerializationMixin, Generic[TResponseModel]):
and isinstance(self._response_format, type)
and issubclass(self._response_format, BaseModel)
):
self._value = cast(TResponseModel, self._response_format.model_validate_json(self.text))
self._value = cast(ResponseModelT, self._response_format.model_validate_json(self.text))
self._value_parsed = True
return self._value
@@ -2306,8 +2306,8 @@ class AgentResponse(SerializationMixin, Generic[TResponseModel]):
cls: type[AgentResponse[Any]],
updates: Sequence[AgentResponseUpdate],
*,
output_format_type: type[TResponseModelT],
) -> AgentResponse[TResponseModelT]: ...
output_format_type: type[ResponseModelBoundT],
) -> AgentResponse[ResponseModelBoundT]: ...
@overload
@classmethod
@@ -2320,11 +2320,11 @@ class AgentResponse(SerializationMixin, Generic[TResponseModel]):
@classmethod
def from_updates(
cls: type[TAgentRunResponse],
cls: type[AgentResponseT],
updates: Sequence[AgentResponseUpdate],
*,
output_format_type: type[BaseModel] | None = None,
) -> TAgentRunResponse:
) -> AgentResponseT:
"""Joins multiple updates into a single AgentResponse.
Args:
@@ -2345,8 +2345,8 @@ class AgentResponse(SerializationMixin, Generic[TResponseModel]):
cls: type[AgentResponse[Any]],
updates: AsyncIterable[AgentResponseUpdate],
*,
output_format_type: type[TResponseModelT],
) -> AgentResponse[TResponseModelT]: ...
output_format_type: type[ResponseModelBoundT],
) -> AgentResponse[ResponseModelBoundT]: ...
@overload
@classmethod
@@ -2359,11 +2359,11 @@ class AgentResponse(SerializationMixin, Generic[TResponseModel]):
@classmethod
async def from_update_generator(
cls: type[TAgentRunResponse],
cls: type[AgentResponseT],
updates: AsyncIterable[AgentResponseUpdate],
*,
output_format_type: type[BaseModel] | None = None,
) -> TAgentRunResponse:
) -> AgentResponseT:
"""Joins multiple updates into a single AgentResponse.
Args:
@@ -2520,23 +2520,23 @@ def map_chat_to_agent_update(update: ChatResponseUpdate, agent_name: str | None)
# Type variables for ResponseStream
TUpdate = TypeVar("TUpdate")
TFinal = TypeVar("TFinal")
TOuterUpdate = TypeVar("TOuterUpdate")
TOuterFinal = TypeVar("TOuterFinal")
UpdateT = TypeVar("UpdateT")
FinalT = TypeVar("FinalT")
OuterUpdateT = TypeVar("OuterUpdateT")
OuterFinalT = TypeVar("OuterFinalT")
class ResponseStream(AsyncIterable[TUpdate], Generic[TUpdate, TFinal]):
class ResponseStream(AsyncIterable[UpdateT], Generic[UpdateT, FinalT]):
"""Async stream wrapper that supports iteration and deferred finalization."""
def __init__(
self,
stream: AsyncIterable[TUpdate] | Awaitable[AsyncIterable[TUpdate]],
stream: AsyncIterable[UpdateT] | Awaitable[AsyncIterable[UpdateT]],
*,
finalizer: Callable[[Sequence[TUpdate]], TFinal | Awaitable[TFinal]] | None = None,
transform_hooks: list[Callable[[TUpdate], TUpdate | Awaitable[TUpdate] | None]] | None = None,
finalizer: Callable[[Sequence[UpdateT]], FinalT | Awaitable[FinalT]] | None = None,
transform_hooks: list[Callable[[UpdateT], UpdateT | Awaitable[UpdateT] | None]] | None = None,
cleanup_hooks: list[Callable[[], Awaitable[None] | None]] | None = None,
result_hooks: list[Callable[[TFinal], TFinal | Awaitable[TFinal | None] | None]] | None = None,
result_hooks: list[Callable[[FinalT], FinalT | Awaitable[FinalT | None] | None]] | None = None,
) -> None:
"""A Async Iterable stream of updates.
@@ -2552,16 +2552,16 @@ class ResponseStream(AsyncIterable[TUpdate], Generic[TUpdate, TFinal]):
"""
self._stream_source = stream
self._finalizer = finalizer
self._stream: AsyncIterable[TUpdate] | None = None
self._iterator: AsyncIterator[TUpdate] | None = None
self._updates: list[TUpdate] = []
self._stream: AsyncIterable[UpdateT] | None = None
self._iterator: AsyncIterator[UpdateT] | None = None
self._updates: list[UpdateT] = []
self._consumed: bool = False
self._finalized: bool = False
self._final_result: TFinal | None = None
self._transform_hooks: list[Callable[[TUpdate], TUpdate | Awaitable[TUpdate] | None]] = (
self._final_result: FinalT | None = None
self._transform_hooks: list[Callable[[UpdateT], UpdateT | Awaitable[UpdateT] | None]] = (
transform_hooks if transform_hooks is not None else []
)
self._result_hooks: list[Callable[[TFinal], TFinal | Awaitable[TFinal | None] | None]] = (
self._result_hooks: list[Callable[[FinalT], FinalT | Awaitable[FinalT | None] | None]] = (
result_hooks if result_hooks is not None else []
)
self._cleanup_hooks: list[Callable[[], Awaitable[None] | None]] = (
@@ -2575,9 +2575,9 @@ class ResponseStream(AsyncIterable[TUpdate], Generic[TUpdate, TFinal]):
def map(
self,
transform: Callable[[TUpdate], TOuterUpdate | Awaitable[TOuterUpdate]],
finalizer: Callable[[Sequence[TOuterUpdate]], TOuterFinal | Awaitable[TOuterFinal]],
) -> ResponseStream[TOuterUpdate, TOuterFinal]:
transform: Callable[[UpdateT], OuterUpdateT | Awaitable[OuterUpdateT]],
finalizer: Callable[[Sequence[OuterUpdateT]], OuterFinalT | Awaitable[OuterFinalT]],
) -> ResponseStream[OuterUpdateT, OuterFinalT]:
"""Create a new stream that transforms each update.
The returned stream delegates iteration to this stream, ensuring single consumption.
@@ -2619,8 +2619,8 @@ class ResponseStream(AsyncIterable[TUpdate], Generic[TUpdate, TFinal]):
def with_finalizer(
self,
finalizer: Callable[[Sequence[TUpdate]], TOuterFinal | Awaitable[TOuterFinal]],
) -> ResponseStream[TUpdate, TOuterFinal]:
finalizer: Callable[[Sequence[UpdateT]], OuterFinalT | Awaitable[OuterFinalT]],
) -> ResponseStream[UpdateT, OuterFinalT]:
"""Create a new stream with a different finalizer.
The returned stream delegates iteration to this stream, ensuring single consumption.
@@ -2647,8 +2647,8 @@ class ResponseStream(AsyncIterable[TUpdate], Generic[TUpdate, TFinal]):
@classmethod
def from_awaitable(
cls,
awaitable: Awaitable[ResponseStream[TUpdate, TFinal]],
) -> ResponseStream[TUpdate, TFinal]:
awaitable: Awaitable[ResponseStream[UpdateT, FinalT]],
) -> ResponseStream[UpdateT, FinalT]:
"""Create a ResponseStream from an awaitable that resolves to a ResponseStream.
This is useful when you have an async function that returns a ResponseStream
@@ -2672,7 +2672,7 @@ class ResponseStream(AsyncIterable[TUpdate], Generic[TUpdate, TFinal]):
stream._wrap_inner = True
return stream # type: ignore[return-value]
async def _get_stream(self) -> AsyncIterable[TUpdate]:
async def _get_stream(self) -> AsyncIterable[UpdateT]:
if self._stream is None:
if hasattr(self._stream_source, "__aiter__"):
self._stream = self._stream_source # type: ignore[assignment]
@@ -2686,10 +2686,10 @@ class ResponseStream(AsyncIterable[TUpdate], Generic[TUpdate, TFinal]):
return self._stream
return self._stream # type: ignore[return-value]
def __aiter__(self) -> ResponseStream[TUpdate, TFinal]:
def __aiter__(self) -> ResponseStream[UpdateT, FinalT]:
return self
async def __anext__(self) -> TUpdate:
async def __anext__(self) -> UpdateT:
if self._iterator is None:
stream = await self._get_stream()
self._iterator = stream.__aiter__()
@@ -2718,19 +2718,19 @@ class ResponseStream(AsyncIterable[TUpdate], Generic[TUpdate, TFinal]):
return update
def __await__(self) -> Any:
async def _wrap() -> ResponseStream[TUpdate, TFinal]:
async def _wrap() -> ResponseStream[UpdateT, FinalT]:
await self._get_stream()
return self
return _wrap().__await__()
async def get_final_response(self) -> TFinal:
async def get_final_response(self) -> FinalT:
"""Get the final response by applying the finalizer to all collected updates.
If a finalizer is configured, it receives the list of updates and returns the final type.
Result hooks are then applied in order to transform the result.
If no finalizer is configured, returns the collected updates as Sequence[TUpdate].
If no finalizer is configured, returns the collected updates as Sequence[UpdateT].
For wrapped streams (created via .map() or .from_awaitable()):
- The inner stream's finalizer is called first to produce the inner final result.
@@ -2815,16 +2815,16 @@ class ResponseStream(AsyncIterable[TUpdate], Generic[TUpdate, TFinal]):
def with_transform_hook(
self,
hook: Callable[[TUpdate], TUpdate | Awaitable[TUpdate] | None],
) -> ResponseStream[TUpdate, TFinal]:
hook: Callable[[UpdateT], UpdateT | Awaitable[UpdateT] | None],
) -> ResponseStream[UpdateT, FinalT]:
"""Register a transform hook executed for each update during iteration."""
self._transform_hooks.append(hook)
return self
def with_result_hook(
self,
hook: Callable[[TFinal], TFinal | Awaitable[TFinal | None] | None],
) -> ResponseStream[TUpdate, TFinal]:
hook: Callable[[FinalT], FinalT | Awaitable[FinalT | None] | None],
) -> ResponseStream[UpdateT, FinalT]:
"""Register a result hook executed after finalization."""
self._result_hooks.append(hook)
self._finalized = False
@@ -2834,7 +2834,7 @@ class ResponseStream(AsyncIterable[TUpdate], Generic[TUpdate, TFinal]):
def with_cleanup_hook(
self,
hook: Callable[[], Awaitable[None] | None],
) -> ResponseStream[TUpdate, TFinal]:
) -> ResponseStream[UpdateT, FinalT]:
"""Register a cleanup hook executed after stream consumption (before finalizer)."""
self._cleanup_hooks.append(hook)
return self
@@ -2849,7 +2849,7 @@ class ResponseStream(AsyncIterable[TUpdate], Generic[TUpdate, TFinal]):
await result
@property
def updates(self) -> Sequence[TUpdate]:
def updates(self) -> Sequence[UpdateT]:
return self._updates
@@ -2944,8 +2944,8 @@ class _ChatOptionsBase(TypedDict, total=False):
if TYPE_CHECKING:
class ChatOptions(_ChatOptionsBase, Generic[TResponseModel], total=False):
response_format: type[TResponseModel] | Mapping[str, Any] | None # type: ignore[misc]
class ChatOptions(_ChatOptionsBase, Generic[ResponseModelT], total=False):
response_format: type[ResponseModelT] | Mapping[str, Any] | None # type: ignore[misc]
else:
ChatOptions = _ChatOptionsBase
@@ -9,7 +9,7 @@ if sys.version_info >= (3, 11):
else:
from typing_extensions import Self # pragma: no cover
TModel = TypeVar("TModel", bound="DictConvertible")
ModelT = TypeVar("ModelT", bound="DictConvertible")
class DictConvertible:
@@ -19,7 +19,7 @@ class DictConvertible:
raise NotImplementedError
@classmethod
def from_dict(cls: type[TModel], data: dict[str, Any]) -> TModel:
def from_dict(cls: type[ModelT], data: dict[str, Any]) -> ModelT:
return cls(**data) # type: ignore[arg-type]
def clone(self, *, deep: bool = True) -> Self:
@@ -31,7 +31,7 @@ class DictConvertible:
return json.dumps(self.to_dict())
@classmethod
def from_json(cls: type[TModel], raw: str) -> TModel:
def from_json(cls: type[ModelT], raw: str) -> ModelT:
import json
data = json.loads(raw)
@@ -32,8 +32,8 @@ __all__ = ["AzureOpenAIAssistantsClient"]
# region Azure OpenAI Assistants Options TypedDict
TAzureOpenAIAssistantsOptions = TypeVar(
"TAzureOpenAIAssistantsOptions",
AzureOpenAIAssistantsOptionsT = TypeVar(
"AzureOpenAIAssistantsOptionsT",
bound=TypedDict, # type: ignore[valid-type]
default="OpenAIAssistantsOptions",
covariant=True,
@@ -44,7 +44,7 @@ TAzureOpenAIAssistantsOptions = TypeVar(
class AzureOpenAIAssistantsClient(
OpenAIAssistantsClient[TAzureOpenAIAssistantsOptions], Generic[TAzureOpenAIAssistantsOptions]
OpenAIAssistantsClient[AzureOpenAIAssistantsOptionsT], Generic[AzureOpenAIAssistantsOptionsT]
):
"""Azure OpenAI Assistants client."""
@@ -53,7 +53,7 @@ logger: logging.Logger = logging.getLogger(__name__)
__all__ = ["AzureOpenAIChatClient", "AzureOpenAIChatOptions", "AzureUserSecurityContext"]
TResponseModel = TypeVar("TResponseModel", bound=BaseModel | None, default=None)
ResponseModelT = TypeVar("ResponseModelT", bound=BaseModel | None, default=None)
# region Azure OpenAI Chat Options TypedDict
@@ -81,7 +81,7 @@ class AzureUserSecurityContext(TypedDict, total=False):
"""The original client's IP address."""
class AzureOpenAIChatOptions(OpenAIChatOptions[TResponseModel], Generic[TResponseModel], total=False):
class AzureOpenAIChatOptions(OpenAIChatOptions[ResponseModelT], Generic[ResponseModelT], total=False):
"""Azure OpenAI-specific chat options dict.
Extends OpenAIChatOptions with Azure-specific options including
@@ -136,8 +136,8 @@ class AzureOpenAIChatOptions(OpenAIChatOptions[TResponseModel], Generic[TRespons
Note: You will be charged based on tokens across all choices. Keep n=1 to minimize costs."""
TAzureOpenAIChatOptions = TypeVar(
"TAzureOpenAIChatOptions",
AzureOpenAIChatOptionsT = TypeVar(
"AzureOpenAIChatOptionsT",
bound=TypedDict, # type: ignore[valid-type]
default="AzureOpenAIChatOptions",
covariant=True,
@@ -146,17 +146,17 @@ TAzureOpenAIChatOptions = TypeVar(
# endregion
TChatResponse = TypeVar("TChatResponse", ChatResponse, ChatResponseUpdate)
TAzureOpenAIChatClient = TypeVar("TAzureOpenAIChatClient", bound="AzureOpenAIChatClient")
ChatResponseT = TypeVar("ChatResponseT", ChatResponse, ChatResponseUpdate)
AzureOpenAIChatClientT = TypeVar("AzureOpenAIChatClientT", bound="AzureOpenAIChatClient")
class AzureOpenAIChatClient( # type: ignore[misc]
AzureOpenAIConfigMixin,
ChatMiddlewareLayer[TAzureOpenAIChatOptions],
FunctionInvocationLayer[TAzureOpenAIChatOptions],
ChatTelemetryLayer[TAzureOpenAIChatOptions],
RawOpenAIChatClient[TAzureOpenAIChatOptions],
Generic[TAzureOpenAIChatOptions],
ChatMiddlewareLayer[AzureOpenAIChatOptionsT],
FunctionInvocationLayer[AzureOpenAIChatOptionsT],
ChatTelemetryLayer[AzureOpenAIChatOptionsT],
RawOpenAIChatClient[AzureOpenAIChatOptionsT],
Generic[AzureOpenAIChatOptionsT],
):
"""Azure OpenAI Chat completion class with middleware, telemetry, and function invocation support."""
@@ -41,8 +41,8 @@ if TYPE_CHECKING:
__all__ = ["AzureOpenAIResponsesClient"]
TAzureOpenAIResponsesOptions = TypeVar(
"TAzureOpenAIResponsesOptions",
AzureOpenAIResponsesOptionsT = TypeVar(
"AzureOpenAIResponsesOptionsT",
bound=TypedDict, # type: ignore[valid-type]
default="OpenAIResponsesOptions",
covariant=True,
@@ -51,11 +51,11 @@ TAzureOpenAIResponsesOptions = TypeVar(
class AzureOpenAIResponsesClient( # type: ignore[misc]
AzureOpenAIConfigMixin,
ChatMiddlewareLayer[TAzureOpenAIResponsesOptions],
FunctionInvocationLayer[TAzureOpenAIResponsesOptions],
ChatTelemetryLayer[TAzureOpenAIResponsesOptions],
RawOpenAIResponsesClient[TAzureOpenAIResponsesOptions],
Generic[TAzureOpenAIResponsesOptions],
ChatMiddlewareLayer[AzureOpenAIResponsesOptionsT],
FunctionInvocationLayer[AzureOpenAIResponsesOptionsT],
ChatTelemetryLayer[AzureOpenAIResponsesOptionsT],
RawOpenAIResponsesClient[AzureOpenAIResponsesOptionsT],
Generic[AzureOpenAIResponsesOptionsT],
):
"""Azure Responses completion class with middleware, telemetry, and function invocation support."""
@@ -54,7 +54,7 @@ if TYPE_CHECKING: # pragma: no cover
ResponseStream,
)
TResponseModelT = TypeVar("TResponseModelT", bound=BaseModel)
ResponseModelBoundT = TypeVar("ResponseModelBoundT", bound=BaseModel)
__all__ = [
"OBSERVABILITY_SETTINGS",
@@ -71,7 +71,7 @@ __all__ = [
AgentT = TypeVar("AgentT", bound="SupportsAgentRun")
TChatClient = TypeVar("TChatClient", bound="ChatClientProtocol[Any]")
ChatClientT = TypeVar("ChatClientT", bound="ChatClientProtocol[Any]")
logger = get_logger()
@@ -1049,15 +1049,15 @@ def _get_token_usage_histogram() -> metrics.Histogram:
)
TOptions_co = TypeVar(
"TOptions_co",
OptionsCoT = TypeVar(
"OptionsCoT",
bound=TypedDict, # type: ignore[valid-type]
default="ChatOptions[None]",
covariant=True,
)
class ChatTelemetryLayer(Generic[TOptions_co]):
class ChatTelemetryLayer(Generic[OptionsCoT]):
"""Layer that wraps chat client get_response with OpenTelemetry tracing."""
def __init__(self, *args: Any, otel_provider_name: str | None = None, **kwargs: Any) -> None:
@@ -1073,9 +1073,9 @@ class ChatTelemetryLayer(Generic[TOptions_co]):
messages: str | ChatMessage | Sequence[str | ChatMessage],
*,
stream: Literal[False] = ...,
options: ChatOptions[TResponseModelT],
options: ChatOptions[ResponseModelBoundT],
**kwargs: Any,
) -> Awaitable[ChatResponse[TResponseModelT]]: ...
) -> Awaitable[ChatResponse[ResponseModelBoundT]]: ...
@overload
def get_response(
@@ -1083,7 +1083,7 @@ class ChatTelemetryLayer(Generic[TOptions_co]):
messages: str | ChatMessage | Sequence[str | ChatMessage],
*,
stream: Literal[False] = ...,
options: TOptions_co | ChatOptions[None] | None = None,
options: OptionsCoT | ChatOptions[None] | None = None,
**kwargs: Any,
) -> Awaitable[ChatResponse[Any]]: ...
@@ -1093,7 +1093,7 @@ class ChatTelemetryLayer(Generic[TOptions_co]):
messages: str | ChatMessage | Sequence[str | ChatMessage],
*,
stream: Literal[True],
options: TOptions_co | ChatOptions[Any] | None = None,
options: OptionsCoT | ChatOptions[Any] | None = None,
**kwargs: Any,
) -> ResponseStream[ChatResponseUpdate, ChatResponse[Any]]: ...
@@ -1102,7 +1102,7 @@ class ChatTelemetryLayer(Generic[TOptions_co]):
messages: str | ChatMessage | Sequence[str | ChatMessage],
*,
stream: bool = False,
options: TOptions_co | ChatOptions[Any] | None = None,
options: OptionsCoT | ChatOptions[Any] | None = None,
**kwargs: Any,
) -> Awaitable[ChatResponse[Any]] | ResponseStream[ChatResponseUpdate, ChatResponse[Any]]:
"""Trace chat responses with OpenTelemetry spans and metrics."""
@@ -33,10 +33,10 @@ else:
__all__ = ["OpenAIAssistantProvider"]
# Type variable for options - allows typed ChatAgent[TOptions] returns
# Type variable for options - allows typed OpenAIAssistantProvider[OptionsCoT] returns
# Default matches OpenAIAssistantsClient's default options type
TOptions_co = TypeVar(
"TOptions_co",
OptionsCoT = TypeVar(
"OptionsCoT",
bound=TypedDict, # type: ignore[valid-type]
default="OpenAIAssistantsOptions",
covariant=True,
@@ -50,7 +50,7 @@ _ToolsType = (
)
class OpenAIAssistantProvider(Generic[TOptions_co]):
class OpenAIAssistantProvider(Generic[OptionsCoT]):
"""Provider for creating ChatAgent instances from OpenAI Assistants API.
This provider allows you to create, retrieve, and wrap OpenAI Assistants
@@ -205,10 +205,10 @@ class OpenAIAssistantProvider(Generic[TOptions_co]):
description: str | None = None,
tools: _ToolsType | None = None,
metadata: dict[str, str] | None = None,
default_options: TOptions_co | None = None,
default_options: OptionsCoT | None = None,
middleware: Sequence[MiddlewareTypes] | None = None,
context_provider: ContextProvider | None = None,
) -> ChatAgent[TOptions_co]:
) -> ChatAgent[OptionsCoT]:
"""Create a new assistant on OpenAI and return a ChatAgent.
This method creates a new assistant on the OpenAI service and wraps it
@@ -313,10 +313,10 @@ class OpenAIAssistantProvider(Generic[TOptions_co]):
*,
tools: _ToolsType | None = None,
instructions: str | None = None,
default_options: TOptions_co | None = None,
default_options: OptionsCoT | None = None,
middleware: Sequence[MiddlewareTypes] | None = None,
context_provider: ContextProvider | None = None,
) -> ChatAgent[TOptions_co]:
) -> ChatAgent[OptionsCoT]:
"""Retrieve an existing assistant by ID and return a ChatAgent.
This method fetches an existing assistant from OpenAI by its ID
@@ -379,10 +379,10 @@ class OpenAIAssistantProvider(Generic[TOptions_co]):
*,
tools: _ToolsType | None = None,
instructions: str | None = None,
default_options: TOptions_co | None = None,
default_options: OptionsCoT | None = None,
middleware: Sequence[MiddlewareTypes] | None = None,
context_provider: ContextProvider | None = None,
) -> ChatAgent[TOptions_co]:
) -> ChatAgent[OptionsCoT]:
"""Wrap an existing SDK Assistant object as a ChatAgent.
This method does NOT make any HTTP calls. It simply wraps an already-
@@ -524,9 +524,9 @@ class OpenAIAssistantProvider(Generic[TOptions_co]):
instructions: str | None,
middleware: Sequence[MiddlewareTypes] | None,
context_provider: ContextProvider | None,
default_options: TOptions_co | None = None,
default_options: OptionsCoT | None = None,
**kwargs: Any,
) -> ChatAgent[TOptions_co]:
) -> ChatAgent[OptionsCoT]:
"""Create a ChatAgent from an Assistant.
Args:
@@ -79,7 +79,7 @@ __all__ = [
# region OpenAI Assistants Options TypedDict
TResponseModel = TypeVar("TResponseModel", bound=BaseModel | None, default=None)
ResponseModelT = TypeVar("ResponseModelT", bound=BaseModel | None, default=None)
class VectorStoreToolResource(TypedDict, total=False):
@@ -109,7 +109,7 @@ class AssistantToolResources(TypedDict, total=False):
"""Resources for file search tool, including vector store IDs."""
class OpenAIAssistantsOptions(ChatOptions[TResponseModel], Generic[TResponseModel], total=False):
class OpenAIAssistantsOptions(ChatOptions[ResponseModelT], Generic[ResponseModelT], total=False):
"""OpenAI Assistants API-specific options dict.
Extends base ChatOptions with Assistants API-specific parameters
@@ -193,8 +193,8 @@ ASSISTANTS_OPTION_TRANSLATIONS: dict[str, str] = {
}
"""Maps ChatOptions keys to OpenAI Assistants API parameter names."""
TOpenAIAssistantsOptions = TypeVar(
"TOpenAIAssistantsOptions",
OpenAIAssistantsOptionsT = TypeVar(
"OpenAIAssistantsOptionsT",
bound=TypedDict, # type: ignore[valid-type]
default="OpenAIAssistantsOptions",
covariant=True,
@@ -206,11 +206,11 @@ TOpenAIAssistantsOptions = TypeVar(
class OpenAIAssistantsClient( # type: ignore[misc]
OpenAIConfigMixin,
ChatMiddlewareLayer[TOpenAIAssistantsOptions],
FunctionInvocationLayer[TOpenAIAssistantsOptions],
ChatTelemetryLayer[TOpenAIAssistantsOptions],
BaseChatClient[TOpenAIAssistantsOptions],
Generic[TOpenAIAssistantsOptions],
ChatMiddlewareLayer[OpenAIAssistantsOptionsT],
FunctionInvocationLayer[OpenAIAssistantsOptionsT],
ChatTelemetryLayer[OpenAIAssistantsOptionsT],
BaseChatClient[OpenAIAssistantsOptionsT],
Generic[OpenAIAssistantsOptionsT],
):
"""OpenAI Assistants client with middleware, telemetry, and function invocation support."""
@@ -65,7 +65,7 @@ __all__ = ["OpenAIChatClient", "OpenAIChatOptions"]
logger = get_logger("agent_framework.openai")
TResponseModel = TypeVar("TResponseModel", bound=BaseModel | None, default=None)
ResponseModelT = TypeVar("ResponseModelT", bound=BaseModel | None, default=None)
# region OpenAI Chat Options TypedDict
@@ -85,7 +85,7 @@ class Prediction(TypedDict, total=False):
content: str | list[PredictionTextContent]
class OpenAIChatOptions(ChatOptions[TResponseModel], Generic[TResponseModel], total=False):
class OpenAIChatOptions(ChatOptions[ResponseModelT], Generic[ResponseModelT], total=False):
"""OpenAI-specific chat options dict.
Extends ChatOptions with options specific to OpenAI's Chat Completions API.
@@ -124,7 +124,7 @@ class OpenAIChatOptions(ChatOptions[TResponseModel], Generic[TResponseModel], to
prediction: Prediction
TOpenAIChatOptions = TypeVar("TOpenAIChatOptions", bound=TypedDict, default="OpenAIChatOptions", covariant=True) # type: ignore[valid-type]
OpenAIChatOptionsT = TypeVar("OpenAIChatOptionsT", bound=TypedDict, default="OpenAIChatOptions", covariant=True) # type: ignore[valid-type]
OPTION_TRANSLATIONS: dict[str, str] = {
"model_id": "model",
@@ -136,8 +136,8 @@ OPTION_TRANSLATIONS: dict[str, str] = {
# region Base Client
class RawOpenAIChatClient( # type: ignore[misc]
OpenAIBase,
BaseChatClient[TOpenAIChatOptions],
Generic[TOpenAIChatOptions],
BaseChatClient[OpenAIChatOptionsT],
Generic[OpenAIChatOptionsT],
):
"""Raw OpenAI Chat completion class without middleware, telemetry, or function invocation.
@@ -593,11 +593,11 @@ class RawOpenAIChatClient( # type: ignore[misc]
class OpenAIChatClient( # type: ignore[misc]
OpenAIConfigMixin,
ChatMiddlewareLayer[TOpenAIChatOptions],
FunctionInvocationLayer[TOpenAIChatOptions],
ChatTelemetryLayer[TOpenAIChatOptions],
RawOpenAIChatClient[TOpenAIChatOptions],
Generic[TOpenAIChatOptions],
ChatMiddlewareLayer[OpenAIChatOptionsT],
FunctionInvocationLayer[OpenAIChatOptionsT],
ChatTelemetryLayer[OpenAIChatOptionsT],
RawOpenAIChatClient[OpenAIChatOptionsT],
Generic[OpenAIChatOptionsT],
):
"""OpenAI Chat completion class with middleware, telemetry, and function invocation support."""
@@ -124,10 +124,10 @@ class StreamOptions(TypedDict, total=False):
"""Whether to include usage statistics in stream events."""
TResponseFormat = TypeVar("TResponseFormat", bound=BaseModel | None, default=None)
ResponseFormatT = TypeVar("ResponseFormatT", bound=BaseModel | None, default=None)
class OpenAIResponsesOptions(ChatOptions[TResponseFormat], Generic[TResponseFormat], total=False):
class OpenAIResponsesOptions(ChatOptions[ResponseFormatT], Generic[ResponseFormatT], total=False):
"""OpenAI Responses API-specific chat options.
Extends ChatOptions with options specific to OpenAI's Responses API.
@@ -191,8 +191,8 @@ class OpenAIResponsesOptions(ChatOptions[TResponseFormat], Generic[TResponseForm
- 'disabled': Fail with 400 error if exceeds context"""
TOpenAIResponsesOptions = TypeVar(
"TOpenAIResponsesOptions",
OpenAIResponsesOptionsT = TypeVar(
"OpenAIResponsesOptionsT",
bound=TypedDict, # type: ignore[valid-type]
default="OpenAIResponsesOptions",
covariant=True,
@@ -207,8 +207,8 @@ TOpenAIResponsesOptions = TypeVar(
class RawOpenAIResponsesClient( # type: ignore[misc]
OpenAIBase,
BaseChatClient[TOpenAIResponsesOptions],
Generic[TOpenAIResponsesOptions],
BaseChatClient[OpenAIResponsesOptionsT],
Generic[OpenAIResponsesOptionsT],
):
"""Raw OpenAI Responses client without middleware, telemetry, or function invocation.
@@ -1437,11 +1437,11 @@ class RawOpenAIResponsesClient( # type: ignore[misc]
class OpenAIResponsesClient( # type: ignore[misc]
OpenAIConfigMixin,
ChatMiddlewareLayer[TOpenAIResponsesOptions],
FunctionInvocationLayer[TOpenAIResponsesOptions],
ChatTelemetryLayer[TOpenAIResponsesOptions],
RawOpenAIResponsesClient[TOpenAIResponsesOptions],
Generic[TOpenAIResponsesOptions],
ChatMiddlewareLayer[OpenAIResponsesOptionsT],
FunctionInvocationLayer[OpenAIResponsesOptionsT],
ChatTelemetryLayer[OpenAIResponsesOptionsT],
RawOpenAIResponsesClient[OpenAIResponsesOptionsT],
Generic[OpenAIResponsesOptionsT],
):
"""OpenAI Responses client class with middleware, telemetry, and function invocation support."""
+6 -6
View File
@@ -27,7 +27,7 @@ from agent_framework import (
ToolProtocol,
tool,
)
from agent_framework._clients import TOptions_co
from agent_framework._clients import OptionsCoT
from agent_framework.observability import ChatTelemetryLayer
if sys.version_info >= (3, 12):
@@ -135,11 +135,11 @@ class MockChatClient:
class MockBaseChatClient(
ChatMiddlewareLayer[TOptions_co],
FunctionInvocationLayer[TOptions_co],
ChatTelemetryLayer[TOptions_co],
BaseChatClient[TOptions_co],
Generic[TOptions_co],
ChatMiddlewareLayer[OptionsCoT],
FunctionInvocationLayer[OptionsCoT],
ChatTelemetryLayer[OptionsCoT],
BaseChatClient[OptionsCoT],
Generic[OptionsCoT],
):
"""Mock implementation of a full-featured ChatClient."""