beam-pi-programbench / source /study /budget_proxy.py
burtenshaw's picture
burtenshaw HF Staff
feat: publish beam pi study source
5741b22 verified
Raw History Blame Contribute Delete
20.6 kB
"""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()