Skip to content

Commit 6fb573e

Browse files
fix(ratelimit): downgrade redis deadline logs (#26)
1 parent bed5eb9 commit 6fb573e

2 files changed

Lines changed: 106 additions & 4 deletions

File tree

internal/middleware/ratelimit.go

Lines changed: 18 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -10,6 +10,7 @@ import (
1010
"net"
1111
"net/http"
1212
"net/netip"
13+
"os"
1314
"strconv"
1415
"strings"
1516
"time"
@@ -41,12 +42,12 @@ func RateLimit(store kv.Store, cfg config.RateLimit, logger *slog.Logger) Middle
4142
key := rateLimitKey(ip, cfg.Window, now)
4243
count, err := store.Increment(r.Context(), key, cfg.Window+time.Second)
4344
if err != nil {
44-
if outcome, ok := rateLimitContextOutcome(err); ok {
45+
if outcome, ok := rateLimitContextOutcome(r.Context(), err); ok {
4546
trace.SpanFromContext(r.Context()).SetAttributes(observability.RateLimitAttrs{Outcome: outcome, Limit: cfg.Requests, Window: cfg.Window}.Attributes()...)
4647
return
4748
}
4849
trace.SpanFromContext(r.Context()).SetAttributes(observability.RateLimitAttrs{Outcome: observability.RateLimitOutcomeStoreError, Limit: cfg.Requests, Window: cfg.Window}.Attributes()...)
49-
logger.ErrorContext(r.Context(), "rate limit check failed", slog.Any("error", err), slog.String("client_ip", ip))
50+
logger.LogAttrs(r.Context(), rateLimitStoreErrorLevel(err), "rate limit check failed", slog.Any("error", err), slog.String("client_ip", ip))
5051
http.Error(w, http.StatusText(http.StatusServiceUnavailable), http.StatusServiceUnavailable)
5152
return
5253
}
@@ -66,7 +67,14 @@ func RateLimit(store kv.Store, cfg config.RateLimit, logger *slog.Logger) Middle
6667
}
6768
}
6869

69-
func rateLimitContextOutcome(err error) (observability.RateLimitOutcome, bool) {
70+
func rateLimitContextOutcome(ctx context.Context, err error) (observability.RateLimitOutcome, bool) {
71+
switch ctx.Err() {
72+
case context.Canceled:
73+
return observability.RateLimitOutcomeCanceled, true
74+
case context.DeadlineExceeded:
75+
return observability.RateLimitOutcomeDeadline, true
76+
}
77+
7078
switch {
7179
case errors.Is(err, context.Canceled):
7280
return observability.RateLimitOutcomeCanceled, true
@@ -77,6 +85,13 @@ func rateLimitContextOutcome(err error) (observability.RateLimitOutcome, bool) {
7785
}
7886
}
7987

88+
func rateLimitStoreErrorLevel(err error) slog.Level {
89+
if errors.Is(err, os.ErrDeadlineExceeded) {
90+
return slog.LevelWarn
91+
}
92+
return slog.LevelError
93+
}
94+
8095
func rateLimitKey(ip string, window time.Duration, now time.Time) string {
8196
sum := sha256.Sum256([]byte(ip))
8297
bucket := now.UnixNano() / window.Nanoseconds()

internal/middleware/ratelimit_test.go

Lines changed: 88 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -6,6 +6,7 @@ import (
66
"log/slog"
77
"net/http"
88
"net/http/httptest"
9+
"os"
910
"strconv"
1011
"sync/atomic"
1112
"testing"
@@ -133,7 +134,7 @@ func TestRateLimitContextOutcome(t *testing.T) {
133134
t.Run(tt.name, func(t *testing.T) {
134135
t.Parallel()
135136

136-
got, ok := rateLimitContextOutcome(tt.err)
137+
got, ok := rateLimitContextOutcome(context.Background(), tt.err)
137138
if ok != tt.ok {
138139
t.Fatalf("ok = %t, want %t", ok, tt.ok)
139140
}
@@ -144,6 +145,63 @@ func TestRateLimitContextOutcome(t *testing.T) {
144145
}
145146
}
146147

148+
func TestRateLimitDoesNotLogCanceledRequestWhenRedisReturnsSocketTimeout(t *testing.T) {
149+
t.Parallel()
150+
151+
var errorLogs atomic.Int64
152+
logger := slog.New(countingErrorHandler{count: &errorLogs})
153+
store := incrementFunc(func(context.Context, string, time.Duration) (int64, error) {
154+
// os.ErrDeadlineExceeded is the public identity of the runtime's
155+
// poll.DeadlineExceededError, which formats as "i/o timeout".
156+
return 0, os.ErrDeadlineExceeded
157+
})
158+
handler := RateLimit(store, config.RateLimit{Enabled: true, Requests: 2, Window: time.Minute}, logger)(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
159+
w.WriteHeader(http.StatusNoContent)
160+
}))
161+
162+
ctx, cancel := context.WithCancel(t.Context())
163+
cancel()
164+
req, err := http.NewRequestWithContext(ctx, http.MethodPost, "http://example.test/mcp", nil)
165+
if err != nil {
166+
t.Fatal(err)
167+
}
168+
req.RemoteAddr = "203.0.113.10:1234"
169+
170+
rec := httptest.NewRecorder()
171+
handler.ServeHTTP(rec, req)
172+
173+
if got := errorLogs.Load(); got != 0 {
174+
t.Fatalf("error logs = %d, want 0", got)
175+
}
176+
}
177+
178+
func TestRateLimitLogsRedisSocketTimeoutAtWarnAndFailsClosed(t *testing.T) {
179+
t.Parallel()
180+
181+
var warningLogs atomic.Int64
182+
var errorLogs atomic.Int64
183+
logger := slog.New(recordingLogHandler{warnings: &warningLogs, errors: &errorLogs})
184+
store := incrementFunc(func(context.Context, string, time.Duration) (int64, error) {
185+
return 0, os.ErrDeadlineExceeded
186+
})
187+
handler := RateLimit(store, config.RateLimit{Enabled: true, Requests: 2, Window: time.Minute}, logger)(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
188+
w.WriteHeader(http.StatusNoContent)
189+
}))
190+
191+
rec := httptest.NewRecorder()
192+
handler.ServeHTTP(rec, requestWithIP(t, "203.0.113.10:1234"))
193+
194+
if rec.Code != http.StatusServiceUnavailable {
195+
t.Fatalf("status = %d, want %d", rec.Code, http.StatusServiceUnavailable)
196+
}
197+
if got := warningLogs.Load(); got != 1 {
198+
t.Fatalf("warning logs = %d, want 1", got)
199+
}
200+
if got := errorLogs.Load(); got != 0 {
201+
t.Fatalf("error logs = %d, want 0", got)
202+
}
203+
}
204+
147205
func TestClientIPPrefersFlyHeaderAndNormalizes(t *testing.T) {
148206
t.Parallel()
149207

@@ -205,3 +263,32 @@ func (h countingErrorHandler) WithAttrs([]slog.Attr) slog.Handler {
205263
func (h countingErrorHandler) WithGroup(string) slog.Handler {
206264
return h
207265
}
266+
267+
type recordingLogHandler struct {
268+
warnings *atomic.Int64
269+
errors *atomic.Int64
270+
}
271+
272+
var _ slog.Handler = recordingLogHandler{}
273+
274+
func (h recordingLogHandler) Enabled(_ context.Context, level slog.Level) bool {
275+
return level >= slog.LevelWarn
276+
}
277+
278+
func (h recordingLogHandler) Handle(_ context.Context, record slog.Record) error {
279+
switch {
280+
case record.Level >= slog.LevelError:
281+
h.errors.Add(1)
282+
case record.Level >= slog.LevelWarn:
283+
h.warnings.Add(1)
284+
}
285+
return nil
286+
}
287+
288+
func (h recordingLogHandler) WithAttrs([]slog.Attr) slog.Handler {
289+
return h
290+
}
291+
292+
func (h recordingLogHandler) WithGroup(string) slog.Handler {
293+
return h
294+
}

0 commit comments

Comments
 (0)