diff --git a/litellm-rust/Cargo.lock b/litellm-rust/Cargo.lock index 768f1c8cc7b..710d9c5399c 100644 --- a/litellm-rust/Cargo.lock +++ b/litellm-rust/Cargo.lock @@ -3836,6 +3836,7 @@ dependencies = [ "litellm-host-http", "litellm-http", "litellm-inference", + "litellm-inference-responses", "litellm-inference-transcription", "litellm-llms", "litellm-llms-types", @@ -4040,6 +4041,33 @@ dependencies = [ "wiremock", ] +[[package]] +name = "litellm-inference-responses" +version = "0.1.0" +dependencies = [ + "bytes", + "futures-util", + "litellm-auth", + "litellm-cache", + "litellm-cache-memory", + "litellm-cache-response", + "litellm-host", + "litellm-host-native", + "litellm-http", + "litellm-inference", + "litellm-llms", + "litellm-llms-types", + "litellm-secrets", + "litellm-tracing", + "reqwest 0.12.28", + "rstest", + "serde_json", + "tokio", + "tokio-tungstenite", + "tracing", + "wiremock", +] + [[package]] name = "litellm-inference-transcription" version = "0.1.0" @@ -4142,6 +4170,7 @@ dependencies = [ "litellm-host-python", "litellm-http", "litellm-inference", + "litellm-inference-responses", "litellm-inference-transcription", "litellm-llms", "litellm-llms-types", diff --git a/litellm-rust/Cargo.toml b/litellm-rust/Cargo.toml index 79c4b520c96..6fc94aedf27 100644 --- a/litellm-rust/Cargo.toml +++ b/litellm-rust/Cargo.toml @@ -18,6 +18,7 @@ litellm-traces-clickhouse = { path = "crates/traces-clickhouse" } litellm-storage-clickhouse = { path = "crates/storage-clickhouse" } litellm-inference = { path = "crates/inference" } litellm-inference-transcription = { path = "crates/inference-transcription" } +litellm-inference-responses = { path = "crates/inference-responses" } litellm-gateway-mcp = { path = "crates/gateway-mcp" } litellm-gateway = { path = "crates/gateway" } litellm-gateway-inference = { path = "crates/gateway-inference" } diff --git a/litellm-rust/crates/gateway-inference/Cargo.toml b/litellm-rust/crates/gateway-inference/Cargo.toml index 235dd3d4b72..dd5f179793f 100644 --- a/litellm-rust/crates/gateway-inference/Cargo.toml +++ b/litellm-rust/crates/gateway-inference/Cargo.toml @@ -14,6 +14,7 @@ litellm-auth.workspace = true litellm-gateway-auth.workspace = true litellm-inference.workspace = true litellm-inference-transcription.workspace = true +litellm-inference-responses.workspace = true litellm-host-http.workspace = true litellm-host.workspace = true litellm-http.workspace = true diff --git a/litellm-rust/crates/gateway-inference/src/lib.rs b/litellm-rust/crates/gateway-inference/src/lib.rs index 48c6666d8a1..059cad93e1e 100644 --- a/litellm-rust/crates/gateway-inference/src/lib.rs +++ b/litellm-rust/crates/gateway-inference/src/lib.rs @@ -18,8 +18,9 @@ use axum::{Router, routing::post}; use litellm_http::{ClientVariant, HttpClientConfig, media::UrlPolicy}; use litellm_inference::{ chat_completions::ChatCompletionsRoute, messages::MessagesRoute, ocr::OcrRoute, - resources::CoreResources, responses::ResponsesRoute, + resources::CoreResources, }; +use litellm_inference_responses::ResponsesRoute; use litellm_inference_transcription::AudioTranscriptionRoute; use litellm_llms::base_llm::ocr::{handler::OcrClient, settings::OcrSettings}; use litellm_secrets::source::SecretSource; diff --git a/litellm-rust/crates/gateway-inference/src/responses.rs b/litellm-rust/crates/gateway-inference/src/responses.rs index 67e409790fd..77bebf610e1 100644 --- a/litellm-rust/crates/gateway-inference/src/responses.rs +++ b/litellm-rust/crates/gateway-inference/src/responses.rs @@ -3,7 +3,7 @@ use std::sync::Arc; use axum::{Json, body::Bytes, extract::State, response::Response}; use litellm_gateway_auth::AuthenticatedRequest; use litellm_host_http::Sse; -use litellm_inference::responses::{route::Responses, types::ResponsesCall}; +use litellm_inference_responses::{route::Responses, types::ResponsesCall}; use serde_json::json; use crate::{Error, Gateway, JsonObject, request}; diff --git a/litellm-rust/crates/inference-responses/Cargo.toml b/litellm-rust/crates/inference-responses/Cargo.toml new file mode 100644 index 00000000000..5aaa4a37b4c --- /dev/null +++ b/litellm-rust/crates/inference-responses/Cargo.toml @@ -0,0 +1,36 @@ +[package] +name = "litellm-inference-responses" +version = "0.1.0" +edition.workspace = true +license.workspace = true +repository.workspace = true + +[dependencies] +bytes.workspace = true +futures-util.workspace = true +litellm-auth = { workspace = true, features = ["aws", "azure", "gcp"] } +litellm-cache-response.workspace = true +litellm-host.workspace = true +litellm-http.workspace = true +litellm-inference.workspace = true +litellm-llms.workspace = true +litellm-llms-types.workspace = true +litellm-secrets.workspace = true +reqwest.workspace = true +serde_json = { workspace = true, features = ["preserve_order"] } +tokio = { workspace = true, features = ["sync"] } +tokio-tungstenite.workspace = true +tracing.workspace = true + +[dev-dependencies] +futures-util.workspace = true +litellm-cache.workspace = true +litellm-cache-memory.workspace = true +litellm-cache-response.workspace = true +litellm-host-native.workspace = true +litellm-http.workspace = true +litellm-inference = { workspace = true, features = ["test-support"] } +litellm-tracing.workspace = true +rstest.workspace = true +tokio.workspace = true +wiremock.workspace = true diff --git a/litellm-rust/crates/inference/src/responses/handler.rs b/litellm-rust/crates/inference-responses/src/handler.rs similarity index 91% rename from litellm-rust/crates/inference/src/responses/handler.rs rename to litellm-rust/crates/inference-responses/src/handler.rs index a4b89c19d8a..eadedca7d3f 100644 --- a/litellm-rust/crates/inference/src/responses/handler.rs +++ b/litellm-rust/crates/inference-responses/src/handler.rs @@ -36,9 +36,11 @@ pub(super) async fn execute( ) .await?; let cache = cache.filter(|_| authenticated.signer.is_none()); - let cache_request = - crate::caching::CacheRequest::from_wire(identity, cache.as_ref().map(|_| &wire)); - crate::caching::execute_streaming::( + let cache_request = litellm_inference::caching::CacheRequest::from_wire( + identity, + cache.as_ref().map(|_| &wire), + ); + litellm_inference::caching::execute_streaming::( cache_request, cache.as_ref().map(|cache| cache.service.clone()), cache.as_ref().map(|cache| cache.options(cache_options)), @@ -50,7 +52,7 @@ pub(super) async fn execute( Some(serde_json::Value::Bool(value)) => *value, Some(_) => return Err(Error::InvalidRequest("stream must be a boolean".into())), }; - let outbound = crate::outbound::outbound_request( + let outbound = litellm_inference::outbound::outbound_request( Authenticated { headers: wire.headers, signer: authenticated.signer, @@ -59,7 +61,7 @@ pub(super) async fn execute( &wire.body, Some(request.timeout.unwrap_or(Duration::from_secs(600))), )?; - let response = crate::outbound::send(outbound, http) + let response = litellm_inference::outbound::send(outbound, http) .await .map_err(network)?; let status = response.status().as_u16(); diff --git a/litellm-rust/crates/inference/src/responses/mod.rs b/litellm-rust/crates/inference-responses/src/lib.rs similarity index 89% rename from litellm-rust/crates/inference/src/responses/mod.rs rename to litellm-rust/crates/inference-responses/src/lib.rs index fb48050184f..431982457f6 100644 --- a/litellm-rust/crates/inference/src/responses/mod.rs +++ b/litellm-rust/crates/inference-responses/src/lib.rs @@ -1,11 +1,12 @@ -pub use crate::error::RouteError as Error; use litellm_host::observation::ObservationSender; +pub use litellm_inference::RouteError as Error; + +pub mod route; +pub mod types; pub mod websocket; mod handler; mod prepare; -pub mod route; -pub mod types; use std::sync::Arc; @@ -47,9 +48,9 @@ impl ResponsesRoute { &self, call: ResponsesCall, interceptors: &impl Interceptors, - options: impl Into, + options: impl Into, ) -> Result { - let crate::CallOptions { + let litellm_inference::CallOptions { cache: cache_options, observers, } = options.into(); @@ -75,7 +76,7 @@ impl ResponsesRoute { interceptors: &impl litellm_host::interceptors::Interceptors, observers: Option<&ObservationSender>, ) -> Result { - crate::diagnostic::call(async { + litellm_inference::diagnostic::call(async { self.run_provider(call, cache_options, interceptors, observers) .await }) @@ -90,7 +91,10 @@ impl ResponsesRoute { observers: Option<&ObservationSender>, ) -> Result { let request = prepare::prepare(call, self.secrets.as_ref()).await?; - crate::diagnostic::provider(&request.context.model, &request.context.custom_llm_provider); + litellm_inference::diagnostic::provider( + &request.context.model, + &request.context.custom_llm_provider, + ); let execute: futures_util::future::BoxFuture<'_, Result> = Box::pin(handler::execute( &self.http, diff --git a/litellm-rust/crates/inference/src/responses/prepare.rs b/litellm-rust/crates/inference-responses/src/prepare.rs similarity index 100% rename from litellm-rust/crates/inference/src/responses/prepare.rs rename to litellm-rust/crates/inference-responses/src/prepare.rs diff --git a/litellm-rust/crates/inference/src/responses/route.rs b/litellm-rust/crates/inference-responses/src/route.rs similarity index 88% rename from litellm-rust/crates/inference/src/responses/route.rs rename to litellm-rust/crates/inference-responses/src/route.rs index 2cdd143ee2b..c62a13244ce 100644 --- a/litellm-rust/crates/inference/src/responses/route.rs +++ b/litellm-rust/crates/inference-responses/src/route.rs @@ -27,9 +27,9 @@ impl ResponsesRoute { pub fn machine( self, call: ResponsesCall, - options: impl Into, + options: impl Into, ) -> HostedMachine { - let crate::CallOptions { + let litellm_inference::CallOptions { cache: cache_options, observers, } = options.into(); @@ -44,7 +44,7 @@ impl ResponsesRoute { } } -impl crate::caching::Cachable for Responses { +impl litellm_inference::caching::Cachable for Responses { const SURFACE: &'static str = "responses"; fn reusable(response: &Self::Response) -> bool { @@ -56,7 +56,7 @@ impl crate::caching::Cachable for Responses { } } -impl crate::caching::StreamCachable for Responses { +impl litellm_inference::caching::StreamCachable for Responses { const TERMINAL_EVENT: &'static str = "response.completed"; fn replay(data: bytes::Bytes) -> Option> { diff --git a/litellm-rust/crates/inference/src/responses/types.rs b/litellm-rust/crates/inference-responses/src/types.rs similarity index 100% rename from litellm-rust/crates/inference/src/responses/types.rs rename to litellm-rust/crates/inference-responses/src/types.rs diff --git a/litellm-rust/crates/inference/src/responses/websocket.rs b/litellm-rust/crates/inference-responses/src/websocket.rs similarity index 93% rename from litellm-rust/crates/inference/src/responses/websocket.rs rename to litellm-rust/crates/inference-responses/src/websocket.rs index 4ceff787a66..5d088bd2d63 100644 --- a/litellm-rust/crates/inference/src/responses/websocket.rs +++ b/litellm-rust/crates/inference-responses/src/websocket.rs @@ -40,7 +40,7 @@ impl ResponsesWebSocketConnection { headers: &HashMap, timeout: Option, ) -> Result { - crate::diagnostic::operation("litellm.websocket.connect_url", async { + litellm_inference::diagnostic::operation("litellm.websocket.connect_url", async { let mut request = url.into_client_request().map_err(|error| { Error::Transport(litellm_http::transport::Error::Network(error.to_string())) })?; @@ -86,7 +86,7 @@ impl ResponsesWebSocketConnection { fields(outcome) )] pub async fn send_text(&self, text: String) -> Result<(), Error> { - crate::diagnostic::operation("litellm.websocket.send_text", async { + litellm_inference::diagnostic::operation("litellm.websocket.send_text", async { let mut socket = self.socket.lock().await; let Some(socket) = socket.as_mut() else { return Err(Error::Transport(litellm_http::transport::Error::Network( @@ -107,7 +107,7 @@ impl ResponsesWebSocketConnection { fields(outcome) )] pub async fn recv_text(&self) -> Result, Error> { - crate::diagnostic::operation("litellm.websocket.recv_text", async { + litellm_inference::diagnostic::operation("litellm.websocket.recv_text", async { let mut socket = self.socket.lock().await; let Some(socket) = socket.as_mut() else { return Ok(None); @@ -134,7 +134,7 @@ impl ResponsesWebSocketConnection { fields(outcome) )] pub async fn close(&self) -> Result<(), Error> { - crate::diagnostic::operation("litellm.websocket.close", async { + litellm_inference::diagnostic::operation("litellm.websocket.close", async { let mut socket = self.socket.lock().await; if let Some(socket) = socket.as_mut() { socket.close(None).await.map_err(|error| { diff --git a/litellm-rust/crates/inference-responses/tests/caching.rs b/litellm-rust/crates/inference-responses/tests/caching.rs new file mode 100644 index 00000000000..f45cefdc13b --- /dev/null +++ b/litellm-rust/crates/inference-responses/tests/caching.rs @@ -0,0 +1,327 @@ +mod support; + +use std::{ + sync::{ + Arc, + atomic::{AtomicUsize, Ordering}, + }, + time::Duration, +}; + +use litellm_cache_memory::InMemoryCache; +use litellm_cache_response::{ + CacheOptions, CacheScope, ResponseCache, ResponseCacheConfig, ResponseCacheService, + ResponseEnvelope, ScopedCache, +}; +use litellm_host::interceptors::{ + ExecutionFacts, Interceptors, ProviderIdentity, RawResponse, RequestContext, ResultSource, + WireRequest, +}; +use litellm_inference::{ + RouteError, + caching::{CacheRequest, execute_unary}, +}; +use rstest::{fixture, rstest}; +use serde_json::{Value, json}; + +fn cache_request(input: Value) -> CacheRequest { + CacheRequest { + identity: ProviderIdentity { + model: "test-model".into(), + provider: "test-provider".into(), + }, + input, + } +} + +#[fixture] +fn cache() -> Arc { + cache_with_limit(4096) +} + +fn cache_with_limit(max_entry_bytes: usize) -> Arc { + Arc::new( + ResponseCache::new(Arc::new(InMemoryCache::new( + Some(100), + Some(Duration::from_secs(60)), + ))) + .with_config(ResponseCacheConfig { + namespace: "test".into(), + max_entry_bytes, + }), + ) +} + +struct InvalidEntryCache( + ResponseCache>, + Value, +); + +impl ResponseCacheService for InvalidEntryCache { + fn config(&self) -> &ResponseCacheConfig { + self.0.config() + } + + fn lookup<'a>( + &'a self, + request: &'a litellm_cache_response::ResponseCacheRequest, + now: Duration, + ) -> futures_util::future::BoxFuture<'a, Result, litellm_cache::Error>> { + Box::pin(async move { + Ok(self + .0 + .async_lookup(request, now) + .await? + .or_else(|| Some(self.1.clone()))) + }) + } + + fn store<'a>( + &'a self, + request: &'a litellm_cache_response::ResponseCacheRequest, + response: Value, + now: Duration, + ) -> futures_util::future::BoxFuture<'a, Result<(), litellm_cache::Error>> { + Box::pin(self.0.async_store(request, response, now)) + } +} + +struct ChangingSecrets { + revision: AtomicUsize, + endpoints: [String; 2], + change_credentials: bool, +} + +impl litellm_secrets::source::SecretSource for ChangingSecrets { + fn get_secret_str<'a>( + &'a self, + name: &'a str, + ) -> futures_util::future::BoxFuture< + 'a, + Result, litellm_secrets::Error>, + > { + Box::pin(async move { + let revision = self.revision.load(Ordering::SeqCst); + let value = if name.ends_with("_API_KEY") { + Some(format!( + "key-{}", + if self.change_credentials { revision } else { 0 } + )) + } else if name.ends_with("_API_BASE") { + Some(self.endpoints[revision].clone()) + } else { + None + }; + Ok(value.map(litellm_secrets::SecretValue::new)) + }) + } +} + +#[derive(Default)] +struct ChangingHooks { + calls: AtomicUsize, + rewrite: bool, + facts: std::sync::Mutex>, +} + +impl Interceptors for ChangingHooks { + async fn before_provider_request( + &self, + mut wire: WireRequest, + _: RequestContext, + ) -> Result { + let call = self.calls.fetch_add(1, Ordering::SeqCst); + if self.rewrite { + wire.body["temperature"] = json!(if call < 2 { 0.1 } else { 0.8 }); + } + Ok(wire) + } + + async fn after_provider_response(&self, _: RawResponse) -> Result<(), RouteError> { + Ok(()) + } + + async fn result_ready(&self, facts: ExecutionFacts) -> Result<(), RouteError> { + self.facts.lock().unwrap().push(facts); + Ok(()) + } +} + +#[rstest] +#[case::chat_completion(json!({"kind":"Response","value":{"id":"chat-1","model":"test","choices":[],"usage":{"prompt_tokens":3,"completion_tokens":2}}}))] +#[case::wrong_envelope(json!({"kind":"Stream","value":"data: [DONE]\n\n"}))] +#[tokio::test] +async fn responses_refetches_instead_of_deserializing_another_api_response( + #[case] poisoned: Value, +) { + use litellm_inference_responses::route::Responses; + use litellm_llms_types::formats::responses::ResponsesApiResponse; + + let cache: Arc = Arc::new(InvalidEntryCache( + ResponseCache::new(Arc::new(InMemoryCache::default())), + serde_json::to_value(ResponseEnvelope::new("responses", poisoned)).unwrap(), + )); + let calls = AtomicUsize::new(0); + for _ in 0..2 { + let response = execute_unary::( + cache_request(json!({"input":"hello"})), + Some(cache.clone()), + Some(CacheOptions::new(CacheScope::Shared)), + &(), + None, + || async { + calls.fetch_add(1, Ordering::SeqCst); + Ok(ResponsesApiResponse { + id: "fresh-response".into(), + model: "test".into(), + output: vec![ + json!({"type":"message","content":[{"type":"output_text","text":"fresh"}]}), + ], + extra: [("status".into(), json!("completed"))] + .into_iter() + .collect(), + }) + }, + ) + .await + .unwrap(); + assert_eq!(response.id, "fresh-response"); + assert_eq!(response.output[0]["content"][0]["text"], "fresh"); + } + assert_eq!(calls.load(Ordering::SeqCst), 1); +} + +#[rstest] +#[case::completed("completed", 1)] +#[case::incomplete("incomplete", 2)] +#[tokio::test] +async fn responses_cache_only_reuses_completed_responses( + cache: Arc, + #[case] status: &str, + #[case] expected_calls: usize, +) { + use litellm_inference_responses::route::Responses; + use litellm_llms_types::formats::responses::ResponsesApiResponse; + + let calls = AtomicUsize::new(0); + for _ in 0..2 { + let response = execute_unary::( + cache_request(json!({"input":"hello"})), + Some(cache.clone()), + Some(CacheOptions::new(CacheScope::Shared)), + &(), + None, + || async { + let call = calls.fetch_add(1, Ordering::SeqCst); + Ok(ResponsesApiResponse { + id: call.to_string(), + model: "test".into(), + output: Vec::new(), + extra: [("status".into(), json!(status))].into_iter().collect(), + }) + }, + ) + .await + .unwrap(); + assert_eq!(response.extra.get("status"), Some(&json!(status))); + } + assert_eq!(calls.load(Ordering::SeqCst), expected_calls); +} + +#[rstest] +#[case::credentials("credentials")] +#[case::endpoint("endpoint")] +#[case::callback("callback")] +#[tokio::test] +async fn responses_cache_identity_follows_resolved_configuration_and_request_callbacks( + cache: Arc, + #[case] change: &str, +) { + use litellm_inference_responses::types::ResponsesCall; + use wiremock::{Mock, MockServer, ResponseTemplate, matchers::method}; + + let first = MockServer::start().await; + let second = MockServer::start().await; + let response = json!({"id":"response-test", "model":"test", "output":[], "status":"completed"}); + Mock::given(method("POST")) + .respond_with(ResponseTemplate::new(200).set_body_json(response.clone())) + .expect(if change == "endpoint" { 1 } else { 2 }) + .mount(&first) + .await; + Mock::given(method("POST")) + .respond_with(ResponseTemplate::new(200).set_body_json(response)) + .expect(if change == "endpoint" { 1 } else { 0 }) + .mount(&second) + .await; + let secrets = Arc::new(ChangingSecrets { + revision: AtomicUsize::new(0), + endpoints: [ + first.uri(), + if change == "endpoint" { + second.uri() + } else { + first.uri() + }, + ], + change_credentials: change == "credentials", + }); + let hooks = ChangingHooks { + rewrite: change == "callback", + ..Default::default() + }; + for call in 0..4 { + secrets + .revision + .store(usize::from(call >= 2), Ordering::SeqCst); + let cache = ScopedCache::new(cache.clone(), CacheScope::Shared); + support::responses_route(secrets.clone()) + .with_cache(cache) + .execute( + ResponsesCall { + model: "test".into(), + input: json!("hello"), + optional_params: Default::default(), + api_key: None, + api_base: None, + custom_llm_provider: None, + extra_headers: None, + timeout: None, + }, + &hooks, + None, + ) + .await + .unwrap(); + } + assert_eq!(hooks.calls.load(Ordering::SeqCst), 4); + { + let facts = hooks.facts.lock().unwrap(); + assert_eq!(facts[0].source, ResultSource::Provider); + assert_eq!(facts[2].source, ResultSource::Provider); + let (ResultSource::Cache { key: first_key }, ResultSource::Cache { key: second_key }) = + (&facts[1].source, &facts[3].source) + else { + panic!("unchanged effective requests must hit the cache"); + }; + assert_ne!(first_key, second_key); + } + let requests = first.received_requests().await.unwrap(); + if change == "credentials" { + assert_ne!( + requests[0].headers["authorization"], + requests[1].headers["authorization"] + ); + } + if change == "callback" { + assert_eq!( + serde_json::from_slice::(&requests[0].body).unwrap()["temperature"], + 0.1 + ); + assert_eq!( + serde_json::from_slice::(&requests[1].body).unwrap()["temperature"], + 0.8 + ); + } + first.verify().await; + second.verify().await; +} diff --git a/litellm-rust/crates/inference/tests/responses.rs b/litellm-rust/crates/inference-responses/tests/responses.rs similarity index 99% rename from litellm-rust/crates/inference/tests/responses.rs rename to litellm-rust/crates/inference-responses/tests/responses.rs index 0a15be38479..969b650106a 100644 --- a/litellm-rust/crates/inference/tests/responses.rs +++ b/litellm-rust/crates/inference-responses/tests/responses.rs @@ -6,11 +6,11 @@ use std::sync::Arc; use futures_util::TryStreamExt; use litellm_host::{call::HostedCompletion, lifecycle::CallEvent}; -use litellm_inference::responses::{ +use litellm_inference::test_support::{RecordingSecrets, no_secrets}; +use litellm_inference_responses::{ route::Responses, types::{ResponsesCall, ResponsesOutput}, }; -use litellm_inference::test_support::{RecordingSecrets, no_secrets}; use rstest::{fixture, rstest}; use serde_json::json; use wiremock::ResponseTemplate; @@ -418,7 +418,7 @@ async fn websocket_operations_trace_outcomes_without_capturing_frames_or_credent traces: TraceCapture, ) { use futures_util::{SinkExt, StreamExt}; - use litellm_inference::responses::websocket::ResponsesWebSocketConnection; + use litellm_inference_responses::websocket::ResponsesWebSocketConnection; let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); let address = listener.local_addr().unwrap(); diff --git a/litellm-rust/crates/inference-responses/tests/support/mod.rs b/litellm-rust/crates/inference-responses/tests/support/mod.rs new file mode 100644 index 00000000000..5e8de46e12d --- /dev/null +++ b/litellm-rust/crates/inference-responses/tests/support/mod.rs @@ -0,0 +1,279 @@ +//! Shared fixtures for route integration tests: a scripted upstream and a recording +//! secret source. + +#![allow(dead_code)] // each test binary compiles this module on its own and uses a different subset + +use std::{ + ops::ControlFlow, + sync::{Arc, Mutex}, +}; + +use litellm_inference::test_support::{http_config, provider_http, resources}; +use litellm_secrets::source::SecretSource; +use serde_json::Value; +use wiremock::{Mock, MockServer, Request, ResponseTemplate, matchers::any}; + +pub fn responses_route( + secrets: Arc, +) -> litellm_inference_responses::ResponsesRoute { + let resources = resources(); + litellm_inference_responses::ResponsesRoute::new( + provider_http(&resources, &http_config()), + resources.auth, + secrets, + ) +} + +pub async fn upstream(responses: impl IntoIterator) -> MockServer { + let server = MockServer::start().await; + respond_in_order(&server, responses).await; + server +} + +pub async fn respond_in_order( + server: &MockServer, + responses: impl IntoIterator, +) { + for response in responses { + Mock::given(any()) + .respond_with(response) + .up_to_n_times(1) + .mount(server) + .await; + } +} + +pub async fn received(server: &MockServer) -> Vec { + server + .received_requests() + .await + .expect("request recording is on") +} + +pub async fn only_request(server: &MockServer) -> Request { + let [request] = <[Request; 1]>::try_from(received(server).await) + .unwrap_or_else(|requests| panic!("expected one request, got {}", requests.len())); + request +} + +pub fn json_response(body: Value) -> ResponseTemplate { + ResponseTemplate::new(200).set_body_json(body) +} + +pub fn status_response(status: u16, body: Value) -> ResponseTemplate { + ResponseTemplate::new(status).set_body_json(body) +} + +pub trait ReceivedRequest { + fn header(&self, name: &str) -> Option<&str>; + fn header_values(&self, name: &str) -> Vec<&str>; + fn json(&self) -> Value; + fn body_text(&self) -> String; + /// The path and query, as the request line carried them. + fn target(&self) -> String; + fn query(&self, name: &str) -> Option; +} + +impl ReceivedRequest for Request { + fn header(&self, name: &str) -> Option<&str> { + self.headers.get(name).and_then(|value| value.to_str().ok()) + } + + fn header_values(&self, name: &str) -> Vec<&str> { + self.headers + .get_all(name) + .iter() + .filter_map(|value| value.to_str().ok()) + .collect() + } + + fn json(&self) -> Value { + serde_json::from_slice(&self.body).expect("request body is json") + } + + fn body_text(&self) -> String { + String::from_utf8_lossy(&self.body).into_owned() + } + + fn target(&self) -> String { + match self.url.query() { + Some(query) => format!("{}?{query}", self.url.path()), + None => self.url.path().to_string(), + } + } + + fn query(&self, name: &str) -> Option { + self.url + .query_pairs() + .find_map(|(key, value)| (key == name).then(|| value.into_owned())) + } +} + +pub struct RecordingCall { + pub request: Mutex>, + pub events: Arc, + pub chunks: Mutex>, + pub head: Mutex>, +} + +#[derive(Default)] +pub struct CallEvents(pub Observations); +pub struct Observations { + pub sender: litellm_host::observation::ObservationSender, + receiver: Mutex>, + recorded: Mutex>, +} + +impl Default for Observations { + fn default() -> Self { + let (sender, receiver) = litellm_host::observation::observation_channel( + std::num::NonZeroUsize::new(128).unwrap(), + ); + Self { + sender, + receiver: Mutex::new(receiver), + recorded: Mutex::new(Vec::new()), + } + } +} + +impl Observations { + pub fn lock( + &self, + ) -> std::sync::LockResult>> + { + let mut events = self.recorded.lock()?; + let mut receiver = self.receiver.lock().unwrap(); + while let Ok(event) = receiver.try_recv() { + events.push(event); + } + Ok(events) + } +} + +impl litellm_host::lifecycle::CallObserver for CallEvents { + fn observe(&self, event: litellm_host::lifecycle::CallEvent) { + self.0.sender.emit(event); + } +} + +impl RecordingCall

{ + pub fn new(request: P::Request) -> Self { + Self { + request: Mutex::new(Some(request)), + events: Arc::new(CallEvents::default()), + chunks: Mutex::new(Vec::new()), + head: Mutex::new(None), + } + } +} + +impl litellm_host::interceptors::Interceptors + for RecordingCall

+{ + async fn before_provider_request( + &self, + wire: litellm_host::interceptors::WireRequest, + _: litellm_host::interceptors::RequestContext, + ) -> Result { + Ok(litellm_host::interceptors::WireRequest { + headers: wire + .headers + .into_iter() + .chain([("x-hook".into(), "called".into())]) + .collect(), + ..wire + }) + } + + async fn after_provider_response( + &self, + _: litellm_host::interceptors::RawResponse, + ) -> Result<(), P::Error> { + Ok(()) + } +} + +impl

RecordingCall

+where + P: litellm_host::protocol::Protocol, + P::Error: From, +{ + pub fn request(&self) -> Result { + self.request + .lock() + .unwrap() + .take() + .ok_or_else(|| litellm_host::machine::MachineFault::Abandoned.into()) + } + pub fn runtime(&self) -> litellm_host_native::in_process::Host<'_, (), Self, Self> { + litellm_host_native::in_process::Host { + services: &(), + interceptors: self, + stream: self, + observers: Some(&self.events.0.sender), + } + } +} + +impl

litellm_host_native::in_process::StreamConsumer

for RecordingCall

+where + P: litellm_host::protocol::Protocol, + P::Error: From, +{ + async fn open_stream(&self, head: P::StreamHead) -> Result, P::Error> { + *self.head.lock().unwrap() = Some(head); + Ok(ControlFlow::Continue(())) + } + async fn send_chunk(&self, chunk: P::Chunk) -> Result, P::Error> { + self.chunks.lock().unwrap().push(chunk); + Ok(ControlFlow::Continue(())) + } +} +impl

litellm_host::lifecycle::CallObserver for RecordingCall

+where + P: litellm_host::protocol::Protocol, + P::Error: From, +{ + fn observe(&self, event: litellm_host::lifecycle::CallEvent) { + self.events.0.sender.emit(event); + } +} + +#[derive(Clone, Default)] +pub struct TraceCapture(Arc>>); + +impl TraceCapture { + pub fn logger(&self) -> litellm_tracing::Logger { + litellm_tracing::Logger::new(self.clone()) + } + + pub fn records(&self) -> Vec { + self.0.lock().unwrap().clone() + } + + pub fn summaries(&self, name: &str) -> Vec { + self.records() + .into_iter() + .filter(|record| record["span_name"] == name) + .collect() + } +} + +impl litellm_tracing::Sink for TraceCapture { + fn enabled(&self, metadata: &litellm_tracing::Metadata<'_>) -> bool { + metadata.target().starts_with("litellm_inference") + } + + fn emit(&self, record: &litellm_tracing::Record) { + self.0 + .lock() + .unwrap() + .push(Value::Object(record.fields.clone())); + } +} + +#[rstest::fixture] +pub fn traces() -> TraceCapture { + TraceCapture::default() +} diff --git a/litellm-rust/crates/inference/src/error.rs b/litellm-rust/crates/inference/src/error.rs index d42b38afe75..e5f02754f3d 100644 --- a/litellm-rust/crates/inference/src/error.rs +++ b/litellm-rust/crates/inference/src/error.rs @@ -52,7 +52,7 @@ impl From for RouteError { } impl RouteError { - pub(crate) fn post_call(error: Self) -> Self { + pub fn post_call(error: Self) -> Self { Self::PostCallHook(Arc::new(error)) } diff --git a/litellm-rust/crates/inference/src/lib.rs b/litellm-rust/crates/inference/src/lib.rs index cb308c7478d..19d62776bd8 100644 --- a/litellm-rust/crates/inference/src/lib.rs +++ b/litellm-rust/crates/inference/src/lib.rs @@ -10,7 +10,6 @@ pub mod ocr; pub mod outbound; pub mod provider; pub mod resources; -pub mod responses; #[cfg(feature = "test-support")] pub mod test_support; diff --git a/litellm-rust/crates/inference/tests/caching.rs b/litellm-rust/crates/inference/tests/caching.rs index ba356a62533..ac1aabdd723 100644 --- a/litellm-rust/crates/inference/tests/caching.rs +++ b/litellm-rust/crates/inference/tests/caching.rs @@ -12,8 +12,7 @@ use bytes::Bytes; use futures_util::{StreamExt, TryStreamExt, stream}; use litellm_cache_memory::InMemoryCache; use litellm_cache_response::{ - CacheOptions, CachePolicy, CacheScope, ResponseCache, ResponseCacheConfig, - ResponseCacheService, ResponseEnvelope, + CacheOptions, CachePolicy, CacheScope, ResponseCache, ResponseCacheConfig, ResponseCacheService, }; use litellm_host::{ call::{CallOutput, OutputOf}, @@ -403,50 +402,6 @@ async fn an_invalid_cached_envelope_is_replaced_by_a_provider_result(#[case] poi assert_eq!(calls.load(Ordering::SeqCst), 1); } -#[rstest] -#[case::chat_completion(json!({"kind":"Response","value":{"id":"chat-1","model":"test","choices":[],"usage":{"prompt_tokens":3,"completion_tokens":2}}}))] -#[case::wrong_envelope(json!({"kind":"Stream","value":"data: [DONE]\n\n"}))] -#[tokio::test] -async fn responses_refetches_instead_of_deserializing_another_api_response( - #[case] poisoned: Value, -) { - use litellm_inference::responses::route::Responses; - use litellm_llms_types::formats::responses::ResponsesApiResponse; - - let cache: Arc = Arc::new(InvalidEntryCache( - ResponseCache::new(Arc::new(InMemoryCache::default())), - serde_json::to_value(ResponseEnvelope::new("responses", poisoned)).unwrap(), - )); - let calls = AtomicUsize::new(0); - for _ in 0..2 { - let response = execute_unary::( - cache_request(json!({"input":"hello"})), - Some(cache.clone()), - Some(CacheOptions::new(CacheScope::Shared)), - &(), - None, - || async { - calls.fetch_add(1, Ordering::SeqCst); - Ok(ResponsesApiResponse { - id: "fresh-response".into(), - model: "test".into(), - output: vec![ - json!({"type":"message","content":[{"type":"output_text","text":"fresh"}]}), - ], - extra: [("status".into(), json!("completed"))] - .into_iter() - .collect(), - }) - }, - ) - .await - .unwrap(); - assert_eq!(response.id, "fresh-response"); - assert_eq!(response.output[0]["content"][0]["text"], "fresh"); - } - assert_eq!(calls.load(Ordering::SeqCst), 1); -} - #[rstest] #[case::system("system", json!("answer ALPHA"), json!("answer BETA"))] #[case::stop_sequences("stop_sequences", json!(["STOP"]), json!(["END"]))] @@ -833,43 +788,6 @@ async fn namespaces_and_surfaces_isolate_entries_on_shared_storage() { assert_eq!(calls.load(Ordering::SeqCst), 3); } -#[rstest] -#[case::completed("completed", 1)] -#[case::incomplete("incomplete", 2)] -#[tokio::test] -async fn responses_cache_only_reuses_completed_responses( - cache: Arc, - #[case] status: &str, - #[case] expected_calls: usize, -) { - use litellm_inference::responses::route::Responses; - use litellm_llms_types::formats::responses::ResponsesApiResponse; - - let calls = AtomicUsize::new(0); - for _ in 0..2 { - let response = execute_unary::( - cache_request(json!({"input":"hello"})), - Some(cache.clone()), - Some(CacheOptions::new(CacheScope::Shared)), - &(), - None, - || async { - let call = calls.fetch_add(1, Ordering::SeqCst); - Ok(ResponsesApiResponse { - id: call.to_string(), - model: "test".into(), - output: Vec::new(), - extra: [("status".into(), json!(status))].into_iter().collect(), - }) - }, - ) - .await - .unwrap(); - assert_eq!(response.extra.get("status"), Some(&json!(status))); - } - assert_eq!(calls.load(Ordering::SeqCst), expected_calls); -} - mod support; use support::traces; @@ -1028,9 +946,6 @@ impl Interceptors for ChangingHooks { #[case::messages_credentials("messages", "credentials")] #[case::messages_endpoint("messages", "endpoint")] #[case::messages_callback("messages", "callback")] -#[case::responses_credentials("responses", "credentials")] -#[case::responses_endpoint("responses", "endpoint")] -#[case::responses_callback("responses", "callback")] #[tokio::test] async fn cache_identity_follows_resolved_configuration_and_request_callbacks( cache: Arc, @@ -1041,19 +956,14 @@ async fn cache_identity_follows_resolved_configuration_and_request_callbacks( use litellm_inference::{ chat_completions::{ChatCompletionsRoute, types::ChatCompletionsRequest}, messages::MessagesCall, - responses::types::ResponsesCall, }; use wiremock::{Mock, MockServer, ResponseTemplate, matchers::method}; let first = MockServer::start().await; let second = MockServer::start().await; - let response = if surface == "responses" { - json!({"id":"response-test", "model":"test", "output":[], "status":"completed"}) - } else { - json!({"id":"message-test", "type":"message", "role":"assistant", "model":"test", - "content":[{"type":"text", "text":"answer"}], "stop_reason":"end_turn", "stop_sequence":null, - "usage":{"input_tokens":3,"output_tokens":2}}) - }; + let response = json!({"id":"message-test", "type":"message", "role":"assistant", "model":"test", + "content":[{"type":"text", "text":"answer"}], "stop_reason":"end_turn", "stop_sequence":null, + "usage":{"input_tokens":3,"output_tokens":2}}); Mock::given(method("POST")) .respond_with(ResponseTemplate::new(200).set_body_json(response.clone())) .expect(if change == "endpoint" { 1 } else { 2 }) @@ -1116,26 +1026,6 @@ async fn cache_identity_follows_resolved_configuration_and_request_callbacks( api_key:None,api_base:None,custom_llm_provider:None,extra_headers:None,provider_specific_header:None,timeout:None,shaping:Default::default(), }, &hooks, None).await.unwrap(); } - "responses" => { - support::responses_route(secrets.clone()) - .with_cache(cache) - .execute( - ResponsesCall { - model: "test".into(), - input: json!("hello"), - optional_params: Default::default(), - api_key: None, - api_base: None, - custom_llm_provider: None, - extra_headers: None, - timeout: None, - }, - &hooks, - None, - ) - .await - .unwrap(); - } _ => unreachable!(), } } @@ -1153,12 +1043,10 @@ async fn cache_identity_follows_resolved_configuration_and_request_callbacks( } let requests = first.received_requests().await.unwrap(); if change == "credentials" { - let header = if surface == "responses" { - "authorization" - } else { - "x-api-key" - }; - assert_ne!(requests[0].headers[header], requests[1].headers[header]); + assert_ne!( + requests[0].headers["x-api-key"], + requests[1].headers["x-api-key"] + ); } if change == "callback" { assert_eq!( diff --git a/litellm-rust/crates/inference/tests/support/mod.rs b/litellm-rust/crates/inference/tests/support/mod.rs index 196f3043da0..41e08e6dcb7 100644 --- a/litellm-rust/crates/inference/tests/support/mod.rs +++ b/litellm-rust/crates/inference/tests/support/mod.rs @@ -37,17 +37,6 @@ pub fn chat_completions_route() -> litellm_inference::chat_completions::ChatComp ) } -pub fn responses_route( - secrets: Arc, -) -> litellm_inference::responses::ResponsesRoute { - let resources = resources(); - litellm_inference::responses::ResponsesRoute::new( - provider_http(&resources, &http_config()), - resources.auth, - secrets, - ) -} - pub fn build_ocr_route( resources: &litellm_inference::resources::CoreResources, config: &HttpClientConfig, diff --git a/litellm-rust/crates/python-bridge/Cargo.toml b/litellm-rust/crates/python-bridge/Cargo.toml index 591935ffd98..69991e7c756 100644 --- a/litellm-rust/crates/python-bridge/Cargo.toml +++ b/litellm-rust/crates/python-bridge/Cargo.toml @@ -38,6 +38,7 @@ litellm-auth-aws.workspace = true litellm-callbacks-legacy-python.workspace = true litellm-inference.workspace = true litellm-inference-transcription.workspace = true +litellm-inference-responses.workspace = true litellm-core-utils.workspace = true litellm-http.workspace = true litellm-llms.workspace = true diff --git a/litellm-rust/crates/python-bridge/src/routes/responses.rs b/litellm-rust/crates/python-bridge/src/routes/responses.rs index 527bf51feab..840bca89b51 100644 --- a/litellm-rust/crates/python-bridge/src/routes/responses.rs +++ b/litellm-rust/crates/python-bridge/src/routes/responses.rs @@ -1,6 +1,6 @@ mod host; -use litellm_inference::responses::websocket::ResponsesWebSocketConnection as RustResponsesWebSocketConnection; +use litellm_inference_responses::websocket::ResponsesWebSocketConnection as RustResponsesWebSocketConnection; use pyo3::{ prelude::*, types::{PyDict, PyTuple}, @@ -81,7 +81,7 @@ fn run_public( py, arguments, move |py, arguments, request| { - let route = litellm_inference::responses::ResponsesRoute::new( + let route = litellm_inference_responses::ResponsesRoute::new( crate::http::provider_client(py, arguments, asynchronous)? .map_err(crate::http::client_error)?, crate::http::resources().auth.clone(), diff --git a/litellm-rust/crates/python-bridge/src/routes/responses/host.rs b/litellm-rust/crates/python-bridge/src/routes/responses/host.rs index cb2750e9398..25d3fc07332 100644 --- a/litellm-rust/crates/python-bridge/src/routes/responses/host.rs +++ b/litellm-rust/crates/python-bridge/src/routes/responses/host.rs @@ -2,7 +2,7 @@ use std::convert::Infallible; use super::super::inference::InferenceHost; use litellm_host_python::{InvokeError, PythonBinding, PythonHostCalls, PythonOwned}; -use litellm_inference::responses::{Error, route::Responses, types::ResponsesCall}; +use litellm_inference_responses::{Error, route::Responses, types::ResponsesCall}; use pyo3::{ gc::{PyTraverseError, PyVisit}, prelude::*,