mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
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:
parent
e735317782
commit
91a72101a1
3 changed files with 116 additions and 55 deletions
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue