diff --git a/helpers.md b/helpers.md index c6df1faf42..0ac348cc25 100644 --- a/helpers.md +++ b/helpers.md @@ -535,3 +535,37 @@ client.vector_stores.file_batches.create_and_poll(...) client.vector_stores.file_batches.upload_and_poll(...) client.videos.create_and_poll(...) ``` + +# Beta Agents turn results + +Both streamed session creation and the one-turn session helper can collect the final +answer. Call `get_final_result()` directly, or enable `with_result_collection()` +before iterating to display progress. Existing follow-up tool handlers continue to run while it drains. + +```python +with client.beta.agents.sessions.create( + agent={"model": MODEL}, + environment={"type": "none"}, + input="Explain this policy.", + stream=True, +) as stream: + result = stream.get_final_result() + +print(result.output_text) + +with client.beta.agents.sessions.stream( + result.session_id, + input="Give me an example.", + tool_handlers=handlers, +).with_result_collection() as stream: + for event in stream: + show_progress(event) + followup = stream.get_final_result() + +print(followup.output_text) +``` + +With `AsyncOpenAI`, await creation, use `async with` / `async for`, and await +`get_final_result()`. The result exposes `output_text`, `turn`, final `messages`, +`session_id`, and `turn_id`. Collection raises `AgentTurnResultError` when a complete +successful answer cannot be established. diff --git a/src/openai/lib/beta/__init__.py b/src/openai/lib/beta/__init__.py new file mode 100644 index 0000000000..dff30161f2 --- /dev/null +++ b/src/openai/lib/beta/__init__.py @@ -0,0 +1 @@ +"""Beta SDK runtime helpers.""" diff --git a/src/openai/lib/beta/agents/__init__.py b/src/openai/lib/beta/agents/__init__.py new file mode 100644 index 0000000000..6d543dda57 --- /dev/null +++ b/src/openai/lib/beta/agents/__init__.py @@ -0,0 +1,7 @@ +"""Beta helpers for collecting hosted Agents API turn results.""" + +from ._result import AgentTurnResult as AgentTurnResult, AgentTurnResultError as AgentTurnResultError +from ._stream import ( + AgentSessionEventStream as AgentSessionEventStream, + AsyncAgentSessionEventStream as AsyncAgentSessionEventStream, +) diff --git a/src/openai/lib/beta/agents/_result.py b/src/openai/lib/beta/agents/_result.py new file mode 100644 index 0000000000..d42955ae2f --- /dev/null +++ b/src/openai/lib/beta/agents/_result.py @@ -0,0 +1,187 @@ +from __future__ import annotations + +from copy import deepcopy +from typing import Iterable +from dataclasses import dataclass +from typing_extensions import Literal + +from ...._exceptions import OpenAIError +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 + + +@dataclass(frozen=True) +class AgentTurnResult: + """Beta: the completed final assistant messages from one successful root turn.""" + + turn: Turn + messages: list[AgentSessionMessage] + + @property + def session_id(self) -> str: + return self.turn.session_id + + @property + def turn_id(self) -> str: + return self.turn.id + + @property + def output_text(self) -> str: + """Join final output text without adding separators or performing I/O.""" + return "".join(message.output_text for message in self.messages) + + +ResultErrorReason = Literal["failed", "cancelled", "requires_action", "incomplete", "observation_failed"] + + +class AgentTurnResultError(OpenAIError): + """Beta: collection did not establish a successful, complete turn result. + + ``turn`` and ``messages`` preserve available partial state. An observation + error does not mean the hosted turn failed. Transport errors are chained. + """ + + def __init__( + self, + reason: ResultErrorReason, + *, + turn: Turn | None, + session_id: str | None, + messages: list[AgentSessionMessage], + required_actions: list[RequiredAction], + ) -> None: + super().__init__(f"Could not collect the agent turn result: {reason}") + self.reason = reason + self.turn = turn + self.session_id = session_id + self.turn_id = turn.id if turn is not None else None + self.messages = messages + self.required_actions = required_actions + + +class AgentTurnResultCollector: + def __init__(self, session_id: str | None = None) -> None: + self.session_id = session_id + self.turn: Turn | None = None + self.boundary = False + self.session_failed = False + self.cause: Exception | None = None + self.required_actions: list[RequiredAction] = [] + self._messages: dict[str, tuple[int, AgentSessionMessage]] = {} + self._error: AgentTurnResultError | None = None + self._result: AgentTurnResult | None = None + + def accept(self, event: AgentSessionEvent) -> None: + if self.boundary: + return + if event.type == "agent.session.created": + self.session_id = event.session.id + elif event.type == "agent.session.turn.created": + if self.turn is None and event.turn.subagent_id is None: + self.turn = deepcopy(event.turn) + self.session_id = event.session_id + elif ( + event.type == "agent.session.turn.completed" + or event.type == "agent.session.turn.failed" + or event.type == "agent.session.turn.cancelled" + ): + if self.turn is not None and event.turn_id == self.turn.id: + self.turn = deepcopy(event.turn) + self.required_actions = [] + elif event.type == "agent.session.turn.item.done": + # Completed items are authoritative on these ordered, uninterrupted streams. + item = event.item + if ( + self.turn is not None + and item.turn_id == self.turn.id + and item.type == "message" + and item.phase != "commentary" + and item.status == "completed" + ): + message = AgentSessionMessage.construct(_fields_set=None, **deepcopy(item.to_dict())) + self._messages[item.id] = (event.output_index, message) + elif event.type == "agent.session.in_progress": + self.required_actions = [] + elif event.type == "agent.session.requires_action": + self.session_id = event.session.id + self.required_actions = deepcopy(event.session.required_actions) + elif event.type == "agent.session.failed": + self.session_id = event.session.id + self.session_failed = True + self.boundary = True + elif event.type == "agent.session.idle": + if self.turn is not None and self.turn.status in ("completed", "failed", "cancelled"): + self.required_actions = [] + self.boundary = True + + def is_done(self) -> bool: + return self.boundary + + def messages(self) -> list[AgentSessionMessage]: + return [message for _, message in sorted(self._messages.values(), key=lambda pair: pair[0])] + + def error(self, reason: ResultErrorReason) -> AgentTurnResultError: + if self._error is None: + self._error = AgentTurnResultError( + reason, + turn=self.turn, + session_id=self.session_id, + messages=self.messages(), + required_actions=self.required_actions, + ) + self.turn = None + self.required_actions = [] + self._messages.clear() + self.boundary = True + return self._error + + def check_outcome(self, handled_tools: Iterable[str] = ()) -> None: + if self._error is not None: + raise self._error + if self.session_failed or self.turn is not None and self.turn.status == "failed": + raise self.error("failed") + if self.turn is not None and self.turn.status == "cancelled": + raise self.error("cancelled") + if any(action.type != "function_call" or action.name not in handled_tools for action in self.required_actions): + raise self.error("requires_action") + if self.cause is not None: + raise self.error("observation_failed") from self.cause + + def result(self) -> AgentTurnResult: + if self._result is not None: + return self._result + self.check_outcome() + if not self.boundary or self.turn is None or self.turn.status != "completed": + raise self.error("incomplete") + self._result = AgentTurnResult(turn=deepcopy(self.turn), messages=self.messages()) + self._messages.clear() + return self._result + + +class AgentTurnResultCollection: + """Keep ordinary event iteration incremental until collection is requested.""" + + def __init__(self, session_id: str | None = None) -> None: + self.collector: AgentTurnResultCollector | None = None + self._session_id = session_id + self._started = False + + def enable(self) -> AgentTurnResultCollector: + if self.collector is None: + if self._started: + raise RuntimeError( + "Call with_result_collection() before consuming events, or call get_final_result() on a fresh stream" + ) + self.collector = AgentTurnResultCollector(self._session_id) + return self.collector + + def accept(self, event: AgentSessionEvent) -> None: + self._started = True + if self.collector is not None: + self.collector.accept(event) + + def record_error(self, error: Exception) -> None: + if self.collector is not None and not self.collector.is_done(): + self.collector.cause = error diff --git a/src/openai/lib/beta/agents/_stream.py b/src/openai/lib/beta/agents/_stream.py new file mode 100644 index 0000000000..10d5c3030b --- /dev/null +++ b/src/openai/lib/beta/agents/_stream.py @@ -0,0 +1,108 @@ +from __future__ import annotations + +from typing import Iterator, AsyncIterator +from typing_extensions import Self, override + +from ._result import AgentTurnResult, AgentTurnResultError, AgentTurnResultCollection +from ...._compat import cached_property +from ...._streaming import Stream, AsyncStream +from ....types.beta.agent_session_event import AgentSessionEvent + + +class AgentSessionEventStream(Stream[AgentSessionEvent]): + """Beta: a creation event stream that can collect its initial turn's result. + + Event iteration and response access behave like ``Stream``. Collection does + not execute local tools; required actions are available on the result error. + Supply initial input when creating a session to collect its turn result. + """ + + @cached_property + def _collection(self) -> AgentTurnResultCollection: + return AgentTurnResultCollection() + + @override + def __stream__(self) -> Iterator[AgentSessionEvent]: + try: + for event in super().__stream__(): + self._collection.accept(event) + yield event + except Exception as error: + self._collection.record_error(error) + raise + + def with_result_collection(self) -> Self: + """Beta: retain final messages for a result; call before consuming events.""" + self._collection.enable() + return self + + def get_final_result(self) -> AgentTurnResult: + """Consume remaining events and return the successful initial turn result. + + Raises AgentTurnResultError for an unsuccessful or unobserved outcome. + Collection starts automatically on a fresh stream. To iterate first, call + with_result_collection() before consuming events. Repeated calls return + the cached result without consuming more events. + """ + collector = self._collection.enable() + try: + collector.check_outcome() + if not collector.is_done(): + for _ in self: + collector.check_outcome() + if collector.is_done(): + break + return collector.result() + except AgentTurnResultError: + raise + except Exception as error: + collector.cause = error + raise collector.error("observation_failed") from error + finally: + self.close() + + +class AsyncAgentSessionEventStream(AsyncStream[AgentSessionEvent]): + """Beta: asynchronous counterpart of AgentSessionEventStream.""" + + @cached_property + def _collection(self) -> AgentTurnResultCollection: + return AgentTurnResultCollection() + + @override + async def __stream__(self) -> AsyncIterator[AgentSessionEvent]: + try: + async for event in super().__stream__(): + self._collection.accept(event) + yield event + except Exception as error: + self._collection.record_error(error) + raise + + def with_result_collection(self) -> Self: + """Beta: retain final messages for a result; call before consuming events.""" + self._collection.enable() + return self + + async def get_final_result(self) -> AgentTurnResult: + """Consume remaining events and return the successful initial turn result. + + Enables collection on a fresh stream. To iterate first, call + with_result_collection() before consuming events. + """ + collector = self._collection.enable() + try: + collector.check_outcome() + if not collector.is_done(): + async for _ in self: + collector.check_outcome() + if collector.is_done(): + break + return collector.result() + except AgentTurnResultError: + raise + except Exception as error: + collector.cause = error + raise collector.error("observation_failed") from error + finally: + await self.close() diff --git a/src/openai/lib/streaming/agents/_streams.py b/src/openai/lib/streaming/agents/_streams.py index 9edb22a9da..d71918ef0a 100644 --- a/src/openai/lib/streaming/agents/_streams.py +++ b/src/openai/lib/streaming/agents/_streams.py @@ -15,6 +15,7 @@ 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 ....types.beta.agent_session_event import AgentSessionEvent from ....types.beta.agent_function_call_item import AgentFunctionCallItem from ....types.beta.agent_session_input_param import ( @@ -144,6 +145,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._stream: Stream[AgentSessionEvent] | None = None self._iterator: Iterator[AgentSessionEvent] | None = None self._entered = False @@ -189,6 +191,36 @@ def until_done(self) -> None: for _ in self: pass + def with_result_collection(self) -> Self: + """Beta: retain final messages for a result; call before consuming events.""" + self._collection.enable() + return self + + def get_final_result(self) -> AgentTurnResult: + """Beta: drain this turn, executing handlers, and collect its final answer. + + Raises AgentTurnResultError when a successful complete result cannot be + established. Draining executes registered handlers and submits their + outputs. Result properties themselves do not execute handlers. + To iterate first, call with_result_collection() before consuming events. + """ + collector = self._collection.enable() + try: + collector.check_outcome(self._handlers) + if not collector.is_done(): + for _ in self: + collector.check_outcome(self._handlers) + if collector.is_done(): + break + return collector.result() + except AgentTurnResultError: + raise + except Exception as error: + self._collection.record_error(error) + raise collector.error("observation_failed") from error + finally: + self.close() + def close(self) -> None: """Close the event connection without cancelling the backend turn.""" self._closed = True @@ -201,6 +233,7 @@ def _iterate(self) -> Iterator[AgentSessionEvent]: for event in self._stream: if not self._state.accept(event): continue + self._collection.accept(event) terminal = self._state.terminal(event) if terminal: self.close() @@ -219,6 +252,9 @@ def _iterate(self) -> Iterator[AgentSessionEvent]: result = failed_event(call) self._submit_result(result) raise RuntimeError("Session event stream ended before the turn reached idle or failed") + except Exception as error: + self._collection.record_error(error) + raise finally: self.close() @@ -265,6 +301,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._stream: AsyncStream[AgentSessionEvent] | None = None self._iterator: AsyncIterator[AgentSessionEvent] | None = None self._entered = False @@ -310,6 +347,36 @@ async def until_done(self) -> None: async for _ in self: pass + def with_result_collection(self) -> Self: + """Beta: retain final messages for a result; call before consuming events.""" + self._collection.enable() + return self + + async def get_final_result(self) -> AgentTurnResult: + """Beta: drain this turn, executing handlers, and collect its final answer. + + Raises AgentTurnResultError when a successful complete result cannot be + established. Draining executes registered handlers and submits their + outputs. Result properties themselves do not execute handlers. + To iterate first, call with_result_collection() before consuming events. + """ + collector = self._collection.enable() + try: + collector.check_outcome(self._handlers) + if not collector.is_done(): + async for _ in self: + collector.check_outcome(self._handlers) + if collector.is_done(): + break + return collector.result() + except AgentTurnResultError: + raise + except Exception as error: + self._collection.record_error(error) + raise collector.error("observation_failed") from error + finally: + await self.close() + async def close(self) -> None: """Close the event connection without cancelling the backend turn.""" self._closed = True @@ -322,6 +389,7 @@ async def _iterate(self) -> AsyncIterator[AgentSessionEvent]: async for event in self._stream: if not self._state.accept(event): continue + self._collection.accept(event) terminal = self._state.terminal(event) if terminal: await self.close() @@ -343,6 +411,9 @@ async def _iterate(self) -> AsyncIterator[AgentSessionEvent]: result = failed_event(call) await self._submit_result(result) raise RuntimeError("Session event stream ended before the turn reached idle or failed") + except Exception as error: + self._collection.record_error(error) + raise finally: await self.close() diff --git a/src/openai/resources/beta/agents/sessions/sessions.py b/src/openai/resources/beta/agents/sessions/sessions.py index 1a596a4fe8..c3b4f4dbb3 100644 --- a/src/openai/resources/beta/agents/sessions/sessions.py +++ b/src/openai/resources/beta/agents/sessions/sessions.py @@ -53,9 +53,9 @@ from ....._compat import cached_property from ....._resource import SyncAPIResource, AsyncAPIResource from ....._response import to_streamed_response_wrapper, async_to_streamed_response_wrapper -from ....._streaming import Stream, AsyncStream from .....pagination import SyncCursorPage, AsyncCursorPage from ....._base_client import AsyncPaginator, make_request_options +from .....lib.beta.agents import AgentSessionEventStream, AsyncAgentSessionEventStream from .subagents.subagents import ( Subagents, AsyncSubagents, @@ -68,7 +68,6 @@ from .....lib.streaming.agents import ToolHandler, AsyncToolHandler, AgentSessionStream, AsyncAgentSessionStream from .....types.beta.agent_session import AgentSession from .....types.beta.environment_param import EnvironmentParam -from .....types.beta.agent_session_event import AgentSessionEvent from .....types.beta.agent_session_deleted import AgentSessionDeleted from .....types.beta.agent_session_input_message_param import AgentSessionInputMessageParam @@ -216,7 +215,7 @@ def create( extra_query: Query | None = None, extra_body: Body | None = None, timeout: float | httpx2.Timeout | None | NotGiven = not_given, - ) -> Stream[AgentSessionEvent]: + ) -> AgentSessionEventStream: """ Creates a managed agent session, optionally submits initial input, and returns the session or streams its events when stream is true. See @@ -270,7 +269,7 @@ def create( extra_query: Query | None = None, extra_body: Body | None = None, timeout: float | httpx2.Timeout | None | NotGiven = not_given, - ) -> AgentSession | Stream[AgentSessionEvent]: + ) -> AgentSession | AgentSessionEventStream: """ Creates a managed agent session, optionally submits initial input, and returns the session or streams its events when stream is true. See @@ -324,7 +323,7 @@ def create( extra_query: Query | None = None, extra_body: Body | None = None, timeout: float | httpx2.Timeout | None | NotGiven = not_given, - ) -> AgentSession | Stream[AgentSessionEvent]: + ) -> AgentSession | AgentSessionEventStream: extra_headers = {"OpenAI-Beta": "agents=v1", **(extra_headers or {})} return self._post( "/agents/sessions", @@ -351,7 +350,7 @@ def create( ), cast_to=AgentSession, stream=stream or False, - stream_cls=Stream[AgentSessionEvent], + stream_cls=AgentSessionEventStream, ) def retrieve( @@ -698,7 +697,7 @@ async def create( extra_query: Query | None = None, extra_body: Body | None = None, timeout: float | httpx2.Timeout | None | NotGiven = not_given, - ) -> AsyncStream[AgentSessionEvent]: + ) -> AsyncAgentSessionEventStream: """ Creates a managed agent session, optionally submits initial input, and returns the session or streams its events when stream is true. See @@ -752,7 +751,7 @@ async def create( extra_query: Query | None = None, extra_body: Body | None = None, timeout: float | httpx2.Timeout | None | NotGiven = not_given, - ) -> AgentSession | AsyncStream[AgentSessionEvent]: + ) -> AgentSession | AsyncAgentSessionEventStream: """ Creates a managed agent session, optionally submits initial input, and returns the session or streams its events when stream is true. See @@ -806,7 +805,7 @@ async def create( extra_query: Query | None = None, extra_body: Body | None = None, timeout: float | httpx2.Timeout | None | NotGiven = not_given, - ) -> AgentSession | AsyncStream[AgentSessionEvent]: + ) -> AgentSession | AsyncAgentSessionEventStream: extra_headers = {"OpenAI-Beta": "agents=v1", **(extra_headers or {})} return await self._post( "/agents/sessions", @@ -833,7 +832,7 @@ async def create( ), cast_to=AgentSession, stream=stream or False, - stream_cls=AsyncStream[AgentSessionEvent], + stream_cls=AsyncAgentSessionEventStream, ) async def retrieve( diff --git a/tests/lib/streaming/agents/test_results.py b/tests/lib/streaming/agents/test_results.py new file mode 100644 index 0000000000..958fc3b030 --- /dev/null +++ b/tests/lib/streaming/agents/test_results.py @@ -0,0 +1,523 @@ +from __future__ import annotations + +from typing import Any, Iterator, cast +from typing_extensions import override + +import httpx2 +import pytest + +from openai import OpenAI, AsyncOpenAI +from openai.lib.beta.agents import AgentTurnResult, AgentTurnResultError +from openai.lib.streaming.agents import AgentSessionStream, AsyncAgentSessionStream +from tests.lib.streaming.agents.test_streams import Server, EventBody, sdk as sdk, call, idle, session, turn_event + + +class ResultServer(Server): + @override + def handle(self, request: httpx2.Request) -> httpx2.Response: + if request.method == "POST" and request.url.path.endswith("/sessions"): + self.requests.append(request) + return httpx2.Response(200, headers={"content-type": "text/event-stream"}, stream=self.body) + return super().handle(request) + + +@pytest.fixture +def server() -> ResultServer: + return ResultServer() + + +@pytest.fixture(params=[False, True], ids=["followup", "creation"]) +def creation(request: pytest.FixtureRequest) -> bool: + return bool(request.param) + + +def message( + text: str = "answer", + *, + item_id: str = "message_a", + index: int = 0, + turn_id: str = "turn_root", + phase: str | None = "final_answer", + kind: str = "done", +) -> dict[str, Any]: + return { + "type": f"agent.session.turn.item.{kind}", + "session_id": "session_test", + "turn_id": turn_id, + "output_index": index, + "item": { + "id": item_id, + "turn_id": turn_id, + "type": "message", + "role": "assistant", + "phase": phase, + "status": "completed" if kind == "done" else "in_progress", + "content": [{"type": "output_text", "text": text, "annotations": []}], + }, + } + + +async def collect(sdk: OpenAI | AsyncOpenAI, creation: bool, *, mode: str = "getter", **kwargs: Any) -> AgentTurnResult: + if isinstance(sdk, AsyncOpenAI): + async_stream = ( + ( + await sdk.beta.agents.sessions.create( + agent={"model": "test-model"}, environment={"type": "none"}, input="Question", stream=True + ) + ) + if creation + else sdk.beta.agents.sessions.stream("session_test", input="Question", **kwargs) + ) + async with async_stream: + if mode != "getter": + async_stream.with_result_collection() + if mode == "iterate": + async for _ in async_stream: + pass + elif mode == "resume": + async for event in async_stream: + if event.type in ("agent.session.in_progress", "agent.session.turn.completed"): + break + elif mode == "drain": + assert isinstance(async_stream, AsyncAgentSessionStream) + await async_stream.until_done() + result = await async_stream.get_final_result() + assert await async_stream.get_final_result() is result + return result + stream = ( + sdk.beta.agents.sessions.create( + agent={"model": "test-model"}, environment={"type": "none"}, input="Question", stream=True + ) + if creation + else sdk.beta.agents.sessions.stream("session_test", input="Question", **kwargs) + ) + with stream: + if mode != "getter": + stream.with_result_collection() + if mode == "iterate": + for _ in stream: + pass + elif mode == "resume": + for event in stream: + if event.type in ("agent.session.in_progress", "agent.session.turn.completed"): + break + elif mode == "drain": + assert isinstance(stream, AgentSessionStream) + stream.until_done() + result = stream.get_final_result() + assert stream.get_final_result() is result + return result + + +@pytest.mark.parametrize("mode", ["getter", "iterate", "drain"]) +async def test_result_selects_ordered_final_messages( + sdk: OpenAI | AsyncOpenAI, server: ResultServer, creation: bool, mode: str +) -> None: + if creation and mode == "drain": + pytest.skip("Only the existing follow-up helper exposes until_done") + server.body = EventBody( + [ + idle(), + turn_event("created"), + message("unfinished", kind="added"), + message("ignored commentary", item_id="commentary", phase="commentary"), + turn_event("created", "child", subagent_id="child_agent"), + message("ignored child", item_id="child_answer", turn_id="child"), + turn_event("completed", "child", subagent_id="child_agent"), + message("second", item_id="message_b", index=2), + message("first", index=1), + message("first", index=1), + turn_event("completed"), + idle(), + ] + ) + result = await collect(sdk, creation, mode=mode) + assert result.session_id == "session_test" + assert result.turn_id == "turn_root" + assert result.turn.status == "completed" + assert result.output_text == "firstsecond" + assert [m.id for m in result.messages] == ["message_a", "message_b"] + assert result.messages[0].content[0].type == "output_text" + assert server.body.closed + + +async def test_empty_success(sdk: OpenAI | AsyncOpenAI, creation: bool) -> None: + result = await collect(sdk, creation) + assert result.output_text == "" + assert result.messages == [] + + +@pytest.mark.parametrize("status", ["failed", "cancelled"]) +async def test_unsuccessful_turn(sdk: OpenAI | AsyncOpenAI, server: ResultServer, creation: bool, status: str) -> None: + server.body = EventBody([turn_event("created"), message(), turn_event(status), idle()]) + with pytest.raises(AgentTurnResultError) as exc: + await collect(sdk, creation) + assert exc.value.reason == status + assert exc.value.turn_id == "turn_root" + assert exc.value.messages[0].output_text == "answer" + assert server.body.closed + + +@pytest.mark.parametrize("tail", [[], [turn_event("completed")]]) +async def test_incomplete_answer( + sdk: OpenAI | AsyncOpenAI, server: ResultServer, creation: bool, tail: list[dict[str, Any]] +) -> None: + server.body = EventBody([turn_event("created"), *tail]) + with pytest.raises(AgentTurnResultError) as exc: + await collect(sdk, creation) + assert exc.value.reason in ("incomplete", "observation_failed") + assert exc.value.turn_id == "turn_root" + + +async def test_legacy_null_phase_is_collected(sdk: OpenAI | AsyncOpenAI, server: ResultServer, creation: bool) -> None: + server.body = EventBody([turn_event("created"), message(phase=None), turn_event("completed"), idle()]) + result = await collect(sdk, creation) + assert result.output_text == "answer" + assert result.messages[0].phase is None + + +async def test_unhandled_action(sdk: OpenAI | AsyncOpenAI, server: ResultServer, creation: bool) -> None: + state = session("requires_action") + state["required_actions"] = [ + { + "type": "function_call", + "turn_id": "turn_root", + "name": "search", + "call_id": "call_test", + "arguments": {"query": "test"}, + } + ] + server.body = EventBody( + [turn_event("created"), call(), {"type": "agent.session.requires_action", "session": state}] + ) + with pytest.raises(AgentTurnResultError) as exc: + await collect(sdk, creation) + assert exc.value.reason == "requires_action" + assert exc.value.required_actions[0].type == "function_call" + assert server.body.closed + + +async def test_getter_dispatches_existing_handler_once(sdk: OpenAI | AsyncOpenAI, server: ResultServer) -> None: + seen: list[object] = [] + + def handler(arguments: object) -> str: + seen.append(arguments) + return "found" + + server.body = EventBody([turn_event("created"), call(), message(), turn_event("completed"), idle()]) + result = await collect(sdk, False, tool_handlers={"search": handler}) + assert result.output_text == "answer" + assert seen == [{"query": "test"}] + assert len(server.inputs()) == 2 + + +async def test_explicit_close_is_not_success(sdk: OpenAI | AsyncOpenAI, creation: bool) -> None: + if isinstance(sdk, AsyncOpenAI): + async_stream = ( + await sdk.beta.agents.sessions.create( + agent={"model": "test-model"}, environment={"type": "none"}, input="Question", stream=True + ) + if creation + else sdk.beta.agents.sessions.stream("session_test", input="Question") + ) + async with async_stream: + await async_stream.close() + with pytest.raises(AgentTurnResultError): + await async_stream.get_final_result() + else: + stream = ( + sdk.beta.agents.sessions.create( + agent={"model": "test-model"}, environment={"type": "none"}, input="Question", stream=True + ) + if creation + else sdk.beta.agents.sessions.stream("session_test", input="Question") + ) + with stream: + stream.close() + with pytest.raises(AgentTurnResultError): + stream.get_final_result() + + +class BrokenBody(EventBody): + @override + def __iter__(self) -> Iterator[bytes]: + yield from super().__iter__() + raise httpx2.ReadError("synthetic interrupted connection") + + +async def test_transport_error_retains_partial_answer( + sdk: OpenAI | AsyncOpenAI, server: ResultServer, creation: bool +) -> None: + server.body = BrokenBody([turn_event("created"), message()]) + with pytest.raises(AgentTurnResultError) as exc: + await collect(sdk, creation) + assert exc.value.reason == "observation_failed" + assert exc.value.__cause__ is not None + assert exc.value.turn is not None and exc.value.turn.status == "in_progress" + assert exc.value.messages[0].output_text == "answer" + + +async def test_result_stops_at_first_turn_boundary( + sdk: OpenAI | AsyncOpenAI, server: ResultServer, creation: bool +) -> None: + server.body = EventBody( + [ + turn_event("created"), + message(), + turn_event("completed"), + idle(), + turn_event("created", "later"), + message("later answer", turn_id="later"), + ] + ) + result = await collect(sdk, creation) + assert result.output_text == "answer" + assert server.body.read_count == 4 + + +async def test_late_added_snapshot_does_not_reopen_completed_message( + sdk: OpenAI | AsyncOpenAI, server: ResultServer, creation: bool +) -> None: + server.body = EventBody( + [turn_event("created"), message(), message("stale", kind="added", phase=None), turn_event("completed"), idle()] + ) + result = await collect(sdk, creation) + assert result.output_text == "answer" + + +async def test_later_stream_error_does_not_poison_completed_result( + sdk: OpenAI | AsyncOpenAI, server: ResultServer +) -> None: + from openai import APIConnectionError + + server.body = BrokenBody( + [ + turn_event("created"), + message(), + turn_event("completed"), + idle(), + turn_event("created", "later"), + turn_event("failed", "later"), + ] + ) + if isinstance(sdk, AsyncOpenAI): + async with await sdk.beta.agents.sessions.create( + agent={"model": "test-model"}, environment={"type": "none"}, input="Question", stream=True + ) as async_stream: + async_stream.with_result_collection() + with pytest.raises(APIConnectionError): + async for _ in async_stream: + pass + result = await async_stream.get_final_result() + else: + with sdk.beta.agents.sessions.create( + agent={"model": "test-model"}, environment={"type": "none"}, input="Question", stream=True + ) as stream: + stream.with_result_collection() + with pytest.raises(APIConnectionError): + for _ in stream: + pass + result = stream.get_final_result() + assert result.turn_id == "turn_root" + assert result.output_text == "answer" + + +@pytest.mark.parametrize("initial_phase", [None, "final_answer"]) +async def test_completed_commentary_is_excluded( + sdk: OpenAI | AsyncOpenAI, server: ResultServer, creation: bool, initial_phase: str | None +) -> None: + server.body = EventBody( + [ + turn_event("created"), + message(kind="added", phase=initial_phase), + message(phase="commentary"), + turn_event("completed"), + idle(), + ] + ) + result = await collect(sdk, creation) + assert result.messages == [] + assert result.output_text == "" + + +def test_collector_transfers_messages_into_cached_result() -> None: + from openai._models import construct_type + from openai.lib.beta.agents._result import AgentTurnResultCollector + from openai.types.beta.agent_session_event import AgentSessionEvent + + collector = AgentTurnResultCollector() + for index, data in enumerate([turn_event("created"), message(), turn_event("completed"), idle()]): + collector.accept( + cast(AgentSessionEvent, construct_type(type_=AgentSessionEvent, value={"event_id": str(index), **data})) + ) + original_message = collector.messages()[0] + result = collector.result() + assert result.messages[0] is original_message + assert collector.messages() == [] # The stream no longer retains a duplicate message collection. + assert collector.result() is result + assert result.output_text == "answer" + + +@pytest.mark.parametrize( + "phase,turn_id,expected", + [("final_answer", "turn_root", "answer"), ("commentary", "turn_root", ""), ("final_answer", "child", "")], +) +async def test_only_completed_message_snapshots_supply_text( + sdk: OpenAI | AsyncOpenAI, server: ResultServer, creation: bool, phase: str, turn_id: str, expected: str +) -> None: + text_event: dict[str, object] = { + "type": "agent.session.turn.output_text.delta", + "session_id": "session_test", + "turn_id": None, + "item_id": "message_a", + "output_index": 0, + "content_index": 0, + "delta": "answer", + } + server.body = EventBody( + [ + turn_event("created"), + text_event, + message(phase=phase, turn_id=turn_id), + {**text_event, "event_id": "late-text"}, + turn_event("completed"), + idle(), + ] + ) + result = await collect(sdk, creation) + assert result.output_text == expected + + +def test_error_transfers_partial_payloads_and_is_cached() -> None: + from openai._models import construct_type + from openai.lib.beta.agents._result import AgentTurnResultCollector + from openai.types.beta.agent_session_event import AgentSessionEvent + + state = session("requires_action") + state["required_actions"] = [ + { + "type": "function_call", + "turn_id": "turn_root", + "name": "search", + "call_id": "call_test", + "arguments": {"query": "test"}, + } + ] + collector = AgentTurnResultCollector() + for index, data in enumerate( + [turn_event("created"), message(), {"type": "agent.session.requires_action", "session": state}] + ): + collector.accept( + cast(AgentSessionEvent, construct_type(type_=AgentSessionEvent, value={"event_id": str(index), **data})) + ) + original_message = collector.messages()[0] + original_actions = collector.required_actions + original_turn = collector.turn + with pytest.raises(AgentTurnResultError) as exc: + collector.check_outcome() + error = exc.value + assert error.messages[0] is original_message + assert error.required_actions is original_actions + assert error.turn is original_turn + assert collector.turn is None and collector.messages() == [] and collector.required_actions == [] + with pytest.raises(AgentTurnResultError) as repeated: + collector.result() + assert repeated.value is error + + +@pytest.mark.parametrize("resumed", [True, False]) +async def test_resolved_action_is_not_reported_again( + sdk: OpenAI | AsyncOpenAI, server: ResultServer, creation: bool, resumed: bool +) -> None: + state = session("requires_action") + state["required_actions"] = [ + { + "type": "function_call", + "turn_id": "turn_root", + "name": "search", + "call_id": "call_test", + "arguments": {"query": "test"}, + } + ] + events: list[dict[str, object]] = [ + turn_event("created"), + {"type": "agent.session.requires_action", "session": state}, + ] + if resumed: + events.append({"type": "agent.session.in_progress", "session": session("in_progress")}) + events.extend([message(), turn_event("completed"), idle()]) + server.body = EventBody(events) + result = await collect(sdk, creation, mode="resume") + assert result.output_text == "answer" + + +async def test_raw_iteration_does_not_collect_output( + sdk: OpenAI | AsyncOpenAI, server: ResultServer, creation: bool +) -> None: + events = [message("x" * 8192, item_id=f"message_{index}", index=index) for index in range(1024)] + server.body = EventBody([turn_event("created"), *events, turn_event("completed"), idle()]) + if isinstance(sdk, AsyncOpenAI): + async_stream = ( + await sdk.beta.agents.sessions.create( + agent={"model": "test-model"}, environment={"type": "none"}, input="Question", stream=True + ) + if creation + else sdk.beta.agents.sessions.stream("session_test", input="Question") + ) + async with async_stream: + async for _ in async_stream: + assert async_stream._collection.collector is None + read_count = server.body.read_count + with pytest.raises(RuntimeError, match="with_result_collection"): + await async_stream.get_final_result() + assert server.body.read_count == read_count + else: + stream = ( + sdk.beta.agents.sessions.create( + agent={"model": "test-model"}, environment={"type": "none"}, input="Question", stream=True + ) + if creation + else sdk.beta.agents.sessions.stream("session_test", input="Question") + ) + with stream: + for _ in stream: + assert stream._collection.collector is None + read_count = server.body.read_count + with pytest.raises(RuntimeError, match="with_result_collection"): + stream.get_final_result() + assert server.body.read_count == read_count + + +async def test_late_collection_opt_in_does_not_consume_events( + sdk: OpenAI | AsyncOpenAI, server: ResultServer, creation: bool +) -> None: + if isinstance(sdk, AsyncOpenAI): + async_stream = ( + await sdk.beta.agents.sessions.create( + agent={"model": "test-model"}, environment={"type": "none"}, input="Question", stream=True + ) + if creation + else sdk.beta.agents.sessions.stream("session_test", input="Question") + ) + async with async_stream: + await async_stream.__anext__() + with pytest.raises(RuntimeError, match="before consuming events"): + async_stream.with_result_collection() + with pytest.raises(RuntimeError, match="before consuming events"): + await async_stream.get_final_result() + assert server.body.read_count == 1 + else: + stream = ( + sdk.beta.agents.sessions.create( + agent={"model": "test-model"}, environment={"type": "none"}, input="Question", stream=True + ) + if creation + else sdk.beta.agents.sessions.stream("session_test", input="Question") + ) + with stream: + next(stream) + with pytest.raises(RuntimeError, match="before consuming events"): + stream.with_result_collection() + with pytest.raises(RuntimeError, match="before consuming events"): + stream.get_final_result() + assert server.body.read_count == 1