use std::convert::Infallible;
use std::net::SocketAddr;
use axum::Router;
use tokio::net::TcpListener;
use tower::Service;
pub async fn serve(router: Router, addr: &str) -> std::io::Result<()> {
let listener = TcpListener::bind(addr).await?;
axum::serve(listener, router).await?;
Ok(())
}
pub async fn serve_with_graceful_shutdown(router: Router, addr: &str) -> std::io::Result<()> {
let listener = TcpListener::bind(addr).await?;
axum::serve(listener, router)
.with_graceful_shutdown(shutdown_signal())
.await?;
Ok(())
}
pub async fn serve_with_listener(router: Router, listener: TcpListener) -> std::io::Result<()> {
axum::serve(listener, router).await?;
Ok(())
}
pub async fn build_tcp_listener(addr: &str) -> std::io::Result<(TcpListener, SocketAddr)> {
let listener = TcpListener::bind(addr).await?;
let local_addr = listener.local_addr()?;
Ok((listener, local_addr))
}
async fn shutdown_signal() {
let ctrl_c = async {
tokio::signal::ctrl_c()
.await
.expect("failed to install Ctrl+C handler");
};
#[cfg(unix)]
let terminate = async {
tokio::signal::unix::signal(tokio::signal::unix::SignalKind::terminate())
.expect("failed to install signal handler")
.recv()
.await;
};
#[cfg(not(unix))]
let terminate = std::future::pending::<()>();
tokio::select! {
_ = ctrl_c => {},
_ = terminate => {},
}
}
#[allow(dead_code)]
fn _assert_router_is_service()
where
Router: Service<
http::Request<axum::body::Body>,
Response = axum::response::Response,
Error = Infallible,
>,
{
}
#[cfg(test)]
mod tests {
use super::*;
use axum::body::Body;
use axum::http::{Method, Request, StatusCode};
use http_body_util::BodyExt;
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::TcpStream;
use tower::ServiceExt;
#[tokio::test]
async fn test_build_tcp_listener_with_random_port() {
let (listener, addr) = build_tcp_listener("127.0.0.1:0").await.unwrap();
assert!(addr.port() > 0);
assert_eq!(addr.ip().to_string(), "127.0.0.1");
let _ = listener.local_addr().unwrap();
}
#[tokio::test]
async fn test_build_tcp_listener_bind_error_for_invalid_addr() {
let result = build_tcp_listener("127.0.0.1:99999").await;
assert!(result.is_err());
}
#[tokio::test]
async fn test_serve_with_listener_responds_to_request() {
let router = Router::new().route("/", axum::routing::get(|| async { "hello from server" }));
let (listener, addr) = build_tcp_listener("127.0.0.1:0").await.unwrap();
tokio::spawn(async move {
let _ = serve_with_listener(router, listener).await;
});
tokio::time::sleep(std::time::Duration::from_millis(50)).await;
let body = http_get_body(addr.to_string().as_str(), "/").await;
assert!(body.contains("hello from server"));
}
#[tokio::test]
async fn test_router_responds_via_oneshot() {
let router = Router::new().route("/ping", axum::routing::get(|| async { "pong" }));
let request = Request::builder()
.method(Method::GET)
.uri("/ping")
.body(Body::empty())
.unwrap();
let response = router.oneshot(request).await.unwrap();
assert_eq!(response.status(), StatusCode::OK);
let bytes = response.into_body().collect().await.unwrap().to_bytes();
assert_eq!(&bytes[..], b"pong");
}
async fn http_get_body(host: &str, path: &str) -> String {
let mut stream = TcpStream::connect(host).await.unwrap();
let request = format!("GET {path} HTTP/1.1\r\nHost: {host}\r\nConnection: close\r\n\r\n");
stream.write_all(request.as_bytes()).await.unwrap();
let mut buf = Vec::new();
stream.read_to_end(&mut buf).await.unwrap();
let response = String::from_utf8_lossy(&buf).to_string();
if let Some(idx) = response.find("\r\n\r\n") {
response[idx + 4..].to_string()
} else {
response
}
}
#[tokio::test]
async fn test_serve_responds_to_request() {
let addr = {
let (_, addr) = build_tcp_listener("127.0.0.1:0").await.unwrap();
addr
};
let addr_str = addr.to_string();
let router = Router::new().route("/", axum::routing::get(|| async { "serve ok" }));
tokio::spawn(async move {
let _ = serve(router, &addr_str).await;
});
tokio::time::sleep(std::time::Duration::from_millis(100)).await;
let host = addr.to_string();
let body = http_get_body(&host, "/").await;
assert!(body.contains("serve ok"));
}
#[tokio::test]
async fn test_serve_with_graceful_shutdown_responds_to_request() {
let addr = {
let (_, addr) = build_tcp_listener("127.0.0.1:0").await.unwrap();
addr
};
let addr_str = addr.to_string();
let router = Router::new().route("/", axum::routing::get(|| async { "graceful ok" }));
tokio::spawn(async move {
let _ = serve_with_graceful_shutdown(router, &addr_str).await;
});
tokio::time::sleep(std::time::Duration::from_millis(100)).await;
let host = addr.to_string();
let body = http_get_body(&host, "/").await;
assert!(body.contains("graceful ok"));
}
#[tokio::test]
async fn test_build_tcp_listener_wildcard_addr() {
let (listener, addr) = build_tcp_listener("0.0.0.0:0").await.unwrap();
assert!(addr.port() > 0);
let _ = listener.local_addr().unwrap();
}
#[tokio::test]
async fn test_build_tcp_listener_invalid_ip() {
let result = build_tcp_listener("invalid_addr:8080").await;
assert!(result.is_err());
}
#[tokio::test]
async fn test_build_tcp_listener_empty_addr() {
let result = build_tcp_listener("").await;
assert!(result.is_err());
}
#[tokio::test]
async fn test_serve_with_listener_multiple_routes() {
let router = Router::new()
.route("/", axum::routing::get(|| async { "home" }))
.route("/api", axum::routing::get(|| async { "api" }));
let (listener, addr) = build_tcp_listener("127.0.0.1:0").await.unwrap();
tokio::spawn(async move {
let _ = serve_with_listener(router, listener).await;
});
tokio::time::sleep(std::time::Duration::from_millis(50)).await;
let host = addr.to_string();
let body1 = http_get_body(&host, "/").await;
assert!(body1.contains("home"));
let body2 = http_get_body(&host, "/api").await;
assert!(body2.contains("api"));
}
}