diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 9012dee8fc0..e9beadbb5e4 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -492,6 +492,8 @@ class LiteLLMRoutes(enum.Enum): "/global/spend/end_users", "/global/activity", "/global/activity/model", + "/v1/models/{model_id}", + "/models/{model_id}", ] + spend_tracking_routes + key_management_routes diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index af9f5d363da..dffb12873ce 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -1728,7 +1728,7 @@ class ProxyConfig: self._load_environment_variables(config=config) ## Callback settings - callback_settings = config.get("callback_settings", None) + callback_settings = config.get("callback_settings", {}) ## LITELLM MODULE SETTINGS (e.g. litellm.drop_params=True,..) litellm_settings = config.get("litellm_settings", None) @@ -3766,100 +3766,34 @@ 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 - all_models = [] - model_access_groups: Dict[str, List[str]] = defaultdict(list) - ## CHECK IF MODEL RESTRICTIONS ARE SET AT KEY/TEAM LEVEL ## - if llm_router is None: - proxy_model_list = [] - else: - proxy_model_list = llm_router.get_model_names() - model_access_groups = llm_router.get_model_access_groups() - - ## if only_model_access_groups is True, - """ - 1. Get all models key/user/team has access to - 2. Filter out models that are not model access groups - 3. Return the models - """ - if only_model_access_groups is True: - include_model_access_groups = True - - key_models = get_key_models( + + from litellm.proxy.utils import get_available_models_for_user, create_model_info_response + + # Get available models for the user + all_models = await get_available_models_for_user( user_api_key_dict=user_api_key_dict, - proxy_model_list=proxy_model_list, - model_access_groups=model_access_groups, - include_model_access_groups=include_model_access_groups, - ) - - team_models: List[str] = user_api_key_dict.team_models - - if team_id: - key_models = [] - team_object = await get_team_object( - team_id=team_id, - prisma_client=prisma_client, - user_api_key_cache=user_api_key_cache, - proxy_logging_obj=proxy_logging_obj, - ) - validate_membership(user_api_key_dict=user_api_key_dict, team_table=team_object) - team_models = team_object.models - - team_models = get_team_models( - team_models=team_models, - proxy_model_list=proxy_model_list, - model_access_groups=model_access_groups, - include_model_access_groups=include_model_access_groups, - ) - - all_models = get_complete_model_list( - key_models=key_models, - team_models=team_models, - proxy_model_list=proxy_model_list, - user_model=user_model, - infer_model_from_keys=general_settings.get("infer_model_from_keys", False), - return_wildcard_routes=return_wildcard_routes, llm_router=llm_router, - model_access_groups=model_access_groups, - include_model_access_groups=include_model_access_groups, - only_model_access_groups=only_model_access_groups, + general_settings=general_settings, + user_model=user_model, + prisma_client=prisma_client, + proxy_logging_obj=proxy_logging_obj, + team_id=team_id, + include_model_access_groups=include_model_access_groups or False, + only_model_access_groups=only_model_access_groups or False, + return_wildcard_routes=return_wildcard_routes or False, + user_api_key_cache=user_api_key_cache, ) # Build response data model_data = [] for model in all_models: - model_info = { - "id": model, - "object": "model", - "created": DEFAULT_MODEL_CREATED_AT_TIME, - "owned_by": "openai", - } - - # Add metadata if requested - if include_metadata: - metadata = {} - - # Default fallback_type to "general" if include_metadata is true - effective_fallback_type = ( - fallback_type if fallback_type is not None else "general" - ) - - # Validate fallback_type - valid_fallback_types = ["general", "context_window", "content_policy"] - if effective_fallback_type not in valid_fallback_types: - raise HTTPException( - status_code=400, - detail=f"Invalid fallback_type. Must be one of: {valid_fallback_types}", - ) - - fallbacks = get_all_fallbacks( - model=model, - llm_router=llm_router, - fallback_type=effective_fallback_type, - ) - metadata["fallbacks"] = fallbacks - - model_info["metadata"] = metadata - + model_info = create_model_info_response( + model_id=model, + provider="openai", + include_metadata=include_metadata or False, + fallback_type=fallback_type, + llm_router=llm_router, + ) model_data.append(model_info) return dict( @@ -3868,6 +3802,60 @@ async def model_list( ) +@router.get( + "/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"] +) +async def model_info( + model_id: str, + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), +): + """ + 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 + + # Get available models for the user + all_models = await get_available_models_for_user( + user_api_key_dict=user_api_key_dict, + llm_router=llm_router, + general_settings=general_settings, + user_model=user_model, + prisma_client=prisma_client, + proxy_logging_obj=proxy_logging_obj, + team_id=None, + include_model_access_groups=False, + only_model_access_groups=False, + return_wildcard_routes=False, + user_api_key_cache=user_api_key_cache, + ) + + # Validate that the requested model is accessible + validate_model_access(model_id=model_id, available_models=all_models) + + # 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, + provider=provider, + include_metadata=False, + fallback_type=None, + llm_router=llm_router, + ) + + @router.post( "/v1/chat/completions", dependencies=[Depends(user_api_key_auth)], @@ -6930,56 +6918,21 @@ async def model_group_info( status_code=500, detail={"error": "LLM Router is not loaded in"} ) - ## CHECK IF MODEL RESTRICTIONS ARE SET AT KEY/TEAM LEVEL ## - model_access_groups: Dict[str, List[str]] = defaultdict(list) - if llm_router is None: - proxy_model_list = [] - else: - proxy_model_list = llm_router.get_model_names() - model_access_groups = llm_router.get_model_access_groups() + from litellm.proxy.utils import get_available_models_for_user - key_models = get_key_models( + # Get available models for the user + all_models_str = await get_available_models_for_user( user_api_key_dict=user_api_key_dict, - proxy_model_list=proxy_model_list, - model_access_groups=model_access_groups, - ) - team_models = [] - if ( - not user_api_key_dict.team_id - and user_api_key_dict.user_id is not None - and not _user_has_admin_view(user_api_key_dict) - ): - if prisma_client is None: - raise HTTPException( - status_code=500, - detail={"error": CommonProxyErrors.db_not_connected_error.value}, - ) - user_object = await prisma_client.db.litellm_usertable.find_first( - where={"user_id": user_api_key_dict.user_id} - ) - user_object_typed = LiteLLM_UserTable(**user_object.model_dump()) - user_models = [] - if user_object is not None: - user_models = get_team_models( - team_models=user_object_typed.models, - proxy_model_list=proxy_model_list, - model_access_groups=model_access_groups, - ) - team_models = user_models - else: - team_models = get_team_models( - team_models=user_api_key_dict.team_models, - proxy_model_list=proxy_model_list, - model_access_groups=model_access_groups, - ) - - all_models_str = get_complete_model_list( - key_models=key_models, - team_models=team_models, - proxy_model_list=proxy_model_list, - user_model=user_model, - infer_model_from_keys=general_settings.get("infer_model_from_keys", False), llm_router=llm_router, + general_settings=general_settings, + user_model=user_model, + prisma_client=prisma_client, + proxy_logging_obj=proxy_logging_obj, + team_id=None, + include_model_access_groups=False, + only_model_access_groups=False, + return_wildcard_routes=False, + user_api_key_cache=user_api_key_cache, ) model_groups: List[ModelGroupInfoProxy] = _get_model_group_info( llm_router=llm_router, all_models_str=all_models_str, model_group=model_group diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index f1e113d8ded..16e151e5e19 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -22,7 +22,7 @@ from typing import ( overload, ) -from litellm.constants import MAX_TEAM_LIST_LIMIT +from litellm.constants import MAX_TEAM_LIST_LIMIT, DEFAULT_MODEL_CREATED_AT_TIME from litellm.proxy._types import ( DB_CONNECTION_ERROR_TYPES, CommonProxyErrors, @@ -3831,3 +3831,176 @@ def construct_database_url_from_env_vars() -> Optional[str]: return database_url return None + + +async def get_available_models_for_user( + user_api_key_dict: "UserAPIKeyAuth", + llm_router: Optional["Router"], + general_settings: dict, + user_model: Optional[str], + prisma_client: Optional["PrismaClient"] = None, + proxy_logging_obj: Optional["ProxyLogging"] = None, + team_id: Optional[str] = None, + include_model_access_groups: bool = False, + only_model_access_groups: bool = False, + return_wildcard_routes: bool = False, + user_api_key_cache: Optional["DualCache"] = None, +) -> List[str]: + """ + Get the list of models available to a user based on their API key and team permissions. + + Args: + user_api_key_dict: User API key authentication object + llm_router: LiteLLM router instance + general_settings: General settings from config + user_model: User-specific model + prisma_client: Prisma client for database operations + proxy_logging_obj: Proxy logging object + team_id: Specific team ID to check (optional) + include_model_access_groups: Whether to include model access groups + only_model_access_groups: Whether to only return model access groups + return_wildcard_routes: Whether to return wildcard routes + + Returns: + List of model names available to the user + """ + from litellm.proxy.auth.model_checks import ( + get_key_models, + get_team_models, + get_complete_model_list, + ) + from litellm.proxy.auth.auth_checks import get_team_object + from litellm.proxy.management_endpoints.team_endpoints import validate_membership + + # Get proxy model list and access groups + if llm_router is None: + proxy_model_list = [] + model_access_groups = {} + else: + proxy_model_list = llm_router.get_model_names() + model_access_groups = llm_router.get_model_access_groups() + + # Get key models + key_models = get_key_models( + user_api_key_dict=user_api_key_dict, + proxy_model_list=proxy_model_list, + model_access_groups=model_access_groups, + include_model_access_groups=include_model_access_groups, + ) + + # Get team models + team_models: List[str] = user_api_key_dict.team_models + + # If specific team_id is provided, validate and get team models + if team_id and prisma_client and proxy_logging_obj and user_api_key_cache: + key_models = [] + team_object = await get_team_object( + team_id=team_id, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + validate_membership(user_api_key_dict=user_api_key_dict, team_table=team_object) + team_models = team_object.models + + team_models = get_team_models( + team_models=team_models, + proxy_model_list=proxy_model_list, + model_access_groups=model_access_groups, + include_model_access_groups=include_model_access_groups, + ) + + # Get complete model list + all_models = get_complete_model_list( + key_models=key_models, + team_models=team_models, + proxy_model_list=proxy_model_list, + user_model=user_model, + infer_model_from_keys=general_settings.get("infer_model_from_keys", False), + return_wildcard_routes=return_wildcard_routes, + llm_router=llm_router, + model_access_groups=model_access_groups, + include_model_access_groups=include_model_access_groups, + only_model_access_groups=only_model_access_groups, + ) + + return all_models + + +def create_model_info_response( + model_id: str, + provider: str, + include_metadata: bool = False, + fallback_type: Optional[str] = None, + llm_router: Optional["Router"] = None, +) -> dict: + """ + Create a standardized model info response. + + Args: + model_id: The model ID + provider: The model provider + include_metadata: Whether to include metadata + fallback_type: Type of fallbacks to include + llm_router: LiteLLM router instance + + Returns: + Dictionary containing model information + """ + from litellm.proxy.auth.model_checks import get_all_fallbacks + + model_info = { + "id": model_id, + "object": "model", + "created": DEFAULT_MODEL_CREATED_AT_TIME, + "owned_by": provider, + } + + # Add metadata if requested + if include_metadata: + metadata = {} + + # Default fallback_type to "general" if include_metadata is true + effective_fallback_type = ( + fallback_type if fallback_type is not None else "general" + ) + + # Validate fallback_type + valid_fallback_types = ["general", "context_window", "content_policy"] + if effective_fallback_type not in valid_fallback_types: + raise HTTPException( + status_code=400, + detail=f"Invalid fallback_type. Must be one of: {valid_fallback_types}", + ) + + fallbacks = get_all_fallbacks( + model=model_id, + llm_router=llm_router, + fallback_type=effective_fallback_type, + ) + metadata["fallbacks"] = fallbacks + + model_info["metadata"] = metadata + + return model_info + + +def validate_model_access( + model_id: str, + available_models: List[str], +) -> None: + """ + Validate that a model is accessible to the user. + + Args: + model_id: The model ID to validate + available_models: List of models available to the user + + Raises: + HTTPException: If the model is not accessible + """ + if model_id not in available_models: + raise HTTPException( + status_code=404, + detail="The model `{}` does not exist or is not accessible".format(model_id) + ) \ No newline at end of file diff --git a/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py index c5fdbf4925a..bd37e9cbe41 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py @@ -451,3 +451,119 @@ class TestClearCache: mock_config.add_deployment.assert_called_once_with( prisma_client=mock_prisma, proxy_logging_obj=mock_logging ) + + +class TestModelInfoEndpoint: + """Test the model_info endpoint for retrieving individual model information""" + + @pytest.mark.asyncio + async def test_model_info_accessible_model_success(self): + """Test model_info returns model data for accessible models""" + from litellm.proxy.proxy_server import model_info + + # Mock user with access to specific models + user_api_key_dict = UserAPIKeyAuth( + user_id="test_user", + api_key="test_key", + models=["gpt-4", "claude-3"], + team_models=["gpt-3.5-turbo"] + ) + + with patch("litellm.proxy.proxy_server.llm_router") as mock_router, \ + patch("litellm.proxy.proxy_server.get_key_models") as mock_get_key_models, \ + patch("litellm.proxy.proxy_server.get_team_models") as mock_get_team_models, \ + patch("litellm.proxy.proxy_server.get_complete_model_list") as mock_get_complete_models, \ + patch("litellm.get_llm_provider") as mock_get_provider: + + # Setup mocks + mock_router.get_model_names.return_value = ["gpt-4", "claude-3", "gpt-3.5-turbo"] + mock_router.get_model_access_groups.return_value = {} + mock_get_key_models.return_value = ["gpt-4", "claude-3"] + mock_get_team_models.return_value = ["gpt-3.5-turbo"] + mock_get_complete_models.return_value = ["gpt-4", "claude-3", "gpt-3.5-turbo"] + mock_get_provider.return_value = (None, "openai", None, None) + + # Test accessible model + result = await model_info( + model_id="gpt-4", + user_api_key_dict=user_api_key_dict + ) + + assert result["id"] == "gpt-4" + assert result["object"] == "model" + assert result["owned_by"] == "openai" + assert "created" in result + + @pytest.mark.asyncio + async def test_model_info_inaccessible_model_returns_404(self): + """Test model_info returns 404 for inaccessible models""" + from litellm.proxy.proxy_server import model_info + from fastapi import HTTPException + + # Mock user with limited access + user_api_key_dict = UserAPIKeyAuth( + user_id="test_user", + api_key="test_key", + models=["gpt-4"], # Only has access to gpt-4 + team_models=[] + ) + + with patch("litellm.proxy.proxy_server.llm_router") as mock_router, \ + patch("litellm.proxy.proxy_server.get_key_models") as mock_get_key_models, \ + patch("litellm.proxy.proxy_server.get_team_models") as mock_get_team_models, \ + patch("litellm.proxy.proxy_server.get_complete_model_list") as mock_get_complete_models: + + # Setup mocks - user only has access to gpt-4 + mock_router.get_model_names.return_value = ["gpt-4", "claude-3"] + mock_router.get_model_access_groups.return_value = {} + mock_get_key_models.return_value = ["gpt-4"] + mock_get_team_models.return_value = [] + mock_get_complete_models.return_value = ["gpt-4"] # Only gpt-4 accessible + + # Test inaccessible model should raise 404 + with pytest.raises(HTTPException) as exc_info: + await model_info( + model_id="claude-3", # Not in user's accessible models + user_api_key_dict=user_api_key_dict + ) + + assert exc_info.value.status_code == 404 + assert "does not exist or is not accessible" in exc_info.value.detail + + @pytest.mark.asyncio + async def test_model_info_team_model_access(self): + """Test model_info works with team model access""" + from litellm.proxy.proxy_server import model_info + + # Mock user with team access + user_api_key_dict = UserAPIKeyAuth( + user_id="test_user", + api_key="test_key", + team_id="test_team", + models=[], # No direct key models + team_models=["team-model-1"] + ) + + with patch("litellm.proxy.proxy_server.llm_router") as mock_router, \ + patch("litellm.proxy.proxy_server.get_key_models") as mock_get_key_models, \ + patch("litellm.proxy.proxy_server.get_team_models") as mock_get_team_models, \ + patch("litellm.proxy.proxy_server.get_complete_model_list") as mock_get_complete_models, \ + patch("litellm.get_llm_provider") as mock_get_provider: + + # Setup mocks + mock_router.get_model_names.return_value = ["team-model-1"] + mock_router.get_model_access_groups.return_value = {} + mock_get_key_models.return_value = [] + mock_get_team_models.return_value = ["team-model-1"] + mock_get_complete_models.return_value = ["team-model-1"] + mock_get_provider.return_value = (None, "custom", None, None) + + # Test team model access + result = await model_info( + model_id="team-model-1", + user_api_key_dict=user_api_key_dict + ) + + assert result["id"] == "team-model-1" + assert result["object"] == "model" + assert result["owned_by"] == "custom"