diff --git a/litellm/router.py b/litellm/router.py index 53aa5041784..383fc1f87b0 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -6490,8 +6490,14 @@ class Router: async def try_retrieve_batch(model: DeploymentTypedDict): try: - # Update kwargs with the current model name or any other model-specific adjustments - return await litellm.alist_batches(**{**model["litellm_params"], **kwargs}) + litellm_params: Final = model["litellm_params"] + _, custom_llm_provider, _, _ = get_llm_provider( + model=litellm_params["model"], + custom_llm_provider=litellm_params.get("custom_llm_provider"), + ) + return await litellm.alist_batches( + **{**litellm_params, "custom_llm_provider": custom_llm_provider, **kwargs} + ) except Exception: return None diff --git a/tests/unit/test_router_batch_list_pagination.py b/tests/unit/test_router_batch_list_pagination.py index 87b956aea9d..7952db151f1 100644 --- a/tests/unit/test_router_batch_list_pagination.py +++ b/tests/unit/test_router_batch_list_pagination.py @@ -18,6 +18,8 @@ def _deployment(api_key: str) -> dict: async def _fake_alist_batches(**kwargs: object) -> OpenAIBatchListResponse: + if kwargs.get("custom_llm_provider") != "mistral": + raise ValueError("a mistral key sent down the default openai list path") token: Final = TOKEN_BY_KEY[str(kwargs["api_key"])] return OpenAIBatchListResponse( data=(), first_id=None, last_id=None, has_more=token is not None, next_page_token=token @@ -33,3 +35,13 @@ async def test_alist_batches_keeps_the_page_token_a_deployment_returned(monkeypa assert result["has_more"] is True assert result["next_page_token"] == "1" + + +@pytest.mark.asyncio +async def test_alist_batches_lists_with_each_deployments_own_provider(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(litellm, "alist_batches", _fake_alist_batches) + router: Final = Router(model_list=[_deployment("key-with-more-pages")]) + + result: Final = await router.alist_batches(model="mistral-ocr", limit=3) + + assert result["has_more"] is True