Files
silk/server-go/internal/handler/auth.go
T
2026-08-17 21:43:26 +08:00

470 lines
14 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
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 签发 JWTpayload: sub/username/role/tokenType/jtiHS256 + 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)
}