imag vector generation with dimentionality reduction

This commit is contained in:
sriram
2026-09-18 15:26:25 +05:30
parent afa0bfa743
commit d6296bd1f0
16 changed files with 1724 additions and 48 deletions

View File

@@ -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()),
)

View File

@@ -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)

View File

@@ -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
# ---------------------------------------------------------------------------

View File

@@ -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
# ---------------------------------------------------------------------------

View File

@@ -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.

View File

@@ -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:

296
app/services/image_match.py Normal file
View File

@@ -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,
)

View File

@@ -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`.

View File

@@ -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,

141
docs/IMAGE_SEARCH_API.md Normal file
View File

@@ -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`.

View File

@@ -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.

View File

@@ -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())

352
tests/test_image_match.py Normal file
View File

@@ -0,0 +1,352 @@
"""Search by image: the ranking, the scope rules, and the SQL behind them.
No database and no model: `image_vector_search` is patched with canned rows,
and the SQL-shape test uses a recorder cursor. The numbers are the ones the
feasibility run produced - three Marie Gold pack sizes sharing one photo tie
at 0.631 and a Good Day trails at 0.561 - so the tie-break is pinned to the
case it was built for.
"""
from __future__ import annotations
import inspect
import math
from typing import Any, Dict, List
import pytest
from app.services import image_match as im
from app.services import vector_store
from app.services.vector_store import fetch_products_by_image_ids as real_fetch_products_by_image_ids
def _unit(seed: float = 1.0) -> List[float]:
v = [math.sin(seed * (i + 1)) for i in range(1024)]
n = math.sqrt(sum(x * x for x in v))
return [x / n for x in v]
def _row(name: str, distance: float, image_id: str = "", sizes=None, table="brand_britannia") -> Dict[str, Any]:
return {
"image_id": image_id or name.lower().replace(" ", "_"),
"product_name": name,
"title": name,
"brand": "Britannia",
"brand_table": table,
"size_variants": sizes or [],
"distance": distance,
}
MARIE = [
_row("Britannia Marie Gold 89g", 0.369),
_row("Britannia Marie Gold 300g", 0.369),
_row("Britannia Marie Gold 1kg", 0.369),
_row("Britannia Good Day 100g", 0.439),
]
@pytest.fixture(autouse=True)
def hydrate_from_the_light_rows(monkeypatch):
"""Stand-in for the second read: the card for an image_id is the light
row plus a `barcode`, so tests can see hydration happened and where."""
def fake(table, image_ids):
return [{**dict(r), "barcode": f"890-{r['image_id']}"} for r in MARIE
if r["brand_table"] == table and r["image_id"] in image_ids]
monkeypatch.setattr(vector_store, "fetch_products_by_image_ids", fake)
@pytest.fixture
def calls(monkeypatch):
"""Patch the store; return the list of (kwargs) it was called with."""
seen: List[Dict[str, Any]] = []
def fake(vector, brand=None, top_k=10, category=None, ef_search=None):
seen.append({"brand": brand, "top_k": top_k, "category": category, "ef_search": ef_search})
return [dict(r) for r in MARIE]
monkeypatch.setattr(vector_store, "image_vector_search", fake)
return seen
# ---------------------------------------------------------------------------
# the vector
# ---------------------------------------------------------------------------
def test_a_unit_vector_passes_through_untouched():
v = _unit()
assert im.normalise_vector(v) == v
def test_an_unnormalised_vector_is_rescaled_to_unit_length():
out = im.normalise_vector([2.0] + [0.0] * 1023)
assert out[0] == 1.0 and abs(math.sqrt(sum(x * x for x in out)) - 1.0) < 1e-9
@pytest.mark.parametrize("bad, msg", [
([0.1] * 1023, "1024 values"),
([float("nan")] + [0.0] * 1023, "NaN"),
([0.0] * 1024, "all zeros"),
(["x"] * 1024, "list of numbers"),
])
def test_unsearchable_vectors_are_refused_with_a_reason(bad, msg):
with pytest.raises(im.InvalidVectorError, match=msg):
im.normalise_vector(bad)
# ---------------------------------------------------------------------------
# the label text
# ---------------------------------------------------------------------------
def test_sizes_are_normalised_and_words_cleaned():
words, sizes = im.tokens("Britannia Marie Gold 300 g Net Wt")
assert words == {"britannia", "marie", "gold"} and sizes == {"300g"}
@pytest.mark.parametrize("text, size", [("89gm", "89g"), ("89 GMS", "89g"), ("1 LTR", "1l"), ("1 litre", "1l"),
("2.50 kg", "2.5kg"), ("500ml", "500ml"), ("6 pcs", "6pcs")])
def test_every_way_a_label_writes_a_size_collapses_to_one_token(text, size):
assert im.tokens(text)[1] == {size}
def test_empty_text_has_no_tokens_and_no_overlap():
assert im.tokens(None) == (set(), set())
assert im.text_overlap(set(), set(), MARIE[0]) == 0.0
def test_size_matches_outweigh_word_matches():
words, sizes = im.tokens("Marie Gold 300 g")
assert im.text_overlap(words, sizes, MARIE[1]) == 3.0 + 2.0 # 300g + marie + gold
assert im.text_overlap(words, sizes, MARIE[0]) == 2.0 # marie + gold only
# ---------------------------------------------------------------------------
# ranking
# ---------------------------------------------------------------------------
def test_label_text_picks_the_pack_size_among_tied_siblings(calls):
result = im.search_by_vector(_unit(), text="Marie Gold 300 g")
names = [r["product_name"] for r in result.rows]
assert names == [
"Britannia Marie Gold 300g", # size token wins the tie
"Britannia Marie Gold 1kg", # then name order among the rest
"Britannia Marie Gold 89g",
"Britannia Good Day 100g", # lower score, whatever the text
]
assert result.rows[0]["score"] == pytest.approx(0.631)
assert result.rows[0]["text_overlap"] == 5.0
assert result.rows[0]["barcode"] == "890-britannia_marie_gold_300g" # hydrated, and to the right row
def test_without_text_ties_fall_back_to_name_order_and_overlap_is_zero(calls):
result = im.search_by_vector(_unit())
assert [r["product_name"] for r in result.rows][:3] == [
"Britannia Marie Gold 1kg", "Britannia Marie Gold 300g", "Britannia Marie Gold 89g",
]
assert all(r["text_overlap"] == 0.0 for r in result.rows)
def test_min_score_drops_everything_below_it(calls):
result = im.search_by_vector(_unit(), min_score=0.65)
assert result.rows == [] and result.min_score == 0.65
def test_top_k_truncates_after_ranking_and_only_winners_are_hydrated(calls, monkeypatch):
asked = []
monkeypatch.setattr(vector_store, "fetch_products_by_image_ids",
lambda table, ids: asked.append((table, sorted(ids))) or [])
result = im.search_by_vector(_unit(), text="300 g", top_k=1)
assert [r["product_name"] for r in result.rows] == ["Britannia Marie Gold 300g"]
assert result.top_k == 1
assert asked == [("brand_britannia", ["britannia_marie_gold_300g"])]
def test_a_row_whose_card_cannot_be_read_keeps_its_light_form(calls, monkeypatch):
monkeypatch.setattr(vector_store, "fetch_products_by_image_ids", lambda table, ids: [])
result = im.search_by_vector(_unit(), top_k=2)
assert len(result.rows) == 2 and "barcode" not in result.rows[0] and result.rows[0]["score"] == pytest.approx(0.631)
def test_duplicate_rows_from_the_same_table_are_collapsed(monkeypatch):
monkeypatch.setattr(vector_store, "image_vector_search",
lambda *a, **k: [dict(MARIE[0]), dict(MARIE[0])])
assert len(im.search_by_vector(_unit()).rows) == 1
# ---------------------------------------------------------------------------
# scope
# ---------------------------------------------------------------------------
def test_an_explicit_brand_is_a_hard_filter_that_never_falls_back(monkeypatch):
seen = []
monkeypatch.setattr(vector_store, "image_vector_search",
lambda vector, brand=None, **k: seen.append(brand) or [])
result = im.search_by_vector(_unit(), brand="Cadbury", text="britannia marie")
assert seen == ["Cadbury"]
assert result.rows == [] and result.scoped_to_brand and not result.scope_fallback
assert result.detected_brand == "Cadbury"
def test_a_brand_read_off_the_label_scopes_the_search(monkeypatch):
seen = []
monkeypatch.setattr(vector_store, "image_vector_search",
lambda vector, brand=None, **k: seen.append(brand) or [dict(r) for r in MARIE])
from app.services import query_intent
monkeypatch.setattr(query_intent, "extract_brand_mention", lambda text: "Britannia")
result = im.search_by_vector(_unit(), text="Britannia Marie Gold")
assert seen == ["Britannia"]
assert result.detected_brand == "Britannia" and result.scoped_to_brand and not result.scope_fallback
def test_an_ocr_brand_that_finds_nothing_retries_across_every_brand(monkeypatch):
seen = []
def fake(vector, brand=None, **k):
seen.append(brand)
return [] if brand else [dict(r) for r in MARIE]
monkeypatch.setattr(vector_store, "image_vector_search", fake)
from app.services import query_intent
monkeypatch.setattr(query_intent, "extract_brand_mention", lambda text: "Britannia")
result = im.search_by_vector(_unit(), text="Britannia marie")
assert seen == ["Britannia", None]
assert result.scope_fallback and not result.scoped_to_brand and len(result.rows) == 4
def test_no_text_and_no_brand_means_one_unscoped_query(calls):
im.search_by_vector(_unit())
assert len(calls) == 1 and calls[0]["brand"] is None
def test_brand_detection_failures_do_not_break_the_search(calls, monkeypatch):
from app.services import query_intent
def boom(text):
raise RuntimeError("brand index unavailable")
monkeypatch.setattr(query_intent, "extract_brand_mention", boom)
result = im.search_by_vector(_unit(), text="something")
assert result.detected_brand is None and len(result.rows) == 4
# ---------------------------------------------------------------------------
# candidate width and ef_search
# ---------------------------------------------------------------------------
@pytest.mark.parametrize("top_k, fetch_k, ef", [(1, 30, 40), (10, 30, 40), (20, 60, 60), (50, 100, 100), (500, 100, 100)])
def test_the_store_is_asked_for_enough_candidates_to_hold_the_ties(calls, top_k, fetch_k, ef):
im.search_by_vector(_unit(), top_k=top_k)
assert calls[0]["top_k"] == fetch_k and calls[0]["ef_search"] == ef
# ---------------------------------------------------------------------------
# the SQL
# ---------------------------------------------------------------------------
class _Cursor:
def __init__(self):
self.statements: List[str] = []
self.description = None
self._pending: Any = None
def execute(self, sql, params=None):
text = " ".join(str(sql).split())
self.statements.append(text)
if "information_schema.columns" in text:
self._pending = [("brand_x", c) for c in ("id", "product_name", "embedding", "img_vector", "img_vector_src")]
elif "information_schema.tables" in text and "EXISTS" in text:
self._pending = (True,)
elif "information_schema.tables" in text:
self._pending = [("brand_x",)]
elif text.startswith("SET"):
self._pending = None
elif "AS distance" in text:
self.description = [("image_id",), ("product_name",), ("title",), ("size_variants",), ("distance",)]
self._pending = [("a", "P", "P", [], 0.2)]
else:
self.description = [("id",), ("product_name",), ("img_vector_src",)]
self._pending = [(1, "P", None)]
def fetchall(self):
return self._pending if isinstance(self._pending, list) else []
def fetchone(self):
return self._pending if isinstance(self._pending, tuple) else None
def __enter__(self):
return self
def __exit__(self, *a):
return False
class _Conn:
def __init__(self, cur):
self._cur = cur
def cursor(self):
return self._cur
def __enter__(self):
return self
def __exit__(self, *a):
return False
def close(self):
pass
def test_the_query_reads_img_vector_by_name_and_sets_ef_search(monkeypatch):
vector_store._invalidate_product_columns_cache()
cur = _Cursor()
monkeypatch.setattr(vector_store, "_connect", lambda: _Conn(cur))
rows = vector_store.image_vector_search([0.0] * 1024, top_k=30, ef_search=40)
select = [s for s in cur.statements if s.startswith("SELECT") and "FROM brand_x" in s and "distance" in s]
assert select, cur.statements
sql = select[0]
# Light projection only: ranking must not drag 7KB product cards per candidate.
assert sql.startswith("SELECT image_id, product_name, title, size_variants, img_vector <=> %s::vector AS distance")
assert "WHERE img_vector IS NOT NULL" in sql and sql.endswith("ORDER BY distance ASC LIMIT %s")
assert "embedding" not in sql and "description" not in sql
assert cur.statements.index("SET hnsw.ef_search = 40") < cur.statements.index(sql)
assert rows[0]["brand_table"] == "brand_x" and rows[0]["brand"] == "X" and rows[0]["distance"] == 0.2
def test_hydration_reads_full_cards_by_image_id_without_the_vectors(monkeypatch):
vector_store._invalidate_product_columns_cache()
cur = _Cursor()
monkeypatch.setattr(vector_store, "_connect", lambda: _Conn(cur))
real_fetch_products_by_image_ids("brand_x", ["a", "b"]) # the autouse fixture patches the module attribute
sql = [s for s in cur.statements if "WHERE image_id = ANY(%s)" in s]
assert sql and sql[0].startswith('SELECT "id", "product_name", "img_vector_src" FROM brand_x')
assert "embedding" not in sql[0] and '"img_vector"' not in sql[0]
assert real_fetch_products_by_image_ids("brand_x", []) == []
def test_the_queries_are_never_select_star():
assert "SELECT *" not in inspect.getsource(vector_store.image_vector_search)
assert "SELECT *" not in inspect.getsource(vector_store.fetch_products_by_image_ids)
def test_a_category_filter_is_added_inside_the_where(monkeypatch):
vector_store._invalidate_product_columns_cache()
cur = _Cursor()
monkeypatch.setattr(vector_store, "_connect", lambda: _Conn(cur))
vector_store.image_vector_search([0.0] * 1024, brand="X", category="Biscuits")
sql = [s for s in cur.statements if "AS distance" in s][0]
assert "WHERE img_vector IS NOT NULL AND category ILIKE %s ORDER BY" in sql

View File

@@ -0,0 +1,174 @@
"""POST /api/search/image-vector and POST /api/search/image - the HTTP contract.
The ranking is tested in tests/test_image_match.py; here `search_by_vector` is
patched on the router module and the assertions are about status codes,
validation and the response shape the app reads. Both routes are public.
"""
from __future__ import annotations
import io
import math
import pytest
from app.api.routers import search as search_router
from app.services.image_match import ImageSearchResult
def _unit():
v = [math.cos(i / 7.0) for i in range(1024)]
n = math.sqrt(sum(x * x for x in v))
return [x / n for x in v]
def _row():
return {
"image_id": "britannia_marie_gold_300g",
"product_name": "Britannia Marie Gold 300g",
"title": "Britannia Marie Gold 300g",
"brand": "Britannia",
"brand_table": "brand_britannia",
"category": "Biscuits",
"image_url": "https://cdn.example/marie.jpg",
"image_urls": ["https://cdn.example/marie.jpg"],
"size_variants": ["300g"],
"barcode": "8901063010512",
"barcode_type": "EAN13",
"final_selling_price": 45.0,
"selling_price": 42.0,
"hsn_code": "1905",
"fssai_license": "10012021000123",
"distance": 0.369,
"score": 0.631,
"text_overlap": 5.0,
}
@pytest.fixture
def fake_search(monkeypatch):
calls = []
def fake(vector, text=None, brand=None, category=None, top_k=10, min_score=0.0):
calls.append({"vector": list(vector), "text": text, "brand": brand,
"category": category, "top_k": top_k, "min_score": min_score})
return ImageSearchResult(rows=[_row()], detected_brand="Britannia", scoped_to_brand=True,
min_score=min_score, query_text=text, top_k=top_k)
monkeypatch.setattr(search_router, "search_by_vector", fake)
return calls
# ---------------------------------------------------------------------------
# /search/image-vector
# ---------------------------------------------------------------------------
def test_a_vector_returns_the_product_card_with_its_score(client, fake_search):
res = client.post("/api/search/image-vector",
json={"vector": _unit(), "text": "Britannia Marie Gold 300 g", "top_k": 5})
assert res.status_code == 200, res.text
body = res.json()
assert body["total"] == 1 and body["detected_brand"] == "Britannia" and body["scoped_to_brand"] is True
hit = body["results"][0]
assert hit["product_name"] == "Britannia Marie Gold 300g"
assert hit["score"] == 0.631 and hit["text_overlap"] == 5.0
for key in ("barcode", "final_selling_price", "selling_price", "category", "image_url",
"size_variants", "hsn_code", "fssai_license", "image_id", "brand"):
assert key in hit, key
assert hit["barcode"] == "8901063010512" and hit["image_url"] == "https://cdn.example/marie.jpg"
assert fake_search[0]["text"] == "Britannia Marie Gold 300 g" and fake_search[0]["top_k"] == 5
def test_the_route_is_public(client, fake_search):
assert client.post("/api/search/image-vector", json={"vector": _unit()}).status_code == 200
@pytest.mark.parametrize("payload, fragment", [
({"vector": [0.1] * 1023}, "vector"),
({"vector": [0.1] * 1025}, "vector"),
({"vector": [0.0] * 1024}, "all zeros"),
({"vector": _unit(), "top_k": 51}, "top_k"),
({"vector": _unit(), "top_k": 0}, "top_k"),
({"vector": _unit(), "min_score": 1.5}, "min_score"),
({}, "vector"),
])
def test_bad_requests_are_422_and_name_the_field(client, fake_search, payload, fragment):
res = client.post("/api/search/image-vector", json=payload)
assert res.status_code == 422
assert fragment in res.text
assert fake_search == []
def test_a_nan_in_the_vector_is_422(client, fake_search):
body = '{"vector": [' + ",".join(["NaN"] + ["0.1"] * 1023) + "]}"
res = client.post("/api/search/image-vector", content=body, headers={"content-type": "application/json"})
assert res.status_code == 422 and fake_search == []
def test_a_service_rejection_is_422_not_500(client, monkeypatch):
from app.services.image_match import InvalidVectorError
def refuse(*a, **k):
raise InvalidVectorError("vector is all zeros")
monkeypatch.setattr(search_router, "search_by_vector", refuse)
res = client.post("/api/search/image-vector", json={"vector": _unit()})
assert res.status_code == 422 and "all zeros" in res.text
# ---------------------------------------------------------------------------
# /search/image
# ---------------------------------------------------------------------------
def _post_image(client, data: bytes, **form):
return client.post("/api/search/image", files={"file": ("photo.jpg", io.BytesIO(data), "image/jpeg")}, data=form)
def test_a_photo_is_embedded_and_searched_with_its_form_fields(client, fake_search, monkeypatch):
monkeypatch.setattr(search_router.image_embedder, "available", lambda: True)
monkeypatch.setattr(search_router.image_embedder, "embedding_for_bytes", lambda data: _unit())
res = _post_image(client, b"\xff\xd8" + b"x" * 5000, text="Marie Gold 300 g", top_k="3", min_score="0.2")
assert res.status_code == 200, res.text
assert res.json()["results"][0]["product_name"] == "Britannia Marie Gold 300g"
call = fake_search[0]
assert call["text"] == "Marie Gold 300 g" and call["top_k"] == 3 and call["min_score"] == 0.2
assert len(call["vector"]) == 1024
def test_without_a_model_the_photo_route_says_503_and_points_at_the_vector_route(client, fake_search, monkeypatch):
monkeypatch.setattr(search_router.image_embedder, "available", lambda: False)
res = _post_image(client, b"x" * 5000)
assert res.status_code == 503 and "/api/search/image-vector" in res.text
assert fake_search == []
def test_an_oversized_photo_is_413_before_any_model_work(client, fake_search, monkeypatch):
monkeypatch.setattr(search_router, "IMAGE_VECTOR_MAX_BYTES", 100)
monkeypatch.setattr(search_router.image_embedder, "available", lambda: pytest.fail("must not load the model"))
res = _post_image(client, b"x" * 101)
assert res.status_code == 413 and "limit is" in res.text
def test_an_empty_upload_is_400(client, fake_search):
assert _post_image(client, b"").status_code == 400
def test_an_undecodable_photo_is_422(client, fake_search, monkeypatch):
monkeypatch.setattr(search_router.image_embedder, "available", lambda: True)
monkeypatch.setattr(search_router.image_embedder, "embedding_for_bytes", lambda data: None)
res = _post_image(client, b"not an image at all" * 100)
assert res.status_code == 422 and "decode" in res.text and fake_search == []
def test_both_routes_are_documented(client):
paths = client.get("/openapi.json").json()["paths"]
assert "/api/search/image-vector" in paths and "/api/search/image" in paths
assert "post" in paths["/api/search/image"] and "get" in paths["/api/search"]

View File

@@ -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]"