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,201 @@
1
+ from __future__ import annotations
2
+ from collections.abc import Callable
3
+ from typing import Any
4
+ from ..cancellation import CancelToken
5
+ from ..messages import Message
6
+ from ..models import ModelCatalog
7
+
8
+ import time
9
+ from copy import deepcopy
10
+ from ..errors import ConfigurationError
11
+ from ..models import ModelInfo
12
+ from ..provider import ModelRequest
13
+ from ..messages import (
14
+ AssistantMessage,
15
+ SystemMessage,
16
+ TextContent,
17
+ ThinkingContent,
18
+ ToolCall,
19
+ ToolResultMessage,
20
+ UserMessage,
21
+ )
22
+ from ..tools import invoke
23
+ from .transport import HTTPTransport
24
+
25
+
26
+ def same_model(message: AssistantMessage, provider: str, api: str, model: str) -> bool:
27
+ return message.provider == provider and message.api == api and message.model == model
28
+
29
+
30
+ def transform_messages(
31
+ messages: list[Message],
32
+ provider: str,
33
+ api: str,
34
+ model: str,
35
+ normalize_id: Callable[[str, AssistantMessage], str] | None = None,
36
+ ) -> list[Message]:
37
+ """Pi transformMessages: replay-safe history for one target model.
38
+
39
+ Cross-model thinking becomes text and opaque signatures are dropped. Failed or
40
+ aborted assistant turns are not replayed. A system message between a tool call
41
+ and its results moves after them. Python history validation already rejects
42
+ dangling calls; the synthetic result only covers hand-built requests.
43
+ """
44
+ ids: dict[str, str] = {}
45
+ first = []
46
+ for message in deepcopy(messages):
47
+ if isinstance(message, ToolResultMessage):
48
+ message.call_id = ids.get(message.call_id, message.call_id)
49
+ elif isinstance(message, AssistantMessage):
50
+ same = same_model(message, provider, api, model)
51
+ content: list[TextContent | ThinkingContent | ToolCall] = []
52
+ for block in message.content:
53
+ if isinstance(block, ThinkingContent):
54
+ if block.redacted:
55
+ if same:
56
+ content.append(block)
57
+ elif same and block.thinking_signature:
58
+ content.append(block)
59
+ elif block.thinking.strip():
60
+ content.append(block if same else TextContent(block.thinking))
61
+ elif isinstance(block, TextContent):
62
+ content.append(block if same else TextContent(block.text))
63
+ elif isinstance(block, ToolCall):
64
+ if not same:
65
+ block.thought_signature = None
66
+ if normalize_id is not None:
67
+ new = normalize_id(block.id, message)
68
+ if new != block.id:
69
+ ids[block.id] = new
70
+ block.id = new
71
+ content.append(block)
72
+ message.content = content
73
+ first.append(message)
74
+ result: list = []
75
+ pending: list[ToolCall] = []
76
+ answered: set[str] = set()
77
+ held: list[SystemMessage] = []
78
+
79
+ def close() -> None:
80
+ nonlocal pending, answered
81
+ for call in pending:
82
+ if call.id not in answered:
83
+ result.append(
84
+ ToolResultMessage(call.id, call.name, [TextContent("No result provided")], True)
85
+ )
86
+ pending, answered = [], set()
87
+ result.extend(held)
88
+ held.clear()
89
+
90
+ for message in first:
91
+ if isinstance(message, AssistantMessage):
92
+ close()
93
+ if message.stop_reason in {"error", "aborted"}:
94
+ continue
95
+ if message.tool_calls:
96
+ pending, answered = message.tool_calls, set()
97
+ result.append(message)
98
+ elif isinstance(message, ToolResultMessage):
99
+ answered.add(message.call_id)
100
+ result.append(message)
101
+ elif isinstance(message, SystemMessage) and pending:
102
+ held.append(message)
103
+ else:
104
+ if isinstance(message, UserMessage):
105
+ close()
106
+ result.append(message)
107
+ close()
108
+ return result
109
+
110
+
111
+ class RemoteProvider:
112
+ name = "remote"
113
+
114
+ def __init__(
115
+ self,
116
+ *,
117
+ api_key: str | Callable[[], Any] | None = None,
118
+ credentials: Any = None,
119
+ transport: HTTPTransport | None = None,
120
+ catalog: ModelCatalog | None = None,
121
+ ) -> None:
122
+ if api_key is not None and credentials is not None:
123
+ raise ConfigurationError("Pass api_key or credentials")
124
+ self.api_key = api_key
125
+ self.credentials = credentials
126
+ self.transport = transport or HTTPTransport()
127
+ self.catalog = catalog if catalog is not None else ModelCatalog.bundled()
128
+
129
+ def model_info(self, request: ModelRequest) -> ModelInfo:
130
+ """The request's model record, or the catalog's; an unknown model is an error."""
131
+ model = request.model_info or self.catalog.get(self.name, request.model)
132
+ if model is None:
133
+ raise ConfigurationError(
134
+ f"Unknown model {self.name}/{request.model}: pass a ModelInfo as the model "
135
+ "or register one in the provider catalog"
136
+ )
137
+ if model.provider != self.name or model.id != request.model:
138
+ raise ConfigurationError(f"Model {model.provider}/{model.id} is not {self.name}")
139
+ model.validate_request(request)
140
+ return deepcopy(model)
141
+
142
+ async def aclose(self) -> None:
143
+ await self.transport.aclose()
144
+
145
+ async def credential(self, request: ModelRequest, cancel: CancelToken) -> str:
146
+ key = request.api_key
147
+ if key is None and self.credentials is not None:
148
+ credential = (
149
+ await self.credentials.get(cancel)
150
+ if hasattr(self.credentials, "get")
151
+ else self.credentials
152
+ )
153
+ expected = "openai-chatgpt" if self.name == "openai" else self.name
154
+ if credential.provider != expected:
155
+ raise ConfigurationError("OAuth credential belongs to a different provider")
156
+ if (
157
+ expected == "openai-chatgpt"
158
+ and "chatgpt.tokens.use.direct" not in credential.scopes
159
+ ):
160
+ raise ConfigurationError("OAuth grant lacks chatgpt.tokens.use.direct")
161
+ if credential.expires_at <= time.time():
162
+ raise ConfigurationError("OAuth credential expired; use RefreshingCredentials")
163
+ key = credential.access_token
164
+ if key is None:
165
+ key = await invoke(self.api_key) if callable(self.api_key) else self.api_key
166
+ if not isinstance(key, str) or not key:
167
+ raise ConfigurationError(f"Missing credentials for {self.name}")
168
+ return key
169
+
170
+ async def payload(self, request: ModelRequest, body: dict[str, Any]) -> dict[str, Any]:
171
+ payload = deepcopy(body)
172
+ replacement = await invoke(request.on_payload, payload)
173
+ if replacement is not None:
174
+ payload = replacement
175
+ if not isinstance(payload, dict):
176
+ raise ConfigurationError("on_payload must return a dict or None")
177
+ return payload
178
+
179
+
180
+ def normalize_usage(raw: dict[str, Any], provider: str) -> dict[str, int]:
181
+ """Pi token accounting in snake_case. Prices are not inferred from model names."""
182
+ if provider == "anthropic":
183
+ read = raw.get("cache_read_input_tokens", 0)
184
+ write = raw.get("cache_creation_input_tokens", 0)
185
+ input_tokens = raw.get("input_tokens", 0)
186
+ reasoning = 0
187
+ else:
188
+ details = raw.get("input_tokens_details") or {}
189
+ read = details.get("cached_tokens", 0)
190
+ write = details.get("cache_write_tokens", 0)
191
+ input_tokens = max(0, raw.get("input_tokens", 0) - read - write)
192
+ reasoning = (raw.get("output_tokens_details") or {}).get("reasoning_tokens", 0)
193
+ output = raw.get("output_tokens", 0)
194
+ return {
195
+ "input": input_tokens,
196
+ "output": output,
197
+ "cache_read": read,
198
+ "cache_write": write,
199
+ "reasoning": reasoning,
200
+ "total_tokens": input_tokens + output + read + write,
201
+ }