Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
96 changes: 66 additions & 30 deletions litellm/router.py
Original file line number Diff line number Diff line change
Expand Up @@ -4042,57 +4042,93 @@ async def aspeech(self, model: str, input: str, voice: str, **kwargs):
```
"""
try:
kwargs["model"] = model
kwargs["input"] = input
kwargs["voice"] = voice
kwargs["original_function"] = self._aspeech
self._update_kwargs_before_fallbacks(model=model, kwargs=kwargs)
response = await self.async_function_with_fallbacks(**kwargs)

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 _aspeech(self, model: str, input: str, voice: str, **kwargs):
model_name = model
try:
verbose_router_logger.debug(
f"Inside _aspeech()- model: {model}; kwargs: {kwargs}"
)
parent_otel_span = _get_parent_otel_span_from_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,
)
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)

potential_model_client = self._get_client(
deployment=deployment, kwargs=kwargs, client_type="async"
self._update_kwargs_with_deployment(deployment=deployment, kwargs=kwargs)
data = deployment["litellm_params"].copy()
model_client = self._get_async_openai_model_client(
Comment thread
greptile-apps[bot] marked this conversation as resolved.
deployment=deployment,
kwargs=kwargs,
)
# 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(
self.total_calls[model_name] += 1
response = litellm.aspeech(
**{
**data,
"input": input,
"voice": voice,
"client": model_client,
**kwargs,
}
)
Comment thread
greptile-apps[bot] marked this conversation as resolved.

### CONCURRENCY-SAFE RPM CHECKS ###
rpm_semaphore = self._get_client(
deployment=deployment,
kwargs=kwargs,
client_type="max_parallel_requests",
)

if rpm_semaphore is not None and isinstance(
rpm_semaphore, asyncio.Semaphore
):
async with rpm_semaphore:
"""
- Check rpm limits before making the call
- If allowed, increment the rpm limit (allows global value to be updated, concurrency-safe)
"""
await self.async_routing_strategy_pre_call_checks(
deployment=deployment, parent_otel_span=parent_otel_span
)
response = await response
else:
await self.async_routing_strategy_pre_call_checks(
deployment=deployment, parent_otel_span=parent_otel_span
)
response = await response

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:
asyncio.create_task(
send_llm_exception_alert(
litellm_router_instance=self,
request_kwargs=kwargs,
error_traceback_str=traceback.format_exc(),
original_exception=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):
Expand Down
90 changes: 90 additions & 0 deletions tests/router_unit_tests/test_router_endpoints.py
Original file line number Diff line number Diff line change
Expand Up @@ -198,6 +198,96 @@ 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_aspeech_sets_deployment_metadata():
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
metadata = mock_aspeech.call_args.kwargs["metadata"]
assert metadata["deployment"] == "openai/tts-1"
assert metadata["deployment_model_name"] == "tts"
assert metadata["model_info"]["id"] is not None


@pytest.mark.asyncio()
async def test_rerank_endpoint(model_list):
from litellm.types.utils import RerankResponse
Expand Down
Loading