package handler import ( "crypto/rand" "encoding/hex" "log/slog" "net/http" "strconv" "strings" "time" "silk-server-go/internal/config" "silk-server-go/internal/middleware" "silk-server-go/internal/model" "github.com/gin-gonic/gin" "github.com/golang-jwt/jwt/v5" "golang.org/x/crypto/bcrypt" "gorm.io/gorm" ) // 默认令牌有效期:访问令牌 2 小时,刷新令牌 7 天 const ( defaultAccessExpiry = 2 * time.Hour defaultRefreshExpiry = 7 * 24 * time.Hour ) // RegisterAuthRoutes 注册认证相关路由 func RegisterAuthRoutes(rg *gin.RouterGroup, db *gorm.DB, cfg *config.Config) { // 启动时确保默认 admin 用户存在 ensureDefaultAdmin(db, cfg) rg.POST("/auth/register", registerHandler(db, cfg)) rg.POST("/auth/login", loginHandler(db, cfg)) rg.POST("/auth/refresh", refreshHandler(db, cfg)) rg.POST("/auth/logout", logoutHandler(cfg)) rg.POST("/auth/change-password", changePasswordHandler(db, cfg)) rg.GET("/auth/me", meHandler(db)) } // registerHandler 注册用户(bcrypt 哈希密码),返回 accessToken + refreshToken + user func registerHandler(db *gorm.DB, cfg *config.Config) gin.HandlerFunc { return func(c *gin.Context) { var body struct { Username string `json:"username"` Email string `json:"email"` Password string `json:"password"` FullName *string `json:"fullName"` Role string `json:"role"` } if err := c.ShouldBindJSON(&body); err != nil { c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()}) return } if len(body.Username) < 3 || len(body.Password) < 6 { c.JSON(http.StatusBadRequest, gin.H{"error": "username 至少3位,password 至少6位"}) return } // 检查用户名或邮箱是否已存在 var exists model.User if db.Where("username = ? OR email = ?", body.Username, body.Email).First(&exists).Error == nil { c.JSON(http.StatusBadRequest, gin.H{"error": "username or email already taken"}) return } // 哈希密码 hash, err := bcrypt.GenerateFromPassword([]byte(body.Password), 10) if err != nil { c.JSON(http.StatusBadRequest, gin.H{"error": "密码哈希失败"}) return } // 强制角色为 viewer,防止垂直越权(注册接口不允许自选角色) role := model.RoleViewer user := model.User{ Username: body.Username, Email: body.Email, PasswordHash: string(hash), FullName: body.FullName, Role: role, Active: true, } if err := db.Create(&user).Error; err != nil { c.JSON(http.StatusBadRequest, gin.H{"error": "创建用户失败"}) return } assignDefaultOrganization(db, user.ID) // 记录审计日志 uid := user.ID uname := user.Username res := "users" desc := "user registered" recordAudit(db, &model.AuditLog{ UserID: &uid, Username: &uname, Action: "create", Resource: &res, TargetID: &uid, Description: &desc, }) c.JSON(http.StatusCreated, buildLoginPayload(db, user, cfg)) } } // loginHandler 登录(用户名或邮箱 + 密码),返回 token func loginHandler(db *gorm.DB, cfg *config.Config) gin.HandlerFunc { return func(c *gin.Context) { var body struct { Username string `json:"username"` Password string `json:"password"` } if err := c.ShouldBindJSON(&body); err != nil { c.JSON(http.StatusUnauthorized, gin.H{"error": "用户名或密码错误"}) return } // 登录限流:检查 IP+用户名是否被锁定 if middleware.CheckLoginLock(c, body.Username) { return } var user model.User if db.Where("username = ? OR email = ?", body.Username, body.Username).First(&user).Error != nil { if err := middleware.RecordLoginFail(c, body.Username); err != nil { c.JSON(http.StatusServiceUnavailable, gin.H{"error": err.Error()}) return } c.JSON(http.StatusUnauthorized, gin.H{"error": "用户名或密码错误"}) return } if err := bcrypt.CompareHashAndPassword([]byte(user.PasswordHash), []byte(body.Password)); err != nil { if err := middleware.RecordLoginFail(c, body.Username); err != nil { c.JSON(http.StatusServiceUnavailable, gin.H{"error": err.Error()}) return } c.JSON(http.StatusUnauthorized, gin.H{"error": "用户名或密码错误"}) return } if !user.Active { c.JSON(http.StatusUnauthorized, gin.H{"error": "账号已禁用"}) return } // 登录成功,清空失败计数 if err := middleware.RecordLoginSuccess(c, body.Username); err != nil { c.JSON(http.StatusServiceUnavailable, gin.H{"error": err.Error()}) return } // 记录审计日志 uid := user.ID uname := user.Username res := "users" desc := "user login" ip := c.ClientIP() ua := c.GetHeader("User-Agent") recordAudit(db, &model.AuditLog{ UserID: &uid, Username: &uname, Action: "login", Resource: &res, TargetID: &uid, Description: &desc, IPAddress: &ip, UserAgent: &ua, }) c.JSON(http.StatusOK, buildLoginPayload(db, user, cfg)) } } // refreshHandler 刷新访问令牌(body 传 refreshToken) func refreshHandler(db *gorm.DB, cfg *config.Config) gin.HandlerFunc { return func(c *gin.Context) { var body struct { RefreshToken string `json:"refreshToken"` } if err := c.ShouldBindJSON(&body); err != nil || strings.TrimSpace(body.RefreshToken) == "" { c.JSON(http.StatusBadRequest, gin.H{"error": "refreshToken 不能为空"}) return } claims, token, err := middleware.ExtractClaims(body.RefreshToken, cfg.JWTSecret) if err != nil || !token.Valid { c.JSON(http.StatusUnauthorized, gin.H{"error": "刷新令牌无效"}) return } if claims.TokenType != "refresh" { c.JSON(http.StatusUnauthorized, gin.H{"error": "刷新令牌无效"}) return } revoked, err := middleware.IsRevoked(claims) if err != nil { c.JSON(http.StatusServiceUnavailable, gin.H{"error": "认证状态服务不可用"}) return } if revoked { c.JSON(http.StatusUnauthorized, gin.H{"error": "刷新令牌已注销"}) return } // 查询用户,确保仍然有效 var user model.User if db.Where("id = ?", claims.Subject).First(&user).Error != nil { c.JSON(http.StatusUnauthorized, gin.H{"error": "用户不存在"}) return } if !user.Active { c.JSON(http.StatusUnauthorized, gin.H{"error": "账号已禁用"}) return } // 吊销旧刷新令牌(一次性使用),签发新令牌对 if err := middleware.RevokeToken(claims, body.RefreshToken, claims.ExpiresAt.Time); err != nil { c.JSON(http.StatusServiceUnavailable, gin.H{"error": err.Error()}) return } c.JSON(http.StatusOK, buildLoginPayload(db, user, cfg)) } } // logoutHandler 登出:将当前访问令牌加入黑名单 func logoutHandler(cfg *config.Config) gin.HandlerFunc { return func(c *gin.Context) { authHeader := c.GetHeader("Authorization") parts := strings.SplitN(authHeader, " ", 2) if len(parts) == 2 && strings.EqualFold(parts[0], "Bearer") { if claims, token, err := middleware.ExtractClaims(parts[1], cfg.JWTSecret); err == nil && token.Valid { if claims.ExpiresAt != nil { if err := middleware.RevokeToken(claims, parts[1], claims.ExpiresAt.Time); err != nil { c.JSON(http.StatusServiceUnavailable, gin.H{"error": err.Error()}) return } } } } c.JSON(http.StatusOK, gin.H{"ok": true}) } } // changePasswordHandler 修改当前用户密码 func changePasswordHandler(db *gorm.DB, cfg *config.Config) gin.HandlerFunc { return func(c *gin.Context) { userVal, exists := c.Get("user") if !exists { c.JSON(http.StatusUnauthorized, gin.H{"error": "未认证"}) return } userMap, _ := userVal.(map[string]interface{}) uid, _ := userMap["sub"].(string) var body struct { OldPassword string `json:"oldPassword"` NewPassword string `json:"newPassword"` } if err := c.ShouldBindJSON(&body); err != nil { c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()}) return } if len(body.NewPassword) < 6 { c.JSON(http.StatusBadRequest, gin.H{"error": "新密码至少 6 位"}) return } var user model.User if db.Where("id = ?", uid).First(&user).Error != nil { c.JSON(http.StatusUnauthorized, gin.H{"error": "用户不存在"}) return } if err := bcrypt.CompareHashAndPassword([]byte(user.PasswordHash), []byte(body.OldPassword)); err != nil { c.JSON(http.StatusBadRequest, gin.H{"error": "原密码错误"}) return } hash, err := bcrypt.GenerateFromPassword([]byte(body.NewPassword), 10) if err != nil { c.JSON(http.StatusInternalServerError, gin.H{"error": "密码哈希失败"}) return } if err := db.Model(&model.User{}).Where("id = ?", uid).Update("password_hash", string(hash)).Error; err != nil { c.JSON(http.StatusInternalServerError, gin.H{"error": "更新密码失败"}) return } // 记录审计日志 uname := user.Username res := "users" desc := "password changed" recordAudit(db, &model.AuditLog{ UserID: &uid, Username: &uname, Action: "update", Resource: &res, TargetID: &uid, Description: &desc, }) c.JSON(http.StatusOK, gin.H{"ok": true}) } } // meHandler 返回当前用户信息(从 context 获取 user)+ 权限列表 func meHandler(db *gorm.DB) gin.HandlerFunc { return func(c *gin.Context) { userVal, exists := c.Get("user") if !exists { c.JSON(http.StatusUnauthorized, gin.H{"error": "未认证"}) return } userMap, _ := userVal.(map[string]interface{}) role, _ := userMap["role"].(string) permissions := getUserPermissionCodes(db, role) c.JSON(http.StatusOK, gin.H{ "sub": userMap["sub"], "username": userMap["username"], "role": role, "permissions": permissions, }) } } // getUserPermissionCodes 查询指定角色的权限码列表 func getUserPermissionCodes(db *gorm.DB, role string) []string { if role == model.RoleAdmin { // admin 拥有全部权限 codes := make([]string, 0, len(model.AllPermissions)) for _, p := range model.AllPermissions { codes = append(codes, p.Code) } return codes } var codes []string db.Table("role_permissions"). Select("permissions.code"). Joins("JOIN permissions ON permissions.id = role_permissions.permission_id"). Where("role_permissions.role = ?", role). Scan(&codes) return codes } // buildLoginPayload 构建登录返回数据(accessToken + refreshToken + user + permissions) func buildLoginPayload(db *gorm.DB, user model.User, cfg *config.Config) gin.H { return gin.H{ "accessToken": signToken(user, cfg, parseDuration(cfg.JWTExpiresIn, defaultAccessExpiry), "access"), "refreshToken": signToken(user, cfg, defaultRefreshExpiry, "refresh"), "user": gin.H{ "id": user.ID, "username": user.Username, "email": user.Email, "fullName": user.FullName, "role": user.Role, "permissions": getUserPermissionCodes(db, user.Role), }, } } // signToken 签发 JWT(payload: sub/username/role/tokenType/jti,HS256 + JWTSecret) func signToken(user model.User, cfg *config.Config, expiry time.Duration, tokenType string) string { claims := middleware.JWTClaims{ Username: user.Username, Role: user.Role, TokenType: tokenType, RegisteredClaims: jwt.RegisteredClaims{ ID: randomJTI(), Subject: user.ID, IssuedAt: jwt.NewNumericDate(time.Now()), ExpiresAt: jwt.NewNumericDate(time.Now().Add(expiry)), }, } token := jwt.NewWithClaims(jwt.SigningMethodHS256, claims) tokenStr, err := token.SignedString([]byte(cfg.JWTSecret)) if err != nil { slog.Error("签发 JWT 失败", "error", err) return "" } return tokenStr } // randomJTI 生成 16 字节随机 token ID func randomJTI() string { b := make([]byte, 16) if _, err := rand.Read(b); err != nil { return strconv.FormatInt(time.Now().UnixNano(), 16) } return hex.EncodeToString(b) } // parseDuration 解析过期时间字符串(支持 "7d"、"1h"、"30m" 等),失败返回默认值 func parseDuration(s string, fallback time.Duration) time.Duration { s = strings.TrimSpace(s) if s == "" { return fallback } // 支持 "7d" 格式(Go 原生 time.ParseDuration 不支持天) if strings.HasSuffix(s, "d") { days, err := strconv.Atoi(strings.TrimSuffix(s, "d")) if err == nil { return time.Duration(days) * 24 * time.Hour } } d, err := time.ParseDuration(s) if err != nil { return fallback } return d } // ensureDefaultAdmin 确保默认 admin 用户存在 func ensureDefaultAdmin(db *gorm.DB, cfg *config.Config) { username := cfg.DefaultAdminUsername if username == "" { username = "admin" } var existing model.User if db.Where("username = ?", username).First(&existing).Error == nil { return // 已存在 } password := cfg.DefaultAdminPassword if password == "" { password = "silk@123" } email := cfg.DefaultAdminEmail if email == "" { email = "admin@silk.local" } hash, err := bcrypt.GenerateFromPassword([]byte(password), 10) if err != nil { slog.Error("默认 admin 密码哈希失败", "error", err) return } fullName := "系统管理员" admin := model.User{ Username: username, Email: email, PasswordHash: string(hash), FullName: &fullName, Role: "admin", Active: true, } if err := db.Create(&admin).Error; err != nil { slog.Error("创建默认 admin 失败", "error", err) return } assignDefaultOrganization(db, admin.ID) // 记录审计日志 uid := admin.ID uname := admin.Username res := "users" desc := "default admin created" recordAudit(db, &model.AuditLog{ UserID: &uid, Username: &uname, Action: "create", Resource: &res, TargetID: &uid, Description: &desc, }) slog.Info("默认 admin 用户已创建", "username", username) }