"""Tests for turnstone.core.memory_relevance — scoring, formatting, context extraction."""
from typing import Any
from unittest.mock import patch
from turnstone.core import auth
from turnstone.core.memory_relevance import (
MemoryConfig,
build_memory_context,
extract_recent_context,
score_memories,
)
from turnstone.core.trajectory import turns_from_dicts
# ---------------------------------------------------------------------------
# score_memories
# ---------------------------------------------------------------------------
class TestScoreMemories:
def test_empty_memories(self):
assert score_memories([], "query") == []
def test_empty_query_returns_recent(self):
mems = [
{"name": "a", "description": "", "content": "alpha"},
{"name": "b", "description": "", "content": "beta"},
{"name": "c", "description": "", "content": "gamma"},
]
result = score_memories(mems, "", k=2)
assert len(result) == 2
assert result[0]["name"] == "a"
def test_whitespace_query_returns_recent(self):
mems = [{"name": "a", "description": "", "content": "alpha"}]
assert score_memories(mems, " ", k=5) == mems
def test_relevance_ranking(self):
mems = [
{"name": "cooking", "description": "recipes", "content": "pasta sauce tomato"},
{"name": "python", "description": "programming", "content": "python file io disk"},
{
"name": "disk_io",
"description": "file operations",
"content": "read write file disk",
},
]
result = score_memories(mems, "file disk", k=2)
names = [m["name"] for m in result]
assert "disk_io" in names
assert "python" in names
def test_k_limits_results(self):
mems = [{"name": f"m{i}", "description": "", "content": f"word{i}"} for i in range(10)]
result = score_memories(mems, "word0 word1 word2", k=2)
assert len(result) <= 2
def test_no_match_returns_empty(self):
mems = [{"name": "a", "description": "", "content": "hello world"}]
result = score_memories(mems, "zzzznotfound")
assert result == []
def test_uses_name_for_scoring(self):
mems = [
{"name": "database_config", "description": "", "content": "host=localhost"},
{"name": "unrelated", "description": "", "content": "nothing here"},
]
result = score_memories(mems, "database", k=1)
assert len(result) == 1
assert result[0]["name"] == "database_config"
def test_uses_description_for_scoring(self):
mems = [
{"name": "x", "description": "postgresql connection settings", "content": "host=db"},
{"name": "y", "description": "unrelated", "content": "nothing"},
]
result = score_memories(mems, "postgresql", k=1)
assert result[0]["name"] == "x"
class TestScoreMemoriesReranking:
"""``score_memories`` forwards a reranker into the BM25 recall pool.
The reranker is a deterministic callable over POSITIONS in the recall pool
(the matched memories, BM25-ordered); the result is the corresponding memory
dicts, best-first. The existing 7 tests above pass no reranker (default
None) and exercise the unchanged BM25-only path.
"""
_MEMS = [
{"name": "alpha", "description": "shared topic", "content": "shared topic alpha"},
{"name": "beta", "description": "shared topic", "content": "shared topic beta"},
{"name": "gamma", "description": "shared topic", "content": "shared topic gamma"},
]
def test_reranker_reorders_memories(self):
# All three match "shared topic" -> pool covers them. The reranker
# reverses the pool positions, so the returned memory order is the
# BM25 order reversed.
baseline = score_memories(self._MEMS, "shared topic", k=3)
reranked = score_memories(
self._MEMS,
"shared topic",
k=3,
reranker=lambda q, d: list(range(len(d)))[::-1],
)
assert [m["name"] for m in reranked] == [m["name"] for m in baseline][::-1]
# Still the same set of memories, just reordered.
assert {m["name"] for m in reranked} == {m["name"] for m in baseline}
def test_floor_empties_returns_nothing(self):
# FILTER MODE (rerank_filters=True): a relevance floor that rejects
# everything (reranker returns []) means "inject no memory" ->
# score_memories returns []. This is the proactive memory floor the
# threshold setting drives (an active floor -> rerank_filters=True).
result = score_memories(
self._MEMS,
"shared topic",
k=3,
reranker=lambda q, d: [],
rerank_filters=True,
)
assert result == []
def test_reorder_mode_empty_does_not_suppress(self):
# REORDER MODE (rerank_filters=False, the disabled-floor default): an
# empty reranker result means the endpoint failed, NOT "suppress all".
# Memories fall back to BM25 top-k -- never silently dropped. Guards the
# threshold<=0 -> reorder-mode wiring in the memory call site.
result = score_memories(
self._MEMS,
"shared topic",
k=3,
reranker=lambda q, d: [],
rerank_filters=False,
)
baseline = score_memories(self._MEMS, "shared topic", k=3)
assert [m["name"] for m in result] == [m["name"] for m in baseline]
assert len(result) == 3
def test_default_none_unchanged(self):
# No reranker kwarg -> identical to passing reranker=None -> BM25-only.
assert score_memories(self._MEMS, "shared topic", k=2) == score_memories(
self._MEMS, "shared topic", k=2, reranker=None
)
# ---------------------------------------------------------------------------
# build_memory_context
# ---------------------------------------------------------------------------
class TestBuildMemoryContext:
def test_empty_memories(self):
assert build_memory_context([]) == ""
def test_single_memory(self):
mems = [{"name": "test", "type": "general", "scope": "global", "content": "hello"}]
ctx = build_memory_context(mems)
assert "" in ctx
assert "" in ctx
assert 'name="test"' in ctx
assert "hello" in ctx
def test_html_escaping(self):
mems = [
{
"name": "a dict[str, str]:
return {
"name": name,
"memory_id": memory_id or f"mid_{name}",
"type": "general",
"scope": "global",
"scope_id": "",
"description": "",
"content": content or name,
"updated": "2024-01-01T00:00:00",
}
def _make_session(fetch_limit: int = 5, relevance_k: int = 3, **kwargs: object):
"""Composition tests need a real ChatSession (constructor calls
``_init_system_messages`` once, unpatched, before the test gets a chance
to install patches). ``tmp_db`` initializes the storage singleton that
constructor needs; tests then patch the visibility helpers and call
``_init_system_messages`` a second time to exercise the new logic.
"""
from tests._helpers import make_chat_session
return make_chat_session(
memory_config=MemoryConfig(fetch_limit=fetch_limit, relevance_k=relevance_k),
**kwargs,
)
def _execute_prepared_tool(session: Any, item: dict[str, Any]) -> tuple[str, str]:
item.setdefault("_principal_id", session._tool_prepare_principal_id())
return item["execute"](item)
class TestCompositionCandidateSelection:
"""Verify the query-aware candidate set in _init_system_messages."""
def test_recency_ceiling_regression(self, tmp_db):
"""Old relevant memory not in recency top-N still injected via search path."""
session = _make_session(fetch_limit=5, relevance_k=3)
session.messages = turns_from_dicts(
[{"role": "user", "content": "postgres database configuration"}]
)
old_mem = _make_mem(
"ancient_db_config",
content="postgres database configuration connection host port",
memory_id="m_old",
)
# Recency top-5 do not include old_mem
recent = [_make_mem(f"recent_{i}", memory_id=f"mr{i}") for i in range(5)]
with (
patch.object(session, "_search_visible_memories", return_value=[old_mem]),
patch.object(session, "_list_visible_memories", return_value=recent),
):
session._init_system_messages()
joined = "\n".join(m["content"] for m in session.system_messages if m["role"] == "system")
# With the fix, old_mem enters the candidate pool via search and wins BM25
assert "ancient_db_config" in joined
def test_empty_query_falls_back_to_recency(self, tmp_db):
"""No user messages → empty context → recency path, search never called."""
session = _make_session()
session.messages = [] # extract_recent_context returns ""
recency = [_make_mem("note_alpha"), _make_mem("note_beta")]
with (
patch.object(session, "_list_visible_memories", return_value=recency),
patch.object(session, "_search_visible_memories") as search_mock,
):
session._init_system_messages()
search_mock.assert_not_called()
joined = "\n".join(m["content"] for m in session.system_messages if m["role"] == "system")
assert "note_alpha" in joined
def test_sparse_match_union_fills_candidate_pool(self, tmp_db):
"""Search returning < fetch_limit results unions with recency fillers."""
session = _make_session(fetch_limit=5, relevance_k=4)
session.messages = turns_from_dicts([{"role": "user", "content": "unique_term xyzzy"}])
hit_a = _make_mem("hit_alpha", content="unique_term xyzzy alpha", memory_id="m_ha")
hit_b = _make_mem("hit_beta", content="unique_term xyzzy beta", memory_id="m_hb")
search_hits = [hit_a, hit_b] # 2 < fetch_limit=5 → triggers union
# Recency overlaps on hit_a/hit_b and adds 3 fillers
filler = [_make_mem(f"filler_{i}", memory_id=f"mf{i}") for i in range(3)]
recency = [hit_a, hit_b] + filler
with (
patch.object(session, "_search_visible_memories", return_value=search_hits),
patch.object(session, "_list_visible_memories", return_value=recency),
):
session._init_system_messages()
joined = "\n".join(m["content"] for m in session.system_messages if m["role"] == "system")
# Both hits match "unique_term xyzzy" well → appear after BM25 ranking
assert "hit_alpha" in joined
assert "hit_beta" in joined
def test_recency_preserved_when_search_returns_noise_above_relevance_k(self, tmp_db):
"""Pool guarantee: recency-50 always reaches BM25, even when search
returns enough noise hits to clear ``relevance_k``.
Closes the narrow regression vs. the original bug — without the
``fetch_limit`` threshold, a stopword-dominated cap-search that
returned >= relevance_k irrelevant hits would short-circuit and
evict the recency-only memory the bug had been surfacing.
"""
session = _make_session(fetch_limit=10, relevance_k=3)
session.messages = turns_from_dicts([{"role": "user", "content": "configure host"}])
# Search returns relevance_k=3 noise hits — enough to skip recency
# under the OLD threshold, not enough to fill fetch_limit=10.
noise = [
_make_mem(f"noise_{i}", content="generic content", memory_id=f"mn{i}") for i in range(3)
]
# The memory the user actually wants — distinctive, in recency,
# but its content doesn't share any token with the noise hits.
wanted = _make_mem(
"host_config_v2",
content="host=localhost port=5432 db=production",
memory_id="m_wanted",
)
recency = [wanted] + [_make_mem(f"recent_{i}", memory_id=f"mr{i}") for i in range(5)]
with (
patch.object(session, "_search_visible_memories", return_value=noise),
patch.object(session, "_list_visible_memories", return_value=recency),
):
session._init_system_messages()
joined = "\n".join(m["content"] for m in session.system_messages if m["role"] == "system")
# ``wanted`` reached BM25 via the union and matched "host" → injected.
assert "host_config_v2" in joined
def test_recency_tail_preserved_when_search_adds_distinct_hits(self, tmp_db):
"""SUPERSET invariant: every recency item is in the candidate pool
when search adds hits, even if the resulting union exceeds
fetch_limit. Truncating the union at fetch_limit (the prior
behavior) evicted the recency tail — which is exactly where
ancient-but-recently-touched memories live, the recall this PR
sets out to improve.
"""
session = _make_session(fetch_limit=10, relevance_k=3)
session.messages = turns_from_dicts([{"role": "user", "content": "alpha"}])
# 5 search hits, none of which appear in recency.
search_hits = [
_make_mem(f"search_{i}", content="alpha", memory_id=f"ms{i}") for i in range(5)
]
# 10 recency items; without the union uncap, the 5 oldest of these
# would be displaced by the 5 search hits.
recency = [_make_mem(f"recency_{i}", memory_id=f"mr{i}") for i in range(10)]
with (
patch.object(session, "_search_visible_memories", return_value=search_hits),
patch.object(session, "_list_visible_memories", return_value=recency),
):
candidates, source = session._select_memory_candidates("alpha")
candidate_ids = {c["memory_id"] for c in candidates}
# Pool is search_hits ∪ recency — 15 items, no truncation.
assert len(candidates) == 15
assert source == "union"
# Every recency item present (no tail eviction).
for i in range(10):
assert f"mr{i}" in candidate_ids, f"recency item {i} evicted"
# And every search hit is also in the pool.
for i in range(5):
assert f"ms{i}" in candidate_ids, f"search hit {i} missing"
def test_coord_scope_isolated_visibility(self, tmp_db):
"""Coord composition queries the coord scope alone, never the
global/workstream/user union."""
from turnstone.core.workstream import WorkstreamKind
coord = _make_session(
fetch_limit=5,
relevance_k=3,
ws_id="coord-1",
user_id="user-1",
kind=WorkstreamKind.COORDINATOR,
)
scopes = coord._visible_scopes()
# Keyed by the creator user_id (durable per-user namespace),
# not the session's ws_id.
assert scopes == [("coordinator", "user-1")]
# And: search uses those same scopes (no global/user fan-in)
coord.messages = turns_from_dicts([{"role": "user", "content": "anything"}])
with patch(
"turnstone.core.session.search_visible_structured_memories",
return_value=[],
) as search_mock:
coord._search_visible_memories("anything", limit=5)
search_mock.assert_called_once()
# Second positional arg is the scopes list
assert search_mock.call_args.args[1] == [("coordinator", "user-1")]
class TestCompositionRerankFiltersWiring:
"""The memory composition call site maps ``threshold > 0`` to ``rerank_filters``.
A disabled floor (threshold <= 0) -> reorder mode (rerank_filters=False) so an
empty/failed reranker falls back to BM25 (memories not suppressed); an active
floor (threshold > 0) -> filter mode (rerank_filters=True) so the floor may
legitimately empty the injection. Drives the real ``_init_system_messages``
call site, capturing the kwarg ``score_memories`` actually receives.
"""
def _capture_rerank_filters(self, session: object, threshold: float) -> bool:
captured: dict[str, bool] = {}
def _fake_score(*_args: object, rerank_filters: bool = False, **_kw: object):
captured["rerank_filters"] = rerank_filters
return []
mem = _make_mem("m_one", content="alpha")
with (
patch("turnstone.core.session.score_memories", _fake_score),
patch.object(session, "_bm25_rerank_threshold", return_value=threshold),
patch.object(session, "_bm25_reranker", return_value=None),
patch.object(session, "_select_memory_candidates", return_value=([mem], "list")),
):
session._init_system_messages()
assert "rerank_filters" in captured, "score_memories was not reached"
return captured["rerank_filters"]
def test_threshold_zero_uses_reorder_mode(self, tmp_db):
session = _make_session()
session.messages = turns_from_dicts([{"role": "user", "content": "alpha"}])
# threshold 0 (disabled floor) -> reorder mode -> no suppression.
assert self._capture_rerank_filters(session, 0.0) is False
def test_positive_threshold_uses_filter_mode(self, tmp_db):
session = _make_session()
session.messages = turns_from_dicts([{"role": "user", "content": "alpha"}])
# An active floor -> filter mode -> the reranker may empty the injection.
assert self._capture_rerank_filters(session, 0.5) is True
class TestMemorySearchToolExecution:
"""End-to-end test of ``memory(action='search')`` through _exec_memory.
Drives the actual tool dispatch (not just the storage facade) so the
OR-of-terms fix and the coalesced ``memory.search`` log get exercised
together.
"""
def test_search_action_returns_or_of_terms_results(self, tmp_db):
"""Multi-word query returns rows where ANY term matches — not all."""
from turnstone.core.memory import save_structured_memory
save_structured_memory(
"postgres_notes", "host=localhost port=5432", description="Postgres notes"
)
save_structured_memory("redis_notes", "host=redis port=6379", description="Redis notes")
save_structured_memory("unrelated", "completely different", description="Unrelated notes")
session = _make_session()
item = session._prepare_memory(
"call-1",
{"action": "search", "query": "postgres no_such_word_a no_such_word_b"},
)
# Sanity: prepare returned a search-ready dispatch (not an error item)
assert item.get("action") == "search"
call_id, msg = _execute_prepared_tool(session, item)
assert call_id == "call-1"
assert "postgres_notes" in msg
# Other memories don't match any query term
assert "unrelated" not in msg
def test_search_and_list_guidance_carries_the_displayed_scope(self, tmp_db, monkeypatch):
"""Follow-up guidance must not drop a project result's scope."""
from turnstone.core.memory import save_structured_memory
save_structured_memory(
"july_digest",
"project day digest",
description="July project digest",
scope="project",
scope_id="p1",
)
monkeypatch.setattr(
auth,
"resolve_project_access",
lambda *_a, **_k: auth.ProjectAccess(True, True, "P", "active"),
)
session = _make_session(user_id="u1", project_id="p1")
for args in (
{"action": "search", "query": "digest"},
{"action": "list"},
):
item = session._prepare_memory("call-1", args)
_, msg = _execute_prepared_tool(session, item)
assert "[general:project] july_digest" in msg
assert "call memory(action='get') with the displayed name and scope" in msg
class TestPerTurnSearchCache:
"""The per-turn cache spares redundant SQL across mid-turn rebuilds."""
def test_repeated_search_in_same_turn_hits_cache(self, tmp_db):
from turnstone.core.memory import save_structured_memory
save_structured_memory("hello_mem", "alpha beta gamma", description="Greeting memory")
session = _make_session()
with patch(
"turnstone.core.session.search_visible_structured_memories",
return_value=[],
) as backend_mock:
session._search_visible_memories("alpha beta", limit=5)
session._search_visible_memories("alpha beta", limit=5)
session._search_visible_memories("alpha beta", limit=5)
# 3 calls but only 1 backend hit — cache absorbed the rest
assert backend_mock.call_count == 1
def test_user_turn_invalidates_cache(self, tmp_db):
from turnstone.core.memory import save_structured_memory
save_structured_memory("hello_mem", "alpha", description="Greeting memory")
session = _make_session()
with patch(
"turnstone.core.session.search_visible_structured_memories",
return_value=[],
) as backend_mock:
session._search_visible_memories("alpha", limit=5)
session._invalidate_memory_cache() # simulates new user turn
session._search_visible_memories("alpha", limit=5)
assert backend_mock.call_count == 2