Files
waitfordb/internal/config/config_test.go
T
unkin-agent 2b37986218 waitfordb: initial tool — env-configured wait-for-DB init container
Small Go tool + distroless container used as a K8s initContainer to block an
app until its database (Postgres/MySQL) is reachable. Env-var configured
(WAITFORDB_* + libpq PG* fallback), configurable timeout/interval, redacted
logs, exit codes. Woodpecker CI publishes docker-internal/waitfordb on tag.
2026-08-22 13:31:44 +10:00

174 lines
4.9 KiB
Go

package config
import (
"strings"
"testing"
"time"
)
// envFrom returns a Getenv backed by a map.
func envFrom(m map[string]string) Getenv {
return func(k string) string { return m[k] }
}
func TestLoadDefaults(t *testing.T) {
c, err := Load(envFrom(map[string]string{
"WAITFORDB_USER": "sonarr",
"WAITFORDB_DATABASE": "sonarr-main",
"WAITFORDB_HOST": "db",
}))
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if c.Driver != "postgres" {
t.Errorf("driver = %q, want postgres", c.Driver)
}
if c.Port != "5432" {
t.Errorf("port = %q, want default 5432", c.Port)
}
if c.Interval != 2*time.Second {
t.Errorf("interval = %v, want 2s", c.Interval)
}
if c.ConnectTimeout != 5*time.Second {
t.Errorf("connect_timeout = %v, want 5s", c.ConnectTimeout)
}
if c.Timeout != 0 {
t.Errorf("timeout = %v, want 0 (forever)", c.Timeout)
}
}
func TestPrecedenceDSNWins(t *testing.T) {
c, err := Load(envFrom(map[string]string{
"WAITFORDB_DSN": "postgres://u:p@h:5432/d",
"WAITFORDB_HOST": "ignored",
"WAITFORDB_USER": "ignored",
"WAITFORDB_DATABASE": "ignored",
}))
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if c.DSN == "" {
t.Fatal("DSN should be set")
}
}
func TestPrecedenceWaitfordbOverPG(t *testing.T) {
c, err := Load(envFrom(map[string]string{
"WAITFORDB_HOST": "native-host",
"PGHOST": "pg-host",
"WAITFORDB_USER": "native-user",
"PGUSER": "pg-user",
"PGDATABASE": "pg-db", // only PG* set for database -> fallback used
}))
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if c.Host != "native-host" {
t.Errorf("host = %q, want native-host (WAITFORDB_* wins)", c.Host)
}
if c.User != "native-user" {
t.Errorf("user = %q, want native-user", c.User)
}
if c.Database != "pg-db" {
t.Errorf("database = %q, want pg-db (PG* fallback)", c.Database)
}
}
func TestPGFallbackPostgresOnly(t *testing.T) {
// For a non-postgres driver the PG* fallback must not apply.
_, err := Load(envFrom(map[string]string{
"WAITFORDB_DRIVER": "mysql",
"PGUSER": "pg-user",
"PGDATABASE": "pg-db",
// No WAITFORDB_USER/DATABASE -> must be a config error, PG* ignored.
}))
if err == nil {
t.Fatal("expected config error: PG* must not satisfy mysql config")
}
}
func TestMySQLDefaultPort(t *testing.T) {
c, err := Load(envFrom(map[string]string{
"WAITFORDB_DRIVER": "mysql",
"WAITFORDB_USER": "u",
"WAITFORDB_DATABASE": "d",
}))
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if c.Port != "3306" {
t.Errorf("port = %q, want 3306", c.Port)
}
}
func TestValidationErrors(t *testing.T) {
cases := map[string]map[string]string{
"missing database": {"WAITFORDB_USER": "u"},
"missing user": {"WAITFORDB_DATABASE": "d"},
"bad driver": {"WAITFORDB_DRIVER": "oracle", "WAITFORDB_USER": "u", "WAITFORDB_DATABASE": "d"},
"bad timeout": {"WAITFORDB_USER": "u", "WAITFORDB_DATABASE": "d", "WAITFORDB_TIMEOUT": "nope"},
"zero interval": {"WAITFORDB_USER": "u", "WAITFORDB_DATABASE": "d", "WAITFORDB_INTERVAL": "0"},
}
for name, env := range cases {
t.Run(name, func(t *testing.T) {
if _, err := Load(envFrom(env)); err == nil {
t.Fatalf("expected error for %s", name)
} else if _, ok := err.(*Error); !ok {
t.Fatalf("expected *config.Error, got %T", err)
}
})
}
}
func TestRedactedHidesPassword(t *testing.T) {
c, err := Load(envFrom(map[string]string{
"WAITFORDB_HOST": "h",
"WAITFORDB_USER": "u",
"WAITFORDB_PASSWORD": "sup3r-s3cret",
"WAITFORDB_DATABASE": "d",
"WAITFORDB_TIMEOUT": "5m",
}))
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
got := c.Redacted()
if strings.Contains(got, "sup3r-s3cret") {
t.Fatalf("redacted output leaked password: %q", got)
}
for _, want := range []string{"driver=postgres", "database=d", "user=u", "password=***", "timeout=5m", "interval=2s"} {
if !strings.Contains(got, want) {
t.Errorf("redacted output missing %q: %q", want, got)
}
}
}
func TestRedactedForeverTimeout(t *testing.T) {
c, _ := Load(envFrom(map[string]string{"WAITFORDB_USER": "u", "WAITFORDB_DATABASE": "d"}))
if !strings.Contains(c.Redacted(), "timeout=forever") {
t.Errorf("want timeout=forever, got %q", c.Redacted())
}
}
func TestRedactedDSNMasksPassword(t *testing.T) {
cases := []struct {
name string
dsn string
}{
{"url", "postgres://user:topsecret@host:5432/db?sslmode=disable"},
{"keyword", "host=h user=u password=topsecret dbname=d"},
{"mysql", "user:topsecret@tcp(h:3306)/d?timeout=5s"},
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
c := Config{DSN: tc.dsn, Timeout: 0, Interval: time.Second, ConnectTimeout: time.Second, Driver: "postgres"}
got := c.Redacted()
if strings.Contains(got, "topsecret") {
t.Fatalf("leaked password: %q", got)
}
if !strings.Contains(got, "***") {
t.Errorf("expected redaction marker in %q", got)
}
})
}
}