Exclude none fields on /chat/completion - fixes n8n bug + Allow calling /v1/models when end user over budget (#13320)

* fix(proxy_server.py): exclude none fields before returning

Fixes https://github.com/BerriAI/litellm/issues/13055

* test: add unit tests

* feat(auth_checks.py): allow info routes to work when end user over budget

Fixes https://github.com/BerriAI/litellm/issues/13286
This commit is contained in:
Krish Dholakia 2025-08-05 21:39:46 -07:00 • committed by GitHub
parent 92c525ddfe
commit 0da25fadc0
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
7 changed files with 212 additions and 32 deletions

View file

@ -380,6 +380,8 @@ class LiteLLMRoutes(enum.Enum):
"/health",
"/key/list",
"/user/filter/ui",
"/models",
"/v1/models",
]
# NOTE: ROUTES ONLY FOR MASTER KEY - only the Master Key should be able to Reset Spend

View file

@ -437,6 +437,7 @@ async def get_end_user_object(
end_user_id: Optional[str],
prisma_client: Optional[PrismaClient],
user_api_key_cache: DualCache,
route: str,
parent_otel_span: Optional[Span] = None,
proxy_logging_obj: Optional[ProxyLogging] = None,
) -> Optional[LiteLLM_EndUserTable]:
@ -453,6 +454,8 @@ async def get_end_user_object(
_key = "end_user_id:{}".format(end_user_id)
def check_in_budget(end_user_obj: LiteLLM_EndUserTable):
if route in LiteLLMRoutes.info_routes.value: # allow calling info routes
return
if end_user_obj.litellm_budget_table is None:
return
end_user_budget = end_user_obj.litellm_budget_table.max_budget

View file

@ -83,7 +83,7 @@ class JWTHandler:
self.user_api_key_cache = user_api_key_cache
self.litellm_jwtauth = litellm_jwtauth
self.leeway = leeway
@staticmethod
def is_jwt(token: str):
parts = token.split(".")
@ -844,6 +844,7 @@ class JWTAuthManager:
user_api_key_cache: DualCache,
parent_otel_span: Optional[Span],
proxy_logging_obj: ProxyLogging,
route: str,
) -> Tuple[
Optional[LiteLLM_UserTable],
Optional[LiteLLM_OrganizationTable],
@ -892,6 +893,7 @@ class JWTAuthManager:
user_api_key_cache=user_api_key_cache,
parent_otel_span=parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
route=route,
)
if end_user_id
else None
@ -1133,6 +1135,7 @@ class JWTAuthManager:
user_api_key_cache=user_api_key_cache,
parent_otel_span=parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
route=route,
)
await JWTAuthManager.sync_user_role_and_teams(

View file

@ -615,6 +615,7 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
user_api_key_cache=user_api_key_cache,
parent_otel_span=parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
route=route,
)
if _end_user_object is not None:
end_user_params["allowed_model_region"] = (

View file

@ -679,7 +679,6 @@ app = FastAPI(
)
### CUSTOM API DOCS [ENTERPRISE FEATURE] ###
# Custom OpenAPI schema generator to include only selected routes
from fastapi.routing import APIWebSocketRoute
@ -1593,7 +1592,9 @@ class ProxyConfig:
litellm.cache = Cache(**cache_params)
if litellm.cache is not None and isinstance(litellm.cache.cache, (RedisCache, RedisClusterCache)):
if litellm.cache is not None and isinstance(
litellm.cache.cache, (RedisCache, RedisClusterCache)
):
## INIT PROXY REDIS USAGE CLIENT ##
redis_usage_cache = litellm.cache.cache
@ -3768,9 +3769,12 @@ async def model_list(
Defaults to "general" when include_metadata=true
"""
global llm_model_list, general_settings, llm_router, prisma_client, user_api_key_cache, proxy_logging_obj
from litellm.proxy.utils import get_available_models_for_user, create_model_info_response
from litellm.proxy.utils import (
create_model_info_response,
get_available_models_for_user,
)
# Get available models for the user
all_models = await get_available_models_for_user(
user_api_key_dict=user_api_key_dict,
@ -3805,10 +3809,14 @@ async def model_list(
@router.get(
"/v1/models/{model_id}", dependencies=[Depends(user_api_key_auth)], tags=["model management"]
"/v1/models/{model_id}",
dependencies=[Depends(user_api_key_auth)],
tags=["model management"],
)
@router.get(
"/models/{model_id}", dependencies=[Depends(user_api_key_auth)], tags=["model management"]
"/models/{model_id}",
dependencies=[Depends(user_api_key_auth)],
tags=["model management"],
)
async def model_info(
model_id: str,
@ -3816,17 +3824,21 @@ async def model_info(
):
"""
Retrieve information about a specific model accessible to your API key.
Returns model details only if the model is available to your API key/team.
Returns 404 if the model doesn't exist or is not accessible.
Follows OpenAI API specification for individual model retrieval.
https://platform.openai.com/docs/api-reference/models/retrieve
"""
global llm_model_list, general_settings, llm_router, prisma_client, user_api_key_cache, proxy_logging_obj
from litellm.proxy.utils import get_available_models_for_user, validate_model_access, create_model_info_response
from litellm.proxy.utils import (
create_model_info_response,
get_available_models_for_user,
validate_model_access,
)
# Get available models for the user
all_models = await get_available_models_for_user(
user_api_key_dict=user_api_key_dict,
@ -3847,7 +3859,7 @@ async def model_info(
# Get provider information
_, provider, _, _ = litellm.get_llm_provider(model=model_id)
# Return the model information in the same format as the list endpoint
return create_model_info_response(
model_id=model_id,
@ -3913,7 +3925,7 @@ async def chat_completion( # noqa: PLR0915
data = await _read_request_body(request=request)
base_llm_response_processor = ProxyBaseLLMRequestProcessing(data=data)
try:
return await base_llm_response_processor.base_process_llm_request(
result = await base_llm_response_processor.base_process_llm_request(
request=request,
fastapi_response=fastapi_response,
user_api_key_dict=user_api_key_dict,
@ -3931,6 +3943,10 @@ async def chat_completion( # noqa: PLR0915
user_api_base=user_api_base,
version=version,
)
if isinstance(result, BaseModel):
return result.model_dump(exclude_none=True, exclude_unset=True)
else:
return result
except RejectedRequestError as e:
_data = e.request_data
await proxy_logging_obj.post_call_failure_hook(
@ -5611,40 +5627,45 @@ def _get_provider_token_counter(deployment: dict, model_to_use: str):
"""
if deployment is None:
return None
from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
full_model = deployment.get("litellm_params", {}).get("model", "")
try:
# Use existing LiteLLM logic to determine provider
model, provider, dynamic_api_key, api_base = get_llm_provider(
model=full_model,
custom_llm_provider=deployment.get("litellm_params", {}).get("custom_llm_provider"),
custom_llm_provider=deployment.get("litellm_params", {}).get(
"custom_llm_provider"
),
api_base=deployment.get("litellm_params", {}).get("api_base"),
api_key=deployment.get("litellm_params", {}).get("api_key")
api_key=deployment.get("litellm_params", {}).get("api_key"),
)
# Switch case pattern using existing get_provider_model_info
from litellm.utils import ProviderConfigManager
from litellm.types.utils import LlmProviders
from litellm.utils import ProviderConfigManager
# Convert string provider to LlmProviders enum
llm_provider_enum = LlmProviders(provider)
# Add more provider mappings as needed
if llm_provider_enum:
provider_model_info = ProviderConfigManager.get_provider_model_info(model=full_model, provider=llm_provider_enum)
provider_model_info = ProviderConfigManager.get_provider_model_info(
model=full_model, provider=llm_provider_enum
)
if provider_model_info is not None:
return provider_model_info.get_token_counter()
except Exception:
# If provider detection fails, fall back to manual checks
if full_model.startswith("anthropic/") or "anthropic" in full_model.lower():
from litellm.llms.anthropic.common_utils import AnthropicModelInfo
anthropic_model_info = AnthropicModelInfo()
return anthropic_model_info.get_token_counter()
return None
@ -5679,7 +5700,7 @@ async def token_counter(request: TokenCountRequest, is_direct_request: bool = Tr
break
if deployment is not None:
litellm_model_name = deployment.get("litellm_params", {}).get("model")
# remove the custom_llm_provider_prefix in the litellm_model_name
# remove the custom_llm_provider_prefix in the litellm_model_name
if "/" in litellm_model_name:
litellm_model_name = litellm_model_name.split("/", 1)[1]
@ -5692,15 +5713,17 @@ async def token_counter(request: TokenCountRequest, is_direct_request: bool = Tr
if deployment is not None and not is_direct_request:
# Auto-route to the correct provider based on model
provider_counter = _get_provider_token_counter(deployment, model_to_use)
if provider_counter is not None and provider_counter.supports_provider(deployment=deployment, from_endpoint=not is_direct_request):
if provider_counter is not None and provider_counter.supports_provider(
deployment=deployment, from_endpoint=not is_direct_request
):
result = await provider_counter.count_tokens(
model_to_use=model_to_use,
messages=messages, # type: ignore
messages=messages, # type: ignore
deployment=deployment,
request_model=request.model,
)
if result is not None:
return TokenCountResponse(
total_tokens=result["total_tokens"],

View file

@ -54,6 +54,7 @@ async def test_get_end_user_object(customer_spend, customer_budget):
end_user_id=end_user_id,
prisma_client="RANDOM VALUE", # type: ignore
user_api_key_cache=_cache,
route="/v1/chat/completions",
)
if customer_spend > customer_budget:
pytest.fail(

View file

@ -1043,3 +1043,150 @@ async def test_async_data_generator_midstream_error():
# Verify that post_call_failure_hook was NOT called (since this is not an exception case)
mock_proxy_logging_obj.post_call_failure_hook.assert_not_called()
def _has_nested_none_values(obj, path="root"):
"""
Recursively check if an object contains nested None values.
Args:
obj: The object to check
path: Current path in the object tree (for debugging)
Returns:
List of paths where None values were found
"""
none_paths = []
if obj is None:
none_paths.append(path)
elif isinstance(obj, dict):
for key, value in obj.items():
none_paths.extend(_has_nested_none_values(value, f"{path}.{key}"))
elif isinstance(obj, (list, tuple)):
for i, item in enumerate(obj):
none_paths.extend(_has_nested_none_values(item, f"{path}[{i}]"))
elif hasattr(obj, "__dict__"):
# Handle object attributes
for key, value in obj.__dict__.items():
if not key.startswith("_"): # Skip private attributes
none_paths.extend(_has_nested_none_values(value, f"{path}.{key}"))
return none_paths
@pytest.mark.asyncio
async def test_chat_completion_result_no_nested_none_values():
"""
Test that chat_completion result doesn't have nested None values when using exclude_none=True
"""
from unittest.mock import AsyncMock, MagicMock, patch
from fastapi import Request, Response
from pydantic import BaseModel
import litellm
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.proxy_server import chat_completion
# Create a mock ModelResponse with nested None values
mock_model_response = litellm.ModelResponse()
mock_model_response.id = "test-id"
mock_model_response.model = "gpt-3.5-turbo"
mock_model_response.object = "chat.completion"
mock_model_response.created = 1234567890
# Create message with None values that should be excluded
mock_message = litellm.Message(
content="Hello, world!",
role="assistant",
function_call=None, # This should be excluded
tool_calls=None, # This should be excluded
audio=None, # This should be excluded
reasoning_content=None, # This should be excluded
thinking_blocks=None, # This should be excluded
annotations=None, # This should be excluded
)
# Create choice with potential None values
mock_choice = litellm.Choices(
finish_reason="stop",
index=0,
message=mock_message,
logprobs=None, # This should be excluded when exclude_none=True
)
mock_model_response.choices = [mock_choice]
mock_model_response.usage = litellm.Usage(
prompt_tokens=10, completion_tokens=5, total_tokens=15
)
# Verify the mock has None values before serialization
raw_dict = mock_model_response.model_dump()
none_paths_before = _has_nested_none_values(raw_dict)
assert (
len(none_paths_before) > 0
), "Mock should have None values before exclude_none=True"
# Mock the request processing to return our mock response
mock_base_processor = MagicMock()
mock_base_processor.base_process_llm_request = AsyncMock(
return_value=mock_model_response
)
# Mock other dependencies
mock_request = MagicMock(spec=Request)
mock_response = MagicMock(spec=Response)
mock_user_api_key_dict = MagicMock(spec=UserAPIKeyAuth)
with patch(
"litellm.proxy.proxy_server._read_request_body",
return_value={"model": "gpt-3.5-turbo", "messages": []},
), patch(
"litellm.proxy.proxy_server.ProxyBaseLLMRequestProcessing",
return_value=mock_base_processor,
):
# Call the chat_completion function
result = await chat_completion(
request=mock_request,
fastapi_response=mock_response,
user_api_key_dict=mock_user_api_key_dict,
)
# Verify the result is a dict (since isinstance(result, BaseModel) was True)
assert isinstance(result, dict), f"Expected dict result, got {type(result)}"
# Check that there are no nested None values in the result
none_paths_after = _has_nested_none_values(result)
assert (
len(none_paths_after) == 0
), f"Result should not contain nested None values. Found None at: {none_paths_after}"
# Verify essential fields are present
assert "id" in result
assert "model" in result
assert "object" in result
assert "created" in result
assert "choices" in result
assert "usage" in result
# Verify that the choices contain the expected message content
assert len(result["choices"]) == 1
assert result["choices"][0]["message"]["content"] == "Hello, world!"
assert result["choices"][0]["message"]["role"] == "assistant"
# Verify that None fields were excluded (should not be present in the dict)
message = result["choices"][0]["message"]
excluded_fields = [
"function_call",
"tool_calls",
"audio",
"reasoning_content",
"thinking_blocks",
"annotations",
]
for field in excluded_fields:
assert (
field not in message
), f"Field '{field}' should be excluded when it's None"