mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
* refactor(rust): move tests.rs files inline or under tests/ Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * refactor(rust): move tests.rs files inline or under tests/ Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * refactor(rust): inline path-included test files into their owning src files Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * chore(rust): drop stray proptest regression file Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(rust): cover lowercase, empty and non-authorization headers in bearer detection Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: Yujong Lee <yujong@berri.ai> Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
471 lines
16 KiB
Rust
471 lines
16 KiB
Rust
use std::{sync::Arc, time::Duration};
|
|
|
|
use futures_util::future::BoxFuture;
|
|
use litellm_core::messages::{
|
|
Error, messages,
|
|
route::{LocalMessagesHost, MessagesCall, messages_machine},
|
|
types::{MessagesRequest, MessagesShaping},
|
|
};
|
|
use litellm_secrets::{SecretValue, source::SecretSource};
|
|
use serde_json::{Map, Value, json};
|
|
use tokio::{
|
|
io::{AsyncReadExt, AsyncWriteExt},
|
|
net::{TcpListener, TcpStream},
|
|
};
|
|
|
|
struct RecordingSecrets {
|
|
values: Vec<(&'static str, String)>,
|
|
fails: bool,
|
|
requested: std::sync::Mutex<Vec<String>>,
|
|
}
|
|
|
|
impl RecordingSecrets {
|
|
fn new(values: Vec<(&'static str, String)>, fails: bool) -> Self {
|
|
Self {
|
|
values,
|
|
fails,
|
|
requested: std::sync::Mutex::new(Vec::new()),
|
|
}
|
|
}
|
|
}
|
|
|
|
impl SecretSource for RecordingSecrets {
|
|
fn get_secret_str<'a>(
|
|
&'a self,
|
|
name: &'a str,
|
|
) -> BoxFuture<'a, Result<Option<SecretValue>, litellm_secrets::Error>> {
|
|
Box::pin(async move {
|
|
self.requested.lock().unwrap().push(name.to_string());
|
|
if self.fails {
|
|
return Err(litellm_secrets::Error::ManagedSecretMissing);
|
|
}
|
|
Ok(self
|
|
.values
|
|
.iter()
|
|
.find(|(key, _)| *key == name)
|
|
.map(|(_, value)| SecretValue::new(value.clone())))
|
|
})
|
|
}
|
|
}
|
|
|
|
fn secrets_call() -> MessagesCall {
|
|
let Value::Object(body) = json!({
|
|
"model": "claude-sonnet-4-5",
|
|
"max_tokens": 16,
|
|
"messages": [{"role": "user", "content": "hi"}]
|
|
}) else {
|
|
unreachable!("literal object")
|
|
};
|
|
MessagesCall {
|
|
model: "claude-sonnet-4-5".into(),
|
|
body,
|
|
api_key: None,
|
|
api_base: None,
|
|
custom_llm_provider: Some("anthropic".into()),
|
|
extra_headers: None,
|
|
provider_specific_header: None,
|
|
timeout: Some(Duration::from_secs(5)),
|
|
shaping: MessagesShaping::default(),
|
|
}
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn route_surfaces_a_secret_manager_failure_before_the_call() {
|
|
let Err(error) = litellm_host::run::run(
|
|
messages_machine(Arc::new(RecordingSecrets::new(Vec::new(), true))),
|
|
&LocalMessagesHost::new(secrets_call()),
|
|
)
|
|
.await
|
|
else {
|
|
panic!("a secret manager failure fails the call");
|
|
};
|
|
assert!(
|
|
matches!(&error, Error::Secret(source) if matches!(source.source_error(), litellm_secrets::Error::ManagedSecretMissing)),
|
|
"{error:?}"
|
|
);
|
|
}
|
|
|
|
async fn read_http_request(socket: &mut TcpStream) -> String {
|
|
let mut request = Vec::new();
|
|
let mut buffer = [0_u8; 1024];
|
|
let header_end = loop {
|
|
let n = socket.read(&mut buffer).await.expect("reads request");
|
|
if n == 0 {
|
|
break request.len();
|
|
}
|
|
request.extend_from_slice(&buffer[..n]);
|
|
if let Some(position) = request.windows(4).position(|window| window == b"\r\n\r\n") {
|
|
break position + 4;
|
|
}
|
|
};
|
|
let headers = String::from_utf8_lossy(&request[..header_end]);
|
|
let content_length = headers
|
|
.lines()
|
|
.find_map(|line| {
|
|
let (name, value) = line.split_once(':')?;
|
|
name.eq_ignore_ascii_case("content-length")
|
|
.then(|| value.trim().parse::<usize>().ok())
|
|
.flatten()
|
|
})
|
|
.unwrap_or(0);
|
|
while request.len().saturating_sub(header_end) < content_length {
|
|
let n = socket.read(&mut buffer).await.expect("reads body");
|
|
if n == 0 {
|
|
break;
|
|
}
|
|
request.extend_from_slice(&buffer[..n]);
|
|
}
|
|
String::from_utf8(request).expect("request is utf8")
|
|
}
|
|
|
|
fn write_response(body: &str) -> String {
|
|
format!(
|
|
"HTTP/1.1 200 OK\r\ncontent-type: application/json\r\ncontent-length: {}\r\nconnection: close\r\n\r\n{}",
|
|
body.len(),
|
|
body
|
|
)
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn messages_round_trip_builds_azure_request_and_passes_response_through() {
|
|
let listener = TcpListener::bind("127.0.0.1:0").await.expect("binds");
|
|
let addr = listener.local_addr().expect("addr");
|
|
|
|
let server = tokio::spawn(async move {
|
|
let (mut socket, _) = listener.accept().await.expect("accepts request");
|
|
let request = read_http_request(&mut socket).await;
|
|
let response_body = r#"{"id":"msg_1","type":"message","role":"assistant","content":[{"type":"text","text":"hi"}],"model":"claude-sonnet-4-5","stop_reason":"end_turn","usage":{"input_tokens":1,"output_tokens":2}}"#;
|
|
socket
|
|
.write_all(write_response(response_body).as_bytes())
|
|
.await
|
|
.expect("writes response");
|
|
request
|
|
});
|
|
|
|
let response = messages(MessagesRequest {
|
|
model: "claude-sonnet-4-5",
|
|
body: json!({
|
|
"model": "claude-sonnet-4-5",
|
|
"max_tokens": 1024,
|
|
"messages": [{
|
|
"role": "user",
|
|
"content": [{
|
|
"type": "text",
|
|
"text": "hi",
|
|
"cache_control": {"type": "ephemeral", "scope": "global"}
|
|
}]
|
|
}]
|
|
}),
|
|
api_key: Some("sk-azure"),
|
|
api_base: Some(&format!("http://{addr}")),
|
|
custom_llm_provider: Some("azure_ai"),
|
|
extra_headers: None,
|
|
provider_specific_header: None,
|
|
timeout: Some(Duration::from_secs(5)),
|
|
shaping: MessagesShaping::default(),
|
|
})
|
|
.await
|
|
.expect("messages request succeeds");
|
|
|
|
assert_eq!(response.content[0]["text"], "hi");
|
|
assert_eq!(response.stop_reason.as_deref(), Some("end_turn"));
|
|
|
|
let request = server.await.expect("server task completes");
|
|
let (head, body) = request.split_once("\r\n\r\n").expect("has body");
|
|
assert!(head.starts_with("POST /anthropic/v1/messages "), "{head}");
|
|
let head_lower = head.to_ascii_lowercase();
|
|
assert!(head_lower.contains("x-api-key: sk-azure"), "{head}");
|
|
assert!(
|
|
head_lower.contains("anthropic-version: 2023-06-01"),
|
|
"{head}"
|
|
);
|
|
assert!(
|
|
head_lower.contains("content-type: application/json"),
|
|
"{head}"
|
|
);
|
|
|
|
let sent_body: Value = serde_json::from_str(body).expect("body is json");
|
|
assert_eq!(
|
|
sent_body["messages"][0]["content"][0]["cache_control"],
|
|
json!({"type": "ephemeral"})
|
|
);
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn messages_round_trip_builds_native_anthropic_request() {
|
|
let listener = TcpListener::bind("127.0.0.1:0").await.expect("binds");
|
|
let addr = listener.local_addr().expect("addr");
|
|
|
|
let server = tokio::spawn(async move {
|
|
let (mut socket, _) = listener.accept().await.expect("accepts request");
|
|
let request = read_http_request(&mut socket).await;
|
|
let response_body = r#"{"id":"msg_1","type":"message","role":"assistant","content":[{"type":"text","text":"hi"}],"model":"claude-sonnet-4-5","stop_reason":"end_turn","usage":{"input_tokens":1,"output_tokens":2}}"#;
|
|
socket
|
|
.write_all(write_response(response_body).as_bytes())
|
|
.await
|
|
.expect("writes response");
|
|
request
|
|
});
|
|
|
|
let response = messages(MessagesRequest {
|
|
model: "claude-sonnet-4-5",
|
|
body: json!({
|
|
"model": "claude-sonnet-4-5",
|
|
"max_tokens": 1024,
|
|
"messages": [{"role": "user", "content": "hi"}]
|
|
}),
|
|
api_key: Some("sk-ant"),
|
|
api_base: Some(&format!("http://{addr}")),
|
|
custom_llm_provider: Some("anthropic"),
|
|
extra_headers: None,
|
|
provider_specific_header: None,
|
|
timeout: Some(Duration::from_secs(5)),
|
|
shaping: MessagesShaping::default(),
|
|
})
|
|
.await
|
|
.expect("messages request succeeds");
|
|
|
|
assert_eq!(response.content[0]["text"], "hi");
|
|
assert_eq!(response.stop_reason.as_deref(), Some("end_turn"));
|
|
|
|
let request = server.await.expect("server task completes");
|
|
let (head, _) = request.split_once("\r\n\r\n").expect("has body");
|
|
assert!(head.starts_with("POST /v1/messages "), "{head}");
|
|
let head_lower = head.to_ascii_lowercase();
|
|
assert!(head_lower.contains("x-api-key: sk-ant"), "{head}");
|
|
assert!(
|
|
head_lower.contains("anthropic-version: 2023-06-01"),
|
|
"{head}"
|
|
);
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn messages_does_not_duplicate_auth_when_x_api_key_supplied() {
|
|
let listener = TcpListener::bind("127.0.0.1:0").await.expect("binds");
|
|
let addr = listener.local_addr().expect("addr");
|
|
|
|
let server = tokio::spawn(async move {
|
|
let (mut socket, _) = listener.accept().await.expect("accepts request");
|
|
let request = read_http_request(&mut socket).await;
|
|
let response_body =
|
|
r#"{"id":"msg_2","type":"message","role":"assistant","content":[],"model":"m"}"#;
|
|
socket
|
|
.write_all(write_response(response_body).as_bytes())
|
|
.await
|
|
.expect("writes response");
|
|
request
|
|
});
|
|
|
|
let mut headers = Map::new();
|
|
headers.insert(
|
|
"x-api-key".to_string(),
|
|
Value::String("from-python".to_string()),
|
|
);
|
|
headers.insert(
|
|
"anthropic-beta".to_string(),
|
|
Value::String("token-efficient-tools-2025-02-19".to_string()),
|
|
);
|
|
|
|
messages(MessagesRequest {
|
|
model: "claude-sonnet-4-5",
|
|
body: json!({"model": "claude-sonnet-4-5", "max_tokens": 8, "messages": []}),
|
|
api_key: Some("rust-fallback-key"),
|
|
api_base: Some(&format!("http://{addr}")),
|
|
custom_llm_provider: Some("azure_ai"),
|
|
extra_headers: Some(headers),
|
|
provider_specific_header: None,
|
|
timeout: Some(Duration::from_secs(5)),
|
|
shaping: MessagesShaping::default(),
|
|
})
|
|
.await
|
|
.expect("messages request succeeds");
|
|
|
|
let request = server.await.expect("server task completes");
|
|
let head = request
|
|
.split_once("\r\n\r\n")
|
|
.expect("has body")
|
|
.0
|
|
.to_ascii_lowercase();
|
|
let api_key_count = head
|
|
.lines()
|
|
.filter(|line| line.starts_with("x-api-key:"))
|
|
.count();
|
|
assert_eq!(api_key_count, 1, "{head}");
|
|
assert!(head.contains("x-api-key: from-python"), "{head}");
|
|
assert!(
|
|
head.contains("anthropic-beta: token-efficient-tools-2025-02-19"),
|
|
"{head}"
|
|
);
|
|
assert!(!head.contains("rust-fallback-key"), "{head}");
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn messages_forwards_entra_id_bearer_without_requiring_api_key() {
|
|
let listener = TcpListener::bind("127.0.0.1:0").await.expect("binds");
|
|
let addr = listener.local_addr().expect("addr");
|
|
|
|
let server = tokio::spawn(async move {
|
|
let (mut socket, _) = listener.accept().await.expect("accepts request");
|
|
let request = read_http_request(&mut socket).await;
|
|
let response_body =
|
|
r#"{"id":"msg_3","type":"message","role":"assistant","content":[],"model":"m"}"#;
|
|
socket
|
|
.write_all(write_response(response_body).as_bytes())
|
|
.await
|
|
.expect("writes response");
|
|
request
|
|
});
|
|
|
|
let mut headers = Map::new();
|
|
headers.insert(
|
|
"Authorization".to_string(),
|
|
Value::String("Bearer entra-token".to_string()),
|
|
);
|
|
|
|
messages(MessagesRequest {
|
|
model: "claude-sonnet-4-5",
|
|
body: json!({"model": "claude-sonnet-4-5", "max_tokens": 8, "messages": []}),
|
|
api_key: None,
|
|
api_base: Some(&format!("http://{addr}")),
|
|
custom_llm_provider: Some("azure_ai"),
|
|
extra_headers: Some(headers),
|
|
provider_specific_header: None,
|
|
timeout: Some(Duration::from_secs(5)),
|
|
shaping: MessagesShaping::default(),
|
|
})
|
|
.await
|
|
.expect("entra id request succeeds without api key");
|
|
|
|
let request = server.await.expect("server task completes");
|
|
let head = request
|
|
.split_once("\r\n\r\n")
|
|
.expect("has body")
|
|
.0
|
|
.to_ascii_lowercase();
|
|
assert!(head.contains("authorization: bearer entra-token"), "{head}");
|
|
assert!(!head.contains("x-api-key"), "{head}");
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn messages_requires_auth_when_no_key_and_no_header() {
|
|
let err = messages(MessagesRequest {
|
|
model: "claude-sonnet-4-5",
|
|
body: json!({"model": "claude-sonnet-4-5", "max_tokens": 8, "messages": []}),
|
|
api_key: None,
|
|
api_base: Some("http://127.0.0.1:1"),
|
|
custom_llm_provider: Some("azure_ai"),
|
|
extra_headers: None,
|
|
provider_specific_header: None,
|
|
timeout: Some(Duration::from_millis(50)),
|
|
shaping: MessagesShaping::default(),
|
|
})
|
|
.await
|
|
.expect_err("missing auth errors");
|
|
|
|
assert!(matches!(err, Error::Auth(_)));
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn messages_ignores_malformed_authorization_and_uses_api_key() {
|
|
let listener = TcpListener::bind("127.0.0.1:0").await.expect("binds");
|
|
let addr = listener.local_addr().expect("addr");
|
|
|
|
let server = tokio::spawn(async move {
|
|
let (mut socket, _) = listener.accept().await.expect("accepts request");
|
|
let request = read_http_request(&mut socket).await;
|
|
let response_body =
|
|
r#"{"id":"msg_4","type":"message","role":"assistant","content":[],"model":"m"}"#;
|
|
socket
|
|
.write_all(write_response(response_body).as_bytes())
|
|
.await
|
|
.expect("writes response");
|
|
request
|
|
});
|
|
|
|
let mut headers = Map::new();
|
|
headers.insert(
|
|
"Authorization".to_string(),
|
|
Value::String("Bearer ".to_string()),
|
|
);
|
|
|
|
messages(MessagesRequest {
|
|
model: "claude-sonnet-4-5",
|
|
body: json!({"model": "claude-sonnet-4-5", "max_tokens": 8, "messages": []}),
|
|
api_key: Some("sk-azure"),
|
|
api_base: Some(&format!("http://{addr}")),
|
|
custom_llm_provider: Some("azure_ai"),
|
|
extra_headers: Some(headers),
|
|
provider_specific_header: None,
|
|
timeout: Some(Duration::from_secs(5)),
|
|
shaping: MessagesShaping::default(),
|
|
})
|
|
.await
|
|
.expect("falls back to api key");
|
|
|
|
let request = server.await.expect("server task completes");
|
|
let head = request
|
|
.split_once("\r\n\r\n")
|
|
.expect("has body")
|
|
.0
|
|
.to_ascii_lowercase();
|
|
assert!(head.contains("x-api-key: sk-azure"), "{head}");
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn messages_maps_provider_error_status_to_http_error() {
|
|
let listener = TcpListener::bind("127.0.0.1:0").await.expect("binds");
|
|
let addr = listener.local_addr().expect("addr");
|
|
|
|
tokio::spawn(async move {
|
|
let (mut socket, _) = listener.accept().await.expect("accepts request");
|
|
let _ = read_http_request(&mut socket).await;
|
|
let body = "unauthorized";
|
|
let response = format!(
|
|
"HTTP/1.1 401 Unauthorized\r\ncontent-length: {}\r\nconnection: close\r\n\r\n{}",
|
|
body.len(),
|
|
body
|
|
);
|
|
socket
|
|
.write_all(response.as_bytes())
|
|
.await
|
|
.expect("writes response");
|
|
});
|
|
|
|
let err = messages(MessagesRequest {
|
|
model: "claude-sonnet-4-5",
|
|
body: json!({"model": "claude-sonnet-4-5", "max_tokens": 8, "messages": []}),
|
|
api_key: Some("sk-azure"),
|
|
api_base: Some(&format!("http://{addr}")),
|
|
custom_llm_provider: Some("azure_ai"),
|
|
extra_headers: None,
|
|
provider_specific_header: None,
|
|
timeout: Some(Duration::from_secs(5)),
|
|
shaping: MessagesShaping::default(),
|
|
})
|
|
.await
|
|
.expect_err("provider error propagates");
|
|
|
|
assert!(matches!(
|
|
err,
|
|
Error::Transport(litellm_http::transport::Error::Http { status: 401, .. })
|
|
));
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn messages_rejects_unsupported_provider() {
|
|
let err = messages(MessagesRequest {
|
|
model: "claude-3-5-sonnet",
|
|
body: json!({"model": "claude-3-5-sonnet", "max_tokens": 8, "messages": []}),
|
|
api_key: Some("sk"),
|
|
api_base: Some("http://127.0.0.1:1"),
|
|
custom_llm_provider: Some("openai"),
|
|
extra_headers: None,
|
|
provider_specific_header: None,
|
|
timeout: Some(Duration::from_millis(50)),
|
|
shaping: MessagesShaping::default(),
|
|
})
|
|
.await
|
|
.expect_err("unsupported provider errors");
|
|
|
|
assert!(matches!(err, Error::InvalidProvider(provider) if provider == "openai"));
|
|
}
|