diff --git a/backend/internal/config/config_test.go b/backend/internal/config/config_test.go new file mode 100644 index 0000000..dfa3b35 --- /dev/null +++ b/backend/internal/config/config_test.go @@ -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) + } + }) + } +} diff --git a/backend/pkg/logger/logger_test.go b/backend/pkg/logger/logger_test.go new file mode 100644 index 0000000..aea314b --- /dev/null +++ b/backend/pkg/logger/logger_test.go @@ -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 +} diff --git a/backend/pkg/utils/cache_test.go b/backend/pkg/utils/cache_test.go new file mode 100644 index 0000000..28eadb7 --- /dev/null +++ b/backend/pkg/utils/cache_test.go @@ -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) + } +} diff --git a/backend/pkg/utils/retry_test.go b/backend/pkg/utils/retry_test.go new file mode 100644 index 0000000..7166b09 --- /dev/null +++ b/backend/pkg/utils/retry_test.go @@ -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 +}