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}