This commit is contained in:
Yujong Lee 2026-09-24 16:04:30 -07:00
parent ef42347675
commit ab9b8b6a07
15 changed files with 1288 additions and 51 deletions

View file

@ -0,0 +1,442 @@
use std::{sync::OnceLock, time::Duration};
use litellm_http::{request::truncate_error_body, transport::Error as TransportError};
use litellm_llms::{
anthropic::batches::transformation::{
ANTHROPIC_BATCHES_TRANSFORMATION, AnthropicBatchesConfig, AnthropicMessageBatch,
LiteLlmMessageBatch,
},
base_llm::{anthropic_messages::transformation::Headers, chat::transformation::Error as LlmError},
};
use reqwest::Method;
use time::OffsetDateTime;
#[derive(Debug, thiserror::Error)]
pub enum Error {
#[error(transparent)]
Request(#[from] LlmError),
#[error(transparent)]
Transport(#[from] TransportError),
#[error("invalid Anthropic batch response: {0}")]
InvalidResponse(String),
}
pub struct Connection<'a> {
pub api_key: Option<&'a str>,
pub api_base: Option<&'a str>,
pub extra_headers: Headers,
pub timeout: Option<Duration>,
}
pub struct RetrieveBatchRequest<'a> {
pub batch_id: &'a str,
pub connection: Connection<'a>,
}
pub struct CreateBatchRequest<'a> {
pub model: Option<&'a str>,
pub input_jsonl: &'a str,
pub connection: Connection<'a>,
}
pub fn http_client() -> &'static reqwest::Client {
static CLIENT: OnceLock<reqwest::Client> = OnceLock::new();
CLIENT.get_or_init(reqwest::Client::new)
}
pub async fn retrieve_batch(
client: &reqwest::Client,
request: RetrieveBatchRequest<'_>,
env_lookup: &(dyn Fn(&str) -> Option<String> + Sync),
) -> Result<LiteLlmMessageBatch, Error> {
let config = &ANTHROPIC_BATCHES_TRANSFORMATION;
let connection = request.connection;
let url = config.retrieve_batch_url(connection.api_base, request.batch_id, env_lookup)?;
let headers =
config.validate_environment(connection.extra_headers, connection.api_key, env_lookup)?;
let batch = send(client, Method::GET, url, headers, None, connection.timeout).await?;
Ok(config.transform_retrieve_batch_response(batch, now()))
}
pub async fn create_batch(
client: &reqwest::Client,
request: CreateBatchRequest<'_>,
env_lookup: &(dyn Fn(&str) -> Option<String> + Sync),
) -> Result<LiteLlmMessageBatch, Error> {
let config = &ANTHROPIC_BATCHES_TRANSFORMATION;
let connection = request.connection;
let body = config.transform_create_batch_request(request.model, request.input_jsonl)?;
let url = config.create_batch_url(connection.api_base, env_lookup)?;
let headers =
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?;
Ok(config.transform_create_batch_response(batch, now()))
}
fn now() -> i64 {
OffsetDateTime::now_utc().unix_timestamp()
}
async fn send(
client: &reqwest::Client,
method: Method,
url: String,
headers: Headers,
body: Option<Vec<u8>>,
timeout: Option<Duration>,
) -> Result<AnthropicMessageBatch, Error> {
let request = headers
.iter()
.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 response = request
.send()
.await
.map_err(TransportError::from_reqwest_before_dispatch)?;
let status = response.status();
let text = response.text().await.map_err(TransportError::from)?;
if !status.is_success() {
return Err(TransportError::Http {
status: status.as_u16(),
body: truncate_error_body(&text),
}
.into());
}
serde_json::from_str(&text).map_err(|error| Error::InvalidResponse(error.to_string()))
}
#[cfg(test)]
mod tests {
use litellm_auth::Error as AuthError;
use litellm_llms::anthropic::batches::transformation::{BatchRequestCounts, BatchStatus};
use rstest::rstest;
use serde_json::{Value, json};
use tokio::{
io::{AsyncReadExt, AsyncWriteExt},
net::TcpListener,
task::JoinHandle,
};
use super::*;
const CREATED: i64 = 1_727_172_000;
const BATCH: &str = r#"{"id":"msgbatch_1","type":"message_batch","processing_status":"in_progress","created_at":"2024-09-24T10:00:00Z","request_counts":{"processing":2,"succeeded":1}}"#;
#[derive(Debug, PartialEq)]
struct Received {
request_line: String,
headers: Vec<(String, String)>,
body: String,
}
async fn stub(status: &'static str, body: &'static str) -> (String, JoinHandle<Received>) {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let base = format!("http://{}", listener.local_addr().unwrap());
let handle = tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.unwrap();
let mut raw = Vec::new();
let mut buffer = [0_u8; 4096];
let received = loop {
let n = socket.read(&mut buffer).await.unwrap();
raw.extend_from_slice(&buffer[..n]);
let text = String::from_utf8_lossy(&raw).to_string();
let Some((head, body)) = text.split_once("\r\n\r\n") else {
continue;
};
let mut lines = head.lines();
let request_line = lines.next().unwrap().to_string();
let mut headers: Vec<(String, String)> = lines
.filter_map(|line| line.split_once(": "))
.map(|(name, value)| (name.to_lowercase(), value.to_string()))
.filter(|(name, _)| !matches!(name.as_str(), "host" | "content-length"))
.collect();
headers.sort();
let length = head
.lines()
.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 {
request_line,
headers,
body: body.to_string(),
};
}
};
let response = format!(
"HTTP/1.1 {status}\r\ncontent-type: application/json\r\ncontent-length: {}\r\nconnection: close\r\n\r\n{body}",
body.len()
);
socket.write_all(response.as_bytes()).await.unwrap();
received
});
(base, handle)
}
fn env(pairs: &[(&str, &str)]) -> impl Fn(&str) -> Option<String> + Sync + use<> {
let pairs: Vec<(String, String)> = pairs
.iter()
.map(|(key, value)| (key.to_string(), value.to_string()))
.collect();
move |name| {
pairs
.iter()
.find(|(key, _)| key == name)
.map(|(_, value)| value.clone())
}
}
fn connection<'a>(api_key: Option<&'a str>, api_base: Option<&'a str>) -> Connection<'a> {
Connection {
api_key,
api_base,
extra_headers: vec![],
timeout: Some(Duration::from_secs(5)),
}
}
fn headers(auth: (&str, &str)) -> Vec<(String, String)> {
let mut headers: Vec<(String, String)> = [
("accept", "application/json"),
("anthropic-beta", "message-batches-2024-09-24"),
("anthropic-version", "2023-06-01"),
("content-type", "application/json"),
auth,
]
.map(|(name, value)| (name.to_string(), value.to_string()))
.into();
headers.sort();
headers
}
fn in_progress_batch() -> LiteLlmMessageBatch {
LiteLlmMessageBatch {
id: "msgbatch_1".into(),
object: "batch".into(),
endpoint: "/v1/messages".into(),
input_file_id: "None".into(),
completion_window: "24h".into(),
status: BatchStatus::InProgress,
output_file_id: "msgbatch_1".into(),
created_at: CREATED,
in_progress_at: Some(CREATED),
expires_at: None,
completed_at: None,
expired_at: None,
cancelling_at: None,
cancelled_at: None,
request_counts: BatchRequestCounts {
total: 3,
completed: 1,
failed: 0,
},
}
}
#[rstest]
#[case::explicit_key(Some("sk-ant-param"), &[], ("x-api-key", "sk-ant-param"))]
#[case::key_from_environment(None, &[("ANTHROPIC_API_KEY", "sk-ant-env")], ("x-api-key", "sk-ant-env"))]
#[tokio::test]
async fn retrieve_gets_the_batch_and_maps_it(
#[case] api_key: Option<&'static str>,
#[case] environment: &'static [(&'static str, &'static str)],
#[case] auth: (&'static str, &'static str),
) {
let (base, server) = stub("200 OK", BATCH).await;
let batch = retrieve_batch(
http_client(),
RetrieveBatchRequest {
batch_id: "msgbatch_1",
connection: connection(api_key, Some(&base)),
},
&env(environment),
)
.await
.unwrap();
assert_eq!(batch, in_progress_batch());
assert_eq!(
server.await.unwrap(),
Received {
request_line: "GET /v1/messages/batches/msgbatch_1 HTTP/1.1".into(),
headers: headers(auth),
body: String::new(),
}
);
}
#[tokio::test]
async fn retrieve_falls_back_to_the_base_from_the_environment() {
let (base, server) = stub("200 OK", BATCH).await;
retrieve_batch(
http_client(),
RetrieveBatchRequest {
batch_id: "msgbatch_1",
connection: connection(Some("sk"), None),
},
&env(&[("ANTHROPIC_API_BASE", &base)]),
)
.await
.unwrap();
assert_eq!(
server.await.unwrap().request_line,
"GET /v1/messages/batches/msgbatch_1 HTTP/1.1"
);
}
#[rstest]
#[case::missing_key(
"msgbatch_1",
None,
Error::Request(LlmError::Auth(AuthError::MissingApiKey {
provider: "Anthropic",
environment_variable: "ANTHROPIC_API_KEY",
})),
)]
#[case::dot_segment(
"..",
Some("sk"),
Error::Request(LlmError::InvalidRequest("batch_id cannot be a dot path segment".into())),
)]
#[tokio::test]
async fn retrieve_rejects_requests_before_calling_out(
#[case] batch_id: &str,
#[case] api_key: Option<&str>,
#[case] expected: Error,
) {
let error = retrieve_batch(
http_client(),
RetrieveBatchRequest {
batch_id,
connection: connection(api_key, Some("http://127.0.0.1:9")),
},
&env(&[]),
)
.await
.unwrap_err();
assert_eq!(error.to_string(), expected.to_string());
assert!(matches!(error, Error::Request(_)));
}
#[rstest]
#[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\"}}",
)]
#[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,
#[case] body: &'static str,
#[case] message: &str,
) {
let (base, _server) = stub(status, body).await;
let error = retrieve_batch(
http_client(),
RetrieveBatchRequest {
batch_id: "msgbatch_1",
connection: connection(Some("sk"), Some(&base)),
},
&env(&[]),
)
.await
.unwrap_err();
assert_eq!(error.to_string(), message);
assert!(!matches!(error, Error::Request(_)));
}
#[tokio::test]
async fn retrieve_rejects_a_success_body_that_is_not_a_batch() {
let (base, _server) = stub("200 OK", "[]").await;
let error = retrieve_batch(
http_client(),
RetrieveBatchRequest {
batch_id: "msgbatch_1",
connection: connection(Some("sk"), Some(&base)),
},
&env(&[]),
)
.await
.unwrap_err();
assert!(matches!(error, Error::InvalidResponse(_)), "{error:?}");
}
#[tokio::test]
async fn create_posts_the_translated_requests_and_maps_the_batch() {
let (base, server) = stub("200 OK", BATCH).await;
let input = json!({
"custom_id": "r1",
"method": "POST",
"url": "/v1/chat/completions",
"body": {"model": "alias", "messages": [{"role": "user", "content": "hi"}], "max_tokens": 4},
})
.to_string();
let batch = create_batch(
http_client(),
CreateBatchRequest {
model: Some("claude-deployed"),
input_jsonl: &input,
connection: Connection {
extra_headers: vec![("x-trace".into(), "t1".into())],
..connection(Some("sk-ant"), Some(&base))
},
},
&env(&[]),
)
.await
.unwrap();
assert_eq!(batch, in_progress_batch());
let received = server.await.unwrap();
let mut expected_headers = headers(("x-api-key", "sk-ant"));
expected_headers.push(("x-trace".into(), "t1".into()));
expected_headers.sort();
assert_eq!(
(
received.request_line,
received.headers,
serde_json::from_str::<Value>(&received.body).unwrap()
),
(
"POST /v1/messages/batches HTTP/1.1".to_string(),
expected_headers,
json!({"requests": [{"custom_id": "r1", "params": {
"model": "claude-deployed",
"messages": [{"role": "user", "content": [{"type": "text", "text": "hi"}]}],
"max_tokens": 4,
}}]}),
)
);
}
#[tokio::test]
async fn create_rejects_an_untranslatable_input_before_calling_out() {
let error = create_batch(
http_client(),
CreateBatchRequest {
model: None,
input_jsonl: "",
connection: connection(Some("sk"), Some("http://127.0.0.1:9")),
},
&env(&[]),
)
.await
.unwrap_err();
assert_eq!(
error.to_string(),
"invalid request: batch input file has no requests"
);
assert!(matches!(error, Error::Request(_)));
}
}

View file

@ -1,4 +1,5 @@
pub mod audio_transcription;
pub mod batches;
pub mod chat_completions;
pub mod constants;
pub mod error;

View file

@ -1,15 +1,20 @@
use litellm_types::llms::anthropic_messages::anthropic_response::AnthropicMessagesResponse;
use litellm_types::llms::openai::ChatMessage;
use serde::{Deserialize, Serialize};
use serde_json::Value;
use serde_json::{Map, Value, json};
use time::OffsetDateTime;
use url::Url;
use crate::{
anthropic::{
chat::transformation::ANTHROPIC_CHAT_COMPLETIONS_CONFIG,
common_utils::is_anthropic_oauth_key,
experimental_pass_through::messages::transformation::resolve_anthropic_api_base,
},
base_llm::{anthropic_messages::transformation::Headers, chat::transformation::Error},
base_llm::{
anthropic_messages::transformation::Headers,
chat::transformation::{BaseConfig, Error},
},
};
const BATCHES_PATH_SUFFIX: &str = "/v1/messages/batches";
@ -17,6 +22,16 @@ const BATCHES_BETA: &str = "message-batches-2024-09-24";
const BETA_HEADER: &str = "anthropic-beta";
const ANTHROPIC_API_KEY_ENV: &str = "ANTHROPIC_API_KEY";
const ANTHROPIC_AUTH_TOKEN_ENV: &str = "ANTHROPIC_AUTH_TOKEN";
const CHAT_COMPLETIONS_URL: &str = "/v1/chat/completions";
const MODEL_PREFIX: &str = "anthropic/";
#[derive(Deserialize)]
struct BatchInputLine {
custom_id: String,
method: String,
url: String,
body: Map<String, Value>,
}
#[derive(Clone, Debug, Default, PartialEq, Eq, Serialize, Deserialize)]
pub struct AnthropicBatchRequestCounts {
@ -95,13 +110,17 @@ pub trait AnthropicBatchesConfig {
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> Result<String, Error>;
fn transform_create_batch_request(&self) -> Result<Value, Error>;
fn transform_create_batch_request(
&self,
model: Option<&str>,
input_jsonl: &str,
) -> Result<Value, Error>;
fn transform_create_batch_response(
&self,
response: AnthropicMessageBatch,
now: i64,
) -> Result<LiteLlmMessageBatch, Error>;
) -> LiteLlmMessageBatch;
fn retrieve_batch_url(
&self,
@ -170,6 +189,55 @@ fn batches_base_url(
.map_err(|error| Error::InvalidRequest(format!("invalid Anthropic API base: {error}")))
}
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))
};
if line.method != "POST" || line.url != CHAT_COMPLETIONS_URL {
return Err(invalid(&format!(
"{} {} is not supported, only POST {CHAT_COMPLETIONS_URL}",
line.method, line.url
)));
}
let model = model
.or_else(|| line.body.get("model").and_then(Value::as_str))
.ok_or_else(|| invalid("model is required"))?;
let model = model.strip_prefix(MODEL_PREFIX).unwrap_or(model);
let messages: Vec<ChatMessage> = line
.body
.get("messages")
.cloned()
.map(serde_json::from_value)
.transpose()
.map_err(|error| invalid(&format!("invalid messages: {error}")))?
.filter(|messages: &Vec<ChatMessage>| !messages.is_empty())
.ok_or_else(|| invalid("messages is required"))?;
let config = &ANTHROPIC_CHAT_COMPLETIONS_CONFIG;
let params = line
.body
.iter()
.filter(|(name, _)| !matches!(name.as_str(), "model" | "messages"))
.map(|(name, value)| {
config
.supported_openai_param_mappings()
.iter()
.find(|(openai, _)| openai == name)
.map(|(_, anthropic)| ((*anthropic).to_string(), value.clone()))
.ok_or_else(|| invalid(&format!("parameter {name} is not supported")))
})
.collect::<Result<Map<_, _>, _>>()?;
if !params.contains_key("max_tokens") {
return Err(invalid("max_tokens is required"));
}
if let Some(reason) = config.unsupported_reason(&messages, &params) {
return Err(invalid(reason.0));
}
let params = config.transform_request(model, messages, params)?.body;
Ok(json!({ "custom_id": line.custom_id, "params": params }))
}
impl AnthropicBatchesConfig for AnthropicBatchesTransformation {
fn validate_environment(
&self,
@ -217,16 +285,31 @@ impl AnthropicBatchesConfig for AnthropicBatchesTransformation {
Ok(batches_base_url(api_base, env_lookup)?.into())
}
fn transform_create_batch_request(&self) -> Result<Value, Error> {
Err(Error::Unsupported("Anthropic message batch creation"))
fn transform_create_batch_request(
&self,
model: Option<&str>,
input_jsonl: &str,
) -> Result<Value, Error> {
let requests = input_jsonl
.lines()
.map(str::trim)
.filter(|line| !line.is_empty())
.map(|line| batch_request(model, line))
.collect::<Result<Vec<_>, _>>()?;
if requests.is_empty() {
return Err(Error::InvalidRequest(
"batch input file has no requests".into(),
));
}
Ok(json!({ "requests": requests }))
}
fn transform_create_batch_response(
&self,
_response: AnthropicMessageBatch,
_now: i64,
) -> Result<LiteLlmMessageBatch, Error> {
Err(Error::Unsupported("Anthropic message batch creation"))
response: AnthropicMessageBatch,
now: i64,
) -> LiteLlmMessageBatch {
self.transform_retrieve_batch_response(response, now)
}
fn retrieve_batch_url(
@ -580,17 +663,156 @@ mod tests {
);
}
fn line(custom_id: &str, body: Value) -> String {
json!({"custom_id": custom_id, "method": "POST", "url": "/v1/chat/completions", "body": body})
.to_string()
}
fn user_turn(text: &str) -> Value {
json!({"role": "user", "content": [{"type": "text", "text": text}]})
}
#[rstest]
#[case::maps_openai_params_and_folds_system(
None,
line("r1", json!({
"model": "anthropic/claude-sonnet",
"messages": [
{"role": "system", "content": "be brief"},
{"role": "user", "content": "hi"},
],
"max_tokens": 16,
"stop": ["END"],
"temperature": 0.5,
"top_p": 0.9,
})),
json!({"requests": [{"custom_id": "r1", "params": {
"model": "claude-sonnet",
"messages": [user_turn("hi")],
"system": [{"type": "text", "text": "be brief"}],
"max_tokens": 16,
"stop_sequences": ["END"],
"temperature": 0.5,
"top_p": 0.9,
}}]}),
)]
#[case::deployment_model_overrides_the_line_model(
Some("claude-deployed"),
line("r1", json!({"model": "alias", "messages": [{"role": "user", "content": "hi"}], "max_tokens": 8})),
json!({"requests": [{"custom_id": "r1", "params": {
"model": "claude-deployed", "messages": [user_turn("hi")], "max_tokens": 8,
}}]}),
)]
#[case::keeps_line_order_and_skips_blank_lines(
None,
format!(
"\n{}\r\n \n{}\n",
line("b", json!({"model": "m", "messages": [{"role": "user", "content": "two"}], "max_tokens": 2})),
line("a", json!({"model": "m", "messages": [{"role": "user", "content": "one"}], "max_tokens": 1})),
),
json!({"requests": [
{"custom_id": "b", "params": {"model": "m", "messages": [user_turn("two")], "max_tokens": 2}},
{"custom_id": "a", "params": {"model": "m", "messages": [user_turn("one")], "max_tokens": 1}},
]}),
)]
fn create_batch_request_translates_each_chat_line(
#[case] model: Option<&str>,
#[case] input: String,
#[case] expected: Value,
) {
assert_eq!(
ANTHROPIC_BATCHES_TRANSFORMATION
.transform_create_batch_request(model, &input)
.unwrap(),
expected
);
}
#[rstest]
#[case::empty_input(String::new(), "invalid request: batch input file has no requests")]
#[case::blank_input("\n \n".to_string(), "invalid request: batch input file has no requests")]
#[case::not_json(
"not-json".to_string(),
"invalid request: invalid batch input line: expected ident at line 1 column 2",
)]
#[case::other_endpoint(
json!({"custom_id": "r1", "method": "POST", "url": "/v1/embeddings", "body": {}}).to_string(),
"invalid request: batch request r1: POST /v1/embeddings is not supported, only POST /v1/chat/completions",
)]
#[case::other_method(
json!({"custom_id": "r1", "method": "GET", "url": "/v1/chat/completions", "body": {}}).to_string(),
"invalid request: batch request r1: GET /v1/chat/completions is not supported, only POST /v1/chat/completions",
)]
#[case::no_model(
line("r1", json!({"messages": [{"role": "user", "content": "hi"}], "max_tokens": 1})),
"invalid request: batch request r1: model is required",
)]
#[case::no_messages(
line("r1", json!({"model": "m", "max_tokens": 1})),
"invalid request: batch request r1: messages is required",
)]
#[case::empty_messages(
line("r1", json!({"model": "m", "messages": [], "max_tokens": 1})),
"invalid request: batch request r1: messages is required",
)]
#[case::malformed_messages(
line("r1", json!({"model": "m", "messages": "hi", "max_tokens": 1})),
"invalid request: batch request r1: invalid messages: invalid type: string \"hi\", expected a sequence",
)]
#[case::unsupported_param(
line("r1", json!({"model": "m", "messages": [{"role": "user", "content": "hi"}], "max_tokens": 1, "tools": []})),
"invalid request: batch request r1: parameter tools is not supported",
)]
#[case::no_max_tokens(
line("r1", json!({"model": "m", "messages": [{"role": "user", "content": "hi"}]})),
"invalid request: batch request r1: max_tokens is required",
)]
#[case::streaming(
line("r1", json!({"model": "m", "messages": [{"role": "user", "content": "hi"}], "max_tokens": 1, "stream": true})),
"invalid request: batch request r1: parameter stream is not supported",
)]
#[case::tool_turn(
line("r1", json!({"model": "m", "messages": [{"role": "tool", "content": "x", "tool_call_id": "t"}], "max_tokens": 1})),
"invalid request: batch request r1: unrecognized message field",
)]
#[case::opens_on_assistant_turn(
line("r1", json!({"model": "m", "messages": [{"role": "assistant", "content": "hi"}], "max_tokens": 1})),
"invalid request: batch request r1: conversation does not open on a user turn",
)]
#[case::one_bad_line_fails_the_batch(
format!(
"{}\n{}",
line("ok", json!({"model": "m", "messages": [{"role": "user", "content": "hi"}], "max_tokens": 1})),
line("bad", json!({"model": "m", "messages": [{"role": "user", "content": "hi"}]})),
),
"invalid request: batch request bad: max_tokens is required",
)]
fn create_batch_request_rejects_lines_it_cannot_translate(
#[case] input: String,
#[case] message: &str,
) {
assert_eq!(
ANTHROPIC_BATCHES_TRANSFORMATION
.transform_create_batch_request(None, &input)
.unwrap_err()
.to_string(),
message
);
}
#[test]
fn batch_creation_is_unsupported() {
assert!(matches!(
ANTHROPIC_BATCHES_TRANSFORMATION.transform_create_batch_request(),
Err(Error::Unsupported("Anthropic message batch creation"))
));
let response: AnthropicMessageBatch = serde_json::from_value(json!({})).unwrap();
assert!(matches!(
ANTHROPIC_BATCHES_TRANSFORMATION.transform_create_batch_response(response, 0),
Err(Error::Unsupported("Anthropic message batch creation"))
));
fn create_batch_response_is_the_retrieved_batch() {
let body = json!({
"id": "msgbatch_new",
"processing_status": "in_progress",
"created_at": "2024-09-24T10:00:00Z",
"request_counts": {"processing": 2},
});
let response: AnthropicMessageBatch = serde_json::from_value(body.clone()).unwrap();
assert_eq!(
ANTHROPIC_BATCHES_TRANSFORMATION.transform_create_batch_response(response, NOW),
retrieve(body)
);
}
#[rstest]

View file

@ -26,6 +26,8 @@ mod _native {
#[pymodule_export]
use crate::routes::audio_transcription::{atranscription, transcription};
#[pymodule_export]
use crate::routes::batches::{acreate_batch, aretrieve_batch, create_batch, retrieve_batch};
#[pymodule_export]
use crate::routes::chat_completions::{
achat_completions, acompletion, chat_completions, chat_completions_decline, completion,
};
@ -85,6 +87,10 @@ mod tests {
"ProcessReservedForForking",
"ocr",
"aocr",
"retrieve_batch",
"aretrieve_batch",
"create_batch",
"acreate_batch",
"embedding",
"aembedding",
"transcription",

View file

@ -0,0 +1,251 @@
use std::collections::BTreeMap;
use litellm_core::batches::{
Connection, CreateBatchRequest, Error, RetrieveBatchRequest, create_batch as run_create_batch,
http_client, retrieve_batch as run_retrieve_batch,
};
use litellm_http::transport::Error as TransportError;
use litellm_llms::anthropic::batches::transformation::LiteLlmMessageBatch;
use pyo3::{exceptions::PyValueError, prelude::*};
use crate::{
errors::{RustBridgeDeclined, RustUpstreamError},
logger::{run_async, run_sync},
marshal::optional_timeout,
};
struct OwnedConnection {
api_key: Option<String>,
api_base: Option<String>,
extra_headers: Option<BTreeMap<String, String>>,
timeout_seconds: Option<f64>,
}
impl OwnedConnection {
fn borrow(&self) -> Connection<'_> {
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(),
timeout: optional_timeout(self.timeout_seconds),
}
}
}
fn env_lookup(name: &str) -> Option<String> {
std::env::var(name).ok()
}
async fn retrieve(batch_id: String, connection: OwnedConnection) -> Result<LiteLlmMessageBatch, Error> {
run_retrieve_batch(
http_client(),
RetrieveBatchRequest {
batch_id: &batch_id,
connection: connection.borrow(),
},
&env_lookup,
)
.await
}
async fn create(
input_jsonl: String,
model: Option<String>,
connection: OwnedConnection,
) -> Result<LiteLlmMessageBatch, Error> {
run_create_batch(
http_client(),
CreateBatchRequest {
model: model.as_deref(),
input_jsonl: &input_jsonl,
connection: connection.borrow(),
},
&env_lookup,
)
.await
}
fn upstream_error(error: TransportError) -> PyErr {
match error {
TransportError::Http { status, body } => RustUpstreamError::new_err((status, body)),
TransportError::Network(message) | TransportError::Connect(message) => {
RustUpstreamError::new_err((0u16, message))
}
}
}
/// Python keeps a retrieve implementation, so anything that fails before the
/// request goes out declines and Python raises its own error for it.
fn retrieve_error_to_pyerr(error: Error) -> PyErr {
match error {
Error::Request(_) | Error::Transport(TransportError::Connect(_)) => {
RustBridgeDeclined::new_err(error.to_string())
}
Error::Transport(error) => upstream_error(error),
Error::InvalidResponse(_) => RustUpstreamError::new_err((0u16, error.to_string())),
}
}
/// Python has no Anthropic create to fall back to, so a rejected request is
/// the caller's error.
fn create_error_to_pyerr(error: Error) -> PyErr {
match error {
Error::Request(_) => PyValueError::new_err(error.to_string()),
Error::Transport(error) => upstream_error(error),
Error::InvalidResponse(_) => RustUpstreamError::new_err((0u16, error.to_string())),
}
}
#[pyfunction]
#[pyo3(signature = (batch_id, api_key=None, api_base=None, extra_headers=None, timeout_seconds=None))]
pub(crate) fn retrieve_batch(
py: Python<'_>,
batch_id: String,
api_key: Option<String>,
api_base: Option<String>,
extra_headers: Option<BTreeMap<String, String>>,
timeout_seconds: Option<f64>,
) -> PyResult<Py<PyAny>> {
let connection = OwnedConnection {
api_key,
api_base,
extra_headers,
timeout_seconds,
};
run_sync(py, retrieve(batch_id, connection), retrieve_error_to_pyerr)
}
#[pyfunction]
#[pyo3(signature = (batch_id, api_key=None, api_base=None, extra_headers=None, timeout_seconds=None))]
pub(crate) fn aretrieve_batch(
py: Python<'_>,
batch_id: String,
api_key: Option<String>,
api_base: Option<String>,
extra_headers: Option<BTreeMap<String, String>>,
timeout_seconds: Option<f64>,
) -> PyResult<Bound<'_, PyAny>> {
let connection = OwnedConnection {
api_key,
api_base,
extra_headers,
timeout_seconds,
};
run_async(py, retrieve(batch_id, connection), retrieve_error_to_pyerr)
}
#[pyfunction]
#[pyo3(signature = (input_jsonl, model=None, api_key=None, api_base=None, extra_headers=None, timeout_seconds=None))]
#[expect(
clippy::too_many_arguments,
reason = "one parameter per Python keyword"
)]
pub(crate) fn create_batch(
py: Python<'_>,
input_jsonl: String,
model: Option<String>,
api_key: Option<String>,
api_base: Option<String>,
extra_headers: Option<BTreeMap<String, String>>,
timeout_seconds: Option<f64>,
) -> PyResult<Py<PyAny>> {
let connection = OwnedConnection {
api_key,
api_base,
extra_headers,
timeout_seconds,
};
run_sync(py, create(input_jsonl, model, connection), create_error_to_pyerr)
}
#[pyfunction]
#[pyo3(signature = (input_jsonl, model=None, api_key=None, api_base=None, extra_headers=None, timeout_seconds=None))]
#[expect(
clippy::too_many_arguments,
reason = "one parameter per Python keyword"
)]
pub(crate) fn acreate_batch(
py: Python<'_>,
input_jsonl: String,
model: Option<String>,
api_key: Option<String>,
api_base: Option<String>,
extra_headers: Option<BTreeMap<String, String>>,
timeout_seconds: Option<f64>,
) -> PyResult<Bound<'_, PyAny>> {
let connection = OwnedConnection {
api_key,
api_base,
extra_headers,
timeout_seconds,
};
run_async(py, create(input_jsonl, model, connection), create_error_to_pyerr)
}
#[cfg(test)]
mod tests {
use litellm_llms::base_llm::chat::transformation::Error as LlmError;
use rstest::rstest;
use super::*;
#[derive(Debug, PartialEq)]
enum Raised {
Declined(String),
Upstream(u16, String),
Value(String),
}
fn raised(error: PyErr) -> Raised {
Python::initialize();
Python::attach(|py| {
if error.is_instance_of::<RustBridgeDeclined>(py) {
return Raised::Declined(error.value(py).to_string());
}
if error.is_instance_of::<PyValueError>(py) {
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();
Raised::Upstream(status, message)
})
}
fn request_error() -> Error {
Error::Request(LlmError::InvalidRequest("bad".into()))
}
fn http_error() -> Error {
Error::Transport(TransportError::Http {
status: 404,
body: "missing".into(),
})
}
fn connect_error() -> Error {
Error::Transport(TransportError::Connect("refused".into()))
}
fn invalid_response() -> Error {
Error::InvalidResponse("[]".into())
}
#[rstest]
#[case::request(request_error(), Raised::Declined("invalid request: bad".into()))]
#[case::connect(connect_error(), Raised::Declined("could not reach the provider: refused".into()))]
#[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) {
assert_eq!(raised(retrieve_error_to_pyerr(error)), expected);
}
#[rstest]
#[case::request(request_error(), Raised::Value("invalid request: bad".into()))]
#[case::connect(connect_error(), Raised::Upstream(0, "refused".into()))]
#[case::http(http_error(), Raised::Upstream(404, "missing".into()))]
#[case::invalid_response(invalid_response(), Raised::Upstream(0, "invalid Anthropic batch response: []".into()))]
fn create_never_declines(#[case] error: Error, #[case] expected: Raised) {
assert_eq!(raised(create_error_to_pyerr(error)), expected);
}
}

View file

@ -1,4 +1,5 @@
pub(crate) mod audio_transcription;
pub(crate) mod batches;
pub(crate) mod chat_completions;
pub(crate) mod embeddings;
pub(crate) mod messages;

111
litellm/batches/dispatch.py Normal file
View file

@ -0,0 +1,111 @@
from collections.abc import Callable, Coroutine
from dataclasses import dataclass
from typing import Final
import httpx
from litellm.rust_bridge import runtime
from litellm.rust_bridge.batches.native import (
NATIVE_ACREATE_BATCH,
NATIVE_ARETRIEVE_BATCH,
NATIVE_CREATE_BATCH,
NATIVE_RETRIEVE_BATCH,
RustAcreateBatch,
RustAretrieveBatch,
RustCreateBatch,
RustRetrieveBatch,
)
from litellm.rust_bridge.catalog import Route, RouteContext
from litellm.rust_bridge.timeouts import timeout_to_seconds
from litellm.types.utils import LiteLLMBatch
ANTHROPIC_CONTEXT: Final = RouteContext(Route.BATCHES, provider="anthropic")
ANTHROPIC_BATCH_INPUT_CONTENT_KWARG: Final = "litellm_batch_input_content"
BatchResult = LiteLLMBatch | Coroutine[object, object, LiteLLMBatch]
@dataclass(frozen=True, slots=True)
class AnthropicConnection:
api_key: str | None
api_base: str | None
extra_headers: dict[str, str] | None
timeout: float | httpx.Timeout | None
def retrieve_anthropic_batch(
*,
is_async: bool,
batch_id: str,
connection: AnthropicConnection,
python: Callable[[], BatchResult],
) -> BatchResult:
timeout_seconds: Final = timeout_to_seconds(connection.timeout)
def native(rust: RustRetrieveBatch) -> LiteLLMBatch:
return LiteLLMBatch.model_validate(
rust(batch_id, connection.api_key, connection.api_base, connection.extra_headers, timeout_seconds)
)
async def anative(rust: RustAretrieveBatch) -> LiteLLMBatch:
return LiteLLMBatch.model_validate(
await rust(batch_id, connection.api_key, connection.api_base, connection.extra_headers, timeout_seconds)
)
async def apython() -> LiteLLMBatch:
result: Final = python()
return await result if isinstance(result, Coroutine) else result
if is_async:
return runtime.arun(ANTHROPIC_CONTEXT, binding=NATIVE_ARETRIEVE_BATCH, native=anative, python=apython)
return runtime.run(ANTHROPIC_CONTEXT, binding=NATIVE_RETRIEVE_BATCH, native=native, python=python)
def create_anthropic_batch(
*,
is_async: bool,
input_jsonl: str | None,
model: str | None,
connection: AnthropicConnection,
python: Callable[[], LiteLLMBatch],
) -> BatchResult:
timeout_seconds: Final = timeout_to_seconds(connection.timeout)
def content() -> str:
if input_jsonl is None:
raise ValueError(
"Anthropic batches read their requests from a managed file: upload the input file through the "
"proxy with purpose=batch, target_model_names, and a target_storage backend"
)
return input_jsonl
def native(rust: RustCreateBatch) -> LiteLLMBatch:
return LiteLLMBatch.model_validate(
rust(
content(),
model,
connection.api_key,
connection.api_base,
connection.extra_headers,
timeout_seconds,
)
)
async def anative(rust: RustAcreateBatch) -> LiteLLMBatch:
return LiteLLMBatch.model_validate(
await rust(
content(),
model,
connection.api_key,
connection.api_base,
connection.extra_headers,
timeout_seconds,
)
)
async def apython() -> LiteLLMBatch:
return python()
if is_async:
return runtime.arun(ANTHROPIC_CONTEXT, binding=NATIVE_ACREATE_BATCH, native=anative, python=apython)
return runtime.run(ANTHROPIC_CONTEXT, binding=NATIVE_CREATE_BATCH, native=native, python=python)

View file

@ -15,13 +15,18 @@ import contextvars
import os
from collections.abc import Coroutine
from functools import partial
from typing import Any, Final, Literal, cast
from typing import Any, Final, Literal, NoReturn, cast
import httpx
from openai.types.batch import BatchRequestCounts
import litellm
from litellm._logging import verbose_logger
from litellm.batches.dispatch import (
AnthropicConnection,
create_anthropic_batch,
retrieve_anthropic_batch,
)
from litellm.litellm_core_utils.get_litellm_params import add_trusted_model_credentials_to_litellm_params
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.llms.anthropic.batches.handler import AnthropicBatchesHandler
@ -108,7 +113,7 @@ async def acreate_batch(
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", "mistral"
"openai", "azure", "vertex_ai", "bedrock", "hosted_vllm", "litellm_proxy", "mistral", "anthropic"
] = "openai",
metadata: dict[str, str] | None = None,
extra_headers: dict[str, str] | None = None,
@ -154,18 +159,32 @@ async def acreate_batch(
raise e
def _raise_create_batch_unsupported(custom_llm_provider: str) -> NoReturn:
raise litellm.exceptions.BadRequestError(
message=f"LiteLLM doesn't support custom_llm_provider={custom_llm_provider} for 'create_batch'",
model="n/a",
llm_provider=custom_llm_provider,
response=httpx.Response(
status_code=400,
content="Unsupported provider",
request=httpx.Request(method="create_batch", url="https://github.com/BerriAI/litellm"),
),
)
@client
def create_batch(
completion_window: Literal["24h"],
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", "mistral"
"openai", "azure", "vertex_ai", "bedrock", "hosted_vllm", "litellm_proxy", "mistral", "anthropic"
] = "openai",
metadata: dict[str, str] | None = None,
extra_headers: dict[str, str] | None = None,
extra_body: dict[str, str] | None = None,
output_expires_after: dict[str, Any] | None = None,
litellm_batch_input_content: str | None = None,
**kwargs,
) -> LiteLLMBatch | Coroutine[Any, Any, LiteLLMBatch]:
"""
@ -325,17 +344,21 @@ def create_batch(
create_batch_data=_create_batch_request,
custom_endpoint=optional_params.get("custom_endpoint"),
)
else:
raise litellm.exceptions.BadRequestError(
message=f"LiteLLM doesn't support custom_llm_provider={custom_llm_provider} for 'create_batch'",
model="n/a",
llm_provider=custom_llm_provider,
response=httpx.Response(
status_code=400,
content="Unsupported provider",
request=httpx.Request(method="create_batch", url="https://github.com/BerriAI/litellm"),
elif custom_llm_provider == "anthropic":
response = create_anthropic_batch(
is_async=_is_async,
input_jsonl=litellm_batch_input_content,
model=model if isinstance(model, str) else None,
connection=AnthropicConnection(
api_key=optional_params.api_key or litellm.api_key,
api_base=optional_params.api_base or litellm.api_base,
extra_headers=extra_headers,
timeout=timeout,
),
python=lambda: _raise_create_batch_unsupported(custom_llm_provider),
)
else:
_raise_create_batch_unsupported(custom_llm_provider)
return response
except Exception as e:
raise e
@ -488,13 +511,23 @@ def _handle_retrieve_batch_providers_without_provider_config(
)
api_key = optional_params.api_key or litellm.api_key or litellm.azure_key or get_secret_str("ANTHROPIC_API_KEY")
response = anthropic_batches_instance.retrieve_batch(
_is_async=_is_async,
response = retrieve_anthropic_batch(
is_async=_is_async,
batch_id=batch_id,
api_base=api_base,
api_key=api_key,
timeout=timeout,
max_retries=optional_params.max_retries,
connection=AnthropicConnection(
api_key=optional_params.api_key or litellm.api_key,
api_base=optional_params.api_base or litellm.api_base,
extra_headers=_retrieve_batch_request.get("extra_headers"),
timeout=timeout,
),
python=lambda: anthropic_batches_instance.retrieve_batch(
_is_async=_is_async,
batch_id=batch_id,
api_base=api_base,
api_key=api_key,
timeout=timeout,
max_retries=optional_params.max_retries,
),
)
else:
raise litellm.exceptions.BadRequestError(

View file

@ -16,7 +16,9 @@ from pydantic import TypeAdapter
import litellm
from litellm._logging import verbose_proxy_logger
from litellm.batches.dispatch import ANTHROPIC_BATCH_INPUT_CONTENT_KWARG
from litellm.batches.main import CancelBatchRequest, RetrieveBatchRequest
from litellm.llms.base_llm.files.storage_backend_factory import get_storage_backend
from litellm.proxy._types import *
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.batches_endpoints.common_utils import validate_batch_list_limit
@ -25,9 +27,9 @@ from litellm.proxy.batches_endpoints.litellm_executed_batches import (
LiteLLMExecutedBatchRunner,
ManagedBatchStore,
batch_error,
deployment_provider_of,
executed_batch_runner_lost,
litellm_executed_provider_for,
resolve_litellm_executed_provider,
)
from litellm.proxy.common_request_processing import (
ProxyBaseLLMRequestProcessing,
@ -73,7 +75,7 @@ from litellm.types.llms.openai import LiteLLMBatchCreateRequest
from litellm.types.utils import LiteLLMBatch
if TYPE_CHECKING:
from prisma.models import LiteLLM_ManagedObjectTable
from prisma.models import LiteLLM_ManagedFileTable, LiteLLM_ManagedObjectTable
router: Final = APIRouter()
_METADATA_ADAPTER: Final[TypeAdapter[Mapping[str, object]]] = TypeAdapter(Mapping[str, object])
@ -183,18 +185,39 @@ async def _resolve_managed_input_file_storage_url(input_file_id: str) -> "str |
which the managed-files deployment hook still maps. This adds resolution
without changing behavior on any path that did not resolve before.
"""
db_file: Final = await _managed_input_file_row(input_file_id)
return None if db_file is None else db_file.storage_url or None
async def _managed_input_file_row(input_file_id: str) -> "LiteLLM_ManagedFileTable | None":
from litellm.proxy.proxy_server import prisma_client
if prisma_client is None:
return None
try:
db_file = 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})
except Exception as e:
verbose_proxy_logger.warning("create_batch: managed file lookup failed for %s: %s", input_file_id, e)
return None
if db_file is None:
async def _anthropic_batch_input_content(
deployment_credentials: Mapping[str, object] | None, input_file_id: str
) -> str | None:
"""Anthropic takes batch requests inline and cannot return an uploaded file, so the proxy reads its own copy."""
from litellm.proxy.proxy_server import prisma_client
if deployment_credentials is None or deployment_provider_of(deployment_credentials) != "anthropic":
return None
return db_file.storage_url or None
db_file: Final = await _managed_input_file_row(input_file_id)
if db_file is None or not db_file.storage_backend or not db_file.storage_url:
return None
try:
backend: Final = get_storage_backend(db_file.storage_backend, prisma_client=prisma_client)
content: Final = await backend.download_file(db_file.storage_url)
except ValueError as e:
raise batch_error(400, str(e))
return content.decode("utf-8")
async def _create_provider_batch_for_managed_file(
@ -202,6 +225,7 @@ async def _create_provider_batch_for_managed_file(
create_batch_data: LiteLLMBatchCreateRequest,
input_file_id: str,
unified_file_id: str,
input_content: str | None = None,
) -> LiteLLMBatch:
resolved_storage_url: Final = await _resolve_managed_input_file_storage_url(input_file_id)
request: Final[LiteLLMBatchCreateRequest] = {
@ -209,7 +233,10 @@ async def _create_provider_batch_for_managed_file(
"input_file_id": resolved_storage_url or input_file_id,
"disable_fallbacks": True,
}
response: Final = await llm_router.acreate_batch(**request)
response: Final = await llm_router.acreate_batch(
**request,
**({} if input_content is None else {ANTHROPIC_BATCH_INPUT_CONTENT_KWARG: input_content}),
)
response.input_file_id = input_file_id
response._hidden_params["unified_file_id"] = unified_file_id
return response
@ -416,8 +443,11 @@ async def create_batch(
detail={"error": "LLM Router not initialized. Ensure models added to proxy."},
)
executed_provider: Final = await resolve_litellm_executed_provider(
llm_router, model, user_api_key_dict.team_id
deployment_credentials: Final = llm_router.get_deployment_credentials_with_provider(
model_id=model, team_id=user_api_key_dict.team_id
)
executed_provider: Final = (
None if deployment_credentials is None else await litellm_executed_provider_for(deployment_credentials)
)
response = (
await _litellm_executed_batch_runner(llm_router, proxy_logging_obj).create(
@ -430,7 +460,11 @@ async def create_batch(
)
if executed_provider is not None
else await _create_provider_batch_for_managed_file(
llm_router, _create_batch_data, input_file_id, unified_file_id
llm_router,
_create_batch_data,
input_file_id,
unified_file_id,
input_content=await _anthropic_batch_input_content(deployment_credentials, input_file_id),
)
)
else:

View file

@ -190,11 +190,13 @@ class _RouterCall(Protocol):
def __call__(self, **params: object) -> Awaitable[object]: ... # kwargs-ok: the request body is passed as keywords
def litellm_executed_provider_of(credentials: Mapping[str, object]) -> str | None:
def deployment_provider_of(credentials: Mapping[str, object]) -> str | None:
explicit_provider: Final = credentials.get("custom_llm_provider")
provider: Final = (
explicit_provider if isinstance(explicit_provider, str) else _provider_of(credentials.get("model"))
)
return explicit_provider if isinstance(explicit_provider, str) else _provider_of(credentials.get("model"))
def litellm_executed_provider_of(credentials: Mapping[str, object]) -> str | None:
provider: Final = deployment_provider_of(credentials)
return provider if provider in LITELLM_EXECUTED_BATCH_PROVIDERS else None

View file

View file

@ -0,0 +1,86 @@
from __future__ import annotations
from collections.abc import Awaitable
from typing import Final, Protocol, cast # noqa: TID251 # validates dynamically loaded native callables
from litellm.rust_bridge.bindings import NativeBinding
class RustRetrieveBatch(Protocol):
def __call__(
self,
batch_id: str,
api_key: str | None,
api_base: str | None,
extra_headers: dict[str, str] | None,
timeout_seconds: float | None,
) -> dict[str, object]:
raise NotImplementedError
class RustAretrieveBatch(Protocol):
def __call__(
self,
batch_id: str,
api_key: str | None,
api_base: str | None,
extra_headers: dict[str, str] | None,
timeout_seconds: float | None,
) -> Awaitable[dict[str, object]]:
raise NotImplementedError
class RustCreateBatch(Protocol):
def __call__(
self,
input_jsonl: str,
model: str | None,
api_key: str | None,
api_base: str | None,
extra_headers: dict[str, str] | None,
timeout_seconds: float | None,
) -> dict[str, object]:
raise NotImplementedError
class RustAcreateBatch(Protocol):
def __call__(
self,
input_jsonl: str,
model: str | None,
api_key: str | None,
api_base: str | None,
extra_headers: dict[str, str] | None,
timeout_seconds: float | None,
) -> Awaitable[dict[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
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
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
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
NATIVE_RETRIEVE_BATCH: Final = NativeBinding("retrieve_batch", validate=_retrieve_binding)
NATIVE_ARETRIEVE_BATCH: Final = NativeBinding("aretrieve_batch", validate=_aretrieve_binding)
NATIVE_CREATE_BATCH: Final = NativeBinding("create_batch", validate=_create_binding)
NATIVE_ACREATE_BATCH: Final = NativeBinding("acreate_batch", validate=_acreate_binding)

View file

@ -17,6 +17,7 @@ from litellm.types.secret_managers.main import KeyManagementSystem
class Route(str, Enum):
BATCHES = "batches"
CHAT_COMPLETIONS = "chat_completions"
EMBEDDINGS = "embeddings"
MESSAGES = "messages"
@ -106,6 +107,8 @@ Rules: TypeAlias = tuple[Rule, ...]
RULES: Final[Rules] = (
LoggerRule(Rollout.RUST_OPT_IN),
RouteRule(Route.BATCHES, Rollout.RUST_OPT_IN, providers=frozenset({"anthropic"})),
RouteRule(Route.BATCHES, Rollout.PYTHON_ONLY),
RouteRule(Route.CHAT_COMPLETIONS, Rollout.PYTHON_ONLY),
RouteRule(Route.EMBEDDINGS, Rollout.PYTHON_ONLY),
RouteRule(Route.OCR, Rollout.RUST_REQUIRED, providers=frozenset({"aws_textract"})),

View file

@ -83,6 +83,11 @@ CREDS: Dict[str, Dict[str, str]] = {
"api_base": "http://vllm.test/v1",
"model": "hosted_vllm/qwen",
},
"claude-batch": {
"custom_llm_provider": "anthropic",
"api_key": "sk-ant",
"model": "anthropic/claude-batch",
},
}
# A real model-encoded file id: decodes to "azure/gpt-4o", strips to "file-original123".
@ -675,6 +680,42 @@ async def test_create__unified_file_id_resolves_real_storage_url(harness):
assert resp._hidden_params["unified_file_id"] == "unified-xyz"
@pytest.mark.asyncio
@pytest.mark.parametrize(
("model", "expected_content"),
(("claude-batch", '{"custom_id": "r1"}\n'), ("vertex-model", None)),
)
async def test_create__managed_file_content_is_read_only_for_anthropic_deployments(harness, model, expected_content):
set_body(
harness,
{"input_file_id": "litellm_proxy_unified_id", "endpoint": "/v1/chat/completions", "completion_window": "24h"},
)
fake_repo_instance = MagicMock()
fake_repo_instance.table.find_first = AsyncMock(
return_value=MagicMock(storage_backend="litellm_db", storage_url="litellm-db://file-1")
)
backend = MagicMock()
backend.download_file = AsyncMock(return_value=b'{"custom_id": "r1"}\n')
get_backend = MagicMock(return_value=backend)
prisma_client = MagicMock()
with (
patch.object(endpoints, "_is_base64_encoded_unified_file_id", return_value="unified-xyz"),
patch.object(endpoints, "get_models_from_unified_file_id", return_value=[model]),
patch.object(proxy_server, "prisma_client", prisma_client),
patch.object(endpoints, "ManagedFileRepository", MagicMock(return_value=fake_repo_instance)),
patch.object(endpoints, "get_storage_backend", get_backend),
):
await call_create(harness)
assert harness.router_kwargs().get("litellm_batch_input_content") == expected_content
if expected_content is None:
get_backend.assert_not_called()
else:
get_backend.assert_called_once_with("litellm_db", prisma_client=prisma_client)
backend.download_file.assert_awaited_once_with("litellm-db://file-1")
@pytest.mark.asyncio
async def test_create__unified_file_id_db_error_falls_back_to_raw_id(harness):
"""Resolution is additive and best-effort: a lookup error leaves the id

View file

@ -54,6 +54,10 @@ def test_shipped_decisions(
enabled: Final = environment == "1" if environment is not None else process is not False
assert catalog.rollout(context) is Rollout.RUST_OPT_OUT
assert catalog.decision(context) is (Decision.RUST_WITH_FALLBACK if enabled else Decision.PYTHON)
elif route is Route.BATCHES and provider == "anthropic":
opted_in: Final = environment == "1" if environment is not None else process is True
assert catalog.rollout(context) is Rollout.RUST_OPT_IN
assert catalog.decision(context) is (Decision.RUST_WITH_FALLBACK if opted_in else Decision.PYTHON)
elif route is Route.TRANSCRIPTION and provider == "bedrock":
assert catalog.rollout(context) is Rollout.RUST_REQUIRED
assert catalog.decision(context) is Decision.RUST_REQUIRED