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/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
+ ]