mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-30 01:52:18 +00:00
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:
parent
ab9b8b6a07
commit
7921b01b1a
5 changed files with 82 additions and 39 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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}",
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue