diff options
Diffstat
| -rw-r--r-- | Cargo.lock | 26 | +25 −1 |
| -rw-r--r-- | Cargo.toml | 6 | +2 −4 |
| -rw-r--r-- | nfpm.yaml | 7 | +4 −3 |
| -rw-r--r-- | pkg/zorian.toml | 2 | +1 −1 |
| -rw-r--r-- | src/config.rs | 51 | +33 −18 |
| -rw-r--r-- | src/main.rs | 214 | +71 −143 |
6 files changed, 136 insertions, 170 deletions
diff --git a/Cargo.lock b/Cargo.lock index 2547d9a..55a7227 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -259,6 +259,28 @@ dependencies = [ ] [[package]] +name = "axum-server" +version = "0.8.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b1df331683d982a0b9492b38127151e6453639cd34926eb9c07d4cd8c6d22bfc" +dependencies = [ + "arc-swap", + "bytes", + "either", + "fs-err", + "http", + "http-body", + "hyper", + "hyper-util", + "pin-project-lite", + "rustls", + "rustls-pki-types", + "tokio", + "tokio-rustls", + "tower-service", +] + +[[package]] name = "base64" version = "0.22.1" source = "registry+https://github.com/rust-lang/crates.io-index" @@ -1044,6 +1066,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "baf68cef89750956493a66a10f512b9e58d9db21f2a573c079c0bdf1207a54a7" dependencies = [ "autocfg", + "tokio", ] [[package]] @@ -5525,6 +5548,7 @@ version = "0.1.0" dependencies = [ "axum", "axum-extra", + "axum-server", "bytes", "bytesize", "cargo-deny", @@ -5537,13 +5561,13 @@ dependencies = [ "hyper-tls", "hyper-util", "minijinja", + "rustls", "semver", "serde", "sqlx", "tempfile", "thiserror 2.0.17", "tokio", - "tokio-rustls", "toml", "tower", "tower-http", diff --git a/Cargo.toml b/Cargo.toml index 37ff052..9b4618f 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -10,6 +10,8 @@ repository = "https://github.com/mrdimidium/Zorian" [dependencies] axum = { version = "0.8", features = ["http2", "macros", "multipart"] } +axum-server = { version = "0.8", default-features = false, features = ["tls-rustls-no-provider"] } +rustls = { version = "0.23", default-features = false, features = ["ring", "std", "tls12"] } axum-extra = { version = "0.12", features = [ "async-read-body", "attachment", @@ -55,10 +57,6 @@ bytesize = { version = "2.3", features = ["serde"] } thiserror = "2.0" tokio = { version = "1.49", features = ["full"] } toml = "0.9" -tokio-rustls = { version = "0.26", default-features = false, features = [ - "ring", - "tls12", -] } tower = { version = "0.5", features = ["full"] } tower_governor = { version = "0.8", features = ["axum", "tracing"] } tower-http = { version = "0.6", features = ["full"] } diff --git a/nfpm.yaml b/nfpm.yaml index 356145b..262bde9 100644 --- a/nfpm.yaml +++ b/nfpm.yaml @@ -27,8 +27,9 @@ contents: dst: /etc/zorian.toml type: config|noreplace file_info: - mode: 0644 - owner: zorian + mode: 0640 + owner: root + group: zorian - src: pkg/zorian.service dst: /usr/lib/systemd/system/zorian.service @@ -39,7 +40,7 @@ contents: type: dir file_info: mode: 0755 - owner: zorian + owner: root group: zorian diff --git a/pkg/zorian.toml b/pkg/zorian.toml index b1a60f6..9612063 100644 --- a/pkg/zorian.toml +++ b/pkg/zorian.toml @@ -21,7 +21,7 @@ rate_limit_burst_size = 50 # Listeners configuration. Each listener has an address and optional hostnames. # Empty hostnames means accept all requests (catch-all). [[listen]] -addr = "0.0.0.0:3000" +addr = "0.0.0.0:2025" hostnames = ["localhost", "127.0.0.1", "::1"] # For HTTPS, set tls_cert and tls_key to PEM file paths. diff --git a/src/config.rs b/src/config.rs index c4daaaa..49b5c15 100644 --- a/src/config.rs +++ b/src/config.rs @@ -2,6 +2,7 @@ // SPDX-License-Identifier: AGPL-3.0-or-later use std::fs; +use std::net::SocketAddr; use std::path::{Path, PathBuf}; use std::time::Duration; @@ -17,6 +18,25 @@ where Ok(Duration::from_secs(secs)) } +fn deserialize_listener_addr<'de, D>(deserializer: D) -> Result<SocketAddr, D::Error> +where + D: serde::Deserializer<'de>, +{ + let raw = String::deserialize(deserializer)?; + let raw = raw.trim(); + + // Hostnames are not supported, but localhost shorthands are useful + if let Some(port) = raw.strip_prefix("localhost:") { + return port + .parse::<u16>() + .map(|port| SocketAddr::new(std::net::Ipv4Addr::LOCALHOST.into(), port)) + .map_err(|err| serde::de::Error::custom(format!("invalid port in '{raw}': {err}"))); + } + + raw.parse::<SocketAddr>() + .map_err(|err| serde::de::Error::custom(format!("invalid address '{raw}': {err}",))) +} + #[derive(Debug, Error)] pub enum ConfigError { #[error("failed to read config file: {0}")] @@ -38,16 +58,16 @@ pub enum ConfigError { NotWritable(PathBuf, std::io::Error), #[error("listener '{0}': tls_crt is set but tls_key is missing")] - TlsKeyMissing(String), + TlsKeyMissing(SocketAddr), #[error("listener '{0}': tls_key is set but tls_crt is missing")] - TlsCrtMissing(String), + TlsCrtMissing(SocketAddr), #[error("listener '{0}': TLS crtificate file not found: {1}")] - TlsCrtNotFound(String, PathBuf), + TlsCrtNotFound(SocketAddr, PathBuf), #[error("listener '{0}': TLS key file not found: {1}")] - TlsKeyNotFound(String, PathBuf), + TlsKeyNotFound(SocketAddr, PathBuf), } #[derive(Debug, Clone, Deserialize)] @@ -89,7 +109,8 @@ impl Default for ServerConfig { #[derive(Debug, Clone, Deserialize)] pub struct ListenerConfig { - pub addr: String, + #[serde(deserialize_with = "deserialize_listener_addr")] + pub addr: SocketAddr, /// Hostnames to accept for this listener. Empty means accept all. #[serde(default)] @@ -105,7 +126,7 @@ pub struct ListenerConfig { impl Default for ListenerConfig { fn default() -> Self { Self { - addr: "127.0.0.1:3000".to_string(), + addr: "127.0.0.1:2025".parse().unwrap(), hostnames: vec![ String::from("[::1]"), String::from("127.0.0.1"), @@ -167,27 +188,21 @@ impl ConfigService { .map_err(|e| ConfigError::NotWritable(self.dirname.clone(), e))?; fs::remove_file(&testfile)?; - // Validate TLS configuration for each listener + // Validate listener configuration for listener in &self.listen { match (&listener.tls_crt, &listener.tls_key) { (Some(_crt), None) => { - return Err(ConfigError::TlsKeyMissing(listener.addr.clone())); + return Err(ConfigError::TlsKeyMissing(listener.addr)); } (None, Some(_)) => { - return Err(ConfigError::TlsCrtMissing(listener.addr.clone())); + return Err(ConfigError::TlsCrtMissing(listener.addr)); } (Some(crt), Some(key)) => { if !crt.exists() { - return Err(ConfigError::TlsCrtNotFound( - listener.addr.clone(), - crt.clone(), - )); + return Err(ConfigError::TlsCrtNotFound(listener.addr, crt.clone())); } if !key.exists() { - return Err(ConfigError::TlsKeyNotFound( - listener.addr.clone(), - key.clone(), - )); + return Err(ConfigError::TlsKeyNotFound(listener.addr, key.clone())); } } (None, None) => {} @@ -222,7 +237,7 @@ impl ConfigService { dirname, server: ServerConfig::default(), listen: vec![ListenerConfig { - addr: "127.0.0.1:0".to_string(), + addr: "127.0.0.1:0".parse().unwrap(), hostnames: Vec::new(), tls_crt: None, tls_key: None, diff --git a/src/main.rs b/src/main.rs index 2b01266..405a9a4 100644 --- a/src/main.rs +++ b/src/main.rs @@ -8,19 +8,20 @@ mod storage; mod web; use std::future::Future; -use std::net::SocketAddr; -use std::path::{Path, PathBuf}; +use std::net::{SocketAddr, TcpListener}; +use std::path::PathBuf; use std::pin::Pin; use std::sync::Arc; use std::task::{Context, Poll}; use axum::{ body::Body, - extract::{self, connect_info::Connected}, + extract::{self, ConnectInfo}, http::{self, Request, Response}, }; +use axum_server::Handle; +use axum_server::tls_rustls::RustlsConfig; use tokio::signal; -use tokio_rustls::TlsAcceptor; use tracing::{error, info}; use tracing_subscriber::{layer::SubscriberExt, registry::LookupSpan, util::SubscriberInitExt}; @@ -44,20 +45,6 @@ struct ListenerInfo { hosts: Vec<String>, } -/// Contains metainfo about one client connection -#[derive(Clone, Copy, Debug)] -struct ClientInfo(SocketAddr); -impl Connected<axum::serve::IncomingStream<'_, TlsListener>> for ClientInfo { - fn connect_info(target: axum::serve::IncomingStream<'_, TlsListener>) -> Self { - ClientInfo(*target.remote_addr()) - } -} -impl Connected<axum::serve::IncomingStream<'_, tokio::net::TcpListener>> for ClientInfo { - fn connect_info(target: axum::serve::IncomingStream<'_, tokio::net::TcpListener>) -> Self { - ClientInfo(*target.remote_addr()) - } -} - /// Request info stored in span extensions for logging #[derive(Clone)] struct RequestInfo { @@ -68,10 +55,9 @@ struct RequestInfo { user_agent: Option<String>, } -/// Key extractor for tower-governor that uses ClientInfo to get the client IP. +/// Key extractor for tower-governor that uses ConnectInfo to get the client IP. #[derive(Clone)] struct ClientIpKeyExtractor; - impl tower_governor::key_extractor::KeyExtractor for ClientIpKeyExtractor { type Key = std::net::IpAddr; @@ -81,8 +67,8 @@ impl tower_governor::key_extractor::KeyExtractor for ClientIpKeyExtractor { fn extract<T>(&self, req: &Request<T>) -> Result<Self::Key, tower_governor::GovernorError> { req.extensions() - .get::<extract::ConnectInfo<ClientInfo>>() - .map(|ci| ci.0.0.ip()) + .get::<extract::ConnectInfo<SocketAddr>>() + .map(|ci| ci.0.ip()) .ok_or(tower_governor::GovernorError::UnableToExtractKey) } } @@ -159,8 +145,8 @@ async fn main() { let local_addr = req.extensions().get::<ListenerInfo>().map(|a| a.addr); let remote_addr = req .extensions() - .get::<extract::ConnectInfo<ClientInfo>>() - .map(|ci| ci.0.0.ip()); + .get::<ConnectInfo<SocketAddr>>() + .map(|ci| ci.0.ip()); tracing::info_span!( "http_request", @@ -174,11 +160,7 @@ async fn main() { method: req.method().clone(), path: req.uri().clone(), version: req.version(), - host: req - .headers() - .get(http::header::HOST) - .and_then(|v| v.to_str().ok()) - .map(String::from), + host: extract_host(req), user_agent: req .headers() .get(http::header::USER_AGENT) @@ -272,59 +254,55 @@ async fn main() { )); let mut tasks = tokio::task::JoinSet::new(); - - // The channel is used to broadcast SIGINT/SIGTERM to all listeners. - let (shutdown_tx, _) = tokio::sync::broadcast::channel::<()>(1); + let handle = Handle::new(); for listener_config in config.listeners() { - let tcp_listener = tokio::net::TcpListener::bind(&listener_config.addr) - .await - .unwrap(); + let std_listener = TcpListener::bind(listener_config.addr).unwrap(); + std_listener.set_nonblocking(true).unwrap(); - let tls_enabled = listener_config.tls_crt.is_some(); - info!( - "listening {} on {} (hostnames: {})", - if tls_enabled { "HTTPS" } else { "HTTP" }, - tcp_listener.local_addr().unwrap(), - if listener_config.hostnames.is_empty() { - "*".to_string() - } else { - listener_config.hostnames.join(", ") - }, - ); + let addr = std_listener.local_addr().unwrap(); - let local_addr = tcp_listener.local_addr().unwrap(); let app = app .clone() .layer(axum::Extension(ListenerInfo { - addr: local_addr, + addr, hosts: listener_config.hostnames.clone(), })) - .into_make_service_with_connect_info::<ClientInfo>(); + .into_make_service_with_connect_info::<SocketAddr>(); - let mut shutdown_rx = shutdown_tx.subscribe(); - let shutdown_signal = async move { - shutdown_rx.recv().await.ok(); - }; + let handle = handle.clone(); - if let (Some(crt_path), Some(key_path)) = - (&listener_config.tls_crt, &listener_config.tls_key) - { - let tls_listener = TlsListener::new(tcp_listener, crt_path, key_path); - tasks.spawn(async move { - axum::serve(tls_listener, app) - .with_graceful_shutdown(shutdown_signal) - .await - .unwrap(); - }); - } else { - tasks.spawn(async move { - axum::serve(tcp_listener, app) - .with_graceful_shutdown(shutdown_signal) + let tls_enabled = + if let (Some(crt), Some(key)) = (&listener_config.tls_crt, &listener_config.tls_key) { + let rustls_config = RustlsConfig::from_pem_file(crt, key) .await - .unwrap(); - }); - } + .expect("failed to load TLS config"); + + tasks.spawn(async move { + let server = axum_server::from_tcp_rustls(std_listener, rustls_config).unwrap(); + server.handle(handle).serve(app).await.unwrap(); + }); + + true + } else { + tasks.spawn(async move { + let server = axum_server::from_tcp(std_listener).unwrap(); + server.handle(handle).serve(app).await.unwrap(); + }); + + false + }; + + info!( + "listening {} on {} (hostnames: {})", + if tls_enabled { "HTTPS" } else { "HTTP" }, + addr, + if listener_config.hostnames.is_empty() { + "*".to_string() + } else { + listener_config.hostnames.join(", ") + }, + ); } // Graceful shutdown @@ -362,7 +340,7 @@ async fn main() { } } - drop(shutdown_tx); // broadcast + handle.graceful_shutdown(None); // Wait for all listeners to finish with timeout let shutdown_result = tokio::time::timeout(config.server().shutdown_timeout, async { @@ -386,59 +364,6 @@ async fn main() { } } -/// A TLS listener that wraps a TCP listener and performs TLS handshakes. -struct TlsListener { - inner: tokio::net::TcpListener, - acceptor: TlsAcceptor, -} -impl TlsListener { - fn new(inner: tokio::net::TcpListener, crt_path: &Path, key_path: &Path) -> Self { - use tokio_rustls::rustls::pki_types::{CertificateDer, PrivateKeyDer, pem::PemObject}; - - let certs: Vec<CertificateDer<'static>> = CertificateDer::pem_file_iter(crt_path) - .expect("failed to open certificate file") - .collect::<Result<_, _>>() - .expect("failed to parse certificates"); - - let key = PrivateKeyDer::from_pem_file(key_path).expect("failed to read private key"); - - let config = tokio_rustls::rustls::ServerConfig::builder() - .with_no_client_auth() - .with_single_cert(certs, key) - .expect("failed to build TLS config"); - - let acceptor = TlsAcceptor::from(Arc::new(config)); - Self { inner, acceptor } - } -} -impl axum::serve::Listener for TlsListener { - type Io = tokio_rustls::server::TlsStream<tokio::net::TcpStream>; - type Addr = std::net::SocketAddr; - - fn local_addr(&self) -> std::io::Result<Self::Addr> { - self.inner.local_addr() - } - - async fn accept(&mut self) -> (Self::Io, Self::Addr) { - loop { - let (stream, addr) = match self.inner.accept().await { - Ok(conn) => conn, - Err(e) => { - error!("failed to accept TCP connection: {}", e); - continue; - } - }; - match self.acceptor.accept(stream).await { - Ok(tls_stream) => return (tls_stream, addr), - Err(e) => { - error!("TLS handshake failed from {}: {}", addr, e); - continue; - } - } - } - } -} - /// Layer that validates the Host header against configured hostnames. #[derive(Clone)] struct HostValidationLayer; @@ -473,25 +398,7 @@ where if let Some(iface) = interface && !iface.hosts.is_empty() { - let host = req - .headers() - .get(http::header::HOST) - .and_then(|v| v.to_str().ok()) - .and_then(|raw| { - let without_port = if let Some((host, port)) = raw.rsplit_once(':') - && port.parse::<u16>().is_ok() - && (host.ends_with(']') || !host.contains('[')) - { - host - } else { - raw - }; - url::Host::parse(without_port) - .ok() - .map(|h| h.to_string().trim_end_matches('.').to_string()) - }); - - let is_valid = host + let is_valid = extract_host(&req) .map(|h| { iface .hosts @@ -515,3 +422,24 @@ where Box::pin(async move { inner.call(req).await }) } } + +fn extract_host(req: &Request<Body>) -> Option<String> { + // HTTP/1.1 uses HOST header, HTTP/2 uses :authority (available via URI) + let raw = if let Some(host) = req.headers().get(http::header::HOST) { + host.to_str().ok().map(|raw| { + if let Some((host, port)) = raw.rsplit_once(':') + && port.parse::<u16>().is_ok() + && (host.ends_with(']') || !host.contains('[')) + { + host + } else { + raw + } + }) + } else { + req.uri().host() + }; + + raw.and_then(|raw| url::Host::parse(raw).ok()) + .map(|h| h.to_string().trim_end_matches('.').to_string()) +} |
