use std::{
convert::Infallible,
future::Future,
pin::{Pin, pin},
sync::Arc,
time::Duration,
};
use hyper::{body::Incoming, service::Service};
use hyper_util::{
rt::{TokioExecutor, TokioIo},
server::conn::auto,
};
use tokio::sync::watch;
use crate::{Body, Listener, Router, request::Request, response::Response};
#[derive(Clone)]
pub struct RouterService {
router: Arc<Router>,
pub(crate) shutdown_timeout: Duration,
}
impl RouterService {
#[must_use]
pub fn new(router: Router) -> Self {
const DEFAULT_SHUTDOWN_TIMEOUT: Duration = Duration::from_secs(30);
Self {
router: Arc::new(router),
shutdown_timeout: DEFAULT_SHUTDOWN_TIMEOUT,
}
}
#[must_use]
pub fn shutdown_timeout(mut self, timeout: Duration) -> Self {
self.shutdown_timeout = timeout;
self
}
}
impl From<Router> for RouterService {
fn from(router: Router) -> Self {
Self::new(router)
}
}
impl Service<Request<Incoming>> for RouterService {
type Response = Response;
type Error = Infallible;
type Future = Pin<Box<dyn Future<Output = Result<Self::Response, Self::Error>> + Send>>;
fn call(&self, request: Request<Incoming>) -> Self::Future {
let router = self.router.clone();
Box::pin(async move { Ok(router.handle(request.map(Body::new)).await) })
}
}
pub async fn internal_serve(
mut listener: impl Listener,
service: RouterService,
shutdown: impl Future<Output = ()>,
) -> std::io::Result<()> {
let (drain_tx, drain_rx) = watch::channel(());
let (cutoff_tx, cutoff_rx) = watch::channel(());
let (done_tx, done_rx) = watch::channel(());
let mut shutdown = pin!(shutdown);
loop {
let accepted = tokio::select! {
accepted = listener.accept() => accepted,
() = &mut shutdown => break,
};
let (stream, _remote) = accepted?;
let io = TokioIo::new(stream);
let service = service.clone();
let mut drain_rx = drain_rx.clone();
let mut cutoff_rx = cutoff_rx.clone();
let done_rx = done_rx.clone();
tokio::spawn(async move {
let _done_rx = done_rx;
let builder = auto::Builder::new(TokioExecutor::new());
let mut connection = pin!(builder.serve_connection_with_upgrades(io, service));
let result = tokio::select! {
result = connection.as_mut() => result,
_ = drain_rx.changed() => {
connection.as_mut().graceful_shutdown();
tokio::select! {
result = connection.as_mut() => result,
_ = cutoff_rx.changed() => return,
}
}
};
if let Err(_error) = result {
}
});
}
drop(listener);
drop(drain_rx);
drop(drain_tx);
drop(done_rx);
tokio::select! {
() = done_tx.closed() => {}
() = tokio::time::sleep(service.shutdown_timeout) => {}
}
drop(cutoff_rx);
drop(cutoff_tx);
done_tx.closed().await;
Ok(())
}
#[cfg(test)]
mod tests {
use std::{
borrow::Cow,
convert::Infallible,
net::SocketAddr,
pin::Pin,
task::{Context, Poll},
time::Duration,
};
use http_body::Frame;
use tokio::{
io::{AsyncReadExt, AsyncWriteExt},
net::{TcpListener, TcpStream},
sync::oneshot,
task::JoinHandle,
};
use topcoat_core::context::Cx;
use super::*;
use crate::{
Body, Method, Path, RouteFn, RouteFuture, RouteHandlerFn, Router,
request::Bytes,
response::{IntoResponse, Response},
};
fn router_with(handler: RouteHandlerFn) -> Router {
Router::builder()
.route(RouteFn::new(
Method::GET,
Cow::Borrowed(Path::new("/x")),
handler,
))
.build()
}
fn say_route(cx: &Cx, _body: Body) -> RouteFuture<'_> {
Box::pin(async move { "served".into_response(cx) })
}
fn panic_route(_cx: &Cx, _body: Body) -> RouteFuture<'_> {
Box::pin(async move { panic!("request handler panicked") })
}
struct PanickingBody;
impl http_body::Body for PanickingBody {
type Data = Bytes;
type Error = Infallible;
fn poll_frame(
self: Pin<&mut Self>,
_cx: &mut Context<'_>,
) -> Poll<Option<Result<Frame<Self::Data>, Self::Error>>> {
panic!("response body panicked");
}
}
fn panicking_body_route(_cx: &Cx, _body: Body) -> RouteFuture<'_> {
Box::pin(async move { Ok(Response::new(Body::new(PanickingBody))) })
}
fn slow_route(cx: &Cx, _body: Body) -> RouteFuture<'_> {
Box::pin(async move {
tokio::time::sleep(Duration::from_millis(200)).await;
"slow".into_response(cx)
})
}
fn hang_route(_cx: &Cx, _body: Body) -> RouteFuture<'_> {
Box::pin(std::future::pending())
}
async fn spawn_server(
service: RouterService,
) -> (
SocketAddr,
oneshot::Sender<()>,
JoinHandle<std::io::Result<()>>,
) {
let listener = TcpListener::bind(("127.0.0.1", 0)).await.unwrap();
let addr = listener.local_addr().unwrap();
let (shutdown_tx, shutdown_rx) = oneshot::channel();
let server = tokio::spawn(internal_serve(listener, service, async {
let _ = shutdown_rx.await;
}));
(addr, shutdown_tx, server)
}
async fn shut_down(server: JoinHandle<std::io::Result<()>>) {
tokio::time::timeout(Duration::from_secs(5), server)
.await
.expect("server did not shut down within the grace period")
.unwrap()
.unwrap();
}
async fn get(addr: SocketAddr, path: &str) -> String {
let mut stream = TcpStream::connect(addr).await.unwrap();
stream
.write_all(
format!("GET {path} HTTP/1.1\r\nhost: test\r\nconnection: close\r\n\r\n")
.as_bytes(),
)
.await
.unwrap();
let mut response = String::new();
stream.read_to_string(&mut response).await.unwrap();
response
}
#[tokio::test]
async fn returns_once_the_shutdown_signal_fires() {
let service = RouterService::new(router_with(say_route));
let (addr, shutdown_tx, server) = spawn_server(service).await;
let response = get(addr, "/x").await;
assert!(response.contains("200 OK"));
assert!(response.ends_with("served"));
shutdown_tx.send(()).unwrap();
shut_down(server).await;
assert!(TcpStream::connect(addr).await.is_err());
}
#[tokio::test]
async fn handler_panic_returns_500_and_server_keeps_running() {
let router = Router::builder()
.route(RouteFn::new(
Method::GET,
Cow::Borrowed(Path::new("/panic")),
panic_route,
))
.route(RouteFn::new(
Method::GET,
Cow::Borrowed(Path::new("/x")),
say_route,
))
.build();
let (addr, shutdown_tx, server) = spawn_server(RouterService::new(router)).await;
let response = get(addr, "/panic").await;
assert!(response.contains("500 Internal Server Error"));
assert!(response.ends_with("internal server error"));
let response = get(addr, "/x").await;
assert!(response.contains("200 OK"));
assert!(response.ends_with("served"));
shutdown_tx.send(()).unwrap();
shut_down(server).await;
}
#[tokio::test]
async fn response_body_panic_does_not_stop_server() {
let router = Router::builder()
.route(RouteFn::new(
Method::GET,
Cow::Borrowed(Path::new("/body-panic")),
panicking_body_route,
))
.route(RouteFn::new(
Method::GET,
Cow::Borrowed(Path::new("/x")),
say_route,
))
.build();
let (addr, shutdown_tx, server) = spawn_server(RouterService::new(router)).await;
let mut stream = TcpStream::connect(addr).await.unwrap();
stream
.write_all(b"GET /body-panic HTTP/1.1\r\nhost: test\r\nconnection: close\r\n\r\n")
.await
.unwrap();
let mut response = String::new();
let _ = stream.read_to_string(&mut response).await;
let response = get(addr, "/x").await;
assert!(response.contains("200 OK"));
assert!(response.ends_with("served"));
shutdown_tx.send(()).unwrap();
shut_down(server).await;
}
#[cfg(unix)]
#[tokio::test]
async fn serves_over_a_unix_socket() {
use tokio::net::{UnixListener, UnixStream};
let service = RouterService::new(router_with(say_route));
let path = std::env::temp_dir().join(format!("topcoat-serve-{}.sock", std::process::id()));
let _ = std::fs::remove_file(&path);
let listener = UnixListener::bind(&path).unwrap();
let (shutdown_tx, shutdown_rx) = oneshot::channel();
let server = tokio::spawn(internal_serve(listener, service, async {
let _ = shutdown_rx.await;
}));
let mut stream = UnixStream::connect(&path).await.unwrap();
stream
.write_all(b"GET /x HTTP/1.1\r\nhost: test\r\nconnection: close\r\n\r\n")
.await
.unwrap();
let mut response = String::new();
stream.read_to_string(&mut response).await.unwrap();
assert!(response.contains("200 OK"));
assert!(response.ends_with("served"));
shutdown_tx.send(()).unwrap();
shut_down(server).await;
let _ = std::fs::remove_file(&path);
}
#[tokio::test]
async fn drains_the_in_flight_request_before_returning() {
let service = RouterService::new(router_with(slow_route));
let (addr, shutdown_tx, server) = spawn_server(service).await;
let mut stream = TcpStream::connect(addr).await.unwrap();
stream
.write_all(b"GET /x HTTP/1.1\r\nhost: test\r\n\r\n")
.await
.unwrap();
tokio::time::sleep(Duration::from_millis(50)).await;
shutdown_tx.send(()).unwrap();
let mut response = String::new();
stream.read_to_string(&mut response).await.unwrap();
assert!(response.contains("200 OK"));
assert!(response.ends_with("slow"));
shut_down(server).await;
}
#[tokio::test]
async fn cuts_hung_connections_at_the_shutdown_timeout() {
let service = RouterService::new(router_with(hang_route))
.shutdown_timeout(Duration::from_millis(100));
let (addr, shutdown_tx, server) = spawn_server(service).await;
let mut stream = TcpStream::connect(addr).await.unwrap();
stream
.write_all(b"GET /x HTTP/1.1\r\nhost: test\r\n\r\n")
.await
.unwrap();
tokio::time::sleep(Duration::from_millis(50)).await;
shutdown_tx.send(()).unwrap();
shut_down(server).await;
let mut response = String::new();
let result = stream.read_to_string(&mut response).await;
assert!(result.is_err() || response.is_empty());
}
}