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:
co-authored by
Claude Opus 5
parent
c56864877a
commit
b7e53ae581
@@ -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():
|
||||
|
||||
@@ -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))
|
||||
Reference in New Issue
Block a user