package auth import ( "context" "database/sql" "errors" "fmt" "time" ) // ErrRateLimited indicates that too many requests have been made for a // given key within the configured window. var ErrRateLimited = errors.New("rate limited") // RateLimiter enforces a fixed-window rate limit per email using the // login_attempts table. The DB-backed approach survives restarts and // works across multiple processes sharing the same SQLite file. type RateLimiter struct { db *sql.DB window time.Duration max int } // NewRateLimiter creates a limiter with the given window and maximum // attempts. Default for magic-link login: 5 per 15 min. func NewRateLimiter(db *sql.DB, window time.Duration, max int) *RateLimiter { return &RateLimiter{db: db, window: window, max: max} } // Check records an attempt for email and returns ErrRateLimited if the // number of attempts within the window exceeds max. func (r *RateLimiter) Check(ctx context.Context, email string) error { now := time.Now().UTC() cutoff := now.Add(-r.window).Unix() tx, err := r.db.BeginTx(ctx, nil) if err != nil { return fmt.Errorf("begin: %w", err) } defer tx.Rollback() if _, err := tx.ExecContext(ctx, `DELETE FROM login_attempts WHERE created_at < ?`, cutoff); err != nil { return fmt.Errorf("prune: %w", err) } var count int if err := tx.QueryRowContext(ctx, `SELECT COUNT(*) FROM login_attempts WHERE email = ? AND created_at >= ?`, email, cutoff).Scan(&count); err != nil { return fmt.Errorf("count: %w", err) } if count >= r.max { return ErrRateLimited } if _, err := tx.ExecContext(ctx, `INSERT INTO login_attempts (email, created_at) VALUES (?, ?)`, email, now.Unix()); err != nil { return fmt.Errorf("insert: %w", err) } return tx.Commit() }