imag vector generation with dimentionality reduction
This commit is contained in:
352
tests/test_image_match.py
Normal file
352
tests/test_image_match.py
Normal file
@@ -0,0 +1,352 @@
|
||||
"""Search by image: the ranking, the scope rules, and the SQL behind them.
|
||||
|
||||
No database and no model: `image_vector_search` is patched with canned rows,
|
||||
and the SQL-shape test uses a recorder cursor. The numbers are the ones the
|
||||
feasibility run produced - three Marie Gold pack sizes sharing one photo tie
|
||||
at 0.631 and a Good Day trails at 0.561 - so the tie-break is pinned to the
|
||||
case it was built for.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import inspect
|
||||
import math
|
||||
from typing import Any, Dict, List
|
||||
|
||||
import pytest
|
||||
|
||||
from app.services import image_match as im
|
||||
from app.services import vector_store
|
||||
from app.services.vector_store import fetch_products_by_image_ids as real_fetch_products_by_image_ids
|
||||
|
||||
|
||||
def _unit(seed: float = 1.0) -> List[float]:
|
||||
v = [math.sin(seed * (i + 1)) for i in range(1024)]
|
||||
n = math.sqrt(sum(x * x for x in v))
|
||||
return [x / n for x in v]
|
||||
|
||||
|
||||
def _row(name: str, distance: float, image_id: str = "", sizes=None, table="brand_britannia") -> Dict[str, Any]:
|
||||
return {
|
||||
"image_id": image_id or name.lower().replace(" ", "_"),
|
||||
"product_name": name,
|
||||
"title": name,
|
||||
"brand": "Britannia",
|
||||
"brand_table": table,
|
||||
"size_variants": sizes or [],
|
||||
"distance": distance,
|
||||
}
|
||||
|
||||
|
||||
MARIE = [
|
||||
_row("Britannia Marie Gold 89g", 0.369),
|
||||
_row("Britannia Marie Gold 300g", 0.369),
|
||||
_row("Britannia Marie Gold 1kg", 0.369),
|
||||
_row("Britannia Good Day 100g", 0.439),
|
||||
]
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def hydrate_from_the_light_rows(monkeypatch):
|
||||
"""Stand-in for the second read: the card for an image_id is the light
|
||||
row plus a `barcode`, so tests can see hydration happened and where."""
|
||||
def fake(table, image_ids):
|
||||
return [{**dict(r), "barcode": f"890-{r['image_id']}"} for r in MARIE
|
||||
if r["brand_table"] == table and r["image_id"] in image_ids]
|
||||
monkeypatch.setattr(vector_store, "fetch_products_by_image_ids", fake)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def calls(monkeypatch):
|
||||
"""Patch the store; return the list of (kwargs) it was called with."""
|
||||
seen: List[Dict[str, Any]] = []
|
||||
|
||||
def fake(vector, brand=None, top_k=10, category=None, ef_search=None):
|
||||
seen.append({"brand": brand, "top_k": top_k, "category": category, "ef_search": ef_search})
|
||||
return [dict(r) for r in MARIE]
|
||||
|
||||
monkeypatch.setattr(vector_store, "image_vector_search", fake)
|
||||
return seen
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# the vector
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def test_a_unit_vector_passes_through_untouched():
|
||||
v = _unit()
|
||||
assert im.normalise_vector(v) == v
|
||||
|
||||
|
||||
def test_an_unnormalised_vector_is_rescaled_to_unit_length():
|
||||
out = im.normalise_vector([2.0] + [0.0] * 1023)
|
||||
assert out[0] == 1.0 and abs(math.sqrt(sum(x * x for x in out)) - 1.0) < 1e-9
|
||||
|
||||
|
||||
@pytest.mark.parametrize("bad, msg", [
|
||||
([0.1] * 1023, "1024 values"),
|
||||
([float("nan")] + [0.0] * 1023, "NaN"),
|
||||
([0.0] * 1024, "all zeros"),
|
||||
(["x"] * 1024, "list of numbers"),
|
||||
])
|
||||
def test_unsearchable_vectors_are_refused_with_a_reason(bad, msg):
|
||||
with pytest.raises(im.InvalidVectorError, match=msg):
|
||||
im.normalise_vector(bad)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# the label text
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def test_sizes_are_normalised_and_words_cleaned():
|
||||
words, sizes = im.tokens("Britannia Marie Gold 300 g Net Wt")
|
||||
assert words == {"britannia", "marie", "gold"} and sizes == {"300g"}
|
||||
|
||||
|
||||
@pytest.mark.parametrize("text, size", [("89gm", "89g"), ("89 GMS", "89g"), ("1 LTR", "1l"), ("1 litre", "1l"),
|
||||
("2.50 kg", "2.5kg"), ("500ml", "500ml"), ("6 pcs", "6pcs")])
|
||||
def test_every_way_a_label_writes_a_size_collapses_to_one_token(text, size):
|
||||
assert im.tokens(text)[1] == {size}
|
||||
|
||||
|
||||
def test_empty_text_has_no_tokens_and_no_overlap():
|
||||
assert im.tokens(None) == (set(), set())
|
||||
assert im.text_overlap(set(), set(), MARIE[0]) == 0.0
|
||||
|
||||
|
||||
def test_size_matches_outweigh_word_matches():
|
||||
words, sizes = im.tokens("Marie Gold 300 g")
|
||||
assert im.text_overlap(words, sizes, MARIE[1]) == 3.0 + 2.0 # 300g + marie + gold
|
||||
assert im.text_overlap(words, sizes, MARIE[0]) == 2.0 # marie + gold only
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# ranking
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def test_label_text_picks_the_pack_size_among_tied_siblings(calls):
|
||||
result = im.search_by_vector(_unit(), text="Marie Gold 300 g")
|
||||
|
||||
names = [r["product_name"] for r in result.rows]
|
||||
assert names == [
|
||||
"Britannia Marie Gold 300g", # size token wins the tie
|
||||
"Britannia Marie Gold 1kg", # then name order among the rest
|
||||
"Britannia Marie Gold 89g",
|
||||
"Britannia Good Day 100g", # lower score, whatever the text
|
||||
]
|
||||
assert result.rows[0]["score"] == pytest.approx(0.631)
|
||||
assert result.rows[0]["text_overlap"] == 5.0
|
||||
assert result.rows[0]["barcode"] == "890-britannia_marie_gold_300g" # hydrated, and to the right row
|
||||
|
||||
|
||||
def test_without_text_ties_fall_back_to_name_order_and_overlap_is_zero(calls):
|
||||
result = im.search_by_vector(_unit())
|
||||
|
||||
assert [r["product_name"] for r in result.rows][:3] == [
|
||||
"Britannia Marie Gold 1kg", "Britannia Marie Gold 300g", "Britannia Marie Gold 89g",
|
||||
]
|
||||
assert all(r["text_overlap"] == 0.0 for r in result.rows)
|
||||
|
||||
|
||||
def test_min_score_drops_everything_below_it(calls):
|
||||
result = im.search_by_vector(_unit(), min_score=0.65)
|
||||
assert result.rows == [] and result.min_score == 0.65
|
||||
|
||||
|
||||
def test_top_k_truncates_after_ranking_and_only_winners_are_hydrated(calls, monkeypatch):
|
||||
asked = []
|
||||
monkeypatch.setattr(vector_store, "fetch_products_by_image_ids",
|
||||
lambda table, ids: asked.append((table, sorted(ids))) or [])
|
||||
result = im.search_by_vector(_unit(), text="300 g", top_k=1)
|
||||
|
||||
assert [r["product_name"] for r in result.rows] == ["Britannia Marie Gold 300g"]
|
||||
assert result.top_k == 1
|
||||
assert asked == [("brand_britannia", ["britannia_marie_gold_300g"])]
|
||||
|
||||
|
||||
def test_a_row_whose_card_cannot_be_read_keeps_its_light_form(calls, monkeypatch):
|
||||
monkeypatch.setattr(vector_store, "fetch_products_by_image_ids", lambda table, ids: [])
|
||||
result = im.search_by_vector(_unit(), top_k=2)
|
||||
assert len(result.rows) == 2 and "barcode" not in result.rows[0] and result.rows[0]["score"] == pytest.approx(0.631)
|
||||
|
||||
|
||||
def test_duplicate_rows_from_the_same_table_are_collapsed(monkeypatch):
|
||||
monkeypatch.setattr(vector_store, "image_vector_search",
|
||||
lambda *a, **k: [dict(MARIE[0]), dict(MARIE[0])])
|
||||
assert len(im.search_by_vector(_unit()).rows) == 1
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# scope
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def test_an_explicit_brand_is_a_hard_filter_that_never_falls_back(monkeypatch):
|
||||
seen = []
|
||||
monkeypatch.setattr(vector_store, "image_vector_search",
|
||||
lambda vector, brand=None, **k: seen.append(brand) or [])
|
||||
|
||||
result = im.search_by_vector(_unit(), brand="Cadbury", text="britannia marie")
|
||||
|
||||
assert seen == ["Cadbury"]
|
||||
assert result.rows == [] and result.scoped_to_brand and not result.scope_fallback
|
||||
assert result.detected_brand == "Cadbury"
|
||||
|
||||
|
||||
def test_a_brand_read_off_the_label_scopes_the_search(monkeypatch):
|
||||
seen = []
|
||||
monkeypatch.setattr(vector_store, "image_vector_search",
|
||||
lambda vector, brand=None, **k: seen.append(brand) or [dict(r) for r in MARIE])
|
||||
from app.services import query_intent
|
||||
monkeypatch.setattr(query_intent, "extract_brand_mention", lambda text: "Britannia")
|
||||
|
||||
result = im.search_by_vector(_unit(), text="Britannia Marie Gold")
|
||||
|
||||
assert seen == ["Britannia"]
|
||||
assert result.detected_brand == "Britannia" and result.scoped_to_brand and not result.scope_fallback
|
||||
|
||||
|
||||
def test_an_ocr_brand_that_finds_nothing_retries_across_every_brand(monkeypatch):
|
||||
seen = []
|
||||
|
||||
def fake(vector, brand=None, **k):
|
||||
seen.append(brand)
|
||||
return [] if brand else [dict(r) for r in MARIE]
|
||||
|
||||
monkeypatch.setattr(vector_store, "image_vector_search", fake)
|
||||
from app.services import query_intent
|
||||
monkeypatch.setattr(query_intent, "extract_brand_mention", lambda text: "Britannia")
|
||||
|
||||
result = im.search_by_vector(_unit(), text="Britannia marie")
|
||||
|
||||
assert seen == ["Britannia", None]
|
||||
assert result.scope_fallback and not result.scoped_to_brand and len(result.rows) == 4
|
||||
|
||||
|
||||
def test_no_text_and_no_brand_means_one_unscoped_query(calls):
|
||||
im.search_by_vector(_unit())
|
||||
assert len(calls) == 1 and calls[0]["brand"] is None
|
||||
|
||||
|
||||
def test_brand_detection_failures_do_not_break_the_search(calls, monkeypatch):
|
||||
from app.services import query_intent
|
||||
|
||||
def boom(text):
|
||||
raise RuntimeError("brand index unavailable")
|
||||
|
||||
monkeypatch.setattr(query_intent, "extract_brand_mention", boom)
|
||||
result = im.search_by_vector(_unit(), text="something")
|
||||
assert result.detected_brand is None and len(result.rows) == 4
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# candidate width and ef_search
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@pytest.mark.parametrize("top_k, fetch_k, ef", [(1, 30, 40), (10, 30, 40), (20, 60, 60), (50, 100, 100), (500, 100, 100)])
|
||||
def test_the_store_is_asked_for_enough_candidates_to_hold_the_ties(calls, top_k, fetch_k, ef):
|
||||
im.search_by_vector(_unit(), top_k=top_k)
|
||||
assert calls[0]["top_k"] == fetch_k and calls[0]["ef_search"] == ef
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# the SQL
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class _Cursor:
|
||||
def __init__(self):
|
||||
self.statements: List[str] = []
|
||||
self.description = None
|
||||
self._pending: Any = None
|
||||
|
||||
def execute(self, sql, params=None):
|
||||
text = " ".join(str(sql).split())
|
||||
self.statements.append(text)
|
||||
if "information_schema.columns" in text:
|
||||
self._pending = [("brand_x", c) for c in ("id", "product_name", "embedding", "img_vector", "img_vector_src")]
|
||||
elif "information_schema.tables" in text and "EXISTS" in text:
|
||||
self._pending = (True,)
|
||||
elif "information_schema.tables" in text:
|
||||
self._pending = [("brand_x",)]
|
||||
elif text.startswith("SET"):
|
||||
self._pending = None
|
||||
elif "AS distance" in text:
|
||||
self.description = [("image_id",), ("product_name",), ("title",), ("size_variants",), ("distance",)]
|
||||
self._pending = [("a", "P", "P", [], 0.2)]
|
||||
else:
|
||||
self.description = [("id",), ("product_name",), ("img_vector_src",)]
|
||||
self._pending = [(1, "P", None)]
|
||||
|
||||
def fetchall(self):
|
||||
return self._pending if isinstance(self._pending, list) else []
|
||||
|
||||
def fetchone(self):
|
||||
return self._pending if isinstance(self._pending, tuple) else None
|
||||
|
||||
def __enter__(self):
|
||||
return self
|
||||
|
||||
def __exit__(self, *a):
|
||||
return False
|
||||
|
||||
|
||||
class _Conn:
|
||||
def __init__(self, cur):
|
||||
self._cur = cur
|
||||
|
||||
def cursor(self):
|
||||
return self._cur
|
||||
|
||||
def __enter__(self):
|
||||
return self
|
||||
|
||||
def __exit__(self, *a):
|
||||
return False
|
||||
|
||||
def close(self):
|
||||
pass
|
||||
|
||||
|
||||
def test_the_query_reads_img_vector_by_name_and_sets_ef_search(monkeypatch):
|
||||
vector_store._invalidate_product_columns_cache()
|
||||
cur = _Cursor()
|
||||
monkeypatch.setattr(vector_store, "_connect", lambda: _Conn(cur))
|
||||
|
||||
rows = vector_store.image_vector_search([0.0] * 1024, top_k=30, ef_search=40)
|
||||
|
||||
select = [s for s in cur.statements if s.startswith("SELECT") and "FROM brand_x" in s and "distance" in s]
|
||||
assert select, cur.statements
|
||||
sql = select[0]
|
||||
# Light projection only: ranking must not drag 7KB product cards per candidate.
|
||||
assert sql.startswith("SELECT image_id, product_name, title, size_variants, img_vector <=> %s::vector AS distance")
|
||||
assert "WHERE img_vector IS NOT NULL" in sql and sql.endswith("ORDER BY distance ASC LIMIT %s")
|
||||
assert "embedding" not in sql and "description" not in sql
|
||||
assert cur.statements.index("SET hnsw.ef_search = 40") < cur.statements.index(sql)
|
||||
assert rows[0]["brand_table"] == "brand_x" and rows[0]["brand"] == "X" and rows[0]["distance"] == 0.2
|
||||
|
||||
|
||||
def test_hydration_reads_full_cards_by_image_id_without_the_vectors(monkeypatch):
|
||||
vector_store._invalidate_product_columns_cache()
|
||||
cur = _Cursor()
|
||||
monkeypatch.setattr(vector_store, "_connect", lambda: _Conn(cur))
|
||||
|
||||
real_fetch_products_by_image_ids("brand_x", ["a", "b"]) # the autouse fixture patches the module attribute
|
||||
|
||||
sql = [s for s in cur.statements if "WHERE image_id = ANY(%s)" in s]
|
||||
assert sql and sql[0].startswith('SELECT "id", "product_name", "img_vector_src" FROM brand_x')
|
||||
assert "embedding" not in sql[0] and '"img_vector"' not in sql[0]
|
||||
assert real_fetch_products_by_image_ids("brand_x", []) == []
|
||||
|
||||
|
||||
def test_the_queries_are_never_select_star():
|
||||
assert "SELECT *" not in inspect.getsource(vector_store.image_vector_search)
|
||||
assert "SELECT *" not in inspect.getsource(vector_store.fetch_products_by_image_ids)
|
||||
|
||||
|
||||
def test_a_category_filter_is_added_inside_the_where(monkeypatch):
|
||||
vector_store._invalidate_product_columns_cache()
|
||||
cur = _Cursor()
|
||||
monkeypatch.setattr(vector_store, "_connect", lambda: _Conn(cur))
|
||||
|
||||
vector_store.image_vector_search([0.0] * 1024, brand="X", category="Biscuits")
|
||||
|
||||
sql = [s for s in cur.statements if "AS distance" in s][0]
|
||||
assert "WHERE img_vector IS NOT NULL AND category ILIKE %s ORDER BY" in sql
|
||||
174
tests/test_image_search_api.py
Normal file
174
tests/test_image_search_api.py
Normal file
@@ -0,0 +1,174 @@
|
||||
"""POST /api/search/image-vector and POST /api/search/image - the HTTP contract.
|
||||
|
||||
The ranking is tested in tests/test_image_match.py; here `search_by_vector` is
|
||||
patched on the router module and the assertions are about status codes,
|
||||
validation and the response shape the app reads. Both routes are public.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import io
|
||||
import math
|
||||
|
||||
import pytest
|
||||
|
||||
from app.api.routers import search as search_router
|
||||
from app.services.image_match import ImageSearchResult
|
||||
|
||||
|
||||
def _unit():
|
||||
v = [math.cos(i / 7.0) for i in range(1024)]
|
||||
n = math.sqrt(sum(x * x for x in v))
|
||||
return [x / n for x in v]
|
||||
|
||||
|
||||
def _row():
|
||||
return {
|
||||
"image_id": "britannia_marie_gold_300g",
|
||||
"product_name": "Britannia Marie Gold 300g",
|
||||
"title": "Britannia Marie Gold 300g",
|
||||
"brand": "Britannia",
|
||||
"brand_table": "brand_britannia",
|
||||
"category": "Biscuits",
|
||||
"image_url": "https://cdn.example/marie.jpg",
|
||||
"image_urls": ["https://cdn.example/marie.jpg"],
|
||||
"size_variants": ["300g"],
|
||||
"barcode": "8901063010512",
|
||||
"barcode_type": "EAN13",
|
||||
"final_selling_price": 45.0,
|
||||
"selling_price": 42.0,
|
||||
"hsn_code": "1905",
|
||||
"fssai_license": "10012021000123",
|
||||
"distance": 0.369,
|
||||
"score": 0.631,
|
||||
"text_overlap": 5.0,
|
||||
}
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def fake_search(monkeypatch):
|
||||
calls = []
|
||||
|
||||
def fake(vector, text=None, brand=None, category=None, top_k=10, min_score=0.0):
|
||||
calls.append({"vector": list(vector), "text": text, "brand": brand,
|
||||
"category": category, "top_k": top_k, "min_score": min_score})
|
||||
return ImageSearchResult(rows=[_row()], detected_brand="Britannia", scoped_to_brand=True,
|
||||
min_score=min_score, query_text=text, top_k=top_k)
|
||||
|
||||
monkeypatch.setattr(search_router, "search_by_vector", fake)
|
||||
return calls
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# /search/image-vector
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def test_a_vector_returns_the_product_card_with_its_score(client, fake_search):
|
||||
res = client.post("/api/search/image-vector",
|
||||
json={"vector": _unit(), "text": "Britannia Marie Gold 300 g", "top_k": 5})
|
||||
|
||||
assert res.status_code == 200, res.text
|
||||
body = res.json()
|
||||
assert body["total"] == 1 and body["detected_brand"] == "Britannia" and body["scoped_to_brand"] is True
|
||||
hit = body["results"][0]
|
||||
assert hit["product_name"] == "Britannia Marie Gold 300g"
|
||||
assert hit["score"] == 0.631 and hit["text_overlap"] == 5.0
|
||||
for key in ("barcode", "final_selling_price", "selling_price", "category", "image_url",
|
||||
"size_variants", "hsn_code", "fssai_license", "image_id", "brand"):
|
||||
assert key in hit, key
|
||||
assert hit["barcode"] == "8901063010512" and hit["image_url"] == "https://cdn.example/marie.jpg"
|
||||
assert fake_search[0]["text"] == "Britannia Marie Gold 300 g" and fake_search[0]["top_k"] == 5
|
||||
|
||||
|
||||
def test_the_route_is_public(client, fake_search):
|
||||
assert client.post("/api/search/image-vector", json={"vector": _unit()}).status_code == 200
|
||||
|
||||
|
||||
@pytest.mark.parametrize("payload, fragment", [
|
||||
({"vector": [0.1] * 1023}, "vector"),
|
||||
({"vector": [0.1] * 1025}, "vector"),
|
||||
({"vector": [0.0] * 1024}, "all zeros"),
|
||||
({"vector": _unit(), "top_k": 51}, "top_k"),
|
||||
({"vector": _unit(), "top_k": 0}, "top_k"),
|
||||
({"vector": _unit(), "min_score": 1.5}, "min_score"),
|
||||
({}, "vector"),
|
||||
])
|
||||
def test_bad_requests_are_422_and_name_the_field(client, fake_search, payload, fragment):
|
||||
res = client.post("/api/search/image-vector", json=payload)
|
||||
assert res.status_code == 422
|
||||
assert fragment in res.text
|
||||
assert fake_search == []
|
||||
|
||||
|
||||
def test_a_nan_in_the_vector_is_422(client, fake_search):
|
||||
body = '{"vector": [' + ",".join(["NaN"] + ["0.1"] * 1023) + "]}"
|
||||
res = client.post("/api/search/image-vector", content=body, headers={"content-type": "application/json"})
|
||||
assert res.status_code == 422 and fake_search == []
|
||||
|
||||
|
||||
def test_a_service_rejection_is_422_not_500(client, monkeypatch):
|
||||
from app.services.image_match import InvalidVectorError
|
||||
|
||||
def refuse(*a, **k):
|
||||
raise InvalidVectorError("vector is all zeros")
|
||||
|
||||
monkeypatch.setattr(search_router, "search_by_vector", refuse)
|
||||
res = client.post("/api/search/image-vector", json={"vector": _unit()})
|
||||
assert res.status_code == 422 and "all zeros" in res.text
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# /search/image
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def _post_image(client, data: bytes, **form):
|
||||
return client.post("/api/search/image", files={"file": ("photo.jpg", io.BytesIO(data), "image/jpeg")}, data=form)
|
||||
|
||||
|
||||
def test_a_photo_is_embedded_and_searched_with_its_form_fields(client, fake_search, monkeypatch):
|
||||
monkeypatch.setattr(search_router.image_embedder, "available", lambda: True)
|
||||
monkeypatch.setattr(search_router.image_embedder, "embedding_for_bytes", lambda data: _unit())
|
||||
|
||||
res = _post_image(client, b"\xff\xd8" + b"x" * 5000, text="Marie Gold 300 g", top_k="3", min_score="0.2")
|
||||
|
||||
assert res.status_code == 200, res.text
|
||||
assert res.json()["results"][0]["product_name"] == "Britannia Marie Gold 300g"
|
||||
call = fake_search[0]
|
||||
assert call["text"] == "Marie Gold 300 g" and call["top_k"] == 3 and call["min_score"] == 0.2
|
||||
assert len(call["vector"]) == 1024
|
||||
|
||||
|
||||
def test_without_a_model_the_photo_route_says_503_and_points_at_the_vector_route(client, fake_search, monkeypatch):
|
||||
monkeypatch.setattr(search_router.image_embedder, "available", lambda: False)
|
||||
|
||||
res = _post_image(client, b"x" * 5000)
|
||||
|
||||
assert res.status_code == 503 and "/api/search/image-vector" in res.text
|
||||
assert fake_search == []
|
||||
|
||||
|
||||
def test_an_oversized_photo_is_413_before_any_model_work(client, fake_search, monkeypatch):
|
||||
monkeypatch.setattr(search_router, "IMAGE_VECTOR_MAX_BYTES", 100)
|
||||
monkeypatch.setattr(search_router.image_embedder, "available", lambda: pytest.fail("must not load the model"))
|
||||
|
||||
res = _post_image(client, b"x" * 101)
|
||||
|
||||
assert res.status_code == 413 and "limit is" in res.text
|
||||
|
||||
|
||||
def test_an_empty_upload_is_400(client, fake_search):
|
||||
assert _post_image(client, b"").status_code == 400
|
||||
|
||||
|
||||
def test_an_undecodable_photo_is_422(client, fake_search, monkeypatch):
|
||||
monkeypatch.setattr(search_router.image_embedder, "available", lambda: True)
|
||||
monkeypatch.setattr(search_router.image_embedder, "embedding_for_bytes", lambda data: None)
|
||||
|
||||
res = _post_image(client, b"not an image at all" * 100)
|
||||
|
||||
assert res.status_code == 422 and "decode" in res.text and fake_search == []
|
||||
|
||||
|
||||
def test_both_routes_are_documented(client):
|
||||
paths = client.get("/openapi.json").json()["paths"]
|
||||
assert "/api/search/image-vector" in paths and "/api/search/image" in paths
|
||||
assert "post" in paths["/api/search/image"] and "get" in paths["/api/search"]
|
||||
@@ -105,13 +105,57 @@ def test_a_solid_png_becomes_a_224_tensor_dominated_by_its_colour():
|
||||
assert t.max() > t.min()
|
||||
|
||||
|
||||
def test_the_default_recipe_is_raw_0_to_255():
|
||||
"""Pinned separately from the invariant tests above: this is the one
|
||||
assertion that is EXPECTED to change when the colleague's preprocessing
|
||||
replaces the default. Update it deliberately, not by accident."""
|
||||
def test_the_recipe_is_0_to_1_rgb_like_the_flutter_app():
|
||||
"""Pinned separately from the invariant tests above: the app divides by
|
||||
255 and nothing else (no mean, no std, no -1..1), RGB order."""
|
||||
t = emb.preprocess(_png(Image.new("RGB", (10, 10), (255, 128, 0))))
|
||||
|
||||
assert float(t[0, 0, 0, 0]) == 255.0 and float(t[0, 0, 0, 2]) == 0.0
|
||||
assert float(t[0, 0, 0, 0]) == 1.0
|
||||
assert abs(float(t[0, 0, 0, 1]) - 128 / 255) < 1e-6
|
||||
assert float(t[0, 0, 0, 2]) == 0.0
|
||||
|
||||
|
||||
def _reference_image_to_tensor(path: str) -> np.ndarray:
|
||||
"""The colleague's Python reference, verbatim up to the model call."""
|
||||
import cv2
|
||||
img = cv2.imread(path, cv2.IMREAD_COLOR) # 1. BGR
|
||||
h, w = img.shape[:2]
|
||||
side = min(w, h)
|
||||
x, y = (w - side) // 2, (h - side) // 2
|
||||
img = img[y:y + side, x:x + side] # 2. center crop
|
||||
img = cv2.resize(img, (224, 224), interpolation=cv2.INTER_AREA) # 3. resize
|
||||
img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB) # 4. RGB
|
||||
return (img.astype(np.float32) / 255.0)[None] # 5. 0-1, batch
|
||||
|
||||
|
||||
@pytest.mark.parametrize("size", [(640, 480), (300, 700), (224, 224), (37, 91)])
|
||||
def test_our_tensor_is_identical_to_the_colleagues_reference_code(tmp_path, size):
|
||||
"""Same bytes in, same floats out - not 'close', identical. This is the
|
||||
guarantee that a catalog vector and an app vector describe the same
|
||||
pixels; any drift here shows up as a lower cosine between the two."""
|
||||
noisy = Image.effect_noise(size, 60).convert("RGB")
|
||||
noisy.paste((200, 30, 30), (0, 0, size[0] // 3, size[1] // 2))
|
||||
data = _jpeg(noisy)
|
||||
path = tmp_path / "photo.jpg"
|
||||
path.write_bytes(data)
|
||||
|
||||
ours = emb.preprocess(data)
|
||||
theirs = _reference_image_to_tensor(str(path))
|
||||
|
||||
assert ours is not None and ours.shape == theirs.shape == (1, 224, 224, 3)
|
||||
assert np.array_equal(ours, theirs)
|
||||
|
||||
|
||||
def test_a_wide_image_is_centre_cropped_not_squashed():
|
||||
"""Left third red, middle third green, right third blue, 300x100. The app
|
||||
crops the central 100x100 before resizing, so only green survives."""
|
||||
im = Image.new("RGB", (300, 100), (255, 0, 0))
|
||||
im.paste((0, 255, 0), (100, 0, 200, 100))
|
||||
im.paste((0, 0, 255), (200, 0, 300, 100))
|
||||
|
||||
t = emb.preprocess(_png(im))
|
||||
|
||||
assert np.all(t[0, :, :, 1] == 1.0) and np.all(t[0, :, :, 0] == 0.0) and np.all(t[0, :, :, 2] == 0.0)
|
||||
|
||||
|
||||
def test_a_solid_jpeg_is_uniform_within_lossy_tolerance():
|
||||
@@ -127,10 +171,11 @@ def test_exif_orientation_is_applied_because_the_browser_applies_it():
|
||||
"""Left half red, right half blue, tagged 'rotate 90 CW to display'.
|
||||
|
||||
Without the transpose the bottom-left pixel is red (the left half). With
|
||||
it, the left half has become the top half and bottom-left is blue.
|
||||
it, the left half has become the top half and bottom-left is blue. The
|
||||
image is 64x64 so the centre crop keeps all of it.
|
||||
"""
|
||||
im = Image.new("RGB", (64, 32), (255, 0, 0))
|
||||
im.paste((0, 0, 255), (32, 0, 64, 32))
|
||||
im = Image.new("RGB", (64, 64), (255, 0, 0))
|
||||
im.paste((0, 0, 255), (32, 0, 64, 64))
|
||||
exif = Image.Exif()
|
||||
exif[0x0112] = 6
|
||||
t = emb.preprocess(_jpeg(im, exif=exif.tobytes()))
|
||||
@@ -285,6 +330,41 @@ def test_inference_is_serialised_on_the_module_lock(monkeypatch):
|
||||
assert not overlap and len(fake.inputs) == 6
|
||||
|
||||
|
||||
def test_status_reports_a_missing_model_without_loading_it(monkeypatch, tmp_path):
|
||||
"""The deploy-time failure: code shipped, .tflite did not. /api/health
|
||||
must say so, and asking must not itself trigger a load attempt."""
|
||||
monkeypatch.setattr(emb, "IMAGE_EMBED_MODEL_PATH", tmp_path / "missing.tflite")
|
||||
loads = []
|
||||
monkeypatch.setattr(emb, "_load", lambda: loads.append(1))
|
||||
|
||||
st = emb.status()
|
||||
|
||||
assert st["model_present"] is False and st["model_path"].endswith("missing.tflite")
|
||||
assert st["state"].startswith("not loaded yet")
|
||||
assert loads == []
|
||||
|
||||
|
||||
def test_status_reflects_a_recorded_failure_and_a_ready_interpreter(monkeypatch):
|
||||
monkeypatch.setattr(emb, "_disabled_reason", "model file not found at x")
|
||||
assert emb.status()["state"] == "disabled: model file not found at x"
|
||||
|
||||
emb._reset()
|
||||
_install_fake(monkeypatch)
|
||||
assert emb.status()["state"] == "ready"
|
||||
|
||||
|
||||
def test_health_carries_the_image_vector_block(monkeypatch, tmp_path):
|
||||
from fastapi.testclient import TestClient
|
||||
from app.main import app
|
||||
monkeypatch.setattr(emb, "IMAGE_EMBED_MODEL_PATH", tmp_path / "missing.tflite")
|
||||
|
||||
body = TestClient(app).get("/api/health").json()
|
||||
|
||||
iv_block = body["image_vectors"]
|
||||
assert iv_block["model_present"] is False
|
||||
assert set(iv_block) == {"enabled", "model_path", "model_present", "runtime_importable", "state"}
|
||||
|
||||
|
||||
def test_to_pg_is_the_same_text_form_the_upsert_uses_for_embedding():
|
||||
assert iv.to_pg([0, 0.5, 1]) == "[0.0,0.5,1.0]"
|
||||
|
||||
|
||||
Reference in New Issue
Block a user