"""/compare?ids=a,b,… — side-by-side of 2–6 entities of the same type with per-type dimension sets and provenance. API 1.1: benchmark dimensions are keyed by (benchmark, canonical metric, config_key) and only appear when EVERY compared model has a current result in that comparability group; `comparability` labels each benchmark dimension; `diff_only=1` keeps only differing dimensions; `mode=models|providers|hardware|companies|frameworks` asserts the expected type.""" from __future__ import annotations from typing import Any from fastapi import APIRouter, Query, Request from aiatlas.api.common import ( COMPANY_TYPES, PRICE_COLS, PRICE_FROM, ApiError, cached, csv, enrich_provenance, entity_summary, price_row, resolve_entity, ) from aiatlas.db import connection, fetch_all from aiatlas.ontology.benchmarks import TRUST_LABELS, comparability from aiatlas.services.frontier import best_per_model, config_summary, group_rows, load_results router = APIRouter(prefix="/api/v1/compare", tags=["compare"]) D = lambda key, label, kind="text", unit=None, source="attr": {"key": key, "label": label, "kind": kind, **({"unit": unit} if unit else {}), "source": source} DIMENSIONS: dict[str, list[dict[str, Any]]] = { "model": [D("parameter_count", "Parameters", "number", "params"), D("active_parameter_count", "Active parameters", "number", "params"), D("context_length", "Context window", "number", "tokens"), D("max_output_tokens", "Max output", "number", "tokens"), D("openness", "Openness"), D("license", "License"), D("modalities", "Modalities", "list"), D("release_date", "Release date", "date"), D("knowledge_cutoff", "Knowledge cutoff", "date"), D("status", "Status"), D("family", "Family"), D("architecture", "Architecture"), D("reasoning", "Reasoning", "bool"), D("tool_calling", "Tool calling", "bool"), D("best_input_per_mtok", "Best input price", "number", "USD / 1M tokens", "prices"), D("best_output_per_mtok", "Best output price", "number", "USD / 1M tokens", "prices"), D("provider_count", "Providers", "number", None, "prices")], "provider": [D("website", "Website"), D("pricing_url", "Pricing page"), D("model_count", "Models priced", "number", None, "prices"), D("min_input_per_mtok", "Min input price", "number", "USD / 1M tokens", "prices"), D("min_output_per_mtok", "Min output price", "number", "USD / 1M tokens", "prices"), D("features", "Features", "list", None, "prices")], "hardware": [D("kind", "Kind"), D("manufacturer", "Manufacturer"), D("memory_gb", "Memory", "number", "GB"), D("memory_bandwidth_gbs", "Memory bandwidth", "number", "GB/s"), D("tdp_watts", "TDP", "number", "W"), D("architecture", "Architecture"), D("memory_type", "Memory type"), D("runtimes", "Runtimes", "list"), D("release_date", "Release", "date")], "framework": [D("latest_version", "Latest version"), D("latest_release_at", "Latest release", "date"), D("license", "License"), D("language", "Language"), D("metric.stars", "GitHub stars", "number"), D("metric.forks", "Forks", "number"), D("repository_url", "Repository")], "company": [D("country", "Country"), D("founded", "Founded", "date"), D("headquarters", "Headquarters"), D("org_kind", "Kind"), D("website", "Website"), D("model_count", "Models", "number", None, "graph"), D("paper_count", "Papers", "number", None, "graph")], "benchmark": [D("category", "Category"), D("metric", "Metric"), D("unit", "Unit"), D("task", "Task"), D("result_count", "Results", "number", None, "results")], } MODES = {"models": "model", "providers": "provider", "hardware": "hardware", "companies": "company", "frameworks": "framework", "benchmarks": "benchmark"} TYPE_ALIASES = {"library": "framework", "runtime": "framework", "artifact": "model"} def _norm(v: Any) -> Any: if isinstance(v, list): return sorted(str(x) for x in v) if isinstance(v, float) and v.is_integer(): return int(v) return v async def compare_entities(keys: list[str], *, diff_only: bool = False, mode: str | None = None) -> dict[str, Any]: if not 2 <= len(keys) <= 6: raise ApiError(400, "ids must list between 2 and 6 entities") if mode and mode not in MODES: raise ApiError(400, f"mode must be one of {', '.join(MODES)}") async with connection() as conn: rows = [await resolve_entity(conn, k) for k in keys] types = {TYPE_ALIASES.get(r["entity_type"], r["entity_type"]) for r in rows} etype = TYPE_ALIASES.get(rows[0]["entity_type"], rows[0]["entity_type"]) if etype in COMPANY_TYPES: etype = "company" types = {"company"} if len(types) != 1: raise ApiError(400, f"all entities must share one type, got {', '.join(sorted(types))}") if mode and MODES[mode] != etype: raise ApiError(400, f"mode={mode} expects {MODES[mode]} entities, got {etype}") dims = [dict(d) for d in DIMENSIONS.get(etype, [])] eids = [r["id"] for r in rows] items: list[dict[str, Any]] = [] extra: dict[str, dict[str, Any]] = {eid: {} for eid in eids} prices_by: dict[str, list[dict[str, Any]]] = {eid: [] for eid in eids} results_by: dict[str, list[dict[str, Any]]] = {eid: [] for eid in eids} comp: dict[str, dict[str, Any]] = {} if etype == "model": for p in await fetch_all(conn, f"select {PRICE_COLS} from {PRICE_FROM} where p.model_id = any(cast(:ids as text[])) and p.valid_to is null order by p.input_per_mtok nulls last", ids=eids): prices_by[p["m_id"]].append(price_row(p)) for eid, plist in prices_by.items(): ins = [x["input_per_mtok"] for x in plist if x["input_per_mtok"] is not None and x["input_per_mtok"] > 0] outs = [x["output_per_mtok"] for x in plist if x["output_per_mtok"] is not None and x["output_per_mtok"] > 0] extra[eid] = {"best_input_per_mtok": min(ins) if ins else None, "best_output_per_mtok": min(outs) if outs else None, "provider_count": len({x["provider"]["id"] for x in plist if x["provider"]})} res = await load_results(conn, model_ids=eids, canonical_models_only=False) for r in res: results_by[r["model_id"]].append({"id": r["id"], "benchmark": {"id": r["benchmark_id"], "slug": r["benchmark_slug"], "name": r["benchmark_name"]}, "score": r["score"], "metric": r["metric_canonical"], "unit": r.get("unit"), "higher_is_better": r.get("higher_is_better"), "config": config_summary(r.get("config")), "config_key": r["config_key"], "trust_level": r["trust_level"], "evaluated_at": r.get("evaluated_at"), "observed_at": r["observed_at"], "source_url": r.get("source_url"), "tier": r.get("tier")}) for g in sorted(group_rows(res).values(), key=lambda g: (g["rows"][0]["benchmark_name"] or "", g["label"])): if g["models"] < set(eids) and not set(eids) <= g["models"]: continue best = {r["model_id"]: r for r in best_per_model(g["rows"], g["higher_is_better"])} if not all(eid in best for eid in eids): continue first = g["rows"][0] key = f"bench:{first['benchmark_slug']}:{g['metric']}:{g['config_key']}" ref = best[eids[0]] worst_level, reasons = "comparable", [] for eid in eids[1:]: level, why = comparability(ref.get("config"), best[eid].get("config"), ref.get("metric"), best[eid].get("metric")) if level == "not-comparable" or (level == "partially-comparable" and worst_level == "comparable"): worst_level = level reasons += [w for w in why if w not in reasons and level != "comparable"] dims.append({"key": key, "label": f"{first['benchmark_name']} · {g['label']}", "kind": "number", "unit": first.get("unit"), "source": "results", "higher_is_better": g["higher_is_better"], "benchmark": first["benchmark_slug"], "metric": g["metric"], "config_key": g["config_key"], "comparability": worst_level, "trust_levels": sorted({best[e]["trust_level"] for e in eids})}) comp[key] = {"level": worst_level, "reasons": reasons or ["same variant, metric and evaluation conditions"], "trust": {e: {"level": best[e]["trust_level"], "label": TRUST_LABELS.get(best[e]["trust_level"], best[e]["trust_level"])} for e in eids}} for eid in eids: extra[eid][key] = best[eid]["score"] elif etype == "provider": agg = await fetch_all(conn, """select provider_id, count(distinct model_id) as model_count, min(nullif(input_per_mtok, 0)) as min_input_per_mtok, min(nullif(output_per_mtok, 0)) as min_output_per_mtok, jsonb_agg(distinct k.key) filter (where k.key is not null) as features from prices p left join lateral jsonb_object_keys(p.features) k(key) on true where provider_id = any(cast(:ids as text[])) and valid_to is null group by 1""", ids=eids) for a in agg: extra[a["provider_id"]] = {"model_count": int(a["model_count"]), "min_input_per_mtok": a["min_input_per_mtok"], "min_output_per_mtok": a["min_output_per_mtok"], "features": a["features"] or []} elif etype == "company": agg = await fetch_all(conn, """select e.id, (select count(*) from entities m where m.organization_id = e.id and m.entity_type = 'model' and m.merged_into is null) as model_count, (select count(*) from entities m where m.organization_id = e.id and m.entity_type = 'paper' and m.merged_into is null) as paper_count from entities e where e.id = any(cast(:ids as text[]))""", ids=eids) for a in agg: extra[a["id"]] = {"model_count": int(a["model_count"]), "paper_count": int(a["paper_count"])} elif etype == "benchmark": agg = await fetch_all(conn, "select benchmark_id, count(*) as n from benchmark_results where benchmark_id = any(cast(:ids as text[])) and valid_to is null group by 1", ids=eids) for a in agg: extra[a["benchmark_id"]] = {"result_count": int(a["n"])} await enrich_provenance(conn, *(r.get("provenance") for r in rows)) for r in rows: attrs = r.get("attributes") or {} prov = r.get("provenance") or {} values = {d["key"]: (extra[r["id"]].get(d["key"]) if d.get("source") != "attr" else attrs.get(d["key"])) for d in dims} item: dict[str, Any] = {"entity": entity_summary(r), "values": values, "provenance": {k: prov[k] for k in values if k in prov}} if etype == "model": item["prices"] = prices_by[r["id"]] item["results"] = results_by[r["id"]] items.append(item) if diff_only: keep = [d for d in dims if len({repr(_norm(it["values"].get(d["key"]))) for it in items}) > 1] kept = {d["key"] for d in keep} dims = keep for it in items: it["values"] = {k: v for k, v in it["values"].items() if k in kept} it["provenance"] = {k: v for k, v in it["provenance"].items() if k in kept} return {"entity_type": etype, "dimensions": dims, "items": items, "comparability": comp, "diff_only": diff_only, "note": "Benchmark dimensions appear only when every compared model has a current result in the same comparability group (benchmark × metric × config_key)."} @router.get("") @cached(300) async def compare(request: Request, ids: str = Query(..., description="2–6 slugs or ids, comma-separated"), diff_only: int = Query(0, ge=0, le=1), mode: str | None = Query(None, description="models|providers|hardware|companies|frameworks|benchmarks")) -> dict[str, Any]: return await compare_entities(csv(ids), diff_only=bool(diff_only), mode=mode) __all__ = ["DIMENSIONS", "compare_entities", "router"]