feat(rust): add the HTTP host driver (#43462)

Co-authored-by: Yujong Lee <yujong@berri.ai>
This commit is contained in:
devin-ai-integration[bot] 2026-09-27 16:18:10 -07:00 • committed by GitHub
parent 36784e3b79
commit 18933c8a21
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
11 changed files with 1224 additions and 0 deletions

View file

@ -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"

View file

@ -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" }

View file

@ -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

View file

@ -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

View file

@ -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<P> = MachineStep<P, HostedCompletion<<P as Protocol>::Response>>;
type Output<E> = CallOutput<Response, http::Response<()>, Bytes, E>;
pub async fn serve_unary<P, A, H, S>(
machine: HostedMachine<P>,
services: S,
hooks: H,
encoder: A,
) -> Result<Response, Error<P::Error>>
where
P: Protocol<Chunk = Infallible, StreamHead = Infallible>,
P::Error: From<MachineFault>,
H: RouteHooks<P::Error>,
S: HostCallHandler<P>,
A: ResponseEncoder<Protocol = P>,
{
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<P, A, H, S>(
machine: HostedMachine<P>,
services: S,
hooks: H,
encoder: A,
) -> Result<Response, Error<P::Error>>
where
P: Protocol,
P::Error: From<MachineFault>,
A: StreamEncoder<Protocol = P>,
H: RouteHooks<P::Error> + 'static,
S: HostCallHandler<P> + '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<P: Protocol, H, S> {
machine: HostedMachine<P>,
services: S,
hooks: H,
demand: Option<Reply<Demand>>,
}
impl<P, H, S> Driver<P, H, S>
where
P: Protocol,
P::Error: From<MachineFault>,
H: RouteHooks<P::Error>,
S: HostCallHandler<P>,
{
fn new(machine: HostedMachine<P>, services: S, hooks: H) -> Self {
Self {
machine,
services,
hooks,
demand: None,
}
}
async fn advance(&mut self) -> Result<StepOf<P>, Error<P::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<A>(mut self, encoder: Arc<A>) -> Result<Output<Error<P::Error>>, Error<P::Error>>
where
A: StreamEncoder<Protocol = P>,
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),
}
}
}

View file

@ -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: <Self::Protocol as Protocol>::Response,
) -> Result<Response, <Self::Protocol as Protocol>::Error>;
}
pub trait StreamEncoder: ResponseEncoder + 'static {
fn encode_stream_head(
&self,
head: <Self::Protocol as Protocol>::StreamHead,
) -> Result<http::Response<()>, <Self::Protocol as Protocol>::Error>;
fn encode_chunk(
&self,
chunk: <Self::Protocol as Protocol>::Chunk,
) -> Result<Bytes, <Self::Protocol as Protocol>::Error>;
fn encode_stream_error(&self, error: Error<<Self::Protocol as Protocol>::Error>) -> Bytes;
}
pub struct Unary<P, F> {
response: F,
protocol: PhantomData<fn() -> P>,
}
impl<P, F> Unary<P, F> {
pub fn new(response: F) -> Self {
Self {
response,
protocol: PhantomData,
}
}
}
impl<P, F, R> ResponseEncoder for Unary<P, F>
where
P: Protocol,
F: Fn(P::Response) -> R + Send + Sync,
R: IntoResponse,
{
type Protocol = P;
fn encode_response(&self, response: P::Response) -> Result<Response, P::Error> {
Ok((self.response)(response).into_response())
}
}

View file

@ -0,0 +1,7 @@
#[derive(Debug, PartialEq, thiserror::Error)]
pub enum Error<E> {
#[error(transparent)]
Call(E),
#[error("unexpected HTTP host operation")]
Protocol,
}

View file

@ -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;

View file

@ -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<P, C, F> {
response: Unary<P, C>,
stream_error: F,
}
impl<P, C, F> Sse<P, C, F> {
pub fn new(response: C, stream_error: F) -> Self {
Self {
response: Unary::new(response),
stream_error,
}
}
}
impl<P, C, F, R> ResponseEncoder for Sse<P, C, F>
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<Response, P::Error> {
self.response.encode_response(response)
}
}
impl<P, C, F, R> StreamEncoder for Sse<P, C, F>
where
P: Protocol<Chunk = Bytes>,
C: Fn(P::Response) -> R + Send + Sync + 'static,
F: Fn(Error<P::Error>) -> Bytes + Send + Sync + 'static,
R: IntoResponse,
{
fn encode_stream_head(&self, _: P::StreamHead) -> Result<http::Response<()>, 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<Bytes, P::Error> {
Ok(chunk)
}
fn encode_stream_error(&self, error: Error<P::Error>) -> Bytes {
(self.stream_error)(error)
}
}

View file

@ -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<MachineFault> 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<Response, TestError> {
if self.0 == Rejection::Complete {
return Err(TestError::Adapter);
}
Ok((StatusCode::CREATED, [("x-converted", "yes")], value).into_response())
}
}
impl litellm_host::services::HostCallHandler<TestProtocol> 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<http::Response<()>, 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<Bytes, TestError> {
if self.0 == Rejection::Chunk {
return Err(TestError::Adapter);
}
Ok(chunk)
}
fn encode_stream_error(&self, error: Error<TestError>) -> Bytes {
Bytes::from(format!("event: error\ndata: {error:?}\n\n"))
}
}
#[derive(Default)]
struct Observer(Mutex<Vec<CallEvent>>);
impl CallObserver for Observer {
fn observe(&self, event: CallEvent) {
self.0.lock().unwrap().push(event);
}
}
struct Hooks {
observer: Arc<Observer>,
reject: bool,
}
impl RouteHooks<TestError> for Hooks {
fn observer(&self) -> Option<Arc<dyn CallObserver>> {
Some(self.observer.clone())
}
async fn before_provider_request(
&self,
wire: WireRequest,
context: RequestContext,
) -> Result<WireRequest, TestError> {
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<Observer> {
Arc::new(Observer::default())
}
#[fixture]
fn hooks(observer: Arc<Observer>) -> 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::<TestProtocol, _, _>(
"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<AtomicBool>);
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<usize>,
) {
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::<TestProtocol, _, _>("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::<TestProtocol, _, _>("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::<TestProtocol, _, _>("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::<TestProtocol, _, _>("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<Observer>,
#[case] event: bool,
) {
let continued = Arc::new(AtomicBool::new(false));
let executed = continued.clone();
let machine = hosted_call::<TestProtocol, _, _>("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::<TestProtocol, HostedCompletion<Bytes>>::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::<UnaryProtocol, _, _>("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<Observer>,
#[case] reject_hook: bool,
#[case] expected: TestError,
) {
let machine = hosted_call::<UnaryProtocol, _, _>("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::<UnaryProtocol, _, _>("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<AtomicBool>,
}
impl ResponseEncoder for CustomUnaryAdapter {
type Protocol = CustomUnaryProtocol;
fn encode_response(&self, response: Bytes) -> Result<Response, TestError> {
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<CustomUnaryProtocol> 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::<CustomUnaryProtocol, _, _>(
"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 { .. }]
));
}
}

View file

@ -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<MachineFault> 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::<TestProtocol, _, _>((), 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::<TestProtocol, _, _>((), |(), _, _| 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"
);
}