From f933ea10a155e646938bd476126f66ed4a47b40b Mon Sep 17 00:00:00 2001 From: sriram Date: Sat, 19 Sep 2026 15:39:53 +0530 Subject: [PATCH] Image vector to product details --- .env.example | 13 ++ Dockerfile | 9 + app/api/routers/health.py | 7 +- app/api/routers/search.py | 94 +++++++- app/api/schemas.py | 39 ++++ app/infrastructure/settings.py | 36 ++++ app/services/image_embedder.py | 57 +++-- app/services/label_match.py | 360 +++++++++++++++++++++++++++++++ app/services/ocr_service.py | 245 +++++++++++++++++++++ app/services/product_identify.py | 162 ++++++++++++++ docs/IMAGE_SEARCH_API.md | 102 ++++++++- requirements-ocr.txt | 17 ++ requirements.txt | 20 ++ tests/conftest.py | 5 + tests/test_identify_api.py | 294 +++++++++++++++++++++++++ tests/test_image_search_api.py | 1 + tests/test_label_match.py | 330 ++++++++++++++++++++++++++++ tests/test_ocr_engine_real.py | 62 ++++++ tests/test_ocr_service.py | 316 +++++++++++++++++++++++++++ tests/test_product_identify.py | 249 +++++++++++++++++++++ 20 files changed, 2391 insertions(+), 27 deletions(-) create mode 100644 app/services/label_match.py create mode 100644 app/services/ocr_service.py create mode 100644 app/services/product_identify.py create mode 100644 requirements-ocr.txt create mode 100644 tests/test_identify_api.py create mode 100644 tests/test_label_match.py create mode 100644 tests/test_ocr_engine_real.py create mode 100644 tests/test_ocr_service.py create mode 100644 tests/test_product_identify.py diff --git a/.env.example b/.env.example index 17c4cdc..dd4d9c2 100644 --- a/.env.example +++ b/.env.example @@ -326,6 +326,19 @@ ENABLE_IMAGE_VECTORS=true #IMAGE_EMBED_MODEL_PATH=/app/app/services/models/mobilenet/mobilenet_v3_small_embedder.tflite #IMAGE_EMBED_NUM_THREADS=2 +# Identify a product from a phone photo (POST /api/search/identify). A phone +# photo scores ~0.63 on img_vector against the catalog's render of the same +# pack, so below IMAGE_IDENTIFY_MIN_IMAGE_SCORE the label text decides: the +# client's OCR `text` if it sent one, else what the server reads off the photo +# with rapidocr (models ship in the wheel; see requirements-ocr.txt). Set +# ENABLE_SERVER_OCR=false to rely on client text only. +ENABLE_SERVER_OCR=true +#OCR_MIN_CONFIDENCE=0.5 +#OCR_MAX_SIDE_PX=1280 +#OCR_NUM_THREADS=2 +#IMAGE_IDENTIFY_MIN_IMAGE_SCORE=0.70 +#IMAGE_IDENTIFY_MIN_TEXT_SCORE=0.60 + # Product SKU: try a live web search for a real marketplace product ID # (Amazon ASIN, Flipkart PID, etc.) before falling back to an internal SKU. # Set to false to always generate internal SKUs only (faster, offline-safe). diff --git a/Dockerfile b/Dockerfile index 32de7ea..e6bd778 100644 --- a/Dockerfile +++ b/Dockerfile @@ -21,6 +21,7 @@ WORKDIR /app # image-search fallback (see requirements.txt); run `playwright install # chromium` in the container if you need that specific fallback tier. COPY requirements.txt . +COPY requirements-ocr.txt . # torch is installed FIRST, from PyTorch's CPU-only index, and that ordering is # the point. sentence-transformers pulls torch in transitively, and pip's @@ -63,6 +64,14 @@ RUN /opt/venv/bin/pip install --no-cache-dir \ RUN /opt/venv/bin/pip install --no-cache-dir -r requirements.txt +# The OCR engine goes in AFTER requirements.txt and WITHOUT its declared +# dependencies. rapidocr's metadata asks for opencv-python (the GUI build); +# letting pip honour that would unpack it over opencv-python-headless and the +# result fails on libGL.so.1 in this image - taking the img_vector embedder +# down with it. Its real dependencies are in requirements.txt already; its +# PP-OCR models are inside the wheel, so nothing is downloaded at runtime. +RUN /opt/venv/bin/pip install --no-cache-dir --no-deps -r requirements-ocr.txt + # Strip payload the running service can never execute. Doing this in the build # stage is what makes it count: the runtime stage copies /opt/venv as one layer, # so anything deleted after that COPY would still occupy space in the layer diff --git a/app/api/routers/health.py b/app/api/routers/health.py index a832066..d1cc6ae 100644 --- a/app/api/routers/health.py +++ b/app/api/routers/health.py @@ -5,12 +5,12 @@ import logging import requests from fastapi import APIRouter -from app.api.schemas import AuthConfigOut, HealthOut, ImageVectorsOut +from app.api.schemas import AuthConfigOut, HealthOut, ImageVectorsOut, OcrOut from app.infrastructure.security import auth_config_summary from app.infrastructure.settings import ( ENABLE_IMAGE_VECTORS, OLLAMA_BASE_URL, OLLAMA_MODEL_NAME, EMBEDDINGS_MODEL, ) -from app.services import image_embedder +from app.services import image_embedder, ocr_service from app.services.vector_store import _connect # internal, but handy for a connectivity probe logger = logging.getLogger(__name__) @@ -63,4 +63,7 @@ def health() -> HealthOut: # usual cause (the model file is not in the image), and it must be # visible from outside the container. status() never loads the model. image_vectors=ImageVectorsOut(enabled=ENABLE_IMAGE_VECTORS, **image_embedder.status()), + # And for server-side OCR behind /search/identify: "ocr_unavailable" + # in a response has one of three causes, and this names it. + ocr=OcrOut(**ocr_service.status()), ) diff --git a/app/api/routers/search.py b/app/api/routers/search.py index 7a1acb8..248a1e5 100644 --- a/app/api/routers/search.py +++ b/app/api/routers/search.py @@ -3,10 +3,12 @@ from __future__ import annotations from typing import Optional from fastapi import APIRouter, File, Form, HTTPException, Query, UploadFile +from fastapi.responses import JSONResponse from starlette.concurrency import run_in_threadpool from app.api.routers.brands import _row_to_product_out from app.api.schemas import ( + IdentifyOut, ImageMatchOut, ImageSearchOut, ImageVectorSearchRequest, @@ -30,6 +32,7 @@ from app.services.image_match import ( parse_vector_param, search_by_vector, ) +from app.services.product_identify import IdentifyResult, identify_product router = APIRouter(tags=["search"]) @@ -101,8 +104,19 @@ def _to_image_search_out(result: ImageSearchResult) -> ImageSearchOut: ) +def _to_identify_out(result: IdentifyResult) -> IdentifyOut: + return IdentifyOut( + **_to_image_search_out(result.search).model_dump(), + matched_by=result.matched_by, + ocr_text=result.ocr_text, + ocr_source=result.ocr_source, + image_top_score=None if result.image_top_score is None else round(float(result.image_top_score), 4), + fallback_reason=result.fallback_reason, + ) + + @router.post("/search/image-vector", response_model=ImageSearchOut) -def image_vector_search_endpoint(body: ImageVectorSearchRequest) -> ImageSearchOut: +def image_vector_search_endpoint(body: ImageVectorSearchRequest): """Products that look like the photo whose embedding is `vector`. `vector` is the 1024-float, L2-normalised MobileNetV3-Small embedding the @@ -112,8 +126,22 @@ def image_vector_search_endpoint(body: ImageVectorSearchRequest) -> ImageSearchO among products that share one photo. An explicit `brand` is a hard filter; a brand recognised from `text` falls back to every brand when it finds nothing (`scope_fallback`). + + With `text_fallback: true` the request runs the /search/identify ladder + instead: when the best image score is under IMAGE_IDENTIFY_MIN_IMAGE_SCORE + the label `text` is resolved against the catalogue, and the response is + an IdentifyOut (this response plus matched_by, fallback_reason, ...). + Off by default so the app's existing calls are answered exactly as before. """ try: + if body.text_fallback: + identified = identify_product( + vector=body.vector, image_bytes=None, text=body.text, brand=body.brand, + category=body.category, top_k=body.top_k, min_score=body.min_score, + ) + # A Response bypasses response_model, which is the point: the + # default path keeps its declared ImageSearchOut contract. + return JSONResponse(content=_to_identify_out(identified).model_dump(mode="json")) result = search_by_vector( body.vector, text=body.text, brand=body.brand, category=body.category, top_k=body.top_k, min_score=body.min_score, @@ -199,3 +227,67 @@ async def image_search_endpoint( except InvalidVectorError as exc: # cannot happen for a model output, but the route must not 500 raise HTTPException(status_code=422, detail=str(exc)) return _to_image_search_out(result) + + +# --------------------------------------------------------------------------- +# Identify: image first, label text second +# --------------------------------------------------------------------------- +# A phone photo of a pack scores ~0.63 against the catalog's render of the +# same pack, so /search/image alone cannot confirm a product. This route +# runs the ladder in app/services/product_identify.py: the image match when +# it clears IMAGE_IDENTIFY_MIN_IMAGE_SCORE, otherwise the label text - the +# client's `text`, else what the server reads off the photo (ocr_service) - +# resolved through the catalogue's text embeddings and product names. + +@router.post("/search/identify", response_model=IdentifyOut) +async def identify_endpoint( + file: UploadFile = File(..., description="The product photo (JPEG/PNG/WebP), ideally cropped to the pack"), + text: Optional[str] = Form(None, max_length=500, + description="OCR text read off the label; when absent the server reads it"), + brand: Optional[str] = Form(None, max_length=120), + category: Optional[str] = Form(None, max_length=120), + top_k: int = Form(IMAGE_SEARCH_DEFAULT_TOP_K, ge=1, le=IMAGE_SEARCH_MAX_TOP_K), + min_score: float = Form(IMAGE_SEARCH_DEFAULT_MIN_SCORE, ge=-1.0, le=1.0), +) -> IdentifyOut: + """Which catalog product is in this photo. + + `matched_by` says which rung answered and therefore which space each + result's `score` is in; `fallback_reason` is set whenever the answer is + a best effort rather than a confirmed match. Works text-only on a + deployment without the image model (fallback_reason + "image_embedder_unavailable"); 503 only when neither a vector nor server + OCR is possible and no `text` was sent - GET /api/health -> image_vectors + and ocr say which is missing. + """ + content = await file.read() + if not content: + raise HTTPException(status_code=400, detail="The uploaded image is empty.") + if len(content) > IMAGE_VECTOR_MAX_BYTES: + raise HTTPException( + status_code=413, + detail=f"Image is {len(content) / 1_048_576:.1f} MB; the limit is " + f"{IMAGE_VECTOR_MAX_BYTES // 1_048_576} MB. Crop or downscale it.", + ) + vector = None + if await run_in_threadpool(image_embedder.available): + vector = await run_in_threadpool(image_embedder.embedding_for_bytes, content) + if vector is None: + raise HTTPException( + status_code=422, + detail="Could not decode the image (unsupported format, corrupt data, or too many pixels).", + ) + try: + result = await run_in_threadpool( + identify_product, vector=vector, image_bytes=content, text=text, brand=brand, + category=category, top_k=top_k, min_score=min_score, + ) + except InvalidVectorError as exc: # cannot happen for a model output, but the route must not 500 + raise HTTPException(status_code=422, detail=str(exc)) + if vector is None and result.fallback_reason == "ocr_unavailable": + raise HTTPException( + status_code=503, + detail="Neither the image embedding model nor server-side OCR is available on this " + "deployment. Send the label as `text`, or embed the photo client-side and POST " + "the vector to /api/search/image-vector.", + ) + return _to_identify_out(result) diff --git a/app/api/schemas.py b/app/api/schemas.py index bd207f8..9901514 100644 --- a/app/api/schemas.py +++ b/app/api/schemas.py @@ -138,6 +138,17 @@ class ImageVectorsOut(BaseModel): state: str = "unknown" +class OcrOut(BaseModel): + """Whether POST /api/search/identify can read a label off a photo itself. + Reported without loading the engine. `runtime_importable=false` after a + deploy means the rapidocr wheel was not installed (requirements-ocr.txt); + `onnxruntime_importable=false` means its engine was not.""" + enabled: bool = True + runtime_importable: bool = False + onnxruntime_importable: bool = False + state: str = "unknown" + + class HealthOut(BaseModel): status: str database: bool @@ -148,6 +159,7 @@ class HealthOut(BaseModel): # Defaulted so a client of this schema still validates against a # deployment predating the field. image_vectors: ImageVectorsOut = Field(default_factory=ImageVectorsOut) + ocr: OcrOut = Field(default_factory=OcrOut) # --------------------------------------------------------------------------- @@ -251,6 +263,16 @@ class ImageVectorSearchRequest(BaseModel): top_k: int = Field(IMAGE_SEARCH_DEFAULT_TOP_K, ge=1, le=IMAGE_SEARCH_MAX_TOP_K) min_score: float = Field(IMAGE_SEARCH_DEFAULT_MIN_SCORE, ge=-1.0, le=1.0, description="Drop matches with cosine similarity below this") + # Opt-in, so a client that never asked for it gets exactly the response + # it always did. With it, the route runs the identify ladder: when the + # best image score is under IMAGE_IDENTIFY_MIN_IMAGE_SCORE, `text` is + # resolved against the catalogue instead, and the response is an + # IdentifyOut (ImageSearchOut plus matched_by / fallback_reason / ...). + text_fallback: bool = Field( + False, + description="Fall back to resolving `text` when the image match is below the identify floor; " + "the response then carries the IdentifyOut fields", + ) @field_validator("vector") @classmethod @@ -281,6 +303,23 @@ class ImageSearchOut(BaseModel): query_text: Optional[str] = None +class IdentifyOut(ImageSearchOut): + """POST /search/identify: an ImageSearchOut plus which rung answered. + + `matched_by` names the space each result's `score` is in: "image_vector" + (cosine of the 1024-d photo embedding) or "text" (cosine of the 384-d + MiniLM embedding of the label). `image_top_score` always carries the + image side. The answer is confirmed when matched_by is "text", or + "image_vector" with no `fallback_reason`; otherwise `fallback_reason` + says why the best effort shown is unconfirmed (see product_identify.py). + """ + matched_by: str = "none" # "image_vector" | "text" | "none" + ocr_text: Optional[str] = None # the label text the ladder used + ocr_source: Optional[str] = None # "client" | "server" + image_top_score: Optional[float] = None + fallback_reason: Optional[str] = None + + # --------------------------------------------------------------------------- # RAG chat # --------------------------------------------------------------------------- diff --git a/app/infrastructure/settings.py b/app/infrastructure/settings.py index 22a2c37..a188189 100644 --- a/app/infrastructure/settings.py +++ b/app/infrastructure/settings.py @@ -433,6 +433,42 @@ IMAGE_SEARCH_DEFAULT_MIN_SCORE = float(os.getenv("IMAGE_SEARCH_DEFAULT_MIN_SCORE # the floor for hnsw.ef_search on that query so the index does not drop them. IMAGE_SEARCH_MAX_FETCH_K = int(os.getenv("IMAGE_SEARCH_MAX_FETCH_K", "100")) +# --------------------------------------------------------------------------- +# Identify a product from a phone photo - POST /api/search/identify +# (app/services/product_identify.py), with server-side OCR +# (app/services/ocr_service.py) and the label-text resolver +# (app/services/label_match.py). Public, read-only. +# --------------------------------------------------------------------------- +# A phone photo of a pack against the catalog's render of it scored 0.63 on +# img_vector (feasibility test), so the image alone cannot confirm a product. +# Below this floor the ladder falls through to the label text: whatever the +# client OCR'd, else what the server reads off the photo itself. +IMAGE_IDENTIFY_MIN_IMAGE_SCORE = float(os.getenv("IMAGE_IDENTIFY_MIN_IMAGE_SCORE", "0.70")) +# The text side is a MiniLM cosine (1 - embedding <=> q) in a different space +# from the image score. A nearest-neighbour query always returns SOMETHING, so +# a text match counts as found only when the label shares a name/size token +# with the row, or the cosine clears this floor. +IMAGE_IDENTIFY_MIN_TEXT_SCORE = float(os.getenv("IMAGE_IDENTIFY_MIN_TEXT_SCORE", "0.60")) +# Server-side OCR (rapidocr, PP-OCR models on onnxruntime CPU; models ship +# inside the wheel, nothing is downloaded). Loaded lazily on the first photo +# that needs it, never at boot; a missing wheel means "no server OCR" and +# GET /api/health -> ocr says so. tests/conftest.py pins it OFF. +ENABLE_SERVER_OCR = _bool("ENABLE_SERVER_OCR", "true") +# rapidocr's Global.text_score: recognised lines below this confidence are +# dropped before the label is assembled. +OCR_MIN_CONFIDENCE = float(os.getenv("OCR_MIN_CONFIDENCE", "0.5")) +# The photo is downscaled so its longer side is at most this before +# detection. The latency knob: a 12MP capture takes ~3x longer than 1280px +# and reads no better off a pack label. +OCR_MAX_SIDE_PX = int(os.getenv("OCR_MAX_SIDE_PX", "1280")) +# onnxruntime intra-op threads per session; the read itself is serialised by +# a lock (the engine is not thread-safe). +OCR_NUM_THREADS = int(os.getenv("OCR_NUM_THREADS", "2")) +# The assembled label is capped here, on a word boundary. 500 is the limit of +# the `text` field on the image-search routes, so client and server text are +# bounded alike. +OCR_MAX_CHARS = int(os.getenv("OCR_MAX_CHARS", "500")) + # --------------------------------------------------------------------------- # USDA FoodData Central - nutrition for loose, unbranded commodities # --------------------------------------------------------------------------- diff --git a/app/services/image_embedder.py b/app/services/image_embedder.py index d9e8e30..828a08b 100644 --- a/app/services/image_embedder.py +++ b/app/services/image_embedder.py @@ -81,6 +81,43 @@ def _flatten_alpha_to_bgr(image_bytes: bytes) -> Optional[np.ndarray]: return np.ascontiguousarray(np.asarray(rgb, dtype=np.uint8)[:, :, ::-1]) +def decode_bgr(image_bytes: bytes) -> Optional[np.ndarray]: + """`image_bytes` as a full-resolution uint8 BGR array, or None. + + Step 1 of `preprocess()`, on its own because the OCR service + (app/services/ocr_service.py) needs the same decode - pixel-bomb guard, + EXIF orientation, transparency flattened onto white - but the WHOLE frame: + the 224px centre crop that follows here would throw away the label. + + Never raises: an undecodable or oversized file is None (see preprocess). + """ + if not image_bytes: + return None + try: + import cv2 + from PIL import Image + + # Header only: the pixel-bomb guard and the transparency check both + # come from the file's metadata, no decode yet. + probe = Image.open(io.BytesIO(image_bytes)) + if probe.width * probe.height > IMAGE_VECTOR_MAX_PIXELS: + logger.debug("image rejected: %dx%d exceeds pixel cap", probe.width, probe.height) + return None + has_alpha = "A" in probe.getbands() or "transparency" in probe.info + + if has_alpha: + bgr = _flatten_alpha_to_bgr(image_bytes) + else: + bgr = cv2.imdecode(np.frombuffer(image_bytes, dtype=np.uint8), cv2.IMREAD_COLOR) + if bgr is None or bgr.ndim != 3 or bgr.shape[2] != 3: + logger.debug("image rejected: OpenCV could not decode it to BGR") + return None + return bgr + except Exception as exc: # noqa: BLE001 - see preprocess() + logger.debug("image undecodable: %s", exc) + return None + + def preprocess(image_bytes: bytes) -> Optional[np.ndarray]: """Decode `image_bytes` into the model's input tensor, or None. @@ -111,27 +148,11 @@ def preprocess(image_bytes: bytes) -> Optional[np.ndarray]: Never raises - pytest runs warnings as errors, and a DecompressionBombWarning is one of the things the broad except absorbs. """ - if not image_bytes: + bgr = decode_bgr(image_bytes) # 1 + if bgr is None: return None try: import cv2 - from PIL import Image - - # Header only: the pixel-bomb guard and the transparency check both - # come from the file's metadata, no decode yet. - probe = Image.open(io.BytesIO(image_bytes)) - if probe.width * probe.height > IMAGE_VECTOR_MAX_PIXELS: - logger.debug("image rejected: %dx%d exceeds pixel cap", probe.width, probe.height) - return None - has_alpha = "A" in probe.getbands() or "transparency" in probe.info - - if has_alpha: - bgr = _flatten_alpha_to_bgr(image_bytes) - else: - bgr = cv2.imdecode(np.frombuffer(image_bytes, dtype=np.uint8), cv2.IMREAD_COLOR) # 1 - if bgr is None or bgr.ndim != 3 or bgr.shape[2] != 3: - logger.debug("image rejected: OpenCV could not decode it to BGR") - return None h, w = bgr.shape[:2] # 2 side = min(w, h) diff --git a/app/services/label_match.py b/app/services/label_match.py new file mode 100644 index 0000000..6c4b902 --- /dev/null +++ b/app/services/label_match.py @@ -0,0 +1,360 @@ +"""Find catalog products from the TEXT on a pack: OCR label -> `embedding` + names. + +WHERE THIS SITS +--------------- +The second rung of POST /api/search/identify (app/services/product_identify.py). +When the photo's img_vector cannot confirm a product - a phone photo scores +~0.63 against the catalog's render of the same pack - the label text takes +over: the client's OCR `text`, or what app/services/ocr_service.py read off +the photo. This module turns that text into ranked catalog rows. + +WHY NOT catalog_search.search_catalog() +--------------------------------------- +GET /api/search is tuned for a human typing in a search box, and two of its +habits are wrong for OCR output: + +* it passes the FUZZY category detection through as a hard filter, and that + detector scores "taste" as Oral Care (matches "paste") and "Colgate" as + Chocolates - a misread word would delete the right product from the result; +* its lexical arm ANDs every term, so one smudged word means zero rows. + +So this module goes to the store functions directly, never filters on a +category it inferred, and asks the lexical arm for only the two most +distinctive words. + +HOW A MATCH IS RANKED +--------------------- +Both arms run against the brand the label names (query_intent's +`extract_brand_mention`), else every active brand: + +* semantic - `embed_texts([f"{brand} {label}"])` (the stored vector is of + "{brand} {name} {category} {description}", store_catalog_pipeline.py, so + the brand prefix keeps the query in-distribution) -> `semantic_search`; +* lexical - `lexical_search` on the two longest non-brand, non-size words, + which tolerates a misread elsewhere on the label. + +Merged rows are ordered by `text_overlap` FIRST (image_match's scorer: 3 per +shared size token, 1 per shared word - the size is what separates the 89g, +300g and 1kg rows of one product), then by `text_extra` ascending (name +words the label never said - what separates "Dairy Milk" from "Dairy Milk +Fruit & Nut"), then by cosine. `score` on a text row is +`1 - (embedding <=> q)` in MiniLM's 384-d space: a different space from the +image score, and the response says which one it is (`matched_by`). + +A nearest-neighbour query always returns something, so `is_confident()` +draws the line: the best row shares a token with the label, or its cosine +clears IMAGE_IDENTIFY_MIN_TEXT_SCORE. Below that the caller treats the text +rung as a miss. Read-only: nothing here writes. +""" +from __future__ import annotations + +import logging +import re +from typing import Any, Dict, List, Optional, Set, Tuple + +from app.infrastructure.settings import ( + IMAGE_IDENTIFY_MIN_TEXT_SCORE, + IMAGE_SEARCH_DEFAULT_TOP_K, + IMAGE_SEARCH_MAX_FETCH_K, + IMAGE_SEARCH_MAX_TOP_K, + OCR_MAX_CHARS, +) +from app.services.image_match import _STOP, ImageSearchResult, _row_text, text_overlap, tokens + +logger = logging.getLogger(__name__) + +# Words a pack prints that name no product: the nutrition panel, the +# regulatory block, storage advice, the price line. Dropped before either +# arm sees the label so "ENERGY 480 kcal PROTEIN 7g" does not become the +# query. Lowercase; compared after image_match.tokens' own normalisation. +_NOISE: Set[str] = { + "nutrition", "nutritional", "information", "facts", "typical", "values", "value", + "energy", "kcal", "kj", "calories", "protein", "carbohydrate", "carbohydrates", + "carbs", "sugar", "sugars", "fat", "fats", "saturated", "trans", "sodium", + "cholesterol", "fibre", "fiber", "dietary", "serving", "servings", "serve", + "serves", "amount", "ingredients", "ingredient", "contains", "allergen", + "allergens", "allergy", "advice", "mfg", "mfd", "manufactured", "manufacturer", + "marketed", "packed", "packer", "lot", "batch", "best", "before", "expiry", + "exp", "use", "date", "fssai", "lic", "license", "licence", "mrp", "price", + "incl", "inclusive", "all", "taxes", "tax", "ltd", "pvt", "limited", "india", + "www", "com", "http", "https", "customer", "care", "helpline", "email", "call", + "store", "cool", "dry", "place", "keep", "away", "sunlight", "hygienic", + "net", "wt", "weight", "qty", "quantity", "approx", "veg", "vegetarian", + "non", "product", "code", "unit", "units", +} +# A nutrition-panel basis ("per 100 g", "per serve") carries a quantity that +# is NOT the pack size; left in, "100g" would score 3 points of overlap and +# hand the match to the 100g sibling. Removed as a span, before tokenising. +_PER_BASIS_RE = re.compile( + r"\bper\s*\d+(?:\.\d+)?\s*(?:kg|gms|gm|g|ml|ltr|litre|l)\b|\bper\s+serv\w*\b", + re.I, +) +# A nutrient with its value ("Protein 7 g", "Sugars: 20.5g", "Energy 480 +# kcal") is also a quantity that is not the pack size. The name and the +# number go together, so the pair is removed as one span. +_NUTRIENT_VALUE_RE = re.compile( + r"\b(?:energy|calories|protein|carbohydrates?|carbs|sugars?|added\s+sugars?|" + r"total\s+fat|fat|saturated(?:\s+fat)?|trans(?:\s+fat)?|cholesterol|sodium|salt|" + r"dietary\s+fibre|dietary\s+fiber|fibre|fiber|calcium|iron|potassium|vitamin\s*\w*)" + r"\s*[:\-]?\s*(?:<\s*)?\d+(?:\.\d+)?\s*(?:kg|mg|mcg|g|kcal|kj|cal|ml|%)?\b", + re.I, +) +# What is left of a nutrition line once the names are gone: values in units +# no pack is sold in. +_NON_PACK_VALUE_RE = re.compile(r"\b\d+(?:\.\d+)?\s*(?:mg|mcg|kcal|kj|cal)\b|\d+(?:\.\d+)?\s*%", re.I) +_URL_RE = re.compile(r"(?:https?://|www\.)\S+|\b\S+\.(?:com|in|org|net)\b", re.I) +_PRICE_RE = re.compile(r"(?:₹|rs\.?|inr)\s*\d+(?:[.,]\d+)?", re.I) +_SIZE_RE = re.compile(r"^\d+(?:\.\d+)?(?:kg|gms|gm|g|ml|ltr|litre|l|pcs|pc|n)$", re.I) +_WORD_OR_SIZE_RE = re.compile(r"\d+(?:\.\d+)?\s*(?:kg|gms|gm|g|ml|ltr|litre|l|pcs|pc|n)\b|[a-z0-9]+(?:[-'&][a-z0-9]+)*", re.I) + +_LEXICAL_TERMS = 2 +_LONG_NUMBER = 6 # digits; an FSSAI licence is 14, a phone number 10, a barcode 8-13 + + +# --------------------------------------------------------------------------- +# the label +# --------------------------------------------------------------------------- + +def clean_label(text: Optional[str]) -> str: + """The label as a query: boilerplate out, sizes kept, repeats folded. + + Spans first (a per-100g basis, a URL, a price), then tokens: a size such + as "300 g" is kept whole (as "300 g" - image_match.tokens normalises it + later), a word in _NOISE or shorter than two characters is dropped, a + repeat keeps its first position. Capped at OCR_MAX_CHARS on a word + boundary like the OCR output it usually is. + """ + if not text: + return "" + lowered = " ".join(str(text).split()) + lowered = _PER_BASIS_RE.sub(" ", lowered) + lowered = _NUTRIENT_VALUE_RE.sub(" ", lowered) + lowered = _NON_PACK_VALUE_RE.sub(" ", lowered) + lowered = _URL_RE.sub(" ", lowered) + lowered = _PRICE_RE.sub(" ", lowered) + + kept: List[str] = [] + seen: Set[str] = set() + for match in _WORD_OR_SIZE_RE.finditer(lowered): + piece = match.group(0) + key = piece.lower().replace(" ", "") + if _SIZE_RE.match(key): + pass # a size stays, whatever its length + elif len(key) < 2 or key in _NOISE or key in _STOP: + continue + elif key.isdigit() and len(key) >= _LONG_NUMBER: + continue # a licence, phone or barcode number names no product + if key in seen: + continue + seen.add(key) + kept.append(" ".join(piece.split())) + + out = " ".join(kept) + if len(out) > OCR_MAX_CHARS: + cut = out.rfind(" ", 0, OCR_MAX_CHARS + 1) + out = out[:cut if cut > 0 else OCR_MAX_CHARS].rstrip() + return out + + +def _lexical_terms(words: Set[str], brand: Optional[str]) -> List[str]: + """The two longest label words that are not the brand: length is a cheap + proxy for distinctiveness ("britannia" > "gold" > "go"), and two ANDed + terms survive one misread elsewhere on the label.""" + brand_words = set(tokens(brand)[0]) if brand else set() + candidates = sorted((w for w in words if w not in brand_words and not w.isdigit()), + key=lambda w: (-len(w), w)) + return candidates[:_LEXICAL_TERMS] + + +# --------------------------------------------------------------------------- +# ranking +# --------------------------------------------------------------------------- + +def text_extra(words: Set[str], row: Dict[str, Any]) -> int: + """How many words of the row's NAME the label does not account for. + + `text_overlap` is recall - how much of the label the row explains - and + on its own it cannot tell "Dairy Milk 50g" from "Dairy Milk Fruit & Nut + 50g": both explain every word of a "Dairy Milk 50 g" label. The variant + carries words the label never said, and this counts them, so the plain + product outranks its variants and a "sugar free" or "family pack" row + does not win on a label that mentions neither. Only the name and the + sizes are looked at, never the description. + """ + if not words: + return 0 + row_words, _ = tokens(_row_text(row)) + return len(row_words - words) + + +def label_rank_key(row: Dict[str, Any]) -> tuple: + """Best first: what the label says, what the row adds, then the vector. + + Overlap outranks cosine because the size token is the only thing that + tells the pack sizes of one product apart, and their descriptions - and + so their vectors - are near-identical. Unexplained name words come next + for the same reason (a variant's vector is as close as the original's). + """ + return ( + -float(row.get("text_overlap", 0.0)), + int(row.get("text_extra", 0)), + -round(float(row.get("score", 0.0)), 3), + str(row.get("product_name") or row.get("title") or "").lower(), + str(row.get("image_id") or ""), + ) + + +def is_confident(rows: List[Dict[str, Any]]) -> bool: + """Whether the best text row is a find rather than the nearest stranger.""" + if not rows: + return False + best = rows[0] + try: + overlap = float(best.get("text_overlap", 0.0)) + score = float(best.get("score", 0.0)) + except (TypeError, ValueError): + return False + return overlap >= 1.0 or score >= IMAGE_IDENTIFY_MIN_TEXT_SCORE + + +def _normalise_row(row: Dict[str, Any], scoped_table: Optional[str]) -> Dict[str, Any]: + """Give a store row the keys image_match's rows have. + + Brand tables carry no `brand` column, so semantic_search / lexical_search + fill `brand` with the caller's brand when scoped and with the TABLE + SUFFIX ("hindustan_unilever") when not; neither adds `brand_table`. The + identify response is built by the same card builder as the image rows, + so both get the display name and the table here. + """ + from app.services.vector_store import display_name_for_suffix + + row = dict(row) + raw_brand = str(row.get("brand") or "") + if scoped_table: + row["brand_table"] = scoped_table + else: + row["brand_table"] = f"brand_{raw_brand}" if raw_brand else "" + row["brand"] = display_name_for_suffix(raw_brand) if raw_brand else raw_brand + distance = row.get("distance") + try: + row["score"] = 1.0 - float(distance) if distance is not None else 0.0 + except (TypeError, ValueError): + row["score"] = 0.0 + return row + + +def _candidates( + query_text: str, + embedding: Optional[List[float]], + words: Set[str], + brand: Optional[str], + category: Optional[str], + fetch_k: int, +) -> List[Dict[str, Any]]: + """Both arms, merged on (brand_table, image_id); the semantic row wins a + tie because it carries the real distance for the full query.""" + from app.services.vector_store import _table_name, lexical_search, semantic_search + + scoped_table = _table_name(brand) if brand else None + merged: Dict[Tuple[str, str], Dict[str, Any]] = {} + + if embedding is not None: + try: + for row in semantic_search(query_embedding=embedding, brand=brand, + top_k=fetch_k, category=category): + norm = _normalise_row(row, scoped_table) + merged.setdefault((norm["brand_table"], str(norm.get("image_id") or "")), norm) + except Exception as exc: # noqa: BLE001 - one arm failing must not lose the other + logger.warning("label semantic arm failed: %s", exc) + + terms = _lexical_terms(words, brand) + for attempt in (terms, terms[:1]): + if not attempt: + break + try: + rows = lexical_search(attempt, brand=brand, limit=fetch_k, category=category, + query_embedding=embedding, exact_phrase=" ".join(attempt)) + except Exception as exc: # noqa: BLE001 + logger.warning("label lexical arm failed: %s", exc) + rows = [] + for row in rows: + norm = _normalise_row(row, scoped_table) + merged.setdefault((norm["brand_table"], str(norm.get("image_id") or "")), norm) + if rows: + break + return list(merged.values()) + + +# --------------------------------------------------------------------------- +# the search +# --------------------------------------------------------------------------- + +def resolve_label( + text: Optional[str], + brand: Optional[str] = None, + category: Optional[str] = None, + top_k: int = IMAGE_SEARCH_DEFAULT_TOP_K, +) -> ImageSearchResult: + """The catalog rows the label text names, best first, at most `top_k`. + + Never raises: no text, no model and no database all give an empty result. + `query_text` on the result is the cleaned label the search actually used. + """ + cleaned = clean_label(text) + top_k = max(1, min(int(top_k), IMAGE_SEARCH_MAX_TOP_K)) + if not cleaned: + return ImageSearchResult(query_text=None, top_k=top_k) + fetch_k = min(max(top_k * 3, 30), IMAGE_SEARCH_MAX_FETCH_K) + + explicit = (brand or "").strip() or None + detected = explicit + if detected is None: + try: + from app.services.query_intent import extract_brand_mention + detected = extract_brand_mention(cleaned) + except Exception as exc: # noqa: BLE001 - brand detection is an optimisation + logger.debug("brand detection skipped: %s", exc) + detected = None + + query = cleaned + if detected and detected.lower() not in cleaned.lower(): + query = f"{detected} {cleaned}" + + embedding: Optional[List[float]] = None + try: + from app.services.embeddings_service import embed_texts + vectors = embed_texts([query]) + embedding = list(vectors[0]) if vectors else None + except Exception as exc: # noqa: BLE001 - lexical-only, like catalog_search does + logger.warning("label embedding failed: %s. Lexical-only.", exc) + embedding = None + + words, sizes = tokens(cleaned) + + def _run(scope: Optional[str]) -> List[Dict[str, Any]]: + rows = _candidates(query, embedding, words, scope, category, fetch_k) + for row in rows: + row["text_overlap"] = text_overlap(words, sizes, row) + row["text_extra"] = text_extra(words, row) + rows.sort(key=label_rank_key) + return rows + + rows = _run(detected) + scoped = detected is not None + fallback = False + if detected and not explicit and not is_confident(rows): + # The OCR named a brand whose table does not hold this label - a + # misread, or a sub-brand mapped to the wrong parent. Try everywhere. + wider = _run(None) + if is_confident(wider) or not rows: + rows, scoped, fallback = wider, False, True + + return ImageSearchResult( + rows=rows[:top_k], + detected_brand=detected, + scoped_to_brand=scoped, + scope_fallback=fallback, + min_score=0.0, + query_text=cleaned, + top_k=top_k, + ) diff --git a/app/services/ocr_service.py b/app/services/ocr_service.py new file mode 100644 index 0000000..ff9b493 --- /dev/null +++ b/app/services/ocr_service.py @@ -0,0 +1,245 @@ +"""Read the label text off a product photo, server-side: rapidocr on onnxruntime. + +WHY THIS EXISTS +--------------- +POST /api/search/identify falls through to the label text when the image +vector cannot confirm a product (a phone photo of a pack scores ~0.63 against +the catalog's render of it). The Nearle app OCRs on-device and sends `text`; +a plain camera upload has no such text, so this module reads it here. + +WHAT IT PRODUCES +---------------- +One string: the recognised lines in detection order (top-to-bottom, which is +reading order for a pack front), each kept only above OCR_MIN_CONFIDENCE, +exact repeats dropped, capped at OCR_MAX_CHARS on a word boundary. It is +deliberately the same shape as the `text` field the routes already accept, so +app/services/label_match.py cannot tell the two apart. + +THE ENGINE +---------- +rapidocr (PP-OCR detection + classification + recognition, ONNX). The models +ship inside the wheel and are verified by checksum at construction, so the +first call makes no network request. Two things to know about it: + +* Its metadata demands opencv-python (the GUI build), which would unpack over + this project's opencv-python-headless and fail on libGL in the slim image - + breaking image_embedder too. It is therefore installed with `--no-deps` + (requirements-ocr.txt) and its real dependencies are listed in + requirements.txt. Nothing here needs the GUI build. +* `RapidOCR.__call__` mutates the engine's own state, so one lock serialises + every read, exactly as image_embedder does for its interpreter. + +FAILURE IS SILENT BY DESIGN +--------------------------- +Imported lazily, on the first photo that needs it. A missing wheel, a missing +onnxruntime, a broken model: `available()` is False after ONE warning, +`read_text()` is None, and the identify route reports `ocr_unavailable` +instead of failing. GET /api/health -> ocr says which of those it was. +""" +from __future__ import annotations + +import logging +import threading +import warnings +from typing import Any, List, Optional + +import numpy as np + +from app.infrastructure.settings import ( + ENABLE_SERVER_OCR, + OCR_MAX_CHARS, + OCR_MAX_SIDE_PX, + OCR_MIN_CONFIDENCE, + OCR_NUM_THREADS, +) +from app.services import image_embedder + +logger = logging.getLogger(__name__) + +_lock = threading.Lock() +_engine: Any = None +_disabled_reason: Optional[str] = None +_warned = False + + +# --------------------------------------------------------------------------- +# the engine +# --------------------------------------------------------------------------- + +def _disable(reason: str, *, warn: bool = True) -> None: + """Record why OCR is off; say so once. Caller holds _lock.""" + global _disabled_reason, _engine, _warned + _disabled_reason = reason + _engine = None + if warn and not _warned: + _warned = True + logger.warning("server OCR disabled: %s (identify falls back to client text)", reason) + + +def _load() -> None: + """Construct the engine. Caller holds _lock. Never raises.""" + global _engine + if not ENABLE_SERVER_OCR: + # A deliberate setting, not a fault: no warning line for it. + _disable("ENABLE_SERVER_OCR is false", warn=False) + return + try: + with warnings.catch_warnings(): + # Import-time deprecation chatter from onnxruntime / omegaconf / + # shapely would be a hard error under pytest's filterwarnings=error. + warnings.simplefilter("ignore") + from rapidocr import RapidOCR + + engine = RapidOCR(params={ + "Global.text_score": float(OCR_MIN_CONFIDENCE), + # Its logger does not propagate and installs its own INFO + # handler; keep the per-call chatter out of the container log. + "Global.log_level": "warning", + "EngineConfig.onnxruntime.intra_op_num_threads": max(1, int(OCR_NUM_THREADS)), + }) + except Exception as exc: # noqa: BLE001 - ImportError, a missing engine, a bad model, libGL + _disable(f"rapidocr is not usable ({exc})") + return + _engine = engine + logger.info("server OCR ready: rapidocr (%d thread(s))", OCR_NUM_THREADS) + + +def _ensure_loaded() -> bool: + """Caller holds _lock.""" + if _engine is None and _disabled_reason is None: + _load() + return _engine is not None + + +def available() -> bool: + """True when the engine can be used. Loads it on first call.""" + with _lock: + return _ensure_loaded() + + +def status() -> dict: + """Diagnostics for /api/health - NEVER loads the engine. + + Same idea as image_embedder.status(): the failure this module absorbs is + one log line at first use, so the same facts go on the health endpoint + where "why does identify say ocr_unavailable?" is one curl away. + """ + try: + import importlib.util + runtime = importlib.util.find_spec("rapidocr") is not None + onnx = importlib.util.find_spec("onnxruntime") is not None + except Exception: # noqa: BLE001 - a broken finder counts as "not importable" + runtime = onnx = False + with _lock: + if _engine is not None: + state = "ready" + elif _disabled_reason: + state = f"disabled: {_disabled_reason}" + elif not ENABLE_SERVER_OCR: + state = "disabled: ENABLE_SERVER_OCR is false" + else: + state = "not loaded yet (loads on first identify)" + return { + "enabled": bool(ENABLE_SERVER_OCR), + "runtime_importable": runtime, + "onnxruntime_importable": onnx, + "state": state, + } + + +# --------------------------------------------------------------------------- +# the read +# --------------------------------------------------------------------------- + +def _downscale(bgr: np.ndarray, max_side: int) -> np.ndarray: + """`bgr` with its longer side at most `max_side` (INTER_AREA), else as is.""" + h, w = bgr.shape[:2] + longest = max(h, w) + if max_side <= 0 or longest <= max_side: + return bgr + import cv2 + + scale = max_side / float(longest) + size = (max(1, int(round(w * scale))), max(1, int(round(h * scale)))) + return cv2.resize(bgr, size, interpolation=cv2.INTER_AREA) + + +def _join_lines(output: Any) -> Optional[str]: + """The engine's lines as one label string, or None when it read nothing. + + Detection order is kept (top-to-bottom is reading order on a pack front); + a line below OCR_MIN_CONFIDENCE is dropped (the engine already filters on + Global.text_score - this is belt and braces for a fake or a future + engine); an exact repeat is dropped (a pack often prints its name twice); + the whole is capped at OCR_MAX_CHARS on a word boundary so the server's + text is bounded exactly like the routes' `text` field. + """ + if output is None: + return None + txts = getattr(output, "txts", None) + scores = getattr(output, "scores", None) + if not txts: + return None + if scores is None or len(scores) != len(txts): + scores = [1.0] * len(txts) + + lines: List[str] = [] + seen = set() + for raw, score in zip(txts, scores): + try: + if float(score) < OCR_MIN_CONFIDENCE: + continue + except (TypeError, ValueError): + continue + line = " ".join(str(raw or "").split()) + if not line: + continue + key = line.lower() + if key in seen: + continue + seen.add(key) + lines.append(line) + + text = " ".join(lines) + if len(text) > OCR_MAX_CHARS: + cut = text.rfind(" ", 0, OCR_MAX_CHARS + 1) + text = text[:cut if cut > 0 else OCR_MAX_CHARS].rstrip() + return text or None + + +def read_text(image_bytes: bytes) -> Optional[str]: + """The label text in `image_bytes`, or None (undecodable, nothing read, no engine). + + The frame goes in as a BGR array - decoded by image_embedder.decode_bgr so + the pixel cap, EXIF orientation and alpha handling are the same as the + vector path's - and downscaled to OCR_MAX_SIDE_PX first. It is never + handed to the engine as raw bytes: that would bypass both guards. + """ + bgr = image_embedder.decode_bgr(image_bytes) + if bgr is None: + return None + try: + bgr = _downscale(bgr, OCR_MAX_SIDE_PX) + except Exception as exc: # noqa: BLE001 - a resize failure is "nothing read" + logger.debug("ocr downscale failed: %s", exc) + return None + with _lock: + if not _ensure_loaded(): + return None + try: + with warnings.catch_warnings(): + warnings.simplefilter("ignore") + output = _engine(bgr) + except Exception as exc: # noqa: BLE001 - one bad photo must not poison the engine + logger.debug("ocr failed: %s", exc) + return None + return _join_lines(output) + + +def _reset() -> None: + """Tests only: forget the engine and the disabled state.""" + global _engine, _disabled_reason, _warned + with _lock: + _engine = None + _disabled_reason = None + _warned = False diff --git a/app/services/product_identify.py b/app/services/product_identify.py new file mode 100644 index 0000000..b2b57a1 --- /dev/null +++ b/app/services/product_identify.py @@ -0,0 +1,162 @@ +"""Identify the product in a phone photo: image vector first, label text second. + +THE LADDER +---------- + photo ─► img_vector kNN (image_match.search_by_vector) + │ best score >= IMAGE_IDENTIFY_MIN_IMAGE_SCORE + │ -> matched_by="image_vector", fallback_reason=None + ▼ else + label text: the client's `text` (ocr_source="client"), + else what ocr_service reads off the photo ("server") + ▼ + label_match.resolve_label + │ confident -> matched_by="text", fallback_reason= + ▼ else + the image rows as they were (they may still be right - the 0.63 + in the feasibility test WAS the product), matched_by= + "image_vector" or "none", fallback_reason="text_no_match" + +`fallback_reason` is the reason for the LAST step the ladder took: + + image_below_threshold the image arm ran but did not clear the floor + no_image_match the image arm ran and found nothing at all + image_embedder_unavailable no vector could be computed (no model here) + no_text no label to fall back to and no photo to read + ocr_unavailable no client text, and server OCR is off / missing + ocr_empty server OCR ran and read nothing usable + text_no_match the label was resolved but nothing was confident + +The rule for a client: the answer is confirmed when `matched_by == "text"`, +or `matched_by == "image_vector"` with no `fallback_reason`. Everything else +is a best effort the client should show as such. + +The two `score`s are in different spaces (1024-d image, 384-d text); +`image_top_score` always carries the image side so the client can see both. +Server OCR is invoked only when it is needed - never when the image is +confident and never when the client sent text - because it costs a second. +Read-only: every function this calls only reads. +""" +from __future__ import annotations + +import logging +from dataclasses import dataclass +from typing import Optional, Sequence + +from app.infrastructure.settings import ( + IMAGE_IDENTIFY_MIN_IMAGE_SCORE, + IMAGE_SEARCH_DEFAULT_MIN_SCORE, + IMAGE_SEARCH_DEFAULT_TOP_K, +) +from app.services import ocr_service +from app.services.image_match import ImageSearchResult, search_by_vector +from app.services.label_match import is_confident, resolve_label + +logger = logging.getLogger(__name__) + +MATCHED_BY_IMAGE = "image_vector" +MATCHED_BY_TEXT = "text" +MATCHED_BY_NONE = "none" + + +@dataclass +class IdentifyResult: + search: ImageSearchResult + matched_by: str = MATCHED_BY_NONE + ocr_text: Optional[str] = None + ocr_source: Optional[str] = None # "client" | "server" + image_top_score: Optional[float] = None + fallback_reason: Optional[str] = None + + +def _image_rung( + vector: Optional[Sequence[float]], + text: Optional[str], + brand: Optional[str], + category: Optional[str], + top_k: int, + min_score: float, + min_image_score: float, +) -> tuple[ImageSearchResult, Optional[float], Optional[str]]: + """(image result, its best score, why it is not the answer - or None).""" + if vector is None: + empty = ImageSearchResult(min_score=min_score, top_k=top_k, query_text=(text or "").strip() or None) + return empty, None, "image_embedder_unavailable" + result = search_by_vector(vector, text=text, brand=brand, category=category, + top_k=top_k, min_score=min_score) + if not result.rows: + return result, None, "no_image_match" + try: + top = float(result.rows[0].get("score", 0.0)) + except (TypeError, ValueError): + top = 0.0 + if top >= min_image_score: + return result, top, None + return result, top, "image_below_threshold" + + +def identify_product( + *, + vector: Optional[Sequence[float]], + image_bytes: Optional[bytes], + text: Optional[str] = None, + brand: Optional[str] = None, + category: Optional[str] = None, + top_k: int = IMAGE_SEARCH_DEFAULT_TOP_K, + min_score: float = IMAGE_SEARCH_DEFAULT_MIN_SCORE, + min_image_score: float = IMAGE_IDENTIFY_MIN_IMAGE_SCORE, +) -> IdentifyResult: + """Run the ladder. `vector` is the photo's embedding (None when this + deployment cannot compute one); `image_bytes` the photo for server OCR + (None on the vector-only route). Raises only InvalidVectorError, from + search_by_vector, for a vector that cannot be searched with. + """ + label = (text or "").strip() or None + + image, top, reason = _image_rung(vector, label, brand, category, top_k, min_score, min_image_score) + if reason is None: + return IdentifyResult(search=image, matched_by=MATCHED_BY_IMAGE, image_top_score=top, + ocr_text=label, ocr_source="client" if label else None) + + # --- the label ---------------------------------------------------------- + ocr_source: Optional[str] = None + if label: + ocr_source = "client" + elif image_bytes: + if ocr_service.available(): + label = ocr_service.read_text(image_bytes) + if label: + ocr_source = "server" + else: + reason = "ocr_empty" + else: + reason = "ocr_unavailable" + else: + reason = "no_text" + + if not label: + return IdentifyResult( + search=image, + matched_by=MATCHED_BY_IMAGE if image.rows else MATCHED_BY_NONE, + image_top_score=top, + fallback_reason=reason, + ) + + resolved = resolve_label(label, brand=brand, category=category, top_k=top_k) + if is_confident(resolved.rows): + return IdentifyResult( + search=resolved, + matched_by=MATCHED_BY_TEXT, + ocr_text=label, + ocr_source=ocr_source, + image_top_score=top, + fallback_reason=reason, # why the image rung was not the answer + ) + + return IdentifyResult( + search=image, + matched_by=MATCHED_BY_IMAGE if image.rows else MATCHED_BY_NONE, + ocr_text=label, + ocr_source=ocr_source, + image_top_score=top, + fallback_reason="text_no_match", + ) diff --git a/docs/IMAGE_SEARCH_API.md b/docs/IMAGE_SEARCH_API.md index 545702a..1954e5b 100644 --- a/docs/IMAGE_SEARCH_API.md +++ b/docs/IMAGE_SEARCH_API.md @@ -10,13 +10,16 @@ index). These two endpoints turn the app's vector - or a photo - into product cards. ``` -POST /api/search/image-vector JSON {vector[1024], text?, brand?, category?, top_k?, min_score?} +POST /api/search/image-vector JSON {vector[1024], text?, brand?, category?, top_k?, min_score?, text_fallback?} GET /api/search/image-vector query string: vector=&text=...&top_k=... (see "GET variant") POST /api/search/image multipart file + the same optional fields as form fields +POST /api/search/identify multipart file + the same fields; image first, label text second (see "Identify") ``` -Both run the same ranking. Use the first from the app (it already has the -vector); use the second from anything without the model, or for testing. +The first three run the same ranking. Use the first from the app (it already +has the vector); use `/image` from anything without the model, or for +testing. `/identify` is for a phone photo that the vector alone cannot +confirm - see the section below. ## How a match is found @@ -159,14 +162,97 @@ final b64 = base64Url.encode(vector.buffer.asUint8List()).replaceAll('=', ''); A malformed `vector` (wrong count, not a number, bad base64, all zeros) is a 422 whose `detail` says which value or what length was wrong. +## Identify — image first, label text second + +`POST /api/search/identify` exists because a **phone photo of a pack scores +only ~0.63 against the catalogue's render of the same pack** (measured), so +`img_vector` alone cannot confirm which product it is. This route runs a +ladder and tells you which rung answered: + +``` +photo ─► img_vector search ─► best score ≥ IMAGE_IDENTIFY_MIN_IMAGE_SCORE (0.70)? + │ yes → matched_by "image_vector", fallback_reason null (confirmed) + ▼ no + label text = your `text` (ocr_source "client") + else the server reads it off the photo (ocr_source "server") + ▼ + resolve the label: brand from the text, MiniLM embedding vs the rows' + `embedding` + product-name match, ranked by shared size/word tokens + first, cosine second + │ confident → matched_by "text" (confirmed) + ▼ not → the low image rows, if any, with fallback_reason "text_no_match" +``` + +Same multipart fields as `/image`. `text` is optional: send it when your +client already OCR'd the label (it is also used to scope and tie-break the +image search); leave it out and the server reads the label itself. + +```bash +curl -s -X POST https://mcp.nearle.ai.in/api/search/identify \ + -F "file=@marie_gold_phone.jpg" -F top_k=5 +``` + +```jsonc +{ + "matched_by": "text", // "image_vector" | "text" | "none" + "fallback_reason": "image_below_threshold", // why the image rung was not the answer + "ocr_source": "server", // "client" | "server" | null + "ocr_text": "Britannia MARIE GOLD Original Tea Time Biscuit 300 g", // the cleaned label used + "image_top_score": 0.6312, // the image side, always reported + "detected_brand": "britannia", "scoped_to_brand": true, "scope_fallback": false, + "results": [ { "product_name": "Britannia Marie Gold 300g", "score": 0.71, "text_overlap": 5.0, ... } ], + "total": 5, "top_k": 5, "min_score": 0.0, "query_text": "Britannia MARIE GOLD ... 300 g" +} +``` + +**Reading the answer.** It is confirmed when `matched_by` is `"text"`, or +`"image_vector"` with `fallback_reason` null. Anything else is a best effort: +show it as such, or ask for another shot. + +| `fallback_reason` | Meaning | +|---|---| +| `null` | the image match cleared the floor | +| `image_below_threshold` | image ran, best score under the floor → the label decided (on a `"text"` answer) | +| `no_image_match` | image ran and found nothing → the label decided | +| `image_embedder_unavailable` | this deployment has no image model → text only | +| `no_text` | `/image-vector` with `text_fallback` but no `text` | +| `ocr_unavailable` | no `text` sent and server OCR is off / not installed (`GET /api/health` → `ocr`) | +| `ocr_empty` | server OCR ran and read nothing usable | +| `text_no_match` | the label was resolved but nothing was confident; the image rows are returned as they were | + +**`score` is not one number.** On an `"image_vector"` answer it is the +1024-d photo cosine as on `/image`. On a `"text"` answer it is +`1 - (embedding <=> q)` in the 384-d MiniLM space - a different scale, not +comparable to the first. `image_top_score` always carries the image side so +you can see both. On text answers `text_overlap` (3 per shared pack-size +token, 1 per shared word) is what ranked the rows, then - among rows that +explain the label equally - the fewest name words the label never said, so +"Dairy Milk 50g" outranks "Dairy Milk Fruit & Nut 50g" on a plain label and +the reverse on a label that says "Fruit & Nut"; cosine only breaks what is +left. + +**Server OCR** is rapidocr (PP-OCR models on onnxruntime, CPU). It is loaded +on the first photo that needs it (~2 s), then costs ~1–2 s per photo at the +1280 px it downscales to; it runs only when the image rung failed *and* no +`text` was sent. `GET /api/health` → `ocr` reports whether it is installed +and loaded. `ENABLE_SERVER_OCR=false` turns it off; the route then relies on +client `text`. + +**The app's route.** `POST /api/search/image-vector` accepts +`"text_fallback": true`. Off (the default) the response is exactly what it +has always been. On, the same ladder runs with the app's vector and `text` +(no photo, so never server OCR), and the response is the identify shape +above. Not on the GET variant. + ## Errors | Code | Cause | What to do | |---|---|---| -| 400 | `/image`: the upload is empty | send the file | -| 413 | `/image`: file over 8 MB | crop or downscale | +| 400 | `/image`, `/identify`: the upload is empty | send the file | +| 413 | `/image`, `/identify`: file over 8 MB | crop or downscale | | 422 | wrong vector length, NaN, all zeros, unparseable GET `vector`, `top_k` out of 1–50, `min_score` out of −1…1, undecodable image | `detail` names the field | | 503 | `/image`: this deployment has no embedding model | embed client-side and use `/image-vector`; `GET /api/health` → `image_vectors.model_present` says whether this can happen | +| 503 | `/identify`: no embedding model AND no server OCR AND no `text` | send `text`, or a vector to `/image-vector`; `GET /api/health` → `image_vectors`, `ocr` | ## Good to know @@ -183,4 +269,8 @@ is a 422 whose `detail` says which value or what length was wrong. Implementation: `app/services/image_match.py` (ranking), `vector_store.image_vector_search` / `fetch_products_by_image_ids` (reads), `app/api/routers/search.py` (routes), -`tests/test_image_match.py`, `tests/test_image_search_api.py`. +`tests/test_image_match.py`, `tests/test_image_search_api.py`. Identify: +`app/services/product_identify.py` (the ladder), `app/services/label_match.py` +(label → rows), `app/services/ocr_service.py` (server OCR; `requirements-ocr.txt`), +`tests/test_product_identify.py`, `tests/test_label_match.py`, +`tests/test_ocr_service.py`, `tests/test_identify_api.py`. diff --git a/requirements-ocr.txt b/requirements-ocr.txt new file mode 100644 index 0000000..c2825c1 --- /dev/null +++ b/requirements-ocr.txt @@ -0,0 +1,17 @@ +# Server-side OCR engine for POST /api/search/identify (app/services/ocr_service.py). +# +# Installed SEPARATELY and WITHOUT its declared dependencies: +# +# pip install --no-deps -r requirements-ocr.txt +# +# because rapidocr's metadata demands opencv-python (the GUI build). A normal +# install would unpack it over the opencv-python-headless this project pins, +# and on the slim container image the GUI build fails to import (no libGL) - +# which would take the img_vector embedder down with it. Everything rapidocr +# actually needs (onnxruntime, Shapely, pyclipper, omegaconf, ...) is declared +# in requirements.txt and installed the normal way first. +# +# Pinned exactly: the PP-OCR models ship inside the wheel (~30MB, verified by +# checksum at construction, no download at runtime), so the version IS the +# model version. Bump deliberately, and re-run tests/test_ocr_engine_real.py. +rapidocr==3.9.2 diff --git a/requirements.txt b/requirements.txt index 67352b6..fdab8e2 100644 --- a/requirements.txt +++ b/requirements.txt @@ -74,6 +74,26 @@ ai-edge-litert>=2.2.0 # backend preprocesses with the same library. Headless: no GUI, no libGL. # abi3 wheel, ~55MB. Pinned below 5 to stay on the app's major version. opencv-python-headless>=4.8,<5 +# Server-side OCR for POST /api/search/identify (app/services/ocr_service.py): +# rapidocr's PP-OCR models on onnxruntime, CPU. rapidocr ITSELF IS NOT LISTED +# HERE, on purpose - its metadata demands opencv-python (the GUI build), which +# a normal install unpacks over opencv-python-headless above and which then +# fails on libGL.so.1 in the slim image, breaking the image embedder too. It is +# installed with `--no-deps` from requirements-ocr.txt (see the Dockerfile), +# and the dependencies it really needs are declared below instead: +# - onnxruntime: the inference engine, which rapidocr does not declare (it is +# one of several it can use). Imported lazily; missing means "no server +# OCR", never a failed boot. +# - the rest are pure-Python helpers from its metadata (Shapely and pyclipper +# for the detection polygons, omegaconf/PyYAML for its config). +onnxruntime>=1.17,<2 +omegaconf>=2.1,!=2.2.1 +pyclipper>=1.2.0 +Shapely>=1.7.1,!=2.0.4 +PyYAML>=6.0 +tqdm>=4.60 +colorlog>=6.0 +six>=1.15.0 aiofiles>=23.2.1 # Playwright (Python) - last-resort image-search fallback only. diff --git a/tests/conftest.py b/tests/conftest.py index 2cb09cf..43662cb 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -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" diff --git a/tests/test_identify_api.py b/tests/test_identify_api.py new file mode 100644 index 0000000..445c98d --- /dev/null +++ b/tests/test_identify_api.py @@ -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} diff --git a/tests/test_image_search_api.py b/tests/test_image_search_api.py index 462c530..baec36a 100644 --- a/tests/test_image_search_api.py +++ b/tests/test_image_search_api.py @@ -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"] diff --git a/tests/test_label_match.py b/tests/test_label_match.py new file mode 100644 index 0000000..4788eda --- /dev/null +++ b/tests/test_label_match.py @@ -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 diff --git a/tests/test_ocr_engine_real.py b/tests/test_ocr_engine_real.py new file mode 100644 index 0000000..a1cc6ba --- /dev/null +++ b/tests/test_ocr_engine_real.py @@ -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 diff --git a/tests/test_ocr_service.py b/tests/test_ocr_service.py new file mode 100644 index 0000000..810975e --- /dev/null +++ b/tests/test_ocr_service.py @@ -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 == [] diff --git a/tests/test_product_identify.py b/tests/test_product_identify.py new file mode 100644 index 0000000..11836c8 --- /dev/null +++ b/tests/test_product_identify.py @@ -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")