Image vector embedding
This commit is contained in:
@@ -314,6 +314,13 @@ 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
|
||||
# `python -m scripts.backfill_image_vectors --all --apply`. Set false to stop
|
||||
# the API process from downloading images at all.
|
||||
ENABLE_IMAGE_VECTORS=true
|
||||
|
||||
# Product SKU: try a live web search for a real marketplace product ID
|
||||
# (Amazon ASIN, Flipkart PID, etc.) before falling back to an internal SKU.
|
||||
# Set to false to always generate internal SKUs only (faster, offline-safe).
|
||||
|
||||
@@ -477,13 +477,15 @@ def _brand_defaults(brand_parent: str) -> Dict[str, Any]:
|
||||
"""
|
||||
# Two reads on purpose, and the split matters on a memory-capped host.
|
||||
#
|
||||
# `_brand_sample` is SELECT * limit 1 - one row, embedding and all, because
|
||||
# the fields it seeds (category, price band, size) need the whole row.
|
||||
# `_brand_sample` is one full row (limit 1), because the fields it seeds
|
||||
# (category, price band, size) need the whole row. The product readers now
|
||||
# project the vector columns out, so "full" no longer means "with the
|
||||
# embedding".
|
||||
#
|
||||
# The consensus read is 300 rows, so it takes only the two columns it
|
||||
# actually inspects. Measured: SELECT * over 244 rows costs 3.0 MB, of
|
||||
# which 4.7 KB per row is an embedding string nothing here reads. Two named
|
||||
# columns is roughly 50 KB for the same rows.
|
||||
# actually inspects. Measured when reads were still SELECT *: 244 rows
|
||||
# cost 3.0 MB, of which 4.7 KB per row was an embedding string nothing
|
||||
# here reads. Two named columns is roughly 50 KB for the same rows.
|
||||
rows = consensus_rows(brand_parent, ["fssai_license", "providers"],
|
||||
limit=_CONSENSUS_SAMPLE_LIMIT)
|
||||
|
||||
|
||||
@@ -380,6 +380,29 @@ USE_PLAYWRIGHT_FALLBACK = _bool("USE_PLAYWRIGHT_FALLBACK", "true")
|
||||
# out 1x1 tracking pixels / broken placeholder images)
|
||||
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)
|
||||
# ---------------------------------------------------------------------------
|
||||
# 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.
|
||||
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"))
|
||||
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"))
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# USDA FoodData Central - nutrition for loose, unbranded commodities
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@@ -154,7 +154,8 @@ def _slim(product: Dict[str, Any]) -> Dict[str, Any]:
|
||||
return {
|
||||
k: v
|
||||
for k, v in product.items()
|
||||
if k not in {"embedding", "embedding_text", "distance"} and v not in (None, "", [], {})
|
||||
if k not in {"embedding", "embedding_text", "distance", "img_vector", "img_vector_src"}
|
||||
and v not in (None, "", [], {})
|
||||
}
|
||||
|
||||
|
||||
|
||||
417
app/services/image_vector.py
Normal file
417
app/services/image_vector.py
Normal file
@@ -0,0 +1,417 @@
|
||||
"""A pixel thumbnail 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.
|
||||
|
||||
`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.
|
||||
|
||||
WHY THIS MODULE, AND NOT THE UPSERT
|
||||
-----------------------------------
|
||||
`upsert_brand_products` never names these columns, for the reason documented
|
||||
next to its INSERT: every product dict that reaches it comes from a
|
||||
spreadsheet, the seed JSON or a SELECT round-trip, none of which can carry a
|
||||
vector, so naming the column there would only ever NULL it. The writes here
|
||||
are targeted UPDATEs by image_id, the same shape as nutrition_score_sync.
|
||||
|
||||
WHY A SINGLE BOUNDED WORKER
|
||||
---------------------------
|
||||
The upsert calls `schedule()`, which enqueues (brand, image_ids) and returns.
|
||||
One daemon thread drains the queue, so at most one image is being downloaded
|
||||
and decoded in the API process at any moment, a write never waits on a CDN,
|
||||
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.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import io
|
||||
import logging
|
||||
import queue
|
||||
import threading
|
||||
import time
|
||||
from typing import Any, Dict, Iterable, List, Mapping, Optional, Tuple
|
||||
from urllib.parse import urlparse
|
||||
|
||||
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,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
IMG_VECTOR_SIZE = 32
|
||||
IMG_VECTOR_DIM = IMG_VECTOR_SIZE * IMG_VECTOR_SIZE * 3 # 3072
|
||||
IMG_VECTOR_COLUMNS = ("img_vector", "img_vector_src")
|
||||
|
||||
_IMAGE_EXTENSIONS = (".jpg", ".jpeg", ".png", ".webp", ".gif", ".bmp")
|
||||
_CHUNK = 64 * 1024
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Which image
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def primary_image_url(row: Mapping[str, Any]) -> Optional[str]:
|
||||
"""The URL the product card displays, or None if it shows nothing.
|
||||
|
||||
`image_url` first, else `image_urls[0]` - ProductCard.jsx builds its list in
|
||||
that order. Normalised through `_usable_image_url` so blanks, bare paths and
|
||||
hosts in `_DEAD_IMAGE_HOSTS` never cost a request. The string returned here
|
||||
is what goes into `img_vector_src`, so staleness is always judged against
|
||||
the same definition.
|
||||
"""
|
||||
from app.services.vector_store import _usable_image_url
|
||||
|
||||
candidate = _usable_image_url(row.get("image_url"))
|
||||
if candidate:
|
||||
return candidate
|
||||
urls = row.get("image_urls") or []
|
||||
if isinstance(urls, (list, tuple)) and urls:
|
||||
return _usable_image_url(urls[0])
|
||||
return None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Bytes -> vector
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
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:
|
||||
"""The text form pgvector accepts, identical to what the upsert uses for `embedding`."""
|
||||
return "[" + ",".join(str(int(v)) for v in vector) + "]"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# URL -> bytes
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def download_image_bytes(url: str, timeout: Optional[float] = None) -> Optional[bytes]:
|
||||
"""Fetch `url` as image bytes, or None.
|
||||
|
||||
Same two-step Referer strategy as `image_search.validate_image_url_live`:
|
||||
first with a same-site Referer (hotlink-protected CDNs), then with none
|
||||
(hosts that reject a Referer pretending to be same-site). Streams in
|
||||
chunks and gives up past IMAGE_VECTOR_MAX_BYTES, so a wrong URL to a video
|
||||
cannot fill memory; anything under MIN_IMAGE_BYTES is a placeholder.
|
||||
"""
|
||||
if not url or not url.startswith(("http://", "https://")):
|
||||
return None
|
||||
import requests
|
||||
from app.services.image_search import _BROWSER_UA
|
||||
|
||||
parsed = urlparse(url)
|
||||
same_site = f"{parsed.scheme}://{parsed.netloc}/" if parsed.netloc else None
|
||||
timeout = timeout or IMAGE_VECTOR_TIMEOUT_SECONDS
|
||||
|
||||
for referer in (same_site, None):
|
||||
headers = {"User-Agent": _BROWSER_UA, "Accept": "image/*,*/*;q=0.8"}
|
||||
if referer:
|
||||
headers["Referer"] = referer
|
||||
resp = None
|
||||
try:
|
||||
resp = requests.get(url, headers=headers, timeout=timeout, stream=True)
|
||||
if resp.status_code != 200:
|
||||
logger.debug("HTTP %s for %s", resp.status_code, url)
|
||||
continue
|
||||
content_type = (resp.headers.get("content-type") or "").lower()
|
||||
if "image" not in content_type:
|
||||
# A server that says text/html means it - typically a CDN that
|
||||
# now redirects every dead path to a homepage (uat.amul.com does
|
||||
# exactly this), and the .png in the URL is no evidence at all.
|
||||
# Only an unlabelled body (octet-stream, blank) gets the benefit
|
||||
# of the extension.
|
||||
if content_type.startswith("text/") or not parsed.path.lower().endswith(_IMAGE_EXTENSIONS):
|
||||
logger.debug("not an image (%s): %s", content_type, url)
|
||||
return None
|
||||
buf = bytearray()
|
||||
for chunk in resp.iter_content(chunk_size=_CHUNK):
|
||||
if not chunk:
|
||||
continue
|
||||
buf.extend(chunk)
|
||||
if len(buf) > IMAGE_VECTOR_MAX_BYTES:
|
||||
logger.debug("image over %d bytes, abandoned: %s", IMAGE_VECTOR_MAX_BYTES, url)
|
||||
return None
|
||||
if len(buf) < MIN_IMAGE_BYTES:
|
||||
logger.debug("image too small (%d bytes): %s", len(buf), url)
|
||||
return None
|
||||
return bytes(buf)
|
||||
except Exception as exc: # noqa: BLE001 - a bad host must not stop the batch
|
||||
logger.debug("download failed for %s: %s", url, exc)
|
||||
continue
|
||||
finally:
|
||||
if resp is not None:
|
||||
try:
|
||||
resp.close()
|
||||
except Exception: # noqa: BLE001
|
||||
pass
|
||||
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)."""
|
||||
data = download_image_bytes(url, timeout=timeout)
|
||||
if data is None:
|
||||
return None, url
|
||||
return pixel_vector(data), url
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Table access
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def table_has_columns(cur, table_name: str) -> bool:
|
||||
"""True when the table carries both img_vector columns.
|
||||
|
||||
Every writer checks this first, so a table the migration has not reached -
|
||||
or one whose pgvector refused `vector(3072)` - is skipped, never errored.
|
||||
"""
|
||||
cur.execute(
|
||||
"SELECT column_name FROM information_schema.columns "
|
||||
"WHERE table_schema = 'public' AND table_name = %s AND column_name = ANY(%s)",
|
||||
(table_name, list(IMG_VECTOR_COLUMNS)),
|
||||
)
|
||||
present = {row[0] for row in cur.fetchall()}
|
||||
return all(col in present for col in IMG_VECTOR_COLUMNS)
|
||||
|
||||
|
||||
def rows_needing_vectors(
|
||||
cur,
|
||||
table_name: str,
|
||||
image_ids: Optional[List[str]] = None,
|
||||
*,
|
||||
recompute_stale: bool = True,
|
||||
limit: Optional[int] = None,
|
||||
) -> List[Dict[str, Any]]:
|
||||
"""Rows that have a primary image and no current vector for it.
|
||||
|
||||
"Current" means `img_vector_src` equals the primary URL; a row whose
|
||||
`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.
|
||||
"""
|
||||
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]))"
|
||||
)
|
||||
params: List[Any] = []
|
||||
if image_ids is not None:
|
||||
sql += " AND image_id = ANY(%s)"
|
||||
params.append(list(image_ids))
|
||||
sql += " ORDER BY updated_at DESC"
|
||||
if limit is not None:
|
||||
sql += " LIMIT %s"
|
||||
params.append(int(limit))
|
||||
cur.execute(sql, params)
|
||||
colnames = [d[0] for d in cur.description]
|
||||
|
||||
out: List[Dict[str, Any]] = []
|
||||
for raw in cur.fetchall():
|
||||
row = dict(zip(colnames, raw))
|
||||
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
|
||||
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:
|
||||
"""Write one vector. Deliberately leaves `updated_at` alone: the catalog
|
||||
listing orders on it, and a background enrichment must not reshuffle it."""
|
||||
cur.execute(
|
||||
f"UPDATE {table_name} SET img_vector = %s::vector, img_vector_src = %s WHERE image_id = %s",
|
||||
(to_pg(vector), src, image_id),
|
||||
)
|
||||
|
||||
|
||||
def backfill_rows(
|
||||
brand: str,
|
||||
image_ids: Optional[List[str]] = None,
|
||||
*,
|
||||
recompute_stale: bool = True,
|
||||
limit: Optional[int] = None,
|
||||
dry_run: bool = False,
|
||||
pause_seconds: Optional[float] = None,
|
||||
) -> 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.
|
||||
"""
|
||||
from app.services.vector_store import _connect, _table_name
|
||||
|
||||
totals = {"candidates": 0, "computed": 0, "failed": 0, "skipped": 0}
|
||||
conn = _connect()
|
||||
if conn is None:
|
||||
logger.warning("image vectors: no database connection for %s", brand)
|
||||
return totals
|
||||
pause = IMAGE_VECTOR_HOST_PAUSE_SECONDS if pause_seconds is None else pause_seconds
|
||||
table_name = _table_name(brand)
|
||||
last_by_host: Dict[str, float] = {}
|
||||
try:
|
||||
with conn.cursor() as cur:
|
||||
if not table_has_columns(cur, table_name):
|
||||
logger.info("image vectors: %s has no img_vector columns yet - skipped", table_name)
|
||||
totals["skipped"] = 1
|
||||
return totals
|
||||
rows = rows_needing_vectors(
|
||||
cur, table_name, image_ids, recompute_stale=recompute_stale, limit=limit
|
||||
)
|
||||
totals["candidates"] = len(rows)
|
||||
for row in rows:
|
||||
url = row["primary_url"]
|
||||
host = urlparse(url).netloc.lower()
|
||||
wait = last_by_host.get(host, 0.0) + pause - time.monotonic()
|
||||
if wait > 0:
|
||||
time.sleep(wait)
|
||||
last_by_host[host] = time.monotonic()
|
||||
|
||||
vector, src = vector_for_url(url)
|
||||
if vector is None:
|
||||
totals["failed"] += 1
|
||||
continue
|
||||
if not dry_run:
|
||||
store_vector(cur, table_name, row["image_id"], vector, src)
|
||||
totals["computed"] += 1
|
||||
except Exception: # noqa: BLE001 - enrichment must never surface as a failure
|
||||
logger.exception("image vectors: backfill failed for %s", brand)
|
||||
finally:
|
||||
conn.close()
|
||||
if totals["candidates"]:
|
||||
logger.info(
|
||||
"image vectors: %s - %d candidate(s), %d computed, %d failed%s",
|
||||
table_name, totals["candidates"], totals["computed"], totals["failed"],
|
||||
" (dry run)" if dry_run else "",
|
||||
)
|
||||
return totals
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# The hook the upsert calls
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
_queue: "queue.Queue[Tuple[str, List[str]]]" = queue.Queue(maxsize=max(1, IMAGE_VECTOR_QUEUE_MAX))
|
||||
_worker: Optional[threading.Thread] = None
|
||||
_lock = threading.Lock()
|
||||
|
||||
|
||||
def queue_depth() -> int:
|
||||
return _queue.qsize()
|
||||
|
||||
|
||||
def is_running() -> bool:
|
||||
return _worker is not None and _worker.is_alive()
|
||||
|
||||
|
||||
def schedule(brand: str, image_ids: Iterable[str]) -> bool:
|
||||
"""Queue the rows just written. Returns True if queued.
|
||||
|
||||
Off when ENABLE_IMAGE_VECTORS is false. `put_nowait`, never `put`: this
|
||||
runs inside the write path and must not wait. A full queue is logged and
|
||||
dropped - the backfill script picks those rows up, because they are the
|
||||
ones with no vector.
|
||||
"""
|
||||
if not ENABLE_IMAGE_VECTORS:
|
||||
return False
|
||||
ids = [i for i in image_ids if i]
|
||||
if not ids:
|
||||
return False
|
||||
try:
|
||||
_queue.put_nowait((brand, ids))
|
||||
except queue.Full:
|
||||
logger.warning(
|
||||
"image vectors: queue full (%d), %d row(s) for %s left for the backfill script",
|
||||
_queue.maxsize, len(ids), brand,
|
||||
)
|
||||
return False
|
||||
_ensure_worker()
|
||||
return True
|
||||
|
||||
|
||||
def _ensure_worker() -> None:
|
||||
global _worker
|
||||
with _lock:
|
||||
if _worker is not None and _worker.is_alive():
|
||||
return
|
||||
_worker = threading.Thread(target=_loop, name="image-vector-worker", daemon=True)
|
||||
_worker.start()
|
||||
|
||||
|
||||
def _loop() -> None:
|
||||
"""Drain forever; one bad batch must not strand the ones behind it."""
|
||||
while True:
|
||||
brand, ids = _queue.get()
|
||||
try:
|
||||
backfill_rows(brand, ids, recompute_stale=True)
|
||||
except Exception: # noqa: BLE001 - see docstring
|
||||
logger.exception("image vectors: worker failed on %s", brand)
|
||||
finally:
|
||||
_queue.task_done()
|
||||
@@ -5,6 +5,7 @@ from typing import List, Optional, Dict, Any, Tuple
|
||||
import logging
|
||||
import os
|
||||
import re
|
||||
import threading
|
||||
import time
|
||||
from datetime import datetime
|
||||
|
||||
@@ -206,7 +207,16 @@ 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),
|
||||
img_vector_src TEXT,
|
||||
|
||||
-- Timestamps
|
||||
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
|
||||
updated_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP
|
||||
@@ -243,6 +253,11 @@ def _connect() -> Optional[psycopg.Connection]:
|
||||
return None
|
||||
|
||||
|
||||
# Columns whose ADD COLUMN may be refused by an older server and must not
|
||||
# fail the write - see the loop in _ensure_columns.
|
||||
_TYPE_OPTIONAL_COLUMNS = frozenset({"img_vector"})
|
||||
|
||||
|
||||
def _ensure_columns(cur, table_name: str) -> None:
|
||||
"""Add missing columns and relax legacy NOT NULL constraints for smooth schema migration."""
|
||||
col_defs = {
|
||||
@@ -300,6 +315,11 @@ def _ensure_columns(cur, table_name: str) -> None:
|
||||
"nutrition_score": "NUMERIC",
|
||||
"health_score": "NUMERIC",
|
||||
"embedding": "vector(384)",
|
||||
# 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_src": "TEXT",
|
||||
}
|
||||
cur.execute(
|
||||
"SELECT column_name, is_nullable, column_default FROM information_schema.columns "
|
||||
@@ -312,7 +332,22 @@ def _ensure_columns(cur, table_name: str) -> None:
|
||||
# 1. Add missing columns
|
||||
for col, col_type in col_defs.items():
|
||||
if col not in existing:
|
||||
cur.execute(f"ALTER TABLE {table_name} ADD COLUMN IF NOT EXISTS {col} {col_type}")
|
||||
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.
|
||||
# Safe to continue because _connect() is autocommit.
|
||||
try:
|
||||
cur.execute(f"ALTER TABLE {table_name} ADD COLUMN IF NOT EXISTS {col} {col_type}")
|
||||
except Exception as e: # noqa: BLE001 - optional column, see above
|
||||
logger.warning(
|
||||
f"Could not add {col} {col_type} to {table_name} "
|
||||
f"(pgvector too old for this type?): {e}"
|
||||
)
|
||||
continue
|
||||
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}")
|
||||
|
||||
@@ -327,7 +362,7 @@ def _ensure_columns(cur, table_name: str) -> None:
|
||||
# `nutrients_per_100g` joins them: it is mirrored from nutrition_facts by
|
||||
# the same sync, not written by the INSERT, and is likewise a current
|
||||
# column rather than legacy debris.
|
||||
inserted_cols = {"id", "product_name", "title", "description", "category", "image_id", "image_url", "image_urls", "price_range", "size_variants", "providers", "fssai_license", "product_sku", "sku_source", "hsn_code", "final_selling_price", "selling_price", "barcode", "barcode_type", "gtin", "ean13", "upc", "barcode_source", "barcode_verified", "barcode_lookup_status", "barcode_last_updated", "gst_percent", "tax_amount", "hsn_gst_needs_review", "highlights", "nutrients", "nutrients_per_100g", "field_sources", "search_query", "nutrition_score", "health_score", "embedding", "created_at", "updated_at"}
|
||||
inserted_cols = {"id", "product_name", "title", "description", "category", "image_id", "image_url", "image_urls", "price_range", "size_variants", "providers", "fssai_license", "product_sku", "sku_source", "hsn_code", "final_selling_price", "selling_price", "barcode", "barcode_type", "gtin", "ean13", "upc", "barcode_source", "barcode_verified", "barcode_lookup_status", "barcode_last_updated", "gst_percent", "tax_amount", "hsn_gst_needs_review", "highlights", "nutrients", "nutrients_per_100g", "field_sources", "search_query", "nutrition_score", "health_score", "embedding", "img_vector", "img_vector_src", "created_at", "updated_at"}
|
||||
for col, is_nullable, col_def in col_info:
|
||||
if col not in inserted_cols and is_nullable == 'NO' and col_def is None:
|
||||
cur.execute(f"ALTER TABLE {table_name} ALTER COLUMN {col} DROP NOT NULL")
|
||||
@@ -343,6 +378,10 @@ def _ensure_columns(cur, table_name: str) -> None:
|
||||
except Exception as e:
|
||||
logger.warning(f"Unique index creation on {table_name}.image_id: {e}")
|
||||
|
||||
# The projected reads cache each table's column list; a column added here
|
||||
# must be visible to the next SELECT.
|
||||
_invalidate_product_columns_cache()
|
||||
|
||||
|
||||
def ensure_brand_schema(brand: str) -> str:
|
||||
"""Ensure brand-specific table exists and return table name"""
|
||||
@@ -608,8 +647,13 @@ def upsert_brand_products(brand: str, products: List[Dict[str, Any]], cleanup: b
|
||||
#
|
||||
# They are written only by app/services/nutrition_score_sync.py,
|
||||
# which mirrors nutrition_insights. Extra score keys present in an
|
||||
# incoming product dict (get_products_by_brand is SELECT *, so
|
||||
# stage_11's merge carries them back in) are simply ignored here.
|
||||
# incoming product dict (stage_11's merge carries prior rows back
|
||||
# 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
|
||||
# owns them and is scheduled below, after the write.
|
||||
cur.executemany(
|
||||
f"""
|
||||
INSERT INTO {table_name}
|
||||
@@ -765,6 +809,15 @@ def upsert_brand_products(brand: str, products: List[Dict[str, Any]], cleanup: b
|
||||
except Exception: # noqa: BLE001 - cache invalidation must never fail a write
|
||||
pass
|
||||
|
||||
# Queue the rows just written for their img_vector. Off the request path
|
||||
# (one bounded worker thread), fill-only-missing-or-stale, and a failure
|
||||
# here is a missing thumbnail, never a failed write.
|
||||
try:
|
||||
from app.services.image_vector import schedule as _schedule_image_vectors
|
||||
_schedule_image_vectors(brand, expected_ids)
|
||||
except Exception: # noqa: BLE001 - enrichment must never fail a write
|
||||
logger.debug("image-vector scheduling skipped for %s", brand, exc_info=True)
|
||||
|
||||
return persisted
|
||||
|
||||
|
||||
@@ -1127,7 +1180,7 @@ def get_products_by_brand(brand: str, limit: Optional[int] = None, offset: int =
|
||||
logger.warning(f"Table {table_name} does not exist")
|
||||
return []
|
||||
|
||||
sql = f"SELECT * FROM {table_name}"
|
||||
sql = f"SELECT {_product_columns(cur, table_name)} FROM {table_name}"
|
||||
params: List[Any] = []
|
||||
if category:
|
||||
sql += " WHERE category ILIKE %s"
|
||||
@@ -1184,7 +1237,10 @@ def get_product_by_image_id(brand: str, image_id: str) -> Optional[Dict[str, Any
|
||||
with conn, conn.cursor() as cur:
|
||||
if not _table_exists(cur, table_name):
|
||||
return None
|
||||
cur.execute(f"SELECT * FROM {table_name} WHERE image_id = %s LIMIT 1", (image_id,))
|
||||
cur.execute(
|
||||
f"SELECT {_product_columns(cur, table_name)} FROM {table_name} WHERE image_id = %s LIMIT 1",
|
||||
(image_id,),
|
||||
)
|
||||
row = cur.fetchone()
|
||||
if not row:
|
||||
return None
|
||||
@@ -1319,8 +1375,8 @@ def semantic_search(
|
||||
if not _table_exists(cur, table_name):
|
||||
continue
|
||||
sql = (
|
||||
f"SELECT *, embedding <=> %s::vector AS distance FROM {table_name} "
|
||||
f"WHERE embedding IS NOT NULL"
|
||||
f"SELECT {_product_columns(cur, table_name)}, embedding <=> %s::vector AS distance "
|
||||
f"FROM {table_name} WHERE embedding IS NOT NULL"
|
||||
)
|
||||
params: List[Any] = [embedding_str]
|
||||
if category:
|
||||
@@ -1389,7 +1445,7 @@ def text_search(
|
||||
for brand_label, table_name in tables:
|
||||
if not _table_exists(cur, table_name):
|
||||
continue
|
||||
sql = f"SELECT * FROM {table_name}"
|
||||
sql = f"SELECT {_product_columns(cur, table_name)} FROM {table_name}"
|
||||
where_clauses: List[str] = []
|
||||
params: List[Any] = []
|
||||
|
||||
@@ -1567,7 +1623,7 @@ def lexical_search(
|
||||
params.append(f"%{category}%")
|
||||
|
||||
sql = (
|
||||
f"SELECT *{distance_select}, {tier_sql} FROM {table_name} "
|
||||
f"SELECT {_product_columns(cur, table_name)}{distance_select}, {tier_sql} FROM {table_name} "
|
||||
f"WHERE {' AND '.join(where_clauses)} "
|
||||
f"ORDER BY lex_tier ASC, length(coalesce(product_name, title, '')) ASC, "
|
||||
f"updated_at DESC LIMIT %s"
|
||||
@@ -1621,6 +1677,71 @@ def _table_exists(cur, table_name: str) -> bool:
|
||||
return bool(row and row[0])
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Projected reads
|
||||
# ---------------------------------------------------------------------------
|
||||
# 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
|
||||
# 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
|
||||
# getting a pgvector string it could not parse. So the read helpers select
|
||||
# every column EXCEPT these, by name.
|
||||
#
|
||||
# Column lists are probed from information_schema (brand tables written under
|
||||
# older DDLs have divergent columns) and cached for a short TTL; a table the
|
||||
# cache does not know falls back to `*`, so a table created between refreshes
|
||||
# cannot 500.
|
||||
PRODUCT_READ_EXCLUDE = frozenset({"embedding", "img_vector"})
|
||||
PRODUCT_COLUMNS_TTL_SECONDS = int(os.getenv("PRODUCT_COLUMNS_TTL_SECONDS", "60"))
|
||||
_COLUMNS_CACHE: Dict[str, Any] = {"at": 0.0, "data": {}}
|
||||
_COLUMNS_LOCK = threading.Lock()
|
||||
|
||||
|
||||
def _invalidate_product_columns_cache() -> None:
|
||||
with _COLUMNS_LOCK:
|
||||
_COLUMNS_CACHE["at"] = 0.0
|
||||
_COLUMNS_CACHE["data"] = {}
|
||||
|
||||
|
||||
def _refresh_product_columns(cur) -> None:
|
||||
cur.execute(
|
||||
"""
|
||||
SELECT table_name, column_name FROM information_schema.columns
|
||||
WHERE table_schema = 'public' AND table_name LIKE 'brand\\_%'
|
||||
ORDER BY table_name, ordinal_position
|
||||
"""
|
||||
)
|
||||
data: Dict[str, List[str]] = {}
|
||||
for table, col in cur.fetchall():
|
||||
data.setdefault(table, []).append(col)
|
||||
with _COLUMNS_LOCK:
|
||||
_COLUMNS_CACHE["at"] = time.monotonic()
|
||||
_COLUMNS_CACHE["data"] = data
|
||||
|
||||
|
||||
def _product_columns(cur, table_name: str, exclude: frozenset = PRODUCT_READ_EXCLUDE) -> str:
|
||||
"""The SELECT list for `table_name` without the vector columns, or `*`."""
|
||||
with _COLUMNS_LOCK:
|
||||
stale = time.monotonic() - _COLUMNS_CACHE["at"] > PRODUCT_COLUMNS_TTL_SECONDS
|
||||
cols = None if stale else _COLUMNS_CACHE["data"].get(table_name)
|
||||
if cols is None:
|
||||
try:
|
||||
_refresh_product_columns(cur)
|
||||
except Exception as e: # noqa: BLE001 - a failed probe means SELECT *, not a failed read
|
||||
logger.warning("Could not probe columns for %s: %s", table_name, e)
|
||||
return "*"
|
||||
with _COLUMNS_LOCK:
|
||||
cols = _COLUMNS_CACHE["data"].get(table_name)
|
||||
if not cols:
|
||||
return "*"
|
||||
kept = [c for c in cols if c not in exclude]
|
||||
if not kept:
|
||||
return "*"
|
||||
return ", ".join(f'"{c}"' for c in kept)
|
||||
|
||||
|
||||
def _list_brand_table_suffixes(cur, *, include_inactive: bool = False) -> List[str]:
|
||||
"""Return brand-table suffixes (e.g. 'parle' from 'brand_parle') for every brand table.
|
||||
|
||||
|
||||
218
scripts/backfill_image_vectors.py
Normal file
218
scripts/backfill_image_vectors.py
Normal file
@@ -0,0 +1,218 @@
|
||||
#!/usr/bin/env python3
|
||||
"""
|
||||
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`
|
||||
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.
|
||||
|
||||
WHAT IT DOES, PER TABLE
|
||||
-----------------------
|
||||
1. Skips the table if the migration has not reached it yet
|
||||
(`python -m scripts.migrate_brand_schema --apply` first).
|
||||
2. Selects rows with a usable primary image and no vector - or, with
|
||||
`--recompute-stale`, a vector computed from a URL that is no longer the
|
||||
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.
|
||||
|
||||
USAGE
|
||||
-----
|
||||
python -m scripts.backfill_image_vectors --brands Amul # dry run
|
||||
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
|
||||
|
||||
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.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import logging
|
||||
import sys
|
||||
import threading
|
||||
import time
|
||||
from concurrent.futures import ThreadPoolExecutor, as_completed
|
||||
from pathlib import Path
|
||||
from typing import Any, Dict, List, Optional, Tuple
|
||||
from urllib.parse import urlparse
|
||||
|
||||
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
|
||||
|
||||
from app.infrastructure.settings import ( # noqa: E402
|
||||
DB_HOST,
|
||||
DB_NAME,
|
||||
IMAGE_VECTOR_HOST_PAUSE_SECONDS,
|
||||
)
|
||||
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
|
||||
|
||||
logging.basicConfig(level=logging.INFO, format="%(message)s")
|
||||
logger = logging.getLogger("backfill_image_vectors")
|
||||
|
||||
MAX_WORKERS = 6
|
||||
|
||||
|
||||
class _HostPacer:
|
||||
"""At most one request per host every `pause` seconds, across threads."""
|
||||
|
||||
def __init__(self, pause: float):
|
||||
self._pause = pause
|
||||
self._lock = threading.Lock()
|
||||
self._next_ok: Dict[str, float] = {}
|
||||
self._busy: Dict[str, threading.Lock] = {}
|
||||
|
||||
def _host_lock(self, host: str) -> threading.Lock:
|
||||
with self._lock:
|
||||
return self._busy.setdefault(host, threading.Lock())
|
||||
|
||||
def fetch(self, url: str, timeout: Optional[float]) -> Tuple[Optional[List[int]], str]:
|
||||
host = urlparse(url).netloc.lower()
|
||||
with self._host_lock(host):
|
||||
with self._lock:
|
||||
wait = self._next_ok.get(host, 0.0) - time.monotonic()
|
||||
if wait > 0:
|
||||
time.sleep(wait)
|
||||
try:
|
||||
return iv.vector_for_url(url, timeout=timeout)
|
||||
finally:
|
||||
with self._lock:
|
||||
self._next_ok[host] = time.monotonic() + self._pause
|
||||
|
||||
|
||||
def backfill_table(
|
||||
cur,
|
||||
conn,
|
||||
table: str,
|
||||
*,
|
||||
apply: bool,
|
||||
recompute_stale: bool,
|
||||
limit: Optional[int],
|
||||
workers: int,
|
||||
timeout: Optional[float],
|
||||
pacer: _HostPacer,
|
||||
) -> Dict[str, Any]:
|
||||
entry: Dict[str, Any] = {
|
||||
"table": table, "candidates": 0, "computed": 0, "failed": 0, "skipped": False,
|
||||
}
|
||||
if not iv.table_has_columns(cur, table):
|
||||
logger.info("%s: no img_vector columns yet - run migrate_brand_schema --apply first", table)
|
||||
entry["skipped"] = True
|
||||
return entry
|
||||
|
||||
rows = iv.rows_needing_vectors(cur, table, recompute_stale=recompute_stale, limit=limit)
|
||||
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)
|
||||
if not apply:
|
||||
for r in rows[:5]:
|
||||
logger.info(" %-42s %s", (r.get("product_name") or "")[:42], r["primary_url"][:90])
|
||||
if len(rows) > 5:
|
||||
logger.info(" ... and %d more", len(rows) - 5)
|
||||
return entry
|
||||
|
||||
results: Dict[str, Tuple[Optional[List[int]], str]] = {}
|
||||
with ThreadPoolExecutor(max_workers=workers) as pool:
|
||||
futures = {
|
||||
pool.submit(pacer.fetch, r["primary_url"], timeout): r["image_id"] for r in rows
|
||||
}
|
||||
for fut in as_completed(futures):
|
||||
image_id = futures[fut]
|
||||
try:
|
||||
results[image_id] = fut.result()
|
||||
except Exception as exc: # noqa: BLE001 - one bad URL must not stop the table
|
||||
logger.warning(" %s: %s", image_id, exc)
|
||||
results[image_id] = (None, "")
|
||||
|
||||
# Writes happen here, on the caller's connection, in row order.
|
||||
for r in rows:
|
||||
vector, src = results.get(r["image_id"], (None, ""))
|
||||
if vector is None:
|
||||
entry["failed"] += 1
|
||||
logger.info(" failed %-42s %s", (r.get("product_name") or "")[:42], r["primary_url"][:80])
|
||||
continue
|
||||
iv.store_vector(cur, table, r["image_id"], vector, src)
|
||||
entry["computed"] += 1
|
||||
conn.commit()
|
||||
logger.info("%s: %d written, %d failed", table, entry["computed"], entry["failed"])
|
||||
return entry
|
||||
|
||||
|
||||
def main() -> int:
|
||||
parser = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
|
||||
mode = parser.add_mutually_exclusive_group(required=True)
|
||||
mode.add_argument("--brands", help="comma-separated brands to fill")
|
||||
mode.add_argument("--all", action="store_true", help="every brand table, including inactive brands")
|
||||
parser.add_argument("--apply", action="store_true", help="download and write (default: dry run)")
|
||||
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("--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,
|
||||
help="seconds between two requests to the same host")
|
||||
parser.add_argument("--json", action="store_true", help="machine-readable summary")
|
||||
args = parser.parse_args()
|
||||
|
||||
workers = max(1, min(MAX_WORKERS, args.workers))
|
||||
|
||||
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")
|
||||
|
||||
conn = _connect()
|
||||
if conn is None:
|
||||
logger.error("No database connection.")
|
||||
return 1
|
||||
|
||||
report: List[Dict[str, Any]] = []
|
||||
pacer = _HostPacer(args.pause)
|
||||
try:
|
||||
with conn.cursor() as cur:
|
||||
brands = None if args.all else [b for b in args.brands.split(",") if b.strip()]
|
||||
for _suffix, table in _brand_tables(cur, brands):
|
||||
report.append(backfill_table(
|
||||
cur, conn, table,
|
||||
apply=args.apply, recompute_stale=args.recompute_stale,
|
||||
limit=args.limit, workers=workers, timeout=args.timeout, pacer=pacer,
|
||||
))
|
||||
finally:
|
||||
conn.close()
|
||||
|
||||
candidates = sum(e["candidates"] for e in report)
|
||||
computed = sum(e["computed"] for e in report)
|
||||
failed = sum(e["failed"] for e in report)
|
||||
skipped = sum(1 for e in report if e["skipped"])
|
||||
|
||||
if args.json:
|
||||
print(json.dumps({"apply": args.apply, "tables": report,
|
||||
"candidates": candidates, "computed": computed,
|
||||
"failed": failed, "skipped_tables": skipped}, indent=2))
|
||||
else:
|
||||
logger.info("")
|
||||
if args.apply:
|
||||
logger.info("Applied: %d candidate(s), %d written, %d failed, %d table(s) skipped",
|
||||
candidates, computed, failed, skipped)
|
||||
else:
|
||||
logger.info("Dry run: %d row(s) would be attempted across %d table(s), %d table(s) skipped",
|
||||
candidates, len(report) - skipped, skipped)
|
||||
if candidates:
|
||||
logger.info("Re-run with --apply to download and write.")
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
sys.exit(main())
|
||||
@@ -108,8 +108,33 @@ _TYPE_ALIASES = {
|
||||
"JSONB": {"jsonb"},
|
||||
"DOUBLE PRECISION": {"double precision"},
|
||||
"vector(384)": {"USER-DEFINED"},
|
||||
"vector(3072)": {"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)
|
||||
|
||||
|
||||
def _pgvector_version(cur) -> Optional[str]:
|
||||
try:
|
||||
cur.execute("SELECT extversion FROM pg_extension WHERE extname = 'vector'")
|
||||
row = cur.fetchone()
|
||||
return str(row[0]) if row and row[0] else None
|
||||
except Exception: # noqa: BLE001 - informational; a failed probe is reported as unknown
|
||||
return None
|
||||
|
||||
|
||||
def _version_tuple(text: str) -> tuple:
|
||||
parts = []
|
||||
for piece in text.split("."):
|
||||
digits = "".join(ch for ch in piece if ch.isdigit())
|
||||
if not digits:
|
||||
break
|
||||
parts.append(int(digits))
|
||||
return tuple(parts)
|
||||
|
||||
|
||||
class RecordingCursor:
|
||||
"""Wraps a real cursor so a dry run can see the statements without
|
||||
@@ -155,6 +180,13 @@ def main() -> int:
|
||||
|
||||
logger.info("database : %s / %s", DB_HOST, DB_NAME)
|
||||
logger.info("mode : %s", "APPLY (writing)" if apply else "dry run (no writes)")
|
||||
with conn.cursor() as cur:
|
||||
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).",
|
||||
pgv, ".".join(map(str, _MIN_PGVECTOR_FOR_IMG_VECTOR)))
|
||||
logger.info("")
|
||||
|
||||
declared_types = _declared_types()
|
||||
|
||||
@@ -85,6 +85,12 @@ os.environ["ACTIVE_BRANDS"] = ""
|
||||
# turns it on explicitly, the same arrangement ACTIVE_BRANDS has above.
|
||||
os.environ["AUTO_ENRICH_ON_UPLOAD"] = "false"
|
||||
|
||||
# Same arrangement for the image-vector worker: every call to
|
||||
# `upsert_brand_products` would otherwise start a thread that opens a
|
||||
# database connection and downloads product images. Off here; the hook and
|
||||
# the worker are covered explicitly in tests/test_image_vector.py.
|
||||
os.environ["ENABLE_IMAGE_VECTORS"] = "false"
|
||||
|
||||
# Auth is set unconditionally (not setdefault): the suite asserts on the real
|
||||
# guards, so it must never inherit a developer's AUTH_ENABLED=false.
|
||||
os.environ["AUTH_ENABLED"] = "true"
|
||||
|
||||
649
tests/test_image_vector.py
Normal file
649
tests/test_image_vector.py
Normal file
@@ -0,0 +1,649 @@
|
||||
"""img_vector: the pixel thumbnail of each product's primary image.
|
||||
|
||||
Four 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
|
||||
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
|
||||
when disabled, and a full queue is a log line, not an error.
|
||||
|
||||
No database is involved anywhere; cursors are recorders.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import io
|
||||
import queue
|
||||
import re
|
||||
import threading
|
||||
from typing import Any, Dict, List
|
||||
|
||||
import pytest
|
||||
from PIL import Image
|
||||
|
||||
from app.services import image_vector as iv
|
||||
from app.services import vector_store
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def _png(im: Image.Image) -> bytes:
|
||||
buf = io.BytesIO()
|
||||
im.save(buf, "PNG")
|
||||
return buf.getvalue()
|
||||
|
||||
|
||||
def _jpeg(im: Image.Image, **kw) -> bytes:
|
||||
buf = io.BytesIO()
|
||||
im.save(buf, "JPEG", quality=95, **kw)
|
||||
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():
|
||||
vector_store._invalidate_product_columns_cache()
|
||||
yield
|
||||
vector_store._invalidate_product_columns_cache()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 1. bytes -> vector
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def test_a_solid_png_becomes_3072_values_of_that_colour():
|
||||
vec = iv.pixel_vector(_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}
|
||||
|
||||
|
||||
def test_a_solid_jpeg_is_within_lossy_tolerance():
|
||||
vec = iv.pixel_vector(_jpeg(Image.new("RGB", (300, 300), (40, 120, 200))))
|
||||
|
||||
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])
|
||||
|
||||
|
||||
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.
|
||||
"""
|
||||
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()))
|
||||
|
||||
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)
|
||||
|
||||
|
||||
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))
|
||||
|
||||
assert vec is not None
|
||||
assert set(vec) == {255}
|
||||
|
||||
|
||||
@pytest.mark.parametrize("mode", ["P", "L", "LA", "CMYK", "I;16"])
|
||||
def test_every_pillow_mode_a_product_photo_could_arrive_in_decodes(mode):
|
||||
base = Image.new("RGB", (40, 40), (10, 200, 30))
|
||||
im = base.convert(mode) if mode != "P" else base.quantize(16)
|
||||
buf = io.BytesIO()
|
||||
im.save(buf, "TIFF" if mode in ("CMYK", "I;16") else "PNG")
|
||||
|
||||
vec = iv.pixel_vector(buf.getvalue())
|
||||
|
||||
assert vec is not None and len(vec) == 3072
|
||||
assert all(0 <= v <= 255 for v in vec)
|
||||
|
||||
|
||||
def test_the_first_frame_of_an_animated_gif_is_used():
|
||||
frames = [Image.new("P", (20, 20), c) for c in (1, 2, 3)]
|
||||
buf = io.BytesIO()
|
||||
frames[0].save(buf, "GIF", save_all=True, append_images=frames[1:])
|
||||
|
||||
assert iv.pixel_vector(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
|
||||
|
||||
|
||||
def test_a_truncated_file_is_none():
|
||||
noisy = Image.effect_noise((64, 64), 50).convert("RGB")
|
||||
whole = _png(noisy)
|
||||
assert len(whole) > 400, "need a file big enough to cut"
|
||||
|
||||
assert iv.pixel_vector(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)
|
||||
|
||||
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
|
||||
|
||||
|
||||
def test_to_pg_is_the_same_text_form_the_upsert_uses_for_embedding():
|
||||
assert iv.to_pg([0, 128, 255]) == "[0,128,255]"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 1b. which image
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def test_primary_is_image_url_then_image_urls_zero_like_the_card():
|
||||
assert iv.primary_image_url({"image_url": "https://a/x.jpg", "image_urls": ["https://b/y.jpg"]}) == "https://a/x.jpg"
|
||||
assert iv.primary_image_url({"image_url": "", "image_urls": ["https://b/y.jpg", "https://c/z.jpg"]}) == "https://b/y.jpg"
|
||||
assert iv.primary_image_url({"image_url": None, "image_urls": None}) is None
|
||||
assert iv.primary_image_url({}) is None
|
||||
|
||||
|
||||
def test_primary_is_normalised_the_way_the_card_fallbacks_are():
|
||||
assert iv.primary_image_url({"image_url": "//cdn.example/x.jpg"}) == "https://cdn.example/x.jpg"
|
||||
dead = next(iter(vector_store._DEAD_IMAGE_HOSTS))
|
||||
assert iv.primary_image_url({"image_url": f"https://{dead}/daily/x.jpg"}) is None
|
||||
assert iv.primary_image_url({"image_url": "not a url"}) is None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 1c. URL -> bytes
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class _Resp:
|
||||
def __init__(self, status: int, body: bytes, content_type: str = "image/jpeg"):
|
||||
self.status_code = status
|
||||
self.headers = {"content-type": content_type}
|
||||
self._body = body
|
||||
self.closed = False
|
||||
|
||||
def iter_content(self, chunk_size=1):
|
||||
for i in range(0, len(self._body), chunk_size):
|
||||
yield self._body[i:i + chunk_size]
|
||||
|
||||
def close(self):
|
||||
self.closed = True
|
||||
|
||||
|
||||
def _patch_requests(monkeypatch, responses: List[_Resp]):
|
||||
calls: List[Dict[str, Any]] = []
|
||||
|
||||
def fake_get(url, headers=None, timeout=None, stream=False):
|
||||
calls.append({"url": url, "headers": headers})
|
||||
return responses.pop(0)
|
||||
|
||||
import requests
|
||||
monkeypatch.setattr(requests, "get", fake_get)
|
||||
return calls
|
||||
|
||||
|
||||
def test_a_403_with_referer_is_retried_without_one(monkeypatch):
|
||||
body = b"x" * (iv.MIN_IMAGE_BYTES + 10)
|
||||
calls = _patch_requests(monkeypatch, [_Resp(403, b""), _Resp(200, body)])
|
||||
|
||||
assert iv.download_image_bytes("https://cdn.example/p/1.jpg") == body
|
||||
assert "Referer" in calls[0]["headers"]
|
||||
assert "Referer" not in calls[1]["headers"]
|
||||
|
||||
|
||||
def test_a_download_past_the_byte_cap_is_abandoned(monkeypatch):
|
||||
monkeypatch.setattr(iv, "IMAGE_VECTOR_MAX_BYTES", 1000)
|
||||
_patch_requests(monkeypatch, [_Resp(200, b"x" * 5000)])
|
||||
|
||||
assert iv.download_image_bytes("https://cdn.example/p/1.jpg") is None
|
||||
|
||||
|
||||
def test_a_placeholder_below_min_image_bytes_is_rejected(monkeypatch):
|
||||
_patch_requests(monkeypatch, [_Resp(200, b"x" * 10), _Resp(200, b"x" * 10)])
|
||||
|
||||
assert iv.download_image_bytes("https://cdn.example/p/1.jpg") is None
|
||||
|
||||
|
||||
def test_a_non_image_content_type_without_an_image_extension_is_rejected(monkeypatch):
|
||||
_patch_requests(monkeypatch, [_Resp(200, b"<html>" * 1000, "text/html")])
|
||||
|
||||
assert iv.download_image_bytes("https://cdn.example/p/1") is None
|
||||
|
||||
|
||||
def test_text_html_is_rejected_even_when_the_url_ends_in_png(monkeypatch):
|
||||
"""uat.amul.com redirects every dead path to its homepage. The .png in the
|
||||
URL is no evidence; the server's own content-type is."""
|
||||
_patch_requests(monkeypatch, [_Resp(200, b"<html>" * 1000, "text/html; charset=utf-8")])
|
||||
|
||||
assert iv.download_image_bytes("https://uat.amul.com/files/products/amul-ghee1.png") is None
|
||||
|
||||
|
||||
def test_an_unlabelled_body_with_an_image_extension_is_accepted(monkeypatch):
|
||||
body = b"x" * (iv.MIN_IMAGE_BYTES + 10)
|
||||
_patch_requests(monkeypatch, [_Resp(200, body, "application/octet-stream")])
|
||||
|
||||
assert iv.download_image_bytes("https://cdn.example/p/1.jpg") == body
|
||||
|
||||
|
||||
def test_a_raised_request_is_none_not_an_exception(monkeypatch):
|
||||
import requests
|
||||
|
||||
def boom(*a, **k):
|
||||
raise requests.ConnectionError("no route")
|
||||
|
||||
monkeypatch.setattr(requests, "get", boom)
|
||||
|
||||
assert iv.download_image_bytes("https://cdn.example/p/1.jpg") is None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 2. the columns exist, and an old pgvector loses only the vector
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class MigrationCursor:
|
||||
def __init__(self, refuse: str = ""):
|
||||
self.statements: List[str] = []
|
||||
self._refuse = refuse
|
||||
|
||||
def execute(self, sql, params=None):
|
||||
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')
|
||||
|
||||
def fetchall(self):
|
||||
return []
|
||||
|
||||
|
||||
def test_the_migration_adds_both_columns_to_an_existing_table():
|
||||
cur = MigrationCursor()
|
||||
|
||||
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_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)")
|
||||
|
||||
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
|
||||
|
||||
|
||||
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_src TEXT" in ddl
|
||||
for line in ddl.splitlines():
|
||||
if "CREATE INDEX" in line:
|
||||
assert "img_vector" not in line, line
|
||||
|
||||
|
||||
def test_the_migration_script_can_still_read_col_defs_as_literals():
|
||||
from scripts.migrate_brand_schema import _declared_types
|
||||
|
||||
declared = _declared_types()
|
||||
|
||||
assert declared["img_vector"] == "vector(3072)"
|
||||
assert declared["img_vector_src"] == "TEXT"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 3. the upsert never names them; the readers never select them
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def _insert_statement() -> str:
|
||||
import inspect
|
||||
body = inspect.getsource(vector_store.upsert_brand_products)
|
||||
match = re.search(r"INSERT INTO \{table_name\}.*?updated_at = CURRENT_TIMESTAMP", body, re.S)
|
||||
assert match
|
||||
return " ".join(match.group(0).split())
|
||||
|
||||
|
||||
def test_the_insert_does_not_name_the_vector_columns():
|
||||
"""A column the statement never names is a column it cannot damage. Every
|
||||
writer reaching the upsert builds its dict from a sheet, the seed JSON or
|
||||
a read-back row; none can carry a pixel vector, so naming the column here
|
||||
would only ever NULL it on the next re-seed."""
|
||||
statement = _insert_statement()
|
||||
|
||||
assert "img_vector" not in statement
|
||||
|
||||
|
||||
def test_the_write_hook_runs_after_the_write_and_cannot_raise_into_it():
|
||||
import inspect
|
||||
body = inspect.getsource(vector_store.upsert_brand_products)
|
||||
|
||||
hook = body.index("_schedule_image_vectors(brand, expected_ids)")
|
||||
assert body.index("return persisted") > hook > body.index("invalidate_brand_overview_cache()")
|
||||
assert "except Exception" in body[body.rindex("try:", 0, hook):hook + 200]
|
||||
|
||||
|
||||
class ColumnsCursor:
|
||||
"""Answers the information_schema probe, the table-exists probe and one
|
||||
product SELECT; records every statement."""
|
||||
|
||||
def __init__(self, tables: Dict[str, List[str]]):
|
||||
self.tables = tables
|
||||
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 = [(t, c) for t, cols in self.tables.items() for c in cols]
|
||||
elif "information_schema.tables" in text:
|
||||
self._pending = (True,)
|
||||
else:
|
||||
self.description = [("id",), ("product_name",)]
|
||||
self._pending = []
|
||||
|
||||
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 test_product_columns_drops_both_vectors_and_keeps_the_provenance_url():
|
||||
cur = ColumnsCursor({"brand_x": ["id", "embedding", "img_vector", "img_vector_src", "size"]})
|
||||
|
||||
assert vector_store._product_columns(cur, "brand_x") == '"id", "img_vector_src", "size"'
|
||||
|
||||
|
||||
def test_an_unknown_table_falls_back_to_star_rather_than_failing():
|
||||
cur = ColumnsCursor({"brand_x": ["id"]})
|
||||
|
||||
assert vector_store._product_columns(cur, "brand_new") == "*"
|
||||
|
||||
|
||||
def test_the_column_probe_is_cached_until_invalidated():
|
||||
cur = ColumnsCursor({"brand_x": ["id", "embedding"]})
|
||||
|
||||
vector_store._product_columns(cur, "brand_x")
|
||||
vector_store._product_columns(cur, "brand_x")
|
||||
probes = [s for s in cur.statements if "information_schema.columns" in s]
|
||||
assert len(probes) == 1
|
||||
|
||||
vector_store._invalidate_product_columns_cache()
|
||||
vector_store._product_columns(cur, "brand_x")
|
||||
probes = [s for s in cur.statements if "information_schema.columns" in s]
|
||||
assert len(probes) == 2
|
||||
|
||||
|
||||
class _Conn:
|
||||
def __init__(self, cur):
|
||||
self._cur = cur
|
||||
|
||||
def cursor(self):
|
||||
return self
|
||||
|
||||
def __enter__(self):
|
||||
return self._cur
|
||||
|
||||
def __exit__(self, *a):
|
||||
return False
|
||||
|
||||
def close(self):
|
||||
pass
|
||||
|
||||
|
||||
def test_get_products_by_brand_selects_by_name_and_never_the_vectors(monkeypatch):
|
||||
cur = ColumnsCursor({"brand_cadbury": ["id", "product_name", "image_url", "embedding", "img_vector", "img_vector_src"]})
|
||||
monkeypatch.setattr(vector_store, "_connect", lambda: _Conn(cur))
|
||||
|
||||
vector_store.get_products_by_brand("Cadbury")
|
||||
|
||||
select = [s for s in cur.statements if s.startswith("SELECT") and "FROM brand_cadbury" in s]
|
||||
assert select, cur.statements
|
||||
assert select[0].startswith('SELECT "id", "product_name", "image_url", "img_vector_src" FROM brand_cadbury')
|
||||
assert "embedding" not in select[0]
|
||||
assert '"img_vector"' not in select[0]
|
||||
|
||||
|
||||
def test_get_product_by_image_id_uses_the_same_projection(monkeypatch):
|
||||
cur = ColumnsCursor({"brand_cadbury": ["id", "embedding", "img_vector"]})
|
||||
monkeypatch.setattr(vector_store, "_connect", lambda: _Conn(cur))
|
||||
|
||||
vector_store.get_product_by_image_id("Cadbury", "cadbury_x")
|
||||
|
||||
select = [s for s in cur.statements if "WHERE image_id = %s LIMIT 1" in s]
|
||||
assert select and select[0].startswith('SELECT "id" FROM brand_cadbury')
|
||||
|
||||
|
||||
def test_no_product_reader_is_select_star_any_more():
|
||||
import inspect
|
||||
for fn in (vector_store.get_products_by_brand, vector_store.get_product_by_image_id,
|
||||
vector_store.semantic_search, vector_store.text_search, vector_store.lexical_search):
|
||||
assert "SELECT *" not in inspect.getsource(fn), fn.__name__
|
||||
|
||||
|
||||
def test_the_mcp_slimmer_hides_the_vector_columns():
|
||||
from app.mcp_server import _slim
|
||||
|
||||
out = _slim({"title": "x", "img_vector": "[1,2]", "img_vector_src": "https://a/x.jpg", "embedding": "[0.1]"})
|
||||
|
||||
assert out == {"title": "x"}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 4. the hook
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def test_schedule_is_a_no_op_when_disabled(monkeypatch):
|
||||
monkeypatch.setattr(iv, "ENABLE_IMAGE_VECTORS", False)
|
||||
q: "queue.Queue" = queue.Queue(maxsize=4)
|
||||
monkeypatch.setattr(iv, "_queue", q)
|
||||
started = []
|
||||
monkeypatch.setattr(iv, "_ensure_worker", lambda: started.append(1))
|
||||
|
||||
assert iv.schedule("Amul", ["a", "b"]) is False
|
||||
assert q.qsize() == 0 and not started
|
||||
|
||||
|
||||
def test_schedule_enqueues_and_the_worker_calls_backfill(monkeypatch):
|
||||
monkeypatch.setattr(iv, "ENABLE_IMAGE_VECTORS", True)
|
||||
q: "queue.Queue" = queue.Queue(maxsize=4)
|
||||
monkeypatch.setattr(iv, "_queue", q)
|
||||
monkeypatch.setattr(iv, "_worker", None)
|
||||
done = threading.Event()
|
||||
calls: List[Any] = []
|
||||
|
||||
def fake_backfill(brand, ids, **kw):
|
||||
calls.append((brand, ids, kw))
|
||||
done.set()
|
||||
return {}
|
||||
|
||||
monkeypatch.setattr(iv, "backfill_rows", fake_backfill)
|
||||
|
||||
assert iv.schedule("Amul", ["a", "", "b"]) is True
|
||||
assert done.wait(5), "worker never ran"
|
||||
assert calls == [("Amul", ["a", "b"], {"recompute_stale": True})]
|
||||
assert q.unfinished_tasks == 0 or q.join() is None
|
||||
|
||||
|
||||
def test_a_full_queue_is_logged_and_dropped_not_raised(monkeypatch, caplog):
|
||||
monkeypatch.setattr(iv, "ENABLE_IMAGE_VECTORS", True)
|
||||
q: "queue.Queue" = queue.Queue(maxsize=1)
|
||||
q.put_nowait(("X", ["1"]))
|
||||
monkeypatch.setattr(iv, "_queue", q)
|
||||
monkeypatch.setattr(iv, "_ensure_worker", lambda: None)
|
||||
|
||||
with caplog.at_level("WARNING", logger=iv.__name__):
|
||||
assert iv.schedule("Amul", ["a"]) is False
|
||||
assert "queue full" in caplog.text
|
||||
|
||||
|
||||
def test_an_empty_id_list_is_not_queued(monkeypatch):
|
||||
monkeypatch.setattr(iv, "ENABLE_IMAGE_VECTORS", True)
|
||||
q: "queue.Queue" = queue.Queue(maxsize=4)
|
||||
monkeypatch.setattr(iv, "_queue", q)
|
||||
monkeypatch.setattr(iv, "_ensure_worker", lambda: None)
|
||||
|
||||
assert iv.schedule("Amul", ["", None]) is False
|
||||
assert q.qsize() == 0
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 5. the backfill: which rows, and dry run writes nothing
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class RowsCursor:
|
||||
"""Returns canned rows for the candidate SELECT and records UPDATEs."""
|
||||
|
||||
def __init__(self, rows: List[Dict[str, Any]], has_columns: bool = True):
|
||||
self._rows = rows
|
||||
self._has = has_columns
|
||||
self.statements: List[tuple] = []
|
||||
self.description = None
|
||||
self._pending: Any = None
|
||||
|
||||
def execute(self, sql, params=None):
|
||||
text = " ".join(str(sql).split())
|
||||
self.statements.append((text, params))
|
||||
if "information_schema.columns" in text:
|
||||
self._pending = [("img_vector",), ("img_vector_src",)] if self._has else []
|
||||
elif text.startswith("SELECT image_id"):
|
||||
keys = ["image_id", "product_name", "image_url", "image_urls", "img_vector_src", "missing"]
|
||||
self.description = [(k,) for k in keys]
|
||||
self._pending = [tuple(r.get(k) for k in keys) for r in self._rows]
|
||||
else:
|
||||
self._pending = []
|
||||
|
||||
def fetchall(self):
|
||||
return self._pending or []
|
||||
|
||||
def __enter__(self):
|
||||
return self
|
||||
|
||||
def __exit__(self, *a):
|
||||
return False
|
||||
|
||||
|
||||
_ROWS = [
|
||||
{"image_id": "missing", "image_url": "https://a/1.jpg", "image_urls": [], "img_vector_src": None, "missing": True},
|
||||
{"image_id": "stale", "image_url": "https://a/2-new.jpg", "image_urls": [], "img_vector_src": "https://a/2-old.jpg", "missing": False},
|
||||
{"image_id": "fresh", "image_url": "https://a/3.jpg", "image_urls": [], "img_vector_src": "https://a/3.jpg", "missing": False},
|
||||
{"image_id": "gallery-only", "image_url": "", "image_urls": ["https://a/4.jpg"], "img_vector_src": None, "missing": True},
|
||||
{"image_id": "no-image", "image_url": None, "image_urls": [], "img_vector_src": None, "missing": True},
|
||||
]
|
||||
|
||||
|
||||
def test_rows_needing_vectors_picks_missing_and_stale_with_a_primary():
|
||||
cur = RowsCursor(_ROWS)
|
||||
|
||||
got = iv.rows_needing_vectors(cur, "brand_a", recompute_stale=True)
|
||||
|
||||
assert [r["image_id"] for r in got] == ["missing", "stale", "gallery-only"]
|
||||
assert got[2]["primary_url"] == "https://a/4.jpg"
|
||||
|
||||
|
||||
def test_rows_needing_vectors_leaves_stale_rows_alone_unless_asked():
|
||||
cur = RowsCursor(_ROWS)
|
||||
|
||||
got = iv.rows_needing_vectors(cur, "brand_a", recompute_stale=False)
|
||||
|
||||
assert [r["image_id"] for r in got] == ["missing", "gallery-only"]
|
||||
|
||||
|
||||
def test_rows_needing_vectors_scopes_to_the_ids_just_written():
|
||||
cur = RowsCursor(_ROWS)
|
||||
|
||||
iv.rows_needing_vectors(cur, "brand_a", ["missing", "stale"])
|
||||
|
||||
text, params = cur.statements[-1]
|
||||
assert "image_id = ANY(%s)" in text and params == [["missing", "stale"]]
|
||||
|
||||
|
||||
def test_store_vector_does_not_touch_updated_at():
|
||||
cur = RowsCursor([])
|
||||
|
||||
iv.store_vector(cur, "brand_a", "x", [1, 2, 3], "https://a/x.jpg")
|
||||
|
||||
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 "updated_at" not in text
|
||||
|
||||
|
||||
class _RowsConn:
|
||||
def __init__(self, cur):
|
||||
self.cur = cur
|
||||
|
||||
def cursor(self):
|
||||
return self.cur
|
||||
|
||||
def close(self):
|
||||
pass
|
||||
|
||||
|
||||
def test_backfill_dry_run_downloads_and_writes_nothing(monkeypatch):
|
||||
cur = RowsCursor(_ROWS)
|
||||
monkeypatch.setattr(vector_store, "_connect", lambda: _RowsConn(cur))
|
||||
fetched: List[str] = []
|
||||
monkeypatch.setattr(iv, "vector_for_url", lambda url, timeout=None: (fetched.append(url) or [0] * 3072, url))
|
||||
|
||||
totals = iv.backfill_rows("A", dry_run=True, pause_seconds=0)
|
||||
|
||||
assert totals["candidates"] == 3 and totals["computed"] == 3 and totals["failed"] == 0
|
||||
assert not any(t.startswith("UPDATE") for t, _ in cur.statements)
|
||||
|
||||
|
||||
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(
|
||||
iv, "vector_for_url",
|
||||
lambda url, timeout=None: (None, url) if "4.jpg" in url else ([7] * 3072, url),
|
||||
)
|
||||
|
||||
totals = iv.backfill_rows("A", pause_seconds=0)
|
||||
|
||||
updates = [(t, p) for t, p in cur.statements if t.startswith("UPDATE")]
|
||||
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"
|
||||
|
||||
|
||||
def test_backfill_skips_a_table_the_migration_has_not_reached(monkeypatch):
|
||||
cur = RowsCursor(_ROWS, has_columns=False)
|
||||
monkeypatch.setattr(vector_store, "_connect", lambda: _RowsConn(cur))
|
||||
|
||||
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_backfill_never_raises(monkeypatch):
|
||||
def boom():
|
||||
raise RuntimeError("db exploded")
|
||||
|
||||
monkeypatch.setattr(vector_store, "_connect", lambda: _RowsConn(type("C", (), {"execute": lambda *a: boom()})()))
|
||||
|
||||
assert iv.backfill_rows("A", pause_seconds=0)["computed"] == 0
|
||||
Reference in New Issue
Block a user