搭建用户体系

This commit is contained in:
2026-07-03 23:40:33 +08:00
parent 37a9f865ed
commit 2df19fcd02
40 changed files with 2622 additions and 0 deletions
+64
View File
@@ -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
}
+27
View File
@@ -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
}
+229
View File
@@ -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)
}
+178
View File
@@ -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
}
+116
View File
@@ -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
}
+127
View File
@@ -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
}
+58
View File
@@ -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
}