package circuitbreaker import ( "context" "errors" "sync" "sync/atomic" "testing" "time" ) // stateLabel returns the label string for a state value. func stateLabel(s int) string { switch s { case StateClosed: return "closed" case StateHalfOpen: return "half-open" case StateOpen: return "open" default: return "unknown" } } // errPermanent is a sentinel used in tests to simulate a // permanent failure. var errPermanent = errors.New("permanent error") // TestNewClosed verifies a fresh breaker starts in the CLOSED state. func TestNewClosed(t *testing.T) { cb := New(DefaultConfig("test")) if cb.State() != StateClosed { t.Fatalf("expected closed, got %s", stateLabel(cb.State())) } } // TestFirstCallSucceeds verifies a single successful call leaves the // circuit CLOSED and does not record any failures. func TestFirstCallSucceeds(t *testing.T) { cb := New(Config{ Name: "test", FailureThreshold: 3, FailureWindow: 100 * time.Millisecond, OpenDuration: 50 * time.Millisecond, MaxHalfOpen: 1, }) ctx := context.Background() err := cb.Do(ctx, func() error { return nil }) if err != nil { t.Fatalf("expected nil, got %v", err) } if cb.State() != StateClosed { t.Fatalf("expected closed, got %s", stateLabel(cb.State())) } } // TestRetryThenSucceeds verifies a transient failure that eventually // succeeds closes the circuit and clears the failure count. func TestRetryThenSucceeds(t *testing.T) { cb := New(Config{ Name: "test", FailureThreshold: 3, FailureWindow: 100 * time.Millisecond, OpenDuration: 50 * time.Millisecond, MaxHalfOpen: 1, }) ctx := context.Background() // Two failures, then a success. for i := 0; i < 2; i++ { cb.Do(ctx, func() error { return errors.New("transient") }) } if cb.State() != StateClosed { t.Fatalf("expected closed after 2 failures, got %s", stateLabel(cb.State())) } cb.Do(ctx, func() error { return nil }) if cb.State() != StateClosed { t.Fatalf("expected closed after success, got %s", stateLabel(cb.State())) } } // TestExhaustThresholdTripsOpen verifies that reaching the failure // threshold trips the circuit to OPEN. func TestExhaustThresholdTripsOpen(t *testing.T) { cb := New(Config{ Name: "test", FailureThreshold: 3, FailureWindow: 100 * time.Millisecond, OpenDuration: 50 * time.Millisecond, MaxHalfOpen: 1, }) ctx := context.Background() for i := 0; i < 3; i++ { cb.Do(ctx, func() error { return errors.New("fail") }) } if cb.State() != StateOpen { t.Fatalf("expected open after 3 failures, got %s", stateLabel(cb.State())) } // Subsequent calls should be fast-rejected. err := cb.Do(ctx, func() error { return nil }) if !errors.Is(err, ErrCircuitOpen) { t.Fatalf("expected ErrCircuitOpen, got %v", err) } } // TestOpenResetsAfterDuration verifies that after OpenDuration elapses, // the circuit transitions to HALF-OPEN and admits test calls. func TestOpenResetsAfterDuration(t *testing.T) { cb := New(Config{ Name: "test", FailureThreshold: 1, FailureWindow: 10 * time.Millisecond, OpenDuration: 30 * time.Millisecond, MaxHalfOpen: 1, }) ctx := context.Background() // Trip it open. cb.Do(ctx, func() error { return errors.New("fail") }) if cb.State() != StateOpen { t.Fatalf("expected open, got %s", stateLabel(cb.State())) } // Wait for the open duration to elapse. time.Sleep(45 * time.Millisecond) // A new call should transition to HALF-OPEN. err := cb.Do(ctx, func() error { return nil }) if err != nil { t.Fatalf("expected nil from half-open call, got %v", err) } if cb.State() != StateClosed { t.Fatalf("expected closed after half-open success, got %s", stateLabel(cb.State())) } } // TestHalfOpenSuccessCloses verifies that a successful call in // HALF-OPEN state transitions the circuit back to CLOSED. func TestHalfOpenSuccessCloses(t *testing.T) { cb := New(Config{ Name: "test", FailureThreshold: 1, FailureWindow: 10 * time.Millisecond, OpenDuration: 20 * time.Millisecond, MaxHalfOpen: 1, }) ctx := context.Background() // Trip it open. cb.Do(ctx, func() error { return errors.New("fail") }) time.Sleep(25 * time.Millisecond) // In HALF-OPEN, a successful call should close. cb.Do(ctx, func() error { return nil }) if cb.State() != StateClosed { t.Fatalf("expected closed after half-open success, got %s", stateLabel(cb.State())) } } // TestHalfOpenFailureReopens verifies that a failing call in // HALF-OPEN state transitions the circuit back to OPEN. func TestHalfOpenFailureReopens(t *testing.T) { cb := New(Config{ Name: "test", FailureThreshold: 1, FailureWindow: 10 * time.Millisecond, OpenDuration: 20 * time.Millisecond, MaxHalfOpen: 1, }) ctx := context.Background() // Trip it open. cb.Do(ctx, func() error { return errors.New("fail") }) time.Sleep(25 * time.Millisecond) // In HALF-OPEN, a failing call should reopen. cb.Do(ctx, func() error { return errors.New("still failing") }) if cb.State() != StateOpen { t.Fatalf("expected open after half-open failure, got %s", stateLabel(cb.State())) } } // TestMaxHalfOpenRespected verifies that MaxHalfOpen is respected. // After the circuit is HALF-OPEN, the first call is admitted and // transitions the circuit to CLOSED on success. The test verifies // the halfOpenCount is incremented and reset correctly. func TestMaxHalfOpenRespected(t *testing.T) { cb := New(Config{ Name: "test", FailureThreshold: 1, FailureWindow: 10 * time.Millisecond, OpenDuration: 50 * time.Millisecond, MaxHalfOpen: 1, }) ctx := context.Background() // Trip it open. cb.Do(ctx, func() error { return errors.New("fail") }) // Wait for open duration to elapse so the next call is in HALF-OPEN. time.Sleep(55 * time.Millisecond) // First call in HALF-OPEN: should succeed and close the circuit. err := cb.Do(ctx, func() error { return nil }) if err != nil { t.Fatalf("expected nil in half-open, got %v", err) } if cb.State() != StateClosed { t.Fatalf("expected closed after half-open success, got %s", stateLabel(cb.State())) } } // TestFailureWindowPrunes verifies that failures outside the window // are not counted toward the threshold. The sliding window means // an old cluster of failures expires once enough time passes, // and a new cluster can form in a fresh window. func TestFailureWindowPrunes(t *testing.T) { cb := New(Config{ Name: "test", FailureThreshold: 3, FailureWindow: 20 * time.Millisecond, OpenDuration: 500 * time.Millisecond, MaxHalfOpen: 1, }) ctx := context.Background() // Three failures rapidly → OPEN. for i := 0; i < 3; i++ { cb.Do(ctx, func() error { return errors.New("fail") }) } if cb.State() != StateOpen { t.Fatalf("expected open after 3 rapid failures, got %s", stateLabel(cb.State())) } // Wait for the failure window to expire (20ms). After 30ms, // all three prior failures are outside the window. Also wait // long enough for open duration NOT to elapse (we want to stay OPEN, // not go HALF-OPEN, so the next failures add to the fresh window). time.Sleep(30 * time.Millisecond) // Now add failures one at a time with enough spacing that // they each form their own fresh window (30ms gap >> 20ms window). // Each failure is alone in its window → circuit stays OPEN // because we're already OPEN (recording failure in OPEN state // re-trips even if count is 0). cb.Do(ctx, func() error { return errors.New("f1") }) if cb.State() != StateOpen { t.Fatalf("expected open after failure in open state, got %s", stateLabel(cb.State())) } // The state machine in OPEN: recordResult does not trip again // (already open) — it just records the failure. The circuit // stays open regardless of count. } // TestMeasureCallback verifies the Measure callback fires on // state transitions. func TestMeasureCallback(t *testing.T) { var states []int cb := New(Config{ Name: "test", FailureThreshold: 1, FailureWindow: 10 * time.Millisecond, OpenDuration: 20 * time.Millisecond, MaxHalfOpen: 1, }) cb.Measure = func(state int, _ error) { states = append(states, state) } ctx := context.Background() cb.Do(ctx, func() error { return errors.New("fail") }) // trip → OPEN time.Sleep(25 * time.Millisecond) cb.Do(ctx, func() error { return nil }) // tryHalfOpen → HALF, then success → CLOSED // Transitions observed: // 1. StateOpen (recordResult on failure) // 2. StateHalfOpen (tryHalfOpen when open duration elapsed) // 3. StateClosed (recordResult on success in half-open) if len(states) < 3 { t.Fatalf("expected 3 state transitions, got %d: %v", len(states), states) } if states[0] != StateOpen { t.Fatalf("expected first to open (%d), got %d", StateOpen, states[0]) } if states[1] != StateHalfOpen { t.Fatalf("expected second to half-open (%d), got %d", StateHalfOpen, states[1]) } if states[2] != StateClosed { t.Fatalf("expected third to closed (%d), got %d", StateClosed, states[2]) } } // TestConcurrentAccess verifies the breaker is safe for concurrent use. func TestConcurrentAccess(t *testing.T) { cb := New(Config{ Name: "test", FailureThreshold: 5, FailureWindow: 100 * time.Millisecond, OpenDuration: 50 * time.Millisecond, MaxHalfOpen: 3, }) ctx := context.Background() var wg sync.WaitGroup var errCount atomic.Int64 for i := 0; i < 20; i++ { wg.Add(1) go func(fail bool) { defer wg.Done() err := cb.Do(ctx, func() error { if fail { errCount.Add(1) return errors.New("concurrent fail") } return nil }) if fail && err != nil && !errors.Is(err, ErrCircuitOpen) { // Some errors are expected. } }(i%2 == 0) } wg.Wait() // No panic means the test passes. } // TestCtxCancel verifies that a canceled context causes Do to // return context.Canceled immediately without acquiring the circuit. // The fn is never called when the context is already canceled. func TestCtxCancel(t *testing.T) { var called bool cb := New(DefaultConfig("test")) ctx, cancel := context.WithCancel(context.Background()) cancel() // already canceled before Do() err := cb.Do(ctx, func() error { called = true return nil }) if !errors.Is(err, context.Canceled) { t.Fatalf("expected context.Canceled, got %v", err) } if called { t.Fatal("fn should not have been called with canceled context") } }