mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
fix(router): route aspeech through async_function_with_fallbacks
Router.aspeech selected a deployment and awaited litellm.aspeech directly, so TTS requests got no retry on failure and no failover to backup deployments; the except block only fired an exception alert and re-raised. Every other router endpoint (acompletion, aembedding, atranscription, arerank) already delegates to async_function_with_fallbacks Mirror the atranscription pattern: move deployment selection and the litellm.aspeech call into a private _aspeech method, then have the public aspeech set kwargs["original_function"] = self._aspeech and await self.async_function_with_fallbacks(**kwargs). _aspeech also picks up the shared _get_async_openai_model_client helper and the same total/success/fail call accounting the sibling endpoints use Fixes #27778.
This commit is contained in:
parent
e15b37a18e
commit
031cc1b887
2 changed files with 117 additions and 37 deletions
|
|
@ -4042,47 +4042,13 @@ class Router:
|
|||
```
|
||||
"""
|
||||
try:
|
||||
kwargs["model"] = model
|
||||
kwargs["input"] = input
|
||||
kwargs["voice"] = voice
|
||||
|
||||
deployment = await self.async_get_available_deployment(
|
||||
model=model,
|
||||
messages=[{"role": "user", "content": "prompt"}],
|
||||
specific_deployment=kwargs.pop("specific_deployment", None),
|
||||
request_kwargs=kwargs,
|
||||
)
|
||||
kwargs["original_function"] = self._aspeech
|
||||
self._update_kwargs_before_fallbacks(model=model, kwargs=kwargs)
|
||||
data = deployment["litellm_params"].copy()
|
||||
data["model"]
|
||||
for k, v in self.default_litellm_params.items():
|
||||
if (
|
||||
k not in kwargs
|
||||
): # prioritize model-specific params > default router params
|
||||
kwargs[k] = v
|
||||
elif k == "metadata":
|
||||
kwargs[k].update(v)
|
||||
response = await self.async_function_with_fallbacks(**kwargs)
|
||||
|
||||
potential_model_client = self._get_client(
|
||||
deployment=deployment, kwargs=kwargs, client_type="async"
|
||||
)
|
||||
# check if provided keys == client keys #
|
||||
dynamic_api_key = kwargs.get("api_key", None)
|
||||
if (
|
||||
dynamic_api_key is not None
|
||||
and potential_model_client is not None
|
||||
and dynamic_api_key != potential_model_client.api_key
|
||||
):
|
||||
model_client = None
|
||||
else:
|
||||
model_client = potential_model_client
|
||||
|
||||
response = await litellm.aspeech(
|
||||
**{
|
||||
**data,
|
||||
"client": model_client,
|
||||
**kwargs,
|
||||
}
|
||||
)
|
||||
return response
|
||||
except Exception as e:
|
||||
asyncio.create_task(
|
||||
|
|
@ -4095,6 +4061,56 @@ class Router:
|
|||
)
|
||||
raise e
|
||||
|
||||
async def _aspeech(self, model: str, input: str, voice: str, **kwargs):
|
||||
model_name = model
|
||||
try:
|
||||
verbose_router_logger.debug(
|
||||
f"Inside _aspeech()- model: {model}; kwargs: {kwargs}"
|
||||
)
|
||||
deployment = await self.async_get_available_deployment(
|
||||
model=model,
|
||||
messages=[{"role": "user", "content": "prompt"}],
|
||||
specific_deployment=kwargs.pop("specific_deployment", None),
|
||||
request_kwargs=kwargs,
|
||||
)
|
||||
data = deployment["litellm_params"].copy()
|
||||
for k, v in self.default_litellm_params.items():
|
||||
if (
|
||||
k not in kwargs
|
||||
): # prioritize model-specific params > default router params
|
||||
kwargs[k] = v
|
||||
elif k == "metadata":
|
||||
kwargs[k].update(v)
|
||||
|
||||
model_client = self._get_async_openai_model_client(
|
||||
deployment=deployment,
|
||||
kwargs=kwargs,
|
||||
)
|
||||
|
||||
self.total_calls[model_name] += 1
|
||||
response = await litellm.aspeech(
|
||||
**{
|
||||
**data,
|
||||
"input": input,
|
||||
"voice": voice,
|
||||
"client": model_client,
|
||||
**kwargs,
|
||||
}
|
||||
)
|
||||
|
||||
self.success_calls[model_name] += 1
|
||||
verbose_router_logger.info(
|
||||
f"litellm.aspeech(model={model_name})\033[32m 200 OK\033[0m"
|
||||
)
|
||||
return response
|
||||
except Exception as e:
|
||||
verbose_router_logger.info(
|
||||
f"litellm.aspeech(model={model_name})\033[31m Exception {str(e)}\033[0m"
|
||||
)
|
||||
if model_name is not None:
|
||||
self.fail_calls[model_name] += 1
|
||||
raise e
|
||||
|
||||
async def arerank(self, model: str, **kwargs):
|
||||
try:
|
||||
kwargs["model"] = model
|
||||
|
|
|
|||
|
|
@ -198,6 +198,70 @@ async def test_audio_speech_router(mode):
|
|||
assert test_logger.standard_logging_object["model_group"] == "tts"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_aspeech_fallbacks_on_deployment_failure():
|
||||
router = Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "tts-main",
|
||||
"litellm_params": {"model": "openai/tts-1", "api_key": "fake-key"},
|
||||
},
|
||||
{
|
||||
"model_name": "tts-backup",
|
||||
"litellm_params": {"model": "openai/tts-1-hd", "api_key": "fake-key"},
|
||||
},
|
||||
],
|
||||
fallbacks=[{"tts-main": ["tts-backup"]}],
|
||||
num_retries=0,
|
||||
)
|
||||
|
||||
called_models = []
|
||||
|
||||
async def mock_aspeech(*args, **kwargs):
|
||||
called_models.append(kwargs["model"])
|
||||
if kwargs["model"] == "openai/tts-1":
|
||||
raise litellm.InternalServerError(
|
||||
message="deployment down",
|
||||
llm_provider="openai",
|
||||
model="tts-1",
|
||||
)
|
||||
return MagicMock()
|
||||
|
||||
with patch("litellm.aspeech", side_effect=mock_aspeech):
|
||||
response = await router.aspeech(
|
||||
model="tts-main",
|
||||
input="the quick brown fox jumped over the lazy dogs",
|
||||
voice="alloy",
|
||||
)
|
||||
|
||||
assert response is not None
|
||||
assert called_models == ["openai/tts-1", "openai/tts-1-hd"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_aspeech_success_returns_response():
|
||||
router = Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "tts",
|
||||
"litellm_params": {"model": "openai/tts-1", "api_key": "fake-key"},
|
||||
},
|
||||
]
|
||||
)
|
||||
|
||||
mock_response = MagicMock()
|
||||
with patch("litellm.aspeech", return_value=mock_response) as mock_aspeech:
|
||||
response = await router.aspeech(
|
||||
model="tts",
|
||||
input="the quick brown fox jumped over the lazy dogs",
|
||||
voice="alloy",
|
||||
)
|
||||
|
||||
assert response is mock_response
|
||||
mock_aspeech.assert_called_once()
|
||||
assert mock_aspeech.call_args.kwargs["model"] == "openai/tts-1"
|
||||
|
||||
|
||||
@pytest.mark.asyncio()
|
||||
async def test_rerank_endpoint(model_list):
|
||||
from litellm.types.utils import RerankResponse
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue