mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-09-30 06:16:37 +08:00
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:
@@ -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)
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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"`
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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": {
|
||||
|
||||
@@ -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": {
|
||||
|
||||
@@ -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 下发任务
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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 = "邮件发送失败"
|
||||
)
|
||||
|
||||
@@ -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}))
|
||||
}
|
||||
|
||||
|
||||
@@ -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"`
|
||||
|
||||
@@ -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{
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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:]
|
||||
}
|
||||
@@ -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:
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user