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.
This commit is contained in:
@@ -0,0 +1,49 @@
|
||||
// Package driver abstracts the per-database connection details behind a small
|
||||
// interface so new engines can be added without touching the wait loop.
|
||||
package driver
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"sort"
|
||||
"strings"
|
||||
|
||||
"git.unkin.net/unkin/waitfordb/internal/config"
|
||||
)
|
||||
|
||||
// Pinger is a live handle to a database that can answer a trivial liveness
|
||||
// query. A successful Ping proves the server is up, auth succeeded, and the
|
||||
// target database/role exist.
|
||||
type Pinger interface {
|
||||
Ping(ctx context.Context) error
|
||||
Close() error
|
||||
}
|
||||
|
||||
// Driver knows how to open a Pinger for one database engine.
|
||||
type Driver interface {
|
||||
Name() string
|
||||
Open(cfg config.Config) (Pinger, error)
|
||||
}
|
||||
|
||||
var registry = map[string]Driver{}
|
||||
|
||||
func register(d Driver) { registry[d.Name()] = d }
|
||||
|
||||
// Get returns the registered driver for name.
|
||||
func Get(name string) (Driver, error) {
|
||||
d, ok := registry[strings.ToLower(name)]
|
||||
if !ok {
|
||||
return nil, fmt.Errorf("unsupported driver %q (supported: %s)", name, strings.Join(Names(), ", "))
|
||||
}
|
||||
return d, nil
|
||||
}
|
||||
|
||||
// Names lists the registered driver names, sorted.
|
||||
func Names() []string {
|
||||
out := make([]string, 0, len(registry))
|
||||
for n := range registry {
|
||||
out = append(out, n)
|
||||
}
|
||||
sort.Strings(out)
|
||||
return out
|
||||
}
|
||||
@@ -0,0 +1,61 @@
|
||||
package driver
|
||||
|
||||
import (
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"git.unkin.net/unkin/waitfordb/internal/config"
|
||||
)
|
||||
|
||||
func TestGetRegisteredDrivers(t *testing.T) {
|
||||
for _, name := range []string{"postgres", "mysql"} {
|
||||
if _, err := Get(name); err != nil {
|
||||
t.Errorf("Get(%q) failed: %v", name, err)
|
||||
}
|
||||
}
|
||||
if _, err := Get("oracle"); err == nil {
|
||||
t.Error("Get(oracle) should fail")
|
||||
}
|
||||
}
|
||||
|
||||
func TestPostgresDSN(t *testing.T) {
|
||||
cfg := config.Config{
|
||||
Host: "db.example", Port: "5432", User: "sonarr",
|
||||
Password: "p@ss/w:rd", Database: "sonarr-main",
|
||||
SSLMode: "disable", ConnectTimeout: 5e9,
|
||||
}
|
||||
dsn := postgresDSN(cfg)
|
||||
if !strings.HasPrefix(dsn, "postgres://sonarr:") {
|
||||
t.Errorf("unexpected prefix: %s", dsn)
|
||||
}
|
||||
// The special-character password must be percent-escaped, not raw.
|
||||
if strings.Contains(dsn, "p@ss/w:rd") {
|
||||
t.Errorf("password not escaped in DSN: %s", dsn)
|
||||
}
|
||||
if !strings.Contains(dsn, "db.example:5432") {
|
||||
t.Errorf("missing host:port: %s", dsn)
|
||||
}
|
||||
if !strings.Contains(dsn, "sslmode=disable") {
|
||||
t.Errorf("missing sslmode: %s", dsn)
|
||||
}
|
||||
if !strings.Contains(dsn, "connect_timeout=5") {
|
||||
t.Errorf("missing connect_timeout: %s", dsn)
|
||||
}
|
||||
if !strings.Contains(dsn, "/sonarr-main") {
|
||||
t.Errorf("missing dbname: %s", dsn)
|
||||
}
|
||||
}
|
||||
|
||||
func TestMySQLDSN(t *testing.T) {
|
||||
cfg := config.Config{
|
||||
Host: "db", Port: "3306", User: "u", Password: "pw",
|
||||
Database: "app", ConnectTimeout: 5e9,
|
||||
}
|
||||
dsn := mysqlDSN(cfg)
|
||||
if !strings.Contains(dsn, "@tcp(db:3306)/app") {
|
||||
t.Errorf("unexpected mysql dsn: %s", dsn)
|
||||
}
|
||||
if !strings.Contains(dsn, "timeout=5s") {
|
||||
t.Errorf("missing timeout: %s", dsn)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,35 @@
|
||||
package driver
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
|
||||
"git.unkin.net/unkin/waitfordb/internal/config"
|
||||
"github.com/go-sql-driver/mysql"
|
||||
)
|
||||
|
||||
func init() { register(mysqlDriver{}) }
|
||||
|
||||
type mysqlDriver struct{}
|
||||
|
||||
func (mysqlDriver) Name() string { return "mysql" }
|
||||
|
||||
func (mysqlDriver) Open(cfg config.Config) (Pinger, error) {
|
||||
dsn := cfg.DSN
|
||||
if dsn == "" {
|
||||
dsn = mysqlDSN(cfg)
|
||||
}
|
||||
return openSQL("mysql", dsn)
|
||||
}
|
||||
|
||||
// mysqlDSN builds a driver-native DSN via mysql.Config so credentials and the
|
||||
// address are escaped correctly.
|
||||
func mysqlDSN(cfg config.Config) string {
|
||||
c := mysql.NewConfig()
|
||||
c.User = cfg.User
|
||||
c.Passwd = cfg.Password
|
||||
c.Net = "tcp"
|
||||
c.Addr = fmt.Sprintf("%s:%s", cfg.Host, cfg.Port)
|
||||
c.DBName = cfg.Database
|
||||
c.Timeout = cfg.ConnectTimeout
|
||||
return c.FormatDSN()
|
||||
}
|
||||
@@ -0,0 +1,48 @@
|
||||
package driver
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"net/url"
|
||||
"strconv"
|
||||
|
||||
"git.unkin.net/unkin/waitfordb/internal/config"
|
||||
_ "github.com/jackc/pgx/v5/stdlib" // registers the "pgx" database/sql driver
|
||||
)
|
||||
|
||||
func init() { register(postgres{}) }
|
||||
|
||||
type postgres struct{}
|
||||
|
||||
func (postgres) Name() string { return "postgres" }
|
||||
|
||||
func (postgres) Open(cfg config.Config) (Pinger, error) {
|
||||
dsn := cfg.DSN
|
||||
if dsn == "" {
|
||||
dsn = postgresDSN(cfg)
|
||||
}
|
||||
return openSQL("pgx", dsn)
|
||||
}
|
||||
|
||||
// postgresDSN builds a URL-style DSN. net/url escapes the userinfo and query so
|
||||
// passwords with special characters are handled safely.
|
||||
func postgresDSN(cfg config.Config) string {
|
||||
u := url.URL{
|
||||
Scheme: "postgres",
|
||||
Host: fmt.Sprintf("%s:%s", cfg.Host, cfg.Port),
|
||||
Path: "/" + cfg.Database,
|
||||
}
|
||||
if cfg.User != "" {
|
||||
u.User = url.UserPassword(cfg.User, cfg.Password)
|
||||
}
|
||||
q := url.Values{}
|
||||
if cfg.SSLMode != "" {
|
||||
q.Set("sslmode", cfg.SSLMode)
|
||||
}
|
||||
// connect_timeout is a per-attempt safety net in addition to the context
|
||||
// deadline the wait loop applies; it is in whole seconds.
|
||||
if secs := int(cfg.ConnectTimeout.Seconds()); secs > 0 {
|
||||
q.Set("connect_timeout", strconv.Itoa(secs))
|
||||
}
|
||||
u.RawQuery = q.Encode()
|
||||
return u.String()
|
||||
}
|
||||
@@ -0,0 +1,31 @@
|
||||
package driver
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
)
|
||||
|
||||
// sqlPinger runs the liveness query against a database/sql handle. It is shared
|
||||
// by every engine whose Go driver plugs into database/sql.
|
||||
type sqlPinger struct {
|
||||
db *sql.DB
|
||||
}
|
||||
|
||||
func openSQL(driverName, dsn string) (*sqlPinger, error) {
|
||||
db, err := sql.Open(driverName, dsn)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
// A readiness check only ever needs one connection; keep the pool tiny so a
|
||||
// failed attempt does not leave idle half-open connections behind.
|
||||
db.SetMaxOpenConns(1)
|
||||
db.SetMaxIdleConns(0)
|
||||
return &sqlPinger{db: db}, nil
|
||||
}
|
||||
|
||||
func (p *sqlPinger) Ping(ctx context.Context) error {
|
||||
var one int
|
||||
return p.db.QueryRowContext(ctx, "SELECT 1").Scan(&one)
|
||||
}
|
||||
|
||||
func (p *sqlPinger) Close() error { return p.db.Close() }
|
||||
Reference in New Issue
Block a user