1use std::convert::Infallible;
26use std::net::SocketAddr;
27
28use axum::Router;
29use tokio::net::TcpListener;
30use tower::Service;
31
32pub async fn serve(router: Router, addr: &str) -> std::io::Result<()> {
45 let listener = TcpListener::bind(addr).await?;
46 axum::serve(listener, router).await?;
47 Ok(())
48}
49
50pub async fn serve_with_graceful_shutdown(router: Router, addr: &str) -> std::io::Result<()> {
59 let listener = TcpListener::bind(addr).await?;
60 axum::serve(listener, router)
61 .with_graceful_shutdown(shutdown_signal())
62 .await?;
63 Ok(())
64}
65
66pub async fn serve_with_listener(router: Router, listener: TcpListener) -> std::io::Result<()> {
76 axum::serve(listener, router).await?;
77 Ok(())
78}
79
80pub async fn build_tcp_listener(addr: &str) -> std::io::Result<(TcpListener, SocketAddr)> {
89 let listener = TcpListener::bind(addr).await?;
90 let local_addr = listener.local_addr()?;
91 Ok((listener, local_addr))
92}
93
94async fn shutdown_signal() {
98 let ctrl_c = async {
99 tokio::signal::ctrl_c()
100 .await
101 .expect("failed to install Ctrl+C handler");
102 };
103
104 #[cfg(unix)]
105 let terminate = async {
106 tokio::signal::unix::signal(tokio::signal::unix::SignalKind::terminate())
107 .expect("failed to install signal handler")
108 .recv()
109 .await;
110 };
111
112 #[cfg(not(unix))]
113 let terminate = std::future::pending::<()>();
114
115 tokio::select! {
116 _ = ctrl_c => {},
117 _ = terminate => {},
118 }
119}
120
121#[allow(dead_code)]
123fn _assert_router_is_service()
124where
125 Router: Service<
126 http::Request<axum::body::Body>,
127 Response = axum::response::Response,
128 Error = Infallible,
129 >,
130{
131}
132
133#[cfg(test)]
134mod tests {
135 use super::*;
136 use axum::body::Body;
137 use axum::http::{Method, Request, StatusCode};
138 use http_body_util::BodyExt;
139 use tokio::io::{AsyncReadExt, AsyncWriteExt};
140 use tokio::net::TcpStream;
141 use tower::ServiceExt;
142
143 #[tokio::test]
144 async fn test_build_tcp_listener_with_random_port() {
145 let (listener, addr) = build_tcp_listener("127.0.0.1:0").await.unwrap();
146 assert!(addr.port() > 0);
147 assert_eq!(addr.ip().to_string(), "127.0.0.1");
148 let _ = listener.local_addr().unwrap();
150 }
151
152 #[tokio::test]
153 async fn test_build_tcp_listener_bind_error_for_invalid_addr() {
154 let result = build_tcp_listener("127.0.0.1:99999").await;
156 assert!(result.is_err());
157 }
158
159 #[tokio::test]
160 async fn test_serve_with_listener_responds_to_request() {
161 let router = Router::new().route("/", axum::routing::get(|| async { "hello from server" }));
162 let (listener, addr) = build_tcp_listener("127.0.0.1:0").await.unwrap();
163
164 tokio::spawn(async move {
165 let _ = serve_with_listener(router, listener).await;
166 });
167
168 tokio::time::sleep(std::time::Duration::from_millis(50)).await;
170
171 let body = http_get_body(addr.to_string().as_str(), "/").await;
172 assert!(body.contains("hello from server"));
173 }
174
175 #[tokio::test]
176 async fn test_router_responds_via_oneshot() {
177 let router = Router::new().route("/ping", axum::routing::get(|| async { "pong" }));
179 let request = Request::builder()
180 .method(Method::GET)
181 .uri("/ping")
182 .body(Body::empty())
183 .unwrap();
184 let response = router.oneshot(request).await.unwrap();
185 assert_eq!(response.status(), StatusCode::OK);
186 let bytes = response.into_body().collect().await.unwrap().to_bytes();
187 assert_eq!(&bytes[..], b"pong");
188 }
189
190 async fn http_get_body(host: &str, path: &str) -> String {
194 let mut stream = TcpStream::connect(host).await.unwrap();
195 let request = format!("GET {path} HTTP/1.1\r\nHost: {host}\r\nConnection: close\r\n\r\n");
196 stream.write_all(request.as_bytes()).await.unwrap();
197
198 let mut buf = Vec::new();
199 stream.read_to_end(&mut buf).await.unwrap();
200 let response = String::from_utf8_lossy(&buf).to_string();
201
202 if let Some(idx) = response.find("\r\n\r\n") {
204 response[idx + 4..].to_string()
205 } else {
206 response
207 }
208 }
209
210 #[tokio::test]
213 async fn test_serve_responds_to_request() {
214 let addr = {
216 let (_, addr) = build_tcp_listener("127.0.0.1:0").await.unwrap();
217 addr
218 };
219 let addr_str = addr.to_string();
220
221 let router = Router::new().route("/", axum::routing::get(|| async { "serve ok" }));
222 tokio::spawn(async move {
223 let _ = serve(router, &addr_str).await;
224 });
225
226 tokio::time::sleep(std::time::Duration::from_millis(100)).await;
227 let host = addr.to_string();
228 let body = http_get_body(&host, "/").await;
229 assert!(body.contains("serve ok"));
230 }
231
232 #[tokio::test]
233 async fn test_serve_with_graceful_shutdown_responds_to_request() {
234 let addr = {
235 let (_, addr) = build_tcp_listener("127.0.0.1:0").await.unwrap();
236 addr
237 };
238 let addr_str = addr.to_string();
239
240 let router = Router::new().route("/", axum::routing::get(|| async { "graceful ok" }));
241 tokio::spawn(async move {
242 let _ = serve_with_graceful_shutdown(router, &addr_str).await;
243 });
244
245 tokio::time::sleep(std::time::Duration::from_millis(100)).await;
246 let host = addr.to_string();
247 let body = http_get_body(&host, "/").await;
248 assert!(body.contains("graceful ok"));
249 }
250
251 #[tokio::test]
252 async fn test_build_tcp_listener_wildcard_addr() {
253 let (listener, addr) = build_tcp_listener("0.0.0.0:0").await.unwrap();
254 assert!(addr.port() > 0);
255 let _ = listener.local_addr().unwrap();
256 }
257
258 #[tokio::test]
259 async fn test_build_tcp_listener_invalid_ip() {
260 let result = build_tcp_listener("invalid_addr:8080").await;
261 assert!(result.is_err());
262 }
263
264 #[tokio::test]
265 async fn test_build_tcp_listener_empty_addr() {
266 let result = build_tcp_listener("").await;
267 assert!(result.is_err());
268 }
269
270 #[tokio::test]
271 async fn test_serve_with_listener_multiple_routes() {
272 let router = Router::new()
273 .route("/", axum::routing::get(|| async { "home" }))
274 .route("/api", axum::routing::get(|| async { "api" }));
275 let (listener, addr) = build_tcp_listener("127.0.0.1:0").await.unwrap();
276
277 tokio::spawn(async move {
278 let _ = serve_with_listener(router, listener).await;
279 });
280
281 tokio::time::sleep(std::time::Duration::from_millis(50)).await;
282
283 let host = addr.to_string();
284 let body1 = http_get_body(&host, "/").await;
285 assert!(body1.contains("home"));
286
287 let body2 = http_get_body(&host, "/api").await;
288 assert!(body2.contains("api"));
289 }
290}