Skip to main content

isb_server/server/
http.rs

1//! A minimal synchronous HTTP/1.1 server: one request per connection, a thread
2//! per connection, a hard cap on connections, and every read and write bounded.
3//!
4//! It serves a JSON API to a tunnel and a local CLI, nothing else, so it speaks
5//! just enough HTTP for that: Content-Length bodies only, `Connection: close` on
6//! every response. Anything it does not understand is refused with a status
7//! rather than guessed at.
8
9use 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/// Bounds on what one connection may cost.
23#[derive(Debug, Clone)]
24pub struct Limits {
25    /// Request line plus headers.
26    pub max_header_bytes: usize,
27    pub max_body_bytes: usize,
28    /// Reading the whole request, start to finish.
29    pub read_timeout: Duration,
30    pub write_timeout: Duration,
31    /// Across every listener of one [`HttpServer`]; beyond it new connections
32    /// get a 503 instead of a thread.
33    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/// Who is on the other end of a connection.
49#[derive(Debug, Clone, PartialEq, Eq)]
50pub enum Peer {
51    Tcp(SocketAddr),
52    /// `uid` from SO_PEERCRED (getpeereid on macOS), when the platform has it.
53    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    /// In arrival order, names as sent; look them up with [`Request::header`].
64    pub headers: Vec<(String, String)>,
65    pub body: Vec<u8>,
66    pub peer: Peer,
67}
68
69impl Request {
70    /// The first header named `name`, case-insensitively.
71    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
79/// Writes a streamed body (server-sent events) until it returns; the
80/// connection closes after it, which is what delimits the body.
81pub type StreamFn = Box<dyn FnOnce(&mut dyn Write) -> std::io::Result<()> + Send>;
82
83/// A connection a handler takes over after `101 Switching Protocols` (a
84/// websocket): both directions, and a read deadline it can shorten to poll.
85pub 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
101/// Runs on the connection after a 101; the connection closes when it returns.
102pub 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    /// When set, written after the headers in place of `body`.
110    pub stream: Option<Arc<std::sync::Mutex<Option<StreamFn>>>>,
111    /// When set, the status is 101 and this takes the connection over.
112    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    /// `101 Switching Protocols` to `protocol`, then `f` owns the connection.
139    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    /// A streamed response: `f` writes the body and the connection closes
146    /// when it returns.
147    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/// A flag every accept loop polls. Cloning shares it.
188#[derive(Debug, Clone, Default)]
189pub struct Shutdown(Arc<AtomicBool>);
190
191impl Shutdown {
192    pub fn new() -> Self {
193        Self::default()
194    }
195
196    /// A flag that SIGINT and SIGTERM set.
197    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/// A bound listening socket.
215#[derive(Debug)]
216pub enum HttpListener {
217    Tcp(TcpListener),
218    Unix(UnixSocket),
219    /// TLS with client certificates (an `isb serve --agent` listener): any
220    /// address, since the handshake is the gate.
221    Tls(TcpListener, TlsConfig),
222}
223
224/// A TLS server config that can be swapped while serving (certificate
225/// rotation); each connection takes the one current when it arrives.
226pub type TlsConfig = Arc<std::sync::RwLock<Arc<rustls::ServerConfig>>>;
227
228/// A unix listener that removes its socket file on drop, but only while the
229/// file is still the one it bound, so a successor's socket is never deleted.
230#[derive(Debug)]
231pub struct UnixSocket {
232    listener: UnixListener,
233    path: PathBuf,
234    ino: (u64, u64),
235}
236
237impl Drop for UnixSocket {
238    fn drop(&mut self) {
239        if std::fs::symlink_metadata(&self.path).is_ok_and(|m| (m.dev(), m.ino()) == self.ino) {
240            let _ = std::fs::remove_file(&self.path);
241        }
242    }
243}
244
245impl HttpListener {
246    /// Bind a TCP address that must resolve to loopback only. Remote access
247    /// belongs behind the tunnel, which reaches loopback; a public bind would
248    /// skip Cloudflare Access entirely.
249    pub fn bind_tcp(addr: &str) -> Result<Self> {
250        let addrs: Vec<SocketAddr> = addr
251            .to_socket_addrs()
252            .map_err(|e| Error::invalid(format!("listen address {addr:?}: {e}")))?
253            .collect();
254        if addrs.is_empty() || addrs.iter().any(|a| !a.ip().is_loopback()) {
255            return Err(Error::invalid(format!(
256                "listen address {addr:?} is not loopback; expose it through a Cloudflare Tunnel instead"
257            )));
258        }
259        let l = TcpListener::bind(addrs[0])?;
260        l.set_nonblocking(true)?;
261        Ok(HttpListener::Tcp(l))
262    }
263
264    /// Bind a tailnet address (100.64.0.0/10, fd7a:115c:a1e0::/48): only
265    /// tailnet peers reach it, so it is served without a tunnel in front
266    /// (`isb serve --superadmin-tailnet`).
267    pub fn bind_tcp_tailnet(addr: &str) -> Result<Self> {
268        let addrs: Vec<SocketAddr> = addr
269            .to_socket_addrs()
270            .map_err(|e| Error::invalid(format!("listen address {addr:?}: {e}")))?
271            .collect();
272        if addrs.is_empty() || addrs.iter().any(|a| !super::tailnet::is_tailnet_ip(a.ip())) {
273            return Err(Error::invalid(format!(
274                "listen address {addr:?} is not a tailnet address (100.64.0.0/10, fd7a:115c:a1e0::/48)"
275            )));
276        }
277        let l = TcpListener::bind(addrs[0])?;
278        l.set_nonblocking(true)?;
279        Ok(HttpListener::Tcp(l))
280    }
281
282    /// Bind a private IPv4 address (10/8, 172.16/12, 192.168/16) that is
283    /// not loopback: an org bridge's gateway, which only that org's
284    /// instances route to. The handler serving it does the gating.
285    pub fn bind_tcp_private(addr: SocketAddr) -> Result<Self> {
286        let ok = match addr.ip() {
287            std::net::IpAddr::V4(v4) => v4.is_private(),
288            std::net::IpAddr::V6(_) => false,
289        };
290        if !ok {
291            return Err(Error::invalid(format!(
292                "listen address {addr} is not a private IPv4 address"
293            )));
294        }
295        let l = TcpListener::bind(addr)?;
296        l.set_nonblocking(true)?;
297        Ok(HttpListener::Tcp(l))
298    }
299
300    /// Bind a unix socket, mode 0600. A parent directory isb creates is 0700. A
301    /// stale socket is replaced; a live one (something answers) is an error,
302    /// and so is a path that is not a socket.
303    pub fn bind_unix(path: &Path) -> Result<Self> {
304        if let Some(dir) = path.parent().filter(|d| !d.as_os_str().is_empty()) {
305            std::fs::DirBuilder::new()
306                .recursive(true)
307                .mode(0o700)
308                .create(dir)?;
309        }
310        match std::fs::symlink_metadata(path) {
311            Ok(m) if m.file_type().is_socket() => {
312                if UnixStream::connect(path).is_ok() {
313                    return Err(Error::invalid(format!(
314                        "{} is in use by another server",
315                        path.display()
316                    )));
317                }
318                std::fs::remove_file(path)?;
319            }
320            Ok(_) => {
321                return Err(Error::invalid(format!(
322                    "{} exists and is not a socket",
323                    path.display()
324                )));
325            }
326            Err(e) if e.kind() == ErrorKind::NotFound => {}
327            Err(e) => return Err(e.into()),
328        }
329        let listener = UnixListener::bind(path)?;
330        std::fs::set_permissions(path, std::fs::Permissions::from_mode(0o600))?;
331        listener.set_nonblocking(true)?;
332        let m = std::fs::symlink_metadata(path)?;
333        Ok(HttpListener::Unix(UnixSocket {
334            listener,
335            path: path.to_path_buf(),
336            ino: (m.dev(), m.ino()),
337        }))
338    }
339
340    /// Bind `addr` (any address) for TLS. Only for a config that requires
341    /// client certificates: nothing else stands between it and the network.
342    pub fn bind_tls(addr: &str, config: TlsConfig) -> Result<Self> {
343        let a = addr
344            .to_socket_addrs()
345            .map_err(|e| Error::invalid(format!("listen address {addr:?}: {e}")))?
346            .next()
347            .ok_or_else(|| Error::invalid(format!("listen address {addr:?} does not resolve")))?;
348        let l = TcpListener::bind(a)?;
349        l.set_nonblocking(true)?;
350        Ok(HttpListener::Tls(l, config))
351    }
352
353    pub fn local_addr(&self) -> Option<SocketAddr> {
354        match self {
355            HttpListener::Tcp(l) | HttpListener::Tls(l, _) => l.local_addr().ok(),
356            HttpListener::Unix(_) => None,
357        }
358    }
359
360    fn poll(&self, timeout: Duration) -> bool {
361        use rustix::event::{PollFd, PollFlags, Timespec, poll};
362        let ts = Timespec {
363            tv_sec: timeout.as_secs() as _,
364            tv_nsec: timeout.subsec_nanos() as _,
365        };
366        let r = match self {
367            HttpListener::Tcp(l) | HttpListener::Tls(l, _) => {
368                poll(&mut [PollFd::new(l, PollFlags::IN)], Some(&ts))
369            }
370            HttpListener::Unix(u) => {
371                poll(&mut [PollFd::new(&u.listener, PollFlags::IN)], Some(&ts))
372            }
373        };
374        matches!(r, Ok(n) if n > 0)
375    }
376}
377
378/// Owns the connection cap and the shutdown flag shared by every listener.
379#[derive(Debug)]
380pub struct HttpServer {
381    limits: Limits,
382    active: AtomicUsize,
383    shutdown: Shutdown,
384}
385
386/// Holds one slot of the connection cap; released on drop, including when the
387/// thread that would have owned it fails to spawn.
388struct Slot(Arc<HttpServer>);
389
390impl Drop for Slot {
391    fn drop(&mut self) {
392        self.0.active.fetch_sub(1, Ordering::SeqCst);
393    }
394}
395
396impl HttpServer {
397    pub fn new(limits: Limits, shutdown: Shutdown) -> Arc<Self> {
398        Arc::new(HttpServer {
399            limits,
400            active: AtomicUsize::new(0),
401            shutdown,
402        })
403    }
404
405    pub fn limits(&self) -> &Limits {
406        &self.limits
407    }
408
409    /// Connections currently being served.
410    pub fn active(&self) -> usize {
411        self.active.load(Ordering::SeqCst)
412    }
413
414    /// Accept connections until shutdown. The listener (and a unix socket's
415    /// file) is closed on return.
416    pub fn run(self: &Arc<Self>, listener: HttpListener, handler: Handler) -> Result<()> {
417        while !self.shutdown.is_triggered() {
418            // Poll rather than block in accept, so shutdown is noticed promptly.
419            if !listener.poll(Duration::from_millis(250)) {
420                continue;
421            }
422            let accepted = match &listener {
423                HttpListener::Tcp(l) => l.accept().map(|(s, a)| Conn::Tcp(s, a)),
424                HttpListener::Tls(l, c) => l.accept().map(|(s, a)| {
425                    let cfg = c.read().unwrap_or_else(|p| p.into_inner()).clone();
426                    Conn::Tls(s, a, cfg)
427                }),
428                HttpListener::Unix(u) => u.listener.accept().map(|(s, _)| Conn::Unix(s)),
429            };
430            let conn = match accepted {
431                Ok(c) => c,
432                Err(e) if matches!(e.kind(), ErrorKind::WouldBlock | ErrorKind::Interrupted) => {
433                    continue;
434                }
435                Err(e) => {
436                    // EMFILE and friends: back off instead of spinning.
437                    eprintln!("isb serve: accept: {e}");
438                    std::thread::sleep(Duration::from_millis(100));
439                    continue;
440                }
441            };
442            if self.active.fetch_add(1, Ordering::SeqCst) >= self.limits.max_connections {
443                self.active.fetch_sub(1, Ordering::SeqCst);
444                conn.reject_busy();
445                continue;
446            }
447            let slot = Slot(self.clone());
448            let h = handler.clone();
449            let spawned = std::thread::Builder::new()
450                .name("isb-http".into())
451                .spawn(move || {
452                    let s = slot;
453                    conn.serve(&s.0.limits, &h);
454                });
455            if let Err(e) = spawned {
456                eprintln!("isb serve: cannot spawn connection thread: {e}");
457            }
458        }
459        Ok(())
460    }
461
462    /// Wait, at most `timeout`, for in-flight connections to finish.
463    pub fn drain(&self, timeout: Duration) {
464        let started = Instant::now();
465        while self.active() > 0 && started.elapsed() < timeout {
466            std::thread::sleep(Duration::from_millis(20));
467        }
468    }
469}
470
471enum Conn {
472    Tcp(TcpStream, SocketAddr),
473    Unix(UnixStream),
474    Tls(TcpStream, SocketAddr, Arc<rustls::ServerConfig>),
475}
476
477/// A server-side TLS connection, as a [`Duplex`].
478struct TlsConn(rustls::StreamOwned<rustls::ServerConnection, TcpStream>);
479
480impl Read for TlsConn {
481    fn read(&mut self, b: &mut [u8]) -> std::io::Result<usize> {
482        self.0.read(b)
483    }
484}
485
486impl Write for TlsConn {
487    fn write(&mut self, b: &[u8]) -> std::io::Result<usize> {
488        self.0.write(b)
489    }
490    fn flush(&mut self) -> std::io::Result<()> {
491        self.0.flush()
492    }
493}
494
495impl Duplex for TlsConn {
496    fn set_read_timeout(&mut self, t: Option<Duration>) -> std::io::Result<()> {
497        self.0.sock.set_read_timeout(t)
498    }
499}
500
501impl Conn {
502    /// Runs on the accept thread, so nothing here may block for long: a short
503    /// write, then discard whatever request bytes already arrived (closing on
504    /// unread input resets the connection and loses the 503).
505    fn reject_busy(self) {
506        fn refuse<S: Read + Write>(s: &mut S) {
507            let r = Response::text(503, "server busy").header("Retry-After", "1");
508            let _ = write_response(s, &r);
509            let mut sink = [0u8; 16384];
510            for _ in 0..4 {
511                if !matches!(s.read(&mut sink), Ok(n) if n > 0) {
512                    break;
513                }
514            }
515        }
516        // Non-blocking: a 503 fits in any socket buffer, and a client that
517        // has not finished sending is not worth waiting for.
518        match self {
519            Conn::Tcp(mut s, _) => {
520                let _ = s.set_nonblocking(true);
521                refuse(&mut s);
522                let _ = s.shutdown(NetShutdown::Write);
523            }
524            // No TLS session yet to say it in: just close.
525            Conn::Tls(s, _, _) => {
526                let _ = s.shutdown(NetShutdown::Both);
527            }
528            Conn::Unix(mut s) => {
529                let _ = s.set_nonblocking(true);
530                refuse(&mut s);
531                let _ = s.shutdown(NetShutdown::Write);
532            }
533        }
534    }
535
536    fn serve(self, limits: &Limits, handler: &Handler) {
537        match self {
538            Conn::Tcp(s, addr) => {
539                let _ = s.set_nonblocking(false);
540                let _ = s.set_nodelay(true);
541                let _ = s.set_read_timeout(Some(limits.read_timeout));
542                let _ = s.set_write_timeout(Some(limits.write_timeout));
543                let mut s = s;
544                handle(&mut s, Peer::Tcp(addr), limits, handler);
545                let _ = s.shutdown(NetShutdown::Write);
546            }
547            Conn::Tls(s, addr, cfg) => {
548                let _ = s.set_nonblocking(false);
549                let _ = s.set_nodelay(true);
550                let _ = s.set_read_timeout(Some(limits.read_timeout));
551                let _ = s.set_write_timeout(Some(limits.write_timeout));
552                let Ok(conn) = rustls::ServerConnection::new(cfg) else {
553                    return;
554                };
555                let mut t = TlsConn(rustls::StreamOwned::new(conn, s));
556                // The handshake (and the client certificate check) first,
557                // bounded by the socket timeouts.
558                while t.0.conn.is_handshaking() {
559                    if let Err(e) = t.0.conn.complete_io(&mut t.0.sock) {
560                        eprintln!("isb serve: TLS handshake from {addr}: {e}");
561                        let _ = t.0.sock.shutdown(NetShutdown::Both);
562                        return;
563                    }
564                }
565                handle(&mut t, Peer::Tcp(addr), limits, handler);
566                t.0.conn.send_close_notify();
567                let _ = t.0.flush();
568                let _ = t.0.sock.shutdown(NetShutdown::Write);
569            }
570            Conn::Unix(s) => {
571                let _ = s.set_nonblocking(false);
572                let _ = s.set_read_timeout(Some(limits.read_timeout));
573                let _ = s.set_write_timeout(Some(limits.write_timeout));
574                let uid = peer_uid(&s);
575                let mut s = s;
576                handle(&mut s, Peer::Unix { uid }, limits, handler);
577                let _ = s.shutdown(NetShutdown::Write);
578            }
579        }
580    }
581}
582
583#[cfg(target_os = "linux")]
584fn peer_uid(s: &UnixStream) -> Option<u32> {
585    rustix::net::sockopt::socket_peercred(s)
586        .ok()
587        .map(|c| c.uid.as_raw())
588}
589
590#[cfg(target_os = "macos")]
591fn peer_uid(s: &UnixStream) -> Option<u32> {
592    use std::os::fd::AsRawFd;
593    let (mut uid, mut gid) = (0, 0);
594    // SAFETY: a valid socket fd and two out-parameters.
595    (unsafe { libc::getpeereid(s.as_raw_fd(), &mut uid, &mut gid) } == 0).then_some(uid)
596}
597
598#[cfg(not(any(target_os = "linux", target_os = "macos")))]
599fn peer_uid(_: &UnixStream) -> Option<u32> {
600    None
601}
602
603/// Serve one request on `stream`. Generic so tests can drive it in memory.
604pub(crate) fn handle<S: Duplex>(stream: &mut S, peer: Peer, limits: &Limits, handler: &Handler) {
605    match read_request(stream, peer, limits) {
606        Ok(req) => {
607            let resp = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| handler(&req)))
608                .unwrap_or_else(|_| {
609                    eprintln!("isb serve: handler panicked on {} {}", req.method, req.path);
610                    Response::text(500, "internal error")
611                });
612            if let Some(up) = &resp.upgrade {
613                let f = up.lock().unwrap().take();
614                if write_upgrade(stream, &resp).is_ok() {
615                    if let Some(f) = f {
616                        f(stream);
617                    }
618                }
619                return;
620            }
621            let _ = write_response(stream, &resp);
622        }
623        Err(Some(resp)) => {
624            let _ = write_response(stream, &resp);
625            // The client may still be sending a body we refused. Closing with
626            // unread input makes the kernel send RST, which can destroy the
627            // response before the client reads it; swallow a bounded amount.
628            let deadline = Instant::now() + Duration::from_secs(2);
629            let mut sink = [0u8; 16384];
630            let mut left = limits.max_body_bytes;
631            while left > 0 && Instant::now() < deadline {
632                match stream.read(&mut sink) {
633                    Ok(0) | Err(_) => break,
634                    Ok(n) => left = left.saturating_sub(n),
635                }
636            }
637        }
638        Err(None) => {}
639    }
640}
641
642/// Read and parse one request. `Err(Some(resp))` is a refusal to send back;
643/// `Err(None)` means the peer went away and there is nobody to answer.
644#[expect(
645    clippy::too_many_lines,
646    reason = "predates the lint ratchet; split it when next changed"
647)]
648pub(crate) fn read_request<S: Read + Write>(
649    stream: &mut S,
650    peer: Peer,
651    limits: &Limits,
652) -> std::result::Result<Request, Option<Response>> {
653    let started = Instant::now();
654    let mut buf: Vec<u8> = Vec::with_capacity(4096);
655    let mut chunk = [0u8; 8192];
656    let head_len = loop {
657        if !buf.is_empty() {
658            let mut hs = [httparse::EMPTY_HEADER; 100];
659            let mut r = httparse::Request::new(&mut hs);
660            match r.parse(&buf) {
661                Ok(httparse::Status::Complete(n)) => break n,
662                Ok(httparse::Status::Partial) => {}
663                Err(httparse::Error::TooManyHeaders) => {
664                    return Err(Some(Response::text(431, "too many headers")));
665                }
666                Err(e) => return Err(Some(Response::text(400, &format!("bad request: {e}")))),
667            }
668        }
669        if buf.len() > limits.max_header_bytes {
670            return Err(Some(Response::text(431, "request headers too large")));
671        }
672        if started.elapsed() > limits.read_timeout {
673            return Err(Some(Response::text(408, "request timeout")));
674        }
675        match stream.read(&mut chunk) {
676            Ok(0) => return Err(None),
677            Ok(n) => buf.extend_from_slice(&chunk[..n]),
678            Err(e) if e.kind() == ErrorKind::Interrupted => {}
679            Err(e) if matches!(e.kind(), ErrorKind::WouldBlock | ErrorKind::TimedOut) => {
680                return Err(if buf.is_empty() {
681                    None
682                } else {
683                    Some(Response::text(408, "request timeout"))
684                });
685            }
686            Err(_) => return Err(None),
687        }
688    };
689    if head_len > limits.max_header_bytes {
690        return Err(Some(Response::text(431, "request headers too large")));
691    }
692
693    let mut hs = [httparse::EMPTY_HEADER; 100];
694    let mut r = httparse::Request::new(&mut hs);
695    // Parsed successfully above on the same bytes.
696    let _ = r.parse(&buf[..head_len]);
697    let method = r.method.unwrap_or_default().to_string();
698    let target = r.path.unwrap_or_default();
699    let headers: Vec<(String, String)> = r
700        .headers
701        .iter()
702        .map(|h| {
703            (
704                h.name.to_string(),
705                String::from_utf8_lossy(h.value).trim().to_string(),
706            )
707        })
708        .collect();
709    let (path, query) =
710        split_target(target).ok_or_else(|| Some(Response::text(400, "bad request target")))?;
711
712    fn find<'a>(
713        headers: &'a [(String, String)],
714        name: &'a str,
715    ) -> impl Iterator<Item = &'a (String, String)> {
716        headers
717            .iter()
718            .filter(move |(k, _)| k.eq_ignore_ascii_case(name))
719    }
720    if find(&headers, "transfer-encoding").next().is_some() {
721        return Err(Some(Response::text(
722            411,
723            "chunked request bodies are not supported; send Content-Length",
724        )));
725    }
726    let mut length: Option<usize> = None;
727    for (_, v) in find(&headers, "content-length") {
728        let n: usize = v
729            .parse()
730            .map_err(|_| Some(Response::text(400, "bad Content-Length")))?;
731        if length.is_some_and(|l| l != n) {
732            return Err(Some(Response::text(400, "conflicting Content-Length")));
733        }
734        length = Some(n);
735    }
736    let length = length.unwrap_or(0);
737    if length > limits.max_body_bytes {
738        return Err(Some(Response::text(413, "request body too large")));
739    }
740
741    let mut body = buf.split_off(head_len);
742    body.truncate(length);
743    if body.len() < length
744        && find(&headers, "expect").any(|(_, v)| v.eq_ignore_ascii_case("100-continue"))
745    {
746        // curl waits for this (up to a second) before sending a large body.
747        stream
748            .write_all(b"HTTP/1.1 100 Continue\r\n\r\n")
749            .map_err(|_| None)?;
750    }
751    while body.len() < length {
752        if started.elapsed() > limits.read_timeout {
753            return Err(Some(Response::text(408, "request timeout")));
754        }
755        let want = (length - body.len()).min(chunk.len());
756        match stream.read(&mut chunk[..want]) {
757            Ok(0) => return Err(None),
758            Ok(n) => body.extend_from_slice(&chunk[..n]),
759            Err(e) if e.kind() == ErrorKind::Interrupted => {}
760            Err(e) if matches!(e.kind(), ErrorKind::WouldBlock | ErrorKind::TimedOut) => {
761                return Err(Some(Response::text(408, "request timeout")));
762            }
763            Err(_) => return Err(None),
764        }
765    }
766    Ok(Request {
767        method,
768        path,
769        query,
770        headers,
771        body,
772        peer,
773    })
774}
775
776/// Origin-form `/path?query`, or absolute-form `http://host/path?query` with
777/// the authority dropped.
778fn split_target(target: &str) -> Option<(String, Option<String>)> {
779    let t = if target.starts_with('/') {
780        target
781    } else {
782        let rest = target
783            .strip_prefix("http://")
784            .or_else(|| target.strip_prefix("https://"))?;
785        match rest.find('/') {
786            Some(i) => &rest[i..],
787            None => "/",
788        }
789    };
790    Some(match t.split_once('?') {
791        Some((p, q)) => (p.to_string(), Some(q.to_string())),
792        None => (t.to_string(), None),
793    })
794}
795
796/// The head of a `101 Switching Protocols`; no body follows.
797fn write_upgrade<W: Write>(w: &mut W, r: &Response) -> std::io::Result<()> {
798    let mut head = String::from("HTTP/1.1 101 Switching Protocols\r\nConnection: Upgrade\r\n");
799    for (k, v) in &r.headers {
800        if k.eq_ignore_ascii_case("content-length")
801            || k.eq_ignore_ascii_case("connection")
802            || k.eq_ignore_ascii_case("transfer-encoding")
803            || k.contains(['\r', '\n', ':'])
804            || v.contains(['\r', '\n'])
805        {
806            continue;
807        }
808        head.push_str(&format!("{k}: {v}\r\n"));
809    }
810    head.push_str("\r\n");
811    w.write_all(head.as_bytes())?;
812    w.flush()
813}
814
815pub(crate) fn write_response<W: Write>(w: &mut W, r: &Response) -> std::io::Result<()> {
816    let mut head = format!("HTTP/1.1 {} {}\r\n", r.status, reason(r.status));
817    for (k, v) in &r.headers {
818        // Framing is ours; and a value with a line break would split the header.
819        if k.eq_ignore_ascii_case("content-length")
820            || k.eq_ignore_ascii_case("connection")
821            || k.eq_ignore_ascii_case("transfer-encoding")
822            || k.contains(['\r', '\n', ':'])
823            || v.contains(['\r', '\n'])
824        {
825            continue;
826        }
827        head.push_str(&format!("{k}: {v}\r\n"));
828    }
829    if let Some(stream) = &r.stream {
830        // No length: the body runs until the connection closes.
831        head.push_str("Cache-Control: no-store\r\nConnection: close\r\n\r\n");
832        w.write_all(head.as_bytes())?;
833        w.flush()?;
834        let f = stream.lock().unwrap().take();
835        return match f {
836            Some(f) => {
837                f(w)?;
838                w.flush()
839            }
840            None => Ok(()),
841        };
842    }
843    head.push_str(&format!(
844        "Content-Length: {}\r\nConnection: close\r\n\r\n",
845        r.body.len()
846    ));
847    w.write_all(head.as_bytes())?;
848    w.write_all(&r.body)?;
849    w.flush()
850}
851
852fn reason(status: u16) -> &'static str {
853    match status {
854        100 => "Continue",
855        101 => "Switching Protocols",
856        200 => "OK",
857        202 => "Accepted",
858        204 => "No Content",
859        400 => "Bad Request",
860        401 => "Unauthorized",
861        403 => "Forbidden",
862        404 => "Not Found",
863        405 => "Method Not Allowed",
864        406 => "Not Acceptable",
865        408 => "Request Timeout",
866        411 => "Length Required",
867        413 => "Content Too Large",
868        415 => "Unsupported Media Type",
869        431 => "Request Header Fields Too Large",
870        500 => "Internal Server Error",
871        501 => "Not Implemented",
872        503 => "Service Unavailable",
873        _ => "Status",
874    }
875}
876
877#[cfg(test)]
878pub(crate) mod tests {
879    use super::*;
880
881    /// An in-memory connection: each read returns at most one of `chunks`,
882    /// and everything written is recorded.
883    pub(crate) struct Mock {
884        pub chunks: std::collections::VecDeque<Vec<u8>>,
885        pub output: Vec<u8>,
886    }
887
888    impl Mock {
889        pub fn new(input: impl Into<Vec<u8>>) -> Self {
890            Mock::chunked(vec![input.into()])
891        }
892
893        pub fn chunked(chunks: Vec<Vec<u8>>) -> Self {
894            Mock {
895                chunks: chunks.into(),
896                output: Vec::new(),
897            }
898        }
899    }
900
901    impl Read for Mock {
902        fn read(&mut self, b: &mut [u8]) -> std::io::Result<usize> {
903            let Some(front) = self.chunks.front_mut() else {
904                return Ok(0);
905            };
906            let n = front.len().min(b.len());
907            b[..n].copy_from_slice(&front[..n]);
908            front.drain(..n);
909            if front.is_empty() {
910                self.chunks.pop_front();
911            }
912            Ok(n)
913        }
914    }
915
916    impl Duplex for Mock {
917        fn set_read_timeout(&mut self, _: Option<Duration>) -> std::io::Result<()> {
918            Ok(())
919        }
920    }
921
922    impl Write for Mock {
923        fn write(&mut self, b: &[u8]) -> std::io::Result<usize> {
924            self.output.extend_from_slice(b);
925            Ok(b.len())
926        }
927        fn flush(&mut self) -> std::io::Result<()> {
928            Ok(())
929        }
930    }
931
932    fn parse(raw: &[u8], limits: &Limits) -> std::result::Result<Request, Option<Response>> {
933        read_request(&mut Mock::new(raw), Peer::Unix { uid: None }, limits)
934    }
935
936    fn status(raw: &[u8], limits: &Limits) -> u16 {
937        match parse(raw, limits) {
938            Ok(_) => 0,
939            Err(Some(r)) => r.status,
940            Err(None) => 1,
941        }
942    }
943
944    #[cfg(any(target_os = "linux", target_os = "macos"))]
945    #[test]
946    fn unix_peer_is_this_uid() {
947        let (a, _b) = UnixStream::pair().unwrap();
948        assert_eq!(peer_uid(&a), Some(rustix::process::getuid().as_raw()));
949    }
950
951    #[test]
952    fn parses_a_post() {
953        let r = parse(
954            b"POST /mcp?x=1 HTTP/1.1\r\nHost: a\r\nCONTENT-type: application/json\r\nContent-Length: 4\r\n\r\nabcdEXTRA",
955            &Limits::default(),
956        )
957        .unwrap();
958        assert_eq!(r.method, "POST");
959        assert_eq!(r.path, "/mcp");
960        assert_eq!(r.query.as_deref(), Some("x=1"));
961        assert_eq!(r.header("content-type"), Some("application/json"));
962        assert_eq!(r.body, b"abcd");
963    }
964
965    #[test]
966    fn absolute_form_target() {
967        let r = parse(
968            b"GET http://localhost:1/healthz HTTP/1.1\r\n\r\n",
969            &Limits::default(),
970        )
971        .unwrap();
972        assert_eq!(r.path, "/healthz");
973    }
974
975    #[test]
976    fn limits_and_malformed_input() {
977        let small = Limits {
978            max_header_bytes: 64,
979            max_body_bytes: 8,
980            ..Limits::default()
981        };
982        let big_header = format!("GET / HTTP/1.1\r\nX: {}\r\n\r\n", "a".repeat(200));
983        assert_eq!(status(big_header.as_bytes(), &small), 431);
984        // Never completes and never stops: still bounded by the header limit.
985        let endless = format!("GET / HTTP/1.1\r\nX: {}", "a".repeat(10_000));
986        assert_eq!(status(endless.as_bytes(), &small), 431);
987        assert_eq!(
988            status(b"POST / HTTP/1.1\r\nContent-Length: 9\r\n\r\n", &small),
989            413
990        );
991        let l = Limits::default();
992        assert_eq!(
993            status(b"POST / HTTP/1.1\r\nTransfer-Encoding: chunked\r\n\r\n", &l),
994            411
995        );
996        assert_eq!(
997            status(b"POST / HTTP/1.1\r\nContent-Length: x\r\n\r\n", &l),
998            400
999        );
1000        assert_eq!(
1001            status(
1002                b"POST / HTTP/1.1\r\nContent-Length: 1\r\nContent-Length: 2\r\n\r\nab",
1003                &l
1004            ),
1005            400
1006        );
1007        assert_eq!(status(b"\x00\x01garbage\r\n\r\n", &l), 400);
1008        assert_eq!(status(b"GET nope HTTP/1.1\r\n\r\n", &l), 400);
1009        let many: String = (0..150).map(|i| format!("H{i}: v\r\n")).collect();
1010        assert_eq!(
1011            status(format!("GET / HTTP/1.1\r\n{many}\r\n").as_bytes(), &l),
1012            431
1013        );
1014        // Peer hung up mid-request or mid-body: nobody to answer.
1015        assert_eq!(status(b"GET / HTTP/1.1\r\nHost:", &l), 1);
1016        assert_eq!(
1017            status(b"POST / HTTP/1.1\r\nContent-Length: 5\r\n\r\nab", &l),
1018            1
1019        );
1020        assert_eq!(status(b"", &l), 1);
1021    }
1022
1023    #[test]
1024    fn expect_continue_is_answered() {
1025        let mut m = Mock::chunked(vec![
1026            b"POST / HTTP/1.1\r\nExpect: 100-continue\r\nContent-Length: 2\r\n\r\n".to_vec(),
1027            b"hi".to_vec(),
1028        ]);
1029        let r = read_request(&mut m, Peer::Unix { uid: None }, &Limits::default()).unwrap();
1030        assert_eq!(r.body, b"hi");
1031        assert_eq!(m.output, b"HTTP/1.1 100 Continue\r\n\r\n");
1032    }
1033
1034    #[test]
1035    fn writes_framed_responses() {
1036        let h: Handler = Arc::new(|r: &Request| {
1037            Response::text(200, &r.path)
1038                .header("Content-Length", "999")
1039                .header("X-Bad", "a\r\nInjected: 1")
1040        });
1041        let mut m = Mock::new(&b"GET /hello HTTP/1.1\r\n\r\n"[..]);
1042        handle(&mut m, Peer::Unix { uid: None }, &Limits::default(), &h);
1043        let out = String::from_utf8(m.output).unwrap();
1044        assert!(out.starts_with("HTTP/1.1 200 OK\r\n"), "{out}");
1045        assert!(out.contains("Content-Length: 7\r\n"), "{out}");
1046        assert!(out.contains("Connection: close\r\n"));
1047        assert!(!out.contains("999") && !out.contains("Injected"));
1048        assert!(out.ends_with("\r\n\r\n/hello\n"));
1049    }
1050
1051    #[test]
1052    fn an_upgrade_hands_over_the_connection() {
1053        let h: Handler = Arc::new(|_: &Request| {
1054            Response::upgrade(
1055                "websocket",
1056                Box::new(|s: &mut dyn Duplex| {
1057                    let _ = s.write_all(b"after");
1058                }),
1059            )
1060            .header("Sec-WebSocket-Accept", "k")
1061        });
1062        let mut m = Mock::new(&b"GET /ws HTTP/1.1\r\nUpgrade: websocket\r\n\r\n"[..]);
1063        handle(&mut m, Peer::Unix { uid: None }, &Limits::default(), &h);
1064        let out = String::from_utf8(m.output).unwrap();
1065        assert!(
1066            out.starts_with("HTTP/1.1 101 Switching Protocols\r\n"),
1067            "{out}"
1068        );
1069        assert!(out.contains("Connection: Upgrade\r\n") && out.contains("Upgrade: websocket\r\n"));
1070        assert!(out.contains("Sec-WebSocket-Accept: k\r\n"));
1071        assert!(!out.contains("Content-Length") && !out.contains("close"));
1072        assert!(out.ends_with("\r\n\r\nafter"), "{out}");
1073    }
1074
1075    #[test]
1076    fn handler_panic_is_a_500() {
1077        let h: Handler = Arc::new(|_: &Request| panic!("boom"));
1078        let mut m = Mock::new(&b"GET / HTTP/1.1\r\n\r\n"[..]);
1079        handle(&mut m, Peer::Unix { uid: None }, &Limits::default(), &h);
1080        assert!(m.output.starts_with(b"HTTP/1.1 500 "));
1081    }
1082
1083    #[test]
1084    fn tcp_bind_is_loopback_only() {
1085        assert!(HttpListener::bind_tcp("0.0.0.0:0").is_err());
1086        assert!(HttpListener::bind_tcp("192.0.2.1:0").is_err());
1087        let l = HttpListener::bind_tcp("127.0.0.1:0").unwrap();
1088        assert!(l.local_addr().unwrap().ip().is_loopback());
1089    }
1090
1091    #[test]
1092    fn unix_bind_permissions_and_stale_socket() {
1093        let dir = tempfile::tempdir().unwrap();
1094        let path = dir.path().join("sub/isb.sock");
1095        let l = HttpListener::bind_unix(&path).unwrap();
1096        let mode = |p: &Path| std::fs::metadata(p).unwrap().permissions().mode() & 0o777;
1097        assert_eq!(mode(&path), 0o600);
1098        assert_eq!(mode(path.parent().unwrap()), 0o700);
1099        // A live socket is not stolen.
1100        assert!(HttpListener::bind_unix(&path).is_err());
1101        drop(l);
1102        // A stale one (left behind by a crash) is replaced.
1103        drop(UnixListener::bind(&path).unwrap());
1104        assert!(path.exists());
1105        let l2 = HttpListener::bind_unix(&path).unwrap();
1106        drop(l2);
1107        assert!(!path.exists(), "socket removed on drop");
1108        // Something that is not a socket is never deleted.
1109        std::fs::write(&path, b"x").unwrap();
1110        assert!(HttpListener::bind_unix(&path).is_err());
1111        assert!(path.exists());
1112    }
1113
1114    #[test]
1115    fn connection_cap_and_shutdown() {
1116        let dir = tempfile::tempdir().unwrap();
1117        let path = dir.path().join("cap.sock");
1118        let listener = HttpListener::bind_unix(&path).unwrap();
1119        let shutdown = Shutdown::new();
1120        let srv = HttpServer::new(
1121            Limits {
1122                max_connections: 1,
1123                ..Limits::default()
1124            },
1125            shutdown.clone(),
1126        );
1127        let gate = Arc::new(std::sync::Barrier::new(2));
1128        let g = gate.clone();
1129        let h: Handler = Arc::new(move |_: &Request| {
1130            g.wait();
1131            Response::text(200, "ok")
1132        });
1133        let s2 = srv.clone();
1134        let t = std::thread::spawn(move || s2.run(listener, h));
1135
1136        let send = |p: &Path| {
1137            let mut s = UnixStream::connect(p).unwrap();
1138            s.set_read_timeout(Some(Duration::from_secs(5))).unwrap();
1139            s.write_all(b"GET / HTTP/1.1\r\n\r\n").unwrap();
1140            s
1141        };
1142        let mut first = send(&path);
1143        // Wait until the first connection holds the only slot.
1144        let t0 = Instant::now();
1145        while srv.active() == 0 && t0.elapsed() < Duration::from_secs(5) {
1146            std::thread::sleep(Duration::from_millis(5));
1147        }
1148        let mut second = send(&path);
1149        let mut out = String::new();
1150        second.read_to_string(&mut out).unwrap();
1151        assert!(out.starts_with("HTTP/1.1 503 "), "{out}");
1152        gate.wait();
1153        out.clear();
1154        first.read_to_string(&mut out).unwrap();
1155        assert!(out.starts_with("HTTP/1.1 200 "), "{out}");
1156
1157        shutdown.trigger();
1158        t.join().unwrap().unwrap();
1159        srv.drain(Duration::from_secs(5));
1160        assert_eq!(srv.active(), 0);
1161        assert!(!path.exists());
1162    }
1163}