From 5439a2acb74a1eba8347e74c1f83f71de4d2a66a Mon Sep 17 00:00:00 2001 From: moe-berri Date: Mon, 14 Sep 2026 17:16:23 -0700 Subject: [PATCH] feat(memory): simplify admin activation and reuse team record permissions --- .../migration.sql | 1 + .../migration.sql | 5 + .../litellm_proxy_extras/schema.prisma | 6 + .../prompt_templates/server_tools.py | 8 +- .../messages/agentic_streaming_iterator.py | 45 +- litellm/proxy/_types.py | 9 +- .../proxy/common_utils/user_api_key_cache.py | 3 +- .../internal_user_endpoints.py | 12 +- .../management_helpers/record_permissions.py | 15 + litellm/proxy/memory/continuation.py | 21 +- litellm/proxy/memory/gateway.py | 40 +- litellm/proxy/memory/knowledge.py | 43 +- litellm/proxy/memory/management.py | 361 +++------ litellm/proxy/memory/memory_endpoints.py | 12 +- litellm/proxy/memory/policy.py | 209 +++--- litellm/proxy/memory/store.py | 244 ++++-- litellm/proxy/memory/transport.py | 2 +- litellm/proxy/schema.prisma | 6 + .../spend_management_endpoints.py | 32 +- litellm/repositories/table_repositories.py | 8 - litellm/repositories/unit_of_work.py | 9 +- litellm/types/memory_v2.py | 50 +- schema.prisma | 6 + tests/e2e/coverage_registry/mgmt.yaml | 8 +- tests/e2e/management/memory_client.py | 105 ++- tests/e2e/management/test_memory_v2_e2e.py | 125 ++-- tests/e2e/models.py | 39 +- .../proxy/memory/test_memory_v2_boundaries.py | 209 +++--- .../proxy/memory/test_memory_v2_management.py | 707 ++++++++---------- .../proxy/memory/test_memory_v2_protocols.py | 10 +- .../_components/AutomaticMemoryEntries.tsx | 277 ++++--- .../memory/_components/MemoryEntriesTable.tsx | 10 +- .../memory/_components/MemorySettings.tsx | 391 ++++------ .../memory/_components/MemoryTargetPicker.tsx | 117 +-- .../memory/page.integration.test.tsx | 343 ++++----- .../src/app/(dashboard)/memory/page.tsx | 39 +- .../team/permission_definitions.tsx | 4 +- ui/litellm-dashboard/src/lib/http/schema.d.ts | 366 ++++----- 38 files changed, 1697 insertions(+), 2200 deletions(-) create mode 100644 litellm-proxy-extras/litellm_proxy_extras/migrations/20260914223000_memory_gateway_pause/migration.sql create mode 100644 litellm-proxy-extras/litellm_proxy_extras/migrations/20260915001500_memory_owner_visibility/migration.sql create mode 100644 litellm/proxy/management_helpers/record_permissions.py diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260914223000_memory_gateway_pause/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260914223000_memory_gateway_pause/migration.sql new file mode 100644 index 00000000000..e322abcabb8 --- /dev/null +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260914223000_memory_gateway_pause/migration.sql @@ -0,0 +1 @@ +ALTER TABLE "LiteLLM_MemoryPolicy" ADD COLUMN "paused" BOOLEAN NOT NULL DEFAULT false; diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260915001500_memory_owner_visibility/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260915001500_memory_owner_visibility/migration.sql new file mode 100644 index 00000000000..0d4c2961c55 --- /dev/null +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260915001500_memory_owner_visibility/migration.sql @@ -0,0 +1,5 @@ +ALTER TABLE "LiteLLM_MemoryTable" ADD COLUMN IF NOT EXISTS "organization_id" TEXT; +ALTER TABLE "LiteLLM_MemoryTable" ADD COLUMN IF NOT EXISTS "owner_key_id" TEXT; +CREATE INDEX IF NOT EXISTS "LiteLLM_MemoryTable_organization_id_user_id_updated_at_idx" ON "LiteLLM_MemoryTable"("organization_id", "user_id", "updated_at"); +CREATE INDEX IF NOT EXISTS "LiteLLM_MemoryTable_team_id_updated_at_idx" ON "LiteLLM_MemoryTable"("team_id", "updated_at"); +CREATE INDEX IF NOT EXISTS "LiteLLM_MemoryTable_owner_key_id_idx" ON "LiteLLM_MemoryTable"("owner_key_id"); diff --git a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma index 8aaf4b48bbb..e2c293d939d 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma +++ b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma @@ -1436,6 +1436,8 @@ model LiteLLM_MemoryTable { metadata Json? user_id String? team_id String? + organization_id String? + owner_key_id String? created_at DateTime @default(now()) created_by String? updated_at DateTime @default(now()) @updatedAt @@ -1444,6 +1446,9 @@ model LiteLLM_MemoryTable { @@index([user_id]) @@index([team_id]) @@index([namespace, updated_at]) + @@index([organization_id, user_id, updated_at]) + @@index([team_id, updated_at]) + @@index([owner_key_id]) } model LiteLLM_MemoryPolicy { @@ -1452,6 +1457,7 @@ model LiteLLM_MemoryPolicy { target_id String activation String scope String + paused Boolean @default(false) created_at DateTime @default(now()) updated_at DateTime @default(now()) @updatedAt updated_by String diff --git a/litellm/litellm_core_utils/prompt_templates/server_tools.py b/litellm/litellm_core_utils/prompt_templates/server_tools.py index 392b036cf87..b9cfb8daf95 100644 --- a/litellm/litellm_core_utils/prompt_templates/server_tools.py +++ b/litellm/litellm_core_utils/prompt_templates/server_tools.py @@ -153,11 +153,15 @@ def trailing_system_messages(data: Mapping[str, object], route: ServerToolRoute) def uncached_system_directive(message: Mapping[str, object]) -> Mapping[str, object]: return { # mutable-ok: Provider wire format requires native JSON containers. - **{key: value for key, value in message.items() if key != "cache_control"}, + **{ # mutable-ok: Provider requires native JSON. + key: value for key, value in message.items() if key != "cache_control" + }, # mutable-ok: Provider requires native JSON. **( { # mutable-ok: Provider wire format requires native JSON containers. "content": [ # mutable-ok: Provider wire format requires native JSON containers. - {key: value for key, value in _OBJECT.validate_python(block).items() if key != "cache_control"} + { # mutable-ok: Provider requires native JSON. + key: value for key, value in _OBJECT.validate_python(block).items() if key != "cache_control" + } # mutable-ok: Provider requires native JSON. for block in _items(message.get("content")) ], } diff --git a/litellm/llms/anthropic/experimental_pass_through/messages/agentic_streaming_iterator.py b/litellm/llms/anthropic/experimental_pass_through/messages/agentic_streaming_iterator.py index bb7ecd09cb7..e401c7dacd5 100644 --- a/litellm/llms/anthropic/experimental_pass_through/messages/agentic_streaming_iterator.py +++ b/litellm/llms/anthropic/experimental_pass_through/messages/agentic_streaming_iterator.py @@ -54,7 +54,7 @@ def _parse_sse_events(raw: bytes) -> list[tuple]: """Return a list of (event_type, parsed_data_dict) from raw SSE bytes.""" text: Final = raw.decode("utf-8", errors="replace") lines: Final = text.split("\n") - events: Final[list[tuple]] = [] + events: Final[list[tuple]] = [] # mutable-ok: Provider requires native JSON. current_event_type: str | None = None for line in lines: @@ -76,11 +76,11 @@ def _parse_sse_events(raw: bytes) -> list[tuple]: def _handle_message_start(data: dict, response: dict) -> None: - msg: Final = data.get("message", {}) + msg: Final = data.get("message", {}) # mutable-ok: Provider requires native JSON. response["id"] = msg.get("id", response["id"]) response["model"] = msg.get("model", response["model"]) response["role"] = msg.get("role", response["role"]) - usage: Final = msg.get("usage", {}) + usage: Final = msg.get("usage", {}) # mutable-ok: Provider requires native JSON. if usage: response["usage"]["input_tokens"] = usage.get("input_tokens", 0) for key in ("cache_creation_input_tokens", "cache_read_input_tokens"): @@ -90,38 +90,41 @@ def _handle_message_start(data: dict, response: dict) -> None: def _handle_content_block_start(data: dict, content_blocks: dict[int, dict]) -> None: idx: Final = data.get("index", len(content_blocks)) - block: Final = data.get("content_block", {}) + block: Final = data.get("content_block", {}) # mutable-ok: Provider requires native JSON. block_type: Final = block.get("type", "text") - _BLOCK_TEMPLATES: Final[dict[str, dict]] = { - "text": {"type": "text", "text": block.get("text", "")}, - "thinking": { + _BLOCK_TEMPLATES: Final[dict[str, dict]] = { # mutable-ok: Provider requires native JSON. + "text": { # mutable-ok: Provider requires native JSON. + "type": "text", + "text": block.get("text", ""), + }, # mutable-ok: Accumulate Anthropic stream deltas. + "thinking": { # mutable-ok: Accumulate Anthropic stream deltas. "type": "thinking", "thinking": block.get("thinking", ""), "signature": block.get("signature", ""), }, - "redacted_thinking": { + "redacted_thinking": { # mutable-ok: Provider requires native JSON. "type": "redacted_thinking", "data": block.get("data", ""), }, } if block_type == "tool_use": - content_blocks[idx] = { + content_blocks[idx] = { # mutable-ok: Accumulate Anthropic stream deltas. "type": "tool_use", "id": block.get("id", ""), "name": block.get("name", ""), - "input": block.get("input", {}), + "input": block.get("input", {}), # mutable-ok: Provider requires native JSON. "_partial_json": "", } elif block_type in _BLOCK_TEMPLATES: - content_blocks[idx] = dict(_BLOCK_TEMPLATES[block_type]) + content_blocks[idx] = dict(_BLOCK_TEMPLATES[block_type]) # mutable-ok: Provider requires native JSON. else: - content_blocks[idx] = dict(block) + content_blocks[idx] = dict(block) # mutable-ok: Provider requires native JSON. def _handle_content_block_delta(data: dict, content_blocks: dict[int, dict]) -> None: idx: Final = data.get("index", 0) - delta: Final = data.get("delta", {}) + delta: Final = data.get("delta", {}) # mutable-ok: Provider requires native JSON. delta_type: Final = delta.get("type", "") block: Final = content_blocks.get(idx) if block is None: @@ -146,16 +149,16 @@ def _handle_content_block_stop(data: dict, content_blocks: dict[int, dict]) -> N try: block["input"] = json.loads(partial) except (json.JSONDecodeError, ValueError): - block["input"] = {"_raw": partial} + block["input"] = {"_raw": partial} # mutable-ok: Provider requires native JSON. def _handle_message_delta(data: dict, response: dict) -> None: - delta: Final = data.get("delta", {}) + delta: Final = data.get("delta", {}) # mutable-ok: Provider requires native JSON. if "stop_reason" in delta: response["stop_reason"] = delta["stop_reason"] if "stop_sequence" in delta: response["stop_sequence"] = delta["stop_sequence"] - usage: Final = data.get("usage", {}) + usage: Final = data.get("usage", {}) # mutable-ok: Provider requires native JSON. if usage.get("output_tokens") is not None: response["usage"]["output_tokens"] = usage["output_tokens"] for key in ( @@ -208,7 +211,7 @@ class AgenticAnthropicStreamingIterator: self._server_fulfilled_tool_names = server_fulfilled_tool_names self._ping_interval_seconds = ping_interval_seconds - self._collected_bytes: list[bytes] = [] + self._collected_bytes: list[bytes] = [] # mutable-ok: Provider requires native JSON. self._stream_exhausted = False self._hook_processing_done = False self._follow_up_iterator: AsyncIterator | None = None @@ -420,17 +423,17 @@ class AgenticAnthropicStreamingIterator: """ events: Final = _parse_sse_events(b"".join(raw_bytes)) - response: Final[dict[str, Any]] = { + response: Final[dict[str, Any]] = { # mutable-ok: Provider requires native JSON. "id": "", "type": "message", "role": "assistant", "model": "", - "content": [], + "content": [], # mutable-ok: Provider requires native JSON. "stop_reason": None, "stop_sequence": None, - "usage": {"input_tokens": 0, "output_tokens": 0}, + "usage": {"input_tokens": 0, "output_tokens": 0}, # mutable-ok: Provider requires native JSON. } - content_blocks: Final[dict[int, dict[str, Any]]] = {} + content_blocks: Final[dict[int, dict[str, Any]]] = {} # mutable-ok: Provider requires native JSON. saw_message_start = False for event_type, data in events: diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 8af2a82a868..6b199e8b87e 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -300,6 +300,7 @@ class KeyManagementRoutes(str, enum.Enum): # team spend-log viewing SPEND_LOGS = "/spend/logs" SPEND_LOGS_V2 = "/spend/logs/v2" + MEMORY_READ = "/v2/memory/entries" class LiteLLMRoutes(enum.Enum): @@ -647,6 +648,7 @@ class LiteLLMRoutes(enum.Enum): KeyManagementRoutes.TEAM_DAILY_ACTIVITY_AGGREGATED.value, KeyManagementRoutes.SPEND_LOGS.value, KeyManagementRoutes.SPEND_LOGS_V2.value, + KeyManagementRoutes.MEMORY_READ.value, KeyManagementRoutes.KEY_RESET_SPEND.value, KeyManagementRoutes.KEY_ALIASES.value, KeyManagementRoutes.KEY_ACCESS_GROUP_ASSIGNMENT.value, @@ -836,9 +838,7 @@ class LiteLLMRoutes(enum.Enum): ) self_managed_routes = [ - "/v2/memory/policies", - "/v2/memory/policies/{policy_id}", - "/v2/memory/preference", + "/v2/memory/settings", "/v2/memory/status", "/v2/memory/entries", "/v2/memory/entries/{memory_id}", @@ -929,8 +929,7 @@ class LiteLLMRoutes(enum.Enum): # updating this list — the default-allow behavior covers it automatically. admin_viewer_routes = ( [ - "/v2/memory/policies", - "/v2/memory/preference", + "/v2/memory/settings", "/v2/memory/status", "/v2/memory/entries", "/user/list", diff --git a/litellm/proxy/common_utils/user_api_key_cache.py b/litellm/proxy/common_utils/user_api_key_cache.py index 1c7a379897f..098690c9a3f 100644 --- a/litellm/proxy/common_utils/user_api_key_cache.py +++ b/litellm/proxy/common_utils/user_api_key_cache.py @@ -188,8 +188,7 @@ class UserApiKeyCache(DualCache): decoded: Final = CacheCodec.deserialize(cached, model_type=model_type) if decoded is None: verbose_proxy_logger.error( - "UserApiKeyCache.async_get_cache failed to deserialize cached value for key=%r model_type=%s", - key, + "UserApiKeyCache.async_get_cache failed to deserialize cached value for model_type=%s", getattr(model_type, "__name__", str(model_type)), ) return None diff --git a/litellm/proxy/management_endpoints/internal_user_endpoints.py b/litellm/proxy/management_endpoints/internal_user_endpoints.py index e3efda507f6..ca7cef28669 100644 --- a/litellm/proxy/management_endpoints/internal_user_endpoints.py +++ b/litellm/proxy/management_endpoints/internal_user_endpoints.py @@ -350,9 +350,9 @@ async def _add_user_to_team( verbose_proxy_logger.error( "litellm.proxy.management_endpoints.internal_user_endpoints._add_user_to_team(): " "failed to add user %s to team %s - %s", - user_id, - team_id, - str(e), + user_id.replace("\r", "").replace("\n", ""), + team_id.replace("\r", "").replace("\n", ""), + str(e).replace("\r", "").replace("\n", ""), ) except Exception as e: if ( @@ -369,9 +369,9 @@ async def _add_user_to_team( verbose_proxy_logger.error( "litellm.proxy.management_endpoints.internal_user_endpoints._add_user_to_team(): " "failed to add user %s to team %s - %s", - user_id, - team_id, - str(e), + user_id.replace("\r", "").replace("\n", ""), + team_id.replace("\r", "").replace("\n", ""), + str(e).replace("\r", "").replace("\n", ""), ) raise e diff --git a/litellm/proxy/management_helpers/record_permissions.py b/litellm/proxy/management_helpers/record_permissions.py new file mode 100644 index 00000000000..4483bb72952 --- /dev/null +++ b/litellm/proxy/management_helpers/record_permissions.py @@ -0,0 +1,15 @@ +from collections.abc import Sequence + +from litellm.proxy._types import KeyManagementRoutes, LiteLLM_TeamTable, UserAPIKeyAuth + + +def can_read_team_records(auth: UserAPIKeyAuth, team: LiteLLM_TeamTable, permission: KeyManagementRoutes) -> bool: + from litellm.proxy.management_endpoints.common_utils import _is_user_team_admin, _team_member_has_permission + + return _is_user_team_admin(auth, team) or _team_member_has_permission(auth, team, permission.value) + + +def permitted_record_teams( + auth: UserAPIKeyAuth, teams: Sequence[LiteLLM_TeamTable], permission: KeyManagementRoutes +) -> tuple[str, ...]: + return tuple(team.team_id for team in teams if can_read_team_records(auth, team, permission)) diff --git a/litellm/proxy/memory/continuation.py b/litellm/proxy/memory/continuation.py index 8cd6e0ea49e..6d3195fa853 100644 --- a/litellm/proxy/memory/continuation.py +++ b/litellm/proxy/memory/continuation.py @@ -43,6 +43,7 @@ class MemoryContinuation(BaseModel): upstream_ids: tuple[str, ...] = () pending_results: tuple[Mapping[str, object], ...] = () transcript_anchor: str | None = None + permission_revision: str | None = None def _empty_array(value: object) -> bool: @@ -143,6 +144,12 @@ class MemoryContinuations: self.route: Final[ServerToolRoute] = route self.table = MemoryContinuationRepository(store.prisma_client).table + def validate_patch(self, payload: object) -> MemoryContinuation: + patch: Final = MemoryContinuation.model_validate(payload) + if patch.permission_revision != self.store.access.permission_revision: + raise HTTPException(status_code=403, detail="Memory permissions changed; start a new conversation") + return patch + def identifier(self, anchor: str) -> str: return memory_digest( self.store.access.namespace, @@ -168,7 +175,7 @@ class MemoryContinuations: }, } ) - patches: Final = MappingProxyType({row.id: MemoryContinuation.model_validate(row.payload) for row in rows}) + patches: Final = MappingProxyType({row.id: self.validate_patch(row.payload) for row in rows}) def apply(result: tuple[Mapping[str, object], ...], index: int) -> tuple[Mapping[str, object], ...]: patch: Final = patches.get(self.identifier(anchors[index])) @@ -207,14 +214,22 @@ class MemoryContinuations: }, } ) - return MemoryContinuation.model_validate(row.payload) if row is not None else None + return self.validate_patch(row.payload) if row is not None else None async def save(self, anchor: str, patch: MemoryContinuation) -> None: await self.save_many(((anchor, patch),)) async def save_many(self, patches: tuple[tuple[str, MemoryContinuation], ...]) -> None: namespace: Final = await self.store.authorize_namespace() - payloads: Final = tuple((self.identifier(anchor), patch.model_dump_json()) for anchor, patch in patches) + payloads: Final = tuple( + ( + self.identifier(anchor), + patch.model_copy( + update=MappingProxyType({"permission_revision": self.store.access.permission_revision}) + ).model_dump_json(), + ) + for anchor, patch in patches + ) if any(len(payload.encode()) > _MAX_PATCH_BYTES for _, payload in payloads): raise HTTPException(status_code=413, detail="Memory continuation exceeds one megabyte") key_id: Final = self.store.access.identity.key_id or self.store.access.identity.user_id or "" diff --git a/litellm/proxy/memory/gateway.py b/litellm/proxy/memory/gateway.py index 29ae52d0118..deb508c8c26 100644 --- a/litellm/proxy/memory/gateway.py +++ b/litellm/proxy/memory/gateway.py @@ -11,6 +11,7 @@ from pydantic import TypeAdapter from starlette.responses import JSONResponse, Response from starlette.types import ASGIApp +from litellm._logging import verbose_proxy_logger from litellm.litellm_core_utils.prompt_templates.server_tool_responses import ( executable_server_calls, object_value, @@ -342,6 +343,20 @@ class GatewayMemoryLoop: yield chunk +def validate_memory_request(data: Mapping[str, object], request: Request) -> None: + if request.url.path.startswith("/cursor/"): + raise HTTPException( + status_code=400, + detail="Gateway memory requires a standard /v1/chat/completions, /v1/messages, or /v1/responses endpoint", + ) + if data.get("functions") is not None or data.get("function_call") is not None: + raise HTTPException( + status_code=400, detail="Gateway memory requires tools and tool_choice instead of legacy functions" + ) + if data.get("background") is True or data.get("n", 1) != 1: + raise HTTPException(status_code=400, detail="Gateway memory requires a foreground request with one completion") + + async def process_gateway_memory( data: Mapping[str, object], request: Request, auth: UserAPIKeyAuth, route: str ) -> Response | None: @@ -355,18 +370,11 @@ async def process_gateway_memory( return None store: Final = await gateway_memory_store(auth) if store is None: + previous: Final = data.get("previous_response_id") + if isinstance(previous, str) and previous.startswith("resp_litellm_memory_"): + raise HTTPException(status_code=404, detail="Memory response not found or expired") return None - if request.url.path.startswith("/cursor/"): - raise HTTPException( - status_code=400, - detail="Gateway memory requires a standard /v1/chat/completions, /v1/messages, or /v1/responses endpoint", - ) - if data.get("functions") is not None or data.get("function_call") is not None: - raise HTTPException( - status_code=400, detail="Gateway memory requires tools and tool_choice instead of legacy functions" - ) - if data.get("background") is True or data.get("n", 1) != 1: - raise HTTPException(status_code=400, detail="Gateway memory requires a foreground request with one completion") + validate_memory_request(data, request) from litellm.proxy.proxy_server import app, llm_router loop: Final = GatewayMemoryLoop(app, request, data, route, store) @@ -422,7 +430,11 @@ async def gateway_memory_store(auth: UserAPIKeyAuth) -> MemoryStore | None: identity: Final = MemoryIdentity.from_auth(auth) if not identity.user_id and not identity.key_id: return None - if not await gateway_memory_is_configured(prisma_client, user_api_key_cache): + try: + if not await gateway_memory_is_configured(prisma_client, user_api_key_cache): + return None + access: Final = await resolve_memory_access(prisma_client, identity) + return MemoryStore(prisma_client, access) if access.active else None + except Exception: + verbose_proxy_logger.warning("Memory access is unavailable; continuing without automatic memory") return None - access: Final = await resolve_memory_access(prisma_client, identity) - return MemoryStore(prisma_client, access) if access.active else None diff --git a/litellm/proxy/memory/knowledge.py b/litellm/proxy/memory/knowledge.py index 64ab8b64b9a..2cf265afdcf 100644 --- a/litellm/proxy/memory/knowledge.py +++ b/litellm/proxy/memory/knowledge.py @@ -1,13 +1,13 @@ -import asyncio from collections.abc import Mapping from dataclasses import dataclass +from types import MappingProxyType from typing import Final from fastapi import HTTPException from pydantic import ValidationError from litellm.litellm_core_utils.prompt_templates.factory import NormalizedToolCall -from litellm.proxy.memory.content import fuzzy_memories, redact_memory +from litellm.proxy.memory.content import redact_memory from litellm.proxy.memory.policy import memory_digest from litellm.proxy.memory.store import MemoryStore from litellm.types.memory_v2 import ( @@ -79,6 +79,8 @@ def _preview(entry: MemoryEntry) -> Mapping[str, object]: "title": entry.title, "when_to_use": entry.when_to_use, "scope": entry.scope, + "contributed_by": entry.actor, + "team_id": entry.team_id, } @@ -87,14 +89,14 @@ def _revision(entries: tuple[MemoryEntry, ...]) -> str: async def memory_catalog(store: MemoryStore, request: MemoryCatalogRequest) -> Mapping[str, object]: - entries: Final = await store.entries() + entries, total, revision = await store.catalog(request) end: Final = request.offset + request.limit return { # mutable-ok: Tool results are JSON objects. - "revision": _revision(entries), - "total": len(entries), - "next_offset": end if end < len(entries) else None, + "revision": revision, + "total": total, + "next_offset": end if end < total else None, "observations": [ # mutable-ok: Native provider JSON containers. - _preview(entry) for entry in entries[request.offset : end] + _preview(entry) for entry in entries ], # mutable-ok: Tool results are JSON. } @@ -108,17 +110,11 @@ async def execute_memory_tool(store: MemoryStore, call: NormalizedToolCall, chec ) case "litellm_memory_search": query: Final = MemoryRecallRequest.model_validate(call["arguments"]) - entries: Final = await store.entries() - candidates: Final = tuple( - entry - for entry in entries - if query.scope is None or query.scope.casefold() in entry.scope.casefold() - ) - ranked: Final = await asyncio.to_thread(fuzzy_memories, query.query, candidates) + ranked, total_matches = await store.recall(query) return MemoryToolResult( { # mutable-ok: Tool results are JSON objects. - "revision": _revision(entries), - "total_matches": len(ranked), + "revision": _revision(tuple(entry for entry, _, _ in ranked)), + "total_matches": total_matches, "results": [ # mutable-ok: Tool results are JSON arrays. { # mutable-ok: Tool results are JSON objects. **_preview(entry), @@ -181,11 +177,7 @@ async def execute_memory_tool(store: MemoryStore, call: NormalizedToolCall, chec reflected=batch.checkpoint == checkpoint, ) case _: - return MemoryToolResult( - { # mutable-ok: Native provider JSON containers. - "error": "Unknown memory tool" - } - ) + pass except ValidationError: return MemoryToolResult( { # mutable-ok: Native provider JSON containers. @@ -199,3 +191,12 @@ async def execute_memory_tool(store: MemoryStore, call: NormalizedToolCall, chec "status": exc.status_code, } ) + except Exception: + return MemoryToolResult( + MappingProxyType({"error": "Memory is temporarily unavailable. Continue the task without memory."}) + ) + return MemoryToolResult( + { # mutable-ok: Native provider JSON containers. + "error": "Unknown memory tool" + } + ) diff --git a/litellm/proxy/memory/management.py b/litellm/proxy/memory/management.py index 6f1ce72671a..49015bfb324 100644 --- a/litellm/proxy/memory/management.py +++ b/litellm/proxy/memory/management.py @@ -4,320 +4,155 @@ from typing import Annotated, Final from fastapi import APIRouter, Depends, HTTPException, Query, Response from pydantic import AwareDatetime -from litellm.proxy._types import UI_TEAM_ID, LitellmUserRoles, UserAPIKeyAuth, user_api_key_has_admin_view +from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth, user_api_key_has_admin_view from litellm.proxy.auth.user_api_key_auth import user_api_key_auth -from litellm.proxy.memory.memory_endpoints import is_memory_team_admin, require_memory_prisma +from litellm.proxy.memory.memory_endpoints import require_memory_prisma from litellm.proxy.memory.policy import ( - MemoryAccess, + MEMORY_CONFIG_PARAM, MemoryIdentity, invalidate_memory_configuration, - memory_digest, memory_primary_client, + memory_settings, resolve_memory_access, ) from litellm.proxy.memory.store import MemoryStore -from litellm.repositories.organization_repository import OrganizationRepository -from litellm.repositories.project_repository import ProjectRepository -from litellm.repositories.table_repositories import ( - MemoryPolicyRepository, - MemoryPreferenceRepository, - OrganizationMembershipRepository, -) +from litellm.repositories.config_repository import ConfigRepository from litellm.repositories.team_repository import TeamRepository from litellm.repositories.user_repository import UserRepository -from litellm.repositories.verification_token_repository import VerificationTokenRepository -from litellm.types.memory_v2 import ( - MemoryCapture, - MemoryEntry, - MemoryPolicy, - MemoryPolicyInput, - MemoryPreference, - MemoryQuery, - MemorySearch, - MemoryStatus, - MemoryTarget, -) +from litellm.types.memory_v2 import MemoryCapture, MemoryEntry, MemoryQuery, MemorySearch, MemorySettings, MemoryStatus _AUTH: Final = Depends(user_api_key_auth) - -router: Final = APIRouter( - prefix="/v2/memory", - tags=[ # mutable-ok: Prisma serializes these as native JSON containers. - "memory management" - ], -) +router: Final = APIRouter(prefix="/v2/memory", tags=["memory management"]) # mutable-ok: FastAPI requires native tags. -async def require_policy_admin( - auth: UserAPIKeyAuth, target_type: MemoryTarget, target_id: str, *, write: bool = True -) -> None: +def require_memory_admin(auth: UserAPIKeyAuth, *, write: bool = False) -> None: + if not user_api_key_has_admin_view(auth) or write and auth.user_role != LitellmUserRoles.PROXY_ADMIN: + raise HTTPException(status_code=403, detail="Only proxy administrators can configure gateway memory") + + +@router.get("/settings", response_model=MemorySettings) +async def get_settings(auth: UserAPIKeyAuth = _AUTH) -> MemorySettings: + require_memory_admin(auth) + return await memory_settings(require_memory_prisma()) + + +@router.put("/settings", response_model=MemorySettings) +async def set_settings(settings: MemorySettings, auth: UserAPIKeyAuth = _AUTH) -> MemorySettings: + require_memory_admin(auth, write=True) prisma: Final = memory_primary_client(require_memory_prisma()) - if write and auth.user_role in (LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY, LitellmUserRoles.INTERNAL_USER_VIEW_ONLY): - raise HTTPException(status_code=403, detail="Memory policies require administrator write access") - proxy_admin: Final = ( - auth.user_role == LitellmUserRoles.PROXY_ADMIN or not write and user_api_key_has_admin_view(auth) - ) - if target_type == "gateway": - if not proxy_admin or target_id != "*": - raise HTTPException(status_code=403, detail="Gateway memory policies require a proxy administrator") - return - if target_type == "organization": - organization: Final = await OrganizationRepository(prisma).find_by_id(target_id) - membership: Final = await OrganizationMembershipRepository(prisma).table.find_first( - where={ # mutable-ok: Prisma serializes these as native JSON containers. - "organization_id": target_id, - "user_id": auth.user_id or "", - "user_role": "org_admin", - } + selected: Final = tuple(sorted(frozenset(settings.user_ids))) if not settings.everyone else () + if settings.enabled and not settings.everyone and not selected: + raise HTTPException(status_code=422, detail="Select at least one user or enable memory for everyone") + if selected: + users: Final = await UserRepository(prisma).table.find_many( + where={"user_id": {"in": list(selected)}}, # mutable-ok: Prisma requires native query JSON. + take=len(selected), ) - if organization and (proxy_admin or membership): - return - if target_type == "team": - team: Final = await TeamRepository(prisma).find_by_id(target_id) - if team and (proxy_admin or await is_memory_team_admin(prisma, auth, target_id)): - return - if target_type == "project": - project: Final = await ProjectRepository(prisma).find_by_id(target_id) - if project and (proxy_admin or project.team_id and await is_memory_team_admin(prisma, auth, project.team_id)): - return - if target_type == "key": - key: Final = await VerificationTokenRepository(prisma).find_by_id(target_id, id_field="token") - if key and (proxy_admin or key.team_id and await is_memory_team_admin(prisma, auth, key.team_id)): - return - if target_type == "user" and proxy_admin and await UserRepository(prisma).find_by_id(target_id): - return - raise HTTPException(status_code=403, detail="You cannot administer memory for this target") - - -@router.get("/policies", response_model=list[MemoryPolicy]) -async def list_policies( - target_type: MemoryTarget | None = None, - target_id: str | None = None, - offset: int = Query(0, ge=0), - auth: UserAPIKeyAuth = _AUTH, -) -> list[MemoryPolicy]: - if target_type is not None and target_id is not None: - await require_policy_admin(auth, target_type, target_id, write=False) - elif not user_api_key_has_admin_view(auth): - raise HTTPException(status_code=403, detail="Select a target you administer") - elif target_type is not None or target_id is not None: - raise HTTPException(status_code=400, detail="Provide both target_type and target_id") - rows: Final = await MemoryPolicyRepository(memory_primary_client(require_memory_prisma())).table.find_many( - where={ # mutable-ok: Prisma serializes these as native JSON containers. - "target_type": target_type, - "target_id": target_id, - } - if target_type and target_id - else None, - take=100, - skip=offset, - order={ # mutable-ok: Prisma serializes these as native JSON containers. - "policy_id": "asc" - }, - ) - return [ # mutable-ok: Prisma serializes these as native JSON containers. - MemoryPolicy.model_validate(row, from_attributes=True) for row in rows - ] - - -@router.put("/policies", response_model=MemoryPolicy) -async def set_policy(policy: MemoryPolicyInput, auth: UserAPIKeyAuth = _AUTH) -> MemoryPolicy: - await require_policy_admin(auth, policy.target_type, policy.target_id) - if auth.user_role != LitellmUserRoles.PROXY_ADMIN and policy.scope in ("user", "organization"): - if policy.target_type != "organization" or policy.scope != "organization": - raise HTTPException(status_code=403, detail="This shared scope requires a proxy administrator") - policy_id: Final = memory_digest(policy.target_type, policy.target_id) - fields: Final = { # mutable-ok: Prisma serializes these as native JSON containers. - **policy.model_dump(), - "updated_by": auth.user_id or "proxy-admin", - } - row: Final = await MemoryPolicyRepository(memory_primary_client(require_memory_prisma())).table.upsert( - where={ # mutable-ok: Prisma serializes these as native JSON containers. - "policy_id": policy_id - }, - data={ # mutable-ok: Prisma serializes these as native JSON containers. - "create": { # mutable-ok: Prisma serializes these as native JSON containers. - **fields, - "policy_id": policy_id, - }, - "update": fields, - }, - ) + if frozenset(user.user_id for user in users) != frozenset(selected): + raise HTTPException(status_code=422, detail="One or more selected users no longer exist") + saved: Final = settings.model_copy(update=MappingProxyType({"user_ids": selected})) + await ConfigRepository(prisma).set_param(MEMORY_CONFIG_PARAM, saved.model_dump(mode="json")) await invalidate_memory_configuration() - return MemoryPolicy.model_validate(row, from_attributes=True) + return saved -@router.delete("/policies/{policy_id}", status_code=204) -async def delete_policy(policy_id: str, auth: UserAPIKeyAuth = _AUTH) -> Response: - table: Final = MemoryPolicyRepository(memory_primary_client(require_memory_prisma())).table - row: Final = await table.find_unique( - where={ # mutable-ok: Prisma serializes these as native JSON containers. - "policy_id": policy_id - } - ) - if row is None: - raise HTTPException(status_code=404, detail="Memory policy not found") - policy: Final = MemoryPolicy.model_validate(row, from_attributes=True) - await require_policy_admin(auth, policy.target_type, policy.target_id) - await table.delete( - where={ # mutable-ok: Prisma serializes these as native JSON containers. - "policy_id": policy_id - } - ) - await invalidate_memory_configuration() - return Response(status_code=204) - - -@router.get("/preference", response_model=MemoryPreference) -async def get_preference( - auth: UserAPIKeyAuth = _AUTH, - key_id: Annotated[str | None, Query(pattern=r"^[a-f0-9]{64}$")] = None, -) -> MemoryPreference: - identity: Final = ( - (await access_for_key(auth, key_id)).identity if key_id is not None else MemoryIdentity.from_auth(auth) - ) - subject: Final = identity.preference_subject - row: Final = await MemoryPreferenceRepository(memory_primary_client(require_memory_prisma())).table.find_unique( - where={ # mutable-ok: Prisma serializes these as native JSON containers. - "subject": subject - } - ) - return MemoryPreference(enabled=row.enabled if row else False) - - -@router.put("/preference", response_model=MemoryPreference) -async def set_preference( - preference: MemoryPreference, - auth: UserAPIKeyAuth = _AUTH, - key_id: Annotated[str | None, Query(pattern=r"^[a-f0-9]{64}$")] = None, -) -> MemoryPreference: - identity: Final = ( - (await access_for_key(auth, key_id)).identity if key_id is not None else MemoryIdentity.from_auth(auth) - ) - if identity.read_only: - raise HTTPException(status_code=403, detail="Read-only users cannot change memory preferences") - subject: Final = identity.preference_subject - if not preference.enabled: - await MemoryPreferenceRepository(memory_primary_client(require_memory_prisma())).table.delete_many( - where={ # mutable-ok: Prisma serializes these as native JSON containers. - "subject": subject - } - ) - return preference - await MemoryPreferenceRepository(memory_primary_client(require_memory_prisma())).table.upsert( - where={ # mutable-ok: Prisma serializes these as native JSON containers. - "subject": subject - }, - data={ # mutable-ok: Prisma serializes these as native JSON containers. - "create": { # mutable-ok: Prisma serializes these as native JSON containers. - "subject": subject, - "enabled": preference.enabled, - }, - "update": { # mutable-ok: Prisma serializes these as native JSON containers. - "enabled": preference.enabled - }, - }, - ) - return preference +async def memory_store(auth: UserAPIKeyAuth) -> MemoryStore: + prisma: Final = memory_primary_client(require_memory_prisma()) + return MemoryStore(prisma, await resolve_memory_access(prisma, MemoryIdentity.from_auth(auth))) @router.get("/status", response_model=MemoryStatus) -async def get_status( - key_id: str | None = Query(None, pattern=r"^[a-f0-9]{64}$"), - auth: UserAPIKeyAuth = _AUTH, -) -> MemoryStatus: - access: Final = await access_for_key(auth, key_id) +async def get_status(auth: UserAPIKeyAuth = _AUTH) -> MemoryStatus: + store: Final = await memory_store(auth) user: Final = ( - await UserRepository(memory_primary_client(require_memory_prisma())).find_by_id(access.identity.user_id) - if access.identity.user_id + await UserRepository(store.prisma_client).find_by_id(store.access.identity.user_id) + if store.access.identity.user_id else None ) - return access.status.model_copy( + return store.access.status.model_copy( update=MappingProxyType({"user_name": user.user_alias or user.user_email or user.user_id if user else None}) ) -async def access_for_key(auth: UserAPIKeyAuth, key_id: str | None) -> MemoryAccess: - prisma: Final = memory_primary_client(require_memory_prisma()) - if key_id is None: - return await resolve_memory_access(prisma, MemoryIdentity.from_auth(auth)) - key: Final = await VerificationTokenRepository(prisma).find_by_id(key_id, id_field="token") - if key is None or not ( - user_api_key_has_admin_view(auth) - or key_id == MemoryIdentity.from_auth(auth).key_id - or (auth.is_session_token or auth.team_id == UI_TEAM_ID) - and auth.user_id - and key.user_id == auth.user_id - ): - raise HTTPException(status_code=403, detail="You cannot access memory for this key") - team: Final = await TeamRepository(prisma).find_by_id(key.team_id) if key.team_id else None - identity: Final = MemoryIdentity( - key_id=key_id, - user_id=key.user_id, - team_id=key.team_id, - project_id=key.project_id, - organization_id=key.org_id or (team.organization_id if team else None), - read_only=MemoryIdentity.from_auth(auth).read_only, +async def named_entries(store: MemoryStore, entries: tuple[MemoryEntry, ...]) -> tuple[MemoryEntry, ...]: + actors: Final = tuple(frozenset(entry.actor for entry in entries if entry.actor)) + teams: Final = tuple(frozenset(entry.team_id for entry in entries if entry.team_id)) + users: Final = ( + await UserRepository(store.prisma_client).table.find_many( + where={"user_id": {"in": list(actors)}}, # mutable-ok: Prisma requires native JSON. + take=len(actors), # mutable-ok: Prisma requires native query JSON. + ) + if actors + else () + ) + team_rows: Final = ( + await TeamRepository(store.prisma_client).table.find_many( + where={"team_id": {"in": list(teams)}}, # mutable-ok: Prisma requires native JSON. + take=len(teams), # mutable-ok: Prisma requires native query JSON. + ) + if teams + else () + ) + names: Final = MappingProxyType( + {user.user_id: user.user_alias or user.user_email or user.user_id for user in users} + ) + team_names: Final = MappingProxyType({team.team_id: team.team_alias or team.team_id for team in team_rows}) + return tuple( + entry.model_copy( + update=MappingProxyType( + { + "actor_name": names.get(entry.actor or ""), + "team_name": team_names.get(entry.team_id or ""), + } + ) + ) + for entry in entries ) - return await resolve_memory_access(prisma, identity) @router.get("/entries", response_model=list[MemoryEntry]) async def list_entries( query: Annotated[MemoryQuery, Query(max_length=500)] = "", limit: int = Query(20, ge=1, le=20), - offset: int = Query(0, ge=0), - key_id: str | None = Query(None, pattern=r"^[a-f0-9]{64}$"), - auth: UserAPIKeyAuth = _AUTH, + offset: int = Query(0, ge=0, le=10000), before_updated_at: Annotated[AwareDatetime | None, Query()] = None, before_memory_id: Annotated[str | None, Query(min_length=1, max_length=128)] = None, + team_id: Annotated[str | None, Query(min_length=1, max_length=256)] = None, + user_id: Annotated[str | None, Query(min_length=1, max_length=256)] = None, + auth: UserAPIKeyAuth = _AUTH, ) -> tuple[MemoryEntry, ...]: if (before_updated_at is None) != (before_memory_id is None): raise HTTPException(status_code=422, detail="Provide both memory cursor fields") - prisma: Final = memory_primary_client(require_memory_prisma()) - access: Final = await access_for_key(auth, key_id) - entries: Final = await MemoryStore(prisma, access).search( + store: Final = await memory_store(auth) + entries: Final = await store.search( MemorySearch(query=query, limit=limit, offset=offset), require_active=False, recent_first=True, before=(before_updated_at, before_memory_id) if before_updated_at and before_memory_id else None, + team_id=team_id, + user_id=user_id, ) - actors: Final = tuple(frozenset(entry.actor for entry in entries if entry.actor)) - if not actors: - return entries - users: Final = await UserRepository(prisma).table.find_many( - where={"user_id": {"in": list(actors)}}, # mutable-ok: Prisma serializes the bounded page's contributor IDs. - take=len(actors), - ) - names: Final = MappingProxyType( - {user.user_id: user.user_alias or user.user_email or user.user_id for user in users} - ) - return tuple( - entry.model_copy(update=MappingProxyType({"actor_name": names.get(entry.actor or "")})) for entry in entries - ) + return await named_entries(store, entries) + + +@router.get("/entries/{memory_id}", response_model=MemoryEntry) +async def read_entry(memory_id: str, auth: UserAPIKeyAuth = _AUTH) -> MemoryEntry: + store: Final = await memory_store(auth) + return (await named_entries(store, (await store.read(memory_id, require_active=False),)))[0] @router.post("/entries", response_model=MemoryEntry) -async def capture_entry( - capture: MemoryCapture, - key_id: str | None = Query(None, pattern=r"^[a-f0-9]{64}$"), - auth: UserAPIKeyAuth = _AUTH, -) -> MemoryEntry: - prisma: Final = memory_primary_client(require_memory_prisma()) - access: Final = await access_for_key(auth, key_id) - return await MemoryStore(prisma, access, actor=auth.user_id or MemoryIdentity.from_auth(auth).key_id).capture( - capture - ) +async def capture_entry(capture: MemoryCapture, auth: UserAPIKeyAuth = _AUTH) -> MemoryEntry: + return await (await memory_store(auth)).capture(capture) + + +@router.put("/entries/{memory_id}", response_model=MemoryEntry) +async def update_entry(memory_id: str, capture: MemoryCapture, auth: UserAPIKeyAuth = _AUTH) -> MemoryEntry: + return await (await memory_store(auth)).update(memory_id, capture) @router.delete("/entries/{memory_id}", status_code=204) -async def delete_entry( - memory_id: str, - key_id: str | None = Query(None, pattern=r"^[a-f0-9]{64}$"), - auth: UserAPIKeyAuth = _AUTH, -) -> Response: - prisma: Final = memory_primary_client(require_memory_prisma()) - access: Final = await access_for_key(auth, key_id) - if not await MemoryStore(prisma, access).delete(memory_id): +async def delete_entry(memory_id: str, auth: UserAPIKeyAuth = _AUTH) -> Response: + if not await (await memory_store(auth)).delete(memory_id): raise HTTPException(status_code=404, detail="Memory not found") return Response(status_code=204) diff --git a/litellm/proxy/memory/memory_endpoints.py b/litellm/proxy/memory/memory_endpoints.py index 69bd9c41fc7..f114d655d9d 100644 --- a/litellm/proxy/memory/memory_endpoints.py +++ b/litellm/proxy/memory/memory_endpoints.py @@ -220,7 +220,11 @@ async def is_memory_team_admin(prisma_client: "PrismaClient", user_api_key_dict: try: team_obj: Final = await TeamRepository(prisma_client).find_by_id(team_id, id_field="team_id") except Exception as e: - verbose_proxy_logger.exception("Error loading team for write-auth check (team_id=%s): %s", team_id, e) + verbose_proxy_logger.error( + "Error loading team for write-auth check (team_id=%s): %s", + team_id.replace("\r", "").replace("\n", ""), + str(e).replace("\r", "").replace("\n", ""), + ) return False if team_obj is None: return False @@ -236,7 +240,11 @@ async def is_memory_team_admin(prisma_client: "PrismaClient", user_api_key_dict: if await _is_user_org_admin_for_team(user_api_key_dict=user_api_key_dict, team_obj=team_obj): return True except Exception as e: - verbose_proxy_logger.debug("Org-admin check skipped during write-auth (team_id=%s): %s", team_id, e) + verbose_proxy_logger.debug( + "Org-admin check skipped during write-auth (team_id=%s): %s", + team_id.replace("\r", "").replace("\n", ""), + str(e).replace("\r", "").replace("\n", ""), + ) return False diff --git a/litellm/proxy/memory/policy.py b/litellm/proxy/memory/policy.py index 31a79f06bc1..742d4d26159 100644 --- a/litellm/proxy/memory/policy.py +++ b/litellm/proxy/memory/policy.py @@ -1,46 +1,43 @@ import hashlib import json -from dataclasses import dataclass +from collections.abc import Mapping +from dataclasses import dataclass, replace from typing import Final -from fastapi import HTTPException - from litellm.caching.caching import DualCache -from litellm.proxy._types import UI_TEAM_ID, LitellmUserRoles, UserAPIKeyAuth +from litellm.proxy._types import UI_TEAM_ID, KeyManagementRoutes, LiteLLM_TeamTable, LitellmUserRoles, UserAPIKeyAuth from litellm.proxy.common_utils.auth_cache_invalidation_pubsub import evict_and_broadcast from litellm.proxy.common_utils.config_sync_pubsub import coordination_redis_cache from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache from litellm.proxy.db.routing_prisma_wrapper import RoutingPrismaWrapper, WriterPinnedClient -from litellm.repositories.table_repositories import MemoryPolicyRepository, MemoryPreferenceRepository -from litellm.types.memory_v2 import MemoryPolicy, MemoryScope, MemoryStatus +from litellm.proxy.management_helpers.record_permissions import permitted_record_teams +from litellm.repositories.config_repository import ConfigRepository +from litellm.repositories.team_repository import TeamRepository +from litellm.repositories.user_repository import UserRepository +from litellm.types.memory_v2 import MemorySettings, MemoryStatus _CONFIGURED_CACHE_KEY: Final = "litellm:memory_v2:configured" +MEMORY_CONFIG_PARAM: Final = "memory_v2" + + +async def memory_settings(prisma_client: object) -> MemorySettings: + row: Final = await ConfigRepository(memory_primary_client(prisma_client)).get_param(MEMORY_CONFIG_PARAM) + return MemorySettings.model_validate(row.param_value) if row is not None else MemorySettings() async def gateway_memory_is_configured(prisma_client: object, cache: DualCache) -> bool: - """Avoid database work on ordinary requests when memory is not configured. - - This is only a presence hint. Authorization is always checked against the - primary before memory is used. The short TTL also covers separate workers - without shared Redis when an administrator first enables memory. - """ cached: Final = await cache.async_get_cache(key=_CONFIGURED_CACHE_KEY) if cached is True: return True redis_cache: Final = cache.redis_cache or coordination_redis_cache() - # An empty local view reads shared Redis through DualCache's existing - # circuit-breaker/error handling, falling back to the primary on a miss. shared_cache: Final = DualCache(redis_cache=redis_cache) if redis_cache is not None else None if cached is False: if shared_cache is None: return False - # A backend mutation evicts Redis, but another worker can still hold - # a negative local hint (Redis Cluster may not support pub/sub). shared: Final = await shared_cache.async_get_cache(key=_CONFIGURED_CACHE_KEY) if shared is False: return False - rows: Final = await MemoryPolicyRepository(memory_primary_client(prisma_client)).table.find_many(take=1) - configured: Final = bool(rows) + configured: Final = (await memory_settings(prisma_client)).enabled await cache.async_set_cache(key=_CONFIGURED_CACHE_KEY, value=configured, ttl=30) if shared_cache is not None and cache.redis_cache is None: await shared_cache.async_set_cache(key=_CONFIGURED_CACHE_KEY, value=configured, ttl=30) @@ -58,12 +55,9 @@ async def invalidate_memory_configuration() -> None: def memory_primary_client(prisma_client: object) -> WriterPinnedClient: - """Memory authorization and read-after-write must never use a lagging replica.""" db: Final = getattr(prisma_client, "db", None) if db is None: raise RuntimeError("Memory requires a connected Prisma database") - # Pin the actual writer even while unavailable: memory must fail closed - # instead of authorizing storage or recall from stale policy rows. return WriterPinnedClient(db.writer if isinstance(db, RoutingPrismaWrapper) else db) @@ -79,126 +73,139 @@ class MemoryIdentity: project_id: str | None organization_id: str | None read_only: bool + role: str | None = None + dashboard: bool = False @classmethod def from_auth(cls, auth: UserAPIKeyAuth) -> "MemoryIdentity": token: Final = auth.token or auth.api_key + dashboard: Final = auth.is_session_token or auth.team_id == UI_TEAM_ID key_id: Final = ( token - if token - and len(token) == 64 - and all(c in "0123456789abcdef" for c in token) - and not auth.is_session_token - and auth.team_id != UI_TEAM_ID + if token and len(token) == 64 and all(c in "0123456789abcdef" for c in token) and not dashboard else None ) return cls( key_id=key_id, user_id=auth.user_id, - team_id=auth.team_id, + team_id=None if dashboard else auth.team_id, project_id=auth.project_id, organization_id=auth.org_id, read_only=auth.user_role - in ( - LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY, - LitellmUserRoles.INTERNAL_USER_VIEW_ONLY, - ), + in (LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY, LitellmUserRoles.INTERNAL_USER_VIEW_ONLY), + role=auth.user_role, + dashboard=dashboard, ) @property - def preference_subject(self) -> str: - if self.user_id: - return memory_digest("user", self.user_id) - if self.key_id: - return memory_digest("key", self.key_id) - raise HTTPException(status_code=403, detail="Memory requires an authenticated user or virtual key") - - @property - def policy_targets(self) -> tuple[tuple[str, str], ...]: - return tuple( - (kind, value) - for kind, value in ( - ("key", self.key_id), - ("user", self.user_id), - ("project", self.project_id), - ("team", self.team_id), - ("organization", self.organization_id), - ("gateway", "*"), - ) - if value + def namespace(self) -> str: + return "v2:" + memory_digest( + self.organization_id, self.team_id, self.user_id, None if self.user_id else self.key_id ) - def namespace(self, scope: MemoryScope) -> str | None: - if scope == "key": - return ( - memory_digest(scope, self.organization_id, self.team_id, self.project_id, self.key_id) - if self.key_id - else None - ) - if scope == "user": - return memory_digest(scope, self.organization_id, self.user_id) if self.user_id else None - if scope == "team": - return memory_digest(scope, self.organization_id, self.team_id) if self.team_id else None - if scope == "project": - return ( - memory_digest(scope, self.organization_id, self.team_id, self.project_id) if self.project_id else None - ) - return memory_digest(scope, self.organization_id) if self.organization_id else None - @dataclass(frozen=True) class MemoryAccess: identity: MemoryIdentity - policy: MemoryPolicy | None - opted_in: bool + settings: MemorySettings + team_ids: tuple[str, ...] = () + admin_view: bool = False + permission_revision: str = "" @property - def namespace(self) -> str | None: - return self.identity.namespace(self.policy.scope) if self.policy else None + def namespace(self) -> str: + return self.identity.namespace @property def active(self) -> bool: return bool( - self.namespace - and self.policy - and (self.policy.activation == "automatic" or self.policy.activation == "opt_in" and self.opted_in) + self.settings.enabled + and (self.identity.user_id or self.identity.key_id) + and (self.settings.everyone or self.identity.user_id in self.settings.user_ids) ) @property def status(self) -> MemoryStatus: return MemoryStatus( active=self.active, - activation=self.policy.activation if self.policy else "disabled", - scope=self.policy.scope if self.policy and self.namespace else None, - opted_in=self.opted_in, - policy_id=self.policy.policy_id if self.policy else None, user_id=self.identity.user_id, + enabled=self.settings.enabled, + team_ids=self.team_ids, + admin_view=self.admin_view, ) + def visible_rows(self, *, write: bool = False) -> Mapping[str, object]: + if self.admin_view: + return {"namespace": {"not": None}} # mutable-ok: Prisma requires native query JSON. + owner: Final = ( + {"user_id": self.identity.user_id} # mutable-ok: Prisma requires native query JSON. + if self.identity.user_id + else { # mutable-ok: Prisma requires native JSON. + "owner_key_id": self.identity.key_id, + "user_id": None, + } # mutable-ok: Prisma requires native query JSON. + if self.identity.key_id + else {"memory_id": "__no_match__"} # mutable-ok: Prisma requires native query JSON. + ) + return { # mutable-ok: Prisma requires native query JSON. + "namespace": {"startswith": "v2:"}, # mutable-ok: Prisma requires native query JSON. + **( + {"organization_id": self.identity.organization_id} # mutable-ok: Prisma requires native query JSON. + if not self.identity.dashboard + else {} # mutable-ok: Prisma requires native JSON. + ), + "OR": [ # mutable-ok: Prisma requires native JSON. + owner, + *( + [{"team_id": {"in": list(self.team_ids)}}] if self.team_ids and not write else [] + ), # mutable-ok: Prisma requires native JSON. + ], # mutable-ok: Prisma requires native JSON. + } + async def resolve_memory_access(prisma_client: object, identity: MemoryIdentity) -> MemoryAccess: primary: Final = memory_primary_client(prisma_client) - rows: Final = await MemoryPolicyRepository(primary).table.find_many( - where={ # mutable-ok: Prisma serializes these as native JSON containers. - "OR": [ # mutable-ok: Prisma serializes these as native JSON containers. - { # mutable-ok: Prisma serializes these as native JSON containers. - "target_type": kind, - "target_id": target, - } - for kind, target in identity.policy_targets - ] - }, - take=len(identity.policy_targets), + settings: Final = await memory_settings(primary) + user: Final = await UserRepository(primary).find_by_id(identity.user_id) if identity.user_id else None + # Global roles come from normal authentication, including JWT and master-key grants. + role: Final = identity.role + admin_view: Final = role in (LitellmUserRoles.PROXY_ADMIN, LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY) + ids: Final = tuple(user.teams or ()) if user else () + team_ids: Final = tuple(dict.fromkeys((*ids, *((identity.team_id,) if identity.team_id else ())))) + rows: Final = ( + await TeamRepository(primary).table.find_many( + where={"team_id": {"in": list(team_ids)}}, # mutable-ok: Prisma requires native query JSON. + take=len(team_ids), + ) + if team_ids + else () ) - policies: Final = { # mutable-ok: Prisma serializes these as native JSON containers. - (row.target_type, row.target_id): MemoryPolicy.model_validate(row, from_attributes=True) for row in rows - } - policy: Final = next((policies[target] for target in identity.policy_targets if target in policies), None) - if not identity.user_id and not identity.key_id: - return MemoryAccess(identity=identity, policy=None, opted_in=False) - preference: Final = await MemoryPreferenceRepository(primary).table.find_unique( - where={ # mutable-ok: Prisma serializes these as native JSON containers. - "subject": identity.preference_subject - } + teams: Final = tuple(LiteLLM_TeamTable.model_validate(row.model_dump()) for row in rows) + context_team: Final = next((team for team in teams if team.team_id == identity.team_id), None) + current: Final = replace( + identity, + organization_id=identity.organization_id or (context_team.organization_id if context_team else None), + read_only=identity.read_only + or role in (LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY, LitellmUserRoles.INTERNAL_USER_VIEW_ONLY), + ) + eligible: Final = tuple( + team for team in teams if current.dashboard or admin_view or team.organization_id == current.organization_id + ) + auth: Final = UserAPIKeyAuth(user_id=current.user_id, user_role=role) + permitted: Final = permitted_record_teams(auth, eligible, KeyManagementRoutes.MEMORY_READ) + return MemoryAccess( + identity=current, + settings=settings, + team_ids=permitted, + admin_view=admin_view, + permission_revision=memory_digest( + current.namespace, + role, + str(admin_view), + *( + f"{team.team_id}:{team.organization_id}" + for team in sorted(eligible, key=lambda item: item.team_id) + if team.team_id in permitted + ), + ), ) - return MemoryAccess(identity=identity, policy=policy, opted_in=preference.enabled if preference else False) diff --git a/litellm/proxy/memory/store.py b/litellm/proxy/memory/store.py index e0f897e8fff..db8159ffb93 100644 --- a/litellm/proxy/memory/store.py +++ b/litellm/proxy/memory/store.py @@ -1,7 +1,8 @@ import asyncio import json +from collections.abc import Mapping from datetime import datetime -from types import SimpleNamespace +from types import MappingProxyType, SimpleNamespace from typing import TYPE_CHECKING, Final from fastapi import HTTPException @@ -12,26 +13,25 @@ from litellm.proxy.memory.policy import MemoryAccess, memory_digest, memory_prim from litellm.repositories.prisma_protocols import TableActions from litellm.repositories.table_repositories import MemoryRepository from litellm.repositories.unit_of_work import prisma_transaction -from litellm.types.memory_v2 import MemoryCapture, MemoryEntry, MemorySearch +from litellm.types.memory_v2 import MemoryCapture, MemoryCatalogRequest, MemoryEntry, MemoryRecallRequest, MemorySearch if TYPE_CHECKING: from prisma.models import LiteLLM_MemoryTable _METADATA: Final = TypeAdapter(dict[str, object]) -_MAX_NAMESPACE_ENTRIES: Final = 1000 +_MAX_OWNER_ENTRIES: Final = 1000 +_PAGE_SIZE: Final = 128 +RankedMemory = tuple[MemoryEntry, float, tuple[str, ...]] def memory_entry(row: "LiteLLM_MemoryTable") -> MemoryEntry: metadata: Final = ( - _METADATA.validate_python(row.metadata) - if isinstance(row.metadata, dict) - else { # mutable-ok: Prisma serializes these as native JSON containers. - } - ) + _METADATA.validate_python(row.metadata) if isinstance(row.metadata, dict) else MappingProxyType({}) + ) # mutable-ok: Prisma requires native JSON. title: Final = metadata.get("title") evidence: Final = metadata.get("evidence") return MemoryEntry.model_validate( - { # mutable-ok: Pydantic validates stored JSON metadata and database fields together. + { # mutable-ok: Pydantic validates stored JSON and database fields together. "memory_id": row.memory_id, "key": row.key.rsplit(":", 1)[-1], "title": title if isinstance(title, str) else row.key, @@ -40,7 +40,9 @@ def memory_entry(row: "LiteLLM_MemoryTable") -> MemoryEntry: "updated_at": row.updated_at, "created_at": row.created_at, "actor": row.created_by, - **{ # mutable-ok: Pydantic validates these stored JSON metadata fields. + "user_id": row.user_id, + "team_id": row.team_id, + **{ # mutable-ok: Prisma requires native JSON. name: value for name in ("when_to_use", "scope", "kind", "certainty", "source") if isinstance(value := metadata.get(name), str) @@ -49,6 +51,24 @@ def memory_entry(row: "LiteLLM_MemoryTable") -> MemoryEntry: ) +def _where(*conditions: Mapping[str, object]) -> Mapping[str, object]: + return {"AND": list(conditions)} # mutable-ok: Prisma requires native query JSON. + + +def _before(cursor: tuple[datetime, str] | None) -> Mapping[str, object]: + if cursor is None: + return {} # mutable-ok: Prisma requires native JSON. + return { # mutable-ok: Prisma requires native query JSON. + "OR": [ # mutable-ok: Prisma requires native JSON. + {"updated_at": {"lt": cursor[0]}}, # mutable-ok: Prisma requires native JSON. + { # mutable-ok: Prisma requires native JSON. + "updated_at": cursor[0], + "memory_id": {"gt": cursor[1]}, # mutable-ok: Prisma requires native JSON. + }, # mutable-ok: Prisma requires native JSON. + ] + } + + class MemoryStore: def __init__(self, prisma_client: object, access: MemoryAccess, *, actor: str | None = None) -> None: self.prisma_client = memory_primary_client(prisma_client) @@ -56,18 +76,99 @@ class MemoryStore: self.actor = actor or access.identity.user_id or access.identity.key_id self.table = MemoryRepository(self.prisma_client).table - async def authorize_namespace(self, *, write: bool = False, require_active: bool = True) -> str: + async def authorize(self, *, write: bool = False, require_active: bool = True) -> MemoryAccess: current: Final = await resolve_memory_access(self.prisma_client, self.access.identity) if ( - current.namespace is None - or current.namespace != self.access.namespace + not (current.identity.user_id or current.identity.key_id) + or current.permission_revision != self.access.permission_revision or require_active and not current.active or write and current.identity.read_only ): - raise HTTPException(status_code=403, detail="Memory is not available under the current policy") - return current.namespace + raise HTTPException(status_code=403, detail="Memory access changed or is disabled") + return current + + async def authorize_namespace(self, *, write: bool = False, require_active: bool = True) -> str: + return (await self.authorize(write=write, require_active=require_active)).namespace + + def entry(self, row: "LiteLLM_MemoryTable") -> MemoryEntry: + owned: Final = ( + row.user_id == self.access.identity.user_id + if self.access.identity.user_id + else bool( + row.user_id is None and self.access.identity.key_id and row.owner_key_id == self.access.identity.key_id + ) + ) + return memory_entry(row).model_copy( + update={ # mutable-ok: Prisma requires native JSON. + "can_edit": not self.access.identity.read_only and (owned or self.access.admin_view) + } # mutable-ok: Prisma requires native JSON. + ) + + async def _page( + self, where: Mapping[str, object], *, limit: int, offset: int = 0, before: tuple[datetime, str] | None = None + ) -> tuple[MemoryEntry, ...]: + rows: Final = await self.table.find_many( + where=_where(where, _before(before)), + order=[{"updated_at": "desc"}, {"memory_id": "asc"}], # mutable-ok: Prisma requires native query JSON. + take=limit, + skip=offset, + ) + return tuple(self.entry(row) for row in rows) + + async def catalog(self, request: MemoryCatalogRequest) -> tuple[tuple[MemoryEntry, ...], int, str]: + access: Final = await self.authorize() + where: Final = access.visible_rows() + total: Final = await self.table.count(where=where) + entries: Final = await self._page(where, limit=request.limit, offset=request.offset) + latest: Final = await self._page(where, limit=1) if request.offset else entries[:1] + await self.authorize() + return ( + entries, + total, + memory_digest(str(total), *(entry.memory_id + entry.updated_at.isoformat() for entry in latest)), + ) + + async def _ranked( + self, + query: str, + where: Mapping[str, object], + keep: int, + *, + scope: str | None = None, + recent_first: bool = False, + ) -> tuple[tuple[RankedMemory, ...], int]: + return await asyncio.wait_for( + self._ranked_pages(query, where, keep, scope=scope, recent_first=recent_first), timeout=15 + ) + + async def _ranked_pages( + self, query: str, where: Mapping[str, object], keep: int, *, scope: str | None, recent_first: bool + ) -> tuple[tuple[RankedMemory, ...], int]: + best: tuple[RankedMemory, ...] = () # rebind-ok: Bounded top results are replaced after each database page. + total = 0 # rebind-ok: Count matches across bounded database pages. + cursor: tuple[datetime, str] | None = None # rebind-ok: Advance the database cursor after each page. + while True: + page = await self._page(where, limit=_PAGE_SIZE, before=cursor) + candidates = tuple(entry for entry in page if scope is None or scope.casefold() in entry.scope.casefold()) + matches = await asyncio.to_thread(fuzzy_memories, query, candidates) + total = total + len(matches) + combined = (*best, *matches) + best = tuple( + sorted(combined, key=lambda item: (-item[0].updated_at.timestamp(), item[0].memory_id)) + if recent_first + else sorted(combined, key=lambda item: (-item[1], item[0].memory_id)) + )[:keep] + if len(page) < _PAGE_SIZE: + return best, total + cursor = (page[-1].updated_at, page[-1].memory_id) + + async def recall(self, request: MemoryRecallRequest) -> tuple[tuple[RankedMemory, ...], int]: + access: Final = await self.authorize() + ranked: Final = await self._ranked(request.query, access.visible_rows(), request.limit, scope=request.scope) + await self.authorize() + return ranked async def search( self, @@ -76,54 +177,34 @@ class MemoryStore: require_active: bool = True, recent_first: bool = False, before: tuple[datetime, str] | None = None, + team_id: str | None = None, + user_id: str | None = None, ) -> tuple[MemoryEntry, ...]: - entries: Final = await self.entries(require_active=require_active) - ranked: Final = await asyncio.to_thread(fuzzy_memories, search.query, entries) - matched_ids: Final = frozenset(entry.memory_id for entry, _, _ in ranked) - ordered: Final = ( - tuple(entry for entry in entries if entry.memory_id in matched_ids) - if recent_first - else tuple(entry for entry, _, _ in ranked) + access: Final = await self.authorize(require_active=require_active) + where: Final = _where( + access.visible_rows(), + {"team_id": team_id} if team_id else {}, # mutable-ok: Prisma requires native JSON. + {"user_id": user_id} if user_id else {}, # mutable-ok: Prisma requires native JSON. + _before(before), ) - page: Final = tuple( - entry - for entry in ordered - if before is None - or entry.updated_at < before[0] - or entry.updated_at == before[0] - and entry.memory_id > before[1] - ) - return page[search.offset : search.offset + search.limit] + result: Final[tuple[MemoryEntry, ...]] + if not search.query.strip(): + result = await self._page(where, limit=search.limit, offset=search.offset) + else: + ranked, _ = await self._ranked(search.query, where, search.offset + search.limit, recent_first=recent_first) + result = tuple(entry for entry, _, _ in ranked[search.offset :]) + await self.authorize(require_active=require_active) + return result - async def entries(self, *, require_active: bool = True) -> tuple[MemoryEntry, ...]: - namespace: Final = await self.authorize_namespace(require_active=require_active) - rows: Final = await self.table.find_many( - where={ # mutable-ok: Prisma serializes these as native JSON containers. - "namespace": namespace, - }, - order=[ # mutable-ok: Prisma serializes these as native JSON containers. - { # mutable-ok: Prisma serializes these as native JSON containers. - "updated_at": "desc" - }, - { # mutable-ok: Prisma serializes these as native JSON containers. - "memory_id": "asc" - }, - ], - take=_MAX_NAMESPACE_ENTRIES, - ) - return tuple(memory_entry(row) for row in rows) - - async def read(self, memory_id: str) -> MemoryEntry: - namespace: Final = await self.authorize_namespace() + async def read(self, memory_id: str, *, require_active: bool = True) -> MemoryEntry: + access: Final = await self.authorize(require_active=require_active) row: Final = await self.table.find_first( - where={ # mutable-ok: Prisma serializes these as native JSON containers. - "memory_id": memory_id, - "namespace": namespace, - } - ) + where=_where(access.visible_rows(), {"memory_id": memory_id}) # mutable-ok: Prisma requires native JSON. + ) # mutable-ok: Prisma requires native JSON. if row is None: raise HTTPException(status_code=404, detail="Memory not found") - return memory_entry(row) + await self.authorize(require_active=require_active) + return self.entry(row) async def capture(self, capture: MemoryCapture) -> MemoryEntry: return (await self.capture_many((capture,)))[0] @@ -137,7 +218,7 @@ class MemoryStore: await transaction.execute_raw("SELECT pg_advisory_xact_lock($1::bigint)", lock_key) table: Final = MemoryRepository(SimpleNamespace(db=transaction)).table saved: Final = tuple([await self._capture(capture, namespace, table) for capture in captures]) - await MemoryStore(SimpleNamespace(db=transaction), self.access).authorize_namespace(write=True) + await MemoryStore(SimpleNamespace(db=transaction), self.access).authorize(write=True) return saved async def _capture( @@ -170,7 +251,7 @@ class MemoryStore: if existing is not None: if existing.namespace != namespace: raise HTTPException(status_code=409, detail="Memory key conflict") - entry: Final = memory_entry(existing) + entry: Final = self.entry(existing) if existing.value == content and all(getattr(entry, name) == value for name, value in metadata.items()): return entry if capture.expected_revision != existing.updated_at: @@ -196,7 +277,7 @@ class MemoryStore: ) if updated is None: raise HTTPException(status_code=409, detail="Memory no longer exists") - return memory_entry(updated) + return self.entry(updated) if capture.expected_revision is not None: raise HTTPException(status_code=409, detail="Memory no longer exists") entries: Final = await table.count( @@ -204,9 +285,9 @@ class MemoryStore: "namespace": namespace } ) - if entries >= _MAX_NAMESPACE_ENTRIES: + if entries >= _MAX_OWNER_ENTRIES: raise HTTPException( - status_code=429, detail="Memory scope has reached 1000 entries; delete unused memories first" + status_code=429, detail="Memory storage limit reached for this owner; contact your administrator" ) created: Final = await table.create( data={ # mutable-ok: Prisma query and write JSON. @@ -216,18 +297,39 @@ class MemoryStore: "namespace": namespace, "user_id": self.access.identity.user_id, "team_id": self.access.identity.team_id, + "organization_id": self.access.identity.organization_id, + "owner_key_id": self.access.identity.key_id, "created_by": self.actor, } ) - return memory_entry(created) + return self.entry(created) + + async def update(self, memory_id: str, capture: MemoryCapture) -> MemoryEntry: + access: Final = await self.authorize(write=True, require_active=False) + async with prisma_transaction(self.prisma_client) as transaction: + table: Final = MemoryRepository(SimpleNamespace(db=transaction)).table + row: Final = await table.find_first( + where=_where( + access.visible_rows(write=True), {"memory_id": memory_id} + ) # mutable-ok: Prisma requires native JSON. + ) # mutable-ok: Prisma requires native JSON. + if row is None or row.namespace is None: + raise HTTPException(status_code=404, detail="Memory not found") + if row.key.rsplit(":", 1)[-1] != capture.key: + raise HTTPException(status_code=422, detail="The memory key cannot be changed") + result: Final = await self._capture(capture, row.namespace, table) + await MemoryStore(SimpleNamespace(db=transaction), self.access).authorize(write=True, require_active=False) + return result async def delete(self, memory_id: str) -> bool: - namespace: Final = await self.authorize_namespace(write=True, require_active=False) - return bool( - await self.table.delete_many( - where={ # mutable-ok: Prisma serializes these as native JSON containers. - "memory_id": memory_id, - "namespace": namespace, - } + access: Final = await self.authorize(write=True, require_active=False) + async with prisma_transaction(self.prisma_client) as transaction: + table: Final = MemoryRepository(SimpleNamespace(db=transaction)).table + deleted: Final = await table.delete_many( + where=_where( + access.visible_rows(write=True), + {"memory_id": memory_id}, # mutable-ok: Prisma requires native JSON. + ) # mutable-ok: Prisma requires native JSON. ) - ) + await MemoryStore(SimpleNamespace(db=transaction), self.access).authorize(write=True, require_active=False) + return bool(deleted) diff --git a/litellm/proxy/memory/transport.py b/litellm/proxy/memory/transport.py index 726a1c36854..e827b3a835e 100644 --- a/litellm/proxy/memory/transport.py +++ b/litellm/proxy/memory/transport.py @@ -155,7 +155,7 @@ class GatewayRound: if not self.task.done(): self.task.cancel() with suppress(asyncio.CancelledError): - await self.task + await asyncio.gather(self.task) await self.reader.aclose() diff --git a/litellm/proxy/schema.prisma b/litellm/proxy/schema.prisma index 8aaf4b48bbb..e2c293d939d 100644 --- a/litellm/proxy/schema.prisma +++ b/litellm/proxy/schema.prisma @@ -1436,6 +1436,8 @@ model LiteLLM_MemoryTable { metadata Json? user_id String? team_id String? + organization_id String? + owner_key_id String? created_at DateTime @default(now()) created_by String? updated_at DateTime @default(now()) @updatedAt @@ -1444,6 +1446,9 @@ model LiteLLM_MemoryTable { @@index([user_id]) @@index([team_id]) @@index([namespace, updated_at]) + @@index([organization_id, user_id, updated_at]) + @@index([team_id, updated_at]) + @@index([owner_key_id]) } model LiteLLM_MemoryPolicy { @@ -1452,6 +1457,7 @@ model LiteLLM_MemoryPolicy { target_id String activation String scope String + paused Boolean @default(false) created_at DateTime @default(now()) updated_at DateTime @default(now()) @updatedAt updated_by String diff --git a/litellm/proxy/spend_tracking/spend_management_endpoints.py b/litellm/proxy/spend_tracking/spend_management_endpoints.py index 0ec2788fbaa..00cb08aec38 100644 --- a/litellm/proxy/spend_tracking/spend_management_endpoints.py +++ b/litellm/proxy/spend_tracking/spend_management_endpoints.py @@ -4690,10 +4690,7 @@ async def _can_team_member_view_log( Returns True if the team exists and the user is either a team admin or a team member with the ``/spend/logs`` permission. """ - from litellm.proxy.management_endpoints.common_utils import ( - _is_user_team_admin, - _team_member_has_permission, - ) + from litellm.proxy.management_helpers.record_permissions import can_read_team_records if team_id is None: return False @@ -4701,13 +4698,7 @@ async def _can_team_member_view_log( if team_row is None: return False team_obj: Final = LiteLLM_TeamTable.model_validate(team_row.model_dump()) - if _is_user_team_admin(user_api_key_dict=user_api_key_dict, team_obj=team_obj): - return True - return _team_member_has_permission( - user_api_key_dict=user_api_key_dict, - team_obj=team_obj, - permission=KeyManagementRoutes.SPEND_LOGS.value, - ) + return can_read_team_records(user_api_key_dict, team_obj, KeyManagementRoutes.SPEND_LOGS) def _can_user_view_spend_log(user_api_key_dict: UserAPIKeyAuth) -> bool: @@ -4920,10 +4911,7 @@ async def _get_permitted_team_ids_for_spend_logs( """ # Imported here to avoid circular import: proxy_server imports this module. from litellm.proxy.auth.auth_checks import get_user_object - from litellm.proxy.management_endpoints.common_utils import ( - _is_user_team_admin, - _team_member_has_permission, - ) + from litellm.proxy.management_helpers.record_permissions import permitted_record_teams from litellm.proxy.proxy_server import proxy_logging_obj, user_api_key_cache user_obj: Final = await get_user_object( @@ -4938,16 +4926,10 @@ async def _get_permitted_team_ids_for_spend_logs( team_rows: Final = await _find_team_rows(prisma_client, user_obj.teams) - permitted: Final[list[str]] = [] - for team_row in team_rows: - team_obj = LiteLLM_TeamTable.model_validate(team_row.model_dump()) - if _is_user_team_admin(user_api_key_dict=user_api_key_dict, team_obj=team_obj) or _team_member_has_permission( - user_api_key_dict=user_api_key_dict, - team_obj=team_obj, - permission=KeyManagementRoutes.SPEND_LOGS.value, - ): - permitted.append(team_obj.team_id) - return permitted + teams: Final = tuple(LiteLLM_TeamTable.model_validate(team.model_dump()) for team in team_rows) + return list( # mutable-ok: Preserve the existing Logs helper return contract. + permitted_record_teams(user_api_key_dict, teams, KeyManagementRoutes.SPEND_LOGS) + ) async def _get_permitted_team_ids_for_spend_logs_or_empty( diff --git a/litellm/repositories/table_repositories.py b/litellm/repositories/table_repositories.py index bf2641f6fad..e8e3d54ad6c 100644 --- a/litellm/repositories/table_repositories.py +++ b/litellm/repositories/table_repositories.py @@ -124,14 +124,6 @@ class MemoryRepository(PrismaTableRepository["prisma_models.LiteLLM_MemoryTable" table_name = "litellm_memorytable" -class MemoryPolicyRepository(PrismaTableRepository["prisma_models.LiteLLM_MemoryPolicy"]): - table_name = "litellm_memorypolicy" - - -class MemoryPreferenceRepository(PrismaTableRepository["prisma_models.LiteLLM_MemoryPreference"]): - table_name = "litellm_memorypreference" - - class MemoryContinuationRepository(PrismaTableRepository["prisma_models.LiteLLM_MemoryContinuation"]): table_name = "litellm_memorycontinuation" diff --git a/litellm/repositories/unit_of_work.py b/litellm/repositories/unit_of_work.py index d5faad0a156..21d2aa43c90 100644 --- a/litellm/repositories/unit_of_work.py +++ b/litellm/repositories/unit_of_work.py @@ -38,7 +38,8 @@ def prisma_transaction(client: object) -> AbstractAsyncContextManager["Prisma"]: raise TypeError("A transactional Prisma client is required") transaction_factory: Final = ( cast( # cast-ok: The callable is Prisma's tx factory, dynamically forwarded by supported database wrappers. - Callable[[], AbstractAsyncContextManager["Prisma"]], factory + Callable[[], AbstractAsyncContextManager["Prisma"]], # mutable-ok: Callable type parameter syntax. + factory, # mutable-ok: Callable annotation uses an empty parameter list. ) ) return transaction_factory() @@ -97,7 +98,7 @@ class LinkedSpendResetWrites: table: BatchTable def queue_spend_zero(self, where: Mapping[str, object]) -> None: - self.table.update_many(where=where, data={"spend": 0}) + self.table.update_many(where=where, data={"spend": 0}) # mutable-ok: Callable type parameter syntax. def queue_spend_decrement(self, where: Mapping[str, object], amount: float) -> None: """``decrement`` rather than a read-then-set, so spend written between the @@ -115,7 +116,9 @@ class BudgetWindowWrites: def queue_window_advance(self, budget_id: str, budget_reset_at: datetime) -> None: """``update_many`` so a tier deleted between the read and the commit is a no-op row count instead of a P2025 that aborts the whole chunk.""" - self.table.update_many(where={"budget_id": budget_id}, data={"budget_reset_at": budget_reset_at}) + self.table.update_many( + where={"budget_id": budget_id}, data={"budget_reset_at": budget_reset_at} + ) # mutable-ok: Callable type parameter syntax. @dataclass(frozen=True, slots=True) diff --git a/litellm/types/memory_v2.py b/litellm/types/memory_v2.py index c8ed56f0f97..03c17760b8e 100644 --- a/litellm/types/memory_v2.py +++ b/litellm/types/memory_v2.py @@ -2,12 +2,8 @@ import re from datetime import datetime from typing import Annotated, Literal, TypeAlias -from pydantic import AfterValidator, BaseModel, ConfigDict, Field, model_validator -from typing_extensions import Self +from pydantic import AfterValidator, BaseModel, ConfigDict, Field -MemoryTarget: TypeAlias = Literal["gateway", "organization", "team", "project", "user", "key"] -MemoryScope: TypeAlias = Literal["key", "user", "team", "project", "organization"] -MemoryActivation: TypeAlias = Literal["disabled", "opt_in", "automatic"] MemoryKind: TypeAlias = Literal["workflow", "decision", "correction", "learning", "context", "disagreement"] MemoryCertainty: TypeAlias = Literal["user_stated", "observed", "inferred"] @@ -23,47 +19,23 @@ MemoryQuery: TypeAlias = Annotated[ ] -class MemoryPolicyInput(BaseModel): +class MemorySettings(BaseModel): model_config = ConfigDict(extra="forbid", frozen=True) - target_type: MemoryTarget - target_id: str = Field(min_length=1, max_length=256) - activation: MemoryActivation - scope: MemoryScope = "user" - - @model_validator(mode="after") - def validate_target(self) -> Self: - if self.target_type == "gateway" and self.target_id != "*": - raise ValueError("The gateway target_id must be '*'") - if self.target_type == "key" and ( - len(self.target_id) != 64 or any(c not in "0123456789abcdef" for c in self.target_id) - ): - raise ValueError("Use the key's hash, never its secret value") - return self - - -class MemoryPolicy(MemoryPolicyInput): - policy_id: str - updated_at: datetime - updated_by: str - - -class MemoryPreference(BaseModel): - model_config = ConfigDict(extra="forbid", frozen=True) - - enabled: bool + enabled: bool = False + everyone: bool = True + user_ids: tuple[Annotated[str, Field(min_length=1, max_length=256)], ...] = Field(default=(), max_length=10000) class MemoryStatus(BaseModel): model_config = ConfigDict(frozen=True) active: bool - activation: MemoryActivation - scope: MemoryScope | None - opted_in: bool - policy_id: str | None user_id: str | None = None user_name: str | None = None + enabled: bool = False + team_ids: tuple[str, ...] = () + admin_view: bool = False class MemoryCapture(BaseModel): @@ -93,6 +65,10 @@ class MemoryEntry(BaseModel): created_at: datetime | None = None actor: str | None = None actor_name: str | None = None + user_id: str | None = None + team_id: str | None = None + team_name: str | None = None + can_edit: bool = False when_to_use: str = "" scope: str = "" kind: MemoryKind = "context" @@ -105,7 +81,7 @@ class MemorySearch(BaseModel): query: MemoryQuery = Field(default="", max_length=500) limit: int = Field(default=8, ge=1, le=20) - offset: int = Field(default=0, ge=0) + offset: int = Field(default=0, ge=0, le=10000) class MemoryRead(BaseModel): diff --git a/schema.prisma b/schema.prisma index 8aaf4b48bbb..e2c293d939d 100644 --- a/schema.prisma +++ b/schema.prisma @@ -1436,6 +1436,8 @@ model LiteLLM_MemoryTable { metadata Json? user_id String? team_id String? + organization_id String? + owner_key_id String? created_at DateTime @default(now()) created_by String? updated_at DateTime @default(now()) @updatedAt @@ -1444,6 +1446,9 @@ model LiteLLM_MemoryTable { @@index([user_id]) @@index([team_id]) @@index([namespace, updated_at]) + @@index([organization_id, user_id, updated_at]) + @@index([team_id, updated_at]) + @@index([owner_key_id]) } model LiteLLM_MemoryPolicy { @@ -1452,6 +1457,7 @@ model LiteLLM_MemoryPolicy { target_id String activation String scope String + paused Boolean @default(false) created_at DateTime @default(now()) updated_at DateTime @default(now()) @updatedAt updated_by String diff --git a/tests/e2e/coverage_registry/mgmt.yaml b/tests/e2e/coverage_registry/mgmt.yaml index d043b863ccf..ab9c6a121f0 100644 --- a/tests/e2e/coverage_registry/mgmt.yaml +++ b/tests/e2e/coverage_registry/mgmt.yaml @@ -91,9 +91,11 @@ - {id: mgmt.mcp_toolset.update.clear_persists, module: mgmt, tier: P0, surface: api, assertions: [clear_persists], source: "mcp_management_endpoints.py:3098", fail_before_fix: proven, rationale: "An explicit null clears the stored description; the update used to drop null and keep the old value"} - {id: mgmt.mcp_toolset.delete.persists, module: mgmt, tier: P0, surface: api, assertions: [persists], source: "mcp_management_endpoints.py:3149", rationale: "A deleted toolset is gone by id and from the list on every replica"} - {id: mgmt.memory_v2.gateway.capture_recall, module: mgmt, tier: P0, surface: api, assertions: [capture_recall], source: "memory/gateway.py", rationale: "Existing clients store and recall across fresh conversations and native streams"} -- {id: mgmt.memory_v2.policy.opt_in, module: mgmt, tier: P0, surface: api, assertions: [opt_in], source: "memory/policy.py", rationale: "Administrators choose opt-in or automatic activation and more specific policies win"} -- {id: mgmt.memory_v2.entries.isolation, module: mgmt, tier: P0, surface: api, assertions: [isolation], source: "memory/store.py", rationale: "Private memories never cross virtual keys, including sibling keys and legacy API reads"} -- {id: mgmt.memory_v2.policy.admin_only, module: mgmt, tier: P0, surface: api, assertions: [admin_only], source: "memory/management.py", rationale: "Members cannot enable memory or broaden its sharing scope"} +- {id: mgmt.memory_v2.settings.enrollment, module: mgmt, tier: P0, surface: api, assertions: [enrollment], source: "memory/policy.py", rationale: "Admin enrollment follows owners across keys; disabling stops automatic memory and preserves authorized inspection"} +- {id: mgmt.memory_v2.entries.isolation, module: mgmt, tier: P0, surface: api, assertions: [isolation], source: "memory/store.py", rationale: "Owner keys share records while unrelated users and V1 remain isolated"} +- {id: mgmt.memory_v2.settings.admin_only, module: mgmt, tier: P0, surface: api, assertions: [admin_only], source: "memory/management.py", rationale: "Only proxy administrators can enable memory for everyone or selected users"} - {id: mgmt.memory_v2.entries.correction_delete, module: mgmt, tier: P0, surface: api, assertions: [correction_delete], source: "memory/store.py", rationale: "Corrections reject stale revisions and deletion removes facts from future recall"} - {id: mgmt.memory_v2.gateway.client_tools, module: mgmt, tier: P0, surface: api, assertions: [client_tools], source: "memory/gateway.py", rationale: "Application tool calls and their continuations stay under client control"} - {id: mgmt.memory_v2.gateway.billing, module: mgmt, tier: P0, surface: api, assertions: [billing], source: "memory/gateway.py", rationale: "Preparation and final answer each create distinct billed calls under the authenticated key, user, and team"} + +- {id: mgmt.memory_v2.entries.team_permissions, module: mgmt, tier: P0, surface: api, assertions: [team_permissions], source: "memory/policy.py", rationale: "Existing team grants allow reads without writes; log grants do not confer memory access; revocation takes effect"} diff --git a/tests/e2e/management/memory_client.py b/tests/e2e/management/memory_client.py index d29234bdcf2..4e52423947a 100644 --- a/tests/e2e/management/memory_client.py +++ b/tests/e2e/management/memory_client.py @@ -1,7 +1,4 @@ -import hashlib from dataclasses import dataclass -from typing import Final, Literal -from urllib.parse import quote from e2e_http import NoBody, Result, unwrap from models import ( @@ -9,11 +6,7 @@ from models import ( MemoryEntriesData, MemoryEntryData, MemoryEntryParams, - MemoryLegacyParams, - MemoryLegacyRows, - MemoryPolicyBody, - MemoryPolicyData, - MemoryPreferenceBody, + MemorySettingsBody, MemoryStatusData, ) from proxy_client import ProxyClient @@ -23,43 +16,38 @@ from proxy_client import ProxyClient class MemoryClient: proxy: ProxyClient - def set_policy(self, body: MemoryPolicyBody, *, caller: str | None = None) -> Result[MemoryPolicyData]: + def settings(self) -> MemorySettingsBody: + return unwrap( + self.proxy.transport.get( + "/v2/memory/settings", + headers=self.proxy.transport.master, + params=NoBody(), + response_type=MemorySettingsBody, + ) + ) + + def set_settings(self, body: MemorySettingsBody, *, caller: str | None = None) -> Result[MemorySettingsBody]: return self.proxy.transport.put( - "/v2/memory/policies", + "/v2/memory/settings", headers=self.proxy.transport.bearer(caller) if caller else self.proxy.transport.master, json=body, - response_type=MemoryPolicyData, + response_type=MemorySettingsBody, ) - def policy_for_key(self, key: str, activation: Literal["disabled", "opt_in", "automatic"]) -> MemoryPolicyData: - return unwrap( - self.set_policy( - MemoryPolicyBody( - target_type="key", - target_id=hashlib.sha256(key.encode()).hexdigest(), - activation=activation, - ) - ) + def read(self, key: str, memory_id: str) -> Result[MemoryEntryData]: + return self.proxy.transport.get( + f"/v2/memory/entries/{memory_id}", + headers=self.proxy.transport.bearer(key), + params=NoBody(), + response_type=MemoryEntryData, ) - def delete_policy(self, policy_id: str) -> None: - unwrap( - self.proxy.transport.delete( - f"/v2/memory/policies/{policy_id}", - headers=self.proxy.transport.master, - json=NoBody(), - response_type=NoBody, - ) - ) - - def preference(self, key: str, enabled: bool) -> MemoryPreferenceBody: - return unwrap( - self.proxy.transport.put( - "/v2/memory/preference", - headers=self.proxy.transport.bearer(key), - json=MemoryPreferenceBody(enabled=enabled), - response_type=MemoryPreferenceBody, - ) + def update(self, key: str, memory_id: str, body: MemoryCaptureBody) -> Result[MemoryEntryData]: + return self.proxy.transport.put( + f"/v2/memory/entries/{memory_id}", + headers=self.proxy.transport.bearer(key), + json=body, + response_type=MemoryEntryData, ) def status(self, key: str) -> MemoryStatusData: @@ -99,32 +87,23 @@ class MemoryClient: ) def cleanup_user_entries(self, user_id: str) -> None: - first: Final = unwrap( - self.proxy.transport.get( - "/v1/memory", - headers=self.proxy.transport.master, - params=MemoryLegacyParams(), - response_type=MemoryLegacyRows, - ) - ) - remaining: Final = tuple( - unwrap( + while True: + page = unwrap( self.proxy.transport.get( - "/v1/memory", + "/v2/memory/entries", headers=self.proxy.transport.master, - params=MemoryLegacyParams(page=page), - response_type=MemoryLegacyRows, + params=MemoryEntryParams(user_id=user_id), + response_type=MemoryEntriesData, ) - ) - for page in range(2, (first.total + 499) // 500 + 1) - ) - rows: Final = tuple(row for page in (first, *remaining) for row in page.memories if row.user_id == user_id) - for row in rows: - unwrap( - self.proxy.transport.delete( - f"/v1/memory/{quote(row.key, safe='')}", - headers=self.proxy.transport.master, - json=NoBody(), - response_type=NoBody, + ).root + if not page: + return + for entry in page: + unwrap( + self.proxy.transport.delete( + f"/v2/memory/entries/{entry.memory_id}", + headers=self.proxy.transport.master, + json=NoBody(), + response_type=NoBody, + ) ) - ) diff --git a/tests/e2e/management/test_memory_v2_e2e.py b/tests/e2e/management/test_memory_v2_e2e.py index 55d2d5d822a..1316cabe591 100644 --- a/tests/e2e/management/test_memory_v2_e2e.py +++ b/tests/e2e/management/test_memory_v2_e2e.py @@ -1,11 +1,15 @@ +import fcntl import hashlib import os +import tempfile +from collections.abc import Iterator from dataclasses import dataclass +from pathlib import Path from typing import Final import pytest from e2e_config import unique_marker -from e2e_http import Success, unwrap +from e2e_http import NoBody, Success, unwrap from lifecycle import ResourceManager from management_client import ManagementClient from memory_client import MemoryClient @@ -20,13 +24,13 @@ from models import ( KeyGenerateBody, LiteLLMParamsBody, MemoryCaptureBody, - MemoryEntriesData, MemoryEntryParams, MemoryLegacyParams, MemoryLegacyRows, - MemoryPolicyBody, MemoryResponsesBody, + MemorySettingsBody, MemoryStreamEvent, + MemoryTeamPermissionBody, MemoryWireResponse, TeamNewBody, UserNewBody, @@ -78,8 +82,17 @@ def memory_models(client: ManagementClient, resources: ResourceManager) -> Memor @pytest.fixture -def memory(client: ManagementClient) -> MemoryClient: - return MemoryClient(client.proxy) +def memory(client: ManagementClient) -> Iterator[MemoryClient]: + # Serialize global configuration changes across local pytest workers. + with (Path(tempfile.gettempdir()) / "litellm-memory-v2-e2e.lock").open("w") as lock: + fcntl.flock(lock, fcntl.LOCK_EX) + memory = MemoryClient(client.proxy) + original = memory.settings() + try: + yield memory + finally: + unwrap(memory.set_settings(original)) + fcntl.flock(lock, fcntl.LOCK_UN) @pytest.fixture @@ -110,14 +123,9 @@ def subjects(client: ManagementClient, memory: MemoryClient, resources: Resource return key keys: Final = tuple(create_key(user) for user in (owner, owner, other)) - policy: Final = unwrap( - memory.set_policy(MemoryPolicyBody(target_type="team", target_id=team, activation="automatic")) - ) - resources.defer(lambda: memory.delete_policy(policy.policy_id)) + unwrap(memory.set_settings(MemorySettingsBody(enabled=True, everyone=False, user_ids=[owner, other]))) resources.defer(lambda: memory.cleanup_user_entries(owner)) resources.defer(lambda: memory.cleanup_user_entries(other)) - resources.defer(lambda: memory.preference(keys[0], False)) - resources.defer(lambda: memory.preference(keys[2], False)) return MemorySubjects(owner=keys[0], sibling=keys[1], outsider=keys[2], user_id=owner, team_id=team) @@ -206,21 +214,16 @@ class TestMemoryV2: assert not public.has_memory_tools if endpoint == "responses": assert public.instructions is None - assert memory.entries(subjects.sibling) == [] + assert any(marker in entry.content for entry in memory.entries(subjects.sibling)) assert memory.entries(subjects.outsider) == [] - @pytest.mark.covers("mgmt.memory_v2.policy.opt_in") - def test_admin_selects_opt_in_or_automatic_and_key_override_wins( - self, - client: ManagementClient, - memory: MemoryClient, - subjects: MemorySubjects, - resources: ResourceManager, - memory_models: MemoryModels, + @pytest.mark.covers("mgmt.memory_v2.settings.enrollment") + def test_admin_enrollment_follows_user_and_disable_preserves_dashboard( + self, client: ManagementClient, memory: MemoryClient, subjects: MemorySubjects, memory_models: MemoryModels ) -> None: - unwrap(memory.set_policy(MemoryPolicyBody(target_type="team", target_id=subjects.team_id, activation="opt_in"))) + unwrap(memory.set_settings(MemorySettingsBody())) assert not memory.status(subjects.owner).active - marker: Final = f"unsaved-{unique_marker()}" + marker = f"unsaved-{unique_marker()}" unwrap( client.proxy.chat( subjects.owner, @@ -236,39 +239,26 @@ class TestMemoryV2: ) ) assert memory.entries(subjects.owner) == [] - memory.preference(subjects.owner, True) - assert memory.status(subjects.owner).active - assert memory.status(subjects.sibling).active + unwrap(memory.set_settings(MemorySettingsBody(enabled=True, everyone=False, user_ids=[subjects.user_id]))) + assert memory.status(subjects.owner).active and memory.status(subjects.sibling).active assert not memory.status(subjects.outsider).active - saved: Final = unwrap(memory.capture(subjects.owner, _fact(unique_marker()))) - assert saved.memory_id in [entry.memory_id for entry in memory.entries(subjects.owner)] - key_policy: Final = memory.policy_for_key(subjects.owner, "disabled") - resources.defer(lambda: memory.delete_policy(key_policy.policy_id)) + saved = unwrap(memory.capture(subjects.owner, _fact(unique_marker()))) + unwrap(memory.set_settings(MemorySettingsBody())) assert not memory.status(subjects.owner).active - assert memory.status(subjects.sibling).active - memory.policy_for_key(subjects.owner, "automatic") - memory.preference(subjects.owner, False) - assert memory.status(subjects.owner).active - assert not memory.status(subjects.sibling).active + assert unwrap(memory.read(subjects.sibling, saved.memory_id)).content == saved.content + _assert_denied(memory.capture(subjects.owner, _fact(unique_marker()))) @pytest.mark.covers("mgmt.memory_v2.entries.isolation") - def test_private_entries_are_isolated_even_for_sibling_keys_and_legacy_api( + def test_owner_keys_share_but_other_users_and_legacy_api_do_not( self, client: ManagementClient, memory: MemoryClient, subjects: MemorySubjects ) -> None: - marker: Final = unique_marker() - saved: Final = unwrap(memory.capture(subjects.owner, _fact(marker))) - assert memory.entries(subjects.owner)[0].memory_id == saved.memory_id - for key in (subjects.sibling, subjects.outsider): - assert memory.entries(key) == [] - _assert_denied(memory.delete_entry(key, saved.memory_id)) - result = client.proxy.transport.get( - "/v2/memory/entries", - headers=client.proxy.transport.bearer(key), - params=MemoryEntryParams(key_id=hashlib.sha256(subjects.owner.encode()).hexdigest()), - response_type=MemoryEntriesData, - ) - _assert_denied(result) - legacy: Final = client.proxy.transport.get( + saved = unwrap(memory.capture(subjects.owner, _fact(unique_marker()))) + assert memory.entries(subjects.sibling)[0].memory_id == saved.memory_id + assert memory.entries(subjects.outsider) == [] + assert memory.entries(subjects.outsider, MemoryEntryParams(user_id=subjects.user_id)) == [] + _assert_denied(memory.read(subjects.outsider, saved.memory_id)) + _assert_denied(memory.delete_entry(subjects.outsider, saved.memory_id)) + legacy = client.proxy.transport.get( "/v1/memory", headers=client.proxy.transport.bearer(subjects.sibling), params=MemoryLegacyParams(), @@ -278,15 +268,36 @@ class TestMemoryV2: assert saved.memory_id not in [row.memory_id for row in legacy.data.memories] assert memory.entries(subjects.owner)[0].content == saved.content - @pytest.mark.covers("mgmt.memory_v2.policy.admin_only") - def test_members_cannot_enable_or_broaden_memory(self, memory: MemoryClient, subjects: MemorySubjects) -> None: - before: Final = memory.status(subjects.owner) - for target_type, target_id in (("gateway", "*"), ("team", subjects.team_id), ("user", subjects.user_id)): - body = MemoryPolicyBody.model_validate( - {"target_type": target_type, "target_id": target_id, "activation": "automatic", "scope": "team"} + @pytest.mark.covers("mgmt.memory_v2.settings.admin_only") + def test_members_cannot_enable_memory(self, memory: MemoryClient, subjects: MemorySubjects) -> None: + before = memory.settings() + _assert_denied(memory.set_settings(MemorySettingsBody(enabled=True), caller=subjects.owner)) + assert memory.settings() == before + + @pytest.mark.covers("mgmt.memory_v2.entries.team_permissions") + def test_delegated_team_reads_allow_recall_but_not_edit_and_can_be_revoked( + self, client: ManagementClient, memory: MemoryClient, subjects: MemorySubjects + ) -> None: + saved = unwrap(memory.capture(subjects.owner, _fact(unique_marker()))) + assert memory.entries(subjects.outsider) == [] + for permissions, visible in ((["/spend/logs"], False), (["/v2/memory/entries"], True), ([], False)): + unwrap( + client.proxy.transport.post( + "/team/permissions_update", + headers=client.proxy.transport.master, + json=MemoryTeamPermissionBody(team_id=subjects.team_id, team_member_permissions=permissions), + response_type=NoBody, + ) ) - _assert_denied(memory.set_policy(body, caller=subjects.owner)) - assert memory.status(subjects.owner) == before + entries = memory.entries(subjects.outsider, MemoryEntryParams(team_id=subjects.team_id)) + assert bool(entries) is visible + if visible: + assert entries[0].memory_id == saved.memory_id and not entries[0].can_edit + assert unwrap(memory.read(subjects.outsider, saved.memory_id)).content == saved.content + _assert_denied(memory.update(subjects.outsider, saved.memory_id, _fact(unique_marker()))) + _assert_denied(memory.delete_entry(subjects.outsider, saved.memory_id)) + else: + _assert_denied(memory.read(subjects.outsider, saved.memory_id)) @pytest.mark.covers("mgmt.memory_v2.entries.correction_delete") def test_corrections_require_current_revision_and_deleted_memory_is_not_recalled( diff --git a/tests/e2e/models.py b/tests/e2e/models.py index 90db4a94732..7cd8cc24ac5 100644 --- a/tests/e2e/models.py +++ b/tests/e2e/models.py @@ -1308,34 +1308,33 @@ class ReadinessDetailsResponse(ReadinessResponse): success_callbacks: list[str] = [] -class MemoryPolicyBody(BaseModel): - target_type: Literal["gateway", "organization", "team", "project", "user", "key"] - target_id: str - activation: Literal["disabled", "opt_in", "automatic"] - scope: Literal["key", "user", "team", "project", "organization"] = "key" - - -class MemoryPolicyData(MemoryPolicyBody): - policy_id: str - - -class MemoryPreferenceBody(BaseModel): - enabled: bool +class MemorySettingsBody(BaseModel): + enabled: bool = False + everyone: bool = True + user_ids: list[str] = [] class MemoryStatusData(BaseModel): active: bool - activation: str - scope: str | None - opted_in: bool - policy_id: str | None + enabled: bool + user_id: str | None = None + team_ids: list[str] = [] + admin_view: bool = False class MemoryEntryParams(BaseModel): query: str = "" limit: int = 20 - key_id: str | None = None + user_id: str | None = None + team_id: str | None = None offset: int = 0 + before_updated_at: str | None = None + before_memory_id: str | None = None + + +class MemoryTeamPermissionBody(BaseModel): + team_id: str + team_member_permissions: list[str] class MemoryCaptureBody(BaseModel): @@ -1347,6 +1346,10 @@ class MemoryCaptureBody(BaseModel): class MemoryEntryData(BaseModel): + user_id: str | None = None + team_id: str | None = None + can_edit: bool = False + actor: str | None = None memory_id: str key: str title: str 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 64d0a2be176..91258c1b5de 100644 --- a/tests/test_litellm/proxy/memory/test_memory_v2_boundaries.py +++ b/tests/test_litellm/proxy/memory/test_memory_v2_boundaries.py @@ -21,30 +21,23 @@ from litellm.litellm_core_utils.prompt_templates.server_tool_responses import ( from litellm.proxy.memory.continuation import MemoryContinuation, MemoryContinuations, prefix_hashes from litellm.proxy.memory.gateway import GatewayMemoryLoop from litellm.proxy.memory.knowledge import MEMORY_TOOL_NAMES, execute_memory_tool -from litellm.proxy.memory.policy import MemoryAccess, MemoryIdentity, resolve_memory_access +from litellm.proxy.memory.policy import MemoryAccess, MemoryIdentity, memory_digest, resolve_memory_access from litellm.proxy.memory.responses import serve_memory_response from litellm.proxy.memory.store import MemoryStore -from litellm.types.memory_v2 import MemoryCapture, MemoryPolicy, MemorySearch +from litellm.types.memory_v2 import MemoryCapture, MemorySearch, MemorySettings _NOW: Final = datetime(2026, 9, 12, tzinfo=timezone.utc) _IDENTITY: Final = MemoryIdentity("a" * 64, "owner", "team", "project", "org", False) -_POLICY: Final = MemoryPolicy( - policy_id="policy", - target_type="team", - target_id="team", - activation="automatic", - scope="key", - updated_at=_NOW, - updated_by="admin", -) +_SETTINGS: Final = MemorySettings(enabled=True) _CAPTURE: Final = MemoryCapture(key="demo", title="Demo", content="Use port 8347", evidence="User selected this port") @pytest.fixture def prisma_edge() -> MagicMock: client = MagicMock() - client.db.litellm_memorypolicy.find_many = AsyncMock(return_value=[_POLICY]) - client.db.litellm_memorypreference.find_unique = AsyncMock(return_value=None) + client.db.litellm_config.find_unique = AsyncMock(return_value=SimpleNamespace(param_value=_SETTINGS.model_dump())) + client.db.litellm_usertable.find_unique = AsyncMock(return_value=None) + client.db.litellm_teamtable.find_many = AsyncMock(return_value=[]) table = client.db.litellm_memorytable table.find_unique = AsyncMock(return_value=None) table.find_first = AsyncMock(return_value=None) @@ -66,15 +59,25 @@ def prisma_edge() -> MagicMock: def store(client: MagicMock, identity: MemoryIdentity = _IDENTITY) -> MemoryStore: - return MemoryStore(client, MemoryAccess(identity, _POLICY, False)) + return MemoryStore(client, access_for(identity)) + + +def access_for(identity: MemoryIdentity = _IDENTITY) -> MemoryAccess: + return MemoryAccess( + identity, _SETTINGS, permission_revision=memory_digest(identity.namespace, identity.role, "False") + ) def row(**changes: object) -> LiteLLM_MemoryTable: return LiteLLM_MemoryTable.model_validate( { "memory_id": "entry", - "key": f"memory-v2:{_IDENTITY.namespace('key')}:demo", - "namespace": _IDENTITY.namespace("key"), + "key": f"memory-v2:{_IDENTITY.namespace}:demo", + "namespace": _IDENTITY.namespace, + "user_id": "owner", + "team_id": "team", + "organization_id": "org", + "owner_key_id": "a" * 64, "value": _CAPTURE.content, "metadata": json.dumps({"title": _CAPTURE.title, "evidence": _CAPTURE.evidence}), "created_at": _NOW, @@ -101,6 +104,7 @@ async def test_saved_response_reads_are_scoped_and_never_return_internal_input( prisma_edge: MagicMock, operation: str ) -> None: patch = MemoryContinuation( + permission_revision=access_for().permission_revision, replaces=1, response={"id": "resp_litellm_memory_test", "output": [{"type": "message", "content": []}]}, upstream_ids=("native=one",), @@ -122,7 +126,7 @@ async def test_saved_response_reads_are_scoped_and_never_return_internal_input( await serve_memory_response("resp_litellm_memory_test", request, route, store(prisma_edge), app) assert exc.value.status_code == (404 if operation == "missing" else 501) where = prisma_edge.db.litellm_memorycontinuation.find_first.call_args.kwargs["where"] - assert where["namespace"] == _IDENTITY.namespace("key") and where["key_id"] == _IDENTITY.key_id + assert where["namespace"] == _IDENTITY.namespace and where["key_id"] == _IDENTITY.key_id assert where["expires_at"]["gt"] <= datetime.now(timezone.utc) @@ -130,6 +134,7 @@ async def test_saved_response_reads_are_scoped_and_never_return_internal_input( @pytest.mark.parametrize("outcome", ["success", "already_missing", "upstream_error", "readonly"]) async def test_response_deletion_preserves_auth_paths_and_retry_state(prisma_edge: MagicMock, outcome: str) -> None: patch = MemoryContinuation( + permission_revision=access_for().permission_revision, replaces=1, response={"id": "resp_litellm_memory_test"}, upstream_ids=("native=one", "native=two"), @@ -181,36 +186,33 @@ async def test_response_deletion_preserves_auth_paths_and_retry_state(prisma_edg @pytest.mark.asyncio -async def test_policy_precedence_and_opt_in_are_resolved_from_database(prisma_edge: MagicMock) -> None: - team = _POLICY.model_copy(update={"activation": "opt_in"}) - key = _POLICY.model_copy(update={"target_type": "key", "target_id": "a" * 64, "activation": "disabled"}) - policies = prisma_edge.db.litellm_memorypolicy.find_many - policies.return_value = [team, key] - prisma_edge.db.litellm_memorypreference.find_unique.return_value = SimpleNamespace(enabled=True) - access = await resolve_memory_access(prisma_edge, _IDENTITY) - assert access.policy == key and not access.active and access.opted_in - policies.return_value = [team] - assert (await resolve_memory_access(prisma_edge, _IDENTITY)).active - prisma_edge.db.litellm_memorypreference.find_unique.return_value = None +async def test_flat_enrollment_follows_the_user_across_keys(prisma_edge: MagicMock) -> None: + config = prisma_edge.db.litellm_config.find_unique + config.return_value = None + assert not (await resolve_memory_access(prisma_edge, _IDENTITY)).active + config.return_value = SimpleNamespace( + param_value=MemorySettings(enabled=True, everyone=False, user_ids=("owner",)).model_dump() + ) + assert (await resolve_memory_access(prisma_edge, _IDENTITY)).active + sibling = MemoryIdentity("b" * 64, "owner", "team", "project", "org", False) + assert (await resolve_memory_access(prisma_edge, sibling)).active + config.return_value = SimpleNamespace( + param_value=MemorySettings(enabled=True, everyone=False, user_ids=("other",)).model_dump() + ) assert not (await resolve_memory_access(prisma_edge, _IDENTITY)).active - policies.return_value = [] - assert (await resolve_memory_access(prisma_edge, _IDENTITY)).namespace is None @pytest.mark.asyncio -@pytest.mark.parametrize("change", ["disabled", "scope", "missing", "readonly"]) -async def test_store_rechecks_policy_before_writing(prisma_edge: MagicMock, change: str) -> None: +@pytest.mark.parametrize("change", ["disabled", "unenrolled", "missing", "readonly"]) +async def test_store_rechecks_access_before_writing(prisma_edge: MagicMock, change: str) -> None: identity = _IDENTITY - if change == "disabled": - prisma_edge.db.litellm_memorypolicy.find_many.return_value = [ - _POLICY.model_copy(update={"activation": "disabled"}) - ] - elif change == "scope": - prisma_edge.db.litellm_memorypolicy.find_many.return_value = [_POLICY.model_copy(update={"scope": "team"})] - elif change == "missing": - prisma_edge.db.litellm_memorypolicy.find_many.return_value = [] - else: + if change == "missing": + prisma_edge.db.litellm_config.find_unique.return_value = None + elif change == "readonly": identity = MemoryIdentity("a" * 64, "owner", "team", "project", "org", True) + else: + config = MemorySettings(enabled=change != "disabled", everyone=False, user_ids=("other",)) + prisma_edge.db.litellm_config.find_unique.return_value = SimpleNamespace(param_value=config.model_dump()) with pytest.raises(HTTPException) as exc: await store(prisma_edge, identity).capture(_CAPTURE) assert exc.value.status_code == 403 @@ -227,8 +229,8 @@ async def test_search_applies_fuzzy_ranking_before_pagination_with_namespace_bou entries = await store(prisma_edge).search(MemorySearch(query="prto demo", limit=1, offset=1)) assert [entry.memory_id for entry in entries] == ["second"] query = prisma_edge.db.litellm_memorytable.find_many.call_args.kwargs - assert query["where"]["namespace"] == _IDENTITY.namespace("key") - assert query["take"] == 1000 + assert query["where"]["AND"][0]["AND"][0]["namespace"] == {"startswith": "v2:"} + assert query["take"] == 128 @pytest.mark.asyncio @@ -238,14 +240,12 @@ async def test_read_and_delete_cannot_address_another_namespace(prisma_edge: Mag await memory.read("foreign-entry") assert exc.value.status_code == 404 assert prisma_edge.db.litellm_memorytable.find_first.call_args.kwargs["where"] == { - "namespace": _IDENTITY.namespace("key"), - "memory_id": "foreign-entry", + "AND": [access_for().visible_rows(), {"memory_id": "foreign-entry"}], } - prisma_edge.db.litellm_memorypolicy.find_many.return_value = [_POLICY.model_copy(update={"activation": "disabled"})] + prisma_edge.db.litellm_config.find_unique.return_value = SimpleNamespace(param_value=MemorySettings().model_dump()) assert await memory.delete("entry") assert prisma_edge.db.litellm_memorytable.delete_many.call_args.kwargs["where"] == { - "namespace": _IDENTITY.namespace("key"), - "memory_id": "entry", + "AND": [access_for().visible_rows(write=True), {"memory_id": "entry"}], } with pytest.raises(HTTPException) as inactive: await memory.read("entry") @@ -258,14 +258,12 @@ async def test_identical_capture_is_idempotent_and_new_capture_has_scoped_identi table = prisma_edge.db.litellm_memorytable table.create.return_value = row() - wrapped: Final = MemoryStore( - SimpleNamespace(db=PrismaWrapper(prisma_edge.db)), MemoryAccess(_IDENTITY, _POLICY, False) - ) + wrapped: Final = MemoryStore(SimpleNamespace(db=PrismaWrapper(prisma_edge.db)), access_for()) saved = await wrapped.capture(_CAPTURE) assert saved.content == _CAPTURE.content data = table.create.call_args.kwargs["data"] - assert data["namespace"] == _IDENTITY.namespace("key") and data["user_id"] == "owner" and data["team_id"] == "team" - assert data["key"].startswith("memory-v2:" + _IDENTITY.namespace("key") + ":") + assert data["namespace"] == _IDENTITY.namespace and data["user_id"] == "owner" and data["team_id"] == "team" + assert data["key"].startswith("memory-v2:" + _IDENTITY.namespace + ":") table.find_unique.return_value = row() assert await wrapped.capture(_CAPTURE) == saved table.create.assert_awaited_once() @@ -276,18 +274,19 @@ async def test_identical_capture_is_idempotent_and_new_capture_has_scoped_identi @pytest.mark.parametrize("revoked", [False, True]) async def test_capture_rechecks_policy_on_its_transaction_connection(prisma_edge: MagicMock, revoked: bool) -> None: prisma_edge.db.litellm_memorytable.create.return_value = row() - prisma_edge.db.litellm_memorypolicy.find_many.side_effect = [ - [_POLICY], + prisma_edge.db.litellm_config.find_unique.side_effect = [ + SimpleNamespace(param_value=_SETTINGS.model_dump()), RuntimeError("The only pooled connection belongs to the active transaction"), ] transaction = SimpleNamespace( litellm_memorytable=prisma_edge.db.litellm_memorytable, - litellm_memorypolicy=SimpleNamespace( - find_many=AsyncMock( - return_value=[_POLICY.model_copy(update={"activation": "disabled"})] if revoked else [_POLICY] + litellm_config=SimpleNamespace( + find_unique=AsyncMock( + return_value=SimpleNamespace(param_value=MemorySettings(enabled=not revoked).model_dump()) ) ), - litellm_memorypreference=prisma_edge.db.litellm_memorypreference, + litellm_usertable=prisma_edge.db.litellm_usertable, + litellm_teamtable=prisma_edge.db.litellm_teamtable, execute_raw=AsyncMock(), ) prisma_edge.db.tx.return_value.__aenter__.return_value = transaction @@ -304,7 +303,7 @@ async def test_capture_rechecks_policy_on_its_transaction_connection(prisma_edge async def test_capture_rejects_stale_or_conflicting_replacements(prisma_edge: MagicMock, race: str) -> None: table = prisma_edge.db.litellm_memorytable table.find_unique.return_value = ( - None if race == "missing" else row(namespace="foreign" if race == "foreign" else _IDENTITY.namespace("key")) + None if race == "missing" else row(namespace="foreign" if race == "foreign" else _IDENTITY.namespace) ) table.update_many.return_value = 0 if race == "concurrent" else 1 correction = _CAPTURE.model_copy( @@ -320,7 +319,7 @@ async def test_capture_rejects_stale_or_conflicting_replacements(prisma_edge: Ma if race == "concurrent": where = table.update_many.call_args.kwargs["where"] assert ( - where["namespace"] == _IDENTITY.namespace("key") + where["namespace"] == _IDENTITY.namespace and where["updated_at"] == _NOW and where["value"] == _CAPTURE.content ) @@ -353,7 +352,7 @@ async def test_tool_argument_errors_and_revocation_return_receipts_without_writi assert missing.output == {"error": "Memory not found", "status": 404} unknown = await execute_memory_tool(memory, {"id": "a", "name": "other_tool", "arguments": {}}, "checkpoint") assert unknown.output == {"error": "Unknown memory tool"} - prisma_edge.db.litellm_memorypolicy.find_many.return_value = [] + prisma_edge.db.litellm_config.find_unique.return_value = None revoked = await execute_memory_tool( memory, {"id": "a", "name": "litellm_memory_capture", "arguments": {"observations": []}}, "checkpoint" ) @@ -457,7 +456,9 @@ async def test_restore_preserves_hidden_memory_tool_results_and_client_cache_mar prisma_edge.db.litellm_memorycontinuation.find_many.return_value = [ SimpleNamespace( id=continuations.identifier(anchor), - payload=MemoryContinuation(replaces=1, replacement=replacement).model_dump(), + payload=MemoryContinuation( + replaces=1, replacement=replacement, permission_revision=access_for().permission_revision + ).model_dump(), ) ] restored: Final = await continuations.restore(items) @@ -666,7 +667,9 @@ async def test_duplicate_directives_preserve_each_current_cache_breakpoint(prism items = (first, second, assistant) continuations = MemoryContinuations(store(prisma_edge), "anthropic_messages") patch = MemoryContinuation( - replaces=3, replacement=({"role": "user", "content": "Memory reference"}, second, first, assistant) + permission_revision=access_for().permission_revision, + replaces=3, + replacement=({"role": "user", "content": "Memory reference"}, second, first, assistant), ) prisma_edge.db.litellm_memorycontinuation.find_many.return_value = [ SimpleNamespace( @@ -713,11 +716,12 @@ async def test_replica_lag_cannot_authorize_memory_after_primary_revocation( writer = MagicMock(spec=PrismaWrapper) reader = MagicMock(spec=PrismaWrapper) - writer.litellm_memorypolicy = SimpleNamespace( - find_many=AsyncMock(return_value=[_POLICY.model_copy(update={"activation": "disabled"})]) + writer.litellm_config = SimpleNamespace(find_unique=AsyncMock(return_value=None)) + reader.litellm_config = SimpleNamespace( + find_unique=AsyncMock(return_value=SimpleNamespace(param_value=_SETTINGS.model_dump())) ) - reader.litellm_memorypolicy = SimpleNamespace(find_many=AsyncMock(return_value=[_POLICY])) - writer.litellm_memorypreference = SimpleNamespace(find_unique=AsyncMock(return_value=None)) + writer.litellm_usertable = prisma_edge.db.litellm_usertable + writer.litellm_teamtable = prisma_edge.db.litellm_teamtable writer.litellm_memorytable = prisma_edge.db.litellm_memorytable writer.is_connected = MagicMock(return_value=False) reader.is_connected = MagicMock(return_value=False) @@ -729,9 +733,9 @@ async def test_replica_lag_cannot_authorize_memory_after_primary_revocation( client = SimpleNamespace(db=routed) access = await resolve_memory_access(client, _IDENTITY) assert not access.active - reader.litellm_memorypolicy.find_many.assert_not_awaited() + reader.litellm_config.find_unique.assert_not_awaited() with pytest.raises(HTTPException) as exc: - await MemoryStore(client, MemoryAccess(_IDENTITY, _POLICY, False)).capture(_CAPTURE) + await MemoryStore(client, access_for()).capture(_CAPTURE) assert exc.value.status_code == 403 writer.litellm_memorytable.create.assert_not_awaited() @@ -742,16 +746,16 @@ async def test_unconfigured_gate_caches_presence_without_caching_authorization(p from litellm.proxy.memory.policy import gateway_memory_is_configured cache = DualCache() - policies = prisma_edge.db.litellm_memorypolicy.find_many - policies.return_value = [] + config = prisma_edge.db.litellm_config.find_unique + config.return_value = None assert not await gateway_memory_is_configured(prisma_edge, cache) assert not await gateway_memory_is_configured(prisma_edge, cache) - policies.assert_awaited_once() - assert policies.call_args.kwargs == {"take": 1} + config.assert_awaited_once() + assert config.call_args.kwargs == {"where": {"param_name": "memory_v2"}} enabled_cache = DualCache() - policies.return_value = [_POLICY] + config.return_value = SimpleNamespace(param_value=_SETTINGS.model_dump()) assert await gateway_memory_is_configured(prisma_edge, enabled_cache) - policies.return_value = [_POLICY.model_copy(update={"activation": "disabled"})] + config.return_value = None assert await gateway_memory_is_configured(prisma_edge, enabled_cache) assert not (await resolve_memory_access(prisma_edge, _IDENTITY)).active @@ -996,18 +1000,18 @@ async def test_backend_activation_invalidates_a_gateway_negative_hint_without_pu ) gateway_cache = DualCache(redis_cache=redis if share_auth_cache else None) backend_cache = DualCache(redis_cache=redis if share_auth_cache else None) - policies = prisma_edge.db.litellm_memorypolicy.find_many + config = prisma_edge.db.litellm_config.find_unique with patch.multiple( # test-quality-ok: Inject external worker caches and Redis; run real invalidation. "litellm.proxy.proxy_server", user_api_key_cache=backend_cache, redis_usage_cache=redis ): - policies.return_value = [] + config.return_value = None assert not await gateway_memory_is_configured(prisma_edge, gateway_cache) assert not await gateway_memory_is_configured(prisma_edge, gateway_cache) - policies.assert_awaited_once() - policies.return_value = [_POLICY] + config.assert_awaited_once() + config.return_value = SimpleNamespace(param_value=_SETTINGS.model_dump()) await invalidate_memory_configuration() assert await gateway_memory_is_configured(prisma_edge, gateway_cache) - assert policies.await_count == 2 + assert config.await_count == 2 @pytest.mark.asyncio @@ -1024,15 +1028,15 @@ async def test_redis_circuit_breaker_falls_back_to_primary_configuration(prisma_ async_delete_cache=AsyncMock(side_effect=RedisCircuitBreakerOpenError("open")), ) cache = DualCache() - prisma_edge.db.litellm_memorypolicy.find_many.return_value = [] + prisma_edge.db.litellm_config.find_unique.return_value = None with patch.multiple( # test-quality-ok: Inject external Redis failure and local worker cache; exercise real fallback. "litellm.proxy.proxy_server", user_api_key_cache=cache, redis_usage_cache=redis ): assert not await gateway_memory_is_configured(prisma_edge, cache) - prisma_edge.db.litellm_memorypolicy.find_many.return_value = [_POLICY] + prisma_edge.db.litellm_config.find_unique.return_value = SimpleNamespace(param_value=_SETTINGS.model_dump()) assert await gateway_memory_is_configured(prisma_edge, cache) await invalidate_memory_configuration() - assert prisma_edge.db.litellm_memorypolicy.find_many.await_count == 2 + assert prisma_edge.db.litellm_config.find_unique.await_count == 2 redis.async_get_cache.assert_awaited_once() @@ -1044,7 +1048,7 @@ async def test_full_scope_blocks_creation_but_permits_correction_and_reclaimed_c 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")} + assert table.count.call_args.kwargs["where"] == {"namespace": _IDENTITY.namespace} 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.side_effect = [row(), row(value="Corrected", updated_at=_NOW + timedelta(seconds=1))] @@ -1078,20 +1082,41 @@ async def test_continuation_quota_rejects_excess_without_writing( async def test_continuation_quota_shares_namespace_lock_across_keys_and_allows_replacements( prisma_edge: MagicMock, ) -> None: - user_policy = _POLICY.model_copy(update={"scope": "user"}) - prisma_edge.db.litellm_memorypolicy.find_many.return_value = [user_policy] prisma_edge.db.query_raw.return_value = [{"key_count": 255, "bytes": 32 * 1024 * 1024}] other_key = MemoryIdentity("b" * 64, "owner", "team", "project", "org", False) for identity in (_IDENTITY, other_key): - continuations = MemoryContinuations( - MemoryStore(prisma_edge, MemoryAccess(identity, user_policy, False)), "aresponses" - ) + continuations = MemoryContinuations(MemoryStore(prisma_edge, access_for(identity)), "aresponses") await continuations.save("response", MemoryContinuation(replaces=1, response={"text": "é漢字"})) query = prisma_edge.db.query_raw.call_args.args - assert query[1:4] == (identity.namespace("user"), identity.key_id, [continuations.identifier("response")]) + assert query[1:4] == (identity.namespace, identity.key_id, [continuations.identifier("response")]) assert json.loads(query[4])[0]["response"]["text"] == "é漢字" locks = prisma_edge.db.execute_raw.call_args_list assert locks[0] == locks[1] cleanup = prisma_edge.db.litellm_memorycontinuation.delete_many.call_args.kwargs["where"] - assert cleanup["namespace"] == _IDENTITY.namespace("user") and "key_id" not in cleanup + assert cleanup["namespace"] == _IDENTITY.namespace and "key_id" not in cleanup assert prisma_edge.db.litellm_memorycontinuation.upsert.await_count == 2 + + +@pytest.mark.asyncio +async def test_memory_lookup_failure_leaves_inference_unchanged_but_never_leaks_owned_response_ids( + prisma_edge: MagicMock, +) -> None: + from unittest.mock import patch + + from litellm.caching.caching import DualCache + from litellm.proxy._types import UserAPIKeyAuth + from litellm.proxy.memory.gateway import process_gateway_memory + + prisma_edge.db.litellm_config.find_unique.side_effect = RuntimeError("database unavailable") + caller = UserAPIKeyAuth(user_id="owner", token="a" * 64) + with patch.multiple( # test-quality-ok: Inject unavailable external DB and an empty worker cache. + "litellm.proxy.proxy_server", prisma_client=prisma_edge, user_api_key_cache=DualCache() + ): + assert await process_gateway_memory({"messages": []}, request(), caller, "acompletion") is None + with pytest.raises(HTTPException) as exc: + await process_gateway_memory( + {"previous_response_id": "resp_litellm_memory_private"}, request(), caller, "aresponses" + ) + assert exc.value.status_code == 404 + prisma_edge.db.litellm_memorytable.find_many.assert_not_awaited() + prisma_edge.db.litellm_memorytable.create.assert_not_awaited() 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 654793fbc65..72126d12643 100644 --- a/tests/test_litellm/proxy/memory/test_memory_v2_management.py +++ b/tests/test_litellm/proxy/memory/test_memory_v2_management.py @@ -1,484 +1,385 @@ -"""Exercise real policy administration and self-service logic at the database edge.""" +"""Admin activation and record permissions, replacing only external database/cache edges.""" +import json from collections.abc import Iterator -from datetime import datetime, timezone +from datetime import datetime, timedelta, timezone from types import SimpleNamespace from unittest.mock import AsyncMock, MagicMock, patch import pytest from fastapi import FastAPI, HTTPException from fastapi.testclient import TestClient +from prisma.models import LiteLLM_MemoryTable -from litellm.proxy._types import UI_TEAM_ID, LitellmUserRoles, UserAPIKeyAuth +from litellm.caching.caching import DualCache +from litellm.proxy._types import UI_TEAM_ID, KeyManagementRoutes, LiteLLM_TeamTable, LitellmUserRoles, UserAPIKeyAuth from litellm.proxy.auth.user_api_key_auth import user_api_key_auth +from litellm.proxy.management_helpers.record_permissions import can_read_team_records from litellm.proxy.memory import management -from litellm.proxy.memory.policy import MemoryIdentity, memory_digest, resolve_memory_access -from litellm.types.memory_v2 import MemoryCapture, MemoryPolicy, MemoryPolicyInput, MemoryPreference +from litellm.proxy.memory.continuation import MemoryContinuation, MemoryContinuations +from litellm.proxy.memory.policy import MemoryIdentity, resolve_memory_access +from litellm.proxy.memory.store import MemoryStore +from litellm.types.memory_v2 import ( + MemoryCapture, + MemoryCatalogRequest, + MemoryRecallRequest, + MemorySearch, + MemorySettings, +) + +_NOW = datetime(2026, 9, 12, tzinfo=timezone.utc) +_CAPTURE = MemoryCapture(key="demo", title="Demo port", content="Use port 8123", evidence="User selected port 8123") @pytest.fixture def database() -> Iterator[MagicMock]: client = MagicMock() for name in ( - "litellm_memorypolicy", - "litellm_memorypreference", + "litellm_config", "litellm_memorytable", "litellm_teamtable", - "litellm_projecttable", - "litellm_organizationtable", - "litellm_organizationmembership", - "litellm_verificationtoken", "litellm_usertable", + "litellm_memorycontinuation", ): table = getattr(client.db, name) table.find_unique = AsyncMock(return_value=None) table.find_first = AsyncMock(return_value=None) table.find_many = AsyncMock(return_value=[]) table.upsert = AsyncMock() - table.delete = AsyncMock() table.delete_many = AsyncMock(return_value=0) - table.create = AsyncMock() + table.create = AsyncMock(return_value=row()) table.count = AsyncMock(return_value=0) + table.update_many = AsyncMock(return_value=1) 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, - "members_with_roles": [{"user_id": "team-admin", "role": "admin"}], - } - client.db.litellm_projecttable.find_unique.return_value = {"project_id": "project", "team_id": "team"} - client.db.litellm_verificationtoken.find_unique.return_value = { - "token": "a" * 64, - "user_id": "owner", - "team_id": "team", - "org_id": "explicit-org", - } - client.db.litellm_usertable.find_unique.return_value = {"user_id": "owner"} - client.db.litellm_organizationtable.find_unique.return_value = { - "organization_id": "org", - "budget_id": "budget", - "created_by": "admin", - "updated_by": "admin", - } - with patch.multiple( # test-quality-ok: Replace only the external database and cache; exercise actual authorization and repositories. - "litellm.proxy.proxy_server", prisma_client=client, user_api_key_cache=MagicMock(async_delete_cache=AsyncMock()) + client.db.litellm_usertable.find_unique.return_value = {"user_id": "owner", "teams": ["team"]} + with patch.multiple( # test-quality-ok: Inject the external database/cache; run real handlers and authorization. + "litellm.proxy.proxy_server", prisma_client=client, user_api_key_cache=DualCache() ): yield client -def auth(user: str = "owner", role: LitellmUserRoles = LitellmUserRoles.INTERNAL_USER) -> UserAPIKeyAuth: - return UserAPIKeyAuth(token="a" * 64, user_id=user, user_role=role, team_id="team", org_id="explicit-org") +def auth(user: str | None = "owner", role: LitellmUserRoles = LitellmUserRoles.INTERNAL_USER) -> UserAPIKeyAuth: + return UserAPIKeyAuth(token="a" * 64, user_id=user, user_role=role, team_id="team", org_id="org") -def test_management_search_rejects_excessive_terms_before_database_work(database: MagicMock) -> None: - app = FastAPI() - app.include_router(management.router) - app.dependency_overrides[user_api_key_auth] = auth - with TestClient(app) as client: - response = client.get("/v2/memory/entries", params={"query": ",".join(f"term{i}" for i in range(17))}) - assert response.status_code == 422 - assert "at most 16 distinct search terms" in response.text - database.db.litellm_memorytable.find_many.assert_not_awaited() +def configure(database: MagicMock, **settings: object) -> None: + database.db.litellm_config.find_unique.return_value = SimpleNamespace( + param_value=MemorySettings.model_validate({"enabled": True, **settings}).model_dump() + ) -def policy(**changes: object) -> MemoryPolicy: - return MemoryPolicy.model_validate( +def team( + *, member: str = "owner", role: str = "user", permissions: tuple[str, ...] = (), **changes: object +) -> LiteLLM_TeamTable: + return LiteLLM_TeamTable.model_validate( { - "policy_id": memory_digest("gateway", "*"), - "target_type": "gateway", - "target_id": "*", - "activation": "automatic", - "scope": "key", - "updated_by": "admin", - "updated_at": datetime(2026, 9, 12, tzinfo=timezone.utc), + "team_id": "team", + "team_alias": "Engineering", + "organization_id": "org", + "members_with_roles": [{"user_id": member, "role": role}], + "team_member_permissions": list(permissions), + **changes, + } + ) + + +def row(**changes: object) -> LiteLLM_MemoryTable: + namespace = MemoryIdentity.from_auth(auth()).namespace + return LiteLLM_MemoryTable.model_validate( + { + "memory_id": "entry", + "key": f"memory-v2:{namespace}:demo", + "namespace": namespace, + "value": _CAPTURE.content, + "metadata": json.dumps({"title": _CAPTURE.title, "evidence": _CAPTURE.evidence}), + "user_id": "owner", + "team_id": "team", + "organization_id": "org", + "owner_key_id": "a" * 64, + "created_by": "owner", + "created_at": _NOW, + "updated_at": _NOW, **changes, } ) @pytest.mark.asyncio -async def test_user_preference_and_memories_follow_owner_across_keys(database: MagicMock) -> None: - database.db.litellm_memorypolicy.find_many.return_value = [policy(scope="user", activation="opt_in")] - first = MemoryIdentity.from_auth(auth()) - second = MemoryIdentity.from_auth(auth().model_copy(update={"token": "b" * 64})) - other = MemoryIdentity.from_auth(auth("other").model_copy(update={"token": "c" * 64})) - other_org = MemoryIdentity.from_auth(auth().model_copy(update={"org_id": "other-org"})) - assert not (await resolve_memory_access(database, first)).active - await management.set_preference(MemoryPreference(enabled=True), auth()) - written = database.db.litellm_memorypreference.upsert.call_args.kwargs["data"]["create"] - - async def preference(*, where: dict[str, str]) -> SimpleNamespace | None: - return SimpleNamespace(enabled=written["enabled"]) if where["subject"] == written["subject"] else None - - database.db.litellm_memorypreference.find_unique.side_effect = preference - one = await resolve_memory_access(database, first) - two = await resolve_memory_access(database, second) - assert one.active and two.active and one.namespace == two.namespace - assert not (await resolve_memory_access(database, other)).active - assert one.namespace != (await resolve_memory_access(database, other_org)).namespace - database.db.litellm_memorypolicy.find_many.return_value = [policy(scope="user", activation="disabled")] - assert not (await resolve_memory_access(database, second)).active - - -@pytest.mark.asyncio -async def test_dashboard_resolves_saved_contributor_names_and_identifies_preference_owner(database: MagicMock) -> None: - database.db.litellm_memorypolicy.find_many.return_value = [policy(scope="team")] - database.db.litellm_usertable.find_unique.return_value = {"user_id": "owner", "user_alias": "Alex Rivera"} - status = await management.get_status(None, auth()) - assert status.user_id == "owner" and status.user_name == "Alex Rivera" - now = datetime(2026, 9, 14, tzinfo=timezone.utc) - database.db.litellm_memorytable.find_many.return_value = [ - SimpleNamespace( - memory_id="entry", - key="memory-v2:namespace:demo", - value="Use port 8123", - metadata={}, - updated_at=now, - created_at=now, - created_by="contributor", - ) - ] - database.db.litellm_usertable.find_many.return_value = [ - SimpleNamespace(user_id="contributor", user_alias="Jamie Davis", user_email=None) - ] - entries = await management.list_entries("", 20, 0, None, auth()) - assert entries[0].actor == "contributor" and entries[0].actor_name == "Jamie Davis" - assert database.db.litellm_usertable.find_many.call_args.kwargs["where"] == {"user_id": {"in": ["contributor"]}} - database.db.litellm_usertable.find_many.return_value = [] - legacy = await management.list_entries("", 20, 0, None, auth()) - assert legacy[0].actor == "contributor" and legacy[0].actor_name is None - - -@pytest.mark.asyncio -async def test_admin_manual_capture_uses_actual_author_instead_of_selected_key_owner(database: MagicMock) -> None: - database.db.litellm_memorypolicy.find_many.return_value = [policy(scope="user")] - now = datetime(2026, 9, 14, tzinfo=timezone.utc) - database.db.litellm_memorytable.create.return_value = SimpleNamespace( - memory_id="entry", - key="memory-v2:namespace:demo", - value="Use port 8123", - metadata={}, - updated_at=now, - created_at=now, - created_by="admin", +async def test_default_off_and_proxy_admin_can_enable_selected_users(database: MagicMock) -> None: + admin = auth("admin", LitellmUserRoles.PROXY_ADMIN) + assert (await management.get_settings(admin)) == MemorySettings() + assert not (await management.get_status(auth())).active + database.db.litellm_usertable.find_many.return_value = [SimpleNamespace(user_id="owner")] + saved = await management.set_settings( + MemorySettings(enabled=True, everyone=False, user_ids=("owner", "owner")), admin ) - saved = await management.capture_entry( - MemoryCapture(key="demo", title="Demo", content="Use port 8123", evidence="Admin correction"), - "a" * 64, - auth("admin", LitellmUserRoles.PROXY_ADMIN), - ) - data = database.db.litellm_memorytable.create.call_args.kwargs["data"] - assert data["created_by"] == data["updated_by"] == saved.actor == "admin" - assert data["user_id"] == "owner" - - -@pytest.mark.asyncio -@pytest.mark.parametrize( - "target,target_id", - [ - ("gateway", "*"), - ("organization", "org"), - ("team", "team"), - ("project", "project"), - ("key", "a" * 64), - ("user", "owner"), - ], -) -async def test_proxy_admin_can_configure_each_existing_target(database: MagicMock, target: str, target_id: str) -> None: - request = MemoryPolicyInput.model_validate( - {"target_type": target, "target_id": target_id, "activation": "opt_in", "scope": "key"} - ) - database.db.litellm_memorypolicy.upsert.return_value = policy(**request.model_dump()) - result = await management.set_policy(request, auth("admin", LitellmUserRoles.PROXY_ADMIN)) - assert result.activation == "opt_in" - written = database.db.litellm_memorypolicy.upsert.call_args.kwargs - assert written["where"]["policy_id"] == memory_digest(target, target_id) - assert written["data"]["create"]["updated_by"] == "admin" - assert written["data"]["update"]["target_id"] == target_id - - -@pytest.mark.asyncio -@pytest.mark.parametrize("target,target_id", [("team", "team"), ("project", "project"), ("key", "a" * 64)]) -async def test_team_admin_can_configure_owned_targets_but_cannot_broaden_to_user_scope( - database: MagicMock, target: str, target_id: str -) -> None: - request = MemoryPolicyInput.model_validate( - {"target_type": target, "target_id": target_id, "activation": "automatic", "scope": "team"} - ) - database.db.litellm_memorypolicy.upsert.return_value = policy(**request.model_dump()) - assert (await management.set_policy(request, auth("team-admin"))).scope == "team" - with pytest.raises(HTTPException) as denied: - await management.set_policy(request.model_copy(update={"scope": "user"}), auth("team-admin")) - assert denied.value.status_code == 403 - database.db.litellm_memorypolicy.upsert.assert_awaited_once() - - -@pytest.mark.asyncio -async def test_org_membership_is_checked_for_the_selected_organization(database: MagicMock) -> None: - request = MemoryPolicyInput( - target_type="organization", target_id="org", activation="automatic", scope="organization" - ) - database.db.litellm_organizationmembership.find_first.return_value = SimpleNamespace(user_role="org_admin") - database.db.litellm_memorypolicy.upsert.return_value = policy(**request.model_dump()) - assert (await management.set_policy(request, auth("org-admin"))).scope == "organization" - assert database.db.litellm_organizationmembership.find_first.call_args.kwargs["where"] == { - "organization_id": "org", - "user_id": "org-admin", - "user_role": "org_admin", - } - database.db.litellm_organizationmembership.find_first.return_value = None - with pytest.raises(HTTPException) as denied: - await management.set_policy(request, auth("org-admin")) - assert denied.value.status_code == 403 - database.db.litellm_memorypolicy.upsert.assert_awaited_once() + assert saved.user_ids == ("owner",) + written = database.db.litellm_config.upsert.call_args.kwargs["data"]["update"]["param_value"] + database.db.litellm_config.find_unique.return_value = SimpleNamespace(param_value=written) + assert (await management.get_status(auth())).active + assert (await management.get_status(auth().model_copy(update={"token": "b" * 64}))).active + assert not (await management.get_status(auth("other"))).active + assert not (await management.get_status(auth(None))).active @pytest.mark.asyncio @pytest.mark.parametrize( "role", - [LitellmUserRoles.INTERNAL_USER, LitellmUserRoles.INTERNAL_USER_VIEW_ONLY, LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY], + [ + LitellmUserRoles.INTERNAL_USER, + LitellmUserRoles.INTERNAL_USER_VIEW_ONLY, + LitellmUserRoles.ORG_ADMIN, + LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY, + ], ) -async def test_non_administrator_cannot_change_gateway_policy(database: MagicMock, role: LitellmUserRoles) -> None: - with pytest.raises(HTTPException) as denied: - await management.set_policy( - MemoryPolicyInput(target_type="gateway", target_id="*", activation="automatic"), auth(role=role) - ) - assert denied.value.status_code == 403 - database.db.litellm_memorypolicy.upsert.assert_not_awaited() +async def test_only_proxy_admin_can_change_activation(database: MagicMock, role: LitellmUserRoles) -> None: + database.db.litellm_teamtable.find_many.return_value = [team(role="admin")] + with pytest.raises(HTTPException) as exc: + await management.set_settings(MemorySettings(enabled=True), auth(role=role)) + assert exc.value.status_code == 403 + database.db.litellm_config.upsert.assert_not_awaited() @pytest.mark.asyncio -async def test_policy_listing_requires_owned_target_and_preserves_pagination(database: MagicMock) -> None: - database.db.litellm_memorypolicy.find_many.return_value = [policy(target_type="team", target_id="team")] - assert len(await management.list_policies("team", "team", 100, auth("team-admin"))) == 1 - assert database.db.litellm_memorypolicy.find_many.call_args.kwargs == { - "where": {"target_type": "team", "target_id": "team"}, - "take": 100, - "skip": 100, - "order": {"policy_id": "asc"}, - } - for caller, target, expected in [(auth(), None, 403), (auth("admin", LitellmUserRoles.PROXY_ADMIN), "team", 400)]: - with pytest.raises(HTTPException) as denied: - await management.list_policies(target, None, 0, caller) - assert denied.value.status_code == expected - assert ( - len(await management.list_policies(None, None, 0, auth("admin", LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY))) == 1 - ) - - -@pytest.mark.asyncio -async def test_policy_delete_rechecks_administrator_and_reports_missing(database: MagicMock) -> None: - table = database.db.litellm_memorypolicy - table.find_unique.return_value = policy() - with pytest.raises(HTTPException) as denied: - await management.delete_policy("policy", auth()) - assert denied.value.status_code == 403 - table.delete.assert_not_awaited() - assert (await management.delete_policy("policy", auth("admin", LitellmUserRoles.PROXY_ADMIN))).status_code == 204 - table.find_unique.return_value = None - with pytest.raises(HTTPException) as missing: - await management.delete_policy("policy", auth("admin", LitellmUserRoles.PROXY_ADMIN)) - assert missing.value.status_code == 404 - - -@pytest.mark.asyncio -async def test_preference_updates_are_bound_to_authenticated_subject(database: MagicMock) -> None: - assert not (await management.get_preference(auth())).enabled - table = database.db.litellm_memorypreference - assert (await management.set_preference(MemoryPreference(enabled=True), auth())).enabled - assert table.upsert.call_args.kwargs["where"] == {"subject": memory_digest("user", "owner")} - table.find_unique.return_value = SimpleNamespace(enabled=True) - assert (await management.get_preference(auth())).enabled - assert not (await management.set_preference(MemoryPreference(enabled=False), auth())).enabled - assert table.delete_many.call_args.kwargs["where"] == {"subject": memory_digest("user", "owner")} - with pytest.raises(HTTPException) as denied: - await management.set_preference( - MemoryPreference(enabled=False), auth(role=LitellmUserRoles.INTERNAL_USER_VIEW_ONLY) - ) - assert denied.value.status_code == 403 - - -@pytest.mark.asyncio -async def test_dashboard_session_can_select_own_key_but_not_another_users_key(database: MagicMock) -> None: - database.db.litellm_memorypolicy.find_many.return_value = [policy()] - session = auth().model_copy(update={"token": "session", "team_id": UI_TEAM_ID}) - own = await management.access_for_key(session, "a" * 64) - assert own.active and own.identity.organization_id == "explicit-org" - assert own.namespace == MemoryIdentity.from_auth(auth()).namespace("key") - assert (await management.get_status(None, auth())).active - database.db.litellm_verificationtoken.find_unique.return_value["user_id"] = "someone-else" - with pytest.raises(HTTPException) as denied: - await management.access_for_key(session, "a" * 64) - assert denied.value.status_code == 403 - database.db.litellm_verificationtoken.find_unique.return_value = None - with pytest.raises(HTTPException) as missing: - await management.access_for_key(auth("admin", LitellmUserRoles.PROXY_ADMIN), "a" * 64) - assert missing.value.status_code == 403 - - -@pytest.mark.asyncio -async def test_entry_endpoints_apply_namespace_and_delete_after_disable(database: MagicMock) -> None: - database.db.litellm_memorypolicy.find_many.return_value = [policy()] - namespace = MemoryIdentity.from_auth(auth()).namespace("key") - table = database.db.litellm_memorytable - assert not await management.list_entries("demo", 2, 4, None, auth()) - assert table.find_many.call_args.kwargs["where"]["namespace"] == namespace - now = datetime(2026, 9, 12, tzinfo=timezone.utc) - table.create.return_value = SimpleNamespace( - memory_id="entry", - key=f"memory-v2:{namespace}:demo", - namespace=namespace, - value="Use port 8347", - metadata={"title": "Demo", "evidence": "User said so"}, - updated_at=now, - created_at=now, - created_by="owner", - ) - saved = await management.capture_entry( - MemoryCapture(key="demo", title="Demo", content="Use port 8347", evidence="User said so"), None, auth() - ) - assert saved.content == "Use port 8347" and saved.key == "demo" - database.db.litellm_memorypolicy.find_many.return_value = [policy(activation="disabled")] - table.delete_many.return_value = 1 - assert (await management.delete_entry("entry", None, auth())).status_code == 204 - assert table.delete_many.call_args.kwargs["where"] == {"namespace": namespace, "memory_id": "entry"} - table.delete_many.return_value = 0 - with pytest.raises(HTTPException) as missing: - await management.delete_entry("entry", None, auth()) - assert missing.value.status_code == 404 - - -@pytest.mark.parametrize("key_owner", ["owner", None]) -def test_admin_selected_key_preference_controls_actual_request_activation( - database: MagicMock, key_owner: str | None +@pytest.mark.parametrize("selected", [(), ("missing",)]) +async def test_invalid_selected_users_are_rejected_without_changing_config( + database: MagicMock, selected: tuple[str, ...] ) -> None: - database.db.litellm_verificationtoken.find_unique.return_value["user_id"] = key_owner - database.db.litellm_memorypolicy.find_many.return_value = [policy(activation="opt_in")] - table = database.db.litellm_memorypreference - app = FastAPI() - app.include_router(management.router) - app.dependency_overrides[user_api_key_auth] = lambda: auth("admin", LitellmUserRoles.PROXY_ADMIN) - params = {"key_id": "a" * 64} - subject = memory_digest("user", key_owner) if key_owner else memory_digest("key", "a" * 64) - with TestClient(app) as client: - assert client.get("/v2/memory/status", params=params).json()["active"] is False - assert client.put("/v2/memory/preference", params=params, json={"enabled": True}).status_code == 200 - assert table.upsert.call_args.kwargs["where"] == {"subject": subject} - table.find_unique.return_value = SimpleNamespace(enabled=True) - assert client.get("/v2/memory/status", params=params).json()["active"] is True - assert client.get("/v2/memory/preference", params=params).json() == {"enabled": True} - assert client.put("/v2/memory/preference", params=params, json={"enabled": False}).status_code == 200 - assert table.delete_many.call_args.kwargs["where"] == {"subject": subject} - table.find_unique.return_value = None - assert client.get("/v2/memory/status", params=params).json()["active"] is False + with pytest.raises(HTTPException) as exc: + await management.set_settings( + MemorySettings(enabled=True, everyone=False, user_ids=selected), auth(role=LitellmUserRoles.PROXY_ADMIN) + ) + assert exc.value.status_code == 422 + database.db.litellm_config.upsert.assert_not_awaited() -@pytest.mark.parametrize("role", [LitellmUserRoles.INTERNAL_USER, LitellmUserRoles.INTERNAL_USER_VIEW_ONLY]) -def test_preference_selection_cannot_bypass_key_ownership_or_readonly( +@pytest.mark.asyncio +async def test_everyone_clears_stale_selection_and_includes_service_keys(database: MagicMock) -> None: + settings = await management.set_settings( + MemorySettings(enabled=True, user_ids=("deleted",)), auth(role=LitellmUserRoles.PROXY_ADMIN) + ) + assert settings.user_ids == () + configure(database) + assert (await resolve_memory_access(database, MemoryIdentity.from_auth(auth(None)))).active + database.db.litellm_usertable.find_many.assert_not_awaited() + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "member,member_role,permission,global_role,expected", + [ + ("owner", "user", None, LitellmUserRoles.INTERNAL_USER, ()), + ("owner", "user", "/spend/logs", LitellmUserRoles.INTERNAL_USER, ()), + ("owner", "user", "/v2/memory/entries", LitellmUserRoles.INTERNAL_USER, ("team",)), + ("owner", "admin", None, LitellmUserRoles.INTERNAL_USER, ("team",)), + ("other", "admin", "/v2/memory/entries", LitellmUserRoles.INTERNAL_USER, ()), + ("owner", "user", None, LitellmUserRoles.ORG_ADMIN, ()), + ], +) +async def test_team_memory_access_reuses_membership_but_not_log_permission( + database: MagicMock, + member: str, + member_role: str, + permission: str | None, + global_role: LitellmUserRoles, + expected: tuple[str, ...], +) -> None: + database.db.litellm_teamtable.find_many.return_value = [ + team(member=member, role=member_role, permissions=(permission,) if permission else ()) + ] + access = await resolve_memory_access(database, MemoryIdentity.from_auth(auth(role=global_role))) + assert access.team_ids == expected + assert not access.admin_view + + +@pytest.mark.parametrize( + "permission,log_read,memory_read", [("/spend/logs", True, False), ("/v2/memory/entries", False, True)] +) +def test_common_helper_keeps_resource_permissions_independent( + permission: str, log_read: bool, memory_read: bool +) -> None: + context = team(permissions=(permission,)) + assert can_read_team_records(auth(), context, KeyManagementRoutes.SPEND_LOGS) is log_read + assert can_read_team_records(auth(), context, KeyManagementRoutes.MEMORY_READ) is memory_read + assert not can_read_team_records(auth("outsider"), context, KeyManagementRoutes.MEMORY_READ) + + +@pytest.mark.asyncio +async def test_key_context_restricts_org_while_dashboard_combines_permitted_teams(database: MagicMock) -> None: + database.db.litellm_usertable.find_unique.return_value = {"user_id": "owner", "teams": ["team", "elsewhere"]} + database.db.litellm_teamtable.find_many.return_value = [ + team(role="admin"), + team(role="admin", team_id="elsewhere", organization_id="other-org"), + ] + key = await resolve_memory_access(database, MemoryIdentity.from_auth(auth())) + dashboard = await resolve_memory_access( + database, MemoryIdentity.from_auth(auth().model_copy(update={"is_session_token": True, "team_id": UI_TEAM_ID})) + ) + assert key.team_ids == ("team",) + assert dashboard.team_ids == ("team", "elsewhere") + assert key.visible_rows()["organization_id"] == "org" + assert "organization_id" not in dashboard.visible_rows() + assert dashboard.identity.key_id is None and dashboard.identity.team_id is None + + +@pytest.mark.asyncio +@pytest.mark.parametrize("role", [LitellmUserRoles.PROXY_ADMIN, LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY]) +async def test_authenticated_proxy_roles_inspect_all_without_reinterpreting_auth( database: MagicMock, role: LitellmUserRoles ) -> None: - app = FastAPI() - app.include_router(management.router) - app.dependency_overrides[user_api_key_auth] = lambda: auth(role=role).model_copy( - update={"token": "session", "team_id": UI_TEAM_ID} + database.db.litellm_usertable.find_unique.return_value = None + access = await resolve_memory_access(database, MemoryIdentity.from_auth(auth(role=role))) + assert access.admin_view + assert access.visible_rows() == {"namespace": {"not": None}} + assert access.identity.read_only == (role == LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY) + + +@pytest.mark.asyncio +async def test_team_read_does_not_enable_human_edit_or_delete(database: MagicMock) -> None: + configure(database) + database.db.litellm_teamtable.find_many.return_value = [team(role="admin")] + store = await management.memory_store(auth()) + assert not store.entry(row(user_id="teammate")).can_edit + assert store.entry(row()).can_edit + assert "team_id" not in str(store.access.visible_rows(write=True)) + with pytest.raises(HTTPException) as exc: + await store.update("teammate-record", _CAPTURE) + assert exc.value.status_code == 404 + assert not await store.delete("teammate-record") + database.db.litellm_memorytable.create.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_service_key_write_ownership_never_matches_all_unowned_rows(database: MagicMock) -> None: + identity = MemoryIdentity.from_auth(auth(None)) + store = MemoryStore(database, await resolve_memory_access(database, identity)) + assert store.entry(row(user_id=None)).can_edit + assert not store.entry(row(user_id=None, owner_key_id="b" * 64)).can_edit + assert not store.entry(row(user_id="someone")).can_edit + assert store.access.visible_rows(write=True)["OR"] == [{"owner_key_id": "a" * 64, "user_id": None}] + + +@pytest.mark.asyncio +async def test_revocation_blocks_existing_store_and_private_continuation(database: MagicMock) -> None: + configure(database) + database.db.litellm_teamtable.find_many.return_value = [team(permissions=("/v2/memory/entries",))] + original = await management.memory_store(auth()) + patch = MemoryContinuation( + replaces=1, + replacement=({"role": "assistant", "content": "Team secret"},), + permission_revision=original.access.permission_revision, ) - with TestClient(app) as client: - own = client.put("/v2/memory/preference", params={"key_id": "a" * 64}, json={"enabled": True}) - assert own.status_code == (403 if role == LitellmUserRoles.INTERNAL_USER_VIEW_ONLY else 200) - database.db.litellm_memorypreference.upsert.reset_mock() - database.db.litellm_verificationtoken.find_unique.return_value["user_id"] = "another-owner" - assert client.get("/v2/memory/preference", params={"key_id": "a" * 64}).status_code == 403 - assert ( - client.put("/v2/memory/preference", params={"key_id": "a" * 64}, json={"enabled": True}).status_code == 403 - ) - assert ( - client.put("/v2/memory/preference", params={"key_id": "invalid"}, json={"enabled": True}).status_code == 422 - ) - database.db.litellm_memorypreference.upsert.assert_not_awaited() + database.db.litellm_teamtable.find_many.return_value = [team()] + with pytest.raises(HTTPException) as exc: + await original.read("team-record") + assert exc.value.status_code == 403 + fresh = await management.memory_store(auth()) + with pytest.raises(HTTPException, match="Memory permissions changed"): + MemoryContinuations(fresh, "acompletion").validate_patch(patch.model_dump()) + database.db.litellm_memorytable.find_first.assert_not_awaited() @pytest.mark.asyncio -async def test_dashboard_search_keeps_recency_before_pagination_while_agent_search_ranks_relevance( - database: MagicMock, -) -> None: - from litellm.proxy.memory.store import MemoryStore - from litellm.types.memory_v2 import MemorySearch - - database.db.litellm_memorypolicy.find_many.return_value = [policy()] - access = await management.access_for_key(auth(), None) - now = datetime(2026, 9, 12, tzinfo=timezone.utc) - database.db.litellm_memorytable.find_many.return_value = [ - SimpleNamespace( - memory_id=memory_id, - key=f"memory-v2:{access.namespace}:{memory_id}", - namespace=access.namespace, - value=content, - metadata={"title": title, "evidence": "A user instruction"}, - updated_at=now.replace(day=day), - created_at=now.replace(day=day), - created_by="owner", - ) - for memory_id, title, content, day in [ - ("newer", "Demo configuration", "The demo uses port 8123", 12), - ("unrelated", "Theme", "Use dark mode", 11), - ("older", "port", "port", 10), - ] - ] - assert [entry.memory_id for entry in await management.list_entries("port", 1, 0, None, auth())] == ["newer"] - assert [entry.memory_id for entry in await management.list_entries("port", 1, 1, None, auth())] == ["older"] - assert [entry.memory_id for entry in await MemoryStore(database, access).search(MemorySearch(query="port"))] == [ - "older", - "newer", - ] +async def test_disable_stops_tools_but_keeps_dashboard_and_ownership_checks(database: MagicMock) -> None: + database.db.litellm_memorytable.find_first.return_value = row() + store = await management.memory_store(auth()) + assert (await store.read("entry", require_active=False)).content == _CAPTURE.content + with pytest.raises(HTTPException) as exc: + await store.read("entry") + assert exc.value.status_code == 403 + with pytest.raises(HTTPException): + await store.capture(_CAPTURE) + viewer = await management.memory_store(auth(role=LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY)) + with pytest.raises(HTTPException): + await viewer.delete("entry") + database.db.litellm_memorytable.delete_many.assert_not_awaited() @pytest.mark.asyncio -async def test_dashboard_cursor_survives_deletion_of_earlier_entries(database: MagicMock) -> None: - database.db.litellm_memorypolicy.find_many.return_value = [policy()] - now = datetime(2026, 9, 12, tzinfo=timezone.utc) - rows = [ - SimpleNamespace( - memory_id=identifier, - key=f"memory-v2:namespace:{identifier}", - value=f"Memory {identifier}", - metadata={"title": identifier}, - updated_at=now, - created_at=now, - created_by="owner", - ) - for identifier in ("a", "b", "c") +async def test_names_and_team_attribution_are_loaded_for_returned_page(database: MagicMock) -> None: + database.db.litellm_memorytable.find_many.return_value = [row()] + database.db.litellm_teamtable.find_many.return_value = [team()] + database.db.litellm_usertable.find_many.return_value = [ + SimpleNamespace(user_id="owner", user_alias="Alex", user_email="alex@example.test") ] - database.db.litellm_memorytable.find_many.return_value = rows - first = await management.list_entries("", 1, 0, None, auth()) - assert first[0].memory_id == "a" - database.db.litellm_memorytable.find_many.return_value = rows[1:] - second = await management.list_entries("", 1, 0, None, auth(), first[0].updated_at, first[0].memory_id) - assert second[0].memory_id == "b" - third = await management.list_entries("", 1, 0, None, auth(), second[0].updated_at, second[0].memory_id) - assert third[0].memory_id == "c" + entries = await management.list_entries( + query="", + limit=20, + offset=0, + before_updated_at=None, + before_memory_id=None, + team_id=None, + user_id=None, + auth=auth(), + ) + assert entries[0].actor_name == "Alex" and entries[0].team_name == "Engineering" + assert entries[0].actor == "owner" and entries[0].user_id == "owner" + + +@pytest.mark.asyncio +async def test_capture_records_authenticated_contributor_and_tenant(database: MagicMock) -> None: + configure(database) + await management.capture_entry(_CAPTURE, auth("admin", LitellmUserRoles.PROXY_ADMIN)) + data = database.db.litellm_memorytable.create.call_args.kwargs["data"] + assert data["user_id"] == "admin" and data["created_by"] == "admin" + assert data["organization_id"] == "org" and data["team_id"] == "team" + + +@pytest.mark.asyncio +async def test_search_finds_an_old_record_beyond_the_first_thousand(database: MagicMock) -> None: + configure(database) + newer = [ + row( + memory_id=f"new-{index:04}", + value="Unrelated", + metadata="{}", + updated_at=_NOW + timedelta(seconds=2000 - index), + ) + for index in range(1152) + ] + old = row(memory_id="old", value="The rare quokka uses port 9187") + table = database.db.litellm_memorytable + table.find_many.side_effect = [newer[start : start + 128] for start in range(0, len(newer), 128)] + [[old]] + store = await management.memory_store(auth()) + matches, total = await store.recall(MemoryRecallRequest(query="quokka", limit=1)) + assert total == 1 and matches[0][0].memory_id == "old" + assert table.find_many.await_count == 10 + + +@pytest.mark.asyncio +async def test_catalog_and_search_recheck_permissions_after_fetch(database: MagicMock) -> None: + configure(database) + table = database.db.litellm_memorytable + store = await management.memory_store(auth()) + + async def revoke(**kwargs: object) -> list[LiteLLM_MemoryTable]: + configure(database, enabled=False) + return [row()] + + table.find_many.side_effect = revoke + for operation in (lambda: store.catalog(MemoryCatalogRequest()), lambda: store.search(MemorySearch())): + configure(database) + with pytest.raises(HTTPException) as exc: + await operation() + assert exc.value.status_code == 403 @pytest.mark.parametrize( "params", [ - {"before_memory_id": "a"}, + {"query": ",".join(f"term{i}" for i in range(17))}, + {"before_memory_id": "only-half"}, {"before_updated_at": "2026-09-12T00:00:00Z"}, - {"before_memory_id": "a", "before_updated_at": "2026-09-12T00:00:00"}, + {"before_updated_at": "2026-09-12T00:00:00", "before_memory_id": "entry"}, + {"offset": "10001"}, ], ) -def test_dashboard_cursor_rejects_partial_or_naive_dates(database: MagicMock, params: dict[str, str]) -> None: +def test_invalid_queries_are_rejected_before_database_work(database: MagicMock, params: dict[str, str]) -> None: app = FastAPI() app.include_router(management.router) app.dependency_overrides[user_api_key_auth] = auth with TestClient(app) as client: - assert client.get("/v2/memory/entries", params=params).status_code == 422 + response = client.get("/v2/memory/entries", params=params) + assert response.status_code == 422 database.db.litellm_memorytable.find_many.assert_not_awaited() - - -@pytest.mark.asyncio -@pytest.mark.parametrize("activation", ["automatic", "opt_in"]) -async def test_status_does_not_offer_an_unresolvable_namespace(database: MagicMock, activation: str) -> None: - database.db.litellm_verificationtoken.find_unique.return_value["user_id"] = None - database.db.litellm_memorypolicy.find_many.return_value = [policy(activation=activation, scope="user")] - status = await management.get_status("a" * 64, auth()) - assert not status.active and status.scope is None diff --git a/tests/test_litellm/proxy/memory/test_memory_v2_protocols.py b/tests/test_litellm/proxy/memory/test_memory_v2_protocols.py index b63b4ec1ffa..58e1b6af8c0 100644 --- a/tests/test_litellm/proxy/memory/test_memory_v2_protocols.py +++ b/tests/test_litellm/proxy/memory/test_memory_v2_protocols.py @@ -138,8 +138,8 @@ def test_private_namespaces_follow_authenticated_key_and_organization() -> None: elsewhere: Final = MemoryIdentity.from_auth( UserAPIKeyAuth(token="a" * 64, user_id="owner", team_id="team", org_id="other") ) - assert owner.namespace("key") != sibling.namespace("key") - assert owner.namespace("key") != elsewhere.namespace("key") - assert owner.namespace("user") == sibling.namespace("user") - assert owner.namespace("user") != elsewhere.namespace("user") - assert owner.namespace("team") == sibling.namespace("team") + assert owner.namespace == sibling.namespace + assert owner.namespace != elsewhere.namespace + service = MemoryIdentity.from_auth(UserAPIKeyAuth(token="a" * 64, team_id="team", org_id="org")) + other_service = MemoryIdentity.from_auth(UserAPIKeyAuth(token="b" * 64, team_id="team", org_id="org")) + assert service.namespace != other_service.namespace diff --git a/ui/litellm-dashboard/src/app/(dashboard)/memory/_components/AutomaticMemoryEntries.tsx b/ui/litellm-dashboard/src/app/(dashboard)/memory/_components/AutomaticMemoryEntries.tsx index fc363a83a3d..635df1b8e23 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/memory/_components/AutomaticMemoryEntries.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/memory/_components/AutomaticMemoryEntries.tsx @@ -1,123 +1,57 @@ "use client"; +import { useDebouncedValue } from "@tanstack/react-pacer/debouncer"; import { useInfiniteQuery, useMutation, useQuery, useQueryClient } from "@tanstack/react-query"; import { Brain } from "lucide-react"; import { type FormEvent, useState } from "react"; - -import { useKeys } from "@/app/(dashboard)/hooks/keys/useKeys"; import DeleteResourceModal from "@/components/common_components/DeleteResourceModal"; import { Button } from "@/components/ui/button"; import { Dialog, DialogContent, DialogHeader, DialogTitle } from "@/components/ui/dialog"; import { Label } from "@/components/ui/label"; -import { Skeleton } from "@/components/ui/skeleton"; import { Textarea } from "@/components/ui/textarea"; import { fetchClient } from "@/lib/http/api"; import type { components } from "@/lib/http/schema"; import { toast } from "@/lib/toast"; - -import { MemoryPreference } from "./MemorySettings"; -import { MemoryKeyPicker } from "./MemoryTargetPicker"; +import { DEBOUNCE_WAIT_MS } from "@/utils/debounceConstants"; +import { MemoryTeamPicker, MemoryUserPicker } from "./MemoryTargetPicker"; import { MemoryEntriesTable } from "./MemoryEntriesTable"; type Entry = components["schemas"]["MemoryEntry"]; type Capture = components["schemas"]["MemoryCapture"]; type Cursor = { before_updated_at: string; before_memory_id: string } | null; -type Status = components["schemas"]["MemoryStatus"]; -function memoryDescription(status?: Status) { - if (!status) return "What your assistants remember across conversations."; - if (status.activation === "disabled" || !status.scope) - return "Memory is off. Your administrator can make it available."; - if (status.activation === "automatic") { - if (status.active) return "Memory is on. Your administrator manages this setting."; - return "Memory is off. Your administrator can make it available."; - } - if (status.active) return "Your assistants can save and recall memories. You can turn this off at any time."; - return "Your assistants won't save or recall memories. Turn it on whenever you're ready."; -} - -function emptyDescription(query: string, active: boolean) { - if (query) return "Try a different search."; - if (active) return "Use your assistant as usual. What it remembers will appear here."; - return "Turn on memory, then use your assistant as usual. Your memories will appear here."; -} - -function loadMoreLabel(fetching: boolean, failed: boolean) { - if (fetching) return "Loading…"; - return failed ? "Try again" : "Load more memories"; -} - -type DashboardProps = Readonly<{ userId: string; readOnly: boolean; proxyAdmin: boolean }>; - -export function AutomaticMemoryEntries({ userId, readOnly, proxyAdmin }: DashboardProps) { - const [selection, setSelection] = useState(); - const keyOptions = { - userID: proxyAdmin ? undefined : userId, - sortBy: "created_at", - sortOrder: "desc", - includeTeamKeys: proxyAdmin, - includeCreatedByKeys: proxyAdmin, - }; - const keys = useKeys(1, 1, keyOptions); - const keyId = selection ?? keys.data?.keys[0]?.token ?? ""; - return ( - -
- - - {keys.error && ( -

- {keys.error.message} -

- )} - {keys.isSuccess && keys.data.total_count === 0 && ( -

Create a virtual key to start using memory.

- )} -
-
- ); -} - -function MemoryDashboard({ +export function AutomaticMemoryEntries({ userId, - keyId, readOnly, - children, -}: Readonly<{ - userId: string; - keyId: string; - readOnly: boolean; - children: React.ReactNode; -}>) { + proxyAdmin, +}: Readonly<{ userId: string; readOnly: boolean; proxyAdmin: boolean }>) { const cache = useQueryClient(); - const [query, setQuery] = useState(""); + const [search, setSearch] = useState(""); + const [query] = useDebouncedValue(search, { wait: DEBOUNCE_WAIT_MS }); + const [teamId, setTeamId] = useState(""); + const [filterUserId, setFilterUserId] = useState(""); const [editing, setEditing] = useState(null); const [deleting, setDeleting] = useState(null); - const statusOptions = { - queryKey: ["memoryStatus", userId, keyId], - enabled: !!keyId, - refetchOnMount: true, - queryFn: async ({ signal }: { signal: AbortSignal }) => - (await fetchClient.GET("/v2/memory/status", { params: { query: { key_id: keyId } }, signal })).data, - }; - const status = useQuery(statusOptions); - const entriesOptions = { - queryKey: ["memoryEntries", userId, keyId, query], - enabled: !!keyId && !!status.data?.scope, - refetchOnMount: true, + const status = useQuery({ + queryKey: ["memoryStatus", userId, readOnly, proxyAdmin], + queryFn: async ({ signal }) => (await fetchClient.GET("/v2/memory/status", { signal })).data, + }); + const entries = useInfiniteQuery({ + queryKey: ["memoryEntries", userId, query, teamId, filterUserId, status.data?.team_ids, status.data?.admin_view], + enabled: status.isSuccess, initialPageParam: null as Cursor, - queryFn: async ({ signal, pageParam }: { signal: AbortSignal; pageParam: Cursor }) => + queryFn: async ({ signal, pageParam }) => ( await fetchClient.GET("/v2/memory/entries", { - params: { query: { key_id: keyId, query, limit: 20, ...pageParam } }, + params: { + query: { + query, + limit: 20, + team_id: teamId || undefined, + user_id: filterUserId || undefined, + ...pageParam, + }, + }, signal, }) ).data ?? [], @@ -127,11 +61,10 @@ function MemoryDashboard({ ? { before_updated_at: last.updated_at, before_memory_id: last.memory_id } : undefined; }, - }; - const entries = useInfiniteQuery(entriesOptions); + }); const save = useMutation({ - mutationFn: async (body: Capture) => - fetchClient.POST("/v2/memory/entries", { params: { query: { key_id: keyId } }, body }), + mutationFn: ({ memory_id, body }: { memory_id: string; body: Capture }) => + fetchClient.PUT("/v2/memory/entries/{memory_id}", { params: { path: { memory_id } }, body }), onSuccess: () => { setEditing(null); toast.success("Memory updated"); @@ -140,10 +73,8 @@ function MemoryDashboard({ onError: (error: Error) => toast.error(error.message), }); const remove = useMutation({ - mutationFn: async (memory_id: string) => - fetchClient.DELETE("/v2/memory/entries/{memory_id}", { - params: { path: { memory_id }, query: { key_id: keyId } }, - }), + mutationFn: (memory_id: string) => + fetchClient.DELETE("/v2/memory/entries/{memory_id}", { params: { path: { memory_id } } }), onSuccess: () => { setDeleting(null); toast.success("Memory deleted"); @@ -158,80 +89,114 @@ function MemoryDashboard({ const saveCorrection = (event: FormEvent) => { event.preventDefault(); if (!editing) return; - const body: Capture = { - key: editing.key, - title: editing.title, - content: editing.content, - evidence: editing.evidence, - when_to_use: editing.when_to_use, - scope: editing.scope, - kind: editing.kind, - certainty: editing.certainty, - source: editing.source, - expected_revision: editing.updated_at, - }; - save.mutate(body); + save.mutate({ + memory_id: editing.memory_id, + body: { + key: editing.key, + title: editing.title, + content: editing.content, + evidence: editing.evidence, + when_to_use: editing.when_to_use, + scope: editing.scope, + kind: editing.kind, + certainty: editing.certainty, + source: editing.source, + expected_revision: editing.updated_at, + }, + }); }; - const description = memoryDescription(status.data); + const filtered = Boolean(query || teamId || filterUserId); + const accessDescription = status.data?.active + ? "Your assistant can save and search memories using your gateway permissions." + : "Automatic memory is off. Your administrator can enable it; saved memories remain available here."; + const gettingStarted = status.data?.active + ? "Use your assistant as usual. Its memories will appear here." + : "Your administrator can enable memory. Existing access permissions decide what you can see."; + const moreLabel = entries.isError ? "Try again" : "Load more memories"; return ( -
+

Memory

-

{description}

-
- {status.data && !status.error && ( - - )} - {keyId && status.isPending && } -
-
- {children} - {status.data?.user_id && (

- Memory for{" "} - {status.data.user_name ?? status.data.user_id} - {status.data.scope === "user" && ". Shared across this user's keys in this organization."} + What your assistants remember, with the people who contributed it.

+
+ {status.isSuccess && ( + + {status.data?.active ? "On · Managed by your admin" : "Off · Managed by your admin"} + )} - {status.error && ( -

- Could not load memory status: {status.error.message} + {status.error ? ( +

+ Could not load memory access: {status.error.message}

- )} - {status.data?.scope && ( -
-

Saved memories

+ ) : ( + <> +

+ {status.isPending ? "Loading memory access..." : accessDescription} +

+
+
+ + +
+ {proxyAdmin && ( +
+ + +
+ )} + {(teamId || filterUserId) && ( + + )} +
Could not load memories: {entries.error.message}

+

+ Could not load memories: {entries.error.message} +

) : (
-

{query ? "No matching memories" : "No memories yet"}

-

{emptyDescription(query, status.data.active)}

+

+ {query || teamId || filterUserId ? "No matching memories" : "No memories yet"} +

+

+ {filtered ? "Try another search or clear the filters." : gettingStarted} +

) } footer={ -
- - {memories.length} memories - {!status.data.active && memories.length > 0 && " · Saved memories stay here while memory is off"} - +
+ {memories.length} memories shown {(entries.hasNextPage || entries.isError) && ( )}
} /> {entries.isFetchNextPageError && ( -

+

Could not load more memories: {entries.error.message}

)} -
+ )} +
+ How access works +

+ You can read your own memories and any teams' memories you have permission to view. Team admins can + inspect their team's records, and proxy admins can inspect all. Team permissions also apply when your + assistant searches memory. +

+
{ @@ -293,7 +266,7 @@ function MemoryDashboard({ -
diff --git a/ui/litellm-dashboard/src/app/(dashboard)/memory/_components/MemoryEntriesTable.tsx b/ui/litellm-dashboard/src/app/(dashboard)/memory/_components/MemoryEntriesTable.tsx index 2894c52134b..9a08cbb9e44 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/memory/_components/MemoryEntriesTable.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/memory/_components/MemoryEntriesTable.tsx @@ -22,6 +22,7 @@ function contributor(entry: Entry) { function memoryDetails(entry: Entry) { const details = { "Contributed by": contributor(entry), + Team: entry.team_name ?? entry.team_id, Evidence: entry.evidence, "When to use": entry.when_to_use, Source: entry.source, @@ -100,6 +101,12 @@ export function MemoryEntriesTable({ ), }, + { + id: "team", + header: "Team", + size: 150, + cell: ({ row }) => {row.original.team_name ?? row.original.team_id ?? "No team"}, + }, { id: "details", header: () => Details, @@ -127,6 +134,7 @@ export function MemoryEntriesTable({ onQueryChange(event.target.value)} @@ -161,7 +169,7 @@ export function MemoryEntriesTable({ - {!readOnly && ( + {!readOnly && selected.can_edit && (
+ + ))} + +
+ )} +

+ Turning memory off stops automatic saving and recall. Existing memories remain available to authorized + viewers. Enabling memory can add model calls, latency, and spend. +

+ {save.error && ( +

+ {save.error.message} +

+ )} +
+ {draft && Unsaved changes} + + +
-
- - -
-
- More specific policies take precedence: virtual key, user, project, team, organization, then gateway. Changing - the sharing scope starts using that scope's memories; existing entries remain stored. -
- {!readOnly && ( - - )} - - {policies.error && ( -

- {policies.error.message} + )} +

+

Who can see memories?

+

+ Users see their own memories. Team admins can also see their team's memories, and proxy admins can see + all. To give ordinary members access to their team's memories, allow “Read team memories” in Member + Permissions.

- )} - {policies.isLoading &&

Loading memory policies...

} - {policies.data?.length === 0 && ( -

No policies set for this selection

- )} -
    - {(policies.data ?? []).map((row) => ( -
  • -
    -

    - {targetNames[row.target_type]}: {row.target_id === "*" ? "All requests" : row.target_id} -

    -

    - {activationNames[row.activation]} · {scopeNames[row.scope ?? "key"]} -

    -
    -
    - - -
    -
  • - ))} -
- {(offset > 0 || policies.data?.length === 100) && ( -
- - -
- )} + + Manage team permissions + +
+

These settings do not change Memory API (V1).

); } diff --git a/ui/litellm-dashboard/src/app/(dashboard)/memory/_components/MemoryTargetPicker.tsx b/ui/litellm-dashboard/src/app/(dashboard)/memory/_components/MemoryTargetPicker.tsx index 4a9c042a8d7..d796d7732a1 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/memory/_components/MemoryTargetPicker.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/memory/_components/MemoryTargetPicker.tsx @@ -1,30 +1,29 @@ "use client"; import { useState } from "react"; - -import { useInfiniteKeys } from "@/app/(dashboard)/hooks/keys/useKeys"; -import { useOrganizations } from "@/app/(dashboard)/hooks/organizations/useOrganizations"; -import { useProjects } from "@/app/(dashboard)/hooks/projects/useProjects"; import { useInfiniteTeams } from "@/app/(dashboard)/hooks/teams/useTeams"; import { useInfiniteUsers } from "@/app/(dashboard)/hooks/users/useUsers"; import { PaginatedSearchSelect } from "@/components/shared/PaginatedSearchSelect"; -import { SearchSelect } from "@/components/shared/SearchSelect"; -import type { components } from "@/lib/http/schema"; -type Target = components["schemas"]["MemoryPolicyInput"]["target_type"]; -type PickerProps = Readonly<{ value: string; onChange: (value: string) => void; disabled: boolean }>; +type PickerProps = Readonly<{ + inputId?: string; + value: string; + onChange: (value: string, label?: string) => void; + disabled: boolean; +}>; -function TeamPicker({ value, onChange, disabled }: PickerProps) { +export function MemoryTeamPicker({ value, onChange, disabled, inputId = "memory-team" }: PickerProps) { const [search, setSearch] = useState(""); const query = useInfiniteTeams(25, search); + const options = (query.data?.pages ?? []).flatMap((page) => + page.teams.map((team) => ({ value: team.team_id, label: team.team_alias || team.team_id })), + ); return ( onChange(v ?? "")} - options={(query.data?.pages ?? []).flatMap((page) => - page.teams.map((team) => ({ value: team.team_id, label: team.team_alias || team.team_id })), - )} + onValueChange={(v) => onChange(v ?? "", options.find((option) => option.value === v)?.label)} + options={options} onSearchChange={setSearch} onLoadMore={query.fetchNextPage} hasNextPage={query.hasNextPage} @@ -32,22 +31,23 @@ function TeamPicker({ value, onChange, disabled }: PickerProps) { isFetchingNextPage={query.isFetchingNextPage} disabled={disabled} errorText={query.error?.message} - placeholder="Search teams" + placeholder="All permitted teams" /> ); } -function UserPicker({ value, onChange, disabled }: PickerProps) { +export function MemoryUserPicker({ value, onChange, disabled, inputId = "memory-user" }: PickerProps) { const [search, setSearch] = useState(""); const query = useInfiniteUsers(25, search); + const options = (query.data?.pages ?? []).flatMap((page) => + page.users.map((user) => ({ value: user.user_id, label: user.user_email || user.user_id })), + ); return ( onChange(v ?? "")} - options={(query.data?.pages ?? []).flatMap((page) => - page.users.map((user) => ({ value: user.user_id, label: user.user_email || user.user_id })), - )} + onValueChange={(v) => onChange(v ?? "", options.find((option) => option.value === v)?.label)} + options={options} onSearchChange={setSearch} onLoadMore={query.fetchNextPage} hasNextPage={query.hasNextPage} @@ -59,78 +59,3 @@ function UserPicker({ value, onChange, disabled }: PickerProps) { /> ); } - -export function MemoryKeyPicker({ - value, - onChange, - disabled, - inputId = "memory-target", - userId, -}: PickerProps & Readonly<{ inputId?: string; userId?: string }>) { - const [search, setSearch] = useState(""); - const keyOptions = { search, userID: userId, includeTeamKeys: !userId, includeCreatedByKeys: !userId }; - const query = useInfiniteKeys(25, keyOptions); - return ( - onChange(v ?? "")} - options={(query.data?.pages ?? []).flatMap((page) => - page.keys.map((key) => ({ value: key.token, label: key.key_alias || key.key_name || key.token })), - )} - onSearchChange={setSearch} - onLoadMore={query.fetchNextPage} - hasNextPage={query.hasNextPage} - isLoading={query.isLoading} - isFetchingNextPage={query.isFetchingNextPage} - disabled={disabled} - errorText={query.error?.message} - placeholder="Search virtual keys" - /> - ); -} - -function OrganizationPicker({ value, onChange, disabled }: PickerProps) { - const query = useOrganizations(); - return ( - onChange(v ?? "")} - options={(query.data ?? []).map((org) => ({ - value: org.organization_id, - label: org.organization_alias || org.organization_id, - }))} - disabled={disabled} - placeholder="Select an organization" - emptyText={query.error?.message ?? "No organizations found"} - /> - ); -} - -function ProjectPicker({ value, onChange, disabled }: PickerProps) { - const query = useProjects(); - return ( - onChange(v ?? "")} - options={(query.data ?? []).map((project) => ({ - value: project.project_id, - label: project.project_alias || project.project_id, - }))} - disabled={disabled} - placeholder="Select a project" - emptyText={query.error?.message ?? "No projects found"} - /> - ); -} - -export function MemoryTargetPicker({ target, ...props }: PickerProps & Readonly<{ target: Target }>) { - if (target === "team") return ; - if (target === "user") return ; - if (target === "key") return ; - if (target === "organization") return ; - if (target === "project") return ; - return null; -} diff --git a/ui/litellm-dashboard/src/app/(dashboard)/memory/page.integration.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/memory/page.integration.test.tsx index 3480d3e46b0..849677406d9 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/memory/page.integration.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/memory/page.integration.test.tsx @@ -1,22 +1,19 @@ -import { fireEvent, screen, waitFor, within } from "@testing-library/react"; +import { fireEvent, screen, waitFor } from "@testing-library/react"; import userEvent from "@testing-library/user-event"; import { beforeEach, describe, expect, it, vi } from "vitest"; - import { renderWithProviders, testQueryClient } from "../../../../tests/test-utils"; +import type { components } from "@/lib/http/schema"; import Memory from "./page"; -import { Toaster } from "@/components/ui/sonner"; vi.unmock("@/app/(dashboard)/hooks/useAuthorized"); vi.unmock("@/lib/toast"); const fetchMock = vi.fn(); -const calls: { path: string; method: string; body: unknown; keyId: string | null }[] = []; -let enabled = false; -let activation = "opt_in"; -let available = true; +const calls: { path: string; method: string; body: unknown; params: string }[] = []; +let settings: components["schemas"]["MemorySettings"]; let paginated = false; -let failPreference = false; -let sameUser = false; +let failure = ""; +let canEdit = true; const entry = { memory_id: "entry-1", key: "demo", @@ -32,6 +29,9 @@ const entry = { updated_at: "2026-09-12T00:00:00Z", actor: "u1", actor_name: "Alex Rivera", + user_id: "u1", + team_id: "engineering", + team_name: "Engineering", }; const session = (user_role: string) => { const payload = { key: "sk-test", user_id: "u1", user_role, exp: Math.floor(Date.now() / 1000) + 3600 }; @@ -42,85 +42,148 @@ beforeEach(async () => { await testQueryClient.cancelQueries(); testQueryClient.clear(); calls.length = 0; - enabled = false; - activation = "opt_in"; - available = true; + settings = { enabled: false, everyone: true, user_ids: [] }; paginated = false; - failPreference = false; - sameUser = false; + failure = ""; + canEdit = true; vi.clearAllMocks(); fetchMock.mockImplementation(async (input, init) => { const request = input instanceof Request ? input : new Request(new URL(String(input), window.location.origin), init); - const path = new URL(request.url).pathname; - const text = request.method === "GET" ? "" : await request.text(); const url = new URL(request.url); - const keyId = url.searchParams.get("key_id"); - if (url.pathname === "/key/list" && url.searchParams.get("user_id")) { - expect(url.searchParams.get("include_team_keys")).toBe("false"); - expect(url.searchParams.get("include_created_by_keys")).toBe("false"); - } - const call = { path, method: request.method, body: text ? JSON.parse(text) : undefined, keyId }; + const path = url.pathname; + const text = request.method === "GET" ? "" : await request.text(); + const call = { path, method: request.method, body: text ? JSON.parse(text) : undefined, params: url.search }; calls.push(call); - if (path === "/v2/memory/preference" && request.method === "PUT" && failPreference) { - return new Response(JSON.stringify({ detail: "Preference unavailable" }), { + if (failure === path && (path !== "/v2/memory/settings" || request.method === "PUT")) { + return new Response(JSON.stringify({ detail: "Memory service unavailable" }), { status: 503, headers: { "Content-Type": "application/json" }, }); } const response = () => { - if (path === "/v2/memory/preference") { - if (request.method === "PUT") enabled = JSON.parse(text).enabled; - return { enabled }; + if (path === "/v2/memory/settings") { + if (request.method === "PUT") settings = JSON.parse(text); + return settings; } - if (path === "/v2/memory/policies") return []; - if (path === "/v1/memory") return { memories: [], total: 0 }; if (path === "/v2/memory/status") return { - active: available && (activation === "automatic" || enabled), - opted_in: enabled, - activation, - scope: available ? "user" : null, - user_id: keyId === "b".repeat(64) && !sameUser ? "u2" : "u1", - user_name: keyId === "b".repeat(64) && !sameUser ? "Jamie Davis" : "Alex Rivera", + active: settings.enabled && (settings.everyone || settings.user_ids?.includes("u1")), + enabled: settings.enabled, + user_id: "u1", + user_name: "Alex Rivera", + team_ids: ["engineering"], + admin_view: false, }; if (path === "/v2/memory/entries") { - if (request.method === "POST") return entry; - if (keyId === "b".repeat(64) && !sameUser) - return [{ ...entry, memory_id: "other", title: "Other key memory", content: "Another project" }]; if (paginated) { - const offset = new URL(request.url).searchParams.get("before_memory_id") === "entry-19" ? 20 : 0; + const offset = url.searchParams.get("before_memory_id") === "entry-19" ? 20 : 0; return Array.from({ length: offset ? 1 : 20 }, (_, i) => ({ ...entry, + can_edit: canEdit, memory_id: `entry-${offset + i}`, title: `Memory ${offset + i}`, })); } - return [entry]; + return [{ ...entry, can_edit: canEdit }]; } - if (path.includes("/key/list")) - return { - keys: [ - { token: "a".repeat(64), key_alias: "QA key" }, - { token: "b".repeat(64), key_alias: "Other key" }, - ], - total_count: 2, - current_page: 1, - total_pages: 1, - }; + if (path === "/v2/memory/entries/entry-1") return { ...entry, can_edit: canEdit }; + if (path === "/v1/memory") return { memories: [], total: 0 }; + if (path === "/v2/team/list") + return { teams: [{ team_id: "engineering", team_alias: "Engineering" }], page: 1, total_pages: 1, total: 1 }; + if (path === "/user/list") + return { users: [{ user_id: "u1", user_email: "alex@example.test" }], page: 1, total_pages: 1, total: 1 }; if (path.includes("/team/list") || path.includes("/organization")) return []; return {}; }; - const data = response(); - return new Response(JSON.stringify(data), { status: 200, headers: { "Content-Type": "application/json" } }); + return new Response(JSON.stringify(response()), { status: 200, headers: { "Content-Type": "application/json" } }); }); vi.stubGlobal("fetch", fetchMock); }); describe("Memory dashboard", () => { - it("preserves observation attribution when a person corrects saved content", async () => { + it("shows recent content and attribution while automatic memory is off", async () => { + session("internal_user"); + renderWithProviders(); + expect(await screen.findByText("Off · Managed by your admin")).toBeVisible(); + expect(await screen.findByText("Use port 8123")).toBeVisible(); + expect(screen.getByText("Alex Rivera")).toBeVisible(); + expect(screen.getByText("Engineering")).toBeVisible(); + expect(screen.queryByRole("switch")).not.toBeInTheDocument(); + expect(screen.queryByRole("tab", { name: "Administration" })).not.toBeInTheDocument(); + expect(calls.some(({ path }) => path === "/v2/memory/settings")).toBe(false); + }); + + it("lets a proxy admin enable everyone through one configuration", async () => { + session("proxy_admin"); + const user = userEvent.setup(); + renderWithProviders(); + await user.click(screen.getByRole("tab", { name: "Administration" })); + const toggle = await screen.findByRole("switch", { name: "Gateway memory" }); + expect(toggle).not.toBeChecked(); + await user.click(toggle); + await user.click(screen.getByRole("button", { name: "Save changes" })); + await waitFor(() => expect(settings.enabled).toBe(true)); + expect(settings.everyone).toBe(true); + await waitFor(() => expect(screen.queryByText("Unsaved changes")).not.toBeInTheDocument()); + expect(screen.getByRole("link", { name: "Manage team permissions" })).toHaveAttribute("href", "/teams"); + await user.click(screen.getByRole("tab", { name: "Memories" })); + expect(await screen.findByText("On · Managed by your admin")).toBeVisible(); + expect(calls.some(({ path }) => path.includes("policies") || path.includes("preference"))).toBe(false); + }); + + it("enrolls selected users without a key or sharing selector", async () => { + session("proxy_admin"); + const user = userEvent.setup(); + renderWithProviders(); + await user.click(screen.getByRole("tab", { name: "Administration" })); + await user.click(await screen.findByRole("switch", { name: "Gateway memory" })); + await user.selectOptions(screen.getByLabelText("Enable for"), "selected"); + await user.click(screen.getByLabelText("Add a user")); + await user.click(await screen.findByRole("option", { name: "alex@example.test" })); + expect(screen.getByRole("button", { name: "Remove alex@example.test" })).toBeVisible(); + await user.click(screen.getByRole("button", { name: "Save changes" })); + await waitFor(() => expect(settings).toEqual({ enabled: true, everyone: false, user_ids: ["u1"] })); + }); + + it("keeps an unsuccessful activation unsaved and shows the error", async () => { + session("proxy_admin"); + failure = "/v2/memory/settings"; + const user = userEvent.setup(); + renderWithProviders(); + await user.click(screen.getByRole("tab", { name: "Administration" })); + await user.click(await screen.findByRole("switch", { name: "Gateway memory" })); + await user.click(screen.getByRole("button", { name: "Save changes" })); + expect(await screen.findByRole("alert")).toHaveTextContent("Memory service unavailable"); + expect(settings.enabled).toBe(false); + expect(screen.getByText("Unsaved changes")).toBeVisible(); + }); + + it("lets viewers inspect but disables administration changes", async () => { + session("proxy_admin_viewer"); + const user = userEvent.setup(); + renderWithProviders(); + await user.click(await screen.findByRole("button", { name: "Details for Demo port" })); + expect(screen.queryByRole("button", { name: "Edit memory" })).not.toBeInTheDocument(); + await user.click(screen.getByRole("button", { name: "Close" })); + await user.click(screen.getByRole("tab", { name: "Administration" })); + expect(await screen.findByRole("switch", { name: "Gateway memory" })).toHaveAttribute("aria-disabled", "true"); + expect(screen.getByRole("button", { name: "Save changes" })).toBeDisabled(); + }); + + it("does not show edit controls for a readable teammate record", async () => { + session("internal_user"); + canEdit = false; + const user = userEvent.setup(); + renderWithProviders(); + await user.click(await screen.findByRole("button", { name: "Details for Demo port" })); + expect(screen.queryByRole("button", { name: "Edit memory" })).not.toBeInTheDocument(); + expect(screen.queryByRole("button", { name: "Delete memory" })).not.toBeInTheDocument(); + expect(screen.getAllByText("Use port 8123")[0]).toBeVisible(); + }); + + it("preserves attribution and revision when an owner corrects content while off", async () => { session("internal_user"); - enabled = true; const user = userEvent.setup(); renderWithProviders(); await user.click(await screen.findByRole("button", { name: "Details for Demo port" })); @@ -128,9 +191,9 @@ describe("Memory dashboard", () => { fireEvent.change(screen.getByLabelText("Correct this memory"), { target: { value: "Use port 8124" } }); await user.click(screen.getByRole("button", { name: "Save correction" })); const expectedCall = { - path: "/v2/memory/entries", - keyId: "a".repeat(64), - method: "POST", + path: "/v2/memory/entries/entry-1", + method: "PUT", + params: "", body: { key: entry.key, title: entry.title, @@ -147,155 +210,47 @@ describe("Memory dashboard", () => { await waitFor(() => expect(calls).toContainEqual(expectedCall)); }); - it("lets an administrator choose automatic activation and a sharing scope", async () => { - session("proxy_admin"); - const user = userEvent.setup(); - renderWithProviders(); - expect(screen.queryByLabelText("Activation")).not.toBeInTheDocument(); - expect(calls.filter(({ path }) => path === "/v1/memory" || path === "/v2/memory/policies")).toEqual([]); - await user.click(screen.getByRole("button", { name: "Advanced settings" })); - expect(await screen.findByLabelText("Activation")).toHaveValue("opt_in"); - await user.selectOptions(screen.getByLabelText("Activation"), "automatic"); - await user.selectOptions(screen.getByLabelText("Who shares the memories"), "team"); - await user.click(screen.getByRole("button", { name: "Save memory policy" })); - const expectedCall = { - path: "/v2/memory/policies", - keyId: null, - method: "PUT", - body: { target_type: "gateway", target_id: "*", activation: "automatic", scope: "team" }, - }; - await waitFor(() => expect(calls).toContainEqual(expectedCall)); - await waitFor(() => expect(screen.getByRole("button", { name: "Save memory policy" })).toBeEnabled()); - expect(screen.queryByText(/draft deprecation list/i)).not.toBeInTheDocument(); - }); - - it("offers members their own preference without fetching administrator memory rows", async () => { + it("filters permitted memories by team and can clear the filter", async () => { session("internal_user"); const user = userEvent.setup(); renderWithProviders(); - const preference = await screen.findByRole("switch", { name: "Memory" }); - await waitFor(() => expect(preference).toBeEnabled()); - expect(preference).not.toBeChecked(); - expect(await screen.findByText("Use port 8123")).toBeVisible(); - expect(screen.queryByText("Fixture recommendation")).not.toBeInTheDocument(); - await user.click(preference); - const expectedCall = { - path: "/v2/memory/preference", - method: "PUT", - body: { enabled: true }, - keyId: "a".repeat(64), - }; - await waitFor(() => expect(calls).toContainEqual(expectedCall)); - await waitFor(() => expect(preference).toBeChecked()); - await waitFor(() => expect(preference).not.toHaveAttribute("aria-disabled", "true")); - await user.click(preference); - await waitFor(() => expect(preference).not.toBeChecked()); - expect(screen.getByText("Use port 8123")).toBeVisible(); - expect(calls.filter(({ path }) => path === "/v1/memory" || path === "/v2/memory/policies")).toEqual([]); - expect(screen.queryByRole("button", { name: "Save memory policy" })).not.toBeInTheDocument(); + await screen.findByText("Off · Managed by your admin"); + await user.click(screen.getByLabelText("Team")); + await user.click(await screen.findByRole("option", { name: "Engineering" })); + await waitFor(() => expect(calls.some(({ params }) => params.includes("team_id=engineering"))).toBe(true)); + await user.click(screen.getByRole("button", { name: "Clear filters" })); + expect(screen.queryByRole("button", { name: "Clear filters" })).not.toBeInTheDocument(); + expect(calls.some(({ params }) => params.includes("key_id"))).toBe(false); }); - it("keeps the V1 management view available without loading it into automatic memory", async () => { - session("proxy_admin"); - const user = userEvent.setup(); - renderWithProviders(); - expect(await screen.findByRole("switch", { name: "Memory" })).not.toBeChecked(); - expect(calls.filter(({ path }) => path === "/v1/memory")).toEqual([]); - await user.click(screen.getByRole("tab", { name: "Memory API (V1)" })); - expect(await screen.findByRole("button", { name: "New memory" })).toBeVisible(); - await waitFor(() => expect(calls.some(({ path }) => path === "/v1/memory")).toBe(true)); - expect(screen.queryByRole("switch", { name: "Memory" })).not.toBeInTheDocument(); - await user.click(screen.getByRole("tab", { name: "Automatic memory" })); - expect(await screen.findByRole("switch", { name: "Memory" })).not.toBeChecked(); - }); - - it("allows viewers to inspect memory while disabling preference writes", async () => { - session("internal_user_viewer"); - renderWithProviders(); - expect(await screen.findByRole("switch", { name: "Memory" })).toHaveAttribute("aria-disabled", "true"); - expect(await screen.findByText("Use port 8123")).toBeVisible(); - }); - - it("appends older memories and resets the table when switching contexts", async () => { - session("proxy_admin"); + it("searches and paginates without replacing the visible first page", async () => { + session("internal_user"); paginated = true; const user = userEvent.setup(); renderWithProviders(); - const table = await screen.findByRole("table"); - await waitFor(() => - expect(within(table).getAllByRole("button", { name: /^Details for Memory \d+$/ })).toHaveLength(20), - ); - await user.click(screen.getByRole("button", { name: "Load more memories" })); - await waitFor(() => - expect(within(table).getAllByRole("button", { name: /^Details for Memory \d+$/ })).toHaveLength(21), - ); - expect(screen.getByRole("button", { name: "Details for Memory 0" })).toBeVisible(); - await user.click(screen.getByLabelText("Key context")); - await user.click(await screen.findByRole("option", { name: "Other key" })); - expect(await screen.findByText("Another project")).toBeVisible(); - expect(screen.queryByRole("button", { name: "Details for Memory 0" })).not.toBeInTheDocument(); + await user.click(await screen.findByRole("button", { name: "Load more memories" })); + expect(await screen.findByText("Memory 20")).toBeVisible(); + expect(screen.getByText("Memory 0")).toBeVisible(); + expect(screen.getByText("21 memories shown")).toBeVisible(); + await user.type(screen.getByRole("textbox", { name: "Search memories" }), "demo"); + await waitFor(() => expect(calls.some(({ params }) => params.includes("query=demo"))).toBe(true)); }); - it("shows contributors in the readable table and keeps user activation across that user's keys", async () => { + it("shows a loading failure instead of claiming an empty collection", async () => { session("internal_user"); - sameUser = true; + failure = "/v2/memory/entries"; + renderWithProviders(); + expect(await screen.findByRole("alert")).toHaveTextContent("Memory service unavailable"); + expect(screen.queryByText("No memories yet")).not.toBeInTheDocument(); + }); + + it("keeps V1 separate and loads its records only when selected", async () => { + session("proxy_admin"); const user = userEvent.setup(); renderWithProviders(); - const table = await screen.findByRole("table"); - expect(within(table).getByRole("columnheader", { name: "Contributed by" })).toBeVisible(); - const row = await within(table).findByRole("row", { name: /Demo port/ }); - expect(within(row).getByText("Use port 8123")).toBeVisible(); - expect(within(row).getByText("Alex Rivera")).toBeVisible(); - expect(screen.queryByText("Fixture recommendation")).not.toBeInTheDocument(); - expect(screen.getByText("Applies to this user's keys")).toBeVisible(); - await user.click(screen.getByRole("switch", { name: "Memory" })); - await waitFor(() => expect(screen.getByRole("switch", { name: "Memory" })).toBeChecked()); - await user.click(screen.getByLabelText("Key context")); - await user.click(await screen.findByRole("option", { name: "Other key" })); - await waitFor(() => expect(screen.getByRole("switch", { name: "Memory" })).toBeChecked()); - expect(await screen.findByRole("button", { name: "Details for Demo port" })).toBeVisible(); - await user.click(screen.getByRole("switch", { name: "Memory" })); - await waitFor(() => expect(screen.getByRole("switch", { name: "Memory" })).not.toBeChecked()); - await user.click(screen.getByLabelText("Key context")); - await user.click(await screen.findByRole("option", { name: "QA key" })); - await waitFor(() => expect(screen.getByRole("switch", { name: "Memory" })).not.toBeChecked()); - }); - - it("shows administrator-managed memory as on even when the preference is off", async () => { - session("internal_user"); - activation = "automatic"; - renderWithProviders(); - const toggle = await screen.findByRole("switch", { name: "Memory" }); - expect(toggle).toBeChecked(); - expect(toggle).toHaveAttribute("aria-disabled", "true"); - expect(screen.getByText(/Your administrator manages this setting/)).toBeVisible(); - }); - - it("shows memory as unavailable when an automatic policy cannot resolve a sharing scope", async () => { - session("internal_user"); - activation = "automatic"; - available = false; - renderWithProviders(); - const toggle = await screen.findByRole("switch", { name: "Memory" }); - expect(toggle).not.toBeChecked(); - expect(toggle).toHaveAttribute("aria-disabled", "true"); - expect(screen.getByText("Memory is off. Your administrator can make it available.")).toBeVisible(); - }); - - it("keeps the actual state off when saving a preference fails", async () => { - session("internal_user"); - failPreference = true; - const user = userEvent.setup(); - renderWithProviders( - <> - - - , - ); - const toggle = await screen.findByRole("switch", { name: "Memory" }); - await user.click(toggle); - expect(await screen.findByText("Preference unavailable")).toBeVisible(); - await waitFor(() => expect(toggle).not.toHaveAttribute("aria-disabled", "true")); - expect(toggle).not.toBeChecked(); + expect(calls.some(({ path }) => path === "/v1/memory")).toBe(false); + await user.click(screen.getByRole("tab", { name: "Memory API (V1)" })); + await waitFor(() => expect(calls.some(({ path }) => path === "/v1/memory")).toBe(true)); + expect(screen.queryByText(/draft deprecation list/i)).not.toBeInTheDocument(); }); }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/memory/page.tsx b/ui/litellm-dashboard/src/app/(dashboard)/memory/page.tsx index 720f3007858..215768aadee 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/memory/page.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/memory/page.tsx @@ -8,15 +8,11 @@ import { useTeams } from "@/app/(dashboard)/hooks/teams/useTeams"; import useIsOrgAdmin from "@/app/(dashboard)/hooks/useIsOrgAdmin"; import { isProxyAdminRole, isUserTeamAdminForAnyTeam } from "@/utils/roles"; import { useState } from "react"; -import { ChevronDown } from "lucide-react"; -import { Button } from "@/components/ui/button"; -import { Collapsible, CollapsibleContent, CollapsibleTrigger } from "@/components/ui/collapsible"; -import { MemoryPolicies } from "./_components/MemorySettings"; +import { MemoryAdministration } from "./_components/MemorySettings"; import { AutomaticMemoryEntries } from "./_components/AutomaticMemoryEntries"; import { Tabs, TabsContent, TabsList, TabsTrigger } from "@/components/ui/tabs"; export default function Memory() { - const [advanced, setAdvanced] = useState(false); const [view, setView] = useState("v2"); const { accessToken, userRole, userId, isViewOnly } = useAuthorized(); const canViewMemory = useCan("viewMemory"); @@ -31,33 +27,26 @@ export default function Memory() { } return ( - setView(value === "v1" ? "v1" : "v2")} className="gap-6 px-8 py-8"> - {proxyAdmin && ( - - Automatic memory - Memory API (V1) + setView(value)} className="gap-6 px-8 py-8"> + {(proxyAdmin || canManage) && ( + + Memories + {canManage && Administration} + {proxyAdmin && Memory API (V1)} )} {userId && ( )} - {canManage && userId && ( - - }> - - Advanced settings - - - {advanced && ( -
- -
- )} -
-
- )}
+ {canManage && userId && ( + + {view === "administration" && ( + + )} + + )} {proxyAdmin && ( {view === "v1" && } diff --git a/ui/litellm-dashboard/src/components/team/permission_definitions.tsx b/ui/litellm-dashboard/src/components/team/permission_definitions.tsx index 68ca9d8a2cb..72f62adfdf7 100644 --- a/ui/litellm-dashboard/src/components/team/permission_definitions.tsx +++ b/ui/litellm-dashboard/src/components/team/permission_definitions.tsx @@ -23,6 +23,7 @@ export const PERMISSION_DESCRIPTIONS: Record = { "/key/access_group_assignment": "Member can assign access groups to virtual keys for this team", "/team/daily/activity": "Member can view all team usage data (not just their own)", "/spend/logs": "Member can view spend logs for the entire team (not just their own)", + "/v2/memory/entries": "Read team memories: members can search and read memories contributed within this team", }; /** @@ -33,7 +34,8 @@ export const getMethodForEndpoint = (endpoint: string): string => { endpoint.includes("/info") || endpoint.includes("/list") || endpoint.includes("/activity") || - endpoint === "/spend/logs" + endpoint === "/spend/logs" || + endpoint === "/v2/memory/entries" ) { return "GET"; } diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index 8843c4e4887..590d70a6ffa 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -21280,8 +21280,10 @@ export interface paths { path?: never; cookie?: never; }; - get?: never; - put?: never; + /** Read Entry */ + get: operations["read_entry_v2_memory_entries__memory_id__get"]; + /** Update Entry */ + put: operations["update_entry_v2_memory_entries__memory_id__put"]; post?: never; /** Delete Entry */ delete: operations["delete_entry_v2_memory_entries__memory_id__delete"]; @@ -21290,52 +21292,17 @@ export interface paths { patch?: never; trace?: never; }; - "/v2/memory/policies": { + "/v2/memory/settings": { parameters: { query?: never; header?: never; path?: never; cookie?: never; }; - /** List Policies */ - get: operations["list_policies_v2_memory_policies_get"]; - /** Set Policy */ - put: operations["set_policy_v2_memory_policies_put"]; - post?: never; - delete?: never; - options?: never; - head?: never; - patch?: never; - trace?: never; - }; - "/v2/memory/policies/{policy_id}": { - parameters: { - query?: never; - header?: never; - path?: never; - cookie?: never; - }; - get?: never; - put?: never; - post?: never; - /** Delete Policy */ - delete: operations["delete_policy_v2_memory_policies__policy_id__delete"]; - options?: never; - head?: never; - patch?: never; - trace?: never; - }; - "/v2/memory/preference": { - parameters: { - query?: never; - header?: never; - path?: never; - cookie?: never; - }; - /** Get Preference */ - get: operations["get_preference_v2_memory_preference_get"]; - /** Set Preference */ - put: operations["set_preference_v2_memory_preference_put"]; + /** Get Settings */ + get: operations["get_settings_v2_memory_settings_get"]; + /** Set Settings */ + put: operations["set_settings_v2_memory_settings_put"]; post?: never; delete?: never; options?: never; @@ -28551,7 +28518,7 @@ export interface components { * @description Enum for key management routes * @enum {string} */ - KeyManagementRoutes: "/key/generate" | "/key/update" | "/key/delete" | "/key/regenerate" | "/key/service-account/generate" | "/key/{key_id}/regenerate" | "/key/block" | "/key/unblock" | "/key/bulk_update" | "/team/key/bulk_update" | "/key/{key_id}/reset_spend" | "/key/access_group_assignment" | "/key/info" | "/key/health" | "/key/list" | "/key/aliases" | "/team/daily/activity" | "/team/daily/activity/aggregated" | "/spend/logs" | "/spend/logs/v2"; + KeyManagementRoutes: "/key/generate" | "/key/update" | "/key/delete" | "/key/regenerate" | "/key/service-account/generate" | "/key/{key_id}/regenerate" | "/key/block" | "/key/unblock" | "/key/bulk_update" | "/team/key/bulk_update" | "/key/{key_id}/reset_spend" | "/key/access_group_assignment" | "/key/info" | "/key/health" | "/key/list" | "/key/aliases" | "/team/daily/activity" | "/team/daily/activity/aggregated" | "/spend/logs" | "/spend/logs/v2" | "/v2/memory/entries"; /** * KeyManagementSystem * @enum {string} @@ -32092,6 +32059,11 @@ export interface components { actor?: string | null; /** Actor Name */ actor_name?: string | null; + /** + * Can Edit + * @default false + */ + can_edit: boolean; /** * Certainty * @default observed @@ -32124,6 +32096,10 @@ export interface components { * @default */ source: string; + /** Team Id */ + team_id?: string | null; + /** Team Name */ + team_name?: string | null; /** Title */ title: string; /** @@ -32131,6 +32107,8 @@ export interface components { * Format: date-time */ updated_at: string; + /** User Id */ + user_id?: string | null; /** * When To Use * @default @@ -32144,77 +32122,43 @@ export interface components { /** Total */ total: number; }; - /** MemoryPolicy */ - MemoryPolicy: { + /** MemorySettings */ + MemorySettings: { /** - * Activation - * @enum {string} + * Enabled + * @default false */ - activation: "disabled" | "opt_in" | "automatic"; - /** Policy Id */ - policy_id: string; - /** - * Scope - * @default user - * @enum {string} - */ - scope: "key" | "user" | "team" | "project" | "organization"; - /** Target Id */ - target_id: string; - /** - * Target Type - * @enum {string} - */ - target_type: "gateway" | "organization" | "team" | "project" | "user" | "key"; - /** - * Updated At - * Format: date-time - */ - updated_at: string; - /** Updated By */ - updated_by: string; - }; - /** MemoryPolicyInput */ - MemoryPolicyInput: { - /** - * Activation - * @enum {string} - */ - activation: "disabled" | "opt_in" | "automatic"; - /** - * Scope - * @default user - * @enum {string} - */ - scope: "key" | "user" | "team" | "project" | "organization"; - /** Target Id */ - target_id: string; - /** - * Target Type - * @enum {string} - */ - target_type: "gateway" | "organization" | "team" | "project" | "user" | "key"; - }; - /** MemoryPreference */ - MemoryPreference: { - /** Enabled */ enabled: boolean; + /** + * Everyone + * @default true + */ + everyone: boolean; + /** + * User Ids + * @default [] + */ + user_ids: string[]; }; /** MemoryStatus */ MemoryStatus: { - /** - * Activation - * @enum {string} - */ - activation: "disabled" | "opt_in" | "automatic"; /** Active */ active: boolean; - /** Opted In */ - opted_in: boolean; - /** Policy Id */ - policy_id: string | null; - /** Scope */ - scope: ("key" | "user" | "team" | "project" | "organization") | null; + /** + * Admin View + * @default false + */ + admin_view: boolean; + /** + * Enabled + * @default false + */ + enabled: boolean; + /** + * Team Ids + * @default [] + */ + team_ids: string[]; /** User Id */ user_id?: string | null; /** User Name */ @@ -67492,9 +67436,10 @@ export interface operations { query?: string; limit?: number; offset?: number; - key_id?: string | null; before_updated_at?: string | null; before_memory_id?: string | null; + team_id?: string | null; + user_id?: string | null; }; header?: never; path?: never; @@ -67524,9 +67469,7 @@ export interface operations { }; capture_entry_v2_memory_entries_post: { parameters: { - query?: { - key_id?: string | null; - }; + query?: never; header?: never; path?: never; cookie?: never; @@ -67557,11 +67500,75 @@ export interface operations { }; }; }; + read_entry_v2_memory_entries__memory_id__get: { + parameters: { + query?: never; + header?: never; + path: { + memory_id: string; + }; + cookie?: never; + }; + requestBody?: never; + responses: { + /** @description Successful Response */ + 200: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": components["schemas"]["MemoryEntry"]; + }; + }; + /** @description Validation Error */ + 422: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": components["schemas"]["HTTPValidationError"]; + }; + }; + }; + }; + update_entry_v2_memory_entries__memory_id__put: { + parameters: { + query?: never; + header?: never; + path: { + memory_id: string; + }; + cookie?: never; + }; + requestBody: { + content: { + "application/json": components["schemas"]["MemoryCapture"]; + }; + }; + responses: { + /** @description Successful Response */ + 200: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": components["schemas"]["MemoryEntry"]; + }; + }; + /** @description Validation Error */ + 422: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": components["schemas"]["HTTPValidationError"]; + }; + }; + }; + }; delete_entry_v2_memory_entries__memory_id__delete: { parameters: { - query?: { - key_id?: string | null; - }; + query?: never; header?: never; path: { memory_id: string; @@ -67588,13 +67595,9 @@ export interface operations { }; }; }; - list_policies_v2_memory_policies_get: { + get_settings_v2_memory_settings_get: { parameters: { - query?: { - target_type?: ("gateway" | "organization" | "team" | "project" | "user" | "key") | null; - target_id?: string | null; - offset?: number; - }; + query?: never; header?: never; path?: never; cookie?: never; @@ -67607,21 +67610,12 @@ export interface operations { [name: string]: unknown; }; content: { - "application/json": components["schemas"]["MemoryPolicy"][]; - }; - }; - /** @description Validation Error */ - 422: { - headers: { - [name: string]: unknown; - }; - content: { - "application/json": components["schemas"]["HTTPValidationError"]; + "application/json": components["schemas"]["MemorySettings"]; }; }; }; }; - set_policy_v2_memory_policies_put: { + set_settings_v2_memory_settings_put: { parameters: { query?: never; header?: never; @@ -67630,7 +67624,7 @@ export interface operations { }; requestBody: { content: { - "application/json": components["schemas"]["MemoryPolicyInput"]; + "application/json": components["schemas"]["MemorySettings"]; }; }; responses: { @@ -67640,102 +67634,7 @@ export interface operations { [name: string]: unknown; }; content: { - "application/json": components["schemas"]["MemoryPolicy"]; - }; - }; - /** @description Validation Error */ - 422: { - headers: { - [name: string]: unknown; - }; - content: { - "application/json": components["schemas"]["HTTPValidationError"]; - }; - }; - }; - }; - delete_policy_v2_memory_policies__policy_id__delete: { - parameters: { - query?: never; - header?: never; - path: { - policy_id: string; - }; - cookie?: never; - }; - requestBody?: never; - responses: { - /** @description Successful Response */ - 204: { - headers: { - [name: string]: unknown; - }; - content?: never; - }; - /** @description Validation Error */ - 422: { - headers: { - [name: string]: unknown; - }; - content: { - "application/json": components["schemas"]["HTTPValidationError"]; - }; - }; - }; - }; - get_preference_v2_memory_preference_get: { - parameters: { - query?: { - key_id?: string | null; - }; - header?: never; - path?: never; - cookie?: never; - }; - requestBody?: never; - responses: { - /** @description Successful Response */ - 200: { - headers: { - [name: string]: unknown; - }; - content: { - "application/json": components["schemas"]["MemoryPreference"]; - }; - }; - /** @description Validation Error */ - 422: { - headers: { - [name: string]: unknown; - }; - content: { - "application/json": components["schemas"]["HTTPValidationError"]; - }; - }; - }; - }; - set_preference_v2_memory_preference_put: { - parameters: { - query?: { - key_id?: string | null; - }; - header?: never; - path?: never; - cookie?: never; - }; - requestBody: { - content: { - "application/json": components["schemas"]["MemoryPreference"]; - }; - }; - responses: { - /** @description Successful Response */ - 200: { - headers: { - [name: string]: unknown; - }; - content: { - "application/json": components["schemas"]["MemoryPreference"]; + "application/json": components["schemas"]["MemorySettings"]; }; }; /** @description Validation Error */ @@ -67751,9 +67650,7 @@ export interface operations { }; get_status_v2_memory_status_get: { parameters: { - query?: { - key_id?: string | null; - }; + query?: never; header?: never; path?: never; cookie?: never; @@ -67769,15 +67666,6 @@ export interface operations { "application/json": components["schemas"]["MemoryStatus"]; }; }; - /** @description Validation Error */ - 422: { - headers: { - [name: string]: unknown; - }; - content: { - "application/json": components["schemas"]["HTTPValidationError"]; - }; - }; }; }; model_info_v2_v2_model_info_get: {