diff --git a/src/truefoundry_gateway_sdk/agents/_sse_helpers.py b/src/truefoundry_gateway_sdk/agents/_sse_helpers.py new file mode 100644 index 0000000..5229bf2 --- /dev/null +++ b/src/truefoundry_gateway_sdk/agents/_sse_helpers.py @@ -0,0 +1,88 @@ +from __future__ import annotations + +import logging +import typing +from json.decoder import JSONDecodeError + +import httpx +from ..core.http_sse._api import EventSource +from ..core.pydantic_utilities import parse_sse_obj +from ..types.turn_streaming_event import TurnStreamingEvent +from .turn_stream_data import TurnStreamData + +_logger = logging.getLogger(__name__) + + +def parse_sequence_number(sse_id: str) -> int: + """Parse the SSE ``id`` field as a sequence number. + + Raises ``ValueError`` when the id is absent or not a valid integer — + mirroring the TypeScript ``parseSequenceNumber`` which throws in the same cases. + """ + if not sse_id: + raise ValueError("Missing SSE sequence number id.") + try: + return int(sse_id) + except (ValueError, TypeError): + raise ValueError(f"Invalid SSE sequence number id: {sse_id!r}.") + + +def iter_sse_stream(response: httpx.Response) -> typing.Iterator[TurnStreamData]: + """Iterate a live httpx SSE response, yielding parsed :class:`TurnStreamData` items. + + Skips unparseable events (with a warning) rather than raising, mirroring the + behaviour of the generated raw client. Raises on a missing or malformed SSE + ``id`` field via :func:`parse_sequence_number`. + """ + for _sse in EventSource(response).iter_sse(): + if not _sse.data: + continue + try: + event = typing.cast( + TurnStreamingEvent, + parse_sse_obj(sse=_sse, type_=TurnStreamingEvent), # type: ignore[arg-type] + ) + except JSONDecodeError as e: + _logger.warning("Skipping SSE event with invalid JSON: %s, sse: %r", e, _sse) + continue + except (TypeError, ValueError, KeyError, AttributeError) as e: + _logger.warning( + "Skipping SSE event due to model construction error: %s: %s, sse: %r", + type(e).__name__, e, _sse, + ) + continue + except Exception as e: + _logger.error( + "Unexpected error processing SSE event: %s: %s, sse: %r", + type(e).__name__, e, _sse, + ) + continue + yield TurnStreamData(sequence_number=parse_sequence_number(_sse.id), event=event) + + +async def aiter_sse_stream(response: httpx.Response) -> typing.AsyncIterator[TurnStreamData]: + """Async version of :func:`iter_sse_stream`.""" + async for _sse in EventSource(response).aiter_sse(): + if not _sse.data: + continue + try: + event = typing.cast( + TurnStreamingEvent, + parse_sse_obj(sse=_sse, type_=TurnStreamingEvent), # type: ignore[arg-type] + ) + except JSONDecodeError as e: + _logger.warning("Skipping SSE event with invalid JSON: %s, sse: %r", e, _sse) + continue + except (TypeError, ValueError, KeyError, AttributeError) as e: + _logger.warning( + "Skipping SSE event due to model construction error: %s: %s, sse: %r", + type(e).__name__, e, _sse, + ) + continue + except Exception as e: + _logger.error( + "Unexpected error processing SSE event: %s: %s, sse: %r", + type(e).__name__, e, _sse, + ) + continue + yield TurnStreamData(sequence_number=parse_sequence_number(_sse.id), event=event) diff --git a/src/truefoundry_gateway_sdk/agents/prepared_turn.py b/src/truefoundry_gateway_sdk/agents/prepared_turn.py index 99998fe..521e3a6 100644 --- a/src/truefoundry_gateway_sdk/agents/prepared_turn.py +++ b/src/truefoundry_gateway_sdk/agents/prepared_turn.py @@ -5,6 +5,7 @@ from ..types.turn import Turn as RawTurn from ..types.turn_created_event import TurnCreatedEvent from ..types.turn_done_event import TurnDoneEvent +from ._sse_helpers import aiter_sse_stream, iter_sse_stream from .turn import AsyncTurn, Turn from .turn_stream_data import TurnStreamData @@ -325,17 +326,18 @@ def _start_and_wait(self, poll_interval_ms: int, request_options: typing.Optiona def _consume_stream(self, request_options: typing.Optional[RequestOptions]) -> typing.Iterator[TurnStreamData]: """Consume the create_turn SSE, adopting the inner Turn from the first turn.created.""" - for event in self._client.agents.sessions.create_turn( + with self._client.agents.sessions.with_raw_response.create_turn( self._session_id, input=self._input_param, previous_turn_id=self._previous_turn_id, request_options=request_options, - ): - if isinstance(event, TurnCreatedEvent) and self._turn is None: - self._adopt_turn(event) - elif self._turn is not None and isinstance(event, TurnDoneEvent): - self._replace_turn_state(event.state) - yield TurnStreamData(sequence_number=None, event=event) + ) as r: + for item in iter_sse_stream(r._response): + if isinstance(item.event, TurnCreatedEvent) and self._turn is None: + self._adopt_turn(item.event) + elif self._turn is not None and isinstance(item.event, TurnDoneEvent): + self._replace_turn_state(item.event.state) + yield item def _must_get_turn(self) -> Turn: if self._turn is None: @@ -686,17 +688,18 @@ async def _start_and_wait( async def _consume_stream( self, request_options: typing.Optional[RequestOptions] ) -> typing.AsyncIterator[TurnStreamData]: - async for event in self._client.agents.sessions.create_turn( + async with self._client.agents.sessions.with_raw_response.create_turn( self._session_id, input=self._input_param, previous_turn_id=self._previous_turn_id, request_options=request_options, - ): - if isinstance(event, TurnCreatedEvent) and self._turn is None: - self._adopt_turn(event) - elif self._turn is not None and isinstance(event, TurnDoneEvent): - self._replace_turn_state(event.state) - yield TurnStreamData(sequence_number=None, event=event) + ) as r: + async for item in aiter_sse_stream(r._response): + if isinstance(item.event, TurnCreatedEvent) and self._turn is None: + self._adopt_turn(item.event) + elif self._turn is not None and isinstance(item.event, TurnDoneEvent): + self._replace_turn_state(item.event.state) + yield item def _must_get_turn(self) -> AsyncTurn: if self._turn is None: diff --git a/src/truefoundry_gateway_sdk/agents/turn.py b/src/truefoundry_gateway_sdk/agents/turn.py index d9eb804..2192a08 100644 --- a/src/truefoundry_gateway_sdk/agents/turn.py +++ b/src/truefoundry_gateway_sdk/agents/turn.py @@ -8,6 +8,7 @@ from ..types.turn_state_cancelled import TurnStateCancelled from ..types.turn_state_done import TurnStateDone from ..types.turn_state_error import TurnStateError +from ._sse_helpers import aiter_sse_stream, iter_sse_stream from .turn_stream_data import TurnStreamData # this is used as the default value for optional parameters @@ -216,14 +217,15 @@ def stream( TurnStreamData SSE stream items. """ - for event in self._client.agents.sessions.subscribe_to_turn( + with self._client.agents.sessions.with_raw_response.subscribe_to_turn( self._session_id, self._id, after_sequence_number=after_sequence_number, request_options=request_options, - ): - self._apply_event(event) - yield TurnStreamData(sequence_number=None, event=event) + ) as r: + for item in iter_sse_stream(r._response): + self._apply_event(item.event) + yield item def cancel(self, *, request_options: typing.Optional[RequestOptions] = None) -> None: """ @@ -463,14 +465,15 @@ async def stream( TurnStreamData SSE stream items. """ - async for event in self._client.agents.sessions.subscribe_to_turn( + async with self._client.agents.sessions.with_raw_response.subscribe_to_turn( self._session_id, self._id, after_sequence_number=after_sequence_number, request_options=request_options, - ): - self._apply_event(event) - yield TurnStreamData(sequence_number=None, event=event) + ) as r: + async for item in aiter_sse_stream(r._response): + self._apply_event(item.event) + yield item async def cancel(self, *, request_options: typing.Optional[RequestOptions] = None) -> None: """ diff --git a/src/truefoundry_gateway_sdk/agents/turn_stream_data.py b/src/truefoundry_gateway_sdk/agents/turn_stream_data.py index f84261e..00e4826 100644 --- a/src/truefoundry_gateway_sdk/agents/turn_stream_data.py +++ b/src/truefoundry_gateway_sdk/agents/turn_stream_data.py @@ -12,12 +12,11 @@ class TurnStreamData: """ Attributes ---------- - sequence_number : typing.Optional[int] - SSE event id for resume; None if unavailable from the stream. + sequence_number : int + SSE event id used for resume via ``subscribe_to_turn``. event : TurnStreamingEvent Streaming event payload. """ - # Sequence number from the SSE event id. None when not available from the underlying stream. - sequence_number: typing.Optional[int] + sequence_number: int event: TurnStreamingEvent