Files
blog-server/internal/service/sys_user.go
2026-07-22 17:50:33 +08:00

291 lines
6.6 KiB
Go

package service
import (
"context"
db "server/internal/db/sqlc"
"server/internal/middleware"
"server/internal/model/common"
"server/internal/model/request"
"server/internal/model/response"
"server/internal/pkg/cache"
"server/internal/pkg/dberr"
"server/internal/pkg/errs"
"server/internal/pkg/httputil"
"github.com/jackc/pgx/v5"
"github.com/jackc/pgx/v5/pgxpool"
"golang.org/x/crypto/bcrypt"
"golang.org/x/sync/errgroup"
)
type SysUserService struct {
queries *db.Queries
pool *pgxpool.Pool
jwt *middleware.JWTMiddleware
cache *cache.Caches
}
func NewSysUserService(queries *db.Queries, pool *pgxpool.Pool, jwt *middleware.JWTMiddleware, cache *cache.Caches) *SysUserService {
return &SysUserService{
queries: queries,
pool: pool,
jwt: jwt,
cache: cache,
}
}
// generatePasswordHash 生成密码哈希
func generatePasswordHash(password string) (string, error) {
hashed, err := bcrypt.GenerateFromPassword([]byte(password), bcrypt.DefaultCost)
if err != nil {
return "", err
}
return string(hashed), nil
}
// comparePasswordHash 比较密码哈希
func comparePasswordHash(passwordHash, inputPassword string) error {
return bcrypt.CompareHashAndPassword([]byte(passwordHash), []byte(inputPassword))
}
func (s *SysUserService) Login(ctx context.Context, req request.LoginRequest) (*response.LoginResponse, error) {
user, err := s.queries.GetSysUserByAccount(ctx, req.Account)
if err != nil {
return nil, errs.ErrInvalidCredentials
}
if err = comparePasswordHash(user.PasswordHash, req.Password); err != nil {
return nil, errs.ErrInvalidCredentials
}
accessToken, accessTokenExp, err := s.jwt.GenerateAccessToken(user.ID)
if err != nil {
return nil, err
}
refreshToken, refreshTokenExp, err := s.jwt.GenerateRefreshToken(user.ID)
if err != nil {
return nil, err
}
return &response.LoginResponse{
AccessToken: accessToken,
AccessTokenExp: accessTokenExp,
RefreshToken: refreshToken,
RefreshTokenExp: refreshTokenExp,
}, nil
}
func (s *SysUserService) RefreshToken(ctx context.Context, refreshToken string) (*response.LoginResponse, error) {
claims, err := s.jwt.ParseRefreshToken(refreshToken)
if err != nil {
return nil, errs.ErrInvalidRefreshToken
}
accessToken, accessTokenExp, err := s.jwt.GenerateAccessToken(claims.UserID)
if err != nil {
return nil, err
}
return &response.LoginResponse{
AccessToken: accessToken,
AccessTokenExp: accessTokenExp,
}, nil
}
func (s *SysUserService) GetUserInfo(ctx context.Context, id int32, isAdmin bool) (*response.SysUserInfo, error) {
g, ctx := errgroup.WithContext(ctx)
var (
user db.GetSysUserByIDRow
roles []db.SysRole
menus []db.SysMenu
)
g.Go(func() error {
u, err := s.queries.GetSysUserByID(ctx, id)
if err != nil {
return dberr.MapNoRows(err, errs.ErrUserNotFound)
}
user = u
return nil
})
g.Go(func() error {
r, err := s.queries.GetSysUserRoles(ctx, id)
if err != nil {
return err
}
roles = r
return nil
})
g.Go(func() error {
var (
m []db.SysMenu
err error
)
if isAdmin {
m, err = s.queries.GetSysAdminMenus(ctx)
} else {
m, err = s.queries.GetSysUserMenus(ctx, id)
}
if err != nil {
return err
}
menus = m
return nil
})
if err := g.Wait(); err != nil {
return nil, err
}
// 处理角色
userInfo := response.NewSysUserInfo(user, roles, menus)
return userInfo, nil
}
func (s *SysUserService) ListPage(ctx context.Context, p *common.Pagination) ([]db.ListSysUsersRow, int64, error) {
params := db.ListSysUsersParams{
Limit: p.PageSize,
Offset: (p.Page - 1) * p.PageSize,
}
total, err := s.queries.CountSysUsers(ctx)
if err != nil {
return nil, 0, err
}
users, err := s.queries.ListSysUsers(ctx, params)
if err != nil {
return nil, 0, err
}
// 处理每个用户的头像URL
for i := range users {
url := httputil.BuildFileUrl(users[i].AvatarUrl)
users[i].AvatarUrl = &url
}
return users, total, nil
}
func (s *SysUserService) GetRoles(ctx context.Context, id int32) ([]db.SysRole, error) {
// 先查询用户是否存在
_, err := s.queries.GetSysUserByID(ctx, id)
if err != nil {
return nil, dberr.MapNoRows(err, errs.ErrUserNotFound)
}
return s.queries.GetSysUserRoles(ctx, id)
}
func (s *SysUserService) Create(ctx context.Context, req request.CreateSysUserRequest) error {
passwordHash, err := generatePasswordHash(req.Password)
if err != nil {
return err
}
user := db.CreateSysUserParams{
Account: req.Account,
Username: req.Username,
PasswordHash: passwordHash,
AvatarID: req.AvatarID,
}
if err = s.queries.CreateSysUser(ctx, user); err != nil {
return dberr.MapUniqueViolation(err, dberr.SysUserAccountKey, errs.ErrAccountAlreadyExists)
}
return nil
}
func (s *SysUserService) Update(ctx context.Context, id int32, req request.UpdateSysUserRequest) error {
user := db.UpdateSysUserParams{
Username: req.Username,
ID: id,
}
if req.AvatarID.Set {
user.UpdateAvatarID = true
if req.AvatarID.Valid {
user.AvatarID = &req.AvatarID.Value
}
}
rows, err := s.queries.UpdateSysUser(ctx, user)
return dberr.MapRowsAffected(rows, err, errs.ErrUserNotFound)
}
func (s *SysUserService) SetRoles(ctx context.Context, userID int32, req request.SetSysUserRolesRequest) error {
// 先查询用户是否存在
_, err := s.queries.GetSysUserByID(ctx, userID)
if err != nil {
return dberr.MapNoRows(err, errs.ErrUserNotFound)
}
tx, err := s.pool.BeginTx(ctx, pgx.TxOptions{})
if err != nil {
return err
}
defer func(ctx context.Context) {
_ = tx.Rollback(ctx)
}(ctx)
q := db.New(tx)
if err = q.ClearSysUserRoles(ctx, userID); err != nil {
return err
}
for _, roleID := range req.RoleIDs {
if err = q.CreateSysUserRole(ctx, db.CreateSysUserRoleParams{
UserID: userID,
RoleID: roleID,
}); err != nil {
return err
}
}
if err = tx.Commit(ctx); err != nil {
return err
}
// 清理缓存
s.cache.ClearSysUserCache(userID)
return nil
}
func (s *SysUserService) UpdatePassword(ctx context.Context, id int32, req request.UpdateSysUserPassword) error {
passwordHash, err := generatePasswordHash(req.Password)
if err != nil {
return err
}
params := db.UpdateSysUserPasswordParams{
ID: id,
PasswordHash: passwordHash,
}
rows, err := s.queries.UpdateSysUserPassword(ctx, params)
return dberr.MapRowsAffected(rows, err, errs.ErrUserNotFound)
}
func (s *SysUserService) Delete(ctx context.Context, id int32) error {
if id == 1 {
return errs.ErrCannotDeleteSuperAdmin
}
// 清理缓存
s.cache.ClearSysUserCache(id)
rows, err := s.queries.DeleteSysUser(ctx, id)
return dberr.MapRowsAffected(rows, err, errs.ErrUserNotFound)
}