jevcrates.git / jev-mock / src / forward.rs

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>)>;
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 {

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    }

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}