#![deny(unsafe_code)]
use std::env;
use std::net::{SocketAddr, TcpStream};
use std::process::ExitCode;
use std::sync::atomic::{AtomicU32, Ordering};
use std::sync::Arc;
use std::time::{SystemTime, UNIX_EPOCH};
use zenith_stack::api::CanonicalResponse;
use zenith_stack::web::app::{success_response, App};
use zenith_stack::web::server::{ProtocolServer, ServerConfig, ServerError};
#[cfg(feature = "tls")]
use zenith_stack::tls::{CertGeneration, TlsAcceptor, TlsConfig};
#[cfg(all(feature = "tls", feature = "http3"))]
use zenith_stack::web::quic_server::QuicServerConfig;
#[derive(Debug, Clone, Copy)]
enum Mode {
Plain,
#[cfg(feature = "tls")]
Tls,
#[cfg(all(feature = "tls", feature = "http3"))]
H3,
}
struct Args {
mode: Mode,
bind: SocketAddr,
audit_log: Option<String>,
}
fn parse_args() -> Result<Args, String> {
let mut mode = Mode::Plain;
let mut port: u16 = 18080;
let mut host: String = "::".to_string();
let mut audit_log: Option<String> = None;
let mut it = env::args().skip(1);
while let Some(arg) = it.next() {
match arg.as_str() {
"--mode" => {
#[allow(unreachable_patterns)]
let val = it.next().ok_or("--mode 需要 <plain|tls|h3>")?;
match val.as_str() {
"plain" => {
mode = Mode::Plain;
if port == 18443 {
port = 18080;
}
}
#[cfg(feature = "tls")]
"tls" => {
mode = Mode::Tls;
if port == 18080 {
port = 18443;
}
}
#[cfg(all(feature = "tls", feature = "http3"))]
"h3" => {
mode = Mode::H3;
if port == 18080 {
port = 18443;
}
}
#[cfg(not(feature = "tls"))]
"tls" | "h3" => {
return Err(
"tls/h3 模式需要启用 tls feature: cargo run --features web,tls ..."
.to_string(),
);
}
#[cfg(all(feature = "tls", not(feature = "http3")))]
"h3" => {
return Err(
"h3 模式需要启用 http3 feature: cargo run --features web,tls,http3 ..."
.to_string(),
);
}
other => return Err(format!("未知 --mode {other:?},可选 <plain|tls|h3>")),
}
}
"--port" => {
let val = it.next().ok_or("--port 需要 <PORT>")?;
port = val.parse::<u16>().map_err(|e| format!("--port 解析失败: {e}"))?;
}
"--host" => {
host = it.next().ok_or("--host 需要 <HOST>")?;
}
"--bind" => {
let val = it.next().ok_or("--bind 需要 <HOST:PORT>")?;
let sa: SocketAddr = val.parse().map_err(|e| format!("--bind 解析失败: {e}"))?;
host = sa.ip().to_string();
port = sa.port();
}
"--audit-log" => {
let val = it.next().ok_or("--audit-log 需要 <PATH>")?;
audit_log = Some(val);
}
"-h" | "--help" => {
println!(
"{}",
concat!(
"hello_server — Zenith Web 真实端口监听演示\n",
"USAGE:\n",
" hello_server [--mode plain|tls|h3] [--port PORT] [--host HOST]\n",
"FLAGS:\n",
" -h, --help 显示帮助\n",
"OPTIONS:\n",
" --mode <plain|tls|h3> 运行模式,默认 plain\n",
" plain = HTTP/1.1 over TCP (明文)\n",
" tls = HTTP/1.1 + TLS 1.3 over TCP\n",
" h3 = HTTP/3 over QUIC over UDP\n",
" --port <u16> 监听端口,默认 plain=18080 / tls=18443 / h3=18443\n",
" --host <ip> 监听地址,默认 ::\n",
" --bind <HOST:PORT> 便捷合并写法(等价于同时指定 --host 和 --port)\n",
" --audit-log <PATH> 审计日志落盘(JSON Lines 追加写;缺省仅内存)\n",
)
);
std::process::exit(0);
}
other => return Err(format!("未知参数 {other:?},使用 -h 查看帮助")),
}
}
let bind: SocketAddr = if host.contains(':') {
format!("[{host}]:{port}")
} else {
format!("{host}:{port}")
}
.parse()
.map_err(|e| format!("SocketAddr 解析失败: {e}"))?;
Ok(Args { mode, bind, audit_log })
}
fn build_server(app: App, args: &Args) -> Result<ProtocolServer, String> {
let cfg = ServerConfig::new().with_waf(true);
match &args.audit_log {
Some(path) => ProtocolServer::try_with_config(app, cfg.with_audit_log_path(path.clone()))
.map_err(|e| format!("审计日志落盘初始化失败 ({path}): {e}")),
None => Ok(ProtocolServer::with_config(app, cfg)),
}
}
fn build_app() -> App {
let mut app = App::new();
app.get("/", |_req, _rm| {
let body = concat!(
"<!doctype html><html><body style=\"font-family:sans-serif;padding:2rem\">",
"<h1>Zenith Hello Server ✓</h1>",
"<p>真实链路:std::net::TcpListener → ProtocolServer::serve_std_tcp_conn → ",
"Router → Handler</p>",
"<ul>",
"<li><code>GET /</code> — 本页</li>",
"<li><code>GET /hello/:name</code> — 路径参数问候</li>",
"<li><code>POST /echo</code> — 回显请求体</li>",
"<li><code>GET /status</code> — 服务器状态 JSON</li>",
"<li><code>GET /404</code> — 404 演示</li>",
"</ul></body></html>\n"
);
Ok(success_response(body, "text/html; charset=utf-8"))
});
app.get("/hello/:name", |_req, rm| {
let name = rm.get("name").unwrap_or("stranger");
Ok(success_response(
format!("Hello, {name}! 🚀\n"),
"text/plain; charset=utf-8",
))
});
app.post("/echo", |req, _rm| {
let ct: Vec<u8> = req
.find_header("content-type")
.map(|h| {
let len = h.value_len as usize;
h.value[..len].to_vec()
})
.unwrap_or_else(|| b"application/octet-stream".to_vec());
let mut resp = CanonicalResponse::new(200);
let _ = resp.add_header(b"content-type", &ct);
let _ = resp.add_header(b"x-echoed-length", &req.body().len().to_string().into_bytes());
resp.set_body(req.body().to_vec());
Ok(resp)
});
app.get("/status", |_req, _rm| {
let ts = SystemTime::now()
.duration_since(UNIX_EPOCH)
.map(|d| d.as_secs())
.unwrap_or(0);
let json = format!(
concat!(
"{{",
"\"ok\":true,",
"\"server\":\"zenith-web\",",
"\"protocol\":\"HTTP/1.1 (+TLS 1.3 when tls mode)\",",
"\"timestamp_sec\":{ts}",
"}}\n"
),
ts = ts,
);
Ok(success_response(json, "application/json"))
});
if let Err(e) = app.validate() {
eprintln!("[hello_server] 路由配置错误: {e}");
}
app
}
fn handle_conn_plain(
server: &ProtocolServer,
stream: TcpStream,
peer: SocketAddr,
) -> Result<(), ServerError> {
server.serve_std_tcp_conn(stream, peer, None)
}
#[cfg(feature = "tls")]
fn handle_conn_tls(
server: &ProtocolServer,
stream: TcpStream,
peer: SocketAddr,
tls_acceptor: &mut TlsAcceptor,
) -> Result<(), ServerError> {
server.serve_std_tcp_conn(stream, peer, Some(tls_acceptor))
}
#[cfg(feature = "tls")]
fn build_demo_cert() -> Result<CertGeneration, ExitCode> {
let cert_path = env::var("ZENITH_DEMO_CERT_PATH");
let key_path = env::var("ZENITH_DEMO_KEY_PATH");
if let (Ok(c), Ok(k)) = (cert_path, key_path) {
let cert_pem = std::fs::read(&c).map_err(|e| {
eprintln!("[hello_server] 读取证书文件失败 ({c}): {e}");
ExitCode::from(1)
})?;
let key_pem = std::fs::read(&k).map_err(|e| {
eprintln!("[hello_server] 读取私钥文件失败 ({k}): {e}");
ExitCode::from(1)
})?;
return CertGeneration::from_pem(&cert_pem, &key_pem).map_err(|e| {
eprintln!("[hello_server] 演示证书 PEM 解析失败: {e}");
ExitCode::from(1)
});
}
let ts = SystemTime::now()
.duration_since(UNIX_EPOCH)
.map(|d| d.as_nanos())
.unwrap_or(0);
let suffix = format!("{}_{}", std::process::id(), ts);
let tmp_cert = std::env::temp_dir()
.join(format!("zenith_demo_cert_{suffix}.pem"))
.to_string_lossy()
.into_owned();
let tmp_key = std::env::temp_dir()
.join(format!("zenith_demo_key_{suffix}.pem"))
.to_string_lossy()
.into_owned();
let status = std::process::Command::new("openssl")
.args([
"req", "-x509", "-newkey", "rsa:2048",
"-keyout", tmp_key.as_str(), "-out", tmp_cert.as_str(),
"-days", "1", "-nodes", "-subj", "/CN=localhost",
])
.stdout(std::process::Stdio::null())
.stderr(std::process::Stdio::null())
.status();
match status {
Ok(s) if s.success() => {
let cert_pem = std::fs::read(&tmp_cert).map_err(|e| {
eprintln!("[hello_server] 读取临时证书失败: {e}");
ExitCode::from(1)
})?;
let key_pem = std::fs::read(&tmp_key).map_err(|e| {
eprintln!("[hello_server] 读取临时私钥失败: {e}");
ExitCode::from(1)
})?;
let _ = std::fs::remove_file(&tmp_cert);
let _ = std::fs::remove_file(&tmp_key);
CertGeneration::from_pem(&cert_pem, &key_pem).map_err(|e| {
eprintln!("[hello_server] 临时证书 PEM 解析失败: {e}");
ExitCode::from(1)
})
}
_ => {
eprintln!(
"[hello_server] openssl 不可用,无法生成临时证书。\n \
请设置环境变量:\n \
export ZENITH_DEMO_CERT_PATH=/path/to/cert.pem\n \
export ZENITH_DEMO_KEY_PATH=/path/to/key.pem"
);
Err(ExitCode::from(1))
}
}
}
#[cfg(feature = "tls")]
fn build_tls_acceptor() -> Result<TlsAcceptor, ExitCode> {
let std_transport_alpn: Vec<Vec<u8>> = vec![b"h2".to_vec(), b"http/1.1".to_vec()];
let mut tls_config = TlsConfig::new().with_alpn(std_transport_alpn);
tls_config.min_version = Some(zenith_stack::tls::TlsVersion::Tls12);
tls_config.max_version = Some(zenith_stack::tls::TlsVersion::Tls13);
let mut acceptor = TlsAcceptor::new(tls_config);
let demo_cert = build_demo_cert()?;
acceptor.set_default_cert(demo_cert.clone()).map_err(|e| {
eprintln!("[hello_server] 设置默认证书失败: {e}");
ExitCode::from(4)
})?;
acceptor.set_cert_for_domain("localhost", demo_cert).map_err(|e| {
eprintln!("[hello_server] 设置 localhost SNI 证书失败: {e}");
ExitCode::from(4)
})?;
Ok(acceptor)
}
fn main() -> ExitCode {
let _ = tracing_subscriber::fmt()
.with_env_filter(
tracing_subscriber::EnvFilter::try_from_default_env()
.unwrap_or_else(|_| tracing_subscriber::EnvFilter::new("info")),
)
.try_init();
let _global_rt = zenith_stack::rt::init_global(zenith_stack::rt::RuntimeConfig::auto());
let args = match parse_args() {
Ok(a) => a,
Err(e) => {
eprintln!("参数错误: {e}");
return ExitCode::from(2);
}
};
let app = build_app();
let route_count = app.route_count();
println!(
"[hello_server] mode={mode:?} bind={bind} routes={n}",
mode = args.mode,
bind = args.bind,
n = route_count,
);
#[cfg(all(feature = "tls", feature = "http3"))]
if matches!(args.mode, Mode::H3) {
let mut server = match build_server(app, &args) {
Ok(s) => s,
Err(e) => {
eprintln!("{e}");
return ExitCode::from(2);
}
};
let cert_gen = match build_demo_cert() {
Ok(g) => g,
Err(code) => return code,
};
let quic_cfg = QuicServerConfig {
bind_addr: args.bind,
..QuicServerConfig::default()
};
if let Err(e) = server.bind_quic(quic_cfg, &cert_gen) {
eprintln!("[hello_server] bind_quic({}) 失败: {e}", args.bind);
return ExitCode::from(5);
}
println!(
"[hello_server] QUIC/UDP 监听成功,本地 addr = {},证书 CN=localhost",
server.quic_local_addr().unwrap_or(args.bind)
);
match server.trigger_changeset() {
Ok(generation) => println!("[hello_server] ChangeSet 热更新演示成功,生成号 = {generation}"),
Err(e) => eprintln!("[hello_server] ChangeSet 热更新演示失败: {e}"),
}
println!("[hello_server] Ctrl+C 退出。");
match server.serve_udp_loop(65536) {
Ok(()) => ExitCode::SUCCESS,
Err(e) => {
eprintln!("[hello_server] serve_udp_loop 错误: {e}");
ExitCode::from(6)
}
}
} else {
run_tcp_server(args, app)
}
#[cfg(not(all(feature = "tls", feature = "http3")))]
{
run_tcp_server(args, app)
}
}
fn run_tcp_server(args: Args, app: App) -> ExitCode {
let server = match build_server(app, &args) {
Ok(s) => s,
Err(e) => {
eprintln!("{e}");
return ExitCode::from(2);
}
};
let server_arc: Arc<ProtocolServer> = Arc::new(server);
match server_arc.trigger_changeset() {
Ok(generation) => println!("[hello_server] ChangeSet 热更新演示成功,生成号 = {generation}"),
Err(e) => eprintln!("[hello_server] ChangeSet 热更新演示失败: {e}"),
}
#[cfg(feature = "tls")]
let tls_acceptor_maybe: Option<TlsAcceptor> = match args.mode {
Mode::Plain => None,
#[cfg(feature = "tls")]
Mode::Tls => match build_tls_acceptor() {
Ok(a) => {
println!("[hello_server] TLS Acceptor 已就绪,证书代际已加载,CN=localhost");
Some(a)
}
Err(code) => return code,
},
#[cfg(all(feature = "tls", feature = "http3"))]
Mode::H3 => unreachable!("H3 模式已在上层分派"),
};
println!("[hello_server] Ctrl+C 退出。");
const MAX_CONNECTIONS: u32 = 1024;
let active_conns: Arc<AtomicU32> = Arc::new(AtomicU32::new(0));
let _ = zenith_stack::rt::block_on(async move {
let listener = match tokio::net::TcpListener::bind(args.bind).await {
Ok(l) => l,
Err(e) => {
eprintln!("[hello_server] tokio TcpListener::bind({}) 失败: {e}", args.bind);
return;
}
};
println!(
"[hello_server] tokio TcpListener 监听成功,本地 addr = {}",
listener.local_addr().unwrap_or(args.bind)
);
let server_for_stats = server_arc.clone();
zenith_stack::rt::spawn(async move {
loop {
tokio::time::sleep(std::time::Duration::from_secs(30)).await;
let snapshot = server_for_stats.export_metrics();
println!("[hello_server] 后台统计(30s 周期):\n{snapshot}");
}
});
loop {
let (stream, peer) = match listener.accept().await {
Ok((s, p)) => (s, p),
Err(e) => {
eprintln!("[hello_server] accept 错误: {e}");
continue;
}
};
if active_conns.fetch_add(1, Ordering::Relaxed) >= MAX_CONNECTIONS {
active_conns.fetch_sub(1, Ordering::Relaxed);
eprintln!(
"[hello_server] 连接数已达上限 {MAX_CONNECTIONS},拒绝 peer={peer}"
);
let _ = stream.into_std();
continue;
}
let server_clone: Arc<ProtocolServer> = Arc::clone(&server_arc);
let conn_counter = Arc::clone(&active_conns);
#[cfg(feature = "tls")]
let tls_acc_clone = tls_acceptor_maybe.clone();
let mode = args.mode;
zenith_stack::rt::spawn(async move {
let std_stream = match stream.into_std() {
Ok(s) => {
let _ = s.set_nonblocking(false);
s
}
Err(e) => {
eprintln!("[hello_server] stream 转换失败 peer={peer}: {e}");
conn_counter.fetch_sub(1, Ordering::Relaxed);
return;
}
};
let result = tokio::task::spawn_blocking(move || {
let result: Result<(), ServerError> = match mode {
Mode::Plain => handle_conn_plain(&server_clone, std_stream, peer),
#[cfg(feature = "tls")]
Mode::Tls => {
let mut acc = tls_acc_clone.expect("tls mode should have built acceptor");
handle_conn_tls(&server_clone, std_stream, peer, &mut acc)
}
#[cfg(all(feature = "tls", feature = "http3"))]
Mode::H3 => unreachable!("H3 模式已在上层分派"),
};
match result {
Ok(()) => println!("[hello_server] 连接正常结束 peer={peer}"),
Err(ServerError::ConnectionClosed) => {
println!("[hello_server] 对端关闭 peer={peer}")
}
Err(e) => eprintln!("[hello_server] 连接处理错误 peer={peer}: {e}"),
}
}).await;
if let Err(e) = result {
eprintln!("[hello_server] spawn_blocking join 错误 peer={peer}: {e}");
}
conn_counter.fetch_sub(1, Ordering::Relaxed);
});
}
});
ExitCode::SUCCESS
}