fix(memory): sanitize entity_refs to prevent dict-shaped entries crashing search (#2951)
## 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>
This commit is contained in:
@@ -22,7 +22,7 @@ from typing import TYPE_CHECKING, Any
|
||||
|
||||
import numpy as np
|
||||
|
||||
from ..models import Memory, ScopeLevel
|
||||
from ..models import Memory, ScopeLevel, normalize_entity_refs
|
||||
from ..ports import VectorFilter, VectorSearchResult
|
||||
|
||||
# hnswlib is optional - may not compile on all platforms
|
||||
@@ -139,7 +139,9 @@ class IndexedMemoryMetadata:
|
||||
valid_until=(
|
||||
datetime.fromisoformat(data["valid_until"]) if data.get("valid_until") else None
|
||||
),
|
||||
entity_refs=data.get("entity_refs", []),
|
||||
# Normalized on load so rows written before #2947 was fixed heal
|
||||
# themselves instead of crashing search.
|
||||
entity_refs=normalize_entity_refs(data.get("entity_refs")),
|
||||
content=data["content"],
|
||||
created_at=datetime.fromisoformat(data["created_at"]),
|
||||
importance=data.get("importance", 0.5),
|
||||
|
||||
@@ -16,7 +16,7 @@ from datetime import datetime, timezone
|
||||
from pathlib import Path
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
from ..models import Memory, ScopeLevel
|
||||
from ..models import Memory, ScopeLevel, normalize_entity_refs
|
||||
from ..ports import MemoryFilter
|
||||
|
||||
if TYPE_CHECKING:
|
||||
@@ -227,7 +227,11 @@ class SQLiteMemoryStore:
|
||||
last_accessed=datetime.fromisoformat(row["last_accessed"])
|
||||
if row["last_accessed"]
|
||||
else None,
|
||||
entity_refs=json.loads(row["entity_refs"]) if row["entity_refs"] else [],
|
||||
# Normalized on load so rows written before #2947 was fixed heal
|
||||
# themselves instead of crashing search.
|
||||
entity_refs=normalize_entity_refs(
|
||||
json.loads(row["entity_refs"]) if row["entity_refs"] else []
|
||||
),
|
||||
embedding=self._deserialize_embedding(row["embedding"]),
|
||||
metadata=json.loads(row["metadata"]) if row["metadata"] else {},
|
||||
)
|
||||
|
||||
@@ -29,7 +29,7 @@ from typing import TYPE_CHECKING, Any, cast
|
||||
|
||||
import numpy as np
|
||||
|
||||
from ..models import Memory, ScopeLevel
|
||||
from ..models import Memory, ScopeLevel, normalize_entity_refs
|
||||
from ..ports import VectorFilter, VectorSearchResult
|
||||
|
||||
if TYPE_CHECKING:
|
||||
@@ -135,7 +135,9 @@ class VectorMetadata:
|
||||
valid_until=(
|
||||
datetime.fromisoformat(d["valid_until"]) if d.get("valid_until") else None
|
||||
),
|
||||
entity_refs=d.get("entity_refs", []),
|
||||
# Normalized on load so rows written before #2947 was fixed heal
|
||||
# themselves instead of crashing search.
|
||||
entity_refs=normalize_entity_refs(d.get("entity_refs")),
|
||||
content=d["content"],
|
||||
created_at=datetime.fromisoformat(d["created_at"]),
|
||||
importance=d.get("importance", 0.5),
|
||||
|
||||
@@ -21,7 +21,7 @@ from pathlib import Path
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
from headroom.memory.adapters.graph_models import Entity, Relationship, Subgraph
|
||||
from headroom.memory.models import Memory
|
||||
from headroom.memory.models import Memory, normalize_entity_refs
|
||||
from headroom.memory.ports import MemorySearchResult
|
||||
from headroom.models.config import ML_MODEL_DEFAULTS
|
||||
|
||||
@@ -284,8 +284,15 @@ class LocalBackend:
|
||||
# Determine if using pre-extraction mode
|
||||
has_pre_extraction = bool(facts or extracted_entities or extracted_relationships)
|
||||
|
||||
# Merge entity names from both simple and typed formats
|
||||
all_entity_names: list[str] = list(entities) if entities else []
|
||||
# Merge entity names from both simple and typed formats.
|
||||
#
|
||||
# `entities` is typed list[str], but it is populated straight from
|
||||
# LLM-supplied memory_save tool arguments, and callers do sometimes
|
||||
# pass the typed {"entity": ..., "entity_type": ...} shape here (that
|
||||
# is what extracted_entities is for). Normalizing keeps those dicts out
|
||||
# of entity_refs, where they used to crash every later search that
|
||||
# retrieved the row -- see issue #2947.
|
||||
all_entity_names: list[str] = normalize_entity_refs(entities)
|
||||
entity_types: dict[str, str] = {}
|
||||
|
||||
if extracted_entities:
|
||||
@@ -448,13 +455,17 @@ class LocalBackend:
|
||||
continue
|
||||
|
||||
seen_memory_ids.add(vr.memory.id)
|
||||
all_entity_refs.update(vr.memory.entity_refs)
|
||||
# Defense in depth: the storage adapters normalize entity_refs on
|
||||
# load, but a backend that hands us Memory objects some other way
|
||||
# must not be able to abort the whole search with one bad row.
|
||||
entity_refs = normalize_entity_refs(vr.memory.entity_refs)
|
||||
all_entity_refs.update(entity_refs)
|
||||
|
||||
results.append(
|
||||
MemorySearchResult(
|
||||
memory=vr.memory,
|
||||
score=vr.similarity,
|
||||
related_entities=list(vr.memory.entity_refs),
|
||||
related_entities=entity_refs,
|
||||
related_memories=[],
|
||||
)
|
||||
)
|
||||
@@ -503,15 +514,17 @@ class LocalBackend:
|
||||
MemorySearchResult(
|
||||
memory=memory,
|
||||
score=0.5, # Default score for graph-expanded results
|
||||
related_entities=list(memory.entity_refs),
|
||||
related_entities=normalize_entity_refs(memory.entity_refs),
|
||||
related_memories=[],
|
||||
)
|
||||
)
|
||||
seen_memory_ids.add(mem_id)
|
||||
|
||||
# Filter by specified entities if provided
|
||||
# Filter by specified entities if provided. Like the save path, this
|
||||
# argument arrives from LLM-supplied tool input, so it gets the same
|
||||
# normalization rather than trusting its list[str] annotation.
|
||||
if entities:
|
||||
entities_lower = {e.lower() for e in entities}
|
||||
entities_lower = {e.lower() for e in normalize_entity_refs(entities)}
|
||||
results = [
|
||||
r
|
||||
for r in results
|
||||
@@ -826,7 +839,7 @@ class LocalBackend:
|
||||
MemorySearchResult(
|
||||
memory=tr.memory,
|
||||
score=tr.score,
|
||||
related_entities=list(tr.memory.entity_refs),
|
||||
related_entities=normalize_entity_refs(tr.memory.entity_refs),
|
||||
related_memories=[],
|
||||
)
|
||||
for tr in text_results
|
||||
|
||||
@@ -14,6 +14,46 @@ except ImportError:
|
||||
np = None # type: ignore[assignment]
|
||||
|
||||
|
||||
def normalize_entity_refs(values: Any) -> list[str]:
|
||||
"""Coerce a raw entity-reference list into the plain ``list[str]`` it claims to be.
|
||||
|
||||
``entity_refs`` (and the ``entities`` argument that feeds it) is typed
|
||||
``list[str]``, but nothing enforced that at runtime, so callers have
|
||||
persisted the typed ``{"entity": ..., "entity_type": ...}`` shape -- the
|
||||
format ``extracted_entities`` expects -- into it by mistake. Those dicts
|
||||
then break every consumer that treats a ref as a string: ``set().update()``
|
||||
raises ``TypeError: unhashable type: 'dict'`` and ``ref.lower()`` raises
|
||||
``AttributeError``, which took down whole memory searches rather than the
|
||||
one bad row (see issue #2947).
|
||||
|
||||
Dicts are unwrapped to their ``entity`` name so no information is lost;
|
||||
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.
|
||||
"""
|
||||
if not values:
|
||||
return []
|
||||
|
||||
normalized: list[str] = []
|
||||
seen: set[str] = set()
|
||||
|
||||
for value in values:
|
||||
if isinstance(value, str):
|
||||
name = value
|
||||
elif isinstance(value, dict):
|
||||
# The extracted_entities shape, mistakenly used as a plain name.
|
||||
candidate = value.get("entity") or value.get("name")
|
||||
name = candidate if isinstance(candidate, str) else ""
|
||||
else:
|
||||
name = ""
|
||||
|
||||
if name and name not in seen:
|
||||
seen.add(name)
|
||||
normalized.append(name)
|
||||
|
||||
return normalized
|
||||
|
||||
|
||||
class ScopeLevel(Enum):
|
||||
"""Memory scope hierarchy levels."""
|
||||
|
||||
@@ -132,7 +172,9 @@ class Memory:
|
||||
last_accessed=datetime.fromisoformat(data["last_accessed"])
|
||||
if data.get("last_accessed")
|
||||
else None,
|
||||
entity_refs=data.get("entity_refs", []),
|
||||
# Normalized on load so rows written before #2947 was fixed heal
|
||||
# themselves instead of crashing search.
|
||||
entity_refs=normalize_entity_refs(data.get("entity_refs")),
|
||||
embedding=embedding,
|
||||
metadata=data.get("metadata", {}),
|
||||
)
|
||||
|
||||
@@ -0,0 +1,302 @@
|
||||
"""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"]
|
||||
Reference in New Issue
Block a user