Compare commits

...

2 Commits

Author SHA1 Message Date
Craig Estep
2692913429 test(memory): keep delete calls out of assert expressions
Run each storage.delete() on its own line and assert on the returned
count, so the deletion does not live inside an assert statement.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
2026-09-29 16:07:33 -06:00
Craig Estep
f414bbfcc4 fix(memory): escape scope and id values in LanceDB filters
LanceDBStorage built its where() filters by interpolating scope prefixes
and record ids into SQL strings. Route every interpolated value through a
quote-escaping helper, and match scope prefixes with LIKE ... ESCAPE '\'
so %, _ and \ in scope names match literally while the trailing prefix
wildcard is kept.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
2026-09-29 12:08:13 -06:00
2 changed files with 233 additions and 15 deletions

View File

@@ -40,6 +40,28 @@ _MAX_RETRIES = 5
_RETRY_BASE_DELAY = 0.2 # seconds; doubles on each retry
def _sql_str(value: str) -> str:
"""Escape a value for use inside a single-quoted SQL string literal."""
return value.replace("'", "''")
def _id_in_filter(record_ids: list[str]) -> str:
"""Build an ``id IN (...)`` predicate that matches each ID literally."""
ids_expr = ", ".join(f"'{_sql_str(str(rid))}'" for rid in record_ids)
return f"id IN ({ids_expr})"
def _scope_prefix_filter(prefix: str) -> str:
r"""Build a predicate matching scopes that start with ``prefix``.
``\``, ``%`` and ``_`` in the prefix are escaped so they match literally;
only the appended trailing ``%`` acts as a wildcard. Lance supports only
``\`` as the ``LIKE`` escape character.
"""
pattern = prefix.replace("\\", "\\\\").replace("%", "\\%").replace("_", "\\_")
return f"scope LIKE '{_sql_str(pattern)}%' ESCAPE '\\'"
class LanceDBStorage:
"""LanceDB-backed storage for the unified memory system."""
@@ -331,7 +353,7 @@ class LanceDBStorage:
)
with store_lock(self._lock_name):
self._ensure_table()
safe_id = str(record.id).replace("'", "''")
safe_id = _sql_str(str(record.id))
self._do_write("delete", f"id = '{safe_id}'")
row = self._record_to_row(record)
if row["vector"] is None or len(row["vector"]) != self._vector_dim:
@@ -352,11 +374,9 @@ class LanceDBStorage:
return
with store_lock(self._lock_name):
now = datetime.utcnow().isoformat()
safe_ids = [str(rid).replace("'", "''") for rid in record_ids]
ids_expr = ", ".join(f"'{rid}'" for rid in safe_ids)
self._do_write(
"update",
where=f"id IN ({ids_expr})",
where=_id_in_filter(record_ids),
values={"last_accessed": now},
)
@@ -364,7 +384,7 @@ class LanceDBStorage:
"""Return a single record by ID, or None if not found."""
if self._table is None:
return None
safe_id = str(record_id).replace("'", "''")
safe_id = _sql_str(str(record_id))
rows = self._table.search().where(f"id = '{safe_id}'").limit(1).to_list()
if not rows:
return None
@@ -387,9 +407,7 @@ class LanceDBStorage:
)
query = self._table.search(query_embedding)
if scope_prefix is not None and scope_prefix.strip("/"):
prefix = scope_prefix.rstrip("/")
like_val = prefix + "%"
query = query.where(f"scope LIKE '{like_val}'")
query = query.where(_scope_prefix_filter(scope_prefix.rstrip("/")))
results = query.limit(
limit * 3 if (categories or metadata_filter) else limit
).to_list()
@@ -423,8 +441,7 @@ class LanceDBStorage:
with store_lock(self._lock_name):
if record_ids and not (categories or metadata_filter):
before = int(self._table.count_rows())
ids_expr = ", ".join(f"'{rid}'" for rid in record_ids)
self._do_write("delete", f"id IN ({ids_expr})")
self._do_write("delete", _id_in_filter(record_ids))
return before - int(self._table.count_rows())
if categories or metadata_filter:
rows = self._scan_rows(scope_prefix)
@@ -445,15 +462,14 @@ class LanceDBStorage:
if not to_delete:
return 0
before = int(self._table.count_rows())
ids_expr = ", ".join(f"'{rid}'" for rid in to_delete)
self._do_write("delete", f"id IN ({ids_expr})")
self._do_write("delete", _id_in_filter(to_delete))
return before - int(self._table.count_rows())
conditions = []
if scope_prefix is not None and scope_prefix.strip("/"):
prefix = scope_prefix.rstrip("/")
if not prefix.startswith("/"):
prefix = "/" + prefix
conditions.append(f"scope LIKE '{prefix}%' OR scope = '/'")
conditions.append(f"{_scope_prefix_filter(prefix)} OR scope = '/'")
if older_than is not None:
conditions.append(f"created_at < '{older_than.isoformat()}'")
if not conditions:
@@ -487,7 +503,7 @@ class LanceDBStorage:
return []
q = self._table.search()
if scope_prefix is not None and scope_prefix.strip("/"):
q = q.where(f"scope LIKE '{scope_prefix.rstrip('/')}%'")
q = q.where(_scope_prefix_filter(scope_prefix.rstrip("/")))
if columns is not None:
q = q.select(columns)
result: list[dict[str, Any]] = q.limit(limit).to_list()
@@ -613,8 +629,9 @@ class LanceDBStorage:
return
prefix = scope_prefix.rstrip("/")
if prefix:
literal = _sql_str(prefix)
self._do_write(
"delete", f"scope >= '{prefix}' AND scope < '{prefix}/\uffff'"
"delete", f"scope >= '{literal}' AND scope < '{literal}/\uffff'"
)
def optimize(self) -> None:

View File

@@ -0,0 +1,201 @@
"""Scope and record-id values must be matched literally in LanceDB filters.
LanceDB ``where()`` takes a raw SQL expression, so every caller-supplied value
has to be escaped before it is embedded. These tests run against a real
temporary LanceDB table with two tenants and check that crafted scope prefixes
and record ids cannot read, delete, or reset another tenant's records, and that
legitimate scope names (including ``-``, ``_``, ``/``, ``%``, quotes, and
backslashes) keep working.
"""
from __future__ import annotations
from collections.abc import Callable
from pathlib import Path
from unittest.mock import MagicMock
from crewai.memory.storage.lancedb_storage import LanceDBStorage
from crewai.memory.types import MemoryRecord
from crewai.memory.unified_memory import Memory
import pytest
DIM = 4
TENANT_A = {"a1", "a2"}
TENANT_B = {"b1", "b2"}
CRAFTED_SCOPES = [
pytest.param("/tenant-a' OR scope <> '", id="quote-breakout"),
pytest.param("/tenant-a' OR scope LIKE '/tenant-b", id="predicate-injection"),
pytest.param("%", id="bare-percent"),
pytest.param("_", id="bare-underscore"),
pytest.param("/tenant-_", id="underscore-wildcard"),
]
def _record(
record_id: str, scope: str, categories: list[str] | None = None
) -> MemoryRecord:
return MemoryRecord(
id=record_id,
content=f"content of {record_id}",
scope=scope,
categories=categories or ["note"],
embedding=[0.1] * DIM,
)
def _new_storage(path: Path, records: list[MemoryRecord]) -> LanceDBStorage:
storage = LanceDBStorage(path=path, vector_dim=DIM, compact_every=0)
storage.save(records)
return storage
def _all_ids(storage: LanceDBStorage) -> set[str]:
return {r.id for r in storage.list_records(limit=1000)}
def _search_ids(storage: LanceDBStorage, scope_prefix: str) -> set[str]:
return {
r.id
for r, _ in storage.search([0.1] * DIM, scope_prefix=scope_prefix, limit=100)
}
@pytest.fixture
def storage(tmp_path: Path) -> LanceDBStorage:
return _new_storage(
tmp_path / "mem",
[
_record("a1", "/tenant-a/agent_1"),
_record("a2", "/tenant-a/notes"),
_record("b1", "/tenant-b/agent_1"),
_record("b2", "/tenant-b/secrets"),
],
)
def test_control_scope_prefix_stays_in_tenant(storage: LanceDBStorage) -> None:
assert _search_ids(storage, "/tenant-a") == TENANT_A
assert {r.id for r in storage.list_records("/tenant-a")} == TENANT_A
assert storage.count("/tenant-a") == 2
@pytest.mark.parametrize("payload", CRAFTED_SCOPES)
def test_crafted_scope_prefix_does_not_read_other_tenant(
storage: LanceDBStorage, payload: str
) -> None:
assert _search_ids(storage, payload) <= TENANT_A
assert {r.id for r in storage.list_records(payload)} <= TENANT_A
assert storage.count(payload) <= len(TENANT_A)
@pytest.mark.parametrize("payload", CRAFTED_SCOPES)
@pytest.mark.parametrize(
"remove",
[
pytest.param(lambda s, p: s.delete(scope_prefix=p), id="delete"),
pytest.param(
lambda s, p: s.delete(scope_prefix=p, categories=["note"]),
id="delete-by-category",
),
pytest.param(lambda s, p: s.reset(scope_prefix=p), id="reset"),
],
)
def test_crafted_scope_prefix_does_not_remove_other_tenant(
storage: LanceDBStorage,
payload: str,
remove: Callable[[LanceDBStorage, str], object],
) -> None:
remove(storage, payload)
assert _all_ids(storage) >= TENANT_B
def test_crafted_record_id_does_not_over_delete(storage: LanceDBStorage) -> None:
deleted = storage.delete(record_ids=["a1') OR id <> ('"])
assert deleted == 0
assert _all_ids(storage) == TENANT_A | TENANT_B
def test_record_ids_with_quotes_are_matched_literally(tmp_path: Path) -> None:
storage = _new_storage(
tmp_path / "mem",
[
_record("o'brien-1", "/people/o'brien", ["x"]),
_record("o'brien-2", "/people/o'brien", ["x"]),
_record("other", "/people/other", ["x"]),
],
)
deleted_by_id = storage.delete(record_ids=["o'brien-1"])
assert deleted_by_id == 1
deleted_by_scope = storage.delete(scope_prefix="/people/o'brien", categories=["x"])
assert deleted_by_scope == 1
assert _all_ids(storage) == {"other"}
def test_memory_root_scope_confines_crafted_explicit_scope(
storage: LanceDBStorage,
) -> None:
embedder = MagicMock(side_effect=lambda texts: [[0.1] * DIM for _ in texts])
memory = Memory(
storage=storage, root_scope="/tenant-a", llm=MagicMock(), embedder=embedder
)
for payload in ("x' OR scope <> '", "x' OR scope LIKE '/tenant-b"):
matches = memory.recall("anything", scope=payload, depth="shallow")
assert {m.record.id for m in matches} <= TENANT_A
memory.forget(scope=payload)
assert _all_ids(storage) == TENANT_A | TENANT_B
assert {m.record.id for m in memory.recall("anything", depth="shallow")} == TENANT_A
def test_legitimate_scope_characters_still_match(tmp_path: Path) -> None:
storage = _new_storage(
tmp_path / "mem",
[
_record("senior", "/crew/research-crew/agent/senior_analyst"),
_record("junior", "/crew/research-crew/agent/junior-analyst"),
_record("underscore", "/crew/research_crew"),
_record("lookalike", "/crew/researchXcrew"),
],
)
assert _search_ids(storage, "/crew/research-crew") == {"senior", "junior"}
assert _search_ids(storage, "/crew/research-crew/agent/senior_analyst") == {
"senior"
}
assert storage.get_scope_info("/crew/research-crew").record_count == 2
assert storage.list_scopes("/crew/research-crew") == ["/crew/research-crew/agent"]
assert storage.list_categories("/crew/research-crew") == {"note": 2}
deleted = storage.delete(scope_prefix="/crew/research_crew")
assert deleted == 1
assert _all_ids(storage) == {"senior", "junior", "lookalike"}
storage.reset(scope_prefix="/crew/research-crew")
assert _all_ids(storage) == {"lookalike"}
def test_percent_underscore_and_backslash_in_scope_names_match_literally(
tmp_path: Path,
) -> None:
storage = _new_storage(
tmp_path / "mem",
[
_record("pct", "/metrics/100%"),
_record("pct-child", "/metrics/100%/daily"),
_record("digits", "/metrics/1000"),
_record("underscore", "/a_b"),
_record("lookalike", "/axb"),
_record("backslash", "/share\\docs"),
],
)
assert _search_ids(storage, "/metrics/100%") == {"pct", "pct-child"}
assert {r.id for r in storage.list_records("/metrics/100%")} == {"pct", "pct-child"}
assert _search_ids(storage, "/a_b") == {"underscore"}
assert storage.count("/a_b") == 1
assert _search_ids(storage, "/share\\") == {"backslash"}
deleted = storage.delete(scope_prefix="/a_b")
assert deleted == 1
assert _all_ids(storage) == {"pct", "pct-child", "digits", "lookalike", "backslash"}