mirror of
https://github.com/turnstonelabs/turnstone.git
synced 2026-08-16 08:52:27 -06:00
fix(memory): query-aware candidate selection + OR-of-terms search (#468)
* fix(memory): query-aware candidate selection + OR-of-terms search The system-message memory composition path used a recency-ordered candidate set (`_list_visible_memories(limit=fetch_limit)`). On deployments with more than `fetch_limit` (default 50) visible memories, BM25 only ever ranked the 50 most-recently-touched memories — a relevant memory written months ago was silently invisible regardless of how well it matched the recent context. Multi-word search at the SQL layer used AND-of-terms, killing recall on any multi-word query without an exact field overlap. ## Functional changes - `_init_system_messages` (`turnstone/core/session.py`): extract recent context first, then `_search_visible_memories(context)` to pull query-aware candidates. Search hits below `fetch_limit` union with the recency list (deduped by memory_id) so the BM25 candidate pool is always a SUPERSET of the prior recency-only pool — even on noisy queries where the cap fills with stopwords, the recency-50 the original bug surfaced still reaches BM25. Empty context skips search entirely. Candidate-selection logic extracted into `_select_memory_candidates`. - `search_structured_memories` (PostgreSQL + SQLite): per-term clauses join with OR instead of AND. A row matches if ANY term matches ANY of name/description/content. Downstream BM25 narrows back down by relevance. ## Perf hardening - Collapse the 1-3 fanned scope queries into a single SQL. New backend methods `list_visible_structured_memories` / `search_visible_structured_memories` union the visibility scopes into one WHERE OR-group, so a composition rebuild now hits the DB at most twice (search + recency) instead of up to six times. - Cap and normalize search terms. Composition can hand a multi-KB pasted message to ILIKE-based search; without a cap, every distinct token would emit one unindexable predicate per scope-fanned query. `normalize_search_terms` (`storage/_utils.py`) de-dupes case-insensitively, drops <2-char tokens, and hard-caps at 16. - Per-turn search cache. `_init_system_messages` fires from many call sites within one turn (state transitions, MCP refresh, tool results) and the recent-context query is identical across them. Session-instance cache keyed by (query, mem_type, limit) absorbs the duplicates; invalidated in `_append_user_turn` and after memory save/delete tool actions. - Stable secondary sort by `memory_id`. `updated` is second-precision and `touch_structured_memories` can land a batch on identical timestamps; without a tie-breaker SQL returns rows in implementation-defined order, BM25 input shuffles, and the LLM-side prompt cache misses across calls. All four backend ORDER BYs now break ties on `memory_id ASC`. ## Quality cleanups - Coalesce `memory.search.term_count` + `memory.search.zero_results` into a single `memory.search` log carrying both `term_count` and `result_count`. - New `memory.composition` log: source / candidates / injected. - Promote a shared `make_chat_session` factory to `tests/_helpers.py`. - Rename SQL builder local `extra` -> `scope_filters` for clarity. - Add docstrings on `search_structured_memories` so the AND->OR flip survives future readers. ## Tests Adds 20 tests across `tests/test_structured_memory.py`, `tests/test_structured_memory_storage.py`, and `tests/test_memory_relevance.py`: recency-ceiling regression, empty-query fallback, sparse-match union, recency-preserved-when- search-returns-noise (locks in the pool-superset invariant), OR-of-terms on both backends, scope filtering preserved, search-facade multi-word behavior, term-cap normalization, the new visible-scope helpers (list + search + empty-scopes guard), coord-scope composition isolation, end-to-end `memory(action='search')` tool execution, per-turn cache hit + invalidation, and stable ordering under tied `updated` timestamps. Memory test sweep: 102/102. Broader regression (session, storage, coordinator, load_skill): 411/411. * fix(memory): address Copilot review on PR #468 Three follow-ups from Copilot's inline review: 1. SUPERSET invariant violation (Copilot, session.py:5510). `(search_hits + extra)[:fetch_limit]` capped the union back down to fetch_limit, evicting the recency tail when search added distinct hits. Recency tail is exactly where ancient-but-recently-touched memories live — the recall this PR is supposed to improve — so tail eviction recreated the bug for the narrow case where a query term fell off the 16-cap and the matching memory sat in recency[40-49]. Drop the cap; both halves are already SQL-capped at fetch_limit, so the union is at most 2 × fetch_limit (~100 with defaults). BM25 over 100 candidates in pure Python is sub-ms; irrelevant recency fillers get score=0 and don't pollute ranking. Updates the docstring to actually be honest about the invariant. Adds `test_recency_tail_preserved_when_search_adds_distinct_hits` that locks the behavior in: 5 search hits + 10 recency = 15-item pool, every recency item present, source="union". 2. Unbounded `query.split()` in normalize_search_terms (Copilot, _utils.py:74). `str.split()` allocates the full token list before the cap-after-16 break, so a 100KB pasted query did MB of throwaway work even though only 16 tokens entered SQL. Switch to `re.finditer(r'\S+', query)` — streaming iterator, stops scanning at the first 16 normalized terms regardless of input size. 3. Misleading + unbounded log term_count (Copilot, session.py:8571). `len(item["query"].split())` had two problems: same unbounded split as #2, and the value reported the raw input token count rather than the normalized term count that actually hit the SQL WHERE clause — misleading metric for an operator trying to understand storage-side behavior. Switch to `len(normalize_search_terms(item["query"]))` — accurate count, and bounded for free via #2. Refuted: github-code-quality flagged `...` bodies in the new Protocol methods as "statement has no effect." False positive — `...` is the canonical Protocol body convention, used 213 other times in the same file. Memory test sweep: 103/103. Broader regression: 411/411.
This commit is contained in:
committed by
Patrick Buckley
parent
1207d27363
commit
2fd0c29a92
@@ -0,0 +1,28 @@
|
||||
"""Shared test helpers — kept out of conftest.py since these are factories,
|
||||
not fixtures, and several test files want to import them directly."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
|
||||
def make_chat_session(**overrides: Any) -> Any:
|
||||
"""Build a minimal ``ChatSession`` with sane test defaults.
|
||||
|
||||
Caller passes any constructor arg as a kwarg to override the default —
|
||||
e.g. ``make_chat_session(memory_config=MemoryConfig(fetch_limit=5))``.
|
||||
"""
|
||||
from turnstone.core.session import ChatSession
|
||||
|
||||
defaults: dict[str, Any] = {
|
||||
"client": MagicMock(),
|
||||
"model": "test-model",
|
||||
"ui": MagicMock(),
|
||||
"instructions": None,
|
||||
"temperature": 0.5,
|
||||
"max_tokens": 4096,
|
||||
"tool_timeout": 30,
|
||||
}
|
||||
defaults.update(overrides)
|
||||
return ChatSession(**defaults)
|
||||
@@ -1,6 +1,9 @@
|
||||
"""Tests for turnstone.core.memory_relevance — scoring, formatting, context extraction."""
|
||||
|
||||
from unittest.mock import patch
|
||||
|
||||
from turnstone.core.memory_relevance import (
|
||||
MemoryConfig,
|
||||
build_memory_context,
|
||||
extract_recent_context,
|
||||
score_memories,
|
||||
@@ -192,3 +195,265 @@ class TestExtractRecentContext:
|
||||
|
||||
def test_empty_messages(self):
|
||||
assert extract_recent_context([]) == ""
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Composition candidate-selection (_init_system_messages)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _make_mem(name: str, content: str = "", memory_id: str | None = None) -> dict[str, str]:
|
||||
return {
|
||||
"name": name,
|
||||
"memory_id": memory_id or f"mid_{name}",
|
||||
"type": "project",
|
||||
"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,
|
||||
)
|
||||
|
||||
|
||||
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 = [{"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 = [{"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 = [{"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 = [{"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()
|
||||
assert scopes == [("coordinator", "coord-1")]
|
||||
# And: search uses those same scopes (no global/user fan-in)
|
||||
coord.messages = [{"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", "coord-1")]
|
||||
|
||||
|
||||
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")
|
||||
save_structured_memory("redis_notes", "host=redis port=6379")
|
||||
save_structured_memory("unrelated", "completely different")
|
||||
|
||||
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 = session._exec_memory(item)
|
||||
assert call_id == "call-1"
|
||||
assert "postgres_notes" in msg
|
||||
# Other memories don't match any query term
|
||||
assert "unrelated" not 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")
|
||||
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")
|
||||
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
|
||||
|
||||
@@ -67,6 +67,42 @@ class TestSearchStructuredMemories:
|
||||
assert len(results) >= 1
|
||||
assert any(r["name"] == "db_host" for r in results)
|
||||
|
||||
def test_multiword_or_matches_partial(self, tmp_db):
|
||||
"""OR-of-terms: memory matching only 1 of 3 query terms is returned."""
|
||||
save_structured_memory("postgres_config", "host=localhost port=5432")
|
||||
save_structured_memory("redis_config", "host=redis port=6379")
|
||||
save_structured_memory("unrelated", "nothing relevant here")
|
||||
|
||||
# "postgres missing_word_a missing_word_b": only postgres_config matches "postgres"
|
||||
results = search_structured_memories("postgres missing_word_a missing_word_b")
|
||||
names = {r["name"] for r in results}
|
||||
assert "postgres_config" in names
|
||||
assert "unrelated" not in names
|
||||
|
||||
def test_multiword_or_multiple_partial_matches(self, tmp_db):
|
||||
"""Multiple memories each matching different terms are all returned."""
|
||||
save_structured_memory("key_alpha", "alpha content here")
|
||||
save_structured_memory("key_beta", "beta content here")
|
||||
save_structured_memory("key_other", "completely different")
|
||||
|
||||
results = search_structured_memories("alpha beta")
|
||||
names = {r["name"] for r in results}
|
||||
assert "key_alpha" in names
|
||||
assert "key_beta" in names
|
||||
assert "key_other" not in names
|
||||
|
||||
def test_search_scope_filtering_preserved(self, tmp_db):
|
||||
"""Search with scope filter only returns memories in that scope."""
|
||||
save_structured_memory("ws1_fact", "alpha info", scope="workstream", scope_id="ws1")
|
||||
save_structured_memory("ws2_fact", "alpha info", scope="workstream", scope_id="ws2")
|
||||
save_structured_memory("global_fact", "alpha info", scope="global")
|
||||
|
||||
results = search_structured_memories("alpha", scope="workstream", scope_id="ws1")
|
||||
names = {r["name"] for r in results}
|
||||
assert "ws1_fact" in names
|
||||
assert "ws2_fact" not in names
|
||||
assert "global_fact" not in names
|
||||
|
||||
|
||||
class TestGetStructuredMemoryByName:
|
||||
def test_get_existing(self, tmp_db):
|
||||
|
||||
@@ -126,3 +126,148 @@ class TestCount:
|
||||
backend.create_structured_memory("m2", "b", "", "project", "workstream", "ws1", "2")
|
||||
assert backend.count_structured_memories(scope="global") == 1
|
||||
assert backend.count_structured_memories(scope="workstream") == 1
|
||||
|
||||
|
||||
class TestSearchOrOfTerms:
|
||||
"""Verify that multi-word search uses OR-of-terms (any term matches → row included)."""
|
||||
|
||||
def test_single_matching_term_in_multi_word_query(self, backend):
|
||||
"""Memory with content 'apple' found when query is 'apple banana cherry'."""
|
||||
backend.create_structured_memory("m1", "apple_mem", "", "project", "global", "", "apple")
|
||||
backend.create_structured_memory("m2", "other_mem", "", "project", "global", "", "grape")
|
||||
|
||||
results = backend.search_structured_memories("apple banana cherry")
|
||||
names = {r["name"] for r in results}
|
||||
assert "apple_mem" in names # matches "apple" — OR-of-terms keeps it
|
||||
assert "other_mem" not in names # "grape" matches nothing in the query
|
||||
|
||||
def test_partial_overlap_across_memories(self, backend):
|
||||
"""Each memory matches one of three terms; all three are returned."""
|
||||
backend.create_structured_memory("m1", "alpha_doc", "", "project", "global", "", "alpha")
|
||||
backend.create_structured_memory("m2", "beta_doc", "", "project", "global", "", "beta")
|
||||
backend.create_structured_memory("m3", "gamma_doc", "", "project", "global", "", "gamma")
|
||||
backend.create_structured_memory("m4", "unrelated", "", "project", "global", "", "delta")
|
||||
|
||||
results = backend.search_structured_memories("alpha beta gamma")
|
||||
names = {r["name"] for r in results}
|
||||
assert "alpha_doc" in names
|
||||
assert "beta_doc" in names
|
||||
assert "gamma_doc" in names
|
||||
assert "unrelated" not in names # "delta" doesn't appear in the query
|
||||
|
||||
def test_scope_filter_preserved(self, backend):
|
||||
"""OR-of-terms search still respects scope / scope_id filters."""
|
||||
backend.create_structured_memory(
|
||||
"m1", "ws1_note", "", "project", "workstream", "ws1", "info"
|
||||
)
|
||||
backend.create_structured_memory(
|
||||
"m2", "ws2_note", "", "project", "workstream", "ws2", "info"
|
||||
)
|
||||
backend.create_structured_memory("m3", "global_note", "", "project", "global", "", "info")
|
||||
|
||||
results = backend.search_structured_memories("info", scope="workstream", scope_id="ws1")
|
||||
names = {r["name"] for r in results}
|
||||
assert "ws1_note" in names
|
||||
assert "ws2_note" not in names
|
||||
assert "global_note" not in names
|
||||
|
||||
def test_term_cap_normalizes_unbounded_query(self, backend):
|
||||
"""A multi-KB query collapses to <= MAX terms (de-dupe + length filter)."""
|
||||
backend.create_structured_memory("m1", "alpha_doc", "", "project", "global", "", "alpha")
|
||||
backend.create_structured_memory(
|
||||
"m2", "other_doc", "", "project", "global", "", "irrelevant"
|
||||
)
|
||||
|
||||
# Build a noisy query: same word repeated, plus 1-char tokens that
|
||||
# the normalizer drops, plus the actual signal "alpha".
|
||||
noisy = " ".join(["x"] * 100 + ["alpha"] * 50)
|
||||
results = backend.search_structured_memories(noisy)
|
||||
names = {r["name"] for r in results}
|
||||
assert "alpha_doc" in names
|
||||
|
||||
|
||||
class TestVisibleStructuredMemories:
|
||||
"""Single-query union helpers used by the composition path."""
|
||||
|
||||
def test_list_visible_unions_global_workstream_user(self, backend):
|
||||
backend.create_structured_memory("m1", "g_note", "", "project", "global", "", "g")
|
||||
backend.create_structured_memory("m2", "ws_note", "", "project", "workstream", "ws1", "w")
|
||||
backend.create_structured_memory("m3", "u_note", "", "project", "user", "u1", "u")
|
||||
backend.create_structured_memory("m4", "other_ws", "", "project", "workstream", "ws2", "x")
|
||||
|
||||
scopes = [("global", ""), ("workstream", "ws1"), ("user", "u1")]
|
||||
rows = backend.list_visible_structured_memories(scopes)
|
||||
names = {r["name"] for r in rows}
|
||||
assert names == {"g_note", "ws_note", "u_note"} # ws2 excluded
|
||||
|
||||
def test_search_visible_unions_scopes_and_terms(self, backend):
|
||||
backend.create_structured_memory("m1", "g_alpha", "", "project", "global", "", "alpha")
|
||||
backend.create_structured_memory(
|
||||
"m2", "ws_beta", "", "project", "workstream", "ws1", "beta"
|
||||
)
|
||||
backend.create_structured_memory(
|
||||
"m3", "ws_other", "", "project", "workstream", "ws2", "alpha"
|
||||
)
|
||||
|
||||
scopes = [("global", ""), ("workstream", "ws1")]
|
||||
rows = backend.search_visible_structured_memories("alpha beta", scopes)
|
||||
names = {r["name"] for r in rows}
|
||||
assert "g_alpha" in names # global, matches "alpha"
|
||||
assert "ws_beta" in names # ws1, matches "beta"
|
||||
assert "ws_other" not in names # ws2 -> outside visibility
|
||||
|
||||
def test_visible_helpers_handle_empty_scopes(self, backend):
|
||||
backend.create_structured_memory("m1", "anything", "", "project", "global", "", "x")
|
||||
assert backend.list_visible_structured_memories([]) == []
|
||||
assert backend.search_visible_structured_memories("x", []) == []
|
||||
|
||||
|
||||
class TestStableOrderingOnTimestampTies:
|
||||
"""When two memories share an `updated` timestamp, secondary sort on
|
||||
memory_id keeps the order deterministic across calls.
|
||||
|
||||
`updated` is second-precision, and touch_structured_memories() can bump
|
||||
a batch to identical timestamps — without a tie-breaker BM25 input
|
||||
order shuffles run-to-run, busting the LLM-side prompt cache.
|
||||
"""
|
||||
|
||||
def _seed_with_shared_timestamp(self, backend):
|
||||
# Create three memories then force their `updated` columns equal —
|
||||
# mirrors the real-world case where a touch_structured_memories
|
||||
# batch lands them in the same second.
|
||||
for mid in ("zebra_id", "apple_id", "mango_id"):
|
||||
backend.create_structured_memory(
|
||||
mid, f"name_{mid}", "", "project", "global", "", "shared content"
|
||||
)
|
||||
import sqlalchemy as sa
|
||||
|
||||
with backend._conn() as conn:
|
||||
conn.execute(sa.text("UPDATE structured_memories SET updated = '2024-01-01T00:00:00'"))
|
||||
conn.commit()
|
||||
|
||||
def test_list_stable_order_under_tied_updated(self, backend):
|
||||
self._seed_with_shared_timestamp(backend)
|
||||
first = [r["memory_id"] for r in backend.list_structured_memories()]
|
||||
second = [r["memory_id"] for r in backend.list_structured_memories()]
|
||||
# Deterministic across calls AND sorted by memory_id ASC for ties
|
||||
assert first == second
|
||||
assert first == ["apple_id", "mango_id", "zebra_id"]
|
||||
|
||||
def test_search_stable_order_under_tied_updated(self, backend):
|
||||
self._seed_with_shared_timestamp(backend)
|
||||
first = [r["memory_id"] for r in backend.search_structured_memories("shared")]
|
||||
second = [r["memory_id"] for r in backend.search_structured_memories("shared")]
|
||||
assert first == second
|
||||
assert first == ["apple_id", "mango_id", "zebra_id"]
|
||||
|
||||
def test_visible_search_stable_order_under_tied_updated(self, backend):
|
||||
self._seed_with_shared_timestamp(backend)
|
||||
scopes = [("global", "")]
|
||||
first = [
|
||||
r["memory_id"] for r in backend.search_visible_structured_memories("shared", scopes)
|
||||
]
|
||||
second = [
|
||||
r["memory_id"] for r in backend.search_visible_structured_memories("shared", scopes)
|
||||
]
|
||||
assert first == second
|
||||
assert first == ["apple_id", "mango_id", "zebra_id"]
|
||||
|
||||
@@ -757,6 +757,37 @@ def search_structured_memories(
|
||||
return []
|
||||
|
||||
|
||||
def list_visible_structured_memories(
|
||||
scopes: list[tuple[str, str]],
|
||||
mem_type: str = "",
|
||||
limit: int = 100,
|
||||
) -> list[dict[str, str]]:
|
||||
"""Single-query union across visible (scope, scope_id) pairs."""
|
||||
try:
|
||||
return get_storage().list_visible_structured_memories(
|
||||
scopes, mem_type=mem_type, limit=limit
|
||||
)
|
||||
except Exception:
|
||||
log.warning("Failed to list visible structured memories", exc_info=True)
|
||||
return []
|
||||
|
||||
|
||||
def search_visible_structured_memories(
|
||||
query: str,
|
||||
scopes: list[tuple[str, str]],
|
||||
mem_type: str = "",
|
||||
limit: int = 20,
|
||||
) -> list[dict[str, str]]:
|
||||
"""OR-of-terms search joined with a single visibility OR-group."""
|
||||
try:
|
||||
return get_storage().search_visible_structured_memories(
|
||||
query, scopes, mem_type=mem_type, limit=limit
|
||||
)
|
||||
except Exception:
|
||||
log.warning("Failed to search visible structured memories", exc_info=True)
|
||||
return []
|
||||
|
||||
|
||||
def touch_structured_memories(keys: list[tuple[str, str, str]]) -> int:
|
||||
"""Batch-touch memories (bump last_accessed, increment access_count).
|
||||
|
||||
|
||||
+96
-41
@@ -58,6 +58,7 @@ from turnstone.core.memory import (
|
||||
list_default_skills,
|
||||
list_skills_by_activation,
|
||||
list_structured_memories,
|
||||
list_visible_structured_memories,
|
||||
list_workstreams_with_history,
|
||||
load_messages,
|
||||
load_workstream_config,
|
||||
@@ -71,6 +72,7 @@ from turnstone.core.memory import (
|
||||
search_history,
|
||||
search_history_recent,
|
||||
search_structured_memories,
|
||||
search_visible_structured_memories,
|
||||
set_workstream_alias,
|
||||
unreserve_attachments,
|
||||
update_workstream_title,
|
||||
@@ -92,6 +94,7 @@ from turnstone.core.providers import create_provider
|
||||
from turnstone.core.safety import is_command_blocked, sanitize_command
|
||||
from turnstone.core.sandbox import execute_math_sandboxed
|
||||
from turnstone.core.storage._registry import get_storage
|
||||
from turnstone.core.storage._utils import normalize_search_terms
|
||||
from turnstone.core.tool_advisory import escape_wrapper_tags, render_system_reminder
|
||||
from turnstone.core.tool_search import ToolSearchManager
|
||||
from turnstone.core.tools import (
|
||||
@@ -428,6 +431,11 @@ class ChatSession:
|
||||
except Exception:
|
||||
log.debug("rule_registry.init_failed", exc_info=True)
|
||||
self._memory_config = memory_config or MemoryConfig()
|
||||
# Per-turn cache for _search_visible_memories — _init_system_messages
|
||||
# fires many times within one turn (state transitions, MCP refresh,
|
||||
# tool results) and the recent-context string is identical across
|
||||
# them. Invalidated on user-turn append and on memory write/delete.
|
||||
self._mem_search_cache: dict[tuple[str, str, int], list[dict[str, str]]] = {}
|
||||
self._ws_id = ws_id or uuid.uuid4().hex
|
||||
self._title_generated = False
|
||||
self._read_files: set[str] = set()
|
||||
@@ -1633,10 +1641,16 @@ class ChatSession:
|
||||
if self.instructions:
|
||||
dev_parts.append("")
|
||||
dev_parts.append(self.instructions)
|
||||
visible_mems = self._list_visible_memories(limit=self._mem_cfg.fetch_limit)
|
||||
context = extract_recent_context(self.messages)
|
||||
visible_mems, candidate_source = self._select_memory_candidates(context)
|
||||
if visible_mems:
|
||||
context = extract_recent_context(self.messages)
|
||||
relevant = score_memories(visible_mems, context, k=self._mem_cfg.relevance_k)
|
||||
log.info(
|
||||
"memory.composition",
|
||||
source=candidate_source,
|
||||
candidates=len(visible_mems),
|
||||
injected=len(relevant),
|
||||
)
|
||||
if relevant:
|
||||
dev_parts.append("")
|
||||
dev_parts.append(build_memory_context(relevant))
|
||||
@@ -2194,6 +2208,9 @@ class ChatSession:
|
||||
consume step adds it to the WHERE clause so a stale send can't
|
||||
steal rows reserved to a different one.
|
||||
"""
|
||||
# New user content invalidates the per-turn memory-search cache
|
||||
# (composition will see a different recent-context string).
|
||||
self._invalidate_memory_cache()
|
||||
user_content: str | list[dict[str, Any]]
|
||||
if attachments:
|
||||
parts: list[dict[str, Any]] = [{"type": "text", "text": user_input}]
|
||||
@@ -5413,60 +5430,90 @@ class ChatSession:
|
||||
n += count_structured_memories(scope="user", scope_id=self._user_id)
|
||||
return n
|
||||
|
||||
def _visible_scopes(self) -> list[tuple[str, str]]:
|
||||
"""Return the (scope, scope_id) pairs visible to this session.
|
||||
|
||||
Coord sessions see ONLY their coord-scope; interactive sessions see
|
||||
global + their workstream + their user (when uid present). Drives
|
||||
the single-query visibility helpers.
|
||||
"""
|
||||
if self._kind == WorkstreamKind.COORDINATOR:
|
||||
return [("coordinator", self._ws_id)]
|
||||
scopes: list[tuple[str, str]] = [("global", ""), ("workstream", self._ws_id)]
|
||||
if self._user_id:
|
||||
scopes.append(("user", self._user_id))
|
||||
return scopes
|
||||
|
||||
def _list_visible_memories(self, mem_type: str = "", limit: int = 50) -> list[dict[str, str]]:
|
||||
"""List memories visible to this session with optional type filter.
|
||||
|
||||
Single SQL round-trip — collapses the prior per-scope fan-out.
|
||||
See :meth:`_visible_memory_count` for the coord-isolation rule.
|
||||
"""
|
||||
if self._kind == WorkstreamKind.COORDINATOR:
|
||||
return list_structured_memories(
|
||||
mem_type=mem_type,
|
||||
scope="coordinator",
|
||||
scope_id=self._ws_id,
|
||||
limit=limit,
|
||||
)
|
||||
global_mems = list_structured_memories(mem_type=mem_type, scope="global", limit=limit)
|
||||
ws_mems = list_structured_memories(
|
||||
mem_type=mem_type, scope="workstream", scope_id=self._ws_id, limit=limit
|
||||
return list_visible_structured_memories(
|
||||
self._visible_scopes(), mem_type=mem_type, limit=limit
|
||||
)
|
||||
user_mems: list[dict[str, str]] = []
|
||||
if self._user_id:
|
||||
user_mems = list_structured_memories(
|
||||
mem_type=mem_type, scope="user", scope_id=self._user_id, limit=limit
|
||||
)
|
||||
combined = global_mems + ws_mems + user_mems
|
||||
combined.sort(key=lambda m: m.get("updated", ""), reverse=True)
|
||||
return combined[:limit]
|
||||
|
||||
def _search_visible_memories(
|
||||
self, query: str, mem_type: str = "", limit: int = 20
|
||||
) -> list[dict[str, str]]:
|
||||
"""Search memories visible to this session (scope-filtered).
|
||||
|
||||
Single SQL round-trip with a per-turn cache: ``_init_system_messages``
|
||||
is invoked many times within a turn (state transitions, MCP refresh,
|
||||
tool results) and the recent-context query is identical across them.
|
||||
Cache is cleared on each new user turn and after memory writes/deletes.
|
||||
See :meth:`_visible_memory_count` for the coord-isolation rule.
|
||||
"""
|
||||
if self._kind == WorkstreamKind.COORDINATOR:
|
||||
return search_structured_memories(
|
||||
query,
|
||||
mem_type=mem_type,
|
||||
scope="coordinator",
|
||||
scope_id=self._ws_id,
|
||||
limit=limit,
|
||||
)
|
||||
global_mems = search_structured_memories(
|
||||
query, mem_type=mem_type, scope="global", limit=limit
|
||||
cache_key = (query, mem_type, limit)
|
||||
cached = self._mem_search_cache.get(cache_key)
|
||||
if cached is not None:
|
||||
return cached
|
||||
rows = search_visible_structured_memories(
|
||||
query, self._visible_scopes(), mem_type=mem_type, limit=limit
|
||||
)
|
||||
ws_mems = search_structured_memories(
|
||||
query, mem_type=mem_type, scope="workstream", scope_id=self._ws_id, limit=limit
|
||||
)
|
||||
user_mems: list[dict[str, str]] = []
|
||||
if self._user_id:
|
||||
user_mems = search_structured_memories(
|
||||
query, mem_type=mem_type, scope="user", scope_id=self._user_id, limit=limit
|
||||
)
|
||||
combined = global_mems + ws_mems + user_mems
|
||||
combined.sort(key=lambda m: m.get("updated", ""), reverse=True)
|
||||
return combined[:limit]
|
||||
self._mem_search_cache[cache_key] = rows
|
||||
return rows
|
||||
|
||||
def _invalidate_memory_cache(self) -> None:
|
||||
"""Drop the per-turn search cache; call on user-turn append + memory writes."""
|
||||
self._mem_search_cache.clear()
|
||||
|
||||
def _select_memory_candidates(self, context: str) -> tuple[list[dict[str, str]], str]:
|
||||
"""Pick the candidate set fed into BM25 ranking.
|
||||
|
||||
Returns ``(memories, source_label)`` where source is one of:
|
||||
``recency`` (no context, or search returned nothing),
|
||||
``search`` (search saturated the fetch_limit budget alone), or
|
||||
``union`` (search hits ∪ recency, deduped by memory_id).
|
||||
|
||||
Invariant: the candidate pool is always a SUPERSET of the
|
||||
recency-only pool the original bug used — recency is fully
|
||||
preserved (not truncated) whenever it gets unioned. Worst
|
||||
case the union is 2 × fetch_limit candidates (~100 with
|
||||
defaults), which BM25 ranks in pure Python in well under a
|
||||
millisecond. BM25's score>0 cutoff in bm25.py drops anything
|
||||
that doesn't match the query, so unranked recency tail items
|
||||
cost nothing on irrelevant candidates while saving the
|
||||
relevant ones.
|
||||
|
||||
Capping the union at fetch_limit (the prior behavior) would
|
||||
evict the recency tail when search added distinct hits — and
|
||||
the recency tail is exactly where ancient-but-recently-touched
|
||||
memories live, which is the recall the PR sets out to improve.
|
||||
"""
|
||||
fetch_limit = self._mem_cfg.fetch_limit
|
||||
if not context:
|
||||
return self._list_visible_memories(limit=fetch_limit), "recency"
|
||||
search_hits = self._search_visible_memories(context, limit=fetch_limit)
|
||||
if len(search_hits) >= fetch_limit:
|
||||
return search_hits, "search"
|
||||
recency = self._list_visible_memories(limit=fetch_limit)
|
||||
seen = {m["memory_id"] for m in search_hits}
|
||||
extra = [m for m in recency if m["memory_id"] not in seen]
|
||||
if not search_hits:
|
||||
return extra, "recency"
|
||||
return search_hits + extra, ("union" if extra else "search")
|
||||
|
||||
def _check_metacognitive_nudge(self, user_message: str) -> tuple[str, str] | None:
|
||||
"""Check if a metacognitive nudge should fire for *user_message*.
|
||||
@@ -8451,6 +8498,7 @@ class ChatSession:
|
||||
msg = f"Error: failed to save memory '{item['name']}'"
|
||||
self._report_tool_result(call_id, "memory", msg, is_error=True)
|
||||
return call_id, msg
|
||||
self._invalidate_memory_cache()
|
||||
self._init_system_messages()
|
||||
if old is not None:
|
||||
msg = f"Updated memory '{item['name']}' (type={item['mem_type']}, scope={item['scope']})"
|
||||
@@ -8496,6 +8544,7 @@ class ChatSession:
|
||||
msg = f"Error: memory '{item['name']}' not found (searched scopes: {tried})"
|
||||
self._report_tool_result(call_id, "memory", msg, is_error=True)
|
||||
else:
|
||||
self._invalidate_memory_cache()
|
||||
self._init_system_messages()
|
||||
msg = f"Deleted memory '{item['name']}' (scope={deleted_scope})"
|
||||
self._report_tool_result(call_id, "memory", msg)
|
||||
@@ -8523,6 +8572,12 @@ class ChatSession:
|
||||
mem_type=item.get("mem_type", ""),
|
||||
limit=item["limit"],
|
||||
)
|
||||
log.info(
|
||||
"memory.search",
|
||||
term_count=len(normalize_search_terms(item["query"])),
|
||||
result_count=len(rows),
|
||||
query=item["query"][:120],
|
||||
)
|
||||
if rows:
|
||||
lines = []
|
||||
for m in rows:
|
||||
|
||||
@@ -87,6 +87,9 @@ from turnstone.core.storage._utils import (
|
||||
from turnstone.core.storage._utils import (
|
||||
VERDICT_MUTABLE as _VERDICT_MUTABLE,
|
||||
)
|
||||
from turnstone.core.storage._utils import (
|
||||
normalize_search_terms as _normalize_search_terms,
|
||||
)
|
||||
from turnstone.core.storage._utils import (
|
||||
reconstruct_messages as _reconstruct_messages,
|
||||
)
|
||||
@@ -3315,7 +3318,10 @@ class PostgreSQLBackend:
|
||||
limit: int = 100,
|
||||
) -> list[dict[str, str]]:
|
||||
with self._conn() as conn:
|
||||
q = sa.select(structured_memories).order_by(structured_memories.c.updated.desc())
|
||||
q = sa.select(structured_memories).order_by(
|
||||
structured_memories.c.updated.desc(),
|
||||
structured_memories.c.memory_id.asc(),
|
||||
)
|
||||
if mem_type:
|
||||
q = q.where(structured_memories.c.type == mem_type)
|
||||
if scope:
|
||||
@@ -3334,11 +3340,16 @@ class PostgreSQLBackend:
|
||||
scope_id: str = "",
|
||||
limit: int = 20,
|
||||
) -> list[dict[str, str]]:
|
||||
"""OR-of-terms ILIKE search; ranking is the caller's job (BM25 downstream)."""
|
||||
if not query or not query.strip():
|
||||
return self.list_structured_memories(
|
||||
mem_type=mem_type, scope=scope, scope_id=scope_id, limit=limit
|
||||
)
|
||||
terms = query.split()
|
||||
terms = _normalize_search_terms(query)
|
||||
if not terms:
|
||||
return self.list_structured_memories(
|
||||
mem_type=mem_type, scope=scope, scope_id=scope_id, limit=limit
|
||||
)
|
||||
with self._conn() as conn:
|
||||
clauses = []
|
||||
params: dict[str, str] = {}
|
||||
@@ -3352,25 +3363,116 @@ class PostgreSQLBackend:
|
||||
params[f"n{i}"] = f"%{escaped}%"
|
||||
params[f"d{i}"] = f"%{escaped}%"
|
||||
params[f"c{i}"] = f"%{escaped}%"
|
||||
where = " AND ".join(clauses)
|
||||
term_clause = " OR ".join(clauses)
|
||||
scope_filters = ""
|
||||
if mem_type:
|
||||
where += " AND type = :type_filter"
|
||||
scope_filters += " AND type = :type_filter"
|
||||
params["type_filter"] = mem_type
|
||||
if scope:
|
||||
where += " AND scope = :scope_filter"
|
||||
scope_filters += " AND scope = :scope_filter"
|
||||
params["scope_filter"] = scope
|
||||
if scope_id and scope:
|
||||
where += " AND scope_id = :scope_id_filter"
|
||||
scope_filters += " AND scope_id = :scope_id_filter"
|
||||
params["scope_id_filter"] = scope_id
|
||||
rows = conn.execute(
|
||||
sa.text(
|
||||
f"SELECT * FROM structured_memories WHERE {where} "
|
||||
f"ORDER BY updated DESC LIMIT :lim"
|
||||
f"SELECT * FROM structured_memories WHERE ({term_clause}){scope_filters} "
|
||||
f"ORDER BY updated DESC, memory_id ASC LIMIT :lim"
|
||||
),
|
||||
{**params, "lim": limit},
|
||||
).fetchall()
|
||||
return [dict(r._mapping) for r in rows]
|
||||
|
||||
def list_visible_structured_memories(
|
||||
self,
|
||||
scopes: list[tuple[str, str]],
|
||||
mem_type: str = "",
|
||||
limit: int = 100,
|
||||
) -> list[dict[str, str]]:
|
||||
"""Single-query union across visible (scope, scope_id) pairs.
|
||||
|
||||
Replaces the per-scope fan-out (one query per visible scope) so the
|
||||
composition path issues 1 round-trip instead of 3.
|
||||
"""
|
||||
if not scopes:
|
||||
return []
|
||||
with self._conn() as conn:
|
||||
scope_clauses, params = self._build_scope_or_clause(scopes)
|
||||
extra = ""
|
||||
if mem_type:
|
||||
extra = " AND type = :type_filter"
|
||||
params["type_filter"] = mem_type
|
||||
rows = conn.execute(
|
||||
sa.text(
|
||||
f"SELECT * FROM structured_memories WHERE ({scope_clauses}){extra} "
|
||||
f"ORDER BY updated DESC, memory_id ASC LIMIT :lim"
|
||||
),
|
||||
{**params, "lim": limit},
|
||||
).fetchall()
|
||||
return [dict(r._mapping) for r in rows]
|
||||
|
||||
def search_visible_structured_memories(
|
||||
self,
|
||||
query: str,
|
||||
scopes: list[tuple[str, str]],
|
||||
mem_type: str = "",
|
||||
limit: int = 20,
|
||||
) -> list[dict[str, str]]:
|
||||
"""OR-of-terms search joined with a single visibility OR-group.
|
||||
|
||||
Replaces the per-scope search fan-out; ranking is the caller's job.
|
||||
"""
|
||||
if not scopes:
|
||||
return []
|
||||
if not query or not query.strip():
|
||||
return self.list_visible_structured_memories(scopes, mem_type=mem_type, limit=limit)
|
||||
terms = _normalize_search_terms(query)
|
||||
if not terms:
|
||||
return self.list_visible_structured_memories(scopes, mem_type=mem_type, limit=limit)
|
||||
with self._conn() as conn:
|
||||
scope_clauses, params = self._build_scope_or_clause(scopes)
|
||||
term_clauses = []
|
||||
for i, t in enumerate(terms):
|
||||
escaped = _escape_ilike(t)
|
||||
term_clauses.append(
|
||||
f"(name ILIKE :n{i} ESCAPE '\\' "
|
||||
f"OR description ILIKE :d{i} ESCAPE '\\' "
|
||||
f"OR content ILIKE :c{i} ESCAPE '\\')"
|
||||
)
|
||||
params[f"n{i}"] = f"%{escaped}%"
|
||||
params[f"d{i}"] = f"%{escaped}%"
|
||||
params[f"c{i}"] = f"%{escaped}%"
|
||||
term_clause = " OR ".join(term_clauses)
|
||||
extra = ""
|
||||
if mem_type:
|
||||
extra = " AND type = :type_filter"
|
||||
params["type_filter"] = mem_type
|
||||
rows = conn.execute(
|
||||
sa.text(
|
||||
f"SELECT * FROM structured_memories "
|
||||
f"WHERE ({scope_clauses}) AND ({term_clause}){extra} "
|
||||
f"ORDER BY updated DESC, memory_id ASC LIMIT :lim"
|
||||
),
|
||||
{**params, "lim": limit},
|
||||
).fetchall()
|
||||
return [dict(r._mapping) for r in rows]
|
||||
|
||||
@staticmethod
|
||||
def _build_scope_or_clause(
|
||||
scopes: list[tuple[str, str]],
|
||||
) -> tuple[str, dict[str, str]]:
|
||||
"""Build a parameterized OR-group of (scope[, scope_id]) predicates."""
|
||||
params: dict[str, str] = {}
|
||||
clauses: list[str] = []
|
||||
for i, (s, sid) in enumerate(scopes):
|
||||
params[f"sc{i}"] = s
|
||||
if sid:
|
||||
params[f"sid{i}"] = sid
|
||||
clauses.append(f"(scope = :sc{i} AND scope_id = :sid{i})")
|
||||
else:
|
||||
clauses.append(f"scope = :sc{i}")
|
||||
return " OR ".join(clauses), params
|
||||
|
||||
def touch_structured_memories(self, keys: list[tuple[str, str, str]]) -> int:
|
||||
"""Batch-touch multiple memories by (name, scope, scope_id)."""
|
||||
if not keys:
|
||||
|
||||
@@ -386,6 +386,34 @@ class StorageBackend(Protocol):
|
||||
"""Search structured memories by query. Returns matching memory dicts."""
|
||||
...
|
||||
|
||||
def list_visible_structured_memories(
|
||||
self,
|
||||
scopes: list[tuple[str, str]],
|
||||
mem_type: str = "",
|
||||
limit: int = 100,
|
||||
) -> list[dict[str, str]]:
|
||||
"""List memories matching ANY of the (scope, scope_id) pairs in *scopes*.
|
||||
|
||||
A pair with an empty ``scope_id`` matches the scope alone (used for
|
||||
``("global", "")``). Single SQL query — replaces the per-scope fan-out
|
||||
pattern that issued one query per visible scope.
|
||||
"""
|
||||
...
|
||||
|
||||
def search_visible_structured_memories(
|
||||
self,
|
||||
query: str,
|
||||
scopes: list[tuple[str, str]],
|
||||
mem_type: str = "",
|
||||
limit: int = 20,
|
||||
) -> list[dict[str, str]]:
|
||||
"""OR-of-terms search across memories visible under *scopes*.
|
||||
|
||||
Single SQL query joining the scope OR-group with the term OR-group.
|
||||
Ranking is the caller's job (BM25 downstream).
|
||||
"""
|
||||
...
|
||||
|
||||
def touch_structured_memories(self, keys: list[tuple[str, str, str]]) -> int:
|
||||
"""Batch-touch multiple memories.
|
||||
|
||||
|
||||
@@ -87,6 +87,9 @@ from turnstone.core.storage._utils import (
|
||||
from turnstone.core.storage._utils import (
|
||||
VERDICT_MUTABLE as _VERDICT_MUTABLE,
|
||||
)
|
||||
from turnstone.core.storage._utils import (
|
||||
normalize_search_terms as _normalize_search_terms,
|
||||
)
|
||||
from turnstone.core.storage._utils import (
|
||||
reconstruct_messages as _reconstruct_messages,
|
||||
)
|
||||
@@ -3454,7 +3457,10 @@ class SQLiteBackend:
|
||||
limit: int = 100,
|
||||
) -> list[dict[str, str]]:
|
||||
with self._conn() as conn:
|
||||
q = sa.select(structured_memories).order_by(structured_memories.c.updated.desc())
|
||||
q = sa.select(structured_memories).order_by(
|
||||
structured_memories.c.updated.desc(),
|
||||
structured_memories.c.memory_id.asc(),
|
||||
)
|
||||
if mem_type:
|
||||
q = q.where(structured_memories.c.type == mem_type)
|
||||
if scope:
|
||||
@@ -3473,11 +3479,16 @@ class SQLiteBackend:
|
||||
scope_id: str = "",
|
||||
limit: int = 20,
|
||||
) -> list[dict[str, str]]:
|
||||
"""OR-of-terms LIKE search; ranking is the caller's job (BM25 downstream)."""
|
||||
if not query or not query.strip():
|
||||
return self.list_structured_memories(
|
||||
mem_type=mem_type, scope=scope, scope_id=scope_id, limit=limit
|
||||
)
|
||||
terms = query.split()
|
||||
terms = _normalize_search_terms(query)
|
||||
if not terms:
|
||||
return self.list_structured_memories(
|
||||
mem_type=mem_type, scope=scope, scope_id=scope_id, limit=limit
|
||||
)
|
||||
with self._conn() as conn:
|
||||
clauses = []
|
||||
params: dict[str, str] = {}
|
||||
@@ -3491,25 +3502,109 @@ class SQLiteBackend:
|
||||
params[f"n{i}"] = f"%{escaped}%"
|
||||
params[f"d{i}"] = f"%{escaped}%"
|
||||
params[f"c{i}"] = f"%{escaped}%"
|
||||
where = " AND ".join(clauses)
|
||||
term_clause = " OR ".join(clauses)
|
||||
scope_filters = ""
|
||||
if mem_type:
|
||||
where += " AND type = :type_filter"
|
||||
scope_filters += " AND type = :type_filter"
|
||||
params["type_filter"] = mem_type
|
||||
if scope:
|
||||
where += " AND scope = :scope_filter"
|
||||
scope_filters += " AND scope = :scope_filter"
|
||||
params["scope_filter"] = scope
|
||||
if scope_id and scope:
|
||||
where += " AND scope_id = :scope_id_filter"
|
||||
scope_filters += " AND scope_id = :scope_id_filter"
|
||||
params["scope_id_filter"] = scope_id
|
||||
rows = conn.execute(
|
||||
sa.text(
|
||||
f"SELECT * FROM structured_memories WHERE {where} "
|
||||
f"ORDER BY updated DESC LIMIT :lim"
|
||||
f"SELECT * FROM structured_memories WHERE ({term_clause}){scope_filters} "
|
||||
f"ORDER BY updated DESC, memory_id ASC LIMIT :lim"
|
||||
),
|
||||
{**params, "lim": limit},
|
||||
).fetchall()
|
||||
return [dict(r._mapping) for r in rows]
|
||||
|
||||
def list_visible_structured_memories(
|
||||
self,
|
||||
scopes: list[tuple[str, str]],
|
||||
mem_type: str = "",
|
||||
limit: int = 100,
|
||||
) -> list[dict[str, str]]:
|
||||
"""Single-query union across visible (scope, scope_id) pairs."""
|
||||
if not scopes:
|
||||
return []
|
||||
with self._conn() as conn:
|
||||
scope_clauses, params = self._build_scope_or_clause(scopes)
|
||||
extra = ""
|
||||
if mem_type:
|
||||
extra = " AND type = :type_filter"
|
||||
params["type_filter"] = mem_type
|
||||
rows = conn.execute(
|
||||
sa.text(
|
||||
f"SELECT * FROM structured_memories WHERE ({scope_clauses}){extra} "
|
||||
f"ORDER BY updated DESC, memory_id ASC LIMIT :lim"
|
||||
),
|
||||
{**params, "lim": limit},
|
||||
).fetchall()
|
||||
return [dict(r._mapping) for r in rows]
|
||||
|
||||
def search_visible_structured_memories(
|
||||
self,
|
||||
query: str,
|
||||
scopes: list[tuple[str, str]],
|
||||
mem_type: str = "",
|
||||
limit: int = 20,
|
||||
) -> list[dict[str, str]]:
|
||||
"""OR-of-terms search joined with a single visibility OR-group."""
|
||||
if not scopes:
|
||||
return []
|
||||
if not query or not query.strip():
|
||||
return self.list_visible_structured_memories(scopes, mem_type=mem_type, limit=limit)
|
||||
terms = _normalize_search_terms(query)
|
||||
if not terms:
|
||||
return self.list_visible_structured_memories(scopes, mem_type=mem_type, limit=limit)
|
||||
with self._conn() as conn:
|
||||
scope_clauses, params = self._build_scope_or_clause(scopes)
|
||||
term_clauses = []
|
||||
for i, t in enumerate(terms):
|
||||
escaped = _escape_like(t)
|
||||
term_clauses.append(
|
||||
f"(name LIKE :n{i} ESCAPE '\\' "
|
||||
f"OR description LIKE :d{i} ESCAPE '\\' "
|
||||
f"OR content LIKE :c{i} ESCAPE '\\')"
|
||||
)
|
||||
params[f"n{i}"] = f"%{escaped}%"
|
||||
params[f"d{i}"] = f"%{escaped}%"
|
||||
params[f"c{i}"] = f"%{escaped}%"
|
||||
term_clause = " OR ".join(term_clauses)
|
||||
extra = ""
|
||||
if mem_type:
|
||||
extra = " AND type = :type_filter"
|
||||
params["type_filter"] = mem_type
|
||||
rows = conn.execute(
|
||||
sa.text(
|
||||
f"SELECT * FROM structured_memories "
|
||||
f"WHERE ({scope_clauses}) AND ({term_clause}){extra} "
|
||||
f"ORDER BY updated DESC, memory_id ASC LIMIT :lim"
|
||||
),
|
||||
{**params, "lim": limit},
|
||||
).fetchall()
|
||||
return [dict(r._mapping) for r in rows]
|
||||
|
||||
@staticmethod
|
||||
def _build_scope_or_clause(
|
||||
scopes: list[tuple[str, str]],
|
||||
) -> tuple[str, dict[str, str]]:
|
||||
"""Build a parameterized OR-group of (scope[, scope_id]) predicates."""
|
||||
params: dict[str, str] = {}
|
||||
clauses: list[str] = []
|
||||
for i, (s, sid) in enumerate(scopes):
|
||||
params[f"sc{i}"] = s
|
||||
if sid:
|
||||
params[f"sid{i}"] = sid
|
||||
clauses.append(f"(scope = :sc{i} AND scope_id = :sid{i})")
|
||||
else:
|
||||
clauses.append(f"scope = :sc{i}")
|
||||
return " OR ".join(clauses), params
|
||||
|
||||
def touch_structured_memories(self, keys: list[tuple[str, str, str]]) -> int:
|
||||
"""Batch-touch multiple memories by (name, scope, scope_id)."""
|
||||
if not keys:
|
||||
|
||||
@@ -5,6 +5,7 @@ from __future__ import annotations
|
||||
import base64
|
||||
import contextlib
|
||||
import json
|
||||
import re
|
||||
from typing import Any
|
||||
|
||||
from turnstone.core.attachments import unreadable_placeholder
|
||||
@@ -48,6 +49,39 @@ def _attachment_to_content_part(att: dict[str, Any]) -> dict[str, Any] | None:
|
||||
return None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Search-term normalization
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
# Composition can hand a multi-KB pasted user message to ILIKE-based search;
|
||||
# without a cap, every distinct token would emit one unindexable predicate
|
||||
# per scope-fanned query, producing hundreds of seq-scan clauses on a single
|
||||
# rebuild. Cap + dedupe + length filter keeps the SQL bounded.
|
||||
_MAX_SEARCH_TERMS = 16
|
||||
_MIN_TERM_LEN = 2
|
||||
|
||||
# Streaming tokenizer — finditer doesn't allocate a full list up front,
|
||||
# so a multi-KB pasted query stops being scanned the moment the cap is
|
||||
# hit instead of after splitting every token.
|
||||
_TOKEN_RE = re.compile(r"\S+")
|
||||
|
||||
|
||||
def normalize_search_terms(query: str) -> list[str]:
|
||||
"""De-dupe (case-insensitive), drop short tokens, and cap at MAX terms."""
|
||||
seen: set[str] = set()
|
||||
terms: list[str] = []
|
||||
for match in _TOKEN_RE.finditer(query):
|
||||
raw = match.group()
|
||||
lowered = raw.lower()
|
||||
if len(lowered) < _MIN_TERM_LEN or lowered in seen:
|
||||
continue
|
||||
seen.add(lowered)
|
||||
terms.append(raw)
|
||||
if len(terms) >= _MAX_SEARCH_TERMS:
|
||||
break
|
||||
return terms
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Text sanitization
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
Reference in New Issue
Block a user