mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
feat: extended /v1/models endpoint, now it returns with fallbacks on demand (#12811)
* Extended `/v1/model` endpoint to support fallbacks * unit tests reworked * linting fixes * fix lining error * fix linting
This commit is contained in:
parent
37c626a9a0
commit
a6ddf5c744
4 changed files with 612 additions and 11 deletions
|
|
@ -6,6 +6,7 @@ import litellm
|
|||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.proxy._types import SpecialModelNames, UserAPIKeyAuth
|
||||
from litellm.router import Router
|
||||
from litellm.router_utils.fallback_event_handlers import get_fallback_model_group
|
||||
from litellm.types.router import LiteLLM_Params
|
||||
from litellm.utils import get_valid_models
|
||||
|
||||
|
|
@ -52,7 +53,7 @@ def _get_models_from_access_groups(
|
|||
if model in model_access_groups:
|
||||
if (
|
||||
not include_model_access_groups
|
||||
): # remove access group, unless requested - e.g. when creating a key and trying to see list of models
|
||||
): # remove access group, unless requested - e.g. when creating a key
|
||||
idx_to_remove.append(idx)
|
||||
new_models.extend(model_access_groups[model])
|
||||
|
||||
|
|
@ -104,7 +105,8 @@ def get_key_models(
|
|||
- List of model name strings
|
||||
- Empty list if no models set
|
||||
- If model_access_groups is provided, only return models that are in the access groups
|
||||
- If include_model_access_groups is True, it includes the 'keys' of the model_access_groups in the response - {"beta-models": ["gpt-4", "claude-v1"]} -> returns 'beta-models'
|
||||
- If include_model_access_groups is True, it includes the 'keys' of the model_access_groups
|
||||
in the response - {"beta-models": ["gpt-4", "claude-v1"]} -> returns 'beta-models'
|
||||
"""
|
||||
all_models: List[str] = []
|
||||
if len(user_api_key_dict.models) > 0:
|
||||
|
|
@ -287,3 +289,53 @@ def _get_wildcard_models(
|
|||
unique_models.remove(model)
|
||||
|
||||
return all_wildcard_models
|
||||
|
||||
|
||||
def get_all_fallbacks(
|
||||
model: str,
|
||||
llm_router: Optional[Router] = None,
|
||||
fallback_type: str = "general",
|
||||
) -> List[str]:
|
||||
"""
|
||||
Get all fallbacks for a given model from the router's fallback configuration.
|
||||
|
||||
Args:
|
||||
model: The model name to get fallbacks for
|
||||
llm_router: The LiteLLM router instance
|
||||
fallback_type: Type of fallback ("general", "context_window", "content_policy")
|
||||
|
||||
Returns:
|
||||
List of fallback model names. Empty list if no fallbacks found.
|
||||
"""
|
||||
if llm_router is None:
|
||||
return []
|
||||
|
||||
# Get the appropriate fallback list based on type
|
||||
fallbacks_config: list = []
|
||||
if fallback_type == "general":
|
||||
fallbacks_config = getattr(llm_router, "fallbacks", [])
|
||||
elif fallback_type == "context_window":
|
||||
fallbacks_config = getattr(llm_router, "context_window_fallbacks", [])
|
||||
elif fallback_type == "content_policy":
|
||||
fallbacks_config = getattr(llm_router, "content_policy_fallbacks", [])
|
||||
else:
|
||||
verbose_proxy_logger.warning(f"Unknown fallback_type: {fallback_type}")
|
||||
return []
|
||||
|
||||
if not fallbacks_config:
|
||||
return []
|
||||
|
||||
try:
|
||||
# Use existing function to get fallback model group
|
||||
fallback_model_group, _ = get_fallback_model_group(
|
||||
fallbacks=fallbacks_config,
|
||||
model_group=model
|
||||
)
|
||||
|
||||
if fallback_model_group is None:
|
||||
return []
|
||||
|
||||
return fallback_model_group
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.error(f"Error getting fallbacks for model {model}: {e}")
|
||||
return []
|
||||
|
|
|
|||
|
|
@ -168,6 +168,7 @@ from litellm.proxy.auth.auth_utils import check_response_size_is_safe
|
|||
from litellm.proxy.auth.handle_jwt import JWTHandler
|
||||
from litellm.proxy.auth.litellm_license import LicenseCheck
|
||||
from litellm.proxy.auth.model_checks import (
|
||||
get_all_fallbacks,
|
||||
get_complete_model_list,
|
||||
get_key_models,
|
||||
get_mcp_server_ids,
|
||||
|
|
@ -3663,11 +3664,18 @@ async def model_list(
|
|||
team_id: Optional[str] = None,
|
||||
include_model_access_groups: Optional[bool] = False,
|
||||
only_model_access_groups: Optional[bool] = False,
|
||||
include_metadata: Optional[bool] = False,
|
||||
fallback_type: Optional[str] = None,
|
||||
):
|
||||
"""
|
||||
Use `/model/info` - to get detailed model information, example - pricing, mode, etc.
|
||||
|
||||
This is just for compatibility with openai projects like aider.
|
||||
|
||||
Query Parameters:
|
||||
- include_metadata: Include additional metadata in the response with fallback information
|
||||
- fallback_type: Type of fallbacks to include ("general", "context_window", "content_policy")
|
||||
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 = []
|
||||
|
|
@ -3728,16 +3736,44 @@ async def model_list(
|
|||
only_model_access_groups=only_model_access_groups,
|
||||
)
|
||||
|
||||
# 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_data.append(model_info)
|
||||
|
||||
return dict(
|
||||
data=[
|
||||
{
|
||||
"id": model,
|
||||
"object": "model",
|
||||
"created": DEFAULT_MODEL_CREATED_AT_TIME,
|
||||
"owned_by": "openai",
|
||||
}
|
||||
for model in all_models
|
||||
],
|
||||
data=model_data,
|
||||
object="list",
|
||||
)
|
||||
|
||||
|
|
|
|||
271
tests/proxy_unit_tests/test_models_fallback_endpoint.py
Normal file
271
tests/proxy_unit_tests/test_models_fallback_endpoint.py
Normal file
|
|
@ -0,0 +1,271 @@
|
|||
import pytest
|
||||
from unittest.mock import Mock, patch
|
||||
|
||||
|
||||
def create_mock_user_api_key_auth():
|
||||
"""Create mock user API key authentication."""
|
||||
mock_auth = Mock()
|
||||
mock_auth.api_key = "test-key"
|
||||
mock_auth.user_id = "test-user"
|
||||
mock_auth.team_id = "test-team"
|
||||
mock_auth.team_models = []
|
||||
mock_auth.models = []
|
||||
return mock_auth
|
||||
|
||||
|
||||
def create_mock_router_with_fallbacks():
|
||||
"""Create a mock router with fallback configurations."""
|
||||
router = Mock()
|
||||
router.fallbacks = [
|
||||
{"claude-4-sonnet": ["bedrock-claude-sonnet-4", "google-claude-sonnet-4"]},
|
||||
{"gpt-4": ["gpt-4-turbo", "gpt-3.5-turbo"]}
|
||||
]
|
||||
router.context_window_fallbacks = [
|
||||
{"claude-4-sonnet": ["claude-3-sonnet"]},
|
||||
{"gpt-4": ["gpt-3.5-turbo"]}
|
||||
]
|
||||
router.content_policy_fallbacks = [
|
||||
{"claude-4-sonnet": ["claude-3-haiku"]}
|
||||
]
|
||||
router.get_model_names.return_value = [
|
||||
"claude-4-sonnet", "bedrock-claude-sonnet-4", "google-claude-sonnet-4",
|
||||
"gpt-4", "gpt-4-turbo", "gpt-3.5-turbo"
|
||||
]
|
||||
router.get_model_access_groups.return_value = {}
|
||||
return router
|
||||
|
||||
|
||||
def test_model_list_function_signature():
|
||||
"""Test that model_list function has the correct signature with new parameters."""
|
||||
from litellm.proxy.proxy_server import model_list
|
||||
import inspect
|
||||
|
||||
sig = inspect.signature(model_list)
|
||||
params = list(sig.parameters.keys())
|
||||
|
||||
# Check that our new parameters are present
|
||||
assert 'include_metadata' in params, "include_metadata parameter missing"
|
||||
assert 'fallback_type' in params, "fallback_type parameter missing"
|
||||
|
||||
# Check parameter defaults
|
||||
include_metadata_param = sig.parameters['include_metadata']
|
||||
fallback_type_param = sig.parameters['fallback_type']
|
||||
|
||||
assert include_metadata_param.default is False, "include_metadata should default to False"
|
||||
assert fallback_type_param.default is None, "fallback_type should default to None"
|
||||
|
||||
|
||||
@patch('litellm.proxy.proxy_server.llm_router')
|
||||
@patch('litellm.proxy.proxy_server.get_complete_model_list')
|
||||
@patch('litellm.proxy.proxy_server.get_key_models')
|
||||
@patch('litellm.proxy.proxy_server.get_team_models')
|
||||
@patch('litellm.proxy.proxy_server.get_all_fallbacks')
|
||||
def test_model_list_with_fallback_metadata(
|
||||
mock_get_all_fallbacks, mock_get_team_models, mock_get_key_models,
|
||||
mock_get_complete_model_list, mock_router
|
||||
):
|
||||
"""Test model_list function with fallback metadata."""
|
||||
|
||||
# Setup mocks
|
||||
mock_user_auth = create_mock_user_api_key_auth()
|
||||
mock_router_instance = create_mock_router_with_fallbacks()
|
||||
mock_router.return_value = mock_router_instance
|
||||
|
||||
mock_get_key_models.return_value = []
|
||||
mock_get_team_models.return_value = []
|
||||
mock_get_complete_model_list.return_value = ["claude-4-sonnet", "bedrock-claude-sonnet-4"]
|
||||
|
||||
# Mock fallback responses
|
||||
def fallback_side_effect(model, llm_router, fallback_type):
|
||||
if model == "claude-4-sonnet":
|
||||
return ["bedrock-claude-sonnet-4", "google-claude-sonnet-4"]
|
||||
return []
|
||||
|
||||
mock_get_all_fallbacks.side_effect = fallback_side_effect
|
||||
|
||||
# Test async function call (simplified - just test the logic)
|
||||
# Note: This is a simplified test since we can't easily run the full async endpoint
|
||||
# The important thing is that our function signature and logic are correct
|
||||
|
||||
# Import the constants we need
|
||||
try:
|
||||
from litellm.proxy.proxy_server import DEFAULT_MODEL_CREATED_AT_TIME
|
||||
except ImportError:
|
||||
DEFAULT_MODEL_CREATED_AT_TIME = 1640995200 # Default fallback
|
||||
|
||||
# Test with include_metadata=True (should default to general fallbacks)
|
||||
all_models = ["claude-4-sonnet", "bedrock-claude-sonnet-4"]
|
||||
|
||||
# Build response manually to test our logic
|
||||
model_data = []
|
||||
for model in all_models:
|
||||
model_info = {
|
||||
"id": model,
|
||||
"object": "model",
|
||||
"created": DEFAULT_MODEL_CREATED_AT_TIME,
|
||||
"owned_by": "openai",
|
||||
}
|
||||
|
||||
# Test metadata logic
|
||||
include_metadata = True
|
||||
fallback_type = None # Should default to "general"
|
||||
|
||||
if include_metadata:
|
||||
metadata = {}
|
||||
effective_fallback_type = fallback_type if fallback_type is not None else "general"
|
||||
|
||||
# Validate fallback_type
|
||||
valid_fallback_types = ["general", "context_window", "content_policy"]
|
||||
assert effective_fallback_type in valid_fallback_types
|
||||
|
||||
fallbacks = fallback_side_effect(model, mock_router_instance, effective_fallback_type)
|
||||
metadata["fallbacks"] = fallbacks
|
||||
model_info["metadata"] = metadata
|
||||
|
||||
model_data.append(model_info)
|
||||
|
||||
response = {
|
||||
"data": model_data,
|
||||
"object": "list",
|
||||
}
|
||||
|
||||
# Verify response structure
|
||||
assert "data" in response
|
||||
assert "object" in response
|
||||
assert response["object"] == "list"
|
||||
|
||||
# Find claude-4-sonnet in response
|
||||
claude_model = next((m for m in response["data"] if m["id"] == "claude-4-sonnet"), None)
|
||||
assert claude_model is not None
|
||||
assert "metadata" in claude_model
|
||||
assert "fallbacks" in claude_model["metadata"]
|
||||
assert claude_model["metadata"]["fallbacks"] == [
|
||||
"bedrock-claude-sonnet-4", "google-claude-sonnet-4"
|
||||
]
|
||||
|
||||
# Find bedrock-claude-sonnet-4 in response (should have no fallbacks)
|
||||
bedrock_model = next(
|
||||
(m for m in response["data"] if m["id"] == "bedrock-claude-sonnet-4"), None
|
||||
)
|
||||
assert bedrock_model is not None
|
||||
assert "metadata" in bedrock_model
|
||||
assert "fallbacks" in bedrock_model["metadata"]
|
||||
assert bedrock_model["metadata"]["fallbacks"] == []
|
||||
|
||||
|
||||
def test_model_list_invalid_fallback_type_validation():
|
||||
"""Test that invalid fallback_type raises proper validation error."""
|
||||
# Test the validation logic
|
||||
valid_fallback_types = ["general", "context_window", "content_policy"]
|
||||
|
||||
# Valid types should pass
|
||||
for valid_type in valid_fallback_types:
|
||||
assert valid_type in valid_fallback_types
|
||||
|
||||
# Invalid type should fail validation
|
||||
invalid_type = "invalid"
|
||||
assert invalid_type not in valid_fallback_types
|
||||
|
||||
# Test HTTPException creation logic
|
||||
try:
|
||||
from fastapi import HTTPException
|
||||
|
||||
# This is the logic from our endpoint
|
||||
if invalid_type not in valid_fallback_types:
|
||||
error = HTTPException(
|
||||
status_code=400,
|
||||
detail=f"Invalid fallback_type. Must be one of: {valid_fallback_types}"
|
||||
)
|
||||
assert error.status_code == 400
|
||||
assert "Invalid fallback_type" in error.detail
|
||||
assert "general" in error.detail
|
||||
assert "context_window" in error.detail
|
||||
assert "content_policy" in error.detail
|
||||
except ImportError:
|
||||
# FastAPI not available, skip this part
|
||||
pass
|
||||
|
||||
|
||||
def test_fallback_type_defaults_to_general():
|
||||
"""Test that fallback_type defaults to 'general' when include_metadata=True."""
|
||||
# Test the defaulting logic
|
||||
include_metadata = True
|
||||
fallback_type = None
|
||||
|
||||
if include_metadata:
|
||||
effective_fallback_type = fallback_type if fallback_type is not None else "general"
|
||||
assert effective_fallback_type == "general"
|
||||
|
||||
# Test with explicit general type
|
||||
fallback_type = "general"
|
||||
effective_fallback_type = fallback_type if fallback_type is not None else "general"
|
||||
assert effective_fallback_type == "general"
|
||||
|
||||
# Test with other types
|
||||
fallback_type = "context_window"
|
||||
effective_fallback_type = fallback_type if fallback_type is not None else "general"
|
||||
assert effective_fallback_type == "context_window"
|
||||
|
||||
|
||||
def test_response_structure_compatibility():
|
||||
"""Test that response structure maintains OpenAI compatibility."""
|
||||
# Test basic model structure (without metadata)
|
||||
basic_model = {
|
||||
"id": "claude-4-sonnet",
|
||||
"object": "model",
|
||||
"created": 1640995200,
|
||||
"owned_by": "openai"
|
||||
}
|
||||
|
||||
required_keys = ["id", "object", "created", "owned_by"]
|
||||
for key in required_keys:
|
||||
assert key in basic_model, f"Required OpenAI key '{key}' missing"
|
||||
|
||||
# Test model with metadata
|
||||
metadata_model = {
|
||||
**basic_model,
|
||||
"metadata": {
|
||||
"fallbacks": ["bedrock-claude-sonnet-4", "google-claude-sonnet-4"]
|
||||
}
|
||||
}
|
||||
|
||||
# Should still have all required keys
|
||||
for key in required_keys:
|
||||
assert key in metadata_model, f"Required OpenAI key '{key}' missing from metadata model"
|
||||
|
||||
# Should have metadata
|
||||
assert "metadata" in metadata_model
|
||||
assert "fallbacks" in metadata_model["metadata"]
|
||||
assert isinstance(metadata_model["metadata"]["fallbacks"], list)
|
||||
|
||||
# Test complete response structure
|
||||
response = {
|
||||
"data": [basic_model, metadata_model],
|
||||
"object": "list"
|
||||
}
|
||||
|
||||
assert "data" in response
|
||||
assert "object" in response
|
||||
assert response["object"] == "list"
|
||||
assert isinstance(response["data"], list)
|
||||
assert len(response["data"]) == 2
|
||||
|
||||
|
||||
def test_get_all_fallbacks_integration():
|
||||
"""Test that get_all_fallbacks function can be imported and has correct signature."""
|
||||
from litellm.proxy.auth.model_checks import get_all_fallbacks
|
||||
import inspect
|
||||
|
||||
# Test function signature
|
||||
sig = inspect.signature(get_all_fallbacks)
|
||||
params = list(sig.parameters.keys())
|
||||
expected_params = ['model', 'llm_router', 'fallback_type']
|
||||
|
||||
assert params == expected_params, f"Expected {expected_params}, got {params}"
|
||||
|
||||
# Test default parameter values
|
||||
fallback_type_param = sig.parameters['fallback_type']
|
||||
assert fallback_type_param.default == "general", "fallback_type should default to 'general'"
|
||||
|
||||
llm_router_param = sig.parameters['llm_router']
|
||||
assert llm_router_param.default is None, "llm_router should default to None"
|
||||
242
tests/test_litellm/proxy/auth/test_model_checks_fallbacks.py
Normal file
242
tests/test_litellm/proxy/auth/test_model_checks_fallbacks.py
Normal file
|
|
@ -0,0 +1,242 @@
|
|||
import pytest
|
||||
from unittest.mock import Mock, patch
|
||||
|
||||
|
||||
def create_mock_router(
|
||||
fallbacks=None, context_window_fallbacks=None, content_policy_fallbacks=None
|
||||
):
|
||||
"""Helper function to create a mock router with fallback configurations."""
|
||||
router = Mock()
|
||||
router.fallbacks = fallbacks or []
|
||||
router.context_window_fallbacks = context_window_fallbacks or []
|
||||
router.content_policy_fallbacks = content_policy_fallbacks or []
|
||||
return router
|
||||
|
||||
|
||||
def test_no_router_returns_empty_list():
|
||||
"""Test that None router returns empty list."""
|
||||
from litellm.proxy.auth.model_checks import get_all_fallbacks
|
||||
|
||||
result = get_all_fallbacks("claude-4-sonnet", llm_router=None)
|
||||
assert result == []
|
||||
|
||||
|
||||
def test_no_fallbacks_config_returns_empty_list():
|
||||
"""Test that empty fallbacks config returns empty list."""
|
||||
from litellm.proxy.auth.model_checks import get_all_fallbacks
|
||||
|
||||
router = create_mock_router(fallbacks=[])
|
||||
result = get_all_fallbacks("claude-4-sonnet", llm_router=router)
|
||||
assert result == []
|
||||
|
||||
|
||||
def test_model_with_fallbacks_returns_complete_list():
|
||||
"""Test that model with fallbacks returns complete fallback list."""
|
||||
from litellm.proxy.auth.model_checks import get_all_fallbacks
|
||||
|
||||
fallbacks_config = [
|
||||
{"claude-4-sonnet": ["bedrock-claude-sonnet-4", "google-claude-sonnet-4"]}
|
||||
]
|
||||
router = create_mock_router(fallbacks=fallbacks_config)
|
||||
|
||||
with patch(
|
||||
'litellm.proxy.auth.model_checks.get_fallback_model_group'
|
||||
) as mock_get_fallback:
|
||||
mock_get_fallback.return_value = (
|
||||
["bedrock-claude-sonnet-4", "google-claude-sonnet-4"], None
|
||||
)
|
||||
|
||||
result = get_all_fallbacks("claude-4-sonnet", llm_router=router)
|
||||
assert result == ["bedrock-claude-sonnet-4", "google-claude-sonnet-4"]
|
||||
|
||||
|
||||
def test_model_without_fallbacks_returns_empty_list():
|
||||
"""Test that model without fallbacks returns empty list."""
|
||||
from litellm.proxy.auth.model_checks import get_all_fallbacks
|
||||
|
||||
fallbacks_config = [
|
||||
{"claude-4-sonnet": ["bedrock-claude-sonnet-4", "google-claude-sonnet-4"]}
|
||||
]
|
||||
router = create_mock_router(fallbacks=fallbacks_config)
|
||||
|
||||
with patch(
|
||||
'litellm.proxy.auth.model_checks.get_fallback_model_group'
|
||||
) as mock_get_fallback:
|
||||
mock_get_fallback.return_value = (None, None)
|
||||
|
||||
result = get_all_fallbacks("bedrock-claude-sonnet-4", llm_router=router)
|
||||
assert result == []
|
||||
|
||||
|
||||
def test_general_fallback_type():
|
||||
"""Test general fallback type uses router.fallbacks."""
|
||||
from litellm.proxy.auth.model_checks import get_all_fallbacks
|
||||
|
||||
fallbacks_config = [
|
||||
{"claude-4-sonnet": ["bedrock-claude-sonnet-4"]}
|
||||
]
|
||||
router = create_mock_router(fallbacks=fallbacks_config)
|
||||
|
||||
with patch(
|
||||
'litellm.proxy.auth.model_checks.get_fallback_model_group'
|
||||
) as mock_get_fallback:
|
||||
mock_get_fallback.return_value = (["bedrock-claude-sonnet-4"], None)
|
||||
|
||||
result = get_all_fallbacks("claude-4-sonnet", llm_router=router, fallback_type="general")
|
||||
assert result == ["bedrock-claude-sonnet-4"]
|
||||
|
||||
# Verify it used the general fallbacks config
|
||||
mock_get_fallback.assert_called_once_with(
|
||||
fallbacks=fallbacks_config,
|
||||
model_group="claude-4-sonnet"
|
||||
)
|
||||
|
||||
|
||||
def test_context_window_fallback_type():
|
||||
"""Test context_window fallback type uses router.context_window_fallbacks."""
|
||||
from litellm.proxy.auth.model_checks import get_all_fallbacks
|
||||
|
||||
context_fallbacks_config = [
|
||||
{"gpt-4": ["gpt-3.5-turbo"]}
|
||||
]
|
||||
router = create_mock_router(context_window_fallbacks=context_fallbacks_config)
|
||||
|
||||
with patch(
|
||||
'litellm.proxy.auth.model_checks.get_fallback_model_group'
|
||||
) as mock_get_fallback:
|
||||
mock_get_fallback.return_value = (["gpt-3.5-turbo"], None)
|
||||
|
||||
result = get_all_fallbacks("gpt-4", llm_router=router, fallback_type="context_window")
|
||||
assert result == ["gpt-3.5-turbo"]
|
||||
|
||||
# Verify it used the context window fallbacks config
|
||||
mock_get_fallback.assert_called_once_with(
|
||||
fallbacks=context_fallbacks_config,
|
||||
model_group="gpt-4"
|
||||
)
|
||||
|
||||
|
||||
def test_content_policy_fallback_type():
|
||||
"""Test content_policy fallback type uses router.content_policy_fallbacks."""
|
||||
from litellm.proxy.auth.model_checks import get_all_fallbacks
|
||||
|
||||
content_fallbacks_config = [
|
||||
{"claude-4": ["claude-3"]}
|
||||
]
|
||||
router = create_mock_router(content_policy_fallbacks=content_fallbacks_config)
|
||||
|
||||
with patch(
|
||||
'litellm.proxy.auth.model_checks.get_fallback_model_group'
|
||||
) as mock_get_fallback:
|
||||
mock_get_fallback.return_value = (["claude-3"], None)
|
||||
|
||||
result = get_all_fallbacks("claude-4", llm_router=router, fallback_type="content_policy")
|
||||
assert result == ["claude-3"]
|
||||
|
||||
# Verify it used the content policy fallbacks config
|
||||
mock_get_fallback.assert_called_once_with(
|
||||
fallbacks=content_fallbacks_config,
|
||||
model_group="claude-4"
|
||||
)
|
||||
|
||||
|
||||
def test_invalid_fallback_type_returns_empty_list():
|
||||
"""Test that invalid fallback type returns empty list and logs warning."""
|
||||
from litellm.proxy.auth.model_checks import get_all_fallbacks
|
||||
|
||||
router = create_mock_router(fallbacks=[])
|
||||
|
||||
with patch('litellm.proxy.auth.model_checks.verbose_proxy_logger') as mock_logger:
|
||||
result = get_all_fallbacks("claude-4-sonnet", llm_router=router, fallback_type="invalid")
|
||||
|
||||
assert result == []
|
||||
mock_logger.warning.assert_called_once_with("Unknown fallback_type: invalid")
|
||||
|
||||
|
||||
def test_exception_handling_returns_empty_list():
|
||||
"""Test that exceptions are handled gracefully and return empty list."""
|
||||
from litellm.proxy.auth.model_checks import get_all_fallbacks
|
||||
|
||||
router = create_mock_router(fallbacks=[{"claude-4-sonnet": ["fallback"]}])
|
||||
|
||||
with patch(
|
||||
'litellm.proxy.auth.model_checks.get_fallback_model_group'
|
||||
) as mock_get_fallback:
|
||||
mock_get_fallback.side_effect = Exception("Test exception")
|
||||
|
||||
with patch('litellm.proxy.auth.model_checks.verbose_proxy_logger') as mock_logger:
|
||||
result = get_all_fallbacks("claude-4-sonnet", llm_router=router)
|
||||
|
||||
assert result == []
|
||||
mock_logger.error.assert_called_once()
|
||||
error_call_args = mock_logger.error.call_args[0][0]
|
||||
assert "Error getting fallbacks for model claude-4-sonnet" in error_call_args
|
||||
|
||||
|
||||
def test_multiple_fallbacks_complete_list():
|
||||
"""Test model with multiple fallbacks returns the complete list."""
|
||||
from litellm.proxy.auth.model_checks import get_all_fallbacks
|
||||
|
||||
fallbacks_config = [
|
||||
{"gpt-4": ["gpt-4-turbo", "gpt-3.5-turbo", "claude-3-haiku"]}
|
||||
]
|
||||
router = create_mock_router(fallbacks=fallbacks_config)
|
||||
|
||||
with patch(
|
||||
'litellm.proxy.auth.model_checks.get_fallback_model_group'
|
||||
) as mock_get_fallback:
|
||||
mock_get_fallback.return_value = (["gpt-4-turbo", "gpt-3.5-turbo", "claude-3-haiku"], None)
|
||||
|
||||
result = get_all_fallbacks("gpt-4", llm_router=router)
|
||||
assert result == ["gpt-4-turbo", "gpt-3.5-turbo", "claude-3-haiku"]
|
||||
|
||||
|
||||
def test_wildcard_and_specific_fallbacks():
|
||||
"""Test fallbacks with wildcard and specific model configurations."""
|
||||
from litellm.proxy.auth.model_checks import get_all_fallbacks
|
||||
|
||||
fallbacks_config = [
|
||||
{"*": ["gpt-3.5-turbo"]},
|
||||
{"claude-4-sonnet": ["bedrock-claude-sonnet-4", "google-claude-sonnet-4"]}
|
||||
]
|
||||
router = create_mock_router(fallbacks=fallbacks_config)
|
||||
|
||||
with patch(
|
||||
'litellm.proxy.auth.model_checks.get_fallback_model_group'
|
||||
) as mock_get_fallback:
|
||||
# Test specific model fallbacks
|
||||
mock_get_fallback.return_value = (
|
||||
["bedrock-claude-sonnet-4", "google-claude-sonnet-4"], None
|
||||
)
|
||||
result = get_all_fallbacks("claude-4-sonnet", llm_router=router)
|
||||
assert result == ["bedrock-claude-sonnet-4", "google-claude-sonnet-4"]
|
||||
|
||||
# Test wildcard fallbacks
|
||||
mock_get_fallback.return_value = (["gpt-3.5-turbo"], 0)
|
||||
result = get_all_fallbacks("some-unknown-model", llm_router=router)
|
||||
assert result == ["gpt-3.5-turbo"]
|
||||
|
||||
|
||||
def test_default_fallback_type_is_general():
|
||||
"""Test that default fallback_type is 'general'."""
|
||||
from litellm.proxy.auth.model_checks import get_all_fallbacks
|
||||
|
||||
fallbacks_config = [
|
||||
{"claude-4-sonnet": ["bedrock-claude-sonnet-4"]}
|
||||
]
|
||||
router = create_mock_router(fallbacks=fallbacks_config)
|
||||
|
||||
with patch(
|
||||
'litellm.proxy.auth.model_checks.get_fallback_model_group'
|
||||
) as mock_get_fallback:
|
||||
mock_get_fallback.return_value = (["bedrock-claude-sonnet-4"], None)
|
||||
|
||||
# Call without specifying fallback_type
|
||||
result = get_all_fallbacks("claude-4-sonnet", llm_router=router)
|
||||
|
||||
# Should use general fallbacks (router.fallbacks)
|
||||
mock_get_fallback.assert_called_once_with(
|
||||
fallbacks=fallbacks_config,
|
||||
model_group="claude-4-sonnet"
|
||||
)
|
||||
assert result == ["bedrock-claude-sonnet-4"]
|
||||
Loading…
Add table
Reference in a new issue