feat(batches): wire Mistral list and cancel through the batches dispatch

This commit is contained in:
mateo-berri 2026-09-26 16:49:41 -07:00
parent 8d166258a6
commit ed046fad47
11 changed files with 897 additions and 43 deletions

View file

@ -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(

View file

@ -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."""

View file

@ -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,

View file

@ -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:

View file

@ -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(

View file

@ -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

View file

@ -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:

View file

@ -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)

View file

@ -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()
# =========================================================================== #

View file

@ -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):
"""

View file

@ -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)