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
This commit is contained in:
co-authored by
Claude Opus 5
parent
7ee5ceea6c
commit
c56864877a
+14
-13
@@ -20,14 +20,14 @@ on a later turn because the cursors only advance on success.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import math
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from . import models
|
||||
from . import models, vectors
|
||||
from .context import history, story_actions, truncate_to_last_tokens
|
||||
from .database import SessionLocal
|
||||
from .providers import OpenAICompatibleProvider, ProviderError
|
||||
from .vectors import cosine # re-exported: the ranking lives here, the maths there
|
||||
|
||||
MEMORY_INTERVAL = 6 # actions per memory
|
||||
MEMORY_START = 12 # first memory once the adventure reaches this many actions
|
||||
@@ -79,14 +79,15 @@ def embedding_provider(settings: models.Settings) -> OpenAICompatibleProvider:
|
||||
)
|
||||
|
||||
|
||||
def cosine(a: list[float], b: list[float]) -> float:
|
||||
# Different lengths means the embedding model changed since this vector was
|
||||
# stored; zip() would silently score garbage.
|
||||
if len(a) != len(b):
|
||||
return 0.0
|
||||
dot = sum(x * y for x, y in zip(a, b))
|
||||
norm = math.sqrt(sum(x * x for x in a)) * math.sqrt(sum(y * y for y in b))
|
||||
return dot / norm if norm else 0.0
|
||||
def set_vector(memory: models.Memory, vector: list[float] | None) -> None:
|
||||
"""Store (or clear) a memory's embedding.
|
||||
|
||||
Both columns, always together: `embedding_blob` is what will be read, and
|
||||
the JSON `embedding` stays correct behind it until the follow-up migration
|
||||
drops it. Going through one function is what keeps them from drifting.
|
||||
"""
|
||||
memory.embedding = vector
|
||||
memory.embedding_blob = None if vector is None else vectors.pack(vector)
|
||||
|
||||
|
||||
def settled_count(adventure: models.Adventure) -> int:
|
||||
@@ -395,11 +396,11 @@ async def _embed_pending(
|
||||
if not pending:
|
||||
return
|
||||
try:
|
||||
vectors = await embedding_provider(settings).embed([m.text for m in pending])
|
||||
new = await embedding_provider(settings).embed([m.text for m in pending])
|
||||
except ProviderError:
|
||||
return
|
||||
for memory, vector in zip(pending, vectors):
|
||||
memory.embedding = vector
|
||||
for memory, vector in zip(pending, new):
|
||||
set_vector(memory, vector)
|
||||
db.commit()
|
||||
|
||||
|
||||
|
||||
@@ -17,13 +17,18 @@ starts fresh (created by create_all, stamped LATEST, never replays them), but
|
||||
migrations added from Phase 9 on must run on both dialects.
|
||||
"""
|
||||
|
||||
import json
|
||||
|
||||
from sqlalchemy import inspect, text
|
||||
from sqlalchemy.engine import Engine
|
||||
|
||||
from . import vectors
|
||||
from .database import Base
|
||||
|
||||
# (version, SQL to run when upgrading past it) — append only, never reorder.
|
||||
MIGRATIONS: list[tuple[int, str]] = [
|
||||
# The SQL is a string, or a {dialect: sql} map with a "default" entry when the
|
||||
# two dialects have to be spelled differently (BLOB vs BYTEA and the like).
|
||||
MIGRATIONS: list[tuple[int, str | dict[str, str]]] = [
|
||||
# Phase 6: auto-summarization + memory bank (the `memories` table itself is
|
||||
# created by create_all, which runs for existing DBs too).
|
||||
(2, "ALTER TABLE adventures ADD COLUMN auto_summarize BOOLEAN NOT NULL DEFAULT 0"),
|
||||
@@ -121,6 +126,14 @@ MIGRATIONS: list[tuple[int, str]] = [
|
||||
# ~5 KB to every later load of that adventure. Backfilled by
|
||||
# _backfill_variant_count.
|
||||
(37, "ALTER TABLE actions ADD COLUMN variant_count INTEGER NOT NULL DEFAULT 0"),
|
||||
# Egress, round three: a 1536-dimension embedding written as a JSON list is
|
||||
# ~31 KB, and the whole bank is fetched every turn to rank it. Packed
|
||||
# float32 is 6 KB for the same numbers, exactly (see vectors.py). Dimensions
|
||||
# are unchanged, so this is a format conversion — no re-embedding, no API
|
||||
# calls. Backfilled by _backfill_embedding_blob. The old JSON column is left
|
||||
# in place and keeps being written until a follow-up migration drops it.
|
||||
(38, {"sqlite": "ALTER TABLE memories ADD COLUMN embedding_blob BLOB",
|
||||
"default": "ALTER TABLE memories ADD COLUMN embedding_blob BYTEA"}),
|
||||
]
|
||||
|
||||
LATEST_VERSION = max((v for v, _ in MIGRATIONS), default=1)
|
||||
@@ -128,6 +141,11 @@ LATEST_VERSION = max((v for v, _ in MIGRATIONS), default=1)
|
||||
# Migrations that need a data pass after their DDL, keyed by version.
|
||||
WORLD_DELTA_VERSION = 36
|
||||
VARIANT_COUNT_VERSION = 37
|
||||
EMBEDDING_BLOB_VERSION = 38
|
||||
|
||||
# Vectors converted per round trip. Small enough that the backfill never holds
|
||||
# more than a few megabytes, large enough that it isn't a query per row.
|
||||
BACKFILL_BATCH = 200
|
||||
|
||||
|
||||
def _backfill_world_delta(conn) -> None:
|
||||
@@ -183,6 +201,46 @@ def _backfill_variant_count(conn) -> None:
|
||||
conn.execute(text(sql))
|
||||
|
||||
|
||||
def _for_dialect(sql: str | dict[str, str], dialect: str) -> str:
|
||||
return sql if isinstance(sql, str) else sql.get(dialect, sql["default"])
|
||||
|
||||
|
||||
def _backfill_embedding_blob(conn) -> None:
|
||||
"""Repack memories.embedding (JSON list) into memories.embedding_blob.
|
||||
|
||||
The one backfill here that has to come through Python: struct packing has
|
||||
no portable SQL spelling, so unlike migrations 36 and 37 this pays a
|
||||
one-time read of every vector (~4 MB in production) to stop paying three
|
||||
megabytes every turn. Batched so the read is bounded whatever the bank
|
||||
grows to.
|
||||
|
||||
Reads the JSON defensively — SQLite hands back the raw string while psycopg
|
||||
has already parsed it into a list — and skips anything that isn't a
|
||||
non-empty list, so one malformed row can't strand the migration.
|
||||
"""
|
||||
last_id = 0
|
||||
while True:
|
||||
rows = conn.execute(
|
||||
text("""
|
||||
SELECT id, embedding FROM memories
|
||||
WHERE embedding IS NOT NULL AND embedding_blob IS NULL AND id > :last
|
||||
ORDER BY id LIMIT :batch
|
||||
"""),
|
||||
{"last": last_id, "batch": BACKFILL_BATCH},
|
||||
).all()
|
||||
if not rows:
|
||||
return
|
||||
for row_id, stored in rows:
|
||||
vector = json.loads(stored) if isinstance(stored, str) else stored
|
||||
if not isinstance(vector, list) or not vector:
|
||||
continue
|
||||
conn.execute(
|
||||
text("UPDATE memories SET embedding_blob = :blob WHERE id = :id"),
|
||||
{"blob": vectors.pack(vector), "id": row_id},
|
||||
)
|
||||
last_id = rows[-1][0]
|
||||
|
||||
|
||||
def _get_version(conn) -> int:
|
||||
if conn.dialect.name == "sqlite":
|
||||
return conn.execute(text("PRAGMA user_version")).scalar() or 1
|
||||
@@ -220,11 +278,13 @@ def bootstrap(engine: Engine) -> None:
|
||||
current = _get_version(conn)
|
||||
for version, sql in MIGRATIONS:
|
||||
if version > current:
|
||||
conn.execute(text(sql))
|
||||
conn.execute(text(_for_dialect(sql, conn.dialect.name)))
|
||||
if version == WORLD_DELTA_VERSION:
|
||||
_backfill_world_delta(conn)
|
||||
if version == VARIANT_COUNT_VERSION:
|
||||
_backfill_variant_count(conn)
|
||||
if version == EMBEDDING_BLOB_VERSION:
|
||||
_backfill_embedding_blob(conn)
|
||||
current = version
|
||||
_set_version(conn, current)
|
||||
_encrypt_plaintext_api_keys(conn)
|
||||
|
||||
+21
-4
@@ -1,7 +1,8 @@
|
||||
from datetime import datetime, timezone
|
||||
|
||||
from sqlalchemy import (
|
||||
JSON, Boolean, Column, DateTime, Float, ForeignKey, Integer, String, Table, Text,
|
||||
JSON, Boolean, Column, DateTime, Float, ForeignKey, Integer, LargeBinary, String,
|
||||
Table, Text,
|
||||
)
|
||||
from sqlalchemy.orm import Mapped, mapped_column, relationship
|
||||
|
||||
@@ -141,9 +142,15 @@ class Adventure(Base):
|
||||
class Memory(Base):
|
||||
"""Phase 6: an auto-summarized (or hand-written) fact about the adventure.
|
||||
|
||||
`embedding` is the raw vector as a JSON list (cosine ranking happens in
|
||||
Python — fine at bank sizes of a few hundred). NULL until embedded, which
|
||||
also marks it for backfill when an embedding model becomes available.
|
||||
The vector lives in `embedding_blob` as packed float32 (see vectors.py).
|
||||
NULL until embedded, which also marks it for backfill when an embedding
|
||||
model becomes available.
|
||||
|
||||
Cosine ranking happens in Python, which means the vectors cross the wire.
|
||||
The original comment here sized that by count — "fine at a few hundred" —
|
||||
and it was wrong by the only measure that mattered: a few hundred JSON
|
||||
vectors is ten megabytes, fetched fresh every turn. Weigh new columns in
|
||||
bytes.
|
||||
"""
|
||||
|
||||
__tablename__ = "memories"
|
||||
@@ -151,7 +158,17 @@ class Memory(Base):
|
||||
id: Mapped[int] = mapped_column(primary_key=True)
|
||||
adventure_id: Mapped[int] = mapped_column(ForeignKey("adventures.id", ondelete="CASCADE"))
|
||||
text: Mapped[str] = mapped_column(Text, default="")
|
||||
# Superseded by embedding_blob and written alongside it, so the two stay in
|
||||
# step until a follow-up migration drops this one. Still the column the
|
||||
# ranking path reads; that moves next.
|
||||
embedding: Mapped[list | None] = mapped_column(JSON, nullable=True)
|
||||
# The vector, little-endian float32. Deferred because it is wider than the
|
||||
# rest of the row put together and exactly one code path wants it: anything
|
||||
# bulk-loading memories (the Memories drawer, eviction, the embed queue)
|
||||
# must project the columns it needs rather than load whole entities.
|
||||
embedding_blob: Mapped[bytes | None] = mapped_column(
|
||||
LargeBinary, nullable=True, deferred=True
|
||||
)
|
||||
# Action index range this memory summarizes (null for manual memories).
|
||||
source_start: Mapped[int | None] = mapped_column(Integer, nullable=True)
|
||||
source_end: Mapped[int | None] = mapped_column(Integer, nullable=True)
|
||||
|
||||
@@ -1488,7 +1488,7 @@ def update_memory(
|
||||
raise HTTPException(404, "Memory not found")
|
||||
fields = {k: v for k, v in payload.model_dump(exclude_unset=True).items() if v is not None}
|
||||
if "text" in fields and fields["text"].strip() != memory.text:
|
||||
memory.embedding = None # re-embed on the next post-turn pass
|
||||
memorybank.set_vector(memory, None) # re-embed on the next post-turn pass
|
||||
for field, value in fields.items():
|
||||
setattr(memory, field, value)
|
||||
db.commit()
|
||||
|
||||
@@ -0,0 +1,40 @@
|
||||
"""Storing and comparing embedding vectors.
|
||||
|
||||
A 1536-dimension vector written as a JSON list is about 31 KB, because every
|
||||
component is spelled out as a decimal string of seventeen-odd digits. The same
|
||||
vector as packed float32 is 6,144 bytes — a straight 5x, and the memory bank is
|
||||
read in full on every turn, so those bytes are paid over and over.
|
||||
|
||||
**Float32 is not an approximation here.** The embedding endpoints return
|
||||
vectors computed in float32, rendered into JSON as the shortest decimal string
|
||||
that round-trips through a double; converting that back to float32 recovers the
|
||||
original bits exactly. Nothing is lost that was ever there, which is why the
|
||||
conversion needs no re-embedding and carries no retrieval-quality risk.
|
||||
|
||||
Dimensions are deliberately unchanged. Dropping to 512 or 768 would have saved
|
||||
another 3x and cost an API call per stored memory to re-embed, against a bank
|
||||
that the packing and the in-process cache together already make cheap.
|
||||
"""
|
||||
|
||||
import math
|
||||
import struct
|
||||
|
||||
|
||||
def pack(vector: list[float]) -> bytes:
|
||||
"""A vector as little-endian float32."""
|
||||
return struct.pack(f"<{len(vector)}f", *vector)
|
||||
|
||||
|
||||
def unpack(blob: bytes) -> list[float]:
|
||||
"""The inverse of `pack`. Length is implied: four bytes per component."""
|
||||
return list(struct.unpack(f"<{len(blob) // 4}f", blob))
|
||||
|
||||
|
||||
def cosine(a: list[float], b: list[float]) -> float:
|
||||
# Different lengths means the embedding model changed since this vector was
|
||||
# stored; zip() would silently score garbage.
|
||||
if len(a) != len(b):
|
||||
return 0.0
|
||||
dot = sum(x * y for x, y in zip(a, b))
|
||||
norm = math.sqrt(sum(x * x for x in a)) * math.sqrt(sum(y * y for y in b))
|
||||
return dot / norm if norm else 0.0
|
||||
@@ -0,0 +1,258 @@
|
||||
"""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")
|
||||
@@ -173,13 +173,18 @@ def build_fixture(args, rng: random.Random) -> tuple[int, int]:
|
||||
))
|
||||
|
||||
for i in range(args.memories):
|
||||
db.add(models.Memory(
|
||||
memory = models.Memory(
|
||||
adventure_id=adventure.id,
|
||||
text=f"{MEMORY_TEXT} ({i})",
|
||||
embedding=[rng.uniform(-1.0, 1.0) for _ in range(EMBEDDING_DIMS)],
|
||||
source_start=i * memorybank.MEMORY_INTERVAL,
|
||||
source_end=i * memorybank.MEMORY_INTERVAL + memorybank.MEMORY_INTERVAL - 1,
|
||||
))
|
||||
)
|
||||
# Through the same door the app uses, so the fixture cannot end up
|
||||
# storing vectors in a shape production never produces.
|
||||
memorybank.set_vector(
|
||||
memory, [rng.uniform(-1.0, 1.0) for _ in range(EMBEDDING_DIMS)]
|
||||
)
|
||||
db.add(memory)
|
||||
|
||||
db.commit()
|
||||
return adventure.id, user.id
|
||||
|
||||
Reference in New Issue
Block a user