diff --git a/app/api/routers/health.py b/app/api/routers/health.py index b8bd80f..a832066 100644 --- a/app/api/routers/health.py +++ b/app/api/routers/health.py @@ -5,9 +5,12 @@ import logging import requests from fastapi import APIRouter -from app.api.schemas import AuthConfigOut, HealthOut +from app.api.schemas import AuthConfigOut, HealthOut, ImageVectorsOut from app.infrastructure.security import auth_config_summary -from app.infrastructure.settings import OLLAMA_BASE_URL, OLLAMA_MODEL_NAME, EMBEDDINGS_MODEL +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.vector_store import _connect # internal, but handy for a connectivity probe logger = logging.getLogger(__name__) @@ -56,4 +59,8 @@ def health() -> HealthOut: ollama_model=OLLAMA_MODEL_NAME, embeddings_model=EMBEDDINGS_MODEL, auth=AuthConfigOut(**auth_config_summary()), + # Same idea as `auth`: img_vector staying NULL after a deploy has one + # 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()), ) diff --git a/app/api/routers/search.py b/app/api/routers/search.py index 714b835..b6ae979 100644 --- a/app/api/routers/search.py +++ b/app/api/routers/search.py @@ -2,11 +2,28 @@ from __future__ import annotations from typing import Optional -from fastapi import APIRouter, Query +from fastapi import APIRouter, File, Form, HTTPException, Query, UploadFile +from starlette.concurrency import run_in_threadpool -from app.api.schemas import SearchOut, SourceProductOut -from app.infrastructure.settings import SEARCH_DEFAULT_TOP_K, SEARCH_MAX_TOP_K +from app.api.routers.brands import _row_to_product_out +from app.api.schemas import ( + ImageMatchOut, + ImageSearchOut, + ImageVectorSearchRequest, + SearchOut, + SourceProductOut, +) +from app.infrastructure.settings import ( + IMAGE_SEARCH_DEFAULT_MIN_SCORE, + IMAGE_SEARCH_DEFAULT_TOP_K, + IMAGE_SEARCH_MAX_TOP_K, + IMAGE_VECTOR_MAX_BYTES, + SEARCH_DEFAULT_TOP_K, + SEARCH_MAX_TOP_K, +) +from app.services import image_embedder from app.services.catalog_search import search_catalog +from app.services.image_match import ImageSearchResult, InvalidVectorError, search_by_vector router = APIRouter(tags=["search"]) @@ -46,3 +63,100 @@ def catalog_search_endpoint( detected_brand=result.detected_brand, detected_category=result.detected_category, ) + + +# --------------------------------------------------------------------------- +# Search by image +# --------------------------------------------------------------------------- +# Public like GET /search. Two ways in, one ranking: the Nearle app embeds +# the cropped photo on-device with the same MobileNetV3 model that filled +# img_vector and POSTs the 1024 floats; anything without the model POSTs the +# photo and this API embeds it (bounded: IMAGE_VECTOR_MAX_BYTES, one +# inference at a time behind the embedder's lock). + +def _to_image_search_out(result: ImageSearchResult) -> ImageSearchOut: + matches = [] + for row in result.rows: + card = _row_to_product_out(row, row.get("brand") or "") + matches.append(ImageMatchOut( + **card.model_dump(), + score=round(float(row["score"]), 4), + text_overlap=float(row.get("text_overlap", 0.0)), + )) + return ImageSearchOut( + results=matches, + total=len(matches), + detected_brand=result.detected_brand, + scoped_to_brand=result.scoped_to_brand, + scope_fallback=result.scope_fallback, + min_score=result.min_score, + top_k=result.top_k, + query_text=result.query_text, + ) + + +@router.post("/search/image-vector", response_model=ImageSearchOut) +def image_vector_search_endpoint(body: ImageVectorSearchRequest) -> ImageSearchOut: + """Products that look like the photo whose embedding is `vector`. + + `vector` is the 1024-float, L2-normalised MobileNetV3-Small embedding the + app computes on-device. `score` on each result is cosine similarity + (1 - pgvector distance). Optional `text` - the OCR read of the label - + narrows the search to the brand it names and picks the right pack size + 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`). + """ + try: + 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, + ) + except InvalidVectorError as exc: + raise HTTPException(status_code=422, detail=str(exc)) + return _to_image_search_out(result) + + +@router.post("/search/image", response_model=ImageSearchOut) +async def image_search_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"), + 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), +) -> ImageSearchOut: + """Same as /search/image-vector, but the API embeds the photo itself. + + 503 when this deployment has no embedding model (GET /api/health -> + image_vectors.model_present says so); send a vector instead. + """ + 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.", + ) + if not await run_in_threadpool(image_embedder.available): + raise HTTPException( + status_code=503, + detail="The image embedding model is not available on this deployment. " + "Embed the photo client-side and POST the vector to /api/search/image-vector.", + ) + 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( + search_by_vector, vector, 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)) + return _to_image_search_out(result) diff --git a/app/api/schemas.py b/app/api/schemas.py index 355dd37..bd207f8 100644 --- a/app/api/schemas.py +++ b/app/api/schemas.py @@ -1,9 +1,16 @@ """Pydantic request/response models for the FastAPI layer.""" from __future__ import annotations +import math from typing import List, Optional -from pydantic import BaseModel, Field +from pydantic import BaseModel, Field, field_validator + +from app.infrastructure.settings import ( + IMAGE_SEARCH_DEFAULT_MIN_SCORE, + IMAGE_SEARCH_DEFAULT_TOP_K, + IMAGE_SEARCH_MAX_TOP_K, +) # --------------------------------------------------------------------------- @@ -120,6 +127,17 @@ class AuthConfigOut(BaseModel): api_keys_source: str = "default" +class ImageVectorsOut(BaseModel): + """Why img_vector is (or is not) being filled. Reported without loading + the model. `model_present=false` after a deploy means the .tflite was not + shipped in the image - the one failure this feature absorbs silently.""" + enabled: bool = True + model_path: str = "" + model_present: bool = False + runtime_importable: bool = False + state: str = "unknown" + + class HealthOut(BaseModel): status: str database: bool @@ -127,6 +145,9 @@ class HealthOut(BaseModel): ollama_model: str embeddings_model: str auth: AuthConfigOut + # Defaulted so a client of this schema still validates against a + # deployment predating the field. + image_vectors: ImageVectorsOut = Field(default_factory=ImageVectorsOut) # --------------------------------------------------------------------------- @@ -214,6 +235,52 @@ class SuggestOut(BaseModel): suggestions: List[SuggestionOut] +# --------------------------------------------------------------------------- +# Image search (POST /api/search/image-vector, POST /api/search/image) +# --------------------------------------------------------------------------- + +class ImageVectorSearchRequest(BaseModel): + """A phone photo's embedding, as the Nearle app computes it on-device.""" + vector: List[float] = Field( + ..., min_length=1024, max_length=1024, + description="L2-normalised MobileNetV3-Small embedding, 1024 floats", + ) + text: Optional[str] = Field(None, max_length=500, description="OCR text read off the label") + brand: Optional[str] = Field(None, max_length=120, description="Restrict to one brand (no fallback)") + category: Optional[str] = Field(None, max_length=120) + 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") + + @field_validator("vector") + @classmethod + def _finite_and_nonzero(cls, v: List[float]) -> List[float]: + if not all(math.isfinite(x) for x in v): + raise ValueError("vector contains NaN or infinite values") + if math.sqrt(sum(x * x for x in v)) < 1e-6: + raise ValueError("vector is all zeros") + return v + + +class ImageMatchOut(ProductOut): + """One catalog product that looks like the photo: the product card plus + how close it is. `score` is cosine similarity (1 - pgvector distance); + `text_overlap` is the label-text tie-break weight, 0 when no text was sent.""" + score: float + text_overlap: float = 0.0 + + +class ImageSearchOut(BaseModel): + results: List[ImageMatchOut] + total: int + detected_brand: Optional[str] = None + scoped_to_brand: bool = False + scope_fallback: bool = False + min_score: float = 0.0 + top_k: int = 0 + query_text: Optional[str] = None + + # --------------------------------------------------------------------------- # RAG chat # --------------------------------------------------------------------------- diff --git a/app/infrastructure/settings.py b/app/infrastructure/settings.py index 7abcfbf..22a2c37 100644 --- a/app/infrastructure/settings.py +++ b/app/infrastructure/settings.py @@ -417,6 +417,22 @@ IMAGE_EMBED_MODEL_PATH = _dir( # serialised by a lock (a TFLite interpreter is not thread-safe). IMAGE_EMBED_NUM_THREADS = int(os.getenv("IMAGE_EMBED_NUM_THREADS", "2")) +# --------------------------------------------------------------------------- +# Image search - POST /api/search/image-vector and /api/search/image +# (app/services/image_search.py). Public, read-only. +# --------------------------------------------------------------------------- +IMAGE_SEARCH_DEFAULT_TOP_K = int(os.getenv("IMAGE_SEARCH_DEFAULT_TOP_K", "10")) +IMAGE_SEARCH_MAX_TOP_K = int(os.getenv("IMAGE_SEARCH_MAX_TOP_K", "50")) +# 0.0 on purpose. A simulated phone photo of Marie Gold against its catalog +# render scored 0.63; the app team's "0.7 means the same product" is a +# client-side rule of thumb for phone-vs-phone, so the server does not +# impose it - callers pass min_score when they want a floor. +IMAGE_SEARCH_DEFAULT_MIN_SCORE = float(os.getenv("IMAGE_SEARCH_DEFAULT_MIN_SCORE", "0.0")) +# Candidates fetched PER brand table before re-ranking (pack sizes of one +# product share an image and tie, so more than top_k must come back), and +# 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")) + # --------------------------------------------------------------------------- # USDA FoodData Central - nutrition for loose, unbranded commodities # --------------------------------------------------------------------------- diff --git a/app/main.py b/app/main.py index 063d298..2a247ec 100644 --- a/app/main.py +++ b/app/main.py @@ -14,10 +14,12 @@ import threading from contextlib import asynccontextmanager from pathlib import Path -from fastapi import FastAPI, HTTPException +from fastapi import FastAPI, HTTPException, Request +from fastapi.exceptions import RequestValidationError +from fastapi.encoders import jsonable_encoder from fastapi.middleware.cors import CORSMiddleware from fastapi.staticfiles import StaticFiles -from fastapi.responses import FileResponse +from fastapi.responses import FileResponse, JSONResponse from app.infrastructure.persistence import restore_bundled_assets from app.infrastructure.settings import ( @@ -213,6 +215,27 @@ app.add_middleware( allow_headers=["*"], ) + +def _json_safe(value): + """Replace floats JSON cannot carry (NaN, +/-inf) so an error can be sent.""" + if isinstance(value, float) and (value != value or value in (float("inf"), float("-inf"))): + return str(value) + if isinstance(value, dict): + return {k: _json_safe(v) for k, v in value.items()} + if isinstance(value, (list, tuple)): + return [_json_safe(v) for v in value] + return value + + +@app.exception_handler(RequestValidationError) +async def _validation_error_as_422(request: Request, exc: RequestValidationError) -> JSONResponse: + """FastAPI's own 422 body echoes the rejected input. Python's JSON parser + accepts `NaN` on the way in, pydantic rejects it, and the echo then fails + to serialise - so a client that sent one NaN in a float field got a 500 + instead of the 422 that names the field. Found by the image-vector + search, where the body is 1024 floats; applies to every route.""" + return JSONResponse(status_code=422, content={"detail": _json_safe(jsonable_encoder(exc.errors()))}) + # A wrong origin list fails only in the browser, as an opaque "blocked by CORS" # with a perfectly healthy 200 in the server log - so state the effective list # at startup, where it can actually be compared against the frontend's URL. diff --git a/app/services/image_embedder.py b/app/services/image_embedder.py index 1c12fab..d9e8e30 100644 --- a/app/services/image_embedder.py +++ b/app/services/image_embedder.py @@ -13,7 +13,9 @@ THIS FILE OWNS TWO THINGS AND NOTHING ELSE ------------------------------------------ 1. `preprocess()` - bytes to the input tensor. It is the ONLY place the resize / crop / scaling decisions live, because those decisions are what - make the vectors comparable. See the banner on that function. + make the vectors comparable. It is a port of the Nearle Flutter app's + OpenCV pipeline (centre crop, INTER_AREA to 224, RGB, 0..1) and uses + OpenCV itself so the two agree to a few decimals. 2. The interpreter - one per process, created lazily on first use, and every call into it serialised by one lock. A TFLite interpreter is not thread-safe, and this process has exactly two callers: the single @@ -65,50 +67,81 @@ _warned = False # bytes -> input tensor # --------------------------------------------------------------------------- +def _flatten_alpha_to_bgr(image_bytes: bytes) -> Optional[np.ndarray]: + """Pillow decode for images WITH transparency: EXIF applied, composited on + white, returned as uint8 BGR so the OpenCV steps below see exactly what + `cv2.imread` would have produced from an opaque file.""" + from PIL import Image, ImageOps + + im = Image.open(io.BytesIO(image_bytes)) + im = ImageOps.exif_transpose(im) or im + rgba = im.convert("RGBA") + white = Image.new("RGBA", rgba.size, (255, 255, 255, 255)) + rgb = Image.alpha_composite(white, rgba).convert("RGB") + return np.ascontiguousarray(np.asarray(rgb, dtype=np.uint8)[:, :, ::-1]) + + def preprocess(image_bytes: bytes) -> Optional[np.ndarray]: """Decode `image_bytes` into the model's input tensor, or None. - ==== PREPROCESSING CONTRACT ============================================== - This is the DEFAULT recipe, written before the colleague's own code was - available. When theirs arrives, replace the body of this function with a - line-for-line port and delete this banner. The three decisions that decide - whether two systems' vectors agree are: - * resize method (default: BILINEAR) - * squash vs crop (default: squash the whole image to 224x224, - no aspect-ratio preservation, no centre crop) - * value scaling (default: raw 0..255 as float32 - what a Keras - MobileNetV3 export expects, since the graph - carries its own Rescaling layer) - Fixed in any variant: EXIF orientation applied (browsers apply it, so the - stored vector describes what a person sees), alpha flattened on white, - RGB channel order, NHWC layout, batch of one. - ========================================================================== + A line-for-line port of the Nearle Flutter app's preprocessing + (core/services/image_embed/image_embedder.dart, opencv_dart) and of the + colleague's Python reference for it. The steps, in their order: - No `Image.draft()` here, deliberately: a DCT-downscaled JPEG decode is not - bit-comparable with a full decode followed by a resize, and comparability - is the entire point. Never raises - pytest runs warnings as errors, and a + 1. read as BGR cv2.imdecode(..., IMREAD_COLOR) + 2. centre square crop side = min(w, h); x = (w - side) // 2; y = (h - side) // 2 + 3. resize to 224x224 cv2.resize(..., interpolation=INTER_AREA) + 4. BGR -> RGB cv2.cvtColor(..., COLOR_BGR2RGB) + 5. uint8 -> float 0..1 astype(float32) / 255.0 (ONLY that - no mean, + no std, no -1..1; the model rescales internally) + 6. batch dimension [1, 224, 224, 3], HWC + + OpenCV is used for the decode and the resize rather than Pillow on + purpose: INTER_AREA and Pillow's BOX filter are not the same filter, and + the two JPEG decoders differ by a pixel here and there. Matching the app + to 3-4 decimals is the requirement, so the app's library is the tool. + + The two "backend-only extras" from the same spec: an image WITH + transparency is flattened onto white before step 1 (IMREAD_COLOR would + turn the transparent area black), and a greyscale image comes out of + IMREAD_COLOR as three channels already. cv2.imdecode applies EXIF + orientation like cv2.imread does, which is what the phone camera path + relies on. + + Never raises - pytest runs warnings as errors, and a DecompressionBombWarning is one of the things the broad except absorbs. """ if not image_bytes: return None try: - from PIL import Image, ImageOps + import cv2 + from PIL import Image - im = Image.open(io.BytesIO(image_bytes)) - if im.width * im.height > IMAGE_VECTOR_MAX_PIXELS: - logger.debug("image rejected: %dx%d exceeds pixel cap", im.width, im.height) + # 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 - im = ImageOps.exif_transpose(im) or im + has_alpha = "A" in probe.getbands() or "transparency" in probe.info - if "A" in im.getbands() or "transparency" in im.info: - rgba = im.convert("RGBA") - white = Image.new("RGBA", rgba.size, (255, 255, 255, 255)) - im = Image.alpha_composite(white, rgba) - im = im.convert("RGB") - im = im.resize((IMAGE_EMBED_INPUT_SIZE, IMAGE_EMBED_INPUT_SIZE), Image.Resampling.BILINEAR) + 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 - arr = np.asarray(im, dtype=np.float32) # (224, 224, 3), 0..255 - arr = np.expand_dims(arr, axis=0) # (1, 224, 224, 3) + h, w = bgr.shape[:2] # 2 + side = min(w, h) + x, y = (w - side) // 2, (h - side) // 2 + sq = bgr[y:y + side, x:x + side] + + size = (IMAGE_EMBED_INPUT_SIZE, IMAGE_EMBED_INPUT_SIZE) + resized = cv2.resize(sq, size, interpolation=cv2.INTER_AREA) # 3 + rgb = cv2.cvtColor(resized, cv2.COLOR_BGR2RGB) # 4 + arr = (rgb.astype(np.float32) / 255.0)[None] # 5, 6 if arr.shape != _INPUT_SHAPE: logger.debug("image rejected: tensor shape %s", arr.shape) return None @@ -195,6 +228,36 @@ def available() -> bool: return _ensure_loaded() +def status() -> dict: + """Diagnostics for /api/health - NEVER loads the model. + + Exists because the failure this module is built to absorb is silent by + design: a container without the .tflite writes NULL vectors and says so + once, in a log line nobody reads. This puts the same facts on the health + endpoint, where "why are the vectors NULL after the deploy?" can be + answered with one curl. + """ + path = IMAGE_EMBED_MODEL_PATH + try: + import importlib.util + runtime = importlib.util.find_spec("ai_edge_litert") is not None + except Exception: # noqa: BLE001 - a broken finder counts as "not importable" + runtime = False + with _lock: + if _interpreter is not None: + state = "ready" + elif _disabled_reason: + state = f"disabled: {_disabled_reason}" + else: + state = "not loaded yet (loads on first write)" + return { + "model_path": str(path), + "model_present": path.is_file(), + "runtime_importable": runtime, + "state": state, + } + + def describe() -> str: """One line for script banners: where the model is and whether it works.""" with _lock: diff --git a/app/services/image_match.py b/app/services/image_match.py new file mode 100644 index 0000000..c691893 --- /dev/null +++ b/app/services/image_match.py @@ -0,0 +1,296 @@ +"""Find catalog products from a phone photo: img_vector nearest-neighbour + label text. + +(Named image_match, not image_search: app/services/image_search.py is the +image DISCOVERY module that finds photos for products. This is the reverse.) + +WHAT COMES IN +------------- +The Nearle app photographs a pack, crops it, embeds it on-device with the same +MobileNetV3-Small model and preprocessing that filled `img_vector` +(app/services/image_embedder.py), L2-normalises, and sends the 1024 floats - +plus whatever OCR read off the label. Alternatively a client sends the photo +and the API embeds it. Either way this module gets a unit vector and maybe +some text. + +HOW A MATCH IS SCORED +--------------------- +`score = 1 - (img_vector <=> q)`: cosine similarity, since both sides are unit +length. This is the app team's convention and is NOT +`RetrievedProduct.similarity` (`1 - distance/2`, rag_service.py), which is +the text search's. A photo of a pack against the catalog's render of it +scored 0.63 in the feasibility test; the app doc's "0.7 = same product" is a +phone-vs-phone rule of thumb, so the server's default floor is 0. + +WHY TEXT IS PART OF IT +---------------------- +Two reasons, both measured on the catalog: + +* Every pack size of one product usually shares one photo, so "Marie Gold + 89g / 300g / 1kg" tie to the last decimal. The label text is the only thing + that can pick the 300g row: size tokens are normalised ("300 g", "300gm", + "300G" -> "300g") and weighted above plain words. +* The brand name, when the OCR caught it, turns a 57-table scan into one + table via `query_intent.extract_brand_mention`. If that scope finds nothing + above `min_score` the search is retried unscoped and says so + (`scope_fallback`), because OCR misreads happen. An EXPLICIT `brand` never + falls back - the caller asked for a filter. + +The ranking is deterministic on purpose (`rank_key`): tied rows are ordered +by text overlap, then name, then image_id, so the same request always +returns the same order and the tests can pin it. + +Read-only: nothing here writes, and the pipeline stages are untouched. +""" +from __future__ import annotations + +import logging +import math +import re +from dataclasses import dataclass, field +from typing import Any, Dict, List, Optional, Sequence, Set, Tuple + +from app.infrastructure.settings import ( + IMAGE_SEARCH_DEFAULT_MIN_SCORE, + IMAGE_SEARCH_DEFAULT_TOP_K, + IMAGE_SEARCH_MAX_FETCH_K, + IMAGE_SEARCH_MAX_TOP_K, +) +from app.services.image_embedder import EMBED_DIM + +logger = logging.getLogger(__name__) + +# The floor pgvector's hnsw needs to return LIMIT rows: ef_search < LIMIT +# silently truncates the candidate list. +_MIN_EF_SEARCH = 40 + +_SIZE_RE = re.compile(r"(\d+(?:\.\d+)?)\s*(kg|gms|gm|g|ml|ltr|litre|l|pcs|pc|n)\b", re.I) +_UNIT_ALIAS = {"gm": "g", "gms": "g", "ltr": "l", "litre": "l", "pc": "pcs", "n": "pcs"} +_WORD_RE = re.compile(r"[a-z0-9]+") +_STOP = { + "the", "a", "an", "of", "and", "with", "for", "pack", "new", "net", "wt", + "weight", "mrp", "rs", "inr", "in", "by", "per", "no", "nos", +} +_SIZE_WEIGHT = 3.0 +_WORD_WEIGHT = 1.0 + + +class InvalidVectorError(ValueError): + """The query vector cannot be searched with: wrong length, NaN, or zero.""" + + +@dataclass +class ImageSearchResult: + rows: List[Dict[str, Any]] = field(default_factory=list) # each carries "score" and "text_overlap" + detected_brand: Optional[str] = None # explicit brand, else the OCR-derived one + scoped_to_brand: bool = False # the rows came from one brand table + scope_fallback: bool = False # OCR scope was empty; retried unscoped + min_score: float = 0.0 + query_text: Optional[str] = None + top_k: int = 0 + + +# --------------------------------------------------------------------------- +# the query vector +# --------------------------------------------------------------------------- + +def normalise_vector(vector: Sequence[float]) -> List[float]: + """`vector` as a unit-length list of EMBED_DIM floats, or InvalidVectorError. + + Already-unit input (|norm - 1| < 1e-3) is returned as-is so the app's own + normalisation is not disturbed by float rounding; anything else is + rescaled, because a client that forgot to normalise should still get the + right neighbours rather than distances scaled by its norm. + """ + try: + values = [float(v) for v in vector] + except (TypeError, ValueError) as exc: + raise InvalidVectorError(f"vector must be a list of numbers: {exc}") from None + if len(values) != EMBED_DIM: + raise InvalidVectorError(f"vector must have {EMBED_DIM} values, got {len(values)}") + if not all(math.isfinite(v) for v in values): + raise InvalidVectorError("vector contains NaN or infinite values") + norm = math.sqrt(sum(v * v for v in values)) + if norm < 1e-6: + raise InvalidVectorError("vector is all zeros") + if abs(norm - 1.0) < 1e-3: + return values + return [v / norm for v in values] + + +# --------------------------------------------------------------------------- +# label text +# --------------------------------------------------------------------------- + +def tokens(text: Optional[str]) -> Tuple[Set[str], Set[str]]: + """(words, sizes) from label text. + + Sizes are normalised so the OCR's "300 g", a sheet's "300gm" and a + product name's "300G" all become "300g". Words are lowercase, alnum, + at least two characters, minus stop-words and minus anything that was + part of a size token (so "300" and "g" do not also count as words). + """ + if not text: + return set(), set() + lowered = text.lower() + sizes: Set[str] = set() + for num, unit in _SIZE_RE.findall(lowered): + unit = _UNIT_ALIAS.get(unit, unit) + num = num.rstrip("0").rstrip(".") if "." in num else num + sizes.add(f"{num}{unit}") + without_sizes = _SIZE_RE.sub(" ", lowered) + words = {w for w in _WORD_RE.findall(without_sizes) if len(w) >= 2 and w not in _STOP} + return words, sizes + + +def _row_text(row: Dict[str, Any]) -> str: + parts = [str(row.get("product_name") or ""), str(row.get("title") or "")] + variants = row.get("size_variants") or [] + if isinstance(variants, (list, tuple)): + parts.extend(str(v) for v in variants if v) + return " ".join(parts) + + +def text_overlap(words: Set[str], sizes: Set[str], row: Dict[str, Any]) -> float: + """How much of the label text this row's name accounts for. + + 3.0 per shared size token, 1.0 per shared word. Sizes weigh more because + they are what separates the pack sizes of one product, which is the tie + this exists to break; brand and product words match every sibling alike. + """ + if not words and not sizes: + return 0.0 + row_words, row_sizes = tokens(_row_text(row)) + return _SIZE_WEIGHT * len(sizes & row_sizes) + _WORD_WEIGHT * len(words & row_words) + + +def rank_key(row: Dict[str, Any]) -> tuple: + """Best first. Score rounded to 3 dp so siblings sharing a photo tie.""" + return ( + -round(float(row.get("score", 0.0)), 3), + -float(row.get("text_overlap", 0.0)), + str(row.get("product_name") or row.get("title") or "").lower(), + str(row.get("image_id") or ""), + ) + + +# --------------------------------------------------------------------------- +# the search +# --------------------------------------------------------------------------- + +def _candidates( + vector: List[float], + brand: Optional[str], + category: Optional[str], + fetch_k: int, + ef_search: int, + min_score: float, +) -> List[Dict[str, Any]]: + from app.services.vector_store import image_vector_search + + rows = image_vector_search(vector, brand=brand, top_k=fetch_k, category=category, ef_search=ef_search) + out: List[Dict[str, Any]] = [] + seen: Set[Tuple[str, str]] = set() + for row in rows: + try: + score = 1.0 - float(row.get("distance")) + except (TypeError, ValueError): + continue + if score < min_score: + continue + key = (str(row.get("brand_table") or ""), str(row.get("image_id") or "")) + if key in seen: + continue + seen.add(key) + row["score"] = score + out.append(row) + return out + + +def _hydrate(rows: List[Dict[str, Any]]) -> List[Dict[str, Any]]: + """Swap the light candidate rows for full product rows, one read per table. + + Ranking ran on ~100-byte rows (image_id, name, sizes, distance) because + the unscoped search pulls candidates from every brand table; only the + winners are worth a 7KB product card. `score` and `text_overlap` are + carried over. A row whose card cannot be read keeps its light form, so + a match is never lost to a hydration hiccup. + """ + from app.services.vector_store import fetch_products_by_image_ids + + by_table: Dict[str, List[str]] = {} + for row in rows: + by_table.setdefault(str(row.get("brand_table") or ""), []).append(str(row.get("image_id") or "")) + cards: Dict[Tuple[str, str], Dict[str, Any]] = {} + for table, ids in by_table.items(): + if not table: + continue + for card in fetch_products_by_image_ids(table, ids): + cards[(table, str(card.get("image_id") or ""))] = card + + out: List[Dict[str, Any]] = [] + for row in rows: + key = (str(row.get("brand_table") or ""), str(row.get("image_id") or "")) + card = cards.get(key) + if card is None: + out.append(row) + continue + merged = dict(card) + merged["brand"] = merged.get("brand") or row.get("brand") + merged["brand_table"] = row.get("brand_table") + merged["score"] = row["score"] + merged["text_overlap"] = row.get("text_overlap", 0.0) + out.append(merged) + return out + + +def search_by_vector( + vector: Sequence[float], + 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, +) -> ImageSearchResult: + """The catalog rows most like `vector`, best first, at most `top_k`. + + Raises InvalidVectorError for a vector that cannot be searched with. + Everything else degrades to an empty result (no database, no vectors). + """ + unit = normalise_vector(vector) + top_k = max(1, min(int(top_k), IMAGE_SEARCH_MAX_TOP_K)) + fetch_k = min(max(top_k * 3, 30), IMAGE_SEARCH_MAX_FETCH_K) + ef_search = max(_MIN_EF_SEARCH, fetch_k) + label = (text or "").strip() or None + + explicit = (brand or "").strip() or None + detected = explicit + if detected is None and label: + try: + from app.services.query_intent import extract_brand_mention + detected = extract_brand_mention(label) + except Exception as exc: # noqa: BLE001 - brand detection is an optimisation + logger.debug("brand detection skipped: %s", exc) + detected = None + + rows = _candidates(unit, detected, category, fetch_k, ef_search, min_score) + scoped = detected is not None + fallback = False + if not rows and detected and not explicit: + rows = _candidates(unit, None, category, fetch_k, ef_search, min_score) + scoped, fallback = False, True + + words, sizes = tokens(label) + for row in rows: + row["text_overlap"] = text_overlap(words, sizes, row) + rows.sort(key=rank_key) + winners = _hydrate(rows[:top_k]) + + return ImageSearchResult( + rows=winners, + detected_brand=detected, + scoped_to_brand=scoped, + scope_fallback=fallback, + min_score=min_score, + query_text=label, + top_k=top_k, + ) diff --git a/app/services/models/mobilenet/README.md b/app/services/models/mobilenet/README.md new file mode 100644 index 0000000..84cd313 --- /dev/null +++ b/app/services/models/mobilenet/README.md @@ -0,0 +1,33 @@ +# mobilenet_v3_small_embedder.tflite + +The image model behind `img_vector` (`app/services/image_embedder.py`): +MobileNetV3-Small, input `[1, 224, 224, 3]` float32 in 0..1, output +`[1, 1024]` (the hard-swish output of the 1024-wide `Conv_2` head layer), +L2-normalised by the caller. Override the path with `IMAGE_EMBED_MODEL_PATH`. + +## Where this file came from + +Produced by `scripts/export_mobilenet_embedder.py` on 2026-09-17 (TensorFlow +2.21.0 / Keras 3.15.1) from Keras' pretrained ImageNet weights, with a +`Rescaling(2, -1)` layer inside the graph so the app's "divide by 255 and +nothing else" rule holds. 6.1 MB, float32, builtin TFLite ops only. That +script's docstring has the design and the exact steps to regenerate it. + +It is meant to be equivalent to the Nearle Flutter app's copy at +`assets/models/mobilenet/mobilenet_v3_small_embedder.tflite`. Same +architecture, weights, tap and input convention - but whether the numbers +agree to the last decimal can only be checked against the app: take one +cropped photo from the app, note the first eight values of its +`[VECTOR][IMAGE]` log line, run `image_embedder.embedding_for_bytes()` on the +same file, compare. If they differ, copy the app's file over this one and run +`python -m scripts.backfill_image_vectors --all --force --apply`. + +## Why this directory and not `data/` + +`/app/data` is a Docker named volume on every deployment, so a file baked +into the image under it is invisible on any volume that already exists. +`app/` is `COPY`d into the image and never mounted. + +Commit the file (`.gitattributes` marks `*.tflite` binary). Without it the +deployed container writes NULL into `img_vector` and reports +`"model_present": false` under `image_vectors` on `GET /api/health`. diff --git a/app/services/models/mobilenet/mobilenet_v3_small_embedder.tflite b/app/services/models/mobilenet/mobilenet_v3_small_embedder.tflite new file mode 100644 index 0000000..2a147ff Binary files /dev/null and b/app/services/models/mobilenet/mobilenet_v3_small_embedder.tflite differ diff --git a/app/services/vector_store.py b/app/services/vector_store.py index 195aa12..2382b1a 100644 --- a/app/services/vector_store.py +++ b/app/services/vector_store.py @@ -1503,6 +1503,127 @@ def semantic_search( return results[:top_k] +# The light projection the image search ranks on. Everything the tie-break +# needs and nothing else: ~100 bytes a row instead of the ~7KB product card, +# which matters because the unscoped search pulls `top_k` candidates from +# EVERY brand table (58 x 30 rows) to keep a handful. The winners are then +# hydrated by fetch_products_by_image_ids. +_IMAGE_CANDIDATE_COLUMNS = "image_id, product_name, title, size_variants" + + +def image_vector_search( + query_vector: List[float], + brand: Optional[str] = None, + top_k: int = 10, + category: Optional[str] = None, + ef_search: Optional[int] = None, +) -> List[Dict[str, Any]]: + """Nearest catalog rows to an image embedding, by cosine distance on img_vector. + + The read half of app/services/image_match.py: semantic_search() for the + 1024-d MobileNetV3 column instead of the 384-d text column, but returning + only the LIGHT candidate columns (`_IMAGE_CANDIDATE_COLUMNS`) plus + `distance` (pgvector `<=>`, so 1 - distance is cosine similarity for unit + vectors), `brand` and `brand_table`. Hydrate the winners with + fetch_products_by_image_ids. + + Unlike semantic_search this returns EVERY candidate, `top_k` PER TABLE, + merged and sorted by distance, and does not truncate: the pack sizes of + one product share a photo and tie exactly, and the caller breaks those + ties with the label text before choosing its top_k. + + `ef_search` sets hnsw.ef_search for this connection (one per call, closed + below, so a plain SET is per-request). The index returns at most + ef_search candidates, so it must be >= LIMIT or the tied siblings are + silently dropped. A server that does not know the GUC just logs. + """ + conn = _connect() + if not conn: + return [] + + embedding_str = "[" + ",".join(map(str, query_vector)) + "]" + results: List[Dict[str, Any]] = [] + + try: + with conn, conn.cursor() as cur: + if ef_search: + try: + cur.execute(f"SET hnsw.ef_search = {int(ef_search)}") + except Exception as e: # noqa: BLE001 - no hnsw, or a recorder cursor: exact scan still works + logger.debug("hnsw.ef_search not applied: %s", e) + + if brand: + table_name = _table_name(brand) + tables = [(brand, table_name)] if _table_exists(cur, table_name) else [] + else: + # Straight from information_schema a moment ago; a table that + # vanishes in between fails inside the per-table try below. + tables = [(name, f"brand_{name}") for name in _list_brand_table_suffixes(cur)] + + for brand_label, table_name in tables: + sql = ( + f"SELECT {_IMAGE_CANDIDATE_COLUMNS}, img_vector <=> %s::vector AS distance " + f"FROM {table_name} WHERE img_vector IS NOT NULL" + ) + params: List[Any] = [embedding_str] + if category: + sql += " AND category ILIKE %s" + params.append(f"%{category}%") + sql += " ORDER BY distance ASC LIMIT %s" + params.append(top_k) + + try: + cur.execute(sql, params) + except Exception as e: # noqa: BLE001 - a table without the column must not break search + logger.warning("Image vector search failed for table %s: %s", table_name, e) + continue + + # The unscoped loop labels by table suffix ("hindustan_unilever"); + # give every hit the display name the scoped path already has. + label = brand_label if brand else display_name_for_suffix(brand_label) + colnames = [desc[0] for desc in cur.description] + for row in cur.fetchall(): + record = dict(zip(colnames, row)) + record["brand"] = label + record["brand_table"] = table_name + results.append(record) + finally: + conn.close() + + results.sort(key=lambda r: r.get("distance", 9.0)) + return results + + +def fetch_products_by_image_ids(table_name: str, image_ids: List[str]) -> List[Dict[str, Any]]: + """Full product rows (vector columns projected out) for `image_ids` in one table. + + The hydration step of the image search: only the rows that survived + ranking are read in full. Order is not significant; the caller keys by + image_id. An unknown table or a failed read yields [] rather than raising. + """ + if not image_ids: + return [] + conn = _connect() + if not conn: + return [] + try: + with conn, conn.cursor() as cur: + if not _table_exists(cur, table_name): + return [] + try: + cur.execute( + f"SELECT {_product_columns(cur, table_name)} FROM {table_name} WHERE image_id = ANY(%s)", + (list(image_ids),), + ) + except Exception as e: # noqa: BLE001 + logger.warning("Product hydration failed for table %s: %s", table_name, e) + return [] + colnames = [desc[0] for desc in cur.description] + return [dict(zip(colnames, row)) for row in cur.fetchall()] + finally: + conn.close() + + def text_search( query: str, brand: Optional[str] = None, diff --git a/docs/IMAGE_SEARCH_API.md b/docs/IMAGE_SEARCH_API.md new file mode 100644 index 0000000..c6acb68 --- /dev/null +++ b/docs/IMAGE_SEARCH_API.md @@ -0,0 +1,141 @@ +# Search the catalogue by photo + +Base: `https://mcp.nearle.ai.in` · Auth: **none** (public, like `GET /api/search`) · Read-only + +The Nearle app photographs a pack, crops it, embeds it on-device with +MobileNetV3-Small (1024 floats, L2-normalised) and reads the label with OCR. +Every catalogue row carries the same kind of vector in `img_vector` +(same model, same OpenCV preprocessing, `vector(1024)` with an hnsw cosine +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 multipart file + the same optional fields as form fields +``` + +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. + +## How a match is found + +1. **Scope.** An explicit `brand` searches that one brand table, full stop. + Otherwise the OCR `text` is checked for a brand name ("Britannia Marie + Gold 300 g" → `brand_britannia`); if none is recognised, every brand table + is searched. A brand recognised from OCR that yields nothing above + `min_score` is retried across every brand (`scope_fallback: true`) - OCR + misreads happen; an explicit `brand` never falls back. +2. **Rank by cosine.** `score = 1 - (img_vector <=> vector)`. Both sides are + unit vectors, so this is cosine similarity; 1.0 is the identical picture. +3. **Break ties with the label.** Every pack size of one product usually + shares one catalogue photo, so "Marie Gold 89g / 300g / 1kg" tie exactly. + Size tokens in `text` ("300 g", "300gm", "300G" all read as `300g`) count + 3 points each, other words 1 point, matched against the product name and + size variants. The row with the most points wins the tie; then name order. +4. Return `top_k` cards. + +`score` is the honest number: a phone photo of a pack against the +catalogue's render of it measured **0.63** in testing; identical files give +0.99+. The app team's "above 0.7 means the same product" is a phone-vs-phone +rule of thumb. The server does not impose it - pass `min_score` if you want +a floor, and read `score` on each result. + +## Request fields + +| Field | Where | Default | Notes | +|---|---|---|---| +| `vector` | JSON only | required | exactly 1024 finite floats, not all zero; re-normalised if not unit length | +| `file` | multipart only | required | JPEG/PNG/WebP, ≤ 8 MB, ideally cropped to the pack; transparency is flattened on white | +| `text` | both | – | OCR text from the label, ≤ 500 chars | +| `brand` | both | – | hard filter to one brand | +| `category` | both | – | `ILIKE` filter on the category column | +| `top_k` | both | 10 | 1–50 | +| `min_score` | both | 0.0 | −1…1; drop matches below it | + +## Examples + +```bash +# The app: vector + OCR text +curl -s -X POST https://mcp.nearle.ai.in/api/search/image-vector \ + -H 'Content-Type: application/json' \ + -d '{"vector":[0.0312,0.0682,-0.0223, ... 1024 values ...], + "text":"Britannia Marie Gold 300 g","top_k":5}' + +# A photo +curl -s -X POST https://mcp.nearle.ai.in/api/search/image \ + -F 'file=@marie_gold.jpg' -F 'text=Britannia Marie Gold 300 g' -F 'top_k=5' +``` + +Captured response (a simulated phone shot of Marie Gold, trimmed to two results): + +```json +{ + "results": [ + { + "image_id": "britannia_marie_gold_300g", + "image_url": "https://www.britannia.co.in/_next/image?url=...", + "image_urls": ["https://www.britannia.co.in/_next/image?url=..."], + "brand": "Britannia", + "product_name": "Britannia Marie Gold 300g", + "title": "Britannia Marie Gold 300g", + "category": "General", + "description": "...", + "price_range": "₹45 - ₹55", + "size_variants": ["300g"], + "providers": [], + "highlights": [], + "nutrients": [], + "fssai_license": null, + "product_sku": "BRI-0342", + "sku_source": "internal", + "hsn_code": "1905", + "final_selling_price": 51.0, + "selling_price": 48.0, + "barcode": "8901063023949", + "barcode_type": "EAN13", + "nutrition_score": 6.2, + "health_score": 5.8, + "score": 0.631, + "text_overlap": 6.0 + }, + { "product_name": "Britannia Marie Gold 117g", "score": 0.631, "text_overlap": 3.0, "...": "..." } + ], + "total": 5, + "detected_brand": "Britannia", + "scoped_to_brand": true, + "scope_fallback": false, + "min_score": 0.0, + "top_k": 5, + "query_text": "Britannia Marie Gold 300 g" +} +``` + +Every result is the full product card (`ProductOut`) plus `score` and +`text_overlap`. `detected_brand` is the brand the search was scoped to, +whether it came from `brand` or from the text. + +## Errors + +| Code | Cause | What to do | +|---|---|---| +| 400 | `/image`: the upload is empty | send the file | +| 413 | `/image`: file over 8 MB | crop or downscale | +| 422 | wrong vector length, NaN, all zeros, `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 | + +## Good to know + +- The brand switch `ACTIVE_BRANDS` (blank in production = all brands) limits + the *unscoped* search exactly as it limits `GET /api/search`; an explicit or + OCR-recognised brand is searched even if inactive. +- Rows with no usable photo have no vector and cannot be found this way + (~92% of image-bearing rows are covered; the rest are dead image hosts). +- The unscoped search runs one small index query per brand table (≈60) and + then reads full cards only for the winners; expect tens of milliseconds + with the database co-located, more over a WAN. +- Nothing here writes. The ingestion pipeline, upserts and the image-vector + worker are untouched. + +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`. diff --git a/requirements.txt b/requirements.txt index 2f920f2..67352b6 100644 --- a/requirements.txt +++ b/requirements.txt @@ -69,6 +69,11 @@ Pillow>=10.0.0 # cp311 manylinux and for cp314 Windows. Imported lazily on first use, so a # missing wheel means "no image vectors", never a failed boot. ai-edge-litert>=2.2.0 +# The Nearle Flutter app preprocesses with OpenCV (centre crop, INTER_AREA +# resize, 0..1) and its vectors must match ours to a few decimals, so the +# 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 aiofiles>=23.2.1 # Playwright (Python) - last-resort image-search fallback only. diff --git a/scripts/export_mobilenet_embedder.py b/scripts/export_mobilenet_embedder.py new file mode 100644 index 0000000..660fbdf --- /dev/null +++ b/scripts/export_mobilenet_embedder.py @@ -0,0 +1,184 @@ +#!/usr/bin/env python3 +""" +Build `mobilenet_v3_small_embedder.tflite` - the image model behind img_vector. + +WHY THIS EXISTS +--------------- +The Nearle Flutter app embeds product photos with a MobileNetV3-Small TFLite +model whose output is the 1024-wide penultimate layer (input [1,224,224,3] +float32 in 0..1, output [1,1024], L2-normalised afterwards). That file is a +hand-made export, not a package: nothing installs it. When the app's own copy +is not to hand, this script produces an equivalent one from Keras' pretrained +ImageNet weights, so the catalog can be vectorised at all. + +WHAT IT BUILDS +-------------- + Input(224, 224, 3) float32, 0..1 - the app divides by 255 and + -> Rescaling(scale=2, offset=-1) nothing else; MobileNetV3 wants -1..1, so + -> MobileNetV3Small backbone the rescale lives INSIDE the graph + (imagenet weights, include_top=True, + include_preprocessing=False) + -> "Conv_2" 1x1 conv (576 -> 1024) + hard-swish <- the embedding + -> Flatten [1, 1024] + +`Conv_2` is the only 1024-wide layer in MobileNetV3-Small, so it is the layer +any "1024-d MobileNetV3-Small embedder" taps. Dropout and the 1000-way Logits +layer after it are discarded. No normalisation in the graph - the app +normalises after inference, and so does app/services/image_embedder.py. + +IS IT THE SAME AS THE APP'S FILE? +--------------------------------- +Same architecture, same pretrained weights, same tap, same input convention. +Whether the numbers agree to the last decimal depends on how the app's copy +was exported (TF version, converter flags), and that can only be checked +against the app itself: take one cropped photo from the app, note the first +eight values of its `[VECTOR][IMAGE]` log line, run +`image_embedder.embedding_for_bytes()` on the same file, and compare. Agreement +to 3-4 decimals means the two systems' vectors are interchangeable. If they +differ, catalog<->catalog search still works (the vectors are self-consistent); +drop the app's real file over this one and run +`python -m scripts.backfill_image_vectors --all --force --apply`. + +HOW TO RUN IT +------------- +TensorFlow is NOT a dependency of this project and must not become one (the +container is memory-capped and already carries torch). Run the export in a +throwaway container, from backend/: + + docker run --rm -v "${PWD}:/work" -w /work python:3.11-slim sh -c \\ + "pip install -q tensorflow-cpu && python scripts/export_mobilenet_embedder.py \\ + --out app/services/models/mobilenet/mobilenet_v3_small_embedder.tflite" + +Or, without Docker, in a throwaway venv at a SHORT path (TensorFlow's wheel +trips Windows' path-length limit inside a deep temp directory; TF has no +wheel for Python 3.14, so use 3.13 or 3.11): + + py -3.13 -m venv C:/tfexport_venv + C:/tfexport_venv/Scripts/pip install "tensorflow==2.21.*" "ai-edge-litert>=2.2.0" + C:/tfexport_venv/Scripts/python scripts/export_mobilenet_embedder.py --out ... + +The committed file was produced this way on 2026-09-17 with TensorFlow 2.21.0 +/ Keras 3.15.1. The pretrained weights are fixed, so re-running produces the +same model. The script self-checks the result with the same runtime the +backend uses (ai-edge-litert) when it is importable, else with tf.lite. +""" +from __future__ import annotations + +import argparse +import sys +import tempfile +from pathlib import Path + +INPUT_SIZE = 224 +EMBED_DIM = 1024 +TAP_LAYER = "Conv_2" # Keras' name for the 1x1 conv that widens 576 -> 1024 + + +def build_embedder(): + """The Keras model described in the module docstring.""" + import tensorflow as tf + from tensorflow import keras + + base = keras.applications.MobileNetV3Small( + input_shape=(INPUT_SIZE, INPUT_SIZE, 3), + weights="imagenet", + include_top=True, # we need the head's Conv_2, which include_top=False drops + include_preprocessing=False, # the 0..1 -> -1..1 rescale is added explicitly below + ) + # Keras 2 named it "Conv_2", Keras 3 "conv_2"; match case-insensitively. + names = [layer.name for layer in base.layers] + try: + idx = [n.lower() for n in names].index(TAP_LAYER.lower()) + except ValueError: + raise SystemExit(f"no layer named {TAP_LAYER} in MobileNetV3Small; layers: {names}") + conv = base.layers[idx] + # The layer right after it is the hard-swish activation whose output is + # the embedding (then come dropout and the 1000-way logits, discarded). + act = base.layers[idx + 1] + if int(conv.output.shape[-1]) != EMBED_DIM or int(act.output.shape[-1]) != EMBED_DIM: + raise SystemExit(f"{conv.name}/{act.name} are not {EMBED_DIM} wide: " + f"{conv.output.shape[-1]}/{act.output.shape[-1]}") + if "activation" not in act.name.lower() and "swish" not in act.name.lower(): + raise SystemExit(f"layer after {conv.name} is {act.name}, expected the hard-swish activation") + + inputs = keras.Input(shape=(INPUT_SIZE, INPUT_SIZE, 3), dtype="float32", name="image_0_1") + x = keras.layers.Rescaling(scale=2.0, offset=-1.0, name="rescale_0_1_to_pm1")(inputs) + features = keras.Model(base.input, act.output, name="mobilenet_v3_small_features")(x) + outputs = keras.layers.Flatten(name="embedding")(features) + model = keras.Model(inputs, outputs, name="mobilenet_v3_small_embedder") + if tuple(model.output.shape) != (None, EMBED_DIM): + raise SystemExit(f"unexpected output shape {model.output.shape}") + return model + + +def convert_to_tflite(model) -> bytes: + import tensorflow as tf + + with tempfile.TemporaryDirectory() as tmp: + saved = Path(tmp) / "saved_model" + # Keras 3 exports a SavedModel via .export(); Keras 2 via tf.saved_model.save. + if hasattr(model, "export"): + model.export(str(saved)) + else: + tf.saved_model.save(model, str(saved)) + converter = tf.lite.TFLiteConverter.from_saved_model(str(saved)) + # Float32, builtin ops only, no quantisation: the backend runtime is + # plain ai-edge-litert and must not need SELECT_TF_OPS. + converter.target_spec.supported_ops = [tf.lite.OpsSet.TFLITE_BUILTINS] + return converter.convert() + + +def self_check(path: Path) -> None: + import numpy as np + + try: + from ai_edge_litert.interpreter import Interpreter + runtime = "ai-edge-litert" + except ImportError: + import tensorflow as tf + Interpreter = tf.lite.Interpreter + runtime = "tf.lite" + + it = Interpreter(model_path=str(path)) + it.allocate_tensors() + inp, out = it.get_input_details(), it.get_output_details() + assert len(inp) == 1 and list(inp[0]["shape"]) == [1, INPUT_SIZE, INPUT_SIZE, 3], inp + assert np.dtype(inp[0]["dtype"]) == np.float32, inp[0]["dtype"] + assert len(out) == 1 and list(out[0]["shape"]) == [1, EMBED_DIM], out + + rng = np.random.default_rng(0) + for label, x in ( + ("zeros", np.zeros((1, INPUT_SIZE, INPUT_SIZE, 3), np.float32)), + ("random", rng.random((1, INPUT_SIZE, INPUT_SIZE, 3), dtype=np.float32)), + ): + it.set_tensor(inp[0]["index"], x) + it.invoke() + v = it.get_tensor(out[0]["index"])[0] + norm = float(np.linalg.norm(v)) + assert np.all(np.isfinite(v)) and norm > 0, label + print(f" {label:6s} first 8: {np.round(v[:8], 5).tolist()} norm={norm:.4f} " + f"zeros={int((v == 0).sum())}/{EMBED_DIM}") + print(f" self-check passed with {runtime}: input [1,{INPUT_SIZE},{INPUT_SIZE},3] float32, output [1,{EMBED_DIM}]") + + +def main() -> int: + ap = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter) + ap.add_argument("--out", required=True, help="where to write the .tflite") + args = ap.parse_args() + out = Path(args.out) + out.parent.mkdir(parents=True, exist_ok=True) + + import tensorflow as tf + print(f"tensorflow {tf.__version__}, keras {tf.keras.__version__ if hasattr(tf.keras, '__version__') else '?'}") + + model = build_embedder() + print(f"model: {model.name}, params={model.count_params():,}, tap={TAP_LAYER}+hard-swish") + data = convert_to_tflite(model) + out.write_bytes(data) + print(f"wrote {out} ({len(data) / 1e6:.1f} MB)") + self_check(out) + return 0 + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/tests/test_image_match.py b/tests/test_image_match.py new file mode 100644 index 0000000..bad08f5 --- /dev/null +++ b/tests/test_image_match.py @@ -0,0 +1,352 @@ +"""Search by image: the ranking, the scope rules, and the SQL behind them. + +No database and no model: `image_vector_search` is patched with canned rows, +and the SQL-shape test uses a recorder cursor. The numbers are the ones the +feasibility run produced - three Marie Gold pack sizes sharing one photo tie +at 0.631 and a Good Day trails at 0.561 - so the tie-break is pinned to the +case it was built for. +""" +from __future__ import annotations + +import inspect +import math +from typing import Any, Dict, List + +import pytest + +from app.services import image_match as im +from app.services import vector_store +from app.services.vector_store import fetch_products_by_image_ids as real_fetch_products_by_image_ids + + +def _unit(seed: float = 1.0) -> List[float]: + v = [math.sin(seed * (i + 1)) for i in range(1024)] + n = math.sqrt(sum(x * x for x in v)) + return [x / n for x in v] + + +def _row(name: str, distance: float, image_id: str = "", sizes=None, table="brand_britannia") -> Dict[str, Any]: + return { + "image_id": image_id or name.lower().replace(" ", "_"), + "product_name": name, + "title": name, + "brand": "Britannia", + "brand_table": table, + "size_variants": sizes or [], + "distance": distance, + } + + +MARIE = [ + _row("Britannia Marie Gold 89g", 0.369), + _row("Britannia Marie Gold 300g", 0.369), + _row("Britannia Marie Gold 1kg", 0.369), + _row("Britannia Good Day 100g", 0.439), +] + + +@pytest.fixture(autouse=True) +def hydrate_from_the_light_rows(monkeypatch): + """Stand-in for the second read: the card for an image_id is the light + row plus a `barcode`, so tests can see hydration happened and where.""" + def fake(table, image_ids): + return [{**dict(r), "barcode": f"890-{r['image_id']}"} for r in MARIE + if r["brand_table"] == table and r["image_id"] in image_ids] + monkeypatch.setattr(vector_store, "fetch_products_by_image_ids", fake) + + +@pytest.fixture +def calls(monkeypatch): + """Patch the store; return the list of (kwargs) it was called with.""" + seen: List[Dict[str, Any]] = [] + + def fake(vector, brand=None, top_k=10, category=None, ef_search=None): + seen.append({"brand": brand, "top_k": top_k, "category": category, "ef_search": ef_search}) + return [dict(r) for r in MARIE] + + monkeypatch.setattr(vector_store, "image_vector_search", fake) + return seen + + +# --------------------------------------------------------------------------- +# the vector +# --------------------------------------------------------------------------- + +def test_a_unit_vector_passes_through_untouched(): + v = _unit() + assert im.normalise_vector(v) == v + + +def test_an_unnormalised_vector_is_rescaled_to_unit_length(): + out = im.normalise_vector([2.0] + [0.0] * 1023) + assert out[0] == 1.0 and abs(math.sqrt(sum(x * x for x in out)) - 1.0) < 1e-9 + + +@pytest.mark.parametrize("bad, msg", [ + ([0.1] * 1023, "1024 values"), + ([float("nan")] + [0.0] * 1023, "NaN"), + ([0.0] * 1024, "all zeros"), + (["x"] * 1024, "list of numbers"), +]) +def test_unsearchable_vectors_are_refused_with_a_reason(bad, msg): + with pytest.raises(im.InvalidVectorError, match=msg): + im.normalise_vector(bad) + + +# --------------------------------------------------------------------------- +# the label text +# --------------------------------------------------------------------------- + +def test_sizes_are_normalised_and_words_cleaned(): + words, sizes = im.tokens("Britannia Marie Gold 300 g Net Wt") + assert words == {"britannia", "marie", "gold"} and sizes == {"300g"} + + +@pytest.mark.parametrize("text, size", [("89gm", "89g"), ("89 GMS", "89g"), ("1 LTR", "1l"), ("1 litre", "1l"), + ("2.50 kg", "2.5kg"), ("500ml", "500ml"), ("6 pcs", "6pcs")]) +def test_every_way_a_label_writes_a_size_collapses_to_one_token(text, size): + assert im.tokens(text)[1] == {size} + + +def test_empty_text_has_no_tokens_and_no_overlap(): + assert im.tokens(None) == (set(), set()) + assert im.text_overlap(set(), set(), MARIE[0]) == 0.0 + + +def test_size_matches_outweigh_word_matches(): + words, sizes = im.tokens("Marie Gold 300 g") + assert im.text_overlap(words, sizes, MARIE[1]) == 3.0 + 2.0 # 300g + marie + gold + assert im.text_overlap(words, sizes, MARIE[0]) == 2.0 # marie + gold only + + +# --------------------------------------------------------------------------- +# ranking +# --------------------------------------------------------------------------- + +def test_label_text_picks_the_pack_size_among_tied_siblings(calls): + result = im.search_by_vector(_unit(), text="Marie Gold 300 g") + + names = [r["product_name"] for r in result.rows] + assert names == [ + "Britannia Marie Gold 300g", # size token wins the tie + "Britannia Marie Gold 1kg", # then name order among the rest + "Britannia Marie Gold 89g", + "Britannia Good Day 100g", # lower score, whatever the text + ] + assert result.rows[0]["score"] == pytest.approx(0.631) + assert result.rows[0]["text_overlap"] == 5.0 + assert result.rows[0]["barcode"] == "890-britannia_marie_gold_300g" # hydrated, and to the right row + + +def test_without_text_ties_fall_back_to_name_order_and_overlap_is_zero(calls): + result = im.search_by_vector(_unit()) + + assert [r["product_name"] for r in result.rows][:3] == [ + "Britannia Marie Gold 1kg", "Britannia Marie Gold 300g", "Britannia Marie Gold 89g", + ] + assert all(r["text_overlap"] == 0.0 for r in result.rows) + + +def test_min_score_drops_everything_below_it(calls): + result = im.search_by_vector(_unit(), min_score=0.65) + assert result.rows == [] and result.min_score == 0.65 + + +def test_top_k_truncates_after_ranking_and_only_winners_are_hydrated(calls, monkeypatch): + asked = [] + monkeypatch.setattr(vector_store, "fetch_products_by_image_ids", + lambda table, ids: asked.append((table, sorted(ids))) or []) + result = im.search_by_vector(_unit(), text="300 g", top_k=1) + + assert [r["product_name"] for r in result.rows] == ["Britannia Marie Gold 300g"] + assert result.top_k == 1 + assert asked == [("brand_britannia", ["britannia_marie_gold_300g"])] + + +def test_a_row_whose_card_cannot_be_read_keeps_its_light_form(calls, monkeypatch): + monkeypatch.setattr(vector_store, "fetch_products_by_image_ids", lambda table, ids: []) + result = im.search_by_vector(_unit(), top_k=2) + assert len(result.rows) == 2 and "barcode" not in result.rows[0] and result.rows[0]["score"] == pytest.approx(0.631) + + +def test_duplicate_rows_from_the_same_table_are_collapsed(monkeypatch): + monkeypatch.setattr(vector_store, "image_vector_search", + lambda *a, **k: [dict(MARIE[0]), dict(MARIE[0])]) + assert len(im.search_by_vector(_unit()).rows) == 1 + + +# --------------------------------------------------------------------------- +# scope +# --------------------------------------------------------------------------- + +def test_an_explicit_brand_is_a_hard_filter_that_never_falls_back(monkeypatch): + seen = [] + monkeypatch.setattr(vector_store, "image_vector_search", + lambda vector, brand=None, **k: seen.append(brand) or []) + + result = im.search_by_vector(_unit(), brand="Cadbury", text="britannia marie") + + assert seen == ["Cadbury"] + assert result.rows == [] and result.scoped_to_brand and not result.scope_fallback + assert result.detected_brand == "Cadbury" + + +def test_a_brand_read_off_the_label_scopes_the_search(monkeypatch): + seen = [] + monkeypatch.setattr(vector_store, "image_vector_search", + lambda vector, brand=None, **k: seen.append(brand) or [dict(r) for r in MARIE]) + from app.services import query_intent + monkeypatch.setattr(query_intent, "extract_brand_mention", lambda text: "Britannia") + + result = im.search_by_vector(_unit(), text="Britannia Marie Gold") + + assert seen == ["Britannia"] + assert result.detected_brand == "Britannia" and result.scoped_to_brand and not result.scope_fallback + + +def test_an_ocr_brand_that_finds_nothing_retries_across_every_brand(monkeypatch): + seen = [] + + def fake(vector, brand=None, **k): + seen.append(brand) + return [] if brand else [dict(r) for r in MARIE] + + monkeypatch.setattr(vector_store, "image_vector_search", fake) + from app.services import query_intent + monkeypatch.setattr(query_intent, "extract_brand_mention", lambda text: "Britannia") + + result = im.search_by_vector(_unit(), text="Britannia marie") + + assert seen == ["Britannia", None] + assert result.scope_fallback and not result.scoped_to_brand and len(result.rows) == 4 + + +def test_no_text_and_no_brand_means_one_unscoped_query(calls): + im.search_by_vector(_unit()) + assert len(calls) == 1 and calls[0]["brand"] is None + + +def test_brand_detection_failures_do_not_break_the_search(calls, monkeypatch): + from app.services import query_intent + + def boom(text): + raise RuntimeError("brand index unavailable") + + monkeypatch.setattr(query_intent, "extract_brand_mention", boom) + result = im.search_by_vector(_unit(), text="something") + assert result.detected_brand is None and len(result.rows) == 4 + + +# --------------------------------------------------------------------------- +# candidate width and ef_search +# --------------------------------------------------------------------------- + +@pytest.mark.parametrize("top_k, fetch_k, ef", [(1, 30, 40), (10, 30, 40), (20, 60, 60), (50, 100, 100), (500, 100, 100)]) +def test_the_store_is_asked_for_enough_candidates_to_hold_the_ties(calls, top_k, fetch_k, ef): + im.search_by_vector(_unit(), top_k=top_k) + assert calls[0]["top_k"] == fetch_k and calls[0]["ef_search"] == ef + + +# --------------------------------------------------------------------------- +# the SQL +# --------------------------------------------------------------------------- + +class _Cursor: + def __init__(self): + self.statements: List[str] = [] + self.description = None + self._pending: Any = None + + def execute(self, sql, params=None): + text = " ".join(str(sql).split()) + self.statements.append(text) + if "information_schema.columns" in text: + self._pending = [("brand_x", c) for c in ("id", "product_name", "embedding", "img_vector", "img_vector_src")] + elif "information_schema.tables" in text and "EXISTS" in text: + self._pending = (True,) + elif "information_schema.tables" in text: + self._pending = [("brand_x",)] + elif text.startswith("SET"): + self._pending = None + elif "AS distance" in text: + self.description = [("image_id",), ("product_name",), ("title",), ("size_variants",), ("distance",)] + self._pending = [("a", "P", "P", [], 0.2)] + else: + self.description = [("id",), ("product_name",), ("img_vector_src",)] + self._pending = [(1, "P", None)] + + def fetchall(self): + return self._pending if isinstance(self._pending, list) else [] + + def fetchone(self): + return self._pending if isinstance(self._pending, tuple) else None + + def __enter__(self): + return self + + def __exit__(self, *a): + return False + + +class _Conn: + def __init__(self, cur): + self._cur = cur + + def cursor(self): + return self._cur + + def __enter__(self): + return self + + def __exit__(self, *a): + return False + + def close(self): + pass + + +def test_the_query_reads_img_vector_by_name_and_sets_ef_search(monkeypatch): + vector_store._invalidate_product_columns_cache() + cur = _Cursor() + monkeypatch.setattr(vector_store, "_connect", lambda: _Conn(cur)) + + rows = vector_store.image_vector_search([0.0] * 1024, top_k=30, ef_search=40) + + select = [s for s in cur.statements if s.startswith("SELECT") and "FROM brand_x" in s and "distance" in s] + assert select, cur.statements + sql = select[0] + # Light projection only: ranking must not drag 7KB product cards per candidate. + assert sql.startswith("SELECT image_id, product_name, title, size_variants, img_vector <=> %s::vector AS distance") + assert "WHERE img_vector IS NOT NULL" in sql and sql.endswith("ORDER BY distance ASC LIMIT %s") + assert "embedding" not in sql and "description" not in sql + assert cur.statements.index("SET hnsw.ef_search = 40") < cur.statements.index(sql) + assert rows[0]["brand_table"] == "brand_x" and rows[0]["brand"] == "X" and rows[0]["distance"] == 0.2 + + +def test_hydration_reads_full_cards_by_image_id_without_the_vectors(monkeypatch): + vector_store._invalidate_product_columns_cache() + cur = _Cursor() + monkeypatch.setattr(vector_store, "_connect", lambda: _Conn(cur)) + + real_fetch_products_by_image_ids("brand_x", ["a", "b"]) # the autouse fixture patches the module attribute + + sql = [s for s in cur.statements if "WHERE image_id = ANY(%s)" in s] + assert sql and sql[0].startswith('SELECT "id", "product_name", "img_vector_src" FROM brand_x') + assert "embedding" not in sql[0] and '"img_vector"' not in sql[0] + assert real_fetch_products_by_image_ids("brand_x", []) == [] + + +def test_the_queries_are_never_select_star(): + assert "SELECT *" not in inspect.getsource(vector_store.image_vector_search) + assert "SELECT *" not in inspect.getsource(vector_store.fetch_products_by_image_ids) + + +def test_a_category_filter_is_added_inside_the_where(monkeypatch): + vector_store._invalidate_product_columns_cache() + cur = _Cursor() + monkeypatch.setattr(vector_store, "_connect", lambda: _Conn(cur)) + + vector_store.image_vector_search([0.0] * 1024, brand="X", category="Biscuits") + + sql = [s for s in cur.statements if "AS distance" in s][0] + assert "WHERE img_vector IS NOT NULL AND category ILIKE %s ORDER BY" in sql diff --git a/tests/test_image_search_api.py b/tests/test_image_search_api.py new file mode 100644 index 0000000..4a9ce8b --- /dev/null +++ b/tests/test_image_search_api.py @@ -0,0 +1,174 @@ +"""POST /api/search/image-vector and POST /api/search/image - the HTTP contract. + +The ranking is tested in tests/test_image_match.py; here `search_by_vector` is +patched on the router module and the assertions are about status codes, +validation and the response shape the app reads. Both routes are public. +""" +from __future__ import annotations + +import io +import math + +import pytest + +from app.api.routers import search as search_router +from app.services.image_match import ImageSearchResult + + +def _unit(): + v = [math.cos(i / 7.0) for i in range(1024)] + n = math.sqrt(sum(x * x for x in v)) + return [x / n for x in v] + + +def _row(): + return { + "image_id": "britannia_marie_gold_300g", + "product_name": "Britannia Marie Gold 300g", + "title": "Britannia Marie Gold 300g", + "brand": "Britannia", + "brand_table": "brand_britannia", + "category": "Biscuits", + "image_url": "https://cdn.example/marie.jpg", + "image_urls": ["https://cdn.example/marie.jpg"], + "size_variants": ["300g"], + "barcode": "8901063010512", + "barcode_type": "EAN13", + "final_selling_price": 45.0, + "selling_price": 42.0, + "hsn_code": "1905", + "fssai_license": "10012021000123", + "distance": 0.369, + "score": 0.631, + "text_overlap": 5.0, + } + + +@pytest.fixture +def fake_search(monkeypatch): + calls = [] + + def fake(vector, text=None, brand=None, category=None, top_k=10, min_score=0.0): + calls.append({"vector": list(vector), "text": text, "brand": brand, + "category": category, "top_k": top_k, "min_score": min_score}) + return ImageSearchResult(rows=[_row()], detected_brand="Britannia", scoped_to_brand=True, + min_score=min_score, query_text=text, top_k=top_k) + + monkeypatch.setattr(search_router, "search_by_vector", fake) + return calls + + +# --------------------------------------------------------------------------- +# /search/image-vector +# --------------------------------------------------------------------------- + +def test_a_vector_returns_the_product_card_with_its_score(client, fake_search): + res = client.post("/api/search/image-vector", + json={"vector": _unit(), "text": "Britannia Marie Gold 300 g", "top_k": 5}) + + assert res.status_code == 200, res.text + body = res.json() + assert body["total"] == 1 and body["detected_brand"] == "Britannia" and body["scoped_to_brand"] is True + hit = body["results"][0] + assert hit["product_name"] == "Britannia Marie Gold 300g" + assert hit["score"] == 0.631 and hit["text_overlap"] == 5.0 + for key in ("barcode", "final_selling_price", "selling_price", "category", "image_url", + "size_variants", "hsn_code", "fssai_license", "image_id", "brand"): + assert key in hit, key + assert hit["barcode"] == "8901063010512" and hit["image_url"] == "https://cdn.example/marie.jpg" + assert fake_search[0]["text"] == "Britannia Marie Gold 300 g" and fake_search[0]["top_k"] == 5 + + +def test_the_route_is_public(client, fake_search): + assert client.post("/api/search/image-vector", json={"vector": _unit()}).status_code == 200 + + +@pytest.mark.parametrize("payload, fragment", [ + ({"vector": [0.1] * 1023}, "vector"), + ({"vector": [0.1] * 1025}, "vector"), + ({"vector": [0.0] * 1024}, "all zeros"), + ({"vector": _unit(), "top_k": 51}, "top_k"), + ({"vector": _unit(), "top_k": 0}, "top_k"), + ({"vector": _unit(), "min_score": 1.5}, "min_score"), + ({}, "vector"), +]) +def test_bad_requests_are_422_and_name_the_field(client, fake_search, payload, fragment): + res = client.post("/api/search/image-vector", json=payload) + assert res.status_code == 422 + assert fragment in res.text + assert fake_search == [] + + +def test_a_nan_in_the_vector_is_422(client, fake_search): + body = '{"vector": [' + ",".join(["NaN"] + ["0.1"] * 1023) + "]}" + res = client.post("/api/search/image-vector", content=body, headers={"content-type": "application/json"}) + assert res.status_code == 422 and fake_search == [] + + +def test_a_service_rejection_is_422_not_500(client, monkeypatch): + from app.services.image_match import InvalidVectorError + + def refuse(*a, **k): + raise InvalidVectorError("vector is all zeros") + + monkeypatch.setattr(search_router, "search_by_vector", refuse) + res = client.post("/api/search/image-vector", json={"vector": _unit()}) + assert res.status_code == 422 and "all zeros" in res.text + + +# --------------------------------------------------------------------------- +# /search/image +# --------------------------------------------------------------------------- + +def _post_image(client, data: bytes, **form): + return client.post("/api/search/image", files={"file": ("photo.jpg", io.BytesIO(data), "image/jpeg")}, data=form) + + +def test_a_photo_is_embedded_and_searched_with_its_form_fields(client, fake_search, monkeypatch): + monkeypatch.setattr(search_router.image_embedder, "available", lambda: True) + monkeypatch.setattr(search_router.image_embedder, "embedding_for_bytes", lambda data: _unit()) + + res = _post_image(client, b"\xff\xd8" + b"x" * 5000, text="Marie Gold 300 g", top_k="3", min_score="0.2") + + assert res.status_code == 200, res.text + assert res.json()["results"][0]["product_name"] == "Britannia Marie Gold 300g" + call = fake_search[0] + assert call["text"] == "Marie Gold 300 g" and call["top_k"] == 3 and call["min_score"] == 0.2 + assert len(call["vector"]) == 1024 + + +def test_without_a_model_the_photo_route_says_503_and_points_at_the_vector_route(client, fake_search, monkeypatch): + monkeypatch.setattr(search_router.image_embedder, "available", lambda: False) + + res = _post_image(client, b"x" * 5000) + + assert res.status_code == 503 and "/api/search/image-vector" in res.text + assert fake_search == [] + + +def test_an_oversized_photo_is_413_before_any_model_work(client, fake_search, monkeypatch): + monkeypatch.setattr(search_router, "IMAGE_VECTOR_MAX_BYTES", 100) + monkeypatch.setattr(search_router.image_embedder, "available", lambda: pytest.fail("must not load the model")) + + res = _post_image(client, b"x" * 101) + + assert res.status_code == 413 and "limit is" in res.text + + +def test_an_empty_upload_is_400(client, fake_search): + assert _post_image(client, b"").status_code == 400 + + +def test_an_undecodable_photo_is_422(client, fake_search, monkeypatch): + monkeypatch.setattr(search_router.image_embedder, "available", lambda: True) + monkeypatch.setattr(search_router.image_embedder, "embedding_for_bytes", lambda data: None) + + res = _post_image(client, b"not an image at all" * 100) + + assert res.status_code == 422 and "decode" in res.text and fake_search == [] + + +def test_both_routes_are_documented(client): + paths = client.get("/openapi.json").json()["paths"] + assert "/api/search/image-vector" in paths and "/api/search/image" in paths + assert "post" in paths["/api/search/image"] and "get" in paths["/api/search"] diff --git a/tests/test_image_vector.py b/tests/test_image_vector.py index 623ed96..3f4d905 100644 --- a/tests/test_image_vector.py +++ b/tests/test_image_vector.py @@ -105,13 +105,57 @@ def test_a_solid_png_becomes_a_224_tensor_dominated_by_its_colour(): assert t.max() > t.min() -def test_the_default_recipe_is_raw_0_to_255(): - """Pinned separately from the invariant tests above: this is the one - assertion that is EXPECTED to change when the colleague's preprocessing - replaces the default. Update it deliberately, not by accident.""" +def test_the_recipe_is_0_to_1_rgb_like_the_flutter_app(): + """Pinned separately from the invariant tests above: the app divides by + 255 and nothing else (no mean, no std, no -1..1), RGB order.""" t = emb.preprocess(_png(Image.new("RGB", (10, 10), (255, 128, 0)))) - assert float(t[0, 0, 0, 0]) == 255.0 and float(t[0, 0, 0, 2]) == 0.0 + assert float(t[0, 0, 0, 0]) == 1.0 + assert abs(float(t[0, 0, 0, 1]) - 128 / 255) < 1e-6 + assert float(t[0, 0, 0, 2]) == 0.0 + + +def _reference_image_to_tensor(path: str) -> np.ndarray: + """The colleague's Python reference, verbatim up to the model call.""" + import cv2 + img = cv2.imread(path, cv2.IMREAD_COLOR) # 1. BGR + h, w = img.shape[:2] + side = min(w, h) + x, y = (w - side) // 2, (h - side) // 2 + img = img[y:y + side, x:x + side] # 2. center crop + img = cv2.resize(img, (224, 224), interpolation=cv2.INTER_AREA) # 3. resize + img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB) # 4. RGB + return (img.astype(np.float32) / 255.0)[None] # 5. 0-1, batch + + +@pytest.mark.parametrize("size", [(640, 480), (300, 700), (224, 224), (37, 91)]) +def test_our_tensor_is_identical_to_the_colleagues_reference_code(tmp_path, size): + """Same bytes in, same floats out - not 'close', identical. This is the + guarantee that a catalog vector and an app vector describe the same + pixels; any drift here shows up as a lower cosine between the two.""" + noisy = Image.effect_noise(size, 60).convert("RGB") + noisy.paste((200, 30, 30), (0, 0, size[0] // 3, size[1] // 2)) + data = _jpeg(noisy) + path = tmp_path / "photo.jpg" + path.write_bytes(data) + + ours = emb.preprocess(data) + theirs = _reference_image_to_tensor(str(path)) + + assert ours is not None and ours.shape == theirs.shape == (1, 224, 224, 3) + assert np.array_equal(ours, theirs) + + +def test_a_wide_image_is_centre_cropped_not_squashed(): + """Left third red, middle third green, right third blue, 300x100. The app + crops the central 100x100 before resizing, so only green survives.""" + im = Image.new("RGB", (300, 100), (255, 0, 0)) + im.paste((0, 255, 0), (100, 0, 200, 100)) + im.paste((0, 0, 255), (200, 0, 300, 100)) + + t = emb.preprocess(_png(im)) + + assert np.all(t[0, :, :, 1] == 1.0) and np.all(t[0, :, :, 0] == 0.0) and np.all(t[0, :, :, 2] == 0.0) def test_a_solid_jpeg_is_uniform_within_lossy_tolerance(): @@ -127,10 +171,11 @@ def test_exif_orientation_is_applied_because_the_browser_applies_it(): """Left half red, right half blue, tagged 'rotate 90 CW to display'. Without the transpose the bottom-left pixel is red (the left half). With - it, the left half has become the top half and bottom-left is blue. + it, the left half has become the top half and bottom-left is blue. The + image is 64x64 so the centre crop keeps all of it. """ - im = Image.new("RGB", (64, 32), (255, 0, 0)) - im.paste((0, 0, 255), (32, 0, 64, 32)) + im = Image.new("RGB", (64, 64), (255, 0, 0)) + im.paste((0, 0, 255), (32, 0, 64, 64)) exif = Image.Exif() exif[0x0112] = 6 t = emb.preprocess(_jpeg(im, exif=exif.tobytes())) @@ -285,6 +330,41 @@ def test_inference_is_serialised_on_the_module_lock(monkeypatch): assert not overlap and len(fake.inputs) == 6 +def test_status_reports_a_missing_model_without_loading_it(monkeypatch, tmp_path): + """The deploy-time failure: code shipped, .tflite did not. /api/health + must say so, and asking must not itself trigger a load attempt.""" + monkeypatch.setattr(emb, "IMAGE_EMBED_MODEL_PATH", tmp_path / "missing.tflite") + loads = [] + monkeypatch.setattr(emb, "_load", lambda: loads.append(1)) + + st = emb.status() + + assert st["model_present"] is False and st["model_path"].endswith("missing.tflite") + assert st["state"].startswith("not loaded yet") + assert loads == [] + + +def test_status_reflects_a_recorded_failure_and_a_ready_interpreter(monkeypatch): + monkeypatch.setattr(emb, "_disabled_reason", "model file not found at x") + assert emb.status()["state"] == "disabled: model file not found at x" + + emb._reset() + _install_fake(monkeypatch) + assert emb.status()["state"] == "ready" + + +def test_health_carries_the_image_vector_block(monkeypatch, tmp_path): + from fastapi.testclient import TestClient + from app.main import app + monkeypatch.setattr(emb, "IMAGE_EMBED_MODEL_PATH", tmp_path / "missing.tflite") + + body = TestClient(app).get("/api/health").json() + + iv_block = body["image_vectors"] + assert iv_block["model_present"] is False + assert set(iv_block) == {"enabled", "model_path", "model_present", "runtime_importable", "state"} + + def test_to_pg_is_the_same_text_form_the_upsert_uses_for_embedding(): assert iv.to_pg([0, 0.5, 1]) == "[0.0,0.5,1.0]"