From 61832f4cc5fcb36c051b98f22f17c7420b55e93b Mon Sep 17 00:00:00 2001 From: JOY <5027251+JOY@users.noreply.github.com> Date: Sun, 13 Sep 2026 17:15:37 +0700 Subject: [PATCH] feat(security): rate limit the public endpoints an anonymous caller can drive There was no rate limiting anywhere in the application. Every unauthenticated endpoint could be called in a tight loop, which made three things cheap: flooding the message upload endpoints with 20 MB bodies, spraying login and support registration, and burning LLM spend through the conversation endpoints. Add a small in-process fixed-window limiter and attach it to six public routes: login (20 per window), support registration (10), the widget session exchange (120), message attachment and image upload (30, shared between the two), and documentation feedback (20). The window is 60 seconds by default, and both it and the on/off switch are configurable. The limits are sized for a human driving a browser, not for a machine. The session exchange budget in particular has to absorb a whole support office loading the widget from behind one NAT address, which is why it sits an order of magnitude above the others. Deliberately not limited: - /api/third/* channel webhooks. A platform that receives a 429 from its webhook stops retrying and eventually disables the delivery, which takes an entire channel offline - a far worse outcome than the flood the limit was meant to stop. Those endpoints authenticate by signature instead. - /api/ws/* websockets, and /api/webhooks/* which are HMAC verified. - /api/dashboard/* - already behind AuthMiddleware, and staff sharing one office address would throttle each other. Rejections return 429 with a Retry-After header and an ordinary JsonResult body. The status is a real 429 rather than the 200-with-error-code the auth middleware uses, because web/lib/api/client.ts parses the payload before it inspects response.ok and surfaces payload.message, so the localized text still reaches the user - and Retry-After is only meaningful on a 429 or 503. The message is error.e0354 in both backend locales and does not disclose the limit or the remaining budget, which would only tell a caller how much room they have left. Retry-After rounds the remaining window up rather than down. Telling a caller to come back sooner than the window actually resets just earns another 429. Rejections are not logged separately. requestLogMiddleware already records path, status and client address for every request, so a 429 is visible there without giving a flood a second way to fill the log. The limiter keys on ctx.ClientIP(), which is only trustworthy because of the trusted-proxy configuration that landed immediately before this. Without it a caller would pick their own bucket with an X-Forwarded-For header. Counters are per process, and expired buckets are swept lazily from inside Allow rather than by a goroutine, so there is no background lifetime to own and no unbounded growth. Running several replicas gives each its own budget, weakening the bound by the replica count; config.example.yaml says so explicitly. These limits exist to make flooding expensive, not to meter a quota. Tests: seven for the limiter, including an exact-count check across 16 goroutines making 3200 calls against one key; three for the middleware, covering the 429 shape, per-address isolation and a nil limiter allowing everything; and five end to end against a server built by NewServer, including one that fires 1000 requests at the health, config, org-sync and two channel webhook routes and fails if any of them returns 429. The concurrency test could not be run under -race: this repository builds with CGO disabled and go test -race requires cgo. It still has teeth, because an unguarded concurrent map write panics rather than merely miscounting. --- .env.example | 7 + config/config.example.yaml | 15 ++ internal/bootstrap/routes.go | 75 ++++++-- internal/bootstrap/server.go | 10 +- internal/bootstrap/server_ratelimit_test.go | 162 ++++++++++++++++++ internal/middleware/ratelimit_middleware.go | 56 ++++++ .../middleware/ratelimit_middleware_test.go | 103 +++++++++++ internal/pkg/config/config.go | 33 +++- internal/pkg/config/config_test.go | 51 ++++++ internal/pkg/i18nx/locales/en-US.yml | 1 + internal/pkg/i18nx/locales/zh-CN.yml | 1 + internal/pkg/ratelimit/ratelimit.go | 101 +++++++++++ internal/pkg/ratelimit/ratelimit_test.go | 132 ++++++++++++++ 13 files changed, 732 insertions(+), 15 deletions(-) create mode 100644 internal/bootstrap/server_ratelimit_test.go create mode 100644 internal/middleware/ratelimit_middleware.go create mode 100644 internal/middleware/ratelimit_middleware_test.go create mode 100644 internal/pkg/ratelimit/ratelimit.go create mode 100644 internal/pkg/ratelimit/ratelimit_test.go diff --git a/.env.example b/.env.example index 2bcd02ab..4e9d1af4 100644 --- a/.env.example +++ b/.env.example @@ -20,6 +20,13 @@ PORT=8083 # TRUSTED_PROXIES="10.0.0.0/8,172.16.0.0/12" TRUSTED_PLATFORM=cloudflare +# Rate limiting for the public, unauthenticated endpoints (login, support +# registration, widget session exchange, message uploads, doc feedback). Enabled +# by default; the counters are per process. Channel webhooks, websockets and the +# authenticated dashboard are deliberately exempt. +# RATE_LIMIT_ENABLED=true +# RATE_LIMIT_WINDOW_SECONDS=60 + # Database Configuration # Driver options: sqlite, mysql, postgres # Supabase PostgreSQL (DOS): diff --git a/config/config.example.yaml b/config/config.example.yaml index 29b76021..9a6722b4 100644 --- a/config/config.example.yaml +++ b/config/config.example.yaml @@ -17,6 +17,21 @@ server: # header name) when an edge overwrites rather than appends the real client # address. When set it takes precedence over X-Forwarded-For entirely. trustedPlatform: "" + rateLimit: + # Bounds how often one client address may call the public, unauthenticated + # endpoints: login, support registration, the widget session exchange, message + # uploads and documentation feedback. Enabled by default. + # + # Channel webhooks under /api/third, the websocket routes, /api/webhooks and + # the whole authenticated dashboard are deliberately NOT covered. A platform + # that receives a 429 from its webhook stops retrying and eventually disables + # the delivery, which takes an entire channel offline. + # + # The counters live in this process. Running several replicas gives each one + # its own budget, which weakens the bound by the replica count; these limits + # exist to make flooding expensive, not to meter a quota. + enabled: true + windowSeconds: 60 cors: # Browser CORS allowlist. In production, replace this with the actual frontend or embedded-site domains, such as https://support.example.com. # Leave it empty to reject cross-origin browser requests. Same-origin and non-browser calls are still supported. diff --git a/internal/bootstrap/routes.go b/internal/bootstrap/routes.go index c4d44903..47b5fd45 100644 --- a/internal/bootstrap/routes.go +++ b/internal/bootstrap/routes.go @@ -1,15 +1,66 @@ package bootstrap import ( + "time" + "agent-desk/internal/handlers/api" "agent-desk/internal/handlers/dashboard" "agent-desk/internal/handlers/third" + "agent-desk/internal/middleware" + "agent-desk/internal/pkg/config" + "agent-desk/internal/pkg/ratelimit" "github.com/gin-gonic/gin" ) -func registerApiAuthRoutes(group *gin.RouterGroup) { - group.POST("/login", api.Login) +// publicRateLimits holds one limiter per abuse-prone unauthenticated endpoint. +// They are built once per server rather than per request, and the middleware keys +// them by client address. +// +// Channel webhooks under /api/third, the websocket routes and /api/webhooks are +// deliberately absent. A platform that receives a 429 from its webhook stops +// retrying and eventually disables the delivery, which would take a whole channel +// offline; those endpoints authenticate by signature instead. The dashboard is +// absent too - it is already behind AuthMiddleware, and staff sharing one office +// address would throttle each other. +type publicRateLimits struct { + login *ratelimit.Limiter + supportRegister *ratelimit.Limiter + sessionExchange *ratelimit.Limiter + upload *ratelimit.Limiter + docFeedback *ratelimit.Limiter +} + +// Limits are sized for a human driving a browser, not for a machine. They are +// deliberately generous: the point is to make flooding expensive, not to police +// legitimate use. The session exchange budget in particular has to absorb a whole +// support office loading the widget from behind one NAT address. +const ( + limitLogin = 20 + limitSupportRegister = 10 + limitSessionExchange = 120 + limitUpload = 30 + limitDocFeedback = 20 +) + +func newPublicRateLimits(cfg config.RateLimitConfig) publicRateLimits { + if !cfg.IsEnabled() { + // Every field stays nil and a nil limiter allows everything, so "disabled" + // needs no branch at any call site. + return publicRateLimits{} + } + window := time.Duration(cfg.WindowSecondsOrDefault()) * time.Second + return publicRateLimits{ + login: ratelimit.New(limitLogin, window), + supportRegister: ratelimit.New(limitSupportRegister, window), + sessionExchange: ratelimit.New(limitSessionExchange, window), + upload: ratelimit.New(limitUpload, window), + docFeedback: ratelimit.New(limitDocFeedback, window), + } +} + +func registerApiAuthRoutes(group *gin.RouterGroup, limits publicRateLimits) { + group.POST("/login", middleware.RateLimit(limits.login), api.Login) group.POST("/logout", api.Logout) group.GET("/profile", api.Profile) group.POST("/profile/update", api.UpdateProfile) @@ -37,8 +88,8 @@ func registerApiWebhookRoutes(group *gin.RouterGroup) { group.POST("/ecosystem", api.OrgSyncWebhook) } -func registerApiCustomerRoutes(group *gin.RouterGroup) { - group.POST("/session_exchange", api.CustomerPostSession_exchange) +func registerApiCustomerRoutes(group *gin.RouterGroup, limits publicRateLimits) { + group.POST("/session_exchange", middleware.RateLimit(limits.sessionExchange), api.CustomerPostSession_exchange) } func registerApiConversationRoutes(group *gin.RouterGroup) { @@ -47,22 +98,26 @@ func registerApiConversationRoutes(group *gin.RouterGroup) { group.POST("/create_or_match", api.ConversationPostCreate_or_match) } -func registerApiMessageRoutes(group *gin.RouterGroup) { +func registerApiMessageRoutes(group *gin.RouterGroup, limits publicRateLimits) { group.Any("/list", api.MessageAnyList) group.POST("/read", api.MessagePostRead) group.POST("/send", api.MessagePostSend) - group.POST("/upload_attachment", api.MessagePostUpload_attachment) - group.POST("/upload_image", api.MessagePostUpload_image) + // One shared budget for both upload routes: what matters is how many bytes an + // unauthenticated caller can push at the storage layer, not which of the two + // endpoints they used. + uploadLimit := middleware.RateLimit(limits.upload) + group.POST("/upload_attachment", uploadLimit, api.MessagePostUpload_attachment) + group.POST("/upload_image", uploadLimit, api.MessagePostUpload_image) } -func registerApiSupportRoutes(group *gin.RouterGroup) { +func registerApiSupportRoutes(group *gin.RouterGroup, limits publicRateLimits) { group.GET("/config", api.SupportConfigGetConfig) - group.POST("/auth/register", api.SupportAuthPostRegister) + group.POST("/auth/register", middleware.RateLimit(limits.supportRegister), api.SupportAuthPostRegister) group.GET("/me", api.SupportGetMe) group.Any("/doc-page/list", api.DocPageAnyList) group.GET("/doc-page/navigation", api.DocPageGetNavigation) group.GET("/doc-page/:id", api.DocPageGetBy) - group.POST("/doc-page/feedback", api.DocPagePostFeedback) + group.POST("/doc-page/feedback", middleware.RateLimit(limits.docFeedback), api.DocPagePostFeedback) group.Any("/community/categories/list", api.CategoryAnyList) group.Any("/community/posts/list", api.PostAnyList) group.GET("/community/posts/:id", api.PostGetBy) diff --git a/internal/bootstrap/server.go b/internal/bootstrap/server.go index f750266a..7c1442d5 100644 --- a/internal/bootstrap/server.go +++ b/internal/bootstrap/server.go @@ -171,17 +171,19 @@ func isWebsocketUpgrade(ctx *gin.Context) bool { } func addRouter(app *gin.Engine) { + limits := newPublicRateLimits(config.Current().Server.RateLimit) + app.Any("/api/mcp", gin.WrapH(mcps.NewHTTPHandler())) apiGroup := app.Group("/api") apiGroup.GET("/health", api.Health) apiGroup.GET("/config", api.PublicConfig) - registerApiAuthRoutes(apiGroup.Group("/auth")) + registerApiAuthRoutes(apiGroup.Group("/auth"), limits) registerApiChannelRoutes(apiGroup.Group("/channel")) - registerApiCustomerRoutes(apiGroup.Group("/customer")) + registerApiCustomerRoutes(apiGroup.Group("/customer"), limits) registerApiConversationRoutes(apiGroup.Group("/conversation", middleware.ExternalUserMiddleware)) - registerApiMessageRoutes(apiGroup.Group("/message", middleware.ExternalUserMiddleware)) - registerApiSupportRoutes(apiGroup.Group("/support")) + registerApiMessageRoutes(apiGroup.Group("/message", middleware.ExternalUserMiddleware), limits) + registerApiSupportRoutes(apiGroup.Group("/support"), limits) registerApiWebhookRoutes(apiGroup.Group("/webhooks")) wsGroup := app.Group("/api/ws") diff --git a/internal/bootstrap/server_ratelimit_test.go b/internal/bootstrap/server_ratelimit_test.go new file mode 100644 index 00000000..f1ae3f14 --- /dev/null +++ b/internal/bootstrap/server_ratelimit_test.go @@ -0,0 +1,162 @@ +package bootstrap + +import ( + "net/http" + "net/http/httptest" + "strconv" + "strings" + "testing" + + "agent-desk/internal/pkg/config" + + "github.com/gin-gonic/gin" +) + +func parseInt(t *testing.T, value string) int { + t.Helper() + parsed, err := strconv.Atoi(value) + if err != nil { + t.Fatalf("%q is not an integer", value) + } + return parsed +} + +// setRateLimitTestConfig disables password login so /api/auth/login returns +// before it reaches the database. The limiter runs ahead of the handler either +// way, which is what these tests are about. +func setRateLimitTestConfig(rateLimit config.RateLimitConfig) { + disabled := false + config.SetCurrent(&config.Config{ + Server: config.ServerConfig{ + RateLimit: rateLimit, + CORS: config.CORSConfig{AllowedOrigins: []string{}}, + }, + Auth: config.AuthConfig{PasswordLoginEnabled: &disabled}, + Storage: config.StorageConfig{Local: config.LocalStorageConfig{Root: "storage", BaseURL: "/storage"}}, + }) +} + +func postJSON(app *gin.Engine, path, clientIP string) *httptest.ResponseRecorder { + rec := httptest.NewRecorder() + req := httptest.NewRequest(http.MethodPost, path, strings.NewReader(`{"username":"admin","password":"secret"}`)) + req.Header.Set("Content-Type", "application/json") + req.RemoteAddr = clientIP + ":52000" + app.ServeHTTP(rec, req) + return rec +} + +func TestNewServerRateLimitsThePublicLoginEndpoint(t *testing.T) { + setRateLimitTestConfig(config.RateLimitConfig{}) + app, err := NewServer() + if err != nil { + t.Fatalf("NewServer() error = %v", err) + } + + for i := 1; i <= limitLogin; i++ { + if rec := postJSON(app, "/api/auth/login", "203.0.113.7"); rec.Code == http.StatusTooManyRequests { + t.Fatalf("request %d was throttled although the login limit is %d per window", i, limitLogin) + } + } + rec := postJSON(app, "/api/auth/login", "203.0.113.7") + if rec.Code != http.StatusTooManyRequests { + t.Fatalf("request %d got status %d, want 429", limitLogin+1, rec.Code) + } + if rec.Header().Get("Retry-After") == "" { + t.Error("the 429 carried no Retry-After header") + } + + // A second address must keep working, otherwise one attacker could lock the + // login page for every user behind their own NAT. + if rec := postJSON(app, "/api/auth/login", "198.51.100.9"); rec.Code == http.StatusTooManyRequests { + t.Fatal("an unrelated address was throttled by another address's requests") + } +} + +func TestNewServerRateLimitsTheSupportRegistrationEndpoint(t *testing.T) { + setRateLimitTestConfig(config.RateLimitConfig{}) + app, err := NewServer() + if err != nil { + t.Fatalf("NewServer() error = %v", err) + } + + for i := 1; i <= limitSupportRegister; i++ { + postJSON(app, "/api/support/auth/register", "203.0.113.7") + } + if rec := postJSON(app, "/api/support/auth/register", "203.0.113.7"); rec.Code != http.StatusTooManyRequests { + t.Fatalf("registration attempt %d got status %d, want 429", limitSupportRegister+1, rec.Code) + } +} + +// TestNewServerLeavesWebhooksAndPublicReadsUnthrottled guards the exemption that +// matters most. A channel platform that receives a 429 from its webhook stops +// retrying and eventually disables the delivery, which takes a whole channel +// offline - a far worse outcome than the flood the limit was meant to stop. +func TestNewServerLeavesWebhooksAndPublicReadsUnthrottled(t *testing.T) { + setRateLimitTestConfig(config.RateLimitConfig{}) + app, err := NewServer() + if err != nil { + t.Fatalf("NewServer() error = %v", err) + } + + requests := []struct { + method string + path string + }{ + {http.MethodGet, "/api/health"}, + {http.MethodGet, "/api/config"}, + {http.MethodPost, "/api/webhooks/org-sync"}, + {http.MethodGet, "/api/third/telegram/webhook"}, + {http.MethodPost, "/api/third/whatsapp/webhook"}, + } + for _, r := range requests { + for i := 0; i < 200; i++ { + rec := httptest.NewRecorder() + req := httptest.NewRequest(r.method, r.path, strings.NewReader(`{}`)) + req.Header.Set("Content-Type", "application/json") + req.RemoteAddr = "203.0.113.7:52000" + app.ServeHTTP(rec, req) + // Any status is acceptable here except 429: these routes either + // succeed, fail validation, or error out, but they must never be + // throttled. + if rec.Code == http.StatusTooManyRequests { + t.Fatalf("%s %s returned 429 on request %d; this route must stay unthrottled", r.method, r.path, i+1) + } + } + } +} + +func TestNewServerRateLimitCanBeDisabled(t *testing.T) { + disabled := false + setRateLimitTestConfig(config.RateLimitConfig{Enabled: &disabled}) + app, err := NewServer() + if err != nil { + t.Fatalf("NewServer() error = %v", err) + } + + for i := 0; i < limitLogin*4; i++ { + if rec := postJSON(app, "/api/auth/login", "203.0.113.7"); rec.Code == http.StatusTooManyRequests { + t.Fatalf("request %d was throttled although rate limiting is disabled", i+1) + } + } +} + +func TestNewServerRateLimitWindowIsConfigurable(t *testing.T) { + setRateLimitTestConfig(config.RateLimitConfig{WindowSeconds: 3600}) + app, err := NewServer() + if err != nil { + t.Fatalf("NewServer() error = %v", err) + } + + for i := 1; i <= limitLogin; i++ { + postJSON(app, "/api/auth/login", "203.0.113.7") + } + rec := postJSON(app, "/api/auth/login", "203.0.113.7") + if rec.Code != http.StatusTooManyRequests { + t.Fatalf("got status %d, want 429", rec.Code) + } + if retryAfter := rec.Header().Get("Retry-After"); retryAfter == "" { + t.Fatal("no Retry-After header") + } else if seconds := parseInt(t, retryAfter); seconds < 3500 || seconds > 3600 { + t.Errorf("Retry-After = %s, want roughly the configured 3600 second window", retryAfter) + } +} diff --git a/internal/middleware/ratelimit_middleware.go b/internal/middleware/ratelimit_middleware.go new file mode 100644 index 00000000..506b31b7 --- /dev/null +++ b/internal/middleware/ratelimit_middleware.go @@ -0,0 +1,56 @@ +package middleware + +import ( + "net/http" + "strconv" + "time" + + "agent-desk/internal/pkg/errorsx" + "agent-desk/internal/pkg/httpx" + "agent-desk/internal/pkg/ratelimit" + + "github.com/gin-gonic/gin" +) + +// rateLimitMessageKey is the localized explanation returned with a 429. It +// deliberately does not disclose the limit or the remaining budget, which would +// only tell a caller how much room they have left. +const rateLimitMessageKey = "error.e0354" + +// RateLimit bounds how often one client address may call a single public +// endpoint. +// +// The address comes from ctx.ClientIP(), which is only meaningful because +// NewServer configures Gin's trusted proxies and platform before any middleware +// is registered. Without that a caller picks their own bucket by setting +// X-Forwarded-For, and the limit is decoration. +// +// A nil limiter allows everything, so "disabled" is expressed by handing out nil +// limiters rather than by branching at every call site. +// +// Rejections are not logged here. requestLogMiddleware already records the path, +// status and client address of every request, so a 429 is visible there without +// giving a flood a second way to fill the log. +func RateLimit(limiter *ratelimit.Limiter) gin.HandlerFunc { + return func(ctx *gin.Context) { + allowed, retryAfter := limiter.Allow(ctx.ClientIP()) + if allowed { + ctx.Next() + return + } + // Retry-After is a whole number of seconds, so round up: telling a caller + // to come back sooner than the window actually resets just earns another + // 429. Rounding down and adding one overshoots when the remaining time is + // already an exact number of seconds. + seconds := int64(retryAfter / time.Second) + if time.Duration(seconds)*time.Second < retryAfter { + seconds++ + } + if seconds < 1 { + seconds = 1 + } + ctx.Header("Retry-After", strconv.FormatInt(seconds, 10)) + httpx.WriteHttpStatusJSON(ctx, http.StatusTooManyRequests, errorsx.InvalidParamI18n(rateLimitMessageKey)) + ctx.Abort() + } +} diff --git a/internal/middleware/ratelimit_middleware_test.go b/internal/middleware/ratelimit_middleware_test.go new file mode 100644 index 00000000..e2aa1cbf --- /dev/null +++ b/internal/middleware/ratelimit_middleware_test.go @@ -0,0 +1,103 @@ +package middleware + +import ( + "encoding/json" + "net/http" + "net/http/httptest" + "strconv" + "testing" + "time" + + "agent-desk/internal/pkg/ratelimit" + + "github.com/gin-gonic/gin" +) + +func newRateLimitTestEngine(limiter *ratelimit.Limiter) *gin.Engine { + gin.SetMode(gin.TestMode) + app := gin.New() + app.POST("/probe", RateLimit(limiter), func(ctx *gin.Context) { + ctx.JSON(http.StatusOK, gin.H{"success": true}) + }) + return app +} + +func postProbe(app *gin.Engine, clientIP string) *httptest.ResponseRecorder { + rec := httptest.NewRecorder() + req := httptest.NewRequest(http.MethodPost, "/probe", nil) + req.RemoteAddr = clientIP + ":52000" + app.ServeHTTP(rec, req) + return rec +} + +func TestRateLimitRefusesAfterTheLimitAndSetsRetryAfter(t *testing.T) { + app := newRateLimitTestEngine(ratelimit.New(2, time.Minute)) + + for i := 1; i <= 2; i++ { + if rec := postProbe(app, "203.0.113.7"); rec.Code != http.StatusOK { + t.Fatalf("request %d got status %d, want 200", i, rec.Code) + } + } + + rec := postProbe(app, "203.0.113.7") + if rec.Code != http.StatusTooManyRequests { + t.Fatalf("the third request got status %d, want 429", rec.Code) + } + + retryAfter := rec.Header().Get("Retry-After") + if retryAfter == "" { + t.Fatal("a 429 without Retry-After gives the caller nothing to act on") + } + seconds, err := strconv.Atoi(retryAfter) + if err != nil { + t.Fatalf("Retry-After = %q is not a number of seconds", retryAfter) + } + // The window is one minute and the requests are microseconds apart, so the + // answer is 60 give or take a clock tick. What must not happen is a value + // above the window, which would tell the caller to wait longer than needed, or + // zero, which would have them retry immediately and collect another 429. + if seconds < 1 || seconds > 60 { + t.Errorf("Retry-After = %d seconds, want between 1 and the 60 second window", seconds) + } + + // The body has to stay a JsonResult, because web/lib/api/client.ts parses the + // payload before it looks at response.ok and surfaces payload.message. + var body struct { + Success bool `json:"success"` + Message string `json:"message"` + } + if err := json.Unmarshal(rec.Body.Bytes(), &body); err != nil { + t.Fatalf("429 body is not JSON: %v (%s)", err, rec.Body.String()) + } + if body.Success { + t.Error("429 body reported success=true") + } + if body.Message == "" { + t.Error("429 body carried no message, so the caller would see a blank error") + } +} + +func TestRateLimitIsKeyedByClientAddress(t *testing.T) { + app := newRateLimitTestEngine(ratelimit.New(1, time.Minute)) + + if rec := postProbe(app, "203.0.113.7"); rec.Code != http.StatusOK { + t.Fatalf("first request got %d, want 200", rec.Code) + } + if rec := postProbe(app, "203.0.113.7"); rec.Code != http.StatusTooManyRequests { + t.Fatalf("second request from the same address got %d, want 429", rec.Code) + } + // Somebody else must not inherit the first caller's exhaustion. Behind a NAT + // this is the difference between a working product and a shared lockout. + if rec := postProbe(app, "198.51.100.9"); rec.Code != http.StatusOK { + t.Fatalf("a different address got %d, want 200", rec.Code) + } +} + +func TestRateLimitWithNilLimiterAllowsEverything(t *testing.T) { + app := newRateLimitTestEngine(nil) + for i := 0; i < 50; i++ { + if rec := postProbe(app, "203.0.113.7"); rec.Code != http.StatusOK { + t.Fatalf("request %d got %d with a nil limiter, want 200", i, rec.Code) + } + } +} diff --git a/internal/pkg/config/config.go b/internal/pkg/config/config.go index 827af873..d3549cae 100644 --- a/internal/pkg/config/config.go +++ b/internal/pkg/config/config.go @@ -69,7 +69,8 @@ type ServerConfig struct { // TrustedPlatform names an edge that overwrites rather than appends the real // client address, for example "cloudflare". When set it takes precedence over // X-Forwarded-For entirely. - TrustedPlatform string `yaml:"trustedPlatform"` + TrustedPlatform string `yaml:"trustedPlatform"` + RateLimit RateLimitConfig `yaml:"rateLimit"` } // defaultTrustedProxies covers loopback, RFC1918, IPv6 unique-local and @@ -147,6 +148,33 @@ type CORSConfig struct { AllowedOrigins []string `yaml:"allowedOrigins"` } +// RateLimitConfig bounds how often one client address may call the public, +// unauthenticated endpoints. It deliberately does not cover channel webhooks, +// websockets or authenticated dashboard routes: a platform that receives a 429 +// from a webhook endpoint stops retrying and eventually disables the webhook. +type RateLimitConfig struct { + // Enabled defaults to true. Set it to false to switch the limits off without + // taking them out of the route table. + Enabled *bool `yaml:"enabled"` + // WindowSeconds is the length of the counting window. Zero or negative means + // one minute. + WindowSeconds int `yaml:"windowSeconds"` +} + +func (r RateLimitConfig) IsEnabled() bool { + if r.Enabled == nil { + return true + } + return *r.Enabled +} + +func (r RateLimitConfig) WindowSecondsOrDefault() int { + if r.WindowSeconds <= 0 { + return 60 + } + return r.WindowSeconds +} + type DBConfig struct { Type string `yaml:"type"` DSN string `yaml:"dsn"` @@ -510,6 +538,7 @@ func bindConfigDefaults(v *viper.Viper) { v.SetDefault("server.cors.allowedOrigins", []string{}) v.SetDefault("server.trustedProxies", []string{}) v.SetDefault("server.trustedPlatform", "") + v.SetDefault("server.rateLimit.windowSeconds", 60) v.SetDefault("db.type", "sqlite") v.SetDefault("db.dsn", "file:./data/app.db?_busy_timeout=5000") v.SetDefault("db.maxIdleConns", 5) @@ -581,6 +610,8 @@ func bindEnvironmentAliases(v *viper.Viper) { _ = v.BindEnv("server.companyFaviconUrl", "AGENT_DESK_SERVER_COMPANYFAVICONURL", "COMPANY_FAVICON_URL", "NEXT_PUBLIC_COMPANY_FAVICON_URL", "BRAND_FAVICON_URL", "FAVICON_URL") _ = v.BindEnv("server.trustedProxies", "AGENT_DESK_SERVER_TRUSTEDPROXIES", "TRUSTED_PROXIES") _ = v.BindEnv("server.trustedPlatform", "AGENT_DESK_SERVER_TRUSTEDPLATFORM", "TRUSTED_PLATFORM") + _ = v.BindEnv("server.rateLimit.enabled", "AGENT_DESK_SERVER_RATELIMIT_ENABLED", "RATE_LIMIT_ENABLED") + _ = v.BindEnv("server.rateLimit.windowSeconds", "AGENT_DESK_SERVER_RATELIMIT_WINDOWSECONDS", "RATE_LIMIT_WINDOW_SECONDS") _ = v.BindEnv("db.type", "AGENT_DESK_DB_TYPE", "DATABASE_TYPE", "DB_TYPE") _ = v.BindEnv("db.dsn", "AGENT_DESK_DB_DSN", "DATABASE_URL", "DB_DSN") _ = v.BindEnv("auth.passwordLoginEnabled", "AGENT_DESK_AUTH_PASSWORDLOGINENABLED", "PASSWORD_LOGIN_ENABLED") diff --git a/internal/pkg/config/config_test.go b/internal/pkg/config/config_test.go index 5e600d9c..fbc71194 100644 --- a/internal/pkg/config/config_test.go +++ b/internal/pkg/config/config_test.go @@ -530,3 +530,54 @@ func TestLoadReadsTrustedProxySettings(t *testing.T) { t.Errorf("TrustedPlatformHeader() = %q want CF-Connecting-IP", cfg.Server.TrustedPlatformHeader()) } } + +func TestRateLimitConfigDefaultsToEnabledWithAMinuteWindow(t *testing.T) { + var cfg RateLimitConfig + if !cfg.IsEnabled() { + t.Error("rate limiting must default to enabled; a zero-value config should not silently turn protection off") + } + if got := cfg.WindowSecondsOrDefault(); got != 60 { + t.Errorf("WindowSecondsOrDefault() = %d want 60", got) + } + + off := false + cfg = RateLimitConfig{Enabled: &off, WindowSeconds: -5} + if cfg.IsEnabled() { + t.Error("Enabled=false must disable the limits") + } + if got := cfg.WindowSecondsOrDefault(); got != 60 { + t.Errorf("a non-positive window must fall back to 60, got %d", got) + } + + on := true + cfg = RateLimitConfig{Enabled: &on, WindowSeconds: 30} + if !cfg.IsEnabled() { + t.Error("Enabled=true must enable the limits") + } + if got := cfg.WindowSecondsOrDefault(); got != 30 { + t.Errorf("WindowSecondsOrDefault() = %d want 30", got) + } +} + +func TestLoadReadsRateLimitSettings(t *testing.T) { + path := filepath.Join(t.TempDir(), "config.yaml") + if err := os.WriteFile(path, []byte("server:\n port: 8083\n"), 0600); err != nil { + t.Fatalf("WriteFile() error = %v", err) + } + + t.Setenv("ENV_FILE", os.DevNull) + t.Setenv("AGENT_DESK_ENV_FILE", os.DevNull) + t.Setenv("RATE_LIMIT_ENABLED", "false") + t.Setenv("RATE_LIMIT_WINDOW_SECONDS", "30") + + cfg, err := Load(path) + if err != nil { + t.Fatalf("Load() error = %v", err) + } + if cfg.Server.RateLimit.IsEnabled() { + t.Error("RATE_LIMIT_ENABLED=false did not disable the limits") + } + if got := cfg.Server.RateLimit.WindowSecondsOrDefault(); got != 30 { + t.Errorf("WindowSecondsOrDefault() = %d want 30", got) + } +} diff --git a/internal/pkg/i18nx/locales/en-US.yml b/internal/pkg/i18nx/locales/en-US.yml index b74ee78d..1ba2cbe5 100644 --- a/internal/pkg/i18nx/locales/en-US.yml +++ b/internal/pkg/i18nx/locales/en-US.yml @@ -352,6 +352,7 @@ error.e0350: "WhatsApp webhook rejected: the X-Hub-Signature-256 header is missi error.e0351: "WhatsApp authorization needs a Meta App ID and App Secret. Set WHATSAPP_APP_ID and WHATSAPP_APP_SECRET, or appId and appSecret on the channel." error.e0352: "WhatsApp authorization failed: %s" error.e0353: "No Meta App ID is configured for this channel. Set %s, or the app id on the channel." +error.e0354: "Too many requests. Please wait a moment and try again." error.whatsapp.oauth.tokenInspectFailed: "Could not inspect the access token: %s" error.whatsapp.oauth.tokenInvalid: "Meta reports this access token is not valid." error.whatsapp.oauth.scopeMissing: "The granted token is missing the \"%s\" permission." diff --git a/internal/pkg/i18nx/locales/zh-CN.yml b/internal/pkg/i18nx/locales/zh-CN.yml index dc038fe9..699be781 100644 --- a/internal/pkg/i18nx/locales/zh-CN.yml +++ b/internal/pkg/i18nx/locales/zh-CN.yml @@ -352,6 +352,7 @@ error.e0350: "WhatsApp 回调已被拒绝:缺少 X-Hub-Signature-256 请求头 error.e0351: "WhatsApp 授权需要 Meta App ID 和 App Secret。请设置 WHATSAPP_APP_ID 和 WHATSAPP_APP_SECRET,或在渠道中填写 appId 和 appSecret。" error.e0352: "WhatsApp 授权失败:%s" error.e0353: "该渠道未配置 Meta App ID。请设置 %s,或在渠道中填写 app id。" +error.e0354: "请求过于频繁,请稍后再试。" error.whatsapp.oauth.tokenInspectFailed: "无法校验 Access Token:%s" error.whatsapp.oauth.tokenInvalid: "Meta 返回该 Access Token 无效。" error.whatsapp.oauth.scopeMissing: "授权令牌缺少 \"%s\" 权限。" diff --git a/internal/pkg/ratelimit/ratelimit.go b/internal/pkg/ratelimit/ratelimit.go new file mode 100644 index 00000000..47f4e6ea --- /dev/null +++ b/internal/pkg/ratelimit/ratelimit.go @@ -0,0 +1,101 @@ +// Package ratelimit provides a small in-process fixed-window counter for the +// public endpoints that an unauthenticated caller can drive. +// +// It is deliberately not distributed. A deployment running several replicas gets +// the configured limit per replica rather than per cluster, which weakens the +// bound by a factor of the replica count. That is the right trade here because +// these limits exist to make flooding expensive, not to meter a quota, and +// because reaching for a shared store would put a network round trip on the +// hottest unauthenticated paths in the application. If this ever needs to be +// exact across replicas, the Limiter interface is small enough to swap. +package ratelimit + +import ( + "sync" + "time" +) + +type bucket struct { + count int + resetAt time.Time +} + +// Limiter counts requests per key inside a fixed window. The zero value is not +// usable; call New. A nil *Limiter allows everything, so a caller can represent +// "disabled" by holding nil instead of branching at every use site. +type Limiter struct { + limit int + window time.Duration + + mu sync.Mutex + buckets map[string]bucket + sweptAt time.Time +} + +// New returns a Limiter that allows limit requests per window for each distinct +// key. A limit of zero or less, or a non-positive window, disables the limiter. +func New(limit int, window time.Duration) *Limiter { + if limit <= 0 || window <= 0 { + return &Limiter{limit: limit, window: window, buckets: make(map[string]bucket)} + } + return &Limiter{limit: limit, window: window, buckets: make(map[string]bucket)} +} + +// Allow records one request for key and reports whether it may proceed. When it +// returns false, retryAfter is how long the caller should wait before the window +// resets, and is meant for a Retry-After header. +func (l *Limiter) Allow(key string) (bool, time.Duration) { + if l == nil || l.limit <= 0 || l.window <= 0 { + return true, 0 + } + + now := time.Now() + + l.mu.Lock() + defer l.mu.Unlock() + + l.sweep(now) + + current, ok := l.buckets[key] + if !ok || !now.Before(current.resetAt) { + l.buckets[key] = bucket{count: 1, resetAt: now.Add(l.window)} + return true, 0 + } + + current.count++ + l.buckets[key] = current + if current.count > l.limit { + retryAfter := current.resetAt.Sub(now) + if retryAfter < 0 { + retryAfter = 0 + } + return false, retryAfter + } + return true, 0 +} + +// Len reports how many keys are currently tracked. It exists for tests. +func (l *Limiter) Len() int { + if l == nil { + return 0 + } + l.mu.Lock() + defer l.mu.Unlock() + return len(l.buckets) +} + +// sweep drops buckets whose window has elapsed, so a long-running process does +// not accumulate one entry per address it has ever seen. It runs at most once per +// window and under the same lock as Allow, which keeps the cost amortised and +// avoids a background goroutine whose lifetime somebody would have to own. +func (l *Limiter) sweep(now time.Time) { + if !l.sweptAt.IsZero() && now.Sub(l.sweptAt) < l.window { + return + } + for key, current := range l.buckets { + if !now.Before(current.resetAt) { + delete(l.buckets, key) + } + } + l.sweptAt = now +} diff --git a/internal/pkg/ratelimit/ratelimit_test.go b/internal/pkg/ratelimit/ratelimit_test.go new file mode 100644 index 00000000..dfa78ea6 --- /dev/null +++ b/internal/pkg/ratelimit/ratelimit_test.go @@ -0,0 +1,132 @@ +package ratelimit + +import ( + "fmt" + "sync" + "sync/atomic" + "testing" + "time" +) + +func TestLimiterAllowsUpToTheLimitThenRefuses(t *testing.T) { + l := New(3, time.Minute) + for i := 1; i <= 3; i++ { + if ok, retryAfter := l.Allow("203.0.113.7"); !ok { + t.Fatalf("request %d was refused although the limit is 3 (retryAfter=%v)", i, retryAfter) + } + } + + ok, retryAfter := l.Allow("203.0.113.7") + if ok { + t.Fatal("the fourth request was allowed although the limit is 3") + } + if retryAfter <= 0 || retryAfter > time.Minute { + t.Fatalf("retryAfter = %v, want a positive duration within the window", retryAfter) + } +} + +func TestLimiterKeysAreIndependent(t *testing.T) { + l := New(1, time.Minute) + if ok, _ := l.Allow("203.0.113.7"); !ok { + t.Fatal("the first request for a key was refused") + } + if ok, _ := l.Allow("203.0.113.7"); ok { + t.Fatal("the second request for the same key was allowed") + } + // A different caller must not inherit somebody else's exhaustion. Without key + // isolation one blocked address would lock out every user behind it. + if ok, _ := l.Allow("198.51.100.9"); !ok { + t.Fatal("a different key was refused because another key was exhausted") + } +} + +func TestLimiterWindowResets(t *testing.T) { + l := New(1, 20*time.Millisecond) + if ok, _ := l.Allow("k"); !ok { + t.Fatal("the first request was refused") + } + if ok, _ := l.Allow("k"); ok { + t.Fatal("the second request inside the window was allowed") + } + time.Sleep(40 * time.Millisecond) + if ok, _ := l.Allow("k"); !ok { + t.Fatal("the window did not reset") + } +} + +func TestNilLimiterAllowsEverything(t *testing.T) { + var l *Limiter + for i := 0; i < 100; i++ { + if ok, retryAfter := l.Allow("k"); !ok { + t.Fatalf("a nil limiter refused request %d (retryAfter=%v)", i, retryAfter) + } + } + if l.Len() != 0 { + t.Fatalf("Len() on a nil limiter = %d want 0", l.Len()) + } +} + +func TestDisabledLimiterAllowsEverything(t *testing.T) { + cases := []struct { + name string + limiter *Limiter + }{ + {"zero limit", New(0, time.Minute)}, + {"negative limit", New(-1, time.Minute)}, + {"zero window", New(5, 0)}, + {"negative window", New(5, -time.Second)}, + } + for _, tc := range cases { + for i := 0; i < 50; i++ { + if ok, _ := tc.limiter.Allow("k"); !ok { + t.Fatalf("%s: request %d was refused although the limiter is disabled", tc.name, i) + } + } + } +} + +func TestSweepReleasesExpiredKeys(t *testing.T) { + l := New(10, 20*time.Millisecond) + for i := 0; i < 200; i++ { + l.Allow(fmt.Sprintf("key-%d", i)) + } + if got := l.Len(); got != 200 { + t.Fatalf("Len() = %d before the sweep, want 200", got) + } + + time.Sleep(40 * time.Millisecond) + // The sweep is driven by Allow rather than by a goroutine, so one more call + // has to reclaim everything that expired. + l.Allow("fresh") + if got := l.Len(); got != 1 { + t.Fatalf("Len() = %d after the sweep, want only the fresh key", got) + } +} + +// TestLimiterCountsExactlyUnderConcurrency is the reason Allow holds a single +// mutex across the read-modify-write rather than checking and then recording. +func TestLimiterCountsExactlyUnderConcurrency(t *testing.T) { + const goroutines = 16 + const perGoroutine = 200 + const limit = 1000 + + l := New(limit, time.Minute) + var allowed atomic.Int64 + var wg sync.WaitGroup + for i := 0; i < goroutines; i++ { + wg.Add(1) + go func() { + defer wg.Done() + for j := 0; j < perGoroutine; j++ { + if ok, _ := l.Allow("shared"); ok { + allowed.Add(1) + } + } + }() + } + wg.Wait() + + if got := allowed.Load(); got != limit { + t.Fatalf("allowed %d of %d requests, want exactly the limit of %d", got, goroutines*perGoroutine, limit) + } +}