mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
fix(router): apply deployment kwargs and rpm semaphore in _aspeech
Bring _aspeech fully in line with _atranscription: call _update_kwargs_with_deployment so deployment metadata, model_info, timeout, and default litellm params flow into the request, and wrap the litellm.aspeech call with the max_parallel_requests semaphore plus async_routing_strategy_pre_call_checks so TTS respects rpm limits the same way the other router endpoints do Also add a unit test that exercises _aspeech directly and asserts the deployment metadata reaches the underlying call
This commit is contained in:
parent
031cc1b887
commit
467c03392c
2 changed files with 55 additions and 9 deletions
|
|
@ -4067,28 +4067,23 @@ class Router:
|
|||
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,
|
||||
)
|
||||
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)
|
||||
|
||||
self._update_kwargs_with_deployment(deployment=deployment, kwargs=kwargs)
|
||||
data = deployment["litellm_params"].copy()
|
||||
model_client = self._get_async_openai_model_client(
|
||||
deployment=deployment,
|
||||
kwargs=kwargs,
|
||||
)
|
||||
|
||||
self.total_calls[model_name] += 1
|
||||
response = await litellm.aspeech(
|
||||
response = litellm.aspeech(
|
||||
**{
|
||||
**data,
|
||||
"input": input,
|
||||
|
|
@ -4098,6 +4093,31 @@ class Router:
|
|||
}
|
||||
)
|
||||
|
||||
### 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"
|
||||
|
|
|
|||
|
|
@ -262,6 +262,32 @@ async def test_aspeech_success_returns_response():
|
|||
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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue