fix(batches): resilient per-row token accounting; no hard-block on count failure

The batch input-file pass iterated a generator whose json.loads raised on a
malformed line; the outer except caught it and stopped the loop, so any body.model
on rows after a bad line was never collected and the model allowlist check ran
against a partial set. It also hard-blocked the batch with a 400 whenever token
counting raised, a backwards-incompatible change from the prior swallow-and-proceed
behavior that breaks legitimate rows the token counter cannot measure (e.g. some
multimodal content).

Iterate the JSONL line-by-line and account each row independently. A malformed line
is skipped (its request cannot run upstream anyway) and a row the counter cannot
measure falls back to a conservative size-based estimate. The loop never aborts, so
the allowlist check always sees every parseable model, and the token total is never
zeroed, so a crafted uncountable row still cannot evade the TPM limit, without
hard-rejecting a legitimate batch.
This commit is contained in:
mubashir1osmani 2026-06-22 18:41:15 -07:00
parent e735317782
commit 91a72101a1
3 changed files with 116 additions and 55 deletions

View file

@ -314,10 +314,11 @@ def _get_file_content_as_dictionary(file_content: bytes) -> List[dict]:
raise e
def _iter_batch_input_entries(file_content: bytes) -> Iterator[dict]:
def _iter_batch_input_lines(file_content: bytes) -> Iterator[bytes]:
"""
Yield batch input JSONL entries one at a time without materializing the whole
file as a list, so peak memory stays bounded when counting a large batch file.
Yield non-empty JSONL lines (unparsed) one at a time, so a caller can parse
each row in its own try/except and a single malformed line cannot abort the
whole pass. Peak memory stays bounded for large batch files.
"""
start, length, newline = 0, len(file_content), ord("\n")
while start < length:
@ -328,7 +329,30 @@ def _iter_batch_input_entries(file_content: bytes) -> Iterator[dict]:
chunk, start = file_content[start:idx], idx + 1
line = chunk.strip()
if line:
yield json.loads(line)
yield line
def _iter_batch_input_entries(file_content: bytes) -> Iterator[dict]:
"""
Yield parsed batch input JSONL entries one at a time without materializing the
whole file as a list, so peak memory stays bounded. Raises on a malformed line;
callers that must survive bad rows should iterate ``_iter_batch_input_lines``
and parse per-row instead.
"""
for line in _iter_batch_input_lines(file_content):
yield json.loads(line)
# A batch request's input tokens scale roughly with its serialized size, so this
# is a conservative per-row fallback when the token counter cannot measure a row.
_BATCH_TOKEN_ESTIMATE_BYTES_PER_TOKEN = 4
def _estimate_batch_entry_tokens(raw_line: bytes) -> int:
"""Conservative token estimate for a batch row the token counter cannot measure
(or that cannot be parsed). Keeps the batch token total non-zero so a crafted
row cannot evade the TPM limit, without hard-rejecting a legitimate batch."""
return max(1, len(raw_line) // _BATCH_TOKEN_ESTIMATE_BYTES_PER_TOKEN)
def _count_entry_tokens(

View file

@ -33,12 +33,15 @@ from typing import (
from fastapi import HTTPException
from pydantic import BaseModel
import json
import litellm
from litellm._logging import verbose_proxy_logger
from litellm.batches.batch_utils import (
_count_entry_tokens,
_estimate_batch_entry_tokens,
_extract_file_access_credentials,
_iter_batch_input_entries,
_iter_batch_input_lines,
)
from litellm.exceptions import RateLimitErrorCategory
from litellm.integrations.custom_logger import CustomLogger
@ -571,37 +574,37 @@ class _PROXY_BatchRateLimiter(CustomLogger):
f"got {type(file_content_bytes)}"
)
# Single streaming pass over the JSONL entries: accumulate request
# count, distinct models, and token total without ever holding all
# entries in a list, so peak memory does not scale with file size.
#
# Token counting and line parsing are best-effort accounting and must
# never gate the security check below: a restricted caller could craft
# a row (e.g. provider-valid `input_audio`/`file` content the token
# counter rejects, or a malformed line) that makes counting/parsing
# raise. If that aborted the loop, async_pre_call_hook would swallow
# the non-HTTP exception and submit the batch unchecked. So failures
# are captured, model collection always completes, and the failure is
# re-raised as a blocking HTTPException only after the allowlist check.
# Single streaming pass over the JSONL lines, accounting each row
# independently. One bad row can never abort the pass: a malformed
# line is skipped (its request can't run upstream anyway) and a row
# the token counter can't measure falls back to a conservative
# size-based estimate. This guarantees two things a restricted caller
# must not be able to break by crafting a row that raises:
# 1. The allowlist check below always sees every parseable
# ``body.model`` (the loop never stops early), so models can't be
# smuggled in after a bad row.
# 2. The token total is never silently zeroed, so the TPM limit
# can't be evaded by sending uncountable rows.
# Counting stays best-effort, so a legitimate (e.g. multimodal) row
# the counter can't measure is estimated, not hard-rejected.
models: set = set()
total_tokens = 0
request_count = 0
processing_error: Optional[Exception] = None
try:
for entry in _iter_batch_input_entries(file_content_bytes):
request_count += 1
if isinstance(entry, dict):
model = (entry.get("body") or {}).get("model")
if model:
models.add(model)
try:
total_tokens += _count_entry_tokens(entry)
except Exception as e:
if processing_error is None:
processing_error = e
except Exception as e:
if processing_error is None:
processing_error = e
for raw_line in _iter_batch_input_lines(file_content_bytes):
request_count += 1
try:
entry = json.loads(raw_line)
except Exception:
total_tokens += _estimate_batch_entry_tokens(raw_line)
continue
if isinstance(entry, dict):
model = (entry.get("body") or {}).get("model")
if model:
models.add(model)
try:
total_tokens += _count_entry_tokens(entry)
except Exception:
total_tokens += _estimate_batch_entry_tokens(raw_line)
# Validate every model named in the batch JSONL against the
# caller's per-key model allowlist. Without this, a caller
@ -615,21 +618,6 @@ class _PROXY_BatchRateLimiter(CustomLogger):
target_model_names=target_model_names or None,
)
# A counting/parsing failure blocks the batch (HTTPException is
# re-raised by async_pre_call_hook, unlike other exceptions), so it
# can neither bypass the allowlist nor evade token rate limits by
# undercounting. Runs only after the allowlist check above.
if processing_error is not None:
raise HTTPException(
status_code=400,
detail={
"error": (
"Could not account for all tokens in the batch input "
f"file: {processing_error}"
)
},
)
return BatchFileUsage(
total_tokens=total_tokens,
request_count=request_count,

View file

@ -1546,11 +1546,12 @@ async def test_count_input_file_usage_enforces_models_when_token_counting_fails(
@pytest.mark.asyncio
async def test_count_input_file_usage_blocks_when_counting_fails_for_allowed_model():
"""Even for an allowed model, a token-counting failure must block with an
HTTPException rather than return. Returning would let async_pre_call_hook
treat it as success, so a caller could evade token rate limits by sending
rows the counter cannot measure."""
async def test_count_input_file_usage_estimates_tokens_when_counting_fails_for_allowed_model():
"""A token-counting failure for an allowed model must not hard-block the batch
(the pre-streaming behavior let such batches through), but it also must not
zero the token total, which would let a caller evade the TPM limit by sending
rows the counter cannot measure. The row falls back to a conservative
size-based estimate so the batch proceeds with a non-zero count."""
from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter
rate_limiter = _PROXY_BatchRateLimiter(
@ -1576,6 +1577,55 @@ async def test_count_input_file_usage_blocks_when_counting_fails_for_allowed_mod
patch("litellm.proxy.hooks.batch_rate_limiter._count_entry_tokens", new=_boom),
patch("litellm.proxy.auth.auth_checks.can_key_call_model", new=allow),
patch("litellm.proxy.proxy_server.llm_router", MagicMock(model_list=[])),
):
usage = await rate_limiter.count_input_file_usage(
file_id="file-not-managed",
custom_llm_provider="openai",
user_api_key_dict=user,
)
allow.assert_awaited()
assert usage.request_count == 1
# Estimated, not zeroed: a crafted uncountable row can't evade the TPM limit.
assert usage.total_tokens > 0
@pytest.mark.asyncio
async def test_count_input_file_usage_collects_models_after_malformed_line():
"""A malformed JSONL line must not abort model collection. A restricted model
named on a row AFTER a malformed line must still be collected and denied by the
allowlist check, otherwise a caller could hide a restricted model behind a bad
row."""
from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter
rate_limiter = _PROXY_BatchRateLimiter(
internal_usage_cache=MagicMock(),
parallel_request_limiter=MagicMock(),
)
fake_content = MagicMock()
fake_content.content = (
_one_row_batch_bytes("only-allowed")
+ b"{ this is not valid json\n"
+ _one_row_batch_bytes("restricted-model")
)
user = UserAPIKeyAuth(
api_key="sk-x",
user_id="bob",
models=["only-allowed"],
user_role=LitellmUserRoles.INTERNAL_USER.value,
)
async def _deny_restricted(model, **kwargs):
if model == "restricted-model":
raise Exception("model not in allowlist")
return True
deny = AsyncMock(side_effect=_deny_restricted)
with (
patch("litellm.afile_content", new=AsyncMock(return_value=fake_content)),
patch("litellm.proxy.auth.auth_checks.can_key_call_model", new=deny),
patch("litellm.proxy.proxy_server.llm_router", MagicMock(model_list=[])),
):
with pytest.raises(HTTPException) as exc:
await rate_limiter.count_input_file_usage(
@ -1584,5 +1634,4 @@ async def test_count_input_file_usage_blocks_when_counting_fails_for_allowed_mod
user_api_key_dict=user,
)
allow.assert_awaited()
assert exc.value.status_code == 400
assert exc.value.status_code == 403