fix(memory): enforce an atomic namespace storage limit

This commit is contained in:
moe-berri 2026-09-12 03:17:40 -07:00
parent ada17a3430
commit 5022cc2960
4 changed files with 61 additions and 11 deletions

View file

@ -86,6 +86,8 @@ correct, or delete entries in Memory; callers can use the self-service API.
- On gateway/backend deployments without shared Redis, first-time activation
can take up to 30 seconds to reach another process. Policy revocation is
checked against the primary database before memory operations.
- Each memory scope can hold up to 1,000 entries. Creation checks this limit
atomically; correction and deletion remain available when the scope is full.
- Stored references are untrusted data. They cannot grant API permissions or
change the namespace derived from authentication. Current user corrections
take precedence. Replacements require the current revision.

View file

@ -1,4 +1,6 @@
import json
from contextlib import AbstractAsyncContextManager
from types import SimpleNamespace
from typing import TYPE_CHECKING, Final
from fastapi import HTTPException
@ -9,9 +11,11 @@ from litellm.repositories.table_repositories import MemoryRepository
from litellm.types.memory_v2 import MemoryCapture, MemoryEntry, MemorySearch
if TYPE_CHECKING:
from prisma import Prisma
from prisma.models import LiteLLM_MemoryTable
_METADATA: Final = TypeAdapter(dict[str, object])
_MAX_NAMESPACE_ENTRIES: Final = 1000
def memory_entry(row: "LiteLLM_MemoryTable") -> MemoryEntry:
@ -157,17 +161,29 @@ class MemoryStore:
if capture.expected_revision is not None:
raise HTTPException(status_code=409, detail="Memory no longer exists")
try:
created: Final = await self.table.create(
data={ # mutable-ok: Prisma serializes these as native JSON containers.
**data,
"memory_id": memory_digest(namespace, capture.key),
"key": key,
"namespace": namespace,
"user_id": self.access.identity.user_id,
"team_id": self.access.identity.team_id,
"created_by": self.access.identity.user_id or self.access.identity.key_id,
}
)
manager: Final[AbstractAsyncContextManager[Prisma]] = self.prisma_client.db.tx()
async with manager as transaction:
lock_key: Final = int(memory_digest("memory-quota", namespace)[:16], 16) - (1 << 63)
await transaction.execute_raw("SELECT pg_advisory_xact_lock($1::bigint)", lock_key)
table: Final = MemoryRepository(SimpleNamespace(db=transaction)).table
entries: Final = await table.count(
where={"namespace": namespace} # mutable-ok: Prisma accepts native query containers.
)
if entries >= _MAX_NAMESPACE_ENTRIES:
raise HTTPException(
status_code=429, detail="Memory scope has reached 1000 entries; delete unused memories first"
)
created: Final = await table.create(
data={ # mutable-ok: Prisma serializes these as native JSON containers.
**data,
"memory_id": memory_digest(namespace, capture.key),
"key": key,
"namespace": namespace,
"user_id": self.access.identity.user_id,
"team_id": self.access.identity.team_id,
"created_by": self.access.identity.user_id or self.access.identity.key_id,
}
)
except UniqueViolationError as exc:
raise HTTPException(status_code=409, detail="Memory changed; read it again before replacing it") from exc
return memory_entry(created)

View file

@ -39,6 +39,9 @@ def prisma_edge() -> MagicMock:
table.find_first = AsyncMock(return_value=None)
table.find_many = AsyncMock(return_value=[])
table.create = AsyncMock()
table.count = AsyncMock(return_value=0)
client.db.tx.return_value.__aenter__.return_value = client.db
client.db.execute_raw = AsyncMock()
table.update_many = AsyncMock(return_value=1)
table.delete_many = AsyncMock(return_value=1)
return client
@ -441,3 +444,29 @@ async def test_redis_circuit_breaker_falls_back_to_primary_configuration(prisma_
await invalidate_memory_configuration()
assert prisma_edge.db.litellm_memorypolicy.find_many.await_count == 2
redis.async_get_cache.assert_awaited_once()
@pytest.mark.asyncio
async def test_full_scope_blocks_creation_but_permits_correction_and_reclaimed_capacity(prisma_edge: MagicMock) -> None:
table = prisma_edge.db.litellm_memorytable
table.count.return_value = 1000
with pytest.raises(HTTPException) as full:
await store(prisma_edge).capture(_CAPTURE)
assert full.value.status_code == 429
table.create.assert_not_awaited()
assert table.count.call_args.kwargs["where"] == {"namespace": _IDENTITY.namespace("key")}
prisma_edge.db.execute_raw.assert_awaited_once()
assert "pg_advisory_xact_lock" in prisma_edge.db.execute_raw.call_args.args[0]
table.find_unique.return_value = row()
table.find_first.return_value = row(value="Corrected", updated_at=_NOW + timedelta(seconds=1))
corrected = await store(prisma_edge).capture(
_CAPTURE.model_copy(update={"content": "Corrected", "expected_revision": _NOW})
)
assert corrected.content == "Corrected"
table.count.assert_awaited_once()
assert await store(prisma_edge).delete("entry")
table.find_unique.return_value = None
table.count.return_value = 999
table.create.return_value = row()
assert (await store(prisma_edge).capture(_CAPTURE)).memory_id == "entry"
table.create.assert_awaited_once()

View file

@ -36,6 +36,9 @@ def database() -> Iterator[MagicMock]:
table.delete = AsyncMock()
table.delete_many = AsyncMock(return_value=0)
table.create = AsyncMock()
table.count = AsyncMock(return_value=0)
client.db.tx.return_value.__aenter__.return_value = client.db
client.db.execute_raw = AsyncMock()
client.db.litellm_teamtable.find_unique.return_value = {
"team_id": "team",
"organization_id": None,