use std::env;
use std::sync::{
Arc,
atomic::{AtomicUsize, Ordering},
};
use std::time::{Duration, Instant};
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::{TcpListener, TcpStream};
use tokio::sync::{Barrier, Semaphore, oneshot};
use tokio::time::sleep;
use zerust::datapack::DataPack;
use zerust::{DefaultRouter, Response, Server};
#[tokio::main]
async fn main() -> Result<(), Box<dyn std::error::Error>> {
let args: Vec<String> = env::args().collect();
match args.get(1).map(|s| s.as_str()) {
Some("server") => run_server().await?,
Some("client") => {
let connections = args
.get(2)
.and_then(|s| s.parse::<usize>().ok())
.unwrap_or(100);
let requests_per_conn = args
.get(3)
.and_then(|s| s.parse::<usize>().ok())
.unwrap_or(1000);
run_client(connections, requests_per_conn).await?
}
_ => {
println!(
"用法: cargo run --release --example benchmark_server -- [server|client] [连接数] [每连接请求数]"
);
println!(" server - 启动基准测试服务器");
println!(" client [连接数] [每连接请求数] - 启动客户端测试");
}
}
Ok(())
}
async fn run_server() -> Result<(), Box<dyn std::error::Error>> {
let (shutdown_tx, shutdown_rx) = oneshot::channel::<()>();
let router = Arc::new(DefaultRouter::new());
let router_clone = router.clone();
let request_counter = Arc::new(AtomicUsize::new(0));
let counter_clone = request_counter.clone();
router_clone.add_route(1, move |req| {
counter_clone.fetch_add(1, Ordering::Relaxed);
Response::new(req.msg_id(), req.data().to_vec())
});
let server_addr = "127.0.0.1:8888";
let server = Server::new(server_addr, router);
println!("[Server] 基准测试服务器启动在 {}", server_addr);
let stats_handle = tokio::spawn(async move {
let mut last_count = 0;
let mut last_time = Instant::now();
loop {
sleep(Duration::from_secs(1)).await;
let current_count = request_counter.load(Ordering::Relaxed);
let current_time = Instant::now();
let elapsed = current_time.duration_since(last_time).as_secs_f64();
let rps = (current_count - last_count) as f64 / elapsed;
println!(
"[Stats] 当前RPS: {:.2} req/s, 总请求数: {}",
rps, current_count
);
last_count = current_count;
last_time = current_time;
}
});
let server_handle = tokio::spawn(async move {
if let Err(e) = server.run(shutdown_rx).await {
eprintln!("[Server] 运行时错误: {}", e);
}
});
println!("[Server] 按 Ctrl+C 停止服务器...");
tokio::signal::ctrl_c().await?;
println!("[Server] 接收到停止信号,正在关闭...");
let _ = shutdown_tx.send(());
let _ = server_handle.await;
stats_handle.abort();
println!("[Server] 服务器已关闭");
Ok(())
}
async fn run_client(
connections: usize,
requests_per_conn: usize,
) -> Result<(), Box<dyn std::error::Error>> {
println!(
"[Client] 开始基准测试: {} 并发连接, 每连接 {} 请求",
connections, requests_per_conn
);
let semaphore = Arc::new(Semaphore::new(connections));
let barrier = Arc::new(Barrier::new(connections + 1));
let total_requests = connections * requests_per_conn;
let completed_requests = Arc::new(AtomicUsize::new(0));
let total_latency = Arc::new(AtomicUsize::new(0));
let mut handles = Vec::with_capacity(connections);
let start_time = Instant::now();
for i in 0..connections {
let semaphore_clone = semaphore.clone();
let barrier_clone = barrier.clone();
let completed_clone = completed_requests.clone();
let latency_clone = total_latency.clone();
let handle = tokio::spawn(async move {
let _permit = semaphore_clone.acquire().await.unwrap();
let mut stream = match TcpStream::connect("127.0.0.1:8888").await {
Ok(stream) => stream,
Err(e) => {
eprintln!("[Client {}] 连接失败: {}", i, e);
return;
}
};
barrier_clone.wait().await;
for _ in 0..requests_per_conn {
let payload = vec![b'A'; 64]; let request = DataPack::pack(1, &payload);
let request_start = Instant::now();
if let Err(e) = stream.write_all(&request).await {
eprintln!("[Client {}] 发送请求失败: {}", i, e);
break;
}
let mut header = [0u8; 8];
if let Err(e) = stream.read_exact(&mut header).await {
eprintln!("[Client {}] 读取响应头失败: {}", i, e);
break;
}
let (msg_id, data_len) = match DataPack::unpack_header(&header) {
Ok(result) => result,
Err(e) => {
eprintln!("[Client {}] 解析响应头失败: {}", i, e);
break;
}
};
let mut data = vec![0u8; data_len as usize];
if let Err(e) = stream.read_exact(&mut data).await {
eprintln!("[Client {}] 读取响应数据失败: {}", i, e);
break;
}
let latency = request_start.elapsed().as_micros() as usize;
latency_clone.fetch_add(latency, Ordering::Relaxed);
completed_clone.fetch_add(1, Ordering::Relaxed);
}
});
handles.push(handle);
}
let progress_completed = completed_requests.clone();
let progress_handle = tokio::spawn(async move {
loop {
sleep(Duration::from_secs(1)).await;
let completed = progress_completed.load(Ordering::Relaxed);
let progress = (completed as f64 / total_requests as f64) * 100.0;
println!(
"[Progress] {:.2}% ({}/{})",
progress, completed, total_requests
);
if completed >= total_requests {
break;
}
}
});
println!("[Client] 所有连接已就绪,开始测试...");
barrier.wait().await;
for handle in handles {
let _ = handle.await;
}
progress_handle.abort();
let elapsed = start_time.elapsed();
let completed = completed_requests.load(Ordering::Relaxed);
let avg_latency = if completed > 0 {
total_latency.load(Ordering::Relaxed) as f64 / completed as f64
} else {
0.0
};
println!("\n===== 基准测试结果 =====");
println!("总连接数: {}", connections);
println!("每连接请求数: {}", requests_per_conn);
println!("总请求数: {}", total_requests);
println!("完成请求数: {}", completed);
println!("总耗时: {:.2} 秒", elapsed.as_secs_f64());
println!("平均延迟: {:.2} 微秒", avg_latency);
println!(
"吞吐量: {:.2} 请求/秒",
completed as f64 / elapsed.as_secs_f64()
);
Ok(())
}