millforge 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.
- millforge/__init__.py +1174 -0
- millforge/_forge/LICENSE +21 -0
- millforge/_forge/PROVENANCE.json +295 -0
- millforge/_forge/UPDATE_POLICY.md +24 -0
- millforge/_forge/__init__.py +14 -0
- millforge/_forge/adapter.py +2232 -0
- millforge/_forge/base_runner.py +121 -0
- millforge/_forge/clients/__init__.py +10 -0
- millforge/_forge/clients/base.py +200 -0
- millforge/_forge/context/__init__.py +23 -0
- millforge/_forge/context/manager.py +178 -0
- millforge/_forge/context/strategies.py +335 -0
- millforge/_forge/core/__init__.py +16 -0
- millforge/_forge/core/inference.py +433 -0
- millforge/_forge/core/messages.py +119 -0
- millforge/_forge/core/runner.py +479 -0
- millforge/_forge/core/steps.py +108 -0
- millforge/_forge/core/workflow.py +400 -0
- millforge/_forge/errors.py +222 -0
- millforge/_forge/guardrails/__init__.py +21 -0
- millforge/_forge/guardrails/error_tracker.py +71 -0
- millforge/_forge/guardrails/guardrails.py +194 -0
- millforge/_forge/guardrails/nudge.py +47 -0
- millforge/_forge/guardrails/response_validator.py +119 -0
- millforge/_forge/guardrails/step_enforcer.py +183 -0
- millforge/_forge/prompts/__init__.py +16 -0
- millforge/_forge/prompts/nudges.py +95 -0
- millforge/_forge/prompts/templates.py +285 -0
- millforge/_version.py +3 -0
- millforge/artifacts.py +570 -0
- millforge/base/__init__.py +97 -0
- millforge/base/composition.py +402 -0
- millforge/base/context.py +285 -0
- millforge/base/harness.py +138 -0
- millforge/base/identity.py +465 -0
- millforge/base/options.py +34 -0
- millforge/base/platform.py +17 -0
- millforge/base/prompt.py +317 -0
- millforge/base/runner.py +546 -0
- millforge/compiled_plan.py +970 -0
- millforge/compiler/__init__.py +231 -0
- millforge/compiler/artifact_validation.py +257 -0
- millforge/compiler/canonicalization.py +169 -0
- millforge/compiler/capabilities.py +66 -0
- millforge/compiler/catalogs.py +500 -0
- millforge/compiler/diagnostics.py +491 -0
- millforge/compiler/graph.py +678 -0
- millforge/compiler/lowering.py +198 -0
- millforge/compiler/output.py +692 -0
- millforge/compiler/parsing.py +1424 -0
- millforge/compiler/requests.py +1180 -0
- millforge/compiler/schema_validation.py +272 -0
- millforge/compiler/semantic.py +490 -0
- millforge/compiler/service.py +448 -0
- millforge/compiler/source.py +375 -0
- millforge/compiler/validators.py +184 -0
- millforge/connectors/__init__.py +95 -0
- millforge/connectors/admission.py +801 -0
- millforge/connectors/broker.py +202 -0
- millforge/connectors/contracts.py +1159 -0
- millforge/connectors/diagnostics.py +189 -0
- millforge/connectors/fake.py +66 -0
- millforge/connectors/runtime.py +236 -0
- millforge/contracts.py +2860 -0
- millforge/custom_tools/__init__.py +67 -0
- millforge/custom_tools/compiler.py +724 -0
- millforge/custom_tools/contracts.py +1093 -0
- millforge/custom_tools/diagnostics.py +205 -0
- millforge/eval_artifacts.py +952 -0
- millforge/eval_boundary.py +2435 -0
- millforge/eval_fixtures/__init__.py +1 -0
- millforge/eval_fixtures/default_pack/__init__.py +1 -0
- millforge/eval_fixtures/default_pack/fixtures/fixture.08a.bug_diagnosis.traceback.v1.json +52 -0
- millforge/eval_fixtures/default_pack/fixtures/fixture.08a.direct_edit.import_sort.v1.json +52 -0
- millforge/eval_fixtures/default_pack/fixtures/fixture.08a.evidence_discipline.no_source_change.v1.json +51 -0
- millforge/eval_fixtures/default_pack/fixtures/fixture.08a.false_closure.visible_green.v1.json +52 -0
- millforge/eval_fixtures/default_pack/fixtures/fixture.08a.multi_file.api_contract.v1.json +54 -0
- millforge/eval_fixtures/default_pack/fixtures/fixture.08a.recovery.malformed_artifact.v1.json +54 -0
- millforge/eval_fixtures/default_pack/manifest.json +12 -0
- millforge/eval_modes.py +1282 -0
- millforge/eval_presets.py +1398 -0
- millforge/eval_reports.py +2517 -0
- millforge/eval_suite.py +2429 -0
- millforge/eval_trials.py +2632 -0
- millforge/eval_workflow.py +794 -0
- millforge/exceptions.py +122 -0
- millforge/model_backend.py +2098 -0
- millforge/protocols.py +340 -0
- millforge/py.typed +0 -0
- millforge/runtime.py +1791 -0
- millforge/testing/__init__.py +1089 -0
- millforge/tools/__init__.py +83 -0
- millforge/tools/builtin_runtime.py +1339 -0
- millforge/tools/builtins.py +773 -0
- millforge/tools/execution.py +1545 -0
- millforge/tools/path_policy.py +155 -0
- millforge/tools/pi_compat/PI_LICENSE +21 -0
- millforge/tools/pi_compat/PROVENANCE.json +55 -0
- millforge/tools/pi_compat/UPDATE_POLICY.md +36 -0
- millforge/tools/pi_compat/__init__.py +34 -0
- millforge/tools/pi_compat/contracts.py +49 -0
- millforge/tools/pi_compat/editing.py +390 -0
- millforge/tools/pi_compat/mutations.py +57 -0
- millforge/tools/pi_compat/operations.py +401 -0
- millforge/tools/pi_compat/paths.py +155 -0
- millforge/tools/pi_compat/process.py +1375 -0
- millforge/tools/pi_compat/search.py +738 -0
- millforge/tools/pi_compat/truncation.py +267 -0
- millforge/tools/pi_compat_catalog.py +396 -0
- millforge/tools/pi_compat_runtime.py +460 -0
- millforge/tools/registry.py +553 -0
- millforge/tools/results.py +533 -0
- millforge-0.1.0.dist-info/METADATA +844 -0
- millforge-0.1.0.dist-info/RECORD +116 -0
- millforge-0.1.0.dist-info/WHEEL +4 -0
- millforge-0.1.0.dist-info/licenses/LICENSE +201 -0
|
@@ -0,0 +1,2098 @@
|
|
|
1
|
+
"""Provider-neutral model backend contracts and private orchestration.
|
|
2
|
+
|
|
3
|
+
The immutable profile, policy, timeout, and secret-resolution contracts needed
|
|
4
|
+
by the supported live factory are re-exported deliberately from ``millforge``.
|
|
5
|
+
Concrete model clients, HTTP transports, wire records, and orchestration helpers
|
|
6
|
+
remain private implementation details and are not public package exports.
|
|
7
|
+
"""
|
|
8
|
+
|
|
9
|
+
from __future__ import annotations
|
|
10
|
+
|
|
11
|
+
import asyncio
|
|
12
|
+
import inspect
|
|
13
|
+
import json
|
|
14
|
+
import math
|
|
15
|
+
import time
|
|
16
|
+
from dataclasses import dataclass, field
|
|
17
|
+
from enum import Enum
|
|
18
|
+
from types import MappingProxyType
|
|
19
|
+
from typing import Any, Literal, Mapping, Protocol, TypeAlias, cast, runtime_checkable
|
|
20
|
+
from urllib.parse import parse_qsl, urlsplit, urlunsplit
|
|
21
|
+
|
|
22
|
+
import httpx
|
|
23
|
+
from pydantic import (
|
|
24
|
+
BaseModel,
|
|
25
|
+
ConfigDict,
|
|
26
|
+
Field,
|
|
27
|
+
field_validator,
|
|
28
|
+
model_serializer,
|
|
29
|
+
model_validator,
|
|
30
|
+
)
|
|
31
|
+
|
|
32
|
+
from millforge.contracts import (
|
|
33
|
+
AssistantMessage,
|
|
34
|
+
InvalidToolArguments,
|
|
35
|
+
JsonObject,
|
|
36
|
+
JsonValue,
|
|
37
|
+
ModelCapabilityRequirements,
|
|
38
|
+
ModelCompletionRequest,
|
|
39
|
+
ModelCompletionResponse,
|
|
40
|
+
ModelMessage,
|
|
41
|
+
ModelToolCall,
|
|
42
|
+
ParsedToolArguments,
|
|
43
|
+
RedactionPolicy,
|
|
44
|
+
SanitizedMetadataValue,
|
|
45
|
+
SamplingRequest,
|
|
46
|
+
SecretRef,
|
|
47
|
+
TokenUsage,
|
|
48
|
+
ToolResultMessage,
|
|
49
|
+
redact_diagnostic_mapping,
|
|
50
|
+
redact_diagnostic_text,
|
|
51
|
+
redact_diagnostic_value,
|
|
52
|
+
)
|
|
53
|
+
from millforge.exceptions import (
|
|
54
|
+
MillforgeConfigError,
|
|
55
|
+
ModelTransportError,
|
|
56
|
+
)
|
|
57
|
+
from millforge.protocols import AsyncHttpTransport
|
|
58
|
+
|
|
59
|
+
_CHAT_COMPLETIONS_SUFFIX = "/chat/completions"
|
|
60
|
+
DEFAULT_REDACTION_POLICY = RedactionPolicy()
|
|
61
|
+
_MAX_SANITIZED_VALUE_LENGTH = min(DEFAULT_REDACTION_POLICY.max_string_length, 512)
|
|
62
|
+
_MAX_SANITIZED_FIELDS = min(DEFAULT_REDACTION_POLICY.max_collection_items, 24)
|
|
63
|
+
_FORBIDDEN_CUSTOM_AUTH_HEADERS = {
|
|
64
|
+
"accept",
|
|
65
|
+
"authorization",
|
|
66
|
+
"proxy-authorization",
|
|
67
|
+
"cookie",
|
|
68
|
+
"set-cookie",
|
|
69
|
+
"host",
|
|
70
|
+
"content-type",
|
|
71
|
+
"user-agent",
|
|
72
|
+
}
|
|
73
|
+
JsonHeaders: TypeAlias = dict[str, str]
|
|
74
|
+
_FINISH_REASON_MAP = {
|
|
75
|
+
"stop": "stop",
|
|
76
|
+
"tool_calls": "tool_calls",
|
|
77
|
+
"function_call": "tool_calls",
|
|
78
|
+
"length": "length",
|
|
79
|
+
"content_filter": "content_filter",
|
|
80
|
+
"cancelled": "cancelled",
|
|
81
|
+
}
|
|
82
|
+
_SUCCESS_BODY_LIMIT_BYTES = 4 * 1024 * 1024
|
|
83
|
+
_ERROR_BODY_LIMIT_BYTES = 64 * 1024
|
|
84
|
+
_DIRECT_SAMPLING_BODY_FIELDS = {
|
|
85
|
+
"temperature",
|
|
86
|
+
"top_p",
|
|
87
|
+
"presence_penalty",
|
|
88
|
+
"frequency_penalty",
|
|
89
|
+
"seed",
|
|
90
|
+
"stop",
|
|
91
|
+
}
|
|
92
|
+
|
|
93
|
+
|
|
94
|
+
def _nonblank(value: str, field_name: str) -> str:
|
|
95
|
+
if not value.strip():
|
|
96
|
+
raise ValueError(f"{field_name} must be a non-empty string")
|
|
97
|
+
return value
|
|
98
|
+
|
|
99
|
+
|
|
100
|
+
def _unique(values: tuple[str, ...], field_name: str) -> None:
|
|
101
|
+
if len(set(values)) != len(values):
|
|
102
|
+
raise ValueError(f"{field_name} values must be unique")
|
|
103
|
+
|
|
104
|
+
|
|
105
|
+
def _bounded(value: str, *, length: int = _MAX_SANITIZED_VALUE_LENGTH) -> str:
|
|
106
|
+
return value if len(value) <= length else f"{value[:length]}...[truncated]"
|
|
107
|
+
|
|
108
|
+
|
|
109
|
+
class ModelBackendConfigError(MillforgeConfigError):
|
|
110
|
+
"""Invalid internal model backend configuration."""
|
|
111
|
+
|
|
112
|
+
|
|
113
|
+
@dataclass(frozen=True, slots=True)
|
|
114
|
+
class OpenAICompatibleTimeouts:
|
|
115
|
+
"""Explicit positive finite timeout authority for a live composition.
|
|
116
|
+
|
|
117
|
+
The four phase bounds configure the HTTP transport. ``local_total_seconds``
|
|
118
|
+
is an additional composition-owned ceiling for the complete model call; it
|
|
119
|
+
may narrow, but never widen, the request deadline or resolved profile bound.
|
|
120
|
+
"""
|
|
121
|
+
|
|
122
|
+
connect_seconds: float
|
|
123
|
+
read_seconds: float
|
|
124
|
+
write_seconds: float
|
|
125
|
+
pool_seconds: float
|
|
126
|
+
local_total_seconds: float
|
|
127
|
+
|
|
128
|
+
def __post_init__(self) -> None:
|
|
129
|
+
for field_name in (
|
|
130
|
+
"connect_seconds",
|
|
131
|
+
"read_seconds",
|
|
132
|
+
"write_seconds",
|
|
133
|
+
"pool_seconds",
|
|
134
|
+
"local_total_seconds",
|
|
135
|
+
):
|
|
136
|
+
value = getattr(self, field_name)
|
|
137
|
+
if isinstance(value, bool) or not isinstance(value, (int, float)):
|
|
138
|
+
raise ModelBackendConfigError(
|
|
139
|
+
f"{field_name} must be a positive finite number"
|
|
140
|
+
)
|
|
141
|
+
if value <= 0 or not math.isfinite(value):
|
|
142
|
+
raise ModelBackendConfigError(
|
|
143
|
+
f"{field_name} must be a positive finite number"
|
|
144
|
+
)
|
|
145
|
+
|
|
146
|
+
@classmethod
|
|
147
|
+
def uniform(cls, timeout_seconds: float) -> OpenAICompatibleTimeouts:
|
|
148
|
+
"""Create equal phase and local bounds for compatibility callers."""
|
|
149
|
+
return cls(
|
|
150
|
+
connect_seconds=timeout_seconds,
|
|
151
|
+
read_seconds=timeout_seconds,
|
|
152
|
+
write_seconds=timeout_seconds,
|
|
153
|
+
pool_seconds=timeout_seconds,
|
|
154
|
+
local_total_seconds=timeout_seconds,
|
|
155
|
+
)
|
|
156
|
+
|
|
157
|
+
|
|
158
|
+
class UnsupportedModelCapabilityError(ModelBackendConfigError):
|
|
159
|
+
"""A required model capability is not supported by the resolved profile."""
|
|
160
|
+
|
|
161
|
+
|
|
162
|
+
class SecretResolutionError(ModelBackendConfigError):
|
|
163
|
+
"""A configured model secret cannot be safely resolved."""
|
|
164
|
+
|
|
165
|
+
|
|
166
|
+
class ProviderErrorCategory(str, Enum):
|
|
167
|
+
"""Stable provider-neutral provider error categories."""
|
|
168
|
+
|
|
169
|
+
AUTHENTICATION = "authentication"
|
|
170
|
+
AUTHORIZATION = "authorization"
|
|
171
|
+
RATE_LIMIT = "rate_limit"
|
|
172
|
+
TIMEOUT = "timeout"
|
|
173
|
+
CONNECTION = "connection"
|
|
174
|
+
INVALID_REQUEST = "invalid_request"
|
|
175
|
+
UNSUPPORTED_CAPABILITY = "unsupported_capability"
|
|
176
|
+
MALFORMED_RESPONSE = "malformed_response"
|
|
177
|
+
SERVER_ERROR = "server_error"
|
|
178
|
+
CANCELLED = "cancelled"
|
|
179
|
+
UNKNOWN = "unknown"
|
|
180
|
+
|
|
181
|
+
|
|
182
|
+
_RETRYABLE_CATEGORIES = {
|
|
183
|
+
ProviderErrorCategory.RATE_LIMIT,
|
|
184
|
+
ProviderErrorCategory.TIMEOUT,
|
|
185
|
+
ProviderErrorCategory.CONNECTION,
|
|
186
|
+
ProviderErrorCategory.SERVER_ERROR,
|
|
187
|
+
}
|
|
188
|
+
|
|
189
|
+
|
|
190
|
+
@dataclass(frozen=True, slots=True)
|
|
191
|
+
class ModelProviderError(ModelTransportError):
|
|
192
|
+
"""Sanitized provider error data.
|
|
193
|
+
|
|
194
|
+
``message`` and ``fields`` must already be redacted. Raw response bodies,
|
|
195
|
+
headers, exception reprs, and secret values are not accepted here.
|
|
196
|
+
"""
|
|
197
|
+
|
|
198
|
+
category: ProviderErrorCategory
|
|
199
|
+
message: str
|
|
200
|
+
retryable: bool | None = None
|
|
201
|
+
provider_request_id: str | None = None
|
|
202
|
+
fields: Mapping[str, SanitizedMetadataValue] = field(default_factory=dict)
|
|
203
|
+
|
|
204
|
+
def __post_init__(self) -> None:
|
|
205
|
+
object.__setattr__(self, "message", _bounded(redact_text(self.message)))
|
|
206
|
+
if self.provider_request_id is not None:
|
|
207
|
+
object.__setattr__(
|
|
208
|
+
self,
|
|
209
|
+
"provider_request_id",
|
|
210
|
+
_bounded(redact_text(self.provider_request_id), length=128),
|
|
211
|
+
)
|
|
212
|
+
if self.retryable is None:
|
|
213
|
+
object.__setattr__(
|
|
214
|
+
self, "retryable", self.category in _RETRYABLE_CATEGORIES
|
|
215
|
+
)
|
|
216
|
+
sanitized = sanitize_provider_error_fields(self.fields)
|
|
217
|
+
object.__setattr__(self, "fields", MappingProxyType(sanitized))
|
|
218
|
+
Exception.__init__(self, self.message)
|
|
219
|
+
|
|
220
|
+
def __str__(self) -> str:
|
|
221
|
+
return self.message
|
|
222
|
+
|
|
223
|
+
def __repr__(self) -> str:
|
|
224
|
+
return (
|
|
225
|
+
"ModelProviderError("
|
|
226
|
+
f"category={self.category.value!r}, retryable={self.retryable!r}, "
|
|
227
|
+
f"provider_request_id={self.provider_request_id!r}, fields={dict(self.fields)!r})"
|
|
228
|
+
)
|
|
229
|
+
|
|
230
|
+
|
|
231
|
+
class ModelRequestDeadlineExceededError(ModelProviderError):
|
|
232
|
+
"""The client-owned effective request deadline expired."""
|
|
233
|
+
|
|
234
|
+
__slots__ = ()
|
|
235
|
+
|
|
236
|
+
def __init__(self) -> None:
|
|
237
|
+
super().__init__(
|
|
238
|
+
category=ProviderErrorCategory.TIMEOUT,
|
|
239
|
+
message="model request deadline expired",
|
|
240
|
+
retryable=False,
|
|
241
|
+
)
|
|
242
|
+
|
|
243
|
+
|
|
244
|
+
class AuthenticationScheme(str, Enum):
|
|
245
|
+
"""Supported internal authentication policies."""
|
|
246
|
+
|
|
247
|
+
NONE = "none"
|
|
248
|
+
BEARER = "bearer"
|
|
249
|
+
HEADER = "header"
|
|
250
|
+
|
|
251
|
+
|
|
252
|
+
class CapabilitySupport(str, Enum):
|
|
253
|
+
"""Tri-state capability support declaration."""
|
|
254
|
+
|
|
255
|
+
SUPPORTED = "supported"
|
|
256
|
+
UNSUPPORTED = "unsupported"
|
|
257
|
+
UNKNOWN = "unknown"
|
|
258
|
+
|
|
259
|
+
|
|
260
|
+
class ReasoningSupport(str, Enum):
|
|
261
|
+
"""Provider-neutral reasoning support declaration."""
|
|
262
|
+
|
|
263
|
+
UNSUPPORTED = "unsupported"
|
|
264
|
+
OPTIONAL = "optional"
|
|
265
|
+
REQUIRED = "required"
|
|
266
|
+
|
|
267
|
+
|
|
268
|
+
class ReasoningMode(str, Enum):
|
|
269
|
+
"""Canonical provider-neutral reasoning intent."""
|
|
270
|
+
|
|
271
|
+
DISABLED = "disabled"
|
|
272
|
+
ENABLED = "enabled"
|
|
273
|
+
REQUIRED = "required"
|
|
274
|
+
|
|
275
|
+
|
|
276
|
+
class ReasoningEffort(str, Enum):
|
|
277
|
+
"""Canonical provider-neutral reasoning effort levels."""
|
|
278
|
+
|
|
279
|
+
LOW = "low"
|
|
280
|
+
MEDIUM = "medium"
|
|
281
|
+
HIGH = "high"
|
|
282
|
+
MAX = "max"
|
|
283
|
+
|
|
284
|
+
|
|
285
|
+
class EndpointConfig(BaseModel):
|
|
286
|
+
"""Immutable normalized endpoint configuration."""
|
|
287
|
+
|
|
288
|
+
model_config = ConfigDict(extra="forbid", frozen=True)
|
|
289
|
+
|
|
290
|
+
base_url: str
|
|
291
|
+
allow_insecure_local: bool = False
|
|
292
|
+
success_content_types: tuple[str, ...] = ("application/json",)
|
|
293
|
+
allow_missing_success_content_type: bool = False
|
|
294
|
+
|
|
295
|
+
@field_validator("base_url")
|
|
296
|
+
@classmethod
|
|
297
|
+
def _base_url_nonblank(cls, value: str) -> str:
|
|
298
|
+
return _nonblank(value, "base_url").rstrip("/")
|
|
299
|
+
|
|
300
|
+
@field_validator("success_content_types")
|
|
301
|
+
@classmethod
|
|
302
|
+
def _content_types_valid(cls, value: tuple[str, ...]) -> tuple[str, ...]:
|
|
303
|
+
if not value:
|
|
304
|
+
raise ValueError("success_content_types must not be empty")
|
|
305
|
+
normalized = tuple(item.strip().lower() for item in value)
|
|
306
|
+
for item in normalized:
|
|
307
|
+
_nonblank(item, "success_content_types")
|
|
308
|
+
_unique(normalized, "success_content_types")
|
|
309
|
+
return normalized
|
|
310
|
+
|
|
311
|
+
@model_validator(mode="after")
|
|
312
|
+
def _endpoint_safe(self) -> EndpointConfig:
|
|
313
|
+
split = urlsplit(
|
|
314
|
+
self.base_url if "://" in self.base_url else f"https://{self.base_url}"
|
|
315
|
+
)
|
|
316
|
+
if split.scheme not in {"https", "http"}:
|
|
317
|
+
raise ValueError(
|
|
318
|
+
"base_url scheme must be https or explicitly allowed local http"
|
|
319
|
+
)
|
|
320
|
+
if split.username or split.password:
|
|
321
|
+
raise ValueError("base_url must not contain userinfo")
|
|
322
|
+
if split.query or split.fragment:
|
|
323
|
+
raise ValueError("base_url must not contain query strings or fragments")
|
|
324
|
+
if not split.netloc:
|
|
325
|
+
raise ValueError("base_url must include a host")
|
|
326
|
+
if split.path.rstrip("/").endswith(_CHAT_COMPLETIONS_SUFFIX):
|
|
327
|
+
raise ValueError("base_url must not include /chat/completions")
|
|
328
|
+
if split.scheme == "http" and not (
|
|
329
|
+
self.allow_insecure_local
|
|
330
|
+
and split.hostname in {"localhost", "127.0.0.1", "::1"}
|
|
331
|
+
):
|
|
332
|
+
raise ValueError("http base_url is allowed only for explicit local testing")
|
|
333
|
+
normalized = urlunsplit(
|
|
334
|
+
(split.scheme, split.netloc, split.path.rstrip("/"), "", "")
|
|
335
|
+
)
|
|
336
|
+
object.__setattr__(self, "base_url", normalized)
|
|
337
|
+
return self
|
|
338
|
+
|
|
339
|
+
@property
|
|
340
|
+
def chat_completions_url(self) -> str:
|
|
341
|
+
"""Return the exact Chat Completions URL for this API prefix."""
|
|
342
|
+
return f"{self.base_url}{_CHAT_COMPLETIONS_SUFFIX}"
|
|
343
|
+
|
|
344
|
+
|
|
345
|
+
class HeaderValuePolicy(BaseModel):
|
|
346
|
+
"""Configured non-secret headers admitted for transport requests."""
|
|
347
|
+
|
|
348
|
+
model_config = ConfigDict(extra="forbid", frozen=True)
|
|
349
|
+
|
|
350
|
+
values: dict[str, str] = Field(default_factory=dict)
|
|
351
|
+
|
|
352
|
+
@field_validator("values")
|
|
353
|
+
@classmethod
|
|
354
|
+
def _headers_safe(cls, value: dict[str, str]) -> dict[str, str]:
|
|
355
|
+
seen: set[str] = set()
|
|
356
|
+
for name, header_value in value.items():
|
|
357
|
+
normalized = _nonblank(name, "header name").lower()
|
|
358
|
+
if normalized in seen:
|
|
359
|
+
raise ValueError("header names must be unique case-insensitively")
|
|
360
|
+
seen.add(normalized)
|
|
361
|
+
if normalized in _FORBIDDEN_CUSTOM_AUTH_HEADERS:
|
|
362
|
+
raise ValueError(f"configured header {name!r} is protected")
|
|
363
|
+
_nonblank(header_value, "header value")
|
|
364
|
+
return dict(value)
|
|
365
|
+
|
|
366
|
+
|
|
367
|
+
class AuthenticationPolicy(BaseModel):
|
|
368
|
+
"""Secret-safe internal authentication policy."""
|
|
369
|
+
|
|
370
|
+
model_config = ConfigDict(extra="forbid", frozen=True)
|
|
371
|
+
|
|
372
|
+
scheme: AuthenticationScheme
|
|
373
|
+
secret_ref: SecretRef | None = None
|
|
374
|
+
header_name: str | None = None
|
|
375
|
+
allowed_custom_header_names: tuple[str, ...] = ()
|
|
376
|
+
|
|
377
|
+
@field_validator("header_name")
|
|
378
|
+
@classmethod
|
|
379
|
+
def _header_name_valid(cls, value: str | None) -> str | None:
|
|
380
|
+
if value is None:
|
|
381
|
+
return None
|
|
382
|
+
return _nonblank(value, "header_name")
|
|
383
|
+
|
|
384
|
+
@field_validator("allowed_custom_header_names")
|
|
385
|
+
@classmethod
|
|
386
|
+
def _custom_names_valid(cls, value: tuple[str, ...]) -> tuple[str, ...]:
|
|
387
|
+
normalized = tuple(
|
|
388
|
+
_nonblank(item, "allowed_custom_header_names").lower() for item in value
|
|
389
|
+
)
|
|
390
|
+
_unique(normalized, "allowed_custom_header_names")
|
|
391
|
+
if any(item in _FORBIDDEN_CUSTOM_AUTH_HEADERS for item in normalized):
|
|
392
|
+
raise ValueError("allowed custom authentication header is protected")
|
|
393
|
+
return normalized
|
|
394
|
+
|
|
395
|
+
@model_validator(mode="after")
|
|
396
|
+
def _auth_consistent(self) -> AuthenticationPolicy:
|
|
397
|
+
if self.scheme is AuthenticationScheme.NONE:
|
|
398
|
+
if self.secret_ref is not None or self.header_name is not None:
|
|
399
|
+
raise ValueError(
|
|
400
|
+
"none authentication must not configure a secret or header"
|
|
401
|
+
)
|
|
402
|
+
elif self.scheme is AuthenticationScheme.BEARER:
|
|
403
|
+
if self.secret_ref is None:
|
|
404
|
+
raise ValueError("bearer authentication requires secret_ref")
|
|
405
|
+
if self.header_name is not None:
|
|
406
|
+
raise ValueError("bearer authentication always uses Authorization")
|
|
407
|
+
else:
|
|
408
|
+
if self.secret_ref is None or self.header_name is None:
|
|
409
|
+
raise ValueError(
|
|
410
|
+
"header authentication requires secret_ref and header_name"
|
|
411
|
+
)
|
|
412
|
+
normalized = self.header_name.lower()
|
|
413
|
+
if normalized in _FORBIDDEN_CUSTOM_AUTH_HEADERS:
|
|
414
|
+
raise ValueError("custom authentication header is protected")
|
|
415
|
+
if normalized not in self.allowed_custom_header_names:
|
|
416
|
+
raise ValueError("custom authentication header must be allowlisted")
|
|
417
|
+
return self
|
|
418
|
+
|
|
419
|
+
|
|
420
|
+
class SamplingPolicy(BaseModel):
|
|
421
|
+
"""Resolved default sampling policy and request override allowlist."""
|
|
422
|
+
|
|
423
|
+
model_config = ConfigDict(extra="forbid", frozen=True)
|
|
424
|
+
|
|
425
|
+
temperature: float | None = Field(default=None, ge=0, le=2)
|
|
426
|
+
top_p: float | None = Field(default=None, ge=0, le=1)
|
|
427
|
+
presence_penalty: float | None = Field(default=None, ge=-2, le=2)
|
|
428
|
+
frequency_penalty: float | None = Field(default=None, ge=-2, le=2)
|
|
429
|
+
seed: int | None = None
|
|
430
|
+
stop: tuple[str, ...] | None = None
|
|
431
|
+
allowed_overrides: tuple[
|
|
432
|
+
str,
|
|
433
|
+
...,
|
|
434
|
+
] = (
|
|
435
|
+
"temperature",
|
|
436
|
+
"top_p",
|
|
437
|
+
"presence_penalty",
|
|
438
|
+
"frequency_penalty",
|
|
439
|
+
"seed",
|
|
440
|
+
"stop",
|
|
441
|
+
)
|
|
442
|
+
allow_maximum_output_tokens_override: bool = True
|
|
443
|
+
|
|
444
|
+
@model_validator(mode="before")
|
|
445
|
+
@classmethod
|
|
446
|
+
def _accept_legacy_defaults(cls, data: Any) -> Any:
|
|
447
|
+
if not isinstance(data, dict):
|
|
448
|
+
return data
|
|
449
|
+
copied = dict(data)
|
|
450
|
+
defaults = copied.pop("defaults", None)
|
|
451
|
+
copied.pop("default_maximum_output_tokens", None)
|
|
452
|
+
if defaults is not None:
|
|
453
|
+
default_values = (
|
|
454
|
+
defaults.model_dump(exclude_none=True)
|
|
455
|
+
if isinstance(defaults, SamplingRequest)
|
|
456
|
+
else SamplingRequest.model_validate(defaults).model_dump(
|
|
457
|
+
exclude_none=True
|
|
458
|
+
)
|
|
459
|
+
)
|
|
460
|
+
for name in _DIRECT_SAMPLING_BODY_FIELDS:
|
|
461
|
+
if name in default_values and name not in copied:
|
|
462
|
+
copied[name] = default_values[name]
|
|
463
|
+
return copied
|
|
464
|
+
|
|
465
|
+
@field_validator("allowed_overrides")
|
|
466
|
+
@classmethod
|
|
467
|
+
def _allowed_overrides_valid(cls, value: tuple[str, ...]) -> tuple[str, ...]:
|
|
468
|
+
allowed = set(_DIRECT_SAMPLING_BODY_FIELDS)
|
|
469
|
+
for item in value:
|
|
470
|
+
if item not in allowed:
|
|
471
|
+
raise ValueError(f"unknown sampling override {item!r}")
|
|
472
|
+
_unique(value, "allowed_overrides")
|
|
473
|
+
return value
|
|
474
|
+
|
|
475
|
+
@field_validator("stop")
|
|
476
|
+
@classmethod
|
|
477
|
+
def _stop_values_nonblank(
|
|
478
|
+
cls, value: tuple[str, ...] | None
|
|
479
|
+
) -> tuple[str, ...] | None:
|
|
480
|
+
if value is None:
|
|
481
|
+
return None
|
|
482
|
+
for item in value:
|
|
483
|
+
_nonblank(item, "stop")
|
|
484
|
+
return value
|
|
485
|
+
|
|
486
|
+
|
|
487
|
+
class ReasoningPolicy(BaseModel):
|
|
488
|
+
"""Provider-neutral reasoning intent with configured wire mappings."""
|
|
489
|
+
|
|
490
|
+
model_config = ConfigDict(extra="forbid", frozen=True)
|
|
491
|
+
|
|
492
|
+
mode: ReasoningMode = ReasoningMode.DISABLED
|
|
493
|
+
effort: ReasoningEffort | None = None
|
|
494
|
+
mode_field: str | None = None
|
|
495
|
+
effort_field: str | None = None
|
|
496
|
+
mode_values: dict[ReasoningMode, JsonValue] = Field(default_factory=dict)
|
|
497
|
+
effort_values: dict[ReasoningEffort, JsonValue] = Field(default_factory=dict)
|
|
498
|
+
tool_call_replay_field: Literal["reasoning_content"] | None = None
|
|
499
|
+
|
|
500
|
+
@model_validator(mode="before")
|
|
501
|
+
@classmethod
|
|
502
|
+
def _accept_legacy_reasoning(cls, data: Any) -> Any:
|
|
503
|
+
if not isinstance(data, dict):
|
|
504
|
+
return data
|
|
505
|
+
copied = dict(data)
|
|
506
|
+
support = copied.pop("support", None)
|
|
507
|
+
mode_parameter = copied.pop("mode_parameter", None)
|
|
508
|
+
effort_parameter = copied.pop("effort_parameter", None)
|
|
509
|
+
allowed_modes = tuple(copied.pop("allowed_modes", ()) or ())
|
|
510
|
+
allowed_efforts = tuple(copied.pop("allowed_efforts", ()) or ())
|
|
511
|
+
if support is not None and "mode" not in copied:
|
|
512
|
+
support_value = getattr(support, "value", support)
|
|
513
|
+
copied["mode"] = (
|
|
514
|
+
ReasoningMode.DISABLED
|
|
515
|
+
if support_value == ReasoningSupport.UNSUPPORTED.value
|
|
516
|
+
else ReasoningMode.REQUIRED
|
|
517
|
+
if support_value == ReasoningSupport.REQUIRED.value
|
|
518
|
+
else ReasoningMode.ENABLED
|
|
519
|
+
)
|
|
520
|
+
if mode_parameter is not None and "mode_field" not in copied:
|
|
521
|
+
copied["mode_field"] = mode_parameter
|
|
522
|
+
if effort_parameter is not None and "effort_field" not in copied:
|
|
523
|
+
copied["effort_field"] = effort_parameter
|
|
524
|
+
if allowed_modes and "mode_values" not in copied:
|
|
525
|
+
copied["mode_values"] = {
|
|
526
|
+
ReasoningMode.ENABLED: allowed_modes[0],
|
|
527
|
+
ReasoningMode.REQUIRED: allowed_modes[0],
|
|
528
|
+
}
|
|
529
|
+
if allowed_efforts and "effort_values" not in copied:
|
|
530
|
+
effort_values: dict[ReasoningEffort, str] = {}
|
|
531
|
+
for index, item in enumerate(allowed_efforts):
|
|
532
|
+
try:
|
|
533
|
+
effort = ReasoningEffort(item)
|
|
534
|
+
except ValueError:
|
|
535
|
+
effort_keys = tuple(ReasoningEffort)
|
|
536
|
+
if index >= len(effort_keys):
|
|
537
|
+
break
|
|
538
|
+
effort = effort_keys[index]
|
|
539
|
+
effort_values[effort] = item
|
|
540
|
+
copied["effort_values"] = effort_values
|
|
541
|
+
if allowed_efforts and "effort" not in copied:
|
|
542
|
+
try:
|
|
543
|
+
copied["effort"] = ReasoningEffort(allowed_efforts[0])
|
|
544
|
+
except ValueError:
|
|
545
|
+
copied["effort"] = ReasoningEffort.LOW
|
|
546
|
+
return copied
|
|
547
|
+
|
|
548
|
+
@model_validator(mode="after")
|
|
549
|
+
def _reasoning_mapping_consistent(self) -> ReasoningPolicy:
|
|
550
|
+
if self.mode is not ReasoningMode.DISABLED and self.mode_field is not None:
|
|
551
|
+
if self.mode not in self.mode_values:
|
|
552
|
+
raise ValueError("reasoning mode needs a configured wire value")
|
|
553
|
+
if self.effort is not None:
|
|
554
|
+
if self.effort_field is None:
|
|
555
|
+
raise ValueError("reasoning effort needs a configured wire field")
|
|
556
|
+
if self.effort not in self.effort_values:
|
|
557
|
+
raise ValueError("reasoning effort needs a configured wire value")
|
|
558
|
+
if self.tool_call_replay_field is not None:
|
|
559
|
+
if self.mode not in (ReasoningMode.ENABLED, ReasoningMode.REQUIRED):
|
|
560
|
+
raise ValueError(
|
|
561
|
+
"reasoning replay requires enabled or required reasoning"
|
|
562
|
+
)
|
|
563
|
+
if self.mode_field is None or not self.mode_field.strip():
|
|
564
|
+
raise ValueError("reasoning replay needs a configured wire field")
|
|
565
|
+
if self.mode not in self.mode_values:
|
|
566
|
+
raise ValueError("reasoning replay needs a configured wire value")
|
|
567
|
+
return self
|
|
568
|
+
|
|
569
|
+
@model_serializer(mode="wrap")
|
|
570
|
+
def _omit_absent_replay_field(self, handler: Any) -> dict[str, Any]:
|
|
571
|
+
payload = handler(self)
|
|
572
|
+
if self.tool_call_replay_field is None:
|
|
573
|
+
payload.pop("tool_call_replay_field", None)
|
|
574
|
+
return payload
|
|
575
|
+
|
|
576
|
+
|
|
577
|
+
class CapabilityDeclarations(BaseModel):
|
|
578
|
+
"""Resolved tri-state capability declarations."""
|
|
579
|
+
|
|
580
|
+
model_config = ConfigDict(extra="forbid", frozen=True)
|
|
581
|
+
|
|
582
|
+
support: dict[str, CapabilitySupport] = Field(default_factory=dict)
|
|
583
|
+
|
|
584
|
+
@field_validator("support")
|
|
585
|
+
@classmethod
|
|
586
|
+
def _capabilities_valid(
|
|
587
|
+
cls, value: dict[str, CapabilitySupport]
|
|
588
|
+
) -> dict[str, CapabilitySupport]:
|
|
589
|
+
for name in value:
|
|
590
|
+
_nonblank(name, "capability name")
|
|
591
|
+
return dict(value)
|
|
592
|
+
|
|
593
|
+
def state_for(self, capability_name: str) -> CapabilitySupport:
|
|
594
|
+
"""Return the explicit support state or ``unknown`` when absent."""
|
|
595
|
+
return self.support.get(capability_name, CapabilitySupport.UNKNOWN)
|
|
596
|
+
|
|
597
|
+
|
|
598
|
+
class RequestOptionAllowlist(BaseModel):
|
|
599
|
+
"""Provider-neutral extra request option allowlist."""
|
|
600
|
+
|
|
601
|
+
model_config = ConfigDict(extra="forbid", frozen=True)
|
|
602
|
+
|
|
603
|
+
allowed_options: tuple[str, ...] = ()
|
|
604
|
+
|
|
605
|
+
@field_validator("allowed_options")
|
|
606
|
+
@classmethod
|
|
607
|
+
def _options_valid(cls, value: tuple[str, ...]) -> tuple[str, ...]:
|
|
608
|
+
allowed = {"tool_choice", "parallel_tool_calls", "response_format", "user"}
|
|
609
|
+
protected = {
|
|
610
|
+
"model",
|
|
611
|
+
"messages",
|
|
612
|
+
"tools",
|
|
613
|
+
"stream",
|
|
614
|
+
"endpoint",
|
|
615
|
+
"authentication",
|
|
616
|
+
"timeout",
|
|
617
|
+
"host",
|
|
618
|
+
"headers",
|
|
619
|
+
"content_type",
|
|
620
|
+
"user_agent",
|
|
621
|
+
"max_tokens",
|
|
622
|
+
"maximum_output_tokens",
|
|
623
|
+
"temperature",
|
|
624
|
+
"top_p",
|
|
625
|
+
"presence_penalty",
|
|
626
|
+
"frequency_penalty",
|
|
627
|
+
"seed",
|
|
628
|
+
"stop",
|
|
629
|
+
}
|
|
630
|
+
for item in value:
|
|
631
|
+
_nonblank(item, "allowed_options")
|
|
632
|
+
if item not in allowed and item not in protected:
|
|
633
|
+
raise ValueError(f"request option {item!r} is not supported")
|
|
634
|
+
if item in protected:
|
|
635
|
+
raise ValueError(f"request option {item!r} is protected")
|
|
636
|
+
_unique(value, "allowed_options")
|
|
637
|
+
return value
|
|
638
|
+
|
|
639
|
+
|
|
640
|
+
class ErrorFieldMappings(BaseModel):
|
|
641
|
+
"""Immutable provider error field mapping diagnostics."""
|
|
642
|
+
|
|
643
|
+
model_config = ConfigDict(extra="forbid", frozen=True)
|
|
644
|
+
|
|
645
|
+
request_id_paths: tuple[str, ...] = ()
|
|
646
|
+
message_paths: tuple[str, ...] = ()
|
|
647
|
+
code_paths: tuple[str, ...] = ()
|
|
648
|
+
|
|
649
|
+
@field_validator("request_id_paths", "message_paths", "code_paths")
|
|
650
|
+
@classmethod
|
|
651
|
+
def _paths_valid(cls, value: tuple[str, ...]) -> tuple[str, ...]:
|
|
652
|
+
for item in value:
|
|
653
|
+
_nonblank(item, "error field path")
|
|
654
|
+
if any(part in {"", "__class__", "__dict__"} for part in item.split(".")):
|
|
655
|
+
raise ValueError("error field paths must be simple dot paths")
|
|
656
|
+
_unique(value, "error field paths")
|
|
657
|
+
return value
|
|
658
|
+
|
|
659
|
+
|
|
660
|
+
class TransportConfig(BaseModel):
|
|
661
|
+
"""Internal transport safety configuration."""
|
|
662
|
+
|
|
663
|
+
model_config = ConfigDict(extra="forbid", frozen=True)
|
|
664
|
+
|
|
665
|
+
timeout_seconds: float = Field(default=60.0, gt=0)
|
|
666
|
+
success_body_limit_bytes: int = Field(default=4 * 1024 * 1024, gt=0)
|
|
667
|
+
error_body_limit_bytes: int = Field(default=64 * 1024, gt=0)
|
|
668
|
+
follow_redirects: bool = False
|
|
669
|
+
trust_env: bool = False
|
|
670
|
+
tls_verify: bool = True
|
|
671
|
+
|
|
672
|
+
@model_validator(mode="after")
|
|
673
|
+
def _transport_safe_defaults(self) -> TransportConfig:
|
|
674
|
+
if not math.isfinite(self.timeout_seconds):
|
|
675
|
+
raise ValueError("timeout_seconds must be finite")
|
|
676
|
+
if self.follow_redirects:
|
|
677
|
+
raise ValueError("model transport must not follow redirects")
|
|
678
|
+
if self.trust_env:
|
|
679
|
+
raise ValueError("model transport must not trust environment proxies")
|
|
680
|
+
if not self.tls_verify:
|
|
681
|
+
raise ValueError("model transport must verify HTTPS TLS by default")
|
|
682
|
+
return self
|
|
683
|
+
|
|
684
|
+
|
|
685
|
+
class ResolvedModelProfile(BaseModel):
|
|
686
|
+
"""Immutable provider-neutral resolved model profile."""
|
|
687
|
+
|
|
688
|
+
model_config = ConfigDict(extra="forbid", frozen=True)
|
|
689
|
+
|
|
690
|
+
profile_id: str
|
|
691
|
+
provider_id: str
|
|
692
|
+
model_id: str
|
|
693
|
+
transport_id: str = "openai-chat-completions"
|
|
694
|
+
endpoint: EndpointConfig
|
|
695
|
+
authentication: AuthenticationPolicy
|
|
696
|
+
timeout_seconds: float = Field(default=60.0, gt=0)
|
|
697
|
+
maximum_output_tokens: int = Field(default=4096, gt=0)
|
|
698
|
+
configured_headers: HeaderValuePolicy = Field(default_factory=HeaderValuePolicy)
|
|
699
|
+
sampling: SamplingPolicy = Field(default_factory=SamplingPolicy)
|
|
700
|
+
reasoning: ReasoningPolicy = Field(default_factory=ReasoningPolicy)
|
|
701
|
+
capabilities: CapabilityDeclarations = Field(default_factory=CapabilityDeclarations)
|
|
702
|
+
request_options: RequestOptionAllowlist = Field(
|
|
703
|
+
default_factory=RequestOptionAllowlist
|
|
704
|
+
)
|
|
705
|
+
error_mappings: ErrorFieldMappings = Field(default_factory=ErrorFieldMappings)
|
|
706
|
+
transport: TransportConfig = Field(default_factory=TransportConfig)
|
|
707
|
+
source_name: str = "static"
|
|
708
|
+
source_digest: str
|
|
709
|
+
|
|
710
|
+
@model_validator(mode="before")
|
|
711
|
+
@classmethod
|
|
712
|
+
def _accept_legacy_profile_shape(cls, data: Any) -> Any:
|
|
713
|
+
if not isinstance(data, dict):
|
|
714
|
+
return data
|
|
715
|
+
copied = dict(data)
|
|
716
|
+
sampling = copied.get("sampling")
|
|
717
|
+
if "maximum_output_tokens" not in copied and isinstance(sampling, dict):
|
|
718
|
+
legacy_max = sampling.get("default_maximum_output_tokens")
|
|
719
|
+
if legacy_max is not None:
|
|
720
|
+
copied["maximum_output_tokens"] = legacy_max
|
|
721
|
+
return copied
|
|
722
|
+
|
|
723
|
+
@field_validator(
|
|
724
|
+
"profile_id",
|
|
725
|
+
"provider_id",
|
|
726
|
+
"model_id",
|
|
727
|
+
"transport_id",
|
|
728
|
+
"source_name",
|
|
729
|
+
"source_digest",
|
|
730
|
+
)
|
|
731
|
+
@classmethod
|
|
732
|
+
def _strings_nonblank(cls, value: str, info: Any) -> str:
|
|
733
|
+
return _nonblank(value, info.field_name)
|
|
734
|
+
|
|
735
|
+
@field_validator("timeout_seconds")
|
|
736
|
+
@classmethod
|
|
737
|
+
def _timeout_finite(cls, value: float) -> float:
|
|
738
|
+
if not math.isfinite(value):
|
|
739
|
+
raise ValueError("timeout_seconds must be finite")
|
|
740
|
+
return value
|
|
741
|
+
|
|
742
|
+
@property
|
|
743
|
+
def diagnostics(self) -> dict[str, str]:
|
|
744
|
+
"""Return sanitized profile source diagnostics only."""
|
|
745
|
+
return {"source_name": self.source_name, "source_digest": self.source_digest}
|
|
746
|
+
|
|
747
|
+
|
|
748
|
+
AuthenticationConfig = AuthenticationPolicy
|
|
749
|
+
ModelCapabilities = CapabilityDeclarations
|
|
750
|
+
RequestOptions = RequestOptionAllowlist
|
|
751
|
+
|
|
752
|
+
|
|
753
|
+
class TransportRequest(BaseModel):
|
|
754
|
+
"""Private transport request record."""
|
|
755
|
+
|
|
756
|
+
model_config = ConfigDict(extra="forbid", frozen=True)
|
|
757
|
+
|
|
758
|
+
request_id: str
|
|
759
|
+
profile: ResolvedModelProfile
|
|
760
|
+
public_request: ModelCompletionRequest
|
|
761
|
+
url: str
|
|
762
|
+
headers: dict[str, str]
|
|
763
|
+
body: JsonObject
|
|
764
|
+
timeout_seconds: float = Field(gt=0)
|
|
765
|
+
|
|
766
|
+
@field_validator("request_id", "url")
|
|
767
|
+
@classmethod
|
|
768
|
+
def _request_strings_nonblank(cls, value: str, info: Any) -> str:
|
|
769
|
+
return _nonblank(value, info.field_name)
|
|
770
|
+
|
|
771
|
+
@field_validator("timeout_seconds")
|
|
772
|
+
@classmethod
|
|
773
|
+
def _timeout_is_finite(cls, value: float) -> float:
|
|
774
|
+
if not math.isfinite(value):
|
|
775
|
+
raise ValueError("timeout_seconds must be finite")
|
|
776
|
+
return value
|
|
777
|
+
|
|
778
|
+
|
|
779
|
+
class TransportResponse(BaseModel):
|
|
780
|
+
"""Private transport response record after bounded parsing."""
|
|
781
|
+
|
|
782
|
+
model_config = ConfigDict(extra="forbid", frozen=True)
|
|
783
|
+
|
|
784
|
+
status_code: int = Field(ge=100, le=599)
|
|
785
|
+
provider_request_id: str | None = None
|
|
786
|
+
body: JsonObject = Field(default_factory=dict)
|
|
787
|
+
normalized_response: ModelCompletionResponse | None = None
|
|
788
|
+
|
|
789
|
+
|
|
790
|
+
class ResolvedSecret:
|
|
791
|
+
"""Non-Pydantic wrapper for resolved raw secret values."""
|
|
792
|
+
|
|
793
|
+
__slots__ = ("_value",)
|
|
794
|
+
|
|
795
|
+
def __init__(self, value: str) -> None:
|
|
796
|
+
if not value:
|
|
797
|
+
raise SecretResolutionError("resolved secret must not be empty")
|
|
798
|
+
self._value = value
|
|
799
|
+
|
|
800
|
+
def __repr__(self) -> str:
|
|
801
|
+
return "ResolvedSecret(**redacted**)"
|
|
802
|
+
|
|
803
|
+
def __str__(self) -> str:
|
|
804
|
+
return "**redacted**"
|
|
805
|
+
|
|
806
|
+
def reveal_for_header(self) -> str:
|
|
807
|
+
"""Expose the raw value at the authentication header construction point."""
|
|
808
|
+
return self._value
|
|
809
|
+
|
|
810
|
+
|
|
811
|
+
@runtime_checkable
|
|
812
|
+
class SecretResolver(Protocol):
|
|
813
|
+
"""Resolve an admitted secret reference without exposing raw values."""
|
|
814
|
+
|
|
815
|
+
def resolve(self, ref: SecretRef) -> ResolvedSecret:
|
|
816
|
+
"""Resolve ``ref`` to a non-serializable secret wrapper."""
|
|
817
|
+
...
|
|
818
|
+
|
|
819
|
+
|
|
820
|
+
@runtime_checkable
|
|
821
|
+
class ModelProfileResolver(Protocol):
|
|
822
|
+
"""Resolve a logical profile ID into an immutable internal profile."""
|
|
823
|
+
|
|
824
|
+
def resolve(self, profile_id: str) -> ResolvedModelProfile:
|
|
825
|
+
"""Resolve ``profile_id`` exactly and without network I/O."""
|
|
826
|
+
...
|
|
827
|
+
|
|
828
|
+
def diagnostics(self, profile_id: str) -> dict[str, str]:
|
|
829
|
+
"""Return sanitized source diagnostics for a known profile."""
|
|
830
|
+
...
|
|
831
|
+
|
|
832
|
+
|
|
833
|
+
@runtime_checkable
|
|
834
|
+
class ModelTransport(Protocol):
|
|
835
|
+
"""Private model transport boundary."""
|
|
836
|
+
|
|
837
|
+
async def send(self, request: TransportRequest) -> TransportResponse:
|
|
838
|
+
"""Send one already-normalized transport request."""
|
|
839
|
+
...
|
|
840
|
+
|
|
841
|
+
|
|
842
|
+
@runtime_checkable
|
|
843
|
+
class ModelCancellationToken(Protocol):
|
|
844
|
+
"""Minimal cancellation token shape consumed by the default client."""
|
|
845
|
+
|
|
846
|
+
@property
|
|
847
|
+
def cancellation_id(self) -> str:
|
|
848
|
+
"""Return the cancellation identifier."""
|
|
849
|
+
...
|
|
850
|
+
|
|
851
|
+
def is_cancelled(self) -> bool:
|
|
852
|
+
"""Return whether cancellation has been requested."""
|
|
853
|
+
...
|
|
854
|
+
|
|
855
|
+
async def wait(self) -> None:
|
|
856
|
+
"""Wait until cancellation is requested."""
|
|
857
|
+
...
|
|
858
|
+
|
|
859
|
+
@property
|
|
860
|
+
def reason(self) -> str | None:
|
|
861
|
+
"""Return a bounded cancellation reason, when present."""
|
|
862
|
+
...
|
|
863
|
+
|
|
864
|
+
|
|
865
|
+
@runtime_checkable
|
|
866
|
+
class ModelCancellationResolver(Protocol):
|
|
867
|
+
"""Resolve a public cancellation reference for backend checks."""
|
|
868
|
+
|
|
869
|
+
def resolve(self, ref: Any) -> ModelCancellationToken:
|
|
870
|
+
"""Resolve a cancellation reference into a checkable token."""
|
|
871
|
+
...
|
|
872
|
+
|
|
873
|
+
|
|
874
|
+
@runtime_checkable
|
|
875
|
+
class ModelBackendClock(Protocol):
|
|
876
|
+
"""Clock shape needed for deadline checks."""
|
|
877
|
+
|
|
878
|
+
def monotonic(self) -> float:
|
|
879
|
+
"""Return monotonic seconds."""
|
|
880
|
+
...
|
|
881
|
+
|
|
882
|
+
|
|
883
|
+
class _SystemClock:
|
|
884
|
+
def monotonic(self) -> float:
|
|
885
|
+
return time.monotonic()
|
|
886
|
+
|
|
887
|
+
|
|
888
|
+
class StaticModelProfileResolver:
|
|
889
|
+
"""In-process exact profile resolver used by composition roots and tests."""
|
|
890
|
+
|
|
891
|
+
def __init__(self, profiles: Mapping[str, ResolvedModelProfile]) -> None:
|
|
892
|
+
copied = dict(profiles)
|
|
893
|
+
for profile_id, profile in copied.items():
|
|
894
|
+
if profile_id != profile.profile_id:
|
|
895
|
+
raise ModelBackendConfigError(
|
|
896
|
+
"profile mapping key must match profile_id"
|
|
897
|
+
)
|
|
898
|
+
self._profiles = copied
|
|
899
|
+
|
|
900
|
+
def resolve(self, profile_id: str) -> ResolvedModelProfile:
|
|
901
|
+
"""Return an immutable profile or reject before secret/transport work."""
|
|
902
|
+
try:
|
|
903
|
+
return self._profiles[profile_id]
|
|
904
|
+
except KeyError as exc:
|
|
905
|
+
raise ModelBackendConfigError(
|
|
906
|
+
f"unknown model profile {redact_text(profile_id)!r}"
|
|
907
|
+
) from exc
|
|
908
|
+
|
|
909
|
+
def diagnostics(self, profile_id: str) -> dict[str, str]:
|
|
910
|
+
"""Return only sanitized profile source diagnostics."""
|
|
911
|
+
return self.resolve(profile_id).diagnostics
|
|
912
|
+
|
|
913
|
+
|
|
914
|
+
class StaticSecretResolver:
|
|
915
|
+
"""Small deterministic resolver for tests and local composition."""
|
|
916
|
+
|
|
917
|
+
def __init__(self, values: Mapping[str, str]) -> None:
|
|
918
|
+
self._values = dict(values)
|
|
919
|
+
|
|
920
|
+
def resolve(self, ref: SecretRef) -> ResolvedSecret:
|
|
921
|
+
try:
|
|
922
|
+
return ResolvedSecret(self._values[ref.secret_id])
|
|
923
|
+
except KeyError as exc:
|
|
924
|
+
raise SecretResolutionError(f"missing secret {ref.secret_id!r}") from exc
|
|
925
|
+
|
|
926
|
+
|
|
927
|
+
class CapabilityNegotiator:
|
|
928
|
+
"""Evaluate public capability requirements against a resolved profile."""
|
|
929
|
+
|
|
930
|
+
_FIELDS = tuple(ModelCapabilityRequirements.model_fields)
|
|
931
|
+
|
|
932
|
+
def negotiate(
|
|
933
|
+
self,
|
|
934
|
+
required: ModelCapabilityRequirements,
|
|
935
|
+
declarations: CapabilityDeclarations,
|
|
936
|
+
) -> None:
|
|
937
|
+
"""Raise stable provider-neutral failure data for unsupported requirements."""
|
|
938
|
+
failures: dict[str, SanitizedMetadataValue] = {}
|
|
939
|
+
for field_name in self._FIELDS:
|
|
940
|
+
if getattr(required, field_name) is True:
|
|
941
|
+
state = declarations.state_for(field_name)
|
|
942
|
+
if state is not CapabilitySupport.SUPPORTED:
|
|
943
|
+
failures[field_name] = state.value
|
|
944
|
+
if failures:
|
|
945
|
+
raise ModelProviderError(
|
|
946
|
+
category=ProviderErrorCategory.UNSUPPORTED_CAPABILITY,
|
|
947
|
+
message="required model capabilities are unsupported",
|
|
948
|
+
retryable=False,
|
|
949
|
+
fields=failures,
|
|
950
|
+
)
|
|
951
|
+
|
|
952
|
+
|
|
953
|
+
class _CallerOwnedAsyncTransport(httpx.AsyncBaseTransport):
|
|
954
|
+
"""Delegate requests without transferring close ownership to httpx."""
|
|
955
|
+
|
|
956
|
+
def __init__(self, transport: AsyncHttpTransport) -> None:
|
|
957
|
+
self._transport = transport
|
|
958
|
+
|
|
959
|
+
async def handle_async_request(self, request: httpx.Request) -> httpx.Response:
|
|
960
|
+
return await self._transport.handle_async_request(request)
|
|
961
|
+
|
|
962
|
+
async def aclose(self) -> None:
|
|
963
|
+
"""Leave the caller-injected transport open."""
|
|
964
|
+
|
|
965
|
+
|
|
966
|
+
class OpenAIChatCompletionsTransport:
|
|
967
|
+
"""Non-streaming OpenAI-compatible Chat Completions HTTP transport."""
|
|
968
|
+
|
|
969
|
+
def __init__(
|
|
970
|
+
self,
|
|
971
|
+
*,
|
|
972
|
+
http_transport: AsyncHttpTransport | None = None,
|
|
973
|
+
timeouts: OpenAICompatibleTimeouts | None = None,
|
|
974
|
+
timeout_seconds: float | None = None,
|
|
975
|
+
) -> None:
|
|
976
|
+
if timeouts is not None and timeout_seconds is not None:
|
|
977
|
+
raise ModelBackendConfigError(
|
|
978
|
+
"configure either explicit timeouts or timeout_seconds, not both"
|
|
979
|
+
)
|
|
980
|
+
self._timeouts = timeouts or OpenAICompatibleTimeouts.uniform(
|
|
981
|
+
60.0 if timeout_seconds is None else timeout_seconds
|
|
982
|
+
)
|
|
983
|
+
timeout = _phase_timeout(self._timeouts)
|
|
984
|
+
if http_transport is not None and not isinstance(
|
|
985
|
+
http_transport, AsyncHttpTransport
|
|
986
|
+
):
|
|
987
|
+
raise ModelBackendConfigError(
|
|
988
|
+
"http_transport does not implement AsyncHttpTransport"
|
|
989
|
+
)
|
|
990
|
+
client_transport = None
|
|
991
|
+
if http_transport is not None:
|
|
992
|
+
client_transport = _CallerOwnedAsyncTransport(http_transport)
|
|
993
|
+
self._client = httpx.AsyncClient(
|
|
994
|
+
follow_redirects=False,
|
|
995
|
+
trust_env=False,
|
|
996
|
+
verify=True,
|
|
997
|
+
timeout=timeout,
|
|
998
|
+
transport=client_transport,
|
|
999
|
+
)
|
|
1000
|
+
self._client.headers.clear()
|
|
1001
|
+
self._closed = False
|
|
1002
|
+
|
|
1003
|
+
async def __aenter__(self) -> OpenAIChatCompletionsTransport:
|
|
1004
|
+
return self
|
|
1005
|
+
|
|
1006
|
+
async def __aexit__(
|
|
1007
|
+
self,
|
|
1008
|
+
exc_type: type[BaseException] | None,
|
|
1009
|
+
exc: BaseException | None,
|
|
1010
|
+
traceback: object | None,
|
|
1011
|
+
) -> None:
|
|
1012
|
+
await self.aclose()
|
|
1013
|
+
|
|
1014
|
+
async def aclose(self) -> None:
|
|
1015
|
+
"""Close the owned HTTP client at most once."""
|
|
1016
|
+
if self._closed:
|
|
1017
|
+
return
|
|
1018
|
+
self._closed = True
|
|
1019
|
+
await self._client.aclose()
|
|
1020
|
+
|
|
1021
|
+
async def send(self, request: TransportRequest) -> TransportResponse:
|
|
1022
|
+
"""POST one bounded non-streaming Chat Completions request."""
|
|
1023
|
+
if self._closed:
|
|
1024
|
+
raise ModelBackendConfigError("model transport is closed")
|
|
1025
|
+
try:
|
|
1026
|
+
response = await self._client.post(
|
|
1027
|
+
request.url,
|
|
1028
|
+
headers=request.headers,
|
|
1029
|
+
json=request.body,
|
|
1030
|
+
timeout=_phase_timeout(
|
|
1031
|
+
self._timeouts,
|
|
1032
|
+
total_timeout_seconds=request.timeout_seconds,
|
|
1033
|
+
),
|
|
1034
|
+
)
|
|
1035
|
+
except httpx.TimeoutException as exc:
|
|
1036
|
+
raise ModelProviderError(
|
|
1037
|
+
category=ProviderErrorCategory.TIMEOUT,
|
|
1038
|
+
message=f"model transport timed out for {redact_url(request.url)}",
|
|
1039
|
+
) from exc
|
|
1040
|
+
except httpx.NetworkError as exc:
|
|
1041
|
+
raise ModelProviderError(
|
|
1042
|
+
category=ProviderErrorCategory.CONNECTION,
|
|
1043
|
+
message=f"model transport connection failed for {redact_url(request.url)}",
|
|
1044
|
+
) from exc
|
|
1045
|
+
except httpx.HTTPError as exc:
|
|
1046
|
+
raise ModelProviderError(
|
|
1047
|
+
category=ProviderErrorCategory.UNKNOWN,
|
|
1048
|
+
message=f"model transport failed for {redact_url(request.url)}",
|
|
1049
|
+
) from exc
|
|
1050
|
+
|
|
1051
|
+
provider_request_id = _provider_request_id(
|
|
1052
|
+
request.profile.error_mappings, response.headers, None
|
|
1053
|
+
)
|
|
1054
|
+
if response.status_code >= 400:
|
|
1055
|
+
body = await _read_limited_response(
|
|
1056
|
+
response, request.profile.transport.error_body_limit_bytes
|
|
1057
|
+
)
|
|
1058
|
+
parsed = _decode_json_body(body, error_body=True)
|
|
1059
|
+
provider_request_id = _provider_request_id(
|
|
1060
|
+
request.profile.error_mappings, response.headers, parsed
|
|
1061
|
+
)
|
|
1062
|
+
raise _provider_http_error(
|
|
1063
|
+
response.status_code, parsed, provider_request_id
|
|
1064
|
+
)
|
|
1065
|
+
|
|
1066
|
+
raw_content_type = response.headers.get("content-type")
|
|
1067
|
+
content_type = (
|
|
1068
|
+
None
|
|
1069
|
+
if raw_content_type is None
|
|
1070
|
+
else raw_content_type.split(";", 1)[0].strip().lower()
|
|
1071
|
+
)
|
|
1072
|
+
if content_type is None:
|
|
1073
|
+
if not request.profile.endpoint.allow_missing_success_content_type:
|
|
1074
|
+
raise _malformed_response("provider response content type is invalid")
|
|
1075
|
+
elif content_type not in request.profile.endpoint.success_content_types:
|
|
1076
|
+
raise _malformed_response("provider response content type is invalid")
|
|
1077
|
+
body = await _read_limited_response(
|
|
1078
|
+
response, request.profile.transport.success_body_limit_bytes
|
|
1079
|
+
)
|
|
1080
|
+
parsed = _decode_json_body(body, error_body=False)
|
|
1081
|
+
provider_request_id = _provider_request_id(
|
|
1082
|
+
request.profile.error_mappings, response.headers, parsed
|
|
1083
|
+
)
|
|
1084
|
+
transport_response = TransportResponse(
|
|
1085
|
+
status_code=response.status_code,
|
|
1086
|
+
provider_request_id=provider_request_id,
|
|
1087
|
+
body=parsed,
|
|
1088
|
+
)
|
|
1089
|
+
normalized = normalize_transport_response(request.profile, transport_response)
|
|
1090
|
+
return transport_response.model_copy(update={"normalized_response": normalized})
|
|
1091
|
+
|
|
1092
|
+
|
|
1093
|
+
class DefaultModelClient:
|
|
1094
|
+
"""Default provider-neutral model client orchestration.
|
|
1095
|
+
|
|
1096
|
+
The client owns backend policy checks and delegates exactly one already-built
|
|
1097
|
+
request to the configured private transport.
|
|
1098
|
+
"""
|
|
1099
|
+
|
|
1100
|
+
def __init__(
|
|
1101
|
+
self,
|
|
1102
|
+
*,
|
|
1103
|
+
profile_resolver: ModelProfileResolver,
|
|
1104
|
+
secret_resolver: SecretResolver,
|
|
1105
|
+
transport: ModelTransport,
|
|
1106
|
+
cancellation_resolver: ModelCancellationResolver | None = None,
|
|
1107
|
+
clock: ModelBackendClock | None = None,
|
|
1108
|
+
capability_negotiator: CapabilityNegotiator | None = None,
|
|
1109
|
+
local_timeout_seconds: float | None = None,
|
|
1110
|
+
) -> None:
|
|
1111
|
+
self._profile_resolver = profile_resolver
|
|
1112
|
+
self._secret_resolver = secret_resolver
|
|
1113
|
+
self._transport = transport
|
|
1114
|
+
self._cancellation_resolver = cancellation_resolver
|
|
1115
|
+
self._clock = clock or _SystemClock()
|
|
1116
|
+
self._capability_negotiator = capability_negotiator or CapabilityNegotiator()
|
|
1117
|
+
if local_timeout_seconds is not None and (
|
|
1118
|
+
isinstance(local_timeout_seconds, bool)
|
|
1119
|
+
or not isinstance(local_timeout_seconds, (int, float))
|
|
1120
|
+
or local_timeout_seconds <= 0
|
|
1121
|
+
or not math.isfinite(local_timeout_seconds)
|
|
1122
|
+
):
|
|
1123
|
+
raise ModelBackendConfigError(
|
|
1124
|
+
"local_timeout_seconds must be a positive finite number"
|
|
1125
|
+
)
|
|
1126
|
+
self._local_timeout_seconds = local_timeout_seconds
|
|
1127
|
+
self._closed = False
|
|
1128
|
+
|
|
1129
|
+
async def __aenter__(self) -> DefaultModelClient:
|
|
1130
|
+
return self
|
|
1131
|
+
|
|
1132
|
+
async def __aexit__(
|
|
1133
|
+
self,
|
|
1134
|
+
exc_type: type[BaseException] | None,
|
|
1135
|
+
exc: BaseException | None,
|
|
1136
|
+
traceback: object | None,
|
|
1137
|
+
) -> None:
|
|
1138
|
+
await self.aclose()
|
|
1139
|
+
|
|
1140
|
+
async def aclose(self) -> None:
|
|
1141
|
+
"""Close owned transport resources at most once."""
|
|
1142
|
+
if self._closed:
|
|
1143
|
+
return
|
|
1144
|
+
self._closed = True
|
|
1145
|
+
close = getattr(self._transport, "aclose", None)
|
|
1146
|
+
if close is None:
|
|
1147
|
+
return
|
|
1148
|
+
result = close()
|
|
1149
|
+
if inspect.isawaitable(result):
|
|
1150
|
+
await result
|
|
1151
|
+
|
|
1152
|
+
async def complete(
|
|
1153
|
+
self, request: ModelCompletionRequest
|
|
1154
|
+
) -> ModelCompletionResponse:
|
|
1155
|
+
"""Resolve backend policy, call transport once, and normalize the response."""
|
|
1156
|
+
if self._closed:
|
|
1157
|
+
raise ModelBackendConfigError("model client is closed")
|
|
1158
|
+
profile = self._profile_resolver.resolve(request.model_profile_id)
|
|
1159
|
+
if profile.profile_id != request.model_profile_id:
|
|
1160
|
+
raise ModelBackendConfigError("resolved profile identity mismatch")
|
|
1161
|
+
|
|
1162
|
+
self._capability_negotiator.negotiate(
|
|
1163
|
+
request.required_capabilities, profile.capabilities
|
|
1164
|
+
)
|
|
1165
|
+
validate_reasoning_policy(profile.reasoning)
|
|
1166
|
+
sampling = merge_sampling_policy(profile.sampling, request.sampling_overrides)
|
|
1167
|
+
maximum_output_tokens = merge_maximum_output_tokens(
|
|
1168
|
+
profile, request.maximum_output_tokens_override
|
|
1169
|
+
)
|
|
1170
|
+
body = build_transport_body(
|
|
1171
|
+
profile,
|
|
1172
|
+
request,
|
|
1173
|
+
sampling,
|
|
1174
|
+
maximum_output_tokens,
|
|
1175
|
+
)
|
|
1176
|
+
|
|
1177
|
+
token = self._resolve_cancellation(request)
|
|
1178
|
+
self._check_cancelled(token)
|
|
1179
|
+
timeout_seconds, call_deadline = self._effective_timeout(profile, request)
|
|
1180
|
+
self._validate_secret_refs(profile.authentication, request.secret_refs)
|
|
1181
|
+
resolved_secret = resolve_authentication_secret(
|
|
1182
|
+
profile.authentication,
|
|
1183
|
+
request.secret_refs,
|
|
1184
|
+
self._secret_resolver,
|
|
1185
|
+
)
|
|
1186
|
+
self._check_cancelled(token)
|
|
1187
|
+
|
|
1188
|
+
transport_request = TransportRequest(
|
|
1189
|
+
request_id=request.request_id,
|
|
1190
|
+
profile=profile,
|
|
1191
|
+
public_request=request,
|
|
1192
|
+
url=profile.endpoint.chat_completions_url,
|
|
1193
|
+
headers=self._headers(profile, resolved_secret),
|
|
1194
|
+
body=body,
|
|
1195
|
+
timeout_seconds=timeout_seconds,
|
|
1196
|
+
)
|
|
1197
|
+
transport_response = await self._send_with_cancellation(
|
|
1198
|
+
transport_request,
|
|
1199
|
+
token=token,
|
|
1200
|
+
timeout_seconds=timeout_seconds,
|
|
1201
|
+
call_deadline=call_deadline,
|
|
1202
|
+
)
|
|
1203
|
+
return normalize_transport_response(profile, transport_response)
|
|
1204
|
+
|
|
1205
|
+
async def _send_with_cancellation(
|
|
1206
|
+
self,
|
|
1207
|
+
request: TransportRequest,
|
|
1208
|
+
*,
|
|
1209
|
+
token: ModelCancellationToken | None,
|
|
1210
|
+
timeout_seconds: float,
|
|
1211
|
+
call_deadline: float,
|
|
1212
|
+
) -> TransportResponse:
|
|
1213
|
+
transport_task = asyncio.create_task(
|
|
1214
|
+
self._transport.send(request),
|
|
1215
|
+
name=f"millforge-model-transport:{request.request_id}",
|
|
1216
|
+
)
|
|
1217
|
+
cancellation_task = (
|
|
1218
|
+
asyncio.create_task(
|
|
1219
|
+
token.wait(),
|
|
1220
|
+
name=f"millforge-model-cancellation:{request.request_id}",
|
|
1221
|
+
)
|
|
1222
|
+
if token is not None
|
|
1223
|
+
else None
|
|
1224
|
+
)
|
|
1225
|
+
owned_tasks = tuple(
|
|
1226
|
+
task for task in (transport_task, cancellation_task) if task is not None
|
|
1227
|
+
)
|
|
1228
|
+
|
|
1229
|
+
try:
|
|
1230
|
+
done, _ = await asyncio.wait(
|
|
1231
|
+
owned_tasks,
|
|
1232
|
+
timeout=timeout_seconds,
|
|
1233
|
+
return_when=asyncio.FIRST_COMPLETED,
|
|
1234
|
+
)
|
|
1235
|
+
self._check_cancelled(token)
|
|
1236
|
+
if not done or self._clock.monotonic() >= call_deadline:
|
|
1237
|
+
self._raise_deadline_expired()
|
|
1238
|
+
|
|
1239
|
+
if transport_task not in done:
|
|
1240
|
+
# A conforming waiter returns only after cancellation. Keep the
|
|
1241
|
+
# state check authoritative if an implementation wakes early.
|
|
1242
|
+
done, _ = await asyncio.wait(
|
|
1243
|
+
(transport_task,),
|
|
1244
|
+
timeout=max(0.0, call_deadline - self._clock.monotonic()),
|
|
1245
|
+
)
|
|
1246
|
+
self._check_cancelled(token)
|
|
1247
|
+
if (
|
|
1248
|
+
transport_task not in done
|
|
1249
|
+
or self._clock.monotonic() >= call_deadline
|
|
1250
|
+
):
|
|
1251
|
+
self._raise_deadline_expired()
|
|
1252
|
+
|
|
1253
|
+
if cancellation_task is not None:
|
|
1254
|
+
await self._cancel_and_await(cancellation_task)
|
|
1255
|
+
self._check_cancelled(token)
|
|
1256
|
+
if self._clock.monotonic() >= call_deadline:
|
|
1257
|
+
self._raise_deadline_expired()
|
|
1258
|
+
return transport_task.result()
|
|
1259
|
+
except asyncio.CancelledError:
|
|
1260
|
+
await self._cancel_and_await(*owned_tasks)
|
|
1261
|
+
raise
|
|
1262
|
+
finally:
|
|
1263
|
+
await self._cancel_and_await(*owned_tasks)
|
|
1264
|
+
|
|
1265
|
+
@staticmethod
|
|
1266
|
+
async def _cancel_and_await(*tasks: asyncio.Task[Any]) -> None:
|
|
1267
|
+
for task in tasks:
|
|
1268
|
+
if not task.done():
|
|
1269
|
+
task.cancel()
|
|
1270
|
+
if tasks:
|
|
1271
|
+
await asyncio.gather(*tasks, return_exceptions=True)
|
|
1272
|
+
|
|
1273
|
+
def _resolve_cancellation(
|
|
1274
|
+
self, request: ModelCompletionRequest
|
|
1275
|
+
) -> ModelCancellationToken | None:
|
|
1276
|
+
if self._cancellation_resolver is None:
|
|
1277
|
+
return None
|
|
1278
|
+
return self._cancellation_resolver.resolve(request.cancellation)
|
|
1279
|
+
|
|
1280
|
+
def _check_cancelled(self, token: ModelCancellationToken | None) -> None:
|
|
1281
|
+
if token is None or not token.is_cancelled():
|
|
1282
|
+
return
|
|
1283
|
+
reason = redact_text(token.reason or "model request cancelled")
|
|
1284
|
+
raise ModelProviderError(
|
|
1285
|
+
category=ProviderErrorCategory.CANCELLED,
|
|
1286
|
+
message=reason,
|
|
1287
|
+
retryable=False,
|
|
1288
|
+
)
|
|
1289
|
+
|
|
1290
|
+
def _effective_timeout(
|
|
1291
|
+
self,
|
|
1292
|
+
profile: ResolvedModelProfile,
|
|
1293
|
+
request: ModelCompletionRequest,
|
|
1294
|
+
) -> tuple[float, float]:
|
|
1295
|
+
now = self._clock.monotonic()
|
|
1296
|
+
remaining = request.deadline.effective_deadline_monotonic - now
|
|
1297
|
+
if remaining <= 0:
|
|
1298
|
+
self._raise_deadline_expired()
|
|
1299
|
+
admitted_bounds = [profile.timeout_seconds, remaining]
|
|
1300
|
+
if self._local_timeout_seconds is not None:
|
|
1301
|
+
admitted_bounds.append(self._local_timeout_seconds)
|
|
1302
|
+
timeout_seconds = min(admitted_bounds)
|
|
1303
|
+
return timeout_seconds, now + timeout_seconds
|
|
1304
|
+
|
|
1305
|
+
@staticmethod
|
|
1306
|
+
def _raise_deadline_expired() -> None:
|
|
1307
|
+
raise ModelRequestDeadlineExceededError()
|
|
1308
|
+
|
|
1309
|
+
def _headers(
|
|
1310
|
+
self,
|
|
1311
|
+
profile: ResolvedModelProfile,
|
|
1312
|
+
resolved_secret: ResolvedSecret | None,
|
|
1313
|
+
) -> dict[str, str]:
|
|
1314
|
+
return assemble_transport_headers(profile, resolved_secret)
|
|
1315
|
+
|
|
1316
|
+
def _validate_secret_refs(
|
|
1317
|
+
self,
|
|
1318
|
+
authentication: AuthenticationPolicy,
|
|
1319
|
+
admitted: tuple[SecretRef, ...],
|
|
1320
|
+
) -> None:
|
|
1321
|
+
secret_ids: set[str] = set()
|
|
1322
|
+
env_vars: set[str] = set()
|
|
1323
|
+
for secret in admitted:
|
|
1324
|
+
if secret.secret_id in secret_ids:
|
|
1325
|
+
raise SecretResolutionError(
|
|
1326
|
+
"duplicate secret_id in request secret_refs"
|
|
1327
|
+
)
|
|
1328
|
+
if secret.env_var in env_vars:
|
|
1329
|
+
raise SecretResolutionError("duplicate env_var in request secret_refs")
|
|
1330
|
+
secret_ids.add(secret.secret_id)
|
|
1331
|
+
env_vars.add(secret.env_var)
|
|
1332
|
+
|
|
1333
|
+
configured = authentication.secret_ref
|
|
1334
|
+
if configured is None:
|
|
1335
|
+
if admitted:
|
|
1336
|
+
raise SecretResolutionError("unknown secret reference was admitted")
|
|
1337
|
+
return
|
|
1338
|
+
if len(admitted) != 1:
|
|
1339
|
+
raise SecretResolutionError("exactly one authentication secret is required")
|
|
1340
|
+
if admitted[0] != configured:
|
|
1341
|
+
raise SecretResolutionError("admitted secret does not match configuration")
|
|
1342
|
+
|
|
1343
|
+
|
|
1344
|
+
def merge_sampling_policy(
|
|
1345
|
+
policy: SamplingPolicy,
|
|
1346
|
+
overrides: SamplingRequest,
|
|
1347
|
+
) -> SamplingRequest:
|
|
1348
|
+
"""Merge sampling defaults with explicitly allowlisted caller overrides."""
|
|
1349
|
+
merged = {
|
|
1350
|
+
"temperature": policy.temperature,
|
|
1351
|
+
"top_p": policy.top_p,
|
|
1352
|
+
"presence_penalty": policy.presence_penalty,
|
|
1353
|
+
"frequency_penalty": policy.frequency_penalty,
|
|
1354
|
+
"seed": policy.seed,
|
|
1355
|
+
"stop": policy.stop,
|
|
1356
|
+
}
|
|
1357
|
+
override_values = overrides.model_dump(exclude_none=True)
|
|
1358
|
+
for name, value in override_values.items():
|
|
1359
|
+
if name not in _DIRECT_SAMPLING_BODY_FIELDS:
|
|
1360
|
+
raise ModelBackendConfigError(f"sampling override {name!r} is not allowed")
|
|
1361
|
+
if name not in policy.allowed_overrides:
|
|
1362
|
+
raise ModelBackendConfigError(f"sampling override {name!r} is not allowed")
|
|
1363
|
+
merged[name] = value
|
|
1364
|
+
return SamplingRequest.model_validate(merged)
|
|
1365
|
+
|
|
1366
|
+
|
|
1367
|
+
def merge_maximum_output_tokens(
|
|
1368
|
+
profile: ResolvedModelProfile,
|
|
1369
|
+
override: int | None,
|
|
1370
|
+
) -> int:
|
|
1371
|
+
"""Return the effective output-token cap under profile override policy."""
|
|
1372
|
+
if override is None:
|
|
1373
|
+
return profile.maximum_output_tokens
|
|
1374
|
+
if not profile.sampling.allow_maximum_output_tokens_override:
|
|
1375
|
+
raise ModelBackendConfigError("maximum output token override is not allowed")
|
|
1376
|
+
return override
|
|
1377
|
+
|
|
1378
|
+
|
|
1379
|
+
def build_transport_body(
|
|
1380
|
+
profile: ResolvedModelProfile,
|
|
1381
|
+
request: ModelCompletionRequest,
|
|
1382
|
+
sampling: SamplingRequest,
|
|
1383
|
+
maximum_output_tokens: int | None,
|
|
1384
|
+
) -> JsonObject:
|
|
1385
|
+
"""Build the provider-neutral request body consumed by private transports."""
|
|
1386
|
+
_validate_reasoning_replay_history(profile, request.messages)
|
|
1387
|
+
body: JsonObject = {
|
|
1388
|
+
"model": profile.model_id,
|
|
1389
|
+
"messages": [
|
|
1390
|
+
_message_payload(
|
|
1391
|
+
message,
|
|
1392
|
+
replay_field=profile.reasoning.tool_call_replay_field,
|
|
1393
|
+
)
|
|
1394
|
+
for message in request.messages
|
|
1395
|
+
],
|
|
1396
|
+
"stream": False,
|
|
1397
|
+
}
|
|
1398
|
+
if request.tools:
|
|
1399
|
+
body["tools"] = [
|
|
1400
|
+
{
|
|
1401
|
+
"type": "function",
|
|
1402
|
+
"function": {
|
|
1403
|
+
"name": tool.name,
|
|
1404
|
+
"description": tool.description,
|
|
1405
|
+
"parameters": tool.input_schema,
|
|
1406
|
+
},
|
|
1407
|
+
}
|
|
1408
|
+
for tool in request.tools
|
|
1409
|
+
]
|
|
1410
|
+
for name, value in sampling.model_dump(
|
|
1411
|
+
include=_DIRECT_SAMPLING_BODY_FIELDS,
|
|
1412
|
+
exclude_none=True,
|
|
1413
|
+
).items():
|
|
1414
|
+
body[name] = value
|
|
1415
|
+
if maximum_output_tokens is not None:
|
|
1416
|
+
body["max_tokens"] = maximum_output_tokens
|
|
1417
|
+
_add_reasoning_controls(profile, sampling, body)
|
|
1418
|
+
_add_request_options(profile, request, body)
|
|
1419
|
+
return body
|
|
1420
|
+
|
|
1421
|
+
|
|
1422
|
+
def validate_reasoning_policy(policy: ReasoningPolicy) -> None:
|
|
1423
|
+
"""Reject un-mappable required reasoning before secret resolution and HTTP."""
|
|
1424
|
+
if policy.mode is ReasoningMode.REQUIRED and (
|
|
1425
|
+
policy.mode_field is None or policy.mode not in policy.mode_values
|
|
1426
|
+
):
|
|
1427
|
+
raise ModelBackendConfigError("required reasoning has no faithful mapping")
|
|
1428
|
+
if policy.tool_call_replay_field is None:
|
|
1429
|
+
return
|
|
1430
|
+
if policy.mode not in (ReasoningMode.ENABLED, ReasoningMode.REQUIRED):
|
|
1431
|
+
raise ModelBackendConfigError(
|
|
1432
|
+
"reasoning replay requires enabled or required reasoning"
|
|
1433
|
+
)
|
|
1434
|
+
if (
|
|
1435
|
+
policy.mode_field is None
|
|
1436
|
+
or not policy.mode_field.strip()
|
|
1437
|
+
or policy.mode not in policy.mode_values
|
|
1438
|
+
):
|
|
1439
|
+
raise ModelBackendConfigError("reasoning replay has no faithful mapping")
|
|
1440
|
+
|
|
1441
|
+
|
|
1442
|
+
def _add_reasoning_controls(
|
|
1443
|
+
profile: ResolvedModelProfile,
|
|
1444
|
+
sampling: SamplingRequest,
|
|
1445
|
+
body: JsonObject,
|
|
1446
|
+
) -> None:
|
|
1447
|
+
del sampling
|
|
1448
|
+
if (
|
|
1449
|
+
profile.reasoning.mode is not ReasoningMode.DISABLED
|
|
1450
|
+
and profile.reasoning.mode_field is not None
|
|
1451
|
+
):
|
|
1452
|
+
body[profile.reasoning.mode_field] = profile.reasoning.mode_values[
|
|
1453
|
+
profile.reasoning.mode
|
|
1454
|
+
]
|
|
1455
|
+
if profile.reasoning.effort is not None:
|
|
1456
|
+
if profile.reasoning.effort_field is None:
|
|
1457
|
+
raise ModelBackendConfigError("reasoning effort has no configured mapping")
|
|
1458
|
+
body[profile.reasoning.effort_field] = profile.reasoning.effort_values[
|
|
1459
|
+
profile.reasoning.effort
|
|
1460
|
+
]
|
|
1461
|
+
|
|
1462
|
+
|
|
1463
|
+
def _add_request_options(
|
|
1464
|
+
profile: ResolvedModelProfile,
|
|
1465
|
+
request: ModelCompletionRequest,
|
|
1466
|
+
body: JsonObject,
|
|
1467
|
+
) -> None:
|
|
1468
|
+
allowed = set(profile.request_options.allowed_options)
|
|
1469
|
+
for name, value in request.request_options.items():
|
|
1470
|
+
if name not in allowed:
|
|
1471
|
+
raise ModelBackendConfigError(f"request option {name!r} is not allowed")
|
|
1472
|
+
body[name] = value
|
|
1473
|
+
|
|
1474
|
+
|
|
1475
|
+
def assemble_transport_headers(
|
|
1476
|
+
profile: ResolvedModelProfile,
|
|
1477
|
+
resolved_secret: ResolvedSecret | None,
|
|
1478
|
+
) -> dict[str, str]:
|
|
1479
|
+
"""Assemble headers in stable base/configured/auth order."""
|
|
1480
|
+
headers: dict[str, str] = {}
|
|
1481
|
+
_append_header(headers, "Content-Type", "application/json")
|
|
1482
|
+
_append_header(headers, "User-Agent", "millforge-model-backend/1")
|
|
1483
|
+
for name, value in profile.configured_headers.values.items():
|
|
1484
|
+
_append_header(headers, name, value)
|
|
1485
|
+
for name, value in build_auth_headers(
|
|
1486
|
+
profile.authentication, resolved_secret
|
|
1487
|
+
).items():
|
|
1488
|
+
_append_header(headers, name, value)
|
|
1489
|
+
return headers
|
|
1490
|
+
|
|
1491
|
+
|
|
1492
|
+
def _append_header(headers: dict[str, str], name: str, value: str) -> None:
|
|
1493
|
+
normalized = name.lower()
|
|
1494
|
+
if any(existing.lower() == normalized for existing in headers):
|
|
1495
|
+
raise ModelBackendConfigError(
|
|
1496
|
+
f"duplicate transport header {redact_text(name)!r}"
|
|
1497
|
+
)
|
|
1498
|
+
headers[name] = value
|
|
1499
|
+
|
|
1500
|
+
|
|
1501
|
+
def _message_payload(
|
|
1502
|
+
message: ModelMessage,
|
|
1503
|
+
*,
|
|
1504
|
+
replay_field: Literal["reasoning_content"] | None,
|
|
1505
|
+
) -> JsonObject:
|
|
1506
|
+
if message.role == "assistant":
|
|
1507
|
+
payload: JsonObject = {"role": "assistant"}
|
|
1508
|
+
if message.content is not None:
|
|
1509
|
+
payload["content"] = message.content
|
|
1510
|
+
if message.tool_calls:
|
|
1511
|
+
payload["tool_calls"] = [
|
|
1512
|
+
{
|
|
1513
|
+
"id": call.call_id,
|
|
1514
|
+
"type": "function",
|
|
1515
|
+
"function": {
|
|
1516
|
+
"name": call.name,
|
|
1517
|
+
"arguments": _tool_arguments_json(call),
|
|
1518
|
+
},
|
|
1519
|
+
}
|
|
1520
|
+
for call in message.tool_calls
|
|
1521
|
+
]
|
|
1522
|
+
if replay_field is not None and message.reasoning_content is not None:
|
|
1523
|
+
payload[replay_field] = message.reasoning_content
|
|
1524
|
+
return payload
|
|
1525
|
+
if message.role == "tool":
|
|
1526
|
+
return {
|
|
1527
|
+
"role": "tool",
|
|
1528
|
+
"tool_call_id": message.tool_call_id,
|
|
1529
|
+
"name": message.tool_name,
|
|
1530
|
+
"content": message.content,
|
|
1531
|
+
}
|
|
1532
|
+
return {"role": message.role, "content": message.content}
|
|
1533
|
+
|
|
1534
|
+
|
|
1535
|
+
def _validate_reasoning_replay_history(
|
|
1536
|
+
profile: ResolvedModelProfile,
|
|
1537
|
+
messages: tuple[ModelMessage, ...],
|
|
1538
|
+
) -> None:
|
|
1539
|
+
replay_field = profile.reasoning.tool_call_replay_field
|
|
1540
|
+
if replay_field is None:
|
|
1541
|
+
if any(
|
|
1542
|
+
isinstance(message, AssistantMessage)
|
|
1543
|
+
and message.reasoning_content is not None
|
|
1544
|
+
for message in messages
|
|
1545
|
+
):
|
|
1546
|
+
raise ModelBackendConfigError(
|
|
1547
|
+
"reasoning continuation requires a selected replay field"
|
|
1548
|
+
)
|
|
1549
|
+
return
|
|
1550
|
+
|
|
1551
|
+
pending: tuple[ModelToolCall, ...] = ()
|
|
1552
|
+
pending_index = 0
|
|
1553
|
+
seen_call_ids: set[str] = set()
|
|
1554
|
+
for message in messages:
|
|
1555
|
+
if pending:
|
|
1556
|
+
if not isinstance(message, ToolResultMessage):
|
|
1557
|
+
raise ModelBackendConfigError(
|
|
1558
|
+
"reasoning replay tool results are incomplete"
|
|
1559
|
+
)
|
|
1560
|
+
expected = pending[pending_index]
|
|
1561
|
+
if (
|
|
1562
|
+
message.tool_call_id != expected.call_id
|
|
1563
|
+
or message.tool_name != expected.name
|
|
1564
|
+
):
|
|
1565
|
+
raise ModelBackendConfigError(
|
|
1566
|
+
"reasoning replay tool results are not in call order"
|
|
1567
|
+
)
|
|
1568
|
+
pending_index += 1
|
|
1569
|
+
if pending_index == len(pending):
|
|
1570
|
+
pending = ()
|
|
1571
|
+
pending_index = 0
|
|
1572
|
+
continue
|
|
1573
|
+
|
|
1574
|
+
if isinstance(message, ToolResultMessage):
|
|
1575
|
+
raise ModelBackendConfigError("reasoning replay tool result is orphaned")
|
|
1576
|
+
if not isinstance(message, AssistantMessage):
|
|
1577
|
+
continue
|
|
1578
|
+
if message.reasoning_content is not None:
|
|
1579
|
+
_validate_outbound_reasoning_content(profile, message.reasoning_content)
|
|
1580
|
+
if not message.tool_calls:
|
|
1581
|
+
continue
|
|
1582
|
+
if message.reasoning_content is None:
|
|
1583
|
+
raise ModelBackendConfigError(
|
|
1584
|
+
"reasoning replay tool call is missing continuation"
|
|
1585
|
+
)
|
|
1586
|
+
for call in message.tool_calls:
|
|
1587
|
+
if call.call_id in seen_call_ids:
|
|
1588
|
+
raise ModelBackendConfigError(
|
|
1589
|
+
"reasoning replay tool-call IDs are duplicated"
|
|
1590
|
+
)
|
|
1591
|
+
seen_call_ids.add(call.call_id)
|
|
1592
|
+
pending = message.tool_calls
|
|
1593
|
+
|
|
1594
|
+
if pending:
|
|
1595
|
+
raise ModelBackendConfigError("reasoning replay tool results are incomplete")
|
|
1596
|
+
|
|
1597
|
+
|
|
1598
|
+
def _validate_outbound_reasoning_content(
|
|
1599
|
+
profile: ResolvedModelProfile,
|
|
1600
|
+
value: object,
|
|
1601
|
+
) -> None:
|
|
1602
|
+
if not isinstance(value, str) or not value.strip():
|
|
1603
|
+
raise ModelBackendConfigError("reasoning replay continuation is malformed")
|
|
1604
|
+
if len(value.encode("utf-8")) > profile.transport.success_body_limit_bytes:
|
|
1605
|
+
raise ModelBackendConfigError(
|
|
1606
|
+
"reasoning replay continuation exceeded configured limit"
|
|
1607
|
+
)
|
|
1608
|
+
|
|
1609
|
+
|
|
1610
|
+
def _tool_arguments_json(call: ModelToolCall) -> str:
|
|
1611
|
+
if isinstance(call.arguments, ParsedToolArguments):
|
|
1612
|
+
_reject_non_finite_json(call.arguments.value, path="tool arguments")
|
|
1613
|
+
return json.dumps(call.arguments.value, sort_keys=True, separators=(",", ":"))
|
|
1614
|
+
return str(call.arguments.raw)
|
|
1615
|
+
|
|
1616
|
+
|
|
1617
|
+
def _reject_non_finite_json(value: JsonValue, *, path: str) -> None:
|
|
1618
|
+
if isinstance(value, float) and not math.isfinite(value):
|
|
1619
|
+
raise ModelBackendConfigError(f"{path} contains non-finite numeric value")
|
|
1620
|
+
if isinstance(value, dict):
|
|
1621
|
+
for key, item in value.items():
|
|
1622
|
+
_reject_non_finite_json(item, path=f"{path}.{key}")
|
|
1623
|
+
elif isinstance(value, (list, tuple)):
|
|
1624
|
+
for index, item in enumerate(value):
|
|
1625
|
+
_reject_non_finite_json(item, path=f"{path}[{index}]")
|
|
1626
|
+
|
|
1627
|
+
|
|
1628
|
+
def normalize_transport_response(
|
|
1629
|
+
profile: ResolvedModelProfile,
|
|
1630
|
+
response: TransportResponse,
|
|
1631
|
+
) -> ModelCompletionResponse:
|
|
1632
|
+
"""Return only owned response data from a transport response."""
|
|
1633
|
+
if response.normalized_response is not None:
|
|
1634
|
+
source = response.normalized_response
|
|
1635
|
+
reasoning_content = _admit_response_reasoning_content(
|
|
1636
|
+
profile,
|
|
1637
|
+
getattr(source.message, "reasoning_content", None),
|
|
1638
|
+
has_tool_calls=bool(source.message.tool_calls),
|
|
1639
|
+
)
|
|
1640
|
+
source = source.model_copy(
|
|
1641
|
+
update={
|
|
1642
|
+
"message": source.message.model_copy(
|
|
1643
|
+
update={"reasoning_content": reasoning_content}
|
|
1644
|
+
)
|
|
1645
|
+
}
|
|
1646
|
+
)
|
|
1647
|
+
normalized = ModelCompletionResponse.model_validate(source.model_dump())
|
|
1648
|
+
if normalized.model_id != profile.model_id:
|
|
1649
|
+
raise ModelProviderError(
|
|
1650
|
+
category=ProviderErrorCategory.MALFORMED_RESPONSE,
|
|
1651
|
+
message="provider response model did not match resolved profile",
|
|
1652
|
+
retryable=False,
|
|
1653
|
+
)
|
|
1654
|
+
if normalized.provider_request_id is None and response.provider_request_id:
|
|
1655
|
+
normalized = normalized.model_copy(
|
|
1656
|
+
update={"provider_request_id": response.provider_request_id}
|
|
1657
|
+
)
|
|
1658
|
+
return normalized
|
|
1659
|
+
return _normalize_openai_chat_body(profile, response)
|
|
1660
|
+
|
|
1661
|
+
|
|
1662
|
+
def _normalize_openai_chat_body(
|
|
1663
|
+
profile: ResolvedModelProfile,
|
|
1664
|
+
response: TransportResponse,
|
|
1665
|
+
) -> ModelCompletionResponse:
|
|
1666
|
+
body = response.body
|
|
1667
|
+
choices = body.get("choices")
|
|
1668
|
+
if not isinstance(choices, list) or len(choices) != 1:
|
|
1669
|
+
raise _malformed_response("provider response requires exactly one choice")
|
|
1670
|
+
choice = choices[0]
|
|
1671
|
+
if not isinstance(choice, dict):
|
|
1672
|
+
raise _malformed_response("provider response choice is malformed")
|
|
1673
|
+
message = choice.get("message")
|
|
1674
|
+
if not isinstance(message, dict):
|
|
1675
|
+
raise _malformed_response("provider response message is malformed")
|
|
1676
|
+
if message.get("role") != "assistant":
|
|
1677
|
+
raise _malformed_response("provider response message role is malformed")
|
|
1678
|
+
|
|
1679
|
+
model_id = body.get("model", profile.model_id)
|
|
1680
|
+
if model_id != profile.model_id:
|
|
1681
|
+
raise _malformed_response("provider response model did not match profile")
|
|
1682
|
+
|
|
1683
|
+
content = message.get("content")
|
|
1684
|
+
if content is not None and not isinstance(content, str):
|
|
1685
|
+
raise _malformed_response("provider response content is malformed")
|
|
1686
|
+
tool_calls = _normalize_tool_calls(message.get("tool_calls", ()))
|
|
1687
|
+
if content == "" and tool_calls:
|
|
1688
|
+
content = None
|
|
1689
|
+
assistant = AssistantMessage(
|
|
1690
|
+
content=content,
|
|
1691
|
+
tool_calls=tool_calls,
|
|
1692
|
+
reasoning_content=_admit_response_reasoning_content(
|
|
1693
|
+
profile,
|
|
1694
|
+
message.get("reasoning_content"),
|
|
1695
|
+
has_tool_calls=bool(tool_calls),
|
|
1696
|
+
),
|
|
1697
|
+
)
|
|
1698
|
+
finish_reason = _FINISH_REASON_MAP.get(str(choice.get("finish_reason")), "unknown")
|
|
1699
|
+
usage = _normalize_usage(body.get("usage"))
|
|
1700
|
+
return ModelCompletionResponse(
|
|
1701
|
+
provider_request_id=response.provider_request_id,
|
|
1702
|
+
model_id=profile.model_id,
|
|
1703
|
+
message=assistant,
|
|
1704
|
+
finish_reason=cast(
|
|
1705
|
+
Any,
|
|
1706
|
+
finish_reason,
|
|
1707
|
+
),
|
|
1708
|
+
usage=usage,
|
|
1709
|
+
)
|
|
1710
|
+
|
|
1711
|
+
|
|
1712
|
+
def _admit_response_reasoning_content(
|
|
1713
|
+
profile: ResolvedModelProfile,
|
|
1714
|
+
value: object,
|
|
1715
|
+
*,
|
|
1716
|
+
has_tool_calls: bool,
|
|
1717
|
+
) -> str | None:
|
|
1718
|
+
if profile.reasoning.tool_call_replay_field is None or not has_tool_calls:
|
|
1719
|
+
return None
|
|
1720
|
+
if not isinstance(value, str) or not value.strip():
|
|
1721
|
+
raise _malformed_response("provider reasoning continuation is malformed")
|
|
1722
|
+
if len(value.encode("utf-8")) > profile.transport.success_body_limit_bytes:
|
|
1723
|
+
raise _malformed_response(
|
|
1724
|
+
"provider reasoning continuation exceeded configured limit"
|
|
1725
|
+
)
|
|
1726
|
+
return value
|
|
1727
|
+
|
|
1728
|
+
|
|
1729
|
+
def _normalize_tool_calls(value: object) -> tuple[ModelToolCall, ...]:
|
|
1730
|
+
if value in (None, ()):
|
|
1731
|
+
return ()
|
|
1732
|
+
if not isinstance(value, list):
|
|
1733
|
+
raise _malformed_response("provider response tool_calls is malformed")
|
|
1734
|
+
calls: list[ModelToolCall] = []
|
|
1735
|
+
seen_call_ids: set[str] = set()
|
|
1736
|
+
for item in value:
|
|
1737
|
+
if not isinstance(item, dict):
|
|
1738
|
+
raise _malformed_response("provider response tool call is malformed")
|
|
1739
|
+
function = item.get("function")
|
|
1740
|
+
if not isinstance(function, dict):
|
|
1741
|
+
raise _malformed_response("provider response tool function is malformed")
|
|
1742
|
+
if item.get("type") not in (None, "function"):
|
|
1743
|
+
raise _malformed_response("provider response tool call type is malformed")
|
|
1744
|
+
call_id = item.get("id")
|
|
1745
|
+
name = function.get("name")
|
|
1746
|
+
if not isinstance(call_id, str) or not isinstance(name, str):
|
|
1747
|
+
raise _malformed_response(
|
|
1748
|
+
"provider response tool call identity is malformed"
|
|
1749
|
+
)
|
|
1750
|
+
if call_id in seen_call_ids:
|
|
1751
|
+
raise _malformed_response("provider response tool call IDs are duplicated")
|
|
1752
|
+
seen_call_ids.add(call_id)
|
|
1753
|
+
raw_arguments = function.get("arguments", "")
|
|
1754
|
+
calls.append(
|
|
1755
|
+
ModelToolCall(
|
|
1756
|
+
call_id=call_id,
|
|
1757
|
+
name=name,
|
|
1758
|
+
arguments=_normalize_tool_arguments(raw_arguments),
|
|
1759
|
+
)
|
|
1760
|
+
)
|
|
1761
|
+
return tuple(calls)
|
|
1762
|
+
|
|
1763
|
+
|
|
1764
|
+
def _normalize_tool_arguments(
|
|
1765
|
+
raw: JsonValue,
|
|
1766
|
+
) -> ParsedToolArguments | InvalidToolArguments:
|
|
1767
|
+
if isinstance(raw, dict):
|
|
1768
|
+
try:
|
|
1769
|
+
_reject_non_finite_json(raw, path="tool arguments")
|
|
1770
|
+
except ModelBackendConfigError:
|
|
1771
|
+
return InvalidToolArguments(raw=raw, error_code="non_finite_json")
|
|
1772
|
+
return ParsedToolArguments(value=raw)
|
|
1773
|
+
if not isinstance(raw, str):
|
|
1774
|
+
return InvalidToolArguments(raw=raw, error_code="not_json_object")
|
|
1775
|
+
try:
|
|
1776
|
+
parsed = json.loads(
|
|
1777
|
+
raw,
|
|
1778
|
+
object_pairs_hook=_reject_duplicate_tool_argument_keys,
|
|
1779
|
+
parse_constant=_reject_tool_argument_constant,
|
|
1780
|
+
)
|
|
1781
|
+
except json.JSONDecodeError:
|
|
1782
|
+
return InvalidToolArguments(raw=raw, error_code="malformed_json")
|
|
1783
|
+
except ValueError:
|
|
1784
|
+
return InvalidToolArguments(raw=raw, error_code="ambiguous_json")
|
|
1785
|
+
if not isinstance(parsed, dict):
|
|
1786
|
+
return InvalidToolArguments(raw=raw, error_code="not_json_object")
|
|
1787
|
+
try:
|
|
1788
|
+
_reject_non_finite_json(parsed, path="tool arguments")
|
|
1789
|
+
except ModelBackendConfigError:
|
|
1790
|
+
return InvalidToolArguments(raw=raw, error_code="non_finite_json")
|
|
1791
|
+
return ParsedToolArguments(value=parsed)
|
|
1792
|
+
|
|
1793
|
+
|
|
1794
|
+
def _reject_duplicate_tool_argument_keys(
|
|
1795
|
+
pairs: list[tuple[str, JsonValue]],
|
|
1796
|
+
) -> JsonObject:
|
|
1797
|
+
parsed: JsonObject = {}
|
|
1798
|
+
for key, value in pairs:
|
|
1799
|
+
if key in parsed:
|
|
1800
|
+
raise ValueError(f"duplicate tool argument key {key!r}")
|
|
1801
|
+
parsed[key] = value
|
|
1802
|
+
return parsed
|
|
1803
|
+
|
|
1804
|
+
|
|
1805
|
+
def _reject_tool_argument_constant(value: str) -> None:
|
|
1806
|
+
raise ValueError(f"non-finite tool argument number {value}")
|
|
1807
|
+
|
|
1808
|
+
|
|
1809
|
+
def _normalize_usage(value: object) -> TokenUsage | None:
|
|
1810
|
+
if value is None:
|
|
1811
|
+
return None
|
|
1812
|
+
if not isinstance(value, dict):
|
|
1813
|
+
raise _malformed_response("provider usage is malformed")
|
|
1814
|
+
input_tokens = value.get("prompt_tokens")
|
|
1815
|
+
output_tokens = value.get("completion_tokens")
|
|
1816
|
+
total_tokens = value.get("total_tokens")
|
|
1817
|
+
if not isinstance(input_tokens, int):
|
|
1818
|
+
raise _malformed_response("provider usage token counts are malformed")
|
|
1819
|
+
if not isinstance(output_tokens, int):
|
|
1820
|
+
raise _malformed_response("provider usage token counts are malformed")
|
|
1821
|
+
if not isinstance(total_tokens, int):
|
|
1822
|
+
raise _malformed_response("provider usage token counts are malformed")
|
|
1823
|
+
try:
|
|
1824
|
+
return TokenUsage(
|
|
1825
|
+
input_tokens=input_tokens,
|
|
1826
|
+
output_tokens=output_tokens,
|
|
1827
|
+
total_tokens=total_tokens,
|
|
1828
|
+
provider_reported=True,
|
|
1829
|
+
)
|
|
1830
|
+
except ValueError as exc:
|
|
1831
|
+
raise _malformed_response(
|
|
1832
|
+
"provider usage token counts are inconsistent"
|
|
1833
|
+
) from exc
|
|
1834
|
+
|
|
1835
|
+
|
|
1836
|
+
def _malformed_response(message: str) -> ModelProviderError:
|
|
1837
|
+
return ModelProviderError(
|
|
1838
|
+
category=ProviderErrorCategory.MALFORMED_RESPONSE,
|
|
1839
|
+
message=message,
|
|
1840
|
+
retryable=False,
|
|
1841
|
+
)
|
|
1842
|
+
|
|
1843
|
+
|
|
1844
|
+
async def _read_limited_response(
|
|
1845
|
+
response: httpx.Response,
|
|
1846
|
+
limit_bytes: int,
|
|
1847
|
+
) -> bytes:
|
|
1848
|
+
body = await response.aread()
|
|
1849
|
+
if len(body) > limit_bytes:
|
|
1850
|
+
raise ModelProviderError(
|
|
1851
|
+
category=ProviderErrorCategory.MALFORMED_RESPONSE,
|
|
1852
|
+
message="provider response body exceeded configured limit",
|
|
1853
|
+
retryable=False,
|
|
1854
|
+
)
|
|
1855
|
+
return body
|
|
1856
|
+
|
|
1857
|
+
|
|
1858
|
+
def _decode_json_body(body: bytes, *, error_body: bool) -> JsonObject:
|
|
1859
|
+
if not body:
|
|
1860
|
+
return {}
|
|
1861
|
+
try:
|
|
1862
|
+
parsed = json.loads(body, object_pairs_hook=_reject_duplicate_json_keys)
|
|
1863
|
+
except json.JSONDecodeError as exc:
|
|
1864
|
+
if error_body:
|
|
1865
|
+
return {}
|
|
1866
|
+
raise ModelProviderError(
|
|
1867
|
+
category=ProviderErrorCategory.MALFORMED_RESPONSE,
|
|
1868
|
+
message="provider response body was not valid JSON",
|
|
1869
|
+
retryable=False,
|
|
1870
|
+
) from exc
|
|
1871
|
+
except ValueError as exc:
|
|
1872
|
+
raise ModelProviderError(
|
|
1873
|
+
category=ProviderErrorCategory.MALFORMED_RESPONSE,
|
|
1874
|
+
message="provider response JSON contained duplicate keys",
|
|
1875
|
+
retryable=False,
|
|
1876
|
+
) from exc
|
|
1877
|
+
if not isinstance(parsed, dict):
|
|
1878
|
+
raise _malformed_response("provider response JSON must be an object")
|
|
1879
|
+
return cast(JsonObject, parsed)
|
|
1880
|
+
|
|
1881
|
+
|
|
1882
|
+
def _reject_duplicate_json_keys(pairs: list[tuple[str, JsonValue]]) -> JsonObject:
|
|
1883
|
+
result: JsonObject = {}
|
|
1884
|
+
for key, value in pairs:
|
|
1885
|
+
if key in result:
|
|
1886
|
+
raise ValueError(f"duplicate JSON key {key!r}")
|
|
1887
|
+
result[key] = value
|
|
1888
|
+
return result
|
|
1889
|
+
|
|
1890
|
+
|
|
1891
|
+
def _provider_http_error(
|
|
1892
|
+
status_code: int,
|
|
1893
|
+
body: JsonObject,
|
|
1894
|
+
provider_request_id: str | None,
|
|
1895
|
+
) -> ModelProviderError:
|
|
1896
|
+
category = _category_for_status(status_code, body)
|
|
1897
|
+
message = _error_message(body) or f"provider returned HTTP {status_code}"
|
|
1898
|
+
fields: dict[str, SanitizedMetadataValue] = {"status_code": status_code}
|
|
1899
|
+
code = _dot_path(body, "error.code")
|
|
1900
|
+
if isinstance(code, str):
|
|
1901
|
+
fields["provider_code"] = code
|
|
1902
|
+
raise ModelProviderError(
|
|
1903
|
+
category=category,
|
|
1904
|
+
message=message,
|
|
1905
|
+
provider_request_id=provider_request_id,
|
|
1906
|
+
fields=fields,
|
|
1907
|
+
)
|
|
1908
|
+
|
|
1909
|
+
|
|
1910
|
+
def _phase_timeout(
|
|
1911
|
+
timeouts: OpenAICompatibleTimeouts | float,
|
|
1912
|
+
*,
|
|
1913
|
+
total_timeout_seconds: float | None = None,
|
|
1914
|
+
) -> httpx.Timeout:
|
|
1915
|
+
if isinstance(timeouts, (int, float)):
|
|
1916
|
+
timeouts = OpenAICompatibleTimeouts.uniform(timeouts)
|
|
1917
|
+
if total_timeout_seconds is not None and (
|
|
1918
|
+
total_timeout_seconds <= 0 or not math.isfinite(total_timeout_seconds)
|
|
1919
|
+
):
|
|
1920
|
+
raise ModelBackendConfigError("transport timeout must be positive and finite")
|
|
1921
|
+
|
|
1922
|
+
def bounded(value: float) -> float:
|
|
1923
|
+
return (
|
|
1924
|
+
value
|
|
1925
|
+
if total_timeout_seconds is None
|
|
1926
|
+
else min(value, total_timeout_seconds)
|
|
1927
|
+
)
|
|
1928
|
+
|
|
1929
|
+
return httpx.Timeout(
|
|
1930
|
+
connect=bounded(timeouts.connect_seconds),
|
|
1931
|
+
read=bounded(timeouts.read_seconds),
|
|
1932
|
+
write=bounded(timeouts.write_seconds),
|
|
1933
|
+
pool=bounded(timeouts.pool_seconds),
|
|
1934
|
+
)
|
|
1935
|
+
|
|
1936
|
+
|
|
1937
|
+
def _category_for_status(
|
|
1938
|
+
status_code: int,
|
|
1939
|
+
body: JsonObject,
|
|
1940
|
+
) -> ProviderErrorCategory:
|
|
1941
|
+
code = _dot_path(body, "error.code")
|
|
1942
|
+
text = str(code).lower() if code is not None else ""
|
|
1943
|
+
if status_code == 401:
|
|
1944
|
+
return ProviderErrorCategory.AUTHENTICATION
|
|
1945
|
+
if status_code == 403:
|
|
1946
|
+
return ProviderErrorCategory.AUTHORIZATION
|
|
1947
|
+
if status_code == 408:
|
|
1948
|
+
return ProviderErrorCategory.TIMEOUT
|
|
1949
|
+
if status_code == 429:
|
|
1950
|
+
return ProviderErrorCategory.RATE_LIMIT
|
|
1951
|
+
if status_code in {400, 409, 422}:
|
|
1952
|
+
return ProviderErrorCategory.INVALID_REQUEST
|
|
1953
|
+
if status_code == 501 or "unsupported" in text:
|
|
1954
|
+
return ProviderErrorCategory.UNSUPPORTED_CAPABILITY
|
|
1955
|
+
if status_code >= 500:
|
|
1956
|
+
return ProviderErrorCategory.SERVER_ERROR
|
|
1957
|
+
return ProviderErrorCategory.UNKNOWN
|
|
1958
|
+
|
|
1959
|
+
|
|
1960
|
+
def _error_message(body: JsonObject) -> str | None:
|
|
1961
|
+
message = _dot_path(body, "error.message")
|
|
1962
|
+
if isinstance(message, str) and message.strip():
|
|
1963
|
+
return message
|
|
1964
|
+
return None
|
|
1965
|
+
|
|
1966
|
+
|
|
1967
|
+
def _provider_request_id(
|
|
1968
|
+
mappings: ErrorFieldMappings,
|
|
1969
|
+
headers: httpx.Headers,
|
|
1970
|
+
body: JsonObject | None,
|
|
1971
|
+
) -> str | None:
|
|
1972
|
+
for header_name in ("x-request-id", "request-id", "openai-request-id"):
|
|
1973
|
+
value = headers.get(header_name)
|
|
1974
|
+
if value:
|
|
1975
|
+
return _bounded(redact_text(value), length=128)
|
|
1976
|
+
if body is None:
|
|
1977
|
+
return None
|
|
1978
|
+
for path in mappings.request_id_paths:
|
|
1979
|
+
value = _dot_path(body, path)
|
|
1980
|
+
if isinstance(value, str) and value.strip():
|
|
1981
|
+
return _bounded(redact_text(value), length=128)
|
|
1982
|
+
return None
|
|
1983
|
+
|
|
1984
|
+
|
|
1985
|
+
def _dot_path(body: Mapping[str, JsonValue], path: str) -> JsonValue:
|
|
1986
|
+
current: JsonValue = body
|
|
1987
|
+
for part in path.split("."):
|
|
1988
|
+
if not isinstance(current, dict) or part not in current:
|
|
1989
|
+
return None
|
|
1990
|
+
current = current[part]
|
|
1991
|
+
return current
|
|
1992
|
+
|
|
1993
|
+
|
|
1994
|
+
def assert_secret_admitted(
|
|
1995
|
+
configured: SecretRef, admitted: tuple[SecretRef, ...]
|
|
1996
|
+
) -> None:
|
|
1997
|
+
"""Ensure a configured secret reference is present in the public request."""
|
|
1998
|
+
if configured not in admitted:
|
|
1999
|
+
raise SecretResolutionError(f"secret {configured.secret_id!r} was not admitted")
|
|
2000
|
+
|
|
2001
|
+
|
|
2002
|
+
def resolve_authentication_secret(
|
|
2003
|
+
policy: AuthenticationPolicy,
|
|
2004
|
+
admitted: tuple[SecretRef, ...],
|
|
2005
|
+
resolver: SecretResolver,
|
|
2006
|
+
) -> ResolvedSecret | None:
|
|
2007
|
+
"""Resolve a configured auth secret only after admission succeeds."""
|
|
2008
|
+
if policy.secret_ref is None:
|
|
2009
|
+
return None
|
|
2010
|
+
assert_secret_admitted(policy.secret_ref, admitted)
|
|
2011
|
+
try:
|
|
2012
|
+
resolved = resolver.resolve(policy.secret_ref)
|
|
2013
|
+
except Exception:
|
|
2014
|
+
raise SecretResolutionError("authentication secret resolution failed") from None
|
|
2015
|
+
if not isinstance(resolved, ResolvedSecret):
|
|
2016
|
+
raise SecretResolutionError(
|
|
2017
|
+
"secret resolver returned an unsupported resolved secret"
|
|
2018
|
+
)
|
|
2019
|
+
return resolved
|
|
2020
|
+
|
|
2021
|
+
|
|
2022
|
+
def build_auth_headers(
|
|
2023
|
+
policy: AuthenticationPolicy,
|
|
2024
|
+
resolved_secret: ResolvedSecret | None,
|
|
2025
|
+
) -> dict[str, str]:
|
|
2026
|
+
"""Build only authentication headers from an already-resolved secret."""
|
|
2027
|
+
if policy.scheme is AuthenticationScheme.NONE:
|
|
2028
|
+
return {}
|
|
2029
|
+
if resolved_secret is None:
|
|
2030
|
+
raise SecretResolutionError("authentication secret was not resolved")
|
|
2031
|
+
raw = resolved_secret.reveal_for_header()
|
|
2032
|
+
if policy.scheme is AuthenticationScheme.BEARER:
|
|
2033
|
+
return {"Authorization": f"Bearer {raw}"}
|
|
2034
|
+
if policy.header_name is None:
|
|
2035
|
+
raise SecretResolutionError("custom authentication header is missing")
|
|
2036
|
+
return {policy.header_name: raw}
|
|
2037
|
+
|
|
2038
|
+
|
|
2039
|
+
def redact_url(value: str) -> str:
|
|
2040
|
+
"""Redact URL userinfo and query values for diagnostics."""
|
|
2041
|
+
split = urlsplit(value)
|
|
2042
|
+
if not split.scheme or not split.netloc:
|
|
2043
|
+
return redact_diagnostic_text(value)
|
|
2044
|
+
host = split.hostname or ""
|
|
2045
|
+
if split.port:
|
|
2046
|
+
host = f"{host}:{split.port}"
|
|
2047
|
+
query = "&".join(
|
|
2048
|
+
f"{key}=**redacted**"
|
|
2049
|
+
for key, _ in parse_qsl(split.query, keep_blank_values=True)
|
|
2050
|
+
)
|
|
2051
|
+
return urlunsplit(
|
|
2052
|
+
(
|
|
2053
|
+
split.scheme,
|
|
2054
|
+
host,
|
|
2055
|
+
split.path,
|
|
2056
|
+
query,
|
|
2057
|
+
DEFAULT_REDACTION_POLICY.replacement if split.fragment else "",
|
|
2058
|
+
)
|
|
2059
|
+
)
|
|
2060
|
+
|
|
2061
|
+
|
|
2062
|
+
def redact_text(value: object, *, secret_values: tuple[str, ...] = ()) -> str:
|
|
2063
|
+
"""Redact auth headers, explicit secrets, URL details, and token/key patterns."""
|
|
2064
|
+
if isinstance(value, str):
|
|
2065
|
+
return redact_diagnostic_text(value, secret_values=secret_values)
|
|
2066
|
+
redacted = redact_diagnostic_value(value, secret_values=secret_values)
|
|
2067
|
+
return (
|
|
2068
|
+
redacted if isinstance(redacted, str) else redact_diagnostic_text(str(redacted))
|
|
2069
|
+
)
|
|
2070
|
+
|
|
2071
|
+
|
|
2072
|
+
def sanitize_provider_error_fields(
|
|
2073
|
+
fields: Mapping[str, SanitizedMetadataValue],
|
|
2074
|
+
*,
|
|
2075
|
+
secret_values: tuple[str, ...] = (),
|
|
2076
|
+
) -> dict[str, SanitizedMetadataValue]:
|
|
2077
|
+
"""Bound and redact provider error fields for safe persistence."""
|
|
2078
|
+
bounded_fields = {
|
|
2079
|
+
key: value
|
|
2080
|
+
for index, (key, value) in enumerate(fields.items())
|
|
2081
|
+
if index < _MAX_SANITIZED_FIELDS
|
|
2082
|
+
}
|
|
2083
|
+
return cast(
|
|
2084
|
+
dict[str, SanitizedMetadataValue],
|
|
2085
|
+
redact_diagnostic_mapping(bounded_fields, secret_values=secret_values),
|
|
2086
|
+
)
|
|
2087
|
+
|
|
2088
|
+
|
|
2089
|
+
def redact_mapping(
|
|
2090
|
+
values: Mapping[str, object],
|
|
2091
|
+
*,
|
|
2092
|
+
secret_values: tuple[str, ...] = (),
|
|
2093
|
+
) -> dict[str, SanitizedMetadataValue]:
|
|
2094
|
+
"""Redact debug summaries, events, traces, metrics, manifests, and diagnostics."""
|
|
2095
|
+
return cast(
|
|
2096
|
+
dict[str, SanitizedMetadataValue],
|
|
2097
|
+
redact_diagnostic_mapping(values, secret_values=secret_values),
|
|
2098
|
+
)
|