[LLM Translation] Support /v1/models/{model_id} retrieval (#13268)

* added model id endpoint

* fix test

* add route to internal users

* make the functions reusable

* fixed mypy
This commit is contained in:
Jugal D. Bhatt 2025-08-04 18:03:59 -07:00 • committed by GitHub
parent de7108b5f8
commit efd34966dc
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 380 additions and 136 deletions

View file

@ -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

View file

@ -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

View file

@ -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)
)

View file

@ -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"