1use 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#[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#[derive(Debug, Clone)]
32pub struct Seen {
33 pub method: String,
34 pub path: String,
36 pub headers: Vec<(String, String)>,
38 pub body: String,
40 pub body_bytes: Vec<u8>,
42}
43
44pub type Request = Seen;
46
47impl Seen {
48 pub fn route(&self) -> &str {
50 self.path.split('?').next().unwrap_or("")
51 }
52 pub fn query(&self) -> Option<&str> {
54 self.path.split_once('?').map(|(_, q)| q)
55 }
56 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 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 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#[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 pub fn json(v: Value) -> Self {
101 Self::status(200, v)
102 }
103 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 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 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 pub fn text(body: &str) -> Self {
133 Self::bytes("text/plain; charset=utf-8", body.as_bytes().to_vec())
134 }
135 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
149impl 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 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 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 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 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}