diff --git a/cmd/server/main.go b/cmd/server/main.go index 9f7e8dc..1e67ad3 100644 --- a/cmd/server/main.go +++ b/cmd/server/main.go @@ -70,10 +70,11 @@ func main() { // 5. 装配各层依赖 var ( - articleRepo = repository.NewArticleRepository(db) - brandRepo = repository.NewBrandRepository(db) - userRepo = repository.NewUserRepository(db) - snapRepo = repository.NewStreetSnapRepository(db) + articleRepo = repository.NewArticleRepository(db) + brandRepo = repository.NewBrandRepository(db) + userRepo = repository.NewUserRepository(db) + snapRepo = repository.NewStreetSnapRepository(db) + refreshRepo = repository.NewRefreshTokenRepository(db) jwtManager = jwt.NewManager(cfg.JWT.Secret, cfg.JWT.ExpireHours) @@ -81,7 +82,7 @@ func main() { brandSvc = service.NewBrandService(brandRepo) indexSvc = service.NewIndexService(brandRepo) snapSvc = service.NewStreetSnapService(snapRepo) - authSvc = service.NewAuthService(userRepo, jwtManager) + authSvc = service.NewAuthService(userRepo, refreshRepo, jwtManager, cfg.JWT.ExpireHours, cfg.JWT.RefreshExpireHours) ) engine := router.New(router.Options{ diff --git a/internal/config/config.go b/internal/config/config.go index 3b6bd9e..30e71b5 100644 --- a/internal/config/config.go +++ b/internal/config/config.go @@ -84,8 +84,9 @@ func (d DatabaseConfig) Addr() string { // JWTConfig 令牌签发配置。 type JWTConfig struct { - Secret string `yaml:"secret"` - ExpireHours int `yaml:"expire_hours"` + Secret string `yaml:"secret"` + ExpireHours int `yaml:"expire_hours"` // access token 有效期(小时) + RefreshExpireHours int `yaml:"refresh_expire_hours"` // refresh token 有效期(小时) } // UploadConfig 图片静态资源配置。 @@ -145,8 +146,9 @@ func defaultConfig() *Config { ConnMaxLifetime: 3600, }, JWT: JWTConfig{ - Secret: "dev-secret-change-me-fashion-2026", - ExpireHours: 168, + Secret: "dev-secret-change-me-fashion-2026", + ExpireHours: 2, // access token 2 小时(短命,泄漏窗口小) + RefreshExpireHours: 720, // refresh token 30 天(长命,落库可吊销) }, Upload: UploadConfig{ Dir: "./uploads", @@ -238,6 +240,7 @@ func (c *Config) applyEnv() { envStr("JWT_SECRET", &c.JWT.Secret) envInt("JWT_EXPIRE_HOURS", &c.JWT.ExpireHours) + envInt("JWT_REFRESH_EXPIRE_HOURS", &c.JWT.RefreshExpireHours) envBool("CLIENT_SIGN_ENABLED", &c.ClientSign.Enabled) envStr("CLIENT_SIGN_SECRET", &c.ClientSign.Secret) @@ -272,6 +275,9 @@ func (c *Config) normalize() { if c.JWT.ExpireHours <= 0 { c.JWT.ExpireHours = 168 } + if c.JWT.RefreshExpireHours <= 0 { + c.JWT.RefreshExpireHours = 720 + } if c.Upload.Dir == "" { c.Upload.Dir = "./uploads" } diff --git a/internal/dto/auth.go b/internal/dto/auth.go index 7edb520..0a94111 100644 --- a/internal/dto/auth.go +++ b/internal/dto/auth.go @@ -8,3 +8,30 @@ type UserPayload struct { Username string `json:"username"` Email string `json:"email"` } + +// ---------- 登录 ---------- + +// LoginRequest 登录请求体。account 支持用户名或邮箱。 +type LoginRequest struct { + Account string `json:"account" binding:"required"` + Password string `json:"password" binding:"required"` +} + +// LoginResponse 登录成功响应(双令牌)。 +type LoginResponse struct { + AccessToken string `json:"access_token"` + RefreshToken string `json:"refresh_token"` + ExpiresIn int `json:"expires_in"` // access token 有效期(秒) + User UserPayload `json:"user"` +} + +// RefreshRequest 续期请求体。 +type RefreshRequest struct { + RefreshToken string `json:"refresh_token" binding:"required"` +} + +// RefreshResponse 续期成功响应。 +type RefreshResponse struct { + AccessToken string `json:"access_token"` + ExpiresIn int `json:"expires_in"` +} diff --git a/internal/handler/auth_handler.go b/internal/handler/auth_handler.go index e6898f1..3545a19 100644 --- a/internal/handler/auth_handler.go +++ b/internal/handler/auth_handler.go @@ -3,6 +3,7 @@ package handler import ( "net/http" + "fashionapi/internal/dto" "fashionapi/internal/middleware" "fashionapi/internal/pkg/response" "fashionapi/internal/service" @@ -38,3 +39,85 @@ func (h *AuthHandler) Me(c *gin.Context) { } c.JSON(http.StatusOK, gin.H{"user": user}) } + +// Login 账号登录,校验账号密码后签发 access + refresh 双令牌。 +// +// POST /api/auth/login +// body: { account, password } +// 成功:200 { access_token, refresh_token, expires_in, user } +func (h *AuthHandler) Login(c *gin.Context) { + var req dto.LoginRequest + if err := c.ShouldBindJSON(&req); err != nil { + response.Error(c, http.StatusBadRequest, "invalid request") + return + } + + access, refresh, expiresIn, user, err := h.auth.Login(c.Request.Context(), req.Account, req.Password) + if err != nil { + fail(c, err) + return + } + c.JSON(http.StatusOK, dto.LoginResponse{ + AccessToken: access, + RefreshToken: refresh, + ExpiresIn: expiresIn, + User: *user, + }) +} + +// Refresh 用 refresh token 换取新的 access token。 +// +// POST /api/auth/refresh +// body: { refresh_token } +// 成功:200 { access_token, expires_in } +func (h *AuthHandler) Refresh(c *gin.Context) { + var req dto.RefreshRequest + if err := c.ShouldBindJSON(&req); err != nil || req.RefreshToken == "" { + response.Error(c, http.StatusBadRequest, "invalid request") + return + } + access, expiresIn, err := h.auth.Refresh(c.Request.Context(), req.RefreshToken) + if err != nil { + fail(c, err) + return + } + c.JSON(http.StatusOK, dto.RefreshResponse{AccessToken: access, ExpiresIn: expiresIn}) +} + +// Logout 吊销当前 refresh token(单设备登出)。 +// +// POST /api/auth/logout +// body: { refresh_token } +// 成功:200 { ok: true } +func (h *AuthHandler) Logout(c *gin.Context) { + var req dto.RefreshRequest + _ = c.ShouldBindJSON(&req) + if req.RefreshToken == "" { + response.Error(c, http.StatusBadRequest, "invalid request") + return + } + if err := h.auth.Logout(c.Request.Context(), req.RefreshToken); err != nil { + fail(c, err) + return + } + response.Data(c, http.StatusOK, gin.H{"ok": true}) +} + +// LogoutAll 吊销当前用户全部 refresh token(踢下线 / 全设备登出)。 +// +// POST /api/auth/logout-all +// body: { refresh_token } +// 成功:200 { revoked: <受影响行数> } +func (h *AuthHandler) LogoutAll(c *gin.Context) { + var req dto.RefreshRequest + if err := c.ShouldBindJSON(&req); err != nil || req.RefreshToken == "" { + response.Error(c, http.StatusBadRequest, "invalid request") + return + } + n, err := h.auth.RevokeAllByRefresh(c.Request.Context(), req.RefreshToken) + if err != nil { + fail(c, err) + return + } + response.Data(c, http.StatusOK, gin.H{"revoked": n}) +} diff --git a/internal/model/refresh_token.go b/internal/model/refresh_token.go new file mode 100644 index 0000000..bf16192 --- /dev/null +++ b/internal/model/refresh_token.go @@ -0,0 +1,28 @@ +package model + +import "time" + +// RefreshToken 刷新令牌(服务端状态),用于无状态 JWT 的续期与吊销(踢下线)。 +// +// refresh_token 本身是随机串,这里只存其 SHA256 哈希,原始串仅返回给客户端一次。 +// revoked=1 表示已吊销(登出 / 全设备登出 / 改密)。 +type RefreshToken struct { + ID uint64 `gorm:"primaryKey;column:id" json:"-"` + UserID uint32 `gorm:"column:user_id;index:idx_user_id" json:"-"` + TokenHash string `gorm:"column:token_hash;size:64;uniqueIndex:uk_token_hash" json:"-"` + ExpiresAt uint32 `gorm:"column:expires_at" json:"-"` + Revoked uint8 `gorm:"column:revoked" json:"-"` + UserAgent string `gorm:"column:user_agent;size:255" json:"-"` + IP string `gorm:"column:ip;size:64" json:"-"` + CreatedAt uint32 `gorm:"column:created_at" json:"-"` + + // 与 GORM 自动时间字段区分:CreatedAt 由业务层显式写入 Unix 秒。 +} + +// TableName 指定表名。 +func (RefreshToken) TableName() string { return "refresh_tokens" } + +// IsValid 未被吊销且未过期。 +func (r *RefreshToken) IsValid() bool { + return r.Revoked == 0 && r.ExpiresAt >= uint32(time.Now().Unix()) +} diff --git a/internal/pkg/jwt/jwt.go b/internal/pkg/jwt/jwt.go index ab7349c..ff31152 100644 --- a/internal/pkg/jwt/jwt.go +++ b/internal/pkg/jwt/jwt.go @@ -44,6 +44,21 @@ func NewManager(secret string, expireHours int) *Manager { // 将来接入登录时恢复:构造 jwtlib.MapClaims{"uid","username","email","exp"}, // 再用 jwtlib.NewWithClaims(jwtlib.SigningMethodHS256, claims).SignedString(m.secret) 签名即可。 +// Generate 签发令牌(登录成功后调用)。 +// +// 构造与 Parse 对称的 Claims(uid / username / email / exp),用 HS256 签名。 +// exp 由 Manager 的 expire 决定(默认 7 天,见 NewManager)。 +func (m *Manager) Generate(userID uint32, username, email string) (string, error) { + claims := jwtlib.MapClaims{ + "uid": float64(userID), + "username": username, + "email": email, + "exp": time.Now().Add(m.expire).Unix(), + } + token := jwtlib.NewWithClaims(jwtlib.SigningMethodHS256, claims) + return token.SignedString(m.secret) +} + // Parse 校验并解析令牌。 func (m *Manager) Parse(tokenStr string) (*Claims, error) { token, err := jwtlib.Parse(tokenStr, func(t *jwtlib.Token) (any, error) { diff --git a/internal/repository/refresh_token_repository.go b/internal/repository/refresh_token_repository.go new file mode 100644 index 0000000..978f484 --- /dev/null +++ b/internal/repository/refresh_token_repository.go @@ -0,0 +1,72 @@ +package repository + +import ( + "context" + "errors" + "time" + + "fashionapi/internal/model" + + "gorm.io/gorm" +) + +// RefreshTokenRepository 刷新令牌数据访问接口。 +// +// refresh token 落库是该账号体系能做「续期 + 踢下线」的关键: +// 服务端持有每个 login 会话的行,吊销它即可强行让对应用户重新登录。 +type RefreshTokenRepository interface { + // Create 写入一条刷新令牌记录。 + Create(ctx context.Context, rt *model.RefreshToken) error + // FindByHash 按 SHA256(raw) 查找,不存在时返回 ErrNotFound。 + FindByHash(ctx context.Context, hash string) (*model.RefreshToken, error) + // Revoke 将指定 id 的令牌置为已吊销。 + Revoke(ctx context.Context, id uint64) error + // RevokeAllByUser 吊销某用户的所有令牌(全设备登出 / 踢下线),返回受影响行数。 + RevokeAllByUser(ctx context.Context, userID uint32) (int64, error) + // DeleteExpired 清理已过期或已吊销的令牌,避免表无限增长。 + DeleteExpired(ctx context.Context) (int64, error) +} + +type refreshTokenRepository struct { + db *gorm.DB +} + +// NewRefreshTokenRepository 创建刷新令牌仓储。 +func NewRefreshTokenRepository(db *gorm.DB) RefreshTokenRepository { + return &refreshTokenRepository{db: db} +} + +func (r *refreshTokenRepository) Create(ctx context.Context, rt *model.RefreshToken) error { + return r.db.WithContext(ctx).Create(rt).Error +} + +func (r *refreshTokenRepository) FindByHash(ctx context.Context, hash string) (*model.RefreshToken, error) { + var rt model.RefreshToken + err := r.db.WithContext(ctx).Where("token_hash = ?", hash).First(&rt).Error + if err != nil { + if errors.Is(err, gorm.ErrRecordNotFound) { + return nil, ErrNotFound + } + return nil, err + } + return &rt, nil +} + +func (r *refreshTokenRepository) Revoke(ctx context.Context, id uint64) error { + return r.db.WithContext(ctx).Model(&model.RefreshToken{}). + Where("id = ?", id).Update("revoked", 1).Error +} + +func (r *refreshTokenRepository) RevokeAllByUser(ctx context.Context, userID uint32) (int64, error) { + res := r.db.WithContext(ctx).Model(&model.RefreshToken{}). + Where("user_id = ?", userID).Update("revoked", 1) + return res.RowsAffected, res.Error +} + +func (r *refreshTokenRepository) DeleteExpired(ctx context.Context) (int64, error) { + now := uint32(time.Now().Unix()) + res := r.db.WithContext(ctx). + Where("expires_at < ? OR revoked = ?", now, 1). + Delete(&model.RefreshToken{}) + return res.RowsAffected, res.Error +} diff --git a/internal/repository/user_repository.go b/internal/repository/user_repository.go index 17e1454..6fe4e1d 100644 --- a/internal/repository/user_repository.go +++ b/internal/repository/user_repository.go @@ -16,6 +16,8 @@ import ( type UserRepository interface { // FindByID 按主键查找用户。不存在时返回 ErrNotFound。 FindByID(ctx context.Context, id uint32) (*model.User, error) + // FindByAccount 按用户名或邮箱查找(忽略已删除账号),供登录校验。 + FindByAccount(ctx context.Context, account string) (*model.User, error) } type userRepository struct { @@ -33,6 +35,15 @@ func (r *userRepository) FindByID(ctx context.Context, id uint32) (*model.User, return wrapUser(&u, err) } +func (r *userRepository) FindByAccount(ctx context.Context, account string) (*model.User, error) { + var u model.User + err := r.db.WithContext(ctx). + Where("username = ? OR email = ?", account, account). + Where("is_deleted = ?", 0). + First(&u).Error + return wrapUser(&u, err) +} + func wrapUser(u *model.User, err error) (*model.User, error) { if err != nil { if errors.Is(err, gorm.ErrRecordNotFound) { diff --git a/internal/router/router.go b/internal/router/router.go index d68b64e..393b8f6 100644 --- a/internal/router/router.go +++ b/internal/router/router.go @@ -61,13 +61,21 @@ func New(opt Options) *gin.Engine { // 街拍详情:同走秀详情,单条枚举防护由 HashID 承担,不走前端签名。 api.GET("/public/street-snaps/:id", opt.StreetSnap.Detail) - // 账号体系:当前只保留「取当前用户」(JWT 无状态,无服务端登出路由)。 + // 账号体系:双令牌(access 短命无状态 + refresh 落库可吊销)。 // - // 注册 / 登录接口已从路由移除:前端没有登录注册入口,对外开放即是垃圾账号与撞库入口。 - // handler / service / repository / dto 中对应实现一并删除。 - // 将来要接入时:补回 POST /auth/register、POST /auth/login,并同步前端 src/lib/api.ts 与登录页。 + // 登录/刷新/登出均为公开端点(无需 Bearer,否则拿不到 token 或刷新不了)。 + // /me 需 Bearer(middleware.Auth);踢下线(logout-all)由 refresh token 反查用户,亦公开。 auth := api.Group("/auth") { + // 登录:公开端点,无需鉴权中间件(否则永远进不来) + auth.POST("/login", opt.Auth.Login) + // 续期:用 refresh token 换新的 access token + auth.POST("/refresh", opt.Auth.Refresh) + // 单设备登出:吊销当前 refresh token + auth.POST("/logout", opt.Auth.Logout) + // 全设备登出 / 踢下线:吊销该用户全部 refresh token + auth.POST("/logout-all", opt.Auth.LogoutAll) + // 取当前用户(需 Bearer) auth.GET("/me", middleware.Auth(opt.JWT), opt.Auth.Me) } } diff --git a/internal/service/auth_service.go b/internal/service/auth_service.go index 62fb94c..1f49347 100644 --- a/internal/service/auth_service.go +++ b/internal/service/auth_service.go @@ -2,30 +2,59 @@ package service import ( "context" + "crypto/rand" + "crypto/sha256" + "encoding/hex" + "time" "fashionapi/internal/dto" "fashionapi/internal/model" "fashionapi/internal/pkg/jwt" "fashionapi/internal/repository" + "golang.org/x/crypto/bcrypt" ) // AuthService 账号体系业务接口。 // -// 当前只保留 Me(取当前用户)。注册 / 登录已随路由一并下线:前端没有登录注册入口, -// 对外开放即是垃圾账号与撞库入口。将来接入时补回 Register / Login 与 issue 令牌签发逻辑, -// 并同步前端 src/lib/api.ts 与登录页。 +// 采用「access token(短命,无状态 JWT)+ refresh token(长命,落库可吊销)」双令牌模式: +// - Login 签发两者;access 用于日常请求,refresh 仅在续期时使用。 +// - Refresh 用 refresh token 换发新 access(refresh 复用、不轮换,足够内部工具)。 +// - Logout 吊销当前 refresh;RevokeAllByRefresh 吊销该用户全部 refresh(踢下线)。 +// 由于 refresh 落库,服务端得以单方面作废会话——这是纯无状态 JWT 做不到的。 type AuthService interface { Me(ctx context.Context, userID uint32) (*dto.UserPayload, error) + // Login 校验账号密码,签发 access + refresh;并吊销该用户既有会话(同账号互斥登录)。account 可为用户名或邮箱。 + Login(ctx context.Context, account, password string) (access, refresh string, expiresIn int, user *dto.UserPayload, err error) + // Refresh 用 refresh token 换取新的 access token。 + Refresh(ctx context.Context, refreshToken string) (access string, expiresIn int, err error) + // Logout 吊销当前 refresh token(单设备登出),幂等。 + Logout(ctx context.Context, refreshToken string) error + // RevokeAllByRefresh 根据 refresh token 找到所属用户,吊销其全部 refresh(踢下线)。 + RevokeAllByRefresh(ctx context.Context, refreshToken string) (int64, error) } type authService struct { - users repository.UserRepository - jwt *jwt.Manager + users repository.UserRepository + refresh repository.RefreshTokenRepository + jwt *jwt.Manager + accessTTLSeconds int + refreshExpireHours int } // NewAuthService 创建账号服务。 -func NewAuthService(users repository.UserRepository, jwtManager *jwt.Manager) AuthService { - return &authService{users: users, jwt: jwtManager} +func NewAuthService( + users repository.UserRepository, + refresh repository.RefreshTokenRepository, + jwtManager *jwt.Manager, + accessExpireHours, refreshExpireHours int, +) AuthService { + return &authService{ + users: users, + refresh: refresh, + jwt: jwtManager, + accessTTLSeconds: accessExpireHours * 3600, + refreshExpireHours: refreshExpireHours, + } } // Me 用令牌中的 uid 取回最新用户信息。 @@ -39,6 +68,135 @@ func (s *authService) Me(ctx context.Context, userID uint32) (*dto.UserPayload, return &p, nil } +// Login 校验账号密码并签发双令牌。account 可以是用户名或邮箱。 +// +// 登录即顶旧:签发前先吊销该用户全部既有 refresh 会话,保证同一账号同一时刻 +// 只有本条新会话有效(互斥登录)。账号不存在与密码错误统一返回 unauthorized +// (避免泄露账号是否存在)。 +func (s *authService) Login(ctx context.Context, account, password string) (string, string, int, *dto.UserPayload, error) { + user, err := s.users.FindByAccount(ctx, account) + if err != nil { + return "", "", 0, nil, unauthorized("invalid credentials") + } + if err := bcrypt.CompareHashAndPassword([]byte(user.PasswordHash), []byte(password)); err != nil { + return "", "", 0, nil, unauthorized("invalid credentials") + } + // 同账号互斥登录:本次登录先吊销该用户所有既有 refresh 会话, + // 仅保留随后新签发的这一条,从而实现「同一账号同时只能有一个人登录」。 + if _, err := s.refresh.RevokeAllByUser(ctx, user.ID); err != nil { + return "", "", 0, nil, internalErr("revoke stale sessions: " + err.Error()) + } + access, err := s.jwt.Generate(user.ID, user.Username, user.Email) + if err != nil { + return "", "", 0, nil, err + } + refresh, err := s.issueRefresh(ctx, user.ID) + if err != nil { + return "", "", 0, nil, err + } + p := toUserPayload(*user) + return access, refresh, s.accessTTLSeconds, &p, nil +} + +// Refresh 用 refresh token 换取新的 access token(refresh 复用)。 +func (s *authService) Refresh(ctx context.Context, refreshToken string) (string, int, error) { + if refreshToken == "" { + return "", 0, unauthorized("missing refresh token") + } + rt, err := s.lookupRefresh(ctx, refreshToken) + if err != nil { + return "", 0, err + } + user, err := s.users.FindByID(ctx, rt.UserID) + if err != nil { + return "", 0, unauthorized("invalid refresh token") + } + access, err := s.jwt.Generate(user.ID, user.Username, user.Email) + if err != nil { + return "", 0, err + } + return access, s.accessTTLSeconds, nil +} + +// Logout 吊销当前 refresh token;令牌不存在/已失效时视为幂等成功。 +func (s *authService) Logout(ctx context.Context, refreshToken string) error { + if refreshToken == "" { + return nil + } + rt, err := s.findRefresh(ctx, refreshToken) + if err != nil { + return nil + } + return s.refresh.Revoke(ctx, rt.ID) +} + +// RevokeAllByRefresh 找到 refresh 所属用户,吊销其全部 refresh(踢下线)。 +func (s *authService) RevokeAllByRefresh(ctx context.Context, refreshToken string) (int64, error) { + if refreshToken == "" { + return 0, unauthorized("missing refresh token") + } + rt, err := s.findRefresh(ctx, refreshToken) + if err != nil { + return 0, unauthorized("invalid refresh token") + } + return s.refresh.RevokeAllByUser(ctx, rt.UserID) +} + +// issueRefresh 生成随机 refresh token,存哈希入表,返回原始串。 +func (s *authService) issueRefresh(ctx context.Context, userID uint32) (string, error) { + raw, hash, err := newRefreshToken() + if err != nil { + return "", internalErr("generate refresh token: " + err.Error()) + } + now := uint32(time.Now().Unix()) + rt := &model.RefreshToken{ + UserID: userID, + TokenHash: hash, + ExpiresAt: now + uint32(s.refreshExpireHours)*3600, + Revoked: 0, + CreatedAt: now, + } + if err := s.refresh.Create(ctx, rt); err != nil { + return "", internalErr("store refresh token: " + err.Error()) + } + return raw, nil +} + +// findRefresh 按原始串取哈希查库(用于登出/吊销),找不到返回 ErrNotFound。 +func (s *authService) findRefresh(ctx context.Context, refreshToken string) (*model.RefreshToken, error) { + hash := hashToken(refreshToken) + return s.refresh.FindByHash(ctx, hash) +} + +// lookupRefresh 同 findRefresh,但额外校验有效性(未吊销且未过期)。 +func (s *authService) lookupRefresh(ctx context.Context, refreshToken string) (*model.RefreshToken, error) { + rt, err := s.findRefresh(ctx, refreshToken) + if err != nil { + return nil, unauthorized("invalid refresh token") + } + if !rt.IsValid() { + return nil, unauthorized("refresh token expired or revoked") + } + return rt, nil +} + +// newRefreshToken 生成 32 字节随机串及其 SHA256 哈希(仅哈希落库)。 +func newRefreshToken() (raw, hash string, err error) { + b := make([]byte, 32) + if _, err = rand.Read(b); err != nil { + return "", "", err + } + raw = hex.EncodeToString(b) + sum := sha256.Sum256([]byte(raw)) + return raw, hex.EncodeToString(sum[:]), nil +} + +// hashToken 计算 refresh token 的存储哈希。 +func hashToken(raw string) string { + sum := sha256.Sum256([]byte(raw)) + return hex.EncodeToString(sum[:]) +} + func toUserPayload(u model.User) dto.UserPayload { return dto.UserPayload{ID: u.ID, Username: u.Username, Email: u.Email} } diff --git a/scripts/migrate_refresh/main.go b/scripts/migrate_refresh/main.go new file mode 100644 index 0000000..4bca971 --- /dev/null +++ b/scripts/migrate_refresh/main.go @@ -0,0 +1,41 @@ +// Command migrate_refresh 一次性建表脚本(与 scripts/sql/*.sql 同源,手写 DDL 以保持与项目"SQL 托管表结构"约定一致)。 +// +// 运行:go run ./scripts/migrate_refresh +// 可用环境变量 DB_DSN 覆盖连接串(默认 root:root@127.0.0.1:3306/db_dev)。 +package main + +import ( + "log" + "os" + + "gorm.io/driver/mysql" + "gorm.io/gorm" +) + +func main() { + dsn := os.Getenv("DB_DSN") + if dsn == "" { + dsn = "root:root@tcp(127.0.0.1:3306)/db_dev?charset=utf8mb4&parseTime=True&loc=Local" + } + db, err := gorm.Open(mysql.Open(dsn), &gorm.Config{}) + if err != nil { + log.Fatalf("open db: %v", err) + } + stmt := `CREATE TABLE IF NOT EXISTS refresh_tokens ( + id BIGINT UNSIGNED NOT NULL AUTO_INCREMENT, + user_id INT UNSIGNED NOT NULL, + token_hash VARCHAR(64) NOT NULL, + expires_at INT UNSIGNED NOT NULL DEFAULT 0, + revoked TINYINT UNSIGNED NOT NULL DEFAULT 0, + user_agent VARCHAR(255) NOT NULL DEFAULT '', + ip VARCHAR(64) NOT NULL DEFAULT '', + created_at INT UNSIGNED NOT NULL DEFAULT 0, + PRIMARY KEY (id), + UNIQUE KEY uk_token_hash (token_hash), + KEY idx_user_id (user_id) + ) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4;` + if err := db.Exec(stmt).Error; err != nil { + log.Fatalf("create refresh_tokens: %v", err) + } + log.Println("✓ refresh_tokens ready") +} diff --git a/scripts/seed_users/main.go b/scripts/seed_users/main.go new file mode 100644 index 0000000..f7229e2 --- /dev/null +++ b/scripts/seed_users/main.go @@ -0,0 +1,95 @@ +// Command seed_users 预置内部账号(运营/编辑用,不开放公开注册)。 +// +// 项目已移除 AutoMigrate,表结构由 scripts/sql/005_create_users.sql 托管; +// 本脚本用 CREATE TABLE IF NOT EXISTS 兜底建表,再 upsert 一个内部账号。 +// +// 用法: +// go run ./scripts/seed_users +// SEED_ADMIN_PASSWORD='你的强密码' go run ./scripts/seed_users +// +// 默认账号 admin / admin@studio.local,密码取环境变量 SEED_ADMIN_PASSWORD, +// 缺省回落强密码常量。生产部署请务必通过环境变量指定并尽快修改。 +package main + +import ( + "errors" + "flag" + "log" + "os" + + "golang.org/x/crypto/bcrypt" + "gorm.io/gorm" + + "fashionapi/internal/config" + "fashionapi/internal/database" + "fashionapi/internal/model" +) + +func main() { + configPath := flag.String("config", "", "配置文件路径,默认查找 configs/config.yml") + flag.Parse() + + cfg, err := config.Load(*configPath) + if err != nil { + log.Fatalf("✗ 加载配置失败: %v", err) + } + db, err := database.New(cfg.Database) + if err != nil { + log.Fatalf("✗ 连接数据库失败: %v", err) + } + defer func() { + if e := database.Close(db); e != nil { + log.Printf("! 关闭数据库连接失败: %v", e) + } + }() + + // 1. 兜底建表(与 scripts/sql/005_create_users.sql 完全一致) + createSQL := `CREATE TABLE IF NOT EXISTS users ( + id int unsigned NOT NULL AUTO_INCREMENT, + created_at int unsigned NOT NULL DEFAULT 0, + updated_at int unsigned NOT NULL DEFAULT 0, + username varchar(64) NOT NULL, + email varchar(191) NOT NULL, + password_hash varchar(255) NOT NULL DEFAULT '', + is_deleted tinyint unsigned NOT NULL DEFAULT 0, + PRIMARY KEY (id), + UNIQUE KEY uk_username (username), + UNIQUE KEY uk_email (email) + ) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4;` + if e := db.Exec(createSQL).Error; e != nil { + log.Fatalf("✗ 建表失败: %v", e) + } + log.Println("✓ users 表就绪") + + // 2. 预置内部账号 + const username = "admin" + const email = "admin@studio.local" + password := os.Getenv("SEED_ADMIN_PASSWORD") + if password == "" { + password = "Studio#2026!Admin" + } + hash, e := bcrypt.GenerateFromPassword([]byte(password), bcrypt.DefaultCost) + if e != nil { + log.Fatalf("✗ 生成密码哈希失败: %v", e) + } + + var existing model.User + qerr := db.Where("username = ?", username).First(&existing).Error + switch { + case errors.Is(qerr, gorm.ErrRecordNotFound): + user := model.User{Username: username, Email: email, PasswordHash: string(hash)} + if cerr := db.Create(&user).Error; cerr != nil { + log.Fatalf("✗ 创建内部账号失败: %v", cerr) + } + log.Printf("✓ 已创建内部账号: %s / %s", username, email) + case qerr != nil: + log.Fatalf("✗ 查询账号失败: %v", qerr) + default: + existing.PasswordHash = string(hash) + if uerr := db.Save(&existing).Error; uerr != nil { + log.Fatalf("✗ 更新账号密码失败: %v", uerr) + } + log.Printf("✓ 内部账号已存在,已刷新密码哈希: %s", username) + } + log.Printf("ℹ 登录账号: %s 密码: %s(请尽快修改默认密码)", username, password) +} diff --git a/scripts/sql/005_create_users.sql b/scripts/sql/005_create_users.sql new file mode 100644 index 0000000..2d501da --- /dev/null +++ b/scripts/sql/005_create_users.sql @@ -0,0 +1,14 @@ +-- 账号体系用户表(内部运营/编辑账号,不开放公开注册)。 +-- 与 scripts/seed_users/main.go 中内置的建表语句保持一致。 +CREATE TABLE IF NOT EXISTS `users` ( + `id` int unsigned NOT NULL AUTO_INCREMENT, + `created_at` int unsigned NOT NULL DEFAULT 0, + `updated_at` int unsigned NOT NULL DEFAULT 0, + `username` varchar(64) NOT NULL, + `email` varchar(191) NOT NULL, + `password_hash` varchar(255) NOT NULL DEFAULT '', + `is_deleted` tinyint unsigned NOT NULL DEFAULT 0, + PRIMARY KEY (`id`), + UNIQUE KEY `uk_username` (`username`), + UNIQUE KEY `uk_email` (`email`) +) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4; diff --git a/scripts/sql/006_create_refresh_tokens.sql b/scripts/sql/006_create_refresh_tokens.sql new file mode 100644 index 0000000..65132bc --- /dev/null +++ b/scripts/sql/006_create_refresh_tokens.sql @@ -0,0 +1,15 @@ +-- 006: refresh_tokens —— 刷新令牌表(双令牌续期 + 踢下线的服务端状态) +-- 与项目约定一致:表结构由 SQL 托管,不依赖 GORM AutoMigrate。 +CREATE TABLE IF NOT EXISTS refresh_tokens ( + id BIGINT UNSIGNED NOT NULL AUTO_INCREMENT, + user_id INT UNSIGNED NOT NULL, + token_hash VARCHAR(64) NOT NULL, + expires_at INT UNSIGNED NOT NULL DEFAULT 0, + revoked TINYINT UNSIGNED NOT NULL DEFAULT 0, + user_agent VARCHAR(255) NOT NULL DEFAULT '', + ip VARCHAR(64) NOT NULL DEFAULT '', + created_at INT UNSIGNED NOT NULL DEFAULT 0, + PRIMARY KEY (id), + UNIQUE KEY uk_token_hash (token_hash), + KEY idx_user_id (user_id) +) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4;