@@ -80,30 +80,37 @@ func TestRateLimitFailsClosedOnStoreError(t *testing.T) {
8080 }
8181}
8282
83- func TestRateLimitDoesNotLogCanceledRequestsAsStoreErrors (t * testing.T ) {
83+ func TestRateLimitDoesNotLogEndedRequestsAsStoreErrors (t * testing.T ) {
8484 t .Parallel ()
8585
86- var errorLogs atomic.Int64
87- logger := slog .New (countingErrorHandler {count : & errorLogs })
88- store := incrementFunc (func (ctx context.Context , _ string , _ time.Duration ) (int64 , error ) {
89- return 0 , ctx .Err ()
90- })
91- handler := RateLimit (store , config.RateLimit {Enabled : true , Requests : 2 , Window : time .Minute }, logger )(http .HandlerFunc (func (w http.ResponseWriter , _ * http.Request ) {
92- w .WriteHeader (http .StatusNoContent )
93- }))
94- ctx , cancel := context .WithCancel (t .Context ())
95- cancel ()
96- req , err := http .NewRequestWithContext (ctx , http .MethodPost , "http://example.test/mcp" , nil )
97- if err != nil {
98- t .Fatal (err )
86+ tests := []struct {
87+ name string
88+ err error
89+ }{
90+ {name : "canceled" , err : context .Canceled },
91+ {name : "deadline_exceeded" , err : context .DeadlineExceeded },
9992 }
100- req .RemoteAddr = "203.0.113.10:1234"
101-
102- rec := httptest .NewRecorder ()
103- handler .ServeHTTP (rec , req )
10493
105- if got := errorLogs .Load (); got != 0 {
106- t .Fatalf ("error logs = %d, want 0" , got )
94+ for _ , tt := range tests {
95+ t .Run (tt .name , func (t * testing.T ) {
96+ t .Parallel ()
97+
98+ var errorLogs atomic.Int64
99+ logger := slog .New (countingErrorHandler {count : & errorLogs })
100+ store := incrementFunc (func (context.Context , string , time.Duration ) (int64 , error ) {
101+ return 0 , tt .err
102+ })
103+ handler := RateLimit (store , config.RateLimit {Enabled : true , Requests : 2 , Window : time .Minute }, logger )(http .HandlerFunc (func (w http.ResponseWriter , _ * http.Request ) {
104+ w .WriteHeader (http .StatusNoContent )
105+ }))
106+
107+ rec := httptest .NewRecorder ()
108+ handler .ServeHTTP (rec , requestWithIP (t , "203.0.113.10:1234" ))
109+
110+ if got := errorLogs .Load (); got != 0 {
111+ t .Fatalf ("error logs = %d, want 0" , got )
112+ }
113+ })
107114 }
108115}
109116
0 commit comments