use crate::courierust_body::Body;
use crate::courierust_bytes::Bytes;
use crate::courierust_error::{Error, Result};
use crate::courierust_h1;
use crate::courierust_http::header::{HeaderMap, HeaderName, HeaderValue};
use crate::courierust_http::request::Request;
use crate::courierust_http::response::Response;
use crate::courierust_http::version::Version;
use crate::courierust_net::poller::{fd_of, Fd, Poller, WAKE_ID};
use crate::courierust_server::{Handler, ServerConfig};
use std::collections::{HashMap, HashSet};
use std::net::TcpStream;
use std::sync::mpsc::{channel, Receiver, Sender, TryRecvError};
use std::sync::Arc;
use std::thread;
use std::time::{Duration, Instant};
const MAX_LINE: usize = 64 * 1024;
const MAX_HEADERS: usize = 1024;
const MAX_HEADER_BLOCK: usize = 1024 * 1024;
const DISPATCH_BATCH: usize = 16;
enum EventMsg {
NewConn { id: usize, stream: TcpStream },
Register { id: usize, fd: Fd, want_write: bool },
Closed { id: usize },
}
#[derive(Debug)]
enum StepOutcome {
Idle,
NeedWrite,
Close,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum Class {
Tls,
H2,
H1,
NeedMore,
Closed,
}
const H2_PREFACE: &[u8] = b"PRI * HTTP/2.0\r\n\r\nSM\r\n\r\n";
fn classify(buf: &[u8]) -> Class {
if buf.is_empty() {
return Class::Closed;
}
if buf[0] == 0x16 {
return Class::Tls;
}
let n = buf.len().min(H2_PREFACE.len());
if buf[..n] != H2_PREFACE[..n] {
return Class::H1;
}
if buf.len() < H2_PREFACE.len() {
return Class::NeedMore;
}
Class::H2
}
#[derive(Clone, Copy)]
enum Phase {
RequestLine,
Headers,
BodyFixed { remaining: usize },
BodyChunked(Chunked),
Done,
}
#[derive(Clone, Copy)]
enum ChunkState {
Size,
Data,
Crlf,
Trailers,
}
#[derive(Clone, Copy)]
struct Chunked {
state: ChunkState,
remaining: usize,
trailer_bytes: usize,
}
struct IncrRequest {
buf: Vec<u8>,
pos: usize,
line: Vec<u8>,
req_line: Vec<u8>,
headers: HeaderMap,
body: Vec<u8>,
header_bytes: usize,
phase: Phase,
body_limit: usize,
}
impl IncrRequest {
fn new(body_limit: usize) -> Self {
Self {
buf: Vec::with_capacity(8192),
pos: 0,
line: Vec::with_capacity(128),
req_line: Vec::new(),
headers: HeaderMap::new(),
body: Vec::new(),
header_bytes: 0,
phase: Phase::RequestLine,
body_limit,
}
}
fn fill(&mut self, socket: &TcpStream) -> Result<bool> {
let mut tmp = [0u8; 8192];
let mut got = false;
loop {
let mut r: &TcpStream = socket;
match std::io::Read::read(&mut r, &mut tmp) {
Ok(0) => return Err(Error::eof()),
Ok(n) => {
got = true;
self.buf.extend_from_slice(&tmp[..n]);
}
Err(e) if e.kind() == std::io::ErrorKind::WouldBlock => return Ok(got),
Err(e) => return Err(Error::io(e.to_string())),
}
if self.buf.len() - self.pos >= 8192 {
break;
}
}
Ok(got)
}
fn read_line(&mut self, delim: u8, max: usize) -> Option<()> {
let window = &self.buf[self.pos..];
match window.iter().position(|&b| b == delim) {
Some(i) => {
self.line.extend_from_slice(&window[..i + 1]);
self.pos += i + 1;
if self.line.len() > max {
self.line.truncate(max);
}
Some(())
}
None => {
self.line.extend_from_slice(window);
self.pos = self.buf.len();
if self.line.len() > max {
self.line.truncate(max);
}
None
}
}
}
fn compact(&mut self) {
if self.pos >= 64 * 1024 {
self.buf.drain(..self.pos);
self.pos = 0;
}
}
pub(crate) fn next_request(&mut self, socket: &TcpStream) -> Result<Option<Request<Body>>> {
loop {
if let Phase::Done = self.phase {
return Ok(Some(self.finish_request()?));
}
if self.parse_step()? {
continue;
}
self.compact();
if !self.fill(socket)? {
return Ok(None);
}
}
}
fn parse_step(&mut self) -> Result<bool> {
match self.phase {
Phase::RequestLine => match self.read_line(b'\n', MAX_LINE) {
Some(()) => {
if self.line.len() >= MAX_LINE {
return Err(Error::overflow("request line too long"));
}
self.req_line = core::mem::take(&mut self.line);
self.phase = Phase::Headers;
Ok(true)
}
None => Ok(false),
},
Phase::Headers => match self.read_line(b'\n', MAX_LINE) {
Some(()) => {
if self.line.len() >= MAX_LINE {
return Err(Error::overflow("header line too long"));
}
self.header_bytes += self.line.len();
if self.header_bytes > MAX_HEADER_BLOCK {
return Err(Error::overflow("header block too large"));
}
let trimmed = courierust_h1::trim_crlf(&self.line);
if trimmed.is_empty() {
let rl = courierust_h1::parse_request_line(&self.req_line)?;
let bl = courierust_h1::body_length(&self.headers, Some(&rl.method), None)?;
self.phase = match bl {
courierust_h1::BodyLen::None => Phase::Done,
courierust_h1::BodyLen::Length(n) => {
if n > self.body_limit {
return Err(Error::overflow("request body too large"));
}
Phase::BodyFixed { remaining: n }
}
courierust_h1::BodyLen::Chunked => Phase::BodyChunked(Chunked {
state: ChunkState::Size,
remaining: 0,
trailer_bytes: 0,
}),
};
} else {
if self.headers.len() >= MAX_HEADERS {
return Err(Error::overflow("too many header fields"));
}
let (name, value) = courierust_h1::split_header(trimmed)?;
self.headers.append(name, value);
}
self.line.clear();
Ok(true)
}
None => Ok(false),
},
Phase::BodyFixed { remaining } => {
let avail = self.buf.len() - self.pos;
if avail == 0 {
return Ok(false);
}
let take = core::cmp::min(remaining, avail);
if self.body.len() + take > self.body_limit {
return Err(Error::overflow("request body too large"));
}
self.body
.extend_from_slice(&self.buf[self.pos..self.pos + take]);
self.pos += take;
let left = remaining - take;
self.phase = if left == 0 {
Phase::Done
} else {
Phase::BodyFixed { remaining: left }
};
Ok(true)
}
Phase::BodyChunked(mut ch) => {
let progressed = self.parse_chunked(&mut ch)?;
self.phase = Phase::BodyChunked(ch);
Ok(progressed)
}
Phase::Done => Ok(true),
}
}
fn parse_chunked(&mut self, ch: &mut Chunked) -> Result<bool> {
match ch.state {
ChunkState::Size => match self.read_line(b'\n', 1024) {
Some(()) => {
if self.line.len() >= 1024 {
return Err(Error::protocol("chunk size line too long"));
}
let line = core::mem::take(&mut self.line);
let sz = courierust_h1::parse_chunk_size(courierust_h1::trim_crlf(&line))
.ok_or_else(|| Error::protocol("invalid chunk size"))?;
if sz == 0 {
ch.state = ChunkState::Trailers;
} else {
ch.remaining = sz;
ch.state = ChunkState::Data;
}
Ok(true)
}
None => Ok(false),
},
ChunkState::Data => {
let avail = self.buf.len() - self.pos;
if avail == 0 {
return Ok(false);
}
let take = core::cmp::min(ch.remaining, avail);
if self.body.len() + take > self.body_limit {
return Err(Error::overflow("request body too large"));
}
self.body
.extend_from_slice(&self.buf[self.pos..self.pos + take]);
self.pos += take;
ch.remaining -= take;
if ch.remaining == 0 {
ch.state = ChunkState::Crlf;
}
Ok(true)
}
ChunkState::Crlf => {
let avail = self.buf.len() - self.pos;
if avail >= 2 {
if &self.buf[self.pos..self.pos + 2] == b"\r\n" {
self.pos += 2;
ch.state = ChunkState::Size;
Ok(true)
} else {
Err(Error::protocol("chunk terminator missing"))
}
} else {
Ok(false)
}
}
ChunkState::Trailers => match self.read_line(b'\n', MAX_LINE) {
Some(()) => {
if self.line.len() >= MAX_LINE {
return Err(Error::overflow("trailer line too long"));
}
ch.trailer_bytes += self.line.len();
if ch.trailer_bytes > MAX_HEADER_BLOCK {
return Err(Error::overflow("trailer section too large"));
}
let line = core::mem::take(&mut self.line);
if courierust_h1::trim_crlf(&line).is_empty() {
self.phase = Phase::Done;
}
Ok(true)
}
None => Ok(false),
},
}
}
fn finish_request(&mut self) -> Result<Request<Body>> {
let rl = courierust_h1::parse_request_line(&self.req_line)?;
self.req_line.clear();
let headers = core::mem::take(&mut self.headers);
self.header_bytes = 0;
let body = core::mem::take(&mut self.body);
self.phase = Phase::RequestLine;
Ok(Request {
method: rl.method,
uri: rl.target,
version: rl.version,
headers,
body: if body.is_empty() {
Body::Empty
} else {
Body::Bytes(Bytes::from(body))
},
})
}
}
struct EventConn {
socket: Arc<TcpStream>,
reader: IncrRequest,
out: Vec<u8>,
out_pos: usize,
keep_alive: bool,
}
impl EventConn {
fn new(socket: TcpStream, body_limit: usize) -> Self {
Self {
socket: Arc::new(socket),
reader: IncrRequest::new(body_limit),
out: Vec::new(),
out_pos: 0,
keep_alive: true,
}
}
fn step(&mut self, handler: &dyn Handler, config: &ServerConfig) -> Result<StepOutcome> {
loop {
if self.out_pos < self.out.len() {
return self.write_more();
}
match self.reader.next_request(&self.socket)? {
Some(req) => {
let resp = handler.handle(req);
let (wire, keep_alive) = build_response(resp, config)?;
self.out = wire;
self.out_pos = 0;
self.keep_alive = keep_alive;
match self.write_more()? {
StepOutcome::Idle => {
continue;
}
other => return Ok(other),
}
}
None => return Ok(StepOutcome::Idle),
}
}
}
fn write_more(&mut self) -> Result<StepOutcome> {
while self.out_pos < self.out.len() {
let mut w: &TcpStream = &self.socket;
match std::io::Write::write(&mut w, &self.out[self.out_pos..]) {
Ok(0) => return Err(Error::eof()),
Ok(n) => self.out_pos += n,
Err(e) if e.kind() == std::io::ErrorKind::WouldBlock => {
return Ok(StepOutcome::NeedWrite);
}
Err(e) => return Err(Error::io(e.to_string())),
}
}
self.out = Vec::new();
self.out_pos = 0;
if self.keep_alive {
Ok(StepOutcome::Idle)
} else {
Ok(StepOutcome::Close)
}
}
}
fn build_response(resp: Response<Body>, config: &ServerConfig) -> Result<(Vec<u8>, bool)> {
let keep_alive = courierust_h1::keep_alive_requested(resp.version, &resp.headers)
&& resp.version != Version::HTTP_10;
let mut out_headers = HeaderMap::with_capacity(resp.headers.len() + 3);
for (n, v) in resp.headers.iter() {
if courierust_h1::is_hop_by_hop(n.as_str()) {
continue;
}
out_headers.append(n.clone(), v.clone());
}
let chunked = matches!(resp.body, Body::Channel(_));
let body_len = match &resp.body {
Body::Bytes(b) => Some(b.len()),
_ => None,
};
if chunked {
out_headers.insert(
HeaderName::from_lowercase("transfer-encoding"),
HeaderValue::from_static("chunked"),
);
} else if let Some(n) = body_len {
let cl = courierust_h1::IToA::new(n);
out_headers.insert(
HeaderName::from_lowercase("content-length"),
HeaderValue::from_bytes(cl.as_slice())?,
);
} else if !(resp.status.is_informational()
|| resp.status == crate::courierust_http::status::StatusCode::NO_CONTENT
|| resp.status == crate::courierust_http::status::StatusCode::NOT_MODIFIED)
{
out_headers.insert(
HeaderName::from_lowercase("content-length"),
HeaderValue::from_static("0"),
);
}
out_headers.insert(
HeaderName::from_lowercase("connection"),
HeaderValue::from_static(if keep_alive { "keep-alive" } else { "close" }),
);
let mut wire = Vec::with_capacity(1024);
courierust_h1::write_response_head(&mut wire, resp.status, Version::HTTP_11, &out_headers)?;
match resp.body {
Body::Empty => {}
Body::Bytes(b) => wire.extend_from_slice(&b),
Body::Channel(rx) => {
let timeout = config.read_timeout;
loop {
let chunk = match timeout {
Some(t) => rx.recv_timeout(t).map_err(|_| ()),
None => rx.recv().map_err(|_| ()),
};
match chunk {
Ok(c) => {
let b = c?;
if b.is_empty() {
continue;
}
let sz = courierust_h1::IToA::new(b.len());
wire.extend_from_slice(sz.as_slice());
wire.extend_from_slice(b"\r\n");
wire.extend_from_slice(&b);
wire.extend_from_slice(b"\r\n");
}
Err(()) => break,
}
}
wire.extend_from_slice(b"0\r\n\r\n");
}
}
Ok((wire, keep_alive))
}
pub(crate) fn serve_event(
listener: std::net::TcpListener,
handler: Arc<dyn Handler>,
config: ServerConfig,
pool: Arc<crate::courierust_pool::ThreadPool>,
) -> std::io::Result<()> {
let (msg_tx, msg_rx) = channel::<EventMsg>();
let (ready_tx, ready_rx): (Sender<Vec<usize>>, Receiver<Vec<usize>>) = channel();
let ready_rx = Arc::new(std::sync::Mutex::new(ready_rx));
let registry: Arc<std::sync::Mutex<HashMap<usize, EventConn>>> =
Arc::new(std::sync::Mutex::new(HashMap::new()));
let (wake_reader, wake_writer) = wakeup_pair()?;
let wake_writer = Arc::new(wake_writer);
let loop_handler = handler.clone();
let loop_config = config.clone();
let loop_pool = pool.clone();
let loop_registry = registry.clone();
let event_thread = thread::Builder::new()
.name("courierust-event".into())
.spawn(move || {
event_loop(
msg_rx,
ready_tx,
loop_handler,
loop_config,
loop_pool,
loop_registry,
wake_reader,
);
})?;
let workers = if config.event_workers == 0 {
std::thread::available_parallelism()
.map(|n| n.get())
.unwrap_or(4)
} else {
config.event_workers
};
let mut worker_handles = Vec::new();
for _ in 0..workers {
let w_registry = registry.clone();
let w_handler = handler.clone();
let w_config = config.clone();
let w_ready_rx = ready_rx.clone();
let w_msg_tx = msg_tx.clone();
let w_wake = wake_writer.clone();
worker_handles.push(
thread::Builder::new()
.name("courierust-event-worker".into())
.spawn(move || {
event_worker(
w_ready_rx,
w_registry,
&*w_handler,
&w_config,
&w_msg_tx,
&w_wake,
);
})?,
);
}
let a_msg_tx = msg_tx.clone();
let a_wake = wake_writer.clone();
let accept_thread = thread::Builder::new()
.name("courierust-accept".into())
.spawn(move || {
accept_loop(listener, a_msg_tx, &a_wake);
})?;
let _ = accept_thread.join();
let _ = event_thread.join();
for h in worker_handles {
let _ = h.join();
}
Ok(())
}
fn wakeup_pair() -> std::io::Result<(TcpStream, TcpStream)> {
let listener = std::net::TcpListener::bind("127.0.0.1:0")?;
let writer = TcpStream::connect(listener.local_addr()?)?;
let (reader, _) = listener.accept()?;
reader.set_nonblocking(true)?;
writer.set_nonblocking(true)?;
Ok((reader, writer))
}
fn wake_nudge(w: &TcpStream) {
let mut s: &TcpStream = w;
let _ = std::io::Write::write(&mut s, &[1]);
}
fn drain_wake(r: &TcpStream) {
let mut buf = [0u8; 64];
loop {
let mut s: &TcpStream = r;
match std::io::Read::read(&mut s, &mut buf) {
Ok(0) | Err(_) => break,
Ok(_) => {}
}
}
}
fn handle_msg(
msg: EventMsg,
poller: &mut Poller,
pending: &mut HashMap<usize, TcpStream>,
activity: &mut HashMap<usize, Instant>,
max_connections: usize,
) {
match msg {
EventMsg::NewConn { id, stream } => {
if max_connections > 0 && activity.len() >= max_connections {
drop(stream);
return;
}
if stream.set_nonblocking(true).is_err() {
return;
}
let fd = fd_of(&stream);
pending.insert(id, stream);
activity.insert(id, Instant::now());
poller.register(id, fd, false);
}
EventMsg::Register { id, fd, want_write } => {
activity.insert(id, Instant::now());
poller.register(id, fd, want_write);
}
EventMsg::Closed { id } => {
activity.remove(&id);
}
}
}
fn event_loop(
msg_rx: Receiver<EventMsg>,
ready_tx: Sender<Vec<usize>>,
handler: Arc<dyn Handler>,
config: ServerConfig,
pool: Arc<crate::courierust_pool::ThreadPool>,
registry: Arc<std::sync::Mutex<HashMap<usize, EventConn>>>,
wake_reader: TcpStream,
) {
let mut poller = Poller::new();
let mut pending: HashMap<usize, TcpStream> = HashMap::new();
let mut activity: HashMap<usize, Instant> = HashMap::new();
let wake_fd = fd_of(&wake_reader);
let poll_timeout = config.event_poll_timeout_ms.clamp(1, 1000) as i32;
let idle_timeout = config.idle_timeout;
loop {
loop {
match msg_rx.try_recv() {
Ok(msg) => handle_msg(
msg,
&mut poller,
&mut pending,
&mut activity,
config.max_connections,
),
Err(TryRecvError::Disconnected) => return,
Err(TryRecvError::Empty) => break,
}
}
if poller.is_empty() {
match msg_rx.recv() {
Ok(msg) => handle_msg(
msg,
&mut poller,
&mut pending,
&mut activity,
config.max_connections,
),
Err(_) => return,
}
continue;
}
let wait_ms = match idle_timeout {
Some(t) => {
let now = Instant::now();
let next = activity
.values()
.map(|at| {
t.checked_sub(now.duration_since(*at))
.unwrap_or(Duration::ZERO)
})
.min()
.unwrap_or(Duration::from_secs(3600));
next.as_millis().min(poll_timeout as u128).max(1) as i32
}
None => poll_timeout,
};
let ready = match poller.wait(wait_ms, Some(wake_fd)) {
Ok(r) => r,
Err(_) => continue,
};
if ready.contains(&WAKE_ID) {
drain_wake(&wake_reader);
loop {
match msg_rx.try_recv() {
Ok(msg) => handle_msg(
msg,
&mut poller,
&mut pending,
&mut activity,
config.max_connections,
),
Err(TryRecvError::Disconnected) => return,
Err(TryRecvError::Empty) => break,
}
}
}
let mut to_dispatch: Vec<usize> = Vec::new();
for id in ready {
if id == WAKE_ID {
continue;
}
poller.unregister(id);
activity.insert(id, Instant::now());
if let Some(stream) = pending.remove(&id) {
let mut prefix = [0u8; 24];
let n = match stream.peek(&mut prefix) {
Ok(n) => n,
Err(e) if e.kind() == std::io::ErrorKind::WouldBlock => {
let fd = fd_of(&stream);
pending.insert(id, stream);
poller.register(id, fd, false);
continue;
}
Err(_) => {
activity.remove(&id);
continue;
}
};
if n == 0 {
activity.remove(&id);
continue;
}
match classify(&prefix[..n]) {
Class::Tls => {
let _ = stream.set_nonblocking(false);
let h = handler.clone();
let c = config.clone();
let p = pool.clone();
p.spawn(move || {
let _ = crate::courierust_server::serve_accepted(stream, &*h, &c);
});
activity.remove(&id);
}
Class::H2 => {
let _ = stream.set_nonblocking(false);
let h = handler.clone();
let c = config.clone();
let p = pool.clone();
p.spawn(move || {
let _ = crate::courierust_server::serve_connection(
crate::courierust_net::ConnStream::plain(stream),
&*h,
&c,
);
});
activity.remove(&id);
}
Class::H1 => {
let conn = EventConn::new(stream, config.max_body);
registry.lock().unwrap().insert(id, conn);
to_dispatch.push(id);
}
Class::NeedMore => {
let fd = fd_of(&stream);
pending.insert(id, stream);
poller.register(id, fd, false);
}
Class::Closed => {
activity.remove(&id);
}
}
} else {
to_dispatch.push(id);
}
}
if !to_dispatch.is_empty() {
for chunk in to_dispatch.chunks(DISPATCH_BATCH) {
let _ = ready_tx.send(chunk.to_vec());
}
}
if let Some(t) = idle_timeout {
let now = Instant::now();
let mut expired = Vec::new();
let registered: HashSet<usize> = registry.lock().unwrap().keys().copied().collect();
for (&id, &at) in &activity {
if now.duration_since(at) < t {
continue;
}
if pending.contains_key(&id) || registered.contains(&id) {
expired.push(id);
}
}
for id in expired {
poller.unregister(id);
pending.remove(&id);
registry.lock().unwrap().remove(&id);
activity.remove(&id);
}
}
}
}
fn event_worker(
ready_rx: Arc<std::sync::Mutex<Receiver<Vec<usize>>>>,
registry: Arc<std::sync::Mutex<HashMap<usize, EventConn>>>,
handler: &dyn Handler,
config: &ServerConfig,
msg_tx: &Sender<EventMsg>,
wake_writer: &Arc<TcpStream>,
) {
loop {
let ids = match ready_rx.lock().unwrap().recv() {
Ok(ids) => ids,
Err(_) => return,
};
for id in ids {
let mut conn = match registry.lock().unwrap().remove(&id) {
Some(c) => c,
None => continue,
};
let step = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
conn.step(handler, config)
}));
let outcome = match step {
Ok(Ok(o)) => o,
_ => StepOutcome::Close,
};
match outcome {
StepOutcome::Idle | StepOutcome::NeedWrite => {
let fd = fd_of(&conn.socket);
let want_write = matches!(outcome, StepOutcome::NeedWrite);
registry.lock().unwrap().insert(id, conn);
let _ = msg_tx.send(EventMsg::Register { id, fd, want_write });
wake_nudge(wake_writer);
}
StepOutcome::Close => {
let _ = msg_tx.send(EventMsg::Closed { id });
wake_nudge(wake_writer);
}
}
}
}
}
fn accept_loop(
listener: std::net::TcpListener,
msg_tx: Sender<EventMsg>,
wake_writer: &Arc<TcpStream>,
) {
let mut next_id = 1usize;
for stream in listener.incoming() {
let Ok(stream) = stream else { continue };
let id = next_id;
next_id += 1;
let _ = msg_tx.send(EventMsg::NewConn { id, stream });
wake_nudge(wake_writer);
}
}