diff --git a/litellm-rust/Cargo.lock b/litellm-rust/Cargo.lock index 2d96efa6077..e7d911f5fd9 100644 --- a/litellm-rust/Cargo.lock +++ b/litellm-rust/Cargo.lock @@ -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", +] diff --git a/litellm-rust/Cargo.toml b/litellm-rust/Cargo.toml index 0c7236e807e..022e8f13311 100644 --- a/litellm-rust/Cargo.toml +++ b/litellm-rust/Cargo.toml @@ -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" diff --git a/litellm-rust/crates/core/src/messages/handler.rs b/litellm-rust/crates/core/src/messages/handler.rs index fe7e8bb4b80..de1a5f476ed 100644 --- a/litellm-rust/crates/core/src/messages/handler.rs +++ b/litellm-rust/crates/core/src/messages/handler.rs @@ -17,8 +17,10 @@ pub(super) async fn send( body: &Value, timeout: Option, ) -> Result { + 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 { diff --git a/litellm-rust/crates/core/src/messages/route.rs b/litellm-rust/crates/core/src/messages/route.rs index fc1a9b63252..40aff185e81 100644 --- a/litellm-rust/crates/core/src/messages/route.rs +++ b/litellm-rust/crates/core/src/messages/route.rs @@ -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 for Error { @@ -77,19 +80,6 @@ impl From for Error { pub type MessagesHost = HostChannel; pub type MessagesMachine = CallMachine; -/// 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 { - 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)? { diff --git a/litellm-rust/crates/core/tests/messages/host.rs b/litellm-rust/crates/core/tests/messages/host.rs new file mode 100644 index 00000000000..ca2aece5ebd --- /dev/null +++ b/litellm-rust/crates/core/tests/messages/host.rs @@ -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 Result + 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>, + optional_params: Mutex>, +} + +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 { + self.events + .lock() + .unwrap() + .iter() + .filter_map(|event| match event { + CallEvent::Machine(MachineEvent::ResponseReceived { raw }) => { + Some(raw.body.clone()) + } + _ => None, + }) + .collect() + } +} + +impl Host for RecordingHost { + async fn project(&self) -> Result { + 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 { + 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 { + 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::(&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 = 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})); +} diff --git a/litellm-rust/crates/core/tests/messages/main.rs b/litellm-rust/crates/core/tests/messages/main.rs index 4e549bae309..21ee678ced3 100644 --- a/litellm-rust/crates/core/tests/messages/main.rs +++ b/litellm-rust/crates/core/tests/messages/main.rs @@ -14,6 +14,7 @@ use wiremock::ResponseTemplate; mod support; use support::*; +mod host; mod request; mod response; mod secrets; diff --git a/litellm-rust/crates/core/tests/messages/request.rs b/litellm-rust/crates/core/tests/messages/request.rs index 9353324d370..2927356b773 100644 --- a/litellm-rust/crates/core/tests/messages/request.rs +++ b/litellm-rust/crates/core/tests/messages/request.rs @@ -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 = call.body.into_iter().chain(object(fields)).collect(); + MessagesCall { body, ..call } +} + +fn sent_betas(request: &wiremock::Request) -> Vec { + 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 = 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 = ["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); +} diff --git a/litellm-rust/crates/core/tests/messages/response.rs b/litellm-rust/crates/core/tests/messages/response.rs index 38a18c415ba..133b7d2b162 100644 --- a/litellm-rust/crates/core/tests/messages/response.rs +++ b/litellm-rust/crates/core/tests/messages/response.rs @@ -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::(&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)] diff --git a/litellm-rust/crates/core/tests/messages/secrets.rs b/litellm-rust/crates/core/tests/messages/secrets.rs index 419b6d6c753..55e510d00d3 100644 --- a/litellm-rust/crates/core/tests/messages/secrets.rs +++ b/litellm-rust/crates/core/tests/messages/secrets.rs @@ -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 { + 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)); +} diff --git a/litellm-rust/crates/core/tests/messages/stream.rs b/litellm-rust/crates/core/tests/messages/stream.rs index ea23a9e8e38..c4be3127d66 100644 --- a/litellm-rust/crates/core/tests/messages/stream.rs +++ b/litellm-rust/crates/core/tests/messages/stream.rs @@ -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 for RecordingStreamHost { match op {} } - async fn open(&self, (): ()) -> Result { - Ok(self.record(Seen::Open)) + async fn open(&self, head: MessagesStreamHead) -> Result { + Ok(self.record(Seen::Open(head.headers))) } async fn deliver(&self, chunk: Bytes) -> Result { @@ -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 { @@ -80,7 +92,7 @@ async fn stream_through(host: &RecordingStreamHost) -> Result = 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 = 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 = 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) { diff --git a/litellm-rust/crates/host-python/src/adapter.rs b/litellm-rust/crates/host-python/src/adapter.rs index 87481aa89b7..7f07475bc4c 100644 --- a/litellm-rust/crates/host-python/src/adapter.rs +++ b/litellm-rust/crates/host-python/src/adapter.rs @@ -134,6 +134,13 @@ pub trait ProtocolHost: Send + Sync { response: ::Response, ) -> PyResult>; + /// What the stream carries at hand-off, as the caller's stream receives it. + fn head( + &mut self, + py: Python<'_>, + head: ::StreamHead, + ) -> PyResult>; + /// One streamed chunk as the caller receives it. fn chunk( &mut self, diff --git a/litellm-rust/crates/host-python/src/driver.rs b/litellm-rust/crates/host-python/src/driver.rs index aaa0752522b..372af2843bd 100644 --- a/litellm-rust/crates/host-python/src/driver.rs +++ b/litellm-rust/crates/host-python/src/driver.rs @@ -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) -> PyResult { + fn opened( + &mut self, + py: Python<'_>, + head: as Protocol>::StreamHead, + reply: Reply, + ) -> PyResult { 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> { + match head {} + } + fn chunk(&mut self, _: Python<'_>, chunk: std::convert::Infallible) -> PyResult> { 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> { + Ok(()) + } + + fn invoke( + &mut self, + _: Python<'_>, + op: std::convert::Infallible, + ) -> Result<(), InvokeError> { + match op {} + } + + fn head( + &mut self, + py: Python<'_>, + head: Vec<(&'static str, &'static str)>, + ) -> PyResult> { + 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> { + Ok(pyo3::types::PyString::new(py, chunk).into_any().unbind()) + } + + fn complete(&mut self, py: Python<'_>, (): ()) -> PyResult> { + Ok(py.None()) + } + + fn classify(&self, _: Python<'_>, error: Error) -> PyResult { + 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 { + 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 { + 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::(py) { + return None; + } + assert!(stop.is_instance_of::(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, + > = 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 { CallMachine::new(|host| { Box::pin(async move { @@ -1202,6 +1372,13 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri ) -> Result<(), InvokeError> { Err(missing_state().into()) } + fn head( + &mut self, + _: Python<'_>, + head: std::convert::Infallible, + ) -> PyResult> { + match head {} + } fn chunk( &mut self, _: Python<'_>, diff --git a/litellm-rust/crates/host-python/src/handle.rs b/litellm-rust/crates/host-python/src/handle.rs index 10abbadbda5..24adfd404d7 100644 --- a/litellm-rust/crates/host-python/src/handle.rs +++ b/litellm-rust/crates/host-python/src/handle.rs @@ -8,9 +8,9 @@ use pyo3::prelude::*; pub enum ExecutionStep { Return(Py), Await(Py), - /// 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), Yield(Py), } @@ -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), }; diff --git a/litellm-rust/crates/python-bridge/src/routes/messages/host.rs b/litellm-rust/crates/python-bridge/src/routes/messages/host.rs index 6de4e1320e1..a253f4f5670 100644 --- a/litellm-rust/crates/python-bridge/src/routes/messages/host.rs +++ b/litellm-rust/crates/python-bridge/src/routes/messages/host.rs @@ -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.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> { Ok(PyBytes::new(py, &chunk).into_any().unbind()) } diff --git a/litellm-rust/crates/python-bridge/src/routes/messages/mod.rs b/litellm-rust/crates/python-bridge/src/routes/messages/mod.rs index 65040f31684..dae8623979a 100644 --- a/litellm-rust/crates/python-bridge/src/routes/messages/mod.rs +++ b/litellm-rust/crates/python-bridge/src/routes/messages/mod.rs @@ -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> { - let model: String = request.getattr("model")?.extract()?; - let provider: Option = request.getattr("custom_llm_provider")?.extract()?; - let stream = request - .getattr("stream")? - .extract::>()? - .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, diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/host.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/host.rs index 5a3806e61e3..dc01ced15a0 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/host.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/host.rs @@ -117,6 +117,10 @@ impl ProtocolHost for OcrPythonHost { .map(Bound::unbind) } + fn head(&mut self, _: Python<'_>, head: std::convert::Infallible) -> PyResult> { + match head {} + } + fn chunk(&mut self, _: Python<'_>, chunk: std::convert::Infallible) -> PyResult> { match chunk {} } diff --git a/litellm-rust/crates/testkit/Cargo.toml b/litellm-rust/crates/testkit/Cargo.toml new file mode 100644 index 00000000000..98a36a1e87f --- /dev/null +++ b/litellm-rust/crates/testkit/Cargo.toml @@ -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 diff --git a/litellm-rust/crates/testkit/src/agent/claude.rs b/litellm-rust/crates/testkit/src/agent/claude.rs new file mode 100644 index 00000000000..6870fff6bf8 --- /dev/null +++ b/litellm-rust/crates/testkit/src/agent/claude.rs @@ -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, +} + +#[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 { + 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 { + 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, +} + +#[derive(Deserialize)] +struct Block { + #[serde(rename = "type")] + kind: String, + name: Option, +} + +#[derive(Deserialize)] +struct Finished { + is_error: bool, + result: Option, + usage: Option, +} + +#[derive(Deserialize)] +struct TokenUsage { + input_tokens: u64, + output_tokens: u64, +} + +impl Drive for ClaudeCode { + fn args(&self, _version: &Version, settings: &Settings, prompt: &Prompt) -> Vec { + 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 = 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, + } + } +} diff --git a/litellm-rust/crates/testkit/src/agent/codex.rs b/litellm-rust/crates/testkit/src/agent/codex.rs new file mode 100644 index 00000000000..6749e81a471 --- /dev/null +++ b/litellm-rust/crates/testkit/src/agent/codex.rs @@ -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 { + 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 { + 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, + usage: Option, + error: Option, +} + +#[derive(Deserialize)] +struct Item { + #[serde(rename = "type")] + kind: String, + text: Option, +} + +#[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 { + 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 = 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, + } + } +} diff --git a/litellm-rust/crates/testkit/src/agent/configure.rs b/litellm-rust/crates/testkit/src/agent/configure.rs new file mode 100644 index 00000000000..a095aacc92f --- /dev/null +++ b/litellm-rust/crates/testkit/src/agent/configure.rs @@ -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, + pub files: BTreeMap, +} + +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; +} + +pub(crate) fn env( + pairs: impl IntoIterator, +) -> BTreeMap { + 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('/')) +} diff --git a/litellm-rust/crates/testkit/src/agent/drive.rs b/litellm-rust/crates/testkit/src/agent/drive.rs new file mode 100644 index 00000000000..2c238843ed7 --- /dev/null +++ b/litellm-rust/crates/testkit/src/agent/drive.rs @@ -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, + pub usage: Usage, + pub errors: Vec, + pub exit_code: Option, +} + +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; + + fn parse(&self, version: &Version, stdout: &str) -> Outcome; +} + +pub(crate) fn json_lines<'a, T: serde::de::DeserializeOwned + 'a>( + stdout: &'a str, +) -> impl Iterator + 'a { + stdout + .lines() + .filter_map(|line| serde_json::from_str(line).ok()) +} diff --git a/litellm-rust/crates/testkit/src/agent/install.rs b/litellm-rust/crates/testkit/src/agent/install.rs new file mode 100644 index 00000000000..3f1a0f8950d --- /dev/null +++ b/litellm-rust/crates/testkit/src/agent/install.rs @@ -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> + Send; +} diff --git a/litellm-rust/crates/testkit/src/agent/mod.rs b/litellm-rust/crates/testkit/src/agent/mod.rs new file mode 100644 index 00000000000..03a70583f66 --- /dev/null +++ b/litellm-rust/crates/testkit/src/agent/mod.rs @@ -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 Agent for T {} diff --git a/litellm-rust/crates/testkit/src/agent/opencode.rs b/litellm-rust/crates/testkit/src/agent/opencode.rs new file mode 100644 index 00000000000..a1b2fa01eef --- /dev/null +++ b/litellm-rust/crates/testkit/src/agent/opencode.rs @@ -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 { + 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 { + 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, +} + +#[derive(Deserialize)] +struct FailureData { + message: Option, +} + +impl Drive for Opencode { + fn args(&self, _version: &Version, _settings: &Settings, prompt: &Prompt) -> Vec { + ["run", "--format", "json", &prompt.text] + .map(str::to_owned) + .to_vec() + } + + fn parse(&self, _version: &Version, stdout: &str) -> Outcome { + let events: Vec = 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, + } + } +} diff --git a/litellm-rust/crates/testkit/src/error.rs b/litellm-rust/crates/testkit/src/error.rs new file mode 100644 index 00000000000..03520d827f1 --- /dev/null +++ b/litellm-rust/crates/testkit/src/error.rs @@ -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), +} diff --git a/litellm-rust/crates/testkit/src/install/archive.rs b/litellm-rust/crates/testkit/src/install/archive.rs new file mode 100644 index 00000000000..c8d08f66f47 --- /dev/null +++ b/litellm-rust/crates/testkit/src/install/archive.rs @@ -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, 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, 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, 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) +} diff --git a/litellm-rust/crates/testkit/src/install/fetch.rs b/litellm-rust/crates/testkit/src/install/fetch.rs new file mode 100644 index 00000000000..73008f7a0da --- /dev/null +++ b/litellm-rust/crates/testkit/src/install/fetch.rs @@ -0,0 +1,55 @@ +use std::future::Future; + +use crate::Error; + +pub trait Fetch: Sync { + fn get(&self, url: &str) -> impl Future, Error>> + Send; +} + +pub struct HttpFetch { + client: reqwest::Client, + github_token: Option, +} + +impl HttpFetch { + pub fn new(github_token: Option) -> 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, 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()) + } +} diff --git a/litellm-rust/crates/testkit/src/install/mod.rs b/litellm-rust/crates/testkit/src/install/mod.rs new file mode 100644 index 00000000000..1104bcec102 --- /dev/null +++ b/litellm-rust/crates/testkit/src/install/mod.rs @@ -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 { + fetch: F, + cache_root: PathBuf, + target: Target, +} + +impl Installer { + pub fn new(fetch: F, cache_root: impl Into, target: Target) -> Self { + Self { + fetch, + cache_root: cache_root.into(), + target, + } + } + + pub async fn install( + &self, + agent: &impl Install, + version: &Version, + ) -> Result { + 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}; diff --git a/litellm-rust/crates/testkit/src/install/release.rs b/litellm-rust/crates/testkit/src/install/release.rs new file mode 100644 index 00000000000..a21b9a14f1b --- /dev/null +++ b/litellm-rust/crates/testkit/src/install/release.rs @@ -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, +} + +#[derive(Deserialize)] +struct GithubAsset { + name: String, + digest: Option, + browser_download_url: String, +} + +pub(crate) async fn github_release( + fetch: &impl Fetch, + releases_url: &str, + tag: &str, + asset_name: &str, + packaging: Packaging, +) -> Result { + 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 Deserialize<'de>>(url: &str, body: &[u8]) -> Result { + serde_json::from_slice(body).map_err(|source| Error::Metadata { + url: url.to_owned(), + source, + }) +} diff --git a/litellm-rust/crates/testkit/src/lib.rs b/litellm-rust/crates/testkit/src/lib.rs new file mode 100644 index 00000000000..9ea6123a176 --- /dev/null +++ b/litellm-rust/crates/testkit/src/lib.rs @@ -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}; diff --git a/litellm-rust/crates/testkit/src/session.rs b/litellm-rust/crates/testkit/src/session.rs new file mode 100644 index 00000000000..6b06e756cca --- /dev/null +++ b/litellm-rust/crates/testkit/src/session.rs @@ -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, +} + +impl Session { + pub fn prepare( + agent: &impl Configure, + installed: &Installed, + settings: Settings, + home: impl Into, + ) -> Result { + 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 { + 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 + }) + } +} diff --git a/litellm-rust/crates/testkit/src/target.rs b/litellm-rust/crates/testkit/src/target.rs new file mode 100644 index 00000000000..a9d4b012d52 --- /dev/null +++ b/litellm-rust/crates/testkit/src/target.rs @@ -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::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 { + 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, + }) + } +} diff --git a/litellm-rust/crates/testkit/tests/configure.rs b/litellm-rust/crates/testkit/tests/configure.rs new file mode 100644 index 00000000000..ca3587c3474 --- /dev/null +++ b/litellm-rust/crates/testkit/tests/configure.rs @@ -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::>() + .len(), + packages.len() + ); +} diff --git a/litellm-rust/crates/testkit/tests/install.rs b/litellm-rust/crates/testkit/tests/install.rs new file mode 100644 index 00000000000..7edaf27de6c --- /dev/null +++ b/litellm-rust/crates/testkit/tests/install.rs @@ -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) -> Vec { + 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)> { + 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, digest: Option) -> Vec<(String, Vec)> { + 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 + ); +} diff --git a/litellm-rust/crates/testkit/tests/live.rs b/litellm-rust/crates/testkit/tests/live.rs new file mode 100644 index 00000000000..b805596879a --- /dev/null +++ b/litellm-rust/crates/testkit/tests/live.rs @@ -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_` per wire +//! (`MESSAGES`, `RESPONSES`, `CHAT_COMPLETIONS`) and one `TESTKIT__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:?}"); +} diff --git a/litellm-rust/crates/testkit/tests/session.rs b/litellm-rust/crates/testkit/tests/session.rs new file mode 100644 index 00000000000..cd5e0dcbc71 --- /dev/null +++ b/litellm-rust/crates/testkit/tests/session.rs @@ -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 { + 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 { + 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(_)))); +} diff --git a/litellm-rust/crates/testkit/tests/support/mod.rs b/litellm-rust/crates/testkit/tests/support/mod.rs new file mode 100644 index 00000000000..f4a13759941 --- /dev/null +++ b/litellm-rust/crates/testkit/tests/support/mod.rs @@ -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>, + calls: AtomicUsize, +} + +impl FakeFetch { + pub fn new(routes: impl IntoIterator)>) -> 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, 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, Error> { + (*self).get(url).await + } +} + +pub fn sha256(bytes: &[u8]) -> String { + format!("{:x}", Sha256::digest(bytes)) +} + +pub fn script_printing(output: &str) -> Vec { + format!("#!/bin/sh\necho '{output}'\n").into_bytes() +} + +pub fn tar_gz(member: &str, contents: &[u8]) -> Vec { + 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 { + 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() +} diff --git a/litellm/messages/dispatch.py b/litellm/messages/dispatch.py index a0c791a136c..13a030e7ebe 100644 --- a/litellm/messages/dispatch.py +++ b/litellm/messages/dispatch.py @@ -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, ) diff --git a/litellm/rust_bridge/catalog.py b/litellm/rust_bridge/catalog.py index d9834adc7e8..6e455817194 100644 --- a/litellm/rust_bridge/catalog.py +++ b/litellm/rust_bridge/catalog.py @@ -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), diff --git a/litellm/rust_bridge/lifecycle.py b/litellm/rust_bridge/lifecycle.py index 2f243e8c212..7a6485a5f2c 100644 --- a/litellm/rust_bridge/lifecycle.py +++ b/litellm/rust_bridge/lifecycle.py @@ -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 diff --git a/litellm/rust_bridge/messages/route_host.py b/litellm/rust_bridge/messages/route_host.py index d49d7b75a6f..0a23989a59c 100644 --- a/litellm/rust_bridge/messages/route_host.py +++ b/litellm/rust_bridge/messages/route_host.py @@ -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 diff --git a/tests/test_litellm/rust_bridge/messages/test_route_host.py b/tests/test_litellm/rust_bridge/messages/test_route_host.py index f47333a45d9..c5a442e0709 100644 --- a/tests/test_litellm/rust_bridge/messages/test_route_host.py +++ b/tests/test_litellm/rust_bridge/messages/test_route_host.py @@ -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 diff --git a/tests/test_litellm_rust/messages/test_callbacks.py b/tests/test_litellm_rust/messages/test_callbacks.py index 19043780eb6..dc66852d214 100644 --- a/tests/test_litellm_rust/messages/test_callbacks.py +++ b/tests/test_litellm_rust/messages/test_callbacks.py @@ -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"