# -*- coding: utf-8 -*- """ Serving tests: model bundle loading, box back-projection, label mapping, weight priors, and the live API surface (via FastAPI TestClient). Run: pytest tests/ -v Requires the deploy/latest bundle (model.onnx + names.json) in the repo. """ from __future__ import annotations import io import sys from pathlib import Path import numpy as np import pytest from PIL import Image REPO = Path(__file__).resolve().parents[1] sys.path.insert(0, str(REPO)) from ml.serving.labels import to_bucket, BUCKETS # noqa: E402 from ml.serving.weights import DEFAULT_PRIORS_G, estimate_weight_g, load_priors # noqa: E402 # --------------------------------------------------------------------------- # Pure logic # --------------------------------------------------------------------------- def test_bucket_identity_for_current_classes(): for b in BUCKETS: assert to_bucket(b) == b def test_bucket_mapping_taco_style(): assert to_bucket("Plastic bottle") == "plastic" assert to_bucket("Drink can") == "metal" assert to_bucket("Glass jar") == "glass" assert to_bucket("Battery") == "ewaste" assert to_bucket("Cardboard") == "paper" assert to_bucket("mystery item") == "other" assert to_bucket("") == "other" def test_weight_priors_cover_all_buckets(): priors = load_priors() for b in BUCKETS: assert estimate_weight_g(b, priors) > 0 assert priors == DEFAULT_PRIORS_G # --------------------------------------------------------------------------- # Model bundle # --------------------------------------------------------------------------- BUNDLE_DIR = REPO / "deploy" / "latest" @pytest.fixture(scope="session") def bundle(): from ml.serving.model_loader import ModelBundle return ModelBundle(BUNDLE_DIR) def test_bundle_loads_names(bundle): assert bundle.names == ["plastic", "paper", "glass", "metal", "organic", "ewaste", "other"] assert bundle.imgsz == 640 def test_backprojection_boxes_span_original_image(bundle, tmp_path): # Non-square image: with the old bug, boxes stayed in 640-space and could # never reach x > 640 on a 1920-wide original. w0, h0 = 1920, 1080 rng = np.random.default_rng(42) arr = (rng.random((h0, w0, 3)) * 255).astype(np.uint8) img_path = tmp_path / "noise.jpg" Image.fromarray(arr).save(img_path) tensor, meta = bundle.load_image(img_path) assert tensor.shape == (1, 3, 640, 640) assert meta[0] == w0 and meta[1] == h0 # Synthetic imgsz-space boxes at the letterbox edges must map to original corners. r = meta[2] pad_w, pad_h = meta[3], meta[4] boxes_imgsz = np.array([ [pad_w, pad_h, 640 - pad_w, 640 - pad_h], # full frame [320.0, 320.0, 480.0, 400.0], # center-right box ], dtype=np.float32) out = bundle._backproject_boxes(boxes_imgsz.copy(), meta) # Full-frame box must span (0,0)-(w0,h0), not stay in 640-space assert out[0][2] == pytest.approx(w0, abs=2.0) assert out[0][3] == pytest.approx(h0, abs=2.0) # Center box must be scaled by 1/r beyond the 640 grid assert out[1][2] == pytest.approx((480.0 - pad_w) / r, abs=2.0) assert out[1][2] > 640 # the old bug capped this at 640 def test_end_to_end_predict_returns_valid_boxes(bundle, tmp_path): w0, h0 = 1280, 960 img = Image.new("RGB", (w0, h0), (140, 150, 160)) p = tmp_path / "plain.jpg" img.save(p) result = bundle.predict(p, return_masks=False) assert result["orig_shape"] == [h0, w0] for b in result["boxes"]: x1, y1, x2, y2 = b["xyxy"] assert 0 <= x1 <= x2 <= w0 + 1 assert 0 <= y1 <= y2 <= h0 + 1 assert 0.0 <= b["conf"] <= 1.0 # --------------------------------------------------------------------------- # API surface # --------------------------------------------------------------------------- @pytest.fixture(scope="session") def client(tmp_path_factory): import os os.environ["ALAMI_AI_BUNDLE"] = str(BUNDLE_DIR) os.environ["ALAMI_FEEDBACK_DIR"] = str(tmp_path_factory.mktemp("fb")) from fastapi.testclient import TestClient from ml.serving.server import app return TestClient(app) def _jpeg_bytes(w=800, h=600) -> bytes: buf = io.BytesIO() Image.new("RGB", (w, h), (120, 130, 140)).save(buf, format="JPEG") return buf.getvalue() def test_healthz(client): r = client.get("/healthz") assert r.status_code == 200 body = r.json() assert body["ok"] is True assert body["names"][0] == "plastic" class _FakeQuery: """Chainable fake for the supabase query builder used by the dashboard.""" def __init__(self, counts, raise_on_or=False): self._counts = counts # dict: key -> count self._raise_on_or = raise_on_or self._corrected = False self._or = False def select(self, *a, **k): return self def limit(self, *a, **k): return self @property def not_(self): return self def is_(self, col, val): self._corrected = True return self def or_(self, expr): if self._raise_on_or: raise Exception("column trash_predictions.source does not exist") self._or = True return self def execute(self): class R: pass r = R() if self._corrected and self._or: r.count = self._counts["corrected_litter"] elif self._corrected: r.count = self._counts["corrected_all"] elif self._or: r.count = self._counts["litter"] else: r.count = self._counts["total"] return r class _FakeSB: def __init__(self, counts, raise_on_or=False): self._counts, self._raise = counts, raise_on_or def table(self, name): return _FakeQuery(dict(self._counts), self._raise) def test_dashboard_corpus_splits_pools(): from ml.serving import dashboard as d counts = {"total": 100, "litter": 80, "corrected_litter": 20, "corrected_all": 25} data = d.gather_metrics(REPO / "deploy" / "latest", REPO, lambda: _FakeSB(counts), "trash_predictions") c = data["corpus"] assert c["litter_predictions"] == 80 assert c["product_scan_pool"] == 20 # 100 - 80 assert c["corrected"] == 20 # litter-only corrections assert c["labeled_pct"] == 25.0 # 20/80 def test_dashboard_corpus_fallback_without_source_column(): from ml.serving import dashboard as d counts = {"total": 50, "litter": 0, "corrected_litter": 0, "corrected_all": 5} data = d.gather_metrics(REPO / "deploy" / "latest", REPO, lambda: _FakeSB(counts, raise_on_or=True), "trash_predictions") c = data["corpus"] assert c["litter_predictions"] == 50 # everything counts as litter assert c["product_scan_pool"] == 0 assert c["corrected"] == 5 # fallback corrected count # and the page must still render html = d.render_html(data, "vtest") assert "Brand-Radar-Pool" in html def test_dashboard_renders(client): r = client.get("/dashboard") assert r.status_code == 200 assert "text/html" in r.headers["content-type"] html = r.text # current model metrics from model_card.json are present assert "Alami Vision" in html and "ML Dashboard" in html assert "mAP@50" in html # per-class F1 table lists the material classes assert "plastic" in html and "organic" in html # corpus monitor section exists (Supabase not configured in test -> n/a, must not error) assert "Trainings-Korpus" in html def test_analyze_upload_contract(client): r = client.post( "/v1/analyze/upload", files={"file": ("test.jpg", _jpeg_bytes(), "image/jpeg")}, data={"user_id": "test-user", "domain": "trash"}, ) assert r.status_code == 200 body = r.json() assert body["model_version"] assert body["domain"] == "trash" assert body["image"] == {"width": 800, "height": 600} assert isinstance(body["objects"], list) s = body["summary"] assert s["item_count"] == len(body["objects"]) assert s["trash_detected"] == (s["item_count"] > 0) for obj in body["objects"]: assert obj["label"] in BUCKETS assert obj["raw_label"] assert obj["weight_source"] == "material_prior_v0" assert 0.0 <= obj["area_fraction"] <= 1.0 # AR-overlay fields: normalized bbox in 0..1, German label, hex colour assert len(obj["bbox_norm"]) == 4 assert all(0.0 <= v <= 1.0 for v in obj["bbox_norm"]) assert obj["label_de"] and obj["label_en"] and obj["color"].startswith("#") def test_materials_catalog(client): r = client.get("/v1/materials") assert r.status_code == 200 mats = r.json()["materials"] assert [m["bucket"] for m in mats] == list(BUCKETS) for m in mats: assert m["label_de"] and m["label_en"] assert m["color"].startswith("#") and len(m["color"]) == 7 by_bucket = {m["bucket"]: m for m in mats} assert by_bucket["glass"]["label_en"] == "Glass" assert by_bucket["glass"]["label_de"] == "Glas" def test_analyze_upload_source_tag_reaches_logging(client, monkeypatch): import ml.serving.server as srv captured = {} def fake_log(prediction_id, image_ref, user_id, mv, preds, endpoint, source=None): captured["source"] = source monkeypatch.setattr(srv, "log_prediction", fake_log) r = client.post("/v1/analyze/upload", files={"file": ("t.jpg", _jpeg_bytes(), "image/jpeg")}, data={"log": "true", "source": "product-scan"}) assert r.status_code == 200 assert captured["source"] == "product-scan" def test_materials_catalog_has_disposal_hints(client): r = client.get("/v1/materials") mats = {m["bucket"]: m for m in r.json()["materials"]} for b in BUCKETS: assert mats[b]["disposal_hint_en"] and mats[b]["disposal_hint_de"] assert "Pfand" in mats["plastic"]["disposal_hint_de"] assert "deposit" in mats["plastic"]["disposal_hint_en"] def test_analyze_upload_preview_mode_not_logged(client, monkeypatch): # log=false (preview) must NOT call log_prediction; log=true (default) must. import ml.serving.server as srv calls = {"n": 0} monkeypatch.setattr(srv, "log_prediction", lambda *a, **k: calls.__setitem__("n", calls["n"] + 1)) r = client.post("/v1/analyze/upload", files={"file": ("t.jpg", _jpeg_bytes(), "image/jpeg")}, data={"log": "false"}) assert r.status_code == 200 and calls["n"] == 0 # preview: not logged r = client.post("/v1/analyze/upload", files={"file": ("t.jpg", _jpeg_bytes(), "image/jpeg")}, data={"log": "true"}) assert r.status_code == 200 and calls["n"] == 1 # snapped: logged def test_analyze_rejects_unknown_domain(client): r = client.post( "/v1/analyze/upload", files={"file": ("t.jpg", _jpeg_bytes(), "image/jpeg")}, data={"domain": "faces"}, ) assert r.status_code == 400 def test_feedback_roundtrip(client): r = client.post("/feedback", json={ "prediction_id": "00000000-0000-0000-0000-000000000001", "corrected_type": "plastic", "corrected_weight_kg": 0.5, "source": "pytest", "corrected_items": [{"index": 0, "corrected_label": "metal", "corrected_weight_g": 15.0}], }) assert r.status_code == 200 assert r.json()["ok"] is True def test_feedback_persists_notes_to_supabase(client, monkeypatch): """notes must reach the Supabase update — the local JSONL is ephemeral.""" import ml.serving.server as srv captured = {} class _Q: def update(self, payload): captured.update(payload) return self def eq(self, *a): return self def execute(self): return None class _SB: def table(self, name): return _Q() monkeypatch.setattr(srv, "get_supabase", lambda: _SB()) r = client.post("/feedback", json={ "prediction_id": "00000000-0000-0000-0000-000000000002", "corrected_type": "glass", "source": "pytest", "notes": "was 2 bottles, model saw 1", }) assert r.status_code == 200 and r.json()["ok"] is True assert captured["notes"] == "was 2 bottles, model saw 1" assert captured["feedback_source"] == "pytest" assert captured["corrected_type"] == "glass" def test_feedback_noop_without_changes(client): r = client.post("/feedback", json={"prediction_id": "x"}) assert r.status_code == 200 assert "ignored" in r.json()["message"] def _capture_sb(monkeypatch): """Mock Supabase that records the update() payload it receives.""" import ml.serving.server as srv captured = {} class _Q: def update(self, payload): captured.update(payload); return self def eq(self, *a): return self def execute(self): return None class _SB: def table(self, name): return _Q() monkeypatch.setattr(srv, "get_supabase", lambda: _SB()) return captured def test_feedback_v2_signals_persist(client, monkeypatch): """v2 (#141): added_items (missed objects), reasons (failure chips) and corrected_items[].action must survive to the Supabase payload — the mobile client sends them live and they used to be silently dropped.""" captured = _capture_sb(monkeypatch) r = client.post("/feedback", json={ "prediction_id": "00000000-0000-0000-0000-00000000000a", "source": "pytest", "corrected_items": [{"index": 0, "action": "reject"}], "added_items": [ {"label": "plastic", "point": {"x": 0.4, "y": 0.6}}, {"label": "glass", "box": {"x": 0.1, "y": 0.1, "w": 0.2, "h": 0.3}, "count": 2}, ], "reasons": ["too_dark", "occluded"], }) assert r.status_code == 200 and r.json()["ok"] is True # action rides along inside the corrected_items JSONB (no own column) assert captured["corrected_items"][0]["action"] == "reject" assert captured["added_items"][0]["label"] == "plastic" assert captured["added_items"][0]["point"] == {"x": 0.4, "y": 0.6} assert captured["added_items"][1]["count"] == 2 assert captured["feedback_reasons"] == ["too_dark", "occluded"] def test_feedback_added_items_alone_is_not_noop(client, monkeypatch): """A submission where the AI missed everything carries only added_items — it is a real recall signal, not a no-op.""" captured = _capture_sb(monkeypatch) r = client.post("/feedback", json={ "prediction_id": "00000000-0000-0000-0000-00000000000b", "added_items": [{"label": "metal", "count": 1}], }) assert r.status_code == 200 and r.json()["ok"] is True assert "ignored" not in (r.json().get("message") or "") assert captured["added_items"][0]["label"] == "metal" def test_feedback_fallback_drops_only_missing_column(client, monkeypatch): """If a v2 column isn't migrated yet, only that column is dropped — the already-migrated ones (corrected_items) must still be written.""" import ml.serving.server as srv seen = {"attempts": []} class _Q: def __init__(self): self._payload = None def update(self, payload): self._payload = dict(payload); return self def eq(self, *a): return self def execute(self): seen["attempts"].append(self._payload) if "added_items" in self._payload: raise RuntimeError("PGRST204: Could not find the 'added_items' column in the schema cache") return None class _SB: def table(self, name): return _Q() monkeypatch.setattr(srv, "get_supabase", lambda: _SB()) r = client.post("/feedback", json={ "prediction_id": "00000000-0000-0000-0000-00000000000c", "corrected_type": "paper", "corrected_items": [{"index": 1, "corrected_label": "paper"}], "added_items": [{"label": "organic"}], }) assert r.status_code == 200 and r.json()["ok"] is True final = seen["attempts"][-1] assert "added_items" not in final # missing column dropped assert final["corrected_items"][0]["corrected_label"] == "paper" # kept assert final["corrected_type"] == "paper" # base field kept def test_feedback_weight_bounds(client): r = client.post("/feedback", json={"prediction_id": "x", "corrected_weight_kg": 99.0}) assert r.status_code == 422 or r.status_code == 400 def test_predict_legacy_contract_has_raw_label(client, tmp_path, monkeypatch): # Serve bytes without network: monkeypatch the fetcher import ml.serving.server as srv monkeypatch.setattr(srv, "fetch_image_bytes", lambda url: _jpeg_bytes()) r = client.post("/predict", json={"image_url": "https://example.test/img.jpg", "user_id": "u1"}) assert r.status_code == 200 body = r.json() assert set(body.keys()) == {"model_version", "inference_ms", "predictions", "prediction_id"} for p in body["predictions"]: assert set(p.keys()) == {"xyxy", "cls", "conf", "label", "raw_label"}