callm-toolkit 0.1.0__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.
- callm/__about__.py +1 -0
- callm/__init__.py +103 -0
- callm/__main__.py +3 -0
- callm/api.py +411 -0
- callm/budgets.py +198 -0
- callm/classify.py +262 -0
- callm/cli.py +349 -0
- callm/config.py +516 -0
- callm/data/__init__.py +1 -0
- callm/data/pricing.json +666 -0
- callm/decorator.py +361 -0
- callm/embeddings.py +135 -0
- callm/errors.py +142 -0
- callm/interception.py +323 -0
- callm/middleware/__init__.py +128 -0
- callm/middleware/cache.py +222 -0
- callm/middleware/cost.py +96 -0
- callm/middleware/fallback.py +75 -0
- callm/middleware/retry.py +52 -0
- callm/middleware/security.py +115 -0
- callm/middleware/telemetry.py +139 -0
- callm/middleware/validator.py +60 -0
- callm/pipeline.py +216 -0
- callm/pricing.py +389 -0
- callm/providers/__init__.py +18 -0
- callm/providers/anthropic.py +256 -0
- callm/providers/base.py +245 -0
- callm/providers/google.py +332 -0
- callm/providers/openai.py +243 -0
- callm/providers/registry.py +125 -0
- callm/py.typed +0 -0
- callm/reports.py +55 -0
- callm/security/__init__.py +20 -0
- callm/security/injection.py +263 -0
- callm/security/pii.py +259 -0
- callm/storage/__init__.py +27 -0
- callm/storage/base.py +193 -0
- callm/storage/memory.py +128 -0
- callm/storage/redis.py +171 -0
- callm/storage/sqlite.py +408 -0
- callm/types.py +376 -0
- callm/validation.py +159 -0
- callm_toolkit-0.1.0.dist-info/METADATA +363 -0
- callm_toolkit-0.1.0.dist-info/RECORD +47 -0
- callm_toolkit-0.1.0.dist-info/WHEEL +4 -0
- callm_toolkit-0.1.0.dist-info/entry_points.txt +2 -0
- callm_toolkit-0.1.0.dist-info/licenses/LICENSE +21 -0
callm/__about__.py
ADDED
|
@@ -0,0 +1 @@
|
|
|
1
|
+
__version__ = "0.1.0"
|
callm/__init__.py
ADDED
|
@@ -0,0 +1,103 @@
|
|
|
1
|
+
"""callm - the production toolkit for LLM calls.
|
|
2
|
+
|
|
3
|
+
Wrap any function that calls OpenAI, Anthropic or Gemini and get caching, retries,
|
|
4
|
+
provider fallback, cost tracking and budgets, PII redaction, prompt injection detection,
|
|
5
|
+
structured output validation and telemetry - with zero required dependencies::
|
|
6
|
+
|
|
7
|
+
from callm import callm
|
|
8
|
+
|
|
9
|
+
@callm(cache=True, retry=3, fallback=["anthropic/claude-sonnet-5"], max_cost=0.25)
|
|
10
|
+
def summarize(text: str):
|
|
11
|
+
return openai.chat.completions.create(
|
|
12
|
+
model="gpt-4o", messages=[{"role": "user", "content": text}]
|
|
13
|
+
)
|
|
14
|
+
"""
|
|
15
|
+
|
|
16
|
+
import logging as _logging
|
|
17
|
+
|
|
18
|
+
from callm.__about__ import __version__
|
|
19
|
+
from callm.api import acomplete, complete, shield
|
|
20
|
+
from callm.budgets import Budget, budget, budget_for
|
|
21
|
+
from callm.config import (
|
|
22
|
+
CacheConfig,
|
|
23
|
+
InjectionConfig,
|
|
24
|
+
PIIConfig,
|
|
25
|
+
RetryConfig,
|
|
26
|
+
Settings,
|
|
27
|
+
configure,
|
|
28
|
+
get_settings,
|
|
29
|
+
reset_settings,
|
|
30
|
+
)
|
|
31
|
+
from callm.decorator import callm
|
|
32
|
+
from callm.embeddings import HashingEmbedder, OpenAIEmbedder, SentenceTransformerEmbedder
|
|
33
|
+
from callm.errors import (
|
|
34
|
+
AllProvidersFailedError,
|
|
35
|
+
BudgetExceeded,
|
|
36
|
+
BudgetExceededError,
|
|
37
|
+
CallmError,
|
|
38
|
+
ConfigurationError,
|
|
39
|
+
MissingDependencyError,
|
|
40
|
+
OutputValidationError,
|
|
41
|
+
PIIDetectedError,
|
|
42
|
+
PromptInjectionError,
|
|
43
|
+
ProviderNotAvailableError,
|
|
44
|
+
SecurityError,
|
|
45
|
+
)
|
|
46
|
+
from callm.middleware.telemetry import last_call
|
|
47
|
+
from callm.pricing import ModelPrice, get_price, set_price
|
|
48
|
+
from callm.providers import OpenAIProvider, Provider, register_provider
|
|
49
|
+
from callm.reports import clear_cache, stats
|
|
50
|
+
from callm.security import detect_injection, redact_pii
|
|
51
|
+
from callm.types import CallRecord, LLMRequest, LLMResponse, Message, Target, Usage
|
|
52
|
+
|
|
53
|
+
_logging.getLogger("callm").addHandler(_logging.NullHandler())
|
|
54
|
+
|
|
55
|
+
__all__ = [
|
|
56
|
+
"AllProvidersFailedError",
|
|
57
|
+
"Budget",
|
|
58
|
+
"BudgetExceeded",
|
|
59
|
+
"BudgetExceededError",
|
|
60
|
+
"CacheConfig",
|
|
61
|
+
"CallRecord",
|
|
62
|
+
"CallmError",
|
|
63
|
+
"ConfigurationError",
|
|
64
|
+
"HashingEmbedder",
|
|
65
|
+
"InjectionConfig",
|
|
66
|
+
"LLMRequest",
|
|
67
|
+
"LLMResponse",
|
|
68
|
+
"Message",
|
|
69
|
+
"MissingDependencyError",
|
|
70
|
+
"ModelPrice",
|
|
71
|
+
"OpenAIEmbedder",
|
|
72
|
+
"OpenAIProvider",
|
|
73
|
+
"OutputValidationError",
|
|
74
|
+
"PIIConfig",
|
|
75
|
+
"PIIDetectedError",
|
|
76
|
+
"PromptInjectionError",
|
|
77
|
+
"Provider",
|
|
78
|
+
"ProviderNotAvailableError",
|
|
79
|
+
"RetryConfig",
|
|
80
|
+
"SecurityError",
|
|
81
|
+
"SentenceTransformerEmbedder",
|
|
82
|
+
"Settings",
|
|
83
|
+
"Target",
|
|
84
|
+
"Usage",
|
|
85
|
+
"__version__",
|
|
86
|
+
"acomplete",
|
|
87
|
+
"budget",
|
|
88
|
+
"budget_for",
|
|
89
|
+
"callm",
|
|
90
|
+
"clear_cache",
|
|
91
|
+
"complete",
|
|
92
|
+
"configure",
|
|
93
|
+
"detect_injection",
|
|
94
|
+
"get_price",
|
|
95
|
+
"get_settings",
|
|
96
|
+
"last_call",
|
|
97
|
+
"redact_pii",
|
|
98
|
+
"register_provider",
|
|
99
|
+
"reset_settings",
|
|
100
|
+
"set_price",
|
|
101
|
+
"shield",
|
|
102
|
+
"stats",
|
|
103
|
+
]
|
callm/__main__.py
ADDED
callm/api.py
ADDED
|
@@ -0,0 +1,411 @@
|
|
|
1
|
+
"""``shield`` context manager and the direct ``complete`` / ``acomplete`` API."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
from collections.abc import Callable, Mapping, Sequence
|
|
6
|
+
from types import TracebackType
|
|
7
|
+
from typing import TYPE_CHECKING, Any
|
|
8
|
+
|
|
9
|
+
from callm.config import (
|
|
10
|
+
CacheConfig,
|
|
11
|
+
CallConfig,
|
|
12
|
+
InjectionConfig,
|
|
13
|
+
PIIConfig,
|
|
14
|
+
RetryConfig,
|
|
15
|
+
build_call_config,
|
|
16
|
+
merge_scopes,
|
|
17
|
+
)
|
|
18
|
+
from callm.errors import OutputValidationError
|
|
19
|
+
from callm.interception import ACTIVE, clean_kwargs, ensure_installed, enter_scope
|
|
20
|
+
from callm.middleware import build_handler
|
|
21
|
+
from callm.pipeline import CallState, run_async, run_sync
|
|
22
|
+
from callm.providers.anthropic import AnthropicProvider
|
|
23
|
+
from callm.providers.base import to_plain
|
|
24
|
+
from callm.providers.registry import get_provider, parse_target
|
|
25
|
+
from callm.types import CallRecord, LLMRequest, LLMResponse, Message, Target
|
|
26
|
+
from callm.validation import ValidationFailure, validate_text
|
|
27
|
+
|
|
28
|
+
if TYPE_CHECKING:
|
|
29
|
+
from callm.budgets import Budget
|
|
30
|
+
|
|
31
|
+
CONFIG_FIELDS = (
|
|
32
|
+
"cache",
|
|
33
|
+
"retry",
|
|
34
|
+
"fallback",
|
|
35
|
+
"fallback_on",
|
|
36
|
+
"max_cost",
|
|
37
|
+
"budget",
|
|
38
|
+
"block_pii",
|
|
39
|
+
"detect_injection",
|
|
40
|
+
"output_schema",
|
|
41
|
+
"validation_retries",
|
|
42
|
+
"name",
|
|
43
|
+
"tags",
|
|
44
|
+
"telemetry",
|
|
45
|
+
)
|
|
46
|
+
|
|
47
|
+
|
|
48
|
+
def _to_messages(messages: str | Sequence[Any]) -> list[Message]:
|
|
49
|
+
if isinstance(messages, str):
|
|
50
|
+
return [Message("user", messages)]
|
|
51
|
+
result: list[Message] = []
|
|
52
|
+
for item in messages:
|
|
53
|
+
if isinstance(item, Message):
|
|
54
|
+
result.append(item)
|
|
55
|
+
continue
|
|
56
|
+
plain = to_plain(item)
|
|
57
|
+
if not isinstance(plain, dict) or "role" not in plain:
|
|
58
|
+
raise TypeError(f"messages must be dicts with 'role' and 'content', got {item!r}")
|
|
59
|
+
content = plain.get("content", "")
|
|
60
|
+
if not isinstance(content, (str, list)):
|
|
61
|
+
content = "" if content is None else str(content)
|
|
62
|
+
result.append(Message(str(plain["role"]), content))
|
|
63
|
+
if not result:
|
|
64
|
+
raise ValueError("messages must not be empty")
|
|
65
|
+
return result
|
|
66
|
+
|
|
67
|
+
|
|
68
|
+
def anthropic_sampling_kwargs(
|
|
69
|
+
sampling: dict[str, Any], extra_body: dict[str, Any] | None
|
|
70
|
+
) -> dict[str, Any]:
|
|
71
|
+
"""Pass ``temperature``/``top_p`` the way the installed ``anthropic`` SDK accepts them.
|
|
72
|
+
|
|
73
|
+
``anthropic>=1.0`` removed the keyword arguments (newer Claude models reject them), but
|
|
74
|
+
older models still honour them when sent in the request body via ``extra_body``.
|
|
75
|
+
"""
|
|
76
|
+
try:
|
|
77
|
+
import inspect
|
|
78
|
+
|
|
79
|
+
from anthropic.resources.messages import Messages
|
|
80
|
+
|
|
81
|
+
accepted = set(inspect.signature(Messages.create).parameters)
|
|
82
|
+
except Exception:
|
|
83
|
+
accepted = set()
|
|
84
|
+
direct = {k: v for k, v in sampling.items() if k in accepted}
|
|
85
|
+
body = {k: v for k, v in sampling.items() if k not in accepted}
|
|
86
|
+
result: dict[str, Any] = dict(direct)
|
|
87
|
+
if body or extra_body:
|
|
88
|
+
result["extra_body"] = {**(extra_body or {}), **body}
|
|
89
|
+
return result
|
|
90
|
+
|
|
91
|
+
|
|
92
|
+
def build_request(
|
|
93
|
+
model: str | Target,
|
|
94
|
+
messages: str | Sequence[Any],
|
|
95
|
+
*,
|
|
96
|
+
provider: str | None = None,
|
|
97
|
+
max_tokens: int | None = None,
|
|
98
|
+
temperature: float | None = None,
|
|
99
|
+
top_p: float | None = None,
|
|
100
|
+
stop: str | Sequence[str] | None = None,
|
|
101
|
+
**provider_kwargs: Any,
|
|
102
|
+
) -> LLMRequest:
|
|
103
|
+
"""Build a provider-native request from canonical arguments plus extra SDK kwargs."""
|
|
104
|
+
provider_kwargs = clean_kwargs(provider_kwargs)
|
|
105
|
+
target = parse_target(model, provider)
|
|
106
|
+
adapter = get_provider(target.provider)
|
|
107
|
+
params: dict[str, Any] = {}
|
|
108
|
+
if max_tokens is not None:
|
|
109
|
+
params["max_tokens"] = max_tokens
|
|
110
|
+
if temperature is not None:
|
|
111
|
+
params["temperature"] = temperature
|
|
112
|
+
if top_p is not None:
|
|
113
|
+
params["top_p"] = top_p
|
|
114
|
+
if stop is not None:
|
|
115
|
+
params["stop"] = [stop] if isinstance(stop, str) else list(stop)
|
|
116
|
+
canonical = LLMRequest(
|
|
117
|
+
provider=target.provider,
|
|
118
|
+
model=target.model,
|
|
119
|
+
messages=_to_messages(messages),
|
|
120
|
+
params=params,
|
|
121
|
+
)
|
|
122
|
+
kwargs = adapter.to_kwargs(canonical)
|
|
123
|
+
if isinstance(adapter, AnthropicProvider):
|
|
124
|
+
# Explicit sampling parameters are the caller's choice on models that accept them.
|
|
125
|
+
sampling: dict[str, Any] = {k: params[k] for k in ("temperature", "top_p") if k in params}
|
|
126
|
+
if sampling:
|
|
127
|
+
kwargs.update(anthropic_sampling_kwargs(sampling, provider_kwargs.get("extra_body")))
|
|
128
|
+
provider_kwargs.pop("extra_body", None)
|
|
129
|
+
if "config" in provider_kwargs and isinstance(kwargs.get("config"), dict):
|
|
130
|
+
user_config = provider_kwargs.pop("config")
|
|
131
|
+
if isinstance(user_config, dict):
|
|
132
|
+
kwargs["config"] = {**kwargs["config"], **user_config}
|
|
133
|
+
elif user_config is not None:
|
|
134
|
+
kwargs["config"] = user_config.model_copy(update=kwargs["config"])
|
|
135
|
+
kwargs.update(provider_kwargs)
|
|
136
|
+
return adapter.parse_native_request(kwargs)
|
|
137
|
+
|
|
138
|
+
|
|
139
|
+
def _finalize(response: LLMResponse, config: CallConfig) -> LLMResponse:
|
|
140
|
+
if response.raw is None:
|
|
141
|
+
try:
|
|
142
|
+
response.raw = get_provider(response.provider).native_for(response)
|
|
143
|
+
except Exception:
|
|
144
|
+
response.raw = None
|
|
145
|
+
if config.output_schema is not None and response.parsed is None:
|
|
146
|
+
try:
|
|
147
|
+
response.parsed = validate_text(response.text, config.output_schema)
|
|
148
|
+
except ValidationFailure as failure:
|
|
149
|
+
raise OutputValidationError(
|
|
150
|
+
str(failure), errors=failure.errors, raw_text=failure.text, attempts=1
|
|
151
|
+
) from None
|
|
152
|
+
return response
|
|
153
|
+
|
|
154
|
+
|
|
155
|
+
def _state(request: LLMRequest, config: CallConfig, mode: str) -> CallState:
|
|
156
|
+
record = CallRecord(
|
|
157
|
+
function=config.name or "callm.complete", provider=request.provider, model=request.model
|
|
158
|
+
)
|
|
159
|
+
return CallState(request=request, config=config, record=record, mode=mode)
|
|
160
|
+
|
|
161
|
+
|
|
162
|
+
def _inherit(config: CallConfig) -> CallConfig:
|
|
163
|
+
parent = ACTIVE.get()
|
|
164
|
+
return merge_scopes(parent.config, config) if parent is not None else config
|
|
165
|
+
|
|
166
|
+
|
|
167
|
+
def run_request(config: CallConfig, request: LLMRequest) -> LLMResponse:
|
|
168
|
+
"""Run a canonical request through the pipeline synchronously."""
|
|
169
|
+
config = _inherit(config)
|
|
170
|
+
response = run_sync(build_handler(config)(_state(request, config, "sync")))
|
|
171
|
+
return _finalize(response, config)
|
|
172
|
+
|
|
173
|
+
|
|
174
|
+
async def arun_request(config: CallConfig, request: LLMRequest) -> LLMResponse:
|
|
175
|
+
"""Run a canonical request through the pipeline on the event loop."""
|
|
176
|
+
config = _inherit(config)
|
|
177
|
+
response = await run_async(build_handler(config)(_state(request, config, "async")))
|
|
178
|
+
return _finalize(response, config)
|
|
179
|
+
|
|
180
|
+
|
|
181
|
+
def complete(
|
|
182
|
+
model: str | Target,
|
|
183
|
+
messages: str | Sequence[Any],
|
|
184
|
+
*,
|
|
185
|
+
provider: str | None = None,
|
|
186
|
+
max_tokens: int | None = None,
|
|
187
|
+
temperature: float | None = None,
|
|
188
|
+
top_p: float | None = None,
|
|
189
|
+
stop: str | Sequence[str] | None = None,
|
|
190
|
+
cache: bool | str | CacheConfig | None = False,
|
|
191
|
+
retry: bool | int | RetryConfig | None = None,
|
|
192
|
+
fallback: str | Sequence[str | Target] | None = None,
|
|
193
|
+
fallback_on: Callable[[BaseException], bool] | None = None,
|
|
194
|
+
max_cost: float | None = None,
|
|
195
|
+
budget: Budget | None = None,
|
|
196
|
+
block_pii: bool | PIIConfig | None = False,
|
|
197
|
+
detect_injection: bool | InjectionConfig | None = False,
|
|
198
|
+
output_schema: Any = None,
|
|
199
|
+
validation_retries: int = 2,
|
|
200
|
+
name: str | None = None,
|
|
201
|
+
tags: Mapping[str, str] | None = None,
|
|
202
|
+
telemetry: bool = True,
|
|
203
|
+
**provider_kwargs: Any,
|
|
204
|
+
) -> LLMResponse:
|
|
205
|
+
"""Call a model through the callm pipeline and get a provider-neutral response.
|
|
206
|
+
|
|
207
|
+
``model`` is ``"provider/model"`` (``"anthropic/claude-sonnet-5"``) or a model id whose
|
|
208
|
+
provider can be inferred (``"gpt-4o-mini"``). ``messages`` is a string or a list of
|
|
209
|
+
``{"role", "content"}`` dicts. Extra keyword arguments go to the provider SDK unchanged.
|
|
210
|
+
|
|
211
|
+
Example::
|
|
212
|
+
|
|
213
|
+
response = callm.complete("gpt-4o-mini", "Say hi", cache=True, retry=3)
|
|
214
|
+
print(response.text, response.cost)
|
|
215
|
+
"""
|
|
216
|
+
config = build_call_config(
|
|
217
|
+
cache=cache,
|
|
218
|
+
retry=retry,
|
|
219
|
+
fallback=fallback,
|
|
220
|
+
fallback_on=fallback_on,
|
|
221
|
+
max_cost=max_cost,
|
|
222
|
+
budget=budget,
|
|
223
|
+
block_pii=block_pii,
|
|
224
|
+
detect_injection=detect_injection,
|
|
225
|
+
output_schema=output_schema,
|
|
226
|
+
validation_retries=validation_retries,
|
|
227
|
+
name=name,
|
|
228
|
+
tags=tags,
|
|
229
|
+
telemetry=telemetry,
|
|
230
|
+
)
|
|
231
|
+
request = build_request(
|
|
232
|
+
model,
|
|
233
|
+
messages,
|
|
234
|
+
provider=provider,
|
|
235
|
+
max_tokens=max_tokens,
|
|
236
|
+
temperature=temperature,
|
|
237
|
+
top_p=top_p,
|
|
238
|
+
stop=stop,
|
|
239
|
+
**provider_kwargs,
|
|
240
|
+
)
|
|
241
|
+
return run_request(config, request)
|
|
242
|
+
|
|
243
|
+
|
|
244
|
+
async def acomplete(
|
|
245
|
+
model: str | Target,
|
|
246
|
+
messages: str | Sequence[Any],
|
|
247
|
+
*,
|
|
248
|
+
provider: str | None = None,
|
|
249
|
+
max_tokens: int | None = None,
|
|
250
|
+
temperature: float | None = None,
|
|
251
|
+
top_p: float | None = None,
|
|
252
|
+
stop: str | Sequence[str] | None = None,
|
|
253
|
+
cache: bool | str | CacheConfig | None = False,
|
|
254
|
+
retry: bool | int | RetryConfig | None = None,
|
|
255
|
+
fallback: str | Sequence[str | Target] | None = None,
|
|
256
|
+
fallback_on: Callable[[BaseException], bool] | None = None,
|
|
257
|
+
max_cost: float | None = None,
|
|
258
|
+
budget: Budget | None = None,
|
|
259
|
+
block_pii: bool | PIIConfig | None = False,
|
|
260
|
+
detect_injection: bool | InjectionConfig | None = False,
|
|
261
|
+
output_schema: Any = None,
|
|
262
|
+
validation_retries: int = 2,
|
|
263
|
+
name: str | None = None,
|
|
264
|
+
tags: Mapping[str, str] | None = None,
|
|
265
|
+
telemetry: bool = True,
|
|
266
|
+
**provider_kwargs: Any,
|
|
267
|
+
) -> LLMResponse:
|
|
268
|
+
"""Async version of :func:`complete`."""
|
|
269
|
+
config = build_call_config(
|
|
270
|
+
cache=cache,
|
|
271
|
+
retry=retry,
|
|
272
|
+
fallback=fallback,
|
|
273
|
+
fallback_on=fallback_on,
|
|
274
|
+
max_cost=max_cost,
|
|
275
|
+
budget=budget,
|
|
276
|
+
block_pii=block_pii,
|
|
277
|
+
detect_injection=detect_injection,
|
|
278
|
+
output_schema=output_schema,
|
|
279
|
+
validation_retries=validation_retries,
|
|
280
|
+
name=name,
|
|
281
|
+
tags=tags,
|
|
282
|
+
telemetry=telemetry,
|
|
283
|
+
)
|
|
284
|
+
request = build_request(
|
|
285
|
+
model,
|
|
286
|
+
messages,
|
|
287
|
+
provider=provider,
|
|
288
|
+
max_tokens=max_tokens,
|
|
289
|
+
temperature=temperature,
|
|
290
|
+
top_p=top_p,
|
|
291
|
+
stop=stop,
|
|
292
|
+
**provider_kwargs,
|
|
293
|
+
)
|
|
294
|
+
return await arun_request(config, request)
|
|
295
|
+
|
|
296
|
+
|
|
297
|
+
class shield:
|
|
298
|
+
"""Apply callm protections to a block of code, or make one-off calls.
|
|
299
|
+
|
|
300
|
+
As a context manager, every supported SDK call inside the block goes through the
|
|
301
|
+
pipeline::
|
|
302
|
+
|
|
303
|
+
with shield(block_pii=True, detect_injection=True):
|
|
304
|
+
client.chat.completions.create(model="gpt-4o", messages=messages)
|
|
305
|
+
|
|
306
|
+
The object also offers :meth:`complete` / :meth:`acomplete` for direct calls::
|
|
307
|
+
|
|
308
|
+
with shield(block_pii=True, detect_injection=True) as s:
|
|
309
|
+
response = s.complete(provider="anthropic", model="claude-sonnet-5", messages=messages)
|
|
310
|
+
"""
|
|
311
|
+
|
|
312
|
+
def __init__(
|
|
313
|
+
self,
|
|
314
|
+
*,
|
|
315
|
+
cache: bool | str | CacheConfig | None = False,
|
|
316
|
+
retry: bool | int | RetryConfig | None = None,
|
|
317
|
+
fallback: str | Sequence[str | Target] | None = None,
|
|
318
|
+
fallback_on: Callable[[BaseException], bool] | None = None,
|
|
319
|
+
max_cost: float | None = None,
|
|
320
|
+
budget: Budget | None = None,
|
|
321
|
+
block_pii: bool | PIIConfig | None = False,
|
|
322
|
+
detect_injection: bool | InjectionConfig | None = False,
|
|
323
|
+
output_schema: Any = None,
|
|
324
|
+
validation_retries: int = 2,
|
|
325
|
+
name: str | None = "shield",
|
|
326
|
+
tags: Mapping[str, str] | None = None,
|
|
327
|
+
telemetry: bool = True,
|
|
328
|
+
) -> None:
|
|
329
|
+
self._options: dict[str, Any] = {
|
|
330
|
+
"cache": cache,
|
|
331
|
+
"retry": retry,
|
|
332
|
+
"fallback": fallback,
|
|
333
|
+
"fallback_on": fallback_on,
|
|
334
|
+
"max_cost": max_cost,
|
|
335
|
+
"budget": budget,
|
|
336
|
+
"block_pii": block_pii,
|
|
337
|
+
"detect_injection": detect_injection,
|
|
338
|
+
"output_schema": output_schema,
|
|
339
|
+
"validation_retries": validation_retries,
|
|
340
|
+
"name": name,
|
|
341
|
+
"tags": tags,
|
|
342
|
+
"telemetry": telemetry,
|
|
343
|
+
}
|
|
344
|
+
self.config = build_call_config(**self._options)
|
|
345
|
+
self._handler = build_handler(self.config)
|
|
346
|
+
|
|
347
|
+
# ----------------------------------------------------------------- context manager
|
|
348
|
+
|
|
349
|
+
def __enter__(self) -> shield:
|
|
350
|
+
ensure_installed()
|
|
351
|
+
ACTIVE.set(enter_scope(self.config, self._handler, self.config.name, owner=self))
|
|
352
|
+
return self
|
|
353
|
+
|
|
354
|
+
def __exit__(
|
|
355
|
+
self,
|
|
356
|
+
exc_type: type[BaseException] | None,
|
|
357
|
+
exc: BaseException | None,
|
|
358
|
+
tb: TracebackType | None,
|
|
359
|
+
) -> None:
|
|
360
|
+
# Restore the enclosing scope. Context variables are per task/thread, so one shield
|
|
361
|
+
# object can safely be entered concurrently from several tasks.
|
|
362
|
+
ctx = ACTIVE.get()
|
|
363
|
+
while ctx is not None and ctx.owner is not self:
|
|
364
|
+
ctx = ctx.parent
|
|
365
|
+
if ctx is not None:
|
|
366
|
+
ACTIVE.set(ctx.parent)
|
|
367
|
+
|
|
368
|
+
async def __aenter__(self) -> shield:
|
|
369
|
+
return self.__enter__()
|
|
370
|
+
|
|
371
|
+
async def __aexit__(
|
|
372
|
+
self,
|
|
373
|
+
exc_type: type[BaseException] | None,
|
|
374
|
+
exc: BaseException | None,
|
|
375
|
+
tb: TracebackType | None,
|
|
376
|
+
) -> None:
|
|
377
|
+
self.__exit__(exc_type, exc, tb)
|
|
378
|
+
|
|
379
|
+
# ----------------------------------------------------------------- direct calls
|
|
380
|
+
|
|
381
|
+
def _split(self, kwargs: dict[str, Any]) -> tuple[CallConfig, dict[str, Any]]:
|
|
382
|
+
overrides = {key: kwargs.pop(key) for key in CONFIG_FIELDS if key in kwargs}
|
|
383
|
+
config = build_call_config(**{**self._options, **overrides}) if overrides else self.config
|
|
384
|
+
return config, kwargs
|
|
385
|
+
|
|
386
|
+
def complete(
|
|
387
|
+
self,
|
|
388
|
+
*,
|
|
389
|
+
model: str | Target,
|
|
390
|
+
messages: str | Sequence[Any],
|
|
391
|
+
provider: str | None = None,
|
|
392
|
+
**kwargs: Any,
|
|
393
|
+
) -> LLMResponse:
|
|
394
|
+
"""Direct call with this shield's protections (same arguments as :func:`callm.complete`)."""
|
|
395
|
+
config, rest = self._split(dict(kwargs))
|
|
396
|
+
return run_request(config, build_request(model, messages, provider=provider, **rest))
|
|
397
|
+
|
|
398
|
+
async def acomplete(
|
|
399
|
+
self,
|
|
400
|
+
*,
|
|
401
|
+
model: str | Target,
|
|
402
|
+
messages: str | Sequence[Any],
|
|
403
|
+
provider: str | None = None,
|
|
404
|
+
**kwargs: Any,
|
|
405
|
+
) -> LLMResponse:
|
|
406
|
+
"""Async version of :meth:`complete`."""
|
|
407
|
+
config, rest = self._split(dict(kwargs))
|
|
408
|
+
return await arun_request(config, build_request(model, messages, provider=provider, **rest))
|
|
409
|
+
|
|
410
|
+
|
|
411
|
+
__all__ = ["acomplete", "arun_request", "build_request", "complete", "run_request", "shield"]
|