diff --git a/deploy/memory-pilot/README.md b/deploy/memory-pilot/README.md index ab10d2f0e4a..fce7e39872b 100644 --- a/deploy/memory-pilot/README.md +++ b/deploy/memory-pilot/README.md @@ -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. diff --git a/litellm/proxy/memory/store.py b/litellm/proxy/memory/store.py index db3e61bd174..ec5b285113e 100644 --- a/litellm/proxy/memory/store.py +++ b/litellm/proxy/memory/store.py @@ -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) diff --git a/tests/test_litellm/proxy/memory/test_memory_v2_boundaries.py b/tests/test_litellm/proxy/memory/test_memory_v2_boundaries.py index e3d1793d11e..421a8105fbb 100644 --- a/tests/test_litellm/proxy/memory/test_memory_v2_boundaries.py +++ b/tests/test_litellm/proxy/memory/test_memory_v2_boundaries.py @@ -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() diff --git a/tests/test_litellm/proxy/memory/test_memory_v2_management.py b/tests/test_litellm/proxy/memory/test_memory_v2_management.py index b680b4a162e..558eae244c9 100644 --- a/tests/test_litellm/proxy/memory/test_memory_v2_management.py +++ b/tests/test_litellm/proxy/memory/test_memory_v2_management.py @@ -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,