diff --git a/.agents/skills/implementation-strategy/SKILL.md b/.agents/skills/implementation-strategy/SKILL.md index df697dd7d3..a18cf574f2 100644 --- a/.agents/skills/implementation-strategy/SKILL.md +++ b/.agents/skills/implementation-strategy/SKILL.md @@ -43,6 +43,8 @@ Use this skill before editing code when the task changes runtime behavior or any - When unsupported OpenAI API or provider-adapter behavior already has a released default path, avoid turning it into a default hard error unless the latest release boundary justifies that break. Prefer an opt-in strict mode such as `strict_feature_validation=True`, while keeping the default path compatible through warning, ignoring unsupported data, or a clearly non-empty placeholder. - For OpenAI API feature gaps, evaluate streaming and non-streaming paths together. Custom tool calls, multi-choice Chat Completions chunks, non-text tool outputs, and similar provider payload differences must not be strict in one path and permissive or malformed in the other. - When a change creates new public SDK behavior, do not expose it only through hard-coded module globals. Prefer an explicit public configuration object or parameter, preserve the existing default behavior when compatibility-sensitive, and make opt-in SDK defaults explicit. +- For SDK-owned public configuration, accept existing typed objects and equivalent dictionaries at the public input boundary while preserving the internal typed representation. Respect the owning model's validation and extra-field policy instead of recreating arbitrary third-party schema semantics. +- Keep model-specific settings inside the existing `model_settings` parameter. Preserve released constructor arguments, typed-object behavior, and provider request payloads when adding dictionary support. - Append new optional fields or constructor parameters to public dataclasses and constructors. Do not insert them before existing public fields unless you also provide a compatibility layer and regression coverage for the old positional call shape. - Treat threshold and quota values as part of the API design when they affect runtime behavior. Distinguish OpenAI platform quota-derived values from defensive SDK defaults; if the value is not anchored in a documented platform limit, avoid making it an unconditional default-on behavior. - Define `None` semantics deliberately for public configuration. For example, use separate meanings for "feature disabled or no SDK limit", "use SDK default limits", and "disable only this specific limit" rather than relying on implicit truthiness checks. diff --git a/src/agents/__init__.py b/src/agents/__init__.py index 2ec18c2460..515916c2d1 100644 --- a/src/agents/__init__.py +++ b/src/agents/__init__.py @@ -314,7 +314,7 @@ def set_default_openai_responses_transport(transport: Literal["http", "websocket def set_default_openai_agent_registration( - config: OpenAIAgentRegistrationConfig | None, + config: OpenAIAgentRegistrationConfig | dict[str, Any] | None, ) -> None: """Set the default OpenAI agent registration config. diff --git a/src/agents/_config.py b/src/agents/_config.py index e5bdd3d0d7..846debd43b 100644 --- a/src/agents/_config.py +++ b/src/agents/_config.py @@ -1,4 +1,4 @@ -from typing import Literal +from typing import Any, Literal from openai import AsyncOpenAI @@ -40,7 +40,7 @@ def set_default_openai_responses_transport(transport: Literal["http", "websocket def set_default_openai_agent_registration( - config: OpenAIAgentRegistrationConfig | None, + config: OpenAIAgentRegistrationConfig | dict[str, Any] | None, ) -> None: set_default_openai_agent_registration_config(config) diff --git a/src/agents/_config_coercion.py b/src/agents/_config_coercion.py new file mode 100644 index 0000000000..2d0722f3d6 --- /dev/null +++ b/src/agents/_config_coercion.py @@ -0,0 +1,97 @@ +from __future__ import annotations + +from dataclasses import fields, is_dataclass +from types import UnionType +from typing import Any, TypeVar, Union, cast, get_args, get_origin, get_type_hints + +from pydantic import AliasChoices, BaseModel + +ConfigT = TypeVar("ConfigT") +DataclassConfigT = TypeVar("DataclassConfigT") +PydanticConfigT = TypeVar("PydanticConfigT", bound=BaseModel) + + +def _declared_dataclass_type( + owner_type: type[Any], + field_name: str, + default_type: type[DataclassConfigT], +) -> type[DataclassConfigT]: + try: + annotation = get_type_hints(owner_type).get(field_name) + except (NameError, TypeError): + return default_type + + candidates = ( + get_args(annotation) if get_origin(annotation) in (Union, UnionType) else (annotation,) + ) + for candidate in candidates: + if ( + isinstance(candidate, type) + and is_dataclass(candidate) + and issubclass(candidate, default_type) + ): + return candidate + return default_type + + +def _dataclass_input_values( + value: dict[str, Any], + config_type: type[Any], +) -> dict[str, Any]: + field_names = {config_field.name for config_field in fields(config_type)} + return {name: field_value for name, field_value in value.items() if name in field_names} + + +def coerce_dataclass_config( + value: ConfigT | dict[str, Any], + config_type: type[ConfigT], + *, + parameter_name: str, +) -> ConfigT: + """Normalize an SDK-owned dataclass configuration at its public input boundary.""" + if isinstance(value, config_type): + return value + if not isinstance(value, dict): + raise TypeError( + f"{parameter_name} must be a {config_type.__name__} instance or a dict, " + f"got {type(value).__name__}" + ) + + field_names = { + config_field.name for config_field in fields(cast(Any, config_type)) if config_field.init + } + unknown_fields = sorted(str(name) for name in value if name not in field_names) + if unknown_fields: + raise TypeError(f"Unknown {parameter_name} settings: {', '.join(unknown_fields)}") + return config_type(**value) + + +def coerce_pydantic_config( + value: PydanticConfigT | dict[str, Any], + config_type: type[PydanticConfigT], + *, + parameter_name: str, +) -> PydanticConfigT: + """Normalize an SDK-owned Pydantic configuration using its declared extra policy.""" + if isinstance(value, config_type): + return value + if not isinstance(value, dict): + raise TypeError( + f"{parameter_name} must be a {config_type.__name__} instance or a dict, " + f"got {type(value).__name__}" + ) + + if config_type.model_config.get("extra") != "allow": + accepted_fields: set[str] = set(config_type.model_fields) + for field_info in config_type.model_fields.values(): + if isinstance(field_info.validation_alias, str): + accepted_fields.add(field_info.validation_alias) + elif isinstance(field_info.validation_alias, AliasChoices): + accepted_fields.update( + alias for alias in field_info.validation_alias.choices if isinstance(alias, str) + ) + unknown_fields = sorted(str(name) for name in value if name not in accepted_fields) + if unknown_fields: + raise TypeError(f"Unknown {parameter_name} settings: {', '.join(unknown_fields)}") + + return config_type.model_validate(value) diff --git a/src/agents/agent.py b/src/agents/agent.py index e29d56801a..4f3c54a074 100644 --- a/src/agents/agent.py +++ b/src/agents/agent.py @@ -31,7 +31,7 @@ from .handoffs import Handoff from .logger import logger from .mcp import MCPUtil -from .model_settings import ModelSettings +from .model_settings import ModelSettings, _coerce_model_settings, _declared_model_settings_type from .models.default_models import ( get_default_model_settings, ) @@ -317,6 +317,8 @@ class Agent(AgentBase, Generic[TContext]): model_settings: ModelSettings = field(default_factory=get_default_model_settings) """Configures model-specific tuning parameters (e.g. temperature, top_p). + + Accepts a ``ModelSettings`` instance or a dictionary containing its fields. """ input_guardrails: list[InputGuardrail[TContext]] = field(default_factory=list) @@ -368,6 +370,39 @@ class Agent(AgentBase, Generic[TContext]): """Whether to reset the tool choice to the default value after a tool has been called. Defaults to True. This ensures that the agent doesn't enter an infinite loop of tool usage.""" + if TYPE_CHECKING: + + def __init__( + self, + name: str, + handoff_description: str | None = None, + tools: list[Tool] = ..., + mcp_servers: list[MCPServer] = ..., + mcp_config: MCPConfig = ..., + instructions: ( + str + | Callable[ + [RunContextWrapper[TContext], Agent[TContext]], + MaybeAwaitable[str], + ] + | None + ) = None, + prompt: Prompt | DynamicPromptFunction | None = None, + handoffs: list[Agent[Any] | Handoff[TContext, Any]] = ..., + model: str | Model | None = None, + model_settings: ModelSettings | dict[str, Any] = ..., + input_guardrails: list[InputGuardrail[TContext]] = ..., + output_guardrails: list[OutputGuardrail[TContext]] = ..., + output_type: type[Any] | AgentOutputSchemaBase | None = None, + hooks: AgentHooks[TContext] | None = None, + tool_use_behavior: ( + Literal["run_llm_again", "stop_on_first_tool"] + | StopAtTools + | ToolsToFinalOutputFunction + ) = "run_llm_again", + reset_tool_choice: bool = True, + ) -> None: ... + def __post_init__(self): from typing import get_origin @@ -424,11 +459,11 @@ def __post_init__(self): f"Agent model must be a string, Model, or None, got {type(self.model).__name__}" ) - if not isinstance(self.model_settings, ModelSettings): - raise TypeError( - f"Agent model_settings must be a ModelSettings instance, " - f"got {type(self.model_settings).__name__}" - ) + self.model_settings = _coerce_model_settings( + self.model_settings, + parameter_name="Agent model_settings", + model_settings_type=_declared_model_settings_type(type(self), "model_settings"), + ) if self.model is not None and self.model_settings == get_default_model_settings(): self.model_settings = _initial_model_settings_for_model(self.model) @@ -503,6 +538,13 @@ def clone(self, **kwargs: Any) -> Agent[TContext]: and _model_settings_match_implicit_model_defaults(self.model, self.model_settings) ): kwargs["model_settings"] = _initial_model_settings_for_model(kwargs["model"]) + if "model_settings" in kwargs: + kwargs["model_settings"] = _coerce_model_settings( + kwargs["model_settings"], + parameter_name="Agent model_settings", + model_settings_type=type(self.model_settings), + inherited_model_settings=self.model_settings, + ) return dataclasses.replace(self, **kwargs) def as_tool( @@ -515,7 +557,7 @@ def as_tool( is_enabled: bool | Callable[[RunContextWrapper[Any], AgentBase[Any]], MaybeAwaitable[bool]] = True, on_stream: Callable[[AgentToolStreamEvent], MaybeAwaitable[None]] | None = None, - run_config: RunConfig | None = None, + run_config: RunConfig | dict[str, Any] | None = None, max_turns: int | None = None, hooks: RunHooks[TContext] | None = None, previous_response_id: str | None = None, @@ -558,6 +600,11 @@ def as_tool( include_input_schema: Whether to include the full JSON schema in structured input. """ + if run_config is not None: + from .run_config import _coerce_run_config + + run_config = _coerce_run_config(run_config) + def _is_supported_parameters(value: Any) -> bool: if not isinstance(value, type): return False diff --git a/src/agents/extensions/memory/advanced_sqlite_session.py b/src/agents/extensions/memory/advanced_sqlite_session.py index 98a54e3123..2dd1e947fe 100644 --- a/src/agents/extensions/memory/advanced_sqlite_session.py +++ b/src/agents/extensions/memory/advanced_sqlite_session.py @@ -41,7 +41,7 @@ def __init__( db_path: str | Path = ":memory:", create_tables: bool = False, logger: logging.Logger | None = None, - session_settings: SessionSettings | None = None, + session_settings: SessionSettings | dict[str, Any] | None = None, **kwargs, ): """Initialize the AdvancedSQLiteSession. diff --git a/src/agents/extensions/memory/async_sqlite_session.py b/src/agents/extensions/memory/async_sqlite_session.py index 27a23b1cbe..63ae77081b 100644 --- a/src/agents/extensions/memory/async_sqlite_session.py +++ b/src/agents/extensions/memory/async_sqlite_session.py @@ -5,13 +5,17 @@ from collections.abc import AsyncIterator from contextlib import asynccontextmanager from pathlib import Path -from typing import cast +from typing import Any, cast import aiosqlite from ...items import TResponseInputItem from ...memory import SessionABC -from ...memory.session_settings import SessionSettings, resolve_session_limit +from ...memory.session_settings import ( + SessionSettings, + coerce_session_settings, + resolve_session_limit, +) class AsyncSQLiteSession(SessionABC): @@ -30,7 +34,7 @@ def __init__( db_path: str | Path = ":memory:", sessions_table: str = "agent_sessions", messages_table: str = "agent_messages", - session_settings: SessionSettings | None = None, + session_settings: SessionSettings | dict[str, Any] | None = None, ): """Initialize the async SQLite session. @@ -44,7 +48,11 @@ def __init__( retrieving items. If None, uses default SessionSettings(). """ self.session_id = session_id - self.session_settings = session_settings or SessionSettings() + self.session_settings = ( + coerce_session_settings(session_settings) + if session_settings is not None + else SessionSettings() + ) self.db_path = db_path self.sessions_table = sessions_table self.messages_table = messages_table diff --git a/src/agents/extensions/memory/dapr_session.py b/src/agents/extensions/memory/dapr_session.py index 6ac68f6020..eaed2574f5 100644 --- a/src/agents/extensions/memory/dapr_session.py +++ b/src/agents/extensions/memory/dapr_session.py @@ -45,7 +45,11 @@ from ...items import TResponseInputItem from ...logger import logger from ...memory.session import SessionABC -from ...memory.session_settings import SessionSettings, resolve_session_limit +from ...memory.session_settings import ( + SessionSettings, + coerce_session_settings, + resolve_session_limit, +) # Type alias for consistency levels ConsistencyLevel = Literal["eventual", "strong"] @@ -72,7 +76,7 @@ def __init__( dapr_client: DaprClient, ttl: int | None = None, consistency: ConsistencyLevel = DAPR_CONSISTENCY_EVENTUAL, - session_settings: SessionSettings | None = None, + session_settings: SessionSettings | dict[str, Any] | None = None, ): """Initializes a new DaprSession. @@ -90,7 +94,11 @@ def __init__( default limit for retrieving items. If None, uses default SessionSettings(). """ self.session_id = session_id - self.session_settings = session_settings or SessionSettings() + self.session_settings = ( + coerce_session_settings(session_settings) + if session_settings is not None + else SessionSettings() + ) self._dapr_client = dapr_client self._state_store_name = state_store_name self._ttl = ttl @@ -109,7 +117,7 @@ def from_address( *, state_store_name: str, dapr_address: str = "localhost:50001", - session_settings: SessionSettings | None = None, + session_settings: SessionSettings | dict[str, Any] | None = None, **kwargs: Any, ) -> DaprSession: """Create a session from a Dapr sidecar address. diff --git a/src/agents/extensions/memory/mongodb_session.py b/src/agents/extensions/memory/mongodb_session.py index 07354577d6..98f7f26008 100644 --- a/src/agents/extensions/memory/mongodb_session.py +++ b/src/agents/extensions/memory/mongodb_session.py @@ -60,7 +60,11 @@ from ...items import TResponseInputItem from ...memory.session import SessionABC -from ...memory.session_settings import SessionSettings, resolve_session_limit +from ...memory.session_settings import ( + SessionSettings, + coerce_session_settings, + resolve_session_limit, +) # Identifies this library in the MongoDB handshake for server-side telemetry. _DRIVER_INFO = DriverInfo(name="openai-agents", version=_VERSION) @@ -110,7 +114,7 @@ def __init__( database: str = "agents", sessions_collection: str = "agent_sessions", messages_collection: str = "agent_messages", - session_settings: SessionSettings | None = None, + session_settings: SessionSettings | dict[str, Any] | None = None, ): """Initialize a new MongoDBSession. @@ -128,7 +132,11 @@ def __init__( is used (no item limit). """ self.session_id = session_id - self.session_settings = session_settings or SessionSettings() + self.session_settings = ( + coerce_session_settings(session_settings) + if session_settings is not None + else SessionSettings() + ) self._client = client self._owns_client = False @@ -153,7 +161,7 @@ def from_uri( uri: str, database: str = "agents", client_kwargs: dict[str, Any] | None = None, - session_settings: SessionSettings | None = None, + session_settings: SessionSettings | dict[str, Any] | None = None, **kwargs: Any, ) -> MongoDBSession: """Create a session from a MongoDB URI string. diff --git a/src/agents/extensions/memory/redis_session.py b/src/agents/extensions/memory/redis_session.py index 11e2dd838b..3ad261b28e 100644 --- a/src/agents/extensions/memory/redis_session.py +++ b/src/agents/extensions/memory/redis_session.py @@ -41,7 +41,11 @@ from ...items import TResponseInputItem from ...memory.session import SessionABC -from ...memory.session_settings import SessionSettings, resolve_session_limit +from ...memory.session_settings import ( + SessionSettings, + coerce_session_settings, + resolve_session_limit, +) class RedisSession(SessionABC): @@ -56,7 +60,7 @@ def __init__( redis_client: Redis, key_prefix: str = "agents:session", ttl: int | None = None, - session_settings: SessionSettings | None = None, + session_settings: SessionSettings | dict[str, Any] | None = None, ): """Initializes a new RedisSession. @@ -71,7 +75,11 @@ def __init__( default limit for retrieving items. If None, uses default SessionSettings(). """ self.session_id = session_id - self.session_settings = session_settings or SessionSettings() + self.session_settings = ( + coerce_session_settings(session_settings) + if session_settings is not None + else SessionSettings() + ) self._redis = redis_client self._key_prefix = key_prefix self._ttl = ttl @@ -90,7 +98,7 @@ def from_url( *, url: str, redis_kwargs: dict[str, Any] | None = None, - session_settings: SessionSettings | None = None, + session_settings: SessionSettings | dict[str, Any] | None = None, **kwargs: Any, ) -> RedisSession: """Create a session from a Redis URL string. diff --git a/src/agents/extensions/memory/sqlalchemy_session.py b/src/agents/extensions/memory/sqlalchemy_session.py index 89467ad2d2..3fc793d328 100644 --- a/src/agents/extensions/memory/sqlalchemy_session.py +++ b/src/agents/extensions/memory/sqlalchemy_session.py @@ -50,7 +50,11 @@ from ...items import TResponseInputItem from ...memory.session import SessionABC -from ...memory.session_settings import SessionSettings, resolve_session_limit +from ...memory.session_settings import ( + SessionSettings, + coerce_session_settings, + resolve_session_limit, +) class SQLAlchemySession(SessionABC): @@ -135,7 +139,7 @@ def __init__( create_tables: bool = False, sessions_table: str = "agent_sessions", messages_table: str = "agent_messages", - session_settings: SessionSettings | None = None, + session_settings: SessionSettings | dict[str, Any] | None = None, ensure_ascii: bool = True, ): """Initializes a new SQLAlchemySession. @@ -155,7 +159,11 @@ def __init__( session items to JSON. Defaults to True to preserve the historical storage format. """ self.session_id = session_id - self.session_settings = session_settings or SessionSettings() + self.session_settings = ( + coerce_session_settings(session_settings) + if session_settings is not None + else SessionSettings() + ) self._engine = engine self._ensure_ascii = ensure_ascii self._configure_sqlite_engine(engine) @@ -225,7 +233,7 @@ def from_url( *, url: str, engine_kwargs: dict[str, Any] | None = None, - session_settings: SessionSettings | None = None, + session_settings: SessionSettings | dict[str, Any] | None = None, **kwargs: Any, ) -> SQLAlchemySession: """Create a session from a database URL string. diff --git a/src/agents/memory/openai_conversations_session.py b/src/agents/memory/openai_conversations_session.py index 0220eccbb1..9114a7dea0 100644 --- a/src/agents/memory/openai_conversations_session.py +++ b/src/agents/memory/openai_conversations_session.py @@ -1,6 +1,7 @@ from __future__ import annotations import asyncio +from typing import Any from openai import AsyncOpenAI @@ -8,7 +9,7 @@ from ..items import TResponseInputItem from .session import SessionABC -from .session_settings import SessionSettings, resolve_session_limit +from .session_settings import SessionSettings, coerce_session_settings, resolve_session_limit async def start_openai_conversations_session(openai_client: AsyncOpenAI | None = None) -> str: @@ -30,11 +31,15 @@ def __init__( *, conversation_id: str | None = None, openai_client: AsyncOpenAI | None = None, - session_settings: SessionSettings | None = None, + session_settings: SessionSettings | dict[str, Any] | None = None, ): self._session_id: str | None = conversation_id self._session_id_lock = asyncio.Lock() - self.session_settings = session_settings or SessionSettings() + self.session_settings = ( + coerce_session_settings(session_settings) + if session_settings is not None + else SessionSettings() + ) _openai_client = openai_client if _openai_client is None: _openai_client = get_default_openai_client() or AsyncOpenAI() diff --git a/src/agents/memory/session_settings.py b/src/agents/memory/session_settings.py index 03dfbd8d23..eb42f617f2 100644 --- a/src/agents/memory/session_settings.py +++ b/src/agents/memory/session_settings.py @@ -8,16 +8,22 @@ from pydantic.dataclasses import dataclass +from .._config_coercion import ( + _dataclass_input_values, + _declared_dataclass_type, + coerce_dataclass_config, +) + def resolve_session_limit( explicit_limit: int | None, - settings: SessionSettings | None, + settings: SessionSettings | dict[str, Any] | None, ) -> int | None: """Safely resolve the effective limit for session operations.""" if explicit_limit is not None: return explicit_limit if settings is not None: - return settings.limit + return coerce_session_settings(settings).limit return None @@ -32,16 +38,23 @@ class SessionSettings: limit: int | None = None """Maximum number of items to retrieve. If None, retrieves all items.""" - def resolve(self, override: SessionSettings | None) -> SessionSettings: + def resolve(self, override: SessionSettings | dict[str, Any] | None) -> SessionSettings: """Produce a new SessionSettings by overlaying any non-None values from the override on top of this instance.""" if override is None: return self + override_fields = ( + set(_dataclass_input_values(override, type(self))) + if isinstance(override, dict) + else None + ) + override = _coerce_session_settings(override, settings_type=type(self)) changes = { field.name: getattr(override, field.name) for field in fields(self) - if getattr(override, field.name) is not None + if (override_fields is None or field.name in override_fields) + and getattr(override, field.name) is not None } return replace(self, **changes) @@ -49,3 +62,25 @@ def resolve(self, override: SessionSettings | None) -> SessionSettings: def to_dict(self) -> dict[str, Any]: """Convert settings to a dictionary.""" return dataclasses.asdict(self) + + +def coerce_session_settings( + value: SessionSettings | dict[str, Any], +) -> SessionSettings: + """Normalize session settings while preserving existing typed instances.""" + return _coerce_session_settings(value, settings_type=SessionSettings) + + +def _coerce_session_settings( + value: SessionSettings | dict[str, Any], + *, + settings_type: type[SessionSettings], +) -> SessionSettings: + return coerce_dataclass_config(value, settings_type, parameter_name="session") + + +def _declared_session_settings_type( + owner_type: type[Any], + field_name: str, +) -> type[SessionSettings]: + return _declared_dataclass_type(owner_type, field_name, SessionSettings) diff --git a/src/agents/memory/sqlite_session.py b/src/agents/memory/sqlite_session.py index 3a69f9883a..b57f3ebf5a 100644 --- a/src/agents/memory/sqlite_session.py +++ b/src/agents/memory/sqlite_session.py @@ -7,11 +7,11 @@ from collections.abc import Iterator from contextlib import contextmanager from pathlib import Path -from typing import ClassVar +from typing import Any, ClassVar from ..items import TResponseInputItem from .session import SessionABC -from .session_settings import SessionSettings, resolve_session_limit +from .session_settings import SessionSettings, coerce_session_settings, resolve_session_limit class SQLiteSession(SessionABC): @@ -33,7 +33,7 @@ def __init__( db_path: str | Path = ":memory:", sessions_table: str = "agent_sessions", messages_table: str = "agent_messages", - session_settings: SessionSettings | None = None, + session_settings: SessionSettings | dict[str, Any] | None = None, ): """Initialize the SQLite session. @@ -47,7 +47,11 @@ def __init__( retrieving items. If None, uses default SessionSettings(). """ self.session_id = session_id - self.session_settings = session_settings or SessionSettings() + self.session_settings = ( + coerce_session_settings(session_settings) + if session_settings is not None + else SessionSettings() + ) self.db_path = db_path self.sessions_table = sessions_table self.messages_table = messages_table diff --git a/src/agents/model_settings.py b/src/agents/model_settings.py index e35279b3c3..0d6c24b837 100644 --- a/src/agents/model_settings.py +++ b/src/agents/model_settings.py @@ -2,7 +2,7 @@ from collections.abc import Mapping from dataclasses import fields, replace -from typing import Annotated, Any, Literal, TypeAlias, cast +from typing import TYPE_CHECKING, Annotated, Any, Literal, TypeAlias, cast from openai import Omit as _Omit from openai._types import Body, Query @@ -13,6 +13,7 @@ from pydantic.dataclasses import dataclass from pydantic_core import core_schema +from ._config_coercion import _declared_dataclass_type, coerce_dataclass_config from .retry import ( ModelRetryBackoffInput, ModelRetryBackoffSettings, @@ -199,20 +200,58 @@ class ModelSettings: control which prompt prefixes are eligible for caching. """ - def resolve(self, override: ModelSettings | None) -> ModelSettings: + if TYPE_CHECKING: + + def __init__( + self, + temperature: float | None = None, + top_p: float | None = None, + frequency_penalty: float | None = None, + presence_penalty: float | None = None, + tool_choice: ToolChoice | dict[str, Any] = None, + parallel_tool_calls: bool | None = None, + truncation: Literal["auto", "disabled"] | None = None, + max_tokens: int | None = None, + reasoning: Reasoning | dict[str, Any] | None = None, + verbosity: Literal["low", "medium", "high"] | None = None, + metadata: dict[str, str] | None = None, + store: bool | None = None, + prompt_cache_retention: Literal["in_memory", "24h"] | None = None, + include_usage: bool | None = None, + response_include: list[ResponseIncludable | str] | None = None, + top_logprobs: int | None = None, + extra_query: Query | None = None, + extra_body: Body | None = None, + extra_headers: Headers | None = None, + extra_args: dict[str, Any] | None = None, + retry: ModelRetrySettings | dict[str, Any] | None = None, + context_management: list[ContextManagement] | None = None, + prompt_cache_options: PromptCacheOptions | None = None, + ) -> None: ... + + def resolve(self, override: ModelSettings | dict[str, Any] | None) -> ModelSettings: """Produce a new ModelSettings by overlaying any non-None values from the override on top of this instance.""" if override is None: return self + override_fields = set(override) if isinstance(override, dict) else None + override = _coerce_model_settings( + override, + parameter_name="ModelSettings override", + model_settings_type=type(self), + ) changes = { field.name: getattr(override, field.name) for field in fields(self) - if getattr(override, field.name) is not None + if (override_fields is None or field.name in override_fields) + and getattr(override, field.name, None) is not None } # Handle extra_args merging specially - merge dictionaries instead of replacing. - if self.extra_args is not None or override.extra_args is not None: + if (override_fields is None or "extra_args" in override_fields) and ( + self.extra_args is not None or override.extra_args is not None + ): merged_args = {} if self.extra_args: merged_args.update(self.extra_args) @@ -220,7 +259,9 @@ def resolve(self, override: ModelSettings | None) -> ModelSettings: merged_args.update(override.extra_args) changes["extra_args"] = merged_args if merged_args else None - if self.retry is not None or override.retry is not None: + if (override_fields is None or "retry" in override_fields) and ( + self.retry is not None or override.retry is not None + ): changes["retry"] = _merge_retry_settings(self.retry, override.retry) return replace(self, **changes) @@ -234,6 +275,83 @@ def to_traceable_dict(self) -> dict[str, Any]: return {key: payload[key] for key in _TRACEABLE_MODEL_SETTING_FIELDS if key in payload} +def _coerce_model_settings( + value: ModelSettings | dict[str, Any], + *, + parameter_name: str, + model_settings_type: type[ModelSettings] = ModelSettings, + inherited_model_settings: ModelSettings | None = None, +) -> ModelSettings: + """Normalize SDK-owned model settings without changing existing typed instances.""" + del inherited_model_settings + if isinstance(value, ModelSettings): + return value + if not isinstance(value, dict): + raise TypeError( + f"{parameter_name} must be a ModelSettings instance or a dict, " + f"got {type(value).__name__}" + ) + + field_names = {model_field.name for model_field in fields(model_settings_type)} + unknown_fields = sorted(str(name) for name in value if name not in field_names) + if unknown_fields: + raise TypeError(f"Unknown model settings: {', '.join(unknown_fields)}") + + _validate_first_party_model_settings(value) + return coerce_dataclass_config(value, model_settings_type, parameter_name=parameter_name) + + +def _declared_model_settings_type( + owner_type: type[Any], + field_name: str, +) -> type[ModelSettings]: + return _declared_dataclass_type(owner_type, field_name, ModelSettings) + + +def _validate_first_party_model_settings(value: dict[str, Any]) -> None: + """Reject SDK-owned structured-setting typos while preserving OpenAI model extras.""" + + def validate_fields(payload: object, names: set[str], path: str) -> None: + if not isinstance(payload, Mapping): + return + unknown_fields = sorted(str(name) for name in payload if name not in names) + if unknown_fields: + raise TypeError(f"Unknown model settings in {path}: {', '.join(unknown_fields)}") + + validate_fields( + value.get("tool_choice"), + {model_field.name for model_field in fields(MCPToolChoice)}, + "tool_choice", + ) + retry = value.get("retry") + validate_fields( + retry, + {model_field.name for model_field in fields(ModelRetrySettings)}, + "retry", + ) + if isinstance(retry, Mapping): + validate_fields( + retry.get("backoff"), + {model_field.name for model_field in fields(ModelRetryBackoffSettings)}, + "retry.backoff", + ) + + context_management = value.get("context_management") + if isinstance(context_management, list | tuple): + for index, item in enumerate(context_management): + validate_fields( + item, + set(ContextManagement.__annotations__), + f"context_management[{index}]", + ) + + validate_fields( + value.get("prompt_cache_options"), + set(PromptCacheOptions.__annotations__), + "prompt_cache_options", + ) + + def _merge_retry_settings( inherited: ModelRetrySettings | None, override: ModelRetrySettings | None, diff --git a/src/agents/models/multi_provider.py b/src/agents/models/multi_provider.py index 4737bb8c0c..ccb644edc2 100644 --- a/src/agents/models/multi_provider.py +++ b/src/agents/models/multi_provider.py @@ -1,6 +1,6 @@ from __future__ import annotations -from typing import Literal, cast +from typing import Any, Literal, cast from openai import AsyncOpenAI @@ -87,7 +87,7 @@ def __init__( openai_websocket_base_url: str | None = None, openai_prefix_mode: MultiProviderOpenAIPrefixMode = "alias", unknown_prefix_mode: MultiProviderUnknownPrefixMode = "error", - openai_agent_registration: OpenAIAgentRegistrationConfig | None = None, + openai_agent_registration: OpenAIAgentRegistrationConfig | dict[str, Any] | None = None, openai_responses_websocket_options: OpenAIResponsesWebSocketOptions | None = None, openai_buffer_streamed_tool_calls: bool = False, ) -> None: diff --git a/src/agents/models/openai_agent_registration.py b/src/agents/models/openai_agent_registration.py index 12e62d8ba0..e0578739bc 100644 --- a/src/agents/models/openai_agent_registration.py +++ b/src/agents/models/openai_agent_registration.py @@ -4,6 +4,8 @@ from dataclasses import dataclass from typing import Any +from .._config_coercion import coerce_dataclass_config + _ENV_HARNESS_ID = "OPENAI_AGENT_HARNESS_ID" OPENAI_HARNESS_ID_TRACE_METADATA_KEY = "agent_harness_id" @@ -22,10 +24,12 @@ class ResolvedOpenAIAgentRegistrationConfig: def set_default_openai_agent_registration_config( - config: OpenAIAgentRegistrationConfig | None, + config: OpenAIAgentRegistrationConfig | dict[str, Any] | None, ) -> None: global _default_agent_registration - _default_agent_registration = config + _default_agent_registration = ( + _coerce_openai_agent_registration_config(config) if config is not None else None + ) def get_default_openai_agent_registration_config() -> OpenAIAgentRegistrationConfig | None: @@ -33,8 +37,10 @@ def get_default_openai_agent_registration_config() -> OpenAIAgentRegistrationCon def resolve_openai_agent_registration_config( - config: OpenAIAgentRegistrationConfig | None, + config: OpenAIAgentRegistrationConfig | dict[str, Any] | None, ) -> ResolvedOpenAIAgentRegistrationConfig | None: + if config is not None: + config = _coerce_openai_agent_registration_config(config) default = get_default_openai_agent_registration_config() harness_id = _resolve_str( explicit=config.harness_id if config else None, @@ -46,6 +52,16 @@ def resolve_openai_agent_registration_config( return ResolvedOpenAIAgentRegistrationConfig(harness_id=harness_id) +def _coerce_openai_agent_registration_config( + config: OpenAIAgentRegistrationConfig | dict[str, Any], +) -> OpenAIAgentRegistrationConfig: + return coerce_dataclass_config( + config, + OpenAIAgentRegistrationConfig, + parameter_name="OpenAI agent registration", + ) + + def resolve_openai_harness_id_for_model_provider(model_provider: Any) -> str | None: """Return the configured harness ID for OpenAI-backed model providers.""" harness_id = _harness_id_from_model_provider(model_provider) diff --git a/src/agents/models/openai_provider.py b/src/agents/models/openai_provider.py index dd4b888cb2..cc88d14ef1 100644 --- a/src/agents/models/openai_provider.py +++ b/src/agents/models/openai_provider.py @@ -3,6 +3,7 @@ import asyncio import os import weakref +from typing import Any import httpx from openai import AsyncOpenAI, DefaultAsyncHttpxClient @@ -54,7 +55,7 @@ def __init__( use_responses: bool | None = None, use_responses_websocket: bool | None = None, strict_feature_validation: bool = False, - agent_registration: OpenAIAgentRegistrationConfig | None = None, + agent_registration: OpenAIAgentRegistrationConfig | dict[str, Any] | None = None, responses_websocket_options: OpenAIResponsesWebSocketOptions | None = None, buffer_streamed_tool_calls: bool = False, ) -> None: diff --git a/src/agents/responses_websocket_session.py b/src/agents/responses_websocket_session.py index 3d0f18137d..b1ac69d938 100644 --- a/src/agents/responses_websocket_session.py +++ b/src/agents/responses_websocket_session.py @@ -3,7 +3,7 @@ from collections.abc import AsyncIterator, Mapping from contextlib import asynccontextmanager from dataclasses import dataclass -from typing import Any +from typing import TYPE_CHECKING, Any from .agent import Agent from .items import TResponseInputItem @@ -16,7 +16,7 @@ from .models.openai_responses import OpenAIResponsesWebSocketOptions from .result import RunResult, RunResultStreaming from .run import Runner -from .run_config import RunConfig +from .run_config import RunConfig, _coerce_run_config from .run_state import RunState @@ -27,7 +27,16 @@ class ResponsesWebSocketSession: provider: OpenAIProvider run_config: RunConfig + if TYPE_CHECKING: + + def __init__( + self, + provider: OpenAIProvider, + run_config: RunConfig | dict[str, Any], + ) -> None: ... + def __post_init__(self) -> None: + object.__setattr__(self, "run_config", _coerce_run_config(self.run_config)) self._validate_provider_alignment() def _validate_provider_alignment(self) -> MultiProvider: diff --git a/src/agents/run.py b/src/agents/run.py index 3928fc52b5..04fe9c09a0 100644 --- a/src/agents/run.py +++ b/src/agents/run.py @@ -3,7 +3,7 @@ import asyncio import contextlib import warnings -from typing import cast +from typing import Any, cast from typing_extensions import Unpack @@ -42,6 +42,7 @@ ToolErrorFormatterArgs, ToolExecutionConfig, ToolNotFoundBehavior, + _coerce_run_config, ) from .run_context import RunContextWrapper, TContext from .run_error_handlers import RunErrorHandlers @@ -208,7 +209,7 @@ async def run( context: TContext | None = None, max_turns: int | None = DEFAULT_MAX_TURNS, hooks: RunHooks[TContext] | None = None, - run_config: RunConfig | None = None, + run_config: RunConfig | dict[str, Any] | None = None, error_handlers: RunErrorHandlers[TContext] | None = None, previous_response_id: str | None = None, auto_previous_response_id: bool = False, @@ -292,7 +293,7 @@ def run_sync( context: TContext | None = None, max_turns: int | None = DEFAULT_MAX_TURNS, hooks: RunHooks[TContext] | None = None, - run_config: RunConfig | None = None, + run_config: RunConfig | dict[str, Any] | None = None, error_handlers: RunErrorHandlers[TContext] | None = None, previous_response_id: str | None = None, auto_previous_response_id: bool = False, @@ -373,7 +374,7 @@ def run_streamed( context: TContext | None = None, max_turns: int | None = DEFAULT_MAX_TURNS, hooks: RunHooks[TContext] | None = None, - run_config: RunConfig | None = None, + run_config: RunConfig | dict[str, Any] | None = None, previous_response_id: str | None = None, auto_previous_response_id: bool = False, conversation_id: str | None = None, @@ -467,8 +468,7 @@ async def run( conversation_id = kwargs.get("conversation_id") session = kwargs.get("session") - if run_config is None: - run_config = RunConfig() + run_config = RunConfig() if run_config is None else _coerce_run_config(run_config) is_resumed_state = isinstance(input, RunState) run_state: RunState[TContext] | None = None @@ -1728,8 +1728,7 @@ def run_streamed( conversation_id = kwargs.get("conversation_id") session = kwargs.get("session") - if run_config is None: - run_config = RunConfig() + run_config = RunConfig() if run_config is None else _coerce_run_config(run_config) # Handle RunState input is_resumed_state = isinstance(input, RunState) diff --git a/src/agents/run_config.py b/src/agents/run_config.py index 08ee4cff9e..393e6dd039 100644 --- a/src/agents/run_config.py +++ b/src/agents/run_config.py @@ -5,14 +5,24 @@ from dataclasses import dataclass, field from typing import TYPE_CHECKING, Any, Generic, Literal +from pydantic import TypeAdapter from typing_extensions import NotRequired, TypedDict +from ._config_coercion import ( + _declared_dataclass_type, + coerce_dataclass_config, + coerce_pydantic_config, +) from .guardrail import InputGuardrail, OutputGuardrail from .handoffs import HandoffHistoryMapper, HandoffInputFilter from .items import TResponseInputItem from .lifecycle import RunHooks from .memory import Session, SessionInputCallback, SessionSettings -from .model_settings import ModelSettings +from .memory.session_settings import ( + _coerce_session_settings, + _declared_session_settings_type, +) +from .model_settings import ModelSettings, _coerce_model_settings, _declared_model_settings_type from .models.interface import Model, ModelProvider from .models.multi_provider import MultiProvider from .run_context import TContext @@ -207,6 +217,86 @@ class SandboxRunConfig: Use `SandboxArchiveLimits()` to enable SDK defaults. """ + if TYPE_CHECKING: + + def __init__( + self, + client: BaseSandboxClient[Any] | None = None, + options: Any | None = None, + session: BaseSandboxSession | None = None, + session_state: SandboxSessionState | None = None, + manifest: Manifest | dict[str, Any] | None = None, + snapshot: SnapshotSpec | SnapshotBase | dict[str, Any] | None = None, + concurrency_limits: SandboxConcurrencyLimits | dict[str, Any] = ..., + archive_limits: SandboxArchiveLimits | dict[str, Any] | None = None, + ) -> None: ... + + def __post_init__(self) -> None: + if isinstance(self.manifest, dict): + from .sandbox.manifest import _coerce_manifest + + self.manifest = _coerce_manifest(self.manifest, parameter_name="sandbox.manifest") + if isinstance(self.snapshot, dict): + from .sandbox.snapshot import SnapshotBase, SnapshotSpecUnion + + if "id" in self.snapshot: + self.snapshot = SnapshotBase.parse(self.snapshot) + else: + self.snapshot = TypeAdapter(SnapshotSpecUnion).validate_python(self.snapshot) + if isinstance(self.options, dict) and self.client is not None: + from .sandbox.session.sandbox_client import BaseSandboxClientOptions + + options_type = BaseSandboxClientOptions._options_class_for_type(self.client.backend_id) + if options_type is not None: + options = self.options + explicit_type = options.get("type") + if explicit_type is not None and explicit_type != self.client.backend_id: + raise ValueError( + f"sandbox.options type `{explicit_type}` does not match selected " + f"sandbox client backend `{self.client.backend_id}`" + ) + if "type" not in options: + options = { + **options, + "type": options_type.model_fields["type"].default, + } + self.options = coerce_pydantic_config( + options, + options_type, + parameter_name="sandbox.options", + ) + elif self.client.backend_id == "blaxel": + from .extensions.sandbox.blaxel.sandbox import ( + BlaxelSandboxClient, + BlaxelSandboxClientOptions, + ) + + if isinstance(self.client, BlaxelSandboxClient): + self.options = coerce_dataclass_config( + self.options, + BlaxelSandboxClientOptions, + parameter_name="sandbox.options", + ) + self.concurrency_limits = coerce_dataclass_config( + self.concurrency_limits, + _declared_dataclass_type( + type(self), + "concurrency_limits", + SandboxConcurrencyLimits, + ), + parameter_name="sandbox.concurrency_limits", + ) + if self.archive_limits is not None: + self.archive_limits = coerce_dataclass_config( + self.archive_limits, + _declared_dataclass_type( + type(self), + "archive_limits", + SandboxArchiveLimits, + ), + parameter_name="sandbox.archive_limits", + ) + @dataclass class RunConfig: @@ -222,7 +312,7 @@ class RunConfig: model_settings: ModelSettings | None = None """Configure global model settings. Any non-null values will override the agent-specific model - settings. + settings. Accepts a ``ModelSettings`` instance or a dictionary containing its fields. """ handoff_input_filter: HandoffInputFilter | None = None @@ -339,6 +429,64 @@ class RunConfig: the run continue. """ + if TYPE_CHECKING: + + def __init__( + self, + model: str | Model | None = None, + model_provider: ModelProvider = ..., + model_settings: ModelSettings | dict[str, Any] | None = None, + handoff_input_filter: HandoffInputFilter | None = None, + nest_handoff_history: bool = False, + handoff_history_mapper: HandoffHistoryMapper | None = None, + input_guardrails: list[InputGuardrail[Any]] | None = None, + output_guardrails: list[OutputGuardrail[Any]] | None = None, + tracing_disabled: bool = False, + tracing: TracingConfig | None = None, + trace_include_sensitive_data: bool = ..., + workflow_name: str = "Agent workflow", + trace_id: str | None = None, + group_id: str | None = None, + trace_metadata: dict[str, Any] | None = None, + session_input_callback: SessionInputCallback | None = None, + call_model_input_filter: CallModelInputFilter | None = None, + tool_error_formatter: ToolErrorFormatter | None = None, + session_settings: SessionSettings | dict[str, Any] | None = None, + reasoning_item_id_policy: ReasoningItemIdPolicy | None = None, + sandbox: SandboxRunConfig | dict[str, Any] | None = None, + tool_execution: ToolExecutionConfig | dict[str, Any] | None = None, + tool_not_found_behavior: ToolNotFoundBehavior = "raise_error", + ) -> None: ... + + def __post_init__(self) -> None: + if self.model_settings is not None: + self.model_settings = _coerce_model_settings( + self.model_settings, + parameter_name="RunConfig model_settings", + model_settings_type=_declared_model_settings_type(type(self), "model_settings"), + ) + if self.session_settings is not None: + self.session_settings = _coerce_session_settings( + self.session_settings, + settings_type=_declared_session_settings_type(type(self), "session_settings"), + ) + if self.sandbox is not None: + self.sandbox = coerce_dataclass_config( + self.sandbox, + _declared_dataclass_type(type(self), "sandbox", SandboxRunConfig), + parameter_name="run_config.sandbox", + ) + if self.tool_execution is not None: + self.tool_execution = coerce_dataclass_config( + self.tool_execution, + _declared_dataclass_type( + type(self), + "tool_execution", + ToolExecutionConfig, + ), + parameter_name="run_config.tool_execution", + ) + class RunOptions(TypedDict, Generic[TContext]): """Arguments for ``AgentRunner`` methods.""" @@ -352,7 +500,7 @@ class RunOptions(TypedDict, Generic[TContext]): hooks: NotRequired[RunHooks[TContext] | None] """Lifecycle hooks for the run.""" - run_config: NotRequired[RunConfig | None] + run_config: NotRequired[RunConfig | dict[str, Any] | None] """Run configuration.""" previous_response_id: NotRequired[str | None] @@ -371,6 +519,11 @@ class RunOptions(TypedDict, Generic[TContext]): """Error handlers keyed by error kind.""" +def _coerce_run_config(value: RunConfig | dict[str, Any]) -> RunConfig: + """Normalize run configuration dictionaries at public runner boundaries.""" + return coerce_dataclass_config(value, RunConfig, parameter_name="run_config") + + __all__ = [ "DEFAULT_MAX_TURNS", "CallModelData", diff --git a/src/agents/sandbox/config.py b/src/agents/sandbox/config.py index 206ed459f1..1e9dc4acd2 100644 --- a/src/agents/sandbox/config.py +++ b/src/agents/sandbox/config.py @@ -1,11 +1,15 @@ from __future__ import annotations from dataclasses import dataclass, field -from typing import Final +from typing import TYPE_CHECKING, Any, Final from openai.types.shared import Reasoning -from ..model_settings import ModelSettings +from ..model_settings import ( + ModelSettings, + _coerce_model_settings, + _declared_model_settings_type, +) from ..models.interface import Model DEFAULT_PYTHON_SANDBOX_IMAGE: Final = "python:3.14-slim" @@ -47,7 +51,10 @@ class MemoryGenerateConfig: phase_one_model_settings: ModelSettings | None = field( default_factory=_default_memory_phase_one_model_settings ) - """Model settings used for phase-1 single-rollout extraction.""" + """Model settings used for phase-1 single-rollout extraction. + + Accepts a ``ModelSettings`` instance or a dictionary containing its fields. + """ phase_two_model: str | Model = "gpt-5.5" """Model used for phase-2 memory consolidation.""" @@ -55,7 +62,10 @@ class MemoryGenerateConfig: phase_two_model_settings: ModelSettings | None = field( default_factory=_default_memory_phase_two_model_settings ) - """Model settings used for phase-2 memory consolidation.""" + """Model settings used for phase-2 memory consolidation. + + Accepts a ``ModelSettings`` instance or a dictionary containing its fields. + """ extra_prompt: str | None = None """Optional developer-specific guidance appended to memory extraction and consolidation @@ -70,7 +80,36 @@ class MemoryGenerateConfig: evidence you actually want it to summarize. """ + if TYPE_CHECKING: + + def __init__( + self, + max_raw_memories_for_consolidation: int = 256, + phase_one_model: str | Model = "gpt-5.4-mini", + phase_one_model_settings: ModelSettings | dict[str, Any] | None = ..., + phase_two_model: str | Model = "gpt-5.5", + phase_two_model_settings: ModelSettings | dict[str, Any] | None = ..., + extra_prompt: str | None = None, + ) -> None: ... + def __post_init__(self) -> None: + if self.phase_one_model_settings is not None: + self.phase_one_model_settings = _coerce_model_settings( + self.phase_one_model_settings, + parameter_name="MemoryGenerateConfig.phase_one_model_settings", + model_settings_type=_declared_model_settings_type( + type(self), "phase_one_model_settings" + ), + ) + if self.phase_two_model_settings is not None: + self.phase_two_model_settings = _coerce_model_settings( + self.phase_two_model_settings, + parameter_name="MemoryGenerateConfig.phase_two_model_settings", + model_settings_type=_declared_model_settings_type( + type(self), "phase_two_model_settings" + ), + ) + if self.max_raw_memories_for_consolidation <= 0: raise ValueError( "MemoryGenerateConfig.max_raw_memories_for_consolidation must be greater than 0." diff --git a/src/agents/sandbox/manifest.py b/src/agents/sandbox/manifest.py index d4cc014870..9421694ecb 100644 --- a/src/agents/sandbox/manifest.py +++ b/src/agents/sandbox/manifest.py @@ -2,11 +2,12 @@ import asyncio from collections.abc import Iterator, Mapping from pathlib import Path, PurePath, PurePosixPath -from typing import Literal +from typing import Any, Literal from pydantic import BaseModel, Field, field_serializer, field_validator from typing_extensions import assert_never +from .._config_coercion import coerce_pydantic_config from .entries import BaseEntry, Dir, Mount, resolve_workspace_path from .errors import InvalidManifestPathError from .manifest_render import render_manifest_description @@ -256,3 +257,15 @@ def describe(self, depth: int | None = 1) -> str: coerce_rel_path=self._coerce_rel_path, depth=depth, ) + + +def _coerce_manifest(value: Manifest | dict[str, Any], *, parameter_name: str) -> Manifest: + """Normalize manifest dictionaries without granting untrusted host filesystem access.""" + if isinstance(value, dict) and "extra_path_grants" in value: + extra_path_grants = value["extra_path_grants"] + if not isinstance(extra_path_grants, list | tuple) or extra_path_grants: + raise TypeError( + f"{parameter_name}.extra_path_grants must be configured on a trusted " + "Manifest instance, not in a dictionary" + ) + return coerce_pydantic_config(value, Manifest, parameter_name=parameter_name) diff --git a/src/agents/sandbox/sandbox_agent.py b/src/agents/sandbox/sandbox_agent.py index 6021415428..82ccbba1f9 100644 --- a/src/agents/sandbox/sandbox_agent.py +++ b/src/agents/sandbox/sandbox_agent.py @@ -2,14 +2,29 @@ from collections.abc import Awaitable, Callable, Sequence from dataclasses import dataclass, field +from typing import TYPE_CHECKING, Any, Literal +from .._config_coercion import coerce_pydantic_config from ..agent import Agent from ..run_context import RunContextWrapper, TContext from .capabilities import Capability from .capabilities.capabilities import Capabilities -from .manifest import Manifest +from .manifest import Manifest, _coerce_manifest from .types import User +if TYPE_CHECKING: + from ..agent import MCPConfig, StopAtTools, ToolsToFinalOutputFunction + from ..agent_output import AgentOutputSchemaBase + from ..guardrail import InputGuardrail, OutputGuardrail + from ..handoffs import Handoff + from ..lifecycle import AgentHooks + from ..mcp import MCPServer + from ..model_settings import ModelSettings + from ..models.interface import Model + from ..prompts import DynamicPromptFunction, Prompt + from ..tool import Tool + from ..util._types import MaybeAwaitable + @dataclass class SandboxAgent(Agent[TContext]): @@ -39,8 +54,58 @@ class SandboxAgent(Agent[TContext]): _sandbox_concurrency_guard: object | None = field(default=None, init=False, repr=False) + if TYPE_CHECKING: + + def __init__( + self, + name: str, + handoff_description: str | None = None, + tools: list[Tool] = ..., + mcp_servers: list[MCPServer] = ..., + mcp_config: MCPConfig = ..., + instructions: ( + str + | Callable[ + [RunContextWrapper[TContext], Agent[TContext]], + MaybeAwaitable[str], + ] + | None + ) = None, + prompt: Prompt | DynamicPromptFunction | None = None, + handoffs: list[Agent[Any] | Handoff[TContext, Any]] = ..., + model: str | Model | None = None, + model_settings: ModelSettings | dict[str, Any] = ..., + input_guardrails: list[InputGuardrail[TContext]] = ..., + output_guardrails: list[OutputGuardrail[TContext]] = ..., + output_type: type[Any] | AgentOutputSchemaBase | None = None, + hooks: AgentHooks[TContext] | None = None, + tool_use_behavior: ( + Literal["run_llm_again", "stop_on_first_tool"] + | StopAtTools + | ToolsToFinalOutputFunction + ) = "run_llm_again", + reset_tool_choice: bool = True, + default_manifest: Manifest | dict[str, Any] | None = None, + base_instructions: ( + str + | Callable[ + [RunContextWrapper[TContext], Agent[TContext]], + Awaitable[str | None] | str | None, + ] + | None + ) = None, + capabilities: Sequence[Capability] = ..., + run_as: User | dict[str, Any] | str | None = None, + ) -> None: ... + def __post_init__(self) -> None: super().__post_init__() + if isinstance(self.default_manifest, dict): + self.default_manifest = _coerce_manifest( + self.default_manifest, parameter_name="sandbox.default_manifest" + ) + if isinstance(self.run_as, dict): + self.run_as = coerce_pydantic_config(self.run_as, User, parameter_name="sandbox.run_as") if ( self.base_instructions is not None and not isinstance(self.base_instructions, str) diff --git a/src/agents/tool.py b/src/agents/tool.py index def9ea2e66..af9a6c3c5c 100644 --- a/src/agents/tool.py +++ b/src/agents/tool.py @@ -43,6 +43,7 @@ from typing_extensions import NotRequired, ParamSpec, TypedDict from . import _debug +from ._config_coercion import coerce_pydantic_config from ._tool_identity import ( get_explicit_function_tool_namespace, tool_qualified_name, @@ -735,6 +736,22 @@ class WebSearchTool: indexed-only behavior where supported. """ + if TYPE_CHECKING: + + def __init__( + self, + user_location: UserLocation | None = None, + filters: WebSearchToolFilters | dict[str, Any] | None = None, + search_context_size: Literal["low", "medium", "high"] = "medium", + external_web_access: bool | None = None, + ) -> None: ... + + def __post_init__(self) -> None: + if isinstance(self.filters, dict): + self.filters = coerce_pydantic_config( + self.filters, WebSearchToolFilters, parameter_name="web search filters" + ) + @property def name(self): return "web_search" diff --git a/src/agents/tool_context.py b/src/agents/tool_context.py index 75947630cf..b9c753c79a 100644 --- a/src/agents/tool_context.py +++ b/src/agents/tool_context.py @@ -68,7 +68,7 @@ def __init__( *, tool_namespace: str | None = None, agent: AgentBase[Any] | None = None, - run_config: RunConfig | None = None, + run_config: RunConfig | dict[str, Any] | None = None, turn_input: list[TResponseInputItem] | None = None, _approvals: dict[str, _ApprovalRecord] | None = None, tool_input: Any | None = None, @@ -102,7 +102,12 @@ def __init__( else get_tool_call_namespace(tool_call) ) self.agent = agent - self.run_config = run_config + if run_config is not None: + from .run_config import _coerce_run_config + + self.run_config = _coerce_run_config(run_config) + else: + self.run_config = None # Internal adapter hook used to attach SDK-only custom data to the emitted output item. self._custom_data: dict[str, Any] | None = None @@ -122,7 +127,7 @@ def from_agent_context( tool_name: str | None = None, tool_arguments: str | None = None, tool_namespace: str | None = None, - run_config: RunConfig | None = None, + run_config: RunConfig | dict[str, Any] | None = None, ) -> ToolContext: """ Create a ToolContext from a RunContextWrapper. diff --git a/src/agents/voice/models/openai_model_provider.py b/src/agents/voice/models/openai_model_provider.py index b992f9b4ad..2736afbf57 100644 --- a/src/agents/voice/models/openai_model_provider.py +++ b/src/agents/voice/models/openai_model_provider.py @@ -1,5 +1,7 @@ from __future__ import annotations +from typing import Any + import httpx from openai import AsyncOpenAI, DefaultAsyncHttpxClient @@ -41,7 +43,7 @@ def __init__( openai_client: AsyncOpenAI | None = None, organization: str | None = None, project: str | None = None, - agent_registration: OpenAIAgentRegistrationConfig | None = None, + agent_registration: OpenAIAgentRegistrationConfig | dict[str, Any] | None = None, ) -> None: """Create a new OpenAI voice model provider. diff --git a/src/agents/voice/pipeline.py b/src/agents/voice/pipeline.py index 745f0faafb..21220c9921 100644 --- a/src/agents/voice/pipeline.py +++ b/src/agents/voice/pipeline.py @@ -1,7 +1,9 @@ from __future__ import annotations import asyncio +from typing import Any +from .._config_coercion import coerce_dataclass_config from ..exceptions import UserError from ..logger import logger from ..tracing import TraceCtxManager @@ -25,7 +27,7 @@ def __init__( workflow: VoiceWorkflowBase, stt_model: STTModel | str | None = None, tts_model: TTSModel | str | None = None, - config: VoicePipelineConfig | None = None, + config: VoicePipelineConfig | dict[str, Any] | None = None, ): """Create a new voice pipeline. @@ -43,7 +45,11 @@ def __init__( self.tts_model = tts_model if isinstance(tts_model, TTSModel) else None self._stt_model_name = stt_model if isinstance(stt_model, str) else None self._tts_model_name = tts_model if isinstance(tts_model, str) else None - self.config = config or VoicePipelineConfig() + self.config = ( + coerce_dataclass_config(config, VoicePipelineConfig, parameter_name="voice.pipeline") + if config is not None + else VoicePipelineConfig() + ) async def run(self, audio_input: AudioInput | StreamedAudioInput) -> StreamedAudioResult: """Run the voice pipeline. diff --git a/src/agents/voice/pipeline_config.py b/src/agents/voice/pipeline_config.py index eed2ab6940..35c55d093a 100644 --- a/src/agents/voice/pipeline_config.py +++ b/src/agents/voice/pipeline_config.py @@ -1,8 +1,9 @@ from __future__ import annotations from dataclasses import dataclass, field -from typing import Any +from typing import TYPE_CHECKING, Any +from .._config_coercion import _declared_dataclass_type, coerce_dataclass_config from ..tracing import TracingConfig from ..tracing.util import gen_group_id from .model import STTModelSettings, TTSModelSettings, VoiceModelProvider @@ -48,3 +49,31 @@ class VoicePipelineConfig: tts_settings: TTSModelSettings = field(default_factory=TTSModelSettings) """The settings to use for the TTS model.""" + + if TYPE_CHECKING: + + def __init__( + self, + model_provider: VoiceModelProvider = ..., + tracing_disabled: bool = False, + tracing: TracingConfig | None = None, + trace_include_sensitive_data: bool = True, + trace_include_sensitive_audio_data: bool = True, + workflow_name: str = "Voice Agent", + group_id: str = ..., + trace_metadata: dict[str, Any] | None = None, + stt_settings: STTModelSettings | dict[str, Any] = ..., + tts_settings: TTSModelSettings | dict[str, Any] = ..., + ) -> None: ... + + def __post_init__(self) -> None: + self.stt_settings = coerce_dataclass_config( + self.stt_settings, + _declared_dataclass_type(type(self), "stt_settings", STTModelSettings), + parameter_name="voice.stt", + ) + self.tts_settings = coerce_dataclass_config( + self.tts_settings, + _declared_dataclass_type(type(self), "tts_settings", TTSModelSettings), + parameter_name="voice.tts", + ) diff --git a/tests/extensions/memory/test_advanced_sqlite_session.py b/tests/extensions/memory/test_advanced_sqlite_session.py index 28d8f3f6a9..2472a8bf12 100644 --- a/tests/extensions/memory/test_advanced_sqlite_session.py +++ b/tests/extensions/memory/test_advanced_sqlite_session.py @@ -1619,17 +1619,18 @@ async def test_session_settings_default(): session.close() -async def test_session_settings_constructor(): +@pytest.mark.parametrize("use_dictionary", [False, True], ids=["class", "dictionary"]) +async def test_session_settings_constructor(use_dictionary: bool): """Test passing session_settings via constructor.""" from agents.memory import SessionSettings session = AdvancedSQLiteSession( session_id="constructor_settings_test", create_tables=True, - session_settings=SessionSettings(limit=5), + session_settings={"limit": 5} if use_dictionary else SessionSettings(limit=5), ) - assert session.session_settings is not None + assert isinstance(session.session_settings, SessionSettings) assert session.session_settings.limit == 5 session.close() diff --git a/tests/extensions/memory/test_async_sqlite_session.py b/tests/extensions/memory/test_async_sqlite_session.py index 7269951829..6ab3d9feb4 100644 --- a/tests/extensions/memory/test_async_sqlite_session.py +++ b/tests/extensions/memory/test_async_sqlite_session.py @@ -151,14 +151,15 @@ async def test_async_sqlite_session_session_settings_default(): await session.close() -async def test_async_sqlite_session_session_settings_constructor(): +@pytest.mark.parametrize("use_dictionary", [False, True], ids=["class", "dictionary"]) +async def test_async_sqlite_session_session_settings_constructor(use_dictionary: bool): """Test passing session_settings via constructor.""" session = AsyncSQLiteSession( "async_constructor_settings", - session_settings=SessionSettings(limit=5), + session_settings={"limit": 5} if use_dictionary else SessionSettings(limit=5), ) - assert session.session_settings is not None + assert isinstance(session.session_settings, SessionSettings) assert session.session_settings.limit == 5 await session.close() diff --git a/tests/extensions/memory/test_dapr_session.py b/tests/extensions/memory/test_dapr_session.py index 9766f35d40..dd49173a19 100644 --- a/tests/extensions/memory/test_dapr_session.py +++ b/tests/extensions/memory/test_dapr_session.py @@ -894,7 +894,8 @@ async def test_session_settings_default(fake_dapr_client: FakeDaprClient): await session.close() -async def test_session_settings_constructor(fake_dapr_client: FakeDaprClient): +@pytest.mark.parametrize("use_dictionary", [False, True], ids=["class", "dictionary"]) +async def test_session_settings_constructor(fake_dapr_client: FakeDaprClient, use_dictionary: bool): """Test passing session_settings via constructor.""" from agents.memory import SessionSettings @@ -902,11 +903,11 @@ async def test_session_settings_constructor(fake_dapr_client: FakeDaprClient): session_id="settings_test", state_store_name="statestore", dapr_client=fake_dapr_client, # type: ignore[arg-type] - session_settings=SessionSettings(limit=5), + session_settings={"limit": 5} if use_dictionary else SessionSettings(limit=5), ) try: - assert session.session_settings is not None + assert isinstance(session.session_settings, SessionSettings) assert session.session_settings.limit == 5 finally: await session.close() diff --git a/tests/extensions/memory/test_mongodb_session.py b/tests/extensions/memory/test_mongodb_session.py index cd7954e3ae..98cfc2654d 100644 --- a/tests/extensions/memory/test_mongodb_session.py +++ b/tests/extensions/memory/test_mongodb_session.py @@ -396,15 +396,17 @@ async def test_get_items_limit_exceeds_count(session: MongoDBSession) -> None: assert len(result) == 1 -async def test_session_settings_limit_used_as_default() -> None: +@pytest.mark.parametrize("use_dictionary", [False, True], ids=["class", "dictionary"]) +async def test_session_settings_limit_used_as_default(use_dictionary: bool) -> None: """session_settings.limit is applied when no explicit limit is given.""" MongoDBSession._init_state.clear() s = MongoDBSession( "ls-test", client=FakeAsyncMongoClient(), # type: ignore[arg-type] database="agents_test", - session_settings=SessionSettings(limit=2), + session_settings={"limit": 2} if use_dictionary else SessionSettings(limit=2), ) + assert isinstance(s.session_settings, SessionSettings) await s.add_items([{"role": "user", "content": str(i)} for i in range(5)]) result = await s.get_items() diff --git a/tests/extensions/memory/test_redis_session.py b/tests/extensions/memory/test_redis_session.py index b5011cdd4d..0cc4c07d8b 100644 --- a/tests/extensions/memory/test_redis_session.py +++ b/tests/extensions/memory/test_redis_session.py @@ -840,7 +840,8 @@ async def test_session_settings_default(): await session.close() -async def test_session_settings_constructor(): +@pytest.mark.parametrize("use_dictionary", [False, True], ids=["class", "dictionary"]) +async def test_session_settings_constructor(use_dictionary: bool): """Test passing session_settings via constructor.""" from agents.memory import SessionSettings @@ -849,15 +850,17 @@ async def test_session_settings_constructor(): session_id="settings_test", redis_client=fake_redis, key_prefix="test:", - session_settings=SessionSettings(limit=5), + session_settings={"limit": 5} if use_dictionary else SessionSettings(limit=5), ) else: session = RedisSession.from_url( - "settings_test", url=REDIS_URL, session_settings=SessionSettings(limit=5) + "settings_test", + url=REDIS_URL, + session_settings={"limit": 5} if use_dictionary else SessionSettings(limit=5), ) try: - assert session.session_settings is not None + assert isinstance(session.session_settings, SessionSettings) assert session.session_settings.limit == 5 finally: await session.close() diff --git a/tests/extensions/memory/test_sqlalchemy_session.py b/tests/extensions/memory/test_sqlalchemy_session.py index 091f88a482..c75d7d6141 100644 --- a/tests/extensions/memory/test_sqlalchemy_session.py +++ b/tests/extensions/memory/test_sqlalchemy_session.py @@ -836,7 +836,8 @@ async def test_session_settings_default(): assert session.session_settings.limit is None -async def test_session_settings_from_url(): +@pytest.mark.parametrize("use_dictionary", [False, True], ids=["class", "dictionary"]) +async def test_session_settings_from_url(use_dictionary: bool): """Test passing session_settings via from_url.""" from agents.memory import SessionSettings @@ -844,10 +845,10 @@ async def test_session_settings_from_url(): "from_url_settings_test", url=DB_URL, create_tables=True, - session_settings=SessionSettings(limit=5), + session_settings={"limit": 5} if use_dictionary else SessionSettings(limit=5), ) - assert session.session_settings is not None + assert isinstance(session.session_settings, SessionSettings) assert session.session_settings.limit == 5 diff --git a/tests/extensions/sandbox/test_blaxel.py b/tests/extensions/sandbox/test_blaxel.py index 2e77fb8f80..8f189bcd4d 100644 --- a/tests/extensions/sandbox/test_blaxel.py +++ b/tests/extensions/sandbox/test_blaxel.py @@ -14,6 +14,7 @@ import pytest from pydantic import ValidationError +from agents.run_config import SandboxRunConfig from agents.sandbox import Manifest, SandboxPathGrant from agents.sandbox.config import DEFAULT_PYTHON_SANDBOX_IMAGE from agents.sandbox.errors import ( @@ -766,6 +767,26 @@ async def test_create(self, monkeypatch: pytest.MonkeyPatch) -> None: session = await client.create(options=options) assert session is not None + @pytest.mark.asyncio + async def test_create_with_dictionary_run_config_options( + self, monkeypatch: pytest.MonkeyPatch + ) -> None: + from agents.extensions.sandbox.blaxel import sandbox as mod + + monkeypatch.setattr(mod, "_import_blaxel_sdk", lambda: _FakeSandboxInstance) + + client = mod.BlaxelSandboxClient(token="test-token") + config = SandboxRunConfig( + client=client, + options={"name": "dict-options", "timeouts": {"exec_timeout_s": 120}}, + ) + + assert isinstance(config.options, mod.BlaxelSandboxClientOptions) + session = await client.create(options=config.options) + + assert isinstance(session.state, mod.BlaxelSandboxSessionState) + assert session.state.timeouts.exec_timeout_s == 120 + @pytest.mark.asyncio async def test_create_with_image(self, monkeypatch: pytest.MonkeyPatch) -> None: from agents.extensions.sandbox.blaxel import sandbox as mod diff --git a/tests/memory/test_openai_conversations_session.py b/tests/memory/test_openai_conversations_session.py index b2b62950fe..2e241b88b8 100644 --- a/tests/memory/test_openai_conversations_session.py +++ b/tests/memory/test_openai_conversations_session.py @@ -553,3 +553,14 @@ def test_session_settings_constructor(self, mock_openai_client): assert session.session_settings is not None assert session.session_settings.limit == 5 + + def test_session_settings_constructor_normalizes_dictionary(self, mock_openai_client): + from agents.memory import SessionSettings + + session = OpenAIConversationsSession( + openai_client=mock_openai_client, + session_settings={"limit": 0}, + ) + + assert isinstance(session.session_settings, SessionSettings) + assert session.session_settings.limit == 0 diff --git a/tests/memory/test_session.py b/tests/memory/test_session.py index f9cc324d2e..d727991f7d 100644 --- a/tests/memory/test_session.py +++ b/tests/memory/test_session.py @@ -7,7 +7,7 @@ import pytest -from agents import Agent, RunConfig, Runner, SQLiteSession, TResponseInputItem +from agents import Agent, RunConfig, Runner, SessionSettings, SQLiteSession, TResponseInputItem from tests.fake_model import FakeModel from tests.test_responses import get_text_message @@ -694,6 +694,22 @@ async def test_session_settings_constructor(): session.close() +@pytest.mark.asyncio +async def test_session_settings_constructor_normalizes_dictionary() -> None: + session = SQLiteSession("dictionary_settings_test", session_settings={"limit": 0}) + + assert isinstance(session.session_settings, SessionSettings) + assert session.session_settings.limit == 0 + assert session.session_settings.resolve({"limit": 4}).limit == 4 + + session.close() + + +def test_session_settings_rejects_unknown_dictionary_fields() -> None: + with pytest.raises(TypeError, match="Unknown session settings: limitt"): + SQLiteSession("invalid_settings_test", session_settings={"limitt": 1}) + + @pytest.mark.asyncio async def test_get_items_uses_session_settings_limit(): """Test that get_items uses session_settings.limit as default.""" diff --git a/tests/model_settings/test_serialization.py b/tests/model_settings/test_serialization.py index ea59dc55f2..073801bd11 100644 --- a/tests/model_settings/test_serialization.py +++ b/tests/model_settings/test_serialization.py @@ -1,6 +1,7 @@ import json from dataclasses import fields +import pytest from openai.types.shared import Reasoning from pydantic import TypeAdapter from pydantic_core import to_json @@ -30,6 +31,54 @@ def test_basic_serialization() -> None: verify_serialization(model_settings) +def test_model_settings_direct_constructor_preserves_openai_reasoning_extensions() -> None: + settings = ModelSettings( + reasoning={"context": "all_turns", "future_reasoning_option": "enabled"} + ) + + assert isinstance(settings.reasoning, Reasoning) + assert settings.reasoning.context == "all_turns" + assert settings.reasoning.model_extra == {"future_reasoning_option": "enabled"} + + +def test_model_settings_dictionary_override_preserves_omitted_values() -> None: + settings = ModelSettings( + temperature=0.5, + reasoning=Reasoning.model_validate( + {"context": "all_turns", "future_reasoning_option": "enabled"} + ), + retry=ModelRetrySettings(max_retries=2), + ) + + resolved = settings.resolve({"temperature": 0.0}) + + assert resolved.temperature == 0.0 + assert resolved.reasoning is settings.reasoning + assert resolved.retry is settings.retry + + +def test_model_settings_dictionary_override_merges_retry_settings() -> None: + settings = ModelSettings( + retry=ModelRetrySettings( + max_retries=2, + backoff=ModelRetryBackoffSettings(initial_delay=0.1, jitter=True), + ) + ) + + resolved = settings.resolve({"retry": {"max_retries": 0, "backoff": {"jitter": False}}}) + + assert resolved.retry is not None + assert resolved.retry.max_retries == 0 + assert isinstance(resolved.retry.backoff, ModelRetryBackoffSettings) + assert resolved.retry.backoff.initial_delay == 0.1 + assert resolved.retry.backoff.jitter is False + + +def test_model_settings_dictionary_override_rejects_unknown_fields() -> None: + with pytest.raises(TypeError, match="Unknown model settings: temperatur"): + ModelSettings().resolve({"temperatur": 0.5}) + + def test_mcp_tool_choice_serialization() -> None: """Tests whether ModelSettings with MCPToolChoice can be serialized to a JSON string.""" # First, lets create a ModelSettings instance diff --git a/tests/models/test_agent_registration.py b/tests/models/test_agent_registration.py index 4741db8b64..c22f69319a 100644 --- a/tests/models/test_agent_registration.py +++ b/tests/models/test_agent_registration.py @@ -17,6 +17,7 @@ from agents.models.openai_provider import OpenAIProvider from agents.run_internal.agent_runner_helpers import resolve_trace_settings from agents.tracing import agent_span, trace +from agents.voice.models.openai_model_provider import OpenAIVoiceModelProvider def test_agent_registration_config_precedence(monkeypatch: pytest.MonkeyPatch) -> None: @@ -91,6 +92,35 @@ def test_agent_registration_provider_constructor_config() -> None: assert multi_provider.openai_provider.agent_registration.harness_id == "provider-harness" +def test_agent_registration_provider_constructors_normalize_dictionaries() -> None: + config = {"harness_id": "dictionary-harness"} + openai_provider = OpenAIProvider(agent_registration=config) + multi_provider = MultiProvider(openai_agent_registration=config) + voice_provider = OpenAIVoiceModelProvider(agent_registration=config) + + assert openai_provider.agent_registration is not None + assert openai_provider.agent_registration.harness_id == "dictionary-harness" + assert multi_provider.openai_provider.agent_registration is not None + assert multi_provider.openai_provider.agent_registration.harness_id == "dictionary-harness" + assert voice_provider.agent_registration is not None + assert voice_provider.agent_registration.harness_id == "dictionary-harness" + + +def test_default_agent_registration_normalizes_dictionary() -> None: + set_default_openai_agent_registration({"harness_id": "dictionary-default"}) + try: + resolved = resolve_openai_agent_registration_config(None) + assert resolved is not None + assert resolved.harness_id == "dictionary-default" + finally: + set_default_openai_agent_registration(None) + + +def test_agent_registration_rejects_unknown_dictionary_fields() -> None: + with pytest.raises(TypeError, match="Unknown OpenAI agent registration settings: harness_idd"): + OpenAIProvider(agent_registration={"harness_idd": "invalid"}) + + def test_harness_id_resolves_private_agent_registration() -> None: class Provider: _agent_registration = OpenAIAgentRegistrationConfig(harness_id="private-harness") diff --git a/tests/models/test_kwargs_functionality.py b/tests/models/test_kwargs_functionality.py index dc641a75d2..3b8a7cc65d 100644 --- a/tests/models/test_kwargs_functionality.py +++ b/tests/models/test_kwargs_functionality.py @@ -1,3 +1,5 @@ +from typing import Any + import httpx import litellm import pytest @@ -9,6 +11,7 @@ from openai.types.chat.chat_completion_message import ChatCompletionMessage from openai.types.completion_usage import CompletionUsage +from agents import Agent from agents.extensions.models.litellm_model import LitellmModel from agents.model_settings import ModelSettings from agents.models._retry_runtime import provider_managed_retries_disabled @@ -66,6 +69,44 @@ async def fake_acompletion(model, messages=None, **kwargs): assert captured["temperature"] == 0.5 +@pytest.mark.allow_call_model_methods +@pytest.mark.asyncio +@pytest.mark.parametrize("use_dictionary", [False, True], ids=["model-settings", "dictionary"]) +async def test_litellm_normalizes_dictionary_agent_model_settings( + monkeypatch, use_dictionary: bool +): + captured: dict[str, object] = {} + + async def fake_acompletion(model, messages=None, **kwargs): + captured.update(kwargs) + message = Message(role="assistant", content="test response") + return ModelResponse(choices=[Choices(index=0, message=message)], usage=Usage(0, 0, 0)) + + monkeypatch.setattr(litellm, "acompletion", fake_acompletion) + settings: dict[str, Any] = {"temperature": 0.0, "reasoning": {"effort": "low"}} + model = LitellmModel(model="test-model") + agent = Agent( + name="test", + model=model, + model_settings=settings if use_dictionary else ModelSettings(**settings), + ) + + await model.get_response( + system_instructions=None, + input="test input", + model_settings=agent.model_settings, + tools=[], + output_schema=None, + handoffs=[], + tracing=ModelTracing.DISABLED, + previous_response_id=None, + conversation_id=None, + ) + + assert captured["temperature"] == 0.0 + assert captured["reasoning_effort"] == "low" + + @pytest.mark.allow_call_model_methods @pytest.mark.asyncio async def test_openai_chatcompletions_kwargs_forwarded(monkeypatch): diff --git a/tests/models/test_openai_chatcompletions.py b/tests/models/test_openai_chatcompletions.py index 0d2d4e8ed4..7bebfbc6e4 100644 --- a/tests/models/test_openai_chatcompletions.py +++ b/tests/models/test_openai_chatcompletions.py @@ -70,7 +70,7 @@ def _minimal_chat_completion(content: str = "ok") -> ChatCompletion: async def _run_chat_completions_model_with_custom_base_url( - model_settings: ModelSettings | None = None, + model_settings: ModelSettings | dict[str, Any] | None = None, ) -> dict[str, Any]: class DummyCompletions: def __init__(self) -> None: @@ -787,6 +787,54 @@ def test_chat_completions_rejects_responses_only_reasoning_settings_in_strict_mo ) +@pytest.mark.allow_call_model_methods +@pytest.mark.asyncio +@pytest.mark.parametrize("use_dictionary", [False, True], ids=["model-settings", "dictionary"]) +async def test_chat_completions_requests_normalize_dictionary_agent_settings( + use_dictionary: bool, +) -> None: + settings: dict[str, Any] = { + "reasoning": {"effort": "high"}, + "prompt_cache_options": {"mode": "explicit", "ttl": "30m"}, + "prompt_cache_retention": "24h", + "verbosity": "low", + "store": False, + "temperature": 0.3, + "top_p": 1.0, + "frequency_penalty": 0.0, + "presence_penalty": 0.0, + "max_tokens": 64, + "parallel_tool_calls": False, + "extra_headers": {"x-model-settings-parity": "preserved"}, + "extra_query": {"model_settings_parity": "verified"}, + "extra_body": {"prompt_cache_key": "extra-body-cache-key"}, + "retry": { + "max_retries": 0, + "backoff": {"initial_delay": 0.0, "jitter": False}, + }, + } + kwargs = await _run_chat_completions_model_with_custom_base_url( + model_settings=settings if use_dictionary else ModelSettings(**settings) + ) + + assert kwargs["reasoning_effort"] == "high" + assert kwargs["prompt_cache_options"] == settings["prompt_cache_options"] + assert kwargs["prompt_cache_retention"] == "24h" + assert kwargs["verbosity"] == "low" + assert kwargs["store"] is False + assert kwargs["temperature"] == 0.3 + assert kwargs["top_p"] == 1.0 + assert kwargs["frequency_penalty"] == 0.0 + assert kwargs["presence_penalty"] == 0.0 + assert kwargs["max_tokens"] == 64 + assert "max_output_tokens" not in kwargs + assert kwargs["parallel_tool_calls"] is False + assert kwargs["extra_headers"]["x-model-settings-parity"] == "preserved" + assert kwargs["extra_query"] == {"model_settings_parity": "verified"} + assert kwargs["extra_body"] == {"prompt_cache_key": "extra-body-cache-key"} + assert "retry" not in kwargs + + @pytest.mark.allow_call_model_methods @pytest.mark.asyncio async def test_custom_base_url_prompt_cache_key_uses_model_settings_only() -> None: diff --git a/tests/models/test_openai_chatcompletions_stream.py b/tests/models/test_openai_chatcompletions_stream.py index 75919a6a11..90b01572ac 100644 --- a/tests/models/test_openai_chatcompletions_stream.py +++ b/tests/models/test_openai_chatcompletions_stream.py @@ -2,6 +2,7 @@ from collections.abc import AsyncIterator from typing import Any, cast +import httpx import pytest from openai.types.chat.chat_completion import ChatCompletion, Choice as ChatCompletionChoice from openai.types.chat.chat_completion_chunk import ( @@ -99,6 +100,95 @@ async def _collect_buffered_tool_call_chunks( ] +@pytest.mark.allow_call_model_methods +@pytest.mark.asyncio +@pytest.mark.parametrize("use_dictionary", [False, True], ids=["model-settings", "dictionary"]) +async def test_stream_response_forwards_dictionary_agent_model_settings( + use_dictionary: bool, +) -> None: + chunk = ChatCompletionChunk( + id="chunk-id", + created=1, + model="gpt-5.4-mini", + object="chat.completion.chunk", + choices=[ + Choice( + index=0, + delta=ChoiceDelta(role="assistant", content="ok"), + finish_reason="stop", + ) + ], + ) + + class DummyCompletions: + def __init__(self) -> None: + self.kwargs: dict[str, Any] = {} + + async def create(self, **kwargs: Any) -> AsyncIterator[ChatCompletionChunk]: + self.kwargs = kwargs + return _completion_stream(chunk) + + class DummyClient: + def __init__(self, completions: DummyCompletions) -> None: + self.chat = type("_Chat", (), {"completions": completions})() + self.base_url = httpx.URL("https://api.openai.com/v1/") + + completions = DummyCompletions() + model = OpenAIChatCompletionsModel( + model="gpt-5.4-mini", openai_client=cast(Any, DummyClient(completions)) + ) + settings: dict[str, Any] = { + "reasoning": {"effort": "low"}, + "prompt_cache_options": {"mode": "explicit", "ttl": "30m"}, + "prompt_cache_retention": "24h", + "verbosity": "low", + "store": False, + "temperature": 0.0, + "top_p": 1.0, + "frequency_penalty": 0.0, + "presence_penalty": 0.0, + "max_tokens": 64, + "parallel_tool_calls": False, + "include_usage": False, + } + agent = Agent( + name="test", + model=model, + model_settings=settings if use_dictionary else ModelSettings(**settings), + ) + + events = [ + event + async for event in model.stream_response( + system_instructions=None, + input="hi", + model_settings=agent.model_settings, + tools=[], + output_schema=None, + handoffs=[], + tracing=ModelTracing.DISABLED, + previous_response_id=None, + conversation_id=None, + prompt=None, + ) + ] + + assert any(event.type == "response.completed" for event in events) + assert completions.kwargs["reasoning_effort"] == "low" + assert completions.kwargs["prompt_cache_options"] == settings["prompt_cache_options"] + assert completions.kwargs["prompt_cache_retention"] == "24h" + assert completions.kwargs["verbosity"] == "low" + assert completions.kwargs["store"] is False + assert completions.kwargs["temperature"] == 0.0 + assert completions.kwargs["top_p"] == 1.0 + assert completions.kwargs["frequency_penalty"] == 0.0 + assert completions.kwargs["presence_penalty"] == 0.0 + assert completions.kwargs["max_tokens"] == 64 + assert completions.kwargs["parallel_tool_calls"] is False + assert completions.kwargs["stream"] is True + assert completions.kwargs["stream_options"] == {"include_usage": False} + + @pytest.mark.allow_call_model_methods @pytest.mark.asyncio async def test_stream_response_yields_events_for_text_content(monkeypatch) -> None: diff --git a/tests/models/test_openai_responses.py b/tests/models/test_openai_responses.py index 2b4ca111be..00c6a4cf50 100644 --- a/tests/models/test_openai_responses.py +++ b/tests/models/test_openai_responses.py @@ -45,7 +45,7 @@ async def _run_responses_model_with_custom_base_url( - model_settings: ModelSettings | None = None, + model_settings: ModelSettings | dict[str, Any] | None = None, ) -> dict[str, Any]: class DummyResponses: def __init__(self) -> None: @@ -934,6 +934,56 @@ def test_build_response_create_kwargs_includes_gpt_5_6_request_controls(): assert kwargs["previous_response_id"] == "resp-previous" +@pytest.mark.allow_call_model_methods +@pytest.mark.asyncio +@pytest.mark.parametrize("use_dictionary", [False, True], ids=["model-settings", "dictionary"]) +async def test_responses_requests_normalize_dictionary_agent_settings(use_dictionary: bool) -> None: + settings: dict[str, Any] = { + "reasoning": {"effort": "low", "context": "all_turns"}, + "context_management": [{"type": "compaction", "compact_threshold": 200000}], + "prompt_cache_options": {"mode": "explicit", "ttl": "30m"}, + "prompt_cache_retention": "24h", + "store": False, + "metadata": {"request": "example"}, + "temperature": 0.0, + "top_p": 1.0, + "frequency_penalty": 0.0, + "presence_penalty": 0.0, + "max_tokens": 64, + "parallel_tool_calls": False, + "extra_headers": {"x-model-settings-parity": "preserved"}, + "extra_query": {"model_settings_parity": "verified"}, + "extra_body": {"prompt_cache_key": "extra-body-cache-key"}, + "retry": { + "max_retries": 0, + "backoff": {"initial_delay": 0.0, "jitter": False}, + }, + } + kwargs = await _run_responses_model_with_custom_base_url( + model_settings=settings if use_dictionary else ModelSettings(**settings) + ) + + assert isinstance(kwargs["reasoning"], Reasoning) + assert kwargs["reasoning"].effort == "low" + assert kwargs["reasoning"].context == "all_turns" + assert kwargs["context_management"] == settings["context_management"] + assert kwargs["prompt_cache_options"] == settings["prompt_cache_options"] + assert kwargs["prompt_cache_retention"] == "24h" + assert kwargs["store"] is False + assert kwargs["metadata"] == {"request": "example"} + assert kwargs["temperature"] == 0.0 + assert kwargs["top_p"] == 1.0 + assert kwargs["max_output_tokens"] == 64 + assert "max_tokens" not in kwargs + assert kwargs["parallel_tool_calls"] is False + assert kwargs["extra_headers"]["x-model-settings-parity"] == "preserved" + assert kwargs["extra_query"] == {"model_settings_parity": "verified"} + assert kwargs["extra_body"] == {"prompt_cache_key": "extra-body-cache-key"} + assert "retry" not in kwargs + assert "frequency_penalty" not in kwargs + assert "presence_penalty" not in kwargs + + @pytest.mark.allow_call_model_methods def test_build_response_create_kwargs_rejects_duplicate_prompt_cache_options_extra_args(): client = DummyWSClient() @@ -1777,13 +1827,22 @@ async def fake_open( monkeypatch.setattr(model, "_open_websocket_connection", fake_open) + configured_agent = Agent( + name="configured", + model=model, + model_settings={ + "reasoning": {"mode": "pro", "effort": "max", "context": "all_turns"}, + "context_management": [{"type": "compaction", "compact_threshold": 200000}], + "prompt_cache_options": {"mode": "explicit", "ttl": "30m"}, + "prompt_cache_retention": "24h", + "store": False, + "metadata": {"request": "example"}, + }, + ) first = await model.get_response( system_instructions=None, input="hi", - model_settings=ModelSettings( - reasoning=Reasoning(mode="pro", effort="max", context="all_turns"), - prompt_cache_options={"mode": "explicit", "ttl": "30m"}, - ), + model_settings=configured_agent.model_settings, tools=[], output_schema=None, handoffs=[], @@ -1815,6 +1874,12 @@ async def fake_open( "mode": "explicit", "ttl": "30m", } + assert ws.sent_messages[0]["context_management"] == [ + {"type": "compaction", "compact_threshold": 200000} + ] + assert ws.sent_messages[0]["prompt_cache_retention"] == "24h" + assert ws.sent_messages[0]["store"] is False + assert ws.sent_messages[0]["metadata"] == {"request": "example"} assert ws.sent_messages[1]["type"] == "response.create" assert ws.sent_messages[1]["stream"] is True assert ws.sent_messages[1]["previous_response_id"] == "resp-1" diff --git a/tests/models/test_openai_responses_converter.py b/tests/models/test_openai_responses_converter.py index e1c8069ec9..cef2c8b81b 100644 --- a/tests/models/test_openai_responses_converter.py +++ b/tests/models/test_openai_responses_converter.py @@ -27,6 +27,7 @@ import pytest from openai import omit +from openai.types.responses.web_search_tool import Filters as WebSearchToolFilters from pydantic import BaseModel from agents import ( @@ -468,6 +469,32 @@ def test_convert_tools_includes_explicit_false_external_web_access() -> None: ] +@pytest.mark.parametrize("use_dictionary", [False, True], ids=["class", "dictionary"]) +def test_web_search_filters_preserve_existing_provider_payload(use_dictionary: bool) -> None: + filters = {"allowed_domains": ["example.com"]} + tool = WebSearchTool( + filters=filters if use_dictionary else WebSearchToolFilters.model_validate(filters) + ) + + assert isinstance(tool.filters, WebSearchToolFilters) + converted = Converter.convert_tools([tool], handoffs=[], model="gpt-5.4") + assert converted.tools == [ + { + "type": "web_search", + "filters": filters, + "user_location": None, + "search_context_size": "medium", + } + ] + + +def test_web_search_filters_preserve_openai_forward_compatible_fields() -> None: + tool = WebSearchTool(filters={"future_filter": ["example.com"]}) + + assert tool.filters is not None + assert tool.filters.model_extra == {"future_filter": ["example.com"]} + + def test_convert_tools_uses_preview_computer_payload_for_preview_model() -> None: comp_tool = ComputerTool(computer=DummyComputer()) diff --git a/tests/models/test_responses_websocket_session.py b/tests/models/test_responses_websocket_session.py index c1272da156..fc2d339756 100644 --- a/tests/models/test_responses_websocket_session.py +++ b/tests/models/test_responses_websocket_session.py @@ -2,7 +2,7 @@ import pytest -from agents import Agent, responses_websocket_session +from agents import Agent, ResponsesWebSocketSession, RunConfig, responses_websocket_session from agents.models.multi_provider import MultiProvider from agents.models.openai_provider import OpenAIProvider @@ -17,6 +17,35 @@ async def test_responses_websocket_session_builds_shared_run_config(): assert ws.run_config.model_provider.openai_provider is ws.provider +def test_responses_websocket_session_normalizes_dictionary_run_config() -> None: + provider = MultiProvider(openai_api_key="test") + + session = ResponsesWebSocketSession( + provider=provider.openai_provider, + run_config={ + "model_provider": provider, + "model_settings": {"temperature": 0.0, "retry": {"max_retries": 0}}, + }, + ) + + assert isinstance(session.run_config, RunConfig) + assert session.run_config.model_provider is provider + assert session.run_config.model_settings is not None + assert session.run_config.model_settings.temperature == 0.0 + assert session.run_config.model_settings.retry is not None + assert session.run_config.model_settings.retry.max_retries == 0 + + +def test_responses_websocket_session_rejects_unknown_dictionary_run_config_fields() -> None: + provider = MultiProvider(openai_api_key="test") + + with pytest.raises(TypeError, match="Unknown run_config settings: tracin_disabled"): + ResponsesWebSocketSession( + provider=provider.openai_provider, + run_config={"model_provider": provider, "tracin_disabled": True}, + ) + + @pytest.mark.asyncio async def test_responses_websocket_session_preserves_openai_prefix_routing(monkeypatch): captured: dict[str, object] = {} diff --git a/tests/sandbox/test_memory.py b/tests/sandbox/test_memory.py index 5eb843de2d..c917dacd6d 100644 --- a/tests/sandbox/test_memory.py +++ b/tests/sandbox/test_memory.py @@ -2,9 +2,10 @@ import io import json +from dataclasses import dataclass from datetime import datetime from pathlib import Path -from typing import Any, cast +from typing import Any, cast, get_type_hints import pytest from openai.types.responses import ResponseCustomToolCall, ResponseFunctionToolCall @@ -16,6 +17,7 @@ import agents.sandbox.memory.phase_one as phase_one_module from agents import ( Agent, + ModelSettings, ReasoningItem, RunConfig, Runner, @@ -70,6 +72,17 @@ from tests.utils.hitl import make_shell_call +@dataclass +class _DeclaredProviderModelSettings(ModelSettings): + provider_field: str | None = None + + +@dataclass +class _DeclaredProviderMemoryGenerateConfig(MemoryGenerateConfig): + phase_one_model_settings: _DeclaredProviderModelSettings | None = None + phase_two_model_settings: _DeclaredProviderModelSettings | None = None + + class _DeleteTrackingUnixLocalSandboxClient(UnixLocalSandboxClient): def __init__(self) -> None: super().__init__() @@ -696,6 +709,93 @@ def test_memory_generate_config_accepts_renamed_limit_field() -> None: assert config.max_raw_memories_for_consolidation == 123 +def test_memory_generate_config_normalizes_dictionary_model_settings() -> None: + config = MemoryGenerateConfig( + phase_one_model_settings={ + "reasoning": {"effort": "low"}, + "retry": {"max_retries": 0}, + }, + phase_two_model_settings={"temperature": 0.0, "store": False}, + ) + + assert isinstance(config.phase_one_model_settings, ModelSettings) + assert config.phase_one_model_settings.reasoning is not None + assert config.phase_one_model_settings.reasoning.effort == "low" + assert config.phase_one_model_settings.retry is not None + assert config.phase_one_model_settings.retry.max_retries == 0 + assert isinstance(config.phase_two_model_settings, ModelSettings) + assert config.phase_two_model_settings.temperature == 0.0 + assert config.phase_two_model_settings.store is False + + +def test_memory_generate_config_subclass_uses_declared_model_settings_types() -> None: + config = cast(Any, _DeclaredProviderMemoryGenerateConfig)( + phase_one_model_settings={"provider_field": "phase-one"}, + phase_two_model_settings={"provider_field": "phase-two"}, + ) + + assert isinstance(config.phase_one_model_settings, _DeclaredProviderModelSettings) + assert config.phase_one_model_settings.provider_field == "phase-one" + assert isinstance(config.phase_two_model_settings, _DeclaredProviderModelSettings) + assert config.phase_two_model_settings.provider_field == "phase-two" + + +def test_memory_generate_config_model_settings_field_types_describe_normalized_values() -> None: + type_hints = get_type_hints(MemoryGenerateConfig) + + assert type_hints["phase_one_model_settings"] == ModelSettings | None + assert type_hints["phase_two_model_settings"] == ModelSettings | None + + +def test_memory_generate_config_preserves_typed_model_settings() -> None: + phase_one_settings = ModelSettings(reasoning={"effort": "low"}) + phase_two_settings = ModelSettings(temperature=0.2) + config = MemoryGenerateConfig( + phase_one_model_settings=phase_one_settings, + phase_two_model_settings=phase_two_settings, + ) + + assert config.phase_one_model_settings is phase_one_settings + assert config.phase_two_model_settings is phase_two_settings + + +@pytest.mark.parametrize( + "field_name", + ["phase_one_model_settings", "phase_two_model_settings"], +) +def test_memory_generate_config_preserves_forward_compatible_reasoning_settings( + field_name: str, +) -> None: + settings: dict[str, Any] = {field_name: {"reasoning": {"future_reasoning_option": "enabled"}}} + + config = MemoryGenerateConfig(**settings) + model_settings = getattr(config, field_name) + + assert model_settings is not None + assert model_settings.reasoning is not None + assert model_settings.reasoning.model_extra == {"future_reasoning_option": "enabled"} + + +@pytest.mark.parametrize( + "field_name", + ["phase_one_model_settings", "phase_two_model_settings"], +) +def test_memory_generate_config_rejects_invalid_model_settings(field_name: str) -> None: + settings: dict[str, Any] = {field_name: "invalid"} + with pytest.raises( + TypeError, + match=f"MemoryGenerateConfig.{field_name} must be a ModelSettings instance or a dict", + ): + MemoryGenerateConfig(**settings) + + +def test_memory_generate_config_preserves_disabled_model_settings() -> None: + config = MemoryGenerateConfig(phase_one_model_settings=None, phase_two_model_settings=None) + + assert config.phase_one_model_settings is None + assert config.phase_two_model_settings is None + + def test_memory_generate_config_rejects_too_many_raw_memories() -> None: with pytest.raises( ValueError, diff --git a/tests/sandbox/test_runtime_agent_preparation.py b/tests/sandbox/test_runtime_agent_preparation.py index eff4a3131a..c532f7e990 100644 --- a/tests/sandbox/test_runtime_agent_preparation.py +++ b/tests/sandbox/test_runtime_agent_preparation.py @@ -17,6 +17,49 @@ from agents.sandbox.manifest import Manifest from agents.sandbox.sandbox_agent import SandboxAgent from agents.sandbox.session.base_sandbox_session import BaseSandboxSession +from agents.sandbox.types import User + + +def test_sandbox_agent_normalizes_first_party_dictionary_configuration() -> None: + agent = SandboxAgent( + name="sandbox", + model_settings={"reasoning": {"context": "all_turns"}}, + default_manifest={"root": "/workspace"}, + run_as={"name": "agent"}, + ) + + assert agent.model_settings.reasoning is not None + assert agent.model_settings.reasoning.context == "all_turns" + assert isinstance(agent.default_manifest, Manifest) + assert isinstance(agent.run_as, User) + assert agent.run_as.name == "agent" + + +def test_sandbox_agent_rejects_untrusted_manifest_path_grants() -> None: + with pytest.raises( + TypeError, + match=( + r"sandbox\.default_manifest\.extra_path_grants must be configured " + r"on a trusted Manifest" + ), + ): + SandboxAgent(name="sandbox", default_manifest={"extra_path_grants": [{"path": "/tmp"}]}) + + +@pytest.mark.parametrize( + "manifest", + [ + Manifest(root="/workspace").model_dump(), + Manifest(root="/workspace").model_dump(mode="json"), + ], +) +def test_sandbox_agent_accepts_serialized_manifest_without_path_grants( + manifest: dict[str, Any], +) -> None: + agent = SandboxAgent(name="sandbox", default_manifest=manifest) + + assert isinstance(agent.default_manifest, Manifest) + assert agent.default_manifest.extra_path_grants == () class _Capability: diff --git a/tests/test_agent_config.py b/tests/test_agent_config.py index ad77eeb3e2..f935cfd7a7 100644 --- a/tests/test_agent_config.py +++ b/tests/test_agent_config.py @@ -1,9 +1,13 @@ +from typing import Any + import pytest +from openai.types.shared import Reasoning from pydantic import BaseModel from agents import Agent, AgentOutputSchema, Handoff, RunContextWrapper, handoff from agents.lifecycle import AgentHooksBase from agents.model_settings import ModelSettings +from agents.retry import ModelRetryBackoffSettings from agents.run_internal.run_loop import get_handoffs, get_output_schema @@ -216,11 +220,66 @@ def test_list_field_validation(self): def test_model_settings_validation(self): """Test model_settings validation - prevents runtime errors""" - # Valid case + # Typed settings and SDK-owned dictionaries are both valid. Agent(name="test", model_settings=ModelSettings()) + agent = Agent(name="test", model_settings={"temperature": 0.25}) + + assert isinstance(agent.model_settings, ModelSettings) + assert agent.model_settings.temperature == 0.25 - # Invalid case that could cause runtime issues + # Invalid values are rejected before model execution. with pytest.raises( - TypeError, match="Agent model_settings must be a ModelSettings instance" + TypeError, match="Agent model_settings must be a ModelSettings instance or a dict" ): - Agent(name="test", model_settings={}) # type: ignore + Agent(name="test", model_settings="invalid") # type: ignore[arg-type] + + +def test_agent_model_settings_dictionary_preserves_openai_reasoning_extensions() -> None: + agent = Agent( + name="test", + model_settings={ + "reasoning": {"context": "all_turns", "future_reasoning_option": "enabled"}, + "context_management": [{"type": "compaction", "compact_threshold": 244800}], + "retry": {"max_retries": 0, "backoff": {"jitter": False}}, + }, + ) + + assert isinstance(agent.model_settings.reasoning, Reasoning) + assert agent.model_settings.reasoning.context == "all_turns" + assert agent.model_settings.reasoning.model_extra == {"future_reasoning_option": "enabled"} + assert agent.model_settings.context_management == [ + {"type": "compaction", "compact_threshold": 244800} + ] + assert agent.model_settings.retry is not None + assert agent.model_settings.retry.max_retries == 0 + assert isinstance(agent.model_settings.retry.backoff, ModelRetryBackoffSettings) + assert agent.model_settings.retry.backoff.jitter is False + + +@pytest.mark.parametrize( + ("settings", "message"), + [ + ({"temperatur": 0.2}, "Unknown model settings: temperatur"), + ({"retry": {"max_retry": 2}}, "Unknown model settings in retry: max_retry"), + ( + {"retry": {"backoff": {"initial_delai": 1}}}, + "Unknown model settings in retry.backoff: initial_delai", + ), + ( + {"context_management": [{"type": "compaction", "compact_threshold_typo": 1}]}, + r"Unknown model settings in context_management\[0\]: compact_threshold_typo", + ), + ], +) +def test_agent_rejects_unknown_first_party_dictionary_model_settings( + settings: dict[str, Any], message: str +) -> None: + with pytest.raises(TypeError, match=message): + Agent(name="test", model_settings=settings) + + +@pytest.mark.parametrize("setting_name", ["reasoning", "context_management", "temperature"]) +def test_agent_does_not_promote_model_settings_to_constructor(setting_name: str) -> None: + arguments: dict[str, Any] = {setting_name: None} + with pytest.raises(TypeError, match=f"unexpected keyword argument '{setting_name}'"): + Agent(name="test", **arguments) diff --git a/tests/test_run_config.py b/tests/test_run_config.py index e3f78ae88f..7b99b649f2 100644 --- a/tests/test_run_config.py +++ b/tests/test_run_config.py @@ -2,9 +2,19 @@ import pytest -from agents import Agent, RunConfig, Runner, ToolExecutionConfig, ToolNotFoundBehavior +from agents import ( + Agent, + RunConfig, + Runner, + SessionSettings, + ToolExecutionConfig, + ToolNotFoundBehavior, +) from agents.model_settings import ModelSettings from agents.models.interface import Model, ModelProvider +from agents.run_config import SandboxConcurrencyLimits, SandboxRunConfig +from agents.sandbox.manifest import Manifest +from agents.sandbox.snapshot import NoopSnapshotSpec from .fake_model import FakeModel from .test_responses import get_text_message @@ -24,6 +34,99 @@ def get_model(self, model_name: str | None) -> Model: return self.model_to_return +def test_run_config_normalizes_first_party_dictionary_settings() -> None: + config = RunConfig( + model_settings={"reasoning": {"context": "all_turns"}, "temperature": 0.0}, + session_settings={"limit": 5}, + tool_execution={"max_function_tool_concurrency": 2}, + sandbox={ + "manifest": {"root": "/workspace"}, + "snapshot": {"type": "noop"}, + "concurrency_limits": {"manifest_entries": 3}, + }, + ) + + assert isinstance(config.model_settings, ModelSettings) + assert config.model_settings.reasoning is not None + assert config.model_settings.reasoning.context == "all_turns" + assert config.model_settings.temperature == 0.0 + assert isinstance(config.session_settings, SessionSettings) + assert config.session_settings.limit == 5 + assert isinstance(config.tool_execution, ToolExecutionConfig) + assert config.tool_execution.max_function_tool_concurrency == 2 + assert isinstance(config.sandbox, SandboxRunConfig) + assert isinstance(config.sandbox.manifest, Manifest) + assert isinstance(config.sandbox.snapshot, NoopSnapshotSpec) + assert isinstance(config.sandbox.concurrency_limits, SandboxConcurrencyLimits) + assert config.sandbox.concurrency_limits.manifest_entries == 3 + + +def test_run_config_preserves_typed_configuration_instances() -> None: + settings = ModelSettings(temperature=0.2) + session_settings = SessionSettings(limit=3) + config = RunConfig(model_settings=settings, session_settings=session_settings) + + assert config.model_settings is settings + assert config.session_settings is session_settings + + +def test_run_config_rejects_untrusted_manifest_path_grants() -> None: + with pytest.raises( + TypeError, + match=r"sandbox\.manifest\.extra_path_grants must be configured on a trusted Manifest", + ): + RunConfig(sandbox={"manifest": {"extra_path_grants": [{"path": "/tmp"}]}}) + + +@pytest.mark.parametrize( + "manifest", + [ + Manifest(root="/workspace").model_dump(), + Manifest(root="/workspace").model_dump(mode="json"), + ], +) +def test_run_config_accepts_serialized_manifest_without_path_grants( + manifest: dict[str, object], +) -> None: + config = RunConfig(sandbox={"manifest": manifest}) + + assert config.sandbox is not None + assert isinstance(config.sandbox.manifest, Manifest) + assert config.sandbox.manifest.extra_path_grants == () + + +@pytest.mark.parametrize( + ("settings", "message"), + [ + ({"model_settings": {"temperatur": 0.2}}, "Unknown model settings: temperatur"), + ({"session_settings": {"limitt": 2}}, "Unknown session settings: limitt"), + ( + {"tool_execution": {"max_function_tool_concurrenc": 2}}, + "Unknown run_config.tool_execution settings: max_function_tool_concurrenc", + ), + ], +) +def test_run_config_rejects_unknown_first_party_dictionary_fields( + settings: dict[str, object], message: str +) -> None: + with pytest.raises(TypeError, match=message): + RunConfig(**settings) # type: ignore[arg-type] + + +@pytest.mark.asyncio +async def test_runner_accepts_dictionary_run_configuration() -> None: + model = FakeModel(initial_output=[get_text_message("done")]) + agent = Agent(name="test", model=model) + + result = await Runner.run( + agent, + "hello", + run_config={"model_settings": {"temperature": 0.0}}, + ) + + assert result.final_output == "done" + + @pytest.mark.asyncio async def test_model_provider_on_run_config_is_used_for_agent_model_name() -> None: """ diff --git a/tests/test_tool_context.py b/tests/test_tool_context.py index 5f1f9c1976..05f4a8a859 100644 --- a/tests/test_tool_context.py +++ b/tests/test_tool_context.py @@ -95,6 +95,35 @@ def test_tool_context_constructor_accepts_agent_keyword() -> None: assert tool_ctx.agent is agent +def test_tool_context_constructor_normalizes_dictionary_run_config() -> None: + tool_ctx: ToolContext[dict[str, object]] = ToolContext( + context={}, + tool_name="my_tool", + tool_call_id="call-2", + tool_arguments="{}", + run_config={ + "tracing_disabled": True, + "model_settings": {"temperature": 0.0}, + }, + ) + + assert isinstance(tool_ctx.run_config, RunConfig) + assert tool_ctx.run_config.tracing_disabled is True + assert tool_ctx.run_config.model_settings is not None + assert tool_ctx.run_config.model_settings.temperature == 0.0 + + +def test_tool_context_constructor_rejects_unknown_dictionary_run_config_fields() -> None: + with pytest.raises(TypeError, match="Unknown run_config settings: tracin_disabled"): + ToolContext( + context={}, + tool_name="my_tool", + tool_call_id="call-2", + tool_arguments="{}", + run_config={"tracin_disabled": True}, + ) + + def test_tool_context_constructor_infers_namespace_from_tool_call() -> None: tool_call = ResponseFunctionToolCall( type="function_call", @@ -221,6 +250,25 @@ def test_tool_context_from_agent_context_prefers_explicit_run_config() -> None: assert tool_ctx.run_config is explicit_run_config +def test_tool_context_from_agent_context_normalizes_dictionary_run_config() -> None: + tool_call = ResponseFunctionToolCall( + type="function_call", + name="test_tool", + call_id="call-1", + arguments="{}", + ) + + tool_ctx = ToolContext.from_agent_context( + make_context_wrapper(), + tool_call_id="call-1", + tool_call=tool_call, + run_config={"tracing_disabled": True}, + ) + + assert isinstance(tool_ctx.run_config, RunConfig) + assert tool_ctx.run_config.tracing_disabled is True + + @pytest.mark.asyncio async def test_invoke_function_tool_passes_plain_run_context_when_requested() -> None: captured_context: RunContextWrapper[str] | None = None diff --git a/tests/voice/test_pipeline.py b/tests/voice/test_pipeline.py index c60dbf6161..d6f97bb0eb 100644 --- a/tests/voice/test_pipeline.py +++ b/tests/voice/test_pipeline.py @@ -1,6 +1,8 @@ from __future__ import annotations import asyncio +from dataclasses import dataclass, field +from typing import Any import numpy as np import numpy.typing as npt @@ -13,6 +15,7 @@ from agents.voice import ( AudioInput, StreamedAudioResult, + STTModelSettings, TTSModelSettings, VoicePipeline, VoicePipelineConfig, @@ -27,6 +30,22 @@ pass +@dataclass +class _ProviderSTTModelSettings(STTModelSettings): + provider_language: str | None = None + + +@dataclass +class _ProviderTTSModelSettings(TTSModelSettings): + provider_voice: str | None = None + + +@dataclass +class _ProviderVoicePipelineConfig(VoicePipelineConfig): + stt_settings: _ProviderSTTModelSettings = field(default_factory=_ProviderSTTModelSettings) + tts_settings: _ProviderTTSModelSettings = field(default_factory=_ProviderTTSModelSettings) + + def test_streamed_audio_result_odd_length_buffer_int16() -> None: result = StreamedAudioResult( FakeTTS(), @@ -40,6 +59,68 @@ def test_streamed_audio_result_odd_length_buffer_int16() -> None: assert transformed.tolist() == [1] +def test_voice_pipeline_config_normalizes_dictionary_settings() -> None: + config = VoicePipelineConfig( + stt_settings={"language": "ja", "temperature": 0.0}, + tts_settings={"voice": "alloy", "buffer_size": 1}, + ) + + assert isinstance(config.stt_settings, STTModelSettings) + assert config.stt_settings.language == "ja" + assert config.stt_settings.temperature == 0.0 + assert isinstance(config.tts_settings, TTSModelSettings) + assert config.tts_settings.voice == "alloy" + assert config.tts_settings.buffer_size == 1 + + +def test_voice_pipeline_config_subclass_uses_declared_settings_types() -> None: + config = _ProviderVoicePipelineConfig( + stt_settings={"provider_language": "ja"}, # type: ignore[arg-type] + tts_settings={"provider_voice": "voice"}, # type: ignore[arg-type] + ) + + assert isinstance(config.stt_settings, _ProviderSTTModelSettings) + assert config.stt_settings.provider_language == "ja" + assert isinstance(config.tts_settings, _ProviderTTSModelSettings) + assert config.tts_settings.provider_voice == "voice" + + +@pytest.mark.parametrize( + ("settings", "message"), + [ + ({"stt_settings": {"languge": "ja"}}, "Unknown voice.stt settings: languge"), + ({"tts_settings": {"voce": "alloy"}}, "Unknown voice.tts settings: voce"), + ], +) +def test_voice_pipeline_config_rejects_unknown_dictionary_settings( + settings: dict[str, Any], message: str +) -> None: + with pytest.raises(TypeError, match=message): + VoicePipelineConfig(**settings) + + +@pytest.mark.asyncio +async def test_voicepipeline_normalizes_nested_dictionary_config() -> None: + fake_stt = FakeSTT(["first"]) + fake_tts = FakeTTS() + pipeline = VoicePipeline( + workflow=FakeWorkflow([["out_1"]]), + stt_model=fake_stt, + tts_model=fake_tts, + config={ + "stt_settings": {"language": "ja"}, + "tts_settings": {"voice": "alloy", "buffer_size": 1}, + }, + ) + + result = await pipeline.run(AudioInput(buffer=np.zeros(2, dtype=np.int16))) + events, audio_chunks = await extract_events(result) + + assert isinstance(pipeline.config, VoicePipelineConfig) + assert events == ["turn_started", "audio", "turn_ended", "session_ended"] + await fake_tts.verify_audio("out_1", audio_chunks[0]) + + def test_streamed_audio_result_odd_length_buffer_float32() -> None: result = StreamedAudioResult( FakeTTS(),