mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-03 23:06:36 +08:00
wavelet init
This commit is contained in:
@@ -0,0 +1,60 @@
|
||||
// Copyright 2025 linux.do
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package model 定义数据模型与 GORM 实体
|
||||
package model
|
||||
|
||||
import (
|
||||
"crypto/rand"
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"fmt"
|
||||
"time"
|
||||
)
|
||||
|
||||
const (
|
||||
tokenByteLength = 24 // Token 随机字节长度
|
||||
maskThreshold = 8 // 脱敏显示阈值
|
||||
)
|
||||
|
||||
// AccessToken 个人访问令牌实体
|
||||
type AccessToken struct {
|
||||
ID uint64 `json:"id" gorm:"primaryKey;autoIncrement"`
|
||||
UserID uint64 `json:"user_id" gorm:"index;not null"`
|
||||
Name string `json:"name" gorm:"size:128;not null"`
|
||||
TokenHash string `json:"-" gorm:"size:64;uniqueIndex;not null"`
|
||||
MaskedToken string `json:"masked_token" gorm:"size:64;not null"`
|
||||
IsAdmin bool `json:"is_admin" gorm:"default:false"`
|
||||
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime"`
|
||||
UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime"`
|
||||
}
|
||||
|
||||
// TableName 表名
|
||||
func (AccessToken) TableName() string {
|
||||
return "w_access_tokens"
|
||||
}
|
||||
|
||||
// GenerateTokenString 生成加密安全的随机 Token 值
|
||||
func GenerateTokenString() (string, error) {
|
||||
bytes := make([]byte, tokenByteLength)
|
||||
if _, err := rand.Read(bytes); err != nil {
|
||||
return "", err
|
||||
}
|
||||
return fmt.Sprintf("at_%s", hex.EncodeToString(bytes)), nil
|
||||
}
|
||||
|
||||
// HashToken 计算 Token 的 SHA-256 哈希值用于数据库存储与查询
|
||||
func HashToken(token string) string {
|
||||
h := sha256.New()
|
||||
h.Write([]byte(token))
|
||||
return hex.EncodeToString(h.Sum(nil))
|
||||
}
|
||||
|
||||
// MaskTokenString 生成脱敏显示的 Token,仅保留前缀和最后四位
|
||||
func MaskTokenString(token string) string {
|
||||
if len(token) <= maskThreshold {
|
||||
return "at_****"
|
||||
}
|
||||
return fmt.Sprintf("%s...%s", token[:7], token[len(token)-4:])
|
||||
}
|
||||
@@ -0,0 +1,318 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package model
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"regexp"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
// 认证源类型
|
||||
const (
|
||||
AuthSourceTypeOIDC = "oidc"
|
||||
)
|
||||
|
||||
var authSourceNamePattern = regexp.MustCompile(`^[A-Za-z0-9][A-Za-z0-9_-]{0,79}$`)
|
||||
|
||||
// AuthSource 认证源实体
|
||||
type AuthSource struct {
|
||||
ID uint64 `json:"id" gorm:"primaryKey"`
|
||||
Name string `json:"name" gorm:"uniqueIndex;size:80;not null"`
|
||||
Type string `json:"type" gorm:"size:20;not null"`
|
||||
DisplayName string `json:"display_name" gorm:"size:100"`
|
||||
IsActive bool `json:"is_active" gorm:"index;not null;default:false"`
|
||||
ClientID string `json:"client_id" gorm:"size:255"`
|
||||
ClientSecret string `json:"-" gorm:"size:1024"`
|
||||
OpenIDDiscoveryURL string `json:"openid_discovery_url" gorm:"column:openid_discovery_url;size:1024"`
|
||||
Scopes string `json:"scopes" gorm:"size:255"`
|
||||
IconURL string `json:"icon_url" gorm:"size:1024"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
UpdatedAt time.Time `json:"updated_at"`
|
||||
ClientSecretConfigured bool `json:"client_secret_configured" gorm:"-"`
|
||||
}
|
||||
|
||||
// TableName 表名
|
||||
func (AuthSource) TableName() string {
|
||||
return "w_auth_sources"
|
||||
}
|
||||
|
||||
// ExternalAccount 外部账号绑定实体
|
||||
type ExternalAccount struct {
|
||||
ID uint64 `json:"id" gorm:"primaryKey"`
|
||||
AuthSourceID uint64 `json:"auth_source_id" gorm:"uniqueIndex:idx_external_accounts_source_external,priority:1;index"`
|
||||
UserID uint64 `json:"user_id" gorm:"index;not null"`
|
||||
ExternalID string `json:"external_id" gorm:"uniqueIndex:idx_external_accounts_source_external,priority:2;size:255;not null"`
|
||||
ExternalUsername string `json:"external_username" gorm:"size:255"`
|
||||
Email string `json:"email" gorm:"size:255"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
UpdatedAt time.Time `json:"updated_at"`
|
||||
}
|
||||
|
||||
// TableName 表名
|
||||
func (ExternalAccount) TableName() string {
|
||||
return "w_external_accounts"
|
||||
}
|
||||
|
||||
// ExternalAccountView 外部帐号绑定视图(脱敏展示用)
|
||||
type ExternalAccountView struct {
|
||||
ID uint64 `json:"id"`
|
||||
AuthSourceID uint64 `json:"auth_source_id"`
|
||||
AuthSourceName string `json:"auth_source_name"`
|
||||
AuthSourceType string `json:"auth_source_type"`
|
||||
AuthSourceLabel string `json:"auth_source_label"`
|
||||
ExternalUsername string `json:"external_username"`
|
||||
Email string `json:"email"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
}
|
||||
|
||||
// Normalize 对认证源字段进行标准化处理
|
||||
func (source *AuthSource) Normalize() {
|
||||
source.Type = strings.ToLower(strings.TrimSpace(source.Type))
|
||||
source.Name = strings.TrimSpace(source.Name)
|
||||
source.DisplayName = strings.TrimSpace(source.DisplayName)
|
||||
source.ClientID = strings.TrimSpace(source.ClientID)
|
||||
source.ClientSecret = strings.TrimSpace(source.ClientSecret)
|
||||
source.OpenIDDiscoveryURL = strings.TrimSpace(source.OpenIDDiscoveryURL)
|
||||
source.Scopes = strings.TrimSpace(source.Scopes)
|
||||
source.IconURL = strings.TrimSpace(source.IconURL)
|
||||
if source.DisplayName == "" {
|
||||
source.DisplayName = source.Name
|
||||
}
|
||||
if source.Type == AuthSourceTypeOIDC && source.Scopes == "" {
|
||||
source.Scopes = "openid profile email"
|
||||
}
|
||||
}
|
||||
|
||||
// Validate 校验认证源字段合法性
|
||||
func (source *AuthSource) Validate() error {
|
||||
source.Normalize()
|
||||
if source.Name == "" {
|
||||
return errors.New(errAuthSourceNameRequired)
|
||||
}
|
||||
if !authSourceNamePattern.MatchString(source.Name) {
|
||||
return errors.New(errAuthSourceNameInvalid)
|
||||
}
|
||||
if source.Type != AuthSourceTypeOIDC {
|
||||
return errors.New(errAuthSourceTypeUnsupported)
|
||||
}
|
||||
if source.OpenIDDiscoveryURL == "" {
|
||||
return errors.New(errAuthSourceDiscoveryURLRequired)
|
||||
}
|
||||
if source.IsActive && (source.ClientID == "" || source.ClientSecret == "") {
|
||||
return errors.New(errAuthSourceClientCredentialsRequired)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// Sanitize 脱敏处理,将 ClientSecret 清空并设置 ClientSecretConfigured 标志
|
||||
func (source *AuthSource) Sanitize() {
|
||||
source.ClientSecretConfigured = source.ClientSecret != ""
|
||||
source.ClientSecret = ""
|
||||
}
|
||||
|
||||
// GetAuthSources 获取所有认证源(已脱敏)
|
||||
func GetAuthSources(ctx context.Context) ([]AuthSource, error) {
|
||||
var sources []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) ([]AuthSource, error) {
|
||||
var sources []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) (*AuthSource, error) {
|
||||
if id == 0 {
|
||||
return nil, errors.New(errAuthSourceIDRequired)
|
||||
}
|
||||
var source 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) (*AuthSource, error) {
|
||||
name = strings.TrimSpace(name)
|
||||
if name == "" {
|
||||
return nil, errors.New(errAuthSourceNameRequired)
|
||||
}
|
||||
var source 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 *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 *AuthSource, keepSecret bool) error {
|
||||
if source.ID == 0 {
|
||||
return errors.New(errAuthSourceIDRequired)
|
||||
}
|
||||
var current 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{
|
||||
"name": 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(&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(&ExternalAccount{}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
return tx.Delete(&AuthSource{}, "id = ?", id).Error
|
||||
})
|
||||
}
|
||||
|
||||
// FindExternalAccount 查找外部帐号绑定记录
|
||||
func FindExternalAccount(ctx context.Context, sourceID uint64, externalID string) (*ExternalAccount, error) {
|
||||
var account 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 *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 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) ([]ExternalAccountView, error) {
|
||||
if userID == 0 {
|
||||
return nil, errors.New(errUserIDRequired)
|
||||
}
|
||||
var accounts []ExternalAccount
|
||||
if err := db.DB(ctx).Where("user_id = ?", userID).Order("id asc").Find(&accounts).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
views := make([]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, 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(&ExternalAccount{}).Error
|
||||
}
|
||||
@@ -0,0 +1,30 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package model
|
||||
|
||||
const (
|
||||
errRegistrationDisabled = "注册已关闭"
|
||||
errDatabaseNotInitialized = "database not initialized"
|
||||
errUsernameExists = "用户名已存在"
|
||||
errEmailAlreadyBound = "该邮箱已被其他账号绑定"
|
||||
errConfigIntParseFailed = "配置 %s 的值 '%s' 无法转换为整数: %w"
|
||||
errConfigDecimalParseFailed = "配置 %s 的值 '%s' 无法转换为decimal: %w"
|
||||
errConfigBoolParseFailed = "配置 %s 的值 '%s' 无法转换为布尔值: %w"
|
||||
errParseMenuDisplayConfigFailed = "解析目录显示配置失败: %w"
|
||||
errTemplateKeyRequired = "模板标识符不能为空"
|
||||
errTemplateNameRequired = "模板名称不能为空"
|
||||
errTemplateContentRequired = "模板内容不能为空"
|
||||
errTemplateUnavailable = "模板 %s 不存在或不可用: %w"
|
||||
errTemplateRenderFailed = "模板 %s 渲染失败: %w"
|
||||
errAuthSourceNameRequired = "认证源名称不能为空"
|
||||
errAuthSourceNameInvalid = "认证源名称只能包含字母、数字、短横线或下划线,且必须以字母或数字开头"
|
||||
errAuthSourceTypeUnsupported = "认证源类型仅支持 oidc"
|
||||
errAuthSourceDiscoveryURLRequired = "OIDC 认证源必须配置 Discovery URL"
|
||||
errAuthSourceClientCredentialsRequired = "启用认证源前必须配置 Client ID 和 Client Secret" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
|
||||
errAuthSourceIDRequired = "认证源 ID 不能为空"
|
||||
errExternalAccountBindingIncomplete = "外部账号绑定信息不完整"
|
||||
errExternalAccountAlreadyBoundToAnother = "该外部账号已绑定到其他用户"
|
||||
errUserIDRequired = "用户 ID 不能为空"
|
||||
errExternalAccountBindingIDRequired = "绑定记录 ID 不能为空"
|
||||
)
|
||||
@@ -0,0 +1,99 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package model
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"regexp"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
const (
|
||||
// TypeCustom 自定义消息通道类型
|
||||
TypeCustom = "custom"
|
||||
// TypeEmail 邮件推送消息通道类型
|
||||
TypeEmail = "email"
|
||||
// TypeTelegram 电报机器人推送消息通道类型
|
||||
TypeTelegram = "telegram"
|
||||
)
|
||||
|
||||
// PushChannel 消息通道模型
|
||||
type PushChannel struct {
|
||||
ID uint64 `json:"id" gorm:"primaryKey;autoIncrement"`
|
||||
Name string `json:"name" gorm:"uniqueIndex;size:80;not null"` // 通道名称,仅英文字母和下划线,唯一
|
||||
Description string `json:"description" gorm:"size:255"` // 备注
|
||||
Type string `json:"type" gorm:"size:50;not null;default:'custom'"` // 通道类型:custom, lark, email
|
||||
Token string `json:"token" gorm:"size:100"` // 鉴权令牌或发信用户名等
|
||||
URL string `json:"url" gorm:"type:text;not null"` // 请求地址,HTTPS 协议或 SMTP 地址
|
||||
Other string `json:"other" gorm:"type:text;not null"` // 请求体/SMTP 密码等
|
||||
Enabled bool `json:"enabled" gorm:"index;not null;default:true"` // 通道是否启用
|
||||
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime;index"`
|
||||
UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime;index"`
|
||||
}
|
||||
|
||||
// TableName 指定 GORM 表名
|
||||
func (PushChannel) TableName() string {
|
||||
return "w_push_channels"
|
||||
}
|
||||
|
||||
var nameRegex = regexp.MustCompile(`^[a-zA-Z0-9_]+$`)
|
||||
|
||||
// Validate 参数合法性与 JSON 格式校验
|
||||
func (pc *PushChannel) Validate() error {
|
||||
pc.Name = strings.TrimSpace(pc.Name)
|
||||
pc.URL = strings.TrimSpace(pc.URL)
|
||||
pc.Other = strings.TrimSpace(pc.Other)
|
||||
pc.Type = strings.TrimSpace(pc.Type)
|
||||
|
||||
if pc.Type == "" {
|
||||
pc.Type = TypeCustom
|
||||
}
|
||||
|
||||
if pc.Type == TypeTelegram && pc.URL == "" {
|
||||
pc.URL = "https://api.telegram.org"
|
||||
}
|
||||
|
||||
if pc.Name == "" {
|
||||
return errors.New("channel name is required")
|
||||
}
|
||||
if !nameRegex.MatchString(pc.Name) {
|
||||
return errors.New("channel name can only contain letters, numbers, and underscores")
|
||||
}
|
||||
if pc.Type != TypeEmail && pc.URL == "" {
|
||||
return errors.New("request URL/address is required")
|
||||
}
|
||||
|
||||
if pc.Type != TypeEmail && !strings.HasPrefix(pc.URL, "https://") {
|
||||
return errors.New("request URL must use HTTPS protocol for security reasons")
|
||||
}
|
||||
|
||||
switch pc.Type {
|
||||
case TypeCustom:
|
||||
if pc.Other == "" {
|
||||
return errors.New("payload schema (request body) is required")
|
||||
}
|
||||
return validateJSON(pc.Other)
|
||||
case TypeEmail:
|
||||
// Email channel SMTP configs fall back to global settings, so they are not required to be filled.
|
||||
case TypeTelegram:
|
||||
if pc.Token == "" {
|
||||
return errors.New("telegram bot token is required")
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func validateJSON(s string) error {
|
||||
var jsonTest map[string]any
|
||||
if err := json.Unmarshal([]byte(s), &jsonTest); err == nil {
|
||||
return nil
|
||||
}
|
||||
var jsonArr []any
|
||||
if err := json.Unmarshal([]byte(s), &jsonArr); err == nil {
|
||||
return nil
|
||||
}
|
||||
return errors.New("payload schema must be a valid JSON format")
|
||||
}
|
||||
@@ -0,0 +1,50 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package model
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
// PushEvent 系统通知事件模型
|
||||
type PushEvent struct {
|
||||
ID uint64 `json:"id" gorm:"primaryKey;autoIncrement"`
|
||||
EventKey string `json:"event_key" gorm:"uniqueIndex;size:80;not null"` // 如 admin_login
|
||||
Name string `json:"name" gorm:"size:100;not null"` // 如 管理员登录
|
||||
TaskType string `json:"task_type" gorm:"size:100;index;not null;default:''"` // 关联的异步任务类型
|
||||
Channels []string `json:"channels" gorm:"type:text;serializer:json;not null"` // 推送渠道列表,如 ["lark"]
|
||||
Targets []string `json:"targets" gorm:"type:text;serializer:json;not null"` // 推送目标用户/邮箱列表
|
||||
Template string `json:"template" gorm:"type:text;not null"` // 消息模板 JSON
|
||||
Enabled bool `json:"enabled" gorm:"index;not null;default:false"` // 是否启用
|
||||
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime;index"`
|
||||
UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime;index"`
|
||||
}
|
||||
|
||||
// TableName 指定 GORM 表名
|
||||
func (PushEvent) TableName() string {
|
||||
return "w_push_events"
|
||||
}
|
||||
|
||||
// Validate 基础校验
|
||||
func (pe *PushEvent) Validate() error {
|
||||
pe.EventKey = strings.TrimSpace(pe.EventKey)
|
||||
pe.Name = strings.TrimSpace(pe.Name)
|
||||
pe.Template = strings.TrimSpace(pe.Template)
|
||||
|
||||
if pe.EventKey == "" {
|
||||
return errors.New("event key is required")
|
||||
}
|
||||
if pe.Name == "" {
|
||||
return errors.New("event name is required")
|
||||
}
|
||||
if pe.Template == "" {
|
||||
return errors.New("event template is required")
|
||||
}
|
||||
if pe.Enabled && len(pe.Channels) == 0 {
|
||||
return errors.New("cannot enable event without any push channels configured")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,27 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package model
|
||||
|
||||
import (
|
||||
"time"
|
||||
)
|
||||
|
||||
// PushHistory 推送日志/历史实体
|
||||
type PushHistory struct {
|
||||
ID uint64 `json:"id" gorm:"primaryKey;autoIncrement"`
|
||||
EventKey string `json:"event_key" gorm:"size:80;not null;index"`
|
||||
Channel string `json:"channel" gorm:"size:50;not null"`
|
||||
Target string `json:"target" gorm:"size:255;not null"`
|
||||
Title string `json:"title" gorm:"size:255;not null"`
|
||||
Content string `json:"content" gorm:"type:text;not null"`
|
||||
Level string `json:"level" gorm:"size:20;not null"`
|
||||
Status string `json:"status" gorm:"size:20;not null"` // success / failed
|
||||
ErrorMsg string `json:"error_msg" gorm:"type:text"`
|
||||
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime;index"`
|
||||
}
|
||||
|
||||
// TableName 指定表名
|
||||
func (PushHistory) TableName() string {
|
||||
return "w_push_histories"
|
||||
}
|
||||
@@ -0,0 +1,70 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package model
|
||||
|
||||
import (
|
||||
"context"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
)
|
||||
|
||||
// Schedule 定时任务配置表
|
||||
type Schedule struct {
|
||||
ID uint64 `json:"id,string" gorm:"primaryKey"`
|
||||
Name string `json:"name" gorm:"size:128;not null"`
|
||||
TaskType string `json:"task_type" gorm:"size:64;not null"`
|
||||
Cron string `json:"cron" gorm:"size:64;not null"`
|
||||
Payload string `json:"payload" gorm:"type:text"`
|
||||
IsActive bool `json:"is_active" gorm:"not null;default:true"`
|
||||
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime;index"`
|
||||
UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime"`
|
||||
}
|
||||
|
||||
// TableName 表名
|
||||
func (Schedule) TableName() string {
|
||||
return "w_schedules"
|
||||
}
|
||||
|
||||
// CreateSchedule 创建定时任务
|
||||
func CreateSchedule(ctx context.Context, schedule *Schedule) error {
|
||||
return db.DB(ctx).Create(schedule).Error
|
||||
}
|
||||
|
||||
// UpdateSchedule 更新定时任务
|
||||
func UpdateSchedule(ctx context.Context, schedule *Schedule) error {
|
||||
return db.DB(ctx).Save(schedule).Error
|
||||
}
|
||||
|
||||
// DeleteSchedule 删除定时任务
|
||||
func DeleteSchedule(ctx context.Context, id uint64) error {
|
||||
return db.DB(ctx).Delete(&Schedule{}, id).Error
|
||||
}
|
||||
|
||||
// GetScheduleByID 根据 ID 获取定时任务
|
||||
func GetScheduleByID(ctx context.Context, id uint64) (*Schedule, error) {
|
||||
var schedule 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) ([]Schedule, error) {
|
||||
var schedules []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) ([]Schedule, error) {
|
||||
var schedules []Schedule
|
||||
if err := db.DB(ctx).Where("is_active = ?", true).Find(&schedules).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return schedules, nil
|
||||
}
|
||||
@@ -0,0 +1,63 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package model
|
||||
|
||||
import "time"
|
||||
|
||||
// 配置键常量 - 所有系统配置的 key 定义
|
||||
const (
|
||||
ConfigKeyUploadAllowedExtensions = "upload_allowed_extensions" // 允许上传的文件扩展名,逗号分隔
|
||||
ConfigKeySiteName = "site_name" // 站点名称
|
||||
ConfigKeyPasswordLoginEnabled = "password_login_enabled" // 是否允许密码登录
|
||||
ConfigKeyRegistrationEnabled = "registration_enabled" // 是否允许注册
|
||||
ConfigKeyPasswordRegisterEnabled = "password_register_enabled" // 是否允许密码注册
|
||||
ConfigKeyOIDCLoginEnabled = "oidc_login_enabled" // 是否允许 OIDC 登录
|
||||
ConfigKeyMaxAPIKeysPerUser = "max_api_keys_per_user" //nolint:gosec // false positive: config key name. 每个用户最大 API Key 数量
|
||||
ConfigKeyCapLoginEnabled = "cap_login_enabled" // 是否启用登录人机验证
|
||||
ConfigKeyCapAutoSolve = "cap_auto_solve" // 打开页面后是否自动开始计算(false 则需用户手动点击)
|
||||
ConfigKeyCapChallengeCount = "cap_challenge_count" // 客户端需求解的 PoW 难题总数,默认 1,推荐 1~5
|
||||
ConfigKeyCapChallengeSize = "cap_challenge_size" // 人机验证盐值长度
|
||||
ConfigKeyCapChallengeDifficulty = "cap_challenge_difficulty" // 人机验证 PoW 难度(目标前缀长度)
|
||||
ConfigKeyCapChallengeTTL = "cap_challenge_ttl_seconds" // 人机验证难题有效时间(秒)
|
||||
ConfigKeyCapTokenTTL = "cap_token_ttl_seconds" //nolint:gosec // false positive: config key name. 人机验证兑换凭证有效时间(秒)
|
||||
ConfigKeyServerAddress = "server_address" // 服务器地址
|
||||
ConfigKeySMTPHost = "smtp_host" // SMTP 服务器地址
|
||||
ConfigKeySMTPPort = "smtp_port" // SMTP 端口
|
||||
ConfigKeySMTPUsername = "smtp_username" // SMTP 账户
|
||||
ConfigKeySMTPPassword = "smtp_password" // SMTP 访问凭证
|
||||
ConfigKeyEmailLoginVerificationEnabled = "email_login_verification_enabled" // 是否启用邮箱登录验证
|
||||
ConfigKeyEmailRegisterVerificationEnabled = "email_register_verification_enabled" // 是否启用邮箱注册验证
|
||||
ConfigKeyMenuDisplayConfig = "menu_display_config" // 目录显示配置 (JSON 字符串)
|
||||
ConfigKeySearchEngineIndexingEnabled = "search_engine_indexing_enabled" // 是否允许搜索引擎检索
|
||||
ConfigKeyFileAccessWhitelist = "file_access_whitelist" // 免登录访问的文件业务类型白名单 (JSON 数组格式)
|
||||
ConfigKeyDiskCacheMaxSizeMB = "disk_cache_max_size_mb" // 磁盘缓存最大空间大小 (MB)
|
||||
ConfigKeyDiskCacheTTLMinutes = "disk_cache_ttl_minutes" // 磁盘缓存默认有效期 (分钟)
|
||||
ConfigKeyDiskCacheLRUEnabled = "disk_cache_lru_enabled" // 是否启用 LRU 淘汰机制
|
||||
ConfigKeyLoginSessionTTLHours = "login_session_ttl_hours" // 登录会话过期时间 (小时,0表示浏览器关闭后自动退出登录,-1表示永不过期)
|
||||
ConfigKeyUpdateUpstreamRepository = "update_upstream_repository" // GitHub Actions Release 上游仓库
|
||||
ConfigKeyStorageConfig = "storage_config" // 文件存储配置 (JSON)
|
||||
)
|
||||
|
||||
const (
|
||||
// ConfigVisibilityHidden 表示配置不通过公共配置接口暴露
|
||||
ConfigVisibilityHidden = 0
|
||||
// ConfigVisibilityVisible 表示配置通过公共配置接口暴露
|
||||
ConfigVisibilityVisible = 1
|
||||
)
|
||||
|
||||
// SystemConfig 系统配置实体
|
||||
type SystemConfig struct {
|
||||
Key string `json:"key" gorm:"primaryKey;size:64;not null"`
|
||||
Value string `json:"value" gorm:"type:text;not null"`
|
||||
Type string `json:"type" gorm:"size:32;not null;default:'system'"`
|
||||
Visibility int `json:"visibility" gorm:"not null;default:0"`
|
||||
Description string `json:"description" gorm:"size:255"`
|
||||
UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime"`
|
||||
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime"`
|
||||
}
|
||||
|
||||
// TableName 表名
|
||||
func (SystemConfig) TableName() string {
|
||||
return "w_system_configs"
|
||||
}
|
||||
@@ -0,0 +1,296 @@
|
||||
// Copyright 2025 linux.do
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package model
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
"github.com/Rain-kl/Wavelet/internal/db/idgen"
|
||||
"github.com/redis/go-redis/v9"
|
||||
)
|
||||
|
||||
// TaskExecutionStatus 任务执行状态
|
||||
type TaskExecutionStatus string
|
||||
|
||||
// 任务执行状态
|
||||
const (
|
||||
TaskExecutionStatusPending TaskExecutionStatus = "pending"
|
||||
TaskExecutionStatusRunning TaskExecutionStatus = "running"
|
||||
TaskExecutionStatusSucceeded TaskExecutionStatus = "succeeded"
|
||||
TaskExecutionStatusFailed TaskExecutionStatus = "failed"
|
||||
|
||||
taskExecutionLogRedisKeyPrefix = "task:execution:log:"
|
||||
taskExecutionLogExpiration = 24 * time.Hour
|
||||
taskExecutionLogMaxLines = 1000
|
||||
)
|
||||
|
||||
// TaskExecution 任务执行记录
|
||||
type TaskExecution struct {
|
||||
ID uint64 `json:"id,string" gorm:"primaryKey"`
|
||||
TaskID string `json:"task_id" gorm:"size:128;uniqueIndex;not null"`
|
||||
TaskType string `json:"task_type" gorm:"size:64;index;not null"`
|
||||
TaskName string `json:"task_name" gorm:"size:128"`
|
||||
Status TaskExecutionStatus `json:"status" gorm:"size:32;index;not null"`
|
||||
Retryable bool `json:"retryable" gorm:"not null;default:false"`
|
||||
MaxRetry int `json:"max_retry" gorm:"not null;default:0"`
|
||||
RetryCount int `json:"retry_count" gorm:"not null;default:0"`
|
||||
Log string `json:"log" gorm:"type:text"`
|
||||
ErrorMessage string `json:"error_message" gorm:"type:text"`
|
||||
Result string `json:"result" gorm:"type:text"`
|
||||
StartedAt *time.Time `json:"started_at" gorm:"index"`
|
||||
FinishedAt *time.Time `json:"finished_at"`
|
||||
Duration int64 `json:"duration" gorm:"comment:耗时毫秒"`
|
||||
Payload string `json:"payload" gorm:"type:text"`
|
||||
TriggeredBy string `json:"triggered_by" gorm:"size:32;not null;default:system"`
|
||||
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime;index"`
|
||||
UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime"`
|
||||
}
|
||||
|
||||
// TaskExecutionCleanupStats describes task execution log cleanup results.
|
||||
type TaskExecutionCleanupStats struct {
|
||||
HighFrequencyDeleted int64
|
||||
LowFrequencyDeleted int64
|
||||
}
|
||||
|
||||
// TableName 表名
|
||||
func (TaskExecution) TableName() string {
|
||||
return "w_task_executions"
|
||||
}
|
||||
|
||||
// CreateTaskExecution 创建任务执行记录
|
||||
func CreateTaskExecution(ctx context.Context, execution *TaskExecution) error {
|
||||
execution.ID = idgen.NextUint64ID()
|
||||
return db.DB(ctx).Create(execution).Error
|
||||
}
|
||||
|
||||
// UpdateTaskExecution 更新任务执行记录,忽略由 Redis 缓冲和归档流程管理的 log 字段。
|
||||
func UpdateTaskExecution(ctx context.Context, execution *TaskExecution) error {
|
||||
return db.DB(ctx).Omit("log").Save(execution).Error
|
||||
}
|
||||
|
||||
// GetTaskExecutionByTaskID 根据 TaskID 获取执行记录
|
||||
func GetTaskExecutionByTaskID(ctx context.Context, taskID string) (*TaskExecution, error) {
|
||||
var execution 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) (*TaskExecution, error) {
|
||||
var execution 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
|
||||
}
|
||||
|
||||
// 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(&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
|
||||
}
|
||||
|
||||
// ListTaskExecutionsRequest 查询任务执行记录列表请求
|
||||
type ListTaskExecutionsRequest struct {
|
||||
Status string `form:"status"`
|
||||
TaskType string `form:"task_type"`
|
||||
Page int `form:"page"`
|
||||
PageSize int `form:"page_size"`
|
||||
}
|
||||
|
||||
// ListTaskExecutions 分页查询任务执行记录
|
||||
func ListTaskExecutions(ctx context.Context, req ListTaskExecutionsRequest) ([]TaskExecution, int64, error) {
|
||||
if req.Page <= 0 {
|
||||
req.Page = 1
|
||||
}
|
||||
if req.PageSize <= 0 {
|
||||
req.PageSize = 20
|
||||
}
|
||||
|
||||
query := db.DB(ctx).Model(&TaskExecution{})
|
||||
|
||||
if req.Status != "" {
|
||||
query = query.Where("status = ?", req.Status)
|
||||
}
|
||||
if req.TaskType != "" {
|
||||
query = query.Where("task_type = ?", req.TaskType)
|
||||
}
|
||||
|
||||
var total int64
|
||||
if err := query.Count(&total).Error; err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
|
||||
var executions []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
|
||||
}
|
||||
|
||||
// CleanupTaskExecutionLogs removes finished task execution logs according to frequency-based retention.
|
||||
func CleanupTaskExecutionLogs(ctx context.Context, now time.Time) (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 := []TaskExecutionStatus{TaskExecutionStatusSucceeded, TaskExecutionStatusFailed}
|
||||
|
||||
var highFrequencyTaskTypes []string
|
||||
if err := db.DB(ctx).
|
||||
Model(&TaskExecution{}).
|
||||
Select("task_type").
|
||||
Where("created_at >= ?", frequencyWindowStart).
|
||||
Group("task_type").
|
||||
Having("COUNT(*) > ?", highFrequencyThreshold).
|
||||
Pluck("task_type", &highFrequencyTaskTypes).Error; err != nil {
|
||||
return 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(&TaskExecution{})
|
||||
if highFrequencyResult.Error != nil {
|
||||
return 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(&TaskExecution{})
|
||||
if lowFrequencyResult.Error != nil {
|
||||
return TaskExecutionCleanupStats{}, fmt.Errorf("delete low-frequency task execution logs: %w", lowFrequencyResult.Error)
|
||||
}
|
||||
|
||||
return TaskExecutionCleanupStats{
|
||||
HighFrequencyDeleted: highFrequencyDeleted,
|
||||
LowFrequencyDeleted: lowFrequencyResult.RowsAffected,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func taskExecutionLogRedisKey(taskID string) string {
|
||||
return db.PrefixedKey(taskExecutionLogRedisKeyPrefix + taskID)
|
||||
}
|
||||
|
||||
func loadTaskExecutionLog(ctx context.Context, execution *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 []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,484 @@
|
||||
// Copyright 2025 linux.do
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package model
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
"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(&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 := &TaskExecution{
|
||||
TaskID: "manual_cleanup_123",
|
||||
TaskType: "system:cleanup",
|
||||
TaskName: "清理未使用上传",
|
||||
Status: 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 := &TaskExecution{
|
||||
TaskID: "test_task_id_001",
|
||||
TaskType: "system:cleanup",
|
||||
TaskName: "清理未使用上传",
|
||||
Status: 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, 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 := &TaskExecution{
|
||||
TaskID: "test_by_id_001",
|
||||
TaskType: "system:cleanup",
|
||||
TaskName: "清理未使用上传",
|
||||
Status: 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 := &TaskExecution{
|
||||
TaskID: "test_update_001",
|
||||
TaskType: "system:cleanup",
|
||||
TaskName: "清理未使用上传",
|
||||
Status: TaskExecutionStatusPending,
|
||||
TriggeredBy: "manual",
|
||||
}
|
||||
err := CreateTaskExecution(ctx, execution)
|
||||
require.NoError(t, err)
|
||||
|
||||
// 更新状态为 running
|
||||
now := time.Now()
|
||||
execution.Status = 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, TaskExecutionStatusRunning, found.Status)
|
||||
assert.NotNil(t, found.StartedAt)
|
||||
|
||||
// 更新为 succeeded
|
||||
finishTime := time.Now()
|
||||
execution.Status = 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, 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 := &TaskExecution{
|
||||
TaskID: "test_fail_001",
|
||||
TaskType: "system:cleanup",
|
||||
TaskName: "清理未使用上传",
|
||||
Status: TaskExecutionStatusPending,
|
||||
Retryable: true,
|
||||
MaxRetry: 3,
|
||||
TriggeredBy: "manual",
|
||||
}
|
||||
err := CreateTaskExecution(ctx, execution)
|
||||
require.NoError(t, err)
|
||||
|
||||
// 标记为失败
|
||||
now := time.Now()
|
||||
execution.Status = 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, 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 := &TaskExecution{
|
||||
TaskID: "test_omit_log_001",
|
||||
TaskType: "system:cleanup",
|
||||
TaskName: "清理未使用上传",
|
||||
Status: 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 = TaskExecutionStatusSucceeded
|
||||
execution.Duration = 100
|
||||
err = UpdateTaskExecution(ctx, execution)
|
||||
require.NoError(t, err)
|
||||
|
||||
var persisted TaskExecution
|
||||
err = db.DB(ctx).Where("task_id = ?", "test_omit_log_001").First(&persisted).Error
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, 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 := &TaskExecution{
|
||||
TaskID: "test_log_001",
|
||||
TaskType: "system:cleanup",
|
||||
TaskName: "清理未使用上传",
|
||||
Status: 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 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 := &TaskExecution{
|
||||
TaskID: "redis_priority_001",
|
||||
TaskType: "system:cleanup",
|
||||
TaskName: "清理未使用上传",
|
||||
Status: 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 := []*TaskExecution{
|
||||
{TaskID: "list_001", TaskType: "system:cleanup", TaskName: "系统垃圾清理", Status: TaskExecutionStatusSucceeded, TriggeredBy: "manual"},
|
||||
{TaskID: "list_002", TaskType: "system:cleanup", TaskName: "系统垃圾清理", Status: TaskExecutionStatusFailed, TriggeredBy: "system"},
|
||||
{TaskID: "list_003", TaskType: "other:task", TaskName: "其他任务", Status: TaskExecutionStatusPending, TriggeredBy: "manual"},
|
||||
{TaskID: "list_004", TaskType: "system:cleanup", TaskName: "系统垃圾清理", Status: TaskExecutionStatusRunning, TriggeredBy: "manual"},
|
||||
{TaskID: "list_005", TaskType: "other:task", TaskName: "其他任务", Status: 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, 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, 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, ListTaskExecutionsRequest{TaskType: "other:task", Page: 1, PageSize: 10})
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, int64(2), total)
|
||||
|
||||
// 分页测试
|
||||
items, total, err = ListTaskExecutions(ctx, 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, 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, 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)
|
||||
}
|
||||
|
||||
func TestListTaskExecutionsDefaultPaging(t *testing.T) {
|
||||
cleanup := setupTaskExecutionTestEnvironment(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
// 不传分页参数,应使用默认值 page=1, pageSize=20
|
||||
items, total, err := ListTaskExecutions(ctx, 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", TaskExecutionStatusSucceeded, now.Add(-2*time.Hour))
|
||||
}
|
||||
createTaskExecutionForCleanup(t, ctx, "high_old_4d", "high:task", TaskExecutionStatusSucceeded, now.AddDate(0, 0, -4))
|
||||
createTaskExecutionForCleanup(t, ctx, "high_old_40d", "high:task", TaskExecutionStatusFailed, now.AddDate(0, 0, -40))
|
||||
createTaskExecutionForCleanup(t, ctx, "high_running_old", "high:task", TaskExecutionStatusRunning, now.AddDate(0, 0, -10))
|
||||
createTaskExecutionForCleanup(t, ctx, "low_old_31d", "low:task", TaskExecutionStatusSucceeded, now.AddDate(0, 0, -31))
|
||||
createTaskExecutionForCleanup(t, ctx, "low_recent_29d", "low:task", TaskExecutionStatusSucceeded, now.AddDate(0, 0, -29))
|
||||
createTaskExecutionForCleanup(t, ctx, "low_pending_old", "low:task", 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(&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(&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 := TaskExecution{}
|
||||
assert.Equal(t, "w_task_executions", execution.TableName())
|
||||
}
|
||||
|
||||
func createTaskExecutionForCleanup(t *testing.T, ctx context.Context, taskID string, taskType string, status TaskExecutionStatus, createdAt time.Time) {
|
||||
t.Helper()
|
||||
|
||||
execution := &TaskExecution{
|
||||
TaskID: taskID,
|
||||
TaskType: taskType,
|
||||
TaskName: taskType,
|
||||
Status: status,
|
||||
CreatedAt: createdAt,
|
||||
UpdatedAt: createdAt,
|
||||
TriggeredBy: "system",
|
||||
}
|
||||
err := CreateTaskExecution(ctx, execution)
|
||||
require.NoError(t, err)
|
||||
}
|
||||
@@ -0,0 +1,88 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package model
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"errors"
|
||||
"strings"
|
||||
"text/template"
|
||||
"time"
|
||||
)
|
||||
|
||||
// Template 邮件/消息模板实体
|
||||
type Template struct {
|
||||
ID uint64 `json:"id" gorm:"primaryKey;autoIncrement"`
|
||||
Key string `json:"key" gorm:"uniqueIndex;size:80;not null"`
|
||||
Name string `json:"name" gorm:"size:100;not null"`
|
||||
Type string `json:"type" gorm:"size:20;not null;default:'email'"`
|
||||
Subject string `json:"subject" gorm:"size:255"`
|
||||
Content string `json:"content" gorm:"type:text;not null"`
|
||||
Description string `json:"description" gorm:"size:255"`
|
||||
IsSystem bool `json:"is_system" gorm:"index;not null;default:false"`
|
||||
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime;index"`
|
||||
UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime;index"`
|
||||
}
|
||||
|
||||
// TableName 表名
|
||||
func (Template) TableName() string {
|
||||
return "w_templates"
|
||||
}
|
||||
|
||||
// Normalize 规范化模板字段
|
||||
func (t *Template) Normalize() {
|
||||
t.Key = strings.TrimSpace(t.Key)
|
||||
t.Name = strings.TrimSpace(t.Name)
|
||||
t.Type = strings.ToLower(strings.TrimSpace(t.Type))
|
||||
t.Subject = strings.TrimSpace(t.Subject)
|
||||
t.Content = strings.TrimSpace(t.Content)
|
||||
t.Description = strings.TrimSpace(t.Description)
|
||||
if t.Type == "" {
|
||||
t.Type = "email"
|
||||
}
|
||||
}
|
||||
|
||||
// Validate 校验模板必填字段
|
||||
func (t *Template) Validate() error {
|
||||
t.Normalize()
|
||||
if t.Key == "" {
|
||||
return errors.New(errTemplateKeyRequired)
|
||||
}
|
||||
if t.Name == "" {
|
||||
return errors.New(errTemplateNameRequired)
|
||||
}
|
||||
if t.Content == "" {
|
||||
return errors.New(errTemplateContentRequired)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// Render 渲染模板的 Subject 和 Content
|
||||
func (t *Template) Render(data any) (string, string, error) {
|
||||
// Render Subject
|
||||
var subject string
|
||||
if t.Subject != "" {
|
||||
tmplSubject, err := template.New(t.Key + "_subject").Parse(t.Subject)
|
||||
if err != nil {
|
||||
return "", "", err
|
||||
}
|
||||
var subBuf bytes.Buffer
|
||||
if err := tmplSubject.Execute(&subBuf, data); err != nil {
|
||||
return "", "", err
|
||||
}
|
||||
subject = subBuf.String()
|
||||
}
|
||||
|
||||
// Render Content
|
||||
tmplContent, err := template.New(t.Key + "_content").Parse(t.Content)
|
||||
if err != nil {
|
||||
return "", "", err
|
||||
}
|
||||
var bodyBuf bytes.Buffer
|
||||
if err := tmplContent.Execute(&bodyBuf, data); err != nil {
|
||||
return "", "", err
|
||||
}
|
||||
|
||||
return subject, bodyBuf.String(), nil
|
||||
}
|
||||
@@ -0,0 +1,28 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package model
|
||||
|
||||
import "time"
|
||||
|
||||
// Upload stats dimension keys stored in w_upload_stats.dimension.
|
||||
const (
|
||||
UploadStatDimensionTotal = "total"
|
||||
UploadStatDimensionType = "type"
|
||||
UploadStatDimensionCategory = "category"
|
||||
UploadStatDimensionTrend = "trend"
|
||||
)
|
||||
|
||||
// UploadStat stores incremental upload statistics keyed by dimension and stat_key.
|
||||
type UploadStat struct {
|
||||
Dimension string `json:"dimension" gorm:"primaryKey;size:32;not null"`
|
||||
StatKey string `json:"stat_key" gorm:"primaryKey;size:64;not null;default:''"`
|
||||
FileCount int64 `json:"file_count" gorm:"not null;default:0"`
|
||||
FileSize int64 `json:"file_size" gorm:"not null;default:0"`
|
||||
UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime"`
|
||||
}
|
||||
|
||||
// TableName returns the upload stats table name.
|
||||
func (UploadStat) TableName() string {
|
||||
return "w_upload_stats"
|
||||
}
|
||||
@@ -0,0 +1,54 @@
|
||||
// Copyright 2025 linux.do
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package model
|
||||
|
||||
import (
|
||||
"time"
|
||||
)
|
||||
|
||||
// UploadStatus 上传状态
|
||||
type UploadStatus string
|
||||
|
||||
// 上传状态
|
||||
const (
|
||||
UploadStatusPending UploadStatus = "pending" // 待使用
|
||||
UploadStatusUsed UploadStatus = "used" // 已使用
|
||||
UploadStatusDeleted UploadStatus = "deleted" // 已删除
|
||||
)
|
||||
|
||||
// UploadMetadata 自定义可扩展的 JSON 字段存储非核心或可选的文件元数据
|
||||
type UploadMetadata struct {
|
||||
Width int `json:"width,omitempty"` // 图像/视频宽度 (px)
|
||||
Height int `json:"height,omitempty"` // 图像/视频高度 (px)
|
||||
Duration float64 `json:"duration,omitempty"` // 音视频时长 (s)
|
||||
OriginalMime string `json:"original_mime,omitempty"` // 原始 MIME 类型
|
||||
UserAgent string `json:"user_agent,omitempty"` // 上传者的 UA
|
||||
ClientIP string `json:"client_ip,omitempty"` // 上传者 IP
|
||||
Bucket string `json:"bucket,omitempty"` // 存储桶名称 (适用于 S3 等)
|
||||
Extra map[string]any `json:"extra,omitempty"` // 其它任意业务自定义元数据
|
||||
}
|
||||
|
||||
// Upload 上传文件记录
|
||||
type Upload struct {
|
||||
ID uint64 `json:"id,string" gorm:"primaryKey"`
|
||||
UserID uint64 `json:"user_id,string" gorm:"index;not null"`
|
||||
FileName string `json:"file_name" gorm:"size:255;not null"` // 原始文件名 (例如: image.png)
|
||||
FilePath string `json:"file_path" gorm:"size:500;not null;index"` // 文件相对路径 / S3 Key
|
||||
FileSize int64 `json:"file_size" gorm:"not null"` // 文件大小(字节)
|
||||
MimeType string `json:"mime_type" gorm:"size:100;not null"` // 媒体类型 (MIME, 如 image/png)
|
||||
Extension string `json:"extension" gorm:"size:50;not null"` // 文件后缀名 (不含点,如 png, pdf)
|
||||
Hash string `json:"hash" gorm:"size:64;index"` // 文件哈希 (SHA-256/MD5,可用于排重)
|
||||
Type string `json:"type" gorm:"column:type;size:50;not null;index"` // 业务标识类型 (如 avatar, doc, attachment)
|
||||
Status UploadStatus `json:"status" gorm:"type:varchar(20);not null"` // 状态
|
||||
AccessMode int `json:"access_mode" gorm:"column:access_mode;not null;default:0"`
|
||||
Metadata UploadMetadata `json:"metadata" gorm:"serializer:json;type:jsonb"` // 业务扩展元数据
|
||||
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime"`
|
||||
UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime"`
|
||||
}
|
||||
|
||||
// TableName 表名
|
||||
func (Upload) TableName() string {
|
||||
return "w_uploads"
|
||||
}
|
||||
@@ -0,0 +1,195 @@
|
||||
// Copyright 2025 linux.do
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package model
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/common"
|
||||
"github.com/Rain-kl/Wavelet/internal/db/idgen"
|
||||
"github.com/Rain-kl/Wavelet/pkg/util"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
// OAuthUserInfo 用户信息结构(同时支持 OIDC ID Token claims 和 UserEndpoint 响应)
|
||||
type OAuthUserInfo struct {
|
||||
ID uint64 `json:"id"`
|
||||
Sub string `json:"sub"`
|
||||
Username string `json:"username"`
|
||||
PreferredUsername string `json:"preferred_username"`
|
||||
Email string `json:"email"`
|
||||
Name string `json:"name"`
|
||||
Active bool `json:"active"`
|
||||
AvatarURL string `json:"avatar_url"`
|
||||
}
|
||||
|
||||
// GetID 获取用户 ID
|
||||
func (u *OAuthUserInfo) GetID() uint64 {
|
||||
if u.ID != 0 {
|
||||
return u.ID
|
||||
}
|
||||
// 从 sub 解析(OIDC 格式)
|
||||
if u.Sub != "" {
|
||||
if id, err := strconv.ParseUint(u.Sub, 10, 64); err == nil {
|
||||
return id
|
||||
}
|
||||
}
|
||||
return 0
|
||||
}
|
||||
|
||||
// User 用户表实体
|
||||
type User struct {
|
||||
ID uint64 `json:"id,string" gorm:"primaryKey;not null"`
|
||||
Username string `json:"username" gorm:"size:64;uniqueIndex"`
|
||||
Password string `json:"password,omitempty" gorm:"size:255"`
|
||||
Nickname string `json:"nickname" gorm:"size:255"`
|
||||
Email string `json:"email" gorm:"size:255;index"`
|
||||
AvatarURL string `json:"avatar_url" gorm:"size:255"`
|
||||
IsActive bool `json:"is_active" gorm:"default:true;index"`
|
||||
IsAdmin bool `json:"is_admin" gorm:"default:false"`
|
||||
Bio string `json:"bio" gorm:"size:500"`
|
||||
Phone string `json:"phone" gorm:"size:32"`
|
||||
Gender string `json:"gender" gorm:"size:16"`
|
||||
Website string `json:"website" gorm:"size:255"`
|
||||
Location string `json:"location" gorm:"size:255"`
|
||||
LastLoginAt time.Time `json:"last_login_at" gorm:"index"`
|
||||
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime;index"`
|
||||
UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime;index"`
|
||||
}
|
||||
|
||||
// TableName 表名
|
||||
func (User) TableName() string {
|
||||
return "w_users"
|
||||
}
|
||||
|
||||
// SetPassword 设置明文密码
|
||||
func (u *User) SetPassword(password string) error {
|
||||
u.Password = password
|
||||
return nil
|
||||
}
|
||||
|
||||
// SetEncryptedPassword 设置加密密码
|
||||
func (u *User) SetEncryptedPassword(password string) error {
|
||||
if password == "" {
|
||||
u.Password = ""
|
||||
return nil
|
||||
}
|
||||
hashed, err := util.HashPassword(password)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
u.Password = hashed
|
||||
return nil
|
||||
}
|
||||
|
||||
// IsPasswordEncrypted 检查密码是否已加密
|
||||
func (u *User) IsPasswordEncrypted() bool {
|
||||
return strings.HasPrefix(u.Password, "$2a$") || strings.HasPrefix(u.Password, "$2b$") || strings.HasPrefix(u.Password, "$2y$")
|
||||
}
|
||||
|
||||
// CheckPassword 验证密码是否匹配
|
||||
func (u *User) CheckPassword(password string) bool {
|
||||
if u.Password == "" || password == "" {
|
||||
return false
|
||||
}
|
||||
if u.IsPasswordEncrypted() {
|
||||
return util.CheckPasswordHash(u.Password, password)
|
||||
}
|
||||
return u.Password == password
|
||||
}
|
||||
|
||||
// GetByID 根据 ID 查询用户
|
||||
func (u *User) GetByID(tx *gorm.DB, id uint64) error {
|
||||
if err := tx.Where("id = ?", id).First(u).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// UpdateFromOAuthInfo 根据 OAuth 信息更新用户数据
|
||||
func (u *User) UpdateFromOAuthInfo(oauthInfo *OAuthUserInfo) {
|
||||
u.Username = oauthInfo.Username
|
||||
u.Nickname = oauthInfo.Name
|
||||
u.Email = oauthInfo.Email
|
||||
u.AvatarURL = oauthInfo.AvatarURL
|
||||
u.IsActive = oauthInfo.Active
|
||||
u.LastLoginAt = time.Now()
|
||||
}
|
||||
|
||||
// CheckActive 检查用户账户是否激活,未激活则返回错误
|
||||
func (u *User) CheckActive() error {
|
||||
if !u.IsActive {
|
||||
return errors.New(common.BannedAccount)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (u *User) assignIDIfMissing() error {
|
||||
if u.ID != 0 {
|
||||
return nil
|
||||
}
|
||||
u.ID = idgen.NextUint64ID()
|
||||
return nil
|
||||
}
|
||||
|
||||
// CreateUser 创建新用户(用于 OAuth/OIDC 自动注册,含底层权限校验)
|
||||
func (u *User) CreateUser(_ context.Context, tx *gorm.DB, oauthInfo *OAuthUserInfo) error {
|
||||
now := time.Now()
|
||||
userID := oauthInfo.GetID()
|
||||
newUser := User{
|
||||
ID: userID,
|
||||
Username: oauthInfo.Username,
|
||||
Nickname: oauthInfo.Name,
|
||||
Email: oauthInfo.Email,
|
||||
AvatarURL: oauthInfo.AvatarURL,
|
||||
IsActive: oauthInfo.Active,
|
||||
LastLoginAt: now,
|
||||
IsAdmin: false,
|
||||
}
|
||||
if err := newUser.assignIDIfMissing(); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := tx.Create(&newUser).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
*u = newUser
|
||||
return nil
|
||||
}
|
||||
|
||||
// RegisterUser 创建新用户并注册(用于本地密码注册,含全局开关和唯一性多重底层校验)
|
||||
func (u *User) RegisterUser(_ context.Context, tx *gorm.DB) error {
|
||||
// 检查用户名冲突
|
||||
var count int64
|
||||
if err := tx.Model(&User{}).Where("username = ?", u.Username).Count(&count).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if count > 0 {
|
||||
return errors.New(errUsernameExists)
|
||||
}
|
||||
|
||||
// 检查邮箱冲突
|
||||
if u.Email != "" {
|
||||
var emailCount int64
|
||||
if err := tx.Model(&User{}).Where("email = ?", u.Email).Count(&emailCount).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if emailCount > 0 {
|
||||
return errors.New(errEmailAlreadyBound)
|
||||
}
|
||||
}
|
||||
|
||||
if err := u.assignIDIfMissing(); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := tx.Create(u).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
Reference in New Issue
Block a user