Image vector to product details
This commit is contained in:
@@ -91,6 +91,11 @@ os.environ["AUTO_ENRICH_ON_UPLOAD"] = "false"
|
||||
# the worker are covered explicitly in tests/test_image_vector.py.
|
||||
os.environ["ENABLE_IMAGE_VECTORS"] = "false"
|
||||
|
||||
# And for server-side OCR: no test may import onnxruntime or load the 30MB
|
||||
# PP-OCR models. tests/test_ocr_service.py drives the service with a fake
|
||||
# engine and flips the flag on itself.
|
||||
os.environ["ENABLE_SERVER_OCR"] = "false"
|
||||
|
||||
# Auth is set unconditionally (not setdefault): the suite asserts on the real
|
||||
# guards, so it must never inherit a developer's AUTH_ENABLED=false.
|
||||
os.environ["AUTH_ENABLED"] = "true"
|
||||
|
||||
294
tests/test_identify_api.py
Normal file
294
tests/test_identify_api.py
Normal file
@@ -0,0 +1,294 @@
|
||||
"""POST /api/search/identify, and `text_fallback` on POST /api/search/image-vector.
|
||||
|
||||
The ladder's decisions are tested in tests/test_product_identify.py and the
|
||||
two arms in tests/test_image_match.py and tests/test_label_match.py. Here
|
||||
the assertions are about the HTTP contract: status codes, the guards the
|
||||
photo route shares with /search/image, the IdentifyOut fields, and - the
|
||||
one that protects the app in the field - that /search/image-vector without
|
||||
`text_fallback` answers exactly as it did before this route existed.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import io
|
||||
import math
|
||||
from typing import Any, Dict, List
|
||||
|
||||
import pytest
|
||||
|
||||
from app.api.routers import search as search_router
|
||||
from app.services import product_identify
|
||||
from app.services.image_match import ImageSearchResult
|
||||
from app.services.product_identify import IdentifyResult
|
||||
|
||||
|
||||
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(score: float = 0.631, image_id: str = "britannia_marie_gold_300g") -> Dict[str, Any]:
|
||||
return {
|
||||
"image_id": image_id,
|
||||
"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"],
|
||||
"score": score,
|
||||
"text_overlap": 5.0,
|
||||
}
|
||||
|
||||
|
||||
IMAGE_SEARCH_KEYS = {"results", "total", "detected_brand", "scoped_to_brand", "scope_fallback",
|
||||
"min_score", "top_k", "query_text"}
|
||||
IDENTIFY_KEYS = IMAGE_SEARCH_KEYS | {"matched_by", "ocr_text", "ocr_source", "image_top_score", "fallback_reason"}
|
||||
|
||||
|
||||
def _post_identify(client, data: bytes, **form):
|
||||
return client.post("/api/search/identify",
|
||||
files={"file": ("photo.jpg", io.BytesIO(data), "image/jpeg")}, data=form)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def embedder_ok(monkeypatch):
|
||||
monkeypatch.setattr(search_router.image_embedder, "available", lambda: True)
|
||||
monkeypatch.setattr(search_router.image_embedder, "embedding_for_bytes", lambda data: _unit())
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def embedder_missing(monkeypatch):
|
||||
monkeypatch.setattr(search_router.image_embedder, "available", lambda: False)
|
||||
monkeypatch.setattr(search_router.image_embedder, "embedding_for_bytes",
|
||||
lambda data: pytest.fail("must not embed without a model"))
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def fake_identify(monkeypatch):
|
||||
"""Patch the ladder on the router: the contract tests."""
|
||||
calls: List[Dict[str, Any]] = []
|
||||
state = {"result": None}
|
||||
|
||||
def fake(**kw):
|
||||
calls.append(kw)
|
||||
return state["result"]
|
||||
|
||||
def set_result(**fields):
|
||||
fields.setdefault("search", ImageSearchResult(rows=[_row()], detected_brand="Britannia",
|
||||
scoped_to_brand=True, top_k=10, query_text="Marie Gold"))
|
||||
state["result"] = IdentifyResult(**fields)
|
||||
|
||||
set_result(matched_by="image_vector", image_top_score=0.912345)
|
||||
monkeypatch.setattr(search_router, "identify_product", fake)
|
||||
fake.calls = calls
|
||||
fake.set_result = set_result
|
||||
return fake
|
||||
|
||||
|
||||
class _Ladder:
|
||||
"""Patch UNDER the real ladder: the arms and the OCR engine."""
|
||||
|
||||
def __init__(self, monkeypatch):
|
||||
self.image_rows = [_row(0.63)]
|
||||
self.text_rows = [_row(0.55, image_id="txt_marie_300")]
|
||||
self.ocr_available = True
|
||||
self.ocr_text = "BRITANNIA MARIE GOLD 300 g"
|
||||
self.ocr_calls: List[bytes] = []
|
||||
self.text_calls: List[Dict[str, Any]] = []
|
||||
monkeypatch.setattr(product_identify, "search_by_vector", self._image)
|
||||
monkeypatch.setattr(product_identify, "resolve_label", self._text)
|
||||
monkeypatch.setattr(product_identify.ocr_service, "available", lambda: self.ocr_available)
|
||||
monkeypatch.setattr(product_identify.ocr_service, "read_text", self._ocr)
|
||||
monkeypatch.setattr(product_identify, "IMAGE_IDENTIFY_MIN_IMAGE_SCORE", 0.70)
|
||||
|
||||
def _image(self, vector, text=None, brand=None, category=None, top_k=10, min_score=0.0):
|
||||
return ImageSearchResult(rows=[dict(r) for r in self.image_rows], detected_brand=brand,
|
||||
min_score=min_score, top_k=top_k, query_text=text)
|
||||
|
||||
def _text(self, text, brand=None, category=None, top_k=10):
|
||||
self.text_calls.append({"text": text, "brand": brand, "category": category, "top_k": top_k})
|
||||
return ImageSearchResult(rows=[dict(r) for r in self.text_rows], query_text=text, top_k=top_k)
|
||||
|
||||
def _ocr(self, data):
|
||||
self.ocr_calls.append(data)
|
||||
return self.ocr_text
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def ladder(monkeypatch):
|
||||
return _Ladder(monkeypatch)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# /search/identify - the contract
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def test_a_confident_image_match_returns_the_card_with_the_identify_fields(client, embedder_ok, fake_identify):
|
||||
res = _post_identify(client, b"\xff\xd8" + b"x" * 5000, top_k="3", min_score="0.2", brand="Britannia")
|
||||
|
||||
assert res.status_code == 200, res.text
|
||||
body = res.json()
|
||||
assert set(body) == IDENTIFY_KEYS
|
||||
assert body["matched_by"] == "image_vector" and body["fallback_reason"] is None
|
||||
assert body["image_top_score"] == 0.9123
|
||||
assert body["results"][0]["product_name"] == "Britannia Marie Gold 300g"
|
||||
assert body["results"][0]["score"] == 0.631
|
||||
call = fake_identify.calls[0]
|
||||
assert len(call["vector"]) == 1024 and call["image_bytes"].startswith(b"\xff\xd8")
|
||||
assert call["top_k"] == 3 and call["min_score"] == 0.2 and call["brand"] == "Britannia"
|
||||
|
||||
|
||||
def test_the_route_is_public(client, embedder_ok, fake_identify):
|
||||
assert _post_identify(client, b"x" * 5000).status_code == 200
|
||||
|
||||
|
||||
def test_the_photo_guards_match_search_image(client, fake_identify, monkeypatch):
|
||||
assert _post_identify(client, b"").status_code == 400
|
||||
|
||||
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_identify(client, b"x" * 101)
|
||||
assert res.status_code == 413 and "limit is" in res.text
|
||||
assert fake_identify.calls == []
|
||||
|
||||
|
||||
def test_an_undecodable_photo_is_422_when_the_model_is_present(client, fake_identify, monkeypatch):
|
||||
monkeypatch.setattr(search_router.image_embedder, "available", lambda: True)
|
||||
monkeypatch.setattr(search_router.image_embedder, "embedding_for_bytes", lambda data: None)
|
||||
|
||||
res = _post_identify(client, b"not an image" * 100)
|
||||
|
||||
assert res.status_code == 422 and "decode" in res.text and fake_identify.calls == []
|
||||
|
||||
|
||||
def test_without_a_model_the_ladder_still_runs_with_no_vector(client, embedder_missing, fake_identify):
|
||||
fake_identify.set_result(matched_by="text", ocr_text="Marie Gold", ocr_source="client",
|
||||
fallback_reason="image_embedder_unavailable")
|
||||
|
||||
res = _post_identify(client, b"x" * 5000, text="Marie Gold")
|
||||
|
||||
assert res.status_code == 200, res.text
|
||||
assert res.json()["matched_by"] == "text"
|
||||
assert res.json()["fallback_reason"] == "image_embedder_unavailable"
|
||||
assert fake_identify.calls[0]["vector"] is None and fake_identify.calls[0]["text"] == "Marie Gold"
|
||||
|
||||
|
||||
def test_no_model_and_no_ocr_and_no_text_is_503(client, embedder_missing, fake_identify):
|
||||
fake_identify.set_result(search=ImageSearchResult(), matched_by="none", fallback_reason="ocr_unavailable")
|
||||
|
||||
res = _post_identify(client, b"x" * 5000)
|
||||
|
||||
assert res.status_code == 503
|
||||
assert "text" in res.text and "/api/search/image-vector" in res.text
|
||||
|
||||
|
||||
def test_ocr_unavailable_with_a_model_is_200_with_the_reason(client, embedder_ok, fake_identify):
|
||||
fake_identify.set_result(matched_by="image_vector", image_top_score=0.63, fallback_reason="ocr_unavailable")
|
||||
|
||||
res = _post_identify(client, b"x" * 5000)
|
||||
|
||||
assert res.status_code == 200
|
||||
assert res.json()["fallback_reason"] == "ocr_unavailable" and res.json()["matched_by"] == "image_vector"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# /search/identify - through the real ladder
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def test_a_low_score_with_client_text_resolves_by_text(client, embedder_ok, ladder):
|
||||
res = _post_identify(client, b"x" * 5000, text="Marie Gold 300 g")
|
||||
|
||||
body = res.json()
|
||||
assert res.status_code == 200, res.text
|
||||
assert body["matched_by"] == "text" and body["fallback_reason"] == "image_below_threshold"
|
||||
assert body["ocr_source"] == "client" and body["ocr_text"] == "Marie Gold 300 g"
|
||||
assert body["image_top_score"] == 0.63
|
||||
assert body["results"][0]["image_id"] == "txt_marie_300"
|
||||
assert ladder.ocr_calls == []
|
||||
|
||||
|
||||
def test_a_low_score_without_text_uses_server_ocr(client, embedder_ok, ladder):
|
||||
res = _post_identify(client, b"\xff\xd8" + b"x" * 5000)
|
||||
|
||||
body = res.json()
|
||||
assert res.status_code == 200, res.text
|
||||
assert body["matched_by"] == "text" and body["ocr_source"] == "server"
|
||||
assert body["ocr_text"] == "BRITANNIA MARIE GOLD 300 g"
|
||||
assert ladder.ocr_calls and ladder.ocr_calls[0].startswith(b"\xff\xd8")
|
||||
assert ladder.text_calls[0]["text"] == "BRITANNIA MARIE GOLD 300 g"
|
||||
|
||||
|
||||
def test_ocr_unavailable_is_reported_not_500(client, embedder_ok, ladder):
|
||||
ladder.ocr_available = False
|
||||
|
||||
res = _post_identify(client, b"x" * 5000)
|
||||
|
||||
body = res.json()
|
||||
assert res.status_code == 200, res.text
|
||||
assert body["matched_by"] == "image_vector" and body["fallback_reason"] == "ocr_unavailable"
|
||||
assert body["results"][0]["score"] == 0.63 and body["image_top_score"] == 0.63
|
||||
|
||||
|
||||
def test_a_confident_image_never_touches_ocr(client, embedder_ok, ladder):
|
||||
ladder.image_rows = [_row(0.88)]
|
||||
|
||||
res = _post_identify(client, b"x" * 5000)
|
||||
|
||||
assert res.json()["matched_by"] == "image_vector" and res.json()["fallback_reason"] is None
|
||||
assert ladder.ocr_calls == [] and ladder.text_calls == []
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# /search/image-vector - text_fallback
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def test_image_vector_without_the_flag_is_unchanged(client, ladder, fake_identify, monkeypatch):
|
||||
seen = []
|
||||
|
||||
def fake_search(vector, text=None, brand=None, category=None, top_k=10, min_score=0.0):
|
||||
seen.append(text)
|
||||
return ImageSearchResult(rows=[_row(0.63)], 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_search)
|
||||
|
||||
res = client.post("/api/search/image-vector", json={"vector": _unit(), "text": "Marie Gold 300 g"})
|
||||
|
||||
assert res.status_code == 200, res.text
|
||||
assert set(res.json()) == IMAGE_SEARCH_KEYS
|
||||
assert seen == ["Marie Gold 300 g"] and fake_identify.calls == [] and ladder.text_calls == []
|
||||
|
||||
|
||||
def test_image_vector_with_the_flag_runs_the_ladder_and_adds_the_fields(client, ladder):
|
||||
res = client.post("/api/search/image-vector",
|
||||
json={"vector": _unit(), "text": "Marie Gold 300 g", "text_fallback": True, "top_k": 4})
|
||||
|
||||
body = res.json()
|
||||
assert res.status_code == 200, res.text
|
||||
assert set(body) == IDENTIFY_KEYS
|
||||
assert body["matched_by"] == "text" and body["ocr_source"] == "client"
|
||||
assert body["fallback_reason"] == "image_below_threshold"
|
||||
assert body["results"][0]["image_id"] == "txt_marie_300" and body["top_k"] == 4
|
||||
assert ladder.ocr_calls == [] # no photo on this route: never server OCR
|
||||
|
||||
|
||||
def test_image_vector_with_the_flag_but_no_text_says_no_text(client, ladder):
|
||||
res = client.post("/api/search/image-vector", json={"vector": _unit(), "text_fallback": True})
|
||||
|
||||
body = res.json()
|
||||
assert res.status_code == 200, res.text
|
||||
assert body["matched_by"] == "image_vector" and body["fallback_reason"] == "no_text"
|
||||
assert body["results"][0]["score"] == 0.63
|
||||
|
||||
|
||||
def test_image_vector_with_the_flag_still_validates_the_vector(client, ladder):
|
||||
res = client.post("/api/search/image-vector", json={"vector": [0.0] * 1024, "text_fallback": True})
|
||||
assert res.status_code == 422
|
||||
|
||||
|
||||
def test_the_get_route_has_no_fallback_flag(client):
|
||||
params = client.get("/openapi.json").json()["paths"]["/api/search/image-vector"]["get"]["parameters"]
|
||||
assert "text_fallback" not in {p["name"] for p in params}
|
||||
@@ -224,3 +224,4 @@ def test_both_routes_are_documented(client):
|
||||
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"]
|
||||
assert {"get", "post"} <= set(paths["/api/search/image-vector"])
|
||||
assert "post" in paths["/api/search/identify"]
|
||||
|
||||
330
tests/test_label_match.py
Normal file
330
tests/test_label_match.py
Normal file
@@ -0,0 +1,330 @@
|
||||
"""Label text -> catalog rows (app/services/label_match.py).
|
||||
|
||||
The rung of /search/identify that runs when the photo's vector cannot
|
||||
confirm a product. What has to hold:
|
||||
|
||||
1. `clean_label` keeps the pack size and the product words, and drops the
|
||||
nutrition panel, the "per 100 g" basis (which would otherwise score as a
|
||||
pack size), URLs, prices and repeats.
|
||||
2. The brand the label names scopes both arms and prefixes the embedded
|
||||
query; an explicit brand never falls back; a detected one that finds
|
||||
nothing confident is retried across every brand and says so.
|
||||
3. Rows from both arms are merged and deduped, given the display brand and
|
||||
the table the image rows carry, and ranked by size/word overlap FIRST and
|
||||
cosine second - so the 300g sibling wins when the label says 300 g.
|
||||
4. No model means lexical-only with zero scores; no store means an empty
|
||||
result; nothing raises.
|
||||
|
||||
Store and model are patched at their source modules; no database anywhere.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any, Dict, List
|
||||
|
||||
import pytest
|
||||
|
||||
from app.services import embeddings_service, label_match, query_intent, vector_store
|
||||
|
||||
|
||||
def _row(image_id: str, name: str, size: str, distance: float, brand: str = "britannia") -> Dict[str, Any]:
|
||||
return {
|
||||
"image_id": image_id, "product_name": name, "title": name.replace("Britannia ", ""),
|
||||
"size_variants": [size], "distance": distance, "brand": brand, "category": "Biscuits",
|
||||
}
|
||||
|
||||
|
||||
MARIE_89 = _row("marie_89", "Britannia Marie Gold 89g", "89g", 0.30)
|
||||
MARIE_300 = _row("marie_300", "Britannia Marie Gold 300g", "300g", 0.31)
|
||||
MARIE_1KG = _row("marie_1kg", "Britannia Marie Gold 1kg", "1kg", 0.32)
|
||||
GOOD_DAY = _row("gd_250", "Britannia Good Day Cashew 250g", "250g", 0.45)
|
||||
AMUL_BUTTER = _row("amul_b", "Amul Butter 100g", "100g", 0.70, brand="amul")
|
||||
|
||||
|
||||
class _Store:
|
||||
"""Records every call to the two arms; answers with canned rows."""
|
||||
|
||||
def __init__(self, semantic=None, lexical=None):
|
||||
self.semantic_rows = semantic if semantic is not None else []
|
||||
self.lexical_rows = lexical if lexical is not None else []
|
||||
self.semantic_calls: List[Dict[str, Any]] = []
|
||||
self.lexical_calls: List[Dict[str, Any]] = []
|
||||
|
||||
def semantic_search(self, query_embedding, brand=None, top_k=5, category=None, **kw):
|
||||
self.semantic_calls.append({"brand": brand, "top_k": top_k, "category": category})
|
||||
rows = self.semantic_rows(brand) if callable(self.semantic_rows) else self.semantic_rows
|
||||
return [dict(r) for r in rows]
|
||||
|
||||
def lexical_search(self, terms, brand=None, limit=200, category=None, query_embedding=None, exact_phrase=None, **kw):
|
||||
self.lexical_calls.append({"terms": list(terms), "brand": brand, "limit": limit,
|
||||
"category": category, "has_embedding": query_embedding is not None})
|
||||
rows = self.lexical_rows(brand, terms) if callable(self.lexical_rows) else self.lexical_rows
|
||||
return [dict(r) for r in rows]
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def store(monkeypatch):
|
||||
st = _Store()
|
||||
monkeypatch.setattr(vector_store, "semantic_search", st.semantic_search)
|
||||
monkeypatch.setattr(vector_store, "lexical_search", st.lexical_search)
|
||||
return st
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def embedder(monkeypatch):
|
||||
calls: List[List[str]] = []
|
||||
|
||||
def fake_embed(texts):
|
||||
calls.append(list(texts))
|
||||
return [[0.1] * 384 for _ in texts]
|
||||
|
||||
monkeypatch.setattr(embeddings_service, "embed_texts", fake_embed)
|
||||
return calls
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def no_brand_detection(monkeypatch):
|
||||
monkeypatch.setattr(query_intent, "extract_brand_mention", lambda text: None)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 1. clean_label
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def test_clean_label_drops_the_nutrition_panel_and_per_100g_but_keeps_the_pack_size():
|
||||
label = ("Britannia MARIE GOLD Biscuits NET WT 300 g NUTRITION INFORMATION per 100 g "
|
||||
"Energy 480 kcal Protein 7 g Carbohydrate 76 g Sugar 20 g MRP Rs. 45 incl. of all taxes")
|
||||
|
||||
cleaned = label_match.clean_label(label)
|
||||
words, sizes = label_match.tokens(cleaned)
|
||||
|
||||
assert sizes == {"300g"}, cleaned
|
||||
assert {"britannia", "marie", "gold", "biscuits"} <= words
|
||||
assert not ({"energy", "kcal", "protein", "nutrition", "mrp", "taxes"} & words)
|
||||
assert "7 g" not in cleaned and "20 g" not in cleaned
|
||||
|
||||
|
||||
def test_clean_label_drops_urls_prices_and_repeats_keeping_first_order():
|
||||
label = "Marie Gold www.britannia.co.in Marie Gold ₹45 300 g 300 g customer care 1800-xx"
|
||||
|
||||
cleaned = label_match.clean_label(label)
|
||||
|
||||
assert cleaned == "Marie Gold 300 g 1800-xx"
|
||||
|
||||
|
||||
def test_clean_label_drops_licence_phone_and_barcode_numbers_but_not_short_ones():
|
||||
label = "FSSAI Lic. No. 10012021000123 Marie Gold 20 pack 8901063010512 call 1800123456 300 g"
|
||||
|
||||
assert label_match.clean_label(label) == "Marie Gold 20 300 g"
|
||||
|
||||
|
||||
def test_clean_label_is_capped_on_a_word_boundary(monkeypatch):
|
||||
monkeypatch.setattr(label_match, "OCR_MAX_CHARS", 12)
|
||||
assert label_match.clean_label("Britannia Marie Gold Biscuits") == "Britannia"
|
||||
|
||||
|
||||
@pytest.mark.parametrize("value", [None, "", " ", "per 100 g MRP ₹45"])
|
||||
def test_an_empty_or_all_noise_label_gives_an_empty_result_without_touching_the_store(store, embedder, value):
|
||||
result = label_match.resolve_label(value)
|
||||
|
||||
assert result.rows == [] and result.query_text is None
|
||||
assert store.semantic_calls == [] and store.lexical_calls == [] and embedder == []
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 2. brand scoping
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def test_the_brand_the_label_names_scopes_both_arms_and_prefixes_the_query(store, embedder, monkeypatch):
|
||||
monkeypatch.setattr(query_intent, "extract_brand_mention", lambda text: "britannia")
|
||||
store.semantic_rows = [MARIE_300]
|
||||
|
||||
result = label_match.resolve_label("MARIE GOLD 300 g")
|
||||
|
||||
assert embedder == [["britannia MARIE GOLD 300 g"]]
|
||||
assert store.semantic_calls[0]["brand"] == "britannia"
|
||||
assert store.lexical_calls[0]["brand"] == "britannia"
|
||||
assert result.detected_brand == "britannia" and result.scoped_to_brand and not result.scope_fallback
|
||||
|
||||
|
||||
def test_a_brand_already_in_the_label_is_not_prefixed_twice(store, embedder, monkeypatch):
|
||||
monkeypatch.setattr(query_intent, "extract_brand_mention", lambda text: "britannia")
|
||||
store.semantic_rows = [MARIE_300]
|
||||
|
||||
label_match.resolve_label("Britannia Marie Gold 300 g")
|
||||
|
||||
assert embedder == [["Britannia Marie Gold 300 g"]]
|
||||
|
||||
|
||||
def test_a_detected_category_is_never_passed_as_a_filter(store, embedder, no_brand_detection):
|
||||
store.semantic_rows = [MARIE_300]
|
||||
label_match.resolve_label("Marie Gold taste of India 300 g")
|
||||
assert all(c["category"] is None for c in store.semantic_calls + store.lexical_calls)
|
||||
|
||||
|
||||
def test_an_explicit_category_is_passed_through(store, embedder, no_brand_detection):
|
||||
store.semantic_rows = [MARIE_300]
|
||||
label_match.resolve_label("Marie Gold 300 g", category="Biscuits")
|
||||
assert store.semantic_calls[0]["category"] == "Biscuits"
|
||||
|
||||
|
||||
def test_a_detected_brand_that_finds_nothing_confident_is_retried_across_every_brand(store, embedder, monkeypatch):
|
||||
monkeypatch.setattr(query_intent, "extract_brand_mention", lambda text: "amul")
|
||||
store.semantic_rows = lambda brand: [] if brand == "amul" else [MARIE_300]
|
||||
|
||||
result = label_match.resolve_label("Marie Gold 300 g")
|
||||
|
||||
assert [c["brand"] for c in store.semantic_calls] == ["amul", None]
|
||||
assert result.scope_fallback is True and result.scoped_to_brand is False
|
||||
assert [r["image_id"] for r in result.rows] == ["marie_300"]
|
||||
|
||||
|
||||
def test_a_scoped_stranger_is_kept_when_the_wider_search_is_no_better(store, embedder, monkeypatch):
|
||||
monkeypatch.setattr(query_intent, "extract_brand_mention", lambda text: "britannia")
|
||||
monkeypatch.setattr(label_match, "IMAGE_IDENTIFY_MIN_TEXT_SCORE", 0.9)
|
||||
store.semantic_rows = lambda brand: [GOOD_DAY] if brand == "britannia" else [AMUL_BUTTER]
|
||||
|
||||
result = label_match.resolve_label("Zzyzx 500 g")
|
||||
|
||||
assert [r["image_id"] for r in result.rows] == ["gd_250"]
|
||||
assert result.scoped_to_brand is True and result.scope_fallback is False
|
||||
|
||||
|
||||
def test_an_explicit_brand_never_falls_back(store, embedder):
|
||||
store.semantic_rows = lambda brand: [] if brand == "Britannia" else [MARIE_300]
|
||||
|
||||
result = label_match.resolve_label("Zzyzx 500 g", brand="Britannia")
|
||||
|
||||
assert [c["brand"] for c in store.semantic_calls] == ["Britannia"]
|
||||
assert result.rows == [] and result.detected_brand == "Britannia" and not result.scope_fallback
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 3. merging and ranking
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def test_the_300g_sibling_wins_on_size_overlap_before_cosine(store, embedder, no_brand_detection):
|
||||
# 89g has the best cosine; the label says 300 g.
|
||||
store.semantic_rows = [MARIE_89, MARIE_300, MARIE_1KG]
|
||||
|
||||
result = label_match.resolve_label("Britannia Marie Gold 300 g", top_k=3)
|
||||
|
||||
assert [r["image_id"] for r in result.rows] == ["marie_300", "marie_89", "marie_1kg"]
|
||||
assert result.rows[0]["text_overlap"] > result.rows[1]["text_overlap"]
|
||||
assert result.rows[0]["score"] == pytest.approx(1 - 0.31)
|
||||
|
||||
|
||||
def test_the_plain_product_beats_its_variants_when_the_label_names_no_variant(store, embedder, no_brand_detection):
|
||||
# Every row explains the whole label ("Dairy Milk 50 g"); the variants
|
||||
# add words the label never said, and the plain row has the WORST cosine.
|
||||
plain = _row("dm_50", "Cadbury Dairy Milk 50g", "50g", 0.40, brand="cadbury")
|
||||
fruit_nut = _row("dm_fn_50", "Cadbury Dairy Milk Fruit & Nut 50g", "50g", 0.30, brand="cadbury")
|
||||
silk = _row("dm_silk_50", "Cadbury Dairy Milk Silk 50g", "50g", 0.35, brand="cadbury")
|
||||
store.semantic_rows = [fruit_nut, silk, plain]
|
||||
|
||||
result = label_match.resolve_label("Cadbury Dairy Milk 50 g", top_k=3)
|
||||
|
||||
assert [r["image_id"] for r in result.rows] == ["dm_50", "dm_silk_50", "dm_fn_50"]
|
||||
assert [r["text_extra"] for r in result.rows] == [0, 1, 2]
|
||||
|
||||
|
||||
def test_a_variant_named_on_the_label_still_wins(store, embedder, no_brand_detection):
|
||||
plain = _row("dm_50", "Cadbury Dairy Milk 50g", "50g", 0.30, brand="cadbury")
|
||||
silk = _row("dm_silk_50", "Cadbury Dairy Milk Silk 50g", "50g", 0.40, brand="cadbury")
|
||||
store.semantic_rows = [plain, silk]
|
||||
|
||||
result = label_match.resolve_label("Cadbury Dairy Milk Silk 50 g", top_k=2)
|
||||
|
||||
assert [r["image_id"] for r in result.rows] == ["dm_silk_50", "dm_50"]
|
||||
|
||||
|
||||
def test_lexical_and_semantic_hits_are_merged_and_deduped_semantic_row_first(store, embedder, no_brand_detection):
|
||||
store.semantic_rows = [MARIE_300]
|
||||
store.lexical_rows = [dict(MARIE_300, distance=None), GOOD_DAY]
|
||||
|
||||
result = label_match.resolve_label("Marie Gold 300 g", top_k=5)
|
||||
|
||||
ids = [r["image_id"] for r in result.rows]
|
||||
assert ids == ["marie_300", "gd_250"]
|
||||
assert result.rows[0]["score"] == pytest.approx(1 - 0.31) # the semantic row's distance survived
|
||||
|
||||
|
||||
def test_the_lexical_arm_gets_the_two_longest_non_brand_words_and_retries_with_one(store, embedder, monkeypatch):
|
||||
monkeypatch.setattr(query_intent, "extract_brand_mention", lambda text: "britannia")
|
||||
store.lexical_rows = lambda brand, terms: [MARIE_300] if len(terms) == 1 else []
|
||||
|
||||
label_match.resolve_label("Britannia Marie Gold Biscuits 300 g")
|
||||
|
||||
assert [c["terms"] for c in store.lexical_calls] == [["biscuits", "marie"], ["biscuits"]]
|
||||
assert store.lexical_calls[0]["has_embedding"] is True
|
||||
|
||||
|
||||
def test_unscoped_rows_get_the_display_brand_and_the_brand_table(store, embedder, no_brand_detection):
|
||||
store.semantic_rows = [dict(AMUL_BUTTER, brand="hindustan_unilever")]
|
||||
|
||||
result = label_match.resolve_label("Butter 100 g")
|
||||
|
||||
row = result.rows[0]
|
||||
assert row["brand_table"] == "brand_hindustan_unilever"
|
||||
assert row["brand"] == "Hindustan Unilever"
|
||||
|
||||
|
||||
def test_scoped_rows_get_the_table_of_the_scoping_brand(store, embedder):
|
||||
store.semantic_rows = [MARIE_300]
|
||||
|
||||
result = label_match.resolve_label("Marie Gold 300 g", brand="Britannia")
|
||||
|
||||
assert result.rows[0]["brand_table"] == "brand_britannia"
|
||||
|
||||
|
||||
def test_top_k_is_clamped_and_fetch_k_is_wider_than_top_k(store, embedder, no_brand_detection):
|
||||
store.semantic_rows = [MARIE_89, MARIE_300, MARIE_1KG]
|
||||
|
||||
result = label_match.resolve_label("Marie Gold", top_k=2)
|
||||
|
||||
assert len(result.rows) == 2 and result.top_k == 2
|
||||
assert store.semantic_calls[0]["top_k"] >= 30
|
||||
assert label_match.resolve_label("x y", top_k=10_000).top_k == label_match.IMAGE_SEARCH_MAX_TOP_K
|
||||
|
||||
|
||||
def test_is_confident_needs_an_overlap_or_the_cosine_floor(monkeypatch):
|
||||
monkeypatch.setattr(label_match, "IMAGE_IDENTIFY_MIN_TEXT_SCORE", 0.6)
|
||||
assert label_match.is_confident([]) is False
|
||||
assert label_match.is_confident([{"text_overlap": 0.0, "score": 0.59}]) is False
|
||||
assert label_match.is_confident([{"text_overlap": 0.0, "score": 0.60}]) is True
|
||||
assert label_match.is_confident([{"text_overlap": 1.0, "score": 0.10}]) is True
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 4. degradation
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def test_an_embedding_failure_degrades_to_lexical_only_with_zero_scores(store, monkeypatch, no_brand_detection):
|
||||
def boom(texts):
|
||||
raise RuntimeError("torch not installed")
|
||||
|
||||
monkeypatch.setattr(embeddings_service, "embed_texts", boom)
|
||||
store.lexical_rows = [dict(MARIE_300, distance=None)]
|
||||
|
||||
result = label_match.resolve_label("Marie Gold 300 g")
|
||||
|
||||
assert store.semantic_calls == []
|
||||
assert store.lexical_calls and store.lexical_calls[0]["has_embedding"] is False
|
||||
assert [r["image_id"] for r in result.rows] == ["marie_300"]
|
||||
assert result.rows[0]["score"] == 0.0 and result.rows[0]["text_overlap"] >= 4.0
|
||||
|
||||
|
||||
def test_a_store_that_raises_yields_an_empty_result_not_an_exception(monkeypatch, embedder, no_brand_detection):
|
||||
def boom(*a, **k):
|
||||
raise RuntimeError("no database")
|
||||
|
||||
monkeypatch.setattr(vector_store, "semantic_search", boom)
|
||||
monkeypatch.setattr(vector_store, "lexical_search", boom)
|
||||
|
||||
result = label_match.resolve_label("Marie Gold 300 g")
|
||||
|
||||
assert result.rows == [] and result.query_text == "Marie Gold 300 g"
|
||||
|
||||
|
||||
def test_a_row_without_a_distance_scores_zero(store, embedder, no_brand_detection):
|
||||
store.semantic_rows = [dict(MARIE_300, distance=None)]
|
||||
assert label_match.resolve_label("Marie Gold 300 g").rows[0]["score"] == 0.0
|
||||
62
tests/test_ocr_engine_real.py
Normal file
62
tests/test_ocr_engine_real.py
Normal file
@@ -0,0 +1,62 @@
|
||||
"""The real OCR engine, when it is installed (requirements-ocr.txt).
|
||||
|
||||
Skipped cleanly otherwise, like tests/test_image_embedder_model.py: the unit
|
||||
tests in tests/test_ocr_service.py drive a fake and never need the wheel.
|
||||
This file proves the wheel that IS installed constructs with our params,
|
||||
finds its bundled models without a download, and reads printed text.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import io
|
||||
|
||||
import pytest
|
||||
|
||||
pytest.importorskip("rapidocr")
|
||||
pytest.importorskip("onnxruntime")
|
||||
|
||||
from PIL import Image, ImageDraw, ImageFont # noqa: E402
|
||||
|
||||
from app.services import ocr_service as ocr # noqa: E402
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _real_engine(monkeypatch):
|
||||
ocr._reset()
|
||||
monkeypatch.setattr(ocr, "ENABLE_SERVER_OCR", True)
|
||||
yield
|
||||
ocr._reset()
|
||||
|
||||
|
||||
def _label_png(lines, size=(640, 240)) -> bytes:
|
||||
im = Image.new("RGB", size, "white")
|
||||
draw = ImageDraw.Draw(im)
|
||||
try:
|
||||
font = ImageFont.truetype("arial.ttf", 56)
|
||||
except OSError:
|
||||
font = ImageFont.load_default(size=56)
|
||||
y = 30
|
||||
for line in lines:
|
||||
draw.text((30, y), line, fill="black", font=font)
|
||||
y += 90
|
||||
buf = io.BytesIO()
|
||||
im.save(buf, "PNG")
|
||||
return buf.getvalue()
|
||||
|
||||
|
||||
def test_the_engine_loads_and_health_says_ready():
|
||||
assert ocr.available() is True, ocr.status()
|
||||
assert ocr.status()["state"] == "ready"
|
||||
|
||||
|
||||
def test_printed_label_text_is_read_in_reading_order():
|
||||
text = ocr.read_text(_label_png(["MARIE GOLD", "300 g"]))
|
||||
|
||||
assert text, "engine read nothing"
|
||||
upper = text.upper()
|
||||
assert "MARIE" in upper and "GOLD" in upper
|
||||
assert "300" in upper
|
||||
assert upper.index("MARIE") < upper.index("300")
|
||||
|
||||
|
||||
def test_a_blank_image_reads_nothing():
|
||||
assert ocr.read_text(_label_png([])) is None
|
||||
316
tests/test_ocr_service.py
Normal file
316
tests/test_ocr_service.py
Normal file
@@ -0,0 +1,316 @@
|
||||
"""Server-side OCR (app/services/ocr_service.py): the label text off a photo.
|
||||
|
||||
What has to hold, each on its own:
|
||||
|
||||
1. The label is the engine's lines in reading order, low-confidence lines
|
||||
dropped, repeats dropped, capped at OCR_MAX_CHARS on a word boundary - the
|
||||
same shape as the routes' `text` field.
|
||||
2. The frame the engine sees is the FULL photo (never the embedder's 224
|
||||
crop), decoded through the same guards as the vector path, and downscaled
|
||||
to OCR_MAX_SIDE_PX. Undecodable bytes never reach the engine.
|
||||
3. The engine is one lazily-built instance behind one lock; with the flag off
|
||||
it is never imported; without the wheel it says so once and returns None
|
||||
forever after; a read that raises is None, not an exception.
|
||||
4. /api/health carries an `ocr` block and asking never loads the engine.
|
||||
|
||||
No engine is involved anywhere: `FakeEngine` records what it was given and
|
||||
answers with canned lines. tests/test_ocr_engine_real.py runs the real one
|
||||
when it is installed.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import io
|
||||
import threading
|
||||
from types import SimpleNamespace
|
||||
from typing import Any, List
|
||||
|
||||
import numpy as np
|
||||
import pytest
|
||||
from PIL import Image
|
||||
|
||||
from app.services import image_embedder as emb
|
||||
from app.services import ocr_service as ocr
|
||||
|
||||
|
||||
def _png(im: Image.Image) -> bytes:
|
||||
buf = io.BytesIO()
|
||||
im.save(buf, "PNG")
|
||||
return buf.getvalue()
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _fresh_state(monkeypatch):
|
||||
ocr._reset()
|
||||
# conftest pins the env to false so no test can load the real engine;
|
||||
# these tests drive a fake and turn the flag on for themselves.
|
||||
monkeypatch.setattr(ocr, "ENABLE_SERVER_OCR", True)
|
||||
yield
|
||||
ocr._reset()
|
||||
|
||||
|
||||
class FakeOutput:
|
||||
def __init__(self, txts, scores=None):
|
||||
self.txts = tuple(txts)
|
||||
self.scores = tuple(scores if scores is not None else [0.99] * len(txts))
|
||||
self.boxes = None
|
||||
|
||||
def __len__(self):
|
||||
return len(self.txts)
|
||||
|
||||
|
||||
class FakeEngine:
|
||||
"""Stands in for rapidocr.RapidOCR: records each frame's shape, returns
|
||||
canned lines."""
|
||||
|
||||
def __init__(self, txts=("MARIE GOLD", "300 g"), scores=None):
|
||||
self.frames: List[Any] = []
|
||||
self.output = FakeOutput(txts, scores)
|
||||
|
||||
def __call__(self, img, **kwargs):
|
||||
self.frames.append(np.asarray(img).shape)
|
||||
return self.output
|
||||
|
||||
|
||||
def _install_fake(monkeypatch, **kw) -> FakeEngine:
|
||||
fake = FakeEngine(**kw)
|
||||
monkeypatch.setattr(ocr, "_engine", fake)
|
||||
return fake
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 1. lines -> label
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def test_lines_are_joined_in_reading_order_and_low_scores_dropped(monkeypatch):
|
||||
_install_fake(monkeypatch, txts=("Britannia", "MARIE GOLD", "smudge", "300 g"),
|
||||
scores=(0.95, 0.99, 0.2, 0.9))
|
||||
|
||||
assert ocr.read_text(_png(Image.new("RGB", (64, 64)))) == "Britannia MARIE GOLD 300 g"
|
||||
|
||||
|
||||
def test_repeated_lines_are_deduped_case_insensitively(monkeypatch):
|
||||
_install_fake(monkeypatch, txts=("Marie Gold", "MARIE GOLD", " Marie Gold ", "300 g"))
|
||||
|
||||
assert ocr.read_text(_png(Image.new("RGB", (64, 64)))) == "Marie Gold 300 g"
|
||||
|
||||
|
||||
def test_output_is_capped_at_ocr_max_chars_on_a_word_boundary(monkeypatch):
|
||||
monkeypatch.setattr(ocr, "OCR_MAX_CHARS", 20)
|
||||
_install_fake(monkeypatch, txts=("Britannia Marie Gold", "Biscuits 300 g"))
|
||||
|
||||
text = ocr.read_text(_png(Image.new("RGB", (64, 64))))
|
||||
|
||||
assert text == "Britannia Marie Gold"
|
||||
assert len(text) <= 20
|
||||
|
||||
|
||||
def test_a_single_word_longer_than_the_cap_is_hard_cut(monkeypatch):
|
||||
monkeypatch.setattr(ocr, "OCR_MAX_CHARS", 5)
|
||||
_install_fake(monkeypatch, txts=("Britannia",))
|
||||
|
||||
assert ocr.read_text(_png(Image.new("RGB", (64, 64)))) == "Brita"
|
||||
|
||||
|
||||
def test_nothing_read_is_none_not_an_empty_string(monkeypatch):
|
||||
_install_fake(monkeypatch, txts=())
|
||||
assert ocr.read_text(_png(Image.new("RGB", (64, 64)))) is None
|
||||
|
||||
_install_fake(monkeypatch, txts=("", " "))
|
||||
assert ocr.read_text(_png(Image.new("RGB", (64, 64)))) is None
|
||||
|
||||
|
||||
def test_an_engine_answering_none_is_none(monkeypatch):
|
||||
fake = _install_fake(monkeypatch)
|
||||
fake.output = None
|
||||
assert ocr.read_text(_png(Image.new("RGB", (64, 64)))) is None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 2. the frame
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@pytest.mark.parametrize("data", [b"", b"not an image", b"\x89PNG\r\n\x1a\n" + b"\x00" * 20])
|
||||
def test_undecodable_bytes_give_none_without_touching_the_engine(monkeypatch, data):
|
||||
fake = _install_fake(monkeypatch)
|
||||
assert ocr.read_text(data) is None
|
||||
assert fake.frames == []
|
||||
|
||||
|
||||
def test_the_engine_sees_the_whole_frame_not_the_224_crop(monkeypatch):
|
||||
fake = _install_fake(monkeypatch)
|
||||
ocr.read_text(_png(Image.new("RGB", (640, 200))))
|
||||
assert fake.frames == [(200, 640, 3)]
|
||||
|
||||
|
||||
def test_a_large_photo_is_downscaled_to_max_side_before_detection(monkeypatch):
|
||||
monkeypatch.setattr(ocr, "OCR_MAX_SIDE_PX", 1280)
|
||||
fake = _install_fake(monkeypatch)
|
||||
ocr.read_text(_png(Image.new("RGB", (4000, 3000))))
|
||||
h, w, c = fake.frames[0]
|
||||
assert (w, h, c) == (1280, 960, 3)
|
||||
|
||||
|
||||
def test_a_small_photo_is_not_upscaled(monkeypatch):
|
||||
monkeypatch.setattr(ocr, "OCR_MAX_SIDE_PX", 1280)
|
||||
fake = _install_fake(monkeypatch)
|
||||
ocr.read_text(_png(Image.new("RGB", (300, 100))))
|
||||
assert fake.frames == [(100, 300, 3)]
|
||||
|
||||
|
||||
def test_the_pixel_cap_of_the_vector_path_applies(monkeypatch):
|
||||
monkeypatch.setattr(emb, "IMAGE_VECTOR_MAX_PIXELS", 100)
|
||||
fake = _install_fake(monkeypatch)
|
||||
assert ocr.read_text(_png(Image.new("RGB", (20, 20)))) is None
|
||||
assert fake.frames == []
|
||||
|
||||
|
||||
def test_transparency_is_flattened_like_the_vector_path(monkeypatch):
|
||||
fake = _install_fake(monkeypatch)
|
||||
ocr.read_text(_png(Image.new("RGBA", (32, 16), (255, 0, 0, 0))))
|
||||
assert fake.frames == [(16, 32, 3)]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 3. the engine lifecycle
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def test_flag_off_means_unavailable_without_importing_anything(monkeypatch, caplog):
|
||||
monkeypatch.setattr(ocr, "ENABLE_SERVER_OCR", False)
|
||||
import builtins
|
||||
real_import = builtins.__import__
|
||||
|
||||
def no_rapidocr(name, *a, **k):
|
||||
if name.startswith("rapidocr"):
|
||||
raise AssertionError("rapidocr must not be imported when the flag is off")
|
||||
return real_import(name, *a, **k)
|
||||
|
||||
monkeypatch.setattr(builtins, "__import__", no_rapidocr)
|
||||
|
||||
with caplog.at_level("WARNING", logger=ocr.__name__):
|
||||
assert ocr.available() is False
|
||||
assert ocr.read_text(_png(Image.new("RGB", (8, 8)))) is None
|
||||
|
||||
assert [r for r in caplog.records if r.levelname == "WARNING"] == []
|
||||
assert ocr.status()["state"] == "disabled: ENABLE_SERVER_OCR is false"
|
||||
assert ocr.status()["enabled"] is False
|
||||
|
||||
|
||||
def test_a_missing_wheel_disables_with_one_warning(monkeypatch, caplog):
|
||||
import builtins
|
||||
real_import = builtins.__import__
|
||||
|
||||
def no_rapidocr(name, *a, **k):
|
||||
if name.startswith("rapidocr"):
|
||||
raise ImportError("No module named 'rapidocr'")
|
||||
return real_import(name, *a, **k)
|
||||
|
||||
monkeypatch.setattr(builtins, "__import__", no_rapidocr)
|
||||
|
||||
with caplog.at_level("WARNING", logger=ocr.__name__):
|
||||
assert ocr.available() is False
|
||||
assert ocr.available() is False
|
||||
assert ocr.read_text(_png(Image.new("RGB", (8, 8)))) is None
|
||||
|
||||
warnings_ = [r for r in caplog.records if r.levelname == "WARNING"]
|
||||
assert len(warnings_) == 1 and "not usable" in warnings_[0].getMessage()
|
||||
assert ocr.status()["state"].startswith("disabled: rapidocr is not usable")
|
||||
|
||||
|
||||
def test_an_engine_that_raises_gives_none_not_an_exception(monkeypatch):
|
||||
class Boom(FakeEngine):
|
||||
def __call__(self, img, **kw):
|
||||
raise RuntimeError("onnxruntime session died")
|
||||
|
||||
monkeypatch.setattr(ocr, "_engine", Boom())
|
||||
assert ocr.read_text(_png(Image.new("RGB", (8, 8)))) is None
|
||||
# The engine is kept: one bad photo is not a reason to disable OCR.
|
||||
assert ocr.status()["state"] == "ready"
|
||||
|
||||
|
||||
def test_reads_are_serialised_on_the_module_lock(monkeypatch):
|
||||
inside = []
|
||||
overlap = []
|
||||
|
||||
class SlowFake(FakeEngine):
|
||||
def __call__(self, img, **kw):
|
||||
inside.append(1)
|
||||
if len(inside) > 1:
|
||||
overlap.append(1)
|
||||
threading.Event().wait(0.02)
|
||||
inside.pop()
|
||||
return super().__call__(img, **kw)
|
||||
|
||||
fake = SlowFake()
|
||||
monkeypatch.setattr(ocr, "_engine", fake)
|
||||
data = _png(Image.new("RGB", (8, 8)))
|
||||
threads = [threading.Thread(target=ocr.read_text, args=(data,)) for _ in range(6)]
|
||||
for t in threads:
|
||||
t.start()
|
||||
for t in threads:
|
||||
t.join()
|
||||
|
||||
assert not overlap and len(fake.frames) == 6
|
||||
|
||||
|
||||
def test_the_engine_is_built_once_with_the_configured_params(monkeypatch):
|
||||
built = []
|
||||
|
||||
class FakeRapidOCR:
|
||||
def __init__(self, params=None):
|
||||
built.append(params)
|
||||
|
||||
def __call__(self, img, **kw):
|
||||
return FakeOutput(("x",))
|
||||
|
||||
import sys
|
||||
monkeypatch.setitem(sys.modules, "rapidocr", SimpleNamespace(RapidOCR=FakeRapidOCR))
|
||||
monkeypatch.setattr(ocr, "OCR_MIN_CONFIDENCE", 0.42)
|
||||
monkeypatch.setattr(ocr, "OCR_NUM_THREADS", 3)
|
||||
|
||||
assert ocr.available() is True
|
||||
assert ocr.available() is True
|
||||
assert ocr.read_text(_png(Image.new("RGB", (8, 8)))) == "x"
|
||||
|
||||
assert len(built) == 1
|
||||
assert built[0]["Global.text_score"] == 0.42
|
||||
assert built[0]["EngineConfig.onnxruntime.intra_op_num_threads"] == 3
|
||||
assert built[0]["Global.log_level"] == "warning"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 4. status and health
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def test_status_never_loads_the_engine(monkeypatch):
|
||||
loads = []
|
||||
monkeypatch.setattr(ocr, "_load", lambda: loads.append(1))
|
||||
|
||||
st = ocr.status()
|
||||
|
||||
assert st["state"].startswith("not loaded yet")
|
||||
assert st["enabled"] is True
|
||||
assert set(st) == {"enabled", "runtime_importable", "onnxruntime_importable", "state"}
|
||||
assert loads == []
|
||||
|
||||
|
||||
def test_status_reflects_a_recorded_failure_and_a_ready_engine(monkeypatch):
|
||||
monkeypatch.setattr(ocr, "_disabled_reason", "rapidocr is not usable (x)")
|
||||
assert ocr.status()["state"] == "disabled: rapidocr is not usable (x)"
|
||||
|
||||
ocr._reset()
|
||||
_install_fake(monkeypatch)
|
||||
assert ocr.status()["state"] == "ready"
|
||||
|
||||
|
||||
def test_health_carries_the_ocr_block(monkeypatch):
|
||||
from fastapi.testclient import TestClient
|
||||
from app.main import app
|
||||
loads = []
|
||||
monkeypatch.setattr(ocr, "_load", lambda: loads.append(1))
|
||||
|
||||
body = TestClient(app).get("/api/health").json()
|
||||
|
||||
block = body["ocr"]
|
||||
assert set(block) == {"enabled", "runtime_importable", "onnxruntime_importable", "state"}
|
||||
assert isinstance(block["enabled"], bool)
|
||||
assert loads == []
|
||||
249
tests/test_product_identify.py
Normal file
249
tests/test_product_identify.py
Normal file
@@ -0,0 +1,249 @@
|
||||
"""The identify ladder (app/services/product_identify.py), one test per rung.
|
||||
|
||||
The image arm, the text arm and the OCR engine are all patched here; this
|
||||
file checks only the decisions between them: when the image is the answer,
|
||||
when the label takes over, where the label comes from, what the response
|
||||
says when nothing is confirmed, and that the OCR engine is never asked to
|
||||
work when it is not needed.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
import pytest
|
||||
|
||||
from app.services import product_identify as pi
|
||||
from app.services.image_match import ImageSearchResult
|
||||
|
||||
|
||||
def _rows(*scores: float, prefix: str = "img") -> List[Dict[str, Any]]:
|
||||
return [{"image_id": f"{prefix}_{i}", "product_name": f"{prefix} {i}", "brand": "Britannia",
|
||||
"brand_table": "brand_britannia", "score": s, "text_overlap": 0.0}
|
||||
for i, s in enumerate(scores)]
|
||||
|
||||
|
||||
class _Arms:
|
||||
"""Records the calls into both arms and the OCR engine."""
|
||||
|
||||
def __init__(self):
|
||||
self.image_rows: List[Dict[str, Any]] = []
|
||||
self.text_rows: List[Dict[str, Any]] = []
|
||||
self.text_confident = True
|
||||
self.ocr_available = True
|
||||
self.ocr_text: Optional[str] = "MARIE GOLD 300 g"
|
||||
self.image_calls: List[Dict[str, Any]] = []
|
||||
self.text_calls: List[Dict[str, Any]] = []
|
||||
self.ocr_calls: List[bytes] = []
|
||||
self.ocr_probes = 0
|
||||
|
||||
def search_by_vector(self, vector, text=None, brand=None, category=None, top_k=10, min_score=0.0):
|
||||
self.image_calls.append({"text": text, "brand": brand, "category": category, "top_k": top_k,
|
||||
"min_score": min_score})
|
||||
return ImageSearchResult(rows=[dict(r) for r in self.image_rows], detected_brand=brand,
|
||||
min_score=min_score, top_k=top_k, query_text=text)
|
||||
|
||||
def resolve_label(self, text, brand=None, category=None, top_k=10):
|
||||
self.text_calls.append({"text": text, "brand": brand, "category": category, "top_k": top_k})
|
||||
return ImageSearchResult(rows=[dict(r) for r in self.text_rows], query_text=text, top_k=top_k)
|
||||
|
||||
def is_confident(self, rows):
|
||||
return bool(rows) and self.text_confident
|
||||
|
||||
def available(self):
|
||||
self.ocr_probes += 1
|
||||
return self.ocr_available
|
||||
|
||||
def read_text(self, data):
|
||||
self.ocr_calls.append(data)
|
||||
return self.ocr_text
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def arms(monkeypatch):
|
||||
a = _Arms()
|
||||
monkeypatch.setattr(pi, "search_by_vector", a.search_by_vector)
|
||||
monkeypatch.setattr(pi, "resolve_label", a.resolve_label)
|
||||
monkeypatch.setattr(pi, "is_confident", a.is_confident)
|
||||
monkeypatch.setattr(pi.ocr_service, "available", a.available)
|
||||
monkeypatch.setattr(pi.ocr_service, "read_text", a.read_text)
|
||||
monkeypatch.setattr(pi, "IMAGE_IDENTIFY_MIN_IMAGE_SCORE", 0.70)
|
||||
return a
|
||||
|
||||
|
||||
VEC = [1.0] + [0.0] * 1023
|
||||
PHOTO = b"\x89PNG fake"
|
||||
|
||||
|
||||
def _identify(**kw):
|
||||
kw.setdefault("vector", VEC)
|
||||
kw.setdefault("image_bytes", PHOTO)
|
||||
return pi.identify_product(**kw)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# the image is the answer
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def test_a_confident_image_match_is_the_answer_and_ocr_is_not_touched(arms):
|
||||
arms.image_rows = _rows(0.91, 0.80)
|
||||
|
||||
out = _identify()
|
||||
|
||||
assert out.matched_by == "image_vector" and out.fallback_reason is None
|
||||
assert out.image_top_score == pytest.approx(0.91)
|
||||
assert [r["image_id"] for r in out.search.rows] == ["img_0", "img_1"]
|
||||
assert arms.text_calls == [] and arms.ocr_calls == [] and arms.ocr_probes == 0
|
||||
assert out.ocr_text is None and out.ocr_source is None
|
||||
|
||||
|
||||
def test_the_floor_is_inclusive_and_configurable(arms):
|
||||
arms.image_rows = _rows(0.70)
|
||||
arms.text_rows = _rows(0.5, prefix="txt")
|
||||
assert _identify().matched_by == "image_vector"
|
||||
assert _identify(min_image_score=0.71).matched_by == "text"
|
||||
|
||||
|
||||
def test_client_text_is_passed_to_the_image_arm_for_scoping_and_tie_breaks(arms):
|
||||
arms.image_rows = _rows(0.95)
|
||||
|
||||
out = _identify(text="Britannia Marie Gold 300 g", brand="Britannia", category="Biscuits", top_k=3, min_score=0.2)
|
||||
|
||||
assert arms.image_calls == [{"text": "Britannia Marie Gold 300 g", "brand": "Britannia",
|
||||
"category": "Biscuits", "top_k": 3, "min_score": 0.2}]
|
||||
assert out.ocr_text == "Britannia Marie Gold 300 g" and out.ocr_source == "client"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# the label takes over
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def test_a_low_score_with_client_text_resolves_by_text_without_ocr(arms):
|
||||
arms.image_rows = _rows(0.63)
|
||||
arms.text_rows = _rows(0.55, 0.40, prefix="txt")
|
||||
|
||||
out = _identify(text=" MARIE GOLD 300 g ")
|
||||
|
||||
assert out.matched_by == "text" and out.fallback_reason == "image_below_threshold"
|
||||
assert out.ocr_source == "client" and out.ocr_text == "MARIE GOLD 300 g"
|
||||
assert out.image_top_score == pytest.approx(0.63)
|
||||
assert [r["image_id"] for r in out.search.rows] == ["txt_0", "txt_1"]
|
||||
assert arms.text_calls == [{"text": "MARIE GOLD 300 g", "brand": None, "category": None, "top_k": 10}]
|
||||
assert arms.ocr_calls == [] and arms.ocr_probes == 0
|
||||
|
||||
|
||||
def test_a_low_score_without_text_uses_server_ocr(arms):
|
||||
arms.image_rows = _rows(0.63)
|
||||
arms.text_rows = _rows(0.55, prefix="txt")
|
||||
|
||||
out = _identify()
|
||||
|
||||
assert arms.ocr_calls == [PHOTO]
|
||||
assert out.matched_by == "text" and out.ocr_source == "server" and out.ocr_text == "MARIE GOLD 300 g"
|
||||
assert out.fallback_reason == "image_below_threshold"
|
||||
assert arms.text_calls[0]["text"] == "MARIE GOLD 300 g"
|
||||
|
||||
|
||||
def test_no_image_match_at_all_is_its_own_reason(arms):
|
||||
arms.image_rows = []
|
||||
arms.text_rows = _rows(0.55, prefix="txt")
|
||||
|
||||
out = _identify(text="Marie Gold")
|
||||
|
||||
assert out.matched_by == "text" and out.fallback_reason == "no_image_match"
|
||||
assert out.image_top_score is None
|
||||
|
||||
|
||||
def test_no_vector_means_text_only_and_says_so(arms):
|
||||
arms.text_rows = _rows(0.55, prefix="txt")
|
||||
|
||||
out = _identify(vector=None, text="Marie Gold 300 g")
|
||||
|
||||
assert arms.image_calls == []
|
||||
assert out.matched_by == "text" and out.fallback_reason == "image_embedder_unavailable"
|
||||
assert out.image_top_score is None
|
||||
assert out.search.top_k == 10
|
||||
|
||||
|
||||
def test_brand_category_and_top_k_reach_the_text_arm(arms):
|
||||
arms.image_rows = _rows(0.1)
|
||||
arms.text_rows = _rows(0.9, prefix="txt")
|
||||
|
||||
_identify(text="x", brand="Britannia", category="Biscuits", top_k=4)
|
||||
|
||||
assert arms.text_calls == [{"text": "x", "brand": "Britannia", "category": "Biscuits", "top_k": 4}]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# nothing confirmed
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def test_no_text_and_ocr_unavailable_returns_the_low_image_rows_with_the_reason(arms):
|
||||
arms.image_rows = _rows(0.63)
|
||||
arms.ocr_available = False
|
||||
|
||||
out = _identify()
|
||||
|
||||
assert out.matched_by == "image_vector" and out.fallback_reason == "ocr_unavailable"
|
||||
assert [r["image_id"] for r in out.search.rows] == ["img_0"]
|
||||
assert out.image_top_score == pytest.approx(0.63)
|
||||
assert out.ocr_text is None and out.ocr_source is None
|
||||
assert arms.ocr_calls == [] and arms.text_calls == []
|
||||
|
||||
|
||||
def test_ocr_reading_nothing_is_ocr_empty(arms):
|
||||
arms.image_rows = _rows(0.63)
|
||||
arms.ocr_text = None
|
||||
|
||||
out = _identify()
|
||||
|
||||
assert out.matched_by == "image_vector" and out.fallback_reason == "ocr_empty"
|
||||
assert arms.ocr_calls == [PHOTO] and arms.text_calls == []
|
||||
|
||||
|
||||
def test_no_photo_and_no_text_is_no_text(arms):
|
||||
arms.image_rows = _rows(0.63)
|
||||
|
||||
out = _identify(image_bytes=None)
|
||||
|
||||
assert out.fallback_reason == "no_text" and out.matched_by == "image_vector"
|
||||
assert arms.ocr_probes == 0
|
||||
|
||||
|
||||
def test_nothing_anywhere_is_matched_by_none(arms):
|
||||
arms.image_rows = []
|
||||
arms.ocr_available = False
|
||||
|
||||
out = _identify()
|
||||
|
||||
assert out.matched_by == "none" and out.fallback_reason == "ocr_unavailable"
|
||||
assert out.search.rows == []
|
||||
|
||||
|
||||
def test_an_unconfident_text_result_keeps_the_image_rows_and_the_label(arms):
|
||||
arms.image_rows = _rows(0.63)
|
||||
arms.text_rows = _rows(0.2, prefix="txt")
|
||||
arms.text_confident = False
|
||||
|
||||
out = _identify()
|
||||
|
||||
assert out.matched_by == "image_vector" and out.fallback_reason == "text_no_match"
|
||||
assert [r["image_id"] for r in out.search.rows] == ["img_0"]
|
||||
assert out.ocr_text == "MARIE GOLD 300 g" and out.ocr_source == "server"
|
||||
assert out.image_top_score == pytest.approx(0.63)
|
||||
|
||||
|
||||
def test_an_unconfident_text_result_with_no_image_rows_is_none(arms):
|
||||
arms.image_rows = []
|
||||
arms.text_rows = _rows(0.2, prefix="txt")
|
||||
arms.text_confident = False
|
||||
|
||||
out = _identify(text="zzz")
|
||||
|
||||
assert out.matched_by == "none" and out.fallback_reason == "text_no_match"
|
||||
assert out.search.rows == []
|
||||
|
||||
|
||||
def test_an_invalid_vector_propagates_as_invalid_vector_error(monkeypatch):
|
||||
from app.services.image_match import InvalidVectorError
|
||||
with pytest.raises(InvalidVectorError):
|
||||
pi.identify_product(vector=[0.0] * 3, image_bytes=None, text="x")
|
||||
Reference in New Issue
Block a user