diff --git a/backend/app/memorybank.py b/backend/app/memorybank.py index 2747b74..dba48b3 100644 --- a/backend/app/memorybank.py +++ b/backend/app/memorybank.py @@ -20,8 +20,11 @@ on a later turn because the cursors only advance on success. """ import asyncio +from array import array +from collections import OrderedDict -from sqlalchemy.orm import Session +from sqlalchemy import func, select, update +from sqlalchemy.orm import Session, object_session from . import models, vectors from .context import history, story_actions, truncate_to_last_tokens @@ -82,12 +85,73 @@ def embedding_provider(settings: models.Settings) -> OpenAICompatibleProvider: 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. + Every column that describes the vector moves together: `embedding_blob` is + what the ranking reads, `embedded` is the flag everything else reads, and + the JSON `embedding` stays correct behind both until the follow-up + migration drops it. Going through one function is what keeps them in step — + and it is also the only place a stored vector can change, which is what + makes the cache below safe to invalidate here and nowhere else. """ memory.embedding = vector memory.embedding_blob = None if vector is None else vectors.pack(vector) + memory.embedded = vector is not None + cached = _vector_cache.get(memory.adventure_id) + if cached is not None: + cached.pop(memory.id, None) + + +# ---------- The vector cache ---------- + +# adventure id -> {memory id: vector}, most-recently-used last. +# +# Turns for one adventure arrive back to back, and the bank barely changes +# between them, so re-reading every vector each turn is the same 600 KB over +# and over. Vectors are held as array("f") — 4 bytes a component, the same +# 6 KB the column holds. A list of Python floats would be eight times that. +# +# Correctness rests on two things. Anything that *changes* a vector goes +# through set_vector, which drops that one entry. Anything that *removes* a +# memory from play — eviction, deletion, pruning, an edit clearing the vector — +# takes it out of the catalogue query below, and entries missing from the +# catalogue are dropped on the next read. So nothing has to remember to call an +# invalidate, which is the failure this design is chosen to avoid. +# +# In-process, so it assumes one worker. That is what the deploy runs; a second +# worker would each keep their own copy and both would still be correct on +# eviction and deletion, but a vector rewritten by one could go stale in the +# other until that memory next leaves the catalogue. +_vector_cache: OrderedDict[int, dict[int, array]] = OrderedDict() +VECTOR_CACHE_ADVENTURES = 8 # ~600 KB each at a 100-memory bank + + +def forget_cached_vectors(adventure_id: int) -> None: + """Drop an adventure's cached vectors. Only needed when the adventure + itself goes away — everything else self-corrects (see above).""" + _vector_cache.pop(adventure_id, None) + + +def _vectors_for(db: Session, adventure_id: int, ids: list[int]) -> dict[int, array]: + """The vectors for `ids`, reading only the ones not already held.""" + cached = _vector_cache.get(adventure_id) + if cached is None: + cached = _vector_cache[adventure_id] = {} + _vector_cache.move_to_end(adventure_id) + while len(_vector_cache) > VECTOR_CACHE_ADVENTURES: + _vector_cache.popitem(last=False) + + wanted = set(ids) + for gone in set(cached) - wanted: + del cached[gone] + missing = [memory_id for memory_id in ids if memory_id not in cached] + if missing: + rows = db.execute( + select(models.Memory.id, models.Memory.embedding_blob) + .where(models.Memory.id.in_(missing)) + ).all() + for memory_id, blob in rows: + if blob: + cached[memory_id] = vectors.unpack(blob) + return cached def settled_count(adventure: models.Adventure) -> int: @@ -205,9 +269,22 @@ async def retrieve_memories( return None if not settings.embedding_model.strip(): return {"used": [], "error": "No embedding model configured in Settings."} + db = object_session(adventure) + if db is None: + return {"used": [], "error": None} - candidates = [m for m in adventure.memories if not m.forgotten and m.embedding] - if not candidates: + # Which memories are in play, and nothing else about them. This used to + # walk adventure.memories, which loaded every row of the bank *including + # its vector* — ~31 KB a memory, three megabytes a turn, 96% of everything + # a turn read. Two ids and a flag per row is about eight bytes. + catalogue = db.execute( + select(models.Memory.id, models.Memory.pinned).where( + models.Memory.adventure_id == adventure.id, + models.Memory.forgotten.is_(False), + models.Memory.embedded.is_(True), + ) + ).all() + if not catalogue: return {"used": [], "error": None} recent = history.tail(adventure, RETRIEVAL_WINDOW_ACTIONS, exclude_action_id) @@ -222,29 +299,51 @@ async def retrieve_memories( except ProviderError as exc: return {"used": [], "error": str(exc)} + held = _vectors_for(db, adventure.id, [memory_id for memory_id, _ in catalogue]) scored = sorted( - ((cosine(query_vec, m.embedding), m) for m in candidates), - key=lambda pair: pair[0], + ( + (cosine(query_vec, held[memory_id]), memory_id, pinned) + for memory_id, pinned in catalogue + if memory_id in held + ), + key=lambda row: row[0], reverse=True, ) # Pinned memories are always used and count toward top_k, so the injected # set never exceeds the configured budget (unless pinned alone exceed it). top_k = max(1, settings.memory_top_k) - used = [(score, m) for score, m in scored if m.pinned] + used = [row for row in scored if row[2]] remaining = max(0, top_k - len(used)) - used += [(score, m) for score, m in scored if not m.pinned][:remaining] - used.sort(key=lambda pair: pair[0], reverse=True) + used += [row for row in scored if not row[2]][:remaining] + used.sort(key=lambda row: row[0], reverse=True) + if not used: + return {"used": [], "error": None} + + # Only now, for at most top_k rows, is the text worth fetching. + used_ids = [memory_id for _, memory_id, _ in used] + texts = dict( + db.execute( + select(models.Memory.id, models.Memory.text) + .where(models.Memory.id.in_(used_ids)) + ).all() + ) if update_stats: - now = models.utcnow() - for _, m in used: - m.use_count += 1 - m.last_used_at = now + # synchronize_session=False: nothing in this request reads the counters + # back, and matching the UPDATE against loaded objects would mean having + # loaded them, which is the cost this whole path exists to avoid. + db.execute( + update(models.Memory) + .where(models.Memory.id.in_(used_ids)) + .values(use_count=models.Memory.use_count + 1, last_used_at=models.utcnow()) + .execution_options(synchronize_session=False) + ) return { "used": [ - {"id": m.id, "text": m.text, "similarity": round(score, 4), "pinned": m.pinned} - for score, m in used + {"id": memory_id, "text": texts.get(memory_id, ""), + "similarity": round(score, 4), "pinned": pinned} + for score, memory_id, pinned in used ], "error": None, } @@ -391,8 +490,19 @@ async def _update_story_summary( async def _embed_pending( adventure: models.Adventure, settings: models.Settings, db: Session ) -> None: - pending = [m for m in adventure.memories if m.embedding is None and not m.forgotten] - pending = pending[:MAX_EMBED_BATCH] + # A query, not a walk of adventure.memories: this ran every turn and pulled + # the whole bank's vectors to find the handful that had none. + pending = ( + db.query(models.Memory) + .filter( + models.Memory.adventure_id == adventure.id, + models.Memory.embedded.is_(False), + models.Memory.forgotten.is_(False), + ) + .order_by(models.Memory.id) + .limit(MAX_EMBED_BATCH) + .all() + ) if not pending: return try: @@ -407,14 +517,35 @@ async def _embed_pending( def _evict_over_capacity( adventure: models.Adventure, settings: models.Settings, db: Session ) -> None: - active = [m for m in adventure.memories if not m.forgotten] - overflow = len(active) - max(1, settings.memory_bank_capacity) + # Counting and ranking are both things the database does without sending + # anything back. Walking adventure.memories to count them fetched every + # vector in the bank, every turn, whether or not anything was over capacity. + in_this_bank = (models.Memory.adventure_id == adventure.id, + models.Memory.forgotten.is_(False)) + active = db.execute( + select(func.count(models.Memory.id)).where(*in_this_bank) + ).scalar() or 0 + overflow = active - max(1, settings.memory_bank_capacity) if overflow <= 0: return - evictable = sorted( - (m for m in active if not m.pinned), - key=lambda m: (m.use_count, m.last_used_at or m.created_at), + doomed = db.execute( + select(models.Memory.id) + .where(*in_this_bank, models.Memory.pinned.is_(False)) + .order_by( + models.Memory.use_count, + func.coalesce(models.Memory.last_used_at, models.Memory.created_at), + ) + .limit(overflow) + ).scalars().all() + if not doomed: + return # every active memory is pinned; capacity yields to the pins + db.execute( + update(models.Memory) + .where(models.Memory.id.in_(doomed)) + .values(forgotten=True) + .execution_options(synchronize_session=False) ) - for memory in evictable[:overflow]: - memory.forgotten = True db.commit() + # The bulk UPDATE went around any loaded objects, so anything still holding + # the collection would see the evicted memories as active. + db.expire(adventure, ["memories"]) diff --git a/backend/app/migrations.py b/backend/app/migrations.py index e28811f..1339317 100644 --- a/backend/app/migrations.py +++ b/backend/app/migrations.py @@ -134,6 +134,16 @@ MIGRATIONS: list[tuple[int, str | dict[str, str]]] = [ # 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"}), + # ...and the one-bit answer beside it, so the Memories drawer and the embed + # queue can ask "has this got a vector?" without fetching one. Same shape as + # actions.variant_count beside actions.variants. TRUE/FALSE and a boolean + # DEFAULT are spelled the same on both dialects; 0/1 would not be. + (39, "ALTER TABLE memories ADD COLUMN embedded BOOLEAN NOT NULL DEFAULT false"), + (40, "UPDATE memories SET embedded = true WHERE embedding_blob IS NOT NULL"), + # Memory bank capacity 200 -> 80. Only rows still on the old default move, + # so anyone who picked a value keeps it — same rule as migration 29. + # Adventures already over 80 evict down on their next turn. + (41, "UPDATE settings SET memory_bank_capacity = 80 WHERE memory_bank_capacity = 200"), ] LATEST_VERSION = max((v for v, _ in MIGRATIONS), default=1) diff --git a/backend/app/models.py b/backend/app/models.py index 9ecf445..7979e24 100644 --- a/backend/app/models.py +++ b/backend/app/models.py @@ -158,10 +158,10 @@ 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) + # Superseded by embedding_blob and still written alongside it, so a + # rollback finds the vectors, until a follow-up migration drops it. Nothing + # reads it. + embedding: Mapped[list | None] = mapped_column(JSON, nullable=True, deferred=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) @@ -172,6 +172,11 @@ class Memory(Base): # 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) + # Whether embedding_blob is set. Maintained on write by memorybank + # .set_vector, for the same reason actions.variant_count exists beside + # actions.variants: every reader wants the one-bit answer and none of them + # should have to fetch six kilobytes of vector to get it. + embedded: Mapped[bool] = mapped_column(Boolean, default=False) pinned: Mapped[bool] = mapped_column(Boolean, default=False) forgotten: Mapped[bool] = mapped_column(Boolean, default=False) # evicted, kept for UI use_count: Mapped[int] = mapped_column(Integer, default=0) @@ -180,10 +185,6 @@ class Memory(Base): adventure: Mapped[Adventure] = relationship(back_populates="memories") - @property - def embedded(self) -> bool: - return self.embedding is not None - class StoryCard(Base): """Owned by either a scenario or an adventure (exactly one set).""" @@ -378,7 +379,11 @@ class Settings(Base): # Phase 6: auto-summarization + memory bank summary_model: Mapped[str] = mapped_column(String(200), default="") # "" = main model embedding_model: Mapped[str] = mapped_column(String(200), default="") # "" = bank disabled - memory_bank_capacity: Mapped[int] = mapped_column(Integer, default=200) + # Was 200. Lowered on retrieval-quality grounds first: ranking two hundred + # memories to pick five means the five are chosen out of a lot of noise, + # and older memories describe a story the player has moved on from. That it + # also cuts what the bank costs to read is the smaller reason. + memory_bank_capacity: Mapped[int] = mapped_column(Integer, default=80) memory_top_k: Mapped[int] = mapped_column(Integer, default=5) @property diff --git a/backend/app/routers/adventures.py b/backend/app/routers/adventures.py index 3b4bca6..e188f37 100644 --- a/backend/app/routers/adventures.py +++ b/backend/app/routers/adventures.py @@ -427,6 +427,9 @@ def delete_adventure( adventure = get_adventure_or_404(adventure_id, db, user) db.delete(adventure) db.commit() + # Nothing else would ever ask for this adventure's vectors again, so the + # cache would hold them until the process restarted. + memorybank.forget_cached_vectors(adventure_id) # ---------- Turn engine ---------- diff --git a/backend/app/vectors.py b/backend/app/vectors.py index 78b795d..8cdc6cd 100644 --- a/backend/app/vectors.py +++ b/backend/app/vectors.py @@ -18,16 +18,28 @@ that the packing and the in-process cache together already make cheap. import math import struct +import sys +from array import array -def pack(vector: list[float]) -> bytes: +def pack(vector) -> 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 unpack(blob: bytes) -> array: + """The inverse of `pack`. Length is implied: four bytes per component. + + Returns an `array("f")` rather than a list, because these are held in + memory between turns: the array is the same 4 bytes a component the column + is, where a list of Python floats is eight times that. It indexes, zips and + lens like a list, which is all the ranking needs. + """ + vector = array("f") + vector.frombytes(blob) + if sys.byteorder != "little": + vector.byteswap() + return vector def cosine(a: list[float], b: list[float]) -> float: diff --git a/backend/tests/test_embedding_blob.py b/backend/tests/test_embedding_blob.py index 196af7f..0ab4c52 100644 --- a/backend/tests/test_embedding_blob.py +++ b/backend/tests/test_embedding_blob.py @@ -66,7 +66,7 @@ def test_pack_round_trips_exactly(): 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 + assert list(vectors.unpack(vectors.pack(vector))) == vector def test_packed_vector_is_four_bytes_per_dimension(): @@ -79,7 +79,18 @@ def test_packed_vector_is_four_bytes_per_dimension(): 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 + 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(): @@ -113,7 +124,8 @@ def test_set_vector_writes_both_columns(db, adventure): db.expire_all() assert memory.embedding == vector - assert vectors.unpack(memory.embedding_blob) == vector + assert list(vectors.unpack(memory.embedding_blob)) == vector + assert memory.embedded is True def test_set_vector_none_clears_both(db, adventure): @@ -130,6 +142,7 @@ def test_set_vector_none_clears_both(db, adventure): assert memory.embedding is None assert memory.embedding_blob is None + assert memory.embedded is False # --------------------------------------------------------------- the backfill @@ -147,7 +160,7 @@ def seed_json_only(db, adventure, count: int, dims: int = 64) -> dict[int, list[ db.flush() expected[memory.id] = vector db.commit() - db.execute(text("UPDATE memories SET embedding_blob = NULL")) + db.execute(text("UPDATE memories SET embedding_blob = NULL, embedded = false")) db.commit() return expected @@ -160,7 +173,7 @@ def test_backfill_converts_every_existing_vector(db, adventure): db.expire_all() for memory in db.query(models.Memory).all(): - assert vectors.unpack(memory.embedding_blob) == expected[memory.id] + assert list(vectors.unpack(memory.embedding_blob)) == expected[memory.id] def test_backfill_reaches_past_one_batch(db, adventure): @@ -176,7 +189,7 @@ def test_backfill_reaches_past_one_batch(db, adventure): 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) + assert all(list(vectors.unpack(m.embedding_blob)) == expected[m.id] for m in memories) def test_backfill_leaves_unembedded_memories_alone(db, adventure): @@ -225,29 +238,41 @@ def test_backfill_skips_a_malformed_row_without_stopping(db, adventure): 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 + blob = db.get(models.Memory, memory_id).embedding_blob + assert list(vectors.unpack(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.""" +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 FROM memories")).all() - assert len(rows) == len(expected) - for row_id, blob in rows: - assert vectors.unpack(blob) == expected[row_id] + 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(): diff --git a/backend/tests/test_memory_retrieval.py b/backend/tests/test_memory_retrieval.py new file mode 100644 index 0000000..e24046f --- /dev/null +++ b/backend/tests/test_memory_retrieval.py @@ -0,0 +1,402 @@ +"""Ranking the memory bank without reading the memory bank. + +Retrieval used to walk `adventure.memories`, which loaded every row *with its +vector* — 96% of everything a turn read. It now asks SQL which memories are in +play, holds their vectors in process, and fetches text for the five it picks. + +Three things have to stay true for that to be safe, and each is a separate +failure that no error message would ever report: + +* the ranking picks the same memories it always did; +* nothing bulk-reads a vector column again; +* a cached vector is never served after the stored one changed. + + python -m pytest tests/test_memory_retrieval.py -v +""" +import os +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 asyncio +from datetime import timedelta + +import pytest +from sqlalchemy import event + +from app import memorybank, models +from app.database import Base, SessionLocal, engine + + +class StubEmbedder: + """Returns whatever vector the test set, and counts calls.""" + + def __init__(self, vector=(1.0, 0.0, 0.0)): + self.vector = list(vector) + self.calls = 0 + + async def embed(self, texts): + self.calls += 1 + return [list(self.vector) for _ in texts] + + +@pytest.fixture() +def db(): + Base.metadata.create_all(bind=engine) + memorybank._vector_cache.clear() + session = SessionLocal() + try: + yield session + finally: + session.close() + memorybank._vector_cache.clear() + Base.metadata.drop_all(bind=engine) + + +@pytest.fixture() +def settings(db): + user = models.User(is_guest=False, email="rank@example.com") + db.add(user) + db.flush() + row = models.Settings( + user_id=user.id, api_key="enc:dummy", model="m", + embedding_model="text-embedding-3-small", memory_top_k=2, + memory_bank_capacity=80, + ) + db.add(row) + db.commit() + return row + + +@pytest.fixture() +def adventure(db, settings): + adv = models.Adventure( + user_id=settings.user_id, title="Cave", script_state={}, + memory_bank_enabled=True, + ) + db.add(adv) + db.flush() + # Retrieval builds its query from the newest actions; with none, it returns + # before ranking anything. + for i in range(2): + db.add(models.Action( + adventure_id=adv.id, index=i, type="ai", text=f"Something happened {i}." + )) + db.commit() + return adv + + +def add_memory(db, adventure, text, vector, **kwargs): + memory = models.Memory(adventure_id=adventure.id, text=text, **kwargs) + db.add(memory) + db.flush() + if vector is not None: + memorybank.set_vector(memory, list(vector)) + db.commit() + return memory + + +@pytest.fixture() +def bank(db, adventure): + """Three orthogonal vectors, so a query vector picks one unambiguously.""" + return { + "x": add_memory(db, adventure, "about x", (1.0, 0.0, 0.0)), + "y": add_memory(db, adventure, "about y", (0.0, 1.0, 0.0)), + "z": add_memory(db, adventure, "about z", (0.0, 0.0, 1.0)), + } + + +def retrieve(adventure, settings, embedder, **kwargs): + memorybank.embedding_provider = lambda s: embedder + kwargs.setdefault("update_stats", False) + return asyncio.run(memorybank.retrieve_memories(adventure, settings, **kwargs)) + + +@pytest.fixture() +def sql_log(): + statements: list[str] = [] + + def record(conn, cursor, statement, parameters, context, executemany): + statements.append(statement) + + event.listen(engine, "before_cursor_execute", record) + try: + yield statements + finally: + event.remove(engine, "before_cursor_execute", record) + + +# ------------------------------------------------------------------- ranking + +def test_ranks_by_cosine_similarity(db, adventure, settings, bank): + result = retrieve(adventure, settings, StubEmbedder((1.0, 0.0, 0.0))) + assert [m["id"] for m in result["used"]][0] == bank["x"].id + assert result["used"][0]["similarity"] == pytest.approx(1.0) + assert result["used"][0]["text"] == "about x" + + +def test_honours_top_k(db, adventure, settings, bank): + settings.memory_top_k = 1 + db.commit() + assert len(retrieve(adventure, settings, StubEmbedder())["used"]) == 1 + + +def test_pinned_memories_are_always_used(db, adventure, settings, bank): + """A pin means "always in context", however badly it scores.""" + settings.memory_top_k = 1 + bank["z"].pinned = True + db.commit() + + used = retrieve(adventure, settings, StubEmbedder((1.0, 0.0, 0.0)))["used"] + assert [m["id"] for m in used] == [bank["z"].id] + assert used[0]["pinned"] is True + + +def test_forgotten_and_unembedded_memories_are_not_candidates(db, adventure, settings): + live = add_memory(db, adventure, "live", (1.0, 0.0, 0.0)) + add_memory(db, adventure, "evicted", (1.0, 0.0, 0.0), forgotten=True) + add_memory(db, adventure, "no vector yet", None) + + used = retrieve(adventure, settings, StubEmbedder())["used"] + assert [m["id"] for m in used] == [live.id] + + +def test_empty_bank_returns_no_error(db, adventure, settings): + assert retrieve(adventure, settings, StubEmbedder()) == {"used": [], "error": None} + + +def test_missing_embedding_model_is_reported(db, adventure, settings, bank): + settings.embedding_model = "" + db.commit() + result = retrieve(adventure, settings, StubEmbedder()) + assert result["used"] == [] and "embedding model" in result["error"] + + +def test_update_stats_bumps_only_the_used(db, adventure, settings, bank): + settings.memory_top_k = 1 + db.commit() + retrieve(adventure, settings, StubEmbedder((1.0, 0.0, 0.0)), update_stats=True) + db.commit() + db.expire_all() + + assert db.get(models.Memory, bank["x"].id).use_count == 1 + assert db.get(models.Memory, bank["x"].id).last_used_at is not None + assert db.get(models.Memory, bank["y"].id).use_count == 0 + + +def test_dry_runs_do_not_bump_the_counters(db, adventure, settings, bank): + """Insights assembles a context without spending a turn; it must not look + like the memories were used.""" + retrieve(adventure, settings, StubEmbedder(), update_stats=False) + db.commit() + db.expire_all() + assert all(db.get(models.Memory, m.id).use_count == 0 for m in bank.values()) + + +# --------------------------------------------------------------- the egress + +def memory_selects(statements): + return [ + s for s in statements + if "FROM memories" in s and s.lstrip().upper().startswith("SELECT") + ] + + +def test_the_json_column_is_never_selected(db, adventure, settings, bank, sql_log): + """`memories.embedding` is dead weight kept only until a follow-up + migration drops it. If anything still reads it, dropping it breaks.""" + retrieve(adventure, settings, StubEmbedder()) + offenders = [s for s in memory_selects(sql_log) if "memories.embedding " in s + or s.rstrip().endswith("memories.embedding")] + assert offenders == [], f"the JSON column was read:\n{offenders[0][:300]}" + + +def test_the_catalogue_query_carries_no_vectors(db, adventure, settings, bank, sql_log): + """The query that decides *which* memories are in play must stay tiny — + this is the one that used to drag the whole bank across.""" + retrieve(adventure, settings, StubEmbedder()) + catalogue = [s for s in memory_selects(sql_log) if "memories.pinned" in s] + assert catalogue, "expected a catalogue query" + assert not any("embedding_blob" in s for s in catalogue) + + +def test_a_second_turn_reads_no_vectors_at_all(db, adventure, settings, bank, sql_log): + """The point of the cache: back-to-back turns on one adventure pay for the + vectors once.""" + retrieve(adventure, settings, StubEmbedder()) + sql_log.clear() + retrieve(adventure, settings, StubEmbedder()) + assert not any("embedding_blob" in s for s in memory_selects(sql_log)) + + +def test_only_the_new_memory_is_fetched_after_one_is_added(db, adventure, settings, bank, sql_log): + """A growing bank must not re-read the vectors it already holds.""" + retrieve(adventure, settings, StubEmbedder()) + added = add_memory(db, adventure, "about w", (0.5, 0.5, 0.0)) + + sql_log.clear() + retrieve(adventure, settings, StubEmbedder()) + vector_reads = [s for s in memory_selects(sql_log) if "embedding_blob" in s] + assert len(vector_reads) == 1 + # One placeholder means one id: the new memory, and nothing else. + assert vector_reads[0].count("?") == 1 + assert added.id in {m["id"] for m in + retrieve(adventure, settings, StubEmbedder((0.5, 0.5, 0.0)))["used"]} + + +def test_only_top_k_texts_are_fetched(db, adventure, settings, bank, sql_log): + settings.memory_top_k = 1 + db.commit() + retrieve(adventure, settings, StubEmbedder()) + text_reads = [s for s in memory_selects(sql_log) if "memories.text" in s] + assert len(text_reads) == 1 + assert text_reads[0].count("?") == 1 + + +# ---------------------------------------------------------------- staleness + +def test_a_rewritten_vector_is_not_served_from_cache(db, adventure, settings, bank): + """The cache's one genuine hazard: a memory keeps its id while its vector + changes, so an id-set check alone would go on serving the old one. Editing + a memory's text and re-embedding it does exactly that. + """ + settings.memory_top_k = 1 + db.commit() + first = retrieve(adventure, settings, StubEmbedder((0.0, 0.0, 1.0))) + assert [m["id"] for m in first["used"]] == [bank["z"].id] + + # z is re-embedded onto the x axis, all within one gap between turns. + memorybank.set_vector(bank["z"], [1.0, 0.0, 0.0]) + db.commit() + + again = retrieve(adventure, settings, StubEmbedder((0.0, 0.0, 1.0))) + assert [m["id"] for m in again["used"]] != [bank["z"].id] + + +def test_a_deleted_memory_leaves_the_cache(db, adventure, settings, bank): + retrieve(adventure, settings, StubEmbedder()) + db.delete(bank["x"]) + db.commit() + + used = retrieve(adventure, settings, StubEmbedder((1.0, 0.0, 0.0)))["used"] + assert bank["x"].id not in {m["id"] for m in used} + assert memorybank._vector_cache[adventure.id].keys() == {bank["y"].id, bank["z"].id} + + +def test_the_cache_is_bounded(db, adventure, settings, bank): + """It holds vectors indefinitely, so without a bound a long-running process + accumulates every adventure ever played.""" + retrieve(adventure, settings, StubEmbedder()) + for fake_id in range(1000, 1000 + memorybank.VECTOR_CACHE_ADVENTURES + 2): + memorybank._vectors_for(db, fake_id, []) + assert len(memorybank._vector_cache) == memorybank.VECTOR_CACHE_ADVENTURES + assert adventure.id not in memorybank._vector_cache # evicted, least recent + + +def test_forget_cached_vectors_drops_an_adventure(db, adventure, settings, bank): + retrieve(adventure, settings, StubEmbedder()) + assert adventure.id in memorybank._vector_cache + memorybank.forget_cached_vectors(adventure.id) + assert adventure.id not in memorybank._vector_cache + + +# ----------------------------------------------------------------- eviction + +def test_eviction_marks_the_least_used(db, adventure, settings): + settings.memory_bank_capacity = 2 + db.commit() + keep = add_memory(db, adventure, "used often", (1.0, 0.0, 0.0), use_count=5) + also = add_memory(db, adventure, "used sometimes", (0.0, 1.0, 0.0), use_count=3) + drop = add_memory(db, adventure, "never used", (0.0, 0.0, 1.0), use_count=0) + + memorybank._evict_over_capacity(adventure, settings, db) + db.expire_all() + + assert db.get(models.Memory, drop.id).forgotten is True + assert db.get(models.Memory, keep.id).forgotten is False + assert db.get(models.Memory, also.id).forgotten is False + + +def test_eviction_breaks_ties_on_recency(db, adventure, settings): + settings.memory_bank_capacity = 1 + db.commit() + now = models.utcnow() + recent = add_memory(db, adventure, "recent", (1.0, 0.0, 0.0), use_count=1) + stale = add_memory(db, adventure, "stale", (0.0, 1.0, 0.0), use_count=1) + recent.last_used_at = now + stale.last_used_at = now - timedelta(days=30) + db.commit() + + memorybank._evict_over_capacity(adventure, settings, db) + db.expire_all() + + assert db.get(models.Memory, stale.id).forgotten is True + assert db.get(models.Memory, recent.id).forgotten is False + + +def test_eviction_never_touches_a_pin(db, adventure, settings): + """Capacity yields to pins: if everything active is pinned there is nothing + to evict, and the bank is allowed to sit over capacity.""" + settings.memory_bank_capacity = 1 + db.commit() + pins = [add_memory(db, adventure, f"pin {i}", (1.0, 0.0, 0.0), pinned=True) + for i in range(3)] + + memorybank._evict_over_capacity(adventure, settings, db) + db.expire_all() + + assert all(db.get(models.Memory, p.id).forgotten is False for p in pins) + + +def test_eviction_reads_no_vectors(db, adventure, settings, bank, sql_log): + """It ran every turn and pulled the whole bank to count it.""" + settings.memory_bank_capacity = 1 + db.commit() + sql_log.clear() + memorybank._evict_over_capacity(adventure, settings, db) + assert not any("embedding_blob" in s for s in memory_selects(sql_log)) + + +def test_evicted_memories_drop_out_of_the_cache(db, adventure, settings, bank): + retrieve(adventure, settings, StubEmbedder()) + settings.memory_bank_capacity = 1 + db.commit() + memorybank._evict_over_capacity(adventure, settings, db) + + used = retrieve(adventure, settings, StubEmbedder())["used"] + assert len(used) == 1 + assert len(memorybank._vector_cache[adventure.id]) == 1 + + +# ------------------------------------------------------------- embed queue + +def test_embed_pending_picks_only_unembedded_memories(db, adventure, settings): + done = add_memory(db, adventure, "already done", (1.0, 0.0, 0.0)) + todo = add_memory(db, adventure, "needs a vector", None) + add_memory(db, adventure, "evicted, skip", None, forgotten=True) + + embedder = StubEmbedder((0.0, 1.0, 0.0)) + memorybank.embedding_provider = lambda s: embedder + asyncio.run(memorybank._embed_pending(adventure, settings, db)) + db.expire_all() + + assert db.get(models.Memory, todo.id).embedded is True + assert embedder.calls == 1 + # The one already embedded keeps the vector it had. + assert db.get(models.Memory, done.id).embedding_blob == memorybank.vectors.pack( + [1.0, 0.0, 0.0] + ) + + +def test_embed_pending_reads_no_vectors(db, adventure, settings, bank, sql_log): + add_memory(db, adventure, "needs a vector", None) + embedder = StubEmbedder() + memorybank.embedding_provider = lambda s: embedder + + sql_log.clear() + asyncio.run(memorybank._embed_pending(adventure, settings, db)) + assert not any("embedding_blob" in s for s in memory_selects(sql_log)) diff --git a/backend/tools/stress_session.py b/backend/tools/stress_session.py index 8418bd9..64b8faa 100644 --- a/backend/tools/stress_session.py +++ b/backend/tools/stress_session.py @@ -246,6 +246,13 @@ def shape_insights(client, meter, adv_id): _check(r) +def shape_memories(client, meter, adv_id): + """The Memories drawer — every memory, and none of their vectors.""" + with meter.scope("GET /adventures/{id}/memories (drawer)"): + r = client.get(f"/api/adventures/{adv_id}/memories") + _check(r) + + def shape_post_turn(client, meter, adv_id): """Summarization, embedding and eviction, after the turn is saved.""" with meter.scope("run_post_turn (background)"): @@ -257,6 +264,7 @@ SHAPES = { "load": shape_load, "turn": shape_turn, "insights": shape_insights, + "memories": shape_memories, "post_turn": shape_post_turn, } @@ -277,8 +285,11 @@ def parse_args(argv=None): help="story actions in the fixture (default: 200)") p.add_argument("--memories", type=int, default=100, help="memories, all embedded (default: 100)") + # Deliberately not the app's default (80): a measuring instrument should + # hold the fixture at the size asked for rather than evict it mid-run. p.add_argument("--capacity", type=int, default=200, - help="Settings.memory_bank_capacity (default: 200)") + help="Settings.memory_bank_capacity; lower it below " + "--memories to exercise eviction (default: 200)") p.add_argument("--narration-bytes", type=int, default=4000, help="length of an AI action's text; production averages " "~2.1 KB per action alternating with player input " diff --git a/plan/13-memory-embedding-cost.md b/plan/13-memory-embedding-cost.md index f297f19..4cd95a8 100644 --- a/plan/13-memory-embedding-cost.md +++ b/plan/13-memory-embedding-cost.md @@ -88,18 +88,19 @@ no denormalisation to fix. It is a *format* problem plus a *fetch-frequency* pro to running with an embedding model configured**, since that omission is precisely what hid this finding. Do this first so every item below is measured, not assumed. **Done 2026-08-16** — see the baseline below. -2. **Migration 38 — `memories.embedding_blob` (`LargeBinary`).** Backfill in Python +2. **Migration 38 — `memories.embedding_blob` (`LargeBinary`).** **Done.** Backfill in Python (`struct.pack(f"<{n}f", *vec)`); the conversion cannot be expressed in portable SQL, so unlike migration 36/37 this one does pay a one-time 4 MB read. Drop the old JSON column in a follow-up migration once verified, not in the same one. -3. **Read path.** `retrieve_memories` queries `memories` directly with +3. **Read path.** **Done.** `retrieve_memories` queries `memories` directly with `forgotten = false AND embedding_blob IS NOT NULL` in **SQL**, not Python. Unpack with `struct`/`numpy`. Same for `_evict_over_capacity` and `_embed_pending`, which walk the same relationship for a count and for the unembedded rows (see the baseline above). -4. **Vector cache.** Keyed by `adventure_id`, invalidated on memory create, evict and - delete. Must survive the retry/undo paths that prune memories. -5. **Capacity default 200 → 80.** Existing adventures inherit it, so adventure 25 evicts - on its next turn — check that the eviction path is sane at scale before shipping. +4. **Vector cache.** **Done.** Keyed by `adventure_id`, bounded to 8 adventures. It + turned out to need no invalidation callbacks at all — see below. +5. **Capacity default 200 → 80.** **Done** (migration 41, only rows still on the old + default). Eviction checked at scale first: trimming a 100-memory bank to 80 costs + 0.8 kB and eight statements, and reads no vectors at all. 6. **Infinite scroll upward in `Play.jsx`** for the remaining 423 KB page load of a finished adventure — the last open item from round two. Load the newest turns, fetch older ones as the reader scrolls up. @@ -136,6 +137,36 @@ So step 3 below is not just `retrieve_memories`: **every walk of `_embed_pending` wants rows where `embedding IS NULL` — neither needs a single vector, and both are pure SQL. +## After (2026-08-16, same harness, same fixture) + +| shape | before | after (cold) | after (warm) | +|---|---|---|---| +| `POST .../actions` (one turn) | 3,258.7 kB | 723.4 kB | **122.3 kB** | +| `run_post_turn` | 3,139.1 kB | 0.7 kB | 0.7 kB | +| `GET .../context` (Insights) | 3,223.7 kB | 117.9 kB | 117.9 kB | +| `GET .../memories` (drawer) | ~3,138 kB | 23.6 kB | 23.6 kB | + +A played turn is turn + post_turn: **6,398 kB → 123 kB** once warm, a 52x cut. The +targets above were ~700 kB cold and ~130 kB warm, so both were met. + +Steps 2–5 landed together, because they are one deployable unit: the columns are no +use unless something reads them, and deferring them breaks the old readers. Two +additions the plan did not anticipate: + +- **`memories.embedded`, a boolean beside the blob** (migration 39/40). Once the + vectors are deferred, every "does this have an embedding?" check becomes a lazy + load — an N+1 of 6 KB reads down the Memories drawer. Same shape as + `actions.variant_count` beside `actions.variants`, and the same reason. +- **The cache needs no invalidation callbacks.** Vectors only ever change through + `set_vector`, which drops the one entry; everything 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 need no hooks and cannot be + forgotten. Vectors are held as `array("f")` — 6 KB each, matching the column; + a list of Python floats would have been eight times the plan's RAM estimate. + +Remaining: **step 6, infinite scroll upward in `Play.jsx`.** The page load is +unchanged at 426.7 kB and is now the largest single read in the app. + ## Guardrails to add with this work - **Query-count / byte assertions per endpoint**, extending the `test_egress.py` idea: