From ed046fad4767071ec744fabf452b65cebe52c481 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 26 Sep 2026 16:49:41 -0700 Subject: [PATCH 1/4] feat(batches): wire Mistral list and cancel through the batches dispatch --- litellm/batches/main.py | 105 +++++++++- .../llms/base_llm/batches/transformation.py | 61 +++++- litellm/llms/custom_httpx/llm_http_handler.py | 190 +++++++++++++++++- .../llms/mistral/batches/transformation.py | 113 +++++++++-- litellm/llms/xai/batches/transformation.py | 13 +- litellm/types/utils.py | 13 +- litellm/utils.py | 2 +- tests/unit/batches/test_main.py | 98 +++++++++ .../base_llm/batches/test_transformation.py | 64 +++++- .../custom_httpx/test_llm_http_handler.py | 162 +++++++++++++++ .../test_mistral_batches_transformation.py | 119 ++++++++++- 11 files changed, 897 insertions(+), 43 deletions(-) diff --git a/litellm/batches/main.py b/litellm/batches/main.py index f977fc03891..72b8f4c7ca2 100644 --- a/litellm/batches/main.py +++ b/litellm/batches/main.py @@ -13,7 +13,8 @@ https://platform.openai.com/docs/api-reference/batch import asyncio import contextvars import os -from collections.abc import Coroutine +import uuid +from collections.abc import Coroutine, Mapping from functools import partial from typing import Any, Final, Literal, cast @@ -26,6 +27,11 @@ from litellm.litellm_core_utils.get_litellm_params import add_trusted_model_cred from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.llms.anthropic.batches.handler import AnthropicBatchesHandler from litellm.llms.azure.batches.handler import AzureBatchesAPI +from litellm.llms.base_llm.batches.transformation import ( + BaseBatchesCancelConfig, + BaseBatchesConfig, + BaseBatchesListConfig, +) from litellm.llms.bedrock.batches.handler import BedrockBatchesHandler from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler @@ -62,6 +68,51 @@ vertex_ai_batches_instance: Final = VertexAIBatchPrediction(gcs_bucket_name="") anthropic_batches_instance: Final = AnthropicBatchesHandler() xai_batches_instance: Final = XAIBatchesHandler() base_llm_http_handler = BaseLLMHTTPHandler() + + +def _provider_batches_config(custom_llm_provider: str) -> BaseBatchesConfig | None: + provider: Final = next((p for p in LlmProviders if p.value == custom_llm_provider), None) + if provider is None: + return None + return ProviderConfigManager.get_provider_batches_config(model=None, provider=provider) + + +def _batch_logging_obj( + kwargs: dict[str, object], + model: str | None, + custom_llm_provider: str, + call_type: str, + call_id: str, + optional_params: GenericLiteLLMParams, + litellm_params: dict[str, object], +) -> LiteLLMLoggingObj: + logging_obj: Final = kwargs.get("litellm_logging_obj") + if isinstance(logging_obj, LiteLLMLoggingObj): + logging_obj.update_from_kwargs( + kwargs=kwargs, + model=model, + user=None, + optional_params=optional_params.model_dump(), + litellm_params=litellm_params, + custom_llm_provider=custom_llm_provider, + ) + return logging_obj + return LiteLLMLoggingObj( + model=model or f"{custom_llm_provider}/unknown", + messages=[], + stream=False, + call_type=call_type, + start_time=None, + litellm_call_id=call_id, + function_id=call_type, + ) + + +def _batch_http_client(kwargs: Mapping[str, object]) -> HTTPHandler | AsyncHTTPHandler | None: + client: Final = kwargs.get("client") + return client if isinstance(client, (HTTPHandler, AsyncHTTPHandler)) else None + + ################################################# @@ -783,6 +834,29 @@ def list_batches( timeout = 600.0 _is_async: Final = kwargs.pop("alist_batches", False) is True + model: Final = kwargs.get("model") + model_name: Final = model if isinstance(model, str) else None + provider_config: Final = _provider_batches_config(custom_llm_provider) + if isinstance(provider_config, BaseBatchesListConfig): + return base_llm_http_handler.list_batches( + after=after, + limit=limit, + litellm_params=litellm_params, + provider_config=provider_config, + logging_obj=_batch_logging_obj( + kwargs, + model=model_name, + custom_llm_provider=custom_llm_provider, + call_type="batch_list", + call_id=f"batch_list_{uuid.uuid4()}", + optional_params=optional_params, + litellm_params=litellm_params, + ), + _is_async=_is_async, + client=_batch_http_client(kwargs), + timeout=timeout, + model=model_name, + ) if custom_llm_provider == LlmProviders.XAI.value: return xai_batches_instance.list_batches( _is_async=_is_async, @@ -888,7 +962,9 @@ def list_batches( async def acancel_batch( batch_id: str, model: str | None = None, - custom_llm_provider: Literal["openai", "azure", "vertex_ai", "bedrock", "litellm_proxy", "xai"] = "openai", + custom_llm_provider: Literal[ + "openai", "azure", "vertex_ai", "bedrock", "litellm_proxy", "xai", "mistral" + ] = "openai", metadata: dict[str, str] | None = None, extra_headers: dict[str, str] | None = None, extra_body: dict[str, str] | None = None, @@ -934,7 +1010,8 @@ async def acancel_batch( def cancel_batch( batch_id: str, model: str | None = None, - custom_llm_provider: Literal["openai", "azure", "vertex_ai", "bedrock", "litellm_proxy", "xai"] | str = "openai", + custom_llm_provider: Literal["openai", "azure", "vertex_ai", "bedrock", "litellm_proxy", "xai", "mistral"] + | str = "openai", metadata: dict[str, str] | None = None, extra_headers: dict[str, str] | None = None, extra_body: dict[str, str] | None = None, @@ -984,6 +1061,26 @@ def cancel_batch( ) _is_async: Final = kwargs.pop("acancel_batch", False) is True + provider_config: Final = _provider_batches_config(custom_llm_provider) + if isinstance(provider_config, BaseBatchesCancelConfig): + return base_llm_http_handler.cancel_batch( + batch_id=batch_id, + litellm_params=litellm_params, + provider_config=provider_config, + logging_obj=_batch_logging_obj( + kwargs, + model=model, + custom_llm_provider=custom_llm_provider, + call_type="batch_cancel", + call_id="batch_cancel_" + batch_id, + optional_params=optional_params, + litellm_params=litellm_params, + ), + _is_async=_is_async, + client=_batch_http_client(kwargs), + timeout=timeout, + model=model, + ) if custom_llm_provider == LlmProviders.XAI.value: return xai_batches_instance.cancel_batch( _is_async=_is_async, @@ -1070,7 +1167,7 @@ def cancel_batch( ) else: raise litellm.exceptions.BadRequestError( - message=f"LiteLLM doesn't support {custom_llm_provider} for 'cancel_batch'. Only 'openai', 'azure', 'vertex_ai', and 'bedrock' are supported.", + message=f"LiteLLM doesn't support {custom_llm_provider} for 'cancel_batch'. Only 'openai', 'azure', 'vertex_ai', 'bedrock', 'xai', and 'mistral' are supported.", model="n/a", llm_provider=custom_llm_provider, response=httpx.Response( diff --git a/litellm/llms/base_llm/batches/transformation.py b/litellm/llms/base_llm/batches/transformation.py index 34c622d4cf6..fb2864df41e 100644 --- a/litellm/llms/base_llm/batches/transformation.py +++ b/litellm/llms/base_llm/batches/transformation.py @@ -1,15 +1,17 @@ import types from abc import ABC, abstractmethod -from typing import TYPE_CHECKING, Any +from collections.abc import Mapping +from typing import TYPE_CHECKING, Any, Literal import httpx from httpx import Headers +from typing_extensions import ReadOnly, TypedDict from litellm.types.llms.openai import ( AllMessageValues, CreateBatchRequest, ) -from litellm.types.utils import LiteLLMBatch, LlmProviders +from litellm.types.utils import LiteLLMBatch, LlmProviders, OpenAIBatchListResponse if TYPE_CHECKING: from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj @@ -206,3 +208,58 @@ class BaseBatchesConfig(ABC): Returns: Provider-specific exception class """ + + +class BatchHttpRequest(TypedDict): + """A fully-formed request the shared HTTP handler sends as-is: the provider config + resolves the base URL, path, query, and auth headers itself.""" + + method: ReadOnly[Literal["GET", "POST"]] + url: ReadOnly[str] + headers: ReadOnly[Mapping[str, str]] + + +class BaseBatchesListConfig(BaseBatchesConfig): + """Opt-in list capability: ``litellm.list_batches`` routes a provider here when its + batches config implements it, so a provider without a list endpoint carries no stub.""" + + @abstractmethod + def transform_list_batches_request( + self, + after: str | None, + limit: int | None, + litellm_params: Mapping[str, object], + ) -> BatchHttpRequest: + """Build the provider's list-jobs request from the OpenAI ``after``/``limit`` page params.""" + + @abstractmethod + def transform_list_batches_response( + self, + model: str | None, + raw_response: httpx.Response, + logging_obj: LiteLLMLoggingObj, + litellm_params: Mapping[str, object], + ) -> OpenAIBatchListResponse: + """Map the provider's job page onto the OpenAI batch list shape.""" + + +class BaseBatchesCancelConfig(BaseBatchesConfig): + """Opt-in cancel capability, the same way as :class:`BaseBatchesListConfig`.""" + + @abstractmethod + def transform_cancel_batch_request( + self, + batch_id: str, + litellm_params: Mapping[str, object], + ) -> BatchHttpRequest: + """Build the provider's cancel-job request for ``batch_id``.""" + + @abstractmethod + def transform_cancel_batch_response( + self, + model: str | None, + raw_response: httpx.Response, + logging_obj: LiteLLMLoggingObj, + litellm_params: Mapping[str, object], + ) -> LiteLLMBatch: + """Map the provider's cancel response onto a LiteLLM batch.""" diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index 0cb1416db3f..c74791750eb 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -65,7 +65,12 @@ from litellm.llms.base_llm.base_model_iterator import ( BaseModelResponseIterator, MockResponseIterator, ) -from litellm.llms.base_llm.batches.transformation import BaseBatchesConfig +from litellm.llms.base_llm.batches.transformation import ( + BaseBatchesCancelConfig, + BaseBatchesConfig, + BaseBatchesListConfig, + BatchHttpRequest, +) from litellm.llms.base_llm.chat.transformation import BaseConfig, BaseLLMException from litellm.llms.base_llm.containers.transformation import BaseContainerConfig from litellm.llms.base_llm.embedding.transformation import BaseEmbeddingConfig @@ -158,6 +163,7 @@ from litellm.types.utils import ( EmbeddingResponse, FileTypes, LiteLLMBatch, + OpenAIBatchListResponse, TranscriptionResponse, ) from litellm.types.vector_store_files import ( @@ -327,14 +333,16 @@ def _has_pre_call_deployment_hook(logging_obj: LiteLLMLoggingObj) -> bool: return False -def _mask_presigned_request_headers(transformed_request: bytes | str | dict) -> bytes | str | dict: +def _mask_presigned_request_headers( + transformed_request: bytes | str | Mapping[str, object], +) -> bytes | str | Mapping[str, object]: """A pre-signed request carries its auth inside its own ``headers`` key, which logging treats as request body (only the top-level headers channel gets masked), so mask it here before the request is handed to ``pre_call``.""" - if not isinstance(transformed_request, dict): + if not isinstance(transformed_request, Mapping): return transformed_request request_headers: Final = transformed_request.get("headers") - if not isinstance(request_headers, dict): + if not isinstance(request_headers, Mapping): return transformed_request from litellm.litellm_core_utils.litellm_logging import ( @@ -343,10 +351,21 @@ def _mask_presigned_request_headers(transformed_request: bytes | str | dict) -> return { # mutable-ok: logging's curl and raw-request builders take dict **transformed_request, - "headers": _get_masked_values(request_headers), + "headers": _get_masked_values(dict(request_headers)), # mutable-ok: the masking helper takes dict } +def _log_batch_http_request(logging_obj: "LiteLLMLoggingObj", request: BatchHttpRequest) -> None: + logging_obj.pre_call( + input="", + api_key="", + additional_args={ + "complete_input_dict": _mask_presigned_request_headers(request), + "api_base": request["url"], + }, + ) + + def _aws_signing_overrides( optional_params: Mapping[str, object], litellm_params: Mapping[str, object] ) -> Mapping[str, object]: @@ -4048,6 +4067,167 @@ class BaseLLMHTTPHandler: litellm_params=litellm_params, ) + def list_batches( + self, + after: str | None, + limit: int | None, + litellm_params: Mapping[str, object], + provider_config: BaseBatchesListConfig, + logging_obj: "LiteLLMLoggingObj", + _is_async: bool = False, + client: HTTPHandler | AsyncHTTPHandler | None = None, + timeout: float | httpx.Timeout | None = None, + model: str | None = None, + ) -> OpenAIBatchListResponse | Coroutine[object, object, OpenAIBatchListResponse]: + request: Final = provider_config.transform_list_batches_request( + after=after, limit=limit, litellm_params=litellm_params + ) + if _is_async: + return self.async_list_batches( + request=request, + litellm_params=litellm_params, + provider_config=provider_config, + logging_obj=logging_obj, + client=client, + timeout=timeout, + model=model, + ) + raw_response: Final = self._send_batch_http_request( + request=request, provider_config=provider_config, logging_obj=logging_obj, client=client, timeout=timeout + ) + return provider_config.transform_list_batches_response( + model=model, raw_response=raw_response, logging_obj=logging_obj, litellm_params=litellm_params + ) + + async def async_list_batches( + self, + request: BatchHttpRequest, + litellm_params: Mapping[str, object], + provider_config: BaseBatchesListConfig, + logging_obj: "LiteLLMLoggingObj", + client: HTTPHandler | AsyncHTTPHandler | None = None, + timeout: float | httpx.Timeout | None = None, + model: str | None = None, + ) -> OpenAIBatchListResponse: + raw_response: Final = await self._async_send_batch_http_request( + request=request, provider_config=provider_config, logging_obj=logging_obj, client=client, timeout=timeout + ) + return provider_config.transform_list_batches_response( + model=model, raw_response=raw_response, logging_obj=logging_obj, litellm_params=litellm_params + ) + + def cancel_batch( + self, + batch_id: str, + litellm_params: Mapping[str, object], + provider_config: BaseBatchesCancelConfig, + logging_obj: "LiteLLMLoggingObj", + _is_async: bool = False, + client: HTTPHandler | AsyncHTTPHandler | None = None, + timeout: float | httpx.Timeout | None = None, + model: str | None = None, + ) -> LiteLLMBatch | Coroutine[object, object, LiteLLMBatch]: + request: Final = provider_config.transform_cancel_batch_request( + batch_id=batch_id, litellm_params=litellm_params + ) + if _is_async: + return self.async_cancel_batch( + request=request, + litellm_params=litellm_params, + provider_config=provider_config, + logging_obj=logging_obj, + client=client, + timeout=timeout, + model=model, + ) + raw_response: Final = self._send_batch_http_request( + request=request, provider_config=provider_config, logging_obj=logging_obj, client=client, timeout=timeout + ) + return provider_config.transform_cancel_batch_response( + model=model, raw_response=raw_response, logging_obj=logging_obj, litellm_params=litellm_params + ) + + async def async_cancel_batch( + self, + request: BatchHttpRequest, + litellm_params: Mapping[str, object], + provider_config: BaseBatchesCancelConfig, + logging_obj: "LiteLLMLoggingObj", + client: HTTPHandler | AsyncHTTPHandler | None = None, + timeout: float | httpx.Timeout | None = None, + model: str | None = None, + ) -> LiteLLMBatch: + raw_response: Final = await self._async_send_batch_http_request( + request=request, provider_config=provider_config, logging_obj=logging_obj, client=client, timeout=timeout + ) + return provider_config.transform_cancel_batch_response( + model=model, raw_response=raw_response, logging_obj=logging_obj, litellm_params=litellm_params + ) + + def _send_batch_http_request( + self, + request: BatchHttpRequest, + provider_config: BaseBatchesConfig, + logging_obj: "LiteLLMLoggingObj", + client: HTTPHandler | AsyncHTTPHandler | None, + timeout: float | httpx.Timeout | None, + ) -> httpx.Response: + sync_httpx_client: Final = client if isinstance(client, HTTPHandler) else _get_httpx_client() + _log_batch_http_request(logging_obj, request) + try: + response: Final = ( + sync_httpx_client.get( + url=request["url"], + headers=dict(request["headers"]), # mutable-ok: HTTPHandler takes dict + timeout=timeout, + ) + if request["method"] == "GET" + else sync_httpx_client.post( + url=request["url"], + headers=dict(request["headers"]), # mutable-ok: HTTPHandler takes dict + timeout=timeout, + ) + ) + response.raise_for_status() + return response + except Exception as e: + verbose_logger.exception("Error on batch request %s %s: %s", request["method"], request["url"], e) + raise self._handle_error(e=e, provider_config=provider_config) + + async def _async_send_batch_http_request( + self, + request: BatchHttpRequest, + provider_config: BaseBatchesConfig, + logging_obj: "LiteLLMLoggingObj", + client: HTTPHandler | AsyncHTTPHandler | None, + timeout: float | httpx.Timeout | None, + ) -> httpx.Response: + async_httpx_client: Final = ( + client + if isinstance(client, AsyncHTTPHandler) + else get_async_httpx_client(llm_provider=provider_config.custom_llm_provider) + ) + _log_batch_http_request(logging_obj, request) + try: + response: Final = ( + await async_httpx_client.get( + url=request["url"], + headers=dict(request["headers"]), # mutable-ok: AsyncHTTPHandler takes dict + timeout=timeout, + ) + if request["method"] == "GET" + else await async_httpx_client.post( + url=request["url"], + headers=dict(request["headers"]), # mutable-ok: AsyncHTTPHandler takes dict + timeout=timeout, + ) + ) + response.raise_for_status() + return response + except Exception as e: + verbose_logger.exception("Error on batch request %s %s: %s", request["method"], request["url"], e) + raise self._handle_error(e=e, provider_config=provider_config) + def cancel_response_api_handler( self, response_id: str, diff --git a/litellm/llms/mistral/batches/transformation.py b/litellm/llms/mistral/batches/transformation.py index d3ed6a3af62..e967d36bda1 100644 --- a/litellm/llms/mistral/batches/transformation.py +++ b/litellm/llms/mistral/batches/transformation.py @@ -5,11 +5,15 @@ Mistral runs one model per job (set on the job, not per input line) and accepts ``/v1/ocr`` as a batch endpoint, which is how OCR gets its 50% batch discount. Output and error files are OpenAI-shaped JSONL (``{custom_id, response: {status_code, body}}``), so the shared batch cost accounting reads them without a provider branch. + +Job lists are paged by a 0-based ``page`` number rather than an id cursor, so ``after`` carries +the page number the previous response handed back as ``next_page_token``. """ from collections.abc import Mapping, Sequence from types import MappingProxyType from typing import Final, Literal, TypeAlias +from urllib.parse import urlencode import httpx from openai.types.batch import BatchRequestCounts @@ -19,10 +23,14 @@ from pydantic import BaseModel, ConfigDict from typing_extensions import NotRequired, ReadOnly, TypedDict from litellm.litellm_core_utils.url_utils import encode_url_path_segment -from litellm.llms.base_llm.batches.transformation import BaseBatchesConfig +from litellm.llms.base_llm.batches.transformation import ( + BaseBatchesCancelConfig, + BaseBatchesListConfig, + BatchHttpRequest, +) from litellm.llms.base_llm.chat.transformation import BaseLLMException from litellm.types.llms.openai import AllMessageValues, CreateBatchRequest -from litellm.types.utils import LiteLLMBatch, LlmProviders +from litellm.types.utils import LiteLLMBatch, LlmProviders, OpenAIBatchListResponse from ..common_utils import get_mistral_api_base, get_mistral_auth_headers, mistral_error @@ -33,6 +41,7 @@ OpenAIBatchStatus: TypeAlias = Literal[ "validating", "failed", "in_progress", "finalizing", "completed", "expired", "cancelling", "cancelled" ] +OPENAI_LIST_DEFAULT_LIMIT: Final = 20 _NO_HEADERS: Final[Mapping[str, str]] = MappingProxyType({}) _STATUS_MAP: Final[MappingProxyType[MistralBatchStatus, OpenAIBatchStatus]] = MappingProxyType( { @@ -56,12 +65,11 @@ class MistralCreateBatchJobRequest(TypedDict): metadata: NotRequired[ReadOnly[Mapping[str, str]]] -class MistralPresignedRequest(TypedDict): - """A fully-formed request the shared HTTP handler sends as-is (its ``method`` branch).""" +class MistralListBatchJobsQuery(TypedDict): + """Query of ``GET /v1/batch/jobs``.""" - method: ReadOnly[Literal["GET"]] - url: ReadOnly[str] - headers: ReadOnly[Mapping[str, str]] + page: ReadOnly[int] + page_size: ReadOnly[int] class MistralBatchError(BaseModel): @@ -92,6 +100,13 @@ class MistralBatchJob(BaseModel): metadata: dict[str, str] | None = None # mutable-ok: LiteLLMBatch.metadata is typed as dict +class MistralBatchJobList(BaseModel): + model_config = ConfigDict(frozen=True, extra="ignore") + + total: int + data: tuple[MistralBatchJob, ...] = () + + def _to_batch_errors(errors: Sequence[MistralBatchError]) -> BatchErrors | None: if not errors: return None @@ -131,7 +146,31 @@ def _to_litellm_batch(job: MistralBatchJob) -> LiteLLMBatch: ) -class MistralBatchesConfig(BaseBatchesConfig): +def _page_number(after: str | None) -> int: + if after is None: + return 0 + if not after.isdecimal(): + raise mistral_error( + f"Mistral pages batch jobs by number: pass the previous page's next_page_token as 'after', got {after!r}", + 400, + _NO_HEADERS, + ) + return int(after) + + +def _batch_http_request( + method: Literal["GET", "POST"], path: str, litellm_params: Mapping[str, object] +) -> BatchHttpRequest: + api_base: Final = litellm_params.get("api_base") + api_key: Final = litellm_params.get("api_key") + return BatchHttpRequest( + method=method, + url=f"{get_mistral_api_base(api_base if isinstance(api_base, str) else None)}{path}", + headers=get_mistral_auth_headers(_NO_HEADERS, api_key if isinstance(api_key, str) else None), + ) + + +class MistralBatchesConfig(BaseBatchesListConfig, BaseBatchesCancelConfig): @property def custom_llm_provider(self) -> LlmProviders: return LlmProviders.MISTRAL @@ -196,13 +235,7 @@ class MistralBatchesConfig(BaseBatchesConfig): litellm_params: Mapping[str, object], ) -> dict[str, object]: # mutable-ok: BaseBatchesConfig signature encoded_batch_id: Final = encode_url_path_segment(batch_id, field_name="batch_id") - api_base: Final = litellm_params.get("api_base") - api_key: Final = litellm_params.get("api_key") - request: Final = MistralPresignedRequest( - method="GET", - url=f"{get_mistral_api_base(api_base if isinstance(api_base, str) else None)}/v1/batch/jobs/{encoded_batch_id}", - headers=get_mistral_auth_headers(_NO_HEADERS, api_key if isinstance(api_key, str) else None), - ) + request: Final = _batch_http_request("GET", f"/v1/batch/jobs/{encoded_batch_id}", litellm_params) return dict(request) # mutable-ok: BaseBatchesConfig signature def transform_retrieve_batch_response( @@ -214,6 +247,56 @@ class MistralBatchesConfig(BaseBatchesConfig): ) -> LiteLLMBatch: return _to_litellm_batch(MistralBatchJob.model_validate(raw_response.json())) + def transform_list_batches_request( + self, + after: str | None, + limit: int | None, + litellm_params: Mapping[str, object], + ) -> BatchHttpRequest: + query: Final = MistralListBatchJobsQuery( + page=_page_number(after), + page_size=limit if limit is not None else OPENAI_LIST_DEFAULT_LIMIT, + ) + return _batch_http_request("GET", f"/v1/batch/jobs?{urlencode(query)}", litellm_params) + + def transform_list_batches_response( + self, + model: str | None, + raw_response: httpx.Response, + logging_obj: object, + litellm_params: Mapping[str, object], + ) -> OpenAIBatchListResponse: + page: Final = MistralBatchJobList.model_validate(raw_response.json()) + sent_query: Final = raw_response.request.url.params + page_number: Final = int(sent_query["page"]) + page_size: Final = int(sent_query["page_size"]) + data: Final = tuple(_to_litellm_batch(job) for job in page.data) + has_more: Final = (page_number + 1) * page_size < page.total + return OpenAIBatchListResponse( + data=data, + first_id=data[0].id if data else None, + last_id=data[-1].id if data else None, + has_more=has_more, + next_page_token=str(page_number + 1) if has_more else None, + ) + + def transform_cancel_batch_request( + self, + batch_id: str, + litellm_params: Mapping[str, object], + ) -> BatchHttpRequest: + encoded_batch_id: Final = encode_url_path_segment(batch_id, field_name="batch_id") + return _batch_http_request("POST", f"/v1/batch/jobs/{encoded_batch_id}/cancel", litellm_params) + + def transform_cancel_batch_response( + self, + model: str | None, + raw_response: httpx.Response, + logging_obj: object, + litellm_params: Mapping[str, object], + ) -> LiteLLMBatch: + return _to_litellm_batch(MistralBatchJob.model_validate(raw_response.json())) + def get_error_class( self, error_message: str, status_code: int, headers: Mapping[str, str] | httpx.Headers ) -> BaseLLMException: diff --git a/litellm/llms/xai/batches/transformation.py b/litellm/llms/xai/batches/transformation.py index 8f305b8c203..6938069ddd1 100644 --- a/litellm/llms/xai/batches/transformation.py +++ b/litellm/llms/xai/batches/transformation.py @@ -24,7 +24,7 @@ from litellm.llms.base_llm.chat.transformation import BaseLLMException from litellm.llms.xai.common_utils import XAIModelInfo from litellm.secret_managers.main import get_secret_str from litellm.types.llms.openai import CreateBatchRequest -from litellm.types.utils import LiteLLMBatch +from litellm.types.utils import LiteLLMBatch, OpenAIBatchListResponse OpenAIBatchStatus: TypeAlias = Literal[ "validating", "failed", "in_progress", "finalizing", "completed", "expired", "cancelling", "cancelled" @@ -202,17 +202,6 @@ def to_litellm_batch(batch: XAIBatch, endpoint: str = DEFAULT_BATCH_ENDPOINT) -> ) -class OpenAIBatchListResponse(BaseModel): - model_config = ConfigDict(frozen=True) - - object: Literal["list"] = "list" - data: tuple[LiteLLMBatch, ...] - first_id: str | None - last_id: str | None - has_more: bool - next_page_token: str | None = None - - def to_openai_batch_list(page: XAIBatchList) -> OpenAIBatchListResponse: data: Final = tuple(to_litellm_batch(b) for b in page.batches) return OpenAIBatchListResponse( diff --git a/litellm/types/utils.py b/litellm/types/utils.py index f8b57139b37..7346916972d 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -4160,7 +4160,7 @@ FILE_CONTENT_STREAMING_PROVIDERS: Final[frozenset[str]] = frozenset( LITELLM_EXECUTED_BATCH_PROVIDERS: Final[frozenset[str]] = frozenset({LlmProviders.HOSTED_VLLM.value}) -ListBatchesSupportedProvider = Literal["openai", "azure", "hosted_vllm", "litellm_proxy", "vertex_ai", "xai"] +ListBatchesSupportedProvider = Literal["openai", "azure", "hosted_vllm", "litellm_proxy", "vertex_ai", "xai", "mistral"] LIST_BATCHES_SUPPORTED_PROVIDERS: Final[frozenset[str]] = frozenset(get_args(ListBatchesSupportedProvider)) @@ -4307,6 +4307,17 @@ class LiteLLMBatch(Batch): return self.dict() +class OpenAIBatchListResponse(BaseModel): + model_config = ConfigDict(frozen=True) + + object: Literal["list"] = "list" + data: tuple[LiteLLMBatch, ...] + first_id: str | None + last_id: str | None + has_more: bool + next_page_token: str | None = None + + class LiteLLMRealtimeStreamLoggingObject(LiteLLMPydanticObjectBase): # Events are already well-formed provider dicts. Validating them against the # OpenAIRealtimeEvents union makes Pydantic try every member per event, which diff --git a/litellm/utils.py b/litellm/utils.py index 092fe936cf9..dda888aaf02 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -9373,7 +9373,7 @@ class ProviderConfigManager: @staticmethod def get_provider_batches_config( - model: str, + model: str | None, provider: LlmProviders, ) -> BaseBatchesConfig | None: if LlmProviders.BEDROCK == provider: diff --git a/tests/unit/batches/test_main.py b/tests/unit/batches/test_main.py index 26dc4083b0b..590024312e8 100644 --- a/tests/unit/batches/test_main.py +++ b/tests/unit/batches/test_main.py @@ -34,6 +34,8 @@ import pytest import litellm import litellm.batches.main as bm +from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj +from litellm.llms.mistral.batches.transformation import MistralBatchesConfig # --------------------------------------------------------------------------- # @@ -336,6 +338,47 @@ def test_list__unsupported_provider_raises_badrequest(seams): m.assert_not_called() +def test_list__unknown_provider_string_raises_badrequest_not_valueerror(seams): + with pytest.raises(litellm.exceptions.BadRequestError): + bm.list_batches(custom_llm_provider="not-a-provider") # type: ignore[arg-type] + + for m in _all_seam_methods(seams, "list_batches"): + m.assert_not_called() + + +def test_list__provider_config_with_list_capability_routes_to_base_http_handler(seams): + """mistral has no per-provider batches instance: its config implements the list + capability, so list_batches hands it to the generic base_llm_http_handler.""" + result = bm.list_batches(custom_llm_provider="mistral", after="1", limit=5) + + assert result is seams.base_http.list_batches.return_value + _assert_only(seams.base_http.list_batches, seams, "list_batches") + + kw = seams.base_http.list_batches.call_args.kwargs + assert isinstance(kw["provider_config"], MistralBatchesConfig) + assert kw["after"] == "1" + assert kw["limit"] == 5 + assert kw["_is_async"] is False + assert isinstance(kw["logging_obj"], LiteLLMLoggingObj) + assert kw["logging_obj"].model_call_details["custom_llm_provider"] == "mistral" + + +def test_list__provider_config_async_flag_propagates_is_async(seams): + bm.list_batches(custom_llm_provider="mistral", alist_batches=True) + + assert seams.base_http.list_batches.call_args.kwargs["_is_async"] is True + + +def test_list__provider_config_without_list_capability_keeps_legacy_dispatch(seams): + """bedrock has a batches config too, but one that cannot list, so the config-first + lookup must fall through to the legacy switch (which rejects bedrock for list).""" + with pytest.raises(litellm.exceptions.BadRequestError): + bm.list_batches(custom_llm_provider="bedrock") # type: ignore[arg-type] + + for m in _all_seam_methods(seams, "list_batches"): + m.assert_not_called() + + # =========================================================================== # # cancel_batch (supported: openai, hosted_vllm, azure, vertex_ai; no @client) # =========================================================================== # @@ -384,6 +427,45 @@ def test_cancel__async_flag_propagates_is_async(seams): assert seams.openai.cancel_batch.call_args.kwargs["_is_async"] is True +def test_cancel__unknown_provider_string_raises_badrequest_not_valueerror(seams): + with pytest.raises(litellm.exceptions.BadRequestError): + bm.cancel_batch(batch_id="batch-1", custom_llm_provider="not-a-provider") + + for m in _all_seam_methods(seams, "cancel_batch"): + m.assert_not_called() + + +def test_cancel__provider_config_with_cancel_capability_routes_to_base_http_handler(seams): + result = bm.cancel_batch(batch_id="batch-1", custom_llm_provider="mistral") + + assert result is seams.base_http.cancel_batch.return_value + _assert_only(seams.base_http.cancel_batch, seams, "cancel_batch") + seams.bedrock_arn.cancel_batch.assert_not_called() + + kw = seams.base_http.cancel_batch.call_args.kwargs + assert isinstance(kw["provider_config"], MistralBatchesConfig) + assert kw["batch_id"] == "batch-1" + assert kw["_is_async"] is False + assert kw["logging_obj"].call_type == "batch_cancel" + + +def test_cancel__provider_config_async_flag_propagates_is_async(seams): + bm.cancel_batch(batch_id="batch-1", custom_llm_provider="mistral", acancel_batch=True) + + assert seams.base_http.cancel_batch.call_args.kwargs["_is_async"] is True + + +def test_cancel__provider_config_without_cancel_capability_keeps_legacy_dispatch(seams): + """bedrock's batches config cannot cancel, so cancel_batch still lands on the + Bedrock ARN handler rather than the generic HTTP handler.""" + result = bm.cancel_batch(batch_id="batch-1", custom_llm_provider="bedrock") + + assert result is seams.bedrock_arn.cancel_batch.return_value + seams.bedrock_arn.cancel_batch.assert_called_once() + for m in _all_seam_methods(seams, "cancel_batch"): + m.assert_not_called() + + # =========================================================================== # # Async wrappers - delegate to the sync function in an executor, set the right # "_is_async" flag, and return the result untouched. @@ -650,9 +732,25 @@ def test_list__vertex_credentials_passthrough(seams): } +def test_list__mistral_credentials_passthrough(seams): + bm.list_batches(custom_llm_provider="mistral", api_key="sk-user-mistral", api_base="https://mistral.user.test") + + litellm_params = seams.base_http.list_batches.call_args.kwargs["litellm_params"] + assert (litellm_params["api_key"], litellm_params["api_base"]) == ("sk-user-mistral", "https://mistral.user.test") + + # ---- cancel_batch ---------------------------------------------------------- # +def test_cancel__mistral_credentials_passthrough(seams): + bm.cancel_batch( + batch_id="b1", custom_llm_provider="mistral", api_key="sk-user-mistral", api_base="https://mistral.user.test" + ) + + litellm_params = seams.base_http.cancel_batch.call_args.kwargs["litellm_params"] + assert (litellm_params["api_key"], litellm_params["api_base"]) == ("sk-user-mistral", "https://mistral.user.test") + + def test_cancel__openai_credentials_passthrough(seams): bm.cancel_batch(batch_id="b1", custom_llm_provider="openai", **OPENAI_CREDS) diff --git a/tests/unit/llms/base_llm/batches/test_transformation.py b/tests/unit/llms/base_llm/batches/test_transformation.py index 0c360ce2ed9..6a039792cc1 100644 --- a/tests/unit/llms/base_llm/batches/test_transformation.py +++ b/tests/unit/llms/base_llm/batches/test_transformation.py @@ -22,8 +22,12 @@ filter to all single-underscore names) makes a test fail. import pytest -from litellm.llms.base_llm.batches.transformation import BaseBatchesConfig -from litellm.types.utils import LlmProviders +from litellm.llms.base_llm.batches.transformation import ( + BaseBatchesCancelConfig, + BaseBatchesConfig, + BaseBatchesListConfig, +) +from litellm.types.utils import LiteLLMBatch, LlmProviders, OpenAIBatchListResponse # --------------------------------------------------------------------------- # @@ -129,6 +133,62 @@ def test_subclass_missing_any_abstract_member_cannot_instantiate(missing_member) Incomplete() +# =========================================================================== # +# Opt-in list / cancel capability ABCs +# =========================================================================== # + + +class _ListAndCancelBatchesConfig(_ConcreteBatchesConfig, BaseBatchesListConfig, BaseBatchesCancelConfig): + def transform_list_batches_request(self, after, limit, litellm_params): + return {"method": "GET", "url": "https://example.test/list", "headers": {}} + + def transform_list_batches_response(self, model, raw_response, logging_obj, litellm_params): + return OpenAIBatchListResponse(data=(), first_id=None, last_id=None, has_more=False) + + def transform_cancel_batch_request(self, batch_id, litellm_params): + return {"method": "POST", "url": f"https://example.test/{batch_id}/cancel", "headers": {}} + + def transform_cancel_batch_response(self, model, raw_response, logging_obj, litellm_params): + return LiteLLMBatch( + id="b", object="batch", endpoint="/v1/chat/completions", input_file_id="f", completion_window="24h", + status="cancelled", created_at=0, + ) + + +def test_list_and_cancel_capable_subclass_is_a_batches_config(): + instance = _ListAndCancelBatchesConfig() + assert isinstance(instance, BaseBatchesConfig) + assert isinstance(instance, BaseBatchesListConfig) + assert isinstance(instance, BaseBatchesCancelConfig) + + +def test_plain_batches_config_carries_neither_capability(): + instance = _ConcreteBatchesConfig() + assert not isinstance(instance, BaseBatchesListConfig) + assert not isinstance(instance, BaseBatchesCancelConfig) + + +@pytest.mark.parametrize( + "base,missing_member", + [ + (BaseBatchesListConfig, "transform_list_batches_request"), + (BaseBatchesListConfig, "transform_list_batches_response"), + (BaseBatchesCancelConfig, "transform_cancel_batch_request"), + (BaseBatchesCancelConfig, "transform_cancel_batch_response"), + ], +) +def test_capability_subclass_missing_its_member_cannot_instantiate(base, missing_member): + namespace = { + k: v + for k, v in {**_ConcreteBatchesConfig.__dict__, **_ListAndCancelBatchesConfig.__dict__}.items() + if not k.startswith("__") + } + namespace.pop(missing_member) + Incomplete = type("Incomplete", (base,), namespace) + with pytest.raises(TypeError): + Incomplete() + + # =========================================================================== # # get_config() # =========================================================================== # diff --git a/tests/unit/llms/custom_httpx/test_llm_http_handler.py b/tests/unit/llms/custom_httpx/test_llm_http_handler.py index f3332cb513c..3301d3d0481 100644 --- a/tests/unit/llms/custom_httpx/test_llm_http_handler.py +++ b/tests/unit/llms/custom_httpx/test_llm_http_handler.py @@ -2887,6 +2887,168 @@ async def test_async_retrieve_batch_masks_presigned_auth_header_in_raw_request_l assert provider_key not in json.dumps(raw_request_body) +def _mistral_batch_logging_obj(call_type: str): + from litellm.litellm_core_utils.litellm_logging import Logging as LitellmLogging + + logging_obj = LitellmLogging( + model="mistral/mistral-ocr-latest", + messages=[], + stream=False, + call_type=call_type, + start_time=time.time(), + litellm_call_id=f"{call_type}-call-id", + function_id=f"{call_type}-function-id", + log_raw_request_response=True, + ) + logging_obj.update_environment_variables( + model="mistral/mistral-ocr-latest", + optional_params={}, + litellm_params={"litellm_call_id": f"{call_type}-call-id", "metadata": {}}, + ) + return logging_obj + + +def _mistral_job(job_id: str, status: str) -> dict: + return { + "id": job_id, + "input_files": ["file-1"], + "endpoint": "/v1/ocr", + "model": "mistral-ocr-latest", + "status": status, + "created_at": 1_757_400_000, + } + + +def _mock_batch_client(handler, is_async: bool): + sent_requests = [] + + def _respond(request: httpx.Request) -> httpx.Response: + sent_requests.append(request) + return handler(request) + + if is_async: + client = AsyncHTTPHandler() + client.client = httpx.AsyncClient(transport=httpx.MockTransport(_respond)) + else: + client = HTTPHandler() + client.client = httpx.Client(transport=httpx.MockTransport(_respond)) + return client, sent_requests + + +async def _maybe_await(value): + return await value if asyncio.iscoroutine(value) else value + + +@pytest.mark.asyncio +@pytest.mark.parametrize("is_async", [False, True]) +async def test_list_batches_sends_provider_list_request_and_maps_page(is_async): + """The generic handler sends the provider config's fully-formed list request (GET, + query, auth) and hands the raw page back to the config, whose result is returned as-is. + The provider key travels on the wire but never into the raw request log.""" + from litellm.llms.mistral.batches.transformation import MistralBatchesConfig + + provider_key = "mistral-s3cret-provider-key-123456" + client, sent_requests = _mock_batch_client( + lambda request: httpx.Response( + 200, + json={ + "object": "list", + "total": 5, + "data": [_mistral_job("job-c", "RUNNING"), _mistral_job("job-d", "SUCCESS")], + }, + request=request, + ), + is_async, + ) + logging_obj = _mistral_batch_logging_obj("batch_list") + + result = await _maybe_await( + BaseLLMHTTPHandler().list_batches( + after="1", + limit=2, + litellm_params={"api_key": provider_key}, + provider_config=MistralBatchesConfig(), + logging_obj=logging_obj, + _is_async=is_async, + client=client, + model="mistral/mistral-ocr-latest", + ) + ) + + assert sent_requests[0].method == "GET" + assert str(sent_requests[0].url) == "https://api.mistral.ai/v1/batch/jobs?page=1&page_size=2" + assert sent_requests[0].headers["Authorization"] == f"Bearer {provider_key}" + assert [b.id for b in result.data] == ["job-c", "job-d"] + assert result.has_more is True + assert result.next_page_token == "2" + raw_request_body = logging_obj.model_call_details["raw_request_typed_dict"]["raw_request_body"] + assert provider_key not in json.dumps(raw_request_body) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("is_async", [False, True]) +async def test_cancel_batch_posts_provider_cancel_request_and_maps_job(is_async): + from litellm.llms.mistral.batches.transformation import MistralBatchesConfig + + provider_key = "mistral-s3cret-provider-key-123456" + client, sent_requests = _mock_batch_client( + lambda request: httpx.Response(200, json=_mistral_job("job-1", "CANCELLATION_REQUESTED"), request=request), + is_async, + ) + logging_obj = _mistral_batch_logging_obj("batch_cancel") + + result = await _maybe_await( + BaseLLMHTTPHandler().cancel_batch( + batch_id="job-1", + litellm_params={"api_key": provider_key}, + provider_config=MistralBatchesConfig(), + logging_obj=logging_obj, + _is_async=is_async, + client=client, + model="mistral/mistral-ocr-latest", + ) + ) + + assert sent_requests[0].method == "POST" + assert str(sent_requests[0].url) == "https://api.mistral.ai/v1/batch/jobs/job-1/cancel" + assert sent_requests[0].headers["Authorization"] == f"Bearer {provider_key}" + assert result.id == "job-1" + assert result.status == "cancelling" + raw_request_body = logging_obj.model_call_details["raw_request_typed_dict"]["raw_request_body"] + assert provider_key not in json.dumps(raw_request_body) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("is_async", [False, True]) +@pytest.mark.parametrize("operation", ["list_batches", "cancel_batch"]) +async def test_batch_list_and_cancel_map_provider_errors_through_config(is_async, operation): + """A non-2xx provider reply (a GET never raises on its own in HTTPHandler) becomes the + config's error class with the provider's status and body, not a parse failure.""" + from litellm.llms.mistral.batches.transformation import MistralBatchesConfig + from litellm.llms.mistral.common_utils import MistralError + + client, _ = _mock_batch_client( + lambda request: httpx.Response(404, json={"detail": "job not found"}, request=request), is_async + ) + call_kwargs = {"after": None, "limit": None} if operation == "list_batches" else {"batch_id": "job-missing"} + + with pytest.raises(MistralError) as excinfo: + await _maybe_await( + getattr(BaseLLMHTTPHandler(), operation)( + **call_kwargs, + litellm_params={"api_key": "sk-test"}, + provider_config=MistralBatchesConfig(), + logging_obj=_mistral_batch_logging_obj("batch_error"), + _is_async=is_async, + client=client, + model="mistral/mistral-ocr-latest", + ) + ) + + assert excinfo.value.status_code == 404 + assert "job not found" in excinfo.value.message + + @pytest.mark.asyncio async def test_async_anthropic_messages_handler_carries_deployment_vertex_location_for_pricing(monkeypatch): """ diff --git a/tests/unit/llms/mistral/batches/test_mistral_batches_transformation.py b/tests/unit/llms/mistral/batches/test_mistral_batches_transformation.py index 4073879e3b8..c4c150f27a6 100644 --- a/tests/unit/llms/mistral/batches/test_mistral_batches_transformation.py +++ b/tests/unit/llms/mistral/batches/test_mistral_batches_transformation.py @@ -4,7 +4,9 @@ behind ``custom_llm_provider="mistral"`` on /v1/batches. Locks the request shape Mistral's ``POST /v1/batch/jobs`` accepts (input_files list, model set on the job, endpoint passed through untouched so ``/v1/ocr`` batches work), -the Mistral -> OpenAI status mapping, request-count and file-id mapping, and auth. +the Mistral -> OpenAI status mapping, request-count and file-id mapping, auth, the +page-number pagination of ``GET /v1/batch/jobs`` behind OpenAI's ``after``/``limit``, +and ``POST /v1/batch/jobs/{id}/cancel``. Everything runs for real against canned httpx responses; only the API key env var is set. """ @@ -61,6 +63,14 @@ def _response(payload: dict, status_code: int = 200) -> httpx.Response: ) +def _list_response(jobs: list[dict], total: int, page: int, page_size: int) -> httpx.Response: + return httpx.Response( + status_code=200, + content=json.dumps({"object": "list", "total": total, "data": jobs}).encode(), + request=httpx.Request("GET", f"https://api.mistral.ai/v1/batch/jobs?page={page}&page_size={page_size}"), + ) + + @pytest.fixture def config() -> MistralBatchesConfig: return MistralBatchesConfig() @@ -256,3 +266,110 @@ def test_get_error_class(config): assert isinstance(err, MistralError) assert err.status_code == 401 assert err.message == "nope" + + +def test_list_request_maps_after_and_limit_onto_page_query(config, api_key): + req = config.transform_list_batches_request(after="2", limit=5, litellm_params={}) + assert req["method"] == "GET" + assert req["url"] == "https://api.mistral.ai/v1/batch/jobs?page=2&page_size=5" + assert req["headers"] == {"Authorization": f"Bearer {api_key}"} + + +def test_list_request_defaults_to_first_page_of_twenty(config, api_key): + req = config.transform_list_batches_request(after=None, limit=None, litellm_params={}) + assert req["url"] == "https://api.mistral.ai/v1/batch/jobs?page=0&page_size=20" + + +def test_list_request_prefers_litellm_params_credentials(config, api_key): + req = config.transform_list_batches_request( + after=None, limit=3, litellm_params={"api_key": "sk-from-deployment", "api_base": "https://mistral.local/v1"} + ) + assert req["url"] == "https://mistral.local/v1/batch/jobs?page=0&page_size=3" + assert req["headers"]["Authorization"] == "Bearer sk-from-deployment" + + +@pytest.mark.parametrize("after", ["batch_abc", "-1", "1.5", ""]) +def test_list_request_rejects_non_page_number_cursor(config, api_key, after): + with pytest.raises(MistralError) as excinfo: + config.transform_list_batches_request(after=after, limit=None, litellm_params={}) + assert excinfo.value.status_code == 400 + assert "next_page_token" in excinfo.value.message + + +def test_list_response_maps_jobs_and_signals_more_pages(config): + jobs = [_job(id="job-a", status="RUNNING"), _job(id="job-b")] + page = config.transform_list_batches_response( + model=None, raw_response=_list_response(jobs, total=5, page=1, page_size=2), logging_obj=None, litellm_params={} + ) + assert page.object == "list" + assert [b.id for b in page.data] == ["job-a", "job-b"] + assert all(isinstance(b, LiteLLMBatch) for b in page.data) + assert page.data[0].status == "in_progress" + assert page.first_id == "job-a" + assert page.last_id == "job-b" + assert page.has_more is True + assert page.next_page_token == "2" + + +def test_list_response_last_page_has_no_more(config): + page = config.transform_list_batches_response( + model=None, + raw_response=_list_response([_job(id="job-e")], total=5, page=2, page_size=2), + logging_obj=None, + litellm_params={}, + ) + assert [b.id for b in page.data] == ["job-e"] + assert page.has_more is False + assert page.next_page_token is None + + +def test_list_response_exactly_filled_last_page_has_no_more(config): + page = config.transform_list_batches_response( + model=None, + raw_response=_list_response([_job(id="job-c"), _job(id="job-d")], total=4, page=1, page_size=2), + logging_obj=None, + litellm_params={}, + ) + assert page.has_more is False + assert page.next_page_token is None + + +def test_list_response_empty_page(config): + page = config.transform_list_batches_response( + model=None, raw_response=_list_response([], total=0, page=0, page_size=20), logging_obj=None, litellm_params={} + ) + assert page.data == () + assert page.first_id is None + assert page.last_id is None + assert page.has_more is False + + +def test_cancel_request_posts_to_cancel_with_auth(config, api_key): + req = config.transform_cancel_batch_request(batch_id="job/with slash", litellm_params={}) + assert req["method"] == "POST" + assert req["url"] == "https://api.mistral.ai/v1/batch/jobs/job%2Fwith%20slash/cancel" + assert req["headers"] == {"Authorization": f"Bearer {api_key}"} + + +def test_cancel_request_prefers_litellm_params_credentials(config, api_key): + req = config.transform_cancel_batch_request( + batch_id="job-1", litellm_params={"api_key": "sk-from-deployment", "api_base": "https://mistral.local"} + ) + assert req["url"] == "https://mistral.local/v1/batch/jobs/job-1/cancel" + assert req["headers"]["Authorization"] == "Bearer sk-from-deployment" + + +@pytest.mark.parametrize( + "mistral_status,openai_status", [("CANCELLATION_REQUESTED", "cancelling"), ("CANCELLED", "cancelled")] +) +def test_cancel_response_maps_job_onto_openai_batch(config, mistral_status, openai_status): + batch = config.transform_cancel_batch_response( + model=None, + raw_response=_response(_job(status=mistral_status, completed_at=1_757_400_600)), + logging_obj=None, + litellm_params={}, + ) + assert isinstance(batch, LiteLLMBatch) + assert batch.id == "8ff5e0d1-6bc2-4c3a-9f7d-0d1c2e3f4a5b" + assert batch.status == openai_status + assert batch.cancelled_at == (1_757_400_600 if openai_status == "cancelled" else None) From 7a25d226750a3544429c0dae8a9915043b40e844 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 26 Sep 2026 17:21:12 -0700 Subject: [PATCH 2/4] fix(router): carry next_page_token through the batch list aggregation --- litellm/router.py | 4 +++ .../unit/test_router_batch_list_pagination.py | 35 +++++++++++++++++++ 2 files changed, 39 insertions(+) create mode 100644 tests/unit/test_router_batch_list_pagination.py diff --git a/litellm/router.py b/litellm/router.py index 1ef68e60440..53aa5041784 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -6504,6 +6504,7 @@ class Router: "first_id": None, "last_id": None, "has_more": False, + "next_page_token": None, } for result in results: @@ -6513,6 +6514,9 @@ class Router: final_results["first_id"] = getattr(result, "first_id") final_results["last_id"] = getattr(result, "last_id") final_results["data"].extend(result.data) + page_token = getattr(result, "next_page_token", None) + if page_token is not None: + final_results["next_page_token"] = page_token ## check 'has_more' if getattr(result, "has_more", False) is True: diff --git a/tests/unit/test_router_batch_list_pagination.py b/tests/unit/test_router_batch_list_pagination.py new file mode 100644 index 00000000000..87b956aea9d --- /dev/null +++ b/tests/unit/test_router_batch_list_pagination.py @@ -0,0 +1,35 @@ +from typing import Final + +import pytest + +import litellm +from litellm import Router +from litellm.types.utils import OpenAIBatchListResponse + +TOKEN_BY_KEY: Final = {"key-with-more-pages": "1", "key-on-last-page": None} + + +def _deployment(api_key: str) -> dict: + return { + "model_name": "mistral-ocr", + "litellm_params": {"model": "mistral/mistral-ocr-latest", "api_key": api_key}, + "model_info": {"id": api_key}, + } + + +async def _fake_alist_batches(**kwargs: object) -> OpenAIBatchListResponse: + 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 + ) + + +@pytest.mark.asyncio +async def test_alist_batches_keeps_the_page_token_a_deployment_returned(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(litellm, "alist_batches", _fake_alist_batches) + router: Final = Router(model_list=[_deployment("key-with-more-pages"), _deployment("key-on-last-page")]) + + result: Final = await router.alist_batches(model="mistral-ocr", limit=3) + + assert result["has_more"] is True + assert result["next_page_token"] == "1" From cdcbde1d5f1c763c7722f04b3b0b1ecc0065c3a8 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 26 Sep 2026 17:40:45 -0700 Subject: [PATCH 3/4] fix(router): list batches with each deployment's own provider --- litellm/router.py | 10 ++++++++-- tests/unit/test_router_batch_list_pagination.py | 12 ++++++++++++ 2 files changed, 20 insertions(+), 2 deletions(-) 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 From 897e2b192910b6073b1585d07432e3f22e737607 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 26 Sep 2026 18:02:06 -0700 Subject: [PATCH 4/4] test(batches): drop type ignore comments from the dispatch tests --- tests/unit/batches/test_main.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/tests/unit/batches/test_main.py b/tests/unit/batches/test_main.py index 590024312e8..93ddeade882 100644 --- a/tests/unit/batches/test_main.py +++ b/tests/unit/batches/test_main.py @@ -340,7 +340,7 @@ def test_list__unsupported_provider_raises_badrequest(seams): def test_list__unknown_provider_string_raises_badrequest_not_valueerror(seams): with pytest.raises(litellm.exceptions.BadRequestError): - bm.list_batches(custom_llm_provider="not-a-provider") # type: ignore[arg-type] + bm.list_batches(custom_llm_provider="not-a-provider") for m in _all_seam_methods(seams, "list_batches"): m.assert_not_called() @@ -373,7 +373,7 @@ def test_list__provider_config_without_list_capability_keeps_legacy_dispatch(sea """bedrock has a batches config too, but one that cannot list, so the config-first lookup must fall through to the legacy switch (which rejects bedrock for list).""" with pytest.raises(litellm.exceptions.BadRequestError): - bm.list_batches(custom_llm_provider="bedrock") # type: ignore[arg-type] + bm.list_batches(custom_llm_provider="bedrock") for m in _all_seam_methods(seams, "list_batches"): m.assert_not_called()