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