搭建用户体系
This commit is contained in:
@@ -0,0 +1,64 @@
|
||||
package config
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/joho/godotenv"
|
||||
)
|
||||
|
||||
type Config struct {
|
||||
DatabaseURL string
|
||||
JWTSecret string
|
||||
JWTExpirationHours int
|
||||
Port string
|
||||
AllowedOrigins []string
|
||||
}
|
||||
|
||||
func Load() (*Config, error) {
|
||||
_ = godotenv.Load()
|
||||
|
||||
databaseURL := os.Getenv("DATABASE_URL")
|
||||
if databaseURL == "" {
|
||||
return nil, fmt.Errorf("DATABASE_URL must be set")
|
||||
}
|
||||
|
||||
jwtSecret := os.Getenv("JWT_SECRET")
|
||||
if jwtSecret == "" {
|
||||
jwtSecret = "change-me-in-production"
|
||||
fmt.Fprintln(os.Stderr, "WARNING: JWT_SECRET not set, using default secret")
|
||||
}
|
||||
|
||||
expHours, _ := strconv.Atoi(os.Getenv("JWT_EXPIRATION_HOURS"))
|
||||
if expHours == 0 {
|
||||
expHours = 168 // 7 days
|
||||
}
|
||||
|
||||
port := os.Getenv("PORT")
|
||||
if port == "" {
|
||||
port = "3019"
|
||||
}
|
||||
|
||||
allowedOrigins := []string{"*"}
|
||||
if v := os.Getenv("ALLOWED_ORIGINS"); v != "" {
|
||||
allowedOrigins = strings.Split(v, ",")
|
||||
for i := range allowedOrigins {
|
||||
allowedOrigins[i] = strings.TrimSpace(allowedOrigins[i])
|
||||
}
|
||||
}
|
||||
|
||||
return &Config{
|
||||
DatabaseURL: databaseURL,
|
||||
JWTSecret: jwtSecret,
|
||||
JWTExpirationHours: expHours,
|
||||
Port: port,
|
||||
AllowedOrigins: allowedOrigins,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (c *Config) JWTExpiration() time.Duration {
|
||||
return time.Duration(c.JWTExpirationHours) * time.Hour
|
||||
}
|
||||
@@ -0,0 +1,27 @@
|
||||
package db
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
|
||||
"gorm.io/driver/postgres"
|
||||
"gorm.io/gorm"
|
||||
"gorm.io/gorm/logger"
|
||||
)
|
||||
|
||||
func Init(databaseURL string) (*gorm.DB, error) {
|
||||
db, err := gorm.Open(postgres.Open(databaseURL), &gorm.Config{
|
||||
Logger: logger.Default.LogMode(logger.Silent),
|
||||
})
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("open database: %w", err)
|
||||
}
|
||||
|
||||
sqlDB, err := db.DB()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
sqlDB.SetMaxIdleConns(10)
|
||||
sqlDB.SetMaxOpenConns(100)
|
||||
|
||||
return db, nil
|
||||
}
|
||||
@@ -0,0 +1,229 @@
|
||||
package handlers
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"strings"
|
||||
|
||||
"stock-user-system/internal/config"
|
||||
"stock-user-system/internal/middleware"
|
||||
"stock-user-system/internal/models"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
type AdminHandler struct {
|
||||
DB *gorm.DB
|
||||
CFG *config.Config
|
||||
}
|
||||
|
||||
func allowedRolesForCreation(actor models.RoleName) []models.RoleName {
|
||||
switch actor {
|
||||
case models.RoleSystemAdmin:
|
||||
return []models.RoleName{models.RoleAdmin, models.RoleUser, models.RoleGuest}
|
||||
case models.RoleAdmin:
|
||||
return []models.RoleName{models.RoleUser, models.RoleGuest}
|
||||
default:
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
func containsRole(roles []models.RoleName, target models.RoleName) bool {
|
||||
for _, r := range roles {
|
||||
if r == target {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func (h *AdminHandler) ListUsers(c *gin.Context) {
|
||||
var users []models.User
|
||||
if err := h.DB.Preload("Role").Order("created_at DESC").Find(&users).Error; err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"success": false, "error": "internal server error"})
|
||||
return
|
||||
}
|
||||
|
||||
result := make([]models.PublicUserInfo, 0, len(users))
|
||||
for _, u := range users {
|
||||
result = append(result, u.ToPublicInfo())
|
||||
}
|
||||
c.JSON(http.StatusOK, result)
|
||||
}
|
||||
|
||||
func (h *AdminHandler) CreateUser(c *gin.Context) {
|
||||
current, _ := middleware.GetCurrentUser(c)
|
||||
|
||||
var req struct {
|
||||
Username string `json:"username" binding:"required"`
|
||||
Email *string `json:"email"`
|
||||
Password string `json:"password" binding:"required"`
|
||||
Role string `json:"role"`
|
||||
}
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"success": false, "error": "请求参数错误"})
|
||||
return
|
||||
}
|
||||
|
||||
targetRole := models.RoleUser
|
||||
if req.Role != "" {
|
||||
parsed, err := models.ParseRoleName(req.Role)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"success": false, "error": "无效的角色"})
|
||||
return
|
||||
}
|
||||
targetRole = parsed
|
||||
}
|
||||
|
||||
if !containsRole(allowedRolesForCreation(current.Role), targetRole) {
|
||||
c.JSON(http.StatusForbidden, gin.H{"success": false, "error": "权限不足"})
|
||||
return
|
||||
}
|
||||
|
||||
if err := validateUsername(req.Username); err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"success": false, "error": err.Error()})
|
||||
return
|
||||
}
|
||||
if len(req.Password) < 6 {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"success": false, "error": "密码长度至少 6 位"})
|
||||
return
|
||||
}
|
||||
|
||||
var role models.Role
|
||||
if err := h.DB.Where("name = ?", targetRole.String()).First(&role).Error; err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"success": false, "error": "角色不存在"})
|
||||
return
|
||||
}
|
||||
|
||||
hash, err := hashPassword(req.Password)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"success": false, "error": "internal server error"})
|
||||
return
|
||||
}
|
||||
|
||||
username := strings.ToLower(strings.TrimSpace(req.Username))
|
||||
var email *string
|
||||
if req.Email != nil && *req.Email != "" {
|
||||
e := strings.ToLower(strings.TrimSpace(*req.Email))
|
||||
email = &e
|
||||
}
|
||||
|
||||
user := models.User{
|
||||
Username: username,
|
||||
Email: email,
|
||||
PasswordHash: hash,
|
||||
RoleID: role.ID,
|
||||
Status: "active",
|
||||
}
|
||||
|
||||
if err := h.DB.Create(&user).Error; err != nil {
|
||||
if strings.Contains(err.Error(), "duplicate key") {
|
||||
c.JSON(http.StatusConflict, gin.H{"success": false, "error": "用户名或邮箱已存在"})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"success": false, "error": "internal server error"})
|
||||
return
|
||||
}
|
||||
|
||||
h.DB.Preload("Role").First(&user, "id = ?", user.ID)
|
||||
c.JSON(http.StatusOK, user.ToPublicInfo())
|
||||
}
|
||||
|
||||
func (h *AdminHandler) UpdateUser(c *gin.Context) {
|
||||
current, _ := middleware.GetCurrentUser(c)
|
||||
userID := c.Param("id")
|
||||
|
||||
var req struct {
|
||||
Email *string `json:"email"`
|
||||
Role string `json:"role"`
|
||||
Status string `json:"status"`
|
||||
}
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"success": false, "error": "请求参数错误"})
|
||||
return
|
||||
}
|
||||
|
||||
var target models.User
|
||||
if err := h.DB.Preload("Role").Where("id = ?", userID).First(&target).Error; err != nil {
|
||||
c.JSON(http.StatusNotFound, gin.H{"success": false, "error": "not found"})
|
||||
return
|
||||
}
|
||||
|
||||
if !current.Role.CanManage(target.Role.NameEnum()) && current.ID != userID {
|
||||
c.JSON(http.StatusForbidden, gin.H{"success": false, "error": "权限不足"})
|
||||
return
|
||||
}
|
||||
|
||||
updates := map[string]interface{}{}
|
||||
|
||||
if req.Email != nil && *req.Email != "" {
|
||||
updates["email"] = strings.ToLower(strings.TrimSpace(*req.Email))
|
||||
}
|
||||
|
||||
if req.Status != "" {
|
||||
updates["status"] = req.Status
|
||||
}
|
||||
|
||||
if req.Role != "" {
|
||||
newRole, err := models.ParseRoleName(req.Role)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"success": false, "error": "无效的角色"})
|
||||
return
|
||||
}
|
||||
if !containsRole(allowedRolesForCreation(current.Role), newRole) || !current.Role.CanManage(newRole) {
|
||||
c.JSON(http.StatusForbidden, gin.H{"success": false, "error": "权限不足"})
|
||||
return
|
||||
}
|
||||
var role models.Role
|
||||
if err := h.DB.Where("name = ?", newRole.String()).First(&role).Error; err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"success": false, "error": "角色不存在"})
|
||||
return
|
||||
}
|
||||
updates["role_id"] = role.ID
|
||||
}
|
||||
|
||||
if err := h.DB.Model(&target).Updates(updates).Error; err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"success": false, "error": "internal server error"})
|
||||
return
|
||||
}
|
||||
|
||||
h.DB.Preload("Role").First(&target, "id = ?", userID)
|
||||
c.JSON(http.StatusOK, target.ToPublicInfo())
|
||||
}
|
||||
|
||||
func (h *AdminHandler) DeleteUser(c *gin.Context) {
|
||||
current, _ := middleware.GetCurrentUser(c)
|
||||
userID := c.Param("id")
|
||||
|
||||
if current.ID == userID {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"success": false, "error": "不能删除自己"})
|
||||
return
|
||||
}
|
||||
|
||||
var target models.User
|
||||
if err := h.DB.Preload("Role").Where("id = ?", userID).First(&target).Error; err != nil {
|
||||
c.JSON(http.StatusNotFound, gin.H{"success": false, "error": "not found"})
|
||||
return
|
||||
}
|
||||
|
||||
if !current.Role.CanManage(target.Role.NameEnum()) {
|
||||
c.JSON(http.StatusForbidden, gin.H{"success": false, "error": "权限不足"})
|
||||
return
|
||||
}
|
||||
|
||||
if err := h.DB.Delete(&target).Error; err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"success": false, "error": "internal server error"})
|
||||
return
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, gin.H{"message": "用户已删除"})
|
||||
}
|
||||
|
||||
func (h *AdminHandler) ListRoles(c *gin.Context) {
|
||||
var roles []models.Role
|
||||
if err := h.DB.Order("id").Find(&roles).Error; err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"success": false, "error": "internal server error"})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, roles)
|
||||
}
|
||||
@@ -0,0 +1,178 @@
|
||||
package handlers
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"net/http"
|
||||
"strings"
|
||||
|
||||
"stock-user-system/internal/config"
|
||||
"stock-user-system/internal/middleware"
|
||||
"stock-user-system/internal/models"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"golang.org/x/crypto/bcrypt"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
type AuthHandler struct {
|
||||
DB *gorm.DB
|
||||
CFG *config.Config
|
||||
}
|
||||
|
||||
func (h *AuthHandler) Register(c *gin.Context) {
|
||||
var req struct {
|
||||
Username string `json:"username" binding:"required"`
|
||||
Email *string `json:"email"`
|
||||
Password string `json:"password" binding:"required"`
|
||||
}
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"success": false, "error": "请求参数错误"})
|
||||
return
|
||||
}
|
||||
|
||||
if err := validateUsername(req.Username); err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"success": false, "error": err.Error()})
|
||||
return
|
||||
}
|
||||
if len(req.Password) < 6 {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"success": false, "error": "密码长度至少 6 位"})
|
||||
return
|
||||
}
|
||||
|
||||
username := strings.ToLower(strings.TrimSpace(req.Username))
|
||||
|
||||
var role models.Role
|
||||
if err := h.DB.Where("name = ?", models.RoleUser).First(&role).Error; err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"success": false, "error": "internal server error"})
|
||||
return
|
||||
}
|
||||
|
||||
hash, err := hashPassword(req.Password)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"success": false, "error": "internal server error"})
|
||||
return
|
||||
}
|
||||
|
||||
var email *string
|
||||
if req.Email != nil && *req.Email != "" {
|
||||
e := strings.ToLower(strings.TrimSpace(*req.Email))
|
||||
email = &e
|
||||
}
|
||||
|
||||
user := models.User{
|
||||
Username: username,
|
||||
Email: email,
|
||||
PasswordHash: hash,
|
||||
RoleID: role.ID,
|
||||
Status: "active",
|
||||
}
|
||||
|
||||
if err := h.DB.Create(&user).Error; err != nil {
|
||||
if strings.Contains(err.Error(), "duplicate key") {
|
||||
c.JSON(http.StatusConflict, gin.H{"success": false, "error": "用户名或邮箱已存在"})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"success": false, "error": "internal server error"})
|
||||
return
|
||||
}
|
||||
|
||||
if err := h.DB.Preload("Role").First(&user, "id = ?", user.ID).Error; err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"success": false, "error": "internal server error"})
|
||||
return
|
||||
}
|
||||
|
||||
token, err := middleware.GenerateToken(user.ID, models.RoleUser, h.CFG)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"success": false, "error": "internal server error"})
|
||||
return
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, gin.H{
|
||||
"token": token,
|
||||
"user": user.ToPublicInfo(),
|
||||
})
|
||||
}
|
||||
|
||||
func (h *AuthHandler) Login(c *gin.Context) {
|
||||
var req struct {
|
||||
Username string `json:"username" binding:"required"`
|
||||
Password string `json:"password" binding:"required"`
|
||||
}
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"success": false, "error": "请求参数错误"})
|
||||
return
|
||||
}
|
||||
|
||||
username := strings.ToLower(strings.TrimSpace(req.Username))
|
||||
|
||||
var user models.User
|
||||
if err := h.DB.Preload("Role").Where("LOWER(username) = ?", username).First(&user).Error; err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"success": false, "error": "用户名或密码错误"})
|
||||
return
|
||||
}
|
||||
|
||||
if user.Status != "active" {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"success": false, "error": "账号已被禁用"})
|
||||
return
|
||||
}
|
||||
|
||||
if err := bcrypt.CompareHashAndPassword([]byte(user.PasswordHash), []byte(req.Password)); err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"success": false, "error": "用户名或密码错误"})
|
||||
return
|
||||
}
|
||||
|
||||
token, err := middleware.GenerateToken(user.ID, user.Role.NameEnum(), h.CFG)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"success": false, "error": "internal server error"})
|
||||
return
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, gin.H{
|
||||
"token": token,
|
||||
"user": user.ToPublicInfo(),
|
||||
})
|
||||
}
|
||||
|
||||
func (h *AuthHandler) Me(c *gin.Context) {
|
||||
current, ok := middleware.GetCurrentUser(c)
|
||||
if !ok {
|
||||
c.JSON(http.StatusUnauthorized, gin.H{"success": false, "error": "访问令牌无效或已过期"})
|
||||
return
|
||||
}
|
||||
|
||||
var user models.User
|
||||
if err := h.DB.Preload("Role").Where("id = ?", current.ID).First(&user).Error; err != nil {
|
||||
c.JSON(http.StatusNotFound, gin.H{"success": false, "error": "not found"})
|
||||
return
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, user.ToPublicInfo())
|
||||
}
|
||||
|
||||
func (h *AuthHandler) Logout(c *gin.Context) {
|
||||
c.JSON(http.StatusOK, gin.H{"message": "登出成功"})
|
||||
}
|
||||
|
||||
func (h *AuthHandler) PublicInfo(c *gin.Context) {
|
||||
c.JSON(http.StatusOK, gin.H{
|
||||
"message": "欢迎访问 A股工具公开信息",
|
||||
"guest_allowed": true,
|
||||
})
|
||||
}
|
||||
|
||||
func hashPassword(password string) (string, error) {
|
||||
bytes, err := bcrypt.GenerateFromPassword([]byte(password), bcrypt.DefaultCost)
|
||||
return string(bytes), err
|
||||
}
|
||||
|
||||
func validateUsername(username string) error {
|
||||
if len(username) < 3 || len(username) > 32 {
|
||||
return fmt.Errorf("用户名长度需在 3-32 位之间")
|
||||
}
|
||||
for _, r := range username {
|
||||
if !(r >= 'a' && r <= 'z') && !(r >= 'A' && r <= 'Z') && !(r >= '0' && r <= '9') && r != '_' && r != '-' {
|
||||
return fmt.Errorf("用户名只能包含字母、数字、下划线和短横线")
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,116 @@
|
||||
package middleware
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"stock-user-system/internal/config"
|
||||
"stock-user-system/internal/models"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/golang-jwt/jwt/v5"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
type CurrentUser struct {
|
||||
ID string
|
||||
Username string
|
||||
Role models.RoleName
|
||||
}
|
||||
|
||||
func AuthMiddleware(cfg *config.Config, db *gorm.DB) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
authHeader := c.GetHeader("Authorization")
|
||||
token := ""
|
||||
if strings.HasPrefix(authHeader, "Bearer ") {
|
||||
token = strings.TrimPrefix(authHeader, "Bearer ")
|
||||
}
|
||||
|
||||
if token != "" {
|
||||
claims, err := parseToken(token, cfg.JWTSecret)
|
||||
if err == nil {
|
||||
var user models.User
|
||||
if err := db.Preload("Role").Where("id = ? AND status = ?", claims.Subject, "active").First(&user).Error; err == nil {
|
||||
c.Set("currentUser", CurrentUser{
|
||||
ID: user.ID,
|
||||
Username: user.Username,
|
||||
Role: user.Role.NameEnum(),
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
c.Next()
|
||||
}
|
||||
}
|
||||
|
||||
type Claims struct {
|
||||
Subject string `json:"sub"`
|
||||
Role string `json:"role"`
|
||||
jwt.RegisteredClaims
|
||||
}
|
||||
|
||||
func GenerateToken(userID string, role models.RoleName, cfg *config.Config) (string, error) {
|
||||
claims := Claims{
|
||||
Subject: userID,
|
||||
Role: role.String(),
|
||||
RegisteredClaims: jwt.RegisteredClaims{
|
||||
ExpiresAt: jwt.NewNumericDate(time.Now().Add(cfg.JWTExpiration())),
|
||||
},
|
||||
}
|
||||
token := jwt.NewWithClaims(jwt.SigningMethodHS256, claims)
|
||||
return token.SignedString([]byte(cfg.JWTSecret))
|
||||
}
|
||||
|
||||
func parseToken(tokenString string, secret string) (*Claims, error) {
|
||||
token, err := jwt.ParseWithClaims(tokenString, &Claims{}, func(token *jwt.Token) (interface{}, error) {
|
||||
return []byte(secret), nil
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if claims, ok := token.Claims.(*Claims); ok && token.Valid {
|
||||
return claims, nil
|
||||
}
|
||||
return nil, jwt.ErrSignatureInvalid
|
||||
}
|
||||
|
||||
func RequireAuth() gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
if _, exists := c.Get("currentUser"); !exists {
|
||||
c.JSON(http.StatusUnauthorized, gin.H{"success": false, "error": "访问令牌无效或已过期"})
|
||||
c.Abort()
|
||||
return
|
||||
}
|
||||
c.Next()
|
||||
}
|
||||
}
|
||||
|
||||
func RequireRoles(allowed ...models.RoleName) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
val, exists := c.Get("currentUser")
|
||||
if !exists {
|
||||
c.JSON(http.StatusUnauthorized, gin.H{"success": false, "error": "访问令牌无效或已过期"})
|
||||
c.Abort()
|
||||
return
|
||||
}
|
||||
current := val.(CurrentUser)
|
||||
for _, role := range allowed {
|
||||
if current.Role == role {
|
||||
c.Next()
|
||||
return
|
||||
}
|
||||
}
|
||||
c.JSON(http.StatusForbidden, gin.H{"success": false, "error": "权限不足"})
|
||||
c.Abort()
|
||||
}
|
||||
}
|
||||
|
||||
func GetCurrentUser(c *gin.Context) (CurrentUser, bool) {
|
||||
val, exists := c.Get("currentUser")
|
||||
if !exists {
|
||||
return CurrentUser{}, false
|
||||
}
|
||||
return val.(CurrentUser), true
|
||||
}
|
||||
@@ -0,0 +1,127 @@
|
||||
package models
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"time"
|
||||
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
type RoleName string
|
||||
|
||||
const (
|
||||
RoleSystemAdmin RoleName = "system_admin"
|
||||
RoleAdmin RoleName = "admin"
|
||||
RoleUser RoleName = "user"
|
||||
RoleGuest RoleName = "guest"
|
||||
)
|
||||
|
||||
func (r RoleName) String() string { return string(r) }
|
||||
|
||||
func (r RoleName) Rank() int {
|
||||
switch r {
|
||||
case RoleSystemAdmin:
|
||||
return 3
|
||||
case RoleAdmin:
|
||||
return 2
|
||||
case RoleUser:
|
||||
return 1
|
||||
default:
|
||||
return 0
|
||||
}
|
||||
}
|
||||
|
||||
func (r RoleName) CanManage(target RoleName) bool {
|
||||
return r.Rank() > target.Rank()
|
||||
}
|
||||
|
||||
func ParseRoleName(s string) (RoleName, error) {
|
||||
switch s {
|
||||
case "system_admin":
|
||||
return RoleSystemAdmin, nil
|
||||
case "admin":
|
||||
return RoleAdmin, nil
|
||||
case "user":
|
||||
return RoleUser, nil
|
||||
case "guest":
|
||||
return RoleGuest, nil
|
||||
default:
|
||||
return RoleGuest, fmt.Errorf("unknown role: %s", s)
|
||||
}
|
||||
}
|
||||
|
||||
type Role struct {
|
||||
ID int32 `json:"id" gorm:"primaryKey;autoIncrement"`
|
||||
Name string `json:"name" gorm:"uniqueIndex;size:32;not null"`
|
||||
Description *string `json:"description"`
|
||||
Permissions string `json:"permissions" gorm:"type:jsonb;default:'[]'"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
Users []User `json:"-" gorm:"foreignKey:RoleID"`
|
||||
}
|
||||
|
||||
func (r *Role) NameEnum() RoleName {
|
||||
role, _ := ParseRoleName(r.Name)
|
||||
return role
|
||||
}
|
||||
|
||||
type User struct {
|
||||
ID string `json:"id" gorm:"type:uuid;primaryKey;default:gen_random_uuid()"`
|
||||
Username string `json:"username" gorm:"uniqueIndex;size:32;not null"`
|
||||
Email *string `json:"email" gorm:"uniqueIndex;size:128"`
|
||||
PasswordHash string `json:"-" gorm:"size:255"`
|
||||
RoleID int32 `json:"role_id" gorm:"not null"`
|
||||
Role Role `json:"role,omitempty" gorm:"foreignKey:RoleID;references:ID"`
|
||||
Status string `json:"status" gorm:"size:16;default:active"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
UpdatedAt time.Time `json:"updated_at"`
|
||||
}
|
||||
|
||||
type PublicUserInfo struct {
|
||||
ID string `json:"id"`
|
||||
Username string `json:"username"`
|
||||
Email *string `json:"email"`
|
||||
Role RoleName `json:"role"`
|
||||
Status string `json:"status"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
}
|
||||
|
||||
func (u *User) ToPublicInfo() PublicUserInfo {
|
||||
return PublicUserInfo{
|
||||
ID: u.ID,
|
||||
Username: u.Username,
|
||||
Email: u.Email,
|
||||
Role: u.Role.NameEnum(),
|
||||
Status: u.Status,
|
||||
CreatedAt: u.CreatedAt,
|
||||
}
|
||||
}
|
||||
|
||||
func AutoMigrate(db *gorm.DB) error {
|
||||
return db.AutoMigrate(&Role{}, &User{})
|
||||
}
|
||||
|
||||
func SeedRoles(db *gorm.DB) error {
|
||||
roles := []Role{
|
||||
{Name: string(RoleSystemAdmin), Description: strPtr("系统管理员,可管理管理员与系统配置"), Permissions: `["*"]`},
|
||||
{Name: string(RoleAdmin), Description: strPtr("管理员,可管理普通用户"), Permissions: `["users.read", "users.write", "users.create"]`},
|
||||
{Name: string(RoleUser), Description: strPtr("普通用户,可访问业务功能"), Permissions: `["dashboard.read", "profile.write"]`},
|
||||
{Name: string(RoleGuest), Description: strPtr("游客,仅可查看公开内容"), Permissions: `["public.read"]`},
|
||||
}
|
||||
for _, role := range roles {
|
||||
var existing Role
|
||||
if err := db.Where("name = ?", role.Name).First(&existing).Error; err != nil {
|
||||
if err == gorm.ErrRecordNotFound {
|
||||
if err := db.Create(&role).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
} else {
|
||||
return err
|
||||
}
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func strPtr(s string) *string {
|
||||
return &s
|
||||
}
|
||||
@@ -0,0 +1,58 @@
|
||||
package routes
|
||||
|
||||
import (
|
||||
"stock-user-system/internal/config"
|
||||
"stock-user-system/internal/handlers"
|
||||
"stock-user-system/internal/middleware"
|
||||
"stock-user-system/internal/models"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
func Setup(cfg *config.Config, db *gorm.DB) *gin.Engine {
|
||||
authHandler := &handlers.AuthHandler{DB: db, CFG: cfg}
|
||||
adminHandler := &handlers.AdminHandler{DB: db, CFG: cfg}
|
||||
|
||||
r := gin.Default()
|
||||
|
||||
// CORS
|
||||
r.Use(func(c *gin.Context) {
|
||||
c.Writer.Header().Set("Access-Control-Allow-Origin", "*")
|
||||
c.Writer.Header().Set("Access-Control-Allow-Methods", "GET, POST, PUT, DELETE, OPTIONS")
|
||||
c.Writer.Header().Set("Access-Control-Allow-Headers", "Origin, Content-Type, Accept, Authorization")
|
||||
if c.Request.Method == "OPTIONS" {
|
||||
c.AbortWithStatus(204)
|
||||
return
|
||||
}
|
||||
c.Next()
|
||||
})
|
||||
|
||||
r.Use(middleware.AuthMiddleware(cfg, db))
|
||||
|
||||
// 公开接口
|
||||
r.GET("/api/public", authHandler.PublicInfo)
|
||||
r.POST("/api/auth/register", authHandler.Register)
|
||||
r.POST("/api/auth/login", authHandler.Login)
|
||||
|
||||
// 受保护接口
|
||||
auth := r.Group("/api/auth")
|
||||
auth.Use(middleware.RequireAuth())
|
||||
{
|
||||
auth.GET("/me", authHandler.Me)
|
||||
auth.POST("/logout", authHandler.Logout)
|
||||
}
|
||||
|
||||
// 管理员接口
|
||||
admin := r.Group("/api/admin")
|
||||
admin.Use(middleware.RequireAuth(), middleware.RequireRoles(models.RoleAdmin, models.RoleSystemAdmin))
|
||||
{
|
||||
admin.GET("/users", adminHandler.ListUsers)
|
||||
admin.POST("/users", adminHandler.CreateUser)
|
||||
admin.PUT("/users/:id", adminHandler.UpdateUser)
|
||||
admin.DELETE("/users/:id", adminHandler.DeleteUser)
|
||||
admin.GET("/roles", adminHandler.ListRoles)
|
||||
}
|
||||
|
||||
return r
|
||||
}
|
||||
Reference in New Issue
Block a user