feat(simulator): curators X-ray experience simulator (sub-project 3) #3
@@ -0,0 +1,96 @@
|
||||
"""FastAPI service: dials -> real hef.selection -> X-ray (pick + ranked pool)."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from collections import Counter
|
||||
from pathlib import Path
|
||||
from typing import Literal, Optional
|
||||
|
||||
from fastapi import FastAPI
|
||||
from fastapi.staticfiles import StaticFiles
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from hef.catalog import load_catalog, record_to_dict
|
||||
from hef.selection import (
|
||||
Coordinate,
|
||||
Weights,
|
||||
candidates_for_mode,
|
||||
ranked_candidates,
|
||||
select,
|
||||
)
|
||||
from simulator.fixtures import generate_fixture_catalog
|
||||
|
||||
STATIC_DIR = Path(__file__).parent / "static"
|
||||
|
||||
|
||||
class SelectRequest(BaseModel):
|
||||
left: int = Field(ge=0, le=4)
|
||||
right: int = Field(ge=0, le=4)
|
||||
dark: int = Field(ge=0, le=4)
|
||||
light: int = Field(ge=0, le=4)
|
||||
mode: Literal["none", "audio", "video", "av"]
|
||||
pool_size: int = Field(default=4, ge=1, le=25)
|
||||
brain_weight: float = Field(default=1.0, ge=0.0)
|
||||
mood_weight: float = Field(default=1.0, ge=0.0)
|
||||
approved_only: bool = False
|
||||
|
||||
|
||||
def load_catalog_or_fixtures() -> list:
|
||||
"""Use the real catalog if a non-empty one is configured/exists, else fixtures."""
|
||||
configured = os.environ.get("HEF_SIM_CATALOG")
|
||||
path = Path(configured) if configured else Path("catalog/library.jsonl")
|
||||
if path.exists() and path.stat().st_size > 0:
|
||||
return load_catalog(path)
|
||||
return generate_fixture_catalog()
|
||||
|
||||
|
||||
def create_app(records: Optional[list] = None) -> FastAPI:
|
||||
app = FastAPI(title="HEF Experience Simulator")
|
||||
app.state.catalog = records if records is not None else load_catalog_or_fixtures()
|
||||
|
||||
@app.post("/api/select")
|
||||
def api_select(req: SelectRequest):
|
||||
catalog = app.state.catalog
|
||||
coord = Coordinate(req.left, req.right, req.dark, req.light)
|
||||
weights = Weights(brain=req.brain_weight, mood=req.mood_weight)
|
||||
if req.mode == "none":
|
||||
return {"pick": None, "pool": [], "coverage": {"candidates_in_mode": 0}}
|
||||
pool = catalog
|
||||
if req.approved_only:
|
||||
pool = [r for r in pool if r.review_status == "approved"]
|
||||
eligible = candidates_for_mode(pool, req.mode, req.pool_size)
|
||||
ranked = ranked_candidates(
|
||||
catalog, coord, req.mode,
|
||||
pool_size=req.pool_size, weights=weights, approved_only=req.approved_only,
|
||||
)
|
||||
pick = select(
|
||||
catalog, coord, req.mode,
|
||||
pool_size=req.pool_size, weights=weights, approved_only=req.approved_only,
|
||||
rng=None,
|
||||
)
|
||||
return {
|
||||
"pick": record_to_dict(pick) if pick else None,
|
||||
"pool": [
|
||||
{"record": record_to_dict(r), "distance": d, "rank": i + 1}
|
||||
for i, (r, d) in enumerate(ranked)
|
||||
],
|
||||
"coverage": {"candidates_in_mode": len(eligible)},
|
||||
}
|
||||
|
||||
@app.get("/api/catalog/meta")
|
||||
def api_meta():
|
||||
catalog = app.state.catalog
|
||||
return {
|
||||
"total": len(catalog),
|
||||
"by_mode": dict(Counter(r.mode for r in catalog)),
|
||||
"by_status": dict(Counter(r.review_status for r in catalog)),
|
||||
}
|
||||
|
||||
if STATIC_DIR.exists():
|
||||
app.mount("/", StaticFiles(directory=STATIC_DIR, html=True), name="static")
|
||||
|
||||
return app
|
||||
|
||||
|
||||
app = create_app()
|
||||
@@ -0,0 +1,90 @@
|
||||
import pytest
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
from hef.catalog import Record
|
||||
from simulator.app import create_app
|
||||
|
||||
|
||||
def make_record(**overrides):
|
||||
base = dict(
|
||||
id="r",
|
||||
title="t",
|
||||
source_url="u",
|
||||
source_archive="internet_archive",
|
||||
license="public_domain",
|
||||
mode="video",
|
||||
left=0,
|
||||
right=0,
|
||||
dark=0,
|
||||
light=0,
|
||||
duration_s=600,
|
||||
file_path="",
|
||||
)
|
||||
base.update(overrides)
|
||||
return Record(**base)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def client():
|
||||
records = [
|
||||
make_record(id="v-near", mode="video", left=0, right=0, dark=0, light=0),
|
||||
make_record(id="v-far", mode="video", left=4, right=4, dark=4, light=4),
|
||||
make_record(id="a-one", mode="audio", left=1, right=1, dark=1, light=1),
|
||||
make_record(id="prop", mode="video", left=0, right=0, dark=0, light=1,
|
||||
review_status="proposed"),
|
||||
make_record(id="appr", mode="video", left=0, right=0, dark=0, light=1,
|
||||
review_status="approved"),
|
||||
]
|
||||
return TestClient(create_app(records=records))
|
||||
|
||||
|
||||
def _body(**overrides):
|
||||
base = dict(left=0, right=0, dark=0, light=0, mode="video")
|
||||
base.update(overrides)
|
||||
return base
|
||||
|
||||
|
||||
def test_select_returns_pick_and_ranked_pool(client):
|
||||
resp = client.post("/api/select", json=_body(mode="video", pool_size=4))
|
||||
assert resp.status_code == 200
|
||||
data = resp.json()
|
||||
assert data["pick"]["id"] == "v-near"
|
||||
ids = [c["record"]["id"] for c in data["pool"]]
|
||||
assert ids[0] == "v-near"
|
||||
assert all("distance" in c and "rank" in c for c in data["pool"])
|
||||
assert [c["rank"] for c in data["pool"]] == list(range(1, len(data["pool"]) + 1))
|
||||
|
||||
|
||||
def test_none_mode_is_the_void(client):
|
||||
resp = client.post("/api/select", json=_body(mode="none"))
|
||||
assert resp.status_code == 200
|
||||
data = resp.json()
|
||||
assert data["pick"] is None
|
||||
assert data["pool"] == []
|
||||
|
||||
|
||||
def test_dial_out_of_range_is_rejected(client):
|
||||
resp = client.post("/api/select", json=_body(left=7))
|
||||
assert resp.status_code == 422
|
||||
|
||||
|
||||
def test_bad_mode_is_rejected(client):
|
||||
resp = client.post("/api/select", json=_body(mode="banana"))
|
||||
assert resp.status_code == 422
|
||||
|
||||
|
||||
def test_approved_only_narrows_pool(client):
|
||||
resp = client.post("/api/select", json=_body(left=0, right=0, dark=0, light=1,
|
||||
mode="video", approved_only=True))
|
||||
data = resp.json()
|
||||
assert all(c["record"]["review_status"] == "approved" for c in data["pool"])
|
||||
|
||||
|
||||
def test_catalog_meta_reports_counts(client):
|
||||
resp = client.get("/api/catalog/meta")
|
||||
assert resp.status_code == 200
|
||||
data = resp.json()
|
||||
assert data["total"] == 5
|
||||
assert data["by_mode"]["video"] == 4
|
||||
assert data["by_mode"]["audio"] == 1
|
||||
assert set(data["by_status"]) == {"proposed", "approved"}
|
||||
Reference in New Issue
Block a user