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/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)