llm-gateway/scripts/client/anthropic-interceptor.py
Rene Fichtmueller b8543ec191
All checks were successful
changelog-draft / changelog-draft (push) Successful in 31s
feat: proxy subscription clients through gateway
2026-07-18 18:39:18 +02:00

289 lines
11 KiB
Python

#!/usr/bin/env python3
"""Fail-closed local subscription companion for Claude Code and Codex."""
import asyncio
import json
import math
import os
import re
import subprocess
import time
import uuid
import httpx
from fastapi import FastAPI, Request, Response
from fastapi.responses import JSONResponse, StreamingResponse
GATEWAY_URL = os.getenv("LLM_GATEWAY_URL", "https://llm-gateway.context-x.org").rstrip("/")
KEYCHAIN_SERVICE = os.getenv("LLM_GATEWAY_KEYCHAIN_SERVICE", "llm-gateway-api-key")
CLAUDE_CALLER = os.getenv("INTERCEPTOR_CALLER", "claude-code-macbook")
CODEX_CALLER = os.getenv("CODEX_INTERCEPTOR_CALLER", "codex-macbook")
MAX_TOOL_CHARS = int(os.getenv("MAX_TOOL_CHARS", "8000"))
_PII = [
(re.compile(r"\b[A-Za-z0-9._%+\-]+@[A-Za-z0-9.\-]+\.[A-Za-z]{2,}\b"), "[EMAIL]"),
(re.compile(r"\b(?:\d{1,3}\.){3}\d{1,3}\b"), "[IP_ADDRESS]"),
(re.compile(r"\b(?:IBAN[:\s]*)?[A-Z]{2}\d{2}[A-Z0-9 ]{10,30}\b"), "[IBAN]"),
(re.compile(r"\beyJ[A-Za-z0-9_\-]+\.eyJ[A-Za-z0-9_\-]+\.[A-Za-z0-9_\-]+\b"), "[JWT]"),
(re.compile(r"\b(?:sk-proj-|sk-ant-)[A-Za-z0-9_\-]{20,}\b"), "[API_KEY]"),
(re.compile(r"\bAKIA[0-9A-Z]{16}\b"), "[AWS_KEY]"),
(re.compile(r"\b(?:ghp_|github_pat_)[A-Za-z0-9_]{20,}\b"), "[GH_TOKEN]"),
]
def gateway_key() -> str:
configured = os.getenv("LLM_GATEWAY_API_KEY", "").strip()
if configured:
return configured
try:
return subprocess.check_output(
["/usr/bin/security", "find-generic-password", "-s", KEYCHAIN_SERVICE, "-w"],
stderr=subprocess.DEVNULL,
timeout=5,
).decode().strip()
except Exception:
return ""
def anonymize(text: str) -> str:
for pattern, replacement in _PII:
text = pattern.sub(replacement, text)
return text
def anonymize_content(content):
if isinstance(content, str):
return anonymize(content)
if isinstance(content, list):
result = []
for part in content:
if isinstance(part, dict) and part.get("type") == "text":
part = {**part, "text": anonymize(str(part.get("text", "")))}
result.append(part)
return result
return content
def compress_messages(messages: list) -> list:
result = []
for original in messages:
message = dict(original)
content = message.get("content", "")
if message.get("role") == "tool" and isinstance(content, str) and len(content) > MAX_TOOL_CHARS:
message["content"] = content[:MAX_TOOL_CHARS] + "\n[truncated by subscription companion]"
elif message.get("role") == "user" and isinstance(content, list):
parts = []
for original_part in content:
part = dict(original_part) if isinstance(original_part, dict) else original_part
if isinstance(part, dict) and part.get("type") == "tool_result":
inner = part.get("content", "")
if isinstance(inner, str) and len(inner) > MAX_TOOL_CHARS:
part["content"] = inner[:MAX_TOOL_CHARS] + "\n[truncated by subscription companion]"
elif isinstance(inner, list):
trimmed = []
for original_inner_part in inner:
inner_part = dict(original_inner_part) if isinstance(original_inner_part, dict) else original_inner_part
if (
isinstance(inner_part, dict)
and inner_part.get("type") == "text"
and len(str(inner_part.get("text", ""))) > MAX_TOOL_CHARS
):
inner_part["text"] = str(inner_part["text"])[:MAX_TOOL_CHARS] + "\n[truncated by subscription companion]"
trimmed.append(inner_part)
part["content"] = trimmed
parts.append(part)
message["content"] = parts
result.append(message)
return result
def estimate_tokens(payload) -> int:
return max(1, math.ceil(len(json.dumps(payload, ensure_ascii=False)) / 4))
def prepare_anthropic_body(raw_body: bytes) -> tuple[bytes, int, int, bool]:
data = json.loads(raw_body or b"{}")
before = estimate_tokens(data)
is_stream = bool(data.get("stream"))
if isinstance(data.get("system"), str):
data["system"] = anonymize(data["system"])
elif isinstance(data.get("system"), list):
data["system"] = anonymize_content(data["system"])
messages = []
for original in data.get("messages", []):
message = dict(original)
message["content"] = anonymize_content(message.get("content", ""))
messages.append(message)
data["messages"] = compress_messages(messages)
after = estimate_tokens(data)
return json.dumps(data, ensure_ascii=False).encode(), before, after, is_stream
def response_headers(headers: httpx.Headers) -> dict[str, str]:
blocked = {"connection", "content-encoding", "content-length", "keep-alive", "transfer-encoding"}
return {name: value for name, value in headers.items() if name.lower() not in blocked}
async def report_usage(
key: str,
request_id: str,
caller: str,
model: str,
tokens_in: int,
tokens_out: int,
latency_ms: int,
before: int,
after: int,
) -> None:
try:
async with httpx.AsyncClient(timeout=5.0) as client:
await client.post(
f"{GATEWAY_URL}/api/tracking/report",
headers={"X-LLM-Gateway-Key": key},
json={
"request_id": request_id,
"caller": caller,
"model": model,
"tokens_in": tokens_in,
"tokens_out": tokens_out,
"latency_ms": latency_ms,
"source": "local-subscription-companion",
"compression_tokens_before": before,
"compression_tokens_after": after,
"compression_mode": "interceptor:tool-trim" if before > after else "interceptor:none",
},
)
except Exception:
return
app = FastAPI(title="llm-subscription-companion", docs_url=None, redoc_url=None)
@app.get("/health")
async def health():
configured = bool(gateway_key())
return JSONResponse(
status_code=200 if configured else 503,
content={"status": "ok" if configured else "auth_required", "gateway": GATEWAY_URL, "failClosed": True},
)
@app.api_route("/{path:path}", methods=["GET", "POST", "PUT", "DELETE", "OPTIONS", "HEAD", "PATCH"])
async def proxy(request: Request, path: str):
key = gateway_key()
if not key:
return JSONResponse(status_code=503, content={"error": "gateway authentication unavailable"})
started_at = time.monotonic()
raw_body = await request.body()
is_codex = path.startswith("openai/")
gateway_path = path[len("openai/"):] if is_codex else path
target_kind = "codex" if is_codex else "anthropic"
caller = CODEX_CALLER if is_codex else CLAUDE_CALLER
is_messages = not is_codex and gateway_path == "v1/messages" and request.method == "POST"
is_stream = False
before = after = estimate_tokens({})
if is_messages:
try:
raw_body, before, after, is_stream = prepare_anthropic_body(raw_body)
except Exception:
return JSONResponse(status_code=400, content={"error": "invalid Anthropic request body"})
request_id = f"{target_kind}-companion:{int(time.time() * 1000)}:{uuid.uuid4().hex[:8]}"
blocked = {"connection", "content-length", "host", "transfer-encoding", "x-llm-gateway-key"}
headers = {name: value for name, value in request.headers.items() if name.lower() not in blocked}
headers.update({
"X-LLM-Gateway-Key": key,
"X-LLM-Gateway-Upstream": target_kind,
"X-LLM-Interceptor-Caller": caller,
"X-LLM-Request-ID": request_id,
})
target = f"{GATEWAY_URL}/{gateway_path}"
params = dict(request.query_params)
if is_stream:
client = httpx.AsyncClient(timeout=300.0)
upstream_request = client.build_request(request.method, target, headers=headers, content=raw_body, params=params)
upstream = await client.send(upstream_request, stream=True)
async def stream_response():
pending = ""
model_name, tokens_in, tokens_out = "unknown", 0, 0
try:
async for chunk in upstream.aiter_bytes():
pending += chunk.decode(errors="ignore")
lines = pending.split("\n")
pending = lines.pop()
for line in lines:
if not line.startswith("data: "):
continue
try:
event = json.loads(line[6:])
except Exception:
continue
if event.get("type") == "message_start":
message = event.get("message", {})
model_name = message.get("model", model_name)
tokens_in = message.get("usage", {}).get("input_tokens", tokens_in)
elif event.get("type") == "message_delta":
tokens_out = event.get("usage", {}).get("output_tokens", tokens_out)
yield chunk
finally:
await upstream.aclose()
await client.aclose()
if tokens_in or tokens_out:
await report_usage(
key, request_id, caller, model_name, tokens_in, tokens_out,
int((time.monotonic() - started_at) * 1000), before, after,
)
return StreamingResponse(
stream_response(),
status_code=upstream.status_code,
headers=response_headers(upstream.headers),
media_type=upstream.headers.get("content-type", "text/event-stream"),
)
async with httpx.AsyncClient(timeout=300.0) as client:
upstream = await client.request(
request.method,
target,
headers=headers,
content=raw_body if request.method not in {"GET", "HEAD"} else None,
params=params,
)
if is_messages and upstream.is_success:
try:
response_data = upstream.json()
usage = response_data.get("usage", {})
asyncio.create_task(report_usage(
key,
request_id,
caller,
response_data.get("model", "unknown"),
int(usage.get("input_tokens", 0)),
int(usage.get("output_tokens", 0)),
int((time.monotonic() - started_at) * 1000),
before,
after,
))
except Exception:
pass
return Response(
content=upstream.content,
status_code=upstream.status_code,
headers=response_headers(upstream.headers),
)
if __name__ == "__main__":
import uvicorn
uvicorn.run(app, host="127.0.0.1", port=int(os.getenv("PORT", "3090")), log_level="warning")