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) + } +}