mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
Merge 897e2b1929 into b781d157d7
This commit is contained in:
commit
999652da84
13 changed files with 956 additions and 45 deletions
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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 (
|
||||
|
|
@ -349,14 +355,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 (
|
||||
|
|
@ -365,10 +373,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]:
|
||||
|
|
@ -4074,6 +4093,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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -6526,8 +6526,14 @@ class Router:
|
|||
|
||||
async def try_retrieve_batch(model: DeploymentTypedDict):
|
||||
try:
|
||||
# Update kwargs with the current model name or any other model-specific adjustments
|
||||
return await litellm.alist_batches(**{**model["litellm_params"], **kwargs})
|
||||
litellm_params: Final = model["litellm_params"]
|
||||
_, custom_llm_provider, _, _ = get_llm_provider(
|
||||
model=litellm_params["model"],
|
||||
custom_llm_provider=litellm_params.get("custom_llm_provider"),
|
||||
)
|
||||
return await litellm.alist_batches(
|
||||
**{**litellm_params, "custom_llm_provider": custom_llm_provider, **kwargs}
|
||||
)
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
|
|
@ -6540,6 +6546,7 @@ class Router:
|
|||
"first_id": None,
|
||||
"last_id": None,
|
||||
"has_more": False,
|
||||
"next_page_token": None,
|
||||
}
|
||||
|
||||
for result in results:
|
||||
|
|
@ -6549,6 +6556,9 @@ class Router:
|
|||
final_results["first_id"] = getattr(result, "first_id")
|
||||
final_results["last_id"] = getattr(result, "last_id")
|
||||
final_results["data"].extend(result.data)
|
||||
page_token = getattr(result, "next_page_token", None)
|
||||
if page_token is not None:
|
||||
final_results["next_page_token"] = page_token
|
||||
|
||||
## check 'has_more'
|
||||
if getattr(result, "has_more", False) is True:
|
||||
|
|
|
|||
|
|
@ -4172,7 +4172,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))
|
||||
|
||||
|
|
@ -4319,6 +4319,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
|
||||
|
|
|
|||
|
|
@ -9367,7 +9367,7 @@ class ProviderConfigManager:
|
|||
|
||||
@staticmethod
|
||||
def get_provider_batches_config(
|
||||
model: str,
|
||||
model: str | None,
|
||||
provider: LlmProviders,
|
||||
) -> BaseBatchesConfig | None:
|
||||
if LlmProviders.BEDROCK == provider:
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
||||
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")
|
||||
|
||||
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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
# =========================================================================== #
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
47
tests/unit/test_router_batch_list_pagination.py
Normal file
47
tests/unit/test_router_batch_list_pagination.py
Normal file
|
|
@ -0,0 +1,47 @@
|
|||
from typing import Final
|
||||
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
from litellm import Router
|
||||
from litellm.types.utils import OpenAIBatchListResponse
|
||||
|
||||
TOKEN_BY_KEY: Final = {"key-with-more-pages": "1", "key-on-last-page": None}
|
||||
|
||||
|
||||
def _deployment(api_key: str) -> dict:
|
||||
return {
|
||||
"model_name": "mistral-ocr",
|
||||
"litellm_params": {"model": "mistral/mistral-ocr-latest", "api_key": api_key},
|
||||
"model_info": {"id": api_key},
|
||||
}
|
||||
|
||||
|
||||
async def _fake_alist_batches(**kwargs: object) -> OpenAIBatchListResponse:
|
||||
if kwargs.get("custom_llm_provider") != "mistral":
|
||||
raise ValueError("a mistral key sent down the default openai list path")
|
||||
token: Final = TOKEN_BY_KEY[str(kwargs["api_key"])]
|
||||
return OpenAIBatchListResponse(
|
||||
data=(), first_id=None, last_id=None, has_more=token is not None, next_page_token=token
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_alist_batches_keeps_the_page_token_a_deployment_returned(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.setattr(litellm, "alist_batches", _fake_alist_batches)
|
||||
router: Final = Router(model_list=[_deployment("key-with-more-pages"), _deployment("key-on-last-page")])
|
||||
|
||||
result: Final = await router.alist_batches(model="mistral-ocr", limit=3)
|
||||
|
||||
assert result["has_more"] is True
|
||||
assert result["next_page_token"] == "1"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_alist_batches_lists_with_each_deployments_own_provider(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.setattr(litellm, "alist_batches", _fake_alist_batches)
|
||||
router: Final = Router(model_list=[_deployment("key-with-more-pages")])
|
||||
|
||||
result: Final = await router.alist_batches(model="mistral-ocr", limit=3)
|
||||
|
||||
assert result["has_more"] is True
|
||||
Loading…
Add table
Reference in a new issue