diff --git a/src/defs/bytes_counter.rs b/src/defs/bytes_counter.rs index 8588b20..4f635c7 100644 --- a/src/defs/bytes_counter.rs +++ b/src/defs/bytes_counter.rs @@ -1,8 +1,10 @@ //! Counts bytes transferred during the download and upload tests. -use std::sync::Mutex; +use std::sync::{Arc, Mutex}; use std::time::Instant; +use crate::http::WriteMeter; + /// Tracks total bytes transferred and derives the average transfer rate. /// /// The total is a mutex rather than an `AtomicU64` because 32-bit targets such @@ -15,6 +17,16 @@ pub struct BytesCounter { start: Mutex>, mebi: bool, upload_size: usize, + /// The ceiling on the total, for the upload phase; see `set_wire`. + wire: Option, +} + +/// What a client's connections had written when a phase started, and the meter +/// to ask again. +#[derive(Debug)] +struct Wire { + meter: Arc, + start: u64, } impl BytesCounter { @@ -24,6 +36,7 @@ impl BytesCounter { start: Mutex::new(None), mebi: false, upload_size: 0, + wire: None, } } @@ -41,6 +54,26 @@ impl BytesCounter { self.upload_size } + /// Caps the total at what `meter`'s connections write from now on. + /// + /// For the upload phase. Its total is body bytes counted as hyper takes + /// each frame, and hyper takes frames ahead of writing them, so without a + /// ceiling the total includes whatever is still queued -- hundreds of + /// kilobytes per connection, which is a quarter to a half of what a 1 Mbit + /// uplink carries in a whole test. + /// + /// The meter counts request heads and chunk framing along with body bytes, + /// so as a measure of body bytes written it is high by the overhead: a + /// head is 119 bytes for the request this client sends, and chunk framing + /// is 8 bytes per 16 KiB frame. That is the whole error this leaves -- the + /// total lands between what the connections wrote of the body and that + /// plus the overhead -- and a request the peer took whole is counted + /// exactly, its body bytes alone being fewer than head and body together. + pub fn set_wire(&mut self, meter: Arc) { + let start = meter.written(); + self.wire = Some(Wire { meter, start }); + } + /// Starts the clock used for the average. pub fn start(&self) { *self.start.lock().unwrap() = Some(Instant::now()); @@ -51,9 +84,13 @@ impl BytesCounter { *self.total.lock().unwrap() += n; } - /// Total bytes read or written. + /// Total bytes read or written, at most what the wire ceiling allows. pub fn total(&self) -> u64 { - *self.total.lock().unwrap() + let counted = *self.total.lock().unwrap(); + match &self.wire { + Some(wire) => counted.min(wire.meter.written().saturating_sub(wire.start)), + None => counted, + } } fn elapsed_secs(&self) -> f64 { diff --git a/src/defs/server.rs b/src/defs/server.rs index 751eca2..9021763 100644 --- a/src/defs/server.rs +++ b/src/defs/server.rs @@ -26,8 +26,9 @@ use crate::{write_debug, write_ui}; /// The stagger between starting concurrent transfer streams. const RAMP_UP_DELAY: Duration = Duration::from_millis(200); /// The chunk size the upload body is fed to the connection in. hyper takes a -/// whole chunk whenever its write buffer has room, so this is also how much of -/// a connection's counted upload can sit past that buffer, unsent. +/// whole chunk whenever its write buffer has room and writes up to sixteen of +/// them at once, so this is the granularity of the upload's write syscalls, not +/// their size. const UPLOAD_CHUNK: usize = 16 * 1024; /// A speed test server, as described by the server list JSON. @@ -449,6 +450,10 @@ impl Server { let mut counter = BytesCounter::new(); counter.set_mebi(opts.use_mebi); counter.set_upload_size(opts.upload_size); + // The body counts frames as hyper takes them, which runs ahead of what + // the connections have written; this is what keeps the total from + // reporting the difference as sent. + counter.set_wire(client.wire()); let counter = Arc::new(counter); // Pre-allocating one random blob and reusing it keeps the CPU out of the @@ -464,9 +469,6 @@ impl Server { }; let url = url_join_path(&self.get_url()?, &self.upload_url); - // Upload over the pool whose buffers are capped, so little is counted - // but still unsent when the window closes. - let client = client.for_uploads(); counter.start(); let spinner = self.start_transfer_spinner("Uploading... ", opts, &counter); @@ -809,10 +811,9 @@ mod upload_body { /// Each frame is added to the upload total as hyper takes it, so the total /// is request body bytes, which is what the Go client's TeeReader around /// its request body counts: no request heads, chunk framing or TLS - /// overhead. A frame taken is not yet sent. hyper takes one whenever its - /// write buffer has room, and that buffer is capped (see - /// `http::client_builder`), so a connection holds at most the buffer and - /// one frame counted but unsent when the window closes. + /// overhead. A frame taken is not yet sent -- hyper takes one whenever its + /// write buffer has room -- so the phase caps the total at what the + /// connections have written; see `BytesCounter::set_wire`. pub(super) struct UploadBody { payload: Option, pos: usize, @@ -1278,19 +1279,21 @@ mod transfer_tests { ); } - /// hyper takes upload frames while its write buffer has room, and each - /// frame is counted as it is taken, so the buffer decides how much a - /// connection has counted but not sent when the window closes. hyper's - /// default of ~400 KB overstated a slow uplink by a quarter to a half. + /// hyper takes upload frames while its write buffer has room, which runs + /// ahead of writing them, and each frame is counted as it is taken. Without + /// the ceiling that queue is reported as sent: hyper writes sixteen frames + /// at a time, and on a slow uplink filling the queue is most of what a + /// window does. #[tokio::test] - async fn a_connection_holds_at_most_a_buffer_and_a_frame_of_counted_upload() { + async fn an_upload_counts_no_more_than_the_connection_wrote() { + use crate::http::MeteredConnector; use hyper_util::client::legacy::connect::{Connected, Connection}; use std::pin::Pin; use std::task::{Context, Poll}; /// What the peer takes before it stops reading. Not a whole number of /// 16-frame batches: hyper also stops at 16 queued buffers, and a - /// budget ending on that boundary would hide an uncapped buffer. + /// budget ending on that boundary would leave nothing queued. const TAKES: usize = 300_000; /// A connection whose peer takes `TAKES` bytes and then nothing more, @@ -1361,82 +1364,99 @@ mod transfer_tests { } } - let taken = Arc::new(AtomicUsize::new(0)); - let client = crate::http::client_builder(1, false, true) - .build::<_, crate::http::ReqBody>(Connect(taken.clone())); - let counter = Arc::new(BytesCounter::new()); - let body = BodyExt::boxed(upload_body::UploadBody::new(None, counter.clone())); - let request = http::Request::post("http://stalled.invalid/") - .body(body) - .unwrap(); - // Nothing ever answers, so this only ends at the timeout. - let _ = tokio::time::timeout(Duration::from_millis(500), client.request(request)).await; + // The same upload twice: without the ceiling to show the queue is + // counted, and with it. + let mut counted = Vec::new(); + for ceiling in [false, true] { + let taken = Arc::new(AtomicUsize::new(0)); + let meter = Arc::new(crate::http::WriteMeter::new()); + let client = crate::http::client_builder(1, false).build::<_, crate::http::ReqBody>( + MeteredConnector::new(Connect(taken.clone()), meter.clone()), + ); + let mut counter = BytesCounter::new(); + if ceiling { + counter.set_wire(meter.clone()); + } + let counter = Arc::new(counter); + let body = BodyExt::boxed(upload_body::UploadBody::new(None, counter.clone())); + let request = http::Request::post("http://stalled.invalid/") + .body(body) + .unwrap(); + // Nothing ever answers, so this only ends at the timeout. + let _ = tokio::time::timeout(Duration::from_millis(500), client.request(request)).await; + + let taken = taken.load(Ordering::SeqCst) as u64; + assert_eq!( + taken, TAKES as u64, + "the connection took less than it could" + ); + assert_eq!( + meter.written(), + taken, + "the meter counted something other than what the peer took" + ); + counted.push(counter.total()); + } - let taken = taken.load(Ordering::SeqCst) as u64; - assert_eq!( - taken, TAKES as u64, - "the connection took less than it could" - ); - // What the peer took includes the request head and chunk framing, so - // this slightly understates what is held; the bound has room to spare. - let held = counter.total().saturating_sub(taken); - let bound = (crate::http::H1_MAX_BUF + UPLOAD_CHUNK) as u64; assert!( - held <= bound, - "{held} bytes were counted but not sent, more than the {bound} a buffer and a frame hold" + counted[0] > TAKES as u64, + "nothing was queued unsent, so this test cannot show a ceiling working" + ); + // The peer took the request head and the chunk framing along with the + // body, so the ceiling is a little above the body bytes it received; + // what matters is that nothing still in hyper is in the total. + assert_eq!( + counted[1], TAKES as u64, + "{} bytes were counted, more than the {TAKES} the connection wrote", + counted[1] ); } - /// The upload phase sends over the capped pool: a response head bigger - /// than the cap fails its request there, and a failed request is not - /// started again. Over the default pool it would be read and repeated. + /// The phase asks for the ceiling, which is the half of the accounting the + /// body cannot do for itself: with the peer no longer reading, the total is + /// what the connection wrote and not what the body handed over. #[tokio::test] - async fn the_upload_phase_sends_over_the_capped_pool() { - use tokio::io::{AsyncReadExt, AsyncWriteExt}; + async fn the_upload_phase_counts_no_more_than_its_connections_wrote() { + use tokio::io::AsyncReadExt; let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); let addr = listener.local_addr().unwrap(); - let requests = Arc::new(AtomicUsize::new(0)); - let seen = requests.clone(); tokio::spawn(async move { while let Ok((mut stream, _)) = listener.accept().await { - let seen = seen.clone(); tokio::spawn(async move { - // Answer only once the head and the 1 KiB body are in. - let mut request = Vec::new(); + // Read the head, then stop reading and hold the connection + // open, so the client fills the socket and hyper queues the + // rest. + let mut seen = Vec::new(); let mut buf = [0u8; 8192]; - while !request - .windows(4) - .position(|w| w == b"\r\n\r\n") - .is_some_and(|head| request.len() >= head + 4 + 1024) - { + while !seen.windows(4).any(|w| w == b"\r\n\r\n") { match stream.read(&mut buf).await { - Ok(n) if n > 0 => request.extend_from_slice(&buf[..n]), + Ok(n) if n > 0 => seen.extend_from_slice(&buf[..n]), _ => return, } } - seen.fetch_add(1, Ordering::SeqCst); - let reply = format!( - "HTTP/1.1 200 OK\r\nX-Pad: {}\r\nContent-Length: 0\r\nConnection: close\r\n\r\n", - "a".repeat(2 * crate::http::H1_MAX_BUF) - ); - let _ = stream.write_all(reply.as_bytes()).await; - while matches!(stream.read(&mut buf).await, Ok(n) if n > 0) {} + std::future::pending::<()>().await; }); } }); - let mut opts = opts(Duration::from_millis(500)); + + let mut opts = opts(Duration::from_secs(1)); opts.requests = 1; + // An endless body, so the window closes with the stream mid-transfer + // and whatever hyper holds is held for good. + opts.no_prealloc = true; - server_at(addr) - .upload(&client(), &TelemetryLog::new(), &opts) + let client = client_with_timeout(Duration::ZERO); + let (_, total) = server_at(addr) + .upload(&client, &TelemetryLog::new(), &opts) .await .unwrap(); - assert_eq!( - requests.load(Ordering::SeqCst), - 1, - "the upload read a response head bigger than the cap, so its pool is not capped" + let written = client.wire().written(); + assert!(total > 0, "nothing was counted at all"); + assert!( + total <= written, + "the phase counted {total} bytes, more than the {written} its connections wrote" ); } } diff --git a/src/http/meter.rs b/src/http/meter.rs new file mode 100644 index 0000000..ab54d75 --- /dev/null +++ b/src/http/meter.rs @@ -0,0 +1,259 @@ +//! Counting the bytes a client's connections hand to the transport. +//! +//! The upload total is body bytes, counted as hyper takes each frame, and a +//! frame taken is not yet written: hyper queues what it cannot write at once. +//! Counting the queue as sent overstated a slow uplink by a quarter to a half, +//! and the first fix for it shrank the queue -- which bought the accounting at +//! the price of a write syscall every 16 KiB. +//! +//! This is the other half: the connections report what they write, and the +//! upload total is never allowed past it (see `BytesCounter::set_wire`). What +//! the meter counts is one layer below what hyper writes and one above the +//! socket -- plaintext, so it is comparable with body bytes, and returned by +//! the transport, so it excludes anything hyper is still holding. It counts +//! request heads and chunk framing as well as body bytes, which is why it is +//! a ceiling rather than the total itself: it can only ever leave the total a +//! few hundred bytes per request too high, where the queue left it hundreds of +//! kilobytes per connection too high. +//! +//! Over TLS the bytes are counted as the TLS session accepts them, so whatever +//! the session has buffered without writing it to the socket is counted too: +//! with rustls, its plaintext and outgoing-record buffers, 64 KiB each by +//! default. Capping hyper's own buffers never reached those, so this is no +//! looser there than what it replaces. + +use std::future::Future; +use std::pin::Pin; +use std::sync::{Arc, Mutex}; +use std::task::{Context, Poll}; + +use http::Uri; +use hyper_util::client::legacy::connect::{Connected, Connection}; + +/// How many bytes a client's connections have handed to the transport. +#[derive(Debug, Default)] +pub struct WriteMeter { + /// A mutex rather than an `AtomicU64` for the reason `BytesCounter` uses + /// one: 32-bit targets such as the PowerPC in Turris 1.x routers have no + /// 64-bit atomics. This is bumped once per write, not once per byte. + written: Mutex, +} + +impl WriteMeter { + pub fn new() -> Self { + Self::default() + } + + fn add(&self, n: usize) { + *self.written.lock().unwrap() += n as u64; + } + + /// Bytes written since the client was built. + pub fn written(&self) -> u64 { + *self.written.lock().unwrap() + } +} + +/// A connector whose streams report what they write to `meter`. +#[derive(Clone, Debug)] +pub struct MeteredConnector { + inner: C, + meter: Arc, +} + +impl MeteredConnector { + pub fn new(inner: C, meter: Arc) -> Self { + Self { inner, meter } + } +} + +impl tower_service::Service for MeteredConnector +where + C: tower_service::Service, + C::Response: Send + 'static, + C::Future: Send + 'static, + C::Error: Send + 'static, +{ + type Response = MeteredStream; + type Error = C::Error; + type Future = Pin> + Send>>; + + fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll> { + self.inner.poll_ready(cx) + } + + fn call(&mut self, dst: Uri) -> Self::Future { + let meter = self.meter.clone(); + let connect = self.inner.call(dst); + Box::pin(async move { + Ok(MeteredStream { + inner: connect.await?, + meter, + }) + }) + } +} + +/// A connected stream that counts the bytes it accepts for writing. +#[derive(Debug)] +pub struct MeteredStream { + inner: S, + meter: Arc, +} + +impl Connection for MeteredStream { + fn connected(&self) -> Connected { + self.inner.connected() + } +} + +impl hyper::rt::Read for MeteredStream { + fn poll_read( + mut self: Pin<&mut Self>, + cx: &mut Context<'_>, + buf: hyper::rt::ReadBufCursor<'_>, + ) -> Poll> { + Pin::new(&mut self.inner).poll_read(cx, buf) + } +} + +impl hyper::rt::Write for MeteredStream { + fn poll_write( + mut self: Pin<&mut Self>, + cx: &mut Context<'_>, + buf: &[u8], + ) -> Poll> { + let written = Pin::new(&mut self.inner).poll_write(cx, buf); + if let Poll::Ready(Ok(n)) = &written { + self.meter.add(*n); + } + written + } + + fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + Pin::new(&mut self.inner).poll_flush(cx) + } + + fn poll_shutdown(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + Pin::new(&mut self.inner).poll_shutdown(cx) + } + + /// Forwarded, and not merely for tidiness: a stream that says it cannot + /// take a vector of buffers makes hyper copy every byte of every body into + /// a buffer of its own before writing it, which on the upload path is the + /// whole payload, over and over. + fn is_write_vectored(&self) -> bool { + self.inner.is_write_vectored() + } + + fn poll_write_vectored( + mut self: Pin<&mut Self>, + cx: &mut Context<'_>, + bufs: &[std::io::IoSlice<'_>], + ) -> Poll> { + let written = Pin::new(&mut self.inner).poll_write_vectored(cx, bufs); + if let Poll::Ready(Ok(n)) = &written { + self.meter.add(*n); + } + written + } +} + +#[cfg(test)] +mod tests { + use super::*; + + use std::io::IoSlice; + + use hyper::rt::Write as _; + + /// A stream that takes at most `take` bytes of any write. + struct Partial { + take: usize, + vectored: bool, + } + + impl hyper::rt::Read for Partial { + fn poll_read( + self: Pin<&mut Self>, + _: &mut Context<'_>, + _: hyper::rt::ReadBufCursor<'_>, + ) -> Poll> { + Poll::Pending + } + } + + impl hyper::rt::Write for Partial { + fn poll_write( + self: Pin<&mut Self>, + _: &mut Context<'_>, + buf: &[u8], + ) -> Poll> { + Poll::Ready(Ok(buf.len().min(self.take))) + } + + fn poll_flush(self: Pin<&mut Self>, _: &mut Context<'_>) -> Poll> { + Poll::Ready(Ok(())) + } + + fn poll_shutdown(self: Pin<&mut Self>, _: &mut Context<'_>) -> Poll> { + Poll::Ready(Ok(())) + } + + fn is_write_vectored(&self) -> bool { + self.vectored + } + + fn poll_write_vectored( + self: Pin<&mut Self>, + _: &mut Context<'_>, + bufs: &[IoSlice<'_>], + ) -> Poll> { + let total: usize = bufs.iter().map(|b| b.len()).sum(); + Poll::Ready(Ok(total.min(self.take))) + } + } + + fn noop_context() -> Context<'static> { + Context::from_waker(std::task::Waker::noop()) + } + + /// What the meter records is what the stream took, not what was offered: + /// a short write leaves the rest for hyper to write later, and counting it + /// now is counting a byte that is still in a buffer. + #[test] + fn a_short_write_counts_only_what_was_taken() { + let meter = Arc::new(WriteMeter::new()); + let mut stream = MeteredStream { + inner: Partial { + take: 10, + vectored: false, + }, + meter: meter.clone(), + }; + let mut cx = noop_context(); + + let n = Pin::new(&mut stream).poll_write(&mut cx, &[0u8; 100]); + assert!(matches!(n, Poll::Ready(Ok(10)))); + let n = Pin::new(&mut stream).poll_write_vectored(&mut cx, &[IoSlice::new(&[0u8; 100])]); + assert!(matches!(n, Poll::Ready(Ok(10)))); + assert_eq!(meter.written(), 20, "the meter counted bytes not taken"); + } + + /// See `is_write_vectored` above: getting this wrong costs a copy of every + /// uploaded byte and would not fail any other test. + #[test] + fn vectored_writes_are_reported_as_the_inner_stream_reports_them() { + for vectored in [false, true] { + let stream = MeteredStream { + inner: Partial { take: 1, vectored }, + meter: Arc::new(WriteMeter::new()), + }; + assert_eq!( + hyper::rt::Write::is_write_vectored(&stream), + vectored, + "the metered stream does not report the inner stream's vectored writes" + ); + } + } +} diff --git a/src/http/mod.rs b/src/http/mod.rs index 5a7f128..0b9bb15 100644 --- a/src/http/mod.rs +++ b/src/http/mod.rs @@ -2,9 +2,11 @@ //! (`--source`, `--interface`, `--fwmark`) can be honoured. pub mod connector; +pub mod meter; pub mod tls; use std::io; +use std::sync::Arc; use std::time::Duration; use anyhow::{bail, Context as _}; @@ -19,6 +21,7 @@ use hyper_util::rt::TokioExecutor; use url::Url; pub use connector::{BindOptions, IpFamily}; +pub use meter::{MeteredConnector, WriteMeter}; pub use tls::{TlsFacts, TlsSettings}; /// What the transport negotiated for a request, read back off its response. @@ -66,9 +69,6 @@ pub const MAX_TELEMETRY_RESPONSE: usize = 64 * 1024; const H2_STREAM_WINDOW: u32 = 4 * 1024 * 1024; const H2_CONNECTION_WINDOW: u32 = 1024 * 1024 * 1024; -/// The cap on each HTTP/1 connection's buffers: hyper's minimum. -pub(crate) const H1_MAX_BUF: usize = 8192; - pub type ReqBody = BoxBody; /// What a request sends, which decides what a 307 or 308 can send again. @@ -174,28 +174,22 @@ async fn body_prefix(body: Incoming, limit: usize) -> anyhow::Result { /// The hyper client configuration `HttpClient` is built from. /// -/// `bounded` caps a connection's HTTP/1 buffers, and only the upload pool asks -/// for it. The upload total counts body frames as hyper takes them, so what -/// hyper holds when the window closes was counted but never sent: with its -/// default buffer, ~400 KB per connection, that overstated a 1 Mbit uplink by -/// a quarter to a half. The cap is hyper's minimum, leaving 8 KiB and one -/// frame. It binds the read buffer as well, which cost a 65 Gbit loopback -/// download a fifth of its rate and would limit every response head to 8 KiB, -/// so downloads, pings and the control-plane requests keep hyper's defaults. +/// The HTTP/1 buffers are left at hyper's defaults, upload connections +/// included. Capping them to hyper's minimum of 8 KiB used to be what kept the +/// upload total honest -- a frame hyper has taken but not written was counted +/// as sent, so the less it could hold, the less it could overstate. It cost a +/// write syscall per 16 KiB frame, where uncapped hyper writes sixteen frames +/// in one, and on plain HTTP that is most of the upload path's CPU. The +/// ceiling in `meter` bounds the same error without bounding the buffers. pub(crate) fn client_builder( concurrent: usize, http2: bool, - bounded: bool, ) -> hyper_util::client::legacy::Builder { // Keep enough connections alive for every concurrent stream, matching the // Go version's MaxIdleConnsPerHost/MaxConnsPerHost tuning. let mut builder = Client::builder(TokioExecutor::new()); builder.pool_max_idle_per_host(concurrent + 2); - if bounded { - builder.http1_max_buf_size(H1_MAX_BUF); - } - if http2 { builder .http2_initial_stream_window_size(H2_STREAM_WINDOW) @@ -207,9 +201,10 @@ pub(crate) fn client_builder( /// The program's HTTP client. #[derive(Clone)] pub struct HttpClient { - inner: Client, - /// The pool whose buffers are capped, for the upload test only. - upload: Client, + inner: Client, ReqBody>, + /// What this client's connections have written, which the upload total is + /// not allowed past; see `meter`. + wire: Arc, timeout: Duration, user_agent: HeaderValue, } @@ -222,27 +217,24 @@ impl HttpClient { concurrent: usize, user_agent: &str, ) -> anyhow::Result { - let https = tls::build(bind, tls_settings)?; - let inner = client_builder(concurrent, tls_settings.http2, false).build(https.clone()); - let upload = client_builder(concurrent, tls_settings.http2, true).build(https); + let wire = Arc::new(WriteMeter::new()); + let https = MeteredConnector::new(tls::build(bind, tls_settings)?, wire.clone()); + let inner = client_builder(concurrent, tls_settings.http2).build(https); Ok(Self { inner, - upload, + wire, timeout, user_agent: HeaderValue::from_str(user_agent)?, }) } - /// The same client, sending over the pool whose buffers are capped. + /// What this client's connections have written so far. /// - /// Only the upload test wants that cap: it keeps what hyper has counted - /// but not sent small, and it costs a download speed. - pub fn for_uploads(&self) -> Self { - Self { - inner: self.upload.clone(), - ..self.clone() - } + /// The upload test takes this as the ceiling on what it may report as + /// sent, which is what lets the connections buffer freely; see `meter`. + pub fn wire(&self) -> Arc { + self.wire.clone() } /// The configured per-request timeout (`--timeout`). @@ -1710,4 +1702,106 @@ ZD/4gnUj9TooNmtCjXdP4GAKxDoCb0FzQoHhDoi8BXY4DFLQYpbnX4Nu .expect("the connection was not closed") .unwrap(); } + + /// Over TLS the meter counts what the TLS session accepted, which is what + /// makes it comparable with body bytes, and what it excludes is whatever + /// hyper is still holding. + /// + /// What it does not exclude is the transport below: rustls buffers 64 KiB + /// of plaintext and 64 KiB of records by default, and the kernel socket + /// buffers take what they take -- on Linux loopback that autotunes into + /// the megabytes, where a peer that read 64 KiB of a 4 MiB body still let + /// 2.8 MB leave this process. So there is no portable tight bound here, + /// and this pins the two ends that are portable: the meter is in the TLS + /// path at all, and it counts each write once. + /// + /// What it catches, proven by reverting each: a meter absent from the + /// path, and one wrapped below the session, where it would count + /// ciphertext. What it does not catch: counting what a write offered + /// rather than what it took, nor counting a write twice -- rustls accepts + /// a whole buffer up to its limit, and a stalled peer makes hyper wait on + /// a waker rather than retry, so neither shows here. The upper bound below + /// is a sanity guard, not a detector for those. What pins the ceiling this + /// feeds is `an_upload_counts_no_more_than_the_connection_wrote`, whose + /// transport is a fake one and therefore deterministic. + #[cfg(feature = "rustls-tls")] + #[tokio::test] + async fn the_meter_counts_no_more_than_a_stalled_tls_peer_took() { + use std::io::Read as _; + use std::sync::atomic::{AtomicUsize, Ordering}; + + use rustls_pki_types::pem::PemObject as _; + use rustls_pki_types::{CertificateDer, PrivateKeyDer}; + + /// Bigger than every buffer in the path together. + const BODY: usize = 4 * 1024 * 1024; + /// What the peer reads before it stops. + const TAKES: usize = 64 * 1024; + + let (listener, url) = tls_listen(); + let taken = Arc::new(AtomicUsize::new(0)); + let read_by_peer = taken.clone(); + + let provider = Arc::new(rustls::crypto::ring::default_provider()); + let config = rustls::ServerConfig::builder_with_provider(provider) + .with_safe_default_protocol_versions() + .unwrap() + .with_no_client_auth() + .with_single_cert( + vec![CertificateDer::from_pem_slice(TEST_CERT.as_bytes()).unwrap()], + PrivateKeyDer::from_pem_slice(TEST_KEY.as_bytes()).unwrap(), + ) + .unwrap(); + let config = Arc::new(config); + + std::thread::spawn(move || { + let Ok((tcp, _)) = listener.accept() else { + return; + }; + let conn = rustls::ServerConnection::new(config).unwrap(); + let mut tls = rustls::StreamOwned::new(conn, tcp); + + // The head, then part of the body, then nothing: the request is + // never answered and the connection is held open, so the client + // goes on writing into a peer that has stopped reading. + let mut head = Vec::new(); + let mut byte = [0u8; 1]; + while !head.ends_with(b"\r\n\r\n") { + match tls.read(&mut byte) { + Ok(1) => head.push(byte[0]), + _ => return, + } + } + let mut chunk = vec![0u8; 16 * 1024]; + while read_by_peer.load(Ordering::Relaxed) < TAKES { + match tls.read(&mut chunk) { + Ok(0) | Err(_) => break, + Ok(n) => { + read_by_peer.fetch_add(n, Ordering::Relaxed); + } + } + } + std::thread::sleep(Duration::from_secs(30)); + }); + + let client = client(Duration::from_secs(2)); + let wire = client.wire(); + let body = Bytes::from(vec![b'x'; BODY]); + let result = client + .post_bytes(&url, "application/octet-stream", body) + .await; + assert!(result.is_err(), "the stalled peer must not answer"); + + let counted = wire.written(); + let took = taken.load(Ordering::Relaxed) as u64; + assert!( + counted >= took, + "the meter counted {counted} but the peer read {took}" + ); + assert!( + counted <= BODY as u64 + 4096, + "the meter counted {counted} of a {BODY}-byte body, so a write was \ + counted more than once" + ); + } }