From 5da52c4d1e50dc65074accd91e475e070bada570 Mon Sep 17 00:00:00 2001 From: Ryan Fowler Date: Mon, 10 Aug 2026 12:43:30 -0400 Subject: [PATCH 1/3] ci: run integration tests in release mode --- .github/workflows/ci.yml | 4 +++- AGENTS.md | 2 +- 2 files changed, 4 insertions(+), 2 deletions(-) diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index f9b59a41..10a25b8d 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -70,4 +70,6 @@ jobs: run: choco install protoc -y - name: Rust integration tests against Rust binary - run: cargo test --locked --all-features --test cli --test formatting --test grpc --test har --test http --test install --test network --test terminal --test update --test websocket -- --test-threads=2 + # Each test case launches fetch as a child process. Use the optimized + # binary to avoid repeatedly loading the large debuggable Windows binary. + run: cargo test --release --locked --all-features --test cli --test formatting --test grpc --test har --test http --test install --test network --test terminal --test update --test websocket -- --test-threads=2 diff --git a/AGENTS.md b/AGENTS.md index 5fc80254..eb7d802b 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -31,7 +31,7 @@ Run the full CI-equivalent suite before PRs and for shared transport/request/res cargo fmt --check cargo clippy --locked --all-targets --all-features -- -D warnings cargo test --locked --all-features --lib --bins -cargo test --locked --all-features --test cli --test formatting --test grpc --test har --test http --test install --test network --test terminal --test update --test websocket -- --test-threads=2 +cargo test --release --locked --all-features --test cli --test formatting --test grpc --test har --test http --test install --test network --test terminal --test update --test websocket -- --test-threads=2 ``` For docs-only changes, skip Cargo unless examples or generated CLI output changed: From b7fce7fcb840f148326454e2f262c284d0232d07 Mon Sep 17 00:00:00 2001 From: Ryan Fowler Date: Mon, 10 Aug 2026 12:50:48 -0400 Subject: [PATCH 2/3] ci: limit release integration tests to Windows --- .github/workflows/ci.yml | 5 +++++ AGENTS.md | 3 ++- 2 files changed, 7 insertions(+), 1 deletion(-) diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 10a25b8d..356b94cf 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -70,6 +70,11 @@ jobs: run: choco install protoc -y - name: Rust integration tests against Rust binary + if: runner.os != 'Windows' + run: cargo test --locked --all-features --test cli --test formatting --test grpc --test har --test http --test install --test network --test terminal --test update --test websocket -- --test-threads=2 + + - name: Rust integration tests against optimized Windows binary + if: runner.os == 'Windows' # Each test case launches fetch as a child process. Use the optimized # binary to avoid repeatedly loading the large debuggable Windows binary. run: cargo test --release --locked --all-features --test cli --test formatting --test grpc --test har --test http --test install --test network --test terminal --test update --test websocket -- --test-threads=2 diff --git a/AGENTS.md b/AGENTS.md index eb7d802b..f547af22 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -31,7 +31,8 @@ Run the full CI-equivalent suite before PRs and for shared transport/request/res cargo fmt --check cargo clippy --locked --all-targets --all-features -- -D warnings cargo test --locked --all-features --lib --bins -cargo test --release --locked --all-features --test cli --test formatting --test grpc --test har --test http --test install --test network --test terminal --test update --test websocket -- --test-threads=2 +cargo test --locked --all-features --test cli --test formatting --test grpc --test har --test http --test install --test network --test terminal --test update --test websocket -- --test-threads=2 +# Windows CI runs this suite with --release to reduce child-process startup cost. ``` For docs-only changes, skip Cargo unless examples or generated CLI output changed: From 609b91093a39574879e27b9672b08ac518abb00a Mon Sep 17 00:00:00 2001 From: Ryan Fowler Date: Mon, 10 Aug 2026 13:02:49 -0400 Subject: [PATCH 3/3] test: remove local server polling delays --- .github/workflows/ci.yml | 7 - AGENTS.md | 1 - tests/support/http.rs | 57 +++---- tests/support/proxy.rs | 49 +++--- tests/support/tls.rs | 190 ++++++++++----------- tests/support/websocket.rs | 329 ++++++++++++++++--------------------- tests/websocket.rs | 146 ++++++++-------- 7 files changed, 342 insertions(+), 437 deletions(-) diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 356b94cf..f9b59a41 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -70,11 +70,4 @@ jobs: run: choco install protoc -y - name: Rust integration tests against Rust binary - if: runner.os != 'Windows' run: cargo test --locked --all-features --test cli --test formatting --test grpc --test har --test http --test install --test network --test terminal --test update --test websocket -- --test-threads=2 - - - name: Rust integration tests against optimized Windows binary - if: runner.os == 'Windows' - # Each test case launches fetch as a child process. Use the optimized - # binary to avoid repeatedly loading the large debuggable Windows binary. - run: cargo test --release --locked --all-features --test cli --test formatting --test grpc --test har --test http --test install --test network --test terminal --test update --test websocket -- --test-threads=2 diff --git a/AGENTS.md b/AGENTS.md index f547af22..5fc80254 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -32,7 +32,6 @@ cargo fmt --check cargo clippy --locked --all-targets --all-features -- -D warnings cargo test --locked --all-features --lib --bins cargo test --locked --all-features --test cli --test formatting --test grpc --test har --test http --test install --test network --test terminal --test update --test websocket -- --test-threads=2 -# Windows CI runs this suite with --release to reduce child-process startup cost. ``` For docs-only changes, skip Cargo unless examples or generated CLI output changed: diff --git a/tests/support/http.rs b/tests/support/http.rs index 9593fd5f..599aaef8 100644 --- a/tests/support/http.rs +++ b/tests/support/http.rs @@ -1,6 +1,6 @@ use std::collections::HashMap; use std::io::{BufRead, BufReader, Read, Write}; -use std::net::{Shutdown, TcpListener}; +use std::net::{Shutdown, TcpListener, TcpStream}; use std::sync::{Arc, Mutex, mpsc}; use std::thread; use std::time::{Duration, Instant}; @@ -87,6 +87,7 @@ pub(crate) struct TestServer { pub(crate) requests: Arc>>, request_notify: mpsc::Receiver<()>, shutdown: Option>, + shutdown_addr: String, join: Option>, } @@ -101,46 +102,38 @@ impl TestServer { handler: impl Fn(TestRequest) -> TestResponse + Send + Sync + 'static, ) -> Self { let listener = TcpListener::bind("127.0.0.1:0").expect("bind test server"); - listener - .set_nonblocking(true) - .expect("set test listener nonblocking"); - let url = format!("http://{}", listener.local_addr().expect("local addr")); + let addr = listener.local_addr().expect("local addr"); + let url = format!("http://{addr}"); let requests = Arc::new(Mutex::new(Vec::new())); let handler = Arc::new(handler); let (shutdown_tx, shutdown_rx) = mpsc::channel(); let (notify_tx, notify_rx) = mpsc::channel(); let request_log = Arc::clone(&requests); let join = thread::spawn(move || { - loop { + for stream in listener.incoming() { if shutdown_rx.try_recv().is_ok() { break; } - match listener.accept() { - Ok((stream, _)) => { - let _ = stream.set_nonblocking(false); - let handler = Arc::clone(&handler); - let request_log = Arc::clone(&request_log); - let notify = notify_tx.clone(); - thread::spawn(move || { - let mut writer = stream.try_clone().expect("clone response stream"); - let mut reader = BufReader::new(stream); - while let Some(req) = read_request(&mut reader) { - let close = req.header("connection").eq_ignore_ascii_case("close"); - request_log.lock().unwrap().push(req.clone()); - let _ = notify.send(()); - let resp = handler(req); - write_response(&mut writer, resp); - if close { - break; - } - } - }); - } - Err(err) if err.kind() == std::io::ErrorKind::WouldBlock => { - thread::sleep(Duration::from_millis(5)); + let Ok(stream) = stream else { + break; + }; + let handler = Arc::clone(&handler); + let request_log = Arc::clone(&request_log); + let notify = notify_tx.clone(); + thread::spawn(move || { + let mut writer = stream.try_clone().expect("clone response stream"); + let mut reader = BufReader::new(stream); + while let Some(req) = read_request(&mut reader) { + let close = req.header("connection").eq_ignore_ascii_case("close"); + request_log.lock().unwrap().push(req.clone()); + let _ = notify.send(()); + let resp = handler(req); + write_response(&mut writer, resp); + if close { + break; + } } - Err(_) => break, - } + }); } }); Self { @@ -148,6 +141,7 @@ impl TestServer { requests, request_notify: notify_rx, shutdown: Some(shutdown_tx), + shutdown_addr: addr.to_string(), join: Some(join), } } @@ -161,6 +155,7 @@ impl Drop for TestServer { fn drop(&mut self) { if let Some(tx) = self.shutdown.take() { let _ = tx.send(()); + let _ = TcpStream::connect(&self.shutdown_addr); } if let Some(join) = self.join.take() { let _ = join.join(); diff --git a/tests/support/proxy.rs b/tests/support/proxy.rs index 7e6b6785..b3516a32 100644 --- a/tests/support/proxy.rs +++ b/tests/support/proxy.rs @@ -17,6 +17,7 @@ pub(crate) struct HttpsProxyTestServer { pub(crate) client_key_path: PathBuf, pub(crate) requests: Arc>>, pub(crate) shutdown: Option>, + shutdown_addr: String, pub(crate) join: Option>, } @@ -30,6 +31,7 @@ impl Drop for HttpsProxyTestServer { fn drop(&mut self) { if let Some(tx) = self.shutdown.take() { let _ = tx.send(()); + let _ = TcpStream::connect(&self.shutdown_addr); } if let Some(join) = self.join.take() { let _ = join.join(); @@ -232,41 +234,35 @@ pub(crate) fn start_https_proxy(require_client_auth: bool) -> HttpsProxyTestServ let config = Arc::new(config); let listener = TcpListener::bind("127.0.0.1:0").expect("bind HTTPS proxy"); - listener.set_nonblocking(true).unwrap(); let port = listener.local_addr().unwrap().port(); + let shutdown_addr = format!("127.0.0.1:{port}"); let url = format!("https://localhost:{port}"); let requests = Arc::new(Mutex::new(Vec::new())); let requests_for_thread = Arc::clone(&requests); let (tx, rx) = mpsc::channel(); let join = thread::spawn(move || { - loop { + for stream in listener.incoming() { if rx.try_recv().is_ok() { break; } - match listener.accept() { - Ok((stream, _)) => { - let _ = stream.set_nonblocking(false); - let config = Arc::clone(&config); - let requests = Arc::clone(&requests_for_thread); - thread::spawn(move || { - let Ok(conn) = rustls::ServerConnection::new(config) else { - return; - }; - let mut tls = rustls::StreamOwned::new(conn, stream); - let mut reader = BufReader::new(&mut tls); - let Some(req) = read_request(&mut reader) else { - return; - }; - requests.lock().unwrap().push(req); - let tls = reader.into_inner(); - write_response(tls, TestResponse::ok("proxied")); - }); - } - Err(err) if err.kind() == std::io::ErrorKind::WouldBlock => { - thread::sleep(Duration::from_millis(5)); - } - Err(_) => break, - } + let Ok(stream) = stream else { + break; + }; + let config = Arc::clone(&config); + let requests = Arc::clone(&requests_for_thread); + thread::spawn(move || { + let Ok(conn) = rustls::ServerConnection::new(config) else { + return; + }; + let mut tls = rustls::StreamOwned::new(conn, stream); + let mut reader = BufReader::new(&mut tls); + let Some(req) = read_request(&mut reader) else { + return; + }; + requests.lock().unwrap().push(req); + let tls = reader.into_inner(); + write_response(tls, TestResponse::ok("proxied")); + }); } }); @@ -277,6 +273,7 @@ pub(crate) fn start_https_proxy(require_client_auth: bool) -> HttpsProxyTestServ client_key_path, requests, shutdown: Some(tx), + shutdown_addr, join: Some(join), } } diff --git a/tests/support/tls.rs b/tests/support/tls.rs index 939b1c33..003d5141 100644 --- a/tests/support/tls.rs +++ b/tests/support/tls.rs @@ -18,6 +18,7 @@ pub(crate) struct TlsTestServer { pub(crate) url: String, pub(crate) ca_cert_path: PathBuf, pub(crate) shutdown: Option>, + pub(crate) shutdown_addr: String, pub(crate) join: Option>, } @@ -28,6 +29,7 @@ pub(crate) struct MtlsTestServer { pub(crate) client_key_path: PathBuf, pub(crate) client_combined_path: PathBuf, pub(crate) shutdown: Option>, + pub(crate) shutdown_addr: String, pub(crate) join: Option>, } @@ -35,6 +37,7 @@ impl Drop for TlsTestServer { fn drop(&mut self) { if let Some(tx) = self.shutdown.take() { let _ = tx.send(()); + let _ = std::net::TcpStream::connect(&self.shutdown_addr); } if let Some(join) = self.join.take() { let _ = join.join(); @@ -46,6 +49,7 @@ impl Drop for MtlsTestServer { fn drop(&mut self) { if let Some(tx) = self.shutdown.take() { let _ = tx.send(()); + let _ = std::net::TcpStream::connect(&self.shutdown_addr); } if let Some(join) = self.join.take() { let _ = join.join(); @@ -73,46 +77,41 @@ pub(crate) fn start_tls_server( .unwrap(); let config = Arc::new(config); let listener = TcpListener::bind("127.0.0.1:0").expect("bind tls server"); - listener.set_nonblocking(true).unwrap(); let port = listener.local_addr().unwrap().port(); + let shutdown_addr = format!("127.0.0.1:{port}"); let url = format!("https://localhost:{port}"); let handler = Arc::new(handler); let (tx, rx) = mpsc::channel(); let join = thread::spawn(move || { - loop { + for stream in listener.incoming() { if rx.try_recv().is_ok() { break; } - match listener.accept() { - Ok((stream, _)) => { - let _ = stream.set_nonblocking(false); - let config = Arc::clone(&config); - let handler = Arc::clone(&handler); - thread::spawn(move || { - let Ok(conn) = rustls::ServerConnection::new(config) else { - return; - }; - let mut tls = rustls::StreamOwned::new(conn, stream); - let mut reader = BufReader::new(&mut tls); - let Some(req) = read_request(&mut reader) else { - return; - }; - let resp = handler(req); - let tls = reader.into_inner(); - write_response(tls, resp); - }); - } - Err(err) if err.kind() == std::io::ErrorKind::WouldBlock => { - thread::sleep(Duration::from_millis(5)); - } - Err(_) => break, - } + let Ok(stream) = stream else { + break; + }; + let config = Arc::clone(&config); + let handler = Arc::clone(&handler); + thread::spawn(move || { + let Ok(conn) = rustls::ServerConnection::new(config) else { + return; + }; + let mut tls = rustls::StreamOwned::new(conn, stream); + let mut reader = BufReader::new(&mut tls); + let Some(req) = read_request(&mut reader) else { + return; + }; + let resp = handler(req); + let tls = reader.into_inner(); + write_response(tls, resp); + }); } }); TlsTestServer { url, ca_cert_path, shutdown: Some(tx), + shutdown_addr, join: Some(join), } } @@ -178,42 +177,38 @@ pub(crate) fn start_invalid_certificate_verify_server() -> TlsTestServer { .with_cert_resolver(Arc::new(SingleCertAndKey::from(certified_key))); let config = Arc::new(config); let listener = TcpListener::bind("127.0.0.1:0").expect("bind invalid signature TLS server"); - listener.set_nonblocking(true).unwrap(); let port = listener.local_addr().unwrap().port(); + let shutdown_addr = format!("127.0.0.1:{port}"); let url = format!("https://localhost:{port}"); let (tx, rx) = mpsc::channel(); let join = thread::spawn(move || { - loop { + for stream in listener.incoming() { if rx.try_recv().is_ok() { break; } - match listener.accept() { - Ok((stream, _)) => { - let config = Arc::clone(&config); - thread::spawn(move || { - let Ok(conn) = rustls::ServerConnection::new(config) else { - return; - }; - let mut tls = rustls::StreamOwned::new(conn, stream); - let mut reader = BufReader::new(&mut tls); - let Some(_) = read_request(&mut reader) else { - return; - }; - let tls = reader.into_inner(); - write_response(tls, TestResponse::ok("invalid signature accepted")); - }); - } - Err(err) if err.kind() == std::io::ErrorKind::WouldBlock => { - thread::sleep(Duration::from_millis(5)); - } - Err(_) => break, - } + let Ok(stream) = stream else { + break; + }; + let config = Arc::clone(&config); + thread::spawn(move || { + let Ok(conn) = rustls::ServerConnection::new(config) else { + return; + }; + let mut tls = rustls::StreamOwned::new(conn, stream); + let mut reader = BufReader::new(&mut tls); + let Some(_) = read_request(&mut reader) else { + return; + }; + let tls = reader.into_inner(); + write_response(tls, TestResponse::ok("invalid signature accepted")); + }); } }); TlsTestServer { url, ca_cert_path, shutdown: Some(tx), + shutdown_addr, join: Some(join), } } @@ -246,8 +241,8 @@ pub(crate) fn start_h2_tls_server_with_accept_delay( config.alpn_protocols = vec![b"h2".to_vec()]; let acceptor = tokio_rustls::TlsAcceptor::from(Arc::new(config)); let listener = TcpListener::bind("127.0.0.1:0").expect("bind h2 tls server"); - listener.set_nonblocking(true).unwrap(); let port = listener.local_addr().unwrap().port(); + let shutdown_addr = format!("127.0.0.1:{port}"); let url = format!("https://localhost:{port}"); let handler: Arc TestResponse + Send + Sync> = Arc::new(handler); let (tx, rx) = mpsc::channel(); @@ -256,39 +251,35 @@ pub(crate) fn start_h2_tls_server_with_accept_delay( .enable_all() .build() .unwrap(); - loop { + for stream in listener.incoming() { if rx.try_recv().is_ok() { break; } - match listener.accept() { - Ok((stream, _)) => { - let _ = stream.set_nonblocking(true); - let acceptor = acceptor.clone(); - let handler = Arc::clone(&handler); - runtime.block_on(async move { - let Ok(stream) = tokio::net::TcpStream::from_std(stream) else { - return; - }; - if !accept_delay.is_zero() { - tokio::time::sleep(accept_delay).await; - } - let Ok(tls) = acceptor.accept(stream).await else { - return; - }; - serve_test_h2_connection(tls, handler).await; - }); - } - Err(err) if err.kind() == std::io::ErrorKind::WouldBlock => { - thread::sleep(Duration::from_millis(5)); + let Ok(stream) = stream else { + break; + }; + let _ = stream.set_nonblocking(true); + let acceptor = acceptor.clone(); + let handler = Arc::clone(&handler); + runtime.block_on(async move { + let Ok(stream) = tokio::net::TcpStream::from_std(stream) else { + return; + }; + if !accept_delay.is_zero() { + tokio::time::sleep(accept_delay).await; } - Err(_) => break, - } + let Ok(tls) = acceptor.accept(stream).await else { + return; + }; + serve_test_h2_connection(tls, handler).await; + }); } }); TlsTestServer { url, ca_cert_path, shutdown: Some(tx), + shutdown_addr, join: Some(join), } } @@ -451,42 +442,36 @@ pub(crate) fn start_mtls_server() -> MtlsTestServer { .unwrap(); let config = Arc::new(config); let listener = TcpListener::bind("127.0.0.1:0").expect("bind mtls server"); - listener.set_nonblocking(true).unwrap(); let port = listener.local_addr().unwrap().port(); + let shutdown_addr = format!("127.0.0.1:{port}"); let url = format!("https://localhost:{port}"); let (tx, rx) = mpsc::channel(); let join = thread::spawn(move || { - loop { + for stream in listener.incoming() { if rx.try_recv().is_ok() { break; } - match listener.accept() { - Ok((stream, _)) => { - let _ = stream.set_nonblocking(false); - let config = Arc::clone(&config); - thread::spawn(move || { - let Ok(conn) = rustls::ServerConnection::new(config) else { - return; - }; - let mut tls = rustls::StreamOwned::new(conn, stream); - let mut reader = BufReader::new(&mut tls); - let Some(req) = read_request(&mut reader) else { - return; - }; - let resp = if req.path == "/" { - TestResponse::ok("mtls-success") - } else { - TestResponse::status(404, "Not Found", "") - }; - let tls = reader.into_inner(); - write_response(tls, resp); - }); - } - Err(err) if err.kind() == std::io::ErrorKind::WouldBlock => { - thread::sleep(Duration::from_millis(5)); - } - Err(_) => break, - } + let Ok(stream) = stream else { + break; + }; + let config = Arc::clone(&config); + thread::spawn(move || { + let Ok(conn) = rustls::ServerConnection::new(config) else { + return; + }; + let mut tls = rustls::StreamOwned::new(conn, stream); + let mut reader = BufReader::new(&mut tls); + let Some(req) = read_request(&mut reader) else { + return; + }; + let resp = if req.path == "/" { + TestResponse::ok("mtls-success") + } else { + TestResponse::status(404, "Not Found", "") + }; + let tls = reader.into_inner(); + write_response(tls, resp); + }); } }); @@ -497,6 +482,7 @@ pub(crate) fn start_mtls_server() -> MtlsTestServer { client_key_path, client_combined_path, shutdown: Some(tx), + shutdown_addr, join: Some(join), } } diff --git a/tests/support/websocket.rs b/tests/support/websocket.rs index 65fd544b..9e981f99 100644 --- a/tests/support/websocket.rs +++ b/tests/support/websocket.rs @@ -16,58 +16,47 @@ pub(crate) fn start_ws_echo_server( validate: impl Fn(&TestRequest) -> Result<(), String> + Send + Sync + 'static, ) -> (String, mpsc::Receiver) { let listener = TcpListener::bind("127.0.0.1:0").expect("bind websocket server"); - listener.set_nonblocking(true).unwrap(); let url = format!("ws://{}", listener.local_addr().unwrap()); let validate = Arc::new(validate); let (seen_tx, seen_rx) = mpsc::channel(); thread::spawn(move || { - loop { - match listener.accept() { - Ok((mut stream, _)) => { - let _ = stream.set_nonblocking(false); - let validate = Arc::clone(&validate); - let seen_tx = seen_tx.clone(); - thread::spawn(move || { - let _ = stream.set_read_timeout(Some(Duration::from_secs(2))); - let mut reader = BufReader::new(stream.try_clone().unwrap()); - let Some(req) = read_request(&mut reader) else { - return; - }; - if let Err(err) = validate(&req) { - write_response( - &mut stream, - TestResponse::status(400, "Bad Request", err), - ); - return; - } - let key = req.header("sec-websocket-key"); - let mut sha = Sha1::new(); - sha.update(key.as_bytes()); - sha.update(b"258EAFA5-E914-47DA-95CA-C5AB0DC85B11"); - let accept = - base64::engine::general_purpose::STANDARD.encode(sha.finalize()); - let response = format!( - "HTTP/1.1 101 Switching Protocols\r\nUpgrade: websocket\r\nConnection: Upgrade\r\nSec-WebSocket-Accept: {accept}\r\n\r\n" - ); - if stream.write_all(response.as_bytes()).is_err() { - return; - } - let msg = read_ws_text(&mut stream); - let _ = seen_tx.send(msg.clone()); - let reply = if msg.trim().starts_with('{') { - msg - } else { - format!("echo: {msg}") - }; - let _ = stream.write_all(&ws_text_frame(reply.as_bytes())); - write_ws_close_and_drain(&mut stream, b"done"); - }); + for stream in listener.incoming() { + let Ok(mut stream) = stream else { + break; + }; + let validate = Arc::clone(&validate); + let seen_tx = seen_tx.clone(); + thread::spawn(move || { + let _ = stream.set_read_timeout(Some(Duration::from_secs(2))); + let mut reader = BufReader::new(stream.try_clone().unwrap()); + let Some(req) = read_request(&mut reader) else { + return; + }; + if let Err(err) = validate(&req) { + write_response(&mut stream, TestResponse::status(400, "Bad Request", err)); + return; } - Err(err) if err.kind() == std::io::ErrorKind::WouldBlock => { - thread::sleep(Duration::from_millis(5)); + let key = req.header("sec-websocket-key"); + let mut sha = Sha1::new(); + sha.update(key.as_bytes()); + sha.update(b"258EAFA5-E914-47DA-95CA-C5AB0DC85B11"); + let accept = base64::engine::general_purpose::STANDARD.encode(sha.finalize()); + let response = format!( + "HTTP/1.1 101 Switching Protocols\r\nUpgrade: websocket\r\nConnection: Upgrade\r\nSec-WebSocket-Accept: {accept}\r\n\r\n" + ); + if stream.write_all(response.as_bytes()).is_err() { + return; } - Err(_) => break, - } + let msg = read_ws_text(&mut stream); + let _ = seen_tx.send(msg.clone()); + let reply = if msg.trim().starts_with('{') { + msg + } else { + format!("echo: {msg}") + }; + let _ = stream.write_all(&ws_text_frame(reply.as_bytes())); + write_ws_close_and_drain(&mut stream, b"done"); + }); } }); (url, seen_rx) @@ -93,69 +82,62 @@ pub(crate) fn start_wss_echo_server( .unwrap(); let config = Arc::new(config); let listener = TcpListener::bind("127.0.0.1:0").expect("bind wss websocket server"); - listener.set_nonblocking(true).unwrap(); let port = listener.local_addr().unwrap().port(); + let shutdown_addr = format!("127.0.0.1:{port}"); let url = format!("wss://localhost:{port}"); let validate = Arc::new(validate); let (seen_tx, seen_rx) = mpsc::channel(); let (shutdown_tx, shutdown_rx) = mpsc::channel(); let join = thread::spawn(move || { - loop { + for stream in listener.incoming() { if shutdown_rx.try_recv().is_ok() { break; } - match listener.accept() { - Ok((stream, _)) => { - let _ = stream.set_nonblocking(false); - let config = Arc::clone(&config); - let validate = Arc::clone(&validate); - let seen_tx = seen_tx.clone(); - thread::spawn(move || { - let Ok(conn) = rustls::ServerConnection::new(config) else { - return; - }; - let mut tls = rustls::StreamOwned::new(conn, stream); - let _ = tls.sock.set_read_timeout(Some(Duration::from_secs(2))); - let mut reader = BufReader::new(&mut tls); - let Some(req) = read_request(&mut reader) else { - return; - }; - if let Err(err) = validate(&req) { - write_response( - reader.into_inner(), - TestResponse::status(400, "Bad Request", err), - ); - return; - } - let key = req.header("sec-websocket-key"); - let mut sha = Sha1::new(); - sha.update(key.as_bytes()); - sha.update(b"258EAFA5-E914-47DA-95CA-C5AB0DC85B11"); - let accept = - base64::engine::general_purpose::STANDARD.encode(sha.finalize()); - let response = format!( - "HTTP/1.1 101 Switching Protocols\r\nUpgrade: websocket\r\nConnection: Upgrade\r\nSec-WebSocket-Accept: {accept}\r\n\r\n" - ); - let tls = reader.into_inner(); - if tls.write_all(response.as_bytes()).is_err() { - return; - } - let msg = read_ws_text(tls); - let _ = seen_tx.send(msg.clone()); - let reply = if msg.trim().starts_with('{') { - msg - } else { - format!("echo: {msg}") - }; - let _ = tls.write_all(&ws_text_frame(reply.as_bytes())); - write_ws_close_and_drain(tls, b"done"); - }); + let Ok(stream) = stream else { + break; + }; + let config = Arc::clone(&config); + let validate = Arc::clone(&validate); + let seen_tx = seen_tx.clone(); + thread::spawn(move || { + let Ok(conn) = rustls::ServerConnection::new(config) else { + return; + }; + let mut tls = rustls::StreamOwned::new(conn, stream); + let _ = tls.sock.set_read_timeout(Some(Duration::from_secs(2))); + let mut reader = BufReader::new(&mut tls); + let Some(req) = read_request(&mut reader) else { + return; + }; + if let Err(err) = validate(&req) { + write_response( + reader.into_inner(), + TestResponse::status(400, "Bad Request", err), + ); + return; } - Err(err) if err.kind() == std::io::ErrorKind::WouldBlock => { - thread::sleep(Duration::from_millis(5)); + let key = req.header("sec-websocket-key"); + let mut sha = Sha1::new(); + sha.update(key.as_bytes()); + sha.update(b"258EAFA5-E914-47DA-95CA-C5AB0DC85B11"); + let accept = base64::engine::general_purpose::STANDARD.encode(sha.finalize()); + let response = format!( + "HTTP/1.1 101 Switching Protocols\r\nUpgrade: websocket\r\nConnection: Upgrade\r\nSec-WebSocket-Accept: {accept}\r\n\r\n" + ); + let tls = reader.into_inner(); + if tls.write_all(response.as_bytes()).is_err() { + return; } - Err(_) => break, - } + let msg = read_ws_text(tls); + let _ = seen_tx.send(msg.clone()); + let reply = if msg.trim().starts_with('{') { + msg + } else { + format!("echo: {msg}") + }; + let _ = tls.write_all(&ws_text_frame(reply.as_bytes())); + write_ws_close_and_drain(tls, b"done"); + }); } }); @@ -164,6 +146,7 @@ pub(crate) fn start_wss_echo_server( url, ca_cert_path, shutdown: Some(shutdown_tx), + shutdown_addr, join: Some(join), }, seen_rx, @@ -172,45 +155,37 @@ pub(crate) fn start_wss_echo_server( pub(crate) fn start_ws_multi_echo_server(messages: usize) -> String { let listener = TcpListener::bind("127.0.0.1:0").expect("bind websocket server"); - listener.set_nonblocking(true).unwrap(); let url = format!("ws://{}", listener.local_addr().unwrap()); thread::spawn(move || { - loop { - match listener.accept() { - Ok((mut stream, _)) => { - let _ = stream.set_nonblocking(false); - thread::spawn(move || { - let _ = stream.set_read_timeout(Some(Duration::from_secs(2))); - let mut reader = BufReader::new(stream.try_clone().unwrap()); - let Some(req) = read_request(&mut reader) else { - return; - }; - let key = req.header("sec-websocket-key"); - let mut sha = Sha1::new(); - sha.update(key.as_bytes()); - sha.update(b"258EAFA5-E914-47DA-95CA-C5AB0DC85B11"); - let accept = - base64::engine::general_purpose::STANDARD.encode(sha.finalize()); - let response = format!( - "HTTP/1.1 101 Switching Protocols\r\nUpgrade: websocket\r\nConnection: Upgrade\r\nSec-WebSocket-Accept: {accept}\r\n\r\n" - ); - if stream.write_all(response.as_bytes()).is_err() { - return; - } - for _ in 0..messages { - let Some(msg) = read_ws_text_frame(&mut stream) else { - return; - }; - let _ = stream.write_all(&ws_text_frame(msg.as_bytes())); - } - write_ws_close_and_drain(&mut stream, b"done"); - }); + for stream in listener.incoming() { + let Ok(mut stream) = stream else { + break; + }; + thread::spawn(move || { + let _ = stream.set_read_timeout(Some(Duration::from_secs(2))); + let mut reader = BufReader::new(stream.try_clone().unwrap()); + let Some(req) = read_request(&mut reader) else { + return; + }; + let key = req.header("sec-websocket-key"); + let mut sha = Sha1::new(); + sha.update(key.as_bytes()); + sha.update(b"258EAFA5-E914-47DA-95CA-C5AB0DC85B11"); + let accept = base64::engine::general_purpose::STANDARD.encode(sha.finalize()); + let response = format!( + "HTTP/1.1 101 Switching Protocols\r\nUpgrade: websocket\r\nConnection: Upgrade\r\nSec-WebSocket-Accept: {accept}\r\n\r\n" + ); + if stream.write_all(response.as_bytes()).is_err() { + return; } - Err(err) if err.kind() == std::io::ErrorKind::WouldBlock => { - thread::sleep(Duration::from_millis(5)); + for _ in 0..messages { + let Some(msg) = read_ws_text_frame(&mut stream) else { + return; + }; + let _ = stream.write_all(&ws_text_frame(msg.as_bytes())); } - Err(_) => break, - } + write_ws_close_and_drain(&mut stream, b"done"); + }); } }); url @@ -220,51 +195,40 @@ pub(crate) fn start_ws_push_server( validate: impl Fn(&TestRequest) -> Result<(), String> + Send + Sync + 'static, ) -> String { let listener = TcpListener::bind("127.0.0.1:0").expect("bind websocket push server"); - listener.set_nonblocking(true).unwrap(); let url = format!("ws://{}", listener.local_addr().unwrap()); let validate = Arc::new(validate); thread::spawn(move || { - loop { - match listener.accept() { - Ok((mut stream, _)) => { - let _ = stream.set_nonblocking(false); - let validate = Arc::clone(&validate); - thread::spawn(move || { - let _ = stream.set_read_timeout(Some(Duration::from_secs(2))); - let mut reader = BufReader::new(stream.try_clone().unwrap()); - let Some(req) = read_request(&mut reader) else { - return; - }; - if let Err(err) = validate(&req) { - write_response( - &mut stream, - TestResponse::status(400, "Bad Request", err), - ); - return; - } - let key = req.header("sec-websocket-key"); - let mut sha = Sha1::new(); - sha.update(key.as_bytes()); - sha.update(b"258EAFA5-E914-47DA-95CA-C5AB0DC85B11"); - let accept = - base64::engine::general_purpose::STANDARD.encode(sha.finalize()); - let response = format!( - "HTTP/1.1 101 Switching Protocols\r\nUpgrade: websocket\r\nConnection: Upgrade\r\nSec-WebSocket-Accept: {accept}\r\n\r\n" - ); - if stream.write_all(response.as_bytes()).is_err() { - return; - } - let _ = stream.write_all(&ws_text_frame(br#"{"hello":"websocket"}"#)); - let _ = stream.write_all(&ws_binary_frame(b"\x00\x01\x02\x03")); - let _ = stream.write_all(&ws_text_frame(b"plain text")); - write_ws_close_and_drain(&mut stream, b"done"); - }); + for stream in listener.incoming() { + let Ok(mut stream) = stream else { + break; + }; + let validate = Arc::clone(&validate); + thread::spawn(move || { + let _ = stream.set_read_timeout(Some(Duration::from_secs(2))); + let mut reader = BufReader::new(stream.try_clone().unwrap()); + let Some(req) = read_request(&mut reader) else { + return; + }; + if let Err(err) = validate(&req) { + write_response(&mut stream, TestResponse::status(400, "Bad Request", err)); + return; } - Err(err) if err.kind() == std::io::ErrorKind::WouldBlock => { - thread::sleep(Duration::from_millis(5)); + let key = req.header("sec-websocket-key"); + let mut sha = Sha1::new(); + sha.update(key.as_bytes()); + sha.update(b"258EAFA5-E914-47DA-95CA-C5AB0DC85B11"); + let accept = base64::engine::general_purpose::STANDARD.encode(sha.finalize()); + let response = format!( + "HTTP/1.1 101 Switching Protocols\r\nUpgrade: websocket\r\nConnection: Upgrade\r\nSec-WebSocket-Accept: {accept}\r\n\r\n" + ); + if stream.write_all(response.as_bytes()).is_err() { + return; } - Err(_) => break, - } + let _ = stream.write_all(&ws_text_frame(br#"{"hello":"websocket"}"#)); + let _ = stream.write_all(&ws_binary_frame(b"\x00\x01\x02\x03")); + let _ = stream.write_all(&ws_text_frame(b"plain text")); + write_ws_close_and_drain(&mut stream, b"done"); + }); } }); url @@ -274,21 +238,13 @@ pub(crate) fn start_ws_hold_open_push_server( message: impl Into>, ) -> (String, mpsc::Sender<()>, thread::JoinHandle<()>) { let listener = TcpListener::bind("127.0.0.1:0").expect("bind websocket hold-open server"); - listener.set_nonblocking(true).unwrap(); let url = format!("ws://{}", listener.local_addr().unwrap()); let message = message.into(); let (shutdown_tx, shutdown_rx) = mpsc::channel(); let join = thread::spawn(move || { - let mut stream = loop { - match listener.accept() { - Ok((stream, _)) => break stream, - Err(err) if err.kind() == std::io::ErrorKind::WouldBlock => { - thread::sleep(Duration::from_millis(5)); - } - Err(_) => return, - } + let Ok((mut stream, _)) = listener.accept() else { + return; }; - let _ = stream.set_nonblocking(false); let _ = stream.set_read_timeout(Some(Duration::from_secs(2))); let mut reader = BufReader::new(stream.try_clone().unwrap()); let Some(req) = read_request(&mut reader) else { @@ -312,12 +268,7 @@ pub(crate) fn start_ws_hold_open_push_server( { return; } - loop { - match shutdown_rx.try_recv() { - Ok(()) | Err(mpsc::TryRecvError::Disconnected) => break, - Err(mpsc::TryRecvError::Empty) => thread::sleep(Duration::from_millis(10)), - } - } + let _ = shutdown_rx.recv(); let _ = stream.write_all(&ws_close_frame(b"done")); }); (url, shutdown_tx, join) diff --git a/tests/websocket.rs b/tests/websocket.rs index 8ebbb3a0..7fc2b4df 100644 --- a/tests/websocket.rs +++ b/tests/websocket.rs @@ -36,52 +36,44 @@ const WEBSOCKET_RECEIVE_LIMIT_BYTES: usize = 16 * 1024 * 1024; fn start_ws_frame_server(reply: Vec) -> (String, mpsc::Receiver<(u8, Vec)>) { let listener = TcpListener::bind("127.0.0.1:0").expect("bind websocket frame server"); - listener.set_nonblocking(true).unwrap(); let url = format!("ws://{}", listener.local_addr().unwrap()); let (seen_tx, seen_rx) = mpsc::channel(); thread::spawn(move || { - loop { - match listener.accept() { - Ok((mut stream, _)) => { - let _ = stream.set_nonblocking(false); - let reply = reply.clone(); - let seen_tx = seen_tx.clone(); - thread::spawn(move || { - let _ = stream.set_read_timeout(Some(Duration::from_secs(2))); - let mut reader = BufReader::new(&mut stream); - let Some(req) = read_request(&mut reader) else { - return; - }; - let key = req.header("sec-websocket-key"); - let mut sha = Sha1::new(); - sha.update(key.as_bytes()); - sha.update(b"258EAFA5-E914-47DA-95CA-C5AB0DC85B11"); - let accept = - base64::engine::general_purpose::STANDARD.encode(sha.finalize()); - let response = format!( - "HTTP/1.1 101 Switching Protocols\r\nUpgrade: websocket\r\nConnection: Upgrade\r\nSec-WebSocket-Accept: {accept}\r\n\r\n" - ); - if reader - .get_mut() - .write_all(response.as_bytes()) - .and_then(|()| reader.get_mut().flush()) - .is_err() - { - return; - } - if let Some(frame) = read_ws_frame(&mut reader) { - let _ = seen_tx.send(frame); - } - let stream = reader.into_inner(); - let _ = stream.write_all(&ws_binary_frame(&reply)); - write_ws_close_and_drain(stream, b"done"); - }); + for stream in listener.incoming() { + let Ok(mut stream) = stream else { + break; + }; + let reply = reply.clone(); + let seen_tx = seen_tx.clone(); + thread::spawn(move || { + let _ = stream.set_read_timeout(Some(Duration::from_secs(2))); + let mut reader = BufReader::new(&mut stream); + let Some(req) = read_request(&mut reader) else { + return; + }; + let key = req.header("sec-websocket-key"); + let mut sha = Sha1::new(); + sha.update(key.as_bytes()); + sha.update(b"258EAFA5-E914-47DA-95CA-C5AB0DC85B11"); + let accept = base64::engine::general_purpose::STANDARD.encode(sha.finalize()); + let response = format!( + "HTTP/1.1 101 Switching Protocols\r\nUpgrade: websocket\r\nConnection: Upgrade\r\nSec-WebSocket-Accept: {accept}\r\n\r\n" + ); + if reader + .get_mut() + .write_all(response.as_bytes()) + .and_then(|()| reader.get_mut().flush()) + .is_err() + { + return; } - Err(err) if err.kind() == std::io::ErrorKind::WouldBlock => { - thread::sleep(Duration::from_millis(5)); + if let Some(frame) = read_ws_frame(&mut reader) { + let _ = seen_tx.send(frame); } - Err(_) => break, - } + let stream = reader.into_inner(); + let _ = stream.write_all(&ws_binary_frame(&reply)); + write_ws_close_and_drain(stream, b"done"); + }); } }); (url, seen_rx) @@ -149,51 +141,43 @@ fn start_ws_abnormal_close_server() -> String { fn start_ws_session_server() -> (String, mpsc::Receiver) { let listener = TcpListener::bind("127.0.0.1:0").expect("bind websocket session server"); - listener.set_nonblocking(true).unwrap(); let url = format!("ws://{}", listener.local_addr().unwrap()); let (seen_tx, seen_rx) = mpsc::channel(); thread::spawn(move || { - loop { - match listener.accept() { - Ok((mut stream, _)) => { - let _ = stream.set_nonblocking(false); - let seen_tx = seen_tx.clone(); - thread::spawn(move || { - let _ = stream.set_read_timeout(Some(Duration::from_secs(2))); - let mut reader = BufReader::new(stream.try_clone().unwrap()); - let Some(req) = read_request(&mut reader) else { - return; - }; - if !req.header("cookie").contains("sid=abc") { - write_response( - &mut stream, - TestResponse::status(401, "Unauthorized", "missing session"), - ); - return; - } - let key = req.header("sec-websocket-key"); - let mut sha = Sha1::new(); - sha.update(key.as_bytes()); - sha.update(b"258EAFA5-E914-47DA-95CA-C5AB0DC85B11"); - let accept = - base64::engine::general_purpose::STANDARD.encode(sha.finalize()); - let response = format!( - "HTTP/1.1 101 Switching Protocols\r\nUpgrade: websocket\r\nConnection: Upgrade\r\nSec-WebSocket-Accept: {accept}\r\nSet-Cookie: wsid=upgraded; Path=/\r\n\r\n" - ); - if stream.write_all(response.as_bytes()).is_err() { - return; - } - let msg = read_ws_text(&mut stream); - let _ = seen_tx.send(msg.clone()); - let _ = stream.write_all(&ws_text_frame(msg.as_bytes())); - write_ws_close_and_drain(&mut stream, b"done"); - }); + for stream in listener.incoming() { + let Ok(mut stream) = stream else { + break; + }; + let seen_tx = seen_tx.clone(); + thread::spawn(move || { + let _ = stream.set_read_timeout(Some(Duration::from_secs(2))); + let mut reader = BufReader::new(stream.try_clone().unwrap()); + let Some(req) = read_request(&mut reader) else { + return; + }; + if !req.header("cookie").contains("sid=abc") { + write_response( + &mut stream, + TestResponse::status(401, "Unauthorized", "missing session"), + ); + return; } - Err(err) if err.kind() == std::io::ErrorKind::WouldBlock => { - thread::sleep(Duration::from_millis(5)); + let key = req.header("sec-websocket-key"); + let mut sha = Sha1::new(); + sha.update(key.as_bytes()); + sha.update(b"258EAFA5-E914-47DA-95CA-C5AB0DC85B11"); + let accept = base64::engine::general_purpose::STANDARD.encode(sha.finalize()); + let response = format!( + "HTTP/1.1 101 Switching Protocols\r\nUpgrade: websocket\r\nConnection: Upgrade\r\nSec-WebSocket-Accept: {accept}\r\nSet-Cookie: wsid=upgraded; Path=/\r\n\r\n" + ); + if stream.write_all(response.as_bytes()).is_err() { + return; } - Err(_) => break, - } + let msg = read_ws_text(&mut stream); + let _ = seen_tx.send(msg.clone()); + let _ = stream.write_all(&ws_text_frame(msg.as_bytes())); + write_ws_close_and_drain(&mut stream, b"done"); + }); } }); (url, seen_rx)