Spaces:
Running
Running
| # -*- 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" | |
| 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 | |
| # --------------------------------------------------------------------------- | |
| 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 | |
| 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"} | |