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
pi_python/proxy.py
ADDED
|
@@ -0,0 +1,304 @@
|
|
|
1
|
+
"""Pi agent-core streamProxy wire protocol (v1.0.0, MIT)."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
from collections.abc import AsyncGenerator, AsyncIterator, Callable
|
|
6
|
+
from .messages import Message
|
|
7
|
+
from .providers.transport import HTTPTransport
|
|
8
|
+
from typing import Any
|
|
9
|
+
from copy import deepcopy
|
|
10
|
+
import json
|
|
11
|
+
from .cancellation import CancelToken
|
|
12
|
+
from .errors import ConfigurationError, ProviderProtocolError, UnsupportedCapabilityError
|
|
13
|
+
from .messages import AssistantMessage, TextContent, ThinkingContent, ToolCall, message_to_dict
|
|
14
|
+
from .provider import ModelEvent, ModelRequest
|
|
15
|
+
from .tools import invoke
|
|
16
|
+
from .stream import event_contract
|
|
17
|
+
from .providers.common import RemoteProvider
|
|
18
|
+
import asyncio
|
|
19
|
+
from dataclasses import replace
|
|
20
|
+
|
|
21
|
+
_KEYS = {
|
|
22
|
+
"tool_call": "toolCall",
|
|
23
|
+
"tool_result": "toolResult",
|
|
24
|
+
"tool_use": "toolUse",
|
|
25
|
+
"text_signature": "textSignature",
|
|
26
|
+
"thinking_signature": "thinkingSignature",
|
|
27
|
+
"thought_signature": "thoughtSignature",
|
|
28
|
+
"mime_type": "mimeType",
|
|
29
|
+
"call_id": "toolCallId",
|
|
30
|
+
"is_error": "isError",
|
|
31
|
+
"stop_reason": "stopReason",
|
|
32
|
+
"provider_thinking_level": "providerThinkingLevel",
|
|
33
|
+
"tools_added": "toolsAdded",
|
|
34
|
+
"tools_removed": "toolsRemoved",
|
|
35
|
+
"input_schema": "parameters",
|
|
36
|
+
"error": "errorMessage",
|
|
37
|
+
"response_id": "responseId",
|
|
38
|
+
"response_model": "responseModel",
|
|
39
|
+
"thinking_level": "thinkingLevel",
|
|
40
|
+
"raw_stop_reason": "rawStopReason",
|
|
41
|
+
"end_turn": "endTurn",
|
|
42
|
+
"nested_calls": "nestedCalls",
|
|
43
|
+
"duration_ms": "durationMs",
|
|
44
|
+
"arguments_bytes": "argumentsBytes",
|
|
45
|
+
"model_id": "modelId",
|
|
46
|
+
"expires_at": "expiresAt",
|
|
47
|
+
"poll_after_ms": "pollAfterMs",
|
|
48
|
+
}
|
|
49
|
+
_OPTIONS = {
|
|
50
|
+
"max_tokens": "maxTokens",
|
|
51
|
+
"sampling_params": "samplingParams",
|
|
52
|
+
"cache_retention": "cacheRetention",
|
|
53
|
+
"session_id": "sessionId",
|
|
54
|
+
"thinking_budgets": "thinkingBudgets",
|
|
55
|
+
"max_retry_delay_ms": "maxRetryDelayMs",
|
|
56
|
+
}
|
|
57
|
+
_ALLOWED = {
|
|
58
|
+
"temperature",
|
|
59
|
+
"samplingParams",
|
|
60
|
+
"maxTokens",
|
|
61
|
+
"reasoning",
|
|
62
|
+
"cacheRetention",
|
|
63
|
+
"sessionId",
|
|
64
|
+
"headers",
|
|
65
|
+
"metadata",
|
|
66
|
+
"transport",
|
|
67
|
+
"thinkingBudgets",
|
|
68
|
+
"maxRetryDelayMs",
|
|
69
|
+
}
|
|
70
|
+
|
|
71
|
+
|
|
72
|
+
def pi_usage(value: dict[str, Any], *, decode: bool = False) -> dict[str, Any]:
|
|
73
|
+
keys = {
|
|
74
|
+
"cache_read": "cacheRead",
|
|
75
|
+
"cache_write": "cacheWrite",
|
|
76
|
+
"cache_write_1h": "cacheWrite1h",
|
|
77
|
+
"total_tokens": "totalTokens",
|
|
78
|
+
}
|
|
79
|
+
if decode:
|
|
80
|
+
keys = {v: k for k, v in keys.items()}
|
|
81
|
+
return {
|
|
82
|
+
keys.get(k, k): pi_usage(v, decode=decode)
|
|
83
|
+
if k == "cost" and isinstance(v, dict)
|
|
84
|
+
else deepcopy(v)
|
|
85
|
+
for k, v in value.items()
|
|
86
|
+
}
|
|
87
|
+
|
|
88
|
+
|
|
89
|
+
def pi_message(message: Message) -> dict[str, Any]:
|
|
90
|
+
def convert(value: Any) -> Any:
|
|
91
|
+
if isinstance(value, list):
|
|
92
|
+
return [convert(v) for v in value]
|
|
93
|
+
if isinstance(value, dict):
|
|
94
|
+
result = {
|
|
95
|
+
_KEYS.get(k, k): (
|
|
96
|
+
deepcopy(v)
|
|
97
|
+
if k
|
|
98
|
+
in {
|
|
99
|
+
"arguments",
|
|
100
|
+
"input_schema",
|
|
101
|
+
"data",
|
|
102
|
+
"usage",
|
|
103
|
+
"sections",
|
|
104
|
+
"details",
|
|
105
|
+
"structured_content",
|
|
106
|
+
}
|
|
107
|
+
else convert(v)
|
|
108
|
+
)
|
|
109
|
+
for k, v in value.items()
|
|
110
|
+
}
|
|
111
|
+
for key in ("type", "role", "stopReason"):
|
|
112
|
+
if key in result:
|
|
113
|
+
result[key] = _KEYS.get(result[key], result[key])
|
|
114
|
+
if value.get("role") in {"assistant", "tool_result"} and isinstance(
|
|
115
|
+
value.get("usage"), dict
|
|
116
|
+
):
|
|
117
|
+
result["usage"] = pi_usage(value["usage"])
|
|
118
|
+
if value.get("role") == "tool_result":
|
|
119
|
+
result["toolName"] = result.pop("name")
|
|
120
|
+
if value.get("role") == "system":
|
|
121
|
+
result["toolsRemoved"] = [{"name": name} for name in value.get("tools_removed", [])]
|
|
122
|
+
if "timestamp" in result:
|
|
123
|
+
result["timestamp"] *= 1000
|
|
124
|
+
return result
|
|
125
|
+
return value
|
|
126
|
+
|
|
127
|
+
return convert(message_to_dict(message))
|
|
128
|
+
|
|
129
|
+
|
|
130
|
+
class ProxyProvider(RemoteProvider):
|
|
131
|
+
name = "proxy"
|
|
132
|
+
|
|
133
|
+
def __init__(
|
|
134
|
+
self,
|
|
135
|
+
*,
|
|
136
|
+
model: dict[str, Any],
|
|
137
|
+
proxy_url: str,
|
|
138
|
+
auth_token: str | Callable[[], Any],
|
|
139
|
+
transport: HTTPTransport | None = None,
|
|
140
|
+
) -> None:
|
|
141
|
+
super().__init__(api_key=auth_token, transport=transport)
|
|
142
|
+
if not {"id", "provider", "api"} <= model.keys():
|
|
143
|
+
raise ConfigurationError(
|
|
144
|
+
"Proxy model needs id, provider, api; use the server model descriptor"
|
|
145
|
+
)
|
|
146
|
+
self.model = deepcopy(model)
|
|
147
|
+
self.proxy_url = proxy_url.rstrip("/")
|
|
148
|
+
|
|
149
|
+
@event_contract
|
|
150
|
+
async def stream(
|
|
151
|
+
self, request: ModelRequest, cancel: CancelToken
|
|
152
|
+
) -> AsyncGenerator[ModelEvent, None]:
|
|
153
|
+
token = await self.credential(request, cancel)
|
|
154
|
+
options = {
|
|
155
|
+
_OPTIONS.get(k, k): v
|
|
156
|
+
for k, v in request.options.items()
|
|
157
|
+
if _OPTIONS.get(k, k) in _ALLOWED
|
|
158
|
+
}
|
|
159
|
+
# Transcript tool declarations are carried by system messages, as upstream does.
|
|
160
|
+
messages = [pi_message(m) for m in request.messages]
|
|
161
|
+
if request.tools:
|
|
162
|
+
messages.append(
|
|
163
|
+
{
|
|
164
|
+
"role": "system",
|
|
165
|
+
"content": "",
|
|
166
|
+
"sections": {},
|
|
167
|
+
"toolsAdded": [
|
|
168
|
+
{"name": t.name, "description": t.description, "parameters": t.input_schema}
|
|
169
|
+
for t in request.tools
|
|
170
|
+
],
|
|
171
|
+
"toolsRemoved": [],
|
|
172
|
+
"timestamp": 0,
|
|
173
|
+
}
|
|
174
|
+
)
|
|
175
|
+
body = await self.payload(
|
|
176
|
+
request, {"model": self.model, "context": {"messages": messages}, "options": options}
|
|
177
|
+
)
|
|
178
|
+
events = self.transport.stream(
|
|
179
|
+
self.proxy_url + "/api/stream",
|
|
180
|
+
body,
|
|
181
|
+
{"authorization": f"Bearer {token}", "content-type": "application/json"},
|
|
182
|
+
cancel,
|
|
183
|
+
on_response=request.on_response,
|
|
184
|
+
)
|
|
185
|
+
blocks: dict[int, Any] = {}
|
|
186
|
+
arguments = {}
|
|
187
|
+
final = None
|
|
188
|
+
closed_blocks = set()
|
|
189
|
+
yield ModelEvent("start")
|
|
190
|
+
try:
|
|
191
|
+
async for event in events:
|
|
192
|
+
await invoke(request.on_provider_stream_event, deepcopy(event))
|
|
193
|
+
kind = event.get("type")
|
|
194
|
+
raw_index = event.get("contentIndex")
|
|
195
|
+
index = raw_index if type(raw_index) is int else -1
|
|
196
|
+
if str(kind).endswith(("_start", "_delta", "_end")) and index < 0:
|
|
197
|
+
raise ProviderProtocolError(f"Proxy {kind} without contentIndex")
|
|
198
|
+
if final is not None:
|
|
199
|
+
raise ProviderProtocolError("Proxy event after completion")
|
|
200
|
+
if kind in {"text_start", "thinking_start", "toolcall_start"}:
|
|
201
|
+
if type(index) is not int or index != len(blocks):
|
|
202
|
+
raise ProviderProtocolError("Invalid proxy block index")
|
|
203
|
+
if kind == "text_start":
|
|
204
|
+
blocks[index] = TextContent("")
|
|
205
|
+
elif kind == "thinking_start":
|
|
206
|
+
blocks[index] = ThinkingContent("")
|
|
207
|
+
else:
|
|
208
|
+
blocks[index] = ToolCall(event["id"], event["toolName"], {})
|
|
209
|
+
arguments[index] = ""
|
|
210
|
+
yield ModelEvent.boundary("start", index, blocks[index])
|
|
211
|
+
elif (
|
|
212
|
+
kind in {"text_delta", "thinking_delta", "toolcall_delta"}
|
|
213
|
+
and index in closed_blocks
|
|
214
|
+
):
|
|
215
|
+
raise ProviderProtocolError("Proxy delta after block end")
|
|
216
|
+
elif kind == "text_delta":
|
|
217
|
+
blocks[index].text += event["delta"]
|
|
218
|
+
yield ModelEvent.text(event["delta"], index)
|
|
219
|
+
elif kind == "thinking_delta":
|
|
220
|
+
blocks[index].thinking += event["delta"]
|
|
221
|
+
yield ModelEvent.thinking(event["delta"], index)
|
|
222
|
+
elif kind == "toolcall_delta":
|
|
223
|
+
arguments[index] += event["delta"]
|
|
224
|
+
b = blocks[index]
|
|
225
|
+
yield ModelEvent.toolcall(event["delta"], index)
|
|
226
|
+
elif kind in {"text_end", "thinking_end"}:
|
|
227
|
+
if index in closed_blocks:
|
|
228
|
+
raise ProviderProtocolError("Duplicate proxy block end")
|
|
229
|
+
closed_blocks.add(index)
|
|
230
|
+
b = blocks[index]
|
|
231
|
+
if kind == "text_end":
|
|
232
|
+
b.text_signature = event.get("contentSignature")
|
|
233
|
+
else:
|
|
234
|
+
b.thinking_signature = event.get("contentSignature")
|
|
235
|
+
yield ModelEvent.boundary("end", index, b)
|
|
236
|
+
elif kind == "toolcall_end":
|
|
237
|
+
if index in closed_blocks:
|
|
238
|
+
raise ProviderProtocolError("Duplicate proxy block end")
|
|
239
|
+
closed_blocks.add(index)
|
|
240
|
+
value = event["toolCall"]
|
|
241
|
+
b = blocks[index]
|
|
242
|
+
if value["id"] != b.id or value["name"] != b.name:
|
|
243
|
+
raise ProviderProtocolError("Proxy tool identity changed")
|
|
244
|
+
b.arguments = value["arguments"]
|
|
245
|
+
b.thought_signature = value.get("thoughtSignature")
|
|
246
|
+
b.namespace = value.get("namespace")
|
|
247
|
+
if arguments[index] and json.loads(arguments[index]) != b.arguments:
|
|
248
|
+
raise ProviderProtocolError("Proxy tool arguments mismatch")
|
|
249
|
+
yield ModelEvent.boundary("end", index, b)
|
|
250
|
+
elif kind == "done":
|
|
251
|
+
if closed_blocks != set(blocks):
|
|
252
|
+
raise ProviderProtocolError("Proxy completed with unfinished blocks")
|
|
253
|
+
reason = {"stop": "stop", "length": "length", "toolUse": "tool_use"}.get(
|
|
254
|
+
event["reason"]
|
|
255
|
+
)
|
|
256
|
+
if reason is None:
|
|
257
|
+
raise ProviderProtocolError("Invalid proxy stop reason")
|
|
258
|
+
final = AssistantMessage(
|
|
259
|
+
list(blocks.values()),
|
|
260
|
+
reason,
|
|
261
|
+
self.model["provider"],
|
|
262
|
+
self.model["id"],
|
|
263
|
+
pi_usage(event.get("usage", {}), decode=True),
|
|
264
|
+
api=self.model["api"],
|
|
265
|
+
provider_thinking_level=event.get("providerThinkingLevel"),
|
|
266
|
+
)
|
|
267
|
+
elif kind == "error":
|
|
268
|
+
if event.get("reason") == "aborted":
|
|
269
|
+
raise asyncio.CancelledError("Proxy request aborted")
|
|
270
|
+
raise ProviderProtocolError("Proxy stream reported an error")
|
|
271
|
+
elif kind != "start":
|
|
272
|
+
raise UnsupportedCapabilityError(f"Unsupported proxy event: {kind}")
|
|
273
|
+
finally:
|
|
274
|
+
await events.aclose()
|
|
275
|
+
if final is None:
|
|
276
|
+
raise ProviderProtocolError("Proxy stream ended without completion")
|
|
277
|
+
message_to_dict(final)
|
|
278
|
+
yield ModelEvent.done(final)
|
|
279
|
+
|
|
280
|
+
|
|
281
|
+
def stream_proxy(
|
|
282
|
+
model: dict[str, Any],
|
|
283
|
+
context: ModelRequest | list[Message],
|
|
284
|
+
options: dict[str, Any],
|
|
285
|
+
cancel: CancelToken | None = None,
|
|
286
|
+
*,
|
|
287
|
+
transport: HTTPTransport | None = None,
|
|
288
|
+
) -> AsyncIterator[ModelEvent]:
|
|
289
|
+
"""Return an async iterator of ModelEvent; context is ModelRequest or message list."""
|
|
290
|
+
options = dict(options)
|
|
291
|
+
provider = ProxyProvider(
|
|
292
|
+
model=model,
|
|
293
|
+
proxy_url=options.pop("proxy_url"),
|
|
294
|
+
auth_token=options.pop("auth_token"),
|
|
295
|
+
transport=transport,
|
|
296
|
+
)
|
|
297
|
+
request = (
|
|
298
|
+
context
|
|
299
|
+
if isinstance(context, ModelRequest)
|
|
300
|
+
else ModelRequest(context, model=model["id"], options=options)
|
|
301
|
+
)
|
|
302
|
+
if isinstance(context, ModelRequest):
|
|
303
|
+
request = replace(context, options={**context.options, **options})
|
|
304
|
+
return provider.stream(request, cancel or CancelToken())
|
pi_python/py.typed
ADDED
|
File without changes
|
pi_python/queues.py
ADDED
|
@@ -0,0 +1,76 @@
|
|
|
1
|
+
"""Steering and follow-up queues, consumed at Pi's safe boundaries."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
from copy import deepcopy
|
|
5
|
+
|
|
6
|
+
from .errors import ConfigurationError
|
|
7
|
+
from .messages import Message
|
|
8
|
+
|
|
9
|
+
MODES = {"all", "one_at_a_time"}
|
|
10
|
+
|
|
11
|
+
|
|
12
|
+
class MessageQueues:
|
|
13
|
+
"""Two input queues. A taken message is reserved until history commits it.
|
|
14
|
+
|
|
15
|
+
If a run ends before committing a reserved message, `restore` puts it back at the
|
|
16
|
+
front of its queue, so an interrupted run never loses queued input.
|
|
17
|
+
"""
|
|
18
|
+
|
|
19
|
+
def __init__(self, steering_mode: str = "one_at_a_time", follow_up_mode: str = "one_at_a_time"):
|
|
20
|
+
if steering_mode not in MODES or follow_up_mode not in MODES:
|
|
21
|
+
raise ConfigurationError("Invalid queue mode")
|
|
22
|
+
self.steering_mode, self.follow_up_mode = steering_mode, follow_up_mode
|
|
23
|
+
self.steering: list[Message] = []
|
|
24
|
+
self.follow_up: list[Message] = []
|
|
25
|
+
self._reserved: list[tuple[list[Message], Message]] = []
|
|
26
|
+
|
|
27
|
+
def take(self, steering: bool) -> list[Message]:
|
|
28
|
+
queue = self.steering if steering else self.follow_up
|
|
29
|
+
mode = self.steering_mode if steering else self.follow_up_mode
|
|
30
|
+
count = len(queue) if mode == "all" else min(1, len(queue))
|
|
31
|
+
taken = queue[:count]
|
|
32
|
+
del queue[:count]
|
|
33
|
+
self._reserved.extend((queue, message) for message in taken)
|
|
34
|
+
return taken
|
|
35
|
+
|
|
36
|
+
def committed(self, message: Message, pending: list[Message]) -> None:
|
|
37
|
+
"""Release the reservation for a committed message from `pending`.
|
|
38
|
+
|
|
39
|
+
Declaration reconciliation can change a system message's tool fields, so the
|
|
40
|
+
match uses identity within `pending` plus timestamp and role.
|
|
41
|
+
"""
|
|
42
|
+
for i, (_, reserved) in enumerate(self._reserved):
|
|
43
|
+
if (
|
|
44
|
+
any(reserved is item for item in pending)
|
|
45
|
+
and reserved.timestamp == message.timestamp
|
|
46
|
+
and reserved.role == message.role
|
|
47
|
+
):
|
|
48
|
+
del self._reserved[i]
|
|
49
|
+
return
|
|
50
|
+
|
|
51
|
+
def restore(self) -> None:
|
|
52
|
+
for queue, message in reversed(self._reserved):
|
|
53
|
+
queue.insert(0, message)
|
|
54
|
+
self._reserved.clear()
|
|
55
|
+
|
|
56
|
+
def clear(self, *, steering: bool, follow_up: bool) -> None:
|
|
57
|
+
if steering:
|
|
58
|
+
self.steering.clear()
|
|
59
|
+
if follow_up:
|
|
60
|
+
self.follow_up.clear()
|
|
61
|
+
self._reserved = [
|
|
62
|
+
(queue, message)
|
|
63
|
+
for queue, message in self._reserved
|
|
64
|
+
if not (
|
|
65
|
+
(steering and queue is self.steering) or (follow_up and queue is self.follow_up)
|
|
66
|
+
)
|
|
67
|
+
]
|
|
68
|
+
|
|
69
|
+
def peek(self) -> list[Message]:
|
|
70
|
+
"""What the next boundary would take: steering first, else follow-up."""
|
|
71
|
+
steering = self.steering if self.steering_mode == "all" else self.steering[:1]
|
|
72
|
+
follow_up = self.follow_up if self.follow_up_mode == "all" else self.follow_up[:1]
|
|
73
|
+
return deepcopy(steering if steering else follow_up)
|
|
74
|
+
|
|
75
|
+
def __bool__(self) -> bool:
|
|
76
|
+
return bool(self.steering or self.follow_up)
|
pi_python/recovery.py
ADDED
|
@@ -0,0 +1,209 @@
|
|
|
1
|
+
"""Recognize failed responses an application can recover from.
|
|
2
|
+
|
|
3
|
+
Ported from pi-ai (utils/overflow.ts, utils/retry.ts at the pinned Pi revision). The
|
|
4
|
+
functions read a committed assistant message; the application decides what to do, for
|
|
5
|
+
example compact the history and continue after an overflow, or back off and continue
|
|
6
|
+
after a transient error. `Agent.continue_run()` retries a failed or aborted last response.
|
|
7
|
+
"""
|
|
8
|
+
|
|
9
|
+
from __future__ import annotations
|
|
10
|
+
|
|
11
|
+
import re
|
|
12
|
+
|
|
13
|
+
from .messages import AssistantMessage
|
|
14
|
+
|
|
15
|
+
# Context overflow messages, per provider (see the pinned overflow.ts for examples).
|
|
16
|
+
_OVERFLOW = [
|
|
17
|
+
re.compile(pattern, re.I)
|
|
18
|
+
for pattern in (
|
|
19
|
+
r"prompt (?:is )?too long", # Anthropic, z.ai
|
|
20
|
+
r"prompt exceeds max length", # z.ai CN
|
|
21
|
+
r"request_too_large", # Anthropic byte-size overflow (HTTP 413)
|
|
22
|
+
r"input is too long for requested model", # Amazon Bedrock
|
|
23
|
+
r"exceeds the context window", # OpenAI Completions and Responses
|
|
24
|
+
r"exceeds (?:the )?(?:model'?s )?maximum context length(?: of [\d,]+ tokens?|\s*\([\d,]+\))",
|
|
25
|
+
r"input token count.*exceeds the maximum", # Google Gemini
|
|
26
|
+
r"maximum prompt length is \d+", # xAI
|
|
27
|
+
r"reduce the length of the messages", # Groq
|
|
28
|
+
r"maximum context length is \d+ tokens", # OpenRouter
|
|
29
|
+
r"exceeds (?:the )?maximum allowed input length of [\d,]+ tokens?", # OpenRouter/Poolside
|
|
30
|
+
r"input \(\d+ tokens\) is longer than the model'?s context length \(\d+ tokens\)", # Together
|
|
31
|
+
r"exceeds the limit of \d+", # GitHub Copilot
|
|
32
|
+
r"exceeds the available context size", # llama.cpp server
|
|
33
|
+
r"greater than the context length", # LM Studio
|
|
34
|
+
r"context window exceeds limit", # MiniMax
|
|
35
|
+
r"exceeded model token limit", # Kimi For Coding
|
|
36
|
+
r"too large for model with \d+ maximum context length", # Mistral
|
|
37
|
+
r"prompt has [\d,]+ tokens?, but the configured context size is [\d,]+ tokens?", # DS4
|
|
38
|
+
r"model_context_window_exceeded", # z.ai finish reason surfaced as an error
|
|
39
|
+
r"prompt too long; exceeded (?:max )?context length", # Ollama
|
|
40
|
+
r"range of input length should be", # DashScope / Qwen
|
|
41
|
+
r"context[_ ]length[_ ]exceeded",
|
|
42
|
+
r"too many tokens",
|
|
43
|
+
r"token limit exceeded",
|
|
44
|
+
)
|
|
45
|
+
]
|
|
46
|
+
_CEREBRAS_OVERFLOW = re.compile(r"^4(?:00|13)\s*(?:status code)?\s*\(no body\)", re.I)
|
|
47
|
+
# Throttling text that would otherwise match an overflow pattern ("Too many tokens, please wait").
|
|
48
|
+
_NOT_OVERFLOW = [
|
|
49
|
+
re.compile(pattern, re.I)
|
|
50
|
+
for pattern in (
|
|
51
|
+
r"^(Throttling error|Service unavailable):",
|
|
52
|
+
r"rate limit",
|
|
53
|
+
r"too many requests",
|
|
54
|
+
)
|
|
55
|
+
]
|
|
56
|
+
|
|
57
|
+
# Account limits look like throttling but do not clear in seconds.
|
|
58
|
+
_NOT_RETRYABLE = re.compile(
|
|
59
|
+
"|".join(
|
|
60
|
+
(
|
|
61
|
+
"GoUsageLimitError",
|
|
62
|
+
"FreeUsageLimitError",
|
|
63
|
+
"Monthly usage limit reached",
|
|
64
|
+
"available balance",
|
|
65
|
+
"insufficient_quota",
|
|
66
|
+
"out of budget",
|
|
67
|
+
"quota exceeded",
|
|
68
|
+
"billing",
|
|
69
|
+
"subscription_sharing_usage_limit_exceeded",
|
|
70
|
+
)
|
|
71
|
+
),
|
|
72
|
+
re.I,
|
|
73
|
+
)
|
|
74
|
+
_RETRYABLE = re.compile(
|
|
75
|
+
"|".join(
|
|
76
|
+
(
|
|
77
|
+
"overloaded",
|
|
78
|
+
"currently experiencing high demand",
|
|
79
|
+
"rate.?limit",
|
|
80
|
+
"too many requests",
|
|
81
|
+
"429",
|
|
82
|
+
"500",
|
|
83
|
+
"502",
|
|
84
|
+
"503",
|
|
85
|
+
"504",
|
|
86
|
+
"520",
|
|
87
|
+
"524",
|
|
88
|
+
"service.?unavailable",
|
|
89
|
+
"server.?error",
|
|
90
|
+
"internal.?error",
|
|
91
|
+
"provider.?returned.?error",
|
|
92
|
+
"exceeded request buffer limit while retrying upstream",
|
|
93
|
+
"network.?error",
|
|
94
|
+
"connection.?error",
|
|
95
|
+
"connection.?refused",
|
|
96
|
+
"connection.?lost",
|
|
97
|
+
"other side closed",
|
|
98
|
+
"fetch failed",
|
|
99
|
+
"getaddrinfo",
|
|
100
|
+
"ENOTFOUND",
|
|
101
|
+
"EAI_AGAIN",
|
|
102
|
+
"upstream.?connect",
|
|
103
|
+
"reset before headers",
|
|
104
|
+
"socket hang up",
|
|
105
|
+
"socket connection was closed",
|
|
106
|
+
"timed? out",
|
|
107
|
+
"timeout",
|
|
108
|
+
"terminated",
|
|
109
|
+
"websocket.?closed",
|
|
110
|
+
"websocket.?error",
|
|
111
|
+
"ended without",
|
|
112
|
+
"stream ended before message_stop",
|
|
113
|
+
"stream ended before a terminal response event",
|
|
114
|
+
"http2 request did not get a response",
|
|
115
|
+
"retry delay",
|
|
116
|
+
"you can retry your request",
|
|
117
|
+
"try your request again",
|
|
118
|
+
"please retry your request",
|
|
119
|
+
"ResourceExhausted",
|
|
120
|
+
"subscription_sharing_usage_unavailable",
|
|
121
|
+
"subscription_sharing_user_unavailable",
|
|
122
|
+
)
|
|
123
|
+
),
|
|
124
|
+
re.I,
|
|
125
|
+
)
|
|
126
|
+
|
|
127
|
+
|
|
128
|
+
# Python's own transport failures (httpx, websockets, the standard library), which the
|
|
129
|
+
# upstream patterns, written for JavaScript runtimes, do not name.
|
|
130
|
+
_PYTHON_TRANSIENT = re.compile(
|
|
131
|
+
"|".join(
|
|
132
|
+
(
|
|
133
|
+
r"\b(?:Connect|Read|Write|Pool)(?:Error|Timeout)\b",
|
|
134
|
+
r"RemoteProtocolError",
|
|
135
|
+
r"ConnectionResetError|ConnectionAbortedError|BrokenPipeError|IncompleteRead",
|
|
136
|
+
r"connection reset",
|
|
137
|
+
r"ConnectionClosed",
|
|
138
|
+
r"server disconnected",
|
|
139
|
+
r"peer closed",
|
|
140
|
+
r"name resolution",
|
|
141
|
+
r"nodename nor servname",
|
|
142
|
+
)
|
|
143
|
+
),
|
|
144
|
+
re.I,
|
|
145
|
+
)
|
|
146
|
+
# This library's HTTP failures read "... (HTTP 503; server) ...": decide by the status,
|
|
147
|
+
# not by digits that happen to appear in the response body.
|
|
148
|
+
_HTTP_STATUS = re.compile(r"\(HTTP (\d{3});")
|
|
149
|
+
_RETRYABLE_STATUS = {408, 429}
|
|
150
|
+
|
|
151
|
+
|
|
152
|
+
def _input_tokens(message: AssistantMessage) -> int:
|
|
153
|
+
usage = message.usage or {}
|
|
154
|
+
return int(usage.get("input") or 0) + int(usage.get("cache_read") or 0)
|
|
155
|
+
|
|
156
|
+
|
|
157
|
+
def is_context_overflow(message: AssistantMessage, context_window: int | None = None) -> bool:
|
|
158
|
+
"""Whether a response failed because the input exceeded the model's context window.
|
|
159
|
+
|
|
160
|
+
Most providers report it as an error message. Pass `context_window` to also catch
|
|
161
|
+
providers that accept an oversized input silently (input usage above the window) or
|
|
162
|
+
truncate it and stop for length with no output.
|
|
163
|
+
"""
|
|
164
|
+
error = message.error or ""
|
|
165
|
+
if message.stop_reason == "error" and error:
|
|
166
|
+
if not any(p.search(error) for p in _NOT_OVERFLOW):
|
|
167
|
+
if any(p.search(error) for p in _OVERFLOW):
|
|
168
|
+
return True
|
|
169
|
+
if message.provider == "cerebras" and _CEREBRAS_OVERFLOW.search(error):
|
|
170
|
+
return True
|
|
171
|
+
if context_window and message.stop_reason == "stop":
|
|
172
|
+
if _input_tokens(message) > context_window:
|
|
173
|
+
return True
|
|
174
|
+
if context_window and message.stop_reason == "length":
|
|
175
|
+
if int((message.usage or {}).get("output") or 0) == 0:
|
|
176
|
+
if _input_tokens(message) >= context_window * 0.99:
|
|
177
|
+
return True
|
|
178
|
+
return False
|
|
179
|
+
|
|
180
|
+
|
|
181
|
+
def is_recoverable_length(message: AssistantMessage, desired_max_output: int) -> bool:
|
|
182
|
+
"""A length stop below the intended output limit, possibly from context pressure."""
|
|
183
|
+
output = int((message.usage or {}).get("output") or 0)
|
|
184
|
+
return (
|
|
185
|
+
message.stop_reason == "length" and desired_max_output > 0 and output < desired_max_output
|
|
186
|
+
)
|
|
187
|
+
|
|
188
|
+
|
|
189
|
+
def is_retryable_error(message: AssistantMessage) -> bool:
|
|
190
|
+
"""Whether a failed response looks transient (overload, rate limit, network, server error).
|
|
191
|
+
|
|
192
|
+
Check `is_context_overflow` first: an overflow needs a smaller context, not a retry.
|
|
193
|
+
Quota and billing limits are never retryable.
|
|
194
|
+
"""
|
|
195
|
+
error = message.error or ""
|
|
196
|
+
if message.stop_reason != "error" or not error:
|
|
197
|
+
return False
|
|
198
|
+
if _NOT_RETRYABLE.search(error):
|
|
199
|
+
return False
|
|
200
|
+
status = _HTTP_STATUS.search(error)
|
|
201
|
+
if status:
|
|
202
|
+
code = int(status.group(1))
|
|
203
|
+
return code >= 500 or code in _RETRYABLE_STATUS
|
|
204
|
+
return bool(_RETRYABLE.search(error) or _PYTHON_TRANSIENT.search(error))
|
|
205
|
+
|
|
206
|
+
|
|
207
|
+
def retry_delay(attempt: int, base: float = 2.0, max_delay: float = 60.0) -> float:
|
|
208
|
+
"""Exponential backoff in seconds for the 1-based retry `attempt`, as in Pi."""
|
|
209
|
+
return min(base * 2 ** max(0, attempt - 1), max_delay)
|