use std::{
net::SocketAddr,
sync::{Arc, Mutex},
time::Instant,
};
use anyhow::{Context, Result};
use clap::Parser;
use rustls::pki_types::{CertificateDer, PrivatePkcs8KeyDer};
use tokio::sync::Semaphore;
use tracing::{info, trace};
use bench::{
Opt, configure_tracing_subscriber, connect_client, drain_stream, rt, send_data_on_stream,
server_endpoint,
stats::{Stats, TransferResult},
};
fn main() {
let opt = Opt::parse();
configure_tracing_subscriber();
let cert = rcgen::generate_simple_self_signed(vec!["localhost".into()]).unwrap();
let key = PrivatePkcs8KeyDer::from(cert.signing_key.serialize_der());
let cert = CertificateDer::from(cert.cert);
let server_span = tracing::error_span!("server");
let runtime = rt();
let (server_addr, endpoint) = {
let _guard = server_span.enter();
server_endpoint(&runtime, cert.clone(), key.into(), &opt)
};
let server_thread = std::thread::spawn(move || {
let _guard = server_span.entered();
if let Err(e) = runtime.block_on(server(endpoint, opt)) {
eprintln!("server failed: {e:#}");
}
});
let mut handles = Vec::new();
for id in 0..opt.clients {
let cert = cert.clone();
handles.push(std::thread::spawn(move || {
let _guard = tracing::error_span!("client", id).entered();
let runtime = rt();
match runtime.block_on(client(server_addr, cert, opt)) {
Ok(stats) => Ok(stats),
Err(e) => {
eprintln!("client failed: {e:#}");
Err(e)
}
}
}));
}
for (id, handle) in handles.into_iter().enumerate() {
if let Ok(stats) = handle.join().expect("client thread") {
stats.print(id);
}
}
server_thread.join().expect("server thread");
}
async fn server(endpoint: quinn::Endpoint, opt: Opt) -> Result<()> {
let mut server_tasks = Vec::new();
for _ in 0..opt.clients {
let handshake = endpoint.accept().await.unwrap();
let connection = handshake.await.context("handshake failed")?;
server_tasks.push(tokio::spawn(async move {
loop {
let (mut send_stream, recv_stream) = match connection.accept_bi().await {
Err(quinn::ConnectionError::ApplicationClosed(_)) => break,
Err(e) => {
eprintln!("accepting stream failed: {e:?}");
break;
}
Ok(stream) => stream,
};
trace!("stream established");
tokio::spawn(async move {
drain_stream(recv_stream, opt.read_unordered).await?;
send_data_on_stream(&mut send_stream, opt.download_size).await?;
Ok::<_, anyhow::Error>(())
});
}
if opt.stats {
println!("\nServer connection stats:\n{:#?}", connection.stats());
}
}));
}
for handle in server_tasks {
if let Err(e) = handle.await {
eprintln!("Server task error: {e:?}");
};
}
Ok(())
}
async fn client(
server_addr: SocketAddr,
server_cert: CertificateDer<'static>,
opt: Opt,
) -> Result<ClientStats> {
let (endpoint, connection) = connect_client(server_addr, server_cert, opt).await?;
let start = Instant::now();
let connection = Arc::new(connection);
let mut stats = ClientStats::default();
let mut first_error = None;
let sem = Arc::new(Semaphore::new(opt.max_streams));
let results = Arc::new(Mutex::new(Vec::new()));
for _ in 0..opt.streams {
let permit = sem.clone().acquire_owned().await.unwrap();
let results = results.clone();
let connection = connection.clone();
tokio::spawn(async move {
let result =
handle_client_stream(connection, opt.upload_size, opt.read_unordered).await;
info!("stream finished: {:?}", result);
results.lock().unwrap().push(result);
drop(permit);
});
}
let _ = sem.acquire_many(opt.max_streams as u32).await.unwrap();
for result in results.lock().unwrap().drain(..) {
match result {
Ok((upload_result, download_result)) => {
stats.upload_stats.stream_finished(upload_result);
stats.download_stats.stream_finished(download_result);
}
Err(e) => {
if first_error.is_none() {
first_error = Some(e);
}
}
}
}
stats.upload_stats.total_duration = start.elapsed();
stats.download_stats.total_duration = start.elapsed();
connection.close(0u32.into(), b"Benchmark done");
endpoint.wait_idle().await;
if opt.stats {
println!("\nClient connection stats:\n{:#?}", connection.stats());
}
match first_error {
None => Ok(stats),
Some(e) => Err(e),
}
}
async fn handle_client_stream(
connection: Arc<quinn::Connection>,
upload_size: u64,
read_unordered: bool,
) -> Result<(TransferResult, TransferResult)> {
let start = Instant::now();
let (mut send_stream, recv_stream) = connection
.open_bi()
.await
.context("failed to open stream")?;
send_data_on_stream(&mut send_stream, upload_size).await?;
let upload_result = TransferResult::new(start.elapsed(), upload_size);
let start = Instant::now();
let size = drain_stream(recv_stream, read_unordered).await?;
let download_result = TransferResult::new(start.elapsed(), size as u64);
Ok((upload_result, download_result))
}
#[derive(Default)]
struct ClientStats {
upload_stats: Stats,
download_stats: Stats,
}
impl ClientStats {
fn print(&self, client_id: usize) {
println!();
println!("Client {client_id} stats:");
if self.upload_stats.total_size != 0 {
self.upload_stats.print("upload");
}
if self.download_stats.total_size != 0 {
self.download_stats.print("download");
}
}
}