use crate::courierust_body::{Body, ChannelStream};
use crate::courierust_bytes::Bytes;
use crate::courierust_error::{Error, Result};
use crate::courierust_h1;
use crate::courierust_http::header::HeaderMap;
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_net::stats::Stats;
use crate::courierust_server::{Handler, ServerConfig};
use std::collections::{HashMap, HashSet};
use std::net::TcpStream;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::mpsc::{channel, Receiver, RecvTimeoutError, 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;
const BODY_POLL_MIN_MS: u64 = 1;
const BODY_POLL_MAX_MS: u64 = 16;
const BODY_POLL_WAKE_MS: u64 = 25;
fn h1_trace() -> bool {
static ENABLED: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*ENABLED.get_or_init(|| std::env::var_os("COURIERUST_H1_TRACE").is_some())
}
#[inline]
fn seg_start(enabled: bool) -> Option<Instant> {
if enabled {
Some(Instant::now())
} else {
None
}
}
#[inline]
fn seg_end(acc: &mut u64, start: Option<Instant>) {
if let Some(start) = start {
*acc = acc.saturating_add(start.elapsed().as_micros() as u64);
}
}
enum EventMsg {
NewConn {
id: usize,
stream: TcpStream,
accepted_at: Option<Instant>,
},
Register {
id: usize,
fd: Fd,
want_write: bool,
},
BodyChunk {
id: usize,
},
Closed {
id: usize,
socket: Option<Arc<TcpStream>>,
},
Detach {
id: usize,
socket: Option<Arc<TcpStream>>,
},
}
#[derive(Clone)]
struct Registries {
h1: Arc<std::sync::Mutex<HashMap<usize, EventConn>>>,
ws: Arc<std::sync::Mutex<HashMap<usize, crate::courierust_server::ws::WsEventConn>>>,
}
impl Registries {
fn new() -> Self {
Self {
h1: Arc::new(std::sync::Mutex::new(HashMap::new())),
ws: Arc::new(std::sync::Mutex::new(HashMap::new())),
}
}
}
enum StepOutcome {
Idle,
NeedWrite,
Close,
Upgrade(Box<crate::courierust_server::ws::WsEventConn>),
Tunnel(Box<TunnelJob>),
}
struct TunnelJob {
plan: crate::courierust_server::TunnelPlan,
socket: Arc<TcpStream>,
leftover: Vec<u8>,
secure: bool,
}
#[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>,
parsed_req_line: Option<crate::courierust_h1::RequestLine>,
headers: HeaderMap,
body: Vec<u8>,
header_bytes: usize,
header_limit: usize,
phase: Phase,
body_limit: usize,
trace: bool,
first_read_at: Option<Instant>,
header_started_at: Option<Instant>,
}
impl IncrRequest {
fn new(body_limit: usize, header_limit: usize, trace: bool) -> Self {
Self {
buf: Vec::with_capacity(8192),
pos: 0,
line: Vec::with_capacity(128),
req_line: Vec::new(),
parsed_req_line: None,
headers: HeaderMap::new(),
body: Vec::new(),
header_bytes: 0,
header_limit,
phase: Phase::RequestLine,
body_limit,
trace,
first_read_at: None,
header_started_at: None,
}
}
fn fill(&mut self, socket: &TcpStream, reads: Option<&AtomicUsize>) -> Result<bool> {
let mut tmp = [0u8; 8192];
let mut got = false;
loop {
if let Some(reads) = reads {
reads.fetch_add(1, Ordering::Relaxed);
}
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]);
if self.trace && self.first_read_at.is_none() {
self.first_read_at = Some(Instant::now());
}
}
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 take_remaining(&mut self) -> Vec<u8> {
let tail = self.buf[self.pos..].to_vec();
self.buf.clear();
self.pos = 0;
self.line.clear();
self.phase = Phase::RequestLine;
tail
}
pub(crate) fn next_request(
&mut self,
socket: &TcpStream,
reads: Option<&AtomicUsize>,
header_timeout: Option<Duration>,
) -> Result<Option<Request<Body>>> {
loop {
if matches!(self.phase, Phase::RequestLine | Phase::Headers)
&& self.header_started_at.is_none()
{
self.header_started_at = Some(Instant::now());
}
if self.header_timeout_expired(header_timeout, Instant::now()) {
return Err(Error::timeout("request header timeout"));
}
if let Phase::Done = self.phase {
return Ok(Some(self.finish_request()?));
}
if self.parse_step()? {
continue;
}
self.compact();
if !self.fill(socket, reads)? {
return Ok(None);
}
}
}
fn header_timeout_expired(&self, timeout: Option<Duration>, now: Instant) -> bool {
match (self.header_started_at, timeout) {
(Some(start), Some(timeout)) => now.duration_since(start) >= timeout,
_ => false,
}
}
fn header_timeout_remaining(&self, timeout: Duration, now: Instant) -> Option<Duration> {
self.header_started_at
.map(|start| timeout.saturating_sub(now.duration_since(start)))
}
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 > self.header_limit {
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.parsed_req_line = Some(rl);
self.header_started_at = None;
self.phase = match bl {
courierust_h1::BodyLen::None => Phase::Done,
courierust_h1::BodyLen::Length(0) => Phase::Done,
courierust_h1::BodyLen::Length(n) => {
if n > self.body_limit {
return Err(Error::overflow("request body too large"));
}
if n == 0 {
Phase::Done
} else {
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 = self
.parsed_req_line
.take()
.ok_or_else(|| Error::protocol("request line not parsed"))?;
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 BodyRx {
stream: ChannelStream,
wait_started: Option<Instant>,
empty_polls: u32,
poll_at: Option<Instant>,
}
impl BodyRx {
fn new(stream: ChannelStream) -> Self {
Self {
stream,
wait_started: None,
empty_polls: 0,
poll_at: None,
}
}
}
fn body_poll_delay(body: &BodyRx) -> Duration {
if body.stream.has_wake() {
return Duration::from_millis(BODY_POLL_WAKE_MS);
}
let shift = body.empty_polls.min(4);
Duration::from_millis((BODY_POLL_MIN_MS << shift).min(BODY_POLL_MAX_MS))
}
struct EventConn {
socket: Arc<TcpStream>,
peer: std::net::SocketAddr,
pending_upgrade: Option<(
crate::courierust_server::ws::WsPlan,
Arc<dyn crate::courierust_server::ws::WsService>,
)>,
pending_tunnel: Option<crate::courierust_server::TunnelPlan>,
wake_slot: Arc<crate::courierust_server::ws::WakeSlot>,
reader: IncrRequest,
out: Vec<u8>,
out_pos: usize,
stream_rx: Option<BodyRx>,
wake_armed: bool,
keep_alive: bool,
reads: Option<Arc<AtomicUsize>>,
writes: Option<Arc<AtomicUsize>>,
trace: bool,
parse_us: u64,
handler_us: u64,
build_us: u64,
write_us: u64,
dispatch_us: u64,
trace_requests: u64,
parked_at: Option<Instant>,
registered_at: Option<Instant>,
pickup_at: Option<Instant>,
}
impl EventConn {
fn new(
socket: TcpStream,
body_limit: usize,
header_limit: usize,
stats: Option<&Stats>,
wake_slot: Arc<crate::courierust_server::ws::WakeSlot>,
) -> Self {
let (reads, writes) = match stats {
Some(s) => (
Some(s.h1_read_syscalls.clone()),
Some(s.h1_write_syscalls.clone()),
),
None => (None, None),
};
let trace = h1_trace();
let peer = socket
.peer_addr()
.unwrap_or_else(|_| std::net::SocketAddr::from(([0, 0, 0, 0], 0)));
Self {
socket: Arc::new(socket),
peer,
pending_upgrade: None,
pending_tunnel: None,
wake_slot,
reader: IncrRequest::new(body_limit, header_limit, trace),
out: Vec::new(),
out_pos: 0,
stream_rx: None,
wake_armed: false,
keep_alive: true,
reads,
writes,
trace,
parse_us: 0,
handler_us: 0,
build_us: 0,
write_us: 0,
dispatch_us: 0,
trace_requests: 0,
parked_at: None,
registered_at: trace.then(Instant::now),
pickup_at: None,
}
}
fn has_pending_output(&self) -> bool {
self.out_pos < self.out.len()
}
fn arm_wake(&mut self, make: impl FnOnce() -> Arc<dyn Fn() + Send + Sync>) {
if self.wake_armed {
return;
}
self.wake_slot.set(make());
self.wake_armed = true;
}
fn attach_body_wake(&self, stream: &ChannelStream) {
let slot = self.wake_slot.clone();
stream.install_wake(move || slot.fire());
}
fn set_stream(&mut self, stream: Option<ChannelStream>) {
if let Some(stream) = &stream {
self.attach_body_wake(stream);
}
self.stream_rx = stream.map(BodyRx::new);
}
fn body_poll_due(&self, now: Instant) -> bool {
self.body_poll_at().is_some_and(|at| at <= now)
}
fn body_poll_at(&self) -> Option<Instant> {
self.stream_rx.as_ref().and_then(|body| body.poll_at)
}
fn step(&mut self, handler: &dyn Handler, config: &ServerConfig) -> Result<StepOutcome> {
let trace = self.trace;
loop {
if self.out_pos < self.out.len() {
let outcome = self.write_more()?;
if !matches!(outcome, StepOutcome::Idle) {
return Ok(outcome);
}
}
if self.stream_rx.is_some() {
let outcome = self.pump_body(config)?;
if !matches!(outcome, StepOutcome::Idle) {
return Ok(outcome);
}
}
if let Some((plan, service)) = self.pending_upgrade.take() {
let leftover = self.reader.buf[self.reader.pos..].to_vec();
let conn = crate::courierust_server::ws::WsEventConn::new(
self.socket.clone(),
&leftover,
plan,
service,
&config.websocket,
self.wake_slot.clone(),
);
return Ok(StepOutcome::Upgrade(Box::new(conn)));
}
if let Some(plan) = self.pending_tunnel.take() {
return Ok(StepOutcome::Tunnel(Box::new(TunnelJob {
plan,
socket: self.socket.clone(),
leftover: self.reader.take_remaining(),
secure: config.tls.is_some(),
})));
}
let parse = seg_start(trace);
let request = match self.reader.next_request(
&self.socket,
self.reads.as_deref(),
config.request_header_timeout,
) {
Ok(request) => request,
Err(Error {
kind: crate::courierust_error::ErrorKind::Timeout,
..
}) => {
self.out.clear();
let (keep_alive, stream) = build_response(
crate::courierust_server::h1::error_response(408, "request header timeout"),
true,
false,
&mut self.out,
)?;
self.keep_alive = keep_alive;
self.out_pos = 0;
self.set_stream(stream);
return self.write_more();
}
Err(error) => return Err(error),
};
match request {
Some(req) => {
if trace {
if let Some(first_read) = self.reader.first_read_at.take() {
let done = Instant::now();
if let Some(pickup) = self.pickup_at.take() {
if first_read >= pickup {
self.dispatch_us = self.dispatch_us.saturating_add(
first_read.duration_since(pickup).as_micros() as u64,
);
}
}
self.parse_us = self
.parse_us
.saturating_add(done.duration_since(first_read).as_micros() as u64);
} else {
seg_end(&mut self.parse_us, parse);
}
} else {
seg_end(&mut self.parse_us, parse);
}
let request_close = courierust_h1::wants_close(&req.headers);
let is_head = req.method == crate::courierust_http::method::Method::HEAD;
if let Some(reason) =
courierust_h1::host_header_error(req.version, &req.headers)
{
self.out.clear();
let (keep_alive, stream) = build_response(
crate::courierust_server::h1::error_response(400, reason),
request_close,
is_head,
&mut self.out,
)?;
self.out_pos = 0;
self.set_stream(stream);
self.keep_alive = keep_alive;
let outcome = self.write_more()?;
match outcome {
StepOutcome::Idle => continue,
other => return Ok(other),
}
}
match self.websocket_decision(handler, &req, config)? {
WsDecision::Pass => {}
WsDecision::Respond(resp) => {
let handle = seg_start(trace);
seg_end(&mut self.handler_us, handle);
self.out.clear();
let (keep_alive, stream) =
build_response(resp, request_close, is_head, &mut self.out)?;
self.out_pos = 0;
self.set_stream(stream);
self.keep_alive = keep_alive;
let outcome = self.write_more()?;
match outcome {
StepOutcome::Idle => continue,
StepOutcome::Close if self.stream_rx.is_some() => continue,
other => return Ok(other),
}
}
WsDecision::Upgrade(upgrade) => {
let head = upgrade.plan.accept_headers()?;
self.out.clear();
courierust_h1::write_response_head(
&mut self.out,
crate::courierust_http::status::StatusCode::SWITCHING_PROTOCOLS,
Version::HTTP_11,
&head,
)?;
self.out_pos = 0;
self.keep_alive = true;
self.pending_upgrade = Some((upgrade.plan, upgrade.service));
continue;
}
}
match handler.tunnel(&req) {
crate::courierust_server::TunnelReply::Pass => {}
crate::courierust_server::TunnelReply::Refuse(resp) => {
let handle = seg_start(trace);
seg_end(&mut self.handler_us, handle);
self.out.clear();
let (keep_alive, stream) =
build_response(resp, request_close, is_head, &mut self.out)?;
self.out_pos = 0;
self.set_stream(stream);
self.keep_alive = keep_alive;
let outcome = self.write_more()?;
match outcome {
StepOutcome::Idle => continue,
StepOutcome::Close if self.stream_rx.is_some() => continue,
other => return Ok(other),
}
}
crate::courierust_server::TunnelReply::Accept(plan) => {
let handle = seg_start(trace);
seg_end(&mut self.handler_us, handle);
self.out.clear();
courierust_h1::write_response_head(
&mut self.out,
plan.status,
Version::HTTP_11,
&plan.headers,
)?;
self.out_pos = 0;
self.keep_alive = false;
self.pending_tunnel = Some(plan);
continue;
}
}
let handle = seg_start(trace);
let connection_info = crate::courierust_server::ConnectionInfo {
peer: self.peer,
secure: false,
};
let resp = handler.handle_connected(&connection_info, req);
seg_end(&mut self.handler_us, handle);
self.out.clear();
let build = seg_start(trace);
let (keep_alive, stream) =
build_response(resp, request_close, is_head, &mut self.out)?;
seg_end(&mut self.build_us, build);
self.out_pos = 0;
self.set_stream(stream);
self.keep_alive = keep_alive;
let write = seg_start(trace);
let outcome = self.write_more()?;
seg_end(&mut self.write_us, write);
if trace {
self.trace_requests = self.trace_requests.saturating_add(1);
}
match outcome {
StepOutcome::Idle => {
continue;
}
StepOutcome::Close if self.stream_rx.is_some() => continue,
other => return Ok(other),
}
}
None => {
seg_end(&mut self.parse_us, parse);
if trace {
self.reader.first_read_at = None;
}
return Ok(StepOutcome::Idle);
}
}
}
}
fn header_timeout_remaining(&self, timeout: Duration, now: Instant) -> Option<Duration> {
self.reader.header_timeout_remaining(timeout, now)
}
fn pump_body(&mut self, config: &ServerConfig) -> Result<StepOutcome> {
let Some(mut body) = self.stream_rx.take() else {
return Ok(StepOutcome::Idle);
};
loop {
match body.stream.try_recv() {
Ok(Ok(chunk)) => {
if chunk.is_empty() {
continue;
}
body.wait_started = None;
body.empty_polls = 0;
self.out.clear();
self.out_pos = 0;
courierust_h1::encode_chunk(&chunk, &mut self.out);
if matches!(self.write_more()?, StepOutcome::NeedWrite) {
self.stream_rx = Some(body);
return Ok(StepOutcome::NeedWrite);
}
}
Ok(Err(error)) => return Err(error),
Err(TryRecvError::Empty) => {
if let Some(timeout) = config.read_timeout {
let started = *body.wait_started.get_or_insert_with(Instant::now);
if started.elapsed() >= timeout {
return Err(Error::timeout("body stream timed out"));
}
}
body.empty_polls = body.empty_polls.saturating_add(1);
body.poll_at = Some(Instant::now() + body_poll_delay(&body));
self.stream_rx = Some(body);
return Ok(StepOutcome::Idle);
}
Err(TryRecvError::Disconnected) => {
body.stream.clear_wake();
self.out.clear();
self.out_pos = 0;
self.out.extend_from_slice(courierust_h1::CHUNKED_END);
return self.write_more();
}
}
}
}
fn write_more(&mut self) -> Result<StepOutcome> {
while self.out_pos < self.out.len() {
if let Some(writes) = &self.writes {
writes.fetch_add(1, Ordering::Relaxed);
}
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.clear();
self.out_pos = 0;
if self.keep_alive {
Ok(StepOutcome::Idle)
} else {
Ok(StepOutcome::Close)
}
}
}
enum WsDecision {
Pass,
Respond(Response<Body>),
Upgrade(Box<WsUpgrade>),
}
struct WsUpgrade {
plan: crate::courierust_server::ws::WsPlan,
service: Arc<dyn crate::courierust_server::ws::WsService>,
}
impl EventConn {
fn websocket_decision(
&self,
handler: &dyn Handler,
req: &Request<Body>,
config: &ServerConfig,
) -> Result<WsDecision> {
if !config.websocket.enabled || !crate::courierust_ws::is_websocket_upgrade(&req.headers) {
return Ok(WsDecision::Pass);
}
let (service, protocol) = match handler.websocket(req) {
crate::courierust_server::ws::WsUpgradeReply::Pass => return Ok(WsDecision::Pass),
crate::courierust_server::ws::WsUpgradeReply::Refuse(resp) => {
return Ok(WsDecision::Respond(resp))
}
crate::courierust_server::ws::WsUpgradeReply::Accept(service) => (service, None),
crate::courierust_server::ws::WsUpgradeReply::AcceptWith { service, protocol } => {
(service, Some(protocol))
}
};
let peer = self
.socket
.peer_addr()
.map(|a| a.ip())
.unwrap_or(core::net::IpAddr::V4(core::net::Ipv4Addr::UNSPECIFIED));
match crate::courierust_server::ws::plan(req, peer, false, &config.websocket) {
Ok(mut plan) => {
let applied = match protocol {
Some(protocol) => plan.override_protocol(protocol.as_deref()),
None => Ok(()),
};
match applied {
Ok(()) => Ok(WsDecision::Upgrade(Box::new(WsUpgrade { plan, service }))),
Err(refusal) => Ok(WsDecision::Respond(refusal.response())),
}
}
Err(refusal) => Ok(WsDecision::Respond(refusal.response())),
}
}
}
fn build_response(
resp: Response<Body>,
request_close: bool,
is_head: bool,
out: &mut Vec<u8>,
) -> Result<(bool, Option<ChannelStream>)> {
let keep_alive = crate::courierust_server::h1::response_wire_head(&resp, request_close, out)?;
match resp.body {
_ if is_head => Ok((keep_alive, None)),
Body::Empty => Ok((keep_alive, None)),
Body::Bytes(b) => {
out.extend_from_slice(&b);
Ok((keep_alive, None))
}
Body::Channel(rx) => Ok((keep_alive, Some(ChannelStream::raw(rx)))),
Body::Stream(stream) => Ok((keep_alive, Some(stream))),
}
}
pub(crate) fn serve_event(
listener: std::net::TcpListener,
handler: Arc<dyn Handler>,
config: ServerConfig,
pool: Arc<crate::courierust_pool::ThreadPool>,
stop: crate::courierust_server::ServerStop,
) -> 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 registries = Registries::new();
let (wake_reader, wake_writer) = wakeup_pair()?;
let wake_writer = Arc::new(wake_writer);
stop.install_reactor_wake(wake_writer.try_clone()?);
stop.install_listener(listener.try_clone()?);
let loop_handler = handler.clone();
let loop_config = config.clone();
let loop_pool = pool.clone();
let loop_registries = registries.clone();
let loop_stop = stop.clone();
let event_thread = thread::Builder::new()
.name("courierust-event".into())
.spawn(move || {
event_loop(
msg_rx,
wake_reader,
LoopContext {
ready_tx,
handler: loop_handler,
config: loop_config,
pool: loop_pool,
registries: loop_registries,
stop: loop_stop,
},
);
})?;
let workers = if config.event_workers == 0 {
std::thread::available_parallelism()
.map(|n| n.get().clamp(1, 8))
.unwrap_or(4)
} else {
config.event_workers
};
let mut worker_handles = Vec::new();
for _ in 0..workers {
let w_registries = registries.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_registries,
&*w_handler,
&w_config,
&w_msg_tx,
&w_wake,
);
})?,
);
}
let a_msg_tx = msg_tx.clone();
let a_wake = wake_writer.clone();
let a_stats = config.stats.clone();
let a_stop = stop.clone();
let accept_thread = thread::Builder::new()
.name("courierust-accept".into())
.spawn(move || {
accept_loop(listener, a_msg_tx, &a_wake, a_stats.as_deref(), &a_stop);
})?;
let _ = accept_thread.join();
let _ = event_thread.join();
for h in worker_handles {
let _ = h.join();
}
Ok(())
}
pub(crate) 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)?;
let _ = reader.set_nodelay(true);
let _ = writer.set_nodelay(true);
Ok((reader, writer))
}
pub(crate) fn wake_nudge(w: &TcpStream) {
let mut s: &TcpStream = w;
let _ = std::io::Write::write(&mut s, &[1]);
}
pub(crate) 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 ws_due_ids(registries: &Registries, now: Instant) -> Vec<usize> {
let ws = registries.ws.lock().unwrap();
ws.iter()
.filter(|(_, conn)| conn.next_deadline().is_some_and(|d| d <= now))
.map(|(&id, _)| id)
.collect()
}
fn h1_body_due_ids(registries: &Registries, now: Instant) -> Vec<usize> {
registries
.h1
.lock()
.unwrap()
.iter()
.filter(|(_, conn)| conn.body_poll_due(now))
.map(|(&id, _)| id)
.collect()
}
fn h1_body_next_deadline(registries: &Registries) -> Option<Instant> {
registries
.h1
.lock()
.unwrap()
.values()
.filter_map(|conn| conn.body_poll_at())
.min()
}
fn wake_connection(id: usize, poller: &mut Poller, ready_tx: &Sender<Vec<usize>>) {
if id == WAKE_ID {
return;
}
poller.unregister(id);
let _ = ready_tx.send(vec![id]);
}
fn ws_next_deadline(registries: &Registries) -> Option<Instant> {
registries
.ws
.lock()
.unwrap()
.values()
.filter_map(|conn| conn.next_deadline())
.min()
}
fn rebuild_wait_set(
poller: &mut Poller,
pending: &HashMap<usize, TcpStream>,
registries: &Registries,
) {
poller.clear();
for (id, stream) in pending.iter() {
poller.register(*id, fd_of(stream), false);
}
for (id, conn) in crate::lock(®istries.h1).iter() {
poller.register(*id, fd_of(&conn.socket), conn.has_pending_output());
}
for (id, conn) in crate::lock(®istries.ws).iter() {
poller.register(*id, fd_of(conn.socket()), conn.has_queued_output());
}
}
struct Reactor<'a> {
poller: &'a mut Poller,
pending: &'a mut HashMap<usize, TcpStream>,
activity: &'a mut HashMap<usize, Instant>,
}
fn handle_msg(
msg: EventMsg,
reactor: Reactor<'_>,
registries: &Registries,
ready_tx: &Sender<Vec<usize>>,
max_connections: usize,
stats: Option<&Stats>,
) {
let Reactor {
poller,
pending,
activity,
} = reactor;
match msg {
EventMsg::BodyChunk { id } => wake_connection(id, poller, ready_tx),
EventMsg::NewConn {
id,
stream,
accepted_at,
} => {
if max_connections > 0 && activity.len() >= max_connections {
drop(stream);
return;
}
if stream.set_nonblocking(true).is_err() {
return;
}
let _ = stream.set_nodelay(true);
if let Some(s) = stats {
s.connections_active.fetch_add(1, Ordering::Relaxed);
}
let fd = fd_of(&stream);
pending.insert(id, stream);
activity.insert(id, Instant::now());
poller.register(id, fd, false);
if let Some(accepted_at) = accepted_at {
eprintln!(
"H1SEG|id={id}|event=newconn|accept_us={}",
accepted_at.elapsed().as_micros()
);
}
}
EventMsg::Register { id, fd, want_write } => {
let want_write = registries
.ws
.lock()
.unwrap()
.get(&id)
.map(|c| c.has_queued_output())
.unwrap_or(want_write);
activity.insert(id, Instant::now());
poller.register(id, fd, want_write);
}
EventMsg::Closed { id, socket } => {
poller.unregister(id);
pending.remove(&id);
if activity.remove(&id).is_some() {
if let Some(s) = stats {
Stats::decrement(&s.connections_active, 1);
}
}
drop(socket);
}
EventMsg::Detach { id, socket } => {
poller.unregister(id);
pending.remove(&id);
drop(socket);
}
}
}
struct LoopContext {
ready_tx: Sender<Vec<usize>>,
handler: Arc<dyn Handler>,
config: ServerConfig,
pool: Arc<crate::courierust_pool::ThreadPool>,
registries: Registries,
stop: crate::courierust_server::ServerStop,
}
fn event_loop(msg_rx: Receiver<EventMsg>, wake_reader: TcpStream, context: LoopContext) {
let LoopContext {
ready_tx,
handler,
config,
pool,
registries,
stop,
} = context;
let mut poller = Poller::new();
let mut pending: HashMap<usize, TcpStream> = HashMap::new();
let mut activity: HashMap<usize, Instant> = HashMap::new();
let stats = config.stats.clone();
let stats = stats.as_deref();
let mut wait_errors = 0usize;
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;
let header_timeout = config.request_header_timeout;
loop {
if stop.is_requested() {
return;
}
let mut drained = 0usize;
loop {
match msg_rx.try_recv() {
Ok(msg) => {
drained += 1;
handle_msg(
msg,
Reactor {
poller: &mut poller,
pending: &mut pending,
activity: &mut activity,
},
®istries,
&ready_tx,
config.max_connections,
stats,
);
}
Err(TryRecvError::Disconnected) => return,
Err(TryRecvError::Empty) => break,
}
}
if drained > 0 {
if let Some(s) = stats {
Stats::bump_peak(&s.event_queue_depth_peak, drained);
}
}
if poller.is_empty() {
match msg_rx.recv_timeout(Duration::from_millis(poll_timeout as u64)) {
Ok(msg) => handle_msg(
msg,
Reactor {
poller: &mut poller,
pending: &mut pending,
activity: &mut activity,
},
®istries,
&ready_tx,
config.max_connections,
stats,
),
Err(RecvTimeoutError::Timeout) => continue,
Err(RecvTimeoutError::Disconnected) => return,
}
continue;
}
let now = Instant::now();
let next_idle = idle_timeout.map(|t| {
activity
.values()
.map(|at| {
t.checked_sub(now.duration_since(*at))
.unwrap_or(Duration::ZERO)
})
.min()
.unwrap_or(Duration::from_secs(3600))
});
let next_header = header_timeout.and_then(|timeout| {
registries
.h1
.lock()
.unwrap()
.values()
.filter_map(|conn| conn.header_timeout_remaining(timeout, now))
.min()
});
let next_deadline = match (next_idle, next_header) {
(Some(idle), Some(header)) => Some(idle.min(header)),
(Some(deadline), None) | (None, Some(deadline)) => Some(deadline),
(None, None) => None,
};
let wait_ms = match next_deadline {
Some(next) => next.as_millis().min(poll_timeout as u128).max(1) as i32,
None => poll_timeout,
};
let timed = [
ws_next_deadline(®istries),
h1_body_next_deadline(®istries),
]
.into_iter()
.flatten()
.min();
let wait_ms = match timed {
Some(deadline) => {
let remaining = deadline.saturating_duration_since(now).as_millis();
wait_ms.min(remaining.min(poll_timeout as u128).max(1) as i32)
}
None => wait_ms,
};
let ready = match poller.wait(wait_ms, Some(wake_fd)) {
Ok(r) => {
wait_errors = 0;
r
}
Err(_) => {
wait_errors += 1;
if let Some(s) = stats {
s.event_wait_errors.fetch_add(1, Ordering::Relaxed);
}
rebuild_wait_set(&mut poller, &pending, ®istries);
if wait_errors >= 64 {
std::thread::sleep(Duration::from_millis(1));
}
continue;
}
};
if let Some(s) = stats {
s.event_poll_syscalls.fetch_add(1, Ordering::Relaxed);
}
if ready.contains(&WAKE_ID) {
if let Some(s) = stats {
s.event_wakeups.fetch_add(1, Ordering::Relaxed);
}
drain_wake(&wake_reader);
let mut drained = 0usize;
loop {
match msg_rx.try_recv() {
Ok(msg) => {
drained += 1;
handle_msg(
msg,
Reactor {
poller: &mut poller,
pending: &mut pending,
activity: &mut activity,
},
®istries,
&ready_tx,
config.max_connections,
stats,
);
}
Err(TryRecvError::Disconnected) => return,
Err(TryRecvError::Empty) => break,
}
}
if drained > 0 {
if let Some(s) = stats {
Stats::bump_peak(&s.event_queue_depth_peak, drained);
}
}
}
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);
if let Some(s) = stats {
Stats::decrement(&s.connections_active, 1);
}
continue;
}
};
if n == 0 {
activity.remove(&id);
if let Some(s) = stats {
Stats::decrement(&s.connections_active, 1);
}
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_connection(stream, &*h, &c);
});
activity.remove(&id);
if let Some(s) = stats {
Stats::decrement(&s.connections_active, 1);
}
}
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::dispatch(
crate::courierust_net::ConnStream::plain(stream),
&*h,
&c,
);
});
activity.remove(&id);
if let Some(s) = stats {
Stats::decrement(&s.connections_active, 1);
}
}
Class::H1 => {
let conn = EventConn::new(
stream,
config.max_body,
config.max_header_list,
stats,
crate::courierust_server::ws::WakeSlot::new(),
);
if let Some(s) = stats {
s.h1_connections.fetch_add(1, Ordering::Relaxed);
}
crate::lock(®istries.h1).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);
if let Some(s) = stats {
Stats::decrement(&s.connections_active, 1);
}
}
}
} else {
to_dispatch.push(id);
}
}
let header_deadline_due = next_header
.map(|next| next <= Duration::from_millis(poll_timeout as u64))
.unwrap_or(false);
let mut header_expired = HashSet::new();
if header_deadline_due {
if let Some(timeout) = header_timeout {
let now = Instant::now();
for (&id, conn) in crate::lock(®istries.h1).iter() {
if conn.header_timeout_remaining(timeout, now) == Some(Duration::ZERO) {
poller.unregister(id);
header_expired.insert(id);
if !to_dispatch.contains(&id) {
to_dispatch.push(id);
}
}
}
}
}
let now = Instant::now();
for id in ws_due_ids(®istries, now) {
poller.unregister(id);
activity.insert(id, now);
to_dispatch.push(id);
}
for id in h1_body_due_ids(®istries, now) {
poller.unregister(id);
activity.insert(id, now);
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 near_idle = next_idle
.map(|next| next <= Duration::from_millis(poll_timeout as u64))
.unwrap_or(false);
if near_idle {
let now = Instant::now();
let mut expired = Vec::new();
let registered: HashSet<usize> = crate::lock(®istries.h1)
.keys()
.chain(crate::lock(®istries.ws).keys())
.copied()
.collect();
for (&id, &at) in &activity {
if header_expired.contains(&id) {
continue;
}
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);
crate::lock(®istries.h1).remove(&id);
crate::lock(®istries.ws).remove(&id);
if activity.remove(&id).is_some() {
if let Some(s) = stats {
Stats::decrement(&s.connections_active, 1);
}
}
}
}
}
}
}
fn refuse_malformed(conn: &EventConn, e: &Error) {
let Some(status) = crate::courierust_server::h1::refusal_status(e) else {
return;
};
let message = if status == 408 {
"request header timeout"
} else {
"bad request"
};
let resp = crate::courierust_server::h1::error_response(status, message);
let body: &[u8] = match &resp.body {
Body::Bytes(b) => b.as_ref(),
_ => b"",
};
let mut out = Vec::with_capacity(128 + body.len());
if courierust_h1::write_response_head(&mut out, resp.status, Version::HTTP_11, &resp.headers)
.is_err()
{
return;
}
out.extend_from_slice(body);
write_and_linger(&conn.socket, &out);
}
fn write_and_linger(socket: &std::net::TcpStream, bytes: &[u8]) {
use std::io::{Read, Write};
let deadline = Instant::now() + crate::courierust_server::h1::LINGER_DEADLINE;
let mut sent = 0usize;
while sent < bytes.len() && Instant::now() < deadline {
let mut writer: &std::net::TcpStream = socket;
match writer.write(&bytes[sent..]) {
Ok(0) => return,
Ok(n) => sent += n,
Err(e) if e.kind() == std::io::ErrorKind::WouldBlock => {
thread::sleep(Duration::from_millis(1));
}
Err(_) => return,
}
}
let mut sink = [0u8; 8 * 1024];
let mut left = crate::courierust_server::h1::LINGER_BUDGET;
while left > 0 && Instant::now() < deadline {
let mut reader: &std::net::TcpStream = socket;
let want = core::cmp::min(left, sink.len());
match reader.read(&mut sink[..want]) {
Ok(0) => break,
Ok(n) => left = left.saturating_sub(n),
Err(e) if e.kind() == std::io::ErrorKind::WouldBlock => {
thread::sleep(Duration::from_millis(1));
}
Err(_) => break,
}
}
}
fn emit_trace(conn: &mut EventConn, id: usize, handoff_us: u64, fresh_wait_us: u64) {
if conn.trace && conn.trace_requests > 0 {
eprintln!(
"H1SEG|id={id}|reqs={}|fresh_wait_us={fresh_wait_us}|handoff_us={handoff_us}|dispatch_us={}|parse_us={}|handler_us={}|build_us={}|write_us={}",
conn.trace_requests,
conn.dispatch_us,
conn.parse_us,
conn.handler_us,
conn.build_us,
conn.write_us,
);
}
if conn.trace {
conn.trace_requests = 0;
conn.dispatch_us = 0;
conn.parse_us = 0;
conn.handler_us = 0;
conn.build_us = 0;
conn.write_us = 0;
}
}
fn event_worker(
ready_rx: Arc<std::sync::Mutex<Receiver<Vec<usize>>>>,
registries: Registries,
handler: &dyn Handler,
config: &ServerConfig,
msg_tx: &Sender<EventMsg>,
wake_writer: &Arc<TcpStream>,
) {
loop {
let ids = match crate::lock(&ready_rx).recv() {
Ok(ids) => ids,
Err(_) => return,
};
for id in ids {
let ws_conn = crate::lock(®istries.ws).remove(&id);
if let Some(mut ws_conn) = ws_conn {
let step =
std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| ws_conn.step()));
let outcome = match step {
Ok(o) => o,
Err(_) => crate::courierust_server::ws::WsStep::Close,
};
match outcome {
crate::courierust_server::ws::WsStep::Idle
| crate::courierust_server::ws::WsStep::NeedWrite => {
let fd = fd_of(ws_conn.socket());
let want_write =
matches!(outcome, crate::courierust_server::ws::WsStep::NeedWrite);
crate::lock(®istries.ws).insert(id, ws_conn);
let _ = msg_tx.send(EventMsg::Register { id, fd, want_write });
wake_nudge(wake_writer);
}
crate::courierust_server::ws::WsStep::Close => {
let socket = ws_conn.socket().clone();
let _ = msg_tx.send(EventMsg::Closed {
id,
socket: Some(socket),
});
wake_nudge(wake_writer);
}
}
continue;
}
let mut conn = match crate::lock(®istries.h1).remove(&id) {
Some(c) => c,
None => continue,
};
conn.arm_wake(|| {
let tx = msg_tx.clone();
let pipe = wake_writer.clone();
Arc::new(move || {
let _ = tx.send(EventMsg::BodyChunk { id });
wake_nudge(&pipe);
})
});
let (handoff_us, fresh_wait_us) = if conn.trace {
let pickup_at = Instant::now();
let handoff = conn
.parked_at
.take()
.map(|at| at.elapsed().as_micros() as u64)
.unwrap_or(0);
let fresh = conn
.registered_at
.take()
.map(|at| pickup_at.duration_since(at).as_micros() as u64)
.unwrap_or(0);
conn.pickup_at = Some(pickup_at);
(handoff, fresh)
} else {
(0, 0)
};
emit_trace(&mut conn, id, handoff_us, fresh_wait_us);
let step = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
conn.step(handler, config)
}));
let outcome = match step {
Ok(Ok(o)) => o,
Ok(Err(e)) => {
refuse_malformed(&conn, &e);
StepOutcome::Close
}
Err(_) => StepOutcome::Close,
};
match outcome {
StepOutcome::Idle | StepOutcome::NeedWrite => {
let fd = fd_of(&conn.socket);
let want_write = matches!(outcome, StepOutcome::NeedWrite);
if conn.trace {
conn.parked_at = Some(Instant::now());
}
crate::lock(®istries.h1).insert(id, conn);
let _ = msg_tx.send(EventMsg::Register { id, fd, want_write });
wake_nudge(wake_writer);
}
StepOutcome::Close => {
emit_trace(&mut conn, id, 0, 0);
let socket = conn.socket.clone();
let _ = msg_tx.send(EventMsg::Closed {
id,
socket: Some(socket),
});
wake_nudge(wake_writer);
}
StepOutcome::Tunnel(job) => {
emit_trace(&mut conn, id, 0, 0);
let TunnelJob {
plan,
socket,
leftover,
secure,
} = *job;
let _ = msg_tx.send(EventMsg::Detach {
id,
socket: Some(conn.socket.clone()),
});
wake_nudge(wake_writer);
let tx = msg_tx.clone();
let wake = wake_writer.clone();
let spawned = std::thread::Builder::new()
.name("courierust-tunnel".into())
.spawn(move || {
run_tunnel(socket, leftover, secure, plan);
let _ = tx.send(EventMsg::Closed { id, socket: None });
wake_nudge(&wake);
});
if spawned.is_err() {
let _ = msg_tx.send(EventMsg::Closed { id, socket: None });
wake_nudge(wake_writer);
}
}
StepOutcome::Upgrade(upgraded) => {
let mut ws_conn = *upgraded;
let fd = fd_of(ws_conn.socket());
let tx = msg_tx.clone();
let wake = wake_writer.clone();
ws_conn.set_wake(Arc::new(move || {
let _ = tx.send(EventMsg::Register {
id,
fd,
want_write: true,
});
wake_nudge(&wake);
}));
let step =
std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| ws_conn.step()));
let outcome = match step {
Ok(o) => o,
Err(_) => crate::courierust_server::ws::WsStep::Close,
};
match outcome {
crate::courierust_server::ws::WsStep::Idle
| crate::courierust_server::ws::WsStep::NeedWrite => {
let want_write =
matches!(outcome, crate::courierust_server::ws::WsStep::NeedWrite);
registries.ws.lock().unwrap().insert(id, ws_conn);
let _ = msg_tx.send(EventMsg::Register { id, fd, want_write });
wake_nudge(wake_writer);
}
crate::courierust_server::ws::WsStep::Close => {
let socket = ws_conn.socket().clone();
let _ = msg_tx.send(EventMsg::Closed {
id,
socket: Some(socket),
});
wake_nudge(wake_writer);
}
}
}
}
}
}
}
fn run_tunnel(
socket: Arc<TcpStream>,
leftover: Vec<u8>,
secure: bool,
plan: crate::courierust_server::tunnel::TunnelPlan,
) {
let _ = socket.set_nonblocking(false);
let stream = Arc::new(crate::courierust_net::ConnStream::plain_shared(socket));
let mut reader =
crate::courierust_io::BufReader::new(stream.clone(), leftover.len().max(16 * 1024));
if !leftover.is_empty() {
reader.seed(&leftover);
}
let conn = crate::courierust_server::TunnelConn::new(stream, reader, secure);
plan.service.run(conn);
}
fn accept_loop(
listener: std::net::TcpListener,
msg_tx: Sender<EventMsg>,
wake_writer: &Arc<TcpStream>,
stats: Option<&Stats>,
stop: &crate::courierust_server::ServerStop,
) {
let mut next_id = 1usize;
for stream in listener.incoming() {
if stop.is_requested() {
return;
}
let Ok(stream) = stream else { continue };
if let Some(s) = stats {
s.connections_accepted.fetch_add(1, Ordering::Relaxed);
}
let id = next_id;
next_id += 1;
let accepted_at = if h1_trace() {
Some(Instant::now())
} else {
None
};
let _ = msg_tx.send(EventMsg::NewConn {
id,
stream,
accepted_at,
});
wake_nudge(wake_writer);
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::courierust_error::ErrorKind;
use std::sync::atomic::AtomicUsize;
fn conn_pair() -> (EventConn, TcpStream) {
let listener = std::net::TcpListener::bind("127.0.0.1:0").unwrap();
let addr = listener.local_addr().unwrap();
let peer = TcpStream::connect(addr).unwrap();
let (socket, _) = listener.accept().unwrap();
socket.set_nonblocking(true).unwrap();
let conn = EventConn::new(
socket,
1024 * 1024,
1 << 20,
None,
crate::courierust_server::ws::WakeSlot::new(),
);
(conn, peer)
}
fn test_config() -> ServerConfig {
ServerConfig {
read_timeout: Some(Duration::from_millis(50)),
..ServerConfig::default()
}
}
#[test]
fn streaming_body_fires_the_connection_wake() {
let (mut conn, _peer) = conn_pair();
let hits = Arc::new(AtomicUsize::new(0));
let armed = hits.clone();
conn.arm_wake(move || {
let armed = armed.clone();
Arc::new(move || {
armed.fetch_add(1, Ordering::Relaxed);
})
});
let (tx, body) = crate::courierust_body::channel();
conn.set_stream(body.into_stream());
assert!(conn.stream_rx.is_some());
tx.send(Bytes::from_static(b"chunk")).unwrap();
assert_eq!(
hits.load(Ordering::Relaxed),
1,
"a produced chunk must wake the connection"
);
tx.send_bytes(b"more").unwrap();
assert_eq!(hits.load(Ordering::Relaxed), 2);
let (_raw_tx, raw) = std::sync::mpsc::channel();
conn.set_stream(Some(ChannelStream::raw(raw)));
assert!(
!conn.stream_rx.as_ref().unwrap().stream.has_wake(),
"a raw channel must not claim a wake"
);
assert_eq!(
body_poll_delay(conn.stream_rx.as_ref().unwrap()),
Duration::from_millis(BODY_POLL_MIN_MS),
"a raw channel is polled, not waited on"
);
}
#[test]
fn a_parked_stream_does_not_block_the_worker() {
let (mut conn, mut peer) = conn_pair();
peer.set_read_timeout(Some(Duration::from_secs(5))).unwrap();
let config = test_config();
let (tx, body) = crate::courierust_body::channel();
conn.set_stream(body.into_stream());
assert!(matches!(
conn.pump_body(&config).unwrap(),
StepOutcome::Idle
));
assert!(conn.stream_rx.is_some(), "the receiver must survive a park");
let due = conn.body_poll_at().expect("a parked body owes a deadline");
assert!(due > Instant::now() && due <= Instant::now() + Duration::from_secs(1));
tx.send(Bytes::from_static(b"hello")).unwrap();
assert!(matches!(
conn.pump_body(&config).unwrap(),
StepOutcome::Idle
));
drop(tx);
assert!(matches!(
conn.pump_body(&config).unwrap(),
StepOutcome::Idle
));
assert!(conn.stream_rx.is_none(), "a finished body is released");
drop(conn);
let mut written = Vec::new();
std::io::Read::read_to_end(&mut peer, &mut written).unwrap();
assert_eq!(written, b"5\r\nhello\r\n0\r\n\r\n");
}
#[test]
fn a_stalled_stream_times_out_on_its_poll_deadline() {
let (mut conn, _peer) = conn_pair();
let config = ServerConfig {
read_timeout: Some(Duration::from_millis(20)),
..ServerConfig::default()
};
let (tx, body) = crate::courierust_body::channel();
conn.set_stream(body.into_stream());
assert!(matches!(
conn.pump_body(&config).unwrap(),
StepOutcome::Idle
));
std::thread::sleep(Duration::from_millis(30));
let error = match conn.pump_body(&config) {
Err(error) => error,
Ok(_) => panic!("a stalled producer must time out"),
};
assert_eq!(error.kind, ErrorKind::Timeout);
drop(tx);
}
#[test]
fn poll_backoff_only_applies_to_a_body_without_a_wake() {
let (_tx, raw) = std::sync::mpsc::channel();
let mut body = BodyRx::new(ChannelStream::raw(raw));
assert_eq!(
body_poll_delay(&body),
Duration::from_millis(BODY_POLL_MIN_MS)
);
body.empty_polls = 4;
assert_eq!(
body_poll_delay(&body),
Duration::from_millis(BODY_POLL_MAX_MS)
);
body.empty_polls = 100;
assert_eq!(
body_poll_delay(&body),
Duration::from_millis(BODY_POLL_MAX_MS),
"the backoff is capped"
);
let (_tx, body) = crate::courierust_body::channel();
let stream = body.into_stream().unwrap();
let mut body = BodyRx::new(stream);
body.empty_polls = 100;
assert_eq!(
body_poll_delay(&body),
Duration::from_millis(BODY_POLL_MAX_MS),
"an unadopted stream polls on the same backoff as a raw channel"
);
let (tx, body) = crate::courierust_body::channel();
let stream = body.into_stream().unwrap();
stream.install_wake(|| {});
let mut stream = BodyRx::new(stream);
stream.empty_polls = 100;
assert_eq!(
body_poll_delay(&stream),
Duration::from_millis(BODY_POLL_WAKE_MS),
"a wake-capable body needs no fast poll"
);
drop(tx);
}
}