|
14 | 14 | from agents import Agent, WebSearchTool, function_tool |
15 | 15 | from agents.exceptions import UserError |
16 | 16 | from agents.handoffs import handoff |
17 | | -from agents.realtime import RealtimeSessionModelSettings |
| 17 | +from agents.realtime import RealtimeAgent, RealtimeSession, RealtimeSessionModelSettings |
18 | 18 | from agents.realtime.model import RealtimeModelConfig, RealtimePlaybackTracker |
19 | 19 | from agents.realtime.model_events import ( |
20 | 20 | RealtimeModelAudioEvent, |
21 | 21 | RealtimeModelAudioInterruptedEvent, |
| 22 | + RealtimeModelConnectionStatusEvent, |
22 | 23 | RealtimeModelErrorEvent, |
23 | 24 | RealtimeModelOutputTextDeltaEvent, |
24 | 25 | RealtimeModelRawServerEvent, |
|
40 | 41 | ) |
41 | 42 |
|
42 | 43 |
|
| 44 | +async def _collect_session_events(session: RealtimeSession) -> list[Any]: |
| 45 | + return [event async for event in session] |
| 46 | + |
| 47 | + |
43 | 48 | class TestOpenAIRealtimeWebSocketModel: |
44 | 49 | """Test suite for OpenAIRealtimeWebSocketModel connection and event handling.""" |
45 | 50 |
|
@@ -3261,6 +3266,168 @@ async def handler(websocket): |
3261 | 3266 | await model.close() |
3262 | 3267 | assert model._websocket is None |
3263 | 3268 |
|
| 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 | + |
3264 | 3431 | @pytest.mark.asyncio |
3265 | 3432 | async def test_ping_timeout_success_when_server_responds_quickly(self): |
3266 | 3433 | """Test that connection stays alive when server responds to pings within timeout.""" |
|
0 commit comments