chore: initial commit

This commit is contained in:
2026-07-22 17:50:33 +08:00
parent 82936a8987
commit ee9f50b859
107 changed files with 8346 additions and 0 deletions

View File

@@ -0,0 +1,74 @@
package service
import (
"context"
db "server/internal/db/sqlc"
"server/internal/model/common"
"server/internal/model/request"
"server/internal/pkg/dberr"
"server/internal/pkg/errs"
)
type CategoryService struct {
queries *db.Queries
}
func NewCategoryService(queries *db.Queries) *CategoryService {
return &CategoryService{
queries: queries,
}
}
func (s *CategoryService) ListPage(ctx context.Context, p *common.Pagination) ([]db.Category, int64, error) {
params := db.ListCategoriesParams{
Limit: p.PageSize,
Offset: (p.Page - 1) * p.PageSize,
}
total, err := s.queries.CountCategories(ctx)
if err != nil {
return nil, 0, err
}
list, err := s.queries.ListCategories(ctx, params)
if err != nil {
return nil, 0, err
}
return list, total, nil
}
func (s *CategoryService) ListAll(ctx context.Context) ([]db.Category, error) {
return s.queries.ListAllCategories(ctx)
}
func (s *CategoryService) Create(ctx context.Context, req request.CreateCategoryRequest) error {
params := db.CreateCategoryParams{
Name: req.Name,
Code: req.Code,
}
err := s.queries.CreateCategory(ctx, params)
return dberr.MapUniqueViolation(err, dberr.CategoryCodeKey, errs.ErrCategoryCodeAlreadyExists)
}
func (s *CategoryService) Update(ctx context.Context, id int32, req request.UpdateCategoryRequest) error {
params := db.UpdateCategoryParams{
ID: id,
Name: req.Name,
Code: req.Code,
}
rows, err := s.queries.UpdateCategory(ctx, params)
// 先判断数据条目是否存在
if err = dberr.MapRowsAffected(rows, err, errs.ErrCategoryNotFound); err != nil {
// 再判断code是否重复
return dberr.MapUniqueViolation(err, dberr.CategoryCodeKey, errs.ErrCategoryCodeAlreadyExists)
}
return nil
}
func (s *CategoryService) Delete(ctx context.Context, id int32) error {
rows, err := s.queries.DeleteCategory(ctx, id)
return dberr.MapRowsAffected(rows, err, errs.ErrCategoryNotFound)
}

View File

@@ -0,0 +1,19 @@
package service
import (
"go.uber.org/fx"
)
var Module = fx.Module("services",
fx.Provide(
NewSysUserService,
NewSysRoleService,
NewSysMenuService,
NewSysApiService,
NewSysFileService,
NewSysPostService,
NewPostService,
NewCategoryService,
),
)

113
internal/service/post.go Normal file
View File

@@ -0,0 +1,113 @@
package service
import (
"context"
"net/netip"
db "server/internal/db/sqlc"
"server/internal/model/common"
"server/internal/model/response"
"server/internal/pkg/dberr"
"server/internal/pkg/errs"
"server/internal/pkg/httputil"
)
type PostService struct {
queries *db.Queries
}
func NewPostService(queries *db.Queries) *PostService {
return &PostService{
queries: queries,
}
}
func (s *PostService) ListPage(ctx context.Context, p *common.Pagination) ([]db.ListPublishedPostsRow, int64, error) {
params := db.ListPublishedPostsParams{
Limit: p.PageSize,
Offset: (p.Page - 1) * p.PageSize,
}
total, err := s.queries.CountPublishedPosts(ctx)
if err != nil {
return nil, 0, err
}
list, err := s.queries.ListPublishedPosts(ctx, params)
if err != nil {
return nil, 0, err
}
for i := range list {
url := httputil.BuildFileUrl(list[i].Cover)
list[i].Cover = &url
}
return list, total, nil
}
func (s *PostService) GetPost(ctx context.Context, slug string, ip netip.Addr) (*db.GetPublicPostBySlugRow, error) {
post, err := s.queries.GetPublicPostBySlug(ctx, slug)
if err != nil {
return nil, dberr.MapNoRows(err, errs.ErrPostNotFound)
}
_ = s.queries.IncrementPostStatsView(ctx, db.IncrementPostStatsViewParams{
PostID: post.ID,
Ip: ip,
})
return &post, nil
}
func (s *PostService) ListCategoryStats(ctx context.Context) ([]db.ListCategoryStatsRow, error) {
return s.queries.ListCategoryStats(ctx)
}
func (s *PostService) ListArchives(ctx context.Context) ([]response.ArchiveYear, error) {
list, err := s.queries.ListArchives(ctx)
if err != nil {
return nil, err
}
archive := make([]response.ArchiveYear, 0)
for _, item := range list {
y := item.PublishedAt.Year()
m := int(item.PublishedAt.Month())
if len(archive) == 0 || archive[len(archive)-1].Year != y {
archive = append(archive, response.ArchiveYear{
Year: y,
Total: 0,
ArchiveMonth: make([]response.ArchiveMonth, 0),
})
}
// 获取索引
lastYearIndex := len(archive) - 1
months := archive[lastYearIndex].ArchiveMonth
if len(months) == 0 || months[len(months)-1].Month != m {
archive[lastYearIndex].ArchiveMonth = append(archive[lastYearIndex].ArchiveMonth, response.ArchiveMonth{
Month: m,
Archive: make([]response.ArchivePost, 0),
})
}
lastMonthIndex := len(archive[lastYearIndex].ArchiveMonth) - 1
archive[lastYearIndex].ArchiveMonth[lastMonthIndex].Archive =
append(archive[lastYearIndex].ArchiveMonth[lastMonthIndex].Archive, response.ArchivePost{
ID: item.ID,
Slug: item.Slug,
Title: item.Title,
PublishedAt: item.PublishedAt,
PublishedAtDisplay: item.PublishedAt.Format("01-02"),
CategoryName: *item.CategoryName,
})
archive[lastYearIndex].Total++
}
return archive, nil
}

158
internal/service/sys_api.go Normal file
View File

@@ -0,0 +1,158 @@
package service
import (
"context"
db "server/internal/db/sqlc"
"server/internal/model/common"
"server/internal/model/enum"
"server/internal/model/request"
"server/internal/pkg/cache"
"server/internal/pkg/dberr"
"server/internal/pkg/errs"
"github.com/jackc/pgx/v5"
"github.com/jackc/pgx/v5/pgxpool"
)
type SysApiService struct {
queries *db.Queries
pool *pgxpool.Pool
cache *cache.Caches
}
func NewSysApiService(queries *db.Queries, pool *pgxpool.Pool, cache *cache.Caches) *SysApiService {
return &SysApiService{queries: queries, pool: pool, cache: cache}
}
func (s *SysApiService) ListPage(ctx context.Context, p *common.Pagination) ([]db.SysApi, int64, error) {
params := db.GetSysApisParams{
Limit: p.PageSize,
Offset: (p.Page - 1) * p.PageSize,
}
total, err := s.queries.CountSysApis(ctx)
if err != nil {
return nil, 0, err
}
list, err := s.queries.GetSysApis(ctx, params)
if err != nil {
return nil, 0, err
}
return list, total, nil
}
func (s *SysApiService) GetAllSysApis(ctx context.Context) ([]db.SysApi, error) {
return s.queries.GetAllSysApis(ctx)
}
func (s *SysApiService) GetApiGroupNames(ctx context.Context) ([]string, error) {
return s.queries.GetSysApiGroupNames(ctx)
}
func (s *SysApiService) Create(ctx context.Context, req request.CreateSysApiRequest) error {
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)
api := db.CreateSysApiParams{
Name: req.Name,
GroupName: req.GroupName,
Method: req.Method,
Path: req.Path,
Sort: req.Sort,
}
// 创建权限
permissionId, err := q.CreateSysPermission(ctx, db.CreateSysPermissionParams{
Type: int16(enum.PermissionTypeApi),
})
if err != nil {
return err
}
// 创建api
apiId, err := q.CreateSysApi(ctx, api)
if err != nil {
return dberr.MapUniqueViolation(err, dberr.SysApisMethodPathKey, errs.ErrSysApiMethodPathAlreadyExists)
}
// 关联权限
if err = q.CreateSysApiPermission(ctx, db.CreateSysApiPermissionParams{
ApiID: apiId,
PermissionID: permissionId,
}); err != nil {
return err
}
if err = tx.Commit(ctx); err != nil {
return err
}
// 清理缓存
s.cache.ClearAllSysUserCache()
return nil
}
func (s *SysApiService) Update(ctx context.Context, id int32, req request.UpdateSysApiRequest) error {
api := db.UpdateSysApiParams{
ID: id,
Name: req.Name,
GroupName: req.GroupName,
Method: req.Method,
Path: req.Path,
Sort: req.Sort,
}
// 清理缓存
s.cache.ClearAllSysUserCache()
rows, err := s.queries.UpdateSysApi(ctx, api)
if err = dberr.MapRowsAffected(rows, err, errs.ErrSysApiNotFound); err != nil {
return dberr.MapUniqueViolation(err, dberr.SysApisMethodPathKey, errs.ErrSysApiMethodPathAlreadyExists)
}
return nil
}
func (s *SysApiService) Delete(ctx context.Context, id int32) error {
tx, err := s.pool.BeginTx(ctx, pgx.TxOptions{})
if err != nil {
return err
}
defer tx.Rollback(ctx)
q := db.New(tx)
rows, err := q.DeleteSysApi(ctx, id)
if err = dberr.MapRowsAffected(rows, err, errs.ErrSysApiNotFound); err != nil {
return err
}
if err = q.DeleteSysPermissionBySysApiID(ctx, id); err != nil {
return err
}
if err = q.DeleteSysApiPermission(ctx, id); err != nil {
return err
}
// 清理缓存
s.cache.ClearAllSysUserCache()
return tx.Commit(ctx)
}

View File

@@ -0,0 +1,118 @@
package service
import (
"context"
"mime/multipart"
"os"
"path/filepath"
db "server/internal/db/sqlc"
"server/internal/model/common"
"server/internal/model/response"
"server/internal/pkg/httputil"
gonanoid "github.com/matoous/go-nanoid/v2"
)
type SysFileService struct {
queries *db.Queries
}
func NewSysFileService(queries *db.Queries) *SysFileService {
return &SysFileService{
queries: queries,
}
}
// MakeSavedDir 创建目录并返回
func MakeSavedDir(folder string) (string, error) {
rootDir, err := os.Getwd()
if err != nil {
return "", err
}
uploadDir := filepath.Join(rootDir, "uploads", folder)
if err = os.MkdirAll(uploadDir, 0755); err != nil {
return "", err
}
return uploadDir, nil
}
func (s *SysFileService) ListPage(ctx context.Context, p *common.Pagination) ([]db.File, int64, error) {
params := db.GetFilesParams{
Limit: p.PageSize,
Offset: (p.Page - 1) * p.PageSize,
}
total, err := s.queries.CountFiles(ctx)
if err != nil {
return nil, 0, err
}
list, err := s.queries.GetFiles(ctx, params)
if err != nil {
return nil, 0, err
}
return response.ToFiles(list), total, nil
}
func (s *SysFileService) Upload(ctx context.Context, folder string, file *multipart.FileHeader) (*db.CreateFileRow, error) {
// 生成文件名
fileID, err := gonanoid.New()
if err != nil {
return nil, err
}
savedDir, err := MakeSavedDir(folder)
if err != nil {
return nil, err
}
fileExt := filepath.Ext(file.Filename)
filename := fileID + fileExt
filePath := filepath.Join(folder, filename)
// 用于保存文件
savedPath := filepath.Join(savedDir, filename)
// 打开上传的文件
src, err := file.Open()
if err != nil {
return nil, err
}
defer src.Close()
// 创建目标文件
dst, err := os.Create(savedPath)
if err != nil {
return nil, err
}
defer dst.Close()
// 复制文件内容
if _, err = dst.ReadFrom(src); err != nil {
return nil, err
}
params := db.CreateFileParams{
FileName: filename,
FilePath: filePath,
OriginalName: file.Filename,
FolderName: folder,
MimeType: file.Header.Get("Content-Type"),
FileSize: file.Size,
}
result, err := s.queries.CreateFile(ctx, params)
if err != nil {
return nil, err
}
result.FilePath = httputil.BuildFileUrl(&result.FilePath)
return &result, nil
}

View File

@@ -0,0 +1,178 @@
package service
import (
"context"
db "server/internal/db/sqlc"
"server/internal/model/common"
"server/internal/model/enum"
"server/internal/model/request"
"server/internal/pkg/dberr"
"server/internal/pkg/errs"
"github.com/jackc/pgx/v5"
"github.com/jackc/pgx/v5/pgxpool"
)
type SysMenuService struct {
pool *pgxpool.Pool
queries *db.Queries
}
func NewSysMenuService(queries *db.Queries, pool *pgxpool.Pool) *SysMenuService {
return &SysMenuService{
queries: queries,
pool: pool,
}
}
func (s *SysMenuService) Create(ctx context.Context, req request.CreateSysMenuRequest) error {
// 开启事务
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)
menu := db.CreateSysMenuParams{
Name: req.Name,
Path: req.Path,
Component: req.Component,
Type: *req.Type,
Status: *req.Status,
Hidden: req.Hidden,
Sort: req.Sort,
Icon: req.Icon,
}
permissionId, err := q.CreateSysPermission(ctx, db.CreateSysPermissionParams{
Type: int16(enum.PermissionTypeMenu),
Code: &req.PermissionCode,
})
if err != nil {
return dberr.MapUniqueViolation(err, dberr.SysPermissionsCodeKey, errs.ErrPermissionCodeAlreadyExists)
}
menuId, err := q.CreateSysMenu(ctx, menu)
if err != nil {
return err
}
if err = q.CreateSysMenuPermission(ctx, db.CreateSysMenuPermissionParams{
MenuID: menuId,
PermissionID: permissionId,
}); err != nil {
return err
}
if err = tx.Commit(ctx); err != nil {
return err
}
return nil
}
func (s *SysMenuService) Update(ctx context.Context, id int32, req request.UpdateSysMenuRequest) error {
tx, err := s.pool.BeginTx(ctx, pgx.TxOptions{})
if err != nil {
return err
}
defer tx.Rollback(ctx)
q := db.New(tx)
// 构造 menu 参数
menu := db.UpdateSysMenuParams{
ID: id,
Name: req.Name,
Path: req.Path,
Component: req.Component,
Hidden: req.Hidden,
Sort: req.Sort,
Type: req.Type,
Status: req.Status,
Icon: req.Icon,
}
if req.ParentID.Set {
menu.UpdateParentID = true
if req.ParentID.Valid {
menu.ParentID = &req.ParentID.Value
}
}
// 执行更新
rows, err := q.UpdateSysMenu(ctx, menu)
if err = dberr.MapRowsAffected(rows, err, errs.ErrSysMenuNotFound); err != nil {
return err
}
permission := db.UpdateSysMenuPermissionCodeParams{
MenuID: id,
Code: req.PermissionCode,
}
if err = q.UpdateSysMenuPermissionCode(ctx, permission); err != nil {
return dberr.MapUniqueViolation(err, dberr.SysPermissionsCodeKey, errs.ErrPermissionCodeAlreadyExists)
}
return tx.Commit(ctx)
}
func (s *SysMenuService) ListPage(ctx context.Context, p *common.Pagination) ([]db.ListSysMenusRow, int64, error) {
params := db.ListSysMenusParams{
Limit: p.PageSize,
Offset: (p.Page - 1) * p.PageSize,
}
total, err := s.queries.CountSysMenus(ctx)
if err != nil {
return nil, 0, err
}
list, err := s.queries.ListSysMenus(ctx, params)
if err != nil {
return nil, 0, err
}
return list, total, nil
}
func (s *SysMenuService) GetMenus(ctx context.Context) ([]db.GetAllSysMenusRow, error) {
return s.queries.GetAllSysMenus(ctx)
}
func (s *SysMenuService) Delete(ctx context.Context, id int32) error {
tx, err := s.pool.BeginTx(ctx, pgx.TxOptions{})
if err != nil {
return err
}
defer tx.Rollback(ctx)
q := db.New(tx)
rows, err := q.DeleteSysMenu(ctx, id)
if err = dberr.MapRowsAffected(rows, err, errs.ErrSysMenuNotFound); err != nil {
return err
}
if err = q.DeleteSysPermissionByMenuID(ctx, id); err != nil {
return err
}
if err = q.DeleteSysMenuPermission(ctx, id); err != nil {
return err
}
return tx.Commit(ctx)
}

View File

@@ -0,0 +1,163 @@
package service
import (
"context"
db "server/internal/db/sqlc"
"server/internal/model/common"
"server/internal/model/request"
"server/internal/pkg/dberr"
"server/internal/pkg/errs"
"server/internal/pkg/httputil"
"github.com/jackc/pgx/v5"
"github.com/jackc/pgx/v5/pgxpool"
)
type SysPostService struct {
queries *db.Queries
pool *pgxpool.Pool
}
func NewSysPostService(queries *db.Queries, pool *pgxpool.Pool) *SysPostService {
return &SysPostService{
queries: queries,
pool: pool,
}
}
func (s *SysPostService) ListPage(ctx context.Context, p *common.Pagination) ([]db.ListPostsRow, int64, error) {
params := db.ListPostsParams{
Limit: p.PageSize,
Offset: (p.Page - 1) * p.PageSize,
}
total, err := s.queries.CountPosts(ctx)
if err != nil {
return nil, 0, err
}
list, err := s.queries.ListPosts(ctx, params)
if err != nil {
return nil, 0, err
}
for i := range list {
url := httputil.BuildFileUrl(list[i].Cover)
list[i].Cover = &url
}
return list, total, nil
}
func (s *SysPostService) FindByID(ctx context.Context, id int32) (*db.GetPostByIdRow, error) {
post, err := s.queries.GetPostById(ctx, id)
if err != nil {
return nil, dberr.MapNoRows(err, errs.ErrPostNotFound)
}
url := httputil.BuildFileUrl(post.Cover)
post.Cover = &url
return &post, nil
}
func (s *SysPostService) Create(ctx context.Context, req request.CreatePostRequest) (int32, error) {
// 开启事务
tx, err := s.pool.BeginTx(ctx, pgx.TxOptions{})
if err != nil {
return 0, err
}
defer func() {
_ = tx.Rollback(ctx)
}()
q := db.New(tx)
params := db.CreatePostParams{
Title: req.Title,
CoverID: req.CoverID,
Slug: req.Slug,
Content: req.Content,
Summary: req.Summary,
Status: *req.Status,
Sort: req.Sort,
PublishedAt: req.PublishedAt,
}
postId, err := q.CreatePost(ctx, params)
if err != nil {
return 0, dberr.MapUniqueViolation(err, dberr.PostSlugKey, errs.ErrSlugAlreadyExists)
}
if err = q.CreatePostCategory(ctx, db.CreatePostCategoryParams{
PostID: postId,
CategoryID: *req.CategoryID,
}); err != nil {
return 0, err
}
// 提交
if err = tx.Commit(ctx); err != nil {
return 0, err
}
return postId, nil
}
func (s *SysPostService) Update(ctx context.Context, id int32, req request.UpdatePostRequest) error {
// 开启事务
tx, err := s.pool.BeginTx(ctx, pgx.TxOptions{})
if err != nil {
return err
}
defer func() {
_ = tx.Rollback(ctx)
}()
q := db.New(tx)
params := db.UpdatePostParams{
Title: req.Title,
CoverID: req.CoverID,
Slug: req.Slug,
Content: req.Content,
Summary: req.Summary,
Status: req.Status,
Sort: req.Sort,
PublishedAt: req.PublishedAt,
ID: id,
}
rows, err := q.UpdatePost(ctx, params)
if err != nil {
return dberr.MapUniqueViolation(err, dberr.PostSlugKey, errs.ErrSlugAlreadyExists)
}
if err = dberr.MapRowsAffected(rows, nil, errs.ErrPostNotFound); err != nil {
return err
}
if err = q.DeletePostCategory(ctx, id); err != nil {
return err
}
if err = q.CreatePostCategory(ctx, db.CreatePostCategoryParams{
PostID: id,
CategoryID: *req.CategoryID,
}); err != nil {
return err
}
// 提交
if err = tx.Commit(ctx); err != nil {
return err
}
return nil
}
func (s *SysPostService) Delete(ctx context.Context, id int32) error {
rows, err := s.queries.DeletePost(ctx, id)
return dberr.MapRowsAffected(rows, err, errs.ErrPostNotFound)
}

View File

@@ -0,0 +1,211 @@
package service
import (
"context"
db "server/internal/db/sqlc"
"server/internal/model/common"
"server/internal/model/enum"
"server/internal/model/request"
"server/internal/pkg/cache"
"server/internal/pkg/dberr"
"server/internal/pkg/errs"
"github.com/jackc/pgx/v5"
"github.com/jackc/pgx/v5/pgxpool"
)
type SysRoleService struct {
queries *db.Queries
pool *pgxpool.Pool
cache *cache.Caches
}
func NewSysRoleService(queries *db.Queries, pool *pgxpool.Pool, cache *cache.Caches) *SysRoleService {
return &SysRoleService{
queries: queries,
pool: pool,
cache: cache,
}
}
func (s *SysRoleService) ListPage(ctx context.Context, p *common.Pagination) ([]db.SysRole, int64, error) {
params := db.ListSysRolesParams{
Limit: p.PageSize,
Offset: (p.Page - 1) * p.PageSize,
}
total, err := s.queries.CountSysRoles(ctx)
if err != nil {
return nil, 0, err
}
list, err := s.queries.ListSysRoles(ctx, params)
if err != nil {
return nil, 0, err
}
return list, total, nil
}
func (s *SysRoleService) GetRoles(ctx context.Context) ([]db.SysRole, error) {
return s.queries.GetAllSysRoles(ctx)
}
func (s *SysRoleService) GetRoleMenus(ctx context.Context, id int32) ([]db.GetSysRoleMenusRow, error) {
_, err := s.queries.GetSysRoleByID(ctx, id)
if err != nil {
return nil, dberr.MapNoRows(err, errs.ErrSysRoleNotFound)
}
return s.queries.GetSysRoleMenus(ctx, id)
}
func (s *SysRoleService) GetRoleApis(ctx context.Context, id int32) ([]db.GetSysRoleApisRow, error) {
_, err := s.queries.GetSysRoleByID(ctx, id)
if err != nil {
return nil, dberr.MapNoRows(err, errs.ErrSysRoleNotFound)
}
return s.queries.GetSysRoleApis(ctx, id)
}
func (s *SysRoleService) Create(ctx context.Context, req request.CreateSysRoleRequest) error {
params := db.CreateSysRoleParams{
Name: req.Name,
Code: req.Code,
}
// 清理缓存
s.cache.ClearAllSysUserCache()
err := s.queries.CreateSysRole(ctx, params)
if err != nil {
return dberr.MapUniqueViolation(err, dberr.SysRoleCodeKey, errs.ErrCodeAlreadyExists)
}
return nil
}
func (s *SysRoleService) Update(ctx context.Context, id int32, req request.UpdateSysRoleRequest) error {
params := db.UpdateSysRoleParams{
ID: id,
Name: req.Name,
}
// 清理缓存
s.cache.ClearAllSysUserCache()
rows, err := s.queries.UpdateSysRole(ctx, params)
return dberr.MapRowsAffected(rows, err, errs.ErrSysRoleNotFound)
}
func (s *SysRoleService) SetRoleMenus(ctx context.Context, roleID int32, req request.SetSysRoleMenusRequest) error {
_, err := s.queries.GetSysRoleByID(ctx, roleID)
if err != nil {
return dberr.MapNoRows(err, errs.ErrSysRoleNotFound)
}
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)
permissionIds, err := s.queries.GetSysPermissionIdsByMenuIDs(ctx, req.MenuIDs)
if err != nil {
return err
}
params := make([]db.CreateSysRolePermissionParams, 0, len(permissionIds))
for _, id := range permissionIds {
params = append(params, db.CreateSysRolePermissionParams{
RoleID: roleID,
PermissionID: id,
})
}
if err = q.DeleteSysRolePermission(ctx, db.DeleteSysRolePermissionParams{
RoleID: roleID,
Type: int16(enum.PermissionTypeMenu),
}); err != nil {
return err
}
_, err = q.CreateSysRolePermission(ctx, params)
if err != nil {
return err
}
if err = tx.Commit(ctx); err != nil {
return err
}
// 清理缓存
s.cache.ClearAllSysUserCache()
return err
}
func (s *SysRoleService) SetRoleApis(ctx context.Context, roleID int32, req request.SetSysRoleApisRequest) error {
_, err := s.queries.GetSysRoleByID(ctx, roleID)
if err != nil {
return dberr.MapNoRows(err, errs.ErrSysRoleNotFound)
}
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)
permissionIds, err := s.queries.GetSysPermissionIdsByApiIDs(ctx, req.ApiIDs)
if err != nil {
return err
}
params := make([]db.CreateSysRolePermissionParams, 0, len(permissionIds))
for _, id := range permissionIds {
params = append(params, db.CreateSysRolePermissionParams{
RoleID: roleID,
PermissionID: id,
})
}
if err = q.DeleteSysRolePermission(ctx, db.DeleteSysRolePermissionParams{
RoleID: roleID,
Type: int16(enum.PermissionTypeApi),
}); err != nil {
return err
}
_, err = q.CreateSysRolePermission(ctx, params)
if err != nil {
return err
}
if err = tx.Commit(ctx); err != nil {
return err
}
// 清理缓存
s.cache.ClearAllSysUserCache()
return err
}
func (s *SysRoleService) Delete(ctx context.Context, id int32) error {
// 清理缓存
s.cache.ClearAllSysUserCache()
rows, err := s.queries.DeleteSysRole(ctx, id)
return dberr.MapRowsAffected(rows, err, errs.ErrSysRoleNotFound)
}

View File

@@ -0,0 +1,290 @@
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)
}