package api // This file is internal (package api, not api_test) because loginLimiter and // its blocked/fail/clear methods are unexported, and TestLoginLimiter_SharedAcrossReplicas // specifically needs to construct two separate loginLimiter values pointed at // one database — standing in for two replicas — which only this package can // do. It duplicates testdb_test.go's newTestDB/withSearchPath rather than // importing them: those live in the separate api_test package, compiled from // this directory's external test files, and are not visible here. Same // reasoning as advisory_lock_test.go, which makes the same trade for the // same reason. import ( "context" "database/sql" "fmt" "net/url" "os" "strings" "testing" "git.ryuvia.com/niklas/terdut-server/internal/db" _ "github.com/jackc/pgx/v5/stdlib" ) var rateLimiterSchemaSeq int // rateLimiterTestDB returns a migrated database private to this test. func rateLimiterTestDB(t *testing.T) *sql.DB { t.Helper() dsn := os.Getenv("TERDUT_TEST_DSN") if dsn == "" { t.Fatalf("TERDUT_TEST_DSN is not set: these tests need Postgres.\n" + "Run `make test-db` for a local one, then\n" + " export TERDUT_TEST_DSN=postgres://terdut:terdut@localhost:5432/terdut_test?sslmode=disable") } rateLimiterSchemaSeq++ schema := fmt.Sprintf("test_rl_%d_%d", os.Getpid(), rateLimiterSchemaSeq) admin, err := sql.Open("pgx", dsn) if err != nil { t.Fatalf("connect to TERDUT_TEST_DSN: %v", err) } defer admin.Close() if _, err := admin.Exec("CREATE SCHEMA " + schema); err != nil { t.Fatalf("create schema %s: %v", schema, err) } database, err := db.Open(rateLimiterWithSearchPath(dsn, schema)) if err != nil { t.Fatalf("open db: %v", err) } if err := db.Migrate(database); err != nil { t.Fatalf("migrate: %v", err) } t.Cleanup(func() { database.Close() cleanup, err := sql.Open("pgx", dsn) if err != nil { return } defer cleanup.Close() if _, err := cleanup.Exec("DROP SCHEMA " + schema + " CASCADE"); err != nil { t.Logf("drop schema %s: %v", schema, err) } }) return database } func rateLimiterWithSearchPath(dsn, schema string) string { opt := "-csearch_path=" + schema if strings.HasPrefix(dsn, "postgres://") || strings.HasPrefix(dsn, "postgresql://") { u, err := url.Parse(dsn) if err == nil { q := u.Query() q.Set("options", opt) u.RawQuery = q.Encode() return u.String() } } return dsn + " options='" + opt + "'" } func TestLoginLimiter_BlocksAtMax(t *testing.T) { database := rateLimiterTestDB(t) ctx := context.Background() l := newLoginLimiter(database) for range 3 { if l.blocked(ctx, "k", 3) { t.Fatal("blocked before reaching max") } l.fail(ctx, "k") } if !l.blocked(ctx, "k", 3) { t.Fatal("not blocked after reaching max") } } func TestLoginLimiter_ClearResetsTheCount(t *testing.T) { database := rateLimiterTestDB(t) ctx := context.Background() l := newLoginLimiter(database) l.fail(ctx, "k") l.fail(ctx, "k") l.clear(ctx, "k") if l.blocked(ctx, "k", 1) { t.Fatal("still blocked after clear") } } func TestLoginLimiter_KeysAreIndependent(t *testing.T) { database := rateLimiterTestDB(t) ctx := context.Background() l := newLoginLimiter(database) l.fail(ctx, "a") if l.blocked(ctx, "b", 1) { t.Fatal("failing one key blocked an unrelated one") } } // TestLoginLimiter_SharedAcrossReplicas is the regression test for the gap // this migration closes: an in-memory limiter would let each replica count // independently, so a caller hitting two different pods could rack up // max*replicaCount failures before either one blocked. Two loginLimiter // values sharing one database, standing in for two replicas behind the same // load balancer, must instead see one combined count. func TestLoginLimiter_SharedAcrossReplicas(t *testing.T) { database := rateLimiterTestDB(t) ctx := context.Background() replicaA := newLoginLimiter(database) replicaB := newLoginLimiter(database) const max = 4 // Alternate which "replica" records the failure, as a real deployment // would split requests across pods. for i := range max { replica := replicaA if i%2 == 1 { replica = replicaB } if replicaA.blocked(ctx, "k", max) || replicaB.blocked(ctx, "k", max) { t.Fatalf("blocked after only %d of %d failures", i, max) } replica.fail(ctx, "k") } if !replicaA.blocked(ctx, "k", max) { t.Fatal("replica A does not see the combined count as blocked") } if !replicaB.blocked(ctx, "k", max) { t.Fatal("replica B does not see the combined count as blocked") } }