From e4b48a896249470b0bac8b9197a963ccb105bc80 Mon Sep 17 00:00:00 2001 From: Cristian Magherusan-Stanciu Date: Wed, 30 Sep 2026 01:48:10 +0200 Subject: [PATCH 1/2] fix(server): share rate limits across all replicas Select the PostgreSQL limiter after every database connection instead of retaining process-local counters on non-Lambda deployments. Verify shared login and approval budgets through two real HTTP servers, including concurrent attempts and fail-closed database errors. Closes #109 --- internal/api/inmemory_rate_limiter.go | 4 +- internal/server/app.go | 29 +-- .../app_rate_limiter_integration_test.go | 166 ++++++++++++++++++ internal/server/app_test.go | 4 +- 4 files changed, 178 insertions(+), 25 deletions(-) create mode 100644 internal/server/app_rate_limiter_integration_test.go diff --git a/internal/api/inmemory_rate_limiter.go b/internal/api/inmemory_rate_limiter.go index b7ae8183..a570afbc 100644 --- a/internal/api/inmemory_rate_limiter.go +++ b/internal/api/inmemory_rate_limiter.go @@ -16,8 +16,8 @@ import ( // memory growth from rotating attacker source IPs (02-M3). const inMemoryRateLimitMaxEntries = 500 -// InMemoryRateLimiter provides in-memory rate limiting for single-instance deployments (Fargate, ECS) -// This implementation should NOT be used for Lambda (multi-instance) - use DBRateLimiter instead. +// InMemoryRateLimiter holds process-local counters for explicitly single-process use +// or temporary initialization. Replicated deployments must use DBRateLimiter. type InMemoryRateLimiter struct { attempts map[string]*inMemoryRateLimitEntry limits map[string]RateLimitConfig diff --git a/internal/server/app.go b/internal/server/app.go index b6a94656..0da40d4c 100644 --- a/internal/server/app.go +++ b/internal/server/app.go @@ -494,18 +494,10 @@ func NewApplicationFromDeps(ctx context.Context, cfg ApplicationConfig, deps Ext DashboardURL: cfg.DashboardURL, }) - // Initialize rate limiter based on runtime environment. - // Lambda: start with an in-memory limiter immediately so the first cold-start - // request is protected. ensureDB() swaps it for the DB-backed limiter once the - // database connection is established (distributed state across warm containers). - // Fargate/containers: in-memory is the permanent implementation because the - // process is long-lived and single-instance. + // Sensitive requests pass ensureDB, which replaces this temporary limiter + // with shared database counters before dispatch on every runtime. rateLimiter := api.RateLimiterInterface(api.NewInMemoryRateLimiter()) - if !cfg.IsLambda { - log.Println("Initialized in-memory rate limiter for single-instance deployment (Fargate/Container)") - } else { - log.Println("Initialized in-memory rate limiter for Lambda cold-start (will be upgraded to DB-backed on first DB connect)") - } + log.Println("Initialized temporary in-memory rate limiter until database connection") // Initialize API handler apiHandler := api.NewHandler(api.HandlerConfig{ @@ -781,16 +773,11 @@ func (app *Application) reinitializeAfterConnect(ctx context.Context, dbConn *da } app.Auth = authSvc - // Initialize distributed rate limiter for Lambda (multi-instance) - // For Fargate/containers, we already have in-memory rate limiter from startup - if app.appConfig.IsLambda { - dbRL := api.NewDBRateLimiter(dbConn.Pool()) - // Start the scheduled cleanup worker so perpetually-denied keys (whose - // count never resets to 1) are still evicted on a fixed schedule (02-M2). - dbRL.StartCleanupWorker(ctx) - app.RateLimiter = dbRL - log.Println("Initialized database-backed rate limiter for Lambda (distributed state)") - } + dbRL := api.NewDBRateLimiter(dbConn.Pool()) + // Periodic cleanup also evicts expired keys with no subsequent allowed request. + dbRL.StartCleanupWorker(ctx) + app.RateLimiter = dbRL + log.Println("Initialized database-backed rate limiter with shared replica counters") // Initialize analytics store for savings data and materialized views, plus // the snapshot collector behind the scheduled analytics_collect task. diff --git a/internal/server/app_rate_limiter_integration_test.go b/internal/server/app_rate_limiter_integration_test.go new file mode 100644 index 00000000..9bedb032 --- /dev/null +++ b/internal/server/app_rate_limiter_integration_test.go @@ -0,0 +1,166 @@ +//go:build integration + +package server + +import ( + "context" + "fmt" + "io" + "net/http" + "net/http/httptest" + "strings" + "testing" + "time" + + "github.com/LeanerCloud/cloud-commitments-platform/internal/database/postgres/migrations" + "github.com/LeanerCloud/cloud-commitments-platform/internal/database/postgres/testhelpers" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestApplicationRateLimiter(t *testing.T) { + for _, key := range []string{ + "CREDENTIAL_ENCRYPTION_KEY_SECRET_ARN", "CREDENTIAL_ENCRYPTION_KEY_SECRET_NAME", + "CREDENTIAL_ENCRYPTION_KEY_SECRET_ID", "CREDENTIAL_ENCRYPTION_ALLOW_DEV_KEY", + "CUDLY_SIGNING_KEY_ID", "CUDLY_SIGNING_KEY_VAULT_URL", "CUDLY_SIGNING_KEY_NAME", + "CUDLY_SIGNING_KEY_RESOURCE", "AWS_PROFILE", "AWS_LAMBDA_RUNTIME_API", + } { + t.Setenv(key, "") + } + t.Setenv("CREDENTIAL_ENCRYPTION_KEY", strings.Repeat("12", 32)) + t.Setenv("CUDLY_ISSUER_URL", "https://rate-limit.example.test") + t.Setenv("CUDLY_SOURCE_CLOUD", "aws") + t.Setenv("SCHEDULED_TASK_AUTH_MODE", "disabled") + t.Setenv("AWS_EC2_METADATA_DISABLED", "true") + t.Setenv("AWS_ACCESS_KEY_ID", "local-rate-limit-test") + t.Setenv("AWS_SECRET_ACCESS_KEY", "local-rate-limit-test") + t.Setenv("AWS_REGION", "us-east-1") + + ctx, cancel := context.WithTimeout(context.Background(), 2*time.Minute) + defer cancel() + pg, err := testhelpers.SetupPostgresContainer(ctx, t) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, pg.Cleanup(context.Background())) }) + require.NoError(t, migrations.RunMigrations(ctx, pg.DB.Pool(), "../database/postgres/migrations", "", "")) + + for mode, isLambda := range []bool{false, true} { + t.Run(fmt.Sprintf("lambda=%t", isLambda), func(t *testing.T) { + cfg := ApplicationConfig{ + IsLambda: isLambda, DashboardURL: "https://rate-limit.example.test", + DefaultTerm: 3, DefaultCoverage: 80, + } + newApp := func() *Application { + app, appErr := NewApplicationFromDeps(ctx, cfg, ExternalDeps{ + DBConfig: pg.Config, EmailSender: &noopEmailSender{}, + }) + require.NoError(t, appErr) + t.Cleanup(func() { require.NoError(t, app.Close()) }) + return app + } + apps := []*Application{newApp(), newApp()} + servers := make([]*httptest.Server, len(apps)) + for i, app := range apps { + servers[i] = httptest.NewServer(CreateHTTPServer(app, 0).Handler) + t.Cleanup(servers[i].Close) + } + client := &http.Client{Timeout: 10 * time.Second} + ip := func(bucket int) string { return fmt.Sprintf("192.0.2.%d", mode*10+bucket) } + request := func(replica int, path, sourceIP string) (int, error) { + method, body := http.MethodGet, "" + if path == "/api/auth/login" { + method, body = http.MethodPost, `{"email":"absent@example.test","password":"d3JvbmctcGFzc3dvcmQ="}` + } + req, reqErr := http.NewRequestWithContext(ctx, method, servers[replica].URL+path, strings.NewReader(body)) + if reqErr != nil { + return 0, reqErr + } + req.Header.Set("X-Forwarded-For", sourceIP) + req.Header.Set("Content-Type", "application/json") + resp, reqErr := client.Do(req) + if reqErr != nil { + return 0, reqErr + } + defer resp.Body.Close() + _, reqErr = io.Copy(io.Discard, resp.Body) + return resp.StatusCode, reqErr + } + checkRequest := func(replica int, path, sourceIP string, want int) { + t.Helper() + status, reqErr := request(replica, path, sourceIP) + require.NoError(t, reqErr) + assert.Equal(t, want, status, "replica=%d path=%s source=%s", replica, path, sourceIP) + } + checkCount := func(sourceIP, endpoint string, want int) { + t.Helper() + var count int + err = pg.DB.Pool().QueryRow(ctx, "SELECT count FROM rate_limits WHERE id = $1", + "IP#"+sourceIP+"#ENDPOINT#"+endpoint).Scan(&count) + assert.NoError(t, err) + assert.Equal(t, want, count) + } + + for i := range 7 { + want := http.StatusUnauthorized + if i >= 5 { + want = http.StatusTooManyRequests + } + checkRequest(i%2, "/api/auth/login", ip(1), want) + } + checkCount(ip(1), "login", 7) + checkRequest(1, "/api/auth/login", ip(2), http.StatusUnauthorized) + checkCount(ip(2), "login", 1) + + for i := range 32 { + action := "approve" + if i%2 == 1 { + action = "cancel" + } + want := http.StatusNotFound + if i >= 30 { + want = http.StatusTooManyRequests + } + checkRequest(i%2, "/api/purchases/"+action+"/00000000-0000-4000-8000-000000000109?token=invalid", ip(3), want) + } + checkCount(ip(3), "approve_cancel_public", 32) + + type outcome struct { + status int + err error + } + results := make(chan outcome, 12) + for i := range 12 { + go func() { + status, reqErr := request(i%2, "/api/auth/login", ip(4)) + results <- outcome{status, reqErr} + }() + } + statuses := make(map[int]int) + for range 12 { + result := <-results + require.NoError(t, result.err) + statuses[result.status]++ + } + assert.Equal(t, map[int]int{http.StatusUnauthorized: 5, http.StatusTooManyRequests: 7}, statuses) + checkCount(ip(4), "login", 12) + + apps[0].DB.Close() + checkRequest(0, "/api/auth/login", ip(5), http.StatusServiceUnavailable) + checkRequest(1, "/api/auth/login", ip(5), http.StatusUnauthorized) + checkCount(ip(5), "login", 1) + + cold := newApp() + canceled, stop := context.WithCancel(ctx) + stop() + req := httptest.NewRequestWithContext(canceled, http.MethodPost, "/api/auth/login", strings.NewReader(`{}`)) + req.Header.Set("X-Forwarded-For", ip(6)) + response := httptest.NewRecorder() + CreateHTTPServer(cold, 0).Handler.ServeHTTP(response, req) + assert.Equal(t, http.StatusServiceUnavailable, response.Code) + assert.False(t, cold.dbConnected) + var count int + require.NoError(t, pg.DB.Pool().QueryRow(ctx, "SELECT COUNT(*) FROM rate_limits WHERE id = $1", + "IP#"+ip(6)+"#ENDPOINT#login").Scan(&count)) + assert.Zero(t, count) + }) + } +} diff --git a/internal/server/app_test.go b/internal/server/app_test.go index f6e6ff0d..3526bf8b 100644 --- a/internal/server/app_test.go +++ b/internal/server/app_test.go @@ -462,7 +462,7 @@ func TestNewApplicationFromDeps(t *testing.T) { IsLambda: false, } - t.Run("non-Lambda path with in-memory rate limiter", func(t *testing.T) { + t.Run("non-Lambda path with temporary preconnect rate limiter", func(t *testing.T) { deps := ExternalDeps{ EmailSender: &noopEmailSender{}, DBConfig: validDBConfig, @@ -476,7 +476,7 @@ func TestNewApplicationFromDeps(t *testing.T) { testutil.AssertTrue(t, app.Scheduler != nil, "Scheduler should be created") testutil.AssertTrue(t, app.Purchase != nil, "Purchase manager should be created") testutil.AssertTrue(t, app.Auth != nil, "Auth service should be created") - testutil.AssertTrue(t, app.RateLimiter != nil, "Rate limiter should be in-memory for non-Lambda") + testutil.AssertTrue(t, app.RateLimiter != nil, "Rate limiter must be present before database connection") testutil.AssertTrue(t, app.DB == nil, "DB should be nil (lazy init)") }) From ef58fda4852f1816a781bc4219f228a3ceed29d5 Mon Sep 17 00:00:00 2001 From: Cristian Magherusan-Stanciu Date: Wed, 30 Sep 2026 02:33:25 +0200 Subject: [PATCH 2/2] test(server): isolate rate-limit concurrency from password hashing Use malformed login JSON for the concurrent budget assertion, which reaches the same rate limiter before parsing. Keep sequential valid login requests and exact shared database counters. This avoids concurrent cost-12 bcrypt work exceeding the CI client deadline under race and coverage instrumentation. --- .../server/app_rate_limiter_integration_test.go | 16 ++++++++++------ 1 file changed, 10 insertions(+), 6 deletions(-) diff --git a/internal/server/app_rate_limiter_integration_test.go b/internal/server/app_rate_limiter_integration_test.go index 9bedb032..bcb1490a 100644 --- a/internal/server/app_rate_limiter_integration_test.go +++ b/internal/server/app_rate_limiter_integration_test.go @@ -65,10 +65,10 @@ func TestApplicationRateLimiter(t *testing.T) { } client := &http.Client{Timeout: 10 * time.Second} ip := func(bucket int) string { return fmt.Sprintf("192.0.2.%d", mode*10+bucket) } - request := func(replica int, path, sourceIP string) (int, error) { - method, body := http.MethodGet, "" + request := func(replica int, path, sourceIP, body string) (int, error) { + method := http.MethodGet if path == "/api/auth/login" { - method, body = http.MethodPost, `{"email":"absent@example.test","password":"d3JvbmctcGFzc3dvcmQ="}` + method = http.MethodPost } req, reqErr := http.NewRequestWithContext(ctx, method, servers[replica].URL+path, strings.NewReader(body)) if reqErr != nil { @@ -86,7 +86,11 @@ func TestApplicationRateLimiter(t *testing.T) { } checkRequest := func(replica int, path, sourceIP string, want int) { t.Helper() - status, reqErr := request(replica, path, sourceIP) + body := "" + if path == "/api/auth/login" { + body = `{"email":"absent@example.test","password":"d3JvbmctcGFzc3dvcmQ="}` + } + status, reqErr := request(replica, path, sourceIP, body) require.NoError(t, reqErr) assert.Equal(t, want, status, "replica=%d path=%s source=%s", replica, path, sourceIP) } @@ -130,7 +134,7 @@ func TestApplicationRateLimiter(t *testing.T) { results := make(chan outcome, 12) for i := range 12 { go func() { - status, reqErr := request(i%2, "/api/auth/login", ip(4)) + status, reqErr := request(i%2, "/api/auth/login", ip(4), "{") results <- outcome{status, reqErr} }() } @@ -140,7 +144,7 @@ func TestApplicationRateLimiter(t *testing.T) { require.NoError(t, result.err) statuses[result.status]++ } - assert.Equal(t, map[int]int{http.StatusUnauthorized: 5, http.StatusTooManyRequests: 7}, statuses) + assert.Equal(t, map[int]int{http.StatusBadRequest: 5, http.StatusTooManyRequests: 7}, statuses) checkCount(ip(4), "login", 12) apps[0].DB.Close()