mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
Resolve model group alias on Auth + /v1/messages Fallback support (#12440)
* fix(auth_checks.py): resolve a model group alias when key has access to underlying model Fixes LIT-293 * feat(anthropic/): add mock_response to anthropic /v1/messages makes it easy to test fallback logic * fix(router.py): support fallbacks on /v1/messages adds working fallbacks on generic api route * refactor(router.py): point _ageneric_api_call_with_fallbacks to updated function * test: add unit test for new helper on router * fix(router.py): use correct metadata variable name * fix(router.py): use correct metadata field * docs(config_settings.md): document new param
This commit is contained in:
parent
0730f61127
commit
07e8609edb
7 changed files with 365 additions and 39 deletions
|
|
@ -597,6 +597,7 @@ router_settings:
|
|||
| OTEL_EXPORTER | Exporter type for OpenTelemetry
|
||||
| OTEL_EXPORTER_OTLP_PROTOCOL | Exporter type for OpenTelemetry
|
||||
| OTEL_HEADERS | Headers for OpenTelemetry requests
|
||||
| OTEL_MODEL_ID | Model ID for OpenTelemetry tracing
|
||||
| OTEL_EXPORTER_OTLP_HEADERS | Headers for OpenTelemetry requests
|
||||
| OTEL_SERVICE_NAME | Service name identifier for OpenTelemetry
|
||||
| OTEL_TRACER_NAME | Tracer name for OpenTelemetry tracing
|
||||
|
|
|
|||
|
|
@ -25,7 +25,7 @@ from litellm.types.router import GenericLiteLLMParams
|
|||
from litellm.utils import ProviderConfigManager, client
|
||||
|
||||
from ..adapters.handler import LiteLLMMessagesToCompletionTransformationHandler
|
||||
from .utils import AnthropicMessagesRequestUtils
|
||||
from .utils import AnthropicMessagesRequestUtils, mock_response
|
||||
|
||||
####### ENVIRONMENT VARIABLES ###################
|
||||
# Initialize any necessary instances or variables here
|
||||
|
|
@ -92,6 +92,7 @@ async def anthropic_messages(
|
|||
response = init_response
|
||||
return response
|
||||
|
||||
|
||||
def validate_anthropic_api_metadata(metadata: Optional[Dict] = None) -> Optional[Dict]:
|
||||
"""
|
||||
Validate Anthropic API metadata - This is done to ensure only allowed `metadata` fields are passed to Anthropic API
|
||||
|
|
@ -103,6 +104,7 @@ def validate_anthropic_api_metadata(metadata: Optional[Dict] = None) -> Optional
|
|||
anthropic_metadata_obj = AnthropicMetadata(**metadata)
|
||||
return anthropic_metadata_obj.model_dump(exclude_none=True)
|
||||
|
||||
|
||||
def anthropic_messages_handler(
|
||||
max_tokens: int,
|
||||
messages: List[Dict],
|
||||
|
|
@ -131,12 +133,14 @@ def anthropic_messages_handler(
|
|||
Makes Anthropic `/v1/messages` API calls In the Anthropic API Spec
|
||||
"""
|
||||
from litellm.types.utils import LlmProviders
|
||||
|
||||
metadata = validate_anthropic_api_metadata(metadata)
|
||||
|
||||
local_vars = locals()
|
||||
is_async = kwargs.pop("is_async", False)
|
||||
# Use provided client or create a new one
|
||||
litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore
|
||||
|
||||
litellm_params = GenericLiteLLMParams(
|
||||
**kwargs,
|
||||
api_key=api_key,
|
||||
|
|
@ -155,6 +159,15 @@ def anthropic_messages_handler(
|
|||
api_key=litellm_params.api_key,
|
||||
)
|
||||
|
||||
if litellm_params.mock_response and isinstance(litellm_params.mock_response, str):
|
||||
|
||||
return mock_response(
|
||||
model=model,
|
||||
messages=messages,
|
||||
max_tokens=max_tokens,
|
||||
mock_response=litellm_params.mock_response,
|
||||
)
|
||||
|
||||
anthropic_messages_provider_config: Optional[BaseAnthropicMessagesConfig] = None
|
||||
|
||||
if custom_llm_provider is not None and custom_llm_provider in [
|
||||
|
|
|
|||
|
|
@ -1,6 +1,9 @@
|
|||
from typing import Any, Dict, cast, get_type_hints
|
||||
from typing import Any, Dict, List, cast, get_type_hints
|
||||
|
||||
from litellm.types.llms.anthropic import AnthropicMessagesRequestOptionalParams
|
||||
from litellm.types.llms.anthropic_messages.anthropic_response import (
|
||||
AnthropicMessagesResponse,
|
||||
)
|
||||
|
||||
|
||||
class AnthropicMessagesRequestUtils:
|
||||
|
|
@ -22,3 +25,51 @@ class AnthropicMessagesRequestUtils:
|
|||
k: v for k, v in params.items() if k in valid_keys and v is not None
|
||||
}
|
||||
return cast(AnthropicMessagesRequestOptionalParams, filtered_params)
|
||||
|
||||
|
||||
def mock_response(
|
||||
model: str,
|
||||
messages: List[Dict],
|
||||
max_tokens: int,
|
||||
mock_response: str = "Hi! My name is Claude.",
|
||||
**kwargs,
|
||||
) -> AnthropicMessagesResponse:
|
||||
"""
|
||||
Mock response for Anthropic messages
|
||||
"""
|
||||
from litellm.exceptions import (
|
||||
ContextWindowExceededError,
|
||||
InternalServerError,
|
||||
RateLimitError,
|
||||
)
|
||||
|
||||
if mock_response == "litellm.InternalServerError":
|
||||
raise InternalServerError(
|
||||
message="this is a mock internal server error",
|
||||
llm_provider="anthropic",
|
||||
model=model,
|
||||
)
|
||||
elif mock_response == "litellm.ContextWindowExceededError":
|
||||
raise ContextWindowExceededError(
|
||||
message="this is a mock context window exceeded error",
|
||||
llm_provider="anthropic",
|
||||
model=model,
|
||||
)
|
||||
elif mock_response == "litellm.RateLimitError":
|
||||
raise RateLimitError(
|
||||
message="this is a mock rate limit error",
|
||||
llm_provider="anthropic",
|
||||
model=model,
|
||||
)
|
||||
return AnthropicMessagesResponse(
|
||||
**{
|
||||
"content": [{"text": mock_response, "type": "text"}],
|
||||
"id": "msg_013Zva2CMHLNnXjNJJKqJ2EF",
|
||||
"model": "claude-sonnet-4-20250514",
|
||||
"role": "assistant",
|
||||
"stop_reason": "end_turn",
|
||||
"stop_sequence": None,
|
||||
"type": "message",
|
||||
"usage": {"input_tokens": 2095, "output_tokens": 503},
|
||||
}
|
||||
)
|
||||
|
|
|
|||
|
|
@ -1191,6 +1191,10 @@ def _can_object_call_model(
|
|||
|
||||
if model in litellm.model_alias_map:
|
||||
model = litellm.model_alias_map[model]
|
||||
elif llm_router and model in llm_router.model_group_alias:
|
||||
_model = llm_router._get_model_from_alias(model)
|
||||
if _model:
|
||||
model = _model
|
||||
|
||||
## check if model in allowed model names
|
||||
from collections import defaultdict
|
||||
|
|
@ -1199,6 +1203,7 @@ def _can_object_call_model(
|
|||
|
||||
if llm_router:
|
||||
access_groups = llm_router.get_model_access_groups(model_name=model)
|
||||
|
||||
if (
|
||||
len(access_groups) > 0 and llm_router is not None
|
||||
): # check if token contains any model access groups
|
||||
|
|
@ -1211,8 +1216,6 @@ def _can_object_call_model(
|
|||
# Filter out models that are access_groups
|
||||
filtered_models = [m for m in models if m not in access_groups]
|
||||
|
||||
verbose_proxy_logger.debug(f"model: {model}; allowed_models: {filtered_models}")
|
||||
|
||||
if _model_in_team_aliases(model=model, team_model_aliases=team_model_aliases):
|
||||
return True
|
||||
|
||||
|
|
|
|||
|
|
@ -2468,30 +2468,44 @@ class Router:
|
|||
self, model: str, original_function: Callable, **kwargs
|
||||
):
|
||||
"""
|
||||
Make a generic LLM API call through the router, this allows you to use retries/fallbacks with litellm router
|
||||
|
||||
Args:
|
||||
model: The model to use
|
||||
handler_function: The handler function to call (e.g., litellm.anthropic_messages)
|
||||
**kwargs: Additional arguments to pass to the handler function
|
||||
|
||||
Returns:
|
||||
The response from the handler function
|
||||
Helper function to make a generic LLM API call through the router, this allows you to use retries/fallbacks with litellm router
|
||||
"""
|
||||
handler_name = original_function.__name__
|
||||
function_name = "_ageneric_api_call_with_fallbacks"
|
||||
passthrough_on_no_deployment = kwargs.pop("passthrough_on_no_deployment", False)
|
||||
self._update_kwargs_before_fallbacks(
|
||||
model=model,
|
||||
kwargs=kwargs,
|
||||
metadata_variable_name=_get_router_metadata_variable_name(
|
||||
function_name=function_name
|
||||
),
|
||||
)
|
||||
try:
|
||||
verbose_router_logger.debug(
|
||||
f"Inside _ageneric_api_call() - handler: {handler_name}, model: {model}; kwargs: {kwargs}"
|
||||
kwargs["model"] = model
|
||||
kwargs["original_generic_function"] = original_function
|
||||
kwargs["original_function"] = self._ageneric_api_call_with_fallbacks_helper
|
||||
self._update_kwargs_before_fallbacks(
|
||||
model=model, kwargs=kwargs, metadata_variable_name="litellm_metadata"
|
||||
)
|
||||
verbose_router_logger.debug(
|
||||
f"Inside ageneric_api_call_with_fallbacks() - model: {model}; kwargs: {kwargs}"
|
||||
)
|
||||
response = await self.async_function_with_fallbacks(**kwargs)
|
||||
return response
|
||||
|
||||
return response
|
||||
except Exception as e:
|
||||
asyncio.create_task(
|
||||
send_llm_exception_alert(
|
||||
litellm_router_instance=self,
|
||||
request_kwargs=kwargs,
|
||||
error_traceback_str=traceback.format_exc(),
|
||||
original_exception=e,
|
||||
)
|
||||
)
|
||||
raise e
|
||||
|
||||
async def _ageneric_api_call_with_fallbacks_helper(
|
||||
self, model: str, original_generic_function: Callable, **kwargs
|
||||
):
|
||||
"""
|
||||
Helper function to make a generic LLM API call through the router, this allows you to use retries/fallbacks with litellm router
|
||||
"""
|
||||
|
||||
passthrough_on_no_deployment = kwargs.pop("passthrough_on_no_deployment", False)
|
||||
function_name = "_ageneric_api_call_with_fallbacks"
|
||||
try:
|
||||
|
||||
parent_otel_span = _get_parent_otel_span_from_kwargs(kwargs)
|
||||
try:
|
||||
deployment = await self.async_get_available_deployment(
|
||||
|
|
@ -2502,7 +2516,7 @@ class Router:
|
|||
)
|
||||
except Exception as e:
|
||||
if passthrough_on_no_deployment:
|
||||
return await original_function(model=model, **kwargs)
|
||||
return await original_generic_function(model=model, **kwargs)
|
||||
raise e
|
||||
|
||||
self._update_kwargs_with_deployment(
|
||||
|
|
@ -2515,7 +2529,7 @@ class Router:
|
|||
|
||||
### get custom
|
||||
|
||||
response = original_function(
|
||||
response = original_generic_function(
|
||||
**{
|
||||
**data,
|
||||
"caching": self.cache_responses,
|
||||
|
|
@ -2549,12 +2563,12 @@ class Router:
|
|||
|
||||
self.success_calls[model_name] += 1
|
||||
verbose_router_logger.info(
|
||||
f"{handler_name}(model={model_name})\033[32m 200 OK\033[0m"
|
||||
f"ageneric_api_call_with_fallbacks(model={model_name})\033[32m 200 OK\033[0m"
|
||||
)
|
||||
return response
|
||||
except Exception as e:
|
||||
verbose_router_logger.info(
|
||||
f"{handler_name}(model={model})\033[31m Exception {str(e)}\033[0m"
|
||||
f"ageneric_api_call_with_fallbacks(model={model})\033[31m Exception {str(e)}\033[0m"
|
||||
)
|
||||
if model is not None:
|
||||
self.fail_calls[model] += 1
|
||||
|
|
@ -4243,6 +4257,9 @@ class Router:
|
|||
When a retry or fallback happens, log the details of the just failed model call - similar to Sentry breadcrumbing
|
||||
"""
|
||||
try:
|
||||
_metadata_var = (
|
||||
"litellm_metadata" if "litellm_metadata" in kwargs else "metadata"
|
||||
)
|
||||
# Log failed model as the previous model
|
||||
previous_model = {
|
||||
"exception_type": type(e).__name__,
|
||||
|
|
@ -4254,11 +4271,11 @@ class Router:
|
|||
) in (
|
||||
kwargs.items()
|
||||
): # log everything in kwargs except the old previous_models value - prevent nesting
|
||||
if k not in ["metadata", "messages", "original_function"]:
|
||||
if k not in [_metadata_var, "messages", "original_function"]:
|
||||
previous_model[k] = v
|
||||
elif k == "metadata" and isinstance(v, dict):
|
||||
previous_model["metadata"] = {} # type: ignore
|
||||
for metadata_k, metadata_v in kwargs["metadata"].items():
|
||||
elif k == _metadata_var and isinstance(v, dict):
|
||||
previous_model[_metadata_var] = {} # type: ignore
|
||||
for metadata_k, metadata_v in kwargs[_metadata_var].items():
|
||||
if metadata_k != "previous_models":
|
||||
previous_model[k][metadata_k] = metadata_v # type: ignore
|
||||
|
||||
|
|
@ -4267,7 +4284,7 @@ class Router:
|
|||
self.previous_models.pop(0)
|
||||
|
||||
self.previous_models.append(previous_model)
|
||||
kwargs["metadata"]["previous_models"] = self.previous_models
|
||||
kwargs[_metadata_var]["previous_models"] = self.previous_models
|
||||
return kwargs
|
||||
except Exception as e:
|
||||
raise e
|
||||
|
|
|
|||
|
|
@ -334,3 +334,39 @@ async def test_vector_store_access_check_with_permissions():
|
|||
)
|
||||
|
||||
assert exc_info.value.type == ProxyErrorTypes.key_vector_store_access_denied
|
||||
|
||||
|
||||
def test_can_object_call_model_with_alias():
|
||||
"""Test that can_object_call_model works with model aliases"""
|
||||
from litellm import Router
|
||||
from litellm.proxy.auth.auth_checks import _can_object_call_model
|
||||
|
||||
model = "[ip-approved] gpt-4o"
|
||||
llm_router = Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "gpt-3.5-turbo",
|
||||
"litellm_params": {
|
||||
"model": "gpt-3.5-turbo",
|
||||
"api_key": "test-api-key",
|
||||
},
|
||||
}
|
||||
],
|
||||
model_group_alias={
|
||||
"[ip-approved] gpt-4o": {
|
||||
"model": "gpt-3.5-turbo",
|
||||
"hidden": True,
|
||||
},
|
||||
},
|
||||
)
|
||||
|
||||
result = _can_object_call_model(
|
||||
model=model,
|
||||
llm_router=llm_router,
|
||||
models=["gpt-3.5-turbo"],
|
||||
team_model_aliases=None,
|
||||
object_type="key",
|
||||
fallback_depth=0,
|
||||
)
|
||||
|
||||
print(result)
|
||||
|
|
|
|||
|
|
@ -37,7 +37,7 @@ def test_update_kwargs_does_not_mutate_defaults_and_merges_metadata():
|
|||
"metadata": {"baz": 123},
|
||||
}
|
||||
original = copy.deepcopy(router.default_litellm_params)
|
||||
kwargs = {}
|
||||
kwargs: dict = {}
|
||||
|
||||
# invoke the helper
|
||||
router._update_kwargs_with_default_litellm_params(
|
||||
|
|
@ -342,7 +342,7 @@ def test_arouter_ignore_invalid_deployments():
|
|||
router.upsert_deployment(
|
||||
Deployment(
|
||||
model_name="gpt-3.5-turbo",
|
||||
litellm_params={"model": "my-bad-model"},
|
||||
litellm_params={"model": "my-bad-model"}, # type: ignore
|
||||
model_info={"tpm": 1000, "rpm": 1000},
|
||||
)
|
||||
)
|
||||
|
|
@ -468,7 +468,7 @@ async def test_arouter_filter_team_based_models():
|
|||
router.add_deployment(
|
||||
Deployment(
|
||||
model_name="gpt-3.5-turbo",
|
||||
litellm_params={"model": "gpt-3.5-turbo"},
|
||||
litellm_params={"model": "gpt-3.5-turbo"}, # type: ignore
|
||||
model_info={"tpm": 1000, "rpm": 1000},
|
||||
)
|
||||
)
|
||||
|
|
@ -661,16 +661,55 @@ def test_arouter_responses_api_bridge():
|
|||
assert mock_post.call_args.kwargs["json"]["model"] == "webinterface-o3-pro"
|
||||
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_router_v1_messages_fallbacks():
|
||||
"""
|
||||
Test that router.v1_messages_fallbacks returns the correct response
|
||||
"""
|
||||
router = litellm.Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "claude-3-5-sonnet-latest",
|
||||
"litellm_params": {
|
||||
"model": "anthropic/claude-3-5-sonnet-latest",
|
||||
"mock_response": "litellm.InternalServerError",
|
||||
},
|
||||
},
|
||||
{
|
||||
"model_name": "bedrock-claude",
|
||||
"litellm_params": {
|
||||
"model": "anthropic.claude-3-5-sonnet-20240620-v1:0",
|
||||
"mock_response": "Hello, world I am a fallback!",
|
||||
},
|
||||
},
|
||||
],
|
||||
fallbacks=[
|
||||
{"claude-3-5-sonnet-latest": ["bedrock-claude"]},
|
||||
],
|
||||
)
|
||||
|
||||
result = await router.aanthropic_messages(
|
||||
model="claude-3-5-sonnet-latest",
|
||||
messages=[{"role": "user", "content": "Hello, world!"}],
|
||||
max_tokens=256,
|
||||
)
|
||||
assert result is not None
|
||||
|
||||
print(result)
|
||||
assert result["content"][0]["text"] == "Hello, world I am a fallback!"
|
||||
|
||||
|
||||
def test_add_invalid_provider_to_router():
|
||||
"""
|
||||
Test that router.add_deployment raises an error if the provider is invalid
|
||||
"""
|
||||
from litellm.types.router import Deployment
|
||||
|
||||
|
||||
router = litellm.Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "gpt-3.5-turbo",
|
||||
"model_name": "gpt-3.5-turbo",
|
||||
"litellm_params": {"model": "gpt-3.5-turbo"},
|
||||
}
|
||||
],
|
||||
|
|
@ -688,3 +727,169 @@ def test_add_invalid_provider_to_router():
|
|||
)
|
||||
|
||||
assert router.pattern_router.patterns == {}
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_router_ageneric_api_call_with_fallbacks_helper():
|
||||
"""
|
||||
Test the _ageneric_api_call_with_fallbacks_helper method with various scenarios
|
||||
"""
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
router = litellm.Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "gpt-3.5-turbo",
|
||||
"litellm_params": {
|
||||
"model": "gpt-3.5-turbo",
|
||||
"api_key": "test-key",
|
||||
"api_base": "https://api.openai.com/v1",
|
||||
},
|
||||
"model_info": {
|
||||
"tpm": 1000,
|
||||
"rpm": 1000,
|
||||
},
|
||||
},
|
||||
],
|
||||
)
|
||||
|
||||
# Test 1: Successful call
|
||||
async def mock_generic_function(**kwargs):
|
||||
return {"result": "success", "model": kwargs.get("model")}
|
||||
|
||||
with patch.object(router, "async_get_available_deployment") as mock_get_deployment:
|
||||
mock_get_deployment.return_value = {
|
||||
"model_name": "gpt-3.5-turbo",
|
||||
"litellm_params": {
|
||||
"model": "gpt-3.5-turbo",
|
||||
"api_key": "test-key",
|
||||
"api_base": "https://api.openai.com/v1",
|
||||
},
|
||||
}
|
||||
|
||||
with patch.object(
|
||||
router, "_update_kwargs_with_deployment"
|
||||
) as mock_update_kwargs:
|
||||
with patch.object(
|
||||
router, "async_routing_strategy_pre_call_checks"
|
||||
) as mock_pre_call_checks:
|
||||
with patch.object(
|
||||
router, "_get_client", return_value=None
|
||||
) as mock_get_client:
|
||||
result = await router._ageneric_api_call_with_fallbacks_helper(
|
||||
model="gpt-3.5-turbo",
|
||||
original_generic_function=mock_generic_function,
|
||||
messages=[{"role": "user", "content": "test"}],
|
||||
)
|
||||
|
||||
assert result is not None
|
||||
assert result["result"] == "success"
|
||||
mock_get_deployment.assert_called_once()
|
||||
mock_update_kwargs.assert_called_once()
|
||||
mock_pre_call_checks.assert_called_once()
|
||||
|
||||
# Test 2: Passthrough on no deployment (success case)
|
||||
async def mock_passthrough_function(**kwargs):
|
||||
return {"result": "passthrough", "model": kwargs.get("model")}
|
||||
|
||||
with patch.object(router, "async_get_available_deployment") as mock_get_deployment:
|
||||
mock_get_deployment.side_effect = Exception("No deployment available")
|
||||
|
||||
result = await router._ageneric_api_call_with_fallbacks_helper(
|
||||
model="gpt-3.5-turbo",
|
||||
original_generic_function=mock_passthrough_function,
|
||||
passthrough_on_no_deployment=True,
|
||||
messages=[{"role": "user", "content": "test"}],
|
||||
)
|
||||
|
||||
assert result is not None
|
||||
assert result["result"] == "passthrough"
|
||||
assert result["model"] == "gpt-3.5-turbo"
|
||||
|
||||
# Test 3: No deployment available and passthrough=False (should raise exception)
|
||||
with patch.object(router, "async_get_available_deployment") as mock_get_deployment:
|
||||
mock_get_deployment.side_effect = Exception("No deployment available")
|
||||
|
||||
with pytest.raises(Exception) as exc_info:
|
||||
await router._ageneric_api_call_with_fallbacks_helper(
|
||||
model="gpt-3.5-turbo",
|
||||
original_generic_function=mock_generic_function,
|
||||
passthrough_on_no_deployment=False,
|
||||
messages=[{"role": "user", "content": "test"}],
|
||||
)
|
||||
|
||||
assert "No deployment available" in str(exc_info.value)
|
||||
|
||||
# Test 4: Test with semaphore (rate limiting)
|
||||
import asyncio
|
||||
|
||||
async def mock_semaphore_function(**kwargs):
|
||||
return {"result": "semaphore_success", "model": kwargs.get("model")}
|
||||
|
||||
with patch.object(router, "async_get_available_deployment") as mock_get_deployment:
|
||||
mock_get_deployment.return_value = {
|
||||
"model_name": "gpt-3.5-turbo",
|
||||
"litellm_params": {
|
||||
"model": "gpt-3.5-turbo",
|
||||
"api_key": "test-key",
|
||||
"api_base": "https://api.openai.com/v1",
|
||||
},
|
||||
}
|
||||
|
||||
mock_semaphore = asyncio.Semaphore(1)
|
||||
|
||||
with patch.object(
|
||||
router, "_update_kwargs_with_deployment"
|
||||
) as mock_update_kwargs:
|
||||
with patch.object(
|
||||
router, "_get_client", return_value=mock_semaphore
|
||||
) as mock_get_client:
|
||||
with patch.object(
|
||||
router, "async_routing_strategy_pre_call_checks"
|
||||
) as mock_pre_call_checks:
|
||||
result = await router._ageneric_api_call_with_fallbacks_helper(
|
||||
model="gpt-3.5-turbo",
|
||||
original_generic_function=mock_semaphore_function,
|
||||
messages=[{"role": "user", "content": "test"}],
|
||||
)
|
||||
|
||||
assert result is not None
|
||||
assert result["result"] == "semaphore_success"
|
||||
mock_get_client.assert_called_once()
|
||||
mock_pre_call_checks.assert_called_once()
|
||||
|
||||
# Test 5: Test call tracking (success and failure counts)
|
||||
initial_success_count = router.success_calls.get("gpt-3.5-turbo", 0)
|
||||
initial_fail_count = router.fail_calls.get("gpt-3.5-turbo", 0)
|
||||
|
||||
async def mock_failing_function(**kwargs):
|
||||
raise Exception("Mock failure")
|
||||
|
||||
with patch.object(router, "async_get_available_deployment") as mock_get_deployment:
|
||||
mock_get_deployment.return_value = {
|
||||
"model_name": "gpt-3.5-turbo",
|
||||
"litellm_params": {
|
||||
"model": "gpt-3.5-turbo",
|
||||
"api_key": "test-key",
|
||||
"api_base": "https://api.openai.com/v1",
|
||||
},
|
||||
}
|
||||
|
||||
with patch.object(
|
||||
router, "_update_kwargs_with_deployment"
|
||||
) as mock_update_kwargs:
|
||||
with patch.object(
|
||||
router, "_get_client", return_value=None
|
||||
) as mock_get_client:
|
||||
with patch.object(
|
||||
router, "async_routing_strategy_pre_call_checks"
|
||||
) as mock_pre_call_checks:
|
||||
with pytest.raises(Exception) as exc_info:
|
||||
await router._ageneric_api_call_with_fallbacks_helper(
|
||||
model="gpt-3.5-turbo",
|
||||
original_generic_function=mock_failing_function,
|
||||
messages=[{"role": "user", "content": "test"}],
|
||||
)
|
||||
|
||||
assert "Mock failure" in str(exc_info.value)
|
||||
# Check that fail_calls was incremented
|
||||
assert router.fail_calls["gpt-3.5-turbo"] == initial_fail_count + 1
|
||||
Loading…
Add table
Reference in a new issue