Forwards each request to a real System One endpoint and records the
exchange as a [Fixture], so what Replay
later serves is exactly what the endpoint answered to exactly those
bytes. The client under test talks to the mock with any key; the real
key is added here and never recorded.
The responder is synchronous, so each request is sent with
block_in_place: the mock must run on a multi-threaded tokio runtime.
One upstream connection per request keeps this free of pooling; a
recording run sends a handful.
12use std::sync::{Arc, Mutex};
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};
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"];
Request headers passed upstream; the authorization is replaced.
28const PASSED: [&str; 2] = ["content-type", "x-typesafe-retry-count"];
Requests as sent, each with its answers in order.
31type Exchanges = Vec<(String, Vec<Response>)>;
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 }
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 }
Sends body as a JSON request of its own, not through the mock,
and records the exchange: for requests only jev-protocol can build,
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 }
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 }