diff --git a/deploy/memory-pilot/pilot.py b/deploy/memory-pilot/pilot.py index 27a9a040332..ef3f7a8b19f 100644 --- a/deploy/memory-pilot/pilot.py +++ b/deploy/memory-pilot/pilot.py @@ -14,6 +14,7 @@ from starlette.responses import JSONResponse from starlette.types import ASGIApp, Receive, Scope, Send from litellm.caching.caching import DualCache +from litellm.caching.in_memory_cache import InMemoryCache from litellm.integrations.custom_logger import CustomLogger from litellm.litellm_core_utils.prompt_templates.server_tool_responses import object_value from litellm.llms.custom_httpx.http_handler import get_async_httpx_client @@ -68,6 +69,7 @@ class PilotGateway: self.app = app self.registered_validation_slots = asyncio.Semaphore(12) self.enrollment_validation_slots = asyncio.Semaphore(4) + self.recently_validated = InMemoryCache(max_size_in_memory=1000, default_ttl=60) self.upstream = get_async_httpx_client( httpxSpecialProvider.PassThroughEndpoint, params={"timeout": 20, "client_alias": "memory-pilot-upstream"}, @@ -123,7 +125,11 @@ class PilotGateway: {"error": "Upstream keys can only use inference and their own memories"}, status_code=403 )(scope, receive, send) return - validation_slots: Final = self.registered_validation_slots if local_key else self.enrollment_validation_slots + validation_slots: Final = ( + self.registered_validation_slots + if self.recently_validated.get_cache(digest) + else self.enrollment_validation_slots + ) try: await asyncio.wait_for(validation_slots.acquire(), timeout=0.05) except TimeoutError: @@ -137,6 +143,10 @@ class PilotGateway: models: Final = await self.upstream.get( _UPSTREAM + "/v1/models", headers={"Authorization": "Bearer " + credential} ) + if models.status_code in (401, 403): + self.recently_validated.delete_cache(digest) + elif models.is_success: + self.recently_validated.set_cache(digest, True) except httpx.HTTPError: await JSONResponse({"error": "Upstream gateway unavailable"}, status_code=503)(scope, receive, send) return diff --git a/tests/test_litellm/proxy/memory/test_memory_pilot.py b/tests/test_litellm/proxy/memory/test_memory_pilot.py index 72075216bf8..eb64ac01398 100644 --- a/tests/test_litellm/proxy/memory/test_memory_pilot.py +++ b/tests/test_litellm/proxy/memory/test_memory_pilot.py @@ -49,6 +49,8 @@ async def test_unknown_keys_cannot_consume_registered_validation_capacity_and_sl "litellm.proxy.proxy_server", prisma_client=database, master_key="local-admin" ): async with httpx.AsyncClient(transport=httpx.ASGITransport(app=gateway), base_url="http://pilot") as client: + initial = await client.get("/v1/models", headers={"Authorization": "Bearer sk-established"}) + assert initial.status_code == 200 pending = [ asyncio.create_task(client.get("/v1/models", headers={"Authorization": f"Bearer sk-invalid-{i}"})) for i in range(4) @@ -56,11 +58,74 @@ async def test_unknown_keys_cannot_consume_registered_validation_capacity_and_sl await asyncio.wait_for(entered.wait(), timeout=2) refused = await client.get("/v1/models", headers={"Authorization": "Bearer sk-overload"}) assert refused.status_code == 503 and refused.headers["retry-after"] == "1" - assert upstream.get.await_count == 4 + assert upstream.get.await_count == 5 established = await client.get("/v1/models", headers={"Authorization": "Bearer sk-established"}) assert established.status_code == 200 and established.json() == {"data": [{"id": "model"}]} release.set() assert all(response.status_code == 503 for response in await asyncio.gather(*pending)) again = await client.get("/v1/models", headers={"Authorization": "Bearer sk-next"}) assert again.status_code == 503 and "unavailable" in again.text - assert upstream.get.await_count == 6 + assert upstream.get.await_count == 7 + + +@pytest.mark.asyncio +@pytest.mark.parametrize("rejection_status", [401, 403]) +async def test_rejected_enrolled_keys_lose_reserved_capacity_and_can_revalidate( + monkeypatch: pytest.MonkeyPatch, rejection_status: int +) -> None: + monkeypatch.setenv("UPSTREAM_LITELLM_BASE_URL", "https://upstream.example.invalid") + filename = Path(__file__).resolve().parents[4] / "deploy" / "memory-pilot" / "pilot.py" + spec = importlib.util.spec_from_file_location("memory_pilot_revoked_test", filename) + assert spec is not None and spec.loader is not None + module = importlib.util.module_from_spec(spec) + spec.loader.exec_module(module) + gateway = module.PilotGateway(Starlette()) + revoked = asyncio.Event() + block = asyncio.Event() + entered = asyncio.Event() + release = asyncio.Event() + count = 0 + + async def upstream_get(*args: object, **kwargs: object) -> httpx.Response: + if kwargs.get("headers") == {"Authorization": "Bearer sk-revoked"}: + if block.is_set(): + nonlocal count + count += 1 + if count == 4: + entered.set() + await release.wait() + if revoked.is_set(): + return httpx.Response(rejection_status, json={"error": "rejected"}) + return httpx.Response(200, json={"data": []}) + + gateway.upstream = MagicMock(get=AsyncMock(side_effect=upstream_get)) + database = MagicMock() + database.db.litellm_verificationtoken.find_unique = AsyncMock(return_value={"token": "enrolled"}) + with patch.multiple( # test-quality-ok: Keep enrollment rows present while upstream revokes access. + "litellm.proxy.proxy_server", prisma_client=database, master_key="local-admin" + ): + async with httpx.AsyncClient(transport=httpx.ASGITransport(app=gateway), base_url="http://pilot") as client: + for credential in ("sk-established", "sk-revoked"): + initial = await client.get("/v1/models", headers={"Authorization": f"Bearer {credential}"}) + assert initial.status_code == 200 + revoked.set() + rejected = await client.get("/v1/models", headers={"Authorization": "Bearer sk-revoked"}) + assert rejected.status_code == rejection_status + block.set() + pending = [ + asyncio.create_task(client.get("/v1/models", headers={"Authorization": "Bearer sk-revoked"})) + for _ in range(4) + ] + await asyncio.wait_for(entered.wait(), timeout=2) + refused = await client.get("/v1/models", headers={"Authorization": "Bearer sk-revoked"}) + assert refused.status_code == 503 and refused.headers["retry-after"] == "1" + established = await client.get("/v1/models", headers={"Authorization": "Bearer sk-established"}) + assert established.status_code == 200 + release.set() + assert all(response.status_code == rejection_status for response in await asyncio.gather(*pending)) + revoked.clear() + restored = await client.get("/v1/models", headers={"Authorization": "Bearer sk-revoked"}) + assert restored.status_code == 200 + revoked.set() + rechecked = await client.get("/v1/models", headers={"Authorization": "Bearer sk-revoked"}) + assert rechecked.status_code == rejection_status