aboutsummaryrefslogtreecommitdiffstats
diff options
context:
space:
mode:
authorNikolay Govorov <me@govorov.online>2026-01-19 07:41:06 +0000
committerNikolay Govorov <me@govorov.online>2026-01-19 07:41:06 +0000
commitf9f02354b882e009dc30dc735f99928623eb213f (patch)
tree595e7fc9a5ffac89be7ac1e533295c7d46c7f598
parent9a636b1910f3adb45e64a09a896cb40cd1d40805 (diff)
downloadtar
tar.gz
tar.bz2
tar.lz
tar.xz
tar.zst
zip
Replace TLS impl
Diffstat
-rw-r--r--Cargo.lock26+25 −1
-rw-r--r--Cargo.toml6+2 −4
-rw-r--r--nfpm.yaml7+4 −3
-rw-r--r--pkg/zorian.toml2+1 −1
-rw-r--r--src/config.rs51+33 −18
-rw-r--r--src/main.rs214+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())
+}