From 610a5a4579db8e2e04c50b17e1112017dd657d86 Mon Sep 17 00:00:00 2001 From: moe-berri Date: Tue, 15 Sep 2026 13:55:12 -0700 Subject: [PATCH] fix(memory): enforce RPM for every model round --- .../hooks/parallel_request_limiter_v3.py | 15 +--- litellm/proxy/memory/store.py | 7 +- litellm/proxy/memory/transport.py | 4 - .../proxy/memory/test_memory_v2_boundaries.py | 76 +++++++++++++++---- .../proxy/memory/test_memory_v2_management.py | 2 +- .../proxy/memory/test_transport.py | 46 ++++++----- 6 files changed, 93 insertions(+), 57 deletions(-) diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index cd25f1267f0..678d65d260b 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -3540,24 +3540,11 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): request_data: Final = _REQUEST_RATE_LIMIT_DATA.validate_python(data) model_value: Final = request_data.get("model") requested_model: Final = model_value if isinstance(model_value, str) else None - from litellm.proxy.memory.transport import is_memory_continuation_round - - built_descriptors: Final = await self._build_request_rate_limit_descriptors( + descriptors: Final = await self._build_request_rate_limit_descriptors( user_api_key_dict=user_api_key_dict, data=request_data, call_type=call_type, ) - descriptors: Final = [ # mutable-ok: Existing limiter helpers consume native descriptor containers. - { # mutable-ok: Existing limiter helpers consume native descriptor containers. - **descriptor, - "rate_limit": { # mutable-ok: Existing limiter helpers consume native descriptor containers. - key: value for key, value in (descriptor["rate_limit"] or {}).items() if key != "requests_per_unit" - }, - } - if is_memory_continuation_round() - else descriptor - for descriptor in built_descriptors - ] # Only check rate limits if we have descriptors with actual limits if descriptors: diff --git a/litellm/proxy/memory/store.py b/litellm/proxy/memory/store.py index ad44c824f37..0fd2b62a601 100644 --- a/litellm/proxy/memory/store.py +++ b/litellm/proxy/memory/store.py @@ -1,5 +1,4 @@ import asyncio -import json from collections.abc import Mapping from datetime import datetime from types import MappingProxyType, SimpleNamespace @@ -222,6 +221,8 @@ class MemoryStore: async def _capture( self, capture: MemoryCapture, namespace: str, table: TableActions["LiteLLM_MemoryTable"] ) -> MemoryEntry: + from prisma import Json + key: Final = f"memory-v2:{namespace}:{capture.key}" metadata: Final = { # mutable-ok: Prisma query and write JSON. name: redact_memory(value) @@ -238,7 +239,7 @@ class MemoryStore: content: Final = redact_memory(capture.content) data: Final = { # mutable-ok: Prisma query and write JSON. "value": content, - "metadata": json.dumps(metadata), + "metadata": Json(metadata), "updated_by": self.actor, } existing: Final = await table.find_unique( @@ -261,7 +262,7 @@ class MemoryStore: "updated_at": capture.expected_revision, "value": existing.value, "metadata": { # mutable-ok: Prisma query and write JSON. - "equals": json.dumps(existing.metadata) + "equals": Json(existing.metadata) }, }, data=data, diff --git a/litellm/proxy/memory/transport.py b/litellm/proxy/memory/transport.py index 1dab8907fcd..baab92890f6 100644 --- a/litellm/proxy/memory/transport.py +++ b/litellm/proxy/memory/transport.py @@ -41,10 +41,6 @@ def in_gateway_round() -> bool: return _GATEWAY_ROUND.get() is not None -def is_memory_continuation_round() -> bool: - return (_GATEWAY_ROUND.get() or 0) > 0 - - def _round_body(body: Mapping[str, object]) -> bytes: return json.dumps( { # mutable-ok: Starlette and the gateway processor consume native request containers. 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 a87ed8bbd4e..4047a2db20d 100644 --- a/tests/test_litellm/proxy/memory/test_memory_v2_boundaries.py +++ b/tests/test_litellm/proxy/memory/test_memory_v2_boundaries.py @@ -25,7 +25,7 @@ from litellm.proxy.memory.knowledge import MEMORY_TOOL_NAMES, execute_memory_too 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.proxy.memory.transport import is_memory_continuation_round +from litellm.proxy.memory.transport import in_gateway_round from litellm.types.memory_v2 import MemoryCapture, MemoryEnrollment, MemorySearch, MemorySettings _NOW: Final = datetime(2026, 9, 12, tzinfo=timezone.utc) @@ -627,6 +627,44 @@ async def test_memory_lookup_failure_leaves_inference_unchanged_but_never_leaks_ _ROUTES: Final = ("acompletion", "aresponses", "anthropic_messages") +@pytest.mark.asyncio +@pytest.mark.parametrize("route", _ROUTES) +@pytest.mark.parametrize("streaming", (False, True)) +@pytest.mark.parametrize("configuration", ("absent", "disabled", "no_database")) +async def test_disabled_memory_leaves_inference_untouched( + prisma_edge: MagicMock, route: ServerToolRoute, streaming: bool, configuration: str +) -> None: + from unittest.mock import patch + + from litellm.caching.caching import DualCache + from litellm.proxy.memory.gateway import process_gateway_memory + + prisma_edge.db.litellm_config.find_unique.return_value = ( + None if configuration == "absent" else SimpleNamespace(param_value=MemorySettings().model_dump()) + ) + data: Final = { + "messages": [{"role": "system", "content": "Original instructions"}, {"role": "user", "content": "hi"}], + "input": "hi", + "stream": streaming, + "caching": True, + "tools": [{"type": "function", "function": {"name": "litellm_memory_search"}}], + } + original: Final = json.dumps(data) + execute: Final = AsyncMock() + with patch.multiple( # test-quality-ok: Inject the external configuration DB and worker cache. + "litellm.proxy.proxy_server", + prisma_client=None if configuration == "no_database" else prisma_edge, + user_api_key_cache=DualCache(), + ): + assert await process_gateway_memory(data, request(), UserAPIKeyAuth(user_id="owner"), route, execute) is None + assert json.dumps(data) == original + execute.assert_not_awaited() + prisma_edge.db.litellm_memorytable.find_many.assert_not_awaited() + prisma_edge.db.litellm_memorytable.create.assert_not_awaited() + prisma_edge.db.litellm_memorycontinuation.upsert.assert_not_awaited() + prisma_edge.db.litellm_usertable.find_unique.assert_not_awaited() + + @pytest.mark.asyncio @pytest.mark.parametrize("route", _ROUTES) async def test_malformed_nonstreaming_provider_body_returns_502(prisma_edge: MagicMock, route: ServerToolRoute) -> None: @@ -854,7 +892,7 @@ async def test_search_is_private_and_client_tools_keep_their_ids( async def execute(inner: Request, body: dict[str, object], auth: UserAPIKeyAuth) -> Response: observed.append(body) - assert is_memory_continuation_round() == (len(observed) > 1) + assert in_gateway_round() if len(observed) == 1: reply: Final = provider_response( route, "INTERNAL HOUSEKEEPING", (search, application) if client_tool else (search,) @@ -945,7 +983,7 @@ async def test_capture_rejects_fabricated_evidence_and_deduplicates_across_keys( assert saved["saved"] == 1 prisma_edge.db.litellm_memorytable.find_unique.return_value = row( key=prisma_edge.db.litellm_memorytable.create.call_args.kwargs["data"]["key"], - metadata=prisma_edge.db.litellm_memorytable.create.call_args.kwargs["data"]["metadata"], + metadata=json.dumps(prisma_edge.db.litellm_memorytable.create.call_args.kwargs["data"]["metadata"].data), value=observation["content"], ) repeated: Final = await execute_memory_tool(memory, call, ({"role": "user", "content": "Use port 8347"},)) @@ -1027,21 +1065,29 @@ async def test_previous_response_uses_owned_upstream_and_pending_tool_outputs(pr ), ) @pytest.mark.parametrize("status", (429, 502)) +@pytest.mark.parametrize("route", _ROUTES) +@pytest.mark.parametrize("streaming", (False, True)) async def test_error_response_retains_status_and_retry_after_without_parsing_provider_body( - prisma_edge: MagicMock, body: bytes, status: int + prisma_edge: MagicMock, body: bytes, status: int, route: ServerToolRoute, streaming: bool ) -> None: + from unittest.mock import patch + + from litellm.caching.caching import DualCache + from litellm.proxy.memory.gateway import process_gateway_memory + execute = AsyncMock(return_value=Response(body, status_code=status, headers={"retry-after": "7"})) - loop = GatewayMemoryLoop( - execute, - request(), - {"messages": [{"role": "user", "content": "hi"}]}, - "acompletion", - store(prisma_edge), - UserAPIKeyAuth(), - ) - with pytest.raises(HTTPException) as error: - async for _ in loop.run(): - pass + caller = UserAPIKeyAuth(user_id="owner", token="a" * 64) + with patch.multiple( # test-quality-ok: Inject external database, cache, and provider error responses. + "litellm.proxy.proxy_server", prisma_client=prisma_edge, user_api_key_cache=DualCache(), llm_router=None + ): + with pytest.raises(HTTPException) as error: + await process_gateway_memory( + {"messages": [{"role": "user", "content": "hi"}], "input": "hi", "stream": streaming}, + request(), + caller, + route, + execute, + ) assert error.value.status_code == status assert error.value.headers["retry-after"] == "7" assert "private provider" not in error.value.detail 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 b5b4c9774fa..f329ea74e8c 100644 --- a/tests/test_litellm/proxy/memory/test_memory_v2_management.py +++ b/tests/test_litellm/proxy/memory/test_memory_v2_management.py @@ -415,7 +415,7 @@ async def test_capture_redacts_credentials_before_persisting_content_and_metadat fields = ("title", "content", "evidence", "when_to_use", "scope", "source") await management.capture_entry(_CAPTURE.model_copy(update={field: text for field in fields}), auth()) data = database.db.litellm_memorytable.create.call_args.kwargs["data"] - stored = {"content": data["value"], **json.loads(data["metadata"])} + stored = {"content": data["value"], **data["metadata"].data} for field in fields: assert secret not in stored[field] assert "REDACTED" in stored[field] diff --git a/tests/test_litellm/proxy/memory/test_transport.py b/tests/test_litellm/proxy/memory/test_transport.py index dc69c61cd7d..57d210f22f4 100644 --- a/tests/test_litellm/proxy/memory/test_transport.py +++ b/tests/test_litellm/proxy/memory/test_transport.py @@ -72,7 +72,9 @@ async def test_stream_reaches_client_before_model_finishes_and_disconnect_cancel @pytest.mark.asyncio -async def test_memory_rounds_share_one_rpm_admission_but_keep_token_limits(): +@pytest.mark.parametrize("scope", ("api_key", "team")) +@pytest.mark.parametrize("limit_type", ("requests", "tokens")) +async def test_every_memory_round_enforces_rate_limits(scope: str, limit_type: str) -> None: from fastapi import HTTPException from starlette.responses import Response @@ -80,32 +82,36 @@ async def test_memory_rounds_share_one_rpm_admission_but_keep_token_limits(): from litellm.proxy.hooks.parallel_request_limiter_v3 import ( _PROXY_MaxParallelRequestsHandler_v3 as _PROXY_MaxParallelRequestsHandler, ) - from litellm.proxy.memory.transport import gateway_round from litellm.proxy.utils import InternalUsageCache, hash_token - cache = DualCache() - handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(cache)) - auth = UserAPIKeyAuth(api_key=hash_token("memory-round-limits"), rpm_limit=1, tpm_limit=100) - request = Request({"type": "http", "method": "POST", "path": "/v1/chat/completions", "headers": []}) + cache: Final = DualCache() + handler: Final = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(cache)) + auth: Final = UserAPIKeyAuth( + api_key=hash_token("memory-round-limits"), + rpm_limit=2 if scope == "api_key" else None, + tpm_limit=100 if scope == "api_key" else None, + team_id="memory-team", + team_rpm_limit=2 if scope == "team" else None, + team_tpm_limit=100 if scope == "team" else None, + ) + request: Final = Request({"type": "http", "method": "POST", "path": "/v1/chat/completions", "headers": []}) - async def execute(inner, body, round_auth): - await handler.async_pre_call_hook(round_auth, cache, body, "") + async def execute(inner: Request, body: dict[str, object], round_auth: UserAPIKeyAuth) -> Response: + await handler.async_pre_call_hook(round_auth, cache, body, "completion") return Response(b"accepted") - for index in (0, 1): - async with gateway_round(execute, request, {"model": "test"}, auth, index) as round: - assert await round.read() == b"accepted" + admitted_rounds: Final = 2 if limit_type == "requests" else 1 + for index in range(admitted_rounds): + async with gateway_round(execute, request, {"model": "test"}, auth, index) as admitted: + assert await admitted.read() == b"accepted" + if limit_type == "tokens": + owner: Final = auth.api_key if scope == "api_key" else auth.team_id + await cache.async_set_cache(key=f"{{{scope}:{owner}}}:tokens", value=101, ttl=60) with pytest.raises(HTTPException) as rate_limited: - async with gateway_round(execute, request, {"model": "test"}, auth, 0): - pass + async with gateway_round(execute, request, {"model": "test"}, auth, admitted_rounds): + pytest.fail("A memory continuation exceeded its rate limit") assert rate_limited.value.status_code == 429 - assert "requests" in rate_limited.value.detail - await cache.async_set_cache(key=f"{{api_key:{auth.api_key}}}:tokens", value=101, ttl=60) - with pytest.raises(HTTPException) as token_limited: - async with gateway_round(execute, request, {"model": "test"}, auth, 1): - pass - assert token_limited.value.status_code == 429 - assert "tokens" in token_limited.value.detail + assert limit_type in rate_limited.value.detail @pytest.mark.asyncio