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,438 @@
1
+ """Build a Tool from a typed Python function: the schema comes from the signature and
2
+ the description from the docstring, so a tool is written like any other function."""
3
+
4
+ from __future__ import annotations
5
+
6
+ import dataclasses
7
+ import enum
8
+ import functools
9
+ import inspect
10
+ import re
11
+ import types
12
+ import typing
13
+ from collections.abc import Callable, Mapping, Sequence
14
+ from datetime import date, datetime
15
+ from pathlib import Path
16
+ from typing import Any, Literal, Union
17
+ from uuid import UUID
18
+
19
+ from .errors import ConfigurationError
20
+ from .messages import validate_json
21
+ from .tools import Tool, ToolContext
22
+
23
+ _PRIMITIVES: dict[Any, dict[str, Any]] = {
24
+ str: {"type": "string"},
25
+ int: {"type": "integer"},
26
+ float: {"type": "number"},
27
+ bool: {"type": "boolean"},
28
+ type(None): {"type": "null"},
29
+ datetime: {"type": "string", "format": "date-time"},
30
+ date: {"type": "string", "format": "date"},
31
+ UUID: {"type": "string", "format": "uuid"},
32
+ Path: {"type": "string"},
33
+ }
34
+ _PARAM_SECTIONS = {
35
+ "args",
36
+ "arguments",
37
+ "parameters",
38
+ "params",
39
+ "keyword args",
40
+ "keyword arguments",
41
+ "other parameters",
42
+ }
43
+ _SECTION = re.compile(
44
+ r"^(args|arguments|parameters|params|keyword args|keyword arguments|other parameters"
45
+ r"|returns?|raises|yields|examples?|notes?|see also|warnings?|references)\s*:?\s*$",
46
+ re.I,
47
+ )
48
+ _TOOL_NAME = re.compile(r"^[A-Za-z0-9_-]{1,64}$") # what model APIs accept
49
+ # TypedDict and dataclass field wrappers that carry no schema of their own.
50
+ _WRAPPERS = {typing.Annotated, typing.Required, typing.NotRequired}
51
+ if hasattr(typing, "ReadOnly"): # Python 3.13+
52
+ _WRAPPERS.add(typing.ReadOnly)
53
+
54
+
55
+ @dataclasses.dataclass
56
+ class FunctionTool(Tool):
57
+ """A Tool made by `@tool`; calling it calls the original function, as in unit tests."""
58
+
59
+ function: Callable | None = dataclasses.field(default=None, repr=False, compare=False)
60
+
61
+ def __call__(self, *args: Any, **kwargs: Any) -> Any:
62
+ assert self.function is not None
63
+ return self.function(*args, **kwargs)
64
+
65
+ @property
66
+ def __wrapped__(self) -> Callable | None:
67
+ return self.function
68
+
69
+
70
+ def tool(
71
+ function: Callable | None = None,
72
+ *,
73
+ name: str | None = None,
74
+ description: str | None = None,
75
+ execution_mode: str = "parallel",
76
+ output_schema: dict[str, Any] | None = None,
77
+ ) -> Any:
78
+ """Turn a function into a Tool. Use as ``@tool`` or ``@tool(name=...)``.
79
+
80
+ Parameters become the input schema: str, int, float, bool, None, list, tuple, set,
81
+ dict, Literal, Enum, Optional/Union, Annotated[T, "description"], TypedDict,
82
+ dataclasses, datetime, date, UUID, Path and pydantic models. A parameter annotated
83
+ ``ToolContext`` receives the call context instead. Google, NumPy or Sphinx docstrings
84
+ describe the parameters. Sync functions run in a worker thread. The function may
85
+ return a ToolResult, a string, a JSON value, None, or a date, dataclass or pydantic
86
+ model. The resulting tool can still be called like the function itself.
87
+ """
88
+
89
+ def build(target: Callable) -> Tool:
90
+ return _build(target, name, description, execution_mode, output_schema)
91
+
92
+ return build(function) if function is not None else build
93
+
94
+
95
+ def _target(function: Callable) -> Any:
96
+ """The plain function behind a partial or a callable object, for hints and docs."""
97
+ target: Any = function
98
+ while isinstance(target, functools.partial):
99
+ target = target.func
100
+ if not (inspect.isfunction(target) or inspect.ismethod(target)) and hasattr(target, "__call__"):
101
+ target = target.__call__
102
+ return target
103
+
104
+
105
+ def _build(
106
+ function: Callable,
107
+ name: str | None,
108
+ description: str | None,
109
+ execution_mode: str,
110
+ output_schema: dict[str, Any] | None,
111
+ ) -> Tool:
112
+ target = _target(function)
113
+ tool_name = name or getattr(function, "__name__", None) or getattr(target, "__name__", None)
114
+ if not tool_name or not _TOOL_NAME.match(tool_name):
115
+ raise ConfigurationError(
116
+ f"Tool name {tool_name!r} must be 1-64 letters, digits, '_' or '-'; pass @tool(name=...)"
117
+ )
118
+ try:
119
+ hints = typing.get_type_hints(target, include_extras=True)
120
+ except Exception as exc:
121
+ raise ConfigurationError(f"Cannot resolve type hints of {tool_name}: {exc}") from exc
122
+ doc = inspect.getdoc(target) or inspect.getdoc(function) or ""
123
+ summary, documented = _parse_docstring(doc)
124
+ schemas = _Schemas()
125
+ properties: dict[str, Any] = {}
126
+ required: list[str] = []
127
+ converters: dict[str, Any] = {}
128
+ context_name = None
129
+ parameters = list(inspect.signature(function).parameters.values())
130
+ if parameters and parameters[0].name in {"self", "cls"}:
131
+ raise ConfigurationError(
132
+ f"{tool_name}: decorate a bound method instead, for example tool(instance.{tool_name})"
133
+ )
134
+ for parameter in parameters:
135
+ if parameter.kind in (parameter.VAR_POSITIONAL, parameter.VAR_KEYWORD):
136
+ raise ConfigurationError(f"{tool_name}: *args and **kwargs cannot be tool parameters")
137
+ if parameter.kind is parameter.POSITIONAL_ONLY:
138
+ raise ConfigurationError(f"{tool_name}: positional-only parameters are not supported")
139
+ annotation = hints.get(parameter.name, Any)
140
+ if _is_context(annotation):
141
+ context_name = parameter.name
142
+ continue
143
+ schema = schemas.schema(annotation, parameter.name)
144
+ if parameter.name in documented and "description" not in schema:
145
+ schema["description"] = documented[parameter.name]
146
+ if parameter.default is parameter.empty:
147
+ required.append(parameter.name)
148
+ else:
149
+ default = _jsonable(parameter.default)
150
+ if default is not _MISSING:
151
+ schema["default"] = default
152
+ properties[parameter.name] = schema
153
+ converters[parameter.name] = annotation
154
+ input_schema: dict[str, Any] = {
155
+ "type": "object",
156
+ "properties": properties,
157
+ "required": required,
158
+ "additionalProperties": False,
159
+ }
160
+ if schemas.definitions:
161
+ input_schema["$defs"] = schemas.definitions
162
+
163
+ def arguments(args: dict[str, Any], context: ToolContext) -> dict[str, Any]:
164
+ values = {key: _convert(converters[key], value) for key, value in args.items()}
165
+ if context_name is not None:
166
+ values[context_name] = context
167
+ return values
168
+
169
+ if inspect.iscoroutinefunction(function) or inspect.iscoroutinefunction(target):
170
+
171
+ async def execute(args: dict[str, Any], context: ToolContext) -> Any:
172
+ return await function(**arguments(args, context))
173
+
174
+ else:
175
+
176
+ def execute(args: dict[str, Any], context: ToolContext) -> Any: # type: ignore[misc]
177
+ return function(**arguments(args, context))
178
+
179
+ return FunctionTool(
180
+ tool_name,
181
+ description if description is not None else summary,
182
+ input_schema,
183
+ execute,
184
+ output_schema=output_schema,
185
+ execution_mode=execution_mode,
186
+ function=function,
187
+ )
188
+
189
+
190
+ _MISSING = object()
191
+
192
+
193
+ def _jsonable(value: Any) -> Any:
194
+ if isinstance(value, enum.Enum):
195
+ value = value.value
196
+ try:
197
+ validate_json(value)
198
+ except Exception:
199
+ return _MISSING
200
+ return value
201
+
202
+
203
+ def _strip(annotation: Any) -> Any:
204
+ """Remove Annotated, Required, NotRequired and ReadOnly wrappers."""
205
+ while typing.get_origin(annotation) in _WRAPPERS:
206
+ annotation = typing.get_args(annotation)[0]
207
+ return annotation
208
+
209
+
210
+ def _is_context(annotation: Any) -> bool:
211
+ annotation = _strip(annotation)
212
+ if annotation is ToolContext:
213
+ return True
214
+ if typing.get_origin(annotation) in (Union, types.UnionType):
215
+ return ToolContext in typing.get_args(annotation)
216
+ return False
217
+
218
+
219
+ def _is_pydantic(annotation: Any) -> bool:
220
+ return isinstance(annotation, type) and hasattr(annotation, "model_json_schema")
221
+
222
+
223
+ def _is_typeddict(annotation: Any) -> bool:
224
+ return isinstance(annotation, type) and typing.is_typeddict(annotation)
225
+
226
+
227
+ def _fields(annotation: Any) -> list[tuple[str, Any, bool]]:
228
+ """(name, type, required) for the fields of a TypedDict or dataclass."""
229
+ hints = typing.get_type_hints(annotation, include_extras=True)
230
+ if _is_typeddict(annotation):
231
+ keys = annotation.__required_keys__
232
+ return [(key, hint, key in keys) for key, hint in hints.items()]
233
+ return [
234
+ (
235
+ field.name,
236
+ hints.get(field.name, Any),
237
+ field.default is dataclasses.MISSING and field.default_factory is dataclasses.MISSING,
238
+ )
239
+ for field in dataclasses.fields(annotation)
240
+ if field.init # ClassVar and init=False fields are not arguments
241
+ ]
242
+
243
+
244
+ class _Schemas:
245
+ """JSON Schema for parameter types; recursive types go to $defs."""
246
+
247
+ def __init__(self) -> None:
248
+ self.definitions: dict[str, Any] = {}
249
+ self.names: dict[Any, str] = {}
250
+ self.building: list[Any] = []
251
+ self.recursive: set[Any] = set()
252
+
253
+ def _name(self, annotation: Any) -> str:
254
+ if annotation not in self.names:
255
+ base = re.sub(r"[^A-Za-z0-9_.-]", "_", annotation.__qualname__)
256
+ name, n = base, 2
257
+ while name in self.names.values():
258
+ name, n = f"{base}_{n}", n + 1
259
+ self.names[annotation] = name
260
+ return self.names[annotation]
261
+
262
+ def schema(self, annotation: Any, where: str) -> dict[str, Any]:
263
+ origin = typing.get_origin(annotation)
264
+ args = typing.get_args(annotation)
265
+ schema: dict[str, Any]
266
+ if origin is typing.Annotated:
267
+ schema = self.schema(args[0], where)
268
+ text = next((a for a in args[1:] if isinstance(a, str)), None)
269
+ if text:
270
+ schema["description"] = text
271
+ return schema
272
+ if origin in _WRAPPERS:
273
+ return self.schema(args[0], where)
274
+ if annotation is Any or annotation is inspect.Parameter.empty:
275
+ return {}
276
+ if annotation in _PRIMITIVES:
277
+ return dict(_PRIMITIVES[annotation])
278
+ if origin is Literal:
279
+ return {"enum": list(args)}
280
+ if origin in (Union, types.UnionType):
281
+ options = [self.schema(a, where) for a in args]
282
+ return options[0] if len(options) == 1 else {"anyOf": options}
283
+ if isinstance(annotation, type) and issubclass(annotation, enum.Enum):
284
+ return {"enum": [member.value for member in annotation]}
285
+ if origin in (list, set, frozenset, Sequence) or annotation in (list, set, frozenset):
286
+ schema = {"type": "array"}
287
+ if args:
288
+ schema["items"] = self.schema(args[0], where)
289
+ if origin in (set, frozenset):
290
+ schema["uniqueItems"] = True
291
+ return schema
292
+ if origin is tuple or annotation is tuple:
293
+ if len(args) == 2 and args[1] is Ellipsis:
294
+ return {"type": "array", "items": self.schema(args[0], where)}
295
+ if not args:
296
+ return {"type": "array"}
297
+ return {
298
+ "type": "array",
299
+ "prefixItems": [self.schema(a, where) for a in args],
300
+ "minItems": len(args),
301
+ "maxItems": len(args),
302
+ }
303
+ if origin in (dict, Mapping) or annotation is dict:
304
+ schema = {"type": "object"}
305
+ if len(args) == 2:
306
+ if args[0] is not str:
307
+ raise ConfigurationError(f"{where}: dictionary keys must be str")
308
+ schema["additionalProperties"] = self.schema(args[1], where)
309
+ return schema
310
+ if _is_pydantic(annotation):
311
+ # Namespace this parameter's definitions so two models named alike cannot collide.
312
+ scope = re.sub(r"[^A-Za-z0-9_-]", "_", where)
313
+ model = annotation.model_json_schema(ref_template=f"#/$defs/{scope}.{{model}}")
314
+ for key, value in model.pop("$defs", {}).items():
315
+ self.definitions[f"{scope}.{key}"] = value
316
+ return model
317
+ if _is_typeddict(annotation) or dataclasses.is_dataclass(annotation):
318
+ return self._object(annotation, where)
319
+ raise ConfigurationError(f"{where}: unsupported parameter type {annotation!r}")
320
+
321
+ def _object(self, annotation: Any, where: str) -> dict[str, Any]:
322
+ ref = {"$ref": f"#/$defs/{self._name(annotation)}"}
323
+ if annotation in self.building: # a type that contains itself
324
+ self.recursive.add(annotation)
325
+ return ref
326
+ self.building.append(annotation)
327
+ try:
328
+ fields = _fields(annotation)
329
+ schema = {
330
+ "type": "object",
331
+ "properties": {key: self.schema(hint, f"{where}.{key}") for key, hint, _ in fields},
332
+ "required": [key for key, _, required in fields if required],
333
+ "additionalProperties": False,
334
+ }
335
+ finally:
336
+ self.building.pop()
337
+ if annotation in self.recursive:
338
+ self.definitions[self._name(annotation)] = schema
339
+ return ref
340
+ return schema
341
+
342
+
343
+ def _convert(annotation: Any, value: Any) -> Any:
344
+ """Turn validated JSON into the Python value the annotation asks for."""
345
+ annotation = _strip(annotation)
346
+ origin = typing.get_origin(annotation)
347
+ args = typing.get_args(annotation)
348
+ if value is None:
349
+ return None
350
+ if origin in (Union, types.UnionType):
351
+ for option in args:
352
+ if option is type(None):
353
+ continue
354
+ try:
355
+ return _convert(option, value)
356
+ except Exception:
357
+ continue
358
+ return value
359
+ if isinstance(annotation, type):
360
+ if issubclass(annotation, enum.Enum):
361
+ return annotation(value)
362
+ if annotation is datetime:
363
+ return datetime.fromisoformat(value)
364
+ if annotation is date:
365
+ return date.fromisoformat(value)
366
+ if annotation is UUID:
367
+ return UUID(value)
368
+ if annotation is Path:
369
+ return Path(value)
370
+ if annotation is float and isinstance(value, int):
371
+ return float(value)
372
+ if annotation is int and isinstance(value, float) and value.is_integer():
373
+ return int(value) # JSON Schema counts 3.0 as an integer
374
+ if _is_pydantic(annotation):
375
+ return annotation.model_validate(value) # type: ignore[attr-defined]
376
+ if _is_typeddict(annotation):
377
+ hints = {key: hint for key, hint, _ in _fields(annotation)}
378
+ return {key: _convert(hints.get(key, Any), item) for key, item in value.items()}
379
+ if dataclasses.is_dataclass(annotation):
380
+ hints = {key: hint for key, hint, _ in _fields(annotation)}
381
+ return annotation(**{k: _convert(hints.get(k, Any), v) for k, v in value.items()})
382
+ if origin in (list, Sequence) and args:
383
+ return [_convert(args[0], item) for item in value]
384
+ if origin in (set, frozenset):
385
+ return origin(_convert(args[0], item) if args else item for item in value)
386
+ if origin is tuple and args:
387
+ if len(args) == 2 and args[1] is Ellipsis:
388
+ return tuple(_convert(args[0], item) for item in value)
389
+ return tuple(_convert(a, item) for a, item in zip(args, value))
390
+ if origin in (dict, Mapping) and len(args) == 2:
391
+ return {key: _convert(args[1], item) for key, item in value.items()}
392
+ return value
393
+
394
+
395
+ def _parse_docstring(doc: str) -> tuple[str, dict[str, str]]:
396
+ """Summary text and per-parameter descriptions (Google, NumPy or Sphinx style)."""
397
+ lines = doc.splitlines()
398
+ summary: list[str] = []
399
+ params: dict[str, str] = {}
400
+ section = None # None while reading the summary
401
+ numpy = False # NumPy sections are underlined; their entries read "name : type"
402
+ current = None
403
+ base = 0 # indentation of a Google-style entry
404
+ index = 0
405
+ while index < len(lines):
406
+ line = lines[index]
407
+ stripped = line.strip()
408
+ indent = len(line) - len(line.lstrip())
409
+ underline = index + 1 < len(lines) and set(lines[index + 1].strip()) == {"-"}
410
+ sphinx = re.match(r":param\s+(?:[^:]*\s)?(\w+):\s*(.*)", stripped)
411
+ if sphinx:
412
+ section, current = "sphinx", sphinx.group(1)
413
+ params[current] = sphinx.group(2).strip()
414
+ elif _SECTION.match(stripped) and (stripped.endswith(":") or underline):
415
+ section, numpy, current = stripped.rstrip(":").strip().lower(), underline, None
416
+ index += 2 if underline else 1
417
+ continue
418
+ elif stripped.startswith(":"):
419
+ section, current = "other", None
420
+ elif section is None:
421
+ summary.append(line)
422
+ elif not stripped:
423
+ pass
424
+ elif section in _PARAM_SECTIONS:
425
+ entry = re.match(r"(\w+)\s*(?:\([^)]*\))?\s*:\s*(.*)", stripped)
426
+ if numpy and indent == 0:
427
+ current = re.match(r"\w+", stripped).group(0) # type: ignore[union-attr]
428
+ params[current] = ""
429
+ elif not numpy and entry and (current is None or indent <= base):
430
+ current = entry.group(1)
431
+ params[current] = entry.group(2).strip()
432
+ base = indent
433
+ elif current is not None:
434
+ params[current] = (params[current] + " " + stripped).strip()
435
+ elif section == "sphinx" and current is not None:
436
+ params[current] = (params[current] + " " + stripped).strip()
437
+ index += 1
438
+ return "\n".join(summary).strip(), {k: v for k, v in params.items() if v}
pi_python/hooks.py ADDED
@@ -0,0 +1,44 @@
1
+ from dataclasses import dataclass, field
2
+ from typing import Any, Callable
3
+ from .messages import AssistantMessage, Message, ToolResultMessage
4
+ from .models import ModelInfo
5
+ from .tools import Tool
6
+
7
+
8
+ @dataclass
9
+ class AgentConfigUpdate:
10
+ tools: list[Tool] | None = None
11
+ model: str | ModelInfo | None = None
12
+ options: dict[str, Any] | None = None
13
+
14
+
15
+ @dataclass
16
+ class TurnUpdate(AgentConfigUpdate):
17
+ context: list[Message] | None = None
18
+ messages: list[Message] = field(default_factory=list)
19
+
20
+
21
+ @dataclass
22
+ class RunContext:
23
+ messages: list[Message]
24
+ model: str | ModelInfo
25
+ options: dict[str, Any]
26
+ tools: list[Tool]
27
+ message: AssistantMessage | None = None
28
+ tool_results: list[ToolResultMessage] = field(default_factory=list)
29
+ new_messages: list[Message] = field(default_factory=list)
30
+
31
+
32
+ @dataclass
33
+ class Hooks:
34
+ prepare_request: Callable | None = None
35
+ prepare_next_turn: Callable | None = None
36
+ transform_context: Callable | None = None
37
+ convert_to_llm: Callable | None = None
38
+ before_tool_call: Callable | None = None
39
+ after_tool_call: Callable | None = None
40
+ finish_turn: Callable | None = None
41
+ get_api_key: Callable | None = None
42
+ on_payload: Callable | None = None
43
+ on_response: Callable | None = None
44
+ on_provider_stream_event: Callable | None = None
pi_python/limits.py ADDED
@@ -0,0 +1,28 @@
1
+ from dataclasses import dataclass
2
+ import math
3
+ from .errors import ConfigurationError
4
+
5
+
6
+ @dataclass(frozen=True)
7
+ class RunLimits:
8
+ # Like Pi, no request, tool-call or concurrency limit unless the application sets one.
9
+ max_model_requests: int | None = None
10
+ max_tool_calls: int | None = None
11
+ max_concurrency: int | None = None
12
+ tool_timeout: float | None = None
13
+ run_timeout: float | None = None
14
+ cleanup_timeout: float = 1.0
15
+
16
+ def __post_init__(self) -> None:
17
+ for name in ("max_model_requests", "max_tool_calls", "max_concurrency"):
18
+ value = getattr(self, name)
19
+ if value is None:
20
+ continue
21
+ if type(value) is not int or value < (1 if name == "max_concurrency" else 0):
22
+ raise ConfigurationError(f"Invalid {name}")
23
+ for name in ("tool_timeout", "run_timeout", "cleanup_timeout"):
24
+ value = getattr(self, name)
25
+ if value is None and name != "cleanup_timeout":
26
+ continue
27
+ if type(value) not in (int, float) or not math.isfinite(value) or value <= 0:
28
+ raise ConfigurationError(f"Invalid {name}")