diff --git a/litellm-rust/crates/core/src/batches/mod.rs b/litellm-rust/crates/core/src/batches/mod.rs new file mode 100644 index 00000000000..17cee372b0f --- /dev/null +++ b/litellm-rust/crates/core/src/batches/mod.rs @@ -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, +} + +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 = 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 + Sync), +) -> Result { + 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 + Sync), +) -> Result { + 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>, + timeout: Option, +) -> Result { + 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) { + 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::().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 + 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::(&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(_))); + } +} diff --git a/litellm-rust/crates/core/src/lib.rs b/litellm-rust/crates/core/src/lib.rs index afe5ea595aa..8b98676f833 100644 --- a/litellm-rust/crates/core/src/lib.rs +++ b/litellm-rust/crates/core/src/lib.rs @@ -1,4 +1,5 @@ pub mod audio_transcription; +pub mod batches; pub mod chat_completions; pub mod constants; pub mod error; diff --git a/litellm-rust/crates/llms/src/anthropic/batches/transformation.rs b/litellm-rust/crates/llms/src/anthropic/batches/transformation.rs index 326b03e7394..73405b96b3e 100644 --- a/litellm-rust/crates/llms/src/anthropic/batches/transformation.rs +++ b/litellm-rust/crates/llms/src/anthropic/batches/transformation.rs @@ -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, +} #[derive(Clone, Debug, Default, PartialEq, Eq, Serialize, Deserialize)] pub struct AnthropicBatchRequestCounts { @@ -95,13 +110,17 @@ pub trait AnthropicBatchesConfig { env_lookup: &dyn Fn(&str) -> Option, ) -> Result; - fn transform_create_batch_request(&self) -> Result; + fn transform_create_batch_request( + &self, + model: Option<&str>, + input_jsonl: &str, + ) -> Result; fn transform_create_batch_response( &self, response: AnthropicMessageBatch, now: i64, - ) -> Result; + ) -> 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 { + 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 = line + .body + .get("messages") + .cloned() + .map(serde_json::from_value) + .transpose() + .map_err(|error| invalid(&format!("invalid messages: {error}")))? + .filter(|messages: &Vec| !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::, _>>()?; + if !params.contains_key("max_tokens") { + return Err(invalid("max_tokens is required")); + } + if let Some(reason) = config.unsupported_reason(&messages, ¶ms) { + 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 { - Err(Error::Unsupported("Anthropic message batch creation")) + fn transform_create_batch_request( + &self, + model: Option<&str>, + input_jsonl: &str, + ) -> Result { + let requests = input_jsonl + .lines() + .map(str::trim) + .filter(|line| !line.is_empty()) + .map(|line| batch_request(model, line)) + .collect::, _>>()?; + 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 { - 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] diff --git a/litellm-rust/crates/python-bridge/src/lib.rs b/litellm-rust/crates/python-bridge/src/lib.rs index 022de0f9ef7..4c3b2787f12 100644 --- a/litellm-rust/crates/python-bridge/src/lib.rs +++ b/litellm-rust/crates/python-bridge/src/lib.rs @@ -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", diff --git a/litellm-rust/crates/python-bridge/src/routes/batches.rs b/litellm-rust/crates/python-bridge/src/routes/batches.rs new file mode 100644 index 00000000000..42cdcd9634c --- /dev/null +++ b/litellm-rust/crates/python-bridge/src/routes/batches.rs @@ -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, + api_base: Option, + extra_headers: Option>, + timeout_seconds: Option, +} + +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 { + std::env::var(name).ok() +} + +async fn retrieve(batch_id: String, connection: OwnedConnection) -> Result { + run_retrieve_batch( + http_client(), + RetrieveBatchRequest { + batch_id: &batch_id, + connection: connection.borrow(), + }, + &env_lookup, + ) + .await +} + +async fn create( + input_jsonl: String, + model: Option, + connection: OwnedConnection, +) -> Result { + 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, + api_base: Option, + extra_headers: Option>, + timeout_seconds: Option, +) -> PyResult> { + 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, + api_base: Option, + extra_headers: Option>, + timeout_seconds: Option, +) -> PyResult> { + 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, + api_key: Option, + api_base: Option, + extra_headers: Option>, + timeout_seconds: Option, +) -> PyResult> { + 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, + api_key: Option, + api_base: Option, + extra_headers: Option>, + timeout_seconds: Option, +) -> PyResult> { + 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::(py) { + return Raised::Declined(error.value(py).to_string()); + } + if error.is_instance_of::(py) { + 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(); + 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); + } +} diff --git a/litellm-rust/crates/python-bridge/src/routes/mod.rs b/litellm-rust/crates/python-bridge/src/routes/mod.rs index dd694fa589f..f28bd38e301 100644 --- a/litellm-rust/crates/python-bridge/src/routes/mod.rs +++ b/litellm-rust/crates/python-bridge/src/routes/mod.rs @@ -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; diff --git a/litellm/batches/dispatch.py b/litellm/batches/dispatch.py new file mode 100644 index 00000000000..39b2377cf54 --- /dev/null +++ b/litellm/batches/dispatch.py @@ -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) diff --git a/litellm/batches/main.py b/litellm/batches/main.py index 76b6c73b375..37570c3e6b5 100644 --- a/litellm/batches/main.py +++ b/litellm/batches/main.py @@ -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( diff --git a/litellm/proxy/batches_endpoints/endpoints.py b/litellm/proxy/batches_endpoints/endpoints.py index f6c86d75169..4e79f9181f3 100644 --- a/litellm/proxy/batches_endpoints/endpoints.py +++ b/litellm/proxy/batches_endpoints/endpoints.py @@ -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: diff --git a/litellm/proxy/batches_endpoints/litellm_executed_batches.py b/litellm/proxy/batches_endpoints/litellm_executed_batches.py index 67201d99422..f1f17f225cd 100644 --- a/litellm/proxy/batches_endpoints/litellm_executed_batches.py +++ b/litellm/proxy/batches_endpoints/litellm_executed_batches.py @@ -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 diff --git a/litellm/rust_bridge/batches/__init__.py b/litellm/rust_bridge/batches/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/litellm/rust_bridge/batches/native.py b/litellm/rust_bridge/batches/native.py new file mode 100644 index 00000000000..60478f19dde --- /dev/null +++ b/litellm/rust_bridge/batches/native.py @@ -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) diff --git a/litellm/rust_bridge/catalog.py b/litellm/rust_bridge/catalog.py index 91cbed89084..07f82bf39f7 100644 --- a/litellm/rust_bridge/catalog.py +++ b/litellm/rust_bridge/catalog.py @@ -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"})), diff --git a/tests/test_litellm/proxy/batches_endpoints/test_endpoints.py b/tests/test_litellm/proxy/batches_endpoints/test_endpoints.py index 2d597abf3b8..8786cbcbaf2 100644 --- a/tests/test_litellm/proxy/batches_endpoints/test_endpoints.py +++ b/tests/test_litellm/proxy/batches_endpoints/test_endpoints.py @@ -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 diff --git a/tests/test_litellm/rust_bridge/test_catalog.py b/tests/test_litellm/rust_bridge/test_catalog.py index 45c35dc6802..c9b15710b4a 100644 --- a/tests/test_litellm/rust_bridge/test_catalog.py +++ b/tests/test_litellm/rust_bridge/test_catalog.py @@ -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