diff --git a/.env.example b/.env.example index 329b5cd..17c4cdc 100644 --- a/.env.example +++ b/.env.example @@ -314,12 +314,17 @@ USE_PLAYWRIGHT_FALLBACK=true MIN_IMAGE_BYTES=3000 -# img_vector: a 32x32 RGB pixel thumbnail of each product's primary image, -# stored on its brand table and computed on a background thread after every -# catalog write (app/services/image_vector.py). Existing rows are filled by +# img_vector: a MobileNetV3-Small embedding (1024 floats, L2-normalised) of +# each product's primary image, stored on its brand table and computed on a +# background thread after every catalog write (app/services/image_vector.py, +# model in app/services/image_embedder.py). Existing rows are filled by # `python -m scripts.backfill_image_vectors --all --apply`. Set false to stop # the API process from downloading images at all. ENABLE_IMAGE_VECTORS=true +# Where the .tflite lives. The default is inside app/ (shipped with the image +# and never volume-mounted); only override to point at a different file. +#IMAGE_EMBED_MODEL_PATH=/app/app/services/models/mobilenet/mobilenet_v3_small_embedder.tflite +#IMAGE_EMBED_NUM_THREADS=2 # Product SKU: try a live web search for a real marketplace product ID # (Amazon ASIN, Flipkart PID, etc.) before falling back to an internal SKU. diff --git a/.gitattributes b/.gitattributes index a2f4b3e..2ec701b 100644 --- a/.gitattributes +++ b/.gitattributes @@ -15,6 +15,7 @@ # to the wildcard above cannot silently start mangling them. *.joblib binary *.db binary +*.tflite binary # Windows-only scripts genuinely need CRLF. None are tracked in this repo # today (start_app.bat lives above it), but the rule belongs with the policy. diff --git a/app/infrastructure/settings.py b/app/infrastructure/settings.py index d23fb9b..7abcfbf 100644 --- a/app/infrastructure/settings.py +++ b/app/infrastructure/settings.py @@ -381,27 +381,41 @@ USE_PLAYWRIGHT_FALLBACK = _bool("USE_PLAYWRIGHT_FALLBACK", "true") MIN_IMAGE_BYTES = int(os.getenv("MIN_IMAGE_BYTES", "3000")) # --------------------------------------------------------------------------- -# img_vector - a 32x32 RGB pixel thumbnail of each product's primary image, -# stored on its brand table (app/services/image_vector.py) +# img_vector - a MobileNetV3-Small image embedding (1024 floats, L2-normalised) +# of each product's primary image, stored on its brand table +# (app/services/image_vector.py owns the column, image_embedder.py the model) # --------------------------------------------------------------------------- # Computed off the request path by one bounded worker thread after every # catalog write, and by scripts/backfill_image_vectors.py for existing rows. -# Default on: the work is one small download per row just written, never on -# a request. tests/conftest.py pins it OFF so the suite never dials the -# database named in a developer's .env. +# Default on: the work is one small download and one ~30ms inference per row +# just written, never on a request. tests/conftest.py pins it OFF so the +# suite never dials the database named in a developer's .env. ENABLE_IMAGE_VECTORS = _bool("ENABLE_IMAGE_VECTORS", "true") # A download past this many bytes is abandoned - a wrong URL to a video must # not fill the container. IMAGE_VECTOR_MAX_BYTES = int(os.getenv("IMAGE_VECTOR_MAX_BYTES", str(8 * 1024 * 1024))) # Refused before decoding when the header claims more pixels than this -# (decompression-bomb guard; 40MP is well past any product photo). -IMAGE_VECTOR_MAX_PIXELS = int(os.getenv("IMAGE_VECTOR_MAX_PIXELS", "40000000")) +# (decompression-bomb guard; 25MP is well past any product photo, and the +# embedder decodes at full resolution - ~75MB of RGB at this cap). +IMAGE_VECTOR_MAX_PIXELS = int(os.getenv("IMAGE_VECTOR_MAX_PIXELS", "25000000")) IMAGE_VECTOR_TIMEOUT_SECONDS = float(os.getenv("IMAGE_VECTOR_TIMEOUT_SECONDS", "15")) # Minimum gap between two requests to the same image host. IMAGE_VECTOR_HOST_PAUSE_SECONDS = float(os.getenv("IMAGE_VECTOR_HOST_PAUSE_SECONDS", "0.5")) # Writes waiting for the worker; beyond this a write's rows are left for the # backfill script rather than queued. IMAGE_VECTOR_QUEUE_MAX = int(os.getenv("IMAGE_VECTOR_QUEUE_MAX", "64")) +# The TFLite embedder (mobilenet_v3_small_embedder.tflite, input [1,224,224,3] +# float32, output [1,1024]). Lives under app/, NOT data/: /app/data is a named +# volume on every deployment (see BUNDLED_ASSETS_DIR), and a file added to +# the image under a mounted path is invisible on any volume that already +# exists. app/ is copied into the image and never mounted. +IMAGE_EMBED_MODEL_PATH = _dir( + "IMAGE_EMBED_MODEL_PATH", + _BACKEND_ROOT / "app" / "services" / "models" / "mobilenet" / "mobilenet_v3_small_embedder.tflite", +) +# Intra-op threads for one inference. 2 on the prod host; inference itself is +# serialised by a lock (a TFLite interpreter is not thread-safe). +IMAGE_EMBED_NUM_THREADS = int(os.getenv("IMAGE_EMBED_NUM_THREADS", "2")) # --------------------------------------------------------------------------- # USDA FoodData Central - nutrition for loose, unbranded commodities diff --git a/app/services/image_embedder.py b/app/services/image_embedder.py new file mode 100644 index 0000000..1c12fab --- /dev/null +++ b/app/services/image_embedder.py @@ -0,0 +1,250 @@ +"""The image embedding model behind `img_vector`: MobileNetV3-Small, TFLite. + +WHAT IT PRODUCES +---------------- +One 1024-float vector per image, L2-normalised, from +`mobilenet_v3_small_embedder.tflite` (input [1, 224, 224, 3] float32, output +[1, 1024]). It is the same model, runtime and post-processing a colleague uses +for their product images, which is the whole point: two vectors from the same +photo must agree, and `<=>` (cosine distance) between our rows and theirs must +mean something. + +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. +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 + image-vector worker thread and the backfill script's download pool. + +Which column gets the vector, when, and for which rows is +app/services/image_vector.py's business, not this file's. + +FAILURE IS SILENT BY DESIGN +--------------------------- +The runtime (`ai-edge-litert`) is imported here, not at module import, and +the model file is opened here, not at boot. If either is missing, `available()` +returns False after ONE warning, `embed()` returns None, and the catalog keeps +writing rows with a NULL img_vector. A missing wheel or a model left out of an +image must never turn into a failed product write. +""" +from __future__ import annotations + +import io +import logging +import threading +import warnings +from typing import Any, List, Optional + +import numpy as np + +from app.infrastructure.settings import ( + IMAGE_EMBED_MODEL_PATH, + IMAGE_EMBED_NUM_THREADS, + IMAGE_VECTOR_MAX_PIXELS, +) + +logger = logging.getLogger(__name__) + +EMBED_DIM = 1024 +IMAGE_EMBED_INPUT_SIZE = 224 +_INPUT_SHAPE = (1, IMAGE_EMBED_INPUT_SIZE, IMAGE_EMBED_INPUT_SIZE, 3) +_OUTPUT_SHAPE = (1, EMBED_DIM) + +_lock = threading.Lock() +_interpreter: Any = None +_input_index: Optional[int] = None +_output_index: Optional[int] = None +_disabled_reason: Optional[str] = None +_warned = False + + +# --------------------------------------------------------------------------- +# bytes -> input tensor +# --------------------------------------------------------------------------- + +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. + ========================================================================== + + 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 + DecompressionBombWarning is one of the things the broad except absorbs. + """ + if not image_bytes: + return None + try: + from PIL import Image, ImageOps + + 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) + return None + im = ImageOps.exif_transpose(im) or im + + 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) + + arr = np.asarray(im, dtype=np.float32) # (224, 224, 3), 0..255 + arr = np.expand_dims(arr, axis=0) # (1, 224, 224, 3) + if arr.shape != _INPUT_SHAPE: + logger.debug("image rejected: tensor shape %s", arr.shape) + return None + return np.ascontiguousarray(arr) + except Exception as exc: # noqa: BLE001 - see docstring + logger.debug("image undecodable: %s", exc) + return None + + +def l2_normalize(vec: np.ndarray) -> Optional[np.ndarray]: + """Unit-length copy of `vec`, or None for a (near-)zero vector.""" + arr = np.asarray(vec, dtype=np.float32).reshape(-1) + norm = float(np.linalg.norm(arr)) + if not np.isfinite(norm) or norm < 1e-12: + return None + return arr / norm + + +# --------------------------------------------------------------------------- +# the interpreter +# --------------------------------------------------------------------------- + +def _disable(reason: str) -> None: + """Record why the embedder is off and say so once. Caller holds _lock.""" + global _disabled_reason, _interpreter, _warned + _disabled_reason = reason + _interpreter = None + if not _warned: + _warned = True + logger.warning("image embedder disabled: %s (img_vector stays NULL until fixed)", reason) + + +def _load() -> None: + """Create and validate the interpreter. Caller holds _lock. Never raises.""" + global _interpreter, _input_index, _output_index + path = IMAGE_EMBED_MODEL_PATH + if not path.is_file(): + _disable(f"model file not found at {path}") + return + try: + with warnings.catch_warnings(): + # The wheel's import-time deprecation chatter would become a hard + # error under pytest's filterwarnings=error. + warnings.simplefilter("ignore") + from ai_edge_litert.interpreter import Interpreter + except Exception as exc: # noqa: BLE001 - ImportError or a broken wheel + _disable(f"ai-edge-litert is not importable ({exc})") + return + try: + interp = Interpreter(model_path=str(path), num_threads=max(1, IMAGE_EMBED_NUM_THREADS)) + interp.allocate_tensors() + inputs = interp.get_input_details() + outputs = interp.get_output_details() + except Exception as exc: # noqa: BLE001 - unreadable / malformed model + _disable(f"could not load {path.name} ({exc})") + return + + if len(inputs) != 1 or tuple(int(d) for d in inputs[0]["shape"]) != _INPUT_SHAPE \ + or np.dtype(inputs[0]["dtype"]) != np.float32: + got = [(list(i["shape"]), np.dtype(i["dtype"]).name) for i in inputs] + _disable(f"{path.name} input is {got}, expected [1,224,224,3] float32") + return + if len(outputs) != 1 or tuple(int(d) for d in outputs[0]["shape"]) != _OUTPUT_SHAPE: + got = [list(o["shape"]) for o in outputs] + _disable(f"{path.name} output is {got}, expected [1,1024]") + return + + _interpreter = interp + _input_index = int(inputs[0]["index"]) + _output_index = int(outputs[0]["index"]) + logger.info("image embedder ready: %s (%d thread(s))", path.name, IMAGE_EMBED_NUM_THREADS) + + +def _ensure_loaded() -> bool: + """Caller holds _lock.""" + if _interpreter is None and _disabled_reason is None: + _load() + return _interpreter is not None + + +def available() -> bool: + """True when the model can be used. Loads it on first call.""" + with _lock: + return _ensure_loaded() + + +def describe() -> str: + """One line for script banners: where the model is and whether it works.""" + with _lock: + ok = _ensure_loaded() + if ok: + return f"ready - {IMAGE_EMBED_MODEL_PATH} ({IMAGE_EMBED_NUM_THREADS} thread(s), {EMBED_DIM}-d)" + return f"disabled - {_disabled_reason}" + + +def embed(preprocessed: np.ndarray) -> Optional[np.ndarray]: + """Raw model output for one preprocessed tensor, as (1024,) float32, or None. + + Serialised on the module lock: the interpreter's tensors are shared + state, and two threads calling set_tensor/invoke at once corrupt both. + """ + if preprocessed is None or tuple(preprocessed.shape) != _INPUT_SHAPE: + return None + with _lock: + if not _ensure_loaded(): + return None + try: + _interpreter.set_tensor(_input_index, np.asarray(preprocessed, dtype=np.float32)) + _interpreter.invoke() + out = _interpreter.get_tensor(_output_index) + return np.array(out, dtype=np.float32, copy=True).reshape(-1)[:EMBED_DIM] + except Exception as exc: # noqa: BLE001 - one bad tensor must not stop the worker + logger.debug("inference failed: %s", exc) + return None + + +def embedding_for_bytes(image_bytes: bytes) -> Optional[List[float]]: + """bytes -> preprocess -> embed -> L2-normalise -> 1024 floats, or None.""" + tensor = preprocess(image_bytes) + if tensor is None: + return None + raw = embed(tensor) + if raw is None or raw.shape != (EMBED_DIM,): + return None + unit = l2_normalize(raw) + if unit is None: + return None + return [float(v) for v in unit] + + +def _reset() -> None: + """Tests only: forget the interpreter and any recorded failure.""" + global _interpreter, _input_index, _output_index, _disabled_reason, _warned + with _lock: + _interpreter = None + _input_index = None + _output_index = None + _disabled_reason = None + _warned = False diff --git a/app/services/image_vector.py b/app/services/image_vector.py index fe02cc5..6d6164c 100644 --- a/app/services/image_vector.py +++ b/app/services/image_vector.py @@ -1,18 +1,26 @@ -"""A pixel thumbnail of every product's PRIMARY image, stored on its brand table. +"""An image embedding of every product's PRIMARY image, stored on its brand table. WHAT `img_vector` IS -------------------- The image the product card shows - `image_url`, else `image_urls[0]`, which is -exactly the order `frontend/src/components/ProductCard.jsx` walks - decoded, -EXIF-rotated, flattened onto white, resized to 32x32 RGB and unrolled into -3072 uint8 values. It is a literal downscaled picture, not a learned embedding: -two rows with the same photo have (near-)identical vectors, and a row whose -image changed has a different one. Stored as pgvector `vector(3072)` so that a -`<=>` comparison later needs no migration; 0..255 is exact in float4. +exactly the order `frontend/src/components/ProductCard.jsx` walks - run through +MobileNetV3-Small (`app/services/image_embedder.py`): 1024 floats, L2-normalised, +stored as pgvector `vector(1024)` with an hnsw cosine index. It is a learned +embedding, not a picture: two photos of the same product from different angles +land near each other, and `<=>` between rows is a visual-similarity search. The +model, runtime and post-processing are the same a colleague uses for their +product images, so our vectors and theirs are directly comparable. + +(Until 2026-09-17 this column held a 32x32 pixel thumbnail as `vector(3072)`. +`vector_store._ensure_img_vector_type` drops and re-creates a column of the +wrong dimension on the next write, discarding those vectors; the backfill +script fills the new ones.) `img_vector_src` records the URL the vector was computed from. That is what makes recomputation idempotent: a row is due when it has no vector, or when -its current primary URL differs from the one the vector came from. +its current primary URL differs from the one the vector came from. A change +to the model or its preprocessing is neither; `backfill_rows(force=True)` / +`--force` recomputes every row for that. WHY THIS MODULE, AND NOT THE UPSERT ----------------------------------- @@ -31,14 +39,12 @@ and an overflowing queue simply leaves rows for `scripts/backfill_image_vectors.py`. Same shape as app/core/batch_worker.py, for the same reasons. -Pillow, not OpenCV: it is already a dependency, `Image.draft()` decodes a JPEG -at 1/8 scale for a target this small, and it honours EXIF orientation on -in-memory bytes. OpenCV would add ~130MB on disk and always decode at full -resolution, on a container that already holds torch. +The worker is also the only thread in the API process that runs the model, +which matters because a TFLite interpreter is not thread-safe; the embedder +serialises every call on its own lock regardless. """ from __future__ import annotations -import io import logging import queue import threading @@ -50,16 +56,15 @@ from app.infrastructure.settings import ( ENABLE_IMAGE_VECTORS, IMAGE_VECTOR_HOST_PAUSE_SECONDS, IMAGE_VECTOR_MAX_BYTES, - IMAGE_VECTOR_MAX_PIXELS, IMAGE_VECTOR_QUEUE_MAX, IMAGE_VECTOR_TIMEOUT_SECONDS, MIN_IMAGE_BYTES, ) +from app.services import image_embedder logger = logging.getLogger(__name__) -IMG_VECTOR_SIZE = 32 -IMG_VECTOR_DIM = IMG_VECTOR_SIZE * IMG_VECTOR_SIZE * 3 # 3072 +IMG_VECTOR_DIM = image_embedder.EMBED_DIM # 1024 IMG_VECTOR_COLUMNS = ("img_vector", "img_vector_src") _IMAGE_EXTENSIONS = (".jpg", ".jpeg", ".png", ".webp", ".gif", ".bmp") @@ -93,54 +98,12 @@ def primary_image_url(row: Mapping[str, Any]) -> Optional[str]: # --------------------------------------------------------------------------- # Bytes -> vector # --------------------------------------------------------------------------- +# The model, its preprocessing and its lock live in image_embedder.py; this +# module only decides which bytes go in and where the result goes. -def pixel_vector(image_bytes: bytes) -> Optional[List[int]]: - """32x32 RGB pixels of `image_bytes` as 3072 ints in 0..255, or None. - - None for anything Pillow cannot turn into a picture: corrupt or truncated - data, an unknown format, or a header claiming more than - IMAGE_VECTOR_MAX_PIXELS (checked before decoding, so a decompression bomb - costs a header read and nothing more). Never raises - pytest runs with - warnings as errors, and a DecompressionBombWarning is one of the things - the broad except is there to absorb. - """ - if not image_bytes: - return None - try: - from PIL import Image, ImageOps - - 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) - return None - # JPEG only: ask the decoder for a DCT-scaled decode no smaller than - # 4x the target. A no-op for every other format. - try: - im.draft("RGB", (IMG_VECTOR_SIZE * 4, IMG_VECTOR_SIZE * 4)) - except Exception: # noqa: BLE001 - draft is an optimisation, not a requirement - pass - im = ImageOps.exif_transpose(im) or im - - has_alpha = "A" in im.getbands() or "transparency" in im.info - if has_alpha: - 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((IMG_VECTOR_SIZE, IMG_VECTOR_SIZE), Image.Resampling.BILINEAR) - values = list(im.tobytes()) - if len(values) != IMG_VECTOR_DIM: - logger.debug("image rejected: %d values, expected %d", len(values), IMG_VECTOR_DIM) - return None - return values - except Exception as exc: # noqa: BLE001 - see docstring - logger.debug("image undecodable: %s", exc) - return None - - -def to_pg(vector: Iterable[int]) -> str: +def to_pg(vector: Iterable[float]) -> str: """The text form pgvector accepts, identical to what the upsert uses for `embedding`.""" - return "[" + ",".join(str(int(v)) for v in vector) + "]" + return "[" + ",".join(str(float(v)) for v in vector) + "]" # --------------------------------------------------------------------------- @@ -230,12 +193,12 @@ def download_image_bytes(url: str, timeout: Optional[float] = None) -> Optional[ return None -def vector_for_url(url: str, timeout: Optional[float] = None) -> Tuple[Optional[List[int]], str]: - """(vector or None, the URL it was computed from).""" +def vector_for_url(url: str, timeout: Optional[float] = None) -> Tuple[Optional[List[float]], str]: + """(L2-normalised 1024-float embedding or None, the URL it was computed from).""" data = download_image_bytes(url, timeout=timeout) if data is None: return None, url - return pixel_vector(data), url + return image_embedder.embedding_for_bytes(data), url # --------------------------------------------------------------------------- @@ -264,6 +227,7 @@ def rows_needing_vectors( *, recompute_stale: bool = True, limit: Optional[int] = None, + force: bool = False, ) -> List[Dict[str, Any]]: """Rows that have a primary image and no current vector for it. @@ -271,17 +235,27 @@ def rows_needing_vectors( `image_url` was rewritten by a later upsert is stale and is included when `recompute_stale` is on. Rows with no usable primary are dropped here, so the caller never spends a request on them. + + `force` takes every row with a primary image, current or not. That is the + switch for a model or preprocessing change: the URL did not move, so + nothing else would notice that every stored vector is now the wrong one. """ sql = ( f"SELECT image_id, product_name, image_url, image_urls, img_vector_src, " - f"(img_vector IS NULL) AS missing FROM {table_name} " - f"WHERE (img_vector IS NULL OR img_vector_src IS DISTINCT FROM " - f"COALESCE(NULLIF(image_url, ''), image_urls[1]))" + f"(img_vector IS NULL) AS missing FROM {table_name}" ) + clauses: List[str] = [] params: List[Any] = [] + if not force: + clauses.append( + "(img_vector IS NULL OR img_vector_src IS DISTINCT FROM " + "COALESCE(NULLIF(image_url, ''), image_urls[1]))" + ) if image_ids is not None: - sql += " AND image_id = ANY(%s)" + clauses.append("image_id = ANY(%s)") params.append(list(image_ids)) + if clauses: + sql += " WHERE " + " AND ".join(clauses) sql += " ORDER BY updated_at DESC" if limit is not None: sql += " LIMIT %s" @@ -295,16 +269,17 @@ def rows_needing_vectors( primary = primary_image_url(row) if not primary: continue - if not row.get("missing") and not recompute_stale: - continue - if not row.get("missing") and row.get("img_vector_src") == primary: - continue + if not force: + if not row.get("missing") and not recompute_stale: + continue + if not row.get("missing") and row.get("img_vector_src") == primary: + continue row["primary_url"] = primary out.append(row) return out -def store_vector(cur, table_name: str, image_id: str, vector: List[int], src: str) -> None: +def store_vector(cur, table_name: str, image_id: str, vector: List[float], src: str) -> None: """Write one vector. Deliberately leaves `updated_at` alone: the catalog listing orders on it, and a background enrichment must not reshuffle it.""" cur.execute( @@ -321,11 +296,15 @@ def backfill_rows( limit: Optional[int] = None, dry_run: bool = False, pause_seconds: Optional[float] = None, + force: bool = False, ) -> Dict[str, int]: """Compute and store vectors for one brand, sequentially. Never raises. - Returns counts: candidates, computed, failed, skipped (no columns yet). - `pause_seconds` is the gap between two requests to the same host. + Returns counts: candidates, computed, failed, skipped (1 when the table + has no img_vector columns yet, or the embedder is unavailable - in both + cases nothing is downloaded). `pause_seconds` is the gap between two + requests to the same host; `force` recomputes rows that already have a + current vector. """ from app.services.vector_store import _connect, _table_name @@ -343,8 +322,14 @@ def backfill_rows( logger.info("image vectors: %s has no img_vector columns yet - skipped", table_name) totals["skipped"] = 1 return totals + if not image_embedder.available(): + # available() already warned once with the reason. Rows stay + # NULL and are picked up by the next write or the script. + logger.info("image vectors: embedder unavailable - %s left for later", table_name) + totals["skipped"] = 1 + return totals rows = rows_needing_vectors( - cur, table_name, image_ids, recompute_stale=recompute_stale, limit=limit + cur, table_name, image_ids, recompute_stale=recompute_stale, limit=limit, force=force ) totals["candidates"] = len(rows) for row in rows: diff --git a/app/services/vector_store.py b/app/services/vector_store.py index 6c3b23d..195aa12 100644 --- a/app/services/vector_store.py +++ b/app/services/vector_store.py @@ -208,13 +208,19 @@ def get_brand_table_ddl(brand: str) -> str: -- Vector embedding for search embedding vector(384), - -- 32x32 RGB pixel thumbnail of the PRIMARY image (image_url, else - -- image_urls[0]) as 3072 uint8 values, and the URL it was computed - -- from. Written ONLY by app/services/image_vector.py via targeted - -- UPDATEs - never by upsert_brand_products, for the same reason as - -- nutrition_score above. Not indexed: pgvector caps vector indexes - -- at 2000 dims, and an exact scan is fine at this catalog's size. - img_vector vector(3072), + -- MobileNetV3-Small embedding of the PRIMARY image (image_url, else + -- image_urls[0]): 1024 floats, L2-normalised, and the URL it was + -- computed from. Written ONLY by app/services/image_vector.py via + -- targeted UPDATEs - never by upsert_brand_products, for the same + -- reason as nutrition_score above. + -- + -- Its hnsw index is deliberately NOT declared here. This DDL runs + -- before _ensure_columns on every write, and a table still carrying + -- the previous vector(3072) column would make CREATE INDEX fail + -- (2000-dim cap) and take the whole write with it. The index is + -- created by _ensure_img_vector_type, after the column is known to + -- be vector(1024). + img_vector vector(1024), img_vector_src TEXT, -- Timestamps @@ -257,6 +263,82 @@ def _connect() -> Optional[psycopg.Connection]: # fail the write - see the loop in _ensure_columns. _TYPE_OPTIONAL_COLUMNS = frozenset({"img_vector"}) +# The embedder's output width. col_defs above must declare img_vector as +# exactly vector(IMG_VECTOR_DIMS) - a test pins the two together. +IMG_VECTOR_DIMS = 1024 + + +def _ensure_img_vector_type(cur, table_name: str) -> None: + """Self-heal img_vector to vector(IMG_VECTOR_DIMS) and give it an hnsw index. + + The column first shipped as vector(3072) (a pixel thumbnail) and was + re-specified as a 1024-float model embedding the next day. The old data + is meaningless under the new definition, so a column of any other + dimension is dropped and re-added, NULL, and `img_vector_src` is cleared + so the backfill sees every row as missing. This is the one place in the + schema code that discards data - by design, and only for this column. + + The dimension is read from pg_attribute.atttypmod (pgvector stores the + declared width there; -1 for a bare `vector`) with fetchall(), so the + recording cursors the tests and the migrate dry run use - which answer [] + - simply take the "nothing to do" path. + + Never raises: every statement is its own try, and _connect() is + autocommit, so a refusal poisons nothing. + """ + try: + cur.execute( + "SELECT a.atttypmod FROM pg_attribute a " + "JOIN pg_class c ON c.oid = a.attrelid " + "JOIN pg_namespace n ON n.oid = c.relnamespace " + "WHERE n.nspname = 'public' AND c.relname = %s " + "AND a.attname = 'img_vector' AND NOT a.attisdropped", + (table_name,), + ) + rows = cur.fetchall() + except Exception as e: # noqa: BLE001 - a failed probe is a skipped heal, not a failed write + logger.warning(f"Could not read img_vector dimension on {table_name}: {e}") + return + if not rows: + return + dims = rows[0][0] + if dims != IMG_VECTOR_DIMS: + try: + cur.execute(f"ALTER TABLE {table_name} DROP COLUMN img_vector") + # IF NOT EXISTS, although it cannot exist here: migrate_brand_schema + # recognises ADD COLUMN statements by that exact phrase. + cur.execute( + f"ALTER TABLE {table_name} ADD COLUMN IF NOT EXISTS img_vector vector({IMG_VECTOR_DIMS})" + ) + cur.execute(f"UPDATE {table_name} SET img_vector_src = NULL") + logger.warning( + f"Retyped {table_name}.img_vector from vector({dims}) to vector({IMG_VECTOR_DIMS}); " + f"stored vectors discarded, backfill needed" + ) + except Exception as e: # noqa: BLE001 - e.g. two first writes raced; the next one re-probes + logger.warning(f"Could not retype {table_name}.img_vector: {e}") + return + # Probe before CREATE INDEX IF NOT EXISTS rather than relying on it: the + # migrate script's dry run records every non-SELECT statement, and an + # unconditional CREATE would make every current table look like work. + # A cursor that cannot answer (fetchall() == []) is treated as "no index". + try: + cur.execute( + "SELECT 1 FROM pg_indexes WHERE schemaname = 'public' AND tablename = %s AND indexname = %s", + (table_name, f"idx_{table_name}_img_vector"), + ) + if cur.fetchall(): + return + except Exception as e: # noqa: BLE001 - fall through and let CREATE IF NOT EXISTS decide + logger.debug(f"Index probe on {table_name} failed, creating anyway: {e}") + try: + cur.execute( + f"CREATE INDEX IF NOT EXISTS idx_{table_name}_img_vector " + f"ON {table_name} USING hnsw (img_vector vector_cosine_ops)" + ) + except Exception as e: # noqa: BLE001 - pgvector < 0.5 has no hnsw; an exact scan still works + logger.warning(f"hnsw index on {table_name}.img_vector skipped: {e}") + def _ensure_columns(cur, table_name: str) -> None: """Add missing columns and relax legacy NOT NULL constraints for smooth schema migration.""" @@ -318,7 +400,7 @@ def _ensure_columns(cur, table_name: str) -> None: # Written only by app/services/image_vector.py - see get_brand_table_ddl. # Plain string literals, always: scripts/migrate_brand_schema.py reads # this dict back out of the source with ast.literal_eval. - "img_vector": "vector(3072)", + "img_vector": "vector(1024)", "img_vector_src": "TEXT", } cur.execute( @@ -330,13 +412,14 @@ def _ensure_columns(cur, table_name: str) -> None: existing = {row[0] for row in col_info} # 1. Add missing columns + have_img_vector = "img_vector" in existing for col, col_type in col_defs.items(): if col not in existing: if col in _TYPE_OPTIONAL_COLUMNS: - # vector(3072) needs pgvector >= 0.4 (the 16000-dim ceiling). - # The production server's version is not pinned anywhere, so - # a refusal here is logged and the write goes on without the - # column; image_vector.table_has_columns() skips the table. + # The vector type needs pgvector; the production server's + # version is not pinned anywhere, so a refusal here is logged + # and the write goes on without the column - + # image_vector.table_has_columns() skips the table. # Safe to continue because _connect() is autocommit. try: cur.execute(f"ALTER TABLE {table_name} ADD COLUMN IF NOT EXISTS {col} {col_type}") @@ -346,11 +429,18 @@ def _ensure_columns(cur, table_name: str) -> None: f"(pgvector too old for this type?): {e}" ) continue + if col == "img_vector": + have_img_vector = True else: cur.execute(f"ALTER TABLE {table_name} ADD COLUMN IF NOT EXISTS {col} {col_type}") if col not in ("created_at", "updated_at"): logger.info(f"Added missing column '{col}' to {table_name}") + # 1b. A pre-existing img_vector of the wrong width is replaced, and the + # column gets its index. Only when the column is actually there. + if have_img_vector: + _ensure_img_vector_type(cur, table_name) + # 2. Relax legacy NOT NULL constraints on columns not present in standard insert # # `nutrition_score` and `health_score` are listed even though the INSERT @@ -651,8 +741,8 @@ def upsert_brand_products(brand: str, products: List[Dict[str, Any]], cleanup: b # in) are simply ignored here. # # img_vector and img_vector_src are omitted for the same reason: - # no writer reaching this statement can carry a pixel vector, so - # naming them would only ever NULL them. app/services/image_vector.py + # no writer reaching this statement can carry an image embedding, + # so naming them would only ever NULL them. app/services/image_vector.py # owns them and is scheduled below, after the write. cur.executemany( f""" @@ -1682,7 +1772,7 @@ def _table_exists(cur, table_name: str) -> bool: # --------------------------------------------------------------------------- # Every product read used to be `SELECT *`, which dragged the vector(384) # embedding over the wire on every listing: measured at 4.7KB of a 7.2KB row -# (see enrichment/catalog_consensus.py). img_vector is 12KB more per row, and +# (see enrichment/catalog_consensus.py). img_vector is 4KB more per row, and # the "all products" browse pulls the whole catalog in one call. Nothing that # reads these dicts uses either column - the API mappers copy fields by name, # and the one exporter that looked at `embedding` (brand_sync) was already diff --git a/requirements.txt b/requirements.txt index f5940dd..2f920f2 100644 --- a/requirements.txt +++ b/requirements.txt @@ -64,6 +64,11 @@ lxml>=4.9.3 python-slugify>=8.0.4 ddgs>=9.14.4 Pillow>=10.0.0 +# TFLite runtime for the img_vector image embedder (app/services/image_embedder.py). +# Runtime only - no TensorFlow - about 20MB; wheels exist for the container's +# 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 aiofiles>=23.2.1 # Playwright (Python) - last-resort image-search fallback only. diff --git a/scripts/backfill_image_vectors.py b/scripts/backfill_image_vectors.py index 156d2e8..714459b 100644 --- a/scripts/backfill_image_vectors.py +++ b/scripts/backfill_image_vectors.py @@ -4,11 +4,12 @@ Fill `img_vector` for the rows that already have a primary image. WHY THIS EXISTS --------------- -New writes get their pixel vector from the worker `upsert_brand_products` +New writes get their image embedding from the worker `upsert_brand_products` schedules (app/services/image_vector.py). Rows written before the column existed never pass through that hook, and rows whose download failed at the time (a CDN that 403'd, a timeout) are left NULL on purpose. This script is -the second chance for both. +the second chance for both - and, with `--force`, the way to recompute every +row after the model or its preprocessing changes. WHAT IT DOES, PER TABLE ----------------------- @@ -19,8 +20,10 @@ WHAT IT DOES, PER TABLE primary. `image_vector.primary_image_url` decides what "primary" means, the same way the product card does. 3. Downloads each image with a small thread pool, no more than one request - in flight per host and a pause between them, decodes it to 32x32 RGB and - writes the 3072 values with a targeted UPDATE. `updated_at` is untouched. + in flight per host and a pause between them, runs it through the + MobileNetV3 embedder (inference is serialised on the embedder's lock; the + downloads are what run in parallel) and writes the 1024 L2-normalised + values with a targeted UPDATE. `updated_at` is untouched. USAGE ----- @@ -28,10 +31,12 @@ USAGE python -m scripts.backfill_image_vectors --brands Amul --apply python -m scripts.backfill_image_vectors --all --apply python -m scripts.backfill_image_vectors --all --recompute-stale --apply + python -m scripts.backfill_image_vectors --all --force --apply # model changed -A dry run downloads nothing; it reports how many rows would be attempted and -shows the first few URLs. No backup is written: nothing else writes this -column, and `--recompute-stale` repairs any row. +A dry run downloads nothing and needs no model; it reports how many rows would +be attempted and shows the first few URLs. `--apply` refuses to start when the +embedder cannot load. No backup is written: nothing else writes this column, +and `--force` rebuilds any row. """ from __future__ import annotations @@ -53,6 +58,7 @@ from app.infrastructure.settings import ( # noqa: E402 DB_NAME, IMAGE_VECTOR_HOST_PAUSE_SECONDS, ) +from app.services import image_embedder # noqa: E402 from app.services import image_vector as iv # noqa: E402 from app.services.vector_store import _connect # noqa: E402 from scripts.repair_brand_images import _brand_tables # noqa: E402 @@ -76,7 +82,7 @@ class _HostPacer: with self._lock: return self._busy.setdefault(host, threading.Lock()) - def fetch(self, url: str, timeout: Optional[float]) -> Tuple[Optional[List[int]], str]: + def fetch(self, url: str, timeout: Optional[float]) -> Tuple[Optional[List[float]], str]: host = urlparse(url).netloc.lower() with self._host_lock(host): with self._lock: @@ -97,6 +103,7 @@ def backfill_table( *, apply: bool, recompute_stale: bool, + force: bool, limit: Optional[int], workers: int, timeout: Optional[float], @@ -110,14 +117,15 @@ def backfill_table( entry["skipped"] = True return entry - rows = iv.rows_needing_vectors(cur, table, recompute_stale=recompute_stale, limit=limit) + rows = iv.rows_needing_vectors(cur, table, recompute_stale=recompute_stale, limit=limit, force=force) entry["candidates"] = len(rows) if not rows: logger.info("%s: nothing to do", table) return entry stale = sum(1 for r in rows if not r.get("missing")) - logger.info("%s: %d row(s) to compute (%d missing, %d stale)", table, len(rows), len(rows) - stale, stale) + logger.info("%s: %d row(s) to compute (%d missing, %d %s)", table, len(rows), len(rows) - stale, stale, + "forced" if force else "stale") if not apply: for r in rows[:5]: logger.info(" %-42s %s", (r.get("product_name") or "")[:42], r["primary_url"][:90]) @@ -125,7 +133,7 @@ def backfill_table( logger.info(" ... and %d more", len(rows) - 5) return entry - results: Dict[str, Tuple[Optional[List[int]], str]] = {} + results: Dict[str, Tuple[Optional[List[float]], str]] = {} with ThreadPoolExecutor(max_workers=workers) as pool: futures = { pool.submit(pacer.fetch, r["primary_url"], timeout): r["image_id"] for r in rows @@ -161,6 +169,8 @@ def main() -> int: parser.add_argument("--limit", type=int, default=None, help="at most N rows per table") parser.add_argument("--recompute-stale", action="store_true", help="also redo rows whose vector came from a URL that is no longer the primary") + parser.add_argument("--force", action="store_true", + help="redo EVERY row with a primary image - after a model or preprocessing change") parser.add_argument("--workers", type=int, default=3, help=f"download threads (1-{MAX_WORKERS}, default 3)") parser.add_argument("--timeout", type=float, default=None, help="per-request timeout in seconds") parser.add_argument("--pause", type=float, default=IMAGE_VECTOR_HOST_PAUSE_SECONDS, @@ -172,6 +182,13 @@ def main() -> int: logger.info("Target database: %s / %s", DB_HOST, DB_NAME) logger.info("Mode: %s", "APPLY - this writes" if args.apply else "DRY RUN - nothing is downloaded or written") + if args.apply: + # Fail before the first download, not after: a missing model would + # otherwise cost every image fetch and write nothing. + logger.info("Embedder: %s", image_embedder.describe()) + if not image_embedder.available(): + logger.error("The embedder cannot run - fix the above before --apply.") + return 2 conn = _connect() if conn is None: @@ -186,7 +203,7 @@ def main() -> int: for _suffix, table in _brand_tables(cur, brands): report.append(backfill_table( cur, conn, table, - apply=args.apply, recompute_stale=args.recompute_stale, + apply=args.apply, recompute_stale=args.recompute_stale, force=args.force, limit=args.limit, workers=workers, timeout=args.timeout, pacer=pacer, )) finally: diff --git a/scripts/migrate_brand_schema.py b/scripts/migrate_brand_schema.py index 1eeb858..e4a10b9 100644 --- a/scripts/migrate_brand_schema.py +++ b/scripts/migrate_brand_schema.py @@ -33,6 +33,13 @@ WHAT IT WILL NOT DO `barcode_last_updated` and REAL for the tax figures rather than the types those values look like they want. Type drift is REPORTED here, never silently "fixed". + + THE ONE EXCEPTION is `img_vector`. A column of any dimension other than + the embedder's (vector_store.IMG_VECTOR_DIMS) is DROPPED and re-added + empty, `img_vector_src` is cleared, and the column gets its hnsw index - + see `_ensure_img_vector_type`. That discards every stored vector on the + table, on purpose: a vector of the wrong width is not data, and the + backfill script rebuilds them. Those statements are shown as `~` lines. * It will not create a brand table that does not exist. * It will not touch `nutrition_facts` or any non-brand table. @@ -108,13 +115,13 @@ _TYPE_ALIASES = { "JSONB": {"jsonb"}, "DOUBLE PRECISION": {"double precision"}, "vector(384)": {"USER-DEFINED"}, - "vector(3072)": {"USER-DEFINED"}, + "vector(1024)": {"USER-DEFINED"}, } -# vector(3072) needs the 16000-dimension ceiling pgvector raised to in 0.4.0. -# Below that, `_ensure_columns` logs and skips the column rather than failing -# the write, and this script says so up front. -_MIN_PGVECTOR_FOR_IMG_VECTOR = (0, 4, 0) +# The hnsw index on img_vector needs pgvector 0.5.0. Below that, +# `_ensure_img_vector_type` logs and skips the index (an exact scan still +# works), and this script says so up front. +_MIN_PGVECTOR_FOR_IMG_VECTOR = (0, 5, 0) def _pgvector_version(cur) -> Optional[str]: @@ -184,8 +191,8 @@ def main() -> int: pgv = _pgvector_version(cur) logger.info("pgvector : %s", pgv or "unknown (extension not found?)") if pgv and _version_tuple(pgv) < _MIN_PGVECTOR_FOR_IMG_VECTOR: - logger.warning("pgvector %s is older than %s: img_vector vector(3072) will be " - "skipped on this server (every other column still applies).", + logger.warning("pgvector %s is older than %s: the hnsw index on img_vector will be " + "skipped on this server (the column itself still applies).", pgv, ".".join(map(str, _MIN_PGVECTOR_FOR_IMG_VECTOR))) logger.info("") @@ -198,6 +205,7 @@ def main() -> int: report: List[Dict[str, Any]] = [] total_missing = 0 total_drift = 0 + total_vec = 0 try: with conn.cursor() as cur: @@ -213,6 +221,21 @@ def main() -> int: _ensure_columns(recorder, table) adds = [s for s in recorder.statements if "ADD COLUMN" in s] + # The img_vector retype/index (see the docstring's one + # exception). A DROP means the ADD that follows is a re-add, + # not a missing column - keep it out of that list. + vec_ops = [ + s for s in recorder.statements + if "img_vector" in s and ( + ("DROP COLUMN" in s) or s.startswith("UPDATE") or s.startswith("CREATE INDEX") + ) + ] + if any("DROP COLUMN img_vector" in s for s in vec_ops): + readd = [s for s in adds if "IF NOT EXISTS img_vector " in s] + adds = [s for s in adds if s not in readd] + vec_ops = [s for s in vec_ops if "DROP COLUMN" in s] + readd + \ + [s for s in vec_ops if "DROP COLUMN" not in s] + # Type drift: a column that exists but whose type is not what # col_defs would have created. Reported, never altered. drift = [] @@ -228,12 +251,15 @@ def main() -> int: "table": table, "missing_columns": [s.split("ADD COLUMN IF NOT EXISTS ")[1] for s in adds], "type_drift": drift, + "img_vector": vec_ops, } report.append(entry) total_missing += len(adds) total_drift += len(drift) + if vec_ops: + total_vec += 1 - if apply and adds: + if apply and (adds or vec_ops): _ensure_columns(cur, table) if apply: @@ -246,12 +272,14 @@ def main() -> int: return 0 width = max(len(e["table"]) for e in report) - changed = [e for e in report if e["missing_columns"] or e["type_drift"]] + changed = [e for e in report if e["missing_columns"] or e["type_drift"] or e["img_vector"]] for entry in sorted(changed, key=lambda e: -len(e["missing_columns"])): logger.info("%-*s %d column(s) missing", width, entry["table"], len(entry["missing_columns"])) for col in entry["missing_columns"]: logger.info("%-*s + %s", width, "", col) + for stmt in entry["img_vector"]: + logger.info("%-*s ~ %s", width, "", stmt) for d in entry["type_drift"]: logger.info("%-*s ! %s is %s, col_defs declares %s (NOT changed)", width, "", d["column"], d["actual"], d["declared"]) @@ -261,7 +289,11 @@ def main() -> int: len(report), len(report) - len(changed)) logger.info("%d column(s) %s, %d type mismatch(es) reported", total_missing, "added" if apply else "would be added", total_drift) - if not apply and total_missing: + if total_vec: + logger.info("%d table(s) %s img_vector retype/index (stored vectors on a retyped " + "table are discarded - run scripts/backfill_image_vectors afterwards)", + total_vec, "had" if apply else "need") + if not apply and (total_missing or total_vec): logger.info("") logger.info("Re-run with --apply to write these changes.") return 0 diff --git a/tests/test_image_embedder_model.py b/tests/test_image_embedder_model.py new file mode 100644 index 0000000..268e2bf --- /dev/null +++ b/tests/test_image_embedder_model.py @@ -0,0 +1,72 @@ +"""The real MobileNetV3 embedder, when its runtime and model file are present. + +Everything in tests/test_image_vector.py runs against a fake interpreter so +the suite never depends on a 20MB wheel or a binary that is not in git. This +file is the one place the actual model is exercised, and it skips - not +fails - when either piece is missing, so a checkout without the .tflite +stays green. + +What it pins is the contract the column relies on, not the model's opinions: +1024 values, unit length, deterministic for identical bytes, and different +for different pictures. +""" +from __future__ import annotations + +import io +import math + +import pytest +from PIL import Image + +pytest.importorskip("ai_edge_litert") + +from app.services import image_embedder as emb # noqa: E402 + +pytestmark = pytest.mark.skipif( + not emb.IMAGE_EMBED_MODEL_PATH.is_file(), + reason=f"model file not present at {emb.IMAGE_EMBED_MODEL_PATH}", +) + + +def _png(colour) -> bytes: + buf = io.BytesIO() + Image.new("RGB", (160, 120), colour).save(buf, "PNG") + return buf.getvalue() + + +@pytest.fixture(autouse=True) +def _fresh(): + emb._reset() + yield + emb._reset() + + +def _cosine(a, b) -> float: + return sum(x * y for x, y in zip(a, b)) + + +def test_the_model_loads_and_reports_its_shapes(): + assert emb.available() is True + assert emb.describe().startswith("ready") + + +def test_an_embedding_is_1024_unit_length_floats(): + vec = emb.embedding_for_bytes(_png((200, 30, 30))) + + assert vec is not None and len(vec) == 1024 + assert abs(math.sqrt(sum(v * v for v in vec)) - 1.0) < 1e-4 + assert all(math.isfinite(v) for v in vec) + + +def test_identical_bytes_give_identical_vectors(): + data = _png((10, 120, 200)) + + assert emb.embedding_for_bytes(data) == emb.embedding_for_bytes(data) + + +def test_different_pictures_give_different_vectors(): + red = emb.embedding_for_bytes(_png((220, 20, 20))) + blue = emb.embedding_for_bytes(_png((20, 20, 220))) + + assert red is not None and blue is not None + assert -1.0 <= _cosine(red, blue) < 0.999 diff --git a/tests/test_image_vector.py b/tests/test_image_vector.py index 546b888..623ed96 100644 --- a/tests/test_image_vector.py +++ b/tests/test_image_vector.py @@ -1,19 +1,27 @@ -"""img_vector: the pixel thumbnail of each product's primary image. +"""img_vector: the MobileNetV3 embedding of each product's primary image. -Four things have to hold, and each has broken independently for a sibling +Five things have to hold, and each has broken independently for a sibling column before, so each gets its own tests here: -1. The vector is what the card shows: `image_url`, else `image_urls[0]`, - EXIF-rotated, flattened onto white, 32x32 RGB - and anything Pillow cannot - decode is None, never an exception (pytest runs warnings as errors). -2. The columns exist on every brand table via `_ensure_columns`, and a server - that refuses `vector(3072)` loses only that column, not the write. -3. The upsert never names them (the nutrition_score rule), and the product +1. The tensor is what the card shows: `image_url`, else `image_urls[0]`, + EXIF-rotated, flattened onto white, RGB, 224x224 - and anything Pillow + cannot decode is None, never an exception (pytest runs warnings as + errors). The tests are written to survive a change of value scaling, since + that part of `preprocess` is meant to be replaced by the colleague's code. +2. The model is one lazily-loaded interpreter behind one lock; without the + runtime or the file it says so once and returns None forever after. +3. The columns exist on every brand table via `_ensure_columns`, a column of + the wrong width is dropped and re-added, the index is created there and + never in the DDL, and a server that refuses the type loses only that + column, not the write. +4. The upsert never names them (the nutrition_score rule), and the product readers never select them - the browse endpoint pulls the whole catalog. -4. The hook is off the write path: it enqueues and returns, it is a no-op +5. The hook is off the write path: it enqueues and returns, it is a no-op when disabled, and a full queue is a log line, not an error. -No database is involved anywhere; cursors are recorders. +No database and no model file are involved anywhere; cursors are recorders +and the interpreter is a fake. tests/test_image_embedder_model.py runs the +real model when it is present. """ from __future__ import annotations @@ -23,9 +31,11 @@ import re import threading from typing import Any, Dict, List +import numpy as np import pytest from PIL import Image +from app.services import image_embedder as emb from app.services import image_vector as iv from app.services import vector_store @@ -46,65 +56,98 @@ def _jpeg(im: Image.Image, **kw) -> bytes: return buf.getvalue() -def _pixel(vec: List[int], x: int, y: int) -> tuple: - i = (y * iv.IMG_VECTOR_SIZE + x) * 3 - return tuple(vec[i:i + 3]) - - @pytest.fixture(autouse=True) -def _fresh_column_cache(): +def _fresh_state(): vector_store._invalidate_product_columns_cache() + emb._reset() yield vector_store._invalidate_product_columns_cache() + emb._reset() + + +class FakeInterpreter: + """Stands in for ai_edge_litert.Interpreter: records the input, returns a + fixed (1, 1024) output.""" + + def __init__(self, output=None): + self.inputs: List[np.ndarray] = [] + self.output = output if output is not None else np.arange(1, 1025, dtype=np.float32).reshape(1, 1024) + + def set_tensor(self, index, value): + self.inputs.append(np.array(value, copy=True)) + + def invoke(self): + pass + + def get_tensor(self, index): + return self.output + + +def _install_fake(monkeypatch, output=None) -> FakeInterpreter: + fake = FakeInterpreter(output) + monkeypatch.setattr(emb, "_interpreter", fake) + monkeypatch.setattr(emb, "_input_index", 0) + monkeypatch.setattr(emb, "_output_index", 0) + return fake # --------------------------------------------------------------------------- -# 1. bytes -> vector +# 1. bytes -> input tensor (scaling-invariant on purpose) # --------------------------------------------------------------------------- -def test_a_solid_png_becomes_3072_values_of_that_colour(): - vec = iv.pixel_vector(_png(Image.new("RGB", (200, 120), (255, 0, 0)))) +def test_a_solid_png_becomes_a_224_tensor_dominated_by_its_colour(): + t = emb.preprocess(_png(Image.new("RGB", (200, 120), (255, 0, 0)))) - assert vec is not None - assert len(vec) == iv.IMG_VECTOR_DIM == 3072 - assert set(vec[0::3]) == {255} - assert set(vec[1::3]) == {0} - assert set(vec[2::3]) == {0} + assert t is not None + assert t.shape == (1, 224, 224, 3) and t.dtype == np.float32 + r, g, b = t[0, :, :, 0], t[0, :, :, 1], t[0, :, :, 2] + assert np.all(r == t.max()) and np.all(g == t.min()) and np.all(b == t.min()) + assert t.max() > t.min() -def test_a_solid_jpeg_is_within_lossy_tolerance(): - vec = iv.pixel_vector(_jpeg(Image.new("RGB", (300, 300), (40, 120, 200)))) +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.""" + t = emb.preprocess(_png(Image.new("RGB", (10, 10), (255, 128, 0)))) - assert vec is not None and len(vec) == 3072 - for channel, expected in enumerate((40, 120, 200)): - assert all(abs(v - expected) <= 3 for v in vec[channel::3]) + assert float(t[0, 0, 0, 0]) == 255.0 and float(t[0, 0, 0, 2]) == 0.0 + + +def test_a_solid_jpeg_is_uniform_within_lossy_tolerance(): + t = emb.preprocess(_jpeg(Image.new("RGB", (300, 300), (40, 120, 200)))) + + assert t is not None + for c in range(3): + chan = t[0, :, :, c] + assert chan.max() - chan.min() <= 3 * (t.max() / 255.0) + 1e-6 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 (it is the left half). - With it, the left half has become the top half and bottom-left is blue. + 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. """ im = Image.new("RGB", (64, 32), (255, 0, 0)) im.paste((0, 0, 255), (32, 0, 64, 32)) exif = Image.Exif() exif[0x0112] = 6 - vec = iv.pixel_vector(_jpeg(im, exif=exif.tobytes())) + t = emb.preprocess(_jpeg(im, exif=exif.tobytes())) - assert vec is not None - r, g, b = _pixel(vec, 0, iv.IMG_VECTOR_SIZE - 1) - assert b > 200 and r < 60, (r, g, b) - r, g, b = _pixel(vec, 0, 0) - assert r > 200 and b < 60, (r, g, b) + assert t is not None + bottom_left = t[0, 223, 0] + top_left = t[0, 0, 0] + assert bottom_left[2] > bottom_left[0], bottom_left # blue dominates + assert top_left[0] > top_left[2], top_left # red dominates def test_transparency_is_flattened_onto_white_not_black(): - im = Image.new("RGBA", (50, 50), (0, 0, 0, 0)) # fully transparent - vec = iv.pixel_vector(_png(im)) + t = emb.preprocess(_png(Image.new("RGBA", (50, 50), (0, 0, 0, 0)))) - assert vec is not None - assert set(vec) == {255} + assert t is not None + assert np.all(t == t.max()) # white: every channel at the top of the scale + assert t.max() > 0 @pytest.mark.parametrize("mode", ["P", "L", "LA", "CMYK", "I;16"]) @@ -114,10 +157,9 @@ def test_every_pillow_mode_a_product_photo_could_arrive_in_decodes(mode): buf = io.BytesIO() im.save(buf, "TIFF" if mode in ("CMYK", "I;16") else "PNG") - vec = iv.pixel_vector(buf.getvalue()) + t = emb.preprocess(buf.getvalue()) - assert vec is not None and len(vec) == 3072 - assert all(0 <= v <= 255 for v in vec) + assert t is not None and t.shape == (1, 224, 224, 3) and np.all(np.isfinite(t)) def test_the_first_frame_of_an_animated_gif_is_used(): @@ -125,12 +167,12 @@ def test_the_first_frame_of_an_animated_gif_is_used(): buf = io.BytesIO() frames[0].save(buf, "GIF", save_all=True, append_images=frames[1:]) - assert iv.pixel_vector(buf.getvalue()) is not None + assert emb.preprocess(buf.getvalue()) is not None @pytest.mark.parametrize("data", [b"", b"not an image", b"\x89PNG\r\n\x1a\n" + b"\x00" * 40]) def test_undecodable_bytes_are_none_not_an_exception(data): - assert iv.pixel_vector(data) is None + assert emb.preprocess(data) is None def test_a_truncated_file_is_none(): @@ -138,18 +180,113 @@ def test_a_truncated_file_is_none(): whole = _png(noisy) assert len(whole) > 400, "need a file big enough to cut" - assert iv.pixel_vector(whole[:200]) is None + assert emb.preprocess(whole[:200]) is None def test_a_header_claiming_too_many_pixels_is_refused_before_decoding(monkeypatch): - monkeypatch.setattr(iv, "IMAGE_VECTOR_MAX_PIXELS", 100) + monkeypatch.setattr(emb, "IMAGE_VECTOR_MAX_PIXELS", 100) - assert iv.pixel_vector(_png(Image.new("RGB", (20, 20)))) is None - assert iv.pixel_vector(_png(Image.new("RGB", (10, 10)))) is not None + assert emb.preprocess(_png(Image.new("RGB", (20, 20)))) is None + assert emb.preprocess(_png(Image.new("RGB", (10, 10)))) is not None + + +# --------------------------------------------------------------------------- +# 1a. the model: one interpreter, one lock, silent when absent +# --------------------------------------------------------------------------- + +def test_l2_normalize_gives_a_unit_vector_and_refuses_zero(): + unit = emb.l2_normalize(np.array([3.0, 4.0], dtype=np.float32)) + assert unit is not None and abs(float(np.linalg.norm(unit)) - 1.0) < 1e-6 + assert emb.l2_normalize(np.zeros(4, dtype=np.float32)) is None + + +def test_embedding_for_bytes_is_1024_unit_floats_from_the_model_output(monkeypatch): + fake = _install_fake(monkeypatch) + + vec = iv.vector_for_url # noqa: F841 - the public path below is what scripts call + out = emb.embedding_for_bytes(_png(Image.new("RGB", (30, 30), (0, 255, 0)))) + + assert out is not None and len(out) == emb.EMBED_DIM == iv.IMG_VECTOR_DIM == 1024 + assert abs(sum(v * v for v in out) - 1.0) < 1e-4 + assert all(isinstance(v, float) for v in out) + assert len(fake.inputs) == 1 and fake.inputs[0].shape == (1, 224, 224, 3) + assert fake.inputs[0].dtype == np.float32 + + +def test_embed_rejects_a_tensor_of_the_wrong_shape(monkeypatch): + fake = _install_fake(monkeypatch) + + assert emb.embed(np.zeros((1, 32, 32, 3), dtype=np.float32)) is None + assert emb.embed(None) is None + assert fake.inputs == [] + + +def test_a_zero_model_output_is_none_not_a_nan_vector(monkeypatch): + _install_fake(monkeypatch, output=np.zeros((1, 1024), dtype=np.float32)) + + assert emb.embedding_for_bytes(_png(Image.new("RGB", (8, 8)))) is None + + +def test_a_missing_model_warns_once_and_then_stays_quiet(monkeypatch, caplog, tmp_path): + monkeypatch.setattr(emb, "IMAGE_EMBED_MODEL_PATH", tmp_path / "nope.tflite") + + with caplog.at_level("WARNING", logger=emb.__name__): + assert emb.available() is False + assert emb.available() is False + assert emb.embedding_for_bytes(_png(Image.new("RGB", (8, 8)))) is None + + warnings_ = [r for r in caplog.records if r.levelname == "WARNING"] + assert len(warnings_) == 1 and "not found" in warnings_[0].getMessage() + assert "disabled" in emb.describe() + + +def test_a_broken_runtime_import_disables_rather_than_raises(monkeypatch, tmp_path): + model = tmp_path / "m.tflite" + model.write_bytes(b"not a flatbuffer") + monkeypatch.setattr(emb, "IMAGE_EMBED_MODEL_PATH", model) + import builtins + real_import = builtins.__import__ + + def no_litert(name, *a, **k): + if name.startswith("ai_edge_litert"): + raise ImportError("no wheel for this platform") + return real_import(name, *a, **k) + + monkeypatch.setattr(builtins, "__import__", no_litert) + + assert emb.available() is False + assert "not importable" in emb.describe() + + +def test_inference_is_serialised_on_the_module_lock(monkeypatch): + """Two threads embedding at once must not interleave set_tensor/invoke.""" + inside = [] + overlap = [] + + class SlowFake(FakeInterpreter): + def invoke(self): + inside.append(1) + if len(inside) > 1: + overlap.append(1) + threading.Event().wait(0.02) + inside.pop() + + fake = SlowFake() + monkeypatch.setattr(emb, "_interpreter", fake) + monkeypatch.setattr(emb, "_input_index", 0) + monkeypatch.setattr(emb, "_output_index", 0) + tensor = np.zeros((1, 224, 224, 3), dtype=np.float32) + threads = [threading.Thread(target=emb.embed, args=(tensor,)) for _ in range(6)] + for t in threads: + t.start() + for t in threads: + t.join() + + assert not overlap and len(fake.inputs) == 6 def test_to_pg_is_the_same_text_form_the_upsert_uses_for_embedding(): - assert iv.to_pg([0, 128, 255]) == "[0,128,255]" + assert iv.to_pg([0, 0.5, 1]) == "[0.0,0.5,1.0]" # --------------------------------------------------------------------------- @@ -297,7 +434,7 @@ class MigrationCursor: text = " ".join(str(sql).split()) self.statements.append(text) if self._refuse and self._refuse in text and "ADD COLUMN" in text: - raise RuntimeError('type "vector(3072)" does not exist') + raise RuntimeError('type "vector(1024)" does not exist') def fetchall(self): return [] @@ -308,23 +445,26 @@ def test_the_migration_adds_both_columns_to_an_existing_table(): vector_store._ensure_columns(cur, "brand_cadbury") - assert "ALTER TABLE brand_cadbury ADD COLUMN IF NOT EXISTS img_vector vector(3072)" in cur.statements + assert "ALTER TABLE brand_cadbury ADD COLUMN IF NOT EXISTS img_vector vector(1024)" in cur.statements assert "ALTER TABLE brand_cadbury ADD COLUMN IF NOT EXISTS img_vector_src TEXT" in cur.statements def test_a_server_that_refuses_the_vector_type_still_gets_every_other_column(): - cur = MigrationCursor(refuse="img_vector vector(3072)") + cur = MigrationCursor(refuse="img_vector vector(1024)") vector_store._ensure_columns(cur, "brand_cadbury") # must not raise assert "ALTER TABLE brand_cadbury ADD COLUMN IF NOT EXISTS img_vector_src TEXT" in cur.statements assert "ALTER TABLE brand_cadbury ADD COLUMN IF NOT EXISTS embedding vector(384)" in cur.statements + # No column -> nothing to retype and nothing to index. + assert not any("pg_attribute" in s or "CREATE INDEX IF NOT EXISTS idx_brand_cadbury_img_vector" in s + for s in cur.statements) def test_the_create_table_ddl_carries_them_and_does_not_index_the_vector(): ddl = vector_store.get_brand_table_ddl("Cadbury") - assert "img_vector vector(3072)" in ddl + assert "img_vector vector(1024)" in ddl assert "img_vector_src TEXT" in ddl for line in ddl.splitlines(): if "CREATE INDEX" in line: @@ -336,8 +476,96 @@ def test_the_migration_script_can_still_read_col_defs_as_literals(): declared = _declared_types() - assert declared["img_vector"] == "vector(3072)" + assert declared["img_vector"] == f"vector({vector_store.IMG_VECTOR_DIMS})" == "vector(1024)" assert declared["img_vector_src"] == "TEXT" + assert vector_store.IMG_VECTOR_DIMS == emb.EMBED_DIM + + +class RetypeCursor(MigrationCursor): + """A table that already HAS img_vector, of a given width.""" + + def __init__(self, dims, indexed=False): + super().__init__() + self._dims = dims + self._indexed = indexed + + def fetchall(self): + last = self.statements[-1] if self.statements else "" + if "pg_attribute" in last: + return [(self._dims,)] + if "pg_indexes" in last: + return [(1,)] if self._indexed else [] + if "information_schema.columns" in last: + return [("img_vector", "YES", None), ("img_vector_src", "YES", None)] + return [] + + +def test_a_column_of_the_old_width_is_dropped_readded_and_indexed(): + cur = RetypeCursor(3072) + + vector_store._ensure_columns(cur, "brand_x") + + wanted = [ + "ALTER TABLE brand_x DROP COLUMN img_vector", + "ALTER TABLE brand_x ADD COLUMN IF NOT EXISTS img_vector vector(1024)", + "UPDATE brand_x SET img_vector_src = NULL", + "CREATE INDEX IF NOT EXISTS idx_brand_x_img_vector ON brand_x USING hnsw (img_vector vector_cosine_ops)", + ] + positions = [cur.statements.index(w) for w in wanted] + assert positions == sorted(positions), cur.statements + + +def test_a_column_of_the_right_width_only_gets_its_index(): + cur = RetypeCursor(1024) + + vector_store._ensure_columns(cur, "brand_x") + + assert not any("DROP COLUMN" in s or s.startswith("UPDATE") for s in cur.statements) + assert "CREATE INDEX IF NOT EXISTS idx_brand_x_img_vector ON brand_x USING hnsw (img_vector vector_cosine_ops)" in cur.statements + + +def test_a_table_that_is_already_current_emits_no_statement_at_all(): + """What the migrate dry run sees after --apply: nothing to report.""" + cur = RetypeCursor(1024, indexed=True) + + vector_store._ensure_columns(cur, "brand_x") + + assert not any("img_vector" in s and not s.startswith("SELECT") for s in cur.statements), cur.statements + + +def test_a_cursor_that_cannot_answer_the_width_probe_changes_nothing(): + """The migrate dry run and the other schema tests use recorders whose + fetchall() is []. That must read as 'nothing to do', never as 'retype'.""" + cur = MigrationCursor() + + vector_store._ensure_columns(cur, "brand_x") + + assert not any("DROP COLUMN" in s or "CREATE INDEX IF NOT EXISTS idx_brand_x_img_vector" in s + for s in cur.statements) + + +def test_a_refused_index_does_not_fail_the_write(): + class NoHnsw(RetypeCursor): + def execute(self, sql, params=None): + super().execute(sql, params) + if "USING hnsw" in self.statements[-1]: + raise RuntimeError("access method hnsw does not exist") + + cur = NoHnsw(1024) + + vector_store._ensure_columns(cur, "brand_x") # must not raise + + assert any("USING hnsw" in s for s in cur.statements) + + +def test_the_retype_never_reaches_the_ddl(): + """The DDL runs BEFORE _ensure_columns on every write; an index there + would hit a still-3072 column and fail the write.""" + import inspect + statements = [line for line in vector_store.get_brand_table_ddl("Cadbury").splitlines() + if not line.strip().startswith("--")] + assert not any("hnsw" in line for line in statements) + assert "DROP COLUMN" not in inspect.getsource(vector_store.get_brand_table_ddl) # --------------------------------------------------------------------------- @@ -616,7 +844,7 @@ def test_store_vector_does_not_touch_updated_at(): text, params = cur.statements[-1] assert text == "UPDATE brand_a SET img_vector = %s::vector, img_vector_src = %s WHERE image_id = %s" - assert params == ("[1,2,3]", "https://a/x.jpg", "x") + assert params == ("[1.0,2.0,3.0]", "https://a/x.jpg", "x") assert "updated_at" not in text @@ -634,8 +862,9 @@ class _RowsConn: def test_backfill_dry_run_downloads_and_writes_nothing(monkeypatch): cur = RowsCursor(_ROWS) monkeypatch.setattr(vector_store, "_connect", lambda: _RowsConn(cur)) + monkeypatch.setattr(emb, "available", lambda: True) fetched: List[str] = [] - monkeypatch.setattr(iv, "vector_for_url", lambda url, timeout=None: (fetched.append(url) or [0] * 3072, url)) + monkeypatch.setattr(iv, "vector_for_url", lambda url, timeout=None: (fetched.append(url) or [0.0] * 1024, url)) totals = iv.backfill_rows("A", dry_run=True, pause_seconds=0) @@ -646,9 +875,10 @@ def test_backfill_dry_run_downloads_and_writes_nothing(monkeypatch): def test_backfill_apply_writes_one_update_per_success_and_counts_failures(monkeypatch): cur = RowsCursor(_ROWS) monkeypatch.setattr(vector_store, "_connect", lambda: _RowsConn(cur)) + monkeypatch.setattr(emb, "available", lambda: True) monkeypatch.setattr( iv, "vector_for_url", - lambda url, timeout=None: (None, url) if "4.jpg" in url else ([7] * 3072, url), + lambda url, timeout=None: (None, url) if "4.jpg" in url else ([0.03125] * 1024, url), ) totals = iv.backfill_rows("A", pause_seconds=0) @@ -657,6 +887,29 @@ def test_backfill_apply_writes_one_update_per_success_and_counts_failures(monkey assert totals == {"candidates": 3, "computed": 2, "failed": 1, "skipped": 0} assert [p[2] for _, p in updates] == ["missing", "stale"] assert updates[1][1][1] == "https://a/2-new.jpg" + assert updates[0][1][0].startswith("[0.03125,0.03125,") + + +def test_backfill_downloads_nothing_when_the_embedder_is_unavailable(monkeypatch): + cur = RowsCursor(_ROWS) + monkeypatch.setattr(vector_store, "_connect", lambda: _RowsConn(cur)) + monkeypatch.setattr(emb, "available", lambda: False) + monkeypatch.setattr(iv, "vector_for_url", lambda url, timeout=None: pytest.fail("must not download")) + + totals = iv.backfill_rows("A", pause_seconds=0) + + assert totals["skipped"] == 1 and totals["candidates"] == 0 + assert not any(t.startswith("SELECT image_id") for t, _ in cur.statements) + + +def test_force_takes_every_row_with_a_primary_even_the_current_ones(): + cur = RowsCursor(_ROWS) + + got = iv.rows_needing_vectors(cur, "brand_a", force=True) + + assert [r["image_id"] for r in got] == ["missing", "stale", "fresh", "gallery-only"] + text, _ = cur.statements[-1] + assert "WHERE" not in text def test_backfill_skips_a_table_the_migration_has_not_reached(monkeypatch):