File size: 14,657 Bytes
59db361 3396ad2 59db361 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 307 308 309 310 311 312 313 314 315 316 317 318 319 320 321 322 323 324 325 326 327 328 329 330 331 332 333 334 335 336 337 338 339 340 341 342 343 344 345 346 347 | """Console HTTP surface. Loopback-only by deployment (compose binds 127.0.0.1);
this layer adds the Origin check and field-naming validation errors."""
from __future__ import annotations
import asyncio
import json
import queue
import re
import time
from pathlib import Path
import httpx
from fastapi import FastAPI, HTTPException, Request
from fastapi.responses import FileResponse, JSONResponse, StreamingResponse
from fastapi.staticfiles import StaticFiles
from je_validation.console import catalog
from je_validation.console.budgets import default_step_budget, default_token_ceiling
from je_validation.console.config import ConsoleSettings
from je_validation.console.models_proxy import PINNED_MODELS, fetch_models
from je_validation.console.runs import (
BASELINES,
DuplicateRunError,
RunManager,
RunRequest,
)
from je_validation.envir.llm_agent import SYSTEM_PROMPT
from je_validation.envir.run_config import DEFAULT_KNOBS, KNOB_RANGES, MIN_SEEDS_FOR_VARIANCE
BASELINE_ENTRIES = [
{"id": "baseline:flag_everything", "name": "Baseline: flag everything",
"prompt_price": 0.0, "completion_price": 0.0, "launchable": True},
{"id": "baseline:no_evidence", "name": "Baseline: no evidence",
"prompt_price": 0.0, "completion_price": 0.0, "launchable": True},
]
MODELS_CACHE_TTL = 900.0
STATIC_DIR = Path(__file__).parent / "static"
# mirrors runs.py's private _NON_TERMINAL; api.py has no access to it and must
# not import a leading-underscore name across modules, so it is redefined here
_NON_TERMINAL = {"queued", "running"}
# mirrors runs.py's private _RUN_ID_RE (same rule the manager enforces at
# launch); redefined here for the same cross-module-underscore-import reason.
_RUN_ID_RE = re.compile(r"^[A-Za-z0-9_-]{1,64}$")
# model slugs (e.g. "anthropic_claude-sonnet-4.5") legitimately contain dots,
# so the charset allows them; "." and ".." are rejected explicitly below since
# they'd otherwise be valid single path segments under this charset.
_MODELSLUG_RE = re.compile(r"^[A-Za-z0-9._-]{1,128}$")
REQUIRED_RUN_FIELDS = ("run_id", "task_id", "seed_count", "step_budget")
def _validate_run_id(run_id: str) -> None:
if not _RUN_ID_RE.match(run_id):
raise HTTPException(status_code=404, detail="not found")
def _validate_modelslug(modelslug: str) -> None:
if modelslug in (".", "..") or not _MODELSLUG_RE.match(modelslug):
raise HTTPException(status_code=404, detail="not found")
def _within_console_root(path: Path, console_root: Path) -> Path:
"""Belt-and-braces check: even with run_id/modelslug pre-validated, confirm
the joined path did not escape the console run directory. Violations 404
rather than reveal why, matching the charset-rejection responses above."""
resolved = path.resolve()
if not resolved.is_relative_to(console_root.resolve()):
raise HTTPException(status_code=404, detail="not found")
return resolved
def _load_models(settings: ConsoleSettings,
models_transport: httpx.BaseTransport | None) -> list[dict]:
try:
fetched = fetch_models(settings.openrouter_base_url, transport=models_transport)
for m in fetched:
m["launchable"] = True
return BASELINE_ENTRIES + fetched
except httpx.HTTPError:
pinned = [{**m, "launchable": False} for m in PINNED_MODELS]
return BASELINE_ENTRIES + pinned
def _field_for_value_error(msg: str) -> str:
if "knob" in msg:
return "knobs"
if "prompt_version" in msg:
return "prompt_version"
if "tools" in msg:
return "tools_enabled"
if "pricing" in msg:
return "models"
return "request"
def _is_terminal(manager, run_id: str) -> bool:
try:
status = manager.status(run_id)
except (KeyError, FileNotFoundError):
return True
return status.get("status") not in _NON_TERMINAL
def _parse_after(raw: str | None) -> int:
"""Parse the SSE resume position (Last-Event-ID header or ?after= query).
Invalid input falls back to 0 rather than 500ing: the client's seq-dedupe
makes a from-zero replay harmless."""
try:
return int(raw or 0)
except ValueError:
return 0
def create_app(settings: ConsoleSettings, manager: RunManager | None = None,
models_transport: httpx.BaseTransport | None = None) -> FastAPI:
app = FastAPI()
mgr = manager if manager is not None else RunManager(settings)
allowed_origins = {f"http://localhost:{settings.port}",
f"http://127.0.0.1:{settings.port}"}
if settings.space_host:
allowed_origins.add(f"https://{settings.space_host}")
models_cache: dict = {"models": None, "at": 0.0}
def get_models() -> list[dict]:
now = time.monotonic()
cached = models_cache["models"]
if cached is not None and now - models_cache["at"] < MODELS_CACHE_TTL:
return cached
models = _load_models(settings, models_transport)
models_cache["models"] = models
models_cache["at"] = now
return models
def run_dir(run_id: str) -> Path:
return settings.data_dir / "runs" / "console" / run_id
@app.middleware("http")
async def origin_check(request: Request, call_next):
origin = request.headers.get("origin")
if origin is not None and origin not in allowed_origins:
return JSONResponse({"error": "forbidden origin"}, status_code=403)
return await call_next(request)
@app.get("/api/health")
def health():
return {"ok": True}
@app.get("/api/tasks")
def list_tasks():
out = []
for task_id, task in catalog.built_tasks().items():
population = catalog.default_population(task_id)
step_budget = default_step_budget(population)
token_ceiling = default_token_ceiling(step_budget)
knobs = {**DEFAULT_KNOBS, "population": population}
knob_meta = {k: {"min": lo, "max": hi} for k, (lo, hi) in KNOB_RANGES.items()}
entry = {
"id": task_id,
"tier": task["tier"],
"brief": task["brief"],
"tools": list(task["tools"]),
"population_editable": task_id in catalog.POPULATION_EDITABLE,
"knob_meta": knob_meta,
"defaults": {"knobs": knobs, "step_budget": step_budget,
"token_ceiling": token_ceiling},
"weights": dict(task["weight"]),
}
if "dials" in task:
entry["dials"] = list(task["dials"])
out.append(entry)
return out
@app.get("/api/models")
def list_models():
return get_models()
@app.get("/api/tasks/{task_id}/prompt")
def task_prompt(task_id: str, version: str = "standard",
step_budget: int | None = None):
"""The exact text the agent receives: system prompt + assembled brief.
Preview only — the episode path builds its own copy via the same
brief_for(), so this can never drift from what actually runs."""
task = catalog.built_tasks().get(task_id)
if task is None:
raise HTTPException(status_code=404, detail=f"unknown task: {task_id}")
try:
brief = catalog.brief_for(task, version, step_budget=step_budget)
except ValueError as e:
return JSONResponse({"field": "prompt_version", "error": str(e)},
status_code=422)
return {"system_prompt": SYSTEM_PROMPT, "brief": brief}
@app.get("/api/runs")
def list_runs():
"""Run history: one summary per persisted run.json, newest first.
Reads disk directly (like /results and /trajectory) so runs from
earlier server processes are listed too, not just this one's."""
console_root = settings.data_dir / "runs" / "console"
if not console_root.exists():
return []
entries = []
for child in console_root.iterdir():
path = child / "run.json"
try:
data = json.loads(path.read_text(encoding="utf-8"))
request = data["request"]
entries.append((path.stat().st_mtime, {
"run_id": data["run_id"],
"status": data["status"],
"task_id": request["task_id"],
"tier": data.get("contract", {}).get("tier"),
"models": list(request.get("models") or ()),
"seed_count": request["seed_count"],
"prompt_version": request.get("prompt_version"),
"cost_usd": data.get("cost_usd"),
}))
except (OSError, ValueError, KeyError, TypeError):
continue # junk dir or partial write: skip, never 500
entries.sort(key=lambda e: e[0], reverse=True)
return [summary for _, summary in entries]
@app.post("/api/runs")
async def create_run(request: Request):
body = await request.json()
for field_name in REQUIRED_RUN_FIELDS:
if field_name not in body:
return JSONResponse({"field": field_name, "error": "required"},
status_code=422)
model_index = {m["id"]: m for m in get_models()}
pricing = {}
for model_id in body.get("models") or []:
if model_id in BASELINES:
continue
m = model_index.get(model_id)
if m is not None and m.get("launchable", True):
pricing[model_id] = {"prompt_price": m["prompt_price"],
"completion_price": m["completion_price"]}
tools_enabled = body.get("tools_enabled")
if tools_enabled is None:
# None means "the task's full tool list"; look it up via .get() so
# an unknown task_id doesn't crash here -- it still flows through
# unchanged to mgr.launch()'s own task lookup below.
task = catalog.built_tasks().get(body["task_id"])
if task is not None:
tools_enabled = task["tools"]
req = RunRequest(
run_id=body["run_id"],
task_id=body["task_id"],
models=tuple(body.get("models") or ()),
seed_count=body["seed_count"],
knobs=body.get("knobs") or {},
tools_enabled=tuple(tools_enabled) if tools_enabled is not None else None,
step_budget=body["step_budget"],
token_ceiling=body.get("token_ceiling"),
pricing=pricing,
prompt_version=body.get("prompt_version") or "standard",
)
try:
contract = mgr.launch(req)
except DuplicateRunError:
raise HTTPException(status_code=409,
detail=f"duplicate run_id: {req.run_id}") from None
except ValueError as e:
msg = str(e)
return JSONResponse({"field": _field_for_value_error(msg), "error": msg},
status_code=422)
seed_warning = req.seed_count < MIN_SEEDS_FOR_VARIANCE
return JSONResponse({"contract": contract, "seed_warning": seed_warning},
status_code=202)
@app.get("/api/runs/{run_id}")
def get_run(run_id: str):
_validate_run_id(run_id)
try:
return mgr.status(run_id)
except (KeyError, FileNotFoundError):
raise HTTPException(status_code=404,
detail=f"unknown run: {run_id}") from None
@app.get("/api/runs/{run_id}/events")
async def run_events(run_id: str, request: Request):
_validate_run_id(run_id)
header_id = request.headers.get("Last-Event-ID")
after = _parse_after(header_id or request.query_params.get("after"))
backlog, q = mgr.subscribe_with_replay(run_id, after)
async def gen():
try:
for evt in backlog:
yield f"id: {evt['seq']}\ndata: {json.dumps(evt)}\n\n"
while True:
if _is_terminal(mgr, run_id) and q.empty():
break
try:
evt = await asyncio.to_thread(q.get, timeout=15)
except queue.Empty:
yield ": keepalive\n\n"
continue
yield f"id: {evt['seq']}\ndata: {json.dumps(evt)}\n\n"
finally:
mgr.unsubscribe(run_id, q)
return StreamingResponse(gen(), media_type="text/event-stream")
@app.get("/api/runs/{run_id}/results")
def run_results(run_id: str):
_validate_run_id(run_id)
console_root = settings.data_dir / "runs" / "console"
path = _within_console_root(run_dir(run_id) / "results.json", console_root)
if not path.exists():
return JSONResponse({"error": "run not finished"}, status_code=409)
return json.loads(path.read_text(encoding="utf-8"))
@app.get("/api/runs/{run_id}/episodes/{modelslug}/{seed}/trajectory")
def run_trajectory(run_id: str, modelslug: str, seed: int):
_validate_run_id(run_id)
_validate_modelslug(modelslug)
console_root = settings.data_dir / "runs" / "console"
ep_dir = _within_console_root(run_dir(run_id) / modelslug / f"seed{seed}",
console_root)
jsonl_files = sorted(ep_dir.glob("*.jsonl")) if ep_dir.exists() else []
if not jsonl_files:
raise HTTPException(status_code=404, detail="trajectory not found")
lines = jsonl_files[0].read_text(encoding="utf-8").splitlines()
return [json.loads(line) for line in lines if line.strip()]
@app.get("/")
def root():
return FileResponse(STATIC_DIR / "index.html")
# static mount last: added after all API routes above
app.mount("/static", StaticFiles(directory=STATIC_DIR), name="static")
return app
# Module-level app for uvicorn (`uvicorn je_validation.console.api:app`).
# Guarded: on a checkout without data/ (e.g. `import api` during test
# collection, or any environment lacking JE_DATA_DIR/snapshots), instantiating
# ConsoleSettings.from_env() raises FileNotFoundError from latest_snapshot();
# importing this module must never crash, so app falls back to None and the
# real deployment (compose, with data/ present) is what actually serves it.
try:
app = create_app(ConsoleSettings.from_env())
except Exception:
app = None
|