// Package auth implements magic-link authentication and JWT issuance // for the librenotes multi-tenant backend. package auth import ( "context" "crypto/rand" "crypto/sha256" "crypto/subtle" "database/sql" "encoding/hex" "errors" "fmt" "time" ) // TokenLifetime is how long a magic link is valid after creation. const TokenLifetime = 15 * time.Minute // ErrInvalidToken is returned for unknown, expired, or already-used tokens. var ErrInvalidToken = errors.New("invalid or expired token") // MagicToken records a single magic-link issuance. The plaintext token // is never stored — only its SHA-256 hash. The plaintext lives only in // the email sent to the user. type MagicToken struct { UserID string Email string CreatedAt time.Time ExpiresAt time.Time } // TokenStore persists magic-link token hashes. type TokenStore struct { db *sql.DB } // NewTokenStore wraps a database handle. func NewTokenStore(db *sql.DB) *TokenStore { return &TokenStore{db: db} } // Issue generates a new random token, stores its hash bound to userID, // and returns the plaintext token (caller emails it to the user). func (s *TokenStore) Issue(ctx context.Context, userID, email string) (string, error) { plaintext, err := randomToken() if err != nil { return "", err } now := time.Now().UTC() hash := hashToken(plaintext) _, err = s.db.ExecContext(ctx, `INSERT INTO magic_tokens (token_hash, user_id, email, created_at, expires_at) VALUES (?, ?, ?, ?, ?)`, hash, userID, email, now.Unix(), now.Add(TokenLifetime).Unix()) if err != nil { return "", fmt.Errorf("insert magic token: %w", err) } return plaintext, nil } // Consume validates the plaintext token. If valid, it marks the token // used (single-use) and returns the bound user ID. Comparison goes // through subtle.ConstantTimeCompare; the hash lookup gives us O(1) // retrieval without leaking timing for the row scan, and the constant- // time compare protects against any residual differences. func (s *TokenStore) Consume(ctx context.Context, plaintext string) (userID string, err error) { hash := hashToken(plaintext) tx, err := s.db.BeginTx(ctx, nil) if err != nil { return "", fmt.Errorf("begin: %w", err) } defer tx.Rollback() var ( storedHash string uid string expires int64 used sql.NullInt64 ) err = tx.QueryRowContext(ctx, `SELECT token_hash, user_id, expires_at, used_at FROM magic_tokens WHERE token_hash = ?`, hash).Scan(&storedHash, &uid, &expires, &used) if errors.Is(err, sql.ErrNoRows) { return "", ErrInvalidToken } if err != nil { return "", fmt.Errorf("select magic token: %w", err) } if subtle.ConstantTimeCompare([]byte(storedHash), []byte(hash)) != 1 { return "", ErrInvalidToken } if used.Valid { return "", ErrInvalidToken } if time.Now().UTC().Unix() > expires { return "", ErrInvalidToken } res, err := tx.ExecContext(ctx, `UPDATE magic_tokens SET used_at = ? WHERE token_hash = ? AND used_at IS NULL`, time.Now().UTC().Unix(), hash) if err != nil { return "", fmt.Errorf("mark used: %w", err) } n, _ := res.RowsAffected() if n != 1 { return "", ErrInvalidToken } if err := tx.Commit(); err != nil { return "", fmt.Errorf("commit: %w", err) } return uid, nil } // PurgeExpired deletes all expired or used tokens older than retention. // Callers can run this on a timer to keep the table small. func (s *TokenStore) PurgeExpired(ctx context.Context, retention time.Duration) error { cutoff := time.Now().UTC().Add(-retention).Unix() _, err := s.db.ExecContext(ctx, `DELETE FROM magic_tokens WHERE expires_at < ? OR (used_at IS NOT NULL AND used_at < ?)`, cutoff, cutoff) if err != nil { return fmt.Errorf("purge: %w", err) } return nil } func randomToken() (string, error) { buf := make([]byte, 32) if _, err := rand.Read(buf); err != nil { return "", fmt.Errorf("rand: %w", err) } return hex.EncodeToString(buf), nil } func hashToken(plaintext string) string { sum := sha256.Sum256([]byte(plaintext)) return hex.EncodeToString(sum[:]) }