// Package config resolves the waitfordb runtime configuration from environment // variables and formats a secret-free summary for logging. package config import ( "fmt" "strings" "time" ) // Defaults applied when the corresponding env var is unset. const ( DefaultDriver = "postgres" DefaultHost = "localhost" DefaultInterval = 2 * time.Second DefaultConnectTimeout = 5 * time.Second ) // defaultPort maps a driver to the port used when none is configured. var defaultPort = map[string]string{ "postgres": "5432", "mysql": "3306", "valkey": "6379", } // Secret formats natively understood via WAITFORDB_SECRET_FORMAT (or // autodetection): operator-generated secrets injected wholesale with envFrom. const ( FormatCNPG = "cnpg" // CNPG -app secret: host/port/dbname/user/password keys FormatValkey = "valkey" // valkey-operator user secrets: key = username, value = password ) // Config is the resolved, validated configuration for a single run. type Config struct { Driver string // SecretFormat records which operator secret convention supplied the // connection parameters ("" when none). SecretFormat string // DSN, when set, is a full driver-native connection string that overrides // the discrete Host/Port/User/Password/Database fields. DSN string Host string Port string User string Password string Database string SSLMode string // Timeout is the total budget to wait for readiness. Zero means wait // forever. Timeout time.Duration // Interval is the gap between retries. Interval time.Duration // ConnectTimeout bounds a single connect+ping attempt. ConnectTimeout time.Duration } // Error is a configuration error; main maps it to exit code 2. type Error struct{ msg string } func (e *Error) Error() string { return e.msg } func errf(format string, a ...any) *Error { return &Error{msg: fmt.Sprintf(format, a...)} } // Getenv matches os.Getenv; injected in tests. type Getenv func(string) string // Load resolves configuration from env. Connection-parameter precedence is // DSN > WAITFORDB_* > PG* > operator-secret keys (the libpq fallback applies // to the postgres driver only). func Load(get Getenv) (Config, error) { c := Config{ Driver: strings.ToLower(strings.TrimSpace(get("WAITFORDB_DRIVER"))), DSN: get("WAITFORDB_DSN"), ConnectTimeout: DefaultConnectTimeout, Interval: DefaultInterval, } if c.Driver == "redis" { c.Driver = "valkey" } c.Host = get("WAITFORDB_HOST") c.Port = get("WAITFORDB_PORT") c.User = get("WAITFORDB_USER") c.Password = get("WAITFORDB_PASSWORD") c.Database = get("WAITFORDB_DATABASE") c.SSLMode = get("WAITFORDB_SSLMODE") format := strings.ToLower(strings.TrimSpace(get("WAITFORDB_SECRET_FORMAT"))) formatExplicit := format != "" if formatExplicit { if format != FormatCNPG && format != FormatValkey { return Config{}, errf("unsupported WAITFORDB_SECRET_FORMAT %q (supported: %s, %s)", format, FormatCNPG, FormatValkey) } } else { format = detectFormat(get, c) } // An explicit driver wins; otherwise the secret format implies it, falling // back to the historical postgres default. if c.Driver == "" { switch format { case FormatValkey: c.Driver = "valkey" default: c.Driver = DefaultDriver } } if formatExplicit && !formatMatchesDriver(format, c.Driver) { return Config{}, errf("WAITFORDB_SECRET_FORMAT %q does not apply to driver %q", format, c.Driver) } if formatMatchesDriver(format, c.Driver) { c.SecretFormat = format } // The libpq PG* fallback for still-unset postgres fields. if c.Driver == "postgres" { fill(&c.Host, get("PGHOST")) fill(&c.Port, get("PGPORT")) fill(&c.User, get("PGUSER")) fill(&c.Password, get("PGPASSWORD")) fill(&c.Database, get("PGDATABASE")) fill(&c.SSLMode, get("PGSSLMODE")) } // Operator-secret keys fill whatever is still unset. switch c.SecretFormat { case FormatCNPG: fill(&c.Host, get("host")) fill(&c.Port, get("port")) fill(&c.Database, get("dbname")) fill(&c.User, firstNonEmpty(get("user"), get("username"))) fill(&c.Password, get("password")) case FormatValkey: if c.User == "" && get("_operator") != "" { c.User = "_operator" } if c.Password == "" && c.User != "" { c.Password = get(c.User) } } if c.Host == "" { c.Host = DefaultHost } if c.Port == "" { c.Port = defaultPort[c.Driver] } var err error if c.Timeout, err = parseDuration(get("WAITFORDB_TIMEOUT"), 0); err != nil { return Config{}, errf("WAITFORDB_TIMEOUT: %v", err) } if c.Interval, err = parseDuration(get("WAITFORDB_INTERVAL"), DefaultInterval); err != nil { return Config{}, errf("WAITFORDB_INTERVAL: %v", err) } if c.ConnectTimeout, err = parseDuration(get("WAITFORDB_CONNECT_TIMEOUT"), DefaultConnectTimeout); err != nil { return Config{}, errf("WAITFORDB_CONNECT_TIMEOUT: %v", err) } if err := c.validate(); err != nil { return Config{}, err } return c, nil } // detectFormat recognises operator-generated secrets injected via envFrom. // CNPG -app secrets carry lowercase host/dbname/user keys; valkey // user secrets carry one key per username (system users start with "_"). func detectFormat(get Getenv, c Config) string { if get("host") != "" && get("dbname") != "" && (get("user") != "" || get("username") != "") { return FormatCNPG } if get("_operator") != "" || get("_replication") != "" { return FormatValkey } // An explicit ACL user whose password key is present, with no SQL database // configured anywhere, is the valkey users-secret shape. if (c.Driver == "" || c.Driver == "valkey") && c.User != "" && c.Password == "" && c.Database == "" && get("PGDATABASE") == "" && get(c.User) != "" { return FormatValkey } return "" } func formatMatchesDriver(format, driver string) bool { switch format { case FormatCNPG: return driver == "postgres" case FormatValkey: return driver == "valkey" } return false } func fill(dst *string, v string) { if *dst == "" { *dst = v } } func (c Config) validate() error { if _, ok := defaultPort[c.Driver]; !ok { return errf("unsupported WAITFORDB_DRIVER %q (supported: postgres, mysql, valkey/redis)", c.Driver) } if c.Interval <= 0 { return errf("WAITFORDB_INTERVAL must be > 0") } if c.ConnectTimeout <= 0 { return errf("WAITFORDB_CONNECT_TIMEOUT must be > 0") } if c.Timeout < 0 { return errf("WAITFORDB_TIMEOUT must be >= 0") } // With a DSN the discrete fields are optional (the DSN carries them). // Valkey needs neither a database nor a user (unauthenticated PING is valid). if c.DSN == "" && c.Driver != "valkey" { if c.Database == "" { return errf("no database configured: set WAITFORDB_DATABASE (or PGDATABASE for postgres) or WAITFORDB_DSN") } if c.User == "" { return errf("no user configured: set WAITFORDB_USER (or PGUSER for postgres) or WAITFORDB_DSN") } } return nil } // Redacted returns a single-line, password-free summary of the resolved // configuration suitable for the startup log line. func (c Config) Redacted() string { var b strings.Builder fmt.Fprintf(&b, "driver=%s", c.Driver) if c.SecretFormat != "" { fmt.Fprintf(&b, " secret_format=%s", c.SecretFormat) } if c.DSN != "" { fmt.Fprintf(&b, " dsn=%s", redactDSN(c.DSN)) } else { fmt.Fprintf(&b, " addr=%s:%s", c.Host, c.Port) if c.Database != "" { fmt.Fprintf(&b, " database=%s", c.Database) } fmt.Fprintf(&b, " user=%s password=%s", c.User, redactSecret(c.Password)) if c.SSLMode != "" { fmt.Fprintf(&b, " sslmode=%s", c.SSLMode) } } fmt.Fprintf(&b, " timeout=%s interval=%s connect_timeout=%s", timeoutStr(c.Timeout), c.Interval, c.ConnectTimeout) return b.String() } func redactSecret(s string) string { if s == "" { return "(unset)" } return "***" } func timeoutStr(d time.Duration) string { if d == 0 { return "forever" } return d.String() } // redactDSN masks the password in either a URL-style or keyword-style DSN so it // never reaches a log line. func redactDSN(dsn string) string { // URL form: scheme://user:password@host/... if i := strings.Index(dsn, "://"); i >= 0 { rest := dsn[i+3:] if at := strings.Index(rest, "@"); at >= 0 { creds := rest[:at] if colon := strings.Index(creds, ":"); colon >= 0 { return dsn[:i+3] + creds[:colon] + ":***@" + rest[at+1:] } } return dsn } // Keyword form: key=value pairs and mysql user:pass@tcp(...) form. out := dsn for _, key := range []string{"password", "passwd"} { out = redactKeyword(out, key) } // mysql DSN: user:pass@tcp(host)/db if at := strings.Index(out, "@tcp("); at >= 0 { if colon := strings.LastIndex(out[:at], ":"); colon >= 0 { out = out[:colon+1] + "***" + out[at:] } } return out } func redactKeyword(dsn, key string) string { lower := strings.ToLower(dsn) idx := strings.Index(lower, key+"=") if idx < 0 { return dsn } valStart := idx + len(key) + 1 valEnd := valStart for valEnd < len(dsn) && dsn[valEnd] != ' ' { valEnd++ } return dsn[:valStart] + "***" + dsn[valEnd:] } func firstNonEmpty(vals ...string) string { for _, v := range vals { if v != "" { return v } } return "" } func parseDuration(s string, def time.Duration) (time.Duration, error) { s = strings.TrimSpace(s) if s == "" { return def, nil } d, err := time.ParseDuration(s) if err != nil { return 0, fmt.Errorf("invalid duration %q", s) } return d, nil }