From 18933c8a21565b3b8e0689ca8b9428c0e31994d0 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sun, 27 Sep 2026 16:18:10 -0700 Subject: [PATCH] feat(rust): add the HTTP host driver (#43462) Co-authored-by: Yujong Lee --- litellm-rust/Cargo.lock | 15 + litellm-rust/Cargo.toml | 1 + litellm-rust/crates/host-http/AGENTS.md | 7 + litellm-rust/crates/host-http/Cargo.toml | 20 + litellm-rust/crates/host-http/src/driver.rs | 165 ++++ litellm-rust/crates/host-http/src/encoding.rs | 57 ++ litellm-rust/crates/host-http/src/error.rs | 7 + litellm-rust/crates/host-http/src/lib.rs | 9 + litellm-rust/crates/host-http/src/sse.rs | 58 ++ litellm-rust/crates/host-http/tests/serve.rs | 776 ++++++++++++++++++ litellm-rust/crates/host-http/tests/sse.rs | 109 +++ 11 files changed, 1224 insertions(+) create mode 100644 litellm-rust/crates/host-http/AGENTS.md create mode 100644 litellm-rust/crates/host-http/Cargo.toml create mode 100644 litellm-rust/crates/host-http/src/driver.rs create mode 100644 litellm-rust/crates/host-http/src/encoding.rs create mode 100644 litellm-rust/crates/host-http/src/error.rs create mode 100644 litellm-rust/crates/host-http/src/lib.rs create mode 100644 litellm-rust/crates/host-http/src/sse.rs create mode 100644 litellm-rust/crates/host-http/tests/serve.rs create mode 100644 litellm-rust/crates/host-http/tests/sse.rs diff --git a/litellm-rust/Cargo.lock b/litellm-rust/Cargo.lock index a74ec097148..fcf9df25cc8 100644 --- a/litellm-rust/Cargo.lock +++ b/litellm-rust/Cargo.lock @@ -3287,6 +3287,21 @@ dependencies = [ "tokio", ] +[[package]] +name = "litellm-host-http" +version = "0.1.0" +dependencies = [ + "axum", + "bytes", + "futures-util", + "http 1.4.2", + "litellm-host", + "rstest", + "serde_json", + "thiserror 2.0.19", + "tokio", +] + [[package]] name = "litellm-host-python" version = "0.1.0" diff --git a/litellm-rust/Cargo.toml b/litellm-rust/Cargo.toml index ed703396c22..e73cc7cccb4 100644 --- a/litellm-rust/Cargo.toml +++ b/litellm-rust/Cargo.toml @@ -18,6 +18,7 @@ litellm-gateway-inference = { path = "crates/gateway-inference" } litellm-gateway-auth = { path = "crates/gateway-auth" } litellm-coroutine = { path = "crates/coroutine" } litellm-host = { path = "crates/host" } +litellm-host-http = { path = "crates/host-http" } litellm-callbacks-legacy-python = { path = "crates/callbacks-legacy-python" } litellm-framing = { path = "crates/framer" } litellm-auth = { path = "crates/auth" } diff --git a/litellm-rust/crates/host-http/AGENTS.md b/litellm-rust/crates/host-http/AGENTS.md new file mode 100644 index 00000000000..f41af54f6d0 --- /dev/null +++ b/litellm-rust/crates/host-http/AGENTS.md @@ -0,0 +1,7 @@ +Own the HTTP driver for hosted calls, including response-body demand, cancellation, and lifecycle observation + +Keep endpoint paths, request parsing, deployment selection, and API-specific response and error formats in gateway-inference + +Depend on the neutral host protocol, never on core routes, Python, or gateway crates + +Do not spawn producer tasks or buffer chunks ahead of HTTP body demand diff --git a/litellm-rust/crates/host-http/Cargo.toml b/litellm-rust/crates/host-http/Cargo.toml new file mode 100644 index 00000000000..4555d2fba49 --- /dev/null +++ b/litellm-rust/crates/host-http/Cargo.toml @@ -0,0 +1,20 @@ +[package] +name = "litellm-host-http" +version = "0.1.0" +edition.workspace = true +license.workspace = true +repository.workspace = true + +[dependencies] +axum.workspace = true +bytes.workspace = true +futures-util.workspace = true +http.workspace = true +litellm-host.workspace = true +thiserror.workspace = true + +[dev-dependencies] +axum = { workspace = true, features = ["json"] } +rstest.workspace = true +serde_json.workspace = true +tokio.workspace = true diff --git a/litellm-rust/crates/host-http/src/driver.rs b/litellm-rust/crates/host-http/src/driver.rs new file mode 100644 index 00000000000..704b93c6b6b --- /dev/null +++ b/litellm-rust/crates/host-http/src/driver.rs @@ -0,0 +1,165 @@ +use std::{convert::Infallible, sync::Arc}; + +use axum::{body::Body, response::Response}; +use bytes::Bytes; +use futures_util::{StreamExt, stream}; + +use litellm_host::{ + call::{CallOutput, HostedCompletion, HostedMachine}, + hooks::RouteHooks, + lifecycle::{observe_call, observe_unary}, + machine::{Machine, MachineFault, MachineStep}, + protocol::{Demand, HookRequest, Protocol, Reply, StreamDelivery, Suspension}, + services::HostCallHandler, +}; + +use crate::{Error, ResponseEncoder, StreamEncoder}; + +type StepOf

= MachineStep::Response>>; +type Output = CallOutput, Bytes, E>; + +pub async fn serve_unary( + machine: HostedMachine

, + services: S, + hooks: H, + encoder: A, +) -> Result> +where + P: Protocol, + P::Error: From, + H: RouteHooks, + S: HostCallHandler

, + A: ResponseEncoder, +{ + let observer = hooks.observer(); + let mut driver = Driver::new(machine, services, hooks); + observe_unary(observer, async move { + match driver.advance().await? { + MachineStep::Complete(HostedCompletion::Complete(value)) => { + encoder.encode_response(value).map_err(Error::Call) + } + _ => Err(Error::Protocol), + } + }) + .await +} + +pub async fn serve( + machine: HostedMachine

, + services: S, + hooks: H, + encoder: A, +) -> Result> +where + P: Protocol, + P::Error: From, + A: StreamEncoder, + H: RouteHooks + 'static, + S: HostCallHandler

+ 'static, +{ + let encoder = Arc::new(encoder); + let observer = hooks.observer(); + let driver = Driver::new(machine, services, hooks); + match observe_call(observer, driver.start(encoder.clone())).await? { + CallOutput::Complete(response) => Ok(response), + CallOutput::Stream { head, chunks } => { + let body = chunks.map(move |chunk| { + Ok::<_, Infallible>( + chunk.unwrap_or_else(|error| encoder.encode_stream_error(error)), + ) + }); + Ok(head.map(|()| Body::from_stream(body))) + } + } +} + +struct Driver { + machine: HostedMachine

, + services: S, + hooks: H, + demand: Option>, +} + +impl Driver +where + P: Protocol, + P::Error: From, + H: RouteHooks, + S: HostCallHandler

, +{ + fn new(machine: HostedMachine

, services: S, hooks: H) -> Self { + Self { + machine, + services, + hooks, + demand: None, + } + } + + async fn advance(&mut self) -> Result, Error> { + if let Some(reply) = self.demand.take() { + reply.send(Demand::More); + } + loop { + match self.machine.resume().await.map_err(Error::Call)? { + MachineStep::Suspended(Suspension::Hook(HookRequest::BeforeProviderRequest { + wire, + context, + reply, + })) => reply.send( + self.hooks + .before_provider_request(*wire, *context) + .await + .map_err(Error::Call)?, + ), + MachineStep::Suspended(Suspension::Hook(HookRequest::Event(event, reply))) => { + self.hooks.on_event(event).await.map_err(Error::Call)?; + reply.send(()); + } + MachineStep::Suspended(Suspension::HostCall(op)) => { + self.services + .handle_host_call(op) + .await + .map_err(Error::Call)?; + } + boundary => return Ok(boundary), + } + } + } + + async fn start(mut self, encoder: Arc) -> Result>, Error> + where + A: StreamEncoder, + H: 'static, + S: 'static, + { + match self.advance().await? { + MachineStep::Complete(HostedCompletion::Complete(value)) => encoder + .encode_response(value) + .map(CallOutput::Complete) + .map_err(Error::Call), + MachineStep::Suspended(Suspension::Stream(StreamDelivery::Open(head, reply))) => { + let head = encoder.encode_stream_head(head).map_err(Error::Call)?; + self.demand = Some(reply); + let chunks = + stream::try_unfold((self, encoder), |(mut driver, encoder)| async move { + match driver.advance().await? { + MachineStep::Suspended(Suspension::Stream(StreamDelivery::Chunk( + chunk, + reply, + ))) => { + let bytes = encoder.encode_chunk(chunk).map_err(Error::Call)?; + driver.demand = Some(reply); + Ok(Some((bytes, (driver, encoder)))) + } + MachineStep::Complete(HostedCompletion::StreamEnded) => Ok(None), + _ => Err(Error::Protocol), + } + }) + .boxed(); + Ok(CallOutput::Stream { head, chunks }) + } + _ => Err(Error::Protocol), + } + } +} diff --git a/litellm-rust/crates/host-http/src/encoding.rs b/litellm-rust/crates/host-http/src/encoding.rs new file mode 100644 index 00000000000..ebdb572abef --- /dev/null +++ b/litellm-rust/crates/host-http/src/encoding.rs @@ -0,0 +1,57 @@ +use std::marker::PhantomData; + +use axum::response::{IntoResponse, Response}; +use bytes::Bytes; +use litellm_host::protocol::Protocol; + +use crate::Error; + +pub trait ResponseEncoder: Send + Sync { + type Protocol: Protocol; + + fn encode_response( + &self, + response: ::Response, + ) -> Result::Error>; +} + +pub trait StreamEncoder: ResponseEncoder + 'static { + fn encode_stream_head( + &self, + head: ::StreamHead, + ) -> Result, ::Error>; + + fn encode_chunk( + &self, + chunk: ::Chunk, + ) -> Result::Error>; + + fn encode_stream_error(&self, error: Error<::Error>) -> Bytes; +} + +pub struct Unary { + response: F, + protocol: PhantomData P>, +} + +impl Unary { + pub fn new(response: F) -> Self { + Self { + response, + protocol: PhantomData, + } + } +} + +impl ResponseEncoder for Unary +where + P: Protocol, + F: Fn(P::Response) -> R + Send + Sync, + R: IntoResponse, +{ + type Protocol = P; + + fn encode_response(&self, response: P::Response) -> Result { + Ok((self.response)(response).into_response()) + } +} diff --git a/litellm-rust/crates/host-http/src/error.rs b/litellm-rust/crates/host-http/src/error.rs new file mode 100644 index 00000000000..285d08d2fe8 --- /dev/null +++ b/litellm-rust/crates/host-http/src/error.rs @@ -0,0 +1,7 @@ +#[derive(Debug, PartialEq, thiserror::Error)] +pub enum Error { + #[error(transparent)] + Call(E), + #[error("unexpected HTTP host operation")] + Protocol, +} diff --git a/litellm-rust/crates/host-http/src/lib.rs b/litellm-rust/crates/host-http/src/lib.rs new file mode 100644 index 00000000000..a4714ef72b2 --- /dev/null +++ b/litellm-rust/crates/host-http/src/lib.rs @@ -0,0 +1,9 @@ +mod driver; +mod encoding; +mod error; +mod sse; + +pub use driver::{serve, serve_unary}; +pub use encoding::{ResponseEncoder, StreamEncoder, Unary}; +pub use error::Error; +pub use sse::Sse; diff --git a/litellm-rust/crates/host-http/src/sse.rs b/litellm-rust/crates/host-http/src/sse.rs new file mode 100644 index 00000000000..8f7c8ffa3b3 --- /dev/null +++ b/litellm-rust/crates/host-http/src/sse.rs @@ -0,0 +1,58 @@ +use axum::response::{IntoResponse, Response}; +use bytes::Bytes; +use http::{HeaderValue, header::CONTENT_TYPE}; +use litellm_host::protocol::Protocol; + +use crate::{Error, ResponseEncoder, StreamEncoder, Unary}; + +pub struct Sse { + response: Unary, + stream_error: F, +} + +impl Sse { + pub fn new(response: C, stream_error: F) -> Self { + Self { + response: Unary::new(response), + stream_error, + } + } +} + +impl ResponseEncoder for Sse +where + P: Protocol, + C: Fn(P::Response) -> R + Send + Sync, + F: Send + Sync, + R: IntoResponse, +{ + type Protocol = P; + + fn encode_response(&self, response: P::Response) -> Result { + self.response.encode_response(response) + } +} + +impl StreamEncoder for Sse +where + P: Protocol, + C: Fn(P::Response) -> R + Send + Sync + 'static, + F: Fn(Error) -> Bytes + Send + Sync + 'static, + R: IntoResponse, +{ + fn encode_stream_head(&self, _: P::StreamHead) -> Result, P::Error> { + let mut response = http::Response::new(()); + response + .headers_mut() + .insert(CONTENT_TYPE, HeaderValue::from_static("text/event-stream")); + Ok(response) + } + + fn encode_chunk(&self, chunk: Bytes) -> Result { + Ok(chunk) + } + + fn encode_stream_error(&self, error: Error) -> Bytes { + (self.stream_error)(error) + } +} diff --git a/litellm-rust/crates/host-http/tests/serve.rs b/litellm-rust/crates/host-http/tests/serve.rs new file mode 100644 index 00000000000..b3ae805a7f6 --- /dev/null +++ b/litellm-rust/crates/host-http/tests/serve.rs @@ -0,0 +1,776 @@ +use std::sync::{ + Arc, Mutex, + atomic::{AtomicBool, AtomicUsize, Ordering}, +}; + +use axum::{ + Json, + body::to_bytes, + response::{IntoResponse, Response}, +}; +use bytes::Bytes; +use futures_util::{StreamExt, stream}; +use http::{StatusCode, header::CONTENT_TYPE}; +use litellm_host::{ + call::{CallOutput, hosted_call}, + event::{CallEvent, MachineEvent, RawResponse, RequestContext, WireRequest}, + hooks::RouteHooks, + lifecycle::CallObserver, + machine::MachineFault, + protocol::Protocol, + protocol::Reply, +}; +use litellm_host_http::{Error, ResponseEncoder, StreamEncoder, Unary, serve, serve_unary}; +use rstest::{fixture, rstest}; +use serde_json::json; + +#[derive(Clone, Debug, PartialEq)] +enum TestError { + Provider, + Adapter, + Hook, + Machine, +} + +impl From for TestError { + fn from(_: MachineFault) -> Self { + Self::Machine + } +} + +struct TestProtocol; + +impl Protocol for TestProtocol { + type Response = Bytes; + type Error = TestError; + type Request = &'static str; + type HostCall = Reply<&'static str>; + type Chunk = Bytes; + type StreamHead = &'static str; +} + +#[derive(Clone, Copy, PartialEq)] +enum Rejection { + None, + Head, + Chunk, + Custom, + Complete, +} + +struct Adapter(Rejection); + +impl ResponseEncoder for Adapter { + type Protocol = TestProtocol; + + fn encode_response(&self, value: Bytes) -> Result { + if self.0 == Rejection::Complete { + return Err(TestError::Adapter); + } + Ok((StatusCode::CREATED, [("x-converted", "yes")], value).into_response()) + } +} + +impl litellm_host::services::HostCallHandler for Adapter { + async fn handle_host_call(&self, reply: Reply<&'static str>) -> Result<(), TestError> { + if self.0 == Rejection::Custom { + return Err(TestError::Adapter); + } + reply.send("custom"); + Ok(()) + } +} + +impl StreamEncoder for Adapter { + fn encode_stream_head( + &self, + content_type: &'static str, + ) -> Result, TestError> { + if self.0 == Rejection::Head { + return Err(TestError::Adapter); + } + Ok(http::Response::builder() + .status(StatusCode::ACCEPTED) + .header(CONTENT_TYPE, content_type) + .body(()) + .unwrap()) + } + + fn encode_chunk(&self, chunk: Bytes) -> Result { + if self.0 == Rejection::Chunk { + return Err(TestError::Adapter); + } + Ok(chunk) + } + + fn encode_stream_error(&self, error: Error) -> Bytes { + Bytes::from(format!("event: error\ndata: {error:?}\n\n")) + } +} + +#[derive(Default)] +struct Observer(Mutex>); + +impl CallObserver for Observer { + fn observe(&self, event: CallEvent) { + self.0.lock().unwrap().push(event); + } +} + +struct Hooks { + observer: Arc, + reject: bool, +} + +impl RouteHooks for Hooks { + fn observer(&self) -> Option> { + Some(self.observer.clone()) + } + + async fn before_provider_request( + &self, + wire: WireRequest, + context: RequestContext, + ) -> Result { + if self.reject { + return Err(TestError::Hook); + } + Ok(WireRequest { + url: format!("{}/{}", wire.url, context.model), + ..wire + }) + } + + async fn on_event(&self, event: MachineEvent) -> Result<(), TestError> { + if self.reject { + return Err(TestError::Hook); + } + self.observer.observe(CallEvent::Machine(event)); + Ok(()) + } +} + +#[fixture] +fn observer() -> Arc { + Arc::new(Observer::default()) +} + +#[fixture] +fn hooks(observer: Arc) -> Hooks { + Hooks { + observer, + reject: false, + } +} + +#[rstest] +#[tokio::test] +async fn projection_custom_operations_and_hooks_feed_the_http_response(hooks: Hooks) { + let observer = hooks.observer.clone(); + let machine = hosted_call::( + "projected", + |request, services, route_hooks| async move { + let custom = services.call(|reply| reply).await?; + let wire = route_hooks + .before_provider_request( + WireRequest { + url: request.into(), + headers: Vec::new(), + body: json!({}), + }, + RequestContext { + model: custom.into(), + custom_llm_provider: "test".into(), + optional_params: json!({}), + secret_fields: Vec::new(), + api_key: None, + }, + ) + .await?; + route_hooks + .on_event(MachineEvent::ResponseReceived { + raw: RawResponse { + body: wire.url.clone(), + }, + }) + .await?; + Ok(CallOutput::Complete(Bytes::from(wire.url))) + }, + ); + let response = serve( + machine, + Adapter(Rejection::None), + hooks, + Adapter(Rejection::None), + ) + .await + .unwrap(); + assert_eq!(response.status(), StatusCode::CREATED); + assert_eq!(response.headers()["x-converted"], "yes"); + assert_eq!( + to_bytes(response.into_body(), 1024).await.unwrap(), + "projected/custom" + ); + let events = observer.0.lock().unwrap(); + assert!(matches!(events.as_slice(), [ + CallEvent::Started { .. }, + CallEvent::Machine(MachineEvent::ResponseReceived { raw }), + CallEvent::Succeeded { .. }, + ] if raw.body == "projected/custom")); +} + +struct Release(Arc); + +impl Drop for Release { + fn drop(&mut self) { + self.0.store(true, Ordering::SeqCst); + } +} + +#[rstest] +#[case::consumed(None)] +#[case::dropped_at_open(Some(0))] +#[case::dropped_after_chunk(Some(1))] +#[case::dropped_before_eof(Some(2))] +#[tokio::test] +async fn body_demand_controls_polling_and_lifecycle( + hooks: Hooks, + #[case] drop_after: Option, +) { + let observer = hooks.observer.clone(); + let polls = Arc::new(AtomicUsize::new(0)); + let released = Arc::new(AtomicBool::new(false)); + let provider_polls = polls.clone(); + let release = Release(released.clone()); + let machine = hosted_call::("input", move |_, _, _| async move { + let chunks = stream::unfold((0, release), move |(index, release)| { + provider_polls.fetch_add(1, Ordering::SeqCst); + async move { + (index < 2).then(|| (Ok(Bytes::from(index.to_string())), (index + 1, release))) + } + }) + .boxed(); + Ok(CallOutput::Stream { + head: "text/event-stream", + chunks, + }) + }); + let response = serve( + machine, + Adapter(Rejection::None), + hooks, + Adapter(Rejection::None), + ) + .await + .unwrap(); + assert_eq!(response.status(), StatusCode::ACCEPTED); + assert_eq!(response.headers()[CONTENT_TYPE], "text/event-stream"); + assert_eq!(polls.load(Ordering::SeqCst), 0); + assert!(matches!( + observer.0.lock().unwrap().as_slice(), + [CallEvent::Started { .. }] + )); + let mut body = response.into_body().into_data_stream(); + for index in 0..drop_after.unwrap_or(2) { + assert_eq!(body.next().await.unwrap().unwrap(), index.to_string()); + assert_eq!(polls.load(Ordering::SeqCst), index + 1); + assert_eq!(observer.0.lock().unwrap().len(), 1); + } + if drop_after.is_none() { + assert!(body.next().await.is_none()); + assert_eq!(polls.load(Ordering::SeqCst), 3); + } + drop(body); + assert!(released.load(Ordering::SeqCst)); + let events = observer.0.lock().unwrap(); + assert_eq!(events.len(), 2); + assert_eq!( + matches!(events[1], CallEvent::Cancelled { .. }), + drop_after.is_some() + ); + assert_eq!( + matches!(events[1], CallEvent::Succeeded { .. }), + drop_after.is_none() + ); +} + +#[rstest] +#[case::provider(Rejection::None, TestError::Provider, 2)] +#[case::encoding(Rejection::Chunk, TestError::Adapter, 1)] +#[tokio::test] +async fn stream_failure_emits_one_error_frame_and_stops( + hooks: Hooks, + #[case] rejection: Rejection, + #[case] expected: TestError, + #[case] expected_polls: usize, +) { + let observer = hooks.observer.clone(); + let polls = Arc::new(AtomicUsize::new(0)); + let provider_polls = polls.clone(); + let machine = hosted_call::("input", move |_, _, _| async move { + let chunks = stream::iter([ + Ok(Bytes::from_static(b"first")), + Err(TestError::Provider), + Ok(Bytes::from_static(b"must not be delivered")), + ]) + .inspect(move |_| { + provider_polls.fetch_add(1, Ordering::SeqCst); + }) + .boxed(); + Ok(CallOutput::Stream { + head: "text/event-stream", + chunks, + }) + }); + let response = serve(machine, Adapter(rejection), hooks, Adapter(rejection)) + .await + .unwrap(); + let body = to_bytes(response.into_body(), 1024).await.unwrap(); + let prefix = if rejection == Rejection::Chunk { + "" + } else { + "first" + }; + assert_eq!( + body, + format!( + "{prefix}event: error\ndata: {:?}\n\n", + Error::Call(expected) + ) + ); + assert_eq!(polls.load(Ordering::SeqCst), expected_polls); + assert!(matches!( + observer.0.lock().unwrap().as_slice(), + [CallEvent::Started { .. }, CallEvent::Failed { .. },] + )); +} + +#[rstest] +#[case::provider(Rejection::None)] +#[case::headers(Rejection::Head)] +#[case::custom_operation(Rejection::Custom)] +#[case::response_conversion(Rejection::Complete)] +#[tokio::test] +async fn failures_before_open_return_an_error(hooks: Hooks, #[case] rejection: Rejection) { + let observer = hooks.observer.clone(); + let machine = hosted_call::("input", move |_, services, _| async move { + services.call(|reply| reply).await?; + match rejection { + Rejection::Head => Ok(CallOutput::Stream { + head: "text/event-stream", + chunks: stream::pending().boxed(), + }), + Rejection::None => Err(TestError::Provider), + _ => Ok(CallOutput::Complete(Bytes::new())), + } + }); + let expected = if rejection == Rejection::None { + TestError::Provider + } else { + TestError::Adapter + }; + assert_eq!( + serve(machine, Adapter(rejection), hooks, Adapter(rejection)) + .await + .unwrap_err(), + Error::Call(expected) + ); + assert!(matches!( + observer.0.lock().unwrap().as_slice(), + [CallEvent::Started { .. }, CallEvent::Failed { .. },] + )); +} + +#[rstest] +#[case::before_headers(false)] +#[case::awaiting_chunk(true)] +#[tokio::test] +async fn cancelling_pending_work_releases_the_machine(hooks: Hooks, #[case] streaming: bool) { + let observer = hooks.observer.clone(); + let released = Arc::new(AtomicBool::new(false)); + let release = Release(released.clone()); + let machine = hosted_call::("input", move |_, _, _| async move { + if !streaming { + let _release = release; + return std::future::pending().await; + } + let chunks = stream::once(async move { + let _release = release; + std::future::pending().await + }) + .boxed(); + Ok(CallOutput::Stream { + head: "text/event-stream", + chunks, + }) + }); + let mut response = Box::pin(serve( + machine, + Adapter(Rejection::None), + hooks, + Adapter(Rejection::None), + )); + if streaming { + let mut body = response.await.unwrap().into_body().into_data_stream(); + assert!(futures_util::poll!(body.next()).is_pending()); + assert!(!released.load(Ordering::SeqCst)); + drop(body); + } else { + assert!(futures_util::poll!(&mut response).is_pending()); + assert!(!released.load(Ordering::SeqCst)); + drop(response); + } + assert!(released.load(Ordering::SeqCst)); + assert!(matches!( + observer.0.lock().unwrap().as_slice(), + [CallEvent::Started { .. }, CallEvent::Cancelled { .. },] + )); +} + +#[rstest] +#[case::before_provider_request(false)] +#[case::event(true)] +#[tokio::test] +async fn hook_rejection_stops_execution_and_is_reported_once( + observer: Arc, + #[case] event: bool, +) { + let continued = Arc::new(AtomicBool::new(false)); + let executed = continued.clone(); + let machine = hosted_call::("input", move |_, _, route_hooks| async move { + if event { + route_hooks + .on_event(MachineEvent::ResponseReceived { + raw: RawResponse { + body: "response".into(), + }, + }) + .await?; + } else { + route_hooks + .before_provider_request( + WireRequest { + url: "url".into(), + headers: Vec::new(), + body: json!({}), + }, + RequestContext { + model: "model".into(), + custom_llm_provider: "provider".into(), + optional_params: json!({}), + secret_fields: Vec::new(), + api_key: None, + }, + ) + .await?; + } + executed.store(true, Ordering::SeqCst); + Ok(CallOutput::Complete(Bytes::new())) + }); + let hooks = Hooks { + observer: observer.clone(), + reject: true, + }; + assert_eq!( + serve( + machine, + Adapter(Rejection::None), + hooks, + Adapter(Rejection::None) + ) + .await + .unwrap_err(), + Error::Call(TestError::Hook) + ); + assert!(!continued.load(Ordering::SeqCst)); + assert!(matches!( + observer.0.lock().unwrap().as_slice(), + [CallEvent::Started { .. }, CallEvent::Failed { .. },] + )); +} + +#[derive(Clone, Copy)] +enum InvalidFlow { + DeliverBeforeOpen, + OpenTwice, +} + +#[rstest] +#[case::deliver_before_open(InvalidFlow::DeliverBeforeOpen)] +#[case::open_twice(InvalidFlow::OpenTwice)] +#[tokio::test] +async fn invalid_host_operations_fail_without_panicking(hooks: Hooks, #[case] flow: InvalidFlow) { + use litellm_host::{call::HostedCompletion, machine::CallMachine}; + + let observer = hooks.observer.clone(); + let machine = CallMachine::>::new(move |host| { + Box::pin(async move { + match flow { + InvalidFlow::DeliverBeforeOpen => { + host.stream.send_chunk(Bytes::new()).await?; + } + InvalidFlow::OpenTwice => { + host.stream.open_stream("text/event-stream").await?; + host.stream.open_stream("text/event-stream").await?; + } + } + Ok(HostedCompletion::StreamEnded) + }) + }); + let result = serve( + machine, + Adapter(Rejection::None), + hooks, + Adapter(Rejection::None), + ) + .await; + if matches!(flow, InvalidFlow::OpenTwice) { + let body = to_bytes(result.unwrap().into_body(), 1024).await.unwrap(); + assert_eq!(body, "event: error\ndata: Protocol\n\n"); + } else { + assert_eq!(result.unwrap_err(), Error::Protocol); + } + assert!(matches!( + observer.0.lock().unwrap().as_slice(), + [CallEvent::Started { .. }, CallEvent::Failed { .. },] + )); +} + +struct UnaryProtocol; + +impl Protocol for UnaryProtocol { + type Response = serde_json::Value; + type Error = TestError; + type Request = &'static str; + type HostCall = std::convert::Infallible; + type Chunk = std::convert::Infallible; + type StreamHead = std::convert::Infallible; +} + +#[rstest] +#[tokio::test] +async fn unary_calls_use_into_response_after_hooks_and_before_success(hooks: Hooks) { + let observer = hooks.observer.clone(); + let machine = + hosted_call::("projected", |request, _, route_hooks| async move { + let wire = route_hooks + .before_provider_request( + WireRequest { + url: request.into(), + headers: Vec::new(), + body: json!({}), + }, + RequestContext { + model: "rewritten".into(), + custom_llm_provider: "test".into(), + optional_params: json!({}), + secret_fields: Vec::new(), + api_key: None, + }, + ) + .await?; + route_hooks + .on_event(MachineEvent::ResponseReceived { + raw: RawResponse { + body: wire.url.clone(), + }, + }) + .await?; + Ok(CallOutput::Complete(json!({"url": wire.url}))) + }); + let response = serve_unary( + machine, + (), + hooks, + Unary::new(|value| { + assert!(matches!( + observer.0.lock().unwrap().as_slice(), + [CallEvent::Started { .. }, CallEvent::Machine(_),] + )); + (StatusCode::CREATED, [("x-converted", "yes")], Json(value)) + }), + ) + .await + .unwrap(); + assert_eq!(response.status(), StatusCode::CREATED); + assert_eq!(response.headers()["x-converted"], "yes"); + assert_eq!(response.headers()[CONTENT_TYPE], "application/json"); + let body: serde_json::Value = + serde_json::from_slice(&to_bytes(response.into_body(), 1024).await.unwrap()).unwrap(); + assert_eq!(body, json!({"url": "projected/rewritten"})); + assert!(matches!( + observer.0.lock().unwrap().as_slice(), + [ + CallEvent::Started { .. }, + CallEvent::Machine(_), + CallEvent::Succeeded { .. }, + ] + )); +} + +#[rstest] +#[case::provider(false, TestError::Provider)] +#[case::hook(true, TestError::Hook)] +#[tokio::test] +async fn unary_failure_preserves_the_error_without_converting( + observer: Arc, + #[case] reject_hook: bool, + #[case] expected: TestError, +) { + let machine = hosted_call::("input", |_, _, route_hooks| async move { + route_hooks + .on_event(MachineEvent::ResponseReceived { + raw: RawResponse { body: "raw".into() }, + }) + .await?; + Err(TestError::Provider) + }); + let hooks = Hooks { + observer: observer.clone(), + reject: reject_hook, + }; + let converted = AtomicBool::new(false); + let result = serve_unary( + machine, + (), + hooks, + Unary::new(|value| { + converted.store(true, Ordering::SeqCst); + Json(value) + }), + ) + .await; + assert_eq!(result.unwrap_err(), Error::Call(expected)); + assert!(!converted.load(Ordering::SeqCst)); + let events = observer.0.lock().unwrap(); + assert!(matches!(events.first(), Some(CallEvent::Started { .. }))); + assert!(matches!(events.last(), Some(CallEvent::Failed { .. }))); + assert_eq!(events.len(), if reject_hook { 2 } else { 3 }); +} + +#[rstest] +#[tokio::test] +async fn cancelling_unary_execution_releases_work_without_converting(hooks: Hooks) { + let observer = hooks.observer.clone(); + let released = Arc::new(AtomicBool::new(false)); + let release = Release(released.clone()); + let machine = hosted_call::("input", move |_, _, _| async move { + let _release = release; + std::future::pending().await + }); + let converted = AtomicBool::new(false); + let mut call = Box::pin(serve_unary( + machine, + (), + hooks, + Unary::new(|value| { + converted.store(true, Ordering::SeqCst); + Json(value) + }), + )); + assert!(futures_util::poll!(&mut call).is_pending()); + assert!(!released.load(Ordering::SeqCst)); + drop(call); + assert!(released.load(Ordering::SeqCst)); + assert!(!converted.load(Ordering::SeqCst)); + assert!(matches!( + observer.0.lock().unwrap().as_slice(), + [CallEvent::Started { .. }, CallEvent::Cancelled { .. },] + )); +} + +struct CustomUnaryProtocol; + +impl Protocol for CustomUnaryProtocol { + type Response = Bytes; + type Error = TestError; + type Request = &'static str; + type HostCall = Reply<&'static str>; + type Chunk = std::convert::Infallible; + type StreamHead = std::convert::Infallible; +} + +struct CustomUnaryAdapter { + reject_response: bool, + converted: Arc, +} + +impl ResponseEncoder for CustomUnaryAdapter { + type Protocol = CustomUnaryProtocol; + + fn encode_response(&self, response: Bytes) -> Result { + self.converted.store(true, Ordering::SeqCst); + if self.reject_response { + return Err(TestError::Adapter); + } + Ok(response.into_response()) + } +} + +struct Credentials(bool); + +impl litellm_host::services::HostCallHandler for Credentials { + async fn handle_host_call(&self, reply: Reply<&'static str>) -> Result<(), TestError> { + if self.0 { + return Err(TestError::Adapter); + } + reply.send("credential"); + Ok(()) + } +} + +#[rstest] +#[case::success(false, false)] +#[case::custom_operation_fails(true, false)] +#[case::response_conversion_fails(false, true)] +#[tokio::test] +async fn unary_custom_operations_and_conversion_finish_before_terminal_observation( + hooks: Hooks, + #[case] reject_op: bool, + #[case] reject_response: bool, +) { + let observer = hooks.observer.clone(); + let continued = Arc::new(AtomicBool::new(false)); + let executed = continued.clone(); + let released = Arc::new(AtomicBool::new(false)); + let release = Release(released.clone()); + let converted = Arc::new(AtomicBool::new(false)); + let machine = hosted_call::( + "request", + move |request, services, _| async move { + let _release = release; + let credential = services.call(|reply| reply).await?; + executed.store(true, Ordering::SeqCst); + Ok(CallOutput::Complete(Bytes::from(format!( + "{request}/{credential}" + )))) + }, + ); + let result = serve_unary( + machine, + Credentials(reject_op), + hooks, + CustomUnaryAdapter { + reject_response, + converted: converted.clone(), + }, + ) + .await; + assert_eq!(continued.load(Ordering::SeqCst), !reject_op); + assert_eq!(converted.load(Ordering::SeqCst), !reject_op); + assert!(released.load(Ordering::SeqCst)); + if reject_op || reject_response { + assert_eq!(result.unwrap_err(), Error::Call(TestError::Adapter)); + assert!(matches!( + observer.0.lock().unwrap().as_slice(), + [CallEvent::Started { .. }, CallEvent::Failed { .. }] + )); + } else { + let body = to_bytes(result.unwrap().into_body(), 1024).await.unwrap(); + assert_eq!(body, "request/credential"); + assert!(matches!( + observer.0.lock().unwrap().as_slice(), + [CallEvent::Started { .. }, CallEvent::Succeeded { .. }] + )); + } +} diff --git a/litellm-rust/crates/host-http/tests/sse.rs b/litellm-rust/crates/host-http/tests/sse.rs new file mode 100644 index 00000000000..d7bce0615ca --- /dev/null +++ b/litellm-rust/crates/host-http/tests/sse.rs @@ -0,0 +1,109 @@ +use std::{ + convert::Infallible, + sync::Arc, + sync::atomic::{AtomicUsize, Ordering}, +}; + +use axum::body::to_bytes; +use bytes::Bytes; +use futures_util::{StreamExt, stream}; +use http::{StatusCode, header::CONTENT_TYPE}; +use litellm_host::{ + call::{CallOutput, hosted_call}, + machine::MachineFault, + protocol::Protocol, +}; +use litellm_host_http::{Sse, serve}; +use rstest::rstest; + +#[derive(Clone, Debug)] +enum TestError { + Upstream, + Machine, +} + +impl From for TestError { + fn from(_: MachineFault) -> Self { + Self::Machine + } +} + +struct TestProtocol; + +impl Protocol for TestProtocol { + type Response = Bytes; + type Error = TestError; + type Request = (); + type HostCall = Infallible; + type Chunk = Bytes; + type StreamHead = (); +} + +#[rstest] +#[case::complete(false)] +#[case::failed(true)] +#[tokio::test] +async fn sse_preserves_encoded_chunks_and_uses_the_supplied_error_format(#[case] fail: bool) { + let first = Bytes::from_static(b"event: custom\ndata: first\n\n"); + let last = Bytes::from_static(b"data: [DONE]\n\n"); + let machine = hosted_call::((), move |(), _, _| async move { + let chunks = stream::iter([ + Ok(first), + if fail { + Err(TestError::Upstream) + } else { + Ok(last) + }, + ]) + .boxed(); + Ok(CallOutput::Stream { head: (), chunks }) + }); + let errors = Arc::new(AtomicUsize::new(0)); + let formatted_errors = errors.clone(); + let adapter = Sse::new(std::convert::identity, move |error| { + formatted_errors.fetch_add(1, Ordering::SeqCst); + Bytes::from(format!("event: custom_error\ndata: {error:?}\n\n")) + }); + let response = serve(machine, (), (), adapter).await.unwrap(); + assert_eq!(response.status(), StatusCode::OK); + assert_eq!(response.headers()[CONTENT_TYPE], "text/event-stream"); + assert_eq!(errors.load(Ordering::SeqCst), 0); + let body = to_bytes(response.into_body(), 1024).await.unwrap(); + assert_eq!( + body, + if fail { + "event: custom\ndata: first\n\nevent: custom_error\ndata: Call(Upstream)\n\n" + } else { + "event: custom\ndata: first\n\ndata: [DONE]\n\n" + } + ); + assert_eq!(errors.load(Ordering::SeqCst), usize::from(fail)); +} + +#[rstest] +#[tokio::test] +async fn completed_calls_use_the_response_converter_without_sse_headers() { + use axum::response::IntoResponse; + + let machine = hosted_call::((), |(), _, _| async { + Ok(CallOutput::Complete(Bytes::from_static(b"completed"))) + }); + let adapter = Sse::new( + |response| (StatusCode::CREATED, [("x-converted", "yes")], response).into_response(), + |_| panic!("a completed call cannot format a stream error"), + ); + let response = serve(machine, (), (), adapter).await.unwrap(); + assert_eq!(response.status(), StatusCode::CREATED); + assert_eq!(response.headers()["x-converted"], "yes"); + assert_ne!( + response + .headers() + .get(CONTENT_TYPE) + .map(|value| value.as_bytes()), + Some(b"text/event-stream".as_slice()) + ); + assert_eq!( + to_bytes(response.into_body(), 1024).await.unwrap(), + "completed" + ); +}