Restore qualified task identities and finish transport protocol types

Assistant: codex
Assistant-Model: gpt-5.6-luna
Assistant-Session: 01a07ff8-19d0-7820-b4d0-1353833cb7fc
This commit is contained in:
tegwick 2026-09-09 21:45:45 +02:00
parent 4055986c98
commit 68ffe9ef76
9 changed files with 899 additions and 157 deletions

View file

@ -15,7 +15,7 @@ import threading
import time
from dataclasses import asdict, dataclass
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
from typing import Protocol
from typing import Any, NoReturn, Protocol, TypeGuard, cast
from urllib.parse import urlsplit
@ -23,11 +23,11 @@ class RequestRefused(RuntimeError):
"""Bounded refusal; never carries request content or credentials."""
def _integer(value: object, upper: int) -> bool:
def _integer(value: object, upper: int) -> TypeGuard[int]:
return type(value) is int and 0 < value <= upper
def _unique(pairs):
def _unique(pairs: list[tuple[str, Any]]) -> dict[str, Any]:
result = {}
for key, value in pairs:
if key in result:
@ -36,8 +36,8 @@ def _unique(pairs):
return result
def _json(raw: bytes):
def invalid(_):
def _json(raw: bytes) -> Any:
def invalid(_: str) -> NoReturn:
raise ValueError("nonfinite number")
return json.loads(raw, object_pairs_hook=_unique, parse_constant=invalid)
@ -62,13 +62,13 @@ class MessagesPolicy:
max_body_bytes: int = 2_000_000
timeout_seconds: int = 120
def __post_init__(self):
def __post_init__(self) -> None:
for value in (self.tariff_ref, self.model):
if not isinstance(value, str) or not re.fullmatch(
r"[A-Za-z0-9][A-Za-z0-9._:/@+-]{0,199}", value
):
raise RequestRefused("invalid request policy identity")
for value, upper in (
for bound, upper in (
(self.context_tokens, 10_000_000),
(self.max_output_tokens, 1_000_000),
(self.input_microusd_per_token, 1_000_000),
@ -76,7 +76,7 @@ class MessagesPolicy:
(self.max_body_bytes, 2_000_000),
(self.timeout_seconds, 900),
):
if not _integer(value, upper):
if not _integer(bound, upper):
raise RequestRefused("invalid request policy bound")
if not isinstance(self.allowed_betas, tuple) or len(set(self.allowed_betas)) != len(
self.allowed_betas
@ -128,7 +128,7 @@ class MessagesPolicy:
}:
raise RequestRefused("only keep-all thinking context admitted")
def cache(value):
def cache(value: Any) -> None:
if (
not isinstance(value, dict)
or set(value) - {"type", "ttl"}
@ -137,7 +137,7 @@ class MessagesPolicy:
):
raise RequestRefused("cache mode not admitted")
def blocks(value, *, system=False, nested=False):
def blocks(value: Any, *, system: bool = False, nested: bool = False) -> None:
if isinstance(value, str):
return
if not isinstance(value, list) or len(value) > 10000:
@ -259,12 +259,12 @@ class RequestMeter(Protocol):
class _Stream:
"""Observe terminal usage without persisting provider content."""
def __init__(self, policy: MessagesPolicy):
def __init__(self, policy: MessagesPolicy) -> None:
self.policy = policy
self.started = self.stopped = self.delta = False
self.usage = {}
self.usage: dict[str, Any] = {}
def event(self, raw: bytes):
def event(self, raw: bytes) -> None:
data = b"\n".join(
line[5:].lstrip() for line in raw.splitlines() if line.startswith(b"data:")
)
@ -304,7 +304,7 @@ class _Stream:
):
raise RequestRefused("provider content feature not admitted")
def _usage(self, value):
def _usage(self, value: Any) -> None:
if not isinstance(value, dict):
raise RequestRefused("provider usage incomplete")
for key, count in value.items():
@ -323,7 +323,7 @@ class _Stream:
def cost(self) -> int:
if not self.stopped:
raise RequestRefused("provider stream incomplete")
counts = {}
counts: dict[str, int] = {}
for key in (
"input_tokens",
"output_tokens",
@ -341,11 +341,15 @@ class _Stream:
)
class _OwnerHTTPServer(ThreadingHTTPServer):
owner: MessagesServer
class _Handler(BaseHTTPRequestHandler):
def log_message(self, *args):
def log_message(self, format: str, *args: Any) -> None:
pass
def _error(self, status, code):
def _error(self, status: int, code: str) -> None:
raw = json.dumps(
{"type": "error", "error": {"type": "invalid_request_error", "message": code}}
).encode()
@ -355,8 +359,8 @@ class _Handler(BaseHTTPRequestHandler):
self.end_headers()
self.wfile.write(raw)
def do_POST(self):
owner = self.server.owner
def do_POST(self) -> None:
owner = cast(_OwnerHTTPServer, self.server).owner
upstream = None
sent = False
try:
@ -402,7 +406,9 @@ class _Handler(BaseHTTPRequestHandler):
else http.client.HTTPConnection
)
upstream = kind(
owner.endpoint.hostname, owner.endpoint.port, timeout=owner.policy.timeout_seconds
cast(str, owner.endpoint.hostname),
owner.endpoint.port,
timeout=owner.policy.timeout_seconds,
)
headers = {
"Content-Type": "application/json",
@ -490,7 +496,7 @@ class MessagesServer:
host: str = "127.0.0.1",
port: int = 0,
allow_test_http: bool = False,
):
) -> None:
endpoint = urlsplit(upstream_url)
if endpoint.scheme != "https" and not (
allow_test_http
@ -515,19 +521,19 @@ class MessagesServer:
raise RequestRefused("explicit provider credential required")
self.policy, self.meter, self.endpoint = policy, meter, endpoint
self._provider_key = provider_key
self._httpd = ThreadingHTTPServer((host, port), _Handler)
self._httpd = _OwnerHTTPServer((host, port), _Handler)
self._httpd.owner = self
self._thread = None
self._thread: threading.Thread | None = None
@property
def port(self):
return self._httpd.server_address[1]
def port(self) -> int:
return int(self._httpd.server_address[1])
def start(self):
def start(self) -> None:
self._thread = threading.Thread(target=self._httpd.serve_forever, daemon=True)
self._thread.start()
def stop(self):
def stop(self) -> None:
if self._thread is not None:
self._httpd.shutdown()
self._thread.join()