use std::borrow::Cow;
use std::panic::AssertUnwindSafe;
use std::pin::Pin;
use std::task::{Context, Poll};
use std::{convert::Infallible, net::SocketAddr, sync::Arc};
use crate::{
ConnectInfo, DefaultErrorHook, Error, ErrorHook, Handler, ObservedRoute, OuterWrapState,
RawPathExt, RequestHook, Wrap, WrapTarget,
};
use crate::{IntoResponse, Result, WrapState};
use async_trait::async_trait;
use axol_http::body::{BodyComponent, BodyWrapper};
use axol_http::header::HeaderMapConvertError;
use axol_http::{Body, StatusCode};
use axol_http::{request::Request, response::Response};
use futures::{FutureExt, Stream};
use http_body::Body as HttpBody;
use hyper::body::Incoming;
use hyper::{Request as HyperRequest, Response as HyperResponse};
use hyper_util::rt::{TokioExecutor, TokioIo};
use hyper_util::server::conn::auto;
use log::error;
use pin_project_lite::pin_project;
use tokio::io::{AsyncRead, AsyncWrite};
use tokio::net::{TcpListener, TcpStream};
#[cfg(feature = "trace")]
use tracing::Instrument;
use crate::Router;
#[cfg(feature = "tls")]
mod tls_acceptor;
#[cfg(feature = "tls")]
pub use tls_acceptor::*;
#[async_trait]
pub trait Acceptor: Send + 'static {
type Conn: AsyncRead + AsyncWrite + Unpin + Send + 'static;
async fn accept(&mut self) -> std::io::Result<(Self::Conn, SocketAddr)>;
}
pub struct TcpAcceptor {
listener: TcpListener,
nodelay: bool,
}
impl TcpAcceptor {
pub async fn bind(addr: SocketAddr) -> std::io::Result<Self> {
Ok(Self {
listener: TcpListener::bind(addr).await?,
nodelay: false,
})
}
pub fn from_listener(listener: TcpListener) -> Self {
Self {
listener,
nodelay: false,
}
}
pub fn set_nodelay(&mut self, nodelay: bool) {
self.nodelay = nodelay;
}
pub fn local_addr(&self) -> std::io::Result<SocketAddr> {
self.listener.local_addr()
}
}
#[async_trait]
impl Acceptor for TcpAcceptor {
type Conn = TcpStream;
async fn accept(&mut self) -> std::io::Result<(Self::Conn, SocketAddr)> {
let (stream, addr) = self.listener.accept().await?;
if self.nodelay {
stream.set_nodelay(true)?;
}
Ok((stream, addr))
}
}
pub struct Server<A> {
acceptor: A,
router: Router,
}
impl Server<TcpAcceptor> {
pub async fn bind(addr: SocketAddr, router: Router) -> std::io::Result<Self> {
Ok(Self {
acceptor: TcpAcceptor::bind(addr).await?,
router,
})
}
}
impl<A: Acceptor> Server<A> {
pub fn new(acceptor: A, router: Router) -> Self {
Self { acceptor, router }
}
pub fn local_addr(&self) -> std::io::Result<SocketAddr>
where
A: LocalAddr,
{
self.acceptor.local_addr()
}
}
pub trait LocalAddr {
fn local_addr(&self) -> std::io::Result<SocketAddr>;
}
impl LocalAddr for TcpAcceptor {
fn local_addr(&self) -> std::io::Result<SocketAddr> {
TcpAcceptor::local_addr(self)
}
}
impl<A: Acceptor + LocalAddr> LocalAddr for Server<A> {
fn local_addr(&self) -> std::io::Result<SocketAddr> {
self.acceptor.local_addr()
}
}
pin_project! {
struct BodyInputStream {
#[pin]
body: Incoming,
ended: bool,
}
}
impl Stream for BodyInputStream {
type Item = std::result::Result<BodyComponent, anyhow::Error>;
fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
let this = self.project();
if *this.ended {
return Poll::Ready(None);
}
match this.body.poll_frame(cx) {
Poll::Pending => Poll::Pending,
Poll::Ready(None) => {
*this.ended = true;
Poll::Ready(None)
}
Poll::Ready(Some(Err(e))) => {
*this.ended = true;
Poll::Ready(Some(Err(e.into())))
}
Poll::Ready(Some(Ok(frame))) => {
let frame = match frame.into_data() {
Ok(data) => return Poll::Ready(Some(Ok(BodyComponent::Data(data)))),
Err(frame) => frame,
};
match frame.into_trailers() {
Ok(trailers) => Poll::Ready(Some(
trailers
.try_into()
.map_err(|e: HeaderMapConvertError| e.into())
.map(BodyComponent::Trailers),
)),
Err(_) => {
cx.waker().wake_by_ref();
Poll::Pending
}
}
}
}
}
}
#[async_recursion::async_recursion]
pub(crate) async fn inner_handler(
request_hooks: Vec<Arc<dyn RequestHook>>,
wraps: Vec<Arc<dyn Wrap>>,
handler: Arc<dyn Handler>,
request: &mut Request,
) -> Result<Response> {
for middleware in request_hooks {
match middleware.handle_request(&mut *request).await {
Ok(Some(x)) => return Ok(x),
Err(Error::SkipMiddleware) | Ok(None) => (),
Err(e) => return Err(e),
}
}
let state = WrapState {
wraps,
target: WrapTarget::Handler(&*handler),
request,
};
state.next().await
}
async fn request_phase(
request_hooks: Vec<Arc<dyn RequestHook>>,
wraps: Vec<Arc<dyn Wrap>>,
outer_wraps: Vec<Arc<dyn Wrap>>,
handler: Arc<dyn Handler>,
request: &mut Request,
) -> Result<Response> {
let outer_wrap_state = OuterWrapState {
request_hooks,
wraps,
handler,
};
let state = WrapState {
wraps: outer_wraps,
target: WrapTarget::Phase(outer_wrap_state),
request,
};
state.next().await
}
async fn handle_error(
observed: &ObservedRoute<'_>,
request: &mut Request,
mut error: Error,
) -> Response {
for middleware in &observed.error_hooks {
match middleware.handle_error(request.parts(), &mut error).await {
Ok(Some(x)) => return x,
Err(Error::SkipMiddleware) | Ok(None) => (),
Err(e) => {
log::error!("error hook middleware failure: {e}");
}
}
}
DefaultErrorHook
.handle_error(request.parts(), &mut error)
.await
.unwrap()
.unwrap()
}
async fn handle_early_response(
observed: &ObservedRoute<'_>,
request: &mut Request,
mut response: Response,
) -> Response {
for middleware in &observed.early_response_hooks {
match middleware
.handle_response(request.parts(), &mut response)
.await
{
Err(Error::SkipMiddleware) | Ok(()) => (),
Err(error) => {
return handle_error(observed, request, error).await;
}
}
}
response
}
async fn handle_late_response(
observed: &ObservedRoute<'_>,
request: &mut Request,
response: &mut Response,
) {
for middleware in &observed.late_response_hooks {
middleware.handle_response(request.parts(), response).await;
}
}
async fn do_handle_axol_response(
router: Arc<Router>,
address: SocketAddr,
request: HyperRequest<Incoming>,
) -> Result<Response> {
#[cfg_attr(not(feature = "ws"), allow(unused_mut))]
let (mut parts, body) = request.into_parts();
let size_hint = body.size_hint().lower() as usize;
let extensions = axol_http::Extensions::new();
#[cfg(feature = "ws")]
extensions.take_from_http::<hyper::upgrade::OnUpgrade>(&mut parts.extensions);
let mut request = Request {
method: parts
.method
.try_into()
.map_err(|_| Error::Status(axol_http::StatusCode::MethodNotAllowed))?,
uri: parts.uri,
version: parts.version,
headers: parts
.headers
.try_into()
.map_err(|e: HeaderMapConvertError| Error::unprocessable_entity(e.to_string()))?,
extensions,
body: Body::Stream {
size_hint: Some(size_hint),
stream: Box::pin(BodyInputStream { body, ended: false }),
},
};
let mut observed = router.resolve_path(request.method, request.uri.path());
for (_, value) in observed.variables.0.iter_mut() {
let decoded = percent_encoding::percent_decode_str(value)
.decode_utf8()
.map_err(|_| Error::BadUtf8)?;
if let Cow::Owned(decoded) = decoded {
*value = decoded;
}
}
request.extensions.extend(&observed.extensions);
request
.extensions
.insert(RawPathExt(std::mem::take(&mut observed.variables.0)));
request.extensions.insert(ConnectInfo(address));
#[cfg(feature = "trace")]
let remote = address;
#[cfg(feature = "trace")]
let span = tracing::trace_span!("axol_http", %remote, uri = %request.uri);
let wraps = std::mem::take(&mut observed.wraps);
let outer_wraps = std::mem::take(&mut observed.outer_wraps);
let request_hooks = std::mem::take(&mut observed.request_hooks);
let late_response = AssertUnwindSafe(async move {
let mut late_response = match request_phase(
request_hooks,
wraps,
outer_wraps,
observed.route.clone(),
&mut request,
)
.await
{
Ok(x) => handle_early_response(&observed, &mut request, x).await,
Err(error) => handle_error(&observed, &mut request, error).await,
};
handle_late_response(&observed, &mut request, &mut late_response).await;
late_response
})
.catch_unwind();
#[cfg(feature = "trace")]
let late_response = late_response.instrument(span.clone()).await;
#[cfg(not(feature = "trace"))]
let late_response = late_response.await;
#[cfg(feature = "trace")]
let _span = span.enter();
let late_response = match late_response {
Ok(x) => x,
Err(e) => {
let display = e
.downcast::<String>()
.map(|x| *x)
.or_else(|e| e.downcast::<&'static str>().map(|x| x.to_string()))
.unwrap_or_else(|e| format!("{e:?}"));
error!("panic during handler/middlware: {display}");
StatusCode::InternalServerError.into_response().unwrap()
}
};
Ok(late_response)
}
async fn do_handle(
router: Arc<Router>,
address: SocketAddr,
request: HyperRequest<Incoming>,
) -> std::result::Result<HyperResponse<BodyWrapper>, Infallible> {
let is_head = request.method() == http::Method::HEAD;
let mut response = match do_handle_axol_response(router, address, request).await {
Ok(x) => x,
Err(e) => e.into_response(),
};
if is_head {
std::mem::take(&mut response.body);
}
let status: http::StatusCode = response.status.into();
let mut builder = HyperResponse::builder()
.status(status)
.version(response.version);
*builder.headers_mut().unwrap() = response.headers.into();
*builder.extensions_mut().unwrap() = Default::default();
Ok(builder
.body(response.body.into())
.expect("body conversion failed"))
}
impl<A: Acceptor> Server<A> {
pub async fn serve(self) -> std::io::Result<()> {
self.serve_custom(|_| ()).await
}
pub async fn serve_custom(
mut self,
customize: impl FnOnce(&mut auto::Builder<TokioExecutor>),
) -> std::io::Result<()> {
self.router.set_paths("");
let router = Arc::new(self.router);
let mut builder = auto::Builder::new(TokioExecutor::new());
customize(&mut builder);
let builder = Arc::new(builder);
loop {
let (conn, address) = match self.acceptor.accept().await {
Ok(x) => x,
Err(e) => {
error!("error accepting connection: {e}");
continue;
}
};
let router = router.clone();
let builder = builder.clone();
tokio::spawn(async move {
let service =
hyper::service::service_fn(move |req| do_handle(router.clone(), address, req));
if let Err(e) = builder
.serve_connection_with_upgrades(TokioIo::new(conn), service)
.await
{
log::debug!("error serving connection from {address}: {e}");
}
});
}
}
}