mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-30 01:52:18 +00:00
feat(rust): add the HTTP host driver (#43462)
Co-authored-by: Yujong Lee <yujong@berri.ai>
This commit is contained in:
parent
36784e3b79
commit
18933c8a21
11 changed files with 1224 additions and 0 deletions
15
litellm-rust/Cargo.lock
generated
15
litellm-rust/Cargo.lock
generated
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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" }
|
||||
|
|
|
|||
7
litellm-rust/crates/host-http/AGENTS.md
Normal file
7
litellm-rust/crates/host-http/AGENTS.md
Normal 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
|
||||
20
litellm-rust/crates/host-http/Cargo.toml
Normal file
20
litellm-rust/crates/host-http/Cargo.toml
Normal 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
|
||||
165
litellm-rust/crates/host-http/src/driver.rs
Normal file
165
litellm-rust/crates/host-http/src/driver.rs
Normal 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),
|
||||
}
|
||||
}
|
||||
}
|
||||
57
litellm-rust/crates/host-http/src/encoding.rs
Normal file
57
litellm-rust/crates/host-http/src/encoding.rs
Normal 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())
|
||||
}
|
||||
}
|
||||
7
litellm-rust/crates/host-http/src/error.rs
Normal file
7
litellm-rust/crates/host-http/src/error.rs
Normal file
|
|
@ -0,0 +1,7 @@
|
|||
#[derive(Debug, PartialEq, thiserror::Error)]
|
||||
pub enum Error<E> {
|
||||
#[error(transparent)]
|
||||
Call(E),
|
||||
#[error("unexpected HTTP host operation")]
|
||||
Protocol,
|
||||
}
|
||||
9
litellm-rust/crates/host-http/src/lib.rs
Normal file
9
litellm-rust/crates/host-http/src/lib.rs
Normal 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;
|
||||
58
litellm-rust/crates/host-http/src/sse.rs
Normal file
58
litellm-rust/crates/host-http/src/sse.rs
Normal 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)
|
||||
}
|
||||
}
|
||||
776
litellm-rust/crates/host-http/tests/serve.rs
Normal file
776
litellm-rust/crates/host-http/tests/serve.rs
Normal 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 { .. }]
|
||||
));
|
||||
}
|
||||
}
|
||||
109
litellm-rust/crates/host-http/tests/sse.rs
Normal file
109
litellm-rust/crates/host-http/tests/sse.rs
Normal 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"
|
||||
);
|
||||
}
|
||||
Loading…
Add table
Reference in a new issue