diff --git a/basedpyright-code-budget.json b/basedpyright-code-budget.json index df52069e71f..10483d2ed64 100644 --- a/basedpyright-code-budget.json +++ b/basedpyright-code-budget.json @@ -1,6 +1,6 @@ { "reportAny": { - "limit": 14076 + "limit": 14075 }, "reportArgumentType": { "limit": 2216 @@ -57,7 +57,7 @@ "limit": 5601 }, "reportMissingTypeArgument": { - "limit": 15306 + "limit": 15304 }, "reportMissingTypeStubs": { "limit": 40 @@ -105,13 +105,13 @@ "limit": 109 }, "reportUnknownMemberType": { - "limit": 38350 + "limit": 38348 }, "reportUnknownParameterType": { - "limit": 19626 + "limit": 19624 }, "reportUnknownVariableType": { - "limit": 29890 + "limit": 29888 }, "reportUnnecessaryCast": { "limit": 111 diff --git a/litellm/integrations/gcs_bucket/gcs_bucket.py b/litellm/integrations/gcs_bucket/gcs_bucket.py index 31ceb338dcd..ca7bfef4046 100644 --- a/litellm/integrations/gcs_bucket/gcs_bucket.py +++ b/litellm/integrations/gcs_bucket/gcs_bucket.py @@ -14,6 +14,7 @@ from litellm.integrations.gcs_bucket.gcs_bucket_base import GCSBucketBase from litellm.litellm_core_utils.cloud_storage_security import ( sanitize_cloud_object_component, ) +from litellm.litellm_core_utils.spend_log_request_id import get_spend_logs_id from litellm.proxy._types import CommonProxyErrors from litellm.types.integrations.base_health_check import IntegrationHealthCheckStatus from litellm.types.integrations.gcs_bucket import * @@ -289,7 +290,7 @@ class GCSBucketLogger(GCSBucketBase, AdditionalLoggingUtils): else: object_name = self._generate_success_object_name( request_date_str=current_date, - response_id=response_obj.get("id", ""), + response_id=get_spend_logs_id(kwargs.get("call_type") or "acompletion", response_obj, kwargs) or "", ) # used for testing diff --git a/litellm/litellm_core_utils/spend_log_request_id.py b/litellm/litellm_core_utils/spend_log_request_id.py new file mode 100644 index 00000000000..7c3caf9da22 --- /dev/null +++ b/litellm/litellm_core_utils/spend_log_request_id.py @@ -0,0 +1,59 @@ +"""The id a request is filed under in LiteLLM_SpendLogs and in the per-request payload +stores (GCS) the logs viewer reads back through ``/spend/logs/ui/{request_id}``. + +Both sides must derive the id the same way or the viewer looks up a payload under a key +it was never stored under, so the derivation lives here rather than in either caller. +""" + +from __future__ import annotations + +from collections.abc import Mapping +from typing import Final + +from litellm.types.utils import CallTypes + +BATCH_COST_REQUEST_ID_SUFFIX: Final = "_batch_cost" + +_RESPONSE_ID_KEYED_CALL_TYPES: Final = frozenset( + { + CallTypes.acreate_batch.value, + CallTypes.aretrieve_batch.value, + CallTypes.acreate_file.value, + } +) +"""Batch and file rows key off the object's own id so repeated polls of the same +object collapse into one row instead of billing it once per poll. Every other call +type keys off the proxy-generated per-call id: request_id is the LiteLLM_SpendLogs +primary key and the flush inserts with skip_duplicates, so keying off the provider's +response id silently drops every row after the first whenever a provider (commonly a +self-hosted OpenAI-compatible server) reuses completion ids.""" + + +def _standard_logging_id(kwargs: Mapping[str, object]) -> str | None: + match kwargs.get("standard_logging_object"): + case {"id": str() as standard_logging_id}: + return standard_logging_id + case _: + return None + + +def get_spend_logs_id(call_type: str, response_obj: Mapping[str, object], kwargs: Mapping[str, object]) -> str | None: + candidate_ids: Final = ( + ( + response_obj.get("id"), + _standard_logging_id(kwargs), + kwargs.get("litellm_call_id"), + ) + if call_type in _RESPONSE_ID_KEYED_CALL_TYPES + else ( + kwargs.get("litellm_call_id"), + _standard_logging_id(kwargs), + response_obj.get("id"), + ) + ) + resolved_id: Final = next( + (candidate for candidate in candidate_ids if isinstance(candidate, str) and candidate), None + ) + if resolved_id is not None and call_type == CallTypes.aretrieve_batch.value: + return f"{resolved_id}{BATCH_COST_REQUEST_ID_SUFFIX}" + return resolved_id diff --git a/litellm/proxy/spend_tracking/spend_tracking_utils.py b/litellm/proxy/spend_tracking/spend_tracking_utils.py index c52d7f5f435..0da39126e5b 100644 --- a/litellm/proxy/spend_tracking/spend_tracking_utils.py +++ b/litellm/proxy/spend_tracking/spend_tracking_utils.py @@ -30,11 +30,13 @@ from litellm.litellm_core_utils.litellm_logging import ( request_model_access_groups_from_litellm_params, ) from litellm.litellm_core_utils.safe_json_dumps import safe_dumps, strip_null_bytes +from litellm.litellm_core_utils.spend_log_request_id import ( + get_spend_logs_id, +) from litellm.proxy._types import SpendLogsMetadata, SpendLogsPayload, SpendLogsRouterMetadata from litellm.proxy.spend_tracking.spend_log_error_logger import spend_log_error from litellm.proxy.utils import PrismaClient, hash_token from litellm.types.utils import ( - CallTypes, CostBreakdown, StandardLoggingGuardrailInformation, StandardLoggingMCPToolCall, @@ -211,49 +213,6 @@ def _get_spend_logs_metadata( return clean_metadata -BATCH_COST_REQUEST_ID_SUFFIX: Final = "_batch_cost" - -_RESPONSE_ID_KEYED_CALL_TYPES: Final = frozenset( - { - CallTypes.acreate_batch.value, - CallTypes.aretrieve_batch.value, - CallTypes.acreate_file.value, - } -) -"""Batch and file rows key off the object's own id so repeated polls of the same -object collapse into one row instead of billing it once per poll. Every other call -type keys off the proxy-generated per-call id: request_id is the LiteLLM_SpendLogs -primary key and the flush inserts with skip_duplicates, so keying off the provider's -response id silently drops every row after the first whenever a provider (commonly a -self-hosted OpenAI-compatible server) reuses completion ids.""" - - -def get_spend_logs_id(call_type: str, response_obj: dict, kwargs: dict) -> str | None: - standard_logging_payload = kwargs.get("standard_logging_object") - standard_logging_id: Final = ( - standard_logging_payload.get("id") if isinstance(standard_logging_payload, dict) else None - ) - candidate_ids: Final = ( - ( - response_obj.get("id"), - standard_logging_id, - kwargs.get("litellm_call_id"), - ) - if call_type in _RESPONSE_ID_KEYED_CALL_TYPES - else ( - kwargs.get("litellm_call_id"), - standard_logging_id, - response_obj.get("id"), - ) - ) - resolved_id: Final = next( - (candidate for candidate in candidate_ids if isinstance(candidate, str) and candidate), None - ) - if resolved_id is not None and call_type == CallTypes.aretrieve_batch.value: - return f"{resolved_id}{BATCH_COST_REQUEST_ID_SUFFIX}" - return resolved_id - - _MISSING_ATTRIBUTE: Final = object() diff --git a/tests/test_litellm/integrations/gcs_bucket/test_gcs_bucket_base.py b/tests/test_litellm/integrations/gcs_bucket/test_gcs_bucket_base.py index 8d662311da1..a5570408ce8 100644 --- a/tests/test_litellm/integrations/gcs_bucket/test_gcs_bucket_base.py +++ b/tests/test_litellm/integrations/gcs_bucket/test_gcs_bucket_base.py @@ -1,10 +1,13 @@ +import json import os +from datetime import datetime, timezone from unittest.mock import AsyncMock, MagicMock, patch import pytest from litellm.integrations.gcs_bucket.gcs_bucket import GCSBucketLogger from litellm.integrations.gcs_bucket.gcs_bucket_base import GCSBucketBase +from litellm.litellm_core_utils.spend_log_request_id import get_spend_logs_id class TestGCSBucketBase: @@ -128,3 +131,26 @@ class TestGCSBucketBase: assert object_name.endswith("-target_uploadType_media") assert ".." not in object_name assert "?" not in object_name + + @pytest.mark.asyncio + async def test_logs_viewer_finds_payloads_stored_for_requests_sharing_a_provider_response_id(self): + logger = GCSBucketLogger.__new__(GCSBucketLogger) + provider_response = {"id": "chatcmpl-reused"} + stored_objects = { + logger._get_object_name( + kwargs={"call_type": "acompletion", "litellm_call_id": call_id}, + logging_payload={"id": "chatcmpl-reused"}, + response_obj=provider_response, + ): json.dumps({"litellm_call_id": call_id}) + for call_id in ("call-id-1", "call-id-2") + } + assert len(stored_objects) == 2 + logger.download_gcs_object = AsyncMock(side_effect=lambda object_name: stored_objects.get(object_name)) + + payload = await logger.get_request_response_payload( + request_id=get_spend_logs_id("acompletion", provider_response, {"litellm_call_id": "call-id-2"}), + start_time_utc=datetime.now(timezone.utc), + end_time_utc=None, + ) + + assert payload == {"litellm_call_id": "call-id-2"} diff --git a/type-discipline-budget.json b/type-discipline-budget.json index 3d2e97d55a5..180201a68be 100644 --- a/type-discipline-budget.json +++ b/type-discipline-budget.json @@ -1,6 +1,6 @@ { "LIT001": { - "limit": 22367 + "limit": 22365 }, "LIT002": { "limit": 26777 @@ -27,7 +27,7 @@ "limit": 0 }, "LIT010": { - "limit": 16507 + "limit": 16506 }, "LIT011": { "limit": 5535