use std::future::Future;
use std::net::SocketAddr;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::Arc;
use arc_swap::ArcSwap;
use axum::Router;
use boatramp_http::{
Body as HttpBody, BodyError, Handler, Request as HttpRequest, Response as HttpResponse,
};
use futures::StreamExt as _;
use http_body_util::BodyStream;
use rustls::ServerConfig;
use tokio::net::{TcpListener, TcpStream};
use tokio_rustls::TlsAcceptor;
const ACME_TLS_ALPN: &[u8] = b"acme-tls/1";
pub fn alpn_h1_h2() -> Vec<Vec<u8>> {
vec![b"h2".to_vec(), b"http/1.1".to_vec()]
}
pub struct RouterHandler {
router: Router,
peer: SocketAddr,
}
impl RouterHandler {
pub fn new(router: Router, peer: SocketAddr) -> Self {
Self { router, peer }
}
}
impl Handler for RouterHandler {
async fn handle(&self, req: HttpRequest) -> HttpResponse {
let mut request = req.map(axum::body::Body::new);
request
.extensions_mut()
.insert(axum::extract::ConnectInfo(self.peer));
use tower_service::Service as _;
let mut router = self.router.clone();
let resp = match router.call(request).await {
Ok(r) => r,
Err(_) => return boatramp_http::response(502, b"bad gateway".to_vec()),
};
let (parts, body) = resp.into_parts();
let chunks = BodyStream::new(body).filter_map(|frame| {
std::future::ready(match frame {
Ok(f) => f.into_data().ok().filter(|b| !b.is_empty()).map(Ok),
Err(_) => Some(Err(BodyError)),
})
});
axum::http::Response::from_parts(parts, HttpBody::try_stream(chunks))
}
}
pub async fn serve_router_conn<IO>(io: IO, peer: SocketAddr, router: Router)
where
IO: tokio::io::AsyncRead + tokio::io::AsyncWrite + Unpin + Send + 'static,
{
if let Err(err) = boatramp_http::serve_connection(io, RouterHandler::new(router, peer)).await {
tracing::debug!(%peer, %err, "connection served with error");
}
}
async fn serve_tls_stream(
stream: tokio_rustls::server::TlsStream<TcpStream>,
peer: SocketAddr,
router: Router,
) {
let alpn = stream.get_ref().1.alpn_protocol().map(<[u8]>::to_vec);
let handler = RouterHandler::new(router, peer);
let result = match alpn.as_deref() {
Some(p) if p == ACME_TLS_ALPN => {
tracing::debug!(%peer, "completed an ACME tls-alpn-01 challenge");
return;
}
Some(b"h2") => boatramp_http::h2::serve_connection_mux(stream, handler).await,
Some(b"http/1.1") => boatramp_http::h1::serve_connection(stream, handler).await,
_ => boatramp_http::serve_connection(stream, handler).await,
};
if let Err(err) = result {
tracing::debug!(%peer, %err, "TLS connection served with error");
}
}
#[derive(Clone)]
pub struct ReloadableTls(Arc<ArcSwap<ServerConfig>>);
impl ReloadableTls {
pub fn new(config: ServerConfig) -> Self {
Self(Arc::new(ArcSwap::from_pointee(config)))
}
pub fn reload(&self, config: ServerConfig) {
self.0.store(Arc::new(config));
}
fn current(&self) -> Arc<ServerConfig> {
self.0.load_full()
}
}
impl From<ServerConfig> for ReloadableTls {
fn from(config: ServerConfig) -> Self {
Self::new(config)
}
}
const DRAIN_DEADLINE: std::time::Duration = std::time::Duration::from_secs(30);
pub async fn serve_tls<S>(
addr: SocketAddr,
tls: ReloadableTls,
router: Router,
shutdown: S,
) -> std::io::Result<()>
where
S: Future<Output = ()> + Send,
{
let listener = TcpListener::bind(addr).await?;
serve_tls_listener(listener, tls, router, shutdown).await
}
pub async fn serve_tls_listener<S>(
listener: TcpListener,
tls: ReloadableTls,
router: Router,
shutdown: S,
) -> std::io::Result<()>
where
S: Future<Output = ()> + Send,
{
if let Ok(addr) = listener.local_addr() {
tracing::info!(%addr, "serving HTTPS (boatramp-http)");
}
let inflight = Arc::new(AtomicUsize::new(0));
tokio::pin!(shutdown);
loop {
tokio::select! {
_ = &mut shutdown => break,
accepted = listener.accept() => {
let (mut tcp, peer) = match accepted {
Ok(v) => v,
Err(err) => {
tracing::debug!(%err, "TLS serve: accept error");
continue;
}
};
crate::disable_nagle(&mut tcp);
let acceptor = TlsAcceptor::from(tls.current());
let router = router.clone();
let inflight = inflight.clone();
inflight.fetch_add(1, Ordering::SeqCst);
tokio::spawn(async move {
match acceptor.accept(tcp).await {
Ok(stream) => serve_tls_stream(stream, peer, router).await,
Err(err) => tracing::debug!(%peer, %err, "TLS handshake failed"),
}
inflight.fetch_sub(1, Ordering::SeqCst);
});
}
}
}
drain(&inflight).await;
Ok(())
}
pub async fn serve_plaintext<S>(
addr: SocketAddr,
router: Router,
shutdown: S,
) -> std::io::Result<()>
where
S: Future<Output = ()> + Send,
{
let listener = TcpListener::bind(addr).await?;
serve_plaintext_listener(listener, router, shutdown).await
}
pub async fn serve_plaintext_listener<S>(
listener: TcpListener,
router: Router,
shutdown: S,
) -> std::io::Result<()>
where
S: Future<Output = ()> + Send,
{
let inflight = Arc::new(AtomicUsize::new(0));
tokio::pin!(shutdown);
loop {
tokio::select! {
_ = &mut shutdown => break,
accepted = listener.accept() => {
let (mut tcp, peer) = match accepted {
Ok(v) => v,
Err(err) => {
tracing::debug!(%err, "plaintext serve: accept error");
continue;
}
};
crate::disable_nagle(&mut tcp);
let router = router.clone();
let inflight = inflight.clone();
inflight.fetch_add(1, Ordering::SeqCst);
tokio::spawn(async move {
serve_router_conn(tcp, peer, router).await;
inflight.fetch_sub(1, Ordering::SeqCst);
});
}
}
}
drain(&inflight).await;
Ok(())
}
async fn drain(inflight: &AtomicUsize) {
let deadline = tokio::time::Instant::now() + DRAIN_DEADLINE;
while inflight.load(Ordering::SeqCst) > 0 {
if tokio::time::Instant::now() >= deadline {
tracing::warn!("TLS/plaintext drain deadline exceeded; dropping in-flight connections");
break;
}
tokio::time::sleep(std::time::Duration::from_millis(50)).await;
}
}