From ed5b90198391b7773a58b0ad1b095903f2861f2b Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Thu, 9 Jul 2026 16:45:28 -0700 Subject: [PATCH] fix(proxy): atomically enqueue relay collector logs --- .../proxy/collector_endpoints/spend_logs.py | 91 +++++++++++++------ .../collector_endpoints/test_spend_logs.py | 43 ++++++++- 2 files changed, 102 insertions(+), 32 deletions(-) diff --git a/litellm/proxy/collector_endpoints/spend_logs.py b/litellm/proxy/collector_endpoints/spend_logs.py index 11a04681529..a69f9636c1e 100644 --- a/litellm/proxy/collector_endpoints/spend_logs.py +++ b/litellm/proxy/collector_endpoints/spend_logs.py @@ -46,6 +46,24 @@ class CollectorSpendLogsIngestResponse(TypedDict): enqueued: int +async def _enqueue_collector_spend_logs( + prisma_client: Any, + spend_logs: list[CollectorSpendLogRow], +) -> None: + async with prisma_client._spend_log_transactions_lock: + queued_spend_logs = len(prisma_client.spend_log_transactions) + if queued_spend_logs + len(spend_logs) > LITELLM_ASYNCIO_QUEUE_MAXSIZE: + raise HTTPException( + status_code=status.HTTP_429_TOO_MANY_REQUESTS, + detail={ + "error": "Collector spend-log queue is full", + "queued": queued_spend_logs, + "limit": LITELLM_ASYNCIO_QUEUE_MAXSIZE, + }, + ) + prisma_client.spend_log_transactions.extend(spend_logs) + + class CollectorSpendLogTransformer: PASSTHROUGH_FIELDS = { "total_tokens", @@ -97,10 +115,15 @@ class CollectorSpendLogTransformer: now: datetime, ) -> CollectorSpendLogRow: collector_request_id = log.get("request_id") - if not isinstance(collector_request_id, str) or not collector_request_id.strip(): + if ( + not isinstance(collector_request_id, str) + or not collector_request_id.strip() + ): raise HTTPException( status_code=status.HTTP_400_BAD_REQUEST, - detail={"error": "request_id is required for collector spend-log ingestion"}, + detail={ + "error": "request_id is required for collector spend-log ingestion" + }, ) key_hash = CollectorSpendLogTransformer._get_auth_key_hash(user_api_key_dict) @@ -122,7 +145,9 @@ class CollectorSpendLogTransformer: "custom_llm_provider": "", "user": getattr(user_api_key_dict, "user_id", None), "team_id": getattr(user_api_key_dict, "team_id", None), - "organization_id": CollectorSpendLogTransformer._get_auth_organization_id(user_api_key_dict), + "organization_id": CollectorSpendLogTransformer._get_auth_organization_id( + user_api_key_dict + ), "metadata": CollectorSpendLogTransformer._normalize_metadata( log.get("metadata"), user_api_key_dict, @@ -130,7 +155,9 @@ class CollectorSpendLogTransformer: ), "cache_hit": "False", "cache_key": "", - "request_tags": CollectorSpendLogTransformer._normalize_request_tags(log.get("request_tags")), + "request_tags": CollectorSpendLogTransformer._normalize_request_tags( + log.get("request_tags") + ), "messages": {}, "response": {}, "proxy_server_request": {}, @@ -145,11 +172,15 @@ class CollectorSpendLogTransformer: @staticmethod def _get_auth_organization_id(user_api_key_dict: Any) -> Optional[str]: - return getattr(user_api_key_dict, "organization_id", None) or getattr(user_api_key_dict, "org_id", None) + return getattr(user_api_key_dict, "organization_id", None) or getattr( + user_api_key_dict, "org_id", None + ) @staticmethod def _get_auth_key_hash(user_api_key_dict: Any) -> Optional[str]: - return getattr(user_api_key_dict, "api_key", None) or getattr(user_api_key_dict, "token", None) + return getattr(user_api_key_dict, "api_key", None) or getattr( + user_api_key_dict, "token", None + ) @staticmethod def _get_auth_key_alias(user_api_key_dict: Any) -> str: @@ -164,7 +195,9 @@ class CollectorSpendLogTransformer: return getattr(user_api_key_dict, "team_alias", None) or None @staticmethod - def _collector_request_id_for(key_hash: Optional[str], collector_request_id: str) -> str: + def _collector_request_id_for( + key_hash: Optional[str], collector_request_id: str + ) -> str: digest = hmac.new( (key_hash or LITELLM_RELAY_CALL_TYPE).encode(), collector_request_id.encode(), @@ -192,11 +225,17 @@ class CollectorSpendLogTransformer: "collector_request_id": collector_request_id, "relay_request_id": collector_request_id, "user_api_key": key_hash, - "user_api_key_alias": CollectorSpendLogTransformer._get_auth_key_alias(user_api_key_dict), + "user_api_key_alias": CollectorSpendLogTransformer._get_auth_key_alias( + user_api_key_dict + ), "user_api_key_user_id": getattr(user_api_key_dict, "user_id", None), "user_api_key_team_id": getattr(user_api_key_dict, "team_id", None), - "user_api_key_team_alias": CollectorSpendLogTransformer._get_auth_team_alias(user_api_key_dict), - "user_api_key_org_id": CollectorSpendLogTransformer._get_auth_organization_id(user_api_key_dict), + "user_api_key_team_alias": CollectorSpendLogTransformer._get_auth_team_alias( + user_api_key_dict + ), + "user_api_key_org_id": CollectorSpendLogTransformer._get_auth_organization_id( + user_api_key_dict + ), } ) return normalized @@ -273,7 +312,9 @@ class CollectorSpendLogTransformer: if encoded_size > MAX_COLLECTOR_SPEND_LOG_BYTES: raise HTTPException( status_code=status.HTTP_413_CONTENT_TOO_LARGE, - detail={"error": f"Collector spend-log entry exceeds {MAX_COLLECTOR_SPEND_LOG_BYTES} bytes"}, + detail={ + "error": f"Collector spend-log entry exceeds {MAX_COLLECTOR_SPEND_LOG_BYTES} bytes" + }, ) return encoded_size @@ -282,7 +323,9 @@ class CollectorSpendLogTransformer: if total_bytes > MAX_COLLECTOR_SPEND_LOG_BATCH_BYTES: raise HTTPException( status_code=status.HTTP_413_CONTENT_TOO_LARGE, - detail={"error": f"Collector spend-log batch exceeds {MAX_COLLECTOR_SPEND_LOG_BATCH_BYTES} bytes"}, + detail={ + "error": f"Collector spend-log batch exceeds {MAX_COLLECTOR_SPEND_LOG_BATCH_BYTES} bytes" + }, ) @staticmethod @@ -316,7 +359,6 @@ async def ingest_collector_spend_logs( Ingest LiteLLM Relay captures into the existing spend-log batcher so they appear in the Gateway Logs UI without replaying captured traffic. """ - proxy_logging_obj = getattr(request.app.state, "proxy_logging_obj", None) prisma_client = getattr(request.app.state, "prisma_client", None) if prisma_client is None: @@ -324,7 +366,7 @@ async def ingest_collector_spend_logs( status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, detail={"error": "Prisma Client is not initialized"}, ) - if proxy_logging_obj is None: + if getattr(request.app.state, "proxy_logging_obj", None) is None: raise HTTPException( status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, detail={"error": "Proxy logging is not initialized"}, @@ -336,18 +378,8 @@ async def ingest_collector_spend_logs( if len(logs) > MAX_COLLECTOR_SPEND_LOGS: raise HTTPException( status_code=status.HTTP_400_BAD_REQUEST, - detail={"error": f"Collector spend-log ingestion is limited to {MAX_COLLECTOR_SPEND_LOGS} rows"}, - ) - - async with prisma_client._spend_log_transactions_lock: - queued_spend_logs = len(prisma_client.spend_log_transactions) - if queued_spend_logs + len(logs) > LITELLM_ASYNCIO_QUEUE_MAXSIZE: - raise HTTPException( - status_code=status.HTTP_429_TOO_MANY_REQUESTS, detail={ - "error": "Collector spend-log queue is full", - "queued": queued_spend_logs, - "limit": LITELLM_ASYNCIO_QUEUE_MAXSIZE, + "error": f"Collector spend-log ingestion is limited to {MAX_COLLECTOR_SPEND_LOGS} rows" }, ) @@ -356,10 +388,9 @@ async def ingest_collector_spend_logs( user_api_key_dict=user_api_key_dict, now=datetime.now(timezone.utc), ) - for spend_log in spend_logs: - await proxy_logging_obj.db_spend_update_writer._insert_spend_log_to_db( - payload=spend_log, - prisma_client=prisma_client, - ) + await _enqueue_collector_spend_logs( + prisma_client=prisma_client, + spend_logs=spend_logs, + ) return {"enqueued": len(spend_logs)} diff --git a/tests/test_litellm/proxy/collector_endpoints/test_spend_logs.py b/tests/test_litellm/proxy/collector_endpoints/test_spend_logs.py index a7365bf8b43..3e91230bcda 100644 --- a/tests/test_litellm/proxy/collector_endpoints/test_spend_logs.py +++ b/tests/test_litellm/proxy/collector_endpoints/test_spend_logs.py @@ -1,6 +1,7 @@ import asyncio import json +import pytest from fastapi.testclient import TestClient import litellm.proxy.proxy_server as ps @@ -207,7 +208,9 @@ def test_collector_spend_logs_rejects_oversized_log(monkeypatch): "logs": [ { "request_id": "large-relay-request", - "proxy_server_request": {"body_preview": "x" * (MAX_COLLECTOR_SPEND_LOG_BYTES + 1)}, + "proxy_server_request": { + "body_preview": "x" * (MAX_COLLECTOR_SPEND_LOG_BYTES + 1) + }, } ] }, @@ -236,7 +239,9 @@ def test_collector_spend_logs_rejects_normalized_row_over_size_limit(monkeypatch "logs": [ { "request_id": "normalized-large-relay-request", - "proxy_server_request": {"body_preview": "x" * (MAX_COLLECTOR_SPEND_LOG_BYTES - 50)}, + "proxy_server_request": { + "body_preview": "x" * (MAX_COLLECTOR_SPEND_LOG_BYTES - 50) + }, } ] }, @@ -321,3 +326,37 @@ def test_collector_spend_logs_rejects_when_queue_is_full(monkeypatch): assert len(prisma_client.spend_log_transactions) == 3 finally: app.dependency_overrides.pop(ps.user_api_key_auth, None) + + +def test_collector_spend_logs_enqueue_is_capacity_checked_under_lock(monkeypatch): + prisma_client = MockPrismaClient() + monkeypatch.setattr( + collector_spend_logs, + "LITELLM_ASYNCIO_QUEUE_MAXSIZE", + 3, + ) + + async def enqueue_two_batches(): + await collector_spend_logs._enqueue_collector_spend_logs( + prisma_client=prisma_client, + spend_logs=[{"request_id": "one"}, {"request_id": "two"}], + ) + with pytest.raises(collector_spend_logs.HTTPException) as exc_info: + await collector_spend_logs._enqueue_collector_spend_logs( + prisma_client=prisma_client, + spend_logs=[{"request_id": "three"}, {"request_id": "four"}], + ) + return exc_info.value + + error = asyncio.run(enqueue_two_batches()) + + assert error.status_code == 429 + assert error.detail == { + "error": "Collector spend-log queue is full", + "queued": 2, + "limit": 3, + } + assert [row["request_id"] for row in prisma_client.spend_log_transactions] == [ + "one", + "two", + ]