|
| 1 | +from __future__ import annotations |
| 2 | + |
| 3 | +import json |
| 4 | +import asyncio |
| 5 | +import threading |
| 6 | +from urllib.parse import parse_qs, urlsplit |
| 7 | + |
| 8 | +import httpx2 |
| 9 | +import pytest |
| 10 | +from websockets.exceptions import ConnectionClosedError |
| 11 | +from websockets.sync.server import ServerConnection |
| 12 | + |
| 13 | +from openai import OpenAI, AsyncOpenAI, AzureOpenAI, AsyncAzureOpenAI |
| 14 | +from openai.lib.azure import API_KEY_SENTINEL |
| 15 | +from openai.types.websocket_reconnection import ReconnectingEvent, ReconnectingOverrides |
| 16 | +from openai.types.realtime.realtime_error_event import RealtimeErrorEvent |
| 17 | +from openai.types.realtime.input_audio_buffer_cleared_event import InputAudioBufferClearedEvent |
| 18 | + |
| 19 | +from .responses.test_websocket_session import script_server |
| 20 | + |
| 21 | + |
| 22 | +@pytest.fixture(autouse=True) |
| 23 | +def bypass_loopback_proxy(monkeypatch: pytest.MonkeyPatch) -> None: |
| 24 | + monkeypatch.setenv("NO_PROXY", "127.0.0.1") |
| 25 | + monkeypatch.setenv("no_proxy", "127.0.0.1") |
| 26 | + |
| 27 | + |
| 28 | +@pytest.mark.parametrize("mode", ["sync", "async"]) |
| 29 | +@pytest.mark.parametrize("profile", ["direct-recv", "no-callback", "caller-abort", "nonrecoverable", "clean"]) |
| 30 | +async def test_realtime_recovery_stops_without_an_extra_upgrade(mode: str, profile: str) -> None: |
| 31 | + attempts: list[int] = [] |
| 32 | + code = {"nonrecoverable": 1008, "clean": 1000}.get(profile, 1011) |
| 33 | + |
| 34 | + def script(socket: ServerConnection) -> None: |
| 35 | + socket.close(code, "synthetic close") |
| 36 | + |
| 37 | + def on_retry(event: ReconnectingEvent) -> ReconnectingOverrides: |
| 38 | + attempts.append(event.attempt) |
| 39 | + return {"abort": True} |
| 40 | + |
| 41 | + with script_server(script) as url: |
| 42 | + callback = None if profile == "no-callback" else on_retry |
| 43 | + if mode == "sync": |
| 44 | + with OpenAI( |
| 45 | + api_key="fake-realtime-key", base_url=url, http_client=httpx2.Client(trust_env=False) |
| 46 | + ) as client: |
| 47 | + with client.realtime.connect(model="gpt-realtime", on_reconnecting=callback, initial_delay=0) as conn: |
| 48 | + if profile == "clean": |
| 49 | + assert list(conn) == [] |
| 50 | + else: |
| 51 | + with pytest.raises(ConnectionClosedError) as error: |
| 52 | + if profile == "direct-recv": |
| 53 | + conn.recv() |
| 54 | + else: |
| 55 | + next(iter(conn)) |
| 56 | + assert error.value.rcvd is not None and error.value.rcvd.code == code |
| 57 | + else: |
| 58 | + async with AsyncOpenAI( |
| 59 | + api_key="fake-realtime-key", base_url=url, http_client=httpx2.AsyncClient(trust_env=False) |
| 60 | + ) as async_client: |
| 61 | + async with async_client.realtime.connect( |
| 62 | + model="gpt-realtime", on_reconnecting=callback, initial_delay=0 |
| 63 | + ) as async_conn: |
| 64 | + if profile == "clean": |
| 65 | + assert [e async for e in async_conn] == [] |
| 66 | + else: |
| 67 | + with pytest.raises(ConnectionClosedError) as async_error: |
| 68 | + if profile == "direct-recv": |
| 69 | + await asyncio.wait_for(async_conn.recv(), timeout=5) |
| 70 | + else: |
| 71 | + await asyncio.wait_for(async_conn.__aiter__().__anext__(), timeout=5) |
| 72 | + assert async_error.value.rcvd is not None and async_error.value.rcvd.code == code |
| 73 | + assert attempts == ([1] if profile == "caller-abort" else []) |
| 74 | + |
| 75 | + |
| 76 | +@pytest.mark.parametrize("mode", ["sync", "async"]) |
| 77 | +async def test_realtime_admission_errors_do_not_reset_the_retry_budget(mode: str) -> None: |
| 78 | + attempts: list[int] = [] |
| 79 | + seen: list[RealtimeErrorEvent] = [] |
| 80 | + |
| 81 | + def script(socket: ServerConnection) -> None: |
| 82 | + socket.send( |
| 83 | + json.dumps( |
| 84 | + { |
| 85 | + "type": "error", |
| 86 | + "event_id": "fake-error", |
| 87 | + "error": {"type": "server_error", "code": "busy", "message": "Synthetic busy"}, |
| 88 | + } |
| 89 | + ) |
| 90 | + ) |
| 91 | + socket.close(1011, "synthetic close") |
| 92 | + |
| 93 | + def on_retry(event: ReconnectingEvent) -> None: |
| 94 | + assert event.close_code == 1011 |
| 95 | + assert event.max_attempts == 2 |
| 96 | + attempts.append(event.attempt) |
| 97 | + |
| 98 | + with script_server(script, expected_connections=3) as url: |
| 99 | + if mode == "sync": |
| 100 | + with OpenAI( |
| 101 | + api_key="fake-realtime-key", base_url=url, http_client=httpx2.Client(trust_env=False) |
| 102 | + ) as client: |
| 103 | + with client.realtime.connect( |
| 104 | + model="gpt-realtime", on_reconnecting=on_retry, max_retries=2, initial_delay=0 |
| 105 | + ) as conn: |
| 106 | + with pytest.raises(ConnectionClosedError): |
| 107 | + for event in conn: |
| 108 | + assert isinstance(event, RealtimeErrorEvent) |
| 109 | + seen.append(event) |
| 110 | + assert len(seen) <= 3 |
| 111 | + else: |
| 112 | + async with AsyncOpenAI( |
| 113 | + api_key="fake-realtime-key", base_url=url, http_client=httpx2.AsyncClient(trust_env=False) |
| 114 | + ) as client_async: |
| 115 | + async with client_async.realtime.connect( |
| 116 | + model="gpt-realtime", on_reconnecting=on_retry, max_retries=2, initial_delay=0 |
| 117 | + ) as async_conn: |
| 118 | + with pytest.raises(ConnectionClosedError): |
| 119 | + async for event in async_conn: |
| 120 | + assert isinstance(event, RealtimeErrorEvent) |
| 121 | + seen.append(event) |
| 122 | + assert len(seen) <= 3 |
| 123 | + assert attempts == [1, 2] |
| 124 | + assert [e.error.code for e in seen] == ["busy", "busy", "busy"] |
| 125 | + |
| 126 | + |
| 127 | +@pytest.mark.parametrize("mode", ["sync", "async"]) |
| 128 | +@pytest.mark.parametrize("provider", ["openai", "azure"]) |
| 129 | +async def test_realtime_replacement_refreshes_provider_and_preserves_connection_options( |
| 130 | + mode: str, provider: str |
| 131 | +) -> None: |
| 132 | + credential = "fake-realtime-before" |
| 133 | + upgrades: list[str] = [] |
| 134 | + |
| 135 | + def script(socket: ServerConnection) -> None: |
| 136 | + assert socket.request is not None |
| 137 | + target = urlsplit(socket.request.path) |
| 138 | + upgrades.append(socket.request.path) |
| 139 | + assert target.path.endswith("/customer/realtime") |
| 140 | + query = parse_qs(target.query) |
| 141 | + assert query["contract"] == ["socket"] |
| 142 | + assert query["tenant"] == ["sample"] |
| 143 | + assert socket.request.headers.get_all("Authorization") == [f"Bearer {credential}"] |
| 144 | + assert socket.request.headers.get("api-key") is None |
| 145 | + assert socket.request.headers["X-Realtime-Test"] == "connection" |
| 146 | + assert socket.request.headers.get("Sec-WebSocket-Extensions") is None |
| 147 | + if len(upgrades) == 1: |
| 148 | + socket.close(1011, "synthetic close") |
| 149 | + else: |
| 150 | + socket.send('{"type":"input_audio_buffer.cleared","event_id":"recovered"}') |
| 151 | + assert json.loads(socket.recv(timeout=5)) == {"type": "input_audio_buffer.clear", "event_id": "next"} |
| 152 | + |
| 153 | + def on_retry(event: ReconnectingEvent) -> None: |
| 154 | + nonlocal credential |
| 155 | + assert event.attempt == 1 |
| 156 | + credential = "fake-realtime-after" |
| 157 | + |
| 158 | + def get_token() -> str: |
| 159 | + return credential |
| 160 | + |
| 161 | + async def get_async_token() -> str: |
| 162 | + return credential |
| 163 | + |
| 164 | + with script_server(script, expected_connections=2) as url: |
| 165 | + if mode == "sync": |
| 166 | + client = ( |
| 167 | + AzureOpenAI( |
| 168 | + api_key=API_KEY_SENTINEL, |
| 169 | + azure_ad_token_provider=get_token, |
| 170 | + azure_endpoint="https://origin.test", |
| 171 | + websocket_base_url=f"{url.replace('http://', 'ws://')}/customer", |
| 172 | + api_version="2024-01-01", |
| 173 | + http_client=httpx2.Client(trust_env=False), |
| 174 | + ) |
| 175 | + if provider == "azure" |
| 176 | + else OpenAI( |
| 177 | + api_key=get_token, |
| 178 | + base_url=f"{url}/customer?tenant=sample", |
| 179 | + http_client=httpx2.Client(trust_env=False), |
| 180 | + ) |
| 181 | + ) |
| 182 | + with client: |
| 183 | + with client.realtime.connect( |
| 184 | + model="gpt-realtime", |
| 185 | + on_reconnecting=on_retry, |
| 186 | + initial_delay=0, |
| 187 | + extra_query={"contract": "socket", **({"tenant": "sample"} if provider == "azure" else {})}, |
| 188 | + extra_headers={"X-Realtime-Test": "connection"}, |
| 189 | + websocket_connection_options={"compression": None}, |
| 190 | + ) as conn: |
| 191 | + event = next(iter(conn)) |
| 192 | + assert isinstance(event, InputAudioBufferClearedEvent) |
| 193 | + assert event.event_id == "recovered" |
| 194 | + conn.input_audio_buffer.clear(event_id="next") |
| 195 | + else: |
| 196 | + async_client = ( |
| 197 | + AsyncAzureOpenAI( |
| 198 | + api_key=API_KEY_SENTINEL, |
| 199 | + azure_ad_token_provider=get_token, |
| 200 | + azure_endpoint="https://origin.test", |
| 201 | + websocket_base_url=f"{url.replace('http://', 'ws://')}/customer", |
| 202 | + api_version="2024-01-01", |
| 203 | + http_client=httpx2.AsyncClient(trust_env=False), |
| 204 | + ) |
| 205 | + if provider == "azure" |
| 206 | + else AsyncOpenAI( |
| 207 | + api_key=get_async_token, |
| 208 | + base_url=f"{url}/customer?tenant=sample", |
| 209 | + http_client=httpx2.AsyncClient(trust_env=False), |
| 210 | + ) |
| 211 | + ) |
| 212 | + async with async_client: |
| 213 | + async with async_client.realtime.connect( |
| 214 | + model="gpt-realtime", |
| 215 | + on_reconnecting=on_retry, |
| 216 | + initial_delay=0, |
| 217 | + extra_query={"contract": "socket", **({"tenant": "sample"} if provider == "azure" else {})}, |
| 218 | + extra_headers={"X-Realtime-Test": "connection"}, |
| 219 | + websocket_connection_options={"compression": None}, |
| 220 | + ) as async_conn: |
| 221 | + async_event = await asyncio.wait_for(async_conn.__aiter__().__anext__(), timeout=5) |
| 222 | + assert isinstance(async_event, InputAudioBufferClearedEvent) |
| 223 | + assert async_event.event_id == "recovered" |
| 224 | + await async_conn.input_audio_buffer.clear(event_id="next") |
| 225 | + assert len(upgrades) == 2 and upgrades[0] == upgrades[1] |
| 226 | + |
| 227 | + |
| 228 | +async def test_realtime_cancelled_receive_preserves_next_event_and_socket() -> None: |
| 229 | + release = threading.Event() |
| 230 | + |
| 231 | + def script(socket: ServerConnection) -> None: |
| 232 | + assert release.wait(timeout=5) |
| 233 | + socket.send('{"type":"input_audio_buffer.cleared","event_id":"after-cancel"}') |
| 234 | + assert json.loads(socket.recv(timeout=5)) == {"type": "input_audio_buffer.clear", "event_id": "after-cancel"} |
| 235 | + |
| 236 | + with script_server(script) as url: |
| 237 | + async with AsyncOpenAI( |
| 238 | + api_key="fake-realtime-key", base_url=url, http_client=httpx2.AsyncClient(trust_env=False) |
| 239 | + ) as client: |
| 240 | + async with client.realtime.connect(model="gpt-realtime") as conn: |
| 241 | + receiving = asyncio.create_task(conn.recv()) |
| 242 | + await asyncio.sleep(0) |
| 243 | + receiving.cancel() |
| 244 | + try: |
| 245 | + with pytest.raises(asyncio.CancelledError): |
| 246 | + await receiving |
| 247 | + finally: |
| 248 | + release.set() |
| 249 | + event = await asyncio.wait_for(conn.recv(), timeout=5) |
| 250 | + assert isinstance(event, InputAudioBufferClearedEvent) |
| 251 | + assert event.event_id == "after-cancel" |
| 252 | + await conn.input_audio_buffer.clear(event_id="after-cancel") |
0 commit comments