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:
2026-08-22 13:31:44 +10:00
parent 98971eb1b7
commit 2b37986218
19 changed files with 1467 additions and 1 deletions
+236
View File
@@ -0,0 +1,236 @@
// 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",
}
// Config is the resolved, validated configuration for a single run.
type Config struct {
Driver 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* (the libpq fallback applies to the postgres driver
// only).
func Load(get Getenv) (Config, error) {
c := Config{
Driver: firstNonEmpty(get("WAITFORDB_DRIVER"), DefaultDriver),
DSN: get("WAITFORDB_DSN"),
ConnectTimeout: DefaultConnectTimeout,
Interval: DefaultInterval,
}
c.Driver = strings.ToLower(strings.TrimSpace(c.Driver))
pg := c.Driver == "postgres"
// WAITFORDB_* first, then the libpq PG* fallback for postgres.
c.Host = pick(get, pg, "WAITFORDB_HOST", "PGHOST")
c.Port = pick(get, pg, "WAITFORDB_PORT", "PGPORT")
c.User = pick(get, pg, "WAITFORDB_USER", "PGUSER")
c.Password = pick(get, pg, "WAITFORDB_PASSWORD", "PGPASSWORD")
c.Database = pick(get, pg, "WAITFORDB_DATABASE", "PGDATABASE")
c.SSLMode = pick(get, pg, "WAITFORDB_SSLMODE", "PGSSLMODE")
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
}
func (c Config) validate() error {
if _, ok := defaultPort[c.Driver]; !ok {
return errf("unsupported WAITFORDB_DRIVER %q (supported: postgres, mysql)", 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).
if c.DSN == "" {
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.DSN != "" {
fmt.Fprintf(&b, " dsn=%s", redactDSN(c.DSN))
} else {
fmt.Fprintf(&b, " addr=%s:%s database=%s user=%s password=%s",
c.Host, c.Port, c.Database, 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:]
}
// pick returns the WAITFORDB_* value, falling back to the PG* value only when
// fallback is true (postgres).
func pick(get Getenv, fallback bool, primary, secondary string) string {
if v := get(primary); v != "" {
return v
}
if fallback {
return get(secondary)
}
return ""
}
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
}
+173
View File
@@ -0,0 +1,173 @@
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)
}
})
}
}
+49
View File
@@ -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
}
+61
View File
@@ -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)
}
}
+35
View File
@@ -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()
}
+48
View File
@@ -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()
}
+31
View File
@@ -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() }
+113
View File
@@ -0,0 +1,113 @@
// Package wait implements the retry/timeout loop that polls a database until a
// liveness attempt succeeds. The clock and the attempt are injected so the loop
// is fully testable without a real database or real time.
package wait
import (
"context"
"time"
)
// Clock abstracts time so tests can advance it instantly.
type Clock interface {
Now() time.Time
// Sleep blocks for d or until ctx is done, returning ctx.Err() if it was
// cancelled first.
Sleep(ctx context.Context, d time.Duration) error
}
// AttemptFunc performs one connect+ping. The ctx carries the per-attempt
// connect timeout.
type AttemptFunc func(ctx context.Context) error
// OnFailure is invoked after each failed attempt that will be retried.
type OnFailure func(attempt int, err error, elapsed, timeout, retryIn time.Duration)
// Params configures the loop.
type Params struct {
Timeout time.Duration // 0 = wait forever
Interval time.Duration
ConnectTimeout time.Duration
}
// Result reports how a run ended.
type Result struct {
OK bool
TimedOut bool
Cancelled bool
Attempts int
Elapsed time.Duration
LastErr error
}
// Run polls attempt until it succeeds, the timeout is exhausted, or ctx is
// cancelled. It always makes at least one attempt.
func Run(ctx context.Context, p Params, attempt AttemptFunc, onFail OnFailure, clk Clock) Result {
start := clk.Now()
var deadline time.Time
if p.Timeout > 0 {
deadline = start.Add(p.Timeout)
}
res := Result{}
for {
res.Attempts++
actx, cancel := context.WithTimeout(ctx, p.ConnectTimeout)
err := attempt(actx)
cancel()
now := clk.Now()
res.Elapsed = now.Sub(start)
if err == nil {
res.OK = true
return res
}
res.LastErr = err
// A cancelled parent context (SIGTERM/SIGINT) wins over a retry.
if ctx.Err() != nil {
res.Cancelled = true
return res
}
// No time budget left for another attempt.
if p.Timeout > 0 && !now.Before(deadline) {
res.TimedOut = true
return res
}
sleep := p.Interval
if p.Timeout > 0 {
if remaining := deadline.Sub(now); remaining < sleep {
sleep = remaining
}
}
onFail(res.Attempts, err, res.Elapsed, p.Timeout, sleep)
if serr := clk.Sleep(ctx, sleep); serr != nil {
res.Cancelled = true
return res
}
}
}
// RealClock is the production Clock backed by the wall clock.
type RealClock struct{}
func (RealClock) Now() time.Time { return time.Now() }
func (RealClock) Sleep(ctx context.Context, d time.Duration) error {
if d <= 0 {
return ctx.Err()
}
t := time.NewTimer(d)
defer t.Stop()
select {
case <-ctx.Done():
return ctx.Err()
case <-t.C:
return nil
}
}
+125
View File
@@ -0,0 +1,125 @@
package wait
import (
"context"
"errors"
"testing"
"time"
)
// fakeClock advances instantly on Sleep so the loop runs with no real delay.
type fakeClock struct {
t time.Time
cancelAt time.Duration // if >0, cancel the run once elapsed reaches this
cancel context.CancelFunc
start time.Time
sleepCall int
}
func newFakeClock() *fakeClock {
start := time.Unix(0, 0)
return &fakeClock{t: start, start: start}
}
func (c *fakeClock) Now() time.Time { return c.t }
func (c *fakeClock) Sleep(ctx context.Context, d time.Duration) error {
c.sleepCall++
c.t = c.t.Add(d)
if c.cancelAt > 0 && c.t.Sub(c.start) >= c.cancelAt && c.cancel != nil {
c.cancel()
}
return ctx.Err()
}
var errDown = errors.New("connection refused")
// failNThenOK returns an AttemptFunc that fails the first n calls then succeeds.
func failNThenOK(n int, calls *int) AttemptFunc {
return func(ctx context.Context) error {
*calls++
if *calls <= n {
return errDown
}
return nil
}
}
func noFail(int, error, time.Duration, time.Duration, time.Duration) {}
func TestSucceedsFirstAttempt(t *testing.T) {
calls := 0
res := Run(context.Background(),
Params{Timeout: time.Minute, Interval: 2 * time.Second, ConnectTimeout: time.Second},
failNThenOK(0, &calls), noFail, newFakeClock())
if !res.OK || res.Attempts != 1 {
t.Fatalf("want OK after 1 attempt, got %+v", res)
}
}
func TestWaitsThenSucceeds(t *testing.T) {
calls := 0
clk := newFakeClock()
failures := 0
res := Run(context.Background(),
Params{Timeout: time.Minute, Interval: 2 * time.Second, ConnectTimeout: time.Second},
failNThenOK(3, &calls),
func(int, error, time.Duration, time.Duration, time.Duration) { failures++ },
clk)
if !res.OK {
t.Fatalf("want OK, got %+v", res)
}
if res.Attempts != 4 {
t.Errorf("attempts = %d, want 4", res.Attempts)
}
if failures != 3 {
t.Errorf("onFail called %d times, want 3", failures)
}
// 3 sleeps of 2s each.
if got := res.Elapsed; got != 6*time.Second {
t.Errorf("elapsed = %v, want 6s", got)
}
}
func TestTimesOut(t *testing.T) {
calls := 0
alwaysFail := func(ctx context.Context) error { calls++; return errDown }
res := Run(context.Background(),
Params{Timeout: 10 * time.Second, Interval: 3 * time.Second, ConnectTimeout: time.Second},
alwaysFail, noFail, newFakeClock())
if res.OK || !res.TimedOut {
t.Fatalf("want timeout, got %+v", res)
}
if !errors.Is(res.LastErr, errDown) {
t.Errorf("LastErr = %v, want errDown", res.LastErr)
}
// Deadline 10s, interval 3s: attempts at 0,3,6,9, then next check at ~12s > deadline.
if res.Attempts < 3 {
t.Errorf("attempts = %d, want several before timeout", res.Attempts)
}
}
func TestWaitForeverEventuallySucceeds(t *testing.T) {
calls := 0
res := Run(context.Background(),
Params{Timeout: 0, Interval: time.Second, ConnectTimeout: time.Second},
failNThenOK(100, &calls), noFail, newFakeClock())
if !res.OK || res.Attempts != 101 {
t.Fatalf("want OK after 101 attempts with no timeout, got %+v", res)
}
}
func TestCancelledDuringSleep(t *testing.T) {
ctx, cancel := context.WithCancel(context.Background())
clk := newFakeClock()
clk.cancelAt = 4 * time.Second
clk.cancel = cancel
calls := 0
alwaysFail := func(ctx context.Context) error { calls++; return errDown }
res := Run(ctx,
Params{Timeout: time.Hour, Interval: 2 * time.Second, ConnectTimeout: time.Second},
alwaysFail, noFail, clk)
if !res.Cancelled {
t.Fatalf("want Cancelled, got %+v", res)
}
}