use std::collections::BTreeMap;
use std::fmt;
use std::net::SocketAddr;
use std::path::{Path, PathBuf};
use serde::de::DeserializeOwned;
use serde::{Deserialize, Deserializer};
use crate::db::Dialect;
use crate::error::Error;
pub const CONFIG_PATH_ENV: &str = "NBS_CONFIG";
pub const DEFAULT_CONFIG_FILE: &str = "config.toml";
pub const ENV_PREFIX: &str = "NBS__";
#[derive(Clone, Default, PartialEq, Eq)]
pub struct SecretString(String);
impl SecretString {
pub fn new(secret: impl Into<String>) -> Self {
Self(secret.into())
}
pub fn expose(&self) -> &str {
&self.0
}
pub fn is_empty(&self) -> bool {
self.0.is_empty()
}
}
impl fmt::Debug for SecretString {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str(if self.0.is_empty() { "<empty>" } else { "<redacted>" })
}
}
impl<'de> Deserialize<'de> for SecretString {
fn deserialize<D: Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
String::deserialize(deserializer).map(SecretString)
}
}
#[derive(Clone, Default, Deserialize)]
#[serde(default, deny_unknown_fields)]
#[non_exhaustive]
pub struct Config {
pub server: ServerConfig,
pub database: DatabaseConfig,
pub http: HttpConfig,
pub cors: CorsConfig,
pub log: LogConfig,
pub metrics: MetricsConfig,
pub openapi: OpenApiConfig,
pub ws: WsConfig,
pub modules: toml::Table,
}
impl fmt::Debug for Config {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
let modules: BTreeMap<&str, Vec<&str>> = self
.modules
.iter()
.map(|(name, value)| (name.as_str(), value.as_table().map(|t| t.keys().map(String::as_str).collect()).unwrap_or_default()))
.collect();
f.debug_struct("Config")
.field("server", &self.server)
.field("database", &self.database)
.field("http", &self.http)
.field("cors", &self.cors)
.field("log", &self.log)
.field("metrics", &self.metrics)
.field("openapi", &self.openapi)
.field("ws", &self.ws)
.field("modules (keys only)", &modules)
.finish()
}
}
#[derive(Clone, Debug, Deserialize)]
#[serde(default, deny_unknown_fields)]
#[non_exhaustive]
pub struct ServerConfig {
pub bind: SocketAddr,
pub shutdown_grace_secs: u64,
pub hook_timeout_ms: u64,
pub header_read_timeout_secs: u64,
pub module_start_timeout_secs: u64,
pub module_shutdown_timeout_secs: u64,
}
impl Default for ServerConfig {
fn default() -> Self {
Self {
bind: SocketAddr::from(([127, 0, 0, 1], 8080)),
shutdown_grace_secs: 20,
hook_timeout_ms: 2000,
header_read_timeout_secs: 15,
module_start_timeout_secs: 30,
module_shutdown_timeout_secs: 10,
}
}
}
#[derive(Clone, Debug, Deserialize)]
#[serde(default, deny_unknown_fields)]
#[non_exhaustive]
pub struct DatabaseConfig {
pub url: SecretString,
pub url_file: Option<PathBuf>,
pub max_connections: u32,
pub min_connections: u32,
pub acquire_timeout_secs: u64,
pub connect_lazy: bool,
pub migrate_on_start: bool,
pub migrations_dir: PathBuf,
pub migrate_lock_timeout_secs: u64,
}
impl Default for DatabaseConfig {
fn default() -> Self {
Self {
url: SecretString::default(),
url_file: None,
max_connections: 10,
min_connections: 0,
acquire_timeout_secs: 5,
connect_lazy: false,
migrate_on_start: false,
migrations_dir: PathBuf::from("migrations"),
migrate_lock_timeout_secs: 60,
}
}
}
impl DatabaseConfig {
pub fn dialect(&self) -> Option<Dialect> {
Dialect::from_url(self.url.expose())
}
}
#[derive(Clone, Debug, Deserialize)]
#[serde(default, deny_unknown_fields)]
#[non_exhaustive]
pub struct HttpConfig {
pub body_limit_bytes: usize,
pub request_timeout_secs: u64,
pub trust_request_id: bool,
pub max_body_bytes: usize,
pub trusted_proxies: Vec<String>,
}
impl Default for HttpConfig {
fn default() -> Self {
Self {
body_limit_bytes: net_backend_protocol::routes::DEFAULT_BODY_LIMIT_BYTES,
request_timeout_secs: 30,
trust_request_id: false,
max_body_bytes: 32 * 1024 * 1024,
trusted_proxies: Vec::new(),
}
}
}
#[derive(Clone, Debug, Deserialize)]
#[serde(default, deny_unknown_fields)]
#[non_exhaustive]
pub struct CorsConfig {
pub allowed_origins: Vec<String>,
pub max_age_secs: u64,
}
impl Default for CorsConfig {
fn default() -> Self {
Self { allowed_origins: Vec::new(), max_age_secs: 600 }
}
}
#[derive(Clone, Debug, Deserialize)]
#[serde(default, deny_unknown_fields)]
#[non_exhaustive]
pub struct LogConfig {
pub level: String,
pub format: LogFormat,
}
impl Default for LogConfig {
fn default() -> Self {
Self { level: "info".into(), format: LogFormat::Pretty }
}
}
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq, Deserialize)]
#[serde(rename_all = "lowercase")]
#[non_exhaustive]
pub enum LogFormat {
#[default]
Pretty,
Json,
}
#[derive(Clone, Debug, Deserialize)]
#[serde(default, deny_unknown_fields)]
#[non_exhaustive]
pub struct MetricsConfig {
pub enabled: bool,
pub bind: SocketAddr,
}
impl Default for MetricsConfig {
fn default() -> Self {
Self { enabled: false, bind: SocketAddr::from(([127, 0, 0, 1], 9100)) }
}
}
#[derive(Clone, Debug, Deserialize)]
#[serde(default, deny_unknown_fields)]
#[non_exhaustive]
pub struct OpenApiConfig {
pub enabled: bool,
pub ui: bool,
pub ui_script_url: Option<String>,
pub ui_script_integrity: Option<String>,
pub title: String,
pub version: String,
}
impl Default for OpenApiConfig {
fn default() -> Self {
Self { enabled: true, ui: false, ui_script_url: None, ui_script_integrity: None, title: "Game backend API".into(), version: "1".into() }
}
}
#[derive(Clone, Debug, Deserialize)]
#[serde(default, deny_unknown_fields)]
#[non_exhaustive]
pub struct WsConfig {
pub enabled: bool,
pub max_connections: usize,
pub max_connections_per_user: usize,
pub max_connections_per_ip: usize,
pub max_pending_connections: usize,
pub roles_refresh_secs: u64,
pub handshakes_per_ip_per_minute: u32,
pub auth_timeout_secs: u64,
pub ping_interval_secs: u64,
pub idle_timeout_secs: u64,
pub request_timeout_secs: u64,
pub write_timeout_secs: u64,
pub outbox_frames: usize,
pub frames_per_second: u32,
pub frame_burst: u32,
pub max_message_bytes: usize,
pub read_buffer_bytes: usize,
pub max_rooms_per_connection: usize,
pub max_room_members: usize,
pub query_token: bool,
}
impl Default for WsConfig {
fn default() -> Self {
Self {
enabled: true,
max_connections: 10_000,
max_connections_per_user: 5,
max_connections_per_ip: 100,
max_pending_connections: 1000,
roles_refresh_secs: 60,
handshakes_per_ip_per_minute: 60,
auth_timeout_secs: net_backend_protocol::envelope::AUTH_TIMEOUT_SECS,
ping_interval_secs: 20,
idle_timeout_secs: 60,
request_timeout_secs: 10,
write_timeout_secs: 10,
outbox_frames: 256,
frames_per_second: 20,
frame_burst: 40,
max_message_bytes: net_backend_protocol::envelope::MAX_MESSAGE_BYTES,
read_buffer_bytes: 8 * 1024,
max_rooms_per_connection: net_backend_protocol::chat::DEFAULT_MAX_JOINED_ROOMS as usize,
max_room_members: net_backend_protocol::chat::DEFAULT_MAX_ROOM_MEMBERS as usize,
query_token: false,
}
}
}
impl WsConfig {
fn problems(&self, problems: &mut Vec<String>) {
let ranges: [(&str, u64, u64, u64); 7] = [
("ws.auth_timeout_secs", self.auth_timeout_secs, 1, 300),
("ws.ping_interval_secs", self.ping_interval_secs, 1, 3600),
("ws.idle_timeout_secs", self.idle_timeout_secs, 2, 7200),
("ws.request_timeout_secs", self.request_timeout_secs, 1, 3600),
("ws.write_timeout_secs", self.write_timeout_secs, 1, 3600),
("ws.frames_per_second", u64::from(self.frames_per_second), 1, 100_000),
("ws.frame_burst", u64::from(self.frame_burst), 1, 100_000),
];
for (name, value, min, max) in ranges {
if !(min..=max).contains(&value) {
problems.push(format!("{name} must be between {min} and {max}"));
}
}
if self.idle_timeout_secs <= self.ping_interval_secs {
problems.push("ws.idle_timeout_secs must be greater than ws.ping_interval_secs".into());
} else if self.request_timeout_secs >= self.idle_timeout_secs - self.ping_interval_secs {
problems.push("ws.request_timeout_secs must be less than ws.idle_timeout_secs - ws.ping_interval_secs".into());
}
if self.roles_refresh_secs > 86_400 {
problems.push("ws.roles_refresh_secs must be at most 86400".into());
}
let sizes: [(&str, usize, usize, usize); 9] = [
("ws.max_connections", self.max_connections, 1, 10_000_000),
("ws.max_connections_per_user", self.max_connections_per_user, 1, 10_000),
("ws.max_connections_per_ip", self.max_connections_per_ip, 1, 10_000_000),
("ws.max_pending_connections", self.max_pending_connections, 1, 10_000_000),
("ws.outbox_frames", self.outbox_frames, 4, 1 << 20),
("ws.max_message_bytes", self.max_message_bytes, 1024, 64 * 1024 * 1024),
("ws.read_buffer_bytes", self.read_buffer_bytes, 1024, 1024 * 1024),
("ws.max_rooms_per_connection", self.max_rooms_per_connection, 1, 100_000),
("ws.max_room_members", self.max_room_members, 1, 10_000_000),
];
for (name, value, min, max) in sizes {
if !(min..=max).contains(&value) {
problems.push(format!("{name} must be between {min} and {max}"));
}
}
}
}
impl Config {
pub fn load() -> Result<Config, Error> {
let path = match std::env::var_os(CONFIG_PATH_ENV) {
Some(path) => Some(PathBuf::from(path)),
None => Some(PathBuf::from(DEFAULT_CONFIG_FILE)).filter(|p| p.is_file()),
};
let text = match &path {
Some(path) => Some(read_file(path)?),
None => None,
};
Self::from_sources(text.as_deref(), std::env::vars())
}
pub fn from_file(path: impl AsRef<Path>) -> Result<Config, Error> {
let text = read_file(path.as_ref())?;
Self::from_sources(Some(&text), std::env::vars())
}
pub fn from_toml_str(text: &str) -> Result<Config, Error> {
Self::from_sources(Some(text), std::iter::empty::<(String, String)>())
}
pub fn from_sources<I, K, V>(toml_text: Option<&str>, env: I) -> Result<Config, Error>
where
I: IntoIterator<Item = (K, V)>,
K: AsRef<str>,
V: AsRef<str>,
{
let mut config = match toml_text {
Some(text) => toml::from_str::<Config>(text).map_err(|e| Error::Config(vec![redact_toml_error(&e)]))?,
None => Config::default(),
};
let mut problems = Vec::new();
let env: BTreeMap<String, String> =
env.into_iter().filter_map(|(k, v)| k.as_ref().strip_prefix(ENV_PREFIX).map(|k| (k.to_string(), v.as_ref().to_string()))).collect();
for (key, value) in &env {
if let Err(problem) = config.apply_env(key, value) {
problems.push(format!("{ENV_PREFIX}{key}: {problem}"));
}
}
if !problems.is_empty() {
return Err(Error::Config(problems));
}
config.resolve_secret_files()?;
config.validate()?;
Ok(config)
}
fn resolve_secret_files(&mut self) -> Result<(), Error> {
if let Some(path) = &self.database.url_file {
if !self.database.url.is_empty() {
return Err(Error::Config(vec!["database.url and database.url_file are both set; use one".into()]));
}
let text = read_file(path)?;
self.database.url = SecretString::new(text.trim_end());
self.database.url_file = None;
}
Ok(())
}
fn apply_env(&mut self, key: &str, value: &str) -> Result<(), String> {
let lower = key.to_ascii_lowercase();
let path: Vec<&str> = lower.split("__").collect();
match path.as_slice() {
["server", "bind"] => self.server.bind = parse(value)?,
["server", "shutdown_grace_secs"] => self.server.shutdown_grace_secs = parse(value)?,
["server", "hook_timeout_ms"] => self.server.hook_timeout_ms = parse(value)?,
["server", "header_read_timeout_secs"] => self.server.header_read_timeout_secs = parse(value)?,
["server", "module_start_timeout_secs"] => self.server.module_start_timeout_secs = parse(value)?,
["server", "module_shutdown_timeout_secs"] => self.server.module_shutdown_timeout_secs = parse(value)?,
["database", "url"] => self.database.url = SecretString::new(value),
["database", "url_file"] => self.database.url_file = Some(PathBuf::from(value)),
["database", "max_connections"] => self.database.max_connections = parse(value)?,
["database", "min_connections"] => self.database.min_connections = parse(value)?,
["database", "acquire_timeout_secs"] => self.database.acquire_timeout_secs = parse(value)?,
["database", "connect_lazy"] => self.database.connect_lazy = parse(value)?,
["database", "migrate_on_start"] => self.database.migrate_on_start = parse(value)?,
["database", "migrations_dir"] => self.database.migrations_dir = PathBuf::from(value),
["database", "migrate_lock_timeout_secs"] => self.database.migrate_lock_timeout_secs = parse(value)?,
["http", "body_limit_bytes"] => self.http.body_limit_bytes = parse(value)?,
["http", "request_timeout_secs"] => self.http.request_timeout_secs = parse(value)?,
["http", "trust_request_id"] => self.http.trust_request_id = parse(value)?,
["http", "max_body_bytes"] => self.http.max_body_bytes = parse(value)?,
["http", "trusted_proxies"] => self.http.trusted_proxies = value.split(',').map(str::trim).filter(|s| !s.is_empty()).map(String::from).collect(),
["cors", "allowed_origins"] => self.cors.allowed_origins = value.split(',').map(str::trim).filter(|s| !s.is_empty()).map(String::from).collect(),
["cors", "max_age_secs"] => self.cors.max_age_secs = parse(value)?,
["log", "level"] => self.log.level = value.to_string(),
["log", "format"] => {
self.log.format = match value.to_ascii_lowercase().as_str() {
"pretty" => LogFormat::Pretty,
"json" => LogFormat::Json,
_ => return Err("expected `pretty` or `json`".into()),
}
}
["metrics", "enabled"] => self.metrics.enabled = parse(value)?,
["metrics", "bind"] => self.metrics.bind = parse(value)?,
["openapi", "enabled"] => self.openapi.enabled = parse(value)?,
["openapi", "ui"] => self.openapi.ui = parse(value)?,
["openapi", "ui_script_url"] => self.openapi.ui_script_url = Some(value.to_string()),
["openapi", "ui_script_integrity"] => self.openapi.ui_script_integrity = Some(value.to_string()),
["openapi", "title"] => self.openapi.title = value.to_string(),
["openapi", "version"] => self.openapi.version = value.to_string(),
["ws", "enabled"] => self.ws.enabled = parse(value)?,
["ws", "max_connections"] => self.ws.max_connections = parse(value)?,
["ws", "max_connections_per_user"] => self.ws.max_connections_per_user = parse(value)?,
["ws", "max_connections_per_ip"] => self.ws.max_connections_per_ip = parse(value)?,
["ws", "max_pending_connections"] => self.ws.max_pending_connections = parse(value)?,
["ws", "roles_refresh_secs"] => self.ws.roles_refresh_secs = parse(value)?,
["ws", "handshakes_per_ip_per_minute"] => self.ws.handshakes_per_ip_per_minute = parse(value)?,
["ws", "auth_timeout_secs"] => self.ws.auth_timeout_secs = parse(value)?,
["ws", "ping_interval_secs"] => self.ws.ping_interval_secs = parse(value)?,
["ws", "idle_timeout_secs"] => self.ws.idle_timeout_secs = parse(value)?,
["ws", "request_timeout_secs"] => self.ws.request_timeout_secs = parse(value)?,
["ws", "write_timeout_secs"] => self.ws.write_timeout_secs = parse(value)?,
["ws", "outbox_frames"] => self.ws.outbox_frames = parse(value)?,
["ws", "frames_per_second"] => self.ws.frames_per_second = parse(value)?,
["ws", "frame_burst"] => self.ws.frame_burst = parse(value)?,
["ws", "max_message_bytes"] => self.ws.max_message_bytes = parse(value)?,
["ws", "read_buffer_bytes"] => self.ws.read_buffer_bytes = parse(value)?,
["ws", "max_rooms_per_connection"] => self.ws.max_rooms_per_connection = parse(value)?,
["ws", "max_room_members"] => self.ws.max_room_members = parse(value)?,
["ws", "query_token"] => self.ws.query_token = parse(value)?,
["modules", module, rest @ ..] if !module.is_empty() && !rest.is_empty() && rest.iter().all(|s| !s.is_empty()) => {
let mut table = match self.modules.remove(*module) {
Some(toml::Value::Table(table)) => table,
None => toml::Table::new(),
Some(_) => return Err(format!("modules.{module} is not a table")),
};
insert_module_value(&mut table, rest, env_value(value))?;
self.modules.insert((*module).to_string(), toml::Value::Table(table));
}
_ => return Err("unknown setting".into()),
}
Ok(())
}
pub fn validate(&self) -> Result<(), Error> {
let mut problems = Vec::new();
let db = &self.database;
if Dialect::enabled().is_empty() {
problems.push("no database backend is compiled in: enable one of the features `mysql`, `postgres`, `sqlite`".into());
}
if db.url.is_empty() {
problems.push("database.url (or database.url_file) is required".into());
} else {
match db.dialect() {
None => problems.push("database.url: not a mysql://, mariadb://, postgres://, postgresql:// or sqlite: URL".into()),
Some(dialect) if !dialect.is_enabled() => problems.push(format!(
"database.url is a {} URL, but the `{}` feature of net_backend_server is not enabled",
dialect.display_name(),
dialect.name()
)),
Some(_) => {}
}
}
if !(1..=10_000).contains(&db.max_connections) {
problems.push("database.max_connections must be between 1 and 10000".into());
}
if db.min_connections > db.max_connections {
problems.push("database.min_connections must not exceed database.max_connections".into());
}
if !(1..=600).contains(&db.acquire_timeout_secs) {
problems.push("database.acquire_timeout_secs must be between 1 and 600".into());
}
if db.migrations_dir.as_os_str().is_empty() {
problems.push("database.migrations_dir must not be empty".into());
}
if self.server.shutdown_grace_secs > 3600 {
problems.push("server.shutdown_grace_secs must be at most 3600".into());
}
if !(1..=600_000).contains(&self.server.hook_timeout_ms) {
problems.push("server.hook_timeout_ms must be between 1 and 600000".into());
}
if !(1..=3600).contains(&self.server.header_read_timeout_secs) {
problems.push("server.header_read_timeout_secs must be between 1 and 3600".into());
}
if !(1..=3600).contains(&self.server.module_start_timeout_secs) || !(1..=3600).contains(&self.server.module_shutdown_timeout_secs) {
problems.push("server.module_start_timeout_secs and server.module_shutdown_timeout_secs must be between 1 and 3600".into());
}
if !(1..=3600).contains(&self.database.migrate_lock_timeout_secs) {
problems.push("database.migrate_lock_timeout_secs must be between 1 and 3600".into());
}
if !(1024..=64 * 1024 * 1024).contains(&self.http.body_limit_bytes) {
problems.push("http.body_limit_bytes must be between 1024 and 67108864 (raise it per route for uploads)".into());
}
if !(1024..=1024 * 1024 * 1024).contains(&self.http.max_body_bytes) || self.http.max_body_bytes < self.http.body_limit_bytes {
problems.push("http.max_body_bytes must be between 1024 and 1073741824 and not below http.body_limit_bytes".into());
}
for proxy in &self.http.trusted_proxies {
if crate::http::client_ip::IpNet::parse(proxy).is_none() {
problems.push(format!("http.trusted_proxies: `{proxy}` is not an address or address block like 10.0.0.0/8"));
}
}
if self.metrics.enabled && self.metrics.bind == self.server.bind {
problems.push("metrics.bind must differ from server.bind (metrics have their own listener)".into());
}
if self.openapi.ui {
let url_ok =
self.openapi.ui_script_url.as_deref().is_some_and(|u| u.starts_with("https://") && u.bytes().all(|b| b.is_ascii_graphic()) && !u.contains('"'));
let sri_ok = self.openapi.ui_script_integrity.as_deref().is_some_and(|i| {
(i.starts_with("sha256-") || i.starts_with("sha384-") || i.starts_with("sha512-"))
&& i.bytes().all(|b| b.is_ascii_alphanumeric() || matches!(b, b'-' | b'+' | b'/' | b'='))
});
if !url_ok || !sri_ok {
problems.push("openapi.ui needs openapi.ui_script_url (an https:// URL pinned to an exact version) and openapi.ui_script_integrity (sha256-/sha384-/sha512-…)".into());
}
}
if !(1..=3600).contains(&self.http.request_timeout_secs) {
problems.push("http.request_timeout_secs must be between 1 and 3600".into());
}
let origins = &self.cors.allowed_origins;
if origins.iter().any(|o| o == "*") && origins.len() > 1 {
problems.push("cors.allowed_origins: `*` must be the only entry".into());
}
for origin in origins.iter().filter(|o| *o != "*") {
let plain = origin.bytes().all(|b| b.is_ascii_graphic()) && !origin.ends_with('/');
if !(origin.starts_with("https://") || origin.starts_with("http://")) || !plain {
problems.push(format!("cors.allowed_origins: `{origin}` is not an origin like https://example.com"));
}
}
if tracing_subscriber::EnvFilter::try_new(&self.log.level).is_err() {
problems.push(format!("log.level: `{}` is not a valid filter", self.log.level));
}
if self.openapi.title.trim().is_empty() || self.openapi.version.trim().is_empty() {
problems.push("openapi.title and openapi.version must not be empty".into());
}
self.ws.problems(&mut problems);
if problems.is_empty() {
Ok(())
} else {
Err(Error::Config(problems))
}
}
pub fn module_config<T: DeserializeOwned>(&self, name: &str) -> Result<Option<T>, Error> {
match self.modules.get(name) {
None => Ok(None),
Some(value) => value.clone().try_into().map(Some).map_err(|e| Error::Config(vec![format!("modules.{name}: {}", mask_quoted(e.message()))])),
}
}
pub fn unknown_module_sections(&self, registered: &[&str]) -> Vec<String> {
self.modules.keys().filter(|name| !registered.contains(&name.as_str())).cloned().collect()
}
}
pub fn resolve_secret(name: &str, value: Option<SecretString>, file: Option<&Path>) -> Result<Option<SecretString>, Error> {
match (value, file) {
(Some(_), Some(_)) => Err(Error::Config(vec![format!("{name} and {name}_file are both set; use one")])),
(Some(value), None) => Ok(Some(value)),
(None, Some(path)) => Ok(Some(SecretString::new(read_file(path)?.trim_end()))),
(None, None) => Ok(None),
}
}
fn mask_quoted(message: &str) -> String {
let mut out = String::with_capacity(message.len());
let mut quote: Option<char> = None;
for c in message.chars() {
match quote {
Some(q) if c == q => {
out.push_str("<hidden>");
out.push(c);
quote = None;
}
Some(_) => {}
None => {
out.push(c);
if c == '"' || c == '\'' {
quote = Some(c);
}
}
}
}
if quote.is_some() {
out.push_str("<hidden>");
}
out
}
fn read_file(path: &Path) -> Result<String, Error> {
std::fs::read_to_string(path).map_err(|e| Error::io(format!("reading {}", path.display()), e))
}
fn parse<T: std::str::FromStr>(value: &str) -> Result<T, String>
where
T::Err: fmt::Display,
{
value.trim().parse::<T>().map_err(|e| format!("invalid value: {e}"))
}
fn redact_toml_error(error: &toml::de::Error) -> String {
match error.span() {
Some(span) => format!("config file: {} (at byte {})", error.message(), span.start),
None => format!("config file: {}", error.message()),
}
}
fn env_value(value: &str) -> toml::Value {
let doc = format!("v = {value}");
match toml::from_str::<toml::Table>(&doc) {
Ok(mut table) => table.remove("v").unwrap_or_else(|| toml::Value::String(value.to_string())),
Err(_) => toml::Value::String(value.to_string()),
}
}
fn insert_module_value(table: &mut toml::Table, path: &[&str], value: toml::Value) -> Result<(), String> {
match path {
[] => Err("empty key".into()),
[last] => {
table.insert((*last).to_string(), value);
Ok(())
}
[first, rest @ ..] => {
let entry = table.entry((*first).to_string()).or_insert_with(|| toml::Value::Table(toml::Table::new()));
match entry {
toml::Value::Table(inner) => insert_module_value(inner, rest, value),
_ => Err(format!("`{first}` is not a table")),
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
fn url_for_enabled() -> &'static str {
match Dialect::enabled().first() {
Some(Dialect::MySql) => "mysql://u:p@127.0.0.1/db",
Some(Dialect::Postgres) => "postgres://u:p@127.0.0.1/db",
_ => "sqlite::memory:",
}
}
fn problems(result: Result<Config, Error>) -> Vec<String> {
match result {
Err(Error::Config(problems)) => problems,
Err(other) => vec![format!("other error: {other}")],
Ok(_) => Vec::new(),
}
}
#[cfg(any(feature = "mysql", feature = "postgres", feature = "sqlite"))]
#[test]
fn defaults_need_only_a_url() {
let env = [("NBS__DATABASE__URL", url_for_enabled())];
let config = Config::from_sources(None, env);
assert!(config.is_ok(), "{:?}", problems(config));
let problems = problems(Config::from_sources(None, std::iter::empty::<(&str, &str)>()));
assert!(problems.iter().any(|p| p.contains("database.url")), "{problems:?}");
}
#[cfg(any(feature = "mysql", feature = "postgres", feature = "sqlite"))]
#[test]
fn file_then_env_override() {
let toml = format!("[server]\nbind = \"0.0.0.0:9000\"\n[database]\nurl = \"{}\"\nmax_connections = 4\n", url_for_enabled());
let env = [("NBS__DATABASE__MAX_CONNECTIONS", "7"), ("NBS__HTTP__REQUEST_TIMEOUT_SECS", "12"), ("OTHER", "x")];
let config = Config::from_sources(Some(&toml), env).unwrap();
assert_eq!(config.server.bind.port(), 9000);
assert_eq!(config.database.max_connections, 7);
assert_eq!(config.http.request_timeout_secs, 12);
assert_eq!(config.http.body_limit_bytes, 64 * 1024);
}
#[test]
fn unknown_keys_and_bad_values_are_reported_together() {
let toml = format!("[database]\nurl = \"{}\"\n", url_for_enabled());
let env = [("NBS__SERVER__BINDD", "x"), ("NBS__SERVER__BIND", "not an address"), ("NBS__DATABASE__CONNECT_LAZY", "maybe")];
let problems = problems(Config::from_sources(Some(&toml), env));
assert_eq!(problems.len(), 3, "{problems:?}");
assert!(problems[0].starts_with("NBS__DATABASE__CONNECT_LAZY"));
assert!(problems.iter().any(|p| p.contains("unknown setting")));
let problems = problems_of_toml("[server]\nbindd = \"1.2.3.4:5\"\n");
assert!(problems[0].contains("unknown field"), "{problems:?}");
}
fn problems_of_toml(toml: &str) -> Vec<String> {
problems(Config::from_toml_str(toml))
}
#[test]
fn validation_collects_every_problem() {
let mut config = Config::default();
config.database.url = SecretString::new("redis://x");
config.database.max_connections = 0;
config.database.min_connections = 5;
config.http.body_limit_bytes = 10;
config.http.request_timeout_secs = 0;
config.server.hook_timeout_ms = 0;
config.cors.allowed_origins = vec!["*".into(), "example.com".into()];
config.log.level = "info,[".into();
let problems = problems(config.validate().map(|_| Config::default()));
for needle in [
"not a mysql://",
"max_connections",
"min_connections",
"body_limit_bytes",
"request_timeout_secs",
"hook_timeout_ms",
"`*` must be",
"example.com",
"log.level",
] {
assert!(problems.iter().any(|p| p.contains(needle)), "missing {needle}: {problems:?}");
}
}
#[test]
fn disabled_backend_is_named() {
for dialect in Dialect::ALL.iter().copied().filter(|d| !d.is_enabled()) {
let url = match dialect {
Dialect::MySql => "mysql://u:p@h/db",
Dialect::Postgres => "postgres://u:p@h/db",
Dialect::Sqlite => "sqlite::memory:",
};
let problems = problems(Config::from_sources(None, [("NBS__DATABASE__URL", url)]));
assert!(problems.iter().any(|p| p.contains(&format!("`{}` feature", dialect.name()))), "{problems:?}");
}
}
#[cfg(not(any(feature = "mysql", feature = "postgres", feature = "sqlite")))]
#[test]
fn no_backend_is_a_clear_error() {
let problems = problems(Config::from_sources(None, [("NBS__DATABASE__URL", "sqlite::memory:")]));
assert!(problems.iter().any(|p| p.contains("no database backend is compiled in")), "{problems:?}");
}
#[test]
fn module_secrets_never_show() {
#[derive(Deserialize, Debug)]
#[allow(dead_code)]
struct Mail {
port: u32,
}
let mut config = Config::default();
let mut mail = toml::Table::new();
mail.insert("smtp_password".into(), toml::Value::String("hunter2-secret".into()));
mail.insert("port".into(), toml::Value::String("hunter3-secret".into()));
config.modules.insert("mail".into(), toml::Value::Table(mail));
let debug = format!("{config:?}");
assert!(!debug.contains("hunter") && debug.contains("smtp_password"), "{debug}");
let error = config.module_config::<Mail>("mail").err().map(|e| e.to_string()).unwrap_or_default();
assert!(error.contains("modules.mail") && !error.contains("hunter"), "{error}");
assert_eq!(mask_quoted(r#"invalid type: string "x", expected u32 'y"#), r#"invalid type: string "<hidden>", expected u32 '<hidden>"#);
}
#[test]
fn secret_or_file() {
let dir = crate::test_support::temp_dir("config-secret-file");
let path = dir.join("pw");
std::fs::write(&path, "s3cret\n").unwrap();
assert_eq!(resolve_secret("x", None, Some(&path)).unwrap().map(|s| s.expose().to_string()).as_deref(), Some("s3cret"));
assert!(resolve_secret("x", Some(SecretString::new("a")), Some(&path)).is_err());
assert!(resolve_secret("x", None, None).unwrap().is_none());
}
#[test]
fn ui_needs_a_pinned_script() {
let mut config = Config::default();
config.openapi.ui = true;
let found = problems(config.validate().map(|_| Config::default()));
assert!(found.iter().any(|p| p.contains("ui_script_integrity")), "{found:?}");
config.openapi.ui_script_url = Some("https://cdn.example/viewer@1.2.3".into());
config.openapi.ui_script_integrity = Some("sha384-abc+/=".into());
let problems = problems(config.validate().map(|_| Config::default()));
assert!(!problems.iter().any(|p| p.contains("ui_script")), "{problems:?}");
}
#[test]
fn secrets_are_redacted() {
let secret = "mysql://game:hunter2-very-secret@db/game";
let config = Config::from_sources(None, [("NBS__DATABASE__URL", secret)]);
let text = match &config {
Ok(config) => format!("{config:?}"),
Err(error) => format!("{error:?} {error}"),
};
assert!(!text.contains("hunter2"), "{text}");
let mut config = Config::default();
config.database.url = SecretString::new(secret);
assert!(!format!("{config:?}").contains("hunter2"));
assert!(format!("{config:?}").contains("<redacted>"));
let problems = problems_of_toml("[database]\nurl = \"mysql://u:hunter2@h/db\" garbage\n");
assert!(!problems.join(" ").contains("hunter2"), "{problems:?}");
}
#[cfg(any(feature = "mysql", feature = "postgres", feature = "sqlite"))]
#[test]
fn url_file_is_read_and_trimmed() {
let dir = crate::test_support::temp_dir("config-url-file");
let path = dir.join("db_url");
std::fs::write(&path, format!("{}\n", url_for_enabled())).unwrap();
let env = [("NBS__DATABASE__URL_FILE", path.to_string_lossy().to_string())];
let config = Config::from_sources(None, env).unwrap();
assert_eq!(config.database.url.expose(), url_for_enabled());
let both = [("NBS__DATABASE__URL_FILE", path.to_string_lossy().to_string()), ("NBS__DATABASE__URL", url_for_enabled().to_string())];
assert!(problems(Config::from_sources(None, both))[0].contains("both set"));
}
#[cfg(any(feature = "mysql", feature = "postgres", feature = "sqlite"))]
#[test]
fn module_sections() {
#[derive(Deserialize, Debug, PartialEq)]
struct Chat {
max_chars: u32,
name: String,
nested: Option<BTreeMap<String, bool>>,
}
let toml = format!("[database]\nurl = \"{}\"\n[modules.chat]\nmax_chars = 10\nname = \"x\"\n", url_for_enabled());
let env = [("NBS__MODULES__CHAT__MAX_CHARS", "500"), ("NBS__MODULES__CHAT__NESTED__ON", "true")];
let config = Config::from_sources(Some(&toml), env).unwrap();
let chat: Option<Chat> = config.module_config("chat").unwrap();
let chat = chat.unwrap();
assert_eq!(chat.max_chars, 500);
assert_eq!(chat.name, "x");
assert_eq!(chat.nested.and_then(|n| n.get("on").copied()), Some(true));
assert!(config.module_config::<Chat>("storage").unwrap().is_none());
assert!(config.module_config::<u32>("chat").is_err());
assert_eq!(config.unknown_module_sections(&["chat"]), Vec::<String>::new());
assert_eq!(config.unknown_module_sections(&["chta"]), ["chat"]);
assert_eq!(env_value("mysql://x"), toml::Value::String("mysql://x".into()));
}
}