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/sync.py
ADDED
|
@@ -0,0 +1,78 @@
|
|
|
1
|
+
"""Blocking entry points for plain scripts.
|
|
2
|
+
|
|
3
|
+
All blocking calls share one private event loop in a daemon thread. Agents, providers and
|
|
4
|
+
their cached connections therefore stay on the same loop from one call to the next, which
|
|
5
|
+
separate ``asyncio.run`` calls would break. Code that already runs an event loop (servers,
|
|
6
|
+
notebooks with top-level await) should await the async API instead.
|
|
7
|
+
"""
|
|
8
|
+
|
|
9
|
+
from __future__ import annotations
|
|
10
|
+
|
|
11
|
+
import asyncio
|
|
12
|
+
import atexit
|
|
13
|
+
import threading
|
|
14
|
+
from collections.abc import Awaitable, Callable
|
|
15
|
+
from typing import Any, TypeVar
|
|
16
|
+
|
|
17
|
+
T = TypeVar("T")
|
|
18
|
+
|
|
19
|
+
_lock = threading.Lock()
|
|
20
|
+
_loop: asyncio.AbstractEventLoop | None = None
|
|
21
|
+
_thread: threading.Thread | None = None
|
|
22
|
+
|
|
23
|
+
|
|
24
|
+
def _portal() -> asyncio.AbstractEventLoop:
|
|
25
|
+
global _loop, _thread
|
|
26
|
+
with _lock:
|
|
27
|
+
if _loop is None or _thread is None or not _thread.is_alive():
|
|
28
|
+
loop = asyncio.new_event_loop()
|
|
29
|
+
thread = threading.Thread(target=loop.run_forever, name="pi-python-sync", daemon=True)
|
|
30
|
+
thread.start()
|
|
31
|
+
_loop, _thread = loop, thread
|
|
32
|
+
return _loop
|
|
33
|
+
|
|
34
|
+
|
|
35
|
+
@atexit.register
|
|
36
|
+
def _shutdown() -> None:
|
|
37
|
+
if _loop is not None and _thread is not None and _thread.is_alive():
|
|
38
|
+
_loop.call_soon_threadsafe(_loop.stop)
|
|
39
|
+
_thread.join(timeout=1)
|
|
40
|
+
|
|
41
|
+
|
|
42
|
+
def run_sync(awaitable: Awaitable[T], *, on_interrupt: Callable[[], Any] | None = None) -> T:
|
|
43
|
+
"""Run an awaitable on the shared loop and block until it finishes.
|
|
44
|
+
|
|
45
|
+
On Ctrl+C, ``on_interrupt`` runs on that loop (for example ``agent.abort``); the call
|
|
46
|
+
waits for the awaitable to settle, then re-raises KeyboardInterrupt. A second Ctrl+C
|
|
47
|
+
stops waiting.
|
|
48
|
+
"""
|
|
49
|
+
try:
|
|
50
|
+
asyncio.get_running_loop()
|
|
51
|
+
except RuntimeError:
|
|
52
|
+
pass
|
|
53
|
+
else:
|
|
54
|
+
close = getattr(awaitable, "close", None)
|
|
55
|
+
if close is not None:
|
|
56
|
+
close() # never awaited; avoid a "coroutine was never awaited" warning
|
|
57
|
+
raise RuntimeError(
|
|
58
|
+
"Blocking call inside a running event loop; await the async API instead "
|
|
59
|
+
"(for example `await agent.prompt(...)`)"
|
|
60
|
+
)
|
|
61
|
+
|
|
62
|
+
async def wrapper() -> T:
|
|
63
|
+
return await awaitable
|
|
64
|
+
|
|
65
|
+
loop = _portal()
|
|
66
|
+
future = asyncio.run_coroutine_threadsafe(wrapper(), loop)
|
|
67
|
+
try:
|
|
68
|
+
return future.result()
|
|
69
|
+
except KeyboardInterrupt:
|
|
70
|
+
if on_interrupt is not None:
|
|
71
|
+
loop.call_soon_threadsafe(on_interrupt)
|
|
72
|
+
else:
|
|
73
|
+
loop.call_soon_threadsafe(future.cancel)
|
|
74
|
+
try:
|
|
75
|
+
future.exception()
|
|
76
|
+
except BaseException:
|
|
77
|
+
pass
|
|
78
|
+
raise
|
pi_python/testing.py
ADDED
|
@@ -0,0 +1,25 @@
|
|
|
1
|
+
"""Deterministic offline provider; records detached request snapshots."""
|
|
2
|
+
|
|
3
|
+
from copy import deepcopy
|
|
4
|
+
from collections.abc import AsyncIterator, Iterable
|
|
5
|
+
from .cancellation import CancelToken
|
|
6
|
+
from .provider import ModelEvent, ModelRequest
|
|
7
|
+
from .messages import AssistantMessage
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
class ScriptedProvider:
|
|
11
|
+
def __init__(self, responses: Iterable[AssistantMessage | list[ModelEvent] | Exception]):
|
|
12
|
+
self.responses = iter(responses)
|
|
13
|
+
self.requests: list[ModelRequest] = []
|
|
14
|
+
|
|
15
|
+
async def stream(self, request: ModelRequest, cancel: CancelToken) -> AsyncIterator[ModelEvent]:
|
|
16
|
+
self.requests.append(deepcopy(request))
|
|
17
|
+
response = next(self.responses, None)
|
|
18
|
+
if response is None:
|
|
19
|
+
raise RuntimeError("ScriptedProvider exhausted")
|
|
20
|
+
if isinstance(response, Exception):
|
|
21
|
+
raise response
|
|
22
|
+
events = [ModelEvent.done(response)] if isinstance(response, AssistantMessage) else response
|
|
23
|
+
for event in events:
|
|
24
|
+
cancel.raise_if_cancelled()
|
|
25
|
+
yield deepcopy(event)
|
pi_python/tools.py
ADDED
|
@@ -0,0 +1,546 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
import asyncio
|
|
3
|
+
import dataclasses
|
|
4
|
+
import datetime
|
|
5
|
+
import enum
|
|
6
|
+
import inspect
|
|
7
|
+
import json
|
|
8
|
+
from contextvars import copy_context
|
|
9
|
+
from copy import copy, deepcopy
|
|
10
|
+
from decimal import Decimal
|
|
11
|
+
from functools import partial
|
|
12
|
+
from urllib.parse import unquote
|
|
13
|
+
from pathlib import PurePath
|
|
14
|
+
from uuid import UUID
|
|
15
|
+
from dataclasses import asdict, dataclass, field
|
|
16
|
+
from typing import Any, Callable, Protocol
|
|
17
|
+
from jsonschema import (
|
|
18
|
+
Draft4Validator,
|
|
19
|
+
Draft6Validator,
|
|
20
|
+
Draft7Validator,
|
|
21
|
+
Draft201909Validator,
|
|
22
|
+
Draft202012Validator,
|
|
23
|
+
SchemaError,
|
|
24
|
+
)
|
|
25
|
+
from .cancellation import CancelToken
|
|
26
|
+
from .errors import ConfigurationError, SubscriptionError, ToolOutcomeUnknownError
|
|
27
|
+
from .messages import (
|
|
28
|
+
ImageContent,
|
|
29
|
+
TextContent,
|
|
30
|
+
ToolCall,
|
|
31
|
+
ToolDeclaration,
|
|
32
|
+
ToolResultMessage,
|
|
33
|
+
_blocks,
|
|
34
|
+
message_to_dict,
|
|
35
|
+
validate_json,
|
|
36
|
+
)
|
|
37
|
+
|
|
38
|
+
# JSON Schema drafts accepted in tool schemas, keyed by $schema without scheme or "#".
|
|
39
|
+
# Pydantic emits 2020-12; MCP servers commonly declare draft-07. No $schema means 2020-12.
|
|
40
|
+
_DRAFTS: dict[str, Any] = {
|
|
41
|
+
"json-schema.org/draft-04/schema": Draft4Validator,
|
|
42
|
+
"json-schema.org/draft-06/schema": Draft6Validator,
|
|
43
|
+
"json-schema.org/draft-07/schema": Draft7Validator,
|
|
44
|
+
"json-schema.org/draft/2019-09/schema": Draft201909Validator,
|
|
45
|
+
"json-schema.org/draft/2020-12/schema": Draft202012Validator,
|
|
46
|
+
"json-schema.org/schema": Draft202012Validator, # "the latest draft"
|
|
47
|
+
}
|
|
48
|
+
# Where subschemas live; other keywords hold data or vendor extensions and are not followed.
|
|
49
|
+
_SUBSCHEMA = {
|
|
50
|
+
"items",
|
|
51
|
+
"additionalItems",
|
|
52
|
+
"contains",
|
|
53
|
+
"additionalProperties",
|
|
54
|
+
"propertyNames",
|
|
55
|
+
"unevaluatedItems",
|
|
56
|
+
"unevaluatedProperties",
|
|
57
|
+
"not",
|
|
58
|
+
"if",
|
|
59
|
+
"then",
|
|
60
|
+
"else",
|
|
61
|
+
"allOf",
|
|
62
|
+
"anyOf",
|
|
63
|
+
"oneOf",
|
|
64
|
+
"prefixItems",
|
|
65
|
+
}
|
|
66
|
+
_SCHEMA_MAPS = {
|
|
67
|
+
"properties",
|
|
68
|
+
"patternProperties",
|
|
69
|
+
"$defs",
|
|
70
|
+
"definitions",
|
|
71
|
+
"dependentSchemas",
|
|
72
|
+
"dependencies",
|
|
73
|
+
}
|
|
74
|
+
|
|
75
|
+
|
|
76
|
+
def schema_validator(schema: dict[str, Any]) -> Any:
|
|
77
|
+
"""A validator for the schema's draft that also checks `format`, like Pi's validator.
|
|
78
|
+
|
|
79
|
+
Formats are checked when jsonschema can check them; some (for example `uri`)
|
|
80
|
+
need jsonschema's optional format dependencies.
|
|
81
|
+
"""
|
|
82
|
+
uri = schema.get("$schema") if isinstance(schema, dict) else None
|
|
83
|
+
cls: Any = Draft202012Validator
|
|
84
|
+
if uri is not None:
|
|
85
|
+
cls = _DRAFTS.get(str(uri).split("://", 1)[-1].rstrip("#"))
|
|
86
|
+
if cls is None:
|
|
87
|
+
raise ConfigurationError(f"Unsupported $schema: {uri}")
|
|
88
|
+
return cls(schema, format_checker=cls.FORMAT_CHECKER)
|
|
89
|
+
|
|
90
|
+
|
|
91
|
+
def _without_empty_required(node: Any) -> Any:
|
|
92
|
+
"""Draft-04 forbids `required: []`, which many generators emit; it means nothing."""
|
|
93
|
+
if isinstance(node, list):
|
|
94
|
+
return [_without_empty_required(item) for item in node]
|
|
95
|
+
if not isinstance(node, dict):
|
|
96
|
+
return node
|
|
97
|
+
return {
|
|
98
|
+
k: _without_empty_required(v) for k, v in node.items() if not (k == "required" and v == [])
|
|
99
|
+
}
|
|
100
|
+
|
|
101
|
+
|
|
102
|
+
def _resolve(schema: dict[str, Any], ref: str) -> None:
|
|
103
|
+
"""Check that a local reference points somewhere: "#", "#/json/pointer" or "#anchor"."""
|
|
104
|
+
fragment = unquote(ref[1:])
|
|
105
|
+
if fragment and not fragment.startswith("/"): # a plain-name anchor
|
|
106
|
+
found = []
|
|
107
|
+
|
|
108
|
+
def find(node: Any) -> None:
|
|
109
|
+
if isinstance(node, dict):
|
|
110
|
+
if node.get("$anchor") == fragment or node.get("$id") == f"#{fragment}":
|
|
111
|
+
found.append(node)
|
|
112
|
+
for value in node.values():
|
|
113
|
+
find(value)
|
|
114
|
+
elif isinstance(node, list):
|
|
115
|
+
for value in node:
|
|
116
|
+
find(value)
|
|
117
|
+
|
|
118
|
+
find(schema)
|
|
119
|
+
if not found:
|
|
120
|
+
raise ConfigurationError(f"Unresolved local schema reference: {ref}")
|
|
121
|
+
return
|
|
122
|
+
target: Any = schema
|
|
123
|
+
try:
|
|
124
|
+
for part in fragment.split("/")[1:] if fragment else []:
|
|
125
|
+
part = part.replace("~1", "/").replace("~0", "~")
|
|
126
|
+
target = target[int(part)] if isinstance(target, list) else target[part]
|
|
127
|
+
except (KeyError, TypeError, ValueError, IndexError) as exc:
|
|
128
|
+
raise ConfigurationError(f"Unresolved local schema reference: {ref}") from exc
|
|
129
|
+
|
|
130
|
+
|
|
131
|
+
# Counted without recursion before anything recursive walks the schema. Relying on
|
|
132
|
+
# RecursionError alone gave interpreter-dependent limits, and PyPy can crash first.
|
|
133
|
+
_MAX_SCHEMA_DEPTH = 100
|
|
134
|
+
|
|
135
|
+
|
|
136
|
+
def _deeper_than(value: Any, limit: int) -> bool:
|
|
137
|
+
"""Whether objects and arrays nest more than `limit` levels; also ends on cycles."""
|
|
138
|
+
stack = [(value, 1)]
|
|
139
|
+
while stack:
|
|
140
|
+
node, depth = stack.pop()
|
|
141
|
+
if isinstance(node, dict):
|
|
142
|
+
children: Any = node.values()
|
|
143
|
+
elif isinstance(node, list):
|
|
144
|
+
children = node
|
|
145
|
+
else:
|
|
146
|
+
continue
|
|
147
|
+
if depth > limit:
|
|
148
|
+
return True
|
|
149
|
+
stack.extend((child, depth + 1) for child in children)
|
|
150
|
+
return False
|
|
151
|
+
|
|
152
|
+
|
|
153
|
+
def validate_schema(schema: dict[str, Any]) -> None:
|
|
154
|
+
"""Accept standard JSON Schema; reject what cannot be checked without leaving the schema."""
|
|
155
|
+
if _deeper_than(schema, _MAX_SCHEMA_DEPTH):
|
|
156
|
+
raise ConfigurationError(
|
|
157
|
+
f"Tool schema is nested too deeply (more than {_MAX_SCHEMA_DEPTH} levels)"
|
|
158
|
+
)
|
|
159
|
+
try:
|
|
160
|
+
validate_json(schema)
|
|
161
|
+
except RecursionError as exc:
|
|
162
|
+
raise ConfigurationError("Tool schema is nested too deeply") from exc
|
|
163
|
+
if not isinstance(schema, dict):
|
|
164
|
+
raise ConfigurationError("Tool schema must be an object")
|
|
165
|
+
validator = schema_validator(schema)
|
|
166
|
+
try:
|
|
167
|
+
if isinstance(validator, Draft4Validator):
|
|
168
|
+
validator.check_schema(_without_empty_required(schema))
|
|
169
|
+
else:
|
|
170
|
+
validator.check_schema(schema)
|
|
171
|
+
except SchemaError as exc:
|
|
172
|
+
raise ConfigurationError(exc.message) from exc
|
|
173
|
+
except RecursionError as exc:
|
|
174
|
+
raise ConfigurationError("Tool schema is nested too deeply") from exc
|
|
175
|
+
|
|
176
|
+
# Only local references, resolved at registration: never a network or file fetch.
|
|
177
|
+
def visit(node: Any) -> None:
|
|
178
|
+
if isinstance(node, list):
|
|
179
|
+
for item in node:
|
|
180
|
+
visit(item)
|
|
181
|
+
return
|
|
182
|
+
if not isinstance(node, dict):
|
|
183
|
+
return
|
|
184
|
+
for key, value in node.items():
|
|
185
|
+
if key in {"$ref", "$dynamicRef", "$recursiveRef"}:
|
|
186
|
+
if not isinstance(value, str) or not value.startswith("#"):
|
|
187
|
+
raise ConfigurationError("Only local schema references are supported")
|
|
188
|
+
_resolve(schema, value)
|
|
189
|
+
elif key in _SCHEMA_MAPS and isinstance(value, dict):
|
|
190
|
+
for sub in value.values():
|
|
191
|
+
visit(sub)
|
|
192
|
+
elif key in _SUBSCHEMA:
|
|
193
|
+
visit(value)
|
|
194
|
+
|
|
195
|
+
try:
|
|
196
|
+
visit(schema)
|
|
197
|
+
except RecursionError as exc:
|
|
198
|
+
raise ConfigurationError("Tool schema is nested too deeply") from exc
|
|
199
|
+
|
|
200
|
+
|
|
201
|
+
@dataclass
|
|
202
|
+
class ToolResult:
|
|
203
|
+
content: list[TextContent | ImageContent]
|
|
204
|
+
details: Any = None
|
|
205
|
+
structured_content: Any = None
|
|
206
|
+
is_error: bool = False
|
|
207
|
+
terminate: bool = False
|
|
208
|
+
error_code: str | None = None
|
|
209
|
+
usage: dict[str, Any] | None = None
|
|
210
|
+
nested_calls: dict[str, Any] | None = None
|
|
211
|
+
|
|
212
|
+
@classmethod
|
|
213
|
+
def text(cls, text: str, **kwargs: Any) -> ToolResult:
|
|
214
|
+
return cls([TextContent(text)], **kwargs)
|
|
215
|
+
|
|
216
|
+
|
|
217
|
+
def error_result(code: str, text: str) -> ToolResult:
|
|
218
|
+
return ToolResult.text(text, is_error=True, error_code=code)
|
|
219
|
+
|
|
220
|
+
|
|
221
|
+
def aborted_result() -> ToolResult:
|
|
222
|
+
"""Pi's result for a call stopped by abort; the text matches upstream."""
|
|
223
|
+
return error_result("aborted", "Operation aborted")
|
|
224
|
+
|
|
225
|
+
|
|
226
|
+
def _plain(value: Any) -> Any:
|
|
227
|
+
"""Common Python values as JSON values: dates, UUIDs, paths, enums, tuples, sets,
|
|
228
|
+
dataclasses and pydantic models, also nested inside dicts and lists."""
|
|
229
|
+
if isinstance(value, enum.Enum):
|
|
230
|
+
return _plain(value.value)
|
|
231
|
+
if isinstance(value, (datetime.date, datetime.time)):
|
|
232
|
+
return value.isoformat()
|
|
233
|
+
if isinstance(value, (UUID, PurePath, Decimal)):
|
|
234
|
+
return str(value)
|
|
235
|
+
if hasattr(value, "model_dump") and not isinstance(value, type):
|
|
236
|
+
return value.model_dump(mode="json")
|
|
237
|
+
if dataclasses.is_dataclass(value) and not isinstance(value, type):
|
|
238
|
+
return _plain(dataclasses.asdict(value))
|
|
239
|
+
if isinstance(value, dict):
|
|
240
|
+
return {str(k): _plain(v) for k, v in value.items()}
|
|
241
|
+
if isinstance(value, (list, tuple, set, frozenset)):
|
|
242
|
+
return [_plain(v) for v in value]
|
|
243
|
+
return value
|
|
244
|
+
|
|
245
|
+
|
|
246
|
+
def as_tool_result(value: Any) -> Any:
|
|
247
|
+
"""Accept what a plain Python function naturally returns.
|
|
248
|
+
|
|
249
|
+
A string becomes text and None empty content. Other values become their JSON text and
|
|
250
|
+
the structured result; dates, UUIDs, enums, tuples, dataclasses and pydantic models are
|
|
251
|
+
converted first. Anything else is left for check_result to reject.
|
|
252
|
+
"""
|
|
253
|
+
if isinstance(value, ToolResult):
|
|
254
|
+
return value
|
|
255
|
+
if value is None:
|
|
256
|
+
return ToolResult([])
|
|
257
|
+
value = _plain(value)
|
|
258
|
+
if isinstance(value, str):
|
|
259
|
+
return ToolResult.text(value)
|
|
260
|
+
if isinstance(value, (dict, list, int, float)):
|
|
261
|
+
validate_json(value)
|
|
262
|
+
return ToolResult.text(json.dumps(value, ensure_ascii=False), structured_content=value)
|
|
263
|
+
return value
|
|
264
|
+
|
|
265
|
+
|
|
266
|
+
async def call_tool_function(function: Callable, *args: Any) -> Any:
|
|
267
|
+
"""Await an async tool; run a sync one in a worker thread, off the event loop."""
|
|
268
|
+
if inspect.iscoroutinefunction(function) or inspect.iscoroutinefunction(
|
|
269
|
+
getattr(function, "__call__", None)
|
|
270
|
+
):
|
|
271
|
+
return await function(*args)
|
|
272
|
+
future = asyncio.get_running_loop().run_in_executor(
|
|
273
|
+
None, partial(copy_context().run, function, *args)
|
|
274
|
+
)
|
|
275
|
+
# A thread cannot be interrupted. On cancellation keep waiting: if it returns, its
|
|
276
|
+
# real result stands, as when a Pi tool ignores the abort signal; while it runs,
|
|
277
|
+
# cleanup sees the work as unfinished.
|
|
278
|
+
swallowed = 0
|
|
279
|
+
while True:
|
|
280
|
+
try:
|
|
281
|
+
value = await asyncio.shield(future)
|
|
282
|
+
break
|
|
283
|
+
except asyncio.CancelledError:
|
|
284
|
+
if future.cancelled():
|
|
285
|
+
raise
|
|
286
|
+
swallowed += 1
|
|
287
|
+
task = asyncio.current_task()
|
|
288
|
+
for _ in range(swallowed if task is not None else 0):
|
|
289
|
+
task.uncancel() # type: ignore[union-attr]
|
|
290
|
+
return await value if inspect.isawaitable(value) else value
|
|
291
|
+
|
|
292
|
+
|
|
293
|
+
def check_result(result: ToolResult, schema: dict[str, Any] | None = None) -> None:
|
|
294
|
+
if not isinstance(result, ToolResult):
|
|
295
|
+
raise TypeError("Tool must return ToolResult")
|
|
296
|
+
validate_json(asdict(result))
|
|
297
|
+
_blocks([asdict(b) for b in result.content], images=True)
|
|
298
|
+
message_to_dict(
|
|
299
|
+
ToolResultMessage(
|
|
300
|
+
"validation",
|
|
301
|
+
"validation",
|
|
302
|
+
result.content,
|
|
303
|
+
usage=result.usage,
|
|
304
|
+
nested_calls=result.nested_calls,
|
|
305
|
+
)
|
|
306
|
+
)
|
|
307
|
+
if type(result.is_error) is not bool or type(result.terminate) is not bool:
|
|
308
|
+
raise TypeError("Tool result flags must be boolean")
|
|
309
|
+
if schema is not None and not result.is_error:
|
|
310
|
+
schema_validator(schema).validate(result.structured_content)
|
|
311
|
+
|
|
312
|
+
|
|
313
|
+
@dataclass
|
|
314
|
+
class ToolResultUpdate:
|
|
315
|
+
content: list[TextContent | ImageContent] | None = None
|
|
316
|
+
details: Any = None
|
|
317
|
+
structured_content: Any = None
|
|
318
|
+
is_error: bool | None = None
|
|
319
|
+
terminate: bool | None = None
|
|
320
|
+
usage: dict[str, Any] | None = None
|
|
321
|
+
nested_calls: dict[str, Any] | None = None
|
|
322
|
+
|
|
323
|
+
|
|
324
|
+
@dataclass
|
|
325
|
+
class ToolContext:
|
|
326
|
+
run_id: str
|
|
327
|
+
call_id: str
|
|
328
|
+
cancel: CancelToken
|
|
329
|
+
_emit: Callable = field(repr=False)
|
|
330
|
+
assistant_message: Any = None
|
|
331
|
+
agent_context: Any = None
|
|
332
|
+
tool_call: ToolCall | None = None
|
|
333
|
+
args: dict[str, Any] | None = None
|
|
334
|
+
result: ToolResult | None = None
|
|
335
|
+
is_error: bool = False
|
|
336
|
+
|
|
337
|
+
async def emit_update(self, value: Any) -> None:
|
|
338
|
+
validate_json(value)
|
|
339
|
+
await self._emit(value)
|
|
340
|
+
|
|
341
|
+
|
|
342
|
+
class ToolExecutor(Protocol):
|
|
343
|
+
async def __call__(self, args: dict[str, Any], context: ToolContext) -> ToolResult: ...
|
|
344
|
+
|
|
345
|
+
|
|
346
|
+
@dataclass
|
|
347
|
+
class Tool:
|
|
348
|
+
name: str
|
|
349
|
+
description: str
|
|
350
|
+
input_schema: dict[str, Any]
|
|
351
|
+
execute: ToolExecutor
|
|
352
|
+
output_schema: dict[str, Any] | None = None
|
|
353
|
+
execution_mode: str = "parallel"
|
|
354
|
+
prepare_arguments: Callable | None = None
|
|
355
|
+
|
|
356
|
+
def __post_init__(self) -> None:
|
|
357
|
+
if not isinstance(self.name, str) or not self.name or not isinstance(self.description, str):
|
|
358
|
+
raise ConfigurationError("Invalid tool name or description")
|
|
359
|
+
if self.execution_mode not in {"parallel", "sequential"}:
|
|
360
|
+
raise ConfigurationError("Invalid tool execution mode")
|
|
361
|
+
if not callable(self.execute):
|
|
362
|
+
raise ConfigurationError("Tool execute must be callable")
|
|
363
|
+
validate_schema(self.input_schema)
|
|
364
|
+
if self.output_schema is not None:
|
|
365
|
+
validate_schema(self.output_schema)
|
|
366
|
+
|
|
367
|
+
def __deepcopy__(self, memo: dict[int, Any]) -> Tool:
|
|
368
|
+
# Executable callables can own clients, locks or other application resources.
|
|
369
|
+
# Copy declarations, never clone the callable's owner or its connections.
|
|
370
|
+
result = copy(self)
|
|
371
|
+
memo[id(self)] = result
|
|
372
|
+
result.input_schema = deepcopy(self.input_schema, memo)
|
|
373
|
+
result.output_schema = deepcopy(self.output_schema, memo)
|
|
374
|
+
return result
|
|
375
|
+
|
|
376
|
+
def declaration(self) -> ToolDeclaration:
|
|
377
|
+
return ToolDeclaration(self.name, self.description, deepcopy(self.input_schema))
|
|
378
|
+
|
|
379
|
+
|
|
380
|
+
@dataclass
|
|
381
|
+
class ToolOutcome:
|
|
382
|
+
call: ToolCall
|
|
383
|
+
execution_status: str = "not_started"
|
|
384
|
+
original_arguments: dict[str, Any] = field(default_factory=dict)
|
|
385
|
+
prepared_arguments: dict[str, Any] | None = None
|
|
386
|
+
raw_result: ToolResult | None = None
|
|
387
|
+
result: ToolResult | None = None
|
|
388
|
+
error: str | None = None
|
|
389
|
+
_settled: bool = field(default=False, repr=False)
|
|
390
|
+
|
|
391
|
+
|
|
392
|
+
async def invoke(function: Callable | None, *args: Any) -> Any:
|
|
393
|
+
if function is None:
|
|
394
|
+
return None
|
|
395
|
+
value = function(*args)
|
|
396
|
+
return await value if inspect.isawaitable(value) else value
|
|
397
|
+
|
|
398
|
+
|
|
399
|
+
async def prepare_tool_call(
|
|
400
|
+
tool: Tool | None, outcome: ToolOutcome, context: ToolContext, before: Callable | None = None
|
|
401
|
+
) -> None:
|
|
402
|
+
if tool is None:
|
|
403
|
+
outcome.result = error_result("unknown_tool", f"Unknown tool: {outcome.call.name}")
|
|
404
|
+
return
|
|
405
|
+
try:
|
|
406
|
+
args = deepcopy(outcome.original_arguments)
|
|
407
|
+
if tool.prepare_arguments:
|
|
408
|
+
args = await invoke(tool.prepare_arguments, args)
|
|
409
|
+
validate_json(args)
|
|
410
|
+
if not isinstance(args, dict):
|
|
411
|
+
raise ValueError("Tool arguments must be an object")
|
|
412
|
+
outcome.prepared_arguments = deepcopy(args)
|
|
413
|
+
# Name each problem by its path; the full schema would only bloat the model's context.
|
|
414
|
+
problems = [
|
|
415
|
+
f"{error.json_path}: {error.message}"
|
|
416
|
+
for error in schema_validator(tool.input_schema).iter_errors(args)
|
|
417
|
+
]
|
|
418
|
+
if problems:
|
|
419
|
+
more = f" (and {len(problems) - 5} more)" if len(problems) > 5 else ""
|
|
420
|
+
raise ValueError("; ".join(problems[:5]) + more)
|
|
421
|
+
except Exception as exc:
|
|
422
|
+
outcome.result = error_result("invalid_arguments", f"{type(exc).__name__}: {exc}")
|
|
423
|
+
return
|
|
424
|
+
context.args = deepcopy(args)
|
|
425
|
+
context.tool_call = deepcopy(outcome.call)
|
|
426
|
+
try:
|
|
427
|
+
decision = await invoke(before, deepcopy(outcome.call), deepcopy(args), context)
|
|
428
|
+
except SubscriptionError:
|
|
429
|
+
raise
|
|
430
|
+
except Exception as exc:
|
|
431
|
+
# As in Pi, a failing preflight hook fails this call, not the whole run.
|
|
432
|
+
outcome.error = f"{type(exc).__name__}: {exc}"
|
|
433
|
+
outcome.result = error_result("hook_error", outcome.error)
|
|
434
|
+
return
|
|
435
|
+
if decision is False:
|
|
436
|
+
outcome.result = error_result("blocked", "Tool call blocked by before_tool_call")
|
|
437
|
+
elif isinstance(decision, ToolResult):
|
|
438
|
+
check_result(decision)
|
|
439
|
+
outcome.result = deepcopy(decision)
|
|
440
|
+
elif decision is not None and decision is not True:
|
|
441
|
+
raise ConfigurationError("before_tool_call must return bool, ToolResult or None")
|
|
442
|
+
|
|
443
|
+
|
|
444
|
+
async def run_tool_call(
|
|
445
|
+
tool: Tool | None,
|
|
446
|
+
call: ToolCall,
|
|
447
|
+
context: ToolContext,
|
|
448
|
+
*,
|
|
449
|
+
before_tool_call: Callable | None = None,
|
|
450
|
+
after_tool_call: Callable | None = None,
|
|
451
|
+
outcome: ToolOutcome | None = None,
|
|
452
|
+
prepared: bool = False,
|
|
453
|
+
) -> ToolOutcome:
|
|
454
|
+
"""Shared programmatic/agent path. Caller cancellation is never swallowed."""
|
|
455
|
+
outcome = outcome or ToolOutcome(deepcopy(call), original_arguments=deepcopy(call.arguments))
|
|
456
|
+
if not prepared:
|
|
457
|
+
await prepare_tool_call(tool, outcome, context, before_tool_call)
|
|
458
|
+
if outcome.result is not None:
|
|
459
|
+
return outcome
|
|
460
|
+
assert tool is not None
|
|
461
|
+
context.cancel.raise_if_cancelled()
|
|
462
|
+
outcome.execution_status = "running"
|
|
463
|
+
try:
|
|
464
|
+
assert outcome.prepared_arguments is not None
|
|
465
|
+
value = await call_tool_function(
|
|
466
|
+
tool.execute, deepcopy(outcome.prepared_arguments), context
|
|
467
|
+
)
|
|
468
|
+
if outcome._settled:
|
|
469
|
+
return outcome
|
|
470
|
+
# Execution succeeded even if conversion, serializability or output validation fails.
|
|
471
|
+
outcome.execution_status = "succeeded"
|
|
472
|
+
try:
|
|
473
|
+
outcome.raw_result = deepcopy(as_tool_result(value))
|
|
474
|
+
except Exception as exc:
|
|
475
|
+
outcome.error = f"{type(exc).__name__}: {exc}"
|
|
476
|
+
outcome.result = error_result("finalization_error", outcome.error)
|
|
477
|
+
return outcome
|
|
478
|
+
except asyncio.CancelledError:
|
|
479
|
+
if outcome._settled:
|
|
480
|
+
raise
|
|
481
|
+
# Pi's abort: the tool saw the signal and stopped; record an ordinary error result.
|
|
482
|
+
outcome.execution_status = "cancelled"
|
|
483
|
+
outcome.error = "Operation aborted"
|
|
484
|
+
outcome.result = aborted_result()
|
|
485
|
+
raise
|
|
486
|
+
except ToolOutcomeUnknownError as exc:
|
|
487
|
+
if outcome._settled:
|
|
488
|
+
return outcome
|
|
489
|
+
outcome.execution_status = "unknown"
|
|
490
|
+
outcome.error = str(exc)
|
|
491
|
+
outcome.result = error_result("outcome_unknown", str(exc))
|
|
492
|
+
return outcome
|
|
493
|
+
except SubscriptionError:
|
|
494
|
+
raise
|
|
495
|
+
except Exception as exc:
|
|
496
|
+
if outcome._settled:
|
|
497
|
+
return outcome
|
|
498
|
+
outcome.execution_status = "failed"
|
|
499
|
+
outcome.error = f"{type(exc).__name__}: {exc}"
|
|
500
|
+
outcome.result = error_result("tool_error", outcome.error)
|
|
501
|
+
outcome.raw_result = deepcopy(outcome.result)
|
|
502
|
+
try:
|
|
503
|
+
result = deepcopy(outcome.raw_result)
|
|
504
|
+
check_result(result, tool.output_schema)
|
|
505
|
+
context.args = deepcopy(outcome.prepared_arguments)
|
|
506
|
+
context.tool_call = deepcopy(call)
|
|
507
|
+
context.result = deepcopy(result)
|
|
508
|
+
context.is_error = result.is_error
|
|
509
|
+
replacement = await invoke(after_tool_call, deepcopy(call), deepcopy(result), context)
|
|
510
|
+
if outcome._settled:
|
|
511
|
+
return outcome
|
|
512
|
+
if isinstance(replacement, ToolResultUpdate):
|
|
513
|
+
if replacement.content is not None:
|
|
514
|
+
result.content = replacement.content
|
|
515
|
+
result.structured_content = replacement.structured_content
|
|
516
|
+
elif replacement.structured_content is not None:
|
|
517
|
+
result.structured_content = replacement.structured_content
|
|
518
|
+
if replacement.details is not None:
|
|
519
|
+
result.details = replacement.details
|
|
520
|
+
if replacement.is_error is not None:
|
|
521
|
+
result.is_error = replacement.is_error
|
|
522
|
+
for metadata_key in ("usage", "nested_calls"):
|
|
523
|
+
if getattr(replacement, metadata_key) is not None:
|
|
524
|
+
setattr(result, metadata_key, deepcopy(getattr(replacement, metadata_key)))
|
|
525
|
+
if replacement.terminate is not None:
|
|
526
|
+
result.terminate = replacement.terminate
|
|
527
|
+
elif replacement is not None:
|
|
528
|
+
if not isinstance(replacement, ToolResult):
|
|
529
|
+
raise TypeError("after_tool_call must return ToolResult, ToolResultUpdate or None")
|
|
530
|
+
result = replacement
|
|
531
|
+
check_result(result, tool.output_schema)
|
|
532
|
+
outcome.result = deepcopy(result)
|
|
533
|
+
except asyncio.CancelledError:
|
|
534
|
+
if outcome._settled:
|
|
535
|
+
raise
|
|
536
|
+
outcome.result = error_result(
|
|
537
|
+
"finalization_cancelled", "Execution finished; result finalization cancelled"
|
|
538
|
+
)
|
|
539
|
+
raise
|
|
540
|
+
except SubscriptionError:
|
|
541
|
+
raise
|
|
542
|
+
except Exception as exc:
|
|
543
|
+
if not outcome._settled:
|
|
544
|
+
outcome.error = f"{type(exc).__name__}: {exc}"
|
|
545
|
+
outcome.result = error_result("finalization_error", outcome.error)
|
|
546
|
+
return outcome
|