feat(batches): support Mistral files/batches and per-page OCR batch cost tracking

Adds MistralFilesConfig and MistralBatchesConfig so Mistral can be used as a
Files and Batches provider through the shared BaseLLMHTTPHandler path, the
same way Bedrock plugs in. /v1/ocr is now an accepted batch endpoint, and
completed OCR batches are billed per page (ocr_cost_per_page_batches, half
the synchronous rate) instead of per token.

Resolves #29914
This commit is contained in:
mubashir1osmani 2026-09-09 18:06:25 -04:00
parent f529d6d6bd
commit 233337628f
23 changed files with 1216 additions and 41 deletions

View file

@ -9,6 +9,7 @@ import litellm
from litellm._logging import verbose_logger
from litellm.litellm_core_utils.get_litellm_params import AWS_CREDENTIAL_KWARGS_KEYS
from litellm.litellm_core_utils.llm_cost_calc.utils import parse_prompt_tokens_details
from litellm.llms.base_llm.ocr.transformation import OCRUsageInfo
from litellm.types.llms.openai import Batch
from litellm.types.utils import ModelInfo, Usage
from litellm.utils import token_counter
@ -50,7 +51,7 @@ def batch_cost_is_final(batch: Batch) -> bool:
async def calculate_batch_cost_and_usage(
file_content_dictionary: list[dict],
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic"],
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "mistral"],
model_name: str | None = None,
model_info: ModelInfo | None = None,
) -> BatchCostUsageResult:
@ -80,7 +81,7 @@ async def calculate_batch_cost_and_usage(
async def _handle_completed_batch(
batch: Batch,
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic"],
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "mistral"],
model_name: str | None = None,
litellm_params: dict | None = None,
model_info: ModelInfo | None = None,
@ -166,7 +167,7 @@ class _BatchOutputLineStats:
def _classify_output_line_stats(
entries: Iterable[dict],
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "bedrock"],
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "bedrock", "mistral"],
model_name: str | None,
model_info: ModelInfo | None,
) -> Iterator[_BatchOutputLineStats | _LineOutcome]:
@ -185,7 +186,7 @@ def _classify_output_line_stats(
def _safe_output_line_stats(
entry: Mapping[str, object],
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "bedrock"],
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "bedrock", "mistral"],
model_name: str | None,
model_info: ModelInfo | None,
) -> _BatchOutputLineStats | None:
@ -207,7 +208,7 @@ def _safe_output_line_stats(
def _compute_output_line_stats(
entry: Mapping[str, object],
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "bedrock"],
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "bedrock", "mistral"],
model_name: str | None,
model_info: ModelInfo | None,
) -> _BatchOutputLineStats:
@ -218,6 +219,7 @@ def _compute_output_line_stats(
response_model: Final = raw_model if isinstance(raw_model, str) and raw_model else None
completion_details: Final = usage.completion_tokens_details
line_prompt_cost, line_completion_cost = _output_line_cost(
response_body=response_body,
usage=usage,
custom_llm_provider=custom_llm_provider,
model_name=model_name,
@ -237,19 +239,36 @@ def _compute_output_line_stats(
)
def _ocr_usage_info_from_response_body(response_body: Mapping[str, object]) -> OCRUsageInfo | None:
"""OCR results report ``usage_info`` (pages) instead of ``usage`` (tokens); None for non-OCR lines."""
raw_usage_info: Final = response_body.get("usage_info")
if not isinstance(raw_usage_info, Mapping):
return None
return OCRUsageInfo.model_validate(raw_usage_info)
def _output_line_cost(
response_body: Mapping[str, object],
usage: Usage,
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "bedrock"],
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "bedrock", "mistral"],
model_name: str | None,
response_model: str | None,
model_info: ModelInfo | None,
) -> tuple[float, float]:
"""(prompt_cost, completion_cost) for one output line, priced at batch rates."""
from litellm.cost_calculator import batch_cost_calculator
from litellm.cost_calculator import batch_cost_calculator, ocr_batch_cost
cost_model: Final = (
model_name if custom_llm_provider == "bedrock" and model_name else response_model or model_name or ""
)
ocr_usage: Final = _ocr_usage_info_from_response_body(response_body)
if ocr_usage is not None:
return ocr_batch_cost(
model=cost_model,
custom_llm_provider=custom_llm_provider,
usage_info=ocr_usage,
model_info=model_info,
)
return batch_cost_calculator(
usage=usage,
model=cost_model,
@ -260,7 +279,7 @@ def _output_line_cost(
def _aggregate_batch_cost_usage_models(
entries: Iterable[dict],
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "bedrock"],
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "bedrock", "mistral"],
model_name: str | None = None,
model_info: ModelInfo | None = None,
) -> BatchCostUsageResult:
@ -427,7 +446,7 @@ def _provider_output_file_id(output_file_id: str) -> str:
async def _fetch_batch_managed_file_content(
file_id: str,
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic"] = "openai",
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "mistral"] = "openai",
litellm_params: dict | None = None,
) -> bytes:
"""
@ -457,7 +476,7 @@ async def _fetch_batch_managed_file_content(
async def _fetch_batch_output_file_content(
batch: Batch,
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic"] = "openai",
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "mistral"] = "openai",
litellm_params: dict | None = None,
) -> bytes:
"""
@ -479,7 +498,7 @@ async def _fetch_batch_output_file_content(
async def count_error_file_failed_requests(
batch: Batch,
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic"],
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "mistral"],
litellm_params: dict | None,
) -> int:
"""Count failed requests reported only in the batch's separate error file.

View file

@ -105,9 +105,11 @@ def _resolve_timeout(
@client
async def acreate_batch(
completion_window: Literal["24h"],
endpoint: Literal["/v1/chat/completions", "/v1/embeddings", "/v1/completions", "/v1/responses"],
endpoint: Literal["/v1/chat/completions", "/v1/embeddings", "/v1/completions", "/v1/responses", "/v1/ocr"],
input_file_id: str,
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "bedrock", "hosted_vllm", "litellm_proxy"] = "openai",
custom_llm_provider: Literal[
"openai", "azure", "vertex_ai", "bedrock", "hosted_vllm", "litellm_proxy", "mistral"
] = "openai",
metadata: dict[str, str] | None = None,
extra_headers: dict[str, str] | None = None,
extra_body: dict[str, str] | None = None,
@ -155,9 +157,11 @@ async def acreate_batch(
@client
def create_batch(
completion_window: Literal["24h"],
endpoint: Literal["/v1/chat/completions", "/v1/embeddings", "/v1/completions", "/v1/responses"],
endpoint: Literal["/v1/chat/completions", "/v1/embeddings", "/v1/completions", "/v1/responses", "/v1/ocr"],
input_file_id: str,
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "bedrock", "hosted_vllm", "litellm_proxy"] = "openai",
custom_llm_provider: Literal[
"openai", "azure", "vertex_ai", "bedrock", "hosted_vllm", "litellm_proxy", "mistral"
] = "openai",
metadata: dict[str, str] | None = None,
extra_headers: dict[str, str] | None = None,
extra_body: dict[str, str] | None = None,
@ -341,7 +345,7 @@ def create_batch(
async def aretrieve_batch(
batch_id: str,
custom_llm_provider: Literal[
"openai", "azure", "vertex_ai", "bedrock", "hosted_vllm", "litellm_proxy", "anthropic"
"openai", "azure", "vertex_ai", "bedrock", "hosted_vllm", "litellm_proxy", "anthropic", "mistral"
] = "openai",
metadata: dict[str, str] | None = None,
extra_headers: dict[str, str] | None = None,
@ -389,7 +393,7 @@ def _handle_retrieve_batch_providers_without_provider_config(
_retrieve_batch_request: RetrieveBatchRequest,
_is_async: bool,
custom_llm_provider: Literal[
"openai", "azure", "vertex_ai", "bedrock", "hosted_vllm", "litellm_proxy", "anthropic"
"openai", "azure", "vertex_ai", "bedrock", "hosted_vllm", "litellm_proxy", "anthropic", "mistral"
] = "openai",
logging_obj: LiteLLMLoggingObj | None = None,
):
@ -497,7 +501,7 @@ def _handle_retrieve_batch_providers_without_provider_config(
message=(
f"LiteLLM doesn't support custom_llm_provider={custom_llm_provider} for 'retrieve_batch' without a `model` kwarg. "
"Supported via this path: 'openai', 'azure', 'vertex_ai', 'anthropic'. "
"'bedrock' is supported but requires `model` to be passed so the provider config can be loaded."
"'bedrock' and 'mistral' are supported but require `model` to be passed so the provider config can be loaded."
),
model="n/a",
llm_provider=custom_llm_provider,
@ -514,7 +518,7 @@ def _handle_retrieve_batch_providers_without_provider_config(
def retrieve_batch(
batch_id: str,
custom_llm_provider: Literal[
"openai", "azure", "vertex_ai", "bedrock", "hosted_vllm", "litellm_proxy", "anthropic"
"openai", "azure", "vertex_ai", "bedrock", "hosted_vllm", "litellm_proxy", "anthropic", "mistral"
] = "openai",
metadata: dict[str, str] | None = None,
extra_headers: dict[str, str] | None = None,

View file

@ -139,6 +139,7 @@ if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import (
Logging as LitellmLoggingObject,
)
from litellm.llms.base_llm.ocr.transformation import OCRUsageInfo
else:
LitellmLoggingObject = Any
@ -1982,6 +1983,66 @@ def ocr_cost(
return ocr_pages_cost + annotation_pages_cost, 0.0
_OCR_PRICING_KEYS: Final = (
"ocr_cost_per_page",
"ocr_cost_per_page_batches",
"annotation_cost_per_page",
"annotation_cost_per_page_batches",
)
def ocr_batch_cost(
model: str,
custom_llm_provider: str | None,
usage_info: "OCRUsageInfo",
model_info: ModelInfo | None = None,
) -> tuple[float, float]:
"""Per-page cost of one OCR result inside a batch output file.
Batch OCR is billed per page at the ``*_batches`` rate, falling back to the
synchronous per-page rate when a model has no batch price recorded, the same
fallback ``batch_cost_calculator`` applies to per-token batch pricing. Returns
``(prompt_cost, completion_cost)`` with the whole cost in the first slot, like
``ocr_cost``.
"""
has_ocr_pricing: Final = model_info is not None and any(model_info.get(k) is not None for k in _OCR_PRICING_KEYS)
if has_ocr_pricing:
resolved_info: ModelInfo | None = model_info
else:
try:
resolved_info = litellm.get_model_info(model=model, custom_llm_provider=custom_llm_provider)
except Exception:
resolved_info = None
if resolved_info is None:
verbose_logger.warning(
"OCR batch cost: model=%s custom_llm_provider=%s has no pricing entry; returning 0.0 cost.",
model,
custom_llm_provider,
)
return 0.0, 0.0
page_rate: Final = _first_price(resolved_info, "ocr_cost_per_page_batches", "ocr_cost_per_page")
annotation_rate: Final = _first_price(
resolved_info, "annotation_cost_per_page_batches", "annotation_cost_per_page"
)
pages_processed: Final = usage_info.pages_processed or 0
annotation_pages: Final = usage_info.pages_processed_annotation or 0
if page_rate is None and pages_processed > 0:
verbose_logger.warning(
"OCR batch cost: model=%s custom_llm_provider=%s reported pages_processed=%s but no "
"ocr_cost_per_page is configured; returning 0.0 cost for those pages.",
model,
custom_llm_provider,
pages_processed,
)
effective_annotation_rate: Final = annotation_rate if annotation_rate is not None else page_rate
return (page_rate or 0.0) * pages_processed + (effective_annotation_rate or 0.0) * annotation_pages, 0.0
def _first_price(model_info: ModelInfo, *keys: str) -> float | None:
return next((price for price in (model_info.get(k) for k in keys) if isinstance(price, (int, float))), None)
def vector_store_search_cost(
model: str | None,
custom_llm_provider: str,

View file

@ -27,12 +27,15 @@ FileCreateProvider = Literal[
"litellm_proxy",
"manus",
"anthropic",
"mistral",
]
FileRetrieveProvider = Literal[
"openai", "azure", "gemini", "vertex_ai", "hosted_vllm", "litellm_proxy", "manus", "anthropic"
"openai", "azure", "gemini", "vertex_ai", "hosted_vllm", "litellm_proxy", "manus", "anthropic", "mistral"
]
FileDeleteProvider = Literal["openai", "azure", "gemini", "bedrock", "litellm_proxy", "manus", "anthropic"]
FileListProvider = Literal["openai", "azure", "litellm_proxy", "manus", "anthropic"]
FileDeleteProvider = Literal[
"openai", "azure", "gemini", "bedrock", "litellm_proxy", "manus", "anthropic", "mistral"
]
FileListProvider = Literal["openai", "azure", "litellm_proxy", "manus", "anthropic", "mistral"]
import litellm
from litellm import get_secret_str
from litellm.files.streaming import FileContentStreamingResponse

View file

@ -2,7 +2,7 @@ from collections.abc import AsyncIterator, Iterator
from typing import Literal, NamedTuple
FileContentProvider = Literal[
"openai", "azure", "vertex_ai", "bedrock", "hosted_vllm", "litellm_proxy", "anthropic", "manus"
"openai", "azure", "vertex_ai", "bedrock", "hosted_vllm", "litellm_proxy", "anthropic", "manus", "mistral"
]

View file

View file

@ -0,0 +1,186 @@
"""
Mistral Batch API. Reference: https://docs.mistral.ai/api/#tag/batch
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.
"""
from types import MappingProxyType
from typing import Final, Literal
import httpx
from openai.types.batch import BatchRequestCounts
from openai.types.batch import Errors as BatchErrors
from openai.types.batch_error import BatchError
from pydantic import BaseModel, ConfigDict
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.chat.transformation import BaseLLMException
from litellm.types.llms.openai import AllMessageValues, CreateBatchRequest
from litellm.types.utils import LiteLLMBatch, LlmProviders
from ..common_utils import get_mistral_api_base, get_mistral_auth_headers, mistral_error
MistralBatchStatus = Literal[
"QUEUED", "RUNNING", "SUCCESS", "FAILED", "TIMEOUT_EXCEEDED", "CANCELLATION_REQUESTED", "CANCELLED"
]
OpenAIBatchStatus = Literal[
"validating", "failed", "in_progress", "finalizing", "completed", "expired", "cancelling", "cancelled"
]
_STATUS_MAP: Final[MappingProxyType[MistralBatchStatus, OpenAIBatchStatus]] = MappingProxyType(
{
"QUEUED": "validating",
"RUNNING": "in_progress",
"SUCCESS": "completed",
"FAILED": "failed",
"TIMEOUT_EXCEEDED": "expired",
"CANCELLATION_REQUESTED": "cancelling",
"CANCELLED": "cancelled",
}
)
class MistralBatchError(BaseModel):
model_config = ConfigDict(frozen=True, extra="ignore")
message: str
count: int = 1
class MistralBatchJob(BaseModel):
model_config = ConfigDict(frozen=True, extra="ignore")
id: str
input_files: tuple[str, ...] = ()
endpoint: str
model: str | None = None
status: MistralBatchStatus
created_at: int
started_at: int | None = None
completed_at: int | None = None
total_requests: int = 0
completed_requests: int = 0
succeeded_requests: int = 0
failed_requests: int = 0
output_file: str | None = None
error_file: str | None = None
errors: tuple[MistralBatchError, ...] = ()
metadata: dict[str, str] | None = None
def _to_litellm_batch(job: MistralBatchJob) -> LiteLLMBatch:
status: Final = _STATUS_MAP[job.status]
terminal_at: Final = job.completed_at
return LiteLLMBatch(
id=job.id,
object="batch",
endpoint=job.endpoint,
input_file_id=job.input_files[0] if job.input_files else "",
completion_window="24h",
status=status,
created_at=job.created_at,
in_progress_at=job.started_at,
completed_at=terminal_at if status == "completed" else None,
failed_at=terminal_at if status == "failed" else None,
expired_at=terminal_at if status == "expired" else None,
cancelled_at=terminal_at if status == "cancelled" else None,
output_file_id=job.output_file,
error_file_id=job.error_file,
errors=(
BatchErrors(
object="list",
data=[BatchError(message=f"{e.message} (x{e.count})" if e.count > 1 else e.message) for e in job.errors],
)
if job.errors
else None
),
request_counts=BatchRequestCounts(
total=job.total_requests,
completed=job.succeeded_requests,
failed=job.failed_requests,
),
metadata=job.metadata,
)
class MistralBatchesConfig(BaseBatchesConfig):
@property
def custom_llm_provider(self) -> LlmProviders:
return LlmProviders.MISTRAL
def validate_environment(
self,
headers: dict,
model: str,
messages: list[AllMessageValues],
optional_params: dict,
litellm_params: dict,
api_key: str | None = None,
api_base: str | None = None,
) -> dict:
return get_mistral_auth_headers(headers, api_key)
def get_complete_batch_url(
self,
api_base: str | None,
api_key: str | None,
model: str,
optional_params: dict,
litellm_params: dict,
data: CreateBatchRequest,
) -> str:
return f"{get_mistral_api_base(api_base)}/v1/batch/jobs"
def transform_create_batch_request(
self,
model: str,
create_batch_data: CreateBatchRequest,
optional_params: dict,
litellm_params: dict,
) -> dict[str, object]:
metadata: Final = create_batch_data.get("metadata")
return {
"input_files": [create_batch_data["input_file_id"]],
"endpoint": create_batch_data["endpoint"],
"model": model,
**({"metadata": metadata} if metadata else {}),
**(create_batch_data.get("extra_body") or {}),
}
def transform_create_batch_response(
self,
model: str | None,
raw_response: httpx.Response,
logging_obj: object,
litellm_params: dict,
) -> LiteLLMBatch:
return _to_litellm_batch(MistralBatchJob.model_validate(raw_response.json()))
def transform_retrieve_batch_request(
self,
batch_id: str,
optional_params: dict,
litellm_params: dict,
) -> dict[str, object]:
encoded_batch_id: Final = encode_url_path_segment(batch_id, field_name="batch_id")
return {
"method": "GET",
"url": f"{get_mistral_api_base(litellm_params.get('api_base'))}/v1/batch/jobs/{encoded_batch_id}",
"headers": get_mistral_auth_headers({}, litellm_params.get("api_key")),
}
def transform_retrieve_batch_response(
self,
model: str | None,
raw_response: httpx.Response,
logging_obj: object,
litellm_params: dict,
) -> LiteLLMBatch:
return _to_litellm_batch(MistralBatchJob.model_validate(raw_response.json()))
def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> BaseLLMException:
return mistral_error(error_message, status_code, headers)

View file

@ -0,0 +1,36 @@
from typing import Final
import httpx
from litellm.llms.base_llm.chat.transformation import BaseLLMException
from litellm.secret_managers.main import get_secret_str
MISTRAL_API_BASE: Final = "https://api.mistral.ai"
MISTRAL_API_KEY_ENV_VAR: Final = "MISTRAL_API_KEY"
class MistralError(BaseLLMException):
pass
def get_mistral_api_base(api_base: str | None) -> str:
"""Return the Mistral origin without a trailing ``/v1``, so callers can append ``/v1/<route>``."""
resolved: Final = (api_base or get_secret_str("MISTRAL_API_BASE") or MISTRAL_API_BASE).rstrip("/")
return resolved.removesuffix("/v1")
def get_mistral_auth_headers(headers: dict, api_key: str | None) -> dict:
resolved_key: Final = api_key or get_secret_str(MISTRAL_API_KEY_ENV_VAR)
if resolved_key is None:
raise ValueError(
"Missing Mistral API Key - A call is being made to Mistral but no key is set either in the environment variables or via params"
)
return {**headers, "Authorization": f"Bearer {resolved_key}"}
def mistral_error(error_message: str, status_code: int, headers: dict | httpx.Headers) -> MistralError:
return MistralError(
status_code=status_code,
message=error_message,
headers=headers if isinstance(headers, httpx.Headers) else httpx.Headers(headers),
)

View file

View file

@ -0,0 +1,226 @@
"""
Mistral Files API. Reference: https://docs.mistral.ai/api/#tag/files
Mistral's file objects already carry the OpenAI field names (id, bytes, created_at,
filename, purpose), so this config is URL routing, auth, and a purpose mapping:
Mistral only accepts ``fine-tune``, ``batch`` and ``ocr`` as upload purposes.
"""
import time
from typing import Final, Literal
import httpx
from openai.types.file_deleted import FileDeleted
from pydantic import BaseModel, ConfigDict
from litellm.litellm_core_utils.prompt_templates.common_utils import extract_file_data
from litellm.litellm_core_utils.url_utils import encode_url_path_segment
from litellm.llms.base_llm.chat.transformation import BaseLLMException
from litellm.llms.base_llm.files.transformation import BaseFilesConfig, LiteLLMLoggingObj
from litellm.types.llms.openai import (
CreateFileRequest,
FileContentRequest,
HttpxBinaryResponseContent,
OpenAICreateFileRequestOptionalParams,
OpenAIFileObject,
OpenAIFilesPurpose,
)
from litellm.types.utils import LlmProviders
from ..common_utils import get_mistral_api_base, get_mistral_auth_headers, mistral_error
MistralFilePurpose = Literal["fine-tune", "batch", "ocr"]
class MistralFile(BaseModel):
model_config = ConfigDict(frozen=True, extra="ignore")
id: str
bytes: int = 0
created_at: int | None = None
filename: str = ""
purpose: MistralFilePurpose = "batch"
expires_at: int | None = None
class MistralFileList(BaseModel):
model_config = ConfigDict(frozen=True, extra="ignore")
data: tuple[MistralFile, ...] = ()
class MistralFileDeleted(BaseModel):
model_config = ConfigDict(frozen=True, extra="ignore")
id: str
deleted: bool = True
def _to_openai_file_object(file: MistralFile) -> OpenAIFileObject:
return OpenAIFileObject(
id=file.id,
bytes=file.bytes,
created_at=file.created_at if file.created_at is not None else int(time.time()),
filename=file.filename,
object="file",
purpose=_to_openai_purpose(file.purpose),
status="uploaded",
expires_at=file.expires_at,
)
def _to_openai_purpose(purpose: MistralFilePurpose) -> OpenAIFilesPurpose:
match purpose:
case "fine-tune" | "batch":
return purpose
case "ocr":
return "user_data"
def _to_mistral_purpose(purpose: str) -> MistralFilePurpose:
match purpose:
case "fine-tune" | "ocr":
return purpose
case _:
return "batch"
class MistralFilesConfig(BaseFilesConfig):
@property
def custom_llm_provider(self) -> LlmProviders:
return LlmProviders.MISTRAL
def get_complete_url(
self,
api_base: str | None,
api_key: str | None,
model: str,
optional_params: dict,
litellm_params: dict,
stream: bool | None = None,
) -> str:
return f"{get_mistral_api_base(api_base)}/v1/files"
def _file_url(self, file_id: str, litellm_params: dict, suffix: str = "") -> str:
encoded_file_id: Final = encode_url_path_segment(file_id, field_name="file_id")
return f"{get_mistral_api_base(litellm_params.get('api_base'))}/v1/files/{encoded_file_id}{suffix}"
def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> BaseLLMException:
return mistral_error(error_message, status_code, headers)
def validate_environment(
self,
headers: dict,
model: str,
messages: list,
optional_params: dict,
litellm_params: dict,
api_key: str | None = None,
api_base: str | None = None,
) -> dict:
return get_mistral_auth_headers(headers, api_key)
def get_supported_openai_params(self, model: str) -> list[OpenAICreateFileRequestOptionalParams]:
return ["purpose"]
def map_openai_params(
self,
non_default_params: dict,
optional_params: dict,
model: str,
drop_params: bool,
) -> dict:
return optional_params
def transform_create_file_request(
self,
model: str,
create_file_data: CreateFileRequest,
optional_params: dict,
litellm_params: dict,
) -> dict:
file_data: Final = create_file_data.get("file")
if file_data is None:
raise ValueError("File data is required")
extracted: Final = extract_file_data(file_data)
filename: Final = extracted["filename"] or f"file_{int(time.time())}.jsonl"
content_type: Final = extracted.get("content_type") or "application/octet-stream"
return {
"file": (filename, extracted["content"], content_type),
"purpose": (None, _to_mistral_purpose(create_file_data.get("purpose", "batch"))),
}
def transform_create_file_response(
self,
model: str | None,
raw_response: httpx.Response,
logging_obj: LiteLLMLoggingObj,
litellm_params: dict,
) -> OpenAIFileObject:
return _to_openai_file_object(MistralFile.model_validate(raw_response.json()))
def transform_retrieve_file_request(
self,
file_id: str,
optional_params: dict,
litellm_params: dict,
) -> tuple[str, dict]:
return self._file_url(file_id, litellm_params), {}
def transform_retrieve_file_response(
self,
raw_response: httpx.Response,
logging_obj: LiteLLMLoggingObj,
litellm_params: dict,
) -> OpenAIFileObject:
return _to_openai_file_object(MistralFile.model_validate(raw_response.json()))
def transform_delete_file_request(
self,
file_id: str,
optional_params: dict,
litellm_params: dict,
) -> tuple[str, dict]:
return self._file_url(file_id, litellm_params), {}
def transform_delete_file_response(
self,
raw_response: httpx.Response,
logging_obj: LiteLLMLoggingObj,
litellm_params: dict,
) -> FileDeleted:
deleted: Final = MistralFileDeleted.model_validate(raw_response.json())
return FileDeleted(id=deleted.id, deleted=deleted.deleted, object="file")
def transform_list_files_request(
self,
purpose: str | None,
optional_params: dict,
litellm_params: dict,
) -> tuple[str, dict]:
params: Final = {"purpose": _to_mistral_purpose(purpose)} if purpose else {}
return f"{get_mistral_api_base(litellm_params.get('api_base'))}/v1/files", params
def transform_list_files_response(
self,
raw_response: httpx.Response,
logging_obj: LiteLLMLoggingObj,
litellm_params: dict,
) -> list[OpenAIFileObject]:
return [_to_openai_file_object(f) for f in MistralFileList.model_validate(raw_response.json()).data]
def transform_file_content_request(
self,
file_content_request: FileContentRequest,
optional_params: dict,
litellm_params: dict,
) -> tuple[str, dict]:
return self._file_url(file_content_request["file_id"], litellm_params, suffix="/content"), {}
def transform_file_content_response(
self,
raw_response: httpx.Response,
logging_obj: LiteLLMLoggingObj,
litellm_params: dict,
) -> HttpxBinaryResponseContent:
return HttpxBinaryResponseContent(response=raw_response)

View file

@ -35262,51 +35262,66 @@
"mistral/mistral-ocr-latest": {
"litellm_provider": "mistral",
"ocr_cost_per_page": 0.004,
"ocr_cost_per_page_batches": 0.002,
"annotation_cost_per_page": 0.005,
"annotation_cost_per_page_batches": 0.0025,
"mode": "ocr",
"supported_endpoints": [
"/v1/ocr"
"/v1/ocr",
"/v1/batch"
],
"source": "https://mistral.ai/pricing#api-pricing"
},
"mistral/mistral-ocr-4-0": {
"litellm_provider": "mistral",
"ocr_cost_per_page": 0.004,
"ocr_cost_per_page_batches": 0.002,
"annotation_cost_per_page": 0.005,
"annotation_cost_per_page_batches": 0.0025,
"mode": "ocr",
"supported_endpoints": [
"/v1/ocr"
"/v1/ocr",
"/v1/batch"
],
"source": "https://mistral.ai/pricing#api-pricing"
},
"mistral/mistral-ocr-4-1": {
"annotation_cost_per_page": 0.005,
"annotation_cost_per_page_batches": 0.0025,
"litellm_provider": "mistral",
"mode": "ocr",
"ocr_cost_per_page": 0.004,
"ocr_cost_per_page_batches": 0.002,
"source": "https://docs.mistral.ai/models/model-cards/ocr-4-1",
"supported_endpoints": [
"/v1/ocr"
"/v1/ocr",
"/v1/batch"
]
},
"mistral/mistral-ocr-2505-completion": {
"deprecation_date": "2026-05-31",
"litellm_provider": "mistral",
"ocr_cost_per_page": 0.001,
"ocr_cost_per_page_batches": 0.0005,
"annotation_cost_per_page": 0.003,
"annotation_cost_per_page_batches": 0.0015,
"mode": "ocr",
"supported_endpoints": [
"/v1/ocr"
"/v1/ocr",
"/v1/batch"
],
"source": "https://mistral.ai/pricing#api-pricing"
},
"mistral/mistral-ocr-2512": {
"litellm_provider": "mistral",
"ocr_cost_per_page": 0.002,
"ocr_cost_per_page_batches": 0.001,
"annotation_cost_per_page": 0.003,
"annotation_cost_per_page_batches": 0.0015,
"mode": "ocr",
"supported_endpoints": [
"/v1/ocr"
"/v1/ocr",
"/v1/batch"
],
"source": "https://mistral.ai/pricing#api-pricing"
},
@ -59822,31 +59837,40 @@
"mistral/mistral-ocr-3": {
"litellm_provider": "mistral",
"ocr_cost_per_page": 0.002,
"ocr_cost_per_page_batches": 0.001,
"annotation_cost_per_page": 0.003,
"annotation_cost_per_page_batches": 0.0015,
"mode": "ocr",
"supported_endpoints": [
"/v1/ocr"
"/v1/ocr",
"/v1/batch"
],
"source": "https://mistral.ai/pricing#api-pricing"
},
"mistral/mistral-ocr-3-0": {
"litellm_provider": "mistral",
"ocr_cost_per_page": 0.002,
"ocr_cost_per_page_batches": 0.001,
"annotation_cost_per_page": 0.003,
"annotation_cost_per_page_batches": 0.0015,
"mode": "ocr",
"supported_endpoints": [
"/v1/ocr"
"/v1/ocr",
"/v1/batch"
],
"source": "https://mistral.ai/pricing#api-pricing"
},
"mistral/mistral-ocr-4": {
"annotation_cost_per_page": 0.005,
"annotation_cost_per_page_batches": 0.0025,
"litellm_provider": "mistral",
"mode": "ocr",
"ocr_cost_per_page": 0.004,
"ocr_cost_per_page_batches": 0.002,
"source": "https://docs.mistral.ai/models/model-cards/ocr-4-1",
"supported_endpoints": [
"/v1/ocr"
"/v1/ocr",
"/v1/batch"
]
},
"mistral/voxtral-mini-latest": {

View file

@ -498,7 +498,7 @@ class CreateBatchRequest(TypedDict, total=False):
"""
completion_window: Literal["24h"]
endpoint: Literal["/v1/chat/completions", "/v1/embeddings", "/v1/completions", "/v1/responses"]
endpoint: Literal["/v1/chat/completions", "/v1/embeddings", "/v1/completions", "/v1/responses", "/v1/ocr"]
input_file_id: str
metadata: dict[str, str] | None
output_expires_after: FileExpiresAfter

View file

@ -320,8 +320,10 @@ class ModelInfoBase(ProviderSpecificModelInfo, total=False):
output_cost_per_second_480p: ReadOnly[float | None]
output_cost_per_second_4k: ReadOnly[float | None]
ocr_cost_per_page: float | None # for OCR models
ocr_cost_per_page_batches: ReadOnly[float | None]
ocr_cost_per_credit: float | None # for OCR models priced by credit
annotation_cost_per_page: float | None # for OCR models
annotation_cost_per_page_batches: ReadOnly[float | None]
search_context_cost_per_query: SearchContextCostPerQuery | None # Cost for using web search tool
web_search_billing_unit: (
Literal["per_query", "per_prompt"] | None
@ -3598,8 +3600,10 @@ class CustomPricingLiteLLMParams(MirroredPricingParams):
output_cost_per_token_above_512k_tokens: float | None = None
output_vector_size: int | None = None
ocr_cost_per_page: float | None = None
ocr_cost_per_page_batches: float | None = None
ocr_cost_per_credit: float | None = None
annotation_cost_per_page: float | None = None
annotation_cost_per_page_batches: float | None = None
regional_processing_uplift_multiplier_eu: float | None = None
regional_processing_uplift_multiplier_us: float | None = None
regional_endpoint_uplift_multiplier: float | None = None

View file

@ -5963,8 +5963,10 @@ def _get_model_info_helper(
tpm=_model_info.get("tpm", None),
rpm=_model_info.get("rpm", None),
ocr_cost_per_page=_model_info.get("ocr_cost_per_page", None),
ocr_cost_per_page_batches=_model_info.get("ocr_cost_per_page_batches", None),
ocr_cost_per_credit=_model_info.get("ocr_cost_per_credit", None),
annotation_cost_per_page=_model_info.get("annotation_cost_per_page", None),
annotation_cost_per_page_batches=_model_info.get("annotation_cost_per_page_batches", None),
provider_specific_entry=_model_info.get("provider_specific_entry", None),
uses_embed_content=_model_info.get("uses_embed_content", None),
supports_image_size=_model_info.get("supports_image_size", None),
@ -8909,6 +8911,10 @@ class ProviderConfigManager:
from litellm.llms.anthropic.files.transformation import AnthropicFilesConfig
return AnthropicFilesConfig()
elif LlmProviders.MISTRAL == provider:
from litellm.llms.mistral.files.transformation import MistralFilesConfig
return MistralFilesConfig()
return None
@staticmethod
@ -8920,6 +8926,10 @@ class ProviderConfigManager:
from litellm.llms.bedrock.batches.transformation import BedrockBatchesConfig
return BedrockBatchesConfig()
elif LlmProviders.MISTRAL == provider:
from litellm.llms.mistral.batches.transformation import MistralBatchesConfig
return MistralBatchesConfig()
return None
@staticmethod

View file

@ -35262,51 +35262,66 @@
"mistral/mistral-ocr-latest": {
"litellm_provider": "mistral",
"ocr_cost_per_page": 0.004,
"ocr_cost_per_page_batches": 0.002,
"annotation_cost_per_page": 0.005,
"annotation_cost_per_page_batches": 0.0025,
"mode": "ocr",
"supported_endpoints": [
"/v1/ocr"
"/v1/ocr",
"/v1/batch"
],
"source": "https://mistral.ai/pricing#api-pricing"
},
"mistral/mistral-ocr-4-0": {
"litellm_provider": "mistral",
"ocr_cost_per_page": 0.004,
"ocr_cost_per_page_batches": 0.002,
"annotation_cost_per_page": 0.005,
"annotation_cost_per_page_batches": 0.0025,
"mode": "ocr",
"supported_endpoints": [
"/v1/ocr"
"/v1/ocr",
"/v1/batch"
],
"source": "https://mistral.ai/pricing#api-pricing"
},
"mistral/mistral-ocr-4-1": {
"annotation_cost_per_page": 0.005,
"annotation_cost_per_page_batches": 0.0025,
"litellm_provider": "mistral",
"mode": "ocr",
"ocr_cost_per_page": 0.004,
"ocr_cost_per_page_batches": 0.002,
"source": "https://docs.mistral.ai/models/model-cards/ocr-4-1",
"supported_endpoints": [
"/v1/ocr"
"/v1/ocr",
"/v1/batch"
]
},
"mistral/mistral-ocr-2505-completion": {
"deprecation_date": "2026-05-31",
"litellm_provider": "mistral",
"ocr_cost_per_page": 0.001,
"ocr_cost_per_page_batches": 0.0005,
"annotation_cost_per_page": 0.003,
"annotation_cost_per_page_batches": 0.0015,
"mode": "ocr",
"supported_endpoints": [
"/v1/ocr"
"/v1/ocr",
"/v1/batch"
],
"source": "https://mistral.ai/pricing#api-pricing"
},
"mistral/mistral-ocr-2512": {
"litellm_provider": "mistral",
"ocr_cost_per_page": 0.002,
"ocr_cost_per_page_batches": 0.001,
"annotation_cost_per_page": 0.003,
"annotation_cost_per_page_batches": 0.0015,
"mode": "ocr",
"supported_endpoints": [
"/v1/ocr"
"/v1/ocr",
"/v1/batch"
],
"source": "https://mistral.ai/pricing#api-pricing"
},
@ -59822,31 +59837,40 @@
"mistral/mistral-ocr-3": {
"litellm_provider": "mistral",
"ocr_cost_per_page": 0.002,
"ocr_cost_per_page_batches": 0.001,
"annotation_cost_per_page": 0.003,
"annotation_cost_per_page_batches": 0.0015,
"mode": "ocr",
"supported_endpoints": [
"/v1/ocr"
"/v1/ocr",
"/v1/batch"
],
"source": "https://mistral.ai/pricing#api-pricing"
},
"mistral/mistral-ocr-3-0": {
"litellm_provider": "mistral",
"ocr_cost_per_page": 0.002,
"ocr_cost_per_page_batches": 0.001,
"annotation_cost_per_page": 0.003,
"annotation_cost_per_page_batches": 0.0015,
"mode": "ocr",
"supported_endpoints": [
"/v1/ocr"
"/v1/ocr",
"/v1/batch"
],
"source": "https://mistral.ai/pricing#api-pricing"
},
"mistral/mistral-ocr-4": {
"annotation_cost_per_page": 0.005,
"annotation_cost_per_page_batches": 0.0025,
"litellm_provider": "mistral",
"mode": "ocr",
"ocr_cost_per_page": 0.004,
"ocr_cost_per_page_batches": 0.002,
"source": "https://docs.mistral.ai/models/model-cards/ocr-4-1",
"supported_endpoints": [
"/v1/ocr"
"/v1/ocr",
"/v1/batch"
]
},
"mistral/voxtral-mini-latest": {

View file

@ -1787,3 +1787,86 @@ class TestBatchCostIsFinal:
@pytest.mark.parametrize("status", ["failed", "expired", "cancelled"])
def test_other_terminal_statuses_are_final(self, status):
assert bu.batch_cost_is_final(_retrieved_batch(status)) is True
# =========================================================================== #
# OCR batch output lines (Mistral /v1/ocr batches) are billed per page, not per token
# =========================================================================== #
def _ocr_row(pages_processed, annotation_pages=None, model="mistral-ocr-latest"):
usage_info = {"pages_processed": pages_processed, "doc_size_bytes": 4096}
if annotation_pages is not None:
usage_info["pages_processed_annotation"] = annotation_pages
return _success_row(model=model, pages=[{"index": i, "markdown": "x"} for i in range(pages_processed)], usage_info=usage_info)
def test_ocr_rows_are_priced_per_page_at_batch_rate(monkeypatch):
monkeypatch.setattr(
litellm,
"get_model_info",
lambda model, custom_llm_provider=None: {"ocr_cost_per_page": 0.004, "ocr_cost_per_page_batches": 0.002},
)
result = bu._aggregate_batch_cost_usage_models(
entries=[_ocr_row(3), _ocr_row(5), _failed_row(model="mistral-ocr-latest")],
custom_llm_provider="mistral",
model_name="mistral/mistral-ocr-latest",
)
assert result.cost == pytest.approx(8 * 0.002)
assert result.prompt_cost == pytest.approx(8 * 0.002)
assert result.completion_cost == 0.0
assert (result.successful_requests, result.failed_requests) == (2, 1)
assert result.usage.total_tokens == 0
assert result.models == ["mistral/mistral-ocr-latest"]
def test_ocr_rows_fall_back_to_sync_page_rate_without_batch_price(monkeypatch):
monkeypatch.setattr(litellm, "get_model_info", lambda model, custom_llm_provider=None: {"ocr_cost_per_page": 0.004})
result = bu._aggregate_batch_cost_usage_models(entries=[_ocr_row(2)], custom_llm_provider="mistral")
assert result.cost == pytest.approx(2 * 0.004)
def test_ocr_rows_bill_annotation_pages_separately(monkeypatch):
monkeypatch.setattr(
litellm,
"get_model_info",
lambda model, custom_llm_provider=None: {
"ocr_cost_per_page_batches": 0.002,
"annotation_cost_per_page_batches": 0.0025,
},
)
result = bu._aggregate_batch_cost_usage_models(entries=[_ocr_row(4, annotation_pages=4)], custom_llm_provider="mistral")
assert result.cost == pytest.approx(4 * 0.002 + 4 * 0.0025)
def test_ocr_rows_use_deployment_model_info_pricing_over_cost_map(monkeypatch):
monkeypatch.setattr(
litellm, "get_model_info", lambda model, custom_llm_provider=None: pytest.fail("cost map must not be consulted")
)
result = bu._aggregate_batch_cost_usage_models(
entries=[_ocr_row(10)],
custom_llm_provider="mistral",
model_info={"ocr_cost_per_page_batches": 0.001},
)
assert result.cost == pytest.approx(0.01)
def test_ocr_rows_without_pricing_bill_zero_but_count_as_successful(monkeypatch):
monkeypatch.setattr(litellm, "get_model_info", lambda model, custom_llm_provider=None: {"mode": "ocr"})
result = bu._aggregate_batch_cost_usage_models(entries=[_ocr_row(3)], custom_llm_provider="mistral")
assert result.cost == 0.0
assert (result.successful_requests, result.failed_requests) == (1, 0)
def test_chat_rows_from_mistral_still_use_token_pricing(monkeypatch):
monkeypatch.setattr(
litellm,
"get_model_info",
lambda model, custom_llm_provider=None: {"input_cost_per_token": 0.001, "output_cost_per_token": 0.002},
)
result = bu._aggregate_batch_cost_usage_models(
entries=[_success_row(model="mistral-small-latest", usage=_usage(10, 5))],
custom_llm_provider="mistral",
)
assert result.cost == pytest.approx((10 * 0.001 + 5 * 0.002) / 2)
assert result.usage.total_tokens == 15

View file

@ -778,3 +778,45 @@ def test_retrieve__omits_trusted_model_credentials_when_not_supplied(seams):
litellm_params = logging_obj.update_from_kwargs.call_args.kwargs["litellm_params"]
assert "_litellm_internal_model_credentials" not in litellm_params
# =========================================================================== #
# mistral - a provider-config provider, like bedrock, so it requires `model`
# =========================================================================== #
def test_create__mistral_ocr_routes_to_base_http_handler_with_mistral_config(seams):
with patch.object(bm.ProviderConfigManager, "get_provider_batches_config", wraps=bm.ProviderConfigManager.get_provider_batches_config) as get_cfg:
result = bm.create_batch(
completion_window="24h",
endpoint="/v1/ocr",
input_file_id="file-abc",
custom_llm_provider="mistral",
model="mistral/mistral-ocr-latest",
)
assert result is seams.base_http.create_batch.return_value
_assert_only(seams.base_http.create_batch, seams, "create_batch")
get_cfg.assert_called_once()
forwarded = seams.base_http.create_batch.call_args.kwargs
assert type(forwarded["provider_config"]).__name__ == "MistralBatchesConfig"
assert forwarded["model"] == "mistral-ocr-latest"
assert forwarded["create_batch_data"]["endpoint"] == "/v1/ocr"
def test_create__mistral_without_model_raises_badrequest(seams):
with pytest.raises(litellm.exceptions.BadRequestError):
bm.create_batch(**CREATE_KW, custom_llm_provider="mistral")
for m in _all_seam_methods(seams, "create_batch"):
m.assert_not_called()
def test_retrieve__mistral_routes_to_base_http_handler_with_mistral_config(seams):
result = bm.retrieve_batch(batch_id="job-1", custom_llm_provider="mistral", model="mistral/mistral-ocr-latest")
assert result is seams.base_http.retrieve_batch.return_value
_assert_only(seams.base_http.retrieve_batch, seams, "retrieve_batch")
forwarded = seams.base_http.retrieve_batch.call_args.kwargs
assert type(forwarded["provider_config"]).__name__ == "MistralBatchesConfig"
assert forwarded["batch_id"] == "job-1"

View file

@ -0,0 +1,260 @@
"""
Regression tests for ``MistralBatchesConfig``, the BaseBatchesConfig implementation
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.
Everything runs for real against canned httpx responses; only the API key env var is
set.
"""
import json
import httpx
import pytest
from litellm.llms.mistral.batches.transformation import MistralBatchesConfig
from litellm.llms.mistral.common_utils import MistralError
from litellm.types.llms.openai import CreateBatchRequest
from litellm.types.utils import LiteLLMBatch, LlmProviders
STATUS_MAP = {
"QUEUED": "validating",
"RUNNING": "in_progress",
"SUCCESS": "completed",
"FAILED": "failed",
"TIMEOUT_EXCEEDED": "expired",
"CANCELLATION_REQUESTED": "cancelling",
"CANCELLED": "cancelled",
}
def _job(**overrides):
base = {
"id": "8ff5e0d1-6bc2-4c3a-9f7d-0d1c2e3f4a5b",
"object": "batch",
"input_files": ["c1a2b3d4-0000-4000-8000-000000000001"],
"endpoint": "/v1/ocr",
"model": "mistral-ocr-latest",
"status": "SUCCESS",
"created_at": 1_757_400_000,
"started_at": 1_757_400_010,
"completed_at": 1_757_400_500,
"total_requests": 3,
"completed_requests": 3,
"succeeded_requests": 2,
"failed_requests": 1,
"output_file": "out-0000-4000-8000-000000000002",
"error_file": "err-0000-4000-8000-000000000003",
"errors": [],
"metadata": {"job_type": "testing"},
}
return {**base, **overrides}
def _response(payload: dict, status_code: int = 200) -> httpx.Response:
return httpx.Response(
status_code=status_code,
content=json.dumps(payload).encode(),
request=httpx.Request("GET", "https://api.mistral.ai/v1/batch/jobs/x"),
)
@pytest.fixture
def config() -> MistralBatchesConfig:
return MistralBatchesConfig()
@pytest.fixture
def api_key(monkeypatch) -> str:
monkeypatch.setenv("MISTRAL_API_KEY", "sk-mistral-test")
return "sk-mistral-test"
def test_custom_llm_provider(config):
assert config.custom_llm_provider == LlmProviders.MISTRAL
# --------------------------------------------------------------------------- #
# create
# --------------------------------------------------------------------------- #
def test_create_request_maps_openai_fields_onto_mistral_job(config):
data = CreateBatchRequest(
completion_window="24h",
endpoint="/v1/ocr",
input_file_id="file-123",
metadata={"team": "docs"},
)
body = config.transform_create_batch_request(
model="mistral-ocr-latest", create_batch_data=data, optional_params={}, litellm_params={}
)
assert body == {
"input_files": ["file-123"],
"endpoint": "/v1/ocr",
"model": "mistral-ocr-latest",
"metadata": {"team": "docs"},
}
def test_create_request_omits_empty_metadata_and_forwards_extra_body(config):
data = CreateBatchRequest(
completion_window="24h",
endpoint="/v1/chat/completions",
input_file_id="file-123",
metadata=None,
extra_body={"timeout_hours": 48},
)
body = config.transform_create_batch_request(
model="mistral-small-latest", create_batch_data=data, optional_params={}, litellm_params={}
)
assert "metadata" not in body
assert body["timeout_hours"] == 48
@pytest.mark.parametrize(
"api_base,expected",
[
(None, "https://api.mistral.ai/v1/batch/jobs"),
("https://api.mistral.ai/v1", "https://api.mistral.ai/v1/batch/jobs"),
("https://proxy.example.com/", "https://proxy.example.com/v1/batch/jobs"),
],
)
def test_create_url(config, api_base, expected):
url = config.get_complete_batch_url(
api_base=api_base, api_key="k", model="m", optional_params={}, litellm_params={}, data={}
)
assert url == expected
def test_validate_environment_uses_bearer_auth(config, api_key):
headers = config.validate_environment(
headers={"x-extra": "1"}, model="m", messages=[], optional_params={}, litellm_params={}
)
assert headers == {"x-extra": "1", "Authorization": f"Bearer {api_key}"}
def test_validate_environment_explicit_key_wins(config, api_key):
headers = config.validate_environment(
headers={}, model="m", messages=[], optional_params={}, litellm_params={}, api_key="sk-explicit"
)
assert headers["Authorization"] == "Bearer sk-explicit"
def test_validate_environment_without_key_raises(config, monkeypatch):
monkeypatch.delenv("MISTRAL_API_KEY", raising=False)
with pytest.raises(ValueError, match="Missing Mistral API Key"):
config.validate_environment(headers={}, model="m", messages=[], optional_params={}, litellm_params={})
def test_create_response_maps_job_onto_openai_batch(config):
batch = config.transform_create_batch_response(
model="mistral-ocr-latest",
raw_response=_response(_job(status="QUEUED", started_at=None, completed_at=None)),
logging_obj=None,
litellm_params={},
)
assert isinstance(batch, LiteLLMBatch)
assert batch.id == "8ff5e0d1-6bc2-4c3a-9f7d-0d1c2e3f4a5b"
assert batch.endpoint == "/v1/ocr"
assert batch.input_file_id == "c1a2b3d4-0000-4000-8000-000000000001"
assert batch.status == "validating"
assert batch.created_at == 1_757_400_000
assert batch.in_progress_at is None
assert batch.completed_at is None
assert batch.metadata == {"job_type": "testing"}
# --------------------------------------------------------------------------- #
# retrieve
# --------------------------------------------------------------------------- #
def test_retrieve_request_is_presigned_get_with_auth(config, api_key):
req = config.transform_retrieve_batch_request(
batch_id="job/with slash", optional_params={}, litellm_params={"api_base": "https://api.mistral.ai"}
)
assert req["method"] == "GET"
assert req["url"] == "https://api.mistral.ai/v1/batch/jobs/job%2Fwith%20slash"
assert req["headers"] == {"Authorization": f"Bearer {api_key}"}
def test_retrieve_request_prefers_litellm_params_api_key(config, api_key):
req = config.transform_retrieve_batch_request(
batch_id="job-1", optional_params={}, litellm_params={"api_key": "sk-from-deployment"}
)
assert req["headers"]["Authorization"] == "Bearer sk-from-deployment"
@pytest.mark.parametrize("mistral_status,openai_status", sorted(STATUS_MAP.items()))
def test_retrieve_response_status_mapping(config, mistral_status, openai_status):
batch = config.transform_retrieve_batch_response(
model=None, raw_response=_response(_job(status=mistral_status)), logging_obj=None, litellm_params={}
)
assert batch.status == openai_status
@pytest.mark.parametrize(
"mistral_status,populated_field",
[
("SUCCESS", "completed_at"),
("FAILED", "failed_at"),
("TIMEOUT_EXCEEDED", "expired_at"),
("CANCELLED", "cancelled_at"),
],
)
def test_retrieve_response_terminal_timestamp_lands_on_matching_field(config, mistral_status, populated_field):
batch = config.transform_retrieve_batch_response(
model=None, raw_response=_response(_job(status=mistral_status)), logging_obj=None, litellm_params={}
)
terminal_fields = {"completed_at", "failed_at", "expired_at", "cancelled_at"}
assert getattr(batch, populated_field) == 1_757_400_500
for other in terminal_fields - {populated_field}:
assert getattr(batch, other) is None
assert batch.in_progress_at == 1_757_400_010
def test_retrieve_response_maps_counts_and_files(config):
batch = config.transform_retrieve_batch_response(
model=None, raw_response=_response(_job()), logging_obj=None, litellm_params={}
)
assert batch.request_counts.total == 3
assert batch.request_counts.completed == 2
assert batch.request_counts.failed == 1
assert batch.output_file_id == "out-0000-4000-8000-000000000002"
assert batch.error_file_id == "err-0000-4000-8000-000000000003"
assert batch.errors is None
def test_retrieve_response_surfaces_job_errors(config):
batch = config.transform_retrieve_batch_response(
model=None,
raw_response=_response(
_job(status="FAILED", errors=[{"message": "invalid document", "count": 2}, {"message": "timeout"}])
),
logging_obj=None,
litellm_params={},
)
assert [e.message for e in batch.errors.data] == ["invalid document (x2)", "timeout"]
def test_retrieve_response_without_files_or_input(config):
batch = config.transform_retrieve_batch_response(
model=None,
raw_response=_response(_job(input_files=[], output_file=None, error_file=None, metadata=None)),
logging_obj=None,
litellm_params={},
)
assert batch.input_file_id == ""
assert batch.output_file_id is None
assert batch.error_file_id is None
assert batch.metadata is None
def test_get_error_class(config):
err = config.get_error_class("nope", 401, {"x-request-id": "r1"})
assert isinstance(err, MistralError)
assert err.status_code == 401
assert err.message == "nope"

View file

@ -0,0 +1,189 @@
"""
Regression tests for ``MistralFilesConfig``, the BaseFilesConfig implementation behind
``custom_llm_provider="mistral"`` on /v1/files.
Locks the URL routing for each file operation, the multipart upload shape Mistral's
``POST /v1/files`` accepts (purpose restricted to fine-tune/batch/ocr), and the
Mistral -> OpenAI file object mapping. Runs against canned httpx responses.
"""
import json
import httpx
import pytest
from openai.types.file_deleted import FileDeleted
from litellm.llms.mistral.files.transformation import MistralFilesConfig
from litellm.types.llms.openai import CreateFileRequest, FileContentRequest, OpenAIFileObject
from litellm.types.utils import LlmProviders
FILE_ID = "497f6eca-6276-4993-bfeb-53cbbbba6f09"
def _file(**overrides):
base = {
"id": FILE_ID,
"object": "file",
"bytes": 13000,
"created_at": 1_716_963_433,
"filename": "batch_input.jsonl",
"purpose": "batch",
"sample_type": "batch_request",
"num_lines": 3,
"source": "upload",
}
return {**base, **overrides}
def _response(payload) -> httpx.Response:
return httpx.Response(
status_code=200,
content=json.dumps(payload).encode(),
request=httpx.Request("GET", "https://api.mistral.ai/v1/files"),
)
@pytest.fixture
def config() -> MistralFilesConfig:
return MistralFilesConfig()
@pytest.fixture
def api_key(monkeypatch) -> str:
monkeypatch.setenv("MISTRAL_API_KEY", "sk-mistral-test")
return "sk-mistral-test"
def test_custom_llm_provider(config):
assert config.custom_llm_provider == LlmProviders.MISTRAL
@pytest.mark.parametrize(
"api_base,expected",
[
(None, "https://api.mistral.ai/v1/files"),
("https://api.mistral.ai/v1/", "https://api.mistral.ai/v1/files"),
("https://proxy.example.com", "https://proxy.example.com/v1/files"),
],
)
def test_upload_url(config, api_base, expected):
url = config.get_complete_url(api_base=api_base, api_key="k", model="", optional_params={}, litellm_params={})
assert url == expected
def test_validate_environment_uses_bearer_auth(config, api_key):
headers = config.validate_environment(headers={}, model="", messages=[], optional_params={}, litellm_params={})
assert headers == {"Authorization": f"Bearer {api_key}"}
def test_upload_request_is_multipart_with_batch_purpose(config):
body = config.transform_create_file_request(
model="",
create_file_data=CreateFileRequest(file=("in.jsonl", b'{"custom_id":"0"}\n', "application/jsonl"), purpose="batch"),
optional_params={},
litellm_params={},
)
assert body == {
"file": ("in.jsonl", b'{"custom_id":"0"}\n', "application/jsonl"),
"purpose": (None, "batch"),
}
@pytest.mark.parametrize(
"openai_purpose,mistral_purpose",
[("batch", "batch"), ("fine-tune", "fine-tune"), ("ocr", "ocr"), ("assistants", "batch"), ("user_data", "batch")],
)
def test_upload_request_maps_purpose_onto_mistral_enum(config, openai_purpose, mistral_purpose):
body = config.transform_create_file_request(
model="",
create_file_data=CreateFileRequest(file=("f.bin", b"x"), purpose=openai_purpose),
optional_params={},
litellm_params={},
)
assert body["purpose"] == (None, mistral_purpose)
def test_upload_request_requires_file(config):
with pytest.raises(ValueError, match="File data is required"):
config.transform_create_file_request(
model="", create_file_data=CreateFileRequest(purpose="batch"), optional_params={}, litellm_params={}
)
def test_upload_response_maps_onto_openai_file_object(config):
obj = config.transform_create_file_response(
model=None, raw_response=_response(_file()), logging_obj=None, litellm_params={}
)
assert obj == OpenAIFileObject(
id=FILE_ID,
bytes=13000,
created_at=1_716_963_433,
filename="batch_input.jsonl",
object="file",
purpose="batch",
status="uploaded",
)
def test_file_response_with_ocr_purpose_maps_onto_user_data(config):
obj = config.transform_retrieve_file_response(
raw_response=_response(_file(purpose="ocr", expires_at=1_800_000_000)), logging_obj=None, litellm_params={}
)
assert obj.purpose == "user_data"
assert obj.expires_at == 1_800_000_000
@pytest.mark.parametrize(
"method,suffix",
[
("transform_retrieve_file_request", ""),
("transform_delete_file_request", ""),
],
)
def test_single_file_urls_encode_id_and_honor_api_base(config, method, suffix):
url, params = getattr(config, method)(
file_id="id/with slash", optional_params={}, litellm_params={"api_base": "https://mistral.internal/v1"}
)
assert url == f"https://mistral.internal/v1/files/id%2Fwith%20slash{suffix}"
assert params == {}
def test_file_content_url(config):
url, params = config.transform_file_content_request(
file_content_request=FileContentRequest(file_id=FILE_ID), optional_params={}, litellm_params={}
)
assert url == f"https://api.mistral.ai/v1/files/{FILE_ID}/content"
assert params == {}
def test_file_content_response_is_binary_passthrough(config):
raw = httpx.Response(
200, content=b'{"custom_id":"0","response":{"status_code":200}}\n', request=httpx.Request("GET", "https://x")
)
out = config.transform_file_content_response(raw_response=raw, logging_obj=None, litellm_params={})
assert out.content == b'{"custom_id":"0","response":{"status_code":200}}\n'
def test_delete_response(config):
out = config.transform_delete_file_response(
raw_response=_response({"id": FILE_ID, "object": "file", "deleted": True}), logging_obj=None, litellm_params={}
)
assert out == FileDeleted(id=FILE_ID, deleted=True, object="file")
def test_list_request_filters_by_mapped_purpose(config):
url, params = config.transform_list_files_request(purpose="batch", optional_params={}, litellm_params={})
assert url == "https://api.mistral.ai/v1/files"
assert params == {"purpose": "batch"}
_, no_params = config.transform_list_files_request(purpose=None, optional_params={}, litellm_params={})
assert no_params == {}
def test_list_response(config):
out = config.transform_list_files_response(
raw_response=_response({"data": [_file(), _file(id="second", filename="b.jsonl")], "object": "list", "total": 2}),
logging_obj=None,
litellm_params={},
)
assert [f.id for f in out] == [FILE_ID, "second"]
assert out[1].filename == "b.jsonl"

View file

@ -72,9 +72,11 @@ def test_ocr3_pricing_entry(cost_map_path: Path) -> None:
assert info is not None, f"{OCR3_MODEL} missing from {cost_map_path.name}"
assert info["litellm_provider"] == "mistral"
assert info["mode"] == "ocr"
assert info["supported_endpoints"] == ["/v1/ocr"]
assert info["supported_endpoints"] == ["/v1/ocr", "/v1/batch"]
assert info["ocr_cost_per_page"] == OCR3_COST_PER_PAGE
assert info["annotation_cost_per_page"] == OCR3_ANNOTATION_COST_PER_PAGE
assert info["ocr_cost_per_page_batches"] == OCR3_COST_PER_PAGE / 2
assert info["annotation_cost_per_page_batches"] == OCR3_ANNOTATION_COST_PER_PAGE / 2
def test_ocr3_model_info_price(local_model_cost_map) -> None:

View file

@ -992,7 +992,9 @@ def test_aaamodel_prices_and_context_window_json_is_valid():
"input_cost_per_video_per_second_above_128k_tokens": {"type": "number"},
"input_dbu_cost_per_token": {"type": "number"},
"annotation_cost_per_page": {"type": "number"},
"annotation_cost_per_page_batches": {"type": "number"},
"ocr_cost_per_page": {"type": "number"},
"ocr_cost_per_page_batches": {"type": "number"},
"ocr_cost_per_credit": {"type": "number"},
"code_interpreter_cost_per_session": {"type": "number"},
"inference_geo": {"type": "string"},