pi-python-core 0.8.1__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- pi_python/__init__.py +160 -0
- pi_python/_version.py +1 -0
- pi_python/agent.py +396 -0
- pi_python/cancellation.py +24 -0
- pi_python/data/models.json +3315 -0
- pi_python/errors.py +49 -0
- pi_python/estimate.py +144 -0
- pi_python/events.py +138 -0
- pi_python/function_tools.py +438 -0
- pi_python/hooks.py +44 -0
- pi_python/limits.py +28 -0
- pi_python/loop.py +431 -0
- pi_python/lowlevel.py +179 -0
- pi_python/mcp.py +187 -0
- pi_python/messages.py +405 -0
- pi_python/models.py +155 -0
- pi_python/provider.py +123 -0
- pi_python/providers/__init__.py +21 -0
- pi_python/providers/anthropic.py +673 -0
- pi_python/providers/common.py +201 -0
- pi_python/providers/completions.py +1149 -0
- pi_python/providers/oauth.py +542 -0
- pi_python/providers/openai.py +681 -0
- pi_python/providers/transport.py +574 -0
- pi_python/proxy.py +304 -0
- pi_python/py.typed +0 -0
- pi_python/queues.py +76 -0
- pi_python/recovery.py +209 -0
- pi_python/run.py +419 -0
- pi_python/stream.py +251 -0
- pi_python/sync.py +78 -0
- pi_python/testing.py +25 -0
- pi_python/tools.py +546 -0
- pi_python/transcript.py +167 -0
- pi_python_core-0.8.1.dist-info/METADATA +119 -0
- pi_python_core-0.8.1.dist-info/RECORD +39 -0
- pi_python_core-0.8.1.dist-info/WHEEL +4 -0
- pi_python_core-0.8.1.dist-info/licenses/LICENSE +21 -0
- pi_python_core-0.8.1.dist-info/licenses/NOTICE +8 -0
|
@@ -0,0 +1,574 @@
|
|
|
1
|
+
"""HTTP/SSE and WebSocket transport. Optional dependencies load only on use."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
from collections.abc import AsyncGenerator, Awaitable, Callable, Mapping
|
|
5
|
+
from typing import TypeVar
|
|
6
|
+
import asyncio
|
|
7
|
+
import codecs
|
|
8
|
+
import json
|
|
9
|
+
import hashlib
|
|
10
|
+
import time
|
|
11
|
+
from email.utils import parsedate_to_datetime
|
|
12
|
+
from collections.abc import AsyncIterator
|
|
13
|
+
from contextlib import asynccontextmanager
|
|
14
|
+
from dataclasses import dataclass
|
|
15
|
+
from typing import Any
|
|
16
|
+
from ..cancellation import CancelToken
|
|
17
|
+
from ..errors import PiError, ProviderProtocolError
|
|
18
|
+
import math
|
|
19
|
+
import re
|
|
20
|
+
from urllib.parse import quote
|
|
21
|
+
from ..tools import invoke
|
|
22
|
+
|
|
23
|
+
|
|
24
|
+
T = TypeVar("T")
|
|
25
|
+
|
|
26
|
+
|
|
27
|
+
MAX_ERROR_BODY_CHARS = 4000 # the same cap as Pi
|
|
28
|
+
# Header names that carry credentials: authorization, x-api-key, x-goog-api-key, cookie, ...
|
|
29
|
+
_CREDENTIAL_HEADER = re.compile(r"auth|key|token|secret|cookie|password", re.I)
|
|
30
|
+
|
|
31
|
+
|
|
32
|
+
class ProviderHTTPError(PiError):
|
|
33
|
+
"""Structured HTTP failure. As in Pi, a model request's error carries the server's
|
|
34
|
+
response body, capped, so callers can recognize causes such as context overflow;
|
|
35
|
+
credentials sent with the request are redacted from it."""
|
|
36
|
+
|
|
37
|
+
def __init__(
|
|
38
|
+
self,
|
|
39
|
+
status: int,
|
|
40
|
+
message: str = "Provider request failed",
|
|
41
|
+
*,
|
|
42
|
+
retry_after: float | None = None,
|
|
43
|
+
request_id: str | None = None,
|
|
44
|
+
body: str | None = None,
|
|
45
|
+
) -> None:
|
|
46
|
+
self.status = status
|
|
47
|
+
self.retry_after = retry_after
|
|
48
|
+
self.request_id = request_id
|
|
49
|
+
self.body = body
|
|
50
|
+
self.retryable = status in {429, 500, 502, 503, 504}
|
|
51
|
+
self.category = (
|
|
52
|
+
"authentication"
|
|
53
|
+
if status in {401, 403}
|
|
54
|
+
else "rate_limit"
|
|
55
|
+
if status == 429
|
|
56
|
+
else "server"
|
|
57
|
+
if status >= 500
|
|
58
|
+
else "request"
|
|
59
|
+
)
|
|
60
|
+
detail = f": {body}" if body else ""
|
|
61
|
+
super().__init__(f"{message} (HTTP {status}; {self.category}){detail}")
|
|
62
|
+
|
|
63
|
+
|
|
64
|
+
async def error_body(response: Any, headers: Mapping[str, str]) -> str | None:
|
|
65
|
+
"""Read a failed response's body (bounded), redact request credentials, cap its length."""
|
|
66
|
+
data = b""
|
|
67
|
+
try:
|
|
68
|
+
async for chunk in response.aiter_bytes():
|
|
69
|
+
data += chunk
|
|
70
|
+
if len(data) > 4 * MAX_ERROR_BODY_CHARS:
|
|
71
|
+
break
|
|
72
|
+
except Exception:
|
|
73
|
+
return None
|
|
74
|
+
text = data.decode("utf-8", "replace").strip()
|
|
75
|
+
for name, value in headers.items():
|
|
76
|
+
if _CREDENTIAL_HEADER.search(name):
|
|
77
|
+
for part in str(value).split():
|
|
78
|
+
if len(part) >= 8:
|
|
79
|
+
# Also as echoed inside JSON ("/" escaped) or a URL (percent-encoded).
|
|
80
|
+
for form in {part, part.replace("/", "\\/"), quote(part, safe="")}:
|
|
81
|
+
text = text.replace(form, "[redacted]")
|
|
82
|
+
if len(text) > MAX_ERROR_BODY_CHARS:
|
|
83
|
+
text = (
|
|
84
|
+
f"{text[:MAX_ERROR_BODY_CHARS]}... [truncated {len(text) - MAX_ERROR_BODY_CHARS} chars]"
|
|
85
|
+
)
|
|
86
|
+
return text or None
|
|
87
|
+
|
|
88
|
+
|
|
89
|
+
def retry_delay(headers: Mapping[str, str]) -> float | None:
|
|
90
|
+
value = headers.get("retry-after")
|
|
91
|
+
if value is None:
|
|
92
|
+
return None
|
|
93
|
+
try:
|
|
94
|
+
delay = float(value)
|
|
95
|
+
except ValueError:
|
|
96
|
+
try:
|
|
97
|
+
delay = parsedate_to_datetime(value).timestamp() - time.time()
|
|
98
|
+
except (ValueError, TypeError, OverflowError):
|
|
99
|
+
return None
|
|
100
|
+
return max(0.0, delay) if math.isfinite(delay) else None
|
|
101
|
+
|
|
102
|
+
|
|
103
|
+
async def cancellable(awaitable: Awaitable[T], cancel: CancelToken) -> T:
|
|
104
|
+
task = asyncio.ensure_future(awaitable)
|
|
105
|
+
signal = asyncio.create_task(cancel.wait())
|
|
106
|
+
try:
|
|
107
|
+
cancel.raise_if_cancelled()
|
|
108
|
+
done, _ = await asyncio.wait({task, signal}, return_when=asyncio.FIRST_COMPLETED)
|
|
109
|
+
if signal in done:
|
|
110
|
+
raise asyncio.CancelledError(cancel.reason)
|
|
111
|
+
return task.result()
|
|
112
|
+
finally:
|
|
113
|
+
signal.cancel()
|
|
114
|
+
if not task.done():
|
|
115
|
+
task.cancel()
|
|
116
|
+
await asyncio.gather(task, signal, return_exceptions=True)
|
|
117
|
+
|
|
118
|
+
|
|
119
|
+
async def sse_events(chunks: AsyncIterator[bytes]) -> AsyncGenerator[dict[str, Any], None]:
|
|
120
|
+
"""SSE framing across UTF-8 chunks, CR/LF variants, multiline data and EOF."""
|
|
121
|
+
decoder = codecs.getincrementaldecoder("utf-8")()
|
|
122
|
+
buffer = ""
|
|
123
|
+
data: list[str] = []
|
|
124
|
+
event_name = None
|
|
125
|
+
|
|
126
|
+
async def emit() -> dict[str, Any] | None:
|
|
127
|
+
nonlocal data, event_name
|
|
128
|
+
if not data:
|
|
129
|
+
event_name = None
|
|
130
|
+
return None
|
|
131
|
+
payload = "\n".join(data)
|
|
132
|
+
data = []
|
|
133
|
+
if payload == "[DONE]":
|
|
134
|
+
event_name = None
|
|
135
|
+
return {"type": "transport_done"}
|
|
136
|
+
try:
|
|
137
|
+
value = json.loads(payload)
|
|
138
|
+
except ValueError as exc:
|
|
139
|
+
raise ProviderProtocolError("Invalid SSE JSON") from exc
|
|
140
|
+
if not isinstance(value, dict):
|
|
141
|
+
raise ProviderProtocolError("SSE data must be an object")
|
|
142
|
+
if "type" not in value and event_name:
|
|
143
|
+
value["type"] = event_name
|
|
144
|
+
event_name = None
|
|
145
|
+
return value
|
|
146
|
+
|
|
147
|
+
async def lines(final: bool = False) -> AsyncGenerator[dict[str, Any], None]:
|
|
148
|
+
nonlocal buffer, event_name
|
|
149
|
+
while True:
|
|
150
|
+
positions = [p for p in (buffer.find("\n"), buffer.find("\r")) if p >= 0]
|
|
151
|
+
if not positions:
|
|
152
|
+
if not final or not buffer:
|
|
153
|
+
break
|
|
154
|
+
line, buffer = buffer, ""
|
|
155
|
+
else:
|
|
156
|
+
pos = min(positions)
|
|
157
|
+
if buffer[pos] == "\r" and pos + 1 == len(buffer) and not final:
|
|
158
|
+
break
|
|
159
|
+
end = pos + 2 if buffer[pos : pos + 2] == "\r\n" else pos + 1
|
|
160
|
+
line, buffer = buffer[:pos], buffer[end:]
|
|
161
|
+
if not line:
|
|
162
|
+
value = await emit()
|
|
163
|
+
if value is not None:
|
|
164
|
+
yield value
|
|
165
|
+
elif line.startswith("data:"):
|
|
166
|
+
field = line[5:]
|
|
167
|
+
data.append(field[1:] if field.startswith(" ") else field)
|
|
168
|
+
elif line.startswith("event:"):
|
|
169
|
+
event_name = line[6:].lstrip(" ")
|
|
170
|
+
|
|
171
|
+
async for chunk in chunks:
|
|
172
|
+
buffer += decoder.decode(chunk)
|
|
173
|
+
async for value in lines():
|
|
174
|
+
yield value
|
|
175
|
+
buffer += decoder.decode(b"", final=True)
|
|
176
|
+
async for value in lines(final=True):
|
|
177
|
+
yield value
|
|
178
|
+
last = await emit()
|
|
179
|
+
if last is not None:
|
|
180
|
+
yield last
|
|
181
|
+
|
|
182
|
+
|
|
183
|
+
class HTTPTransport:
|
|
184
|
+
def __init__(
|
|
185
|
+
self,
|
|
186
|
+
client: Any = None,
|
|
187
|
+
*,
|
|
188
|
+
timeout: float = 300,
|
|
189
|
+
max_retries: int = 2,
|
|
190
|
+
retry_base: float = 0.5,
|
|
191
|
+
max_retry_delay: float = 30,
|
|
192
|
+
websocket_idle_ttl: float = 5 * 60,
|
|
193
|
+
websocket_max_age: float = 55 * 60,
|
|
194
|
+
) -> None:
|
|
195
|
+
if type(max_retries) is not int or not 0 <= max_retries <= 10:
|
|
196
|
+
raise ValueError("max_retries must be between 0 and 10")
|
|
197
|
+
if retry_base < 0 or max_retry_delay < 0:
|
|
198
|
+
raise ValueError("Retry delays must be nonnegative")
|
|
199
|
+
self.client = client
|
|
200
|
+
self.timeout = timeout
|
|
201
|
+
self.max_retries = max_retries
|
|
202
|
+
self.retry_base = retry_base
|
|
203
|
+
self.max_retry_delay = max_retry_delay
|
|
204
|
+
self.websocket_idle_ttl = websocket_idle_ttl
|
|
205
|
+
self.websocket_max_age = websocket_max_age
|
|
206
|
+
self._clock = time.monotonic
|
|
207
|
+
self._sockets: dict[tuple, CachedSocket] = {}
|
|
208
|
+
self._socket_locks: dict[tuple, asyncio.Lock] = {}
|
|
209
|
+
self._stats = dict(
|
|
210
|
+
http_attempts=0,
|
|
211
|
+
http_retries=0,
|
|
212
|
+
websocket_connects=0,
|
|
213
|
+
websocket_reuses=0,
|
|
214
|
+
websocket_fallbacks=0,
|
|
215
|
+
websocket_expired=0,
|
|
216
|
+
websocket_retries=0,
|
|
217
|
+
websocket_delta_requests=0,
|
|
218
|
+
)
|
|
219
|
+
|
|
220
|
+
@property
|
|
221
|
+
def stats(self) -> dict[str, int]:
|
|
222
|
+
return dict(self._stats)
|
|
223
|
+
|
|
224
|
+
async def aclose(self) -> None:
|
|
225
|
+
sockets, self._sockets = self._sockets, {}
|
|
226
|
+
await asyncio.gather(*(entry.socket.close() for entry in sockets.values()))
|
|
227
|
+
self._socket_locks.clear()
|
|
228
|
+
|
|
229
|
+
async def __aenter__(self) -> HTTPTransport:
|
|
230
|
+
return self
|
|
231
|
+
|
|
232
|
+
async def __aexit__(self, *args: object) -> None:
|
|
233
|
+
await self.aclose()
|
|
234
|
+
|
|
235
|
+
@asynccontextmanager
|
|
236
|
+
async def _client(self) -> AsyncIterator[Any]:
|
|
237
|
+
if self.client is not None:
|
|
238
|
+
yield self.client
|
|
239
|
+
else:
|
|
240
|
+
import httpx
|
|
241
|
+
|
|
242
|
+
async with httpx.AsyncClient(
|
|
243
|
+
timeout=httpx.Timeout(self.timeout, connect=30),
|
|
244
|
+
trust_env=False,
|
|
245
|
+
follow_redirects=False,
|
|
246
|
+
) as client:
|
|
247
|
+
yield client
|
|
248
|
+
|
|
249
|
+
async def request_json(
|
|
250
|
+
self,
|
|
251
|
+
url: str,
|
|
252
|
+
body: dict[str, Any],
|
|
253
|
+
headers: dict[str, str],
|
|
254
|
+
cancel: CancelToken,
|
|
255
|
+
*,
|
|
256
|
+
form: bool = False,
|
|
257
|
+
) -> dict[str, Any]:
|
|
258
|
+
async with self._client() as client:
|
|
259
|
+
response = await cancellable(
|
|
260
|
+
client.post(url, headers=headers, **({"data": body} if form else {"json": body})),
|
|
261
|
+
cancel,
|
|
262
|
+
)
|
|
263
|
+
if response.status_code >= 400:
|
|
264
|
+
raise ProviderHTTPError(response.status_code)
|
|
265
|
+
value = response.json()
|
|
266
|
+
if not isinstance(value, dict):
|
|
267
|
+
raise ProviderProtocolError("Expected JSON object")
|
|
268
|
+
return value
|
|
269
|
+
|
|
270
|
+
async def get_json(self, url: str, cancel: CancelToken) -> Any:
|
|
271
|
+
async with self._client() as client:
|
|
272
|
+
response = await cancellable(client.get(url), cancel)
|
|
273
|
+
if response.status_code >= 400:
|
|
274
|
+
raise ProviderHTTPError(response.status_code)
|
|
275
|
+
return response.json()
|
|
276
|
+
|
|
277
|
+
async def stream(
|
|
278
|
+
self,
|
|
279
|
+
url: str,
|
|
280
|
+
body: dict[str, Any],
|
|
281
|
+
headers: dict[str, str],
|
|
282
|
+
cancel: CancelToken,
|
|
283
|
+
*,
|
|
284
|
+
on_response: Callable | None = None,
|
|
285
|
+
) -> AsyncGenerator[dict[str, Any], None]:
|
|
286
|
+
async with self._client() as client:
|
|
287
|
+
for attempt in range(self.max_retries + 1):
|
|
288
|
+
request = client.build_request("POST", url, headers=headers, json=body)
|
|
289
|
+
# A network exception may follow server acceptance. Never replay it.
|
|
290
|
+
self._stats["http_attempts"] += 1
|
|
291
|
+
response = await cancellable(client.send(request, stream=True), cancel)
|
|
292
|
+
try:
|
|
293
|
+
await invoke(
|
|
294
|
+
on_response,
|
|
295
|
+
{"status": response.status_code, "headers": dict(response.headers)},
|
|
296
|
+
)
|
|
297
|
+
if response.status_code < 400:
|
|
298
|
+
break
|
|
299
|
+
error = ProviderHTTPError(
|
|
300
|
+
response.status_code,
|
|
301
|
+
retry_after=retry_delay(response.headers),
|
|
302
|
+
request_id=response.headers.get("x-request-id"),
|
|
303
|
+
body=await cancellable(error_body(response, headers), cancel),
|
|
304
|
+
)
|
|
305
|
+
except BaseException:
|
|
306
|
+
await response.aclose()
|
|
307
|
+
raise
|
|
308
|
+
await response.aclose()
|
|
309
|
+
if not error.retryable or attempt == self.max_retries:
|
|
310
|
+
raise error
|
|
311
|
+
delay = (
|
|
312
|
+
error.retry_after
|
|
313
|
+
if error.retry_after is not None
|
|
314
|
+
else self.retry_base * 2**attempt
|
|
315
|
+
)
|
|
316
|
+
# Do not retry sooner than a server's requested delay.
|
|
317
|
+
if delay > self.max_retry_delay:
|
|
318
|
+
raise error
|
|
319
|
+
self._stats["http_retries"] += 1
|
|
320
|
+
await cancellable(asyncio.sleep(delay), cancel)
|
|
321
|
+
try:
|
|
322
|
+
iterator = sse_events(response.aiter_bytes())
|
|
323
|
+
while True:
|
|
324
|
+
try:
|
|
325
|
+
event = await cancellable(anext(iterator), cancel)
|
|
326
|
+
except StopAsyncIteration:
|
|
327
|
+
break
|
|
328
|
+
yield event
|
|
329
|
+
finally:
|
|
330
|
+
await response.aclose()
|
|
331
|
+
|
|
332
|
+
async def responses(
|
|
333
|
+
self,
|
|
334
|
+
url: str,
|
|
335
|
+
body: dict[str, Any],
|
|
336
|
+
headers: dict[str, str],
|
|
337
|
+
cancel: CancelToken,
|
|
338
|
+
*,
|
|
339
|
+
mode: str = "sse",
|
|
340
|
+
session_id: str | None = None,
|
|
341
|
+
on_response: Callable | None = None,
|
|
342
|
+
link: WebSocketLink | None = None,
|
|
343
|
+
) -> AsyncGenerator[dict[str, Any], None]:
|
|
344
|
+
if mode == "sse":
|
|
345
|
+
events = self.stream(url, body, headers, cancel, on_response=on_response)
|
|
346
|
+
else:
|
|
347
|
+
events = self.websocket(
|
|
348
|
+
url,
|
|
349
|
+
body,
|
|
350
|
+
headers,
|
|
351
|
+
cancel,
|
|
352
|
+
on_response=on_response,
|
|
353
|
+
session_id=session_id if mode in {"websocket-cached", "auto"} else None,
|
|
354
|
+
fallback=mode == "auto",
|
|
355
|
+
link=link,
|
|
356
|
+
)
|
|
357
|
+
try:
|
|
358
|
+
async for event in events:
|
|
359
|
+
yield event
|
|
360
|
+
finally:
|
|
361
|
+
await events.aclose()
|
|
362
|
+
|
|
363
|
+
def _reusable(self, entry: CachedSocket) -> bool:
|
|
364
|
+
"""Pi acquireWebSocket: reuse only an open, young socket that was not idle too long."""
|
|
365
|
+
now = self._clock()
|
|
366
|
+
state = getattr(entry.socket, "state", None)
|
|
367
|
+
open_ = state is None or getattr(state, "name", "OPEN") == "OPEN"
|
|
368
|
+
return (
|
|
369
|
+
open_
|
|
370
|
+
and now - entry.created < self.websocket_max_age
|
|
371
|
+
and now - entry.released < self.websocket_idle_ttl
|
|
372
|
+
)
|
|
373
|
+
|
|
374
|
+
async def websocket(
|
|
375
|
+
self,
|
|
376
|
+
url: str,
|
|
377
|
+
body: dict[str, Any],
|
|
378
|
+
headers: dict[str, str],
|
|
379
|
+
cancel: CancelToken,
|
|
380
|
+
*,
|
|
381
|
+
on_response: Callable | None = None,
|
|
382
|
+
session_id: str | None = None,
|
|
383
|
+
fallback: bool = False,
|
|
384
|
+
link: WebSocketLink | None = None,
|
|
385
|
+
) -> AsyncGenerator[dict[str, Any], None]:
|
|
386
|
+
from websockets.asyncio.client import connect
|
|
387
|
+
from websockets.exceptions import InvalidHandshake
|
|
388
|
+
from websockets.version import version as websockets_version
|
|
389
|
+
|
|
390
|
+
# websockets 15 started reading proxy settings from the environment; turn that
|
|
391
|
+
# off. Older releases never use a proxy and do not accept the argument.
|
|
392
|
+
no_proxy: dict[str, Any] = (
|
|
393
|
+
{"proxy": None} if int(websockets_version.split(".")[0]) >= 15 else {}
|
|
394
|
+
)
|
|
395
|
+
|
|
396
|
+
endpoint = url.replace("https://", "wss://", 1).replace("http://", "ws://", 1)
|
|
397
|
+
# Include all headers to prevent reuse across accounts or incompatible options.
|
|
398
|
+
key = (
|
|
399
|
+
endpoint,
|
|
400
|
+
session_id,
|
|
401
|
+
hashlib.sha256(json.dumps(headers, sort_keys=True).encode()).hexdigest(),
|
|
402
|
+
)
|
|
403
|
+
lock = self._socket_locks.setdefault(key, asyncio.Lock()) if session_id else asyncio.Lock()
|
|
404
|
+
await cancellable(lock.acquire(), cancel)
|
|
405
|
+
entry: CachedSocket | None = None
|
|
406
|
+
complete = False
|
|
407
|
+
retried: set[str] = set()
|
|
408
|
+
try:
|
|
409
|
+
while True:
|
|
410
|
+
entry = self._sockets.pop(key, None) if session_id else None
|
|
411
|
+
if entry is not None and not self._reusable(entry):
|
|
412
|
+
self._stats["websocket_expired"] += 1
|
|
413
|
+
await entry.socket.close()
|
|
414
|
+
entry = None
|
|
415
|
+
if entry is not None:
|
|
416
|
+
self._stats["websocket_reuses"] += 1
|
|
417
|
+
else:
|
|
418
|
+
try:
|
|
419
|
+
socket = await cancellable(
|
|
420
|
+
connect(
|
|
421
|
+
endpoint,
|
|
422
|
+
additional_headers={
|
|
423
|
+
**headers,
|
|
424
|
+
"OpenAI-Beta": "responses_websockets=2026-02-06",
|
|
425
|
+
},
|
|
426
|
+
open_timeout=30,
|
|
427
|
+
max_size=16 * 1024 * 1024,
|
|
428
|
+
**no_proxy,
|
|
429
|
+
),
|
|
430
|
+
cancel,
|
|
431
|
+
)
|
|
432
|
+
self._stats["websocket_connects"] += 1
|
|
433
|
+
except (
|
|
434
|
+
OSError,
|
|
435
|
+
TimeoutError,
|
|
436
|
+
InvalidHandshake,
|
|
437
|
+
):
|
|
438
|
+
if not fallback:
|
|
439
|
+
raise
|
|
440
|
+
self._stats["websocket_fallbacks"] += 1
|
|
441
|
+
# No response.create has been sent, so SSE fallback is safe.
|
|
442
|
+
events = self.stream(url, body, headers, cancel, on_response=on_response)
|
|
443
|
+
try:
|
|
444
|
+
async for event in events:
|
|
445
|
+
yield event
|
|
446
|
+
finally:
|
|
447
|
+
await events.aclose()
|
|
448
|
+
return
|
|
449
|
+
entry = CachedSocket(socket, self._clock(), self._clock())
|
|
450
|
+
await invoke(
|
|
451
|
+
on_response,
|
|
452
|
+
{"status": 101, "headers": dict(entry.socket.response.headers.raw_items())},
|
|
453
|
+
)
|
|
454
|
+
# A continuation is consumed here; the provider restores one only after a
|
|
455
|
+
# complete response, so any failure falls back to the full context.
|
|
456
|
+
state, entry.continuation = entry.continuation, None
|
|
457
|
+
sent = continued_body(body, state) if link is not None else None
|
|
458
|
+
if sent is not None:
|
|
459
|
+
self._stats["websocket_delta_requests"] += 1
|
|
460
|
+
if link is not None:
|
|
461
|
+
link.entry = entry
|
|
462
|
+
await cancellable(
|
|
463
|
+
entry.socket.send(
|
|
464
|
+
json.dumps(
|
|
465
|
+
{
|
|
466
|
+
"type": "response.create",
|
|
467
|
+
**{k: v for k, v in (sent or body).items() if k != "stream"},
|
|
468
|
+
}
|
|
469
|
+
)
|
|
470
|
+
),
|
|
471
|
+
cancel,
|
|
472
|
+
)
|
|
473
|
+
output = False
|
|
474
|
+
retry = None
|
|
475
|
+
while True:
|
|
476
|
+
raw = await cancellable(entry.socket.recv(), cancel)
|
|
477
|
+
event = json.loads(raw)
|
|
478
|
+
if not isinstance(event, dict):
|
|
479
|
+
raise ProviderProtocolError("WebSocket event must be an object")
|
|
480
|
+
code = error_code(event)
|
|
481
|
+
# Rejected before any model output: Pi retries once with the full
|
|
482
|
+
# context on a fresh connection. Nothing was generated, so this
|
|
483
|
+
# cannot duplicate a response.
|
|
484
|
+
if not output and code in _RETRYABLE_WEBSOCKET_CODES and code not in retried:
|
|
485
|
+
if code != "previous_response_not_found" or sent is not None:
|
|
486
|
+
retry = code
|
|
487
|
+
break
|
|
488
|
+
output = output or str(event.get("type", "")).startswith(
|
|
489
|
+
("response.output", "response.content_part", "response.reasoning")
|
|
490
|
+
)
|
|
491
|
+
terminal = event.get("type") in {
|
|
492
|
+
"response.completed",
|
|
493
|
+
"response.done",
|
|
494
|
+
"response.failed",
|
|
495
|
+
"response.incomplete",
|
|
496
|
+
"error",
|
|
497
|
+
}
|
|
498
|
+
# A response cut at the output limit still ends cleanly (Pi keeps it).
|
|
499
|
+
complete = event.get("type") in {
|
|
500
|
+
"response.completed",
|
|
501
|
+
"response.done",
|
|
502
|
+
"response.incomplete",
|
|
503
|
+
}
|
|
504
|
+
yield event
|
|
505
|
+
if terminal:
|
|
506
|
+
break
|
|
507
|
+
if retry is None:
|
|
508
|
+
break
|
|
509
|
+
retried.add(retry)
|
|
510
|
+
self._stats["websocket_retries"] += 1
|
|
511
|
+
await entry.socket.close()
|
|
512
|
+
entry = None
|
|
513
|
+
finally:
|
|
514
|
+
if entry is not None:
|
|
515
|
+
if session_id and complete and not cancel.cancelled:
|
|
516
|
+
entry.released = self._clock()
|
|
517
|
+
self._sockets[key] = entry
|
|
518
|
+
else:
|
|
519
|
+
await entry.socket.close()
|
|
520
|
+
lock.release()
|
|
521
|
+
|
|
522
|
+
|
|
523
|
+
_RETRYABLE_WEBSOCKET_CODES = {"previous_response_not_found", "websocket_connection_limit_reached"}
|
|
524
|
+
|
|
525
|
+
|
|
526
|
+
def error_code(event: dict) -> str | None:
|
|
527
|
+
"""Pi extractCodexEventError, plus the failed-response code."""
|
|
528
|
+
|
|
529
|
+
def field(value: dict, name: str) -> dict:
|
|
530
|
+
nested = value.get(name)
|
|
531
|
+
return nested if isinstance(nested, dict) else {}
|
|
532
|
+
|
|
533
|
+
if event.get("type") == "error":
|
|
534
|
+
code = event.get("code", field(event, "error").get("code"))
|
|
535
|
+
elif event.get("type") == "response.failed":
|
|
536
|
+
code = field(field(event, "response"), "error").get("code")
|
|
537
|
+
else:
|
|
538
|
+
code = None
|
|
539
|
+
return code if isinstance(code, str) else None
|
|
540
|
+
|
|
541
|
+
|
|
542
|
+
@dataclass
|
|
543
|
+
class CachedSocket:
|
|
544
|
+
socket: Any
|
|
545
|
+
created: float
|
|
546
|
+
released: float
|
|
547
|
+
continuation: dict | None = None
|
|
548
|
+
|
|
549
|
+
|
|
550
|
+
@dataclass
|
|
551
|
+
class WebSocketLink:
|
|
552
|
+
"""Lets a provider record the response items that the next request will repeat."""
|
|
553
|
+
|
|
554
|
+
entry: CachedSocket | None = None
|
|
555
|
+
|
|
556
|
+
|
|
557
|
+
def continued_body(body: dict, state: dict | None) -> dict | None:
|
|
558
|
+
"""Pi buildCachedWebSocketRequestBody: send only new input after the last response.
|
|
559
|
+
|
|
560
|
+
Valid only when every other field is unchanged and the new input starts with the
|
|
561
|
+
previous input followed by the previous response's own items.
|
|
562
|
+
"""
|
|
563
|
+
if not state or not state.get("response_id"):
|
|
564
|
+
return None
|
|
565
|
+
last = state["body"]
|
|
566
|
+
if {k: v for k, v in body.items() if k not in {"input", "previous_response_id"}} != {
|
|
567
|
+
k: v for k, v in last.items() if k not in {"input", "previous_response_id"}
|
|
568
|
+
}:
|
|
569
|
+
return None
|
|
570
|
+
baseline = [*last.get("input", []), *state["items"]]
|
|
571
|
+
current = body.get("input", [])
|
|
572
|
+
if len(current) < len(baseline) or current[: len(baseline)] != baseline:
|
|
573
|
+
return None
|
|
574
|
+
return {**body, "previous_response_id": state["response_id"], "input": current[len(baseline) :]}
|