fix(rust): satisfy parity branch validation gates

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
Yujong Lee 2026-09-24 23:45:17 +00:00
parent ab9b8b6a07
commit 7921b01b1a
5 changed files with 82 additions and 39 deletions

View file

@ -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::<usize>().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,

View file

@ -192,9 +192,8 @@ fn batches_base_url(
fn batch_request(model: Option<&str>, line: &str) -> Result<Value, Error> {
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}",

View file

@ -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<String> {
std::env::var(name).ok()
}
async fn retrieve(batch_id: String, connection: OwnedConnection) -> Result<LiteLlmMessageBatch, Error> {
async fn retrieve(
batch_id: String,
connection: OwnedConnection,
) -> Result<LiteLlmMessageBatch, Error> {
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::<RustUpstreamError>(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);
}

View file

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

View file

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