diff --git a/cbreaker/cbreaker_test.go b/cbreaker/cbreaker_test.go index a9cea71..41d53b2 100644 --- a/cbreaker/cbreaker_test.go +++ b/cbreaker/cbreaker_test.go @@ -304,6 +304,42 @@ func TestCircuitBreaker_sideEffects(t *testing.T) { } } +func TestCircuitBreaker_requestThreshold(t *testing.T) { + handler := http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.WriteHeader(http.StatusGatewayTimeout) + }) + + testutils.FreezeTime(t) + + cb, err := New(handler, triggerNetRatio+" && RequestThreshold() >= 20") + require.NoError(t, err) + + srv := httptest.NewServer(cb) + t.Cleanup(srv.Close) + + cb.metrics = statsResponseCodes(statusCode{Code: http.StatusGatewayTimeout, Count: 18}) + + clock.Advance(defaultCheckPeriod + clock.Millisecond) + + re, _, err := testutils.Get(srv.URL) + require.NoError(t, err) + assert.Equal(t, http.StatusGatewayTimeout, re.StatusCode) + assert.Equal(t, cbState(stateStandby), cb.state) + + cb.metrics = statsResponseCodes(statusCode{Code: http.StatusGatewayTimeout, Count: 19}) + + clock.Advance(defaultCheckPeriod + clock.Millisecond) + + re, _, err = testutils.Get(srv.URL) + require.NoError(t, err) + assert.Equal(t, http.StatusGatewayTimeout, re.StatusCode) + assert.Equal(t, cbState(stateTripped), cb.state) + + re, _, err = testutils.Get(srv.URL) + require.NoError(t, err) + assert.Equal(t, http.StatusServiceUnavailable, re.StatusCode) +} + func statsOK() *memmetrics.RTMetrics { m, err := memmetrics.NewRTMetrics() if err != nil { diff --git a/cbreaker/predicates.go b/cbreaker/predicates.go index dd368fb..1bbd55d 100644 --- a/cbreaker/predicates.go +++ b/cbreaker/predicates.go @@ -25,6 +25,7 @@ func parseExpression(in string) (hpredicate, error) { "LatencyAtQuantileMS": latencyAtQuantile, "NetworkErrorRatio": networkErrorRatio, "ResponseCodeRatio": responseCodeRatio, + "RequestThreshold": requestThreshold, }, }) if err != nil { diff --git a/cbreaker/predicates_functions.go b/cbreaker/predicates_functions.go index 3dfa444..f55945f 100644 --- a/cbreaker/predicates_functions.go +++ b/cbreaker/predicates_functions.go @@ -4,7 +4,7 @@ import ( "github.com/vulcand/oxy/v2/internal/holsterv4/clock" ) -type toType[T int | float64] func(c *CircuitBreaker) T +type toType[T int | int64 | float64] func(c *CircuitBreaker) T func latencyAtQuantile(quantile float64) toType[int] { return func(c *CircuitBreaker) int { @@ -29,3 +29,9 @@ func responseCodeRatio(startA, endA, startB, endB int) toType[float64] { return c.metrics.ResponseCodeRatio(startA, endA, startB, endB) } } + +func requestThreshold() toType[int64] { + return func(c *CircuitBreaker) int64 { + return c.metrics.TotalCount() + } +} diff --git a/cbreaker/predicates_operators.go b/cbreaker/predicates_operators.go index d783ec9..8f4a062 100644 --- a/cbreaker/predicates_operators.go +++ b/cbreaker/predicates_operators.go @@ -42,6 +42,8 @@ func eq(m any, value any) (hpredicate, error) { switch mapper := m.(type) { case toType[int]: return genericEQ(mapper, value) + case toType[int64]: + return int64EQ(mapper, value) case toType[float64]: return genericEQ(mapper, value) } @@ -64,6 +66,8 @@ func lt(m any, value any) (hpredicate, error) { switch mapper := m.(type) { case toType[int]: return genericLT(mapper, value) + case toType[int64]: + return int64LT(mapper, value) case toType[float64]: return genericLT(mapper, value) } @@ -93,6 +97,8 @@ func gt(m any, value any) (hpredicate, error) { switch mapper := m.(type) { case toType[int]: return genericGT(mapper, value) + case toType[int64]: + return int64GT(mapper, value) case toType[float64]: return genericGT(mapper, value) } @@ -149,3 +155,42 @@ func genericGT[T int | float64](m toType[T], val any) (hpredicate, error) { return m(c) > value }, nil } + +func int64EQ(m toType[int64], val any) (hpredicate, error) { + // Use `int` instead of `int64` because `vulcand/predicate` only use `int`. + // Note: using `int64` should be the default type of any integer, but changing this will break `vulcand/predicate`. + value, ok := val.(int) + if !ok { + return nil, fmt.Errorf("expected int, got %T", val) + } + + return func(c *CircuitBreaker) bool { + return m(c) == int64(value) + }, nil +} + +func int64LT(m toType[int64], val any) (hpredicate, error) { + // Use `int` instead of `int64` because `vulcand/predicate` only use `int`. + // Note: using `int64` should be the default type of any integer, but changing this will break `vulcand/predicate`. + value, ok := val.(int) + if !ok { + return nil, fmt.Errorf("expected int, got %T", val) + } + + return func(c *CircuitBreaker) bool { + return m(c) < int64(value) + }, nil +} + +func int64GT(m toType[int64], val any) (hpredicate, error) { + // Use `int` instead of `int64` because `vulcand/predicate` only use `int`. + // Note: using `int64` should be the default type of any integer, but changing this will break `vulcand/predicate`. + value, ok := val.(int) + if !ok { + return nil, fmt.Errorf("expected int, got %T", val) + } + + return func(c *CircuitBreaker) bool { + return m(c) > int64(value) + }, nil +}