//! A hand-written HTTP/1.1 client for talking to `inferproxy`. One request per connection; the //! body is offered as a `std::io::Read`. Everything read from the socket is bounded and cannot //! panic: numbers from the peer are parsed with `try_from`/`from_str_radix` and slices are taken //! with `windows`, never indexing that can run off the end. use std::io::{self, ErrorKind, Read, Write}; use std::os::unix::net::UnixStream; use std::path::Path; use std::time::Duration; pub const MAX_HEAD: usize = 16 * 1024; #[derive(Debug)] pub enum HttpError { Connect(std::io::Error), Io(std::io::Error), Timeout, Closed, TooLarge(&'static str), Malformed(&'static str), } impl std::fmt::Display for HttpError { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { match self { HttpError::Connect(e) => write!(f, "could not connect: {e}"), HttpError::Io(e) => write!(f, "io error: {e}"), HttpError::Timeout => write!(f, "read timed out before the response arrived"), HttpError::Closed => write!(f, "the peer closed the connection"), HttpError::TooLarge(what) => write!(f, "{what} is larger than allowed"), HttpError::Malformed(what) => write!(f, "{what} is malformed"), } } } impl std::error::Error for HttpError { fn source(&self) -> Option<&(dyn std::error::Error + 'static)> { match self { HttpError::Connect(e) | HttpError::Io(e) => Some(e), _ => None, } } } /// True for the two error kinds a read timeout produces: WouldBlock and TimedOut. pub fn is_timeout(e: &std::io::Error) -> bool { matches!(e.kind(), ErrorKind::WouldBlock | ErrorKind::TimedOut) } pub struct Request<'a> { pub method: &'a str, pub path: &'a str, pub body: Option<&'a [u8]>, } #[derive(Debug, Clone, PartialEq, Eq)] pub struct Head { pub status: u16, pub headers: Vec<(String, String)>, } impl Head { pub fn header(&self, name: &str) -> Option<&str> { let name = name.to_ascii_lowercase(); self.headers .iter() .find(|(n, _)| *n == name) .map(|(_, v)| v.as_str()) } } fn find_needle(haystack: &[u8], needle: &[u8]) -> Option { if needle.is_empty() || haystack.len() < needle.len() { return None; } haystack.windows(needle.len()).position(|w| w == needle) } fn parse_status(line: &str) -> Result { let mut parts = line.split_whitespace(); let version = parts.next(); let code = parts.next(); if version != Some("HTTP/1.1") && version != Some("HTTP/1.0") { return Err(HttpError::Malformed("status line")); } let code: u16 = code .and_then(|c| c.parse::().ok()) .ok_or(HttpError::Malformed("status line"))?; if !(100..=599).contains(&code) { return Err(HttpError::Malformed("status line")); } Ok(code) } fn parse_head(bytes: &[u8]) -> Result { let text = std::str::from_utf8(bytes).map_err(|_| HttpError::Malformed("head"))?; let mut lines = text.split("\r\n"); let status_line = lines.next().ok_or(HttpError::Malformed("head"))?; let status = parse_status(status_line)?; let mut headers = Vec::new(); for line in lines { let (name, value) = line .split_once(':') .ok_or(HttpError::Malformed("header line"))?; if name.is_empty() || name.chars().any(char::is_whitespace) { return Err(HttpError::Malformed("header line")); } headers.push((name.to_ascii_lowercase(), value.trim().to_string())); } Ok(Head { status, headers }) } fn extend(out: &mut Vec, src: impl AsRef<[u8]>) { out.extend_from_slice(src.as_ref()); } pub struct Connection { stream: UnixStream, buf: Vec, received: bool, } impl Connection { pub fn open(socket: &Path) -> Result { let stream = UnixStream::connect(socket).map_err(HttpError::Connect)?; Ok(Connection { stream, buf: Vec::new(), received: false, }) } pub fn send(&mut self, req: &Request<'_>) -> Result<(), HttpError> { let mut out: Vec = Vec::with_capacity(256); extend( &mut out, format!("{} {} HTTP/1.1\r\n", req.method, req.path), ); extend(&mut out, "Host: localhost\r\n"); extend(&mut out, "Connection: close\r\n"); extend(&mut out, "Accept: */*\r\n"); if let Some(body) = req.body { extend(&mut out, "Content-Type: application/json\r\n"); extend(&mut out, format!("Content-Length: {}\r\n", body.len())); } extend(&mut out, "\r\n"); if let Some(body) = req.body { extend(&mut out, body); } self.stream.write_all(&out).map_err(HttpError::Io) } pub fn set_read_timeout(&self, timeout: Duration) -> Result<(), HttpError> { crate::socket::set_read_timeout(&self.stream, timeout).map_err(HttpError::Io) } pub fn received_any(&self) -> bool { self.received } pub fn read_head(&mut self) -> Result { let mut tmp = [0u8; 8192]; loop { if let Some(pos) = find_needle(&self.buf, b"\r\n\r\n") { let head = parse_head(&self.buf[..pos])?; self.buf.drain(..pos + 4); return Ok(head); } if self.buf.len() > MAX_HEAD { return Err(HttpError::TooLarge("head")); } match self.stream.read(&mut tmp) { Ok(0) => return Err(HttpError::Closed), Ok(n) => { self.received = true; self.buf.extend_from_slice(&tmp[..n]); } Err(e) if is_timeout(&e) => return Err(HttpError::Timeout), Err(e) => return Err(HttpError::Io(e)), } } } pub fn body(&mut self, head: &Head) -> Result, HttpError> { let framing = if head .header("transfer-encoding") .is_some_and(|v| v.eq_ignore_ascii_case("chunked")) { Framing::Chunked } else if let Some(len) = head.header("content-length") { Framing::Length(parse_length(len)?) } else { Framing::Close }; let buf = std::mem::take(&mut self.buf); let stream = &mut self.stream; Ok(Body { stream, buf, framing, phase: ChunkPhase::Size, chunk_remaining: 0, finished: false, }) } } fn parse_length(value: &str) -> Result { if value.is_empty() { return Err(HttpError::Malformed("content-length")); } let n: u64 = value .parse() .map_err(|_| HttpError::Malformed("content-length"))?; usize::try_from(n).map_err(|_| HttpError::TooLarge("content-length")) } fn bad_chunk() -> io::Error { io::Error::new(ErrorKind::InvalidData, "malformed chunk") } fn parse_chunk_size(token: &[u8]) -> Result { let semi = token.iter().position(|&b| b == b';').unwrap_or(token.len()); let hex = &token[..semi]; if hex.is_empty() { return Err(bad_chunk()); } for &b in hex { if !b.is_ascii_hexdigit() { return Err(bad_chunk()); } } let text = std::str::from_utf8(hex).map_err(|_| bad_chunk())?; usize::from_str_radix(text, 16).map_err(|_| bad_chunk()) } enum Framing { Chunked, Length(usize), Close, } enum ChunkPhase { Size, Data, Crlf, Trailers, } pub struct Body<'a> { stream: &'a mut UnixStream, buf: Vec, framing: Framing, phase: ChunkPhase, chunk_remaining: usize, finished: bool, } impl<'a> Body<'a> { fn read_chunked(&mut self, buf: &mut [u8]) -> io::Result { if self.finished { return Ok(0); } let mut tmp = [0u8; 8192]; let mut written = 0usize; loop { match self.phase { ChunkPhase::Size => loop { match find_needle(&self.buf, b"\r\n") { Some(pos) => { let size = parse_chunk_size(&self.buf[..pos]).map_err(|_| bad_chunk())?; self.buf.drain(..pos + 2); if size == 0 { self.phase = ChunkPhase::Trailers; } else { self.phase = ChunkPhase::Data; self.chunk_remaining = size; } break; } None if self.buf.len() > 64 => return Err(bad_chunk()), None => { let k = self.stream.read(&mut tmp)?; if k == 0 { return Err(ErrorKind::UnexpectedEof.into()); } self.buf.extend_from_slice(&tmp[..k]); } } }, ChunkPhase::Data => { if self.buf.is_empty() { let k = self.stream.read(&mut tmp)?; if k == 0 { return Err(ErrorKind::UnexpectedEof.into()); } self.buf.extend_from_slice(&tmp[..k]); } let want = (buf.len() - written).min(self.chunk_remaining); let take = self.buf.len().min(want); buf[written..written + take].copy_from_slice(&self.buf[..take]); self.buf.drain(..take); written += take; self.chunk_remaining -= take; if self.chunk_remaining == 0 { self.phase = ChunkPhase::Crlf; } return Ok(written); } ChunkPhase::Crlf => { if self.buf.len() >= 2 { if &self.buf[..2] == b"\r\n" { self.buf.drain(..2); self.phase = ChunkPhase::Size; } else { return Err(bad_chunk()); } } else { let k = self.stream.read(&mut tmp)?; if k == 0 { return Err(ErrorKind::UnexpectedEof.into()); } self.buf.extend_from_slice(&tmp[..k]); } } ChunkPhase::Trailers => loop { match find_needle(&self.buf, b"\r\n") { Some(pos) => { if pos == 0 { self.buf.drain(..2); self.finished = true; return Ok(written); } self.buf.drain(..pos + 2); } None if self.buf.len() > 512 => return Err(bad_chunk()), None => { let k = self.stream.read(&mut tmp)?; if k == 0 { return Err(ErrorKind::UnexpectedEof.into()); } self.buf.extend_from_slice(&tmp[..k]); } } }, } } } fn read_length(&mut self, buf: &mut [u8]) -> io::Result { let Framing::Length(remaining) = self.framing else { return Ok(0); }; if remaining == 0 { return Ok(0); } let want = buf.len().min(remaining); let from_buf = self.buf.len().min(want); buf[..from_buf].copy_from_slice(&self.buf[..from_buf]); self.buf.drain(..from_buf); let mut got = from_buf; if got < want { let space = want - got; let k = self.stream.read(&mut buf[got..got + space])?; if k == 0 { return Err(ErrorKind::UnexpectedEof.into()); } got += k; } self.framing = Framing::Length(remaining - got); Ok(got) } fn read_close(&mut self, buf: &mut [u8]) -> io::Result { let take = self.buf.len().min(buf.len()); buf[..take].copy_from_slice(&self.buf[..take]); self.buf.drain(..take); if take > 0 { return Ok(take); } self.stream.read(buf) } } impl std::io::Read for Body<'_> { fn read(&mut self, buf: &mut [u8]) -> io::Result { match self.framing { Framing::Chunked => self.read_chunked(buf), Framing::Length(_) => self.read_length(buf), Framing::Close => self.read_close(buf), } } } pub fn read_capped(body: &mut impl Read, cap: usize) -> Result, HttpError> { let mut out: Vec = Vec::new(); let mut tmp = [0u8; 8192]; loop { let n = body.read(&mut tmp).map_err(|e| { if is_timeout(&e) { HttpError::Timeout } else if e.kind() == ErrorKind::UnexpectedEof { HttpError::Closed } else { HttpError::Io(e) } })?; if n == 0 { return Ok(out); } if out.len() + n > cap { return Err(HttpError::TooLarge("body")); } out.extend_from_slice(&tmp[..n]); } }