lib.rsannotatedlib.rssource487 lines · 17.8 KB · raw
1//! A stand-in for TypeSafe's System One endpoint: HTTP/2 over TLS on
2//! loopback, with a throwaway CA the client under test is told to trust.
3//! It records every request and answers from a closure. Any Jev client
4//! can test against it; nothing here knows about Postgres.
5
6use std::net::SocketAddr;
7use std::sync::atomic::{AtomicUsize, Ordering};
8use std::sync::{Arc, Mutex};
9use std::time::Duration;
10
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;
26
27/// 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}
36
37/// What the mock answers.
38pub struct Reply {
39    pub status: u16,
40    pub headers: Vec<(&'static str, String)>,
41    pub body: serde_json::Value,
42    /// How long to stall before answering.
43    pub delay: Option<Duration>,
44    /// Stall until the mock has recorded this many requests.
45    pub hold_until: Option<usize>,
46    /// Sent byte for byte instead of `body`, as a replayed fixture is.
47    pub raw_body: Option<Bytes>,
48}
49
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    }
54
55    /// Answers only after `delay`; a client that gives up first resets
56    /// the stream, which [`MockJev::abandoned`] counts.
57    pub fn after(mut self, delay: Duration) -> Self {
58        self.delay = Some(delay);
59        self
60    }
61
62    /// Answers only once `n` requests have arrived (then waits any
63    /// [`after`](Self::after) delay), so a test can count what a client
64    /// 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    }
69
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;
77
78/// What the server tasks share with the test.
79struct Shared {
80    requests: Mutex<Vec<Recorded>>,
81    /// How many requests have been recorded, for replies that wait on it.
82    recorded: tokio::sync::watch::Sender<usize>,
83    /// Requests whose handler was dropped before answering: the client
84    /// reset the stream (RST_STREAM) or closed the connection.
85    abandoned: AtomicUsize,
86    /// Requests arrived and not yet answered or abandoned, and the most
87    /// there have been at once.
88    in_flight: AtomicUsize,
89    peak_in_flight: AtomicUsize,
90    /// TCP connections accepted, whether or not TLS then succeeded.
91    accepts: AtomicUsize,
92    /// TLS connections accepted.
93    connections: AtomicUsize,
94    /// PING frames the client sent (not acknowledgements).
95    pings: AtomicUsize,
96    /// Connections numbered up to this one (from 1) never acknowledge a
97    /// PING, as a dead path would not.
98    silent_through: AtomicUsize,
99    /// Bumped to make every open connection send GOAWAY and close.
100    goaway: tokio::sync::watch::Sender<u64>,
101    respond: Box<Responder>,
102}
103
104pub struct MockJev {
105    pub addr: SocketAddr,
106    cert_dir: TempDir,
107    shared: Arc<Shared>,
108    task: tokio::task::JoinHandle<()>,
109}
110
111impl MockJev {
112    /// 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    }
116
117    /// Serves on 127.0.0.1 with a certificate for each of `names` (host
118    /// 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    }
125
126    /// As [`start`](Self::start), advertising
127    /// `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    }
134
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    }
218
219    /// 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    }
223
224    pub fn endpoint(&self) -> String {
225        format!("https://{}/v1/systemone", self.addr)
226    }
227
228    /// 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    }
232
233    pub fn requests(&self) -> Vec<Recorded> {
234        self.shared.requests.lock().unwrap().clone()
235    }
236
237    /// Sends GOAWAY on every open connection and closes it once idle, as
238    /// a server does when it restarts or rotates connections.
239    pub fn goaway(&self) {
240        self.shared.goaway.send_modify(|n| *n += 1);
241    }
242
243    /// The most requests the mock has held unanswered at once: a direct
244    /// 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    }
248
249    /// How many TCP connections the client opened, including those whose
250    /// TLS handshake failed.
251    pub fn accepts(&self) -> usize {
252        self.shared.accepts.load(Ordering::SeqCst)
253    }
254
255    /// How many TLS connections the client opened.
256    pub fn connections(&self) -> usize {
257        self.shared.connections.load(Ordering::SeqCst)
258    }
259
260    /// 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    }
264
265    /// The next `n` connections the client opens never acknowledge its
266    /// 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    }
270
271    /// 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}
276
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}
335
336/// The plaintext HTTP/2 byte stream between TLS and hyper, read as
337/// frames: it counts the client's PINGs, and on a silent connection drops
338/// the server's PING acknowledgements. hyper answers PINGs itself, so this
339/// is the only place to see or withhold them.
340struct Frames<T> {
341    inner: T,
342    shared: Arc<Shared>,
343    /// Client to server: the connection preface comes before any frame.
344    incoming: FrameParser,
345    outgoing: FrameParser,
346    silent: bool,
347    /// Filtered bytes accepted from hyper, not yet written.
348    pending: Vec<u8>,
349}
350
351/// Tracks frame boundaries across arbitrary chunks.
352struct FrameParser {
353    /// Preface bytes still to pass.
354    preface: usize,
355    header: Vec<u8>,
356    /// Payload bytes left in the current frame, and whether they are kept.
357    payload: usize,
358    keep: bool,
359}
360
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    }
368
369    /// Feeds `bytes`; `frame(type, flags)` says whether to keep each frame,
370    /// 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}
401
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}