Skip to content

Commit 2632043

Browse files
seratchayaangazali
andauthored
fix(realtime): end iteration after clean server close (#4461)
Co-authored-by: ayaangazali <ayaangazali.work@gmail.com>
1 parent 4cb461a commit 2632043

7 files changed

Lines changed: 247 additions & 4 deletions

File tree

src/agents/realtime/__init__.py

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -62,6 +62,7 @@
6262
RealtimeModelAudioInterruptedEvent,
6363
RealtimeModelCachedTokensDetails,
6464
RealtimeModelConnectionStatusEvent,
65+
RealtimeModelEndOfStreamEvent,
6566
RealtimeModelErrorEvent,
6667
RealtimeModelEvent,
6768
RealtimeModelExceptionEvent,
@@ -166,6 +167,7 @@
166167
"RealtimeModelAudioInterruptedEvent",
167168
"RealtimeModelCachedTokensDetails",
168169
"RealtimeModelConnectionStatusEvent",
170+
"RealtimeModelEndOfStreamEvent",
169171
"RealtimeModelErrorEvent",
170172
"RealtimeModelEvent",
171173
"RealtimeModelExceptionEvent",

src/agents/realtime/model_events.py

Lines changed: 8 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -144,6 +144,13 @@ class RealtimeModelConnectionStatusEvent:
144144
type: Literal["connection_status"] = "connection_status"
145145

146146

147+
@dataclass
148+
class RealtimeModelEndOfStreamEvent:
149+
"""The model event stream ended permanently and will emit no further events."""
150+
151+
type: Literal["end_of_stream"] = "end_of_stream"
152+
153+
147154
@dataclass
148155
class RealtimeModelTurnStartedEvent:
149156
"""Triggered when the model starts generating a response for a turn."""
@@ -248,6 +255,7 @@ class RealtimeModelRawServerEvent:
248255
| RealtimeModelItemUpdatedEvent
249256
| RealtimeModelItemDeletedEvent
250257
| RealtimeModelConnectionStatusEvent
258+
| RealtimeModelEndOfStreamEvent
251259
| RealtimeModelTurnStartedEvent
252260
| RealtimeModelUsageEvent
253261
| RealtimeModelTurnEndedEvent

src/agents/realtime/openai_realtime.py

Lines changed: 15 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -126,6 +126,8 @@
126126
RealtimeModelAudioEvent,
127127
RealtimeModelAudioInterruptedEvent,
128128
RealtimeModelCachedTokensDetails,
129+
RealtimeModelConnectionStatusEvent,
130+
RealtimeModelEndOfStreamEvent,
129131
RealtimeModelErrorEvent,
130132
RealtimeModelEvent,
131133
RealtimeModelExceptionEvent,
@@ -552,6 +554,7 @@ def __init__(self, *, transport_config: TransportConfig | None = None) -> None:
552554
self._websocket: ClientConnection | None = None
553555
self._websocket_task: asyncio.Task[None] | None = None
554556
self._connection_attempt_active = False
557+
self._close_requested = False
555558
self._response_create_tasks: set[asyncio.Task[None]] = set()
556559
self._user_input_lock = asyncio.Lock()
557560
self._listeners: list[RealtimeModelListener] = []
@@ -639,6 +642,7 @@ async def connect(self, options: RealtimeModelConfig) -> None:
639642
headers=headers,
640643
transport_config=self._transport_config,
641644
)
645+
self._close_requested = False
642646
try:
643647
self._websocket_task = asyncio.create_task(self._listen_for_messages())
644648
await self._update_session_config(model_settings)
@@ -724,6 +728,12 @@ async def _emit_event(self, event: RealtimeModelEvent) -> None:
724728
for listener in list(self._listeners):
725729
await listener.on_event(event)
726730

731+
async def _emit_normal_disconnect(self) -> None:
732+
logger.debug("WebSocket connection closed normally")
733+
if not self._close_requested:
734+
await self._emit_event(RealtimeModelConnectionStatusEvent(status="disconnected"))
735+
await self._emit_event(RealtimeModelEndOfStreamEvent())
736+
727737
async def _listen_for_messages(self):
728738
assert self._websocket is not None, "Not connected"
729739

@@ -746,8 +756,7 @@ async def _listen_for_messages(self):
746756
)
747757

748758
except websockets.exceptions.ConnectionClosedOK:
749-
# Normal connection closure - no exception event needed
750-
logger.debug("WebSocket connection closed normally")
759+
await self._emit_normal_disconnect()
751760
except websockets.exceptions.ConnectionClosed as e:
752761
await self._emit_event(
753762
RealtimeModelExceptionEvent(
@@ -760,6 +769,9 @@ async def _listen_for_messages(self):
760769
exception=e, context="WebSocket error in message listener"
761770
)
762771
)
772+
else:
773+
# ClientConnection.__aiter__ consumes ConnectionClosedOK and returns normally.
774+
await self._emit_normal_disconnect()
763775
finally:
764776
await self._cancel_response_create_tasks()
765777
await self._release_response_waiters()
@@ -1249,6 +1261,7 @@ async def close(self) -> None:
12491261
cleanup_error: BaseException | None = None
12501262

12511263
if self._websocket:
1264+
self._close_requested = True
12521265
try:
12531266
await self._websocket.close()
12541267
except BaseException as exc:

src/agents/realtime/session.py

Lines changed: 7 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -80,6 +80,7 @@
8080
)
8181
from .model import RealtimeModel, RealtimeModelConfig, RealtimeModelListener
8282
from .model_events import (
83+
RealtimeModelEndOfStreamEvent,
8384
RealtimeModelEvent,
8485
RealtimeModelInputAudioTranscriptionCompletedEvent,
8586
RealtimeModelOutputTextDeltaEvent,
@@ -227,6 +228,7 @@ def __init__(
227228
self._event_iterator_waiters = 0
228229
self._closing = False
229230
self._closed = False
231+
self._model_stream_ended = False
230232
self._cleanup_task: asyncio.Task[None] | None = None
231233
self._stored_exception: BaseException | None = None
232234
self._pending_tool_calls: dict[str, _PendingToolCall] = {}
@@ -310,7 +312,7 @@ async def __aexit__(self, _exc_type: Any, _exc_val: Any, _exc_tb: Any) -> None:
310312
async def __aiter__(self) -> AsyncIterator[RealtimeSessionEvent]:
311313
"""Iterate over events from the session."""
312314
while True:
313-
if self._closed and self._event_queue.empty():
315+
if (self._closed or self._model_stream_ended) and self._event_queue.empty():
314316
return
315317

316318
# Check if there's a stored exception to raise
@@ -576,6 +578,10 @@ async def on_event(self, event: RealtimeModelEvent) -> None:
576578
)
577579
elif event.type == "connection_status":
578580
pass
581+
elif event.type == "end_of_stream":
582+
assert isinstance(event, RealtimeModelEndOfStreamEvent)
583+
self._model_stream_ended = True
584+
self._wake_event_iterators()
579585
elif event.type == "turn_started":
580586
is_late_start_for_active_response = (
581587
event.response_id is not None

tests/realtime/test_model_events.py

Lines changed: 7 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -27,6 +27,13 @@ def test_usage_event_types_are_publicly_exported() -> None:
2727
assert getattr(realtime, name) is not None
2828

2929

30+
def test_end_of_stream_event_is_publicly_exported() -> None:
31+
assert "RealtimeModelEndOfStreamEvent" in realtime.__all__
32+
33+
event = realtime.RealtimeModelEndOfStreamEvent()
34+
assert event.type == "end_of_stream"
35+
36+
3037
def test_custom_model_can_construct_typed_usage_without_openai_types() -> None:
3138
event = realtime.RealtimeModelUsageEvent(
3239
usage=Usage(requests=1, input_tokens=8, output_tokens=5, total_tokens=13),

tests/realtime/test_openai_realtime.py

Lines changed: 168 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -14,11 +14,12 @@
1414
from agents import Agent, WebSearchTool, function_tool
1515
from agents.exceptions import UserError
1616
from agents.handoffs import handoff
17-
from agents.realtime import RealtimeSessionModelSettings
17+
from agents.realtime import RealtimeAgent, RealtimeSession, RealtimeSessionModelSettings
1818
from agents.realtime.model import RealtimeModelConfig, RealtimePlaybackTracker
1919
from agents.realtime.model_events import (
2020
RealtimeModelAudioEvent,
2121
RealtimeModelAudioInterruptedEvent,
22+
RealtimeModelConnectionStatusEvent,
2223
RealtimeModelErrorEvent,
2324
RealtimeModelOutputTextDeltaEvent,
2425
RealtimeModelRawServerEvent,
@@ -40,6 +41,10 @@
4041
)
4142

4243

44+
async def _collect_session_events(session: RealtimeSession) -> list[Any]:
45+
return [event async for event in session]
46+
47+
4348
class TestOpenAIRealtimeWebSocketModel:
4449
"""Test suite for OpenAIRealtimeWebSocketModel connection and event handling."""
4550

@@ -3261,6 +3266,168 @@ async def handler(websocket):
32613266
await model.close()
32623267
assert model._websocket is None
32633268

3269+
@pytest.mark.asyncio
3270+
async def test_normal_server_close_ends_session_iteration(self):
3271+
"""A clean server close must end session iteration without an exception."""
3272+
3273+
async def handler(websocket):
3274+
await websocket.recv()
3275+
await websocket.close(code=1000, reason="session ended")
3276+
3277+
async with websockets.serve(handler, "127.0.0.1", 0) as server:
3278+
sockets = list(server.sockets)
3279+
port = sockets[0].getsockname()[1]
3280+
session = RealtimeSession(
3281+
OpenAIRealtimeWebSocketModel(),
3282+
RealtimeAgent(name="agent"),
3283+
None,
3284+
model_config={
3285+
"api_key": "test-key",
3286+
"url": f"ws://127.0.0.1:{port}/v1/realtime",
3287+
"initial_model_settings": {"model_name": "gpt-realtime"},
3288+
},
3289+
)
3290+
3291+
await session.__aenter__()
3292+
try:
3293+
events = await asyncio.wait_for(
3294+
_collect_session_events(session),
3295+
timeout=1,
3296+
)
3297+
assert session._closed is False
3298+
finally:
3299+
await session.close()
3300+
3301+
disconnects = [
3302+
event.data
3303+
for event in events
3304+
if event.type == "raw_model_event" and event.data.type == "connection_status"
3305+
]
3306+
assert disconnects == [RealtimeModelConnectionStatusEvent(status="disconnected")]
3307+
assert any(event.type == "history_updated" for event in events)
3308+
3309+
@pytest.mark.asyncio
3310+
async def test_client_close_does_not_emit_server_disconnect(self):
3311+
"""Caller-owned close must not look like a clean server disconnect."""
3312+
3313+
connection_count = 0
3314+
3315+
async def handler(websocket):
3316+
nonlocal connection_count
3317+
connection_count += 1
3318+
if connection_count == 1:
3319+
await websocket.wait_closed()
3320+
else:
3321+
await websocket.recv()
3322+
await websocket.close(code=1000, reason="session ended")
3323+
3324+
async with websockets.serve(handler, "127.0.0.1", 0) as server:
3325+
sockets = list(server.sockets)
3326+
port = sockets[0].getsockname()[1]
3327+
model = OpenAIRealtimeWebSocketModel()
3328+
listener = AsyncMock()
3329+
model.add_listener(listener)
3330+
3331+
await model.connect(
3332+
{
3333+
"api_key": "test-key",
3334+
"url": f"ws://127.0.0.1:{port}/v1/realtime",
3335+
"initial_model_settings": {"model_name": "gpt-realtime"},
3336+
}
3337+
)
3338+
await model.close()
3339+
3340+
first_connection_events = [call.args[0] for call in listener.on_event.await_args_list]
3341+
assert not any(event.type == "connection_status" for event in first_connection_events)
3342+
3343+
await model.connect(
3344+
{
3345+
"api_key": "test-key",
3346+
"url": f"ws://127.0.0.1:{port}/v1/realtime",
3347+
"initial_model_settings": {"model_name": "gpt-realtime"},
3348+
}
3349+
)
3350+
assert model._websocket_task is not None
3351+
await asyncio.wait_for(model._websocket_task, timeout=1)
3352+
await model.close()
3353+
3354+
emitted_events = [call.args[0] for call in listener.on_event.await_args_list]
3355+
disconnects = [event for event in emitted_events if event.type == "connection_status"]
3356+
assert disconnects == [RealtimeModelConnectionStatusEvent(status="disconnected")]
3357+
3358+
@pytest.mark.asyncio
3359+
async def test_cancelled_close_before_websocket_handshake_preserves_server_disconnect(self):
3360+
"""A cancelled preliminary close must not claim transport-close ownership."""
3361+
allow_server_close = asyncio.Event()
3362+
3363+
async def handler(websocket):
3364+
await websocket.recv()
3365+
await allow_server_close.wait()
3366+
await websocket.close(code=1000, reason="session ended")
3367+
3368+
async with websockets.serve(handler, "127.0.0.1", 0) as server:
3369+
sockets = list(server.sockets)
3370+
port = sockets[0].getsockname()[1]
3371+
model = OpenAIRealtimeWebSocketModel()
3372+
listener = AsyncMock()
3373+
model.add_listener(listener)
3374+
3375+
await model.connect(
3376+
{
3377+
"api_key": "test-key",
3378+
"url": f"ws://127.0.0.1:{port}/v1/realtime",
3379+
"initial_model_settings": {"model_name": "gpt-realtime"},
3380+
}
3381+
)
3382+
3383+
cancellation_started = asyncio.Event()
3384+
3385+
async def wait_until_cancelled():
3386+
cancellation_started.set()
3387+
await asyncio.Future()
3388+
3389+
with patch.object(model, "_cancel_response_create_tasks", wait_until_cancelled):
3390+
close_task = asyncio.create_task(model.close())
3391+
await asyncio.wait_for(cancellation_started.wait(), timeout=1)
3392+
close_task.cancel()
3393+
with pytest.raises(asyncio.CancelledError):
3394+
await close_task
3395+
3396+
allow_server_close.set()
3397+
assert model._websocket_task is not None
3398+
await asyncio.wait_for(model._websocket_task, timeout=1)
3399+
await model.close()
3400+
3401+
emitted_events = [call.args[0] for call in listener.on_event.await_args_list]
3402+
disconnects = [event for event in emitted_events if event.type == "connection_status"]
3403+
assert disconnects == [RealtimeModelConnectionStatusEvent(status="disconnected")]
3404+
3405+
@pytest.mark.asyncio
3406+
async def test_abnormal_server_close_still_raises(self):
3407+
"""An abnormal server close must retain the existing exception behavior."""
3408+
3409+
async def handler(websocket):
3410+
await websocket.recv()
3411+
await websocket.close(code=1011, reason="server failure")
3412+
3413+
async with websockets.serve(handler, "127.0.0.1", 0) as server:
3414+
sockets = list(server.sockets)
3415+
port = sockets[0].getsockname()[1]
3416+
session = RealtimeSession(
3417+
OpenAIRealtimeWebSocketModel(),
3418+
RealtimeAgent(name="agent"),
3419+
None,
3420+
model_config={
3421+
"api_key": "test-key",
3422+
"url": f"ws://127.0.0.1:{port}/v1/realtime",
3423+
"initial_model_settings": {"model_name": "gpt-realtime"},
3424+
},
3425+
)
3426+
3427+
with pytest.raises(websockets.exceptions.ConnectionClosedError):
3428+
async with session:
3429+
await asyncio.wait_for(_collect_session_events(session), timeout=1)
3430+
32643431
@pytest.mark.asyncio
32653432
async def test_ping_timeout_success_when_server_responds_quickly(self):
32663433
"""Test that connection stays alive when server responds to pings within timeout."""

0 commit comments

Comments
 (0)