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 Tls(TcpListener, TlsConfig),
222}
223
224pub type TlsConfig = Arc<std::sync::RwLock<Arc<rustls::ServerConfig>>>;
227
228#[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 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 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 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 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 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#[derive(Debug)]
380pub struct HttpServer {
381 limits: Limits,
382 active: AtomicUsize,
383 shutdown: Shutdown,
384}
385
386struct 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 pub fn active(&self) -> usize {
411 self.active.load(Ordering::SeqCst)
412 }
413
414 pub fn run(self: &Arc<Self>, listener: HttpListener, handler: Handler) -> Result<()> {
417 while !self.shutdown.is_triggered() {
418 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 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 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
477struct 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 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 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 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 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 (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
603pub(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 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#[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 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 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
776fn 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
796fn 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 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 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 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 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 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 assert!(HttpListener::bind_unix(&path).is_err());
1101 drop(l);
1102 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 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 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}