diff --git a/src/truefoundry_gateway_sdk/agents/_sse_helpers.py b/src/truefoundry_gateway_sdk/agents/_sse_helpers.py deleted file mode 100644 index 5229bf2..0000000 --- a/src/truefoundry_gateway_sdk/agents/_sse_helpers.py +++ /dev/null @@ -1,88 +0,0 @@ -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 521e3a6..95ab766 100644 --- a/src/truefoundry_gateway_sdk/agents/prepared_turn.py +++ b/src/truefoundry_gateway_sdk/agents/prepared_turn.py @@ -5,9 +5,8 @@ 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 +from .turn_stream_data import TurnStreamData, parse_sequence_number # this is used as the default value for optional parameters OMIT = typing.cast(typing.Any, ...) @@ -326,18 +325,19 @@ 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.""" - with self._client.agents.sessions.with_raw_response.create_turn( + with self._client.agents.sessions.create_turn( self._session_id, input=self._input_param, previous_turn_id=self._previous_turn_id, request_options=request_options, - ) 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 + ) as sse: + for event in sse.with_metadata(): + sequence_number = parse_sequence_number(event.id) + if isinstance(event.data, TurnCreatedEvent) and self._turn is None: + self._adopt_turn(event.data) + elif self._turn is not None and isinstance(event.data, TurnDoneEvent): + self._replace_turn_state(event.data.state) + yield TurnStreamData(sequence_number=sequence_number, event=event.data) def _must_get_turn(self) -> Turn: if self._turn is None: @@ -688,18 +688,19 @@ async def _start_and_wait( async def _consume_stream( self, request_options: typing.Optional[RequestOptions] ) -> typing.AsyncIterator[TurnStreamData]: - async with self._client.agents.sessions.with_raw_response.create_turn( + async with self._client.agents.sessions.create_turn( self._session_id, input=self._input_param, previous_turn_id=self._previous_turn_id, request_options=request_options, - ) 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 + ) as sse: + async for event in sse.with_metadata(): + sequence_number = parse_sequence_number(event.id) + if isinstance(event.data, TurnCreatedEvent) and self._turn is None: + self._adopt_turn(event.data) + elif self._turn is not None and isinstance(event.data, TurnDoneEvent): + self._replace_turn_state(event.data.state) + yield TurnStreamData(sequence_number=sequence_number, event=event.data) 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 2192a08..11b4df6 100644 --- a/src/truefoundry_gateway_sdk/agents/turn.py +++ b/src/truefoundry_gateway_sdk/agents/turn.py @@ -8,8 +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 +from .turn_stream_data import TurnStreamData, parse_sequence_number # this is used as the default value for optional parameters OMIT = typing.cast(typing.Any, ...) @@ -217,15 +216,16 @@ def stream( TurnStreamData SSE stream items. """ - with self._client.agents.sessions.with_raw_response.subscribe_to_turn( + with self._client.agents.sessions.subscribe_to_turn( self._session_id, self._id, after_sequence_number=after_sequence_number, request_options=request_options, - ) as r: - for item in iter_sse_stream(r._response): - self._apply_event(item.event) - yield item + ) as sse: + for event in sse.with_metadata(): + sequence_number = parse_sequence_number(event.id) + self._apply_event(event.data) + yield TurnStreamData(sequence_number=sequence_number, event=event.data) def cancel(self, *, request_options: typing.Optional[RequestOptions] = None) -> None: """ @@ -465,15 +465,16 @@ async def stream( TurnStreamData SSE stream items. """ - async with self._client.agents.sessions.with_raw_response.subscribe_to_turn( + async with self._client.agents.sessions.subscribe_to_turn( self._session_id, self._id, after_sequence_number=after_sequence_number, request_options=request_options, - ) as r: - async for item in aiter_sse_stream(r._response): - self._apply_event(item.event) - yield item + ) as sse: + async for event in sse.with_metadata(): + sequence_number = parse_sequence_number(event.id) + self._apply_event(event.data) + yield TurnStreamData(sequence_number=sequence_number, event=event.data) 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 00e4826..615a733 100644 --- a/src/truefoundry_gateway_sdk/agents/turn_stream_data.py +++ b/src/truefoundry_gateway_sdk/agents/turn_stream_data.py @@ -20,3 +20,17 @@ class TurnStreamData: sequence_number: int event: TurnStreamingEvent + + +def parse_sequence_number(sse_id: typing.Optional[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``. + """ + if not sse_id: + raise ValueError("Missing SSE sequence number id.") + try: + return int(sse_id) + except (ValueError, TypeError) as exc: + raise ValueError(f"Invalid SSE sequence number id: {sse_id!r}.") from exc