1//! Forwards each request to a real System One endpoint and records the 2//! exchange as a [`Fixture`], so what [`Replay`](crate::replay::Replay) 3//! later serves is exactly what the endpoint answered to exactly those 4//! bytes. The client under test talks to the mock with any key; the real 5//! key is added here and never recorded. 6//! 7//! The responder is synchronous, so each request is sent with 8//! `block_in_place`: the mock must run on a multi-threaded tokio runtime. 9//! One upstream connection per request keeps this free of pooling; a 10//! recording run sends a handful. 11 12use std::sync::{Arc, Mutex}; 13 14use bytes::Bytes; 15use http_body_util::{BodyExt, Full}; 16use hyper_util::rt::{TokioExecutor, TokioIo}; 17use tokio_rustls::TlsConnector; 18use tokio_rustls::rustls::{ClientConfig, RootCertStore}; 19use tokio_rustls::rustls::pki_types::ServerName; 20 21use crate::replay::{Fixture, Response}; 22use crate::{Recorded, Reply}; 23 24/// Response headers kept in a fixture: the ones the client reads. 25const KEPT: [&str; 4] = ["content-type", "x-typesafe-request-id", "retry-after", "retry-after-ms"]; 26 27/// Request headers passed upstream; the authorization is replaced. 28const PASSED: [&str; 2] = ["content-type", "x-typesafe-retry-count"]; 29 30/// Requests as sent, each with its answers in order. 31type Exchanges = Vec<(String, Vec<Response>)>; 32 33#[derive(Clone)] 34pub struct Forward { 35 host: String, 36 port: u16, 37 path: String, 38 authorization: String, 39 model: String, 40 exchanges: Arc<Mutex<Exchanges>>, 41} 42 43impl Forward { 44 /// `upstream` is an `https://host[:port]/path` URL. 45 pub fn new(upstream: &str, api_key: &str, model: &str) -> Result<Self, String> { 46 let rest = upstream.strip_prefix("https://").ok_or("the upstream must be https://")?; 47 let (authority, path) = rest.split_once('/').map(|(a, p)| (a, format!("/{p}"))).unwrap_or((rest, "/".into())); 48 let (host, port) = match authority.rsplit_once(':') { 49 Some((h, p)) => (h.to_owned(), p.parse().map_err(|_| format!("bad port in {upstream}"))?), 50 None => (authority.to_owned(), 443), 51 }; 52 Ok(Forward { 53 host, 54 port, 55 path, 56 authorization: format!("Bearer {api_key}"), 57 model: model.to_owned(), 58 exchanges: Arc::default(), 59 }) 60 } 61 62 pub fn responder(&self) -> impl Fn(&Recorded) -> Reply + Send + Sync + 'static { 63 let this = self.clone(); 64 move |req| this.answer(req) 65 } 66 67 fn answer(&self, req: &Recorded) -> Reply { 68 let sent = 69 tokio::task::block_in_place(|| tokio::runtime::Handle::current().block_on(self.send(&req.headers, &req.body))); 70 let response = match sent { 71 Ok(r) => r, 72 // A failure to reach the endpoint is not an answer: the client 73 // sees a non-retryable 400 and nothing is recorded. 74 Err(e) => { 75 return Reply::json(400, serde_json::json!({ "detail": { "error_type": "forward_failed", "message": e } })); 76 } 77 }; 78 let mut reply = Reply::json(response.status, serde_json::Value::Null); 79 reply.raw_body = Some(Bytes::from(response.body.clone())); 80 for (name, value) in &response.headers { 81 let name: &'static str = KEPT.iter().find(|k| **k == name).expect("only kept headers are recorded"); 82 reply.headers.push((name, value.clone())); 83 } 84 self.record(&req.body, response); 85 reply 86 } 87 88 /// Sends `body` as a JSON request of its own, not through the mock, 89 /// and records the exchange: for requests only jev-protocol can build, 90 /// such as a layout the extension cannot reach yet. 91 pub async fn exchange(&self, body: &[u8]) -> Result<Response, String> { 92 let mut headers = http::HeaderMap::new(); 93 headers.insert("content-type", http::HeaderValue::from_static("application/json")); 94 let response = self.send(&headers, &Bytes::copy_from_slice(body)).await?; 95 self.record(body, response.clone()); 96 Ok(response) 97 } 98 99 fn record(&self, body: &[u8], response: Response) { 100 let request = String::from_utf8_lossy(body).into_owned(); 101 let mut exchanges = self.exchanges.lock().unwrap(); 102 // A retried request (429, 5xx) records its answers in order. 103 match exchanges.iter_mut().find(|(r, _)| *r == request) { 104 Some((_, responses)) => responses.push(response), 105 None => exchanges.push((request, vec![response])), 106 } 107 } 108 109 async fn send(&self, headers: &http::HeaderMap, body: &Bytes) -> Result<Response, String> { 110 let mut roots = RootCertStore::empty(); 111 roots.extend(webpki_roots::TLS_SERVER_ROOTS.iter().cloned()); 112 let mut tls = ClientConfig::builder().with_root_certificates(roots).with_no_client_auth(); 113 tls.alpn_protocols = vec![b"h2".to_vec()]; 114 let name = ServerName::try_from(self.host.clone()).map_err(|e| e.to_string())?; 115 let tcp = tokio::net::TcpStream::connect((self.host.as_str(), self.port)).await.map_err(|e| e.to_string())?; 116 let tls = TlsConnector::from(Arc::new(tls)).connect(name, tcp).await.map_err(|e| e.to_string())?; 117 let (mut sender, conn) = hyper::client::conn::http2::handshake(TokioExecutor::new(), TokioIo::new(tls)) 118 .await 119 .map_err(|e| e.to_string())?; 120 tokio::spawn(conn); 121 122 let mut builder = hyper::Request::post(format!("https://{}:{}{}", self.host, self.port, self.path)) 123 .header("authorization", &self.authorization); 124 for name in PASSED { 125 if let Some(value) = headers.get(name) { 126 builder = builder.header(name, value); 127 } 128 } 129 let request = builder.body(Full::new(body.clone())).map_err(|e| e.to_string())?; 130 let response = sender.send_request(request).await.map_err(|e| e.to_string())?; 131 let status = response.status().as_u16(); 132 let headers = KEPT 133 .iter() 134 .filter_map(|k| response.headers().get(*k).map(|v| (k.to_string(), v.to_str().unwrap_or_default().to_owned()))) 135 .collect(); 136 let body = response.into_body().collect().await.map_err(|e| e.to_string())?.to_bytes(); 137 let body = String::from_utf8(body.to_vec()).map_err(|_| "the endpoint answered with non-UTF-8 bytes".to_owned())?; 138 Ok(Response { status, headers, body }) 139 } 140 141 /// Every exchange so far, in the order first sent. 142 pub fn fixture(&self) -> Fixture { 143 Fixture { model: self.model.clone(), exchanges: self.exchanges.lock().unwrap().clone() } 144 } 145}