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/models.py
ADDED
|
@@ -0,0 +1,155 @@
|
|
|
1
|
+
"""Explicit, replaceable model capabilities. Catalog prices are snapshot estimates."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
from collections.abc import Iterable
|
|
5
|
+
from copy import deepcopy
|
|
6
|
+
from typing import TYPE_CHECKING, Any
|
|
7
|
+
from dataclasses import dataclass, field
|
|
8
|
+
from importlib.resources import files
|
|
9
|
+
import json
|
|
10
|
+
from .errors import ConfigurationError, UnsupportedCapabilityError
|
|
11
|
+
from .messages import ImageContent, validate_json
|
|
12
|
+
|
|
13
|
+
if TYPE_CHECKING:
|
|
14
|
+
from .provider import ModelRequest
|
|
15
|
+
|
|
16
|
+
LEVELS = ("off", "minimal", "low", "medium", "high", "xhigh", "max")
|
|
17
|
+
|
|
18
|
+
|
|
19
|
+
@dataclass(frozen=True)
|
|
20
|
+
class ModelInfo:
|
|
21
|
+
id: str
|
|
22
|
+
provider: str
|
|
23
|
+
api: str
|
|
24
|
+
name: str
|
|
25
|
+
context_window: int
|
|
26
|
+
max_tokens: int
|
|
27
|
+
reasoning: bool = False
|
|
28
|
+
input: tuple[str, ...] = ("text",)
|
|
29
|
+
thinking_level_map: dict[str, str | None] = field(default_factory=dict)
|
|
30
|
+
cost: dict[str, float] = field(default_factory=dict)
|
|
31
|
+
compat: dict = field(default_factory=dict)
|
|
32
|
+
input_limits: dict = field(default_factory=dict)
|
|
33
|
+
prompt_cache: dict = field(default_factory=dict)
|
|
34
|
+
base_url: str = ""
|
|
35
|
+
|
|
36
|
+
def supported_thinking_levels(self) -> tuple[str, ...]:
|
|
37
|
+
if not self.reasoning:
|
|
38
|
+
return ("off",)
|
|
39
|
+
return tuple(
|
|
40
|
+
level
|
|
41
|
+
for level in LEVELS
|
|
42
|
+
if self.thinking_level_map.get(level, "") is not None
|
|
43
|
+
and (level not in {"xhigh", "max"} or level in self.thinking_level_map)
|
|
44
|
+
)
|
|
45
|
+
|
|
46
|
+
def clamp_thinking_level(self, level: str) -> str:
|
|
47
|
+
if level not in LEVELS:
|
|
48
|
+
raise ConfigurationError("Unknown thinking level")
|
|
49
|
+
available = self.supported_thinking_levels()
|
|
50
|
+
index = LEVELS.index(level)
|
|
51
|
+
return next(
|
|
52
|
+
(v for v in (*LEVELS[index:], *reversed(LEVELS[:index])) if v in available), "off"
|
|
53
|
+
)
|
|
54
|
+
|
|
55
|
+
def provider_effort(self, level: str) -> str:
|
|
56
|
+
level = self.clamp_thinking_level(level)
|
|
57
|
+
return self.thinking_level_map.get(level) or ("none" if level == "off" else level)
|
|
58
|
+
|
|
59
|
+
def estimate_cost(self, usage: dict) -> dict[str, float]:
|
|
60
|
+
"""USD at captured base rates; excludes tiers, subscriptions and discounts."""
|
|
61
|
+
result = {
|
|
62
|
+
name: usage.get(name, 0) * self.cost.get(source, 0) / 1_000_000
|
|
63
|
+
for name, source in [
|
|
64
|
+
("input", "input"),
|
|
65
|
+
("output", "output"),
|
|
66
|
+
("cache_read", "cacheRead"),
|
|
67
|
+
("cache_write", "cacheWrite"),
|
|
68
|
+
]
|
|
69
|
+
}
|
|
70
|
+
result["total"] = sum(result.values())
|
|
71
|
+
return result
|
|
72
|
+
|
|
73
|
+
def validate_request(self, request: ModelRequest) -> None:
|
|
74
|
+
if "max_tokens" in request.options:
|
|
75
|
+
value = request.options["max_tokens"]
|
|
76
|
+
if type(value) is not int or not 0 < value <= self.max_tokens:
|
|
77
|
+
raise ConfigurationError(f"max_tokens must be between 1 and {self.max_tokens}")
|
|
78
|
+
images = []
|
|
79
|
+
for message in request.messages:
|
|
80
|
+
content = getattr(message, "content", None)
|
|
81
|
+
if isinstance(content, list):
|
|
82
|
+
images += [b for b in content if isinstance(b, ImageContent)]
|
|
83
|
+
if images and "image" not in self.input:
|
|
84
|
+
raise UnsupportedCapabilityError(f"{self.provider}/{self.id} does not support images")
|
|
85
|
+
limit = self.input_limits.get("images", {}).get("maxPerRequest")
|
|
86
|
+
if limit and len(images) > limit:
|
|
87
|
+
raise UnsupportedCapabilityError("Too many images for model")
|
|
88
|
+
|
|
89
|
+
|
|
90
|
+
class ModelCatalog:
|
|
91
|
+
def __init__(
|
|
92
|
+
self, models: Iterable[ModelInfo] = (), *, provenance: dict[str, Any] | None = None
|
|
93
|
+
) -> None:
|
|
94
|
+
self._models: dict[tuple[str, str], ModelInfo] = {}
|
|
95
|
+
self._provenance = deepcopy(provenance or {})
|
|
96
|
+
for model in models:
|
|
97
|
+
self.register(model)
|
|
98
|
+
|
|
99
|
+
@property
|
|
100
|
+
def provenance(self) -> dict[str, Any]:
|
|
101
|
+
return deepcopy(self._provenance)
|
|
102
|
+
|
|
103
|
+
def register(self, model: ModelInfo, *, replace: bool = False) -> None:
|
|
104
|
+
if not isinstance(model, ModelInfo) or not all((model.id, model.provider, model.api)):
|
|
105
|
+
raise ConfigurationError("Invalid model identity")
|
|
106
|
+
if (
|
|
107
|
+
type(model.max_tokens) is not int
|
|
108
|
+
or type(model.context_window) is not int
|
|
109
|
+
or min(model.max_tokens, model.context_window) <= 0
|
|
110
|
+
):
|
|
111
|
+
raise ConfigurationError("Model limits must be positive integers")
|
|
112
|
+
validate_json(model.compat)
|
|
113
|
+
if any(
|
|
114
|
+
k not in LEVELS or (v is not None and not isinstance(v, str))
|
|
115
|
+
for k, v in model.thinking_level_map.items()
|
|
116
|
+
):
|
|
117
|
+
raise ConfigurationError("Invalid thinking level map")
|
|
118
|
+
key = (model.provider, model.id)
|
|
119
|
+
if key in self._models and not replace:
|
|
120
|
+
raise ConfigurationError("Model already registered")
|
|
121
|
+
self._models[key] = deepcopy(model)
|
|
122
|
+
|
|
123
|
+
def get(self, provider: str, model: str) -> ModelInfo | None:
|
|
124
|
+
return deepcopy(self._models.get((provider, model)))
|
|
125
|
+
|
|
126
|
+
def list(self, provider: str | None = None) -> list[ModelInfo]:
|
|
127
|
+
return deepcopy(
|
|
128
|
+
[m for m in self._models.values() if provider is None or m.provider == provider]
|
|
129
|
+
)
|
|
130
|
+
|
|
131
|
+
@classmethod
|
|
132
|
+
def bundled(cls) -> ModelCatalog:
|
|
133
|
+
data = json.loads(
|
|
134
|
+
files("pi_python").joinpath("data/models.json").read_text(encoding="utf-8")
|
|
135
|
+
)
|
|
136
|
+
models = [
|
|
137
|
+
ModelInfo(
|
|
138
|
+
id=m["id"],
|
|
139
|
+
provider=m["provider"],
|
|
140
|
+
api=m["api"],
|
|
141
|
+
name=m["name"],
|
|
142
|
+
context_window=m["contextWindow"],
|
|
143
|
+
max_tokens=m["maxTokens"],
|
|
144
|
+
reasoning=m["reasoning"],
|
|
145
|
+
input=tuple(m["input"]),
|
|
146
|
+
thinking_level_map=m.get("thinkingLevelMap", {}),
|
|
147
|
+
cost=m.get("cost", {}),
|
|
148
|
+
compat=m.get("compat", {}),
|
|
149
|
+
input_limits=m.get("inputLimits", {}),
|
|
150
|
+
prompt_cache=m.get("promptCache", {}),
|
|
151
|
+
base_url=m.get("baseUrl", ""),
|
|
152
|
+
)
|
|
153
|
+
for m in data["models"]
|
|
154
|
+
]
|
|
155
|
+
return cls(models, provenance=data["provenance"])
|
pi_python/provider.py
ADDED
|
@@ -0,0 +1,123 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
from collections.abc import AsyncIterator
|
|
3
|
+
from dataclasses import dataclass, field
|
|
4
|
+
from typing import Any, Callable, Protocol
|
|
5
|
+
from .cancellation import CancelToken
|
|
6
|
+
from .models import ModelInfo
|
|
7
|
+
from .messages import (
|
|
8
|
+
AssistantMessage,
|
|
9
|
+
Message,
|
|
10
|
+
ToolDeclaration,
|
|
11
|
+
TextContent,
|
|
12
|
+
ThinkingContent,
|
|
13
|
+
ToolCall,
|
|
14
|
+
)
|
|
15
|
+
from .errors import ConfigurationError
|
|
16
|
+
|
|
17
|
+
|
|
18
|
+
@dataclass
|
|
19
|
+
class ModelRequest:
|
|
20
|
+
messages: list[Message]
|
|
21
|
+
tools: list[ToolDeclaration] = field(default_factory=list)
|
|
22
|
+
model: str = "mock"
|
|
23
|
+
options: dict[str, Any] = field(default_factory=dict)
|
|
24
|
+
# The caller's model record, when given; otherwise providers look `model` up.
|
|
25
|
+
model_info: ModelInfo | None = None
|
|
26
|
+
api_key: str | None = field(default=None, repr=False)
|
|
27
|
+
on_payload: Callable | None = field(default=None, repr=False)
|
|
28
|
+
on_response: Callable | None = field(default=None, repr=False)
|
|
29
|
+
on_provider_stream_event: Callable | None = field(default=None, repr=False)
|
|
30
|
+
|
|
31
|
+
|
|
32
|
+
@dataclass
|
|
33
|
+
class ModelEvent:
|
|
34
|
+
"""One event of Pi's assistant stream.
|
|
35
|
+
|
|
36
|
+
Order: an optional `start`, then for each content block `*_start`, any number of
|
|
37
|
+
`*_delta`, and `*_end` (block kinds: text, thinking, toolcall), then exactly one
|
|
38
|
+
`done` carrying the complete message. Streaming is optional: a provider may yield
|
|
39
|
+
only `done`. Failures are raised; consumers see them as an `error` event.
|
|
40
|
+
`index` is the block's position in the final message.
|
|
41
|
+
"""
|
|
42
|
+
|
|
43
|
+
type: str
|
|
44
|
+
delta: str = ""
|
|
45
|
+
call_id: str | None = None
|
|
46
|
+
name: str | None = None
|
|
47
|
+
index: int = 0
|
|
48
|
+
message: AssistantMessage | None = None
|
|
49
|
+
partial: AssistantMessage | None = None
|
|
50
|
+
block: TextContent | ThinkingContent | ToolCall | None = None
|
|
51
|
+
content: str | None = None
|
|
52
|
+
reason: str | None = None
|
|
53
|
+
|
|
54
|
+
@classmethod
|
|
55
|
+
def text(cls, delta: str, index: int = 0) -> ModelEvent:
|
|
56
|
+
return cls("text_delta", delta, index=index)
|
|
57
|
+
|
|
58
|
+
@classmethod
|
|
59
|
+
def thinking(cls, delta: str, index: int = 0) -> ModelEvent:
|
|
60
|
+
return cls("thinking_delta", delta, index=index)
|
|
61
|
+
|
|
62
|
+
@classmethod
|
|
63
|
+
def toolcall(cls, delta: str, index: int = 0) -> ModelEvent:
|
|
64
|
+
"""A fragment of the JSON arguments of the tool call block at `index`."""
|
|
65
|
+
return cls("toolcall_delta", delta, index=index)
|
|
66
|
+
|
|
67
|
+
@classmethod
|
|
68
|
+
def boundary(
|
|
69
|
+
cls, phase: str, index: int, block: TextContent | ThinkingContent | ToolCall
|
|
70
|
+
) -> ModelEvent:
|
|
71
|
+
"""`phase` is "start" or "end"; the end event carries the complete block."""
|
|
72
|
+
kind = "toolcall" if isinstance(block, ToolCall) else block.type
|
|
73
|
+
return cls(
|
|
74
|
+
f"{kind}_{phase}",
|
|
75
|
+
index=index,
|
|
76
|
+
block=block,
|
|
77
|
+
call_id=block.id if isinstance(block, ToolCall) else None,
|
|
78
|
+
name=block.name if isinstance(block, ToolCall) else None,
|
|
79
|
+
)
|
|
80
|
+
|
|
81
|
+
@classmethod
|
|
82
|
+
def done(cls, message: AssistantMessage) -> ModelEvent:
|
|
83
|
+
return cls("done", message=message, reason=message.stop_reason)
|
|
84
|
+
|
|
85
|
+
|
|
86
|
+
class Provider(Protocol):
|
|
87
|
+
def stream(self, request: ModelRequest, cancel: CancelToken) -> AsyncIterator[ModelEvent]: ...
|
|
88
|
+
|
|
89
|
+
|
|
90
|
+
_default_stream: Provider | Callable | None = None
|
|
91
|
+
|
|
92
|
+
|
|
93
|
+
def set_default_stream_fn(stream: Provider | Callable | None) -> None:
|
|
94
|
+
"""Explicit process-wide fallback, matching Pi's setDefaultStreamFn."""
|
|
95
|
+
global _default_stream
|
|
96
|
+
_default_stream = stream
|
|
97
|
+
|
|
98
|
+
|
|
99
|
+
class FunctionProvider:
|
|
100
|
+
def __init__(self, function: Callable):
|
|
101
|
+
self.function = function
|
|
102
|
+
|
|
103
|
+
def stream(self, request: ModelRequest, cancel: CancelToken) -> AsyncIterator[ModelEvent]:
|
|
104
|
+
return self.function(request, cancel)
|
|
105
|
+
|
|
106
|
+
|
|
107
|
+
def has_default_stream() -> bool:
|
|
108
|
+
return _default_stream is not None
|
|
109
|
+
|
|
110
|
+
|
|
111
|
+
class DefaultProvider:
|
|
112
|
+
@property
|
|
113
|
+
def name(self) -> str:
|
|
114
|
+
return getattr(_default_stream, "name", "custom")
|
|
115
|
+
|
|
116
|
+
def stream(self, request: ModelRequest, cancel: CancelToken) -> AsyncIterator[ModelEvent]:
|
|
117
|
+
if _default_stream is None:
|
|
118
|
+
raise ConfigurationError(
|
|
119
|
+
"No default stream configured; pass provider or call set_default_stream_fn"
|
|
120
|
+
)
|
|
121
|
+
if hasattr(_default_stream, "stream"):
|
|
122
|
+
return _default_stream.stream(request, cancel)
|
|
123
|
+
return _default_stream(request, cancel)
|
|
@@ -0,0 +1,21 @@
|
|
|
1
|
+
"""Opt-in network providers; install pi-python-core[providers] for transports."""
|
|
2
|
+
|
|
3
|
+
from .anthropic import AnthropicProvider
|
|
4
|
+
from .completions import OpenAICompletionsProvider
|
|
5
|
+
from .openai import OpenAIProvider, OpenAICodexProvider, DeepSeekProvider
|
|
6
|
+
from .oauth import OAuthClient, OAuthCredential, OAuthAttempt, RefreshingCredentials
|
|
7
|
+
from .transport import HTTPTransport, ProviderHTTPError
|
|
8
|
+
|
|
9
|
+
__all__ = [
|
|
10
|
+
"AnthropicProvider",
|
|
11
|
+
"OpenAIProvider",
|
|
12
|
+
"OpenAICompletionsProvider",
|
|
13
|
+
"DeepSeekProvider",
|
|
14
|
+
"OpenAICodexProvider",
|
|
15
|
+
"OAuthClient",
|
|
16
|
+
"OAuthCredential",
|
|
17
|
+
"OAuthAttempt",
|
|
18
|
+
"RefreshingCredentials",
|
|
19
|
+
"HTTPTransport",
|
|
20
|
+
"ProviderHTTPError",
|
|
21
|
+
]
|