decision-studio / app.py
Xunzhuo's picture
Simplify Decision 2.0 inputs and Arena model selection
a2032aa verified
Raw History Blame Contribute Delete
26.3 kB
"""Same-origin Studio and strict Decision Gateway."""
import asyncio
import json
import logging
import os
import time
from contextlib import asynccontextmanager
from pathlib import Path
from fastapi import FastAPI, HTTPException, Request
from fastapi.responses import JSONResponse, StreamingResponse
from fastapi.staticfiles import StaticFiles
from starlette.concurrency import run_in_threadpool
from contract import MODEL, to_records
from direct_gateway import DirectGateway, DirectGatewayError
from engine import Busy, Engine, Unavailable
from relay import Relay, RelayError
from systemone_api import (
arrange_models,
public_models,
resolve_canonical_model,
resolve_model,
public_response,
sdk_response,
)
from tetris_arena import (
ArenaBusyError,
ArenaValidationError,
HTTPDecisionAdapter,
RaceManager,
encode_sse,
)
ROOT = Path(__file__).resolve().parent
logger = logging.getLogger("decision.studio")
def unique_object(pairs):
result = {}
for key, value in pairs:
if key in result:
raise ValueError("Duplicate JSON key: " + key)
result[key] = value
return result
async def read_json(request, limit=256 * 1024, *, with_size=False):
if request.headers.get("content-type", "").split(";", 1)[0].strip().lower() != "application/json":
raise HTTPException(415, "Use Content-Type: application/json")
body = bytearray()
async for chunk in request.stream():
body.extend(chunk)
if len(body) > limit:
raise HTTPException(413, "Request is too large. No field is truncated.")
try:
document = json.loads(body, object_pairs_hook=unique_object,
parse_constant=lambda x: (_ for _ in ()).throw(ValueError("Nonfinite JSON value")))
return (document, len(body)) if with_size else document
except (ValueError, TypeError, UnicodeError, RecursionError) as exc:
raise HTTPException(422, str(exc)) from None
async def read_input(request):
payload = await read_json(request)
try:
records = to_records(payload)
except (ValueError, TypeError, UnicodeError, RecursionError) as exc:
raise HTTPException(422, str(exc)) from None
return payload, records
def create_app(
mode=None,
*,
native_engine=None,
relay=None,
registry=None,
relays=None,
direct_gateway=None,
tetris_adapter=None,
tetris_manager=None,
):
mode = mode or os.getenv("DECISION_BACKEND", "native")
if mode not in {"native", "pull_queue", "direct"}:
raise ValueError("DECISION_BACKEND must be native, pull_queue, or direct")
from model_registry import model_registry
if relay is not None and registry is None:
registry = [{'id': MODEL, 'label': 'Kai', 'version': '1.0', 'manifest_sha256': relay.manifest}]
registry = arrange_models(model_registry(registry))
studio_generation = os.getenv("DECISION_STUDIO_GENERATION", "").strip()
if studio_generation not in {"", "1.0", "2.0"}:
raise ValueError("DECISION_STUDIO_GENERATION must be 1.0 or 2.0 when set")
studio_registry = {
wire: item for wire, item in registry.items()
if not studio_generation or item['version'] == studio_generation
}
if not studio_registry:
raise ValueError("The Studio generation must include a configured model")
default_model = next(iter(studio_registry))
canonical_models = tuple(item['repo_id'] for item in registry.values())
expected_artifacts = {item['repo_id']: item for item in registry.values()}
owns_direct_gateway = mode == "direct" and direct_gateway is None
if mode == "direct":
if any('revision' not in item for item in registry.values()):
raise ValueError("Direct serving mode requires a pinned Hub revision for every model")
direct_gateway = direct_gateway or DirectGateway.from_environment(expected_artifacts)
if getattr(direct_gateway, "models", None) != frozenset(canonical_models):
raise ValueError(
"The direct Gateway must match every configured canonical model"
)
gateway_pins = getattr(direct_gateway, "expected_artifacts", None)
if gateway_pins is not None and any(
gateway_pins.get(model) != {
"revision": item["revision"],
"manifest_sha256": item["manifest_sha256"],
**({"content_sha256": item["content_sha256"]}
if "content_sha256" in item else {}),
}
for model, item in expected_artifacts.items()
):
raise ValueError("The direct Gateway artifact pins must match the registry")
elif direct_gateway is not None:
raise ValueError("A direct Gateway may only be supplied in direct mode")
owns_tetris_manager = tetris_manager is None
owns_tetris_adapter = owns_tetris_manager and tetris_adapter is None
if tetris_manager is None:
tetris_adapter = tetris_adapter or HTTPDecisionAdapter.from_environment(
allowed_local_models=canonical_models,
include_cloud=studio_generation != "2.0",
local_endpoints=(
direct_gateway.local_systemone_endpoints() if mode == "direct" else None
),
direct_gateway=direct_gateway if mode == "direct" else None,
)
tetris_manager = RaceManager.from_environment(tetris_adapter)
@asynccontextmanager
async def lifespan(_api):
yield
try:
if owns_tetris_manager:
try:
await tetris_manager.aclose()
finally:
if owns_tetris_adapter:
await tetris_adapter.close()
finally:
if owns_direct_gateway:
await direct_gateway.aclose()
api = FastAPI(
title="Decision Studio",
version="0.6.0",
docs_url="/api/docs",
redoc_url=None,
lifespan=lifespan,
)
api.state.engine = native_engine or Engine()
if mode == "native" and (set(registry) != {MODEL} or api.state.engine.model != MODEL
or api.state.engine.manifest != registry[MODEL]['manifest_sha256']):
raise ValueError("Native mode requires only Decision 1.0 Kai with its matching manifest")
if mode == "pull_queue":
relays = relays or ({MODEL: relay} if relay else {
key: Relay(os.getenv("DECISION_WORKER_TOKEN", ""), item['manifest_sha256'], model=key, complete_input_tokens=item['complete_input_tokens'])
for key, item in registry.items()})
if set(relays) != set(registry) or any(r.model != key or r.manifest != registry[key]['manifest_sha256'] or r.complete_input_tokens != registry[key]['complete_input_tokens'] for key, r in relays.items()):
raise ValueError("Every queue must match its registry ID and native manifest")
else:
relays = {}
if relays and len({r.token for r in relays.values()}) != 1:
raise ValueError("This gateway uses one shared private worker token")
api.state.relays = relays
api.state.relay = relays.get(MODEL) # backward-compatible default for local tooling
api.state.registry = registry
api.state.direct = direct_gateway
api.state.tetris = tetris_manager
def select(model):
if not isinstance(model, str) or model not in registry:
raise HTTPException(422, "This model is not available in this Studio")
return relays.get(model)
async def direct_readiness(models):
"""Probe distinct configured Decision models without exposing their origins."""
selected = tuple(dict.fromkeys(models))
observations = await asyncio.gather(
*(direct_gateway.probe(model) for model in selected),
return_exceptions=True,
)
return {
model: isinstance(observed, dict) and observed.get('loaded') is True
for model, observed in zip(selected, observations, strict=True)
}
async def tetris_model_readiness(competitor_ids=None):
if mode not in {"direct", "pull_queue"}:
return {}
competitors = (
api.state.tetris.catalog()
if hasattr(api.state.tetris, 'catalog') else ()
)
selected = {
competitor.id: competitor.request_model
for competitor in competitors
if competitor.family == 'decision'
and competitor.ready
and (competitor_ids is None or competitor.id in competitor_ids)
}
if mode == "direct":
by_model = await direct_readiness(selected.values())
else:
wire_by_canonical = {
item['repo_id']: wire for wire, item in registry.items()
}
wires = tuple(dict.fromkeys(
wire_by_canonical[model]
for model in selected.values()
if model in wire_by_canonical
))
observations = await asyncio.gather(
*(relays[wire].status() for wire in wires),
return_exceptions=True,
)
by_model = {
registry[wire]['repo_id']: isinstance(observed, dict)
and observed.get('loaded') is True
for wire, observed in zip(wires, observations, strict=True)
}
return {
competitor_id: by_model.get(model, False)
for competitor_id, model in selected.items()
}
async def input_for(request, *, public_contract=None):
payload, request_bytes = await read_json(
request, limit=2 * 1024 * 1024, with_size=True,
)
if request_bytes > 256 * 1024 and (
not isinstance(payload, dict) or 'states' not in payload
):
raise HTTPException(413, "Request is too large. No field is truncated.")
if public_contract is not None:
expected = {'model', 'questions', 'states' if public_contract == 'batch' else 'state'}
if not isinstance(payload, dict) or set(payload) != expected:
raise HTTPException(422, f"Provide exactly {', '.join(sorted(expected))}.")
requested_model = payload.get('model') if isinstance(payload, dict) else None
model = (
resolve_canonical_model(requested_model, registry)
if public_contract is not None
else resolve_model(
requested_model if mode == "direct" else requested_model or default_model,
registry,
)
)
selected = select(model)
public_model = registry[model]['repo_id']
canonical_payload = (
dict(payload, model=public_model) if isinstance(payload, dict) else payload
)
if isinstance(payload, dict):
payload = dict(payload, model=model)
try:
records = to_records(payload, model=model)
except (ValueError, TypeError, UnicodeError, RecursionError) as exc:
if (
mode == "direct"
and "Expanded context/question input exceeds" in str(exc)
):
raise HTTPException(413, str(exc)) from None
raise HTTPException(422, str(exc)) from None
return payload, records, selected, public_model, canonical_payload, model
@api.exception_handler(RelayError)
async def relay_error(request, exc):
return JSONResponse({"detail": exc.message}, status_code=exc.code,
headers={"Retry-After": "2"} if exc.code in {429, 503} else None)
@api.exception_handler(DirectGatewayError)
async def direct_error(request, exc):
headers = {"Retry-After": exc.retry_after} if exc.retry_after else None
return JSONResponse({"detail": exc.message}, status_code=exc.code, headers=headers)
@api.get("/api/status")
async def status(model: str = default_model):
model = resolve_model(model, registry)
selected = select(model)
if mode == "direct":
observed = await direct_gateway.probe(registry[model]['repo_id'])
loaded = observed['loaded']
artifact = observed['artifact'] or {}
return {
"configured": True,
"loaded": loaded,
"phase": "ready" if loaded else "unavailable",
"model": model,
"manifest_sha256": registry[model]['manifest_sha256'],
"revision": registry[model]['revision'],
"loaded_revision": artifact.get('revision'),
"loaded_manifest_sha256": artifact.get('manifest_sha256'),
"loaded_content_sha256": artifact.get('content_sha256'),
"complete_input_tokens": registry[model]['complete_input_tokens'],
"context_batch": loaded,
"backend": "direct",
"live_inference": loaded,
"running": False,
"queued": 0,
}
return await selected.status() if selected else api.state.engine.status()
@api.get("/api/ready")
async def ready():
"""Return HTTP 503 when any configured model is unavailable."""
if mode == 'direct':
loaded = await direct_readiness(canonical_models)
elif relays:
observations = await asyncio.gather(
*(relay.status() for relay in relays.values()),
return_exceptions=True,
)
loaded = {
registry[key]['repo_id']: isinstance(observed, dict)
and observed.get('loaded') is True
for key, observed in zip(relays, observations, strict=True)
}
else:
observed = api.state.engine.status()
loaded = {
registry[default_model]['repo_id']: observed.get('loaded') is True
}
healthy = bool(loaded) and all(loaded.values())
return JSONResponse(
{
'backend': mode,
'status': 'ready' if healthy else 'unavailable',
'models': loaded,
},
status_code=200 if healthy else 503,
)
@api.get("/api/examples")
def examples():
return json.loads((ROOT / "examples.json").read_text())
@api.get("/api/tetris/config")
async def tetris_config():
# Credentials and upstream URLs intentionally never enter this response.
config = api.state.tetris.public_config()
if studio_generation:
visible_models = {item['repo_id'] for item in studio_registry.values()}
visible_competitors = {
competitor.id for competitor in api.state.tetris.catalog()
if competitor.family == 'decision'
and competitor.request_model in visible_models
}
config['competitors'] = [
competitor for competitor in config.get('competitors', ())
if (competitor['family'] != 'decision'
or competitor['id'] in visible_competitors)
and (studio_generation != "2.0" or competitor['family'] != 'cloud')
]
readiness = await tetris_model_readiness({
competitor['id'] for competitor in config.get('competitors', ())
})
for competitor in config.get('competitors', ()):
if competitor['id'] in readiness:
competitor['ready'] = competitor['ready'] and readiness[competitor['id']]
return config
@api.post("/api/tetris/races", status_code=201)
async def create_tetris_race(request: Request):
payload = await read_json(request)
if mode in {'direct', 'pull_queue'} and isinstance(payload, dict):
selected_ids = {
value for side in ('left', 'right')
if isinstance(value := payload.get(side), str)
}
readiness = await tetris_model_readiness(selected_ids)
if any(not ready for ready in readiness.values()):
raise HTTPException(503, "The selected Decision model is unavailable.")
try:
race = await api.state.tetris.create(payload)
except ArenaValidationError as exc:
headers = (
{"Retry-After": str(exc.retry_after)}
if isinstance(exc, ArenaBusyError)
else None
)
raise HTTPException(exc.status_code, str(exc), headers=headers) from None
return {
"id": race.id,
"status": race.status,
"events_url": f"/api/tetris/races/{race.id}/events",
"result_url": f"/api/tetris/races/{race.id}",
}
def get_tetris_race(race_id):
try:
return api.state.tetris.get(race_id)
except KeyError:
raise HTTPException(404, "Race not found") from None
@api.get("/api/tetris/races/{race_id}")
async def tetris_race(race_id: str):
return get_tetris_race(race_id).snapshot()
@api.get("/api/tetris/races/{race_id}/trace")
async def tetris_race_trace(race_id: str):
return get_tetris_race(race_id).snapshot(include_traces=True)
@api.get("/api/tetris/races/{race_id}/events")
async def tetris_race_events(race_id: str, request: Request, after: int = 0):
race = get_tetris_race(race_id)
header_cursor = request.headers.get("last-event-id")
if header_cursor:
try:
after = max(after, int(header_cursor))
except ValueError:
raise HTTPException(400, "Last-Event-ID must be an integer") from None
if after < 0:
raise HTTPException(400, "after must not be negative")
async def stream():
async for event in race.iter_events(after):
yield encode_sse(event)
return StreamingResponse(
stream(),
media_type="text/event-stream",
headers={
"Cache-Control": "no-store",
"Connection": "keep-alive",
"X-Accel-Buffering": "no",
},
)
@api.delete("/api/tetris/races/{race_id}", status_code=204)
async def cancel_tetris_race(race_id: str):
try:
await api.state.tetris.cancel(race_id)
except KeyError:
raise HTTPException(404, "Race not found") from None
@api.get("/v1/models")
def models():
synchronous_wait = (
getattr(direct_gateway, "timeout_seconds", 45)
if mode == "direct"
else 45
)
limits = {
"request_bytes": 256 * 1024,
"expanded_input_bytes": 16 * 1024 * 1024,
"questions": None,
"contexts": None,
# direct capacity belongs to each live model instance, not the
# retired Gateway's fixed eight-request/eight-row profile.
"active_requests_per_model": None if mode == "direct" else 8,
"gpu_microbatch": None if mode == "direct" else 8,
"synchronous_wait_seconds": synchronous_wait,
}
result = {
"models": public_models(studio_registry),
"limits": limits,
}
if mode != "direct":
result["default"] = registry[default_model]['repo_id']
return result
@api.post("/v1/systemone")
async def evaluate_systemone(request: Request):
return await evaluate(request, public_contract='single')
@api.post("/v1/systemone/batches")
async def evaluate_decision_batches(request: Request):
return await evaluate(request, public_contract='batch')
@api.post("/api/evaluate")
async def evaluate_studio(request: Request):
return await evaluate(request, public_contract=None)
async def evaluate(request: Request, *, public_contract: str | None):
payload, records, selected, public_model, canonical_payload, _ = await input_for(
request, public_contract=public_contract
)
if mode == "direct":
return await direct_gateway.evaluate(
canonical_payload,
batch='states' in canonical_payload,
)
if selected:
job = await selected.submit(payload, keep_until_read=True)
deadline = time.monotonic() + 45
try:
while time.monotonic() < deadline:
remaining = deadline - time.monotonic()
current = await selected.wait(job["id"], max(0.0, min(0.25, remaining)))
if current["status"] == "succeeded":
if current['result'].get('model') == public_model:
return current['result']
if public_contract is not None:
return public_response(current['result'], public_model=public_model,
batch=public_contract == 'batch')
return sdk_response(current['result'])
if current["status"] not in Relay.ACTIVE:
raise HTTPException(503, current.get("detail", "The request did not finish"))
if await request.is_disconnected():
raise asyncio.CancelledError()
raise HTTPException(504, "Request exceeded the synchronous wait. Use the Studio asynchronous job API.")
finally:
await selected.release(job["id"])
try:
result = await run_in_threadpool(api.state.engine.evaluate, payload, records)
if public_contract is not None:
return public_response(result, public_model=public_model,
batch=public_contract == 'batch')
return sdk_response(result)
except Unavailable as exc:
raise HTTPException(503, str(exc)) from None
except Busy as exc:
raise HTTPException(429, str(exc), headers={"Retry-After": "2"}) from None
except ValueError as exc:
if "tokens" in str(exc) or "room for state" in str(exc):
raise HTTPException(422, str(exc)) from None
logger.warning("Native validation failure: %s", type(exc).__name__)
raise HTTPException(503, "The configured native package could not be validated. Check the server configuration.") from None
except Exception as exc:
logger.warning("Inference failure: %s", type(exc).__name__)
raise HTTPException(503, "The model could not finish this request. No prediction was substituted.") from None
if relays:
@api.post("/api/jobs", status_code=202)
async def submit(request: Request):
payload, _, selected, _, _, _ = await input_for(request)
return await selected.submit(payload)
@api.get("/api/jobs/{job_id}")
async def read_job(job_id: str, model: str = default_model):
return await select(resolve_model(model, registry)).read(job_id)
@api.delete("/api/jobs/{job_id}")
async def cancel_job(job_id: str, model: str = default_model):
return await select(resolve_model(model, registry)).cancel(job_id)
async def worker_body(request, required, optional=(), *, limit=512 * 1024):
# All queues use the same private trust boundary, but identities are mandatory.
# Authentication precedes reading user data. The token is never returned or logged.
next(iter(relays.values())).authenticate(request.headers.get("authorization", ""))
body = await read_json(request, limit=limit)
required = set(required) | {"model", "manifest_sha256"}
if not isinstance(body, dict) or set(body) - required - set(optional) or required - set(body):
raise HTTPException(422, "Invalid worker request fields")
selected = select(body['model'])
selected.authenticate(request.headers.get("authorization", ""))
if body['manifest_sha256'] != selected.manifest:
raise HTTPException(409, "Worker model and native manifest do not match this queue")
return body, selected
@api.post("/internal/worker/heartbeat")
async def heartbeat(request: Request):
body, selected = await worker_body(request, {"worker_id", "phase"}, {"capabilities"})
if any(other is not selected and other.worker and other.worker['id'] == body['worker_id']
for other in relays.values()):
raise HTTPException(409, "A worker identity cannot serve multiple models")
return await selected.heartbeat(body["worker_id"], body["manifest_sha256"], body["phase"], body.get("capabilities"))
@api.post("/internal/worker/claim")
async def claim(request: Request):
body, selected = await worker_body(request, {"worker_id"}, {"wait_seconds"})
return await selected.claim(body["worker_id"], body.get("wait_seconds", 20))
@api.post("/internal/worker/result")
async def complete(request: Request):
# Leave room for the worker/lease envelope around an 8 MiB
# attested runtime response.
body, selected = await worker_body(request, {"worker_id", "id", "lease_token"}, {"result", "error_code"}, limit=8 * 1024 * 1024 + 64 * 1024)
return await selected.complete(body["worker_id"], body["id"], body["lease_token"],
body.get("result"), body.get("error_code"))
@api.middleware("http")
async def response_headers(request, call_next):
response = await call_next(request)
response.headers["X-Content-Type-Options"] = "nosniff"
response.headers["Referrer-Policy"] = "no-referrer"
if request.url.path.startswith(("/api/", "/v1/", "/internal/", "/tetris/")):
response.headers["Cache-Control"] = "no-store"
return response
api.mount("/", StaticFiles(directory=ROOT / "static", html=True), name="studio")
return api
app = create_app()