mlx-model-explorer / tests /test_api.py
codelion's picture
Record one summary per visit instead of every click; hourly writes to data/incoming; monthly compaction
a866fc9 verified
Raw History Blame Contribute Delete
10.2 kB
import time
import httpx
import pytest
from fastapi.testclient import TestClient
from app.catalogue import Catalogue
from app.main import State, create_app
from app.model_detail import DetailCache, ModelDetail
from app.sink import LocalParquetSink
from .test_parsing import QWEN3_8B_CONFIG
LISTING = [
{"id": "mlx-community/Qwen3-8B-4bit", "downloads": 90000, "likes": 50, "lastModified": "2026-08-01T00:00:00Z",
"pipeline_tag": "text-generation", "safetensors": {"total": 8.19e9}, "tags": ["4-bit", "base_model:Qwen/Qwen3-8B"]},
{"id": "mlx-community/Qwen3-8B-8bit", "downloads": 30000, "likes": 10, "lastModified": "2026-07-01T00:00:00Z",
"pipeline_tag": "text-generation", "safetensors": {"total": 8.19e9}, "tags": ["8-bit", "base_model:Qwen/Qwen3-8B"]},
{"id": "mlx-community/Qwen3-14B-4bit", "downloads": 50000, "likes": 20, "lastModified": "2026-06-01T00:00:00Z",
"pipeline_tag": "text-generation", "safetensors": {"total": 14.8e9}, "tags": ["4-bit"]},
{"id": "mlx-community/Qwen3-mystery", "pipeline_tag": "text-generation", "tags": []},
{"id": "mlx-community/gemma-3-12b-it-4bit", "downloads": 70000, "pipeline_tag": "image-text-to-text",
"safetensors": {"total": 12e9}, "tags": ["4-bit"]},
{"id": "mlx-community/whisper-large-v3", "downloads": 99999, "pipeline_tag": "automatic-speech-recognition", "tags": []},
{"broken": True},
{"id": "mlx-community/secret-private-model", "private": True, "pipeline_tag": "text-generation", "tags": ["4-bit"]},
]
def fake_detail(mid):
if mid == "mlx-community/Qwen3-mystery":
return ModelDetail(mid, None, [], None, time.time(), error="config: HTTPStatusError")
return ModelDetail(mid, QWEN3_8B_CONFIG, [{"path": "model.safetensors", "size": 4_600_000_000}],
4_600_000_000, time.time())
@pytest.fixture
def client(tmp_path):
cat = Catalogue(fetcher=lambda: LISTING, cache_dir=tmp_path)
cat.refresh()
st = State(catalogue=cat, details=DetailCache(fetch=fake_detail, cache_dir=tmp_path),
sink=LocalParquetSink(tmp_path / "ds", flush_seconds=999), fetch_live=False)
with TestClient(create_app(st)) as c:
c.st = st
yield c
def test_healthz_and_meta(client):
h = client.get("/healthz").json()
assert h["ok"] and h["catalogue"]["models"] == 6 and h["catalogue"]["skipped_malformed"] == 2
m = client.get("/api/meta").json()
assert "Qwen" in [f["family"] for f in m["families"] + [{"family": "Qwen"}]]
assert 32768 in m["contexts"] and 36 in m["ram_classes"]
def test_index_served_with_csp(client):
r = client.get("/")
assert r.status_code == 200 and "MLX Model Explorer" in r.text
assert "default-src 'self'" in r.headers["content-security-policy"]
def test_recommend_flow(client):
r = client.post("/api/recommend", json={"family": "Qwen", "size": "8-15B", "quant": "4-bit",
"context": 32768, "ram_gb": 36, "ram_source": "confirmed",
"priority": "balanced"})
body = r.json()
ids = [x["model"]["id"] for x in body["results"]]
assert set(ids) == {"mlx-community/Qwen3-8B-4bit", "mlx-community/Qwen3-14B-4bit"}
top = body["results"][0]
assert top["memory"]["fit"] in ("Comfortable", "Likely") and top["reasons"]
assert top["model"]["url"].startswith("https://huggingface.co/mlx-community/")
def test_llm_filter_and_unknowns(client):
body = client.post("/api/recommend", json={"family": "Qwen", "context": 8192}).json()
ids = {x["model"]["id"] for x in body["results"]}
assert "mlx-community/Qwen3-mystery" in ids and "mlx-community/whisper-large-v3" not in ids
mystery = next(x for x in body["results"] if x["model"]["id"] == "mlx-community/Qwen3-mystery")
assert mystery["model"]["quant_bucket"] == "unknown" and mystery["memory"]["total_gb"] is None
assert any("isn't available" in n for n in body["notices"])
allm = client.post("/api/recommend", json={"llm_only": False, "context": 8192}).json()
assert allm["total"] == 6
def test_sorts(client):
pop = client.post("/api/recommend", json={"family": "Qwen", "sort": "popular"}).json()
dls = [x["model"]["downloads"] for x in pop["results"]]
assert dls == sorted(dls, reverse=True)
rec = client.post("/api/recommend", json={"family": "Qwen", "sort": "recent"}).json()
assert rec["results"][0]["model"]["id"] == "mlx-community/Qwen3-8B-4bit"
def test_bad_recommend_inputs(client):
assert client.post("/api/recommend", json={"context": 12345}).status_code == 422
assert client.post("/api/recommend", json={"priority": "evil"}).status_code == 422
assert client.post("/api/recommend", json={"ram_gb": -3}).status_code == 422
assert client.post("/api/recommend", json={"family": "NoSuchFamily"}).json()["total"] == 0
def test_model_detail(client):
d = client.get("/api/model/mlx-community/Qwen3-8B-4bit", params={"ram_gb": 16, "context": 32768}).json()
assert len(d["memory_by_context"]) == 7
assert [s["model"]["id"] for s in d["siblings"]] == ["mlx-community/Qwen3-8B-8bit"]
assert d["memory"]["weights_source"] == "files"
assert client.get("/api/model/mlx-community/nope").status_code == 404
junk = client.get("/api/model/mlx-community/Qwen3-8B-4bit", params={"context": 3, "priority": "x"})
assert junk.status_code == 200
def test_events_ingest_to_parquet(client, tmp_path):
sid = "abcdef0123456789"
batch = {"events": [
{"event_type": "session", "session_id": sid, "model_family": "Qwen", "quantization": "4-bit",
"target_context": 32768, "priority": "balanced", "hardware_source": "confirmed", "hardware_memory_class": 36,
"families_searched": ["Qwen"], "quants_searched": ["4-bit"], "distinct_queries": 2,
"models_viewed": ["mlx-community/Qwen3-8B-4bit"]},
{"event_type": "feedback", "session_id": sid, "selected_model": "mlx-community/Qwen3-8B-4bit",
"tried": "yes", "quality_rating": "good", "notes": "worked fine, email me x@y.com"},
{"event_type": "search", "session_id": sid, "model_family": "Qwen"},
]}
r = client.post("/api/events", json=batch)
assert r.status_code == 202 and r.json()["accepted"] == 2
assert r.json()["flags"][2] == ["dropped_legacy_event"]
# the visit re-sends its summary later: the newest row wins in stats
client.post("/api/events", json={"events": [dict(batch["events"][0], distinct_queries=5,
models_compared=["mlx-community/Qwen3-14B-4bit"])]})
assert client.st.sink.flush()
rows = client.st.sink.read_existing()
assert [x["event_type"] for x in rows] == ["session", "feedback", "session"]
assert rows[1]["hf_downloads_at_selection"] == 90000 and "x@y.com" not in rows[1]["notes"]
client.st.stats.invalidate()
stats = client.get("/api/stats").json()
assert stats["totals"]["sessions"] == 1 and stats["totals"]["engaged_sessions"] == 1
assert stats["totals"]["feedback"] == 1
assert stats["families"] == [{"key": "other", "count": 1, "share": 1.0}] # k-anonymity folds tiny buckets
def test_events_reject_malformed(client):
assert client.post("/api/events", content=b"not json", headers={"content-type": "application/json"}).status_code == 400
assert client.post("/api/events", json={"events": [{"event_type": "hack"}]}).status_code == 422
assert client.post("/api/events", json={"events": [{"event_type": "search", "extra": 1}]}).status_code == 422
big = b'{"events": [' + b'{"event_type":"search"},' * 5000 + b'{"event_type":"search"}]}'
assert client.post("/api/events", content=big, headers={"content-type": "application/json"}).status_code == 413
assert client.st.sink.pending == 0
def test_rate_limit(client):
client.st.limiter.burst = 3
client.st.limiter._buckets.clear()
codes = [client.post("/api/events", json={"events": [{"event_type": "search"}]}).status_code for _ in range(8)]
assert 429 in codes
def test_hub_failure_falls_back_to_snapshot(tmp_path):
def boom():
raise httpx.ConnectError("hub unreachable")
cat = Catalogue(fetcher=boom, cache_dir=tmp_path)
cat.refresh()
assert cat.source == "snapshot" and len(cat.records) > 1000 and cat.last_error
st = State(catalogue=cat, details=DetailCache(fetch=fake_detail, cache_dir=tmp_path),
sink=LocalParquetSink(tmp_path / "ds", flush_seconds=999), fetch_live=False)
with TestClient(create_app(st)) as c:
body = c.post("/api/recommend", json={"family": "Qwen", "quant": "4-bit"}).json()
assert body["total"] > 0 and any("cached copy" in n for n in body["notices"])
assert c.get("/healthz").json()["catalogue"]["degraded"]
def test_collection_disabled(client):
client.st.collection_enabled = False
r = client.post("/api/events", json={"events": [{"event_type": "search"}]})
assert r.json() == {"accepted": 0} and client.st.sink.pending == 0
def test_stats_one_count_per_visit_and_suppression():
from app.stats import compute
rows = []
for i in range(6):
sid = f"{i:016x}"
for n in range(3): # each visit re-sent its summary three times
rows.append({"event_type": "session", "session_id": sid, "timestamp": f"2026-09-14T01:0{n}:00Z",
"families_searched": ["Qwen"], "quants_searched": ["4-bit"], "target_context": 8192,
"hardware_memory_class": 36, "suspicious_flags": []})
rows.append({"event_type": "session", "session_id": "f" * 16, "timestamp": "2026-09-14T02:00:00Z",
"families_searched": ["Rare"], "quants_searched": ["3-bit"], "suspicious_flags": []})
rows.append({"event_type": "session", "session_id": "e" * 16, "families_searched": ["Flagged"],
"suspicious_flags": ["unknown_model_in_session"]})
s = compute(rows)
assert s["totals"]["sessions"] == 7
assert s["families"] == [{"key": "Qwen", "count": 6, "share": 0.8571}, {"key": "other", "count": 1, "share": 0.1429}]
assert s["ram_classes"][0] == {"key": 36, "count": 6, "share": 1.0}
assert s["quant_share"]["4-bit"] == 0.8571