use {
crate::Endpoint,
futures::{Stream, StreamExt, future::BoxFuture},
http_body_util::{BodyExt, Full, StreamBody, combinators::BoxBody},
hyper::{
HeaderMap,
body::{Bytes, Frame, Incoming},
},
hyper_util::rt::TokioIo,
std::{
future::Future,
marker::PhantomData,
net::SocketAddr,
pin::Pin,
task::{Context, Poll},
},
tokio::net::TcpListener,
};
pub type ResponseStream = Pin<Box<dyn Stream<Item = Bytes> + Send + Sync + 'static>>;
pub fn into_response_stream(stream: impl Stream<Item = Bytes> + Send + Sync + 'static) -> ResponseStream {
Box::pin(stream)
}
pub type MilBody = BoxBody<Bytes, std::convert::Infallible>;
pub struct IOTypeNotSend {
_marker: PhantomData<*const ()>,
stream: TokioIo<tokio::net::TcpStream>,
}
impl IOTypeNotSend {
pub fn new(stream: TokioIo<tokio::net::TcpStream>) -> Self { Self { _marker: PhantomData, stream } }
}
impl hyper::rt::Write for IOTypeNotSend {
fn poll_write(mut self: Pin<&mut Self>, cx: &mut Context<'_>, buf: &[u8]) -> Poll<Result<usize, std::io::Error>> {
Pin::new(&mut self.stream).poll_write(cx, buf)
}
fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<(), std::io::Error>> {
Pin::new(&mut self.stream).poll_flush(cx)
}
fn poll_shutdown(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<(), std::io::Error>> {
Pin::new(&mut self.stream).poll_shutdown(cx)
}
}
impl hyper::rt::Read for IOTypeNotSend {
fn poll_read(
mut self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: hyper::rt::ReadBufCursor<'_>,
) -> Poll<std::io::Result<()>> {
Pin::new(&mut self.stream).poll_read(cx, buf)
}
}
#[derive(Default)]
pub struct Body {
_marker: PhantomData<*const ()>,
data: Option<Bytes>,
}
impl From<String> for Body {
fn from(value: String) -> Self { Body { _marker: PhantomData, data: Some(value.into()) } }
}
impl<'a> From<&'a [u8]> for Body {
fn from(value: &'a [u8]) -> Self { Body { _marker: PhantomData, data: Some(Bytes::from_iter(value.iter().cloned())) } }
}
impl Body {
pub fn full(self) -> Full<Bytes> { Full::new(self.data.unwrap_or_default()) }
pub fn boxed(self) -> MilBody { self.full().map_err(|never| match never {}).boxed() }
}
impl hyper::body::Body for Body {
type Data = Bytes;
type Error = hyper::Error;
fn poll_frame(self: Pin<&mut Self>, _: &mut Context<'_>) -> Poll<Option<Result<Frame<Self::Data>, Self::Error>>> {
Poll::Ready(self.get_mut().data.take().map(|d| Ok(Frame::data(d))))
}
}
pub type AsyncHandler<I, O> = Box<dyn Fn(I) -> BoxFuture<'static, O> + Send + 'static>;
pub type AsyncHandler3<I, I2, I3, O> = Box<dyn Fn(I, I2, I3) -> BoxFuture<'static, O> + Send + 'static>;
#[allow(clippy::type_complexity)]
pub trait ServerEndpoint<C>: Endpoint<C> {
fn auth() -> AsyncHandler<HeaderMap, Result<C, anyhow::Error>>;
fn handler() -> AsyncHandler3<C, HeaderMap, <Self as Endpoint<C>>::Data, anyhow::Result<<Self as Endpoint<C>>::Returns>>;
fn is_raw() -> bool { false }
fn stream_handler() -> Option<AsyncHandler3<C, HeaderMap, <Self as Endpoint<C>>::Data, anyhow::Result<ResponseStream>>> {
None
}
}
pub trait TypedEndpoint: Endpoint<Self::Client> {
type Client: Send;
}
pub trait ClientEndpoint<C>: Endpoint<C> {
fn decode_response(bytes: Bytes) -> anyhow::Result<<Self as Endpoint<C>>::Returns>;
}
pub fn serve_local<RouteFut>(addr: SocketAddr, route: fn(hyper::Request<Incoming>) -> RouteFut) -> anyhow::Result<()>
where
RouteFut:
Future<Output = Result<hyper::Response<MilBody>, std::convert::Infallible>> + 'static,
{
use {
hyper::server::conn::http1,
hyper::service::service_fn,
};
let rt = tokio::runtime::Builder::new_current_thread().enable_all().build()?;
let ls = tokio::task::LocalSet::new();
let listener = ls.block_on(&rt, TcpListener::bind(addr))?;
tracing::info!("Listening on http://{}", addr);
loop {
let (stream, _) = ls.block_on(&rt, listener.accept())?;
let io = IOTypeNotSend::new(TokioIo::new(stream));
let service = service_fn(route);
ls.spawn_local(async move {
if let Err(err) = http1::Builder::new().serve_connection(io, service).await {
tracing::warn!("Error serving connection: {:?}", err);
}
});
}
}
pub async fn serve<RouteFut>(addr: SocketAddr, route: fn(hyper::Request<Incoming>) -> RouteFut) -> anyhow::Result<()>
where
RouteFut:
Future<Output = Result<hyper::Response<MilBody>, std::convert::Infallible>>
+ Send
+ 'static,
{
use {
hyper::server::conn::http1,
hyper::service::service_fn,
};
let listener = TcpListener::bind(addr).await?;
tracing::info!("Listening on http://{}", addr);
loop {
let (stream, _) = listener.accept().await?;
let io = TokioIo::new(stream);
let service = service_fn(route);
tokio::spawn(async move {
if let Err(err) = http1::Builder::new().serve_connection(io, service).await {
tracing::warn!("Error serving connection: {:?}", err);
}
});
}
}
pub fn stream_to_body(stream: ResponseStream) -> MilBody {
BodyExt::boxed(StreamBody::new(stream.map(|chunk| Ok(Frame::data(chunk)))))
}
pub use std::fmt as _fmt;