use crate::error_page;
use crate::headers::Headers;
use crate::method::Method;
use crate::panic;
use crate::request::Request;
use crate::response::Response;
use crate::router::Router;
use crate::status::Status;
use crate::url;
use rustlavel_core::{Context, Error, Result};
use std::net::SocketAddr;
use std::sync::Arc;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::time::{Duration, Instant};
use tokio::io::{AsyncReadExt, AsyncWriteExt, BufWriter};
use tokio::net::{TcpListener, TcpStream};
#[derive(Debug, Clone)]
pub struct Limits {
pub max_header_bytes: usize,
pub max_body_bytes: usize,
pub keep_alive_timeout: Duration,
pub header_timeout: Duration,
}
impl Default for Limits {
fn default() -> Self {
Limits {
max_header_bytes: 64 * 1024,
max_body_bytes: 10 * 1024 * 1024,
keep_alive_timeout: Duration::from_secs(15),
header_timeout: Duration::from_secs(10),
}
}
}
fn port_attempts(config: &rustlavel_core::Config) -> u16 {
let default = if config.is_production() { 1 } else { 10 };
config.int("server.port_attempts", default).clamp(1, 1000) as u16
}
fn check_framing(headers: &Headers) -> Result<()> {
let lengths = headers.get_all("content-length");
if lengths.len() > 1 && lengths.iter().any(|value| value != &lengths[0]) {
return Err(Error::Protocol(
"more than one Content-Length, and they disagree".into(),
));
}
if !lengths.is_empty() && headers.get("transfer-encoding").is_some() {
return Err(Error::Protocol(
"both Transfer-Encoding and Content-Length: a message may say where it ends \
once, not twice"
.into(),
));
}
let encodings = headers.get_all("transfer-encoding");
if !encodings.is_empty() {
let listed: Vec<&str> = encodings
.iter()
.flat_map(|value| value.split(','))
.map(str::trim)
.filter(|value| !value.is_empty())
.collect();
if listed.last() != Some(&"chunked")
|| listed.iter().filter(|value| **value == "chunked").count() != 1
{
return Err(Error::Protocol(format!(
"unsupported Transfer-Encoding: {}",
listed.join(", ")
)));
}
}
Ok(())
}
async fn bind_walking(addr: &str, attempts: u16) -> Result<TcpListener> {
let Some((host, first)) = addr
.rsplit_once(':')
.and_then(|(host, port)| port.parse::<u16>().ok().map(|port| (host, port)))
.filter(|(_, port)| *port != 0)
else {
return TcpListener::bind(addr).await.map_err(Error::Io);
};
let mut tried = first;
for offset in 0..attempts {
let Some(port) = first.checked_add(offset) else { break };
tried = port;
match TcpListener::bind(format!("{host}:{port}")).await {
Ok(listener) => {
if offset > 0 {
rustlavel_core::warn!(
"port {first} is in use, so this is serving on {port} instead"
);
}
return Ok(listener);
}
Err(e) if e.kind() == std::io::ErrorKind::AddrInUse => continue,
Err(e) => return Err(Error::Io(e)),
}
}
Err(Error::msg(if first == tried {
format!(
"port {first} is already in use. Something else is listening on it — `lsof -i :{first}` \
says what — so stop that, or set SERVER_PORT to a free port."
)
} else {
format!(
"every port from {first} to {tried} is already in use. Stop whatever is holding them \
— `lsof -i :{first}` names the first — or set SERVER_PORT to a free one."
)
}))
}
#[cfg(unix)]
async fn stop_requested() {
use tokio::signal::unix::{SignalKind, signal};
let mut terminate = match signal(SignalKind::terminate()) {
Ok(stream) => stream,
Err(error) => {
rustlavel_core::warn!("cannot listen for SIGTERM: {error}");
let _ = tokio::signal::ctrl_c().await;
return;
}
};
tokio::select! {
_ = tokio::signal::ctrl_c() => {}
_ = terminate.recv() => {}
}
}
#[cfg(not(unix))]
async fn stop_requested() {
let _ = tokio::signal::ctrl_c().await;
}
pub type OnShutdown = Box<dyn FnOnce() -> crate::handler::BoxFuture<()> + Send + Sync>;
pub struct Server {
router: Arc<Router>,
context: Context,
limits: Limits,
on_shutdown: Vec<OnShutdown>,
}
impl Server {
pub fn new(mut router: Router, context: Context) -> Self {
router.finalize();
let limits = Limits {
max_body_bytes: context.config().int("server.max_body_bytes", 10 * 1024 * 1024) as usize,
..Limits::default()
};
Server { router: Arc::new(router), context, limits, on_shutdown: Vec::new() }
}
pub fn on_shutdown(mut self, work: OnShutdown) -> Self {
self.on_shutdown.push(work);
self
}
pub fn limits(mut self, limits: Limits) -> Self {
self.limits = limits;
self
}
pub async fn listen(mut self, addr: impl Into<String>) -> Result<()> {
let addr = addr.into();
let listener = bind_walking(&addr, port_attempts(self.context.config())).await?;
let local = listener.local_addr().map_err(Error::Io)?;
panic::install_hook();
error_page::set_debug(self.context.config().debug());
rustlavel_core::info!("Rustlavel serving on http://{local}");
rustlavel_core::info!("Press Ctrl-C to stop");
let in_flight = Arc::new(AtomicUsize::new(0));
let mut on_shutdown = std::mem::take(&mut self.on_shutdown);
let shared = Arc::new(self);
loop {
let accepted = tokio::select! {
result = listener.accept() => result,
_ = stop_requested() => break,
};
let (stream, peer) = match accepted {
Ok(pair) => pair,
Err(e) => {
rustlavel_core::warn!("accept failed: {e}");
continue;
}
};
let server = Arc::clone(&shared);
let counter = Arc::clone(&in_flight);
counter.fetch_add(1, Ordering::SeqCst);
tokio::spawn(async move {
if let Err(e) = server.serve_connection(stream, peer).await {
rustlavel_core::debug!("connection closed: {e}");
}
counter.fetch_sub(1, Ordering::SeqCst);
});
}
rustlavel_core::info!("Shutting down, waiting for in-flight requests…");
let deadline = Instant::now() + Duration::from_secs(10);
while in_flight.load(Ordering::SeqCst) > 0 && Instant::now() < deadline {
tokio::time::sleep(Duration::from_millis(25)).await;
}
for work in on_shutdown.drain(..) {
work().await;
}
rustlavel_core::info!("Goodbye.");
Ok(())
}
pub(crate) async fn serve_connection(&self, stream: TcpStream, peer: SocketAddr) -> Result<()> {
let _ = stream.set_nodelay(true);
let (mut reader, writer) = stream.into_split();
let mut writer = BufWriter::new(writer);
let mut buffer: Vec<u8> = Vec::with_capacity(2048);
loop {
let head = match self.read_head(&mut reader, &mut buffer).await? {
Some(head) => head,
None => return Ok(()),
};
let (mut request, keep_alive) = match self.parse(&head, &mut reader, &mut buffer, peer).await {
Ok(parsed) => parsed,
Err(error) => {
let response = Response::new(Status::BAD_REQUEST).with_text(error.to_string());
writer.write_all(&response.to_bytes(true)).await.map_err(Error::Io)?;
writer.flush().await.map_err(Error::Io)?;
return Ok(());
}
};
request.context = self.context.clone();
let is_head = request.method() == Method::Head;
let mut response = self.dispatch(request).await;
if response.upgrades() {
let head = response.to_bytes(false);
let upgrade = response.take_upgrade().expect("checked just above");
writer.write_all(&head).await.map_err(Error::Io)?;
writer.flush().await.map_err(Error::Io)?;
let upgraded = crate::upgrade::Upgraded {
reader: Box::new(reader),
writer: Box::new(writer),
buffered: std::mem::take(&mut buffer),
};
upgrade.run(upgraded).await;
return Ok(());
}
if !keep_alive {
response.headers.set("connection", "close");
}
writer.write_all(&response.to_bytes(!is_head)).await.map_err(Error::Io)?;
writer.flush().await.map_err(Error::Io)?;
if !keep_alive {
return Ok(());
}
}
}
async fn dispatch(&self, request: Request) -> Response {
let started = Instant::now();
let method = request.method();
let path = request.path().to_string();
let response = self.router.dispatch(request).await;
let elapsed = started.elapsed();
if rustlavel_core::log::enabled(rustlavel_core::log::Level::Debug) {
rustlavel_core::debug!(
"{method} {path} → {} ({:.1}ms)",
response.status.code(),
elapsed.as_secs_f64() * 1000.0
);
}
response
}
async fn read_head(
&self,
reader: &mut tokio::net::tcp::OwnedReadHalf,
buffer: &mut Vec<u8>,
) -> Result<Option<Vec<u8>>> {
let mut timeout = self.limits.keep_alive_timeout;
loop {
if let Some(end) = find_head_end(buffer) {
let head = buffer[..end].to_vec();
buffer.drain(..end);
return Ok(Some(head));
}
if buffer.len() > self.limits.max_header_bytes {
return Err(Error::Protocol("request headers are too large".into()));
}
let mut chunk = [0u8; 4096];
let read = match tokio::time::timeout(timeout, reader.read(&mut chunk)).await {
Ok(Ok(0)) if buffer.is_empty() => return Ok(None),
Ok(Ok(0)) => return Err(Error::Protocol("connection closed mid-request".into())),
Ok(Ok(n)) => n,
Ok(Err(e)) => return Err(Error::Io(e)),
Err(_) if buffer.is_empty() => return Ok(None),
Err(_) => return Err(Error::Protocol("timed out reading request headers".into())),
};
buffer.extend_from_slice(&chunk[..read]);
timeout = self.limits.header_timeout;
}
}
async fn parse(
&self,
head: &[u8],
reader: &mut tokio::net::tcp::OwnedReadHalf,
buffer: &mut Vec<u8>,
peer: SocketAddr,
) -> Result<(Request, bool)> {
let text = std::str::from_utf8(head).map_err(|_| Error::Protocol("headers are not UTF-8".into()))?;
let mut lines = text.split("\r\n");
let request_line = lines.next().ok_or_else(|| Error::Protocol("empty request".into()))?;
let mut parts = request_line.split(' ');
let method = parts
.next()
.and_then(Method::parse)
.ok_or_else(|| Error::Protocol("unsupported method".into()))?;
let target = parts.next().ok_or_else(|| Error::Protocol("missing request target".into()))?;
let version = parts.next().unwrap_or("HTTP/1.1");
let mut headers = Headers::new();
for line in lines {
if line.is_empty() {
continue;
}
let (name, value) = line
.split_once(':')
.ok_or_else(|| Error::Protocol(format!("malformed header line: {line}")))?;
if name.ends_with(' ') || name.ends_with('\t') {
return Err(Error::Protocol(
"a header name may not be followed by whitespace before the colon".into(),
));
}
headers.append(name.trim(), value.trim());
}
check_framing(&headers)?;
let target = match target.find("://") {
Some(scheme_end) => match target[scheme_end + 3..].find('/') {
Some(path_start) => &target[scheme_end + 3 + path_start..],
None => "/",
},
None => target,
};
let body = self.read_body(&headers, reader, buffer).await?;
let keep_alive = match headers.get("connection") {
Some(value) if value.eq_ignore_ascii_case("close") => false,
Some(value) if value.eq_ignore_ascii_case("keep-alive") => true,
_ => version != "HTTP/1.0",
};
let (path, query) = url::split_target(target);
let mut request = Request::new(method, target);
request.path = url::decode(path);
request.query = url::parse_query(query);
request.headers = headers;
request.peer = Some(peer);
Ok((request.with_body(body), keep_alive))
}
async fn read_body(
&self,
headers: &Headers,
reader: &mut tokio::net::tcp::OwnedReadHalf,
buffer: &mut Vec<u8>,
) -> Result<Vec<u8>> {
if headers.get("transfer-encoding").is_some() {
return self.read_chunked_body(reader, buffer).await;
}
let Some(length) = headers.content_length() else {
return Ok(Vec::new());
};
if length > self.limits.max_body_bytes {
return Err(Error::Protocol("request body is too large".into()));
}
while buffer.len() < length {
let mut chunk = vec![0u8; (length - buffer.len()).min(64 * 1024)];
let read = tokio::time::timeout(self.limits.header_timeout, reader.read(&mut chunk))
.await
.map_err(|_| Error::Protocol("timed out reading request body".into()))?
.map_err(Error::Io)?;
if read == 0 {
return Err(Error::Protocol("request body ended early".into()));
}
buffer.extend_from_slice(&chunk[..read]);
}
Ok(buffer.drain(..length).collect())
}
async fn read_chunked_body(
&self,
reader: &mut tokio::net::tcp::OwnedReadHalf,
buffer: &mut Vec<u8>,
) -> Result<Vec<u8>> {
let mut body = Vec::new();
loop {
let line_end = loop {
if let Some(at) = find_crlf(buffer) {
break at;
}
if !fill(reader, buffer, self.limits.header_timeout).await? {
return Err(Error::Protocol("chunked body ended early".into()));
}
};
let header: Vec<u8> = buffer.drain(..line_end + 2).collect();
let size_text = String::from_utf8_lossy(&header[..line_end]);
let size = usize::from_str_radix(size_text.split(';').next().unwrap_or("").trim(), 16)
.map_err(|_| Error::Protocol("invalid chunk size".into()))?;
if size == 0 {
loop {
let end = loop {
if let Some(at) = find_crlf(buffer) {
break at;
}
if !fill(reader, buffer, self.limits.header_timeout).await? {
return Ok(body);
}
};
buffer.drain(..end + 2);
if end == 0 {
return Ok(body);
}
}
}
if body.len() + size > self.limits.max_body_bytes {
return Err(Error::Protocol("request body is too large".into()));
}
while buffer.len() < size + 2 {
if !fill(reader, buffer, self.limits.header_timeout).await? {
return Err(Error::Protocol("chunked body ended early".into()));
}
}
body.extend(buffer.drain(..size));
buffer.drain(..2);
}
}
}
async fn fill(
reader: &mut tokio::net::tcp::OwnedReadHalf,
buffer: &mut Vec<u8>,
timeout: Duration,
) -> Result<bool> {
let mut chunk = [0u8; 4096];
let read = tokio::time::timeout(timeout, reader.read(&mut chunk))
.await
.map_err(|_| Error::Protocol("timed out reading request body".into()))?
.map_err(Error::Io)?;
buffer.extend_from_slice(&chunk[..read]);
Ok(read > 0)
}
fn find_head_end(buffer: &[u8]) -> Option<usize> {
buffer.windows(4).position(|w| w == b"\r\n\r\n").map(|at| at + 4)
}
fn find_crlf(buffer: &[u8]) -> Option<usize> {
buffer.windows(2).position(|w| w == b"\r\n")
}
#[cfg(test)]
mod tests {
fn framing_of(raw: &str) -> Result<()> {
let mut headers = Headers::new();
for line in raw.split("\r\n").skip(1) {
if line.is_empty() {
break;
}
let (name, value) = line.split_once(':').expect("a header line");
if name.ends_with(' ') || name.ends_with('\t') {
return Err(Error::Protocol("whitespace before the colon".into()));
}
headers.append(name.trim(), value.trim());
}
check_framing(&headers)
}
#[test]
fn a_message_that_says_where_it_ends_twice_is_refused() {
let ambiguous = [
(
"two lengths that disagree",
"POST / HTTP/1.1\r\nhost: x\r\ncontent-length: 6\r\ncontent-length: 0\r\n\r\nsmuggl",
),
(
"a length and a chunked encoding",
"POST / HTTP/1.1\r\nhost: x\r\ncontent-length: 6\r\ntransfer-encoding: chunked\r\n\r\n0\r\n\r\n",
),
(
"an encoding that is not chunked last",
"POST / HTTP/1.1\r\nhost: x\r\ntransfer-encoding: chunked, identity\r\n\r\n0\r\n\r\n",
),
(
"whitespace before the colon",
"POST / HTTP/1.1\r\nhost: x\r\ncontent-length : 6\r\n\r\nsmuggl",
),
];
for (what, raw) in ambiguous {
assert!(
framing_of(raw).is_err(),
"{what}: accepted a message with two answers for where its body ends"
);
}
}
#[test]
fn an_unambiguous_message_still_parses() {
for raw in [
"GET / HTTP/1.1\r\nhost: x\r\n\r\n",
"POST / HTTP/1.1\r\nhost: x\r\ncontent-length: 3\r\n\r\nabc",
"POST / HTTP/1.1\r\nhost: x\r\ntransfer-encoding: chunked\r\n\r\n0\r\n\r\n",
"POST / HTTP/1.1\r\nhost: x\r\ncontent-length: 3\r\ncontent-length: 3\r\n\r\nabc",
] {
assert!(framing_of(raw).is_ok(), "refused an ordinary message: {raw:?}");
}
}
#[test]
fn production_gets_one_attempt_and_development_gets_more() {
use rustlavel_core::Config;
let production = Config::with_defaults();
production.set("app.env", "production");
assert_eq!(port_attempts(&production), 1);
let local = Config::with_defaults();
local.set("app.env", "local");
assert!(port_attempts(&local) > 1);
}
#[test]
fn the_setting_overrides_the_environment_both_ways() {
use rustlavel_core::Config;
let production = Config::with_defaults();
production.set("app.env", "production");
production.set("server.port_attempts", "5");
assert_eq!(port_attempts(&production), 5);
let local = Config::with_defaults();
local.set("app.env", "local");
local.set("server.port_attempts", "1");
assert_eq!(port_attempts(&local), 1);
}
#[tokio::test]
async fn a_taken_port_moves_to_the_next_one() {
let held = TcpListener::bind("127.0.0.1:0").await.unwrap();
let taken = held.local_addr().unwrap().port();
let listener = bind_walking(&format!("127.0.0.1:{taken}"), 10).await.unwrap();
assert_ne!(listener.local_addr().unwrap().port(), taken);
assert!(listener.local_addr().unwrap().port() > taken);
}
#[tokio::test]
async fn a_single_attempt_does_not_move() {
let held = TcpListener::bind("127.0.0.1:0").await.unwrap();
let taken = held.local_addr().unwrap().port();
let error = bind_walking(&format!("127.0.0.1:{taken}"), 1).await.unwrap_err();
let message = error.to_string();
assert!(message.contains(&taken.to_string()), "{message}");
assert!(message.contains("lsof"), "the message has to say what to do: {message}");
}
#[tokio::test]
async fn port_zero_is_left_to_the_operating_system() {
let listener = bind_walking("127.0.0.1:0", 10).await.unwrap();
assert_ne!(listener.local_addr().unwrap().port(), 0);
}
#[tokio::test]
async fn exhausting_the_range_says_what_it_tried() {
let first = TcpListener::bind("127.0.0.1:0").await.unwrap();
let start = first.local_addr().unwrap().port();
let mut held = vec![first];
for offset in 1..3u16 {
if let Ok(listener) = TcpListener::bind(format!("127.0.0.1:{}", start + offset)).await {
held.push(listener);
}
}
let error = bind_walking(&format!("127.0.0.1:{start}"), 3).await;
if let Err(error) = error {
let message = error.to_string();
assert!(message.contains(&start.to_string()), "{message}");
assert!(message.contains(&(start + 2).to_string()), "{message}");
}
}
use super::*;
#[test]
fn finds_the_end_of_a_header_block() {
assert_eq!(find_head_end(b"GET / HTTP/1.1\r\n\r\nbody"), Some(18));
assert_eq!(find_head_end(b"GET / HTTP/1.1\r\n"), None);
}
#[tokio::test]
async fn parses_a_request_with_a_body() {
let server = Server::new(Router::new(), Context::default());
let head = b"POST /users?page=2 HTTP/1.1\r\nHost: localhost\r\nContent-Type: application/json\r\nContent-Length: 14\r\n\r\n";
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
tokio::spawn(async move {
let (mut stream, _) = listener.accept().await.unwrap();
stream.write_all(br#"{"name":"ada"}"#).await.unwrap();
});
let stream = TcpStream::connect(addr).await.unwrap();
let (mut reader, _writer) = stream.into_split();
let mut buffer = Vec::new();
let (mut request, keep_alive) =
server.parse(head, &mut reader, &mut buffer, addr).await.unwrap();
assert_eq!(request.method(), Method::Post);
assert_eq!(request.path(), "/users");
assert_eq!(request.query("page"), Some("2"));
assert_eq!(request.header("host"), Some("localhost"));
assert_eq!(request.input("name").as_deref(), Some("ada"));
assert!(keep_alive);
}
#[tokio::test]
async fn http_1_0_closes_by_default() {
let server = Server::new(Router::new(), Context::default());
let head = b"GET / HTTP/1.0\r\nHost: localhost\r\n\r\n";
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
tokio::spawn(async move {
let _ = listener.accept().await;
});
let (mut reader, _w) = TcpStream::connect(addr).await.unwrap().into_split();
let mut buffer = Vec::new();
let (_request, keep_alive) =
server.parse(head, &mut reader, &mut buffer, addr).await.unwrap();
assert!(!keep_alive);
}
#[tokio::test]
async fn rejects_a_body_larger_than_the_limit() {
let mut server = Server::new(Router::new(), Context::default());
server.limits.max_body_bytes = 8;
let head = b"POST / HTTP/1.1\r\nContent-Length: 9999\r\n\r\n";
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
tokio::spawn(async move {
let _ = listener.accept().await;
});
let (mut reader, _w) = TcpStream::connect(addr).await.unwrap().into_split();
let mut buffer = Vec::new();
let error = server.parse(head, &mut reader, &mut buffer, addr).await.unwrap_err();
assert!(error.to_string().contains("too large"));
}
}