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}
220
221/// A unix listener that removes its socket file on drop, but only while the
222/// file is still the one it bound, so a successor's socket is never deleted.
223#[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    /// Bind a TCP address that must resolve to loopback only. Remote access
240    /// belongs behind the tunnel, which reaches loopback; a public bind would
241    /// skip Cloudflare Access entirely.
242    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    /// Bind a tailnet address (100.64.0.0/10, fd7a:115c:a1e0::/48): only
258    /// tailnet peers reach it, so it is served without a tunnel in front
259    /// (`isb serve --superadmin-tailnet`).
260    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    /// Bind a private IPv4 address (10/8, 172.16/12, 192.168/16) that is
276    /// not loopback: an org bridge's gateway, which only that org's
277    /// instances route to. The handler serving it does the gating.
278    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    /// Bind a unix socket, mode 0600. A parent directory isb creates is 0700. A
294    /// stale socket is replaced; a live one (something answers) is an error,
295    /// and so is a path that is not a socket.
296    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/// Owns the connection cap and the shutdown flag shared by every listener.
357#[derive(Debug)]
358pub struct HttpServer {
359    limits: Limits,
360    active: AtomicUsize,
361    shutdown: Shutdown,
362}
363
364/// Holds one slot of the connection cap; released on drop, including when the
365/// thread that would have owned it fails to spawn.
366struct 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    /// Connections currently being served.
388    pub fn active(&self) -> usize {
389        self.active.load(Ordering::SeqCst)
390    }
391
392    /// Accept connections until shutdown. The listener (and a unix socket's
393    /// file) is closed on return.
394    pub fn run(self: &Arc<Self>, listener: HttpListener, handler: Handler) -> Result<()> {
395        while !self.shutdown.is_triggered() {
396            // Poll rather than block in accept, so shutdown is noticed promptly.
397            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                    // EMFILE and friends: back off instead of spinning.
411                    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    /// Wait, at most `timeout`, for in-flight connections to finish.
437    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    /// Runs on the accept thread, so nothing here may block for long: a short
452    /// write, then discard whatever request bytes already arrived (closing on
453    /// unread input resets the connection and loses the 503).
454    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        // Non-blocking: a 503 fits in any socket buffer, and a client that
466        // has not finished sending is not worth waiting for.
467        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    // SAFETY: a valid socket fd and two out-parameters.
517    (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
525/// Serve one request on `stream`. Generic so tests can drive it in memory.
526pub(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            // The client may still be sending a body we refused. Closing with
548            // unread input makes the kernel send RST, which can destroy the
549            // response before the client reads it; swallow a bounded amount.
550            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/// Read and parse one request. `Err(Some(resp))` is a refusal to send back;
565/// `Err(None)` means the peer went away and there is nobody to answer.
566#[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    // Parsed successfully above on the same bytes.
618    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        // curl waits for this (up to a second) before sending a large body.
669        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
698/// Origin-form `/path?query`, or absolute-form `http://host/path?query` with
699/// the authority dropped.
700fn 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
718/// The head of a `101 Switching Protocols`; no body follows.
719fn 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        // Framing is ours; and a value with a line break would split the header.
741        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        // No length: the body runs until the connection closes.
753        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    /// An in-memory connection: each read returns at most one of `chunks`,
804    /// and everything written is recorded.
805    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        // Never completes and never stops: still bounded by the header limit.
907        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        // Peer hung up mid-request or mid-body: nobody to answer.
937        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        // A live socket is not stolen.
1022        assert!(HttpListener::bind_unix(&path).is_err());
1023        drop(l);
1024        // A stale one (left behind by a crash) is replaced.
1025        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        // Something that is not a socket is never deleted.
1031        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        // Wait until the first connection holds the only slot.
1066        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}