-
Notifications
You must be signed in to change notification settings - Fork 7k
feat: collect final output from beta Agents streams #4004
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We鈥檒l occasionally send you account related emails.
Already on GitHub? Sign in to your account
Merged
Merged
Changes from all commits
Commits
Show all changes
7 commits
Select commit
Hold shift + click to select a range
cf810d8
feat: collect final output from beta Agents streams
apcha-oai e541046
fix(agents): resolve commentary snapshots and release collected output
apcha-oai 575fe2b
fix(agents): retain incomplete output evidence without duplicate errors
apcha-oai 8d149ea
fix(agents): clear resolved required actions before collection
apcha-oai 70c4b1c
refactor(agents): collect authoritative completed output items
apcha-oai bf515f4
fix(agents): make streamed result retention opt in
apcha-oai 33c9a07
Merge branch 'main' into apcha/beta-agents-final-result
apcha-oai File filter
Filter by extension
Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
There are no files selected for viewing
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -0,0 +1 @@ | ||
| """Beta SDK runtime helpers.""" |
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -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, | ||
| ) |
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -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) | ||
|
apcha-oai marked this conversation as resolved.
|
||
| 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 | ||
|
apcha-oai marked this conversation as resolved.
|
||
| 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 | ||
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -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() |
Oops, something went wrong.
Oops, something went wrong.
Add this suggestion to a batch that can be applied as a single commit.
This suggestion is invalid because no changes were made to the code.
Suggestions cannot be applied while the pull request is closed.
Suggestions cannot be applied while viewing a subset of changes.
Only one suggestion per line can be applied in a batch.
Add this suggestion to a batch that can be applied as a single commit.
Applying suggestions on deleted lines is not supported.
You must change the existing code in this line in order to create a valid suggestion.
Outdated suggestions cannot be applied.
This suggestion has been applied or marked resolved.
Suggestions cannot be applied from pending reviews.
Suggestions cannot be applied on multi-line comments.
Suggestions cannot be applied while the pull request is queued to merge.
Suggestion cannot be applied right now. Please check back later.
Uh oh!
There was an error while loading. Please reload this page.