Skip to main content

sz_rust_core/
server.rs

1//! HTTP 服务器模块 — axum::serve 启动器
2//!
3//! 对齐 PHP `think\swoole` / `think-worker` 启动入口,封装 axum::serve。
4//!
5//! ## 功能
6//!
7//! - `serve()`:基础启动器(不含 graceful shutdown)
8//! - `serve_with_graceful_shutdown()`:带优雅关闭(监听 Ctrl+C)
9//! - `serve_with_listener()`:使用自定义 tokio::net::TcpListener(测试友好)
10//! - `build_tcp_listener()`:构造 TCP listener
11//!
12//! ## 用法
13//!
14//! ```ignore
15//! use sz_rust_core::server::serve;
16//! use axum::Router;
17//!
18//! # async fn run() -> Result<(), Box<dyn std::error::Error>> {
19//! let router = Router::new();
20//! serve(router, "127.0.0.1:8801").await?;
21//! # Ok(())
22//! # }
23//! ```
24
25use std::convert::Infallible;
26use std::net::SocketAddr;
27
28use axum::Router;
29use tokio::net::TcpListener;
30use tower::Service;
31
32/// 启动 HTTP 服务器
33///
34/// 阻塞当前异步任务,直到服务器关闭。
35///
36/// ## 参数
37///
38/// - `router`:axum::Router(实现了 Service)
39/// - `addr`:监听地址,例如 `"127.0.0.1:8801"` 或 `"0.0.0.0:80"`
40///
41/// ## 错误
42///
43/// 绑定端口失败时返回 `std::io::Error`。
44pub 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
50/// 启动 HTTP 服务器(带优雅关闭)
51///
52/// 监听 Ctrl+C 信号,收到后启动 graceful shutdown。
53///
54/// ## 参数
55///
56/// - `router`:axum::Router
57/// - `addr`:监听地址
58pub 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
66/// 启动 HTTP 服务器(使用已有 TcpListener,测试友好)
67///
68/// 适用于测试场景:测试代码可以 `listener.local_addr()` 获取实际端口,
69/// 然后在另一个 task 中连接。也适用于 Unix socket 等自定义 listener。
70///
71/// ## 参数
72///
73/// - `router`:axum::Router
74/// - `listener`:已绑定的 tokio::net::TcpListener
75pub async fn serve_with_listener(router: Router, listener: TcpListener) -> std::io::Result<()> {
76    axum::serve(listener, router).await?;
77    Ok(())
78}
79
80/// 构造 TCP listener
81///
82/// 内部使用 `tokio::net::TcpListener::bind`,返回 listener 和实际绑定的地址。
83/// 适用于测试场景:传入 `"127.0.0.1:0"` 让 OS 分配端口。
84///
85/// ## 返回
86///
87/// `(TcpListener, SocketAddr)`,addr 是实际绑定的地址(端口可能为 0 表示由 OS 分配)。
88pub 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
94/// 优雅关闭信号监听
95///
96/// 监听 Ctrl+C / SIGTERM,返回后触发 axum 的 graceful shutdown。
97async 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// 用于编译期验证 Router 满足 Service trait 约束
122#[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        // listener 必须可用
149        let _ = listener.local_addr().unwrap();
150    }
151
152    #[tokio::test]
153    async fn test_build_tcp_listener_bind_error_for_invalid_addr() {
154        // 端口号超出范围
155        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        // 等服务器就绪
169        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        // 验证 Router 不需要真实 TCP 也能直接 oneshot 测试
178        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    /// 最小 HTTP/1.1 GET 客户端,避免引入 reqwest 依赖
191    ///
192    /// 返回响应体(HTTP body 部分)字符串。
193    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        // 分离 header / body
203        if let Some(idx) = response.find("\r\n\r\n") {
204            response[idx + 4..].to_string()
205        } else {
206            response
207        }
208    }
209
210    // ---- 补充测试:覆盖 serve() 和 serve_with_graceful_shutdown() ----
211
212    #[tokio::test]
213    async fn test_serve_responds_to_request() {
214        // 获取空闲端口(bind 后 drop 释放端口,serve 内部重新 bind)
215        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}