mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
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:
parent
92c525ddfe
commit
0da25fadc0
7 changed files with 212 additions and 32 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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"] = (
|
||||
|
|
|
|||
|
|
@ -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"],
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue