From 081f73f021620bce438fb86a71435ec34a707d54 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Fri, 25 Sep 2026 10:03:15 -0700 Subject: [PATCH 1/2] feat(rust): hand upstream response headers to the native Messages stream (#43178) * ci: drop the ocr_testing job now that tests/ocr_tests is gone Co-Authored-By: Claude Opus 5.5 * test(ocr): restore the live OCR matrix and the ocr_testing job The public litellm.ocr / aocr / Router interface is unchanged by the Rust migration, so the live provider matrix still applies. Drops the stale VCR skip list for the deleted test_rust_bridge.py. Co-Authored-By: Claude Opus 5.5 * test(messages): show streamed upstream headers never reach the native stream The Python handler puts the upstream response headers on the stream's _hidden_params before the first chunk so the proxy can forward them as llm_provider-* headers. The native route drops them, and this test fails on the Rust path while passing on Python. Co-Authored-By: Claude Fable 5.1 * feat(messages): hand upstream response headers to the native stream before its first chunk The Messages route fills MessagesStreamHead from the upstream response and yields it on Open. The Python driver converts it through the protocol host and hands it to Stream and SyncStream as their _hidden_params, so a streamed native call carries additional_headers the same way the Python handler does and the proxy can forward them as llm_provider-* headers. The relay contract lives in the core crate test, the hand-off in the host-python driver test, and the header projection in the route host test, so the recording-server test that showed the gap is dropped. Co-Authored-By: Claude Fable 5.1 * wip --------- Co-authored-by: Yujong Lee Co-authored-by: Claude Opus 5.5 --- .../crates/core/src/messages/handler.rs | 4 +- .../crates/core/src/messages/route.rs | 38 +- .../crates/core/tests/messages/host.rs | 210 +++++++++ .../crates/core/tests/messages/main.rs | 1 + .../crates/core/tests/messages/request.rs | 436 +++++++++++++++++- .../crates/core/tests/messages/response.rs | 87 +++- .../crates/core/tests/messages/secrets.rs | 124 ++++- .../crates/core/tests/messages/stream.rs | 135 +++++- .../crates/host-python/src/adapter.rs | 7 + litellm-rust/crates/host-python/src/driver.rs | 187 +++++++- litellm-rust/crates/host-python/src/handle.rs | 8 +- .../python-bridge/src/routes/messages/host.rs | 9 +- .../python-bridge/src/routes/messages/mod.rs | 15 +- .../python-bridge/src/routes/ocr/host.rs | 4 + litellm/messages/dispatch.py | 11 +- litellm/rust_bridge/catalog.py | 1 + litellm/rust_bridge/lifecycle.py | 14 +- litellm/rust_bridge/messages/route_host.py | 9 + .../rust_bridge/messages/test_route_host.py | 12 + .../messages/test_callbacks.py | 57 ++- 20 files changed, 1274 insertions(+), 95 deletions(-) create mode 100644 litellm-rust/crates/core/tests/messages/host.rs 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/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" From 9c10e0f985787875ddb36818959b0667f41c9cf6 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Fri, 25 Sep 2026 17:27:30 +0000 Subject: [PATCH 2/2] feat(testkit): agent clients for Claude Code, Codex and opencode (#43181) * feat(testkit): install and configure Claude Code, Codex and opencode against a gateway Co-Authored-By: Claude Sonnet 5 * refactor(testkit): derive targets from target-lexicon and split agents behind a trait Co-Authored-By: Claude Sonnet 5 * refactor(testkit): group sources into agent and install folders Co-Authored-By: Claude Sonnet 5 * feat(testkit): split agents into install, configure and drive with semver-aware launch Co-Authored-By: Claude Sonnet 5 * style(testkit): drop a needless borrow Co-Authored-By: Claude Sonnet 5 --------- Co-authored-by: Yujong Lee Co-authored-by: Claude Sonnet 5 --- litellm-rust/Cargo.lock | 152 +++++++++- litellm-rust/Cargo.toml | 6 + litellm-rust/crates/testkit/Cargo.toml | 32 +++ .../crates/testkit/src/agent/claude.rs | 181 ++++++++++++ .../crates/testkit/src/agent/codex.rs | 174 ++++++++++++ .../crates/testkit/src/agent/configure.rs | 69 +++++ .../crates/testkit/src/agent/drive.rs | 57 ++++ .../crates/testkit/src/agent/install.rs | 17 ++ litellm-rust/crates/testkit/src/agent/mod.rs | 20 ++ .../crates/testkit/src/agent/opencode.rs | 187 +++++++++++++ litellm-rust/crates/testkit/src/error.rs | 56 ++++ .../crates/testkit/src/install/archive.rs | 52 ++++ .../crates/testkit/src/install/fetch.rs | 55 ++++ .../crates/testkit/src/install/mod.rs | 118 ++++++++ .../crates/testkit/src/install/release.rs | 65 +++++ litellm-rust/crates/testkit/src/lib.rs | 15 + litellm-rust/crates/testkit/src/session.rs | 76 +++++ litellm-rust/crates/testkit/src/target.rs | 69 +++++ .../crates/testkit/tests/configure.rs | 133 +++++++++ litellm-rust/crates/testkit/tests/install.rs | 262 ++++++++++++++++++ litellm-rust/crates/testkit/tests/live.rs | 133 +++++++++ litellm-rust/crates/testkit/tests/session.rs | 155 +++++++++++ .../crates/testkit/tests/support/mod.rs | 70 +++++ 23 files changed, 2151 insertions(+), 3 deletions(-) create mode 100644 litellm-rust/crates/testkit/Cargo.toml create mode 100644 litellm-rust/crates/testkit/src/agent/claude.rs create mode 100644 litellm-rust/crates/testkit/src/agent/codex.rs create mode 100644 litellm-rust/crates/testkit/src/agent/configure.rs create mode 100644 litellm-rust/crates/testkit/src/agent/drive.rs create mode 100644 litellm-rust/crates/testkit/src/agent/install.rs create mode 100644 litellm-rust/crates/testkit/src/agent/mod.rs create mode 100644 litellm-rust/crates/testkit/src/agent/opencode.rs create mode 100644 litellm-rust/crates/testkit/src/error.rs create mode 100644 litellm-rust/crates/testkit/src/install/archive.rs create mode 100644 litellm-rust/crates/testkit/src/install/fetch.rs create mode 100644 litellm-rust/crates/testkit/src/install/mod.rs create mode 100644 litellm-rust/crates/testkit/src/install/release.rs create mode 100644 litellm-rust/crates/testkit/src/lib.rs create mode 100644 litellm-rust/crates/testkit/src/session.rs create mode 100644 litellm-rust/crates/testkit/src/target.rs create mode 100644 litellm-rust/crates/testkit/tests/configure.rs create mode 100644 litellm-rust/crates/testkit/tests/install.rs create mode 100644 litellm-rust/crates/testkit/tests/live.rs create mode 100644 litellm-rust/crates/testkit/tests/session.rs create mode 100644 litellm-rust/crates/testkit/tests/support/mod.rs 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/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() +}