use std::convert::Infallible;
use std::fs;
use std::net::SocketAddr;
use std::path::{Path, PathBuf};
use std::pin::Pin;
use std::sync::Arc;
use std::task::{Context, Poll};
use std::time::{Duration, SystemTime};
use bytes::Bytes;
use http_body_util::Full;
use hyper::body::Incoming;
use hyper::http::response::Builder;
use hyper::service::service_fn;
use hyper::{HeaderMap, Method, Request, Response, StatusCode};
use hyper_util::rt::TokioExecutor;
use hyper_util::rt::TokioIo;
use hyper_util::server::conn::auto::Builder as AutoBuilder;
use tokio::fs::File;
use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite, ReadBuf};
use tokio::net::{TcpListener, TcpStream};
use tokio::sync::{OwnedSemaphorePermit, Semaphore};
use tokio::time::timeout;
use crate::css::{CssOptions, CssTool};
use crate::error::StaticError;
use crate::handler::{FileBody, ResponseBody};
use crate::js::{JsOptions, JsTool};
use crate::reload::{self, SseBody};
use crate::resolve;
use crate::source::SourcePipeline;
use crate::tool;
use crate::watcher::{start_watching, Broadcaster};
const ACCEPT_BACKOFF_INITIAL: Duration = Duration::from_millis(10);
const ACCEPT_BACKOFF_MAX: Duration = Duration::from_secs(1);
const DEFAULT_MAX_CONNECTIONS: usize = 1024;
trait TcpAccept {
async fn accept(&self) -> std::io::Result<(TcpStream, SocketAddr)>;
}
impl TcpAccept for TcpListener {
async fn accept(&self) -> std::io::Result<(TcpStream, SocketAddr)> {
TcpListener::accept(self).await
}
}
async fn accept_and_permit<L: TcpAccept>(
listener: &L,
backoff: &mut Duration,
semaphore: &Arc<Semaphore>,
) -> Option<(TcpStream, OwnedSemaphorePermit)> {
loop {
let stream = match listener.accept().await {
Ok((stream, _)) => {
*backoff = ACCEPT_BACKOFF_INITIAL;
stream
}
Err(_) => {
tokio::time::sleep(*backoff).await;
*backoff = (*backoff * 2).min(ACCEPT_BACKOFF_MAX);
continue;
}
};
return semaphore
.clone()
.acquire_owned()
.await
.ok()
.map(|permit| (stream, permit));
}
}
type ImmutablePredicate = Arc<dyn Fn(&Path) -> bool + Send + Sync>;
#[derive(Clone)]
pub struct Server {
root_canon: PathBuf,
bundle_roots: Vec<PathBuf>,
max_connections: usize,
live_reload: bool,
broadcaster: Option<Broadcaster>,
immutable_predicate: Option<ImmutablePredicate>,
source_folders: Vec<PathBuf>,
output_dir: PathBuf,
css_tool: Option<(CssTool, CssOptions)>,
js_tool: Option<(JsTool, JsOptions)>,
prune_output: bool,
}
fn paths_overlap(a: &Path, b: &Path) -> bool {
a.starts_with(b) || b.starts_with(a)
}
impl Server {
pub fn new(root: &Path) -> Result<Self, StaticError> {
let root_canon = root.canonicalize().map_err(StaticError::Io)?;
let output_dir = root_canon.clone();
Ok(Server {
root_canon,
bundle_roots: Vec::new(),
max_connections: DEFAULT_MAX_CONNECTIONS,
live_reload: false,
broadcaster: None,
immutable_predicate: None,
source_folders: Vec::new(),
output_dir,
css_tool: None,
js_tool: None,
prune_output: false,
})
}
pub fn with_max_connections(mut self, max: usize) -> Self {
self.max_connections = max;
self
}
pub fn with_live_reload(mut self) -> Self {
self.live_reload = true;
self
}
pub fn with_immutable_assets<F>(mut self, predicate: F) -> Self
where
F: Fn(&Path) -> bool + Send + Sync + 'static,
{
self.immutable_predicate = Some(Arc::new(predicate));
self
}
fn cache_control_for(&self, path: &Path) -> &'static str {
match &self.immutable_predicate {
Some(predicate) if predicate(path) => "public, max-age=31536000, immutable",
_ => "no-cache",
}
}
pub fn with_bundle_root(mut self, path: &Path) -> Result<Self, StaticError> {
let canon = path.canonicalize().map_err(StaticError::Io)?;
self.bundle_roots.push(canon);
Ok(self)
}
pub fn with_source_folder(mut self, dir: &Path) -> Result<Self, StaticError> {
let canon = dir.canonicalize().map_err(StaticError::Io)?;
if paths_overlap(&canon, &self.output_dir) {
return Err(StaticError::Traversal(format!(
"source folder {} overlaps the output dir {}",
canon.display(),
self.output_dir.display()
)));
}
if self
.source_folders
.iter()
.any(|existing| paths_overlap(&canon, existing))
{
return Err(StaticError::Traversal(format!(
"source folder {} overlaps an already-registered source folder",
canon.display()
)));
}
self.source_folders.push(canon);
Ok(self)
}
pub fn with_output_dir(mut self, dir: &Path) -> Result<Self, StaticError> {
let canon = dir.canonicalize().map_err(StaticError::Io)?;
if self
.source_folders
.iter()
.any(|existing| paths_overlap(&canon, existing))
{
return Err(StaticError::Traversal(format!(
"output dir {} overlaps a registered source folder",
canon.display()
)));
}
self.output_dir = canon;
Ok(self)
}
pub fn with_css_tool(mut self, tool: CssTool, options: CssOptions) -> Self {
self.css_tool = Some((tool, options));
self
}
pub fn with_js_tool(mut self, tool: JsTool, options: JsOptions) -> Result<Self, StaticError> {
if let Some(entry) = options.entry() {
let entry_canon = entry.canonicalize().map_err(StaticError::Io)?;
let under_source_folder = self
.source_folders
.iter()
.any(|folder| entry_canon.starts_with(folder));
if !under_source_folder {
return Err(StaticError::Traversal(format!(
"js bundle entry {} is not under any registered source folder",
entry_canon.display()
)));
}
}
self.js_tool = Some((tool, options));
Ok(self)
}
pub fn with_prune_output(mut self) -> Self {
self.prune_output = true;
self
}
fn has_pipeline(&self) -> bool {
self.css_tool.is_some() || self.js_tool.is_some() || !self.source_folders.is_empty()
}
fn required_tool_binaries(&self) -> Vec<(&'static str, &'static str)> {
let mut required = Vec::new();
if let Some((css_tool, options)) = &self.css_tool {
if options.is_bundle() || options.is_minify() {
required.push((css_tool.binary_name(), css_tool.install_hint()));
}
}
if let Some((js_tool, options)) = &self.js_tool {
if options.is_bundle() || options.is_minify() {
required.push((js_tool.binary_name(), js_tool.install_hint()));
}
}
required
}
fn watch_targets(&self) -> Vec<PathBuf> {
let mut targets = Vec::new();
for dir in self.source_folders.iter().chain(self.bundle_roots.iter()) {
if !targets.contains(dir) {
targets.push(dir.clone());
}
}
targets
}
pub fn resolve(&self, request_path: &str) -> Result<PathBuf, StaticError> {
resolve::resolve_with_canonical_root(&self.root_canon, request_path)
}
pub async fn run_on(
&self,
addr: SocketAddr,
header_timeout: Duration,
) -> Result<(u16, ServerHandle), StaticError> {
for (binary, install_hint) in self.required_tool_binaries() {
if !tool::locate_on_path(binary) {
return Err(StaticError::PipelineSetup(format!(
"{binary} not found on PATH ({install_hint})"
)));
}
}
let listener = TcpListener::bind(addr).await.map_err(StaticError::Io)?;
let port = listener.local_addr().map_err(StaticError::Io)?.port();
let mut server = self.clone();
if server.live_reload {
let broadcaster = Broadcaster::new();
if server.has_pipeline() {
let pipeline = Arc::new(SourcePipeline::new(
server.source_folders.clone(),
server.bundle_roots.clone(),
server.output_dir.clone(),
server.css_tool.clone(),
server.js_tool.clone(),
server.prune_output,
broadcaster.clone(),
));
let mut rx = broadcaster.subscribe();
tokio::spawn(async move {
if let Err(e) = pipeline.full_build().await {
eprintln!("source pipeline build error: {e}");
}
while let Some(event) = rx.recv().await {
if let Err(e) = pipeline
.process_change(&event.path, &event.change_type)
.await
{
eprintln!("source pipeline error: {e}");
}
}
});
}
for dir in server.watch_targets() {
start_watching(Arc::new(dir), broadcaster.clone());
}
server.broadcaster = Some(broadcaster);
} else if server.has_pipeline() {
let pipeline = Arc::new(SourcePipeline::new(
server.source_folders.clone(),
server.bundle_roots.clone(),
server.output_dir.clone(),
server.css_tool.clone(),
server.js_tool.clone(),
server.prune_output,
Broadcaster::new(),
));
tokio::spawn(async move {
if let Err(e) = pipeline.full_build().await {
eprintln!("source pipeline build error: {e}");
}
});
}
let semaphore = Arc::new(Semaphore::new(server.max_connections));
let (shutdown_tx, shutdown_rx) = tokio::sync::oneshot::channel();
let accept_task = tokio::spawn(async move {
let mut backoff = ACCEPT_BACKOFF_INITIAL;
let mut join_set: tokio::task::JoinSet<()> = tokio::task::JoinSet::new();
let mut shutdown_pin = std::pin::pin!(shutdown_rx);
let mut shutting_down = false;
loop {
if !shutting_down {
tokio::select! {
accepted = accept_and_permit(&listener, &mut backoff, &semaphore) => {
match accepted {
Some((stream, permit)) => {
let server = server.clone();
join_set.spawn(async move {
let _permit = permit;
serve_connection(stream, server, header_timeout).await;
});
}
None => shutting_down = true,
}
}
_ = shutdown_pin.as_mut() => {
shutting_down = true;
}
}
continue;
}
match join_set.join_next().await {
Some(_) => continue,
None => break,
}
}
});
Ok((
port,
ServerHandle {
shutdown_tx: Some(shutdown_tx),
accept_task,
},
))
}
pub async fn run(&self, header_timeout: Duration) -> Result<(u16, ServerHandle), StaticError> {
self.run_on(([127, 0, 0, 1], 0).into(), header_timeout)
.await
}
pub async fn run_all(
&self,
port: u16,
header_timeout: Duration,
) -> Result<(u16, ServerHandle), StaticError> {
self.run_on(([0, 0, 0, 0], port).into(), header_timeout)
.await
}
pub async fn run_ephemeral(&self) -> Result<(u16, ServerHandle), StaticError> {
self.run(DEFAULT_HEADER_TIMEOUT).await
}
pub async fn handle_request(
&self,
method: &Method,
request_path: &str,
headers: &HeaderMap,
) -> Response<ResponseBody> {
if method != Method::GET && method != Method::HEAD {
return text(
response(StatusCode::METHOD_NOT_ALLOWED).header("Allow", "GET, HEAD"),
"method not allowed\n",
);
}
if *method == Method::GET && request_path == reload::LIVE_RELOAD_PATH {
if let Some(broadcaster) = &self.broadcaster {
return finish(
response(StatusCode::OK)
.header("Content-Type", "text/event-stream")
.header("Cache-Control", "no-cache")
.header("Connection", "keep-alive")
.body(ResponseBody::Sse(SseBody::new(broadcaster.subscribe()))),
);
}
}
let server = self.clone();
let owned_request_path = request_path.to_string();
let resolved =
tokio::task::spawn_blocking(move || server.resolve(&owned_request_path)).await;
let path = match resolved {
Err(_) => return internal_error_response(),
Ok(Err(e)) => {
return text(
response(StatusCode::NOT_FOUND),
format!("{}\n", e.user_message()),
)
}
Ok(Ok(path)) => path,
};
let decoded_request_path = resolve::decode_request_path(request_path);
if path.file_name().is_some_and(|name| name == "index.html")
&& !decoded_request_path.ends_with('/')
&& !decoded_request_path.ends_with("index.html")
{
let location = format!("{}/", request_path.trim_end_matches('/'));
return text(
response(StatusCode::MOVED_PERMANENTLY).header("Location", location),
"moved\n",
);
}
let Ok(file) = File::open(&path).await else {
return internal_error_response();
};
let Ok(metadata) = file.metadata().await else {
return internal_error_response();
};
let content_type = mime_type_for_path(&path);
let html_injection = self.broadcaster.is_some() && content_type.starts_with("text/html");
let accept_encoding = header_str(headers, "accept-encoding");
let sidecar = if html_injection {
None
} else {
select_precompressed_sidecar(&path, accept_encoding).await
};
let (mut file, metadata, content_encoding) = match sidecar {
Some((sidecar_file, sidecar_metadata, encoding)) => {
(sidecar_file, sidecar_metadata, Some(encoding))
}
None => (file, metadata, None),
};
let etag = generate_etag(&metadata);
let cache_control = self.cache_control_for(&path);
if header_str(headers, "if-none-match").is_some_and(|value| is_etag_match(value, &etag)) {
return finish(
Response::builder()
.status(StatusCode::NOT_MODIFIED)
.header("Cache-Control", cache_control)
.header("Vary", "Accept-Encoding")
.header("ETag", etag)
.body(ResponseBody::Buffered(Full::new(Bytes::new()))),
);
}
let transformed: Option<Bytes> = if html_injection {
let mut html = Vec::with_capacity(metadata.len() as usize);
if file.read_to_end(&mut html).await.is_err() {
return internal_error_response();
}
reload::inject_reload_script(&mut html);
Some(Bytes::from(html))
} else {
None
};
let file_size = transformed
.as_ref()
.map_or(metadata.len(), |bytes| bytes.len() as u64);
let body = if *method == Method::HEAD {
ResponseBody::Buffered(Full::new(Bytes::new()))
} else {
match transformed {
Some(bytes) => ResponseBody::Buffered(Full::new(bytes)),
None => ResponseBody::Streamed(FileBody::new(file)),
}
};
let mut builder = response(StatusCode::OK)
.header("Content-Type", content_type)
.header("Content-Length", file_size.to_string())
.header("Cache-Control", cache_control)
.header("Vary", "Accept-Encoding")
.header("ETag", etag);
if let Some(encoding) = content_encoding {
builder = builder.header("Content-Encoding", encoding);
}
finish(builder.body(body))
}
}
const MAX_HEADER_BYTES: usize = 64 * 1024;
#[derive(Debug)]
enum HeaderReadError {
ConnectionClosed,
TooLarge,
#[allow(dead_code)]
Io(std::io::Error),
}
async fn read_header_prefix(stream: &mut TcpStream) -> Result<Vec<u8>, HeaderReadError> {
let mut buf = Vec::new();
let mut chunk = [0u8; 4096];
loop {
let n = stream.read(&mut chunk).await.map_err(HeaderReadError::Io)?;
if n == 0 {
return Err(HeaderReadError::ConnectionClosed);
}
buf.extend_from_slice(&chunk[..n]);
if buf.len() > MAX_HEADER_BYTES {
return Err(HeaderReadError::TooLarge);
}
let scan_from = buf.len().saturating_sub(n + 3);
if buf[scan_from..].windows(4).any(|w| w == b"\r\n\r\n") {
return Ok(buf);
}
}
}
struct PrefixedIo {
prefix: Bytes,
prefix_pos: usize,
inner: TcpStream,
}
impl PrefixedIo {
fn new(prefix: Vec<u8>, inner: TcpStream) -> Self {
PrefixedIo {
prefix: Bytes::from(prefix),
prefix_pos: 0,
inner,
}
}
}
impl AsyncRead for PrefixedIo {
fn poll_read(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &mut ReadBuf<'_>,
) -> Poll<std::io::Result<()>> {
let this = self.get_mut();
if this.prefix_pos < this.prefix.len() {
let remaining = &this.prefix[this.prefix_pos..];
let n = remaining.len().min(buf.remaining());
buf.put_slice(&remaining[..n]);
this.prefix_pos += n;
return Poll::Ready(Ok(()));
}
Pin::new(&mut this.inner).poll_read(cx, buf)
}
}
impl AsyncWrite for PrefixedIo {
fn poll_write(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &[u8],
) -> Poll<std::io::Result<usize>> {
Pin::new(&mut self.get_mut().inner).poll_write(cx, buf)
}
fn poll_flush(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
Pin::new(&mut self.get_mut().inner).poll_flush(cx)
}
fn poll_shutdown(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
Pin::new(&mut self.get_mut().inner).poll_shutdown(cx)
}
}
async fn serve_connection(mut stream: TcpStream, server: Server, header_timeout: Duration) {
let prefix = match timeout(header_timeout, read_header_prefix(&mut stream)).await {
Ok(Ok(prefix)) => prefix,
Ok(Err(_)) | Err(_) => return,
};
let io = TokioIo::new(PrefixedIo::new(prefix, stream));
let svc = service_fn(move |req: Request<Incoming>| {
let server = server.clone();
async move {
let resp = server
.handle_request(req.method(), req.uri().path(), req.headers())
.await;
Ok::<_, Infallible>(resp)
}
});
let _ = AutoBuilder::new(TokioExecutor::new())
.serve_connection(io, svc)
.await;
}
const DEFAULT_HEADER_TIMEOUT: Duration = Duration::from_secs(30);
const DEFAULT_SHUTDOWN_DRAIN_TIMEOUT: Duration = Duration::from_secs(5);
pub struct ServerHandle {
shutdown_tx: Option<tokio::sync::oneshot::Sender<()>>,
accept_task: tokio::task::JoinHandle<()>,
}
impl ServerHandle {
pub async fn shutdown(self) {
self.shutdown_with_timeout(DEFAULT_SHUTDOWN_DRAIN_TIMEOUT)
.await;
}
pub async fn shutdown_with_timeout(mut self, drain_timeout: Duration) {
if let Some(tx) = self.shutdown_tx.take() {
let _ = tx.send(());
}
if timeout(drain_timeout, &mut self.accept_task).await.is_err() {
self.accept_task.abort();
}
}
}
fn header_str<'h>(headers: &'h HeaderMap, name: &str) -> Option<&'h str> {
headers.get(name).and_then(|value| value.to_str().ok())
}
fn response(status: StatusCode) -> Builder {
Response::builder()
.status(status)
.header("X-Content-Type-Options", "nosniff")
}
fn text(builder: Builder, body: impl Into<Bytes>) -> Response<ResponseBody> {
finish(builder.body(ResponseBody::Buffered(Full::new(body.into()))))
}
fn finish(built: Result<Response<ResponseBody>, hyper::http::Error>) -> Response<ResponseBody> {
built.unwrap_or_else(|_| bad_request_response())
}
fn internal_error_response() -> Response<ResponseBody> {
response(StatusCode::INTERNAL_SERVER_ERROR)
.body(ResponseBody::Buffered(Full::new(Bytes::from_static(
b"internal server error\n",
))))
.unwrap()
}
fn bad_request_response() -> Response<ResponseBody> {
response(StatusCode::BAD_REQUEST)
.body(ResponseBody::Buffered(Full::new(Bytes::from_static(
b"bad request\n",
))))
.unwrap()
}
const SIDECAR_ENCODINGS: [(&str, &str); 2] = [("br", ".br"), ("gzip", ".gz")];
fn accepts_encoding(accept_encoding: Option<&str>, encoding: &str) -> bool {
accept_encoding.is_some_and(|header| header.contains(encoding))
}
async fn select_precompressed_sidecar(
path: &Path,
accept_encoding: Option<&str>,
) -> Option<(File, fs::Metadata, &'static str)> {
for (encoding, ext) in SIDECAR_ENCODINGS {
if !accepts_encoding(accept_encoding, encoding) {
continue;
}
let mut sidecar = path.as_os_str().to_os_string();
sidecar.push(ext);
let sidecar_path = PathBuf::from(sidecar);
debug_assert_eq!(
sidecar_path.parent(),
path.parent(),
"sidecar path must stay in the same directory as the already-resolved path"
);
if let Ok(sidecar_file) = File::open(&sidecar_path).await {
if let Ok(sidecar_metadata) = sidecar_file.metadata().await {
return Some((sidecar_file, sidecar_metadata, encoding));
}
}
}
None
}
fn generate_etag(metadata: &fs::Metadata) -> String {
let mtime = metadata
.modified()
.ok()
.and_then(|t| t.duration_since(SystemTime::UNIX_EPOCH).ok())
.map(|d| d.as_secs())
.unwrap_or(0);
format!("\"{}-{}\"", metadata.len(), mtime)
}
fn mime_type_for_path(path: &Path) -> &'static str {
let ext = path
.extension()
.and_then(|ext| ext.to_str())
.unwrap_or_default()
.to_lowercase();
match ext.as_str() {
"html" | "htm" => "text/html; charset=utf-8",
"css" => "text/css; charset=utf-8",
"js" => "application/javascript; charset=utf-8",
"json" => "application/json; charset=utf-8",
"svg" => "image/svg+xml",
"png" => "image/png",
"jpg" | "jpeg" => "image/jpeg",
"gif" => "image/gif",
"webp" => "image/webp",
"ico" => "image/x-icon",
"woff" => "font/woff",
"woff2" => "font/woff2",
"ttf" => "font/ttf",
"md" | "markdown" => "text/markdown; charset=utf-8",
"txt" => "text/plain; charset=utf-8",
"xml" => "application/xml",
"pdf" => "application/pdf",
"zip" => "application/zip",
_ => "application/octet-stream",
}
}
fn is_etag_match(if_none_match: &str, etag: &str) -> bool {
if if_none_match == "*" {
return true;
}
if_none_match.split(',').any(|tag| tag.trim() == etag)
}
#[cfg(test)]
mod precompressed_sidecar_tests {
use super::*;
#[tokio::test]
async fn sidecar_never_leaves_the_resolved_files_directory() {
let root = tempfile::TempDir::new().unwrap();
let sub = root.path().join("assets");
fs::create_dir(&sub).unwrap();
let resolved = sub.join("app.js");
fs::write(&resolved, b"plain").unwrap();
fs::write(sub.join("app.js.br"), b"brotli-bytes").unwrap();
fs::write(sub.join("app.js.gz"), b"gzip-bytes").unwrap();
let (_, _, encoding) = select_precompressed_sidecar(&resolved, Some("br, gzip"))
.await
.expect("both sidecars present, br should be preferred");
assert_eq!(
encoding, "br",
"br must be preferred over gzip when both are accepted"
);
let (_, _, encoding) = select_precompressed_sidecar(&resolved, Some("gzip"))
.await
.expect("gzip sidecar present");
assert_eq!(encoding, "gzip");
assert!(
select_precompressed_sidecar(&resolved, None)
.await
.is_none(),
"no Accept-Encoding header should never select a sidecar"
);
}
#[test]
fn accepts_encoding_matches_only_listed_directives() {
assert!(!accepts_encoding(None, "br"));
assert!(!accepts_encoding(Some("identity"), "br"));
assert!(!accepts_encoding(Some("identity"), "gzip"));
assert!(accepts_encoding(Some("gzip, br"), "br"));
assert!(accepts_encoding(Some("gzip"), "gzip"));
assert!(!accepts_encoding(Some("gzip"), "br"));
}
}
#[cfg(test)]
mod file_body_tests {
use super::*;
use crate::handler::FILE_CHUNK_SIZE;
use http_body_util::BodyExt;
#[tokio::test]
async fn file_body_yields_multiple_bounded_chunks_not_one_buffered_frame() {
let dir = tempfile::TempDir::new().unwrap();
let path = dir.path().join("big.bin");
let content = vec![7u8; FILE_CHUNK_SIZE * 3 + 12_345];
fs::write(&path, &content).unwrap();
let file = File::open(&path).await.unwrap();
let mut body = FileBody::new(file);
let mut frame_count = 0usize;
let mut max_frame_len = 0usize;
let mut reassembled = Vec::new();
while let Some(frame) = body.frame().await {
let frame = frame.unwrap();
let data = frame.into_data().unwrap();
frame_count += 1;
max_frame_len = max_frame_len.max(data.len());
reassembled.extend_from_slice(&data);
}
assert!(
frame_count > 1,
"expected the file to be delivered as multiple frames, got {frame_count}"
);
assert!(
max_frame_len <= FILE_CHUNK_SIZE,
"no single frame should exceed the chunk size ({FILE_CHUNK_SIZE}), got {max_frame_len}"
);
assert_eq!(
reassembled, content,
"reassembled chunks must match original file content exactly"
);
}
}
#[cfg(test)]
mod accept_tests {
use super::*;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::Mutex;
struct FlakyListener {
inner: TcpListener,
remaining_failures: AtomicUsize,
attempts: Mutex<Vec<tokio::time::Instant>>,
}
impl TcpAccept for FlakyListener {
async fn accept(&self) -> std::io::Result<(TcpStream, SocketAddr)> {
self.attempts
.lock()
.unwrap()
.push(tokio::time::Instant::now());
if self.remaining_failures.fetch_sub(1, Ordering::SeqCst) > 0 {
Err(std::io::Error::other("simulated accept error"))
} else {
TcpAccept::accept(&self.inner).await
}
}
}
#[tokio::test(start_paused = true)]
async fn accept_loop_backs_off_between_repeated_errors_instead_of_busy_spinning() {
let inner = TcpListener::bind(("127.0.0.1", 0)).await.unwrap();
let addr = inner.local_addr().unwrap();
let flaky = FlakyListener {
inner,
remaining_failures: AtomicUsize::new(5),
attempts: Mutex::new(Vec::new()),
};
tokio::spawn(async move {
let _ = TcpStream::connect(addr).await;
});
let semaphore = Arc::new(Semaphore::new(1));
let mut backoff = ACCEPT_BACKOFF_INITIAL;
let result = accept_and_permit(&flaky, &mut backoff, &semaphore).await;
assert!(
result.is_some(),
"accept should eventually succeed once the flaky listener stops failing"
);
let recorded = flaky.attempts.lock().unwrap();
assert_eq!(recorded.len(), 6, "5 failures then 1 success");
let expected_gaps = [
ACCEPT_BACKOFF_INITIAL,
ACCEPT_BACKOFF_INITIAL * 2,
ACCEPT_BACKOFF_INITIAL * 4,
ACCEPT_BACKOFF_INITIAL * 8,
ACCEPT_BACKOFF_INITIAL * 16,
];
for (i, expected) in expected_gaps.iter().enumerate() {
let gap = recorded[i + 1] - recorded[i];
assert_eq!(
gap,
*expected,
"gap between attempt {i} and {} should reflect the backoff delay, not a busy spin",
i + 1
);
}
let mut capped = ACCEPT_BACKOFF_MAX;
capped = (capped * 2).min(ACCEPT_BACKOFF_MAX);
assert_eq!(capped, ACCEPT_BACKOFF_MAX);
}
#[tokio::test(start_paused = true)]
async fn a_successful_accept_resets_the_backoff() {
let inner = TcpListener::bind(("127.0.0.1", 0)).await.unwrap();
let addr = inner.local_addr().unwrap();
let flaky = FlakyListener {
inner,
remaining_failures: AtomicUsize::new(3),
attempts: Mutex::new(Vec::new()),
};
tokio::spawn(async move {
let _ = TcpStream::connect(addr).await;
});
let semaphore = Arc::new(Semaphore::new(1));
let mut backoff = ACCEPT_BACKOFF_INITIAL * 32;
accept_and_permit(&flaky, &mut backoff, &semaphore).await;
assert_eq!(
backoff, ACCEPT_BACKOFF_INITIAL,
"the delay must return to its initial value once an accept succeeds"
);
}
}
#[cfg(test)]
mod finish_tests {
use super::*;
#[test]
fn finish_degrades_to_400_on_invalid_header_value_instead_of_panicking() {
let built = Response::builder()
.status(StatusCode::OK)
.header("X-Test", "invalid\r\nvalue")
.body(ResponseBody::Buffered(Full::new(Bytes::new())));
assert!(
built.is_err(),
"CR/LF in a header value should be rejected by the builder"
);
let response = finish(built);
assert_eq!(
response.status(),
StatusCode::BAD_REQUEST,
"finish() should degrade to 400 rather than panicking on an invalid header value"
);
}
}
#[cfg(test)]
mod header_prefix_tests {
use super::*;
use tokio::io::AsyncWriteExt;
async fn connected_pair() -> (TcpStream, TcpStream) {
let listener = TcpListener::bind(("127.0.0.1", 0)).await.unwrap();
let addr = listener.local_addr().unwrap();
let client = TcpStream::connect(addr).await.unwrap();
let (server_side, _) = listener.accept().await.unwrap();
(server_side, client)
}
#[tokio::test]
async fn reads_exactly_up_to_and_including_the_terminating_blank_line() {
let (mut server_side, mut client) = connected_pair().await;
client
.write_all(b"GET / HTTP/1.1\r\nHost: localhost\r\n\r\n")
.await
.unwrap();
let prefix = read_header_prefix(&mut server_side)
.await
.unwrap_or_else(|_| {
panic!("expected a complete header block to be read");
});
assert_eq!(prefix, b"GET / HTTP/1.1\r\nHost: localhost\r\n\r\n");
}
#[tokio::test]
async fn assembles_a_header_block_split_across_multiple_writes() {
let (mut server_side, mut client) = connected_pair().await;
client
.write_all(b"GET /page HTTP/1.1\r\nHost: localhost\r")
.await
.unwrap();
client.write_all(b"\n\r\n").await.unwrap();
let prefix = read_header_prefix(&mut server_side)
.await
.unwrap_or_else(|_| {
panic!("expected a complete header block to be read across multiple writes");
});
assert_eq!(prefix, b"GET /page HTTP/1.1\r\nHost: localhost\r\n\r\n");
}
#[tokio::test]
async fn preserves_bytes_sent_past_the_header_block() {
let (mut server_side, mut client) = connected_pair().await;
let first = b"GET /a HTTP/1.1\r\nHost: localhost\r\n\r\n";
let second = b"GET /b HTTP/1.1\r\nHost: localhost\r\n\r\n";
let mut sent = Vec::new();
sent.extend_from_slice(first);
sent.extend_from_slice(second);
client.write_all(&sent).await.unwrap();
let prefix = read_header_prefix(&mut server_side)
.await
.unwrap_or_else(|_| {
panic!("expected a complete header block to be read");
});
assert_eq!(
&prefix, &sent,
"pipelined bytes past the first header block must survive intact"
);
}
#[tokio::test]
async fn errors_with_connection_closed_when_client_disconnects_before_headers_complete() {
let (mut server_side, client) = connected_pair().await;
drop(client);
match read_header_prefix(&mut server_side).await {
Err(HeaderReadError::ConnectionClosed) => {}
Err(_) => panic!("expected ConnectionClosed, got a different error variant"),
Ok(_) => {
panic!("expected an error, got a complete header block from a closed connection")
}
}
}
#[tokio::test]
async fn errors_with_too_large_once_max_header_bytes_is_exceeded_without_a_terminator() {
let (mut server_side, mut client) = connected_pair().await;
let garbage = vec![b'a'; MAX_HEADER_BYTES + 1];
client.write_all(&garbage).await.unwrap();
match read_header_prefix(&mut server_side).await {
Err(HeaderReadError::TooLarge) => {}
Err(_) => panic!("expected TooLarge, got a different error variant"),
Ok(_) => {
panic!("expected an error, got a complete header block from unterminated garbage")
}
}
}
#[tokio::test]
async fn prefixed_io_replays_the_prefix_before_reading_from_the_live_socket() {
let (server_side, mut client) = connected_pair().await;
let mut io = PrefixedIo::new(b"buffered-prefix".to_vec(), server_side);
client.write_all(b"-live-bytes").await.unwrap();
let mut collected = Vec::new();
let mut chunk = [0u8; 8];
while collected.len() < b"buffered-prefix-live-bytes".len() {
let n = io.read(&mut chunk).await.unwrap();
assert!(n > 0, "read returned 0 before all expected bytes arrived");
collected.extend_from_slice(&chunk[..n]);
}
assert_eq!(collected, b"buffered-prefix-live-bytes");
}
}
#[cfg(test)]
mod css_bundle_tests {
use super::*;
use std::fs;
use std::time::Duration;
use tempfile::TempDir;
use tokio::time::sleep;
#[tokio::test]
async fn source_folder_overlapping_output_dir_is_rejected() {
let root = TempDir::new().unwrap();
let result = Server::new(root.path())
.unwrap()
.with_source_folder(root.path());
assert!(
result.is_err(),
"a source folder equal to the output dir must be rejected"
);
}
#[tokio::test]
async fn source_folder_inside_output_dir_is_rejected() {
let root = TempDir::new().unwrap();
let nested = root.path().join("nested");
fs::create_dir(&nested).unwrap();
let result = Server::new(root.path())
.unwrap()
.with_source_folder(&nested);
assert!(
result.is_err(),
"a source folder nested in the output dir must be rejected"
);
}
#[tokio::test]
async fn output_dir_overlapping_source_folder_is_rejected() {
let root = TempDir::new().unwrap();
let source = TempDir::new().unwrap();
let server = Server::new(root.path())
.unwrap()
.with_source_folder(source.path())
.unwrap();
let result = server.with_output_dir(source.path());
assert!(
result.is_err(),
"an output dir equal to a source folder must be rejected"
);
}
#[tokio::test]
async fn css_bundle_creates_output_on_startup_with_live_reload() {
let src = TempDir::new().unwrap();
let out = TempDir::new().unwrap();
fs::write(src.path().join("style.css"), "body { margin: 0; }").unwrap();
let server = Server::new(out.path())
.unwrap()
.with_live_reload()
.with_source_folder(src.path())
.unwrap()
.with_css_tool(CssTool::TestEcho, CssOptions::new().bundle(true));
let (_port, handle) = server.run_ephemeral().await.unwrap();
sleep(Duration::from_millis(800)).await;
let bundle = out.path().join("styles.css");
assert!(
bundle.exists(),
"bundle should be written to the default <output>/styles.css"
);
let content = fs::read_to_string(&bundle).unwrap();
assert!(!content.is_empty(), "bundle should contain CSS");
handle.shutdown().await;
}
#[tokio::test]
async fn css_bundle_rebuilds_once_and_settles_when_source_css_changes() {
let src = TempDir::new().unwrap();
let out = TempDir::new().unwrap();
let src_path = src.path();
let bundle = out.path().join("styles.css");
fs::write(src_path.join("style.css"), "body { margin: 0; }").unwrap();
let server = Server::new(out.path())
.unwrap()
.with_live_reload()
.with_source_folder(src_path)
.unwrap()
.with_css_tool(CssTool::TestEcho, CssOptions::new().bundle(true));
let (_port, handle) = server.run_ephemeral().await.unwrap();
sleep(Duration::from_millis(800)).await;
assert!(bundle.exists());
fs::write(
src_path.join("style.css"),
"body { margin: 0; color: blue; }",
)
.unwrap();
sleep(Duration::from_millis(1500)).await;
let content_v2 = fs::read_to_string(&bundle).unwrap();
assert!(
content_v2.contains("color"),
"rebundle should contain the new color rule"
);
let mtime_after = fs::metadata(&bundle).unwrap().modified().unwrap();
sleep(Duration::from_millis(1200)).await;
let mtime_later = fs::metadata(&bundle).unwrap().modified().unwrap();
assert_eq!(
mtime_after, mtime_later,
"bundle mtime must settle after one rebuild — an ongoing loop would keep changing it"
);
handle.shutdown().await;
}
#[tokio::test]
async fn css_bundle_creates_output_on_startup_without_live_reload() {
let src = TempDir::new().unwrap();
let out = TempDir::new().unwrap();
fs::write(src.path().join("style.css"), "body { margin: 0; }").unwrap();
let server = Server::new(out.path())
.unwrap()
.with_source_folder(src.path())
.unwrap()
.with_css_tool(CssTool::TestEcho, CssOptions::new().bundle(true));
let (_port, handle) = server.run_ephemeral().await.unwrap();
sleep(Duration::from_millis(200)).await;
let bundle = out.path().join("styles.css");
assert!(
bundle.exists(),
"bundle should be created even without live_reload"
);
let content = fs::read_to_string(&bundle).unwrap();
assert!(!content.is_empty(), "bundle should contain CSS");
handle.shutdown().await;
}
#[tokio::test]
async fn css_bundle_concatenates_multiple_source_css_files() {
let src = TempDir::new().unwrap();
let out = TempDir::new().unwrap();
fs::write(src.path().join("reset.css"), "* { margin: 0; padding: 0; }").unwrap();
fs::write(src.path().join("theme.css"), "body { background: white; }").unwrap();
let server = Server::new(out.path())
.unwrap()
.with_live_reload()
.with_source_folder(src.path())
.unwrap()
.with_css_tool(CssTool::TestEcho, CssOptions::new().bundle(true));
let (_port, handle) = server.run_ephemeral().await.unwrap();
sleep(Duration::from_millis(800)).await;
let content = fs::read_to_string(out.path().join("styles.css")).unwrap();
assert!(
content.contains("margin"),
"output should contain reset CSS"
);
assert!(
content.contains("background"),
"output should contain theme CSS"
);
handle.shutdown().await;
}
}