Files
interactive-story/backend/tests/test_embedding_blob.py
T
parththakkar106andClaude Opus 5 b7e53ae581 Rank the memory bank without reading the memory bank
Retrieval walked adventure.memories, so every turn loaded every row of the
bank with its vector attached -- 3.1 MB, 96% of everything a turn read. It
now asks SQL which memories are in play (an id and a flag per row), ranks
against vectors held in process, and fetches text only for the five it picks.

Two more callers were doing the same thing and the production SQL could not
see them: _evict_over_capacity walked the bank to count it, and _embed_pending
walked it to find the rows with no vector. Both are counts and filters the
database can do without sending anything back.

    one turn      3,258.7 kB -> 723.4 kB cold, 122.3 kB warm
    run_post_turn 3,139.1 kB -> 0.7 kB
    Insights      3,223.7 kB -> 117.9 kB
    Memories drawer  ~3.1 MB -> 23.6 kB

A played turn is turn plus post-turn work: 6.4 MB down to 123 kB.

The cache needs no invalidation callbacks, which is what makes it safe. A
vector can only change through set_vector, which drops that one entry;
anything that removes a memory from play leaves the catalogue query, and
entries missing from the catalogue are dropped on the next read. So eviction,
deletion and pruning have nothing to remember to call.

memories.embedded joins the blob, for the same reason actions.variant_count
sits beside actions.variants: with the vector deferred, every "is this
embedded?" check would otherwise be a 6 KB lazy load, once per row.

Capacity drops 200 -> 80, on retrieval quality as much as cost -- ranking two
hundred memories to pick five buries the five. Eviction was measured at scale
first: trimming 100 to 80 costs 0.8 kB and reads no vectors.

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

284 lines
10 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 list(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 list(vectors.unpack(vectors.pack(vector))) == vector
def test_unpack_returns_a_compact_array():
"""These are held in memory between turns, so the container matters: an
array("f") is the 4 bytes a component the column is, a list of Python
floats is eight times that."""
vector = sample_vector(random.Random(9))
unpacked = vectors.unpack(vectors.pack(vector))
assert unpacked.typecode == "f"
assert unpacked.itemsize == 4
assert len(unpacked) == len(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 list(vectors.unpack(memory.embedding_blob)) == vector
assert memory.embedded is True
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
assert memory.embedded is False
# --------------------------------------------------------------- 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, embedded = false"))
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 list(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(list(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():
blob = db.get(models.Memory, memory_id).embedding_blob
assert list(vectors.unpack(blob)) == vector
# ------------------------------------------------------- the upgrade in full
def test_bootstrap_adds_the_columns_and_backfills_them(db, adventure):
"""The path a deployed database actually takes: sitting at 37 with neither
new column, then started on this build."""
expected = seed_json_only(db, adventure, count=4)
unembedded = models.Memory(adventure_id=adventure.id, text="no vector yet")
db.add(unembedded)
db.commit()
unembedded_id = unembedded.id
db.close()
with engine.begin() as conn:
conn.execute(text("ALTER TABLE memories DROP COLUMN embedding_blob"))
conn.execute(text("ALTER TABLE memories DROP COLUMN embedded"))
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, embedded FROM memories")).all()
by_id = {row[0]: (row[1], row[2]) for row in rows}
assert len(by_id) == len(expected) + 1
for memory_id, vector in expected.items():
blob, embedded = by_id[memory_id]
assert list(vectors.unpack(blob)) == vector
assert embedded
# The flag has to follow the vector, not the row: a memory that was never
# embedded must still read as not embedded afterwards.
assert by_id[unembedded_id] == (None, False)
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")