Files
interactive-story/backend/tests/test_embedding_blob.py
T
parththakkar106andClaude Opus 5 c56864877a Store embeddings as packed float32 instead of a JSON list
A 1536-dimension vector spelled out as JSON decimals is ~31 KB. The same
numbers packed as float32 are 6,144 bytes, and the whole bank is read on
every turn, so those bytes are paid over and over.

It is a format change, not a precision trade: the endpoints compute in
float32 and render that into JSON, so converting back recovers the original
bits exactly. Nothing is re-embedded and no API call is made -- migration 38
is a pure repack of what is already stored.

Unlike migrations 36 and 37 this backfill cannot be expressed in portable
SQL, so it comes through Python, batched, and pays a one-time read of every
vector to stop paying three megabytes a turn.

The JSON column stays, still written through set_vector, so a rollback finds
the vectors intact. Reading from the blob comes next; a follow-up migration
drops the old column once that is verified.

Migration SQL can now be a {dialect: sql} map -- BLOB and BYTEA have no
common spelling, and every Postgres deploy replays this one.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_015CYEJKobJ2Re4Dv7qUoSA7
2026-08-16 21:17:44 +05:30

259 lines
9.0 KiB
Python

"""Migration 38: embeddings move from a JSON list to packed float32.
The conversion has to be exact, because nothing re-embeds — a memory whose
vector shifts is silently ranked wrong forever, with no error anywhere to say
so. So these tests check the numbers survive the round trip bit for bit, and
that the migration reaches every row however many there are.
python -m pytest tests/test_embedding_blob.py -v
"""
import os
import struct
import tempfile
_tmp = tempfile.NamedTemporaryFile(suffix=".db", delete=False)
_tmp.close()
os.environ["AIDND_DB_PATH"] = _tmp.name
os.environ.pop("AIDND_DATABASE_URL", None)
os.environ.pop("DATABASE_URL", None)
import json
import random
import pytest
from sqlalchemy import text
from app import memorybank, migrations, models, vectors
from app.database import Base, SessionLocal, engine
def float32(value: float) -> float:
"""`value` as the double nearest to its float32 truncation — what an
embedding endpoint's JSON actually holds."""
return struct.unpack("<f", struct.pack("<f", value))[0]
def sample_vector(rng: random.Random, dims: int = 1536) -> list[float]:
return [float32(rng.uniform(-1.0, 1.0)) for _ in range(dims)]
@pytest.fixture()
def db():
Base.metadata.create_all(bind=engine)
session = SessionLocal()
try:
yield session
finally:
session.close()
Base.metadata.drop_all(bind=engine)
@pytest.fixture()
def adventure(db):
user = models.User(is_guest=False, email="vectors@example.com")
db.add(user)
db.flush()
adv = models.Adventure(user_id=user.id, title="Cave", script_state={})
db.add(adv)
db.commit()
return adv
# ------------------------------------------------------------------- packing
def test_pack_round_trips_exactly():
"""Not "close enough": embedding endpoints compute in float32 and render
that into JSON, so packing back to float32 must be lossless."""
rng = random.Random(1)
vector = sample_vector(rng)
assert vectors.unpack(vectors.pack(vector)) == vector
def test_packed_vector_is_four_bytes_per_dimension():
"""The whole point: 1536 dims is 6 KB here against ~31 KB as JSON."""
vector = sample_vector(random.Random(2))
blob = vectors.pack(vector)
assert len(blob) == 1536 * 4
assert len(blob) < len(json.dumps(vector).encode()) / 4
def test_pack_handles_the_extremes():
vector = [float32(v) for v in (0.0, -0.0, 1.0, -1.0, 3.4028234663852886e38, 1e-38)]
assert vectors.unpack(vectors.pack(vector)) == vector
def test_pack_rounds_a_value_float32_cannot_hold():
"""The guarantee is exactness for vectors that came from an embedding
model, which computes in float32 — not for arbitrary doubles. Worth
pinning down, because it is the line the round-trip claim sits on."""
assert vectors.unpack(vectors.pack([1e-38]))[0] != 1e-38
assert vectors.unpack(vectors.pack([1e-38]))[0] == pytest.approx(1e-38)
def test_cosine_moved_but_still_reachable_from_memorybank():
"""Callers import it from memorybank; the maths lives in vectors."""
assert memorybank.cosine is vectors.cosine
assert vectors.cosine([1.0, 0.0], [1.0, 0.0]) == pytest.approx(1.0)
assert vectors.cosine([1.0, 0.0], [0.0, 1.0]) == pytest.approx(0.0)
assert vectors.cosine([1.0, 0.0], [1.0, 0.0, 0.0]) == 0.0 # length mismatch
# -------------------------------------------------------------- set_vector
def test_set_vector_writes_both_columns(db, adventure):
"""Until the follow-up migration drops the JSON column, it has to stay
correct — a rollback reads it."""
memory = models.Memory(adventure_id=adventure.id, text="a fact")
db.add(memory)
db.commit()
vector = sample_vector(random.Random(3), dims=8)
memorybank.set_vector(memory, vector)
db.commit()
db.expire_all()
assert memory.embedding == vector
assert vectors.unpack(memory.embedding_blob) == vector
def test_set_vector_none_clears_both(db, adventure):
"""Editing a memory's text drops its vector so the next pass re-embeds. A
blob left behind would keep ranking the old text."""
memory = models.Memory(adventure_id=adventure.id, text="a fact")
db.add(memory)
memorybank.set_vector(memory, sample_vector(random.Random(4), dims=8))
db.commit()
memorybank.set_vector(memory, None)
db.commit()
db.expire_all()
assert memory.embedding is None
assert memory.embedding_blob is None
# --------------------------------------------------------------- the backfill
def seed_json_only(db, adventure, count: int, dims: int = 64) -> dict[int, list[float]]:
"""Memories as they exist before the migration: JSON vector, no blob."""
rng = random.Random(count)
expected = {}
for i in range(count):
vector = sample_vector(rng, dims)
memory = models.Memory(
adventure_id=adventure.id, text=f"fact {i}", embedding=vector
)
db.add(memory)
db.flush()
expected[memory.id] = vector
db.commit()
db.execute(text("UPDATE memories SET embedding_blob = NULL"))
db.commit()
return expected
def test_backfill_converts_every_existing_vector(db, adventure):
expected = seed_json_only(db, adventure, count=5)
with engine.begin() as conn:
migrations._backfill_embedding_blob(conn)
db.expire_all()
for memory in db.query(models.Memory).all():
assert vectors.unpack(memory.embedding_blob) == expected[memory.id]
def test_backfill_reaches_past_one_batch(db, adventure):
"""It loops on id, and an off-by-one there would silently leave the tail
of a big bank unconverted — which reads as "not embedded yet"."""
count = migrations.BACKFILL_BATCH * 2 + 3
expected = seed_json_only(db, adventure, count=count, dims=4)
with engine.begin() as conn:
migrations._backfill_embedding_blob(conn)
db.expire_all()
memories = db.query(models.Memory).all()
assert len(memories) == count
assert all(m.embedding_blob is not None for m in memories)
assert all(vectors.unpack(m.embedding_blob) == expected[m.id] for m in memories)
def test_backfill_leaves_unembedded_memories_alone(db, adventure):
db.add(models.Memory(adventure_id=adventure.id, text="not embedded yet"))
db.commit()
with engine.begin() as conn:
migrations._backfill_embedding_blob(conn)
db.expire_all()
assert db.query(models.Memory).one().embedding_blob is None
def test_backfill_is_idempotent(db, adventure):
"""It runs once from bootstrap, but a half-finished run must be safe to
repeat, and rows already converted must not be rewritten."""
seed_json_only(db, adventure, count=3)
with engine.begin() as conn:
migrations._backfill_embedding_blob(conn)
db.expire_all()
first = {m.id: m.embedding_blob for m in db.query(models.Memory).all()}
with engine.begin() as conn:
migrations._backfill_embedding_blob(conn)
db.expire_all()
assert {m.id: m.embedding_blob for m in db.query(models.Memory).all()} == first
def test_backfill_skips_a_malformed_row_without_stopping(db, adventure):
"""One bad row must not strand every row after it — the loop orders by id,
so an exception here would leave the rest of the bank unconverted."""
expected = seed_json_only(db, adventure, count=2)
broken = models.Memory(adventure_id=adventure.id, text="broken")
db.add(broken)
db.commit()
db.execute(
text("UPDATE memories SET embedding = :bad WHERE id = :id"),
{"bad": '"not a list"', "id": broken.id},
)
db.commit()
with engine.begin() as conn:
migrations._backfill_embedding_blob(conn)
db.expire_all()
assert db.get(models.Memory, broken.id).embedding_blob is None
for memory_id, vector in expected.items():
assert vectors.unpack(db.get(models.Memory, memory_id).embedding_blob) == vector
# ------------------------------------------------------- the upgrade in full
def test_bootstrap_adds_the_column_and_backfills_it(db, adventure):
"""The path a deployed database actually takes: sitting at 37 without the
column, then started on this build."""
expected = seed_json_only(db, adventure, count=4)
db.close()
with engine.begin() as conn:
conn.execute(text("ALTER TABLE memories DROP COLUMN embedding_blob"))
conn.execute(text(f"PRAGMA user_version = {migrations.EMBEDDING_BLOB_VERSION - 1}"))
migrations.bootstrap(engine)
with engine.begin() as conn:
assert conn.execute(text("PRAGMA user_version")).scalar() == migrations.LATEST_VERSION
rows = conn.execute(text("SELECT id, embedding_blob FROM memories")).all()
assert len(rows) == len(expected)
for row_id, blob in rows:
assert vectors.unpack(blob) == expected[row_id]
def test_migration_38_is_spelled_for_both_dialects():
"""Every Postgres deploy replays migrations from 24 on, so a SQLite-only
ALTER here would break the live database and nothing else would notice."""
sql = dict(migrations.MIGRATIONS)[migrations.EMBEDDING_BLOB_VERSION]
assert "BLOB" in migrations._for_dialect(sql, "sqlite")
assert "BYTEA" in migrations._for_dialect(sql, "postgresql")