diff --git a/litellm/batches/batch_line_item_logging.py b/litellm/batches/batch_line_item_logging.py index 3d54625d3fe..98b6f0a1757 100644 --- a/litellm/batches/batch_line_item_logging.py +++ b/litellm/batches/batch_line_item_logging.py @@ -31,6 +31,20 @@ _BatchLineProvider: TypeAlias = Literal[ _SUPPORTED_LINE_PROVIDERS: Final = frozenset(get_args(_BatchLineProvider)) +_SECRET_PARAM_KEYS: Final = frozenset( + { + "api_key", + "_litellm_internal_model_credentials", + "azure_ad_token", + "azure_ad_token_provider", + "vertex_credentials", + "aws_access_key_id", + "aws_secret_access_key", + "aws_session_token", + "aws_web_identity_token", + } +) + def _supported_line_provider(value: str) -> _BatchLineProvider | None: if value in _SUPPORTED_LINE_PROVIDERS: @@ -292,6 +306,10 @@ async def _emit_line_event( model=child.model, custom_llm_provider=custom_llm_provider, ) + for secret_key in _SECRET_PARAM_KEYS: + child.litellm_params.pop( + secret_key, None + ) # mutable-ok: model_call_details holds this same dict # pyright: ignore[reportUnknownMemberType] # Logging.litellm_params is untyped upstream now: Final = datetime.now() # noqa: DTZ005 # naive to match the logging pipeline start_time if result is None: diff --git a/litellm/litellm_core_utils/core_helpers.py b/litellm/litellm_core_utils/core_helpers.py index d7fbe9f7e09..088a4a41428 100644 --- a/litellm/litellm_core_utils/core_helpers.py +++ b/litellm/litellm_core_utils/core_helpers.py @@ -5,7 +5,7 @@ import logging import re from collections.abc import Collection, Iterable, Mapping from types import MappingProxyType -from typing import TYPE_CHECKING, Any, Final, Literal, Protocol +from typing import TYPE_CHECKING, Any, Final, Literal, Protocol, cast import httpx from pydantic import TypeAdapter, ValidationError @@ -765,3 +765,16 @@ def set_response_cost_in_hidden_params(response: _CarriesHiddenParams, cost: flo RESPONSE_COST_HEADER: cost, } hidden_params["additional_headers"] = merged # rebind-ok: the caller's record is the point + + +def is_batch_line_item_event(kwargs: object) -> bool: + if not isinstance(kwargs, Mapping): + return False + litellm_params: Final = cast(Mapping[str, object], kwargs).get( + "litellm_params" + ) # cast-ok: isinstance narrows only to unparameterized Mapping + if not isinstance(litellm_params, Mapping): + return False + return bool( + cast(Mapping[str, object], litellm_params).get("batch_parent_id") + ) # cast-ok: same narrowing limit as above diff --git a/litellm/llms/anthropic/chat/transformation.py b/litellm/llms/anthropic/chat/transformation.py index 7f9294c58df..14c28b0ce68 100644 --- a/litellm/llms/anthropic/chat/transformation.py +++ b/litellm/llms/anthropic/chat/transformation.py @@ -2797,7 +2797,9 @@ def anthropic_message_to_model_response(result: Mapping[str, object], speed: str result_id: Final = result.get("id") return AnthropicConfig().transform_parsed_response( completion_response=pydantic_result.model_dump(), - raw_response=httpx.Response(status_code=200, headers={}), + raw_response=httpx.Response( + status_code=200, headers={} + ), # mutable-ok: httpx.Response wants a plain dict of headers model_response=ModelResponse(id=result_id if isinstance(result_id, str) and result_id else None), json_mode=None, speed=speed, diff --git a/litellm/llms/bedrock/batches/transformation.py b/litellm/llms/bedrock/batches/transformation.py index f8d71bab2c7..c61cd0dfeef 100644 --- a/litellm/llms/bedrock/batches/transformation.py +++ b/litellm/llms/bedrock/batches/transformation.py @@ -106,7 +106,9 @@ def bedrock_batch_line_to_response( embedding: Final = model_output.get("embedding") return EmbeddingResponse( model=model, - data=[{"object": "embedding", "index": 0, "embedding": embedding if isinstance(embedding, list) else []}], + data=[ + {"object": "embedding", "index": 0, "embedding": embedding if isinstance(embedding, list) else []} + ], # mutable-ok: EmbeddingResponse takes a plain data list usage=titan_embedding_usage_from_batch_output(model_output), ) if "output" in model_output: @@ -114,14 +116,14 @@ def bedrock_batch_line_to_response( return AmazonConverseConfig()._transform_response( # pyright: ignore[reportPrivateUsage] # same reconstruction the converse chat path performs on the live response model=model, - response=Response(200, json=dict(model_output)), + response=Response(200, json=dict(model_output)), # mutable-ok: httpx.Response json= wants a plain dict model_response=ModelResponse(), stream=False, logging_obj=None, - optional_params={}, + optional_params={}, # mutable-ok: the converse transform signature takes a plain dict api_key=None, data="", - messages=[], + messages=[], # mutable-ok: the converse transform signature takes a plain list encoding=None, ) if "content" in model_output: diff --git a/litellm/proxy/hooks/dynamic_rate_limiter_v3.py b/litellm/proxy/hooks/dynamic_rate_limiter_v3.py index 0339cf4dfea..6777e0c88f3 100644 --- a/litellm/proxy/hooks/dynamic_rate_limiter_v3.py +++ b/litellm/proxy/hooks/dynamic_rate_limiter_v3.py @@ -14,6 +14,7 @@ from litellm import ModelResponse, Router from litellm._logging import verbose_proxy_logger from litellm.caching.caching import DualCache from litellm.integrations.custom_logger import CustomLogger +from litellm.litellm_core_utils.core_helpers import is_batch_line_item_event from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.common_utils.proxy_rate_limit_error import ( ProxyRateLimitError, @@ -693,6 +694,8 @@ class _PROXY_DynamicRateLimitHandlerV3(CustomLogger): - model_saturation_check: Model-wide token tracking - priority_model: Priority-specific token tracking """ + if is_batch_line_item_event(kwargs): + return from litellm.litellm_core_utils.core_helpers import ( _get_parent_otel_span_from_kwargs, ) diff --git a/litellm/proxy/hooks/max_budget_per_session_limiter.py b/litellm/proxy/hooks/max_budget_per_session_limiter.py index e07b96e5773..a1f60fc4d87 100644 --- a/litellm/proxy/hooks/max_budget_per_session_limiter.py +++ b/litellm/proxy/hooks/max_budget_per_session_limiter.py @@ -23,6 +23,7 @@ from litellm._logging import verbose_proxy_logger from litellm.caching.redis_cache import log_redis_failure from litellm.exceptions import RateLimitType from litellm.integrations.custom_logger import CustomLogger +from litellm.litellm_core_utils.core_helpers import is_batch_line_item_event from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError from litellm.proxy.hooks.rate_limiter_utils import resolve_llm_provider_for_rate_limit @@ -131,6 +132,8 @@ class _PROXY_MaxBudgetPerSessionHandler(CustomLogger): """ After a successful LLM call, increment the session spend by the response cost. """ + if is_batch_line_item_event(kwargs): # pyright: ignore[reportUnknownArgumentType] # hook kwargs arrive untyped from the logging dispatcher + return try: litellm_params: Final = kwargs.get("litellm_params") or {} metadata: Final = litellm_params.get("metadata") or {} diff --git a/litellm/proxy/hooks/model_max_budget_limiter.py b/litellm/proxy/hooks/model_max_budget_limiter.py index cfa54ae01a2..de3d72ddb5b 100644 --- a/litellm/proxy/hooks/model_max_budget_limiter.py +++ b/litellm/proxy/hooks/model_max_budget_limiter.py @@ -11,6 +11,7 @@ import litellm from litellm._logging import verbose_proxy_logger from litellm.caching.caching import DualCache from litellm.integrations.custom_logger import Span +from litellm.litellm_core_utils.core_helpers import is_batch_line_item_event from litellm.litellm_core_utils.duration_parser import duration_in_seconds from litellm.llms.bedrock.common_utils import get_bedrock_base_model from litellm.proxy._types import Litellm_EntityType, UserAPIKeyAuth @@ -486,6 +487,8 @@ class _PROXY_VirtualKeyModelMaxBudgetLimiter(RouterBudgetLimiting): Example: key=sk-1234567890, model=gpt-4o, max_budget=100, time_period=1d """ + if is_batch_line_item_event(kwargs): + return verbose_proxy_logger.debug("in RouterBudgetLimiting.async_log_success_event") standard_logging_payload: Final[StandardLoggingPayload | None] = kwargs.get("standard_logging_object", None) if standard_logging_payload is None: diff --git a/litellm/proxy/hooks/parallel_request_limiter.py b/litellm/proxy/hooks/parallel_request_limiter.py index d41acadc4dd..a6dde585309 100644 --- a/litellm/proxy/hooks/parallel_request_limiter.py +++ b/litellm/proxy/hooks/parallel_request_limiter.py @@ -11,7 +11,10 @@ from litellm import DualCache, EmbeddingResponse, ModelResponse, TextCompletionR from litellm._logging import verbose_proxy_logger from litellm.exceptions import RateLimitType from litellm.integrations.custom_logger import CustomLogger -from litellm.litellm_core_utils.core_helpers import _get_parent_otel_span_from_kwargs +from litellm.litellm_core_utils.core_helpers import ( + _get_parent_otel_span_from_kwargs, + is_batch_line_item_event, +) from litellm.proxy._types import CommonProxyErrors, CurrentItemRateLimit, UserAPIKeyAuth from litellm.proxy.auth.auth_utils import ( get_key_model_rpm_limit, @@ -489,6 +492,8 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger): ) async def async_log_success_event(self, kwargs, response_obj: object, start_time, end_time): + if is_batch_line_item_event(kwargs): + return from litellm.proxy.common_utils.callback_utils import ( get_model_group_from_litellm_kwargs, ) diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index 33744906b13..d5b31624064 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -33,6 +33,7 @@ from litellm._logging import verbose_proxy_logger from litellm.caching.redis_cache import log_redis_failure from litellm.constants import DYNAMIC_RATE_LIMIT_ERROR_THRESHOLD_PER_MINUTE, INTERNAL_CALL_ORIGIN_METADATA_KEY from litellm.integrations.custom_logger import CustomLogger +from litellm.litellm_core_utils.core_helpers import is_batch_line_item_event from litellm.litellm_core_utils.prompt_templates.common_utils import ( get_str_from_messages, ) @@ -4557,6 +4558,8 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): """ Update TPM usage on successful API calls by incrementing counters using pipeline """ + if is_batch_line_item_event(kwargs): + return from litellm.litellm_core_utils.core_helpers import ( _get_parent_otel_span_from_kwargs, ) diff --git a/litellm/proxy/hooks/proxy_track_cost_callback.py b/litellm/proxy/hooks/proxy_track_cost_callback.py index 66df6578bc8..e8d508feb07 100644 --- a/litellm/proxy/hooks/proxy_track_cost_callback.py +++ b/litellm/proxy/hooks/proxy_track_cost_callback.py @@ -13,6 +13,7 @@ from litellm.litellm_core_utils.core_helpers import ( _get_parent_otel_span_from_kwargs, budget_reservation_from_metadata, get_litellm_metadata_from_kwargs, + is_batch_line_item_event, ) from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup from litellm.litellm_core_utils.llm_cost_calc.guardrail_cost import guardrail_information_cost @@ -106,10 +107,7 @@ class _ProxyDBLogger(CustomLogger): async def async_log_success_event( self, kwargs: ObjectMapping, response_obj: object, start_time: datetime, end_time: datetime ) -> None: - # Per-line batch events emitted under store_batch_line_items_in_callbacks - # never touch spend: the aggregate aretrieve_batch event already bills the batch. - litellm_params: Final = kwargs.get("litellm_params") - if isinstance(litellm_params, Mapping) and litellm_params.get("batch_parent_id"): + if is_batch_line_item_event(kwargs): return if self.spend_event_producer is None or not is_offloadable_success(response_obj): await self._PROXY_track_cost_callback(kwargs, response_obj, start_time, end_time) diff --git a/tests/test_litellm/batches/test_batch_line_item_logging.py b/tests/test_litellm/batches/test_batch_line_item_logging.py index 1f37969bcfb..eb8ae26222d 100644 --- a/tests/test_litellm/batches/test_batch_line_item_logging.py +++ b/tests/test_litellm/batches/test_batch_line_item_logging.py @@ -14,7 +14,7 @@ cost, or error propagation fails here. import json import uuid from datetime import datetime -from types import SimpleNamespace +from types import MappingProxyType, SimpleNamespace from typing import Final from unittest.mock import AsyncMock, patch @@ -560,6 +560,39 @@ async def test_line_items_edge_shapes_and_edge_cases(recorder): assert not any(file_id is None for file_id in file_ids_fetched) +@pytest.mark.asyncio +async def test_line_items_child_params_drop_parent_credentials(recorder): + litellm.store_batch_line_items_in_callbacks = True # test-quality-ok: the flag under test is a module global; fixture restores it + file_mock: Final = AsyncMock(side_effect=_file_content) + parent: Final = _parent_logging_with_params( + { + "api_key": "sk-parent-secret", + "_litellm_internal_model_credentials": MappingProxyType({"api_key": "sk-parent-secret"}), + "api_base": "https://api.openai.com/v1", + "metadata": {"model_info": {"id": "dep-1"}, "model_group": "gpt-4o"}, + } + ) + batch: Final = _batch() + with ( + patch("litellm.files.main.afile_content", file_mock), # test-quality-ok: afile_content is the provider boundary; no injection seam for managed file fetch + patch("litellm.cost_calculator.batch_cost_calculator", return_value=(0.01, 0.02)), # test-quality-ok: the pricing table boundary, same seam existing batch_utils tests patch + ): + await _log_completed_batch(parent, batch) + + line_events = [ + e + for e in [*recorder.success_events, *recorder.failure_events] + if _hidden(e).get("batch_custom_id") is not None + ] + assert len(line_events) == 2 + for event in line_events: + params = event["litellm_params"] + assert "api_key" not in params + assert "_litellm_internal_model_credentials" not in params + assert params["batch_parent_id"] == batch.id + assert params["metadata"]["model_group"] == "gpt-4o" + + @pytest.mark.asyncio async def test_line_items_anthropic_shapes(recorder): litellm.store_batch_line_items_in_callbacks = True # test-quality-ok: the flag under test is a module global; fixture restores it diff --git a/tests/test_litellm/litellm_core_utils/test_core_helpers.py b/tests/test_litellm/litellm_core_utils/test_core_helpers.py index 6eeea271127..368248fd956 100644 --- a/tests/test_litellm/litellm_core_utils/test_core_helpers.py +++ b/tests/test_litellm/litellm_core_utils/test_core_helpers.py @@ -11,6 +11,7 @@ from litellm.litellm_core_utils.core_helpers import ( drop_params_env_flag, drop_params_flag, get_or_create_metadata_bucket, + is_batch_line_item_event, map_finish_reason, normalize_drop_params, reconstruct_model_name, @@ -489,3 +490,13 @@ class TestIsExpectedClientError: category=RateLimitErrorCategory.VENDOR_RATE_LIMIT, ) assert is_expected_client_error(vendor_limit) is False + + +def test_is_batch_line_item_event(): + assert is_batch_line_item_event({"litellm_params": {"batch_parent_id": "batch_1"}}) is True + assert is_batch_line_item_event({"litellm_params": {"batch_parent_id": "batch_1", "metadata": {}}}) is True + assert is_batch_line_item_event({"litellm_params": {"metadata": {}}}) is False + assert is_batch_line_item_event({"litellm_params": {"batch_parent_id": None}}) is False + assert is_batch_line_item_event({}) is False + assert is_batch_line_item_event({"litellm_params": "not-a-mapping"}) is False + assert is_batch_line_item_event({"litellm_params": None}) is False diff --git a/tests/test_litellm/proxy/hooks/test_model_max_budget_limiter.py b/tests/test_litellm/proxy/hooks/test_model_max_budget_limiter.py index ffb60fb4651..b062efb3e81 100644 --- a/tests/test_litellm/proxy/hooks/test_model_max_budget_limiter.py +++ b/tests/test_litellm/proxy/hooks/test_model_max_budget_limiter.py @@ -196,3 +196,17 @@ async def test_a_batch_polled_within_every_budget_window_is_never_charged_again( await _poll(limiter, finished, BATCH_COST) assert _local_spend(limiter, KEY_SPEND_KEY) == pytest.approx(BATCH_COST) + + +@pytest.mark.asyncio +async def test_batch_line_item_events_do_not_charge_the_model_budget(): + """Line events carry batch_parent_id; the aggregate aretrieve_batch event is + the one that already bills the batch, so children must not double-charge.""" + limiter: Final = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=DualCache()) + + line_event = _event("acompletion", CHAT_COST) + line_event["litellm_params"]["batch_parent_id"] = "batch_first" + await limiter.async_log_success_event(line_event, response_obj=None, start_time=None, end_time=None) + + assert await _spend(limiter, KEY_SPEND_KEY) == 0.0 + assert await _spend(limiter, USER_SPEND_KEY) == 0.0 diff --git a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py index 8c0dcd3383c..904b5d67556 100644 --- a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py +++ b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py @@ -7150,3 +7150,30 @@ async def test_success_tpm_accounting_keeps_the_admission_target_after_an_alias_ charged: Final = {op["key"]: op["increment_value"] for op in ops} assert charged[admission_bucket] == 150 - stash.reserved_tokens assert not any(":target-b" in key for key in charged) + + +@pytest.mark.asyncio +async def test_async_log_success_event_skips_batch_line_item_events(): + """Per-line batch events already ran through the limiter as the aggregate + aretrieve_batch; children must not increment TPM or request counters.""" + local_cache = DualCache() + handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(local_cache)) + + await handler.async_log_success_event( + kwargs={ + "standard_logging_object": {"metadata": {"user_api_key_hash": hash_token("sk-line-item")}}, + "litellm_params": { + "batch_parent_id": "batch_1", + "metadata": {"user_api_key_hash": hash_token("sk-line-item"), "model_group": "gpt-3.5-turbo"}, + }, + "model": "gpt-3.5-turbo", + }, + response_obj=ModelResponse( + id="x", object="chat.completion", created=1, model="gpt-3.5-turbo", + usage=Usage(prompt_tokens=10, completion_tokens=5, total_tokens=15), choices=[], + ), + start_time=datetime.now(), + end_time=datetime.now(), + ) + + assert local_cache.in_memory_cache.cache_dict == {}