package app import ( "bytes" "context" "crypto/aes" "crypto/cipher" "crypto/rand" "crypto/sha256" "crypto/subtle" "database/sql" "encoding/base64" "encoding/hex" "fmt" "net/http" "strconv" "strings" "time" "github.com/golang-jwt/jwt/v5" "golang.org/x/crypto/argon2" "golang.org/x/crypto/bcrypt" ) type identity struct { ID int64 Role string Name string IssuedAt time.Time Version int } type identityKey struct{} type tokenClaims struct { Role string `json:"role"` Name string `json:"name"` Version int `json:"ver,omitempty"` jwt.RegisteredClaims } func (a *App) token(id int64, role, name string, ttl time.Duration) (string, error) { now := time.Now() version := 0 if role == "user" { _ = a.db.QueryRow(`SELECT token_version FROM user_security_controls WHERE user_id=?`, id).Scan(&version) } else if role == "admin" { _ = a.db.QueryRow(`SELECT token_version FROM admin_users WHERE id=?`, id).Scan(&version) } claims := tokenClaims{ Role: role, Name: name, Version: version, RegisteredClaims: jwt.RegisteredClaims{ Subject: strconv.FormatInt(id, 10), IssuedAt: jwt.NewNumericDate(now), ExpiresAt: jwt.NewNumericDate(now.Add(ttl)), }, } return jwt.NewWithClaims(jwt.SigningMethodHS256, claims).SignedString([]byte(a.config.JWTSecret)) } func (a *App) parseToken(raw string) (identity, error) { parsed, err := jwt.ParseWithClaims(raw, &tokenClaims{}, func(t *jwt.Token) (any, error) { if t.Method != jwt.SigningMethodHS256 { return nil, fmt.Errorf("unexpected signing method") } return []byte(a.config.JWTSecret), nil }) if err != nil || !parsed.Valid { return identity{}, fmt.Errorf("invalid token") } claims, ok := parsed.Claims.(*tokenClaims) if !ok { return identity{}, fmt.Errorf("invalid claims") } id, err := strconv.ParseInt(claims.Subject, 10, 64) if err != nil { return identity{}, fmt.Errorf("invalid subject") } issuedAt := time.Time{} if claims.IssuedAt != nil { issuedAt = claims.IssuedAt.Time } return identity{ID: id, Role: claims.Role, Name: claims.Name, IssuedAt: issuedAt, Version: claims.Version}, nil } func (a *App) requireAuth(role string, next http.HandlerFunc) http.HandlerFunc { return func(w http.ResponseWriter, r *http.Request) { raw := strings.TrimSpace(strings.TrimPrefix(r.Header.Get("Authorization"), "Bearer ")) if raw == "" { fail(w, http.StatusUnauthorized, 10001, "请先登录") return } who, err := a.parseToken(raw) if err != nil || who.Role != role { fail(w, http.StatusUnauthorized, 10001, "登录状态已失效") return } if role == "user" { if a.executeDueAccountClosure(r, who.ID) { fail(w, http.StatusUnauthorized, 10001, "账号已按注销申请完成删除") return } var status int if err = a.db.QueryRowContext(r.Context(), `SELECT status FROM users WHERE id=? AND deleted_at IS NULL`, who.ID).Scan(&status); err == nil { status = a.normalizeUserStatus(r.Context(), who.ID, status) } if err != nil || status != 1 { fail(w, http.StatusForbidden, 10006, "账号已被冻结或封禁") return } var forceLogout sql.NullTime var tokenVersion int _ = a.db.QueryRowContext(r.Context(), `SELECT force_logout_at,token_version FROM user_security_controls WHERE user_id=?`, who.ID).Scan(&forceLogout, &tokenVersion) if who.Version != tokenVersion || (forceLogout.Valid && !who.IssuedAt.IsZero() && who.IssuedAt.Unix() < forceLogout.Time.Unix()) { fail(w, http.StatusUnauthorized, 10001, "登录状态已失效,请重新登录") return } a.touchUserActivity(r.Context(), who.ID) } else if role == "admin" { var status, tokenVersion int if err = a.db.QueryRowContext(r.Context(), `SELECT status,token_version FROM admin_users WHERE id=?`, who.ID).Scan(&status, &tokenVersion); err != nil || status != 1 { fail(w, http.StatusForbidden, 10006, "管理员账号已停用") return } if who.Version != tokenVersion { fail(w, http.StatusUnauthorized, 10001, "密码已修改,请重新登录") return } } next(w, r.WithContext(context.WithValue(r.Context(), identityKey{}, who))) } } // Keep the persisted activity timestamp fresh without rewriting the row for // every API call or every 25-second WebSocket heartbeat. func (a *App) touchUserActivity(ctx context.Context, userID int64) { _, _ = a.db.ExecContext(ctx, `UPDATE user_profiles SET last_active_at=NOW(3) WHERE user_id=? AND (last_active_at IS NULL OR last_active_at 72 { salt := make([]byte, 16) if _, err := rand.Read(salt); err != nil { return "", err } key := argon2.IDKey([]byte(password), salt, 2, 19*1024, 1, 32) return longPasswordHashPrefix + base64.RawStdEncoding.EncodeToString(salt) + "$" + base64.RawStdEncoding.EncodeToString(key), nil } hash, err := bcrypt.GenerateFromPassword([]byte(password), bcrypt.DefaultCost) return string(hash), err } func checkPassword(hash, password string) bool { if strings.HasPrefix(hash, "$argon2id$") { // Only our supported parameters are accepted; never allocate memory from // unchecked parameters in a corrupt or externally supplied hash. payload, ok := strings.CutPrefix(hash, longPasswordHashPrefix) if !ok || len(payload) != 66 { return false } parts := strings.Split(payload, "$") if len(parts) != 2 { return false } salt, err := base64.RawStdEncoding.DecodeString(parts[0]) if err != nil || len(salt) != 16 { return false } expected, err := base64.RawStdEncoding.DecodeString(parts[1]) if err != nil || len(expected) != 32 { return false } key := argon2.IDKey([]byte(password), salt, 2, 19*1024, 1, 32) return subtle.ConstantTimeCompare(key, expected) == 1 } if len(password) > 72 { return false } return bcrypt.CompareHashAndPassword([]byte(hash), []byte(password)) == nil } // 8 characters and not all digits. Phone verification used to be the barrier // to junk accounts; with the SMS switch off the password is what is left, and // a 4-digit password is not one. Only setting a password is checked — existing // accounts keep signing in with whatever they already have. const passwordRule = "密码至少 8 位,且不能全是数字" func validUserPassword(password string) bool { // Both bounds count characters. Counting the upper one in bytes would refuse // a 43-character Chinese password while accepting a 128-character English one. if length := len([]rune(password)); length < 8 || length > 128 { return false } for _, character := range password { if character < '0' || character > '9' { return true } } return false } func validPhone(value string) bool { value = strings.TrimSpace(value) if len(value) != 11 || value[0] != '1' || value[1] < '3' || value[1] > '9' { return false } for _, char := range value { if char < '0' || char > '9' { return false } } return true } func phoneHash(phone string) []byte { sum := sha256.Sum256([]byte(strings.TrimSpace(phone))) return sum[:] } var encryptedPhonePrefix = []byte("enc:v1:") func (a *App) piiEncryptionKey() [32]byte { key := a.config.ConfigEncryptionKey if key == "" { key = a.config.JWTSecret + ":development-pii" } return sha256.Sum256([]byte(key + ":phone")) } func (a *App) encryptPhone(phone string) ([]byte, error) { key := a.piiEncryptionKey() 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 } payload := append(append([]byte{}, encryptedPhonePrefix...), nonce...) payload = gcm.Seal(payload, nonce, []byte(strings.TrimSpace(phone)), nil) return payload, nil } func (a *App) decryptPhone(payload []byte) (string, error) { if !bytes.HasPrefix(payload, encryptedPhonePrefix) { return string(payload), nil } key := a.piiEncryptionKey() block, err := aes.NewCipher(key[:]) if err != nil { return "", err } gcm, err := cipher.NewGCM(block) if err != nil { return "", err } ciphertext := payload[len(encryptedPhonePrefix):] if len(ciphertext) < gcm.NonceSize() { return "", fmt.Errorf("invalid encrypted phone") } plain, err := gcm.Open(nil, ciphertext[:gcm.NonceSize()], ciphertext[gcm.NonceSize():], nil) return string(plain), err } func (a *App) encryptLegacyPhones(ctx context.Context) error { rows, err := a.db.QueryContext(ctx, `SELECT id,phone_cipher FROM users WHERE phone_cipher IS NOT NULL AND LEFT(phone_cipher,7)<>?`, encryptedPhonePrefix) if err != nil { return err } type legacyPhone struct { id int64 payload []byte } items := []legacyPhone{} for rows.Next() { var item legacyPhone if err = rows.Scan(&item.id, &item.payload); err != nil { _ = rows.Close() return err } items = append(items, item) } if err = rows.Close(); err != nil { return err } for _, item := range items { encrypted, encryptErr := a.encryptPhone(string(item.payload)) if encryptErr != nil { return encryptErr } if _, err = a.db.ExecContext(ctx, `UPDATE users SET phone_cipher=? WHERE id=? AND phone_cipher=?`, encrypted, item.id, item.payload); err != nil { return err } } return nil } func randomToken() string { buf := make([]byte, 32) _, _ = rand.Read(buf) return hex.EncodeToString(buf) }