From 2f165d3a45ac03b17db8a98807c6315a5168e842 Mon Sep 17 00:00:00 2001 From: yvan <8574526@qq.com> Date: Sat, 15 Aug 2026 01:28:51 +0800 Subject: [PATCH] =?UTF-8?q?=E4=BC=98=E5=8C=96?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- cmd/kratos-admin/wire_gen.go | 45 ++++++++------- internal/biz/biz.go | 2 +- internal/biz/system.go | 34 ----------- internal/biz/system_init.go | 22 +++++-- internal/biz/user.go | 34 ++++++----- internal/data/data.go | 2 +- internal/data/system.go | 4 +- internal/server/gin.go | 4 +- internal/server/handler/export.go | 14 ++--- internal/server/handler/navigation.go | 4 +- internal/server/handler/public.go | 35 ++++++------ internal/server/handler/system_config.go | 4 +- internal/server/handler/user.go | 11 +++- internal/server/middleware/rate_limit.go | 4 +- internal/server/server.go | 2 +- internal/service/authentication.go | 37 +++++++----- internal/service/service.go | 2 +- internal/service/settings.go | 20 ++++++- internal/service/system.go | 13 ++--- internal/service/system_config.go | 14 +++-- internal/service/system_init.go | 4 +- internal/service/task.go | 18 ++++-- internal/service/user.go | 49 ++++++++-------- .../task_executor.go} | 57 ++++++++++--------- .../task_executor_test.go} | 2 +- internal/worker/task_scheduler.go | 21 ++++--- 26 files changed, 242 insertions(+), 216 deletions(-) delete mode 100644 internal/biz/system.go rename internal/{service/task_execution.go => worker/task_executor.go} (71%) rename internal/{service/task_execution_test.go => worker/task_executor_test.go} (95%) diff --git a/cmd/kratos-admin/wire_gen.go b/cmd/kratos-admin/wire_gen.go index d4a5aad..8eca50e 100644 --- a/cmd/kratos-admin/wire_gen.go +++ b/cmd/kratos-admin/wire_gen.go @@ -30,18 +30,6 @@ func wireApp(confServer *conf.Server, runtime *conf.Runtime, logger *slog.Logger if err != nil { return nil, nil, err } - systemRepo := data.NewSystemRepo(dataData) - cache := data.NewCache(dataData) - fileStorage, err := data.NewFileStorage(dataData) - if err != nil { - cleanup() - return nil, nil, err - } - systemUsecase := biz.NewSystemUsecase(systemRepo, cache, fileStorage) - settingsRepo := data.NewSettingsRepo(dataData) - settingsUsecase := biz.NewSettingsUsecase(settingsRepo) - settingsService := service.NewSettingsService(settingsUsecase, runtime, cache) - systemService := service.NewSystemService(systemUsecase, runtime, settingsService) accessRepo := data.NewAccessRepo(dataData) accessUsecase := biz.NewAccessUsecase(accessRepo) accessService := service.NewAccessService(accessUsecase, runtime) @@ -66,10 +54,16 @@ func wireApp(confServer *conf.Server, runtime *conf.Runtime, logger *slog.Logger email := handler.NewEmail(emailService) taskRepo := data.NewTaskRepo(dataData) taskUsecase := biz.NewTaskUsecase(taskRepo) + taskService := service.NewTaskService(taskUsecase) mediaRepo := data.NewMediaRepo(dataData) + fileStorage, err := data.NewFileStorage(dataData) + if err != nil { + cleanup() + return nil, nil, err + } mediaUsecase := biz.NewMediaUsecase(mediaRepo, fileStorage) - taskService := service.NewTaskService(taskUsecase, mediaUsecase, runtime) - taskScheduler := worker.NewTaskScheduler(taskService, logger) + taskExecutor := worker.NewTaskExecutor(taskUsecase, mediaUsecase, runtime) + taskScheduler := worker.NewTaskScheduler(taskUsecase, taskExecutor, logger) task := handler.NewTask(taskService, taskScheduler) mediaService := service.NewMediaService(mediaUsecase, runtime) media := handler.NewMedia(mediaService) @@ -77,10 +71,14 @@ func wireApp(confServer *conf.Server, runtime *conf.Runtime, logger *slog.Logger auditUsecase := biz.NewAuditUsecase(auditRepo) auditService := service.NewAuditService(auditUsecase) audit := handler.NewAudit(auditService) + settingsRepo := data.NewSettingsRepo(dataData) + settingsUsecase := biz.NewSettingsUsecase(settingsRepo) + cache := data.NewCache(dataData) + settingsService := service.NewSettingsService(settingsUsecase, runtime, cache) exportRepo := data.NewExportRepo(dataData) exportUsecase := biz.NewExportUsecase(exportRepo) exportService := service.NewExportService(exportUsecase) - export := handler.NewExport(systemService, exportService) + export := handler.NewExport(settingsService, exportService) versionRepo := data.NewVersionRepo(dataData) versionUsecase := biz.NewVersionUsecase(versionRepo) versionService := service.NewVersionService(versionUsecase) @@ -88,12 +86,19 @@ func wireApp(confServer *conf.Server, runtime *conf.Runtime, logger *slog.Logger dictionary := handler.NewDictionary(settingsService) parameter := handler.NewParameter(settingsService) apiToken := handler.NewAPIToken(settingsService) - systemConfig := handler.NewSystemConfig(systemService, settingsService, taskScheduler) - public := handler.NewPublic(runtime, systemService, settingsService, auditService, taskScheduler) - user := handler.NewUser(systemService) - navigation := handler.NewNavigation(systemService) + initializationRepo := data.NewInitializationRepo(dataData) + systemConfigUsecase := biz.NewSystemConfigUsecase(initializationRepo) + systemConfigService := service.NewSystemConfigService(systemConfigUsecase, runtime) + systemConfig := handler.NewSystemConfig(systemConfigService, settingsService, taskScheduler) + userRepo := data.NewUserRepo(dataData) + userUsecase := biz.NewUserUsecase(userRepo) + authService := service.NewAuthService(userUsecase, runtime, settingsService) + public := handler.NewPublic(runtime, authService, systemConfigService, settingsService, auditService, taskScheduler) + userService := service.NewUserService(userUsecase, settingsService) + user := handler.NewUser(userService, authService) + navigation := handler.NewNavigation(userService) session := handler.NewSession(settingsService) - engine := server.NewGinEngine(runtime, systemService, accessService, authority, menu, api, permission, organization, announcement, email, task, media, audit, export, version, dictionary, parameter, apiToken, systemConfig, public, user, navigation, session, settingsService, auditService, logger) + engine := server.NewGinEngine(runtime, accessService, authority, menu, api, permission, organization, announcement, email, task, media, audit, export, version, dictionary, parameter, apiToken, systemConfig, public, user, navigation, session, settingsService, auditService, logger) httpServer := server.NewGinServer(confServer, engine) app := newApp(logger, httpServer, taskScheduler) return app, func() { diff --git a/internal/biz/biz.go b/internal/biz/biz.go index 4981f41..6499bfb 100644 --- a/internal/biz/biz.go +++ b/internal/biz/biz.go @@ -3,4 +3,4 @@ package biz import "github.com/google/wire" // ProviderSet is biz providers. -var ProviderSet = wire.NewSet(NewSystemUsecase, NewAccessUsecase, NewMenuUsecase, NewOrganizationUsecase, NewSettingsUsecase, NewVersionUsecase, NewExportUsecase, NewAuditUsecase, NewTaskUsecase, NewMediaUsecase, NewAnnouncementUsecase, NewEmailUsecase) +var ProviderSet = wire.NewSet(NewUserUsecase, NewSystemConfigUsecase, NewAccessUsecase, NewMenuUsecase, NewOrganizationUsecase, NewSettingsUsecase, NewVersionUsecase, NewExportUsecase, NewAuditUsecase, NewTaskUsecase, NewMediaUsecase, NewAnnouncementUsecase, NewEmailUsecase) diff --git a/internal/biz/system.go b/internal/biz/system.go deleted file mode 100644 index 589ff7c..0000000 --- a/internal/biz/system.go +++ /dev/null @@ -1,34 +0,0 @@ -package biz - -import ( - "context" - "time" -) - -type SystemRepo interface { - InitializationRepo - UserRepo -} - -type SystemUsecase struct { - repo SystemRepo - cache Cache - files FileStorage -} - -func NewSystemUsecase(repo SystemRepo, cache Cache, files FileStorage) *SystemUsecase { - return &SystemUsecase{repo: repo, cache: cache, files: files} -} - -func (uc *SystemUsecase) CacheGet(ctx context.Context, key string) (string, bool, error) { - return uc.cache.Get(ctx, key) -} -func (uc *SystemUsecase) CacheSet(ctx context.Context, key, value string, expiration time.Duration) error { - return uc.cache.Set(ctx, key, value, expiration) -} -func (uc *SystemUsecase) CacheDelete(ctx context.Context, key string) error { - return uc.cache.Delete(ctx, key) -} -func (uc *SystemUsecase) CacheIncrement(ctx context.Context, key string, expiration time.Duration) (int64, error) { - return uc.cache.Increment(ctx, key, expiration) -} diff --git a/internal/biz/system_init.go b/internal/biz/system_init.go index ac2a8a4..c69a310 100644 --- a/internal/biz/system_init.go +++ b/internal/biz/system_init.go @@ -16,20 +16,30 @@ type InitializationRepo interface { ReloadConfig(context.Context) error } -func (uc *SystemUsecase) IsInitialized(ctx context.Context) (bool, error) { +type SystemConfigUsecase struct{ repo InitializationRepo } + +func NewSystemConfigUsecase(repo InitializationRepo) *SystemConfigUsecase { + return &SystemConfigUsecase{repo: repo} +} + +func (uc *SystemConfigUsecase) IsInitialized(ctx context.Context) (bool, error) { return uc.repo.IsInitialized(ctx) } -func (uc *SystemUsecase) Initialize(ctx context.Context, config *DatabaseConfig) error { +func (uc *SystemConfigUsecase) Initialize(ctx context.Context, config *DatabaseConfig) error { if config == nil || len(config.AdminPassword) < 6 { return ErrInvalidCredentials } return uc.repo.Initialize(ctx, config) } -func (uc *SystemUsecase) PersistConfig(ctx context.Context) error { return uc.repo.PersistConfig(ctx) } -func (uc *SystemUsecase) PersistAdminConfig(ctx context.Context, value []byte) error { +func (uc *SystemConfigUsecase) PersistConfig(ctx context.Context) error { + return uc.repo.PersistConfig(ctx) +} +func (uc *SystemConfigUsecase) PersistAdminConfig(ctx context.Context, value []byte) error { return uc.repo.PersistAdminConfig(ctx, value) } -func (uc *SystemUsecase) PersistRuntimeConfig(ctx context.Context, data, admin []byte) error { +func (uc *SystemConfigUsecase) PersistRuntimeConfig(ctx context.Context, data, admin []byte) error { return uc.repo.PersistRuntimeConfig(ctx, data, admin) } -func (uc *SystemUsecase) ReloadConfig(ctx context.Context) error { return uc.repo.ReloadConfig(ctx) } +func (uc *SystemConfigUsecase) ReloadConfig(ctx context.Context) error { + return uc.repo.ReloadConfig(ctx) +} diff --git a/internal/biz/user.go b/internal/biz/user.go index 85f0849..cbe3ef7 100644 --- a/internal/biz/user.go +++ b/internal/biz/user.go @@ -63,7 +63,11 @@ type UserRepo interface { FillDepartmentNamePaths(context.Context, *User) error } -func (uc *SystemUsecase) Login(ctx context.Context, username, password string) (*User, error) { +type UserUsecase struct{ repo UserRepo } + +func NewUserUsecase(repo UserRepo) *UserUsecase { return &UserUsecase{repo: repo} } + +func (uc *UserUsecase) Login(ctx context.Context, username, password string) (*User, error) { u, err := uc.repo.FindUserByUsername(ctx, username) if err != nil || bcrypt.CompareHashAndPassword([]byte(u.Password), []byte(password)) != nil { return nil, ErrInvalidCredentials @@ -72,7 +76,7 @@ func (uc *SystemUsecase) Login(ctx context.Context, username, password string) ( return u, nil } -func (uc *SystemUsecase) User(ctx context.Context, id uint) (*User, error) { +func (uc *UserUsecase) User(ctx context.Context, id uint) (*User, error) { user, err := uc.repo.FindUserByID(ctx, id) if err != nil { return nil, err @@ -82,7 +86,7 @@ func (uc *SystemUsecase) User(ctx context.Context, id uint) (*User, error) { return user, nil } -func (uc *SystemUsecase) fallbackDefaultRouter(ctx context.Context, user *User) { +func (uc *UserUsecase) fallbackDefaultRouter(ctx context.Context, user *User) { if user == nil || user.Authority.DefaultRouter == "" { return } @@ -101,14 +105,14 @@ func menuNameExists(menus []*Menu, name string) bool { return false } -func (uc *SystemUsecase) Menus(ctx context.Context, authorityID uint) ([]*Menu, error) { +func (uc *UserUsecase) Menus(ctx context.Context, authorityID uint) ([]*Menu, error) { return uc.repo.MenusByAuthority(ctx, authorityID) } -func (uc *SystemUsecase) ListUsers(ctx context.Context, page, pageSize int, filter *UserListFilter) ([]*User, int64, error) { +func (uc *UserUsecase) ListUsers(ctx context.Context, page, pageSize int, filter *UserListFilter) ([]*User, int64, error) { return uc.repo.ListUsers(ctx, page, pageSize, filter) } -func (uc *SystemUsecase) CreateUser(ctx context.Context, user *User, authorityIDs []uint) (*User, error) { +func (uc *UserUsecase) CreateUser(ctx context.Context, user *User, authorityIDs []uint) (*User, error) { if user.AuthorityID == 0 && len(authorityIDs) > 0 { user.AuthorityID = authorityIDs[0] } @@ -120,12 +124,12 @@ func (uc *SystemUsecase) CreateUser(ctx context.Context, user *User, authorityID user.Password = string(hash) return uc.repo.CreateUserWithAuthorities(ctx, user, authorityIDs) } -func (uc *SystemUsecase) UpdateUser(ctx context.Context, user *User, authorityIDs []uint) error { +func (uc *UserUsecase) UpdateUser(ctx context.Context, user *User, authorityIDs []uint) error { authorityIDs = includeAuthority(authorityIDs, user.AuthorityID) return uc.repo.UpdateUserWithAuthorities(ctx, user, authorityIDs) } -func (uc *SystemUsecase) UpdateSelfUser(ctx context.Context, user *User) error { +func (uc *UserUsecase) UpdateSelfUser(ctx context.Context, user *User) error { return uc.repo.UpdateSelfUser(ctx, user) } @@ -140,17 +144,17 @@ func includeAuthority(ids []uint, primary uint) []uint { } return append(ids, primary) } -func (uc *SystemUsecase) DeleteUser(ctx context.Context, id uint) error { +func (uc *UserUsecase) DeleteUser(ctx context.Context, id uint) error { return uc.repo.DeleteUser(ctx, id) } -func (uc *SystemUsecase) ResetPassword(ctx context.Context, id uint, password string) error { +func (uc *UserUsecase) ResetPassword(ctx context.Context, id uint, password string) error { hash, err := bcrypt.GenerateFromPassword([]byte(password), bcrypt.DefaultCost) if err != nil { return err } return uc.repo.UpdatePassword(ctx, id, string(hash), false) } -func (uc *SystemUsecase) ChangePassword(ctx context.Context, id uint, oldPassword, newPassword string) error { +func (uc *UserUsecase) ChangePassword(ctx context.Context, id uint, oldPassword, newPassword string) error { user, err := uc.repo.FindUserByID(ctx, id) if err != nil { return err @@ -164,18 +168,18 @@ func (uc *SystemUsecase) ChangePassword(ctx context.Context, id uint, oldPasswor } return uc.repo.UpdatePassword(ctx, id, string(hash), true) } -func (uc *SystemUsecase) Authorities(ctx context.Context) ([]*Authority, error) { +func (uc *UserUsecase) Authorities(ctx context.Context) ([]*Authority, error) { return uc.repo.ListAuthorities(ctx) } -func (uc *SystemUsecase) SetUserAuthorities(ctx context.Context, id uint, authorityIDs []uint) error { +func (uc *UserUsecase) SetUserAuthorities(ctx context.Context, id uint, authorityIDs []uint) error { if len(authorityIDs) == 0 { return ErrInvalidCredentials } return uc.repo.SetUserAuthorities(ctx, id, authorityIDs) } -func (uc *SystemUsecase) SetUserAuthority(ctx context.Context, id, authorityID uint) error { +func (uc *UserUsecase) SetUserAuthority(ctx context.Context, id, authorityID uint) error { return uc.repo.SetUserAuthority(ctx, id, authorityID) } -func (uc *SystemUsecase) SetUserSetting(ctx context.Context, id uint, setting map[string]any) error { +func (uc *UserUsecase) SetUserSetting(ctx context.Context, id uint, setting map[string]any) error { return uc.repo.SetUserSetting(ctx, id, setting) } diff --git a/internal/data/data.go b/internal/data/data.go index 67dee5d..9f7f4d7 100644 --- a/internal/data/data.go +++ b/internal/data/data.go @@ -13,7 +13,7 @@ import ( "kra/internal/conf" ) -var ProviderSet = wire.NewSet(NewData, NewSystemRepo, NewAccessRepo, NewMenuRepo, NewOrganizationRepo, NewSettingsRepo, NewVersionRepo, NewExportRepo, NewAuditRepo, NewTaskRepo, NewMediaRepo, NewAnnouncementRepo, NewEmailRepo, NewCache, NewFileStorage) +var ProviderSet = wire.NewSet(NewData, NewUserRepo, NewInitializationRepo, NewAccessRepo, NewMenuRepo, NewOrganizationRepo, NewSettingsRepo, NewVersionRepo, NewExportRepo, NewAuditRepo, NewTaskRepo, NewMediaRepo, NewAnnouncementRepo, NewEmailRepo, NewCache, NewFileStorage) type Data struct { initMu sync.Mutex diff --git a/internal/data/system.go b/internal/data/system.go index 0968229..c6fd56c 100644 --- a/internal/data/system.go +++ b/internal/data/system.go @@ -92,4 +92,6 @@ func (menuParameterPO) TableName() string { return "sys_base_menu_parameters" } type systemRepo struct{ data *Data } -func NewSystemRepo(data *Data) biz.SystemRepo { return &systemRepo{data: data} } +func NewUserRepo(data *Data) biz.UserRepo { return &systemRepo{data: data} } + +func NewInitializationRepo(data *Data) biz.InitializationRepo { return &systemRepo{data: data} } diff --git a/internal/server/gin.go b/internal/server/gin.go index e202d43..9a924a4 100644 --- a/internal/server/gin.go +++ b/internal/server/gin.go @@ -20,10 +20,10 @@ import ( kratoshttp "github.com/go-kratos/kratos/v3/transport/http" ) -func NewGinEngine(runtime *conf.Runtime, system *service.SystemService, access *service.AccessService, authority *handler.Authority, menu *handler.Menu, api *handler.API, permission *handler.Permission, organization *handler.Organization, announcement *handler.Announcement, email *handler.Email, task *handler.Task, media *handler.Media, auditHandler *handler.Audit, export *handler.Export, version *handler.Version, dictionary *handler.Dictionary, parameter *handler.Parameter, apiToken *handler.APIToken, systemConfig *handler.SystemConfig, publicHandler *handler.Public, user *handler.User, navigation *handler.Navigation, session *handler.Session, settings *service.SettingsService, audit *service.AuditService, logger *slog.Logger) *gin.Engine { +func NewGinEngine(runtime *conf.Runtime, access *service.AccessService, authority *handler.Authority, menu *handler.Menu, api *handler.API, permission *handler.Permission, organization *handler.Organization, announcement *handler.Announcement, email *handler.Email, task *handler.Task, media *handler.Media, auditHandler *handler.Audit, export *handler.Export, version *handler.Version, dictionary *handler.Dictionary, parameter *handler.Parameter, apiToken *handler.APIToken, systemConfig *handler.SystemConfig, publicHandler *handler.Public, user *handler.User, navigation *handler.Navigation, session *handler.Session, settings *service.SettingsService, audit *service.AuditService, logger *slog.Logger) *gin.Engine { gin.SetMode(gin.ReleaseMode) engine := gin.New() - engine.Use(servermiddleware.RequestMeta(), servermiddleware.Recovery(audit, logger), servermiddleware.AccessLog(runtime, logger), servermiddleware.ErrorAudit(audit), servermiddleware.SecurityRateLimit(system, settings), servermiddleware.OperationAudit(runtime, audit)) + engine.Use(servermiddleware.RequestMeta(), servermiddleware.Recovery(audit, logger), servermiddleware.AccessLog(runtime, logger), servermiddleware.ErrorAudit(audit), servermiddleware.SecurityRateLimit(settings), servermiddleware.OperationAudit(runtime, audit)) prefix := "" config := runtime.Admin() diff --git a/internal/server/handler/export.go b/internal/server/handler/export.go index f598abc..5c3f6ee 100644 --- a/internal/server/handler/export.go +++ b/internal/server/handler/export.go @@ -16,12 +16,12 @@ import ( ) type Export struct { - system *service.SystemService - service *service.ExportService + settings *service.SettingsService + service *service.ExportService } -func NewExport(system *service.SystemService, service *service.ExportService) *Export { - return &Export{system: system, service: service} +func NewExport(settings *service.SettingsService, export *service.ExportService) *Export { + return &Export{settings: settings, service: export} } type exportToken struct { @@ -153,7 +153,7 @@ func (h *Export) Issue(blank bool) gin.HandlerFunc { } token := strings.ReplaceAll(uuid.NewString(), "-", "") raw, _ := json.Marshal(exportToken{TemplateID: templateID, Params: exportParams(c.Request.URL.Query()), Blank: blank}) - if err := h.system.CacheSet(c.Request.Context(), "export:"+token, string(raw), 30*time.Minute); err != nil { + if err := h.settings.CacheSet(c.Request.Context(), "export:"+token, string(raw), 30*time.Minute); err != nil { httpx.Fail(c, "导出令牌创建失败") return } @@ -196,7 +196,7 @@ func (h *Export) Download(expectBlank bool) gin.HandlerFunc { httpx.Fail(c, "导出token不能为空") return } - raw, ok, err := h.system.CacheGet(c.Request.Context(), "export:"+token) + raw, ok, err := h.settings.CacheGet(c.Request.Context(), "export:"+token) if err != nil || !ok { httpx.Fail(c, "导出token无效或已过期") return @@ -210,7 +210,7 @@ func (h *Export) Download(expectBlank bool) gin.HandlerFunc { httpx.Fail(c, "token类型错误") return } - _ = h.system.CacheDelete(c.Request.Context(), "export:"+token) + _ = h.settings.CacheDelete(c.Request.Context(), "export:"+token) var data []byte var name string if expectBlank { diff --git a/internal/server/handler/navigation.go b/internal/server/handler/navigation.go index 0785837..3668a4a 100644 --- a/internal/server/handler/navigation.go +++ b/internal/server/handler/navigation.go @@ -8,9 +8,9 @@ import ( "github.com/gin-gonic/gin" ) -type Navigation struct{ service *service.SystemService } +type Navigation struct{ service *service.UserService } -func NewNavigation(service *service.SystemService) *Navigation { return &Navigation{service: service} } +func NewNavigation(service *service.UserService) *Navigation { return &Navigation{service: service} } func (h *Navigation) Menu(c *gin.Context) { claims := middleware.Claims(c) if claims == nil { diff --git a/internal/server/handler/public.go b/internal/server/handler/public.go index 5205d85..f87820d 100644 --- a/internal/server/handler/public.go +++ b/internal/server/handler/public.go @@ -19,7 +19,8 @@ import ( ) type Public struct { - system *service.SystemService + auth *service.AuthService + system *service.SystemConfigService settings *service.SettingsService audit *service.AuditService scheduler *worker.TaskScheduler @@ -27,8 +28,8 @@ type Public struct { store *captchaStore } -func NewPublic(runtime *conf.Runtime, system *service.SystemService, settings *service.SettingsService, audit *service.AuditService, scheduler *worker.TaskScheduler) *Public { - return &Public{system: system, settings: settings, audit: audit, scheduler: scheduler, runtime: runtime, store: &captchaStore{service: system, runtime: runtime}} +func NewPublic(runtime *conf.Runtime, auth *service.AuthService, system *service.SystemConfigService, settings *service.SettingsService, audit *service.AuditService, scheduler *worker.TaskScheduler) *Public { + return &Public{auth: auth, system: system, settings: settings, audit: audit, scheduler: scheduler, runtime: runtime, store: &captchaStore{service: settings, runtime: runtime}} } func (h *Public) captchaConfig() (int, int, int) { @@ -55,7 +56,7 @@ func (h *Public) Captcha(c *gin.Context) { if security != nil { keyLong, width, height = security.KeyLong, security.ImgWidth, security.ImgHeight ttl := time.Duration(security.CaptchaTimeout) * time.Second - failures, _ := ensureLoginIPCounter(c.Request.Context(), h.system, c.ClientIP(), ttl) + failures, _ := ensureLoginIPCounter(c.Request.Context(), h.settings, c.ClientIP(), ttl) openCaptcha = security.CaptchaOpen == 0 || failures > security.CaptchaOpen } driver := base64Captcha.NewDriverDigit(height, width, keyLong, 0.7, 80) @@ -75,7 +76,7 @@ func (h *Public) Login(c *gin.Context) { } security, _ := h.settings.CurrentSecurity(c.Request.Context()) if security != nil && security.LockEnable { - if _, locked, _ := h.system.CacheGet(c.Request.Context(), "login:lock:"+req.Username); locked { + if _, locked, _ := h.settings.CacheGet(c.Request.Context(), "login:lock:"+req.Username); locked { httpx.Fail(c, "账号已锁定,请 "+strconv.Itoa(security.LockDuration)+" 分钟后再试") _ = h.audit.RecordLoginRequest(c.Request.Context(), &dto.LoginLogRequest{Username: req.Username, IP: c.ClientIP(), Status: false, ErrorMessage: "账号已锁定", Agent: c.Request.UserAgent()}) return @@ -87,18 +88,18 @@ func (h *Public) Login(c *gin.Context) { if security.CaptchaTimeout > 0 { ipTTL = time.Duration(security.CaptchaTimeout) * time.Second } - failures, _ := ensureLoginIPCounter(c.Request.Context(), h.system, c.ClientIP(), ipTTL) + failures, _ := ensureLoginIPCounter(c.Request.Context(), h.settings, c.ClientIP(), ipTTL) requireCaptcha = security.CaptchaOpen == 0 || failures > security.CaptchaOpen } if requireCaptcha && (req.CaptchaID == "" || req.Captcha == "" || !h.store.Verify(req.CaptchaID, req.Captcha, true)) { - _, _ = h.system.CacheIncrement(c.Request.Context(), c.ClientIP(), ipTTL) + _, _ = h.settings.CacheIncrement(c.Request.Context(), c.ClientIP(), ipTTL) _ = h.audit.RecordLoginRequest(c.Request.Context(), &dto.LoginLogRequest{Username: req.Username, IP: c.ClientIP(), Status: false, ErrorMessage: "验证码错误", Agent: c.Request.UserAgent()}) httpx.Fail(c, "验证码错误") return } - result, err := h.system.Login(c.Request.Context(), req.Username, req.Password) + result, err := h.auth.Login(c.Request.Context(), req.Username, req.Password) if err != nil { - _, _ = h.system.CacheIncrement(c.Request.Context(), c.ClientIP(), ipTTL) + _, _ = h.settings.CacheIncrement(c.Request.Context(), c.ClientIP(), ipTTL) if errors.Is(err, biz.ErrUserDisabled) { var disabled *service.UserDisabledError errors.As(err, &disabled) @@ -113,9 +114,9 @@ func (h *Public) Login(c *gin.Context) { _ = h.audit.RecordLoginRequest(c.Request.Context(), &dto.LoginLogRequest{Username: req.Username, IP: c.ClientIP(), Status: false, ErrorMessage: "用户名不存在或者密码错误", Agent: c.Request.UserAgent()}) if security != nil && security.LockEnable { lockTTL := time.Duration(security.LockDuration) * time.Minute - failures, _ := h.system.CacheIncrement(c.Request.Context(), "login:fail:"+req.Username, lockTTL) + failures, _ := h.settings.CacheIncrement(c.Request.Context(), "login:fail:"+req.Username, lockTTL) if int(failures) >= security.LockThreshold { - _ = h.system.CacheSet(c.Request.Context(), "login:lock:"+req.Username, "1", lockTTL) + _ = h.settings.CacheSet(c.Request.Context(), "login:lock:"+req.Username, "1", lockTTL) } } httpx.Fail(c, "用户名不存在或者密码错误") @@ -123,15 +124,15 @@ func (h *Public) Login(c *gin.Context) { } userID, _ := result.User["ID"].(uint) _ = h.audit.RecordLoginRequest(c.Request.Context(), &dto.LoginLogRequest{Username: req.Username, IP: c.ClientIP(), Status: true, Agent: c.Request.UserAgent(), UserID: userID}) - _ = h.system.CacheDelete(c.Request.Context(), "login:fail:"+req.Username) - _ = h.system.CacheDelete(c.Request.Context(), "login:lock:"+req.Username) + _ = h.settings.CacheDelete(c.Request.Context(), "login:fail:"+req.Username) + _ = h.settings.CacheDelete(c.Request.Context(), "login:lock:"+req.Username) maxAge := int(time.Until(time.UnixMilli(result.ExpiresAt)).Seconds()) httpx.SetTokenCookie(c, result.Token, maxAge) httpx.Write(c, httpx.CodeSuccess, result, "登录成功") } -func ensureLoginIPCounter(ctx context.Context, system *service.SystemService, ip string, expiration time.Duration) (int, error) { - value, exists, err := system.CacheGet(ctx, ip) +func ensureLoginIPCounter(ctx context.Context, settings *service.SettingsService, ip string, expiration time.Duration) (int, error) { + value, exists, err := settings.CacheGet(ctx, ip) if err != nil { return 0, err } @@ -141,7 +142,7 @@ func ensureLoginIPCounter(ctx context.Context, system *service.SystemService, ip if expiration <= 0 { expiration = time.Hour } - if err = system.CacheSet(ctx, ip, "1", expiration); err != nil { + if err = settings.CacheSet(ctx, ip, "1", expiration); err != nil { return 0, err } return 1, nil @@ -194,7 +195,7 @@ func (h *Public) InitializeDatabase(engine *gin.Engine) gin.HandlerFunc { } type captchaStore struct { - service *service.SystemService + service *service.SettingsService runtime *conf.Runtime } diff --git a/internal/server/handler/system_config.go b/internal/server/handler/system_config.go index 4649485..98df491 100644 --- a/internal/server/handler/system_config.go +++ b/internal/server/handler/system_config.go @@ -16,12 +16,12 @@ import ( ) type SystemConfig struct { - system *service.SystemService + system *service.SystemConfigService settings *service.SettingsService scheduler *worker.TaskScheduler } -func NewSystemConfig(system *service.SystemService, settings *service.SettingsService, scheduler *worker.TaskScheduler) *SystemConfig { +func NewSystemConfig(system *service.SystemConfigService, settings *service.SettingsService, scheduler *worker.TaskScheduler) *SystemConfig { return &SystemConfig{system: system, settings: settings, scheduler: scheduler} } diff --git a/internal/server/handler/user.go b/internal/server/handler/user.go index 7480aec..8fcedeb 100644 --- a/internal/server/handler/user.go +++ b/internal/server/handler/user.go @@ -12,9 +12,14 @@ import ( "github.com/gin-gonic/gin" ) -type User struct{ service *service.SystemService } +type User struct { + service *service.UserService + auth *service.AuthService +} -func NewUser(service *service.SystemService) *User { return &User{service: service} } +func NewUser(user *service.UserService, auth *service.AuthService) *User { + return &User{service: user, auth: auth} +} func (h *User) List(c *gin.Context) { var req dto.UserListRequest @@ -155,7 +160,7 @@ func (h *User) SwitchAuthority(c *gin.Context) { httpx.Fail(c, "参数错误") return } - login, err := h.service.SwitchAuthority(c.Request.Context(), claims.ID, req.AuthorityID) + login, err := h.auth.SwitchAuthority(c.Request.Context(), claims.ID, req.AuthorityID) if err != nil { httpx.Fail(c, err.Error()) return diff --git a/internal/server/middleware/rate_limit.go b/internal/server/middleware/rate_limit.go index 72fcb31..cc3c1d2 100644 --- a/internal/server/middleware/rate_limit.go +++ b/internal/server/middleware/rate_limit.go @@ -10,7 +10,7 @@ import ( "github.com/gin-gonic/gin" ) -func SecurityRateLimit(system *service.SystemService, settings *service.SettingsService) gin.HandlerFunc { +func SecurityRateLimit(settings *service.SettingsService) gin.HandlerFunc { return func(c *gin.Context) { path := strings.TrimSuffix(c.Request.URL.Path, "/") if !strings.HasSuffix(path, "/base/login") && !strings.HasSuffix(path, "/base/captcha") { @@ -27,7 +27,7 @@ func SecurityRateLimit(system *service.SystemService, settings *service.Settings window = 60 } key := "KRA_SecLimit" + c.ClientIP() + c.FullPath() - count, cacheErr := system.CacheIncrement(c.Request.Context(), key, time.Duration(window)*time.Second) + count, cacheErr := settings.CacheIncrement(c.Request.Context(), key, time.Duration(window)*time.Second) if cacheErr == nil && int(count) > config.LimitCount { httpx.Fail(c, "请求太过频繁,请稍后再试") c.Abort() diff --git a/internal/server/server.go b/internal/server/server.go index aa38095..d4df517 100644 --- a/internal/server/server.go +++ b/internal/server/server.go @@ -8,4 +8,4 @@ import ( ) // ProviderSet is server providers. -var ProviderSet = wire.NewSet(NewGinEngine, NewGinServer, handler.NewAuthority, handler.NewMenu, handler.NewAPI, handler.NewPermission, handler.NewOrganization, handler.NewAnnouncement, handler.NewEmail, handler.NewTask, handler.NewMedia, handler.NewAudit, handler.NewExport, handler.NewVersion, handler.NewDictionary, handler.NewParameter, handler.NewAPIToken, handler.NewSystemConfig, handler.NewPublic, handler.NewUser, handler.NewNavigation, handler.NewSession, worker.NewTaskScheduler) +var ProviderSet = wire.NewSet(NewGinEngine, NewGinServer, handler.NewAuthority, handler.NewMenu, handler.NewAPI, handler.NewPermission, handler.NewOrganization, handler.NewAnnouncement, handler.NewEmail, handler.NewTask, handler.NewMedia, handler.NewAudit, handler.NewExport, handler.NewVersion, handler.NewDictionary, handler.NewParameter, handler.NewAPIToken, handler.NewSystemConfig, handler.NewPublic, handler.NewUser, handler.NewNavigation, handler.NewSession, worker.NewTaskExecutor, worker.NewTaskScheduler) diff --git a/internal/service/authentication.go b/internal/service/authentication.go index 718728f..6f75a38 100644 --- a/internal/service/authentication.go +++ b/internal/service/authentication.go @@ -5,6 +5,7 @@ import ( "time" "kra/internal/biz" + "kra/internal/conf" "kra/pkg/adminauth" ) @@ -20,20 +21,17 @@ type LoginResult struct { NeedChangePassword bool `json:"needChangePassword"` } -func (s *SystemService) CacheGet(ctx context.Context, key string) (string, bool, error) { - return s.uc.CacheGet(ctx, key) -} -func (s *SystemService) CacheSet(ctx context.Context, key, value string, expiration time.Duration) error { - return s.uc.CacheSet(ctx, key, value, expiration) -} -func (s *SystemService) CacheDelete(ctx context.Context, key string) error { - return s.uc.CacheDelete(ctx, key) -} -func (s *SystemService) CacheIncrement(ctx context.Context, key string, expiration time.Duration) (int64, error) { - return s.uc.CacheIncrement(ctx, key, expiration) +type AuthService struct { + uc *biz.UserUsecase + runtime *conf.Runtime + settings *SettingsService } -func (s *SystemService) Login(ctx context.Context, username, password string) (*LoginResult, error) { +func NewAuthService(uc *biz.UserUsecase, runtime *conf.Runtime, settings *SettingsService) *AuthService { + return &AuthService{uc: uc, runtime: runtime, settings: settings} +} + +func (s *AuthService) Login(ctx context.Context, username, password string) (*LoginResult, error) { u, err := s.uc.Login(ctx, username, password) if err != nil { return nil, err @@ -47,7 +45,7 @@ func (s *SystemService) Login(ctx context.Context, username, password string) (* return s.issueLogin(ctx, u, u.AuthorityID) } -func (s *SystemService) issueLogin(ctx context.Context, user *biz.User, authorityID uint) (*LoginResult, error) { +func (s *AuthService) issueLogin(ctx context.Context, user *biz.User, authorityID uint) (*LoginResult, error) { expires, buffer := 7*24*time.Hour, 24*time.Hour secret, issuer := "", "kra" config := s.runtime.Admin() @@ -65,7 +63,7 @@ func (s *SystemService) issueLogin(ctx context.Context, user *biz.User, authorit return nil, err } if s.settings.UseMultipoint() { - oldToken, _, cacheErr := s.settings.cache.Get(ctx, activeTokenKey(user.Username)) + oldToken, _, cacheErr := s.settings.CacheGet(ctx, activeTokenKey(user.Username)) if cacheErr != nil { return nil, cacheErr } @@ -75,3 +73,14 @@ func (s *SystemService) issueLogin(ctx context.Context, user *biz.User, authorit } return &LoginResult{User: convertUser(user), Token: token, ExpiresAt: claims.ExpiresAt.UnixMilli(), NeedChangePassword: user.MustChangePassword}, nil } + +func (s *AuthService) SwitchAuthority(ctx context.Context, id, authorityID uint) (*LoginResult, error) { + if err := s.uc.SetUserAuthority(ctx, id, authorityID); err != nil { + return nil, err + } + user, err := s.uc.User(ctx, id) + if err != nil { + return nil, err + } + return s.issueLogin(ctx, user, authorityID) +} diff --git a/internal/service/service.go b/internal/service/service.go index 45cb7f6..c73090e 100644 --- a/internal/service/service.go +++ b/internal/service/service.go @@ -3,4 +3,4 @@ package service import "github.com/google/wire" // ProviderSet is service providers. -var ProviderSet = wire.NewSet(NewSystemService, NewAccessService, NewMenuService, NewOrganizationService, NewSettingsService, NewVersionService, NewExportService, NewAuditService, NewTaskService, NewMediaService, NewAnnouncementService, NewEmailService) +var ProviderSet = wire.NewSet(NewAuthService, NewUserService, NewSystemConfigService, NewAccessService, NewMenuService, NewOrganizationService, NewSettingsService, NewVersionService, NewExportService, NewAuditService, NewTaskService, NewMediaService, NewAnnouncementService, NewEmailService) diff --git a/internal/service/settings.go b/internal/service/settings.go index 708c1e6..e863b23 100644 --- a/internal/service/settings.go +++ b/internal/service/settings.go @@ -21,6 +21,22 @@ func NewSettingsService(uc *biz.SettingsUsecase, runtime *conf.Runtime, cache bi return &SettingsService{uc: uc, runtime: runtime, cache: cache} } +func (s *SettingsService) CacheGet(ctx context.Context, key string) (string, bool, error) { + return s.cache.Get(ctx, key) +} + +func (s *SettingsService) CacheSet(ctx context.Context, key, value string, expiration time.Duration) error { + return s.cache.Set(ctx, key, value, expiration) +} + +func (s *SettingsService) CacheDelete(ctx context.Context, key string) error { + return s.cache.Delete(ctx, key) +} + +func (s *SettingsService) CacheIncrement(ctx context.Context, key string, expiration time.Duration) (int64, error) { + return s.cache.Increment(ctx, key, expiration) +} + func (s *SettingsService) UseMultipoint() bool { config := s.runtime.Admin() return config != nil && config.System != nil && config.System.UseMultipoint @@ -32,7 +48,7 @@ func (s *SettingsService) ActiveTokenMatches(ctx context.Context, username, toke if !s.UseMultipoint() { return true, nil } - active, ok, err := s.cache.Get(ctx, activeTokenKey(username)) + active, ok, err := s.CacheGet(ctx, activeTokenKey(username)) return ok && active == token, err } @@ -45,5 +61,5 @@ func (s *SettingsService) RotateActiveToken(ctx context.Context, username, oldTo return err } } - return s.cache.Set(ctx, activeTokenKey(username), newToken, expiration) + return s.CacheSet(ctx, activeTokenKey(username), newToken, expiration) } diff --git a/internal/service/system.go b/internal/service/system.go index ba0f352..8d24b67 100644 --- a/internal/service/system.go +++ b/internal/service/system.go @@ -7,16 +7,15 @@ import ( "kra/internal/conf" ) -type SystemService struct { - uc *biz.SystemUsecase - runtime *conf.Runtime - settings *SettingsService +type SystemConfigService struct { + uc *biz.SystemConfigUsecase + runtime *conf.Runtime } -func NewSystemService(uc *biz.SystemUsecase, runtime *conf.Runtime, settings *SettingsService) *SystemService { - return &SystemService{uc: uc, runtime: runtime, settings: settings} +func NewSystemConfigService(uc *biz.SystemConfigUsecase, runtime *conf.Runtime) *SystemConfigService { + return &SystemConfigService{uc: uc, runtime: runtime} } -func (s *SystemService) IsInitialized(ctx context.Context) (bool, error) { +func (s *SystemConfigService) IsInitialized(ctx context.Context) (bool, error) { return s.uc.IsInitialized(ctx) } diff --git a/internal/service/system_config.go b/internal/service/system_config.go index 736d3ff..cefeafc 100644 --- a/internal/service/system_config.go +++ b/internal/service/system_config.go @@ -13,9 +13,11 @@ import ( "google.golang.org/protobuf/types/known/durationpb" ) -func (s *SystemService) PersistConfig(ctx context.Context) error { return s.uc.PersistConfig(ctx) } +func (s *SystemConfigService) PersistConfig(ctx context.Context) error { + return s.uc.PersistConfig(ctx) +} -func (s *SystemService) PersistAdminConfig(ctx context.Context, value *conf.AdminBackend) error { +func (s *SystemConfigService) PersistAdminConfig(ctx context.Context, value *conf.AdminBackend) error { raw, err := protojson.MarshalOptions{UseProtoNames: true}.Marshal(value) if err != nil { return err @@ -23,9 +25,9 @@ func (s *SystemService) PersistAdminConfig(ctx context.Context, value *conf.Admi return s.uc.PersistAdminConfig(ctx, raw) } -func (s *SystemService) ReloadConfig(ctx context.Context) error { return s.uc.ReloadConfig(ctx) } +func (s *SystemConfigService) ReloadConfig(ctx context.Context) error { return s.uc.ReloadConfig(ctx) } -func (s *SystemService) DiskMountPoints() []string { +func (s *SystemConfigService) DiskMountPoints() []string { config := s.runtime.Admin() if config == nil { return nil @@ -39,7 +41,7 @@ func (s *SystemService) DiskMountPoints() []string { return points } -func (s *SystemService) SystemConfig() map[string]any { +func (s *SystemConfigService) SystemConfig() map[string]any { admin := map[string]any{"routerPrefix": ""} email := map[string]any{} config := s.runtime.Admin() @@ -108,7 +110,7 @@ func (s *SystemService) SystemConfig() map[string]any { return map[string]any{"config": map[string]any{"admin": admin, "email": email, "data": dataMap}} } -func (s *SystemService) SaveSystemConfig(ctx context.Context, req *dto.SetSystemConfigRequest) error { +func (s *SystemConfigService) SaveSystemConfig(ctx context.Context, req *dto.SetSystemConfigRequest) error { config := s.runtime.Admin() if config == nil { return nil diff --git a/internal/service/system_init.go b/internal/service/system_init.go index 4acb173..7399afe 100644 --- a/internal/service/system_init.go +++ b/internal/service/system_init.go @@ -20,7 +20,7 @@ type DatabaseInit struct { AdminPassword string `json:"adminPassword" binding:"required"` } -func (s *SystemService) Initialize(ctx context.Context, input *DatabaseInit, apis []*biz.API) error { +func (s *SystemConfigService) Initialize(ctx context.Context, input *DatabaseInit, apis []*biz.API) error { config := "" switch input.DBType { case "mysql": @@ -31,7 +31,7 @@ func (s *SystemService) Initialize(ctx context.Context, input *DatabaseInit, api return s.uc.Initialize(ctx, &biz.DatabaseConfig{Driver: input.DBType, Host: input.Host, Port: input.Port, User: input.UserName, Password: input.Password, Name: input.DBName, Path: input.DBPath, Config: config, Template: input.Template, AdminPassword: input.AdminPassword, APIs: apis}) } -func (s *SystemService) InitializeRoutes(ctx context.Context, input *DatabaseInit, routes []dto.Route) error { +func (s *SystemConfigService) InitializeRoutes(ctx context.Context, input *DatabaseInit, routes []dto.Route) error { apis := make([]*biz.API, 0, len(routes)) for _, route := range routes { path := route.Path diff --git a/internal/service/task.go b/internal/service/task.go index a231ca5..0a2c92b 100644 --- a/internal/service/task.go +++ b/internal/service/task.go @@ -6,18 +6,15 @@ import ( "time" "kra/internal/biz" - "kra/internal/conf" "kra/internal/service/dto" ) type TaskService struct { - uc *biz.TaskUsecase - media *biz.MediaUsecase - runtime *conf.Runtime + uc *biz.TaskUsecase } -func NewTaskService(uc *biz.TaskUsecase, media *biz.MediaUsecase, runtime *conf.Runtime) *TaskService { - return &TaskService{uc: uc, media: media, runtime: runtime} +func NewTaskService(uc *biz.TaskUsecase) *TaskService { + return &TaskService{uc: uc} } func taskDomain(v *dto.TaskRequest) *biz.TimedTask { return &biz.TimedTask{ID: v.ID, Name: v.Name, Description: v.Description, Spec: v.Spec, WithSeconds: v.WithSeconds, ExecutorType: v.ExecutorType, MethodName: v.MethodName, Params: v.Params, HTTPURL: v.HTTPURL, HTTPMethod: v.HTTPMethod, HTTPHeader: v.HTTPHeader, HTTPBody: v.HTTPBody, HTTPAllowPrivate: v.HTTPAllowPrivate, Enabled: v.Enabled} @@ -87,3 +84,12 @@ func (s *TaskService) Logs(ctx context.Context, page, size int, taskID uint, sta } return out, total, nil } + +func (s *TaskService) RegisteredMethods() []map[string]any { + methods := biz.RegisteredTaskMethods() + out := make([]map[string]any, 0, len(methods)) + for _, method := range methods { + out = append(out, map[string]any{"name": method.Name, "description": method.Description}) + } + return out +} diff --git a/internal/service/user.go b/internal/service/user.go index 9c2ab99..2970e8b 100644 --- a/internal/service/user.go +++ b/internal/service/user.go @@ -15,20 +15,29 @@ type UserInput struct { AuthorityIDs []uint } +type UserService struct { + uc *biz.UserUsecase + settings *SettingsService +} + +func NewUserService(uc *biz.UserUsecase, settings *SettingsService) *UserService { + return &UserService{uc: uc, settings: settings} +} + func userInput(value *dto.UserRequest) UserInput { return UserInput{ID: value.ID, Username: value.Username, Password: value.Password, NickName: value.NickName, HeaderImg: value.HeaderImg, AuthorityID: value.AuthorityID, AuthorityIDs: value.AuthorityIDs, Enable: value.Enable, Phone: value.Phone, Email: value.Email} } -func (s *SystemService) ListUsersRequest(ctx context.Context, value *dto.UserListRequest) ([]map[string]any, int64, error) { +func (s *UserService) ListUsersRequest(ctx context.Context, value *dto.UserListRequest) ([]map[string]any, int64, error) { return s.ListUsers(ctx, value.Page, value.PageSize, &biz.UserListFilter{Username: value.Username, NickName: value.NickName, Phone: value.Phone, Email: value.Email, OrderKey: value.OrderKey, Desc: value.Desc}) } -func (s *SystemService) CreateUserRequest(ctx context.Context, value *dto.UserRequest) (map[string]any, error) { +func (s *UserService) CreateUserRequest(ctx context.Context, value *dto.UserRequest) (map[string]any, error) { return s.CreateUser(ctx, userInput(value)) } -func (s *SystemService) UpdateUserRequest(ctx context.Context, value *dto.UserRequest) error { +func (s *UserService) UpdateUserRequest(ctx context.Context, value *dto.UserRequest) error { return s.UpdateUser(ctx, userInput(value)) } -func (s *SystemService) User(ctx context.Context, id uint) (map[string]any, error) { +func (s *UserService) User(ctx context.Context, id uint) (map[string]any, error) { value, err := s.uc.User(ctx, id) if err != nil { return nil, err @@ -36,7 +45,7 @@ func (s *SystemService) User(ctx context.Context, id uint) (map[string]any, erro return convertUser(value), nil } -func (s *SystemService) Menus(ctx context.Context, authorityID uint) ([]map[string]any, error) { +func (s *UserService) Menus(ctx context.Context, authorityID uint) ([]map[string]any, error) { menus, err := s.uc.Menus(ctx, authorityID) if err != nil { return nil, err @@ -48,7 +57,7 @@ func (s *SystemService) Menus(ctx context.Context, authorityID uint) ([]map[stri return result, nil } -func (s *SystemService) ListUsers(ctx context.Context, page, pageSize int, filter *biz.UserListFilter) ([]map[string]any, int64, error) { +func (s *UserService) ListUsers(ctx context.Context, page, pageSize int, filter *biz.UserListFilter) ([]map[string]any, int64, error) { users, total, err := s.uc.ListUsers(ctx, page, pageSize, filter) if err != nil { return nil, 0, err @@ -60,7 +69,7 @@ func (s *SystemService) ListUsers(ctx context.Context, page, pageSize int, filte return result, total, nil } -func (s *SystemService) Authorities(ctx context.Context) ([]map[string]any, error) { +func (s *UserService) Authorities(ctx context.Context) ([]map[string]any, error) { values, err := s.uc.Authorities(ctx) if err != nil { return nil, err @@ -72,7 +81,7 @@ func (s *SystemService) Authorities(ctx context.Context) ([]map[string]any, erro return result, nil } -func (s *SystemService) CreateUser(ctx context.Context, input UserInput) (map[string]any, error) { +func (s *UserService) CreateUser(ctx context.Context, input UserInput) (map[string]any, error) { if err := s.settings.ValidatePassword(ctx, input.Password); err != nil { return nil, err } @@ -84,41 +93,31 @@ func (s *SystemService) CreateUser(ctx context.Context, input UserInput) (map[st } return convertUser(user), nil } -func (s *SystemService) UpdateUser(ctx context.Context, input UserInput) error { +func (s *UserService) UpdateUser(ctx context.Context, input UserInput) error { return s.uc.UpdateUser(ctx, &biz.User{ID: input.ID, NickName: input.NickName, HeaderImg: input.HeaderImg, AuthorityID: input.AuthorityID, Phone: input.Phone, Email: input.Email, Enable: input.Enable}, input.AuthorityIDs) } -func (s *SystemService) UpdateSelfUser(ctx context.Context, input UserInput) error { +func (s *UserService) UpdateSelfUser(ctx context.Context, input UserInput) error { return s.uc.UpdateSelfUser(ctx, &biz.User{ID: input.ID, NickName: input.NickName, HeaderImg: input.HeaderImg, Phone: input.Phone, Email: input.Email, Enable: input.Enable}) } -func (s *SystemService) DeleteUser(ctx context.Context, id uint) error { +func (s *UserService) DeleteUser(ctx context.Context, id uint) error { return s.uc.DeleteUser(ctx, id) } -func (s *SystemService) ResetPassword(ctx context.Context, id uint, password string) error { +func (s *UserService) ResetPassword(ctx context.Context, id uint, password string) error { if err := s.settings.ValidatePassword(ctx, password); err != nil { return err } return s.uc.ResetPassword(ctx, id, password) } -func (s *SystemService) ChangePassword(ctx context.Context, id uint, oldPassword, newPassword string) error { +func (s *UserService) ChangePassword(ctx context.Context, id uint, oldPassword, newPassword string) error { if err := s.settings.ValidatePassword(ctx, newPassword); err != nil { return err } return s.uc.ChangePassword(ctx, id, oldPassword, newPassword) } -func (s *SystemService) SetUserAuthorities(ctx context.Context, id uint, authorityIDs []uint) error { +func (s *UserService) SetUserAuthorities(ctx context.Context, id uint, authorityIDs []uint) error { return s.uc.SetUserAuthorities(ctx, id, authorityIDs) } -func (s *SystemService) SwitchAuthority(ctx context.Context, id, authorityID uint) (*LoginResult, error) { - if err := s.uc.SetUserAuthority(ctx, id, authorityID); err != nil { - return nil, err - } - user, err := s.uc.User(ctx, id) - if err != nil { - return nil, err - } - return s.issueLogin(ctx, user, authorityID) -} -func (s *SystemService) SetUserSetting(ctx context.Context, id uint, setting map[string]any) error { +func (s *UserService) SetUserSetting(ctx context.Context, id uint, setting map[string]any) error { return s.uc.SetUserSetting(ctx, id, setting) } diff --git a/internal/service/task_execution.go b/internal/worker/task_executor.go similarity index 71% rename from internal/service/task_execution.go rename to internal/worker/task_executor.go index 8ea28e4..3d100c1 100644 --- a/internal/service/task_execution.go +++ b/internal/worker/task_executor.go @@ -1,4 +1,4 @@ -package service +package worker import ( "bytes" @@ -15,8 +15,19 @@ import ( "time" "kra/internal/biz" + "kra/internal/conf" ) +type TaskExecutor struct { + tasks *biz.TaskUsecase + media *biz.MediaUsecase + runtime *conf.Runtime +} + +func NewTaskExecutor(tasks *biz.TaskUsecase, media *biz.MediaUsecase, runtime *conf.Runtime) *TaskExecutor { + return &TaskExecutor{tasks: tasks, media: media, runtime: runtime} +} + func privateIP(ip net.IP) bool { return ip.IsLoopback() || ip.IsPrivate() || ip.IsLinkLocalMulticast() || ip.IsLinkLocalUnicast() || ip.IsUnspecified() } @@ -42,28 +53,28 @@ func taskHTTPClient(allowPrivate bool) *http.Client { return &http.Client{Timeout: 30 * time.Second, Transport: &http.Transport{Proxy: nil, DialContext: dialer.DialContext}} } -func (s *TaskService) runHTTP(ctx context.Context, v *biz.TimedTask) (string, error) { - parsed, err := url.Parse(v.HTTPURL) +func (e *TaskExecutor) runHTTP(ctx context.Context, task *biz.TimedTask) (string, error) { + parsed, err := url.Parse(task.HTTPURL) if err != nil { return "", err } if parsed.Scheme != "http" && parsed.Scheme != "https" { return "", fmt.Errorf("仅允许 http/https, 实际为 %q", parsed.Scheme) } - method := strings.ToUpper(v.HTTPMethod) + method := strings.ToUpper(task.HTTPMethod) if method == "" { method = http.MethodGet } - request, err := http.NewRequestWithContext(ctx, method, v.HTTPURL, bytes.NewBufferString(v.HTTPBody)) + request, err := http.NewRequestWithContext(ctx, method, task.HTTPURL, bytes.NewBufferString(task.HTTPBody)) if err != nil { return "", err } headers := map[string]string{} - _ = json.Unmarshal(v.HTTPHeader, &headers) + _ = json.Unmarshal(task.HTTPHeader, &headers) for key, value := range headers { request.Header.Set(key, value) } - response, err := taskHTTPClient(v.HTTPAllowPrivate).Do(request) + response, err := taskHTTPClient(task.HTTPAllowPrivate).Do(request) if err != nil { return "", err } @@ -86,7 +97,7 @@ func truncateTaskText(value string) string { return value[:limit] + "...(截断)" } -func (s *TaskService) runMethod(v *biz.TimedTask) error { +func (e *TaskExecutor) runMethod(task *biz.TimedTask) error { ctx, cancel := context.WithTimeout(context.Background(), 5*time.Minute) defer cancel() done := make(chan error, 1) @@ -97,16 +108,16 @@ func (s *TaskService) runMethod(v *biz.TimedTask) error { } }() var err error - switch v.MethodName { + switch task.MethodName { case biz.TaskMethodClearDB: - err = s.uc.CleanupLogs(ctx) + err = e.tasks.CleanupLogs(ctx) case biz.TaskMethodUploads: ttl := 24 - config := s.runtime.Admin() + config := e.runtime.Admin() if config != nil && config.Media != nil && config.Media.SessionTtl > 0 { ttl = int(config.Media.SessionTtl) } - err = s.media.CleanupStale(ctx, ttl) + err = e.media.CleanupStale(ctx, ttl) } done <- err }() @@ -121,9 +132,9 @@ func (s *TaskService) runMethod(v *biz.TimedTask) error { } } -func (s *TaskService) Run(ctx context.Context, v *biz.TimedTask, trigger string) (log *biz.TimedTaskLog) { +func (e *TaskExecutor) Run(ctx context.Context, task *biz.TimedTask, trigger string) (log *biz.TimedTaskLog) { started := time.Now() - log = &biz.TimedTaskLog{TaskID: v.ID, TaskName: v.Name, TriggerType: trigger, StartedAt: started, Status: "success"} + log = &biz.TimedTaskLog{TaskID: task.ID, TaskName: task.Name, TriggerType: trigger, StartedAt: started, Status: "success"} defer func() { if recovered := recover(); recovered != nil { log.Status = "fail" @@ -134,18 +145,18 @@ func (s *TaskService) Run(ctx context.Context, v *biz.TimedTask, trigger string) log.DurationMS = log.FinishedAt.Sub(started).Milliseconds() logCtx, cancel := context.WithTimeout(context.WithoutCancel(ctx), 3*time.Second) defer cancel() - _ = s.uc.RecordTaskLog(logCtx, log) + _ = e.tasks.RecordTaskLog(logCtx, log) }() var err error - switch v.ExecutorType { + switch task.ExecutorType { case biz.TaskExecutorMethod: - err = s.runMethod(v) + err = e.runMethod(task) case biz.TaskExecutorHTTP: runCtx, cancel := context.WithTimeout(context.Background(), 30*time.Second) - log.Output, err = s.runHTTP(runCtx, v) + log.Output, err = e.runHTTP(runCtx, task) cancel() default: - err = fmt.Errorf("未知执行器类型: %s", v.ExecutorType) + err = fmt.Errorf("未知执行器类型: %s", task.ExecutorType) } if err != nil { if errors.Is(err, errTaskTimeout) || errors.Is(err, context.DeadlineExceeded) { @@ -157,11 +168,3 @@ func (s *TaskService) Run(ctx context.Context, v *biz.TimedTask, trigger string) } return log } -func (s *TaskService) RegisteredMethods() []map[string]any { - methods := biz.RegisteredTaskMethods() - out := make([]map[string]any, 0, len(methods)) - for _, method := range methods { - out = append(out, map[string]any{"name": method.Name, "description": method.Description}) - } - return out -} diff --git a/internal/service/task_execution_test.go b/internal/worker/task_executor_test.go similarity index 95% rename from internal/service/task_execution_test.go rename to internal/worker/task_executor_test.go index a8c323f..051f911 100644 --- a/internal/service/task_execution_test.go +++ b/internal/worker/task_executor_test.go @@ -1,4 +1,4 @@ -package service +package worker import ( "net" diff --git a/internal/worker/task_scheduler.go b/internal/worker/task_scheduler.go index a9ca4f9..4436010 100644 --- a/internal/worker/task_scheduler.go +++ b/internal/worker/task_scheduler.go @@ -8,14 +8,13 @@ import ( "sync" "time" - "kra/internal/biz" - "kra/internal/service" - "github.com/robfig/cron/v3" + "kra/internal/biz" ) type TaskScheduler struct { - service *service.TaskService + tasks *biz.TaskUsecase + executor *TaskExecutor logger *slog.Logger standard *cron.Cron seconds *cron.Cron @@ -33,8 +32,8 @@ type scheduledEntry struct { entry cron.EntryID } -func NewTaskScheduler(service *service.TaskService, logger *slog.Logger) *TaskScheduler { - return &TaskScheduler{service: service, logger: logger, standard: cron.New(), seconds: cron.New(cron.WithSeconds()), entries: map[uint]scheduledEntry{}, subscribers: map[chan []byte]struct{}{}} +func NewTaskScheduler(tasks *biz.TaskUsecase, executor *TaskExecutor, logger *slog.Logger) *TaskScheduler { + return &TaskScheduler{tasks: tasks, executor: executor, logger: logger, standard: cron.New(), seconds: cron.New(cron.WithSeconds()), entries: map[uint]scheduledEntry{}, subscribers: map[chan []byte]struct{}{}} } func (s *TaskScheduler) Start(ctx context.Context) error { @@ -44,7 +43,7 @@ func (s *TaskScheduler) Start(ctx context.Context) error { s.ctxMu.Unlock() s.standard.Start() s.seconds.Start() - items, err := s.service.ScheduledTasks(ctx) + items, _, err := s.tasks.ListTasks(ctx, 0, 0, nil) if err == nil { for _, task := range items { if task.Enabled { @@ -104,7 +103,7 @@ func (s *TaskScheduler) Reload(ctx context.Context) error { delete(s.entries, id) } s.mu.Unlock() - items, err := s.service.ScheduledTasks(ctx) + items, _, err := s.tasks.ListTasks(ctx, 0, 0, nil) if err != nil { return err } @@ -128,7 +127,7 @@ func (s *TaskScheduler) executionContext() context.Context { } func (s *TaskScheduler) run(task *biz.TimedTask, trigger string) { - log := s.service.Run(s.executionContext(), task, trigger) + log := s.executor.Run(s.executionContext(), task, trigger) if log.Status != "success" { s.Broadcast(map[string]any{"taskId": task.ID, "taskName": task.Name, "status": log.Status, "errorMsg": log.ErrorMsg, "time": time.Now()}) } @@ -160,7 +159,7 @@ func (s *TaskScheduler) Schedule(task *biz.TimedTask) error { } func (s *TaskScheduler) ScheduleID(ctx context.Context, id uint) error { - task, err := s.service.Task(ctx, id) + task, err := s.tasks.FindTask(ctx, id) if err != nil { return err } @@ -168,7 +167,7 @@ func (s *TaskScheduler) ScheduleID(ctx context.Context, id uint) error { } func (s *TaskScheduler) TriggerID(ctx context.Context, id uint) error { - task, err := s.service.Task(ctx, id) + task, err := s.tasks.FindTask(ctx, id) if err != nil { return err }