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 asyncio
|
||||||
import math
|
|
||||||
|
|
||||||
from sqlalchemy.orm import Session
|
from sqlalchemy.orm import Session
|
||||||
|
|
||||||
from . import models
|
from . import models, vectors
|
||||||
from .context import history, story_actions, truncate_to_last_tokens
|
from .context import history, story_actions, truncate_to_last_tokens
|
||||||
from .database import SessionLocal
|
from .database import SessionLocal
|
||||||
from .providers import OpenAICompatibleProvider, ProviderError
|
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_INTERVAL = 6 # actions per memory
|
||||||
MEMORY_START = 12 # first memory once the adventure reaches this many actions
|
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:
|
def set_vector(memory: models.Memory, vector: list[float] | None) -> None:
|
||||||
# Different lengths means the embedding model changed since this vector was
|
"""Store (or clear) a memory's embedding.
|
||||||
# stored; zip() would silently score garbage.
|
|
||||||
if len(a) != len(b):
|
Both columns, always together: `embedding_blob` is what will be read, and
|
||||||
return 0.0
|
the JSON `embedding` stays correct behind it until the follow-up migration
|
||||||
dot = sum(x * y for x, y in zip(a, b))
|
drops it. Going through one function is what keeps them from drifting.
|
||||||
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
|
memory.embedding = vector
|
||||||
|
memory.embedding_blob = None if vector is None else vectors.pack(vector)
|
||||||
|
|
||||||
|
|
||||||
def settled_count(adventure: models.Adventure) -> int:
|
def settled_count(adventure: models.Adventure) -> int:
|
||||||
@@ -395,11 +396,11 @@ async def _embed_pending(
|
|||||||
if not pending:
|
if not pending:
|
||||||
return
|
return
|
||||||
try:
|
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:
|
except ProviderError:
|
||||||
return
|
return
|
||||||
for memory, vector in zip(pending, vectors):
|
for memory, vector in zip(pending, new):
|
||||||
memory.embedding = vector
|
set_vector(memory, vector)
|
||||||
db.commit()
|
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.
|
migrations added from Phase 9 on must run on both dialects.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
import json
|
||||||
|
|
||||||
from sqlalchemy import inspect, text
|
from sqlalchemy import inspect, text
|
||||||
from sqlalchemy.engine import Engine
|
from sqlalchemy.engine import Engine
|
||||||
|
|
||||||
|
from . import vectors
|
||||||
from .database import Base
|
from .database import Base
|
||||||
|
|
||||||
# (version, SQL to run when upgrading past it) — append only, never reorder.
|
# (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
|
# Phase 6: auto-summarization + memory bank (the `memories` table itself is
|
||||||
# created by create_all, which runs for existing DBs too).
|
# created by create_all, which runs for existing DBs too).
|
||||||
(2, "ALTER TABLE adventures ADD COLUMN auto_summarize BOOLEAN NOT NULL DEFAULT 0"),
|
(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
|
# ~5 KB to every later load of that adventure. Backfilled by
|
||||||
# _backfill_variant_count.
|
# _backfill_variant_count.
|
||||||
(37, "ALTER TABLE actions ADD COLUMN variant_count INTEGER NOT NULL DEFAULT 0"),
|
(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)
|
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.
|
# Migrations that need a data pass after their DDL, keyed by version.
|
||||||
WORLD_DELTA_VERSION = 36
|
WORLD_DELTA_VERSION = 36
|
||||||
VARIANT_COUNT_VERSION = 37
|
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:
|
def _backfill_world_delta(conn) -> None:
|
||||||
@@ -183,6 +201,46 @@ def _backfill_variant_count(conn) -> None:
|
|||||||
conn.execute(text(sql))
|
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:
|
def _get_version(conn) -> int:
|
||||||
if conn.dialect.name == "sqlite":
|
if conn.dialect.name == "sqlite":
|
||||||
return conn.execute(text("PRAGMA user_version")).scalar() or 1
|
return conn.execute(text("PRAGMA user_version")).scalar() or 1
|
||||||
@@ -220,11 +278,13 @@ def bootstrap(engine: Engine) -> None:
|
|||||||
current = _get_version(conn)
|
current = _get_version(conn)
|
||||||
for version, sql in MIGRATIONS:
|
for version, sql in MIGRATIONS:
|
||||||
if version > current:
|
if version > current:
|
||||||
conn.execute(text(sql))
|
conn.execute(text(_for_dialect(sql, conn.dialect.name)))
|
||||||
if version == WORLD_DELTA_VERSION:
|
if version == WORLD_DELTA_VERSION:
|
||||||
_backfill_world_delta(conn)
|
_backfill_world_delta(conn)
|
||||||
if version == VARIANT_COUNT_VERSION:
|
if version == VARIANT_COUNT_VERSION:
|
||||||
_backfill_variant_count(conn)
|
_backfill_variant_count(conn)
|
||||||
|
if version == EMBEDDING_BLOB_VERSION:
|
||||||
|
_backfill_embedding_blob(conn)
|
||||||
current = version
|
current = version
|
||||||
_set_version(conn, current)
|
_set_version(conn, current)
|
||||||
_encrypt_plaintext_api_keys(conn)
|
_encrypt_plaintext_api_keys(conn)
|
||||||
|
|||||||
+21
-4
@@ -1,7 +1,8 @@
|
|||||||
from datetime import datetime, timezone
|
from datetime import datetime, timezone
|
||||||
|
|
||||||
from sqlalchemy import (
|
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
|
from sqlalchemy.orm import Mapped, mapped_column, relationship
|
||||||
|
|
||||||
@@ -141,9 +142,15 @@ class Adventure(Base):
|
|||||||
class Memory(Base):
|
class Memory(Base):
|
||||||
"""Phase 6: an auto-summarized (or hand-written) fact about the adventure.
|
"""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
|
The vector lives in `embedding_blob` as packed float32 (see vectors.py).
|
||||||
Python — fine at bank sizes of a few hundred). NULL until embedded, which
|
NULL until embedded, which also marks it for backfill when an embedding
|
||||||
also marks it for backfill when an embedding model becomes available.
|
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"
|
__tablename__ = "memories"
|
||||||
@@ -151,7 +158,17 @@ class Memory(Base):
|
|||||||
id: Mapped[int] = mapped_column(primary_key=True)
|
id: Mapped[int] = mapped_column(primary_key=True)
|
||||||
adventure_id: Mapped[int] = mapped_column(ForeignKey("adventures.id", ondelete="CASCADE"))
|
adventure_id: Mapped[int] = mapped_column(ForeignKey("adventures.id", ondelete="CASCADE"))
|
||||||
text: Mapped[str] = mapped_column(Text, default="")
|
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)
|
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).
|
# Action index range this memory summarizes (null for manual memories).
|
||||||
source_start: Mapped[int | None] = mapped_column(Integer, nullable=True)
|
source_start: Mapped[int | None] = mapped_column(Integer, nullable=True)
|
||||||
source_end: 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")
|
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}
|
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:
|
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():
|
for field, value in fields.items():
|
||||||
setattr(memory, field, value)
|
setattr(memory, field, value)
|
||||||
db.commit()
|
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):
|
for i in range(args.memories):
|
||||||
db.add(models.Memory(
|
memory = models.Memory(
|
||||||
adventure_id=adventure.id,
|
adventure_id=adventure.id,
|
||||||
text=f"{MEMORY_TEXT} ({i})",
|
text=f"{MEMORY_TEXT} ({i})",
|
||||||
embedding=[rng.uniform(-1.0, 1.0) for _ in range(EMBEDDING_DIMS)],
|
|
||||||
source_start=i * memorybank.MEMORY_INTERVAL,
|
source_start=i * memorybank.MEMORY_INTERVAL,
|
||||||
source_end=i * memorybank.MEMORY_INTERVAL + memorybank.MEMORY_INTERVAL - 1,
|
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()
|
db.commit()
|
||||||
return adventure.id, user.id
|
return adventure.id, user.id
|
||||||
|
|||||||
Reference in New Issue
Block a user