Unified TLS certificate management

This commit is contained in:
mdecimus
2024-03-28 11:12:46 +01:00
parent 37eb6483ca
commit 16b0465933
52 changed files with 1480 additions and 1075 deletions

View File

@@ -55,6 +55,7 @@ unicode-security = "0.1.0"
infer = "0.15.0"
bincode = "1.3.1"
[target.'cfg(unix)'.dependencies]
privdrop = "0.5.3"
tracing-journald = "0.3"

View File

@@ -0,0 +1,436 @@
use std::{
collections::{btree_map::Entry, BTreeMap},
path::PathBuf,
sync::Arc,
};
use arc_swap::ArcSwap;
use store::{
write::{BatchBuilder, ValueClass},
Deserialize, IterateParams, Store, Stores, ValueKey,
};
use tracing_appender::non_blocking::WorkerGuard;
use utils::{
config::{Config, ConfigKey},
failed,
glob::GlobPattern,
UnwrapFailure,
};
use crate::{Core, SharedCore};
use super::{server::Servers, tracers::Tracers};
#[derive(Default)]
pub struct ConfigManager {
cfg_local: ArcSwap<BTreeMap<String, String>>,
cfg_local_path: PathBuf,
cfg_local_patterns: Arc<Patterns>,
cfg_store: Store,
}
#[derive(Default)]
pub struct Patterns {
patterns: Vec<Pattern>,
}
enum Pattern {
Include(MatchType),
Exclude(MatchType),
}
enum MatchType {
Equal(String),
StartsWith(String),
EndsWith(String),
Matches(GlobPattern),
All,
}
pub struct BootManager {
pub config: Config,
pub core: SharedCore,
pub servers: Servers,
pub guards: Option<Vec<WorkerGuard>>,
}
impl BootManager {
pub async fn init() -> Self {
let mut config_path = std::env::var("CONFIG_PATH").ok();
let mut found_param = false;
if config_path.is_none() {
for arg in std::env::args().skip(1) {
if let Some((key, value)) = arg.split_once('=') {
if key.starts_with("--config") {
config_path = value.trim().to_string().into();
break;
} else {
failed(&format!("Invalid command line argument: {key}"));
}
} else if found_param {
config_path = arg.into();
break;
} else if arg.starts_with("--config") {
found_param = true;
} else {
failed(&format!("Invalid command line argument: {arg}"));
}
}
}
// Read main configuration file
let cfg_local_path =
PathBuf::from(config_path.failed("Missing parameter --config=<path-to-config>."));
let mut config = Config::default();
match std::fs::read_to_string(&cfg_local_path) {
Ok(value) => {
config.parse(&value).failed("Invalid configuration file");
}
Err(err) => {
config.new_build_error("*", format!("Could not read configuration file: {err}"));
}
}
let cfg_local = config.keys.clone();
// Resolve macros
config.resolve_macros().await;
// Parse include/exclude patterns
let mut cfg_local_patterns = Vec::new();
for (key, value) in &config.keys {
if !key.starts_with("config.local-keys") {
if cfg_local_patterns.is_empty() {
continue;
} else {
break;
}
};
let value = value.trim();
let (value, is_include) = value
.strip_prefix('!')
.map_or((value, true), |value| (value, false));
let value = value.trim().to_ascii_lowercase();
if value.is_empty() {
continue;
}
let match_type = if value == "*" {
MatchType::All
} else if let Some(value) = value.strip_prefix('*') {
MatchType::StartsWith(value.to_string())
} else if let Some(value) = value.strip_suffix('*') {
MatchType::EndsWith(value.to_string())
} else if value.contains('*') {
MatchType::Matches(GlobPattern::compile(&value, false))
} else {
MatchType::Equal(value.to_string())
};
cfg_local_patterns.push(if is_include {
Pattern::Include(match_type)
} else {
Pattern::Exclude(match_type)
});
}
if cfg_local_patterns.is_empty() {
cfg_local_patterns = vec![
Pattern::Include(MatchType::StartsWith("store.".to_string())),
Pattern::Include(MatchType::StartsWith("server.listener.".to_string())),
Pattern::Include(MatchType::StartsWith("cluster.".to_string())),
Pattern::Include(MatchType::Equal("storage.data".to_string())),
Pattern::Include(MatchType::Equal("storage.blob".to_string())),
Pattern::Include(MatchType::Equal("storage.lookup".to_string())),
Pattern::Include(MatchType::Equal("storage.fts".to_string())),
Pattern::Include(MatchType::Equal("server.run-as.user".to_string())),
Pattern::Include(MatchType::Equal("server.run-as.group".to_string())),
Pattern::Exclude(MatchType::Matches(GlobPattern::compile(
"store.*.query.*",
false,
))),
];
}
// Parser servers
let mut servers = Servers::parse(&mut config);
// Bind ports and drop privileges
servers.bind_and_drop_priv(&mut config);
// Load stores
let mut stores = Stores::parse(&mut config).await;
// Build manager
let manager = ConfigManager {
cfg_local: ArcSwap::from_pointee(cfg_local),
cfg_local_path,
cfg_local_patterns: Patterns {
patterns: cfg_local_patterns,
}
.into(),
cfg_store: config
.value("storage.data")
.and_then(|id| stores.stores.get(id))
.cloned()
.unwrap_or_default(),
};
// Extend configuration with settings stored in the db
if !manager.cfg_store.is_none() {
manager
.extend_config(&mut config)
.await
.failed("Failed to read configuration");
}
// Parse lookup stores
stores.parse_lookups(&mut config).await;
// Parse settings and build shared core
let core = Core::parse(&mut config, stores, manager)
.await
.into_shared();
// Parse TCP acceptors
servers.parse_tcp_acceptors(&mut config, core.clone());
BootManager {
core,
guards: Tracers::parse(&mut config).enable(&mut config),
config,
servers,
}
}
}
impl ConfigManager {
async fn extend_config(&self, config: &mut Config) -> store::Result<()> {
for (key, value) in self.db_list("", false).await? {
config.keys.entry(key).or_insert(value);
}
Ok(())
}
pub async fn get(&self, key: impl AsRef<str>) -> store::Result<Option<String>> {
let key = key.as_ref();
match self.cfg_local.load().get(key) {
Some(value) => Ok(Some(value.to_string())),
None => {
self.cfg_store
.get_value(ValueKey::from(ValueClass::Config(
key.to_string().into_bytes(),
)))
.await
}
}
}
pub async fn list(
&self,
prefix: &str,
strip_prefix: bool,
) -> store::Result<Vec<(String, String)>> {
let mut results = self.db_list(prefix, strip_prefix).await?;
for (key, value) in self.cfg_local.load().iter() {
if !strip_prefix || prefix.is_empty() {
results.push((key.clone(), value.clone()));
} else if key.starts_with(prefix) {
if let Some(key) = key.strip_prefix(prefix) {
results.push((key.to_string(), value.clone()));
}
}
}
Ok(results)
}
async fn db_list(
&self,
prefix: &str,
strip_prefix: bool,
) -> store::Result<Vec<(String, String)>> {
let key = prefix.as_bytes();
let from_key = ValueKey::from(ValueClass::Config(key.to_vec()));
let to_key = ValueKey::from(ValueClass::Config(
key.iter()
.copied()
.chain([u8::MAX, u8::MAX, u8::MAX, u8::MAX, u8::MAX])
.collect::<Vec<_>>(),
));
let mut results = Vec::new();
let patterns = self.cfg_local_patterns.clone();
self.cfg_store
.iterate(
IterateParams::new(from_key, to_key).ascending(),
|key, value| {
let mut key =
std::str::from_utf8(key.get(1..).unwrap_or_default()).map_err(|_| {
store::Error::InternalError(
"Failed to deserialize config key".to_string(),
)
})?;
if !patterns.is_local_key(key) {
if strip_prefix && !prefix.is_empty() {
key = key.strip_prefix(prefix).unwrap_or(key);
}
results.push((key.to_string(), String::deserialize(value)?));
}
Ok(true)
},
)
.await?;
Ok(results)
}
pub async fn set(&self, keys: impl IntoIterator<Item = ConfigKey>) -> store::Result<()> {
let mut batch = BatchBuilder::new();
let mut local_batch = Vec::new();
for key in keys {
if self.cfg_local_patterns.is_local_key(&key.key) {
local_batch.push(key);
} else {
batch.set(ValueClass::Config(key.key.into_bytes()), key.value);
}
}
if !batch.is_empty() {
self.cfg_store.write(batch.build()).await?;
}
if !local_batch.is_empty() {
let mut local = self.cfg_local.load().as_ref().clone();
let mut has_changes = false;
for key in local_batch {
match local.entry(key.key) {
Entry::Vacant(v) => {
v.insert(key.value);
has_changes = true;
}
Entry::Occupied(mut v) => {
if v.get() != &key.value {
v.insert(key.value);
has_changes = true;
}
}
}
}
if has_changes {
self.update_local(local).await?;
}
}
Ok(())
}
pub async fn clear(&self, key: impl AsRef<str>) -> store::Result<()> {
let key = key.as_ref();
if self.cfg_local_patterns.is_local_key(key) {
let mut local = self.cfg_local.load().as_ref().clone();
if local.remove(key).is_some() {
self.update_local(local).await
} else {
Ok(())
}
} else {
let mut batch = BatchBuilder::new();
batch.clear(ValueClass::Config(key.to_string().into_bytes()));
self.cfg_store.write(batch.build()).await.map(|_| ())
}
}
pub async fn clear_prefix(&self, key: impl AsRef<str>) -> store::Result<()> {
let key = key.as_ref();
// Delete local keys
let local = self.cfg_local.load();
if local.keys().any(|k| k.starts_with(key)) {
let mut local = local.as_ref().clone();
local.retain(|k, _| !k.starts_with(key));
self.update_local(local).await?;
}
// Delete db keys
self.cfg_store
.delete_range(
ValueKey::from(ValueClass::Config(key.as_bytes().to_vec())),
ValueKey::from(ValueClass::Config(
key.as_bytes()
.iter()
.copied()
.chain([u8::MAX, u8::MAX, u8::MAX, u8::MAX, u8::MAX])
.collect::<Vec<_>>(),
)),
)
.await
}
async fn update_local(&self, map: BTreeMap<String, String>) -> store::Result<()> {
let mut cfg_text = String::with_capacity(1024);
for (key, value) in &map {
cfg_text.push_str(key);
cfg_text.push_str(" = ");
if value == "true" || value == "false" || value.parse::<f64>().is_ok() {
cfg_text.push_str(value);
} else {
cfg_text.push('"');
cfg_text.push_str(&value.replace('"', "\\\""));
cfg_text.push('"');
}
cfg_text.push_str(value);
cfg_text.push('\n');
}
self.cfg_local.store(map.into());
tokio::fs::write(&self.cfg_local_path, cfg_text)
.await
.map_err(|err| {
store::Error::InternalError(format!(
"Failed to write local configuration file: {err}"
))
})
}
}
impl Patterns {
pub fn is_local_key(&self, key: &str) -> bool {
let mut is_local = false;
for pattern in &self.patterns {
match pattern {
Pattern::Include(pattern) => {
if !is_local && pattern.matches(key) {
is_local = true;
}
}
Pattern::Exclude(pattern) => {
if pattern.matches(key) {
return false;
}
}
}
}
is_local
}
}
impl MatchType {
fn matches(&self, value: &str) -> bool {
match self {
MatchType::Equal(pattern) => value == pattern,
MatchType::StartsWith(pattern) => value.starts_with(pattern),
MatchType::EndsWith(pattern) => value.ends_with(pattern),
MatchType::Matches(pattern) => pattern.matches(value),
MatchType::All => true,
}
}
}

View File

@@ -5,15 +5,16 @@ use directory::{Directories, Directory};
use store::{BlobBackend, BlobStore, FtsStore, LookupStore, Store, Stores};
use utils::config::Config;
use crate::{Core, Network};
use crate::{listener::tls::TlsManager, Core, Network};
use self::{
imap::ImapConfig, jmap::settings::JmapConfig, scripts::Scripting, smtp::SmtpConfig,
storage::Storage,
imap::ImapConfig, jmap::settings::JmapConfig, manager::ConfigManager, scripts::Scripting,
smtp::SmtpConfig, storage::Storage,
};
pub mod imap;
pub mod jmap;
pub mod manager;
pub mod network;
pub mod scripts;
pub mod server;
@@ -22,7 +23,7 @@ pub mod storage;
pub mod tracers;
impl Core {
pub async fn parse(config: &mut Config, stores: Stores) -> Self {
pub async fn parse(config: &mut Config, stores: Stores, config_manager: ConfigManager) -> Self {
let mut data = config
.value_require("storage.data")
.map(|id| id.to_string())
@@ -116,6 +117,7 @@ impl Core {
smtp: SmtpConfig::parse(config).await,
jmap: JmapConfig::parse(config),
imap: ImapConfig::parse(config),
tls: TlsManager::parse(config),
storage: Storage {
data,
blob,
@@ -125,6 +127,7 @@ impl Core {
directory,
directories: directories.directories,
purge_schedules: stores.purge_schedules,
config: config_manager,
},
}
}

View File

@@ -25,7 +25,6 @@ use std::{net::SocketAddr, sync::Arc};
use rustls::{
crypto::ring::{default_provider, ALL_CIPHER_SUITES},
server::ResolvesServerCert,
ServerConfig, SupportedCipherSuite, ALL_VERSIONS,
};
@@ -36,7 +35,14 @@ use utils::config::{
Config,
};
use crate::listener::{acme::directory::ACME_TLS_ALPN_NAME, tls::CertificateResolver, TcpAcceptor};
use crate::{
listener::{
acme::{directory::ACME_TLS_ALPN_NAME, AcmeResolver},
tls::CertificateResolver,
TcpAcceptor,
},
SharedCore,
};
use super::{
tls::{TLS12_VERSION, TLS13_VERSION},
@@ -45,17 +51,15 @@ use super::{
impl Servers {
pub fn parse(config: &mut Config) -> Self {
// Parse certificates and ACME managers
// Parse ACME managers
let mut servers = Servers::default();
servers.parse_certificates(config);
servers.parse_acmes(config);
// Parse servers
let ids = config
for id in config
.sub_keys("server.listener", ".protocol")
.map(|s| s.to_string())
.collect::<Vec<_>>();
for id in ids {
.collect::<Vec<_>>()
{
servers.parse_server(config, id);
}
servers
@@ -170,154 +174,6 @@ impl Servers {
return;
}
// Build TLS config
let (acceptor, tls_implicit) = if config
.property_or_else(("server.listener", id, "tls.enable"), "server.tls.enable")
.unwrap_or(false)
{
// Parse protocol versions
let mut tls_v2 = true;
let mut tls_v3 = true;
let mut proto_err = None;
for (_, protocol) in config.values_or_else(
("server.listener", id, "tls.disable-protocols"),
"server.tls.disable-protocols",
) {
match protocol {
"TLSv1.2" | "0x0303" => tls_v2 = false,
"TLSv1.3" | "0x0304" => tls_v3 = false,
protocol => {
proto_err = format!("Unsupported TLS protocol {protocol:?}").into();
}
}
}
if let Some(proto_err) = proto_err {
config.new_parse_error(("server.listener", id, "tls.disable-protocols"), proto_err);
}
// Parse cipher suites
let mut disabled_ciphers: Vec<SupportedCipherSuite> = Vec::new();
let cipher_keys = if config.has_prefix(("server.listener", id, "tls.disable-ciphers")) {
("server.listener", id, "tls.disable-ciphers").as_key()
} else {
"server.tls.disable-ciphers".as_key()
};
for (_, protocol) in config.properties::<SupportedCipherSuite>(cipher_keys) {
disabled_ciphers.push(protocol);
}
// Build resolver
let mut acme_acceptor = None;
let resolver: Arc<dyn ResolvesServerCert> = if let Some(acme_id) =
config.value_or_else(("server.listener", id, "tls.acme"), "server.tls.acme")
{
let acme = if let Some(acme) = self.acme_managers.get(acme_id) {
acme
} else {
config.new_parse_error(
("server.listener", id, "tls.acme"),
format!("Undefined ACME manager id {acme_id:?}"),
);
return;
};
// Check if this port is used to receive ACME challenges
let port_key = ("acme", acme_id, "port").as_key();
let acme_port = config
.property_or_default::<u16>(port_key, "443")
.unwrap_or(443);
if listeners.iter().any(|l| l.addr.port() == acme_port) {
acme_acceptor = Some(acme.clone());
}
acme.clone()
} else if let Some(cert) = config
.value_or_else(
("server.listener", id, "tls.certificate"),
"server.tls.certificate",
)
.and_then(|cert_id| self.certificates.get(cert_id))
.cloned()
{
Arc::new(CertificateResolver {
sni: self.certificates_sni.clone(),
cert,
})
} else {
config.new_parse_error(
("server.listener", id, "tls.certificate"),
"Undefined certificate id",
);
return;
};
// Build cert provider
let mut provider = default_provider();
if !disabled_ciphers.is_empty() {
provider.cipher_suites = ALL_CIPHER_SUITES
.iter()
.filter(|suite| !disabled_ciphers.contains(suite))
.copied()
.collect();
}
// Build server config
let mut server_config = match ServerConfig::builder_with_provider(provider.into())
.with_protocol_versions(if tls_v3 == tls_v2 {
ALL_VERSIONS
} else if tls_v3 {
TLS13_VERSION
} else {
TLS12_VERSION
}) {
Ok(server_config) => server_config
.with_no_client_auth()
.with_cert_resolver(resolver.clone()),
Err(err) => {
config.new_build_error(
("server.listener", id, "tls"),
format!("Failed to build TLS server config: {err}"),
);
return;
}
};
server_config.ignore_client_order = config
.property_or_else(
("server.listener", id, "tls.ignore-client-order"),
"server.tls.ignore-client-order",
)
.unwrap_or(true);
// Build acceptor
let acceptor = if let Some(manager) = acme_acceptor {
let mut challenge = ServerConfig::builder()
.with_no_client_auth()
.with_cert_resolver(resolver);
challenge.alpn_protocols.push(ACME_TLS_ALPN_NAME.to_vec());
TcpAcceptor::Acme {
challenge: Arc::new(challenge),
default: Arc::new(server_config),
manager,
}
} else {
TcpAcceptor::Tls(TlsAcceptor::from(Arc::new(server_config)))
};
(
acceptor,
config
.property_or_else(
("server.listener", id, "tls.implicit"),
"server.tls.implicit",
)
.unwrap_or(true),
)
} else {
(TcpAcceptor::Plain, false)
};
// Parse proxy networks
let mut proxy_networks = Vec::new();
let proxy_keys = if config.has_prefix(("server.listener", id, "proxy.trusted-networks")) {
@@ -339,11 +195,126 @@ impl Servers {
id: id_,
protocol,
listeners,
acceptor,
tls_implicit,
proxy_networks,
});
}
pub fn parse_tcp_acceptors(&mut self, config: &mut Config, core: SharedCore) {
let resolver = Arc::new(CertificateResolver::new(core.clone()));
let acme_config = {
let mut challenge = ServerConfig::builder()
.with_no_client_auth()
.with_cert_resolver(Arc::new(AcmeResolver::new(core)));
challenge.alpn_protocols.push(ACME_TLS_ALPN_NAME.to_vec());
Arc::new(challenge)
};
for id_ in config
.sub_keys("server.listener", ".protocol")
.map(|s| s.to_string())
.collect::<Vec<_>>()
{
let id = id_.as_str();
// Build TLS config
let acceptor = if config
.property_or_else(("server.listener", id, "tls.enable"), "server.tls.enable")
.unwrap_or(false)
{
// Parse protocol versions
let mut tls_v2 = true;
let mut tls_v3 = true;
let mut proto_err = None;
for (_, protocol) in config.values_or_else(
("server.listener", id, "tls.disable-protocols"),
"server.tls.disable-protocols",
) {
match protocol {
"TLSv1.2" | "0x0303" => tls_v2 = false,
"TLSv1.3" | "0x0304" => tls_v3 = false,
protocol => {
proto_err = format!("Unsupported TLS protocol {protocol:?}").into();
}
}
}
if let Some(proto_err) = proto_err {
config.new_parse_error(
("server.listener", id, "tls.disable-protocols"),
proto_err,
);
}
// Parse cipher suites
let mut disabled_ciphers: Vec<SupportedCipherSuite> = Vec::new();
let cipher_keys =
if config.has_prefix(("server.listener", id, "tls.disable-ciphers")) {
("server.listener", id, "tls.disable-ciphers").as_key()
} else {
"server.tls.disable-ciphers".as_key()
};
for (_, protocol) in config.properties::<SupportedCipherSuite>(cipher_keys) {
disabled_ciphers.push(protocol);
}
// Build cert provider
let mut provider = default_provider();
if !disabled_ciphers.is_empty() {
provider.cipher_suites = ALL_CIPHER_SUITES
.iter()
.filter(|suite| !disabled_ciphers.contains(suite))
.copied()
.collect();
}
// Build server config
let mut server_config = match ServerConfig::builder_with_provider(provider.into())
.with_protocol_versions(if tls_v3 == tls_v2 {
ALL_VERSIONS
} else if tls_v3 {
TLS13_VERSION
} else {
TLS12_VERSION
}) {
Ok(server_config) => server_config
.with_no_client_auth()
.with_cert_resolver(resolver.clone()),
Err(err) => {
config.new_build_error(
("server.listener", id, "tls"),
format!("Failed to build TLS server config: {err}"),
);
return;
}
};
server_config.ignore_client_order = config
.property_or_else(
("server.listener", id, "tls.ignore-client-order"),
"server.tls.ignore-client-order",
)
.unwrap_or(true);
// Build acceptor
let default_config = Arc::new(server_config);
TcpAcceptor::Tls {
acceptor: TlsAcceptor::from(default_config.clone()),
acme_config: acme_config.clone(),
default_config,
implicit: config
.property_or_else(
("server.listener", id, "tls.implicit"),
"server.tls.implicit",
)
.unwrap_or(true),
}
} else {
TcpAcceptor::Plain
};
self.tcp_acceptors.insert(id_, acceptor);
}
}
}
impl ParseValue for ServerProtocol {

View File

@@ -1,10 +1,10 @@
use std::{fmt::Display, net::SocketAddr, sync::Arc, time::Duration};
use std::{fmt::Display, net::SocketAddr, time::Duration};
use ahash::AHashMap;
use tokio::net::TcpSocket;
use utils::config::ipmask::IpAddrMask;
use crate::listener::{acme::AcmeManager, tls::Certificate, TcpAcceptor};
use crate::listener::TcpAcceptor;
pub mod listener;
pub mod tls;
@@ -12,9 +12,7 @@ pub mod tls;
#[derive(Default)]
pub struct Servers {
pub servers: Vec<Server>,
pub certificates: AHashMap<String, Arc<Certificate>>,
pub certificates_sni: AHashMap<String, Arc<Certificate>>,
pub acme_managers: AHashMap<String, Arc<AcmeManager>>,
pub tcp_acceptors: AHashMap<String, TcpAcceptor>,
}
#[derive(Debug, Default)]
@@ -23,8 +21,6 @@ pub struct Server {
pub protocol: ServerProtocol,
pub listeners: Vec<Listener>,
pub proxy_networks: Vec<IpAddrMask>,
pub acceptor: TcpAcceptor,
pub tls_implicit: bool,
pub max_connections: u64,
}

View File

@@ -21,39 +21,51 @@
* for more details.
*/
use std::{io::Cursor, sync::Arc, time::Duration};
use std::{
io::Cursor,
net::{Ipv4Addr, Ipv6Addr},
sync::Arc,
time::Duration,
};
use ahash::{AHashMap, AHashSet};
use arc_swap::ArcSwap;
use rcgen::generate_simple_self_signed;
use rustls::{
client::verify_server_name,
crypto::ring::sign::any_supported_type,
server::ParsedCertificate,
sign::CertifiedKey,
version::{TLS12, TLS13},
Error, SupportedProtocolVersion,
SupportedProtocolVersion,
};
use rustls_pemfile::{certs, read_one, Item};
use rustls_pki_types::{DnsName, PrivateKeyDer, ServerName};
use rustls_pki_types::PrivateKeyDer;
use utils::config::Config;
use crate::listener::{
acme::{directory::LETS_ENCRYPT_PRODUCTION_DIRECTORY, AcmeManager},
tls::Certificate,
use x509_parser::{
certificate::X509Certificate,
der_parser::asn1_rs::FromDer,
extensions::{GeneralName, ParsedExtension},
};
use super::Servers;
use crate::listener::{
acme::{directory::LETS_ENCRYPT_PRODUCTION_DIRECTORY, AcmeProvider},
tls::TlsManager,
};
pub static TLS13_VERSION: &[&SupportedProtocolVersion] = &[&TLS13];
pub static TLS12_VERSION: &[&SupportedProtocolVersion] = &[&TLS12];
impl Servers {
pub fn parse_certificates(&mut self, config: &mut Config) {
let cert_ids = config
impl TlsManager {
pub fn parse(config: &mut Config) -> Self {
let mut certificates = AHashMap::new();
let mut acme_providers = AHashMap::new();
let mut subject_names = AHashSet::new();
// Parse certificates
for cert_id in config
.sub_keys("certificate", ".cert")
.map(|s| s.to_string())
.collect::<Vec<_>>();
for cert_id in cert_ids {
.collect::<Vec<_>>()
{
let cert_id = cert_id.as_str();
let key_cert = ("certificate", cert_id, "cert");
let key_pk = ("certificate", cert_id, "private-key");
@@ -66,53 +78,89 @@ impl Servers {
if let (Some(cert), Some(pk)) = (cert, pk) {
match build_certified_key(cert, pk) {
Ok(cert) => {
// Parse alternative names
let subjects = config
.values(("certificate", cert_id, "sni-subjects"))
.map(|(_, v)| v.to_string())
.collect::<Vec<_>>();
let mut sni_names = Vec::new();
for subject in subjects {
match DnsName::try_from(subject)
.map_err(|_| Error::General("Bad DNS name".into()))
.map(|name| ServerName::DnsName(name.to_lowercase_owned()))
.and_then(|name| {
cert.end_entity_cert()
.and_then(ParsedCertificate::try_from)
.and_then(|cert| verify_server_name(&cert, &name))
.map(|_| name)
}) {
Ok(ServerName::DnsName(server_name)) => {
sni_names.push(server_name.as_ref().to_string());
match cert
.end_entity_cert()
.map_err(|err| format!("Failed to obtain end entity cert: {err}"))
.and_then(|cert| {
X509Certificate::from_der(cert.as_ref()).map_err(|err| {
format!("Failed to parse end entity cert: {err}")
})
}) {
Ok((_, parsed)) => {
// Add CNs and SANs to the list of names
let mut names = AHashSet::new();
for name in parsed.subject().iter_common_name() {
if let Ok(name) = name.as_str() {
names.insert(name.to_string());
}
}
Ok(_) => {}
Err(err) => {
config.new_parse_error(
("certificate", cert_id, "sni-subjects"),
err.to_string(),
for ext in parsed.extensions() {
if let ParsedExtension::SubjectAlternativeName(san) =
ext.parsed_extension()
{
for name in &san.general_names {
let name = match name {
GeneralName::DNSName(name) => name.to_string(),
GeneralName::IPAddress(ip) => match ip.len() {
4 => Ipv4Addr::from(
<[u8; 4]>::try_from(*ip).unwrap(),
)
.to_string(),
16 => Ipv6Addr::from(
<[u8; 16]>::try_from(*ip).unwrap(),
)
.to_string(),
_ => continue,
},
_ => {
continue;
}
};
names.insert(name);
}
}
}
// Add custom SNIs
names.extend(
config
.values(("certificate", cert_id, "subjects"))
.map(|(_, v)| v.trim().to_string()),
);
// Add domain names
subject_names.extend(names.iter().cloned());
// Add certificates
let cert = Arc::new(cert);
for name in names {
certificates.insert(
name.strip_prefix("*.")
.map(|name| name.to_string())
.unwrap_or(name),
cert.clone(),
);
}
// Add default certificate
if config
.property::<bool>(("certificate", cert_id, "default"))
.unwrap_or_default()
{
certificates.insert("*".to_string(), cert.clone());
}
}
Err(err) => {
config.new_build_error(format!("certificate.{cert_id}"), err)
}
}
let cert = Arc::new(Certificate {
cert: ArcSwap::from(Arc::new(cert)),
cert_id: cert_id.to_string(),
});
for sni_name in sni_names {
self.certificates_sni.insert(sni_name, cert.clone());
}
self.certificates.insert(cert_id.to_string(), cert);
}
Err(err) => config.new_build_error(format!("certificate.{cert_id}"), err),
}
}
}
}
pub fn parse_acmes(&mut self, config: &mut Config) {
// Parse ACME providers
for acme_id in config
.sub_keys("acme", ".directory")
.map(|s| s.to_string())
@@ -154,17 +202,19 @@ impl Servers {
.map(|(_, v)| v.to_string())
.collect::<Vec<_>>();
// Add domains for self-signed certificate
subject_names.extend(domains.iter().cloned());
if !domains.is_empty() {
match AcmeManager::new(
match AcmeProvider::new(
acme_id.to_string(),
directory,
domains,
contact,
renew_before,
) {
Ok(acme_manager) => {
self.acme_managers
.insert(acme_id.to_string(), Arc::new(acme_manager));
Ok(acme_provider) => {
acme_providers.insert(acme_id.to_string(), acme_provider);
}
Err(err) => {
config.new_build_error(format!("acme.{acme_id}"), err);
@@ -172,6 +222,24 @@ impl Servers {
}
}
}
if subject_names.is_empty() {
subject_names.insert("localhost".to_string());
}
TlsManager {
certificates: ArcSwap::from_pointee(certificates),
acme_providers,
acme_auth_keys: Default::default(),
acme_in_progress: false.into(),
self_signed_cert: build_self_signed_cert(subject_names.into_iter().collect::<Vec<_>>())
.or_else(|err| {
config.new_build_error("certificate.self-signed", err);
build_self_signed_cert(vec!["localhost".to_string()])
})
.ok()
.map(Arc::new),
}
}
}
@@ -205,13 +273,11 @@ pub(crate) fn build_certified_key(
})
}
pub(crate) fn build_self_signed_cert(domains: &[String]) -> utils::config::Result<CertifiedKey> {
let cert = generate_simple_self_signed(domains).map_err(|err| {
format!(
"Failed to generate self-signed certificate for {domains:?}: {err}",
domains = domains
)
})?;
pub(crate) fn build_self_signed_cert(
domains: impl Into<Vec<String>>,
) -> utils::config::Result<CertifiedKey> {
let cert = generate_simple_self_signed(domains)
.map_err(|err| format!("Failed to generate self-signed certificate: {err}",))?;
build_certified_key(
cert.serialize_pem().unwrap().into_bytes(),
cert.serialize_private_key_pem().into_bytes(),

View File

@@ -4,6 +4,8 @@ use ahash::AHashMap;
use directory::Directory;
use store::{write::purge::PurgeSchedule, BlobStore, FtsStore, LookupStore, Store};
use super::manager::ConfigManager;
#[derive(Default)]
pub struct Storage {
pub data: Store,
@@ -14,4 +16,5 @@ pub struct Storage {
pub directory: Arc<Directory>,
pub directories: AHashMap<String, Arc<Directory>>,
pub purge_schedules: Vec<PurgeSchedule>,
pub config: ConfigManager,
}

View File

@@ -15,7 +15,7 @@ use config::{
};
use directory::{Directory, Principal, QueryBy};
use expr::if_block::IfBlock;
use listener::blocked::BlockedIps;
use listener::{blocked::BlockedIps, tls::TlsManager};
use mail_send::Credentials;
use opentelemetry::KeyValue;
use opentelemetry_sdk::{
@@ -47,6 +47,7 @@ pub struct Core {
pub storage: Storage,
pub sieve: Scripting,
pub network: Network,
pub tls: TlsManager,
pub smtp: SmtpConfig,
pub jmap: JmapConfig,
pub imap: ImapConfig,

View File

@@ -24,45 +24,66 @@
use std::io::ErrorKind;
use base64::{engine::general_purpose::URL_SAFE_NO_PAD, Engine};
use ring::digest::{Context, SHA512};
use utils::config::ConfigKey;
use super::{AcmeError, AcmeManager};
use crate::Core;
impl AcmeManager {
pub(crate) async fn load_cert(&self) -> Result<Option<Vec<u8>>, AcmeError> {
self.read_if_exists("private-key", self.domains.as_slice())
use super::{AcmeError, AcmeProvider};
impl Core {
pub(crate) async fn load_cert(
&self,
provider: &AcmeProvider,
) -> Result<Option<Vec<u8>>, AcmeError> {
self.read_if_exists(provider, "cert", provider.domains.as_slice())
.await
.map_err(AcmeError::CertCacheLoad)
}
pub(crate) async fn store_cert(&self, cert: &[u8]) -> Result<(), AcmeError> {
self.write("private-key", self.domains.as_slice(), cert)
pub(crate) async fn store_cert(
&self,
provider: &AcmeProvider,
cert: &[u8],
) -> Result<(), AcmeError> {
self.write(provider, "cert", provider.domains.as_slice(), cert)
.await
.map_err(AcmeError::CertCacheStore)
}
pub(crate) async fn load_account(&self) -> Result<Option<Vec<u8>>, AcmeError> {
self.read_if_exists("cert", self.contact.as_slice())
pub(crate) async fn load_account(
&self,
provider: &AcmeProvider,
) -> Result<Option<Vec<u8>>, AcmeError> {
self.read_if_exists(provider, "account-key", provider.contact.as_slice())
.await
.map_err(AcmeError::AccountCacheLoad)
}
pub(crate) async fn store_account(&self, account: &[u8]) -> Result<(), AcmeError> {
self.write("cert", self.contact.as_slice(), account)
.await
.map_err(AcmeError::AccountCacheStore)
pub(crate) async fn store_account(
&self,
provider: &AcmeProvider,
account: &[u8],
) -> Result<(), AcmeError> {
self.write(
provider,
"account-key",
provider.contact.as_slice(),
account,
)
.await
.map_err(AcmeError::AccountCacheStore)
}
async fn read_if_exists(
&self,
provider: &AcmeProvider,
class: &str,
items: &[String],
) -> Result<Option<Vec<u8>>, std::io::Error> {
match self
.store
.load()
.config_get(self.build_key(class, items))
.storage
.config
.get(self.build_key(provider, class, items))
.await
{
Ok(Some(content)) => match URL_SAFE_NO_PAD.decode(content.as_bytes()) {
@@ -76,33 +97,36 @@ impl AcmeManager {
async fn write(
&self,
provider: &AcmeProvider,
class: &str,
items: &[String],
contents: impl AsRef<[u8]>,
) -> Result<(), std::io::Error> {
self.store
.load()
.config_set([ConfigKey {
key: self.build_key(class, items),
self.storage
.config
.set([ConfigKey {
key: self.build_key(provider, class, items),
value: URL_SAFE_NO_PAD.encode(contents.as_ref()),
}])
.await
.map_err(|err| std::io::Error::new(ErrorKind::Other, err))
}
fn build_key(&self, class: &str, items: &[String]) -> String {
let mut ctx = Context::new(&SHA512);
fn build_key(&self, provider: &AcmeProvider, class: &str, _: &[String]) -> String {
/*let mut ctx = Context::new(&SHA512);
for el in items {
ctx.update(el.as_ref());
ctx.update(&[0])
}
ctx.update(self.directory_url.as_bytes());
ctx.update(provider.directory_url.as_bytes());
format!(
"certificate.acme-{}-{}.{}",
self.id,
provider.id,
URL_SAFE_NO_PAD.encode(ctx.finish()),
class
)
)*/
format!("acme.{}.{}", provider.id, class)
}
}

View File

@@ -29,38 +29,30 @@ pub mod resolver;
use std::{
fmt::Debug,
sync::{
atomic::{AtomicBool, Ordering},
Arc,
},
sync::{atomic::Ordering, Arc},
time::Duration,
};
use ahash::AHashMap;
use arc_swap::ArcSwap;
use parking_lot::Mutex;
use rustls::sign::CertifiedKey;
use store::Store;
use tokio::sync::watch;
use crate::config::server::tls::build_self_signed_cert;
use crate::{Core, SharedCore};
use self::{
directory::Account,
order::{CertParseError, OrderError},
};
pub struct AcmeManager {
id: String,
pub(crate) directory_url: String,
pub(crate) domains: Vec<String>,
contact: Vec<String>,
pub struct AcmeProvider {
pub id: String,
pub directory_url: String,
pub domains: Vec<String>,
pub contact: Vec<String>,
renew_before: chrono::Duration,
store: ArcSwap<Store>,
account_key: ArcSwap<Vec<u8>>,
auth_keys: Mutex<AHashMap<String, Arc<CertifiedKey>>>,
order_in_progress: AtomicBool,
cert: ArcSwap<CertifiedKey>,
}
pub struct AcmeResolver {
pub core: SharedCore,
}
#[derive(Debug)]
@@ -74,7 +66,7 @@ pub enum AcmeError {
NewCertParse(CertParseError),
}
impl AcmeManager {
impl AcmeProvider {
pub fn new(
id: String,
directory_url: String,
@@ -82,7 +74,7 @@ impl AcmeManager {
contact: Vec<String>,
renew_before: Duration,
) -> utils::config::Result<Self> {
Ok(AcmeManager {
Ok(AcmeProvider {
id,
directory_url,
contact: contact
@@ -96,115 +88,44 @@ impl AcmeManager {
})
.collect(),
renew_before: chrono::Duration::from_std(renew_before).unwrap(),
store: ArcSwap::from_pointee(Store::None),
account_key: ArcSwap::from_pointee(Vec::new()),
auth_keys: Mutex::new(AHashMap::new()),
order_in_progress: false.into(),
cert: ArcSwap::from_pointee(build_self_signed_cert(&domains)?),
domains,
account_key: Default::default(),
})
}
}
pub async fn init(&self, store: Store) -> Result<Duration, AcmeError> {
// Update data store
self.store.store(Arc::new(store));
impl Core {
pub async fn init_acme(&self, provider: &AcmeProvider) -> Result<Duration, AcmeError> {
// Load account key from cache or generate a new one
if let Some(account_key) = self.load_account().await? {
self.account_key.store(Arc::new(account_key));
if let Some(account_key) = self.load_account(provider).await? {
provider.account_key.store(Arc::new(account_key));
} else {
let account_key = Account::generate_key_pair();
self.store_account(&account_key).await?;
self.account_key.store(Arc::new(account_key));
self.store_account(provider, &account_key).await?;
provider.account_key.store(Arc::new(account_key));
}
// Load certificate from cache or request a new one
Ok(if let Some(pem) = self.load_cert().await? {
self.process_cert(pem, true).await?
Ok(if let Some(pem) = self.load_cert(provider).await? {
self.process_cert(provider, pem, true).await?
} else {
Duration::from_millis(1000)
})
}
pub fn has_order_in_progress(&self) -> bool {
self.order_in_progress.load(Ordering::Relaxed)
pub fn has_acme_order_in_progress(&self) -> bool {
self.tls.acme_in_progress.load(Ordering::Relaxed)
}
}
pub trait SpawnAcme {
fn spawn(self, store: Store, shutdown_rx: watch::Receiver<bool>);
}
impl SpawnAcme for Arc<AcmeManager> {
fn spawn(self, store: Store, mut shutdown_rx: watch::Receiver<bool>) {
tokio::spawn(async move {
let acme = self;
let mut renew_at = match acme.init(store).await {
Ok(renew_at) => renew_at,
Err(err) => {
tracing::error!(
context = "acme",
event = "error",
error = ?err,
"Failed to initialize ACME certificate manager.");
return;
}
};
loop {
tokio::select! {
_ = tokio::time::sleep(renew_at) => {
tracing::info!(
context = "acme",
event = "order",
domains = ?acme.domains,
"Ordering certificates.");
match acme.renew().await {
Ok(renew_at_) => {
renew_at = renew_at_;
tracing::info!(
context = "acme",
event = "success",
domains = ?acme.domains,
next_renewal = ?renew_at,
"Certificates renewed.");
},
Err(err) => {
tracing::error!(
context = "acme",
event = "error",
error = ?err,
"Failed to renew certificates.");
renew_at = Duration::from_secs(3600);
},
}
},
_ = shutdown_rx.changed() => {
tracing::debug!(
context = "acme",
event = "shutdown",
domains = ?acme.domains,
"ACME certificate manager shutting down.");
break;
}
};
}
});
impl AcmeResolver {
pub fn new(core: SharedCore) -> Self {
Self { core }
}
}
impl Debug for AcmeManager {
impl Debug for AcmeResolver {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("AcmeManager")
.field("directory_url", &self.directory_url)
.field("domains", &self.domains)
.field("contact", &self.contact)
.field("account_key", &self.account_key)
.finish()
f.debug_struct("AcmeResolver").finish()
}
}

View File

@@ -13,10 +13,11 @@ use std::time::Duration;
use x509_parser::parse_x509_certificate;
use crate::listener::acme::directory::Identifier;
use crate::Core;
use super::directory::{Account, Auth, AuthStatus, Directory, DirectoryError, Order, OrderStatus};
use super::jose::JoseError;
use super::{AcmeError, AcmeManager};
use super::{AcmeError, AcmeProvider};
#[derive(Debug)]
pub enum OrderError {
@@ -36,9 +37,10 @@ pub enum CertParseError {
InvalidPrivateKey,
}
impl AcmeManager {
impl Core {
pub(crate) async fn process_cert(
&self,
provider: &AcmeProvider,
pem: Vec<u8>,
cached: bool,
) -> Result<Duration, AcmeError> {
@@ -52,13 +54,13 @@ impl AcmeManager {
}
};
self.set_cert(Arc::new(cert));
self.set_cert(provider, Arc::new(cert));
let renew_at = (validity[1] - self.renew_before - Utc::now())
let renew_at = (validity[1] - provider.renew_before - Utc::now())
.max(chrono::Duration::zero())
.to_std()
.unwrap_or_default();
let renewal_date = validity[1] - self.renew_before;
let renewal_date = validity[1] - provider.renew_before;
tracing::info!(
context = "acme",
@@ -66,27 +68,27 @@ impl AcmeManager {
valid_not_before = %validity[0],
valid_not_after = %validity[1],
renewal_date = ?renewal_date,
domains = ?self.domains,
"Loaded certificate for domains {:?}", self.domains);
domains = ?provider.domains,
"Loaded certificate for domains {:?}", provider.domains);
if !cached {
self.store_cert(&pem).await?;
self.store_cert(provider, &pem).await?;
}
Ok(renew_at)
}
pub async fn renew(&self) -> Result<Duration, AcmeError> {
pub async fn renew(&self, provider: &AcmeProvider) -> Result<Duration, AcmeError> {
let mut backoff = 0;
self.order_in_progress.store(true, Ordering::Relaxed);
self.tls.acme_in_progress.store(true, Ordering::Relaxed);
loop {
match self.order().await {
Ok(pem) => return self.process_cert(pem, false).await,
match self.order(provider).await {
Ok(pem) => return self.process_cert(provider, pem, false).await,
Err(err) if backoff < 16 => {
tracing::debug!(
context = "acme",
event = "renew-backoff",
domains = ?self.domains,
domains = ?provider.domains,
attempt = backoff,
reason = ?err,
"Failed to renew certificate, backing off for {} seconds",
@@ -99,33 +101,33 @@ impl AcmeManager {
}
}
async fn order(&self) -> Result<Vec<u8>, OrderError> {
let directory = Directory::discover(&self.directory_url).await?;
async fn order(&self, provider: &AcmeProvider) -> Result<Vec<u8>, OrderError> {
let directory = Directory::discover(&provider.directory_url).await?;
let account = Account::create_with_keypair(
directory,
&self.contact,
self.account_key.load().as_slice(),
&provider.contact,
provider.account_key.load().as_slice(),
)
.await?;
let mut params = CertificateParams::new(self.domains.clone());
let mut params = CertificateParams::new(provider.domains.clone());
params.distinguished_name = DistinguishedName::new();
params.alg = &PKCS_ECDSA_P256_SHA256;
let cert = rcgen::Certificate::from_params(params)?;
let (order_url, mut order) = account.new_order(self.domains.clone()).await?;
let (order_url, mut order) = account.new_order(provider.domains.clone()).await?;
loop {
match order.status {
OrderStatus::Pending => {
let auth_futures = order
.authorizations
.iter()
.map(|url| self.authorize(&account, url));
.map(|url| self.authorize(provider, &account, url));
try_join_all(auth_futures).await?;
tracing::info!(
context = "acme",
event = "auth-complete",
domains = ?self.domains.as_slice(),
domains = ?provider.domains.as_slice(),
"Completed all authorizations"
);
order = account.order(&order_url).await?;
@@ -135,7 +137,7 @@ impl AcmeManager {
tracing::info!(
context = "acme",
event = "processing",
domains = ?self.domains.as_slice(),
domains = ?provider.domains.as_slice(),
attempt = i,
"Processing order"
);
@@ -153,7 +155,7 @@ impl AcmeManager {
tracing::info!(
context = "acme",
event = "csr-send",
domains = ?self.domains.as_slice(),
domains = ?provider.domains.as_slice(),
"Sending CSR"
);
@@ -164,7 +166,7 @@ impl AcmeManager {
tracing::info!(
context = "acme",
event = "download",
domains = ?self.domains.as_slice(),
domains = ?provider.domains.as_slice(),
"Downloading certificate"
);
@@ -181,7 +183,7 @@ impl AcmeManager {
context = "acme",
event = "error",
reason = "invalid-order",
domains = ?self.domains.as_slice(),
domains = ?provider.domains.as_slice(),
"Invalid order"
);
@@ -191,7 +193,12 @@ impl AcmeManager {
}
}
async fn authorize(&self, account: &Account, url: &String) -> Result<(), OrderError> {
async fn authorize(
&self,
provider: &AcmeProvider,
account: &Account,
url: &String,
) -> Result<(), OrderError> {
let auth = account.auth(url).await?;
let (domain, challenge_url) = match auth.status {
AuthStatus::Pending => {
@@ -204,7 +211,7 @@ impl AcmeManager {
);
let (challenge, auth_key) =
account.tls_alpn_01(&auth.challenges, domain.clone())?;
self.set_auth_key(domain.clone(), Arc::new(auth_key));
self.set_auth_key(provider, domain.clone(), Arc::new(auth_key));
account.challenge(&challenge.url).await?;
(domain, challenge.url.clone())
}

View File

@@ -28,23 +28,63 @@ use rustls::{
sign::CertifiedKey,
};
use super::{directory::ACME_TLS_ALPN_NAME, AcmeManager};
use crate::{listener::tls::AcmeAuthKey, Core};
impl AcmeManager {
pub(crate) fn set_cert(&self, cert: Arc<CertifiedKey>) {
self.cert.store(cert);
self.order_in_progress.store(false, Ordering::Relaxed);
self.auth_keys.lock().clear();
use super::{directory::ACME_TLS_ALPN_NAME, AcmeProvider, AcmeResolver};
impl Core {
pub(crate) fn set_cert(&self, provider: &AcmeProvider, cert: Arc<CertifiedKey>) {
// Add certificates
let mut certificates = self.tls.certificates.load().as_ref().clone();
for domain in provider.domains.iter() {
certificates.insert(
domain
.strip_prefix("*.")
.unwrap_or(domain.as_str())
.to_string(),
cert.clone(),
);
}
self.tls.certificates.store(certificates.into());
// Remove auth keys
let mut auth_keys = self.tls.acme_auth_keys.lock();
auth_keys.retain(|_, v| v.provider_id != provider.id);
self.tls
.acme_in_progress
.store(!auth_keys.is_empty(), Ordering::Relaxed);
}
pub(crate) fn set_auth_key(&self, domain: String, cert: Arc<CertifiedKey>) {
self.auth_keys.lock().insert(domain, cert);
pub(crate) fn set_auth_key(
&self,
provider: &AcmeProvider,
domain: String,
cert: Arc<CertifiedKey>,
) {
self.tls
.acme_auth_keys
.lock()
.insert(domain, AcmeAuthKey::new(provider.id.clone(), cert));
}
}
impl ResolvesServerCert for AcmeManager {
impl ResolvesServerCert for AcmeResolver {
fn resolve(&self, client_hello: ClientHello) -> Option<Arc<CertifiedKey>> {
if self.has_order_in_progress() && client_hello.is_tls_alpn_challenge() {
let core = self.core.load();
if core.has_acme_order_in_progress() && client_hello.is_tls_alpn_challenge() {
match client_hello.server_name() {
Some(domain) => {
tracing::trace!(
context = "acme",
event = "auth-key",
domain = %domain,
"Found client supplied SNI");
core.tls
.acme_auth_keys
.lock()
.get(domain)
.map(|ak| ak.key.clone())
}
None => {
tracing::debug!(
context = "acme",
@@ -54,18 +94,9 @@ impl ResolvesServerCert for AcmeManager {
);
None
}
Some(domain) => {
tracing::trace!(
context = "acme",
event = "auth-key",
domain = %domain,
"Found client supplied SNI");
self.auth_keys.lock().get(domain).cloned()
}
}
} else {
self.cert.load().clone().into()
core.resolve_certificate(client_hello.server_name())
}
}
}

View File

@@ -97,8 +97,8 @@ impl Core {
// Write blocked IP to config
self.storage
.data
.config_set([ConfigKey {
.config
.set([ConfigKey {
key: format!("{}.{}", BLOCKED_IP_KEY, ip),
value: String::new(),
}])

View File

@@ -30,7 +30,6 @@ use std::{
use arc_swap::ArcSwap;
use proxy_header::io::ProxiedStream;
use rustls::crypto::ring::cipher_suite::TLS13_AES_128_GCM_SHA256;
use store::Store;
use tokio::{
net::{TcpListener, TcpStream},
sync::watch,
@@ -40,13 +39,13 @@ use tracing::Span;
use utils::{config::Config, UnwrapFailure};
use crate::{
config::server::{Listener, Server, Servers},
config::server::{Listener, Server, ServerProtocol, Servers},
Core,
};
use super::{
acme::SpawnAcme, limiter::ConcurrencyLimiter, ServerInstance, SessionData, SessionManager,
SessionStream, TcpAcceptorResult,
limiter::ConcurrencyLimiter, ServerInstance, SessionData, SessionManager, SessionStream,
TcpAcceptor,
};
impl Server {
@@ -54,18 +53,20 @@ impl Server {
self,
manager: impl SessionManager,
core: Arc<ArcSwap<Core>>,
acceptor: TcpAcceptor,
shutdown_rx: watch::Receiver<bool>,
) {
// Prepare instance
let instance = Arc::new(ServerInstance {
id: self.id,
protocol: self.protocol,
acceptor: self.acceptor,
proxy_networks: self.proxy_networks,
limiter: ConcurrencyLimiter::new(self.max_connections),
acceptor,
shutdown_rx,
});
let is_tls = self.tls_implicit;
let is_tls = matches!(instance.acceptor, TcpAcceptor::Tls { implicit, .. } if implicit);
let is_https = is_tls && self.protocol == ServerProtocol::Http;
let has_proxies = !instance.proxy_networks.is_empty();
// Spawn listeners
@@ -114,6 +115,7 @@ impl Server {
match stream {
Ok((stream, remote_addr)) => {
let core = core.as_ref().load();
let enable_acme = is_https && core.has_acme_order_in_progress();
if has_proxies && instance.proxy_networks.iter().any(|network| network.matches(&remote_addr.ip())) {
let instance = instance.clone();
@@ -131,7 +133,7 @@ impl Server {
.unwrap_or(remote_addr);
if let Some(session) = instance.build_session(stream, local_ip, remote_addr, &core) {
// Spawn session
manager.spawn(session, is_tls);
manager.spawn(session, is_tls, enable_acme);
}
}
Err(err) => {
@@ -149,7 +151,7 @@ impl Server {
opts.apply(&session.stream);
// Spawn session
manager.spawn(session, is_tls);
manager.spawn(session, is_tls, enable_acme);
}
}
Err(err) => {
@@ -308,9 +310,17 @@ impl Servers {
// Drop privileges
#[cfg(not(target_env = "msvc"))]
{
if let Some(run_as_user) = config.value("server.run-as.user") {
if let Some(run_as_user) = config
.value("server.run-as.user")
.map(|s| s.to_string())
.or_else(|| std::env::var("RUN_AS_USER").ok())
{
let mut pd = privdrop::PrivDrop::default().user(run_as_user);
if let Some(run_as_group) = config.value("server.run-as.group") {
if let Some(run_as_group) = config
.value("server.run-as.group")
.map(|s| s.to_string())
.or_else(|| std::env::var("RUN_AS_GROUP").ok())
{
pd = pd.group(run_as_group);
}
pd.apply().failed("Failed to drop privileges");
@@ -319,21 +329,19 @@ impl Servers {
}
pub fn spawn(
self,
spawn: impl Fn(Server, watch::Receiver<bool>),
store: Store,
mut self,
spawn: impl Fn(Server, TcpAcceptor, watch::Receiver<bool>),
) -> watch::Sender<bool> {
// Spawn listeners
let (shutdown_tx, shutdown_rx) = watch::channel(false);
for server in self.servers {
spawn(server, shutdown_rx.clone());
}
let acceptor = self
.tcp_acceptors
.remove(&server.id)
.unwrap_or(TcpAcceptor::Plain);
// Spawn ACME managers
for (_, acme_manager) in self.acme_managers {
acme_manager.spawn(store.clone(), shutdown_rx.clone());
spawn(server, acceptor, shutdown_rx.clone());
}
shutdown_tx
}
}
@@ -352,8 +360,8 @@ impl ServerInstance {
stream: T,
span: &Span,
) -> Result<TlsStream<T>, ()> {
match self.acceptor.accept(stream).await {
TcpAcceptorResult::Tls(accept) => match accept.await {
match &self.acceptor {
TcpAcceptor::Tls { acceptor, .. } => match acceptor.accept(stream).await {
Ok(stream) => {
tracing::info!(
parent: span,
@@ -375,7 +383,7 @@ impl ServerInstance {
Err(())
}
},
TcpAcceptorResult::Plain(_) | TcpAcceptorResult::Close => {
TcpAcceptor::Plain => {
tracing::debug!(
parent: span,
context = "tls",

View File

@@ -40,10 +40,7 @@ use crate::{
expr::functions::ResolveVariable,
};
use self::{
acme::AcmeManager,
limiter::{ConcurrencyLimiter, InFlight},
};
use self::limiter::{ConcurrencyLimiter, InFlight};
pub mod acme;
pub mod blocked;
@@ -63,11 +60,11 @@ pub struct ServerInstance {
#[derive(Default)]
pub enum TcpAcceptor {
Tls(TlsAcceptor),
Acme {
challenge: Arc<ServerConfig>,
default: Arc<ServerConfig>,
manager: Arc<AcmeManager>,
Tls {
acme_config: Arc<ServerConfig>,
default_config: Arc<ServerConfig>,
acceptor: TlsAcceptor,
implicit: bool,
},
#[default]
Plain,
@@ -99,12 +96,22 @@ pub trait SessionStream: AsyncRead + AsyncWrite + Unpin + 'static + Sync + Send
}
pub trait SessionManager: Sync + Send + 'static + Clone {
fn spawn<T: SessionStream>(&self, mut session: SessionData<T>, is_tls: bool) {
fn spawn<T: SessionStream>(
&self,
mut session: SessionData<T>,
is_tls: bool,
enable_acme: bool,
) {
let manager = self.clone();
tokio::spawn(async move {
if is_tls {
match session.instance.acceptor.accept(session.stream).await {
match session
.instance
.acceptor
.accept(session.stream, enable_acme)
.await
{
TcpAcceptorResult::Tls(accept) => match accept.await {
Ok(stream) => {
let session = SessionData {
@@ -173,17 +180,7 @@ impl<T: SessionStream> ResolveVariable for SessionData<T> {
impl Debug for TcpAcceptor {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::Tls(_) => f.debug_tuple("Tls").finish(),
Self::Acme {
challenge,
default,
manager,
} => f
.debug_struct("Acme")
.field("challenge", challenge)
.field("default", default)
.field("manager", manager)
.finish(),
Self::Tls { .. } => f.debug_tuple("Tls").finish(),
Self::Plain => write!(f, "Plain"),
}
}

View File

@@ -22,91 +22,140 @@
*/
use std::{
cmp::Ordering,
fmt::{self, Formatter},
sync::Arc,
sync::{atomic::AtomicBool, Arc},
};
use ahash::AHashMap;
use arc_swap::ArcSwap;
use parking_lot::Mutex;
use rustls::{
client::verify_server_name,
server::{ClientHello, ParsedCertificate, ResolvesServerCert},
server::{ClientHello, ResolvesServerCert},
sign::CertifiedKey,
version::{TLS12, TLS13},
Error, SupportedProtocolVersion,
SupportedProtocolVersion,
};
use rustls_pki_types::{DnsName, ServerName};
use store::Store;
use tokio::io::{AsyncRead, AsyncWrite, AsyncWriteExt};
use tokio_rustls::{Accept, LazyConfigAcceptor, TlsAcceptor};
use tokio_rustls::{Accept, LazyConfigAcceptor};
use crate::config::server::tls::build_certified_key;
use crate::{Core, SharedCore};
use super::{acme::resolver::IsTlsAlpnChallenge, SessionStream, TcpAcceptor, TcpAcceptorResult};
use super::{
acme::{resolver::IsTlsAlpnChallenge, AcmeProvider},
SessionStream, TcpAcceptor, TcpAcceptorResult,
};
pub static TLS13_VERSION: &[&SupportedProtocolVersion] = &[&TLS13];
pub static TLS12_VERSION: &[&SupportedProtocolVersion] = &[&TLS12];
pub struct CertificateResolver {
pub sni: AHashMap<String, Arc<Certificate>>,
pub cert: Arc<Certificate>,
#[derive(Default)]
pub struct TlsManager {
pub certificates: ArcSwap<AHashMap<String, Arc<CertifiedKey>>>,
pub acme_providers: AHashMap<String, AcmeProvider>,
pub(crate) acme_auth_keys: Mutex<AHashMap<String, AcmeAuthKey>>,
pub acme_in_progress: AtomicBool,
pub self_signed_cert: Option<Arc<CertifiedKey>>,
}
pub struct Certificate {
pub cert: ArcSwap<CertifiedKey>,
pub cert_id: String,
pub(crate) struct AcmeAuthKey {
pub provider_id: String,
pub key: Arc<CertifiedKey>,
}
#[derive(Clone)]
pub struct CertificateResolver {
pub core: SharedCore,
}
impl CertificateResolver {
pub fn add(&mut self, name: &str, ck: Arc<Certificate>) -> Result<(), Error> {
let server_name = {
let checked_name = DnsName::try_from(name)
.map_err(|_| Error::General("Bad DNS name".into()))
.map(|name| name.to_lowercase_owned())?;
ServerName::DnsName(checked_name)
};
pub fn new(core: SharedCore) -> Self {
Self { core }
}
}
ck.cert
.load()
.end_entity_cert()
.and_then(ParsedCertificate::try_from)
.and_then(|cert| verify_server_name(&cert, &server_name))?;
if let ServerName::DnsName(name) = server_name {
self.sni.insert(name.as_ref().to_string(), ck);
}
Ok(())
impl AcmeAuthKey {
pub fn new(provider_id: String, key: Arc<CertifiedKey>) -> Self {
Self { provider_id, key }
}
}
impl ResolvesServerCert for CertificateResolver {
fn resolve(&self, hello: ClientHello<'_>) -> Option<Arc<CertifiedKey>> {
if !self.sni.is_empty() {
if let Some(cert) = hello.server_name().and_then(|name| self.sni.get(name)) {
return cert.cert.load().clone().into();
self.core
.as_ref()
.load()
.resolve_certificate(hello.server_name())
}
}
impl Core {
pub(crate) fn resolve_certificate(&self, name: Option<&str>) -> Option<Arc<CertifiedKey>> {
let certs = self.tls.certificates.load();
name.map_or_else(
|| certs.get("*"),
|name| {
certs
.get(name)
.or_else(|| {
// Try with a wildcard certificate
name.split_once('.')
.and_then(|(_, domain)| certs.get(domain))
})
.or_else(|| {
tracing::debug!(
context = "tls",
event = "not-found",
client_name = name,
"No SNI certificate found by name, using default."
);
certs.get("*")
})
},
)
.or_else(|| match certs.len().cmp(&1) {
Ordering::Equal => certs.values().next(),
Ordering::Greater => {
tracing::debug!(
context = "tls",
event = "error",
"Multiple certificates available and no default certificate configured."
);
certs.values().next()
}
}
self.cert.cert.load().clone().into()
Ordering::Less => {
tracing::warn!(
context = "tls",
event = "error",
"No certificates available, using self-signed."
);
self.tls.self_signed_cert.as_ref()
}
})
.cloned()
}
}
impl TcpAcceptor {
pub async fn accept<IO>(&self, stream: IO) -> TcpAcceptorResult<IO>
pub async fn accept<IO>(&self, stream: IO, enable_acme: bool) -> TcpAcceptorResult<IO>
where
IO: SessionStream,
{
match self {
TcpAcceptor::Tls(acceptor) => TcpAcceptorResult::Tls(acceptor.accept(stream)),
TcpAcceptor::Acme {
challenge,
default,
manager,
} => {
if manager.has_order_in_progress() {
TcpAcceptor::Tls {
acme_config,
default_config,
acceptor,
implicit,
} if *implicit => {
if !enable_acme {
TcpAcceptorResult::Tls(acceptor.accept(stream))
} else {
match LazyConfigAcceptor::new(Default::default(), stream).await {
Ok(start_handshake) => {
if start_handshake.client_hello().is_tls_alpn_challenge() {
match start_handshake.into_stream(challenge.clone()).await {
match start_handshake.into_stream(acme_config.clone()).await {
Ok(mut tls) => {
tracing::debug!(
context = "acme",
@@ -126,7 +175,7 @@ impl TcpAcceptor {
}
} else {
return TcpAcceptorResult::Tls(
start_handshake.into_stream(default.clone()),
start_handshake.into_stream(default_config.clone()),
);
}
}
@@ -141,16 +190,14 @@ impl TcpAcceptor {
}
TcpAcceptorResult::Close
} else {
TcpAcceptorResult::Tls(TlsAcceptor::from(default.clone()).accept(stream))
}
}
TcpAcceptor::Plain => TcpAcceptorResult::Plain(stream),
_ => TcpAcceptorResult::Plain(stream),
}
}
pub fn is_tls(&self) -> bool {
matches!(self, TcpAcceptor::Tls(_) | TcpAcceptor::Acme { .. })
matches!(self, TcpAcceptor::Tls { .. })
}
}
@@ -166,37 +213,8 @@ where
}
}
impl Certificate {
pub async fn reload(&self, store: &Store) -> utils::config::Result<()> {
match (
store
.config_get(format!("certificate.{}.cert", self.cert_id))
.await,
store
.config_get(format!("certificate.{}.private-key", self.cert_id))
.await,
) {
(Ok(Some(cert)), Ok(Some(pk))) => {
match build_certified_key(cert.into_bytes(), pk.into_bytes()) {
Ok(cert) => {
self.cert.store(Arc::new(cert));
Ok(())
}
Err(err) => Err(err),
}
}
(Ok(None), _) | (_, Ok(None)) => Err("Certificate or private key not found".into()),
(Err(err), _) | (_, Err(err)) => Err(err.to_string()),
}
}
}
impl std::fmt::Debug for CertificateResolver {
fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
f.debug_struct("CertificateResolver")
.field("sni", &self.sni.keys())
.field("id", &self.cert.cert_id)
.finish()
f.debug_struct("CertificateResolver").finish()
}
}

View File

@@ -380,12 +380,12 @@ impl JMAP {
.into_http_response(),
}
}
/*("reload", Some("settings"), &Method::GET) => {
let _ = self
.inner
.housekeeper_tx
.send(housekeeper::Event::ReloadConfig)
.await;
("reload", Some("settings"), &Method::GET) => {
/*let _ = self
.inner
.housekeeper_tx
.send(housekeeper::Event::ReloadConfig)
.await;*/
JsonResponse::new(json!({
"data": (),
@@ -393,17 +393,17 @@ impl JMAP {
.into_http_response()
}
("reload", Some("certificates"), &Method::GET) => {
let _ = self
.inner
.housekeeper_tx
.send(housekeeper::Event::ReloadCertificates)
.await;
/*let _ = self
.inner
.housekeeper_tx
.send(housekeeper::Event::ReloadCertificates)
.await;*/
JsonResponse::new(json!({
"data": (),
}))
.into_http_response()
}*/
}
("settings", Some("group"), &Method::GET) => {
// List settings
let params = UrlParams::new(req.uri().query());
@@ -434,7 +434,7 @@ impl JMAP {
params.parse::<usize>("page").unwrap_or(0).saturating_sub(1) * limit;
let has_filter = !filter.is_empty();
match self.core.storage.data.config_list(&prefix, true).await {
match self.core.storage.config.list(&prefix, true).await {
Ok(settings) => if !suffix.is_empty() && !settings.is_empty() {
// Obtain record ids
let mut total = 0;
@@ -548,7 +548,7 @@ impl JMAP {
let limit: usize = params.parse("limit").unwrap_or(0);
let offset = params.parse::<usize>("page").unwrap_or(0).saturating_sub(1) * limit;
match self.core.storage.data.config_list(&prefix, true).await {
match self.core.storage.config.list(&prefix, true).await {
Ok(settings) => {
let total = settings.len();
let items = settings
@@ -588,7 +588,7 @@ impl JMAP {
let mut results = AHashMap::with_capacity(keys.len());
for key in keys {
match self.core.storage.data.config_get(key).await {
match self.core.storage.config.get(key).await {
Ok(Some(value)) => {
results.insert(key.to_string(), value);
}
@@ -605,7 +605,7 @@ impl JMAP {
} else {
prefix.to_string()
};
match self.core.storage.data.config_list(&prefix, false).await {
match self.core.storage.config.list(&prefix, false).await {
Ok(values) => {
results.extend(values);
}
@@ -631,7 +631,7 @@ impl JMAP {
}
}
("settings", Some(prefix), &Method::DELETE) if !prefix.is_empty() => {
match self.core.storage.data.config_clear(prefix).await {
match self.core.storage.config.clear(prefix).await {
Ok(_) => JsonResponse::new(json!({
"data": (),
}))
@@ -654,13 +654,8 @@ impl JMAP {
match change {
UpdateSettings::Delete { keys } => {
for key in keys {
result = self
.core
.storage
.data
.config_clear(key)
.await
.map(|_| true);
result =
self.core.storage.config.clear(key).await.map(|_| true);
if result.is_err() {
break 'next;
}
@@ -670,8 +665,8 @@ impl JMAP {
result = self
.core
.storage
.data
.config_clear_prefix(&prefix)
.config
.clear_prefix(&prefix)
.await
.map(|_| true);
if result.is_err() {
@@ -688,8 +683,8 @@ impl JMAP {
result = self
.core
.storage
.data
.config_list(&format!("{prefix}."), true)
.config
.list(&format!("{prefix}."), true)
.await
.map(|items| items.is_empty());
@@ -700,8 +695,8 @@ impl JMAP {
result = self
.core
.storage
.data
.config_get(key)
.config
.get(key)
.await
.map(|items| items.is_none());
@@ -714,8 +709,8 @@ impl JMAP {
result = self
.core
.storage
.data
.config_set(values.into_iter().map(|(key, value)| ConfigKey {
.config
.set(values.into_iter().map(|(key, value)| ConfigKey {
key: if let Some(prefix) = &prefix {
format!("{prefix}.{key}")
} else {

View File

@@ -21,7 +21,10 @@
* for more details.
*/
use std::{collections::BinaryHeap, time::Instant};
use std::{
collections::BinaryHeap,
time::{Duration, Instant},
};
use store::write::purge::PurgeStore;
use tokio::sync::mpsc;
@@ -34,21 +37,26 @@ use super::IPC_CHANNEL_BUFFER;
pub enum Event {
IndexStart,
IndexDone,
AcmeReschedule {
provider_id: String,
renew_at: Instant,
},
#[cfg(feature = "test_mode")]
IndexIsActive(tokio::sync::oneshot::Sender<bool>),
Exit,
}
#[derive(PartialEq, Eq)]
struct PurgeEvent {
struct Action {
due: Instant,
event: PurgeClass,
event: ActionClass,
}
#[derive(PartialEq, Eq)]
enum PurgeClass {
enum ActionClass {
Session,
Store(usize),
Acme(String),
}
pub fn spawn_housekeeper(core: JmapInstance, mut rx: mpsc::Receiver<Event>) {
@@ -67,17 +75,36 @@ pub fn spawn_housekeeper(core: JmapInstance, mut rx: mpsc::Receiver<Event>) {
// Add all purge events to heap
let core_ = core.core.load();
heap.push(PurgeEvent {
heap.push(Action {
due: Instant::now() + core_.jmap.session_purge_frequency.time_to_next(),
event: PurgeClass::Session,
event: ActionClass::Session,
});
for (idx, schedule) in core_.storage.purge_schedules.iter().enumerate() {
heap.push(PurgeEvent {
heap.push(Action {
due: Instant::now() + schedule.cron.time_to_next(),
event: PurgeClass::Store(idx),
event: ActionClass::Store(idx),
});
}
// Add all ACME renewals to heap
for provider in core_.tls.acme_providers.values() {
match core_.init_acme(provider).await {
Ok(renew_at) => {
heap.push(Action {
due: Instant::now() + renew_at,
event: ActionClass::Acme(provider.id.clone()),
});
}
Err(err) => {
tracing::error!(
context = "acme",
event = "error",
error = ?err,
"Failed to initialize ACME certificate manager.");
}
};
}
loop {
let time_to_next = heap
.peek()
@@ -86,6 +113,15 @@ pub fn spawn_housekeeper(core: JmapInstance, mut rx: mpsc::Receiver<Event>) {
match tokio::time::timeout(time_to_next, rx.recv()).await {
Ok(Some(event)) => match event {
Event::AcmeReschedule {
provider_id,
renew_at,
} => {
heap.push(Action {
due: renew_at,
event: ActionClass::Acme(provider_id),
});
}
Event::IndexStart => {
if !index_busy {
index_busy = true;
@@ -129,25 +165,70 @@ pub fn spawn_housekeeper(core: JmapInstance, mut rx: mpsc::Receiver<Event>) {
}
let event = heap.pop().unwrap();
match event.event {
PurgeClass::Session => {
ActionClass::Acme(provider_id) => {
let inner = core.jmap_inner.clone();
let core = core_.clone();
tokio::spawn(async move {
if let Some(provider) =
core.tls.acme_providers.get(&provider_id)
{
tracing::info!(
context = "acme",
event = "order",
domains = ?provider.domains,
"Ordering certificates.");
let renew_at = match core.renew(provider).await {
Ok(renew_at) => {
tracing::info!(
context = "acme",
event = "success",
domains = ?provider.domains,
next_renewal = ?renew_at,
"Certificates renewed.");
renew_at
}
Err(err) => {
tracing::error!(
context = "acme",
event = "error",
error = ?err,
"Failed to renew certificates.");
Duration::from_secs(3600)
}
};
inner
.housekeeper_tx
.send(Event::AcmeReschedule {
provider_id: provider_id.clone(),
renew_at: Instant::now() + renew_at,
})
.await
.ok();
}
});
}
ActionClass::Session => {
let inner = core.jmap_inner.clone();
tokio::spawn(async move {
tracing::debug!("Purging session cache.");
inner.purge();
});
heap.push(PurgeEvent {
heap.push(Action {
due: Instant::now()
+ core_.jmap.session_purge_frequency.time_to_next(),
event: PurgeClass::Session,
event: ActionClass::Session,
});
}
PurgeClass::Store(idx) => {
ActionClass::Store(idx) => {
if let Some(schedule) =
core_.storage.purge_schedules.get(idx).cloned()
{
heap.push(PurgeEvent {
heap.push(Action {
due: Instant::now() + schedule.cron.time_to_next(),
event: PurgeClass::Store(idx),
event: ActionClass::Store(idx),
});
tokio::spawn(async move {
let (class, result) = match schedule.store {
@@ -187,13 +268,13 @@ pub fn spawn_housekeeper(core: JmapInstance, mut rx: mpsc::Receiver<Event>) {
});
}
impl Ord for PurgeEvent {
impl Ord for Action {
fn cmp(&self, other: &Self) -> std::cmp::Ordering {
self.due.cmp(&other.due).reverse()
}
}
impl PartialOrd for PurgeEvent {
impl PartialOrd for Action {
fn partial_cmp(&self, other: &Self) -> Option<std::cmp::Ordering> {
Some(self.cmp(other))
}

View File

@@ -23,23 +23,13 @@
use std::time::Duration;
use common::{
config::{
server::{ServerProtocol, Servers},
tracers::Tracers,
},
Core,
};
use common::config::{manager::BootManager, server::ServerProtocol};
use imap::core::{ImapSessionManager, IMAP};
use jmap::{api::JmapSessionManager, services::IPC_CHANNEL_BUFFER, JMAP};
use managesieve::core::ManageSieveSessionManager;
use smtp::core::{SmtpSessionManager, SMTP};
use store::Stores;
use tokio::sync::mpsc;
use utils::{
config::{Config, ConfigError},
wait_for_shutdown,
};
use utils::wait_for_shutdown;
#[cfg(not(target_env = "msvc"))]
use jemallocator::Jemalloc;
@@ -51,75 +41,55 @@ static GLOBAL: Jemalloc = Jemalloc;
#[tokio::main]
async fn main() -> std::io::Result<()> {
// Load config and apply macros
let mut config = Config::init();
config.resolve_macros().await;
let init = BootManager::init().await;
// Parse servers
let servers = Servers::parse(&mut config);
// Parse core
let mut config = init.config;
let core = init.core;
// Bind ports and drop privileges
servers.bind_and_drop_priv(&mut config);
// Build stores
let stores = Stores::parse(&mut config).await;
let todo = "merge config with data store, resolve macros";
// Enable tracing
let guards = Tracers::parse(&mut config).enable(&mut config);
// Init servers
tracing::info!(
"Starting Stalwart Mail Server v{}...",
env!("CARGO_PKG_VERSION")
);
// Parse core
let core = Core::parse(&mut config, stores).await;
let store = core.storage.data.clone();
let shared_core = core.into_shared();
// Init servers
let (delivery_tx, delivery_rx) = mpsc::channel(IPC_CHANNEL_BUFFER);
let smtp = SMTP::init(&mut config, shared_core.clone(), delivery_tx).await;
let jmap = JMAP::init(
&mut config,
delivery_rx,
shared_core.clone(),
smtp.inner.clone(),
)
.await;
let smtp = SMTP::init(&mut config, core.clone(), delivery_tx).await;
let jmap = JMAP::init(&mut config, delivery_rx, core.clone(), smtp.inner.clone()).await;
let imap = IMAP::init(&mut config, jmap.clone()).await;
// Log configuration errors
config.log_errors(guards.is_none());
config.log_warnings(guards.is_none());
config.log_errors(init.guards.is_none());
config.log_warnings(init.guards.is_none());
// Spawn servers
let shutdown_tx = servers.spawn(
|server, shutdown_rx| {
match &server.protocol {
ServerProtocol::Smtp | ServerProtocol::Lmtp => server.spawn(
SmtpSessionManager::new(smtp.clone()),
shared_core.clone(),
shutdown_rx,
),
ServerProtocol::Http => server.spawn(
JmapSessionManager::new(jmap.clone()),
shared_core.clone(),
shutdown_rx,
),
ServerProtocol::Imap => server.spawn(
ImapSessionManager::new(imap.clone()),
shared_core.clone(),
shutdown_rx,
),
ServerProtocol::ManageSieve => server.spawn(
ManageSieveSessionManager::new(imap.clone()),
shared_core.clone(),
shutdown_rx,
),
};
},
store,
);
let shutdown_tx = init.servers.spawn(|server, acceptor, shutdown_rx| {
match &server.protocol {
ServerProtocol::Smtp | ServerProtocol::Lmtp => server.spawn(
SmtpSessionManager::new(smtp.clone()),
core.clone(),
acceptor,
shutdown_rx,
),
ServerProtocol::Http => server.spawn(
JmapSessionManager::new(jmap.clone()),
core.clone(),
acceptor,
shutdown_rx,
),
ServerProtocol::Imap => server.spawn(
ImapSessionManager::new(imap.clone()),
core.clone(),
acceptor,
shutdown_rx,
),
ServerProtocol::ManageSieve => server.spawn(
ManageSieveSessionManager::new(imap.clone()),
core.clone(),
acceptor,
shutdown_rx,
),
};
});
// Wait for shutdown signal
wait_for_shutdown(&format!(

View File

@@ -42,103 +42,104 @@ impl MemoryStore {
}
}
pub fn parse_memory_stores(config: &mut Config, stores: &mut Stores) {
let mut lookups = AHashMap::new();
let mut errors = Vec::new();
impl Stores {
pub fn parse_memory_stores(&mut self, config: &mut Config) {
let mut lookups = AHashMap::new();
let mut errors = Vec::new();
for (key, value) in &config.keys {
if let Some(key) = key.strip_prefix("lookup.") {
if let Some((id, key)) = key
.split_once('.')
.filter(|(id, key)| !id.is_empty() && !key.is_empty())
{
// Detect if the key is a glob pattern
let mut last_ch = '\0';
let mut has_escape = false;
let mut is_glob = false;
for ch in key.chars() {
match ch {
'\\' => {
has_escape = true;
}
'*' | '?' if last_ch != '\\' => {
is_glob = true;
}
_ => {}
}
last_ch = ch;
}
// Detect value type
let value = if !value.is_empty() {
let mut has_integers = false;
let mut has_floats = false;
let mut has_others = false;
for (pos, ch) in value.as_bytes().iter().enumerate() {
for (key, value) in &config.keys {
if let Some(key) = key.strip_prefix("lookup.") {
if let Some((id, key)) = key
.split_once('.')
.filter(|(id, key)| !id.is_empty() && !key.is_empty())
{
// Detect if the key is a glob pattern
let mut last_ch = '\0';
let mut has_escape = false;
let mut is_glob = false;
for ch in key.chars() {
match ch {
b'.' if !has_floats && has_integers => {
has_floats = true;
'\\' => {
has_escape = true;
}
b'0'..=b'9' => {
has_integers = true;
'*' | '?' if last_ch != '\\' => {
is_glob = true;
}
b'-' if pos == 0 && value.len() > 1 => {}
_ => {
has_others = true;
_ => {}
}
last_ch = ch;
}
// Detect value type
let value = if !value.is_empty() {
let mut has_integers = false;
let mut has_floats = false;
let mut has_others = false;
for (pos, ch) in value.as_bytes().iter().enumerate() {
match ch {
b'.' if !has_floats && has_integers => {
has_floats = true;
}
b'0'..=b'9' => {
has_integers = true;
}
b'-' if pos == 0 && value.len() > 1 => {}
_ => {
has_others = true;
}
}
}
}
if has_others {
Value::Text(value.to_string().into())
} else if has_floats {
value
.parse()
.map(Value::Float)
.unwrap_or_else(|_| Value::Text(value.to_string().into()))
} else {
value
.parse()
.map(Value::Integer)
.unwrap_or_else(|_| Value::Text(value.to_string().into()))
}
} else {
Value::Text("".into())
};
// Add entry
let store = lookups
.entry(id.to_string())
.or_insert_with(MemoryStore::default);
if is_glob {
store.globs.push((GlobPattern::compile(key, false), value));
} else {
store.entries.insert(
if has_escape {
key.replace('\\', "")
if has_others {
Value::Text(value.to_string().into())
} else if has_floats {
value
.parse()
.map(Value::Float)
.unwrap_or_else(|_| Value::Text(value.to_string().into()))
} else {
key.to_string()
},
value,
);
value
.parse()
.map(Value::Integer)
.unwrap_or_else(|_| Value::Text(value.to_string().into()))
}
} else {
Value::Text("".into())
};
// Add entry
let store = lookups
.entry(id.to_string())
.or_insert_with(MemoryStore::default);
if is_glob {
store.globs.push((GlobPattern::compile(key, false), value));
} else {
store.entries.insert(
if has_escape {
key.replace('\\', "")
} else {
key.to_string()
},
value,
);
}
} else {
errors.push(key.to_string());
}
} else {
errors.push(key.to_string());
} else if !lookups.is_empty() {
break;
}
} else if !lookups.is_empty() {
break;
}
for error in errors {
config.new_parse_error(error, "Invalid lookup key format");
}
for (id, store) in lookups {
self.lookup_stores
.insert(id, LookupStore::Memory(store.into()));
}
}
for error in errors {
config.new_parse_error(error, "Invalid lookup key format");
}
for (id, store) in lookups {
stores
.lookup_stores
.insert(id, LookupStore::Memory(store.into()));
}
}

View File

@@ -26,7 +26,7 @@ use std::sync::Arc;
use utils::config::{cron::SimpleCron, Config};
use crate::{
backend::{fs::FsStore, memory::parse_memory_stores},
backend::fs::FsStore,
write::purge::{PurgeSchedule, PurgeStore},
BlobStore, CompressionAlgo, FtsStore, LookupStore, QueryStore, Store, Stores,
};
@@ -56,13 +56,19 @@ use crate::backend::elastic::ElasticSearchStore;
use crate::backend::redis::RedisStore;
impl Stores {
pub async fn parse_all(config: &mut Config) -> Self {
let mut stores = Self::parse(config).await;
stores.parse_lookups(config).await;
stores
}
pub async fn parse(config: &mut Config) -> Self {
let mut stores = Stores::default();
let ids = config
for id in config
.sub_keys("store", ".type")
.map(|id| id.to_string())
.collect::<Vec<_>>();
for id in ids {
.collect::<Vec<_>>()
{
let id = id.as_str();
// Parse store
#[cfg(feature = "test_mode")]
@@ -86,7 +92,7 @@ impl Stores {
.property_or_default::<CompressionAlgo>(("store", id, "compression"), "none")
.unwrap_or(CompressionAlgo::None);
let lookup_store: Store = match protocol.as_str() {
match protocol.as_str() {
#[cfg(feature = "rocks")]
"rocksdb" => {
if let Some(db) = RocksDbStore::open(config, prefix).await.map(Store::from) {
@@ -100,7 +106,6 @@ impl Stores {
);
stores.lookup_stores.insert(store_id, db.into());
}
continue;
}
#[cfg(feature = "foundation")]
"foundationdb" => {
@@ -115,7 +120,6 @@ impl Stores {
);
stores.lookup_stores.insert(store_id, db.into());
}
continue;
}
#[cfg(feature = "postgres")]
"postgresql" => {
@@ -128,9 +132,7 @@ impl Stores {
store_id.clone(),
BlobStore::from(db.clone()).with_compression(compression_algo),
);
db
} else {
continue;
stores.lookup_stores.insert(store_id.clone(), db.into());
}
}
#[cfg(feature = "mysql")]
@@ -144,9 +146,7 @@ impl Stores {
store_id.clone(),
BlobStore::from(db.clone()).with_compression(compression_algo),
);
db
} else {
continue;
stores.lookup_stores.insert(store_id.clone(), db.into());
}
}
#[cfg(feature = "sqlite")]
@@ -160,9 +160,7 @@ impl Stores {
store_id.clone(),
BlobStore::from(db.clone()).with_compression(compression_algo),
);
db
} else {
continue;
stores.lookup_stores.insert(store_id.clone(), db.into());
}
}
"fs" => {
@@ -171,7 +169,6 @@ impl Stores {
.blob_stores
.insert(store_id, db.with_compression(compression_algo));
}
continue;
}
#[cfg(feature = "s3")]
"s3" => {
@@ -180,7 +177,6 @@ impl Stores {
.blob_stores
.insert(store_id, db.with_compression(compression_algo));
}
continue;
}
#[cfg(feature = "elastic")]
"elasticsearch" => {
@@ -190,7 +186,6 @@ impl Stores {
{
stores.fts_stores.insert(store_id, db);
}
continue;
}
#[cfg(feature = "redis")]
"redis" => {
@@ -200,19 +195,36 @@ impl Stores {
{
stores.lookup_stores.insert(store_id, db);
}
continue;
}
unknown => {
tracing::debug!("Unknown directory type: {unknown:?}");
continue;
}
};
}
}
stores
}
pub async fn parse_lookups(&mut self, config: &mut Config) {
// Parse memory stores
self.parse_memory_stores(config);
// Add SQL queries as lookup stores
for (store_id, lookup_store) in self.stores.iter().filter_map(|(id, store)| {
if matches!(
store,
Store::MySQL(_) | Store::PostgreSQL(_) | Store::SQLite(_)
) {
Some((id.clone(), LookupStore::from(store.clone())))
} else {
None
}
}) {
// Add queries as lookup stores
let lookup_store: LookupStore = lookup_store.into();
for lookup_id in config.sub_keys(("store", id, "query"), "") {
if let Some(query) = config.value(("store", id, "query", lookup_id)) {
stores.lookup_stores.insert(
for lookup_id in config.sub_keys(("store", store_id.as_str(), "query"), "") {
if let Some(query) = config.value(("store", store_id.as_str(), "query", lookup_id))
{
self.lookup_stores.insert(
format!("{store_id}/{lookup_id}"),
LookupStore::Query(Arc::new(QueryStore {
store: lookup_store.clone(),
@@ -221,17 +233,16 @@ impl Stores {
);
}
}
stores.lookup_stores.insert(store_id, lookup_store.clone());
// Run init queries on database
for query in config
.values(("store", id, "init.execute"))
.values(("store", store_id.as_str(), "init.execute"))
.map(|(_, s)| s.to_string())
.collect::<Vec<_>>()
{
if let Err(err) = lookup_store.query::<usize>(&query, Vec::new()).await {
config.new_build_error(
("store", id),
("store", store_id.as_str()),
format!("Failed to initialize store: {err}"),
);
}
@@ -241,13 +252,13 @@ impl Stores {
// Parse purge schedules
if let Some(store) = config
.value("storage.data")
.and_then(|store_id| stores.stores.get(store_id))
.and_then(|store_id| self.stores.get(store_id))
{
let store_id = config.value("storage.data").unwrap().to_string();
if let Some(cron) =
config.property::<SimpleCron>(("store", store_id.as_str(), "purge.frequency"))
{
stores.purge_schedules.push(PurgeSchedule {
self.purge_schedules.push(PurgeSchedule {
cron,
store_id,
store: PurgeStore::Data(store.clone()),
@@ -256,13 +267,13 @@ impl Stores {
if let Some(blob_store) = config
.value("storage.blob")
.and_then(|blob_store_id| stores.blob_stores.get(blob_store_id))
.and_then(|blob_store_id| self.blob_stores.get(blob_store_id))
{
let store_id = config.value("storage.blob").unwrap().to_string();
if let Some(cron) =
config.property::<SimpleCron>(("store", store_id.as_str(), "purge.frequency"))
{
stores.purge_schedules.push(PurgeSchedule {
self.purge_schedules.push(PurgeSchedule {
cron,
store_id,
store: PurgeStore::Blobs {
@@ -273,22 +284,17 @@ impl Stores {
}
}
}
for (store_id, store) in &stores.lookup_stores {
for (store_id, store) in &self.lookup_stores {
if let Some(cron) =
config.property::<SimpleCron>(("store", store_id.as_str(), "purge.frequency"))
{
stores.purge_schedules.push(PurgeSchedule {
self.purge_schedules.push(PurgeSchedule {
cron,
store_id: store_id.clone(),
store: PurgeStore::Lookup(store.clone()),
});
}
}
// Parse memory stores
parse_memory_stores(config, &mut stores);
stores
}
}

View File

@@ -1,100 +0,0 @@
/*
* Copyright (c) 2023 Stalwart Labs Ltd.
*
* This file is part of the Stalwart Mail Server.
*
* This program is free software: you can redistribute it and/or modify
* it under the terms of the GNU Affero General Public License as
* published by the Free Software Foundation, either version 3 of
* the License, or (at your option) any later version.
*
* This program is distributed in the hope that it will be useful,
* but WITHOUT ANY WARRANTY; without even the implied warranty of
* MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the
* GNU Affero General Public License for more details.
* in the LICENSE file at the top-level directory of this distribution.
* You should have received a copy of the GNU Affero General Public License
* along with this program. If not, see <http://www.gnu.org/licenses/>.
*
* You can be released from the requirements of the AGPLv3 license by
* purchasing a commercial license. Please contact licensing@stalw.art
* for more details.
*/
use utils::config::ConfigKey;
use crate::{
write::{BatchBuilder, ValueClass},
Deserialize, IterateParams, Store, ValueKey,
};
impl Store {
pub async fn config_get(&self, key: impl Into<String>) -> crate::Result<Option<String>> {
self.get_value(ValueKey::from(ValueClass::Config(key.into().into_bytes())))
.await
}
pub async fn config_list(
&self,
prefix: &str,
strip_prefix: bool,
) -> crate::Result<Vec<(String, String)>> {
let key = prefix.as_bytes();
let from_key = ValueKey::from(ValueClass::Config(key.to_vec()));
let to_key = ValueKey::from(ValueClass::Config(
key.iter()
.copied()
.chain([u8::MAX, u8::MAX, u8::MAX, u8::MAX, u8::MAX])
.collect::<Vec<_>>(),
));
let mut results = Vec::new();
self.iterate(
IterateParams::new(from_key, to_key).ascending(),
|key, value| {
let mut key =
std::str::from_utf8(key.get(1..).unwrap_or_default()).map_err(|_| {
crate::Error::InternalError("Failed to deserialize config key".to_string())
})?;
if strip_prefix && !prefix.is_empty() {
key = key.strip_prefix(prefix).unwrap_or(key);
}
results.push((key.to_string(), String::deserialize(value)?));
Ok(true)
},
)
.await?;
Ok(results)
}
pub async fn config_set(&self, keys: impl IntoIterator<Item = ConfigKey>) -> crate::Result<()> {
let mut batch = BatchBuilder::new();
for key in keys {
batch.set(ValueClass::Config(key.key.into_bytes()), key.value);
}
self.write(batch.build()).await.map(|_| ())
}
pub async fn config_clear(&self, key: impl Into<String>) -> crate::Result<()> {
let mut batch = BatchBuilder::new();
batch.clear(ValueClass::Config(key.into().into_bytes()));
self.write(batch.build()).await.map(|_| ())
}
pub async fn config_clear_prefix(&self, key: impl AsRef<str>) -> crate::Result<()> {
self.delete_range(
ValueKey::from(ValueClass::Config(key.as_ref().as_bytes().to_vec())),
ValueKey::from(ValueClass::Config(
key.as_ref()
.as_bytes()
.iter()
.copied()
.chain([u8::MAX, u8::MAX, u8::MAX, u8::MAX, u8::MAX])
.collect::<Vec<_>>(),
)),
)
.await
}
}

View File

@@ -22,7 +22,6 @@
*/
pub mod blob;
pub mod config;
pub mod fts;
pub mod lookup;
pub mod store;

View File

@@ -673,6 +673,12 @@ impl From<Rows> for Vec<u32> {
}
}
impl Store {
pub fn is_none(&self) -> bool {
matches!(self, Self::None)
}
}
impl std::fmt::Debug for Store {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {

View File

@@ -30,8 +30,6 @@ use std::{collections::BTreeMap, time::Duration};
use ahash::AHashMap;
use crate::{failed, UnwrapFailure};
#[derive(Debug, Default, Clone, PartialEq, Eq)]
pub struct Config {
pub keys: BTreeMap<String, String>,
@@ -61,42 +59,6 @@ pub struct Rate {
pub type Result<T> = std::result::Result<T, String>;
impl Config {
pub fn init() -> Self {
let mut config_path = None;
let mut found_param = false;
for arg in std::env::args().skip(1) {
if let Some((key, value)) = arg.split_once('=') {
if key.starts_with("--config") {
config_path = value.trim().to_string().into();
break;
} else {
failed(&format!("Invalid command line argument: {key}"));
}
} else if found_param {
config_path = arg.into();
break;
} else if arg.starts_with("--config") {
found_param = true;
} else {
failed(&format!("Invalid command line argument: {arg}"));
}
}
// Read main configuration file
let mut config = Config::default();
config
.parse(
&std::fs::read_to_string(
config_path.failed("Missing parameter --config=<path-to-config>."),
)
.failed("Could not read configuration file"),
)
.failed("Invalid configuration file");
config
}
pub async fn resolve_macros(&mut self) {
for macro_class in ["env", "file", "cfg"] {
self.resolve_macro_type(macro_class).await;

View File

@@ -292,7 +292,7 @@ impl DirectoryTest {
config_file.replace("type = \"memory\"", "type = \"memory\"\ndisable = true")
}
let mut config = utils::config::Config::new(&config_file).unwrap();
let stores = Stores::parse(&mut config).await;
let stores = Stores::parse_all(&mut config).await;
let directories = Directories::parse(
&mut config,
&stores,
@@ -570,7 +570,7 @@ async fn lookup_local() {
)
.unwrap();
let lookups = Stores::parse(&mut config).await.lookup_stores;
let lookups = Stores::parse_all(&mut config).await.lookup_stores;
for (lookup, item, expect) in [
("glob", "user@example.org", true),

View File

@@ -273,19 +273,22 @@ async fn init_imap_tests(store_id: &str, delete_if_exists: bool) -> IMAPTest {
config.resolve_macros().await;
// Parse servers
let servers = Servers::parse(&mut config);
let mut servers = Servers::parse(&mut config);
// Bind ports and drop privileges
servers.bind_and_drop_priv(&mut config);
// Build stores
let stores = Stores::parse(&mut config).await;
let stores = Stores::parse_all(&mut config).await;
// Parse core
let core = Core::parse(&mut config, stores).await;
let core = Core::parse(&mut config, stores, Default::default()).await;
let store = core.storage.data.clone();
let shared_core = core.into_shared();
// Parse acceptors
servers.parse_tcp_acceptors(&mut config, shared_core.clone());
// Init servers
let (delivery_tx, delivery_rx) = mpsc::channel(IPC_CHANNEL_BUFFER);
let smtp = SMTP::init(&mut config, shared_core.clone(), delivery_tx).await;
@@ -300,33 +303,34 @@ async fn init_imap_tests(store_id: &str, delete_if_exists: bool) -> IMAPTest {
config.assert_no_errors();
// Spawn servers
let shutdown_tx = servers.spawn(
|server, shutdown_rx| {
match &server.protocol {
ServerProtocol::Smtp | ServerProtocol::Lmtp => server.spawn(
SmtpSessionManager::new(smtp.clone()),
shared_core.clone(),
shutdown_rx,
),
ServerProtocol::Http => server.spawn(
JmapSessionManager::new(jmap.clone()),
shared_core.clone(),
shutdown_rx,
),
ServerProtocol::Imap => server.spawn(
ImapSessionManager::new(imap.clone()),
shared_core.clone(),
shutdown_rx,
),
ServerProtocol::ManageSieve => server.spawn(
ManageSieveSessionManager::new(imap.clone()),
shared_core.clone(),
shutdown_rx,
),
};
},
store.clone(),
);
let shutdown_tx = servers.spawn(|server, acceptor, shutdown_rx| {
match &server.protocol {
ServerProtocol::Smtp | ServerProtocol::Lmtp => server.spawn(
SmtpSessionManager::new(smtp.clone()),
shared_core.clone(),
acceptor,
shutdown_rx,
),
ServerProtocol::Http => server.spawn(
JmapSessionManager::new(jmap.clone()),
shared_core.clone(),
acceptor,
shutdown_rx,
),
ServerProtocol::Imap => server.spawn(
ImapSessionManager::new(imap.clone()),
shared_core.clone(),
acceptor,
shutdown_rx,
),
ServerProtocol::ManageSieve => server.spawn(
ManageSieveSessionManager::new(imap.clone()),
shared_core.clone(),
acceptor,
shutdown_rx,
),
};
});
// Create tables and test accounts
let lookup = DirectoryStore {
store: shared_core

View File

@@ -127,8 +127,8 @@ pub async fn test(params: &mut JMAPTest) {
server
.core
.storage
.data
.config_get(format!("{BLOCKED_IP_KEY}.127.0.0.1"))
.config
.get(format!("{BLOCKED_IP_KEY}.127.0.0.1"))
.await
.unwrap(),
None
@@ -149,8 +149,8 @@ pub async fn test(params: &mut JMAPTest) {
server
.core
.storage
.data
.config_get(format!("{BLOCKED_IP_KEY}.127.0.0.1"))
.config
.get(format!("{BLOCKED_IP_KEY}.127.0.0.1"))
.await
.unwrap(),
Some(String::new())
@@ -164,8 +164,8 @@ pub async fn test(params: &mut JMAPTest) {
server
.core
.storage
.data
.config_clear(format!("{BLOCKED_IP_KEY}.127.0.0.1"))
.config
.clear(format!("{BLOCKED_IP_KEY}.127.0.0.1"))
.await
.unwrap();
server

View File

@@ -405,19 +405,22 @@ async fn init_jmap_tests(store_id: &str, delete_if_exists: bool) -> JMAPTest {
config.resolve_macros().await;
// Parse servers
let servers = Servers::parse(&mut config);
let mut servers = Servers::parse(&mut config);
// Bind ports and drop privileges
servers.bind_and_drop_priv(&mut config);
// Build stores
let stores = Stores::parse(&mut config).await;
let stores = Stores::parse_all(&mut config).await;
// Parse core
let core = Core::parse(&mut config, stores).await;
let core = Core::parse(&mut config, stores, Default::default()).await;
let store = core.storage.data.clone();
let shared_core = core.into_shared();
// Parse acceptors
servers.parse_tcp_acceptors(&mut config, shared_core.clone());
// Init servers
let (delivery_tx, delivery_rx) = mpsc::channel(IPC_CHANNEL_BUFFER);
let smtp = SMTP::init(&mut config, shared_core.clone(), delivery_tx).await;
@@ -432,33 +435,34 @@ async fn init_jmap_tests(store_id: &str, delete_if_exists: bool) -> JMAPTest {
config.assert_no_errors();
// Spawn servers
let shutdown_tx = servers.spawn(
|server, shutdown_rx| {
match &server.protocol {
ServerProtocol::Smtp | ServerProtocol::Lmtp => server.spawn(
SmtpSessionManager::new(smtp.clone()),
shared_core.clone(),
shutdown_rx,
),
ServerProtocol::Http => server.spawn(
JmapSessionManager::new(jmap.clone()),
shared_core.clone(),
shutdown_rx,
),
ServerProtocol::Imap => server.spawn(
ImapSessionManager::new(imap.clone()),
shared_core.clone(),
shutdown_rx,
),
ServerProtocol::ManageSieve => server.spawn(
ManageSieveSessionManager::new(imap.clone()),
shared_core.clone(),
shutdown_rx,
),
};
},
store.clone(),
);
let shutdown_tx = servers.spawn(|server, acceptor, shutdown_rx| {
match &server.protocol {
ServerProtocol::Smtp | ServerProtocol::Lmtp => server.spawn(
SmtpSessionManager::new(smtp.clone()),
shared_core.clone(),
acceptor,
shutdown_rx,
),
ServerProtocol::Http => server.spawn(
JmapSessionManager::new(jmap.clone()),
shared_core.clone(),
acceptor,
shutdown_rx,
),
ServerProtocol::Imap => server.spawn(
ImapSessionManager::new(imap.clone()),
shared_core.clone(),
acceptor,
shutdown_rx,
),
ServerProtocol::ManageSieve => server.spawn(
ManageSieveSessionManager::new(imap.clone()),
shared_core.clone(),
acceptor,
shutdown_rx,
),
};
});
// Create tables
let directory = DirectoryStore {

View File

@@ -45,7 +45,7 @@ use jmap::{
};
use jmap_client::{mailbox::Role, push_subscription::Keys};
use jmap_proto::types::{id::Id, type_state::DataType};
use store::{ahash::AHashSet, Store};
use store::ahash::AHashSet;
use tokio::sync::mpsc;
use utils::config::Config;
@@ -119,18 +119,17 @@ pub async fn test(params: &mut JMAPTest) {
// Start mock push server
let mut settings = Config::new(add_test_certs(SERVER)).unwrap();
settings.resolve_macros().await;
let servers = Servers::parse(&mut settings);
let mock_core = Core::default().into_shared();
let mut servers = Servers::parse(&mut settings);
servers.parse_tcp_acceptors(&mut settings, mock_core.clone());
// Start JMAP server
let manager = SessionManager::from(push_server.clone());
servers.bind_and_drop_priv(&mut settings);
settings.assert_no_errors();
let _shutdown_tx = servers.spawn(
|server, shutdown_rx| {
server.spawn(manager.clone(), Core::default().into_shared(), shutdown_rx);
},
Store::default(),
);
let _shutdown_tx = servers.spawn(|server, acceptor, shutdown_rx| {
server.spawn(manager.clone(), mock_core.clone(), acceptor, shutdown_rx);
});
// Register push notification (no encryption)
let push_id = client
@@ -309,7 +308,7 @@ impl common::listener::SessionManager for SessionManager {
session
.instance
.acceptor
.accept(session.stream)
.accept(session.stream, false)
.await
.unwrap_tls()
.await

View File

@@ -29,7 +29,6 @@ use common::{
smtp::{throttle::parse_throttle, *},
},
expr::{functions::ResolveVariable, if_block::*, tokenizer::TokenMap, *},
listener::TcpAcceptor,
Core,
};
use tokio::net::TcpSocket;
@@ -330,8 +329,6 @@ fn parse_servers() {
linger: None,
nodelay: true,
}],
acceptor: TcpAcceptor::Plain,
tls_implicit: false,
max_connections: 8192,
proxy_networks: vec![],
},
@@ -356,8 +353,6 @@ fn parse_servers() {
nodelay: true,
},
],
acceptor: TcpAcceptor::Plain,
tls_implicit: true,
max_connections: 1024,
proxy_networks: vec![],
},
@@ -372,8 +367,6 @@ fn parse_servers() {
linger: None,
nodelay: true,
}],
acceptor: TcpAcceptor::Plain,
tls_implicit: true,
max_connections: 8192,
proxy_networks: vec![],
},
@@ -390,11 +383,6 @@ fn parse_servers() {
"failed for {}",
expected_server.id
);
assert_eq!(
server.tls_implicit, expected_server.tls_implicit,
"failed for {}",
expected_server.id
);
for (listener, expected_listener) in
server.listeners.into_iter().zip(expected_server.listeners)
{

View File

@@ -210,8 +210,8 @@ async fn antispam() {
// Parse config
let mut config = Config::new(&config).unwrap();
config.resolve_macros().await;
let stores = Stores::parse(&mut config).await;
let core = Core::parse(&mut config, stores).await;
let stores = Stores::parse_all(&mut config).await;
let core = Core::parse(&mut config, stores, Default::default()).await;
// Add mock DNS entries
for (domain, ip) in [

View File

@@ -94,8 +94,8 @@ async fn auth() {
let tmp_dir = TempDir::new("smtp_auth_test", true);
let mut config = Config::new(tmp_dir.update_config(CONFIG)).unwrap();
let stores = Stores::parse(&mut config).await;
let core = Core::parse(&mut config, stores).await;
let stores = Stores::parse_all(&mut config).await;
let core = Core::parse(&mut config, stores, Default::default()).await;
// EHLO should not advertise plain text auth without TLS
let mut session = Session::test(build_smtp(core, Inner::default()));

View File

@@ -126,8 +126,8 @@ async fn data() {
let mut inner = Inner::default();
let tmp_dir = TempDir::new("smtp_data_test", true);
let mut config = Config::new(tmp_dir.update_config(CONFIG)).unwrap();
let stores = Stores::parse(&mut config).await;
let core = Core::parse(&mut config, stores).await;
let stores = Stores::parse_all(&mut config).await;
let core = Core::parse(&mut config, stores, Default::default()).await;
let mut qr = inner.init_test_queue(&core);
// Test queue message builder

View File

@@ -112,8 +112,8 @@ async fn dmarc() {
let mut inner = Inner::default();
let tmp_dir = TempDir::new("smtp_dmarc_test", true);
let mut config = Config::new(tmp_dir.update_config(CONFIG.to_string() + SIGNATURES)).unwrap();
let stores = Stores::parse(&mut config).await;
let core = Core::parse(&mut config, stores).await;
let stores = Stores::parse_all(&mut config).await;
let core = Core::parse(&mut config, stores, Default::default()).await;
// Create temp dir for queue
let mut qr = inner.init_test_queue(&core);

View File

@@ -56,7 +56,7 @@ ehlo = [{if = "remote_ip = '10.0.0.2'", then = 'strict'},
#[tokio::test]
async fn ehlo() {
let mut config = Config::new(CONFIG).unwrap();
let core = Core::parse(&mut config, Default::default()).await;
let core = Core::parse(&mut config, Default::default(), Default::default()).await;
core.smtp.resolvers.dns.txt_add(
"mx1.foobar.org",
Spf::parse(b"v=spf1 ip4:10.0.0.1 -all").unwrap(),

View File

@@ -47,7 +47,7 @@ duration = [{if = "remote_ip = '10.0.0.3'", then = '500ms'},
#[tokio::test]
async fn limits() {
let mut config = Config::new(CONFIG).unwrap();
let core = Core::parse(&mut config, Default::default()).await;
let core = Core::parse(&mut config, Default::default(), Default::default()).await;
let (_tx, rx) = watch::channel(true);

View File

@@ -89,8 +89,8 @@ enable = true
async fn mail() {
let tmp_dir = TempDir::new("smtp_mail_test", true);
let mut config = Config::new(tmp_dir.update_config(CONFIG)).unwrap();
let stores = Stores::parse(&mut config).await;
let core = Core::parse(&mut config, stores).await;
let stores = Stores::parse_all(&mut config).await;
let core = Core::parse(&mut config, stores, Default::default()).await;
core.smtp.resolvers.dns.txt_add(
"foobar.org",
Spf::parse(b"v=spf1 ip4:10.0.0.1 -all").unwrap(),

View File

@@ -98,8 +98,8 @@ async fn milter_session() {
// Configure tests
let tmp_dir = TempDir::new("smtp_milter_test", true);
let mut config = Config::new(tmp_dir.update_config(CONFIG)).unwrap();
let stores = Stores::parse(&mut config).await;
let core = Core::parse(&mut config, stores).await;
let stores = Stores::parse_all(&mut config).await;
let core = Core::parse(&mut config, stores, Default::default()).await;
let _rx = spawn_mock_milter_server();
tokio::time::sleep(Duration::from_millis(100)).await;
let mut inner = Inner::default();

View File

@@ -112,8 +112,8 @@ async fn rcpt() {
.unwrap();*/
let tmp_dir = TempDir::new("smtp_rcpt_test", true);
let mut config = Config::new(tmp_dir.update_config(CONFIG)).unwrap();
let stores = Stores::parse(&mut config).await;
let core = Core::parse(&mut config, stores).await;
let stores = Stores::parse_all(&mut config).await;
let core = Core::parse(&mut config, stores, Default::default()).await;
// RCPT without MAIL FROM
let mut session = Session::test(build_smtp(core, Inner::default()));

View File

@@ -92,7 +92,7 @@ async fn address_rewrite() {
// Prepare config
let mut config = Config::new(CONFIG).unwrap();
let core = Core::parse(&mut config, Default::default()).await;
let core = Core::parse(&mut config, Default::default(), Default::default()).await;
// Init session
let mut session = Session::test(build_smtp(core, Inner::default()));

View File

@@ -150,8 +150,8 @@ async fn sieve_scripts() {
)
.unwrap();
config.resolve_macros().await;
let stores = Stores::parse(&mut config).await;
let core = Core::parse(&mut config, stores).await;
let stores = Stores::parse_all(&mut config).await;
let core = Core::parse(&mut config, stores, Default::default()).await;
let mut qr = inner.init_test_queue(&core);
// Build session

View File

@@ -144,8 +144,8 @@ verify = "relaxed"
async fn sign_and_seal() {
let tmp_dir = TempDir::new("smtp_sign_test", true);
let mut config = Config::new(tmp_dir.update_config(CONFIG.to_string() + SIGNATURES)).unwrap();
let stores = Stores::parse(&mut config).await;
let core = Core::parse(&mut config, stores).await;
let stores = Stores::parse_all(&mut config).await;
let core = Core::parse(&mut config, stores, Default::default()).await;
let mut inner = Inner::default();
// Create temp dir for queue

View File

@@ -72,8 +72,8 @@ async fn throttle_inbound() {
let tmp_dir = TempDir::new("smtp_inbound_throttle", true);
let mut config = Config::new(tmp_dir.update_config(CONFIG)).unwrap();
let stores = Stores::parse(&mut config).await;
let core = Core::parse(&mut config, stores).await;
let stores = Stores::parse_all(&mut config).await;
let core = Core::parse(&mut config, stores, Default::default()).await;
let inner = Inner::default();
// Test connection concurrency limit

View File

@@ -84,8 +84,8 @@ expn = [{if = "remote_ip = '10.0.0.1'", then = true},
async fn vrfy_expn() {
let tmp_dir = TempDir::new("smtp_vrfy_test", true);
let mut config = Config::new(tmp_dir.update_config(CONFIG)).unwrap();
let stores = Stores::parse(&mut config).await;
let core = Core::parse(&mut config, stores).await;
let stores = Stores::parse_all(&mut config).await;
let core = Core::parse(&mut config, stores, Default::default()).await;
// EHLO should not advertise VRFY/EXPN to 10.0.0.2
let mut session = Session::test(build_smtp(core, Inner::default()));

View File

@@ -127,10 +127,10 @@ async fn lookup_sql() {
// Parse settings
let temp_dir = TempDir::new("smtp_lookup_tests", true);
let mut config = Config::new(temp_dir.update_config(CONFIG)).unwrap();
let stores = Stores::parse(&mut config).await;
let stores = Stores::parse_all(&mut config).await;
let inner = Inner::default();
let core = Core::parse(&mut config, stores).await;
let core = Core::parse(&mut config, stores, Default::default()).await;
core.smtp.resolvers.dns.mx_add(
"test.org",

View File

@@ -80,7 +80,7 @@ async fn lookup_ip() {
];
let mut config = Config::new(CONFIG_V4).unwrap();
let core = build_smtp(
Core::parse(&mut config, Default::default()).await,
Core::parse(&mut config, Default::default(), Default::default()).await,
Inner::default(),
);
core.core.smtp.resolvers.dns.ipv4_add(
@@ -117,7 +117,7 @@ async fn lookup_ip() {
// Ipv6 strategy
let mut config = Config::new(CONFIG_V6).unwrap();
let core = build_smtp(
Core::parse(&mut config, Default::default()).await,
Core::parse(&mut config, Default::default(), Default::default()).await,
Inner::default(),
);
core.core.smtp.resolvers.dns.ipv4_add(

View File

@@ -50,7 +50,7 @@ pub mod smtp;
pub mod throttle;
pub mod tls;
const SERVER: &str = "
const CONFIG: &str = r#"
[server]
hostname = 'mx.example.org'
greeting = 'Test SMTP instance'
@@ -80,9 +80,7 @@ certificate = 'default'
[certificate.default]
cert = '%{file:{CERT}}%'
private-key = '%{file:{PK}}%'
";
const STORES: &str = r#"
[storage]
data = "sqlite"
lookup = "sqlite"
@@ -106,9 +104,10 @@ impl TestServer {
pub async fn new(name: &str, config: impl AsRef<str>, with_receiver: bool) -> TestServer {
let temp_dir = TempDir::new(name, true);
let mut config =
Config::new(temp_dir.update_config(STORES.to_string() + config.as_ref())).unwrap();
let stores = Stores::parse(&mut config).await;
let core = Core::parse(&mut config, stores).await;
Config::new(temp_dir.update_config(add_test_certs(CONFIG) + config.as_ref())).unwrap();
config.resolve_macros().await;
let stores = Stores::parse_all(&mut config).await;
let core = Core::parse(&mut config, stores, Default::default()).await;
let mut inner = Inner::default();
let qr = if with_receiver {
inner.init_test_queue(&core)
@@ -137,9 +136,9 @@ impl TestServer {
pub async fn start(&self, protocols: &[ServerProtocol]) -> watch::Sender<bool> {
// Spawn listeners
let mut config = Config::new(add_test_certs(SERVER)).unwrap();
config.resolve_macros().await;
let mut config = Config::new(CONFIG).unwrap();
let mut servers = Servers::parse(&mut config);
servers.parse_tcp_acceptors(&mut config, self.instance.core.clone());
// Filter out protocols
servers
@@ -152,24 +151,25 @@ impl TestServer {
let instance = self.instance.clone();
let smtp_manager = SmtpSessionManager::new(instance.clone());
let smtp_admin_manager = SmtpAdminSessionManager::new(instance.clone());
servers.spawn(
|server, shutdown_rx| {
match &server.protocol {
ServerProtocol::Smtp | ServerProtocol::Lmtp => {
server.spawn(smtp_manager.clone(), instance.core.clone(), shutdown_rx)
}
ServerProtocol::Http => server.spawn(
smtp_admin_manager.clone(),
instance.core.clone(),
shutdown_rx,
),
ServerProtocol::Imap | ServerProtocol::ManageSieve => {
unreachable!()
}
};
},
instance.core.load().storage.data.clone(),
)
servers.spawn(|server, acceptor, shutdown_rx| {
match &server.protocol {
ServerProtocol::Smtp | ServerProtocol::Lmtp => server.spawn(
smtp_manager.clone(),
instance.core.clone(),
acceptor,
shutdown_rx,
),
ServerProtocol::Http => server.spawn(
smtp_admin_manager.clone(),
instance.core.clone(),
acceptor,
shutdown_rx,
),
ServerProtocol::Imap | ServerProtocol::ManageSieve => {
unreachable!()
}
};
})
}
pub fn new_session(&self) -> Session<DummyIo> {

View File

@@ -357,14 +357,21 @@ pub trait TestServerInstance {
impl TestServerInstance for ServerInstance {
fn test_with_shutdown(shutdown_rx: watch::Receiver<bool>) -> Self {
let tls_config = Arc::new(
ServerConfig::builder()
.with_no_client_auth()
.with_cert_resolver(Arc::new(DummyCertResolver)),
);
Self {
id: "smtp".to_string(),
protocol: ServerProtocol::Smtp,
acceptor: TcpAcceptor::Tls(TlsAcceptor::from(Arc::new(
ServerConfig::builder()
.with_no_client_auth()
.with_cert_resolver(Arc::new(DummyCertResolver)),
))),
acceptor: TcpAcceptor::Tls {
acme_config: tls_config.clone(),
default_config: tls_config.clone(),
acceptor: TlsAcceptor::from(tls_config),
implicit: false,
},
limiter: ConcurrencyLimiter::new(100),
shutdown_rx,
proxy_networks: vec![],

View File

@@ -35,7 +35,7 @@ pub async fn blob_tests() {
let temp_dir = TempDir::new("blob_tests", true);
let mut config =
Config::new(CONFIG.replace("{TMP}", temp_dir.path.as_path().to_str().unwrap())).unwrap();
let stores = Stores::parse(&mut config).await;
let stores = Stores::parse_all(&mut config).await;
for (store_id, blob_store) in &stores.blob_stores {
println!("Testing blob store {}...", store_id);

View File

@@ -38,7 +38,7 @@ pub async fn lookup_tests() {
Config::new(CONFIG.replace("{TMP}", temp_dir.path.as_path().to_str().unwrap()))
.unwrap()
.assert_no_errors();
let stores = Stores::parse(&mut config).await;
let stores = Stores::parse_all(&mut config).await;
let rate = Rate {
requests: 1,
period: Duration::from_secs(1),

View File

@@ -92,7 +92,7 @@ pub async fn store_tests() {
let mut config = Config::new(CONFIG.replace("{TMP}", &temp_dir.path.to_string_lossy()))
.unwrap()
.assert_no_errors();
let stores = Stores::parse(&mut config).await;
let stores = Stores::parse_all(&mut config).await;
let store_id = std::env::var("STORE")
.expect("Missing store type. Try running `STORE=<store_type> cargo test`");