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..bcb1490a --- /dev/null +++ b/internal/server/app_rate_limiter_integration_test.go @@ -0,0 +1,170 @@ +//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, body string) (int, error) { + method := http.MethodGet + if path == "/api/auth/login" { + method = http.MethodPost + } + 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() + 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) + } + 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.StatusBadRequest: 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)") })