mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-11 01:36: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,
|
// 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.
|
// 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)) {
|
func When[T any](ctx *Context, fn func(s T)) {
|
||||||
if ctx == nil {
|
if ctx == nil {
|
||||||
panic("core: nil context provided to When")
|
panic("core: nil context provided to When")
|
||||||
}
|
}
|
||||||
|
|
||||||
targetType := reflect.TypeFor[T]()
|
targetType := reflect.TypeFor[T]()
|
||||||
c := ctx.Container()
|
c := ctx.Root().Container()
|
||||||
|
|
||||||
// If already ready, execute immediately
|
// If already ready, execute immediately
|
||||||
if s, err := Inject[T](ctx); err == nil {
|
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)
|
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) {
|
func TestContextDisposerLifecycle(t *testing.T) {
|
||||||
parent := core.NewContext(context.Background())
|
parent := core.NewContext(context.Background())
|
||||||
child := parent.Fork()
|
child := parent.Fork()
|
||||||
|
|||||||
@@ -43,6 +43,12 @@ type TaskResultDTO struct {
|
|||||||
Detail any `json:"detail,omitempty"`
|
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.
|
// TaskExecutionDTO represents a single task execution record.
|
||||||
type TaskExecutionDTO struct {
|
type TaskExecutionDTO struct {
|
||||||
ID uint64 `json:"id,string"`
|
ID uint64 `json:"id,string"`
|
||||||
|
|||||||
@@ -229,6 +229,22 @@ func TestTaskExtension(t *testing.T) {
|
|||||||
assert.False(t, ok)
|
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) {
|
func TestScheduleExtension(t *testing.T) {
|
||||||
sr := extpoints.NewScheduleRegistry()
|
sr := extpoints.NewScheduleRegistry()
|
||||||
require.NotNil(t, sr)
|
require.NotNil(t, sr)
|
||||||
|
|||||||
@@ -5,6 +5,8 @@ package extpoints
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"Wavelet/core/contracts"
|
"Wavelet/core/contracts"
|
||||||
|
"fmt"
|
||||||
|
"reflect"
|
||||||
"sync"
|
"sync"
|
||||||
"time"
|
"time"
|
||||||
)
|
)
|
||||||
@@ -230,10 +232,15 @@ func NewTaskRegistry() *TaskRegistry {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Register registers a task pattern and its handler with optional configuration.
|
// 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) {
|
func (t *TaskRegistry) Register(pattern string, handler any, opts ...TaskOption) {
|
||||||
t.mu.Lock()
|
t.mu.Lock()
|
||||||
defer t.mu.Unlock()
|
defer t.mu.Unlock()
|
||||||
|
|
||||||
|
if isNilTaskHandler(handler) {
|
||||||
|
panic(fmt.Sprintf("extpoints: nil handler for task pattern %q", pattern))
|
||||||
|
}
|
||||||
|
|
||||||
td := TaskDefinition{
|
td := TaskDefinition{
|
||||||
Pattern: pattern,
|
Pattern: pattern,
|
||||||
Handler: handler,
|
Handler: handler,
|
||||||
@@ -245,8 +252,23 @@ func (t *TaskRegistry) Register(pattern string, handler any, opts ...TaskOption)
|
|||||||
opt(&td)
|
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 {
|
for i, item := range t.tasks {
|
||||||
if item.Pattern == pattern {
|
if item.Pattern == pattern {
|
||||||
t.tasks[i] = td
|
t.tasks[i] = td
|
||||||
@@ -258,13 +280,44 @@ func (t *TaskRegistry) Register(pattern string, handler any, opts ...TaskOption)
|
|||||||
}
|
}
|
||||||
|
|
||||||
t.lookup[pattern] = td
|
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.
|
// Unregister removes a registered task definition by its pattern.
|
||||||
func (t *TaskRegistry) Unregister(pattern string) bool {
|
func (t *TaskRegistry) Unregister(pattern string) bool {
|
||||||
return unregisterEntry(&t.mu, t.lookup, &t.tasks, pattern, func(item TaskDefinition) bool {
|
t.mu.Lock()
|
||||||
return item.Pattern == pattern
|
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.
|
// 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) {
|
func (t *TaskRegistry) Get(pattern string) (TaskDefinition, bool) {
|
||||||
t.mu.RLock()
|
t.mu.RLock()
|
||||||
defer t.mu.RUnlock()
|
defer t.mu.RUnlock()
|
||||||
if td, ok := t.lookup[pattern]; ok {
|
td, ok := t.lookup[pattern]
|
||||||
return td, true
|
return td, ok
|
||||||
}
|
|
||||||
for _, td := range t.tasks {
|
|
||||||
if td.Type == pattern {
|
|
||||||
return td, true
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return TaskDefinition{}, false
|
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -5925,6 +5925,17 @@ const docTemplate = `{
|
|||||||
"user"
|
"user"
|
||||||
],
|
],
|
||||||
"summary": "发送邮箱验证码",
|
"summary": "发送邮箱验证码",
|
||||||
|
"parameters": [
|
||||||
|
{
|
||||||
|
"description": "目标邮箱",
|
||||||
|
"name": "request",
|
||||||
|
"in": "body",
|
||||||
|
"required": true,
|
||||||
|
"schema": {
|
||||||
|
"$ref": "#/definitions/user.sendEmailCodeRequest"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
],
|
||||||
"responses": {
|
"responses": {
|
||||||
"200": {
|
"200": {
|
||||||
"description": "发送成功",
|
"description": "发送成功",
|
||||||
@@ -5937,6 +5948,12 @@ const docTemplate = `{
|
|||||||
"schema": {
|
"schema": {
|
||||||
"$ref": "#/definitions/response.Any"
|
"$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": {
|
"user.updateProfileRequest": {
|
||||||
"type": "object",
|
"type": "object",
|
||||||
"properties": {
|
"properties": {
|
||||||
|
|||||||
@@ -5918,6 +5918,17 @@
|
|||||||
"user"
|
"user"
|
||||||
],
|
],
|
||||||
"summary": "发送邮箱验证码",
|
"summary": "发送邮箱验证码",
|
||||||
|
"parameters": [
|
||||||
|
{
|
||||||
|
"description": "目标邮箱",
|
||||||
|
"name": "request",
|
||||||
|
"in": "body",
|
||||||
|
"required": true,
|
||||||
|
"schema": {
|
||||||
|
"$ref": "#/definitions/user.sendEmailCodeRequest"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
],
|
||||||
"responses": {
|
"responses": {
|
||||||
"200": {
|
"200": {
|
||||||
"description": "发送成功",
|
"description": "发送成功",
|
||||||
@@ -5930,6 +5941,12 @@
|
|||||||
"schema": {
|
"schema": {
|
||||||
"$ref": "#/definitions/response.Any"
|
"$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": {
|
"user.updateProfileRequest": {
|
||||||
"type": "object",
|
"type": "object",
|
||||||
"properties": {
|
"properties": {
|
||||||
|
|||||||
@@ -1363,6 +1363,13 @@ definitions:
|
|||||||
- password
|
- password
|
||||||
- username
|
- username
|
||||||
type: object
|
type: object
|
||||||
|
user.sendEmailCodeRequest:
|
||||||
|
properties:
|
||||||
|
email:
|
||||||
|
type: string
|
||||||
|
required:
|
||||||
|
- email
|
||||||
|
type: object
|
||||||
user.updateProfileRequest:
|
user.updateProfileRequest:
|
||||||
properties:
|
properties:
|
||||||
avatar_url:
|
avatar_url:
|
||||||
@@ -4937,6 +4944,13 @@ paths:
|
|||||||
consumes:
|
consumes:
|
||||||
- application/json
|
- application/json
|
||||||
description: 向指定邮箱发送验证码(用于注册场景)
|
description: 向指定邮箱发送验证码(用于注册场景)
|
||||||
|
parameters:
|
||||||
|
- description: 目标邮箱
|
||||||
|
in: body
|
||||||
|
name: request
|
||||||
|
required: true
|
||||||
|
schema:
|
||||||
|
$ref: '#/definitions/user.sendEmailCodeRequest'
|
||||||
produces:
|
produces:
|
||||||
- application/json
|
- application/json
|
||||||
responses:
|
responses:
|
||||||
@@ -4948,6 +4962,10 @@ paths:
|
|||||||
description: 参数错误
|
description: 参数错误
|
||||||
schema:
|
schema:
|
||||||
$ref: '#/definitions/response.Any'
|
$ref: '#/definitions/response.Any'
|
||||||
|
"500":
|
||||||
|
description: 发送失败
|
||||||
|
schema:
|
||||||
|
$ref: '#/definitions/response.Any'
|
||||||
summary: 发送邮箱验证码
|
summary: 发送邮箱验证码
|
||||||
tags:
|
tags:
|
||||||
- user
|
- user
|
||||||
|
|||||||
@@ -48,7 +48,7 @@ func abortTaskLogicError(c *gin.Context, err error) bool {
|
|||||||
// @Failure 403 {object} response.Any "无管理员权限"
|
// @Failure 403 {object} response.Any "无管理员权限"
|
||||||
// @Router /api/v1/admin/tasks/types [get]
|
// @Router /api/v1/admin/tasks/types [get]
|
||||||
func ListTaskTypes(c *gin.Context) {
|
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 下发任务
|
// DispatchTask 下发任务
|
||||||
|
|||||||
@@ -12,7 +12,6 @@ import (
|
|||||||
"Wavelet/plugins/domain/admin/handler"
|
"Wavelet/plugins/domain/admin/handler"
|
||||||
"Wavelet/plugins/domain/admin/model"
|
"Wavelet/plugins/domain/admin/model"
|
||||||
"Wavelet/plugins/domain/admin/service"
|
"Wavelet/plugins/domain/admin/service"
|
||||||
"context"
|
|
||||||
"embed"
|
"embed"
|
||||||
"reflect"
|
"reflect"
|
||||||
|
|
||||||
@@ -85,56 +84,13 @@ func (p *Plugin) Apply(ctx *core.Context) error {
|
|||||||
_ = ctx.Config().Bind("clickhouse", &chCfg)
|
_ = ctx.Config().Bind("clickhouse", &chCfg)
|
||||||
service.SetClickHouseConfig(chCfg)
|
service.SetClickHouseConfig(chCfg)
|
||||||
|
|
||||||
// 0. Bind Services reactively
|
core.Bind[contracts.DBService](ctx, service.SetDBService)
|
||||||
if db, err := core.Inject[contracts.DBService](ctx); err == nil && db != nil {
|
core.Bind[contracts.CacheService](ctx, service.SetCacheService)
|
||||||
service.SetDBService(db)
|
core.Bind[contracts.UserService](ctx, service.SetUserService)
|
||||||
} else {
|
core.Bind[contracts.AuthService](ctx, service.SetAuthService)
|
||||||
core.When[contracts.DBService](ctx, func(db contracts.DBService) {
|
core.Bind[contracts.TaskService](ctx, service.SetTaskService)
|
||||||
service.SetDBService(db)
|
core.Bind[contracts.StorageService](ctx, service.SetStorageService)
|
||||||
})
|
core.Bind[contracts.RiskControlService](ctx, service.SetRiskControlService)
|
||||||
}
|
|
||||||
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)
|
|
||||||
})
|
|
||||||
}
|
|
||||||
service.SetEventEmitter(ctx.Events().Emit)
|
service.SetEventEmitter(ctx.Events().Emit)
|
||||||
|
|
||||||
ctx.OnDispose(func() error {
|
ctx.OnDispose(func() error {
|
||||||
@@ -175,11 +131,7 @@ func (p *Plugin) Apply(ctx *core.Context) error {
|
|||||||
ctx.Router().RegisterWhitelist("/robots.txt")
|
ctx.Router().RegisterWhitelist("/robots.txt")
|
||||||
|
|
||||||
// 2. Register Background Tasks
|
// 2. Register Background Tasks
|
||||||
logSwitchHandler := &service.LogDBSwitchHandler{}
|
ctx.Task().Register(service.LogDBSwitchTask, &service.LogDBSwitchHandler{}, extpoints.WithTaskMeta(service.LogDBSwitchMeta))
|
||||||
ctx.Task().Register(service.LogDBSwitchTask, func(c context.Context, payload []byte) error {
|
|
||||||
_, err := logSwitchHandler.Execute(c, payload)
|
|
||||||
return err
|
|
||||||
}, extpoints.WithTaskMeta(service.LogDBSwitchMeta))
|
|
||||||
|
|
||||||
// 3. Register Settings Schemas
|
// 3. Register Settings Schemas
|
||||||
ctx.Settings().Register(extpoints.SettingSchema{
|
ctx.Settings().Register(extpoints.SettingSchema{
|
||||||
|
|||||||
@@ -5,6 +5,7 @@
|
|||||||
package repository
|
package repository
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"Wavelet/core"
|
||||||
"Wavelet/core/contracts"
|
"Wavelet/core/contracts"
|
||||||
"Wavelet/pkg/cache/ram"
|
"Wavelet/pkg/cache/ram"
|
||||||
"Wavelet/pkg/logger"
|
"Wavelet/pkg/logger"
|
||||||
@@ -56,6 +57,9 @@ func ResetServices() {
|
|||||||
|
|
||||||
// GetDB returns the GORM DB instance bound to the context if available.
|
// GetDB returns the GORM DB instance bound to the context if available.
|
||||||
func GetDB(ctx context.Context) *gorm.DB {
|
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()
|
repoMu.RLock()
|
||||||
defer repoMu.RUnlock()
|
defer repoMu.RUnlock()
|
||||||
if dbService == nil {
|
if dbService == nil {
|
||||||
@@ -65,7 +69,10 @@ func GetDB(ctx context.Context) *gorm.DB {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// GetCache returns the unified CacheService instance.
|
// 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()
|
repoMu.RLock()
|
||||||
defer repoMu.RUnlock()
|
defer repoMu.RUnlock()
|
||||||
return cacheService
|
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.
|
// AccessLogs queries the analytical access log store and decorates rows with user names.
|
||||||
func AccessLogs(ctx context.Context, q model.AccessLogQuery) (model.AccessLogsResponse, error) {
|
func AccessLogs(ctx context.Context, q model.AccessLogQuery) (model.AccessLogsResponse, error) {
|
||||||
rc := GetRiskControlService()
|
rc := GetRiskControlService(ctx)
|
||||||
if rc == nil {
|
if rc == nil {
|
||||||
return model.AccessLogsResponse{}, errs.ErrLogStoreUnavailable
|
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.
|
// AccessLogAnalytics aggregates the daily trend of the access log store.
|
||||||
func AccessLogAnalytics(ctx context.Context) (model.LogsAnalyticsResponse, error) {
|
func AccessLogAnalytics(ctx context.Context) (model.LogsAnalyticsResponse, error) {
|
||||||
rc := GetRiskControlService()
|
rc := GetRiskControlService(ctx)
|
||||||
if rc == nil {
|
if rc == nil {
|
||||||
return model.LogsAnalyticsResponse{}, errs.ErrLogStoreUnavailable
|
return model.LogsAnalyticsResponse{}, errs.ErrLogStoreUnavailable
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -109,7 +109,7 @@ func (h *LogDBSwitchHandler) Execute(ctx context.Context, payload []byte) (*cont
|
|||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
taskSvc := GetTaskService()
|
taskSvc := GetTaskService(ctx)
|
||||||
if taskSvc != nil {
|
if taskSvc != nil {
|
||||||
taskSvc.AppendLog(ctx, "开始切换日志数据库:%s -> %s", source, p.Target)
|
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 rc != nil {
|
||||||
if err := rc.SwitchLogEngine(ctx, p.Target); err != nil {
|
if err := rc.SwitchLogEngine(ctx, p.Target); err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
|
|||||||
@@ -5,6 +5,7 @@
|
|||||||
package service
|
package service
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"Wavelet/core"
|
||||||
"Wavelet/core/contracts"
|
"Wavelet/core/contracts"
|
||||||
"Wavelet/plugins/domain/admin/errs"
|
"Wavelet/plugins/domain/admin/errs"
|
||||||
"Wavelet/plugins/domain/admin/repository"
|
"Wavelet/plugins/domain/admin/repository"
|
||||||
@@ -112,6 +113,9 @@ func ResetServices() {
|
|||||||
|
|
||||||
// GetDB returns the GORM DB instance bound to the context if available.
|
// GetDB returns the GORM DB instance bound to the context if available.
|
||||||
func GetDB(ctx context.Context) *gorm.DB {
|
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()
|
servicesMu.RLock()
|
||||||
defer servicesMu.RUnlock()
|
defer servicesMu.RUnlock()
|
||||||
if dbService == nil {
|
if dbService == nil {
|
||||||
@@ -121,42 +125,60 @@ func GetDB(ctx context.Context) *gorm.DB {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// GetCache returns the unified CacheService instance.
|
// 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()
|
servicesMu.RLock()
|
||||||
defer servicesMu.RUnlock()
|
defer servicesMu.RUnlock()
|
||||||
return cacheService
|
return cacheService
|
||||||
}
|
}
|
||||||
|
|
||||||
// GetUserService returns the UserService instance.
|
// 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()
|
servicesMu.RLock()
|
||||||
defer servicesMu.RUnlock()
|
defer servicesMu.RUnlock()
|
||||||
return userService
|
return userService
|
||||||
}
|
}
|
||||||
|
|
||||||
// GetAuthService returns the AuthService instance.
|
// 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()
|
servicesMu.RLock()
|
||||||
defer servicesMu.RUnlock()
|
defer servicesMu.RUnlock()
|
||||||
return authService
|
return authService
|
||||||
}
|
}
|
||||||
|
|
||||||
// GetTaskService returns the TaskService instance.
|
// 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()
|
servicesMu.RLock()
|
||||||
defer servicesMu.RUnlock()
|
defer servicesMu.RUnlock()
|
||||||
return taskService
|
return taskService
|
||||||
}
|
}
|
||||||
|
|
||||||
// GetStorageService returns the StorageService instance.
|
// 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()
|
servicesMu.RLock()
|
||||||
defer servicesMu.RUnlock()
|
defer servicesMu.RUnlock()
|
||||||
return storageSvc
|
return storageSvc
|
||||||
}
|
}
|
||||||
|
|
||||||
// GetRiskControlService returns the RiskControlService instance.
|
// 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()
|
servicesMu.RLock()
|
||||||
defer servicesMu.RUnlock()
|
defer servicesMu.RUnlock()
|
||||||
return riskControlService
|
return riskControlService
|
||||||
@@ -195,8 +217,8 @@ func requireAuthService(ctx context.Context) (contracts.AuthService, error) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// requireTaskService resolves the injected task contract service.
|
// requireTaskService resolves the injected task contract service.
|
||||||
func requireTaskService() (contracts.TaskService, error) {
|
func requireTaskService(ctx context.Context) (contracts.TaskService, error) {
|
||||||
taskSvc := GetTaskService()
|
taskSvc := GetTaskService(ctx)
|
||||||
if taskSvc == nil {
|
if taskSvc == nil {
|
||||||
return nil, errs.ErrTaskServiceUnavailable
|
return nil, errs.ErrTaskServiceUnavailable
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -117,7 +117,7 @@ func formatDuration(d time.Duration) string {
|
|||||||
func LogDatabaseStatus(ctx context.Context) model.LogDatabaseStatus {
|
func LogDatabaseStatus(ctx context.Context) model.LogDatabaseStatus {
|
||||||
activeDB := logDBNameSQLite
|
activeDB := logDBNameSQLite
|
||||||
migration := logMigrationIdle
|
migration := logMigrationIdle
|
||||||
if rc := GetRiskControlService(); rc != nil {
|
if rc := GetRiskControlService(ctx); rc != nil {
|
||||||
activeDB = rc.ActiveLogEngine(ctx)
|
activeDB = rc.ActiveLogEngine(ctx)
|
||||||
if rc.IsLogEngineMigrating(ctx) {
|
if rc.IsLogEngineMigrating(ctx) {
|
||||||
migration = logMigrationInProgress
|
migration = logMigrationInProgress
|
||||||
|
|||||||
@@ -17,8 +17,8 @@ import (
|
|||||||
)
|
)
|
||||||
|
|
||||||
// ListTaskTypes returns every dispatchable task type declared in the task registry.
|
// ListTaskTypes returns every dispatchable task type declared in the task registry.
|
||||||
func ListTaskTypes() []contracts.TaskMetaDTO {
|
func ListTaskTypes(ctx context.Context) []contracts.TaskMetaDTO {
|
||||||
taskSvc := GetTaskService()
|
taskSvc := GetTaskService(ctx)
|
||||||
if taskSvc == nil {
|
if taskSvc == nil {
|
||||||
return []contracts.TaskMetaDTO{}
|
return []contracts.TaskMetaDTO{}
|
||||||
}
|
}
|
||||||
@@ -27,7 +27,7 @@ func ListTaskTypes() []contracts.TaskMetaDTO {
|
|||||||
|
|
||||||
// DispatchTask validates and enqueues a manual task run, returning the new task id.
|
// DispatchTask validates and enqueues a manual task run, returning the new task id.
|
||||||
func DispatchTask(ctx context.Context, req model.DispatchTaskRequest) (string, error) {
|
func DispatchTask(ctx context.Context, req model.DispatchTaskRequest) (string, error) {
|
||||||
taskSvc, err := requireTaskService()
|
taskSvc, err := requireTaskService(ctx)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return "", err
|
return "", err
|
||||||
}
|
}
|
||||||
@@ -67,29 +67,75 @@ func ListTaskExecutions(
|
|||||||
ctx context.Context,
|
ctx context.Context,
|
||||||
req model.ListTaskExecutionsRequest,
|
req model.ListTaskExecutionsRequest,
|
||||||
) ([]model.TaskExecution, int64, error) {
|
) ([]model.TaskExecution, int64, error) {
|
||||||
if req.TaskType != "" {
|
taskSvc, err := requireTaskService(ctx)
|
||||||
if taskSvc := GetTaskService(); taskSvc != nil {
|
if err != nil {
|
||||||
if meta, ok := taskSvc.GetTaskMeta(req.TaskType); ok {
|
return nil, 0, err
|
||||||
req.TaskType = meta.Name
|
}
|
||||||
}
|
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 {
|
if err != nil {
|
||||||
return nil, 0, err
|
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
|
return executions, total, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// TaskExecution loads a single execution record including its buffered log.
|
// TaskExecution loads a single execution record including its buffered log.
|
||||||
func TaskExecution(ctx context.Context, id uint64) (*model.TaskExecution, error) {
|
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.
|
// RetryTask re-dispatches a failed execution as a new task run.
|
||||||
func RetryTask(ctx context.Context, id uint64) (string, error) {
|
func RetryTask(ctx context.Context, id uint64) (string, error) {
|
||||||
taskSvc, err := requireTaskService()
|
taskSvc, err := requireTaskService(ctx)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return "", err
|
return "", err
|
||||||
}
|
}
|
||||||
@@ -126,7 +172,7 @@ func CreateSchedule(ctx context.Context, req model.CreateScheduleRequest) (*mode
|
|||||||
return nil, errs.ErrInvalidCronExpression
|
return nil, errs.ErrInvalidCronExpression
|
||||||
}
|
}
|
||||||
|
|
||||||
taskSvc, err := requireTaskService()
|
taskSvc, err := requireTaskService(ctx)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
@@ -168,7 +214,7 @@ func UpdateSchedule(ctx context.Context, id uint64, req model.UpdateScheduleRequ
|
|||||||
return nil, errs.ErrInvalidCronExpression
|
return nil, errs.ErrInvalidCronExpression
|
||||||
}
|
}
|
||||||
|
|
||||||
taskSvc, err := requireTaskService()
|
taskSvc, err := requireTaskService(ctx)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
@@ -203,7 +249,7 @@ func DeleteSchedule(ctx context.Context, id uint64) error {
|
|||||||
return fmt.Errorf("%s: %w", errs.ScheduleDeleteFailed, err)
|
return fmt.Errorf("%s: %w", errs.ScheduleDeleteFailed, err)
|
||||||
}
|
}
|
||||||
|
|
||||||
if taskSvc := GetTaskService(); taskSvc != nil {
|
if taskSvc := GetTaskService(ctx); taskSvc != nil {
|
||||||
reloadScheduler(ctx, taskSvc)
|
reloadScheduler(ctx, taskSvc)
|
||||||
}
|
}
|
||||||
return nil
|
return nil
|
||||||
|
|||||||
@@ -87,21 +87,8 @@ func (p *Plugin) Apply(ctx *core.Context) error {
|
|||||||
SetSessionConfig(cfg)
|
SetSessionConfig(cfg)
|
||||||
}
|
}
|
||||||
|
|
||||||
// 0. Bind DBService & CacheService from Context
|
core.Bind[contracts.DBService](ctx, setDBService)
|
||||||
if db, err := core.Inject[contracts.DBService](ctx); err == nil && db != nil {
|
core.Bind[contracts.CacheService](ctx, setCacheService)
|
||||||
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)
|
|
||||||
})
|
|
||||||
}
|
|
||||||
ctx.OnDispose(func() error {
|
ctx.OnDispose(func() error {
|
||||||
setDBService(nil)
|
setDBService(nil)
|
||||||
setCacheService(nil)
|
setCacheService(nil)
|
||||||
|
|||||||
@@ -59,14 +59,7 @@ func (p *Plugin) Apply(ctx *core.Context) error {
|
|||||||
SetSecret([]byte(cfg.SessionSecret))
|
SetSecret([]byte(cfg.SessionSecret))
|
||||||
}
|
}
|
||||||
|
|
||||||
// 0. Bind DBService from Context
|
core.Bind[contracts.DBService](ctx, setDBService)
|
||||||
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)
|
|
||||||
})
|
|
||||||
}
|
|
||||||
ctx.OnDispose(func() error {
|
ctx.OnDispose(func() error {
|
||||||
setDBService(nil)
|
setDBService(nil)
|
||||||
return nil
|
return nil
|
||||||
|
|||||||
@@ -236,7 +236,7 @@ func TestUserPlugin(t *testing.T) {
|
|||||||
assert.Len(t, list, 1)
|
assert.Len(t, list, 1)
|
||||||
assert.Equal(t, "bob", list[0].Username)
|
assert.Equal(t, "bob", list[0].Username)
|
||||||
|
|
||||||
// 9. Tasks & Schedules
|
// 9. Tasks
|
||||||
taskDef, ok := ctx.Tasks().Get("user:send_email_code")
|
taskDef, ok := ctx.Tasks().Get("user:send_email_code")
|
||||||
require.True(t, ok)
|
require.True(t, ok)
|
||||||
assert.Equal(t, 3, taskDef.Retry)
|
assert.Equal(t, 3, taskDef.Retry)
|
||||||
|
|||||||
@@ -28,13 +28,15 @@ var (
|
|||||||
|
|
||||||
// User-facing validation and error message constants.
|
// User-facing validation and error message constants.
|
||||||
const (
|
const (
|
||||||
ErrNameRequired = "name is required"
|
ErrNameRequired = "name is required"
|
||||||
ErrTypeInvalid = "type must be telegram or qq"
|
ErrTypeInvalid = "type must be telegram or qq"
|
||||||
ErrTelegramTokenRequired = "telegram bot secret is required" //nolint:gosec // user-facing validation text
|
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
|
ErrQQCredentialsRequired = "qq app id and secret are required" //nolint:gosec // user-facing validation text
|
||||||
ErrChannelNotFound = "channel not found"
|
ErrChannelNotFound = "channel not found"
|
||||||
ErrChannelProbeFailed = "channel probe failed"
|
ErrChannelProbeFailed = "channel probe failed"
|
||||||
MaskedSecret = "********"
|
ErrBotDispatchTextRequired = "message text is required"
|
||||||
|
ErrBotChannelNotRegistered = "channel adapter is not registered"
|
||||||
|
MaskedSecret = "********"
|
||||||
|
|
||||||
ErrLoginRequired = "login required"
|
ErrLoginRequired = "login required"
|
||||||
ErrInvalidBindingID = "invalid binding id"
|
ErrInvalidBindingID = "invalid binding id"
|
||||||
|
|||||||
@@ -10,6 +10,8 @@ import (
|
|||||||
"Wavelet/core/extpoints"
|
"Wavelet/core/extpoints"
|
||||||
"Wavelet/pkg/ginutil"
|
"Wavelet/pkg/ginutil"
|
||||||
"Wavelet/pkg/util"
|
"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/handler"
|
||||||
"Wavelet/plugins/domain/message_gateway/model"
|
"Wavelet/plugins/domain/message_gateway/model"
|
||||||
"Wavelet/plugins/domain/message_gateway/repository"
|
"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 != "" {
|
if err := ctx.Config().Bind("app", &cfg); err == nil && cfg.SessionSecret != "" {
|
||||||
service.SetCredentialSecret(cfg.SessionSecret)
|
service.SetCredentialSecret(cfg.SessionSecret)
|
||||||
}
|
}
|
||||||
// 0. Bind DBService, CacheService, TaskService, UserService
|
core.Bind[contracts.DBService](ctx, repository.SetDBService)
|
||||||
if db, err := core.Inject[contracts.DBService](ctx); err == nil && db != nil {
|
core.Bind[contracts.CacheService](ctx, func(cache contracts.CacheService) {
|
||||||
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 {
|
|
||||||
repository.SetCacheService(cache)
|
repository.SetCacheService(cache)
|
||||||
service.SetCacheService(cache)
|
service.SetCacheService(cache)
|
||||||
} else {
|
})
|
||||||
core.When[contracts.CacheService](ctx, func(cache contracts.CacheService) {
|
core.Bind[contracts.TaskService](ctx, service.SetTaskService)
|
||||||
repository.SetCacheService(cache)
|
core.Bind[contracts.UserService](ctx, service.SetUserService)
|
||||||
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)
|
|
||||||
})
|
|
||||||
}
|
|
||||||
ctx.OnDispose(func() error {
|
ctx.OnDispose(func() error {
|
||||||
repository.SetDBService(nil)
|
repository.SetDBService(nil)
|
||||||
repository.SetCacheService(nil)
|
repository.SetCacheService(nil)
|
||||||
@@ -159,6 +137,9 @@ func (p *Plugin) Apply(ctx *core.Context) error {
|
|||||||
// 4. Register Admin Push HTTP Routes
|
// 4. Register Admin Push HTTP Routes
|
||||||
handler.RegisterAdminPushRoutes(ctx.Router().Group("/api/v1/admin"), loginMW, adminMW)
|
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
|
const defaultTaskRetry = 3
|
||||||
pushHandler := &service.PushHandler{}
|
pushHandler := &service.PushHandler{}
|
||||||
|
|
||||||
@@ -179,15 +160,8 @@ func (p *Plugin) Apply(ctx *core.Context) error {
|
|||||||
return pushHandler.Execute(c, payload)
|
return pushHandler.Execute(c, payload)
|
||||||
}, extpoints.WithTaskMeta(service.SendNotificationMeta), extpoints.WithTaskRetry(defaultTaskRetry))
|
}, extpoints.WithTaskMeta(service.SendNotificationMeta), extpoints.WithTaskRetry(defaultTaskRetry))
|
||||||
|
|
||||||
ctx.Task().Register("message_gateway:dispatch_bot_msg", func(_ context.Context, _ []byte) error {
|
ctx.Task().Register(service.TaskDispatchBotMsg, &service.BotDispatchHandler{},
|
||||||
return nil
|
extpoints.WithTaskMeta(service.BotDispatchMeta))
|
||||||
},
|
|
||||||
extpoints.WithTaskType("dispatch_bot_msg"),
|
|
||||||
extpoints.WithTaskName("分发 Bot 消息"),
|
|
||||||
extpoints.WithTaskDescription("异步处理与分发 Bot 下行消息"),
|
|
||||||
extpoints.WithTaskCategory("messaging"),
|
|
||||||
extpoints.WithTaskQueue("default"),
|
|
||||||
)
|
|
||||||
|
|
||||||
ctx.Task().Register("message_gateway:cleanup_pairing_codes", func(c context.Context, _ []byte) error {
|
ctx.Task().Register("message_gateway:cleanup_pairing_codes", func(c context.Context, _ []byte) error {
|
||||||
return repository.DeleteExpiredPairingCodes(c)
|
return repository.DeleteExpiredPairingCodes(c)
|
||||||
|
|||||||
@@ -132,6 +132,15 @@ func ListBindingsByUser(ctx context.Context, userID uint64) ([]model.MessageBind
|
|||||||
return rows, nil
|
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.
|
// GetMessageBinding loads a binding by id.
|
||||||
func GetMessageBinding(ctx context.Context, id uint64) (*model.MessageBinding, error) {
|
func GetMessageBinding(ctx context.Context, id uint64) (*model.MessageBinding, error) {
|
||||||
var b model.MessageBinding
|
var b model.MessageBinding
|
||||||
|
|||||||
@@ -27,14 +27,14 @@ func ListDefinitions() []model.Definition {
|
|||||||
{
|
{
|
||||||
Type: model.MessageChannelTypeTelegram,
|
Type: model.MessageChannelTypeTelegram,
|
||||||
Fields: []model.Field{
|
Fields: []model.Field{
|
||||||
{Key: "token", Type: "password", Required: true},
|
{Key: "token", Type: model.TypePassword, Required: true},
|
||||||
{Key: "api_base", Type: "text", Required: false},
|
{Key: "api_base", Type: model.TypeText, Required: false},
|
||||||
},
|
},
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
Type: model.MessageChannelTypeQQ,
|
Type: model.MessageChannelTypeQQ,
|
||||||
Fields: []model.Field{
|
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},
|
{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.
|
// PushRegistryAdapter adapts contracts.PushRegistry onto the built-in event store.
|
||||||
type PushRegistryAdapter struct{}
|
type PushRegistryAdapter struct{}
|
||||||
|
|
||||||
|
// RegisterBuiltInEvent records a built-in push event definition.
|
||||||
func (PushRegistryAdapter) RegisterBuiltInEvent(meta contracts.PushEventMeta) {
|
func (PushRegistryAdapter) RegisterBuiltInEvent(meta contracts.PushEventMeta) {
|
||||||
RegisterBuiltInEvent(eventMetadataFromContract(meta))
|
RegisterBuiltInEvent(eventMetadataFromContract(meta))
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// SyncEvents persists registered built-in events into the database.
|
||||||
func (PushRegistryAdapter) SyncEvents(ctx context.Context) error {
|
func (PushRegistryAdapter) SyncEvents(ctx context.Context) error {
|
||||||
return SyncEvents(ctx)
|
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.
|
// 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) {
|
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 {
|
if err != nil {
|
||||||
return model.PushEvent{}, err
|
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
|
// GetEventInfo derives the event key, display name and default template for a
|
||||||
// task-completion based event or a registered built-in event key.
|
// 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 != "" {
|
if req.TaskType != "" {
|
||||||
taskName := req.TaskType
|
taskName := req.TaskType
|
||||||
if taskSvc := GetTaskService(); taskSvc != nil {
|
if taskSvc := GetTaskService(ctx); taskSvc != nil {
|
||||||
if meta, ok := taskSvc.GetTaskMeta(req.TaskType); ok {
|
if meta, ok := taskSvc.GetTaskMeta(req.TaskType); ok {
|
||||||
taskName = meta.DisplayName
|
taskName = meta.DisplayName
|
||||||
}
|
}
|
||||||
@@ -694,7 +696,7 @@ func EnqueuePushTask(ctx context.Context, payload model.SendPayload) error {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
if taskSvc := GetTaskService(); taskSvc != nil {
|
if taskSvc := GetTaskService(ctx); taskSvc != nil {
|
||||||
_, err = taskSvc.Dispatch(ctx, "send_notification", payloadBytes, "system")
|
_, err = taskSvc.Dispatch(ctx, "send_notification", payloadBytes, "system")
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
@@ -749,13 +751,13 @@ var SendNotificationMeta = contracts.TaskMetaDTO{
|
|||||||
Category: "push",
|
Category: "push",
|
||||||
SupportsTime: false,
|
SupportsTime: false,
|
||||||
MaxRetry: 3,
|
MaxRetry: 3,
|
||||||
Queue: "default",
|
Queue: taskQueueDefault,
|
||||||
Retryable: true,
|
Retryable: true,
|
||||||
Params: []contracts.TaskParamDTO{
|
Params: []contracts.TaskParamDTO{
|
||||||
{
|
{
|
||||||
Name: "event_key",
|
Name: "event_key",
|
||||||
Label: "事件标识",
|
Label: "事件标识",
|
||||||
Type: "string",
|
Type: taskParamTypeString,
|
||||||
Required: true,
|
Required: true,
|
||||||
Placeholder: "admin_login",
|
Placeholder: "admin_login",
|
||||||
Description: "事件标识 (如 admin_login)",
|
Description: "事件标识 (如 admin_login)",
|
||||||
@@ -763,7 +765,7 @@ var SendNotificationMeta = contracts.TaskMetaDTO{
|
|||||||
{
|
{
|
||||||
Name: "target",
|
Name: "target",
|
||||||
Label: "目标接收者",
|
Label: "目标接收者",
|
||||||
Type: "string",
|
Type: taskParamTypeString,
|
||||||
Required: false,
|
Required: false,
|
||||||
Description: "目标接收者",
|
Description: "目标接收者",
|
||||||
},
|
},
|
||||||
|
|||||||
@@ -251,10 +251,8 @@ func SetUserService(s contracts.UserService) {
|
|||||||
|
|
||||||
// GetCache resolves the cache service for the context.
|
// GetCache resolves the cache service for the context.
|
||||||
func GetCache(ctx context.Context) contracts.CacheService {
|
func GetCache(ctx context.Context) contracts.CacheService {
|
||||||
if c, ok := ctx.(*core.Context); ok && c != nil {
|
if s, err := core.InjectFrom[contracts.CacheService](ctx); err == nil && s != nil {
|
||||||
if s, err := core.Inject[contracts.CacheService](c); err == nil && s != nil {
|
return s
|
||||||
return s
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
cacheMu.RLock()
|
cacheMu.RLock()
|
||||||
s := cacheSvc
|
s := cacheSvc
|
||||||
@@ -263,7 +261,10 @@ func GetCache(ctx context.Context) contracts.CacheService {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// GetTaskService returns the task service.
|
// 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()
|
taskMu.RLock()
|
||||||
defer taskMu.RUnlock()
|
defer taskMu.RUnlock()
|
||||||
return taskSvc
|
return taskSvc
|
||||||
@@ -271,10 +272,8 @@ func GetTaskService() contracts.TaskService {
|
|||||||
|
|
||||||
// GetUserService resolves the user service for the context.
|
// GetUserService resolves the user service for the context.
|
||||||
func GetUserService(ctx context.Context) contracts.UserService {
|
func GetUserService(ctx context.Context) contracts.UserService {
|
||||||
if c, ok := ctx.(*core.Context); ok && c != nil {
|
if s, err := core.InjectFrom[contracts.UserService](ctx); err == nil && s != nil {
|
||||||
if s, err := core.Inject[contracts.UserService](c); err == nil && s != nil {
|
return s
|
||||||
return s
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
userMu.RLock()
|
userMu.RLock()
|
||||||
s := userSvc
|
s := userSvc
|
||||||
|
|||||||
@@ -94,14 +94,7 @@ func (p *Plugin) Apply(ctx *core.Context) error {
|
|||||||
SetAccessLogEnabled(chCfg.Enabled)
|
SetAccessLogEnabled(chCfg.Enabled)
|
||||||
logstore.SetDefaultDatabases(dbCfg.Enabled, chCfg.Enabled)
|
logstore.SetDefaultDatabases(dbCfg.Enabled, chCfg.Enabled)
|
||||||
|
|
||||||
// 0. Bind DBService
|
core.Bind[contracts.DBService](ctx, logstore.SetDBService)
|
||||||
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)
|
|
||||||
})
|
|
||||||
}
|
|
||||||
ctx.OnDispose(func() error {
|
ctx.OnDispose(func() error {
|
||||||
logstore.SetDBService(nil)
|
logstore.SetDBService(nil)
|
||||||
return nil
|
return nil
|
||||||
|
|||||||
@@ -12,7 +12,6 @@ import (
|
|||||||
"Wavelet/plugins/domain/upload/handler"
|
"Wavelet/plugins/domain/upload/handler"
|
||||||
"Wavelet/plugins/domain/upload/shared"
|
"Wavelet/plugins/domain/upload/shared"
|
||||||
"Wavelet/plugins/domain/upload/task"
|
"Wavelet/plugins/domain/upload/task"
|
||||||
"context"
|
|
||||||
"embed"
|
"embed"
|
||||||
"reflect"
|
"reflect"
|
||||||
|
|
||||||
@@ -56,50 +55,11 @@ func (p *Plugin) Manifest() core.Manifest {
|
|||||||
|
|
||||||
// Apply registers upload routes, tasks, and settings into the Context.
|
// Apply registers upload routes, tasks, and settings into the Context.
|
||||||
func (p *Plugin) Apply(ctx *core.Context) error {
|
func (p *Plugin) Apply(ctx *core.Context) error {
|
||||||
// Bind DBService
|
core.Bind[contracts.DBService](ctx, shared.SetDBService)
|
||||||
if db, err := core.Inject[contracts.DBService](ctx); err == nil && db != nil {
|
core.Bind[contracts.CacheService](ctx, shared.SetCacheService)
|
||||||
shared.SetDBService(db)
|
core.Bind[contracts.StorageService](ctx, shared.SetStorageService)
|
||||||
} else {
|
core.Bind[contracts.TaskService](ctx, shared.SetTaskService)
|
||||||
core.When[contracts.DBService](ctx, func(db contracts.DBService) {
|
core.Bind[contracts.AuthService](ctx, shared.SetAuthService)
|
||||||
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)
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
ctx.OnDispose(func() error {
|
ctx.OnDispose(func() error {
|
||||||
shared.ResetServices()
|
shared.ResetServices()
|
||||||
@@ -147,31 +107,10 @@ func (p *Plugin) Apply(ctx *core.Context) error {
|
|||||||
defaultSingleRetry = 1
|
defaultSingleRetry = 1
|
||||||
)
|
)
|
||||||
|
|
||||||
// 3. Register tasks. Handlers take raw payload bytes rather than a driver
|
ctx.Task().Register(task.SystemCleanupTask, &task.SystemCleanupHandler{}, extpoints.WithTaskMeta(task.SystemCleanupMeta), extpoints.WithTaskRetry(defaultCleanupRetry))
|
||||||
// specific task type so they run under both the asynq and in-process workers.
|
ctx.Task().Register(task.RebuildUploadStatsTask, &task.RebuildUploadStatsHandler{}, extpoints.WithTaskMeta(task.RebuildUploadStatsMeta), extpoints.WithTaskRetry(defaultStatsRetry))
|
||||||
cleanupHandler := &task.SystemCleanupHandler{}
|
ctx.Task().Register(task.StorageMigrationTask, &task.MigrationHandler{}, extpoints.WithTaskMeta(task.StorageMigrationMeta), extpoints.WithTaskRetry(defaultSingleRetry))
|
||||||
ctx.Task().Register(task.SystemCleanupTask, func(c context.Context, payload []byte) error {
|
ctx.Task().Register(task.WarmImageCacheTask, &task.WarmImageCacheHandler{}, extpoints.WithTaskMeta(task.WarmImageCacheMeta), extpoints.WithTaskRetry(1))
|
||||||
_, 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))
|
|
||||||
|
|
||||||
// 4. Register Cron Schedule
|
// 4. Register Cron Schedule
|
||||||
ctx.Schedule().RegisterCron("0 3 * * *", task.SystemCleanupTask, nil)
|
ctx.Schedule().RegisterCron("0 3 * * *", task.SystemCleanupTask, nil)
|
||||||
|
|||||||
@@ -69,10 +69,8 @@ func ResetServices() {
|
|||||||
|
|
||||||
// GetDB resolves the GORM DB instance.
|
// GetDB resolves the GORM DB instance.
|
||||||
func GetDB(ctx context.Context) *gorm.DB {
|
func GetDB(ctx context.Context) *gorm.DB {
|
||||||
if c, ok := ctx.(*core.Context); ok && c != nil {
|
if s, err := core.InjectFrom[contracts.DBService](ctx); err == nil && s != nil {
|
||||||
if s, err := core.Inject[contracts.DBService](c); err == nil && s != nil {
|
return s.DB(ctx)
|
||||||
return s.DB(ctx)
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
svcMu.RLock()
|
svcMu.RLock()
|
||||||
s := dbSvc
|
s := dbSvc
|
||||||
@@ -85,10 +83,8 @@ func GetDB(ctx context.Context) *gorm.DB {
|
|||||||
|
|
||||||
// GetCache resolves the CacheService instance.
|
// GetCache resolves the CacheService instance.
|
||||||
func GetCache(ctx context.Context) contracts.CacheService {
|
func GetCache(ctx context.Context) contracts.CacheService {
|
||||||
if c, ok := ctx.(*core.Context); ok && c != nil {
|
if s, err := core.InjectFrom[contracts.CacheService](ctx); err == nil && s != nil {
|
||||||
if s, err := core.Inject[contracts.CacheService](c); err == nil && s != nil {
|
return s
|
||||||
return s
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
svcMu.RLock()
|
svcMu.RLock()
|
||||||
s := cacheSvc
|
s := cacheSvc
|
||||||
@@ -98,10 +94,8 @@ func GetCache(ctx context.Context) contracts.CacheService {
|
|||||||
|
|
||||||
// GetStorage resolves the StorageService instance.
|
// GetStorage resolves the StorageService instance.
|
||||||
func GetStorage(ctx context.Context) contracts.StorageService {
|
func GetStorage(ctx context.Context) contracts.StorageService {
|
||||||
if c, ok := ctx.(*core.Context); ok && c != nil {
|
if s, err := core.InjectFrom[contracts.StorageService](ctx); err == nil && s != nil {
|
||||||
if s, err := core.Inject[contracts.StorageService](c); err == nil && s != nil {
|
return s
|
||||||
return s
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
svcMu.RLock()
|
svcMu.RLock()
|
||||||
s := storageSvc
|
s := storageSvc
|
||||||
@@ -110,7 +104,10 @@ func GetStorage(ctx context.Context) contracts.StorageService {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// GetTaskService resolves the TaskService instance.
|
// 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()
|
svcMu.RLock()
|
||||||
defer svcMu.RUnlock()
|
defer svcMu.RUnlock()
|
||||||
return taskSvc
|
return taskSvc
|
||||||
@@ -118,10 +115,8 @@ func GetTaskService() contracts.TaskService {
|
|||||||
|
|
||||||
// GetAuthService resolves the AuthService instance.
|
// GetAuthService resolves the AuthService instance.
|
||||||
func GetAuthService(ctx context.Context) contracts.AuthService {
|
func GetAuthService(ctx context.Context) contracts.AuthService {
|
||||||
if c, ok := ctx.(*core.Context); ok && c != nil {
|
if s, err := core.InjectFrom[contracts.AuthService](ctx); err == nil && s != nil {
|
||||||
if s, err := core.Inject[contracts.AuthService](c); err == nil && s != nil {
|
return s
|
||||||
return s
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
svcMu.RLock()
|
svcMu.RLock()
|
||||||
s := authSvc
|
s := authSvc
|
||||||
|
|||||||
@@ -40,4 +40,12 @@ const (
|
|||||||
//nolint:gosec // error message, not hardcoded credentials
|
//nolint:gosec // error message, not hardcoded credentials
|
||||||
errServicePasswordTooShort = "密码长度至少为 8 位"
|
errServicePasswordTooShort = "密码长度至少为 8 位"
|
||||||
errUniqueUsernameFailed = "failed to generate unique username"
|
errUniqueUsernameFailed = "failed to generate unique username"
|
||||||
|
errInvalidEmail = "邮箱地址无效"
|
||||||
|
errInvalidEmailCode = "验证码必须是 6 位数字"
|
||||||
|
errInvalidTaskPayload = "任务参数无效"
|
||||||
|
errMailSubjectRequired = "邮件主题不能为空"
|
||||||
|
errMailBodyRequired = "邮件内容不能为空"
|
||||||
|
errSMTPNotConfigured = "SMTP 未配置"
|
||||||
|
errEmailCacheUnavailable = "缓存服务不可用,无法保存验证码"
|
||||||
|
errSendEmailFailed = "邮件发送失败"
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -11,6 +11,7 @@ import (
|
|||||||
"crypto/rand"
|
"crypto/rand"
|
||||||
"crypto/sha256"
|
"crypto/sha256"
|
||||||
"encoding/hex"
|
"encoding/hex"
|
||||||
|
"encoding/json"
|
||||||
"net/http"
|
"net/http"
|
||||||
"strconv"
|
"strconv"
|
||||||
"sync"
|
"sync"
|
||||||
@@ -190,10 +191,36 @@ func Logout(c *gin.Context) {
|
|||||||
// @Tags user
|
// @Tags user
|
||||||
// @Accept json
|
// @Accept json
|
||||||
// @Produce json
|
// @Produce json
|
||||||
|
// @Param request body user.sendEmailCodeRequest true "目标邮箱"
|
||||||
// @Success 200 {object} response.Any "发送成功"
|
// @Success 200 {object} response.Any "发送成功"
|
||||||
// @Failure 400 {object} response.Any "参数错误"
|
// @Failure 400 {object} response.Any "参数错误"
|
||||||
|
// @Failure 500 {object} response.Any "发送失败"
|
||||||
// @Router /api/v1/user/send-email-code [post]
|
// @Router /api/v1/user/send-email-code [post]
|
||||||
func SendEmailCode(c *gin.Context) {
|
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}))
|
c.JSON(http.StatusOK, response.OK(gin.H{"sent": true}))
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -109,6 +109,10 @@ type registerRequest struct {
|
|||||||
Email string `json:"email"`
|
Email string `json:"email"`
|
||||||
}
|
}
|
||||||
|
|
||||||
|
type sendEmailCodeRequest struct {
|
||||||
|
Email string `json:"email" binding:"required"`
|
||||||
|
}
|
||||||
|
|
||||||
// changePasswordRequest 修改密码请求参数
|
// changePasswordRequest 修改密码请求参数
|
||||||
type changePasswordRequest struct {
|
type changePasswordRequest struct {
|
||||||
OldPassword string `json:"old_password" binding:"required"`
|
OldPassword string `json:"old_password" binding:"required"`
|
||||||
|
|||||||
@@ -9,7 +9,6 @@ import (
|
|||||||
"Wavelet/core/contracts"
|
"Wavelet/core/contracts"
|
||||||
"Wavelet/core/extpoints"
|
"Wavelet/core/extpoints"
|
||||||
"Wavelet/pkg/ginutil"
|
"Wavelet/pkg/ginutil"
|
||||||
"context"
|
|
||||||
"embed"
|
"embed"
|
||||||
"reflect"
|
"reflect"
|
||||||
|
|
||||||
@@ -76,16 +75,13 @@ func (p *Plugin) Manifest() core.Manifest {
|
|||||||
|
|
||||||
// Apply registers user migrations, services, routes, tasks, schedules, and settings into the Context.
|
// Apply registers user migrations, services, routes, tasks, schedules, and settings into the Context.
|
||||||
func (p *Plugin) Apply(ctx *core.Context) error {
|
func (p *Plugin) Apply(ctx *core.Context) error {
|
||||||
// 0. Bind DBService from Context
|
core.Bind[contracts.DBService](ctx, SetDBService)
|
||||||
if db, err := core.Inject[contracts.DBService](ctx); err == nil && db != nil {
|
core.Bind[contracts.CacheService](ctx, SetCacheService)
|
||||||
SetDBService(db)
|
core.Bind[contracts.TaskService](ctx, SetTaskService)
|
||||||
} else {
|
|
||||||
core.When[contracts.DBService](ctx, func(db contracts.DBService) {
|
|
||||||
SetDBService(db)
|
|
||||||
})
|
|
||||||
}
|
|
||||||
ctx.OnDispose(func() error {
|
ctx.OnDispose(func() error {
|
||||||
SetDBService(nil)
|
SetDBService(nil)
|
||||||
|
SetCacheService(nil)
|
||||||
|
SetTaskService(nil)
|
||||||
return nil
|
return nil
|
||||||
})
|
})
|
||||||
|
|
||||||
@@ -101,11 +97,8 @@ func (p *Plugin) Apply(ctx *core.Context) error {
|
|||||||
if mw, ok := authSvc.DisallowTokenAuthMiddleware().(gin.HandlerFunc); ok {
|
if mw, ok := authSvc.DisallowTokenAuthMiddleware().(gin.HandlerFunc); ok {
|
||||||
noTokenMW = mw
|
noTokenMW = mw
|
||||||
}
|
}
|
||||||
} else {
|
|
||||||
core.When[contracts.AuthService](ctx, func(svc contracts.AuthService) {
|
|
||||||
SetAuthService(svc)
|
|
||||||
})
|
|
||||||
}
|
}
|
||||||
|
core.Bind[contracts.AuthService](ctx, SetAuthService)
|
||||||
ctx.OnDispose(func() error {
|
ctx.OnDispose(func() error {
|
||||||
SetAuthService(nil)
|
SetAuthService(nil)
|
||||||
return nil
|
return nil
|
||||||
@@ -155,90 +148,12 @@ func (p *Plugin) Apply(ctx *core.Context) error {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
const (
|
ctx.Task().Register(TaskSendEmailCode, &SendEmailCodeHandler{},
|
||||||
defaultUserTaskRetry = 3
|
extpoints.WithTaskMeta(SendEmailCodeMeta), extpoints.WithTaskRetry(defaultUserTaskRetry))
|
||||||
paramTypeString = "string"
|
ctx.Task().Register(TaskSendMail, &SendMailHandler{},
|
||||||
paramNameEmail = "email"
|
extpoints.WithTaskMeta(SendMailMeta), extpoints.WithTaskRetry(defaultUserTaskRetry))
|
||||||
)
|
ctx.Task().Register(TaskCleanupInactive, &CleanupInactiveHandler{},
|
||||||
|
extpoints.WithTaskMeta(CleanupInactiveMeta))
|
||||||
// 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"),
|
|
||||||
)
|
|
||||||
|
|
||||||
// 5. Register Settings Schemas
|
// 5. Register Settings Schemas
|
||||||
ctx.Settings().Register(extpoints.SettingSchema{
|
ctx.Settings().Register(extpoints.SettingSchema{
|
||||||
|
|||||||
@@ -8,8 +8,10 @@ import (
|
|||||||
"Wavelet/core/contracts"
|
"Wavelet/core/contracts"
|
||||||
"Wavelet/pkg/util"
|
"Wavelet/pkg/util"
|
||||||
"context"
|
"context"
|
||||||
|
"errors"
|
||||||
"strings"
|
"strings"
|
||||||
"sync"
|
"sync"
|
||||||
|
"time"
|
||||||
|
|
||||||
"gorm.io/gorm"
|
"gorm.io/gorm"
|
||||||
)
|
)
|
||||||
@@ -27,10 +29,8 @@ func SetDBService(s contracts.DBService) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func getDB(ctx context.Context) *gorm.DB {
|
func getDB(ctx context.Context) *gorm.DB {
|
||||||
if c, ok := ctx.(*core.Context); ok && c != nil {
|
if s, err := core.InjectFrom[contracts.DBService](ctx); err == nil && s != nil {
|
||||||
if s, err := core.Inject[contracts.DBService](c); err == nil && s != nil {
|
return s.DB(ctx)
|
||||||
return s.DB(ctx)
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
dbMu.RLock()
|
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 获取第一个管理员用户
|
// GetFirstAdminUser 获取第一个管理员用户
|
||||||
func GetFirstAdminUser(ctx context.Context) (*User, error) {
|
func GetFirstAdminUser(ctx context.Context) (*User, error) {
|
||||||
var u User
|
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()
|
p.mu.Unlock()
|
||||||
|
|
||||||
// Bind DBService
|
core.Bind[contracts.DBService](ctx, setDBService)
|
||||||
if db, err := core.Inject[contracts.DBService](ctx); err == nil && db != nil {
|
core.Bind[contracts.TaskService](ctx, setTaskService)
|
||||||
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)
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
ctx.OnDispose(func() error {
|
ctx.OnDispose(func() error {
|
||||||
setDBService(nil)
|
setDBService(nil)
|
||||||
|
|||||||
@@ -34,10 +34,8 @@ func SetRedisClient(c redis.UniversalClient) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func getDB(ctx context.Context) *gorm.DB {
|
func getDB(ctx context.Context) *gorm.DB {
|
||||||
if c, ok := ctx.(*core.Context); ok && c != nil {
|
if s, err := core.InjectFrom[contracts.DBService](ctx); err == nil && s != nil {
|
||||||
if s, err := core.Inject[contracts.DBService](c); err == nil && s != nil {
|
return s.DB(ctx)
|
||||||
return s.DB(ctx)
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
dbMu.RLock()
|
dbMu.RLock()
|
||||||
s := dbSvc
|
s := dbSvc
|
||||||
|
|||||||
@@ -8,6 +8,7 @@ import (
|
|||||||
"Wavelet/core"
|
"Wavelet/core"
|
||||||
"Wavelet/core/contracts"
|
"Wavelet/core/contracts"
|
||||||
"context"
|
"context"
|
||||||
|
"encoding/json"
|
||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"sync"
|
"sync"
|
||||||
@@ -144,14 +145,7 @@ func (p *Plugin) Apply(ctx *core.Context) error {
|
|||||||
ResetAsynqClient()
|
ResetAsynqClient()
|
||||||
p.mu.Unlock()
|
p.mu.Unlock()
|
||||||
|
|
||||||
// 0. Bind DBService
|
core.Bind[contracts.DBService](ctx, setDBService)
|
||||||
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)
|
|
||||||
})
|
|
||||||
}
|
|
||||||
ctx.OnDispose(func() error {
|
ctx.OnDispose(func() error {
|
||||||
setDBService(nil)
|
setDBService(nil)
|
||||||
return nil
|
return nil
|
||||||
@@ -191,12 +185,15 @@ func (p *Plugin) Start(_ context.Context) error {
|
|||||||
mux := asynq.NewServeMux()
|
mux := asynq.NewServeMux()
|
||||||
|
|
||||||
if p.coreCtx != nil && p.coreCtx.Tasks() != nil {
|
if p.coreCtx != nil && p.coreCtx.Tasks() != nil {
|
||||||
|
appCtx := p.coreCtx.Root()
|
||||||
for _, td := range p.coreCtx.Tasks().Tasks() {
|
for _, td := range p.coreCtx.Tasks().Tasks() {
|
||||||
handler, err := toAsynqHandler(td.Pattern, td.Handler)
|
handler, err := toAsynqHandler(td.Pattern, td.Handler)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("driver_asynq_worker: invalid handler for task pattern %q: %w", td.Pattern, err)
|
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)
|
RegisterHandler(pattern, th)
|
||||||
return asynq.HandlerFunc(ProcessTask), nil
|
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)
|
inner, err := toRawAsynqHandler(h)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -294,6 +299,58 @@ func toAsynqHandler(pattern string, h any) (asynq.Handler, error) {
|
|||||||
return asynq.HandlerFunc(ProcessTask), nil
|
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) {
|
func toRawAsynqHandler(h any) (asynq.Handler, error) {
|
||||||
switch fn := h.(type) {
|
switch fn := h.(type) {
|
||||||
case asynq.HandlerFunc:
|
case asynq.HandlerFunc:
|
||||||
|
|||||||
@@ -121,27 +121,13 @@ func (p *Plugin) Apply(ctx *core.Context) error {
|
|||||||
}
|
}
|
||||||
p.mu.Unlock()
|
p.mu.Unlock()
|
||||||
|
|
||||||
// Bind DBService from Context
|
core.Bind[contracts.DBService](ctx, setDBService)
|
||||||
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)
|
|
||||||
})
|
|
||||||
}
|
|
||||||
ctx.OnDispose(func() error {
|
ctx.OnDispose(func() error {
|
||||||
setDBService(nil)
|
setDBService(nil)
|
||||||
return nil
|
return nil
|
||||||
})
|
})
|
||||||
|
|
||||||
// Bind CacheService from Context
|
core.Bind[contracts.CacheService](ctx, setCacheService)
|
||||||
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)
|
|
||||||
})
|
|
||||||
}
|
|
||||||
ctx.OnDispose(func() error {
|
ctx.OnDispose(func() error {
|
||||||
setCacheService(nil)
|
setCacheService(nil)
|
||||||
return nil
|
return nil
|
||||||
@@ -184,30 +170,8 @@ func (p *Plugin) Start(ctx context.Context) error {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// Mount routes collected in Context RouterExtension
|
if err := p.mountContextRoutes(ctx); err != nil {
|
||||||
if p.coreCtx != nil && p.coreCtx.Router() != nil {
|
return err
|
||||||
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...)
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// Mount Swagger in non-production environments
|
// Mount Swagger in non-production environments
|
||||||
@@ -282,6 +246,41 @@ func (p *Plugin) Stop(ctx context.Context) error {
|
|||||||
return err
|
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).
|
// Addr returns the current listening address (or configured address if not yet started).
|
||||||
func (p *Plugin) Addr() string {
|
func (p *Plugin) Addr() string {
|
||||||
p.mu.RLock()
|
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)
|
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.
|
// Stop terminates the in-process cron scheduler.
|
||||||
|
|||||||
@@ -24,10 +24,8 @@ func setDBService(s contracts.DBService) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func getDB(ctx context.Context) *gorm.DB {
|
func getDB(ctx context.Context) *gorm.DB {
|
||||||
if c, ok := ctx.(*core.Context); ok && c != nil {
|
if s, err := core.InjectFrom[contracts.DBService](ctx); err == nil && s != nil {
|
||||||
if s, err := core.Inject[contracts.DBService](c); err == nil && s != nil {
|
return s.DB(ctx)
|
||||||
return s.DB(ctx)
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
dbMu.RLock()
|
dbMu.RLock()
|
||||||
s := dbSvc
|
s := dbSvc
|
||||||
|
|||||||
@@ -4,11 +4,14 @@
|
|||||||
package driver_inproc_worker
|
package driver_inproc_worker
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"Wavelet/core"
|
||||||
|
"Wavelet/core/contracts"
|
||||||
"Wavelet/core/extpoints"
|
"Wavelet/core/extpoints"
|
||||||
"Wavelet/pkg/idgen"
|
"Wavelet/pkg/idgen"
|
||||||
"Wavelet/pkg/logger"
|
"Wavelet/pkg/logger"
|
||||||
"Wavelet/pkg/util"
|
"Wavelet/pkg/util"
|
||||||
"context"
|
"context"
|
||||||
|
"encoding/json"
|
||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"sync"
|
"sync"
|
||||||
@@ -40,6 +43,7 @@ type InprocQueue struct {
|
|||||||
// baseCtx is the app-lifetime context captured at Start; task handlers
|
// baseCtx is the app-lifetime context captured at Start; task handlers
|
||||||
// derive their timeouts from it so shutdown cancellation propagates.
|
// derive their timeouts from it so shutdown cancellation propagates.
|
||||||
baseCtx context.Context
|
baseCtx context.Context
|
||||||
|
appCtx *core.Context
|
||||||
}
|
}
|
||||||
|
|
||||||
// NewInprocQueue creates a new InprocQueue with a given concurrency and queue capacity.
|
// 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)
|
taskCtx, cancel := context.WithTimeout(ctx, timeout)
|
||||||
defer cancel()
|
defer cancel()
|
||||||
|
if q.appCtx != nil {
|
||||||
|
taskCtx = core.WithAppContext(taskCtx, q.appCtx)
|
||||||
|
ctx = core.WithAppContext(ctx, q.appCtx)
|
||||||
|
}
|
||||||
|
|
||||||
q.markRunning(ctx, msg)
|
q.markRunning(ctx, msg)
|
||||||
start := time.Now()
|
start := time.Now()
|
||||||
err := invokeHandler(taskCtx, td.Handler, msg.Payload)
|
result, err := invokeHandler(taskCtx, td.Handler, msg.Payload)
|
||||||
duration := time.Since(start)
|
duration := time.Since(start)
|
||||||
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -210,25 +218,29 @@ func (q *InprocQueue) executeTask(ctx context.Context, msg TaskMessage) {
|
|||||||
}
|
}
|
||||||
return
|
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 {
|
if handler == nil {
|
||||||
return errors.New("nil task handler")
|
return nil, errors.New("nil task handler")
|
||||||
}
|
}
|
||||||
|
|
||||||
switch fn := handler.(type) {
|
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)
|
return fn(ctx, payload)
|
||||||
|
case func(context.Context, []byte) error:
|
||||||
|
return nil, fn(ctx, payload)
|
||||||
case func(context.Context) error:
|
case func(context.Context) error:
|
||||||
return fn(ctx)
|
return nil, fn(ctx)
|
||||||
case func([]byte) error:
|
case func([]byte) error:
|
||||||
return fn(payload)
|
return nil, fn(payload)
|
||||||
case func() error:
|
case func() error:
|
||||||
return fn()
|
return nil, fn()
|
||||||
default:
|
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))
|
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)
|
db := getDB(ctx)
|
||||||
if db == nil {
|
if db == nil {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
now := time.Now()
|
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{
|
updates := map[string]any{
|
||||||
taskExecutionColStatus: taskExecutionStatusSucceeded,
|
taskExecutionColStatus: taskExecutionStatusSucceeded,
|
||||||
"error_message": "",
|
"error_message": "",
|
||||||
"result": "ok",
|
"result": resultText,
|
||||||
"finished_at": now,
|
"finished_at": now,
|
||||||
"duration": duration.Milliseconds(),
|
"duration": duration.Milliseconds(),
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -119,13 +119,7 @@ func (p *Plugin) ConfigEnabled(view core.ConfigView) bool {
|
|||||||
func (p *Plugin) Apply(ctx *core.Context) error {
|
func (p *Plugin) Apply(ctx *core.Context) error {
|
||||||
p.coreCtx = ctx
|
p.coreCtx = ctx
|
||||||
|
|
||||||
if db, err := core.Inject[contracts.DBService](ctx); err == nil && db != nil {
|
core.Bind[contracts.DBService](ctx, setDBService)
|
||||||
setDBService(db)
|
|
||||||
} else {
|
|
||||||
core.When[contracts.DBService](ctx, func(db contracts.DBService) {
|
|
||||||
setDBService(db)
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
taskSvc := newInprocTaskService(ctx.Tasks())
|
taskSvc := newInprocTaskService(ctx.Tasks())
|
||||||
core.Provide[contracts.TaskService](ctx, taskSvc)
|
core.Provide[contracts.TaskService](ctx, taskSvc)
|
||||||
@@ -151,6 +145,9 @@ func (p *Plugin) Start(ctx context.Context) error {
|
|||||||
if p.queue == nil {
|
if p.queue == nil {
|
||||||
p.queue = NewInprocQueue(p.concurrency, p.queueCapacity, p.coreCtx.Tasks())
|
p.queue = NewInprocQueue(p.concurrency, p.queueCapacity, p.coreCtx.Tasks())
|
||||||
}
|
}
|
||||||
|
if p.coreCtx != nil {
|
||||||
|
p.queue.appCtx = p.coreCtx.Root()
|
||||||
|
}
|
||||||
|
|
||||||
globalMu.Lock()
|
globalMu.Lock()
|
||||||
globalQueue = p.queue
|
globalQueue = p.queue
|
||||||
|
|||||||
@@ -80,9 +80,9 @@ func TestInprocWorkerDispatchByTypeTracksExecution(t *testing.T) {
|
|||||||
require.NoError(t, p.Apply(ctx))
|
require.NoError(t, p.Apply(ctx))
|
||||||
|
|
||||||
var executedCount atomic.Int32
|
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)
|
executedCount.Add(1)
|
||||||
return nil
|
return &contracts.TaskResultDTO{Message: "cleaned 3 files"}, nil
|
||||||
},
|
},
|
||||||
extpoints.WithTaskType("system_cleanup"),
|
extpoints.WithTaskType("system_cleanup"),
|
||||||
extpoints.WithTaskName("系统垃圾清理"),
|
extpoints.WithTaskName("系统垃圾清理"),
|
||||||
@@ -111,6 +111,6 @@ func TestInprocWorkerDispatchByTypeTracksExecution(t *testing.T) {
|
|||||||
if listErr != nil || total == 0 || len(execs) == 0 {
|
if listErr != nil || total == 0 || len(execs) == 0 {
|
||||||
return false
|
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")
|
}, 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})
|
core.Provide[contracts.DBService](ctx, &testDBService{db: testDB})
|
||||||
|
|
||||||
var processed atomic.Bool
|
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)
|
processed.Store(true)
|
||||||
return nil
|
return &contracts.TaskResultDTO{Message: "cleaned 3 files"}, nil
|
||||||
},
|
},
|
||||||
extpoints.WithTaskType("system_cleanup"),
|
extpoints.WithTaskType("system_cleanup"),
|
||||||
extpoints.WithTaskName("系统垃圾清理"),
|
extpoints.WithTaskName("系统垃圾清理"),
|
||||||
@@ -277,7 +277,7 @@ func TestAsynqWorkerDispatchTracksExecution(t *testing.T) {
|
|||||||
if listErr != nil || len(execs) == 0 {
|
if listErr != nil || len(execs) == 0 {
|
||||||
return false
|
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")
|
}, 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.
|
// Apply mounts the storage service into the Context.
|
||||||
func (p *Plugin) Apply(ctx *core.Context) error {
|
func (p *Plugin) Apply(ctx *core.Context) error {
|
||||||
// Bind DBService
|
core.Bind[contracts.DBService](ctx, func(db contracts.DBService) {
|
||||||
if db, err := core.Inject[contracts.DBService](ctx); err == nil && db != nil {
|
|
||||||
objectstore.SetDBService(db)
|
objectstore.SetDBService(db)
|
||||||
diskcache.SetDBService(db)
|
diskcache.SetDBService(db)
|
||||||
} else {
|
})
|
||||||
core.When[contracts.DBService](ctx, func(db contracts.DBService) {
|
core.Bind[contracts.CacheService](ctx, objectstore.SetCacheService)
|
||||||
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)
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
ctx.OnDispose(func() error {
|
ctx.OnDispose(func() error {
|
||||||
objectstore.SetDBService(nil)
|
objectstore.SetDBService(nil)
|
||||||
|
|||||||
Reference in New Issue
Block a user