diff --git a/headroom/memory/adapters/sqlite.py b/headroom/memory/adapters/sqlite.py index 66ffb56f9..42654c12c 100644 --- a/headroom/memory/adapters/sqlite.py +++ b/headroom/memory/adapters/sqlite.py @@ -342,6 +342,31 @@ class SQLiteMemoryStore: return [self._row_to_memory(row) for row in cursor] + async def record_access( + self, + memory_ids: list[str], + accessed_at: datetime | None = None, + ) -> int: + """Atomically record one retrieval for each distinct memory ID.""" + unique_ids = list(dict.fromkeys(memory_ids)) + if not unique_ids: + return 0 + + timestamp = accessed_at or datetime.utcnow() + placeholders = ", ".join("?" for _ in unique_ids) + with self._get_conn() as conn: + cursor = conn.execute( + f""" + UPDATE memories + SET access_count = access_count + 1, + last_accessed = ? + WHERE id IN ({placeholders}) + """, # nosec B608 + [timestamp.isoformat(), *unique_ids], + ) + conn.commit() + return cursor.rowcount + async def delete(self, memory_id: str) -> bool: """Delete a memory by ID. diff --git a/headroom/memory/backends/local.py b/headroom/memory/backends/local.py index a74fa12b8..757ccdbed 100644 --- a/headroom/memory/backends/local.py +++ b/headroom/memory/backends/local.py @@ -516,6 +516,12 @@ class LocalBackend: results.sort(key=lambda x: x.score, reverse=True) return results[:top_k] + async def record_access(self, memory_ids: list[str]) -> int: + """Record retrieval metadata for memories returned to a caller.""" + await self._ensure_initialized() + assert self._hierarchical_memory is not None + return await self._hierarchical_memory.record_access(memory_ids) + async def update_memory( self, memory_id: str, diff --git a/headroom/memory/core.py b/headroom/memory/core.py index 3a26f1e9a..746a2e5a9 100644 --- a/headroom/memory/core.py +++ b/headroom/memory/core.py @@ -304,6 +304,21 @@ class HierarchicalMemory: return memory + async def record_access( + self, + memory_ids: list[str], + accessed_at: datetime | None = None, + ) -> int: + """Record retrieval metadata for memories returned to a caller.""" + unique_ids = list(dict.fromkeys(memory_ids)) + if not unique_ids: + return 0 + + updated = await self._store.record_access(unique_ids, accessed_at) + if self._cache is not None: + await self._cache.invalidate_batch(unique_ids) + return updated + async def query(self, filter: MemoryFilter) -> list[Memory]: """Query memories with filtering. diff --git a/headroom/memory/mcp_server.py b/headroom/memory/mcp_server.py index d5fecab97..e71fe9b54 100644 --- a/headroom/memory/mcp_server.py +++ b/headroom/memory/mcp_server.py @@ -245,6 +245,12 @@ async def _handle_search( # Trim to requested top_k active_results = active_results[:top_k] + try: + await backend.record_access([r.memory.id for r in active_results]) + except Exception as e: + # Usage metadata must never make a successful retrieval fail. + logger.warning(f"Memory MCP: failed to record access: {e}") + lines = [] for i, r in enumerate(active_results, 1): score = f"{r.score:.2f}" if hasattr(r, "score") else "?" diff --git a/headroom/memory/ports.py b/headroom/memory/ports.py index fd2790cf2..7fd0b7bad 100644 --- a/headroom/memory/ports.py +++ b/headroom/memory/ports.py @@ -311,6 +311,22 @@ class MemoryStore(Protocol): """ ... + async def record_access( + self, + memory_ids: list[str], + accessed_at: datetime | None = None, + ) -> int: + """Record one retrieval for each distinct memory ID. + + Args: + memory_ids: IDs of memories actually returned to a caller. + accessed_at: Retrieval time (defaults to now). + + Returns: + Number of existing memories updated. + """ + ... + async def delete(self, memory_id: str) -> bool: """ Delete a memory by ID. diff --git a/tests/test_memory/test_hierarchical.py b/tests/test_memory/test_hierarchical.py index 5394fedaf..7533f1a24 100644 --- a/tests/test_memory/test_hierarchical.py +++ b/tests/test_memory/test_hierarchical.py @@ -182,6 +182,36 @@ class TestSQLiteMemoryStore: assert retrieved is not None assert retrieved.content == memory.content + @pytest.mark.asyncio + async def test_record_access_is_atomic_and_deduplicates_ids(self, store): + memories = [Memory(content=f"Memory {i}", user_id="alice") for i in range(2)] + await store.save_batch(memories) + + first_access = datetime(2026, 7, 12, 9, 30) + updated = await store.record_access( + [memories[0].id, memories[0].id, memories[1].id, "missing"], + first_access, + ) + + assert updated == 2 + first = await store.get(memories[0].id) + second = await store.get(memories[1].id) + assert first is not None + assert second is not None + assert first.access_count == 1 + assert second.access_count == 1 + assert first.last_accessed == first_access + assert second.last_accessed == first_access + + second_access = datetime(2026, 7, 12, 9, 31) + assert await store.record_access([memories[0].id], second_access) == 1 + first = await store.get(memories[0].id) + assert first is not None + assert first.access_count == 2 + assert first.last_accessed == second_access + + assert await store.record_access([]) == 0 + @pytest.mark.asyncio async def test_delete(self, store, sample_memory): """Test deleting a memory.""" diff --git a/tests/test_memory/test_mcp_server.py b/tests/test_memory/test_mcp_server.py index bd56174ad..34339c31d 100644 --- a/tests/test_memory/test_mcp_server.py +++ b/tests/test_memory/test_mcp_server.py @@ -163,3 +163,68 @@ def test_main_logs_memory_mcp_startup_context(monkeypatch, tmp_path, caplog) -> and "resolution=dynamic-cwd" in record.message for record in caplog.records ) + + +def test_search_records_access_only_for_returned_memories() -> None: + active = Memory(content="Active preference", user_id="alice") + extra = Memory(content="Lower-ranked preference", user_id="alice") + backend = SimpleNamespace( + search_memories=AsyncMock( + return_value=[ + SimpleNamespace(memory=active, score=0.9, related_entities=[]), + SimpleNamespace(memory=extra, score=0.8, related_entities=[]), + ] + ), + get_memory=AsyncMock( + side_effect=lambda memory_id: { + active.id: active, + extra.id: extra, + }[memory_id] + ), + record_access=AsyncMock(return_value=1), + ) + + result = asyncio.run( + mcp_server_mod._handle_search( + backend, + {"query": "preference", "top_k": 1}, + "alice", + ) + ) + + backend.record_access.assert_awaited_once_with([active.id]) + assert "Active preference" in result[0].kwargs["text"] + assert "Lower-ranked preference" not in result[0].kwargs["text"] + + +def test_search_does_not_record_superseded_memories() -> None: + superseded = Memory(content="Old preference", user_id="alice") + replacement = Memory(content="Current preference", user_id="alice") + superseded.superseded_by = replacement.id + backend = SimpleNamespace( + search_memories=AsyncMock( + return_value=[SimpleNamespace(memory=superseded, score=0.9, related_entities=[])] + ), + get_memory=AsyncMock(return_value=superseded), + record_access=AsyncMock(), + ) + + result = asyncio.run(mcp_server_mod._handle_search(backend, {"query": "preference"}, "alice")) + + backend.record_access.assert_not_awaited() + assert result[0].kwargs["text"] == "No memories found." + + +def test_search_fails_open_when_access_tracking_fails() -> None: + memory = Memory(content="Useful preference", user_id="alice") + backend = SimpleNamespace( + search_memories=AsyncMock( + return_value=[SimpleNamespace(memory=memory, score=0.9, related_entities=[])] + ), + get_memory=AsyncMock(return_value=memory), + record_access=AsyncMock(side_effect=RuntimeError("write failed")), + ) + + result = asyncio.run(mcp_server_mod._handle_search(backend, {"query": "preference"}, "alice")) + + assert "Useful preference" in result[0].kwargs["text"]