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