Fix aembedding retry propagation

This commit is contained in:
Sambhram1 2026-05-07 22:46:37 +05:30
parent fee5900acc
commit 75be23873e
2 changed files with 35 additions and 0 deletions

View file

@ -4383,6 +4383,7 @@ class Router:
kwargs["model"] = model
kwargs["input"] = input
kwargs["original_function"] = self._aembedding
kwargs["num_retries"] = kwargs.get("num_retries", self.num_retries)
self._update_kwargs_before_fallbacks(model=model, kwargs=kwargs)
response = await self.async_function_with_fallbacks(**kwargs)
return response

View file

@ -103,6 +103,40 @@ class TestRouterEmbeddingHeaders:
assert call_kwargs["model"] == "text-embedding-ada-002"
assert "kwargs" in call_kwargs
@pytest.mark.asyncio
async def test_aembedding_passes_num_retries_to_fallbacks(self):
"""
Test that router.aembedding() seeds num_retries before fallback handling.
Regression test for #27363.
"""
model_list = [
{
"model_name": "text-embedding-ada-002",
"litellm_params": {
"model": "text-embedding-ada-002",
"api_key": "fake-key",
},
}
]
router = Router(model_list=model_list, num_retries=3)
with patch.object(
router, "async_function_with_fallbacks", new_callable=AsyncMock
) as mock_async_fallbacks:
mock_async_fallbacks.return_value = MagicMock(
data=[{"embedding": [0.1, 0.2, 0.3]}]
)
await router.aembedding(
model="text-embedding-ada-002", input=["test input"]
)
mock_async_fallbacks.assert_called_once()
call_kwargs = mock_async_fallbacks.call_args.kwargs
assert call_kwargs["num_retries"] == 3
def test_embedding_propagates_default_litellm_params(self):
"""
Test that embedding calls properly propagate default_litellm_params including headers.