mirror of
https://github.com/headroomlabs-ai/headroom.git
synced 2026-08-27 14:17:10 -04:00
## Description Closes #2947 `entity_refs` is annotated `list[str]` everywhere, but nothing enforced that at runtime. `LocalBackend.save_memory`'s `entities` argument is filled straight from LLM-supplied `memory_save` tool input (`headroom/memory/system.py:575` into `memory_handler.py:1242`), so a caller can pass the typed `{"entity": ..., "entity_type": ...}` shape, which is the format `extracted_entities` expects, into it by mistake. Those dicts were then persisted verbatim into `entity_refs`, both in the `memories` table and in the duplicated copy the vector index keeps for post-filtering. Every later `search_memories` call does `set().update(memory.entity_refs)` while collecting entities for graph expansion. Hashing a dict raises `TypeError: unhashable type: 'dict'`, and because that happens inside the vector-result loop rather than per-item, **one** poisoned row aborted the **entire** search. The proxy's memory handler catches the exception and returns no memories, so recall went quietly dark rather than failing loudly, and the bad row kept re-appearing in top-k for related queries, so it stayed dark. The issue reporter hit this in production: 4 bad rows disabled memory search for a whole project for a day, with nothing visible to the end user beyond a swallowed warning in `proxy.log`. The same root cause has two more crash modes, both confirmed below: `AttributeError: 'dict' object has no attribute 'lower'` during graph linking on the save path, and the same error in the `entities` search filter (`ref.lower()`). The fix adds one helper and applies it at both ends of the data flow. Dicts are **unwrapped to their `entity` name** rather than dropped, so rows that are already corrupted keep contributing to graph expansion instead of silently losing their entities. ## Type of Change - [x] Bug fix (non-breaking change that fixes an issue) - [ ] New feature (non-breaking change that adds functionality) - [ ] Breaking change (fix or feature that would cause existing functionality to change) - [ ] Documentation update - [ ] Performance improvement - [ ] Code refactoring (no functional changes) ## Changes Made - **New helper `normalize_entity_refs()` in `headroom/memory/models.py`.** Coerces a raw entity-reference list into the `list[str]` it claims to be: strings pass through, dicts are unwrapped via their `entity` (or `name`) key, and anything with no recoverable name is dropped rather than stringified, since a ref like `"{'entity_type': 'project'}"` would only pollute the graph. Order is preserved and duplicate names are collapsed. - **Write path, to stop new corruption at the door.** `LocalBackend.save_memory` normalizes `entities` before it reaches `entity_refs` and graph linking. `LocalBackend.search_memories` normalizes the `entities` *filter* argument too, since it arrives from the same untrusted tool input (`memory_handler.py:1320`). - **Read path, to heal rows that were written before this fix.** Applied at the three deserialization boundaries, so no data migration is needed and corrupted rows normalize themselves the next time they are loaded: `Memory.from_dict` (`headroom/memory/models.py`), `SQLiteMemoryStore._row_to_memory` (`headroom/memory/adapters/sqlite.py`), and the vector indexes' own `entity_refs` copies used for post-filtering, `VectorMetadata.from_json` (`headroom/memory/adapters/sqlite_vector.py`) and `IndexedMemoryMetadata.from_dict` (`headroom/memory/adapters/hnsw.py`). - **Defensive normalization on emitted results.** `search_memories` and `text_search` normalize the refs they return as `related_entities`, so a backend that produces `Memory` objects by some path not covered above still cannot take a whole query down, and callers never receive a dict where they expect an entity name. **Note on scope versus the patch proposed in the issue.** The issue proposed normalizing in two places (`save_memory` plus the `set().update()` line). I widened it slightly because that pair leaves three related failures live: the `entities` filter still crashes on `ref.lower()`, `related_entities` still hands dicts back to the caller, and, most importantly, already-poisoned rows stay poisoned in storage. Normalizing at the deserialization boundaries fixes all three at once and is what makes existing corrupted databases recover on their own. **Behavior change worth flagging.** `entity_refs` is now de-duplicated (case-sensitively) on both save and load. Refs were already treated as a set for graph expansion, so this is semantically a no-op, but it is a visible difference if anything asserts on exact list contents. ## Testing - [x] Unit tests pass (`pytest`) - [x] Linting passes (`ruff check .`) - [x] Type checking passes (`mypy headroom`) - [x] New tests added for new functionality - [ ] Manual testing performed New file `tests/test_memory/test_entity_ref_sanitization.py` adds 10 tests covering the helper, both write paths, all three deserialization boundaries, and the three crash modes. ### Test Output ```text $ python -m pytest tests/test_memory/test_entity_ref_sanitization.py -q .......... [100%] 10 passed, 17 warnings in 0.34s ``` Full memory suite, plus a before/after comparison of the failure set to prove no regressions: ```text $ python -m pytest tests/test_memory/ -q 13 failed, 576 passed, 3 skipped, 1072 warnings, 25 errors in 44.74s # the same run with the source changes stashed (baseline on upstream/main @941c25d3): 13 failed, 566 passed, 3 skipped, 1055 warnings, 25 errors in 44.92s # diff of the failing/erroring test IDs, before versus after: $ diff baseline.txt after.txt && echo "NO NEW FAILURES vs baseline" NO NEW FAILURES vs baseline ``` 576 passed equals the 566 baseline plus the 10 new tests. The 13 failures and 25 errors are pre-existing on `upstream/main` and unrelated to this change: they are Windows-only temp-directory cleanup failures in this local environment. ```text E PermissionError: [WinError 32] The process cannot access the file because it is being used by another process: 'C:\Users\...\Temp\tmp_qnqtami\test.db' ``` Adjacent suites that construct `Memory` objects: ```text $ python -m pytest tests/test_memory_system.py tests/test_memory_eval.py tests/test_critical_gaps.py -q 158 passed, 1 skipped, 514 warnings in 20.58s ``` Lint and format on the changed files: ```text $ ruff check headroom/memory tests/test_memory/test_entity_ref_sanitization.py All checks passed! $ ruff format --check headroom/memory tests/test_memory/test_entity_ref_sanitization.py 48 files already formatted ``` `mypy headroom --ignore-missing-imports --python-version 3.13` reports 12 errors, all pre-existing on `upstream/main` and all in files this PR does not touch (`headroom/ccr/mcp_server.py`, `headroom/memory/mcp_server.py`, `headroom/release_version.py`; they come from a local MCP SDK version mismatch). Zero errors in any changed file. ## Real Behavior Proof - Environment: Windows 11, Python 3.13.11, pytest 9.1.1, ruff 0.15.17, local `headroom._core` built. Branched from `upstream/main` at `941c25d3`, the same branch point the issue reports. - Exact command / steps: Ran a standalone script (not a mock-only test) driving `LocalBackend.search_memories` and `LocalBackend.save_memory` with `entity_refs=[{"entity": "Project X", "entity_type": "project"}]`, first against unmodified `941c25d3` and then against this branch. Three scenarios: vector search with graph expansion, search with an `entities` filter, and a save carrying dict-shaped `entities`. - Observed result: on unmodified `941c25d3` all three crashed, printing `SEARCH: TypeError: unhashable type: 'dict'`, `FILTER: TypeError: unhashable type: 'dict'`, and `SAVE: AttributeError: 'dict' object has no attribute 'lower'`. With this branch applied all three succeed: search returns both the poisoned and the clean memory with `related_entities == ["Project X"]`, the filter matches the recovered name, and the save persists `entity_refs == ["Project X"]`. Those three scenarios are now the regression tests in `test_entity_ref_sanitization.py`. - Not tested: no live end-to-end run through the MCP `memory_save` tool against a real LLM, and no test against a real pre-existing SQLite database containing dict-shaped rows. The healing-on-load path is covered at the deserialization functions (`Memory.from_dict`, `VectorMetadata.from_json`, `IndexedMemoryMetadata.from_dict`) rather than through an actual corrupted `.db` file. The non-local backends (`mem0`, `direct_mem0`, `qdrant-neo4j`, `cognee`) were not exercised; this PR only changes the local backend and the shared models and adapters. ## Review Readiness - [x] I have performed a self-review - [x] This PR is ready for human review ## Checklist - [x] My code follows the project's style guidelines - [x] I have performed a self-review of my code - [x] I have commented my code, particularly in hard-to-understand areas - [ ] I have made corresponding changes to the documentation - [x] My changes generate no new warnings - [x] I have added tests that prove my fix is effective or that my feature works - [x] New and existing unit tests pass locally with my changes - [x] I did **not** edit `CHANGELOG.md`, it is generated by release-please from my Conventional Commit PR title (a CI guard enforces this) ## Additional Notes Documentation was not changed because `normalize_entity_refs()` is an internal helper and no public API or user-facing behavior changes; `entity_refs` still behaves exactly as its existing `list[str]` contract always documented. Credit for the diagnosis, the root-cause analysis, and the original repro goes to @apacheco-RT in #2947, who could not open a PR directly because GitHub blocks Enterprise Managed User accounts from forking outside their enterprise. Co-authored-by: Claude Opus 5 <noreply@anthropic.com> Co-authored-by: JD Davis <mxjerrett@gmail.com>
302 lines
11 KiB
Python
302 lines
11 KiB
Python
"""Regression tests for entity_refs type safety.
|
|
|
|
`entity_refs` is typed `list[str]` everywhere, but nothing enforced that at
|
|
runtime. A caller that mistakenly passed the typed
|
|
`{"entity": ..., "entity_type": ...}` shape (the format `extracted_entities`
|
|
expects) into the plain `entities` field of `save_memory` got those dicts
|
|
persisted verbatim into `entity_refs` -- both in the `memories` table and in
|
|
the duplicated copy the vector index keeps for post-filtering.
|
|
|
|
Every later `search_memories` call does `set().update(memory.entity_refs)`
|
|
while collecting entities for graph expansion. Hashing a dict raises
|
|
`TypeError: unhashable type: 'dict'`, and because that happens inside the
|
|
vector-result loop (not guarded per-item) it aborted the *entire* search for
|
|
any query whose top-k included one poisoned row. The proxy's memory handler
|
|
swallows the exception and returns no memories, so recall went quietly dark
|
|
rather than failing loudly.
|
|
|
|
See https://github.com/headroomlabs-ai/headroom/issues/2947.
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
import json
|
|
from datetime import datetime, timezone
|
|
from types import SimpleNamespace
|
|
from unittest.mock import AsyncMock
|
|
|
|
import pytest
|
|
|
|
from headroom.memory.adapters.hnsw import IndexedMemoryMetadata
|
|
from headroom.memory.adapters.sqlite_vector import VectorMetadata
|
|
from headroom.memory.backends.local import LocalBackend
|
|
from headroom.memory.models import Memory, normalize_entity_refs
|
|
|
|
# The malformed shape that started all of this: the extracted_entities format
|
|
# passed into a field that expects plain names.
|
|
DICT_REF = {"entity": "Project X", "entity_type": "project"}
|
|
|
|
|
|
# =============================================================================
|
|
# The helper itself
|
|
# =============================================================================
|
|
|
|
|
|
def test_normalize_entity_refs_unwraps_dicts_and_drops_junk() -> None:
|
|
"""Dicts are unwrapped to their name; anything unusable is dropped."""
|
|
assert normalize_entity_refs(["Alice", DICT_REF]) == ["Alice", "Project X"]
|
|
|
|
# Nothing usable in these: no name to recover, so they are dropped rather
|
|
# than stringified into garbage entity names like "{'foo': 'bar'}".
|
|
assert normalize_entity_refs([{"entity_type": "project"}, {}, None, 42, ""]) == []
|
|
|
|
# Common no-op cases stay untouched.
|
|
assert normalize_entity_refs(["Alice", "Bob"]) == ["Alice", "Bob"]
|
|
assert normalize_entity_refs(None) == []
|
|
assert normalize_entity_refs([]) == []
|
|
|
|
|
|
def test_normalize_entity_refs_preserves_order_and_deduplicates() -> None:
|
|
"""A name already present is not appended twice, and order is stable."""
|
|
assert normalize_entity_refs(["Alice", DICT_REF, "Alice", "Project X"]) == [
|
|
"Alice",
|
|
"Project X",
|
|
]
|
|
|
|
|
|
# =============================================================================
|
|
# Write path: stop new corruption at the door
|
|
# =============================================================================
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_save_memory_sanitizes_dict_shaped_entities_param() -> None:
|
|
"""`entities` items that are dicts get coerced to plain names before storage."""
|
|
backend = LocalBackend()
|
|
backend._initialized = True
|
|
|
|
saved: list[Memory] = []
|
|
|
|
async def fake_add(**kwargs: object) -> Memory:
|
|
memory = Memory(
|
|
id="new-memory",
|
|
content=str(kwargs["content"]),
|
|
user_id=str(kwargs["user_id"]),
|
|
entity_refs=list(kwargs["entity_refs"]), # type: ignore[arg-type]
|
|
)
|
|
saved.append(memory)
|
|
return memory
|
|
|
|
backend._hierarchical_memory = SimpleNamespace(add=AsyncMock(side_effect=fake_add))
|
|
backend._graph = SimpleNamespace(
|
|
get_entity_by_name=AsyncMock(return_value=None),
|
|
add_entity=AsyncMock(return_value=SimpleNamespace(id="entity-id")),
|
|
add_relationship=AsyncMock(),
|
|
)
|
|
|
|
# Without the fix this raises AttributeError: 'dict' object has no
|
|
# attribute 'lower' during graph linking.
|
|
await backend.save_memory(
|
|
content="Alice manages Project X",
|
|
user_id="alice",
|
|
entities=[DICT_REF], # type: ignore[list-item]
|
|
)
|
|
|
|
assert saved[0].entity_refs == ["Project X"]
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_save_memory_merges_dict_entities_with_extracted_entities() -> None:
|
|
"""A name arriving through both `entities` and `extracted_entities` is stored once."""
|
|
backend = LocalBackend()
|
|
backend._initialized = True
|
|
|
|
saved: list[Memory] = []
|
|
|
|
async def fake_add(**kwargs: object) -> Memory:
|
|
memory = Memory(
|
|
id="new-memory",
|
|
content=str(kwargs["content"]),
|
|
user_id=str(kwargs["user_id"]),
|
|
entity_refs=list(kwargs["entity_refs"]), # type: ignore[arg-type]
|
|
)
|
|
saved.append(memory)
|
|
return memory
|
|
|
|
backend._hierarchical_memory = SimpleNamespace(add=AsyncMock(side_effect=fake_add))
|
|
backend._graph = SimpleNamespace(
|
|
get_entity_by_name=AsyncMock(return_value=None),
|
|
add_entity=AsyncMock(return_value=SimpleNamespace(id="entity-id")),
|
|
add_relationship=AsyncMock(),
|
|
)
|
|
|
|
await backend.save_memory(
|
|
content="Alice manages Project X",
|
|
user_id="alice",
|
|
entities=[DICT_REF], # type: ignore[list-item]
|
|
extracted_entities=[{"entity": "Project X", "entity_type": "project"}],
|
|
)
|
|
|
|
assert saved[0].entity_refs == ["Project X"]
|
|
|
|
|
|
# =============================================================================
|
|
# Read path: heal rows that were already written before the fix
|
|
# =============================================================================
|
|
|
|
|
|
def test_memory_from_dict_heals_stored_dict_refs() -> None:
|
|
"""Rows persisted before the fix load as plain names instead of dicts."""
|
|
now = datetime.now(timezone.utc).isoformat()
|
|
memory = Memory.from_dict(
|
|
{
|
|
"id": "poisoned-memory",
|
|
"content": "Alice manages Project X",
|
|
"user_id": "alice",
|
|
"created_at": now,
|
|
"valid_from": now,
|
|
"importance": 0.5,
|
|
"entity_refs": [DICT_REF, "Alice"],
|
|
}
|
|
)
|
|
|
|
assert memory.entity_refs == ["Project X", "Alice"]
|
|
|
|
|
|
def test_vector_metadata_from_json_heals_stored_dict_refs() -> None:
|
|
"""The vector index keeps its own copy of entity_refs; heal that one too."""
|
|
now = datetime.now(timezone.utc).isoformat()
|
|
metadata = VectorMetadata.from_json(
|
|
json.dumps(
|
|
{
|
|
"memory_id": "poisoned-memory",
|
|
"user_id": "alice",
|
|
"session_id": None,
|
|
"agent_id": None,
|
|
"valid_until": None,
|
|
"entity_refs": [DICT_REF],
|
|
"content": "Alice manages Project X",
|
|
"created_at": now,
|
|
"importance": 0.5,
|
|
"metadata": {},
|
|
}
|
|
)
|
|
)
|
|
|
|
assert metadata.entity_refs == ["Project X"]
|
|
assert metadata.to_memory().entity_refs == ["Project X"]
|
|
|
|
|
|
def test_indexed_memory_metadata_from_dict_heals_stored_dict_refs() -> None:
|
|
"""Same for the HNSW index's metadata copy."""
|
|
now = datetime.now(timezone.utc).isoformat()
|
|
metadata = IndexedMemoryMetadata.from_dict(
|
|
{
|
|
"memory_id": "poisoned-memory",
|
|
"user_id": "alice",
|
|
"session_id": None,
|
|
"agent_id": None,
|
|
"valid_until": None,
|
|
"entity_refs": [DICT_REF],
|
|
"content": "Alice manages Project X",
|
|
"created_at": now,
|
|
"importance": 0.5,
|
|
"metadata": {},
|
|
}
|
|
)
|
|
|
|
assert metadata.entity_refs == ["Project X"]
|
|
|
|
|
|
# =============================================================================
|
|
# Search: a single bad row must not take the whole query down
|
|
# =============================================================================
|
|
|
|
|
|
def _backend_with_results(memories: list[Memory]) -> LocalBackend:
|
|
backend = LocalBackend()
|
|
backend._initialized = True
|
|
backend._hierarchical_memory = SimpleNamespace(
|
|
search=AsyncMock(return_value=[SimpleNamespace(memory=m, similarity=0.9) for m in memories])
|
|
)
|
|
backend._graph = SimpleNamespace(
|
|
get_entity_by_name=AsyncMock(return_value=None),
|
|
query_subgraph=AsyncMock(return_value=SimpleNamespace(entities=[], relationships=[])),
|
|
)
|
|
return backend
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_search_memories_tolerates_dict_shaped_entity_refs() -> None:
|
|
"""A single legacy/corrupted row with dict entity_refs must not crash search."""
|
|
poisoned = Memory(
|
|
id="poisoned-memory",
|
|
content="Alice manages Project X",
|
|
user_id="alice",
|
|
entity_refs=[DICT_REF], # type: ignore[list-item]
|
|
)
|
|
clean = Memory(
|
|
id="clean-memory",
|
|
content="Bob manages Project Y",
|
|
user_id="alice",
|
|
entity_refs=["Project Y"],
|
|
)
|
|
backend = _backend_with_results([poisoned, clean])
|
|
|
|
# Without the fix this raises TypeError: unhashable type: 'dict'.
|
|
results = await backend.search_memories("Alice's work", "alice", include_related=True)
|
|
|
|
assert [r.memory.id for r in results] == ["poisoned-memory", "clean-memory"]
|
|
# The recovered name is still usable for graph expansion and is reported
|
|
# back to the caller as a plain string, not a dict.
|
|
assert results[0].related_entities == ["Project X"]
|
|
backend._graph.get_entity_by_name.assert_awaited()
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_search_memories_entity_filter_matches_healed_refs() -> None:
|
|
"""The `entities` filter lowercases each ref, which dicts also break.
|
|
|
|
On unfixed code this never gets that far -- the unconditional
|
|
`set().update()` above raises first -- but once refs are strings again the
|
|
filter has to actually match the recovered name.
|
|
"""
|
|
poisoned = Memory(
|
|
id="poisoned-memory",
|
|
content="Alice manages Project X",
|
|
user_id="alice",
|
|
entity_refs=[DICT_REF], # type: ignore[list-item]
|
|
)
|
|
backend = _backend_with_results([poisoned])
|
|
|
|
results = await backend.search_memories(
|
|
"Alice's work",
|
|
"alice",
|
|
include_related=False,
|
|
entities=["project x"],
|
|
)
|
|
|
|
assert [r.memory.id for r in results] == ["poisoned-memory"]
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_search_memories_tolerates_dict_shaped_entities_filter() -> None:
|
|
"""The filter argument comes from LLM tool input too, so it can be malformed."""
|
|
clean = Memory(
|
|
id="clean-memory",
|
|
content="Alice manages Project X",
|
|
user_id="alice",
|
|
entity_refs=["Project X"],
|
|
)
|
|
backend = _backend_with_results([clean])
|
|
|
|
# Without normalization this raises AttributeError: 'dict' object has no
|
|
# attribute 'lower' while building the filter set.
|
|
results = await backend.search_memories(
|
|
"Alice's work",
|
|
"alice",
|
|
include_related=False,
|
|
entities=[DICT_REF], # type: ignore[list-item]
|
|
)
|
|
|
|
assert [r.memory.id for r in results] == ["clean-memory"]
|