fix(batches): keep line item events out of proxy limiters and strip credentials from child params

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
yucheng 2026-09-23 15:38:04 +00:00
parent 271b6c61ca
commit c0b69aedb9
14 changed files with 147 additions and 12 deletions

View file

@ -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:

View file

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

View file

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

View file

@ -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:

View file

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

View file

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

View file

@ -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:

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

@ -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 == {}