Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
88 changes: 88 additions & 0 deletions src/truefoundry_gateway_sdk/agents/_sse_helpers.py
Original file line number Diff line number Diff line change
@@ -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)
31 changes: 17 additions & 14 deletions src/truefoundry_gateway_sdk/agents/prepared_turn.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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:
Expand Down
19 changes: 11 additions & 8 deletions src/truefoundry_gateway_sdk/agents/turn.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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:
"""
Expand Down Expand Up @@ -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:
"""
Expand Down
7 changes: 3 additions & 4 deletions src/truefoundry_gateway_sdk/agents/turn_stream_data.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Loading