mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
fix: address code review feedback on max_tokens capping logic
- Simplify two-try-block approach into a single try block: if no explicit provider in extra_kwargs, infer it via get_llm_provider() in the same path rather than a separate fallback block. This eliminates the misleading _capped/_lookup_succeeded flag entirely. - Use explicit `is not None` check instead of truthy check on max_output_tokens, consistent with get_modified_max_tokens pattern in token_counter.py and correct for hypothetical zero-limit models. - Add unit tests covering all capping scenarios: - max_tokens capped when it exceeds the model limit - max_tokens unchanged when within the limit - max_tokens unchanged when equal to the limit - provider inferred from model string when absent from extra_kwargs - resilient when get_model_info raises - no cap when max_output_tokens key is missing from model_info - no cap when max_output_tokens is explicitly None - explicit provider is used without calling get_llm_provider
This commit is contained in:
parent
68c891cf2f
commit
0b65e0a2b2
2 changed files with 104 additions and 18 deletions
|
|
@ -129,32 +129,21 @@ class LiteLLMMessagesToCompletionTransformationHandler:
|
|||
# By this point get_llm_provider() has already been called in the outer
|
||||
# anthropic_messages_handler, so `model` is the stripped model name
|
||||
# (e.g. "converse/us.amazon.nova-pro-v1:0") and `custom_llm_provider`
|
||||
# is passed via extra_kwargs (e.g. "bedrock"). We try the lookup with
|
||||
# the provider first, then fall back to inferring it from the model
|
||||
# string for cases where extra_kwargs doesn't carry the provider.
|
||||
# is passed via extra_kwargs (e.g. "bedrock"). If no explicit provider
|
||||
# is available we infer it from the model string.
|
||||
_custom_llm_provider = (extra_kwargs or {}).get("custom_llm_provider")
|
||||
_capped = False
|
||||
try:
|
||||
_lookup_provider = _custom_llm_provider
|
||||
if _lookup_provider is None:
|
||||
_, _lookup_provider, _, _ = litellm.utils.get_llm_provider(model)
|
||||
model_info = litellm.get_model_info(
|
||||
model=model, custom_llm_provider=_custom_llm_provider
|
||||
model=model, custom_llm_provider=_lookup_provider
|
||||
)
|
||||
model_max_output = model_info.get("max_output_tokens")
|
||||
if model_max_output and max_tokens > model_max_output:
|
||||
if model_max_output is not None and max_tokens > model_max_output:
|
||||
max_tokens = model_max_output
|
||||
_capped = True
|
||||
except Exception:
|
||||
pass
|
||||
if not _capped:
|
||||
try:
|
||||
_, inferred_provider, _, _ = litellm.utils.get_llm_provider(model)
|
||||
model_info = litellm.get_model_info(
|
||||
model=model, custom_llm_provider=inferred_provider
|
||||
)
|
||||
model_max_output = model_info.get("max_output_tokens")
|
||||
if model_max_output and max_tokens > model_max_output:
|
||||
max_tokens = model_max_output
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
request_data = {
|
||||
"model": model,
|
||||
|
|
|
|||
|
|
@ -0,0 +1,97 @@
|
|||
"""
|
||||
Unit tests for LiteLLMMessagesToCompletionTransformationHandler._prepare_completion_kwargs,
|
||||
specifically the max_tokens capping logic added to prevent HTTP 400 errors from providers
|
||||
with strict output token limits (e.g. Amazon Nova Pro: 10,000 tokens).
|
||||
"""
|
||||
import os
|
||||
import sys
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
sys.path.insert(0, os.path.abspath("../../../../.."))
|
||||
|
||||
from litellm.llms.anthropic.experimental_pass_through.adapters.handler import (
|
||||
LiteLLMMessagesToCompletionTransformationHandler,
|
||||
)
|
||||
|
||||
|
||||
MESSAGES = [{"role": "user", "content": "hello"}]
|
||||
MODEL = "converse/us.amazon.nova-pro-v1:0"
|
||||
PROVIDER = "bedrock"
|
||||
MODEL_MAX_OUTPUT = 10_000
|
||||
|
||||
|
||||
def _call(max_tokens, extra_kwargs=None):
|
||||
"""Helper: call _prepare_completion_kwargs and return the resolved max_tokens."""
|
||||
kwargs, _ = LiteLLMMessagesToCompletionTransformationHandler._prepare_completion_kwargs(
|
||||
max_tokens=max_tokens,
|
||||
messages=MESSAGES,
|
||||
model=MODEL,
|
||||
extra_kwargs=extra_kwargs or {"custom_llm_provider": PROVIDER},
|
||||
)
|
||||
return kwargs["max_tokens"]
|
||||
|
||||
|
||||
class TestMaxTokensCapping:
|
||||
def test_caps_when_exceeds_limit(self):
|
||||
"""max_tokens above the model limit is silently capped to max_output_tokens."""
|
||||
with patch("litellm.get_model_info") as mock_info:
|
||||
mock_info.return_value = {"max_output_tokens": MODEL_MAX_OUTPUT}
|
||||
result = _call(max_tokens=16_000)
|
||||
assert result == MODEL_MAX_OUTPUT
|
||||
|
||||
def test_unchanged_when_within_limit(self):
|
||||
"""max_tokens at or below the model limit is left unchanged."""
|
||||
with patch("litellm.get_model_info") as mock_info:
|
||||
mock_info.return_value = {"max_output_tokens": MODEL_MAX_OUTPUT}
|
||||
result = _call(max_tokens=8_000)
|
||||
assert result == 8_000
|
||||
|
||||
def test_unchanged_when_equal_to_limit(self):
|
||||
"""max_tokens exactly equal to the model limit is left unchanged."""
|
||||
with patch("litellm.get_model_info") as mock_info:
|
||||
mock_info.return_value = {"max_output_tokens": MODEL_MAX_OUTPUT}
|
||||
result = _call(max_tokens=MODEL_MAX_OUTPUT)
|
||||
assert result == MODEL_MAX_OUTPUT
|
||||
|
||||
def test_fallback_infers_provider_when_not_in_extra_kwargs(self):
|
||||
"""When custom_llm_provider is absent from extra_kwargs, max_tokens is still
|
||||
capped correctly by inferring the provider from the model string."""
|
||||
with patch("litellm.utils.get_llm_provider") as mock_provider, \
|
||||
patch("litellm.get_model_info") as mock_info:
|
||||
mock_provider.return_value = (MODEL, PROVIDER, None, None)
|
||||
mock_info.return_value = {"max_output_tokens": MODEL_MAX_OUTPUT}
|
||||
|
||||
result = _call(max_tokens=16_000, extra_kwargs={})
|
||||
|
||||
assert result == MODEL_MAX_OUTPUT
|
||||
|
||||
def test_resilient_when_get_model_info_raises(self):
|
||||
"""If get_model_info raises, max_tokens is passed through unchanged."""
|
||||
with patch("litellm.get_model_info", side_effect=Exception("model not found")), \
|
||||
patch("litellm.utils.get_llm_provider", side_effect=Exception("no provider")):
|
||||
result = _call(max_tokens=16_000)
|
||||
assert result == 16_000
|
||||
|
||||
def test_no_cap_when_max_output_tokens_missing(self):
|
||||
"""If model_info has no max_output_tokens key, max_tokens is unchanged."""
|
||||
with patch("litellm.get_model_info") as mock_info:
|
||||
mock_info.return_value = {}
|
||||
result = _call(max_tokens=16_000)
|
||||
assert result == 16_000
|
||||
|
||||
def test_no_cap_when_max_output_tokens_is_none(self):
|
||||
"""Explicit None max_output_tokens does not trigger capping."""
|
||||
with patch("litellm.get_model_info") as mock_info:
|
||||
mock_info.return_value = {"max_output_tokens": None}
|
||||
result = _call(max_tokens=16_000)
|
||||
assert result == 16_000
|
||||
|
||||
def test_explicit_provider_used_before_inference(self):
|
||||
"""When custom_llm_provider is present, get_llm_provider is not called."""
|
||||
with patch("litellm.utils.get_llm_provider") as mock_provider, \
|
||||
patch("litellm.get_model_info") as mock_info:
|
||||
mock_info.return_value = {"max_output_tokens": MODEL_MAX_OUTPUT}
|
||||
_call(max_tokens=16_000, extra_kwargs={"custom_llm_provider": PROVIDER})
|
||||
mock_provider.assert_not_called()
|
||||
Loading…
Add table
Reference in a new issue