420 lines
14 KiB
Rust
420 lines
14 KiB
Rust
//! 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<usize> {
|
|
if needle.is_empty() || haystack.len() < needle.len() {
|
|
return None;
|
|
}
|
|
haystack.windows(needle.len()).position(|w| w == needle)
|
|
}
|
|
|
|
fn parse_status(line: &str) -> Result<u16, HttpError> {
|
|
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::<u16>().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<Head, HttpError> {
|
|
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<u8>, src: impl AsRef<[u8]>) {
|
|
out.extend_from_slice(src.as_ref());
|
|
}
|
|
|
|
pub struct Connection {
|
|
stream: UnixStream,
|
|
buf: Vec<u8>,
|
|
received: bool,
|
|
}
|
|
|
|
impl Connection {
|
|
pub fn open(socket: &Path) -> Result<Self, HttpError> {
|
|
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<u8> = 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> {
|
|
self.stream
|
|
.set_read_timeout(Some(timeout))
|
|
.map_err(HttpError::Io)
|
|
}
|
|
|
|
pub fn received_any(&self) -> bool {
|
|
self.received
|
|
}
|
|
|
|
pub fn read_head(&mut self) -> Result<Head, HttpError> {
|
|
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<Body<'_>, 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<usize, HttpError> {
|
|
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<usize, io::Error> {
|
|
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<u8>,
|
|
framing: Framing,
|
|
phase: ChunkPhase,
|
|
chunk_remaining: usize,
|
|
finished: bool,
|
|
}
|
|
|
|
impl<'a> Body<'a> {
|
|
fn read_chunked(&mut self, buf: &mut [u8]) -> io::Result<usize> {
|
|
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<usize> {
|
|
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<usize> {
|
|
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<usize> {
|
|
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<Vec<u8>, HttpError> {
|
|
let mut out: Vec<u8> = 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]);
|
|
}
|
|
}
|