mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
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:
parent
271b6c61ca
commit
c0b69aedb9
14 changed files with 147 additions and 12 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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 {}
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 == {}
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue