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 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 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 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 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 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}