summaryrefslogtreecommitdiff
path: root/internal/controlplane/mfa.go
diff options
context:
space:
mode:
Diffstat (limited to '')
-rw-r--r--internal/controlplane/mfa.go344
1 files changed, 344 insertions, 0 deletions
diff --git a/internal/controlplane/mfa.go b/internal/controlplane/mfa.go
new file mode 100644
index 0000000..c96cac1
--- /dev/null
+++ b/internal/controlplane/mfa.go
@@ -0,0 +1,344 @@
+package controlplane
+
+import (
+ "bytes"
+ "context"
+ "crypto/sha256"
+ "encoding/base64"
+ "errors"
+ "fmt"
+ "image/png"
+ "strings"
+ "time"
+
+ "aigw/internal/security"
+
+ "github.com/jackc/pgx/v5"
+ "github.com/pquerna/otp"
+ "github.com/pquerna/otp/totp"
+)
+
+func (s *Store) MFAMethods(ctx context.Context, userID string) ([]string, error) {
+ var totpEnabled, passkeyEnabled, recoveryEnabled bool
+ if err := s.db.QueryRow(ctx, `SELECT EXISTS(SELECT 1 FROM console_totp_credentials WHERE user_id=$1 AND confirmed_at IS NOT NULL),
+ EXISTS(SELECT 1 FROM console_passkeys WHERE user_id=$1),
+ EXISTS(SELECT 1 FROM console_recovery_codes WHERE user_id=$1 AND used_at IS NULL)`, userID).
+ Scan(&totpEnabled, &passkeyEnabled, &recoveryEnabled); err != nil {
+ return nil, err
+ }
+ methods := make([]string, 0, 3)
+ if totpEnabled {
+ methods = append(methods, "totp")
+ }
+ if passkeyEnabled {
+ methods = append(methods, "passkey")
+ }
+ if recoveryEnabled {
+ methods = append(methods, "recovery")
+ }
+ return methods, nil
+}
+
+func (s *Store) MFAStatus(ctx context.Context, userID string) (MFAStatus, error) {
+ var result MFAStatus
+ if err := s.db.QueryRow(ctx, `SELECT EXISTS(SELECT 1 FROM console_totp_credentials
+ WHERE user_id=$1 AND confirmed_at IS NOT NULL)`, userID).Scan(&result.TOTPEnabled); err != nil {
+ return MFAStatus{}, err
+ }
+ passkeys, err := s.ListPasskeys(ctx, userID)
+ if err != nil {
+ return MFAStatus{}, err
+ }
+ result.Passkeys = passkeys
+ return result, nil
+}
+
+func (s *Store) FindActiveConsoleUser(ctx context.Context, email string) (ConsoleActor, error) {
+ var actor ConsoleActor
+ err := s.db.QueryRow(ctx, `SELECT id::text,COALESCE(tenant_id::text,''),email,display_name,role
+ FROM console_users WHERE lower(email)=$1 AND status='active' AND email_verified_at IS NOT NULL`,
+ strings.ToLower(strings.TrimSpace(email))).Scan(&actor.ID, &actor.TenantID, &actor.Email, &actor.DisplayName, &actor.Role)
+ if errors.Is(err, pgx.ErrNoRows) {
+ return ConsoleActor{}, ErrConsoleUnauthorized
+ }
+ return actor, err
+}
+
+func (s *Store) VerifyConsolePassword(ctx context.Context, userID, password string) error {
+ var expectedHash, salt []byte
+ var iterations int
+ if err := s.db.QueryRow(ctx, `SELECT password_hash,password_salt,password_iterations FROM console_users
+ WHERE id=$1 AND status='active'`, userID).Scan(&expectedHash, &salt, &iterations); err != nil {
+ return ErrConsoleUnauthorized
+ }
+ if !security.VerifyPassword(password, expectedHash, salt, iterations) {
+ return ErrConsoleUnauthorized
+ }
+ return nil
+}
+
+func (s *Store) ResolveMFAChallenge(ctx context.Context, rawToken, remoteIP string) (string, string, error) {
+ hash := sha256.Sum256([]byte(strings.TrimSpace(rawToken)))
+ var id, userID, expectedIP string
+ err := s.db.QueryRow(ctx, `SELECT id::text,user_id::text,COALESCE(host(remote_ip),'')
+ FROM console_auth_challenges WHERE token_hash=$1 AND purpose='mfa_login'
+ AND consumed_at IS NULL AND expires_at>now()`, hash[:]).Scan(&id, &userID, &expectedIP)
+ if errors.Is(err, pgx.ErrNoRows) {
+ return "", "", ErrActionTokenInvalid
+ }
+ if err != nil {
+ return "", "", err
+ }
+ if expectedIP != "" && remoteIP != "" && expectedIP != remoteIP {
+ return "", "", ErrActionTokenInvalid
+ }
+ return id, userID, nil
+}
+
+func (s *Store) VerifyTOTPForUser(ctx context.Context, userID, code string) error {
+ tx, err := s.db.Begin(ctx)
+ if err != nil {
+ return err
+ }
+ defer tx.Rollback(ctx)
+ var ciphertext []byte
+ var lastStep int64
+ if err := tx.QueryRow(ctx, `SELECT secret_ciphertext,last_used_step FROM console_totp_credentials
+ WHERE user_id=$1 AND confirmed_at IS NOT NULL FOR UPDATE`, userID).Scan(&ciphertext, &lastStep); err != nil {
+ return ErrConsoleUnauthorized
+ }
+ secret, err := s.cipher.Decrypt(ciphertext)
+ if err != nil {
+ return err
+ }
+ step, ok := matchingTOTPStep(secret, strings.TrimSpace(code), time.Now().UTC())
+ if !ok || step <= lastStep {
+ return ErrConsoleUnauthorized
+ }
+ if _, err := tx.Exec(ctx, `UPDATE console_totp_credentials SET last_used_step=$2 WHERE user_id=$1`, userID, step); err != nil {
+ return err
+ }
+ return tx.Commit(ctx)
+}
+
+func (s *Store) BeginMFAChallenge(ctx context.Context, actor ConsoleActor, remoteIP, userAgent string) (AuthChallenge, error) {
+ methods, err := s.MFAMethods(ctx, actor.ID)
+ if err != nil {
+ return AuthChallenge{}, err
+ }
+ if len(methods) == 0 {
+ return AuthChallenge{}, ErrActionTokenInvalid
+ }
+ raw, hash, err := randomCredential("mfa-aigw-")
+ if err != nil {
+ return AuthChallenge{}, err
+ }
+ expires := time.Now().UTC().Add(10 * time.Minute)
+ _, err = s.db.Exec(ctx, `INSERT INTO console_auth_challenges
+ (user_id,token_hash,purpose,allowed_methods,expires_at,remote_ip,user_agent)
+ VALUES ($1,$2,'mfa_login',$3,$4,NULLIF($5,'')::inet,$6)`, actor.ID, hash, methods, expires, remoteIP, userAgent)
+ if err != nil {
+ return AuthChallenge{}, fmt.Errorf("create MFA challenge: %w", err)
+ }
+ return AuthChallenge{Token: raw, Methods: methods, ExpiresAt: expires}, nil
+}
+
+func (s *Store) BeginTOTP(ctx context.Context, actor ConsoleActor) (TOTPEnrollment, error) {
+ key, err := totp.Generate(totp.GenerateOpts{Issuer: "AIGW", AccountName: actor.Email, Period: 30, SecretSize: 20, Digits: otp.DigitsSix, Algorithm: otp.AlgorithmSHA1})
+ if err != nil {
+ return TOTPEnrollment{}, fmt.Errorf("generate TOTP secret: %w", err)
+ }
+ ciphertext, err := s.cipher.Encrypt(key.Secret())
+ if err != nil {
+ return TOTPEnrollment{}, err
+ }
+ result, err := s.db.Exec(ctx, `INSERT INTO console_totp_credentials (user_id,secret_ciphertext,confirmed_at,last_used_step)
+ VALUES ($1,$2,NULL,-1) ON CONFLICT (user_id) DO UPDATE SET secret_ciphertext=EXCLUDED.secret_ciphertext,
+ confirmed_at=NULL,last_used_step=-1,created_at=now()
+ WHERE console_totp_credentials.confirmed_at IS NULL`, actor.ID, ciphertext)
+ if err != nil {
+ return TOTPEnrollment{}, fmt.Errorf("store TOTP enrollment: %w", err)
+ }
+ if result.RowsAffected() == 0 {
+ return TOTPEnrollment{}, errors.New("TOTP is already enabled")
+ }
+ var image bytes.Buffer
+ qr, err := key.Image(256, 256)
+ if err == nil {
+ err = png.Encode(&image, qr)
+ }
+ if err != nil {
+ return TOTPEnrollment{}, fmt.Errorf("render TOTP QR: %w", err)
+ }
+ return TOTPEnrollment{Secret: key.Secret(), URI: key.URL(), QRCode: "data:image/png;base64," + base64.StdEncoding.EncodeToString(image.Bytes())}, nil
+}
+
+func (s *Store) ConfirmTOTP(ctx context.Context, actorID, code string) ([]string, error) {
+ code = strings.TrimSpace(code)
+ tx, err := s.db.Begin(ctx)
+ if err != nil {
+ return nil, err
+ }
+ defer tx.Rollback(ctx)
+ var ciphertext []byte
+ var confirmed *time.Time
+ if err := tx.QueryRow(ctx, `SELECT secret_ciphertext,confirmed_at FROM console_totp_credentials WHERE user_id=$1 FOR UPDATE`, actorID).Scan(&ciphertext, &confirmed); err != nil {
+ if errors.Is(err, pgx.ErrNoRows) {
+ return nil, errors.New("TOTP enrollment has not started")
+ }
+ return nil, err
+ }
+ secret, err := s.cipher.Decrypt(ciphertext)
+ if err != nil {
+ return nil, err
+ }
+ step, ok := matchingTOTPStep(secret, code, time.Now().UTC())
+ if !ok {
+ return nil, errors.New("invalid TOTP code")
+ }
+ if confirmed != nil {
+ return nil, errors.New("TOTP is already enabled")
+ }
+ if _, err := tx.Exec(ctx, `UPDATE console_totp_credentials SET confirmed_at=now(),last_used_step=$2 WHERE user_id=$1`, actorID, step); err != nil {
+ return nil, err
+ }
+ if _, err := tx.Exec(ctx, `DELETE FROM console_recovery_codes WHERE user_id=$1`, actorID); err != nil {
+ return nil, err
+ }
+ codes := make([]string, 0, 10)
+ for i := 0; i < 10; i++ {
+ code, hash, err := newRecoveryCode()
+ if err != nil {
+ return nil, err
+ }
+ if _, err := tx.Exec(ctx, `INSERT INTO console_recovery_codes (user_id,code_hash) VALUES ($1,$2)`, actorID, hash); err != nil {
+ return nil, err
+ }
+ codes = append(codes, code)
+ }
+ if err := tx.Commit(ctx); err != nil {
+ return nil, err
+ }
+ return codes, nil
+}
+
+func (s *Store) DisableTOTP(ctx context.Context, actorID string) error {
+ tx, err := s.db.Begin(ctx)
+ if err != nil {
+ return err
+ }
+ defer tx.Rollback(ctx)
+ if _, err := tx.Exec(ctx, `DELETE FROM console_totp_credentials WHERE user_id=$1`, actorID); err != nil {
+ return err
+ }
+ if _, err := tx.Exec(ctx, `DELETE FROM console_recovery_codes WHERE user_id=$1`, actorID); err != nil {
+ return err
+ }
+ return tx.Commit(ctx)
+}
+
+func (s *Store) CompleteMFAChallenge(ctx context.Context, input MFACodeInput, remoteIP string) (ConsoleActor, string, error) {
+ code := strings.TrimSpace(input.Code)
+ tokenHash := sha256.Sum256([]byte(strings.TrimSpace(input.ChallengeToken)))
+ tx, err := s.db.Begin(ctx)
+ if err != nil {
+ return ConsoleActor{}, "", err
+ }
+ defer tx.Rollback(ctx)
+ var challengeID, userID, expectedIP string
+ var allowed []string
+ err = tx.QueryRow(ctx, `SELECT c.id::text,c.user_id::text,COALESCE(host(c.remote_ip),''),c.allowed_methods
+ FROM console_auth_challenges c WHERE c.token_hash=$1 AND c.purpose='mfa_login'
+ AND c.consumed_at IS NULL AND c.expires_at>now() FOR UPDATE`, tokenHash[:]).Scan(&challengeID, &userID, &expectedIP, &allowed)
+ if errors.Is(err, pgx.ErrNoRows) {
+ return ConsoleActor{}, "", ErrActionTokenInvalid
+ }
+ if err != nil {
+ return ConsoleActor{}, "", err
+ }
+ if expectedIP != "" && remoteIP != "" && expectedIP != remoteIP {
+ return ConsoleActor{}, "", ErrActionTokenInvalid
+ }
+ var method string
+ if method = "totp"; len(code) != 6 {
+ method = "recovery"
+ }
+ valid := false
+ if method == "totp" {
+ var ciphertext []byte
+ var lastStep int64
+ if err := tx.QueryRow(ctx, `SELECT secret_ciphertext,last_used_step FROM console_totp_credentials WHERE user_id=$1 AND confirmed_at IS NOT NULL FOR UPDATE`, userID).Scan(&ciphertext, &lastStep); err == nil {
+ secret, decryptErr := s.cipher.Decrypt(ciphertext)
+ if decryptErr == nil {
+ if step, ok := matchingTOTPStep(secret, code, time.Now().UTC()); ok && step > lastStep {
+ valid = true
+ if _, err := tx.Exec(ctx, `UPDATE console_totp_credentials SET last_used_step=$2 WHERE user_id=$1`, userID, step); err != nil {
+ return ConsoleActor{}, "", err
+ }
+ }
+ }
+ }
+ }
+ if !valid && method == "recovery" {
+ hash := sha256.Sum256([]byte(normalizeRecoveryCode(code)))
+ result, updateErr := tx.Exec(ctx, `UPDATE console_recovery_codes SET used_at=now()
+ WHERE user_id=$1 AND code_hash=$2 AND used_at IS NULL`, userID, hash[:])
+ valid = updateErr == nil && result.RowsAffected() == 1
+ }
+ if !valid {
+ _, _ = tx.Exec(ctx, `UPDATE console_auth_challenges SET attempts=attempts+1,
+ consumed_at=CASE WHEN attempts+1 >= 5 THEN now() ELSE consumed_at END WHERE id=$1`, challengeID)
+ if err := tx.Commit(ctx); err != nil {
+ return ConsoleActor{}, "", err
+ }
+ return ConsoleActor{}, "", ErrConsoleUnauthorized
+ }
+ if _, err := tx.Exec(ctx, `UPDATE console_auth_challenges SET consumed_at=now() WHERE id=$1`, challengeID); err != nil {
+ return ConsoleActor{}, "", err
+ }
+ var actor ConsoleActor
+ if err := tx.QueryRow(ctx, `SELECT id::text,COALESCE(tenant_id::text,''),email,display_name,role FROM console_users WHERE id=$1 AND status='active'`, userID).
+ Scan(&actor.ID, &actor.TenantID, &actor.Email, &actor.DisplayName, &actor.Role); err != nil {
+ return ConsoleActor{}, "", err
+ }
+ if err := tx.Commit(ctx); err != nil {
+ return ConsoleActor{}, "", err
+ }
+ return actor, method, nil
+}
+
+func matchingTOTPStep(secret, code string, now time.Time) (int64, bool) {
+ if len(strings.TrimSpace(code)) != 6 {
+ return 0, false
+ }
+ current := now.Unix() / 30
+ opts := totp.ValidateOpts{Period: 30, Skew: 0, Digits: otp.DigitsSix, Algorithm: otp.AlgorithmSHA1}
+ for _, step := range []int64{current - 1, current, current + 1} {
+ generated, err := totp.GenerateCodeCustom(secret, time.Unix(step*30, 0), opts)
+ if err == nil && generated == code {
+ return step, true
+ }
+ }
+ return 0, false
+}
+
+func newRecoveryCode() (string, []byte, error) {
+ raw, _, err := randomCredential("")
+ if err != nil {
+ return "", nil, err
+ }
+ value := strings.ToUpper(strings.TrimRight(raw, "="))
+ if len(value) > 10 {
+ value = value[:10]
+ }
+ if len(value) < 10 {
+ return newRecoveryCode()
+ }
+ code := value[:5] + "-" + value[5:]
+ hash := sha256.Sum256([]byte(normalizeRecoveryCode(code)))
+ return code, hash[:], nil
+}
+
+func normalizeRecoveryCode(value string) string {
+ return strings.ToUpper(strings.ReplaceAll(strings.TrimSpace(value), "-", ""))
+}