fix(proxy): atomically enqueue relay collector logs

This commit is contained in:
Ishaan Jaff 2026-07-09 16:45:28 -07:00
parent 9975fd4683
commit ed5b901983
No known key found for this signature in database
2 changed files with 102 additions and 32 deletions

View file

@ -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)}

View file

@ -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",
]