package service import ( "crypto/rand" "encoding/base64" "errors" "time" "github.com/freefire/jiang13-bbs/model" "github.com/golang-jwt/jwt/v5" "golang.org/x/crypto/bcrypt" "gorm.io/gorm" ) const ( RoleUser = "user" RoleAdmin = "admin" // CookieName JWT access token 存储的 HttpOnly cookie 名 CookieName = "j13_token" // RefreshCookieName refresh token 存储的 HttpOnly cookie 名 RefreshCookieName = "j13_refresh" // CSRFCookieName CSRF token cookie 名(非 HttpOnly,前端可读) CSRFCookieName = "j13_csrf" // CSRFHeaderName 前端传递 CSRF token 的 header 名 CSRFHeaderName = "X-CSRF-Token" // AccessTokenTTL access token 有效期(短期,降低被盗窗口) AccessTokenTTL = 15 * time.Minute // RefreshTokenTTL refresh token 有效期(长期) RefreshTokenTTL = 7 * 24 * time.Hour ) // UserClaims JWT 中携带的用户信息 type UserClaims struct { ID uint `json:"id"` Username string `json:"username"` Role string `json:"role"` Banned bool `json:"banned"` TokenVersion int `json:"tv"` // token 版本号,用于服务端撤销 } // AuthService 认证服务 type AuthService struct { db *gorm.DB jwtSecret []byte } func NewAuthService(db *gorm.DB, jwtSecret string) *AuthService { return &AuthService{db: db, jwtSecret: []byte(jwtSecret)} } // Register 用户注册 func (s *AuthService) Register(username, email, password string) (*model.User, error) { // 检查用户名是否已存在 var count int64 s.db.Model(&model.User{}).Where("username = ?", username).Count(&count) if count > 0 { return nil, errors.New("用户名已被使用") } hashed, err := bcrypt.GenerateFromPassword([]byte(password), bcrypt.DefaultCost) if err != nil { return nil, err } user := &model.User{ Username: username, Email: email, Password: string(hashed), Nickname: username, Role: model.RoleUser, } if err := s.db.Create(user).Error; err != nil { return nil, err } return user, nil } // Login 用户登录,返回 access token + refresh token + user func (s *AuthService) Login(username, password string) (string, string, *model.User, error) { var user model.User if err := s.db.Where("username = ?", username).First(&user).Error; err != nil { return "", "", nil, errors.New("用户名或密码错误") } if user.Banned { return "", "", nil, errors.New("账号已被封禁") } if err := bcrypt.CompareHashAndPassword([]byte(user.Password), []byte(password)); err != nil { return "", "", nil, errors.New("用户名或密码错误") } accessToken, err := s.generateToken(&user) if err != nil { return "", "", nil, err } refreshToken, err := s.CreateRefreshToken(user.ID) if err != nil { return "", "", nil, err } return accessToken, refreshToken, &user, nil } // generateToken 签发 access token(短期) func (s *AuthService) generateToken(user *model.User) (string, error) { claims := &tokenClaims{ UserClaims: UserClaims{ ID: user.ID, Username: user.Username, Role: string(user.Role), Banned: user.Banned, TokenVersion: user.TokenVersion, }, RegisteredClaims: jwt.RegisteredClaims{ Subject: user.Username, ExpiresAt: jwt.NewNumericDate(time.Now().Add(AccessTokenTTL)), IssuedAt: jwt.NewNumericDate(time.Now()), }, } token := jwt.NewWithClaims(jwt.SigningMethodHS256, claims) return token.SignedString(s.jwtSecret) } type tokenClaims struct { UserClaims jwt.RegisteredClaims } // ParseToken 解析 JWT(仅校验签名和过期,不查 DB) func (s *AuthService) ParseToken(tokenStr string) (*UserClaims, error) { claims := &tokenClaims{} _, err := jwt.ParseWithClaims(tokenStr, claims, func(t *jwt.Token) (interface{}, error) { return s.jwtSecret, nil }) if err != nil { return nil, err } return &claims.UserClaims, nil } // ValidateClaims 校验 claims 是否仍然有效(查 DB:token_version 匹配且未封禁) // 用于中间件在每个请求上做实时权限校验 func (s *AuthService) ValidateClaims(claims *UserClaims) (*model.User, error) { var user model.User if err := s.db.First(&user, claims.ID).Error; err != nil { return nil, errors.New("用户不存在") } // token 版本不匹配 → 已被撤销(改密码/封禁/管理员操作) if user.TokenVersion != claims.TokenVersion { return nil, errors.New("token 已失效") } // 实时校验封禁状态(不依赖 JWT 中的缓存值) if user.Banned { return nil, errors.New("账号已被封禁") } return &user, nil } // GetUserByID 根据 ID 获取用户 func (s *AuthService) GetUserByID(id uint) (*model.User, error) { var user model.User if err := s.db.First(&user, id).Error; err != nil { return nil, err } return &user, nil } // GenerateCSRFToken 生成随机 CSRF token func GenerateCSRFToken() string { b := make([]byte, 32) _, _ = rand.Read(b) return base64.URLEncoding.EncodeToString(b) } // generateRandomToken 生成随机令牌字符串 func generateRandomToken() string { b := make([]byte, 32) _, _ = rand.Read(b) return base64.URLEncoding.EncodeToString(b) } // CreateRefreshToken 创建并存储 refresh token func (s *AuthService) CreateRefreshToken(userID uint) (string, error) { token := generateRandomToken() rt := &model.RefreshToken{ UserID: userID, Token: token, ExpiresAt: time.Now().Add(RefreshTokenTTL), } if err := s.db.Create(rt).Error; err != nil { return "", err } return token, nil } // ValidateRefreshToken 校验 refresh token 并返回所属用户 func (s *AuthService) ValidateRefreshToken(token string) (*model.User, error) { var rt model.RefreshToken if err := s.db.Where("token = ?", token).First(&rt).Error; err != nil { return nil, errors.New("refresh token 无效") } if rt.Revoked { return nil, errors.New("refresh token 已撤销") } if time.Now().After(rt.ExpiresAt) { return nil, errors.New("refresh token 已过期") } var user model.User if err := s.db.First(&user, rt.UserID).Error; err != nil { return nil, errors.New("用户不存在") } if user.Banned { return nil, errors.New("账号已被封禁") } return &user, nil } // RotateRefreshToken 轮转 refresh token:撤销旧的,签发新的 func (s *AuthService) RotateRefreshToken(oldToken string) (string, string, *model.User, error) { user, err := s.ValidateRefreshToken(oldToken) if err != nil { return "", "", nil, err } // 撤销旧 token s.db.Model(&model.RefreshToken{}).Where("token = ?", oldToken).Update("revoked", true) // 签发新 access + refresh accessToken, err := s.generateToken(user) if err != nil { return "", "", nil, err } newRefresh, err := s.CreateRefreshToken(user.ID) if err != nil { return "", "", nil, err } return accessToken, newRefresh, user, nil } // RevokeRefreshToken 撤销单个 refresh token(登出时用) func (s *AuthService) RevokeRefreshToken(token string) { s.db.Model(&model.RefreshToken{}).Where("token = ?", token).Update("revoked", true) } // RevokeAllUserRefreshTokens 撤销用户所有 refresh token(改密码/封禁时用) func (s *AuthService) RevokeAllUserRefreshTokens(userID uint) { s.db.Model(&model.RefreshToken{}).Where("user_id = ?", userID).Update("revoked", true) } // IncrementTokenVersion 递增用户 token 版本,使所有已有 JWT 失效 // 用于:改密码、封禁用户、管理员强制下线 func (s *AuthService) IncrementTokenVersion(userID uint) error { result := s.db.Model(&model.User{}).Where("id = ?", userID).UpdateColumn("token_version", gorm.Expr("token_version + 1")) if result.Error != nil { return result.Error } s.RevokeAllUserRefreshTokens(userID) return nil } // ChangePassword 修改密码:校验旧密码,更新新密码,递增 token_version 使旧 token 失效 func (s *AuthService) ChangePassword(userID uint, oldPassword, newPassword string) error { var user model.User if err := s.db.First(&user, userID).Error; err != nil { return errors.New("用户不存在") } if err := bcrypt.CompareHashAndPassword([]byte(user.Password), []byte(oldPassword)); err != nil { return errors.New("旧密码错误") } if len(newPassword) < 6 { return errors.New("新密码至少 6 位") } hashed, err := bcrypt.GenerateFromPassword([]byte(newPassword), bcrypt.DefaultCost) if err != nil { return err } result := s.db.Model(&model.User{}).Where("id = ?", userID).Update("password", string(hashed)) if result.Error != nil { return result.Error } if result.RowsAffected == 0 { return errors.New("密码更新失败") } // 递增 token_version,使所有旧 JWT 和 refresh token 失效 return s.IncrementTokenVersion(userID) }