use std::collections::HashMap;
use std::io::{BufRead, BufReader, Read, Write};
use std::net::TcpStream;
use std::time::Duration;
const MAX_HEADER_BYTES: usize = 16 * 1024;
const READ_TIMEOUT: Duration = Duration::from_secs(10);
const WRITE_TIMEOUT: Duration = Duration::from_secs(30);
pub struct Request {
pub method: String,
pub path: String,
pub query: HashMap<String, String>,
}
impl Request {
pub fn parse(stream: &TcpStream) -> Result<Request, (u16, &'static str)> {
let _ = stream.set_read_timeout(Some(READ_TIMEOUT));
let _ = stream.set_write_timeout(Some(WRITE_TIMEOUT));
let mut reader = BufReader::new(stream.take(MAX_HEADER_BYTES as u64));
let mut line = String::new();
if reader.read_line(&mut line).is_err() || line.is_empty() {
return Err((400, "malformed request line"));
}
let mut parts = line.split_whitespace();
let method = parts.next().unwrap_or_default().to_string();
let target = parts.next().unwrap_or_default();
if method != "GET" && method != "HEAD" {
return Err((405, "this server only answers GET"));
}
if target.is_empty() {
return Err((400, "malformed request line"));
}
let (raw_path, raw_query) = match target.split_once('?') {
Some((p, q)) => (p, q),
None => (target, ""),
};
let path = percent_decode(raw_path);
if !path.starts_with('/') || path.contains("..") || path.contains('\0') {
return Err((400, "unacceptable path"));
}
let query = raw_query
.split('&')
.filter(|p| !p.is_empty())
.map(|pair| match pair.split_once('=') {
Some((k, v)) => (percent_decode(k), percent_decode(v)),
None => (percent_decode(pair), String::new()),
})
.collect();
loop {
let mut header = String::new();
match reader.read_line(&mut header) {
Ok(0) => break,
Ok(_) if header.trim().is_empty() => break,
Ok(_) => {}
Err(_) => return Err((431, "request headers too large")),
}
}
Ok(Request {
method,
path,
query,
})
}
pub fn token(&self) -> &str {
self.query.get("t").map_or("", String::as_str)
}
}
fn percent_decode(s: &str) -> String {
let bytes = s.as_bytes();
let mut out: Vec<u8> = Vec::with_capacity(bytes.len());
let mut i = 0;
while i < bytes.len() {
match bytes[i] {
b'%' if i + 2 < bytes.len() => {
let hex = std::str::from_utf8(&bytes[i + 1..i + 3]).ok();
match hex.and_then(|h| u8::from_str_radix(h, 16).ok()) {
Some(byte) => {
out.push(byte);
i += 3;
}
None => {
out.push(b'%');
i += 1;
}
}
}
b'+' => {
out.push(b' ');
i += 1;
}
byte => {
out.push(byte);
i += 1;
}
}
}
String::from_utf8_lossy(&out).into_owned()
}
fn reason(status: u16) -> &'static str {
match status {
200 => "OK",
400 => "Bad Request",
403 => "Forbidden",
404 => "Not Found",
405 => "Method Not Allowed",
431 => "Request Header Fields Too Large",
503 => "Service Unavailable",
_ => "Error",
}
}
fn common_headers(out: &mut String) {
out.push_str(
"X-Content-Type-Options: nosniff\r\n\
Referrer-Policy: no-referrer\r\n\
Content-Security-Policy: default-src 'none'; \
style-src 'unsafe-inline'; \
script-src 'unsafe-inline'; \
img-src data:; \
connect-src 'self'; \
base-uri 'none'; \
form-action 'none'; \
frame-ancestors 'none'\r\n",
);
}
pub fn respond(
stream: &mut TcpStream,
request: Option<&Request>,
status: u16,
content_type: &str,
body: &[u8],
) {
let mut head = format!(
"HTTP/1.1 {status} {}\r\n\
Content-Type: {content_type}\r\n\
Content-Length: {}\r\n\
Connection: close\r\n\
Cache-Control: no-store\r\n",
reason(status),
body.len(),
);
common_headers(&mut head);
head.push_str("\r\n");
let head_only = request.is_some_and(|r| r.method == "HEAD");
let mut buf = head.into_bytes();
if !head_only {
buf.extend_from_slice(body);
}
let _ = stream.write_all(&buf);
let _ = stream.flush();
}
pub fn respond_error(stream: &mut TcpStream, request: Option<&Request>, status: u16, msg: &str) {
respond(
stream,
request,
status,
"text/plain; charset=utf-8",
format!("{status} {}: {msg}\n", reason(status)).as_bytes(),
);
}
pub struct EventStream<'a> {
stream: &'a mut TcpStream,
}
impl<'a> EventStream<'a> {
pub fn open(stream: &'a mut TcpStream) -> std::io::Result<EventStream<'a>> {
let mut head = String::from(
"HTTP/1.1 200 OK\r\n\
Content-Type: text/event-stream; charset=utf-8\r\n\
Cache-Control: no-store\r\n\
Connection: close\r\n",
);
common_headers(&mut head);
head.push_str("\r\n");
stream.write_all(head.as_bytes())?;
stream.flush()?;
Ok(EventStream { stream })
}
pub fn send(&mut self, event: &str, data: &str) -> std::io::Result<()> {
let mut frame = format!("event: {event}\n");
for line in data.split('\n') {
frame.push_str("data: ");
frame.push_str(line);
frame.push('\n');
}
frame.push('\n');
self.stream.write_all(frame.as_bytes())?;
self.stream.flush()
}
pub fn keepalive(&mut self) -> std::io::Result<()> {
self.stream.write_all(b": keepalive\n\n")?;
self.stream.flush()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn percent_decoding_handles_escapes_and_plus() {
assert_eq!(percent_decode("/a%2Fb"), "/a/b");
assert_eq!(percent_decode("hello+world"), "hello world");
assert_eq!(percent_decode("100%"), "100%");
assert_eq!(percent_decode("%zz"), "%zz");
assert_eq!(percent_decode("caf%C3%A9"), "café");
}
#[test]
fn a_truncated_escape_stays_literal() {
assert_eq!(percent_decode("/x%2"), "/x%2");
}
#[test]
fn every_status_the_router_sends_has_a_reason() {
for status in [200, 400, 403, 404, 405, 431, 503] {
assert_ne!(reason(status), "Error", "status {status} has no reason");
}
}
#[test]
fn the_policy_forbids_loading_anything_off_the_network() {
let mut headers = String::new();
common_headers(&mut headers);
assert!(headers.contains("default-src 'none'"));
assert!(headers.contains("frame-ancestors 'none'"));
assert!(headers.contains("nosniff"));
}
}