diff --git a/litellm-rust/crates/core/src/batches/mod.rs b/litellm-rust/crates/core/src/batches/mod.rs index 17cee372b0f..58a23f82694 100644 --- a/litellm-rust/crates/core/src/batches/mod.rs +++ b/litellm-rust/crates/core/src/batches/mod.rs @@ -6,7 +6,9 @@ use litellm_llms::{ ANTHROPIC_BATCHES_TRANSFORMATION, AnthropicBatchesConfig, AnthropicMessageBatch, LiteLlmMessageBatch, }, - base_llm::{anthropic_messages::transformation::Headers, chat::transformation::Error as LlmError}, + base_llm::{ + anthropic_messages::transformation::Headers, chat::transformation::Error as LlmError, + }, }; use reqwest::Method; use time::OffsetDateTime; @@ -71,7 +73,15 @@ pub async fn create_batch( config.validate_environment(connection.extra_headers, connection.api_key, env_lookup)?; let body = serde_json::to_vec(&body) .map_err(|error| LlmError::InvalidRequest(format!("unserializable batch: {error}")))?; - let batch = send(client, Method::POST, url, headers, Some(body), connection.timeout).await?; + let batch = send( + client, + Method::POST, + url, + headers, + Some(body), + connection.timeout, + ) + .await?; Ok(config.transform_create_batch_response(batch, now())) } @@ -92,8 +102,12 @@ async fn send( .fold(client.request(method, url), |builder, (name, value)| { builder.header(name, value) }); - let request = body.into_iter().fold(request, reqwest::RequestBuilder::body); - let request = timeout.into_iter().fold(request, reqwest::RequestBuilder::timeout); + let request = body + .into_iter() + .fold(request, reqwest::RequestBuilder::body); + let request = timeout + .into_iter() + .fold(request, reqwest::RequestBuilder::timeout); let response = request .send() .await @@ -158,7 +172,11 @@ mod tests { headers.sort(); let length = head .lines() - .find_map(|line| line.to_lowercase().strip_prefix("content-length: ").map(str::to_string)) + .find_map(|line| { + line.to_lowercase() + .strip_prefix("content-length: ") + .map(str::to_string) + }) .map_or(0, |value| value.parse::().unwrap()); if body.len() >= length || n == 0 { break Received { @@ -329,9 +347,13 @@ mod tests { #[case::not_found( "404 Not Found", r#"{"type":"error","error":{"type":"not_found_error"}}"#, - "upstream request failed with status 404: {\"type\":\"error\",\"error\":{\"type\":\"not_found_error\"}}", + "upstream request failed with status 404: {\"type\":\"error\",\"error\":{\"type\":\"not_found_error\"}}" + )] + #[case::server_error( + "500 Internal Server Error", + "boom", + "upstream request failed with status 500: boom" )] - #[case::server_error("500 Internal Server Error", "boom", "upstream request failed with status 500: boom")] #[tokio::test] async fn retrieve_surfaces_upstream_failures( #[case] status: &'static str, diff --git a/litellm-rust/crates/llms/src/anthropic/batches/transformation.rs b/litellm-rust/crates/llms/src/anthropic/batches/transformation.rs index 73405b96b3e..ba70555661c 100644 --- a/litellm-rust/crates/llms/src/anthropic/batches/transformation.rs +++ b/litellm-rust/crates/llms/src/anthropic/batches/transformation.rs @@ -192,9 +192,8 @@ fn batches_base_url( fn batch_request(model: Option<&str>, line: &str) -> Result { let line: BatchInputLine = serde_json::from_str(line) .map_err(|error| Error::InvalidRequest(format!("invalid batch input line: {error}")))?; - let invalid = |reason: &str| { - Error::InvalidRequest(format!("batch request {}: {reason}", line.custom_id)) - }; + let invalid = + |reason: &str| Error::InvalidRequest(format!("batch request {}: {reason}", line.custom_id)); if line.method != "POST" || line.url != CHAT_COMPLETIONS_URL { return Err(invalid(&format!( "{} {} is not supported, only POST {CHAT_COMPLETIONS_URL}", diff --git a/litellm-rust/crates/python-bridge/src/routes/batches.rs b/litellm-rust/crates/python-bridge/src/routes/batches.rs index 42cdcd9634c..1f1cf9871ff 100644 --- a/litellm-rust/crates/python-bridge/src/routes/batches.rs +++ b/litellm-rust/crates/python-bridge/src/routes/batches.rs @@ -26,7 +26,12 @@ impl OwnedConnection { Connection { api_key: self.api_key.as_deref(), api_base: self.api_base.as_deref(), - extra_headers: self.extra_headers.clone().unwrap_or_default().into_iter().collect(), + extra_headers: self + .extra_headers + .clone() + .unwrap_or_default() + .into_iter() + .collect(), timeout: optional_timeout(self.timeout_seconds), } } @@ -36,7 +41,10 @@ fn env_lookup(name: &str) -> Option { std::env::var(name).ok() } -async fn retrieve(batch_id: String, connection: OwnedConnection) -> Result { +async fn retrieve( + batch_id: String, + connection: OwnedConnection, +) -> Result { run_retrieve_batch( http_client(), RetrieveBatchRequest { @@ -155,7 +163,11 @@ pub(crate) fn create_batch( extra_headers, timeout_seconds, }; - run_sync(py, create(input_jsonl, model, connection), create_error_to_pyerr) + run_sync( + py, + create(input_jsonl, model, connection), + create_error_to_pyerr, + ) } #[pyfunction] @@ -179,7 +191,11 @@ pub(crate) fn acreate_batch( extra_headers, timeout_seconds, }; - run_async(py, create(input_jsonl, model, connection), create_error_to_pyerr) + run_async( + py, + create(input_jsonl, model, connection), + create_error_to_pyerr, + ) } #[cfg(test)] @@ -206,7 +222,8 @@ mod tests { return Raised::Value(error.value(py).to_string()); } assert!(error.is_instance_of::(py), "{error}"); - let (status, message): (u16, String) = error.value(py).getattr("args").unwrap().extract().unwrap(); + let (status, message): (u16, String) = + error.value(py).getattr("args").unwrap().extract().unwrap(); Raised::Upstream(status, message) }) } @@ -236,7 +253,10 @@ mod tests { #[case::http(http_error(), Raised::Upstream(404, "missing".into()))] #[case::network(Error::Transport(TransportError::Network("reset".into())), Raised::Upstream(0, "reset".into()))] #[case::invalid_response(invalid_response(), Raised::Upstream(0, "invalid Anthropic batch response: []".into()))] - fn retrieve_declines_only_before_the_request_is_sent(#[case] error: Error, #[case] expected: Raised) { + fn retrieve_declines_only_before_the_request_is_sent( + #[case] error: Error, + #[case] expected: Raised, + ) { assert_eq!(raised(retrieve_error_to_pyerr(error)), expected); } diff --git a/litellm/proxy/batches_endpoints/endpoints.py b/litellm/proxy/batches_endpoints/endpoints.py index 4e79f9181f3..9f6a868d7f2 100644 --- a/litellm/proxy/batches_endpoints/endpoints.py +++ b/litellm/proxy/batches_endpoints/endpoints.py @@ -195,7 +195,9 @@ async def _managed_input_file_row(input_file_id: str) -> "LiteLLM_ManagedFileTab if prisma_client is None: return None try: - return await ManagedFileRepository(prisma_client).table.find_first(where={"unified_file_id": input_file_id}) + return await ManagedFileRepository(prisma_client).table.find_first( + where={"unified_file_id": input_file_id} # mutable-ok: Prisma query filters require a concrete mapping + ) except Exception as e: verbose_proxy_logger.warning("create_batch: managed file lookup failed for %s: %s", input_file_id, e) return None @@ -235,7 +237,7 @@ async def _create_provider_batch_for_managed_file( } response: Final = await llm_router.acreate_batch( **request, - **({} if input_content is None else {ANTHROPIC_BATCH_INPUT_CONTENT_KWARG: input_content}), + **({} if input_content is None else {ANTHROPIC_BATCH_INPUT_CONTENT_KWARG: input_content}), # mutable-ok: kwargs require a concrete mapping ) response.input_file_id = input_file_id response._hidden_params["unified_file_id"] = unified_file_id diff --git a/litellm/rust_bridge/batches/native.py b/litellm/rust_bridge/batches/native.py index 60478f19dde..2f2e2227ec5 100644 --- a/litellm/rust_bridge/batches/native.py +++ b/litellm/rust_bridge/batches/native.py @@ -1,6 +1,6 @@ from __future__ import annotations -from collections.abc import Awaitable +from collections.abc import Awaitable, Mapping from typing import Final, Protocol, cast # noqa: TID251 # validates dynamically loaded native callables from litellm.rust_bridge.bindings import NativeBinding @@ -12,9 +12,9 @@ class RustRetrieveBatch(Protocol): batch_id: str, api_key: str | None, api_base: str | None, - extra_headers: dict[str, str] | None, + extra_headers: Mapping[str, str] | None, timeout_seconds: float | None, - ) -> dict[str, object]: + ) -> Mapping[str, object]: raise NotImplementedError @@ -24,9 +24,9 @@ class RustAretrieveBatch(Protocol): batch_id: str, api_key: str | None, api_base: str | None, - extra_headers: dict[str, str] | None, + extra_headers: Mapping[str, str] | None, timeout_seconds: float | None, - ) -> Awaitable[dict[str, object]]: + ) -> Awaitable[Mapping[str, object]]: raise NotImplementedError @@ -37,9 +37,9 @@ class RustCreateBatch(Protocol): model: str | None, api_key: str | None, api_base: str | None, - extra_headers: dict[str, str] | None, + extra_headers: Mapping[str, str] | None, timeout_seconds: float | None, - ) -> dict[str, object]: + ) -> Mapping[str, object]: raise NotImplementedError @@ -50,34 +50,34 @@ class RustAcreateBatch(Protocol): model: str | None, api_key: str | None, api_base: str | None, - extra_headers: dict[str, str] | None, + extra_headers: Mapping[str, str] | None, timeout_seconds: float | None, - ) -> Awaitable[dict[str, object]]: + ) -> Awaitable[Mapping[str, object]]: raise NotImplementedError def _retrieve_binding(value: object) -> RustRetrieveBatch | None: - return ( - cast("RustRetrieveBatch", value) if callable(value) else None - ) # cast-ok: callable validated at the native binding boundary + if not callable(value): + return None + return cast("RustRetrieveBatch", value) # cast-ok: callable validated at the native binding boundary def _aretrieve_binding(value: object) -> RustAretrieveBatch | None: - return ( - cast("RustAretrieveBatch", value) if callable(value) else None - ) # cast-ok: callable validated at the native binding boundary + if not callable(value): + return None + return cast("RustAretrieveBatch", value) # cast-ok: callable validated at the native binding boundary def _create_binding(value: object) -> RustCreateBatch | None: - return ( - cast("RustCreateBatch", value) if callable(value) else None - ) # cast-ok: callable validated at the native binding boundary + if not callable(value): + return None + return cast("RustCreateBatch", value) # cast-ok: callable validated at the native binding boundary def _acreate_binding(value: object) -> RustAcreateBatch | None: - return ( - cast("RustAcreateBatch", value) if callable(value) else None - ) # cast-ok: callable validated at the native binding boundary + if not callable(value): + return None + return cast("RustAcreateBatch", value) # cast-ok: callable validated at the native binding boundary NATIVE_RETRIEVE_BATCH: Final = NativeBinding("retrieve_batch", validate=_retrieve_binding)