Merge remote-tracking branch 'origin/main' into litellm_agent365_fail_open_default

This commit is contained in:
yucheng 2026-09-25 17:56:16 +00:00
commit 6e03a36e16
43 changed files with 3425 additions and 98 deletions

152
litellm-rust/Cargo.lock generated
View file

@ -73,6 +73,15 @@ version = "1.0.104"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "330a5ed07fa54e4702c9d6c4174f74427fc0ef6e214bbd677ae50a5099946470"
[[package]]
name = "arbitrary"
version = "1.4.2"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "c3d036a3c4ab069c7b410a2ce876bd74808d2d0888a82667669f8e783a898bf1"
dependencies = [
"derive_arbitrary",
]
[[package]]
name = "arc-swap"
version = "1.9.2"
@ -1470,6 +1479,17 @@ dependencies = [
"serde_core",
]
[[package]]
name = "derive_arbitrary"
version = "1.4.2"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "1e567bd82dcff979e4b03460c307b3cdc9e96fde3d73bed1496d2bc75d9dd62a"
dependencies = [
"proc-macro2",
"quote",
"syn 2.0.119",
]
[[package]]
name = "derive_builder"
version = "0.20.2"
@ -1643,6 +1663,16 @@ version = "2.5.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "da7c62ceae207dd37ea5b845da6a0696c799f85e97da1ab5b7910be3c1c80223"
[[package]]
name = "filetime"
version = "0.2.29"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "5c287a33c7f0a620c38e641e7f60827713987b3c0f26e8ddc9462cc69cf75759"
dependencies = [
"cfg-if",
"libc",
]
[[package]]
name = "find-msvc-tools"
version = "0.1.9"
@ -3453,6 +3483,27 @@ dependencies = [
"veil",
]
[[package]]
name = "litellm-testkit"
version = "0.1.0"
dependencies = [
"flate2",
"futures-util",
"reqwest 0.12.28",
"rstest",
"semver",
"serde",
"serde_json",
"sha2 0.10.9",
"tar",
"target-lexicon",
"tempfile",
"thiserror 2.0.19",
"tokio",
"toml",
"zip",
]
[[package]]
name = "litellm-token-counter"
version = "0.1.0"
@ -5206,6 +5257,15 @@ dependencies = [
"serde_core",
]
[[package]]
name = "serde_spanned"
version = "1.1.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "6662b5879511e06e8999a8a235d848113e942c9124f211511b16466ee2995f26"
dependencies = [
"serde_core",
]
[[package]]
name = "serde_urlencoded"
version = "0.7.1"
@ -5537,6 +5597,17 @@ version = "0.2.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "7b2093cf4c8eb1e67749a6762251bc9cd836b6fc171623bd0a9d324d37af2417"
[[package]]
name = "tar"
version = "0.4.46"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "3f6221d9a6003c78398e3b239969f352578258df48c8eb051caadae0015bc840"
dependencies = [
"filetime",
"libc",
"xattr",
]
[[package]]
name = "target-lexicon"
version = "0.13.5"
@ -5806,6 +5877,30 @@ dependencies = [
"tokio",
]
[[package]]
name = "toml"
version = "0.9.12+spec-1.1.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "cf92845e79fc2e2def6a5d828f0801e29a2f8acc037becc5ab08595c7d5e9863"
dependencies = [
"indexmap 2.14.0",
"serde_core",
"serde_spanned",
"toml_datetime 0.7.5+spec-1.1.0",
"toml_parser",
"toml_writer",
"winnow 0.7.15",
]
[[package]]
name = "toml_datetime"
version = "0.7.5+spec-1.1.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "92e1cfed4a3038bc5a127e35a2d360f145e1f4b971b551a2ba5fd7aedf7e1347"
dependencies = [
"serde_core",
]
[[package]]
name = "toml_datetime"
version = "1.1.1+spec-1.1.0"
@ -5822,9 +5917,9 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "6975367e4d2ef766d86af01ffad14b622fecc8d4357a998fbc4deb6e9bacaf9b"
dependencies = [
"indexmap 2.14.0",
"toml_datetime",
"toml_datetime 1.1.1+spec-1.1.0",
"toml_parser",
"winnow",
"winnow 1.0.4",
]
[[package]]
@ -5833,9 +5928,15 @@ version = "1.1.3+spec-1.1.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "1d38ac1cf9b95face32296c0a3ede1fdc270627c9d9c02a7274dd6d960dc4d56"
dependencies = [
"winnow",
"winnow 1.0.4",
]
[[package]]
name = "toml_writer"
version = "1.1.2+spec-1.1.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "7d56353a2a665ad0f41a421187180aab746c8c325620617ad883a99a1cbe66d2"
[[package]]
name = "tonic"
version = "0.14.6"
@ -6597,6 +6698,12 @@ version = "0.52.6"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "589f6da84c646204747d1270a2a5661ea66ed1cced2631d546fdfb155959f9ec"
[[package]]
name = "winnow"
version = "0.7.15"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "df79d97927682d2fd8adb29682d1140b343be4ac0f08fd68b7765d9c059d3945"
[[package]]
name = "winnow"
version = "1.0.4"
@ -6659,6 +6766,16 @@ dependencies = [
"time",
]
[[package]]
name = "xattr"
version = "1.6.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "32e45ad4206f6d2479085147f02bc2ef834ac85886624a23575ae137c8aa8156"
dependencies = [
"libc",
"rustix",
]
[[package]]
name = "xmlparser"
version = "0.13.6"
@ -6784,6 +6901,23 @@ dependencies = [
"syn 2.0.119",
]
[[package]]
name = "zip"
version = "2.4.2"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "fabe6324e908f85a1c52063ce7aa26b68dcb7eb6dbc83a2d148403c9bc3eba50"
dependencies = [
"arbitrary",
"crc32fast",
"crossbeam-utils",
"displaydoc",
"flate2",
"indexmap 2.14.0",
"memchr",
"thiserror 2.0.19",
"zopfli",
]
[[package]]
name = "zlib-rs"
version = "0.6.7"
@ -6795,3 +6929,15 @@ name = "zmij"
version = "1.0.23"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "29666d0abbfad1e3dc4dcf6144730dd3a3ab225bbbdac83319345b1b44ccfc1b"
[[package]]
name = "zopfli"
version = "0.8.3"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "f05cd8797d63865425ff89b5c4a48804f35ba0ce8d125800027ad6017d2b5249"
dependencies = [
"bumpalo",
"crc32fast",
"log",
"simd-adler32",
]

View file

@ -81,6 +81,12 @@ tokio = { version = "1", features = ["rt-multi-thread", "macros", "time", "net"]
tokio-tungstenite = { version = "0.24", default-features = false, features = ["connect", "rustls-tls-native-roots"] }
futures-util = { version = "0.3", default-features = false, features = ["sink", "std"] }
base64 = "0.22"
flate2 = "1"
semver = "1"
tar = "0.4"
target-lexicon = "0.13.5"
tempfile = "3"
zip = { version = "2", default-features = false, features = ["deflate"] }
moka = { version = "0.12.16", features = ["future"] }
strum = { version = "0.28.0", features = ["derive"] }
url = "2.5.8"

View file

@ -17,8 +17,10 @@ pub(super) async fn send(
body: &Value,
timeout: Option<Duration>,
) -> Result<reqwest::Response, Error> {
let encoded = serde_json::to_vec(body)
.map_err(|err| Error::InvalidRequest(format!("failed to encode messages body: {err}")))?;
let builder = headers.iter().fold(
http_client().post(url).json(body),
http_client().post(url).body(encoded),
|builder, (key, value)| builder.header(key, value),
);
let builder = match timeout {

View file

@ -6,7 +6,6 @@ use std::{
use bytes::Bytes;
use litellm_auth::SecretValue;
use litellm_core_utils::get_llm_provider_logic::get_custom_llm_provider;
use litellm_host::{
event::{MachineEvent, RawResponse, RequestContext, WireRequest},
host::{Demand, Host},
@ -22,7 +21,6 @@ use serde_json::{Map, Value};
use super::{
Error,
common_utils::messages_provider_config,
handler::{decode_response, network, provider_error, send},
prepare::{prepare_provider_request, resolve_provider},
types::{MessagesRequest, MessagesShaping},
@ -54,6 +52,11 @@ pub enum MessagesOutput {
Streamed,
}
/// The upstream response as the caller sees it at stream hand-off, before any chunk.
pub struct MessagesStreamHead {
pub headers: Vec<(String, String)>,
}
pub struct Messages;
impl Protocol for Messages {
@ -62,7 +65,7 @@ impl Protocol for Messages {
type Projection = MessagesCall;
type Op = Infallible;
type Chunk = Bytes;
type StreamHead = ();
type StreamHead = MessagesStreamHead;
}
impl From<MachineFault> for Error {
@ -77,19 +80,6 @@ impl From<MachineFault> for Error {
pub type MessagesHost = HostChannel<Messages>;
pub type MessagesMachine = CallMachine<Messages>;
/// Whether this route serves the request, decided before any callback runs so a host
/// can still run its own path.
pub fn supports(model: &str, custom_llm_provider: Option<&str>, stream: bool) -> bool {
let provider = get_custom_llm_provider(model, custom_llm_provider)
.map(|resolved| resolved.custom_llm_provider)
.or(custom_llm_provider);
match provider {
Some(ANTHROPIC_MESSAGES_PROVIDER) => true,
Some(provider) => !stream && messages_provider_config(provider).is_some(),
None => false,
}
}
/// The in-process host for a request already in hand. It answers projection once and
/// observes nothing.
pub struct LocalMessagesHost {
@ -152,8 +142,11 @@ async fn execute(
model: request.model.clone(),
custom_llm_provider: request.provider.clone(),
optional_params: Value::Object(
call.body
.iter()
request
.body
.as_object()
.into_iter()
.flatten()
.filter(|(name, _)| !matches!(name.as_str(), "model" | "messages"))
.map(|(name, value)| (name.clone(), value.clone()))
.collect(),
@ -193,7 +186,14 @@ async fn relay(
host: &MessagesHost,
mut response: reqwest::Response,
) -> Result<MessagesOutput, Error> {
if host.open(()).await? == Demand::Detached {
let head = MessagesStreamHead {
headers: response
.headers()
.iter()
.filter_map(|(name, value)| Some((name.to_string(), value.to_str().ok()?.to_string())))
.collect(),
};
if host.open(head).await? == Demand::Detached {
return Ok(MessagesOutput::Streamed);
}
while let Some(chunk) = response.chunk().await.map_err(network)? {

View file

@ -0,0 +1,210 @@
use std::{convert::Infallible, sync::Mutex};
use litellm_core::messages::route::Messages;
use litellm_host::{
event::{CallEvent, MachineEvent, RequestContext, WireRequest},
host::Host,
};
use litellm_llms::anthropic::common_utils::AnthropicModelCapabilities;
use rstest::rstest;
use super::*;
type Rewrite = Box<dyn Fn(WireRequest) -> Result<WireRequest, Error> + Send + Sync>;
/// Projects like `LocalMessagesHost`, answers `before_send` through `rewrite`, and keeps
/// every event the driver emits.
struct RecordingHost {
call: LocalMessagesHost,
rewrite: Rewrite,
events: Mutex<Vec<CallEvent>>,
optional_params: Mutex<Vec<Value>>,
}
impl RecordingHost {
fn new(call: MessagesCall, rewrite: Rewrite) -> Self {
Self {
call: LocalMessagesHost::new(call),
rewrite,
events: Mutex::new(Vec::new()),
optional_params: Mutex::new(Vec::new()),
}
}
fn passthrough(call: MessagesCall) -> Self {
Self::new(call, Box::new(Ok))
}
fn raw_responses(&self) -> Vec<String> {
self.events
.lock()
.unwrap()
.iter()
.filter_map(|event| match event {
CallEvent::Machine(MachineEvent::ResponseReceived { raw }) => {
Some(raw.body.clone())
}
_ => None,
})
.collect()
}
}
impl Host<Messages> for RecordingHost {
async fn project(&self) -> Result<MessagesCall, Error> {
self.call.project().await
}
async fn custom_op(&self, op: Infallible) -> Result<(), Error> {
match op {}
}
async fn before_send(
&self,
wire: WireRequest,
context: &RequestContext,
) -> Result<WireRequest, Error> {
self.optional_params
.lock()
.unwrap()
.push(context.optional_params.clone());
(self.rewrite)(wire)
}
async fn emit(&self, event: &CallEvent) -> Result<(), Error> {
self.events.lock().unwrap().push(event.clone());
Ok(())
}
}
async fn run_through(host: &RecordingHost) -> Result<MessagesOutput, Error> {
litellm_host::run::run(messages_machine(Arc::new(RecordingSecrets::empty())), host).await
}
fn authenticated(call: MessagesCall, api_base: String) -> MessagesCall {
MessagesCall {
api_key: Some("sk-ant".into()),
api_base: Some(api_base),
..call
}
}
#[rstest]
#[tokio::test]
async fn what_before_send_returns_is_what_the_provider_receives(call: MessagesCall) {
let upstream = upstream([message_response()]).await;
let host = RecordingHost::new(
authenticated(call, upstream.uri()),
Box::new(|wire| {
let mut body = wire.body;
body["system"] = json!("added by the host");
Ok(WireRequest {
headers: wire
.headers
.into_iter()
.chain([("x-host".to_string(), "seen".to_string())])
.collect(),
body,
..wire
})
}),
);
run_through(&host).await.expect("messages call succeeds");
let request = only_request(&upstream).await;
assert_eq!(request.json()["system"], "added by the host");
assert_eq!(request.header("x-host"), Some("seen"));
assert_eq!(request.header("x-api-key"), Some("sk-ant"));
}
#[rstest]
#[tokio::test]
async fn a_before_send_failure_never_sends(call: MessagesCall) {
let upstream = upstream([message_response()]).await;
let host = RecordingHost::new(
authenticated(call, upstream.uri()),
Box::new(|_| Err(Error::InvalidRequest("vetoed by the host".into()))),
);
let error = run_through(&host)
.await
.err()
.expect("the host failure fails the call");
assert_eq!(error, Error::InvalidRequest("vetoed by the host".into()));
assert!(received(&upstream).await.is_empty());
assert!(host.raw_responses().is_empty());
}
#[rstest]
#[tokio::test]
async fn the_raw_upstream_text_is_emitted_once_for_a_message(call: MessagesCall) {
let raw = message_body();
let upstream = upstream([json_response(raw.clone())]).await;
let host = RecordingHost::passthrough(authenticated(call, upstream.uri()));
let output = run_through(&host).await.expect("messages call succeeds");
assert!(matches!(output, MessagesOutput::Message(_)));
let [emitted] = <[String; 1]>::try_from(host.raw_responses())
.unwrap_or_else(|raws| panic!("expected one raw response, got {}", raws.len()));
assert_eq!(serde_json::from_str::<Value>(&emitted).unwrap(), raw);
}
#[rstest]
#[case::upstream_error(ResponseTemplate::new(500).set_body_string("boom"))]
#[case::stream(ResponseTemplate::new(200).set_body_raw("event: message_stop\ndata: {}\n\n", "text/event-stream"))]
#[tokio::test]
async fn no_raw_response_is_emitted_for_a_stream_or_a_failure(
call: MessagesCall,
#[case] response: ResponseTemplate,
) {
let upstream = upstream([response]).await;
let mut body = call.body.clone();
body.insert("stream".into(), json!(true));
let host =
RecordingHost::passthrough(authenticated(MessagesCall { body, ..call }, upstream.uri()));
let _ = run_through(&host).await;
assert_eq!(received(&upstream).await.len(), 1);
assert!(host.raw_responses().is_empty());
}
/// Python logs `optional_params` as what it is about to send, so a dropped param must
/// not resurface in callbacks.
#[rstest]
#[tokio::test]
async fn the_request_context_carries_the_shaped_params_without_model_or_messages(
call: MessagesCall,
) {
let upstream = upstream([message_response()]).await;
let body: Map<String, Value> = call
.body
.clone()
.into_iter()
.chain([("temperature".to_string(), json!(0.2))])
.collect();
let host = RecordingHost::passthrough(authenticated(
MessagesCall {
body,
shaping: MessagesShaping {
capabilities: AnthropicModelCapabilities {
supports_sampling_params: false,
..AnthropicModelCapabilities::default()
},
drop_params: true,
..MessagesShaping::default()
},
..call
},
upstream.uri(),
));
run_through(&host).await.expect("messages call succeeds");
let [optional_params] = <[Value; 1]>::try_from(host.optional_params.into_inner().unwrap())
.unwrap_or_else(|seen| panic!("before_send runs once, saw {}", seen.len()));
assert_eq!(optional_params, json!({"max_tokens": 16}));
}

View file

@ -14,6 +14,7 @@ use wiremock::ResponseTemplate;
mod support;
use support::*;
mod host;
mod request;
mod response;
mod secrets;

View file

@ -1,3 +1,7 @@
use litellm_llms::anthropic::common_utils::{
ANTHROPIC_ADVISOR_TOOL_TYPE, ANTHROPIC_OAUTH_BETA_HEADER, AnthropicModelCapabilities,
SupportedEffortTiers, beta,
};
use litellm_types::utils::{ProviderSpecificHeader, ProviderSpecificHeaders};
use rstest::rstest;
@ -132,8 +136,8 @@ async fn each_provider_posts_to_its_messages_endpoint(
assert_eq!(request.method.as_str(), "POST");
assert_eq!(request.url.path(), path);
assert_eq!(request.json()["model"], MODEL);
assert_eq!(request.header("anthropic-version"), Some("2023-06-01"));
assert_eq!(request.header("content-type"), Some("application/json"));
assert_eq!(request.header_values("anthropic-version"), ["2023-06-01"]);
assert_eq!(request.header_values("content-type"), ["application/json"]);
}
#[rstest]
@ -249,21 +253,423 @@ async fn additional_drop_params_remove_fields_before_sending(call: MessagesCall)
assert_eq!(sent["top_k"], 3);
}
fn with_fields(call: MessagesCall, fields: Value) -> MessagesCall {
let body: Map<String, Value> = call.body.into_iter().chain(object(fields)).collect();
MessagesCall { body, ..call }
}
fn sent_betas(request: &wiremock::Request) -> Vec<String> {
let [header] = <[&str; 1]>::try_from(request.header_values("anthropic-beta"))
.unwrap_or_else(|values| panic!("expected one anthropic-beta header, got {values:?}"));
header
.split(',')
.map(str::trim)
.map(str::to_string)
.collect()
}
#[rstest]
#[case::anthropic_streams(MODEL, Some("anthropic"), true, true)]
#[case::anthropic_prefix_streams("anthropic/claude-sonnet-4-5", None, true, true)]
#[case::azure_without_stream(MODEL, Some("azure_ai"), false, true)]
#[case::azure_stream(MODEL, Some("azure_ai"), true, false)]
#[case::other_provider(MODEL, Some("openai"), false, false)]
#[case::unresolvable_model("no-such-model", None, false, false)]
fn supports_matches_what_the_route_can_serve(
#[case] model: &str,
#[case] provider: Option<&str>,
#[case] stream: bool,
#[case] supported: bool,
#[case::structured_output(json!({"output_format": {"type": "json_schema"}}), &[beta::STRUCTURED_OUTPUT])]
#[case::fast_mode(json!({"speed": "fast"}), &[beta::FAST_MODE_2026_02_01])]
#[case::compaction(json!({"compaction": {"enabled": true}}), &[beta::COMPACT_2026_09_04])]
#[case::context_management_edits(
json!({"context_management": {"edits": [{"type": "clear_tool_uses_20250919"}]}}),
&[beta::CONTEXT_MANAGEMENT_2025_06_27]
)]
#[case::per_message_output_config(
json!({"messages": [{"role": "user", "content": "hi", "output_config": {"effort": "low"}}]}),
&[beta::PER_TURN_CONTROL_2026_07_01]
)]
#[case::advisor_tool(
json!({"tools": [{"type": ANTHROPIC_ADVISOR_TOOL_TYPE, "name": "advisor", "model": MODEL}]}),
&[beta::ADVISOR_TOOL_2026_03_01]
)]
#[case::several_features_at_once(
json!({"speed": "fast", "output_format": {"type": "json_schema"}}),
&[beta::STRUCTURED_OUTPUT, beta::FAST_MODE_2026_02_01]
)]
#[tokio::test]
async fn feature_betas_join_the_callers_betas_in_one_sorted_header(
call: MessagesCall,
#[case] fields: Value,
#[case] features: &[&str],
) {
let upstream = upstream([message_response()]).await;
let capabilities = AnthropicModelCapabilities {
supports_speed: true,
..AnthropicModelCapabilities::default()
};
run_message(with_fields(
MessagesCall {
api_key: Some("sk".into()),
api_base: Some(upstream.uri()),
extra_headers: headers([("Anthropic-Beta", "caller-beta-2025-01-01")]),
shaping: MessagesShaping {
capabilities,
..MessagesShaping::default()
},
..call
},
fields,
))
.await;
let sent = sent_betas(&only_request(&upstream).await);
let mut expected: Vec<String> = features
.iter()
.map(|feature| feature.to_string())
.chain(["caller-beta-2025-01-01".to_string()])
.collect();
expected.sort();
assert_eq!(sent, expected);
}
#[rstest]
#[tokio::test]
async fn an_oauth_key_sends_the_browser_access_header_and_the_oauth_beta(call: MessagesCall) {
let upstream = upstream([message_response()]).await;
run_message(MessagesCall {
api_key: Some("sk-ant-oat01-token".into()),
api_base: Some(upstream.uri()),
..call
})
.await;
let request = only_request(&upstream).await;
assert_eq!(
litellm_core::messages::route::supports(model, provider, stream),
supported
request.header("anthropic-dangerous-direct-browser-access"),
Some("true")
);
assert_eq!(sent_betas(&request), [ANTHROPIC_OAUTH_BETA_HEADER]);
assert_eq!(request.header("x-api-key"), None);
}
#[rstest]
#[case::anthropic("anthropic")]
#[case::azure_ai("azure_ai")]
#[tokio::test]
async fn caller_protocol_headers_win_over_the_defaults(call: MessagesCall, #[case] provider: &str) {
let upstream = upstream([message_response()]).await;
run_message(MessagesCall {
custom_llm_provider: Some(provider.into()),
api_key: Some("sk".into()),
api_base: Some(upstream.uri()),
extra_headers: headers([
("Anthropic-Version", "2024-01-01"),
("Content-Type", "application/json; charset=utf-8"),
]),
..call
})
.await;
let request = only_request(&upstream).await;
assert_eq!(request.header_values("anthropic-version"), ["2024-01-01"]);
assert_eq!(
request.header_values("content-type"),
["application/json; charset=utf-8"]
);
}
fn sampling_removed() -> AnthropicModelCapabilities {
AnthropicModelCapabilities {
supports_sampling_params: false,
..AnthropicModelCapabilities::default()
}
}
#[rstest]
#[case::sampling_params(sampling_removed(), json!({"temperature": 0.2, "top_p": 0.9, "top_k": 5}), &["temperature", "top_p", "top_k"], "temperature=0.2")]
#[case::speed(AnthropicModelCapabilities::default(), json!({"speed": "fast"}), &["speed"], "speed='fast'")]
#[tokio::test]
async fn unsupported_params_are_dropped_under_drop_params_and_rejected_without_it(
call: MessagesCall,
#[case] capabilities: AnthropicModelCapabilities,
#[case] fields: Value,
#[case] dropped: &[&str],
#[case] rejected_as: &str,
) {
let upstream = upstream([message_response(), message_response()]).await;
let shaped = |drop_params: bool| {
with_fields(
MessagesCall {
api_key: Some("sk".into()),
api_base: Some(upstream.uri()),
shaping: MessagesShaping {
capabilities: capabilities.clone(),
drop_params,
..MessagesShaping::default()
},
body: call.body.clone(),
custom_llm_provider: call.custom_llm_provider.clone(),
extra_headers: None,
provider_specific_header: None,
model: call.model.clone(),
timeout: call.timeout,
},
fields.clone(),
)
};
let error = run(shaped(false))
.await
.err()
.expect("an unsupported param is rejected without drop_params");
assert!(
matches!(&error, Error::InvalidRequest(message) if message.contains(rejected_as)),
"{error:?}"
);
assert!(received(&upstream).await.is_empty());
run_message(shaped(true)).await;
let sent = only_request(&upstream).await.json();
for name in dropped {
assert_eq!(sent.get(*name), None, "{name} must be dropped");
}
assert_eq!(sent["max_tokens"], 16);
}
#[rstest]
#[case::adaptive_thinking(json!({"type": "adaptive"}), json!({"type": "adaptive", "display": "summarized"}))]
#[case::disabled_thinking(json!({"type": "disabled"}), json!({"type": "disabled"}))]
#[tokio::test]
async fn reasoning_auto_summary_marks_active_thinking_on_the_wire(
call: MessagesCall,
#[case] thinking: Value,
#[case] expected: Value,
) {
let upstream = upstream([message_response()]).await;
run_message(with_fields(
MessagesCall {
api_key: Some("sk".into()),
api_base: Some(upstream.uri()),
shaping: MessagesShaping {
capabilities: AnthropicModelCapabilities {
supports_reasoning: true,
supports_adaptive_thinking: true,
..AnthropicModelCapabilities::default()
},
reasoning_auto_summary: true,
..MessagesShaping::default()
},
..call
},
json!({"thinking": thinking}),
))
.await;
assert_eq!(only_request(&upstream).await.json()["thinking"], expected);
}
#[rstest]
#[case::reasoning_effort_on_an_adaptive_model(
AnthropicModelCapabilities {
supports_reasoning: true,
supports_adaptive_thinking: true,
supports_output_config: true,
effort_tiers: SupportedEffortTiers { high: true, ..SupportedEffortTiers::default() },
..AnthropicModelCapabilities::default()
},
json!({"reasoning_effort": "high"}),
json!({"thinking": {"type": "adaptive", "display": "summarized"}, "output_config": {"effort": "high"}})
)]
#[case::reasoning_effort_on_a_legacy_model_caps_the_budget_below_max_tokens(
AnthropicModelCapabilities {
supports_reasoning: true,
..AnthropicModelCapabilities::default()
},
json!({"reasoning_effort": "high"}),
json!({"thinking": {"type": "enabled", "budget_tokens": 2999}})
)]
#[case::adaptive_payload_on_a_legacy_model_becomes_a_capped_budget(
AnthropicModelCapabilities {
supports_reasoning: true,
..AnthropicModelCapabilities::default()
},
json!({"thinking": {"type": "adaptive"}, "output_config": {"effort": "high"}, "temperature": 0}),
json!({"thinking": {"type": "enabled", "budget_tokens": 2999}})
)]
#[case::adaptive_payload_on_a_model_without_reasoning_is_dropped(
AnthropicModelCapabilities::default(),
json!({"thinking": {"type": "adaptive"}, "output_config": {"effort": "high"}}),
json!({})
)]
#[tokio::test]
async fn reasoning_is_translated_by_the_model_capabilities(
call: MessagesCall,
#[case] capabilities: AnthropicModelCapabilities,
#[case] fields: Value,
#[case] expected: Value,
) {
let upstream = upstream([message_response()]).await;
run_message(with_fields(
MessagesCall {
api_key: Some("sk".into()),
api_base: Some(upstream.uri()),
shaping: MessagesShaping {
capabilities,
..MessagesShaping::default()
},
..call
},
[("max_tokens".to_string(), json!(3000))]
.into_iter()
.chain(object(fields))
.collect(),
))
.await;
let sent = only_request(&upstream).await.json();
assert_eq!(sent.get("reasoning_effort"), None);
assert_eq!(sent.get("temperature"), None);
let reasoning: Map<String, Value> = ["thinking", "output_config"]
.into_iter()
.filter_map(|name| Some((name.to_string(), sent.get(name)?.clone())))
.collect();
assert_eq!(Value::Object(reasoning), expected);
}
#[rstest]
#[case::empty_text_blocks(
json!([{"role": "assistant", "content": [{"type": "text", "text": " "}, {"type": "text", "text": "kept"}]}]),
json!([{"role": "assistant", "content": [{"type": "text", "text": "kept"}]}])
)]
#[case::provider_specific_fields(
json!([{"role": "assistant", "content": [{"type": "text", "text": "kept", "provider_specific_fields": {"x": 1}}]}]),
json!([{"role": "assistant", "content": [{"type": "text", "text": "kept"}]}])
)]
#[case::unencrypted_web_search_results_become_text(
json!([{"role": "assistant", "content": [{
"type": "web_search_tool_result",
"tool_use_id": "srvtoolu_1",
"content": [{"type": "web_search_result", "title": "T", "url": "https://e.x", "page_age": null}]
}]}]),
json!([{"role": "assistant", "content": [{"type": "text", "text": "Web search results:\n\nTitle: T\nURL: https://e.x"}]}])
)]
#[tokio::test]
async fn replayed_history_is_cleaned_before_sending(
call: MessagesCall,
#[case] history: Value,
#[case] expected: Value,
) {
let upstream = upstream([message_response()]).await;
run_message(with_fields(
MessagesCall {
api_key: Some("sk".into()),
api_base: Some(upstream.uri()),
..call
},
json!({"messages": history}),
))
.await;
assert_eq!(only_request(&upstream).await.json()["messages"], expected);
}
#[rstest]
#[tokio::test]
async fn metadata_is_reduced_to_the_user_id(call: MessagesCall) {
let upstream = upstream([message_response()]).await;
run_message(with_fields(
MessagesCall {
api_key: Some("sk".into()),
api_base: Some(upstream.uri()),
..call
},
json!({"metadata": {"user_id": "u-1", "trace_id": "internal", "tags": ["a"]}}),
))
.await;
assert_eq!(
only_request(&upstream).await.json()["metadata"],
json!({"user_id": "u-1"})
);
}
#[rstest]
#[case::numeric_user_id(json!({"metadata": {"user_id": 7}}))]
#[case::missing_max_tokens(json!({"max_tokens": null}))]
#[tokio::test]
async fn an_invalid_request_fails_before_sending(call: MessagesCall, #[case] fields: Value) {
let upstream = upstream([message_response()]).await;
let error = run(with_fields(
MessagesCall {
api_key: Some("sk".into()),
api_base: Some(upstream.uri()),
..call
},
fields,
))
.await
.err()
.expect("the request is rejected");
assert!(error.is_request(), "{error:?}");
assert!(received(&upstream).await.is_empty());
}
#[rstest]
#[tokio::test]
async fn azure_folds_system_role_messages_into_the_system_prompt(call: MessagesCall) {
let upstream = upstream([message_response()]).await;
run_message(with_fields(
MessagesCall {
custom_llm_provider: Some("azure_ai".into()),
api_key: Some("sk-azure".into()),
api_base: Some(upstream.uri()),
..call
},
json!({
"system": "top level",
"messages": [
{"role": "system", "content": "from a message"},
{"role": "user", "content": "hi"}
]
}),
))
.await;
let sent = only_request(&upstream).await.json();
assert_eq!(
sent["system"],
json!([
{"type": "text", "text": "top level"},
{"type": "text", "text": "from a message"}
])
);
assert_eq!(sent["messages"], json!([{"role": "user", "content": "hi"}]));
}
#[rstest]
#[case::bare_model(MODEL, MODEL)]
#[case::one_prefix("anthropic/claude-sonnet-4-5", MODEL)]
#[case::doubled_prefix_loses_one_segment(
"anthropic/anthropic/claude-sonnet-4-5",
"anthropic/claude-sonnet-4-5"
)]
#[tokio::test]
async fn the_provider_prefix_is_stripped_exactly_once(
call: MessagesCall,
#[case] model: &str,
#[case] sent_model: &str,
) {
let upstream = upstream([message_response()]).await;
run_message(MessagesCall {
model: model.into(),
api_key: Some("sk".into()),
api_base: Some(upstream.uri()),
..call
})
.await;
assert_eq!(only_request(&upstream).await.json()["model"], sent_model);
}

View file

@ -5,11 +5,14 @@ use rstest::rstest;
use super::*;
#[rstest]
#[case::anthropic("anthropic")]
#[case::azure_ai("azure_ai")]
#[tokio::test]
async fn the_provider_message_is_returned(call: MessagesCall) {
async fn the_provider_message_is_returned(call: MessagesCall, #[case] provider: &str) {
let upstream = upstream([message_response()]).await;
let message = run_message(MessagesCall {
custom_llm_provider: Some(provider.into()),
api_key: Some("sk".into()),
api_base: Some(upstream.uri()),
..call
@ -21,6 +24,88 @@ async fn the_provider_message_is_returned(call: MessagesCall) {
assert_eq!(message.stop_reason.as_deref(), Some("end_turn"));
}
/// A refusal and fields the route does not model come back exactly as the provider sent
/// them, since the Python side returns the raw message and the router decides what to do.
#[rstest]
#[tokio::test]
async fn the_message_passes_through_losslessly(call: MessagesCall) {
let upstream_body = json!({
"id": "msg_2",
"type": "message",
"role": "assistant",
"model": MODEL,
"content": [
{"type": "server_tool_use", "id": "srvtoolu_1", "name": "web_search", "input": {"query": "q"}},
{"type": "text", "text": "no", "citations": [{"type": "web_search_result_location", "url": "https://e.x"}]}
],
"stop_reason": "refusal",
"stop_sequence": null,
"stop_details": {"type": "safeguard", "safeguard_types": ["dangerous_tool_use"]},
"container": {"id": "container_1", "expires_at": "2026-01-01T00:00:00Z"},
"context_management": {"applied_edits": []},
"usage": {"input_tokens": 1, "output_tokens": 2, "server_tool_use": {"web_search_requests": 1}},
"unknown_future_field": {"nested": true}
});
let upstream = upstream([json_response(upstream_body.clone())]).await;
let message = run_message(MessagesCall {
api_key: Some("sk".into()),
api_base: Some(upstream.uri()),
..call
})
.await;
assert_eq!(message.stop_reason.as_deref(), Some("refusal"));
assert_eq!(serde_json::to_value(&message).unwrap(), upstream_body);
}
#[rstest]
#[tokio::test]
async fn a_json_error_envelope_is_kept_verbatim(call: MessagesCall) {
let envelope =
json!({"type": "error", "error": {"type": "invalid_request_error", "message": "bad"}});
let upstream = upstream([status_response(400, envelope.clone())]).await;
let error = run(MessagesCall {
api_key: Some("sk".into()),
api_base: Some(upstream.uri()),
..call
})
.await
.err()
.expect("upstream error propagates");
let Error::Transport(TransportError::Http { status, body }) = error else {
panic!("{error:?}");
};
assert_eq!(status, 400);
assert_eq!(serde_json::from_str::<Value>(&body).unwrap(), envelope);
}
#[rstest]
#[tokio::test]
async fn a_long_error_body_is_truncated_at_the_documented_cap(call: MessagesCall) {
let long = "x".repeat(600);
let upstream = upstream([ResponseTemplate::new(500).set_body_string(long.clone())]).await;
let error = run(MessagesCall {
api_key: Some("sk".into()),
api_base: Some(upstream.uri()),
..call
})
.await
.err()
.expect("upstream error propagates");
assert_eq!(
error,
Error::Transport(TransportError::Http {
status: 500,
body: format!("{}... (truncated)", &long[..256])
})
);
}
#[rstest]
#[case::bad_request(400)]
#[case::unauthorized(401)]

View file

@ -1,23 +1,30 @@
use litellm_llms::{
anthropic::experimental_pass_through::messages::transformation::ANTHROPIC_MESSAGES_CONFIG,
azure_ai::anthropic::messages_transformation::AZURE_ANTHROPIC_MESSAGES_CONFIG,
base_llm::anthropic_messages::transformation::BaseAnthropicMessagesConfig,
};
use rstest::rstest;
use super::*;
#[rstest]
#[case::anthropic("anthropic", &ANTHROPIC_MESSAGES_CONFIG, "ANTHROPIC_API_KEY", "ANTHROPIC_BASE_URL", "/v1/messages")]
#[case::azure_ai("azure_ai", &AZURE_ANTHROPIC_MESSAGES_CONFIG, "AZURE_API_KEY", "AZURE_API_BASE", "/anthropic/v1/messages")]
#[case::anthropic(
"anthropic",
"ANTHROPIC_API_KEY",
"ANTHROPIC_BASE_URL",
"/v1/messages",
&["ANTHROPIC_API_KEY", "ANTHROPIC_AUTH_TOKEN", "ANTHROPIC_API_BASE", "ANTHROPIC_BASE_URL"]
)]
#[case::azure_ai(
"azure_ai",
"AZURE_API_KEY",
"AZURE_API_BASE",
"/anthropic/v1/messages",
&["AZURE_API_KEY", "AZURE_API_BASE"]
)]
#[tokio::test]
async fn the_credential_and_base_come_from_the_secret_source(
call: MessagesCall,
#[case] provider: &str,
#[case] config: &dyn BaseAnthropicMessagesConfig,
#[case] key_name: &str,
#[case] base_name: &str,
#[case] path: &str,
#[case] looked_up: &[&str],
) {
let upstream = upstream([message_response()]).await;
let base = upstream.uri();
@ -40,7 +47,7 @@ async fn the_credential_and_base_come_from_the_secret_source(
let request = only_request(&upstream).await;
assert_eq!(request.url.path(), path);
assert_eq!(request.header("x-api-key"), Some("sk-from-manager"));
assert_eq!(secrets.requested(), config.secret_names());
assert_eq!(secrets.requested(), looked_up);
}
#[rstest]
@ -92,3 +99,102 @@ async fn a_secret_manager_failure_fails_the_call_before_sending(call: MessagesCa
);
assert!(received(&upstream).await.is_empty());
}
#[derive(Clone, Copy)]
enum Base {
Upstream,
Unreachable,
Blank,
Absent,
}
fn base_value(base: Base, upstream: &str) -> Option<String> {
match base {
Base::Upstream => Some(upstream.to_string()),
Base::Unreachable => Some(UNREACHABLE_BASE.to_string()),
Base::Blank => Some(" ".to_string()),
Base::Absent => None,
}
}
#[rstest]
#[case::api_base_beats_base_url(Base::Upstream, Base::Unreachable)]
#[case::blank_api_base_falls_through_to_base_url(Base::Blank, Base::Upstream)]
#[case::base_url_alone(Base::Absent, Base::Upstream)]
#[tokio::test]
async fn the_anthropic_base_env_precedence_picks_the_upstream(
call: MessagesCall,
#[case] api_base: Base,
#[case] base_url: Base,
) {
let upstream = upstream([message_response()]).await;
let uri = upstream.uri();
let values: Vec<(&str, &str)> = [
("ANTHROPIC_API_KEY", Some("sk-env".to_string())),
("ANTHROPIC_API_BASE", base_value(api_base, &uri)),
("ANTHROPIC_BASE_URL", base_value(base_url, &uri)),
]
.iter()
.filter_map(|(name, value)| Some((*name, value.as_deref()?)))
.map(|(name, value)| (name, Box::leak(value.to_string().into_boxed_str()) as &str))
.collect();
run_with(Arc::new(RecordingSecrets::new(values)), call)
.await
.expect("messages call reaches the upstream the precedence picks");
assert_eq!(only_request(&upstream).await.url.path(), "/v1/messages");
}
#[rstest]
#[case::auth_token_alone(
&[("ANTHROPIC_AUTH_TOKEN", "tok")],
("authorization", "Bearer tok"),
"x-api-key"
)]
#[case::api_key_beats_the_auth_token(
&[("ANTHROPIC_API_KEY", "sk-env"), ("ANTHROPIC_AUTH_TOKEN", "tok")],
("x-api-key", "sk-env"),
"authorization"
)]
#[tokio::test]
async fn the_auth_token_env_is_a_bearer_only_without_a_key(
call: MessagesCall,
#[case] values: &[(&str, &str)],
#[case] expected: (&str, &str),
#[case] absent: &str,
) {
let upstream = upstream([message_response()]).await;
run_with(
Arc::new(RecordingSecrets::new(values.iter().copied())),
MessagesCall {
api_base: Some(upstream.uri()),
..call
},
)
.await
.expect("messages call succeeds");
let request = only_request(&upstream).await;
let (name, value) = expected;
assert_eq!(request.header_values(name), [value]);
assert_eq!(request.header(absent), None);
}
#[rstest]
#[tokio::test]
async fn azure_without_a_base_anywhere_fails_before_sending(call: MessagesCall) {
let error = run_with(
Arc::new(RecordingSecrets::new([("AZURE_API_KEY", "sk-azure")])),
MessagesCall {
custom_llm_provider: Some("azure_ai".into()),
..call
},
)
.await
.err()
.expect("azure needs a base");
assert_eq!(error, Error::Auth(litellm_auth::Error::MissingAzureApiBase));
}

View file

@ -1,16 +1,25 @@
use std::{convert::Infallible, sync::Mutex};
use bytes::Bytes;
use litellm_core::messages::route::Messages;
use litellm_core::messages::route::{Messages, MessagesStreamHead};
use litellm_host::host::{Demand, Host};
use rstest::rstest;
use tokio::{
io::{AsyncReadExt, AsyncWriteExt},
net::TcpListener,
};
use super::*;
const UPSTREAM_HEADERS: [(&str, &str); 2] = [
("request-id", "req_upstream_123"),
("anthropic-ratelimit-requests-remaining", "41"),
];
const SSE_BODY: &str = "event: message_start\ndata: {\"type\":\"message_start\"}\n\nevent: message_stop\ndata: {\"type\":\"message_stop\"}\n\n";
enum Seen {
Open,
Open(Vec<(String, String)>),
Deliver(Bytes),
}
@ -50,8 +59,8 @@ impl Host<Messages> for RecordingStreamHost {
match op {}
}
async fn open(&self, (): ()) -> Result<Demand, Error> {
Ok(self.record(Seen::Open))
async fn open(&self, head: MessagesStreamHead) -> Result<Demand, Error> {
Ok(self.record(Seen::Open(head.headers)))
}
async fn deliver(&self, chunk: Bytes) -> Result<Demand, Error> {
@ -71,7 +80,10 @@ fn streaming(call: MessagesCall, api_base: String) -> MessagesCall {
}
fn sse_response() -> ResponseTemplate {
ResponseTemplate::new(200).set_body_raw(SSE_BODY, "text/event-stream")
UPSTREAM_HEADERS.iter().fold(
ResponseTemplate::new(200).set_body_raw(SSE_BODY, "text/event-stream"),
|response, (name, value)| response.insert_header(*name, *value),
)
}
async fn stream_through(host: &RecordingStreamHost) -> Result<MessagesOutput, Error> {
@ -80,7 +92,7 @@ async fn stream_through(host: &RecordingStreamHost) -> Result<MessagesOutput, Er
#[rstest]
#[tokio::test]
async fn the_stream_opens_once_before_relaying_the_upstream_body(call: MessagesCall) {
async fn upstream_headers_are_on_the_stream_head_before_the_first_chunk(call: MessagesCall) {
let upstream = upstream([sse_response()]).await;
let host = RecordingStreamHost::new(streaming(call, upstream.uri()), usize::MAX);
@ -88,14 +100,24 @@ async fn the_stream_opens_once_before_relaying_the_upstream_body(call: MessagesC
assert!(matches!(outcome, MessagesOutput::Streamed));
let seen = host.seen.into_inner().unwrap();
let [Seen::Open, chunks @ ..] = seen.as_slice() else {
let [Seen::Open(headers), chunks @ ..] = seen.as_slice() else {
panic!("the stream opens before any chunk is delivered");
};
let surfaced: Vec<(&str, &str)> = headers
.iter()
.filter(|(name, _)| {
UPSTREAM_HEADERS
.iter()
.any(|(upstream, _)| upstream == name)
})
.map(|(name, value)| (name.as_str(), value.as_str()))
.collect();
assert_eq!(surfaced, UPSTREAM_HEADERS);
let delivered: Vec<u8> = chunks
.iter()
.flat_map(|step| match step {
Seen::Deliver(chunk) => chunk.to_vec(),
Seen::Open => panic!("the stream opens exactly once"),
Seen::Open(_) => panic!("the stream opens exactly once"),
})
.collect();
assert_eq!(delivered, SSE_BODY.as_bytes());
@ -118,9 +140,18 @@ async fn a_detached_caller_receives_nothing_more(call: MessagesCall, #[case] det
}
#[rstest]
#[case::text_body(ResponseTemplate::new(429).set_body_string("slow down"), "slow down")]
#[case::json_envelope(
status_response(429, json!({"type": "error", "error": {"type": "rate_limit_error", "message": "slow down"}})),
r#"{"type":"error","error":{"type":"rate_limit_error","message":"slow down"}}"#
)]
#[tokio::test]
async fn an_upstream_error_fails_the_call_without_opening_the_stream(call: MessagesCall) {
let upstream = upstream([ResponseTemplate::new(429).set_body_string("slow down")]).await;
async fn an_upstream_error_fails_the_call_without_opening_the_stream(
call: MessagesCall,
#[case] response: ResponseTemplate,
#[case] body: &str,
) {
let upstream = upstream([response]).await;
let host = RecordingStreamHost::new(streaming(call, upstream.uri()), usize::MAX);
let error = stream_through(&host)
@ -128,16 +159,88 @@ async fn an_upstream_error_fails_the_call_without_opening_the_stream(call: Messa
.err()
.expect("upstream error propagates");
assert!(
matches!(
error,
Error::Transport(litellm_http::transport::Error::Http { status: 429, .. })
),
"{error:?}"
assert_eq!(
error,
Error::Transport(litellm_http::transport::Error::Http {
status: 429,
body: body.into()
})
);
assert!(host.seen.into_inner().unwrap().is_empty());
}
/// The native route relays bytes as they are. Python's synthetic `api_error` for a stream
/// that never reaches `message_stop` lives in its SSE wrapper, above this route.
#[rstest]
#[tokio::test]
async fn a_stream_that_ends_without_message_stop_is_relayed_as_is(call: MessagesCall) {
const INCOMPLETE: &str = "event: message_start\ndata: {\"type\":\"message_start\"}\n\n";
let upstream =
upstream([ResponseTemplate::new(200).set_body_raw(INCOMPLETE, "text/event-stream")]).await;
let host = RecordingStreamHost::new(streaming(call, upstream.uri()), usize::MAX);
stream_through(&host).await.expect("streamed call succeeds");
let delivered: Vec<u8> = host
.seen
.into_inner()
.unwrap()
.iter()
.flat_map(|step| match step {
Seen::Deliver(chunk) => chunk.to_vec(),
Seen::Open(_) => Vec::new(),
})
.collect();
assert_eq!(delivered, INCOMPLETE.as_bytes());
}
/// Serves one SSE chunk and then holds the connection open without ever finishing.
async fn stalling_upstream() -> String {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let base = format!("http://{}", listener.local_addr().unwrap());
tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.unwrap();
let mut request = vec![0; 4096];
let _ = socket.read(&mut request).await;
socket
.write_all(
b"HTTP/1.1 200 OK\r\ncontent-type: text/event-stream\r\ntransfer-encoding: chunked\r\n\r\n\
1f\r\nevent: message_start\ndata: {}\n\n\r\n",
)
.await
.unwrap();
std::future::pending::<()>().await;
});
base
}
#[rstest]
#[tokio::test]
async fn the_timeout_covers_a_stalled_stream_body(call: MessagesCall) {
let base = stalling_upstream().await;
let host = RecordingStreamHost::new(
MessagesCall {
timeout: Some(Duration::from_millis(300)),
..streaming(call, base)
},
usize::MAX,
);
let error = tokio::time::timeout(Duration::from_secs(5), stream_through(&host))
.await
.expect("the stalled stream gives up within the timeout")
.err()
.expect("a stalled body fails the call");
assert!(matches!(error, Error::Transport(_)), "{error:?}");
let seen = host.seen.into_inner().unwrap();
assert!(
matches!(seen.as_slice(), [Seen::Open(_), Seen::Deliver(chunk)] if chunk.as_ref() == b"event: message_start\ndata: {}\n\n"),
"the chunk before the stall reached the caller, saw {} ops",
seen.len()
);
}
#[rstest]
#[tokio::test]
async fn streaming_is_refused_for_providers_that_cannot_stream(call: MessagesCall) {

View file

@ -134,6 +134,13 @@ pub trait ProtocolHost: Send + Sync {
response: <Self::Protocol as Protocol>::Response,
) -> PyResult<Py<PyAny>>;
/// What the stream carries at hand-off, as the caller's stream receives it.
fn head(
&mut self,
py: Python<'_>,
head: <Self::Protocol as Protocol>::StreamHead,
) -> PyResult<Py<PyAny>>;
/// One streamed chunk as the caller receives it.
fn chunk(
&mut self,

View file

@ -134,10 +134,10 @@ where
}
match driver.resume(None)? {
ExecutionStep::Return(value) => Ok(value),
ExecutionStep::Open => py
ExecutionStep::Open(head) => py
.import("litellm.rust_bridge.lifecycle")?
.getattr("SyncStream")?
.call1((Py::new(py, Execution::suspended(driver))?,))
.call1((Py::new(py, Execution::suspended(driver))?, head))
.map(Bound::unbind),
ExecutionStep::Await(_) | ExecutionStep::Yield(_) => {
Err(PyRuntimeError::new_err("sync call suspended"))
@ -312,7 +312,7 @@ where
Ok(_) => return Err(missing_state()),
Err(error) => Err(error),
},
HostOp::Open(_, reply) => return self.opened(py, reply).map(Next::Return),
HostOp::Open(head, reply) => return self.opened(py, head, reply).map(Next::Return),
HostOp::Deliver(chunk, reply) => {
return self.delivered(py, chunk, reply).map(Next::Return);
}
@ -340,12 +340,21 @@ where
}
}
fn opened(&mut self, py: Python<'_>, reply: Reply<Demand>) -> PyResult<ExecutionStep> {
fn opened(
&mut self,
py: Python<'_>,
head: <ProtocolOf<H> as Protocol>::StreamHead,
reply: Reply<Demand>,
) -> PyResult<ExecutionStep> {
self.stage = Stage::Streaming;
let head = match self.host.head(py, head) {
Ok(head) => head,
Err(error) => return self.interrupt(py, error),
};
match self.adapter.opened(py) {
Ok(()) => {
self.pending = Some(Pending::Consumer(reply));
Ok(ExecutionStep::Open)
Ok(ExecutionStep::Open(head))
}
Err(error) => self.interrupt(py, error),
}
@ -699,6 +708,10 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri
.map(|answer| reply.send(answer))
}
fn head(&mut self, _: Python<'_>, head: std::convert::Infallible) -> PyResult<Py<PyAny>> {
match head {}
}
fn chunk(&mut self, _: Python<'_>, chunk: std::convert::Infallible) -> PyResult<Py<PyAny>> {
match chunk {}
}
@ -945,6 +958,163 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri
});
}
struct Streaming;
impl Protocol for Streaming {
type Response = ();
type Error = Error;
type Projection = ();
type Op = std::convert::Infallible;
type Chunk = &'static str;
type StreamHead = Vec<(&'static str, &'static str)>;
}
struct StreamingHost;
impl ProtocolHost for StreamingHost {
type Protocol = Streaming;
type Failure = Classified;
fn project(
&mut self,
_: Python<'_>,
_: &Bound<'_, PyDict>,
) -> Result<(), InvokeError<Error>> {
Ok(())
}
fn invoke(
&mut self,
_: Python<'_>,
op: std::convert::Infallible,
) -> Result<(), InvokeError<Error>> {
match op {}
}
fn head(
&mut self,
py: Python<'_>,
head: Vec<(&'static str, &'static str)>,
) -> PyResult<Py<PyAny>> {
let headers = PyDict::new(py);
for (name, value) in head {
headers.set_item(name, value)?;
}
let hidden = PyDict::new(py);
hidden.set_item("additional_headers", headers)?;
Ok(hidden.into_any().unbind())
}
fn chunk(&mut self, py: Python<'_>, chunk: &'static str) -> PyResult<Py<PyAny>> {
Ok(pyo3::types::PyString::new(py, chunk).into_any().unbind())
}
fn complete(&mut self, py: Python<'_>, (): ()) -> PyResult<Py<PyAny>> {
Ok(py.None())
}
fn classify(&self, _: Python<'_>, error: Error) -> PyResult<Classified> {
Ok(Classified(error.0))
}
fn host_error(error: &PyErr) -> Error {
Error(error.to_string())
}
fn close(&mut self, _: Python<'_>) {}
fn traverse(&self, _: &PyVisit<'_>) -> Result<(), PyTraverseError> {
Ok(())
}
}
fn streaming_machine() -> CallMachine<Streaming> {
CallMachine::new(|host| {
Box::pin(async move {
host.project().await?;
if host.open(vec![("request-id", "req_1")]).await? == Demand::Detached {
return Ok(());
}
for chunk in ["first", "second"] {
if host.deliver(chunk).await? == Demand::Detached {
break;
}
}
Ok(())
})
})
}
/// Drives a `Stream` (async) or `SyncStream` to completion from a sync test.
fn read_all(py: Python<'_>, stream: &Bound<'_, PyAny>, asynchronous: bool) -> Vec<String> {
if !asynchronous {
return stream
.try_iter()
.unwrap()
.map(|chunk| chunk.unwrap().extract().unwrap())
.collect();
}
std::iter::from_fn(|| {
let stop = stream
.call_method0("__anext__")
.unwrap()
.call_method1("send", (py.None(),))
.unwrap_err();
if stop.is_instance_of::<pyo3::exceptions::PyStopAsyncIteration>(py) {
return None;
}
assert!(stop.is_instance_of::<pyo3::exceptions::PyStopIteration>(py));
Some(stop.value(py).getattr("value").unwrap().extract().unwrap())
})
.collect()
}
#[test]
fn a_stream_carries_its_head_as_hidden_params_before_the_first_chunk() {
let _guard = PYTHON_GLOBALS
.lock()
.unwrap_or_else(|error| error.into_inner());
crate::initialize_python();
Python::attach(|py| {
install_lifecycle_module(py);
for asynchronous in [false, true] {
let log = Log::default();
let adapter = SyntheticAdapter {
log: Log(log.0.clone()),
script: AdapterScript::Plain,
};
let handed = run_call(
py,
streaming_machine(),
StreamingHost,
Box::new(adapter),
PyDict::new(py).unbind(),
asynchronous,
)
.unwrap();
let stream = if asynchronous {
let stop = handed.call_method1(py, "send", (py.None(),)).unwrap_err();
stop.value(py).getattr("value").unwrap()
} else {
handed.into_bound(py)
};
let hidden: std::collections::HashMap<
String,
std::collections::HashMap<String, String>,
> = stream.getattr("_hidden_params").unwrap().extract().unwrap();
assert_eq!(
hidden["additional_headers"],
std::collections::HashMap::from([(
"request-id".to_string(),
"req_1".to_string()
)])
);
assert_eq!(log.entries(), ["started", "begin", "opened"]);
assert_eq!(read_all(py, &stream, asynchronous), ["first", "second"]);
}
});
}
fn failing_machine() -> CallMachine<Synthetic> {
CallMachine::new(|host| {
Box::pin(async move {
@ -1202,6 +1372,13 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri
) -> Result<(), InvokeError<Error>> {
Err(missing_state().into())
}
fn head(
&mut self,
_: Python<'_>,
head: std::convert::Infallible,
) -> PyResult<Py<PyAny>> {
match head {}
}
fn chunk(
&mut self,
_: Python<'_>,

View file

@ -8,9 +8,9 @@ use pyo3::prelude::*;
pub enum ExecutionStep {
Return(Py<PyAny>),
Await(Py<PyAny>),
/// The call streams: the caller gets a stream over this execution, which stays
/// suspended until the stream asks for a chunk.
Open,
/// The call streams: the caller gets a stream over this execution carrying this head,
/// and the execution stays suspended until the stream asks for a chunk.
Open(Py<PyAny>),
Yield(Py<PyAny>),
}
@ -75,7 +75,7 @@ impl Execution {
let step = body.resume(result)?;
let (tag, value, suspended) = match step {
ExecutionStep::Await(value) => ("Await", value, true),
ExecutionStep::Open => ("Open", py.None(), true),
ExecutionStep::Open(head) => ("Open", head, true),
ExecutionStep::Yield(value) => ("Yield", value, true),
ExecutionStep::Return(value) => ("Complete", value, false),
};

View file

@ -3,7 +3,7 @@ use std::convert::Infallible;
use bytes::Bytes;
use litellm_core::messages::{
Error,
route::{Messages, MessagesCall, MessagesOutput},
route::{Messages, MessagesCall, MessagesOutput, MessagesStreamHead},
types::MessagesShaping,
};
use litellm_host_python::{InvokeError, ProtocolHost, from_py, lookup, to_py};
@ -238,6 +238,13 @@ impl ProtocolHost for MessagesPythonHost {
}
}
fn head(&mut self, py: Python<'_>, head: MessagesStreamHead) -> PyResult<Py<PyAny>> {
py.import(ROUTE_HOST_MODULE)?
.getattr("stream_hidden_params")?
.call1((to_py(py, &head.headers)?,))
.map(Bound::unbind)
}
fn chunk(&mut self, py: Python<'_>, chunk: Bytes) -> PyResult<Py<PyAny>> {
Ok(PyBytes::new(py, &chunk).into_any().unbind())
}

View file

@ -4,14 +4,12 @@ use host::MessagesPythonHost;
use litellm_callbacks_legacy_python::{
LegacySurface, PassThroughStream, PublicCall, run_legacy_call,
};
use litellm_core::messages::route::{messages_machine, supports};
use litellm_core::messages::route::messages_machine;
use pyo3::{
prelude::*,
types::{PyDict, PyTuple},
};
use crate::errors::RustBridgeDeclined;
const SURFACE: LegacySurface = LegacySurface {
call_type: "anthropic_messages",
input_description: "Messages",
@ -28,17 +26,6 @@ fn run_messages(
kwargs: Bound<'_, PyDict>,
asynchronous: bool,
) -> PyResult<Py<PyAny>> {
let model: String = request.getattr("model")?.extract()?;
let provider: Option<String> = request.getattr("custom_llm_provider")?.extract()?;
let stream = request
.getattr("stream")?
.extract::<Option<bool>>()?
.unwrap_or(false);
if !supports(&model, provider.as_deref(), stream) {
return Err(RustBridgeDeclined::new_err(
"the Rust Messages route does not serve this provider",
));
}
let secrets = crate::secrets::source(py)?;
run_legacy_call(
py,

View file

@ -117,6 +117,10 @@ impl ProtocolHost for OcrPythonHost {
.map(Bound::unbind)
}
fn head(&mut self, _: Python<'_>, head: std::convert::Infallible) -> PyResult<Py<PyAny>> {
match head {}
}
fn chunk(&mut self, _: Python<'_>, chunk: std::convert::Infallible) -> PyResult<Py<PyAny>> {
match chunk {}
}

View file

@ -0,0 +1,32 @@
[package]
name = "litellm-testkit"
version = "0.1.0"
edition.workspace = true
license.workspace = true
repository.workspace = true
publish = false
[dependencies]
flate2.workspace = true
reqwest.workspace = true
serde.workspace = true
semver.workspace = true
serde_json.workspace = true
sha2.workspace = true
tar.workspace = true
target-lexicon.workspace = true
thiserror.workspace = true
tokio = { workspace = true, features = ["fs", "process"] }
zip.workspace = true
[dev-dependencies]
flate2.workspace = true
rstest.workspace = true
sha2.workspace = true
tar.workspace = true
target-lexicon.workspace = true
futures-util.workspace = true
tempfile.workspace = true
tokio.workspace = true
toml = "0.9"
zip.workspace = true

View file

@ -0,0 +1,181 @@
use std::collections::BTreeMap;
use std::path::Path;
use semver::Version;
use serde::Deserialize;
use super::{
Configure, Drive, Install, LaunchSpec, Outcome, Prompt, Settings, Usage, Wire, env, json_lines,
path_string,
};
use crate::install::release::parse;
use crate::install::{Packaging, Release};
use crate::{Error, Fetch, Target};
const RELEASES: &str = "https://downloads.claude.ai/claude-code-releases";
pub struct ClaudeCode;
#[derive(Deserialize)]
struct Manifest {
platforms: BTreeMap<String, Platform>,
}
#[derive(Deserialize)]
struct Platform {
checksum: String,
}
impl Install for ClaudeCode {
fn binary(&self) -> &'static str {
"claude"
}
async fn release(
&self,
fetch: &impl Fetch,
version: &Version,
target: Target,
) -> Result<Release, Error> {
let manifest_url = format!("{RELEASES}/{version}/manifest.json");
let manifest: Manifest = parse(&manifest_url, &fetch.get(&manifest_url).await?)?;
let key = format!(
"{}-{}{}",
target.os_name(),
target.arch_name(),
target.musl_suffix()
);
let platform = manifest
.platforms
.get(&key)
.ok_or_else(|| Error::AssetNotFound(key.clone()))?;
Ok(Release {
url: format!("{RELEASES}/{version}/{key}/claude"),
asset: key,
sha256: platform.checksum.clone(),
packaging: Packaging::Bare,
})
}
}
impl Configure for ClaudeCode {
fn configure(
&self,
_version: &Version,
settings: &Settings,
home: &Path,
) -> Result<LaunchSpec, Error> {
if settings.wire != Wire::Messages {
return Err(Error::UnsupportedWire {
agent: "claude",
wire: settings.wire,
});
}
Ok(LaunchSpec {
env: env([
("HOME", path_string(home)),
("CLAUDE_CONFIG_DIR", path_string(&home.join(".claude"))),
("ANTHROPIC_BASE_URL", settings.base_url.clone()),
("ANTHROPIC_AUTH_TOKEN", settings.api_key.clone()),
("ANTHROPIC_MODEL", settings.model.clone()),
("DISABLE_AUTOUPDATER", "1".to_owned()),
]),
files: BTreeMap::new(),
})
}
}
#[derive(Deserialize)]
#[serde(tag = "type", rename_all = "snake_case")]
enum Event {
Assistant {
message: AssistantMessage,
},
Result(Finished),
#[serde(other)]
Other,
}
#[derive(Deserialize)]
struct AssistantMessage {
content: Vec<Block>,
}
#[derive(Deserialize)]
struct Block {
#[serde(rename = "type")]
kind: String,
name: Option<String>,
}
#[derive(Deserialize)]
struct Finished {
is_error: bool,
result: Option<String>,
usage: Option<TokenUsage>,
}
#[derive(Deserialize)]
struct TokenUsage {
input_tokens: u64,
output_tokens: u64,
}
impl Drive for ClaudeCode {
fn args(&self, _version: &Version, settings: &Settings, prompt: &Prompt) -> Vec<String> {
let base = [
"-p",
&prompt.text,
"--output-format",
"stream-json",
"--verbose",
"--model",
&settings.model,
];
let tools = ["--allowedTools", "Bash,Read,Write,Edit"];
base.into_iter()
.chain(tools.into_iter().filter(|_| prompt.allow_tools))
.map(str::to_owned)
.collect()
}
fn parse(&self, _version: &Version, stdout: &str) -> Outcome {
let events: Vec<Event> = json_lines(stdout).collect();
let tool_calls = events
.iter()
.filter_map(|event| match event {
Event::Assistant { message } => Some(&message.content),
_ => None,
})
.flatten()
.filter(|block| block.kind == "tool_use")
.filter_map(|block| block.name.clone())
.collect();
let finished = events.into_iter().find_map(|event| match event {
Event::Result(finished) => Some(finished),
_ => None,
});
let Some(finished) = finished else {
return Outcome {
tool_calls,
..Outcome::default()
};
};
let result = finished.result.unwrap_or_default();
let (text, errors) = if finished.is_error {
(String::new(), vec![result])
} else {
(result, Vec::new())
};
Outcome {
text,
tool_calls,
usage: finished.usage.map_or_else(Usage::default, |usage| Usage {
input_tokens: usage.input_tokens,
output_tokens: usage.output_tokens,
}),
errors,
exit_code: None,
}
}
}

View file

@ -0,0 +1,174 @@
use std::collections::BTreeMap;
use std::path::{Path, PathBuf};
use semver::Version;
use serde::Deserialize;
use super::{
Configure, Drive, Install, LaunchSpec, Outcome, Prompt, Settings, Usage, Wire, env, json_lines,
path_string, quoted, v1,
};
use crate::install::release::github_release;
use crate::install::{Packaging, Release};
use crate::target::{Arch, Os};
use crate::{Error, Fetch, Target};
const RELEASES: &str = "https://api.github.com/repos/openai/codex/releases/tags";
pub struct Codex;
fn triple(target: Target) -> String {
let arch = match target.arch {
Arch::Aarch64 => "aarch64",
Arch::X86_64 => "x86_64",
};
match target.os {
Os::Macos => format!("{arch}-apple-darwin"),
Os::Linux => format!("{arch}-unknown-linux-musl"),
}
}
impl Install for Codex {
fn binary(&self) -> &'static str {
"codex"
}
async fn release(
&self,
fetch: &impl Fetch,
version: &Version,
target: Target,
) -> Result<Release, Error> {
let triple = triple(target);
github_release(
fetch,
RELEASES,
&format!("rust-v{version}"),
&format!("codex-{triple}.tar.gz"),
Packaging::TarGz {
member: format!("codex-{triple}"),
},
)
.await
}
}
impl Configure for Codex {
fn configure(
&self,
_version: &Version,
settings: &Settings,
home: &Path,
) -> Result<LaunchSpec, Error> {
if settings.wire != Wire::Responses {
return Err(Error::UnsupportedWire {
agent: "codex",
wire: settings.wire,
});
}
let config = format!(
"model = {model}\nmodel_provider = \"litellm\"\n\n[model_providers.litellm]\nname = \"LiteLLM\"\nbase_url = {base_url}\nenv_key = \"LITELLM_API_KEY\"\nwire_api = \"responses\"\n",
model = quoted(&settings.model),
base_url = quoted(&v1(settings)),
);
Ok(LaunchSpec {
env: env([
("HOME", path_string(home)),
("CODEX_HOME", path_string(&home.join(".codex"))),
("LITELLM_API_KEY", settings.api_key.clone()),
]),
files: BTreeMap::from([(PathBuf::from(".codex/config.toml"), config)]),
})
}
}
#[derive(Deserialize)]
enum EventKind {
#[serde(rename = "item.completed")]
ItemCompleted,
#[serde(rename = "turn.completed")]
TurnCompleted,
#[serde(rename = "turn.failed")]
TurnFailed,
#[serde(other)]
Other,
}
#[derive(Deserialize)]
struct Event {
#[serde(rename = "type")]
kind: EventKind,
item: Option<Item>,
usage: Option<TokenUsage>,
error: Option<Failure>,
}
#[derive(Deserialize)]
struct Item {
#[serde(rename = "type")]
kind: String,
text: Option<String>,
}
#[derive(Deserialize)]
struct TokenUsage {
input_tokens: u64,
output_tokens: u64,
}
#[derive(Deserialize)]
struct Failure {
message: String,
}
const NON_TOOL_ITEMS: [&str; 3] = ["agent_message", "reasoning", "error"];
impl Drive for Codex {
fn args(&self, _version: &Version, _settings: &Settings, prompt: &Prompt) -> Vec<String> {
let sandbox = ["--sandbox", "workspace-write"];
["exec", "--json", "--skip-git-repo-check"]
.into_iter()
.chain(sandbox.into_iter().filter(|_| prompt.allow_tools))
.chain([prompt.text.as_str()])
.map(str::to_owned)
.collect()
}
fn parse(&self, _version: &Version, stdout: &str) -> Outcome {
let events: Vec<Event> = json_lines(stdout).collect();
let items: Vec<&Item> = events
.iter()
.filter(|event| matches!(event.kind, EventKind::ItemCompleted))
.filter_map(|event| event.item.as_ref())
.collect();
Outcome {
text: items
.iter()
.rev()
.find(|item| item.kind == "agent_message")
.and_then(|item| item.text.clone())
.unwrap_or_default(),
tool_calls: items
.iter()
.filter(|item| !NON_TOOL_ITEMS.contains(&item.kind.as_str()))
.map(|item| item.kind.clone())
.collect(),
usage: events
.iter()
.filter(|event| matches!(event.kind, EventKind::TurnCompleted))
.filter_map(|event| event.usage.as_ref())
.map(|usage| Usage {
input_tokens: usage.input_tokens,
output_tokens: usage.output_tokens,
})
.fold(Usage::default(), |total, turn| total + turn),
errors: events
.iter()
.filter(|event| matches!(event.kind, EventKind::TurnFailed))
.filter_map(|event| event.error.as_ref())
.map(|failure| failure.message.clone())
.collect(),
exit_code: None,
}
}
}

View file

@ -0,0 +1,69 @@
use std::collections::BTreeMap;
use std::path::{Path, PathBuf};
use semver::Version;
use crate::Error;
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum Wire {
ChatCompletions,
Messages,
Responses,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct Settings {
pub base_url: String,
pub api_key: String,
pub model: String,
pub wire: Wire,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct LaunchSpec {
pub env: BTreeMap<String, String>,
pub files: BTreeMap<PathBuf, String>,
}
impl LaunchSpec {
pub fn write_files(&self, home: &Path) -> std::io::Result<()> {
self.files.iter().try_for_each(|(relative, contents)| {
let path = home.join(relative);
if let Some(parent) = path.parent() {
std::fs::create_dir_all(parent)?;
}
std::fs::write(path, contents)
})
}
}
pub trait Configure {
fn configure(
&self,
version: &Version,
settings: &Settings,
home: &Path,
) -> Result<LaunchSpec, Error>;
}
pub(crate) fn env(
pairs: impl IntoIterator<Item = (&'static str, String)>,
) -> BTreeMap<String, String> {
pairs
.into_iter()
.map(|(key, value)| (key.to_owned(), value))
.collect()
}
pub(crate) fn path_string(path: &Path) -> String {
path.to_string_lossy().into_owned()
}
pub(crate) fn quoted(value: &str) -> String {
serde_json::Value::from(value).to_string()
}
pub(crate) fn v1(settings: &Settings) -> String {
format!("{}/v1", settings.base_url.trim_end_matches('/'))
}

View file

@ -0,0 +1,57 @@
use std::ops::Add;
use semver::Version;
use crate::Settings;
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct Prompt {
pub text: String,
pub allow_tools: bool,
}
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
pub struct Usage {
pub input_tokens: u64,
pub output_tokens: u64,
}
impl Add for Usage {
type Output = Self;
fn add(self, other: Self) -> Self {
Self {
input_tokens: self.input_tokens + other.input_tokens,
output_tokens: self.output_tokens + other.output_tokens,
}
}
}
#[derive(Clone, Debug, Default, PartialEq, Eq)]
pub struct Outcome {
pub text: String,
pub tool_calls: Vec<String>,
pub usage: Usage,
pub errors: Vec<String>,
pub exit_code: Option<i32>,
}
impl Outcome {
pub fn succeeded(&self) -> bool {
self.exit_code == Some(0) && self.errors.is_empty()
}
}
pub trait Drive {
fn args(&self, version: &Version, settings: &Settings, prompt: &Prompt) -> Vec<String>;
fn parse(&self, version: &Version, stdout: &str) -> Outcome;
}
pub(crate) fn json_lines<'a, T: serde::de::DeserializeOwned + 'a>(
stdout: &'a str,
) -> impl Iterator<Item = T> + 'a {
stdout
.lines()
.filter_map(|line| serde_json::from_str(line).ok())
}

View file

@ -0,0 +1,17 @@
use std::future::Future;
use semver::Version;
use crate::install::Release;
use crate::{Error, Fetch, Target};
pub trait Install: Sync {
fn binary(&self) -> &'static str;
fn release(
&self,
fetch: &impl Fetch,
version: &Version,
target: Target,
) -> impl Future<Output = Result<Release, Error>> + Send;
}

View file

@ -0,0 +1,20 @@
mod claude;
mod codex;
mod configure;
mod drive;
mod install;
mod opencode;
pub use claude::ClaudeCode;
pub use codex::Codex;
pub use configure::{Configure, LaunchSpec, Settings, Wire};
pub use drive::{Drive, Outcome, Prompt, Usage};
pub use install::Install;
pub use opencode::Opencode;
pub(crate) use configure::{env, path_string, quoted, v1};
pub(crate) use drive::json_lines;
pub trait Agent: Install + Configure + Drive {}
impl<T: Install + Configure + Drive> Agent for T {}

View file

@ -0,0 +1,187 @@
use std::collections::BTreeMap;
use std::path::{Path, PathBuf};
use semver::Version;
use serde::Deserialize;
use super::{
Configure, Drive, Install, LaunchSpec, Outcome, Prompt, Settings, Usage, Wire, env, json_lines,
path_string, v1,
};
use crate::install::release::github_release;
use crate::install::{Packaging, Release};
use crate::target::Os;
use crate::{Error, Fetch, Target};
const RELEASES: &str = "https://api.github.com/repos/sst/opencode/releases/tags";
pub struct Opencode;
impl Install for Opencode {
fn binary(&self) -> &'static str {
"opencode"
}
async fn release(
&self,
fetch: &impl Fetch,
version: &Version,
target: Target,
) -> Result<Release, Error> {
let stem = format!(
"opencode-{}-{}{}",
target.os_name(),
target.arch_name(),
target.musl_suffix()
);
let member = "opencode".to_owned();
let (asset, packaging) = match target.os {
Os::Macos => (format!("{stem}.zip"), Packaging::Zip { member }),
Os::Linux => (format!("{stem}.tar.gz"), Packaging::TarGz { member }),
};
github_release(fetch, RELEASES, &format!("v{version}"), &asset, packaging).await
}
}
impl Configure for Opencode {
fn configure(
&self,
_version: &Version,
settings: &Settings,
home: &Path,
) -> Result<LaunchSpec, Error> {
let npm = match settings.wire {
Wire::ChatCompletions => "@ai-sdk/openai-compatible",
Wire::Responses => "@ai-sdk/openai",
Wire::Messages => "@ai-sdk/anthropic",
};
let config = serde_json::json!({
"$schema": "https://opencode.ai/config.json",
"model": format!("litellm/{}", settings.model),
"provider": {
"litellm": {
"npm": npm,
"name": "LiteLLM",
"options": { "baseURL": v1(settings), "apiKey": settings.api_key },
"models": { settings.model.clone(): { "name": settings.model } },
}
},
});
Ok(LaunchSpec {
env: env([
("HOME", path_string(home)),
("XDG_CONFIG_HOME", path_string(&home.join(".config"))),
("XDG_DATA_HOME", path_string(&home.join(".local/share"))),
("OPENCODE_DISABLE_AUTOUPDATE", "true".to_owned()),
]),
files: BTreeMap::from([(
PathBuf::from(".config/opencode/opencode.json"),
config.to_string(),
)]),
})
}
}
#[derive(Deserialize)]
#[serde(tag = "type", rename_all = "snake_case")]
enum Event {
Text {
part: TextPart,
},
ToolUse {
part: ToolPart,
},
StepFinish {
part: StepFinish,
},
Error {
error: Failure,
},
#[serde(other)]
Other,
}
#[derive(Deserialize)]
struct TextPart {
text: String,
}
#[derive(Deserialize)]
struct ToolPart {
tool: String,
}
#[derive(Deserialize)]
struct StepFinish {
tokens: Tokens,
}
#[derive(Deserialize)]
struct Tokens {
input: u64,
output: u64,
}
#[derive(Deserialize)]
struct Failure {
name: String,
data: Option<FailureData>,
}
#[derive(Deserialize)]
struct FailureData {
message: Option<String>,
}
impl Drive for Opencode {
fn args(&self, _version: &Version, _settings: &Settings, prompt: &Prompt) -> Vec<String> {
["run", "--format", "json", &prompt.text]
.map(str::to_owned)
.to_vec()
}
fn parse(&self, _version: &Version, stdout: &str) -> Outcome {
let events: Vec<Event> = json_lines(stdout).collect();
Outcome {
text: events
.iter()
.rev()
.find_map(|event| match event {
Event::Text { part } => Some(part.text.clone()),
_ => None,
})
.unwrap_or_default(),
tool_calls: events
.iter()
.filter_map(|event| match event {
Event::ToolUse { part } => Some(part.tool.clone()),
_ => None,
})
.collect(),
usage: events
.iter()
.filter_map(|event| match event {
Event::StepFinish { part } => Some(Usage {
input_tokens: part.tokens.input,
output_tokens: part.tokens.output,
}),
_ => None,
})
.fold(Usage::default(), |total, step| total + step),
errors: events
.iter()
.filter_map(|event| match event {
Event::Error { error } => Some(
error
.data
.as_ref()
.and_then(|data| data.message.clone())
.unwrap_or_else(|| error.name.clone()),
),
_ => None,
})
.collect(),
exit_code: None,
}
}
}

View file

@ -0,0 +1,56 @@
use std::io;
use std::path::PathBuf;
use thiserror::Error;
use crate::Wire;
#[derive(Debug, Error)]
pub enum Error {
#[error("unsupported target {0}")]
UnsupportedTarget(String),
#[error("{0} is not a plain x.y.z release version")]
InvalidVersion(String),
#[error("request to {url} failed")]
Request {
url: String,
#[source]
source: reqwest::Error,
},
#[error("{url} answered with status {status}")]
Status { url: String, status: u16 },
#[error("release metadata at {url} is malformed")]
Metadata {
url: String,
#[source]
source: serde_json::Error,
},
#[error("release has no asset named {0}")]
AssetNotFound(String),
#[error("release publishes no sha256 for {0}")]
MissingChecksum(String),
#[error("sha256 mismatch for {asset}: expected {expected}, got {actual}")]
ChecksumMismatch {
asset: String,
expected: String,
actual: String,
},
#[error("archive does not contain {0}")]
ArchiveMemberNotFound(String),
#[error("archive is unreadable")]
Archive(#[source] io::Error),
#[error("zip archive is unreadable")]
Zip(#[from] zip::result::ZipError),
#[error("{binary} reports version '{reported}', expected {expected}")]
VersionMismatch {
binary: PathBuf,
expected: String,
reported: String,
},
#[error("{agent} cannot talk to the gateway over {wire:?}")]
UnsupportedWire { agent: &'static str, wire: Wire },
#[error("agent did not finish within {0:?}")]
Timeout(std::time::Duration),
#[error("io failure")]
Io(#[from] io::Error),
}

View file

@ -0,0 +1,52 @@
use std::io::{Cursor, Read};
use flate2::read::GzDecoder;
use sha2::{Digest, Sha256};
use super::release::Packaging;
use crate::Error;
pub(crate) fn verify_sha256(asset: &str, expected: &str, bytes: &[u8]) -> Result<(), Error> {
let actual = format!("{:x}", Sha256::digest(bytes));
if actual.eq_ignore_ascii_case(expected) {
return Ok(());
}
Err(Error::ChecksumMismatch {
asset: asset.to_owned(),
expected: expected.to_owned(),
actual,
})
}
pub(crate) fn extract_binary(packaging: &Packaging, bytes: &[u8]) -> Result<Vec<u8>, Error> {
match packaging {
Packaging::Bare => Ok(bytes.to_vec()),
Packaging::TarGz { member } => extract_tar_gz(member, bytes),
Packaging::Zip { member } => extract_zip(member, bytes),
}
}
fn extract_tar_gz(member: &str, bytes: &[u8]) -> Result<Vec<u8>, Error> {
let mut archive = tar::Archive::new(GzDecoder::new(bytes));
for entry in archive.entries().map_err(Error::Archive)? {
let mut entry = entry.map_err(Error::Archive)?;
let path = entry.path().map_err(Error::Archive)?;
if path.file_name().is_some_and(|name| name == member) {
let mut binary = Vec::new();
entry.read_to_end(&mut binary).map_err(Error::Archive)?;
return Ok(binary);
}
}
Err(Error::ArchiveMemberNotFound(member.to_owned()))
}
fn extract_zip(member: &str, bytes: &[u8]) -> Result<Vec<u8>, Error> {
let mut archive = zip::ZipArchive::new(Cursor::new(bytes))?;
let mut file = archive.by_name(member).map_err(|error| match error {
zip::result::ZipError::FileNotFound => Error::ArchiveMemberNotFound(member.to_owned()),
other => Error::Zip(other),
})?;
let mut binary = Vec::new();
file.read_to_end(&mut binary).map_err(Error::Archive)?;
Ok(binary)
}

View file

@ -0,0 +1,55 @@
use std::future::Future;
use crate::Error;
pub trait Fetch: Sync {
fn get(&self, url: &str) -> impl Future<Output = Result<Vec<u8>, Error>> + Send;
}
pub struct HttpFetch {
client: reqwest::Client,
github_token: Option<String>,
}
impl HttpFetch {
pub fn new(github_token: Option<String>) -> Self {
Self {
client: reqwest::Client::new(),
github_token,
}
}
pub fn from_env() -> Self {
Self::new(std::env::var("GITHUB_TOKEN").ok())
}
}
impl Fetch for HttpFetch {
async fn get(&self, url: &str) -> Result<Vec<u8>, Error> {
let request = self
.client
.get(url)
.header("user-agent", "litellm-testkit")
.header("accept", "application/json, application/octet-stream");
let request = match (
&self.github_token,
url.starts_with("https://api.github.com/"),
) {
(Some(token), true) => request.bearer_auth(token),
_ => request,
};
let request_error = |source| Error::Request {
url: url.to_owned(),
source,
};
let response = request.send().await.map_err(request_error)?;
let status = response.status();
if !status.is_success() {
return Err(Error::Status {
url: url.to_owned(),
status: status.as_u16(),
});
}
Ok(response.bytes().await.map_err(request_error)?.to_vec())
}
}

View file

@ -0,0 +1,118 @@
mod archive;
mod fetch;
pub(crate) mod release;
use std::os::unix::fs::PermissionsExt;
use std::path::{Path, PathBuf};
use std::process::Stdio;
use std::sync::atomic::{AtomicU64, Ordering};
use semver::Version;
use tokio::fs;
use tokio::process::Command;
use crate::{Error, Install, Target};
use archive::{extract_binary, verify_sha256};
static STAGING_COUNTER: AtomicU64 = AtomicU64::new(0);
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct Installed {
pub version: Version,
pub binary: PathBuf,
}
pub struct Installer<F> {
fetch: F,
cache_root: PathBuf,
target: Target,
}
impl<F: Fetch> Installer<F> {
pub fn new(fetch: F, cache_root: impl Into<PathBuf>, target: Target) -> Self {
Self {
fetch,
cache_root: cache_root.into(),
target,
}
}
pub async fn install(
&self,
agent: &impl Install,
version: &Version,
) -> Result<Installed, Error> {
validate_release(version)?;
let dir = self
.cache_root
.join(agent.binary())
.join(version.to_string());
let binary = dir.join(agent.binary());
let installed = Installed {
version: version.clone(),
binary: binary.clone(),
};
if fs::try_exists(&binary).await? && probe_version(&binary, version).await.is_ok() {
return Ok(installed);
}
let release = agent.release(&self.fetch, version, self.target).await?;
let archive = self.fetch.get(&release.url).await?;
verify_sha256(&release.asset, &release.sha256, &archive)?;
let contents = extract_binary(&release.packaging, &archive)?;
fs::create_dir_all(&dir).await?;
let staging = dir.join(format!(
".{}.{}.{}.partial",
agent.binary(),
std::process::id(),
STAGING_COUNTER.fetch_add(1, Ordering::Relaxed)
));
fs::write(&staging, contents).await?;
fs::set_permissions(&staging, std::fs::Permissions::from_mode(0o755)).await?;
fs::rename(&staging, &binary).await?;
match probe_version(&binary, version).await {
Ok(()) => Ok(installed),
Err(error) => {
fs::remove_file(&binary).await?;
Err(error)
}
}
}
}
fn validate_release(version: &Version) -> Result<(), Error> {
if version.pre.is_empty() && version.build.is_empty() {
return Ok(());
}
Err(Error::InvalidVersion(version.to_string()))
}
async fn probe_version(binary: &Path, expected: &Version) -> Result<(), Error> {
let home = std::env::temp_dir();
let output = Command::new(binary)
.arg("--version")
.env_clear()
.env("HOME", home)
.env("DISABLE_AUTOUPDATER", "1")
.stdin(Stdio::null())
.output()
.await?;
let stdout = String::from_utf8_lossy(&output.stdout);
if stdout
.split_whitespace()
.filter_map(|token| Version::parse(token).ok())
.any(|reported| &reported == expected)
{
return Ok(());
}
Err(Error::VersionMismatch {
binary: binary.to_owned(),
expected: expected.to_string(),
reported: stdout.trim().to_owned(),
})
}
pub use fetch::{Fetch, HttpFetch};
pub use release::{Packaging, Release};

View file

@ -0,0 +1,65 @@
use serde::Deserialize;
use crate::{Error, Fetch};
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum Packaging {
Bare,
TarGz { member: String },
Zip { member: String },
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct Release {
pub asset: String,
pub url: String,
pub sha256: String,
pub packaging: Packaging,
}
#[derive(Deserialize)]
struct GithubRelease {
assets: Vec<GithubAsset>,
}
#[derive(Deserialize)]
struct GithubAsset {
name: String,
digest: Option<String>,
browser_download_url: String,
}
pub(crate) async fn github_release(
fetch: &impl Fetch,
releases_url: &str,
tag: &str,
asset_name: &str,
packaging: Packaging,
) -> Result<Release, Error> {
let url = format!("{releases_url}/{tag}");
let release: GithubRelease = parse(&url, &fetch.get(&url).await?)?;
let asset = release
.assets
.into_iter()
.find(|asset| asset.name == asset_name)
.ok_or_else(|| Error::AssetNotFound(asset_name.to_owned()))?;
let sha256 = asset
.digest
.as_deref()
.and_then(|digest| digest.strip_prefix("sha256:"))
.ok_or_else(|| Error::MissingChecksum(asset_name.to_owned()))?
.to_owned();
Ok(Release {
asset: asset.name,
url: asset.browser_download_url,
sha256,
packaging,
})
}
pub(crate) fn parse<T: for<'de> Deserialize<'de>>(url: &str, body: &[u8]) -> Result<T, Error> {
serde_json::from_slice(body).map_err(|source| Error::Metadata {
url: url.to_owned(),
source,
})
}

View file

@ -0,0 +1,15 @@
mod agent;
mod error;
mod install;
mod session;
mod target;
pub use agent::{
Agent, ClaudeCode, Codex, Configure, Drive, Install, LaunchSpec, Opencode, Outcome, Prompt,
Settings, Usage, Wire,
};
pub use error::Error;
pub use install::{Fetch, HttpFetch, Installed, Installer, Packaging, Release};
pub use semver::Version;
pub use session::Session;
pub use target::{Arch, Os, Target};

View file

@ -0,0 +1,76 @@
use std::collections::BTreeMap;
use std::path::PathBuf;
use std::process::Stdio;
use std::time::Duration;
use semver::Version;
use tokio::process::Command;
use tokio::time::timeout;
use crate::{Configure, Drive, Error, Installed, Outcome, Prompt, Settings};
const STDERR_LIMIT_CHARS: usize = 2000;
pub struct Session {
binary: PathBuf,
home: PathBuf,
version: Version,
settings: Settings,
env: BTreeMap<String, String>,
}
impl Session {
pub fn prepare(
agent: &impl Configure,
installed: &Installed,
settings: Settings,
home: impl Into<PathBuf>,
) -> Result<Self, Error> {
let home = home.into();
let spec = agent.configure(&installed.version, &settings, &home)?;
spec.write_files(&home)?;
Ok(Self {
binary: installed.binary.clone(),
home,
version: installed.version.clone(),
settings,
env: spec.env,
})
}
pub async fn run(
&self,
agent: &impl Drive,
prompt: &Prompt,
limit: Duration,
) -> Result<Outcome, Error> {
let child = Command::new(&self.binary)
.args(agent.args(&self.version, &self.settings, prompt))
.env_clear()
.env("PATH", "/usr/bin:/bin")
.envs(&self.env)
.current_dir(&self.home)
.stdin(Stdio::null())
.kill_on_drop(true)
.output();
let output = timeout(limit, child)
.await
.map_err(|_| Error::Timeout(limit))??;
let parsed = agent.parse(&self.version, &String::from_utf8_lossy(&output.stdout));
let failed_silently = !output.status.success() && parsed.errors.is_empty();
Ok(Outcome {
errors: if failed_silently {
vec![
String::from_utf8_lossy(&output.stderr)
.chars()
.take(STDERR_LIMIT_CHARS)
.collect(),
]
} else {
parsed.errors
},
exit_code: output.status.code(),
..parsed
})
}
}

View file

@ -0,0 +1,69 @@
use target_lexicon::{Architecture, Environment, OperatingSystem, Triple};
use crate::Error;
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum Os {
Macos,
Linux,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum Arch {
Aarch64,
X86_64,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct Target {
pub os: Os,
pub arch: Arch,
pub musl: bool,
}
impl Target {
pub fn host() -> Result<Self, Error> {
Self::try_from(&Triple::host())
}
pub(crate) const fn os_name(self) -> &'static str {
match self.os {
Os::Macos => "darwin",
Os::Linux => "linux",
}
}
pub(crate) const fn arch_name(self) -> &'static str {
match self.arch {
Arch::Aarch64 => "arm64",
Arch::X86_64 => "x64",
}
}
pub(crate) const fn musl_suffix(self) -> &'static str {
if self.musl { "-musl" } else { "" }
}
}
impl TryFrom<&Triple> for Target {
type Error = Error;
fn try_from(triple: &Triple) -> Result<Self, Error> {
let unsupported = || Error::UnsupportedTarget(triple.to_string());
let os = match triple.operating_system {
OperatingSystem::Darwin(_) | OperatingSystem::MacOSX(_) => Os::Macos,
OperatingSystem::Linux => Os::Linux,
_ => return Err(unsupported()),
};
let arch = match triple.architecture {
Architecture::Aarch64(_) => Arch::Aarch64,
Architecture::X86_64 => Arch::X86_64,
_ => return Err(unsupported()),
};
Ok(Self {
os,
arch,
musl: triple.environment == Environment::Musl,
})
}
}

View file

@ -0,0 +1,133 @@
use std::path::Path;
use litellm_testkit::{ClaudeCode, Codex, Configure, Error, Opencode, Settings, Version, Wire};
use rstest::rstest;
fn settings(wire: Wire) -> Settings {
Settings {
base_url: "http://localhost:4000/".to_owned(),
api_key: "sk-test \"quoted\"".to_owned(),
model: "some-model".to_owned(),
wire,
}
}
fn version() -> Version {
Version::new(1, 2, 3)
}
#[rstest]
#[case(&ClaudeCode, Wire::Messages)]
#[case(&Codex, Wire::Responses)]
#[case(&Opencode, Wire::ChatCompletions)]
fn every_agent_runs_inside_the_given_home(#[case] agent: &impl Configure, #[case] wire: Wire) {
let home = Path::new("/scratch/home");
let spec = agent.configure(&version(), &settings(wire), home).unwrap();
assert_eq!(spec.env["HOME"], "/scratch/home");
assert!(
spec.env
.iter()
.filter(|(key, _)| key.ends_with("_HOME") || key.as_str() == "CLAUDE_CONFIG_DIR")
.all(|(_, value)| value.starts_with("/scratch/home"))
);
assert!(spec.files.keys().all(|path| path.is_relative()));
}
#[rstest]
#[case::claude_code(&ClaudeCode, &[Wire::ChatCompletions, Wire::Responses])]
#[case::codex(&Codex, &[Wire::ChatCompletions, Wire::Messages])]
fn wires_an_agent_cannot_speak_are_refused(
#[case] agent: &impl Configure,
#[case] refused: &[Wire],
) {
refused.iter().for_each(|wire| {
let result = agent.configure(&version(), &settings(*wire), Path::new("/h"));
assert!(matches!(result, Err(Error::UnsupportedWire { wire: got, .. }) if got == *wire));
});
}
#[test]
fn claude_code_points_at_the_gateway_root_with_the_key_and_model() {
let spec = ClaudeCode
.configure(&version(), &settings(Wire::Messages), Path::new("/h"))
.unwrap();
assert_eq!(spec.env["ANTHROPIC_BASE_URL"], "http://localhost:4000/");
assert_eq!(spec.env["ANTHROPIC_AUTH_TOKEN"], "sk-test \"quoted\"");
assert_eq!(spec.env["ANTHROPIC_MODEL"], "some-model");
}
#[test]
fn codex_config_is_valid_toml_routing_the_responses_api_to_the_gateway() {
let dir = tempfile::tempdir().unwrap();
let spec = Codex
.configure(&version(), &settings(Wire::Responses), dir.path())
.unwrap();
spec.write_files(dir.path()).unwrap();
let config: toml::Table =
toml::from_str(&std::fs::read_to_string(dir.path().join(".codex/config.toml")).unwrap())
.unwrap();
let provider = &config["model_providers"]["litellm"];
assert_eq!(config["model"].as_str(), Some("some-model"));
assert_eq!(config["model_provider"].as_str(), Some("litellm"));
assert_eq!(
provider["base_url"].as_str(),
Some("http://localhost:4000/v1")
);
assert_eq!(provider["wire_api"].as_str(), Some("responses"));
let key_var = provider["env_key"].as_str().unwrap();
assert_eq!(spec.env[key_var], "sk-test \"quoted\"");
}
#[rstest]
#[case(Wire::ChatCompletions)]
#[case(Wire::Responses)]
#[case(Wire::Messages)]
fn opencode_config_is_valid_json_registering_the_gateway_model(#[case] wire: Wire) {
let dir = tempfile::tempdir().unwrap();
let spec = Opencode
.configure(&version(), &settings(wire), dir.path())
.unwrap();
spec.write_files(dir.path()).unwrap();
let config: serde_json::Value = serde_json::from_str(
&std::fs::read_to_string(dir.path().join(".config/opencode/opencode.json")).unwrap(),
)
.unwrap();
let provider = &config["provider"]["litellm"];
assert_eq!(config["model"], "litellm/some-model");
assert_eq!(provider["options"]["baseURL"], "http://localhost:4000/v1");
assert_eq!(provider["options"]["apiKey"], "sk-test \"quoted\"");
assert!(provider["models"]["some-model"].is_object());
}
#[test]
fn opencode_uses_a_different_provider_package_for_every_wire() {
let package = |wire| {
let dir = tempfile::tempdir().unwrap();
let spec = Opencode
.configure(&version(), &settings(wire), dir.path())
.unwrap();
let config: serde_json::Value =
serde_json::from_str(spec.files.values().next().unwrap()).unwrap();
config["provider"]["litellm"]["npm"]
.as_str()
.unwrap()
.to_owned()
};
let packages = [Wire::ChatCompletions, Wire::Responses, Wire::Messages].map(package);
assert_eq!(
packages
.iter()
.collect::<std::collections::BTreeSet<_>>()
.len(),
packages.len()
);
}

View file

@ -0,0 +1,262 @@
mod support;
use std::str::FromStr;
use litellm_testkit::{ClaudeCode, Codex, Error, Installer, Opencode, Target, Version};
use rstest::rstest;
use serde_json::json;
use support::{FakeFetch, script_printing, sha256, tar_gz, zip_archive};
use target_lexicon::Triple;
fn target(triple: &str) -> Target {
Target::try_from(&Triple::from_str(triple).unwrap()).unwrap()
}
fn linux() -> Target {
target("x86_64-unknown-linux-gnu")
}
fn version() -> Version {
Version::new(9, 8, 7)
}
fn github_release(asset: &str, download_url: &str, digest: Option<String>) -> Vec<u8> {
json!({
"assets": [
{ "name": "unrelated.txt", "digest": "sha256:00", "browser_download_url": "https://example.test/unrelated" },
{ "name": asset, "digest": digest, "browser_download_url": download_url },
]
})
.to_string()
.into_bytes()
}
fn claude_routes(binary: &[u8], checksum: &str) -> Vec<(String, Vec<u8>)> {
let base = "https://downloads.claude.ai/claude-code-releases/9.8.7";
let manifest = json!({ "platforms": { "linux-x64": { "checksum": checksum } } });
vec![
(
format!("{base}/manifest.json"),
manifest.to_string().into_bytes(),
),
(format!("{base}/linux-x64/claude"), binary.to_vec()),
]
}
fn codex_routes(archive: Vec<u8>, digest: Option<String>) -> Vec<(String, Vec<u8>)> {
let release = github_release(
"codex-x86_64-unknown-linux-musl.tar.gz",
"https://example.test/codex.tar.gz",
digest,
);
vec![
(
"https://api.github.com/repos/openai/codex/releases/tags/rust-v9.8.7".to_owned(),
release,
),
("https://example.test/codex.tar.gz".to_owned(), archive),
]
}
#[tokio::test]
async fn claude_bare_binary_is_installed_and_runnable() {
let binary = script_printing("9.8.7 (Claude Code)");
let fetch = FakeFetch::new(claude_routes(&binary, &sha256(&binary)));
let cache = tempfile::tempdir().unwrap();
let installed = Installer::new(&fetch, cache.path(), linux())
.install(&ClaudeCode, &version())
.await
.unwrap();
assert_eq!(installed.binary, cache.path().join("claude/9.8.7/claude"));
assert_eq!(std::fs::read(&installed.binary).unwrap(), binary);
}
#[tokio::test]
async fn codex_binary_is_extracted_from_the_tarball_under_its_own_name() {
let binary = script_printing("codex-cli 9.8.7");
let archive = tar_gz("codex-x86_64-unknown-linux-musl", &binary);
let fetch = FakeFetch::new(codex_routes(
archive.clone(),
Some(format!("sha256:{}", sha256(&archive))),
));
let cache = tempfile::tempdir().unwrap();
let installed = Installer::new(&fetch, cache.path(), linux())
.install(&Codex, &version())
.await
.unwrap();
assert_eq!(std::fs::read(&installed.binary).unwrap(), binary);
assert_eq!(installed.binary, cache.path().join("codex/9.8.7/codex"));
}
#[tokio::test]
async fn opencode_binary_is_extracted_from_the_darwin_zip() {
let binary = script_printing("9.8.7");
let archive = zip_archive("opencode", &binary);
let release = github_release(
"opencode-darwin-arm64.zip",
"https://example.test/opencode.zip",
Some(format!("sha256:{}", sha256(&archive))),
);
let fetch = FakeFetch::new([
(
"https://api.github.com/repos/sst/opencode/releases/tags/v9.8.7".to_owned(),
release,
),
("https://example.test/opencode.zip".to_owned(), archive),
]);
let cache = tempfile::tempdir().unwrap();
let installed = Installer::new(&fetch, cache.path(), target("aarch64-apple-darwin"))
.install(&Opencode, &version())
.await
.unwrap();
assert_eq!(std::fs::read(&installed.binary).unwrap(), binary);
}
#[tokio::test]
async fn tampered_download_is_rejected_and_nothing_is_left_behind() {
let binary = script_printing("9.8.7 (Claude Code)");
let fetch = FakeFetch::new(claude_routes(&binary, &sha256(b"what the vendor signed")));
let cache = tempfile::tempdir().unwrap();
let result = Installer::new(&fetch, cache.path(), linux())
.install(&ClaudeCode, &version())
.await;
assert!(matches!(result, Err(Error::ChecksumMismatch { .. })));
assert!(!cache.path().join("claude/9.8.7").exists());
}
#[tokio::test]
async fn github_asset_without_a_digest_is_refused() {
let archive = tar_gz(
"codex-x86_64-unknown-linux-musl",
&script_printing("codex-cli 9.8.7"),
);
let fetch = FakeFetch::new(codex_routes(archive, None));
let cache = tempfile::tempdir().unwrap();
let result = Installer::new(&fetch, cache.path(), linux())
.install(&Codex, &version())
.await;
assert!(matches!(result, Err(Error::MissingChecksum(_))));
}
#[tokio::test]
async fn binary_reporting_a_different_version_is_removed() {
let binary = script_printing("1.0.0 (Claude Code)");
let fetch = FakeFetch::new(claude_routes(&binary, &sha256(&binary)));
let cache = tempfile::tempdir().unwrap();
let result = Installer::new(&fetch, cache.path(), linux())
.install(&ClaudeCode, &version())
.await;
assert!(matches!(result, Err(Error::VersionMismatch { .. })));
assert!(!cache.path().join("claude/9.8.7/claude").exists());
}
#[tokio::test]
async fn second_install_reuses_the_cached_binary_without_downloading() {
let binary = script_printing("9.8.7 (Claude Code)");
let fetch = FakeFetch::new(claude_routes(&binary, &sha256(&binary)));
let cache = tempfile::tempdir().unwrap();
let installer = Installer::new(&fetch, cache.path(), linux());
let first = installer.install(&ClaudeCode, &version()).await.unwrap();
let calls_after_first = fetch.calls();
let second = installer.install(&ClaudeCode, &version()).await.unwrap();
assert_eq!(first, second);
assert_eq!(fetch.calls(), calls_after_first);
}
#[tokio::test]
async fn corrupted_cache_entry_is_replaced_by_a_fresh_download() {
let binary = script_printing("9.8.7 (Claude Code)");
let fetch = FakeFetch::new(claude_routes(&binary, &sha256(&binary)));
let cache = tempfile::tempdir().unwrap();
let installer = Installer::new(&fetch, cache.path(), linux());
let installed = installer.install(&ClaudeCode, &version()).await.unwrap();
std::fs::write(&installed.binary, script_printing("0.0.1")).unwrap();
installer.install(&ClaudeCode, &version()).await.unwrap();
assert_eq!(std::fs::read(&installed.binary).unwrap(), binary);
}
#[rstest]
#[case("9.8.7-beta.1")]
#[case("9.8.7+build.5")]
#[tokio::test]
async fn pre_releases_never_reach_the_network_or_the_filesystem(#[case] version: &str) {
let fetch = FakeFetch::new([]);
let cache = tempfile::tempdir().unwrap();
let result = Installer::new(&fetch, cache.path(), linux())
.install(&ClaudeCode, &Version::parse(version).unwrap())
.await;
assert!(matches!(result, Err(Error::InvalidVersion(_))));
assert_eq!(fetch.calls(), 0);
assert_eq!(std::fs::read_dir(cache.path()).unwrap().count(), 0);
}
#[tokio::test]
async fn musl_linux_picks_the_musl_claude_build() {
let binary = script_printing("9.8.7 (Claude Code)");
let base = "https://downloads.claude.ai/claude-code-releases/9.8.7";
let manifest = json!({ "platforms": {
"linux-x64": { "checksum": sha256(b"glibc build") },
"linux-x64-musl": { "checksum": sha256(&binary) },
} });
let fetch = FakeFetch::new([
(
format!("{base}/manifest.json"),
manifest.to_string().into_bytes(),
),
(format!("{base}/linux-x64-musl/claude"), binary.clone()),
]);
let cache = tempfile::tempdir().unwrap();
let installed = Installer::new(&fetch, cache.path(), target("x86_64-unknown-linux-musl"))
.install(&ClaudeCode, &version())
.await
.unwrap();
assert_eq!(std::fs::read(&installed.binary).unwrap(), binary);
}
#[rstest]
#[case("x86_64-pc-windows-msvc")]
#[case("riscv64gc-unknown-linux-gnu")]
#[case("wasm32-unknown-unknown")]
fn targets_no_agent_ships_for_are_rejected(#[case] triple: &str) {
let result = Target::try_from(&Triple::from_str(triple).unwrap());
assert!(matches!(result, Err(Error::UnsupportedTarget(_))));
}
#[tokio::test]
async fn concurrent_installs_of_the_same_version_both_succeed() {
let binary = script_printing("9.8.7 (Claude Code)");
let fetch = FakeFetch::new(claude_routes(&binary, &sha256(&binary)));
let cache = tempfile::tempdir().unwrap();
let installer = Installer::new(&fetch, cache.path(), linux());
let wanted = version();
let installs =
futures_util::future::join_all((0..8).map(|_| installer.install(&ClaudeCode, &wanted)))
.await;
assert!(installs.iter().all(Result::is_ok));
assert_eq!(
std::fs::read(&installs[0].as_ref().unwrap().binary).unwrap(),
binary
);
}

View file

@ -0,0 +1,133 @@
//! Drives the real agents through a real gateway. Run with `cargo test -p litellm-testkit --test live -- --ignored`
//! after exporting `TESTKIT_GATEWAY_URL`, `TESTKIT_GATEWAY_KEY`, one `TESTKIT_MODEL_<WIRE>` per wire
//! (`MESSAGES`, `RESPONSES`, `CHAT_COMPLETIONS`) and one `TESTKIT_<AGENT>_VERSION` per agent
//! (`CLAUDE`, `CODEX`, `OPENCODE`). `TESTKIT_CACHE_DIR` and `GITHUB_TOKEN` are optional.
use std::path::PathBuf;
use std::time::Duration;
use litellm_testkit::{
Agent, ClaudeCode, Codex, HttpFetch, Installer, Opencode, Outcome, Prompt, Session, Settings,
Target, Version, Wire,
};
use rstest::rstest;
const LIMIT: Duration = Duration::from_secs(180);
fn required(name: &str) -> String {
std::env::var(name).unwrap_or_else(|_| panic!("{name} must be set to run the live tests"))
}
fn model_var(wire: Wire) -> &'static str {
match wire {
Wire::Messages => "TESTKIT_MODEL_MESSAGES",
Wire::Responses => "TESTKIT_MODEL_RESPONSES",
Wire::ChatCompletions => "TESTKIT_MODEL_CHAT_COMPLETIONS",
}
}
async fn drive(
agent: &impl Agent,
version_var: &str,
wire: Wire,
model: Option<&str>,
prompt: Prompt,
) -> Outcome {
let cache = std::env::var("TESTKIT_CACHE_DIR")
.map(PathBuf::from)
.unwrap_or_else(|_| std::env::temp_dir().join("litellm-testkit-cache"));
let installer = Installer::new(HttpFetch::from_env(), cache, Target::host().unwrap());
let installed = installer
.install(agent, &Version::parse(&required(version_var)).unwrap())
.await
.unwrap();
let settings = Settings {
base_url: required("TESTKIT_GATEWAY_URL"),
api_key: required("TESTKIT_GATEWAY_KEY"),
model: model.map_or_else(|| required(model_var(wire)), str::to_owned),
wire,
};
let home = tempfile::tempdir().unwrap();
let session = Session::prepare(agent, &installed, settings, home.path()).unwrap();
session.run(agent, &prompt, LIMIT).await.unwrap()
}
fn text_prompt() -> Prompt {
Prompt {
text: "Reply with the single word: pong".to_owned(),
allow_tools: false,
}
}
fn tool_prompt() -> Prompt {
Prompt {
text: "Run the shell command 'echo tool-ok' and reply with exactly its output.".to_owned(),
allow_tools: true,
}
}
#[rstest]
#[case::claude_messages(&ClaudeCode, "TESTKIT_CLAUDE_VERSION", Wire::Messages)]
#[case::codex_responses(&Codex, "TESTKIT_CODEX_VERSION", Wire::Responses)]
#[case::opencode_chat(&Opencode, "TESTKIT_OPENCODE_VERSION", Wire::ChatCompletions)]
#[case::opencode_responses(&Opencode, "TESTKIT_OPENCODE_VERSION", Wire::Responses)]
#[case::opencode_messages(&Opencode, "TESTKIT_OPENCODE_VERSION", Wire::Messages)]
#[ignore = "needs a live gateway, see the module docs"]
#[tokio::test]
async fn plain_prompt_gets_an_answer_and_token_usage(
#[case] agent: &impl Agent,
#[case] version_var: &str,
#[case] wire: Wire,
) {
let outcome = drive(agent, version_var, wire, None, text_prompt()).await;
assert!(outcome.succeeded(), "{outcome:?}");
assert!(outcome.text.to_lowercase().contains("pong"), "{outcome:?}");
assert!(outcome.usage.output_tokens > 0, "{outcome:?}");
}
#[rstest]
#[case::claude_messages(&ClaudeCode, "TESTKIT_CLAUDE_VERSION", Wire::Messages)]
#[case::codex_responses(&Codex, "TESTKIT_CODEX_VERSION", Wire::Responses)]
#[case::opencode_chat(&Opencode, "TESTKIT_OPENCODE_VERSION", Wire::ChatCompletions)]
#[case::opencode_responses(&Opencode, "TESTKIT_OPENCODE_VERSION", Wire::Responses)]
#[case::opencode_messages(&Opencode, "TESTKIT_OPENCODE_VERSION", Wire::Messages)]
#[ignore = "needs a live gateway, see the module docs"]
#[tokio::test]
async fn tool_use_is_reported_and_its_result_reaches_the_answer(
#[case] agent: &impl Agent,
#[case] version_var: &str,
#[case] wire: Wire,
) {
let outcome = drive(agent, version_var, wire, None, tool_prompt()).await;
assert!(outcome.succeeded(), "{outcome:?}");
assert!(!outcome.tool_calls.is_empty(), "{outcome:?}");
assert!(outcome.text.contains("tool-ok"), "{outcome:?}");
}
#[rstest]
#[case::claude_messages(&ClaudeCode, "TESTKIT_CLAUDE_VERSION", Wire::Messages)]
#[case::codex_responses(&Codex, "TESTKIT_CODEX_VERSION", Wire::Responses)]
#[case::opencode_chat(&Opencode, "TESTKIT_OPENCODE_VERSION", Wire::ChatCompletions)]
#[case::opencode_responses(&Opencode, "TESTKIT_OPENCODE_VERSION", Wire::Responses)]
#[case::opencode_messages(&Opencode, "TESTKIT_OPENCODE_VERSION", Wire::Messages)]
#[ignore = "needs a live gateway, see the module docs"]
#[tokio::test]
async fn model_the_gateway_rejects_is_reported_as_an_error(
#[case] agent: &impl Agent,
#[case] version_var: &str,
#[case] wire: Wire,
) {
let outcome = drive(
agent,
version_var,
wire,
Some("testkit-no-such-model"),
text_prompt(),
)
.await;
assert!(!outcome.succeeded(), "{outcome:?}");
assert!(!outcome.errors.is_empty(), "{outcome:?}");
}

View file

@ -0,0 +1,155 @@
use std::os::unix::fs::PermissionsExt;
use std::path::{Path, PathBuf};
use std::time::Duration;
use litellm_testkit::{
Configure, Drive, Error, Installed, LaunchSpec, Outcome, Prompt, Session, Settings, Version,
Wire,
};
struct Scripted;
impl Configure for Scripted {
fn configure(
&self,
version: &Version,
_settings: &Settings,
home: &Path,
) -> Result<LaunchSpec, Error> {
Ok(LaunchSpec {
env: [
("AGENT_HOME".to_owned(), home.to_string_lossy().into_owned()),
("AGENT_SAW_VERSION".to_owned(), version.to_string()),
]
.into(),
files: [(
PathBuf::from("conf/agent.toml"),
"configured = true\n".to_owned(),
)]
.into(),
})
}
}
impl Drive for Scripted {
fn args(&self, _version: &Version, _settings: &Settings, prompt: &Prompt) -> Vec<String> {
vec!["--prompt".to_owned(), prompt.text.clone()]
}
fn parse(&self, _version: &Version, stdout: &str) -> Outcome {
Outcome {
text: stdout.to_owned(),
..Outcome::default()
}
}
}
fn settings() -> Settings {
Settings {
base_url: "http://gateway.test".to_owned(),
api_key: "sk-test".to_owned(),
model: "some-model".to_owned(),
wire: Wire::Messages,
}
}
fn prompt(text: &str) -> Prompt {
Prompt {
text: text.to_owned(),
allow_tools: false,
}
}
fn session(script: &str) -> (Session, tempfile::TempDir) {
let dir = tempfile::tempdir().unwrap();
let binary = dir.path().join("agent");
std::fs::write(&binary, format!("#!/bin/sh\n{script}\n")).unwrap();
std::fs::set_permissions(&binary, std::fs::Permissions::from_mode(0o755)).unwrap();
let home = dir.path().join("home");
std::fs::create_dir(&home).unwrap();
let installed = Installed {
version: Version::new(4, 5, 6),
binary,
};
(
Session::prepare(&Scripted, &installed, settings(), home).unwrap(),
dir,
)
}
const LIMIT: Duration = Duration::from_secs(20);
#[tokio::test]
async fn prepare_writes_the_config_files_under_home() {
let (_session, dir) = session("true");
let written = std::fs::read_to_string(dir.path().join("home/conf/agent.toml")).unwrap();
assert_eq!(written, "configured = true\n");
}
#[tokio::test]
async fn configure_and_drive_are_given_the_installed_version() {
let (session, _dir) = session("echo \"$AGENT_SAW_VERSION\"");
let outcome = session.run(&Scripted, &prompt("hi"), LIMIT).await.unwrap();
assert_eq!(outcome.text.trim(), "4.5.6");
}
#[tokio::test]
async fn agent_runs_in_home_with_only_its_own_environment() {
let (session, dir) = session("pwd -P; env");
let outcome = session.run(&Scripted, &prompt("hi"), LIMIT).await.unwrap();
let home = dir.path().join("home").canonicalize().unwrap();
assert_eq!(outcome.text.lines().next().unwrap(), home.to_string_lossy());
assert!(outcome.text.contains("AGENT_HOME="));
assert!(
!outcome.text.contains("CARGO_"),
"test runner environment leaked into the agent"
);
}
#[tokio::test]
async fn prompt_reaches_the_agent_as_one_untouched_argument() {
let (session, _dir) = session("printf '%s|' \"$@\"");
let text = "two spaces; $(echo injected) 'quoted'";
let outcome = session.run(&Scripted, &prompt(text), LIMIT).await.unwrap();
assert_eq!(outcome.text, format!("--prompt|{text}|"));
}
#[tokio::test]
async fn clean_exit_is_a_success() {
let (session, _dir) = session("echo done");
let outcome = session.run(&Scripted, &prompt("hi"), LIMIT).await.unwrap();
assert_eq!(outcome.exit_code, Some(0));
assert!(outcome.succeeded());
}
#[tokio::test]
async fn failing_exit_without_a_parsed_error_reports_stderr() {
let (session, _dir) = session("echo boom >&2; exit 3");
let outcome = session.run(&Scripted, &prompt("hi"), LIMIT).await.unwrap();
assert_eq!(outcome.exit_code, Some(3));
assert!(!outcome.succeeded());
assert_eq!(outcome.errors, ["boom\n"]);
}
#[tokio::test]
async fn agent_that_outlives_the_limit_is_stopped() {
let (session, _dir) = session("sleep 30");
let result = session
.run(&Scripted, &prompt("hi"), Duration::from_millis(200))
.await;
assert!(matches!(result, Err(Error::Timeout(_))));
}

View file

@ -0,0 +1,70 @@
use std::collections::HashMap;
use std::io::Write;
use std::sync::atomic::{AtomicUsize, Ordering};
use litellm_testkit::{Error, Fetch};
use sha2::{Digest, Sha256};
pub struct FakeFetch {
routes: HashMap<String, Vec<u8>>,
calls: AtomicUsize,
}
impl FakeFetch {
pub fn new(routes: impl IntoIterator<Item = (String, Vec<u8>)>) -> Self {
Self {
routes: routes.into_iter().collect(),
calls: AtomicUsize::new(0),
}
}
pub fn calls(&self) -> usize {
self.calls.load(Ordering::SeqCst)
}
}
impl Fetch for FakeFetch {
async fn get(&self, url: &str) -> Result<Vec<u8>, Error> {
self.calls.fetch_add(1, Ordering::SeqCst);
self.routes.get(url).cloned().ok_or_else(|| Error::Status {
url: url.to_owned(),
status: 404,
})
}
}
impl Fetch for &FakeFetch {
async fn get(&self, url: &str) -> Result<Vec<u8>, Error> {
(*self).get(url).await
}
}
pub fn sha256(bytes: &[u8]) -> String {
format!("{:x}", Sha256::digest(bytes))
}
pub fn script_printing(output: &str) -> Vec<u8> {
format!("#!/bin/sh\necho '{output}'\n").into_bytes()
}
pub fn tar_gz(member: &str, contents: &[u8]) -> Vec<u8> {
let mut builder = tar::Builder::new(Vec::new());
let mut header = tar::Header::new_gnu();
header.set_size(contents.len() as u64);
header.set_mode(0o755);
header.set_cksum();
builder.append_data(&mut header, member, contents).unwrap();
let tarball = builder.into_inner().unwrap();
let mut encoder = flate2::write::GzEncoder::new(Vec::new(), flate2::Compression::default());
encoder.write_all(&tarball).unwrap();
encoder.finish().unwrap()
}
pub fn zip_archive(member: &str, contents: &[u8]) -> Vec<u8> {
let mut writer = zip::ZipWriter::new(std::io::Cursor::new(Vec::new()));
writer
.start_file(member, zip::write::SimpleFileOptions::default())
.unwrap();
writer.write_all(contents).unwrap();
writer.finish().unwrap().into_inner()
}

View file

@ -3,6 +3,8 @@ from collections.abc import AsyncIterator, Awaitable, Callable, Coroutine, Itera
from types import MappingProxyType
from typing import Final, TypeAlias, cast # noqa: TID251 # native binding selects a sync result or an async awaitable
from litellm.exceptions import BadRequestError
from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
from litellm.llms.anthropic.experimental_pass_through.messages import handler as main
from litellm.rust_bridge.catalog import Delivery, Route, RouteContext
from litellm.rust_bridge.dispatch import PublicDispatch, call_hook
@ -71,10 +73,17 @@ def _public_request(
)
def _resolved_provider(request: LiteLLMMessagesRequest) -> str | None:
try:
return get_llm_provider(request.model, request.custom_llm_provider)[1]
except BadRequestError:
return request.custom_llm_provider
def _context(request: LiteLLMMessagesRequest) -> RouteContext:
return RouteContext(
Route.MESSAGES,
provider=request.custom_llm_provider,
provider=_resolved_provider(request),
model=request.model,
delivery=Delivery.STREAMING if request.stream else Delivery.COMPLETED,
)

View file

@ -109,6 +109,7 @@ RULES: Final[Rules] = (
RouteRule(Route.CHAT_COMPLETIONS, Rollout.PYTHON_ONLY),
RouteRule(Route.EMBEDDINGS, Rollout.PYTHON_ONLY),
RouteRule(Route.OCR, Rollout.RUST_REQUIRED),
RouteRule(Route.MESSAGES, Rollout.PYTHON_ONLY, providers=frozenset({"anthropic"})),
RouteRule(Route.MESSAGES, Rollout.PYTHON_ONLY),
RouteRule(Route.RESPONSES, Rollout.PYTHON_ONLY),
RouteRule(Route.TOKEN_COUNTER, Rollout.PYTHON_ONLY),

View file

@ -1,6 +1,6 @@
from __future__ import annotations
from collections.abc import AsyncIterator, Awaitable, Iterator
from collections.abc import AsyncIterator, Awaitable, Iterator, Mapping
from dataclasses import dataclass
from typing import Final, Protocol
@ -17,7 +17,7 @@ class Complete:
@dataclass(frozen=True, slots=True)
class Open:
value: None
value: Mapping[str, object] | None
@dataclass(frozen=True, slots=True)
@ -68,7 +68,7 @@ async def drive(execution: Execution) -> object:
step: Final = await _settle(execution, execution.start())
if isinstance(step, Open):
handed_off = True
return Stream(execution)
return Stream(execution, step.value)
return step.value
finally:
if not handed_off:
@ -78,10 +78,10 @@ async def drive(execution: Execution) -> object:
class Stream(AsyncIterator[object]):
"""A streamed native call: each read resumes the execution until its next chunk."""
def __init__(self, execution: Execution) -> None:
def __init__(self, execution: Execution, hidden_params: Mapping[str, object] | None = None) -> None:
self._execution: Final = execution
self._done = False
self._hidden_params: dict[str, object] = {} # mutable-ok: header writers mutate _hidden_params in place
self._hidden_params: dict[str, object] = dict(hidden_params or {}) # mutable-ok: header writers mutate it
def __aiter__(self) -> Stream:
return self
@ -115,10 +115,10 @@ class Stream(AsyncIterator[object]):
class SyncStream(Iterator[object]):
"""The sync form of `Stream`; its execution never suspends on an awaitable."""
def __init__(self, execution: Execution) -> None:
def __init__(self, execution: Execution, hidden_params: Mapping[str, object] | None = None) -> None:
self._execution: Final = execution
self._done = False
self._hidden_params: dict[str, object] = {} # mutable-ok: header writers mutate _hidden_params in place
self._hidden_params: dict[str, object] = dict(hidden_params or {}) # mutable-ok: header writers mutate it
def __iter__(self) -> SyncStream:
return self

View file

@ -4,6 +4,7 @@ from collections.abc import Mapping, Sequence
from dataclasses import asdict, dataclass
from typing import Final, cast # noqa: TID251 # narrows the normalized native payload to the public TypedDict
import httpx
from pydantic import TypeAdapter, ValidationError
import litellm
@ -53,6 +54,14 @@ def response(value: Mapping[str, object]) -> AnthropicMessagesResponse:
)
def stream_hidden_params(headers: Sequence[tuple[str, str]]) -> Mapping[str, object]:
from litellm.llms.anthropic.experimental_pass_through.messages.streaming_iterator import (
anthropic_messages_stream_hidden_params,
)
return anthropic_messages_stream_hidden_params(httpx.Headers(list(headers)))
def arguments(request: LiteLLMMessagesRequest) -> Mapping[str, object]:
return request.kwargs

View file

@ -110,3 +110,15 @@ def test_native_request_rejections_map_to_the_public_400() -> None:
assert "does not support top_k=5" in mapped.message
assert mapped.model == "claude-sonnet-5"
assert not isinstance(route_host.map_failure(ValueError("plain"), request, "anthropic"), litellm.BadRequestError)
def test_stream_hidden_params_projects_upstream_headers_the_way_the_python_handler_does() -> None:
hidden: Final = route_host.stream_hidden_params(
(("request-id", "req_upstream_123"), ("x-ratelimit-remaining-requests", "41"))
)
additional: Final = hidden["additional_headers"]
assert isinstance(additional, dict)
assert additional["llm_provider-request-id"] == "req_upstream_123"
assert additional["x-ratelimit-remaining-requests"] == "41"
assert "request-id" not in additional

View file

@ -127,7 +127,7 @@ async def test_native_messages_stream_relays_provider_events_and_logs_success_on
**arguments(messages_server, stream=True, callbacks=[recorder])
)
assert isinstance(stream, AsyncIterator)
assert get_hidden_params_dict(stream) == {"additional_headers": {"x-litellm-rust": "true"}}
assert get_hidden_params_dict(stream)["additional_headers"]["x-litellm-rust"] == "true"
first: Final = await anext(stream)
await drain_logging()
assert "async_log_success_event" not in recorder.names
@ -171,7 +171,7 @@ def test_native_sync_messages_stream_relays_provider_events_and_logs_success_onc
stream: Final = litellm.anthropic.messages.create(**arguments(messages_server, stream=True, callbacks=[recorder]))
assert isinstance(stream, Iterator)
assert get_hidden_params_dict(stream) == {"additional_headers": {"x-litellm-rust": "true"}}
assert get_hidden_params_dict(stream)["additional_headers"]["x-litellm-rust"] == "true"
assert b"".join(stream) == sse_payload()
assert_served_natively(messages_server)
@ -186,3 +186,56 @@ def test_native_sync_messages_returns_the_provider_message(messages_server: Reco
assert_served_natively(messages_server)
assert response["content"] == MESSAGES_RESPONSE["content"]
assert len(recorder.wait_for("log_success_event")) == 1
@pytest.mark.asyncio
async def test_native_messages_pre_call_sees_the_shaped_optional_params(
messages_server: RecordingServer,
) -> None:
recorder: Final = RecordingLogger()
await litellm.anthropic.messages.acreate(
**arguments(messages_server, callbacks=[recorder], temperature=0.2, top_k=3, drop_params=True)
)
sent: Final = messages_server.requests[0].body
assert not {"temperature", "top_k"} & sent.keys()
pre_call: Final = recorder.wait_for("log_pre_api_call")[0].kwargs
assert isinstance(pre_call, dict)
optional_params: Final = pre_call["optional_params"]
assert isinstance(optional_params, dict)
assert not {"model", "messages", "temperature", "top_k"} & optional_params.keys()
assert optional_params["max_tokens"] == sent["max_tokens"]
@pytest.mark.asyncio
async def test_native_messages_failing_pre_call_logger_does_not_fail_the_call(messages_server: RecordingServer) -> None:
class Broken(CustomLogger):
def log_pre_api_call(self, model, messages, kwargs):
raise RuntimeError("logger exploded")
response: Final = await litellm.anthropic.messages.acreate(**arguments(messages_server, callbacks=[Broken()]))
assert_served_natively(messages_server)
assert response["content"] == MESSAGES_RESPONSE["content"]
@pytest.mark.asyncio
async def test_native_messages_stream_success_log_carries_usage_rebuilt_from_the_relayed_events(
messages_server: RecordingServer,
) -> None:
messages_server.enqueue(STREAM)
recorder: Final = RecordingLogger()
stream: Final = await litellm.anthropic.messages.acreate(
**arguments(messages_server, stream=True, callbacks=[recorder])
)
assert isinstance(stream, AsyncIterator)
async for _ in stream:
pass
success: Final = await recorder.wait_for_async("async_log_success_event")
usage: Final = success[0].response.usage
assert usage.completion_tokens == MESSAGES_EVENTS[4][1]["usage"]["output_tokens"]
assert usage.prompt_tokens == MESSAGES_RESPONSE["usage"]["input_tokens"]
assert success[0].response.choices[0].message.content == "Hello from native Messages"