use std::collections::HashMap;
use std::fmt;
use std::path::PathBuf;
use serde::Deserialize;
use crate::error::{Result, WaypointError};
macro_rules! apply_option {
($opt:expr => $target:expr) => {
if let Some(v) = $opt {
$target = v;
}
};
}
macro_rules! apply_option_some {
($opt:expr => $target:expr) => {
if let Some(v) = $opt {
$target = Some(v);
}
};
}
macro_rules! apply_option_clone {
($opt:expr => $target:expr) => {
if let Some(ref v) = $opt {
$target = v.clone();
}
};
}
macro_rules! apply_option_some_clone {
($opt:expr => $target:expr) => {
if let Some(ref v) = $opt {
$target = Some(v.clone());
}
};
}
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
pub enum SslMode {
Disable,
#[default]
Prefer,
Require,
VerifyCa,
VerifyFull,
}
impl SslMode {
pub fn verifies_certificate(&self) -> bool {
matches!(self, SslMode::VerifyCa | SslMode::VerifyFull)
}
pub fn requires_tls(&self) -> bool {
matches!(
self,
SslMode::Require | SslMode::VerifyCa | SslMode::VerifyFull
)
}
pub fn as_str(&self) -> &'static str {
match self {
SslMode::Disable => "disable",
SslMode::Prefer => "prefer",
SslMode::Require => "require",
SslMode::VerifyCa => "verify-ca",
SslMode::VerifyFull => "verify-full",
}
}
}
impl fmt::Display for SslMode {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str(self.as_str())
}
}
impl std::str::FromStr for SslMode {
type Err = WaypointError;
fn from_str(s: &str) -> std::result::Result<Self, Self::Err> {
let normalized = s.to_lowercase().replace(['-', '_'], "");
match normalized.as_str() {
"disable" | "disabled" => Ok(SslMode::Disable),
"prefer" => Ok(SslMode::Prefer),
"require" | "required" => Ok(SslMode::Require),
"verifyca" => Ok(SslMode::VerifyCa),
"verifyfull" => Ok(SslMode::VerifyFull),
"allow" => Err(WaypointError::ConfigError(
"SSL mode 'allow' is not supported. Use 'prefer' to try TLS first \
and fall back to plaintext."
.to_string(),
)),
_ => Err(WaypointError::ConfigError(format!(
"Invalid SSL mode '{}'. Use 'disable', 'prefer', 'require', \
'verify-ca', or 'verify-full'.",
s
))),
}
}
}
fn apply_ssl_mode(target: &mut SslMode, value: &str, source: &str) {
match value.parse() {
Ok(mode) => *target = mode,
Err(e) => log::warn!("{} (from {}); keeping '{}'.", e, source, target),
}
}
fn apply_env_number<T>(target: &mut T, value: &str, var: &str)
where
T: std::str::FromStr + std::fmt::Display,
{
match value.parse::<T>() {
Ok(n) => *target = n,
Err(_) => log::warn!(
"Invalid {} '{}': expected a whole number; keeping '{}'.",
var,
value,
target
),
}
}
fn parse_env_bool(value: &str, var: &str) -> Option<bool> {
match value.trim().to_ascii_lowercase().as_str() {
"1" | "true" | "yes" | "on" => Some(true),
"0" | "false" | "no" | "off" => Some(false),
other => {
log::warn!(
"Invalid {} '{}': expected one of 1/true/yes/on or 0/false/no/off; \
leaving the setting unchanged.",
var,
other
);
None
}
}
}
#[derive(Debug, Clone, Default)]
pub struct WaypointConfig {
pub database: DatabaseConfig,
pub migrations: MigrationSettings,
pub hooks: HooksConfig,
pub placeholders: HashMap<String, String>,
pub lint: LintConfig,
pub snapshots: crate::commands::snapshot::SnapshotConfig,
pub preflight: crate::preflight::PreflightConfig,
pub multi_database: Option<Vec<crate::multi::NamedDatabaseConfig>>,
pub guards: crate::guard::GuardsConfig,
pub reversals: crate::reversal::ReversalConfig,
pub safety: crate::safety::SafetyConfig,
pub advisor: crate::advisor::AdvisorConfig,
pub simulation: SimulationConfig,
}
#[derive(Clone)]
pub struct DatabaseConfig {
pub url: Option<String>,
pub host: Option<String>,
pub port: Option<u16>,
pub user: Option<String>,
pub password: Option<String>,
pub database: Option<String>,
pub connect_retries: u32,
pub ssl_mode: SslMode,
pub ssl_root_cert: Option<PathBuf>,
pub connect_timeout_secs: u32,
pub statement_timeout_secs: u32,
pub keepalive_secs: u32,
pub engine: crate::dialect::DialectKind,
}
impl Default for DatabaseConfig {
fn default() -> Self {
Self {
url: None,
host: None,
port: None,
user: None,
password: None,
database: None,
connect_retries: 0,
ssl_mode: SslMode::Prefer,
ssl_root_cert: None,
connect_timeout_secs: 30,
statement_timeout_secs: 0,
keepalive_secs: 120,
engine: crate::dialect::DialectKind::Postgres,
}
}
}
impl fmt::Debug for DatabaseConfig {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("DatabaseConfig")
.field("url", &self.url.as_ref().map(|_| "[REDACTED]"))
.field("host", &self.host)
.field("port", &self.port)
.field("user", &self.user)
.field("password", &self.password.as_ref().map(|_| "[REDACTED]"))
.field("database", &self.database)
.field("connect_retries", &self.connect_retries)
.field("ssl_mode", &self.ssl_mode)
.field("ssl_root_cert", &self.ssl_root_cert)
.field("connect_timeout_secs", &self.connect_timeout_secs)
.field("statement_timeout_secs", &self.statement_timeout_secs)
.field("keepalive_secs", &self.keepalive_secs)
.field("engine", &self.engine)
.finish()
}
}
#[derive(Debug, Clone, Default)]
pub struct HooksConfig {
pub before_migrate: Vec<PathBuf>,
pub after_migrate: Vec<PathBuf>,
pub before_each_migrate: Vec<PathBuf>,
pub after_each_migrate: Vec<PathBuf>,
}
#[derive(Debug, Clone, Default)]
pub struct LintConfig {
pub disabled_rules: Vec<String>,
}
#[derive(Debug, Clone)]
pub struct MigrationSettings {
pub locations: Vec<PathBuf>,
pub table: String,
pub schema: String,
pub out_of_order: bool,
pub validate_on_migrate: bool,
pub clean_enabled: bool,
pub baseline_version: String,
pub installed_by: Option<String>,
pub environment: Option<String>,
pub dependency_ordering: bool,
pub show_progress: bool,
pub batch_transaction: bool,
}
impl Default for MigrationSettings {
fn default() -> Self {
Self {
locations: vec![PathBuf::from("db/migrations")],
table: "waypoint_schema_history".to_string(),
schema: "public".to_string(),
out_of_order: false,
validate_on_migrate: true,
clean_enabled: false,
baseline_version: "1".to_string(),
installed_by: None,
environment: None,
dependency_ordering: false,
show_progress: true,
batch_transaction: false,
}
}
}
#[derive(Debug, Clone, Default)]
pub struct SimulationConfig {
pub simulate_before_migrate: bool,
}
#[derive(Deserialize, Default)]
struct TomlConfig {
database: Option<TomlDatabaseConfig>,
migrations: Option<TomlMigrationSettings>,
hooks: Option<TomlHooksConfig>,
placeholders: Option<HashMap<String, String>>,
lint: Option<TomlLintConfig>,
snapshots: Option<TomlSnapshotConfig>,
preflight: Option<TomlPreflightConfig>,
databases: Option<Vec<TomlNamedDatabaseConfig>>,
guards: Option<TomlGuardsConfig>,
reversals: Option<TomlReversalConfig>,
safety: Option<TomlSafetyConfig>,
advisor: Option<TomlAdvisorConfig>,
simulation: Option<TomlSimulationConfig>,
}
#[derive(Deserialize, Default)]
struct TomlDatabaseConfig {
url: Option<String>,
host: Option<String>,
port: Option<u16>,
user: Option<String>,
password: Option<String>,
database: Option<String>,
connect_retries: Option<u32>,
ssl_mode: Option<String>,
ssl_root_cert: Option<String>,
connect_timeout: Option<u32>,
statement_timeout: Option<u32>,
keepalive: Option<u32>,
engine: Option<String>,
}
#[derive(Deserialize, Default)]
struct TomlMigrationSettings {
locations: Option<Vec<String>>,
table: Option<String>,
schema: Option<String>,
out_of_order: Option<bool>,
validate_on_migrate: Option<bool>,
clean_enabled: Option<bool>,
baseline_version: Option<String>,
installed_by: Option<String>,
environment: Option<String>,
dependency_ordering: Option<bool>,
show_progress: Option<bool>,
batch_transaction: Option<bool>,
}
#[derive(Deserialize, Default)]
struct TomlLintConfig {
disabled_rules: Option<Vec<String>>,
}
#[derive(Deserialize, Default)]
struct TomlSnapshotConfig {
directory: Option<String>,
auto_snapshot_on_migrate: Option<bool>,
max_snapshots: Option<usize>,
strip_definer_mysql: Option<bool>,
}
#[derive(Deserialize, Default)]
struct TomlPreflightConfig {
enabled: Option<bool>,
max_replication_lag_mb: Option<i64>,
max_replication_lag_secs: Option<i64>,
long_query_threshold_secs: Option<i64>,
}
#[derive(Deserialize, Default)]
struct TomlNamedDatabaseConfig {
name: Option<String>,
url: Option<String>,
depends_on: Option<Vec<String>>,
migrations: Option<TomlMigrationSettings>,
hooks: Option<TomlHooksConfig>,
placeholders: Option<HashMap<String, String>>,
}
#[derive(Deserialize, Default)]
struct TomlHooksConfig {
before_migrate: Option<Vec<String>>,
after_migrate: Option<Vec<String>>,
before_each_migrate: Option<Vec<String>>,
after_each_migrate: Option<Vec<String>>,
}
#[derive(Deserialize, Default)]
struct TomlGuardsConfig {
enabled: Option<bool>,
on_require_fail: Option<String>,
}
#[derive(Deserialize, Default)]
struct TomlReversalConfig {
enabled: Option<bool>,
warn_data_loss: Option<bool>,
}
#[derive(Deserialize, Default)]
struct TomlSafetyConfig {
enabled: Option<bool>,
block_on_danger: Option<bool>,
large_table_threshold: Option<i64>,
huge_table_threshold: Option<i64>,
refresh_stats_mysql: Option<bool>,
}
#[derive(Deserialize, Default)]
struct TomlAdvisorConfig {
run_after_migrate: Option<bool>,
disabled_rules: Option<Vec<String>>,
}
#[derive(Deserialize, Default)]
struct TomlSimulationConfig {
simulate_before_migrate: Option<bool>,
}
#[derive(Debug, Default, Clone)]
pub struct CliOverrides {
pub url: Option<String>,
pub schema: Option<String>,
pub table: Option<String>,
pub locations: Option<Vec<PathBuf>>,
pub out_of_order: Option<bool>,
pub validate_on_migrate: Option<bool>,
pub baseline_version: Option<String>,
pub connect_retries: Option<u32>,
pub ssl_mode: Option<String>,
pub ssl_root_cert: Option<PathBuf>,
pub connect_timeout: Option<u32>,
pub statement_timeout: Option<u32>,
pub environment: Option<String>,
pub dependency_ordering: Option<bool>,
pub keepalive: Option<u32>,
pub batch_transaction: Option<bool>,
}
impl WaypointConfig {
pub fn load(config_path: Option<&str>, overrides: &CliOverrides) -> Result<Self> {
let mut config = WaypointConfig::default();
let toml_path = config_path.unwrap_or("waypoint.toml");
if let Ok(content) = std::fs::read_to_string(toml_path) {
#[cfg(unix)]
{
use std::os::unix::fs::PermissionsExt;
if let Ok(meta) = std::fs::metadata(toml_path) {
let mode = meta.permissions().mode();
if mode & 0o077 != 0 {
log::warn!(
"Config file has overly permissive permissions. Consider chmod 600.; path={}, mode={:o}",
toml_path,
mode
);
}
}
}
let toml_config: TomlConfig = toml::from_str(&content).map_err(|e| {
WaypointError::ConfigError(format!(
"Failed to parse config file '{}': {}",
toml_path, e
))
})?;
config.apply_toml(toml_config);
} else if config_path.is_some() {
return Err(WaypointError::ConfigError(format!(
"Config file '{}' not found",
toml_path
)));
}
config.apply_env();
config.apply_cli(overrides);
crate::db::validate_identifier(&config.migrations.schema)?;
crate::db::validate_identifier(&config.migrations.table)?;
if config.database.connect_retries > 20 {
config.database.connect_retries = 20;
log::warn!("connect_retries capped at 20");
}
Ok(config)
}
fn apply_toml(&mut self, toml: TomlConfig) {
if let Some(db) = toml.database {
apply_option_some!(db.url => self.database.url);
apply_option_some!(db.host => self.database.host);
apply_option_some!(db.port => self.database.port);
apply_option_some!(db.user => self.database.user);
apply_option_some!(db.password => self.database.password);
apply_option_some!(db.database => self.database.database);
apply_option!(db.connect_retries => self.database.connect_retries);
if let Some(v) = db.ssl_mode {
apply_ssl_mode(&mut self.database.ssl_mode, &v, "waypoint.toml");
}
if let Some(v) = db.ssl_root_cert {
self.database.ssl_root_cert = Some(PathBuf::from(v));
}
apply_option!(db.connect_timeout => self.database.connect_timeout_secs);
apply_option!(db.statement_timeout => self.database.statement_timeout_secs);
apply_option!(db.keepalive => self.database.keepalive_secs);
if let Some(v) = db.engine {
match v.parse() {
Ok(kind) => self.database.engine = kind,
Err(_) => log::warn!(
"Invalid engine '{}' in config, using default 'postgres'. Valid values: postgres, mysql",
v
),
}
}
}
if let Some(m) = toml.migrations {
if let Some(v) = m.locations {
self.migrations.locations = v.into_iter().map(|s| normalize_location(&s)).collect();
}
apply_option!(m.table => self.migrations.table);
apply_option!(m.schema => self.migrations.schema);
apply_option!(m.out_of_order => self.migrations.out_of_order);
apply_option!(m.validate_on_migrate => self.migrations.validate_on_migrate);
apply_option!(m.clean_enabled => self.migrations.clean_enabled);
apply_option!(m.baseline_version => self.migrations.baseline_version);
apply_option_some!(m.installed_by => self.migrations.installed_by);
apply_option_some!(m.environment => self.migrations.environment);
apply_option!(m.dependency_ordering => self.migrations.dependency_ordering);
apply_option!(m.show_progress => self.migrations.show_progress);
apply_option!(m.batch_transaction => self.migrations.batch_transaction);
}
if let Some(h) = toml.hooks {
if let Some(v) = h.before_migrate {
self.hooks.before_migrate = v.into_iter().map(PathBuf::from).collect();
}
if let Some(v) = h.after_migrate {
self.hooks.after_migrate = v.into_iter().map(PathBuf::from).collect();
}
if let Some(v) = h.before_each_migrate {
self.hooks.before_each_migrate = v.into_iter().map(PathBuf::from).collect();
}
if let Some(v) = h.after_each_migrate {
self.hooks.after_each_migrate = v.into_iter().map(PathBuf::from).collect();
}
}
if let Some(p) = toml.placeholders {
self.placeholders.extend(p);
}
if let Some(l) = toml.lint {
apply_option!(l.disabled_rules => self.lint.disabled_rules);
}
if let Some(s) = toml.snapshots {
if let Some(v) = s.directory {
self.snapshots.directory = PathBuf::from(v);
}
apply_option!(s.auto_snapshot_on_migrate => self.snapshots.auto_snapshot_on_migrate);
apply_option!(s.max_snapshots => self.snapshots.max_snapshots);
apply_option!(s.strip_definer_mysql => self.snapshots.strip_definer_mysql);
}
if let Some(p) = toml.preflight {
apply_option!(p.enabled => self.preflight.enabled);
apply_option!(p.max_replication_lag_mb => self.preflight.max_replication_lag_mb);
apply_option!(p.max_replication_lag_secs => self.preflight.max_replication_lag_secs);
apply_option!(p.long_query_threshold_secs => self.preflight.long_query_threshold_secs);
}
if let Some(g) = toml.guards {
apply_option!(g.enabled => self.guards.enabled);
if let Some(v) = g.on_require_fail {
match v.parse() {
Ok(policy) => self.guards.on_require_fail = policy,
Err(_) => log::warn!(
"Invalid on_require_fail '{}' in config, using default 'error'. Valid values: error, warn, skip",
v
),
}
}
}
if let Some(r) = toml.reversals {
apply_option!(r.enabled => self.reversals.enabled);
apply_option!(r.warn_data_loss => self.reversals.warn_data_loss);
}
if let Some(s) = toml.safety {
apply_option!(s.enabled => self.safety.enabled);
apply_option!(s.block_on_danger => self.safety.block_on_danger);
apply_option!(s.large_table_threshold => self.safety.large_table_threshold);
apply_option!(s.huge_table_threshold => self.safety.huge_table_threshold);
apply_option!(s.refresh_stats_mysql => self.safety.refresh_stats_mysql);
}
if let Some(a) = toml.advisor {
apply_option!(a.run_after_migrate => self.advisor.run_after_migrate);
apply_option!(a.disabled_rules => self.advisor.disabled_rules);
}
if let Some(s) = toml.simulation {
apply_option!(s.simulate_before_migrate => self.simulation.simulate_before_migrate);
}
if let Some(databases) = toml.databases {
let mut named_dbs = Vec::new();
for db in databases {
let name = db.name.unwrap_or_default();
let mut db_config = DatabaseConfig {
url: None,
host: None,
port: None,
user: None,
password: None,
database: None,
..self.database.clone()
};
apply_option_some!(db.url => db_config.url);
let env_url_key = format!("WAYPOINT_DB_{}_URL", name.to_uppercase());
if let Ok(url) = std::env::var(&env_url_key) {
db_config.url = Some(url);
}
let mut mig_settings = self.migrations.clone();
if let Some(m) = db.migrations {
if let Some(v) = m.locations {
mig_settings.locations =
v.into_iter().map(|s| normalize_location(&s)).collect();
}
apply_option!(m.table => mig_settings.table);
apply_option!(m.schema => mig_settings.schema);
apply_option!(m.out_of_order => mig_settings.out_of_order);
apply_option!(m.validate_on_migrate => mig_settings.validate_on_migrate);
apply_option!(m.clean_enabled => mig_settings.clean_enabled);
apply_option!(m.baseline_version => mig_settings.baseline_version);
apply_option_some!(m.installed_by => mig_settings.installed_by);
apply_option_some!(m.environment => mig_settings.environment);
apply_option!(m.dependency_ordering => mig_settings.dependency_ordering);
apply_option!(m.show_progress => mig_settings.show_progress);
apply_option!(m.batch_transaction => mig_settings.batch_transaction);
}
let mut hooks_config = HooksConfig::default();
if let Some(h) = db.hooks {
if let Some(v) = h.before_migrate {
hooks_config.before_migrate = v.into_iter().map(PathBuf::from).collect();
}
if let Some(v) = h.after_migrate {
hooks_config.after_migrate = v.into_iter().map(PathBuf::from).collect();
}
if let Some(v) = h.before_each_migrate {
hooks_config.before_each_migrate =
v.into_iter().map(PathBuf::from).collect();
}
if let Some(v) = h.after_each_migrate {
hooks_config.after_each_migrate =
v.into_iter().map(PathBuf::from).collect();
}
}
named_dbs.push(crate::multi::NamedDatabaseConfig {
name,
database: db_config,
migrations: mig_settings,
hooks: hooks_config,
placeholders: db.placeholders.unwrap_or_default(),
depends_on: db.depends_on.unwrap_or_default(),
});
}
self.multi_database = Some(named_dbs);
}
}
fn apply_env(&mut self) {
if let Ok(v) = std::env::var("WAYPOINT_DATABASE_URL") {
self.database.url = Some(v);
}
if let Ok(v) = std::env::var("WAYPOINT_DATABASE_HOST") {
self.database.host = Some(v);
}
if let Ok(v) = std::env::var("WAYPOINT_DATABASE_PORT") {
match v.parse::<u16>() {
Ok(port) => self.database.port = Some(port),
Err(_) => log::warn!(
"Invalid WAYPOINT_DATABASE_PORT '{}': expected a port number; \
keeping the existing setting.",
v
),
}
}
if let Ok(v) = std::env::var("WAYPOINT_DATABASE_USER") {
self.database.user = Some(v);
}
if let Ok(v) = std::env::var("WAYPOINT_DATABASE_PASSWORD") {
self.database.password = Some(v);
}
if let Ok(v) = std::env::var("WAYPOINT_DATABASE_NAME") {
self.database.database = Some(v);
}
if let Ok(v) = std::env::var("WAYPOINT_CONNECT_RETRIES") {
apply_env_number(
&mut self.database.connect_retries,
&v,
"WAYPOINT_CONNECT_RETRIES",
);
}
if let Ok(v) = std::env::var("WAYPOINT_SSL_MODE") {
apply_ssl_mode(&mut self.database.ssl_mode, &v, "WAYPOINT_SSL_MODE");
}
if let Ok(v) = std::env::var("WAYPOINT_SSL_ROOT_CERT") {
self.database.ssl_root_cert = Some(PathBuf::from(v));
}
if let Ok(v) = std::env::var("WAYPOINT_CONNECT_TIMEOUT") {
apply_env_number(
&mut self.database.connect_timeout_secs,
&v,
"WAYPOINT_CONNECT_TIMEOUT",
);
}
if let Ok(v) = std::env::var("WAYPOINT_STATEMENT_TIMEOUT") {
apply_env_number(
&mut self.database.statement_timeout_secs,
&v,
"WAYPOINT_STATEMENT_TIMEOUT",
);
}
if let Ok(v) = std::env::var("WAYPOINT_MIGRATIONS_LOCATIONS") {
self.migrations.locations =
v.split(',').map(|s| normalize_location(s.trim())).collect();
}
if let Ok(v) = std::env::var("WAYPOINT_MIGRATIONS_TABLE") {
self.migrations.table = v;
}
if let Ok(v) = std::env::var("WAYPOINT_MIGRATIONS_SCHEMA") {
self.migrations.schema = v;
}
if let Ok(v) = std::env::var("WAYPOINT_DATABASE_ENGINE") {
match v.parse() {
Ok(kind) => self.database.engine = kind,
Err(_) => log::warn!(
"Invalid WAYPOINT_DATABASE_ENGINE '{}', using default 'postgres'",
v
),
}
}
if let Ok(v) = std::env::var("WAYPOINT_KEEPALIVE") {
apply_env_number(&mut self.database.keepalive_secs, &v, "WAYPOINT_KEEPALIVE");
}
if let Ok(v) = std::env::var("WAYPOINT_BATCH_TRANSACTION")
&& let Some(b) = parse_env_bool(&v, "WAYPOINT_BATCH_TRANSACTION")
{
self.migrations.batch_transaction = b;
}
if let Ok(v) = std::env::var("WAYPOINT_ENVIRONMENT") {
self.migrations.environment = Some(v);
}
for (key, value) in std::env::vars() {
if let Some(placeholder_key) = key.strip_prefix("WAYPOINT_PLACEHOLDER_") {
self.placeholders
.insert(placeholder_key.to_lowercase(), value);
}
}
}
fn apply_cli(&mut self, overrides: &CliOverrides) {
apply_option_some_clone!(overrides.url => self.database.url);
apply_option_clone!(overrides.schema => self.migrations.schema);
apply_option_clone!(overrides.table => self.migrations.table);
apply_option_clone!(overrides.locations => self.migrations.locations);
apply_option!(overrides.out_of_order => self.migrations.out_of_order);
apply_option!(overrides.validate_on_migrate => self.migrations.validate_on_migrate);
apply_option_clone!(overrides.baseline_version => self.migrations.baseline_version);
apply_option!(overrides.connect_retries => self.database.connect_retries);
if let Some(ref v) = overrides.ssl_mode {
apply_ssl_mode(&mut self.database.ssl_mode, v, "--ssl-mode");
}
apply_option_some_clone!(overrides.ssl_root_cert => self.database.ssl_root_cert);
apply_option!(overrides.connect_timeout => self.database.connect_timeout_secs);
apply_option!(overrides.statement_timeout => self.database.statement_timeout_secs);
apply_option_some_clone!(overrides.environment => self.migrations.environment);
apply_option!(overrides.dependency_ordering => self.migrations.dependency_ordering);
apply_option!(overrides.keepalive => self.database.keepalive_secs);
apply_option!(overrides.batch_transaction => self.migrations.batch_transaction);
}
pub fn connection_string(&self) -> Result<String> {
if let Some(ref url) = self.database.url {
return Ok(normalize_jdbc_url(url));
}
let engine = self.database.engine;
let host = self.database.host.as_deref().unwrap_or("localhost");
let default_port = match engine {
crate::dialect::DialectKind::Postgres => 5432,
crate::dialect::DialectKind::Mysql => 3306,
};
let port = self.database.port.unwrap_or(default_port);
let user =
self.database.user.as_deref().ok_or_else(|| {
WaypointError::ConfigError("Database user is required".to_string())
})?;
let database =
self.database.database.as_deref().ok_or_else(|| {
WaypointError::ConfigError("Database name is required".to_string())
})?;
match engine {
crate::dialect::DialectKind::Mysql => {
let auth = match self.database.password {
Some(ref password) => format!(
"{}:{}",
percent_encode_userinfo(user),
percent_encode_userinfo(password)
),
None => percent_encode_userinfo(user),
};
Ok(format!("mysql://{}@{}:{}/{}", auth, host, port, database))
}
crate::dialect::DialectKind::Postgres => {
let mut url = format!(
"host={} port={} user={} dbname={}",
host, port, user, database
);
if let Some(ref password) = self.database.password {
let escaped = password.replace('\\', "\\\\").replace('\'', "\\'");
url.push_str(&format!(" password='{}'", escaped));
}
Ok(url)
}
}
}
}
fn percent_encode_userinfo(s: &str) -> String {
let mut out = String::with_capacity(s.len());
for b in s.bytes() {
match b {
b'A'..=b'Z' | b'a'..=b'z' | b'0'..=b'9' | b'-' | b'.' | b'_' | b'~' => {
out.push(b as char)
}
_ => out.push_str(&format!("%{:02X}", b)),
}
}
out
}
fn normalize_jdbc_url(url: &str) -> String {
let url = url.strip_prefix("jdbc:").unwrap_or(url);
if let Some((base, query)) = url.split_once('?') {
let mut user = None;
let mut password = None;
let mut other_params = Vec::new();
for param in query.split('&') {
if let Some((key, value)) = param.split_once('=') {
match key.to_lowercase().as_str() {
"user" => user = Some(value.to_string()),
"password" => password = Some(value.to_string()),
_ => other_params.push(param.to_string()),
}
}
}
if (user.is_some() || password.is_some())
&& let Some(rest) = base
.strip_prefix("postgresql://")
.or_else(|| base.strip_prefix("postgres://"))
{
let scheme = if base.starts_with("postgresql://") {
"postgresql"
} else {
"postgres"
};
let auth = match (user, password) {
(Some(u), Some(p)) => format!(
"{}:{}@",
percent_encode_userinfo(&u),
percent_encode_userinfo(&p)
),
(Some(u), None) => format!("{}@", percent_encode_userinfo(&u)),
(None, Some(p)) => format!(":{}@", percent_encode_userinfo(&p)),
(None, None) => String::new(),
};
let mut result = format!("{}://{}{}", scheme, auth, rest);
if !other_params.is_empty() {
result.push('?');
result.push_str(&other_params.join("&"));
}
return result;
}
if other_params.is_empty() {
return base.to_string();
}
return format!("{}?{}", base, other_params.join("&"));
}
url.to_string()
}
pub fn normalize_location(location: &str) -> PathBuf {
let stripped = location.strip_prefix("filesystem:").unwrap_or(location);
PathBuf::from(stripped)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_default_config() {
let config = WaypointConfig::default();
assert_eq!(config.migrations.table, "waypoint_schema_history");
assert_eq!(config.migrations.schema, "public");
assert!(!config.migrations.out_of_order);
assert!(config.migrations.validate_on_migrate);
assert!(!config.migrations.clean_enabled);
assert_eq!(config.migrations.baseline_version, "1");
assert_eq!(
config.migrations.locations,
vec![PathBuf::from("db/migrations")]
);
}
#[test]
fn test_connection_string_from_url() {
let mut config = WaypointConfig::default();
config.database.url = Some("postgres://user:pass@localhost/db".to_string());
assert_eq!(
config.connection_string().unwrap(),
"postgres://user:pass@localhost/db"
);
}
#[test]
fn test_connection_string_from_fields() {
let mut config = WaypointConfig::default();
config.database.host = Some("myhost".to_string());
config.database.port = Some(5433);
config.database.user = Some("myuser".to_string());
config.database.database = Some("mydb".to_string());
config.database.password = Some("secret".to_string());
let conn = config.connection_string().unwrap();
assert!(conn.contains("host=myhost"));
assert!(conn.contains("port=5433"));
assert!(conn.contains("user=myuser"));
assert!(conn.contains("dbname=mydb"));
assert!(conn.contains("password='secret'"));
}
#[test]
fn test_connection_string_missing_user() {
let mut config = WaypointConfig::default();
config.database.database = Some("mydb".to_string());
assert!(config.connection_string().is_err());
}
#[test]
fn test_cli_overrides() {
let mut config = WaypointConfig::default();
let overrides = CliOverrides {
url: Some("postgres://override@localhost/db".to_string()),
schema: Some("custom_schema".to_string()),
table: Some("custom_table".to_string()),
locations: Some(vec![PathBuf::from("custom/path")]),
out_of_order: Some(true),
validate_on_migrate: Some(false),
baseline_version: Some("5".to_string()),
connect_retries: None,
ssl_mode: None,
ssl_root_cert: None,
connect_timeout: None,
statement_timeout: None,
environment: None,
dependency_ordering: None,
keepalive: None,
batch_transaction: None,
};
config.apply_cli(&overrides);
assert_eq!(
config.database.url.as_deref(),
Some("postgres://override@localhost/db")
);
assert_eq!(config.migrations.schema, "custom_schema");
assert_eq!(config.migrations.table, "custom_table");
assert_eq!(
config.migrations.locations,
vec![PathBuf::from("custom/path")]
);
assert!(config.migrations.out_of_order);
assert!(!config.migrations.validate_on_migrate);
assert_eq!(config.migrations.baseline_version, "5");
}
#[test]
fn test_toml_parsing() {
let toml_str = r#"
[database]
url = "postgres://user:pass@localhost/mydb"
[migrations]
table = "my_history"
schema = "app"
out_of_order = true
locations = ["sql/migrations", "sql/seeds"]
[placeholders]
env = "production"
app_name = "myapp"
"#;
let toml_config: TomlConfig = toml::from_str(toml_str).unwrap();
let mut config = WaypointConfig::default();
config.apply_toml(toml_config);
assert_eq!(
config.database.url.as_deref(),
Some("postgres://user:pass@localhost/mydb")
);
assert_eq!(config.migrations.table, "my_history");
assert_eq!(config.migrations.schema, "app");
assert!(config.migrations.out_of_order);
assert_eq!(
config.migrations.locations,
vec![PathBuf::from("sql/migrations"), PathBuf::from("sql/seeds")]
);
assert_eq!(config.placeholders.get("env").unwrap(), "production");
assert_eq!(config.placeholders.get("app_name").unwrap(), "myapp");
}
#[test]
fn test_normalize_jdbc_url_with_credentials() {
let url = "jdbc:postgresql://myhost:5432/mydb?user=admin&password=secret";
assert_eq!(
normalize_jdbc_url(url),
"postgresql://admin:secret@myhost:5432/mydb"
);
}
#[test]
fn test_normalize_jdbc_url_user_only() {
let url = "jdbc:postgresql://myhost:5432/mydb?user=admin";
assert_eq!(
normalize_jdbc_url(url),
"postgresql://admin@myhost:5432/mydb"
);
}
#[test]
fn test_normalize_jdbc_url_strips_jdbc_prefix() {
let url = "jdbc:postgresql://myhost:5432/mydb";
assert_eq!(normalize_jdbc_url(url), "postgresql://myhost:5432/mydb");
}
#[test]
fn test_normalize_jdbc_url_passthrough() {
let url = "postgresql://user:pass@myhost:5432/mydb";
assert_eq!(normalize_jdbc_url(url), url);
}
#[test]
fn test_normalize_jdbc_url_preserves_other_params() {
let url = "jdbc:postgresql://myhost:5432/mydb?user=admin&password=secret&sslmode=require";
assert_eq!(
normalize_jdbc_url(url),
"postgresql://admin:secret@myhost:5432/mydb?sslmode=require"
);
}
#[test]
fn test_normalize_location_filesystem_prefix() {
assert_eq!(
normalize_location("filesystem:/flyway/sql"),
PathBuf::from("/flyway/sql")
);
}
#[test]
fn test_normalize_location_plain_path() {
assert_eq!(
normalize_location("/my/migrations"),
PathBuf::from("/my/migrations")
);
}
#[test]
fn test_normalize_location_relative() {
assert_eq!(
normalize_location("filesystem:db/migrations"),
PathBuf::from("db/migrations")
);
}
#[test]
fn test_connection_string_password_special_chars() {
let config = WaypointConfig {
database: DatabaseConfig {
host: Some("localhost".to_string()),
port: Some(5432),
user: Some("admin".to_string()),
database: Some("mydb".to_string()),
password: Some("p@ss'w ord".to_string()),
..Default::default()
},
..Default::default()
};
let conn = config.connection_string().unwrap();
assert!(conn.contains("password='p@ss\\'w ord'"));
}
#[test]
fn test_connection_string_mysql_from_fields() {
let config = WaypointConfig {
database: DatabaseConfig {
engine: crate::dialect::DialectKind::Mysql,
host: Some("db.internal".to_string()),
user: Some("app".to_string()),
password: Some("s3cr3t".to_string()),
database: Some("shop".to_string()),
..Default::default()
},
..Default::default()
};
assert_eq!(
config.connection_string().unwrap(),
"mysql://app:s3cr3t@db.internal:3306/shop"
);
assert_eq!(
crate::dialect::DialectKind::from_url(&config.connection_string().unwrap()),
Some(crate::dialect::DialectKind::Mysql)
);
}
#[test]
fn test_connection_string_mysql_percent_encodes_password() {
let config = WaypointConfig {
database: DatabaseConfig {
engine: crate::dialect::DialectKind::Mysql,
host: Some("h".to_string()),
port: Some(13306),
user: Some("u".to_string()),
password: Some("p@ss/word".to_string()),
database: Some("d".to_string()),
..Default::default()
},
..Default::default()
};
assert_eq!(
config.connection_string().unwrap(),
"mysql://u:p%40ss%2Fword@h:13306/d"
);
}
#[test]
fn test_normalize_jdbc_url_percent_encodes_credentials() {
let url = "jdbc:postgresql://myhost:5432/mydb?user=adm%69n&password=p@ss";
assert_eq!(
normalize_jdbc_url(url),
"postgresql://adm%2569n:p%40ss@myhost:5432/mydb"
);
}
#[test]
fn test_engine_defaults_to_postgres() {
let config = WaypointConfig::default();
assert_eq!(
config.database.engine,
crate::dialect::DialectKind::Postgres
);
}
#[test]
fn test_toml_engine_key() {
let toml_str = r#"
[database]
engine = "mysql"
host = "localhost"
user = "root"
database = "app"
"#;
let toml_config: TomlConfig = toml::from_str(toml_str).unwrap();
let mut config = WaypointConfig::default();
config.apply_toml(toml_config);
assert_eq!(config.database.engine, crate::dialect::DialectKind::Mysql);
assert!(config.connection_string().unwrap().starts_with("mysql://"));
}
#[test]
fn test_ssl_mode_parses_the_libpq_ladder() {
for (input, want) in [
("disable", SslMode::Disable),
("DISABLE", SslMode::Disable),
("disabled", SslMode::Disable),
("prefer", SslMode::Prefer),
("require", SslMode::Require),
("required", SslMode::Require),
("verify-ca", SslMode::VerifyCa),
("verify_ca", SslMode::VerifyCa),
("verifyca", SslMode::VerifyCa),
("Verify-Full", SslMode::VerifyFull),
("verify_full", SslMode::VerifyFull),
] {
assert_eq!(input.parse::<SslMode>().unwrap(), want, "input: {}", input);
}
}
#[test]
fn test_ssl_mode_rejects_allow_and_names_the_alternative() {
let err = "allow".parse::<SslMode>().unwrap_err().to_string();
assert!(err.contains("not supported"), "got: {}", err);
assert!(err.contains("prefer"), "must name the replacement: {}", err);
}
#[test]
fn test_ssl_mode_rejects_unknown_values() {
assert!("verify".parse::<SslMode>().is_err());
assert!("".parse::<SslMode>().is_err());
let err = "banana".parse::<SslMode>().unwrap_err().to_string();
assert!(
err.contains("verify-full"),
"should list valid modes: {}",
err
);
}
#[test]
fn test_ssl_mode_predicates_form_a_ladder() {
use SslMode::*;
let ladder = [Disable, Prefer, Require, VerifyCa, VerifyFull];
assert_eq!(
ladder.map(|m| m.requires_tls()),
[false, false, true, true, true]
);
assert_eq!(
ladder.map(|m| m.verifies_certificate()),
[false, false, false, true, true]
);
}
#[test]
fn test_ssl_root_cert_defaults_to_none() {
assert_eq!(WaypointConfig::default().database.ssl_root_cert, None);
}
#[test]
fn test_toml_ssl_root_cert_key() {
let toml_str = r#"
[database]
url = "postgres://user@localhost/mydb"
ssl_mode = "verify-full"
ssl_root_cert = "/etc/ssl/certs/internal-ca.pem"
"#;
let toml_config: TomlConfig = toml::from_str(toml_str).unwrap();
let mut config = WaypointConfig::default();
config.apply_toml(toml_config);
assert_eq!(config.database.ssl_mode, SslMode::VerifyFull);
assert_eq!(
config.database.ssl_root_cert,
Some(PathBuf::from("/etc/ssl/certs/internal-ca.pem"))
);
}
#[test]
fn test_invalid_toml_ssl_mode_keeps_the_default() {
let toml_str = r#"
[database]
url = "postgres://user@localhost/mydb"
ssl_mode = "verify-fulll"
"#;
let toml_config: TomlConfig = toml::from_str(toml_str).unwrap();
let mut config = WaypointConfig::default();
config.apply_toml(toml_config);
assert_eq!(config.database.ssl_mode, SslMode::Prefer);
}
#[test]
fn test_multi_database_inherits_tls_settings() {
let toml_str = r#"
[database]
ssl_mode = "verify-full"
ssl_root_cert = "/etc/ssl/ca.pem"
[[databases]]
name = "orders"
url = "postgres://user@localhost/orders"
"#;
let toml_config: TomlConfig = toml::from_str(toml_str).unwrap();
let mut config = WaypointConfig::default();
config.apply_toml(toml_config);
let multi = config.multi_database.expect("expected a multi-db config");
let db = &multi[0];
assert_eq!(db.database.ssl_mode, SslMode::VerifyFull);
assert_eq!(
db.database.ssl_root_cert,
Some(PathBuf::from("/etc/ssl/ca.pem"))
);
}
#[test]
fn test_multi_database_inherits_top_level_migration_settings() {
let toml_str = r#"
[migrations]
table = "custom_history"
schema = "app"
out_of_order = true
validate_on_migrate = false
[[databases]]
name = "orders"
url = "postgres://user@localhost/orders"
[databases.migrations]
locations = ["db/orders"]
"#;
let toml_config: TomlConfig = toml::from_str(toml_str).unwrap();
let mut config = WaypointConfig::default();
config.apply_toml(toml_config);
let multi = config.multi_database.expect("expected a multi-db config");
let db = &multi[0];
assert_eq!(db.migrations.locations, vec![PathBuf::from("db/orders")]);
assert_eq!(db.migrations.table, "custom_history");
assert_eq!(db.migrations.schema, "app");
assert!(db.migrations.out_of_order);
assert!(!db.migrations.validate_on_migrate);
}
#[test]
fn test_apply_env_number_keeps_previous_value_on_garbage() {
let mut timeout: u32 = 30;
apply_env_number(&mut timeout, "45", "WAYPOINT_CONNECT_TIMEOUT");
assert_eq!(timeout, 45, "a valid value must be applied");
apply_env_number(&mut timeout, "30s", "WAYPOINT_CONNECT_TIMEOUT");
assert_eq!(
timeout, 45,
"a unit suffix must not silently reset the value"
);
apply_env_number(&mut timeout, "", "WAYPOINT_CONNECT_TIMEOUT");
assert_eq!(timeout, 45);
}
#[test]
fn test_parse_env_bool_accepts_common_spellings() {
for v in ["1", "true", "TRUE", "yes", "on", " True "] {
assert_eq!(parse_env_bool(v, "T"), Some(true), "input: {v:?}");
}
for v in ["0", "false", "FALSE", "no", "off"] {
assert_eq!(parse_env_bool(v, "T"), Some(false), "input: {v:?}");
}
}
#[test]
fn test_parse_env_bool_rejects_unknown_instead_of_defaulting_to_false() {
for v in ["y", "enabled", "maybe", "tru"] {
assert_eq!(
parse_env_bool(v, "WAYPOINT_BATCH_TRANSACTION"),
None,
"input: {v:?}"
);
}
}
}