Rank the memory bank without reading the memory bank

Retrieval walked adventure.memories, so every turn loaded every row of the
bank with its vector attached -- 3.1 MB, 96% of everything a turn read. It
now asks SQL which memories are in play (an id and a flag per row), ranks
against vectors held in process, and fetches text only for the five it picks.

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

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

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

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

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

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

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_015CYEJKobJ2Re4Dv7qUoSA7
This commit is contained in:
parththakkar106
2026-08-16 21:28:00 +05:30
co-authored by Claude Opus 5
parent c56864877a
commit b7e53ae581
9 changed files with 690 additions and 60 deletions
+157 -26
View File
@@ -20,8 +20,11 @@ on a later turn because the cursors only advance on success.
""" """
import asyncio 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 . import models, vectors
from .context import history, story_actions, truncate_to_last_tokens 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: def set_vector(memory: models.Memory, vector: list[float] | None) -> None:
"""Store (or clear) a memory's embedding. """Store (or clear) a memory's embedding.
Both columns, always together: `embedding_blob` is what will be read, and Every column that describes the vector moves together: `embedding_blob` is
the JSON `embedding` stays correct behind it until the follow-up migration what the ranking reads, `embedded` is the flag everything else reads, and
drops it. Going through one function is what keeps them from drifting. 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 = vector
memory.embedding_blob = None if vector is None else vectors.pack(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: def settled_count(adventure: models.Adventure) -> int:
@@ -205,9 +269,22 @@ async def retrieve_memories(
return None return None
if not settings.embedding_model.strip(): if not settings.embedding_model.strip():
return {"used": [], "error": "No embedding model configured in Settings."} 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] # Which memories are in play, and nothing else about them. This used to
if not candidates: # 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} return {"used": [], "error": None}
recent = history.tail(adventure, RETRIEVAL_WINDOW_ACTIONS, exclude_action_id) recent = history.tail(adventure, RETRIEVAL_WINDOW_ACTIONS, exclude_action_id)
@@ -222,29 +299,51 @@ async def retrieve_memories(
except ProviderError as exc: except ProviderError as exc:
return {"used": [], "error": str(exc)} return {"used": [], "error": str(exc)}
held = _vectors_for(db, adventure.id, [memory_id for memory_id, _ in catalogue])
scored = sorted( 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, reverse=True,
) )
# Pinned memories are always used and count toward top_k, so the injected # Pinned memories are always used and count toward top_k, so the injected
# set never exceeds the configured budget (unless pinned alone exceed it). # set never exceeds the configured budget (unless pinned alone exceed it).
top_k = max(1, settings.memory_top_k) 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)) remaining = max(0, top_k - len(used))
used += [(score, m) for score, m in scored if not m.pinned][:remaining] used += [row for row in scored if not row[2]][:remaining]
used.sort(key=lambda pair: pair[0], reverse=True) 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: if update_stats:
now = models.utcnow() # synchronize_session=False: nothing in this request reads the counters
for _, m in used: # back, and matching the UPDATE against loaded objects would mean having
m.use_count += 1 # loaded them, which is the cost this whole path exists to avoid.
m.last_used_at = now 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 { return {
"used": [ "used": [
{"id": m.id, "text": m.text, "similarity": round(score, 4), "pinned": m.pinned} {"id": memory_id, "text": texts.get(memory_id, ""),
for score, m in used "similarity": round(score, 4), "pinned": pinned}
for score, memory_id, pinned in used
], ],
"error": None, "error": None,
} }
@@ -391,8 +490,19 @@ async def _update_story_summary(
async def _embed_pending( async def _embed_pending(
adventure: models.Adventure, settings: models.Settings, db: Session adventure: models.Adventure, settings: models.Settings, db: Session
) -> None: ) -> None:
pending = [m for m in adventure.memories if m.embedding is None and not m.forgotten] # A query, not a walk of adventure.memories: this ran every turn and pulled
pending = pending[:MAX_EMBED_BATCH] # 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: if not pending:
return return
try: try:
@@ -407,14 +517,35 @@ async def _embed_pending(
def _evict_over_capacity( def _evict_over_capacity(
adventure: models.Adventure, settings: models.Settings, db: Session adventure: models.Adventure, settings: models.Settings, db: Session
) -> None: ) -> None:
active = [m for m in adventure.memories if not m.forgotten] # Counting and ranking are both things the database does without sending
overflow = len(active) - max(1, settings.memory_bank_capacity) # 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: if overflow <= 0:
return return
evictable = sorted( doomed = db.execute(
(m for m in active if not m.pinned), select(models.Memory.id)
key=lambda m: (m.use_count, m.last_used_at or m.created_at), .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() 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"])
+10
View File
@@ -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. # in place and keeps being written until a follow-up migration drops it.
(38, {"sqlite": "ALTER TABLE memories ADD COLUMN embedding_blob BLOB", (38, {"sqlite": "ALTER TABLE memories ADD COLUMN embedding_blob BLOB",
"default": "ALTER TABLE memories ADD COLUMN embedding_blob BYTEA"}), "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) LATEST_VERSION = max((v for v, _ in MIGRATIONS), default=1)
+14 -9
View File
@@ -158,10 +158,10 @@ 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 # Superseded by embedding_blob and still written alongside it, so a
# step until a follow-up migration drops this one. Still the column the # rollback finds the vectors, until a follow-up migration drops it. Nothing
# ranking path reads; that moves next. # reads it.
embedding: Mapped[list | None] = mapped_column(JSON, nullable=True) embedding: Mapped[list | None] = mapped_column(JSON, nullable=True, deferred=True)
# The vector, little-endian float32. Deferred because it is wider than the # 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 # rest of the row put together and exactly one code path wants it: anything
# bulk-loading memories (the Memories drawer, eviction, the embed queue) # 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). # 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)
# 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) pinned: Mapped[bool] = mapped_column(Boolean, default=False)
forgotten: Mapped[bool] = mapped_column(Boolean, default=False) # evicted, kept for UI forgotten: Mapped[bool] = mapped_column(Boolean, default=False) # evicted, kept for UI
use_count: Mapped[int] = mapped_column(Integer, default=0) use_count: Mapped[int] = mapped_column(Integer, default=0)
@@ -180,10 +185,6 @@ class Memory(Base):
adventure: Mapped[Adventure] = relationship(back_populates="memories") adventure: Mapped[Adventure] = relationship(back_populates="memories")
@property
def embedded(self) -> bool:
return self.embedding is not None
class StoryCard(Base): class StoryCard(Base):
"""Owned by either a scenario or an adventure (exactly one set).""" """Owned by either a scenario or an adventure (exactly one set)."""
@@ -378,7 +379,11 @@ class Settings(Base):
# Phase 6: auto-summarization + memory bank # Phase 6: auto-summarization + memory bank
summary_model: Mapped[str] = mapped_column(String(200), default="") # "" = main model summary_model: Mapped[str] = mapped_column(String(200), default="") # "" = main model
embedding_model: Mapped[str] = mapped_column(String(200), default="") # "" = bank disabled 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) memory_top_k: Mapped[int] = mapped_column(Integer, default=5)
@property @property
+3
View File
@@ -427,6 +427,9 @@ def delete_adventure(
adventure = get_adventure_or_404(adventure_id, db, user) adventure = get_adventure_or_404(adventure_id, db, user)
db.delete(adventure) db.delete(adventure)
db.commit() 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 ---------- # ---------- Turn engine ----------
+16 -4
View File
@@ -18,16 +18,28 @@ that the packing and the in-process cache together already make cheap.
import math import math
import struct import struct
import sys
from array import array
def pack(vector: list[float]) -> bytes: def pack(vector) -> bytes:
"""A vector as little-endian float32.""" """A vector as little-endian float32."""
return struct.pack(f"<{len(vector)}f", *vector) return struct.pack(f"<{len(vector)}f", *vector)
def unpack(blob: bytes) -> list[float]: def unpack(blob: bytes) -> array:
"""The inverse of `pack`. Length is implied: four bytes per component.""" """The inverse of `pack`. Length is implied: four bytes per component.
return list(struct.unpack(f"<{len(blob) // 4}f", blob))
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: def cosine(a: list[float], b: list[float]) -> float:
+39 -14
View File
@@ -66,7 +66,7 @@ def test_pack_round_trips_exactly():
that into JSON, so packing back to float32 must be lossless.""" that into JSON, so packing back to float32 must be lossless."""
rng = random.Random(1) rng = random.Random(1)
vector = sample_vector(rng) 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(): 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(): def test_pack_handles_the_extremes():
vector = [float32(v) for v in (0.0, -0.0, 1.0, -1.0, 3.4028234663852886e38, 1e-38)] 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(): 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() db.expire_all()
assert memory.embedding == vector 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): 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 is None
assert memory.embedding_blob is None assert memory.embedding_blob is None
assert memory.embedded is False
# --------------------------------------------------------------- the backfill # --------------------------------------------------------------- the backfill
@@ -147,7 +160,7 @@ def seed_json_only(db, adventure, count: int, dims: int = 64) -> dict[int, list[
db.flush() db.flush()
expected[memory.id] = vector expected[memory.id] = vector
db.commit() db.commit()
db.execute(text("UPDATE memories SET embedding_blob = NULL")) db.execute(text("UPDATE memories SET embedding_blob = NULL, embedded = false"))
db.commit() db.commit()
return expected return expected
@@ -160,7 +173,7 @@ def test_backfill_converts_every_existing_vector(db, adventure):
db.expire_all() db.expire_all()
for memory in db.query(models.Memory).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): 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() memories = db.query(models.Memory).all()
assert len(memories) == count assert len(memories) == count
assert all(m.embedding_blob is not None for m in memories) 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): 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() db.expire_all()
assert db.get(models.Memory, broken.id).embedding_blob is None assert db.get(models.Memory, broken.id).embedding_blob is None
for memory_id, vector in expected.items(): 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 # ------------------------------------------------------- the upgrade in full
def test_bootstrap_adds_the_column_and_backfills_it(db, adventure): def test_bootstrap_adds_the_columns_and_backfills_them(db, adventure):
"""The path a deployed database actually takes: sitting at 37 without the """The path a deployed database actually takes: sitting at 37 with neither
column, then started on this build.""" new column, then started on this build."""
expected = seed_json_only(db, adventure, count=4) 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() db.close()
with engine.begin() as conn: with engine.begin() as conn:
conn.execute(text("ALTER TABLE memories DROP COLUMN embedding_blob")) 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}")) conn.execute(text(f"PRAGMA user_version = {migrations.EMBEDDING_BLOB_VERSION - 1}"))
migrations.bootstrap(engine) migrations.bootstrap(engine)
with engine.begin() as conn: with engine.begin() as conn:
assert conn.execute(text("PRAGMA user_version")).scalar() == migrations.LATEST_VERSION assert conn.execute(text("PRAGMA user_version")).scalar() == migrations.LATEST_VERSION
rows = conn.execute(text("SELECT id, embedding_blob FROM memories")).all() rows = conn.execute(text("SELECT id, embedding_blob, embedded FROM memories")).all()
assert len(rows) == len(expected) by_id = {row[0]: (row[1], row[2]) for row in rows}
for row_id, blob in rows: assert len(by_id) == len(expected) + 1
assert vectors.unpack(blob) == expected[row_id] 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(): def test_migration_38_is_spelled_for_both_dialects():
+402
View File
@@ -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))
+12 -1
View File
@@ -246,6 +246,13 @@ def shape_insights(client, meter, adv_id):
_check(r) _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): def shape_post_turn(client, meter, adv_id):
"""Summarization, embedding and eviction, after the turn is saved.""" """Summarization, embedding and eviction, after the turn is saved."""
with meter.scope("run_post_turn (background)"): with meter.scope("run_post_turn (background)"):
@@ -257,6 +264,7 @@ SHAPES = {
"load": shape_load, "load": shape_load,
"turn": shape_turn, "turn": shape_turn,
"insights": shape_insights, "insights": shape_insights,
"memories": shape_memories,
"post_turn": shape_post_turn, "post_turn": shape_post_turn,
} }
@@ -277,8 +285,11 @@ def parse_args(argv=None):
help="story actions in the fixture (default: 200)") help="story actions in the fixture (default: 200)")
p.add_argument("--memories", type=int, default=100, p.add_argument("--memories", type=int, default=100,
help="memories, all embedded (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, 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, p.add_argument("--narration-bytes", type=int, default=4000,
help="length of an AI action's text; production averages " help="length of an AI action's text; production averages "
"~2.1 KB per action alternating with player input " "~2.1 KB per action alternating with player input "
+37 -6
View File
@@ -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 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. hid this finding. Do this first so every item below is measured, not assumed.
**Done 2026-08-16** — see the baseline below. **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 (`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 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. 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 `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 `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). 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 4. **Vector cache.** **Done.** Keyed by `adventure_id`, bounded to 8 adventures. It
delete. Must survive the retry/undo paths that prune memories. turned out to need no invalidation callbacks at all — see below.
5. **Capacity default 200 → 80.** Existing adventures inherit it, so adventure 25 evicts 5. **Capacity default 200 → 80.** **Done** (migration 41, only rows still on the old
on its next turn — check that the eviction path is sane at scale before shipping. 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 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 finished adventure — the last open item from round two. Load the newest turns, fetch
older ones as the reader scrolls up. 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, `_embed_pending` wants rows where `embedding IS NULL` — neither needs a single vector,
and both are pure SQL. 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 ## Guardrails to add with this work
- **Query-count / byte assertions per endpoint**, extending the `test_egress.py` idea: - **Query-count / byte assertions per endpoint**, extending the `test_egress.py` idea: