bulk/
bulk.rs

1use std::{
2    net::SocketAddr,
3    sync::{Arc, Mutex},
4    time::Instant,
5};
6
7use anyhow::{Context, Result};
8use clap::Parser;
9use rustls::pki_types::{CertificateDer, PrivatePkcs8KeyDer};
10use tokio::sync::Semaphore;
11use tracing::{info, trace};
12
13use bench::{
14    Opt, configure_tracing_subscriber, connect_client, drain_stream, rt, send_data_on_stream,
15    server_endpoint,
16    stats::{Stats, TransferResult},
17};
18
19fn main() {
20    let opt = Opt::parse();
21    configure_tracing_subscriber();
22
23    let cert = rcgen::generate_simple_self_signed(vec!["localhost".into()]).unwrap();
24    let key = PrivatePkcs8KeyDer::from(cert.signing_key.serialize_der());
25    let cert = CertificateDer::from(cert.cert);
26
27    let server_span = tracing::error_span!("server");
28    let runtime = rt(opt.runtime_type);
29    let (server_addr, endpoint) = {
30        let _guard = server_span.enter();
31        server_endpoint(&runtime, cert.clone(), key.into(), &opt)
32    };
33
34    let server_thread = std::thread::spawn(move || {
35        let _guard = server_span.entered();
36        if let Err(e) = runtime.block_on(server(endpoint, opt)) {
37            eprintln!("server failed: {e:#}");
38        }
39    });
40
41    let mut handles = Vec::new();
42    for id in 0..opt.clients {
43        let cert = cert.clone();
44        handles.push(std::thread::spawn(move || {
45            let _guard = tracing::error_span!("client", id).entered();
46            let runtime = rt(opt.runtime_type);
47            match runtime.block_on(client(server_addr, cert, opt)) {
48                Ok(stats) => Ok(stats),
49                Err(e) => {
50                    eprintln!("client failed: {e:#}");
51                    Err(e)
52                }
53            }
54        }));
55    }
56
57    for (id, handle) in handles.into_iter().enumerate() {
58        // We print all stats at the end of the test sequentially to avoid
59        // them being garbled due to being printed concurrently
60        if let Ok(stats) = handle.join().expect("client thread") {
61            stats.print(id);
62        }
63    }
64
65    server_thread.join().expect("server thread");
66}
67
68async fn server(endpoint: noq::Endpoint, opt: Opt) -> Result<()> {
69    let mut server_tasks = Vec::new();
70
71    // Handle only the expected amount of clients
72    for _ in 0..opt.clients {
73        let handshake = endpoint.accept().await.unwrap();
74        let connection = handshake.await.context("handshake failed")?;
75
76        server_tasks.push(tokio::spawn(async move {
77            loop {
78                let (mut send_stream, recv_stream) = match connection.accept_bi().await {
79                    Err(noq::ConnectionError::ApplicationClosed(_)) => break,
80                    Err(e) => {
81                        eprintln!("accepting stream failed: {e:?}");
82                        break;
83                    }
84                    Ok(stream) => stream,
85                };
86                trace!("stream established");
87
88                tokio::spawn(async move {
89                    drain_stream(recv_stream, opt.read_unordered).await?;
90                    send_data_on_stream(&mut send_stream, opt.download_size).await?;
91                    Ok::<_, anyhow::Error>(())
92                });
93            }
94
95            if opt.stats {
96                println!("\nServer connection stats:\n{:#?}", connection.stats());
97            }
98        }));
99    }
100
101    // Await all the tasks. We have to do this to prevent the runtime getting dropped
102    // and all server tasks to be cancelled
103    for handle in server_tasks {
104        if let Err(e) = handle.await {
105            eprintln!("Server task error: {e:?}");
106        };
107    }
108
109    Ok(())
110}
111
112async fn client(
113    server_addr: SocketAddr,
114    server_cert: CertificateDer<'static>,
115    opt: Opt,
116) -> Result<ClientStats> {
117    let (endpoint, connection) = connect_client(server_addr, server_cert, opt).await?;
118
119    let start = Instant::now();
120
121    let connection = Arc::new(connection);
122
123    let mut stats = ClientStats::default();
124    let mut first_error = None;
125
126    let sem = Arc::new(Semaphore::new(opt.max_streams));
127    let results = Arc::new(Mutex::new(Vec::new()));
128    for _ in 0..opt.streams {
129        let permit = sem.clone().acquire_owned().await.unwrap();
130        let results = results.clone();
131        let connection = connection.clone();
132        tokio::spawn(async move {
133            let result =
134                handle_client_stream(connection, opt.upload_size, opt.read_unordered).await;
135            info!("stream finished: {:?}", result);
136            results.lock().unwrap().push(result);
137            drop(permit);
138        });
139    }
140
141    // Wait for remaining streams to finish
142    let _ = sem.acquire_many(opt.max_streams as u32).await.unwrap();
143
144    for result in results.lock().unwrap().drain(..) {
145        match result {
146            Ok((upload_result, download_result)) => {
147                stats.upload_stats.stream_finished(upload_result);
148                stats.download_stats.stream_finished(download_result);
149            }
150            Err(e) => {
151                if first_error.is_none() {
152                    first_error = Some(e);
153                }
154            }
155        }
156    }
157
158    stats.upload_stats.total_duration = start.elapsed();
159    stats.download_stats.total_duration = start.elapsed();
160
161    // Explicit close of the connection, since handles can still be around due
162    // to `Arc`ing them
163    connection.close(0u32.into(), b"Benchmark done");
164
165    endpoint.wait_all_draining().await;
166
167    if opt.stats {
168        println!("\nClient connection stats:\n{:#?}", connection.stats());
169    }
170
171    match first_error {
172        None => Ok(stats),
173        Some(e) => Err(e),
174    }
175}
176
177async fn handle_client_stream(
178    connection: Arc<noq::Connection>,
179    upload_size: u64,
180    read_unordered: bool,
181) -> Result<(TransferResult, TransferResult)> {
182    let start = Instant::now();
183
184    let (mut send_stream, recv_stream) = connection
185        .open_bi()
186        .await
187        .context("failed to open stream")?;
188
189    send_data_on_stream(&mut send_stream, upload_size).await?;
190
191    let upload_result = TransferResult::new(start.elapsed(), upload_size);
192
193    let start = Instant::now();
194    let size = drain_stream(recv_stream, read_unordered).await?;
195    let download_result = TransferResult::new(start.elapsed(), size as u64);
196
197    Ok((upload_result, download_result))
198}
199
200#[derive(Default)]
201struct ClientStats {
202    upload_stats: Stats,
203    download_stats: Stats,
204}
205
206impl ClientStats {
207    fn print(&self, client_id: usize) {
208        println!();
209        println!("Client {client_id} stats:");
210
211        if self.upload_stats.total_size != 0 {
212            self.upload_stats.print("upload");
213        }
214
215        if self.download_stats.total_size != 0 {
216            self.download_stats.print("download");
217        }
218    }
219}