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:
Mihidum Hettiyahandi 2026-08-27 14:56:07 +10:00
parent 7b1c639bd7
commit 88f36b5b81
3 changed files with 130 additions and 0 deletions

View file

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

View file

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

View file

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