feat(core): bind request services and implement registered tasks

Wire plugin services through Bind/InjectFrom and AppContext so HTTP and
workers resolve dependencies after Apply. Register TaskHandler objects
with persisted results, and implement send_email_code, mail:send,
cleanup_inactive_users, and dispatch_bot_msg.
This commit is contained in:
ryan
2026-09-02 16:59:00 +08:00
parent 30bbe965bf
commit 4f50f6a8f9
49 changed files with 1406 additions and 499 deletions
+43
View File
@@ -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)
}
+11 -1
View File
@@ -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)
}
+40
View File
@@ -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()
+6
View File
@@ -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"`
+16
View File
@@ -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)
+59 -13
View File
@@ -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
}
+28
View File
@@ -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": {
+28
View File
@@ -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": {
+18
View File
@@ -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
@@ -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 下发任务
+8 -56
View File
@@ -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{
@@ -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
+2 -2
View File
@@ -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
}
@@ -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
@@ -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
}
@@ -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
+60 -14
View File
@@ -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
+2 -15
View File
@@ -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)
+1 -8
View File
@@ -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
+1 -1
View File
@@ -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)
@@ -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"
@@ -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)
@@ -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
@@ -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},
},
},
@@ -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
}
@@ -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")
}
@@ -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: "目标接收者",
},
@@ -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
@@ -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
+9 -70
View File
@@ -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)
@@ -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
+8
View File
@@ -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 = "邮件发送失败"
)
+27
View File
@@ -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}))
}
+4
View File
@@ -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"`
+12 -97
View File
@@ -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{
+24 -4
View File
@@ -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
+385
View File
@@ -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("<p>您的验证码是 <b>%s</b>,%d 分钟内有效。</p>", 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:]
}
+109
View File
@@ -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":"<p>ok</p>"}`))
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":"<p>ok</p>"}`))
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{})
}
@@ -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)
@@ -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
@@ -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:
+39 -40
View File
@@ -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()
@@ -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.
@@ -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
@@ -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(),
}
@@ -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
@@ -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")
}
+3 -3
View File
@@ -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")
}
+3 -17
View File
@@ -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)