image vector dimensionality reduction
This commit is contained in:
72
tests/test_image_embedder_model.py
Normal file
72
tests/test_image_embedder_model.py
Normal file
@@ -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
|
||||
@@ -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):
|
||||
|
||||
Reference in New Issue
Block a user