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.
@@ -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) :]}