Spaces:
Running
Running
Download source/study/budget_proxy.py from burtenshaw/beam-pi-programbench: direct link, hf CLI and curl.
- Browser
- Download file 20.6 kB
-
https://huggingface.co/spaces/burtenshaw/beam-pi-programbench/resolve/main/source/study/budget_proxy.py
- Command line
-
hf download hf://spaces/burtenshaw/beam-pi-programbench/source/study/budget_proxy.py
-
curl -L -o budget_proxy.py https://huggingface.co/spaces/burtenshaw/beam-pi-programbench/resolve/main/source/study/budget_proxy.py
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'") | |
| 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() | |