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) }