diff --git a/backend/core/appctx.go b/backend/core/appctx.go new file mode 100644 index 00000000..1c1e2201 --- /dev/null +++ b/backend/core/appctx.go @@ -0,0 +1,43 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package core + +import "context" + +type appContextKey struct{} + +// WithAppContext attaches the micro-kernel Context to a standard context.Context +// so request and worker handlers can Inject services without package-level setters. +func WithAppContext(ctx context.Context, app *Context) context.Context { + if ctx == nil { + ctx = context.Background() + } + if app == nil { + return ctx + } + return context.WithValue(ctx, appContextKey{}, app.Root()) +} + +// AppContext extracts the micro-kernel Context from ctx, if present. +func AppContext(ctx context.Context) *Context { + if ctx == nil { + return nil + } + if c, ok := ctx.(*Context); ok { + return c + } + app, _ := ctx.Value(appContextKey{}).(*Context) + return app +} + +// InjectFrom resolves T from ctx when it carries a micro-kernel Context +// (*Context itself, or a value attached by WithAppContext). +func InjectFrom[T any](ctx context.Context) (T, error) { + var zero T + app := AppContext(ctx) + if app == nil { + return zero, ErrNilContext + } + return Inject[T](app) +} diff --git a/backend/core/container.go b/backend/core/container.go index ed783572..64d8c030 100644 --- a/backend/core/container.go +++ b/backend/core/container.go @@ -201,13 +201,17 @@ func Using3[T1, T2, T3 any](ctx *Context, fn func(s1 T1, s2 T2, s3 T3)) error { // When registers a reactive hook that is called immediately if T is already provided, // or called as soon as T is provided in the future. +// +// Listeners are stored on the root container so they observe core.Provide, which +// always writes to the root. Registering on a Fiber child container would miss +// services provided by plugins that load later. func When[T any](ctx *Context, fn func(s T)) { if ctx == nil { panic("core: nil context provided to When") } targetType := reflect.TypeFor[T]() - c := ctx.Container() + c := ctx.Root().Container() // If already ready, execute immediately if s, err := Inject[T](ctx); err == nil { @@ -223,3 +227,9 @@ func When[T any](ctx *Context, fn func(s T)) { } }) } + +// Bind is When with a name that matches plugin wiring: fill a dependency as +// soon as the root container provides it. +func Bind[T any](ctx *Context, fn func(s T)) { + When(ctx, fn) +} diff --git a/backend/core/context_test.go b/backend/core/context_test.go index d6c66bc6..9a9bce93 100644 --- a/backend/core/context_test.go +++ b/backend/core/context_test.go @@ -313,6 +313,46 @@ func TestContextReactiveWhen(t *testing.T) { assert.True(t, immediateCalled) } +func TestWhenObservesProvideFromForkedFiberContext(t *testing.T) { + root := core.NewContext(context.Background()) + adminFiber := root.Fork() + lateFiber := root.Fork() + + var got atomic.Bool + core.When[SampleService](adminFiber, func(s SampleService) { + if s != nil { + got.Store(true) + } + }) + assert.False(t, got.Load()) + + core.Provide[SampleService](lateFiber, &sampleServiceImpl{}) + assert.True(t, got.Load(), "When on a Fiber child must observe Provide on the root") +} + +func TestBindIsWhen(t *testing.T) { + ctx := core.NewContext(context.Background()) + var called atomic.Bool + core.Bind[SampleService](ctx, func(s SampleService) { + called.Store(true) + }) + core.Provide[SampleService](ctx, &sampleServiceImpl{}) + assert.True(t, called.Load()) +} + +func TestInjectFromAppContext(t *testing.T) { + app := core.NewContext(context.Background()) + core.Provide[SampleService](app, &sampleServiceImpl{prefix: "Hi:"}) + + req := core.WithAppContext(context.Background(), app) + svc, err := core.InjectFrom[SampleService](req) + require.NoError(t, err) + assert.Equal(t, "Hi: Ada", svc.Greet("Ada")) + + _, err = core.InjectFrom[SampleService](context.Background()) + assert.ErrorIs(t, err, core.ErrNilContext) +} + func TestContextDisposerLifecycle(t *testing.T) { parent := core.NewContext(context.Background()) child := parent.Fork() diff --git a/backend/core/contracts/task.go b/backend/core/contracts/task.go index 3fe1c01c..838d6afa 100644 --- a/backend/core/contracts/task.go +++ b/backend/core/contracts/task.go @@ -43,6 +43,12 @@ type TaskResultDTO struct { Detail any `json:"detail,omitempty"` } +// TaskHandler is the preferred background task handler. Drivers invoke Execute +// and persist Message/Detail onto the execution record. +type TaskHandler interface { + Execute(ctx context.Context, payload []byte) (*TaskResultDTO, error) +} + // TaskExecutionDTO represents a single task execution record. type TaskExecutionDTO struct { ID uint64 `json:"id,string"` diff --git a/backend/core/extpoints/extpoints_test.go b/backend/core/extpoints/extpoints_test.go index d0c438f4..fe684d75 100644 --- a/backend/core/extpoints/extpoints_test.go +++ b/backend/core/extpoints/extpoints_test.go @@ -229,6 +229,22 @@ func TestTaskExtension(t *testing.T) { assert.False(t, ok) } +func TestTaskRegisterRejectsNilHandler(t *testing.T) { + tr := extpoints.NewTaskRegistry() + assert.Panics(t, func() { + tr.Register("broken:task", nil) + }) +} + +func TestTaskRegisterRejectsDuplicateType(t *testing.T) { + tr := extpoints.NewTaskRegistry() + handler := func(ctx context.Context, payload []byte) error { return nil } + tr.Register("system:cleanup", handler, extpoints.WithTaskType("system_cleanup")) + assert.Panics(t, func() { + tr.Register("admin:system_cleanup", handler, extpoints.WithTaskType("system_cleanup")) + }) +} + func TestScheduleExtension(t *testing.T) { sr := extpoints.NewScheduleRegistry() require.NotNil(t, sr) diff --git a/backend/core/extpoints/task.go b/backend/core/extpoints/task.go index 09270158..3e22b1e3 100644 --- a/backend/core/extpoints/task.go +++ b/backend/core/extpoints/task.go @@ -5,6 +5,8 @@ package extpoints import ( "Wavelet/core/contracts" + "fmt" + "reflect" "sync" "time" ) @@ -230,10 +232,15 @@ func NewTaskRegistry() *TaskRegistry { } // Register registers a task pattern and its handler with optional configuration. +// A nil handler panics. A non-empty Type that is already used by another pattern panics. func (t *TaskRegistry) Register(pattern string, handler any, opts ...TaskOption) { t.mu.Lock() defer t.mu.Unlock() + if isNilTaskHandler(handler) { + panic(fmt.Sprintf("extpoints: nil handler for task pattern %q", pattern)) + } + td := TaskDefinition{ Pattern: pattern, Handler: handler, @@ -245,8 +252,23 @@ func (t *TaskRegistry) Register(pattern string, handler any, opts ...TaskOption) opt(&td) } } + if td.Type == "" { + td.Type = pattern + } - if _, exists := t.lookup[pattern]; exists { + for _, item := range t.tasks { + if item.Pattern == pattern { + continue + } + if item.Type == td.Type { + panic(fmt.Sprintf("extpoints: duplicate task type %q (patterns %q and %q)", td.Type, item.Pattern, pattern)) + } + } + + if existing, exists := t.lookup[pattern]; exists { + if existing.Type != "" && existing.Type != pattern { + delete(t.lookup, existing.Type) + } for i, item := range t.tasks { if item.Pattern == pattern { t.tasks[i] = td @@ -258,13 +280,44 @@ func (t *TaskRegistry) Register(pattern string, handler any, opts ...TaskOption) } t.lookup[pattern] = td + if td.Type != pattern { + t.lookup[td.Type] = td + } +} + +func isNilTaskHandler(handler any) bool { + if handler == nil { + return true + } + v := reflect.ValueOf(handler) + switch v.Kind() { + case reflect.Chan, reflect.Func, reflect.Map, reflect.Pointer, reflect.UnsafePointer, reflect.Interface, reflect.Slice: + return v.IsNil() + default: + return false + } } // Unregister removes a registered task definition by its pattern. func (t *TaskRegistry) Unregister(pattern string) bool { - return unregisterEntry(&t.mu, t.lookup, &t.tasks, pattern, func(item TaskDefinition) bool { - return item.Pattern == pattern - }) + t.mu.Lock() + defer t.mu.Unlock() + td, ok := t.lookup[pattern] + if !ok { + return false + } + delete(t.lookup, td.Pattern) + if td.Type != "" && td.Type != td.Pattern { + delete(t.lookup, td.Type) + } + filtered := t.tasks[:0] + for _, item := range t.tasks { + if item.Pattern != td.Pattern { + filtered = append(filtered, item) + } + } + t.tasks = filtered + return true } // Tasks returns a copy of all registered TaskDefinitions. @@ -280,13 +333,6 @@ func (t *TaskRegistry) Tasks() []TaskDefinition { func (t *TaskRegistry) Get(pattern string) (TaskDefinition, bool) { t.mu.RLock() defer t.mu.RUnlock() - if td, ok := t.lookup[pattern]; ok { - return td, true - } - for _, td := range t.tasks { - if td.Type == pattern { - return td, true - } - } - return TaskDefinition{}, false + td, ok := t.lookup[pattern] + return td, ok } diff --git a/backend/docs/docs.go b/backend/docs/docs.go index 21314fc6..028c21b1 100644 --- a/backend/docs/docs.go +++ b/backend/docs/docs.go @@ -5925,6 +5925,17 @@ const docTemplate = `{ "user" ], "summary": "发送邮箱验证码", + "parameters": [ + { + "description": "目标邮箱", + "name": "request", + "in": "body", + "required": true, + "schema": { + "$ref": "#/definitions/user.sendEmailCodeRequest" + } + } + ], "responses": { "200": { "description": "发送成功", @@ -5937,6 +5948,12 @@ const docTemplate = `{ "schema": { "$ref": "#/definitions/response.Any" } + }, + "500": { + "description": "发送失败", + "schema": { + "$ref": "#/definitions/response.Any" + } } } } @@ -8047,6 +8064,17 @@ const docTemplate = `{ } } }, + "user.sendEmailCodeRequest": { + "type": "object", + "required": [ + "email" + ], + "properties": { + "email": { + "type": "string" + } + } + }, "user.updateProfileRequest": { "type": "object", "properties": { diff --git a/backend/docs/swagger.json b/backend/docs/swagger.json index 84f27bc9..a01fa772 100644 --- a/backend/docs/swagger.json +++ b/backend/docs/swagger.json @@ -5918,6 +5918,17 @@ "user" ], "summary": "发送邮箱验证码", + "parameters": [ + { + "description": "目标邮箱", + "name": "request", + "in": "body", + "required": true, + "schema": { + "$ref": "#/definitions/user.sendEmailCodeRequest" + } + } + ], "responses": { "200": { "description": "发送成功", @@ -5930,6 +5941,12 @@ "schema": { "$ref": "#/definitions/response.Any" } + }, + "500": { + "description": "发送失败", + "schema": { + "$ref": "#/definitions/response.Any" + } } } } @@ -8040,6 +8057,17 @@ } } }, + "user.sendEmailCodeRequest": { + "type": "object", + "required": [ + "email" + ], + "properties": { + "email": { + "type": "string" + } + } + }, "user.updateProfileRequest": { "type": "object", "properties": { diff --git a/backend/docs/swagger.yaml b/backend/docs/swagger.yaml index 555b8a9c..52dbb9ce 100644 --- a/backend/docs/swagger.yaml +++ b/backend/docs/swagger.yaml @@ -1363,6 +1363,13 @@ definitions: - password - username type: object + user.sendEmailCodeRequest: + properties: + email: + type: string + required: + - email + type: object user.updateProfileRequest: properties: avatar_url: @@ -4937,6 +4944,13 @@ paths: consumes: - application/json description: 向指定邮箱发送验证码(用于注册场景) + parameters: + - description: 目标邮箱 + in: body + name: request + required: true + schema: + $ref: '#/definitions/user.sendEmailCodeRequest' produces: - application/json responses: @@ -4948,6 +4962,10 @@ paths: description: 参数错误 schema: $ref: '#/definitions/response.Any' + "500": + description: 发送失败 + schema: + $ref: '#/definitions/response.Any' summary: 发送邮箱验证码 tags: - user diff --git a/backend/plugins/domain/admin/handler/tasks.go b/backend/plugins/domain/admin/handler/tasks.go index 438cae3f..2e396c67 100644 --- a/backend/plugins/domain/admin/handler/tasks.go +++ b/backend/plugins/domain/admin/handler/tasks.go @@ -48,7 +48,7 @@ func abortTaskLogicError(c *gin.Context, err error) bool { // @Failure 403 {object} response.Any "无管理员权限" // @Router /api/v1/admin/tasks/types [get] func ListTaskTypes(c *gin.Context) { - c.JSON(http.StatusOK, response.OK(service.ListTaskTypes())) + c.JSON(http.StatusOK, response.OK(service.ListTaskTypes(c.Request.Context()))) } // DispatchTask 下发任务 diff --git a/backend/plugins/domain/admin/plugin.go b/backend/plugins/domain/admin/plugin.go index 835e2e80..418f1d11 100644 --- a/backend/plugins/domain/admin/plugin.go +++ b/backend/plugins/domain/admin/plugin.go @@ -12,7 +12,6 @@ import ( "Wavelet/plugins/domain/admin/handler" "Wavelet/plugins/domain/admin/model" "Wavelet/plugins/domain/admin/service" - "context" "embed" "reflect" @@ -85,56 +84,13 @@ func (p *Plugin) Apply(ctx *core.Context) error { _ = ctx.Config().Bind("clickhouse", &chCfg) service.SetClickHouseConfig(chCfg) - // 0. Bind Services reactively - if db, err := core.Inject[contracts.DBService](ctx); err == nil && db != nil { - service.SetDBService(db) - } else { - core.When[contracts.DBService](ctx, func(db contracts.DBService) { - service.SetDBService(db) - }) - } - if cache, err := core.Inject[contracts.CacheService](ctx); err == nil && cache != nil { - service.SetCacheService(cache) - } else { - core.When[contracts.CacheService](ctx, func(cache contracts.CacheService) { - service.SetCacheService(cache) - }) - } - if user, err := core.Inject[contracts.UserService](ctx); err == nil && user != nil { - service.SetUserService(user) - } else { - core.When[contracts.UserService](ctx, func(user contracts.UserService) { - service.SetUserService(user) - }) - } - if auth, err := core.Inject[contracts.AuthService](ctx); err == nil && auth != nil { - service.SetAuthService(auth) - } else { - core.When[contracts.AuthService](ctx, func(auth contracts.AuthService) { - service.SetAuthService(auth) - }) - } - if task, err := core.Inject[contracts.TaskService](ctx); err == nil && task != nil { - service.SetTaskService(task) - } else { - core.When[contracts.TaskService](ctx, func(task contracts.TaskService) { - service.SetTaskService(task) - }) - } - if storage, err := core.Inject[contracts.StorageService](ctx); err == nil && storage != nil { - service.SetStorageService(storage) - } else { - core.When[contracts.StorageService](ctx, func(storage contracts.StorageService) { - service.SetStorageService(storage) - }) - } - if rc, err := core.Inject[contracts.RiskControlService](ctx); err == nil && rc != nil { - service.SetRiskControlService(rc) - } else { - core.When[contracts.RiskControlService](ctx, func(rc contracts.RiskControlService) { - service.SetRiskControlService(rc) - }) - } + core.Bind[contracts.DBService](ctx, service.SetDBService) + core.Bind[contracts.CacheService](ctx, service.SetCacheService) + core.Bind[contracts.UserService](ctx, service.SetUserService) + core.Bind[contracts.AuthService](ctx, service.SetAuthService) + core.Bind[contracts.TaskService](ctx, service.SetTaskService) + core.Bind[contracts.StorageService](ctx, service.SetStorageService) + core.Bind[contracts.RiskControlService](ctx, service.SetRiskControlService) service.SetEventEmitter(ctx.Events().Emit) ctx.OnDispose(func() error { @@ -175,11 +131,7 @@ func (p *Plugin) Apply(ctx *core.Context) error { ctx.Router().RegisterWhitelist("/robots.txt") // 2. Register Background Tasks - logSwitchHandler := &service.LogDBSwitchHandler{} - ctx.Task().Register(service.LogDBSwitchTask, func(c context.Context, payload []byte) error { - _, err := logSwitchHandler.Execute(c, payload) - return err - }, extpoints.WithTaskMeta(service.LogDBSwitchMeta)) + ctx.Task().Register(service.LogDBSwitchTask, &service.LogDBSwitchHandler{}, extpoints.WithTaskMeta(service.LogDBSwitchMeta)) // 3. Register Settings Schemas ctx.Settings().Register(extpoints.SettingSchema{ diff --git a/backend/plugins/domain/admin/repository/repository.go b/backend/plugins/domain/admin/repository/repository.go index 6118614a..c0c2331a 100644 --- a/backend/plugins/domain/admin/repository/repository.go +++ b/backend/plugins/domain/admin/repository/repository.go @@ -5,6 +5,7 @@ package repository import ( + "Wavelet/core" "Wavelet/core/contracts" "Wavelet/pkg/cache/ram" "Wavelet/pkg/logger" @@ -56,6 +57,9 @@ func ResetServices() { // GetDB returns the GORM DB instance bound to the context if available. func GetDB(ctx context.Context) *gorm.DB { + if s, err := core.InjectFrom[contracts.DBService](ctx); err == nil && s != nil { + return s.DB(ctx) + } repoMu.RLock() defer repoMu.RUnlock() if dbService == nil { @@ -65,7 +69,10 @@ func GetDB(ctx context.Context) *gorm.DB { } // GetCache returns the unified CacheService instance. -func GetCache(_ context.Context) contracts.CacheService { +func GetCache(ctx context.Context) contracts.CacheService { + if s, err := core.InjectFrom[contracts.CacheService](ctx); err == nil && s != nil { + return s + } repoMu.RLock() defer repoMu.RUnlock() return cacheService diff --git a/backend/plugins/domain/admin/service/log.go b/backend/plugins/domain/admin/service/log.go index 82a6dd84..e0cca4a6 100644 --- a/backend/plugins/domain/admin/service/log.go +++ b/backend/plugins/domain/admin/service/log.go @@ -75,7 +75,7 @@ func IsAllowedLogOrigin(ctx context.Context, origin, host string) bool { // AccessLogs queries the analytical access log store and decorates rows with user names. func AccessLogs(ctx context.Context, q model.AccessLogQuery) (model.AccessLogsResponse, error) { - rc := GetRiskControlService() + rc := GetRiskControlService(ctx) if rc == nil { return model.AccessLogsResponse{}, errs.ErrLogStoreUnavailable } @@ -117,7 +117,7 @@ func AccessLogs(ctx context.Context, q model.AccessLogQuery) (model.AccessLogsRe // AccessLogAnalytics aggregates the daily trend of the access log store. func AccessLogAnalytics(ctx context.Context) (model.LogsAnalyticsResponse, error) { - rc := GetRiskControlService() + rc := GetRiskControlService(ctx) if rc == nil { return model.LogsAnalyticsResponse{}, errs.ErrLogStoreUnavailable } diff --git a/backend/plugins/domain/admin/service/log_switch.go b/backend/plugins/domain/admin/service/log_switch.go index 792eb95f..a6ae1091 100644 --- a/backend/plugins/domain/admin/service/log_switch.go +++ b/backend/plugins/domain/admin/service/log_switch.go @@ -109,7 +109,7 @@ func (h *LogDBSwitchHandler) Execute(ctx context.Context, payload []byte) (*cont return nil, err } - taskSvc := GetTaskService() + taskSvc := GetTaskService(ctx) if taskSvc != nil { taskSvc.AppendLog(ctx, "开始切换日志数据库:%s -> %s", source, p.Target) } @@ -123,7 +123,7 @@ func (h *LogDBSwitchHandler) Execute(ctx context.Context, payload []byte) (*cont } }() - rc := GetRiskControlService() + rc := GetRiskControlService(ctx) if rc != nil { if err := rc.SwitchLogEngine(ctx, p.Target); err != nil { return nil, err diff --git a/backend/plugins/domain/admin/service/service.go b/backend/plugins/domain/admin/service/service.go index ba4a2e98..db474f87 100644 --- a/backend/plugins/domain/admin/service/service.go +++ b/backend/plugins/domain/admin/service/service.go @@ -5,6 +5,7 @@ package service import ( + "Wavelet/core" "Wavelet/core/contracts" "Wavelet/plugins/domain/admin/errs" "Wavelet/plugins/domain/admin/repository" @@ -112,6 +113,9 @@ func ResetServices() { // GetDB returns the GORM DB instance bound to the context if available. func GetDB(ctx context.Context) *gorm.DB { + if s, err := core.InjectFrom[contracts.DBService](ctx); err == nil && s != nil { + return s.DB(ctx) + } servicesMu.RLock() defer servicesMu.RUnlock() if dbService == nil { @@ -121,42 +125,60 @@ func GetDB(ctx context.Context) *gorm.DB { } // GetCache returns the unified CacheService instance. -func GetCache(_ context.Context) contracts.CacheService { +func GetCache(ctx context.Context) contracts.CacheService { + if s, err := core.InjectFrom[contracts.CacheService](ctx); err == nil && s != nil { + return s + } servicesMu.RLock() defer servicesMu.RUnlock() return cacheService } // GetUserService returns the UserService instance. -func GetUserService(_ context.Context) contracts.UserService { +func GetUserService(ctx context.Context) contracts.UserService { + if s, err := core.InjectFrom[contracts.UserService](ctx); err == nil && s != nil { + return s + } servicesMu.RLock() defer servicesMu.RUnlock() return userService } // GetAuthService returns the AuthService instance. -func GetAuthService(_ context.Context) contracts.AuthService { +func GetAuthService(ctx context.Context) contracts.AuthService { + if s, err := core.InjectFrom[contracts.AuthService](ctx); err == nil && s != nil { + return s + } servicesMu.RLock() defer servicesMu.RUnlock() return authService } // GetTaskService returns the TaskService instance. -func GetTaskService() contracts.TaskService { +func GetTaskService(ctx context.Context) contracts.TaskService { + if s, err := core.InjectFrom[contracts.TaskService](ctx); err == nil && s != nil { + return s + } servicesMu.RLock() defer servicesMu.RUnlock() return taskService } // GetStorageService returns the StorageService instance. -func GetStorageService() contracts.StorageService { +func GetStorageService(ctx context.Context) contracts.StorageService { + if s, err := core.InjectFrom[contracts.StorageService](ctx); err == nil && s != nil { + return s + } servicesMu.RLock() defer servicesMu.RUnlock() return storageSvc } // GetRiskControlService returns the RiskControlService instance. -func GetRiskControlService() contracts.RiskControlService { +func GetRiskControlService(ctx context.Context) contracts.RiskControlService { + if s, err := core.InjectFrom[contracts.RiskControlService](ctx); err == nil && s != nil { + return s + } servicesMu.RLock() defer servicesMu.RUnlock() return riskControlService @@ -195,8 +217,8 @@ func requireAuthService(ctx context.Context) (contracts.AuthService, error) { } // requireTaskService resolves the injected task contract service. -func requireTaskService() (contracts.TaskService, error) { - taskSvc := GetTaskService() +func requireTaskService(ctx context.Context) (contracts.TaskService, error) { + taskSvc := GetTaskService(ctx) if taskSvc == nil { return nil, errs.ErrTaskServiceUnavailable } diff --git a/backend/plugins/domain/admin/service/status.go b/backend/plugins/domain/admin/service/status.go index dd3ef455..77c25647 100644 --- a/backend/plugins/domain/admin/service/status.go +++ b/backend/plugins/domain/admin/service/status.go @@ -117,7 +117,7 @@ func formatDuration(d time.Duration) string { func LogDatabaseStatus(ctx context.Context) model.LogDatabaseStatus { activeDB := logDBNameSQLite migration := logMigrationIdle - if rc := GetRiskControlService(); rc != nil { + if rc := GetRiskControlService(ctx); rc != nil { activeDB = rc.ActiveLogEngine(ctx) if rc.IsLogEngineMigrating(ctx) { migration = logMigrationInProgress diff --git a/backend/plugins/domain/admin/service/task.go b/backend/plugins/domain/admin/service/task.go index 2543ee85..3226c31f 100644 --- a/backend/plugins/domain/admin/service/task.go +++ b/backend/plugins/domain/admin/service/task.go @@ -17,8 +17,8 @@ import ( ) // ListTaskTypes returns every dispatchable task type declared in the task registry. -func ListTaskTypes() []contracts.TaskMetaDTO { - taskSvc := GetTaskService() +func ListTaskTypes(ctx context.Context) []contracts.TaskMetaDTO { + taskSvc := GetTaskService(ctx) if taskSvc == nil { return []contracts.TaskMetaDTO{} } @@ -27,7 +27,7 @@ func ListTaskTypes() []contracts.TaskMetaDTO { // DispatchTask validates and enqueues a manual task run, returning the new task id. func DispatchTask(ctx context.Context, req model.DispatchTaskRequest) (string, error) { - taskSvc, err := requireTaskService() + taskSvc, err := requireTaskService(ctx) if err != nil { return "", err } @@ -67,29 +67,75 @@ func ListTaskExecutions( ctx context.Context, req model.ListTaskExecutionsRequest, ) ([]model.TaskExecution, int64, error) { - if req.TaskType != "" { - if taskSvc := GetTaskService(); taskSvc != nil { - if meta, ok := taskSvc.GetTaskMeta(req.TaskType); ok { - req.TaskType = meta.Name - } + taskSvc, err := requireTaskService(ctx) + if err != nil { + return nil, 0, err + } + if req.Page <= 0 { + req.Page = 1 + } + if req.PageSize <= 0 { + req.PageSize = 20 + } + + filterType := req.TaskType + if filterType != "" { + if meta, ok := taskSvc.GetTaskMeta(filterType); ok { + filterType = meta.AsynqTask } } - executions, total, err := repository.ListTaskExecutionRecords(ctx, req) + rows, total, err := taskSvc.ListExecutions(ctx, filterType, req.Status, req.Page, req.PageSize) if err != nil { return nil, 0, err } + executions := make([]model.TaskExecution, 0, len(rows)) + for i := range rows { + executions = append(executions, executionFromDTO(rows[i])) + } return executions, total, nil } // TaskExecution loads a single execution record including its buffered log. func TaskExecution(ctx context.Context, id uint64) (*model.TaskExecution, error) { - return repository.GetTaskExecutionByID(ctx, id) + taskSvc, err := requireTaskService(ctx) + if err != nil { + return nil, err + } + dto, err := taskSvc.GetExecution(ctx, id) + if err != nil || dto == nil { + return nil, err + } + row := executionFromDTO(*dto) + return &row, nil +} + +func executionFromDTO(dto contracts.TaskExecutionDTO) model.TaskExecution { + return model.TaskExecution{ + ID: dto.ID, + TaskID: dto.TaskID, + TaskType: dto.TaskType, + TaskName: dto.TaskName, + Status: model.TaskExecutionStatus(dto.Status), + Retryable: dto.Retryable, + MaxRetry: dto.MaxRetry, + RetryCount: dto.RetryCount, + Log: dto.Log, + ErrorMessage: dto.ErrorMessage, + Result: dto.Result, + StartedAt: dto.StartedAt, + FinishedAt: dto.FinishedAt, + Duration: dto.Duration, + Payload: dto.Payload, + TriggeredBy: dto.TriggeredBy, + CreatedAt: dto.CreatedAt, + UpdatedAt: dto.UpdatedAt, + } } // RetryTask re-dispatches a failed execution as a new task run. func RetryTask(ctx context.Context, id uint64) (string, error) { - taskSvc, err := requireTaskService() + taskSvc, err := requireTaskService(ctx) if err != nil { return "", err } @@ -126,7 +172,7 @@ func CreateSchedule(ctx context.Context, req model.CreateScheduleRequest) (*mode return nil, errs.ErrInvalidCronExpression } - taskSvc, err := requireTaskService() + taskSvc, err := requireTaskService(ctx) if err != nil { return nil, err } @@ -168,7 +214,7 @@ func UpdateSchedule(ctx context.Context, id uint64, req model.UpdateScheduleRequ return nil, errs.ErrInvalidCronExpression } - taskSvc, err := requireTaskService() + taskSvc, err := requireTaskService(ctx) if err != nil { return nil, err } @@ -203,7 +249,7 @@ func DeleteSchedule(ctx context.Context, id uint64) error { return fmt.Errorf("%s: %w", errs.ScheduleDeleteFailed, err) } - if taskSvc := GetTaskService(); taskSvc != nil { + if taskSvc := GetTaskService(ctx); taskSvc != nil { reloadScheduler(ctx, taskSvc) } return nil diff --git a/backend/plugins/domain/auth/plugin.go b/backend/plugins/domain/auth/plugin.go index b836bb1d..eff4d138 100644 --- a/backend/plugins/domain/auth/plugin.go +++ b/backend/plugins/domain/auth/plugin.go @@ -87,21 +87,8 @@ func (p *Plugin) Apply(ctx *core.Context) error { SetSessionConfig(cfg) } - // 0. Bind DBService & CacheService from Context - if db, err := core.Inject[contracts.DBService](ctx); err == nil && db != nil { - setDBService(db) - } else { - core.When[contracts.DBService](ctx, func(db contracts.DBService) { - setDBService(db) - }) - } - if cache, err := core.Inject[contracts.CacheService](ctx); err == nil && cache != nil { - setCacheService(cache) - } else { - core.When[contracts.CacheService](ctx, func(cache contracts.CacheService) { - setCacheService(cache) - }) - } + core.Bind[contracts.DBService](ctx, setDBService) + core.Bind[contracts.CacheService](ctx, setCacheService) ctx.OnDispose(func() error { setDBService(nil) setCacheService(nil) diff --git a/backend/plugins/domain/cap/plugin.go b/backend/plugins/domain/cap/plugin.go index 2c4703dc..148b02ed 100644 --- a/backend/plugins/domain/cap/plugin.go +++ b/backend/plugins/domain/cap/plugin.go @@ -59,14 +59,7 @@ func (p *Plugin) Apply(ctx *core.Context) error { SetSecret([]byte(cfg.SessionSecret)) } - // 0. Bind DBService from Context - if db, err := core.Inject[contracts.DBService](ctx); err == nil && db != nil { - setDBService(db) - } else { - core.When[contracts.DBService](ctx, func(db contracts.DBService) { - setDBService(db) - }) - } + core.Bind[contracts.DBService](ctx, setDBService) ctx.OnDispose(func() error { setDBService(nil) return nil diff --git a/backend/plugins/domain/domain_test.go b/backend/plugins/domain/domain_test.go index efa54eb4..32a031a1 100644 --- a/backend/plugins/domain/domain_test.go +++ b/backend/plugins/domain/domain_test.go @@ -236,7 +236,7 @@ func TestUserPlugin(t *testing.T) { assert.Len(t, list, 1) assert.Equal(t, "bob", list[0].Username) - // 9. Tasks & Schedules + // 9. Tasks taskDef, ok := ctx.Tasks().Get("user:send_email_code") require.True(t, ok) assert.Equal(t, 3, taskDef.Retry) diff --git a/backend/plugins/domain/message_gateway/errs/errs.go b/backend/plugins/domain/message_gateway/errs/errs.go index e6fa4280..6d4c1538 100644 --- a/backend/plugins/domain/message_gateway/errs/errs.go +++ b/backend/plugins/domain/message_gateway/errs/errs.go @@ -28,13 +28,15 @@ var ( // User-facing validation and error message constants. const ( - ErrNameRequired = "name is required" - ErrTypeInvalid = "type must be telegram or qq" - ErrTelegramTokenRequired = "telegram bot secret is required" //nolint:gosec // user-facing validation text - ErrQQCredentialsRequired = "qq app id and secret are required" //nolint:gosec // user-facing validation text - ErrChannelNotFound = "channel not found" - ErrChannelProbeFailed = "channel probe failed" - MaskedSecret = "********" + ErrNameRequired = "name is required" + ErrTypeInvalid = "type must be telegram or qq" + ErrTelegramTokenRequired = "telegram bot secret is required" //nolint:gosec // user-facing validation text + ErrQQCredentialsRequired = "qq app id and secret are required" //nolint:gosec // user-facing validation text + ErrChannelNotFound = "channel not found" + ErrChannelProbeFailed = "channel probe failed" + ErrBotDispatchTextRequired = "message text is required" + ErrBotChannelNotRegistered = "channel adapter is not registered" + MaskedSecret = "********" ErrLoginRequired = "login required" ErrInvalidBindingID = "invalid binding id" diff --git a/backend/plugins/domain/message_gateway/plugin.go b/backend/plugins/domain/message_gateway/plugin.go index 482a6c26..038e5aa3 100644 --- a/backend/plugins/domain/message_gateway/plugin.go +++ b/backend/plugins/domain/message_gateway/plugin.go @@ -10,6 +10,8 @@ import ( "Wavelet/core/extpoints" "Wavelet/pkg/ginutil" "Wavelet/pkg/util" + "Wavelet/plugins/domain/message_gateway/channels/qq" + "Wavelet/plugins/domain/message_gateway/channels/telegram" "Wavelet/plugins/domain/message_gateway/handler" "Wavelet/plugins/domain/message_gateway/model" "Wavelet/plugins/domain/message_gateway/repository" @@ -94,37 +96,13 @@ func (p *Plugin) Apply(ctx *core.Context) error { if err := ctx.Config().Bind("app", &cfg); err == nil && cfg.SessionSecret != "" { service.SetCredentialSecret(cfg.SessionSecret) } - // 0. Bind DBService, CacheService, TaskService, UserService - if db, err := core.Inject[contracts.DBService](ctx); err == nil && db != nil { - repository.SetDBService(db) - } else { - core.When[contracts.DBService](ctx, func(db contracts.DBService) { - repository.SetDBService(db) - }) - } - if cache, err := core.Inject[contracts.CacheService](ctx); err == nil && cache != nil { + core.Bind[contracts.DBService](ctx, repository.SetDBService) + core.Bind[contracts.CacheService](ctx, func(cache contracts.CacheService) { repository.SetCacheService(cache) service.SetCacheService(cache) - } else { - core.When[contracts.CacheService](ctx, func(cache contracts.CacheService) { - repository.SetCacheService(cache) - service.SetCacheService(cache) - }) - } - if taskSvc, err := core.Inject[contracts.TaskService](ctx); err == nil && taskSvc != nil { - service.SetTaskService(taskSvc) - } else { - core.When[contracts.TaskService](ctx, func(taskSvc contracts.TaskService) { - service.SetTaskService(taskSvc) - }) - } - if uSvc, err := core.Inject[contracts.UserService](ctx); err == nil && uSvc != nil { - service.SetUserService(uSvc) - } else { - core.When[contracts.UserService](ctx, func(uSvc contracts.UserService) { - service.SetUserService(uSvc) - }) - } + }) + core.Bind[contracts.TaskService](ctx, service.SetTaskService) + core.Bind[contracts.UserService](ctx, service.SetUserService) ctx.OnDispose(func() error { repository.SetDBService(nil) repository.SetCacheService(nil) @@ -159,6 +137,9 @@ func (p *Plugin) Apply(ctx *core.Context) error { // 4. Register Admin Push HTTP Routes handler.RegisterAdminPushRoutes(ctx.Router().Group("/api/v1/admin"), loginMW, adminMW) + service.Register(model.MessageChannelTypeTelegram, telegram.New) + service.Register(model.MessageChannelTypeQQ, qq.New) + const defaultTaskRetry = 3 pushHandler := &service.PushHandler{} @@ -179,15 +160,8 @@ func (p *Plugin) Apply(ctx *core.Context) error { return pushHandler.Execute(c, payload) }, extpoints.WithTaskMeta(service.SendNotificationMeta), extpoints.WithTaskRetry(defaultTaskRetry)) - ctx.Task().Register("message_gateway:dispatch_bot_msg", func(_ context.Context, _ []byte) error { - return nil - }, - extpoints.WithTaskType("dispatch_bot_msg"), - extpoints.WithTaskName("分发 Bot 消息"), - extpoints.WithTaskDescription("异步处理与分发 Bot 下行消息"), - extpoints.WithTaskCategory("messaging"), - extpoints.WithTaskQueue("default"), - ) + ctx.Task().Register(service.TaskDispatchBotMsg, &service.BotDispatchHandler{}, + extpoints.WithTaskMeta(service.BotDispatchMeta)) ctx.Task().Register("message_gateway:cleanup_pairing_codes", func(c context.Context, _ []byte) error { return repository.DeleteExpiredPairingCodes(c) diff --git a/backend/plugins/domain/message_gateway/repository/repository.go b/backend/plugins/domain/message_gateway/repository/repository.go index b6e5302a..52a78eab 100644 --- a/backend/plugins/domain/message_gateway/repository/repository.go +++ b/backend/plugins/domain/message_gateway/repository/repository.go @@ -132,6 +132,15 @@ func ListBindingsByUser(ctx context.Context, userID uint64) ([]model.MessageBind return rows, nil } +// ListBindingsByChannel lists bindings on one messaging channel. +func ListBindingsByChannel(ctx context.Context, channelID uint64) ([]model.MessageBinding, error) { + var rows []model.MessageBinding + if err := GetDB(ctx).Where("channel_id = ?", channelID).Order("id DESC").Find(&rows).Error; err != nil { + return nil, err + } + return rows, nil +} + // GetMessageBinding loads a binding by id. func GetMessageBinding(ctx context.Context, id uint64) (*model.MessageBinding, error) { var b model.MessageBinding diff --git a/backend/plugins/domain/message_gateway/service/admin.go b/backend/plugins/domain/message_gateway/service/admin.go index 50633dea..d785d3a9 100644 --- a/backend/plugins/domain/message_gateway/service/admin.go +++ b/backend/plugins/domain/message_gateway/service/admin.go @@ -27,14 +27,14 @@ func ListDefinitions() []model.Definition { { Type: model.MessageChannelTypeTelegram, Fields: []model.Field{ - {Key: "token", Type: "password", Required: true}, - {Key: "api_base", Type: "text", Required: false}, + {Key: "token", Type: model.TypePassword, Required: true}, + {Key: "api_base", Type: model.TypeText, Required: false}, }, }, { Type: model.MessageChannelTypeQQ, Fields: []model.Field{ - {Key: "app_id", Type: "text", Required: true}, + {Key: "app_id", Type: model.TypeText, Required: true}, {Key: "client_secret", Type: "password", Required: true}, }, }, diff --git a/backend/plugins/domain/message_gateway/service/dispatch.go b/backend/plugins/domain/message_gateway/service/dispatch.go new file mode 100644 index 00000000..9e0c8f78 --- /dev/null +++ b/backend/plugins/domain/message_gateway/service/dispatch.go @@ -0,0 +1,190 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package service + +import ( + "Wavelet/core/contracts" + "Wavelet/pkg/logger" + "Wavelet/plugins/domain/message_gateway/errs" + "Wavelet/plugins/domain/message_gateway/model" + "Wavelet/plugins/domain/message_gateway/repository" + "context" + "encoding/json" + "errors" + "fmt" + "strings" +) + +const ( + // TaskDispatchBotMsg is the queue pattern for bot downlink dispatch. + TaskDispatchBotMsg = "message_gateway:dispatch_bot_msg" + // TaskTypeDispatchBotMsg is the admin type identifier for bot downlink dispatch. + TaskTypeDispatchBotMsg = "dispatch_bot_msg" + + taskQueueDefault = "default" + taskParamTypeString = "string" + paramNameText = "text" +) + +// BotDispatchMeta describes the bot downlink dispatch task. +var BotDispatchMeta = contracts.TaskMetaDTO{ + Type: TaskTypeDispatchBotMsg, + AsynqTask: TaskDispatchBotMsg, + Name: "分发 Bot 消息", + DisplayName: "分发 Bot 消息", + Description: "向已绑定的平台账号异步下发 Bot 文本消息", + Category: "messaging", + Queue: taskQueueDefault, + Retryable: true, + Params: []contracts.TaskParamDTO{ + {Name: paramNameText, Label: "消息内容", Type: model.TypeText, Required: true, Placeholder: "要发送的文本", Description: "下发给绑定用户的文本"}, + {Name: "channel_id", Label: "频道 ID", Type: taskParamTypeString, Required: false, Placeholder: "留空表示全部启用频道", Description: "仅向指定频道的绑定发送"}, + {Name: "user_id", Label: "用户 ID", Type: taskParamTypeString, Required: false, Placeholder: "留空表示频道下全部绑定", Description: "仅向指定 Wavelet 用户的绑定发送"}, + }, +} + +type botDispatchPayload struct { + Text string `json:"text"` + ChannelID uint64 `json:"channel_id,string"` + UserID uint64 `json:"user_id,string"` +} + +// BotDispatchHandler sends a text message through enabled bot channels. +type BotDispatchHandler struct{} + +// ValidatePayload requires a non-empty message body. +func (h *BotDispatchHandler) ValidatePayload(payload []byte) ([]byte, error) { + p, err := parseBotDispatchPayload(payload) + if err != nil { + return nil, err + } + return json.Marshal(p) +} + +// Execute delivers the text to matching channel bindings. +func (h *BotDispatchHandler) Execute(ctx context.Context, payload []byte) (*contracts.TaskResultDTO, error) { + p, err := parseBotDispatchPayload(payload) + if err != nil { + return nil, err + } + + channels, err := repository.ListEnabledMessageChannels(ctx) + if err != nil { + return nil, err + } + if p.ChannelID != 0 { + filtered := channels[:0] + for i := range channels { + if channels[i].ID == p.ChannelID { + filtered = append(filtered, channels[i]) + } + } + channels = filtered + if len(channels) == 0 { + return nil, errors.New(errs.ErrChannelNotFound) + } + } + + sent := 0 + failed := 0 + for i := range channels { + n, ferr := dispatchOnChannel(ctx, &channels[i], p.UserID, p.Text) + sent += n + failed += ferr + } + msg := fmt.Sprintf("Bot 消息已尝试发送,成功 %d,失败 %d", sent, failed) + if svc := GetTaskService(ctx); svc != nil { + svc.AppendLog(ctx, "%s", msg) + } + if sent == 0 && failed > 0 { + return nil, errors.New(msg) + } + return &contracts.TaskResultDTO{Message: msg}, nil +} + +func parseBotDispatchPayload(payload []byte) (botDispatchPayload, error) { + var p botDispatchPayload + if len(payload) > 0 { + if err := json.Unmarshal(payload, &p); err != nil { + return p, fmt.Errorf("%s: %w", errs.ErrInvalidJSONFormat, err) + } + } + p.Text = strings.TrimSpace(p.Text) + if p.Text == "" { + return p, errors.New(errs.ErrBotDispatchTextRequired) + } + return p, nil +} + +func dispatchOnChannel(ctx context.Context, row *model.MessageChannel, userID uint64, text string) (sent, failed int) { + factory, ok := Lookup(row.Type) + if !ok { + logger.ErrorF(ctx, "bot dispatch: %s type=%s", errs.ErrBotChannelNotRegistered, row.Type) + return 0, 1 + } + cfg, err := channelConfigFromRow(row) + if err != nil { + logger.ErrorF(ctx, "bot dispatch: decode channel %d: %v", row.ID, err) + return 0, 1 + } + ch, err := factory(cfg, nil) + if err != nil { + logger.ErrorF(ctx, "bot dispatch: create adapter %d: %v", row.ID, err) + return 0, 1 + } + if err := ch.Connect(ctx); err != nil { + logger.ErrorF(ctx, "bot dispatch: connect channel %d: %v", row.ID, err) + return 0, 1 + } + defer func() { _ = ch.Disconnect(ctx) }() + + bindings, err := repository.ListBindingsByChannel(ctx, row.ID) + if err != nil { + logger.ErrorF(ctx, "bot dispatch: list bindings %d: %v", row.ID, err) + return 0, 1 + } + for i := range bindings { + if userID != 0 && bindings[i].UserID != userID { + continue + } + to := model.Recipient{ + ChatID: bindings[i].PlatformUserID, + PlatformUserID: bindings[i].PlatformUserID, + } + if err := ch.Send(ctx, to, model.OutboundMessage{Text: text}); err != nil { + logger.ErrorF(ctx, "bot dispatch: send channel=%d user=%d: %v", row.ID, bindings[i].UserID, err) + failed++ + continue + } + sent++ + } + return sent, failed +} + +func channelConfigFromRow(row *model.MessageChannel) (model.ChannelConfig, error) { + creds, err := DecryptCredentials(row.Credentials) + if err != nil { + return model.ChannelConfig{}, err + } + if creds == nil { + creds = map[string]string{} + } + if creds["bot_token"] == "" && creds["token"] != "" { + creds["bot_token"] = creds["token"] + } + if creds["app_secret"] == "" && creds["client_secret"] != "" { + creds["app_secret"] = creds["client_secret"] + } + extra := ParseExtra(row.Extra) + if extra["base_url"] == "" && creds["api_base"] != "" { + extra["base_url"] = creds["api_base"] + } + return model.ChannelConfig{ + ID: row.ID, + Type: row.Type, + Name: row.Name, + Credentials: creds, + Extra: extra, + }, nil +} diff --git a/backend/plugins/domain/message_gateway/service/dispatch_test.go b/backend/plugins/domain/message_gateway/service/dispatch_test.go new file mode 100644 index 00000000..9ef76d5a --- /dev/null +++ b/backend/plugins/domain/message_gateway/service/dispatch_test.go @@ -0,0 +1,47 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package service_test + +import ( + "Wavelet/plugins/domain/message_gateway/service" + "context" + "path/filepath" + "testing" + + "github.com/glebarez/sqlite" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "gorm.io/gorm" + + "Wavelet/plugins/domain/message_gateway/model" + "Wavelet/plugins/domain/message_gateway/repository" +) + +type dispatchTestDB struct{ db *gorm.DB } + +func (m *dispatchTestDB) GORM() *gorm.DB { return m.db } +func (m *dispatchTestDB) DB(ctx context.Context) *gorm.DB { return m.db.WithContext(ctx) } +func (m *dispatchTestDB) Named(_ string) *gorm.DB { return m.db } + +func TestBotDispatchValidatePayload(t *testing.T) { + h := &service.BotDispatchHandler{} + _, err := h.ValidatePayload([]byte(`{}`)) + require.Error(t, err) + _, err = h.ValidatePayload([]byte(`{"text":"hello"}`)) + require.NoError(t, err) +} + +func TestBotDispatchNoChannels(t *testing.T) { + testDB, err := gorm.Open(sqlite.Open(filepath.Join(t.TempDir(), "dispatch.db")), &gorm.Config{}) + require.NoError(t, err) + require.NoError(t, testDB.AutoMigrate(&model.MessageChannel{}, &model.MessageBinding{})) + repository.SetDBServiceForTest(&dispatchTestDB{db: testDB}) + t.Cleanup(func() { repository.SetDBServiceForTest(nil) }) + + h := &service.BotDispatchHandler{} + res, err := h.Execute(context.Background(), []byte(`{"text":"hello"}`)) + require.NoError(t, err) + require.NotNil(t, res) + assert.Contains(t, res.Message, "成功 0") +} diff --git a/backend/plugins/domain/message_gateway/service/push.go b/backend/plugins/domain/message_gateway/service/push.go index 0ea7081f..e7d6c646 100644 --- a/backend/plugins/domain/message_gateway/service/push.go +++ b/backend/plugins/domain/message_gateway/service/push.go @@ -52,10 +52,12 @@ func GetBuiltInEvents() []model.EventMetadata { // PushRegistryAdapter adapts contracts.PushRegistry onto the built-in event store. type PushRegistryAdapter struct{} +// RegisterBuiltInEvent records a built-in push event definition. func (PushRegistryAdapter) RegisterBuiltInEvent(meta contracts.PushEventMeta) { RegisterBuiltInEvent(eventMetadataFromContract(meta)) } +// SyncEvents persists registered built-in events into the database. func (PushRegistryAdapter) SyncEvents(ctx context.Context) error { return SyncEvents(ctx) } @@ -108,7 +110,7 @@ func ListPushEvents(ctx context.Context) ([]model.PushEvent, error) { // CreatePushEvent stores a push event configuration for a built-in event or task type. func CreatePushEvent(ctx context.Context, req model.CreatePushEventRequest) (model.PushEvent, error) { - eventKey, eventName, defaultTemplateBytes, err := GetEventInfo(req) + eventKey, eventName, defaultTemplateBytes, err := GetEventInfo(ctx, req) if err != nil { return model.PushEvent{}, err } @@ -650,10 +652,10 @@ func FindBuiltInEvent(key string) (model.EventMetadata, bool) { // GetEventInfo derives the event key, display name and default template for a // task-completion based event or a registered built-in event key. -func GetEventInfo(req model.CreatePushEventRequest) (string, string, []byte, error) { +func GetEventInfo(ctx context.Context, req model.CreatePushEventRequest) (string, string, []byte, error) { if req.TaskType != "" { taskName := req.TaskType - if taskSvc := GetTaskService(); taskSvc != nil { + if taskSvc := GetTaskService(ctx); taskSvc != nil { if meta, ok := taskSvc.GetTaskMeta(req.TaskType); ok { taskName = meta.DisplayName } @@ -694,7 +696,7 @@ func EnqueuePushTask(ctx context.Context, payload model.SendPayload) error { if err != nil { return err } - if taskSvc := GetTaskService(); taskSvc != nil { + if taskSvc := GetTaskService(ctx); taskSvc != nil { _, err = taskSvc.Dispatch(ctx, "send_notification", payloadBytes, "system") return err } @@ -749,13 +751,13 @@ var SendNotificationMeta = contracts.TaskMetaDTO{ Category: "push", SupportsTime: false, MaxRetry: 3, - Queue: "default", + Queue: taskQueueDefault, Retryable: true, Params: []contracts.TaskParamDTO{ { Name: "event_key", Label: "事件标识", - Type: "string", + Type: taskParamTypeString, Required: true, Placeholder: "admin_login", Description: "事件标识 (如 admin_login)", @@ -763,7 +765,7 @@ var SendNotificationMeta = contracts.TaskMetaDTO{ { Name: "target", Label: "目标接收者", - Type: "string", + Type: taskParamTypeString, Required: false, Description: "目标接收者", }, diff --git a/backend/plugins/domain/message_gateway/service/service.go b/backend/plugins/domain/message_gateway/service/service.go index dfd33db5..b298c5ec 100644 --- a/backend/plugins/domain/message_gateway/service/service.go +++ b/backend/plugins/domain/message_gateway/service/service.go @@ -251,10 +251,8 @@ func SetUserService(s contracts.UserService) { // GetCache resolves the cache service for the context. func GetCache(ctx context.Context) contracts.CacheService { - if c, ok := ctx.(*core.Context); ok && c != nil { - if s, err := core.Inject[contracts.CacheService](c); err == nil && s != nil { - return s - } + if s, err := core.InjectFrom[contracts.CacheService](ctx); err == nil && s != nil { + return s } cacheMu.RLock() s := cacheSvc @@ -263,7 +261,10 @@ func GetCache(ctx context.Context) contracts.CacheService { } // GetTaskService returns the task service. -func GetTaskService() contracts.TaskService { +func GetTaskService(ctx context.Context) contracts.TaskService { + if s, err := core.InjectFrom[contracts.TaskService](ctx); err == nil && s != nil { + return s + } taskMu.RLock() defer taskMu.RUnlock() return taskSvc @@ -271,10 +272,8 @@ func GetTaskService() contracts.TaskService { // GetUserService resolves the user service for the context. func GetUserService(ctx context.Context) contracts.UserService { - if c, ok := ctx.(*core.Context); ok && c != nil { - if s, err := core.Inject[contracts.UserService](c); err == nil && s != nil { - return s - } + if s, err := core.InjectFrom[contracts.UserService](ctx); err == nil && s != nil { + return s } userMu.RLock() s := userSvc diff --git a/backend/plugins/domain/risk_control/plugin.go b/backend/plugins/domain/risk_control/plugin.go index 3976fc3a..c879c381 100644 --- a/backend/plugins/domain/risk_control/plugin.go +++ b/backend/plugins/domain/risk_control/plugin.go @@ -94,14 +94,7 @@ func (p *Plugin) Apply(ctx *core.Context) error { SetAccessLogEnabled(chCfg.Enabled) logstore.SetDefaultDatabases(dbCfg.Enabled, chCfg.Enabled) - // 0. Bind DBService - if db, err := core.Inject[contracts.DBService](ctx); err == nil && db != nil { - logstore.SetDBService(db) - } else { - core.When[contracts.DBService](ctx, func(db contracts.DBService) { - logstore.SetDBService(db) - }) - } + core.Bind[contracts.DBService](ctx, logstore.SetDBService) ctx.OnDispose(func() error { logstore.SetDBService(nil) return nil diff --git a/backend/plugins/domain/upload/plugin.go b/backend/plugins/domain/upload/plugin.go index 5b4cc281..310eef97 100644 --- a/backend/plugins/domain/upload/plugin.go +++ b/backend/plugins/domain/upload/plugin.go @@ -12,7 +12,6 @@ import ( "Wavelet/plugins/domain/upload/handler" "Wavelet/plugins/domain/upload/shared" "Wavelet/plugins/domain/upload/task" - "context" "embed" "reflect" @@ -56,50 +55,11 @@ func (p *Plugin) Manifest() core.Manifest { // Apply registers upload routes, tasks, and settings into the Context. func (p *Plugin) Apply(ctx *core.Context) error { - // Bind DBService - if db, err := core.Inject[contracts.DBService](ctx); err == nil && db != nil { - shared.SetDBService(db) - } else { - core.When[contracts.DBService](ctx, func(db contracts.DBService) { - shared.SetDBService(db) - }) - } - - // Bind CacheService - if cache, err := core.Inject[contracts.CacheService](ctx); err == nil && cache != nil { - shared.SetCacheService(cache) - } else { - core.When[contracts.CacheService](ctx, func(cache contracts.CacheService) { - shared.SetCacheService(cache) - }) - } - - // Bind StorageService - if storage, err := core.Inject[contracts.StorageService](ctx); err == nil && storage != nil { - shared.SetStorageService(storage) - } else { - core.When[contracts.StorageService](ctx, func(storage contracts.StorageService) { - shared.SetStorageService(storage) - }) - } - - // Bind TaskService - if taskSvc, err := core.Inject[contracts.TaskService](ctx); err == nil && taskSvc != nil { - shared.SetTaskService(taskSvc) - } else { - core.When[contracts.TaskService](ctx, func(taskSvc contracts.TaskService) { - shared.SetTaskService(taskSvc) - }) - } - - // Bind AuthService - if authSvc, err := core.Inject[contracts.AuthService](ctx); err == nil && authSvc != nil { - shared.SetAuthService(authSvc) - } else { - core.When[contracts.AuthService](ctx, func(authSvc contracts.AuthService) { - shared.SetAuthService(authSvc) - }) - } + core.Bind[contracts.DBService](ctx, shared.SetDBService) + core.Bind[contracts.CacheService](ctx, shared.SetCacheService) + core.Bind[contracts.StorageService](ctx, shared.SetStorageService) + core.Bind[contracts.TaskService](ctx, shared.SetTaskService) + core.Bind[contracts.AuthService](ctx, shared.SetAuthService) ctx.OnDispose(func() error { shared.ResetServices() @@ -147,31 +107,10 @@ func (p *Plugin) Apply(ctx *core.Context) error { defaultSingleRetry = 1 ) - // 3. Register tasks. Handlers take raw payload bytes rather than a driver - // specific task type so they run under both the asynq and in-process workers. - cleanupHandler := &task.SystemCleanupHandler{} - ctx.Task().Register(task.SystemCleanupTask, func(c context.Context, payload []byte) error { - _, err := cleanupHandler.Execute(c, payload) - return err - }, extpoints.WithTaskMeta(task.SystemCleanupMeta), extpoints.WithTaskRetry(defaultCleanupRetry)) - - rebuildStatsHandler := &task.RebuildUploadStatsHandler{} - ctx.Task().Register(task.RebuildUploadStatsTask, func(c context.Context, payload []byte) error { - _, err := rebuildStatsHandler.Execute(c, payload) - return err - }, extpoints.WithTaskMeta(task.RebuildUploadStatsMeta), extpoints.WithTaskRetry(defaultStatsRetry)) - - migrationHandler := &task.MigrationHandler{} - ctx.Task().Register(task.StorageMigrationTask, func(c context.Context, payload []byte) error { - _, err := migrationHandler.Execute(c, payload) - return err - }, extpoints.WithTaskMeta(task.StorageMigrationMeta), extpoints.WithTaskRetry(defaultSingleRetry)) - - warmHandler := &task.WarmImageCacheHandler{} - ctx.Task().Register(task.WarmImageCacheTask, func(c context.Context, payload []byte) error { - _, err := warmHandler.Execute(c, payload) - return err - }, extpoints.WithTaskMeta(task.WarmImageCacheMeta), extpoints.WithTaskRetry(1)) + ctx.Task().Register(task.SystemCleanupTask, &task.SystemCleanupHandler{}, extpoints.WithTaskMeta(task.SystemCleanupMeta), extpoints.WithTaskRetry(defaultCleanupRetry)) + ctx.Task().Register(task.RebuildUploadStatsTask, &task.RebuildUploadStatsHandler{}, extpoints.WithTaskMeta(task.RebuildUploadStatsMeta), extpoints.WithTaskRetry(defaultStatsRetry)) + ctx.Task().Register(task.StorageMigrationTask, &task.MigrationHandler{}, extpoints.WithTaskMeta(task.StorageMigrationMeta), extpoints.WithTaskRetry(defaultSingleRetry)) + ctx.Task().Register(task.WarmImageCacheTask, &task.WarmImageCacheHandler{}, extpoints.WithTaskMeta(task.WarmImageCacheMeta), extpoints.WithTaskRetry(1)) // 4. Register Cron Schedule ctx.Schedule().RegisterCron("0 3 * * *", task.SystemCleanupTask, nil) diff --git a/backend/plugins/domain/upload/shared/context_services.go b/backend/plugins/domain/upload/shared/context_services.go index cb6c2415..7e9494cd 100644 --- a/backend/plugins/domain/upload/shared/context_services.go +++ b/backend/plugins/domain/upload/shared/context_services.go @@ -69,10 +69,8 @@ func ResetServices() { // GetDB resolves the GORM DB instance. func GetDB(ctx context.Context) *gorm.DB { - if c, ok := ctx.(*core.Context); ok && c != nil { - if s, err := core.Inject[contracts.DBService](c); err == nil && s != nil { - return s.DB(ctx) - } + if s, err := core.InjectFrom[contracts.DBService](ctx); err == nil && s != nil { + return s.DB(ctx) } svcMu.RLock() s := dbSvc @@ -85,10 +83,8 @@ func GetDB(ctx context.Context) *gorm.DB { // GetCache resolves the CacheService instance. func GetCache(ctx context.Context) contracts.CacheService { - if c, ok := ctx.(*core.Context); ok && c != nil { - if s, err := core.Inject[contracts.CacheService](c); err == nil && s != nil { - return s - } + if s, err := core.InjectFrom[contracts.CacheService](ctx); err == nil && s != nil { + return s } svcMu.RLock() s := cacheSvc @@ -98,10 +94,8 @@ func GetCache(ctx context.Context) contracts.CacheService { // GetStorage resolves the StorageService instance. func GetStorage(ctx context.Context) contracts.StorageService { - if c, ok := ctx.(*core.Context); ok && c != nil { - if s, err := core.Inject[contracts.StorageService](c); err == nil && s != nil { - return s - } + if s, err := core.InjectFrom[contracts.StorageService](ctx); err == nil && s != nil { + return s } svcMu.RLock() s := storageSvc @@ -110,7 +104,10 @@ func GetStorage(ctx context.Context) contracts.StorageService { } // GetTaskService resolves the TaskService instance. -func GetTaskService() contracts.TaskService { +func GetTaskService(ctx context.Context) contracts.TaskService { + if s, err := core.InjectFrom[contracts.TaskService](ctx); err == nil && s != nil { + return s + } svcMu.RLock() defer svcMu.RUnlock() return taskSvc @@ -118,10 +115,8 @@ func GetTaskService() contracts.TaskService { // GetAuthService resolves the AuthService instance. func GetAuthService(ctx context.Context) contracts.AuthService { - if c, ok := ctx.(*core.Context); ok && c != nil { - if s, err := core.Inject[contracts.AuthService](c); err == nil && s != nil { - return s - } + if s, err := core.InjectFrom[contracts.AuthService](ctx); err == nil && s != nil { + return s } svcMu.RLock() s := authSvc diff --git a/backend/plugins/domain/user/errs.go b/backend/plugins/domain/user/errs.go index df4102f2..9f132948 100644 --- a/backend/plugins/domain/user/errs.go +++ b/backend/plugins/domain/user/errs.go @@ -40,4 +40,12 @@ const ( //nolint:gosec // error message, not hardcoded credentials errServicePasswordTooShort = "密码长度至少为 8 位" errUniqueUsernameFailed = "failed to generate unique username" + errInvalidEmail = "邮箱地址无效" + errInvalidEmailCode = "验证码必须是 6 位数字" + errInvalidTaskPayload = "任务参数无效" + errMailSubjectRequired = "邮件主题不能为空" + errMailBodyRequired = "邮件内容不能为空" + errSMTPNotConfigured = "SMTP 未配置" + errEmailCacheUnavailable = "缓存服务不可用,无法保存验证码" + errSendEmailFailed = "邮件发送失败" ) diff --git a/backend/plugins/domain/user/handlers.go b/backend/plugins/domain/user/handlers.go index 8ddea90b..4df8b824 100644 --- a/backend/plugins/domain/user/handlers.go +++ b/backend/plugins/domain/user/handlers.go @@ -11,6 +11,7 @@ import ( "crypto/rand" "crypto/sha256" "encoding/hex" + "encoding/json" "net/http" "strconv" "sync" @@ -190,10 +191,36 @@ func Logout(c *gin.Context) { // @Tags user // @Accept json // @Produce json +// @Param request body user.sendEmailCodeRequest true "目标邮箱" // @Success 200 {object} response.Any "发送成功" // @Failure 400 {object} response.Any "参数错误" +// @Failure 500 {object} response.Any "发送失败" // @Router /api/v1/user/send-email-code [post] func SendEmailCode(c *gin.Context) { + var req sendEmailCodeRequest + if err := c.ShouldBindJSON(&req); err != nil { + response.AbortBadRequest(c, errInvalidParams) + return + } + payload, err := json.Marshal(sendEmailCodePayload{Email: req.Email}) + if err != nil { + response.AbortInternal(c, errSendEmailFailed) + return + } + ctx := c.Request.Context() + if taskSvc := getTaskService(ctx); taskSvc != nil { + if _, err := taskSvc.Dispatch(ctx, TaskTypeSendEmailCode, payload, "http"); err != nil { + logger.ErrorF(ctx, "dispatch send_email_code failed: %v", err) + response.AbortInternal(c, errSendEmailFailed) + return + } + c.JSON(http.StatusOK, response.OK(gin.H{"sent": true})) + return + } + if _, err := (&SendEmailCodeHandler{}).Execute(ctx, payload); err != nil { + response.AbortBadRequest(c, err.Error()) + return + } c.JSON(http.StatusOK, response.OK(gin.H{"sent": true})) } diff --git a/backend/plugins/domain/user/models.go b/backend/plugins/domain/user/models.go index b742f9a9..9e367c57 100644 --- a/backend/plugins/domain/user/models.go +++ b/backend/plugins/domain/user/models.go @@ -109,6 +109,10 @@ type registerRequest struct { Email string `json:"email"` } +type sendEmailCodeRequest struct { + Email string `json:"email" binding:"required"` +} + // changePasswordRequest 修改密码请求参数 type changePasswordRequest struct { OldPassword string `json:"old_password" binding:"required"` diff --git a/backend/plugins/domain/user/plugin.go b/backend/plugins/domain/user/plugin.go index a46c37b4..42e93383 100644 --- a/backend/plugins/domain/user/plugin.go +++ b/backend/plugins/domain/user/plugin.go @@ -9,7 +9,6 @@ import ( "Wavelet/core/contracts" "Wavelet/core/extpoints" "Wavelet/pkg/ginutil" - "context" "embed" "reflect" @@ -76,16 +75,13 @@ func (p *Plugin) Manifest() core.Manifest { // Apply registers user migrations, services, routes, tasks, schedules, and settings into the Context. func (p *Plugin) Apply(ctx *core.Context) error { - // 0. Bind DBService from Context - if db, err := core.Inject[contracts.DBService](ctx); err == nil && db != nil { - SetDBService(db) - } else { - core.When[contracts.DBService](ctx, func(db contracts.DBService) { - SetDBService(db) - }) - } + core.Bind[contracts.DBService](ctx, SetDBService) + core.Bind[contracts.CacheService](ctx, SetCacheService) + core.Bind[contracts.TaskService](ctx, SetTaskService) ctx.OnDispose(func() error { SetDBService(nil) + SetCacheService(nil) + SetTaskService(nil) return nil }) @@ -101,11 +97,8 @@ func (p *Plugin) Apply(ctx *core.Context) error { if mw, ok := authSvc.DisallowTokenAuthMiddleware().(gin.HandlerFunc); ok { noTokenMW = mw } - } else { - core.When[contracts.AuthService](ctx, func(svc contracts.AuthService) { - SetAuthService(svc) - }) } + core.Bind[contracts.AuthService](ctx, SetAuthService) ctx.OnDispose(func() error { SetAuthService(nil) return nil @@ -155,90 +148,12 @@ func (p *Plugin) Apply(ctx *core.Context) error { } } - const ( - defaultUserTaskRetry = 3 - paramTypeString = "string" - paramNameEmail = "email" - ) - - // 4. Register background tasks - ctx.Task().Register("user:send_email_code", func(_ context.Context, _ []byte) error { - return nil - }, - extpoints.WithTaskType("send_email_code"), - extpoints.WithTaskName("发送邮箱验证码"), - extpoints.WithTaskDescription("异步发送用户注册与验证邮箱验证码"), - extpoints.WithTaskCategory("user"), - extpoints.WithTaskRetry(defaultUserTaskRetry), - extpoints.WithTaskQueue("default"), - extpoints.WithTaskRetryable(true), - extpoints.WithTaskParams( - contracts.TaskParamDTO{ - Name: paramNameEmail, - Label: "目标邮箱", - Type: paramTypeString, - Required: true, - Placeholder: "user@example.com", - Description: "接收验证码的目标邮箱", - }, - contracts.TaskParamDTO{ - Name: "code", - Label: "验证码", - Type: paramTypeString, - Required: true, - Placeholder: "123456", - Description: "6 位数字验证码", - }, - ), - ) - - ctx.Task().Register("mail:send", func(_ context.Context, _ []byte) error { - return nil - }, - extpoints.WithTaskType("send_email"), - extpoints.WithTaskName("发送邮件"), - extpoints.WithTaskDescription("异步发送系统邮件"), - extpoints.WithTaskCategory("mail"), - extpoints.WithTaskRetry(defaultUserTaskRetry), - extpoints.WithTaskQueue("default"), - extpoints.WithTaskRetryable(true), - extpoints.WithTaskParams( - contracts.TaskParamDTO{ - Name: "to", - Label: "接收邮箱 (To)", - Type: paramTypeString, - Required: true, - Placeholder: "receiver@example.com", - Description: "接收邮件的目标邮箱地址", - }, - contracts.TaskParamDTO{ - Name: "subject", - Label: "邮件主题 (Subject)", - Type: paramTypeString, - Required: true, - Placeholder: "请输入邮件主题", - Description: "发送邮件的主题标题", - }, - contracts.TaskParamDTO{ - Name: "body", - Label: "邮件内容 (Body)", - Type: "text", - Required: true, - Placeholder: "请输入邮件内容(支持 HTML格式)", - Description: "发送邮件的内容主体", - }, - ), - ) - - ctx.Task().Register("user:cleanup_inactive", func(_ context.Context, _ []byte) error { - return nil - }, - extpoints.WithTaskType("cleanup_inactive_users"), - extpoints.WithTaskName("清理未激活用户"), - extpoints.WithTaskDescription("清理长期未激活的注册用户与临时凭据"), - extpoints.WithTaskCategory("user"), - extpoints.WithTaskQueue("default"), - ) + ctx.Task().Register(TaskSendEmailCode, &SendEmailCodeHandler{}, + extpoints.WithTaskMeta(SendEmailCodeMeta), extpoints.WithTaskRetry(defaultUserTaskRetry)) + ctx.Task().Register(TaskSendMail, &SendMailHandler{}, + extpoints.WithTaskMeta(SendMailMeta), extpoints.WithTaskRetry(defaultUserTaskRetry)) + ctx.Task().Register(TaskCleanupInactive, &CleanupInactiveHandler{}, + extpoints.WithTaskMeta(CleanupInactiveMeta)) // 5. Register Settings Schemas ctx.Settings().Register(extpoints.SettingSchema{ diff --git a/backend/plugins/domain/user/repository.go b/backend/plugins/domain/user/repository.go index 2bb02a9e..94445a9c 100644 --- a/backend/plugins/domain/user/repository.go +++ b/backend/plugins/domain/user/repository.go @@ -8,8 +8,10 @@ import ( "Wavelet/core/contracts" "Wavelet/pkg/util" "context" + "errors" "strings" "sync" + "time" "gorm.io/gorm" ) @@ -27,10 +29,8 @@ func SetDBService(s contracts.DBService) { } func getDB(ctx context.Context) *gorm.DB { - if c, ok := ctx.(*core.Context); ok && c != nil { - if s, err := core.Inject[contracts.DBService](c); err == nil && s != nil { - return s.DB(ctx) - } + if s, err := core.InjectFrom[contracts.DBService](ctx); err == nil && s != nil { + return s.DB(ctx) } dbMu.RLock() @@ -180,6 +180,26 @@ func DeleteUserWithRelations(ctx context.Context, id uint64) error { }) } +// ListInactiveNeverLoggedInUserIDs returns non-admin users created before cutoff +// who have never logged in. Seeded admin/system accounts are excluded. +func ListInactiveNeverLoggedInUserIDs(ctx context.Context, cutoff time.Time) ([]uint64, error) { + db := getDB(ctx) + if db == nil { + return nil, errors.New("database not available") + } + var ids []uint64 + unixEpoch := time.Unix(0, 0).UTC() + err := db.Model(&User{}). + Where("is_admin = ? AND username NOT IN ?", false, []string{"admin", "system"}). + Where("created_at < ?", cutoff). + Where("last_login_at IS NULL OR last_login_at < ?", unixEpoch). + Pluck("id", &ids).Error + if err != nil { + return nil, err + } + return ids, nil +} + // GetFirstAdminUser 获取第一个管理员用户 func GetFirstAdminUser(ctx context.Context) (*User, error) { var u User diff --git a/backend/plugins/domain/user/task.go b/backend/plugins/domain/user/task.go new file mode 100644 index 00000000..439079aa --- /dev/null +++ b/backend/plugins/domain/user/task.go @@ -0,0 +1,385 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package user + +import ( + "Wavelet/core" + "Wavelet/core/contracts" + "Wavelet/pkg/logger" + pkgmail "Wavelet/pkg/mail" + "context" + "crypto/rand" + "encoding/binary" + "encoding/json" + "errors" + "fmt" + "net/mail" + "strconv" + "strings" + "sync" + "time" + "unicode" +) + +const ( + // TaskSendEmailCode is the queue pattern for email verification codes. + TaskSendEmailCode = "user:send_email_code" + // TaskTypeSendEmailCode is the admin type identifier for email verification codes. + TaskTypeSendEmailCode = "send_email_code" + // TaskSendMail is the queue pattern for generic outbound mail. + TaskSendMail = "mail:send" + // TaskTypeSendMail is the admin type identifier for generic outbound mail. + TaskTypeSendMail = "send_email" + // TaskCleanupInactive is the queue pattern for inactive-user cleanup. + TaskCleanupInactive = "user:cleanup_inactive" + // TaskTypeCleanupInactive is the admin type identifier for inactive-user cleanup. + TaskTypeCleanupInactive = "cleanup_inactive_users" + + defaultUserTaskRetry = 3 + emailCodeTTL = 10 * time.Minute + emailCodeCacheKeyPrefix = "user:email_code:" + inactiveRetentionDays = 30 + hoursPerDay = 24 + inactiveRetention = inactiveRetentionDays * hoursPerDay * time.Hour + smtpConfigKeyHost = "smtp_host" + smtpConfigKeyPort = "smtp_port" + smtpConfigKeyUsername = "smtp_username" + smtpConfigKeyPassword = "smtp_password" + defaultSMTPPort = 587 + emailCodeLength = 6 + emailCodeModulo = 1000000 + taskQueueDefault = "default" + taskParamTypeString = "string" + taskParamTypeText = "text" + paramNameEmail = "email" +) + +var smtpConfigKeys = []string{ + smtpConfigKeyHost, smtpConfigKeyPort, smtpConfigKeyUsername, smtpConfigKeyPassword, +} + +var ( + cacheMu sync.RWMutex + cacheSvc contracts.CacheService + taskMu sync.RWMutex + taskSvc contracts.TaskService +) + +// SetCacheService sets the cache contract used to store email verification codes. +func SetCacheService(s contracts.CacheService) { + cacheMu.Lock() + defer cacheMu.Unlock() + cacheSvc = s +} + +// SetTaskService sets the task contract used by HTTP handlers to enqueue mail jobs. +func SetTaskService(s contracts.TaskService) { + taskMu.Lock() + defer taskMu.Unlock() + taskSvc = s +} + +func getCache(ctx context.Context) contracts.CacheService { + if s, err := core.InjectFrom[contracts.CacheService](ctx); err == nil && s != nil { + return s + } + cacheMu.RLock() + defer cacheMu.RUnlock() + return cacheSvc +} + +func getTaskService(ctx context.Context) contracts.TaskService { + if s, err := core.InjectFrom[contracts.TaskService](ctx); err == nil && s != nil { + return s + } + taskMu.RLock() + defer taskMu.RUnlock() + return taskSvc +} + +func appendTaskLog(ctx context.Context, format string, args ...any) { + if svc := getTaskService(ctx); svc != nil { + svc.AppendLog(ctx, format, args...) + } +} + +// SendEmailCodeMeta describes the email verification-code task. +var SendEmailCodeMeta = contracts.TaskMetaDTO{ + Type: TaskTypeSendEmailCode, + AsynqTask: TaskSendEmailCode, + Name: "发送邮箱验证码", + DisplayName: "发送邮箱验证码", + Description: "异步发送用户注册与验证邮箱验证码", + Category: "user", + MaxRetry: defaultUserTaskRetry, + Queue: taskQueueDefault, + Retryable: true, + Params: []contracts.TaskParamDTO{ + {Name: paramNameEmail, Label: "目标邮箱", Type: taskParamTypeString, Required: true, Placeholder: "user@example.com", Description: "接收验证码的目标邮箱"}, + {Name: "code", Label: "验证码", Type: taskParamTypeString, Required: false, Placeholder: "123456", Description: "6 位数字验证码,留空则自动生成"}, + }, +} + +// SendMailMeta describes the generic outbound-mail task. +var SendMailMeta = contracts.TaskMetaDTO{ + Type: TaskTypeSendMail, + AsynqTask: TaskSendMail, + Name: "发送邮件", + DisplayName: "发送邮件", + Description: "异步发送系统邮件", + Category: "mail", + MaxRetry: defaultUserTaskRetry, + Queue: taskQueueDefault, + Retryable: true, + Params: []contracts.TaskParamDTO{ + {Name: "to", Label: "接收邮箱 (To)", Type: taskParamTypeString, Required: true, Placeholder: "receiver@example.com", Description: "接收邮件的目标邮箱地址"}, + {Name: "subject", Label: "邮件主题 (Subject)", Type: taskParamTypeString, Required: true, Placeholder: "请输入邮件主题", Description: "发送邮件的主题标题"}, + {Name: "body", Label: "邮件内容 (Body)", Type: taskParamTypeText, Required: true, Placeholder: "请输入邮件内容(支持 HTML格式)", Description: "发送邮件的内容主体"}, + }, +} + +// CleanupInactiveMeta describes the inactive-user cleanup task. +var CleanupInactiveMeta = contracts.TaskMetaDTO{ + Type: TaskTypeCleanupInactive, + AsynqTask: TaskCleanupInactive, + Name: "清理未激活用户", + DisplayName: "清理未激活用户", + Description: "清理长期未登录的注册用户及其访问令牌", + Category: "user", + Queue: taskQueueDefault, + Retryable: true, +} + +type sendEmailCodePayload struct { + Email string `json:"email"` + Code string `json:"code"` +} + +type sendMailPayload struct { + To string `json:"to"` + Subject string `json:"subject"` + Body string `json:"body"` +} + +// SendEmailCodeHandler sends a 6-digit email verification code and caches it. +type SendEmailCodeHandler struct{} + +// ValidatePayload checks the destination address and optional code. +func (h *SendEmailCodeHandler) ValidatePayload(payload []byte) ([]byte, error) { + p, err := parseSendEmailCodePayload(payload) + if err != nil { + return nil, err + } + return json.Marshal(p) +} + +// Execute generates (if needed), caches, and emails the verification code. +func (h *SendEmailCodeHandler) Execute(ctx context.Context, payload []byte) (*contracts.TaskResultDTO, error) { + p, err := parseSendEmailCodePayload(payload) + if err != nil { + return nil, err + } + if p.Code == "" { + p.Code, err = generateEmailCode() + if err != nil { + return nil, err + } + } + + cache := getCache(ctx) + if cache == nil { + return nil, errors.New(errEmailCacheUnavailable) + } + if err := cache.Set(ctx, emailCodeCacheKey(p.Email), p.Code, emailCodeTTL); err != nil { + return nil, fmt.Errorf("store email code: %w", err) + } + + cfg, err := loadSMTPConfig(ctx) + if err != nil { + return nil, err + } + subject := "邮箱验证码" + body := fmt.Sprintf("

您的验证码是 %s,%d 分钟内有效。

", p.Code, int(emailCodeTTL.Minutes())) + appendTaskLog(ctx, "发送邮箱验证码到 %s", maskEmail(p.Email)) + if err := pkgmail.SendMail(ctx, cfg, p.Email, subject, body); err != nil { + logger.ErrorF(ctx, "send email code failed: %v", err) + return nil, errors.New(errSendEmailFailed) + } + return &contracts.TaskResultDTO{Message: fmt.Sprintf("验证码已发送至 %s", maskEmail(p.Email))}, nil +} + +// SendMailHandler sends a generic HTML email through the configured SMTP server. +type SendMailHandler struct{} + +// ValidatePayload checks to/subject/body. +func (h *SendMailHandler) ValidatePayload(payload []byte) ([]byte, error) { + p, err := parseSendMailPayload(payload) + if err != nil { + return nil, err + } + return json.Marshal(p) +} + +// Execute sends the mail. +func (h *SendMailHandler) Execute(ctx context.Context, payload []byte) (*contracts.TaskResultDTO, error) { + p, err := parseSendMailPayload(payload) + if err != nil { + return nil, err + } + cfg, err := loadSMTPConfig(ctx) + if err != nil { + return nil, err + } + appendTaskLog(ctx, "发送邮件到 %s,主题: %s", maskEmail(p.To), p.Subject) + if err := pkgmail.SendMail(ctx, cfg, p.To, p.Subject, p.Body); err != nil { + logger.ErrorF(ctx, "send mail failed: %v", err) + return nil, errors.New(errSendEmailFailed) + } + return &contracts.TaskResultDTO{Message: fmt.Sprintf("邮件已发送至 %s", maskEmail(p.To))}, nil +} + +// CleanupInactiveHandler deletes users who registered long ago and never logged in. +type CleanupInactiveHandler struct{} + +// Execute removes stale never-logged-in non-admin users and their access tokens. +func (h *CleanupInactiveHandler) Execute(ctx context.Context, _ []byte) (*contracts.TaskResultDTO, error) { + cutoff := time.Now().Add(-inactiveRetention) + ids, err := ListInactiveNeverLoggedInUserIDs(ctx, cutoff) + if err != nil { + return nil, err + } + appendTaskLog(ctx, "扫描到 %d 个超过 %d 天未登录的注册用户", len(ids), int(inactiveRetention.Hours()/float64(hoursPerDay))) + deleted := 0 + for _, id := range ids { + if err := DeleteUserWithRelations(ctx, id); err != nil { + logger.ErrorF(ctx, "cleanup inactive user %d failed: %v", id, err) + continue + } + deleted++ + } + msg := fmt.Sprintf("已清理 %d 个长期未登录用户及其访问令牌", deleted) + appendTaskLog(ctx, "%s", msg) + return &contracts.TaskResultDTO{Message: msg}, nil +} + +func parseSendEmailCodePayload(payload []byte) (sendEmailCodePayload, error) { + var p sendEmailCodePayload + if len(payload) > 0 { + if err := json.Unmarshal(payload, &p); err != nil { + return p, errors.New(errInvalidTaskPayload) + } + } + p.Email = normalizeEmail(p.Email) + if err := validateEmail(p.Email); err != nil { + return p, err + } + p.Code = strings.TrimSpace(p.Code) + if p.Code != "" && !isSixDigitCode(p.Code) { + return p, errors.New(errInvalidEmailCode) + } + return p, nil +} + +func parseSendMailPayload(payload []byte) (sendMailPayload, error) { + var p sendMailPayload + if err := json.Unmarshal(payload, &p); err != nil { + return p, errors.New(errInvalidTaskPayload) + } + p.To = normalizeEmail(p.To) + p.Subject = strings.TrimSpace(p.Subject) + if err := validateEmail(p.To); err != nil { + return p, err + } + if p.Subject == "" { + return p, errors.New(errMailSubjectRequired) + } + if strings.TrimSpace(p.Body) == "" { + return p, errors.New(errMailBodyRequired) + } + return p, nil +} + +func loadSMTPConfig(ctx context.Context) (pkgmail.Config, error) { + db := getDB(ctx) + if db == nil { + return pkgmail.Config{}, errors.New(errSMTPNotConfigured) + } + var rows []struct { + Key string + Value string + } + if err := db.Table("w_system_configs"). + Select("key", "value"). + Where("key IN ?", smtpConfigKeys). + Find(&rows).Error; err != nil { + return pkgmail.Config{}, fmt.Errorf("read smtp config: %w", err) + } + cfg := pkgmail.Config{Port: defaultSMTPPort} + for _, row := range rows { + switch row.Key { + case smtpConfigKeyHost: + cfg.Host = strings.TrimSpace(row.Value) + case smtpConfigKeyPort: + if n, err := strconv.Atoi(strings.TrimSpace(row.Value)); err == nil && n > 0 { + cfg.Port = n + } + case smtpConfigKeyUsername: + cfg.Username = strings.TrimSpace(row.Value) + case smtpConfigKeyPassword: + cfg.Password = row.Value + } + } + if cfg.Host == "" || cfg.Username == "" { + return pkgmail.Config{}, errors.New(errSMTPNotConfigured) + } + return cfg, nil +} + +func generateEmailCode() (string, error) { + var buf [4]byte + if _, err := rand.Read(buf[:]); err != nil { + return "", err + } + n := binary.BigEndian.Uint32(buf[:]) % emailCodeModulo + return fmt.Sprintf("%06d", n), nil +} + +func emailCodeCacheKey(email string) string { + return emailCodeCacheKeyPrefix + normalizeEmail(email) +} + +func normalizeEmail(email string) string { + return strings.ToLower(strings.TrimSpace(email)) +} + +func validateEmail(email string) error { + if email == "" { + return errors.New(errEmailEmpty) + } + addr, err := mail.ParseAddress(email) + if err != nil || !strings.EqualFold(addr.Address, email) { + return errors.New(errInvalidEmail) + } + return nil +} + +func isSixDigitCode(code string) bool { + if len(code) != emailCodeLength { + return false + } + for _, r := range code { + if !unicode.IsDigit(r) { + return false + } + } + return true +} + +func maskEmail(email string) string { + at := strings.IndexByte(email, '@') + if at <= 1 { + return "***" + } + return email[:1] + "***" + email[at:] +} diff --git a/backend/plugins/domain/user/task_test.go b/backend/plugins/domain/user/task_test.go new file mode 100644 index 00000000..8c4ccb2b --- /dev/null +++ b/backend/plugins/domain/user/task_test.go @@ -0,0 +1,109 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package user_test + +import ( + "Wavelet/core/contracts" + "Wavelet/pkg/idgen" + "Wavelet/plugins/domain/user" + "context" + "path/filepath" + "testing" + "time" + + "github.com/glebarez/sqlite" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "gorm.io/gorm" +) + +type taskTestDB struct{ db *gorm.DB } + +func (m *taskTestDB) GORM() *gorm.DB { return m.db } +func (m *taskTestDB) DB(ctx context.Context) *gorm.DB { return m.db.WithContext(ctx) } +func (m *taskTestDB) Named(_ string) *gorm.DB { return m.db } + +type sysConfigRow struct { + Key string `gorm:"primaryKey;size:64"` + Value string `gorm:"type:text"` +} + +func (sysConfigRow) TableName() string { return "w_system_configs" } + +func setupUserTaskDB(t *testing.T) *gorm.DB { + t.Helper() + _ = idgen.Init(1) + testDB, err := gorm.Open(sqlite.Open(filepath.Join(t.TempDir(), "user_task.db")), &gorm.Config{}) + require.NoError(t, err) + require.NoError(t, testDB.AutoMigrate(&user.User{}, &user.AccessToken{}, &sysConfigRow{})) + user.SetDBService(&taskTestDB{db: testDB}) + t.Cleanup(func() { user.SetDBService(nil) }) + return testDB +} + +func TestSendEmailCodeValidatePayload(t *testing.T) { + h := &user.SendEmailCodeHandler{} + _, err := h.ValidatePayload([]byte(`{"email":"not-an-email"}`)) + require.Error(t, err) + + out, err := h.ValidatePayload([]byte(`{"email":"User@Example.com"}`)) + require.NoError(t, err) + assert.Contains(t, string(out), `"user@example.com"`) +} + +func TestSendMailValidatePayload(t *testing.T) { + h := &user.SendMailHandler{} + _, err := h.ValidatePayload([]byte(`{"to":"a@b.com","subject":"","body":"x"}`)) + require.Error(t, err) + _, err = h.ValidatePayload([]byte(`{"to":"a@b.com","subject":"Hi","body":"

ok

"}`)) + require.NoError(t, err) +} + +func TestSendMailRequiresSMTP(t *testing.T) { + setupUserTaskDB(t) + h := &user.SendMailHandler{} + _, err := h.Execute(context.Background(), []byte(`{"to":"a@b.com","subject":"Hi","body":"

ok

"}`)) + require.Error(t, err) + assert.Contains(t, err.Error(), "SMTP") +} + +func TestCleanupInactiveNeverLoggedInUsers(t *testing.T) { + db := setupUserTaskDB(t) + old := time.Now().Add(-40 * 24 * time.Hour) + stale := user.User{ID: 42, Username: "stale", Password: "x", IsActive: true, CreatedAt: old} + require.NoError(t, db.Create(&stale).Error) + require.NoError(t, db.Model(&stale).Updates(map[string]any{ + "created_at": old, + "last_login_at": time.Time{}, + }).Error) + + fresh := user.User{ID: 43, Username: "fresh", Password: "x", IsActive: true, LastLoginAt: time.Now()} + require.NoError(t, db.Create(&fresh).Error) + + admin := user.User{ID: 1, Username: "admin", Password: "x", IsAdmin: true, CreatedAt: old} + require.NoError(t, db.Create(&admin).Error) + require.NoError(t, db.Model(&admin).Updates(map[string]any{ + "created_at": old, + "last_login_at": time.Time{}, + }).Error) + + h := &user.CleanupInactiveHandler{} + res, err := h.Execute(context.Background(), nil) + require.NoError(t, err) + require.NotNil(t, res) + assert.Contains(t, res.Message, "1") + + _, err = user.GetUserByID(context.Background(), 42) + assert.Error(t, err) + _, err = user.GetUserByID(context.Background(), 43) + require.NoError(t, err) + _, err = user.GetUserByID(context.Background(), 1) + require.NoError(t, err) +} + +func TestSendEmailCodeMetaExported(t *testing.T) { + assert.Equal(t, "send_email_code", user.SendEmailCodeMeta.Type) + assert.Equal(t, "user:send_email_code", user.SendEmailCodeMeta.AsynqTask) + _ = contracts.TaskHandler(&user.SendEmailCodeHandler{}) +} diff --git a/backend/plugins/drivers/driver_asynq_cron/plugin.go b/backend/plugins/drivers/driver_asynq_cron/plugin.go index d822d894..5439583c 100644 --- a/backend/plugins/drivers/driver_asynq_cron/plugin.go +++ b/backend/plugins/drivers/driver_asynq_cron/plugin.go @@ -120,23 +120,8 @@ func (p *Plugin) Apply(ctx *core.Context) error { } p.mu.Unlock() - // Bind DBService - if db, err := core.Inject[contracts.DBService](ctx); err == nil && db != nil { - setDBService(db) - } else { - core.When[contracts.DBService](ctx, func(db contracts.DBService) { - setDBService(db) - }) - } - - // Bind TaskService - if taskSvc, err := core.Inject[contracts.TaskService](ctx); err == nil && taskSvc != nil { - setTaskService(taskSvc) - } else { - core.When[contracts.TaskService](ctx, func(taskSvc contracts.TaskService) { - setTaskService(taskSvc) - }) - } + core.Bind[contracts.DBService](ctx, setDBService) + core.Bind[contracts.TaskService](ctx, setTaskService) ctx.OnDispose(func() error { setDBService(nil) diff --git a/backend/plugins/drivers/driver_asynq_worker/db_helper.go b/backend/plugins/drivers/driver_asynq_worker/db_helper.go index 025fbca1..a212de18 100644 --- a/backend/plugins/drivers/driver_asynq_worker/db_helper.go +++ b/backend/plugins/drivers/driver_asynq_worker/db_helper.go @@ -34,10 +34,8 @@ func SetRedisClient(c redis.UniversalClient) { } func getDB(ctx context.Context) *gorm.DB { - if c, ok := ctx.(*core.Context); ok && c != nil { - if s, err := core.Inject[contracts.DBService](c); err == nil && s != nil { - return s.DB(ctx) - } + if s, err := core.InjectFrom[contracts.DBService](ctx); err == nil && s != nil { + return s.DB(ctx) } dbMu.RLock() s := dbSvc diff --git a/backend/plugins/drivers/driver_asynq_worker/plugin.go b/backend/plugins/drivers/driver_asynq_worker/plugin.go index d16b1a6c..fee807f8 100644 --- a/backend/plugins/drivers/driver_asynq_worker/plugin.go +++ b/backend/plugins/drivers/driver_asynq_worker/plugin.go @@ -8,6 +8,7 @@ import ( "Wavelet/core" "Wavelet/core/contracts" "context" + "encoding/json" "errors" "fmt" "sync" @@ -144,14 +145,7 @@ func (p *Plugin) Apply(ctx *core.Context) error { ResetAsynqClient() p.mu.Unlock() - // 0. Bind DBService - if db, err := core.Inject[contracts.DBService](ctx); err == nil && db != nil { - setDBService(db) - } else { - core.When[contracts.DBService](ctx, func(db contracts.DBService) { - setDBService(db) - }) - } + core.Bind[contracts.DBService](ctx, setDBService) ctx.OnDispose(func() error { setDBService(nil) return nil @@ -191,12 +185,15 @@ func (p *Plugin) Start(_ context.Context) error { mux := asynq.NewServeMux() if p.coreCtx != nil && p.coreCtx.Tasks() != nil { + appCtx := p.coreCtx.Root() for _, td := range p.coreCtx.Tasks().Tasks() { handler, err := toAsynqHandler(td.Pattern, td.Handler) if err != nil { return fmt.Errorf("driver_asynq_worker: invalid handler for task pattern %q: %w", td.Pattern, err) } - mux.Handle(td.Pattern, handler) + mux.Handle(td.Pattern, asynq.HandlerFunc(func(c context.Context, t *asynq.Task) error { + return handler.ProcessTask(core.WithAppContext(c, appCtx), t) + })) } } @@ -285,6 +282,14 @@ func toAsynqHandler(pattern string, h any) (asynq.Handler, error) { RegisterHandler(pattern, th) return asynq.HandlerFunc(ProcessTask), nil } + if th, ok := h.(contracts.TaskHandler); ok { + RegisterHandler(pattern, contractTaskAdapter{inner: th}) + return asynq.HandlerFunc(ProcessTask), nil + } + if fn, ok := h.(func(context.Context, []byte) (*contracts.TaskResultDTO, error)); ok { + RegisterHandler(pattern, contractFuncAdapter{fn: fn}) + return asynq.HandlerFunc(ProcessTask), nil + } inner, err := toRawAsynqHandler(h) if err != nil { @@ -294,6 +299,58 @@ func toAsynqHandler(pattern string, h any) (asynq.Handler, error) { return asynq.HandlerFunc(ProcessTask), nil } +type contractTaskAdapter struct { + inner contracts.TaskHandler +} + +func (a contractTaskAdapter) Execute(ctx context.Context, payload []byte) (*TaskResult, error) { + res, err := a.inner.Execute(ctx, payload) + if err != nil { + return nil, err + } + return dtoToTaskResult(res), nil +} + +func (a contractTaskAdapter) ValidatePayload(payload []byte) ([]byte, error) { + if v, ok := a.inner.(PayloadValidator); ok { + return v.ValidatePayload(payload) + } + return payload, nil +} + +type contractFuncAdapter struct { + fn func(context.Context, []byte) (*contracts.TaskResultDTO, error) +} + +func (a contractFuncAdapter) Execute(ctx context.Context, payload []byte) (*TaskResult, error) { + res, err := a.fn(ctx, payload) + if err != nil { + return nil, err + } + return dtoToTaskResult(res), nil +} + +func dtoToTaskResult(res *contracts.TaskResultDTO) *TaskResult { + if res == nil { + return &TaskResult{Message: "ok"} + } + out := &TaskResult{Message: res.Message} + if res.Detail == nil { + return out + } + if s, ok := res.Detail.(string); ok { + out.Detail = s + return out + } + b, err := json.Marshal(res.Detail) + if err != nil { + out.Detail = fmt.Sprint(res.Detail) + return out + } + out.Detail = string(b) + return out +} + func toRawAsynqHandler(h any) (asynq.Handler, error) { switch fn := h.(type) { case asynq.HandlerFunc: diff --git a/backend/plugins/drivers/driver_http/plugin.go b/backend/plugins/drivers/driver_http/plugin.go index 72ebaa34..fd3fb5dc 100644 --- a/backend/plugins/drivers/driver_http/plugin.go +++ b/backend/plugins/drivers/driver_http/plugin.go @@ -121,27 +121,13 @@ func (p *Plugin) Apply(ctx *core.Context) error { } p.mu.Unlock() - // Bind DBService from Context - if db, err := core.Inject[contracts.DBService](ctx); err == nil && db != nil { - setDBService(db) - } else { - core.When[contracts.DBService](ctx, func(db contracts.DBService) { - setDBService(db) - }) - } + core.Bind[contracts.DBService](ctx, setDBService) ctx.OnDispose(func() error { setDBService(nil) return nil }) - // Bind CacheService from Context - if cache, err := core.Inject[contracts.CacheService](ctx); err == nil && cache != nil { - setCacheService(cache) - } else { - core.When[contracts.CacheService](ctx, func(cache contracts.CacheService) { - setCacheService(cache) - }) - } + core.Bind[contracts.CacheService](ctx, setCacheService) ctx.OnDispose(func() error { setCacheService(nil) return nil @@ -184,30 +170,8 @@ func (p *Plugin) Start(ctx context.Context) error { } } - // Mount routes collected in Context RouterExtension - if p.coreCtx != nil && p.coreCtx.Router() != nil { - SetWhitelist(p.coreCtx.Router().Whitelist()) - for _, rd := range p.coreCtx.Router().Routes() { - allHandlers := make([]gin.HandlerFunc, 0, len(rd.Middlewares)+len(rd.Handlers)) - - for _, m := range rd.Middlewares { - gh, err := toGinHandler(m) - if err != nil { - return fmt.Errorf("driver_http: invalid middleware for route %s %s: %w", rd.Method, rd.Path, err) - } - allHandlers = append(allHandlers, gh) - } - - for _, h := range rd.Handlers { - gh, err := toGinHandler(h) - if err != nil { - return fmt.Errorf("driver_http: invalid handler for route %s %s: %w", rd.Method, rd.Path, err) - } - allHandlers = append(allHandlers, gh) - } - - p.engine.Handle(rd.Method, rd.Path, allHandlers...) - } + if err := p.mountContextRoutes(ctx); err != nil { + return err } // Mount Swagger in non-production environments @@ -282,6 +246,41 @@ func (p *Plugin) Stop(ctx context.Context) error { return err } +func (p *Plugin) mountContextRoutes(ctx context.Context) error { + if p.coreCtx == nil || p.coreCtx.Router() == nil || p.engine == nil { + return nil + } + p.engine.Use(appContextMiddleware(ctx, p.coreCtx.Root())) + SetWhitelist(p.coreCtx.Router().Whitelist()) + for _, rd := range p.coreCtx.Router().Routes() { + allHandlers := make([]gin.HandlerFunc, 0, len(rd.Middlewares)+len(rd.Handlers)) + for _, m := range rd.Middlewares { + gh, err := toGinHandler(m) + if err != nil { + return fmt.Errorf("driver_http: invalid middleware for route %s %s: %w", rd.Method, rd.Path, err) + } + allHandlers = append(allHandlers, gh) + } + for _, h := range rd.Handlers { + gh, err := toGinHandler(h) + if err != nil { + return fmt.Errorf("driver_http: invalid handler for route %s %s: %w", rd.Method, rd.Path, err) + } + allHandlers = append(allHandlers, gh) + } + p.engine.Handle(rd.Method, rd.Path, allHandlers...) + } + return nil +} + +//nolint:contextcheck // middleware must wrap the gin request context, not Start's ctx +func appContextMiddleware(_ context.Context, appCtx *core.Context) gin.HandlerFunc { + return func(c *gin.Context) { + c.Request = c.Request.WithContext(core.WithAppContext(c.Request.Context(), appCtx)) + c.Next() + } +} + // Addr returns the current listening address (or configured address if not yet started). func (p *Plugin) Addr() string { p.mu.RLock() diff --git a/backend/plugins/drivers/driver_inproc_cron/plugin.go b/backend/plugins/drivers/driver_inproc_cron/plugin.go index 49fbc3fe..40caf4d9 100644 --- a/backend/plugins/drivers/driver_inproc_cron/plugin.go +++ b/backend/plugins/drivers/driver_inproc_cron/plugin.go @@ -83,7 +83,11 @@ func (p *Plugin) Start(ctx context.Context) error { p.scheduler = newInprocScheduler(p.coreCtx.Schedules(), p.coreCtx.Tasks(), taskSvc) } - return p.scheduler.Start(ctx) + runCtx := ctx + if p.coreCtx != nil { + runCtx = core.WithAppContext(ctx, p.coreCtx.Root()) + } + return p.scheduler.Start(runCtx) } // Stop terminates the in-process cron scheduler. diff --git a/backend/plugins/drivers/driver_inproc_worker/db_helper.go b/backend/plugins/drivers/driver_inproc_worker/db_helper.go index d564c491..32077db9 100644 --- a/backend/plugins/drivers/driver_inproc_worker/db_helper.go +++ b/backend/plugins/drivers/driver_inproc_worker/db_helper.go @@ -24,10 +24,8 @@ func setDBService(s contracts.DBService) { } func getDB(ctx context.Context) *gorm.DB { - if c, ok := ctx.(*core.Context); ok && c != nil { - if s, err := core.Inject[contracts.DBService](c); err == nil && s != nil { - return s.DB(ctx) - } + if s, err := core.InjectFrom[contracts.DBService](ctx); err == nil && s != nil { + return s.DB(ctx) } dbMu.RLock() s := dbSvc diff --git a/backend/plugins/drivers/driver_inproc_worker/executor.go b/backend/plugins/drivers/driver_inproc_worker/executor.go index fe211612..d22b2363 100644 --- a/backend/plugins/drivers/driver_inproc_worker/executor.go +++ b/backend/plugins/drivers/driver_inproc_worker/executor.go @@ -4,11 +4,14 @@ package driver_inproc_worker import ( + "Wavelet/core" + "Wavelet/core/contracts" "Wavelet/core/extpoints" "Wavelet/pkg/idgen" "Wavelet/pkg/logger" "Wavelet/pkg/util" "context" + "encoding/json" "errors" "fmt" "sync" @@ -40,6 +43,7 @@ type InprocQueue struct { // baseCtx is the app-lifetime context captured at Start; task handlers // derive their timeouts from it so shutdown cancellation propagates. baseCtx context.Context + appCtx *core.Context } // NewInprocQueue creates a new InprocQueue with a given concurrency and queue capacity. @@ -182,10 +186,14 @@ func (q *InprocQueue) executeTask(ctx context.Context, msg TaskMessage) { taskCtx, cancel := context.WithTimeout(ctx, timeout) defer cancel() + if q.appCtx != nil { + taskCtx = core.WithAppContext(taskCtx, q.appCtx) + ctx = core.WithAppContext(ctx, q.appCtx) + } q.markRunning(ctx, msg) start := time.Now() - err := invokeHandler(taskCtx, td.Handler, msg.Payload) + result, err := invokeHandler(taskCtx, td.Handler, msg.Payload) duration := time.Since(start) if err != nil { @@ -210,25 +218,29 @@ func (q *InprocQueue) executeTask(ctx context.Context, msg TaskMessage) { } return } - q.succeedExecution(ctx, msg, duration) + q.succeedExecution(ctx, msg, duration, result) } -func invokeHandler(ctx context.Context, handler any, payload []byte) error { +func invokeHandler(ctx context.Context, handler any, payload []byte) (*contracts.TaskResultDTO, error) { if handler == nil { - return errors.New("nil task handler") + return nil, errors.New("nil task handler") } switch fn := handler.(type) { - case func(context.Context, []byte) error: + case contracts.TaskHandler: + return fn.Execute(ctx, payload) + case func(context.Context, []byte) (*contracts.TaskResultDTO, error): return fn(ctx, payload) + case func(context.Context, []byte) error: + return nil, fn(ctx, payload) case func(context.Context) error: - return fn(ctx) + return nil, fn(ctx) case func([]byte) error: - return fn(payload) + return nil, fn(payload) case func() error: - return fn() + return nil, fn() default: - return fmt.Errorf("unsupported handler type: %T", handler) + return nil, fmt.Errorf("unsupported handler type: %T", handler) } } @@ -280,16 +292,27 @@ func (q *InprocQueue) markRunning(ctx context.Context, msg TaskMessage) { q.appendExecutionLog(ctx, msg.ID, fmt.Sprintf("[系统] 开始执行异步任务 [类型: %s]", msg.TaskType)) } -func (q *InprocQueue) succeedExecution(ctx context.Context, msg TaskMessage, duration time.Duration) { +func (q *InprocQueue) succeedExecution(ctx context.Context, msg TaskMessage, duration time.Duration, result *contracts.TaskResultDTO) { db := getDB(ctx) if db == nil { return } now := time.Now() + resultText := "ok" + if result != nil { + resultText = result.Message + if result.Detail != nil { + if s, ok := result.Detail.(string); ok && s != "" { + resultText = result.Message + "\n" + s + } else if b, err := json.Marshal(result.Detail); err == nil && len(b) > 0 && string(b) != "null" { + resultText = result.Message + "\n" + string(b) + } + } + } updates := map[string]any{ taskExecutionColStatus: taskExecutionStatusSucceeded, "error_message": "", - "result": "ok", + "result": resultText, "finished_at": now, "duration": duration.Milliseconds(), } diff --git a/backend/plugins/drivers/driver_inproc_worker/plugin.go b/backend/plugins/drivers/driver_inproc_worker/plugin.go index 2525955a..119c3b1d 100644 --- a/backend/plugins/drivers/driver_inproc_worker/plugin.go +++ b/backend/plugins/drivers/driver_inproc_worker/plugin.go @@ -119,13 +119,7 @@ func (p *Plugin) ConfigEnabled(view core.ConfigView) bool { func (p *Plugin) Apply(ctx *core.Context) error { p.coreCtx = ctx - if db, err := core.Inject[contracts.DBService](ctx); err == nil && db != nil { - setDBService(db) - } else { - core.When[contracts.DBService](ctx, func(db contracts.DBService) { - setDBService(db) - }) - } + core.Bind[contracts.DBService](ctx, setDBService) taskSvc := newInprocTaskService(ctx.Tasks()) core.Provide[contracts.TaskService](ctx, taskSvc) @@ -151,6 +145,9 @@ func (p *Plugin) Start(ctx context.Context) error { if p.queue == nil { p.queue = NewInprocQueue(p.concurrency, p.queueCapacity, p.coreCtx.Tasks()) } + if p.coreCtx != nil { + p.queue.appCtx = p.coreCtx.Root() + } globalMu.Lock() globalQueue = p.queue diff --git a/backend/plugins/drivers/driver_inproc_worker/plugin_test.go b/backend/plugins/drivers/driver_inproc_worker/plugin_test.go index 2555fd76..85feaff7 100644 --- a/backend/plugins/drivers/driver_inproc_worker/plugin_test.go +++ b/backend/plugins/drivers/driver_inproc_worker/plugin_test.go @@ -80,9 +80,9 @@ func TestInprocWorkerDispatchByTypeTracksExecution(t *testing.T) { require.NoError(t, p.Apply(ctx)) var executedCount atomic.Int32 - ctx.Tasks().Register("system:cleanup", func(_ context.Context, _ []byte) error { + ctx.Tasks().Register("system:cleanup", func(_ context.Context, _ []byte) (*contracts.TaskResultDTO, error) { executedCount.Add(1) - return nil + return &contracts.TaskResultDTO{Message: "cleaned 3 files"}, nil }, extpoints.WithTaskType("system_cleanup"), extpoints.WithTaskName("系统垃圾清理"), @@ -111,6 +111,6 @@ func TestInprocWorkerDispatchByTypeTracksExecution(t *testing.T) { if listErr != nil || total == 0 || len(execs) == 0 { return false } - return execs[0].TaskID == taskID && execs[0].Status == "succeeded" && execs[0].TaskName == "系统垃圾清理" + return execs[0].TaskID == taskID && execs[0].Status == "succeeded" && execs[0].TaskName == "系统垃圾清理" && execs[0].Result == "cleaned 3 files" }, 2*time.Second, 20*time.Millisecond, "inproc worker should persist a succeeded execution record") } diff --git a/backend/plugins/drivers/drivers_test.go b/backend/plugins/drivers/drivers_test.go index 2a323e83..09149511 100644 --- a/backend/plugins/drivers/drivers_test.go +++ b/backend/plugins/drivers/drivers_test.go @@ -239,9 +239,9 @@ func TestAsynqWorkerDispatchTracksExecution(t *testing.T) { core.Provide[contracts.DBService](ctx, &testDBService{db: testDB}) var processed atomic.Bool - ctx.Tasks().Register("system:cleanup", func(_ context.Context, _ []byte) error { + ctx.Tasks().Register("system:cleanup", func(_ context.Context, _ []byte) (*contracts.TaskResultDTO, error) { processed.Store(true) - return nil + return &contracts.TaskResultDTO{Message: "cleaned 3 files"}, nil }, extpoints.WithTaskType("system_cleanup"), extpoints.WithTaskName("系统垃圾清理"), @@ -277,7 +277,7 @@ func TestAsynqWorkerDispatchTracksExecution(t *testing.T) { if listErr != nil || len(execs) == 0 { return false } - return execs[0].TaskID == taskID && execs[0].Status == "succeeded" + return execs[0].TaskID == taskID && execs[0].Status == "succeeded" && execs[0].Result == "cleaned 3 files" }, 5*time.Second, 50*time.Millisecond, "task execution should become succeeded after worker runs") } diff --git a/backend/plugins/infra/storage/plugin.go b/backend/plugins/infra/storage/plugin.go index 3cd105af..1fd7fff9 100644 --- a/backend/plugins/infra/storage/plugin.go +++ b/backend/plugins/infra/storage/plugin.go @@ -48,25 +48,11 @@ func (p *Plugin) Name() string { // Apply mounts the storage service into the Context. func (p *Plugin) Apply(ctx *core.Context) error { - // Bind DBService - if db, err := core.Inject[contracts.DBService](ctx); err == nil && db != nil { + core.Bind[contracts.DBService](ctx, func(db contracts.DBService) { objectstore.SetDBService(db) diskcache.SetDBService(db) - } else { - core.When[contracts.DBService](ctx, func(db contracts.DBService) { - objectstore.SetDBService(db) - diskcache.SetDBService(db) - }) - } - - // Bind CacheService - if cache, err := core.Inject[contracts.CacheService](ctx); err == nil && cache != nil { - objectstore.SetCacheService(cache) - } else { - core.When[contracts.CacheService](ctx, func(cache contracts.CacheService) { - objectstore.SetCacheService(cache) - }) - } + }) + core.Bind[contracts.CacheService](ctx, objectstore.SetCacheService) ctx.OnDispose(func() error { objectstore.SetDBService(nil)