470 lines
14 KiB
Go
470 lines
14 KiB
Go
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)
|
||
}
|