"""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