llm-connect/llm_connect/messages_gate.py
tegwick 718e6730e4
All checks were successful
CI Smoke / host-smoke (push) Successful in 0s
CI Smoke / container-smoke (push) Successful in 1s
Provide a private Unix listener for owner-metered Messages
Assistant: codex
Assistant-Model: gpt-6-astra
Assistant-Session: 01a07ff8-19d0-7820-b4d0-1353833cb7fc
2026-09-09 22:20:55 +02:00

603 lines
24 KiB
Python

"""Opt-in Anthropic Messages transport with owner-supplied request admission.
This listener has no /execute route, environment credential discovery, proxy,
redirect or automatic retry. Hosting, accepted tariffs, custody and egress remain
the caller's responsibility. No production listener is installed by this module.
"""
from __future__ import annotations
import hashlib
import http.client
import json
import os
import re
import socket
import stat
import threading
import time
from dataclasses import asdict, dataclass
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
from pathlib import Path
from socketserver import ThreadingUnixStreamServer
from typing import Any, NoReturn, Protocol, TypeGuard, cast
from urllib.parse import urlsplit
class RequestRefused(RuntimeError):
"""Bounded refusal; never carries request content or credentials."""
def _integer(value: object, upper: int) -> TypeGuard[int]:
return type(value) is int and 0 < value <= upper
def _unique(pairs: list[tuple[str, Any]]) -> dict[str, Any]:
result = {}
for key, value in pairs:
if key in result:
raise ValueError("duplicate field")
result[key] = value
return result
def _json(raw: bytes) -> Any:
def invalid(_: str) -> NoReturn:
raise ValueError("nonfinite number")
return json.loads(raw, object_pairs_hook=_unique, parse_constant=invalid)
@dataclass(frozen=True)
class MessagesPolicy:
"""Owner-accepted upper bounds, not a built-in price list or token estimate.
input_microusd_per_token must cover the highest admitted input/cache rate,
including residency/tier multipliers. Reserve the entire admitted context.
One microusd = USD 0.000001; rounding is the accepting owner's responsibility.
"""
tariff_ref: str
model: str
context_tokens: int
max_output_tokens: int
input_microusd_per_token: int
output_microusd_per_token: int
allowed_betas: tuple[str, ...] = ()
max_body_bytes: int = 2_000_000
timeout_seconds: int = 120
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 bound, upper in (
(self.context_tokens, 10_000_000),
(self.max_output_tokens, 1_000_000),
(self.input_microusd_per_token, 1_000_000),
(self.output_microusd_per_token, 1_000_000),
(self.max_body_bytes, 2_000_000),
(self.timeout_seconds, 900),
):
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
):
raise RequestRefused("invalid beta policy")
if any(
not isinstance(b, str) or not re.fullmatch(r"[a-z0-9-]{1,100}", b)
for b in self.allowed_betas
):
raise RequestRefused("invalid beta policy")
@property
def sha256(self) -> str:
return hashlib.sha256(
json.dumps(asdict(self), sort_keys=True, separators=(",", ":")).encode()
).hexdigest()
def validate(self, data: object, betas: str) -> int:
"""Validate a narrow text/custom-tool Messages subset, then bound cost."""
allowed = {
"model",
"messages",
"max_tokens",
"stream",
"system",
"tools",
"tool_choice",
"metadata",
"thinking",
"output_config",
"temperature",
"top_p",
"top_k",
"stop_sequences",
"cache_control",
"context_management",
}
if not isinstance(data, dict) or set(data) - allowed:
raise RequestRefused("unsupported request fields")
if data.get("model") != self.model or data.get("stream") is not True:
raise RequestRefused("request model or stream not admitted")
output = data.get("max_tokens")
if not _integer(output, self.max_output_tokens):
raise RequestRefused("output limit not admitted")
if set(filter(None, (b.strip() for b in betas.split(",")))) - set(self.allowed_betas):
raise RequestRefused("beta feature not admitted")
if "context_management" in data and data["context_management"] != {
"edits": [{"type": "clear_thinking_20251015", "keep": "all"}]
}:
raise RequestRefused("only keep-all thinking context admitted")
def cache(value: Any) -> None:
if (
not isinstance(value, dict)
or set(value) - {"type", "ttl"}
or value.get("type") != "ephemeral"
or value.get("ttl", "5m") not in ("5m", "1h")
):
raise RequestRefused("cache mode not admitted")
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:
raise RequestRefused("content not admitted")
for block in value:
if not isinstance(block, dict):
raise RequestRefused("content block not admitted")
kind = block.get("type")
fields = {
"text": {"type", "text", "cache_control"},
"tool_use": {"type", "id", "name", "input", "cache_control"},
"tool_result": {"type", "tool_use_id", "content", "is_error", "cache_control"},
"thinking": {"type", "thinking", "signature"},
"redacted_thinking": {"type", "data"},
}
if (
kind not in fields
or set(block) - fields[kind]
or ((system or nested) and kind != "text")
):
raise RequestRefused("content feature not admitted")
if kind == "text" and not isinstance(block.get("text"), str):
raise RequestRefused("invalid text block")
if kind == "tool_use" and (
not isinstance(block.get("input"), dict)
or not isinstance(block.get("name"), str)
or not isinstance(block.get("id"), str)
):
raise RequestRefused("invalid custom tool use")
if kind == "tool_result":
if (
not isinstance(block.get("tool_use_id"), str)
or type(block.get("is_error", False)) is not bool
):
raise RequestRefused("invalid custom tool result")
blocks(block.get("content", ""), nested=True)
if "cache_control" in block:
cache(block["cache_control"])
messages = data.get("messages")
if not isinstance(messages, list) or not 1 <= len(messages) <= 10000:
raise RequestRefused("messages not admitted")
for message in messages:
if (
not isinstance(message, dict)
or set(message) != {"role", "content"}
or message["role"] not in ("user", "assistant")
):
raise RequestRefused("message not admitted")
blocks(message["content"])
if "system" in data:
blocks(data["system"], system=True)
tools = data.get("tools", [])
if not isinstance(tools, list) or len(tools) > 100:
raise RequestRefused("tools not admitted")
for tool in tools:
if (
not isinstance(tool, dict)
or set(tool) - {"name", "description", "input_schema", "cache_control"}
or not isinstance(tool.get("name"), str)
or not isinstance(tool.get("input_schema"), dict)
):
raise RequestRefused("only custom client tools admitted")
if "cache_control" in tool:
cache(tool["cache_control"])
if "cache_control" in data:
cache(data["cache_control"])
if "thinking" in data:
value = data["thinking"]
if (
not isinstance(value, dict)
or set(value) - {"type", "budget_tokens", "display"}
or value.get("type") not in ("disabled", "adaptive", "enabled")
):
raise RequestRefused("thinking mode not admitted")
if value.get("type") == "enabled" and not _integer(value.get("budget_tokens"), output):
raise RequestRefused("thinking budget not admitted")
if "display" in value and value["display"] not in ("summarized", "omitted"):
raise RequestRefused("thinking display not admitted")
if "output_config" in data:
value = data["output_config"]
if (
not isinstance(value, dict)
or set(value) != {"effort"}
or value["effort"] not in ("low", "medium", "high", "max")
):
raise RequestRefused("output feature not admitted")
if "tool_choice" in data:
value = data["tool_choice"]
if (
not isinstance(value, dict)
or set(value) - {"type", "name", "disable_parallel_tool_use"}
or value.get("type") not in ("auto", "any", "tool", "none")
):
raise RequestRefused("tool choice not admitted")
if "metadata" in data:
value = data["metadata"]
if (
not isinstance(value, dict)
or set(value) - {"user_id"}
or not isinstance(value.get("user_id", ""), str)
or len(value.get("user_id", "")) > 1000
):
raise RequestRefused("metadata not admitted")
return (
self.context_tokens * self.input_microusd_per_token
+ output * self.output_microusd_per_token
)
class RequestMeter(Protocol):
"""Implemented by the trusted caller; methods must commit before returning."""
def reserve_request(self, token: str, policy_sha256: str, liability_microusd: int) -> str: ...
def request_active(self, receipt: str) -> bool: ...
def complete_request(self, receipt: str, observed_microusd: int) -> None: ...
class _Stream:
"""Observe terminal usage without persisting provider content."""
def __init__(self, policy: MessagesPolicy) -> None:
self.policy = policy
self.started = self.stopped = self.delta = False
self.usage: dict[str, Any] = {}
def event(self, raw: bytes) -> None:
data = b"\n".join(
line[5:].lstrip() for line in raw.splitlines() if line.startswith(b"data:")
)
if not data:
return
value = _json(data)
kind = value.get("type")
if kind == "ping":
return
if self.stopped or kind == "error":
raise RequestRefused("provider stream incomplete")
if kind == "message_start":
message = value.get("message", {})
if self.started or message.get("model") != self.policy.model:
raise RequestRefused("provider model mismatch")
self.started = True
self._usage(message.get("usage", {}))
elif kind == "message_delta":
if not self.started:
raise RequestRefused("provider stream order invalid")
self._usage(value.get("usage", {}))
self.delta = self.delta or bool(value.get("delta", {}).get("stop_reason"))
elif kind == "message_stop":
if not self.started or not self.delta:
raise RequestRefused("provider stream incomplete")
self.stopped = True
elif (
kind not in ("content_block_start", "content_block_delta", "content_block_stop")
or not self.started
):
raise RequestRefused("provider stream feature not admitted")
elif kind == "content_block_start" and value.get("content_block", {}).get("type") not in (
"text",
"thinking",
"redacted_thinking",
"tool_use",
):
raise RequestRefused("provider content feature not admitted")
def _usage(self, value: Any) -> None:
if not isinstance(value, dict):
raise RequestRefused("provider usage incomplete")
for key, count in value.items():
if key in (
"input_tokens",
"output_tokens",
"cache_creation_input_tokens",
"cache_read_input_tokens",
):
if type(count) is not int or count < self.usage.get(key, 0):
raise RequestRefused("provider usage regressed")
if key == "server_tool_use" and count and any(count.values()):
raise RequestRefused("provider fee feature not admitted")
self.usage.update(value)
def cost(self) -> int:
if not self.stopped:
raise RequestRefused("provider stream incomplete")
counts: dict[str, int] = {}
for key in (
"input_tokens",
"output_tokens",
"cache_creation_input_tokens",
"cache_read_input_tokens",
):
value = self.usage.get(key, 0 if key.startswith("cache_") else None)
if type(value) is not int or not 0 <= value <= 100_000_000:
raise RequestRefused("provider usage incomplete")
counts[key] = value
inputs = sum(counts[k] for k in counts if k != "output_tokens")
return (
inputs * self.policy.input_microusd_per_token
+ counts["output_tokens"] * self.policy.output_microusd_per_token
)
class _OwnerHTTPServer(ThreadingHTTPServer):
owner: MessagesServer
class _OwnerUnixServer(ThreadingUnixStreamServer):
owner: MessagesServer
daemon_threads = True
def server_bind(self) -> None:
self._slots = threading.BoundedSemaphore(16)
super().server_bind()
def process_request(self, request: socket.socket | tuple[bytes, socket.socket], client_address: Any) -> None:
if not isinstance(request, socket.socket):
raise TypeError("Unix stream socket required")
if not self._slots.acquire(blocking=False):
request.close()
return
request.settimeout(10)
try:
super().process_request(request, client_address)
except BaseException:
self._slots.release()
raise
def process_request_thread(self, request: socket.socket | tuple[bytes, socket.socket], client_address: Any) -> None:
try:
super().process_request_thread(request, client_address)
finally:
self._slots.release()
class _Handler(BaseHTTPRequestHandler):
def log_message(self, format: str, *args: Any) -> None:
pass
def _error(self, status: int, code: str) -> None:
raw = json.dumps(
{"type": "error", "error": {"type": "invalid_request_error", "message": code}}
).encode()
self.send_response(status)
self.send_header("Content-Type", "application/json")
self.send_header("Content-Length", str(len(raw)))
self.end_headers()
self.wfile.write(raw)
def do_POST(self) -> None:
owner = cast(_OwnerHTTPServer | _OwnerUnixServer, self.server).owner
upstream = None
sent = False
try:
if self.path not in ("/v1/messages", "/v1/messages?beta=true"):
self._error(404, "route_not_admitted")
return
for header in (
"Content-Length",
"x-api-key",
"anthropic-beta",
"Authorization",
"Content-Type",
):
if len(self.headers.get_all(header, [])) > 1:
raise RequestRefused("ambiguous headers")
if (
self.headers.get("Transfer-Encoding")
or self.headers.get("Content-Encoding")
or self.headers.get("Authorization")
):
raise RequestRefused("unsupported request encoding or authentication")
if self.headers.get_content_type() != "application/json":
raise RequestRefused("JSON required")
length = int(self.headers.get("Content-Length", "0"))
if not 0 < length <= owner.policy.max_body_bytes:
raise RequestRefused("request body limit")
self.connection.settimeout(10)
raw = self.rfile.read(length)
if len(raw) != length:
raise RequestRefused("incomplete request")
data = _json(raw)
betas = self.headers.get("anthropic-beta", "")
liability = owner.policy.validate(data, betas)
receipt = owner.meter.reserve_request(
self.headers.get("x-api-key", ""), owner.policy.sha256, liability
)
if not owner.meter.request_active(receipt):
raise RequestRefused("request lease lost")
deadline = time.monotonic() + owner.policy.timeout_seconds
kind = (
http.client.HTTPSConnection
if owner.endpoint.scheme == "https"
else http.client.HTTPConnection
)
upstream = kind(
cast(str, owner.endpoint.hostname),
owner.endpoint.port,
timeout=owner.policy.timeout_seconds,
)
headers = {
"Content-Type": "application/json",
"Accept": "text/event-stream",
"Accept-Encoding": "identity",
"anthropic-version": "2023-06-01",
"x-api-key": owner._provider_key,
}
if betas:
headers["anthropic-beta"] = betas
upstream.request(
"POST",
"/v1/messages",
body=json.dumps(data, allow_nan=False).encode(),
headers=headers,
)
response = upstream.getresponse()
if (
response.status != 200
or response.headers.get_content_type() != "text/event-stream"
or response.getheader("Content-Encoding")
):
raise RequestRefused("provider outcome uncertain")
self.send_response(200)
self.send_header("Content-Type", "text/event-stream")
self.send_header("Connection", "close")
self.end_headers()
sent = True
self.close_connection = True
stream = _Stream(owner.policy)
event = bytearray()
total = 0
terminal_bytes = bytearray()
while True:
if not owner.meter.request_active(receipt) or time.monotonic() >= deadline:
raise RequestRefused("request lease or deadline lost")
# read1 avoids waiting for a full buffer and preserves streaming.
if upstream.sock is not None:
upstream.sock.settimeout(max(0.001, deadline - time.monotonic()))
chunk = response.read1(16384)
if not chunk:
break
total += len(chunk)
event.extend(chunk)
if total > 32_000_000 or len(event) > 1_000_000:
raise RequestRefused("provider stream size exceeded")
# SSE accepts CRLF too. Preserve original bytes on the client wire.
normalized = bytes(event).replace(b"\r\n", b"\n")
while b"\n\n" in normalized:
frame, normalized = normalized.split(b"\n\n", 1)
stream.event(frame)
event = bytearray(normalized)
if stream.stopped:
terminal_bytes.extend(chunk)
else:
self.wfile.write(chunk)
self.wfile.flush()
if bytes(event).strip():
raise RequestRefused("provider stream truncated")
owner.meter.complete_request(receipt, stream.cost())
# Do not expose message_stop until durable completion; the CLI may
# immediately issue its next tool-loop request after this event.
self.wfile.write(terminal_bytes)
self.wfile.flush()
except Exception:
# Any post-reservation failure leaves a durable unknown hold.
# Neither upstream errors nor caller exceptions reach logs or clients.
if not sent:
self._error(400, "request_not_admitted_or_outcome_uncertain")
finally:
if upstream is not None:
upstream.close()
class MessagesServer:
"""Explicit owner-hosted listener; default loopback, never enabled by serve mode."""
def __init__(
self,
policy: MessagesPolicy,
meter: RequestMeter,
*,
provider_key: str,
upstream_url: str = "https://api.anthropic.com",
host: str = "127.0.0.1",
port: int = 0,
allow_test_http: bool = False,
unix_path: Path | None = None,
) -> None:
endpoint = urlsplit(upstream_url)
if endpoint.scheme != "https" and not (
allow_test_http
and endpoint.scheme == "http"
and endpoint.hostname in ("127.0.0.1", "::1")
):
raise RequestRefused("provider transport requires HTTPS")
if (
not endpoint.hostname
or endpoint.username
or endpoint.password
or endpoint.path not in ("", "/")
or endpoint.query
or endpoint.fragment
):
raise RequestRefused("fixed provider origin required")
if (
not isinstance(provider_key, str)
or not provider_key
or any(ord(c) < 33 or ord(c) > 126 for c in provider_key)
):
raise RequestRefused("explicit provider credential required")
self.policy, self.meter, self.endpoint = policy, meter, endpoint
self._provider_key = provider_key
self._unix_path = unix_path
self._httpd: _OwnerHTTPServer | _OwnerUnixServer
if unix_path is not None:
if host != "127.0.0.1" or port != 0:
raise RequestRefused("Unix route cannot also select a TCP listener")
parent = unix_path.parent
metadata = parent.lstat()
if (
not unix_path.is_absolute()
or parent.resolve() != parent
or not stat.S_ISDIR(metadata.st_mode)
or metadata.st_uid != os.getuid()
or stat.S_IMODE(metadata.st_mode) != 0o700
):
raise RequestRefused("Unix route requires a private owner directory")
# bind refuses existing paths, including dangling symlinks. Never unlink
# someone else's socket to recover a crashed or duplicated owner.
self._httpd = _OwnerUnixServer(str(unix_path), _Handler)
unix_path.chmod(0o600)
self._unix_inode = unix_path.stat().st_ino
else:
self._httpd = _OwnerHTTPServer((host, port), _Handler)
self._httpd.owner = self
self._thread: threading.Thread | None = None
@property
def port(self) -> int:
if self._unix_path is not None:
raise RequestRefused("Unix route has no TCP port")
return int(cast(_OwnerHTTPServer, self._httpd).server_address[1])
def start(self) -> None:
self._thread = threading.Thread(target=self._httpd.serve_forever, daemon=True)
self._thread.start()
def stop(self) -> None:
if self._thread is not None:
self._httpd.shutdown()
self._thread.join()
self._httpd.server_close()
if self._unix_path is not None:
try:
if self._unix_path.lstat().st_ino == self._unix_inode:
self._unix_path.unlink()
except FileNotFoundError:
pass