Skip to main content

rightkit_qa/
mock.rs

1//! Local HTTP mock for PAID providers: paid-tier scenarios assert the exact
2//! requests the product would send, and never make a billed call.
3//!
4//! Two ways to answer:
5//! * [`MockServer::start`]: static [`Route`]s with JSON bodies (the declarative `mock_start` step).
6//! * [`MockServer::start_with`]: a handler closure over the request and the mock's own base URL,
7//!   returning a [`Reply`] with any status, content type, headers and binary body. Use it for
8//!   stateful flows (poll until done, fail the Nth call) and generated media (PNG/MP3/MP4).
9//!
10//! Every connection is served on its own thread, so a slow handler or a client holding a
11//! connection open never blocks the next request. Every request is recorded ([`MockServer::seen`]).
12use crate::util::Result;
13use serde_json::{json, Value};
14use std::io::{Read, Write};
15use std::net::{SocketAddr, TcpListener, TcpStream};
16use std::sync::atomic::{AtomicBool, Ordering};
17use std::sync::{Arc, Mutex};
18use std::time::Duration;
19
20/// Static route: first route whose method (if any) and path (query ignored) match answers.
21#[derive(Debug, Clone)]
22pub struct Route {
23    pub method: Option<String>,
24    pub path: String,
25    pub status: u16,
26    pub body: Value,
27    pub content_type: Option<String>,
28}
29
30/// One recorded request.
31#[derive(Debug, Clone)]
32pub struct Seen {
33    pub method: String,
34    /// Request target as sent, including any query string.
35    pub path: String,
36    /// Header names as sent; use [`Seen::header`] for case-insensitive lookup.
37    pub headers: Vec<(String, String)>,
38    /// Body as (lossy) UTF-8 text.
39    pub body: String,
40    /// Raw body bytes (multipart uploads, binary posts).
41    pub body_bytes: Vec<u8>,
42}
43
44/// The request a [`MockServer::start_with`] handler receives.
45pub type Request = Seen;
46
47impl Seen {
48    /// Path without the query string.
49    pub fn route(&self) -> &str {
50        self.path.split('?').next().unwrap_or("")
51    }
52    /// The query string, without `?`.
53    pub fn query(&self) -> Option<&str> {
54        self.path.split_once('?').map(|(_, q)| q)
55    }
56    /// One query parameter (no percent-decoding).
57    pub fn query_param(&self, name: &str) -> Option<&str> {
58        self.query()?
59            .split('&')
60            .filter_map(|kv| kv.split_once('=').or(Some((kv, ""))))
61            .find(|(k, _)| *k == name)
62            .map(|(_, v)| v)
63    }
64    /// Case-insensitive header value.
65    pub fn header(&self, name: &str) -> Option<&str> {
66        self.headers
67            .iter()
68            .find(|(k, _)| k.eq_ignore_ascii_case(name))
69            .map(|(_, v)| v.as_str())
70    }
71    /// Body parsed as JSON, or `Null`.
72    pub fn json(&self) -> Value {
73        serde_json::from_slice(&self.body_bytes).unwrap_or(Value::Null)
74    }
75    pub fn is(&self, method: &str, route: &str) -> bool {
76        self.method.eq_ignore_ascii_case(method) && self.route() == route
77    }
78    pub fn to_value(&self) -> Value {
79        let headers: serde_json::Map<String, Value> = self
80            .headers
81            .iter()
82            .map(|(k, v)| (k.to_ascii_lowercase(), json!(v)))
83            .collect();
84        json!({"method": self.method, "path": self.path, "headers": headers, "body": self.body,
85               "json": self.json()})
86    }
87}
88
89/// A handler's answer: any status, content type, extra headers and binary body.
90#[derive(Debug, Clone, PartialEq, Eq)]
91pub struct Reply {
92    pub status: u16,
93    pub content_type: String,
94    pub headers: Vec<(String, String)>,
95    pub body: Vec<u8>,
96}
97
98impl Reply {
99    /// 200 with a JSON body.
100    pub fn json(v: Value) -> Self {
101        Self::status(200, v)
102    }
103    /// Any status with a JSON body.
104    pub fn status(status: u16, v: Value) -> Self {
105        Self {
106            status,
107            content_type: "application/json".into(),
108            headers: vec![],
109            body: v.to_string().into_bytes(),
110        }
111    }
112    /// 200 with a binary body (generated media: `image/png`, `audio/mpeg`, `video/mp4`, ...).
113    pub fn bytes(content_type: &str, body: impl Into<Vec<u8>>) -> Self {
114        Self {
115            status: 200,
116            content_type: content_type.into(),
117            headers: vec![],
118            body: body.into(),
119        }
120    }
121    /// 200 with a file's bytes; an unreadable file is a 500 naming it.
122    pub fn file(content_type: &str, path: &std::path::Path) -> Self {
123        match std::fs::read(path) {
124            Ok(b) => Self::bytes(content_type, b),
125            Err(e) => Self::status(
126                500,
127                json!({"error": format!("mock could not read {}: {e}", path.display())}),
128            ),
129        }
130    }
131    /// 200 `text/plain`.
132    pub fn text(body: &str) -> Self {
133        Self::bytes("text/plain; charset=utf-8", body.as_bytes().to_vec())
134    }
135    /// 404 `{"error":"unmocked"}` (what an unmatched request gets).
136    pub fn unmocked() -> Self {
137        Self::status(404, json!({"error": "unmocked"}))
138    }
139    pub fn with_status(mut self, status: u16) -> Self {
140        self.status = status;
141        self
142    }
143    pub fn with_header(mut self, name: &str, value: &str) -> Self {
144        self.headers.push((name.into(), value.into()));
145        self
146    }
147}
148
149/// `None` from a handler means "not mocked" (404).
150impl From<Option<Reply>> for Reply {
151    fn from(r: Option<Reply>) -> Self {
152        r.unwrap_or_else(Reply::unmocked)
153    }
154}
155
156type Handler = dyn Fn(&Request, &str) -> Reply + Send + Sync;
157
158pub struct MockServer {
159    pub base: String,
160    addr: SocketAddr,
161    seen: Arc<Mutex<Vec<Seen>>>,
162    stop: Arc<AtomicBool>,
163}
164
165impl MockServer {
166    /// Static JSON routes; unmatched requests get 404 `{"error":"unmocked"}`.
167    pub fn start(routes: Vec<Route>) -> Result<Self> {
168        Self::start_with(move |req: &Request, _base: &str| {
169            routes
170                .iter()
171                .find(|r| {
172                    r.path == req.route()
173                        && r.method
174                            .as_deref()
175                            .map(|m| m.eq_ignore_ascii_case(&req.method))
176                            .unwrap_or(true)
177                })
178                .map(|r| Reply {
179                    status: r.status,
180                    content_type: r
181                        .content_type
182                        .clone()
183                        .unwrap_or_else(|| "application/json".into()),
184                    headers: vec![],
185                    body: r.body.to_string().into_bytes(),
186                })
187        })
188    }
189
190    /// Dynamic handler: `handler(request, base_url)` answers every request. `base_url`
191    /// (`http://127.0.0.1:<port>`) lets a reply point the client back at the mock (a job's
192    /// result URL). Return a [`Reply`], or an `Option<Reply>` where `None` means 404. A
193    /// panicking handler yields a 500 instead of killing the server.
194    pub fn start_with<R: Into<Reply>>(
195        handler: impl Fn(&Request, &str) -> R + Send + Sync + 'static,
196    ) -> Result<Self> {
197        let listener = TcpListener::bind(("127.0.0.1", 0))?;
198        let addr = listener.local_addr()?;
199        let base = format!("http://{addr}");
200        let seen = Arc::new(Mutex::new(Vec::new()));
201        let stop = Arc::new(AtomicBool::new(false));
202        let handler: Arc<Handler> = Arc::new(move |r: &Request, b: &str| handler(r, b).into());
203        let (s2, st2, b2) = (seen.clone(), stop.clone(), base.clone());
204        std::thread::spawn(move || {
205            for conn in listener.incoming().flatten() {
206                if st2.load(Ordering::SeqCst) {
207                    break;
208                }
209                let (s3, h, b) = (s2.clone(), handler.clone(), b2.clone());
210                std::thread::spawn(move || {
211                    let _ = serve(conn, &s3, &b, &*h);
212                });
213            }
214        });
215        Ok(Self {
216            base,
217            addr,
218            seen,
219            stop,
220        })
221    }
222
223    pub fn port(&self) -> u16 {
224        self.addr.port()
225    }
226    /// `base + path`.
227    pub fn url(&self, path: &str) -> String {
228        format!("{}{path}", self.base)
229    }
230    pub fn seen(&self) -> Vec<Seen> {
231        self.seen.lock().unwrap().clone()
232    }
233    pub fn close(&self) {
234        self.stop.store(true, Ordering::SeqCst);
235        let _ = TcpStream::connect_timeout(&self.addr, Duration::from_millis(200));
236    }
237}
238
239impl Drop for MockServer {
240    fn drop(&mut self) {
241        self.close();
242    }
243}
244
245fn serve(
246    mut s: TcpStream,
247    seen: &Mutex<Vec<Seen>>,
248    base: &str,
249    handler: &Handler,
250) -> std::io::Result<()> {
251    s.set_read_timeout(Some(Duration::from_secs(10)))?;
252    let mut buf = Vec::new();
253    let mut chunk = [0u8; 16384];
254    let (head_end, content_len) = loop {
255        let n = s.read(&mut chunk)?;
256        if n == 0 {
257            return Ok(());
258        }
259        buf.extend_from_slice(&chunk[..n]);
260        if let Some(p) = buf.windows(4).position(|w| w == b"\r\n\r\n") {
261            let head = String::from_utf8_lossy(&buf[..p]).to_ascii_lowercase();
262            let len = head
263                .lines()
264                .find_map(|l| {
265                    l.strip_prefix("content-length:")
266                        .and_then(|v| v.trim().parse::<usize>().ok())
267                })
268                .unwrap_or(0);
269            break (p + 4, len);
270        }
271    };
272    while buf.len() < head_end + content_len {
273        let n = s.read(&mut chunk)?;
274        if n == 0 {
275            break;
276        }
277        buf.extend_from_slice(&chunk[..n]);
278    }
279    let head = String::from_utf8_lossy(&buf[..head_end]).to_string();
280    let mut lines = head.lines();
281    let mut first = lines.next().unwrap_or("").split_whitespace();
282    let (method, path) = (
283        first.next().unwrap_or("").to_string(),
284        first.next().unwrap_or("").to_string(),
285    );
286    let headers = lines
287        .filter_map(|l| l.split_once(':'))
288        .map(|(k, v)| (k.trim().to_string(), v.trim().to_string()))
289        .collect();
290    let body_bytes = buf[head_end..(head_end + content_len).min(buf.len())].to_vec();
291    let req = Seen {
292        method,
293        path,
294        headers,
295        body: String::from_utf8_lossy(&body_bytes).into_owned(),
296        body_bytes,
297    };
298    seen.lock().unwrap().push(req.clone());
299    let reply = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| handler(&req, base)))
300        .unwrap_or_else(|_| Reply::status(500, json!({"error": "mock handler panicked"})));
301    let mut head = format!(
302        "HTTP/1.1 {} X\r\nContent-Type: {}\r\nContent-Length: {}\r\nConnection: close\r\n",
303        reply.status,
304        reply.content_type,
305        reply.body.len()
306    );
307    for (k, v) in &reply.headers {
308        head.push_str(&format!("{k}: {v}\r\n"));
309    }
310    head.push_str("\r\n");
311    s.write_all(head.as_bytes())?;
312    if req.method != "HEAD" {
313        s.write_all(&reply.body)?;
314    }
315    s.flush()
316}
317
318#[cfg(test)]
319mod tests {
320    use super::*;
321    use crate::http::request;
322    use std::sync::atomic::AtomicUsize;
323
324    fn get(m: &MockServer, method: &str, path: &str, body: Option<&str>) -> crate::http::Response {
325        request(
326            m.addr,
327            method,
328            path,
329            Some("Bearer k"),
330            body,
331            Duration::from_secs(5),
332        )
333        .unwrap()
334    }
335
336    #[test]
337    fn static_routes_still_answer_json_and_404() {
338        let m = MockServer::start(vec![Route {
339            method: Some("POST".into()),
340            path: "/v1/gen".into(),
341            status: 201,
342            body: json!({"id": "j1"}),
343            content_type: None,
344        }])
345        .unwrap();
346        let r = get(&m, "POST", "/v1/gen?x=1", Some(r#"{"prompt":"p"}"#));
347        assert_eq!(r.status, 201);
348        assert_eq!(r.content_type, "application/json");
349        assert_eq!(
350            serde_json::from_slice::<Value>(&r.body).unwrap()["id"],
351            "j1"
352        );
353        assert_eq!(get(&m, "GET", "/v1/gen", None).status, 404);
354        let seen = m.seen();
355        assert_eq!(seen.len(), 2);
356        assert_eq!(seen[0].json()["prompt"], "p");
357        assert_eq!(seen[0].header("authorization"), Some("Bearer k"));
358        assert_eq!(seen[0].query_param("x"), Some("1"));
359        assert_eq!(seen[0].to_value()["json"]["prompt"], "p");
360    }
361
362    #[test]
363    fn dynamic_handler_is_stateful_binary_and_knows_its_base() {
364        let png: Vec<u8> = (0..=255u8).cycle().take(70_000).collect();
365        let polls = Arc::new(AtomicUsize::new(0));
366        let (p2, png2) = (polls.clone(), png.clone());
367        let m = MockServer::start_with(move |req: &Request, base: &str| {
368            if req.is("POST", "/jobs") {
369                Some(Reply::json(json!({"poll": format!("{base}/jobs/1")})))
370            } else if req.is("GET", "/jobs/1") {
371                let n = p2.fetch_add(1, Ordering::SeqCst);
372                Some(if n < 2 {
373                    Reply::json(json!({"status": "running"}))
374                } else {
375                    Reply::json(json!({"status": "done", "url": format!("{base}/out.png")}))
376                })
377            } else if req.is("GET", "/out.png") {
378                Some(Reply::bytes("image/png", png2.clone()).with_header("X-Mock", "1"))
379            } else if req.is("GET", "/boom") {
380                panic!("handler bug")
381            } else {
382                None
383            }
384        })
385        .unwrap();
386        let first: Value =
387            serde_json::from_slice(&get(&m, "POST", "/jobs", Some("{}")).body).unwrap();
388        assert_eq!(first["poll"], m.url("/jobs/1"));
389        let states: Vec<String> = (0..3)
390            .map(|_| {
391                serde_json::from_slice::<Value>(&get(&m, "GET", "/jobs/1", None).body).unwrap()
392                    ["status"]
393                    .as_str()
394                    .unwrap()
395                    .to_string()
396            })
397            .collect();
398        assert_eq!(states, ["running", "running", "done"]);
399        let media = get(&m, "GET", "/out.png", None);
400        assert_eq!(media.status, 200);
401        assert_eq!(media.content_type, "image/png");
402        assert_eq!(media.body, png, "binary body survives byte-for-byte");
403        assert_eq!(get(&m, "GET", "/boom", None).status, 500);
404        assert_eq!(get(&m, "GET", "/nope", None).status, 404);
405        assert_eq!(m.seen().len(), 7);
406    }
407
408    #[test]
409    fn a_held_connection_does_not_block_other_requests() {
410        let m = MockServer::start_with(|_: &Request, _: &str| Reply::text("ok")).unwrap();
411        // A client that connects and never sends a request.
412        let _idle = TcpStream::connect(m.addr).unwrap();
413        let started = std::time::Instant::now();
414        let r = get(&m, "GET", "/", None);
415        assert_eq!(r.text(), "ok");
416        assert!(started.elapsed() < Duration::from_secs(3));
417    }
418}