imag vector generation with dimentionality reduction
This commit is contained in:
@@ -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()),
|
||||
)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@@ -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
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
27
app/main.py
27
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.
|
||||
|
||||
@@ -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
296
app/services/image_match.py
Normal 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,
|
||||
)
|
||||
33
app/services/models/mobilenet/README.md
Normal file
33
app/services/models/mobilenet/README.md
Normal 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`.
|
||||
BIN
app/services/models/mobilenet/mobilenet_v3_small_embedder.tflite
Normal file
BIN
app/services/models/mobilenet/mobilenet_v3_small_embedder.tflite
Normal file
Binary file not shown.
@@ -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
141
docs/IMAGE_SEARCH_API.md
Normal 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`.
|
||||
@@ -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.
|
||||
|
||||
184
scripts/export_mobilenet_embedder.py
Normal file
184
scripts/export_mobilenet_embedder.py
Normal 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
352
tests/test_image_match.py
Normal 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
|
||||
174
tests/test_image_search_api.py
Normal file
174
tests/test_image_search_api.py
Normal 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"]
|
||||
@@ -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]"
|
||||
|
||||
|
||||
Reference in New Issue
Block a user