From 91a72101a123d64156db7d955f12508a1ceff0ba Mon Sep 17 00:00:00 2001 From: mubashir1osmani Date: Mon, 22 Jun 2026 18:41:15 -0700 Subject: [PATCH] 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. --- litellm/batches/batch_utils.py | 32 +++++++- litellm/proxy/hooks/batch_rate_limiter.py | 76 ++++++++----------- .../proxy/hooks/test_batch_file_validation.py | 63 +++++++++++++-- 3 files changed, 116 insertions(+), 55 deletions(-) diff --git a/litellm/batches/batch_utils.py b/litellm/batches/batch_utils.py index 0ce2d02f39f..aeec58f1dfc 100644 --- a/litellm/batches/batch_utils.py +++ b/litellm/batches/batch_utils.py @@ -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( diff --git a/litellm/proxy/hooks/batch_rate_limiter.py b/litellm/proxy/hooks/batch_rate_limiter.py index 69d27fd38cf..91c604e6204 100644 --- a/litellm/proxy/hooks/batch_rate_limiter.py +++ b/litellm/proxy/hooks/batch_rate_limiter.py @@ -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, diff --git a/tests/test_litellm/proxy/hooks/test_batch_file_validation.py b/tests/test_litellm/proxy/hooks/test_batch_file_validation.py index e986f1582f8..1d4d39ec140 100644 --- a/tests/test_litellm/proxy/hooks/test_batch_file_validation.py +++ b/tests/test_litellm/proxy/hooks/test_batch_file_validation.py @@ -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