pub mod h1;
pub mod h2;
pub mod reverse_proxy;
pub mod tunnel;
pub mod ws;
pub(crate) mod event;
pub use tunnel::{TunnelConn, TunnelPlan, TunnelReply, TunnelService};
use crate::courierust_body::Body;
use crate::courierust_http::request::Request;
use crate::courierust_http::response::Response;
use crate::courierust_net::stats::Stats;
use crate::courierust_pool::ThreadPool;
use std::net::{SocketAddr, TcpListener, TcpStream, ToSocketAddrs, UdpSocket};
use std::sync::Arc;
use std::time::Duration;
#[derive(Debug, Clone)]
pub struct TlsSettings {
pub identity: crate::courierust_tls::Identity,
pub alpn: Vec<Vec<u8>>,
pub min_version: crate::courierust_tls::TlsVersion,
pub max_version: crate::courierust_tls::TlsVersion,
pub session_ticket_key: [u8; 32],
pub client_auth: Option<crate::courierust_tls::ClientAuth>,
}
impl Default for TlsSettings {
fn default() -> Self {
let mut ticket_key = [0u8; 32];
let _ = crate::courierust_tls::crypto::rng::fill_random(&mut ticket_key);
Self {
identity: crate::courierust_tls::Identity::empty(),
alpn: Vec::new(),
min_version: crate::courierust_tls::TlsVersion::Tls12,
max_version: crate::courierust_tls::TlsVersion::Tls13,
session_ticket_key: ticket_key,
client_auth: None,
}
}
}
impl TlsSettings {
pub fn from_pem(cert_pem: &str, key_pem: &str) -> crate::courierust_tls::TlsResult<Self> {
Ok(Self {
identity: crate::courierust_tls::Identity::from_pem(cert_pem, key_pem)?,
alpn: default_alpn(),
..Self::default()
})
}
pub fn from_pem_file(
cert_path: impl AsRef<std::path::Path>,
key_path: impl AsRef<std::path::Path>,
) -> crate::courierust_tls::TlsResult<Self> {
Ok(Self {
identity: crate::courierust_tls::Identity::from_pem_file(cert_path, key_path)?,
alpn: default_alpn(),
..Self::default()
})
}
}
fn default_alpn() -> Vec<Vec<u8>> {
vec![b"h2".to_vec(), b"http/1.1".to_vec()]
}
#[derive(Debug, Clone)]
pub struct ServerConfig {
pub read_timeout: Option<Duration>,
pub request_header_timeout: Option<Duration>,
pub max_header_list: usize,
pub max_body: usize,
pub http2: bool,
pub http3: bool,
pub threads: usize,
pub tls: Option<TlsSettings>,
pub event_driven: bool,
pub event_workers: usize,
pub event_poll_timeout_ms: u64,
pub max_connections: usize,
pub handshake_timeout: Option<Duration>,
pub idle_timeout: Option<Duration>,
pub h2_settings_timeout: Option<Duration>,
pub h2_ping_interval: Option<Duration>,
pub h2_ping_timeout: Option<Duration>,
pub h2_idle_timeout: Option<Duration>,
pub h2_max_concurrent_streams: u32,
pub auto_release_credit: bool,
pub stats: Option<Arc<crate::courierust_net::stats::Stats>>,
pub websocket: ws::WsConfig,
}
impl Default for ServerConfig {
fn default() -> Self {
Self {
read_timeout: Some(Duration::from_secs(120)),
request_header_timeout: Some(Duration::from_secs(15)),
max_header_list: 1 << 20,
max_body: 16 * 1024 * 1024,
http2: true,
http3: false,
threads: 0, tls: None,
event_driven: true,
event_workers: 0,
event_poll_timeout_ms: 50,
max_connections: 1024,
handshake_timeout: Some(Duration::from_secs(10)),
idle_timeout: Some(Duration::from_secs(300)),
h2_settings_timeout: Some(Duration::from_secs(10)),
h2_ping_interval: Some(Duration::from_secs(30)),
h2_ping_timeout: Some(Duration::from_secs(15)),
h2_idle_timeout: Some(Duration::from_secs(300)),
h2_max_concurrent_streams: 1024,
auto_release_credit: true,
stats: None,
websocket: ws::WsConfig::default(),
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct ConnectionInfo {
pub peer: std::net::SocketAddr,
pub secure: bool,
}
pub trait Handler: Send + Sync + 'static {
fn handle(&self, req: Request<Body>) -> Response<Body>;
fn handle_connected(&self, _info: &ConnectionInfo, req: Request<Body>) -> Response<Body> {
self.handle(req)
}
fn websocket(&self, _req: &Request<Body>) -> ws::WsUpgradeReply {
ws::WsUpgradeReply::Pass
}
fn tunnel(&self, _req: &Request<Body>) -> TunnelReply {
TunnelReply::Pass
}
}
impl<F> Handler for F
where
F: Fn(Request<Body>) -> Response<Body> + Send + Sync + 'static,
{
fn handle(&self, req: Request<Body>) -> Response<Body> {
self(req)
}
}
const H3_PORT_ATTEMPTS: usize = 16;
fn bind_listeners(
addr: impl ToSocketAddrs,
http3: bool,
) -> std::io::Result<(TcpListener, Option<UdpSocket>)> {
let addrs: Vec<SocketAddr> = addr.to_socket_addrs()?.collect();
let ephemeral = addrs.iter().any(|a| a.port() == 0);
let attempts = if http3 && ephemeral {
H3_PORT_ATTEMPTS
} else {
1
};
let mut last: Option<std::io::Error> = None;
for _ in 0..attempts {
for resolved in &addrs {
match bind_pair(*resolved, http3) {
Ok(pair) => return Ok(pair),
Err(error) => last = Some(error),
}
}
}
Err(last.unwrap_or_else(|| {
std::io::Error::new(std::io::ErrorKind::AddrNotAvailable, "no address to bind")
}))
}
fn bind_pair(addr: SocketAddr, http3: bool) -> std::io::Result<(TcpListener, Option<UdpSocket>)> {
if !http3 {
return Ok((TcpListener::bind(addr)?, None));
}
let udp = crate::courierust_net::udp::bind_udp(addr)?;
let mut tcp_addr = addr;
tcp_addr.set_port(udp.local_addr()?.port());
let listener = TcpListener::bind(tcp_addr)?;
Ok((listener, Some(udp)))
}
pub struct Server {
listener: TcpListener,
pool: Arc<ThreadPool>,
config: ServerConfig,
h3_socket: Option<UdpSocket>,
}
impl Server {
pub fn bind(addr: impl std::net::ToSocketAddrs) -> std::io::Result<Self> {
Self::bind_with_config(addr, ServerConfig::default())
}
pub fn bind_with_config(
addr: impl std::net::ToSocketAddrs,
config: ServerConfig,
) -> std::io::Result<Self> {
let (listener, h3_socket) = bind_listeners(addr, config.http3)?;
Self::adopt(listener, h3_socket, config)
}
pub fn from_listener(listener: TcpListener, config: ServerConfig) -> std::io::Result<Self> {
let h3_socket = if config.http3 {
let addr = listener.local_addr()?;
Some(crate::courierust_net::udp::bind_udp(addr)?)
} else {
None
};
Self::adopt(listener, h3_socket, config)
}
fn adopt(
listener: TcpListener,
h3_socket: Option<UdpSocket>,
config: ServerConfig,
) -> std::io::Result<Self> {
if let Some(message) = identity_error(&config) {
return Err(std::io::Error::new(
std::io::ErrorKind::InvalidInput,
message,
));
}
let threads = if config.threads == 0 {
recommended_workers()
} else {
config.threads
};
Ok(Self {
listener,
pool: pool_for(threads),
config,
h3_socket,
})
}
pub fn local_addr(&self) -> std::io::Result<SocketAddr> {
self.listener.local_addr()
}
pub fn serve<H: Handler>(self, handler: H) -> std::io::Result<()> {
self.serve_with_config(handler)
}
pub fn serve_with_config<H: Handler>(self, handler: H) -> std::io::Result<()> {
self.serve_inner(handler, None, ServerStop::new())
}
pub fn serve_with_stop<H: Handler>(self, handler: H, stop: ServerStop) -> std::io::Result<()> {
self.serve_inner(handler, None, stop)
}
fn serve_inner<H: Handler>(
self,
handler: H,
ready: Option<&std::sync::mpsc::Sender<std::io::Result<()>>>,
stop: ServerStop,
) -> std::io::Result<()> {
let handler = Arc::new(handler);
let config = self.config;
let pool = self.pool;
let h3_socket = self.h3_socket;
let setup: std::io::Result<Option<_>> = (|| {
if !config.http3 {
return Ok(None);
}
let tls = config.tls.as_ref().ok_or_else(|| {
std::io::Error::new(
std::io::ErrorKind::InvalidInput,
"ServerConfig.http3 requires a TLS identity",
)
})?;
if tls.client_auth.is_some() {
return Err(std::io::Error::new(
std::io::ErrorKind::InvalidInput,
"ServerConfig.http3 cannot enforce client authentication \
(mTLS is implemented for TLS 1.3 over TCP)",
));
}
let socket = h3_socket.ok_or_else(|| {
std::io::Error::new(
std::io::ErrorKind::InvalidInput,
"HTTP/3 was enabled after bind; rebuild the server with ServerConfig.http3",
)
})?;
Ok(Some(
crate::courierust_h3::runtime::spawn_server_with_socket(
socket,
tls,
handler.clone(),
config.clone(),
)?,
))
})();
if let Some(ready) = ready {
let _ = ready.send(match &setup {
Ok(_) => Ok(()),
Err(error) => Err(std::io::Error::new(error.kind(), error.to_string())),
});
}
let _http3 = setup?;
if config.event_driven {
return event::serve_event(self.listener, handler, config, pool, stop);
}
let active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
stop.install_listener(self.listener.try_clone()?);
for stream in self.listener.incoming() {
if stop.is_requested() {
break;
}
match stream {
Ok(stream) => {
if let Some(s) = config.stats.as_deref() {
s.connections_accepted
.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
}
if !try_reserve(&active, config.max_connections) {
drop(stream);
continue;
}
let h = handler.clone();
let c = config.clone();
let p = pool.clone();
let active = active.clone();
if let Some(s) = config.stats.as_deref() {
s.connections_active
.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
}
let permit = ConnectionPermit {
active: active.clone(),
stats: config.stats.clone(),
};
p.spawn(move || {
let _permit = permit;
let _ = serve_connection(stream, h.as_ref(), &c);
});
}
Err(_) => continue,
}
}
Ok(())
}
pub fn serve_background<H: Handler>(self, handler: H) -> std::io::Result<ServerHandle> {
let (ready_tx, ready_rx) = std::sync::mpsc::channel();
let (tx, rx) = std::sync::mpsc::channel();
let stop = ServerStop::new();
let thread_stop = stop.clone();
std::thread::Builder::new()
.name("courierust-server".into())
.spawn(move || {
let res = self.serve_inner(handler, Some(&ready_tx), thread_stop);
let _ = tx.send(res);
})?;
ready_rx.recv().unwrap_or(Ok(()))?;
Ok(ServerHandle { done: rx, stop })
}
}
fn recommended_workers() -> usize {
std::thread::available_parallelism()
.map(|n| n.get().clamp(1, 8))
.unwrap_or(4)
}
fn pool_for(threads: usize) -> Arc<ThreadPool> {
Arc::new(
ThreadPool::with_size(threads).unwrap_or_else(|_| ThreadPool::with_size(2).expect("pool")),
)
}
fn identity_error(config: &ServerConfig) -> Option<&'static str> {
match config.tls.as_ref() {
Some(tls) if tls.identity.is_empty() => Some(
"TLS is enabled but the identity is empty: load a certificate/key pair with \
Identity::from_pem_file (or Identity::from_pem)",
),
_ => None,
}
}
fn try_reserve(active: &std::sync::atomic::AtomicUsize, limit: usize) -> bool {
if limit == 0 {
active.fetch_add(1, std::sync::atomic::Ordering::AcqRel);
return true;
}
let mut current = active.load(std::sync::atomic::Ordering::Acquire);
loop {
if current >= limit {
return false;
}
match active.compare_exchange_weak(
current,
current + 1,
std::sync::atomic::Ordering::AcqRel,
std::sync::atomic::Ordering::Acquire,
) {
Ok(_) => return true,
Err(observed) => current = observed,
}
}
}
struct ConnectionPermit {
active: Arc<std::sync::atomic::AtomicUsize>,
stats: Option<Arc<Stats>>,
}
impl Drop for ConnectionPermit {
fn drop(&mut self) {
Stats::decrement(&self.active, 1);
if let Some(stats) = self.stats.as_deref() {
Stats::decrement(&stats.connections_active, 1);
}
}
}
#[derive(Clone)]
pub struct ServerStop {
requested: Arc<std::sync::atomic::AtomicBool>,
reactor_wake: Arc<std::sync::Mutex<Option<TcpStream>>>,
listener: Arc<std::sync::Mutex<Option<TcpListener>>>,
}
impl ServerStop {
pub fn new() -> Self {
Self {
requested: Arc::new(std::sync::atomic::AtomicBool::new(false)),
reactor_wake: Arc::new(std::sync::Mutex::new(None)),
listener: Arc::new(std::sync::Mutex::new(None)),
}
}
pub fn request(&self) {
if !self
.requested
.swap(true, std::sync::atomic::Ordering::AcqRel)
{
self.wake();
}
}
pub fn is_requested(&self) -> bool {
self.requested.load(std::sync::atomic::Ordering::Relaxed)
}
fn wake(&self) {
if let Ok(guard) = self.reactor_wake.lock() {
if let Some(writer) = guard.as_ref() {
event::wake_nudge(writer);
}
}
if let Ok(guard) = self.listener.lock() {
if let Some(listener) = guard.as_ref() {
if let Ok(addr) = listener.local_addr() {
let target = match addr.ip() {
std::net::IpAddr::V4(ip) if ip.is_unspecified() => {
SocketAddr::from(([127, 0, 0, 1], addr.port()))
}
std::net::IpAddr::V6(ip) if ip.is_unspecified() => {
SocketAddr::from((std::net::Ipv6Addr::LOCALHOST, addr.port()))
}
_ => addr,
};
let _ = TcpStream::connect_timeout(&target, Duration::from_millis(200));
}
}
}
}
pub(crate) fn install_reactor_wake(&self, writer: TcpStream) {
if let Ok(mut guard) = self.reactor_wake.lock() {
*guard = Some(writer);
}
}
pub(crate) fn install_listener(&self, listener: TcpListener) {
if let Ok(mut guard) = self.listener.lock() {
*guard = Some(listener);
}
}
}
impl Default for ServerStop {
fn default() -> Self {
Self::new()
}
}
pub struct ServerHandle {
done: std::sync::mpsc::Receiver<std::io::Result<()>>,
stop: ServerStop,
}
impl ServerHandle {
pub fn stop(&self) {
self.stop.request();
}
pub fn is_stopping(&self) -> bool {
self.stop.is_requested()
}
pub fn stop_signal(&self) -> ServerStop {
self.stop.clone()
}
pub fn join(self) -> std::io::Result<()> {
self.done.recv().unwrap_or(Ok(()))
}
}
impl std::fmt::Debug for ServerHandle {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("ServerHandle")
.field("stopping", &self.is_stopping())
.finish_non_exhaustive()
}
}
pub fn serve_connection(
stream: TcpStream,
handler: &dyn Handler,
config: &ServerConfig,
) -> crate::Result<()> {
if let Some(message) = identity_error(config) {
return Err(crate::courierust_error::Error::with_message(
crate::courierust_error::ErrorKind::Other,
message,
));
}
if config.http3 {
return Err(crate::courierust_error::Error::with_message(
crate::courierust_error::ErrorKind::Other,
"HTTP/3 is served by the server's QUIC reactor (Server::serve_background); \
serve_connection drives TCP only",
));
}
if config.tls.is_some() {
crate::courierust_net::configure(&stream, config.handshake_timeout)?;
} else {
crate::courierust_net::configure(&stream, config.read_timeout)?;
}
match &config.tls {
Some(t) => {
let acceptor =
crate::courierust_tls::TlsAcceptor::new(crate::courierust_tls::ServerConfig {
identity: t.identity.clone(),
alpn: t.alpn.clone(),
min_version: t.min_version,
max_version: t.max_version,
session_ticket_key: Some(t.session_ticket_key),
client_auth: t.client_auth.clone(),
});
let arc = Arc::new(stream);
let peer = arc
.peer_addr()
.unwrap_or_else(|_| SocketAddr::from(([0, 0, 0, 0], 0)));
let tls = acceptor.accept(arc.clone(), arc.clone()).map_err(|e| {
crate::courierust_error::Error::with_message(
crate::courierust_error::ErrorKind::Other,
e.to_string(),
)
})?;
let conn = crate::courierust_net::ConnStream::tls_server(tls, peer);
let _ = conn.configure(config.read_timeout);
dispatch(conn, handler, config)
}
None => dispatch(
crate::courierust_net::ConnStream::plain(stream),
handler,
config,
),
}
}
pub(crate) fn dispatch(
stream: crate::courierust_net::ConnStream,
handler: &dyn Handler,
config: &ServerConfig,
) -> crate::Result<()> {
let stream = Arc::new(stream);
if let Some(alpn) = stream.alpn() {
if config.http2 && alpn == b"h2" {
return h2::serve(&stream, handler, config);
}
return h1::serve(&stream, handler, config);
}
let mut prefix = [0u8; 24];
let n = stream.peek(&mut prefix).unwrap_or(0);
if config.http2 && n == 24 && crate::courierust_h2::connection::is_preface(&prefix) {
h2::serve(&stream, handler, config)
} else {
h1::serve(&stream, handler, config)
}
}