package store import ( "context" "crypto/sha256" "crypto/subtle" "database/sql" "encoding/hex" "fmt" "time" "atlas9.dev/c/core" "atlas9.dev/c/core/dbi" "atlas9.dev/c/core/tokens" "atlas9.dev/c/demo/lib/access" "atlas9.dev/c/demo/lib/mfa" ) // mfaChallengeTTL is how long a texted code stays valid. Short: the user is // waiting on the code-entry screen. const mfaChallengeTTL = 5 * time.Minute var ( mfaChallengeID = tokens.RandomString(16) mfaChallengeCode = tokens.NumericCode(6) ) type SqliteMfaChallengeStore struct { db dbi.DBI guard access.Guard } var _ mfa.ChallengeStore = (*SqliteMfaChallengeStore)(nil) func NewSqliteMfaChallengeStore(db dbi.DBI, guard access.Guard) *SqliteMfaChallengeStore { return &SqliteMfaChallengeStore{db: db, guard: guard} } func (s *SqliteMfaChallengeStore) Create(ctx context.Context, userID core.ID) (string, string, error) { if err := s.guard.System(ctx, mfa.Cap_Mfa_Challenge); err != nil { return "", "", err } id := mfaChallengeID() code := mfaChallengeCode() expiresAt := time.Now().UTC().Add(mfaChallengeTTL) _, err := s.db.Exec(ctx, ` INSERT INTO mfa_challenges (id, user_id, code_hash, expires_at, attempts) VALUES ($1, $2, $3, $4, 0) `, id, userID, hashMfaCode(code), expiresAt.Format("2006-01-02 15:04:05")) if err != nil { return "", "", fmt.Errorf("storing challenge: %w", err) } return id, code, nil } func (s *SqliteMfaChallengeStore) Verify(ctx context.Context, id, code string) (core.ID, error) { if err := s.guard.System(ctx, mfa.Cap_Mfa_Challenge); err != nil { return core.ID{}, err } var userID core.ID var codeHash string var attempts int err := s.db.QueryRow(ctx, ` SELECT user_id, code_hash, attempts FROM mfa_challenges WHERE id = $1 AND expires_at > datetime('now') `, id).Scan(&userID, &codeHash, &attempts) if err == sql.ErrNoRows { return core.ID{}, core.ErrNotFound } if err != nil { return core.ID{}, fmt.Errorf("getting challenge: %w", err) } // Correct code: single-use, delete and return the user. if subtle.ConstantTimeCompare([]byte(codeHash), []byte(hashMfaCode(code))) == 1 { if _, err := s.db.Exec(ctx, `DELETE FROM mfa_challenges WHERE id = $1`, id); err != nil { return core.ID{}, err } return userID, nil } // Wrong code: burn an attempt, discarding the challenge once exhausted so a // stolen id can't be brute-forced. attempts++ if attempts >= mfa.MaxAttempts { _, err = s.db.Exec(ctx, `DELETE FROM mfa_challenges WHERE id = $1`, id) } else { _, err = s.db.Exec(ctx, `UPDATE mfa_challenges SET attempts = $1 WHERE id = $2`, attempts, id) } if err != nil { return core.ID{}, err } return core.ID{}, mfa.ErrInvalidCode } func (s *SqliteMfaChallengeStore) DeleteExpired(ctx context.Context) (int64, error) { if err := s.guard.System(ctx, tokens.CapTokensPrune); err != nil { return 0, err } res, err := s.db.Exec(ctx, `DELETE FROM mfa_challenges WHERE expires_at <= datetime('now')`) if err != nil { return 0, err } n, _ := res.RowsAffected() return n, nil } // hashMfaCode hashes a code the way session secrets are hashed: the code is // short but high-entropy enough for an unsalted SHA-256 to keep the table useless // if leaked, and constant-time compare guards verification. func hashMfaCode(code string) string { sum := sha256.Sum256([]byte(code)) return hex.EncodeToString(sum[:]) }