mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
fix(memory): enforce an atomic namespace storage limit
This commit is contained in:
parent
ada17a3430
commit
5022cc2960
4 changed files with 61 additions and 11 deletions
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue