mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-09 00:56:37 +08:00
feat(framework): 回灌 OpenFlare 分层、安全与运行时改进
将平台域持久化收敛为 repository 唯一入口,model 去掉 IO。 邮件头写入前清除 CR/LF,防止 header 注入。 httppool 支持可配置 Transport;batchwriter 增加 MinBatchSize/Stats,flush 失败交回批次;任务 PermanentError 作为 SkipRetry 终态。 设置与推送页的确认改为 AlertDialog;axios 去尾斜杠并按 Gin 数组序列化查询参数。 升级共享 Go 依赖(Gin、Asynq、OTel、GORM、Redis 等)。
This commit is contained in:
@@ -0,0 +1,69 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package repository
|
||||
|
||||
import (
|
||||
"context"
|
||||
|
||||
db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
)
|
||||
|
||||
// ListAccessTokensByUserID returns all access tokens for a user ordered by created_at desc.
|
||||
func ListAccessTokensByUserID(ctx context.Context, userID uint64) ([]model.AccessToken, error) {
|
||||
var tokens []model.AccessToken
|
||||
if err := db.DB(ctx).Where("user_id = ?", userID).Order("created_at desc").Find(&tokens).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return tokens, nil
|
||||
}
|
||||
|
||||
// CountAccessTokensByUserID returns how many access tokens a user owns.
|
||||
func CountAccessTokensByUserID(ctx context.Context, userID uint64) (int64, error) {
|
||||
var count int64
|
||||
if err := db.DB(ctx).Model(&model.AccessToken{}).Where("user_id = ?", userID).Count(&count).Error; err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return count, nil
|
||||
}
|
||||
|
||||
// CreateAccessToken inserts a new access token record.
|
||||
func CreateAccessToken(ctx context.Context, record *model.AccessToken) error {
|
||||
return db.DB(ctx).Create(record).Error
|
||||
}
|
||||
|
||||
// GetAccessTokenByIDAndUserID loads a token owned by the given user.
|
||||
func GetAccessTokenByIDAndUserID(ctx context.Context, id, userID uint64) (model.AccessToken, error) {
|
||||
var token model.AccessToken
|
||||
if err := db.DB(ctx).Where("id = ? AND user_id = ?", id, userID).First(&token).Error; err != nil {
|
||||
return model.AccessToken{}, err
|
||||
}
|
||||
return token, nil
|
||||
}
|
||||
|
||||
// DeleteAccessTokenForUser deletes a token if it belongs to the user.
|
||||
// Returns the number of rows affected.
|
||||
func DeleteAccessTokenForUser(ctx context.Context, id, userID uint64) (int64, error) {
|
||||
tx := db.DB(ctx).Where("id = ? AND user_id = ?", id, userID).Delete(&model.AccessToken{})
|
||||
return tx.RowsAffected, tx.Error
|
||||
}
|
||||
|
||||
// GetAccessTokenByHash loads an access token by its token hash.
|
||||
func GetAccessTokenByHash(ctx context.Context, tokenHash string) (model.AccessToken, error) {
|
||||
var token model.AccessToken
|
||||
if err := db.DB(ctx).Where("token_hash = ?", tokenHash).First(&token).Error; err != nil {
|
||||
return model.AccessToken{}, err
|
||||
}
|
||||
return token, nil
|
||||
}
|
||||
|
||||
// SaveAccessToken persists all fields of an existing access token.
|
||||
func SaveAccessToken(ctx context.Context, record *model.AccessToken) error {
|
||||
return db.DB(ctx).Save(record).Error
|
||||
}
|
||||
|
||||
// DeleteAccessTokensByUserID deletes all access tokens for a user.
|
||||
func DeleteAccessTokensByUserID(ctx context.Context, userID uint64) error {
|
||||
return db.DB(ctx).Where("user_id = ?", userID).Delete(&model.AccessToken{}).Error
|
||||
}
|
||||
@@ -5,6 +5,7 @@ package analytics
|
||||
|
||||
import (
|
||||
"context"
|
||||
"io"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
@@ -169,6 +170,12 @@ func (m *mockConn) Exec(_ context.Context, _ string, _ ...any) error { return ni
|
||||
|
||||
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{} }
|
||||
|
||||
@@ -0,0 +1,215 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package repository
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"strings"
|
||||
|
||||
"gorm.io/gorm"
|
||||
|
||||
db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
)
|
||||
|
||||
// GetAuthSources 获取所有认证源(已脱敏)
|
||||
func GetAuthSources(ctx context.Context) ([]model.AuthSource, error) {
|
||||
var sources []model.AuthSource
|
||||
if err := db.DB(ctx).Order("id asc").Find(&sources).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
for i := range sources {
|
||||
sources[i].Sanitize()
|
||||
}
|
||||
return sources, nil
|
||||
}
|
||||
|
||||
// GetActiveAuthSources 获取所有已启用的认证源(已脱敏)
|
||||
func GetActiveAuthSources(ctx context.Context) ([]model.AuthSource, error) {
|
||||
var sources []model.AuthSource
|
||||
if err := db.DB(ctx).Where("is_active = ?", true).Order("id asc").Find(&sources).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
for i := range sources {
|
||||
sources[i].Sanitize()
|
||||
}
|
||||
return sources, nil
|
||||
}
|
||||
|
||||
// GetAuthSourceByID 根据 ID 获取认证源
|
||||
func GetAuthSourceByID(ctx context.Context, id uint64) (*model.AuthSource, error) {
|
||||
if id == 0 {
|
||||
return nil, errors.New(errAuthSourceIDRequired)
|
||||
}
|
||||
var source model.AuthSource
|
||||
if err := db.DB(ctx).First(&source, "id = ?", id).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
source.ClientSecretConfigured = source.ClientSecret != ""
|
||||
return &source, nil
|
||||
}
|
||||
|
||||
// GetAuthSourceByName 根据名称获取认证源(名称比较不区分大小写)
|
||||
func GetAuthSourceByName(ctx context.Context, name string) (*model.AuthSource, error) {
|
||||
name = strings.TrimSpace(name)
|
||||
if name == "" {
|
||||
return nil, errors.New(errAuthSourceNameRequired)
|
||||
}
|
||||
var source model.AuthSource
|
||||
if err := db.DB(ctx).First(&source, "LOWER(name) = LOWER(?)", name).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
source.ClientSecretConfigured = source.ClientSecret != ""
|
||||
return &source, nil
|
||||
}
|
||||
|
||||
// CreateAuthSource 创建认证源
|
||||
func CreateAuthSource(ctx context.Context, source *model.AuthSource) error {
|
||||
if err := source.Validate(); err != nil {
|
||||
return err
|
||||
}
|
||||
return db.DB(ctx).Create(source).Error
|
||||
}
|
||||
|
||||
// UpdateAuthSource 更新认证源,keepSecret 为 true 时保留原密钥
|
||||
func UpdateAuthSource(ctx context.Context, source *model.AuthSource, keepSecret bool) error {
|
||||
if source.ID == 0 {
|
||||
return errors.New(errAuthSourceIDRequired)
|
||||
}
|
||||
var current model.AuthSource
|
||||
if err := db.DB(ctx).First(¤t, "id = ?", source.ID).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if keepSecret {
|
||||
source.ClientSecret = current.ClientSecret
|
||||
}
|
||||
if err := source.Validate(); err != nil {
|
||||
return err
|
||||
}
|
||||
return db.DB(ctx).Model(¤t).Updates(map[string]any{
|
||||
colName: source.Name,
|
||||
"type": source.Type,
|
||||
"display_name": source.DisplayName,
|
||||
"is_active": source.IsActive,
|
||||
"client_id": source.ClientID,
|
||||
"client_secret": source.ClientSecret,
|
||||
"openid_discovery_url": source.OpenIDDiscoveryURL,
|
||||
"scopes": source.Scopes,
|
||||
"icon_url": source.IconURL,
|
||||
}).Error
|
||||
}
|
||||
|
||||
// ToggleAuthSource 切换认证源启用状态
|
||||
func ToggleAuthSource(ctx context.Context, id uint64, isActive bool) error {
|
||||
source, err := GetAuthSourceByID(ctx, id)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
source.IsActive = isActive
|
||||
if err := source.Validate(); err != nil {
|
||||
return err
|
||||
}
|
||||
return db.DB(ctx).Model(&model.AuthSource{}).Where("id = ?", id).Update("is_active", isActive).Error
|
||||
}
|
||||
|
||||
// DeleteAuthSource 删除认证源及其关联的外部帐号绑定
|
||||
func DeleteAuthSource(ctx context.Context, id uint64) error {
|
||||
if id == 0 {
|
||||
return errors.New(errAuthSourceIDRequired)
|
||||
}
|
||||
return db.DB(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
if err := tx.Where("auth_source_id = ?", id).Delete(&model.ExternalAccount{}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
return tx.Delete(&model.AuthSource{}, "id = ?", id).Error
|
||||
})
|
||||
}
|
||||
|
||||
// FindExternalAccount 查找外部帐号绑定记录
|
||||
func FindExternalAccount(ctx context.Context, sourceID uint64, externalID string) (*model.ExternalAccount, error) {
|
||||
var account model.ExternalAccount
|
||||
if err := db.DB(ctx).Where("auth_source_id = ? AND external_id = ?", sourceID, externalID).First(&account).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &account, nil
|
||||
}
|
||||
|
||||
// BindExternalAccount 绑定外部帐号(已存在时更新用户名和邮箱)
|
||||
func BindExternalAccount(ctx context.Context, account *model.ExternalAccount) error {
|
||||
if account.UserID == 0 || strings.TrimSpace(account.ExternalID) == "" {
|
||||
return errors.New(errExternalAccountBindingIncomplete)
|
||||
}
|
||||
account.ExternalID = strings.TrimSpace(account.ExternalID)
|
||||
account.ExternalUsername = strings.TrimSpace(account.ExternalUsername)
|
||||
account.Email = strings.TrimSpace(account.Email)
|
||||
|
||||
return db.DB(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
var current model.ExternalAccount
|
||||
err := tx.Where("auth_source_id = ? AND external_id = ?", account.AuthSourceID, account.ExternalID).First(¤t).Error
|
||||
if err == nil {
|
||||
if current.UserID != account.UserID {
|
||||
return errors.New(errExternalAccountAlreadyBoundToAnother)
|
||||
}
|
||||
return tx.Model(¤t).Updates(map[string]any{
|
||||
"external_username": account.ExternalUsername,
|
||||
"email": account.Email,
|
||||
}).Error
|
||||
}
|
||||
if !errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return err
|
||||
}
|
||||
return tx.Create(account).Error
|
||||
})
|
||||
}
|
||||
|
||||
// ListExternalAccountsByUserID 获取指定用户的所有外部帐号绑定视图
|
||||
func ListExternalAccountsByUserID(ctx context.Context, userID uint64) ([]model.ExternalAccountView, error) {
|
||||
if userID == 0 {
|
||||
return nil, errors.New(errUserIDRequired)
|
||||
}
|
||||
var accounts []model.ExternalAccount
|
||||
if err := db.DB(ctx).Where("user_id = ?", userID).Order("id asc").Find(&accounts).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
views := make([]model.ExternalAccountView, 0, len(accounts))
|
||||
for _, account := range accounts {
|
||||
var name, sourceType, label string
|
||||
if account.AuthSourceID == 0 {
|
||||
name = "default"
|
||||
sourceType = "oidc"
|
||||
label = "历史认证源"
|
||||
} else {
|
||||
source, err := GetAuthSourceByID(ctx, account.AuthSourceID)
|
||||
if err != nil {
|
||||
continue
|
||||
}
|
||||
name = source.Name
|
||||
sourceType = source.Type
|
||||
label = source.DisplayName
|
||||
if label == "" {
|
||||
label = source.Name
|
||||
}
|
||||
}
|
||||
views = append(views, model.ExternalAccountView{
|
||||
ID: account.ID,
|
||||
AuthSourceID: account.AuthSourceID,
|
||||
AuthSourceName: name,
|
||||
AuthSourceType: sourceType,
|
||||
AuthSourceLabel: label,
|
||||
ExternalUsername: account.ExternalUsername,
|
||||
Email: account.Email,
|
||||
CreatedAt: account.CreatedAt,
|
||||
})
|
||||
}
|
||||
return views, nil
|
||||
}
|
||||
|
||||
// DeleteExternalAccountForUser 删除指定用户的外部帐号绑定
|
||||
func DeleteExternalAccountForUser(ctx context.Context, id uint64, userID uint64) error {
|
||||
if id == 0 || userID == 0 {
|
||||
return errors.New(errExternalAccountBindingIDRequired)
|
||||
}
|
||||
return db.DB(ctx).Where("id = ? AND user_id = ?", id, userID).Delete(&model.ExternalAccount{}).Error
|
||||
}
|
||||
@@ -178,7 +178,7 @@ func GetActiveAuthSourcesCached(ctx context.Context) ([]model.AuthSource, error)
|
||||
}
|
||||
}
|
||||
|
||||
sources, err := model.GetActiveAuthSources(ctx)
|
||||
sources, err := GetActiveAuthSources(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -192,7 +192,7 @@ func GetAuthSourceByNameCached(ctx context.Context, name string) (*model.AuthSou
|
||||
|
||||
normalized := normalizeAuthSourceName(name)
|
||||
if normalized == "" {
|
||||
return model.GetAuthSourceByName(ctx, name)
|
||||
return GetAuthSourceByName(ctx, name)
|
||||
}
|
||||
|
||||
if source, ok := authSourceByNameRAM.GetIfPresent(normalized); ok {
|
||||
@@ -210,7 +210,7 @@ func GetAuthSourceByNameCached(ctx context.Context, name string) (*model.AuthSou
|
||||
}
|
||||
}
|
||||
|
||||
source, err := model.GetAuthSourceByName(ctx, name)
|
||||
source, err := GetAuthSourceByName(ctx, name)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
@@ -74,7 +74,7 @@ func TestGetActiveAuthSourcesCached_LoadsFromRedisBeforeDB(t *testing.T) {
|
||||
ClientSecret: "client-secret",
|
||||
OpenIDDiscoveryURL: "https://issuer.example.com",
|
||||
}
|
||||
if err := model.CreateAuthSource(ctx, &source); err != nil {
|
||||
if err := CreateAuthSource(ctx, &source); err != nil {
|
||||
t.Fatalf("CreateAuthSource() error = %v", err)
|
||||
}
|
||||
|
||||
@@ -119,7 +119,7 @@ func TestGetAuthSourceByNameCached_LoadsFromRedisBeforeDB(t *testing.T) {
|
||||
ClientSecret: "client-secret",
|
||||
OpenIDDiscoveryURL: "https://issuer.example.com",
|
||||
}
|
||||
if err := model.CreateAuthSource(ctx, &source); err != nil {
|
||||
if err := CreateAuthSource(ctx, &source); err != nil {
|
||||
t.Fatalf("CreateAuthSource() error = %v", err)
|
||||
}
|
||||
|
||||
@@ -164,7 +164,7 @@ func TestInvalidateAuthSourceCache_ClearsRedisKeys(t *testing.T) {
|
||||
ClientSecret: "client-secret",
|
||||
OpenIDDiscoveryURL: "https://issuer.example.com",
|
||||
}
|
||||
if err := model.CreateAuthSource(ctx, &source); err != nil {
|
||||
if err := CreateAuthSource(ctx, &source); err != nil {
|
||||
t.Fatalf("CreateAuthSource() error = %v", err)
|
||||
}
|
||||
if _, err := GetActiveAuthSourcesCached(ctx); err != nil {
|
||||
@@ -209,7 +209,7 @@ func TestAuthSourceInvalidationPubSubClearsPeerRAM(t *testing.T) {
|
||||
ClientSecret: "client-secret",
|
||||
OpenIDDiscoveryURL: "https://issuer.example.com",
|
||||
}
|
||||
if err := model.CreateAuthSource(ctx, &source); err != nil {
|
||||
if err := CreateAuthSource(ctx, &source); err != nil {
|
||||
t.Fatalf("CreateAuthSource() error = %v", err)
|
||||
}
|
||||
|
||||
|
||||
@@ -0,0 +1,25 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package repository
|
||||
|
||||
// Persistence and repository-layer parameter messages live here (unexported).
|
||||
// Domain field validation used by model.Validate stays in internal/model/errs.go;
|
||||
// repository may call model.Validate and return those errors as-is.
|
||||
// Keep wording aligned with model where the same user-facing phrase applies,
|
||||
// but do not import or re-export model unexported consts (would require exporting).
|
||||
const (
|
||||
errDatabaseNotInitialized = "database not initialized"
|
||||
errConfigIntParseFailed = "配置 %s 的值 '%s' 无法转换为整数: %w"
|
||||
errConfigDecimalParseFailed = "配置 %s 的值 '%s' 无法转换为decimal: %w"
|
||||
errConfigBoolParseFailed = "配置 %s 的值 '%s' 无法转换为布尔值: %w"
|
||||
errParseMenuDisplayConfigFailed = "解析目录显示配置失败: %w"
|
||||
errAuthSourceNameRequired = "认证源名称不能为空"
|
||||
errAuthSourceIDRequired = "认证源 ID 不能为空"
|
||||
errExternalAccountBindingIncomplete = "外部账号绑定信息不完整"
|
||||
errExternalAccountAlreadyBoundToAnother = "该外部账号已绑定到其他用户"
|
||||
errUserIDRequired = "用户 ID 不能为空"
|
||||
errExternalAccountBindingIDRequired = "绑定记录 ID 不能为空"
|
||||
)
|
||||
|
||||
const colName = "name"
|
||||
@@ -0,0 +1,53 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package repository
|
||||
|
||||
import (
|
||||
"context"
|
||||
|
||||
db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
)
|
||||
|
||||
// CreateSchedule 创建定时任务
|
||||
func CreateSchedule(ctx context.Context, schedule *model.Schedule) error {
|
||||
return db.DB(ctx).Create(schedule).Error
|
||||
}
|
||||
|
||||
// UpdateSchedule 更新定时任务
|
||||
func UpdateSchedule(ctx context.Context, schedule *model.Schedule) error {
|
||||
return db.DB(ctx).Save(schedule).Error
|
||||
}
|
||||
|
||||
// DeleteSchedule 删除定时任务
|
||||
func DeleteSchedule(ctx context.Context, id uint64) error {
|
||||
return db.DB(ctx).Delete(&model.Schedule{}, id).Error
|
||||
}
|
||||
|
||||
// GetScheduleByID 根据 ID 获取定时任务
|
||||
func GetScheduleByID(ctx context.Context, id uint64) (*model.Schedule, error) {
|
||||
var schedule model.Schedule
|
||||
if err := db.DB(ctx).Where("id = ?", id).First(&schedule).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &schedule, nil
|
||||
}
|
||||
|
||||
// ListSchedules 获取所有定时任务
|
||||
func ListSchedules(ctx context.Context) ([]model.Schedule, error) {
|
||||
var schedules []model.Schedule
|
||||
if err := db.DB(ctx).Order("id DESC").Find(&schedules).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return schedules, nil
|
||||
}
|
||||
|
||||
// ListActiveSchedules 获取所有启用的定时任务
|
||||
func ListActiveSchedules(ctx context.Context) ([]model.Schedule, error) {
|
||||
var schedules []model.Schedule
|
||||
if err := db.DB(ctx).Where("is_active = ?", true).Find(&schedules).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return schedules, nil
|
||||
}
|
||||
@@ -18,14 +18,7 @@ import (
|
||||
"github.com/Rain-kl/Wavelet/pkg/cache/ram"
|
||||
)
|
||||
|
||||
const (
|
||||
configTypeSystem = "system"
|
||||
errDatabaseNotInitialized = "database not initialized"
|
||||
errConfigIntParseFailed = "配置 %s 的值 '%s' 无法转换为整数: %w"
|
||||
errConfigDecimalParseFailed = "配置 %s 的值 '%s' 无法转换为decimal: %w"
|
||||
errConfigBoolParseFailed = "配置 %s 的值 '%s' 无法转换为布尔值: %w"
|
||||
errParseMenuDisplayConfigFailed = "解析目录显示配置失败: %w"
|
||||
)
|
||||
const configTypeSystem = "system"
|
||||
|
||||
// PreheatSystemConfigs loads all system configs from database.
|
||||
// This function strictly performs database read and does not perform any cache read or write operations.
|
||||
|
||||
@@ -0,0 +1,304 @@
|
||||
// Copyright 2025 linux.do
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package repository
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/redis/go-redis/v9"
|
||||
"gorm.io/gorm"
|
||||
|
||||
db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
|
||||
"github.com/Rain-kl/Wavelet/internal/infra/persistence/idgen"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
)
|
||||
|
||||
const (
|
||||
taskExecutionLogRedisKeyPrefix = "task:execution:log:"
|
||||
taskExecutionLogExpiration = 24 * time.Hour
|
||||
taskExecutionLogMaxLines = 1000
|
||||
)
|
||||
|
||||
// CreateTaskExecution 创建任务执行记录
|
||||
func CreateTaskExecution(ctx context.Context, execution *model.TaskExecution) error {
|
||||
execution.ID = idgen.NextUint64ID()
|
||||
return db.DB(ctx).Create(execution).Error
|
||||
}
|
||||
|
||||
// UpdateTaskExecution 更新任务执行记录,忽略由 Redis 缓冲和归档流程管理的 log 字段。
|
||||
func UpdateTaskExecution(ctx context.Context, execution *model.TaskExecution) error {
|
||||
return db.DB(ctx).Omit("log").Save(execution).Error
|
||||
}
|
||||
|
||||
// GetTaskExecutionByTaskID 根据 TaskID 获取执行记录
|
||||
func GetTaskExecutionByTaskID(ctx context.Context, taskID string) (*model.TaskExecution, error) {
|
||||
var execution model.TaskExecution
|
||||
if err := db.DB(ctx).Where("task_id = ?", taskID).First(&execution).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := loadTaskExecutionLog(ctx, &execution); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &execution, nil
|
||||
}
|
||||
|
||||
// GetTaskExecutionByID 根据 ID 获取执行记录
|
||||
func GetTaskExecutionByID(ctx context.Context, id uint64) (*model.TaskExecution, error) {
|
||||
var execution model.TaskExecution
|
||||
if err := db.DB(ctx).Where("id = ?", id).First(&execution).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := loadTaskExecutionLog(ctx, &execution); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &execution, nil
|
||||
}
|
||||
|
||||
// GetLatestTaskExecutionByTaskType returns the most recent execution for a task type.
|
||||
// ok is false when no row exists.
|
||||
func GetLatestTaskExecutionByTaskType(ctx context.Context, taskType string) (*model.TaskExecution, bool, error) {
|
||||
var execution model.TaskExecution
|
||||
err := db.DB(ctx).
|
||||
Where("task_type = ?", taskType).
|
||||
Order("id DESC").
|
||||
First(&execution).Error
|
||||
if err == nil {
|
||||
if loadErr := loadTaskExecutionLog(ctx, &execution); loadErr != nil {
|
||||
return nil, false, loadErr
|
||||
}
|
||||
return &execution, true, nil
|
||||
}
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return nil, false, nil
|
||||
}
|
||||
return nil, false, err
|
||||
}
|
||||
|
||||
// AppendTaskExecutionLog 将日志追加到 Redis 缓冲,任务完成后再持久化到数据库。
|
||||
func AppendTaskExecutionLog(ctx context.Context, taskID string, logLine string) error {
|
||||
if db.Redis == nil {
|
||||
return errors.New("redis client is not initialized")
|
||||
}
|
||||
|
||||
now := time.Now().Format("15:04:05")
|
||||
line := fmt.Sprintf("[%s] %s\n", now, logLine)
|
||||
key := taskExecutionLogRedisKey(taskID)
|
||||
|
||||
_, err := db.Redis.TxPipelined(ctx, func(pipe redis.Pipeliner) error {
|
||||
pipe.RPush(ctx, key, line)
|
||||
pipe.LTrim(ctx, key, -taskExecutionLogMaxLines, -1)
|
||||
pipe.Expire(ctx, key, taskExecutionLogExpiration)
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
return fmt.Errorf("append task execution log to redis: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// FlushTaskExecutionLog 将 Redis 中的完整任务日志写入数据库,并在成功后清理缓存。
|
||||
func FlushTaskExecutionLog(ctx context.Context, taskID string) error {
|
||||
if db.Redis == nil {
|
||||
return errors.New("redis client is not initialized")
|
||||
}
|
||||
|
||||
key := taskExecutionLogRedisKey(taskID)
|
||||
logLines, err := db.Redis.LRange(ctx, key, 0, -1).Result()
|
||||
if err != nil {
|
||||
return fmt.Errorf("get task execution log from redis: %w", err)
|
||||
}
|
||||
if len(logLines) == 0 {
|
||||
return nil
|
||||
}
|
||||
logText := strings.Join(logLines, "")
|
||||
|
||||
result := db.DB(ctx).Model(&model.TaskExecution{}).
|
||||
Where("task_id = ?", taskID).
|
||||
Update("log", logText)
|
||||
if result.Error != nil {
|
||||
return fmt.Errorf("persist task execution log: %w", result.Error)
|
||||
}
|
||||
if result.RowsAffected == 0 {
|
||||
return fmt.Errorf("persist task execution log: task %q not found", taskID)
|
||||
}
|
||||
|
||||
if err := db.Redis.Del(ctx, key).Err(); err != nil {
|
||||
return fmt.Errorf("delete persisted task execution log from redis: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// ListTaskExecutions 分页查询任务执行记录
|
||||
func ListTaskExecutions(ctx context.Context, req model.ListTaskExecutionsRequest) ([]model.TaskExecution, int64, error) {
|
||||
if req.Page <= 0 {
|
||||
req.Page = 1
|
||||
}
|
||||
if req.PageSize <= 0 {
|
||||
req.PageSize = 20
|
||||
}
|
||||
|
||||
query := db.DB(ctx).Model(&model.TaskExecution{})
|
||||
|
||||
if req.Status != "" {
|
||||
query = query.Where("status = ?", req.Status)
|
||||
}
|
||||
if req.TaskType != "" {
|
||||
query = query.Where("task_type = ?", req.TaskType)
|
||||
} else if types := parseTaskTypesFilter(req.TaskTypes); len(types) > 0 {
|
||||
query = query.Where("task_type IN ?", types)
|
||||
} else if req.TaskTypePrefix != "" {
|
||||
query = query.Where("task_type LIKE ?", req.TaskTypePrefix+"%")
|
||||
}
|
||||
|
||||
var total int64
|
||||
if err := query.Count(&total).Error; err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
|
||||
var executions []model.TaskExecution
|
||||
offset := (req.Page - 1) * req.PageSize
|
||||
if err := query.Order("id DESC").Offset(offset).Limit(req.PageSize).Find(&executions).Error; err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
if err := loadTaskExecutionLogs(ctx, executions); err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
|
||||
return executions, total, nil
|
||||
}
|
||||
|
||||
func parseTaskTypesFilter(raw string) []string {
|
||||
if strings.TrimSpace(raw) == "" {
|
||||
return nil
|
||||
}
|
||||
parts := strings.Split(raw, ",")
|
||||
out := make([]string, 0, len(parts))
|
||||
for _, part := range parts {
|
||||
part = strings.TrimSpace(part)
|
||||
if part != "" {
|
||||
out = append(out, part)
|
||||
}
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// MarkFailedTaskExecutionsSucceededTx marks failed executions of a task type as succeeded within a transaction.
|
||||
func MarkFailedTaskExecutionsSucceededTx(
|
||||
tx *gorm.DB,
|
||||
taskType string,
|
||||
result string,
|
||||
finishedAt time.Time,
|
||||
) error {
|
||||
return tx.Model(&model.TaskExecution{}).
|
||||
Where("task_type = ? AND status = ?", taskType, model.TaskExecutionStatusFailed).
|
||||
Updates(map[string]any{
|
||||
"status": model.TaskExecutionStatusSucceeded,
|
||||
"result": result,
|
||||
"finished_at": finishedAt,
|
||||
}).Error
|
||||
}
|
||||
|
||||
// CleanupTaskExecutionLogs removes finished task execution logs according to frequency-based retention.
|
||||
func CleanupTaskExecutionLogs(ctx context.Context, now time.Time) (model.TaskExecutionCleanupStats, error) {
|
||||
const (
|
||||
frequencyWindowDays = 30
|
||||
highFrequencyThreshold = frequencyWindowDays
|
||||
)
|
||||
|
||||
frequencyWindowStart := now.AddDate(0, 0, -frequencyWindowDays)
|
||||
highFrequencyCutoff := now.AddDate(0, 0, -3)
|
||||
lowFrequencyCutoff := now.AddDate(0, 0, -30)
|
||||
terminalStatuses := []model.TaskExecutionStatus{model.TaskExecutionStatusSucceeded, model.TaskExecutionStatusFailed}
|
||||
|
||||
var highFrequencyTaskTypes []string
|
||||
if err := db.DB(ctx).
|
||||
Model(&model.TaskExecution{}).
|
||||
Select("task_type").
|
||||
Where("created_at >= ?", frequencyWindowStart).
|
||||
Group("task_type").
|
||||
Having("COUNT(*) > ?", highFrequencyThreshold).
|
||||
Pluck("task_type", &highFrequencyTaskTypes).Error; err != nil {
|
||||
return model.TaskExecutionCleanupStats{}, fmt.Errorf("query high-frequency task types: %w", err)
|
||||
}
|
||||
|
||||
var highFrequencyDeleted int64
|
||||
if len(highFrequencyTaskTypes) > 0 {
|
||||
highFrequencyResult := db.DB(ctx).
|
||||
Where("status IN ?", terminalStatuses).
|
||||
Where("created_at < ?", highFrequencyCutoff).
|
||||
Where("task_type IN ?", highFrequencyTaskTypes).
|
||||
Delete(&model.TaskExecution{})
|
||||
if highFrequencyResult.Error != nil {
|
||||
return model.TaskExecutionCleanupStats{}, fmt.Errorf("delete high-frequency task execution logs: %w", highFrequencyResult.Error)
|
||||
}
|
||||
highFrequencyDeleted = highFrequencyResult.RowsAffected
|
||||
}
|
||||
|
||||
lowFrequencyQuery := db.DB(ctx).
|
||||
Where("status IN ?", terminalStatuses).
|
||||
Where("created_at < ?", lowFrequencyCutoff)
|
||||
if len(highFrequencyTaskTypes) > 0 {
|
||||
lowFrequencyQuery = lowFrequencyQuery.Where("task_type NOT IN ?", highFrequencyTaskTypes)
|
||||
}
|
||||
lowFrequencyResult := lowFrequencyQuery.Delete(&model.TaskExecution{})
|
||||
if lowFrequencyResult.Error != nil {
|
||||
return model.TaskExecutionCleanupStats{}, fmt.Errorf("delete low-frequency task execution logs: %w", lowFrequencyResult.Error)
|
||||
}
|
||||
|
||||
return model.TaskExecutionCleanupStats{
|
||||
HighFrequencyDeleted: highFrequencyDeleted,
|
||||
LowFrequencyDeleted: lowFrequencyResult.RowsAffected,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func taskExecutionLogRedisKey(taskID string) string {
|
||||
return db.PrefixedKey(taskExecutionLogRedisKeyPrefix + taskID)
|
||||
}
|
||||
|
||||
func loadTaskExecutionLog(ctx context.Context, execution *model.TaskExecution) error {
|
||||
if db.Redis == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
logLines, err := db.Redis.LRange(ctx, taskExecutionLogRedisKey(execution.TaskID), 0, -1).Result()
|
||||
if err != nil {
|
||||
return fmt.Errorf("get task execution log from redis: %w", err)
|
||||
}
|
||||
if len(logLines) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
execution.Log = strings.Join(logLines, "")
|
||||
return nil
|
||||
}
|
||||
|
||||
func loadTaskExecutionLogs(ctx context.Context, executions []model.TaskExecution) error {
|
||||
if db.Redis == nil || len(executions) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
commands := make([]*redis.StringSliceCmd, len(executions))
|
||||
_, err := db.Redis.Pipelined(ctx, func(pipe redis.Pipeliner) error {
|
||||
for i := range executions {
|
||||
commands[i] = pipe.LRange(ctx, taskExecutionLogRedisKey(executions[i].TaskID), 0, -1)
|
||||
}
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
return fmt.Errorf("get task execution logs from redis: %w", err)
|
||||
}
|
||||
|
||||
for i := range executions {
|
||||
logLines := commands[i].Val()
|
||||
if len(logLines) > 0 {
|
||||
executions[i].Log = strings.Join(logLines, "")
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,514 @@
|
||||
// Copyright 2025 linux.do
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package repository
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
|
||||
db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
|
||||
"github.com/alicebob/miniredis/v2"
|
||||
"github.com/glebarez/sqlite"
|
||||
"github.com/redis/go-redis/v9"
|
||||
"github.com/redis/go-redis/v9/maintnotifications"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
func setupTaskExecutionTestEnvironment(t *testing.T) func() {
|
||||
sqliteDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{
|
||||
DisableForeignKeyConstraintWhenMigrating: true,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
err = sqliteDB.AutoMigrate(&model.TaskExecution{})
|
||||
require.NoError(t, err)
|
||||
|
||||
miniRedis, err := miniredis.Run()
|
||||
require.NoError(t, err)
|
||||
redisClient := redis.NewClient(&redis.Options{
|
||||
Addr: miniRedis.Addr(),
|
||||
MaintNotificationsConfig: &maintnotifications.Config{
|
||||
Mode: maintnotifications.ModeDisabled,
|
||||
},
|
||||
})
|
||||
|
||||
db.SetDB(sqliteDB)
|
||||
db.Redis = redisClient
|
||||
|
||||
return func() {
|
||||
require.NoError(t, redisClient.Close())
|
||||
miniRedis.Close()
|
||||
db.SetDB(nil)
|
||||
db.Redis = nil
|
||||
}
|
||||
}
|
||||
|
||||
func TestCreateTaskExecution(t *testing.T) {
|
||||
cleanup := setupTaskExecutionTestEnvironment(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
execution := &model.TaskExecution{
|
||||
TaskID: "manual_cleanup_123",
|
||||
TaskType: "system:cleanup",
|
||||
TaskName: "清理未使用上传",
|
||||
Status: model.TaskExecutionStatusPending,
|
||||
Retryable: true,
|
||||
MaxRetry: 3,
|
||||
RetryCount: 0,
|
||||
Payload: `{"test": true}`,
|
||||
TriggeredBy: "manual",
|
||||
}
|
||||
|
||||
err := CreateTaskExecution(ctx, execution)
|
||||
require.NoError(t, err)
|
||||
assert.NotZero(t, execution.ID, "ID should be generated")
|
||||
assert.NotZero(t, execution.CreatedAt, "CreatedAt should be set")
|
||||
assert.NotZero(t, execution.UpdatedAt, "UpdatedAt should be set")
|
||||
}
|
||||
|
||||
func TestGetTaskExecutionByTaskID(t *testing.T) {
|
||||
cleanup := setupTaskExecutionTestEnvironment(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
// 创建记录
|
||||
execution := &model.TaskExecution{
|
||||
TaskID: "test_task_id_001",
|
||||
TaskType: "system:cleanup",
|
||||
TaskName: "清理未使用上传",
|
||||
Status: model.TaskExecutionStatusPending,
|
||||
Retryable: true,
|
||||
MaxRetry: 3,
|
||||
TriggeredBy: "manual",
|
||||
}
|
||||
err := CreateTaskExecution(ctx, execution)
|
||||
require.NoError(t, err)
|
||||
|
||||
// 按 TaskID 查询
|
||||
found, err := GetTaskExecutionByTaskID(ctx, "test_task_id_001")
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, execution.ID, found.ID)
|
||||
assert.Equal(t, "test_task_id_001", found.TaskID)
|
||||
assert.Equal(t, model.TaskExecutionStatusPending, found.Status)
|
||||
assert.True(t, found.Retryable)
|
||||
assert.Equal(t, 3, found.MaxRetry)
|
||||
|
||||
// 查询不存在的 TaskID
|
||||
_, err = GetTaskExecutionByTaskID(ctx, "nonexistent")
|
||||
assert.Error(t, err, "should return error for non-existent taskID")
|
||||
}
|
||||
|
||||
func TestGetTaskExecutionByID(t *testing.T) {
|
||||
cleanup := setupTaskExecutionTestEnvironment(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
execution := &model.TaskExecution{
|
||||
TaskID: "test_by_id_001",
|
||||
TaskType: "system:cleanup",
|
||||
TaskName: "清理未使用上传",
|
||||
Status: model.TaskExecutionStatusPending,
|
||||
TriggeredBy: "system",
|
||||
}
|
||||
err := CreateTaskExecution(ctx, execution)
|
||||
require.NoError(t, err)
|
||||
|
||||
// 按主键查询
|
||||
found, err := GetTaskExecutionByID(ctx, execution.ID)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, execution.TaskID, found.TaskID)
|
||||
}
|
||||
|
||||
func TestUpdateTaskExecution(t *testing.T) {
|
||||
cleanup := setupTaskExecutionTestEnvironment(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
// 创建记录
|
||||
execution := &model.TaskExecution{
|
||||
TaskID: "test_update_001",
|
||||
TaskType: "system:cleanup",
|
||||
TaskName: "清理未使用上传",
|
||||
Status: model.TaskExecutionStatusPending,
|
||||
TriggeredBy: "manual",
|
||||
}
|
||||
err := CreateTaskExecution(ctx, execution)
|
||||
require.NoError(t, err)
|
||||
|
||||
// 更新状态为 running
|
||||
now := time.Now()
|
||||
execution.Status = model.TaskExecutionStatusRunning
|
||||
execution.StartedAt = &now
|
||||
err = UpdateTaskExecution(ctx, execution)
|
||||
require.NoError(t, err)
|
||||
|
||||
// 验证更新
|
||||
found, err := GetTaskExecutionByTaskID(ctx, "test_update_001")
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, model.TaskExecutionStatusRunning, found.Status)
|
||||
assert.NotNil(t, found.StartedAt)
|
||||
|
||||
// 更新为 succeeded
|
||||
finishTime := time.Now()
|
||||
execution.Status = model.TaskExecutionStatusSucceeded
|
||||
execution.FinishedAt = &finishTime
|
||||
execution.Duration = 1500
|
||||
execution.Result = "共清理 50 个文件"
|
||||
err = UpdateTaskExecution(ctx, execution)
|
||||
require.NoError(t, err)
|
||||
|
||||
found, err = GetTaskExecutionByTaskID(ctx, "test_update_001")
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, model.TaskExecutionStatusSucceeded, found.Status)
|
||||
assert.Equal(t, int64(1500), found.Duration)
|
||||
assert.Equal(t, "共清理 50 个文件", found.Result)
|
||||
}
|
||||
|
||||
func TestUpdateTaskExecutionFailed(t *testing.T) {
|
||||
cleanup := setupTaskExecutionTestEnvironment(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
execution := &model.TaskExecution{
|
||||
TaskID: "test_fail_001",
|
||||
TaskType: "system:cleanup",
|
||||
TaskName: "清理未使用上传",
|
||||
Status: model.TaskExecutionStatusPending,
|
||||
Retryable: true,
|
||||
MaxRetry: 3,
|
||||
TriggeredBy: "manual",
|
||||
}
|
||||
err := CreateTaskExecution(ctx, execution)
|
||||
require.NoError(t, err)
|
||||
|
||||
// 标记为失败
|
||||
now := time.Now()
|
||||
execution.Status = model.TaskExecutionStatusFailed
|
||||
execution.StartedAt = &now
|
||||
execution.FinishedAt = &now
|
||||
execution.Duration = 200
|
||||
execution.ErrorMessage = "S3 连接超时"
|
||||
err = UpdateTaskExecution(ctx, execution)
|
||||
require.NoError(t, err)
|
||||
|
||||
found, err := GetTaskExecutionByTaskID(ctx, "test_fail_001")
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, model.TaskExecutionStatusFailed, found.Status)
|
||||
assert.Equal(t, "S3 连接超时", found.ErrorMessage)
|
||||
assert.Equal(t, int64(200), found.Duration)
|
||||
}
|
||||
|
||||
func TestUpdateTaskExecutionDoesNotPersistBufferedLog(t *testing.T) {
|
||||
cleanup := setupTaskExecutionTestEnvironment(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
execution := &model.TaskExecution{
|
||||
TaskID: "test_omit_log_001",
|
||||
TaskType: "system:cleanup",
|
||||
TaskName: "清理未使用上传",
|
||||
Status: model.TaskExecutionStatusPending,
|
||||
TriggeredBy: "manual",
|
||||
}
|
||||
err := CreateTaskExecution(ctx, execution)
|
||||
require.NoError(t, err)
|
||||
|
||||
// 运行中的日志仅缓存在 Redis。
|
||||
err = AppendTaskExecutionLog(ctx, "test_omit_log_001", "第一条执行日志")
|
||||
require.NoError(t, err)
|
||||
|
||||
assert.Empty(t, execution.Log)
|
||||
|
||||
execution.Status = model.TaskExecutionStatusSucceeded
|
||||
execution.Duration = 100
|
||||
err = UpdateTaskExecution(ctx, execution)
|
||||
require.NoError(t, err)
|
||||
|
||||
var persisted model.TaskExecution
|
||||
err = db.DB(ctx).Where("task_id = ?", "test_omit_log_001").First(&persisted).Error
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, model.TaskExecutionStatusSucceeded, persisted.Status)
|
||||
assert.Empty(t, persisted.Log)
|
||||
|
||||
found, err := GetTaskExecutionByTaskID(ctx, "test_omit_log_001")
|
||||
require.NoError(t, err)
|
||||
assert.Contains(t, found.Log, "第一条执行日志")
|
||||
}
|
||||
|
||||
func TestAppendTaskExecutionLog(t *testing.T) {
|
||||
cleanup := setupTaskExecutionTestEnvironment(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
execution := &model.TaskExecution{
|
||||
TaskID: "test_log_001",
|
||||
TaskType: "system:cleanup",
|
||||
TaskName: "清理未使用上传",
|
||||
Status: model.TaskExecutionStatusPending,
|
||||
TriggeredBy: "manual",
|
||||
}
|
||||
err := CreateTaskExecution(ctx, execution)
|
||||
require.NoError(t, err)
|
||||
|
||||
// 追加多条日志
|
||||
err = AppendTaskExecutionLog(ctx, "test_log_001", "开始扫描未使用上传文件")
|
||||
require.NoError(t, err)
|
||||
|
||||
err = AppendTaskExecutionLog(ctx, "test_log_001", "本批次找到 42 个待清理文件")
|
||||
require.NoError(t, err)
|
||||
|
||||
err = AppendTaskExecutionLog(ctx, "test_log_001", "清理完成,共删除 42 个文件")
|
||||
require.NoError(t, err)
|
||||
|
||||
// 读取时优先返回 Redis 中的在途日志。
|
||||
found, err := GetTaskExecutionByTaskID(ctx, "test_log_001")
|
||||
require.NoError(t, err)
|
||||
assert.Contains(t, found.Log, "开始扫描未使用上传文件")
|
||||
assert.Contains(t, found.Log, "本批次找到 42 个待清理文件")
|
||||
assert.Contains(t, found.Log, "清理完成,共删除 42 个文件")
|
||||
|
||||
var persisted model.TaskExecution
|
||||
err = db.DB(ctx).Where("task_id = ?", "test_log_001").First(&persisted).Error
|
||||
require.NoError(t, err)
|
||||
assert.Empty(t, persisted.Log)
|
||||
|
||||
err = FlushTaskExecutionLog(ctx, "test_log_001")
|
||||
require.NoError(t, err)
|
||||
|
||||
err = db.DB(ctx).Where("task_id = ?", "test_log_001").First(&persisted).Error
|
||||
require.NoError(t, err)
|
||||
assert.Contains(t, persisted.Log, "开始扫描未使用上传文件")
|
||||
|
||||
exists, err := db.Redis.Exists(ctx, taskExecutionLogRedisKey("test_log_001")).Result()
|
||||
require.NoError(t, err)
|
||||
assert.Zero(t, exists)
|
||||
}
|
||||
|
||||
func TestAppendTaskExecutionLogLimitsLinesAndRefreshesTTL(t *testing.T) {
|
||||
cleanup := setupTaskExecutionTestEnvironment(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
const taskID = "limited_log_001"
|
||||
for i := 0; i < taskExecutionLogMaxLines+5; i++ {
|
||||
err := AppendTaskExecutionLog(ctx, taskID, fmt.Sprintf("日志-%04d", i))
|
||||
require.NoError(t, err)
|
||||
}
|
||||
|
||||
key := taskExecutionLogRedisKey(taskID)
|
||||
logLines, err := db.Redis.LRange(ctx, key, 0, -1).Result()
|
||||
require.NoError(t, err)
|
||||
assert.Len(t, logLines, taskExecutionLogMaxLines)
|
||||
assert.Contains(t, logLines[0], "日志-0005")
|
||||
assert.Contains(t, logLines[len(logLines)-1], "日志-1004")
|
||||
|
||||
ttl, err := db.Redis.TTL(ctx, key).Result()
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, taskExecutionLogExpiration, ttl)
|
||||
}
|
||||
|
||||
func TestAppendTaskExecutionLogNonExistent(t *testing.T) {
|
||||
cleanup := setupTaskExecutionTestEnvironment(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
// Redis 缓冲不依赖数据库记录是否已经创建。
|
||||
err := AppendTaskExecutionLog(ctx, "nonexistent_task", "测试日志")
|
||||
assert.NoError(t, err)
|
||||
|
||||
err = FlushTaskExecutionLog(ctx, "nonexistent_task")
|
||||
assert.Error(t, err)
|
||||
}
|
||||
|
||||
func TestGetTaskExecutionLogPrefersRedis(t *testing.T) {
|
||||
cleanup := setupTaskExecutionTestEnvironment(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
execution := &model.TaskExecution{
|
||||
TaskID: "redis_priority_001",
|
||||
TaskType: "system:cleanup",
|
||||
TaskName: "清理未使用上传",
|
||||
Status: model.TaskExecutionStatusRunning,
|
||||
Log: "数据库旧日志",
|
||||
TriggeredBy: "manual",
|
||||
}
|
||||
err := CreateTaskExecution(ctx, execution)
|
||||
require.NoError(t, err)
|
||||
err = AppendTaskExecutionLog(ctx, execution.TaskID, "Redis 最新日志")
|
||||
require.NoError(t, err)
|
||||
|
||||
found, err := GetTaskExecutionByID(ctx, execution.ID)
|
||||
require.NoError(t, err)
|
||||
assert.Contains(t, found.Log, "Redis 最新日志")
|
||||
assert.NotContains(t, found.Log, "数据库旧日志")
|
||||
}
|
||||
|
||||
func TestListTaskExecutions(t *testing.T) {
|
||||
cleanup := setupTaskExecutionTestEnvironment(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
// 创建多条记录,包含不同状态和类型
|
||||
records := []*model.TaskExecution{
|
||||
{TaskID: "list_001", TaskType: "system:cleanup", TaskName: "系统垃圾清理", Status: model.TaskExecutionStatusSucceeded, TriggeredBy: "manual"},
|
||||
{TaskID: "list_002", TaskType: "system:cleanup", TaskName: "系统垃圾清理", Status: model.TaskExecutionStatusFailed, TriggeredBy: "system"},
|
||||
{TaskID: "list_003", TaskType: "other:task", TaskName: "其他任务", Status: model.TaskExecutionStatusPending, TriggeredBy: "manual"},
|
||||
{TaskID: "list_004", TaskType: "system:cleanup", TaskName: "系统垃圾清理", Status: model.TaskExecutionStatusRunning, TriggeredBy: "manual"},
|
||||
{TaskID: "list_005", TaskType: "other:task", TaskName: "其他任务", Status: model.TaskExecutionStatusSucceeded, TriggeredBy: "system"},
|
||||
}
|
||||
for _, r := range records {
|
||||
err := CreateTaskExecution(ctx, r)
|
||||
require.NoError(t, err)
|
||||
}
|
||||
err := AppendTaskExecutionLog(ctx, "list_004", "运行中的 Redis 日志")
|
||||
require.NoError(t, err)
|
||||
|
||||
// 查询全部(分页)
|
||||
items, total, err := ListTaskExecutions(ctx, model.ListTaskExecutionsRequest{Page: 1, PageSize: 10})
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, int64(5), total)
|
||||
assert.Len(t, items, 5)
|
||||
for _, item := range items {
|
||||
if item.TaskID == "list_004" {
|
||||
assert.Contains(t, item.Log, "运行中的 Redis 日志")
|
||||
}
|
||||
}
|
||||
|
||||
// 按状态筛选:failed
|
||||
items, total, err = ListTaskExecutions(ctx, model.ListTaskExecutionsRequest{Status: "failed", Page: 1, PageSize: 10})
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, int64(1), total)
|
||||
assert.Len(t, items, 1)
|
||||
assert.Equal(t, "list_002", items[0].TaskID)
|
||||
|
||||
// 按类型筛选
|
||||
_, total, err = ListTaskExecutions(ctx, model.ListTaskExecutionsRequest{TaskType: "other:task", Page: 1, PageSize: 10})
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, int64(2), total)
|
||||
|
||||
// 分页测试
|
||||
items, total, err = ListTaskExecutions(ctx, model.ListTaskExecutionsRequest{Page: 1, PageSize: 2})
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, int64(5), total)
|
||||
assert.Len(t, items, 2)
|
||||
|
||||
items2, total2, err := ListTaskExecutions(ctx, model.ListTaskExecutionsRequest{Page: 2, PageSize: 2})
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, int64(5), total2)
|
||||
assert.Len(t, items2, 2)
|
||||
|
||||
// 确保分页数据不重复
|
||||
assert.NotEqual(t, items[0].ID, items2[0].ID)
|
||||
|
||||
// 状态 + 类型组合筛选
|
||||
items, total, err = ListTaskExecutions(ctx, model.ListTaskExecutionsRequest{Status: "succeeded", TaskType: "system:cleanup", Page: 1, PageSize: 10})
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, int64(1), total)
|
||||
assert.Equal(t, "list_001", items[0].TaskID)
|
||||
|
||||
// 按类型前缀筛选
|
||||
items, total, err = ListTaskExecutions(ctx, model.ListTaskExecutionsRequest{TaskTypePrefix: "system:", Page: 1, PageSize: 10})
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, int64(3), total)
|
||||
assert.Len(t, items, 3)
|
||||
|
||||
// 按多类型 IN 筛选
|
||||
items, total, err = ListTaskExecutions(ctx, model.ListTaskExecutionsRequest{
|
||||
TaskTypes: "system:cleanup,other:task",
|
||||
Page: 1,
|
||||
PageSize: 10,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, int64(5), total)
|
||||
assert.Len(t, items, 5)
|
||||
|
||||
// 精确类型优先于 task_types / 前缀
|
||||
items, total, err = ListTaskExecutions(ctx, model.ListTaskExecutionsRequest{
|
||||
TaskType: "other:task",
|
||||
TaskTypes: "system:cleanup",
|
||||
TaskTypePrefix: "system:",
|
||||
Page: 1,
|
||||
PageSize: 10,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, int64(2), total)
|
||||
assert.Len(t, items, 2)
|
||||
}
|
||||
|
||||
func TestListTaskExecutionsDefaultPaging(t *testing.T) {
|
||||
cleanup := setupTaskExecutionTestEnvironment(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
// 不传分页参数,应使用默认值 page=1, pageSize=20
|
||||
items, total, err := ListTaskExecutions(ctx, model.ListTaskExecutionsRequest{})
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, int64(0), total)
|
||||
assert.Len(t, items, 0)
|
||||
}
|
||||
|
||||
func TestCleanupTaskExecutionLogs(t *testing.T) {
|
||||
cleanup := setupTaskExecutionTestEnvironment(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
now := time.Date(2026, 6, 17, 12, 0, 0, 0, time.UTC)
|
||||
for i := 0; i < 31; i++ {
|
||||
createTaskExecutionForCleanup(t, ctx, fmt.Sprintf("high_recent_%02d", i), "high:task", model.TaskExecutionStatusSucceeded, now.Add(-2*time.Hour))
|
||||
}
|
||||
createTaskExecutionForCleanup(t, ctx, "high_old_4d", "high:task", model.TaskExecutionStatusSucceeded, now.AddDate(0, 0, -4))
|
||||
createTaskExecutionForCleanup(t, ctx, "high_old_40d", "high:task", model.TaskExecutionStatusFailed, now.AddDate(0, 0, -40))
|
||||
createTaskExecutionForCleanup(t, ctx, "high_running_old", "high:task", model.TaskExecutionStatusRunning, now.AddDate(0, 0, -10))
|
||||
createTaskExecutionForCleanup(t, ctx, "low_old_31d", "low:task", model.TaskExecutionStatusSucceeded, now.AddDate(0, 0, -31))
|
||||
createTaskExecutionForCleanup(t, ctx, "low_recent_29d", "low:task", model.TaskExecutionStatusSucceeded, now.AddDate(0, 0, -29))
|
||||
createTaskExecutionForCleanup(t, ctx, "low_pending_old", "low:task", model.TaskExecutionStatusPending, now.AddDate(0, 0, -45))
|
||||
|
||||
stats, err := CleanupTaskExecutionLogs(ctx, now)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, int64(2), stats.HighFrequencyDeleted)
|
||||
assert.Equal(t, int64(1), stats.LowFrequencyDeleted)
|
||||
|
||||
for _, taskID := range []string{"high_old_4d", "high_old_40d", "low_old_31d"} {
|
||||
var count int64
|
||||
err := db.DB(ctx).Model(&model.TaskExecution{}).Where("task_id = ?", taskID).Count(&count).Error
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, int64(0), count, "CleanupTaskExecutionLogs(%s) should delete expired log", taskID)
|
||||
}
|
||||
for _, taskID := range []string{"high_recent_00", "high_running_old", "low_recent_29d", "low_pending_old"} {
|
||||
var count int64
|
||||
err := db.DB(ctx).Model(&model.TaskExecution{}).Where("task_id = ?", taskID).Count(&count).Error
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, int64(1), count, "CleanupTaskExecutionLogs(%s) should keep retained log", taskID)
|
||||
}
|
||||
}
|
||||
|
||||
func TestTaskExecutionTableName(t *testing.T) {
|
||||
execution := model.TaskExecution{}
|
||||
assert.Equal(t, "w_task_executions", execution.TableName())
|
||||
}
|
||||
|
||||
func createTaskExecutionForCleanup(t *testing.T, ctx context.Context, taskID string, taskType string, status model.TaskExecutionStatus, createdAt time.Time) {
|
||||
t.Helper()
|
||||
|
||||
execution := &model.TaskExecution{
|
||||
TaskID: taskID,
|
||||
TaskType: taskType,
|
||||
TaskName: taskType,
|
||||
Status: status,
|
||||
CreatedAt: createdAt,
|
||||
UpdatedAt: createdAt,
|
||||
TriggeredBy: "system",
|
||||
}
|
||||
err := CreateTaskExecution(ctx, execution)
|
||||
require.NoError(t, err)
|
||||
}
|
||||
+114
-1
@@ -5,8 +5,11 @@ package repository
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/infra/persistence"
|
||||
db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
|
||||
"github.com/Rain-kl/Wavelet/internal/infra/persistence/idgen"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
@@ -175,7 +178,117 @@ func ListUsersByIDs(ctx context.Context, ids []uint64) ([]model.User, error) {
|
||||
return users, nil
|
||||
}
|
||||
|
||||
// ListUserIDsByUsernameContains returns user IDs whose username contains the given fragment.
|
||||
func ListUserIDsByUsernameContains(ctx context.Context, username string) ([]uint64, error) {
|
||||
if username == "" {
|
||||
return []uint64{}, nil
|
||||
}
|
||||
var userIDs []uint64
|
||||
if err := db.DB(ctx).Model(&model.User{}).
|
||||
Where("username LIKE ?", "%"+username+"%").
|
||||
Pluck("id", &userIDs).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return userIDs, nil
|
||||
}
|
||||
|
||||
// UpdateUser updates all fields of an existing user.
|
||||
func UpdateUser(ctx context.Context, user *model.User) error {
|
||||
return db.DB(ctx).Save(user).Error
|
||||
}
|
||||
|
||||
// CreateUserFromOAuth creates a user from OAuth profile data and fills userOut.
|
||||
func CreateUserFromOAuth(ctx context.Context, userOut *model.User, oauthInfo *model.OAuthUserInfo) error {
|
||||
now := time.Now()
|
||||
userID := oauthInfo.GetID()
|
||||
newUser := model.User{
|
||||
ID: userID,
|
||||
Username: oauthInfo.Username,
|
||||
Nickname: oauthInfo.Name,
|
||||
Email: oauthInfo.Email,
|
||||
AvatarURL: oauthInfo.AvatarURL,
|
||||
IsActive: oauthInfo.Active,
|
||||
LastLoginAt: now,
|
||||
IsAdmin: false,
|
||||
}
|
||||
if newUser.ID == 0 {
|
||||
newUser.ID = idgen.NextUint64ID()
|
||||
}
|
||||
if err := db.DB(ctx).Create(&newUser).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
*userOut = newUser
|
||||
return nil
|
||||
}
|
||||
|
||||
// ListUsernamesMatchingBase returns usernames equal to base or prefixed with base+"-".
|
||||
func ListUsernamesMatchingBase(ctx context.Context, base string) ([]string, error) {
|
||||
var names []string
|
||||
if err := db.DB(ctx).Model(&model.User{}).
|
||||
Where("username = ? OR username LIKE ?", base, base+"-%").
|
||||
Pluck("username", &names).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return names, nil
|
||||
}
|
||||
|
||||
// GetActiveUserByID loads a user by ID who is active.
|
||||
func GetActiveUserByID(ctx context.Context, id uint64) (model.User, error) {
|
||||
var user model.User
|
||||
if err := db.DB(ctx).Where("id = ? AND is_active = ?", id, true).First(&user).Error; err != nil {
|
||||
return model.User{}, err
|
||||
}
|
||||
return user, nil
|
||||
}
|
||||
|
||||
// GetUserByUsernameOrEmail loads a user by username or email.
|
||||
func GetUserByUsernameOrEmail(ctx context.Context, input string) (model.User, error) {
|
||||
var user model.User
|
||||
if err := db.DB(ctx).Where("username = ? OR email = ?", input, input).First(&user).Error; err != nil {
|
||||
return model.User{}, err
|
||||
}
|
||||
return user, nil
|
||||
}
|
||||
|
||||
// CountUsersByEmailExceptID counts users with the email excluding a given user id.
|
||||
func CountUsersByEmailExceptID(ctx context.Context, email string, exceptID uint64) (int64, error) {
|
||||
var count int64
|
||||
if err := db.DB(ctx).Model(&model.User{}).Where("email = ? AND id != ?", email, exceptID).Count(&count).Error; err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return count, nil
|
||||
}
|
||||
|
||||
// UpdateUserLastLoginAt updates only last_login_at for a user.
|
||||
func UpdateUserLastLoginAt(ctx context.Context, userID uint64, at time.Time) error {
|
||||
return db.DB(ctx).Model(&model.User{}).Where("id = ?", userID).Update("last_login_at", at).Error
|
||||
}
|
||||
|
||||
// UpdateUserPassword updates only the password hash for a user.
|
||||
func UpdateUserPassword(ctx context.Context, userID uint64, passwordHash string) error {
|
||||
return db.DB(ctx).Model(&model.User{}).Where("id = ?", userID).Update("password", passwordHash).Error
|
||||
}
|
||||
|
||||
// RegisterUserWithChecks validates username/email uniqueness then creates the user.
|
||||
func RegisterUserWithChecks(ctx context.Context, user *model.User) error {
|
||||
var count int64
|
||||
if err := db.DB(ctx).Model(&model.User{}).Where("username = ?", user.Username).Count(&count).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if count > 0 {
|
||||
return errors.New("用户名已存在")
|
||||
}
|
||||
if user.Email != "" {
|
||||
var emailCount int64
|
||||
if err := db.DB(ctx).Model(&model.User{}).Where("email = ?", user.Email).Count(&emailCount).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if emailCount > 0 {
|
||||
return errors.New("该邮箱已被其他账号绑定")
|
||||
}
|
||||
}
|
||||
if user.ID == 0 {
|
||||
user.ID = idgen.NextUint64ID()
|
||||
}
|
||||
return db.DB(ctx).Create(user).Error
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user