use std::env;
use std::net::{IpAddr, Ipv4Addr, SocketAddr};
use std::path::PathBuf;
use std::time::Duration;
use anyhow::{Context, bail};
use crate::mail::MailConfig;
use crate::storage::StorageConfig;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[non_exhaustive]
pub enum Environment {
Local,
Testing,
Production,
}
impl Environment {
fn parse(value: &str) -> anyhow::Result<Self> {
match value.to_ascii_lowercase().as_str() {
"local" | "dev" | "development" => Ok(Self::Local),
"testing" | "test" => Ok(Self::Testing),
"production" | "prod" => Ok(Self::Production),
other => bail!("APP_ENV must be local, testing or production, got `{other}`"),
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
#[non_exhaustive]
pub enum CspMode {
#[default]
Relaxed,
Strict,
Off,
}
impl CspMode {
fn parse(value: &str) -> anyhow::Result<Self> {
match value.to_ascii_lowercase().as_str() {
"relaxed" | "" => Ok(Self::Relaxed),
"strict" => Ok(Self::Strict),
"off" | "false" | "none" => Ok(Self::Off),
other => bail!("CSP must be relaxed, strict or off, got `{other}`"),
}
}
}
setting_enum! {
pub enum SessionDriver ("SESSION_DRIVER") {
Cookie = "cookie",
Database = "database",
}
}
setting_enum! {
pub enum LogFormat ("LOG_FORMAT") {
Text = "text",
Json = "json",
}
}
setting_enum! {
pub enum CacheStore ("CACHE_STORE") {
Memory = "memory",
Database = "database",
}
}
#[derive(Clone)]
#[non_exhaustive]
pub struct Config {
pub name: String,
pub env: Environment,
pub debug: bool,
pub url: String,
pub key: Option<String>,
pub host: IpAddr,
pub port: u16,
pub views_path: PathBuf,
pub public_path: PathBuf,
pub session_lifetime: Duration,
pub session_cookie: String,
pub session_driver: SessionDriver,
pub remember_lifetime: Duration,
pub database_url: String,
pub database_pool_size: u32,
pub database_acquire_timeout: Duration,
pub database_statement_timeout: Option<Duration>,
pub request_timeout: Option<Duration>,
pub locale: String,
pub fallback_locale: String,
pub currency: String,
pub lang_path: PathBuf,
pub mail: MailConfig,
pub queue_workers: usize,
pub scheduler: bool,
pub log_format: LogFormat,
pub log_file: Option<PathBuf>,
pub timezone: crate::timezone::Zone,
pub cache_store: CacheStore,
pub storage_path: PathBuf,
pub storage: StorageConfig,
pub upload_max_size: usize,
pub csp: CspMode,
pub vars: std::collections::HashMap<String, String>,
pub analytics: AnalyticsConfig,
pub trusted_proxies: crate::TrustedProxies,
pub trusted_hosts: Vec<String>,
}
impl std::fmt::Debug for Config {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("Config")
.field("name", &self.name)
.field("env", &self.env)
.field("debug", &self.debug)
.field("url", &self.url)
.field("key", &self.key.as_ref().map(|_| "[hidden]"))
.field("host", &self.host)
.field("port", &self.port)
.field("views_path", &self.views_path)
.field("public_path", &self.public_path)
.field("session_lifetime", &self.session_lifetime)
.field("session_cookie", &self.session_cookie)
.field("session_driver", &self.session_driver)
.field("remember_lifetime", &self.remember_lifetime)
.field("database_url", &self.database_url)
.field("database_pool_size", &self.database_pool_size)
.field("database_acquire_timeout", &self.database_acquire_timeout)
.field(
"database_statement_timeout",
&self.database_statement_timeout,
)
.field("request_timeout", &self.request_timeout)
.field("locale", &self.locale)
.field("fallback_locale", &self.fallback_locale)
.field("currency", &self.currency)
.field("lang_path", &self.lang_path)
.field("mail", &self.mail)
.field("queue_workers", &self.queue_workers)
.field("scheduler", &self.scheduler)
.field("log_format", &self.log_format)
.field("log_file", &self.log_file)
.field("timezone", &self.timezone)
.field("cache_store", &self.cache_store)
.field("storage_path", &self.storage_path)
.field("storage", &self.storage)
.field("upload_max_size", &self.upload_max_size)
.field("csp", &self.csp)
.field(
"vars",
&self.vars.keys().collect::<std::collections::BTreeSet<_>>(),
)
.field("analytics", &self.analytics)
.field("trusted_proxies", &self.trusted_proxies)
.field("trusted_hosts", &self.trusted_hosts)
.finish()
}
}
#[derive(Clone, Default)]
#[non_exhaustive]
pub struct AnalyticsConfig {
pub google_site_verification: Option<String>,
pub ga4_measurement_id: Option<String>,
pub ga4_api_secret: Option<String>,
pub gtm_container_id: Option<String>,
}
impl std::fmt::Debug for AnalyticsConfig {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("AnalyticsConfig")
.field("google_site_verification", &self.google_site_verification)
.field("ga4_measurement_id", &self.ga4_measurement_id)
.field(
"ga4_api_secret",
&self.ga4_api_secret.as_ref().map(|_| "[hidden]"),
)
.field("gtm_container_id", &self.gtm_container_id)
.finish()
}
}
impl Config {
pub fn load() -> crate::Result<Self> {
let _ = dotenvy::dotenv();
Self::from_env()
}
pub fn from_env() -> crate::Result<Self> {
Self::from_vars(|name| env::var(name).ok())
}
pub fn from_vars(get: impl Fn(&str) -> Option<String>) -> crate::Result<Self> {
Ok(Self::read(get)?)
}
fn read(get: impl Fn(&str) -> Option<String>) -> anyhow::Result<Self> {
let v = Vars(get);
let env = Environment::parse(&v.or("APP_ENV", "local"))?;
let debug = v.bool("APP_DEBUG", env == Environment::Local)?;
let host = v
.or("APP_HOST", "127.0.0.1")
.parse()
.context("APP_HOST must be an IP address")?;
let port = v
.or("APP_PORT", "3000")
.parse()
.context("APP_PORT must be a port number")?;
let key = v.get("APP_KEY").filter(|k| !k.is_empty());
if env == Environment::Production && key.is_none() {
bail!("APP_KEY must be set in production");
}
Ok(Self {
name: v.or("APP_NAME", "Renox"),
env,
debug,
url: v.or("APP_URL", &format!("http://{host}:{port}")),
key,
host,
port,
views_path: v.or("VIEWS_PATH", "resources/views").into(),
public_path: v.or("PUBLIC_PATH", "public").into(),
session_lifetime: Duration::from_secs(
v.or("SESSION_LIFETIME", "120")
.parse::<u64>()
.context("SESSION_LIFETIME must be a number of minutes")?
.saturating_mul(60),
),
session_cookie: v.or("SESSION_COOKIE", "renox_session"),
session_driver: SessionDriver::parse(&v.or("SESSION_DRIVER", "cookie"))?,
remember_lifetime: Duration::from_secs(
v.or("REMEMBER_LIFETIME", "43200")
.parse::<u64>()
.context("REMEMBER_LIFETIME must be a number of minutes")?
.saturating_mul(60),
),
database_url: v.or("DATABASE_URL", "sqlite://storage/app.db"),
database_pool_size: v
.or("DATABASE_POOL_SIZE", "8")
.parse()
.ok()
.filter(|n| *n > 0)
.context("DATABASE_POOL_SIZE must be a number above 0")?,
database_acquire_timeout: v
.seconds("DATABASE_ACQUIRE_TIMEOUT", 5)?
.unwrap_or(Duration::from_secs(5)),
database_statement_timeout: v.seconds("DATABASE_STATEMENT_TIMEOUT", 30)?,
request_timeout: v.seconds("REQUEST_TIMEOUT", 60)?,
locale: v.or("APP_LOCALE", "en"),
fallback_locale: v.or("APP_FALLBACK_LOCALE", "en"),
currency: match v.or("APP_CURRENCY", "USD").trim().to_ascii_uppercase() {
code if code.len() == 3 && code.bytes().all(|b| b.is_ascii_uppercase()) => code,
other => {
bail!("APP_CURRENCY must be an ISO 4217 code like USD or EUR, got `{other}`")
}
},
lang_path: v.or("LANG_PATH", "resources/lang").into(),
mail: MailConfig {
mailer: crate::mail::MailDriver::parse(&v.or("MAIL_MAILER", "log"))?,
host: v.or("MAIL_HOST", "localhost"),
port: v
.get("MAIL_PORT")
.filter(|p| !p.is_empty())
.map(|p| p.parse())
.transpose()
.context("MAIL_PORT must be a port number")?,
username: v.get("MAIL_USERNAME").filter(|v| !v.is_empty()),
password: v.get("MAIL_PASSWORD").filter(|v| !v.is_empty()),
encryption: crate::mail::MailEncryption::parse(
&v.or("MAIL_ENCRYPTION", "starttls"),
)?,
from_address: v.or("MAIL_FROM_ADDRESS", "hello@example.com"),
from_name: v.get("MAIL_FROM_NAME").filter(|v| !v.is_empty()),
timeout: v
.seconds("MAIL_TIMEOUT", 10)?
.unwrap_or(Duration::from_secs(10)),
failover: v
.or("MAIL_FAILOVER", "")
.split(',')
.map(|name| name.trim().to_owned())
.filter(|name| !name.is_empty())
.collect(),
},
queue_workers: v
.or("QUEUE_WORKERS", "2")
.parse()
.context("QUEUE_WORKERS must be a number")?,
scheduler: v.bool("SCHEDULER", true)?,
log_format: LogFormat::parse(&v.or("LOG_FORMAT", "text"))?,
log_file: v
.get("LOG_FILE")
.filter(|p| !p.is_empty())
.map(PathBuf::from),
timezone: v
.or("APP_TIMEZONE", "UTC")
.parse()
.map_err(|err| anyhow::anyhow!("APP_TIMEZONE: {err}"))?,
cache_store: CacheStore::parse(&v.or("CACHE_STORE", "memory"))?,
storage_path: v.or("STORAGE_PATH", "storage").into(),
storage: StorageConfig {
disk: crate::storage::DiskDriver::parse(&v.or("STORAGE_DISK", "local"))?,
bucket: v.optional("S3_BUCKET"),
region: v.optional("S3_REGION"),
endpoint: v.optional("S3_ENDPOINT"),
access_key_id: v.optional("S3_ACCESS_KEY_ID"),
secret_access_key: v.optional("S3_SECRET_ACCESS_KEY"),
url: v.optional("STORAGE_URL"),
root: None,
},
upload_max_size: v
.or("UPLOAD_MAX_SIZE", "10")
.parse::<usize>()
.ok()
.and_then(|mb| mb.checked_mul(1024 * 1024))
.context("UPLOAD_MAX_SIZE must be a number of megabytes")?,
csp: CspMode::parse(&v.or("CSP", "relaxed"))?,
trusted_proxies: crate::TrustedProxies::read(&v.or("TRUSTED_PROXIES", ""))?,
trusted_hosts: v
.or("TRUSTED_HOSTS", "")
.split(',')
.map(|host| host.trim().to_ascii_lowercase())
.filter(|host| !host.is_empty())
.collect(),
vars: Default::default(),
analytics: AnalyticsConfig {
google_site_verification: v.optional("GOOGLE_SITE_VERIFICATION"),
ga4_measurement_id: v.optional("GA4_MEASUREMENT_ID"),
ga4_api_secret: v.optional("GA4_API_SECRET"),
gtm_container_id: v.optional("GTM_CONTAINER_ID"),
},
})
}
pub fn var(&self, name: &str) -> Option<String> {
self.vars
.get(name)
.cloned()
.or_else(|| env::var(name).ok())
.filter(|v| !v.is_empty())
}
pub fn addr(&self) -> SocketAddr {
SocketAddr::new(self.host, self.port)
}
}
impl Default for Config {
fn default() -> Self {
Self {
name: "Renox".into(),
env: Environment::Local,
debug: true,
url: "http://127.0.0.1:3000".into(),
key: None,
host: IpAddr::V4(Ipv4Addr::LOCALHOST),
port: 3000,
views_path: "resources/views".into(),
public_path: "public".into(),
session_lifetime: Duration::from_secs(120 * 60),
session_cookie: "renox_session".into(),
session_driver: SessionDriver::Cookie,
remember_lifetime: Duration::from_secs(43_200 * 60),
database_url: "sqlite::memory:".into(),
database_pool_size: 8,
database_acquire_timeout: Duration::from_secs(30),
database_statement_timeout: Some(Duration::from_secs(30)),
request_timeout: Some(Duration::from_secs(60)),
locale: "en".into(),
fallback_locale: "en".into(),
currency: "USD".into(),
lang_path: "resources/lang".into(),
mail: MailConfig {
mailer: crate::mail::MailDriver::Memory,
..MailConfig::default()
},
queue_workers: 0,
scheduler: false,
log_format: LogFormat::Text,
log_file: None,
timezone: crate::timezone::Zone::default(),
cache_store: CacheStore::Memory,
storage_path: "storage".into(),
storage: StorageConfig::default(),
upload_max_size: 10 * 1024 * 1024,
csp: CspMode::Relaxed,
vars: Default::default(),
analytics: AnalyticsConfig::default(),
trusted_proxies: Default::default(),
trusted_hosts: Vec::new(),
}
}
}
struct Vars<F>(F);
impl<F: Fn(&str) -> Option<String>> Vars<F> {
fn get(&self, name: &str) -> Option<String> {
(self.0)(name)
}
fn optional(&self, name: &str) -> Option<String> {
self.get(name).filter(|v| !v.is_empty())
}
fn or(&self, name: &str, default: &str) -> String {
self.get(name).unwrap_or_else(|| default.to_owned())
}
fn seconds(&self, name: &str, default: u64) -> anyhow::Result<Option<Duration>> {
let value: u64 = self
.or(name, &default.to_string())
.parse()
.with_context(|| format!("{name} must be a number of seconds"))?;
Ok((value > 0).then(|| Duration::from_secs(value)))
}
fn bool(&self, name: &str, default: bool) -> anyhow::Result<bool> {
match self.get(name) {
None => Ok(default),
Some(v) => match v.to_ascii_lowercase().as_str() {
"1" | "true" | "yes" | "on" => Ok(true),
"0" | "false" | "no" | "off" | "" => Ok(false),
other => bail!("{name} must be true or false, got `{other}`"),
},
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn parses_environment_aliases() {
assert_eq!(Environment::parse("local").unwrap(), Environment::Local);
assert_eq!(Environment::parse("DEV").unwrap(), Environment::Local);
assert_eq!(Environment::parse("test").unwrap(), Environment::Testing);
assert_eq!(Environment::parse("prod").unwrap(), Environment::Production);
assert!(Environment::parse("staging").is_err());
}
fn load(pairs: &[(&str, &str)]) -> anyhow::Result<Config> {
let vars: std::collections::HashMap<String, String> = pairs
.iter()
.map(|(k, v)| ((*k).to_owned(), (*v).to_owned()))
.collect();
Config::read(|name| vars.get(name).cloned())
}
#[test]
fn defaults_without_any_variable() {
let c = load(&[]).unwrap();
assert_eq!(c.env, Environment::Local);
assert!(c.debug, "local defaults to debug");
assert_eq!(
(c.port, c.database_pool_size, c.upload_max_size),
(3000, 8, 10 * 1024 * 1024)
);
assert_eq!(c.database_acquire_timeout, Duration::from_secs(5));
assert_eq!(c.request_timeout, Some(Duration::from_secs(60)));
assert_eq!(c.cache_store, CacheStore::Memory);
assert_eq!(c.csp, CspMode::Relaxed);
assert!(c.key.is_none() && c.mail.port.is_none());
}
#[test]
fn debug_hides_secrets() {
let mut c = load(&[("APP_KEY", "base64:c2VjcmV0"), ("MAIL_PASSWORD", "hunter2")]).unwrap();
c.vars
.insert("XENDIT_SECRET_KEY".into(), "xnd_live_123".into());
let shown = format!("{c:?}");
for secret in ["c2VjcmV0", "hunter2", "xnd_live_123"] {
assert!(!shown.contains(secret), "{secret} in {shown}");
}
assert!(
shown.contains("XENDIT_SECRET_KEY"),
"variable names stay: {shown}"
);
assert!(shown.contains("[hidden]"));
}
#[test]
fn session_driver() {
assert_eq!(load(&[]).unwrap().session_driver, SessionDriver::Cookie);
assert_eq!(
load(&[("SESSION_DRIVER", "database")])
.unwrap()
.session_driver,
SessionDriver::Database
);
assert!(load(&[("SESSION_DRIVER", "redis")]).is_err());
}
#[test]
fn log_format_and_file() {
let c = load(&[]).unwrap();
assert_eq!((c.log_format.as_str(), c.log_file), ("text", None));
let c = load(&[("LOG_FORMAT", "json"), ("LOG_FILE", "storage/logs/app.log")]).unwrap();
assert_eq!(c.log_format, LogFormat::Json);
assert_eq!(c.log_file, Some(PathBuf::from("storage/logs/app.log")));
let err = load(&[("LOG_FORMAT", "xml")]).unwrap_err();
assert!(err.to_string().contains("LOG_FORMAT"), "{err}");
}
#[test]
fn currency() {
assert_eq!(load(&[]).unwrap().currency, "USD");
assert_eq!(load(&[("APP_CURRENCY", " idr ")]).unwrap().currency, "IDR");
for bad in ["US", "rupiah", "U$D"] {
let err = load(&[("APP_CURRENCY", bad)]).unwrap_err();
assert!(err.to_string().contains("APP_CURRENCY"), "{err}");
}
}
#[test]
fn reads_every_kind_of_value() {
let c = load(&[
("APP_ENV", "production"),
("APP_KEY", "base64:abc"),
("APP_DEBUG", "off"),
("APP_PORT", "8080"),
("DATABASE_POOL_SIZE", "3"),
("REQUEST_TIMEOUT", "0"),
("DATABASE_STATEMENT_TIMEOUT", "5"),
("UPLOAD_MAX_SIZE", "2"),
("MAIL_PORT", "2525"),
("MAIL_USERNAME", ""),
("MAIL_FAILOVER", "backup, log"),
("CSP", "strict"),
("TRUSTED_PROXIES", "10.0.0.0/8"),
("TRUSTED_HOSTS", "shop.example.com, *.Example.org"),
("S3_BUCKET", "files"),
])
.unwrap();
assert_eq!(c.env, Environment::Production);
assert!(!c.debug);
assert_eq!((c.port, c.database_pool_size), (8080, 3));
assert_eq!(c.request_timeout, None, "0 means no limit");
assert_eq!(c.database_statement_timeout, Some(Duration::from_secs(5)));
assert_eq!(c.upload_max_size, 2 * 1024 * 1024);
assert_eq!(c.mail.port, Some(2525));
assert_eq!(c.mail.username, None, "empty means unset");
assert_eq!(c.mail.failover, ["backup", "log"]);
assert_eq!(c.csp, CspMode::Strict);
assert!(c.trusted_proxies.contains("10.1.2.3".parse().unwrap()));
assert_eq!(c.trusted_hosts, ["shop.example.com", "*.example.org"]);
assert_eq!(c.storage.bucket.as_deref(), Some("files"));
}
#[test]
fn refuses_bad_values_with_the_variable_name() {
for (name, value) in [
("APP_ENV", "staging"),
("APP_DEBUG", "maybe"),
("APP_PORT", "http"),
("APP_HOST", "localhost"),
("DATABASE_POOL_SIZE", "0"),
("REQUEST_TIMEOUT", "-1"),
("UPLOAD_MAX_SIZE", "99999999999999999"),
("MAIL_PORT", "99999"),
("CSP", "loose"),
("TRUSTED_PROXIES", "proxy.local"),
("SESSION_DRIVER", "redis"),
("CACHE_STORE", "redis"),
("LOG_FORMAT", "xml"),
("MAIL_MAILER", "sendgrid"),
("MAIL_ENCRYPTION", "ssl3"),
("STORAGE_DISK", "ftp"),
("APP_TIMEZONE", "Mars/Olympus"),
] {
let err = load(&[(name, value)])
.err()
.map(|e| format!("{e:#}"))
.unwrap_or_default();
assert!(
err.contains(name) || err.contains(value),
"{name}={value}: {err}"
);
}
let err = load(&[("APP_ENV", "production")]).unwrap_err();
assert!(
format!("{err}").contains("APP_KEY"),
"production needs a key"
);
}
#[test]
fn settings_print_and_serialize_as_their_word() {
assert_eq!(SessionDriver::Database.to_string(), "database");
assert_eq!(
serde_json::to_string(&CacheStore::Database).unwrap(),
r#""database""#
);
assert_eq!(format!("{}", LogFormat::Json), "json");
}
}