mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-30 01:52:18 +00:00
wip
This commit is contained in:
parent
ef42347675
commit
ab9b8b6a07
15 changed files with 1288 additions and 51 deletions
442
litellm-rust/crates/core/src/batches/mod.rs
Normal file
442
litellm-rust/crates/core/src/batches/mod.rs
Normal 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(_)));
|
||||
}
|
||||
}
|
||||
|
|
@ -1,4 +1,5 @@
|
|||
pub mod audio_transcription;
|
||||
pub mod batches;
|
||||
pub mod chat_completions;
|
||||
pub mod constants;
|
||||
pub mod error;
|
||||
|
|
|
|||
|
|
@ -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, ¶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<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]
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
251
litellm-rust/crates/python-bridge/src/routes/batches.rs
Normal file
251
litellm-rust/crates/python-bridge/src/routes/batches.rs
Normal 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);
|
||||
}
|
||||
}
|
||||
|
|
@ -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
111
litellm/batches/dispatch.py
Normal 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)
|
||||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
0
litellm/rust_bridge/batches/__init__.py
Normal file
0
litellm/rust_bridge/batches/__init__.py
Normal file
86
litellm/rust_bridge/batches/native.py
Normal file
86
litellm/rust_bridge/batches/native.py
Normal 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)
|
||||
|
|
@ -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"})),
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue