diff --git a/server/apiv1/gate.go b/server/apiv1/gate.go index ee25b15a2..35e7841dd 100644 --- a/server/apiv1/gate.go +++ b/server/apiv1/gate.go @@ -93,6 +93,8 @@ func newGate(doc *openapi3.T, mux chi.Routes, auth authenticator, rules gateRule } if conf.Server.AuthRequestLimit > 0 { g.limiter = server.ClientIPRateLimiter(conf.Server.AuthRequestLimit, conf.Server.AuthWindowLength, + // Counts are per node, so X-RateLimit-Remaining would mislead clients of a scaled-out server. + httprate.WithResponseHeaders(httprate.ResponseHeaders{RetryAfter: "Retry-After"}), httprate.WithLimitHandler(func(w http.ResponseWriter, r *http.Request) { writeProblemStatus(w, r, http.StatusTooManyRequests, ProblemCodeRateLimited, "too many requests") })) diff --git a/server/apiv1/gate_test.go b/server/apiv1/gate_test.go index 37cc67448..a2f6fc0a2 100644 --- a/server/apiv1/gate_test.go +++ b/server/apiv1/gate_test.go @@ -370,10 +370,13 @@ var _ = Describe("spec gate", func() { var err error mux, err = build(gateSpec) Expect(err).ToNot(HaveOccurred()) - Expect(do(http.MethodPost, "/limited", "", "").Code).To(Equal(http.StatusOK)) w := do(http.MethodPost, "/limited", "", "") + Expect(w.Code).To(Equal(http.StatusOK)) + expectNoXRateLimitHeaders(w) + w = do(http.MethodPost, "/limited", "", "") Expect(w.Code).To(Equal(http.StatusTooManyRequests)) Expect(w.Header().Get("Retry-After")).ToNot(BeEmpty()) + expectNoXRateLimitHeaders(w) Expect(decodeProblem(w).Code).To(Equal(ProblemCodeRateLimited)) }) @@ -436,3 +439,10 @@ func captureLogs() *bytes.Buffer { }) return buf } + +func expectNoXRateLimitHeaders(w *httptest.ResponseRecorder) { + GinkgoHelper() + for _, h := range []string{"X-RateLimit-Limit", "X-RateLimit-Remaining", "X-RateLimit-Increment", "X-RateLimit-Reset"} { + Expect(w.Header().Values(h)).To(BeEmpty(), h) + } +} diff --git a/server/middlewares_test.go b/server/middlewares_test.go index 58ffa2e37..55f842914 100644 --- a/server/middlewares_test.go +++ b/server/middlewares_test.go @@ -549,6 +549,12 @@ var _ = Describe("middlewares", func() { Entry("True-Client-IP", "True-Client-IP"), ) + It("sends the X-RateLimit headers by default", func() { + w := httptest.NewRecorder() + handler.ServeHTTP(w, httptest.NewRequestWithContext(GinkgoT().Context(), "POST", "/auth/login", nil)) + Expect(w.Header().Get("X-RateLimit-Limit")).To(Equal("2")) + }) + Context("behind a trusted proxy", func() { BeforeEach(func() { conf.Server.ExtAuth.TrustedSources = "10.0.0.0/8"