mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-07 08:06:37 +08:00
fix(persistence): migrate all pkg/persistence imports to plugins/infra/database and plugins/infra/cache
- Replace db.DB(ctx) with database.DB(ctx) from plugins/infra/database
- Replace db.Redis/db.PrefixedKey/db.GetJSON/db.SetJSON with cachepkg.* from plugins/infra/cache
- Replace pkg/persistence/idgen with pkg/idgen (already exists)
- Replace pkg/persistence/batchwriter with pkg/batchwriter (already exists)
- Replace pkg/persistence/migrator with pkg/migrator (already exists)
- Replace pkg/persistence/logstore with plugins/domain/risk_control/logstore
- Delete defunct pkg/{persistence,cap,message_gateway,push,shared,task}
- Fix vet issues: db alias in domain_test.go, driver_asynq_worker.TaskHandler reference
- Update Makefile architecture guard
- Update docs and skill references
- Update go.mod: gorilla/sessions promotion to direct dependency
This commit is contained in:
@@ -7,16 +7,16 @@ import (
|
||||
"net/http"
|
||||
"strconv"
|
||||
|
||||
persistence "github.com/Rain-kl/Wavelet/pkg/persistence"
|
||||
"github.com/Rain-kl/Wavelet/pkg/response"
|
||||
"github.com/Rain-kl/Wavelet/plugins/domain/auth"
|
||||
"github.com/Rain-kl/Wavelet/plugins/infra/database"
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
// ListAuthSources lists all configured authentication sources.
|
||||
func ListAuthSources(c *gin.Context) {
|
||||
var sources []auth.AuthSource
|
||||
gormDB := persistence.DB(c.Request.Context())
|
||||
gormDB := database.DB(c.Request.Context())
|
||||
if err := gormDB.Order("id ASC").Find(&sources).Error; err != nil {
|
||||
response.AbortInternal(c, "获取认证源列表失败")
|
||||
return
|
||||
@@ -51,7 +51,7 @@ func CreateAuthSource(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
|
||||
gormDB := persistence.DB(c.Request.Context())
|
||||
gormDB := database.DB(c.Request.Context())
|
||||
if err := gormDB.Create(&source).Error; err != nil {
|
||||
response.AbortBadRequest(c, "创建认证源失败: "+err.Error())
|
||||
return
|
||||
@@ -70,7 +70,7 @@ func UpdateAuthSource(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
|
||||
gormDB := persistence.DB(c.Request.Context())
|
||||
gormDB := database.DB(c.Request.Context())
|
||||
var existing auth.AuthSource
|
||||
if err := gormDB.First(&existing, id).Error; err != nil {
|
||||
response.AbortNotFound(c, "认证源不存在")
|
||||
@@ -115,7 +115,7 @@ func ToggleAuthSource(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
|
||||
gormDB := persistence.DB(c.Request.Context())
|
||||
gormDB := database.DB(c.Request.Context())
|
||||
var existing auth.AuthSource
|
||||
if err := gormDB.First(&existing, id).Error; err != nil {
|
||||
response.AbortNotFound(c, "认证源不存在")
|
||||
@@ -147,7 +147,7 @@ func DeleteAuthSource(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
|
||||
gormDB := persistence.DB(c.Request.Context())
|
||||
gormDB := database.DB(c.Request.Context())
|
||||
if err := gormDB.Delete(&auth.AuthSource{}, id).Error; err != nil {
|
||||
response.AbortInternal(c, "删除认证源失败")
|
||||
return
|
||||
|
||||
@@ -15,9 +15,10 @@ import (
|
||||
|
||||
"github.com/Rain-kl/Wavelet/pkg/logger"
|
||||
mail "github.com/Rain-kl/Wavelet/pkg/mail"
|
||||
db "github.com/Rain-kl/Wavelet/pkg/persistence"
|
||||
"github.com/Rain-kl/Wavelet/pkg/response"
|
||||
"github.com/Rain-kl/Wavelet/plugins/domain/cap"
|
||||
cachepkg "github.com/Rain-kl/Wavelet/plugins/infra/cache"
|
||||
db "github.com/Rain-kl/Wavelet/plugins/infra/database"
|
||||
"github.com/Rain-kl/Wavelet/plugins/infra/storage/objectstore"
|
||||
"github.com/gin-gonic/gin"
|
||||
"gorm.io/gorm"
|
||||
@@ -350,15 +351,15 @@ func invalidateCachesAfterConfigUpdate(ctx context.Context, key string) {
|
||||
invalidateSystemConfigCaches(ctx, key)
|
||||
|
||||
if key == ConfigKeyStorageConfig {
|
||||
if db.Redis != nil {
|
||||
_ = db.Redis.Publish(ctx, "upload:access_cache:invalidate", "reset").Err()
|
||||
if cachepkg.Redis != nil {
|
||||
_ = cachepkg.Redis.Publish(ctx, "upload:access_cache:invalidate", "reset").Err()
|
||||
}
|
||||
objectstore.ResetCache()
|
||||
objectstore.PublishCacheInvalidation(ctx)
|
||||
}
|
||||
if key == ConfigKeyFileAccessWhitelist {
|
||||
if db.Redis != nil {
|
||||
_ = db.Redis.Publish(ctx, "upload:access_cache:invalidate", "reset").Err()
|
||||
if cachepkg.Redis != nil {
|
||||
_ = cachepkg.Redis.Publish(ctx, "upload:access_cache:invalidate", "reset").Err()
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -20,8 +20,8 @@ import (
|
||||
"gorm.io/gorm"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/pkg/config"
|
||||
db "github.com/Rain-kl/Wavelet/pkg/persistence"
|
||||
"github.com/Rain-kl/Wavelet/pkg/response"
|
||||
db "github.com/Rain-kl/Wavelet/plugins/infra/database"
|
||||
)
|
||||
|
||||
const (
|
||||
|
||||
@@ -17,12 +17,12 @@ import (
|
||||
|
||||
"github.com/Rain-kl/Wavelet/pkg/config"
|
||||
"github.com/Rain-kl/Wavelet/pkg/logger"
|
||||
db "github.com/Rain-kl/Wavelet/pkg/persistence"
|
||||
"github.com/Rain-kl/Wavelet/pkg/persistence/logstore"
|
||||
"github.com/Rain-kl/Wavelet/pkg/response"
|
||||
"github.com/Rain-kl/Wavelet/pkg/task"
|
||||
"github.com/Rain-kl/Wavelet/pkg/util"
|
||||
"github.com/Rain-kl/Wavelet/plugins/domain/risk_control"
|
||||
"github.com/Rain-kl/Wavelet/plugins/domain/risk_control/logstore"
|
||||
db "github.com/Rain-kl/Wavelet/plugins/infra/database"
|
||||
"github.com/Rain-kl/Wavelet/plugins/drivers/driver_asynq_worker"
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/gorilla/websocket"
|
||||
)
|
||||
@@ -497,16 +497,16 @@ const (
|
||||
)
|
||||
|
||||
// LogDBSwitchMeta 描述切换日志数据库任务。
|
||||
var LogDBSwitchMeta = task.TaskMeta{
|
||||
var LogDBSwitchMeta = driver_asynq_worker.TaskMeta{
|
||||
Type: TaskTypeLogDBSwitch,
|
||||
AsynqTask: LogDBSwitchTask,
|
||||
Name: "切换日志数据库",
|
||||
Description: "复制迁移用户访问日志并在成功后切换日志主库(期间禁止日志写入)",
|
||||
SupportsTime: false,
|
||||
MaxRetry: task.DefaultMaxRetry,
|
||||
Queue: task.QueueDefault,
|
||||
MaxRetry: driver_asynq_worker.DefaultMaxRetry,
|
||||
Queue: driver_asynq_worker.QueueDefault,
|
||||
Retryable: true,
|
||||
Params: []task.TaskParam{
|
||||
Params: []driver_asynq_worker.TaskParam{
|
||||
{Name: "target", Label: "目标日志库", Type: "string", Required: true,
|
||||
Placeholder: "postgres|sqlite|clickhouse", Description: "迁移目标:postgres(主库为 PG 时)、sqlite(主库为 SQLite 时)或 clickhouse"},
|
||||
},
|
||||
@@ -553,7 +553,7 @@ func validTarget(v string) bool {
|
||||
}
|
||||
|
||||
// Execute 执行迁移。
|
||||
func (h *LogDBSwitchHandler) Execute(ctx context.Context, payload []byte) (*task.TaskResult, error) {
|
||||
func (h *LogDBSwitchHandler) Execute(ctx context.Context, payload []byte) (*driver_asynq_worker.TaskResult, error) {
|
||||
var p logDBSwitchPayload
|
||||
if err := json.Unmarshal(payload, &p); err != nil {
|
||||
return nil, fmt.Errorf("参数解析失败: %w", err)
|
||||
@@ -565,10 +565,10 @@ func (h *LogDBSwitchHandler) Execute(ctx context.Context, payload []byte) (*task
|
||||
|
||||
source, err := currentLogDatabase(ctx)
|
||||
if err != nil {
|
||||
task.AppendLog(ctx, "读取日志主库失败: %v", err)
|
||||
driver_asynq_worker.AppendLog(ctx, "读取日志主库失败: %v", err)
|
||||
return nil, err
|
||||
}
|
||||
task.AppendLog(ctx, "开始切换日志数据库:%s -> %s", source, p.Target)
|
||||
driver_asynq_worker.AppendLog(ctx, "开始切换日志数据库:%s -> %s", source, p.Target)
|
||||
|
||||
if err := setMigrationFlag(ctx, "migrating"); err != nil {
|
||||
return nil, err
|
||||
@@ -612,8 +612,8 @@ func (h *LogDBSwitchHandler) Execute(ctx context.Context, payload []byte) (*task
|
||||
return nil, err
|
||||
}
|
||||
logstore.InvalidateCache()
|
||||
task.AppendLog(ctx, "日志数据库已切换为 %s,写入恢复", p.Target)
|
||||
return &task.TaskResult{Message: fmt.Sprintf("日志数据库已从 %s 切换为 %s", source, p.Target)}, nil
|
||||
driver_asynq_worker.AppendLog(ctx, "日志数据库已切换为 %s,写入恢复", p.Target)
|
||||
return &driver_asynq_worker.TaskResult{Message: fmt.Sprintf("日志数据库已从 %s 切换为 %s", source, p.Target)}, nil
|
||||
}
|
||||
|
||||
func validateSwitch(ctx context.Context, target string) error {
|
||||
@@ -676,7 +676,7 @@ func copyUserAccessLogs(ctx context.Context, src, dst *logstore.Store) error {
|
||||
}
|
||||
afterID = rows[len(rows)-1].ID
|
||||
copied += len(rows)
|
||||
task.AppendLog(ctx, "已复制用户访问日志 %d 条", copied)
|
||||
driver_asynq_worker.AppendLog(ctx, "已复制用户访问日志 %d 条", copied)
|
||||
if len(rows) < copyBatchSize {
|
||||
break
|
||||
}
|
||||
|
||||
@@ -18,8 +18,8 @@ import (
|
||||
|
||||
"github.com/Rain-kl/Wavelet/pkg/config"
|
||||
"github.com/Rain-kl/Wavelet/pkg/logger"
|
||||
"github.com/Rain-kl/Wavelet/pkg/persistence/logstore"
|
||||
"github.com/Rain-kl/Wavelet/pkg/response"
|
||||
"github.com/Rain-kl/Wavelet/plugins/domain/risk_control/logstore"
|
||||
)
|
||||
|
||||
var startTime = time.Now()
|
||||
|
||||
@@ -16,8 +16,8 @@ import (
|
||||
|
||||
"github.com/Rain-kl/Wavelet/pkg/logger"
|
||||
"github.com/Rain-kl/Wavelet/pkg/response"
|
||||
"github.com/Rain-kl/Wavelet/pkg/task"
|
||||
"github.com/Rain-kl/Wavelet/pkg/task/scheduler"
|
||||
"github.com/Rain-kl/Wavelet/plugins/drivers/driver_asynq_cron"
|
||||
"github.com/Rain-kl/Wavelet/plugins/drivers/driver_asynq_worker"
|
||||
)
|
||||
|
||||
// ListTaskTypes 获取支持的任务类型列表
|
||||
@@ -26,12 +26,12 @@ import (
|
||||
// @Tags admin
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Success 200 {object} response.Any{data=[]task.TaskMeta} "任务类型列表"
|
||||
// @Success 200 {object} response.Any{data=[]driver_asynq_worker.TaskMeta} "任务类型列表"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 403 {object} response.Any "无管理员权限"
|
||||
// @Router /api/v1/admin/tasks/types [get]
|
||||
func ListTaskTypes(c *gin.Context) {
|
||||
c.JSON(http.StatusOK, response.OK(task.GetDispatchableTasks()))
|
||||
c.JSON(http.StatusOK, response.OK(driver_asynq_worker.GetDispatchableTasks()))
|
||||
}
|
||||
|
||||
// DispatchTaskRequest 下发任务请求
|
||||
@@ -64,7 +64,7 @@ func DispatchTask(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
|
||||
meta := task.GetTaskMeta(req.TaskType)
|
||||
meta := driver_asynq_worker.GetTaskMeta(req.TaskType)
|
||||
if meta == nil {
|
||||
response.AbortBadRequest(c, InvalidTaskType)
|
||||
return
|
||||
@@ -75,13 +75,13 @@ func DispatchTask(c *gin.Context) {
|
||||
payloadBytes = []byte(req.Payload)
|
||||
}
|
||||
|
||||
validated, err := task.ValidateAndNormalizePayload(meta.AsynqTask, payloadBytes)
|
||||
validated, err := driver_asynq_worker.ValidateAndNormalizePayload(meta.AsynqTask, payloadBytes)
|
||||
if err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
taskID, err := task.DispatchTask(c.Request.Context(), req.TaskType, validated, "manual")
|
||||
taskID, err := driver_asynq_worker.DispatchTask(c.Request.Context(), req.TaskType, validated, "manual")
|
||||
if err != nil {
|
||||
response.AbortInternal(c, fmt.Sprintf("%s: %v", TaskDispatchFailed, err))
|
||||
return
|
||||
@@ -112,7 +112,7 @@ func ListTaskExecutions(c *gin.Context) {
|
||||
}
|
||||
|
||||
if req.TaskType != "" {
|
||||
if meta := task.GetTaskMeta(req.TaskType); meta != nil {
|
||||
if meta := driver_asynq_worker.GetTaskMeta(req.TaskType); meta != nil {
|
||||
req.TaskType = meta.AsynqTask
|
||||
}
|
||||
}
|
||||
@@ -181,7 +181,7 @@ func RetryTask(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
|
||||
newTaskID, err := task.RetryTask(c.Request.Context(), id)
|
||||
newTaskID, err := driver_asynq_worker.RetryTask(c.Request.Context(), id)
|
||||
if err != nil {
|
||||
errMsg := err.Error()
|
||||
switch {
|
||||
@@ -254,7 +254,7 @@ func CreateSchedule(c *gin.Context) {
|
||||
}
|
||||
|
||||
// 校验关联的异步任务类型
|
||||
meta := task.GetTaskMeta(req.TaskType)
|
||||
meta := driver_asynq_worker.GetTaskMeta(req.TaskType)
|
||||
if meta == nil {
|
||||
response.AbortBadRequest(c, InvalidTaskType)
|
||||
return
|
||||
@@ -265,7 +265,7 @@ func CreateSchedule(c *gin.Context) {
|
||||
if strings.TrimSpace(req.Payload) != "" {
|
||||
payloadBytes = []byte(req.Payload)
|
||||
}
|
||||
validated, err := task.ValidateAndNormalizePayload(meta.AsynqTask, payloadBytes)
|
||||
validated, err := driver_asynq_worker.ValidateAndNormalizePayload(meta.AsynqTask, payloadBytes)
|
||||
if err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
@@ -285,7 +285,7 @@ func CreateSchedule(c *gin.Context) {
|
||||
}
|
||||
|
||||
// 触发调度服务重载
|
||||
if err := scheduler.ReloadScheduler(); err != nil {
|
||||
if err := driver_asynq_cron.ReloadScheduler(); err != nil {
|
||||
logger.ErrorF(c.Request.Context(), "[TaskAdmin] 重载调度器失败: %v", err)
|
||||
}
|
||||
|
||||
@@ -343,7 +343,7 @@ func UpdateSchedule(c *gin.Context) {
|
||||
}
|
||||
|
||||
// 校验关联的异步任务类型
|
||||
meta := task.GetTaskMeta(req.TaskType)
|
||||
meta := driver_asynq_worker.GetTaskMeta(req.TaskType)
|
||||
if meta == nil {
|
||||
response.AbortBadRequest(c, InvalidTaskType)
|
||||
return
|
||||
@@ -354,7 +354,7 @@ func UpdateSchedule(c *gin.Context) {
|
||||
if strings.TrimSpace(req.Payload) != "" {
|
||||
payloadBytes = []byte(req.Payload)
|
||||
}
|
||||
validated, err := task.ValidateAndNormalizePayload(meta.AsynqTask, payloadBytes)
|
||||
validated, err := driver_asynq_worker.ValidateAndNormalizePayload(meta.AsynqTask, payloadBytes)
|
||||
if err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
@@ -372,7 +372,7 @@ func UpdateSchedule(c *gin.Context) {
|
||||
}
|
||||
|
||||
// 触发调度服务重载
|
||||
if err := scheduler.ReloadScheduler(); err != nil {
|
||||
if err := driver_asynq_cron.ReloadScheduler(); err != nil {
|
||||
logger.ErrorF(c.Request.Context(), "[TaskAdmin] 重载调度器失败: %v", err)
|
||||
}
|
||||
|
||||
@@ -405,7 +405,7 @@ func DeleteSchedule(c *gin.Context) {
|
||||
}
|
||||
|
||||
// 触发调度服务重载
|
||||
if err := scheduler.ReloadScheduler(); err != nil {
|
||||
if err := driver_asynq_cron.ReloadScheduler(); err != nil {
|
||||
logger.ErrorF(c.Request.Context(), "[TaskAdmin] 重载调度器失败: %v", err)
|
||||
}
|
||||
|
||||
|
||||
@@ -16,12 +16,12 @@ import (
|
||||
"gorm.io/gorm"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/core/contracts"
|
||||
"github.com/Rain-kl/Wavelet/pkg/idgen"
|
||||
"github.com/Rain-kl/Wavelet/pkg/logger"
|
||||
db "github.com/Rain-kl/Wavelet/pkg/persistence"
|
||||
"github.com/Rain-kl/Wavelet/pkg/persistence/idgen"
|
||||
"github.com/Rain-kl/Wavelet/pkg/response"
|
||||
"github.com/Rain-kl/Wavelet/pkg/util"
|
||||
"github.com/Rain-kl/Wavelet/plugins/domain/auth"
|
||||
db "github.com/Rain-kl/Wavelet/plugins/infra/database"
|
||||
)
|
||||
|
||||
const minPasswordLength = 8
|
||||
|
||||
@@ -17,9 +17,10 @@ import (
|
||||
"gorm.io/gorm"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/pkg/cache/ram"
|
||||
db "github.com/Rain-kl/Wavelet/pkg/persistence"
|
||||
"github.com/Rain-kl/Wavelet/pkg/persistence/idgen"
|
||||
"github.com/Rain-kl/Wavelet/pkg/idgen"
|
||||
"github.com/Rain-kl/Wavelet/pkg/util"
|
||||
cachepkg "github.com/Rain-kl/Wavelet/plugins/infra/cache"
|
||||
db "github.com/Rain-kl/Wavelet/plugins/infra/database"
|
||||
)
|
||||
|
||||
const (
|
||||
@@ -482,7 +483,7 @@ func GetLatestTaskExecutionByTaskType(ctx context.Context, taskType string) (*Ta
|
||||
|
||||
// AppendTaskExecutionLog 将日志追加到 Redis 缓冲,任务完成后再持久化到数据库。
|
||||
func AppendTaskExecutionLog(ctx context.Context, taskID string, logLine string) error {
|
||||
if db.Redis == nil {
|
||||
if cachepkg.Redis == nil {
|
||||
return errors.New("redis client is not initialized")
|
||||
}
|
||||
|
||||
@@ -490,7 +491,7 @@ func AppendTaskExecutionLog(ctx context.Context, taskID string, logLine string)
|
||||
line := fmt.Sprintf("[%s] %s\n", now, logLine)
|
||||
key := taskExecutionLogRedisKey(taskID)
|
||||
|
||||
_, err := db.Redis.TxPipelined(ctx, func(pipe redis.Pipeliner) error {
|
||||
_, err := cachepkg.Redis.TxPipelined(ctx, func(pipe redis.Pipeliner) error {
|
||||
pipe.RPush(ctx, key, line)
|
||||
pipe.LTrim(ctx, key, -taskExecutionLogMaxLines, -1)
|
||||
pipe.Expire(ctx, key, taskExecutionLogExpiration)
|
||||
@@ -504,12 +505,12 @@ func AppendTaskExecutionLog(ctx context.Context, taskID string, logLine string)
|
||||
|
||||
// FlushTaskExecutionLog 将 Redis 中的完整任务日志写入数据库,并在成功后清理缓存。
|
||||
func FlushTaskExecutionLog(ctx context.Context, taskID string) error {
|
||||
if db.Redis == nil {
|
||||
if cachepkg.Redis == nil {
|
||||
return errors.New("redis client is not initialized")
|
||||
}
|
||||
|
||||
key := taskExecutionLogRedisKey(taskID)
|
||||
logLines, err := db.Redis.LRange(ctx, key, 0, -1).Result()
|
||||
logLines, err := cachepkg.Redis.LRange(ctx, key, 0, -1).Result()
|
||||
if err != nil {
|
||||
return fmt.Errorf("get task execution log from redis: %w", err)
|
||||
}
|
||||
@@ -528,7 +529,7 @@ func FlushTaskExecutionLog(ctx context.Context, taskID string) error {
|
||||
return fmt.Errorf("persist task execution log: task %q not found", taskID)
|
||||
}
|
||||
|
||||
if err := db.Redis.Del(ctx, key).Err(); err != nil {
|
||||
if err := cachepkg.Redis.Del(ctx, key).Err(); err != nil {
|
||||
return fmt.Errorf("delete persisted task execution log from redis: %w", err)
|
||||
}
|
||||
return nil
|
||||
@@ -658,15 +659,15 @@ func CleanupTaskExecutionLogs(ctx context.Context, now time.Time) (TaskExecution
|
||||
}
|
||||
|
||||
func taskExecutionLogRedisKey(taskID string) string {
|
||||
return db.PrefixedKey(taskExecutionLogRedisKeyPrefix + taskID)
|
||||
return cachepkg.PrefixedKey(taskExecutionLogRedisKeyPrefix + taskID)
|
||||
}
|
||||
|
||||
func loadTaskExecutionLog(ctx context.Context, execution *TaskExecution) error {
|
||||
if db.Redis == nil {
|
||||
if cachepkg.Redis == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
logLines, err := db.Redis.LRange(ctx, taskExecutionLogRedisKey(execution.TaskID), 0, -1).Result()
|
||||
logLines, err := cachepkg.Redis.LRange(ctx, taskExecutionLogRedisKey(execution.TaskID), 0, -1).Result()
|
||||
if err != nil {
|
||||
return fmt.Errorf("get task execution log from redis: %w", err)
|
||||
}
|
||||
@@ -679,12 +680,12 @@ func loadTaskExecutionLog(ctx context.Context, execution *TaskExecution) error {
|
||||
}
|
||||
|
||||
func loadTaskExecutionLogs(ctx context.Context, executions []TaskExecution) error {
|
||||
if db.Redis == nil || len(executions) == 0 {
|
||||
if cachepkg.Redis == nil || len(executions) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
commands := make([]*redis.StringSliceCmd, len(executions))
|
||||
_, err := db.Redis.Pipelined(ctx, func(pipe redis.Pipeliner) error {
|
||||
_, err := cachepkg.Redis.Pipelined(ctx, func(pipe redis.Pipeliner) error {
|
||||
for i := range executions {
|
||||
commands[i] = pipe.LRange(ctx, taskExecutionLogRedisKey(executions[i].TaskID), 0, -1)
|
||||
}
|
||||
|
||||
@@ -13,8 +13,8 @@ import (
|
||||
"gorm.io/gorm"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/pkg/cache/ram"
|
||||
db "github.com/Rain-kl/Wavelet/pkg/persistence"
|
||||
"github.com/Rain-kl/Wavelet/pkg/util"
|
||||
cachepkg "github.com/Rain-kl/Wavelet/plugins/infra/cache"
|
||||
)
|
||||
|
||||
const (
|
||||
@@ -104,14 +104,14 @@ func ensureSystemConfigCacheListener() {
|
||||
}
|
||||
|
||||
func startSystemConfigCacheInvalidationListener() {
|
||||
if db.Redis == nil {
|
||||
if cachepkg.Redis == nil {
|
||||
return
|
||||
}
|
||||
|
||||
systemConfigListenerCtx, systemConfigListenerCancel = context.WithCancel(context.Background())
|
||||
systemConfigListenerDone = make(chan struct{})
|
||||
|
||||
redisClient := db.Redis // 捕获当前客户端:goroutine 不读可变全局,避免与测试置空 db.Redis 竞争
|
||||
redisClient := cachepkg.Redis // 捕获当前客户端:goroutine 不读可变全局,避免与测试置空 cachepkg.Redis 竞争
|
||||
util.Go(func() {
|
||||
listenerCtx := systemConfigListenerCtx
|
||||
defer close(systemConfigListenerDone)
|
||||
@@ -169,8 +169,8 @@ func InvalidateSystemConfigCache(ctx context.Context, key string) error {
|
||||
ram.Delete(ConfigCacheType, key)
|
||||
|
||||
// Broadcast to other nodes and clean legacy Redis cache key
|
||||
if db.Redis != nil {
|
||||
_ = db.HDel(ctx, SystemConfigRedisHashKey, key)
|
||||
if cachepkg.Redis != nil {
|
||||
_ = cachepkg.HDel(ctx, SystemConfigRedisHashKey, key)
|
||||
publishSystemConfigBroadcast(ctx, ConfigCacheType, key)
|
||||
}
|
||||
return nil
|
||||
@@ -184,22 +184,22 @@ func InvalidateAllSystemConfigCaches(ctx context.Context) error {
|
||||
ram.UpdateTypeItems(ConfigCacheType, nil)
|
||||
|
||||
// Broadcast to other nodes and clean legacy Redis cache keys
|
||||
if db.Redis != nil {
|
||||
_ = db.Redis.Del(ctx, db.PrefixedKey(SystemConfigRedisHashKey), db.PrefixedKey(SystemConfigVisibleListRedisKey)).Err()
|
||||
if cachepkg.Redis != nil {
|
||||
_ = cachepkg.Redis.Del(ctx, cachepkg.PrefixedKey(SystemConfigRedisHashKey), cachepkg.PrefixedKey(SystemConfigVisibleListRedisKey)).Err()
|
||||
publishSystemConfigBroadcast(ctx, ConfigCacheType, "*")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func publishSystemConfigBroadcast(ctx context.Context, configType string, key string) {
|
||||
if db.Redis == nil {
|
||||
if cachepkg.Redis == nil {
|
||||
return
|
||||
}
|
||||
payload, err := json.Marshal(systemConfigBroadcastMessage{Type: configType, Key: key})
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
_ = db.Redis.Publish(ctx, SystemConfigBroadcastChannel, payload).Err()
|
||||
_ = cachepkg.Redis.Publish(ctx, SystemConfigBroadcastChannel, payload).Err()
|
||||
}
|
||||
|
||||
// ResetSystemConfigRAMCacheForTest clears only the process-local RAM cache.
|
||||
|
||||
@@ -14,7 +14,8 @@ import (
|
||||
"github.com/redis/go-redis/v9/maintnotifications"
|
||||
"gorm.io/gorm"
|
||||
|
||||
db "github.com/Rain-kl/Wavelet/pkg/persistence"
|
||||
"github.com/Rain-kl/Wavelet/plugins/infra/cache"
|
||||
"github.com/Rain-kl/Wavelet/plugins/infra/database"
|
||||
)
|
||||
|
||||
func setupSystemConfigTest(t *testing.T) (*gorm.DB, func()) {
|
||||
@@ -51,15 +52,15 @@ func setupSystemConfigTest(t *testing.T) (*gorm.DB, func()) {
|
||||
},
|
||||
})
|
||||
|
||||
previousRedis := db.Redis
|
||||
db.SetDB(sqliteDB)
|
||||
db.Redis = redisClient
|
||||
previousRedis := cache.Redis
|
||||
database.SetDB(sqliteDB)
|
||||
cache.Redis = redisClient
|
||||
|
||||
cleanup := func() {
|
||||
StopSystemConfigCacheListener()
|
||||
ResetSystemConfigRAMCacheForTest()
|
||||
db.SetDB(nil)
|
||||
db.Redis = previousRedis
|
||||
database.SetDB(nil)
|
||||
cache.Redis = previousRedis
|
||||
_ = redisClient.Close()
|
||||
mr.Close()
|
||||
}
|
||||
|
||||
@@ -11,8 +11,7 @@ import (
|
||||
"strings"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/core/contracts"
|
||||
|
||||
db "github.com/Rain-kl/Wavelet/pkg/persistence"
|
||||
db "github.com/Rain-kl/Wavelet/plugins/infra/database"
|
||||
|
||||
"github.com/coreos/go-oidc/v3/oidc"
|
||||
"golang.org/x/oauth2"
|
||||
|
||||
@@ -12,8 +12,8 @@ import (
|
||||
|
||||
"github.com/Rain-kl/Wavelet/core/contracts"
|
||||
"github.com/Rain-kl/Wavelet/pkg/cache/ram"
|
||||
db "github.com/Rain-kl/Wavelet/pkg/persistence"
|
||||
"github.com/Rain-kl/Wavelet/pkg/util"
|
||||
db "github.com/Rain-kl/Wavelet/plugins/infra/cache"
|
||||
)
|
||||
|
||||
const (
|
||||
|
||||
@@ -12,8 +12,8 @@ import (
|
||||
"github.com/redis/go-redis/v9/maintnotifications"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/core/contracts"
|
||||
db "github.com/Rain-kl/Wavelet/pkg/persistence"
|
||||
"github.com/Rain-kl/Wavelet/plugins/domain/auth"
|
||||
db "github.com/Rain-kl/Wavelet/plugins/infra/cache"
|
||||
)
|
||||
|
||||
func setupOauthCacheTest(t *testing.T) (*miniredis.Miniredis, func()) {
|
||||
|
||||
@@ -34,4 +34,6 @@ const (
|
||||
errAdminRequired = "无权访问"
|
||||
//nolint:gosec // error message, not hardcoded credentials
|
||||
errTokenAdminRequired = "令牌无管理员权限"
|
||||
errBannedAccount = "账号已被封禁"
|
||||
errUnAuthorized = "未登录"
|
||||
)
|
||||
|
||||
@@ -15,14 +15,12 @@ import (
|
||||
|
||||
"github.com/Rain-kl/Wavelet/core/contracts"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/pkg/idgen"
|
||||
"github.com/Rain-kl/Wavelet/pkg/logger"
|
||||
|
||||
db "github.com/Rain-kl/Wavelet/pkg/persistence"
|
||||
"github.com/Rain-kl/Wavelet/pkg/persistence/idgen"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/pkg/response"
|
||||
"github.com/Rain-kl/Wavelet/pkg/shared"
|
||||
"github.com/Rain-kl/Wavelet/pkg/util"
|
||||
cachepkg "github.com/Rain-kl/Wavelet/plugins/infra/cache"
|
||||
db "github.com/Rain-kl/Wavelet/plugins/infra/database"
|
||||
"github.com/coreos/go-oidc/v3/oidc"
|
||||
"github.com/gin-contrib/sessions"
|
||||
"github.com/gin-gonic/gin"
|
||||
@@ -81,7 +79,7 @@ func GetLoginURL(c *gin.Context) {
|
||||
response.AbortInternal(c, err.Error())
|
||||
return
|
||||
}
|
||||
if err := db.Redis.Set(ctx, db.PrefixedKey(fmt.Sprintf(OAuthStateCacheKeyFormat, state)), payloadValue, OAuthStateCacheKeyExpiration).Err(); err != nil {
|
||||
if err := cachepkg.Redis.Set(ctx, cachepkg.PrefixedKey(fmt.Sprintf(OAuthStateCacheKeyFormat, state)), payloadValue, OAuthStateCacheKeyExpiration).Err(); err != nil {
|
||||
response.AbortInternal(c, err.Error())
|
||||
return
|
||||
}
|
||||
@@ -110,16 +108,16 @@ func buildAuthorizeURL(ctx context.Context, source *AuthSource, state string) (s
|
||||
}
|
||||
|
||||
func reserveOAuthStateSlot(ctx context.Context, sessionHash string) error {
|
||||
if db.Redis == nil || sessionHash == "" {
|
||||
if cachepkg.Redis == nil || sessionHash == "" {
|
||||
return nil
|
||||
}
|
||||
key := db.PrefixedKey(fmt.Sprintf(oauthStateLimitKeyFormat, sessionHash))
|
||||
n, err := db.Redis.Incr(ctx, key).Result()
|
||||
key := cachepkg.PrefixedKey(fmt.Sprintf(oauthStateLimitKeyFormat, sessionHash))
|
||||
n, err := cachepkg.Redis.Incr(ctx, key).Result()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if n == 1 {
|
||||
_ = db.Redis.Expire(ctx, key, OAuthStateCacheKeyExpiration).Err()
|
||||
_ = cachepkg.Redis.Expire(ctx, key, OAuthStateCacheKeyExpiration).Err()
|
||||
}
|
||||
if n > oauthStateLimitMax {
|
||||
return errors.New(errOAuthStateRateLimited)
|
||||
@@ -153,7 +151,7 @@ func Authorize(c *gin.Context) {
|
||||
session := sessions.Default(c)
|
||||
userID := GetUserIDFromSession(session)
|
||||
if purpose == OAuthPurposeBind && userID == 0 {
|
||||
response.AbortUnauthorized(c, shared.UnAuthorized)
|
||||
response.AbortUnauthorized(c, errUnAuthorized)
|
||||
return
|
||||
}
|
||||
|
||||
@@ -182,7 +180,7 @@ func Authorize(c *gin.Context) {
|
||||
response.AbortInternal(c, err.Error())
|
||||
return
|
||||
}
|
||||
if err := db.Redis.Set(ctx, db.PrefixedKey(fmt.Sprintf(OAuthStateCacheKeyFormat, state)), payloadValue, OAuthStateCacheKeyExpiration).Err(); err != nil {
|
||||
if err := cachepkg.Redis.Set(ctx, cachepkg.PrefixedKey(fmt.Sprintf(OAuthStateCacheKeyFormat, state)), payloadValue, OAuthStateCacheKeyExpiration).Err(); err != nil {
|
||||
response.AbortInternal(c, err.Error())
|
||||
return
|
||||
}
|
||||
@@ -204,13 +202,13 @@ func Callback(c *gin.Context) {
|
||||
}
|
||||
|
||||
ctx := c.Request.Context()
|
||||
stateKey := db.PrefixedKey(fmt.Sprintf(OAuthStateCacheKeyFormat, req.State))
|
||||
payloadRaw, err := db.Redis.Get(ctx, stateKey).Result()
|
||||
stateKey := cachepkg.PrefixedKey(fmt.Sprintf(OAuthStateCacheKeyFormat, req.State))
|
||||
payloadRaw, err := cachepkg.Redis.Get(ctx, stateKey).Result()
|
||||
if err != nil {
|
||||
response.AbortBadRequest(c, errInvalidState)
|
||||
return
|
||||
}
|
||||
_ = db.Redis.Del(ctx, stateKey)
|
||||
_ = cachepkg.Redis.Del(ctx, stateKey)
|
||||
|
||||
payload, err := decodeOAuthStatePayload(payloadRaw)
|
||||
if err != nil {
|
||||
@@ -222,7 +220,7 @@ func Callback(c *gin.Context) {
|
||||
currentUserID := GetUserIDFromSession(session)
|
||||
|
||||
if payload.Purpose == OAuthPurposeBind && currentUserID == 0 {
|
||||
response.AbortUnauthorized(c, shared.UnAuthorized)
|
||||
response.AbortUnauthorized(c, errUnAuthorized)
|
||||
return
|
||||
}
|
||||
|
||||
@@ -288,7 +286,7 @@ func Callback(c *gin.Context) {
|
||||
func handleCallbackBind(ctx context.Context, c *gin.Context, source *AuthSource, userInfo *contracts.OAuthUserInfoDTO) {
|
||||
userID := GetUserIDFromContext(c)
|
||||
if userID == 0 {
|
||||
response.AbortUnauthorized(c, shared.UnAuthorized)
|
||||
response.AbortUnauthorized(c, errUnAuthorized)
|
||||
return
|
||||
}
|
||||
var user contracts.UserDTO
|
||||
@@ -477,7 +475,7 @@ func ListExternalAccounts(c *gin.Context) {
|
||||
func DeleteExternalAccount(c *gin.Context) {
|
||||
userID := GetUserIDFromContext(c)
|
||||
if userID == 0 {
|
||||
response.AbortUnauthorized(c, shared.UnAuthorized)
|
||||
response.AbortUnauthorized(c, errUnAuthorized)
|
||||
return
|
||||
}
|
||||
rawID := strings.TrimSpace(c.Param("id"))
|
||||
|
||||
@@ -11,10 +11,9 @@ import (
|
||||
"errors"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/core/contracts"
|
||||
db "github.com/Rain-kl/Wavelet/pkg/persistence"
|
||||
"github.com/Rain-kl/Wavelet/pkg/response"
|
||||
"github.com/Rain-kl/Wavelet/pkg/shared"
|
||||
otel_trace "github.com/Rain-kl/Wavelet/pkg/trace"
|
||||
db "github.com/Rain-kl/Wavelet/plugins/infra/database"
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
@@ -132,7 +131,7 @@ func LoginRequired() gin.HandlerFunc {
|
||||
|
||||
user, err := GetUserFromRequest(c)
|
||||
if err != nil {
|
||||
response.AbortUnauthorized(c, shared.UnAuthorized)
|
||||
response.AbortUnauthorized(c, errUnAuthorized)
|
||||
return
|
||||
}
|
||||
|
||||
@@ -150,7 +149,7 @@ func AdminRequired() gin.HandlerFunc {
|
||||
|
||||
user, err := GetUserFromRequest(c)
|
||||
if err != nil {
|
||||
response.AbortUnauthorized(c, shared.UnAuthorized)
|
||||
response.AbortUnauthorized(c, errUnAuthorized)
|
||||
return
|
||||
}
|
||||
|
||||
|
||||
@@ -18,8 +18,8 @@ import (
|
||||
|
||||
"github.com/Rain-kl/Wavelet/core"
|
||||
"github.com/Rain-kl/Wavelet/core/contracts"
|
||||
db "github.com/Rain-kl/Wavelet/pkg/persistence"
|
||||
"github.com/Rain-kl/Wavelet/plugins/domain/auth"
|
||||
db "github.com/Rain-kl/Wavelet/plugins/infra/database"
|
||||
)
|
||||
|
||||
type testUser struct {
|
||||
|
||||
@@ -7,7 +7,7 @@ package auth
|
||||
import (
|
||||
"context"
|
||||
|
||||
db "github.com/Rain-kl/Wavelet/pkg/persistence"
|
||||
db "github.com/Rain-kl/Wavelet/plugins/infra/database"
|
||||
)
|
||||
|
||||
// GetAuthSourceByID 根据 ID 获取认证源
|
||||
|
||||
@@ -9,7 +9,7 @@ import (
|
||||
"sync"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/core/contracts"
|
||||
db "github.com/Rain-kl/Wavelet/pkg/persistence"
|
||||
db "github.com/Rain-kl/Wavelet/plugins/infra/database"
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
|
||||
@@ -13,7 +13,7 @@ import (
|
||||
|
||||
"github.com/Rain-kl/Wavelet/core/contracts"
|
||||
"github.com/Rain-kl/Wavelet/pkg/config"
|
||||
db "github.com/Rain-kl/Wavelet/pkg/persistence"
|
||||
db "github.com/Rain-kl/Wavelet/plugins/infra/database"
|
||||
"github.com/gin-contrib/sessions"
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/google/uuid"
|
||||
|
||||
@@ -6,14 +6,14 @@ package cap
|
||||
import (
|
||||
"net/http"
|
||||
|
||||
pkgcap "github.com/Rain-kl/Wavelet/pkg/cap"
|
||||
"github.com/Rain-kl/Wavelet/pkg/logger"
|
||||
"github.com/Rain-kl/Wavelet/pkg/response"
|
||||
"github.com/Rain-kl/Wavelet/plugins/domain/cap/pow"
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
// ChallengeResponse is a local type alias for the pkg/cap.ChallengeResponse struct
|
||||
type ChallengeResponse = pkgcap.ChallengeResponse
|
||||
// ChallengeResponse is a local type alias for the pow.ChallengeResponse struct
|
||||
type ChallengeResponse = pow.ChallengeResponse
|
||||
|
||||
type challengeRequest struct {
|
||||
Scope string `json:"scope" form:"scope"`
|
||||
|
||||
@@ -13,9 +13,9 @@ import (
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
pkgcap "github.com/Rain-kl/Wavelet/pkg/cap"
|
||||
"github.com/Rain-kl/Wavelet/pkg/config"
|
||||
"github.com/Rain-kl/Wavelet/pkg/persistence"
|
||||
db "github.com/Rain-kl/Wavelet/plugins/infra/cache"
|
||||
"github.com/Rain-kl/Wavelet/plugins/domain/cap/pow"
|
||||
)
|
||||
|
||||
const (
|
||||
@@ -28,11 +28,11 @@ const (
|
||||
// Manager orchestrates challenge generation and solution validation.
|
||||
type Manager struct {
|
||||
secret []byte
|
||||
store pkgcap.Store
|
||||
store pow.Store
|
||||
}
|
||||
|
||||
// NewManager creates a new CAPTCHA Manager.
|
||||
func NewManager(secret []byte, store pkgcap.Store) *Manager {
|
||||
func NewManager(secret []byte, store pow.Store) *Manager {
|
||||
return &Manager{
|
||||
secret: secret,
|
||||
store: store,
|
||||
@@ -40,19 +40,19 @@ func NewManager(secret []byte, store pkgcap.Store) *Manager {
|
||||
}
|
||||
|
||||
// Generate creates a challenge response.
|
||||
func (m *Manager) Generate(ctx context.Context, scope string) (*pkgcap.ChallengeResponse, error) {
|
||||
func (m *Manager) Generate(ctx context.Context, scope string) (*pow.ChallengeResponse, error) {
|
||||
settings, err := CurrentSettings(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
challengeConfig := pkgcap.ChallengeConfig{
|
||||
challengeConfig := pow.ChallengeConfig{
|
||||
Count: settings.ChallengeCount,
|
||||
Size: settings.ChallengeSize,
|
||||
Difficulty: settings.ChallengeDifficulty,
|
||||
Expires: settings.ChallengeTTL,
|
||||
}
|
||||
return pkgcap.GenerateChallenge(m.secret, challengeConfig, scope)
|
||||
return pow.GenerateChallenge(m.secret, challengeConfig, scope)
|
||||
}
|
||||
|
||||
// RedeemResponse is returned to the client on redeem.
|
||||
@@ -65,14 +65,14 @@ type RedeemResponse struct {
|
||||
|
||||
// Redeem verifies PoW solutions and returns a one-time redeem token.
|
||||
func (m *Manager) Redeem(ctx context.Context, token string, solutions []int, scope string) (*RedeemResponse, error) {
|
||||
sigHex := pkgcap.JwtSigHex(token)
|
||||
sigHex := pow.JwtSigHex(token)
|
||||
if sigHex == "" {
|
||||
return &RedeemResponse{Success: false, Error: "invalid_token"}, nil
|
||||
}
|
||||
|
||||
nonceKey := "cap:nonce:" + sigHex
|
||||
|
||||
payload, err := pkgcap.VerifyChallengeSolutions(token, solutions, m.secret, scope)
|
||||
payload, err := pow.VerifyChallengeSolutions(token, solutions, m.secret, scope)
|
||||
if err != nil {
|
||||
return &RedeemResponse{Success: false, Error: err.Error()}, nil //nolint:nilerr // validation errors are returned as response, not system errors
|
||||
}
|
||||
@@ -96,8 +96,8 @@ func (m *Manager) Redeem(ctx context.Context, token string, solutions []int, sco
|
||||
return &RedeemResponse{Success: false, Error: "settings_load_error"}, err
|
||||
}
|
||||
|
||||
id := pkgcap.RandomHex(redeemTokenIDLength)
|
||||
verToken := pkgcap.RandomHex(redeemVerTokenLength)
|
||||
id := pow.RandomHex(redeemTokenIDLength)
|
||||
verToken := pow.RandomHex(redeemVerTokenLength)
|
||||
verHashBytes := sha256.Sum256([]byte(verToken))
|
||||
verHashHex := hex.EncodeToString(verHashBytes[:])
|
||||
|
||||
@@ -163,7 +163,7 @@ func (m *Manager) VerifyToken(ctx context.Context, token string, expectedScope s
|
||||
return true, nil
|
||||
}
|
||||
|
||||
func sGetAndDelete(ctx context.Context, store pkgcap.Store, key string) (string, bool, error) {
|
||||
func sGetAndDelete(ctx context.Context, store pow.Store, key string) (string, bool, error) {
|
||||
if store == nil {
|
||||
return "", false, nil
|
||||
}
|
||||
@@ -186,11 +186,11 @@ func GetDefaultManager() *Manager {
|
||||
return
|
||||
}
|
||||
|
||||
var store pkgcap.Store
|
||||
var store pow.Store
|
||||
if config.Config != nil && config.Config.Redis.Enabled && db.Redis != nil {
|
||||
store = pkgcap.NewRedisStore(db.Redis)
|
||||
store = pow.NewRedisStore(db.Redis)
|
||||
} else {
|
||||
store = pkgcap.NewMemoryStore(1 * time.Minute)
|
||||
store = pow.NewMemoryStore(1 * time.Minute)
|
||||
}
|
||||
|
||||
defaultManager = NewManager(secret, store)
|
||||
|
||||
@@ -0,0 +1,259 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package pow provides proof-of-work challenge generation and verification.
|
||||
package pow
|
||||
|
||||
import (
|
||||
"crypto/hmac"
|
||||
"crypto/rand"
|
||||
"crypto/sha256"
|
||||
"encoding/base64"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
const (
|
||||
jwtHeaderB64 = "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9"
|
||||
jwtPartsCount = 3 // JWT 三段结构
|
||||
defaultChallengeCount = 50 // 默认 PoW 难题数
|
||||
defaultChallengeSize = 32 // 默认盐值长度
|
||||
defaultDifficulty = 4 // 默认难度
|
||||
defaultNonceLength = 25 // 随机 Nonce 字节长度
|
||||
defaultExpires = 10 * time.Minute // 默认过期时间
|
||||
)
|
||||
|
||||
// ChallengeConfig holds parameters for the PoW challenge
|
||||
type ChallengeConfig struct {
|
||||
Count int // Number of puzzles (c)
|
||||
Size int // Salt length (s)
|
||||
Difficulty int // Difficulty prefix length (d)
|
||||
Expires time.Duration // Challenge TTL
|
||||
}
|
||||
|
||||
// ChallengeResponse is returned to the client
|
||||
type ChallengeResponse struct {
|
||||
Challenge struct {
|
||||
C int `json:"c"`
|
||||
S int `json:"s"`
|
||||
D int `json:"d"`
|
||||
} `json:"challenge"`
|
||||
Token string `json:"token"`
|
||||
Expires int64 `json:"expires"` // ms timestamp
|
||||
}
|
||||
|
||||
// ChallengePayload represents the signed JWT payload
|
||||
type ChallengePayload struct {
|
||||
Nonce string `json:"n"`
|
||||
Count int `json:"c"`
|
||||
Size int `json:"s"`
|
||||
Difficulty int `json:"d"`
|
||||
Expires int64 `json:"exp"` // ms timestamp
|
||||
IssuedAt int64 `json:"iat"` // ms timestamp
|
||||
Scope string `json:"sk,omitempty"`
|
||||
}
|
||||
|
||||
// RedeemRequest payload sent by client
|
||||
type RedeemRequest struct {
|
||||
Token string `json:"token"`
|
||||
Solutions []int `json:"solutions"`
|
||||
}
|
||||
|
||||
// RedeemResponse returned to client after verification
|
||||
type RedeemResponse struct {
|
||||
Success bool `json:"success"`
|
||||
Token string `json:"token,omitempty"`
|
||||
Expires int64 `json:"expires,omitempty"`
|
||||
Error string `json:"error,omitempty"`
|
||||
}
|
||||
|
||||
func b64urlEncode(data []byte) string {
|
||||
return base64.RawURLEncoding.EncodeToString(data)
|
||||
}
|
||||
|
||||
func b64urlDecode(str string) ([]byte, error) {
|
||||
return base64.RawURLEncoding.DecodeString(str)
|
||||
}
|
||||
|
||||
// RandomHex generates a cryptographically secure random hexadecimal string of the specified byte length.
|
||||
func RandomHex(byteLen int) string {
|
||||
bytes := make([]byte, byteLen)
|
||||
if _, err := rand.Read(bytes); err != nil {
|
||||
panic(err)
|
||||
}
|
||||
return hex.EncodeToString(bytes)
|
||||
}
|
||||
|
||||
func jwtSign(payload []byte, secret []byte) string {
|
||||
body := b64urlEncode(payload)
|
||||
sigInput := jwtHeaderB64 + "." + body
|
||||
|
||||
mac := hmac.New(sha256.New, secret)
|
||||
mac.Write([]byte(sigInput))
|
||||
sig := mac.Sum(nil)
|
||||
|
||||
return sigInput + "." + b64urlEncode(sig)
|
||||
}
|
||||
|
||||
func jwtVerify(token string, secret []byte) ([]byte, error) {
|
||||
parts := strings.Split(token, ".")
|
||||
if len(parts) != jwtPartsCount {
|
||||
return nil, errors.New(errInvalidTokenFormat)
|
||||
}
|
||||
if parts[0] != jwtHeaderB64 {
|
||||
return nil, errors.New(errInvalidHeader)
|
||||
}
|
||||
|
||||
sigInput := parts[0] + "." + parts[1]
|
||||
mac := hmac.New(sha256.New, secret)
|
||||
mac.Write([]byte(sigInput))
|
||||
expectedSig := mac.Sum(nil)
|
||||
|
||||
actualSig, err := b64urlDecode(parts[2])
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if !hmac.Equal(expectedSig, actualSig) {
|
||||
return nil, errors.New(errSignatureMismatch)
|
||||
}
|
||||
|
||||
payload, err := b64urlDecode(parts[1])
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return payload, nil
|
||||
}
|
||||
|
||||
// JwtSigHex extracts the signature part of a JWT token and returns it as a hexadecimal string.
|
||||
func JwtSigHex(token string) string {
|
||||
parts := strings.Split(token, ".")
|
||||
if len(parts) != jwtPartsCount {
|
||||
return ""
|
||||
}
|
||||
sigBytes, err := b64urlDecode(parts[2])
|
||||
if err != nil {
|
||||
return ""
|
||||
}
|
||||
return hex.EncodeToString(sigBytes)
|
||||
}
|
||||
|
||||
// GenerateChallenge produces a new challenge and signed token
|
||||
func GenerateChallenge(secret []byte, conf ChallengeConfig, scope string) (*ChallengeResponse, error) {
|
||||
if conf.Count <= 0 {
|
||||
conf.Count = defaultChallengeCount
|
||||
}
|
||||
if conf.Size <= 0 {
|
||||
conf.Size = defaultChallengeSize
|
||||
}
|
||||
if conf.Difficulty <= 0 {
|
||||
conf.Difficulty = defaultDifficulty
|
||||
}
|
||||
if conf.Expires <= 0 {
|
||||
conf.Expires = defaultExpires
|
||||
}
|
||||
|
||||
now := time.Now().UnixNano() / int64(time.Millisecond)
|
||||
expires := now + int64(conf.Expires/time.Millisecond)
|
||||
|
||||
payload := ChallengePayload{
|
||||
Nonce: RandomHex(defaultNonceLength),
|
||||
Count: conf.Count,
|
||||
Size: conf.Size,
|
||||
Difficulty: conf.Difficulty,
|
||||
Expires: expires,
|
||||
IssuedAt: now,
|
||||
Scope: scope,
|
||||
}
|
||||
|
||||
payloadBytes, err := json.Marshal(payload)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
token := jwtSign(payloadBytes, secret)
|
||||
|
||||
resp := &ChallengeResponse{
|
||||
Token: token,
|
||||
Expires: expires,
|
||||
}
|
||||
resp.Challenge.C = conf.Count
|
||||
resp.Challenge.S = conf.Size
|
||||
resp.Challenge.D = conf.Difficulty
|
||||
|
||||
return resp, nil
|
||||
}
|
||||
|
||||
// VerifyChallengeSolutions verifies client submitted solutions
|
||||
func VerifyChallengeSolutions(token string, solutions []int, secret []byte, expectedScope string) (*ChallengePayload, error) {
|
||||
payloadBytes, err := jwtVerify(token, secret)
|
||||
if err != nil {
|
||||
return nil, errors.New(errInvalidToken)
|
||||
}
|
||||
|
||||
var payload ChallengePayload
|
||||
if err := json.Unmarshal(payloadBytes, &payload); err != nil {
|
||||
return nil, errors.New(errInvalidToken)
|
||||
}
|
||||
|
||||
if expectedScope != "" && payload.Scope != expectedScope {
|
||||
return nil, errors.New(errScopeMismatch)
|
||||
}
|
||||
|
||||
now := time.Now().UnixNano() / int64(time.Millisecond)
|
||||
if payload.Expires < now {
|
||||
return nil, errors.New(errExpired)
|
||||
}
|
||||
|
||||
if len(solutions) != payload.Count {
|
||||
return nil, errors.New(errInvalidSolutions)
|
||||
}
|
||||
|
||||
tokenFnv := fnv1a(token)
|
||||
for i := 0; i < payload.Count; i++ {
|
||||
idxStr := strconv.Itoa(i + 1)
|
||||
saltSeed := fnv1aResume(tokenFnv, idxStr)
|
||||
targetSeed := fnv1aResume(saltSeed, "d")
|
||||
salt := prngFromHash(saltSeed, payload.Size)
|
||||
target := prngFromHash(targetSeed, payload.Difficulty)
|
||||
|
||||
hashInput := salt + strconv.Itoa(solutions[i])
|
||||
hashBytes := sha256.Sum256([]byte(hashInput))
|
||||
hashHex := hex.EncodeToString(hashBytes[:])
|
||||
|
||||
if !strings.HasPrefix(hashHex, target) {
|
||||
return nil, errors.New(errInvalidSolution)
|
||||
}
|
||||
}
|
||||
|
||||
return &payload, nil
|
||||
}
|
||||
|
||||
// Solve is a utility function to solve a challenge (mainly used for tests and reference implementation)
|
||||
func Solve(token string, count, size, difficulty int) []int {
|
||||
solutions := make([]int, count)
|
||||
tokenFnv := fnv1a(token)
|
||||
for i := 0; i < count; i++ {
|
||||
idxStr := strconv.Itoa(i + 1)
|
||||
saltSeed := fnv1aResume(tokenFnv, idxStr)
|
||||
targetSeed := fnv1aResume(saltSeed, "d")
|
||||
salt := prngFromHash(saltSeed, size)
|
||||
target := prngFromHash(targetSeed, difficulty)
|
||||
|
||||
for nonce := 0; nonce < 1000000; nonce++ {
|
||||
hashInput := salt + strconv.Itoa(nonce)
|
||||
hashBytes := sha256.Sum256([]byte(hashInput))
|
||||
hashHex := hex.EncodeToString(hashBytes[:])
|
||||
if strings.HasPrefix(hashHex, target) {
|
||||
solutions[i] = nonce
|
||||
break
|
||||
}
|
||||
}
|
||||
}
|
||||
return solutions
|
||||
}
|
||||
@@ -0,0 +1,15 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package pow
|
||||
|
||||
const (
|
||||
errInvalidTokenFormat = "invalid token format"
|
||||
errInvalidHeader = "invalid header"
|
||||
errSignatureMismatch = "signature mismatch"
|
||||
errInvalidToken = "invalid_token"
|
||||
errScopeMismatch = "scope_mismatch"
|
||||
errExpired = "expired"
|
||||
errInvalidSolutions = "invalid_solutions"
|
||||
errInvalidSolution = "invalid_solution"
|
||||
)
|
||||
@@ -0,0 +1,111 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package pow
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func TestPowChallengeFlow(t *testing.T) {
|
||||
secret := []byte("test-secret-key-1234567890123456")
|
||||
conf := ChallengeConfig{
|
||||
Count: 2,
|
||||
Size: 16,
|
||||
Difficulty: 1,
|
||||
Expires: 1 * time.Minute,
|
||||
}
|
||||
scope := "login"
|
||||
|
||||
resp, err := GenerateChallenge(secret, conf, scope)
|
||||
if err != nil {
|
||||
t.Fatalf("GenerateChallenge failed: %v", err)
|
||||
}
|
||||
if resp.Token == "" {
|
||||
t.Fatal("expected non-empty token")
|
||||
}
|
||||
if resp.Challenge.C != 2 {
|
||||
t.Fatalf("expected count 2, got %d", resp.Challenge.C)
|
||||
}
|
||||
|
||||
sigHex := JwtSigHex(resp.Token)
|
||||
if sigHex == "" {
|
||||
t.Fatal("expected non-empty sigHex")
|
||||
}
|
||||
|
||||
solutions := Solve(resp.Token, resp.Challenge.C, resp.Challenge.S, resp.Challenge.D)
|
||||
if len(solutions) != 2 {
|
||||
t.Fatalf("expected 2 solutions, got %d", len(solutions))
|
||||
}
|
||||
|
||||
payload, err := VerifyChallengeSolutions(resp.Token, solutions, secret, scope)
|
||||
if err != nil {
|
||||
t.Fatalf("VerifyChallengeSolutions failed: %v", err)
|
||||
}
|
||||
if payload.Scope != scope {
|
||||
t.Fatalf("expected scope %s, got %s", scope, payload.Scope)
|
||||
}
|
||||
|
||||
// Scope mismatch test
|
||||
_, err = VerifyChallengeSolutions(resp.Token, solutions, secret, "other_scope")
|
||||
if err == nil {
|
||||
t.Fatal("expected scope mismatch error")
|
||||
}
|
||||
|
||||
// Invalid solutions test
|
||||
_, err = VerifyChallengeSolutions(resp.Token, []int{9999999, 9999999}, secret, scope)
|
||||
if err == nil {
|
||||
t.Fatal("expected invalid solution error")
|
||||
}
|
||||
|
||||
// Invalid token test
|
||||
_, err = VerifyChallengeSolutions("invalid.jwt.token", solutions, secret, scope)
|
||||
if err == nil {
|
||||
t.Fatal("expected invalid token error")
|
||||
}
|
||||
}
|
||||
|
||||
func TestMemoryStore(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
store := NewMemoryStore(100 * time.Millisecond)
|
||||
|
||||
// Set and Get
|
||||
err := store.Set(ctx, "k1", "v1", 200*time.Millisecond)
|
||||
if err != nil {
|
||||
t.Fatalf("Set failed: %v", err)
|
||||
}
|
||||
val, ok, err := store.Get(ctx, "k1")
|
||||
if err != nil || !ok || val != "v1" {
|
||||
t.Fatalf("Get failed: val=%s, ok=%v, err=%v", val, ok, err)
|
||||
}
|
||||
|
||||
// SetNX
|
||||
set, err := store.SetNX(ctx, "k1", "v2", 200*time.Millisecond)
|
||||
if err != nil || set {
|
||||
t.Fatalf("SetNX should have failed because key exists: set=%v, err=%v", set, err)
|
||||
}
|
||||
|
||||
set, err = store.SetNX(ctx, "k2", "v2", 200*time.Millisecond)
|
||||
if err != nil || !set {
|
||||
t.Fatalf("SetNX should have succeeded: set=%v, err=%v", set, err)
|
||||
}
|
||||
|
||||
// GetAndDelete
|
||||
val, ok, err = store.GetAndDelete(ctx, "k2")
|
||||
if err != nil || !ok || val != "v2" {
|
||||
t.Fatalf("GetAndDelete failed: val=%s, ok=%v, err=%v", val, ok, err)
|
||||
}
|
||||
_, ok, _ = store.Get(ctx, "k2")
|
||||
if ok {
|
||||
t.Fatal("k2 should be deleted")
|
||||
}
|
||||
|
||||
// Delete
|
||||
_ = store.Delete(ctx, "k1")
|
||||
_, ok, _ = store.Get(ctx, "k1")
|
||||
if ok {
|
||||
t.Fatal("k1 should be deleted")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,56 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package pow
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// fnv1a returns the 32-bit FNV-1a hash of a string
|
||||
//
|
||||
//nolint:mnd // FNV-1a 算法位移常量
|
||||
func fnv1a(str string) uint32 {
|
||||
var hash uint32 = 2166136261
|
||||
for i := 0; i < len(str); i++ {
|
||||
hash ^= uint32(str[i])
|
||||
hash += (hash << 1) + (hash << 4) + (hash << 7) + (hash << 8) + (hash << 24)
|
||||
}
|
||||
return hash
|
||||
}
|
||||
|
||||
// fnv1aResume resumes FNV-1a hashing from a given state
|
||||
//
|
||||
//nolint:mnd // FNV-1a 算法位移常量
|
||||
func fnv1aResume(state uint32, str string) uint32 {
|
||||
h := state
|
||||
for i := 0; i < len(str); i++ {
|
||||
h ^= uint32(str[i])
|
||||
h += (hashShift(h))
|
||||
}
|
||||
return h
|
||||
}
|
||||
|
||||
// hashShift computes FNV-1a mix additions
|
||||
//
|
||||
//nolint:mnd
|
||||
func hashShift(h uint32) uint32 {
|
||||
return (h << 1) + (h << 4) + (h << 7) + (h << 8) + (h << 24)
|
||||
}
|
||||
|
||||
// prngFromHash generates a hex string of specified length using an initial hash state
|
||||
//
|
||||
//nolint:mnd // xorshift 算法位移常量
|
||||
func prngFromHash(initialHash uint32, length int) string {
|
||||
state := initialHash
|
||||
var result strings.Builder
|
||||
for result.Len() < length {
|
||||
state ^= state << 13
|
||||
state ^= state >> 17
|
||||
state ^= state << 5
|
||||
hexStr := fmt.Sprintf("%08x", state)
|
||||
result.WriteString(hexStr)
|
||||
}
|
||||
return result.String()[:length]
|
||||
}
|
||||
@@ -0,0 +1,186 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package pow
|
||||
|
||||
import (
|
||||
"context"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/redis/go-redis/v9"
|
||||
)
|
||||
|
||||
// Store defines the storage interface for challenge nonces and verification tokens
|
||||
type Store interface {
|
||||
Get(ctx context.Context, key string) (string, bool, error)
|
||||
Set(ctx context.Context, key string, val string, ttl time.Duration) error
|
||||
Delete(ctx context.Context, key string) error
|
||||
// SetNX atomically sets key=val with the given TTL only when the key does not
|
||||
// exist yet. It returns true when the key was actually written (i.e. this
|
||||
// caller "won" the race), and false when the key already existed.
|
||||
SetNX(ctx context.Context, key string, val string, ttl time.Duration) (bool, error)
|
||||
// GetAndDelete atomically retrieves the value of key and removes it in a
|
||||
// single operation. Returns ("", false, nil) when the key does not exist.
|
||||
GetAndDelete(ctx context.Context, key string) (string, bool, error)
|
||||
}
|
||||
|
||||
type memoryItem struct {
|
||||
value string
|
||||
expiresAt time.Time
|
||||
}
|
||||
|
||||
// MemoryStore is a thread-safe in-memory implementation of Store
|
||||
type MemoryStore struct {
|
||||
items map[string]memoryItem
|
||||
mu sync.Mutex // unified write-lock; promotes to exclusive for all ops
|
||||
}
|
||||
|
||||
// NewMemoryStore creates and initializes a new MemoryStore
|
||||
func NewMemoryStore(cleanupInterval time.Duration) *MemoryStore {
|
||||
store := &MemoryStore{
|
||||
items: make(map[string]memoryItem),
|
||||
}
|
||||
if cleanupInterval > 0 {
|
||||
go store.startCleanupLoop(cleanupInterval)
|
||||
}
|
||||
return store
|
||||
}
|
||||
|
||||
// Get 从 MemoryStore 获取指定 key 的值
|
||||
func (s *MemoryStore) Get(_ context.Context, key string) (string, bool, error) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
return s.getLocked(key)
|
||||
}
|
||||
|
||||
// getLocked is the internal helper – caller must hold s.mu.
|
||||
func (s *MemoryStore) getLocked(key string) (string, bool, error) {
|
||||
item, found := s.items[key]
|
||||
if !found {
|
||||
return "", false, nil
|
||||
}
|
||||
if time.Now().After(item.expiresAt) {
|
||||
delete(s.items, key)
|
||||
return "", false, nil
|
||||
}
|
||||
return item.value, true, nil
|
||||
}
|
||||
|
||||
// Set 向 MemoryStore 写入指定 key 的值
|
||||
func (s *MemoryStore) Set(_ context.Context, key string, val string, ttl time.Duration) error {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
s.items[key] = memoryItem{
|
||||
value: val,
|
||||
expiresAt: time.Now().Add(ttl),
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// Delete 从 MemoryStore 删除指定 key
|
||||
func (s *MemoryStore) Delete(_ context.Context, key string) error {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
delete(s.items, key)
|
||||
return nil
|
||||
}
|
||||
|
||||
// SetNX atomically sets key only when it is absent (or expired).
|
||||
// Returns true if the key was written by this call.
|
||||
func (s *MemoryStore) SetNX(_ context.Context, key string, val string, ttl time.Duration) (bool, error) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
|
||||
_, exists, _ := s.getLocked(key)
|
||||
if exists {
|
||||
return false, nil
|
||||
}
|
||||
s.items[key] = memoryItem{
|
||||
value: val,
|
||||
expiresAt: time.Now().Add(ttl),
|
||||
}
|
||||
return true, nil
|
||||
}
|
||||
|
||||
// GetAndDelete atomically retrieves and removes key in one critical section.
|
||||
func (s *MemoryStore) GetAndDelete(_ context.Context, key string) (string, bool, error) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
|
||||
val, exists, err := s.getLocked(key)
|
||||
if err != nil || !exists {
|
||||
return "", false, err
|
||||
}
|
||||
delete(s.items, key)
|
||||
return val, true, nil
|
||||
}
|
||||
|
||||
func (s *MemoryStore) startCleanupLoop(interval time.Duration) {
|
||||
ticker := time.NewTicker(interval)
|
||||
for range ticker.C {
|
||||
s.cleanupExpired()
|
||||
}
|
||||
}
|
||||
|
||||
func (s *MemoryStore) cleanupExpired() {
|
||||
now := time.Now()
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
for k, v := range s.items {
|
||||
if now.After(v.expiresAt) {
|
||||
delete(s.items, k)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// RedisStore is a GORM-compatible/standalone Redis-backed implementation of Store
|
||||
type RedisStore struct {
|
||||
client redis.UniversalClient
|
||||
}
|
||||
|
||||
// NewRedisStore creates a new RedisStore wrapping a redis.UniversalClient
|
||||
func NewRedisStore(client redis.UniversalClient) *RedisStore {
|
||||
return &RedisStore{
|
||||
client: client,
|
||||
}
|
||||
}
|
||||
|
||||
// Get 从 RedisStore 获取指定 key 的值
|
||||
func (s *RedisStore) Get(ctx context.Context, key string) (string, bool, error) {
|
||||
val, err := s.client.Get(ctx, key).Result()
|
||||
if err == redis.Nil {
|
||||
return "", false, nil
|
||||
}
|
||||
if err != nil {
|
||||
return "", false, err
|
||||
}
|
||||
return val, true, nil
|
||||
}
|
||||
|
||||
// Set 向 RedisStore 写入指定 key 的值
|
||||
func (s *RedisStore) Set(ctx context.Context, key string, val string, ttl time.Duration) error {
|
||||
return s.client.Set(ctx, key, val, ttl).Err()
|
||||
}
|
||||
|
||||
// Delete 从 RedisStore 删除指定 key
|
||||
func (s *RedisStore) Delete(ctx context.Context, key string) error {
|
||||
return s.client.Del(ctx, key).Err()
|
||||
}
|
||||
|
||||
// SetNX wraps Redis SET NX – returns true only when the key was newly created.
|
||||
func (s *RedisStore) SetNX(ctx context.Context, key string, val string, ttl time.Duration) (bool, error) {
|
||||
return s.client.SetNX(ctx, key, val, ttl).Result()
|
||||
}
|
||||
|
||||
// GetAndDelete wraps Redis GETDEL (available since Redis 6.2).
|
||||
func (s *RedisStore) GetAndDelete(ctx context.Context, key string) (string, bool, error) {
|
||||
val, err := s.client.GetDel(ctx, key).Result()
|
||||
if err == redis.Nil {
|
||||
return "", false, nil
|
||||
}
|
||||
if err != nil {
|
||||
return "", false, err
|
||||
}
|
||||
return val, true, nil
|
||||
}
|
||||
@@ -14,9 +14,9 @@ import (
|
||||
|
||||
"golang.org/x/sync/singleflight"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/pkg/persistence"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/pkg/util"
|
||||
cachepkg "github.com/Rain-kl/Wavelet/plugins/infra/cache"
|
||||
database "github.com/Rain-kl/Wavelet/plugins/infra/database"
|
||||
)
|
||||
|
||||
const (
|
||||
@@ -148,7 +148,7 @@ func loadRuntimeSettings(ctx context.Context) (RuntimeSettings, error) {
|
||||
Value string `gorm:"column:value"`
|
||||
}
|
||||
var records []configRecord
|
||||
if err := db.DB(ctx).Table("w_system_configs").Where("key IN ?", runtimeConfigKeys).Find(&records).Error; err != nil {
|
||||
if err := database.DB(ctx).Table("w_system_configs").Where("key IN ?", runtimeConfigKeys).Find(&records).Error; err != nil {
|
||||
return RuntimeSettings{}, err
|
||||
}
|
||||
configs := make(map[string]string, len(records))
|
||||
@@ -209,7 +209,7 @@ func (s *runtimeSettingsStore) ensureInvalidationListener() {
|
||||
const SystemConfigInvalidationChannel = "system_config:invalidation"
|
||||
|
||||
func startRuntimeSettingsInvalidationListener() {
|
||||
rdb := db.Redis
|
||||
rdb := cachepkg.Redis
|
||||
if rdb == nil {
|
||||
return
|
||||
}
|
||||
|
||||
@@ -18,15 +18,13 @@ import (
|
||||
|
||||
"github.com/Rain-kl/Wavelet/core"
|
||||
"github.com/Rain-kl/Wavelet/core/contracts"
|
||||
|
||||
db "github.com/Rain-kl/Wavelet/pkg/persistence"
|
||||
"github.com/Rain-kl/Wavelet/plugins/domain/admin"
|
||||
"github.com/Rain-kl/Wavelet/plugins/domain/auth"
|
||||
"github.com/Rain-kl/Wavelet/plugins/domain/message_gateway"
|
||||
"github.com/Rain-kl/Wavelet/plugins/domain/risk_control"
|
||||
"github.com/Rain-kl/Wavelet/plugins/domain/user"
|
||||
"github.com/Rain-kl/Wavelet/plugins/infra/cache"
|
||||
"github.com/Rain-kl/Wavelet/plugins/infra/database"
|
||||
db "github.com/Rain-kl/Wavelet/plugins/infra/database"
|
||||
"github.com/Rain-kl/Wavelet/plugins/infra/logger"
|
||||
"github.com/Rain-kl/Wavelet/plugins/infra/storage"
|
||||
)
|
||||
@@ -80,7 +78,7 @@ func TestAuthPlugin(t *testing.T) {
|
||||
ctx := core.NewContext(context.Background())
|
||||
testDB := setupTestDB(t)
|
||||
|
||||
require.NoError(t, database.New(database.WithDB(testDB)).Apply(ctx))
|
||||
require.NoError(t, db.New(db.WithDB(testDB)).Apply(ctx))
|
||||
require.NoError(t, cache.New().Apply(ctx))
|
||||
require.NoError(t, logger.New().Apply(ctx))
|
||||
|
||||
@@ -144,7 +142,7 @@ func TestUserPlugin(t *testing.T) {
|
||||
ctx := core.NewContext(context.Background())
|
||||
testDB := setupTestDB(t)
|
||||
|
||||
require.NoError(t, database.New(database.WithDB(testDB)).Apply(ctx))
|
||||
require.NoError(t, db.New(db.WithDB(testDB)).Apply(ctx))
|
||||
require.NoError(t, cache.New().Apply(ctx))
|
||||
require.NoError(t, logger.New().Apply(ctx))
|
||||
|
||||
@@ -250,7 +248,7 @@ func TestMessageGatewayPlugin(t *testing.T) {
|
||||
ctx := core.NewContext(context.Background())
|
||||
testDB := setupTestDB(t)
|
||||
|
||||
require.NoError(t, database.New(database.WithDB(testDB)).Apply(ctx))
|
||||
require.NoError(t, db.New(db.WithDB(testDB)).Apply(ctx))
|
||||
require.NoError(t, cache.New().Apply(ctx))
|
||||
require.NoError(t, logger.New().Apply(ctx))
|
||||
|
||||
@@ -338,7 +336,7 @@ func TestAdminPlugin(t *testing.T) {
|
||||
ctx := core.NewContext(context.Background())
|
||||
testDB := setupTestDB(t)
|
||||
|
||||
require.NoError(t, database.New(database.WithDB(testDB)).Apply(ctx))
|
||||
require.NoError(t, db.New(db.WithDB(testDB)).Apply(ctx))
|
||||
require.NoError(t, cache.New().Apply(ctx))
|
||||
require.NoError(t, logger.New().Apply(ctx))
|
||||
|
||||
@@ -398,7 +396,7 @@ func TestAllDomainPluginsCombined(t *testing.T) {
|
||||
testDB := setupTestDB(t)
|
||||
|
||||
// Apply Infra plugins
|
||||
require.NoError(t, database.New(database.WithDB(testDB)).Apply(ctx))
|
||||
require.NoError(t, db.New(db.WithDB(testDB)).Apply(ctx))
|
||||
require.NoError(t, cache.New(cache.WithRedis(rdb)).Apply(ctx))
|
||||
require.NoError(t, logger.New().Apply(ctx))
|
||||
require.NoError(t, storage.New().Apply(ctx))
|
||||
|
||||
@@ -0,0 +1,21 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package message_gateway
|
||||
|
||||
import "context"
|
||||
|
||||
// Handler processes one inbound message.
|
||||
type Handler func(ctx context.Context, msg InboundMessage) error
|
||||
|
||||
// Factory constructs a Channel from decrypted config.
|
||||
type Factory func(cfg ChannelConfig, onInbound Handler) (Channel, error)
|
||||
|
||||
// Channel is one connected messaging adapter.
|
||||
type Channel interface {
|
||||
Type() string
|
||||
Connect(ctx context.Context) error
|
||||
Disconnect(ctx context.Context) error
|
||||
Send(ctx context.Context, to Recipient, msg OutboundMessage) error
|
||||
Capabilities() Capability
|
||||
}
|
||||
@@ -0,0 +1,161 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package qq implements the official QQ Bot C2C adapter.
|
||||
package qq
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/pkg/logger"
|
||||
"github.com/Rain-kl/Wavelet/plugins/domain/message_gateway"
|
||||
"github.com/tencent-connect/botgo"
|
||||
"github.com/tencent-connect/botgo/dto"
|
||||
"github.com/tencent-connect/botgo/event"
|
||||
"github.com/tencent-connect/botgo/openapi"
|
||||
"github.com/tencent-connect/botgo/token"
|
||||
"golang.org/x/oauth2"
|
||||
)
|
||||
|
||||
// qqEvent is a testable inbound envelope.
|
||||
type qqEvent struct {
|
||||
Kind string
|
||||
UserID string
|
||||
Text string
|
||||
MessageID string
|
||||
}
|
||||
|
||||
// Adapter is an official QQ Bot C2C channel.
|
||||
type Adapter struct {
|
||||
cfg message_gateway.ChannelConfig
|
||||
onInbound message_gateway.Handler
|
||||
api openapi.OpenAPI
|
||||
tokenSrc oauth2.TokenSource
|
||||
cancel context.CancelFunc
|
||||
mu sync.Mutex
|
||||
disconnected bool
|
||||
}
|
||||
|
||||
// New constructs a QQ adapter.
|
||||
func New(cfg message_gateway.ChannelConfig, onInbound message_gateway.Handler) (message_gateway.Channel, error) {
|
||||
if strings.TrimSpace(cfg.Credentials["app_id"]) == "" || strings.TrimSpace(cfg.Credentials["app_secret"]) == "" {
|
||||
return nil, fmt.Errorf("qq: app_id and app_secret are required")
|
||||
}
|
||||
return &Adapter{cfg: cfg, onInbound: onInbound}, nil
|
||||
}
|
||||
|
||||
// Type returns qq.
|
||||
func (a *Adapter) Type() string { return message_gateway.ChannelTypeQQ }
|
||||
|
||||
// Capabilities reports C2C text/media support.
|
||||
func (a *Adapter) Capabilities() message_gateway.Capability {
|
||||
return message_gateway.Capability{Text: true, Image: true, File: true, Reply: true}
|
||||
}
|
||||
|
||||
// Connect starts the official WebSocket session (C2C intent).
|
||||
func (a *Adapter) Connect(ctx context.Context) error {
|
||||
credentials := &token.QQBotCredentials{
|
||||
AppID: a.cfg.Credentials["app_id"],
|
||||
AppSecret: a.cfg.Credentials["app_secret"],
|
||||
}
|
||||
tokSrc := token.NewQQBotTokenSource(credentials)
|
||||
runCtx, cancel := context.WithCancel(ctx)
|
||||
if err := token.StartRefreshAccessToken(runCtx, tokSrc); err != nil {
|
||||
cancel()
|
||||
return fmt.Errorf("qq: refresh token: %w", err)
|
||||
}
|
||||
|
||||
var api openapi.OpenAPI
|
||||
const apiTimeout = 5 * time.Second
|
||||
if strings.EqualFold(strings.TrimSpace(a.cfg.Extra["sandbox"]), "true") {
|
||||
api = botgo.NewSandboxOpenAPI(credentials.AppID, tokSrc).WithTimeout(apiTimeout)
|
||||
} else {
|
||||
api = botgo.NewOpenAPI(credentials.AppID, tokSrc).WithTimeout(apiTimeout)
|
||||
}
|
||||
|
||||
wsAP, err := api.WS(ctx, nil, "")
|
||||
if err != nil {
|
||||
cancel()
|
||||
return fmt.Errorf("qq: websocket ap: %w", err)
|
||||
}
|
||||
|
||||
intent := event.RegisterHandlers(event.C2CMessageEventHandler(func(_ *dto.WSPayload, data *dto.WSC2CMessageData) error {
|
||||
authorID := ""
|
||||
if data != nil && data.Author != nil {
|
||||
authorID = data.Author.ID
|
||||
}
|
||||
text := ""
|
||||
id := ""
|
||||
if data != nil {
|
||||
text = data.Content
|
||||
id = data.ID
|
||||
}
|
||||
a.handleEvent(runCtx, qqEvent{Kind: "c2c", UserID: authorID, Text: text, MessageID: id})
|
||||
return nil
|
||||
}))
|
||||
|
||||
a.mu.Lock()
|
||||
a.api = api
|
||||
a.tokenSrc = tokSrc
|
||||
a.cancel = cancel
|
||||
a.disconnected = false
|
||||
a.mu.Unlock()
|
||||
|
||||
go func() {
|
||||
if err := botgo.NewSessionManager().Start(wsAP, tokSrc, &intent); err != nil {
|
||||
logger.ErrorF(runCtx, "qq session stopped: %v", err)
|
||||
}
|
||||
}()
|
||||
return nil
|
||||
}
|
||||
|
||||
// Disconnect stops token refresh and drops further inbound events.
|
||||
func (a *Adapter) Disconnect(_ context.Context) error {
|
||||
a.mu.Lock()
|
||||
defer a.mu.Unlock()
|
||||
a.disconnected = true
|
||||
if a.cancel != nil {
|
||||
a.cancel()
|
||||
a.cancel = nil
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// Send posts a C2C text reply.
|
||||
func (a *Adapter) Send(ctx context.Context, to message_gateway.Recipient, msg message_gateway.OutboundMessage) error {
|
||||
a.mu.Lock()
|
||||
api := a.api
|
||||
a.mu.Unlock()
|
||||
if api == nil {
|
||||
return fmt.Errorf("qq: not connected")
|
||||
}
|
||||
_, err := api.PostC2CMessage(ctx, to.PlatformUserID, &dto.MessageToCreate{
|
||||
Content: msg.Text,
|
||||
MsgID: msg.ReplyToID,
|
||||
})
|
||||
return err
|
||||
}
|
||||
|
||||
func (a *Adapter) handleEvent(ctx context.Context, ev qqEvent) {
|
||||
if ev.Kind != "c2c" {
|
||||
return
|
||||
}
|
||||
a.mu.Lock()
|
||||
disconnected := a.disconnected
|
||||
a.mu.Unlock()
|
||||
if disconnected || a.onInbound == nil {
|
||||
return
|
||||
}
|
||||
_ = a.onInbound(ctx, message_gateway.InboundMessage{
|
||||
ChannelID: a.cfg.ID,
|
||||
ChannelType: message_gateway.ChannelTypeQQ,
|
||||
PlatformUserID: ev.UserID,
|
||||
ChatID: ev.UserID,
|
||||
MessageID: ev.MessageID,
|
||||
Text: ev.Text,
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,42 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package qq
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/plugins/domain/message_gateway"
|
||||
)
|
||||
|
||||
func TestHandleEvent_DropsNonC2C(t *testing.T) {
|
||||
var got int
|
||||
a := &Adapter{onInbound: func(ctx context.Context, msg message_gateway.InboundMessage) error {
|
||||
got++
|
||||
return nil
|
||||
}}
|
||||
a.handleEvent(context.Background(), qqEvent{Kind: "group", UserID: "u1", Text: "hi"})
|
||||
if got != 0 {
|
||||
t.Fatal("non-C2C must be ignored")
|
||||
}
|
||||
}
|
||||
|
||||
func TestHandleEvent_C2CText(t *testing.T) {
|
||||
var got message_gateway.InboundMessage
|
||||
a := &Adapter{cfg: message_gateway.ChannelConfig{ID: 3}, onInbound: func(ctx context.Context, msg message_gateway.InboundMessage) error {
|
||||
got = msg
|
||||
return nil
|
||||
}}
|
||||
a.handleEvent(context.Background(), qqEvent{Kind: "c2c", UserID: "openid-1", Text: "hello", MessageID: "m1"})
|
||||
if got.Text != "hello" || got.PlatformUserID != "openid-1" || got.ChannelID != 3 {
|
||||
t.Fatalf("%+v", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNew_RequiresCreds(t *testing.T) {
|
||||
_, err := New(message_gateway.ChannelConfig{}, nil)
|
||||
if err == nil {
|
||||
t.Fatal("expected error")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,153 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package telegram implements the Telegram private-chat adapter.
|
||||
package telegram
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/plugins/domain/message_gateway"
|
||||
tele "gopkg.in/telebot.v4"
|
||||
)
|
||||
|
||||
// Adapter is a Telegram private-chat channel.
|
||||
type Adapter struct {
|
||||
cfg message_gateway.ChannelConfig
|
||||
onInbound message_gateway.Handler
|
||||
bot *tele.Bot
|
||||
}
|
||||
|
||||
// New constructs a Telegram adapter. Call message_gateway.Register from the runner.
|
||||
func New(cfg message_gateway.ChannelConfig, onInbound message_gateway.Handler) (message_gateway.Channel, error) {
|
||||
if strings.TrimSpace(cfg.Credentials["bot_token"]) == "" {
|
||||
return nil, fmt.Errorf("telegram: bot_token is required")
|
||||
}
|
||||
return &Adapter{cfg: cfg, onInbound: onInbound}, nil
|
||||
}
|
||||
|
||||
// Type returns telegram.
|
||||
func (a *Adapter) Type() string { return message_gateway.ChannelTypeTelegram }
|
||||
|
||||
// Capabilities reports private-chat media support.
|
||||
func (a *Adapter) Capabilities() message_gateway.Capability {
|
||||
return message_gateway.Capability{Text: true, Image: true, File: true, Reply: true}
|
||||
}
|
||||
|
||||
// Connect starts long polling.
|
||||
func (a *Adapter) Connect(ctx context.Context) error {
|
||||
pref := tele.Settings{
|
||||
Token: a.cfg.Credentials["bot_token"],
|
||||
Poller: &tele.LongPoller{Timeout: 10},
|
||||
}
|
||||
if base := strings.TrimSpace(a.cfg.Extra["base_url"]); base != "" {
|
||||
pref.URL = strings.TrimSuffix(base, "/")
|
||||
}
|
||||
bot, err := tele.NewBot(pref)
|
||||
if err != nil {
|
||||
return fmt.Errorf("telegram: new bot: %w", err)
|
||||
}
|
||||
a.bot = bot
|
||||
bot.Handle(tele.OnText, func(c tele.Context) error {
|
||||
a.handleTeleMessage(ctx, c.Message())
|
||||
return nil
|
||||
})
|
||||
bot.Handle(tele.OnPhoto, func(c tele.Context) error {
|
||||
a.handleTeleMessage(ctx, c.Message())
|
||||
return nil
|
||||
})
|
||||
bot.Handle(tele.OnDocument, func(c tele.Context) error {
|
||||
a.handleTeleMessage(ctx, c.Message())
|
||||
return nil
|
||||
})
|
||||
go bot.Start()
|
||||
go func() {
|
||||
<-ctx.Done()
|
||||
bot.Stop()
|
||||
}()
|
||||
return nil
|
||||
}
|
||||
|
||||
// Disconnect stops the bot.
|
||||
func (a *Adapter) Disconnect(_ context.Context) error {
|
||||
if a.bot != nil {
|
||||
a.bot.Stop()
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// Send replies to a private chat.
|
||||
func (a *Adapter) Send(_ context.Context, to message_gateway.Recipient, msg message_gateway.OutboundMessage) error {
|
||||
if a.bot == nil {
|
||||
return fmt.Errorf("telegram: not connected")
|
||||
}
|
||||
chatID, err := strconv.ParseInt(to.ChatID, 10, 64)
|
||||
if err != nil {
|
||||
return fmt.Errorf("telegram: chat id: %w", err)
|
||||
}
|
||||
_, err = a.bot.Send(tele.ChatID(chatID), msg.Text)
|
||||
return err
|
||||
}
|
||||
|
||||
func (a *Adapter) handleTeleMessage(ctx context.Context, m *tele.Message) {
|
||||
if m == nil || m.Chat == nil || m.Chat.Type != tele.ChatPrivate {
|
||||
return
|
||||
}
|
||||
if a.onInbound == nil {
|
||||
return
|
||||
}
|
||||
msg := message_gateway.InboundMessage{
|
||||
ChannelID: a.cfg.ID,
|
||||
ChannelType: message_gateway.ChannelTypeTelegram,
|
||||
PlatformUserID: strconv.FormatInt(m.Sender.ID, 10),
|
||||
ChatID: strconv.FormatInt(m.Chat.ID, 10),
|
||||
MessageID: strconv.Itoa(m.ID),
|
||||
Text: m.Text,
|
||||
}
|
||||
if m.Caption != "" && msg.Text == "" {
|
||||
msg.Text = m.Caption
|
||||
}
|
||||
if a.bot != nil {
|
||||
msg.Attachments = a.downloadMedia(m)
|
||||
}
|
||||
_ = a.onInbound(ctx, msg)
|
||||
}
|
||||
|
||||
func (a *Adapter) downloadMedia(m *tele.Message) []message_gateway.Attachment {
|
||||
var files []*tele.File
|
||||
var names []string
|
||||
if m.Photo != nil {
|
||||
files = append(files, m.Photo.MediaFile())
|
||||
names = append(names, "photo.jpg")
|
||||
}
|
||||
if m.Document != nil {
|
||||
files = append(files, &m.Document.File)
|
||||
name := m.Document.FileName
|
||||
if name == "" {
|
||||
name = "file"
|
||||
}
|
||||
names = append(names, name)
|
||||
}
|
||||
if len(files) == 0 {
|
||||
return nil
|
||||
}
|
||||
dir, err := os.MkdirTemp("", "wg-tg-*")
|
||||
if err != nil {
|
||||
return []message_gateway.Attachment{{Error: err.Error()}}
|
||||
}
|
||||
out := make([]message_gateway.Attachment, 0, len(files))
|
||||
for i, f := range files {
|
||||
path := filepath.Join(dir, names[i])
|
||||
if err := a.bot.Download(f, path); err != nil {
|
||||
out = append(out, message_gateway.Attachment{FileName: names[i], Error: err.Error()})
|
||||
continue
|
||||
}
|
||||
out = append(out, message_gateway.Attachment{Path: path, FileName: names[i]})
|
||||
}
|
||||
return out
|
||||
}
|
||||
@@ -0,0 +1,56 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package telegram
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/plugins/domain/message_gateway"
|
||||
tele "gopkg.in/telebot.v4"
|
||||
)
|
||||
|
||||
func TestHandleUpdate_DropsGroups(t *testing.T) {
|
||||
var got int
|
||||
a := &Adapter{onInbound: func(ctx context.Context, msg message_gateway.InboundMessage) error {
|
||||
got++
|
||||
return nil
|
||||
}}
|
||||
a.handleTeleMessage(context.Background(), &tele.Message{
|
||||
ID: 1,
|
||||
Text: "hi",
|
||||
Chat: &tele.Chat{ID: -100, Type: tele.ChatGroup},
|
||||
Sender: &tele.User{ID: 1},
|
||||
})
|
||||
if got != 0 {
|
||||
t.Fatalf("group must be ignored")
|
||||
}
|
||||
}
|
||||
|
||||
func TestHandleUpdate_PrivateText(t *testing.T) {
|
||||
var got message_gateway.InboundMessage
|
||||
a := &Adapter{
|
||||
cfg: message_gateway.ChannelConfig{ID: 7, Type: "telegram"},
|
||||
onInbound: func(ctx context.Context, msg message_gateway.InboundMessage) error {
|
||||
got = msg
|
||||
return nil
|
||||
},
|
||||
}
|
||||
a.handleTeleMessage(context.Background(), &tele.Message{
|
||||
ID: 9,
|
||||
Text: "hi",
|
||||
Chat: &tele.Chat{ID: 42, Type: tele.ChatPrivate},
|
||||
Sender: &tele.User{ID: 42},
|
||||
})
|
||||
if got.Text != "hi" || got.PlatformUserID != "42" || got.ChannelID != 7 {
|
||||
t.Fatalf("%+v", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNew_RequiresToken(t *testing.T) {
|
||||
_, err := New(message_gateway.ChannelConfig{}, nil)
|
||||
if err == nil {
|
||||
t.Fatal("expected error")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,62 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package message_gateway defines channel adapters, pairing codes, and inbound types.
|
||||
package message_gateway
|
||||
|
||||
// ChannelTypeTelegram is the Telegram private-chat adapter type.
|
||||
const ChannelTypeTelegram = "telegram"
|
||||
|
||||
// ChannelTypeQQ is the official QQ Bot C2C adapter type.
|
||||
const ChannelTypeQQ = "qq"
|
||||
|
||||
// Capability describes what an adapter can send and receive.
|
||||
type Capability struct {
|
||||
Text bool
|
||||
Image bool
|
||||
File bool
|
||||
Reply bool
|
||||
Group bool
|
||||
}
|
||||
|
||||
// ChannelConfig is the decrypted runtime config passed to a factory.
|
||||
type ChannelConfig struct {
|
||||
ID uint64
|
||||
Type string
|
||||
Name string
|
||||
Credentials map[string]string
|
||||
Extra map[string]string
|
||||
}
|
||||
|
||||
// Recipient is the outbound destination on a platform.
|
||||
type Recipient struct {
|
||||
ChatID string
|
||||
PlatformUserID string
|
||||
}
|
||||
|
||||
// Attachment is a downloaded inbound file sitting on local disk.
|
||||
type Attachment struct {
|
||||
Path string
|
||||
FileName string
|
||||
MIME string
|
||||
Error string
|
||||
}
|
||||
|
||||
// InboundMessage is a normalized private-chat message.
|
||||
type InboundMessage struct {
|
||||
ChannelID uint64
|
||||
ChannelType string
|
||||
PlatformUserID string
|
||||
ChatID string
|
||||
MessageID string
|
||||
Text string
|
||||
Attachments []Attachment
|
||||
BindingUserID *uint64
|
||||
}
|
||||
|
||||
// OutboundMessage is a reply or probe send.
|
||||
type OutboundMessage struct {
|
||||
Text string
|
||||
ReplyToID string
|
||||
Attachments []Attachment
|
||||
}
|
||||
@@ -10,8 +10,6 @@ import (
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
pkgmg "github.com/Rain-kl/Wavelet/pkg/message_gateway"
|
||||
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
@@ -37,7 +35,7 @@ func bindChannel(ctx context.Context, userID uint64, req BindRequest) (BindingDT
|
||||
if err != nil || channelID == 0 {
|
||||
return BindingDTO{}, errChannelIDRequired
|
||||
}
|
||||
code := pkgmg.NormalizeCode(req.Code)
|
||||
code := NormalizeCode(req.Code)
|
||||
if code == "" {
|
||||
return BindingDTO{}, errCodeInvalid
|
||||
}
|
||||
|
||||
@@ -0,0 +1,50 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package message_gateway
|
||||
|
||||
import (
|
||||
"crypto/rand"
|
||||
"strings"
|
||||
"unicode"
|
||||
)
|
||||
|
||||
// CodeAlphabet excludes easily confused runes 0/O/1/I.
|
||||
const CodeAlphabet = "ABCDEFGHJKLMNPQRSTUVWXYZ23456789"
|
||||
|
||||
// CodeLength is the raw pairing code size.
|
||||
const CodeLength = 8
|
||||
|
||||
// GenerateCode returns an 8-character pairing code.
|
||||
func GenerateCode() (string, error) {
|
||||
buf := make([]byte, CodeLength)
|
||||
if _, err := rand.Read(buf); err != nil {
|
||||
return "", err
|
||||
}
|
||||
out := make([]byte, CodeLength)
|
||||
for i, b := range buf {
|
||||
out[i] = CodeAlphabet[int(b)%len(CodeAlphabet)]
|
||||
}
|
||||
return string(out), nil
|
||||
}
|
||||
|
||||
// NormalizeCode strips separators and uppercases.
|
||||
func NormalizeCode(s string) string {
|
||||
var b strings.Builder
|
||||
for _, r := range s {
|
||||
if r == '-' || unicode.IsSpace(r) {
|
||||
continue
|
||||
}
|
||||
b.WriteRune(unicode.ToUpper(r))
|
||||
}
|
||||
return b.String()
|
||||
}
|
||||
|
||||
// FormatCode renders ABCD-EFGH.
|
||||
func FormatCode(s string) string {
|
||||
s = NormalizeCode(s)
|
||||
if len(s) != CodeLength {
|
||||
return s
|
||||
}
|
||||
return s[:4] + "-" + s[4:]
|
||||
}
|
||||
@@ -0,0 +1,33 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package message_gateway
|
||||
|
||||
import (
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestGenerateCode_AlphabetAndLength(t *testing.T) {
|
||||
code, err := GenerateCode()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(code) != 8 {
|
||||
t.Fatalf("len=%d", len(code))
|
||||
}
|
||||
for _, r := range code {
|
||||
if !strings.ContainsRune(CodeAlphabet, r) {
|
||||
t.Fatalf("bad rune %q", r)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestNormalizeAndFormat(t *testing.T) {
|
||||
if got := NormalizeCode("ab-cd-ef-gh"); got != "ABCDEFGH" {
|
||||
t.Fatalf("got %q", got)
|
||||
}
|
||||
if got := FormatCode("ABCDEFGH"); got != "ABCD-EFGH" {
|
||||
t.Fatalf("got %q", got)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,98 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package push
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"strings"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/pkg/httppool"
|
||||
)
|
||||
|
||||
func init() {
|
||||
Register("custom", &CustomPusher{})
|
||||
}
|
||||
|
||||
// maxCustomResponseBytes 限制读取 Webhook 响应体的最大字节数,防止无界读取。
|
||||
const maxCustomResponseBytes = 4096
|
||||
|
||||
// CustomPusher 自定义 Webhook 发送实现
|
||||
type CustomPusher struct{}
|
||||
|
||||
// Send 发送自定义 webhook
|
||||
func (p *CustomPusher) Send(ctx context.Context, cfg Config, _ string, body map[string]any, template string, _ map[string]any) (string, error) {
|
||||
if cfg.URL == "" {
|
||||
return "", errors.New("custom: URL is required")
|
||||
}
|
||||
|
||||
var reqBody []byte
|
||||
|
||||
if template != "" {
|
||||
// 替换模板中的 {{key}} 占位符
|
||||
rendered := ParseTemplate(template, body)
|
||||
reqBody = []byte(rendered)
|
||||
} else {
|
||||
// 兜底:直接把 body 转为 JSON 字符串发送
|
||||
var err error
|
||||
reqBody, err = json.Marshal(body)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("custom: marshal body failed: %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
httpReq, err := http.NewRequestWithContext(ctx, http.MethodPost, cfg.URL, bytes.NewReader(reqBody))
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("custom: create http request failed: %w", err)
|
||||
}
|
||||
httpReq.Header.Set("Content-Type", "application/json")
|
||||
|
||||
// 如果配置了 Key 且格式为 "HeaderName:HeaderValue",我们可以附加测试用 Header
|
||||
if cfg.Key != "" && strings.Contains(cfg.Key, ":") {
|
||||
parts := strings.SplitN(cfg.Key, ":", 2) //nolint:mnd
|
||||
httpReq.Header.Set(strings.TrimSpace(parts[0]), strings.TrimSpace(parts[1]))
|
||||
}
|
||||
|
||||
client := httppool.NewClient(defaultHTTPClientTimeout)
|
||||
resp, err := client.Do(httpReq)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("custom: http request failed: %w", err)
|
||||
}
|
||||
defer func() { _ = resp.Body.Close() }()
|
||||
|
||||
bodyBytes, _ := io.ReadAll(io.LimitReader(resp.Body, maxCustomResponseBytes))
|
||||
upstreamResp := strings.TrimSpace(string(bodyBytes))
|
||||
|
||||
if resp.StatusCode < 200 || resp.StatusCode >= 300 { //nolint:mnd
|
||||
return upstreamResp, fmt.Errorf("custom: http status %s", resp.Status)
|
||||
}
|
||||
|
||||
// 部分 Webhook(如企业微信、钉钉)即使业务失败也返回 HTTP 200,
|
||||
// 仅当响应体包含非零 errcode 时才判定为发送失败,避免审计记录误报成功。
|
||||
var apiResp struct {
|
||||
ErrCode int `json:"errcode"`
|
||||
ErrMsg string `json:"errmsg"`
|
||||
}
|
||||
if err := json.Unmarshal(bodyBytes, &apiResp); err == nil && apiResp.ErrCode != 0 {
|
||||
return upstreamResp, fmt.Errorf("custom: webhook rejected: errcode=%d errmsg=%q", apiResp.ErrCode, apiResp.ErrMsg)
|
||||
}
|
||||
|
||||
return upstreamResp, nil
|
||||
}
|
||||
|
||||
// ValidateConfig 校验自定义配置
|
||||
func (p *CustomPusher) ValidateConfig(cfg Config) error {
|
||||
if cfg.URL == "" {
|
||||
return errors.New("webhook URL is required")
|
||||
}
|
||||
if !strings.HasPrefix(cfg.URL, "http://") && !strings.HasPrefix(cfg.URL, "https://") {
|
||||
return errors.New("webhook URL must start with http:// or https://")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,91 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package push
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestCustomPusherSend_ResponseBodyErrcode(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
statusCode int
|
||||
body string
|
||||
wantErr bool
|
||||
wantErrMsg string
|
||||
}{
|
||||
{
|
||||
name: "wechat business error returns HTTP 200 with non-zero errcode",
|
||||
statusCode: http.StatusOK,
|
||||
body: `{"errcode":93000,"errmsg":"invalid request data"}`,
|
||||
wantErr: true,
|
||||
wantErrMsg: "errcode=93000",
|
||||
},
|
||||
{
|
||||
name: "wechat success returns errcode 0",
|
||||
statusCode: http.StatusOK,
|
||||
body: `{"errcode":0,"errmsg":"ok"}`,
|
||||
wantErr: false,
|
||||
},
|
||||
{
|
||||
name: "json response without errcode is tolerated",
|
||||
statusCode: http.StatusOK,
|
||||
body: `{"success":true}`,
|
||||
wantErr: false,
|
||||
},
|
||||
{
|
||||
name: "non-json response body is tolerated",
|
||||
statusCode: http.StatusOK,
|
||||
body: "ok",
|
||||
wantErr: false,
|
||||
},
|
||||
{
|
||||
name: "empty response body is tolerated",
|
||||
statusCode: http.StatusNoContent,
|
||||
body: "",
|
||||
wantErr: false,
|
||||
},
|
||||
{
|
||||
name: "http error status still fails",
|
||||
statusCode: http.StatusInternalServerError,
|
||||
body: `{"errcode":0,"errmsg":"ok"}`,
|
||||
wantErr: true,
|
||||
wantErrMsg: "http status",
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||
w.WriteHeader(tt.statusCode)
|
||||
_, _ = w.Write([]byte(tt.body))
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
pusher := &CustomPusher{}
|
||||
upstreamResp, err := pusher.Send(context.Background(),
|
||||
Config{Channel: "custom", URL: srv.URL},
|
||||
"",
|
||||
map[string]any{"title": "t", "content": "c"},
|
||||
`{"title":"$title","content":"$content"}`,
|
||||
nil,
|
||||
)
|
||||
if tt.wantErr {
|
||||
require.Error(t, err)
|
||||
assert.Contains(t, err.Error(), tt.wantErrMsg)
|
||||
return
|
||||
}
|
||||
assert.NoError(t, err)
|
||||
if tt.body != "" {
|
||||
assert.Contains(t, upstreamResp, tt.body)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,119 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package push
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net"
|
||||
"net/smtp"
|
||||
"strings"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/pkg/util"
|
||||
)
|
||||
|
||||
func init() {
|
||||
Register("email", &EmailPusher{})
|
||||
}
|
||||
|
||||
// EmailPusher 极简 SMTP 邮件推送实现 (静态、解耦)
|
||||
type EmailPusher struct{}
|
||||
|
||||
// sanitizeEmailHeader removes CR/LF bytes so untrusted values cannot inject
|
||||
// additional email headers (email header injection).
|
||||
func sanitizeEmailHeader(v string) string {
|
||||
v = strings.ReplaceAll(v, "\r", "")
|
||||
v = strings.ReplaceAll(v, "\n", "")
|
||||
return v
|
||||
}
|
||||
|
||||
// Send 发送邮件
|
||||
func (p *EmailPusher) Send(ctx context.Context, cfg Config, target string, body map[string]any, _ string, ext map[string]any) (string, error) {
|
||||
if cfg.URL == "" || cfg.Key == "" || cfg.Secret == "" {
|
||||
return "", errors.New("email: SMTP configuration (url, key, secret) is incomplete")
|
||||
}
|
||||
if target == "" {
|
||||
return "", errors.New("email: target email address is required")
|
||||
}
|
||||
|
||||
title := defaultTitle
|
||||
if t, ok := body["title"].(string); ok && t != "" {
|
||||
title = t
|
||||
}
|
||||
|
||||
content := ""
|
||||
if c, ok := body["content"].(string); ok && c != "" {
|
||||
content = c
|
||||
} else {
|
||||
// 自动格式化 map
|
||||
var parts []string
|
||||
for k, v := range body {
|
||||
parts = append(parts, fmt.Sprintf("<p><b>%s</b>: %v</p>", k, v))
|
||||
}
|
||||
content = strings.Join(parts, "")
|
||||
}
|
||||
|
||||
// 邮件头和体
|
||||
from := cfg.Key
|
||||
to := target
|
||||
|
||||
// 如果 ext 中指定了 from_name,我们在 From 头部包含它
|
||||
fromName := "System Notification"
|
||||
if ext != nil {
|
||||
if fn, ok := ext["from_name"].(string); ok && fn != "" {
|
||||
fromName = fn
|
||||
}
|
||||
}
|
||||
|
||||
subjectHeader := fmt.Sprintf("Subject: %s\r\n", sanitizeEmailHeader(title))
|
||||
fromHeader := fmt.Sprintf("From: %s <%s>\r\n", sanitizeEmailHeader(fromName), sanitizeEmailHeader(from))
|
||||
toHeader := fmt.Sprintf("To: %s\r\n", sanitizeEmailHeader(to))
|
||||
mimeHeader := "MIME-version: 1.0;\r\nContent-Type: text/html; charset=\"UTF-8\";\r\n\r\n"
|
||||
|
||||
// 拼装完整的邮件报文
|
||||
// 简单的 HTML 正文渲染
|
||||
htmlBody := fmt.Sprintf(`<html><body><h2>%s</h2><div>%s</div></body></html>`, title, content)
|
||||
msg := []byte(fromHeader + toHeader + subjectHeader + mimeHeader + htmlBody + "\r\n")
|
||||
|
||||
// 解析 Host 和 Port
|
||||
host, port, err := net.SplitHostPort(cfg.URL)
|
||||
if err != nil {
|
||||
host = cfg.URL
|
||||
port = "25" // 默认 SMTP 端口
|
||||
}
|
||||
|
||||
auth := smtp.PlainAuth("", cfg.Key, cfg.Secret, host)
|
||||
|
||||
// 异步超时处理
|
||||
errChan := make(chan error, 1)
|
||||
util.Go(func() {
|
||||
errChan <- smtp.SendMail(host+":"+port, auth, from, []string{to}, msg)
|
||||
})
|
||||
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return "", ctx.Err()
|
||||
case err := <-errChan:
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("email: send smtp mail failed: %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
return "", nil
|
||||
}
|
||||
|
||||
// ValidateConfig 校验邮件 SMTP 配置
|
||||
func (p *EmailPusher) ValidateConfig(cfg Config) error {
|
||||
if cfg.URL == "" {
|
||||
return errors.New("SMTP host:port is required")
|
||||
}
|
||||
if cfg.Key == "" {
|
||||
return errors.New("SMTP username is required")
|
||||
}
|
||||
if cfg.Secret == "" {
|
||||
return errors.New("SMTP password is required")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,26 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package push
|
||||
|
||||
import "testing"
|
||||
|
||||
func TestSanitizeEmailHeader(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
input string
|
||||
want string
|
||||
}{
|
||||
{"plain", "System Notification", "System Notification"},
|
||||
{"crlf stripped", "alert\r\nBcc: attacker@example.com", "alertBcc: attacker@example.com"},
|
||||
{"cr stripped", "a\rb", "ab"},
|
||||
{"lf stripped", "a\nb", "ab"},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
if got := sanitizeEmailHeader(tt.input); got != tt.want {
|
||||
t.Errorf("sanitizeEmailHeader(%q) = %q, want %q", tt.input, got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,275 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package push
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"crypto/hmac"
|
||||
"crypto/sha256"
|
||||
"encoding/base64"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/pkg/httppool"
|
||||
)
|
||||
|
||||
func init() {
|
||||
Register("lark", &LarkPusher{})
|
||||
}
|
||||
|
||||
const (
|
||||
msgTypeInteractive = "interactive"
|
||||
)
|
||||
|
||||
// LarkPusher 飞书 Webhook 机器人推送实现
|
||||
type LarkPusher struct{}
|
||||
|
||||
type larkTextContent struct {
|
||||
Text string `json:"text"`
|
||||
}
|
||||
|
||||
type larkCardHeaderTitle struct {
|
||||
Content string `json:"content"`
|
||||
Tag string `json:"tag"`
|
||||
}
|
||||
|
||||
type larkCardHeader struct {
|
||||
Template string `json:"template"` // "blue", "orange", "red" etc.
|
||||
Title larkCardHeaderTitle `json:"title"`
|
||||
}
|
||||
|
||||
type larkCardElementText struct {
|
||||
Content string `json:"content"`
|
||||
Tag string `json:"tag"` // "lark_md"
|
||||
}
|
||||
|
||||
type larkCardElement struct {
|
||||
Tag string `json:"tag"` // "div"
|
||||
Text larkCardElementText `json:"text"`
|
||||
}
|
||||
|
||||
type larkCardContent struct {
|
||||
Header larkCardHeader `json:"header"`
|
||||
Elements []larkCardElement `json:"elements"`
|
||||
}
|
||||
|
||||
type larkMessageRequest struct {
|
||||
MessageType string `json:"msg_type"`
|
||||
Timestamp string `json:"timestamp,omitempty"`
|
||||
Sign string `json:"sign,omitempty"`
|
||||
Content larkTextContent `json:"content,omitempty"`
|
||||
Card *larkCardContent `json:"card,omitempty"`
|
||||
}
|
||||
|
||||
type larkMessageResponse struct {
|
||||
Code int `json:"code"`
|
||||
Msg string `json:"msg"`
|
||||
}
|
||||
|
||||
// Send 执行飞书消息发送
|
||||
//
|
||||
//nolint:nestif,cyclop
|
||||
func (p *LarkPusher) Send(ctx context.Context, cfg Config, _ string, body map[string]any, template string, _ map[string]any) (string, error) {
|
||||
if cfg.URL == "" {
|
||||
return "", errors.New("lark: URL is required")
|
||||
}
|
||||
|
||||
var req larkMessageRequest
|
||||
|
||||
// 1. 如果有自定义模板,我们尝试进行解析
|
||||
if template != "" {
|
||||
rendered := ParseTemplate(template, body)
|
||||
|
||||
// 尝试解析原生的 Lark Card
|
||||
var customCard larkCardContent
|
||||
var rawMap map[string]any
|
||||
_ = json.Unmarshal([]byte(rendered), &rawMap)
|
||||
|
||||
if rawMap != nil && rawMap["elements"] != nil {
|
||||
// 如果包含 elements 字段,说明是用户定制的原生飞书卡片 JSON
|
||||
if err := json.Unmarshal([]byte(rendered), &customCard); err == nil {
|
||||
req.MessageType = msgTypeInteractive
|
||||
req.Card = &customCard
|
||||
} else {
|
||||
req.MessageType = "text"
|
||||
req.Content.Text = rendered
|
||||
}
|
||||
} else {
|
||||
// 说明配置的是系统统一通知消息 of JSON 模板:{"title": "...", "content": "...", "level": "..."}
|
||||
type larkNotificationMessage struct {
|
||||
Title string `json:"title"`
|
||||
Content string `json:"content"`
|
||||
Level string `json:"level"`
|
||||
}
|
||||
var msg larkNotificationMessage
|
||||
if err := json.Unmarshal([]byte(rendered), &msg); err == nil && (msg.Title != "" || msg.Content != "") {
|
||||
title := msg.Title
|
||||
if title == "" {
|
||||
title = defaultTitle
|
||||
}
|
||||
content := msg.Content
|
||||
level := strings.ToUpper(msg.Level)
|
||||
if level == "" {
|
||||
level = levelInfo
|
||||
}
|
||||
|
||||
headerColor := "blue"
|
||||
switch level {
|
||||
case "IMPORTANT":
|
||||
headerColor = "orange"
|
||||
case "CRITICAL":
|
||||
headerColor = "red"
|
||||
}
|
||||
|
||||
req.MessageType = msgTypeInteractive
|
||||
req.Card = &larkCardContent{
|
||||
Header: larkCardHeader{
|
||||
Template: headerColor,
|
||||
Title: larkCardHeaderTitle{
|
||||
Content: title,
|
||||
Tag: "plain_text",
|
||||
},
|
||||
},
|
||||
Elements: []larkCardElement{
|
||||
{
|
||||
Tag: "div",
|
||||
Text: larkCardElementText{
|
||||
Content: content,
|
||||
Tag: "lark_md",
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
} else {
|
||||
// 兜底:如果无法按 JSON 解析出结构化字段,当做普通文本发送
|
||||
req.MessageType = "text"
|
||||
req.Content.Text = rendered
|
||||
}
|
||||
}
|
||||
} else {
|
||||
// 2. 如果无模板,默认生成一个精美的飞书互动卡片
|
||||
title := defaultTitle
|
||||
if t, ok := body["title"].(string); ok && t != "" {
|
||||
title = t
|
||||
}
|
||||
|
||||
content := ""
|
||||
if c, ok := body["content"].(string); ok && c != "" {
|
||||
content = c
|
||||
} else {
|
||||
// 兜底:如果连 content 都没有,把 body 里的所有值拼成 markdown
|
||||
var parts []string
|
||||
for k, v := range body {
|
||||
parts = append(parts, fmt.Sprintf("**%s**: %v", k, v))
|
||||
}
|
||||
content = strings.Join(parts, "\n")
|
||||
}
|
||||
|
||||
level := levelInfo
|
||||
if l, ok := body["level"].(string); ok && l != "" {
|
||||
level = strings.ToUpper(l)
|
||||
}
|
||||
|
||||
// 根据级别确定飞书卡片头部的背景色模板
|
||||
headerColor := "blue"
|
||||
switch level {
|
||||
case "IMPORTANT":
|
||||
headerColor = "orange"
|
||||
case "CRITICAL":
|
||||
headerColor = "red"
|
||||
}
|
||||
|
||||
req.MessageType = msgTypeInteractive
|
||||
req.Card = &larkCardContent{
|
||||
Header: larkCardHeader{
|
||||
Template: headerColor,
|
||||
Title: larkCardHeaderTitle{
|
||||
Content: title,
|
||||
Tag: "plain_text",
|
||||
},
|
||||
},
|
||||
Elements: []larkCardElement{
|
||||
{
|
||||
Tag: "div",
|
||||
Text: larkCardElementText{
|
||||
Content: content,
|
||||
Tag: "lark_md",
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
// 3. 计算签名 (如果配置了 secret)
|
||||
if cfg.Secret != "" {
|
||||
timestamp := time.Now().Unix()
|
||||
sign, err := larkSign(cfg.Secret, timestamp)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("lark: sign failed: %w", err)
|
||||
}
|
||||
req.Timestamp = strconv.FormatInt(timestamp, 10)
|
||||
req.Sign = sign
|
||||
}
|
||||
|
||||
jsonData, err := json.Marshal(req)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("lark: marshal request failed: %w", err)
|
||||
}
|
||||
|
||||
// 4. 发送 POST 请求
|
||||
httpReq, err := http.NewRequestWithContext(ctx, http.MethodPost, cfg.URL, bytes.NewBuffer(jsonData))
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("lark: create http request failed: %w", err)
|
||||
}
|
||||
httpReq.Header.Set("Content-Type", "application/json")
|
||||
|
||||
client := httppool.NewClient(defaultHTTPClientTimeout)
|
||||
resp, err := client.Do(httpReq)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("lark: http request failed: %w", err)
|
||||
}
|
||||
defer func() { _ = resp.Body.Close() }()
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return "", fmt.Errorf("lark: http status %s", resp.Status)
|
||||
}
|
||||
|
||||
var res larkMessageResponse
|
||||
if err := json.NewDecoder(resp.Body).Decode(&res); err != nil {
|
||||
return "", fmt.Errorf("lark: decode response failed: %w", err)
|
||||
}
|
||||
|
||||
if res.Code != 0 {
|
||||
return "", fmt.Errorf("lark: send message failed, code %d: %s", res.Code, res.Msg)
|
||||
}
|
||||
|
||||
return "", nil
|
||||
}
|
||||
|
||||
// ValidateConfig 校验飞书配置
|
||||
func (p *LarkPusher) ValidateConfig(cfg Config) error {
|
||||
if cfg.URL == "" {
|
||||
return errors.New("webhook URL is required")
|
||||
}
|
||||
if !strings.HasPrefix(cfg.URL, "http://") && !strings.HasPrefix(cfg.URL, "https://") {
|
||||
return errors.New("webhook URL must start with http:// or https://")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func larkSign(secret string, timestamp int64) (string, error) {
|
||||
stringToSign := fmt.Sprintf("%v", timestamp) + "\n" + secret
|
||||
h := hmac.New(sha256.New, []byte(stringToSign))
|
||||
_, err := h.Write(nil)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return base64.StdEncoding.EncodeToString(h.Sum(nil)), nil
|
||||
}
|
||||
@@ -0,0 +1,67 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package push 提供解耦的、无外部业务依赖 of 通知推送底层实现
|
||||
package push
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"sync"
|
||||
"time"
|
||||
)
|
||||
|
||||
const (
|
||||
defaultTitle = "系统通知"
|
||||
levelInfo = "INFO"
|
||||
defaultHTTPClientTimeout = 10 * time.Second
|
||||
)
|
||||
|
||||
// Config 基础通知渠道配置
|
||||
type Config struct {
|
||||
Channel string `json:"channel"` // 渠道名称,例如 "lark", "custom", "email" 等,唯一标识
|
||||
URL string `json:"url,omitempty"` // Webhook 地址或 SMTP 地址
|
||||
Secret string `json:"secret,omitempty"` // 签名密钥或 SMTP 密码/Token
|
||||
Key string `json:"key,omitempty"` // AppID 或 SMTP 用户名
|
||||
Ext map[string]any `json:"ext,omitempty"` // 预留拓展 JSON 配置
|
||||
}
|
||||
|
||||
// Pusher 通知推送渠道接口
|
||||
type Pusher interface {
|
||||
// Send 发送通知消息
|
||||
// target: 发送目标 (如邮箱地址或特定用户标识;若为 bot 机器人此项为空)
|
||||
// body: 消息体数据 (含默认字段如 title, content, level)
|
||||
// template: 消息卡片/模板 JSON (可选)
|
||||
// ext: 预留的单次发送拓展数据
|
||||
// 返回 upstreamResp: 上游服务返回的响应内容(如 Webhook 响应体),用于任务日志审计;无响应时为空字符串
|
||||
Send(ctx context.Context, cfg Config, target string, body map[string]any, template string, ext map[string]any) (upstreamResp string, err error)
|
||||
|
||||
// ValidateConfig 校验渠道配置合法性
|
||||
ValidateConfig(cfg Config) error
|
||||
}
|
||||
|
||||
var (
|
||||
pushersMu sync.RWMutex
|
||||
pushers = make(map[string]Pusher)
|
||||
)
|
||||
|
||||
// Register 注册一个推送渠道实现
|
||||
func Register(channelType string, pusher Pusher) {
|
||||
pushersMu.Lock()
|
||||
defer pushersMu.Unlock()
|
||||
if pusher == nil {
|
||||
panic("push: Register pusher is nil")
|
||||
}
|
||||
pushers[channelType] = pusher
|
||||
}
|
||||
|
||||
// GetPusher 获取指定类型的推送渠道实现
|
||||
func GetPusher(channelType string) (Pusher, error) {
|
||||
pushersMu.RLock()
|
||||
defer pushersMu.RUnlock()
|
||||
pusher, ok := pushers[channelType]
|
||||
if !ok {
|
||||
return nil, fmt.Errorf("push: unknown channel type %q", channelType)
|
||||
}
|
||||
return pusher, nil
|
||||
}
|
||||
@@ -0,0 +1,158 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package push
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"strings"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/pkg/httppool"
|
||||
)
|
||||
|
||||
func init() {
|
||||
Register("telegram", &TelegramPusher{})
|
||||
}
|
||||
|
||||
// TelegramPusher Telegram 机器人推送实现
|
||||
type TelegramPusher struct{}
|
||||
|
||||
type telegramMessageRequest struct {
|
||||
ChatID string `json:"chat_id"`
|
||||
Text string `json:"text"`
|
||||
ParseMode string `json:"parse_mode,omitempty"`
|
||||
}
|
||||
|
||||
type telegramErrorResponse struct {
|
||||
Ok bool `json:"ok"`
|
||||
ErrorCode int `json:"error_code"`
|
||||
Description string `json:"description"`
|
||||
}
|
||||
|
||||
// Send 执行 Telegram 消息发送
|
||||
//
|
||||
//nolint:cyclop
|
||||
func (p *TelegramPusher) Send(ctx context.Context, cfg Config, target string, body map[string]any, template string, _ map[string]any) (string, error) {
|
||||
if cfg.Secret == "" {
|
||||
return "", errors.New("telegram: Bot Token (Secret) is required")
|
||||
}
|
||||
|
||||
chatID := target
|
||||
if chatID == "" {
|
||||
chatID = cfg.Key // Use default chat ID (Key) if target is blank
|
||||
}
|
||||
if chatID == "" {
|
||||
return "", errors.New("telegram: chat_id (target or default Key) is required")
|
||||
}
|
||||
|
||||
baseURL := cfg.URL
|
||||
if baseURL == "" {
|
||||
baseURL = "https://api.telegram.org"
|
||||
}
|
||||
baseURL = strings.TrimSuffix(baseURL, "/")
|
||||
|
||||
title := defaultTitle
|
||||
if t, ok := body["title"].(string); ok && t != "" {
|
||||
title = t
|
||||
}
|
||||
content := ""
|
||||
if c, ok := body["content"].(string); ok && c != "" {
|
||||
content = c
|
||||
} else {
|
||||
var parts []string
|
||||
for k, v := range body {
|
||||
parts = append(parts, fmt.Sprintf("<b>%s</b>: %v", k, v))
|
||||
}
|
||||
content = strings.Join(parts, "\n")
|
||||
}
|
||||
level := levelInfo
|
||||
if l, ok := body["level"].(string); ok && l != "" {
|
||||
level = strings.ToUpper(l)
|
||||
}
|
||||
|
||||
var text string
|
||||
if template != "" {
|
||||
text = ParseTemplate(template, body)
|
||||
} else {
|
||||
text = fmt.Sprintf("<b>[%s] %s</b>\n\n%s", escapeHTML(level), escapeHTML(title), escapeHTML(content))
|
||||
}
|
||||
|
||||
// Try sending with HTML parse mode
|
||||
err := p.sendMessage(ctx, baseURL, cfg.Secret, chatID, text, "HTML")
|
||||
if err != nil {
|
||||
// Fallback: send as plain text without parse mode
|
||||
plainText := text
|
||||
if template == "" {
|
||||
plainText = fmt.Sprintf("[%s] %s\n\n%s", level, title, content)
|
||||
}
|
||||
fallbackErr := p.sendMessage(ctx, baseURL, cfg.Secret, chatID, plainText, "")
|
||||
if fallbackErr != nil {
|
||||
return "", fmt.Errorf("telegram: send message failed (fallback also failed): %w (original HTML error: %v)", fallbackErr, err)
|
||||
}
|
||||
}
|
||||
|
||||
return "", nil
|
||||
}
|
||||
|
||||
// ValidateConfig 校验 Telegram 配置
|
||||
func (p *TelegramPusher) ValidateConfig(cfg Config) error {
|
||||
if cfg.Secret == "" {
|
||||
return errors.New("bot Token (Secret) is required")
|
||||
}
|
||||
if cfg.URL != "" {
|
||||
if !strings.HasPrefix(cfg.URL, "http://") && !strings.HasPrefix(cfg.URL, "https://") {
|
||||
return errors.New("API base URL must start with http:// or https://")
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (p *TelegramPusher) sendMessage(ctx context.Context, baseURL, token, chatID, text, parseMode string) error {
|
||||
apiURL := fmt.Sprintf("%s/bot%s/sendMessage", baseURL, token)
|
||||
|
||||
reqPayload := telegramMessageRequest{
|
||||
ChatID: chatID,
|
||||
Text: text,
|
||||
ParseMode: parseMode,
|
||||
}
|
||||
|
||||
jsonData, err := json.Marshal(reqPayload)
|
||||
if err != nil {
|
||||
return fmt.Errorf("marshal request failed: %w", err)
|
||||
}
|
||||
|
||||
httpReq, err := http.NewRequestWithContext(ctx, http.MethodPost, apiURL, bytes.NewBuffer(jsonData))
|
||||
if err != nil {
|
||||
return fmt.Errorf("create http request failed: %w", err)
|
||||
}
|
||||
httpReq.Header.Set("Content-Type", "application/json")
|
||||
|
||||
client := httppool.NewClient(defaultHTTPClientTimeout)
|
||||
resp, err := client.Do(httpReq)
|
||||
if err != nil {
|
||||
return fmt.Errorf("http request failed: %w", err)
|
||||
}
|
||||
defer func() { _ = resp.Body.Close() }()
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
var errRes telegramErrorResponse
|
||||
if decodeErr := json.NewDecoder(resp.Body).Decode(&errRes); decodeErr == nil {
|
||||
return fmt.Errorf("http status %d: %s", resp.StatusCode, errRes.Description)
|
||||
}
|
||||
return fmt.Errorf("http status %s", resp.Status)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func escapeHTML(s string) string {
|
||||
s = strings.ReplaceAll(s, "&", "&")
|
||||
s = strings.ReplaceAll(s, "<", "<")
|
||||
s = strings.ReplaceAll(s, ">", ">")
|
||||
return s
|
||||
}
|
||||
@@ -0,0 +1,116 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package push
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestTelegramPusher_Send(t *testing.T) {
|
||||
t.Run("successful send with HTML parse mode", func(t *testing.T) {
|
||||
var receivedReq telegramMessageRequest
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
assert.Equal(t, "/botmy-token/sendMessage", r.URL.Path)
|
||||
assert.Equal(t, http.MethodPost, r.Method)
|
||||
assert.Equal(t, "application/json", r.Header.Get("Content-Type"))
|
||||
|
||||
err := json.NewDecoder(r.Body).Decode(&receivedReq)
|
||||
require.NoError(t, err)
|
||||
|
||||
w.WriteHeader(http.StatusOK)
|
||||
_, _ = w.Write([]byte(`{"ok": true}`))
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
pusher := &TelegramPusher{}
|
||||
cfg := Config{
|
||||
Channel: "telegram",
|
||||
URL: server.URL,
|
||||
Secret: "my-token",
|
||||
}
|
||||
body := map[string]any{
|
||||
"title": "Alert",
|
||||
"content": "Host down",
|
||||
"level": "CRITICAL",
|
||||
}
|
||||
_, err := pusher.Send(context.Background(), cfg, "123456", body, "", nil)
|
||||
require.NoError(t, err)
|
||||
|
||||
assert.Equal(t, "123456", receivedReq.ChatID)
|
||||
assert.Contains(t, receivedReq.Text, "[CRITICAL] Alert")
|
||||
assert.Contains(t, receivedReq.Text, "Host down")
|
||||
assert.Equal(t, "HTML", receivedReq.ParseMode)
|
||||
})
|
||||
|
||||
t.Run("fallback to plain text on HTML error", func(t *testing.T) {
|
||||
var requests []*telegramMessageRequest
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
var req telegramMessageRequest
|
||||
err := json.NewDecoder(r.Body).Decode(&req)
|
||||
require.NoError(t, err)
|
||||
requests = append(requests, &req)
|
||||
|
||||
if len(requests) == 1 {
|
||||
w.WriteHeader(http.StatusBadRequest)
|
||||
_, _ = w.Write([]byte(`{"ok": false, "error_code": 400, "description": "Bad Request: can't parse entities"}`))
|
||||
} else {
|
||||
w.WriteHeader(http.StatusOK)
|
||||
_, _ = w.Write([]byte(`{"ok": true}`))
|
||||
}
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
pusher := &TelegramPusher{}
|
||||
cfg := Config{
|
||||
Channel: "telegram",
|
||||
URL: server.URL,
|
||||
Secret: "my-token",
|
||||
}
|
||||
body := map[string]any{
|
||||
"title": "Alert & Info",
|
||||
"content": "A < B comparison",
|
||||
"level": "INFO",
|
||||
}
|
||||
_, err := pusher.Send(context.Background(), cfg, "123456", body, "", nil)
|
||||
require.NoError(t, err)
|
||||
|
||||
require.Len(t, requests, 2)
|
||||
assert.Equal(t, "HTML", requests[0].ParseMode)
|
||||
assert.Equal(t, "", requests[1].ParseMode)
|
||||
assert.Contains(t, requests[1].Text, "[INFO] Alert & Info")
|
||||
assert.Contains(t, requests[1].Text, "A < B comparison")
|
||||
})
|
||||
|
||||
t.Run("validation error", func(t *testing.T) {
|
||||
pusher := &TelegramPusher{}
|
||||
cfg := Config{
|
||||
Channel: "telegram",
|
||||
URL: "https://api.telegram.org",
|
||||
}
|
||||
err := pusher.ValidateConfig(cfg)
|
||||
assert.Error(t, err)
|
||||
|
||||
cfg = Config{
|
||||
Channel: "telegram",
|
||||
URL: "ftp://api.telegram.org",
|
||||
Secret: "token",
|
||||
}
|
||||
err = pusher.ValidateConfig(cfg)
|
||||
assert.Error(t, err)
|
||||
|
||||
cfg = Config{
|
||||
Channel: "telegram",
|
||||
Secret: "token",
|
||||
}
|
||||
err = pusher.ValidateConfig(cfg)
|
||||
assert.NoError(t, err)
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,78 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package push
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"strconv"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// ParseTemplate parses template strings by replacing {{placeholder}} structures with values from body.
|
||||
// It is a single-pass parser designed for high performance and low allocations.
|
||||
func ParseTemplate(template string, body map[string]any) string {
|
||||
var buf strings.Builder
|
||||
buf.Grow(len(template))
|
||||
|
||||
i := 0
|
||||
for {
|
||||
pos := strings.Index(template[i:], "{{")
|
||||
if pos == -1 {
|
||||
buf.WriteString(template[i:])
|
||||
break
|
||||
}
|
||||
// Write prefix
|
||||
buf.WriteString(template[i : i+pos])
|
||||
i += pos + 2 // skip "{{"
|
||||
|
||||
endPos := strings.Index(template[i:], "}}")
|
||||
if endPos == -1 {
|
||||
// Unbalanced "{{"
|
||||
buf.WriteString("{{")
|
||||
buf.WriteString(template[i:])
|
||||
break
|
||||
}
|
||||
key := template[i : i+endPos]
|
||||
if val, ok := body[key]; ok {
|
||||
buf.WriteString(formatValue(val))
|
||||
} else {
|
||||
// Keep the placeholder if key not found
|
||||
buf.WriteString("{{")
|
||||
buf.WriteString(key)
|
||||
buf.WriteString("}}")
|
||||
}
|
||||
i += endPos + 2 // skip "}}"
|
||||
}
|
||||
return buf.String()
|
||||
}
|
||||
|
||||
func formatValue(v any) string {
|
||||
if v == nil {
|
||||
return ""
|
||||
}
|
||||
switch val := v.(type) {
|
||||
case string:
|
||||
return val
|
||||
case []byte:
|
||||
return string(val)
|
||||
case int:
|
||||
return strconv.Itoa(val)
|
||||
case int32:
|
||||
return strconv.FormatInt(int64(val), 10)
|
||||
case int64:
|
||||
return strconv.FormatInt(val, 10)
|
||||
case float64:
|
||||
return strconv.FormatFloat(val, 'f', -1, 64)
|
||||
case bool:
|
||||
return strconv.FormatBool(val)
|
||||
default:
|
||||
// If it's a map, slice, or struct, marshal it to JSON.
|
||||
b, err := json.Marshal(v)
|
||||
if err == nil {
|
||||
return string(b)
|
||||
}
|
||||
return fmt.Sprintf("%v", v)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,75 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package push
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
)
|
||||
|
||||
func TestParseTemplate(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
template string
|
||||
body map[string]any
|
||||
expected string
|
||||
}{
|
||||
{
|
||||
name: "simple replacement",
|
||||
template: "hello {{name}}",
|
||||
body: map[string]any{"name": "world"},
|
||||
expected: "hello world",
|
||||
},
|
||||
{
|
||||
name: "multiple replacements",
|
||||
template: "{{greeting}} {{name}}!",
|
||||
body: map[string]any{"greeting": "Hello", "name": "Alice"},
|
||||
expected: "Hello Alice!",
|
||||
},
|
||||
{
|
||||
name: "missing key preserves placeholder",
|
||||
template: "hello {{name}} and {{other}}",
|
||||
body: map[string]any{"name": "world"},
|
||||
expected: "hello world and {{other}}",
|
||||
},
|
||||
{
|
||||
name: "unbalanced placeholders",
|
||||
template: "hello {{name",
|
||||
body: map[string]any{"name": "world"},
|
||||
expected: "hello {{name",
|
||||
},
|
||||
{
|
||||
name: "nil value",
|
||||
template: "val: {{val}}",
|
||||
body: map[string]any{"val": nil},
|
||||
expected: "val: ",
|
||||
},
|
||||
{
|
||||
name: "basic types",
|
||||
template: "int: {{i}}, float: {{f}}, bool: {{b}}",
|
||||
body: map[string]any{"i": 123, "f": 45.67, "b": true},
|
||||
expected: "int: 123, float: 45.67, bool: true",
|
||||
},
|
||||
{
|
||||
name: "complex type slice",
|
||||
template: "items: {{items}}",
|
||||
body: map[string]any{"items": []string{"a", "b"}},
|
||||
expected: `items: ["a","b"]`,
|
||||
},
|
||||
{
|
||||
name: "complex type map",
|
||||
template: "obj: {{obj}}",
|
||||
body: map[string]any{"obj": map[string]any{"key": "value"}},
|
||||
expected: `obj: {"key":"value"}`,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
result := ParseTemplate(tt.template, tt.body)
|
||||
assert.Equal(t, tt.expected, result)
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -11,8 +11,8 @@ import (
|
||||
"strings"
|
||||
"sync"
|
||||
|
||||
pkgpush "github.com/Rain-kl/Wavelet/pkg/push"
|
||||
"github.com/Rain-kl/Wavelet/pkg/response"
|
||||
pkgpush "github.com/Rain-kl/Wavelet/plugins/domain/message_gateway/push"
|
||||
"github.com/gin-gonic/gin"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
@@ -11,7 +11,7 @@ import (
|
||||
|
||||
"github.com/Rain-kl/Wavelet/pkg/logger"
|
||||
|
||||
pkgpush "github.com/Rain-kl/Wavelet/pkg/push"
|
||||
pkgpush "github.com/Rain-kl/Wavelet/plugins/domain/message_gateway/push"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/pkg/util"
|
||||
"gorm.io/gorm"
|
||||
|
||||
@@ -9,7 +9,7 @@ import (
|
||||
"net/http"
|
||||
"strconv"
|
||||
|
||||
pkgpush "github.com/Rain-kl/Wavelet/pkg/push"
|
||||
pkgpush "github.com/Rain-kl/Wavelet/plugins/domain/message_gateway/push"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/pkg/response"
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
@@ -12,9 +12,9 @@ import (
|
||||
"strings"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/core/contracts"
|
||||
db "github.com/Rain-kl/Wavelet/pkg/persistence"
|
||||
pkgpush "github.com/Rain-kl/Wavelet/pkg/push"
|
||||
"github.com/Rain-kl/Wavelet/pkg/task"
|
||||
pkgpush "github.com/Rain-kl/Wavelet/plugins/domain/message_gateway/push"
|
||||
"github.com/Rain-kl/Wavelet/plugins/drivers/driver_asynq_worker"
|
||||
db "github.com/Rain-kl/Wavelet/plugins/infra/database"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
@@ -445,7 +445,7 @@ func findBuiltInEvent(key string) (EventMetadata, bool) {
|
||||
|
||||
func getEventInfo(req CreatePushEventRequest) (string, string, []byte, error) {
|
||||
if req.TaskType != "" {
|
||||
meta := task.GetTaskMetaByAsynqTask(req.TaskType)
|
||||
meta := driver_asynq_worker.GetTaskMetaByAsynqTask(req.TaskType)
|
||||
if meta == nil {
|
||||
return "", "", nil, errors.New("unsupported task type")
|
||||
}
|
||||
@@ -484,7 +484,7 @@ func enqueuePushTask(ctx context.Context, payload SendPayload) error {
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
_, err = task.DispatchTask(ctx, "send_notification", payloadBytes, "system")
|
||||
_, err = driver_asynq_worker.DispatchTask(ctx, "send_notification", payloadBytes, "system")
|
||||
return err
|
||||
}
|
||||
|
||||
|
||||
@@ -10,15 +10,15 @@ import (
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/pkg/logger"
|
||||
"github.com/Rain-kl/Wavelet/pkg/task"
|
||||
"github.com/Rain-kl/Wavelet/plugins/drivers/driver_asynq_worker"
|
||||
)
|
||||
|
||||
// RegisterTaskListeners subscribes push notification handlers to task completion events.
|
||||
func RegisterTaskListeners() {
|
||||
task.OnTaskCompleted(handleTaskCompleted)
|
||||
driver_asynq_worker.OnTaskCompleted(handleTaskCompleted)
|
||||
}
|
||||
|
||||
func handleTaskCompleted(ctx context.Context, execution *task.TaskExecution, result *task.TaskResult, execErr error) {
|
||||
func handleTaskCompleted(ctx context.Context, execution *driver_asynq_worker.TaskExecution, result *driver_asynq_worker.TaskResult, execErr error) {
|
||||
events, err := listActivePushEventsByTaskType(ctx, execution.TaskType)
|
||||
if err != nil {
|
||||
logger.ErrorF(ctx, "push_task_completed_listener: failed to query push events for task type %s: %v", execution.TaskType, err)
|
||||
|
||||
@@ -9,8 +9,8 @@ import (
|
||||
"errors"
|
||||
"fmt"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/pkg/push"
|
||||
"github.com/Rain-kl/Wavelet/pkg/task"
|
||||
"github.com/Rain-kl/Wavelet/plugins/domain/message_gateway/push"
|
||||
"github.com/Rain-kl/Wavelet/plugins/drivers/driver_asynq_worker"
|
||||
)
|
||||
|
||||
const (
|
||||
@@ -21,16 +21,16 @@ const (
|
||||
)
|
||||
|
||||
// SendNotificationMeta represents the task metadata.
|
||||
var SendNotificationMeta = task.TaskMeta{
|
||||
var SendNotificationMeta = driver_asynq_worker.TaskMeta{
|
||||
Type: TaskTypeSendNotification,
|
||||
AsynqTask: SendNotificationTask,
|
||||
Name: "推送通知",
|
||||
Description: "异步执行系统通知的多渠道派发与推送",
|
||||
SupportsTime: false,
|
||||
MaxRetry: task.DefaultMaxRetry,
|
||||
Queue: task.QueueDefault,
|
||||
MaxRetry: driver_asynq_worker.DefaultMaxRetry,
|
||||
Queue: driver_asynq_worker.QueueDefault,
|
||||
Retryable: true,
|
||||
Params: []task.TaskParam{
|
||||
Params: []driver_asynq_worker.TaskParam{
|
||||
{
|
||||
Name: "event_key",
|
||||
Label: "事件标识",
|
||||
@@ -69,20 +69,20 @@ func (h *PushHandler) ValidatePayload(payload []byte) ([]byte, error) {
|
||||
}
|
||||
|
||||
// Execute performs the push send and logs delivery history audit.
|
||||
func (h *PushHandler) Execute(ctx context.Context, payload []byte) (*task.TaskResult, error) {
|
||||
func (h *PushHandler) Execute(ctx context.Context, payload []byte) (*driver_asynq_worker.TaskResult, error) {
|
||||
var req SendPayload
|
||||
if err := json.Unmarshal(payload, &req); err != nil {
|
||||
task.AppendLog(ctx, "解析推送参数失败: %v", err)
|
||||
driver_asynq_worker.AppendLog(ctx, "解析推送参数失败: %v", err)
|
||||
return nil, fmt.Errorf("parse payload failed: %w", err)
|
||||
}
|
||||
|
||||
task.AppendLog(ctx, "开始推送通知: 事件 = %s, 渠道 = %s, 接收目标 = %s", req.EventKey, req.Config.Channel, req.Target)
|
||||
driver_asynq_worker.AppendLog(ctx, "开始推送通知: 事件 = %s, 渠道 = %s, 接收目标 = %s", req.EventKey, req.Config.Channel, req.Target)
|
||||
|
||||
pusher, err := push.GetPusher(req.Config.Channel)
|
||||
if err != nil {
|
||||
errWrap := fmt.Errorf("get pusher failed: %w", err)
|
||||
task.AppendLog(ctx, "推送失败: %v", errWrap)
|
||||
if task.IsFinalAttempt(ctx) {
|
||||
driver_asynq_worker.AppendLog(ctx, "推送失败: %v", errWrap)
|
||||
if driver_asynq_worker.IsFinalAttempt(ctx) {
|
||||
h.recordHistory(ctx, req, "failed", errWrap.Error())
|
||||
}
|
||||
return nil, errWrap
|
||||
@@ -95,29 +95,29 @@ func (h *PushHandler) Execute(ctx context.Context, payload []byte) (*task.TaskRe
|
||||
content := req.Body.Content
|
||||
|
||||
if err != nil {
|
||||
task.AppendLog(ctx, "消息推送失败 (标题: %s): %v", title, err)
|
||||
driver_asynq_worker.AppendLog(ctx, "消息推送失败 (标题: %s): %v", title, err)
|
||||
if upstreamResp != "" {
|
||||
task.AppendLog(ctx, "上游返回: %s", upstreamResp)
|
||||
driver_asynq_worker.AppendLog(ctx, "上游返回: %s", upstreamResp)
|
||||
}
|
||||
if task.IsFinalAttempt(ctx) {
|
||||
if driver_asynq_worker.IsFinalAttempt(ctx) {
|
||||
h.recordHistory(ctx, req, "failed", err.Error())
|
||||
}
|
||||
return nil, fmt.Errorf("pusher.Send failed: %w", err)
|
||||
}
|
||||
|
||||
task.AppendLog(ctx, "消息推送成功 (标题: %s, 内容摘要: %s)", title, content)
|
||||
driver_asynq_worker.AppendLog(ctx, "消息推送成功 (标题: %s, 内容摘要: %s)", title, content)
|
||||
if upstreamResp != "" {
|
||||
task.AppendLog(ctx, "上游返回: %s", upstreamResp)
|
||||
driver_asynq_worker.AppendLog(ctx, "上游返回: %s", upstreamResp)
|
||||
}
|
||||
h.recordHistory(ctx, req, "success", "")
|
||||
|
||||
return &task.TaskResult{
|
||||
return &driver_asynq_worker.TaskResult{
|
||||
Message: fmt.Sprintf("推送成功: [%s] -> %s", req.Config.Channel, req.Target),
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (h *PushHandler) recordHistory(ctx context.Context, req SendPayload, status string, errMsg string) {
|
||||
if dbErr := recordPushHistory(ctx, req, status, errMsg); dbErr != nil {
|
||||
task.AppendLog(ctx, "写入推送历史审计记录失败: %v", dbErr)
|
||||
driver_asynq_worker.AppendLog(ctx, "写入推送历史审计记录失败: %v", dbErr)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,26 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package message_gateway
|
||||
|
||||
import "sync"
|
||||
|
||||
var (
|
||||
factoriesMu sync.RWMutex
|
||||
factories = map[string]Factory{}
|
||||
)
|
||||
|
||||
// Register stores a channel factory under typ.
|
||||
func Register(typ string, fn Factory) {
|
||||
factoriesMu.Lock()
|
||||
defer factoriesMu.Unlock()
|
||||
factories[typ] = fn
|
||||
}
|
||||
|
||||
// Lookup returns a previously registered factory.
|
||||
func Lookup(typ string) (Factory, bool) {
|
||||
factoriesMu.RLock()
|
||||
defer factoriesMu.RUnlock()
|
||||
fn, ok := factories[typ]
|
||||
return fn, ok
|
||||
}
|
||||
@@ -0,0 +1,38 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package message_gateway
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
)
|
||||
|
||||
type stubChannel struct{}
|
||||
|
||||
func (stubChannel) Type() string { return "stub" }
|
||||
func (stubChannel) Connect(context.Context) error {
|
||||
return nil
|
||||
}
|
||||
func (stubChannel) Disconnect(context.Context) error { return nil }
|
||||
func (stubChannel) Send(context.Context, Recipient, OutboundMessage) error {
|
||||
return nil
|
||||
}
|
||||
func (stubChannel) Capabilities() Capability { return Capability{Text: true} }
|
||||
|
||||
func TestRegisterLookup(t *testing.T) {
|
||||
Register("stub", func(ChannelConfig, Handler) (Channel, error) {
|
||||
return stubChannel{}, nil
|
||||
})
|
||||
fn, ok := Lookup("stub")
|
||||
if !ok {
|
||||
t.Fatal("expected factory")
|
||||
}
|
||||
ch, err := fn(ChannelConfig{}, nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if ch.Type() != "stub" {
|
||||
t.Fatalf("type=%s", ch.Type())
|
||||
}
|
||||
}
|
||||
@@ -10,8 +10,9 @@ import (
|
||||
|
||||
"gorm.io/gorm"
|
||||
|
||||
db "github.com/Rain-kl/Wavelet/pkg/persistence"
|
||||
"github.com/Rain-kl/Wavelet/pkg/persistence/idgen"
|
||||
"github.com/Rain-kl/Wavelet/pkg/idgen"
|
||||
cachepkg "github.com/Rain-kl/Wavelet/plugins/infra/cache"
|
||||
db "github.com/Rain-kl/Wavelet/plugins/infra/database"
|
||||
)
|
||||
|
||||
const (
|
||||
@@ -223,8 +224,8 @@ func DeletePushChannelRecord(ctx context.Context, channel *PushChannel) error {
|
||||
func GetActivePushChannelByName(ctx context.Context, name string) (*PushChannel, error) {
|
||||
cacheKey := "push:channel:active:" + name
|
||||
var channel PushChannel
|
||||
if db.Redis != nil {
|
||||
if err := db.GetJSON(ctx, cacheKey, &channel); err == nil {
|
||||
if cachepkg.Redis != nil {
|
||||
if err := cachepkg.GetJSON(ctx, cacheKey, &channel); err == nil {
|
||||
return &channel, nil
|
||||
}
|
||||
}
|
||||
@@ -233,8 +234,8 @@ func GetActivePushChannelByName(ctx context.Context, name string) (*PushChannel,
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if db.Redis != nil {
|
||||
_ = db.SetJSON(ctx, cacheKey, channel, activePushChannelCacheTTL)
|
||||
if cachepkg.Redis != nil {
|
||||
_ = cachepkg.SetJSON(ctx, cacheKey, channel, activePushChannelCacheTTL)
|
||||
}
|
||||
|
||||
return &channel, nil
|
||||
@@ -242,8 +243,8 @@ func GetActivePushChannelByName(ctx context.Context, name string) (*PushChannel,
|
||||
|
||||
// DeleteActivePushChannelCache 清理启用消息通道的缓存。
|
||||
func DeleteActivePushChannelCache(ctx context.Context, name string) {
|
||||
if db.Redis != nil {
|
||||
_ = db.Redis.Del(ctx, db.PrefixedKey("push:channel:active:"+name)).Err()
|
||||
if cachepkg.Redis != nil {
|
||||
_ = cachepkg.Redis.Del(ctx, cachepkg.PrefixedKey("push:channel:active:"+name)).Err()
|
||||
}
|
||||
}
|
||||
|
||||
@@ -333,8 +334,8 @@ func ListActivePushEventsByTaskTypeRecord(ctx context.Context, taskType string)
|
||||
func GetActivePushEventByKey(ctx context.Context, key string) (*PushEvent, error) {
|
||||
cacheKey := "push:event:active:" + key
|
||||
var event PushEvent
|
||||
if db.Redis != nil {
|
||||
if err := db.GetJSON(ctx, cacheKey, &event); err == nil {
|
||||
if cachepkg.Redis != nil {
|
||||
if err := cachepkg.GetJSON(ctx, cacheKey, &event); err == nil {
|
||||
return &event, nil
|
||||
}
|
||||
}
|
||||
@@ -343,8 +344,8 @@ func GetActivePushEventByKey(ctx context.Context, key string) (*PushEvent, error
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if db.Redis != nil {
|
||||
_ = db.SetJSON(ctx, cacheKey, event, activePushEventCacheTTL)
|
||||
if cachepkg.Redis != nil {
|
||||
_ = cachepkg.SetJSON(ctx, cacheKey, event, activePushEventCacheTTL)
|
||||
}
|
||||
|
||||
return &event, nil
|
||||
@@ -352,8 +353,8 @@ func GetActivePushEventByKey(ctx context.Context, key string) (*PushEvent, error
|
||||
|
||||
// DeleteActivePushEventCache 清理启用通知事件的缓存。
|
||||
func DeleteActivePushEventCache(ctx context.Context, key string) {
|
||||
if db.Redis != nil {
|
||||
_ = db.Redis.Del(ctx, db.PrefixedKey("push:event:active:"+key)).Err()
|
||||
if cachepkg.Redis != nil {
|
||||
_ = cachepkg.Redis.Del(ctx, cachepkg.PrefixedKey("push:event:active:"+key)).Err()
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -8,10 +8,9 @@ import (
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/pkg/batchwriter"
|
||||
"github.com/Rain-kl/Wavelet/pkg/logger"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/pkg/persistence/batchwriter"
|
||||
"github.com/Rain-kl/Wavelet/pkg/persistence/logstore"
|
||||
"github.com/Rain-kl/Wavelet/plugins/domain/risk_control/logstore"
|
||||
)
|
||||
|
||||
var (
|
||||
|
||||
@@ -0,0 +1,119 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package logstore provides data access for analytics tables.
|
||||
package logstore
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/pkg/util"
|
||||
db "github.com/Rain-kl/Wavelet/plugins/infra/database"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
// CountAccessLogs returns the number of access logs matching filter.
|
||||
func CountAccessLogs(ctx context.Context, filter AccessLogFilter) (uint64, error) {
|
||||
ch := db.ChDB(ctx)
|
||||
if ch == nil {
|
||||
return 0, fmt.Errorf("clickhouse gorm connection is not initialized")
|
||||
}
|
||||
|
||||
var count int64
|
||||
query := applyFilter(ch.Model(&UserAccessLog{}), filter)
|
||||
if err := query.Count(&count).Error; err != nil {
|
||||
return 0, fmt.Errorf("count access logs: %w", err)
|
||||
}
|
||||
return safeUint64Count(count), nil
|
||||
}
|
||||
|
||||
// ListAccessLogs returns paginated access logs and the total match count.
|
||||
func ListAccessLogs(ctx context.Context, filter AccessLogFilter, page, pageSize int) ([]UserAccessLog, uint64, error) {
|
||||
ch := db.ChDB(ctx)
|
||||
if ch == nil {
|
||||
return nil, 0, fmt.Errorf("clickhouse gorm connection is not initialized")
|
||||
}
|
||||
|
||||
if filter.UserIDs != nil && len(filter.UserIDs) == 0 {
|
||||
return []UserAccessLog{}, 0, nil
|
||||
}
|
||||
|
||||
var total int64
|
||||
baseQuery := applyFilter(ch.Model(&UserAccessLog{}), filter)
|
||||
if err := baseQuery.Count(&total).Error; err != nil {
|
||||
return nil, 0, fmt.Errorf("count access logs: %w", err)
|
||||
}
|
||||
if total == 0 {
|
||||
return []UserAccessLog{}, 0, nil
|
||||
}
|
||||
|
||||
if page < 1 {
|
||||
page = 1
|
||||
}
|
||||
if pageSize < 1 {
|
||||
pageSize = 20
|
||||
}
|
||||
offset := (page - 1) * pageSize
|
||||
|
||||
var logs []UserAccessLog
|
||||
err := applyFilter(ch.Model(&UserAccessLog{}), filter).
|
||||
Order("created_at DESC, id DESC").
|
||||
Limit(pageSize).
|
||||
Offset(offset).
|
||||
Find(&logs).Error
|
||||
if err != nil {
|
||||
return nil, 0, fmt.Errorf("list access logs: %w", err)
|
||||
}
|
||||
|
||||
return logs, safeUint64Count(total), nil
|
||||
}
|
||||
|
||||
// DeleteAllUserAccessLogs hard-deletes all user access logs via TRUNCATE.
|
||||
func DeleteAllUserAccessLogs(ctx context.Context) (int64, error) {
|
||||
if db.ChConn == nil {
|
||||
return 0, fmt.Errorf("clickhouse connection is not initialized")
|
||||
}
|
||||
if err := db.ChConn.Exec(ctx, "TRUNCATE TABLE "+UserAccessLog{}.TableName()); err != nil {
|
||||
return 0, fmt.Errorf("truncate user access logs: %w", err)
|
||||
}
|
||||
return 0, nil
|
||||
}
|
||||
|
||||
// DeleteUserAccessLogsBefore deletes user access logs older than cutoff.
|
||||
func DeleteUserAccessLogsBefore(ctx context.Context, cutoff time.Time) (int64, error) {
|
||||
if db.ChConn == nil {
|
||||
return 0, fmt.Errorf("clickhouse connection is not initialized")
|
||||
}
|
||||
if err := db.ChConn.Exec(ctx, "ALTER TABLE "+UserAccessLog{}.TableName()+" DELETE WHERE created_at < ?", cutoff); err != nil {
|
||||
return 0, fmt.Errorf("delete expired user access logs: %w", err)
|
||||
}
|
||||
return 0, nil
|
||||
}
|
||||
|
||||
func safeUint64Count(count int64) uint64 {
|
||||
if count < 0 {
|
||||
return 0
|
||||
}
|
||||
return uint64(count)
|
||||
}
|
||||
|
||||
func applyFilter(query *gorm.DB, filter AccessLogFilter) *gorm.DB {
|
||||
if filter.UserIDs != nil {
|
||||
if len(filter.UserIDs) == 0 {
|
||||
return query.Where("1 = 0")
|
||||
}
|
||||
query = query.Where("user_id IN ?", filter.UserIDs)
|
||||
}
|
||||
if filter.Path != "" {
|
||||
query = query.Where("path LIKE ?", "%"+util.EscapeLike(filter.Path)+"%")
|
||||
}
|
||||
if filter.StartTime != nil {
|
||||
query = query.Where("created_at >= ?", *filter.StartTime)
|
||||
}
|
||||
if filter.EndTime != nil {
|
||||
query = query.Where("created_at <= ?", *filter.EndTime)
|
||||
}
|
||||
return query
|
||||
}
|
||||
@@ -0,0 +1,35 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package logstore
|
||||
|
||||
import "time"
|
||||
|
||||
// AccessLogFilter scopes user access log queries.
|
||||
type AccessLogFilter struct {
|
||||
// UserIDs filters by user IDs. nil means no user filter; an empty slice means no matches.
|
||||
UserIDs []uint64
|
||||
Path string
|
||||
// StartTime filters created_at >= StartTime when non-nil.
|
||||
StartTime *time.Time
|
||||
// EndTime filters created_at <= EndTime when non-nil.
|
||||
EndTime *time.Time
|
||||
}
|
||||
|
||||
// DailyTrend is a single day's access count.
|
||||
type DailyTrend struct {
|
||||
Date string
|
||||
Count uint64
|
||||
}
|
||||
|
||||
// BrowserShare is a browser group's share of access logs.
|
||||
type BrowserShare struct {
|
||||
Browser string
|
||||
Count uint64
|
||||
}
|
||||
|
||||
// TopUser is an active user ranked by access count.
|
||||
type TopUser struct {
|
||||
UserID uint64
|
||||
Count uint64
|
||||
}
|
||||
@@ -0,0 +1,140 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package logstore
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"sort"
|
||||
"time"
|
||||
|
||||
db "github.com/Rain-kl/Wavelet/plugins/infra/database"
|
||||
)
|
||||
|
||||
const hoursInDay = 24
|
||||
|
||||
// GetDailyTrend returns per-day access counts for the last days days (inclusive of today).
|
||||
func GetDailyTrend(ctx context.Context, days int) ([]DailyTrend, error) {
|
||||
if days < 1 {
|
||||
days = 7
|
||||
}
|
||||
|
||||
ch := db.ChDB(ctx)
|
||||
if ch == nil {
|
||||
return nil, fmt.Errorf("clickhouse gorm connection is not initialized")
|
||||
}
|
||||
|
||||
startTime := time.Now().AddDate(0, 0, -(days - 1)).Truncate(hoursInDay * time.Hour)
|
||||
tableName := UserAccessLog{}.TableName()
|
||||
|
||||
query := fmt.Sprintf(`
|
||||
SELECT toDate(created_at) AS date, count() AS count
|
||||
FROM %s
|
||||
WHERE created_at >= ?
|
||||
GROUP BY date
|
||||
ORDER BY date ASC
|
||||
`, tableName)
|
||||
|
||||
type trendRow struct {
|
||||
Date time.Time
|
||||
Count uint64
|
||||
}
|
||||
|
||||
var rows []trendRow
|
||||
if err := ch.Raw(query, startTime).Scan(&rows).Error; err != nil {
|
||||
return nil, fmt.Errorf("get daily trend: %w", err)
|
||||
}
|
||||
|
||||
trendMap := make(map[string]uint64, days)
|
||||
for i := 0; i < days; i++ {
|
||||
dateStr := time.Now().AddDate(0, 0, -i).Format("2006-01-02")
|
||||
trendMap[dateStr] = 0
|
||||
}
|
||||
for _, row := range rows {
|
||||
dateStr := row.Date.Format("2006-01-02")
|
||||
trendMap[dateStr] = row.Count
|
||||
}
|
||||
|
||||
result := make([]DailyTrend, 0, days)
|
||||
for i := days - 1; i >= 0; i-- {
|
||||
dateStr := time.Now().AddDate(0, 0, -i).Format("2006-01-02")
|
||||
result = append(result, DailyTrend{
|
||||
Date: dateStr,
|
||||
Count: trendMap[dateStr],
|
||||
})
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
|
||||
// GetBrowserDistribution returns browser-grouped access counts since startTime.
|
||||
func GetBrowserDistribution(ctx context.Context, startTime time.Time) ([]BrowserShare, error) {
|
||||
ch := db.ChDB(ctx)
|
||||
if ch == nil {
|
||||
return nil, fmt.Errorf("clickhouse gorm connection is not initialized")
|
||||
}
|
||||
|
||||
tableName := UserAccessLog{}.TableName()
|
||||
query := fmt.Sprintf(`
|
||||
SELECT user_agent, count() AS count
|
||||
FROM %s
|
||||
WHERE created_at >= ?
|
||||
GROUP BY user_agent
|
||||
`, tableName)
|
||||
|
||||
type uaRow struct {
|
||||
UserAgent string
|
||||
Count uint64
|
||||
}
|
||||
|
||||
var rows []uaRow
|
||||
if err := ch.Raw(query, startTime).Scan(&rows).Error; err != nil {
|
||||
return nil, fmt.Errorf("get browser distribution: %w", err)
|
||||
}
|
||||
|
||||
browserCounts := make(map[string]uint64)
|
||||
for _, row := range rows {
|
||||
browser := ParseBrowserName(row.UserAgent)
|
||||
browserCounts[browser] += row.Count
|
||||
}
|
||||
|
||||
result := make([]BrowserShare, 0, len(browserCounts))
|
||||
for browser, count := range browserCounts {
|
||||
result = append(result, BrowserShare{
|
||||
Browser: browser,
|
||||
Count: count,
|
||||
})
|
||||
}
|
||||
sort.Slice(result, func(i, j int) bool {
|
||||
return result[i].Count > result[j].Count
|
||||
})
|
||||
return result, nil
|
||||
}
|
||||
|
||||
// GetTopActiveUsers returns the most active users since startTime.
|
||||
func GetTopActiveUsers(ctx context.Context, startTime time.Time, limit int) ([]TopUser, error) {
|
||||
if limit < 1 {
|
||||
limit = 10
|
||||
}
|
||||
|
||||
ch := db.ChDB(ctx)
|
||||
if ch == nil {
|
||||
return nil, fmt.Errorf("clickhouse gorm connection is not initialized")
|
||||
}
|
||||
|
||||
tableName := UserAccessLog{}.TableName()
|
||||
query := fmt.Sprintf(`
|
||||
SELECT user_id, count() AS count
|
||||
FROM %s
|
||||
WHERE created_at >= ? AND user_id > 0
|
||||
GROUP BY user_id
|
||||
ORDER BY count DESC
|
||||
LIMIT ?
|
||||
`, tableName)
|
||||
|
||||
var users []TopUser
|
||||
if err := ch.Raw(query, startTime, limit).Scan(&users).Error; err != nil {
|
||||
return nil, fmt.Errorf("get top active users: %w", err)
|
||||
}
|
||||
return users, nil
|
||||
}
|
||||
@@ -0,0 +1,215 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package logstore
|
||||
|
||||
import (
|
||||
"context"
|
||||
"io"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/ClickHouse/clickhouse-go/v2/lib/column"
|
||||
"github.com/ClickHouse/clickhouse-go/v2/lib/driver"
|
||||
|
||||
"github.com/glebarez/sqlite"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"gorm.io/gorm"
|
||||
|
||||
db "github.com/Rain-kl/Wavelet/plugins/infra/database"
|
||||
)
|
||||
|
||||
func setupChGormDB(t *testing.T) *gorm.DB {
|
||||
t.Helper()
|
||||
|
||||
gormDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{
|
||||
DisableForeignKeyConstraintWhenMigrating: true,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, gormDB.AutoMigrate(&UserAccessLog{}))
|
||||
db.SetChDBForTest(gormDB)
|
||||
return gormDB
|
||||
}
|
||||
|
||||
func TestParseBrowserName(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
ua string
|
||||
want string
|
||||
}{
|
||||
{name: "chrome", ua: "Mozilla/5.0 Chrome/120.0.0.0", want: "Chrome"},
|
||||
{name: "firefox", ua: "Mozilla/5.0 Firefox/121.0", want: "Firefox"},
|
||||
{name: "safari", ua: "Mozilla/5.0 Safari/605.1.15", want: "Safari"},
|
||||
{name: "edge", ua: "Mozilla/5.0 Edg/120.0.0.0", want: "Edge"},
|
||||
{name: "wechat", ua: "MicroMessenger/8.0", want: "WeChat"},
|
||||
{name: "postman", ua: "PostmanRuntime/7.36.0", want: "Postman"},
|
||||
{name: "other", ua: "curl/8.0", want: "Other"},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
assert.Equal(t, tt.want, ParseBrowserName(tt.ua))
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestCountAccessLogs_EmptyUserIDs(t *testing.T) {
|
||||
setupChGormDB(t)
|
||||
t.Cleanup(func() { db.SetChDBForTest(nil) })
|
||||
|
||||
count, err := CountAccessLogs(context.Background(), AccessLogFilter{UserIDs: []uint64{}})
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, uint64(0), count)
|
||||
}
|
||||
|
||||
func TestListAccessLogs_EmptyUserIDs(t *testing.T) {
|
||||
setupChGormDB(t)
|
||||
t.Cleanup(func() { db.SetChDBForTest(nil) })
|
||||
|
||||
logs, total, err := ListAccessLogs(context.Background(), AccessLogFilter{UserIDs: []uint64{}}, 1, 20)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, uint64(0), total)
|
||||
assert.Empty(t, logs)
|
||||
}
|
||||
|
||||
func TestListAccessLogs_WithFilters(t *testing.T) {
|
||||
gormDB := setupChGormDB(t)
|
||||
t.Cleanup(func() { db.SetChDBForTest(nil) })
|
||||
|
||||
now := time.Now().UTC().Truncate(time.Second)
|
||||
logs := []UserAccessLog{
|
||||
{ID: 1, UserID: 10, Path: "/api/v1/users", Method: "GET", Status: 200, CreatedAt: now},
|
||||
{ID: 2, UserID: 20, Path: "/api/v1/admin/logs", Method: "GET", Status: 200, CreatedAt: now},
|
||||
{ID: 3, UserID: 10, Path: "/api/v1/other", Method: "POST", Status: 201, CreatedAt: now},
|
||||
}
|
||||
require.NoError(t, gormDB.Create(&logs).Error)
|
||||
|
||||
start := now.Add(-time.Hour)
|
||||
filter := AccessLogFilter{
|
||||
UserIDs: []uint64{10},
|
||||
Path: "users",
|
||||
StartTime: &start,
|
||||
}
|
||||
|
||||
count, err := CountAccessLogs(context.Background(), filter)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, uint64(1), count)
|
||||
|
||||
result, total, err := ListAccessLogs(context.Background(), filter, 1, 10)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, uint64(1), total)
|
||||
require.Len(t, result, 1)
|
||||
assert.Equal(t, uint64(1), result[0].ID)
|
||||
assert.Equal(t, "/api/v1/users", result[0].Path)
|
||||
}
|
||||
|
||||
func TestBatchInsert_Empty(t *testing.T) {
|
||||
err := BatchInsert(context.Background(), nil)
|
||||
require.NoError(t, err)
|
||||
}
|
||||
|
||||
func TestBatchInsert_UsesModelBatchSQL(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
mockBatch := &mockBatch{}
|
||||
mockConn := &mockConn{
|
||||
batch: mockBatch,
|
||||
batchQuery: UserAccessLog{}.BatchInsertSQL(),
|
||||
}
|
||||
db.SetChConnForTest(mockConn)
|
||||
t.Cleanup(func() { db.SetChConnForTest(nil) })
|
||||
|
||||
createdAt := time.Now().UTC()
|
||||
err := BatchInsert(ctx, []UserAccessLog{
|
||||
{
|
||||
ID: 1,
|
||||
UserID: 42,
|
||||
Path: "/api/v1/test",
|
||||
Method: "GET",
|
||||
IP: "127.0.0.1",
|
||||
UserAgent: "test-agent",
|
||||
Headers: "{}",
|
||||
Status: 200,
|
||||
Latency: 12,
|
||||
CreatedAt: createdAt,
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
assert.True(t, mockConn.prepareCalled)
|
||||
assert.Equal(t, UserAccessLog{}.BatchInsertSQL(), mockConn.preparedQuery)
|
||||
assert.True(t, mockBatch.sendCalled)
|
||||
require.Len(t, mockBatch.rows, 1)
|
||||
assert.Equal(t, uint64(42), mockBatch.rows[0][1])
|
||||
}
|
||||
|
||||
type mockConn struct {
|
||||
batch driver.Batch
|
||||
batchQuery string
|
||||
prepareCalled bool
|
||||
preparedQuery string
|
||||
}
|
||||
|
||||
func (m *mockConn) Contributors() []string { return nil }
|
||||
|
||||
func (m *mockConn) ServerVersion() (*driver.ServerVersion, error) { return nil, nil }
|
||||
|
||||
func (m *mockConn) Select(_ context.Context, _ any, _ string, _ ...any) error { return nil }
|
||||
|
||||
func (m *mockConn) Query(_ context.Context, _ string, _ ...any) (driver.Rows, error) {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
func (m *mockConn) QueryRow(_ context.Context, _ string, _ ...any) driver.Row { return nil }
|
||||
|
||||
func (m *mockConn) PrepareBatch(_ context.Context, query string, _ ...driver.PrepareBatchOption) (driver.Batch, error) {
|
||||
m.prepareCalled = true
|
||||
m.preparedQuery = query
|
||||
return m.batch, nil
|
||||
}
|
||||
|
||||
func (m *mockConn) Exec(_ context.Context, _ string, _ ...any) error { return nil }
|
||||
|
||||
func (m *mockConn) AsyncInsert(_ context.Context, _ string, _ bool, _ ...any) error { return nil }
|
||||
|
||||
func (m *mockConn) InsertFormat(_ context.Context, _ string, _ string, _ io.Reader) error { return nil }
|
||||
|
||||
func (m *mockConn) QueryFormat(_ context.Context, _ string, _ string, _ ...any) (io.ReadCloser, error) {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
func (m *mockConn) Ping(_ context.Context) error { return nil }
|
||||
|
||||
func (m *mockConn) Stats() driver.Stats { return driver.Stats{} }
|
||||
|
||||
func (m *mockConn) Close() error { return nil }
|
||||
|
||||
type mockBatch struct {
|
||||
rows [][]any
|
||||
sendCalled bool
|
||||
}
|
||||
|
||||
func (m *mockBatch) Abort() error { return nil }
|
||||
|
||||
func (m *mockBatch) Append(v ...any) error {
|
||||
m.rows = append(m.rows, v)
|
||||
return nil
|
||||
}
|
||||
|
||||
func (m *mockBatch) AppendStruct(_ any) error { return nil }
|
||||
|
||||
func (m *mockBatch) Column(_ int) driver.BatchColumn { return nil }
|
||||
|
||||
func (m *mockBatch) Flush() error { return nil }
|
||||
|
||||
func (m *mockBatch) Send() error {
|
||||
m.sendCalled = true
|
||||
return nil
|
||||
}
|
||||
|
||||
func (m *mockBatch) IsSent() bool { return m.sendCalled }
|
||||
|
||||
func (m *mockBatch) Rows() int { return len(m.rows) }
|
||||
|
||||
func (m *mockBatch) Columns() []column.Interface { return nil }
|
||||
|
||||
func (m *mockBatch) Close() error { return nil }
|
||||
@@ -0,0 +1,48 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package logstore
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
|
||||
db "github.com/Rain-kl/Wavelet/plugins/infra/database"
|
||||
)
|
||||
|
||||
// BatchInsert writes access logs to ClickHouse using the native batch API.
|
||||
func BatchInsert(ctx context.Context, logs []UserAccessLog) error {
|
||||
if len(logs) == 0 {
|
||||
return nil
|
||||
}
|
||||
if db.ChConn == nil {
|
||||
return fmt.Errorf("clickhouse connection is not initialized")
|
||||
}
|
||||
|
||||
batch, err := db.ChConn.PrepareBatch(ctx, UserAccessLog{}.BatchInsertSQL())
|
||||
if err != nil {
|
||||
return fmt.Errorf("prepare clickhouse batch: %w", err)
|
||||
}
|
||||
|
||||
for _, logItem := range logs {
|
||||
if err := batch.Append(
|
||||
logItem.ID,
|
||||
logItem.UserID,
|
||||
logItem.Path,
|
||||
logItem.Method,
|
||||
logItem.IP,
|
||||
logItem.UserAgent,
|
||||
logItem.Headers,
|
||||
logItem.Status,
|
||||
logItem.Latency,
|
||||
logItem.CreatedAt,
|
||||
); err != nil {
|
||||
return fmt.Errorf("append access log to batch: %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
if err := batch.Send(); err != nil {
|
||||
return fmt.Errorf("send clickhouse batch: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,30 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package logstore
|
||||
|
||||
import "strings"
|
||||
|
||||
// ParseBrowserName performs lightweight User-Agent browser identification.
|
||||
func ParseBrowserName(ua string) string {
|
||||
uaLower := strings.ToLower(ua)
|
||||
if strings.Contains(uaLower, "micromessenger") {
|
||||
return "WeChat"
|
||||
}
|
||||
if strings.Contains(uaLower, "postman") {
|
||||
return "Postman"
|
||||
}
|
||||
if strings.Contains(uaLower, "edg/") || strings.Contains(uaLower, "edge") {
|
||||
return "Edge"
|
||||
}
|
||||
if strings.Contains(uaLower, "firefox") {
|
||||
return "Firefox"
|
||||
}
|
||||
if strings.Contains(uaLower, "chrome") {
|
||||
return "Chrome"
|
||||
}
|
||||
if strings.Contains(uaLower, "safari") {
|
||||
return "Safari"
|
||||
}
|
||||
return "Other"
|
||||
}
|
||||
@@ -0,0 +1,99 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package logstore
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"strconv"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/pkg/logger"
|
||||
)
|
||||
|
||||
const (
|
||||
defaultLogRetentionDays = 30
|
||||
partitionLeadMonths = 2
|
||||
userAccessLogTable = "w_user_access_logs"
|
||||
)
|
||||
|
||||
// CleanupSummary 汇总本次清理结果。
|
||||
type CleanupSummary struct {
|
||||
ActiveDatabase string `json:"active_database"`
|
||||
RetentionDays int `json:"retention_days"`
|
||||
Deleted int64 `json:"deleted"`
|
||||
}
|
||||
|
||||
// CleanupExpired 按当前日志库保留天数删除过期用户访问日志,并预建 PG 分区。
|
||||
func CleanupExpired(ctx context.Context) (CleanupSummary, error) {
|
||||
active, err := ActiveDatabase(ctx)
|
||||
if err != nil {
|
||||
return CleanupSummary{}, err
|
||||
}
|
||||
days := retentionDaysForDatabase(ctx, active)
|
||||
summary := CleanupSummary{ActiveDatabase: active, RetentionDays: days}
|
||||
|
||||
store, err := Active(ctx)
|
||||
if err != nil {
|
||||
return summary, err
|
||||
}
|
||||
now := time.Now().UTC()
|
||||
if err := store.UserAccessLogs.EnsurePartitions(ctx, now, now.AddDate(0, partitionLeadMonths, 0)); err != nil {
|
||||
logger.WarnF(ctx, "logstore: ensure partitions during cleanup failed: %v", err)
|
||||
}
|
||||
cutoff := now.AddDate(0, 0, -days)
|
||||
// 先 DROP 完全过期的整月分区,再对边界月逐行 DeleteBefore。
|
||||
if err := store.UserAccessLogs.DropExpiredPartitions(ctx, cutoff); err != nil {
|
||||
return summary, fmt.Errorf("drop expired partitions: %w", err)
|
||||
}
|
||||
deleted, err := store.UserAccessLogs.DeleteBefore(ctx, cutoff)
|
||||
if err != nil {
|
||||
return summary, fmt.Errorf("delete expired user access logs: %w", err)
|
||||
}
|
||||
summary.Deleted = deleted
|
||||
if err := store.UserAccessLogs.DropEmptyPartitions(ctx, now); err != nil {
|
||||
logger.WarnF(ctx, "drop empty log partitions failed: %v", err)
|
||||
}
|
||||
return summary, nil
|
||||
}
|
||||
|
||||
func retentionDaysForDatabase(ctx context.Context, dbName string) int {
|
||||
key := "log_retention_days_postgres"
|
||||
switch dbName {
|
||||
case dbNameSQLite:
|
||||
key = "log_retention_days_sqlite"
|
||||
case dbNameClickHouse:
|
||||
key = "log_retention_days_clickhouse"
|
||||
}
|
||||
v, err := getConfig(ctx, key)
|
||||
if err != nil {
|
||||
if !errors.Is(err, errConfigReaderNotWired) {
|
||||
logger.ErrorF(ctx, "读取日志保留天数配置失败(key=%s),回退默认 %d 天: %v", key, defaultLogRetentionDays, err)
|
||||
}
|
||||
return defaultLogRetentionDays
|
||||
}
|
||||
days, perr := strconv.Atoi(v)
|
||||
if perr != nil || days <= 0 {
|
||||
logger.ErrorF(ctx, "日志保留天数配置非法(key=%s, value=%q),回退默认 %d 天", key, v, defaultLogRetentionDays)
|
||||
return defaultLogRetentionDays
|
||||
}
|
||||
return days
|
||||
}
|
||||
|
||||
func partitionStatementsRange(from, to time.Time) []string {
|
||||
var out []string
|
||||
start := time.Date(from.Year(), from.Month(), 1, 0, 0, 0, 0, time.UTC)
|
||||
end := time.Date(to.Year(), to.Month(), 1, 0, 0, 0, 0, time.UTC).AddDate(0, 1, 0)
|
||||
for ; start.Before(end); start = start.AddDate(0, 1, 0) {
|
||||
monthEnd := start.AddDate(0, 1, 0)
|
||||
suffix := start.Format("200601")
|
||||
fromDay := start.Format("2006-01-02")
|
||||
toDay := monthEnd.Format("2006-01-02")
|
||||
out = append(out, fmt.Sprintf(
|
||||
"CREATE TABLE IF NOT EXISTS %s_%s PARTITION OF %s FOR VALUES FROM ('%s') TO ('%s')",
|
||||
userAccessLogTable, suffix, userAccessLogTable, fromDay, toDay))
|
||||
}
|
||||
return out
|
||||
}
|
||||
@@ -0,0 +1,152 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package logstore
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"time"
|
||||
|
||||
"github.com/ClickHouse/clickhouse-go/v2/lib/driver"
|
||||
db "github.com/Rain-kl/Wavelet/plugins/infra/database"
|
||||
)
|
||||
|
||||
type clickhouseUserAccessLogStore struct {
|
||||
skipFreeze bool
|
||||
}
|
||||
|
||||
func newClickHouseUserAccessLogStore() *clickhouseUserAccessLogStore {
|
||||
return &clickhouseUserAccessLogStore{}
|
||||
}
|
||||
|
||||
var (
|
||||
_ UserAccessLogStore = (*clickhouseUserAccessLogStore)(nil)
|
||||
_ StatusStore = (*clickhouseUserAccessLogStore)(nil)
|
||||
)
|
||||
|
||||
func (s *clickhouseUserAccessLogStore) ActiveDatabase(_ context.Context) (string, error) {
|
||||
return dbNameClickHouse, nil
|
||||
}
|
||||
|
||||
func (s *clickhouseUserAccessLogStore) ensureWritable(ctx context.Context) error {
|
||||
if !s.skipFreeze && Migrating(ctx) {
|
||||
return ErrMigrating
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *clickhouseUserAccessLogStore) BatchInsert(ctx context.Context, logs []UserAccessLog) error {
|
||||
if len(logs) == 0 {
|
||||
return nil
|
||||
}
|
||||
if err := s.ensureWritable(ctx); err != nil {
|
||||
return err
|
||||
}
|
||||
return BatchInsert(ctx, logs)
|
||||
}
|
||||
|
||||
func (s *clickhouseUserAccessLogStore) DeleteAll(ctx context.Context) (int64, error) {
|
||||
if err := s.ensureWritable(ctx); err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return DeleteAllUserAccessLogs(ctx)
|
||||
}
|
||||
|
||||
func (s *clickhouseUserAccessLogStore) DeleteBefore(ctx context.Context, cutoff time.Time) (int64, error) {
|
||||
if err := s.ensureWritable(ctx); err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return DeleteUserAccessLogsBefore(ctx, cutoff)
|
||||
}
|
||||
|
||||
func (s *clickhouseUserAccessLogStore) Count(ctx context.Context, filter AccessLogFilter) (uint64, error) {
|
||||
return CountAccessLogs(ctx, filter)
|
||||
}
|
||||
|
||||
func (s *clickhouseUserAccessLogStore) List(ctx context.Context, filter AccessLogFilter, page, pageSize int) ([]UserAccessLog, uint64, error) {
|
||||
return ListAccessLogs(ctx, filter, page, pageSize)
|
||||
}
|
||||
|
||||
func (s *clickhouseUserAccessLogStore) GetDailyTrend(ctx context.Context, days int) ([]DailyTrend, error) {
|
||||
return GetDailyTrend(ctx, days)
|
||||
}
|
||||
|
||||
func (s *clickhouseUserAccessLogStore) GetBrowserDistribution(ctx context.Context, startTime time.Time) ([]BrowserShare, error) {
|
||||
return GetBrowserDistribution(ctx, startTime)
|
||||
}
|
||||
|
||||
func (s *clickhouseUserAccessLogStore) GetTopActiveUsers(ctx context.Context, startTime time.Time, limit int) ([]TopUser, error) {
|
||||
return GetTopActiveUsers(ctx, startTime, limit)
|
||||
}
|
||||
|
||||
func (s *clickhouseUserAccessLogStore) EnsurePartitions(_ context.Context, _, _ time.Time) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *clickhouseUserAccessLogStore) DropEmptyPartitions(_ context.Context, _ time.Time) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *clickhouseUserAccessLogStore) DropExpiredPartitions(_ context.Context, _ time.Time) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *clickhouseUserAccessLogStore) MigrationRange(ctx context.Context) (time.Time, time.Time, error) {
|
||||
if db.ChConn == nil {
|
||||
return time.Time{}, time.Time{}, fmt.Errorf("clickhouse connection is not initialized")
|
||||
}
|
||||
table := UserAccessLog{}.TableName()
|
||||
var minTime, maxTime *time.Time
|
||||
if err := db.ChConn.QueryRow(ctx, "SELECT min(created_at), max(created_at) FROM "+table).Scan(&minTime, &maxTime); err != nil {
|
||||
return time.Time{}, time.Time{}, fmt.Errorf("query migration range %s: %w", table, err)
|
||||
}
|
||||
if minTime == nil || maxTime == nil {
|
||||
return time.Time{}, time.Time{}, nil
|
||||
}
|
||||
return minTime.UTC(), maxTime.UTC(), nil
|
||||
}
|
||||
|
||||
func (s *clickhouseUserAccessLogStore) ListForMigration(ctx context.Context, afterID uint64, limit int) ([]UserAccessLog, error) {
|
||||
if db.ChConn == nil {
|
||||
return nil, fmt.Errorf("clickhouse connection is not initialized")
|
||||
}
|
||||
if limit <= 0 {
|
||||
limit = migrationPageSize
|
||||
}
|
||||
table := UserAccessLog{}.TableName()
|
||||
columns := UserAccessLog{}.InsertColumns()
|
||||
rows, err := db.ChConn.Query(ctx, fmt.Sprintf(
|
||||
"SELECT %s FROM %s WHERE id > ? ORDER BY id ASC LIMIT ?",
|
||||
columns, table,
|
||||
), afterID, limit)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("list user access logs for migration: %w", err)
|
||||
}
|
||||
defer func() { _ = rows.Close() }()
|
||||
return scanUserAccessLogs(rows)
|
||||
}
|
||||
|
||||
func scanUserAccessLogs(rows driver.Rows) ([]UserAccessLog, error) {
|
||||
var result []UserAccessLog
|
||||
for rows.Next() {
|
||||
var item UserAccessLog
|
||||
if err := rows.Scan(
|
||||
&item.ID,
|
||||
&item.UserID,
|
||||
&item.Path,
|
||||
&item.Method,
|
||||
&item.IP,
|
||||
&item.UserAgent,
|
||||
&item.Headers,
|
||||
&item.Status,
|
||||
&item.Latency,
|
||||
&item.CreatedAt,
|
||||
); err != nil {
|
||||
return nil, fmt.Errorf("scan user access log row: %w", err)
|
||||
}
|
||||
item.CreatedAt = item.CreatedAt.UTC()
|
||||
result = append(result, item)
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
@@ -0,0 +1,325 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package logstore
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"sort"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/pkg/idgen"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
const (
|
||||
insertBatchSize = 500
|
||||
migrationPageSize = 100
|
||||
defaultPageSize = 20
|
||||
defaultTopN = 10
|
||||
topUserAgents = 100
|
||||
dayDuration = 24 * time.Hour
|
||||
)
|
||||
|
||||
type gormLogStore struct {
|
||||
db *gorm.DB
|
||||
skipFreeze bool
|
||||
}
|
||||
|
||||
func newGormStore(db *gorm.DB) *gormLogStore { return &gormLogStore{db: db} }
|
||||
|
||||
type userAccessLogGormStore struct {
|
||||
*gormLogStore
|
||||
}
|
||||
|
||||
func newUserAccessLogGormStore(db *gorm.DB) *userAccessLogGormStore {
|
||||
return &userAccessLogGormStore{gormLogStore: newGormStore(db)}
|
||||
}
|
||||
|
||||
var (
|
||||
_ UserAccessLogStore = (*userAccessLogGormStore)(nil)
|
||||
_ StatusStore = (*userAccessLogGormStore)(nil)
|
||||
)
|
||||
|
||||
func (s *gormLogStore) ActiveDatabase(_ context.Context) (string, error) {
|
||||
if isPostgresDialect(s.db) {
|
||||
return dbNamePostgres, nil
|
||||
}
|
||||
return dbNameSQLite, nil
|
||||
}
|
||||
|
||||
func (s *gormLogStore) ensureWritable(ctx context.Context) error {
|
||||
if !s.skipFreeze && Migrating(ctx) {
|
||||
return ErrMigrating
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *userAccessLogGormStore) BatchInsert(ctx context.Context, logs []UserAccessLog) error {
|
||||
if len(logs) == 0 {
|
||||
return nil
|
||||
}
|
||||
if err := s.ensureWritable(ctx); err != nil {
|
||||
return err
|
||||
}
|
||||
for i := range logs {
|
||||
if logs[i].ID == 0 {
|
||||
logs[i].ID = idgen.NextUint64ID()
|
||||
}
|
||||
}
|
||||
return s.db.WithContext(ctx).CreateInBatches(logs, insertBatchSize).Error
|
||||
}
|
||||
|
||||
func (s *userAccessLogGormStore) DeleteAll(ctx context.Context) (int64, error) {
|
||||
if err := s.ensureWritable(ctx); err != nil {
|
||||
return 0, err
|
||||
}
|
||||
res := s.db.WithContext(ctx).Where("1 = 1").Delete(&UserAccessLog{})
|
||||
return res.RowsAffected, res.Error
|
||||
}
|
||||
|
||||
func (s *userAccessLogGormStore) DeleteBefore(ctx context.Context, cutoff time.Time) (int64, error) {
|
||||
if err := s.ensureWritable(ctx); err != nil {
|
||||
return 0, err
|
||||
}
|
||||
res := s.db.WithContext(ctx).Where("created_at < ?", cutoff).Delete(&UserAccessLog{})
|
||||
if res.Error != nil && isMissingRelation(res.Error) {
|
||||
return 0, nil
|
||||
}
|
||||
return res.RowsAffected, res.Error
|
||||
}
|
||||
|
||||
func (s *userAccessLogGormStore) ListForMigration(ctx context.Context, afterID uint64, limit int) ([]UserAccessLog, error) {
|
||||
var rows []UserAccessLog
|
||||
q := s.db.WithContext(ctx).Model(&UserAccessLog{}).
|
||||
Where("id > ?", afterID).
|
||||
Order("id ASC").
|
||||
Limit(limitOr(limit, migrationPageSize))
|
||||
if err := q.Find(&rows).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return rows, nil
|
||||
}
|
||||
|
||||
func (s *userAccessLogGormStore) MigrationRange(ctx context.Context) (time.Time, time.Time, error) {
|
||||
return gormMigrationRange(ctx, s.db, "created_at", UserAccessLog{}, func(v *UserAccessLog) time.Time {
|
||||
return v.CreatedAt
|
||||
})
|
||||
}
|
||||
|
||||
func (s *userAccessLogGormStore) Count(ctx context.Context, filter AccessLogFilter) (uint64, error) {
|
||||
where, args, ok := buildUserAccessLogWhere(filter)
|
||||
if !ok {
|
||||
return 0, nil
|
||||
}
|
||||
var total int64
|
||||
if err := s.db.WithContext(ctx).Model(&UserAccessLog{}).Where(where, args...).Count(&total).Error; err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return countToUint64(total), nil
|
||||
}
|
||||
|
||||
func (s *userAccessLogGormStore) List(ctx context.Context, filter AccessLogFilter, page, pageSize int) ([]UserAccessLog, uint64, error) {
|
||||
where, args, ok := buildUserAccessLogWhere(filter)
|
||||
if !ok {
|
||||
return []UserAccessLog{}, 0, nil
|
||||
}
|
||||
var total int64
|
||||
if err := s.db.WithContext(ctx).Model(&UserAccessLog{}).Where(where, args...).Count(&total).Error; err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
if total == 0 {
|
||||
return []UserAccessLog{}, 0, nil
|
||||
}
|
||||
var rows []UserAccessLog
|
||||
q := s.db.WithContext(ctx).Where(where, args...).Order("created_at DESC, id DESC")
|
||||
if err := q.Limit(limitOr(pageSize, defaultPageSize)).Offset(offsetOf(page, pageSize)).Find(&rows).Error; err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
return rows, countToUint64(total), nil
|
||||
}
|
||||
|
||||
func buildUserAccessLogWhere(filter AccessLogFilter) (string, []any, bool) {
|
||||
if filter.UserIDs != nil && len(filter.UserIDs) == 0 {
|
||||
return "", nil, false
|
||||
}
|
||||
var parts []string
|
||||
var args []any
|
||||
if filter.UserIDs != nil {
|
||||
parts = append(parts, "user_id IN ?")
|
||||
args = append(args, filter.UserIDs)
|
||||
}
|
||||
if trimmed := strings.TrimSpace(filter.Path); trimmed != "" {
|
||||
parts = append(parts, "path LIKE ?")
|
||||
args = append(args, "%"+trimmed+"%")
|
||||
}
|
||||
if filter.StartTime != nil {
|
||||
parts = append(parts, "created_at >= ?")
|
||||
args = append(args, *filter.StartTime)
|
||||
}
|
||||
if filter.EndTime != nil {
|
||||
parts = append(parts, "created_at <= ?")
|
||||
args = append(args, *filter.EndTime)
|
||||
}
|
||||
if len(parts) == 0 {
|
||||
return "1 = 1", args, true
|
||||
}
|
||||
return strings.Join(parts, " AND "), args, true
|
||||
}
|
||||
|
||||
func (s *userAccessLogGormStore) GetDailyTrend(ctx context.Context, days int) ([]DailyTrend, error) {
|
||||
if days <= 0 {
|
||||
days = 7
|
||||
}
|
||||
start := time.Now().AddDate(0, 0, -(days - 1)).Truncate(dayDuration)
|
||||
type row struct {
|
||||
Date string
|
||||
Cnt uint64
|
||||
}
|
||||
var rows []row
|
||||
err := s.db.WithContext(ctx).Model(&UserAccessLog{}).
|
||||
Select(dailyTrendDateSQL(s.db)+" AS date, COUNT(*) AS cnt").
|
||||
Where("created_at >= ?", start).
|
||||
Group("date").Order("date ASC").Scan(&rows).Error
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
counts := make(map[string]uint64, len(rows))
|
||||
for _, r := range rows {
|
||||
counts[r.Date] = r.Cnt
|
||||
}
|
||||
out := make([]DailyTrend, 0, days)
|
||||
for i := 0; i < days; i++ {
|
||||
d := start.AddDate(0, 0, i).Format("2006-01-02")
|
||||
out = append(out, DailyTrend{Date: d, Count: counts[d]})
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
func (s *userAccessLogGormStore) GetBrowserDistribution(ctx context.Context, startTime time.Time) ([]BrowserShare, error) {
|
||||
type row struct {
|
||||
UserAgent string
|
||||
Cnt uint64
|
||||
}
|
||||
var rows []row
|
||||
err := s.db.WithContext(ctx).Model(&UserAccessLog{}).
|
||||
Select("user_agent, COUNT(*) AS cnt").
|
||||
Where("created_at >= ?", startTime).
|
||||
Group("user_agent").Order("cnt DESC").Limit(topUserAgents).Scan(&rows).Error
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
counts := make(map[string]uint64)
|
||||
for _, r := range rows {
|
||||
counts[ParseBrowserName(r.UserAgent)] += r.Cnt
|
||||
}
|
||||
out := make([]BrowserShare, 0, len(counts))
|
||||
for label, count := range counts {
|
||||
out = append(out, BrowserShare{Browser: label, Count: count})
|
||||
}
|
||||
sort.Slice(out, func(i, j int) bool { return out[i].Count > out[j].Count })
|
||||
return out, nil
|
||||
}
|
||||
|
||||
func (s *userAccessLogGormStore) GetTopActiveUsers(ctx context.Context, startTime time.Time, limit int) ([]TopUser, error) {
|
||||
type row struct {
|
||||
UserID uint64
|
||||
Cnt uint64
|
||||
}
|
||||
var rows []row
|
||||
err := s.db.WithContext(ctx).Model(&UserAccessLog{}).
|
||||
Select("user_id, COUNT(*) AS cnt").
|
||||
Where("user_id <> 0 AND created_at >= ?", startTime).
|
||||
Group("user_id").Order("cnt DESC").Limit(limitOr(limit, defaultTopN)).Scan(&rows).Error
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
out := make([]TopUser, len(rows))
|
||||
for i, r := range rows {
|
||||
out[i] = TopUser{UserID: r.UserID, Count: r.Cnt}
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
func (s *userAccessLogGormStore) EnsurePartitions(ctx context.Context, from, to time.Time) error {
|
||||
if !isPostgresDialect(s.db) {
|
||||
return nil
|
||||
}
|
||||
for _, sql := range partitionStatementsRange(from, to) {
|
||||
if err := s.db.WithContext(ctx).Exec(sql).Error; err != nil {
|
||||
return fmt.Errorf("ensure partition: %w", err)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func gormMigrationRange[T any](
|
||||
ctx context.Context,
|
||||
gdb *gorm.DB,
|
||||
column string,
|
||||
model T,
|
||||
timeOf func(*T) time.Time,
|
||||
) (time.Time, time.Time, error) {
|
||||
var first, last T
|
||||
found := false
|
||||
for _, order := range []string{"ASC", "DESC"} {
|
||||
out := &first
|
||||
if order == "DESC" {
|
||||
out = &last
|
||||
}
|
||||
res := gdb.WithContext(ctx).Model(model).Order(column + " " + order).Limit(1).Take(out)
|
||||
if res.Error != nil && !errors.Is(res.Error, gorm.ErrRecordNotFound) {
|
||||
return time.Time{}, time.Time{}, fmt.Errorf("query migration range %s: %w", column, res.Error)
|
||||
}
|
||||
if res.Error == nil {
|
||||
found = true
|
||||
}
|
||||
}
|
||||
if !found {
|
||||
return time.Time{}, time.Time{}, nil
|
||||
}
|
||||
return timeOf(&first).UTC(), timeOf(&last).UTC(), nil
|
||||
}
|
||||
|
||||
func limitOr(v, def int) int {
|
||||
if v <= 0 {
|
||||
return def
|
||||
}
|
||||
return v
|
||||
}
|
||||
|
||||
func offsetOf(page, pageSize int) int {
|
||||
if page < 1 {
|
||||
page = 1
|
||||
}
|
||||
return (page - 1) * limitOr(pageSize, defaultPageSize)
|
||||
}
|
||||
|
||||
func countToUint64(v int64) uint64 {
|
||||
if v < 0 {
|
||||
return 0
|
||||
}
|
||||
return uint64(v)
|
||||
}
|
||||
|
||||
func isPostgresDialect(db *gorm.DB) bool {
|
||||
return db != nil && db.Dialector != nil && db.Name() == "postgres"
|
||||
}
|
||||
|
||||
func dailyTrendDateSQL(db *gorm.DB) string {
|
||||
if isPostgresDialect(db) {
|
||||
return "to_char(created_at, 'YYYY-MM-DD')"
|
||||
}
|
||||
return "strftime('%Y-%m-%d', created_at)"
|
||||
}
|
||||
|
||||
func isMissingRelation(err error) bool {
|
||||
if err == nil {
|
||||
return false
|
||||
}
|
||||
msg := strings.ToLower(err.Error())
|
||||
return strings.Contains(msg, "no such table") || strings.Contains(msg, "does not exist")
|
||||
}
|
||||
@@ -0,0 +1,58 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package logstore
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/glebarez/sqlite"
|
||||
"github.com/stretchr/testify/require"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
func newTestUserAccessStore(t *testing.T) *userAccessLogGormStore {
|
||||
t.Helper()
|
||||
gdb, err := gorm.Open(sqlite.Open("file:logstore-"+t.Name()+"?mode=memory&cache=shared"), &gorm.Config{})
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, gdb.AutoMigrate(&UserAccessLog{}))
|
||||
return newUserAccessLogGormStore(gdb)
|
||||
}
|
||||
|
||||
func TestGormUserAccessLogCountList(t *testing.T) {
|
||||
ua := newTestUserAccessStore(t)
|
||||
ctx := context.Background()
|
||||
now := time.Now().UTC().Truncate(time.Second)
|
||||
require.NoError(t, ua.BatchInsert(ctx, []UserAccessLog{
|
||||
{UserID: 10, Path: "/api/v1/users", Method: "GET", Status: 200, CreatedAt: now},
|
||||
{UserID: 20, Path: "/api/v1/admin", Method: "GET", Status: 200, CreatedAt: now},
|
||||
{UserID: 10, Path: "/api/v1/other", Method: "POST", Status: 201, CreatedAt: now},
|
||||
}))
|
||||
|
||||
count, err := ua.Count(ctx, AccessLogFilter{UserIDs: []uint64{10}, Path: "users"})
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, uint64(1), count)
|
||||
|
||||
rows, total, err := ua.List(ctx, AccessLogFilter{UserIDs: []uint64{10}, Path: "users"}, 1, 10)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, uint64(1), total)
|
||||
require.Len(t, rows, 1)
|
||||
require.Equal(t, "/api/v1/users", rows[0].Path)
|
||||
require.NotZero(t, rows[0].ID)
|
||||
}
|
||||
|
||||
func TestGormUserAccessLogFreeze(t *testing.T) {
|
||||
ua := newTestUserAccessStore(t)
|
||||
SetConfigReader(func(_ context.Context, key string) (string, error) {
|
||||
if key == logMigrationKey {
|
||||
return "migrating", nil
|
||||
}
|
||||
return "", nil
|
||||
})
|
||||
t.Cleanup(ResetForTest)
|
||||
|
||||
err := ua.BatchInsert(context.Background(), []UserAccessLog{{UserID: 1, CreatedAt: time.Now()}})
|
||||
require.ErrorIs(t, err, ErrMigrating)
|
||||
}
|
||||
@@ -0,0 +1,41 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package logstore
|
||||
|
||||
import (
|
||||
"os/exec"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
// plugins/domain 不得直接 import 已废弃的 pkg/persistence。
|
||||
var forbiddenImports = []string{
|
||||
"github.com/Rain-kl/Wavelet/pkg/idgen",
|
||||
"github.com/Rain-kl/Wavelet/plugins/domain/risk_control/logstore",
|
||||
}
|
||||
|
||||
func TestAppsMustNotImportPersistenceDirectly(t *testing.T) {
|
||||
t.Chdir("../../../..")
|
||||
out, err := exec.Command("go", "list", "-test", "-f", `{{.ImportPath}} {{join .Imports " "}}`, "./plugins/domain/...").Output()
|
||||
if err != nil {
|
||||
t.Fatalf("go list: %v", err)
|
||||
}
|
||||
for _, line := range strings.Split(string(out), "\n") {
|
||||
fields := strings.Fields(line)
|
||||
if len(fields) == 0 {
|
||||
continue
|
||||
}
|
||||
pkg := fields[0]
|
||||
if !strings.HasPrefix(pkg, "github.com/Rain-kl/Wavelet/plugins/domain") {
|
||||
continue
|
||||
}
|
||||
for _, imp := range fields[1:] {
|
||||
for _, forbidden := range forbiddenImports {
|
||||
if imp == forbidden {
|
||||
t.Errorf("%s must not import forbidden package %s", pkg, forbidden)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,46 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package logstore abstracts user access-log storage across ClickHouse, PostgreSQL and SQLite.
|
||||
package logstore
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"time"
|
||||
)
|
||||
|
||||
// ErrMigrating 表示日志数据库正在迁移,当前禁止写入。
|
||||
var ErrMigrating = errors.New("log database is migrating, writes are disabled")
|
||||
|
||||
// UserAccessLogStore 用户访问日志(w_user_access_logs)。
|
||||
type UserAccessLogStore interface {
|
||||
BatchInsert(ctx context.Context, logs []UserAccessLog) error
|
||||
DeleteAll(ctx context.Context) (int64, error)
|
||||
DeleteBefore(ctx context.Context, cutoff time.Time) (int64, error)
|
||||
Count(ctx context.Context, filter AccessLogFilter) (uint64, error)
|
||||
List(ctx context.Context, filter AccessLogFilter, page, pageSize int) ([]UserAccessLog, uint64, error)
|
||||
GetDailyTrend(ctx context.Context, days int) ([]DailyTrend, error)
|
||||
GetBrowserDistribution(ctx context.Context, startTime time.Time) ([]BrowserShare, error)
|
||||
GetTopActiveUsers(ctx context.Context, startTime time.Time, limit int) ([]TopUser, error)
|
||||
ListForMigration(ctx context.Context, afterID uint64, limit int) ([]UserAccessLog, error)
|
||||
MigrationRange(ctx context.Context) (from, to time.Time, err error)
|
||||
EnsurePartitions(ctx context.Context, from, to time.Time) error
|
||||
// DropEmptyPartitions 幂等清理 PG 空分区表:删除 before 月份之前、且无任何数据的按月分区;
|
||||
// CH/SQLite 为 no-op。
|
||||
DropEmptyPartitions(ctx context.Context, before time.Time) error
|
||||
// DropExpiredPartitions 直接删除完全过期的 PG 整月分区(候选为月份早于 cutoff 月的分区,
|
||||
// 删除前校验分区内无保留期内数据,避免时区偏移下误删;迁移冻结期间拒绝执行);CH/SQLite 为 no-op。
|
||||
DropExpiredPartitions(ctx context.Context, cutoff time.Time) error
|
||||
}
|
||||
|
||||
// StatusStore 日志库状态。
|
||||
type StatusStore interface {
|
||||
ActiveDatabase(ctx context.Context) (string, error)
|
||||
}
|
||||
|
||||
// Store 当前生效日志库。
|
||||
type Store struct {
|
||||
UserAccessLogs UserAccessLogStore
|
||||
Status StatusStore
|
||||
}
|
||||
@@ -0,0 +1,44 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package logstore abstracts user access-log storage across ClickHouse, PostgreSQL and SQLite.
|
||||
package logstore
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"time"
|
||||
)
|
||||
|
||||
const (
|
||||
userAccessLogTableName = "w_user_access_logs"
|
||||
userAccessLogInsertColumns = "id, user_id, path, method, ip, user_agent, headers, status, latency, created_at"
|
||||
)
|
||||
|
||||
// UserAccessLog stores HTTP access records in ClickHouse, PostgreSQL and SQLite.
|
||||
type UserAccessLog struct {
|
||||
ID uint64 `gorm:"column:id"`
|
||||
UserID uint64 `gorm:"column:user_id"`
|
||||
Path string `gorm:"column:path"`
|
||||
Method string `gorm:"column:method"`
|
||||
IP string `gorm:"column:ip"`
|
||||
UserAgent string `gorm:"column:user_agent"`
|
||||
Headers string `gorm:"column:headers"`
|
||||
Status int32 `gorm:"column:status"`
|
||||
Latency int64 `gorm:"column:latency"`
|
||||
CreatedAt time.Time `gorm:"column:created_at"`
|
||||
}
|
||||
|
||||
// TableName returns the table name.
|
||||
func (UserAccessLog) TableName() string {
|
||||
return userAccessLogTableName
|
||||
}
|
||||
|
||||
// InsertColumns returns comma-separated column names for batch insert.
|
||||
func (UserAccessLog) InsertColumns() string {
|
||||
return userAccessLogInsertColumns
|
||||
}
|
||||
|
||||
// BatchInsertSQL returns the INSERT prefix used by native batch writers.
|
||||
func (UserAccessLog) BatchInsertSQL() string {
|
||||
return fmt.Sprintf("INSERT INTO %s (%s)", userAccessLogTableName, userAccessLogInsertColumns)
|
||||
}
|
||||
@@ -0,0 +1,115 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package logstore
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
// listPartitionNames 列出 table 在当前 schema 下的全部直接分区表名(pg_inherits)。
|
||||
func listPartitionNames(ctx context.Context, gdb *gorm.DB, table string) ([]string, error) {
|
||||
var names []string
|
||||
if err := gdb.WithContext(ctx).Raw(`
|
||||
SELECT c.relname
|
||||
FROM pg_inherits i
|
||||
JOIN pg_class c ON c.oid = i.inhrelid
|
||||
JOIN pg_class p ON p.oid = i.inhparent
|
||||
JOIN pg_namespace n ON n.oid = p.relnamespace AND n.nspname = current_schema()
|
||||
WHERE p.relname = ?`, table).Scan(&names).Error; err != nil {
|
||||
return nil, fmt.Errorf("list partitions of %s: %w", table, err)
|
||||
}
|
||||
return names, nil
|
||||
}
|
||||
|
||||
// partitionNameMonth 解析按月分区表名 <table>_YYYYMM 的所属月份;命名不匹配返回 (零值, false)。
|
||||
func partitionNameMonth(table, name string) (time.Time, bool) {
|
||||
suffix, ok := strings.CutPrefix(name, table+"_")
|
||||
if !ok || len(suffix) != 6 {
|
||||
return time.Time{}, false
|
||||
}
|
||||
m, err := time.Parse("200601", suffix)
|
||||
if err != nil {
|
||||
return time.Time{}, false
|
||||
}
|
||||
return m, true
|
||||
}
|
||||
|
||||
// dropEligiblePartitionNames 返回 before 月份之前、命名合法的分区表名(是否为空由调用方校验)。
|
||||
func dropEligiblePartitionNames(table string, names []string, before time.Time) []string {
|
||||
beforeMonth := time.Date(before.Year(), before.Month(), 1, 0, 0, 0, 0, time.UTC)
|
||||
out := make([]string, 0, len(names))
|
||||
for _, name := range names {
|
||||
month, ok := partitionNameMonth(table, name)
|
||||
if !ok || !month.Before(beforeMonth) {
|
||||
continue
|
||||
}
|
||||
out = append(out, name)
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// DropEmptyPartitions 幂等清理 PG 空分区表:仅删除 before 月份之前、且无任何数据的分区。
|
||||
// 非 PG 方言为 no-op。
|
||||
func (s *gormLogStore) DropEmptyPartitions(ctx context.Context, before time.Time) error {
|
||||
if !isPostgresDialect(s.db) {
|
||||
return nil
|
||||
}
|
||||
names, err := listPartitionNames(ctx, s.db, userAccessLogTable)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
for _, name := range dropEligiblePartitionNames(userAccessLogTable, names, before) {
|
||||
var one int
|
||||
if err := s.db.WithContext(ctx).Raw("SELECT 1 FROM " + name + " LIMIT 1").Scan(&one).Error; err != nil {
|
||||
return fmt.Errorf("check partition %s empty: %w", name, err)
|
||||
}
|
||||
if one == 1 {
|
||||
continue
|
||||
}
|
||||
if err := s.db.WithContext(ctx).Exec("DROP TABLE IF EXISTS " + name).Error; err != nil {
|
||||
return fmt.Errorf("drop empty partition %s: %w", name, err)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// DropExpiredPartitions 直接删除完全过期的 PG 整月分区(避免 retention 清理逐行 DELETE)。
|
||||
// 候选 = 月份早于 cutoff 月(按 cutoff 的 UTC 时刻取月);删除前校验分区内不存在 created_at >= cutoff 的行。
|
||||
// 迁移冻结期间返回 ErrMigrating。CH/SQLite 为 no-op。
|
||||
func (s *gormLogStore) DropExpiredPartitions(ctx context.Context, cutoff time.Time) error {
|
||||
if !isPostgresDialect(s.db) {
|
||||
return nil
|
||||
}
|
||||
if err := s.ensureWritable(ctx); err != nil {
|
||||
return err
|
||||
}
|
||||
names, err := listPartitionNames(ctx, s.db, userAccessLogTable)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
cu := cutoff.UTC()
|
||||
cutoffMonth := time.Date(cu.Year(), cu.Month(), 1, 0, 0, 0, 0, time.UTC)
|
||||
for _, name := range names {
|
||||
month, ok := partitionNameMonth(userAccessLogTable, name)
|
||||
if !ok || !month.Before(cutoffMonth) {
|
||||
continue
|
||||
}
|
||||
var hasRetained int
|
||||
if err := s.db.WithContext(ctx).Raw("SELECT 1 FROM "+name+" WHERE created_at >= ? LIMIT 1", cu).Scan(&hasRetained).Error; err != nil {
|
||||
return fmt.Errorf("check partition %s retained rows: %w", name, err)
|
||||
}
|
||||
if hasRetained == 1 {
|
||||
continue
|
||||
}
|
||||
if err := s.db.WithContext(ctx).Exec("DROP TABLE IF EXISTS " + name).Error; err != nil {
|
||||
return fmt.Errorf("drop expired partition %s: %w", name, err)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,69 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package logstore
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestPartitionNameMonth(t *testing.T) {
|
||||
cases := []struct {
|
||||
table string
|
||||
name string
|
||||
want string
|
||||
}{
|
||||
{"w_user_access_logs", "w_user_access_logs_202612", "2026-12"},
|
||||
{"w_user_access_logs", "w_user_access_logs_202608", "2026-08"},
|
||||
{"w_user_access_logs", "of_node_access_logs_202608", ""},
|
||||
{"w_user_access_logs", "w_user_access_logs_20268", ""},
|
||||
{"w_user_access_logs", "w_user_access_logs_202613", ""},
|
||||
{"w_user_access_logs", "w_user_access_logs_default", ""},
|
||||
}
|
||||
for _, c := range cases {
|
||||
got, ok := partitionNameMonth(c.table, c.name)
|
||||
if c.want == "" {
|
||||
if ok {
|
||||
t.Fatalf("partitionNameMonth(%q, %q) ok = true, want false", c.table, c.name)
|
||||
}
|
||||
continue
|
||||
}
|
||||
if !ok || got.Format("2006-01") != c.want {
|
||||
t.Fatalf("partitionNameMonth(%q, %q) = %v, want %s", c.table, c.name, got, c.want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestDropEligiblePartitionNames(t *testing.T) {
|
||||
before := time.Date(2026, 10, 15, 0, 0, 0, 0, time.UTC)
|
||||
names := []string{
|
||||
"w_user_access_logs_202608",
|
||||
"w_user_access_logs_202609",
|
||||
"w_user_access_logs_202610",
|
||||
"w_user_access_logs_202611",
|
||||
"w_user_access_logs_default",
|
||||
}
|
||||
got := dropEligiblePartitionNames(userAccessLogTable, names, before)
|
||||
want := []string{"w_user_access_logs_202608", "w_user_access_logs_202609"}
|
||||
require.Equal(t, want, got)
|
||||
|
||||
first := time.Date(2026, 10, 1, 0, 0, 0, 0, time.UTC)
|
||||
require.Empty(t, dropEligiblePartitionNames(userAccessLogTable, []string{"w_user_access_logs_202610"}, first))
|
||||
}
|
||||
|
||||
func TestDropPartitionHelpersSQLiteNoop(t *testing.T) {
|
||||
ua := newTestUserAccessStore(t)
|
||||
ctx := context.Background()
|
||||
require.NoError(t, ua.BatchInsert(ctx, []UserAccessLog{
|
||||
{UserID: 1, Path: "/x", CreatedAt: time.Now().UTC()},
|
||||
}))
|
||||
require.NoError(t, ua.DropExpiredPartitions(ctx, time.Now().AddDate(0, 0, -90)))
|
||||
require.NoError(t, ua.DropEmptyPartitions(ctx, time.Now()))
|
||||
count, err := ua.Count(ctx, AccessLogFilter{})
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, uint64(1), count)
|
||||
}
|
||||
@@ -0,0 +1,188 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package logstore
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/pkg/config"
|
||||
"github.com/Rain-kl/Wavelet/pkg/logger"
|
||||
db "github.com/Rain-kl/Wavelet/plugins/infra/database"
|
||||
)
|
||||
|
||||
const (
|
||||
logDatabaseKey = "log_database"
|
||||
logMigrationKey = "log_db_migration"
|
||||
)
|
||||
|
||||
const (
|
||||
dbNamePostgres = "postgres"
|
||||
dbNameSQLite = "sqlite"
|
||||
dbNameClickHouse = "clickhouse"
|
||||
)
|
||||
|
||||
var errConfigReaderNotWired = errors.New("logstore: config reader not wired")
|
||||
|
||||
// ConfigReader 读取系统配置字符串值,由 bootstrap 注入(避免 logstore ↔ repository 循环依赖)。
|
||||
type ConfigReader func(ctx context.Context, key string) (string, error)
|
||||
|
||||
const resolveCacheTTL = 1 * time.Second
|
||||
|
||||
var (
|
||||
configReader ConfigReader
|
||||
|
||||
storeMu sync.RWMutex
|
||||
active *Store
|
||||
activeDB string
|
||||
lastResolveDB string
|
||||
lastResolveTime time.Time
|
||||
)
|
||||
|
||||
// SetConfigReader 注入系统配置读取函数(bootstrap 调用,测试可注入内存实现)。
|
||||
func SetConfigReader(fn ConfigReader) { configReader = fn }
|
||||
|
||||
func getConfig(ctx context.Context, key string) (string, error) {
|
||||
if configReader == nil {
|
||||
return "", errConfigReaderNotWired
|
||||
}
|
||||
return configReader(ctx, key)
|
||||
}
|
||||
|
||||
// Active 返回当前生效的日志库 Store。
|
||||
func Active(ctx context.Context) (*Store, error) {
|
||||
current, err := resolveDatabase(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
storeMu.RLock()
|
||||
if active != nil && activeDB == current {
|
||||
s := active
|
||||
storeMu.RUnlock()
|
||||
return s, nil
|
||||
}
|
||||
storeMu.RUnlock()
|
||||
|
||||
storeMu.Lock()
|
||||
defer storeMu.Unlock()
|
||||
if active != nil && activeDB == current {
|
||||
return active, nil
|
||||
}
|
||||
s, err := buildStore(ctx, current, false)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
active = s
|
||||
activeDB = current
|
||||
return s, nil
|
||||
}
|
||||
|
||||
// Build 直接按目标构造 store(不经 Active 缓存)。
|
||||
func Build(ctx context.Context, database string) (*Store, error) {
|
||||
return buildStore(ctx, database, false)
|
||||
}
|
||||
|
||||
// BuildForMigration 构造迁移目标 store,跳过冻结检查。
|
||||
func BuildForMigration(ctx context.Context, database string) (*Store, error) {
|
||||
return buildStore(ctx, database, true)
|
||||
}
|
||||
|
||||
func buildStore(ctx context.Context, database string, skipFreeze bool) (*Store, error) {
|
||||
switch database {
|
||||
case dbNameClickHouse:
|
||||
ual := newClickHouseUserAccessLogStore()
|
||||
ual.skipFreeze = skipFreeze
|
||||
return &Store{UserAccessLogs: ual, Status: ual}, nil
|
||||
case dbNamePostgres, dbNameSQLite:
|
||||
gdb := db.DB(ctx)
|
||||
ual := newUserAccessLogGormStore(gdb)
|
||||
ual.skipFreeze = skipFreeze
|
||||
return &Store{UserAccessLogs: ual, Status: ual}, nil
|
||||
default:
|
||||
return nil, fmt.Errorf("unsupported log database: %s", database)
|
||||
}
|
||||
}
|
||||
|
||||
// Migrating 返回日志库是否处于迁移冻结状态。
|
||||
func Migrating(ctx context.Context) bool {
|
||||
v, err := getConfig(ctx, logMigrationKey)
|
||||
if err != nil {
|
||||
if !errors.Is(err, errConfigReaderNotWired) {
|
||||
logger.ErrorF(ctx, "read log migration config failed: %v", err)
|
||||
}
|
||||
return false
|
||||
}
|
||||
return v == "migrating"
|
||||
}
|
||||
|
||||
// Init 预热激活 store,并兜底预建当前月及未来分区。
|
||||
func Init(ctx context.Context) {
|
||||
s, err := Active(ctx)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
now := time.Now().UTC()
|
||||
if err := s.UserAccessLogs.EnsurePartitions(ctx, now, now.AddDate(0, partitionLeadMonths, 0)); err != nil {
|
||||
logger.WarnF(ctx, "logstore: ensure startup partitions failed: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// InvalidateCache 清空日志库解析缓存。
|
||||
func InvalidateCache() {
|
||||
storeMu.Lock()
|
||||
defer storeMu.Unlock()
|
||||
lastResolveTime = time.Time{}
|
||||
lastResolveDB = ""
|
||||
}
|
||||
|
||||
// ResetForTest 清空缓存的激活 store 与 config reader。
|
||||
func ResetForTest() {
|
||||
storeMu.Lock()
|
||||
active = nil
|
||||
activeDB = ""
|
||||
lastResolveDB = ""
|
||||
lastResolveTime = time.Time{}
|
||||
storeMu.Unlock()
|
||||
configReader = nil
|
||||
}
|
||||
|
||||
// ActiveDatabase 返回当前日志主库名。
|
||||
func ActiveDatabase(ctx context.Context) (string, error) {
|
||||
return resolveDatabase(ctx)
|
||||
}
|
||||
|
||||
func resolveDatabase(ctx context.Context) (string, error) {
|
||||
storeMu.RLock()
|
||||
if active != nil && time.Since(lastResolveTime) < resolveCacheTTL {
|
||||
name := lastResolveDB
|
||||
storeMu.RUnlock()
|
||||
return name, nil
|
||||
}
|
||||
storeMu.RUnlock()
|
||||
|
||||
v, err := getConfig(ctx, logDatabaseKey)
|
||||
if err != nil && !errors.Is(err, errConfigReaderNotWired) {
|
||||
return "", err
|
||||
}
|
||||
|
||||
resolved := v
|
||||
if resolved == "" {
|
||||
resolved = dbNameSQLite
|
||||
if config.Config.Database.Enabled {
|
||||
resolved = dbNamePostgres
|
||||
}
|
||||
if config.Config.ClickHouse.Enabled {
|
||||
resolved = dbNameClickHouse
|
||||
}
|
||||
}
|
||||
|
||||
storeMu.Lock()
|
||||
lastResolveDB = resolved
|
||||
lastResolveTime = time.Now()
|
||||
storeMu.Unlock()
|
||||
return resolved, nil
|
||||
}
|
||||
@@ -11,10 +11,10 @@ import (
|
||||
|
||||
"github.com/Rain-kl/Wavelet/core/contracts"
|
||||
"github.com/Rain-kl/Wavelet/pkg/config"
|
||||
"github.com/Rain-kl/Wavelet/pkg/persistence/idgen"
|
||||
"github.com/Rain-kl/Wavelet/pkg/persistence/logstore"
|
||||
"github.com/Rain-kl/Wavelet/pkg/idgen"
|
||||
"github.com/Rain-kl/Wavelet/pkg/response"
|
||||
"github.com/Rain-kl/Wavelet/plugins/domain/auth"
|
||||
"github.com/Rain-kl/Wavelet/plugins/domain/risk_control/logstore"
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
|
||||
@@ -13,12 +13,12 @@ import (
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/core/contracts"
|
||||
"github.com/Rain-kl/Wavelet/pkg/batchwriter"
|
||||
"github.com/Rain-kl/Wavelet/pkg/config"
|
||||
"github.com/Rain-kl/Wavelet/pkg/persistence/batchwriter"
|
||||
"github.com/Rain-kl/Wavelet/pkg/persistence/logstore"
|
||||
"github.com/Rain-kl/Wavelet/pkg/testhelper"
|
||||
"github.com/Rain-kl/Wavelet/plugins/domain/auth"
|
||||
"github.com/Rain-kl/Wavelet/plugins/domain/risk_control"
|
||||
"github.com/Rain-kl/Wavelet/plugins/domain/risk_control/logstore"
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/stretchr/testify/assert"
|
||||
)
|
||||
|
||||
@@ -9,7 +9,7 @@ import (
|
||||
|
||||
"github.com/Rain-kl/Wavelet/core"
|
||||
"github.com/Rain-kl/Wavelet/pkg/config"
|
||||
db "github.com/Rain-kl/Wavelet/pkg/persistence"
|
||||
database "github.com/Rain-kl/Wavelet/plugins/infra/database"
|
||||
"github.com/Rain-kl/Wavelet/pkg/response"
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
@@ -54,7 +54,7 @@ func (p *Plugin) Apply(ctx *core.Context) error {
|
||||
Value string `json:"value"`
|
||||
}
|
||||
var configs []configItem
|
||||
_ = db.DB(c.Request.Context()).Table("w_system_configs").Where("visibility = ?", "visible").Find(&configs).Error
|
||||
_ = database.DB(c.Request.Context()).Table("w_system_configs").Where("visibility = ?", "visible").Find(&configs).Error
|
||||
c.JSON(http.StatusOK, response.OK(gin.H{
|
||||
"configs": configs,
|
||||
"app": gin.H{
|
||||
|
||||
+6
-5
@@ -11,7 +11,8 @@ import (
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
db "github.com/Rain-kl/Wavelet/pkg/persistence"
|
||||
cachepkg "github.com/Rain-kl/Wavelet/plugins/infra/cache"
|
||||
database "github.com/Rain-kl/Wavelet/plugins/infra/database"
|
||||
"github.com/Rain-kl/Wavelet/pkg/util"
|
||||
"github.com/Rain-kl/Wavelet/plugins/domain/upload/shared"
|
||||
uploadstorage "github.com/Rain-kl/Wavelet/plugins/domain/upload/storage"
|
||||
@@ -41,8 +42,8 @@ func ResetAccessCaches() {
|
||||
|
||||
// PublishAccessCacheInvalidation broadcasts upload access cache eviction to all nodes.
|
||||
func PublishAccessCacheInvalidation(ctx context.Context) {
|
||||
if db.Redis != nil {
|
||||
_ = db.Redis.Publish(ctx, fileAccessInvalidationChannel, "reset").Err()
|
||||
if cachepkg.Redis != nil {
|
||||
_ = cachepkg.Redis.Publish(ctx, fileAccessInvalidationChannel, "reset").Err()
|
||||
}
|
||||
}
|
||||
|
||||
@@ -51,7 +52,7 @@ func ensureAccessCacheListener() {
|
||||
}
|
||||
|
||||
func startAccessCacheInvalidationListener() {
|
||||
rdb := db.Redis
|
||||
rdb := cachepkg.Redis
|
||||
if rdb == nil {
|
||||
return
|
||||
}
|
||||
@@ -114,7 +115,7 @@ func fetchFileAccessWhitelist(ctx context.Context) map[string]struct{} {
|
||||
|
||||
func parseFileAccessWhitelist(ctx context.Context) []string {
|
||||
var sc struct{ Value string }
|
||||
err := db.DB(ctx).Table("w_system_configs").Where("key = ?", "file_access_whitelist").First(&sc).Error
|
||||
err := database.DB(ctx).Table("w_system_configs").Where("key = ?", "file_access_whitelist").First(&sc).Error
|
||||
if err != nil || sc.Value == "" {
|
||||
return []string{shared.DefaultPublicUploadType}
|
||||
}
|
||||
|
||||
+14
-13
@@ -10,7 +10,8 @@ import (
|
||||
"sync"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/pkg/cache/ram"
|
||||
db "github.com/Rain-kl/Wavelet/pkg/persistence"
|
||||
database "github.com/Rain-kl/Wavelet/plugins/infra/database"
|
||||
cachepkg "github.com/Rain-kl/Wavelet/plugins/infra/cache"
|
||||
"github.com/Rain-kl/Wavelet/pkg/util"
|
||||
"github.com/Rain-kl/Wavelet/plugins/domain/upload/models"
|
||||
)
|
||||
@@ -42,7 +43,7 @@ func cloneUpload(u models.Upload) models.Upload {
|
||||
}
|
||||
|
||||
func ensureUploadMetaCacheListener() {
|
||||
if db.Redis == nil {
|
||||
if cachepkg.Redis == nil {
|
||||
return
|
||||
}
|
||||
uploadMetaListenerOnce.Do(startUploadMetaCacheInvalidationListener)
|
||||
@@ -52,7 +53,7 @@ func startUploadMetaCacheInvalidationListener() {
|
||||
uploadMetaListenerCtx, uploadMetaListenerCancel = context.WithCancel(context.Background())
|
||||
uploadMetaListenerDone = make(chan struct{})
|
||||
|
||||
redisClient := db.Redis // 捕获当前客户端:goroutine 不读可变全局,避免与测试置空 db.Redis 竞争
|
||||
redisClient := cachepkg.Redis // 捕获当前客户端:goroutine 不读可变全局,避免与测试置空 cachepkg.Redis 竞争
|
||||
util.Go(func() {
|
||||
defer close(uploadMetaListenerDone)
|
||||
pubsub := redisClient.Subscribe(uploadMetaListenerCtx, uploadMetaInvalidationChan)
|
||||
@@ -77,14 +78,14 @@ func startUploadMetaCacheInvalidationListener() {
|
||||
}
|
||||
|
||||
func publishUploadMetaRAMInvalidation(ctx context.Context, id uint64) {
|
||||
if db.Redis == nil {
|
||||
if cachepkg.Redis == nil {
|
||||
return
|
||||
}
|
||||
payload, err := json.Marshal(uploadMetaInvalidationMessage{ID: id})
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
_ = db.Redis.Publish(ctx, uploadMetaInvalidationChan, payload).Err()
|
||||
_ = cachepkg.Redis.Publish(ctx, uploadMetaInvalidationChan, payload).Err()
|
||||
}
|
||||
|
||||
// GetUploadByID loads upload metadata from RAM, Redis, or the database.
|
||||
@@ -96,16 +97,16 @@ func GetUploadByID(ctx context.Context, id uint64) (models.Upload, error) {
|
||||
}
|
||||
|
||||
key := uploadMetaRedisKey(id)
|
||||
if db.Redis != nil {
|
||||
if cachepkg.Redis != nil {
|
||||
var u models.Upload
|
||||
if err := db.GetJSON(ctx, key, &u); err == nil {
|
||||
if err := cachepkg.GetJSON(ctx, key, &u); err == nil {
|
||||
uploadMetaRAM.Set(id, cloneUpload(u))
|
||||
return u, nil
|
||||
}
|
||||
}
|
||||
|
||||
var u models.Upload
|
||||
if err := db.DB(ctx).
|
||||
if err := database.DB(ctx).
|
||||
Where("id = ? AND status IN (?, ?)", id, models.UploadStatusPending, models.UploadStatusUsed).
|
||||
First(&u).Error; err != nil {
|
||||
return models.Upload{}, err
|
||||
@@ -125,8 +126,8 @@ func SetUploadMetaCache(ctx context.Context, u *models.Upload) {
|
||||
|
||||
cloned := cloneUpload(*u)
|
||||
uploadMetaRAM.Set(u.ID, cloned)
|
||||
if db.Redis != nil {
|
||||
_ = db.SetJSON(ctx, uploadMetaRedisKey(u.ID), cloned, uploadMetaRedisCacheTTL)
|
||||
if cachepkg.Redis != nil {
|
||||
_ = cachepkg.SetJSON(ctx, uploadMetaRedisKey(u.ID), cloned, uploadMetaRedisCacheTTL)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -135,8 +136,8 @@ func InvalidateUploadMetaCache(ctx context.Context, id uint64) {
|
||||
ensureUploadMetaCacheListener()
|
||||
|
||||
uploadMetaRAM.Invalidate(id)
|
||||
if db.Redis != nil {
|
||||
_ = db.Redis.Del(ctx, db.PrefixedKey(uploadMetaRedisKey(id))).Err()
|
||||
if cachepkg.Redis != nil {
|
||||
_ = cachepkg.Redis.Del(ctx, cachepkg.PrefixedKey(uploadMetaRedisKey(id))).Err()
|
||||
publishUploadMetaRAMInvalidation(ctx, id)
|
||||
}
|
||||
}
|
||||
@@ -151,7 +152,7 @@ func StopUploadMetaCacheListener() {
|
||||
if uploadMetaListenerCancel != nil {
|
||||
uploadMetaListenerCancel()
|
||||
if uploadMetaListenerDone != nil {
|
||||
<-uploadMetaListenerDone // 等待 goroutine 退出,保证之后置空 db.Redis 不再竞争
|
||||
<-uploadMetaListenerDone // 等待 goroutine 退出,保证之后置空 cachepkg.Redis 不再竞争
|
||||
}
|
||||
uploadMetaListenerCancel = nil
|
||||
uploadMetaListenerDone = nil
|
||||
|
||||
+8
-8
@@ -9,7 +9,7 @@ import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/pkg/persistence"
|
||||
cachepkg "github.com/Rain-kl/Wavelet/plugins/infra/cache"
|
||||
"github.com/Rain-kl/Wavelet/pkg/testhelper"
|
||||
"github.com/Rain-kl/Wavelet/plugins/domain/upload/models"
|
||||
"gorm.io/gorm"
|
||||
@@ -58,7 +58,7 @@ func TestGetUploadByIDLoadsFromDBAndPopulatesCache(t *testing.T) {
|
||||
}
|
||||
|
||||
var redisUpload models.Upload
|
||||
if err := db.GetJSON(ctx, uploadMetaRedisKey(upload.ID), &redisUpload); err != nil {
|
||||
if err := cachepkg.GetJSON(ctx, uploadMetaRedisKey(upload.ID), &redisUpload); err != nil {
|
||||
t.Fatalf("redis cache miss after DB load: %v", err)
|
||||
}
|
||||
if redisUpload.ID != upload.ID {
|
||||
@@ -137,7 +137,7 @@ func TestInvalidateUploadMetaCacheClearsRAMAndRedis(t *testing.T) {
|
||||
InvalidateUploadMetaCache(ctx, upload.ID)
|
||||
|
||||
var redisUpload models.Upload
|
||||
if err := db.GetJSON(ctx, uploadMetaRedisKey(upload.ID), &redisUpload); err == nil {
|
||||
if err := cachepkg.GetJSON(ctx, uploadMetaRedisKey(upload.ID), &redisUpload); err == nil {
|
||||
t.Fatal("expected redis cache to be invalidated")
|
||||
}
|
||||
|
||||
@@ -188,7 +188,7 @@ func TestUploadMetaInvalidationPubSubClearsPeerRAM(t *testing.T) {
|
||||
if err != nil {
|
||||
t.Fatalf("marshal invalidation payload: %v", err)
|
||||
}
|
||||
if err := db.Redis.Publish(ctx, uploadMetaInvalidationChan, string(payload)).Err(); err != nil {
|
||||
if err := cachepkg.Redis.Publish(ctx, uploadMetaInvalidationChan, string(payload)).Err(); err != nil {
|
||||
t.Fatalf("publish invalidation: %v", err)
|
||||
}
|
||||
|
||||
@@ -205,7 +205,7 @@ func TestUploadMetaInvalidationPubSubClearsPeerRAM(t *testing.T) {
|
||||
t.Fatal("expected peer RAM cache to be cleared by pub/sub")
|
||||
}
|
||||
|
||||
if err := db.Redis.Del(ctx, db.PrefixedKey(uploadMetaRedisKey(upload.ID))).Err(); err != nil {
|
||||
if err := cachepkg.Redis.Del(ctx, cachepkg.PrefixedKey(uploadMetaRedisKey(upload.ID))).Err(); err != nil {
|
||||
t.Fatalf("delete redis cache: %v", err)
|
||||
}
|
||||
if _, err := GetUploadByID(ctx, upload.ID); err == nil {
|
||||
@@ -243,10 +243,10 @@ func TestGetUploadByIDWorksWithRedisDisabled(t *testing.T) {
|
||||
defer cleanup()
|
||||
ResetUploadMetaCacheForTest()
|
||||
|
||||
redisClient := db.Redis
|
||||
db.Redis = nil
|
||||
redisClient := cachepkg.Redis
|
||||
cachepkg.Redis = nil
|
||||
t.Cleanup(func() {
|
||||
db.Redis = redisClient
|
||||
cachepkg.Redis = redisClient
|
||||
StopUploadMetaCacheListener()
|
||||
})
|
||||
|
||||
|
||||
@@ -35,4 +35,5 @@ const (
|
||||
ErrS3KeyStartsWithSlash = shared.ErrS3KeyStartsWithSlash
|
||||
ErrS3KeyContainsNullBytes = shared.ErrS3KeyContainsNullBytes
|
||||
ErrQueryUnusedUploadsFailed = shared.ErrQueryUnusedUploadsFailed
|
||||
ErrUnauthorized = shared.ErrUnauthorized
|
||||
)
|
||||
|
||||
@@ -4,7 +4,6 @@
|
||||
package upload
|
||||
|
||||
import (
|
||||
"github.com/Rain-kl/Wavelet/pkg/task"
|
||||
"github.com/Rain-kl/Wavelet/plugins/domain/upload/cache"
|
||||
"github.com/Rain-kl/Wavelet/plugins/domain/upload/filesrv"
|
||||
"github.com/Rain-kl/Wavelet/plugins/domain/upload/handler"
|
||||
@@ -12,6 +11,7 @@ import (
|
||||
uploadstats "github.com/Rain-kl/Wavelet/plugins/domain/upload/stats"
|
||||
uploadtask "github.com/Rain-kl/Wavelet/plugins/domain/upload/task"
|
||||
"github.com/Rain-kl/Wavelet/plugins/domain/upload/util"
|
||||
"github.com/Rain-kl/Wavelet/plugins/drivers/driver_asynq_worker"
|
||||
)
|
||||
|
||||
// HTTP handlers
|
||||
@@ -116,11 +116,11 @@ type WarmImageCachePayload = uploadtask.WarmImageCachePayload
|
||||
|
||||
// Ensure task handler types implement required interfaces.
|
||||
var (
|
||||
_ task.TaskHandler = (*MigrationHandler)(nil)
|
||||
_ task.TaskHandler = (*SystemCleanupHandler)(nil)
|
||||
_ task.TaskHandler = (*RebuildUploadStatsHandler)(nil)
|
||||
_ driver_asynq_worker.TaskHandler = (*MigrationHandler)(nil)
|
||||
_ driver_asynq_worker.TaskHandler = (*SystemCleanupHandler)(nil)
|
||||
_ driver_asynq_worker.TaskHandler = (*RebuildUploadStatsHandler)(nil)
|
||||
_ interface {
|
||||
task.TaskHandler
|
||||
driver_asynq_worker.TaskHandler
|
||||
ValidatePayload([]byte) ([]byte, error)
|
||||
} = (*WarmImageCacheHandler)(nil)
|
||||
)
|
||||
|
||||
@@ -17,7 +17,6 @@ import (
|
||||
|
||||
"github.com/Rain-kl/Wavelet/core/contracts"
|
||||
"github.com/Rain-kl/Wavelet/pkg/response"
|
||||
appshared "github.com/Rain-kl/Wavelet/pkg/shared"
|
||||
"github.com/Rain-kl/Wavelet/plugins/domain/auth"
|
||||
"github.com/Rain-kl/Wavelet/plugins/domain/upload/cache"
|
||||
"github.com/Rain-kl/Wavelet/plugins/domain/upload/models"
|
||||
@@ -78,7 +77,7 @@ func ServeFileByID(c *gin.Context) {
|
||||
}
|
||||
|
||||
if err := CheckFileAccessPermission(c, upload); err != nil {
|
||||
response.AbortUnauthorized(c, appshared.UnAuthorized)
|
||||
response.AbortUnauthorized(c, shared.ErrUnauthorized)
|
||||
return
|
||||
}
|
||||
|
||||
|
||||
@@ -25,7 +25,6 @@ import (
|
||||
"github.com/Rain-kl/Wavelet/core/contracts"
|
||||
"github.com/Rain-kl/Wavelet/pkg/logger"
|
||||
"github.com/Rain-kl/Wavelet/pkg/response"
|
||||
appshared "github.com/Rain-kl/Wavelet/pkg/shared"
|
||||
"github.com/Rain-kl/Wavelet/plugins/domain/auth"
|
||||
"github.com/Rain-kl/Wavelet/plugins/domain/upload/filesrv"
|
||||
"github.com/Rain-kl/Wavelet/plugins/domain/upload/ingest"
|
||||
@@ -183,7 +182,7 @@ func DownloadFile(c *gin.Context) {
|
||||
}
|
||||
|
||||
if err := filesrv.CheckFileAccessPermission(c, upload); err != nil {
|
||||
response.AbortUnauthorized(c, appshared.UnAuthorized)
|
||||
response.AbortUnauthorized(c, shared.ErrUnauthorized)
|
||||
return
|
||||
}
|
||||
|
||||
|
||||
@@ -13,8 +13,8 @@ import (
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/pkg/logger"
|
||||
"github.com/Rain-kl/Wavelet/pkg/persistence"
|
||||
"github.com/Rain-kl/Wavelet/pkg/persistence/idgen"
|
||||
database "github.com/Rain-kl/Wavelet/plugins/infra/database"
|
||||
"github.com/Rain-kl/Wavelet/pkg/idgen"
|
||||
uploadcache "github.com/Rain-kl/Wavelet/plugins/domain/upload/cache"
|
||||
"github.com/Rain-kl/Wavelet/plugins/domain/upload/models"
|
||||
"github.com/Rain-kl/Wavelet/plugins/domain/upload/repository"
|
||||
@@ -50,7 +50,7 @@ func resolveAccessMode(uploadType string, explicit *int) int {
|
||||
|
||||
func validateAllowedExtension(ctx context.Context, ext string) error {
|
||||
var val string
|
||||
err := db.DB(ctx).Table("w_system_configs").Where("key = ?", "upload_allowed_extensions").Pluck("value", &val).Error
|
||||
err := database.DB(ctx).Table("w_system_configs").Where("key = ?", "upload_allowed_extensions").Pluck("value", &val).Error
|
||||
if err != nil {
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return nil
|
||||
@@ -128,7 +128,7 @@ func persistUploadRecord(ctx context.Context, upload *models.Upload, objectKey s
|
||||
}
|
||||
|
||||
func createUploadWithStats(ctx context.Context, upload *models.Upload) error {
|
||||
return db.DB(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
return database.DB(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
if err := repository.CreateUploadTx(tx, upload); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
@@ -13,7 +13,7 @@ import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/pkg/persistence"
|
||||
database "github.com/Rain-kl/Wavelet/plugins/infra/database"
|
||||
"github.com/Rain-kl/Wavelet/pkg/testhelper"
|
||||
"github.com/Rain-kl/Wavelet/plugins/domain/upload/models"
|
||||
"github.com/Rain-kl/Wavelet/plugins/infra/storage/objectstore"
|
||||
@@ -276,7 +276,7 @@ type totalStatsSnapshot struct {
|
||||
|
||||
func loadTotalStats(ctx context.Context) (totalStatsSnapshot, error) {
|
||||
var rows []models.UploadStat
|
||||
if err := db.DB(ctx).Where("dimension = ?", models.UploadStatDimensionTotal).Find(&rows).Error; err != nil {
|
||||
if err := database.DB(ctx).Where("dimension = ?", models.UploadStatDimensionTotal).Find(&rows).Error; err != nil {
|
||||
return totalStatsSnapshot{}, err
|
||||
}
|
||||
if len(rows) == 0 {
|
||||
|
||||
@@ -6,7 +6,7 @@ package ingest
|
||||
import (
|
||||
"context"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/pkg/persistence"
|
||||
database "github.com/Rain-kl/Wavelet/plugins/infra/database"
|
||||
uploadcache "github.com/Rain-kl/Wavelet/plugins/domain/upload/cache"
|
||||
"github.com/Rain-kl/Wavelet/plugins/domain/upload/models"
|
||||
"github.com/Rain-kl/Wavelet/plugins/domain/upload/repository"
|
||||
@@ -45,7 +45,7 @@ func RemoveOwned(ctx context.Context, userID, uploadID uint64) (models.Upload, e
|
||||
|
||||
func softDeleteUploadWithStats(ctx context.Context, upload *models.Upload) error {
|
||||
statsSnapshot := *upload
|
||||
if err := db.DB(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
if err := database.DB(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
if err := repository.SoftDeleteUploadTx(tx, upload); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
@@ -9,8 +9,8 @@ import (
|
||||
|
||||
"gorm.io/gorm"
|
||||
|
||||
db "github.com/Rain-kl/Wavelet/pkg/persistence"
|
||||
"github.com/Rain-kl/Wavelet/pkg/persistence/idgen"
|
||||
database "github.com/Rain-kl/Wavelet/plugins/infra/database"
|
||||
"github.com/Rain-kl/Wavelet/pkg/idgen"
|
||||
"github.com/Rain-kl/Wavelet/pkg/util"
|
||||
)
|
||||
|
||||
@@ -28,7 +28,7 @@ type UploadListFilter struct {
|
||||
|
||||
// ListUploads returns paginated upload records matching the filter.
|
||||
func ListUploads(ctx context.Context, filter UploadListFilter) (int64, []Upload, error) {
|
||||
query := db.DB(ctx).Model(&Upload{}).
|
||||
query := database.DB(ctx).Model(&Upload{}).
|
||||
Where("status != ?", UploadStatusDeleted)
|
||||
|
||||
if filter.UserID != 0 {
|
||||
@@ -60,7 +60,7 @@ func ListUploads(ctx context.Context, filter UploadListFilter) (int64, []Upload,
|
||||
// GetActiveUploadByID loads a non-deleted upload by ID.
|
||||
func GetActiveUploadByID(ctx context.Context, id uint64) (Upload, error) {
|
||||
var upload Upload
|
||||
if err := db.DB(ctx).Where("id = ? AND status != ?", id, UploadStatusDeleted).First(&upload).Error; err != nil {
|
||||
if err := database.DB(ctx).Where("id = ? AND status != ?", id, UploadStatusDeleted).First(&upload).Error; err != nil {
|
||||
return Upload{}, err
|
||||
}
|
||||
return upload, nil
|
||||
@@ -69,7 +69,7 @@ func GetActiveUploadByID(ctx context.Context, id uint64) (Upload, error) {
|
||||
// SoftDeleteUpload marks an upload as deleted.
|
||||
// External modules must use upload.Remove or upload.RemoveOwned; only internal/apps/upload may call this.
|
||||
func SoftDeleteUpload(ctx context.Context, upload *Upload) error {
|
||||
return SoftDeleteUploadTx(db.DB(ctx), upload)
|
||||
return SoftDeleteUploadTx(database.DB(ctx), upload)
|
||||
}
|
||||
|
||||
// SoftDeleteUploadTx marks an upload as deleted within an existing transaction.
|
||||
@@ -82,13 +82,13 @@ func UpdateUpload(ctx context.Context, upload *Upload, updates map[string]any) e
|
||||
if len(updates) == 0 {
|
||||
return nil
|
||||
}
|
||||
return db.DB(ctx).Model(upload).Updates(updates).Error
|
||||
return database.DB(ctx).Model(upload).Updates(updates).Error
|
||||
}
|
||||
|
||||
// ListDistinctUploadTypes returns all distinct non-empty upload business types.
|
||||
func ListDistinctUploadTypes(ctx context.Context) ([]string, error) {
|
||||
var types []string
|
||||
if err := db.DB(ctx).Model(&Upload{}).
|
||||
if err := database.DB(ctx).Model(&Upload{}).
|
||||
Where("type IS NOT NULL AND type != ''").
|
||||
Distinct().
|
||||
Pluck("type", &types).Error; err != nil {
|
||||
@@ -100,7 +100,7 @@ func ListDistinctUploadTypes(ctx context.Context) ([]string, error) {
|
||||
// FindReusableUploadByHash finds an existing upload with the same hash and size.
|
||||
func FindReusableUploadByHash(ctx context.Context, hash string, size int64) (Upload, error) {
|
||||
var existing Upload
|
||||
err := db.DB(ctx).
|
||||
err := database.DB(ctx).
|
||||
Where("hash = ? AND file_size = ? AND status IN (?, ?)", hash, size, UploadStatusPending, UploadStatusUsed).
|
||||
First(&existing).Error
|
||||
return existing, err
|
||||
@@ -108,7 +108,7 @@ func FindReusableUploadByHash(ctx context.Context, hash string, size int64) (Upl
|
||||
|
||||
// CreateUpload persists a new upload record.
|
||||
func CreateUpload(ctx context.Context, upload *Upload) error {
|
||||
return CreateUploadTx(db.DB(ctx), upload)
|
||||
return CreateUploadTx(database.DB(ctx), upload)
|
||||
}
|
||||
|
||||
// CreateUploadTx persists a new upload record within an existing transaction.
|
||||
@@ -122,7 +122,7 @@ func CreateUploadTx(tx *gorm.DB, upload *Upload) error {
|
||||
// ListUploadsByIDs returns active uploads matching the given IDs.
|
||||
func ListUploadsByIDs(ctx context.Context, ids []uint64) ([]Upload, error) {
|
||||
var uploads []Upload
|
||||
if err := db.DB(ctx).
|
||||
if err := database.DB(ctx).
|
||||
Where("id IN ? AND status IN (?, ?)", ids, UploadStatusPending, UploadStatusUsed).
|
||||
Find(&uploads).Error; err != nil {
|
||||
return nil, err
|
||||
@@ -134,13 +134,13 @@ func ListUploadsByIDs(ctx context.Context, ids []uint64) ([]Upload, error) {
|
||||
//
|
||||
//nolint:revive
|
||||
func UploadQuery(ctx context.Context) *gorm.DB {
|
||||
return db.DB(ctx).Model(&Upload{})
|
||||
return database.DB(ctx).Model(&Upload{})
|
||||
}
|
||||
|
||||
// ListUploadStats returns all upload statistics rows.
|
||||
func ListUploadStats(ctx context.Context) ([]UploadStat, error) {
|
||||
var stats []UploadStat
|
||||
if err := db.DB(ctx).Find(&stats).Error; err != nil {
|
||||
if err := database.DB(ctx).Find(&stats).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return stats, nil
|
||||
|
||||
@@ -9,7 +9,7 @@ import (
|
||||
"context"
|
||||
"strings"
|
||||
|
||||
db "github.com/Rain-kl/Wavelet/pkg/persistence"
|
||||
database "github.com/Rain-kl/Wavelet/plugins/infra/database"
|
||||
"github.com/Rain-kl/Wavelet/pkg/util"
|
||||
"github.com/Rain-kl/Wavelet/plugins/domain/upload/models"
|
||||
"gorm.io/gorm"
|
||||
@@ -27,7 +27,7 @@ type UploadListFilter struct {
|
||||
|
||||
// ListUploads returns paginated upload records matching the filter.
|
||||
func ListUploads(ctx context.Context, filter UploadListFilter) (int64, []models.Upload, error) {
|
||||
query := db.DB(ctx).Model(&models.Upload{}).
|
||||
query := database.DB(ctx).Model(&models.Upload{}).
|
||||
Where("status != ?", models.UploadStatusDeleted)
|
||||
|
||||
if filter.UserID != 0 {
|
||||
@@ -59,7 +59,7 @@ func ListUploads(ctx context.Context, filter UploadListFilter) (int64, []models.
|
||||
// GetActiveUploadByID loads a non-deleted upload by ID.
|
||||
func GetActiveUploadByID(ctx context.Context, id uint64) (models.Upload, error) {
|
||||
var upload models.Upload
|
||||
if err := db.DB(ctx).Where("id = ? AND status != ?", id, models.UploadStatusDeleted).First(&upload).Error; err != nil {
|
||||
if err := database.DB(ctx).Where("id = ? AND status != ?", id, models.UploadStatusDeleted).First(&upload).Error; err != nil {
|
||||
return models.Upload{}, err
|
||||
}
|
||||
return upload, nil
|
||||
@@ -67,7 +67,7 @@ func GetActiveUploadByID(ctx context.Context, id uint64) (models.Upload, error)
|
||||
|
||||
// SoftDeleteUpload marks an upload as deleted.
|
||||
func SoftDeleteUpload(ctx context.Context, upload *models.Upload) error {
|
||||
return SoftDeleteUploadTx(db.DB(ctx), upload)
|
||||
return SoftDeleteUploadTx(database.DB(ctx), upload)
|
||||
}
|
||||
|
||||
// SoftDeleteUploadTx marks an upload as deleted within an existing transaction.
|
||||
@@ -80,13 +80,13 @@ func UpdateUpload(ctx context.Context, upload *models.Upload, updates map[string
|
||||
if len(updates) == 0 {
|
||||
return nil
|
||||
}
|
||||
return db.DB(ctx).Model(upload).Updates(updates).Error
|
||||
return database.DB(ctx).Model(upload).Updates(updates).Error
|
||||
}
|
||||
|
||||
// ListDistinctUploadTypes returns all distinct non-empty upload business types.
|
||||
func ListDistinctUploadTypes(ctx context.Context) ([]string, error) {
|
||||
var types []string
|
||||
if err := db.DB(ctx).Model(&models.Upload{}).
|
||||
if err := database.DB(ctx).Model(&models.Upload{}).
|
||||
Where("type IS NOT NULL AND type != ''").
|
||||
Distinct().
|
||||
Pluck("type", &types).Error; err != nil {
|
||||
@@ -98,7 +98,7 @@ func ListDistinctUploadTypes(ctx context.Context) ([]string, error) {
|
||||
// FindReusableUploadByHash finds an existing upload with the same hash and size.
|
||||
func FindReusableUploadByHash(ctx context.Context, hash string, size int64) (models.Upload, error) {
|
||||
var existing models.Upload
|
||||
err := db.DB(ctx).
|
||||
err := database.DB(ctx).
|
||||
Where("hash = ? AND file_size = ? AND status IN (?, ?)", hash, size, models.UploadStatusPending, models.UploadStatusUsed).
|
||||
First(&existing).Error
|
||||
return existing, err
|
||||
@@ -106,7 +106,7 @@ func FindReusableUploadByHash(ctx context.Context, hash string, size int64) (mod
|
||||
|
||||
// CreateUpload persists a new upload record.
|
||||
func CreateUpload(ctx context.Context, upload *models.Upload) error {
|
||||
return CreateUploadTx(db.DB(ctx), upload)
|
||||
return CreateUploadTx(database.DB(ctx), upload)
|
||||
}
|
||||
|
||||
// CreateUploadTx persists a new upload record within an existing transaction.
|
||||
@@ -117,7 +117,7 @@ func CreateUploadTx(tx *gorm.DB, upload *models.Upload) error {
|
||||
// ListUploadsByIDs returns active uploads matching the given IDs.
|
||||
func ListUploadsByIDs(ctx context.Context, ids []uint64) ([]models.Upload, error) {
|
||||
var uploads []models.Upload
|
||||
if err := db.DB(ctx).
|
||||
if err := database.DB(ctx).
|
||||
Where("id IN ? AND status IN (?, ?)", ids, models.UploadStatusPending, models.UploadStatusUsed).
|
||||
Find(&uploads).Error; err != nil {
|
||||
return nil, err
|
||||
@@ -127,13 +127,13 @@ func ListUploadsByIDs(ctx context.Context, ids []uint64) ([]models.Upload, error
|
||||
|
||||
// UploadQuery returns a scoped GORM query for uploads.
|
||||
func UploadQuery(ctx context.Context) *gorm.DB {
|
||||
return db.DB(ctx).Model(&models.Upload{})
|
||||
return database.DB(ctx).Model(&models.Upload{})
|
||||
}
|
||||
|
||||
// ListUploadStats returns all upload statistics rows.
|
||||
func ListUploadStats(ctx context.Context) ([]models.UploadStat, error) {
|
||||
var stats []models.UploadStat
|
||||
if err := db.DB(ctx).Find(&stats).Error; err != nil {
|
||||
if err := database.DB(ctx).Find(&stats).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return stats, nil
|
||||
|
||||
@@ -40,4 +40,5 @@ const (
|
||||
ErrQueryImagesForCacheWarmup = "查询待预热图片失败: %w"
|
||||
ErrQueryTypeListFailed = "查询文件类型列表失败"
|
||||
ErrUpdateFileFailed = "更新文件失败"
|
||||
ErrUnauthorized = "未登录"
|
||||
)
|
||||
|
||||
@@ -8,7 +8,7 @@ import (
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/pkg/logger"
|
||||
"github.com/Rain-kl/Wavelet/pkg/persistence"
|
||||
database "github.com/Rain-kl/Wavelet/plugins/infra/database"
|
||||
"github.com/Rain-kl/Wavelet/plugins/domain/upload/models"
|
||||
"gorm.io/gorm"
|
||||
"gorm.io/gorm/clause"
|
||||
@@ -26,7 +26,7 @@ func ApplyUploadStatsRemove(ctx context.Context, upload *models.Upload) error {
|
||||
|
||||
// RebuildUploadStats rebuilds all incremental stats from current upload records.
|
||||
func RebuildUploadStats(ctx context.Context) error {
|
||||
return db.DB(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
return database.DB(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
if err := tx.Where("1 = 1").Delete(&models.UploadStat{}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
@@ -49,7 +49,7 @@ func applyUploadStatsDelta(ctx context.Context, upload *models.Upload, sign int6
|
||||
if upload == nil || !isActiveUploadStatus(upload.Status) {
|
||||
return nil
|
||||
}
|
||||
return db.DB(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
return database.DB(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
return ApplyUploadStatsDeltaTx(tx, upload, sign)
|
||||
})
|
||||
}
|
||||
|
||||
@@ -8,7 +8,7 @@ import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/pkg/persistence"
|
||||
database "github.com/Rain-kl/Wavelet/plugins/infra/database"
|
||||
"github.com/Rain-kl/Wavelet/pkg/testhelper"
|
||||
"github.com/Rain-kl/Wavelet/plugins/domain/upload/models"
|
||||
"gorm.io/gorm"
|
||||
@@ -29,7 +29,7 @@ func TestApplyUploadStatsDeltaTxWithinTransaction(t *testing.T) {
|
||||
CreatedAt: time.Now(),
|
||||
}
|
||||
|
||||
if err := db.DB(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
if err := database.DB(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
return ApplyUploadStatsDeltaTx(tx, upload, 1)
|
||||
}); err != nil {
|
||||
t.Fatalf("ApplyUploadStatsDeltaTx returned error: %v", err)
|
||||
@@ -90,7 +90,7 @@ type uploadStatsSnapshot struct {
|
||||
|
||||
func loadUploadStats(ctx context.Context) (uploadStatsSnapshot, error) {
|
||||
var rows []models.UploadStat
|
||||
if err := db.DB(ctx).Where("dimension = ?", models.UploadStatDimensionTotal).Find(&rows).Error; err != nil {
|
||||
if err := database.DB(ctx).Where("dimension = ?", models.UploadStatDimensionTotal).Find(&rows).Error; err != nil {
|
||||
return uploadStatsSnapshot{}, err
|
||||
}
|
||||
if len(rows) == 0 {
|
||||
|
||||
@@ -9,9 +9,8 @@ import (
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/pkg/task"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/plugins/domain/upload/shared"
|
||||
"github.com/Rain-kl/Wavelet/plugins/drivers/driver_asynq_worker"
|
||||
"github.com/Rain-kl/Wavelet/plugins/infra/storage/objectstore"
|
||||
)
|
||||
|
||||
@@ -71,9 +70,9 @@ func buildMigrationAccessState(ctx context.Context) MigrationAccessState {
|
||||
}
|
||||
|
||||
state := MigrationAccessState{
|
||||
ReadOnly: execution.Status != task.TaskExecutionStatusSucceeded,
|
||||
ReadOnly: execution.Status != driver_asynq_worker.TaskExecutionStatusSucceeded,
|
||||
}
|
||||
if execution.Status == task.TaskExecutionStatusSucceeded {
|
||||
if execution.Status == driver_asynq_worker.TaskExecutionStatusSucceeded {
|
||||
return state
|
||||
}
|
||||
|
||||
|
||||
@@ -10,8 +10,7 @@ import (
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/pkg/task"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/plugins/drivers/driver_asynq_worker"
|
||||
"github.com/Rain-kl/Wavelet/plugins/infra/storage/objectstore"
|
||||
)
|
||||
|
||||
@@ -19,8 +18,8 @@ import (
|
||||
const StorageMigrationTask = "storage:migrate"
|
||||
|
||||
// LatestMigrationExecution returns the most recent storage migration task execution.
|
||||
func LatestMigrationExecution(ctx context.Context) (*task.TaskExecution, bool, error) {
|
||||
return task.GetLatestTaskExecutionByTaskType(ctx, StorageMigrationTask)
|
||||
func LatestMigrationExecution(ctx context.Context) (*driver_asynq_worker.TaskExecution, bool, error) {
|
||||
return driver_asynq_worker.GetLatestTaskExecutionByTaskType(ctx, StorageMigrationTask)
|
||||
}
|
||||
|
||||
// ParseMigrationTargetConfig parses and validates a storage migration target payload.
|
||||
|
||||
@@ -10,18 +10,18 @@ import (
|
||||
"fmt"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/pkg/logger"
|
||||
"github.com/Rain-kl/Wavelet/pkg/persistence"
|
||||
"github.com/Rain-kl/Wavelet/plugins/domain/upload/models"
|
||||
"gorm.io/gorm"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/pkg/persistence/logstore"
|
||||
"github.com/Rain-kl/Wavelet/pkg/task"
|
||||
"github.com/Rain-kl/Wavelet/pkg/logger"
|
||||
database "github.com/Rain-kl/Wavelet/plugins/infra/database"
|
||||
logstore "github.com/Rain-kl/Wavelet/plugins/domain/risk_control/logstore"
|
||||
uploadcache "github.com/Rain-kl/Wavelet/plugins/domain/upload/cache"
|
||||
"github.com/Rain-kl/Wavelet/plugins/domain/upload/models"
|
||||
"github.com/Rain-kl/Wavelet/plugins/domain/upload/shared"
|
||||
uploadstats "github.com/Rain-kl/Wavelet/plugins/domain/upload/stats"
|
||||
uploadstorage "github.com/Rain-kl/Wavelet/plugins/domain/upload/storage"
|
||||
"github.com/Rain-kl/Wavelet/plugins/drivers/driver_asynq_worker"
|
||||
"github.com/Rain-kl/Wavelet/plugins/infra/storage/objectstore"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
const (
|
||||
@@ -32,14 +32,14 @@ const (
|
||||
)
|
||||
|
||||
// SystemCleanupMeta represents the task metadata.
|
||||
var SystemCleanupMeta = task.TaskMeta{
|
||||
var SystemCleanupMeta = driver_asynq_worker.TaskMeta{
|
||||
Type: TaskTypeSystemCleanup,
|
||||
AsynqTask: SystemCleanupTask,
|
||||
Name: "系统垃圾清理",
|
||||
Description: "定期清理未使用上传文件、历史推送记录和过期任务执行日志",
|
||||
SupportsTime: false,
|
||||
MaxRetry: task.DefaultMaxRetry,
|
||||
Queue: task.QueueDefault,
|
||||
MaxRetry: driver_asynq_worker.DefaultMaxRetry,
|
||||
Queue: driver_asynq_worker.QueueDefault,
|
||||
Retryable: true,
|
||||
}
|
||||
|
||||
@@ -47,7 +47,7 @@ var SystemCleanupMeta = task.TaskMeta{
|
||||
type SystemCleanupHandler struct{}
|
||||
|
||||
// Execute 执行系统清理(包含文件清理、历史推送日志和任务执行日志清理)
|
||||
func (h *SystemCleanupHandler) Execute(ctx context.Context, _ []byte) (*task.TaskResult, error) {
|
||||
func (h *SystemCleanupHandler) Execute(ctx context.Context, _ []byte) (*driver_asynq_worker.TaskResult, error) {
|
||||
if uploadstorage.ReadOnly(ctx) {
|
||||
return nil, errors.New(shared.ErrStorageReadOnly)
|
||||
}
|
||||
@@ -58,16 +58,16 @@ func (h *SystemCleanupHandler) Execute(ctx context.Context, _ []byte) (*task.Tas
|
||||
|
||||
oneHourAgo := time.Now().Add(-1 * time.Hour)
|
||||
|
||||
task.AppendLog(ctx, "开始扫描未使用上传文件,阈值: %s", oneHourAgo.Format(time.RFC3339))
|
||||
driver_asynq_worker.AppendLog(ctx, "开始扫描未使用上传文件,阈值: %s", oneHourAgo.Format(time.RFC3339))
|
||||
|
||||
for {
|
||||
var unusedUploads []models.Upload
|
||||
if err := db.DB(ctx).
|
||||
if err := database.DB(ctx).
|
||||
Where("id > ? AND status = ? AND created_at < ?", lastID, models.UploadStatusPending, oneHourAgo).
|
||||
Order("id ASC").
|
||||
Limit(batchSize).
|
||||
Find(&unusedUploads).Error; err != nil {
|
||||
task.AppendLog(ctx, "查询未使用的上传文件失败: %v", err)
|
||||
driver_asynq_worker.AppendLog(ctx, "查询未使用的上传文件失败: %v", err)
|
||||
return nil, fmt.Errorf(shared.ErrQueryUnusedUploadsFailed, err)
|
||||
}
|
||||
|
||||
@@ -75,12 +75,12 @@ func (h *SystemCleanupHandler) Execute(ctx context.Context, _ []byte) (*task.Tas
|
||||
break
|
||||
}
|
||||
|
||||
task.AppendLog(ctx, "本批次找到 %d 个需要清理的上传文件", len(unusedUploads))
|
||||
driver_asynq_worker.AppendLog(ctx, "本批次找到 %d 个需要清理的上传文件", len(unusedUploads))
|
||||
|
||||
for _, u := range unusedUploads {
|
||||
totalProcessed++
|
||||
|
||||
if err := db.DB(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
if err := database.DB(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
if err := tx.Model(&models.Upload{}).
|
||||
Where("id = ? AND status = ?", u.ID, models.UploadStatusPending).
|
||||
Update("status", models.UploadStatusDeleted).Error; err != nil {
|
||||
@@ -97,7 +97,7 @@ func (h *SystemCleanupHandler) Execute(ctx context.Context, _ []byte) (*task.Tas
|
||||
|
||||
return nil
|
||||
}); err != nil {
|
||||
task.AppendLog(ctx, "清理上传文件失败 [ID:%d]: %v", u.ID, err)
|
||||
driver_asynq_worker.AppendLog(ctx, "清理上传文件失败 [ID:%d]: %v", u.ID, err)
|
||||
lastID = u.ID
|
||||
continue
|
||||
}
|
||||
@@ -109,28 +109,28 @@ func (h *SystemCleanupHandler) Execute(ctx context.Context, _ []byte) (*task.Tas
|
||||
}
|
||||
}
|
||||
|
||||
task.AppendLog(ctx, "开始清理历史推送审计日志,只保留最近7天数据...")
|
||||
driver_asynq_worker.AppendLog(ctx, "开始清理历史推送审计日志,只保留最近7天数据...")
|
||||
cutoff := time.Now().AddDate(0, 0, -7)
|
||||
var pushHistoryCount int64
|
||||
if err := db.DB(ctx).Table("w_push_histories").Where("created_at < ?", cutoff).Count(&pushHistoryCount).Error; err != nil {
|
||||
task.AppendLog(ctx, "统计待清理的历史推送记录失败: %v", err)
|
||||
if err := database.DB(ctx).Table("w_push_histories").Where("created_at < ?", cutoff).Count(&pushHistoryCount).Error; err != nil {
|
||||
driver_asynq_worker.AppendLog(ctx, "统计待清理的历史推送记录失败: %v", err)
|
||||
} else if pushHistoryCount > 0 {
|
||||
if err := db.DB(ctx).Table("w_push_histories").Where("created_at < ?", cutoff).Delete(map[string]any{}).Error; err != nil {
|
||||
task.AppendLog(ctx, "删除历史推送记录失败: %v", err)
|
||||
if err := database.DB(ctx).Table("w_push_histories").Where("created_at < ?", cutoff).Delete(map[string]any{}).Error; err != nil {
|
||||
driver_asynq_worker.AppendLog(ctx, "删除历史推送记录失败: %v", err)
|
||||
} else {
|
||||
task.AppendLog(ctx, "成功删除 %d 条历史推送记录 (截止时间: %s)", pushHistoryCount, cutoff.Format("2006-01-02 15:04:05"))
|
||||
driver_asynq_worker.AppendLog(ctx, "成功删除 %d 条历史推送记录 (截止时间: %s)", pushHistoryCount, cutoff.Format("2006-01-02 15:04:05"))
|
||||
}
|
||||
} else {
|
||||
task.AppendLog(ctx, "没有需要清理的历史推送记录 (截止时间: %s)", cutoff.Format("2006-01-02 15:04:05"))
|
||||
driver_asynq_worker.AppendLog(ctx, "没有需要清理的历史推送记录 (截止时间: %s)", cutoff.Format("2006-01-02 15:04:05"))
|
||||
}
|
||||
|
||||
task.AppendLog(ctx, "开始清理任务执行日志:高频任务保留最近3天,低频任务保留最近30天...")
|
||||
taskLogStats, err := task.CleanupTaskExecutionLogs(ctx, time.Now())
|
||||
driver_asynq_worker.AppendLog(ctx, "开始清理任务执行日志:高频任务保留最近3天,低频任务保留最近30天...")
|
||||
taskLogStats, err := driver_asynq_worker.CleanupTaskExecutionLogs(ctx, time.Now())
|
||||
if err != nil {
|
||||
task.AppendLog(ctx, "清理任务执行日志失败: %v", err)
|
||||
driver_asynq_worker.AppendLog(ctx, "清理任务执行日志失败: %v", err)
|
||||
logger.ErrorF(ctx, "清理任务执行日志失败: %v", err)
|
||||
} else {
|
||||
task.AppendLog(ctx, "成功清理任务执行日志 %d 条(高频 %d 条,低频 %d 条)",
|
||||
driver_asynq_worker.AppendLog(ctx, "成功清理任务执行日志 %d 条(高频 %d 条,低频 %d 条)",
|
||||
taskLogStats.HighFrequencyDeleted+taskLogStats.LowFrequencyDeleted,
|
||||
taskLogStats.HighFrequencyDeleted,
|
||||
taskLogStats.LowFrequencyDeleted,
|
||||
@@ -140,11 +140,11 @@ func (h *SystemCleanupHandler) Execute(ctx context.Context, _ []byte) (*task.Tas
|
||||
var logDeleted int64
|
||||
logSummary, logErr := logstore.CleanupExpired(ctx)
|
||||
if logErr != nil {
|
||||
task.AppendLog(ctx, "清理过期用户访问日志失败: %v", logErr)
|
||||
driver_asynq_worker.AppendLog(ctx, "清理过期用户访问日志失败: %v", logErr)
|
||||
logger.ErrorF(ctx, "清理过期用户访问日志失败: %v", logErr)
|
||||
} else {
|
||||
logDeleted = logSummary.Deleted
|
||||
task.AppendLog(ctx, "成功清理过期用户访问日志 %d 条(%s 保留 %d 天)",
|
||||
driver_asynq_worker.AppendLog(ctx, "成功清理过期用户访问日志 %d 条(%s 保留 %d 天)",
|
||||
logSummary.Deleted, logSummary.ActiveDatabase, logSummary.RetentionDays)
|
||||
}
|
||||
|
||||
@@ -155,6 +155,6 @@ func (h *SystemCleanupHandler) Execute(ctx context.Context, _ []byte) (*task.Tas
|
||||
taskLogStats.HighFrequencyDeleted+taskLogStats.LowFrequencyDeleted,
|
||||
logDeleted,
|
||||
)
|
||||
task.AppendLog(ctx, "%s", msg)
|
||||
return &task.TaskResult{Message: msg}, nil
|
||||
driver_asynq_worker.AppendLog(ctx, "%s", msg)
|
||||
return &driver_asynq_worker.TaskResult{Message: msg}, nil
|
||||
}
|
||||
|
||||
@@ -7,10 +7,10 @@ import (
|
||||
"context"
|
||||
"fmt"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/pkg/persistence"
|
||||
"github.com/Rain-kl/Wavelet/pkg/task"
|
||||
database "github.com/Rain-kl/Wavelet/plugins/infra/database"
|
||||
"github.com/Rain-kl/Wavelet/plugins/domain/upload/models"
|
||||
uploadstats "github.com/Rain-kl/Wavelet/plugins/domain/upload/stats"
|
||||
"github.com/Rain-kl/Wavelet/plugins/drivers/driver_asynq_worker"
|
||||
)
|
||||
|
||||
const (
|
||||
@@ -21,14 +21,14 @@ const (
|
||||
)
|
||||
|
||||
// RebuildUploadStatsMeta describes the upload stats rebuild task.
|
||||
var RebuildUploadStatsMeta = task.TaskMeta{
|
||||
var RebuildUploadStatsMeta = driver_asynq_worker.TaskMeta{
|
||||
Type: TaskTypeRebuildUploadStats,
|
||||
AsynqTask: RebuildUploadStatsTask,
|
||||
Name: "重算文件存储统计",
|
||||
Description: "根据当前 w_uploads 活跃记录全量重建 w_upload_stats(总量、类型、分类、趋势)",
|
||||
SupportsTime: false,
|
||||
MaxRetry: task.DefaultMaxRetry,
|
||||
Queue: task.QueueDefault,
|
||||
MaxRetry: driver_asynq_worker.DefaultMaxRetry,
|
||||
Queue: driver_asynq_worker.QueueDefault,
|
||||
Retryable: true,
|
||||
}
|
||||
|
||||
@@ -36,28 +36,28 @@ var RebuildUploadStatsMeta = task.TaskMeta{
|
||||
type RebuildUploadStatsHandler struct{}
|
||||
|
||||
// Execute scans active uploads and rebuilds all upload stat dimensions.
|
||||
func (h *RebuildUploadStatsHandler) Execute(ctx context.Context, _ []byte) (*task.TaskResult, error) {
|
||||
func (h *RebuildUploadStatsHandler) Execute(ctx context.Context, _ []byte) (*driver_asynq_worker.TaskResult, error) {
|
||||
var activeCount int64
|
||||
if err := db.DB(ctx).
|
||||
if err := database.DB(ctx).
|
||||
Model(&models.Upload{}).
|
||||
Where("status != ?", models.UploadStatusDeleted).
|
||||
Count(&activeCount).Error; err != nil {
|
||||
task.AppendLog(ctx, "统计活跃上传记录失败: %v", err)
|
||||
driver_asynq_worker.AppendLog(ctx, "统计活跃上传记录失败: %v", err)
|
||||
return nil, fmt.Errorf("count active uploads: %w", err)
|
||||
}
|
||||
|
||||
task.AppendLog(ctx, "开始重算文件存储统计,活跃记录数: %d", activeCount)
|
||||
driver_asynq_worker.AppendLog(ctx, "开始重算文件存储统计,活跃记录数: %d", activeCount)
|
||||
|
||||
if err := uploadstats.RebuildUploadStats(ctx); err != nil {
|
||||
task.AppendLog(ctx, "重算文件存储统计失败: %v", err)
|
||||
driver_asynq_worker.AppendLog(ctx, "重算文件存储统计失败: %v", err)
|
||||
return nil, fmt.Errorf("rebuild upload stats: %w", err)
|
||||
}
|
||||
|
||||
var totalStat models.UploadStat
|
||||
if err := db.DB(ctx).
|
||||
if err := database.DB(ctx).
|
||||
Where("dimension = ? AND stat_key = ?", models.UploadStatDimensionTotal, "").
|
||||
First(&totalStat).Error; err != nil {
|
||||
task.AppendLog(ctx, "读取总量统计失败: %v", err)
|
||||
driver_asynq_worker.AppendLog(ctx, "读取总量统计失败: %v", err)
|
||||
return nil, fmt.Errorf("load total upload stats: %w", err)
|
||||
}
|
||||
|
||||
@@ -67,6 +67,6 @@ func (h *RebuildUploadStatsHandler) Execute(ctx context.Context, _ []byte) (*tas
|
||||
totalStat.FileCount,
|
||||
totalStat.FileSize,
|
||||
)
|
||||
task.AppendLog(ctx, "%s", msg)
|
||||
return &task.TaskResult{Message: msg}, nil
|
||||
driver_asynq_worker.AppendLog(ctx, "%s", msg)
|
||||
return &driver_asynq_worker.TaskResult{Message: msg}, nil
|
||||
}
|
||||
|
||||
@@ -8,7 +8,7 @@ import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/pkg/persistence"
|
||||
database "github.com/Rain-kl/Wavelet/plugins/infra/database"
|
||||
"github.com/Rain-kl/Wavelet/pkg/testhelper"
|
||||
"github.com/Rain-kl/Wavelet/plugins/domain/upload/models"
|
||||
)
|
||||
@@ -33,13 +33,13 @@ func TestRebuildUploadStatsHandler_Execute(t *testing.T) {
|
||||
},
|
||||
}
|
||||
for i := range uploads {
|
||||
if err := db.DB(ctx).Create(&uploads[i]).Error; err != nil {
|
||||
if err := database.DB(ctx).Create(&uploads[i]).Error; err != nil {
|
||||
t.Fatalf("seed upload failed: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// Corrupt stats to ensure rebuild recalculates from uploads.
|
||||
if err := db.DB(ctx).Create(&models.UploadStat{
|
||||
if err := database.DB(ctx).Create(&models.UploadStat{
|
||||
Dimension: models.UploadStatDimensionTotal,
|
||||
StatKey: "",
|
||||
FileCount: 0,
|
||||
@@ -58,7 +58,7 @@ func TestRebuildUploadStatsHandler_Execute(t *testing.T) {
|
||||
}
|
||||
|
||||
var totalStat models.UploadStat
|
||||
if err := db.DB(ctx).
|
||||
if err := database.DB(ctx).
|
||||
Where("dimension = ? AND stat_key = ?", models.UploadStatDimensionTotal, "").
|
||||
First(&totalStat).Error; err != nil {
|
||||
t.Fatalf("load total stat failed: %v", err)
|
||||
|
||||
@@ -15,14 +15,16 @@ import (
|
||||
"sync/atomic"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/pkg/persistence"
|
||||
"github.com/Rain-kl/Wavelet/pkg/task"
|
||||
"golang.org/x/sync/errgroup"
|
||||
|
||||
cache "github.com/Rain-kl/Wavelet/plugins/infra/cache"
|
||||
database "github.com/Rain-kl/Wavelet/plugins/infra/database"
|
||||
"github.com/Rain-kl/Wavelet/pkg/util"
|
||||
"github.com/Rain-kl/Wavelet/plugins/domain/upload/models"
|
||||
uploadstats "github.com/Rain-kl/Wavelet/plugins/domain/upload/stats"
|
||||
uploadstorage "github.com/Rain-kl/Wavelet/plugins/domain/upload/storage"
|
||||
"github.com/Rain-kl/Wavelet/plugins/drivers/driver_asynq_worker"
|
||||
"github.com/Rain-kl/Wavelet/plugins/infra/storage/objectstore"
|
||||
"golang.org/x/sync/errgroup"
|
||||
)
|
||||
|
||||
const (
|
||||
@@ -33,16 +35,16 @@ const (
|
||||
)
|
||||
|
||||
// StorageMigrationMeta describes the manually dispatchable migration task.
|
||||
var StorageMigrationMeta = task.TaskMeta{
|
||||
var StorageMigrationMeta = driver_asynq_worker.TaskMeta{
|
||||
Type: TaskTypeStorageMigration,
|
||||
AsynqTask: StorageMigrationTask,
|
||||
Name: "迁移文件存储",
|
||||
Description: "将活动存储中的文件迁移到待切换的目标存储,迁移期间文件系统保持只读",
|
||||
SupportsTime: false,
|
||||
MaxRetry: task.DefaultMaxRetry,
|
||||
Queue: task.QueueDefault,
|
||||
MaxRetry: driver_asynq_worker.DefaultMaxRetry,
|
||||
Queue: driver_asynq_worker.QueueDefault,
|
||||
Retryable: true,
|
||||
Params: []task.TaskParam{
|
||||
Params: []driver_asynq_worker.TaskParam{
|
||||
{
|
||||
Name: "target",
|
||||
Label: "目标存储配置 (JSON)",
|
||||
@@ -74,15 +76,15 @@ func (h *MigrationHandler) ValidatePayload(payload []byte) ([]byte, error) {
|
||||
}
|
||||
|
||||
// Execute migrates all unique active-storage objects to the pending backend.
|
||||
func (h *MigrationHandler) Execute(ctx context.Context, payload []byte) (*task.TaskResult, error) {
|
||||
if db.Redis != nil {
|
||||
func (h *MigrationHandler) Execute(ctx context.Context, payload []byte) (*driver_asynq_worker.TaskResult, error) {
|
||||
if cache.Redis != nil {
|
||||
const (
|
||||
cleanupTimeout = 5 * time.Second
|
||||
renewalInterval = 10 * time.Minute
|
||||
)
|
||||
|
||||
lockKey := db.PrefixedKey("lock:storage:migrate")
|
||||
ok, err := db.Redis.SetNX(ctx, lockKey, "locked", time.Hour).Result()
|
||||
lockKey := cache.PrefixedKey("lock:storage:migrate")
|
||||
ok, err := cache.Redis.SetNX(ctx, lockKey, "locked", time.Hour).Result()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("acquire migration lock: %w", err)
|
||||
}
|
||||
@@ -96,7 +98,7 @@ func (h *MigrationHandler) Execute(ctx context.Context, payload []byte) (*task.T
|
||||
close(stopRenewal)
|
||||
cleanupCtx, cancel := context.WithTimeout(context.Background(), cleanupTimeout)
|
||||
defer cancel()
|
||||
_ = db.Redis.Del(cleanupCtx, lockKey)
|
||||
_ = cache.Redis.Del(cleanupCtx, lockKey)
|
||||
}()
|
||||
|
||||
//nolint:contextcheck,gosec
|
||||
@@ -107,7 +109,7 @@ func (h *MigrationHandler) Execute(ctx context.Context, payload []byte) (*task.T
|
||||
select {
|
||||
case <-ticker.C:
|
||||
renewCtx, cancel := context.WithTimeout(context.Background(), cleanupTimeout)
|
||||
_ = db.Redis.Expire(renewCtx, lockKey, time.Hour).Err()
|
||||
_ = cache.Redis.Expire(renewCtx, lockKey, time.Hour).Err()
|
||||
cancel()
|
||||
case <-stopRenewal:
|
||||
return
|
||||
@@ -131,8 +133,8 @@ func (h *MigrationHandler) Execute(ctx context.Context, payload []byte) (*task.T
|
||||
return nil, fmt.Errorf("activate same-driver storage config: %w", err)
|
||||
}
|
||||
message := fmt.Sprintf("存储配置已更新,活动存储保持为 %s", target.Driver)
|
||||
task.AppendLog(ctx, "%s", message)
|
||||
return &task.TaskResult{Message: message}, nil
|
||||
driver_asynq_worker.AppendLog(ctx, "%s", message)
|
||||
return &driver_asynq_worker.TaskResult{Message: message}, nil
|
||||
}
|
||||
|
||||
total, err := countStorageObjects(ctx)
|
||||
@@ -144,8 +146,8 @@ func (h *MigrationHandler) Execute(ctx context.Context, payload []byte) (*task.T
|
||||
return nil, fmt.Errorf("activate empty storage config: %w", err)
|
||||
}
|
||||
message := fmt.Sprintf("当前存储没有需要迁移的对象,活动存储已切换为 %s", target.Driver)
|
||||
task.AppendLog(ctx, "%s", message)
|
||||
return &task.TaskResult{Message: message}, nil
|
||||
driver_asynq_worker.AppendLog(ctx, "%s", message)
|
||||
return &driver_asynq_worker.TaskResult{Message: message}, nil
|
||||
}
|
||||
|
||||
sourceBackend, err := objectstore.NewBackend(ctx, active, active.Driver)
|
||||
@@ -157,7 +159,7 @@ func (h *MigrationHandler) Execute(ctx context.Context, payload []byte) (*task.T
|
||||
return nil, fmt.Errorf("create target storage: %w", err)
|
||||
}
|
||||
|
||||
task.AppendLog(ctx, "开始存储迁移: %s -> %s,总对象数: %d", active.Driver, target.Driver, total)
|
||||
driver_asynq_worker.AppendLog(ctx, "开始存储迁移: %s -> %s,总对象数: %d", active.Driver, target.Driver, total)
|
||||
migrated, err := migrateObjects(ctx, sourceBackend, targetBackend, total)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
@@ -167,13 +169,13 @@ func (h *MigrationHandler) Execute(ctx context.Context, payload []byte) (*task.T
|
||||
return nil, fmt.Errorf("activate target storage: %w", err)
|
||||
}
|
||||
message := fmt.Sprintf("存储迁移完成,共迁移 %d 个对象,活动存储已切换为 %s", migrated, target.Driver)
|
||||
task.AppendLog(ctx, "%s", message)
|
||||
return &task.TaskResult{Message: message}, nil
|
||||
driver_asynq_worker.AppendLog(ctx, "%s", message)
|
||||
return &driver_asynq_worker.TaskResult{Message: message}, nil
|
||||
}
|
||||
|
||||
func countStorageObjects(ctx context.Context) (int64, error) {
|
||||
var count int64
|
||||
err := db.DB(ctx).Model(&models.Upload{}).
|
||||
err := database.DB(ctx).Model(&models.Upload{}).
|
||||
Where("status != ?", models.UploadStatusDeleted).
|
||||
Distinct("file_path").
|
||||
Count(&count).Error
|
||||
@@ -185,7 +187,7 @@ func hasUnresolvedMigrationTask(ctx context.Context) (bool, error) {
|
||||
if err != nil || !ok {
|
||||
return false, err
|
||||
}
|
||||
return execution.Status == task.TaskExecutionStatusPending || execution.Status == task.TaskExecutionStatusRunning, nil
|
||||
return execution.Status == driver_asynq_worker.TaskExecutionStatusPending || execution.Status == driver_asynq_worker.TaskExecutionStatusRunning, nil
|
||||
}
|
||||
|
||||
type migrationObject struct {
|
||||
@@ -211,10 +213,10 @@ func migrateObjects(
|
||||
return atomic.LoadInt64(&migrated), fmt.Errorf("storage migration canceled: %w", err)
|
||||
}
|
||||
|
||||
task.AppendLog(ctx, "正在查询待迁移对象批次,当前已完成迁移: %d/%d", atomic.LoadInt64(&migrated), total)
|
||||
driver_asynq_worker.AppendLog(ctx, "正在查询待迁移对象批次,当前已完成迁移: %d/%d", atomic.LoadInt64(&migrated), total)
|
||||
|
||||
var objects []migrationObject
|
||||
query := db.DB(ctx).Model(&models.Upload{}).
|
||||
query := database.DB(ctx).Model(&models.Upload{}).
|
||||
Select("file_path, MAX(file_size) AS file_size, MAX(mime_type) AS mime_type, MAX(hash) AS hash").
|
||||
Where("status != ?", models.UploadStatusDeleted)
|
||||
if lastFilePath != "" {
|
||||
@@ -227,12 +229,12 @@ func migrateObjects(
|
||||
return atomic.LoadInt64(&migrated), fmt.Errorf("query source objects: %w", err)
|
||||
}
|
||||
if len(objects) == 0 {
|
||||
task.AppendLog(ctx, "所有对象迁移完毕")
|
||||
driver_asynq_worker.AppendLog(ctx, "所有对象迁移完毕")
|
||||
break
|
||||
}
|
||||
|
||||
lastFilePath = objects[len(objects)-1].FilePath
|
||||
task.AppendLog(ctx, "获取当前批次迁移对象,批次大小: %d,实际获取对象数: %d", batchSize, len(objects))
|
||||
driver_asynq_worker.AppendLog(ctx, "获取当前批次迁移对象,批次大小: %d,实际获取对象数: %d", batchSize, len(objects))
|
||||
|
||||
var g errgroup.Group
|
||||
g.SetLimit(migrationConcurrency)
|
||||
@@ -252,7 +254,7 @@ func migrateObjects(
|
||||
return atomic.LoadInt64(&migrated), err
|
||||
}
|
||||
|
||||
task.AppendLog(ctx, "当前批次迁移完成。迁移进度: %d/%d", atomic.LoadInt64(&migrated), total)
|
||||
driver_asynq_worker.AppendLog(ctx, "当前批次迁移完成。迁移进度: %d/%d", atomic.LoadInt64(&migrated), total)
|
||||
}
|
||||
return atomic.LoadInt64(&migrated), nil
|
||||
}
|
||||
@@ -265,11 +267,11 @@ func migrateSingleObject(
|
||||
sha256HexLength int,
|
||||
) error {
|
||||
if shouldSkipMigration(ctx, targetBackend, obj) {
|
||||
task.AppendLog(ctx, "[跳过迁移] 目标存储已存在相同文件: %s", obj.FilePath)
|
||||
driver_asynq_worker.AppendLog(ctx, "[跳过迁移] 目标存储已存在相同文件: %s", obj.FilePath)
|
||||
return nil
|
||||
}
|
||||
|
||||
task.AppendLog(ctx, "[迁移开始] 正在从源存储读取文件: %s", obj.FilePath)
|
||||
driver_asynq_worker.AppendLog(ctx, "[迁移开始] 正在从源存储读取文件: %s", obj.FilePath)
|
||||
source, err := sourceBackend.Get(ctx, obj.FilePath)
|
||||
if err != nil {
|
||||
if isNotFoundError(err) {
|
||||
@@ -277,7 +279,7 @@ func migrateSingleObject(
|
||||
}
|
||||
return fmt.Errorf("open source object %q: %w", obj.FilePath, err)
|
||||
}
|
||||
task.AppendLog(ctx, "[传输中] 正在向目标存储上传文件: %s (大小: %d 字节, 类型: %s)", obj.FilePath, obj.FileSize, obj.MimeType)
|
||||
driver_asynq_worker.AppendLog(ctx, "[传输中] 正在向目标存储上传文件: %s (大小: %d 字节, 类型: %s)", obj.FilePath, obj.FileSize, obj.MimeType)
|
||||
targetResult, putErr := targetBackend.Put(ctx, obj.FilePath, source.Body, obj.FileSize, obj.MimeType)
|
||||
closeErr := source.Body.Close()
|
||||
if putErr != nil {
|
||||
@@ -288,7 +290,7 @@ func migrateSingleObject(
|
||||
}
|
||||
|
||||
if len(obj.Hash) == sha256HexLength {
|
||||
task.AppendLog(ctx, "[校验中] 正在对目标文件进行数据一致性校验 (SHA-256): %s", targetResult.Key)
|
||||
driver_asynq_worker.AppendLog(ctx, "[校验中] 正在对目标文件进行数据一致性校验 (SHA-256): %s", targetResult.Key)
|
||||
targetObj, getErr := targetBackend.Get(ctx, targetResult.Key)
|
||||
if getErr != nil {
|
||||
return fmt.Errorf("retrieve target object for verification %q: %w", obj.FilePath, getErr)
|
||||
@@ -306,18 +308,18 @@ func migrateSingleObject(
|
||||
if computedHash != obj.Hash {
|
||||
return fmt.Errorf("integrity check failed for %q: got hash %s, want %s", obj.FilePath, computedHash, obj.Hash)
|
||||
}
|
||||
task.AppendLog(ctx, "[校验通过] 文件一致性校验成功: %s", targetResult.Key)
|
||||
driver_asynq_worker.AppendLog(ctx, "[校验通过] 文件一致性校验成功: %s", targetResult.Key)
|
||||
}
|
||||
|
||||
if targetResult.Key != obj.FilePath {
|
||||
task.AppendLog(ctx, "[更新数据库] 正在更新文件路径: %s -> %s", obj.FilePath, targetResult.Key)
|
||||
if err := db.DB(ctx).Model(&models.Upload{}).
|
||||
driver_asynq_worker.AppendLog(ctx, "[更新数据库] 正在更新文件路径: %s -> %s", obj.FilePath, targetResult.Key)
|
||||
if err := database.DB(ctx).Model(&models.Upload{}).
|
||||
Where("file_path = ? AND status != ?", obj.FilePath, models.UploadStatusDeleted).
|
||||
Update("file_path", targetResult.Key).Error; err != nil {
|
||||
return fmt.Errorf("update migrated object %q: %w", obj.FilePath, err)
|
||||
}
|
||||
}
|
||||
task.AppendLog(ctx, "[迁移成功] 文件已完成迁移: %s", targetResult.Key)
|
||||
driver_asynq_worker.AppendLog(ctx, "[迁移成功] 文件已完成迁移: %s", targetResult.Key)
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -342,15 +344,15 @@ func markMissingMigrationObjectDeleted(
|
||||
filePath string,
|
||||
sourceErr error,
|
||||
) error {
|
||||
task.AppendLog(ctx, "警告: 源存储中物理文件不存在,标记为已删除并跳过: %s (错误: %v)", filePath, sourceErr)
|
||||
driver_asynq_worker.AppendLog(ctx, "警告: 源存储中物理文件不存在,标记为已删除并跳过: %s (错误: %v)", filePath, sourceErr)
|
||||
|
||||
var affectedUploads []models.Upload
|
||||
if err := db.DB(ctx).
|
||||
if err := database.DB(ctx).
|
||||
Where("file_path = ? AND status != ?", filePath, models.UploadStatusDeleted).
|
||||
Find(&affectedUploads).Error; err != nil {
|
||||
return fmt.Errorf("load missing object uploads %q: %w", filePath, err)
|
||||
}
|
||||
if err := db.DB(ctx).Model(&models.Upload{}).
|
||||
if err := database.DB(ctx).Model(&models.Upload{}).
|
||||
Where("file_path = ?", filePath).
|
||||
Update("status", models.UploadStatusDeleted).Error; err != nil {
|
||||
return fmt.Errorf("update missing object %q: %w", filePath, err)
|
||||
|
||||
@@ -16,7 +16,7 @@ import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/pkg/persistence"
|
||||
cache "github.com/Rain-kl/Wavelet/plugins/infra/cache"
|
||||
"github.com/Rain-kl/Wavelet/pkg/testhelper"
|
||||
"github.com/Rain-kl/Wavelet/plugins/domain/upload/models"
|
||||
"github.com/Rain-kl/Wavelet/plugins/infra/storage/objectstore"
|
||||
@@ -239,16 +239,16 @@ func TestMigrationHandlerExecuteWithRedisLock(t *testing.T) {
|
||||
})
|
||||
defer rdb.Close()
|
||||
|
||||
oldRedis := db.Redis
|
||||
db.Redis = rdb
|
||||
oldRedis := cache.Redis
|
||||
cache.Redis = rdb
|
||||
defer func() {
|
||||
db.Redis = oldRedis
|
||||
cache.Redis = oldRedis
|
||||
}()
|
||||
|
||||
ctx := context.Background()
|
||||
|
||||
// Acquire lock manually
|
||||
lockKey := db.PrefixedKey("lock:storage:migrate")
|
||||
lockKey := cache.PrefixedKey("lock:storage:migrate")
|
||||
if err := rdb.Set(ctx, lockKey, "locked", time.Hour).Err(); err != nil {
|
||||
t.Fatalf("Failed to set manual lock in Redis: %v", err)
|
||||
}
|
||||
|
||||
@@ -12,11 +12,11 @@ import (
|
||||
"strings"
|
||||
"sync"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/pkg/persistence"
|
||||
"github.com/Rain-kl/Wavelet/pkg/task"
|
||||
database "github.com/Rain-kl/Wavelet/plugins/infra/database"
|
||||
"github.com/Rain-kl/Wavelet/plugins/domain/upload/filesrv"
|
||||
"github.com/Rain-kl/Wavelet/plugins/domain/upload/models"
|
||||
"github.com/Rain-kl/Wavelet/plugins/domain/upload/shared"
|
||||
"github.com/Rain-kl/Wavelet/plugins/drivers/driver_asynq_worker"
|
||||
)
|
||||
|
||||
const (
|
||||
@@ -29,16 +29,16 @@ const (
|
||||
var warmImageCacheMu sync.Mutex
|
||||
|
||||
// WarmImageCacheMeta represents the image cache warmup task metadata.
|
||||
var WarmImageCacheMeta = task.TaskMeta{
|
||||
var WarmImageCacheMeta = driver_asynq_worker.TaskMeta{
|
||||
Type: TaskTypeWarmImageCache,
|
||||
AsynqTask: WarmImageCacheTask,
|
||||
Name: "预热图片压缩缓存",
|
||||
Description: "串行将文件管理中的图片转换为指定质量的 WebP 并写入永久缓存",
|
||||
SupportsTime: false,
|
||||
MaxRetry: task.DefaultMaxRetry,
|
||||
Queue: task.QueueDefault,
|
||||
MaxRetry: driver_asynq_worker.DefaultMaxRetry,
|
||||
Queue: driver_asynq_worker.QueueDefault,
|
||||
Retryable: true,
|
||||
Params: []task.TaskParam{
|
||||
Params: []driver_asynq_worker.TaskParam{
|
||||
{
|
||||
Name: "quality",
|
||||
Label: "图片质量",
|
||||
@@ -80,10 +80,10 @@ func (h *WarmImageCacheHandler) ValidatePayload(payload []byte) ([]byte, error)
|
||||
}
|
||||
|
||||
// Execute serially converts all managed images to WebP cache entries.
|
||||
func (h *WarmImageCacheHandler) Execute(ctx context.Context, payload []byte) (*task.TaskResult, error) {
|
||||
func (h *WarmImageCacheHandler) Execute(ctx context.Context, payload []byte) (*driver_asynq_worker.TaskResult, error) {
|
||||
normalizedPayload, err := h.ValidatePayload(payload)
|
||||
if err != nil {
|
||||
task.AppendLog(ctx, "图片缓存预热参数无效: %v", err)
|
||||
driver_asynq_worker.AppendLog(ctx, "图片缓存预热参数无效: %v", err)
|
||||
return nil, err
|
||||
}
|
||||
|
||||
@@ -92,7 +92,7 @@ func (h *WarmImageCacheHandler) Execute(ctx context.Context, payload []byte) (*t
|
||||
return nil, fmt.Errorf(shared.ErrParseImageCacheWarmupPayload, err)
|
||||
}
|
||||
|
||||
task.AppendLog(ctx, "等待获取图片缓存预热执行锁,质量: %s", req.Quality)
|
||||
driver_asynq_worker.AppendLog(ctx, "等待获取图片缓存预热执行锁,质量: %s", req.Quality)
|
||||
warmImageCacheMu.Lock()
|
||||
defer warmImageCacheMu.Unlock()
|
||||
|
||||
@@ -106,7 +106,7 @@ func (h *WarmImageCacheHandler) Execute(ctx context.Context, payload []byte) (*t
|
||||
var totalGenerated int
|
||||
var totalFailed int
|
||||
|
||||
task.AppendLog(ctx, "开始串行预热图片压缩缓存,质量: %s,每批: %d", req.Quality, batchSize)
|
||||
driver_asynq_worker.AppendLog(ctx, "开始串行预热图片压缩缓存,质量: %s,每批: %d", req.Quality, batchSize)
|
||||
|
||||
for {
|
||||
if err := ctx.Err(); err != nil {
|
||||
@@ -114,7 +114,7 @@ func (h *WarmImageCacheHandler) Execute(ctx context.Context, payload []byte) (*t
|
||||
}
|
||||
|
||||
var uploads []models.Upload
|
||||
if err := db.DB(ctx).
|
||||
if err := database.DB(ctx).
|
||||
Where("id > ? AND status != ? AND (LOWER(mime_type) LIKE ? OR LOWER(extension) IN ?)",
|
||||
lastID,
|
||||
models.UploadStatusDeleted,
|
||||
@@ -124,7 +124,7 @@ func (h *WarmImageCacheHandler) Execute(ctx context.Context, payload []byte) (*t
|
||||
Order("id ASC").
|
||||
Limit(batchSize).
|
||||
Find(&uploads).Error; err != nil {
|
||||
task.AppendLog(ctx, "查询图片上传记录失败: %v", err)
|
||||
driver_asynq_worker.AppendLog(ctx, "查询图片上传记录失败: %v", err)
|
||||
return nil, fmt.Errorf(shared.ErrQueryImagesForCacheWarmup, err)
|
||||
}
|
||||
|
||||
@@ -149,7 +149,7 @@ func (h *WarmImageCacheHandler) Execute(ctx context.Context, payload []byte) (*t
|
||||
totalFailed++
|
||||
batchFailed++
|
||||
if totalFailed <= maxFailureLogs {
|
||||
task.AppendLog(ctx, "图片处理失败 [ID:%d]: %v", upload.ID, err)
|
||||
driver_asynq_worker.AppendLog(ctx, "图片处理失败 [ID:%d]: %v", upload.ID, err)
|
||||
}
|
||||
continue
|
||||
}
|
||||
@@ -162,7 +162,7 @@ func (h *WarmImageCacheHandler) Execute(ctx context.Context, payload []byte) (*t
|
||||
batchGenerated++
|
||||
}
|
||||
|
||||
task.AppendLog(
|
||||
driver_asynq_worker.AppendLog(
|
||||
ctx,
|
||||
"批次完成,末尾 ID: %d,生成: %d,命中: %d,失败: %d",
|
||||
lastID,
|
||||
@@ -179,6 +179,6 @@ func (h *WarmImageCacheHandler) Execute(ctx context.Context, payload []byte) (*t
|
||||
totalCached,
|
||||
totalFailed,
|
||||
)
|
||||
task.AppendLog(ctx, "%s", msg)
|
||||
return &task.TaskResult{Message: msg}, nil
|
||||
driver_asynq_worker.AppendLog(ctx, "%s", msg)
|
||||
return &driver_asynq_worker.TaskResult{Message: msg}, nil
|
||||
}
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user