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), "-", "")) }