lib.rsannotatedlib.rssource487 lines · 17.8 KB · raw

A stand-in for TypeSafe's System One endpoint: HTTP/2 over TLS on loopback, with a throwaway CA the client under test is told to trust. It records every request and answers from a closure. Any Jev client can test against it; nothing here knows about Postgres.

6use std::net::SocketAddr;
7use std::sync::atomic::{AtomicUsize, Ordering};
8use std::sync::{Arc, Mutex};
9use std::time::Duration;
11use bytes::Bytes;
12use http_body_util::{BodyExt, Full};
13use hyper::body::Incoming;
14use hyper::{Request, Response};
15use hyper_util::rt::{TokioExecutor, TokioIo};
16use rcgen::{BasicConstraints, CertificateParams, IsCa, Issuer, KeyPair};
17use tokio::net::TcpListener;
18use tokio_rustls::TlsAcceptor;
19use tokio_rustls::rustls::ServerConfig;
20use tokio_rustls::rustls::pki_types::{CertificateDer, PrivateKeyDer, PrivatePkcs8KeyDer};
21
22use tempfile::TempDir;
23
24pub mod forward;
25pub mod replay;

One request as the server saw it.

28#[derive(Debug, Clone)]
29pub struct Recorded {
30    pub version: http::Version,
31    pub method: http::Method,
32    pub path: String,
33    pub headers: http::HeaderMap,
34    pub body: Bytes,
35}

What the mock answers.

38pub struct Reply {
39    pub status: u16,
40    pub headers: Vec<(&'static str, String)>,
41    pub body: serde_json::Value,

How long to stall before answering.

43    pub delay: Option<Duration>,

Stall until the mock has recorded this many requests.

45    pub hold_until: Option<usize>,

Sent byte for byte instead of body, as a replayed fixture is.

47    pub raw_body: Option<Bytes>,
48}
50impl Reply {
51    pub fn json(status: u16, body: serde_json::Value) -> Self {
52        Reply { status, headers: Vec::new(), body, delay: None, hold_until: None, raw_body: None }
53    }

Answers only after delay; a client that gives up first resets the stream, which [MockJev::abandoned] counts.

57    pub fn after(mut self, delay: Duration) -> Self {
58        self.delay = Some(delay);
59        self
60    }

Answers only once n requests have arrived (then waits any after delay), so a test can count what a client sends before its first answer without racing it.

65    pub fn after_requests(mut self, n: usize) -> Self {
66        self.hold_until = Some(n);
67        self
68    }
70    pub fn header(mut self, name: &'static str, value: impl Into<String>) -> Self {
71        self.headers.push((name, value.into()));
72        self
73    }
74}
75
76type Responder = dyn Fn(&Recorded) -> Reply + Send + Sync;

What the server tasks share with the test.

79struct Shared {
80    requests: Mutex<Vec<Recorded>>,

How many requests have been recorded, for replies that wait on it.

82    recorded: tokio::sync::watch::Sender<usize>,

Requests whose handler was dropped before answering: the client reset the stream (RST_STREAM) or closed the connection.

85    abandoned: AtomicUsize,

Requests arrived and not yet answered or abandoned, and the most there have been at once.

88    in_flight: AtomicUsize,
89    peak_in_flight: AtomicUsize,

TCP connections accepted, whether or not TLS then succeeded.

91    accepts: AtomicUsize,

TLS connections accepted.

93    connections: AtomicUsize,

PING frames the client sent (not acknowledgements).

95    pings: AtomicUsize,

Connections numbered up to this one (from 1) never acknowledge a PING, as a dead path would not.

98    silent_through: AtomicUsize,

Bumped to make every open connection send GOAWAY and close.

100    goaway: tokio::sync::watch::Sender<u64>,
101    respond: Box<Responder>,
102}
104pub struct MockJev {
105    pub addr: SocketAddr,
106    cert_dir: TempDir,
107    shared: Arc<Shared>,
108    task: tokio::task::JoinHandle<()>,
109}
110
111impl MockJev {

Serves on 127.0.0.1 with a certificate for that address.

113    pub async fn start(respond: impl Fn(&Recorded) -> Reply + Send + Sync + 'static) -> Self {
114        Self::start_for(&["127.0.0.1"], respond).await
115    }

Serves on 127.0.0.1 with a certificate for each of names (host names or IP addresses), for clients that resolve a name to it.

119    pub async fn start_for(
120        names: &[&str],
121        respond: impl Fn(&Recorded) -> Reply + Send + Sync + 'static,
122    ) -> Self {
123        Self::serve(names, None, respond).await
124    }

As start, advertising SETTINGS_MAX_CONCURRENT_STREAMS = max_streams on every connection.

128    pub async fn start_limited(
129        max_streams: u32,
130        respond: impl Fn(&Recorded) -> Reply + Send + Sync + 'static,
131    ) -> Self {
132        Self::serve(&["127.0.0.1"], Some(max_streams), respond).await
133    }
135    async fn serve(
136        names: &[&str],
137        max_streams: Option<u32>,
138        respond: impl Fn(&Recorded) -> Reply + Send + Sync + 'static,
139    ) -> Self {
140        let ca_key = KeyPair::generate().unwrap();
141        let mut ca_params = CertificateParams::new(Vec::<String>::new()).unwrap();
142        ca_params.is_ca = IsCa::Ca(BasicConstraints::Unconstrained);
143        let ca_cert = ca_params.self_signed(&ca_key).unwrap();
144        let issuer = Issuer::new(ca_params, ca_key);
145
146        let leaf_key = KeyPair::generate().unwrap();
147        let leaf = CertificateParams::new(names.iter().map(|n| n.to_string()).collect::<Vec<_>>())
148            .unwrap()
149            .signed_by(&leaf_key, &issuer)
150            .unwrap();
151
152        let cert_dir = tempfile::Builder::new().prefix("jev-mock").tempdir().unwrap();
153        std::fs::write(cert_dir.path().join("ca.pem"), ca_cert.pem()).unwrap();
154
155        let mut tls = ServerConfig::builder()
156            .with_no_client_auth()
157            .with_single_cert(
158                vec![CertificateDer::from(leaf.der().to_vec())],
159                PrivateKeyDer::Pkcs8(PrivatePkcs8KeyDer::from(leaf_key.serialize_der())),
160            )
161            .unwrap();
162        // h2 only: the client must negotiate it (contract §2).
163        tls.alpn_protocols = vec![b"h2".to_vec()];
164        let acceptor = TlsAcceptor::from(Arc::new(tls));
165
166        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
167        let addr = listener.local_addr().unwrap();
168        let shared = Arc::new(Shared {
169            requests: Mutex::new(Vec::new()),
170            recorded: tokio::sync::watch::channel(0).0,
171            abandoned: AtomicUsize::new(0),
172            in_flight: AtomicUsize::new(0),
173            peak_in_flight: AtomicUsize::new(0),
174            accepts: AtomicUsize::new(0),
175            connections: AtomicUsize::new(0),
176            pings: AtomicUsize::new(0),
177            silent_through: AtomicUsize::new(0),
178            goaway: tokio::sync::watch::channel(0).0,
179            respond: Box::new(respond),
180        });
181
182        let task = tokio::spawn({
183            let shared = shared.clone();
184            async move {
185                loop {
186                    let Ok((tcp, _)) = listener.accept().await else { return };
187                    shared.accepts.fetch_add(1, Ordering::SeqCst);
188                    let acceptor = acceptor.clone();
189                    let shared = shared.clone();
190                    tokio::spawn(async move {
191                        let Ok(tls) = acceptor.accept(tcp).await else { return };
192                        let n = shared.connections.fetch_add(1, Ordering::SeqCst) + 1;
193                        let silent = n <= shared.silent_through.load(Ordering::SeqCst);
194                        let tls = Frames::new(tls, shared.clone(), silent);
195                        let mut goaway = shared.goaway.subscribe();
196                        let service = {
197                            let shared = shared.clone();
198                            hyper::service::service_fn(move |req| handle(req, shared.clone()))
199                        };
200                        let conn = hyper::server::conn::http2::Builder::new(TokioExecutor::new())
201                            .max_concurrent_streams(max_streams)
202                            .serve_connection(TokioIo::new(tls), service);
203                        let mut conn = std::pin::pin!(conn);
204                        tokio::select! {
205                            _ = conn.as_mut() => {}
206                            _ = goaway.changed() => {
207                                conn.as_mut().graceful_shutdown();
208                                let _ = conn.await;
209                            }
210                        }
211                    });
212                }
213            }
214        });
215
216        MockJev { addr, cert_dir, shared, task }
217    }

The CA certificate (PEM) the client must trust.

220    pub fn ca_file(&self) -> std::path::PathBuf {
221        self.cert_dir.path().join("ca.pem")
222    }
224    pub fn endpoint(&self) -> String {
225        format!("https://{}/v1/systemone", self.addr)
226    }

The endpoint URL reached through host instead of the address.

229    pub fn endpoint_via(&self, host: &str) -> String {
230        format!("https://{host}:{}/v1/systemone", self.addr.port())
231    }
233    pub fn requests(&self) -> Vec<Recorded> {
234        self.shared.requests.lock().unwrap().clone()
235    }

Sends GOAWAY on every open connection and closes it once idle, as a server does when it restarts or rotates connections.

239    pub fn goaway(&self) {
240        self.shared.goaway.send_modify(|n| *n += 1);
241    }

The most requests the mock has held unanswered at once: a direct measure of how many the client had in flight, whatever the clock.

245    pub fn peak_in_flight(&self) -> usize {
246        self.shared.peak_in_flight.load(Ordering::SeqCst)
247    }

How many TCP connections the client opened, including those whose TLS handshake failed.

251    pub fn accepts(&self) -> usize {
252        self.shared.accepts.load(Ordering::SeqCst)
253    }

How many TLS connections the client opened.

256    pub fn connections(&self) -> usize {
257        self.shared.connections.load(Ordering::SeqCst)
258    }

How many PING frames the client sent (keep-alive and BDP probes).

261    pub fn pings(&self) -> usize {
262        self.shared.pings.load(Ordering::SeqCst)
263    }

The next n connections the client opens never acknowledge its PINGs, so its keep-alive sees a dead path; later ones do.

267    pub fn ignore_pings_on_next(&self, n: usize) {
268        self.shared.silent_through.store(self.connections() + n, Ordering::SeqCst);
269    }

How many requests the client gave up on before they were answered.

272    pub fn abandoned(&self) -> usize {
273        self.shared.abandoned.load(Ordering::SeqCst)
274    }
275}
277impl Drop for MockJev {
278    fn drop(&mut self) {
279        self.task.abort();
280    }
281}
282
283async fn handle(req: Request<Incoming>, shared: Arc<Shared>) -> Result<Response<Full<Bytes>>, hyper::Error> {
284    // hyper drops this future when the client resets the stream, whether
285    // during the upload or while the answer stalls; the guard sees either
286    // as a drop before `answered`.
287    let in_flight = shared.in_flight.fetch_add(1, Ordering::SeqCst) + 1;
288    shared.peak_in_flight.fetch_max(in_flight, Ordering::SeqCst);
289    let mut guard = Unanswered { shared: &shared, answered: false };
290    let (parts, body) = req.into_parts();
291    let recorded = Recorded {
292        version: parts.version,
293        method: parts.method,
294        path: parts.uri.path().to_string(),
295        headers: parts.headers,
296        body: body.collect().await?.to_bytes(),
297    };
298    let reply = (shared.respond)(&recorded);
299    {
300        let mut requests = shared.requests.lock().unwrap();
301        requests.push(recorded);
302        shared.recorded.send_replace(requests.len());
303    }
304    if let Some(n) = reply.hold_until {
305        // Never errs: `shared` keeps the sender alive.
306        let _ = shared.recorded.subscribe().wait_for(|&count| count >= n).await;
307    }
308    if let Some(delay) = reply.delay {
309        tokio::time::sleep(delay).await;
310    }
311    guard.answered = true;
312    let mut response = Response::builder()
313        .status(reply.status)
314        .header("content-type", "application/json");
315    for (name, value) in reply.headers {
316        response = response.header(name, value);
317    }
318    let body = reply.raw_body.unwrap_or_else(|| Bytes::from(reply.body.to_string()));
319    Ok(response.body(Full::new(body)).unwrap())
320}
321
322struct Unanswered<'a> {
323    shared: &'a Shared,
324    answered: bool,
325}
326
327impl Drop for Unanswered<'_> {
328    fn drop(&mut self) {
329        self.shared.in_flight.fetch_sub(1, Ordering::SeqCst);
330        if !self.answered {
331            self.shared.abandoned.fetch_add(1, Ordering::SeqCst);
332        }
333    }
334}

The plaintext HTTP/2 byte stream between TLS and hyper, read as frames: it counts the client's PINGs, and on a silent connection drops the server's PING acknowledgements. hyper answers PINGs itself, so this is the only place to see or withhold them.

340struct Frames<T> {
341    inner: T,
342    shared: Arc<Shared>,

Client to server: the connection preface comes before any frame.

344    incoming: FrameParser,
345    outgoing: FrameParser,
346    silent: bool,

Filtered bytes accepted from hyper, not yet written.

348    pending: Vec<u8>,
349}

Tracks frame boundaries across arbitrary chunks.

352struct FrameParser {

Preface bytes still to pass.

354    preface: usize,
355    header: Vec<u8>,

Payload bytes left in the current frame, and whether they are kept.

357    payload: usize,
358    keep: bool,
359}
361const PING: u8 = 0x6;
362const ACK: u8 = 0x1;
363
364impl FrameParser {
365    fn new(preface: usize) -> Self {
366        FrameParser { preface, header: Vec::with_capacity(9), payload: 0, keep: true }
367    }

Feeds bytes; frame(type, flags) says whether to keep each frame, and kept bytes are appended to out.

371    fn feed(&mut self, mut bytes: &[u8], out: &mut Vec<u8>, mut frame: impl FnMut(u8, u8) -> bool) {
372        while !bytes.is_empty() {
373            if self.preface > 0 {
374                let n = self.preface.min(bytes.len());
375                out.extend_from_slice(&bytes[..n]);
376                self.preface -= n;
377                bytes = &bytes[n..];
378            } else if self.payload > 0 {
379                let n = self.payload.min(bytes.len());
380                if self.keep {
381                    out.extend_from_slice(&bytes[..n]);
382                }
383                self.payload -= n;
384                bytes = &bytes[n..];
385            } else {
386                let n = (9 - self.header.len()).min(bytes.len());
387                self.header.extend_from_slice(&bytes[..n]);
388                bytes = &bytes[n..];
389                if self.header.len() == 9 {
390                    let h = std::mem::take(&mut self.header);
391                    self.payload = usize::from(h[0]) << 16 | usize::from(h[1]) << 8 | usize::from(h[2]);
392                    self.keep = frame(h[3], h[4]);
393                    if self.keep {
394                        out.extend_from_slice(&h);
395                    }
396                }
397            }
398        }
399    }
400}
402impl<T> Frames<T> {
403    fn new(inner: T, shared: Arc<Shared>, silent: bool) -> Self {
404        Frames {
405            inner,
406            shared,
407            incoming: FrameParser::new(24),
408            outgoing: FrameParser::new(0),
409            silent,
410            pending: Vec::new(),
411        }
412    }
413}
414
415impl<T: tokio::io::AsyncWrite + Unpin> Frames<T> {
416    fn poll_drain(&mut self, cx: &mut std::task::Context<'_>) -> std::task::Poll<std::io::Result<()>> {
417        use std::task::Poll;
418        while !self.pending.is_empty() {
419            match std::pin::Pin::new(&mut self.inner).poll_write(cx, &self.pending) {
420                Poll::Ready(Ok(0)) => return Poll::Ready(Err(std::io::ErrorKind::WriteZero.into())),
421                Poll::Ready(Ok(n)) => drop(self.pending.drain(..n)),
422                Poll::Ready(Err(e)) => return Poll::Ready(Err(e)),
423                Poll::Pending => return Poll::Pending,
424            }
425        }
426        Poll::Ready(Ok(()))
427    }
428}
429
430impl<T: tokio::io::AsyncRead + Unpin> tokio::io::AsyncRead for Frames<T> {
431    fn poll_read(
432        self: std::pin::Pin<&mut Self>,
433        cx: &mut std::task::Context<'_>,
434        buf: &mut tokio::io::ReadBuf<'_>,
435    ) -> std::task::Poll<std::io::Result<()>> {
436        let this = self.get_mut();
437        let before = buf.filled().len();
438        let polled = std::pin::Pin::new(&mut this.inner).poll_read(cx, buf);
439        if let std::task::Poll::Ready(Ok(())) = polled {
440            let shared = &this.shared;
441            // Passed through untouched; `sink` only satisfies `feed`.
442            let mut sink = Vec::new();
443            this.incoming.feed(&buf.filled()[before..], &mut sink, |kind, flags| {
444                if kind == PING && flags & ACK == 0 {
445                    shared.pings.fetch_add(1, Ordering::SeqCst);
446                }
447                true
448            });
449        }
450        polled
451    }
452}
453
454impl<T: tokio::io::AsyncWrite + Unpin> tokio::io::AsyncWrite for Frames<T> {
455    fn poll_write(
456        self: std::pin::Pin<&mut Self>,
457        cx: &mut std::task::Context<'_>,
458        buf: &[u8],
459    ) -> std::task::Poll<std::io::Result<usize>> {
460        let this = self.get_mut();
461        if let std::task::Poll::Ready(Err(e)) = this.poll_drain(cx) {
462            return std::task::Poll::Ready(Err(e));
463        }
464        if !this.pending.is_empty() {
465            return std::task::Poll::Pending;
466        }
467        let silent = this.silent;
468        this.outgoing.feed(buf, &mut this.pending, |kind, flags| !(silent && kind == PING && flags & ACK != 0));
469        // Written now if the socket takes it, else by the next write or flush.
470        if let std::task::Poll::Ready(Err(e)) = this.poll_drain(cx) {
471            return std::task::Poll::Ready(Err(e));
472        }
473        std::task::Poll::Ready(Ok(buf.len()))
474    }
475
476    fn poll_flush(self: std::pin::Pin<&mut Self>, cx: &mut std::task::Context<'_>) -> std::task::Poll<std::io::Result<()>> {
477        let this = self.get_mut();
478        std::task::ready!(this.poll_drain(cx))?;
479        std::pin::Pin::new(&mut this.inner).poll_flush(cx)
480    }
481
482    fn poll_shutdown(self: std::pin::Pin<&mut Self>, cx: &mut std::task::Context<'_>) -> std::task::Poll<std::io::Result<()>> {
483        let this = self.get_mut();
484        std::task::ready!(this.poll_drain(cx))?;
485        std::pin::Pin::new(&mut this.inner).poll_shutdown(cx)
486    }
487}