test: add comprehensive test suite for cache, config, logger, and retry functionalities

Signed-off-by: Noooste <83548733+Noooste@users.noreply.github.com>
This commit is contained in:
Noooste
2026-04-17 17:25:59 +02:00
parent f63ce3452e
commit 047a653446
4 changed files with 1032 additions and 0 deletions
+396
View File
@@ -0,0 +1,396 @@
package config
import (
"os"
"path/filepath"
"strings"
"testing"
"github.com/spf13/viper"
)
// writeConfigFile writes yaml content to a temp path and returns it.
func writeConfigFile(t *testing.T, yaml string) string {
t.Helper()
dir := t.TempDir()
path := filepath.Join(dir, "config.yaml")
if err := os.WriteFile(path, []byte(yaml), 0600); err != nil {
t.Fatalf("write config: %v", err)
}
return path
}
// resetViper clears all global viper state between tests.
func resetViper(t *testing.T) {
t.Helper()
viper.Reset()
}
// minimalValidYAML is the smallest configuration that passes Validate.
const minimalValidYAML = `
server:
host: "0.0.0.0"
port: 8080
environment: development
garage:
endpoint: http://garage:3900
admin_endpoint: http://garage:3903
admin_token: supersecret
`
func TestLoad_YAMLOnly(t *testing.T) {
resetViper(t)
path := writeConfigFile(t, minimalValidYAML)
cfg, err := Load(path)
if err != nil {
t.Fatalf("Load: %v", err)
}
if cfg.Server.Host != "0.0.0.0" {
t.Errorf("Server.Host = %q, want 0.0.0.0", cfg.Server.Host)
}
if cfg.Server.Port != 8080 {
t.Errorf("Server.Port = %d, want 8080", cfg.Server.Port)
}
if cfg.Server.Environment != "development" {
t.Errorf("Server.Environment = %q, want development", cfg.Server.Environment)
}
if cfg.Garage.Endpoint != "http://garage:3900" {
t.Errorf("Garage.Endpoint = %q", cfg.Garage.Endpoint)
}
if cfg.Garage.AdminToken != "supersecret" {
t.Errorf("Garage.AdminToken = %q", cfg.Garage.AdminToken)
}
}
func TestLoad_EnvOnly_MissingFile(t *testing.T) {
resetViper(t)
// Point at a path that definitely does not exist. Load tolerates missing
// files and falls back to env + viper defaults.
missing := filepath.Join(t.TempDir(), "does-not-exist.yaml")
// Every required field provided via env.
t.Setenv("GARAGE_UI_SERVER_PORT", "9090")
t.Setenv("GARAGE_UI_GARAGE_ENDPOINT", "http://g:3900")
t.Setenv("GARAGE_UI_GARAGE_ADMIN_ENDPOINT", "http://g:3903")
t.Setenv("GARAGE_UI_GARAGE_ADMIN_TOKEN", "env-token")
cfg, err := Load(missing)
if err != nil {
t.Fatalf("Load with env-only: %v", err)
}
if cfg.Server.Port != 9090 {
t.Errorf("Server.Port = %d, want 9090 (from env)", cfg.Server.Port)
}
if cfg.Garage.AdminToken != "env-token" {
t.Errorf("Garage.AdminToken = %q, want env-token", cfg.Garage.AdminToken)
}
}
func TestLoad_EnvOverridesYAML(t *testing.T) {
resetViper(t)
path := writeConfigFile(t, minimalValidYAML)
// YAML has port=8080; env should win.
t.Setenv("GARAGE_UI_SERVER_PORT", "9090")
t.Setenv("GARAGE_UI_GARAGE_ADMIN_TOKEN", "env-wins")
cfg, err := Load(path)
if err != nil {
t.Fatalf("Load: %v", err)
}
if cfg.Server.Port != 9090 {
t.Errorf("Server.Port = %d, want 9090 (env override)", cfg.Server.Port)
}
if cfg.Garage.AdminToken != "env-wins" {
t.Errorf("Garage.AdminToken = %q, want env-wins", cfg.Garage.AdminToken)
}
// Host was not overridden; YAML value should persist.
if cfg.Server.Host != "0.0.0.0" {
t.Errorf("Server.Host = %q, want 0.0.0.0 (from YAML)", cfg.Server.Host)
}
}
func TestLoad_MalformedYAMLReturnsError(t *testing.T) {
resetViper(t)
// Deliberately broken YAML: unindented key after a mapping start.
path := writeConfigFile(t, "server:\n port: 8080\n:: not: valid ::\n")
_, err := Load(path)
if err == nil {
t.Fatal("expected error for malformed YAML, got nil")
}
if !strings.Contains(err.Error(), "error reading config file") {
t.Errorf("unexpected error: %v", err)
}
}
func TestLoad_ValidationFailurePropagates(t *testing.T) {
resetViper(t)
// Valid YAML syntax but Garage.Endpoint is blank → Validate must fail.
path := writeConfigFile(t, `
server:
port: 8080
garage:
endpoint: ""
admin_endpoint: http://g:3903
admin_token: t
`)
_, err := Load(path)
if err == nil {
t.Fatal("expected validation error, got nil")
}
if !strings.Contains(err.Error(), "invalid configuration") {
t.Errorf("expected wrapped invalid-config error, got %v", err)
}
if !strings.Contains(err.Error(), "garage endpoint is required") {
t.Errorf("expected endpoint-required message, got %v", err)
}
}
// validBaseConfig returns a deep copy of a minimal Config that passes Validate.
func validBaseConfig() Config {
return Config{
Server: ServerConfig{Port: 8080},
Garage: GarageConfig{
Endpoint: "http://g:3900",
AdminEndpoint: "http://g:3903",
AdminToken: "t",
},
}
}
// applyValidOIDC fills OIDC with all required fields.
func applyValidOIDC(c *Config) {
c.Auth.OIDC.Enabled = true
c.Auth.OIDC.ClientID = "client-xyz"
c.Auth.OIDC.IssuerURL = "https://idp.example/realms/test"
c.Auth.OIDC.Scopes = []string{"openid"}
c.Auth.OIDC.AdminRole = "admin"
c.Server.RootURL = "https://garage-ui.example"
}
// Note on spec coverage: spec/2026-04-17-backend-test-suite-design.md lists
// "invalid log level/format" as a Validate case, but the current Validate does
// not check Logging.Level or Logging.Format. That's a code-vs-spec gap to
// resolve in a follow-up plan; Stage 2 tests the current behavior only.
func TestValidate(t *testing.T) {
tests := []struct {
name string
mutate func(*Config)
wantErrContains string // empty = expect no error
}{
{
name: "valid minimal config",
mutate: func(c *Config) {},
},
{
name: "port zero is invalid",
mutate: func(c *Config) { c.Server.Port = 0 },
wantErrContains: "invalid server port",
},
{
name: "port negative is invalid",
mutate: func(c *Config) { c.Server.Port = -1 },
wantErrContains: "invalid server port",
},
{
name: "port above 65535 is invalid",
mutate: func(c *Config) { c.Server.Port = 70000 },
wantErrContains: "invalid server port",
},
{
name: "port at 65535 is valid",
mutate: func(c *Config) { c.Server.Port = 65535 },
wantErrContains: "",
},
{
name: "missing garage endpoint",
mutate: func(c *Config) { c.Garage.Endpoint = "" },
wantErrContains: "garage endpoint is required",
},
{
name: "missing garage admin_endpoint",
mutate: func(c *Config) { c.Garage.AdminEndpoint = "" },
wantErrContains: "admin_endpoint is required",
},
{
name: "missing garage admin_token",
mutate: func(c *Config) { c.Garage.AdminToken = "" },
wantErrContains: "admin_token is required",
},
{
name: "admin auth enabled without username",
mutate: func(c *Config) {
c.Auth.Admin.Enabled = true
c.Auth.Admin.Password = "p"
},
wantErrContains: "admin auth username and password are required",
},
{
name: "admin auth enabled without password",
mutate: func(c *Config) {
c.Auth.Admin.Enabled = true
c.Auth.Admin.Username = "u"
},
wantErrContains: "admin auth username and password are required",
},
{
name: "admin auth enabled with both set is valid",
mutate: func(c *Config) {
c.Auth.Admin.Enabled = true
c.Auth.Admin.Username = "u"
c.Auth.Admin.Password = "p"
},
wantErrContains: "",
},
{
name: "admin auth disabled ignores missing credentials",
mutate: func(c *Config) {
c.Auth.Admin.Enabled = false
c.Auth.Admin.Username = ""
c.Auth.Admin.Password = ""
},
wantErrContains: "",
},
{
name: "oidc enabled without client_id",
mutate: func(c *Config) {
applyValidOIDC(c)
c.Auth.OIDC.ClientID = ""
},
wantErrContains: "oidc client_id is required",
},
{
name: "oidc enabled without issuer_url",
mutate: func(c *Config) {
applyValidOIDC(c)
c.Auth.OIDC.IssuerURL = ""
},
wantErrContains: "oidc issuer_url is required",
},
{
name: "oidc enabled without server.root_url",
mutate: func(c *Config) {
applyValidOIDC(c)
c.Server.RootURL = ""
},
wantErrContains: "server.root_url is required",
},
{
name: "oidc enabled without scopes",
mutate: func(c *Config) {
applyValidOIDC(c)
c.Auth.OIDC.Scopes = nil
},
wantErrContains: "oidc scopes are required",
},
{
name: "oidc enabled without admin_role rejected for safety",
mutate: func(c *Config) {
applyValidOIDC(c)
c.Auth.OIDC.AdminRole = ""
},
wantErrContains: "oidc admin_role is required",
},
{
name: "oidc fully configured is valid",
mutate: applyValidOIDC,
wantErrContains: "",
},
{
name: "oidc disabled ignores missing client_id",
mutate: func(c *Config) {
c.Auth.OIDC.Enabled = false
c.Auth.OIDC.ClientID = ""
},
wantErrContains: "",
},
}
for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
cfg := validBaseConfig()
tc.mutate(&cfg)
err := cfg.Validate()
if tc.wantErrContains == "" {
if err != nil {
t.Errorf("expected no error, got %v", err)
}
return
}
if err == nil {
t.Fatalf("expected error containing %q, got nil", tc.wantErrContains)
}
if !strings.Contains(err.Error(), tc.wantErrContains) {
t.Errorf("error %q does not contain %q", err.Error(), tc.wantErrContains)
}
})
}
}
func TestGetAddress(t *testing.T) {
tests := []struct {
host string
port int
want string
}{
{"localhost", 8080, "localhost:8080"},
{"0.0.0.0", 80, "0.0.0.0:80"},
{"", 443, ":443"},
}
for _, tc := range tests {
t.Run(tc.want, func(t *testing.T) {
cfg := &Config{Server: ServerConfig{Host: tc.host, Port: tc.port}}
if got := cfg.GetAddress(); got != tc.want {
t.Errorf("GetAddress() = %q, want %q", got, tc.want)
}
})
}
}
func TestIsDevelopment(t *testing.T) {
tests := []struct {
env string
want bool
}{
{"development", true},
{"production", false},
{"", false},
// Case-sensitive per current impl; lock in that behavior.
{"Development", false},
{"DEV", false},
}
for _, tc := range tests {
t.Run(tc.env, func(t *testing.T) {
cfg := &Config{Server: ServerConfig{Environment: tc.env}}
if got := cfg.IsDevelopment(); got != tc.want {
t.Errorf("IsDevelopment(%q) = %v, want %v", tc.env, got, tc.want)
}
})
}
}
func TestIsProduction(t *testing.T) {
tests := []struct {
env string
want bool
}{
{"production", true},
{"development", false},
{"", false},
{"Production", false},
{"PROD", false},
}
for _, tc := range tests {
t.Run(tc.env, func(t *testing.T) {
cfg := &Config{Server: ServerConfig{Environment: tc.env}}
if got := cfg.IsProduction(); got != tc.want {
t.Errorf("IsProduction(%q) = %v, want %v", tc.env, got, tc.want)
}
})
}
}
+239
View File
@@ -0,0 +1,239 @@
package logger
import (
"bufio"
"encoding/json"
"io"
"os"
"strings"
"sync"
"testing"
)
// serializeLoggerTests guards the global mutations (os.Stdout, globalLogger,
// zerolog global). These tests cannot run in parallel with each other.
var serializeLoggerTests sync.Mutex
// captureStdout swaps os.Stdout for a pipe, calls fn, restores stdout, and
// returns everything written during fn.
func captureStdout(t *testing.T, fn func()) string {
t.Helper()
r, w, err := os.Pipe()
if err != nil {
t.Fatalf("os.Pipe: %v", err)
}
old := os.Stdout
os.Stdout = w
t.Cleanup(func() { os.Stdout = old })
// Run fn and close writer so the reader unblocks.
doneWrite := make(chan struct{})
go func() {
fn()
_ = w.Close()
close(doneWrite)
}()
var buf strings.Builder
scanner := bufio.NewScanner(r)
scanner.Buffer(make([]byte, 64*1024), 1024*1024)
for scanner.Scan() {
buf.WriteString(scanner.Text())
buf.WriteByte('\n')
}
// Drain any residual (shouldn't happen after Close, but safe):
_, _ = io.Copy(io.Discard, r)
<-doneWrite
return buf.String()
}
func TestInit_JSONFormatProducesParseableOutput(t *testing.T) {
serializeLoggerTests.Lock()
defer serializeLoggerTests.Unlock()
out := captureStdout(t, func() {
Init(Config{Level: "info", Format: "json"})
Info().Str("user", "alice").Msg("hello")
})
// Find the first non-empty line; parse as JSON.
var line string
for l := range strings.SplitSeq(out, "\n") {
if strings.TrimSpace(l) != "" {
line = l
break
}
}
if line == "" {
t.Fatalf("no log output captured; stdout = %q", out)
}
var parsed map[string]any
if err := json.Unmarshal([]byte(line), &parsed); err != nil {
t.Fatalf("log line is not valid JSON: %v\nline: %s", err, line)
}
// Field assertions — zerolog uses "message" for the msg and "level" for level.
if got, _ := parsed["message"].(string); got != "hello" {
t.Errorf("message = %v, want hello", parsed["message"])
}
if got, _ := parsed["user"].(string); got != "alice" {
t.Errorf("user field = %v, want alice", parsed["user"])
}
if got, _ := parsed["level"].(string); got != "info" {
t.Errorf("level = %v, want info", parsed["level"])
}
if _, ok := parsed["time"]; !ok {
t.Errorf("expected time field; got keys %v", keysOf(parsed))
}
if _, ok := parsed["caller"]; !ok {
t.Errorf("expected caller field; got keys %v", keysOf(parsed))
}
}
func TestInit_LevelFilterDropsBelowThreshold(t *testing.T) {
serializeLoggerTests.Lock()
defer serializeLoggerTests.Unlock()
out := captureStdout(t, func() {
Init(Config{Level: "warn", Format: "json"})
Debug().Msg("debug-dropped")
Info().Msg("info-dropped")
Warn().Msg("warn-kept")
Error().Msg("error-kept")
})
if strings.Contains(out, "debug-dropped") {
t.Errorf("debug event leaked through warn filter: %s", out)
}
if strings.Contains(out, "info-dropped") {
t.Errorf("info event leaked through warn filter: %s", out)
}
if !strings.Contains(out, "warn-kept") {
t.Errorf("warn event missing: %s", out)
}
if !strings.Contains(out, "error-kept") {
t.Errorf("error event missing: %s", out)
}
}
func TestInit_UnknownLevelDefaultsToInfo(t *testing.T) {
serializeLoggerTests.Lock()
defer serializeLoggerTests.Unlock()
out := captureStdout(t, func() {
Init(Config{Level: "gibberish", Format: "json"})
Debug().Msg("debug-should-be-dropped")
Info().Msg("info-should-appear")
})
if strings.Contains(out, "debug-should-be-dropped") {
t.Errorf("debug leaked at default info level: %s", out)
}
if !strings.Contains(out, "info-should-appear") {
t.Errorf("info missing at default info level: %s", out)
}
}
func TestInit_TextFormatDoesNotCrashAndIsNotJSON(t *testing.T) {
serializeLoggerTests.Lock()
defer serializeLoggerTests.Unlock()
out := captureStdout(t, func() {
Init(Config{Level: "info", Format: "text"})
Info().Str("k", "v").Msg("plain")
})
if !strings.Contains(out, "plain") {
t.Errorf("text output missing message: %s", out)
}
// Console writer output is ANSI-colored key=value form, not JSON.
var parsed map[string]any
if json.Unmarshal([]byte(strings.Split(out, "\n")[0]), &parsed) == nil {
t.Errorf("text format unexpectedly parsed as JSON: %s", out)
}
}
func TestGet_AutoInitializesWhenUnused(t *testing.T) {
serializeLoggerTests.Lock()
defer serializeLoggerTests.Unlock()
// Forcibly clear the global so Get() hits the lazy-init branch.
globalLogger = nil
l := Get()
if l == nil {
t.Fatal("Get() returned nil; lazy init did not run")
}
if globalLogger == nil {
t.Fatal("globalLogger still nil after Get()")
}
}
func TestWithComponent_AddsComponentField(t *testing.T) {
serializeLoggerTests.Lock()
defer serializeLoggerTests.Unlock()
out := captureStdout(t, func() {
Init(Config{Level: "info", Format: "json"})
comp := WithComponent("buckets")
comp.Info().Msg("tagged")
})
line := firstNonEmptyLine(out)
var parsed map[string]any
if err := json.Unmarshal([]byte(line), &parsed); err != nil {
t.Fatalf("not JSON: %v — %s", err, line)
}
if got, _ := parsed["component"].(string); got != "buckets" {
t.Errorf("component = %v, want buckets", parsed["component"])
}
}
func TestLogger_WithContext_AddsFields(t *testing.T) {
serializeLoggerTests.Lock()
defer serializeLoggerTests.Unlock()
out := captureStdout(t, func() {
Init(Config{Level: "info", Format: "json"})
l := Get().WithContext(map[string]any{
"request_id": "req-42",
"attempt": 2,
})
l.Info().Msg("ctx")
})
line := firstNonEmptyLine(out)
var parsed map[string]any
if err := json.Unmarshal([]byte(line), &parsed); err != nil {
t.Fatalf("not JSON: %v — %s", err, line)
}
if parsed["request_id"] != "req-42" {
t.Errorf("request_id = %v", parsed["request_id"])
}
// JSON numbers decode to float64.
if got, _ := parsed["attempt"].(float64); got != 2 {
t.Errorf("attempt = %v, want 2", parsed["attempt"])
}
}
// --- helpers ---
func firstNonEmptyLine(s string) string {
for l := range strings.SplitSeq(s, "\n") {
if strings.TrimSpace(l) != "" {
return l
}
}
return ""
}
func keysOf(m map[string]any) []string {
out := make([]string, 0, len(m))
for k := range m {
out = append(out, k)
}
return out
}
+129
View File
@@ -0,0 +1,129 @@
package utils
import (
"fmt"
"sync"
"testing"
"time"
)
func TestCache_GetMissReturnsNil(t *testing.T) {
c := NewCache()
if v := c.Get("nope"); v != nil {
t.Errorf("expected nil for missing key, got %v", v)
}
}
func TestCache_SetThenGetReturnsValue(t *testing.T) {
c := NewCache()
c.Set("k", "v", time.Minute)
got := c.Get("k")
if got != "v" {
t.Errorf("Get(k) = %v, want v", got)
}
}
func TestCache_SetWithDifferentTypes(t *testing.T) {
c := NewCache()
c.Set("str", "hello", time.Minute)
c.Set("int", 42, time.Minute)
c.Set("slice", []int{1, 2, 3}, time.Minute)
if got := c.Get("str"); got != "hello" {
t.Errorf("str: got %v", got)
}
if got := c.Get("int"); got != 42 {
t.Errorf("int: got %v", got)
}
if got, ok := c.Get("slice").([]int); !ok || len(got) != 3 {
t.Errorf("slice: got %v", c.Get("slice"))
}
}
func TestCache_GetExpiredReturnsNil(t *testing.T) {
c := NewCache()
c.Set("k", "v", 10*time.Millisecond)
time.Sleep(25 * time.Millisecond)
if got := c.Get("k"); got != nil {
t.Errorf("expected nil after TTL, got %v", got)
}
}
func TestCache_DeleteRemovesItem(t *testing.T) {
c := NewCache()
c.Set("k", "v", time.Minute)
c.Delete("k")
if got := c.Get("k"); got != nil {
t.Errorf("expected nil after Delete, got %v", got)
}
}
func TestCache_DeleteMissingKeyIsNoOp(t *testing.T) {
c := NewCache()
// Should not panic or error.
c.Delete("never-set")
}
func TestCache_ClearRemovesAllItems(t *testing.T) {
c := NewCache()
c.Set("a", 1, time.Minute)
c.Set("b", 2, time.Minute)
c.Set("c", 3, time.Minute)
c.Clear()
if c.Get("a") != nil || c.Get("b") != nil || c.Get("c") != nil {
t.Errorf("expected all items cleared")
}
}
func TestCache_SetOverwrites(t *testing.T) {
c := NewCache()
c.Set("k", "v1", time.Minute)
c.Set("k", "v2", time.Minute)
if got := c.Get("k"); got != "v2" {
t.Errorf("expected v2 after overwrite, got %v", got)
}
}
// TestCache_ConcurrentAccess exercises the RWMutex under load. Run with
// `go test -race` to catch data races. Uses bounded concurrency so the test
// stays deterministic.
func TestCache_ConcurrentAccess(t *testing.T) {
c := NewCache()
const goroutines = 50
const opsPerGoroutine = 100
var wg sync.WaitGroup
wg.Add(goroutines)
for g := range goroutines {
go func(id int) {
defer wg.Done()
for i := range opsPerGoroutine {
key := fmt.Sprintf("k%d", (id+i)%10)
c.Set(key, i, time.Minute)
_ = c.Get(key)
if i%10 == 0 {
c.Delete(key)
}
}
}(g)
}
wg.Wait()
// If we got here without a panic and `-race` is clean, the RWMutex is
// protecting the map correctly.
}
// TestGlobalCache_IsUsable is a smoke test for the package-level var.
// It doesn't Clear() afterwards because the global is shared state that
// other packages may depend on at test time.
func TestGlobalCache_IsUsable(t *testing.T) {
key := "stage2-smoke-key"
GlobalCache.Set(key, "x", time.Minute)
t.Cleanup(func() { GlobalCache.Delete(key) })
if got := GlobalCache.Get(key); got != "x" {
t.Errorf("GlobalCache.Get = %v, want x", got)
}
}
+268
View File
@@ -0,0 +1,268 @@
package utils
import (
"context"
"errors"
"fmt"
"net"
"strings"
"syscall"
"testing"
"time"
)
func TestIsConnectionRefused(t *testing.T) {
tests := []struct {
name string
err error
want bool
}{
{
name: "nil error returns false",
err: nil,
want: false,
},
{
name: "unrelated error returns false",
err: errors.New("something else went wrong"),
want: false,
},
{
name: "bare ECONNREFUSED returns true (fallback errors.Is branch)",
err: syscall.ECONNREFUSED,
want: true,
},
{
name: "wrapped ECONNREFUSED returns true (fallback errors.Is branch)",
err: fmt.Errorf("context: %w", syscall.ECONNREFUSED),
want: true,
},
{
name: "OpError dial+ECONNREFUSED returns true (primary branch)",
err: &net.OpError{
Op: "dial",
Net: "tcp",
Err: syscall.ECONNREFUSED,
},
want: true,
},
{
name: "OpError read+ECONNREFUSED returns true (primary branch)",
err: &net.OpError{
Op: "read",
Net: "tcp",
Err: syscall.ECONNREFUSED,
},
want: true,
},
{
name: "OpError dial+ETIMEDOUT returns false (primary branch, wrong errno)",
err: &net.OpError{
Op: "dial",
Net: "tcp",
Err: syscall.ETIMEDOUT,
},
want: false,
},
{
name: "OpError dial+plain error falls through to errors.Is and returns false (inner As miss)",
err: &net.OpError{
Op: "dial",
Net: "tcp",
Err: errors.New("not a syscall errno"),
},
want: false,
},
{
name: "OpError write+ECONNREFUSED returns true via fallback errors.Is",
err: &net.OpError{
Op: "write",
Net: "tcp",
Err: syscall.ECONNREFUSED,
},
want: true,
},
}
for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
if got := IsConnectionRefused(tc.err); got != tc.want {
t.Errorf("IsConnectionRefused(%v) = %v, want %v", tc.err, got, tc.want)
}
})
}
}
// fastRetryConfig keeps test runtime in the low-millisecond range.
func fastRetryConfig() RetryConfig {
return RetryConfig{
MaxRetries: 3,
InitialBackoff: 1 * time.Millisecond,
MaxBackoff: 5 * time.Millisecond,
BackoffFactor: 2.0,
}
}
func TestRetryWithBackoff_SuccessOnFirstAttempt(t *testing.T) {
calls := 0
err := RetryWithBackoff(context.Background(), fastRetryConfig(), func() error {
calls++
return nil
})
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if calls != 1 {
t.Errorf("want 1 call, got %d", calls)
}
}
func TestRetryWithBackoff_NonRetryableErrorReturnedImmediately(t *testing.T) {
sentinel := errors.New("boom")
calls := 0
err := RetryWithBackoff(context.Background(), fastRetryConfig(), func() error {
calls++
return sentinel
})
if !errors.Is(err, sentinel) {
t.Errorf("want wrapped sentinel, got %v", err)
}
if calls != 1 {
t.Errorf("want 1 call (no retry on non-conn-refused), got %d", calls)
}
}
func TestRetryWithBackoff_SuccessAfterTransientRefusals(t *testing.T) {
cfg := fastRetryConfig()
cfg.MaxRetries = 5 // allow up to 6 attempts
calls := 0
err := RetryWithBackoff(context.Background(), cfg, func() error {
calls++
if calls < 3 {
return syscall.ECONNREFUSED
}
return nil
})
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if calls != 3 {
t.Errorf("want 3 calls (2 failures + 1 success), got %d", calls)
}
}
func TestRetryWithBackoff_MaxRetriesExceededReturnsWrappedError(t *testing.T) {
cfg := fastRetryConfig()
cfg.MaxRetries = 2 // 3 total attempts (attempt 0, 1, 2)
calls := 0
err := RetryWithBackoff(context.Background(), cfg, func() error {
calls++
return syscall.ECONNREFUSED
})
if err == nil {
t.Fatal("expected error after exhausting retries, got nil")
}
if !errors.Is(err, syscall.ECONNREFUSED) {
t.Errorf("expected wrapped ECONNREFUSED, got %v", err)
}
// The loop runs attempt = 0..MaxRetries inclusive.
if calls != cfg.MaxRetries+1 {
t.Errorf("want %d calls, got %d", cfg.MaxRetries+1, calls)
}
// Error message includes the retry count for operator diagnostics.
if !containsAll(err.Error(), "max retries", "2") {
t.Errorf("error message missing retry count: %q", err.Error())
}
}
func TestRetryWithBackoff_ZeroMaxRetriesReturnsImmediately(t *testing.T) {
cfg := RetryConfig{
MaxRetries: 0,
InitialBackoff: 1 * time.Second, // large on purpose; must not sleep
MaxBackoff: 5 * time.Second,
BackoffFactor: 2.0,
}
calls := 0
start := time.Now()
err := RetryWithBackoff(context.Background(), cfg, func() error {
calls++
return syscall.ECONNREFUSED
})
elapsed := time.Since(start)
if err == nil {
t.Fatal("expected error, got nil")
}
if !errors.Is(err, syscall.ECONNREFUSED) {
t.Errorf("expected wrapped ECONNREFUSED, got %v", err)
}
if calls != 1 {
t.Errorf("want 1 call (no retry budget), got %d", calls)
}
// The only sleep would be after the attempt, but attempt == MaxRetries is
// short-circuited before the sleep select. So total runtime must be well
// under InitialBackoff.
if elapsed >= 500*time.Millisecond {
t.Errorf("no-retry path should not have slept; elapsed %v", elapsed)
}
}
func TestRetryWithBackoff_ContextCancelledDuringBackoff(t *testing.T) {
// Use a slow backoff so cancellation is guaranteed to land during the sleep.
cfg := RetryConfig{
MaxRetries: 5,
InitialBackoff: 50 * time.Millisecond,
MaxBackoff: 1 * time.Second,
BackoffFactor: 2.0,
}
ctx, cancel := context.WithCancel(context.Background())
// Cancel shortly after the first failed attempt starts its backoff.
go func() {
time.Sleep(10 * time.Millisecond)
cancel()
}()
calls := 0
err := RetryWithBackoff(ctx, cfg, func() error {
calls++
return syscall.ECONNREFUSED
})
if err == nil {
t.Fatal("expected error from cancelled context, got nil")
}
if !errors.Is(err, context.Canceled) {
t.Errorf("expected wrapped context.Canceled, got %v", err)
}
if calls < 1 {
t.Errorf("expected at least 1 call before cancellation, got %d", calls)
}
}
func TestRetryWithBackoff_WaitsBetweenAttempts(t *testing.T) {
// Lower-bound timing check — with InitialBackoff=20ms and BackoffFactor=2,
// three failed attempts sleep ~20ms + ~40ms = ~60ms before giving up.
// Assert >= 50ms to absorb scheduler jitter.
cfg := RetryConfig{
MaxRetries: 2,
InitialBackoff: 20 * time.Millisecond,
MaxBackoff: 100 * time.Millisecond,
BackoffFactor: 2.0,
}
start := time.Now()
_ = RetryWithBackoff(context.Background(), cfg, func() error {
return syscall.ECONNREFUSED
})
elapsed := time.Since(start)
if elapsed < 50*time.Millisecond {
t.Errorf("expected at least ~60ms of backoff delay, got %v", elapsed)
}
}
// containsAll reports whether s contains every substring in subs.
func containsAll(s string, subs ...string) bool {
for _, sub := range subs {
if !strings.Contains(s, sub) {
return false
}
}
return true
}