package auth import ( "errors" "net/http" "strings" "time" "github.com/gin-gonic/gin" "github.com/golang-jwt/jwt/v5" "golang.org/x/crypto/bcrypt" "frpc-console/internal/db" ) var jwtSecretCache []byte // getJwtSecret 从数据库获取 JWT 密钥 func getJwtSecret() []byte { if len(jwtSecretCache) > 0 { return jwtSecretCache } secret, err := db.GetJwtSecret() if err != nil { // 如果数据库还没有密钥,生成一个默认的(仅用于开发) // 生产环境应该通过数据库初始化时生成 return []byte("frpc-console-default-secret-key-2024") } jwtSecretCache = []byte(secret) return jwtSecretCache } // Claims JWT 声明 type Claims struct { Username string `json:"username"` jwt.RegisteredClaims } // GenerateJWT 生成 JWT func GenerateJWT(username string) (string, error) { claims := Claims{ Username: username, RegisteredClaims: jwt.RegisteredClaims{ ExpiresAt: jwt.NewNumericDate(time.Now().Add(7 * 24 * time.Hour)), IssuedAt: jwt.NewNumericDate(time.Now()), }, } token := jwt.NewWithClaims(jwt.SigningMethodHS256, claims) return token.SignedString(getJwtSecret()) } // ParseJWT 解析 JWT func ParseJWT(tokenString string) (*Claims, error) { token, err := jwt.ParseWithClaims(tokenString, &Claims{}, func(token *jwt.Token) (interface{}, error) { return getJwtSecret(), nil }) if err != nil { return nil, err } if claims, ok := token.Claims.(*Claims); ok && token.Valid { return claims, nil } return nil, errors.New("invalid token") } // ValidatePassword 校验密码复杂度 // 要求: 至少8位,包含大小写字母、数字、特殊字符 func ValidatePassword(pwd string) bool { if len(pwd) < 8 { return false } var hasUpper, hasLower, hasDigit, hasSpecial bool for _, ch := range pwd { switch { case ch >= 'A' && ch <= 'Z': hasUpper = true case ch >= 'a' && ch <= 'z': hasLower = true case ch >= '0' && ch <= '9': hasDigit = true case strings.ContainsAny(string(ch), "!@#$%^&*()_+-=[]{}|;:,.<>?"): hasSpecial = true } } return hasUpper && hasLower && hasDigit && hasSpecial } // HashPassword 加密密码 func HashPassword(password string) (string, error) { hash, err := bcrypt.GenerateFromPassword([]byte(password), bcrypt.DefaultCost) if err != nil { return "", err } return string(hash), nil } // CheckPassword 验证密码 func CheckPassword(password, hash string) bool { err := bcrypt.CompareHashAndPassword([]byte(hash), []byte(password)) return err == nil } // ================================================================ // 登录 / 注册 / 修改密码 的业务逻辑 (供 handler 调用) // ================================================================ // LoginRequest 登录请求 type LoginRequest struct { Username string `json:"username"` Password string `json:"password"` } // LoginResponse 登录响应 type LoginResponse struct { Token string `json:"token"` } // Login 登录业务逻辑 func Login(username, password string) (string, error) { user, err := db.GetUserByUsername(username) if err != nil { return "", errors.New("用户不存在") } if !CheckPassword(password, user.PasswordHash) { return "", errors.New("密码错误") } return GenerateJWT(username) } // RegisterRequest 注册请求 type RegisterRequest struct { Username string `json:"username"` Password string `json:"password"` } // Register 注册业务逻辑 func Register(username, password string) (string, error) { // 检查用户名是否已存在 exists, err := db.UserExists(username) if err != nil { return "", err } if exists { return "", errors.New("用户名已存在") } // 密码复杂度校验 if !ValidatePassword(password) { return "", errors.New("密码不符合复杂度要求") } // 加密密码 hash, err := HashPassword(password) if err != nil { return "", err } // 创建用户 if err := db.CreateUser(username, hash); err != nil { return "", err } // 生成 JWT return GenerateJWT(username) } // ChangePasswordRequest 修改密码请求 type ChangePasswordRequest struct { OldPassword string `json:"oldPassword"` NewPassword string `json:"newPassword"` } // ChangePassword 修改密码业务逻辑 func ChangePassword(username, oldPassword, newPassword string) error { user, err := db.GetUserByUsername(username) if err != nil { return errors.New("用户不存在") } if !CheckPassword(oldPassword, user.PasswordHash) { return errors.New("原密码错误") } if !ValidatePassword(newPassword) { return errors.New("新密码不符合复杂度要求") } hash, err := HashPassword(newPassword) if err != nil { return err } return db.UpdateUserPassword(username, hash) } // HasUsers 检查是否存在用户 func HasUsers() (bool, error) { count, err := db.CountUsers() if err != nil { return false, err } return count > 0, nil } // ================================================================ // Gin 中间件 // ================================================================ // AuthMiddleware JWT 认证中间件 func AuthMiddleware() gin.HandlerFunc { return func(c *gin.Context) { authHeader := c.GetHeader("Authorization") if authHeader == "" { c.JSON(http.StatusUnauthorized, gin.H{"code": 1, "msg": "未提供认证令牌"}) c.Abort() return } parts := strings.SplitN(authHeader, " ", 2) if len(parts) != 2 || parts[0] != "Bearer" { c.JSON(http.StatusUnauthorized, gin.H{"code": 1, "msg": "认证令牌格式错误"}) c.Abort() return } claims, err := ParseJWT(parts[1]) if err != nil { c.JSON(http.StatusUnauthorized, gin.H{"code": 1, "msg": "无效的认证令牌"}) c.Abort() return } c.Set("username", claims.Username) c.Next() } } // GetUsernameFromContext 从 gin.Context 获取当前用户名 func GetUsernameFromContext(c *gin.Context) string { username, exists := c.Get("username") if !exists { return "" } return username.(string) }