mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
Merge remote-tracking branch 'origin/main' into litellm_agent365_fail_open_default
This commit is contained in:
commit
6e03a36e16
43 changed files with 3425 additions and 98 deletions
152
litellm-rust/Cargo.lock
generated
152
litellm-rust/Cargo.lock
generated
|
|
@ -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",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -17,8 +17,10 @@ pub(super) async fn send(
|
|||
body: &Value,
|
||||
timeout: Option<Duration>,
|
||||
) -> Result<reqwest::Response, Error> {
|
||||
let encoded = serde_json::to_vec(body)
|
||||
.map_err(|err| Error::InvalidRequest(format!("failed to encode messages body: {err}")))?;
|
||||
let builder = headers.iter().fold(
|
||||
http_client().post(url).json(body),
|
||||
http_client().post(url).body(encoded),
|
||||
|builder, (key, value)| builder.header(key, value),
|
||||
);
|
||||
let builder = match timeout {
|
||||
|
|
|
|||
|
|
@ -6,7 +6,6 @@ use std::{
|
|||
|
||||
use bytes::Bytes;
|
||||
use litellm_auth::SecretValue;
|
||||
use litellm_core_utils::get_llm_provider_logic::get_custom_llm_provider;
|
||||
use litellm_host::{
|
||||
event::{MachineEvent, RawResponse, RequestContext, WireRequest},
|
||||
host::{Demand, Host},
|
||||
|
|
@ -22,7 +21,6 @@ use serde_json::{Map, Value};
|
|||
|
||||
use super::{
|
||||
Error,
|
||||
common_utils::messages_provider_config,
|
||||
handler::{decode_response, network, provider_error, send},
|
||||
prepare::{prepare_provider_request, resolve_provider},
|
||||
types::{MessagesRequest, MessagesShaping},
|
||||
|
|
@ -54,6 +52,11 @@ pub enum MessagesOutput {
|
|||
Streamed,
|
||||
}
|
||||
|
||||
/// The upstream response as the caller sees it at stream hand-off, before any chunk.
|
||||
pub struct MessagesStreamHead {
|
||||
pub headers: Vec<(String, String)>,
|
||||
}
|
||||
|
||||
pub struct Messages;
|
||||
|
||||
impl Protocol for Messages {
|
||||
|
|
@ -62,7 +65,7 @@ impl Protocol for Messages {
|
|||
type Projection = MessagesCall;
|
||||
type Op = Infallible;
|
||||
type Chunk = Bytes;
|
||||
type StreamHead = ();
|
||||
type StreamHead = MessagesStreamHead;
|
||||
}
|
||||
|
||||
impl From<MachineFault> for Error {
|
||||
|
|
@ -77,19 +80,6 @@ impl From<MachineFault> for Error {
|
|||
pub type MessagesHost = HostChannel<Messages>;
|
||||
pub type MessagesMachine = CallMachine<Messages>;
|
||||
|
||||
/// Whether this route serves the request, decided before any callback runs so a host
|
||||
/// can still run its own path.
|
||||
pub fn supports(model: &str, custom_llm_provider: Option<&str>, stream: bool) -> bool {
|
||||
let provider = get_custom_llm_provider(model, custom_llm_provider)
|
||||
.map(|resolved| resolved.custom_llm_provider)
|
||||
.or(custom_llm_provider);
|
||||
match provider {
|
||||
Some(ANTHROPIC_MESSAGES_PROVIDER) => true,
|
||||
Some(provider) => !stream && messages_provider_config(provider).is_some(),
|
||||
None => false,
|
||||
}
|
||||
}
|
||||
|
||||
/// The in-process host for a request already in hand. It answers projection once and
|
||||
/// observes nothing.
|
||||
pub struct LocalMessagesHost {
|
||||
|
|
@ -152,8 +142,11 @@ async fn execute(
|
|||
model: request.model.clone(),
|
||||
custom_llm_provider: request.provider.clone(),
|
||||
optional_params: Value::Object(
|
||||
call.body
|
||||
.iter()
|
||||
request
|
||||
.body
|
||||
.as_object()
|
||||
.into_iter()
|
||||
.flatten()
|
||||
.filter(|(name, _)| !matches!(name.as_str(), "model" | "messages"))
|
||||
.map(|(name, value)| (name.clone(), value.clone()))
|
||||
.collect(),
|
||||
|
|
@ -193,7 +186,14 @@ async fn relay(
|
|||
host: &MessagesHost,
|
||||
mut response: reqwest::Response,
|
||||
) -> Result<MessagesOutput, Error> {
|
||||
if host.open(()).await? == Demand::Detached {
|
||||
let head = MessagesStreamHead {
|
||||
headers: response
|
||||
.headers()
|
||||
.iter()
|
||||
.filter_map(|(name, value)| Some((name.to_string(), value.to_str().ok()?.to_string())))
|
||||
.collect(),
|
||||
};
|
||||
if host.open(head).await? == Demand::Detached {
|
||||
return Ok(MessagesOutput::Streamed);
|
||||
}
|
||||
while let Some(chunk) = response.chunk().await.map_err(network)? {
|
||||
|
|
|
|||
210
litellm-rust/crates/core/tests/messages/host.rs
Normal file
210
litellm-rust/crates/core/tests/messages/host.rs
Normal file
|
|
@ -0,0 +1,210 @@
|
|||
use std::{convert::Infallible, sync::Mutex};
|
||||
|
||||
use litellm_core::messages::route::Messages;
|
||||
use litellm_host::{
|
||||
event::{CallEvent, MachineEvent, RequestContext, WireRequest},
|
||||
host::Host,
|
||||
};
|
||||
use litellm_llms::anthropic::common_utils::AnthropicModelCapabilities;
|
||||
use rstest::rstest;
|
||||
|
||||
use super::*;
|
||||
|
||||
type Rewrite = Box<dyn Fn(WireRequest) -> Result<WireRequest, Error> + Send + Sync>;
|
||||
|
||||
/// Projects like `LocalMessagesHost`, answers `before_send` through `rewrite`, and keeps
|
||||
/// every event the driver emits.
|
||||
struct RecordingHost {
|
||||
call: LocalMessagesHost,
|
||||
rewrite: Rewrite,
|
||||
events: Mutex<Vec<CallEvent>>,
|
||||
optional_params: Mutex<Vec<Value>>,
|
||||
}
|
||||
|
||||
impl RecordingHost {
|
||||
fn new(call: MessagesCall, rewrite: Rewrite) -> Self {
|
||||
Self {
|
||||
call: LocalMessagesHost::new(call),
|
||||
rewrite,
|
||||
events: Mutex::new(Vec::new()),
|
||||
optional_params: Mutex::new(Vec::new()),
|
||||
}
|
||||
}
|
||||
|
||||
fn passthrough(call: MessagesCall) -> Self {
|
||||
Self::new(call, Box::new(Ok))
|
||||
}
|
||||
|
||||
fn raw_responses(&self) -> Vec<String> {
|
||||
self.events
|
||||
.lock()
|
||||
.unwrap()
|
||||
.iter()
|
||||
.filter_map(|event| match event {
|
||||
CallEvent::Machine(MachineEvent::ResponseReceived { raw }) => {
|
||||
Some(raw.body.clone())
|
||||
}
|
||||
_ => None,
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
}
|
||||
|
||||
impl Host<Messages> for RecordingHost {
|
||||
async fn project(&self) -> Result<MessagesCall, Error> {
|
||||
self.call.project().await
|
||||
}
|
||||
|
||||
async fn custom_op(&self, op: Infallible) -> Result<(), Error> {
|
||||
match op {}
|
||||
}
|
||||
|
||||
async fn before_send(
|
||||
&self,
|
||||
wire: WireRequest,
|
||||
context: &RequestContext,
|
||||
) -> Result<WireRequest, Error> {
|
||||
self.optional_params
|
||||
.lock()
|
||||
.unwrap()
|
||||
.push(context.optional_params.clone());
|
||||
(self.rewrite)(wire)
|
||||
}
|
||||
|
||||
async fn emit(&self, event: &CallEvent) -> Result<(), Error> {
|
||||
self.events.lock().unwrap().push(event.clone());
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
async fn run_through(host: &RecordingHost) -> Result<MessagesOutput, Error> {
|
||||
litellm_host::run::run(messages_machine(Arc::new(RecordingSecrets::empty())), host).await
|
||||
}
|
||||
|
||||
fn authenticated(call: MessagesCall, api_base: String) -> MessagesCall {
|
||||
MessagesCall {
|
||||
api_key: Some("sk-ant".into()),
|
||||
api_base: Some(api_base),
|
||||
..call
|
||||
}
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn what_before_send_returns_is_what_the_provider_receives(call: MessagesCall) {
|
||||
let upstream = upstream([message_response()]).await;
|
||||
let host = RecordingHost::new(
|
||||
authenticated(call, upstream.uri()),
|
||||
Box::new(|wire| {
|
||||
let mut body = wire.body;
|
||||
body["system"] = json!("added by the host");
|
||||
Ok(WireRequest {
|
||||
headers: wire
|
||||
.headers
|
||||
.into_iter()
|
||||
.chain([("x-host".to_string(), "seen".to_string())])
|
||||
.collect(),
|
||||
body,
|
||||
..wire
|
||||
})
|
||||
}),
|
||||
);
|
||||
|
||||
run_through(&host).await.expect("messages call succeeds");
|
||||
|
||||
let request = only_request(&upstream).await;
|
||||
assert_eq!(request.json()["system"], "added by the host");
|
||||
assert_eq!(request.header("x-host"), Some("seen"));
|
||||
assert_eq!(request.header("x-api-key"), Some("sk-ant"));
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn a_before_send_failure_never_sends(call: MessagesCall) {
|
||||
let upstream = upstream([message_response()]).await;
|
||||
let host = RecordingHost::new(
|
||||
authenticated(call, upstream.uri()),
|
||||
Box::new(|_| Err(Error::InvalidRequest("vetoed by the host".into()))),
|
||||
);
|
||||
|
||||
let error = run_through(&host)
|
||||
.await
|
||||
.err()
|
||||
.expect("the host failure fails the call");
|
||||
|
||||
assert_eq!(error, Error::InvalidRequest("vetoed by the host".into()));
|
||||
assert!(received(&upstream).await.is_empty());
|
||||
assert!(host.raw_responses().is_empty());
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn the_raw_upstream_text_is_emitted_once_for_a_message(call: MessagesCall) {
|
||||
let raw = message_body();
|
||||
let upstream = upstream([json_response(raw.clone())]).await;
|
||||
let host = RecordingHost::passthrough(authenticated(call, upstream.uri()));
|
||||
|
||||
let output = run_through(&host).await.expect("messages call succeeds");
|
||||
|
||||
assert!(matches!(output, MessagesOutput::Message(_)));
|
||||
let [emitted] = <[String; 1]>::try_from(host.raw_responses())
|
||||
.unwrap_or_else(|raws| panic!("expected one raw response, got {}", raws.len()));
|
||||
assert_eq!(serde_json::from_str::<Value>(&emitted).unwrap(), raw);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::upstream_error(ResponseTemplate::new(500).set_body_string("boom"))]
|
||||
#[case::stream(ResponseTemplate::new(200).set_body_raw("event: message_stop\ndata: {}\n\n", "text/event-stream"))]
|
||||
#[tokio::test]
|
||||
async fn no_raw_response_is_emitted_for_a_stream_or_a_failure(
|
||||
call: MessagesCall,
|
||||
#[case] response: ResponseTemplate,
|
||||
) {
|
||||
let upstream = upstream([response]).await;
|
||||
let mut body = call.body.clone();
|
||||
body.insert("stream".into(), json!(true));
|
||||
let host =
|
||||
RecordingHost::passthrough(authenticated(MessagesCall { body, ..call }, upstream.uri()));
|
||||
|
||||
let _ = run_through(&host).await;
|
||||
|
||||
assert_eq!(received(&upstream).await.len(), 1);
|
||||
assert!(host.raw_responses().is_empty());
|
||||
}
|
||||
|
||||
/// Python logs `optional_params` as what it is about to send, so a dropped param must
|
||||
/// not resurface in callbacks.
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn the_request_context_carries_the_shaped_params_without_model_or_messages(
|
||||
call: MessagesCall,
|
||||
) {
|
||||
let upstream = upstream([message_response()]).await;
|
||||
let body: Map<String, Value> = call
|
||||
.body
|
||||
.clone()
|
||||
.into_iter()
|
||||
.chain([("temperature".to_string(), json!(0.2))])
|
||||
.collect();
|
||||
let host = RecordingHost::passthrough(authenticated(
|
||||
MessagesCall {
|
||||
body,
|
||||
shaping: MessagesShaping {
|
||||
capabilities: AnthropicModelCapabilities {
|
||||
supports_sampling_params: false,
|
||||
..AnthropicModelCapabilities::default()
|
||||
},
|
||||
drop_params: true,
|
||||
..MessagesShaping::default()
|
||||
},
|
||||
..call
|
||||
},
|
||||
upstream.uri(),
|
||||
));
|
||||
|
||||
run_through(&host).await.expect("messages call succeeds");
|
||||
|
||||
let [optional_params] = <[Value; 1]>::try_from(host.optional_params.into_inner().unwrap())
|
||||
.unwrap_or_else(|seen| panic!("before_send runs once, saw {}", seen.len()));
|
||||
assert_eq!(optional_params, json!({"max_tokens": 16}));
|
||||
}
|
||||
|
|
@ -14,6 +14,7 @@ use wiremock::ResponseTemplate;
|
|||
mod support;
|
||||
use support::*;
|
||||
|
||||
mod host;
|
||||
mod request;
|
||||
mod response;
|
||||
mod secrets;
|
||||
|
|
|
|||
|
|
@ -1,3 +1,7 @@
|
|||
use litellm_llms::anthropic::common_utils::{
|
||||
ANTHROPIC_ADVISOR_TOOL_TYPE, ANTHROPIC_OAUTH_BETA_HEADER, AnthropicModelCapabilities,
|
||||
SupportedEffortTiers, beta,
|
||||
};
|
||||
use litellm_types::utils::{ProviderSpecificHeader, ProviderSpecificHeaders};
|
||||
use rstest::rstest;
|
||||
|
||||
|
|
@ -132,8 +136,8 @@ async fn each_provider_posts_to_its_messages_endpoint(
|
|||
assert_eq!(request.method.as_str(), "POST");
|
||||
assert_eq!(request.url.path(), path);
|
||||
assert_eq!(request.json()["model"], MODEL);
|
||||
assert_eq!(request.header("anthropic-version"), Some("2023-06-01"));
|
||||
assert_eq!(request.header("content-type"), Some("application/json"));
|
||||
assert_eq!(request.header_values("anthropic-version"), ["2023-06-01"]);
|
||||
assert_eq!(request.header_values("content-type"), ["application/json"]);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
|
|
@ -249,21 +253,423 @@ async fn additional_drop_params_remove_fields_before_sending(call: MessagesCall)
|
|||
assert_eq!(sent["top_k"], 3);
|
||||
}
|
||||
|
||||
fn with_fields(call: MessagesCall, fields: Value) -> MessagesCall {
|
||||
let body: Map<String, Value> = call.body.into_iter().chain(object(fields)).collect();
|
||||
MessagesCall { body, ..call }
|
||||
}
|
||||
|
||||
fn sent_betas(request: &wiremock::Request) -> Vec<String> {
|
||||
let [header] = <[&str; 1]>::try_from(request.header_values("anthropic-beta"))
|
||||
.unwrap_or_else(|values| panic!("expected one anthropic-beta header, got {values:?}"));
|
||||
header
|
||||
.split(',')
|
||||
.map(str::trim)
|
||||
.map(str::to_string)
|
||||
.collect()
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::anthropic_streams(MODEL, Some("anthropic"), true, true)]
|
||||
#[case::anthropic_prefix_streams("anthropic/claude-sonnet-4-5", None, true, true)]
|
||||
#[case::azure_without_stream(MODEL, Some("azure_ai"), false, true)]
|
||||
#[case::azure_stream(MODEL, Some("azure_ai"), true, false)]
|
||||
#[case::other_provider(MODEL, Some("openai"), false, false)]
|
||||
#[case::unresolvable_model("no-such-model", None, false, false)]
|
||||
fn supports_matches_what_the_route_can_serve(
|
||||
#[case] model: &str,
|
||||
#[case] provider: Option<&str>,
|
||||
#[case] stream: bool,
|
||||
#[case] supported: bool,
|
||||
#[case::structured_output(json!({"output_format": {"type": "json_schema"}}), &[beta::STRUCTURED_OUTPUT])]
|
||||
#[case::fast_mode(json!({"speed": "fast"}), &[beta::FAST_MODE_2026_02_01])]
|
||||
#[case::compaction(json!({"compaction": {"enabled": true}}), &[beta::COMPACT_2026_09_04])]
|
||||
#[case::context_management_edits(
|
||||
json!({"context_management": {"edits": [{"type": "clear_tool_uses_20250919"}]}}),
|
||||
&[beta::CONTEXT_MANAGEMENT_2025_06_27]
|
||||
)]
|
||||
#[case::per_message_output_config(
|
||||
json!({"messages": [{"role": "user", "content": "hi", "output_config": {"effort": "low"}}]}),
|
||||
&[beta::PER_TURN_CONTROL_2026_07_01]
|
||||
)]
|
||||
#[case::advisor_tool(
|
||||
json!({"tools": [{"type": ANTHROPIC_ADVISOR_TOOL_TYPE, "name": "advisor", "model": MODEL}]}),
|
||||
&[beta::ADVISOR_TOOL_2026_03_01]
|
||||
)]
|
||||
#[case::several_features_at_once(
|
||||
json!({"speed": "fast", "output_format": {"type": "json_schema"}}),
|
||||
&[beta::STRUCTURED_OUTPUT, beta::FAST_MODE_2026_02_01]
|
||||
)]
|
||||
#[tokio::test]
|
||||
async fn feature_betas_join_the_callers_betas_in_one_sorted_header(
|
||||
call: MessagesCall,
|
||||
#[case] fields: Value,
|
||||
#[case] features: &[&str],
|
||||
) {
|
||||
let upstream = upstream([message_response()]).await;
|
||||
let capabilities = AnthropicModelCapabilities {
|
||||
supports_speed: true,
|
||||
..AnthropicModelCapabilities::default()
|
||||
};
|
||||
|
||||
run_message(with_fields(
|
||||
MessagesCall {
|
||||
api_key: Some("sk".into()),
|
||||
api_base: Some(upstream.uri()),
|
||||
extra_headers: headers([("Anthropic-Beta", "caller-beta-2025-01-01")]),
|
||||
shaping: MessagesShaping {
|
||||
capabilities,
|
||||
..MessagesShaping::default()
|
||||
},
|
||||
..call
|
||||
},
|
||||
fields,
|
||||
))
|
||||
.await;
|
||||
|
||||
let sent = sent_betas(&only_request(&upstream).await);
|
||||
let mut expected: Vec<String> = features
|
||||
.iter()
|
||||
.map(|feature| feature.to_string())
|
||||
.chain(["caller-beta-2025-01-01".to_string()])
|
||||
.collect();
|
||||
expected.sort();
|
||||
assert_eq!(sent, expected);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn an_oauth_key_sends_the_browser_access_header_and_the_oauth_beta(call: MessagesCall) {
|
||||
let upstream = upstream([message_response()]).await;
|
||||
|
||||
run_message(MessagesCall {
|
||||
api_key: Some("sk-ant-oat01-token".into()),
|
||||
api_base: Some(upstream.uri()),
|
||||
..call
|
||||
})
|
||||
.await;
|
||||
|
||||
let request = only_request(&upstream).await;
|
||||
assert_eq!(
|
||||
litellm_core::messages::route::supports(model, provider, stream),
|
||||
supported
|
||||
request.header("anthropic-dangerous-direct-browser-access"),
|
||||
Some("true")
|
||||
);
|
||||
assert_eq!(sent_betas(&request), [ANTHROPIC_OAUTH_BETA_HEADER]);
|
||||
assert_eq!(request.header("x-api-key"), None);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::anthropic("anthropic")]
|
||||
#[case::azure_ai("azure_ai")]
|
||||
#[tokio::test]
|
||||
async fn caller_protocol_headers_win_over_the_defaults(call: MessagesCall, #[case] provider: &str) {
|
||||
let upstream = upstream([message_response()]).await;
|
||||
|
||||
run_message(MessagesCall {
|
||||
custom_llm_provider: Some(provider.into()),
|
||||
api_key: Some("sk".into()),
|
||||
api_base: Some(upstream.uri()),
|
||||
extra_headers: headers([
|
||||
("Anthropic-Version", "2024-01-01"),
|
||||
("Content-Type", "application/json; charset=utf-8"),
|
||||
]),
|
||||
..call
|
||||
})
|
||||
.await;
|
||||
|
||||
let request = only_request(&upstream).await;
|
||||
assert_eq!(request.header_values("anthropic-version"), ["2024-01-01"]);
|
||||
assert_eq!(
|
||||
request.header_values("content-type"),
|
||||
["application/json; charset=utf-8"]
|
||||
);
|
||||
}
|
||||
|
||||
fn sampling_removed() -> AnthropicModelCapabilities {
|
||||
AnthropicModelCapabilities {
|
||||
supports_sampling_params: false,
|
||||
..AnthropicModelCapabilities::default()
|
||||
}
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::sampling_params(sampling_removed(), json!({"temperature": 0.2, "top_p": 0.9, "top_k": 5}), &["temperature", "top_p", "top_k"], "temperature=0.2")]
|
||||
#[case::speed(AnthropicModelCapabilities::default(), json!({"speed": "fast"}), &["speed"], "speed='fast'")]
|
||||
#[tokio::test]
|
||||
async fn unsupported_params_are_dropped_under_drop_params_and_rejected_without_it(
|
||||
call: MessagesCall,
|
||||
#[case] capabilities: AnthropicModelCapabilities,
|
||||
#[case] fields: Value,
|
||||
#[case] dropped: &[&str],
|
||||
#[case] rejected_as: &str,
|
||||
) {
|
||||
let upstream = upstream([message_response(), message_response()]).await;
|
||||
let shaped = |drop_params: bool| {
|
||||
with_fields(
|
||||
MessagesCall {
|
||||
api_key: Some("sk".into()),
|
||||
api_base: Some(upstream.uri()),
|
||||
shaping: MessagesShaping {
|
||||
capabilities: capabilities.clone(),
|
||||
drop_params,
|
||||
..MessagesShaping::default()
|
||||
},
|
||||
body: call.body.clone(),
|
||||
custom_llm_provider: call.custom_llm_provider.clone(),
|
||||
extra_headers: None,
|
||||
provider_specific_header: None,
|
||||
model: call.model.clone(),
|
||||
timeout: call.timeout,
|
||||
},
|
||||
fields.clone(),
|
||||
)
|
||||
};
|
||||
|
||||
let error = run(shaped(false))
|
||||
.await
|
||||
.err()
|
||||
.expect("an unsupported param is rejected without drop_params");
|
||||
assert!(
|
||||
matches!(&error, Error::InvalidRequest(message) if message.contains(rejected_as)),
|
||||
"{error:?}"
|
||||
);
|
||||
assert!(received(&upstream).await.is_empty());
|
||||
|
||||
run_message(shaped(true)).await;
|
||||
let sent = only_request(&upstream).await.json();
|
||||
for name in dropped {
|
||||
assert_eq!(sent.get(*name), None, "{name} must be dropped");
|
||||
}
|
||||
assert_eq!(sent["max_tokens"], 16);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::adaptive_thinking(json!({"type": "adaptive"}), json!({"type": "adaptive", "display": "summarized"}))]
|
||||
#[case::disabled_thinking(json!({"type": "disabled"}), json!({"type": "disabled"}))]
|
||||
#[tokio::test]
|
||||
async fn reasoning_auto_summary_marks_active_thinking_on_the_wire(
|
||||
call: MessagesCall,
|
||||
#[case] thinking: Value,
|
||||
#[case] expected: Value,
|
||||
) {
|
||||
let upstream = upstream([message_response()]).await;
|
||||
|
||||
run_message(with_fields(
|
||||
MessagesCall {
|
||||
api_key: Some("sk".into()),
|
||||
api_base: Some(upstream.uri()),
|
||||
shaping: MessagesShaping {
|
||||
capabilities: AnthropicModelCapabilities {
|
||||
supports_reasoning: true,
|
||||
supports_adaptive_thinking: true,
|
||||
..AnthropicModelCapabilities::default()
|
||||
},
|
||||
reasoning_auto_summary: true,
|
||||
..MessagesShaping::default()
|
||||
},
|
||||
..call
|
||||
},
|
||||
json!({"thinking": thinking}),
|
||||
))
|
||||
.await;
|
||||
|
||||
assert_eq!(only_request(&upstream).await.json()["thinking"], expected);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::reasoning_effort_on_an_adaptive_model(
|
||||
AnthropicModelCapabilities {
|
||||
supports_reasoning: true,
|
||||
supports_adaptive_thinking: true,
|
||||
supports_output_config: true,
|
||||
effort_tiers: SupportedEffortTiers { high: true, ..SupportedEffortTiers::default() },
|
||||
..AnthropicModelCapabilities::default()
|
||||
},
|
||||
json!({"reasoning_effort": "high"}),
|
||||
json!({"thinking": {"type": "adaptive", "display": "summarized"}, "output_config": {"effort": "high"}})
|
||||
)]
|
||||
#[case::reasoning_effort_on_a_legacy_model_caps_the_budget_below_max_tokens(
|
||||
AnthropicModelCapabilities {
|
||||
supports_reasoning: true,
|
||||
..AnthropicModelCapabilities::default()
|
||||
},
|
||||
json!({"reasoning_effort": "high"}),
|
||||
json!({"thinking": {"type": "enabled", "budget_tokens": 2999}})
|
||||
)]
|
||||
#[case::adaptive_payload_on_a_legacy_model_becomes_a_capped_budget(
|
||||
AnthropicModelCapabilities {
|
||||
supports_reasoning: true,
|
||||
..AnthropicModelCapabilities::default()
|
||||
},
|
||||
json!({"thinking": {"type": "adaptive"}, "output_config": {"effort": "high"}, "temperature": 0}),
|
||||
json!({"thinking": {"type": "enabled", "budget_tokens": 2999}})
|
||||
)]
|
||||
#[case::adaptive_payload_on_a_model_without_reasoning_is_dropped(
|
||||
AnthropicModelCapabilities::default(),
|
||||
json!({"thinking": {"type": "adaptive"}, "output_config": {"effort": "high"}}),
|
||||
json!({})
|
||||
)]
|
||||
#[tokio::test]
|
||||
async fn reasoning_is_translated_by_the_model_capabilities(
|
||||
call: MessagesCall,
|
||||
#[case] capabilities: AnthropicModelCapabilities,
|
||||
#[case] fields: Value,
|
||||
#[case] expected: Value,
|
||||
) {
|
||||
let upstream = upstream([message_response()]).await;
|
||||
|
||||
run_message(with_fields(
|
||||
MessagesCall {
|
||||
api_key: Some("sk".into()),
|
||||
api_base: Some(upstream.uri()),
|
||||
shaping: MessagesShaping {
|
||||
capabilities,
|
||||
..MessagesShaping::default()
|
||||
},
|
||||
..call
|
||||
},
|
||||
[("max_tokens".to_string(), json!(3000))]
|
||||
.into_iter()
|
||||
.chain(object(fields))
|
||||
.collect(),
|
||||
))
|
||||
.await;
|
||||
|
||||
let sent = only_request(&upstream).await.json();
|
||||
assert_eq!(sent.get("reasoning_effort"), None);
|
||||
assert_eq!(sent.get("temperature"), None);
|
||||
let reasoning: Map<String, Value> = ["thinking", "output_config"]
|
||||
.into_iter()
|
||||
.filter_map(|name| Some((name.to_string(), sent.get(name)?.clone())))
|
||||
.collect();
|
||||
assert_eq!(Value::Object(reasoning), expected);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::empty_text_blocks(
|
||||
json!([{"role": "assistant", "content": [{"type": "text", "text": " "}, {"type": "text", "text": "kept"}]}]),
|
||||
json!([{"role": "assistant", "content": [{"type": "text", "text": "kept"}]}])
|
||||
)]
|
||||
#[case::provider_specific_fields(
|
||||
json!([{"role": "assistant", "content": [{"type": "text", "text": "kept", "provider_specific_fields": {"x": 1}}]}]),
|
||||
json!([{"role": "assistant", "content": [{"type": "text", "text": "kept"}]}])
|
||||
)]
|
||||
#[case::unencrypted_web_search_results_become_text(
|
||||
json!([{"role": "assistant", "content": [{
|
||||
"type": "web_search_tool_result",
|
||||
"tool_use_id": "srvtoolu_1",
|
||||
"content": [{"type": "web_search_result", "title": "T", "url": "https://e.x", "page_age": null}]
|
||||
}]}]),
|
||||
json!([{"role": "assistant", "content": [{"type": "text", "text": "Web search results:\n\nTitle: T\nURL: https://e.x"}]}])
|
||||
)]
|
||||
#[tokio::test]
|
||||
async fn replayed_history_is_cleaned_before_sending(
|
||||
call: MessagesCall,
|
||||
#[case] history: Value,
|
||||
#[case] expected: Value,
|
||||
) {
|
||||
let upstream = upstream([message_response()]).await;
|
||||
|
||||
run_message(with_fields(
|
||||
MessagesCall {
|
||||
api_key: Some("sk".into()),
|
||||
api_base: Some(upstream.uri()),
|
||||
..call
|
||||
},
|
||||
json!({"messages": history}),
|
||||
))
|
||||
.await;
|
||||
|
||||
assert_eq!(only_request(&upstream).await.json()["messages"], expected);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn metadata_is_reduced_to_the_user_id(call: MessagesCall) {
|
||||
let upstream = upstream([message_response()]).await;
|
||||
|
||||
run_message(with_fields(
|
||||
MessagesCall {
|
||||
api_key: Some("sk".into()),
|
||||
api_base: Some(upstream.uri()),
|
||||
..call
|
||||
},
|
||||
json!({"metadata": {"user_id": "u-1", "trace_id": "internal", "tags": ["a"]}}),
|
||||
))
|
||||
.await;
|
||||
|
||||
assert_eq!(
|
||||
only_request(&upstream).await.json()["metadata"],
|
||||
json!({"user_id": "u-1"})
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::numeric_user_id(json!({"metadata": {"user_id": 7}}))]
|
||||
#[case::missing_max_tokens(json!({"max_tokens": null}))]
|
||||
#[tokio::test]
|
||||
async fn an_invalid_request_fails_before_sending(call: MessagesCall, #[case] fields: Value) {
|
||||
let upstream = upstream([message_response()]).await;
|
||||
|
||||
let error = run(with_fields(
|
||||
MessagesCall {
|
||||
api_key: Some("sk".into()),
|
||||
api_base: Some(upstream.uri()),
|
||||
..call
|
||||
},
|
||||
fields,
|
||||
))
|
||||
.await
|
||||
.err()
|
||||
.expect("the request is rejected");
|
||||
|
||||
assert!(error.is_request(), "{error:?}");
|
||||
assert!(received(&upstream).await.is_empty());
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn azure_folds_system_role_messages_into_the_system_prompt(call: MessagesCall) {
|
||||
let upstream = upstream([message_response()]).await;
|
||||
|
||||
run_message(with_fields(
|
||||
MessagesCall {
|
||||
custom_llm_provider: Some("azure_ai".into()),
|
||||
api_key: Some("sk-azure".into()),
|
||||
api_base: Some(upstream.uri()),
|
||||
..call
|
||||
},
|
||||
json!({
|
||||
"system": "top level",
|
||||
"messages": [
|
||||
{"role": "system", "content": "from a message"},
|
||||
{"role": "user", "content": "hi"}
|
||||
]
|
||||
}),
|
||||
))
|
||||
.await;
|
||||
|
||||
let sent = only_request(&upstream).await.json();
|
||||
assert_eq!(
|
||||
sent["system"],
|
||||
json!([
|
||||
{"type": "text", "text": "top level"},
|
||||
{"type": "text", "text": "from a message"}
|
||||
])
|
||||
);
|
||||
assert_eq!(sent["messages"], json!([{"role": "user", "content": "hi"}]));
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::bare_model(MODEL, MODEL)]
|
||||
#[case::one_prefix("anthropic/claude-sonnet-4-5", MODEL)]
|
||||
#[case::doubled_prefix_loses_one_segment(
|
||||
"anthropic/anthropic/claude-sonnet-4-5",
|
||||
"anthropic/claude-sonnet-4-5"
|
||||
)]
|
||||
#[tokio::test]
|
||||
async fn the_provider_prefix_is_stripped_exactly_once(
|
||||
call: MessagesCall,
|
||||
#[case] model: &str,
|
||||
#[case] sent_model: &str,
|
||||
) {
|
||||
let upstream = upstream([message_response()]).await;
|
||||
|
||||
run_message(MessagesCall {
|
||||
model: model.into(),
|
||||
api_key: Some("sk".into()),
|
||||
api_base: Some(upstream.uri()),
|
||||
..call
|
||||
})
|
||||
.await;
|
||||
|
||||
assert_eq!(only_request(&upstream).await.json()["model"], sent_model);
|
||||
}
|
||||
|
|
|
|||
|
|
@ -5,11 +5,14 @@ use rstest::rstest;
|
|||
use super::*;
|
||||
|
||||
#[rstest]
|
||||
#[case::anthropic("anthropic")]
|
||||
#[case::azure_ai("azure_ai")]
|
||||
#[tokio::test]
|
||||
async fn the_provider_message_is_returned(call: MessagesCall) {
|
||||
async fn the_provider_message_is_returned(call: MessagesCall, #[case] provider: &str) {
|
||||
let upstream = upstream([message_response()]).await;
|
||||
|
||||
let message = run_message(MessagesCall {
|
||||
custom_llm_provider: Some(provider.into()),
|
||||
api_key: Some("sk".into()),
|
||||
api_base: Some(upstream.uri()),
|
||||
..call
|
||||
|
|
@ -21,6 +24,88 @@ async fn the_provider_message_is_returned(call: MessagesCall) {
|
|||
assert_eq!(message.stop_reason.as_deref(), Some("end_turn"));
|
||||
}
|
||||
|
||||
/// A refusal and fields the route does not model come back exactly as the provider sent
|
||||
/// them, since the Python side returns the raw message and the router decides what to do.
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn the_message_passes_through_losslessly(call: MessagesCall) {
|
||||
let upstream_body = json!({
|
||||
"id": "msg_2",
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"model": MODEL,
|
||||
"content": [
|
||||
{"type": "server_tool_use", "id": "srvtoolu_1", "name": "web_search", "input": {"query": "q"}},
|
||||
{"type": "text", "text": "no", "citations": [{"type": "web_search_result_location", "url": "https://e.x"}]}
|
||||
],
|
||||
"stop_reason": "refusal",
|
||||
"stop_sequence": null,
|
||||
"stop_details": {"type": "safeguard", "safeguard_types": ["dangerous_tool_use"]},
|
||||
"container": {"id": "container_1", "expires_at": "2026-01-01T00:00:00Z"},
|
||||
"context_management": {"applied_edits": []},
|
||||
"usage": {"input_tokens": 1, "output_tokens": 2, "server_tool_use": {"web_search_requests": 1}},
|
||||
"unknown_future_field": {"nested": true}
|
||||
});
|
||||
let upstream = upstream([json_response(upstream_body.clone())]).await;
|
||||
|
||||
let message = run_message(MessagesCall {
|
||||
api_key: Some("sk".into()),
|
||||
api_base: Some(upstream.uri()),
|
||||
..call
|
||||
})
|
||||
.await;
|
||||
|
||||
assert_eq!(message.stop_reason.as_deref(), Some("refusal"));
|
||||
assert_eq!(serde_json::to_value(&message).unwrap(), upstream_body);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn a_json_error_envelope_is_kept_verbatim(call: MessagesCall) {
|
||||
let envelope =
|
||||
json!({"type": "error", "error": {"type": "invalid_request_error", "message": "bad"}});
|
||||
let upstream = upstream([status_response(400, envelope.clone())]).await;
|
||||
|
||||
let error = run(MessagesCall {
|
||||
api_key: Some("sk".into()),
|
||||
api_base: Some(upstream.uri()),
|
||||
..call
|
||||
})
|
||||
.await
|
||||
.err()
|
||||
.expect("upstream error propagates");
|
||||
|
||||
let Error::Transport(TransportError::Http { status, body }) = error else {
|
||||
panic!("{error:?}");
|
||||
};
|
||||
assert_eq!(status, 400);
|
||||
assert_eq!(serde_json::from_str::<Value>(&body).unwrap(), envelope);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn a_long_error_body_is_truncated_at_the_documented_cap(call: MessagesCall) {
|
||||
let long = "x".repeat(600);
|
||||
let upstream = upstream([ResponseTemplate::new(500).set_body_string(long.clone())]).await;
|
||||
|
||||
let error = run(MessagesCall {
|
||||
api_key: Some("sk".into()),
|
||||
api_base: Some(upstream.uri()),
|
||||
..call
|
||||
})
|
||||
.await
|
||||
.err()
|
||||
.expect("upstream error propagates");
|
||||
|
||||
assert_eq!(
|
||||
error,
|
||||
Error::Transport(TransportError::Http {
|
||||
status: 500,
|
||||
body: format!("{}... (truncated)", &long[..256])
|
||||
})
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::bad_request(400)]
|
||||
#[case::unauthorized(401)]
|
||||
|
|
|
|||
|
|
@ -1,23 +1,30 @@
|
|||
use litellm_llms::{
|
||||
anthropic::experimental_pass_through::messages::transformation::ANTHROPIC_MESSAGES_CONFIG,
|
||||
azure_ai::anthropic::messages_transformation::AZURE_ANTHROPIC_MESSAGES_CONFIG,
|
||||
base_llm::anthropic_messages::transformation::BaseAnthropicMessagesConfig,
|
||||
};
|
||||
use rstest::rstest;
|
||||
|
||||
use super::*;
|
||||
|
||||
#[rstest]
|
||||
#[case::anthropic("anthropic", &ANTHROPIC_MESSAGES_CONFIG, "ANTHROPIC_API_KEY", "ANTHROPIC_BASE_URL", "/v1/messages")]
|
||||
#[case::azure_ai("azure_ai", &AZURE_ANTHROPIC_MESSAGES_CONFIG, "AZURE_API_KEY", "AZURE_API_BASE", "/anthropic/v1/messages")]
|
||||
#[case::anthropic(
|
||||
"anthropic",
|
||||
"ANTHROPIC_API_KEY",
|
||||
"ANTHROPIC_BASE_URL",
|
||||
"/v1/messages",
|
||||
&["ANTHROPIC_API_KEY", "ANTHROPIC_AUTH_TOKEN", "ANTHROPIC_API_BASE", "ANTHROPIC_BASE_URL"]
|
||||
)]
|
||||
#[case::azure_ai(
|
||||
"azure_ai",
|
||||
"AZURE_API_KEY",
|
||||
"AZURE_API_BASE",
|
||||
"/anthropic/v1/messages",
|
||||
&["AZURE_API_KEY", "AZURE_API_BASE"]
|
||||
)]
|
||||
#[tokio::test]
|
||||
async fn the_credential_and_base_come_from_the_secret_source(
|
||||
call: MessagesCall,
|
||||
#[case] provider: &str,
|
||||
#[case] config: &dyn BaseAnthropicMessagesConfig,
|
||||
#[case] key_name: &str,
|
||||
#[case] base_name: &str,
|
||||
#[case] path: &str,
|
||||
#[case] looked_up: &[&str],
|
||||
) {
|
||||
let upstream = upstream([message_response()]).await;
|
||||
let base = upstream.uri();
|
||||
|
|
@ -40,7 +47,7 @@ async fn the_credential_and_base_come_from_the_secret_source(
|
|||
let request = only_request(&upstream).await;
|
||||
assert_eq!(request.url.path(), path);
|
||||
assert_eq!(request.header("x-api-key"), Some("sk-from-manager"));
|
||||
assert_eq!(secrets.requested(), config.secret_names());
|
||||
assert_eq!(secrets.requested(), looked_up);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
|
|
@ -92,3 +99,102 @@ async fn a_secret_manager_failure_fails_the_call_before_sending(call: MessagesCa
|
|||
);
|
||||
assert!(received(&upstream).await.is_empty());
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy)]
|
||||
enum Base {
|
||||
Upstream,
|
||||
Unreachable,
|
||||
Blank,
|
||||
Absent,
|
||||
}
|
||||
|
||||
fn base_value(base: Base, upstream: &str) -> Option<String> {
|
||||
match base {
|
||||
Base::Upstream => Some(upstream.to_string()),
|
||||
Base::Unreachable => Some(UNREACHABLE_BASE.to_string()),
|
||||
Base::Blank => Some(" ".to_string()),
|
||||
Base::Absent => None,
|
||||
}
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::api_base_beats_base_url(Base::Upstream, Base::Unreachable)]
|
||||
#[case::blank_api_base_falls_through_to_base_url(Base::Blank, Base::Upstream)]
|
||||
#[case::base_url_alone(Base::Absent, Base::Upstream)]
|
||||
#[tokio::test]
|
||||
async fn the_anthropic_base_env_precedence_picks_the_upstream(
|
||||
call: MessagesCall,
|
||||
#[case] api_base: Base,
|
||||
#[case] base_url: Base,
|
||||
) {
|
||||
let upstream = upstream([message_response()]).await;
|
||||
let uri = upstream.uri();
|
||||
let values: Vec<(&str, &str)> = [
|
||||
("ANTHROPIC_API_KEY", Some("sk-env".to_string())),
|
||||
("ANTHROPIC_API_BASE", base_value(api_base, &uri)),
|
||||
("ANTHROPIC_BASE_URL", base_value(base_url, &uri)),
|
||||
]
|
||||
.iter()
|
||||
.filter_map(|(name, value)| Some((*name, value.as_deref()?)))
|
||||
.map(|(name, value)| (name, Box::leak(value.to_string().into_boxed_str()) as &str))
|
||||
.collect();
|
||||
|
||||
run_with(Arc::new(RecordingSecrets::new(values)), call)
|
||||
.await
|
||||
.expect("messages call reaches the upstream the precedence picks");
|
||||
|
||||
assert_eq!(only_request(&upstream).await.url.path(), "/v1/messages");
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::auth_token_alone(
|
||||
&[("ANTHROPIC_AUTH_TOKEN", "tok")],
|
||||
("authorization", "Bearer tok"),
|
||||
"x-api-key"
|
||||
)]
|
||||
#[case::api_key_beats_the_auth_token(
|
||||
&[("ANTHROPIC_API_KEY", "sk-env"), ("ANTHROPIC_AUTH_TOKEN", "tok")],
|
||||
("x-api-key", "sk-env"),
|
||||
"authorization"
|
||||
)]
|
||||
#[tokio::test]
|
||||
async fn the_auth_token_env_is_a_bearer_only_without_a_key(
|
||||
call: MessagesCall,
|
||||
#[case] values: &[(&str, &str)],
|
||||
#[case] expected: (&str, &str),
|
||||
#[case] absent: &str,
|
||||
) {
|
||||
let upstream = upstream([message_response()]).await;
|
||||
|
||||
run_with(
|
||||
Arc::new(RecordingSecrets::new(values.iter().copied())),
|
||||
MessagesCall {
|
||||
api_base: Some(upstream.uri()),
|
||||
..call
|
||||
},
|
||||
)
|
||||
.await
|
||||
.expect("messages call succeeds");
|
||||
|
||||
let request = only_request(&upstream).await;
|
||||
let (name, value) = expected;
|
||||
assert_eq!(request.header_values(name), [value]);
|
||||
assert_eq!(request.header(absent), None);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn azure_without_a_base_anywhere_fails_before_sending(call: MessagesCall) {
|
||||
let error = run_with(
|
||||
Arc::new(RecordingSecrets::new([("AZURE_API_KEY", "sk-azure")])),
|
||||
MessagesCall {
|
||||
custom_llm_provider: Some("azure_ai".into()),
|
||||
..call
|
||||
},
|
||||
)
|
||||
.await
|
||||
.err()
|
||||
.expect("azure needs a base");
|
||||
|
||||
assert_eq!(error, Error::Auth(litellm_auth::Error::MissingAzureApiBase));
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,16 +1,25 @@
|
|||
use std::{convert::Infallible, sync::Mutex};
|
||||
|
||||
use bytes::Bytes;
|
||||
use litellm_core::messages::route::Messages;
|
||||
use litellm_core::messages::route::{Messages, MessagesStreamHead};
|
||||
use litellm_host::host::{Demand, Host};
|
||||
use rstest::rstest;
|
||||
use tokio::{
|
||||
io::{AsyncReadExt, AsyncWriteExt},
|
||||
net::TcpListener,
|
||||
};
|
||||
|
||||
use super::*;
|
||||
|
||||
const UPSTREAM_HEADERS: [(&str, &str); 2] = [
|
||||
("request-id", "req_upstream_123"),
|
||||
("anthropic-ratelimit-requests-remaining", "41"),
|
||||
];
|
||||
|
||||
const SSE_BODY: &str = "event: message_start\ndata: {\"type\":\"message_start\"}\n\nevent: message_stop\ndata: {\"type\":\"message_stop\"}\n\n";
|
||||
|
||||
enum Seen {
|
||||
Open,
|
||||
Open(Vec<(String, String)>),
|
||||
Deliver(Bytes),
|
||||
}
|
||||
|
||||
|
|
@ -50,8 +59,8 @@ impl Host<Messages> for RecordingStreamHost {
|
|||
match op {}
|
||||
}
|
||||
|
||||
async fn open(&self, (): ()) -> Result<Demand, Error> {
|
||||
Ok(self.record(Seen::Open))
|
||||
async fn open(&self, head: MessagesStreamHead) -> Result<Demand, Error> {
|
||||
Ok(self.record(Seen::Open(head.headers)))
|
||||
}
|
||||
|
||||
async fn deliver(&self, chunk: Bytes) -> Result<Demand, Error> {
|
||||
|
|
@ -71,7 +80,10 @@ fn streaming(call: MessagesCall, api_base: String) -> MessagesCall {
|
|||
}
|
||||
|
||||
fn sse_response() -> ResponseTemplate {
|
||||
ResponseTemplate::new(200).set_body_raw(SSE_BODY, "text/event-stream")
|
||||
UPSTREAM_HEADERS.iter().fold(
|
||||
ResponseTemplate::new(200).set_body_raw(SSE_BODY, "text/event-stream"),
|
||||
|response, (name, value)| response.insert_header(*name, *value),
|
||||
)
|
||||
}
|
||||
|
||||
async fn stream_through(host: &RecordingStreamHost) -> Result<MessagesOutput, Error> {
|
||||
|
|
@ -80,7 +92,7 @@ async fn stream_through(host: &RecordingStreamHost) -> Result<MessagesOutput, Er
|
|||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn the_stream_opens_once_before_relaying_the_upstream_body(call: MessagesCall) {
|
||||
async fn upstream_headers_are_on_the_stream_head_before_the_first_chunk(call: MessagesCall) {
|
||||
let upstream = upstream([sse_response()]).await;
|
||||
let host = RecordingStreamHost::new(streaming(call, upstream.uri()), usize::MAX);
|
||||
|
||||
|
|
@ -88,14 +100,24 @@ async fn the_stream_opens_once_before_relaying_the_upstream_body(call: MessagesC
|
|||
|
||||
assert!(matches!(outcome, MessagesOutput::Streamed));
|
||||
let seen = host.seen.into_inner().unwrap();
|
||||
let [Seen::Open, chunks @ ..] = seen.as_slice() else {
|
||||
let [Seen::Open(headers), chunks @ ..] = seen.as_slice() else {
|
||||
panic!("the stream opens before any chunk is delivered");
|
||||
};
|
||||
let surfaced: Vec<(&str, &str)> = headers
|
||||
.iter()
|
||||
.filter(|(name, _)| {
|
||||
UPSTREAM_HEADERS
|
||||
.iter()
|
||||
.any(|(upstream, _)| upstream == name)
|
||||
})
|
||||
.map(|(name, value)| (name.as_str(), value.as_str()))
|
||||
.collect();
|
||||
assert_eq!(surfaced, UPSTREAM_HEADERS);
|
||||
let delivered: Vec<u8> = chunks
|
||||
.iter()
|
||||
.flat_map(|step| match step {
|
||||
Seen::Deliver(chunk) => chunk.to_vec(),
|
||||
Seen::Open => panic!("the stream opens exactly once"),
|
||||
Seen::Open(_) => panic!("the stream opens exactly once"),
|
||||
})
|
||||
.collect();
|
||||
assert_eq!(delivered, SSE_BODY.as_bytes());
|
||||
|
|
@ -118,9 +140,18 @@ async fn a_detached_caller_receives_nothing_more(call: MessagesCall, #[case] det
|
|||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::text_body(ResponseTemplate::new(429).set_body_string("slow down"), "slow down")]
|
||||
#[case::json_envelope(
|
||||
status_response(429, json!({"type": "error", "error": {"type": "rate_limit_error", "message": "slow down"}})),
|
||||
r#"{"type":"error","error":{"type":"rate_limit_error","message":"slow down"}}"#
|
||||
)]
|
||||
#[tokio::test]
|
||||
async fn an_upstream_error_fails_the_call_without_opening_the_stream(call: MessagesCall) {
|
||||
let upstream = upstream([ResponseTemplate::new(429).set_body_string("slow down")]).await;
|
||||
async fn an_upstream_error_fails_the_call_without_opening_the_stream(
|
||||
call: MessagesCall,
|
||||
#[case] response: ResponseTemplate,
|
||||
#[case] body: &str,
|
||||
) {
|
||||
let upstream = upstream([response]).await;
|
||||
let host = RecordingStreamHost::new(streaming(call, upstream.uri()), usize::MAX);
|
||||
|
||||
let error = stream_through(&host)
|
||||
|
|
@ -128,16 +159,88 @@ async fn an_upstream_error_fails_the_call_without_opening_the_stream(call: Messa
|
|||
.err()
|
||||
.expect("upstream error propagates");
|
||||
|
||||
assert!(
|
||||
matches!(
|
||||
error,
|
||||
Error::Transport(litellm_http::transport::Error::Http { status: 429, .. })
|
||||
),
|
||||
"{error:?}"
|
||||
assert_eq!(
|
||||
error,
|
||||
Error::Transport(litellm_http::transport::Error::Http {
|
||||
status: 429,
|
||||
body: body.into()
|
||||
})
|
||||
);
|
||||
assert!(host.seen.into_inner().unwrap().is_empty());
|
||||
}
|
||||
|
||||
/// The native route relays bytes as they are. Python's synthetic `api_error` for a stream
|
||||
/// that never reaches `message_stop` lives in its SSE wrapper, above this route.
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn a_stream_that_ends_without_message_stop_is_relayed_as_is(call: MessagesCall) {
|
||||
const INCOMPLETE: &str = "event: message_start\ndata: {\"type\":\"message_start\"}\n\n";
|
||||
let upstream =
|
||||
upstream([ResponseTemplate::new(200).set_body_raw(INCOMPLETE, "text/event-stream")]).await;
|
||||
let host = RecordingStreamHost::new(streaming(call, upstream.uri()), usize::MAX);
|
||||
|
||||
stream_through(&host).await.expect("streamed call succeeds");
|
||||
|
||||
let delivered: Vec<u8> = host
|
||||
.seen
|
||||
.into_inner()
|
||||
.unwrap()
|
||||
.iter()
|
||||
.flat_map(|step| match step {
|
||||
Seen::Deliver(chunk) => chunk.to_vec(),
|
||||
Seen::Open(_) => Vec::new(),
|
||||
})
|
||||
.collect();
|
||||
assert_eq!(delivered, INCOMPLETE.as_bytes());
|
||||
}
|
||||
|
||||
/// Serves one SSE chunk and then holds the connection open without ever finishing.
|
||||
async fn stalling_upstream() -> String {
|
||||
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
|
||||
let base = format!("http://{}", listener.local_addr().unwrap());
|
||||
tokio::spawn(async move {
|
||||
let (mut socket, _) = listener.accept().await.unwrap();
|
||||
let mut request = vec![0; 4096];
|
||||
let _ = socket.read(&mut request).await;
|
||||
socket
|
||||
.write_all(
|
||||
b"HTTP/1.1 200 OK\r\ncontent-type: text/event-stream\r\ntransfer-encoding: chunked\r\n\r\n\
|
||||
1f\r\nevent: message_start\ndata: {}\n\n\r\n",
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
std::future::pending::<()>().await;
|
||||
});
|
||||
base
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn the_timeout_covers_a_stalled_stream_body(call: MessagesCall) {
|
||||
let base = stalling_upstream().await;
|
||||
let host = RecordingStreamHost::new(
|
||||
MessagesCall {
|
||||
timeout: Some(Duration::from_millis(300)),
|
||||
..streaming(call, base)
|
||||
},
|
||||
usize::MAX,
|
||||
);
|
||||
|
||||
let error = tokio::time::timeout(Duration::from_secs(5), stream_through(&host))
|
||||
.await
|
||||
.expect("the stalled stream gives up within the timeout")
|
||||
.err()
|
||||
.expect("a stalled body fails the call");
|
||||
|
||||
assert!(matches!(error, Error::Transport(_)), "{error:?}");
|
||||
let seen = host.seen.into_inner().unwrap();
|
||||
assert!(
|
||||
matches!(seen.as_slice(), [Seen::Open(_), Seen::Deliver(chunk)] if chunk.as_ref() == b"event: message_start\ndata: {}\n\n"),
|
||||
"the chunk before the stall reached the caller, saw {} ops",
|
||||
seen.len()
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn streaming_is_refused_for_providers_that_cannot_stream(call: MessagesCall) {
|
||||
|
|
|
|||
|
|
@ -134,6 +134,13 @@ pub trait ProtocolHost: Send + Sync {
|
|||
response: <Self::Protocol as Protocol>::Response,
|
||||
) -> PyResult<Py<PyAny>>;
|
||||
|
||||
/// What the stream carries at hand-off, as the caller's stream receives it.
|
||||
fn head(
|
||||
&mut self,
|
||||
py: Python<'_>,
|
||||
head: <Self::Protocol as Protocol>::StreamHead,
|
||||
) -> PyResult<Py<PyAny>>;
|
||||
|
||||
/// One streamed chunk as the caller receives it.
|
||||
fn chunk(
|
||||
&mut self,
|
||||
|
|
|
|||
|
|
@ -134,10 +134,10 @@ where
|
|||
}
|
||||
match driver.resume(None)? {
|
||||
ExecutionStep::Return(value) => Ok(value),
|
||||
ExecutionStep::Open => py
|
||||
ExecutionStep::Open(head) => py
|
||||
.import("litellm.rust_bridge.lifecycle")?
|
||||
.getattr("SyncStream")?
|
||||
.call1((Py::new(py, Execution::suspended(driver))?,))
|
||||
.call1((Py::new(py, Execution::suspended(driver))?, head))
|
||||
.map(Bound::unbind),
|
||||
ExecutionStep::Await(_) | ExecutionStep::Yield(_) => {
|
||||
Err(PyRuntimeError::new_err("sync call suspended"))
|
||||
|
|
@ -312,7 +312,7 @@ where
|
|||
Ok(_) => return Err(missing_state()),
|
||||
Err(error) => Err(error),
|
||||
},
|
||||
HostOp::Open(_, reply) => return self.opened(py, reply).map(Next::Return),
|
||||
HostOp::Open(head, reply) => return self.opened(py, head, reply).map(Next::Return),
|
||||
HostOp::Deliver(chunk, reply) => {
|
||||
return self.delivered(py, chunk, reply).map(Next::Return);
|
||||
}
|
||||
|
|
@ -340,12 +340,21 @@ where
|
|||
}
|
||||
}
|
||||
|
||||
fn opened(&mut self, py: Python<'_>, reply: Reply<Demand>) -> PyResult<ExecutionStep> {
|
||||
fn opened(
|
||||
&mut self,
|
||||
py: Python<'_>,
|
||||
head: <ProtocolOf<H> as Protocol>::StreamHead,
|
||||
reply: Reply<Demand>,
|
||||
) -> PyResult<ExecutionStep> {
|
||||
self.stage = Stage::Streaming;
|
||||
let head = match self.host.head(py, head) {
|
||||
Ok(head) => head,
|
||||
Err(error) => return self.interrupt(py, error),
|
||||
};
|
||||
match self.adapter.opened(py) {
|
||||
Ok(()) => {
|
||||
self.pending = Some(Pending::Consumer(reply));
|
||||
Ok(ExecutionStep::Open)
|
||||
Ok(ExecutionStep::Open(head))
|
||||
}
|
||||
Err(error) => self.interrupt(py, error),
|
||||
}
|
||||
|
|
@ -699,6 +708,10 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri
|
|||
.map(|answer| reply.send(answer))
|
||||
}
|
||||
|
||||
fn head(&mut self, _: Python<'_>, head: std::convert::Infallible) -> PyResult<Py<PyAny>> {
|
||||
match head {}
|
||||
}
|
||||
|
||||
fn chunk(&mut self, _: Python<'_>, chunk: std::convert::Infallible) -> PyResult<Py<PyAny>> {
|
||||
match chunk {}
|
||||
}
|
||||
|
|
@ -945,6 +958,163 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri
|
|||
});
|
||||
}
|
||||
|
||||
struct Streaming;
|
||||
|
||||
impl Protocol for Streaming {
|
||||
type Response = ();
|
||||
type Error = Error;
|
||||
type Projection = ();
|
||||
type Op = std::convert::Infallible;
|
||||
type Chunk = &'static str;
|
||||
type StreamHead = Vec<(&'static str, &'static str)>;
|
||||
}
|
||||
|
||||
struct StreamingHost;
|
||||
|
||||
impl ProtocolHost for StreamingHost {
|
||||
type Protocol = Streaming;
|
||||
type Failure = Classified;
|
||||
|
||||
fn project(
|
||||
&mut self,
|
||||
_: Python<'_>,
|
||||
_: &Bound<'_, PyDict>,
|
||||
) -> Result<(), InvokeError<Error>> {
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn invoke(
|
||||
&mut self,
|
||||
_: Python<'_>,
|
||||
op: std::convert::Infallible,
|
||||
) -> Result<(), InvokeError<Error>> {
|
||||
match op {}
|
||||
}
|
||||
|
||||
fn head(
|
||||
&mut self,
|
||||
py: Python<'_>,
|
||||
head: Vec<(&'static str, &'static str)>,
|
||||
) -> PyResult<Py<PyAny>> {
|
||||
let headers = PyDict::new(py);
|
||||
for (name, value) in head {
|
||||
headers.set_item(name, value)?;
|
||||
}
|
||||
let hidden = PyDict::new(py);
|
||||
hidden.set_item("additional_headers", headers)?;
|
||||
Ok(hidden.into_any().unbind())
|
||||
}
|
||||
|
||||
fn chunk(&mut self, py: Python<'_>, chunk: &'static str) -> PyResult<Py<PyAny>> {
|
||||
Ok(pyo3::types::PyString::new(py, chunk).into_any().unbind())
|
||||
}
|
||||
|
||||
fn complete(&mut self, py: Python<'_>, (): ()) -> PyResult<Py<PyAny>> {
|
||||
Ok(py.None())
|
||||
}
|
||||
|
||||
fn classify(&self, _: Python<'_>, error: Error) -> PyResult<Classified> {
|
||||
Ok(Classified(error.0))
|
||||
}
|
||||
|
||||
fn host_error(error: &PyErr) -> Error {
|
||||
Error(error.to_string())
|
||||
}
|
||||
|
||||
fn close(&mut self, _: Python<'_>) {}
|
||||
|
||||
fn traverse(&self, _: &PyVisit<'_>) -> Result<(), PyTraverseError> {
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
fn streaming_machine() -> CallMachine<Streaming> {
|
||||
CallMachine::new(|host| {
|
||||
Box::pin(async move {
|
||||
host.project().await?;
|
||||
if host.open(vec![("request-id", "req_1")]).await? == Demand::Detached {
|
||||
return Ok(());
|
||||
}
|
||||
for chunk in ["first", "second"] {
|
||||
if host.deliver(chunk).await? == Demand::Detached {
|
||||
break;
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
})
|
||||
})
|
||||
}
|
||||
|
||||
/// Drives a `Stream` (async) or `SyncStream` to completion from a sync test.
|
||||
fn read_all(py: Python<'_>, stream: &Bound<'_, PyAny>, asynchronous: bool) -> Vec<String> {
|
||||
if !asynchronous {
|
||||
return stream
|
||||
.try_iter()
|
||||
.unwrap()
|
||||
.map(|chunk| chunk.unwrap().extract().unwrap())
|
||||
.collect();
|
||||
}
|
||||
std::iter::from_fn(|| {
|
||||
let stop = stream
|
||||
.call_method0("__anext__")
|
||||
.unwrap()
|
||||
.call_method1("send", (py.None(),))
|
||||
.unwrap_err();
|
||||
if stop.is_instance_of::<pyo3::exceptions::PyStopAsyncIteration>(py) {
|
||||
return None;
|
||||
}
|
||||
assert!(stop.is_instance_of::<pyo3::exceptions::PyStopIteration>(py));
|
||||
Some(stop.value(py).getattr("value").unwrap().extract().unwrap())
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn a_stream_carries_its_head_as_hidden_params_before_the_first_chunk() {
|
||||
let _guard = PYTHON_GLOBALS
|
||||
.lock()
|
||||
.unwrap_or_else(|error| error.into_inner());
|
||||
crate::initialize_python();
|
||||
Python::attach(|py| {
|
||||
install_lifecycle_module(py);
|
||||
for asynchronous in [false, true] {
|
||||
let log = Log::default();
|
||||
let adapter = SyntheticAdapter {
|
||||
log: Log(log.0.clone()),
|
||||
script: AdapterScript::Plain,
|
||||
};
|
||||
let handed = run_call(
|
||||
py,
|
||||
streaming_machine(),
|
||||
StreamingHost,
|
||||
Box::new(adapter),
|
||||
PyDict::new(py).unbind(),
|
||||
asynchronous,
|
||||
)
|
||||
.unwrap();
|
||||
let stream = if asynchronous {
|
||||
let stop = handed.call_method1(py, "send", (py.None(),)).unwrap_err();
|
||||
stop.value(py).getattr("value").unwrap()
|
||||
} else {
|
||||
handed.into_bound(py)
|
||||
};
|
||||
let hidden: std::collections::HashMap<
|
||||
String,
|
||||
std::collections::HashMap<String, String>,
|
||||
> = stream.getattr("_hidden_params").unwrap().extract().unwrap();
|
||||
assert_eq!(
|
||||
hidden["additional_headers"],
|
||||
std::collections::HashMap::from([(
|
||||
"request-id".to_string(),
|
||||
"req_1".to_string()
|
||||
)])
|
||||
);
|
||||
assert_eq!(log.entries(), ["started", "begin", "opened"]);
|
||||
assert_eq!(read_all(py, &stream, asynchronous), ["first", "second"]);
|
||||
}
|
||||
});
|
||||
}
|
||||
|
||||
fn failing_machine() -> CallMachine<Synthetic> {
|
||||
CallMachine::new(|host| {
|
||||
Box::pin(async move {
|
||||
|
|
@ -1202,6 +1372,13 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri
|
|||
) -> Result<(), InvokeError<Error>> {
|
||||
Err(missing_state().into())
|
||||
}
|
||||
fn head(
|
||||
&mut self,
|
||||
_: Python<'_>,
|
||||
head: std::convert::Infallible,
|
||||
) -> PyResult<Py<PyAny>> {
|
||||
match head {}
|
||||
}
|
||||
fn chunk(
|
||||
&mut self,
|
||||
_: Python<'_>,
|
||||
|
|
|
|||
|
|
@ -8,9 +8,9 @@ use pyo3::prelude::*;
|
|||
pub enum ExecutionStep {
|
||||
Return(Py<PyAny>),
|
||||
Await(Py<PyAny>),
|
||||
/// The call streams: the caller gets a stream over this execution, which stays
|
||||
/// suspended until the stream asks for a chunk.
|
||||
Open,
|
||||
/// The call streams: the caller gets a stream over this execution carrying this head,
|
||||
/// and the execution stays suspended until the stream asks for a chunk.
|
||||
Open(Py<PyAny>),
|
||||
Yield(Py<PyAny>),
|
||||
}
|
||||
|
||||
|
|
@ -75,7 +75,7 @@ impl Execution {
|
|||
let step = body.resume(result)?;
|
||||
let (tag, value, suspended) = match step {
|
||||
ExecutionStep::Await(value) => ("Await", value, true),
|
||||
ExecutionStep::Open => ("Open", py.None(), true),
|
||||
ExecutionStep::Open(head) => ("Open", head, true),
|
||||
ExecutionStep::Yield(value) => ("Yield", value, true),
|
||||
ExecutionStep::Return(value) => ("Complete", value, false),
|
||||
};
|
||||
|
|
|
|||
|
|
@ -3,7 +3,7 @@ use std::convert::Infallible;
|
|||
use bytes::Bytes;
|
||||
use litellm_core::messages::{
|
||||
Error,
|
||||
route::{Messages, MessagesCall, MessagesOutput},
|
||||
route::{Messages, MessagesCall, MessagesOutput, MessagesStreamHead},
|
||||
types::MessagesShaping,
|
||||
};
|
||||
use litellm_host_python::{InvokeError, ProtocolHost, from_py, lookup, to_py};
|
||||
|
|
@ -238,6 +238,13 @@ impl ProtocolHost for MessagesPythonHost {
|
|||
}
|
||||
}
|
||||
|
||||
fn head(&mut self, py: Python<'_>, head: MessagesStreamHead) -> PyResult<Py<PyAny>> {
|
||||
py.import(ROUTE_HOST_MODULE)?
|
||||
.getattr("stream_hidden_params")?
|
||||
.call1((to_py(py, &head.headers)?,))
|
||||
.map(Bound::unbind)
|
||||
}
|
||||
|
||||
fn chunk(&mut self, py: Python<'_>, chunk: Bytes) -> PyResult<Py<PyAny>> {
|
||||
Ok(PyBytes::new(py, &chunk).into_any().unbind())
|
||||
}
|
||||
|
|
|
|||
|
|
@ -4,14 +4,12 @@ use host::MessagesPythonHost;
|
|||
use litellm_callbacks_legacy_python::{
|
||||
LegacySurface, PassThroughStream, PublicCall, run_legacy_call,
|
||||
};
|
||||
use litellm_core::messages::route::{messages_machine, supports};
|
||||
use litellm_core::messages::route::messages_machine;
|
||||
use pyo3::{
|
||||
prelude::*,
|
||||
types::{PyDict, PyTuple},
|
||||
};
|
||||
|
||||
use crate::errors::RustBridgeDeclined;
|
||||
|
||||
const SURFACE: LegacySurface = LegacySurface {
|
||||
call_type: "anthropic_messages",
|
||||
input_description: "Messages",
|
||||
|
|
@ -28,17 +26,6 @@ fn run_messages(
|
|||
kwargs: Bound<'_, PyDict>,
|
||||
asynchronous: bool,
|
||||
) -> PyResult<Py<PyAny>> {
|
||||
let model: String = request.getattr("model")?.extract()?;
|
||||
let provider: Option<String> = request.getattr("custom_llm_provider")?.extract()?;
|
||||
let stream = request
|
||||
.getattr("stream")?
|
||||
.extract::<Option<bool>>()?
|
||||
.unwrap_or(false);
|
||||
if !supports(&model, provider.as_deref(), stream) {
|
||||
return Err(RustBridgeDeclined::new_err(
|
||||
"the Rust Messages route does not serve this provider",
|
||||
));
|
||||
}
|
||||
let secrets = crate::secrets::source(py)?;
|
||||
run_legacy_call(
|
||||
py,
|
||||
|
|
|
|||
|
|
@ -117,6 +117,10 @@ impl ProtocolHost for OcrPythonHost {
|
|||
.map(Bound::unbind)
|
||||
}
|
||||
|
||||
fn head(&mut self, _: Python<'_>, head: std::convert::Infallible) -> PyResult<Py<PyAny>> {
|
||||
match head {}
|
||||
}
|
||||
|
||||
fn chunk(&mut self, _: Python<'_>, chunk: std::convert::Infallible) -> PyResult<Py<PyAny>> {
|
||||
match chunk {}
|
||||
}
|
||||
|
|
|
|||
32
litellm-rust/crates/testkit/Cargo.toml
Normal file
32
litellm-rust/crates/testkit/Cargo.toml
Normal file
|
|
@ -0,0 +1,32 @@
|
|||
[package]
|
||||
name = "litellm-testkit"
|
||||
version = "0.1.0"
|
||||
edition.workspace = true
|
||||
license.workspace = true
|
||||
repository.workspace = true
|
||||
publish = false
|
||||
|
||||
[dependencies]
|
||||
flate2.workspace = true
|
||||
reqwest.workspace = true
|
||||
serde.workspace = true
|
||||
semver.workspace = true
|
||||
serde_json.workspace = true
|
||||
sha2.workspace = true
|
||||
tar.workspace = true
|
||||
target-lexicon.workspace = true
|
||||
thiserror.workspace = true
|
||||
tokio = { workspace = true, features = ["fs", "process"] }
|
||||
zip.workspace = true
|
||||
|
||||
[dev-dependencies]
|
||||
flate2.workspace = true
|
||||
rstest.workspace = true
|
||||
sha2.workspace = true
|
||||
tar.workspace = true
|
||||
target-lexicon.workspace = true
|
||||
futures-util.workspace = true
|
||||
tempfile.workspace = true
|
||||
tokio.workspace = true
|
||||
toml = "0.9"
|
||||
zip.workspace = true
|
||||
181
litellm-rust/crates/testkit/src/agent/claude.rs
Normal file
181
litellm-rust/crates/testkit/src/agent/claude.rs
Normal file
|
|
@ -0,0 +1,181 @@
|
|||
use std::collections::BTreeMap;
|
||||
use std::path::Path;
|
||||
|
||||
use semver::Version;
|
||||
use serde::Deserialize;
|
||||
|
||||
use super::{
|
||||
Configure, Drive, Install, LaunchSpec, Outcome, Prompt, Settings, Usage, Wire, env, json_lines,
|
||||
path_string,
|
||||
};
|
||||
use crate::install::release::parse;
|
||||
use crate::install::{Packaging, Release};
|
||||
use crate::{Error, Fetch, Target};
|
||||
|
||||
const RELEASES: &str = "https://downloads.claude.ai/claude-code-releases";
|
||||
|
||||
pub struct ClaudeCode;
|
||||
|
||||
#[derive(Deserialize)]
|
||||
struct Manifest {
|
||||
platforms: BTreeMap<String, Platform>,
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
struct Platform {
|
||||
checksum: String,
|
||||
}
|
||||
|
||||
impl Install for ClaudeCode {
|
||||
fn binary(&self) -> &'static str {
|
||||
"claude"
|
||||
}
|
||||
|
||||
async fn release(
|
||||
&self,
|
||||
fetch: &impl Fetch,
|
||||
version: &Version,
|
||||
target: Target,
|
||||
) -> Result<Release, Error> {
|
||||
let manifest_url = format!("{RELEASES}/{version}/manifest.json");
|
||||
let manifest: Manifest = parse(&manifest_url, &fetch.get(&manifest_url).await?)?;
|
||||
let key = format!(
|
||||
"{}-{}{}",
|
||||
target.os_name(),
|
||||
target.arch_name(),
|
||||
target.musl_suffix()
|
||||
);
|
||||
let platform = manifest
|
||||
.platforms
|
||||
.get(&key)
|
||||
.ok_or_else(|| Error::AssetNotFound(key.clone()))?;
|
||||
Ok(Release {
|
||||
url: format!("{RELEASES}/{version}/{key}/claude"),
|
||||
asset: key,
|
||||
sha256: platform.checksum.clone(),
|
||||
packaging: Packaging::Bare,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
impl Configure for ClaudeCode {
|
||||
fn configure(
|
||||
&self,
|
||||
_version: &Version,
|
||||
settings: &Settings,
|
||||
home: &Path,
|
||||
) -> Result<LaunchSpec, Error> {
|
||||
if settings.wire != Wire::Messages {
|
||||
return Err(Error::UnsupportedWire {
|
||||
agent: "claude",
|
||||
wire: settings.wire,
|
||||
});
|
||||
}
|
||||
Ok(LaunchSpec {
|
||||
env: env([
|
||||
("HOME", path_string(home)),
|
||||
("CLAUDE_CONFIG_DIR", path_string(&home.join(".claude"))),
|
||||
("ANTHROPIC_BASE_URL", settings.base_url.clone()),
|
||||
("ANTHROPIC_AUTH_TOKEN", settings.api_key.clone()),
|
||||
("ANTHROPIC_MODEL", settings.model.clone()),
|
||||
("DISABLE_AUTOUPDATER", "1".to_owned()),
|
||||
]),
|
||||
files: BTreeMap::new(),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
#[serde(tag = "type", rename_all = "snake_case")]
|
||||
enum Event {
|
||||
Assistant {
|
||||
message: AssistantMessage,
|
||||
},
|
||||
Result(Finished),
|
||||
#[serde(other)]
|
||||
Other,
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
struct AssistantMessage {
|
||||
content: Vec<Block>,
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
struct Block {
|
||||
#[serde(rename = "type")]
|
||||
kind: String,
|
||||
name: Option<String>,
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
struct Finished {
|
||||
is_error: bool,
|
||||
result: Option<String>,
|
||||
usage: Option<TokenUsage>,
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
struct TokenUsage {
|
||||
input_tokens: u64,
|
||||
output_tokens: u64,
|
||||
}
|
||||
|
||||
impl Drive for ClaudeCode {
|
||||
fn args(&self, _version: &Version, settings: &Settings, prompt: &Prompt) -> Vec<String> {
|
||||
let base = [
|
||||
"-p",
|
||||
&prompt.text,
|
||||
"--output-format",
|
||||
"stream-json",
|
||||
"--verbose",
|
||||
"--model",
|
||||
&settings.model,
|
||||
];
|
||||
let tools = ["--allowedTools", "Bash,Read,Write,Edit"];
|
||||
base.into_iter()
|
||||
.chain(tools.into_iter().filter(|_| prompt.allow_tools))
|
||||
.map(str::to_owned)
|
||||
.collect()
|
||||
}
|
||||
|
||||
fn parse(&self, _version: &Version, stdout: &str) -> Outcome {
|
||||
let events: Vec<Event> = json_lines(stdout).collect();
|
||||
let tool_calls = events
|
||||
.iter()
|
||||
.filter_map(|event| match event {
|
||||
Event::Assistant { message } => Some(&message.content),
|
||||
_ => None,
|
||||
})
|
||||
.flatten()
|
||||
.filter(|block| block.kind == "tool_use")
|
||||
.filter_map(|block| block.name.clone())
|
||||
.collect();
|
||||
let finished = events.into_iter().find_map(|event| match event {
|
||||
Event::Result(finished) => Some(finished),
|
||||
_ => None,
|
||||
});
|
||||
let Some(finished) = finished else {
|
||||
return Outcome {
|
||||
tool_calls,
|
||||
..Outcome::default()
|
||||
};
|
||||
};
|
||||
let result = finished.result.unwrap_or_default();
|
||||
let (text, errors) = if finished.is_error {
|
||||
(String::new(), vec![result])
|
||||
} else {
|
||||
(result, Vec::new())
|
||||
};
|
||||
Outcome {
|
||||
text,
|
||||
tool_calls,
|
||||
usage: finished.usage.map_or_else(Usage::default, |usage| Usage {
|
||||
input_tokens: usage.input_tokens,
|
||||
output_tokens: usage.output_tokens,
|
||||
}),
|
||||
errors,
|
||||
exit_code: None,
|
||||
}
|
||||
}
|
||||
}
|
||||
174
litellm-rust/crates/testkit/src/agent/codex.rs
Normal file
174
litellm-rust/crates/testkit/src/agent/codex.rs
Normal file
|
|
@ -0,0 +1,174 @@
|
|||
use std::collections::BTreeMap;
|
||||
use std::path::{Path, PathBuf};
|
||||
|
||||
use semver::Version;
|
||||
use serde::Deserialize;
|
||||
|
||||
use super::{
|
||||
Configure, Drive, Install, LaunchSpec, Outcome, Prompt, Settings, Usage, Wire, env, json_lines,
|
||||
path_string, quoted, v1,
|
||||
};
|
||||
use crate::install::release::github_release;
|
||||
use crate::install::{Packaging, Release};
|
||||
use crate::target::{Arch, Os};
|
||||
use crate::{Error, Fetch, Target};
|
||||
|
||||
const RELEASES: &str = "https://api.github.com/repos/openai/codex/releases/tags";
|
||||
|
||||
pub struct Codex;
|
||||
|
||||
fn triple(target: Target) -> String {
|
||||
let arch = match target.arch {
|
||||
Arch::Aarch64 => "aarch64",
|
||||
Arch::X86_64 => "x86_64",
|
||||
};
|
||||
match target.os {
|
||||
Os::Macos => format!("{arch}-apple-darwin"),
|
||||
Os::Linux => format!("{arch}-unknown-linux-musl"),
|
||||
}
|
||||
}
|
||||
|
||||
impl Install for Codex {
|
||||
fn binary(&self) -> &'static str {
|
||||
"codex"
|
||||
}
|
||||
|
||||
async fn release(
|
||||
&self,
|
||||
fetch: &impl Fetch,
|
||||
version: &Version,
|
||||
target: Target,
|
||||
) -> Result<Release, Error> {
|
||||
let triple = triple(target);
|
||||
github_release(
|
||||
fetch,
|
||||
RELEASES,
|
||||
&format!("rust-v{version}"),
|
||||
&format!("codex-{triple}.tar.gz"),
|
||||
Packaging::TarGz {
|
||||
member: format!("codex-{triple}"),
|
||||
},
|
||||
)
|
||||
.await
|
||||
}
|
||||
}
|
||||
|
||||
impl Configure for Codex {
|
||||
fn configure(
|
||||
&self,
|
||||
_version: &Version,
|
||||
settings: &Settings,
|
||||
home: &Path,
|
||||
) -> Result<LaunchSpec, Error> {
|
||||
if settings.wire != Wire::Responses {
|
||||
return Err(Error::UnsupportedWire {
|
||||
agent: "codex",
|
||||
wire: settings.wire,
|
||||
});
|
||||
}
|
||||
let config = format!(
|
||||
"model = {model}\nmodel_provider = \"litellm\"\n\n[model_providers.litellm]\nname = \"LiteLLM\"\nbase_url = {base_url}\nenv_key = \"LITELLM_API_KEY\"\nwire_api = \"responses\"\n",
|
||||
model = quoted(&settings.model),
|
||||
base_url = quoted(&v1(settings)),
|
||||
);
|
||||
Ok(LaunchSpec {
|
||||
env: env([
|
||||
("HOME", path_string(home)),
|
||||
("CODEX_HOME", path_string(&home.join(".codex"))),
|
||||
("LITELLM_API_KEY", settings.api_key.clone()),
|
||||
]),
|
||||
files: BTreeMap::from([(PathBuf::from(".codex/config.toml"), config)]),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
enum EventKind {
|
||||
#[serde(rename = "item.completed")]
|
||||
ItemCompleted,
|
||||
#[serde(rename = "turn.completed")]
|
||||
TurnCompleted,
|
||||
#[serde(rename = "turn.failed")]
|
||||
TurnFailed,
|
||||
#[serde(other)]
|
||||
Other,
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
struct Event {
|
||||
#[serde(rename = "type")]
|
||||
kind: EventKind,
|
||||
item: Option<Item>,
|
||||
usage: Option<TokenUsage>,
|
||||
error: Option<Failure>,
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
struct Item {
|
||||
#[serde(rename = "type")]
|
||||
kind: String,
|
||||
text: Option<String>,
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
struct TokenUsage {
|
||||
input_tokens: u64,
|
||||
output_tokens: u64,
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
struct Failure {
|
||||
message: String,
|
||||
}
|
||||
|
||||
const NON_TOOL_ITEMS: [&str; 3] = ["agent_message", "reasoning", "error"];
|
||||
|
||||
impl Drive for Codex {
|
||||
fn args(&self, _version: &Version, _settings: &Settings, prompt: &Prompt) -> Vec<String> {
|
||||
let sandbox = ["--sandbox", "workspace-write"];
|
||||
["exec", "--json", "--skip-git-repo-check"]
|
||||
.into_iter()
|
||||
.chain(sandbox.into_iter().filter(|_| prompt.allow_tools))
|
||||
.chain([prompt.text.as_str()])
|
||||
.map(str::to_owned)
|
||||
.collect()
|
||||
}
|
||||
|
||||
fn parse(&self, _version: &Version, stdout: &str) -> Outcome {
|
||||
let events: Vec<Event> = json_lines(stdout).collect();
|
||||
let items: Vec<&Item> = events
|
||||
.iter()
|
||||
.filter(|event| matches!(event.kind, EventKind::ItemCompleted))
|
||||
.filter_map(|event| event.item.as_ref())
|
||||
.collect();
|
||||
Outcome {
|
||||
text: items
|
||||
.iter()
|
||||
.rev()
|
||||
.find(|item| item.kind == "agent_message")
|
||||
.and_then(|item| item.text.clone())
|
||||
.unwrap_or_default(),
|
||||
tool_calls: items
|
||||
.iter()
|
||||
.filter(|item| !NON_TOOL_ITEMS.contains(&item.kind.as_str()))
|
||||
.map(|item| item.kind.clone())
|
||||
.collect(),
|
||||
usage: events
|
||||
.iter()
|
||||
.filter(|event| matches!(event.kind, EventKind::TurnCompleted))
|
||||
.filter_map(|event| event.usage.as_ref())
|
||||
.map(|usage| Usage {
|
||||
input_tokens: usage.input_tokens,
|
||||
output_tokens: usage.output_tokens,
|
||||
})
|
||||
.fold(Usage::default(), |total, turn| total + turn),
|
||||
errors: events
|
||||
.iter()
|
||||
.filter(|event| matches!(event.kind, EventKind::TurnFailed))
|
||||
.filter_map(|event| event.error.as_ref())
|
||||
.map(|failure| failure.message.clone())
|
||||
.collect(),
|
||||
exit_code: None,
|
||||
}
|
||||
}
|
||||
}
|
||||
69
litellm-rust/crates/testkit/src/agent/configure.rs
Normal file
69
litellm-rust/crates/testkit/src/agent/configure.rs
Normal file
|
|
@ -0,0 +1,69 @@
|
|||
use std::collections::BTreeMap;
|
||||
use std::path::{Path, PathBuf};
|
||||
|
||||
use semver::Version;
|
||||
|
||||
use crate::Error;
|
||||
|
||||
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
|
||||
pub enum Wire {
|
||||
ChatCompletions,
|
||||
Messages,
|
||||
Responses,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq)]
|
||||
pub struct Settings {
|
||||
pub base_url: String,
|
||||
pub api_key: String,
|
||||
pub model: String,
|
||||
pub wire: Wire,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq)]
|
||||
pub struct LaunchSpec {
|
||||
pub env: BTreeMap<String, String>,
|
||||
pub files: BTreeMap<PathBuf, String>,
|
||||
}
|
||||
|
||||
impl LaunchSpec {
|
||||
pub fn write_files(&self, home: &Path) -> std::io::Result<()> {
|
||||
self.files.iter().try_for_each(|(relative, contents)| {
|
||||
let path = home.join(relative);
|
||||
if let Some(parent) = path.parent() {
|
||||
std::fs::create_dir_all(parent)?;
|
||||
}
|
||||
std::fs::write(path, contents)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
pub trait Configure {
|
||||
fn configure(
|
||||
&self,
|
||||
version: &Version,
|
||||
settings: &Settings,
|
||||
home: &Path,
|
||||
) -> Result<LaunchSpec, Error>;
|
||||
}
|
||||
|
||||
pub(crate) fn env(
|
||||
pairs: impl IntoIterator<Item = (&'static str, String)>,
|
||||
) -> BTreeMap<String, String> {
|
||||
pairs
|
||||
.into_iter()
|
||||
.map(|(key, value)| (key.to_owned(), value))
|
||||
.collect()
|
||||
}
|
||||
|
||||
pub(crate) fn path_string(path: &Path) -> String {
|
||||
path.to_string_lossy().into_owned()
|
||||
}
|
||||
|
||||
pub(crate) fn quoted(value: &str) -> String {
|
||||
serde_json::Value::from(value).to_string()
|
||||
}
|
||||
|
||||
pub(crate) fn v1(settings: &Settings) -> String {
|
||||
format!("{}/v1", settings.base_url.trim_end_matches('/'))
|
||||
}
|
||||
57
litellm-rust/crates/testkit/src/agent/drive.rs
Normal file
57
litellm-rust/crates/testkit/src/agent/drive.rs
Normal file
|
|
@ -0,0 +1,57 @@
|
|||
use std::ops::Add;
|
||||
|
||||
use semver::Version;
|
||||
|
||||
use crate::Settings;
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq)]
|
||||
pub struct Prompt {
|
||||
pub text: String,
|
||||
pub allow_tools: bool,
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
|
||||
pub struct Usage {
|
||||
pub input_tokens: u64,
|
||||
pub output_tokens: u64,
|
||||
}
|
||||
|
||||
impl Add for Usage {
|
||||
type Output = Self;
|
||||
|
||||
fn add(self, other: Self) -> Self {
|
||||
Self {
|
||||
input_tokens: self.input_tokens + other.input_tokens,
|
||||
output_tokens: self.output_tokens + other.output_tokens,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Default, PartialEq, Eq)]
|
||||
pub struct Outcome {
|
||||
pub text: String,
|
||||
pub tool_calls: Vec<String>,
|
||||
pub usage: Usage,
|
||||
pub errors: Vec<String>,
|
||||
pub exit_code: Option<i32>,
|
||||
}
|
||||
|
||||
impl Outcome {
|
||||
pub fn succeeded(&self) -> bool {
|
||||
self.exit_code == Some(0) && self.errors.is_empty()
|
||||
}
|
||||
}
|
||||
|
||||
pub trait Drive {
|
||||
fn args(&self, version: &Version, settings: &Settings, prompt: &Prompt) -> Vec<String>;
|
||||
|
||||
fn parse(&self, version: &Version, stdout: &str) -> Outcome;
|
||||
}
|
||||
|
||||
pub(crate) fn json_lines<'a, T: serde::de::DeserializeOwned + 'a>(
|
||||
stdout: &'a str,
|
||||
) -> impl Iterator<Item = T> + 'a {
|
||||
stdout
|
||||
.lines()
|
||||
.filter_map(|line| serde_json::from_str(line).ok())
|
||||
}
|
||||
17
litellm-rust/crates/testkit/src/agent/install.rs
Normal file
17
litellm-rust/crates/testkit/src/agent/install.rs
Normal file
|
|
@ -0,0 +1,17 @@
|
|||
use std::future::Future;
|
||||
|
||||
use semver::Version;
|
||||
|
||||
use crate::install::Release;
|
||||
use crate::{Error, Fetch, Target};
|
||||
|
||||
pub trait Install: Sync {
|
||||
fn binary(&self) -> &'static str;
|
||||
|
||||
fn release(
|
||||
&self,
|
||||
fetch: &impl Fetch,
|
||||
version: &Version,
|
||||
target: Target,
|
||||
) -> impl Future<Output = Result<Release, Error>> + Send;
|
||||
}
|
||||
20
litellm-rust/crates/testkit/src/agent/mod.rs
Normal file
20
litellm-rust/crates/testkit/src/agent/mod.rs
Normal file
|
|
@ -0,0 +1,20 @@
|
|||
mod claude;
|
||||
mod codex;
|
||||
mod configure;
|
||||
mod drive;
|
||||
mod install;
|
||||
mod opencode;
|
||||
|
||||
pub use claude::ClaudeCode;
|
||||
pub use codex::Codex;
|
||||
pub use configure::{Configure, LaunchSpec, Settings, Wire};
|
||||
pub use drive::{Drive, Outcome, Prompt, Usage};
|
||||
pub use install::Install;
|
||||
pub use opencode::Opencode;
|
||||
|
||||
pub(crate) use configure::{env, path_string, quoted, v1};
|
||||
pub(crate) use drive::json_lines;
|
||||
|
||||
pub trait Agent: Install + Configure + Drive {}
|
||||
|
||||
impl<T: Install + Configure + Drive> Agent for T {}
|
||||
187
litellm-rust/crates/testkit/src/agent/opencode.rs
Normal file
187
litellm-rust/crates/testkit/src/agent/opencode.rs
Normal file
|
|
@ -0,0 +1,187 @@
|
|||
use std::collections::BTreeMap;
|
||||
use std::path::{Path, PathBuf};
|
||||
|
||||
use semver::Version;
|
||||
use serde::Deserialize;
|
||||
|
||||
use super::{
|
||||
Configure, Drive, Install, LaunchSpec, Outcome, Prompt, Settings, Usage, Wire, env, json_lines,
|
||||
path_string, v1,
|
||||
};
|
||||
use crate::install::release::github_release;
|
||||
use crate::install::{Packaging, Release};
|
||||
use crate::target::Os;
|
||||
use crate::{Error, Fetch, Target};
|
||||
|
||||
const RELEASES: &str = "https://api.github.com/repos/sst/opencode/releases/tags";
|
||||
|
||||
pub struct Opencode;
|
||||
|
||||
impl Install for Opencode {
|
||||
fn binary(&self) -> &'static str {
|
||||
"opencode"
|
||||
}
|
||||
|
||||
async fn release(
|
||||
&self,
|
||||
fetch: &impl Fetch,
|
||||
version: &Version,
|
||||
target: Target,
|
||||
) -> Result<Release, Error> {
|
||||
let stem = format!(
|
||||
"opencode-{}-{}{}",
|
||||
target.os_name(),
|
||||
target.arch_name(),
|
||||
target.musl_suffix()
|
||||
);
|
||||
let member = "opencode".to_owned();
|
||||
let (asset, packaging) = match target.os {
|
||||
Os::Macos => (format!("{stem}.zip"), Packaging::Zip { member }),
|
||||
Os::Linux => (format!("{stem}.tar.gz"), Packaging::TarGz { member }),
|
||||
};
|
||||
github_release(fetch, RELEASES, &format!("v{version}"), &asset, packaging).await
|
||||
}
|
||||
}
|
||||
|
||||
impl Configure for Opencode {
|
||||
fn configure(
|
||||
&self,
|
||||
_version: &Version,
|
||||
settings: &Settings,
|
||||
home: &Path,
|
||||
) -> Result<LaunchSpec, Error> {
|
||||
let npm = match settings.wire {
|
||||
Wire::ChatCompletions => "@ai-sdk/openai-compatible",
|
||||
Wire::Responses => "@ai-sdk/openai",
|
||||
Wire::Messages => "@ai-sdk/anthropic",
|
||||
};
|
||||
let config = serde_json::json!({
|
||||
"$schema": "https://opencode.ai/config.json",
|
||||
"model": format!("litellm/{}", settings.model),
|
||||
"provider": {
|
||||
"litellm": {
|
||||
"npm": npm,
|
||||
"name": "LiteLLM",
|
||||
"options": { "baseURL": v1(settings), "apiKey": settings.api_key },
|
||||
"models": { settings.model.clone(): { "name": settings.model } },
|
||||
}
|
||||
},
|
||||
});
|
||||
Ok(LaunchSpec {
|
||||
env: env([
|
||||
("HOME", path_string(home)),
|
||||
("XDG_CONFIG_HOME", path_string(&home.join(".config"))),
|
||||
("XDG_DATA_HOME", path_string(&home.join(".local/share"))),
|
||||
("OPENCODE_DISABLE_AUTOUPDATE", "true".to_owned()),
|
||||
]),
|
||||
files: BTreeMap::from([(
|
||||
PathBuf::from(".config/opencode/opencode.json"),
|
||||
config.to_string(),
|
||||
)]),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
#[serde(tag = "type", rename_all = "snake_case")]
|
||||
enum Event {
|
||||
Text {
|
||||
part: TextPart,
|
||||
},
|
||||
ToolUse {
|
||||
part: ToolPart,
|
||||
},
|
||||
StepFinish {
|
||||
part: StepFinish,
|
||||
},
|
||||
Error {
|
||||
error: Failure,
|
||||
},
|
||||
#[serde(other)]
|
||||
Other,
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
struct TextPart {
|
||||
text: String,
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
struct ToolPart {
|
||||
tool: String,
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
struct StepFinish {
|
||||
tokens: Tokens,
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
struct Tokens {
|
||||
input: u64,
|
||||
output: u64,
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
struct Failure {
|
||||
name: String,
|
||||
data: Option<FailureData>,
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
struct FailureData {
|
||||
message: Option<String>,
|
||||
}
|
||||
|
||||
impl Drive for Opencode {
|
||||
fn args(&self, _version: &Version, _settings: &Settings, prompt: &Prompt) -> Vec<String> {
|
||||
["run", "--format", "json", &prompt.text]
|
||||
.map(str::to_owned)
|
||||
.to_vec()
|
||||
}
|
||||
|
||||
fn parse(&self, _version: &Version, stdout: &str) -> Outcome {
|
||||
let events: Vec<Event> = json_lines(stdout).collect();
|
||||
Outcome {
|
||||
text: events
|
||||
.iter()
|
||||
.rev()
|
||||
.find_map(|event| match event {
|
||||
Event::Text { part } => Some(part.text.clone()),
|
||||
_ => None,
|
||||
})
|
||||
.unwrap_or_default(),
|
||||
tool_calls: events
|
||||
.iter()
|
||||
.filter_map(|event| match event {
|
||||
Event::ToolUse { part } => Some(part.tool.clone()),
|
||||
_ => None,
|
||||
})
|
||||
.collect(),
|
||||
usage: events
|
||||
.iter()
|
||||
.filter_map(|event| match event {
|
||||
Event::StepFinish { part } => Some(Usage {
|
||||
input_tokens: part.tokens.input,
|
||||
output_tokens: part.tokens.output,
|
||||
}),
|
||||
_ => None,
|
||||
})
|
||||
.fold(Usage::default(), |total, step| total + step),
|
||||
errors: events
|
||||
.iter()
|
||||
.filter_map(|event| match event {
|
||||
Event::Error { error } => Some(
|
||||
error
|
||||
.data
|
||||
.as_ref()
|
||||
.and_then(|data| data.message.clone())
|
||||
.unwrap_or_else(|| error.name.clone()),
|
||||
),
|
||||
_ => None,
|
||||
})
|
||||
.collect(),
|
||||
exit_code: None,
|
||||
}
|
||||
}
|
||||
}
|
||||
56
litellm-rust/crates/testkit/src/error.rs
Normal file
56
litellm-rust/crates/testkit/src/error.rs
Normal file
|
|
@ -0,0 +1,56 @@
|
|||
use std::io;
|
||||
use std::path::PathBuf;
|
||||
|
||||
use thiserror::Error;
|
||||
|
||||
use crate::Wire;
|
||||
|
||||
#[derive(Debug, Error)]
|
||||
pub enum Error {
|
||||
#[error("unsupported target {0}")]
|
||||
UnsupportedTarget(String),
|
||||
#[error("{0} is not a plain x.y.z release version")]
|
||||
InvalidVersion(String),
|
||||
#[error("request to {url} failed")]
|
||||
Request {
|
||||
url: String,
|
||||
#[source]
|
||||
source: reqwest::Error,
|
||||
},
|
||||
#[error("{url} answered with status {status}")]
|
||||
Status { url: String, status: u16 },
|
||||
#[error("release metadata at {url} is malformed")]
|
||||
Metadata {
|
||||
url: String,
|
||||
#[source]
|
||||
source: serde_json::Error,
|
||||
},
|
||||
#[error("release has no asset named {0}")]
|
||||
AssetNotFound(String),
|
||||
#[error("release publishes no sha256 for {0}")]
|
||||
MissingChecksum(String),
|
||||
#[error("sha256 mismatch for {asset}: expected {expected}, got {actual}")]
|
||||
ChecksumMismatch {
|
||||
asset: String,
|
||||
expected: String,
|
||||
actual: String,
|
||||
},
|
||||
#[error("archive does not contain {0}")]
|
||||
ArchiveMemberNotFound(String),
|
||||
#[error("archive is unreadable")]
|
||||
Archive(#[source] io::Error),
|
||||
#[error("zip archive is unreadable")]
|
||||
Zip(#[from] zip::result::ZipError),
|
||||
#[error("{binary} reports version '{reported}', expected {expected}")]
|
||||
VersionMismatch {
|
||||
binary: PathBuf,
|
||||
expected: String,
|
||||
reported: String,
|
||||
},
|
||||
#[error("{agent} cannot talk to the gateway over {wire:?}")]
|
||||
UnsupportedWire { agent: &'static str, wire: Wire },
|
||||
#[error("agent did not finish within {0:?}")]
|
||||
Timeout(std::time::Duration),
|
||||
#[error("io failure")]
|
||||
Io(#[from] io::Error),
|
||||
}
|
||||
52
litellm-rust/crates/testkit/src/install/archive.rs
Normal file
52
litellm-rust/crates/testkit/src/install/archive.rs
Normal file
|
|
@ -0,0 +1,52 @@
|
|||
use std::io::{Cursor, Read};
|
||||
|
||||
use flate2::read::GzDecoder;
|
||||
use sha2::{Digest, Sha256};
|
||||
|
||||
use super::release::Packaging;
|
||||
use crate::Error;
|
||||
|
||||
pub(crate) fn verify_sha256(asset: &str, expected: &str, bytes: &[u8]) -> Result<(), Error> {
|
||||
let actual = format!("{:x}", Sha256::digest(bytes));
|
||||
if actual.eq_ignore_ascii_case(expected) {
|
||||
return Ok(());
|
||||
}
|
||||
Err(Error::ChecksumMismatch {
|
||||
asset: asset.to_owned(),
|
||||
expected: expected.to_owned(),
|
||||
actual,
|
||||
})
|
||||
}
|
||||
|
||||
pub(crate) fn extract_binary(packaging: &Packaging, bytes: &[u8]) -> Result<Vec<u8>, Error> {
|
||||
match packaging {
|
||||
Packaging::Bare => Ok(bytes.to_vec()),
|
||||
Packaging::TarGz { member } => extract_tar_gz(member, bytes),
|
||||
Packaging::Zip { member } => extract_zip(member, bytes),
|
||||
}
|
||||
}
|
||||
|
||||
fn extract_tar_gz(member: &str, bytes: &[u8]) -> Result<Vec<u8>, Error> {
|
||||
let mut archive = tar::Archive::new(GzDecoder::new(bytes));
|
||||
for entry in archive.entries().map_err(Error::Archive)? {
|
||||
let mut entry = entry.map_err(Error::Archive)?;
|
||||
let path = entry.path().map_err(Error::Archive)?;
|
||||
if path.file_name().is_some_and(|name| name == member) {
|
||||
let mut binary = Vec::new();
|
||||
entry.read_to_end(&mut binary).map_err(Error::Archive)?;
|
||||
return Ok(binary);
|
||||
}
|
||||
}
|
||||
Err(Error::ArchiveMemberNotFound(member.to_owned()))
|
||||
}
|
||||
|
||||
fn extract_zip(member: &str, bytes: &[u8]) -> Result<Vec<u8>, Error> {
|
||||
let mut archive = zip::ZipArchive::new(Cursor::new(bytes))?;
|
||||
let mut file = archive.by_name(member).map_err(|error| match error {
|
||||
zip::result::ZipError::FileNotFound => Error::ArchiveMemberNotFound(member.to_owned()),
|
||||
other => Error::Zip(other),
|
||||
})?;
|
||||
let mut binary = Vec::new();
|
||||
file.read_to_end(&mut binary).map_err(Error::Archive)?;
|
||||
Ok(binary)
|
||||
}
|
||||
55
litellm-rust/crates/testkit/src/install/fetch.rs
Normal file
55
litellm-rust/crates/testkit/src/install/fetch.rs
Normal file
|
|
@ -0,0 +1,55 @@
|
|||
use std::future::Future;
|
||||
|
||||
use crate::Error;
|
||||
|
||||
pub trait Fetch: Sync {
|
||||
fn get(&self, url: &str) -> impl Future<Output = Result<Vec<u8>, Error>> + Send;
|
||||
}
|
||||
|
||||
pub struct HttpFetch {
|
||||
client: reqwest::Client,
|
||||
github_token: Option<String>,
|
||||
}
|
||||
|
||||
impl HttpFetch {
|
||||
pub fn new(github_token: Option<String>) -> Self {
|
||||
Self {
|
||||
client: reqwest::Client::new(),
|
||||
github_token,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn from_env() -> Self {
|
||||
Self::new(std::env::var("GITHUB_TOKEN").ok())
|
||||
}
|
||||
}
|
||||
|
||||
impl Fetch for HttpFetch {
|
||||
async fn get(&self, url: &str) -> Result<Vec<u8>, Error> {
|
||||
let request = self
|
||||
.client
|
||||
.get(url)
|
||||
.header("user-agent", "litellm-testkit")
|
||||
.header("accept", "application/json, application/octet-stream");
|
||||
let request = match (
|
||||
&self.github_token,
|
||||
url.starts_with("https://api.github.com/"),
|
||||
) {
|
||||
(Some(token), true) => request.bearer_auth(token),
|
||||
_ => request,
|
||||
};
|
||||
let request_error = |source| Error::Request {
|
||||
url: url.to_owned(),
|
||||
source,
|
||||
};
|
||||
let response = request.send().await.map_err(request_error)?;
|
||||
let status = response.status();
|
||||
if !status.is_success() {
|
||||
return Err(Error::Status {
|
||||
url: url.to_owned(),
|
||||
status: status.as_u16(),
|
||||
});
|
||||
}
|
||||
Ok(response.bytes().await.map_err(request_error)?.to_vec())
|
||||
}
|
||||
}
|
||||
118
litellm-rust/crates/testkit/src/install/mod.rs
Normal file
118
litellm-rust/crates/testkit/src/install/mod.rs
Normal file
|
|
@ -0,0 +1,118 @@
|
|||
mod archive;
|
||||
mod fetch;
|
||||
pub(crate) mod release;
|
||||
|
||||
use std::os::unix::fs::PermissionsExt;
|
||||
use std::path::{Path, PathBuf};
|
||||
use std::process::Stdio;
|
||||
use std::sync::atomic::{AtomicU64, Ordering};
|
||||
|
||||
use semver::Version;
|
||||
use tokio::fs;
|
||||
use tokio::process::Command;
|
||||
|
||||
use crate::{Error, Install, Target};
|
||||
use archive::{extract_binary, verify_sha256};
|
||||
|
||||
static STAGING_COUNTER: AtomicU64 = AtomicU64::new(0);
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq)]
|
||||
pub struct Installed {
|
||||
pub version: Version,
|
||||
pub binary: PathBuf,
|
||||
}
|
||||
|
||||
pub struct Installer<F> {
|
||||
fetch: F,
|
||||
cache_root: PathBuf,
|
||||
target: Target,
|
||||
}
|
||||
|
||||
impl<F: Fetch> Installer<F> {
|
||||
pub fn new(fetch: F, cache_root: impl Into<PathBuf>, target: Target) -> Self {
|
||||
Self {
|
||||
fetch,
|
||||
cache_root: cache_root.into(),
|
||||
target,
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn install(
|
||||
&self,
|
||||
agent: &impl Install,
|
||||
version: &Version,
|
||||
) -> Result<Installed, Error> {
|
||||
validate_release(version)?;
|
||||
let dir = self
|
||||
.cache_root
|
||||
.join(agent.binary())
|
||||
.join(version.to_string());
|
||||
let binary = dir.join(agent.binary());
|
||||
let installed = Installed {
|
||||
version: version.clone(),
|
||||
binary: binary.clone(),
|
||||
};
|
||||
if fs::try_exists(&binary).await? && probe_version(&binary, version).await.is_ok() {
|
||||
return Ok(installed);
|
||||
}
|
||||
|
||||
let release = agent.release(&self.fetch, version, self.target).await?;
|
||||
let archive = self.fetch.get(&release.url).await?;
|
||||
verify_sha256(&release.asset, &release.sha256, &archive)?;
|
||||
let contents = extract_binary(&release.packaging, &archive)?;
|
||||
|
||||
fs::create_dir_all(&dir).await?;
|
||||
let staging = dir.join(format!(
|
||||
".{}.{}.{}.partial",
|
||||
agent.binary(),
|
||||
std::process::id(),
|
||||
STAGING_COUNTER.fetch_add(1, Ordering::Relaxed)
|
||||
));
|
||||
fs::write(&staging, contents).await?;
|
||||
fs::set_permissions(&staging, std::fs::Permissions::from_mode(0o755)).await?;
|
||||
fs::rename(&staging, &binary).await?;
|
||||
|
||||
match probe_version(&binary, version).await {
|
||||
Ok(()) => Ok(installed),
|
||||
Err(error) => {
|
||||
fs::remove_file(&binary).await?;
|
||||
Err(error)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn validate_release(version: &Version) -> Result<(), Error> {
|
||||
if version.pre.is_empty() && version.build.is_empty() {
|
||||
return Ok(());
|
||||
}
|
||||
Err(Error::InvalidVersion(version.to_string()))
|
||||
}
|
||||
|
||||
async fn probe_version(binary: &Path, expected: &Version) -> Result<(), Error> {
|
||||
let home = std::env::temp_dir();
|
||||
let output = Command::new(binary)
|
||||
.arg("--version")
|
||||
.env_clear()
|
||||
.env("HOME", home)
|
||||
.env("DISABLE_AUTOUPDATER", "1")
|
||||
.stdin(Stdio::null())
|
||||
.output()
|
||||
.await?;
|
||||
let stdout = String::from_utf8_lossy(&output.stdout);
|
||||
if stdout
|
||||
.split_whitespace()
|
||||
.filter_map(|token| Version::parse(token).ok())
|
||||
.any(|reported| &reported == expected)
|
||||
{
|
||||
return Ok(());
|
||||
}
|
||||
Err(Error::VersionMismatch {
|
||||
binary: binary.to_owned(),
|
||||
expected: expected.to_string(),
|
||||
reported: stdout.trim().to_owned(),
|
||||
})
|
||||
}
|
||||
|
||||
pub use fetch::{Fetch, HttpFetch};
|
||||
pub use release::{Packaging, Release};
|
||||
65
litellm-rust/crates/testkit/src/install/release.rs
Normal file
65
litellm-rust/crates/testkit/src/install/release.rs
Normal file
|
|
@ -0,0 +1,65 @@
|
|||
use serde::Deserialize;
|
||||
|
||||
use crate::{Error, Fetch};
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq)]
|
||||
pub enum Packaging {
|
||||
Bare,
|
||||
TarGz { member: String },
|
||||
Zip { member: String },
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq)]
|
||||
pub struct Release {
|
||||
pub asset: String,
|
||||
pub url: String,
|
||||
pub sha256: String,
|
||||
pub packaging: Packaging,
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
struct GithubRelease {
|
||||
assets: Vec<GithubAsset>,
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
struct GithubAsset {
|
||||
name: String,
|
||||
digest: Option<String>,
|
||||
browser_download_url: String,
|
||||
}
|
||||
|
||||
pub(crate) async fn github_release(
|
||||
fetch: &impl Fetch,
|
||||
releases_url: &str,
|
||||
tag: &str,
|
||||
asset_name: &str,
|
||||
packaging: Packaging,
|
||||
) -> Result<Release, Error> {
|
||||
let url = format!("{releases_url}/{tag}");
|
||||
let release: GithubRelease = parse(&url, &fetch.get(&url).await?)?;
|
||||
let asset = release
|
||||
.assets
|
||||
.into_iter()
|
||||
.find(|asset| asset.name == asset_name)
|
||||
.ok_or_else(|| Error::AssetNotFound(asset_name.to_owned()))?;
|
||||
let sha256 = asset
|
||||
.digest
|
||||
.as_deref()
|
||||
.and_then(|digest| digest.strip_prefix("sha256:"))
|
||||
.ok_or_else(|| Error::MissingChecksum(asset_name.to_owned()))?
|
||||
.to_owned();
|
||||
Ok(Release {
|
||||
asset: asset.name,
|
||||
url: asset.browser_download_url,
|
||||
sha256,
|
||||
packaging,
|
||||
})
|
||||
}
|
||||
|
||||
pub(crate) fn parse<T: for<'de> Deserialize<'de>>(url: &str, body: &[u8]) -> Result<T, Error> {
|
||||
serde_json::from_slice(body).map_err(|source| Error::Metadata {
|
||||
url: url.to_owned(),
|
||||
source,
|
||||
})
|
||||
}
|
||||
15
litellm-rust/crates/testkit/src/lib.rs
Normal file
15
litellm-rust/crates/testkit/src/lib.rs
Normal file
|
|
@ -0,0 +1,15 @@
|
|||
mod agent;
|
||||
mod error;
|
||||
mod install;
|
||||
mod session;
|
||||
mod target;
|
||||
|
||||
pub use agent::{
|
||||
Agent, ClaudeCode, Codex, Configure, Drive, Install, LaunchSpec, Opencode, Outcome, Prompt,
|
||||
Settings, Usage, Wire,
|
||||
};
|
||||
pub use error::Error;
|
||||
pub use install::{Fetch, HttpFetch, Installed, Installer, Packaging, Release};
|
||||
pub use semver::Version;
|
||||
pub use session::Session;
|
||||
pub use target::{Arch, Os, Target};
|
||||
76
litellm-rust/crates/testkit/src/session.rs
Normal file
76
litellm-rust/crates/testkit/src/session.rs
Normal file
|
|
@ -0,0 +1,76 @@
|
|||
use std::collections::BTreeMap;
|
||||
use std::path::PathBuf;
|
||||
use std::process::Stdio;
|
||||
use std::time::Duration;
|
||||
|
||||
use semver::Version;
|
||||
use tokio::process::Command;
|
||||
use tokio::time::timeout;
|
||||
|
||||
use crate::{Configure, Drive, Error, Installed, Outcome, Prompt, Settings};
|
||||
|
||||
const STDERR_LIMIT_CHARS: usize = 2000;
|
||||
|
||||
pub struct Session {
|
||||
binary: PathBuf,
|
||||
home: PathBuf,
|
||||
version: Version,
|
||||
settings: Settings,
|
||||
env: BTreeMap<String, String>,
|
||||
}
|
||||
|
||||
impl Session {
|
||||
pub fn prepare(
|
||||
agent: &impl Configure,
|
||||
installed: &Installed,
|
||||
settings: Settings,
|
||||
home: impl Into<PathBuf>,
|
||||
) -> Result<Self, Error> {
|
||||
let home = home.into();
|
||||
let spec = agent.configure(&installed.version, &settings, &home)?;
|
||||
spec.write_files(&home)?;
|
||||
Ok(Self {
|
||||
binary: installed.binary.clone(),
|
||||
home,
|
||||
version: installed.version.clone(),
|
||||
settings,
|
||||
env: spec.env,
|
||||
})
|
||||
}
|
||||
|
||||
pub async fn run(
|
||||
&self,
|
||||
agent: &impl Drive,
|
||||
prompt: &Prompt,
|
||||
limit: Duration,
|
||||
) -> Result<Outcome, Error> {
|
||||
let child = Command::new(&self.binary)
|
||||
.args(agent.args(&self.version, &self.settings, prompt))
|
||||
.env_clear()
|
||||
.env("PATH", "/usr/bin:/bin")
|
||||
.envs(&self.env)
|
||||
.current_dir(&self.home)
|
||||
.stdin(Stdio::null())
|
||||
.kill_on_drop(true)
|
||||
.output();
|
||||
let output = timeout(limit, child)
|
||||
.await
|
||||
.map_err(|_| Error::Timeout(limit))??;
|
||||
let parsed = agent.parse(&self.version, &String::from_utf8_lossy(&output.stdout));
|
||||
let failed_silently = !output.status.success() && parsed.errors.is_empty();
|
||||
Ok(Outcome {
|
||||
errors: if failed_silently {
|
||||
vec![
|
||||
String::from_utf8_lossy(&output.stderr)
|
||||
.chars()
|
||||
.take(STDERR_LIMIT_CHARS)
|
||||
.collect(),
|
||||
]
|
||||
} else {
|
||||
parsed.errors
|
||||
},
|
||||
exit_code: output.status.code(),
|
||||
..parsed
|
||||
})
|
||||
}
|
||||
}
|
||||
69
litellm-rust/crates/testkit/src/target.rs
Normal file
69
litellm-rust/crates/testkit/src/target.rs
Normal file
|
|
@ -0,0 +1,69 @@
|
|||
use target_lexicon::{Architecture, Environment, OperatingSystem, Triple};
|
||||
|
||||
use crate::Error;
|
||||
|
||||
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
|
||||
pub enum Os {
|
||||
Macos,
|
||||
Linux,
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
|
||||
pub enum Arch {
|
||||
Aarch64,
|
||||
X86_64,
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
|
||||
pub struct Target {
|
||||
pub os: Os,
|
||||
pub arch: Arch,
|
||||
pub musl: bool,
|
||||
}
|
||||
|
||||
impl Target {
|
||||
pub fn host() -> Result<Self, Error> {
|
||||
Self::try_from(&Triple::host())
|
||||
}
|
||||
|
||||
pub(crate) const fn os_name(self) -> &'static str {
|
||||
match self.os {
|
||||
Os::Macos => "darwin",
|
||||
Os::Linux => "linux",
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) const fn arch_name(self) -> &'static str {
|
||||
match self.arch {
|
||||
Arch::Aarch64 => "arm64",
|
||||
Arch::X86_64 => "x64",
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) const fn musl_suffix(self) -> &'static str {
|
||||
if self.musl { "-musl" } else { "" }
|
||||
}
|
||||
}
|
||||
|
||||
impl TryFrom<&Triple> for Target {
|
||||
type Error = Error;
|
||||
|
||||
fn try_from(triple: &Triple) -> Result<Self, Error> {
|
||||
let unsupported = || Error::UnsupportedTarget(triple.to_string());
|
||||
let os = match triple.operating_system {
|
||||
OperatingSystem::Darwin(_) | OperatingSystem::MacOSX(_) => Os::Macos,
|
||||
OperatingSystem::Linux => Os::Linux,
|
||||
_ => return Err(unsupported()),
|
||||
};
|
||||
let arch = match triple.architecture {
|
||||
Architecture::Aarch64(_) => Arch::Aarch64,
|
||||
Architecture::X86_64 => Arch::X86_64,
|
||||
_ => return Err(unsupported()),
|
||||
};
|
||||
Ok(Self {
|
||||
os,
|
||||
arch,
|
||||
musl: triple.environment == Environment::Musl,
|
||||
})
|
||||
}
|
||||
}
|
||||
133
litellm-rust/crates/testkit/tests/configure.rs
Normal file
133
litellm-rust/crates/testkit/tests/configure.rs
Normal file
|
|
@ -0,0 +1,133 @@
|
|||
use std::path::Path;
|
||||
|
||||
use litellm_testkit::{ClaudeCode, Codex, Configure, Error, Opencode, Settings, Version, Wire};
|
||||
use rstest::rstest;
|
||||
|
||||
fn settings(wire: Wire) -> Settings {
|
||||
Settings {
|
||||
base_url: "http://localhost:4000/".to_owned(),
|
||||
api_key: "sk-test \"quoted\"".to_owned(),
|
||||
model: "some-model".to_owned(),
|
||||
wire,
|
||||
}
|
||||
}
|
||||
|
||||
fn version() -> Version {
|
||||
Version::new(1, 2, 3)
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case(&ClaudeCode, Wire::Messages)]
|
||||
#[case(&Codex, Wire::Responses)]
|
||||
#[case(&Opencode, Wire::ChatCompletions)]
|
||||
fn every_agent_runs_inside_the_given_home(#[case] agent: &impl Configure, #[case] wire: Wire) {
|
||||
let home = Path::new("/scratch/home");
|
||||
|
||||
let spec = agent.configure(&version(), &settings(wire), home).unwrap();
|
||||
|
||||
assert_eq!(spec.env["HOME"], "/scratch/home");
|
||||
assert!(
|
||||
spec.env
|
||||
.iter()
|
||||
.filter(|(key, _)| key.ends_with("_HOME") || key.as_str() == "CLAUDE_CONFIG_DIR")
|
||||
.all(|(_, value)| value.starts_with("/scratch/home"))
|
||||
);
|
||||
assert!(spec.files.keys().all(|path| path.is_relative()));
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::claude_code(&ClaudeCode, &[Wire::ChatCompletions, Wire::Responses])]
|
||||
#[case::codex(&Codex, &[Wire::ChatCompletions, Wire::Messages])]
|
||||
fn wires_an_agent_cannot_speak_are_refused(
|
||||
#[case] agent: &impl Configure,
|
||||
#[case] refused: &[Wire],
|
||||
) {
|
||||
refused.iter().for_each(|wire| {
|
||||
let result = agent.configure(&version(), &settings(*wire), Path::new("/h"));
|
||||
|
||||
assert!(matches!(result, Err(Error::UnsupportedWire { wire: got, .. }) if got == *wire));
|
||||
});
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn claude_code_points_at_the_gateway_root_with_the_key_and_model() {
|
||||
let spec = ClaudeCode
|
||||
.configure(&version(), &settings(Wire::Messages), Path::new("/h"))
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(spec.env["ANTHROPIC_BASE_URL"], "http://localhost:4000/");
|
||||
assert_eq!(spec.env["ANTHROPIC_AUTH_TOKEN"], "sk-test \"quoted\"");
|
||||
assert_eq!(spec.env["ANTHROPIC_MODEL"], "some-model");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn codex_config_is_valid_toml_routing_the_responses_api_to_the_gateway() {
|
||||
let dir = tempfile::tempdir().unwrap();
|
||||
let spec = Codex
|
||||
.configure(&version(), &settings(Wire::Responses), dir.path())
|
||||
.unwrap();
|
||||
spec.write_files(dir.path()).unwrap();
|
||||
|
||||
let config: toml::Table =
|
||||
toml::from_str(&std::fs::read_to_string(dir.path().join(".codex/config.toml")).unwrap())
|
||||
.unwrap();
|
||||
let provider = &config["model_providers"]["litellm"];
|
||||
|
||||
assert_eq!(config["model"].as_str(), Some("some-model"));
|
||||
assert_eq!(config["model_provider"].as_str(), Some("litellm"));
|
||||
assert_eq!(
|
||||
provider["base_url"].as_str(),
|
||||
Some("http://localhost:4000/v1")
|
||||
);
|
||||
assert_eq!(provider["wire_api"].as_str(), Some("responses"));
|
||||
let key_var = provider["env_key"].as_str().unwrap();
|
||||
assert_eq!(spec.env[key_var], "sk-test \"quoted\"");
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case(Wire::ChatCompletions)]
|
||||
#[case(Wire::Responses)]
|
||||
#[case(Wire::Messages)]
|
||||
fn opencode_config_is_valid_json_registering_the_gateway_model(#[case] wire: Wire) {
|
||||
let dir = tempfile::tempdir().unwrap();
|
||||
let spec = Opencode
|
||||
.configure(&version(), &settings(wire), dir.path())
|
||||
.unwrap();
|
||||
spec.write_files(dir.path()).unwrap();
|
||||
|
||||
let config: serde_json::Value = serde_json::from_str(
|
||||
&std::fs::read_to_string(dir.path().join(".config/opencode/opencode.json")).unwrap(),
|
||||
)
|
||||
.unwrap();
|
||||
let provider = &config["provider"]["litellm"];
|
||||
|
||||
assert_eq!(config["model"], "litellm/some-model");
|
||||
assert_eq!(provider["options"]["baseURL"], "http://localhost:4000/v1");
|
||||
assert_eq!(provider["options"]["apiKey"], "sk-test \"quoted\"");
|
||||
assert!(provider["models"]["some-model"].is_object());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn opencode_uses_a_different_provider_package_for_every_wire() {
|
||||
let package = |wire| {
|
||||
let dir = tempfile::tempdir().unwrap();
|
||||
let spec = Opencode
|
||||
.configure(&version(), &settings(wire), dir.path())
|
||||
.unwrap();
|
||||
let config: serde_json::Value =
|
||||
serde_json::from_str(spec.files.values().next().unwrap()).unwrap();
|
||||
config["provider"]["litellm"]["npm"]
|
||||
.as_str()
|
||||
.unwrap()
|
||||
.to_owned()
|
||||
};
|
||||
let packages = [Wire::ChatCompletions, Wire::Responses, Wire::Messages].map(package);
|
||||
|
||||
assert_eq!(
|
||||
packages
|
||||
.iter()
|
||||
.collect::<std::collections::BTreeSet<_>>()
|
||||
.len(),
|
||||
packages.len()
|
||||
);
|
||||
}
|
||||
262
litellm-rust/crates/testkit/tests/install.rs
Normal file
262
litellm-rust/crates/testkit/tests/install.rs
Normal file
|
|
@ -0,0 +1,262 @@
|
|||
mod support;
|
||||
|
||||
use std::str::FromStr;
|
||||
|
||||
use litellm_testkit::{ClaudeCode, Codex, Error, Installer, Opencode, Target, Version};
|
||||
use rstest::rstest;
|
||||
use serde_json::json;
|
||||
use support::{FakeFetch, script_printing, sha256, tar_gz, zip_archive};
|
||||
use target_lexicon::Triple;
|
||||
|
||||
fn target(triple: &str) -> Target {
|
||||
Target::try_from(&Triple::from_str(triple).unwrap()).unwrap()
|
||||
}
|
||||
|
||||
fn linux() -> Target {
|
||||
target("x86_64-unknown-linux-gnu")
|
||||
}
|
||||
fn version() -> Version {
|
||||
Version::new(9, 8, 7)
|
||||
}
|
||||
|
||||
fn github_release(asset: &str, download_url: &str, digest: Option<String>) -> Vec<u8> {
|
||||
json!({
|
||||
"assets": [
|
||||
{ "name": "unrelated.txt", "digest": "sha256:00", "browser_download_url": "https://example.test/unrelated" },
|
||||
{ "name": asset, "digest": digest, "browser_download_url": download_url },
|
||||
]
|
||||
})
|
||||
.to_string()
|
||||
.into_bytes()
|
||||
}
|
||||
|
||||
fn claude_routes(binary: &[u8], checksum: &str) -> Vec<(String, Vec<u8>)> {
|
||||
let base = "https://downloads.claude.ai/claude-code-releases/9.8.7";
|
||||
let manifest = json!({ "platforms": { "linux-x64": { "checksum": checksum } } });
|
||||
vec![
|
||||
(
|
||||
format!("{base}/manifest.json"),
|
||||
manifest.to_string().into_bytes(),
|
||||
),
|
||||
(format!("{base}/linux-x64/claude"), binary.to_vec()),
|
||||
]
|
||||
}
|
||||
|
||||
fn codex_routes(archive: Vec<u8>, digest: Option<String>) -> Vec<(String, Vec<u8>)> {
|
||||
let release = github_release(
|
||||
"codex-x86_64-unknown-linux-musl.tar.gz",
|
||||
"https://example.test/codex.tar.gz",
|
||||
digest,
|
||||
);
|
||||
vec![
|
||||
(
|
||||
"https://api.github.com/repos/openai/codex/releases/tags/rust-v9.8.7".to_owned(),
|
||||
release,
|
||||
),
|
||||
("https://example.test/codex.tar.gz".to_owned(), archive),
|
||||
]
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn claude_bare_binary_is_installed_and_runnable() {
|
||||
let binary = script_printing("9.8.7 (Claude Code)");
|
||||
let fetch = FakeFetch::new(claude_routes(&binary, &sha256(&binary)));
|
||||
let cache = tempfile::tempdir().unwrap();
|
||||
|
||||
let installed = Installer::new(&fetch, cache.path(), linux())
|
||||
.install(&ClaudeCode, &version())
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(installed.binary, cache.path().join("claude/9.8.7/claude"));
|
||||
assert_eq!(std::fs::read(&installed.binary).unwrap(), binary);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn codex_binary_is_extracted_from_the_tarball_under_its_own_name() {
|
||||
let binary = script_printing("codex-cli 9.8.7");
|
||||
let archive = tar_gz("codex-x86_64-unknown-linux-musl", &binary);
|
||||
let fetch = FakeFetch::new(codex_routes(
|
||||
archive.clone(),
|
||||
Some(format!("sha256:{}", sha256(&archive))),
|
||||
));
|
||||
let cache = tempfile::tempdir().unwrap();
|
||||
|
||||
let installed = Installer::new(&fetch, cache.path(), linux())
|
||||
.install(&Codex, &version())
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(std::fs::read(&installed.binary).unwrap(), binary);
|
||||
assert_eq!(installed.binary, cache.path().join("codex/9.8.7/codex"));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn opencode_binary_is_extracted_from_the_darwin_zip() {
|
||||
let binary = script_printing("9.8.7");
|
||||
let archive = zip_archive("opencode", &binary);
|
||||
let release = github_release(
|
||||
"opencode-darwin-arm64.zip",
|
||||
"https://example.test/opencode.zip",
|
||||
Some(format!("sha256:{}", sha256(&archive))),
|
||||
);
|
||||
let fetch = FakeFetch::new([
|
||||
(
|
||||
"https://api.github.com/repos/sst/opencode/releases/tags/v9.8.7".to_owned(),
|
||||
release,
|
||||
),
|
||||
("https://example.test/opencode.zip".to_owned(), archive),
|
||||
]);
|
||||
let cache = tempfile::tempdir().unwrap();
|
||||
|
||||
let installed = Installer::new(&fetch, cache.path(), target("aarch64-apple-darwin"))
|
||||
.install(&Opencode, &version())
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(std::fs::read(&installed.binary).unwrap(), binary);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn tampered_download_is_rejected_and_nothing_is_left_behind() {
|
||||
let binary = script_printing("9.8.7 (Claude Code)");
|
||||
let fetch = FakeFetch::new(claude_routes(&binary, &sha256(b"what the vendor signed")));
|
||||
let cache = tempfile::tempdir().unwrap();
|
||||
|
||||
let result = Installer::new(&fetch, cache.path(), linux())
|
||||
.install(&ClaudeCode, &version())
|
||||
.await;
|
||||
|
||||
assert!(matches!(result, Err(Error::ChecksumMismatch { .. })));
|
||||
assert!(!cache.path().join("claude/9.8.7").exists());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn github_asset_without_a_digest_is_refused() {
|
||||
let archive = tar_gz(
|
||||
"codex-x86_64-unknown-linux-musl",
|
||||
&script_printing("codex-cli 9.8.7"),
|
||||
);
|
||||
let fetch = FakeFetch::new(codex_routes(archive, None));
|
||||
let cache = tempfile::tempdir().unwrap();
|
||||
|
||||
let result = Installer::new(&fetch, cache.path(), linux())
|
||||
.install(&Codex, &version())
|
||||
.await;
|
||||
|
||||
assert!(matches!(result, Err(Error::MissingChecksum(_))));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn binary_reporting_a_different_version_is_removed() {
|
||||
let binary = script_printing("1.0.0 (Claude Code)");
|
||||
let fetch = FakeFetch::new(claude_routes(&binary, &sha256(&binary)));
|
||||
let cache = tempfile::tempdir().unwrap();
|
||||
|
||||
let result = Installer::new(&fetch, cache.path(), linux())
|
||||
.install(&ClaudeCode, &version())
|
||||
.await;
|
||||
|
||||
assert!(matches!(result, Err(Error::VersionMismatch { .. })));
|
||||
assert!(!cache.path().join("claude/9.8.7/claude").exists());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn second_install_reuses_the_cached_binary_without_downloading() {
|
||||
let binary = script_printing("9.8.7 (Claude Code)");
|
||||
let fetch = FakeFetch::new(claude_routes(&binary, &sha256(&binary)));
|
||||
let cache = tempfile::tempdir().unwrap();
|
||||
let installer = Installer::new(&fetch, cache.path(), linux());
|
||||
|
||||
let first = installer.install(&ClaudeCode, &version()).await.unwrap();
|
||||
let calls_after_first = fetch.calls();
|
||||
let second = installer.install(&ClaudeCode, &version()).await.unwrap();
|
||||
|
||||
assert_eq!(first, second);
|
||||
assert_eq!(fetch.calls(), calls_after_first);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn corrupted_cache_entry_is_replaced_by_a_fresh_download() {
|
||||
let binary = script_printing("9.8.7 (Claude Code)");
|
||||
let fetch = FakeFetch::new(claude_routes(&binary, &sha256(&binary)));
|
||||
let cache = tempfile::tempdir().unwrap();
|
||||
let installer = Installer::new(&fetch, cache.path(), linux());
|
||||
let installed = installer.install(&ClaudeCode, &version()).await.unwrap();
|
||||
std::fs::write(&installed.binary, script_printing("0.0.1")).unwrap();
|
||||
|
||||
installer.install(&ClaudeCode, &version()).await.unwrap();
|
||||
|
||||
assert_eq!(std::fs::read(&installed.binary).unwrap(), binary);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case("9.8.7-beta.1")]
|
||||
#[case("9.8.7+build.5")]
|
||||
#[tokio::test]
|
||||
async fn pre_releases_never_reach_the_network_or_the_filesystem(#[case] version: &str) {
|
||||
let fetch = FakeFetch::new([]);
|
||||
let cache = tempfile::tempdir().unwrap();
|
||||
|
||||
let result = Installer::new(&fetch, cache.path(), linux())
|
||||
.install(&ClaudeCode, &Version::parse(version).unwrap())
|
||||
.await;
|
||||
|
||||
assert!(matches!(result, Err(Error::InvalidVersion(_))));
|
||||
assert_eq!(fetch.calls(), 0);
|
||||
assert_eq!(std::fs::read_dir(cache.path()).unwrap().count(), 0);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn musl_linux_picks_the_musl_claude_build() {
|
||||
let binary = script_printing("9.8.7 (Claude Code)");
|
||||
let base = "https://downloads.claude.ai/claude-code-releases/9.8.7";
|
||||
let manifest = json!({ "platforms": {
|
||||
"linux-x64": { "checksum": sha256(b"glibc build") },
|
||||
"linux-x64-musl": { "checksum": sha256(&binary) },
|
||||
} });
|
||||
let fetch = FakeFetch::new([
|
||||
(
|
||||
format!("{base}/manifest.json"),
|
||||
manifest.to_string().into_bytes(),
|
||||
),
|
||||
(format!("{base}/linux-x64-musl/claude"), binary.clone()),
|
||||
]);
|
||||
let cache = tempfile::tempdir().unwrap();
|
||||
|
||||
let installed = Installer::new(&fetch, cache.path(), target("x86_64-unknown-linux-musl"))
|
||||
.install(&ClaudeCode, &version())
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(std::fs::read(&installed.binary).unwrap(), binary);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case("x86_64-pc-windows-msvc")]
|
||||
#[case("riscv64gc-unknown-linux-gnu")]
|
||||
#[case("wasm32-unknown-unknown")]
|
||||
fn targets_no_agent_ships_for_are_rejected(#[case] triple: &str) {
|
||||
let result = Target::try_from(&Triple::from_str(triple).unwrap());
|
||||
|
||||
assert!(matches!(result, Err(Error::UnsupportedTarget(_))));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn concurrent_installs_of_the_same_version_both_succeed() {
|
||||
let binary = script_printing("9.8.7 (Claude Code)");
|
||||
let fetch = FakeFetch::new(claude_routes(&binary, &sha256(&binary)));
|
||||
let cache = tempfile::tempdir().unwrap();
|
||||
let installer = Installer::new(&fetch, cache.path(), linux());
|
||||
|
||||
let wanted = version();
|
||||
let installs =
|
||||
futures_util::future::join_all((0..8).map(|_| installer.install(&ClaudeCode, &wanted)))
|
||||
.await;
|
||||
|
||||
assert!(installs.iter().all(Result::is_ok));
|
||||
assert_eq!(
|
||||
std::fs::read(&installs[0].as_ref().unwrap().binary).unwrap(),
|
||||
binary
|
||||
);
|
||||
}
|
||||
133
litellm-rust/crates/testkit/tests/live.rs
Normal file
133
litellm-rust/crates/testkit/tests/live.rs
Normal file
|
|
@ -0,0 +1,133 @@
|
|||
//! Drives the real agents through a real gateway. Run with `cargo test -p litellm-testkit --test live -- --ignored`
|
||||
//! after exporting `TESTKIT_GATEWAY_URL`, `TESTKIT_GATEWAY_KEY`, one `TESTKIT_MODEL_<WIRE>` per wire
|
||||
//! (`MESSAGES`, `RESPONSES`, `CHAT_COMPLETIONS`) and one `TESTKIT_<AGENT>_VERSION` per agent
|
||||
//! (`CLAUDE`, `CODEX`, `OPENCODE`). `TESTKIT_CACHE_DIR` and `GITHUB_TOKEN` are optional.
|
||||
|
||||
use std::path::PathBuf;
|
||||
use std::time::Duration;
|
||||
|
||||
use litellm_testkit::{
|
||||
Agent, ClaudeCode, Codex, HttpFetch, Installer, Opencode, Outcome, Prompt, Session, Settings,
|
||||
Target, Version, Wire,
|
||||
};
|
||||
use rstest::rstest;
|
||||
|
||||
const LIMIT: Duration = Duration::from_secs(180);
|
||||
|
||||
fn required(name: &str) -> String {
|
||||
std::env::var(name).unwrap_or_else(|_| panic!("{name} must be set to run the live tests"))
|
||||
}
|
||||
|
||||
fn model_var(wire: Wire) -> &'static str {
|
||||
match wire {
|
||||
Wire::Messages => "TESTKIT_MODEL_MESSAGES",
|
||||
Wire::Responses => "TESTKIT_MODEL_RESPONSES",
|
||||
Wire::ChatCompletions => "TESTKIT_MODEL_CHAT_COMPLETIONS",
|
||||
}
|
||||
}
|
||||
|
||||
async fn drive(
|
||||
agent: &impl Agent,
|
||||
version_var: &str,
|
||||
wire: Wire,
|
||||
model: Option<&str>,
|
||||
prompt: Prompt,
|
||||
) -> Outcome {
|
||||
let cache = std::env::var("TESTKIT_CACHE_DIR")
|
||||
.map(PathBuf::from)
|
||||
.unwrap_or_else(|_| std::env::temp_dir().join("litellm-testkit-cache"));
|
||||
let installer = Installer::new(HttpFetch::from_env(), cache, Target::host().unwrap());
|
||||
let installed = installer
|
||||
.install(agent, &Version::parse(&required(version_var)).unwrap())
|
||||
.await
|
||||
.unwrap();
|
||||
let settings = Settings {
|
||||
base_url: required("TESTKIT_GATEWAY_URL"),
|
||||
api_key: required("TESTKIT_GATEWAY_KEY"),
|
||||
model: model.map_or_else(|| required(model_var(wire)), str::to_owned),
|
||||
wire,
|
||||
};
|
||||
let home = tempfile::tempdir().unwrap();
|
||||
let session = Session::prepare(agent, &installed, settings, home.path()).unwrap();
|
||||
session.run(agent, &prompt, LIMIT).await.unwrap()
|
||||
}
|
||||
|
||||
fn text_prompt() -> Prompt {
|
||||
Prompt {
|
||||
text: "Reply with the single word: pong".to_owned(),
|
||||
allow_tools: false,
|
||||
}
|
||||
}
|
||||
|
||||
fn tool_prompt() -> Prompt {
|
||||
Prompt {
|
||||
text: "Run the shell command 'echo tool-ok' and reply with exactly its output.".to_owned(),
|
||||
allow_tools: true,
|
||||
}
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::claude_messages(&ClaudeCode, "TESTKIT_CLAUDE_VERSION", Wire::Messages)]
|
||||
#[case::codex_responses(&Codex, "TESTKIT_CODEX_VERSION", Wire::Responses)]
|
||||
#[case::opencode_chat(&Opencode, "TESTKIT_OPENCODE_VERSION", Wire::ChatCompletions)]
|
||||
#[case::opencode_responses(&Opencode, "TESTKIT_OPENCODE_VERSION", Wire::Responses)]
|
||||
#[case::opencode_messages(&Opencode, "TESTKIT_OPENCODE_VERSION", Wire::Messages)]
|
||||
#[ignore = "needs a live gateway, see the module docs"]
|
||||
#[tokio::test]
|
||||
async fn plain_prompt_gets_an_answer_and_token_usage(
|
||||
#[case] agent: &impl Agent,
|
||||
#[case] version_var: &str,
|
||||
#[case] wire: Wire,
|
||||
) {
|
||||
let outcome = drive(agent, version_var, wire, None, text_prompt()).await;
|
||||
|
||||
assert!(outcome.succeeded(), "{outcome:?}");
|
||||
assert!(outcome.text.to_lowercase().contains("pong"), "{outcome:?}");
|
||||
assert!(outcome.usage.output_tokens > 0, "{outcome:?}");
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::claude_messages(&ClaudeCode, "TESTKIT_CLAUDE_VERSION", Wire::Messages)]
|
||||
#[case::codex_responses(&Codex, "TESTKIT_CODEX_VERSION", Wire::Responses)]
|
||||
#[case::opencode_chat(&Opencode, "TESTKIT_OPENCODE_VERSION", Wire::ChatCompletions)]
|
||||
#[case::opencode_responses(&Opencode, "TESTKIT_OPENCODE_VERSION", Wire::Responses)]
|
||||
#[case::opencode_messages(&Opencode, "TESTKIT_OPENCODE_VERSION", Wire::Messages)]
|
||||
#[ignore = "needs a live gateway, see the module docs"]
|
||||
#[tokio::test]
|
||||
async fn tool_use_is_reported_and_its_result_reaches_the_answer(
|
||||
#[case] agent: &impl Agent,
|
||||
#[case] version_var: &str,
|
||||
#[case] wire: Wire,
|
||||
) {
|
||||
let outcome = drive(agent, version_var, wire, None, tool_prompt()).await;
|
||||
|
||||
assert!(outcome.succeeded(), "{outcome:?}");
|
||||
assert!(!outcome.tool_calls.is_empty(), "{outcome:?}");
|
||||
assert!(outcome.text.contains("tool-ok"), "{outcome:?}");
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::claude_messages(&ClaudeCode, "TESTKIT_CLAUDE_VERSION", Wire::Messages)]
|
||||
#[case::codex_responses(&Codex, "TESTKIT_CODEX_VERSION", Wire::Responses)]
|
||||
#[case::opencode_chat(&Opencode, "TESTKIT_OPENCODE_VERSION", Wire::ChatCompletions)]
|
||||
#[case::opencode_responses(&Opencode, "TESTKIT_OPENCODE_VERSION", Wire::Responses)]
|
||||
#[case::opencode_messages(&Opencode, "TESTKIT_OPENCODE_VERSION", Wire::Messages)]
|
||||
#[ignore = "needs a live gateway, see the module docs"]
|
||||
#[tokio::test]
|
||||
async fn model_the_gateway_rejects_is_reported_as_an_error(
|
||||
#[case] agent: &impl Agent,
|
||||
#[case] version_var: &str,
|
||||
#[case] wire: Wire,
|
||||
) {
|
||||
let outcome = drive(
|
||||
agent,
|
||||
version_var,
|
||||
wire,
|
||||
Some("testkit-no-such-model"),
|
||||
text_prompt(),
|
||||
)
|
||||
.await;
|
||||
|
||||
assert!(!outcome.succeeded(), "{outcome:?}");
|
||||
assert!(!outcome.errors.is_empty(), "{outcome:?}");
|
||||
}
|
||||
155
litellm-rust/crates/testkit/tests/session.rs
Normal file
155
litellm-rust/crates/testkit/tests/session.rs
Normal file
|
|
@ -0,0 +1,155 @@
|
|||
use std::os::unix::fs::PermissionsExt;
|
||||
use std::path::{Path, PathBuf};
|
||||
use std::time::Duration;
|
||||
|
||||
use litellm_testkit::{
|
||||
Configure, Drive, Error, Installed, LaunchSpec, Outcome, Prompt, Session, Settings, Version,
|
||||
Wire,
|
||||
};
|
||||
|
||||
struct Scripted;
|
||||
|
||||
impl Configure for Scripted {
|
||||
fn configure(
|
||||
&self,
|
||||
version: &Version,
|
||||
_settings: &Settings,
|
||||
home: &Path,
|
||||
) -> Result<LaunchSpec, Error> {
|
||||
Ok(LaunchSpec {
|
||||
env: [
|
||||
("AGENT_HOME".to_owned(), home.to_string_lossy().into_owned()),
|
||||
("AGENT_SAW_VERSION".to_owned(), version.to_string()),
|
||||
]
|
||||
.into(),
|
||||
files: [(
|
||||
PathBuf::from("conf/agent.toml"),
|
||||
"configured = true\n".to_owned(),
|
||||
)]
|
||||
.into(),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
impl Drive for Scripted {
|
||||
fn args(&self, _version: &Version, _settings: &Settings, prompt: &Prompt) -> Vec<String> {
|
||||
vec!["--prompt".to_owned(), prompt.text.clone()]
|
||||
}
|
||||
|
||||
fn parse(&self, _version: &Version, stdout: &str) -> Outcome {
|
||||
Outcome {
|
||||
text: stdout.to_owned(),
|
||||
..Outcome::default()
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn settings() -> Settings {
|
||||
Settings {
|
||||
base_url: "http://gateway.test".to_owned(),
|
||||
api_key: "sk-test".to_owned(),
|
||||
model: "some-model".to_owned(),
|
||||
wire: Wire::Messages,
|
||||
}
|
||||
}
|
||||
|
||||
fn prompt(text: &str) -> Prompt {
|
||||
Prompt {
|
||||
text: text.to_owned(),
|
||||
allow_tools: false,
|
||||
}
|
||||
}
|
||||
|
||||
fn session(script: &str) -> (Session, tempfile::TempDir) {
|
||||
let dir = tempfile::tempdir().unwrap();
|
||||
let binary = dir.path().join("agent");
|
||||
std::fs::write(&binary, format!("#!/bin/sh\n{script}\n")).unwrap();
|
||||
std::fs::set_permissions(&binary, std::fs::Permissions::from_mode(0o755)).unwrap();
|
||||
let home = dir.path().join("home");
|
||||
std::fs::create_dir(&home).unwrap();
|
||||
let installed = Installed {
|
||||
version: Version::new(4, 5, 6),
|
||||
binary,
|
||||
};
|
||||
(
|
||||
Session::prepare(&Scripted, &installed, settings(), home).unwrap(),
|
||||
dir,
|
||||
)
|
||||
}
|
||||
|
||||
const LIMIT: Duration = Duration::from_secs(20);
|
||||
|
||||
#[tokio::test]
|
||||
async fn prepare_writes_the_config_files_under_home() {
|
||||
let (_session, dir) = session("true");
|
||||
|
||||
let written = std::fs::read_to_string(dir.path().join("home/conf/agent.toml")).unwrap();
|
||||
|
||||
assert_eq!(written, "configured = true\n");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn configure_and_drive_are_given_the_installed_version() {
|
||||
let (session, _dir) = session("echo \"$AGENT_SAW_VERSION\"");
|
||||
|
||||
let outcome = session.run(&Scripted, &prompt("hi"), LIMIT).await.unwrap();
|
||||
|
||||
assert_eq!(outcome.text.trim(), "4.5.6");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn agent_runs_in_home_with_only_its_own_environment() {
|
||||
let (session, dir) = session("pwd -P; env");
|
||||
|
||||
let outcome = session.run(&Scripted, &prompt("hi"), LIMIT).await.unwrap();
|
||||
|
||||
let home = dir.path().join("home").canonicalize().unwrap();
|
||||
assert_eq!(outcome.text.lines().next().unwrap(), home.to_string_lossy());
|
||||
assert!(outcome.text.contains("AGENT_HOME="));
|
||||
assert!(
|
||||
!outcome.text.contains("CARGO_"),
|
||||
"test runner environment leaked into the agent"
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn prompt_reaches_the_agent_as_one_untouched_argument() {
|
||||
let (session, _dir) = session("printf '%s|' \"$@\"");
|
||||
let text = "two spaces; $(echo injected) 'quoted'";
|
||||
|
||||
let outcome = session.run(&Scripted, &prompt(text), LIMIT).await.unwrap();
|
||||
|
||||
assert_eq!(outcome.text, format!("--prompt|{text}|"));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn clean_exit_is_a_success() {
|
||||
let (session, _dir) = session("echo done");
|
||||
|
||||
let outcome = session.run(&Scripted, &prompt("hi"), LIMIT).await.unwrap();
|
||||
|
||||
assert_eq!(outcome.exit_code, Some(0));
|
||||
assert!(outcome.succeeded());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn failing_exit_without_a_parsed_error_reports_stderr() {
|
||||
let (session, _dir) = session("echo boom >&2; exit 3");
|
||||
|
||||
let outcome = session.run(&Scripted, &prompt("hi"), LIMIT).await.unwrap();
|
||||
|
||||
assert_eq!(outcome.exit_code, Some(3));
|
||||
assert!(!outcome.succeeded());
|
||||
assert_eq!(outcome.errors, ["boom\n"]);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn agent_that_outlives_the_limit_is_stopped() {
|
||||
let (session, _dir) = session("sleep 30");
|
||||
|
||||
let result = session
|
||||
.run(&Scripted, &prompt("hi"), Duration::from_millis(200))
|
||||
.await;
|
||||
|
||||
assert!(matches!(result, Err(Error::Timeout(_))));
|
||||
}
|
||||
70
litellm-rust/crates/testkit/tests/support/mod.rs
Normal file
70
litellm-rust/crates/testkit/tests/support/mod.rs
Normal file
|
|
@ -0,0 +1,70 @@
|
|||
use std::collections::HashMap;
|
||||
use std::io::Write;
|
||||
use std::sync::atomic::{AtomicUsize, Ordering};
|
||||
|
||||
use litellm_testkit::{Error, Fetch};
|
||||
use sha2::{Digest, Sha256};
|
||||
|
||||
pub struct FakeFetch {
|
||||
routes: HashMap<String, Vec<u8>>,
|
||||
calls: AtomicUsize,
|
||||
}
|
||||
|
||||
impl FakeFetch {
|
||||
pub fn new(routes: impl IntoIterator<Item = (String, Vec<u8>)>) -> Self {
|
||||
Self {
|
||||
routes: routes.into_iter().collect(),
|
||||
calls: AtomicUsize::new(0),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn calls(&self) -> usize {
|
||||
self.calls.load(Ordering::SeqCst)
|
||||
}
|
||||
}
|
||||
|
||||
impl Fetch for FakeFetch {
|
||||
async fn get(&self, url: &str) -> Result<Vec<u8>, Error> {
|
||||
self.calls.fetch_add(1, Ordering::SeqCst);
|
||||
self.routes.get(url).cloned().ok_or_else(|| Error::Status {
|
||||
url: url.to_owned(),
|
||||
status: 404,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
impl Fetch for &FakeFetch {
|
||||
async fn get(&self, url: &str) -> Result<Vec<u8>, Error> {
|
||||
(*self).get(url).await
|
||||
}
|
||||
}
|
||||
|
||||
pub fn sha256(bytes: &[u8]) -> String {
|
||||
format!("{:x}", Sha256::digest(bytes))
|
||||
}
|
||||
|
||||
pub fn script_printing(output: &str) -> Vec<u8> {
|
||||
format!("#!/bin/sh\necho '{output}'\n").into_bytes()
|
||||
}
|
||||
|
||||
pub fn tar_gz(member: &str, contents: &[u8]) -> Vec<u8> {
|
||||
let mut builder = tar::Builder::new(Vec::new());
|
||||
let mut header = tar::Header::new_gnu();
|
||||
header.set_size(contents.len() as u64);
|
||||
header.set_mode(0o755);
|
||||
header.set_cksum();
|
||||
builder.append_data(&mut header, member, contents).unwrap();
|
||||
let tarball = builder.into_inner().unwrap();
|
||||
let mut encoder = flate2::write::GzEncoder::new(Vec::new(), flate2::Compression::default());
|
||||
encoder.write_all(&tarball).unwrap();
|
||||
encoder.finish().unwrap()
|
||||
}
|
||||
|
||||
pub fn zip_archive(member: &str, contents: &[u8]) -> Vec<u8> {
|
||||
let mut writer = zip::ZipWriter::new(std::io::Cursor::new(Vec::new()));
|
||||
writer
|
||||
.start_file(member, zip::write::SimpleFileOptions::default())
|
||||
.unwrap();
|
||||
writer.write_all(contents).unwrap();
|
||||
writer.finish().unwrap().into_inner()
|
||||
}
|
||||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue