1use std::io::{ErrorKind, Read, Write};
10use std::net::{Shutdown as NetShutdown, SocketAddr, TcpListener, TcpStream, ToSocketAddrs};
11use std::os::unix::fs::{DirBuilderExt, FileTypeExt, MetadataExt, PermissionsExt};
12use std::os::unix::net::{UnixListener, UnixStream};
13use std::path::{Path, PathBuf};
14use std::sync::Arc;
15use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
16use std::time::{Duration, Instant};
17
18use serde_json::Value;
19
20use crate::error::{Error, Result};
21
22#[derive(Debug, Clone)]
24pub struct Limits {
25 pub max_header_bytes: usize,
27 pub max_body_bytes: usize,
28 pub read_timeout: Duration,
30 pub write_timeout: Duration,
31 pub max_connections: usize,
34}
35
36impl Default for Limits {
37 fn default() -> Self {
38 Limits {
39 max_header_bytes: 64 * 1024,
40 max_body_bytes: 4 * 1024 * 1024,
41 read_timeout: Duration::from_secs(30),
42 write_timeout: Duration::from_secs(30),
43 max_connections: 256,
44 }
45 }
46}
47
48#[derive(Debug, Clone, PartialEq, Eq)]
50pub enum Peer {
51 Tcp(SocketAddr),
52 Unix {
54 uid: Option<u32>,
55 },
56}
57
58#[derive(Debug, Clone)]
59pub struct Request {
60 pub method: String,
61 pub path: String,
62 pub query: Option<String>,
63 pub headers: Vec<(String, String)>,
65 pub body: Vec<u8>,
66 pub peer: Peer,
67}
68
69impl Request {
70 pub fn header(&self, name: &str) -> Option<&str> {
72 self.headers
73 .iter()
74 .find(|(k, _)| k.eq_ignore_ascii_case(name))
75 .map(|(_, v)| v.as_str())
76 }
77}
78
79pub type StreamFn = Box<dyn FnOnce(&mut dyn Write) -> std::io::Result<()> + Send>;
82
83pub trait Duplex: Read + Write + Send {
86 fn set_read_timeout(&mut self, t: Option<Duration>) -> std::io::Result<()>;
87}
88
89impl Duplex for TcpStream {
90 fn set_read_timeout(&mut self, t: Option<Duration>) -> std::io::Result<()> {
91 TcpStream::set_read_timeout(self, t)
92 }
93}
94
95impl Duplex for UnixStream {
96 fn set_read_timeout(&mut self, t: Option<Duration>) -> std::io::Result<()> {
97 UnixStream::set_read_timeout(self, t)
98 }
99}
100
101pub type UpgradeFn = Box<dyn FnOnce(&mut dyn Duplex) + Send>;
103
104#[derive(Clone)]
105pub struct Response {
106 pub status: u16,
107 pub headers: Vec<(String, String)>,
108 pub body: Vec<u8>,
109 pub stream: Option<Arc<std::sync::Mutex<Option<StreamFn>>>>,
111 pub upgrade: Option<Arc<std::sync::Mutex<Option<UpgradeFn>>>>,
113}
114
115impl std::fmt::Debug for Response {
116 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
117 f.debug_struct("Response")
118 .field("status", &self.status)
119 .field("headers", &self.headers)
120 .field("body", &self.body.len())
121 .field("stream", &self.stream.is_some())
122 .field("upgrade", &self.upgrade.is_some())
123 .finish()
124 }
125}
126
127impl Response {
128 pub fn new(status: u16) -> Self {
129 Response {
130 status,
131 headers: Vec::new(),
132 body: Vec::new(),
133 stream: None,
134 upgrade: None,
135 }
136 }
137
138 pub fn upgrade(protocol: &str, f: UpgradeFn) -> Self {
140 let mut r = Response::new(101).header("Upgrade", protocol);
141 r.upgrade = Some(Arc::new(std::sync::Mutex::new(Some(f))));
142 r
143 }
144
145 pub fn stream(status: u16, content_type: &str, f: StreamFn) -> Self {
148 let mut r = Response::new(status).header("Content-Type", content_type);
149 r.stream = Some(Arc::new(std::sync::Mutex::new(Some(f))));
150 r
151 }
152
153 pub fn json(status: u16, v: &Value) -> Self {
154 let mut body = serde_json::to_vec(v).unwrap_or_default();
155 body.push(b'\n');
156 Response::new(status)
157 .header("Content-Type", "application/json")
158 .body(body)
159 }
160
161 pub fn text(status: u16, s: &str) -> Self {
162 Response::new(status)
163 .header("Content-Type", "text/plain; charset=utf-8")
164 .body(format!("{s}\n").into_bytes())
165 }
166
167 pub fn header(mut self, name: impl Into<String>, value: impl Into<String>) -> Self {
168 self.headers.push((name.into(), value.into()));
169 self
170 }
171
172 pub fn body(mut self, body: impl Into<Vec<u8>>) -> Self {
173 self.body = body.into();
174 self
175 }
176
177 pub fn get_header(&self, name: &str) -> Option<&str> {
178 self.headers
179 .iter()
180 .find(|(k, _)| k.eq_ignore_ascii_case(name))
181 .map(|(_, v)| v.as_str())
182 }
183}
184
185pub type Handler = Arc<dyn Fn(&Request) -> Response + Send + Sync>;
186
187#[derive(Debug, Clone, Default)]
189pub struct Shutdown(Arc<AtomicBool>);
190
191impl Shutdown {
192 pub fn new() -> Self {
193 Self::default()
194 }
195
196 pub fn on_signals() -> Result<Self> {
198 let s = Self::new();
199 for sig in [signal_hook::consts::SIGINT, signal_hook::consts::SIGTERM] {
200 signal_hook::flag::register(sig, s.0.clone())?;
201 }
202 Ok(s)
203 }
204
205 pub fn trigger(&self) {
206 self.0.store(true, Ordering::SeqCst);
207 }
208
209 pub fn is_triggered(&self) -> bool {
210 self.0.load(Ordering::SeqCst)
211 }
212}
213
214#[derive(Debug)]
216pub enum HttpListener {
217 Tcp(TcpListener),
218 Unix(UnixSocket),
219}
220
221#[derive(Debug)]
224pub struct UnixSocket {
225 listener: UnixListener,
226 path: PathBuf,
227 ino: (u64, u64),
228}
229
230impl Drop for UnixSocket {
231 fn drop(&mut self) {
232 if std::fs::symlink_metadata(&self.path).is_ok_and(|m| (m.dev(), m.ino()) == self.ino) {
233 let _ = std::fs::remove_file(&self.path);
234 }
235 }
236}
237
238impl HttpListener {
239 pub fn bind_tcp(addr: &str) -> Result<Self> {
243 let addrs: Vec<SocketAddr> = addr
244 .to_socket_addrs()
245 .map_err(|e| Error::invalid(format!("listen address {addr:?}: {e}")))?
246 .collect();
247 if addrs.is_empty() || addrs.iter().any(|a| !a.ip().is_loopback()) {
248 return Err(Error::invalid(format!(
249 "listen address {addr:?} is not loopback; expose it through a Cloudflare Tunnel instead"
250 )));
251 }
252 let l = TcpListener::bind(addrs[0])?;
253 l.set_nonblocking(true)?;
254 Ok(HttpListener::Tcp(l))
255 }
256
257 pub fn bind_tcp_tailnet(addr: &str) -> Result<Self> {
261 let addrs: Vec<SocketAddr> = addr
262 .to_socket_addrs()
263 .map_err(|e| Error::invalid(format!("listen address {addr:?}: {e}")))?
264 .collect();
265 if addrs.is_empty() || addrs.iter().any(|a| !super::tailnet::is_tailnet_ip(a.ip())) {
266 return Err(Error::invalid(format!(
267 "listen address {addr:?} is not a tailnet address (100.64.0.0/10, fd7a:115c:a1e0::/48)"
268 )));
269 }
270 let l = TcpListener::bind(addrs[0])?;
271 l.set_nonblocking(true)?;
272 Ok(HttpListener::Tcp(l))
273 }
274
275 pub fn bind_tcp_private(addr: SocketAddr) -> Result<Self> {
279 let ok = match addr.ip() {
280 std::net::IpAddr::V4(v4) => v4.is_private(),
281 std::net::IpAddr::V6(_) => false,
282 };
283 if !ok {
284 return Err(Error::invalid(format!(
285 "listen address {addr} is not a private IPv4 address"
286 )));
287 }
288 let l = TcpListener::bind(addr)?;
289 l.set_nonblocking(true)?;
290 Ok(HttpListener::Tcp(l))
291 }
292
293 pub fn bind_unix(path: &Path) -> Result<Self> {
297 if let Some(dir) = path.parent().filter(|d| !d.as_os_str().is_empty()) {
298 std::fs::DirBuilder::new()
299 .recursive(true)
300 .mode(0o700)
301 .create(dir)?;
302 }
303 match std::fs::symlink_metadata(path) {
304 Ok(m) if m.file_type().is_socket() => {
305 if UnixStream::connect(path).is_ok() {
306 return Err(Error::invalid(format!(
307 "{} is in use by another server",
308 path.display()
309 )));
310 }
311 std::fs::remove_file(path)?;
312 }
313 Ok(_) => {
314 return Err(Error::invalid(format!(
315 "{} exists and is not a socket",
316 path.display()
317 )));
318 }
319 Err(e) if e.kind() == ErrorKind::NotFound => {}
320 Err(e) => return Err(e.into()),
321 }
322 let listener = UnixListener::bind(path)?;
323 std::fs::set_permissions(path, std::fs::Permissions::from_mode(0o600))?;
324 listener.set_nonblocking(true)?;
325 let m = std::fs::symlink_metadata(path)?;
326 Ok(HttpListener::Unix(UnixSocket {
327 listener,
328 path: path.to_path_buf(),
329 ino: (m.dev(), m.ino()),
330 }))
331 }
332
333 pub fn local_addr(&self) -> Option<SocketAddr> {
334 match self {
335 HttpListener::Tcp(l) => l.local_addr().ok(),
336 HttpListener::Unix(_) => None,
337 }
338 }
339
340 fn poll(&self, timeout: Duration) -> bool {
341 use rustix::event::{PollFd, PollFlags, Timespec, poll};
342 let ts = Timespec {
343 tv_sec: timeout.as_secs() as _,
344 tv_nsec: timeout.subsec_nanos() as _,
345 };
346 let r = match self {
347 HttpListener::Tcp(l) => poll(&mut [PollFd::new(l, PollFlags::IN)], Some(&ts)),
348 HttpListener::Unix(u) => {
349 poll(&mut [PollFd::new(&u.listener, PollFlags::IN)], Some(&ts))
350 }
351 };
352 matches!(r, Ok(n) if n > 0)
353 }
354}
355
356#[derive(Debug)]
358pub struct HttpServer {
359 limits: Limits,
360 active: AtomicUsize,
361 shutdown: Shutdown,
362}
363
364struct Slot(Arc<HttpServer>);
367
368impl Drop for Slot {
369 fn drop(&mut self) {
370 self.0.active.fetch_sub(1, Ordering::SeqCst);
371 }
372}
373
374impl HttpServer {
375 pub fn new(limits: Limits, shutdown: Shutdown) -> Arc<Self> {
376 Arc::new(HttpServer {
377 limits,
378 active: AtomicUsize::new(0),
379 shutdown,
380 })
381 }
382
383 pub fn limits(&self) -> &Limits {
384 &self.limits
385 }
386
387 pub fn active(&self) -> usize {
389 self.active.load(Ordering::SeqCst)
390 }
391
392 pub fn run(self: &Arc<Self>, listener: HttpListener, handler: Handler) -> Result<()> {
395 while !self.shutdown.is_triggered() {
396 if !listener.poll(Duration::from_millis(250)) {
398 continue;
399 }
400 let accepted = match &listener {
401 HttpListener::Tcp(l) => l.accept().map(|(s, a)| Conn::Tcp(s, a)),
402 HttpListener::Unix(u) => u.listener.accept().map(|(s, _)| Conn::Unix(s)),
403 };
404 let conn = match accepted {
405 Ok(c) => c,
406 Err(e) if matches!(e.kind(), ErrorKind::WouldBlock | ErrorKind::Interrupted) => {
407 continue;
408 }
409 Err(e) => {
410 eprintln!("isb serve: accept: {e}");
412 std::thread::sleep(Duration::from_millis(100));
413 continue;
414 }
415 };
416 if self.active.fetch_add(1, Ordering::SeqCst) >= self.limits.max_connections {
417 self.active.fetch_sub(1, Ordering::SeqCst);
418 conn.reject_busy();
419 continue;
420 }
421 let slot = Slot(self.clone());
422 let h = handler.clone();
423 let spawned = std::thread::Builder::new()
424 .name("isb-http".into())
425 .spawn(move || {
426 let s = slot;
427 conn.serve(&s.0.limits, &h);
428 });
429 if let Err(e) = spawned {
430 eprintln!("isb serve: cannot spawn connection thread: {e}");
431 }
432 }
433 Ok(())
434 }
435
436 pub fn drain(&self, timeout: Duration) {
438 let started = Instant::now();
439 while self.active() > 0 && started.elapsed() < timeout {
440 std::thread::sleep(Duration::from_millis(20));
441 }
442 }
443}
444
445enum Conn {
446 Tcp(TcpStream, SocketAddr),
447 Unix(UnixStream),
448}
449
450impl Conn {
451 fn reject_busy(self) {
455 fn refuse<S: Read + Write>(s: &mut S) {
456 let r = Response::text(503, "server busy").header("Retry-After", "1");
457 let _ = write_response(s, &r);
458 let mut sink = [0u8; 16384];
459 for _ in 0..4 {
460 if !matches!(s.read(&mut sink), Ok(n) if n > 0) {
461 break;
462 }
463 }
464 }
465 match self {
468 Conn::Tcp(mut s, _) => {
469 let _ = s.set_nonblocking(true);
470 refuse(&mut s);
471 let _ = s.shutdown(NetShutdown::Write);
472 }
473 Conn::Unix(mut s) => {
474 let _ = s.set_nonblocking(true);
475 refuse(&mut s);
476 let _ = s.shutdown(NetShutdown::Write);
477 }
478 }
479 }
480
481 fn serve(self, limits: &Limits, handler: &Handler) {
482 match self {
483 Conn::Tcp(s, addr) => {
484 let _ = s.set_nonblocking(false);
485 let _ = s.set_nodelay(true);
486 let _ = s.set_read_timeout(Some(limits.read_timeout));
487 let _ = s.set_write_timeout(Some(limits.write_timeout));
488 let mut s = s;
489 handle(&mut s, Peer::Tcp(addr), limits, handler);
490 let _ = s.shutdown(NetShutdown::Write);
491 }
492 Conn::Unix(s) => {
493 let _ = s.set_nonblocking(false);
494 let _ = s.set_read_timeout(Some(limits.read_timeout));
495 let _ = s.set_write_timeout(Some(limits.write_timeout));
496 let uid = peer_uid(&s);
497 let mut s = s;
498 handle(&mut s, Peer::Unix { uid }, limits, handler);
499 let _ = s.shutdown(NetShutdown::Write);
500 }
501 }
502 }
503}
504
505#[cfg(target_os = "linux")]
506fn peer_uid(s: &UnixStream) -> Option<u32> {
507 rustix::net::sockopt::socket_peercred(s)
508 .ok()
509 .map(|c| c.uid.as_raw())
510}
511
512#[cfg(target_os = "macos")]
513fn peer_uid(s: &UnixStream) -> Option<u32> {
514 use std::os::fd::AsRawFd;
515 let (mut uid, mut gid) = (0, 0);
516 (unsafe { libc::getpeereid(s.as_raw_fd(), &mut uid, &mut gid) } == 0).then_some(uid)
518}
519
520#[cfg(not(any(target_os = "linux", target_os = "macos")))]
521fn peer_uid(_: &UnixStream) -> Option<u32> {
522 None
523}
524
525pub(crate) fn handle<S: Duplex>(stream: &mut S, peer: Peer, limits: &Limits, handler: &Handler) {
527 match read_request(stream, peer, limits) {
528 Ok(req) => {
529 let resp = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| handler(&req)))
530 .unwrap_or_else(|_| {
531 eprintln!("isb serve: handler panicked on {} {}", req.method, req.path);
532 Response::text(500, "internal error")
533 });
534 if let Some(up) = &resp.upgrade {
535 let f = up.lock().unwrap().take();
536 if write_upgrade(stream, &resp).is_ok() {
537 if let Some(f) = f {
538 f(stream);
539 }
540 }
541 return;
542 }
543 let _ = write_response(stream, &resp);
544 }
545 Err(Some(resp)) => {
546 let _ = write_response(stream, &resp);
547 let deadline = Instant::now() + Duration::from_secs(2);
551 let mut sink = [0u8; 16384];
552 let mut left = limits.max_body_bytes;
553 while left > 0 && Instant::now() < deadline {
554 match stream.read(&mut sink) {
555 Ok(0) | Err(_) => break,
556 Ok(n) => left = left.saturating_sub(n),
557 }
558 }
559 }
560 Err(None) => {}
561 }
562}
563
564#[expect(
567 clippy::too_many_lines,
568 reason = "predates the lint ratchet; split it when next changed"
569)]
570pub(crate) fn read_request<S: Read + Write>(
571 stream: &mut S,
572 peer: Peer,
573 limits: &Limits,
574) -> std::result::Result<Request, Option<Response>> {
575 let started = Instant::now();
576 let mut buf: Vec<u8> = Vec::with_capacity(4096);
577 let mut chunk = [0u8; 8192];
578 let head_len = loop {
579 if !buf.is_empty() {
580 let mut hs = [httparse::EMPTY_HEADER; 100];
581 let mut r = httparse::Request::new(&mut hs);
582 match r.parse(&buf) {
583 Ok(httparse::Status::Complete(n)) => break n,
584 Ok(httparse::Status::Partial) => {}
585 Err(httparse::Error::TooManyHeaders) => {
586 return Err(Some(Response::text(431, "too many headers")));
587 }
588 Err(e) => return Err(Some(Response::text(400, &format!("bad request: {e}")))),
589 }
590 }
591 if buf.len() > limits.max_header_bytes {
592 return Err(Some(Response::text(431, "request headers too large")));
593 }
594 if started.elapsed() > limits.read_timeout {
595 return Err(Some(Response::text(408, "request timeout")));
596 }
597 match stream.read(&mut chunk) {
598 Ok(0) => return Err(None),
599 Ok(n) => buf.extend_from_slice(&chunk[..n]),
600 Err(e) if e.kind() == ErrorKind::Interrupted => {}
601 Err(e) if matches!(e.kind(), ErrorKind::WouldBlock | ErrorKind::TimedOut) => {
602 return Err(if buf.is_empty() {
603 None
604 } else {
605 Some(Response::text(408, "request timeout"))
606 });
607 }
608 Err(_) => return Err(None),
609 }
610 };
611 if head_len > limits.max_header_bytes {
612 return Err(Some(Response::text(431, "request headers too large")));
613 }
614
615 let mut hs = [httparse::EMPTY_HEADER; 100];
616 let mut r = httparse::Request::new(&mut hs);
617 let _ = r.parse(&buf[..head_len]);
619 let method = r.method.unwrap_or_default().to_string();
620 let target = r.path.unwrap_or_default();
621 let headers: Vec<(String, String)> = r
622 .headers
623 .iter()
624 .map(|h| {
625 (
626 h.name.to_string(),
627 String::from_utf8_lossy(h.value).trim().to_string(),
628 )
629 })
630 .collect();
631 let (path, query) =
632 split_target(target).ok_or_else(|| Some(Response::text(400, "bad request target")))?;
633
634 fn find<'a>(
635 headers: &'a [(String, String)],
636 name: &'a str,
637 ) -> impl Iterator<Item = &'a (String, String)> {
638 headers
639 .iter()
640 .filter(move |(k, _)| k.eq_ignore_ascii_case(name))
641 }
642 if find(&headers, "transfer-encoding").next().is_some() {
643 return Err(Some(Response::text(
644 411,
645 "chunked request bodies are not supported; send Content-Length",
646 )));
647 }
648 let mut length: Option<usize> = None;
649 for (_, v) in find(&headers, "content-length") {
650 let n: usize = v
651 .parse()
652 .map_err(|_| Some(Response::text(400, "bad Content-Length")))?;
653 if length.is_some_and(|l| l != n) {
654 return Err(Some(Response::text(400, "conflicting Content-Length")));
655 }
656 length = Some(n);
657 }
658 let length = length.unwrap_or(0);
659 if length > limits.max_body_bytes {
660 return Err(Some(Response::text(413, "request body too large")));
661 }
662
663 let mut body = buf.split_off(head_len);
664 body.truncate(length);
665 if body.len() < length
666 && find(&headers, "expect").any(|(_, v)| v.eq_ignore_ascii_case("100-continue"))
667 {
668 stream
670 .write_all(b"HTTP/1.1 100 Continue\r\n\r\n")
671 .map_err(|_| None)?;
672 }
673 while body.len() < length {
674 if started.elapsed() > limits.read_timeout {
675 return Err(Some(Response::text(408, "request timeout")));
676 }
677 let want = (length - body.len()).min(chunk.len());
678 match stream.read(&mut chunk[..want]) {
679 Ok(0) => return Err(None),
680 Ok(n) => body.extend_from_slice(&chunk[..n]),
681 Err(e) if e.kind() == ErrorKind::Interrupted => {}
682 Err(e) if matches!(e.kind(), ErrorKind::WouldBlock | ErrorKind::TimedOut) => {
683 return Err(Some(Response::text(408, "request timeout")));
684 }
685 Err(_) => return Err(None),
686 }
687 }
688 Ok(Request {
689 method,
690 path,
691 query,
692 headers,
693 body,
694 peer,
695 })
696}
697
698fn split_target(target: &str) -> Option<(String, Option<String>)> {
701 let t = if target.starts_with('/') {
702 target
703 } else {
704 let rest = target
705 .strip_prefix("http://")
706 .or_else(|| target.strip_prefix("https://"))?;
707 match rest.find('/') {
708 Some(i) => &rest[i..],
709 None => "/",
710 }
711 };
712 Some(match t.split_once('?') {
713 Some((p, q)) => (p.to_string(), Some(q.to_string())),
714 None => (t.to_string(), None),
715 })
716}
717
718fn write_upgrade<W: Write>(w: &mut W, r: &Response) -> std::io::Result<()> {
720 let mut head = String::from("HTTP/1.1 101 Switching Protocols\r\nConnection: Upgrade\r\n");
721 for (k, v) in &r.headers {
722 if k.eq_ignore_ascii_case("content-length")
723 || k.eq_ignore_ascii_case("connection")
724 || k.eq_ignore_ascii_case("transfer-encoding")
725 || k.contains(['\r', '\n', ':'])
726 || v.contains(['\r', '\n'])
727 {
728 continue;
729 }
730 head.push_str(&format!("{k}: {v}\r\n"));
731 }
732 head.push_str("\r\n");
733 w.write_all(head.as_bytes())?;
734 w.flush()
735}
736
737pub(crate) fn write_response<W: Write>(w: &mut W, r: &Response) -> std::io::Result<()> {
738 let mut head = format!("HTTP/1.1 {} {}\r\n", r.status, reason(r.status));
739 for (k, v) in &r.headers {
740 if k.eq_ignore_ascii_case("content-length")
742 || k.eq_ignore_ascii_case("connection")
743 || k.eq_ignore_ascii_case("transfer-encoding")
744 || k.contains(['\r', '\n', ':'])
745 || v.contains(['\r', '\n'])
746 {
747 continue;
748 }
749 head.push_str(&format!("{k}: {v}\r\n"));
750 }
751 if let Some(stream) = &r.stream {
752 head.push_str("Cache-Control: no-store\r\nConnection: close\r\n\r\n");
754 w.write_all(head.as_bytes())?;
755 w.flush()?;
756 let f = stream.lock().unwrap().take();
757 return match f {
758 Some(f) => {
759 f(w)?;
760 w.flush()
761 }
762 None => Ok(()),
763 };
764 }
765 head.push_str(&format!(
766 "Content-Length: {}\r\nConnection: close\r\n\r\n",
767 r.body.len()
768 ));
769 w.write_all(head.as_bytes())?;
770 w.write_all(&r.body)?;
771 w.flush()
772}
773
774fn reason(status: u16) -> &'static str {
775 match status {
776 100 => "Continue",
777 101 => "Switching Protocols",
778 200 => "OK",
779 202 => "Accepted",
780 204 => "No Content",
781 400 => "Bad Request",
782 401 => "Unauthorized",
783 403 => "Forbidden",
784 404 => "Not Found",
785 405 => "Method Not Allowed",
786 406 => "Not Acceptable",
787 408 => "Request Timeout",
788 411 => "Length Required",
789 413 => "Content Too Large",
790 415 => "Unsupported Media Type",
791 431 => "Request Header Fields Too Large",
792 500 => "Internal Server Error",
793 501 => "Not Implemented",
794 503 => "Service Unavailable",
795 _ => "Status",
796 }
797}
798
799#[cfg(test)]
800pub(crate) mod tests {
801 use super::*;
802
803 pub(crate) struct Mock {
806 pub chunks: std::collections::VecDeque<Vec<u8>>,
807 pub output: Vec<u8>,
808 }
809
810 impl Mock {
811 pub fn new(input: impl Into<Vec<u8>>) -> Self {
812 Mock::chunked(vec![input.into()])
813 }
814
815 pub fn chunked(chunks: Vec<Vec<u8>>) -> Self {
816 Mock {
817 chunks: chunks.into(),
818 output: Vec::new(),
819 }
820 }
821 }
822
823 impl Read for Mock {
824 fn read(&mut self, b: &mut [u8]) -> std::io::Result<usize> {
825 let Some(front) = self.chunks.front_mut() else {
826 return Ok(0);
827 };
828 let n = front.len().min(b.len());
829 b[..n].copy_from_slice(&front[..n]);
830 front.drain(..n);
831 if front.is_empty() {
832 self.chunks.pop_front();
833 }
834 Ok(n)
835 }
836 }
837
838 impl Duplex for Mock {
839 fn set_read_timeout(&mut self, _: Option<Duration>) -> std::io::Result<()> {
840 Ok(())
841 }
842 }
843
844 impl Write for Mock {
845 fn write(&mut self, b: &[u8]) -> std::io::Result<usize> {
846 self.output.extend_from_slice(b);
847 Ok(b.len())
848 }
849 fn flush(&mut self) -> std::io::Result<()> {
850 Ok(())
851 }
852 }
853
854 fn parse(raw: &[u8], limits: &Limits) -> std::result::Result<Request, Option<Response>> {
855 read_request(&mut Mock::new(raw), Peer::Unix { uid: None }, limits)
856 }
857
858 fn status(raw: &[u8], limits: &Limits) -> u16 {
859 match parse(raw, limits) {
860 Ok(_) => 0,
861 Err(Some(r)) => r.status,
862 Err(None) => 1,
863 }
864 }
865
866 #[cfg(any(target_os = "linux", target_os = "macos"))]
867 #[test]
868 fn unix_peer_is_this_uid() {
869 let (a, _b) = UnixStream::pair().unwrap();
870 assert_eq!(peer_uid(&a), Some(rustix::process::getuid().as_raw()));
871 }
872
873 #[test]
874 fn parses_a_post() {
875 let r = parse(
876 b"POST /mcp?x=1 HTTP/1.1\r\nHost: a\r\nCONTENT-type: application/json\r\nContent-Length: 4\r\n\r\nabcdEXTRA",
877 &Limits::default(),
878 )
879 .unwrap();
880 assert_eq!(r.method, "POST");
881 assert_eq!(r.path, "/mcp");
882 assert_eq!(r.query.as_deref(), Some("x=1"));
883 assert_eq!(r.header("content-type"), Some("application/json"));
884 assert_eq!(r.body, b"abcd");
885 }
886
887 #[test]
888 fn absolute_form_target() {
889 let r = parse(
890 b"GET http://localhost:1/healthz HTTP/1.1\r\n\r\n",
891 &Limits::default(),
892 )
893 .unwrap();
894 assert_eq!(r.path, "/healthz");
895 }
896
897 #[test]
898 fn limits_and_malformed_input() {
899 let small = Limits {
900 max_header_bytes: 64,
901 max_body_bytes: 8,
902 ..Limits::default()
903 };
904 let big_header = format!("GET / HTTP/1.1\r\nX: {}\r\n\r\n", "a".repeat(200));
905 assert_eq!(status(big_header.as_bytes(), &small), 431);
906 let endless = format!("GET / HTTP/1.1\r\nX: {}", "a".repeat(10_000));
908 assert_eq!(status(endless.as_bytes(), &small), 431);
909 assert_eq!(
910 status(b"POST / HTTP/1.1\r\nContent-Length: 9\r\n\r\n", &small),
911 413
912 );
913 let l = Limits::default();
914 assert_eq!(
915 status(b"POST / HTTP/1.1\r\nTransfer-Encoding: chunked\r\n\r\n", &l),
916 411
917 );
918 assert_eq!(
919 status(b"POST / HTTP/1.1\r\nContent-Length: x\r\n\r\n", &l),
920 400
921 );
922 assert_eq!(
923 status(
924 b"POST / HTTP/1.1\r\nContent-Length: 1\r\nContent-Length: 2\r\n\r\nab",
925 &l
926 ),
927 400
928 );
929 assert_eq!(status(b"\x00\x01garbage\r\n\r\n", &l), 400);
930 assert_eq!(status(b"GET nope HTTP/1.1\r\n\r\n", &l), 400);
931 let many: String = (0..150).map(|i| format!("H{i}: v\r\n")).collect();
932 assert_eq!(
933 status(format!("GET / HTTP/1.1\r\n{many}\r\n").as_bytes(), &l),
934 431
935 );
936 assert_eq!(status(b"GET / HTTP/1.1\r\nHost:", &l), 1);
938 assert_eq!(
939 status(b"POST / HTTP/1.1\r\nContent-Length: 5\r\n\r\nab", &l),
940 1
941 );
942 assert_eq!(status(b"", &l), 1);
943 }
944
945 #[test]
946 fn expect_continue_is_answered() {
947 let mut m = Mock::chunked(vec![
948 b"POST / HTTP/1.1\r\nExpect: 100-continue\r\nContent-Length: 2\r\n\r\n".to_vec(),
949 b"hi".to_vec(),
950 ]);
951 let r = read_request(&mut m, Peer::Unix { uid: None }, &Limits::default()).unwrap();
952 assert_eq!(r.body, b"hi");
953 assert_eq!(m.output, b"HTTP/1.1 100 Continue\r\n\r\n");
954 }
955
956 #[test]
957 fn writes_framed_responses() {
958 let h: Handler = Arc::new(|r: &Request| {
959 Response::text(200, &r.path)
960 .header("Content-Length", "999")
961 .header("X-Bad", "a\r\nInjected: 1")
962 });
963 let mut m = Mock::new(&b"GET /hello HTTP/1.1\r\n\r\n"[..]);
964 handle(&mut m, Peer::Unix { uid: None }, &Limits::default(), &h);
965 let out = String::from_utf8(m.output).unwrap();
966 assert!(out.starts_with("HTTP/1.1 200 OK\r\n"), "{out}");
967 assert!(out.contains("Content-Length: 7\r\n"), "{out}");
968 assert!(out.contains("Connection: close\r\n"));
969 assert!(!out.contains("999") && !out.contains("Injected"));
970 assert!(out.ends_with("\r\n\r\n/hello\n"));
971 }
972
973 #[test]
974 fn an_upgrade_hands_over_the_connection() {
975 let h: Handler = Arc::new(|_: &Request| {
976 Response::upgrade(
977 "websocket",
978 Box::new(|s: &mut dyn Duplex| {
979 let _ = s.write_all(b"after");
980 }),
981 )
982 .header("Sec-WebSocket-Accept", "k")
983 });
984 let mut m = Mock::new(&b"GET /ws HTTP/1.1\r\nUpgrade: websocket\r\n\r\n"[..]);
985 handle(&mut m, Peer::Unix { uid: None }, &Limits::default(), &h);
986 let out = String::from_utf8(m.output).unwrap();
987 assert!(
988 out.starts_with("HTTP/1.1 101 Switching Protocols\r\n"),
989 "{out}"
990 );
991 assert!(out.contains("Connection: Upgrade\r\n") && out.contains("Upgrade: websocket\r\n"));
992 assert!(out.contains("Sec-WebSocket-Accept: k\r\n"));
993 assert!(!out.contains("Content-Length") && !out.contains("close"));
994 assert!(out.ends_with("\r\n\r\nafter"), "{out}");
995 }
996
997 #[test]
998 fn handler_panic_is_a_500() {
999 let h: Handler = Arc::new(|_: &Request| panic!("boom"));
1000 let mut m = Mock::new(&b"GET / HTTP/1.1\r\n\r\n"[..]);
1001 handle(&mut m, Peer::Unix { uid: None }, &Limits::default(), &h);
1002 assert!(m.output.starts_with(b"HTTP/1.1 500 "));
1003 }
1004
1005 #[test]
1006 fn tcp_bind_is_loopback_only() {
1007 assert!(HttpListener::bind_tcp("0.0.0.0:0").is_err());
1008 assert!(HttpListener::bind_tcp("192.0.2.1:0").is_err());
1009 let l = HttpListener::bind_tcp("127.0.0.1:0").unwrap();
1010 assert!(l.local_addr().unwrap().ip().is_loopback());
1011 }
1012
1013 #[test]
1014 fn unix_bind_permissions_and_stale_socket() {
1015 let dir = tempfile::tempdir().unwrap();
1016 let path = dir.path().join("sub/isb.sock");
1017 let l = HttpListener::bind_unix(&path).unwrap();
1018 let mode = |p: &Path| std::fs::metadata(p).unwrap().permissions().mode() & 0o777;
1019 assert_eq!(mode(&path), 0o600);
1020 assert_eq!(mode(path.parent().unwrap()), 0o700);
1021 assert!(HttpListener::bind_unix(&path).is_err());
1023 drop(l);
1024 drop(UnixListener::bind(&path).unwrap());
1026 assert!(path.exists());
1027 let l2 = HttpListener::bind_unix(&path).unwrap();
1028 drop(l2);
1029 assert!(!path.exists(), "socket removed on drop");
1030 std::fs::write(&path, b"x").unwrap();
1032 assert!(HttpListener::bind_unix(&path).is_err());
1033 assert!(path.exists());
1034 }
1035
1036 #[test]
1037 fn connection_cap_and_shutdown() {
1038 let dir = tempfile::tempdir().unwrap();
1039 let path = dir.path().join("cap.sock");
1040 let listener = HttpListener::bind_unix(&path).unwrap();
1041 let shutdown = Shutdown::new();
1042 let srv = HttpServer::new(
1043 Limits {
1044 max_connections: 1,
1045 ..Limits::default()
1046 },
1047 shutdown.clone(),
1048 );
1049 let gate = Arc::new(std::sync::Barrier::new(2));
1050 let g = gate.clone();
1051 let h: Handler = Arc::new(move |_: &Request| {
1052 g.wait();
1053 Response::text(200, "ok")
1054 });
1055 let s2 = srv.clone();
1056 let t = std::thread::spawn(move || s2.run(listener, h));
1057
1058 let send = |p: &Path| {
1059 let mut s = UnixStream::connect(p).unwrap();
1060 s.set_read_timeout(Some(Duration::from_secs(5))).unwrap();
1061 s.write_all(b"GET / HTTP/1.1\r\n\r\n").unwrap();
1062 s
1063 };
1064 let mut first = send(&path);
1065 let t0 = Instant::now();
1067 while srv.active() == 0 && t0.elapsed() < Duration::from_secs(5) {
1068 std::thread::sleep(Duration::from_millis(5));
1069 }
1070 let mut second = send(&path);
1071 let mut out = String::new();
1072 second.read_to_string(&mut out).unwrap();
1073 assert!(out.starts_with("HTTP/1.1 503 "), "{out}");
1074 gate.wait();
1075 out.clear();
1076 first.read_to_string(&mut out).unwrap();
1077 assert!(out.starts_with("HTTP/1.1 200 "), "{out}");
1078
1079 shutdown.trigger();
1080 t.join().unwrap().unwrap();
1081 srv.drain(Duration::from_secs(5));
1082 assert_eq!(srv.active(), 0);
1083 assert!(!path.exists());
1084 }
1085}