Files
pangolin/server/internal/auth/totp_user.go
T
wangjia 0782cf651b feat(server): web 用户中心后端(3/3) — 用户 TOTP 2FA + 登录二段式
- 用户 TOTP(auth/totp_user.go,复用 internal/totp + AES-256-GCM 加密存密钥):
  POST /v1/me/totp/setup(生成密钥+otpauth_uri)、/verify(校验码→启用)、
  /disable(校验码→清空)。仅在 USER_TOTP_ENC_KEY(32B/64hex) 配置时挂载。
- 登录二段式:Login 在 totp_enabled 时不发 token,改发短期 pending token(Redis
  5min)+ 返回 {totp_required, pending_token};POST /v1/auth/login/totp 消费
  pending + 校验码 → 发 token。非 TOTP 用户仍走扁平 TokenPair,app 不受影响。
- User 结构 + GetUserByEmail 补 totp_enabled。
- 单测覆盖 AES seal/open 往返 + 篡改/错误密钥检测 + pending token 唯一性。
- 全量 server 23 包测试通过。

Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>
2026-06-17 08:45:50 +08:00

252 lines
7.8 KiB
Go

package auth
import (
"context"
"crypto/aes"
"crypto/cipher"
"crypto/rand"
"database/sql"
"encoding/hex"
"encoding/json"
"errors"
"net/http"
"strconv"
"time"
"github.com/redis/go-redis/v9"
"github.com/wangjia/pangolin/server/internal/apierr"
"github.com/wangjia/pangolin/server/internal/totp"
)
const totpIssuer = "Pangolin"
// errTOTPInvalid mirrors the web user-center's expected error code.
var errTOTPInvalid = apierr.New("totp_invalid", "动态码不正确,请重试", "Invalid code, try again")
// TOTPHandler serves the user two-factor endpoints (/v1/me/totp/*) and the
// second login step (/v1/auth/login/totp). Secrets are stored AES-256-GCM
// encrypted under encKey; pending login tokens live in Redis.
type TOTPHandler struct {
db *sql.DB
encKey []byte // 32 bytes (AES-256)
tokens *TokenManager
rdb *redis.Client
now func() time.Time
}
// NewTOTPHandler wires the handler. encKey must be 32 bytes.
func NewTOTPHandler(db *sql.DB, encKey []byte, tokens *TokenManager, rdb *redis.Client) *TOTPHandler {
return &TOTPHandler{db: db, encKey: encKey, tokens: tokens, rdb: rdb, now: time.Now}
}
type totpCodeRequest struct {
Code string `json:"code"`
}
// Setup handles POST /v1/me/totp/setup — generates and stores a secret (not yet
// enabled) and returns it plus an otpauth:// URI for QR rendering.
func (h *TOTPHandler) Setup(w http.ResponseWriter, r *http.Request) {
uid, ok := UserIDFromContext(r.Context())
if !ok {
apierr.WriteJSON(w, http.StatusUnauthorized, apierr.ErrUnauthorized)
return
}
var email string
if err := h.db.QueryRowContext(r.Context(),
`SELECT email FROM users WHERE id = ? AND status = 'active'`, uid).Scan(&email); err != nil {
apierr.WriteJSON(w, http.StatusInternalServerError, apierr.ErrInternal)
return
}
secret, err := totp.GenerateSecret()
if err != nil {
apierr.WriteJSON(w, http.StatusInternalServerError, apierr.ErrInternal)
return
}
enc, err := sealAES(h.encKey, secret)
if err != nil {
apierr.WriteJSON(w, http.StatusInternalServerError, apierr.ErrInternal)
return
}
if _, err := h.db.ExecContext(r.Context(),
`UPDATE users SET totp_secret_enc = ?, totp_enabled = FALSE WHERE id = ?`, enc, uid); err != nil {
apierr.WriteJSON(w, http.StatusInternalServerError, apierr.ErrInternal)
return
}
w.Header().Set("Content-Type", "application/json; charset=utf-8")
_ = json.NewEncoder(w).Encode(map[string]string{
"secret": secret,
"otpauth_uri": totp.ProvisioningURI(secret, email, totpIssuer),
})
}
// Verify handles POST /v1/me/totp/verify {code} — enables TOTP after confirming
// the user can produce a valid code for the secret created by Setup.
func (h *TOTPHandler) Verify(w http.ResponseWriter, r *http.Request) {
uid, ok := UserIDFromContext(r.Context())
if !ok {
apierr.WriteJSON(w, http.StatusUnauthorized, apierr.ErrUnauthorized)
return
}
var req totpCodeRequest
if !decodeJSON(w, r, &req) {
return
}
secret, _, err := h.loadSecret(r.Context(), uid)
if err != nil || secret == "" || !totp.Validate(secret, req.Code, h.now().UTC(), 1) {
apierr.WriteJSON(w, http.StatusUnauthorized, errTOTPInvalid)
return
}
if _, err := h.db.ExecContext(r.Context(),
`UPDATE users SET totp_enabled = TRUE WHERE id = ?`, uid); err != nil {
apierr.WriteJSON(w, http.StatusInternalServerError, apierr.ErrInternal)
return
}
w.WriteHeader(http.StatusNoContent)
}
// Disable handles POST /v1/me/totp/disable {code} — turns off TOTP after a valid
// code and clears the stored secret.
func (h *TOTPHandler) Disable(w http.ResponseWriter, r *http.Request) {
uid, ok := UserIDFromContext(r.Context())
if !ok {
apierr.WriteJSON(w, http.StatusUnauthorized, apierr.ErrUnauthorized)
return
}
var req totpCodeRequest
if !decodeJSON(w, r, &req) {
return
}
secret, _, err := h.loadSecret(r.Context(), uid)
if err != nil || secret == "" || !totp.Validate(secret, req.Code, h.now().UTC(), 1) {
apierr.WriteJSON(w, http.StatusUnauthorized, errTOTPInvalid)
return
}
if _, err := h.db.ExecContext(r.Context(),
`UPDATE users SET totp_secret_enc = NULL, totp_enabled = FALSE WHERE id = ?`, uid); err != nil {
apierr.WriteJSON(w, http.StatusInternalServerError, apierr.ErrInternal)
return
}
w.WriteHeader(http.StatusNoContent)
}
type loginTOTPRequest struct {
PendingToken string `json:"pending_token"`
Code string `json:"code"`
}
// LoginTOTP handles POST /v1/auth/login/totp — second login step. Consumes the
// one-time pending token from Login, validates the TOTP code, then issues tokens.
func (h *TOTPHandler) LoginTOTP(w http.ResponseWriter, r *http.Request) {
var req loginTOTPRequest
if !decodeJSON(w, r, &req) {
return
}
if req.PendingToken == "" {
apierr.WriteJSON(w, http.StatusUnauthorized, errTOTPInvalid)
return
}
// One-time consume of the pending token.
uidStr, err := h.rdb.GetDel(r.Context(), totpPendingPrefix+req.PendingToken).Result()
if errors.Is(err, redis.Nil) {
apierr.WriteJSON(w, http.StatusUnauthorized, errTOTPInvalid)
return
} else if err != nil {
apierr.WriteJSON(w, http.StatusInternalServerError, apierr.ErrInternal)
return
}
uid, err := strconv.ParseInt(uidStr, 10, 64)
if err != nil {
apierr.WriteJSON(w, http.StatusInternalServerError, apierr.ErrInternal)
return
}
secret, enabled, err := h.loadSecret(r.Context(), uid)
if err != nil || !enabled || secret == "" || !totp.Validate(secret, req.Code, h.now().UTC(), 1) {
apierr.WriteJSON(w, http.StatusUnauthorized, errTOTPInvalid)
return
}
var uuid string
if err := h.db.QueryRowContext(r.Context(),
`SELECT uuid FROM users WHERE id = ? AND status = 'active'`, uid).Scan(&uuid); err != nil {
apierr.WriteJSON(w, http.StatusInternalServerError, apierr.ErrInternal)
return
}
pair, err := h.tokens.Issue(r.Context(), uid, uuid)
if err != nil {
apierr.WriteJSON(w, http.StatusInternalServerError, apierr.ErrInternal)
return
}
writeTokenPair(w, pair)
}
// loadSecret returns the user's decrypted TOTP secret and enabled flag. A NULL
// secret yields ("", enabled, nil).
func (h *TOTPHandler) loadSecret(ctx context.Context, uid int64) (string, bool, error) {
var enc []byte
var enabled bool
if err := h.db.QueryRowContext(ctx,
`SELECT totp_secret_enc, totp_enabled FROM users WHERE id = ?`, uid).Scan(&enc, &enabled); err != nil {
return "", false, err
}
if len(enc) == 0 {
return "", enabled, nil
}
secret, err := openAES(h.encKey, enc)
if err != nil {
return "", enabled, err
}
return secret, enabled, nil
}
// ── crypto / token helpers ───────────────────────────────────────────────────
// sealAES encrypts plaintext with AES-256-GCM under key (32 bytes); the nonce is
// prepended to the ciphertext.
func sealAES(key []byte, plaintext string) ([]byte, error) {
block, err := aes.NewCipher(key)
if err != nil {
return nil, err
}
gcm, err := cipher.NewGCM(block)
if err != nil {
return nil, err
}
nonce := make([]byte, gcm.NonceSize())
if _, err := rand.Read(nonce); err != nil {
return nil, err
}
return gcm.Seal(nonce, nonce, []byte(plaintext), nil), nil
}
// openAES reverses sealAES.
func openAES(key, blob []byte) (string, error) {
block, err := aes.NewCipher(key)
if err != nil {
return "", err
}
gcm, err := cipher.NewGCM(block)
if err != nil {
return "", err
}
if len(blob) < gcm.NonceSize() {
return "", errors.New("auth: totp ciphertext too short")
}
nonce, ct := blob[:gcm.NonceSize()], blob[gcm.NonceSize():]
plaintext, err := gcm.Open(nil, nonce, ct, nil)
if err != nil {
return "", err
}
return string(plaintext), nil
}
// newToken16 returns a 32-char hex random token (used for TOTP pending logins).
func newToken16() (string, error) {
var b [16]byte
if _, err := rand.Read(b[:]); err != nil {
return "", err
}
return hex.EncodeToString(b[:]), nil
}