diff options
| author | Chia <Chia@93.nz> | 2026-08-05 14:48:00 +1200 |
|---|---|---|
| committer | Chia <Chia@93.nz> | 2026-08-05 14:48:00 +1200 |
| commit | cd0dd91ab93653631904f2ea0e574ccde6d60339 (patch) | |
| tree | c65417b880a3f4a35c504c44edae821bc2122f70 /internal/controlplane/mfa.go | |
| parent | 86b1f42e3c5601ff10621a9779cf0076590797a1 (diff) | |
add passkey, totp.
Diffstat (limited to 'internal/controlplane/mfa.go')
| -rw-r--r-- | internal/controlplane/mfa.go | 344 |
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), "-", "")) +} |
