"""Local-only Reflection gateway with atomic, restart-safe token reservations. Pi never receives the provider credential. Every request, including summaries, passes through this gateway. A failed request with unknown usage keeps its full reservation as a conservative charge. No prompts or credentials enter logs. """ from __future__ import annotations import argparse from contextlib import contextmanager from datetime import datetime, timezone import hashlib from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer import json import math import os from pathlib import Path import re import secrets import select import socket import sqlite3 import threading import time import urllib.error import urllib.request import uuid ENDPOINT = "https://api.reflection.ai/openai/v1/chat/completions" CONTEXT_LIMIT = 262_144 MAX_OUTPUT = 16_000 ROUTE = re.compile(r"^/runs/([A-Za-z0-9_-]+)/agents/([A-Za-z0-9_-]+)/v1/chat/completions$") def utcnow(): return datetime.now(timezone.utc).isoformat() class BudgetExceeded(Exception): pass class Ledger: def __init__(self, path, config): self.path = str(path) Path(path).parent.mkdir(parents=True, exist_ok=True) self.config = config self.lock = threading.RLock() with self.connect() as db: db.executescript(""" CREATE TABLE IF NOT EXISTS requests ( id TEXT PRIMARY KEY, run_id TEXT NOT NULL, agent_id TEXT NOT NULL, category TEXT NOT NULL, reserved INTEGER NOT NULL, charged INTEGER NOT NULL DEFAULT 0, status TEXT NOT NULL, prompt_tokens INTEGER, completion_tokens INTEGER, started_at TEXT NOT NULL, ended_at TEXT, elapsed REAL, request_sha256 TEXT NOT NULL, http_status INTEGER, remaining_daily_tokens INTEGER, context_reserved INTEGER NOT NULL DEFAULT 0, context_charged INTEGER NOT NULL DEFAULT 0, fresh_input_estimate INTEGER NOT NULL DEFAULT 0, new_window INTEGER NOT NULL DEFAULT 0, history_hashes TEXT, client_request_id TEXT, purpose TEXT, context_window TEXT ); CREATE TABLE IF NOT EXISTS state (key TEXT PRIMARY KEY, value TEXT); """) encoded = json.dumps(config, sort_keys=True) previous = db.execute("SELECT value FROM state WHERE key='config'").fetchone() if previous and previous[0] != encoded: raise ValueError("Refusing to change budgets or run identities in an existing ledger") db.execute("INSERT OR IGNORE INTO state VALUES ('config', ?)", (encoded,)) # After a crash, dispatched calls may have completed remotely. Keep # their entire reservation charged; never silently refund them. db.execute("UPDATE requests SET status='unknown_after_restart', charged=reserved, reserved=0, context_charged=context_reserved, context_reserved=0 WHERE status='inflight'") @contextmanager def connect(self): db = sqlite3.connect(self.path, timeout=30) db.row_factory = sqlite3.Row db.execute("PRAGMA journal_mode=WAL") try: with db: yield db finally: db.close() def run_config(self, run_id): for row in self.config["runs"]: if row["run_id"] == run_id: return row raise BudgetExceeded("unregistered_run") def reserve(self, run_id, agent_id, request_hash, output_limit, body=None, metadata=None): run = self.run_config(run_id) # Reflection's model context bound is a conservative prompt bound. It # covers serialization/tool-schema overhead without guessing char/4. reservation = CONTEXT_LIMIT + output_limit with self.lock, self.connect() as db: db.execute("BEGIN IMMEDIATE") rows = db.execute("SELECT run_id,category,charged+reserved AS used FROM requests").fetchall() total = sum(row["used"] for row in rows) per_run = sum(row["used"] for row in rows if row["run_id"] == run_id) category = run.get("category", "benchmark") category_used = sum(row["used"] for row in rows if row["category"] == category) caps = [(total, self.config["global_api_token_cap"], "global_api_budget"), (per_run, run["api_token_cap"], "episode_api_budget"), (category_used, self.config["category_caps"][category], "category_api_budget")] for used, cap, reason in caps: if used + reservation > cap: raise BudgetExceeded(reason) quota = db.execute("SELECT value FROM state WHERE key='quota_blocked'").fetchone() if quota and quota[0] == "true": raise BudgetExceeded("daily_quota_headroom") hashes = [] fresh_input = 0 new_window = False context_reservation = 0 if body is not None and "context_token_cap" in run: incoming = [m for m in body.get("messages", []) if m.get("role") != "assistant"] encoded_messages = [json.dumps(m, sort_keys=True, separators=(",", ":")).encode() for m in incoming] hashes = [hashlib.sha256(m).hexdigest() for m in encoded_messages] key = f"history:{run_id}:{agent_id}" previous = db.execute("SELECT value FROM state WHERE key=?", (key,)).fetchone() previous = json.loads(previous[0]) if previous else [] new_window = not previous or hashes[:len(previous)] != previous or (metadata or {}).get("purpose") == "compaction" if new_window: # Approximation for admission only; reconciled to the actual # full prompt count on a context-window transition. fresh_input = math.ceil(len(json.dumps(body,ensure_ascii=False).encode()) / 4) + 1024 else: fresh_input = sum(math.ceil(len(m) / 4) + 8 for m in encoded_messages[len(previous):]) context_reservation = fresh_input + output_limit agent_used, episode_used = 0, 0 for context_row in db.execute("SELECT agent_id,context_charged+context_reserved AS n FROM requests WHERE run_id=?", (run_id,)): episode_used += context_row["n"] if context_row["agent_id"] == agent_id: agent_used += context_row["n"] if episode_used + context_reservation > run["context_token_cap"]: raise BudgetExceeded("episode_context_budget_estimated") if agent_used + context_reservation > run["agent_context_token_cap"]: raise BudgetExceeded("agent_context_budget_estimated") ident = uuid.uuid4().hex db.execute("INSERT INTO requests(id,run_id,agent_id,category,reserved,status,started_at,request_sha256,context_reserved,fresh_input_estimate,new_window,history_hashes) VALUES(?,?,?,?,?,'inflight',?,?,?,?,?,?)", (ident,run_id,agent_id,category,reservation,utcnow(),request_hash,context_reservation,fresh_input,int(new_window),json.dumps(hashes))) meta = metadata or {} db.execute("UPDATE requests SET client_request_id=?,purpose=?,context_window=? WHERE id=?",(meta.get("request_id"),meta.get("purpose"),meta.get("context_window"),ident)) return ident, reservation def settle(self, ident, usage, elapsed, http_status, remaining_daily=None, rejected=False): with self.lock, self.connect() as db: db.execute("BEGIN IMMEDIATE") row = db.execute("SELECT * FROM requests WHERE id=?", (ident,)).fetchone() if row["status"] != "inflight": raise ValueError("Request already settled") if usage is not None: prompt = int(usage["prompt_tokens"]) output = int(usage["completion_tokens"]) if min(prompt, output) < 0: raise ValueError("Negative provider usage") charge = prompt + output status = "complete" if charge > row["reserved"]: db.execute("INSERT OR REPLACE INTO state VALUES('quota_blocked','true')") status = "reservation_bound_violated" elif rejected: prompt = output = None charge = 0 status = "rejected" else: prompt = output = None charge = row["reserved"] status = "unknown_usage_conservatively_charged" context_charge = 0 if row["context_reserved"]: context_charge = (prompt if row["new_window"] else row["fresh_input_estimate"]) + output if usage else (0 if rejected else row["context_reserved"]) if usage and row["history_hashes"]: db.execute("INSERT OR REPLACE INTO state VALUES(?,?)", (f"history:{row['run_id']}:{row['agent_id']}",row["history_hashes"])) db.execute("UPDATE requests SET charged=?,reserved=0,status=?,prompt_tokens=?,completion_tokens=?,ended_at=?,elapsed=?,http_status=?,remaining_daily_tokens=? WHERE id=?", (charge,status,prompt,output,utcnow(),elapsed,http_status,remaining_daily,ident)) db.execute("UPDATE requests SET context_charged=?,context_reserved=0 WHERE id=?", (context_charge,ident)) # Headers count estimated input at admission, not finalized usage. # Preserve enough space for this completion and all other calls. if remaining_daily is not None: outstanding = db.execute("SELECT COALESCE(SUM(reserved),0) FROM requests").fetchone()[0] if remaining_daily < self.config.get("daily_headroom_tokens", 20_000_000) + outstanding + MAX_OUTPUT: db.execute("INSERT OR REPLACE INTO state VALUES('quota_blocked','true')") return dict(db.execute("SELECT * FROM requests WHERE id=?", (ident,)).fetchone()) def totals(self, run_id=None): with self.connect() as db: rows = db.execute("SELECT * FROM requests" + (" WHERE run_id=?" if run_id else ""), (run_id,) if run_id else ()).fetchall() return { "api_tokens_charged": sum(r["charged"] for r in rows), "api_tokens_reserved": sum(r["reserved"] for r in rows), "api_input_tokens": sum(r["prompt_tokens"] or 0 for r in rows), "api_output_tokens": sum(r["completion_tokens"] or 0 for r in rows), "api_requests": len(rows), "unknown_usage_requests": sum(r["status"].startswith("unknown") for r in rows), "context_tokens_estimated": sum(r["context_charged"] for r in rows), } class Gateway(ThreadingHTTPServer): daemon_threads = False def __init__(self, address, config, ledger_path, event_path, api_key, local_key, endpoint=ENDPOINT): super().__init__(address, Handler) self.config = config self.ledger = Ledger(ledger_path, config) self.events = Path(event_path) self.events.parent.mkdir(parents=True, exist_ok=True) self.event_lock = threading.Lock() self.api_key, self.local_key, self.endpoint = api_key, local_key, endpoint self.slots = threading.BoundedSemaphore(config.get("max_inflight_requests", 5)) def emit(self, run_id, kind, metrics=None, message=None): run = self.ledger.run_config(run_id) event = dict(timestamp=utcnow(),run_id=run_id,task=run.get("task"),condition=run.get("condition"),kind=kind) if metrics is not None: event["metrics"] = metrics if message is not None: event["message"] = message with self.event_lock: with self.events.open("a") as stream: stream.write(json.dumps(event) + "\n") class Handler(BaseHTTPRequestHandler): protocol_version = "HTTP/1.1" def log_message(self, *_): pass # The URL and raw request/response body are not application logs. def send_json(self, status, payload): data = json.dumps(payload).encode() self.send_response(status) self.send_header("Content-Type", "application/json") self.send_header("Content-Length", str(len(data))) self.send_header("Connection", "close") self.end_headers() self.wfile.write(data) self.close_connection = True def do_GET(self): if self.path == "/health": self.send_json(200, {"status":"ready", "model":"Beam-501B-A23B"}) else: self.send_json(404, {"error":"not_found"}) def do_POST(self): match = ROUTE.fullmatch(self.path) auth = self.headers.get("Authorization", "") if not secrets.compare_digest(auth, "Bearer " + self.server.local_key): return self.send_json(401, {"error":{"message":"Invalid local gateway credential"}}) if not match: return self.send_json(404, {"error":{"message":"Unregistered gateway route"}}) run_id, agent_id = match.groups() try: size = int(self.headers.get("Content-Length", "0")) except ValueError: return self.send_json(400, {"error":{"message":"Invalid request size"}}) if not 0 < size <= 8_000_000: return self.send_json(400, {"error":{"message":"Invalid request size"}}) try: body = json.loads(self.rfile.read(size)) if not isinstance(body, dict): raise ValueError("Request must be a JSON object") if body.get("model") != "Beam-501B-A23B": raise ValueError("Only the pinned Beam model is allowed") if not body.get("stream"): raise ValueError("Streaming with final usage is required") body.pop("max_completion_tokens", None) body.update(temperature=0.7, top_p=0.9, reasoning_effort="medium") body["max_tokens"] = min(int(body.get("max_tokens") or MAX_OUTPUT), MAX_OUTPUT) if body["max_tokens"] <= 0: raise ValueError("Positive output cap required") body["stream_options"] = {"include_usage": True} data = json.dumps(body, separators=(",", ":")).encode() except (ValueError, TypeError, KeyError) as error: return self.send_json(400, {"error":{"message": str(error)}}) if not self.server.slots.acquire(timeout=660): return self.send_json(429, {"error":{"message":"Local concurrency limit; no provider request sent"}}) # An aborted Pi request may have waited behind another agent's summary. # Do not spend tokens after its local connection has already closed. try: disconnected=bool(select.select([self.connection],[],[],0)[0]) and not self.connection.recv(1,socket.MSG_PEEK) except (ConnectionResetError,OSError): disconnected=True if disconnected: self.server.slots.release() self.close_connection=True return ident = None try: metadata = {"request_id":self.headers.get("X-Study-Request-ID"),"purpose":self.headers.get("X-Study-Purpose"),"context_window":self.headers.get("X-Study-Context-Window")} metadata = {k:v for k,v in metadata.items() if v and re.fullmatch(r"[A-Za-z0-9_-]{1,100}",v)} ident, _ = self.server.ledger.reserve(run_id,agent_id,hashlib.sha256(data).hexdigest(),body["max_tokens"],body,metadata) except BudgetExceeded as error: self.server.slots.release() if str(error) != "unregistered_run": self.server.emit(run_id,"budget_stop",message=str(error)) return self.send_json(402, {"error":{"message":"Study budget stop: " + str(error),"code":"study_budget_exhausted"}}) except Exception: self.server.slots.release() return self.send_json(500, {"error":{"message":"Budget ledger unavailable; no provider request sent"}}) started = time.monotonic() usage = None status = None remaining = None rejected = False headers_sent = False client_alive = True try: request = urllib.request.Request(self.server.endpoint,data,headers={"Content-Type":"application/json","Authorization":"Bearer " + self.server.api_key}) with urllib.request.urlopen(request, timeout=180) as response: status = response.status header = response.headers.get("x-ratelimit-remaining-tokens-day") remaining = int(header) if header else None self.send_response(status) self.send_header("Content-Type", "text/event-stream") self.send_header("Cache-Control", "no-cache") self.send_header("Connection", "close") self.end_headers() headers_sent = True for line in response: if time.monotonic()-started > 600: raise TimeoutError("Provider request exceeded the ten-minute transport ceiling") if line.startswith(b"data:") and line[5:].strip() != b"[DONE]": event = json.loads(line[5:]) if event.get("usage"): candidate = event["usage"] if isinstance(candidate,dict) and all(isinstance(candidate.get(k),int) and candidate[k]>=0 for k in ("prompt_tokens","completion_tokens")): usage = candidate if client_alive: try: self.wfile.write(line) self.wfile.flush() except (BrokenPipeError, ConnectionResetError): client_alive = False # Keep reading even if Pi exits: usage must still be settled. except urllib.error.HTTPError as error: status = error.code # Reflection documents that rejected requests consume no capacity. rejected = status in (400,401,403,404,413,422,429) if not headers_sent: self.send_json(status, {"error":{"message":f"Reflection HTTP {status}; details withheld from solver logs"}}) error.close() except Exception as error: if not headers_sent: try: self.send_json(502, {"error":{"message":"Reflection transport error; reservation retained"}}) except (BrokenPipeError, ConnectionResetError): pass self.server.emit(run_id,"error",message=f"Provider transport failure: {type(error).__name__}") finally: self.close_connection = True try: record = self.server.ledger.settle(ident,usage,time.monotonic()-started,status,remaining,rejected) totals = self.server.ledger.totals(run_id) totals["last_request_seconds"] = record["elapsed"] if remaining is not None: totals["provider_daily_remaining_at_admission"] = remaining self.server.emit(run_id,"progress",metrics=totals) if record["status"] != "complete": self.server.emit(run_id,"error",message=record["status"]) finally: self.server.slots.release() def main(): parser = argparse.ArgumentParser(description=__doc__) parser.add_argument("--config", type=Path, required=True) parser.add_argument("--ledger", type=Path, required=True) parser.add_argument("--events", type=Path, required=True) parser.add_argument("--port", type=int, default=18093) args = parser.parse_args() config = json.loads(args.config.read_text()) server = Gateway(("127.0.0.1",args.port),config,args.ledger,args.events, os.environ["REFLECTION_API_KEY"],os.environ["PI_STUDY_API_KEY"]) print(json.dumps({"status":"ready","port":server.server_port,"global_api_token_cap":config["global_api_token_cap"]}),flush=True) try: server.serve_forever(poll_interval=0.25) except KeyboardInterrupt: pass finally: server.server_close() if __name__ == "__main__": main()