diff --git a/helpers.md b/helpers.md index 9deaa95000..54902af655 100644 --- a/helpers.md +++ b/helpers.md @@ -604,3 +604,32 @@ lookup = pydantic_function_tool( ``` Callbacks can be async when used with `AsyncOpenAI`. Existing dictionary handlers still work. + +### Typed Agents output (beta) + +Pass a Pydantic model (or a Pydantic v2 dataclass) to generate the Agents output +schema and parse the completed answer. Schemas use the same normalization as +Responses; the API validates which schema features it supports. + +```python +from pydantic import BaseModel + +class Report(BaseModel): + summary: str + findings: list[str] + +with client.beta.agents.sessions.create( + agent={"model": MODEL}, environment={"type": "none"}, + input="Summarize the findings.", stream=True, output_type=Report, +) as stream: + result = stream.get_final_result() +print(result.output_parsed) +``` + +For a session already configured with that schema, use +`sessions.stream(session_id, input="Update the report.", output_type=Report)`. +This only selects the local parser; it does not change the session's schema. +`output_parsed` exposes the first parsed final text part; every final text part is validated. +`result.parse(Report)` parses an existing raw result. `AgentOutputParseError.result` +retains the completed raw answer if validation fails. With `AsyncOpenAI`, await +creation and the result getter, and use `async with`. diff --git a/src/openai/lib/beta/agents/__init__.py b/src/openai/lib/beta/agents/__init__.py index 77e1a1d3f3..6b004a47e7 100644 --- a/src/openai/lib/beta/agents/__init__.py +++ b/src/openai/lib/beta/agents/__init__.py @@ -5,7 +5,12 @@ function_tool as function_tool, pydantic_function_tool as pydantic_function_tool, ) -from ._result import AgentTurnResult as AgentTurnResult, AgentTurnResultError as AgentTurnResultError +from ._output import agent_text_format as agent_text_format +from ._result import ( + AgentTurnResult as AgentTurnResult, + AgentTurnResultError as AgentTurnResultError, + AgentOutputParseError as AgentOutputParseError, +) from ._stream import ( AgentSessionEventStream as AgentSessionEventStream, AsyncAgentSessionEventStream as AsyncAgentSessionEventStream, diff --git a/src/openai/lib/beta/agents/_output.py b/src/openai/lib/beta/agents/_output.py new file mode 100644 index 0000000000..91321ba9d5 --- /dev/null +++ b/src/openai/lib/beta/agents/_output.py @@ -0,0 +1,53 @@ +from __future__ import annotations + +from typing import Any, cast +from typing_extensions import TypeVar + +import pydantic + +from ._schema import model_schema +from ...._types import Omit +from ...._compat import PYDANTIC_V1 +from ..._pydantic import is_basemodel_type, is_dataclass_like_type +from ....types.beta.agent_text_param import AgentTextParam +from ....types.beta.text_format_param import TextFormatParamJSONSchema +from ....types.beta.agents.session_create_params import Agent + + +def validate_output_type(output_type: type[Any]) -> None: + if not is_basemodel_type(output_type) and not (is_dataclass_like_type(output_type) and not PYDANTIC_V1): + raise TypeError("Agents output_type must be a Pydantic model or a Pydantic v2 dataclass") + + +def agent_text_format(output_type: type[Any]) -> TextFormatParamJSONSchema: + """Beta: build the Agents JSON-schema format for a Pydantic model.""" + validate_output_type(output_type) + if is_basemodel_type(output_type): + schema = model_schema(output_type, strict=True) + else: + schema = model_schema(pydantic.TypeAdapter(output_type), strict=True) + return {"type": "json_schema", "schema": schema} + + +def with_output_schema(agent: Agent | Omit, output_type: type[Any] | None) -> Agent | Omit: + if output_type is None: + return agent + config = cast(Agent, {} if isinstance(agent, Omit) else dict(agent)) + text = cast(AgentTextParam, dict(config.get("text") or {})) + if text.get("format") is not None: + raise ValueError("Pass output_type or agent.text.format, not both") + text["format"] = agent_text_format(output_type) + config["text"] = text + return config + + +ResponseT = TypeVar("ResponseT") + + +def bind_output_type(response: ResponseT, output_type: type[Any] | None) -> ResponseT: + from ._stream import AgentSessionEventStream, AsyncAgentSessionEventStream + + if isinstance(response, (AgentSessionEventStream, AsyncAgentSessionEventStream)): + stream = cast("AgentSessionEventStream[Any] | AsyncAgentSessionEventStream[Any]", response) + stream._collection.output_type = output_type + return cast(ResponseT, response) diff --git a/src/openai/lib/beta/agents/_result.py b/src/openai/lib/beta/agents/_result.py index d42955ae2f..9bcb2ba43d 100644 --- a/src/openai/lib/beta/agents/_result.py +++ b/src/openai/lib/beta/agents/_result.py @@ -1,23 +1,50 @@ from __future__ import annotations from copy import deepcopy -from typing import Iterable +from typing import Any, Generic, Iterable from dataclasses import dataclass -from typing_extensions import Literal +from typing_extensions import Literal, TypeVar +from ._output import validate_output_type from ...._exceptions import OpenAIError +from ..._parsing._completions import _parse_content from ....types.beta.agent_session import RequiredAction from ....types.beta.agent_session_event import AgentSessionEvent from ....types.beta.agents.sessions.turn import Turn from ....types.beta.agent_session_message import AgentSessionMessage +OutputT = TypeVar("OutputT", default=Any) +ParseT = TypeVar("ParseT") + @dataclass(frozen=True) -class AgentTurnResult: +class AgentTurnResult(Generic[OutputT]): """Beta: the completed final assistant messages from one successful root turn.""" turn: Turn messages: list[AgentSessionMessage] + output_parsed: OutputT | None = None + + def parse(self, output_type: type[ParseT]) -> AgentTurnResult[ParseT]: + """Beta: parse a completed answer without changing its session configuration.""" + try: + validate_output_type(output_type) + parsed: ParseT | None = None + found = False + for message in self.messages: + for content in message.content: + if content.type != "output_text": + continue + value = _parse_content(output_type, content.text) + if not found: + parsed = value + found = True + if not found: + raise ValueError("No final output text to parse") + except Exception: + # Pydantic errors may include response text in their rendered message. + raise AgentOutputParseError(self) from None + return AgentTurnResult(turn=self.turn, messages=self.messages, output_parsed=parsed) @property def session_id(self) -> str: @@ -33,6 +60,14 @@ def output_text(self) -> str: return "".join(message.output_text for message in self.messages) +class AgentOutputParseError(OpenAIError): + """Beta: output parsing failed after a successful hosted turn.""" + + def __init__(self, result: AgentTurnResult) -> None: + super().__init__("Could not parse the completed agent output") + self.result = result + + ResultErrorReason = Literal["failed", "cancelled", "requires_action", "incomplete", "observation_failed"] @@ -160,10 +195,14 @@ def result(self) -> AgentTurnResult: return self._result -class AgentTurnResultCollection: +class AgentTurnResultCollection(Generic[OutputT]): """Keep ordinary event iteration incremental until collection is requested.""" - def __init__(self, session_id: str | None = None) -> None: + def __init__(self, session_id: str | None = None, output_type: type[OutputT] | None = None) -> None: + if output_type is not None: + validate_output_type(output_type) + self.output_type = output_type + self._parsed_result: AgentTurnResult[OutputT] | None = None self.collector: AgentTurnResultCollector | None = None self._session_id = session_id self._started = False @@ -185,3 +224,11 @@ def accept(self, event: AgentSessionEvent) -> None: def record_error(self, error: Exception) -> None: if self.collector is not None and not self.collector.is_done(): self.collector.cause = error + + def result(self) -> AgentTurnResult[OutputT]: + result = self.enable().result() + if self.output_type is None: + return result + if self._parsed_result is None: + self._parsed_result = result.parse(self.output_type) + return self._parsed_result diff --git a/src/openai/lib/beta/agents/_schema.py b/src/openai/lib/beta/agents/_schema.py new file mode 100644 index 0000000000..35abaedb03 --- /dev/null +++ b/src/openai/lib/beta/agents/_schema.py @@ -0,0 +1,22 @@ +from __future__ import annotations + +from copy import deepcopy +from typing import Any, cast + +import pydantic + +from ...._compat import PYDANTIC_V1, model_json_schema +from ..._pydantic import _ensure_strict_json_schema + + +def model_schema( + model: type[pydantic.BaseModel] | pydantic.TypeAdapter[Any], *, strict: bool = False +) -> dict[str, Any]: + # Pydantic v1 caches schema dictionaries. Input and output policies must not + # mutate each other's schemas or the model's cache. + schema = deepcopy( + model.json_schema() + if not PYDANTIC_V1 and isinstance(model, pydantic.TypeAdapter) + else model_json_schema(cast(type[pydantic.BaseModel], model)) + ) + return _ensure_strict_json_schema(schema, path=(), root=schema) if strict else schema diff --git a/src/openai/lib/beta/agents/_stream.py b/src/openai/lib/beta/agents/_stream.py index 10d5c3030b..3539b8f994 100644 --- a/src/openai/lib/beta/agents/_stream.py +++ b/src/openai/lib/beta/agents/_stream.py @@ -1,15 +1,15 @@ from __future__ import annotations -from typing import Iterator, AsyncIterator +from typing import Generic, Iterator, AsyncIterator from typing_extensions import Self, override -from ._result import AgentTurnResult, AgentTurnResultError, AgentTurnResultCollection +from ._result import OutputT, AgentTurnResult, AgentTurnResultError, AgentOutputParseError, AgentTurnResultCollection from ...._compat import cached_property from ...._streaming import Stream, AsyncStream from ....types.beta.agent_session_event import AgentSessionEvent -class AgentSessionEventStream(Stream[AgentSessionEvent]): +class AgentSessionEventStream(Stream[AgentSessionEvent], Generic[OutputT]): """Beta: a creation event stream that can collect its initial turn's result. Event iteration and response access behave like ``Stream``. Collection does @@ -18,7 +18,7 @@ class AgentSessionEventStream(Stream[AgentSessionEvent]): """ @cached_property - def _collection(self) -> AgentTurnResultCollection: + def _collection(self) -> AgentTurnResultCollection[OutputT]: return AgentTurnResultCollection() @override @@ -36,7 +36,7 @@ def with_result_collection(self) -> Self: self._collection.enable() return self - def get_final_result(self) -> AgentTurnResult: + def get_final_result(self) -> AgentTurnResult[OutputT]: """Consume remaining events and return the successful initial turn result. Raises AgentTurnResultError for an unsuccessful or unobserved outcome. @@ -52,8 +52,8 @@ def get_final_result(self) -> AgentTurnResult: collector.check_outcome() if collector.is_done(): break - return collector.result() - except AgentTurnResultError: + return self._collection.result() + except (AgentTurnResultError, AgentOutputParseError): raise except Exception as error: collector.cause = error @@ -62,11 +62,11 @@ def get_final_result(self) -> AgentTurnResult: self.close() -class AsyncAgentSessionEventStream(AsyncStream[AgentSessionEvent]): +class AsyncAgentSessionEventStream(AsyncStream[AgentSessionEvent], Generic[OutputT]): """Beta: asynchronous counterpart of AgentSessionEventStream.""" @cached_property - def _collection(self) -> AgentTurnResultCollection: + def _collection(self) -> AgentTurnResultCollection[OutputT]: return AgentTurnResultCollection() @override @@ -84,7 +84,7 @@ def with_result_collection(self) -> Self: self._collection.enable() return self - async def get_final_result(self) -> AgentTurnResult: + async def get_final_result(self) -> AgentTurnResult[OutputT]: """Consume remaining events and return the successful initial turn result. Enables collection on a fresh stream. To iterate first, call @@ -98,8 +98,8 @@ async def get_final_result(self) -> AgentTurnResult: collector.check_outcome() if collector.is_done(): break - return collector.result() - except AgentTurnResultError: + return self._collection.result() + except (AgentTurnResultError, AgentOutputParseError): raise except Exception as error: collector.cause = error diff --git a/src/openai/lib/beta/agents/_tools.py b/src/openai/lib/beta/agents/_tools.py index 79e4fc77ea..2c62fc17cf 100644 --- a/src/openai/lib/beta/agents/_tools.py +++ b/src/openai/lib/beta/agents/_tools.py @@ -11,8 +11,9 @@ import pydantic +from ._schema import model_schema from ...._utils import is_dict -from ...._compat import PYDANTIC_V1, model_dump, model_json, model_parse, model_json_schema +from ...._compat import PYDANTIC_V1, model_dump, model_json, model_parse from ..._pydantic import resolve_ref from ...streaming.agents._types import ToolOutput from ....types.beta.agent_tool_param import AgentToolConfigParamFunction @@ -39,7 +40,7 @@ def __init__( ) -> None: if not name: raise ValueError("Tool name must not be empty") - parameters = model_json_schema(model) + parameters = model_schema(model) root = parameters seen: set[str] = set() while isinstance(ref := root.get("$ref"), str) and ref not in seen: diff --git a/src/openai/lib/streaming/agents/_streams.py b/src/openai/lib/streaming/agents/_streams.py index 9eeacb4455..3be252e9e6 100644 --- a/src/openai/lib/streaming/agents/_streams.py +++ b/src/openai/lib/streaming/agents/_streams.py @@ -4,7 +4,7 @@ from copy import deepcopy from uuid import uuid4 from types import TracebackType -from typing import TYPE_CHECKING, Any, Mapping, Iterable, Iterator, AsyncIterator +from typing import TYPE_CHECKING, Any, Generic, Mapping, Iterable, Iterator, AsyncIterator from collections import deque from typing_extensions import Self, TypedDict @@ -15,7 +15,13 @@ from ...._types import Omit, Headers, NotGiven, omit, not_given from ...._streaming import Stream, AsyncStream from ...._exceptions import BadRequestError -from ...beta.agents._result import AgentTurnResult, AgentTurnResultError, AgentTurnResultCollection +from ...beta.agents._result import ( + OutputT, + AgentTurnResult, + AgentTurnResultError, + AgentOutputParseError, + AgentTurnResultCollection, +) from ....types.beta.agent_session_event import AgentSessionEvent from ....types.beta.agent_function_call_item import AgentFunctionCallItem from ....types.beta.agent_session_input_param import ( @@ -110,7 +116,7 @@ def call(self, event: AgentSessionEvent) -> AgentFunctionCallItem | None: return call -class AgentSessionStream: +class AgentSessionStream(Generic[OutputT]): """Submit input to an idle session and stream through the resulting turn's terminal session event. Use as a context manager. Only one caller may submit input to this session while @@ -133,6 +139,7 @@ def __init__( session_id: str, *, input: str | Iterable[AgentSessionInputMessageParam], + output_type: type[OutputT] | None = None, tool_handlers: Mapping[str, ToolHandler] | None = None, idempotency_key: str | Omit = omit, extra_headers: Headers | None = None, @@ -145,7 +152,7 @@ def __init__( self._idempotency_key = _input_key(idempotency_key, extra_headers) self._options = _request_options({"extra_headers": extra_headers, "timeout": timeout}) self._state = _TurnState() - self._collection = AgentTurnResultCollection(session_id) + self._collection: AgentTurnResultCollection[OutputT] = AgentTurnResultCollection(session_id, output_type) self._stream: Stream[AgentSessionEvent] | None = None self._iterator: Iterator[AgentSessionEvent] | None = None self._entered = False @@ -196,7 +203,7 @@ def with_result_collection(self) -> Self: self._collection.enable() return self - def get_final_result(self) -> AgentTurnResult: + def get_final_result(self) -> AgentTurnResult[OutputT]: """Beta: drain this turn, executing handlers, and collect its final answer. Raises AgentTurnResultError when a successful complete result cannot be @@ -212,8 +219,8 @@ def get_final_result(self) -> AgentTurnResult: collector.check_outcome(self._handlers) if collector.is_done(): break - return collector.result() - except AgentTurnResultError: + return self._collection.result() + except (AgentTurnResultError, AgentOutputParseError): raise except Exception as error: self._collection.record_error(error) @@ -280,7 +287,7 @@ def _submit_result(self, result: SessionInputParamAgentSessionInputToolResult) - self._sessions._sleep(delay) -class AsyncAgentSessionStream: +class AsyncAgentSessionStream(Generic[OutputT]): """Async counterpart of AgentSessionStream; use with ``async with``. Requires an idle session with a single input writer. Handlers may return a @@ -295,6 +302,7 @@ def __init__( session_id: str, *, input: str | Iterable[AgentSessionInputMessageParam], + output_type: type[OutputT] | None = None, tool_handlers: Mapping[str, AsyncToolHandler] | None = None, idempotency_key: str | Omit = omit, extra_headers: Headers | None = None, @@ -307,7 +315,7 @@ def __init__( self._idempotency_key = _input_key(idempotency_key, extra_headers) self._options = _request_options({"extra_headers": extra_headers, "timeout": timeout}) self._state = _TurnState() - self._collection = AgentTurnResultCollection(session_id) + self._collection: AgentTurnResultCollection[OutputT] = AgentTurnResultCollection(session_id, output_type) self._stream: AsyncStream[AgentSessionEvent] | None = None self._iterator: AsyncIterator[AgentSessionEvent] | None = None self._entered = False @@ -358,7 +366,7 @@ def with_result_collection(self) -> Self: self._collection.enable() return self - async def get_final_result(self) -> AgentTurnResult: + async def get_final_result(self) -> AgentTurnResult[OutputT]: """Beta: drain this turn, executing handlers, and collect its final answer. Raises AgentTurnResultError when a successful complete result cannot be @@ -374,8 +382,8 @@ async def get_final_result(self) -> AgentTurnResult: collector.check_outcome(self._handlers) if collector.is_done(): break - return collector.result() - except AgentTurnResultError: + return self._collection.result() + except (AgentTurnResultError, AgentOutputParseError): raise except Exception as error: self._collection.record_error(error) diff --git a/src/openai/resources/beta/agents/sessions/sessions.py b/src/openai/resources/beta/agents/sessions/sessions.py index c3b4f4dbb3..80edcbc259 100644 --- a/src/openai/resources/beta/agents/sessions/sessions.py +++ b/src/openai/resources/beta/agents/sessions/sessions.py @@ -66,6 +66,8 @@ ) from .....types.beta.agents import session_list_params, session_create_params, session_update_params from .....lib.streaming.agents import ToolHandler, AsyncToolHandler, AgentSessionStream, AsyncAgentSessionStream +from .....lib.beta.agents._output import bind_output_type, with_output_schema +from .....lib.beta.agents._result import OutputT from .....types.beta.agent_session import AgentSession from .....types.beta.environment_param import EnvironmentParam from .....types.beta.agent_session_deleted import AgentSessionDeleted @@ -80,11 +82,12 @@ def stream( session_id: str, *, input: str | Iterable[AgentSessionInputMessageParam], + output_type: type[OutputT] | None = None, tool_handlers: Mapping[str, ToolHandler] | None = None, idempotency_key: str | Omit = omit, extra_headers: Headers | None = None, timeout: float | httpx2.Timeout | None | NotGiven = not_given, - ) -> AgentSessionStream: + ) -> AgentSessionStream[OutputT]: """Stream one turn of an idle session, subscribing before submitting input. Use as a context manager. Only one caller may submit input to the session @@ -96,6 +99,7 @@ def stream( session_id, input=input, tool_handlers=tool_handlers, + output_type=output_type, idempotency_key=idempotency_key, extra_headers=extra_headers, timeout=timeout, @@ -149,6 +153,7 @@ def create( self, *, environment: EnvironmentParam, + output_type: type[OutputT] | None = None, agent: session_create_params.Agent | Omit = omit, agent_id: str | Omit = omit, input: Union[str, Iterable[AgentSessionInputMessageParam], None] | Omit = omit, @@ -203,6 +208,7 @@ def create( self, *, environment: EnvironmentParam, + output_type: type[OutputT] | None = None, stream: Literal[True], agent: session_create_params.Agent | Omit = omit, agent_id: str | Omit = omit, @@ -215,7 +221,7 @@ def create( extra_query: Query | None = None, extra_body: Body | None = None, timeout: float | httpx2.Timeout | None | NotGiven = not_given, - ) -> AgentSessionEventStream: + ) -> AgentSessionEventStream[OutputT]: """ Creates a managed agent session, optionally submits initial input, and returns the session or streams its events when stream is true. See @@ -257,6 +263,7 @@ def create( self, *, environment: EnvironmentParam, + output_type: type[OutputT] | None = None, stream: bool, agent: session_create_params.Agent | Omit = omit, agent_id: str | Omit = omit, @@ -269,7 +276,7 @@ def create( extra_query: Query | None = None, extra_body: Body | None = None, timeout: float | httpx2.Timeout | None | NotGiven = not_given, - ) -> AgentSession | AgentSessionEventStream: + ) -> AgentSession | AgentSessionEventStream[OutputT]: """ Creates a managed agent session, optionally submits initial input, and returns the session or streams its events when stream is true. See @@ -311,6 +318,7 @@ def create( self, *, environment: EnvironmentParam, + output_type: type[OutputT] | None = None, agent: session_create_params.Agent | Omit = omit, agent_id: str | Omit = omit, input: Union[str, Iterable[AgentSessionInputMessageParam], None] | Omit = omit, @@ -323,8 +331,9 @@ def create( extra_query: Query | None = None, extra_body: Body | None = None, timeout: float | httpx2.Timeout | None | NotGiven = not_given, - ) -> AgentSession | AgentSessionEventStream: + ) -> AgentSession | AgentSessionEventStream[OutputT]: extra_headers = {"OpenAI-Beta": "agents=v1", **(extra_headers or {})} + agent = with_output_schema(agent, output_type) return self._post( "/agents/sessions", body=maybe_transform( @@ -342,6 +351,7 @@ def create( else session_create_params.SessionCreateParamsNonStreaming, ), options=make_request_options( + post_parser=lambda response: bind_output_type(response, output_type), extra_headers=extra_headers, extra_query=extra_query, extra_body=extra_body, @@ -562,11 +572,12 @@ def stream( session_id: str, *, input: str | Iterable[AgentSessionInputMessageParam], + output_type: type[OutputT] | None = None, tool_handlers: Mapping[str, AsyncToolHandler] | None = None, idempotency_key: str | Omit = omit, extra_headers: Headers | None = None, timeout: float | httpx2.Timeout | None | NotGiven = not_given, - ) -> AsyncAgentSessionStream: + ) -> AsyncAgentSessionStream[OutputT]: """Stream one turn of an idle session, subscribing before submitting input. Use as an async context manager. Only one caller may submit input to the session @@ -578,6 +589,7 @@ def stream( session_id, input=input, tool_handlers=tool_handlers, + output_type=output_type, idempotency_key=idempotency_key, extra_headers=extra_headers, timeout=timeout, @@ -631,6 +643,7 @@ async def create( self, *, environment: EnvironmentParam, + output_type: type[OutputT] | None = None, agent: session_create_params.Agent | Omit = omit, agent_id: str | Omit = omit, input: Union[str, Iterable[AgentSessionInputMessageParam], None] | Omit = omit, @@ -685,6 +698,7 @@ async def create( self, *, environment: EnvironmentParam, + output_type: type[OutputT] | None = None, stream: Literal[True], agent: session_create_params.Agent | Omit = omit, agent_id: str | Omit = omit, @@ -697,7 +711,7 @@ async def create( extra_query: Query | None = None, extra_body: Body | None = None, timeout: float | httpx2.Timeout | None | NotGiven = not_given, - ) -> AsyncAgentSessionEventStream: + ) -> AsyncAgentSessionEventStream[OutputT]: """ Creates a managed agent session, optionally submits initial input, and returns the session or streams its events when stream is true. See @@ -739,6 +753,7 @@ async def create( self, *, environment: EnvironmentParam, + output_type: type[OutputT] | None = None, stream: bool, agent: session_create_params.Agent | Omit = omit, agent_id: str | Omit = omit, @@ -751,7 +766,7 @@ async def create( extra_query: Query | None = None, extra_body: Body | None = None, timeout: float | httpx2.Timeout | None | NotGiven = not_given, - ) -> AgentSession | AsyncAgentSessionEventStream: + ) -> AgentSession | AsyncAgentSessionEventStream[OutputT]: """ Creates a managed agent session, optionally submits initial input, and returns the session or streams its events when stream is true. See @@ -793,6 +808,7 @@ async def create( self, *, environment: EnvironmentParam, + output_type: type[OutputT] | None = None, agent: session_create_params.Agent | Omit = omit, agent_id: str | Omit = omit, input: Union[str, Iterable[AgentSessionInputMessageParam], None] | Omit = omit, @@ -805,8 +821,9 @@ async def create( extra_query: Query | None = None, extra_body: Body | None = None, timeout: float | httpx2.Timeout | None | NotGiven = not_given, - ) -> AgentSession | AsyncAgentSessionEventStream: + ) -> AgentSession | AsyncAgentSessionEventStream[OutputT]: extra_headers = {"OpenAI-Beta": "agents=v1", **(extra_headers or {})} + agent = with_output_schema(agent, output_type) return await self._post( "/agents/sessions", body=await async_maybe_transform( @@ -824,6 +841,7 @@ async def create( else session_create_params.SessionCreateParamsNonStreaming, ), options=make_request_options( + post_parser=lambda response: bind_output_type(response, output_type), extra_headers=extra_headers, extra_query=extra_query, extra_body=extra_body, diff --git a/tests/lib/streaming/agents/test_output.py b/tests/lib/streaming/agents/test_output.py new file mode 100644 index 0000000000..ffa5c61f81 --- /dev/null +++ b/tests/lib/streaming/agents/test_output.py @@ -0,0 +1,552 @@ +from __future__ import annotations + +import json +import traceback +from typing import Any, cast +from typing_extensions import Literal + +import httpx2 +import pytest +import pydantic +from pydantic import Field, HttpUrl, BaseModel + +from openai import OpenAI, AsyncOpenAI, BadRequestError +from openai._compat import PYDANTIC_V1 +from openai.lib.beta.agents import ( + AgentTurnResultError, + AgentOutputParseError, + function_tool, + agent_text_format, + pydantic_function_tool, +) +from openai.lib._parsing._responses import type_to_text_format_param +from tests.lib.streaming.agents.test_results import ResultServer, server as server, message +from tests.lib.streaming.agents.test_streams import EventBody, sdk as sdk, call, idle, turn_event + + +def _assert_matches_responses(model: type[Any]) -> dict[str, Any]: + schema = agent_text_format(model)["schema"] + responses_format = type_to_text_format_param(model) + assert responses_format["type"] == "json_schema" + assert schema == responses_format["schema"] + return schema + + +class Report(BaseModel): + summary: str + findings: list[str] + + +class NestedReport(BaseModel): + report: Report + alternative: Report | None + + +def test_agents_format() -> None: + format = agent_text_format(NestedReport) + assert set(format) == {"type", "schema"} + assert format["type"] == "json_schema" + schema = format["schema"] + assert schema["type"] == "object" + assert schema["additionalProperties"] is False + assert schema["required"] == ["report", "alternative"] + assert "name" not in format and "strict" not in format + + +def test_mapping_schema_matches_responses_without_mutating_tools() -> None: + class WithMap(BaseModel): + value: dict[str, str] + + tool = pydantic_function_tool(WithMap, handler=lambda args: args.value) + parameters = tool.definition["parameters"] + schema = _assert_matches_responses(WithMap) + assert schema["properties"]["value"]["additionalProperties"] == {"type": "string"} + assert tool.definition["parameters"] == parameters + assert tool({"value": {"label": "example"}}) == '{"label":"example"}' + + +def test_schema_normalizes_defaults_without_mutating_model() -> None: + class WithDefault(BaseModel): + value: str = "default" + + assert agent_text_format(WithDefault)["schema"]["required"] == ["value"] + assert WithDefault().value == "default" + + +def test_unique_items_matches_responses() -> None: + class WithSet(BaseModel): + value: set[str] + + schema = _assert_matches_responses(WithSet) + assert schema["properties"]["value"]["uniqueItems"] is True + + +@pytest.mark.parametrize("creation", [False, True]) +@pytest.mark.parametrize("progress", [False, True]) +async def test_typed_result(sdk: OpenAI | AsyncOpenAI, server: ResultServer, creation: bool, progress: bool) -> None: + text = '{"summary":"done","findings":["one"]}' + server.body = EventBody([turn_event("created"), message(text), turn_event("completed"), idle()]) + agent: Any = {"model": "test-model", "text": {"verbosity": "low"}} + if isinstance(sdk, AsyncOpenAI): + stream = ( + ( + await sdk.beta.agents.sessions.create( + agent=agent, + environment={"type": "none"}, + input="Question", + stream=True, + output_type=Report, + ) + ) + if creation + else sdk.beta.agents.sessions.stream("session_test", input="Question", output_type=Report) + ) + async with stream: + if progress: + stream.with_result_collection() + async for _ in stream: + pass + result = await stream.get_final_result() + assert await stream.get_final_result() is result + else: + stream = ( + sdk.beta.agents.sessions.create( + agent=agent, + environment={"type": "none"}, + input="Question", + stream=True, + output_type=Report, + ) + if creation + else sdk.beta.agents.sessions.stream("session_test", input="Question", output_type=Report) + ) + with stream: + if progress: + stream.with_result_collection() + for _ in stream: + pass + result = stream.get_final_result() + assert stream.get_final_result() is result + assert result.output_parsed == Report(summary="done", findings=["one"]) + assert result.output_text == text + assert result.turn_id == "turn_root" + assert result.messages[0].output_text == text + assert agent == {"model": "test-model", "text": {"verbosity": "low"}} + bodies = [json.loads(request.content) for request in server.requests if request.method == "POST"] + assert all("output_type" not in body for body in bodies) + if creation: + assert bodies[0]["agent"]["text"] == {"verbosity": "low", "format": agent_text_format(Report)} + else: + assert all("agent" not in body for body in bodies) + + +@pytest.mark.parametrize( + "text", ["SYNTHETIC_RESPONSE_CANARY not json", '{"summary": "SYNTHETIC_RESPONSE_CANARY missing findings"}'] +) +async def test_parse_error_preserves_successful_raw_result( + sdk: OpenAI | AsyncOpenAI, server: ResultServer, text: str +) -> None: + server.body = EventBody([turn_event("created"), message(text), turn_event("completed"), idle()]) + with pytest.raises(AgentOutputParseError) as caught: + if isinstance(sdk, AsyncOpenAI): + async with sdk.beta.agents.sessions.stream("session_test", input="Question", output_type=Report) as stream: + await stream.get_final_result() + else: + with sdk.beta.agents.sessions.stream("session_test", input="Question", output_type=Report) as stream: + stream.get_final_result() + assert caught.value.result.output_text == text + assert caught.value.result.turn.status == "completed" + assert caught.value.result.output_parsed is None + assert caught.value.__cause__ is None + assert text not in "".join(traceback.format_exception(caught.value)) + assert server.body.closed + + +async def test_hosted_failure_not_parse_failure(sdk: OpenAI | AsyncOpenAI, server: ResultServer) -> None: + server.body = EventBody([turn_event("created"), message("not json"), turn_event("failed"), idle()]) + with pytest.raises(AgentTurnResultError, match="failed"): + if isinstance(sdk, AsyncOpenAI): + async with sdk.beta.agents.sessions.stream("session_test", input="Question", output_type=Report) as stream: + await stream.get_final_result() + else: + with sdk.beta.agents.sessions.stream("session_test", input="Question", output_type=Report) as stream: + stream.get_final_result() + + +async def test_conflicting_format_rejected_before_request(sdk: OpenAI | AsyncOpenAI, server: ResultServer) -> None: + with pytest.raises(ValueError, match="not both"): + if isinstance(sdk, AsyncOpenAI): + await sdk.beta.agents.sessions.create( + agent={"model": "test", "text": {"format": {"type": "text"}}}, + environment={"type": "none"}, + input="Question", + stream=True, + output_type=Report, + ) + else: + sdk.beta.agents.sessions.create( + agent={"model": "test", "text": {"format": {"type": "text"}}}, + environment={"type": "none"}, + input="Question", + stream=True, + output_type=Report, + ) + assert server.requests == [] + + +class Cat(BaseModel): + kind: Literal["cat"] + lives: int + + +class Dog(BaseModel): + kind: Literal["dog"] + name: str + + +class PetReport(BaseModel): + pet: Cat | Dog = Field(discriminator="kind") + + +def test_discriminated_union_matches_responses() -> None: + schema = _assert_matches_responses(PetReport) + assert "oneOf" in schema["properties"]["pet"] + + +def test_fixed_tuple_matches_responses() -> None: + class TupleReport(BaseModel): + coordinates: tuple[int, int] + + _assert_matches_responses(TupleReport) + + +class RecursiveReport(BaseModel): + summary: str + children: list[RecursiveReport] + + +def test_recursive_references_preserved() -> None: + schema = agent_text_format(RecursiveReport)["schema"] + assert schema["type"] == "object" + assert "$defs" in schema or "definitions" in schema + + +@pytest.mark.parametrize( + "model", + [ + type("AnyReport", (BaseModel,), {"__annotations__": {"value": Any}}), + type("URLReport", (BaseModel,), {"__annotations__": {"url": HttpUrl}}), + ], +) +def test_field_types_match_responses(model: type[BaseModel]) -> None: + _assert_matches_responses(model) + + +class NullableChild(BaseModel): + value: str | None = None + + +class NullableReport(BaseModel): + optional: str | None = None + child: NullableChild | None = None + values: list[str | None] + choice: int | str | None = None + + +def test_nullable_schema_matches_responses() -> None: + schema = _assert_matches_responses(NullableReport) + assert schema["required"] == ["optional", "child", "values", "choice"] + assert NullableReport(values=[None]).optional is None + + +@pytest.mark.skipif(PYDANTIC_V1, reason="Pydantic dataclass output requires v2, matching Responses") +async def test_pydantic_dataclass_output(sdk: OpenAI | AsyncOpenAI, server: ResultServer) -> None: + from pydantic.dataclasses import dataclass + + @dataclass + class DataclassReport: + summary: str + + server.body = EventBody([turn_event("created"), message('{"summary":"done"}'), turn_event("completed"), idle()]) + assert agent_text_format(DataclassReport)["schema"]["type"] == "object" + if isinstance(sdk, AsyncOpenAI): + async with sdk.beta.agents.sessions.stream( + "session_test", input="Question", output_type=DataclassReport + ) as stream: + result = await stream.get_final_result() + else: + with sdk.beta.agents.sessions.stream("session_test", input="Question", output_type=DataclassReport) as stream: + result = stream.get_final_result() + assert result.output_parsed == DataclassReport(summary="done") + + +class NullableTree(BaseModel): + label: str + parent: NullableTree | None = None + children: list[NullableTree | None] + + +def test_nullable_recursive_refs_and_aliases_match_responses() -> None: + _assert_matches_responses(NullableTree) + + class Aliased(BaseModel): + count: int | None = Field(None, alias="optionalCount") + + alias_schema = _assert_matches_responses(Aliased) + assert alias_schema["required"] == ["optionalCount"] + + +class LocalURLReport(BaseModel): + url: HttpUrl + + +async def test_followup_output_type_only_selects_local_parser(sdk: OpenAI | AsyncOpenAI, server: ResultServer) -> None: + server.body = EventBody( + [turn_event("created"), message('{"url":"https://example.com/"}'), turn_event("completed"), idle()] + ) + if isinstance(sdk, AsyncOpenAI): + async with sdk.beta.agents.sessions.stream( + "session_test", input="Question", output_type=LocalURLReport + ) as stream: + result = await stream.get_final_result() + else: + with sdk.beta.agents.sessions.stream("session_test", input="Question", output_type=LocalURLReport) as stream: + result = stream.get_final_result() + assert result.output_parsed is not None + assert str(result.output_parsed.url) == "https://example.com/" + assert "output_type" not in str(server.inputs()) + _assert_matches_responses(LocalURLReport) + + +@pytest.mark.parametrize("value", ['say "hello"', "two\nlines", (1, 2)]) +def test_enum_literals_match_responses(value: Any) -> None: + from enum import Enum + + enum = Enum("Example", {"A": value, "B": "other"}) + model = type("EnumReport", (BaseModel,), {"__annotations__": {"value": enum}}) + _assert_matches_responses(model) + + +def test_property_literals_match_responses() -> None: + class Aliased(BaseModel): + value: str = Field(alias='a"b') + + _assert_matches_responses(Aliased) + + +def test_pattern_keyed_mapping_matches_responses() -> None: + from pydantic import constr + + key_type = cast(Any, constr)(**{("regex" if PYDANTIC_V1 else "pattern"): "^item_"}) + model = type("PatternMapping", (BaseModel,), {"__annotations__": {"values": dict[key_type, str]}}) + schema = _assert_matches_responses(model) + assert schema["properties"]["values"]["patternProperties"] == {"^item_": {"type": "string"}} + + +def test_bare_mapping_and_empty_fixed_object_match_responses() -> None: + model = type("BareMapping", (BaseModel,), {"__annotations__": {"value": dict}}) + _assert_matches_responses(model) + empty = type("EmptyReport", (BaseModel,), {}) + assert _assert_matches_responses(empty)["properties"] == {} + + +@pytest.mark.parametrize("value", ['say "hello"', "two\nlines"]) +def test_single_literals_match_responses(value: Any) -> None: + model = type("LiteralReport", (BaseModel,), {"__annotations__": {"value": Literal[value]}}) + _assert_matches_responses(model) + + +@pytest.mark.parametrize("explicit_model", [False, True]) +async def test_typed_tools_and_output_keep_distinct_schema_policies( + sdk: OpenAI | AsyncOpenAI, server: ResultServer, explicit_model: bool +) -> None: + class Summary(BaseModel): + summary: str = "default" + + seen: list[str] = [] + + def summarize(summary: str = "default") -> Summary: + seen.append(summary) + return Summary(summary=summary) + + tool = ( + pydantic_function_tool(Summary, name="search", handler=lambda args: summarize(args.summary)) + if explicit_model + else function_tool(summarize, name="search") + ) + parameters = tool.definition["parameters"] + assert "summary" not in cast(list[str], parameters.get("required", [])) + output_format = agent_text_format(Summary) + assert output_format["schema"]["required"] == ["summary"] + assert tool.definition["parameters"] == parameters + later_tool = pydantic_function_tool(Summary, handler=lambda args: args) + assert "summary" not in cast(list[str], later_tool.definition["parameters"].get("required", [])) + text = '{"summary":"default"}' + server.body = EventBody([turn_event("created"), call({}), message(text), turn_event("completed"), idle()]) + if isinstance(sdk, AsyncOpenAI): + async with sdk.beta.agents.sessions.stream( + "session_test", input="Summarize the document.", tool_handlers={tool.name: tool}, output_type=Summary + ) as stream: + result = await stream.get_final_result() + else: + with sdk.beta.agents.sessions.stream( + "session_test", input="Summarize the document.", tool_handlers={tool.name: tool}, output_type=Summary + ) as stream: + result = stream.get_final_result() + assert result.output_parsed == Summary(summary="default") + assert seen == ["default"] + assert server.inputs()[1]["output"] == text + assert server.inputs()[1]["success"] is True + + +@pytest.mark.parametrize("schema_kind", ["mapping", "root", "format"]) +@pytest.mark.parametrize("asynchronous", [False, True]) +async def test_schema_is_sent_and_api_rejection_is_preserved(schema_kind: str, asynchronous: bool) -> None: + if schema_kind == "root": + model = ( + pydantic.create_model("ListOutput", __root__=(list[str], ...)) + if PYDANTIC_V1 + else pydantic.RootModel[list[str]] + ) + elif schema_kind == "mapping": + model = pydantic.create_model("MappingOutput", value=(dict[str, str], ...)) + else: + model = LocalURLReport + expected_schema = _assert_matches_responses(model) + requests: list[httpx2.Request] = [] + error = {"message": "This schema is not supported", "type": "invalid_request_error", "code": "invalid_json_schema"} + + def handle(request: httpx2.Request) -> httpx2.Response: + requests.append(request) + return httpx2.Response(400, json={"error": error}) + + transport = httpx2.MockTransport(handle) + with pytest.raises(BadRequestError) as caught: + if asynchronous: + async with AsyncOpenAI( + api_key="synthetic", max_retries=0, http_client=httpx2.AsyncClient(transport=transport) + ) as client: + await client.beta.agents.sessions.create( + agent={"model": "test-model"}, + environment={"type": "none"}, + input="Summarize the document.", + stream=True, + output_type=model, + ) + else: + with OpenAI(api_key="synthetic", max_retries=0, http_client=httpx2.Client(transport=transport)) as client: + client.beta.agents.sessions.create( + agent={"model": "test-model"}, + environment={"type": "none"}, + input="Summarize the document.", + stream=True, + output_type=model, + ) + assert caught.value.status_code == 400 + assert caught.value.body == error + assert len(requests) == 1 + body = json.loads(requests[0].content) + assert body["agent"]["text"]["format"] == {"type": "json_schema", "schema": expected_schema} + assert "output_type" not in body + + +@pytest.mark.parametrize("model", [str, dict, object]) +async def test_invalid_output_model_rejected_before_request( + sdk: OpenAI | AsyncOpenAI, server: ResultServer, model: type[Any] +) -> None: + with pytest.raises(TypeError, match="Pydantic"): + if isinstance(sdk, AsyncOpenAI): + await sdk.beta.agents.sessions.create( + agent={"model": "test-model"}, + environment={"type": "none"}, + input="Question", + stream=True, + output_type=model, + ) + else: + sdk.beta.agents.sessions.create( + agent={"model": "test-model"}, + environment={"type": "none"}, + input="Question", + stream=True, + output_type=model, + ) + assert server.requests == [] + + +async def test_root_model_result_uses_native_parser(sdk: OpenAI | AsyncOpenAI, server: ResultServer) -> None: + model = ( + pydantic.create_model("ListOutput", __root__=(list[str], ...)) if PYDANTIC_V1 else pydantic.RootModel[list[str]] + ) + server.body = EventBody([turn_event("created"), message('["one"]'), turn_event("completed"), idle()]) + if isinstance(sdk, AsyncOpenAI): + async with sdk.beta.agents.sessions.stream("session_test", input="Question", output_type=model) as stream: + result = await stream.get_final_result() + else: + with sdk.beta.agents.sessions.stream("session_test", input="Question", output_type=model) as stream: + result = stream.get_final_result() + assert getattr(result.output_parsed, "__root__" if PYDANTIC_V1 else "root") == ["one"] + + +def test_inherited_schema_conversion_errors_are_preserved() -> None: + extra: dict[str, Any] = {"$ref": "https://example.com/schema.json"} + config = ( + type("Config", (), {"schema_extra": extra}) if PYDANTIC_V1 else pydantic.ConfigDict(json_schema_extra=extra) + ) + model = pydantic.create_model("ExternalReference", __config__=cast(Any, config), value=(str, ...)) + for convert in (agent_text_format, type_to_text_format_param): + with pytest.raises(ValueError, match="Unexpected.*ref format"): + convert(model) + + +@pytest.mark.parametrize("typed", [False, True], ids=["parse", "getter"]) +@pytest.mark.parametrize("case", ["messages", "parts", "invalid_later", "invalid_first", "empty", "split_json"]) +async def test_final_text_parts_are_parsed_independently( + sdk: OpenAI | AsyncOpenAI, server: ResultServer, typed: bool, case: str +) -> None: + first = '{"summary":"first","findings":["one"]}' + second = '{"summary":"second","findings":["two"]}' + if case == "messages": + parts = [[first], [second]] + elif case == "parts": + parts = [[first, second]] + elif case == "invalid_later": + parts = [[first], ["invalid JSON"]] + elif case == "invalid_first": + parts = [["invalid JSON"], [second]] + elif case == "split_json": + parts = [[first[:10], first[10:]]] + else: + parts = [] + messages: list[dict[str, Any]] = [] + for index, texts in enumerate(parts): + event = message(item_id=f"message_{index}", index=index) + event["item"]["content"] = [{"type": "output_text", "text": text, "annotations": []} for text in texts] + messages.append(event) + server.body = EventBody([turn_event("created"), *messages, turn_event("completed"), idle()]) + should_fail = case not in ("messages", "parts") + try: + if isinstance(sdk, AsyncOpenAI): + async with sdk.beta.agents.sessions.stream( + "session_test", input="Question", output_type=Report if typed else None + ) as stream: + result = await stream.get_final_result() + else: + with sdk.beta.agents.sessions.stream( + "session_test", input="Question", output_type=Report if typed else None + ) as stream: + result = stream.get_final_result() + if not typed: + result = result.parse(Report) + except AgentOutputParseError as exc: + assert should_fail + raw = exc.result + assert raw.output_parsed is None + else: + assert not should_fail + assert result.output_parsed == Report(summary="first", findings=["one"]) + raw = result + assert raw.output_text == "".join(text for texts in parts for text in texts) + assert [ + [content.text for content in item.content if content.type == "output_text"] for item in raw.messages + ] == parts