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:
ryan
2026-08-16 11:07:20 +08:00
parent b9b42e3174
commit 6a53619dd2
58 changed files with 2201 additions and 1175 deletions
+69
View File
@@ -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{} }
+215
View File
@@ -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(&current, "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(&current).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(&current).Error
if err == nil {
if current.UserID != account.UserID {
return errors.New(errExternalAccountAlreadyBoundToAnother)
}
return tx.Model(&current).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
}
+3 -3
View File
@@ -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)
}
+25
View File
@@ -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"
+53
View File
@@ -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
}
+1 -8
View File
@@ -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.
+304
View File
@@ -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
}
+514
View File
@@ -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
View File
@@ -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
}