diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index e9beadbb5e4..b69d27775c1 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -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 diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 89c012c20ae..b306512847f 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -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 diff --git a/litellm/proxy/auth/handle_jwt.py b/litellm/proxy/auth/handle_jwt.py index 12146f0e866..60529f1e2f4 100644 --- a/litellm/proxy/auth/handle_jwt.py +++ b/litellm/proxy/auth/handle_jwt.py @@ -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( diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index 00711afee70..9efa904574a 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -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"] = ( diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 3676420bab0..0732f2eb454 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -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"], diff --git a/tests/proxy_unit_tests/test_auth_checks.py b/tests/proxy_unit_tests/test_auth_checks.py index b78543f2c8c..2c829406c01 100644 --- a/tests/proxy_unit_tests/test_auth_checks.py +++ b/tests/proxy_unit_tests/test_auth_checks.py @@ -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( diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index 5840d52813f..3e2e8fba2e8 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -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"