mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
fix(batch): floor input_audio rows at the size-based estimate so audio payloads cannot evade TPM limits
Before input_audio blocks were countable, a batch row carrying one RAISED inside token_counter and the batch rate limiter fell back to its conservative size-based estimate (raw bytes / 4), which covers the base64 payload. Making the row countable replaced that with the counter's audio estimate (decoded bytes / AUDIO_BYTES_PER_TOKEN, a deliberately low assumed bitrate), so a row carrying a large audio payload reserved ~500x less and a crafted batch could slide under the TPM limit -- the same loophole class raised and fixed for file blocks on #33659. Restore the conservatism at the rate-limiter call site: when a row's messages carry an input_audio block, take max(counted, size-based estimate). Plain rows keep the measured count. Live (non-batch) requests are unaffected -- parallel_request_limiter_v3 strips audio blocks before counting and applies its own identical estimate. Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
parent
7b1c639bd7
commit
88f36b5b81
3 changed files with 130 additions and 0 deletions
|
|
@ -887,6 +887,28 @@ def _count_input_audio_content_block(c: Mapping[str, object]) -> int:
|
|||
return DEFAULT_AUDIO_TOKEN_ESTIMATE
|
||||
|
||||
|
||||
def messages_contain_input_audio_content_blocks(messages: object) -> bool:
|
||||
"""
|
||||
True when any message carries an OpenAI ``input_audio`` content block.
|
||||
|
||||
Callers that use ``token_counter`` for reservations or rate limits must
|
||||
check this first: an audio block's contribution is a size-derived
|
||||
ESTIMATE at a deliberately low assumed bitrate (see
|
||||
``_count_input_audio_content_block``), not a measurement, so the count
|
||||
for such a message must not be trusted as an upper bound.
|
||||
"""
|
||||
if not isinstance(messages, list):
|
||||
return False
|
||||
for message in messages:
|
||||
content = message.get("content") if isinstance(message, dict) else None
|
||||
if not isinstance(content, list):
|
||||
continue
|
||||
for content_item in content:
|
||||
if isinstance(content_item, dict) and content_item.get("type") == "input_audio":
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
def _count_content_list(
|
||||
count_function: TokenCounterFunction,
|
||||
content_list: str
|
||||
|
|
|
|||
|
|
@ -37,6 +37,7 @@ from litellm.batches.batch_utils import (
|
|||
from litellm.constants import BATCH_TPD_DESCRIPTOR_SUFFIX, BATCH_TPD_WINDOW_SECONDS
|
||||
from litellm.exceptions import RateLimitErrorCategory
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.litellm_core_utils.token_counter import messages_contain_input_audio_content_blocks
|
||||
from litellm.proxy._types import (
|
||||
ProxyErrorTypes,
|
||||
ProxyException,
|
||||
|
|
@ -971,6 +972,17 @@ class _PROXY_BatchRateLimiter(CustomLogger):
|
|||
|
||||
try:
|
||||
entry_total_tokens = _count_entry_tokens(entry)
|
||||
# An `input_audio` block's contribution is a size-derived
|
||||
# estimate at a deliberately low assumed bitrate (see
|
||||
# token_counter._count_input_audio_content_block), far
|
||||
# below this fallback's raw-bytes estimate. Before such
|
||||
# blocks were countable (#38459) an audio row RAISED
|
||||
# inside token_counter and fell back to the size-based
|
||||
# estimate; floor it there again so a row carrying a
|
||||
# large base64 audio payload cannot slide the batch
|
||||
# under the TPM limit.
|
||||
if messages_contain_input_audio_content_blocks((entry.get("body") or {}).get("messages")):
|
||||
entry_total_tokens = max(entry_total_tokens, _estimate_batch_entry_tokens(raw_line))
|
||||
except Exception:
|
||||
entry_total_tokens = _estimate_batch_entry_tokens(raw_line)
|
||||
total_tokens += entry_total_tokens
|
||||
|
|
|
|||
|
|
@ -2286,3 +2286,99 @@ async def test_disable_flag_still_skips_batch_processing_with_enqueued_limits():
|
|||
|
||||
assert result is data
|
||||
afile_content_mock.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_input_audio_row_reserves_size_based_floor():
|
||||
"""A chat row carrying an OpenAI `input_audio` content block must reserve
|
||||
at least the size-based estimate (serialized bytes / 4).
|
||||
|
||||
token_counter's audio contribution is a size-derived estimate at a
|
||||
deliberately low assumed bitrate (decoded bytes / AUDIO_BYTES_PER_TOKEN),
|
||||
far below the raw-bytes fallback. Before audio blocks were countable
|
||||
(#38459) such a row RAISED inside token_counter and fell back to the
|
||||
size-based estimate; the floor restores exactly that conservatism so a
|
||||
large base64 audio payload cannot slide the batch under the TPM limit.
|
||||
"""
|
||||
import json as _json
|
||||
|
||||
from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter
|
||||
|
||||
prl = MagicMock()
|
||||
prl.no_max_tokens_output_floor.return_value = 0
|
||||
rate_limiter = _PROXY_BatchRateLimiter(
|
||||
internal_usage_cache=MagicMock(),
|
||||
parallel_request_limiter=prl,
|
||||
)
|
||||
|
||||
blob = "A" * 400_000
|
||||
audio_row_bytes = _json.dumps(
|
||||
{
|
||||
"custom_id": "row-1",
|
||||
"method": "POST",
|
||||
"url": "/v1/chat/completions",
|
||||
"body": {
|
||||
"model": "gpt-4o-audio-preview",
|
||||
"messages": [
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"type": "text", "text": "What does the audio say?"},
|
||||
{"type": "input_audio", "input_audio": {"data": blob, "format": "wav"}},
|
||||
],
|
||||
}
|
||||
],
|
||||
},
|
||||
}
|
||||
).encode("utf-8")
|
||||
fake_content = MagicMock()
|
||||
fake_content.content = audio_row_bytes
|
||||
|
||||
with patch( # test-quality-ok: mirrors the file's established harness — the download, not an HTTP boundary
|
||||
"litellm.afile_content",
|
||||
new=AsyncMock(return_value=fake_content),
|
||||
):
|
||||
usage = await rate_limiter.count_input_file_usage(
|
||||
file_id="file-not-managed",
|
||||
custom_llm_provider="openai",
|
||||
user_api_key_dict=None,
|
||||
)
|
||||
|
||||
size_based_floor = len(audio_row_bytes) // 4
|
||||
assert usage.request_count == 1
|
||||
assert usage.total_tokens >= size_based_floor, (
|
||||
f"audio-block row must reserve at least the size-based estimate "
|
||||
f"({size_based_floor} tokens for {len(audio_row_bytes)} bytes), got "
|
||||
f"{usage.total_tokens} — a large audio payload would evade TPM limits"
|
||||
)
|
||||
|
||||
# Control: a plain-text row must NOT be floored at its serialized size —
|
||||
# measured text rows keep the (smaller) real token count.
|
||||
text_row_bytes = _json.dumps(
|
||||
{
|
||||
"custom_id": "row-1",
|
||||
"method": "POST",
|
||||
"url": "/v1/chat/completions",
|
||||
"body": {
|
||||
"model": "gpt-4o-audio-preview",
|
||||
"messages": [{"role": "user", "content": "What does the audio say?"}],
|
||||
},
|
||||
}
|
||||
).encode("utf-8")
|
||||
fake_text_content = MagicMock()
|
||||
fake_text_content.content = text_row_bytes
|
||||
|
||||
with patch( # test-quality-ok: mirrors the file's established harness — the download, not an HTTP boundary
|
||||
"litellm.afile_content",
|
||||
new=AsyncMock(return_value=fake_text_content),
|
||||
):
|
||||
text_usage = await rate_limiter.count_input_file_usage(
|
||||
file_id="file-not-managed",
|
||||
custom_llm_provider="openai",
|
||||
user_api_key_dict=None,
|
||||
)
|
||||
|
||||
assert text_usage.request_count == 1
|
||||
assert text_usage.total_tokens < len(text_row_bytes) // 4, (
|
||||
"plain-text rows must keep the measured token count, not the size-based floor"
|
||||
)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue