refactor(repository): 收敛 model/repository 分层为唯一持久化入口

将 OpenFlare 与平台业务的数据访问从 model 与 apps 直连迁入 repository,
model 仅保留实体与无 IO 规则;补充 code-check 架构守卫与开发规范。
This commit is contained in:
ryan
2026-07-24 17:00:17 +08:00
parent 23a5488203
commit 943818f7d4
184 changed files with 5592 additions and 4364 deletions
-204
View File
@@ -4,14 +4,10 @@
package model
import (
"context"
"errors"
"regexp"
"strings"
"time"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
"gorm.io/gorm"
)
// 认证源类型
@@ -124,203 +120,3 @@ 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(&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(&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(&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) ([]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
}
+10 -24
View File
@@ -3,29 +3,15 @@
package model
// Domain validation messages used by model.Validate and other no-IO rules.
// Persistence / data-access messages belong in internal/repository (do not import repository).
const (
errRegistrationDisabled = "注册已关闭"
errDatabaseNotInitialized = "database not initialized"
errClickHouseNotInitialized = "clickhouse 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 不能为空"
errTemplateKeyRequired = "模板标识符不能为空"
errTemplateNameRequired = "模板名称不能为空"
errTemplateContentRequired = "模板内容不能为空"
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
)
+18 -464
View File
@@ -3,472 +3,26 @@
package model
import (
"context"
"math"
"sort"
"strings"
"time"
analyticsmodel "github.com/Rain-kl/Wavelet/internal/model/analytics"
)
const (
columnRemoteAddr = "remote_addr"
columnHost = "host"
sortOrderAsc = "asc"
secondsPerMinute = 60
)
type openFlareAccessLogBucketAggregateRow = analyticsmodel.NodeAccessLogBucketAggregate
type openFlareAccessLogBucketDimensionRow = analyticsmodel.NodeAccessLogBucketDimension
type openFlareAccessLogIPAggregateRow = analyticsmodel.NodeAccessLogIPAggregate
type openFlareAccessLogIPSummaryRow = analyticsmodel.NodeAccessLogIPSummary
type openFlareAccessLogIPTrendRow = analyticsmodel.NodeAccessLogIPTrend
type openFlareAccessLogWAFIPAggregateRow = analyticsmodel.NodeAccessLogWAFIPAggregate
// ListOpenFlareAccessLogWAFIPAggregates returns per-IP aggregates for WAF automatic rules.
func ListOpenFlareAccessLogWAFIPAggregates(ctx context.Context, query OpenFlareAccessLogQuery) ([]*OpenFlareAccessLogWAFIPAggregate, error) {
rows, err := currentAccessLogStore().WAFIPAggregates(ctx, query)
if err != nil {
return nil, err
}
result := make([]*OpenFlareAccessLogWAFIPAggregate, 0, len(rows))
for _, row := range rows {
remoteAddr := strings.TrimSpace(row.RemoteAddr)
if remoteAddr == "" {
continue
}
statusCounts := make(map[int]int, len(row.StatusCounts))
for code, count := range row.StatusCounts {
statusCounts[code] = int(count)
}
result = append(result, &OpenFlareAccessLogWAFIPAggregate{
RemoteAddr: remoteAddr,
RequestCount: int(row.RequestCount),
Status404Count: int(row.Status404Count),
ClientErrorCount: int(row.ClientErrorCount),
ServerErrorCount: int(row.ServerErrorCount),
IPHostCount: int(row.IPHostCount),
LastSeenEpoch: row.LastSeenEpoch,
StatusCounts: statusCounts,
})
}
return result, nil
// OpenFlareAccessLogTrafficSummary is a window-level traffic summary from access logs.
type OpenFlareAccessLogTrafficSummary struct {
RequestCount int64
ErrorCount int64
UniqueIPCount int64
BytesSent int64
RequestLength int64
NodeCount int64
}
// InsertOpenFlareAccessLogsBatch inserts access log rows into ClickHouse.
func InsertOpenFlareAccessLogsBatch(ctx context.Context, records []*OpenFlareAccessLog) error {
return currentAccessLogStore().InsertBatch(ctx, records)
// OpenFlareAccessLogValueCount is a dimension value count.
type OpenFlareAccessLogValueCount struct {
Value string
Count int64
}
// ListOpenFlareAccessLogs lists access logs matching the query.
func ListOpenFlareAccessLogs(ctx context.Context, query OpenFlareAccessLogQuery) ([]*OpenFlareAccessLog, error) {
return currentAccessLogStore().List(ctx, query)
}
// CountOpenFlareAccessLogs counts access logs, distinct IPs, and total bytes sent matching the query.
func CountOpenFlareAccessLogs(ctx context.Context, query OpenFlareAccessLogQuery) (int64, int64, int64, error) {
return currentAccessLogStore().Count(ctx, query)
}
// TrafficSummaryOpenFlareAccessLogs returns window-level request/error/UV/bytes summary.
func TrafficSummaryOpenFlareAccessLogs(ctx context.Context, query OpenFlareAccessLogQuery) (OpenFlareAccessLogTrafficSummary, error) {
return currentAccessLogStore().TrafficSummary(ctx, query)
}
// ValueCountsOpenFlareAccessLogs groups logs by status_code, host, path, remote_addr, or user_agent.
func ValueCountsOpenFlareAccessLogs(ctx context.Context, query OpenFlareAccessLogQuery, column string, limit int) ([]OpenFlareAccessLogValueCount, error) {
return currentAccessLogStore().ValueCounts(ctx, query, column, limit)
}
// NodeAggregatesOpenFlareAccessLogs returns per-node request/error/UV for the window.
func NodeAggregatesOpenFlareAccessLogs(ctx context.Context, query OpenFlareAccessLogQuery) ([]OpenFlareAccessLogNodeAggregate, error) {
return currentAccessLogStore().NodeAggregates(ctx, query)
}
// ListOpenFlareAccessLogRegionCounts returns region counts for access logs.
func ListOpenFlareAccessLogRegionCounts(ctx context.Context, nodeID string, since time.Time, limit int) ([]*OpenFlareAccessLogRegionCount, error) {
return currentAccessLogStore().RegionCounts(ctx, nodeID, since, limit)
}
// ListOpenFlareAccessLogBuckets lists folded access log buckets.
func ListOpenFlareAccessLogBuckets(ctx context.Context, query OpenFlareAccessLogBucketQuery) ([]*OpenFlareAccessLogBucketRow, error) {
return buildOpenFlareAccessLogBucketRows(ctx, query)
}
// CountOpenFlareAccessLogBuckets counts folded access log buckets.
func CountOpenFlareAccessLogBuckets(ctx context.Context, query OpenFlareAccessLogBucketQuery) (int64, error) {
filter := openFlareAccessLogQueryFromBucket(query)
bucketSeconds := int64(query.FoldMinutes * secondsPerMinute)
if bucketSeconds <= 0 {
bucketSeconds = 180
}
return currentAccessLogStore().CountBuckets(ctx, filter, bucketSeconds)
}
// ListOpenFlareAccessLogBucketIPs lists folded IP rows for a bucket window.
func ListOpenFlareAccessLogBucketIPs(ctx context.Context, query OpenFlareAccessLogBucketIPQuery) ([]*OpenFlareAccessLogBucketIPRow, error) {
rows, err := buildOpenFlareAccessLogBucketIPRows(ctx, query)
if err != nil {
return nil, err
}
start, end := openFlareAccessLogPaginateBounds(len(rows), query.Page, query.PageSize)
if start >= len(rows) {
return []*OpenFlareAccessLogBucketIPRow{}, nil
}
return rows[start:end], nil
}
// CountOpenFlareAccessLogBucketIPs counts folded IP rows for a bucket window.
func CountOpenFlareAccessLogBucketIPs(ctx context.Context, query OpenFlareAccessLogBucketIPQuery) (int64, error) {
rows, err := buildOpenFlareAccessLogBucketIPRows(ctx, query)
if err != nil {
return 0, err
}
return int64(len(rows)), nil
}
// ListOpenFlareAccessLogIPSummaries lists IP summaries.
func ListOpenFlareAccessLogIPSummaries(ctx context.Context, query OpenFlareAccessLogIPSummaryQuery, recentSince time.Time) ([]*OpenFlareAccessLogIPSummaryRow, error) {
return buildOpenFlareAccessLogIPSummaryRows(ctx, query, recentSince)
}
// CountOpenFlareAccessLogIPSummaries counts IP summaries.
func CountOpenFlareAccessLogIPSummaries(ctx context.Context, query OpenFlareAccessLogIPSummaryQuery) (int64, error) {
filter := openFlareAccessLogQueryFromIPSummary(query)
return currentAccessLogStore().CountIPSummaries(ctx, filter)
}
// ListOpenFlareAccessLogIPTrend lists IP trend points.
func ListOpenFlareAccessLogIPTrend(ctx context.Context, query OpenFlareAccessLogIPTrendQuery) ([]*OpenFlareAccessLogIPTrendRow, error) {
remoteAddr := strings.TrimSpace(query.RemoteAddr)
if remoteAddr == "" {
return []*OpenFlareAccessLogIPTrendRow{}, nil
}
filter := OpenFlareAccessLogQuery{
NodeID: query.NodeID,
RemoteAddr: remoteAddr,
Host: query.Host,
Since: query.Since,
}
bucketSeconds := int64(query.BucketMinutes * secondsPerMinute)
if bucketSeconds <= 0 {
bucketSeconds = 1800
}
rows, err := currentAccessLogStore().IPTrend(ctx, filter, bucketSeconds)
if err != nil {
return nil, err
}
result := make([]*OpenFlareAccessLogIPTrendRow, len(rows))
for index, row := range rows {
result[index] = &OpenFlareAccessLogIPTrendRow{
BucketEpoch: row.BucketEpoch,
RequestCount: row.RequestCount,
}
}
return result, nil
}
// DeleteAllOpenFlareAccessLogs deletes all access logs.
func DeleteAllOpenFlareAccessLogs(ctx context.Context) (int64, error) {
return currentAccessLogStore().DeleteAll(ctx)
}
// DeleteOpenFlareAccessLogsBefore deletes access logs older than cutoff.
func DeleteOpenFlareAccessLogsBefore(ctx context.Context, cutoff time.Time) (int64, error) {
return currentAccessLogStore().DeleteBefore(ctx, cutoff)
}
// DeleteOpenFlareAccessLogsByNodeBefore deletes access logs for a node older than cutoff.
func DeleteOpenFlareAccessLogsByNodeBefore(ctx context.Context, nodeID string, cutoff time.Time) (int64, error) {
return currentAccessLogStore().DeleteByNodeBefore(ctx, nodeID, cutoff)
}
func buildOpenFlareAccessLogBucketRows(ctx context.Context, query OpenFlareAccessLogBucketQuery) ([]*OpenFlareAccessLogBucketRow, error) {
filter := openFlareAccessLogQueryFromBucket(query)
bucketSeconds := int64(query.FoldMinutes * secondsPerMinute)
if bucketSeconds <= 0 {
bucketSeconds = 180
}
partials, err := currentAccessLogStore().BucketAggregates(ctx, filter, bucketSeconds)
if err != nil {
return nil, err
}
rows := make([]*OpenFlareAccessLogBucketRow, 0, len(partials))
for _, partial := range partials {
rows = append(rows, &OpenFlareAccessLogBucketRow{
BucketEpoch: partial.BucketEpoch,
RequestCount: partial.RequestCount,
UniqueIPCount: partial.UniqueIPCount,
UniqueHostCount: partial.UniqueHostCount,
SuccessCount: partial.SuccessCount,
ClientErrorCount: partial.ClientErrorCount,
ServerErrorCount: partial.ServerErrorCount,
BytesSent: partial.BytesSent,
RequestLength: partial.RequestLength,
})
}
return rows, nil
}
func buildOpenFlareAccessLogBucketIPRows(ctx context.Context, query OpenFlareAccessLogBucketIPQuery) ([]*OpenFlareAccessLogBucketIPRow, error) {
if query.BucketStartedAt.IsZero() {
return []*OpenFlareAccessLogBucketIPRow{}, nil
}
foldMinutes := query.FoldMinutes
if foldMinutes <= 0 {
foldMinutes = 3
}
bucketStartedAt := query.BucketStartedAt.UTC()
filter := OpenFlareAccessLogQuery{
NodeID: query.NodeID,
RemoteAddr: query.RemoteAddr,
Host: query.Host,
Path: query.Path,
Since: bucketStartedAt,
Until: bucketStartedAt.Add(time.Duration(foldMinutes) * time.Minute),
}
rows, err := queryOpenFlareAccessLogIPAggregateRows(ctx, filter, false)
if err != nil {
return nil, err
}
sortOpenFlareAccessLogBucketIPRows(rows, query.SortBy, query.SortOrder)
return rows, nil
}
func buildOpenFlareAccessLogIPSummaryRows(ctx context.Context, query OpenFlareAccessLogIPSummaryQuery, recentSince time.Time) ([]*OpenFlareAccessLogIPSummaryRow, error) {
filter := openFlareAccessLogQueryFromIPSummary(query)
partials, err := currentAccessLogStore().IPSummaries(ctx, filter, recentSince)
if err != nil {
return nil, err
}
rows := make([]*OpenFlareAccessLogIPSummaryRow, 0, len(partials))
for _, partial := range partials {
remoteAddr := strings.TrimSpace(partial.RemoteAddr)
if remoteAddr == "" {
continue
}
rows = append(rows, &OpenFlareAccessLogIPSummaryRow{
RemoteAddr: remoteAddr,
Region: strings.TrimSpace(partial.Region),
TotalRequests: partial.TotalRequests,
Success2xxCount: partial.Success2xxCount,
SuccessRatio: partial.SuccessRatio,
BytesReceived: partial.BytesReceived,
BytesSent: partial.BytesSent,
RecentRequests: 0,
LastSeenEpoch: partial.LastSeenEpoch,
})
}
return rows, nil
}
func queryOpenFlareAccessLogIPAggregateRows(ctx context.Context, filter OpenFlareAccessLogQuery, exactRemoteAddr bool) ([]*OpenFlareAccessLogBucketIPRow, error) {
partials, err := currentAccessLogStore().IPAggregates(ctx, filter, exactRemoteAddr)
if err != nil {
return nil, err
}
rows := make([]*OpenFlareAccessLogBucketIPRow, 0, len(partials))
for _, partial := range partials {
remoteAddr := strings.TrimSpace(partial.RemoteAddr)
if remoteAddr == "" {
continue
}
rows = append(rows, &OpenFlareAccessLogBucketIPRow{
RemoteAddr: remoteAddr,
RequestCount: partial.RequestCount,
SuccessCount: partial.SuccessCount,
ClientErrorCount: partial.ClientErrorCount,
ServerErrorCount: partial.ServerErrorCount,
LastSeenEpoch: partial.LastSeenEpoch,
})
}
return rows, nil
}
func openFlareAccessLogQueryFromBucket(query OpenFlareAccessLogBucketQuery) OpenFlareAccessLogQuery {
return OpenFlareAccessLogQuery{
NodeID: query.NodeID,
RemoteAddr: query.RemoteAddr,
Host: query.Host,
Hosts: query.Hosts,
Path: query.Path,
Since: query.Since,
Until: query.Until,
Page: query.Page,
PageSize: query.PageSize,
SortBy: query.SortBy,
SortOrder: query.SortOrder,
}
}
func openFlareAccessLogQueryFromIPSummary(query OpenFlareAccessLogIPSummaryQuery) OpenFlareAccessLogQuery {
return OpenFlareAccessLogQuery{
NodeID: query.NodeID,
RemoteAddr: query.RemoteAddr,
Host: query.Host,
Since: query.Since,
Until: query.Until,
Page: query.Page,
PageSize: query.PageSize,
SortBy: query.SortBy,
SortOrder: query.SortOrder,
}
}
func sortOpenFlareAccessLogBucketIPRows(items []*OpenFlareAccessLogBucketIPRow, sortBy string, sortOrder string) {
desc := openFlareAccessLogNormalizeSortOrder(sortOrder) != sortOrderAsc
sort.Slice(items, func(i, j int) bool {
left := items[i]
right := items[j]
if left == nil || right == nil {
return left != nil
}
var compare int
switch strings.TrimSpace(sortBy) {
case "last_seen_at":
compare = openFlareAccessLogCompareInt64(left.LastSeenEpoch, right.LastSeenEpoch)
case "remote_addr":
compare = strings.Compare(left.RemoteAddr, right.RemoteAddr)
default:
compare = openFlareAccessLogCompareInt64(left.RequestCount, right.RequestCount)
}
if compare == 0 {
compare = openFlareAccessLogCompareInt64(left.LastSeenEpoch, right.LastSeenEpoch)
}
if compare == 0 {
compare = strings.Compare(left.RemoteAddr, right.RemoteAddr)
}
if desc {
return compare > 0
}
return compare < 0
})
}
func sortOpenFlareAccessLogBucketRows(items []*OpenFlareAccessLogBucketRow, sortBy string, sortOrder string) {
desc := openFlareAccessLogNormalizeSortOrder(sortOrder) != sortOrderAsc
sort.Slice(items, func(i, j int) bool {
left := items[i]
right := items[j]
if left == nil || right == nil {
return left != nil
}
var compare int
switch strings.TrimSpace(sortBy) {
case "request_count":
compare = openFlareAccessLogCompareInt64(left.RequestCount, right.RequestCount)
default:
compare = openFlareAccessLogCompareInt64(left.BucketEpoch, right.BucketEpoch)
}
if compare == 0 {
compare = openFlareAccessLogCompareInt64(left.BucketEpoch, right.BucketEpoch)
}
if desc {
return compare > 0
}
return compare < 0
})
}
func sortOpenFlareAccessLogIPSummaryRows(items []*OpenFlareAccessLogIPSummaryRow, sortBy string, sortOrder string) {
desc := openFlareAccessLogNormalizeSortOrder(sortOrder) != sortOrderAsc
sort.Slice(items, func(i, j int) bool {
left := items[i]
right := items[j]
if left == nil || right == nil {
return left != nil
}
var compare int
switch strings.TrimSpace(sortBy) {
case "request_length", "bytes_received":
compare = openFlareAccessLogCompareInt64(left.BytesReceived, right.BytesReceived)
case "bytes_sent":
compare = openFlareAccessLogCompareInt64(left.BytesSent, right.BytesSent)
case "success_ratio":
compare = openFlareAccessLogCompareFloat64(left.SuccessRatio, right.SuccessRatio)
case "last_seen_at":
compare = openFlareAccessLogCompareInt64(left.LastSeenEpoch, right.LastSeenEpoch)
case "remote_addr":
compare = strings.Compare(left.RemoteAddr, right.RemoteAddr)
default:
compare = openFlareAccessLogCompareInt64(left.TotalRequests, right.TotalRequests)
}
if compare == 0 {
compare = openFlareAccessLogCompareInt64(left.LastSeenEpoch, right.LastSeenEpoch)
}
if compare == 0 {
compare = strings.Compare(left.RemoteAddr, right.RemoteAddr)
}
if desc {
return compare > 0
}
return compare < 0
})
}
func openFlareAccessLogCompareFloat64(left, right float64) int {
if left < right {
return -1
}
if left > right {
return 1
}
return 0
}
func openFlareAccessLogPaginateBounds(total int, page int, pageSize int) (int, int) {
if page < 0 {
page = 0
}
if pageSize <= 0 {
return 0, total
}
start := page * pageSize
if start > total {
start = total
}
end := start + pageSize
if end > total {
end = total
}
return start, end
}
func openFlareAccessLogNormalizeSortOrder(sortOrder string) string {
if strings.EqualFold(strings.TrimSpace(sortOrder), sortOrderAsc) {
return sortOrderAsc
}
return "desc"
}
func openFlareAccessLogStatusCodeToInt32(code int) int32 {
switch {
case code > math.MaxInt32:
return math.MaxInt32
case code < math.MinInt32:
return math.MinInt32
default:
return int32(code)
}
}
func openFlareAccessLogUintToInt64(value uint64) int64 {
if value > math.MaxInt64 {
return math.MaxInt64
}
return int64(value)
}
func openFlareAccessLogCompareInt64(left int64, right int64) int {
switch {
case left > right:
return 1
case left < right:
return -1
default:
return 0
}
// OpenFlareAccessLogNodeAggregate is per-node traffic over a window.
type OpenFlareAccessLogNodeAggregate struct {
NodeID string
RequestCount int64
ErrorCount int64
UniqueIPCount int64
}
@@ -1,330 +0,0 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package model
import (
"context"
"math"
"sync"
"time"
analyticsmodel "github.com/Rain-kl/Wavelet/internal/model/analytics"
analyticsrepo "github.com/Rain-kl/Wavelet/internal/repository/analytics"
)
// AccessLogInsertHooks queues node access logs for async ClickHouse write.
// Wired from openflare/chwriter.Init so model never imports the apps layer.
type AccessLogInsertHooks struct {
QueueNodeAccessLogs func(logs []analyticsmodel.NodeAccessLog)
}
var (
accessLogInsertHooksMu sync.RWMutex
accessLogInsertHooks AccessLogInsertHooks
)
// SetAccessLogInsertHooks registers async queue callbacks for access log inserts.
func SetAccessLogInsertHooks(hooks AccessLogInsertHooks) {
accessLogInsertHooksMu.Lock()
accessLogInsertHooks = hooks
accessLogInsertHooksMu.Unlock()
}
func currentAccessLogInsertHooks() AccessLogInsertHooks {
accessLogInsertHooksMu.RLock()
defer accessLogInsertHooksMu.RUnlock()
return accessLogInsertHooks
}
type accessLogStore interface {
InsertBatch(ctx context.Context, records []*OpenFlareAccessLog) error
List(ctx context.Context, query OpenFlareAccessLogQuery) ([]*OpenFlareAccessLog, error)
Count(ctx context.Context, query OpenFlareAccessLogQuery) (int64, int64, int64, error)
RegionCounts(ctx context.Context, nodeID string, since time.Time, limit int) ([]*OpenFlareAccessLogRegionCount, error)
BucketAggregates(ctx context.Context, filter OpenFlareAccessLogQuery, bucketSeconds int64) ([]openFlareAccessLogBucketAggregateRow, error)
CountBuckets(ctx context.Context, filter OpenFlareAccessLogQuery, bucketSeconds int64) (int64, error)
BucketDimensions(ctx context.Context, filter OpenFlareAccessLogQuery, column string, bucketSeconds int64) ([]openFlareAccessLogBucketDimensionRow, error)
IPAggregates(ctx context.Context, filter OpenFlareAccessLogQuery, exactRemoteAddr bool) ([]openFlareAccessLogIPAggregateRow, error)
WAFIPAggregates(ctx context.Context, filter OpenFlareAccessLogQuery) ([]openFlareAccessLogWAFIPAggregateRow, error)
IPSummaries(ctx context.Context, filter OpenFlareAccessLogQuery, recentSince time.Time) ([]openFlareAccessLogIPSummaryRow, error)
CountIPSummaries(ctx context.Context, filter OpenFlareAccessLogQuery) (int64, error)
IPTrend(ctx context.Context, filter OpenFlareAccessLogQuery, bucketSeconds int64) ([]openFlareAccessLogIPTrendRow, error)
TrafficSummary(ctx context.Context, filter OpenFlareAccessLogQuery) (OpenFlareAccessLogTrafficSummary, error)
ValueCounts(ctx context.Context, filter OpenFlareAccessLogQuery, column string, limit int) ([]OpenFlareAccessLogValueCount, error)
NodeAggregates(ctx context.Context, filter OpenFlareAccessLogQuery) ([]OpenFlareAccessLogNodeAggregate, error)
DeleteAll(ctx context.Context) (int64, error)
DeleteBefore(ctx context.Context, cutoff time.Time) (int64, error)
DeleteByNodeBefore(ctx context.Context, nodeID string, before time.Time) (int64, error)
}
// OpenFlareAccessLogTrafficSummary is a window-level traffic summary from access logs.
type OpenFlareAccessLogTrafficSummary struct {
RequestCount int64
ErrorCount int64
UniqueIPCount int64
BytesSent int64
RequestLength int64
NodeCount int64
}
// OpenFlareAccessLogValueCount is a dimension value count.
type OpenFlareAccessLogValueCount struct {
Value string
Count int64
}
// OpenFlareAccessLogNodeAggregate is per-node traffic over a window.
type OpenFlareAccessLogNodeAggregate struct {
NodeID string
RequestCount int64
ErrorCount int64
UniqueIPCount int64
}
var (
accessLogStoreMu sync.RWMutex
accessLogStoreHolder accessLogStore
)
func currentAccessLogStore() accessLogStore {
accessLogStoreMu.RLock()
defer accessLogStoreMu.RUnlock()
if accessLogStoreHolder != nil {
return accessLogStoreHolder
}
return clickhouseAccessLogStore{}
}
// SetAccessLogStoreForTest swaps the access log store implementation for unit tests.
func SetAccessLogStoreForTest(store accessLogStore) func() {
accessLogStoreMu.Lock()
previous := accessLogStoreHolder
accessLogStoreHolder = store
accessLogStoreMu.Unlock()
return func() {
accessLogStoreMu.Lock()
accessLogStoreHolder = previous
accessLogStoreMu.Unlock()
}
}
// NewMemoryAccessLogStore returns an in-memory access log store for unit tests.
func NewMemoryAccessLogStore() accessLogStore {
return &memoryAccessLogStore{
records: make([]*OpenFlareAccessLog, 0),
}
}
type clickhouseAccessLogStore struct{}
func (clickhouseAccessLogStore) InsertBatch(_ context.Context, records []*OpenFlareAccessLog) error {
logs := make([]analyticsmodel.NodeAccessLog, 0, len(records))
for _, record := range records {
if record == nil {
continue
}
logs = append(logs, toAnalyticsNodeAccessLog(record))
}
if hook := currentAccessLogInsertHooks().QueueNodeAccessLogs; hook != nil {
hook(logs)
}
return nil
}
func (clickhouseAccessLogStore) List(ctx context.Context, query OpenFlareAccessLogQuery) ([]*OpenFlareAccessLog, error) {
rows, err := analyticsrepo.ListNodeAccessLogs(ctx, toNodeAccessLogFilter(query))
if err != nil {
return nil, err
}
return fromAnalyticsNodeAccessLogs(rows), nil
}
func (clickhouseAccessLogStore) Count(ctx context.Context, query OpenFlareAccessLogQuery) (int64, int64, int64, error) {
return analyticsrepo.CountNodeAccessLogs(ctx, toNodeAccessLogFilter(query))
}
func (clickhouseAccessLogStore) RegionCounts(ctx context.Context, nodeID string, since time.Time, limit int) ([]*OpenFlareAccessLogRegionCount, error) {
rows, err := analyticsrepo.RegionCountsNodeAccessLogs(ctx, nodeID, since, limit)
if err != nil {
return nil, err
}
result := make([]*OpenFlareAccessLogRegionCount, len(rows))
for index, row := range rows {
result[index] = &OpenFlareAccessLogRegionCount{
Region: row.Region,
Count: row.Count,
}
}
return result, nil
}
func (clickhouseAccessLogStore) BucketAggregates(ctx context.Context, filter OpenFlareAccessLogQuery, bucketSeconds int64) ([]openFlareAccessLogBucketAggregateRow, error) {
return analyticsrepo.BucketAggregatesNodeAccessLogs(ctx, toNodeAccessLogFilter(filter), bucketSeconds)
}
func (clickhouseAccessLogStore) CountBuckets(ctx context.Context, filter OpenFlareAccessLogQuery, bucketSeconds int64) (int64, error) {
return analyticsrepo.CountBucketAggregatesNodeAccessLogs(ctx, toNodeAccessLogFilter(filter), bucketSeconds)
}
func (clickhouseAccessLogStore) BucketDimensions(ctx context.Context, filter OpenFlareAccessLogQuery, column string, bucketSeconds int64) ([]openFlareAccessLogBucketDimensionRow, error) {
return analyticsrepo.BucketDimensionsNodeAccessLogs(ctx, toNodeAccessLogFilter(filter), column, bucketSeconds)
}
func (clickhouseAccessLogStore) IPAggregates(ctx context.Context, filter OpenFlareAccessLogQuery, exactRemoteAddr bool) ([]openFlareAccessLogIPAggregateRow, error) {
return analyticsrepo.IPAggregatesNodeAccessLogs(ctx, toNodeAccessLogFilter(filter), exactRemoteAddr)
}
func (clickhouseAccessLogStore) IPSummaries(ctx context.Context, filter OpenFlareAccessLogQuery, recentSince time.Time) ([]openFlareAccessLogIPSummaryRow, error) {
return analyticsrepo.IPSummariesNodeAccessLogs(ctx, toNodeAccessLogFilter(filter), recentSince)
}
func (clickhouseAccessLogStore) CountIPSummaries(ctx context.Context, filter OpenFlareAccessLogQuery) (int64, error) {
return analyticsrepo.CountIPSummaryNodeAccessLogs(ctx, toNodeAccessLogFilter(filter))
}
func (clickhouseAccessLogStore) WAFIPAggregates(ctx context.Context, filter OpenFlareAccessLogQuery) ([]openFlareAccessLogWAFIPAggregateRow, error) {
return analyticsrepo.IPAggregatesForWAFNodeAccessLogs(ctx, toNodeAccessLogFilter(filter))
}
func (clickhouseAccessLogStore) IPTrend(ctx context.Context, filter OpenFlareAccessLogQuery, bucketSeconds int64) ([]openFlareAccessLogIPTrendRow, error) {
return analyticsrepo.IPTrendNodeAccessLogs(ctx, toNodeAccessLogFilter(filter), bucketSeconds)
}
func (clickhouseAccessLogStore) DeleteAll(ctx context.Context) (int64, error) {
return analyticsrepo.DeleteAllNodeAccessLogs(ctx)
}
func (clickhouseAccessLogStore) DeleteBefore(ctx context.Context, cutoff time.Time) (int64, error) {
return analyticsrepo.DeleteNodeAccessLogsBefore(ctx, cutoff)
}
func (clickhouseAccessLogStore) DeleteByNodeBefore(ctx context.Context, nodeID string, before time.Time) (int64, error) {
return analyticsrepo.DeleteNodeAccessLogsByNodeBefore(ctx, nodeID, before)
}
func (clickhouseAccessLogStore) TrafficSummary(ctx context.Context, filter OpenFlareAccessLogQuery) (OpenFlareAccessLogTrafficSummary, error) {
row, err := analyticsrepo.TrafficSummaryNodeAccessLogs(ctx, toNodeAccessLogFilter(filter))
if err != nil {
return OpenFlareAccessLogTrafficSummary{}, err
}
return OpenFlareAccessLogTrafficSummary{
RequestCount: row.RequestCount,
ErrorCount: row.ErrorCount,
UniqueIPCount: row.UniqueIPCount,
BytesSent: row.BytesSent,
RequestLength: row.RequestLength,
NodeCount: row.NodeCount,
}, nil
}
func (clickhouseAccessLogStore) ValueCounts(ctx context.Context, filter OpenFlareAccessLogQuery, column string, limit int) ([]OpenFlareAccessLogValueCount, error) {
rows, err := analyticsrepo.ValueCountsNodeAccessLogs(ctx, toNodeAccessLogFilter(filter), column, limit)
if err != nil {
return nil, err
}
result := make([]OpenFlareAccessLogValueCount, len(rows))
for i, row := range rows {
result[i] = OpenFlareAccessLogValueCount{Value: row.Value, Count: row.Count}
}
return result, nil
}
func (clickhouseAccessLogStore) NodeAggregates(ctx context.Context, filter OpenFlareAccessLogQuery) ([]OpenFlareAccessLogNodeAggregate, error) {
rows, err := analyticsrepo.NodeAggregatesNodeAccessLogs(ctx, toNodeAccessLogFilter(filter))
if err != nil {
return nil, err
}
result := make([]OpenFlareAccessLogNodeAggregate, len(rows))
for i, row := range rows {
result[i] = OpenFlareAccessLogNodeAggregate{
NodeID: row.NodeID,
RequestCount: row.RequestCount,
ErrorCount: row.ErrorCount,
UniqueIPCount: row.UniqueIPCount,
}
}
return result, nil
}
func toNodeAccessLogFilter(query OpenFlareAccessLogQuery) analyticsrepo.NodeAccessLogFilter {
return analyticsrepo.NodeAccessLogFilter{
NodeID: query.NodeID,
RemoteAddr: query.RemoteAddr,
Host: query.Host,
Hosts: query.Hosts,
Path: query.Path,
Since: query.Since,
Until: query.Until,
Page: query.Page,
PageSize: query.PageSize,
SortBy: query.SortBy,
SortOrder: query.SortOrder,
}
}
func toAnalyticsNodeAccessLog(record *OpenFlareAccessLog) analyticsmodel.NodeAccessLog {
var bytesSent uint64
if record.BytesSent > 0 {
bytesSent = uint64(record.BytesSent)
}
var requestLength uint64
if record.RequestLength > 0 {
requestLength = uint64(record.RequestLength)
}
var requestTimeMs uint32
if record.RequestTimeMs > 0 && record.RequestTimeMs <= int64(math.MaxUint32) {
requestTimeMs = uint32(record.RequestTimeMs)
}
return analyticsmodel.NodeAccessLog{
ID: record.ID,
NodeID: record.NodeID,
LoggedAt: record.LoggedAt,
RemoteAddr: record.RemoteAddr,
Region: record.Region,
Host: record.Host,
Path: record.Path,
UserAgent: record.UserAgent,
CacheStatus: record.CacheStatus,
StatusCode: openFlareAccessLogStatusCodeToInt32(record.StatusCode),
BytesSent: bytesSent,
RequestLength: requestLength,
RequestTimeMs: requestTimeMs,
CreatedAt: record.CreatedAt,
}
}
func fromAnalyticsNodeAccessLogs(rows []analyticsmodel.NodeAccessLog) []*OpenFlareAccessLog {
result := make([]*OpenFlareAccessLog, len(rows))
for index, row := range rows {
var bytesSent int64
if row.BytesSent <= math.MaxInt64 {
bytesSent = int64(row.BytesSent)
} else {
bytesSent = math.MaxInt64
}
var requestLength int64
if row.RequestLength <= math.MaxInt64 {
requestLength = int64(row.RequestLength)
} else {
requestLength = math.MaxInt64
}
result[index] = &OpenFlareAccessLog{
ID: row.ID,
NodeID: row.NodeID,
LoggedAt: row.LoggedAt,
RemoteAddr: row.RemoteAddr,
Region: row.Region,
Host: row.Host,
Path: row.Path,
UserAgent: row.UserAgent,
CacheStatus: row.CacheStatus,
StatusCode: int(row.StatusCode),
BytesSent: bytesSent,
RequestLength: requestLength,
RequestTimeMs: int64(row.RequestTimeMs),
CreatedAt: row.CreatedAt,
}
}
return result
}
@@ -1,708 +0,0 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package model
import (
"context"
"net"
"net/http"
"net/netip"
"sort"
"strconv"
"strings"
"sync"
"time"
"github.com/Rain-kl/Wavelet/internal/infra/persistence/idgen"
)
const (
accessLogColumnStatusCode = "status_code"
accessLogColumnHost = "host"
accessLogColumnPath = "path"
accessLogColumnRemoteAddr = "remote_addr"
accessLogColumnUserAgent = "user_agent"
)
type memoryAccessLogStore struct {
mu sync.RWMutex
records []*OpenFlareAccessLog
}
func (s *memoryAccessLogStore) InsertBatch(_ context.Context, records []*OpenFlareAccessLog) error {
s.mu.Lock()
defer s.mu.Unlock()
now := time.Now().UTC()
for _, record := range records {
if record == nil {
continue
}
copyRecord := *record
if copyRecord.ID == 0 {
copyRecord.ID = idgen.NextUint64ID()
}
if copyRecord.CreatedAt.IsZero() {
copyRecord.CreatedAt = now
}
copyRecord.LoggedAt = copyRecord.LoggedAt.UTC()
copyRecord.CreatedAt = copyRecord.CreatedAt.UTC()
s.records = append(s.records, &copyRecord)
}
return nil
}
func (s *memoryAccessLogStore) List(_ context.Context, query OpenFlareAccessLogQuery) ([]*OpenFlareAccessLog, error) {
s.mu.RLock()
defer s.mu.RUnlock()
rows := s.filterRecords(query)
sortOpenFlareAccessLogRows(rows, query.SortBy, query.SortOrder)
if query.PageSize > 0 {
start, end := openFlareAccessLogPaginateBounds(len(rows), query.Page, query.PageSize)
return cloneAccessLogSlice(rows[start:end]), nil
}
return cloneAccessLogSlice(rows), nil
}
func (s *memoryAccessLogStore) Count(_ context.Context, query OpenFlareAccessLogQuery) (int64, int64, int64, error) {
s.mu.RLock()
defer s.mu.RUnlock()
rows := s.filterRecords(query)
ips := make(map[string]struct{})
var totalBytes int64
for _, row := range rows {
totalBytes += row.BytesSent
remoteAddr := strings.TrimSpace(row.RemoteAddr)
if remoteAddr == "" {
continue
}
ips[remoteAddr] = struct{}{}
}
return int64(len(rows)), int64(len(ips)), totalBytes, nil
}
func (s *memoryAccessLogStore) RegionCounts(_ context.Context, nodeID string, since time.Time, limit int) ([]*OpenFlareAccessLogRegionCount, error) {
s.mu.RLock()
defer s.mu.RUnlock()
rows := s.filterRecords(OpenFlareAccessLogQuery{NodeID: nodeID, Since: since})
counts := make(map[string]int64)
for _, row := range rows {
region := strings.TrimSpace(row.Region)
if region == "" {
continue
}
counts[region]++
}
result := make([]*OpenFlareAccessLogRegionCount, 0, len(counts))
for region, count := range counts {
result = append(result, &OpenFlareAccessLogRegionCount{Region: region, Count: count})
}
sort.Slice(result, func(i, j int) bool {
if result[i].Count == result[j].Count {
return result[i].Region < result[j].Region
}
return result[i].Count > result[j].Count
})
if limit > 0 && len(result) > limit {
result = result[:limit]
}
return result, nil
}
func (s *memoryAccessLogStore) BucketAggregates(_ context.Context, filter OpenFlareAccessLogQuery, bucketSeconds int64) ([]openFlareAccessLogBucketAggregateRow, error) {
s.mu.RLock()
defer s.mu.RUnlock()
rows := s.filterRecords(filter)
type bucketAccumulator struct {
openFlareAccessLogBucketAggregateRow
uniqueIPs map[string]struct{}
uniqueHosts map[string]struct{}
}
aggregates := make(map[int64]*bucketAccumulator)
for _, row := range rows {
bucketEpoch := memoryAccessLogBucketEpoch(row.LoggedAt, bucketSeconds)
item := aggregates[bucketEpoch]
if item == nil {
item = &bucketAccumulator{
openFlareAccessLogBucketAggregateRow: openFlareAccessLogBucketAggregateRow{BucketEpoch: bucketEpoch},
uniqueIPs: make(map[string]struct{}),
uniqueHosts: make(map[string]struct{}),
}
aggregates[bucketEpoch] = item
}
item.RequestCount++
item.BytesSent += row.BytesSent
item.RequestLength += row.RequestLength
switch {
case row.StatusCode < 400:
item.SuccessCount++
case row.StatusCode < 500:
item.ClientErrorCount++
default:
item.ServerErrorCount++
}
if remoteAddr := strings.TrimSpace(row.RemoteAddr); remoteAddr != "" {
item.uniqueIPs[remoteAddr] = struct{}{}
}
if host := strings.TrimSpace(row.Host); host != "" {
item.uniqueHosts[host] = struct{}{}
}
}
result := make([]openFlareAccessLogBucketAggregateRow, 0, len(aggregates))
for _, item := range aggregates {
item.UniqueIPCount = int64(len(item.uniqueIPs))
item.UniqueHostCount = int64(len(item.uniqueHosts))
result = append(result, item.openFlareAccessLogBucketAggregateRow)
}
bucketRows := make([]*OpenFlareAccessLogBucketRow, len(result))
for index := range result {
bucketRows[index] = &OpenFlareAccessLogBucketRow{
BucketEpoch: result[index].BucketEpoch,
RequestCount: result[index].RequestCount,
UniqueIPCount: result[index].UniqueIPCount,
UniqueHostCount: result[index].UniqueHostCount,
SuccessCount: result[index].SuccessCount,
ClientErrorCount: result[index].ClientErrorCount,
ServerErrorCount: result[index].ServerErrorCount,
BytesSent: result[index].BytesSent,
RequestLength: result[index].RequestLength,
}
}
sortOpenFlareAccessLogBucketRows(bucketRows, filter.SortBy, filter.SortOrder)
for index := range result {
result[index] = openFlareAccessLogBucketAggregateRow{
BucketEpoch: bucketRows[index].BucketEpoch,
RequestCount: bucketRows[index].RequestCount,
UniqueIPCount: bucketRows[index].UniqueIPCount,
UniqueHostCount: bucketRows[index].UniqueHostCount,
SuccessCount: bucketRows[index].SuccessCount,
ClientErrorCount: bucketRows[index].ClientErrorCount,
ServerErrorCount: bucketRows[index].ServerErrorCount,
BytesSent: bucketRows[index].BytesSent,
RequestLength: bucketRows[index].RequestLength,
}
}
if filter.PageSize > 0 {
start, end := openFlareAccessLogPaginateBounds(len(result), filter.Page, filter.PageSize)
return result[start:end], nil
}
return result, nil
}
func (s *memoryAccessLogStore) CountBuckets(_ context.Context, filter OpenFlareAccessLogQuery, bucketSeconds int64) (int64, error) {
s.mu.RLock()
defer s.mu.RUnlock()
rows := s.filterRecords(filter)
seen := make(map[int64]struct{})
for _, row := range rows {
seen[memoryAccessLogBucketEpoch(row.LoggedAt, bucketSeconds)] = struct{}{}
}
return int64(len(seen)), nil
}
func (s *memoryAccessLogStore) BucketDimensions(_ context.Context, filter OpenFlareAccessLogQuery, column string, bucketSeconds int64) ([]openFlareAccessLogBucketDimensionRow, error) {
s.mu.RLock()
defer s.mu.RUnlock()
rows := s.filterRecords(filter)
seen := make(map[int64]map[string]struct{})
var result []openFlareAccessLogBucketDimensionRow
for _, row := range rows {
var value string
switch column {
case columnRemoteAddr:
value = strings.TrimSpace(row.RemoteAddr)
case columnHost:
value = strings.TrimSpace(row.Host)
default:
continue
}
if value == "" {
continue
}
bucketEpoch := memoryAccessLogBucketEpoch(row.LoggedAt, bucketSeconds)
if seen[bucketEpoch] == nil {
seen[bucketEpoch] = make(map[string]struct{})
}
if _, ok := seen[bucketEpoch][value]; ok {
continue
}
seen[bucketEpoch][value] = struct{}{}
result = append(result, openFlareAccessLogBucketDimensionRow{BucketEpoch: bucketEpoch, Value: value})
}
return result, nil
}
func (s *memoryAccessLogStore) IPAggregates(_ context.Context, filter OpenFlareAccessLogQuery, exactRemoteAddr bool) ([]openFlareAccessLogIPAggregateRow, error) {
s.mu.RLock()
defer s.mu.RUnlock()
if exactRemoteAddr && strings.TrimSpace(filter.RemoteAddr) == "" {
return []openFlareAccessLogIPAggregateRow{}, nil
}
rows := s.filterRecords(filter)
aggregates := make(map[string]*openFlareAccessLogIPAggregateRow)
for _, row := range rows {
remoteAddr := strings.TrimSpace(row.RemoteAddr)
if remoteAddr == "" {
continue
}
if exactRemoteAddr && remoteAddr != strings.TrimSpace(filter.RemoteAddr) {
continue
}
item := aggregates[remoteAddr]
if item == nil {
item = &openFlareAccessLogIPAggregateRow{RemoteAddr: remoteAddr}
aggregates[remoteAddr] = item
}
item.RequestCount++
epoch := row.LoggedAt.UTC().Unix()
if epoch > item.LastSeenEpoch {
item.LastSeenEpoch = epoch
}
switch {
case row.StatusCode < 400:
item.SuccessCount++
case row.StatusCode < 500:
item.ClientErrorCount++
default:
item.ServerErrorCount++
}
}
result := make([]openFlareAccessLogIPAggregateRow, 0, len(aggregates))
for _, item := range aggregates {
result = append(result, *item)
}
return result, nil
}
func (s *memoryAccessLogStore) IPSummaries(_ context.Context, filter OpenFlareAccessLogQuery, _ time.Time) ([]openFlareAccessLogIPSummaryRow, error) {
s.mu.RLock()
defer s.mu.RUnlock()
rows := s.filterRecords(filter)
type aggregate struct {
RemoteAddr string
Region string
RegionEpoch int64
TotalRequests int64
Success2xxCount int64
BytesReceived int64
BytesSent int64
LastSeenEpoch int64
}
aggregates := make(map[string]*aggregate)
for _, row := range rows {
remoteAddr := strings.TrimSpace(row.RemoteAddr)
if remoteAddr == "" {
continue
}
item := aggregates[remoteAddr]
if item == nil {
item = &aggregate{RemoteAddr: remoteAddr}
aggregates[remoteAddr] = item
}
item.TotalRequests++
if row.StatusCode >= 200 && row.StatusCode < 300 {
item.Success2xxCount++
}
item.BytesReceived += row.RequestLength
item.BytesSent += row.BytesSent
epoch := row.LoggedAt.UTC().Unix()
if epoch > item.LastSeenEpoch {
item.LastSeenEpoch = epoch
}
if epoch >= item.RegionEpoch {
item.RegionEpoch = epoch
item.Region = strings.TrimSpace(row.Region)
}
}
summaryRows := make([]*OpenFlareAccessLogIPSummaryRow, 0, len(aggregates))
for _, item := range aggregates {
ratio := 0.0
if item.TotalRequests > 0 {
ratio = float64(item.Success2xxCount) / float64(item.TotalRequests)
}
summaryRows = append(summaryRows, &OpenFlareAccessLogIPSummaryRow{
RemoteAddr: item.RemoteAddr,
Region: item.Region,
TotalRequests: item.TotalRequests,
Success2xxCount: item.Success2xxCount,
SuccessRatio: ratio,
BytesReceived: item.BytesReceived,
BytesSent: item.BytesSent,
RecentRequests: 0,
LastSeenEpoch: item.LastSeenEpoch,
})
}
sortOpenFlareAccessLogIPSummaryRows(summaryRows, filter.SortBy, filter.SortOrder)
if filter.PageSize > 0 {
start, end := openFlareAccessLogPaginateBounds(len(summaryRows), filter.Page, filter.PageSize)
summaryRows = summaryRows[start:end]
}
result := make([]openFlareAccessLogIPSummaryRow, len(summaryRows))
for index, item := range summaryRows {
result[index] = openFlareAccessLogIPSummaryRow{
RemoteAddr: item.RemoteAddr,
Region: item.Region,
TotalRequests: item.TotalRequests,
Success2xxCount: item.Success2xxCount,
SuccessRatio: item.SuccessRatio,
BytesReceived: item.BytesReceived,
BytesSent: item.BytesSent,
RecentRequests: 0,
LastSeenEpoch: item.LastSeenEpoch,
}
}
return result, nil
}
func (s *memoryAccessLogStore) CountIPSummaries(_ context.Context, filter OpenFlareAccessLogQuery) (int64, error) {
s.mu.RLock()
defer s.mu.RUnlock()
rows := s.filterRecords(filter)
seen := make(map[string]struct{})
for _, row := range rows {
remoteAddr := strings.TrimSpace(row.RemoteAddr)
if remoteAddr == "" {
continue
}
seen[remoteAddr] = struct{}{}
}
return int64(len(seen)), nil
}
func (s *memoryAccessLogStore) WAFIPAggregates(_ context.Context, filter OpenFlareAccessLogQuery) ([]openFlareAccessLogWAFIPAggregateRow, error) {
s.mu.RLock()
defer s.mu.RUnlock()
rows := s.filterRecords(filter)
aggregates := make(map[string]*openFlareAccessLogWAFIPAggregateRow)
order := make([]string, 0)
for _, row := range rows {
remoteAddr := strings.TrimSpace(row.RemoteAddr)
if remoteAddr == "" {
continue
}
item := aggregates[remoteAddr]
if item == nil {
item = &openFlareAccessLogWAFIPAggregateRow{
RemoteAddr: remoteAddr,
StatusCounts: make(map[int]int64),
}
aggregates[remoteAddr] = item
order = append(order, remoteAddr)
}
item.RequestCount++
item.StatusCounts[row.StatusCode]++
if row.StatusCode == http.StatusNotFound {
item.Status404Count++
}
if row.StatusCode >= 400 && row.StatusCode < 500 {
item.ClientErrorCount++
}
if row.StatusCode >= http.StatusInternalServerError {
item.ServerErrorCount++
}
if memoryAccessLogHostIsIPLiteral(row.Host) {
item.IPHostCount++
}
epoch := row.LoggedAt.UTC().Unix()
if epoch > item.LastSeenEpoch {
item.LastSeenEpoch = epoch
}
}
result := make([]openFlareAccessLogWAFIPAggregateRow, 0, len(order))
for _, remoteAddr := range order {
if item := aggregates[remoteAddr]; item != nil {
result = append(result, *item)
}
}
return result, nil
}
func (s *memoryAccessLogStore) IPTrend(_ context.Context, filter OpenFlareAccessLogQuery, bucketSeconds int64) ([]openFlareAccessLogIPTrendRow, error) {
s.mu.RLock()
defer s.mu.RUnlock()
rows := s.filterRecords(filter)
aggregates := make(map[int64]int64)
for _, row := range rows {
bucketEpoch := memoryAccessLogBucketEpoch(row.LoggedAt, bucketSeconds)
aggregates[bucketEpoch]++
}
result := make([]openFlareAccessLogIPTrendRow, 0, len(aggregates))
for bucketEpoch, count := range aggregates {
result = append(result, openFlareAccessLogIPTrendRow{BucketEpoch: bucketEpoch, RequestCount: count})
}
sort.Slice(result, func(i, j int) bool { return result[i].BucketEpoch < result[j].BucketEpoch })
return result, nil
}
func (s *memoryAccessLogStore) DeleteAll(_ context.Context) (int64, error) {
s.mu.Lock()
defer s.mu.Unlock()
count := int64(len(s.records))
s.records = nil
return count, nil
}
func (s *memoryAccessLogStore) DeleteBefore(_ context.Context, cutoff time.Time) (int64, error) {
s.mu.Lock()
defer s.mu.Unlock()
cutoff = cutoff.UTC()
remaining := make([]*OpenFlareAccessLog, 0, len(s.records))
var deleted int64
for _, row := range s.records {
if row.LoggedAt.Before(cutoff) {
deleted++
continue
}
remaining = append(remaining, row)
}
s.records = remaining
return deleted, nil
}
func (s *memoryAccessLogStore) DeleteByNodeBefore(_ context.Context, nodeID string, before time.Time) (int64, error) {
s.mu.Lock()
defer s.mu.Unlock()
before = before.UTC()
remaining := make([]*OpenFlareAccessLog, 0, len(s.records))
var deleted int64
for _, row := range s.records {
if row.NodeID == nodeID && row.LoggedAt.Before(before) {
deleted++
continue
}
remaining = append(remaining, row)
}
s.records = remaining
return deleted, nil
}
func (s *memoryAccessLogStore) TrafficSummary(_ context.Context, filter OpenFlareAccessLogQuery) (OpenFlareAccessLogTrafficSummary, error) {
s.mu.RLock()
defer s.mu.RUnlock()
rows := s.filterRecords(filter)
ips := make(map[string]struct{})
nodes := make(map[string]struct{})
var summary OpenFlareAccessLogTrafficSummary
for _, row := range rows {
summary.RequestCount++
summary.BytesSent += row.BytesSent
summary.RequestLength += row.RequestLength
if row.StatusCode >= http.StatusInternalServerError {
summary.ErrorCount++
}
if ip := strings.TrimSpace(row.RemoteAddr); ip != "" {
ips[ip] = struct{}{}
}
if id := strings.TrimSpace(row.NodeID); id != "" {
nodes[id] = struct{}{}
}
}
summary.UniqueIPCount = int64(len(ips))
summary.NodeCount = int64(len(nodes))
return summary, nil
}
func (s *memoryAccessLogStore) ValueCounts(_ context.Context, filter OpenFlareAccessLogQuery, column string, limit int) ([]OpenFlareAccessLogValueCount, error) {
s.mu.RLock()
defer s.mu.RUnlock()
col := strings.TrimSpace(strings.ToLower(column))
switch col {
case accessLogColumnStatusCode, accessLogColumnHost, accessLogColumnPath, accessLogColumnRemoteAddr, accessLogColumnUserAgent:
default:
return nil, nil
}
rows := s.filterRecords(filter)
counts := make(map[string]int64)
for _, row := range rows {
var value string
switch col {
case accessLogColumnStatusCode:
value = strconv.Itoa(row.StatusCode)
case accessLogColumnHost:
value = strings.TrimSpace(row.Host)
case accessLogColumnPath:
value = strings.TrimSpace(row.Path)
case accessLogColumnRemoteAddr:
value = strings.TrimSpace(row.RemoteAddr)
case accessLogColumnUserAgent:
value = strings.TrimSpace(row.UserAgent)
}
if value == "" {
continue
}
counts[value]++
}
result := make([]OpenFlareAccessLogValueCount, 0, len(counts))
for value, count := range counts {
result = append(result, OpenFlareAccessLogValueCount{Value: value, Count: count})
}
sort.Slice(result, func(i, j int) bool {
if result[i].Count == result[j].Count {
return result[i].Value < result[j].Value
}
return result[i].Count > result[j].Count
})
if limit > 0 && len(result) > limit {
result = result[:limit]
}
return result, nil
}
func (s *memoryAccessLogStore) NodeAggregates(_ context.Context, filter OpenFlareAccessLogQuery) ([]OpenFlareAccessLogNodeAggregate, error) {
s.mu.RLock()
defer s.mu.RUnlock()
rows := s.filterRecords(filter)
type acc struct {
OpenFlareAccessLogNodeAggregate
ips map[string]struct{}
}
byNode := make(map[string]*acc)
for _, row := range rows {
id := strings.TrimSpace(row.NodeID)
if id == "" {
continue
}
item := byNode[id]
if item == nil {
item = &acc{
OpenFlareAccessLogNodeAggregate: OpenFlareAccessLogNodeAggregate{NodeID: id},
ips: make(map[string]struct{}),
}
byNode[id] = item
}
item.RequestCount++
if row.StatusCode >= http.StatusInternalServerError {
item.ErrorCount++
}
if ip := strings.TrimSpace(row.RemoteAddr); ip != "" {
item.ips[ip] = struct{}{}
}
}
result := make([]OpenFlareAccessLogNodeAggregate, 0, len(byNode))
for _, item := range byNode {
item.UniqueIPCount = int64(len(item.ips))
result = append(result, item.OpenFlareAccessLogNodeAggregate)
}
sort.Slice(result, func(i, j int) bool {
if result[i].RequestCount == result[j].RequestCount {
return result[i].NodeID < result[j].NodeID
}
return result[i].RequestCount > result[j].RequestCount
})
return result, nil
}
func (s *memoryAccessLogStore) filterRecords(query OpenFlareAccessLogQuery) []*OpenFlareAccessLog {
result := make([]*OpenFlareAccessLog, 0, len(s.records))
for _, row := range s.records {
if !memoryAccessLogMatches(row, query) {
continue
}
result = append(result, row)
}
return result
}
func memoryAccessLogMatches(row *OpenFlareAccessLog, query OpenFlareAccessLogQuery) bool {
if row == nil {
return false
}
if trimmed := strings.TrimSpace(query.NodeID); trimmed != "" && row.NodeID != trimmed {
return false
}
if trimmed := strings.TrimSpace(query.RemoteAddr); trimmed != "" && !strings.HasPrefix(strings.TrimSpace(row.RemoteAddr), trimmed) {
return false
}
if len(query.Hosts) > 0 {
rowHost := strings.ToLower(strings.TrimSpace(row.Host))
matched := false
for _, host := range query.Hosts {
if strings.ToLower(strings.TrimSpace(host)) == rowHost {
matched = true
break
}
}
if !matched {
return false
}
} else if trimmed := strings.TrimSpace(query.Host); trimmed != "" && !strings.HasPrefix(strings.TrimSpace(row.Host), trimmed) {
return false
}
if trimmed := strings.TrimSpace(query.Path); trimmed != "" && !strings.HasPrefix(strings.TrimSpace(row.Path), trimmed) {
return false
}
if !query.Since.IsZero() && row.LoggedAt.Before(query.Since) {
return false
}
if !query.Until.IsZero() && !row.LoggedAt.Before(query.Until) {
return false
}
return true
}
func memoryAccessLogHostIsIPLiteral(value string) bool {
host := strings.TrimSpace(value)
if host == "" {
return false
}
if parsedHost, _, err := net.SplitHostPort(host); err == nil {
host = parsedHost
}
host = strings.Trim(host, "[]")
_, err := netip.ParseAddr(host)
return err == nil
}
func memoryAccessLogBucketEpoch(loggedAt time.Time, bucketSeconds int64) int64 {
if bucketSeconds <= 0 {
bucketSeconds = 180
}
epoch := loggedAt.UTC().Unix()
return (epoch / bucketSeconds) * bucketSeconds
}
func cloneAccessLogSlice(rows []*OpenFlareAccessLog) []*OpenFlareAccessLog {
result := make([]*OpenFlareAccessLog, len(rows))
for index, row := range rows {
if row == nil {
continue
}
copyRecord := *row
result[index] = &copyRecord
}
return result
}
func sortOpenFlareAccessLogRows(items []*OpenFlareAccessLog, sortBy string, sortOrder string) {
desc := openFlareAccessLogNormalizeSortOrder(sortOrder) != sortOrderAsc
sort.Slice(items, func(i, j int) bool {
left := items[i]
right := items[j]
if left == nil || right == nil {
return left != nil
}
var compare int
switch strings.TrimSpace(sortBy) {
case "status_code":
compare = left.StatusCode - right.StatusCode
case columnRemoteAddr:
compare = strings.Compare(left.RemoteAddr, right.RemoteAddr)
case columnHost:
compare = strings.Compare(left.Host, right.Host)
case "path":
compare = strings.Compare(left.Path, right.Path)
default:
compare = openFlareAccessLogCompareInt64(left.LoggedAt.Unix(), right.LoggedAt.Unix())
}
if compare == 0 {
compare = openFlareAccessLogCompareInt64(left.LoggedAt.Unix(), right.LoggedAt.Unix())
}
if compare == 0 {
compare = openFlareAccessLogCompareInt64(openFlareAccessLogUintToInt64(left.ID), openFlareAccessLogUintToInt64(right.ID))
}
if desc {
return compare > 0
}
return compare < 0
})
}
-119
View File
@@ -1,119 +0,0 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package model
import (
"context"
"fmt"
"testing"
"time"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
func setupOpenFlareAccessLogTestEnvironment(t *testing.T) (context.Context, func()) {
t.Helper()
store := NewMemoryAccessLogStore()
reset := SetAccessLogStoreForTest(store)
return context.Background(), func() {
reset()
}
}
func seedOpenFlareAccessLogs(t *testing.T, ctx context.Context, now time.Time) {
t.Helper()
records := []*OpenFlareAccessLog{
{NodeID: "node-a", LoggedAt: now.Add(-5 * time.Minute), RemoteAddr: "1.1.1.1", Region: "US", Host: "a.example.com", Path: "/alpha", StatusCode: 200},
{NodeID: "node-a", LoggedAt: now.Add(-4 * time.Minute), RemoteAddr: "2.2.2.2", Region: "US", Host: "a.example.com", Path: "/beta", StatusCode: 404},
{NodeID: "node-b", LoggedAt: now.Add(-3 * time.Minute), RemoteAddr: "1.1.1.1", Region: "EU", Host: "b.example.com", Path: "/gamma", StatusCode: 502},
{NodeID: "node-b", LoggedAt: now.Add(-2 * time.Minute), RemoteAddr: "3.3.3.3", Region: "EU", Host: "b.example.com", Path: "/delta", StatusCode: 200},
{NodeID: "node-b", LoggedAt: now.Add(-1 * time.Minute), RemoteAddr: "", Region: "", Host: "b.example.com", Path: "/empty-ip", StatusCode: 200},
}
require.NoError(t, InsertOpenFlareAccessLogsBatch(ctx, records))
}
func TestListOpenFlareAccessLogsPaginated(t *testing.T) {
ctx, cleanup := setupOpenFlareAccessLogTestEnvironment(t)
defer cleanup()
now := time.Now().UTC()
for index := range 15 {
record := &OpenFlareAccessLog{
NodeID: "node-page",
LoggedAt: now.Add(-time.Duration(index) * time.Minute),
RemoteAddr: fmt.Sprintf("203.0.113.%d", (index%5)+1),
Host: "example.com",
Path: fmt.Sprintf("/path-%02d", index),
StatusCode: 200,
}
require.NoError(t, InsertOpenFlareAccessLogsBatch(ctx, []*OpenFlareAccessLog{record}))
}
query := OpenFlareAccessLogQuery{
NodeID: "node-page",
Since: now.Add(-24 * time.Hour),
Page: 1,
PageSize: 5,
SortBy: "logged_at",
SortOrder: "desc",
}
page, err := ListOpenFlareAccessLogs(ctx, query)
require.NoError(t, err)
require.Len(t, page, 5)
assert.Equal(t, "/path-05", page[0].Path)
assert.Equal(t, "/path-09", page[4].Path)
}
func TestCountOpenFlareAccessLogs(t *testing.T) {
ctx, cleanup := setupOpenFlareAccessLogTestEnvironment(t)
defer cleanup()
now := time.Now().UTC()
seedOpenFlareAccessLogs(t, ctx, now)
query := OpenFlareAccessLogQuery{
Since: now.Add(-10 * time.Minute),
}
totalRecords, totalIPs, _, err := CountOpenFlareAccessLogs(ctx, query)
require.NoError(t, err)
assert.Equal(t, int64(5), totalRecords)
assert.Equal(t, int64(3), totalIPs)
}
func TestListOpenFlareAccessLogsFiltersAndSort(t *testing.T) {
ctx, cleanup := setupOpenFlareAccessLogTestEnvironment(t)
defer cleanup()
now := time.Now().UTC()
seedOpenFlareAccessLogs(t, ctx, now)
query := OpenFlareAccessLogQuery{
NodeID: "node-a",
Since: now.Add(-10 * time.Minute),
SortBy: "status_code",
SortOrder: "desc",
}
rows, err := ListOpenFlareAccessLogs(ctx, query)
require.NoError(t, err)
require.Len(t, rows, 2)
assert.Equal(t, 404, rows[0].StatusCode)
assert.Equal(t, 200, rows[1].StatusCode)
}
func TestDeleteOpenFlareAccessLogsBefore(t *testing.T) {
ctx, cleanup := setupOpenFlareAccessLogTestEnvironment(t)
defer cleanup()
now := time.Now().UTC()
seedOpenFlareAccessLogs(t, ctx, now)
deleted, err := DeleteOpenFlareAccessLogsBefore(ctx, now.Add(-2*time.Minute))
require.NoError(t, err)
assert.Equal(t, int64(3), deleted)
totalRecords, _, _, err := CountOpenFlareAccessLogs(ctx, OpenFlareAccessLogQuery{Since: now.Add(-10 * time.Minute)})
require.NoError(t, err)
assert.Equal(t, int64(2), totalRecords)
}
-59
View File
@@ -4,12 +4,7 @@
package model
import (
"context"
"errors"
"time"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
"gorm.io/gorm"
)
// AcmeAccount OpenFlare ACME 账号实体。
@@ -26,57 +21,3 @@ type AcmeAccount struct {
func (AcmeAccount) TableName() string {
return "of_acme_accounts"
}
// GetAcmeAccountByID 按 ID 查询 ACME 账号。
func GetAcmeAccountByID(ctx context.Context, id uint) (*AcmeAccount, error) {
conn := db.DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
var account AcmeAccount
if err := conn.First(&account, id).Error; err != nil {
return nil, err
}
return &account, nil
}
// CreateAcmeAccountRecord 创建 ACME 账号。
func CreateAcmeAccountRecord(ctx context.Context, account *AcmeAccount) error {
conn := db.DB(ctx)
if conn == nil {
return errors.New(errDatabaseNotInitialized)
}
return conn.Create(account).Error
}
// SaveAcmeAccount 保存 ACME 账号。
func SaveAcmeAccount(ctx context.Context, account *AcmeAccount) error {
conn := db.DB(ctx)
if conn == nil {
return errors.New(errDatabaseNotInitialized)
}
return conn.Save(account).Error
}
// GetDefaultAcmeAccount 获取默认 ACME 账号,不存在时创建占位记录。
func GetDefaultAcmeAccount(ctx context.Context) (*AcmeAccount, error) {
conn := db.DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
var account AcmeAccount
err := conn.Order("id asc").First(&account).Error
if err == nil {
return &account, nil
}
if !errors.Is(err, gorm.ErrRecordNotFound) {
return nil, err
}
account = AcmeAccount{
Email: "admin@openflare.dev",
}
if err = conn.Create(&account).Error; err != nil {
return nil, err
}
return &account, nil
}
-121
View File
@@ -4,13 +4,8 @@
package model
import (
"context"
"errors"
"strings"
"time"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
"gorm.io/gorm"
)
// OpenFlareApplyLogQuery filters apply logs for list queries.
@@ -39,74 +34,6 @@ func (OpenFlareApplyLog) TableName() string {
return "of_apply_logs"
}
// ListOpenFlareApplyLogs returns apply logs ordered by id desc with optional pagination.
func ListOpenFlareApplyLogs(ctx context.Context, query OpenFlareApplyLogQuery) ([]*OpenFlareApplyLog, error) {
conn := db.DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
dbQuery := conn.Model(&OpenFlareApplyLog{}).Order("id desc")
if query.NodeID != "" {
dbQuery = dbQuery.Where("node_id = ?", query.NodeID)
}
if query.PageSize > 0 {
offset := 0
if query.PageNo > 1 {
offset = (query.PageNo - 1) * query.PageSize
}
dbQuery = dbQuery.Limit(query.PageSize).Offset(offset)
}
var logs []*OpenFlareApplyLog
if err := dbQuery.Find(&logs).Error; err != nil {
return nil, err
}
return logs, nil
}
// CountOpenFlareApplyLogs returns total apply logs, optionally filtered by node_id.
func CountOpenFlareApplyLogs(ctx context.Context, nodeID string) (int64, error) {
conn := db.DB(ctx)
if conn == nil {
return 0, errors.New(errDatabaseNotInitialized)
}
query := conn.Model(&OpenFlareApplyLog{})
if nodeID != "" {
query = query.Where("node_id = ?", nodeID)
}
var total int64
if err := query.Count(&total).Error; err != nil {
return 0, err
}
return total, nil
}
// GetLatestOpenFlareApplyLogByNodeID returns the most recent apply log for a node.
func GetLatestOpenFlareApplyLogByNodeID(ctx context.Context, nodeID string) (*OpenFlareApplyLog, error) {
nodeID = strings.TrimSpace(nodeID)
if nodeID == "" {
return nil, errors.New("node_id is required")
}
conn := db.DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
var log OpenFlareApplyLog
err := conn.Where("node_id = ?", nodeID).Order("id desc").First(&log).Error
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, nil
}
if err != nil {
return nil, err
}
return &log, nil
}
// IsRepeatSuccessApplyLog reports whether the payload repeats an already-recorded success entry.
func IsRepeatSuccessApplyLog(latest *OpenFlareApplyLog, version, checksum, result string) bool {
if latest == nil || result != "success" {
@@ -116,51 +43,3 @@ func IsRepeatSuccessApplyLog(latest *OpenFlareApplyLog, version, checksum, resul
strings.TrimSpace(latest.Version) == strings.TrimSpace(version) &&
strings.TrimSpace(latest.Checksum) == strings.TrimSpace(checksum)
}
// GetLatestOpenFlareApplyLogsByNodeIDs returns the latest apply log per node id.
func GetLatestOpenFlareApplyLogsByNodeIDs(ctx context.Context, nodeIDs []string) (map[string]*OpenFlareApplyLog, error) {
result := make(map[string]*OpenFlareApplyLog)
if len(nodeIDs) == 0 {
return result, nil
}
conn := db.DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
var logs []*OpenFlareApplyLog
subQuery := conn.Model(&OpenFlareApplyLog{}).
Select("MAX(id) AS id").
Where("node_id IN ?", nodeIDs).
Group("node_id")
if err := conn.Where("id IN (?)", subQuery).Find(&logs).Error; err != nil {
return nil, err
}
for _, log := range logs {
result[log.NodeID] = log
}
return result, nil
}
// DeleteAllOpenFlareApplyLogs removes every apply log record.
func DeleteAllOpenFlareApplyLogs(ctx context.Context) (int64, error) {
conn := db.DB(ctx)
if conn == nil {
return 0, errors.New(errDatabaseNotInitialized)
}
result := conn.Session(&gorm.Session{AllowGlobalUpdate: true}).Delete(&OpenFlareApplyLog{})
return result.RowsAffected, result.Error
}
// DeleteOpenFlareApplyLogsBefore removes apply logs created before the cutoff time.
func DeleteOpenFlareApplyLogsBefore(ctx context.Context, before time.Time) (int64, error) {
conn := db.DB(ctx)
if conn == nil {
return 0, errors.New(errDatabaseNotInitialized)
}
result := conn.Where("created_at < ?", before).Delete(&OpenFlareApplyLog{})
return result.RowsAffected, result.Error
}
@@ -1,73 +0,0 @@
package model
import (
"context"
"testing"
"time"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/glebarez/sqlite"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"gorm.io/gorm"
)
func setupApplyLogModelTestDB(t *testing.T) func() {
t.Helper()
sqliteDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{
DisableForeignKeyConstraintWhenMigrating: true,
})
require.NoError(t, err)
require.NoError(t, sqliteDB.AutoMigrate(&OpenFlareApplyLog{}))
db.SetDB(sqliteDB)
return func() {
db.SetDB(nil)
}
}
func TestIsRepeatSuccessApplyLog(t *testing.T) {
latest := &OpenFlareApplyLog{
Version: "20260615-001",
Checksum: "checksum-a",
Result: "success",
}
assert.True(t, IsRepeatSuccessApplyLog(latest, "20260615-001", "checksum-a", "success"))
assert.False(t, IsRepeatSuccessApplyLog(latest, "20260615-002", "checksum-a", "success"))
assert.False(t, IsRepeatSuccessApplyLog(latest, "20260615-001", "checksum-b", "success"))
assert.False(t, IsRepeatSuccessApplyLog(latest, "20260615-001", "checksum-a", "failed"))
assert.False(t, IsRepeatSuccessApplyLog(nil, "20260615-001", "checksum-a", "success"))
}
func TestGetLatestOpenFlareApplyLogByNodeID(t *testing.T) {
cleanup := setupApplyLogModelTestDB(t)
defer cleanup()
ctx := context.Background()
now := time.Now().UTC()
require.NoError(t, db.DB(ctx).Create(&OpenFlareApplyLog{
NodeID: "node-1",
Version: "v1",
Result: "success",
Checksum: "checksum-1",
CreatedAt: now.Add(-time.Hour),
}).Error)
require.NoError(t, db.DB(ctx).Create(&OpenFlareApplyLog{
NodeID: "node-1",
Version: "v2",
Result: "success",
Checksum: "checksum-2",
CreatedAt: now,
}).Error)
latest, err := GetLatestOpenFlareApplyLogByNodeID(ctx, "node-1")
require.NoError(t, err)
require.NotNil(t, latest)
assert.Equal(t, "v2", latest.Version)
missing, err := GetLatestOpenFlareApplyLogByNodeID(ctx, "node-missing")
require.NoError(t, err)
assert.Nil(t, missing)
}
-124
View File
@@ -4,11 +4,8 @@
package model
import (
"context"
"errors"
"time"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
"gorm.io/gorm"
)
@@ -58,124 +55,3 @@ func (cv *ConfigVersion) AfterCreate(_ *gorm.DB) (err error) {
func (ConfigVersion) TableName() string {
return "of_config_versions"
}
// ListConfigVersionSummaries returns config version summaries ordered by created_at desc.
func ListConfigVersionSummaries(ctx context.Context) ([]*ConfigVersionSummary, error) {
conn := db.DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
var versions []*ConfigVersionSummary
err := conn.Model(&ConfigVersion{}).
Select("version", "checksum", "is_active", "created_by", "created_at").
Order("created_at desc, version desc").
Find(&versions).Error
return versions, err
}
// GetConfigVersionByVersion returns a config version by version string.
func GetConfigVersionByVersion(ctx context.Context, version string) (*ConfigVersion, error) {
conn := db.DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
var cv ConfigVersion
if err := conn.First(&cv, "version = ?", version).Error; err != nil {
return nil, err
}
return &cv, nil
}
// GetActiveConfigVersion returns the currently active config version.
func GetActiveConfigVersion(ctx context.Context) (*ConfigVersion, error) {
conn := db.DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
var version ConfigVersion
if err := conn.Where("is_active = ?", true).Order("version desc").First(&version).Error; err != nil {
return nil, err
}
return &version, nil
}
// GetLatestConfigVersionByPrefix returns the latest version string matching a date prefix.
func GetLatestConfigVersionByPrefix(ctx context.Context, prefix string) (string, error) {
conn := db.DB(ctx)
if conn == nil {
return "", errors.New(errDatabaseNotInitialized)
}
var version ConfigVersion
err := conn.Model(&ConfigVersion{}).
Select("version").
Where("version LIKE ?", prefix+"-%").
Order("version desc").
First(&version).Error
if err != nil {
return "", err
}
return version.Version, nil
}
// CreateConfigVersion inserts a new config version record.
func CreateConfigVersion(ctx context.Context, version *ConfigVersion) error {
conn := db.DB(ctx)
if conn == nil {
return errors.New(errDatabaseNotInitialized)
}
return conn.Create(version).Error
}
// PublishConfigVersionTx deactivates all versions and creates a new active version.
func PublishConfigVersionTx(ctx context.Context, version *ConfigVersion) error {
conn := db.DB(ctx)
if conn == nil {
return errors.New(errDatabaseNotInitialized)
}
return conn.Transaction(func(tx *gorm.DB) error {
if err := tx.Model(&ConfigVersion{}).Where("is_active = ?", true).Update("is_active", false).Error; err != nil {
return err
}
return tx.Create(version).Error
})
}
// ActivateConfigVersionTx marks the given version active and deactivates others.
func ActivateConfigVersionTx(ctx context.Context, version string) error {
conn := db.DB(ctx)
if conn == nil {
return errors.New(errDatabaseNotInitialized)
}
return conn.Transaction(func(tx *gorm.DB) error {
if err := tx.Model(&ConfigVersion{}).Where("is_active = ?", true).Update("is_active", false).Error; err != nil {
return err
}
return tx.Model(&ConfigVersion{}).Where("version = ?", version).Update("is_active", true).Error
})
}
// DeleteConfigVersionsByVersions removes config versions by versions.
func DeleteConfigVersionsByVersions(ctx context.Context, versions []string) (int64, error) {
if len(versions) == 0 {
return 0, nil
}
conn := db.DB(ctx)
if conn == nil {
return 0, errors.New(errDatabaseNotInitialized)
}
result := conn.Where("version IN ?", versions).Delete(&ConfigVersion{})
return result.RowsAffected, result.Error
}
// ListEnabledProxyRoutes returns enabled proxy routes ordered by id asc.
func ListEnabledProxyRoutes(ctx context.Context) ([]*ProxyRoute, error) {
conn := db.DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
var routes []*ProxyRoute
if err := conn.Where("enabled = ?", true).Order("id asc").Find(&routes).Error; err != nil {
return nil, err
}
return routes, nil
}
-57
View File
@@ -4,11 +4,7 @@
package model
import (
"context"
"errors"
"time"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
)
// DNSAccount OpenFlare DNS 账号实体。
@@ -25,56 +21,3 @@ type DNSAccount struct {
func (DNSAccount) TableName() string {
return "of_dns_accounts"
}
// ListDNSAccounts 列出全部 DNS 账号(授权信息不通过 JSON 暴露)。
func ListDNSAccounts(ctx context.Context) ([]DNSAccount, error) {
conn := db.DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
var accounts []DNSAccount
if err := conn.Order("id desc").Find(&accounts).Error; err != nil {
return nil, err
}
return accounts, nil
}
// GetDNSAccountByID 按 ID 查询 DNS 账号。
func GetDNSAccountByID(ctx context.Context, id uint) (*DNSAccount, error) {
conn := db.DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
var account DNSAccount
if err := conn.First(&account, id).Error; err != nil {
return nil, err
}
return &account, nil
}
// CreateDNSAccountRecord 创建 DNS 账号。
func CreateDNSAccountRecord(ctx context.Context, account *DNSAccount) error {
conn := db.DB(ctx)
if conn == nil {
return errors.New(errDatabaseNotInitialized)
}
return conn.Create(account).Error
}
// SaveDNSAccount 保存 DNS 账号。
func SaveDNSAccount(ctx context.Context, account *DNSAccount) error {
conn := db.DB(ctx)
if conn == nil {
return errors.New(errDatabaseNotInitialized)
}
return conn.Save(account).Error
}
// DeleteDNSAccountRecord 删除 DNS 账号。
func DeleteDNSAccountRecord(ctx context.Context, id uint) error {
conn := db.DB(ctx)
if conn == nil {
return errors.New(errDatabaseNotInitialized)
}
return conn.Delete(&DNSAccount{}, id).Error
}
@@ -1,75 +0,0 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package model
import (
"context"
"testing"
"time"
analyticsmodel "github.com/Rain-kl/Wavelet/internal/model/analytics"
)
// Hook setters are process-global; keep these tests serial.
func TestObservabilityInsertHooksAreInvoked(t *testing.T) {
var gotSnapshot analyticsmodel.NodeMetricSnapshot
SetObservabilityInsertHooks(ObservabilityInsertHooks{
QueueMetricSnapshot: func(s analyticsmodel.NodeMetricSnapshot) {
gotSnapshot = s
},
})
t.Cleanup(func() {
SetObservabilityInsertHooks(ObservabilityInsertHooks{})
})
record := &OpenFlareMetricSnapshot{
NodeID: "node-1",
CapturedAt: time.Unix(100, 0).UTC(),
}
if err := (clickhouseObservabilityStore{}).InsertMetricSnapshot(context.Background(), record); err != nil {
t.Fatalf("InsertMetricSnapshot error = %v", err)
}
if gotSnapshot.NodeID != "node-1" {
t.Fatalf("hook node id = %q, want node-1", gotSnapshot.NodeID)
}
}
func TestAccessLogInsertHooksAreInvoked(t *testing.T) {
var got []analyticsmodel.NodeAccessLog
SetAccessLogInsertHooks(AccessLogInsertHooks{
QueueNodeAccessLogs: func(logs []analyticsmodel.NodeAccessLog) {
got = append([]analyticsmodel.NodeAccessLog(nil), logs...)
},
})
t.Cleanup(func() {
SetAccessLogInsertHooks(AccessLogInsertHooks{})
})
records := []*OpenFlareAccessLog{
{NodeID: "n1", Path: "/a"},
{NodeID: "n1", Path: "/b"},
}
if err := (clickhouseAccessLogStore{}).InsertBatch(context.Background(), records); err != nil {
t.Fatalf("InsertBatch error = %v", err)
}
if len(got) != 2 {
t.Fatalf("hook logs = %d, want 2", len(got))
}
if got[0].Path != "/a" || got[1].Path != "/b" {
t.Fatalf("hook paths = %q/%q, want /a /b", got[0].Path, got[1].Path)
}
}
func TestInsertHooksNoopWhenUnset(t *testing.T) {
SetObservabilityInsertHooks(ObservabilityInsertHooks{})
SetAccessLogInsertHooks(AccessLogInsertHooks{})
if err := (clickhouseObservabilityStore{}).InsertMetricSnapshot(context.Background(), &OpenFlareMetricSnapshot{NodeID: "x"}); err != nil {
t.Fatalf("InsertMetricSnapshot with nil hook error = %v", err)
}
if err := (clickhouseAccessLogStore{}).InsertBatch(context.Background(), []*OpenFlareAccessLog{{NodeID: "x"}}); err != nil {
t.Fatalf("InsertBatch with nil hook error = %v", err)
}
}
-111
View File
@@ -4,11 +4,7 @@
package model
import (
"context"
"errors"
"time"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
)
// OpenFlareNode stores an edge, relay, or tunnel client node.
@@ -54,110 +50,3 @@ type OpenFlareNode struct {
func (OpenFlareNode) TableName() string {
return "of_nodes"
}
// ListOpenFlareNodes returns all nodes ordered by id desc.
func ListOpenFlareNodes(ctx context.Context) ([]OpenFlareNode, error) {
conn := db.DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
var nodes []OpenFlareNode
if err := conn.Order("id desc").Find(&nodes).Error; err != nil {
return nil, err
}
return nodes, nil
}
// ListOpenFlareNodesByNodeIDs returns nodes matching the given node ids.
func ListOpenFlareNodesByNodeIDs(ctx context.Context, nodeIDs []string) ([]OpenFlareNode, error) {
if len(nodeIDs) == 0 {
return []OpenFlareNode{}, nil
}
conn := db.DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
var nodes []OpenFlareNode
if err := conn.Where("node_id IN ?", nodeIDs).Find(&nodes).Error; err != nil {
return nil, err
}
return nodes, nil
}
// GetOpenFlareNodeByID returns a node by primary key.
func GetOpenFlareNodeByID(ctx context.Context, id uint) (*OpenFlareNode, error) {
conn := db.DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
var node OpenFlareNode
if err := conn.First(&node, id).Error; err != nil {
return nil, err
}
return &node, nil
}
// GetOpenFlareNodeByNodeID returns a node by node_id.
func GetOpenFlareNodeByNodeID(ctx context.Context, nodeID string) (*OpenFlareNode, error) {
conn := db.DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
var node OpenFlareNode
if err := conn.Where("node_id = ?", nodeID).First(&node).Error; err != nil {
return nil, err
}
return &node, nil
}
// GetOpenFlareNodeByAccessToken returns a node by access token.
func GetOpenFlareNodeByAccessToken(ctx context.Context, token string) (*OpenFlareNode, error) {
conn := db.DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
var node OpenFlareNode
if err := conn.Where("access_token = ?", token).First(&node).Error; err != nil {
return nil, err
}
return &node, nil
}
// CreateOpenFlareNode inserts a new node.
func CreateOpenFlareNode(ctx context.Context, node *OpenFlareNode) error {
conn := db.DB(ctx)
if conn == nil {
return errors.New(errDatabaseNotInitialized)
}
return conn.Create(node).Error
}
// SaveOpenFlareNode persists node changes.
func SaveOpenFlareNode(ctx context.Context, node *OpenFlareNode) error {
conn := db.DB(ctx)
if conn == nil {
return errors.New(errDatabaseNotInitialized)
}
return conn.Save(node).Error
}
// UpdateOpenFlareNodeFields updates selected columns for a node.
func UpdateOpenFlareNodeFields(ctx context.Context, node *OpenFlareNode, fields ...string) error {
conn := db.DB(ctx)
if conn == nil {
return errors.New(errDatabaseNotInitialized)
}
if len(fields) == 0 {
return conn.Save(node).Error
}
return conn.Model(node).Select(fields).Updates(node).Error
}
// DeleteOpenFlareNode removes a node by primary key.
func DeleteOpenFlareNode(ctx context.Context, id uint) error {
conn := db.DB(ctx)
if conn == nil {
return errors.New(errDatabaseNotInitialized)
}
return conn.Delete(&OpenFlareNode{}, id).Error
}
-279
View File
@@ -4,14 +4,7 @@
package model
import (
"context"
"errors"
"strings"
"time"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
analyticsrepo "github.com/Rain-kl/Wavelet/internal/repository/analytics"
"gorm.io/gorm"
)
// OpenFlareMetricSnapshot stores a node capacity snapshot in ClickHouse (database: openflare, table: of_node_metric_snapshots).
@@ -285,80 +278,6 @@ type OpenFlareAccessLogWAFIPAggregate struct {
StatusCounts map[int]int
}
func isMissingTableError(err error) bool {
if err == nil {
return false
}
if errors.Is(err, gorm.ErrRecordNotFound) {
return false
}
msg := strings.ToLower(err.Error())
return strings.Contains(msg, "no such table") ||
strings.Contains(msg, "doesn't exist") ||
strings.Contains(msg, "does not exist")
}
// InsertOpenFlareMetricSnapshot inserts a metric snapshot into ClickHouse.
func InsertOpenFlareMetricSnapshot(ctx context.Context, record *OpenFlareMetricSnapshot) error {
return currentObservabilityStore().InsertMetricSnapshot(ctx, record)
}
// InsertOpenFlareEdgeHealth inserts an L2 edge health snapshot into ClickHouse.
func InsertOpenFlareEdgeHealth(ctx context.Context, record *OpenFlareEdgeHealth) error {
return currentObservabilityStore().InsertEdgeHealth(ctx, record)
}
// InsertOpenFlareNodeObservationFrps inserts an FRPS observation into ClickHouse.
func InsertOpenFlareNodeObservationFrps(ctx context.Context, record *OpenFlareNodeObservationFrps) error {
return currentObservabilityStore().InsertNodeObservationFrps(ctx, record)
}
// InsertOpenFlareNodeObservationFrpc inserts an FRPC observation into ClickHouse.
func InsertOpenFlareNodeObservationFrpc(ctx context.Context, record *OpenFlareNodeObservationFrpc) error {
return currentObservabilityStore().InsertNodeObservationFrpc(ctx, record)
}
// ListOpenFlareMetricSnapshotsSince returns metric snapshots since the given time.
func ListOpenFlareMetricSnapshotsSince(ctx context.Context, nodeID string, since time.Time, limit int) ([]*OpenFlareMetricSnapshot, error) {
return currentObservabilityStore().ListMetricSnapshots(ctx, nodeID, since, limit)
}
// ListOpenFlareLatestMetricSnapshotsSince returns the latest metric snapshot per node.
// Prefer ClickHouse LIMIT 1 BY; on CH unavailability fall back to store list + reduce.
func ListOpenFlareLatestMetricSnapshotsSince(ctx context.Context, nodeID string, since time.Time) ([]*OpenFlareMetricSnapshot, error) {
rows, err := analyticsrepo.ListLatestNodeMetricSnapshots(ctx, analyticsrepo.NodeObservabilityFilter{
NodeID: nodeID,
Since: since,
})
if err == nil {
return fromAnalyticsNodeMetricSnapshots(rows), nil
}
// Fallback for unit tests (memory store) and environments without ClickHouse.
all, listErr := ListOpenFlareMetricSnapshotsSince(ctx, nodeID, since, 0)
if listErr != nil {
return nil, err
}
return openFlareLatestMetricSnapshots(all), nil
}
func openFlareLatestMetricSnapshots(snapshots []*OpenFlareMetricSnapshot) []*OpenFlareMetricSnapshot {
latestByNode := make(map[string]*OpenFlareMetricSnapshot, len(snapshots))
for _, snapshot := range snapshots {
if snapshot == nil || snapshot.NodeID == "" {
continue
}
if existing, ok := latestByNode[snapshot.NodeID]; ok && !snapshot.CapturedAt.After(existing.CapturedAt) {
continue
}
latestByNode[snapshot.NodeID] = snapshot
}
result := make([]*OpenFlareMetricSnapshot, 0, len(latestByNode))
for _, snapshot := range latestByNode {
result = append(result, snapshot)
}
return result
}
// OpenFlareTrafficHourly is an hourly traffic rollup row.
type OpenFlareTrafficHourly struct {
NodeID string `json:"node_id"`
@@ -368,29 +287,6 @@ type OpenFlareTrafficHourly struct {
UniqueVisitorCount int64 `json:"unique_visitor_count"`
}
// ListOpenFlareTrafficHourlySince returns hourly traffic rollup rows since the given time.
// Source: of_access_log_hourly (M5).
func ListOpenFlareTrafficHourlySince(ctx context.Context, nodeID string, since time.Time) ([]*OpenFlareTrafficHourly, error) {
rows, err := analyticsrepo.ListNodeTrafficHourly(ctx, analyticsrepo.NodeObservabilityFilter{
NodeID: nodeID,
Since: since,
})
if err != nil {
return nil, err
}
result := make([]*OpenFlareTrafficHourly, len(rows))
for index, row := range rows {
result[index] = &OpenFlareTrafficHourly{
NodeID: row.NodeID,
Hour: row.Hour,
RequestCount: row.RequestCount,
ErrorCount: row.ErrorCount,
UniqueVisitorCount: row.UniqueVisitorCount,
}
}
return result, nil
}
// OpenFlareAccessLogHourly is a per-node/host hourly access log rollup.
type OpenFlareAccessLogHourly struct {
NodeID string `json:"node_id"`
@@ -402,30 +298,6 @@ type OpenFlareAccessLogHourly struct {
RequestLength int64 `json:"request_length"`
}
// ListOpenFlareAccessLogHourlySince returns of_access_log_hourly rows since the given time.
func ListOpenFlareAccessLogHourlySince(ctx context.Context, nodeID string, since time.Time) ([]*OpenFlareAccessLogHourly, error) {
rows, err := analyticsrepo.ListAccessLogHourly(ctx, analyticsrepo.NodeObservabilityFilter{
NodeID: nodeID,
Since: since,
})
if err != nil {
return nil, err
}
result := make([]*OpenFlareAccessLogHourly, len(rows))
for index, row := range rows {
result[index] = &OpenFlareAccessLogHourly{
NodeID: row.NodeID,
Hour: row.Hour,
Host: row.Host,
RequestCount: row.RequestCount,
ErrorCount: row.ErrorCount,
BytesSent: row.BytesSent,
RequestLength: row.RequestLength,
}
}
return result, nil
}
// OpenFlareMetricHourly is an hourly metric snapshot aggregation row.
type OpenFlareMetricHourly struct {
Hour time.Time `json:"hour"`
@@ -437,154 +309,3 @@ type OpenFlareMetricHourly struct {
DiskWriteBytes int64 `json:"disk_write_bytes"`
ReportedNodes int `json:"reported_nodes"`
}
// ListOpenFlareMetricHourlySince returns hourly metric aggregates since the given time.
func ListOpenFlareMetricHourlySince(ctx context.Context, nodeID string, since time.Time) ([]*OpenFlareMetricHourly, error) {
rows, err := analyticsrepo.ListNodeMetricHourly(ctx, analyticsrepo.NodeObservabilityFilter{
NodeID: nodeID,
Since: since,
})
if err != nil {
return nil, err
}
result := make([]*OpenFlareMetricHourly, len(rows))
for index, row := range rows {
result[index] = &OpenFlareMetricHourly{
Hour: row.Hour,
AverageCPUUsagePercent: row.AverageCPUUsagePercent,
AverageMemoryUsagePercent: row.AverageMemoryUsagePercent,
NetworkRxBytes: row.NetworkRxBytes,
NetworkTxBytes: row.NetworkTxBytes,
DiskReadBytes: row.DiskReadBytes,
DiskWriteBytes: row.DiskWriteBytes,
ReportedNodes: row.ReportedNodes,
}
}
return result, nil
}
// ListOpenFlareActiveHealthEvents returns active health events across all nodes.
func ListOpenFlareActiveHealthEvents(ctx context.Context) ([]*OpenFlareHealthEvent, error) {
conn := db.DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
var rows []*OpenFlareHealthEvent
if err := conn.Where("status = ?", "active").Order("last_triggered_at desc").Find(&rows).Error; err != nil {
if isMissingTableError(err) {
return []*OpenFlareHealthEvent{}, nil
}
return nil, err
}
return rows, nil
}
// ListOpenFlareHealthEvents returns health events for a node.
func ListOpenFlareHealthEvents(ctx context.Context, nodeID string, activeOnly bool, limit int) ([]*OpenFlareHealthEvent, error) {
conn := db.DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
query := conn.Model(&OpenFlareHealthEvent{}).Where("node_id = ?", nodeID).Order("last_triggered_at desc")
if activeOnly {
query = query.Where("status = ?", "active")
}
if limit > 0 {
query = query.Limit(limit)
}
var rows []*OpenFlareHealthEvent
if err := query.Find(&rows).Error; err != nil {
if isMissingTableError(err) {
return []*OpenFlareHealthEvent{}, nil
}
return nil, err
}
return rows, nil
}
// DeleteOpenFlareMetricSnapshotsBefore deletes metric snapshots captured before cutoff.
func DeleteOpenFlareMetricSnapshotsBefore(ctx context.Context, cutoff time.Time) (int64, error) {
return currentObservabilityStore().DeleteMetricSnapshotsBefore(ctx, cutoff)
}
// DeleteAllOpenFlareMetricSnapshots deletes all metric snapshots.
func DeleteAllOpenFlareMetricSnapshots(ctx context.Context) (int64, error) {
return currentObservabilityStore().DeleteAllMetricSnapshots(ctx)
}
// DeleteOpenFlareEdgeHealthBefore deletes edge health rows captured before cutoff.
func DeleteOpenFlareEdgeHealthBefore(ctx context.Context, cutoff time.Time) (int64, error) {
return currentObservabilityStore().DeleteEdgeHealthBefore(ctx, cutoff)
}
// DeleteAllOpenFlareEdgeHealth deletes all edge health snapshots.
func DeleteAllOpenFlareEdgeHealth(ctx context.Context) (int64, error) {
return currentObservabilityStore().DeleteAllEdgeHealth(ctx)
}
// DeleteOpenFlareNodeObservationFrpsBefore deletes FRPS observations captured before cutoff.
func DeleteOpenFlareNodeObservationFrpsBefore(ctx context.Context, cutoff time.Time) (int64, error) {
return currentObservabilityStore().DeleteNodeObservationFrpsBefore(ctx, cutoff)
}
// DeleteAllOpenFlareNodeObservationFrps deletes all FRPS observations.
func DeleteAllOpenFlareNodeObservationFrps(ctx context.Context) (int64, error) {
return currentObservabilityStore().DeleteAllNodeObservationFrps(ctx)
}
// DeleteOpenFlareNodeObservationFrpcBefore deletes FRPC observations captured before cutoff.
func DeleteOpenFlareNodeObservationFrpcBefore(ctx context.Context, cutoff time.Time) (int64, error) {
return currentObservabilityStore().DeleteNodeObservationFrpcBefore(ctx, cutoff)
}
// DeleteAllOpenFlareNodeObservationFrpc deletes all FRPC observations.
func DeleteAllOpenFlareNodeObservationFrpc(ctx context.Context) (int64, error) {
return currentObservabilityStore().DeleteAllNodeObservationFrpc(ctx)
}
// DeleteOpenFlareHealthEventsByNodeID deletes all health events for a node.
func DeleteOpenFlareHealthEventsByNodeID(ctx context.Context, nodeID string) (int64, error) {
conn := db.DB(ctx)
if conn == nil {
return 0, errors.New(errDatabaseNotInitialized)
}
result := conn.Where("node_id = ?", nodeID).Delete(&OpenFlareHealthEvent{})
if result.Error != nil {
if isMissingTableError(result.Error) {
return 0, nil
}
return 0, result.Error
}
return result.RowsAffected, nil
}
// GetOpenFlareNodeSystemProfile returns the system profile for a node.
func GetOpenFlareNodeSystemProfile(ctx context.Context, nodeID string) (*OpenFlareNodeSystemProfile, error) {
conn := db.DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
var profile OpenFlareNodeSystemProfile
if err := conn.Where("node_id = ?", nodeID).First(&profile).Error; err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) || isMissingTableError(err) {
return nil, gorm.ErrRecordNotFound
}
return nil, err
}
return &profile, nil
}
// ListOpenFlareEdgeHealth returns L2 edge health snapshots.
func ListOpenFlareEdgeHealth(ctx context.Context, nodeID string, since time.Time, limit int) ([]*OpenFlareEdgeHealth, error) {
return currentObservabilityStore().ListEdgeHealth(ctx, nodeID, since, limit)
}
// ListOpenFlareNodeObservationFrpc returns frpc observations.
func ListOpenFlareNodeObservationFrpc(ctx context.Context, nodeID string, since time.Time, limit int) ([]*OpenFlareNodeObservationFrpc, error) {
return currentObservabilityStore().ListNodeObservationFrpc(ctx, nodeID, since, limit)
}
// ListOpenFlareNodeObservationFrps returns frps observations.
func ListOpenFlareNodeObservationFrps(ctx context.Context, nodeID string, since time.Time, limit int) ([]*OpenFlareNodeObservationFrps, error) {
return currentObservabilityStore().ListNodeObservationFrps(ctx, nodeID, since, limit)
}
@@ -1,353 +0,0 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package model
import (
"context"
"math"
"strings"
"sync"
"time"
analyticsmodel "github.com/Rain-kl/Wavelet/internal/model/analytics"
analyticsrepo "github.com/Rain-kl/Wavelet/internal/repository/analytics"
)
// ObservabilityInsertHooks queues observability rows for async ClickHouse write.
// Wired from openflare/chwriter.Init so model never imports the apps layer.
type ObservabilityInsertHooks struct {
QueueMetricSnapshot func(analyticsmodel.NodeMetricSnapshot)
QueueEdgeHealth func(analyticsmodel.NodeEdgeHealth)
QueueFrpsObservation func(analyticsmodel.NodeObsFrps)
QueueFrpcObservation func(analyticsmodel.NodeObsFrpc)
}
var (
observabilityInsertHooksMu sync.RWMutex
observabilityInsertHooks ObservabilityInsertHooks
)
// SetObservabilityInsertHooks registers async queue callbacks for observability inserts.
func SetObservabilityInsertHooks(hooks ObservabilityInsertHooks) {
observabilityInsertHooksMu.Lock()
observabilityInsertHooks = hooks
observabilityInsertHooksMu.Unlock()
}
func currentObservabilityInsertHooks() ObservabilityInsertHooks {
observabilityInsertHooksMu.RLock()
defer observabilityInsertHooksMu.RUnlock()
return observabilityInsertHooks
}
type observabilityStore interface {
InsertMetricSnapshot(ctx context.Context, record *OpenFlareMetricSnapshot) error
ListMetricSnapshots(ctx context.Context, nodeID string, since time.Time, limit int) ([]*OpenFlareMetricSnapshot, error)
DeleteAllMetricSnapshots(ctx context.Context) (int64, error)
DeleteMetricSnapshotsBefore(ctx context.Context, cutoff time.Time) (int64, error)
InsertEdgeHealth(ctx context.Context, record *OpenFlareEdgeHealth) error
ListEdgeHealth(ctx context.Context, nodeID string, since time.Time, limit int) ([]*OpenFlareEdgeHealth, error)
DeleteAllEdgeHealth(ctx context.Context) (int64, error)
DeleteEdgeHealthBefore(ctx context.Context, cutoff time.Time) (int64, error)
InsertNodeObservationFrps(ctx context.Context, record *OpenFlareNodeObservationFrps) error
ListNodeObservationFrps(ctx context.Context, nodeID string, since time.Time, limit int) ([]*OpenFlareNodeObservationFrps, error)
DeleteAllNodeObservationFrps(ctx context.Context) (int64, error)
DeleteNodeObservationFrpsBefore(ctx context.Context, cutoff time.Time) (int64, error)
InsertNodeObservationFrpc(ctx context.Context, record *OpenFlareNodeObservationFrpc) error
ListNodeObservationFrpc(ctx context.Context, nodeID string, since time.Time, limit int) ([]*OpenFlareNodeObservationFrpc, error)
DeleteAllNodeObservationFrpc(ctx context.Context) (int64, error)
DeleteNodeObservationFrpcBefore(ctx context.Context, cutoff time.Time) (int64, error)
}
var (
observabilityStoreMu sync.RWMutex
observabilityStoreHolder observabilityStore
)
func currentObservabilityStore() observabilityStore {
observabilityStoreMu.RLock()
defer observabilityStoreMu.RUnlock()
if observabilityStoreHolder != nil {
return observabilityStoreHolder
}
return clickhouseObservabilityStore{}
}
// SetObservabilityStoreForTest swaps the observability store implementation for unit tests.
func SetObservabilityStoreForTest(store observabilityStore) func() {
observabilityStoreMu.Lock()
previous := observabilityStoreHolder
observabilityStoreHolder = store
observabilityStoreMu.Unlock()
return func() {
observabilityStoreMu.Lock()
observabilityStoreHolder = previous
observabilityStoreMu.Unlock()
}
}
// NewMemoryObservabilityStore returns an in-memory observability store for unit tests.
func NewMemoryObservabilityStore() observabilityStore {
return &memoryObservabilityStore{}
}
type clickhouseObservabilityStore struct{}
func (clickhouseObservabilityStore) InsertMetricSnapshot(_ context.Context, record *OpenFlareMetricSnapshot) error {
if record == nil {
return nil
}
if hook := currentObservabilityInsertHooks().QueueMetricSnapshot; hook != nil {
hook(toAnalyticsNodeMetricSnapshot(record))
}
return nil
}
func (clickhouseObservabilityStore) ListMetricSnapshots(ctx context.Context, nodeID string, since time.Time, limit int) ([]*OpenFlareMetricSnapshot, error) {
rows, err := analyticsrepo.ListNodeMetricSnapshots(ctx, toNodeObservabilityFilter(nodeID, since, limit))
if err != nil {
return nil, err
}
return fromAnalyticsNodeMetricSnapshots(rows), nil
}
func (clickhouseObservabilityStore) DeleteAllMetricSnapshots(ctx context.Context) (int64, error) {
return analyticsrepo.DeleteAllNodeMetricSnapshots(ctx)
}
func (clickhouseObservabilityStore) DeleteMetricSnapshotsBefore(ctx context.Context, cutoff time.Time) (int64, error) {
return analyticsrepo.DeleteNodeMetricSnapshotsBefore(ctx, cutoff)
}
const edgeHealthStatusUnknown = "unknown"
func normalizeEdgeHealthStatus(status string) string {
status = strings.TrimSpace(status)
if status == "" {
return edgeHealthStatusUnknown
}
return status
}
func (clickhouseObservabilityStore) InsertEdgeHealth(_ context.Context, record *OpenFlareEdgeHealth) error {
if record == nil {
return nil
}
if hook := currentObservabilityInsertHooks().QueueEdgeHealth; hook != nil {
hook(toAnalyticsNodeEdgeHealth(record))
}
return nil
}
func (clickhouseObservabilityStore) ListEdgeHealth(ctx context.Context, nodeID string, since time.Time, limit int) ([]*OpenFlareEdgeHealth, error) {
rows, err := analyticsrepo.ListNodeEdgeHealth(ctx, toNodeObservabilityFilter(nodeID, since, limit))
if err != nil {
return nil, err
}
return fromAnalyticsNodeEdgeHealth(rows), nil
}
func (clickhouseObservabilityStore) DeleteAllEdgeHealth(ctx context.Context) (int64, error) {
return analyticsrepo.DeleteAllNodeEdgeHealth(ctx)
}
func (clickhouseObservabilityStore) DeleteEdgeHealthBefore(ctx context.Context, cutoff time.Time) (int64, error) {
return analyticsrepo.DeleteNodeEdgeHealthBefore(ctx, cutoff)
}
func (clickhouseObservabilityStore) InsertNodeObservationFrps(_ context.Context, record *OpenFlareNodeObservationFrps) error {
if record == nil {
return nil
}
if hook := currentObservabilityInsertHooks().QueueFrpsObservation; hook != nil {
hook(toAnalyticsNodeObsFrps(record))
}
return nil
}
func (clickhouseObservabilityStore) ListNodeObservationFrps(ctx context.Context, nodeID string, since time.Time, limit int) ([]*OpenFlareNodeObservationFrps, error) {
rows, err := analyticsrepo.ListNodeObsFrps(ctx, toNodeObservabilityFilter(nodeID, since, limit))
if err != nil {
return nil, err
}
return fromAnalyticsNodeObsFrps(rows), nil
}
func (clickhouseObservabilityStore) DeleteAllNodeObservationFrps(ctx context.Context) (int64, error) {
return analyticsrepo.DeleteAllNodeObsFrps(ctx)
}
func (clickhouseObservabilityStore) DeleteNodeObservationFrpsBefore(ctx context.Context, cutoff time.Time) (int64, error) {
return analyticsrepo.DeleteNodeObsFrpsBefore(ctx, cutoff)
}
func (clickhouseObservabilityStore) InsertNodeObservationFrpc(_ context.Context, record *OpenFlareNodeObservationFrpc) error {
if record == nil {
return nil
}
if hook := currentObservabilityInsertHooks().QueueFrpcObservation; hook != nil {
hook(toAnalyticsNodeObsFrpc(record))
}
return nil
}
func (clickhouseObservabilityStore) ListNodeObservationFrpc(ctx context.Context, nodeID string, since time.Time, limit int) ([]*OpenFlareNodeObservationFrpc, error) {
rows, err := analyticsrepo.ListNodeObsFrpc(ctx, toNodeObservabilityFilter(nodeID, since, limit))
if err != nil {
return nil, err
}
return fromAnalyticsNodeObsFrpc(rows), nil
}
func (clickhouseObservabilityStore) DeleteAllNodeObservationFrpc(ctx context.Context) (int64, error) {
return analyticsrepo.DeleteAllNodeObsFrpc(ctx)
}
func (clickhouseObservabilityStore) DeleteNodeObservationFrpcBefore(ctx context.Context, cutoff time.Time) (int64, error) {
return analyticsrepo.DeleteNodeObsFrpcBefore(ctx, cutoff)
}
func toNodeObservabilityFilter(nodeID string, since time.Time, limit int) analyticsrepo.NodeObservabilityFilter {
return analyticsrepo.NodeObservabilityFilter{
NodeID: nodeID,
Since: since,
Limit: limit,
}
}
func toAnalyticsNodeMetricSnapshot(record *OpenFlareMetricSnapshot) analyticsmodel.NodeMetricSnapshot {
return analyticsmodel.NodeMetricSnapshot{
ID: uint64(record.ID),
NodeID: record.NodeID,
CapturedAt: record.CapturedAt,
CPUUsagePercent: record.CPUUsagePercent,
MemoryUsedBytes: record.MemoryUsedBytes,
MemoryTotalBytes: record.MemoryTotalBytes,
StorageUsedBytes: record.StorageUsedBytes,
StorageTotalBytes: record.StorageTotalBytes,
DiskReadBytes: record.DiskReadBytes,
DiskWriteBytes: record.DiskWriteBytes,
NetworkRxBytes: record.NetworkRxBytes,
NetworkTxBytes: record.NetworkTxBytes,
CreatedAt: record.CreatedAt,
}
}
func fromAnalyticsNodeMetricSnapshots(rows []analyticsmodel.NodeMetricSnapshot) []*OpenFlareMetricSnapshot {
result := make([]*OpenFlareMetricSnapshot, len(rows))
for index, row := range rows {
result[index] = &OpenFlareMetricSnapshot{
ID: uint(row.ID),
NodeID: row.NodeID,
CapturedAt: row.CapturedAt,
CPUUsagePercent: row.CPUUsagePercent,
MemoryUsedBytes: row.MemoryUsedBytes,
MemoryTotalBytes: row.MemoryTotalBytes,
StorageUsedBytes: row.StorageUsedBytes,
StorageTotalBytes: row.StorageTotalBytes,
DiskReadBytes: row.DiskReadBytes,
DiskWriteBytes: row.DiskWriteBytes,
NetworkRxBytes: row.NetworkRxBytes,
NetworkTxBytes: row.NetworkTxBytes,
CreatedAt: row.CreatedAt,
}
}
return result
}
func toAnalyticsNodeEdgeHealth(record *OpenFlareEdgeHealth) analyticsmodel.NodeEdgeHealth {
return analyticsmodel.NodeEdgeHealth{
ID: uint64(record.ID),
NodeID: record.NodeID,
CapturedAt: record.CapturedAt,
Status: normalizeEdgeHealthStatus(record.Status),
Connections: record.Connections,
CreatedAt: record.CreatedAt,
}
}
func fromAnalyticsNodeEdgeHealth(rows []analyticsmodel.NodeEdgeHealth) []*OpenFlareEdgeHealth {
result := make([]*OpenFlareEdgeHealth, len(rows))
for index, row := range rows {
result[index] = &OpenFlareEdgeHealth{
ID: uint(row.ID),
NodeID: row.NodeID,
CapturedAt: row.CapturedAt,
Status: normalizeEdgeHealthStatus(row.Status),
Connections: row.Connections,
CreatedAt: row.CreatedAt,
}
}
return result
}
func toAnalyticsNodeObsFrps(record *OpenFlareNodeObservationFrps) analyticsmodel.NodeObsFrps {
return analyticsmodel.NodeObsFrps{
ID: uint64(record.ID),
NodeID: record.NodeID,
CapturedAt: record.CapturedAt,
FrpsConnections: openFlareObservabilityIntToInt32(record.FrpsConnections),
FrpsProxyCount: openFlareObservabilityIntToInt32(record.FrpsProxyCount),
FrpsClientCount: openFlareObservabilityIntToInt32(record.FrpsClientCount),
FrpsProxies: record.FrpsProxies,
CreatedAt: record.CreatedAt,
}
}
func fromAnalyticsNodeObsFrps(rows []analyticsmodel.NodeObsFrps) []*OpenFlareNodeObservationFrps {
result := make([]*OpenFlareNodeObservationFrps, len(rows))
for index, row := range rows {
result[index] = &OpenFlareNodeObservationFrps{
ID: uint(row.ID),
NodeID: row.NodeID,
CapturedAt: row.CapturedAt,
FrpsConnections: int(row.FrpsConnections),
FrpsProxyCount: int(row.FrpsProxyCount),
FrpsClientCount: int(row.FrpsClientCount),
FrpsProxies: row.FrpsProxies,
CreatedAt: row.CreatedAt,
}
}
return result
}
func toAnalyticsNodeObsFrpc(record *OpenFlareNodeObservationFrpc) analyticsmodel.NodeObsFrpc {
return analyticsmodel.NodeObsFrpc{
ID: uint64(record.ID),
NodeID: record.NodeID,
CapturedAt: record.CapturedAt,
TunnelStatus: record.TunnelStatus,
ConnectedRelaysCount: openFlareObservabilityIntToInt32(record.ConnectedRelaysCount),
CreatedAt: record.CreatedAt,
}
}
func openFlareObservabilityIntToInt32(value int) int32 {
switch {
case value > math.MaxInt32:
return math.MaxInt32
case value < math.MinInt32:
return math.MinInt32
default:
return int32(value)
}
}
func fromAnalyticsNodeObsFrpc(rows []analyticsmodel.NodeObsFrpc) []*OpenFlareNodeObservationFrpc {
result := make([]*OpenFlareNodeObservationFrpc, len(rows))
for index, row := range rows {
result[index] = &OpenFlareNodeObservationFrpc{
ID: uint(row.ID),
NodeID: row.NodeID,
CapturedAt: row.CapturedAt,
TunnelStatus: row.TunnelStatus,
ConnectedRelaysCount: int(row.ConnectedRelaysCount),
CreatedAt: row.CreatedAt,
}
}
return result
}
@@ -1,407 +0,0 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package model
import (
"context"
"sort"
"strings"
"sync"
"time"
"github.com/Rain-kl/Wavelet/internal/infra/persistence/idgen"
)
type memoryObservabilityStore struct {
mu sync.RWMutex
metricSnapshots []*OpenFlareMetricSnapshot
edgeHealth []*OpenFlareEdgeHealth
frpsObs []*OpenFlareNodeObservationFrps
frpcObs []*OpenFlareNodeObservationFrpc
}
func (s *memoryObservabilityStore) InsertMetricSnapshot(_ context.Context, record *OpenFlareMetricSnapshot) error {
if record == nil {
return nil
}
s.mu.Lock()
defer s.mu.Unlock()
copyRecord := cloneOpenFlareMetricSnapshot(record)
if memoryMetricSnapshotExists(s.metricSnapshots, copyRecord.NodeID, copyRecord.CapturedAt) {
return nil
}
s.metricSnapshots = append(s.metricSnapshots, copyRecord)
return nil
}
func (s *memoryObservabilityStore) ListMetricSnapshots(_ context.Context, nodeID string, since time.Time, limit int) ([]*OpenFlareMetricSnapshot, error) {
s.mu.RLock()
defer s.mu.RUnlock()
rows := memoryFilterMetricSnapshots(s.metricSnapshots, nodeID, since)
sortOpenFlareMetricSnapshots(rows)
return memoryLimitObservabilityRows(rows, limit), nil
}
func (s *memoryObservabilityStore) DeleteAllMetricSnapshots(_ context.Context) (int64, error) {
s.mu.Lock()
defer s.mu.Unlock()
count := int64(len(s.metricSnapshots))
s.metricSnapshots = nil
return count, nil
}
func (s *memoryObservabilityStore) DeleteMetricSnapshotsBefore(_ context.Context, cutoff time.Time) (int64, error) {
s.mu.Lock()
defer s.mu.Unlock()
cutoff = cutoff.UTC()
remaining := make([]*OpenFlareMetricSnapshot, 0, len(s.metricSnapshots))
var deleted int64
for _, row := range s.metricSnapshots {
if row.CapturedAt.Before(cutoff) {
deleted++
continue
}
remaining = append(remaining, row)
}
s.metricSnapshots = remaining
return deleted, nil
}
func (s *memoryObservabilityStore) InsertEdgeHealth(_ context.Context, record *OpenFlareEdgeHealth) error {
if record == nil {
return nil
}
s.mu.Lock()
defer s.mu.Unlock()
s.edgeHealth = append(s.edgeHealth, cloneOpenFlareEdgeHealth(record))
return nil
}
func (s *memoryObservabilityStore) ListEdgeHealth(_ context.Context, nodeID string, since time.Time, limit int) ([]*OpenFlareEdgeHealth, error) {
s.mu.RLock()
defer s.mu.RUnlock()
rows := memoryFilterEdgeHealth(s.edgeHealth, nodeID, since)
sortOpenFlareEdgeHealth(rows)
return memoryLimitObservabilityRows(rows, limit), nil
}
func (s *memoryObservabilityStore) DeleteAllEdgeHealth(_ context.Context) (int64, error) {
s.mu.Lock()
defer s.mu.Unlock()
count := int64(len(s.edgeHealth))
s.edgeHealth = nil
return count, nil
}
func (s *memoryObservabilityStore) DeleteEdgeHealthBefore(_ context.Context, cutoff time.Time) (int64, error) {
s.mu.Lock()
defer s.mu.Unlock()
cutoff = cutoff.UTC()
remaining := make([]*OpenFlareEdgeHealth, 0, len(s.edgeHealth))
var deleted int64
for _, row := range s.edgeHealth {
if row.CapturedAt.Before(cutoff) {
deleted++
continue
}
remaining = append(remaining, row)
}
s.edgeHealth = remaining
return deleted, nil
}
func (s *memoryObservabilityStore) InsertNodeObservationFrps(_ context.Context, record *OpenFlareNodeObservationFrps) error {
if record == nil {
return nil
}
s.mu.Lock()
defer s.mu.Unlock()
s.frpsObs = append(s.frpsObs, cloneOpenFlareNodeObservationFrps(record))
return nil
}
func (s *memoryObservabilityStore) ListNodeObservationFrps(_ context.Context, nodeID string, since time.Time, limit int) ([]*OpenFlareNodeObservationFrps, error) {
s.mu.RLock()
defer s.mu.RUnlock()
rows := memoryFilterFrpsObservations(s.frpsObs, nodeID, since)
sortOpenFlareNodeObservationFrps(rows)
return memoryLimitObservabilityRows(rows, limit), nil
}
func (s *memoryObservabilityStore) DeleteAllNodeObservationFrps(_ context.Context) (int64, error) {
s.mu.Lock()
defer s.mu.Unlock()
count := int64(len(s.frpsObs))
s.frpsObs = nil
return count, nil
}
func (s *memoryObservabilityStore) DeleteNodeObservationFrpsBefore(_ context.Context, cutoff time.Time) (int64, error) {
s.mu.Lock()
defer s.mu.Unlock()
cutoff = cutoff.UTC()
remaining := make([]*OpenFlareNodeObservationFrps, 0, len(s.frpsObs))
var deleted int64
for _, row := range s.frpsObs {
if row.CapturedAt.Before(cutoff) {
deleted++
continue
}
remaining = append(remaining, row)
}
s.frpsObs = remaining
return deleted, nil
}
func (s *memoryObservabilityStore) InsertNodeObservationFrpc(_ context.Context, record *OpenFlareNodeObservationFrpc) error {
if record == nil {
return nil
}
s.mu.Lock()
defer s.mu.Unlock()
s.frpcObs = append(s.frpcObs, cloneOpenFlareNodeObservationFrpc(record))
return nil
}
func (s *memoryObservabilityStore) ListNodeObservationFrpc(_ context.Context, nodeID string, since time.Time, limit int) ([]*OpenFlareNodeObservationFrpc, error) {
s.mu.RLock()
defer s.mu.RUnlock()
rows := memoryFilterFrpcObservations(s.frpcObs, nodeID, since)
sortOpenFlareNodeObservationFrpc(rows)
return memoryLimitObservabilityRows(rows, limit), nil
}
func (s *memoryObservabilityStore) DeleteAllNodeObservationFrpc(_ context.Context) (int64, error) {
s.mu.Lock()
defer s.mu.Unlock()
count := int64(len(s.frpcObs))
s.frpcObs = nil
return count, nil
}
func (s *memoryObservabilityStore) DeleteNodeObservationFrpcBefore(_ context.Context, cutoff time.Time) (int64, error) {
s.mu.Lock()
defer s.mu.Unlock()
cutoff = cutoff.UTC()
remaining := make([]*OpenFlareNodeObservationFrpc, 0, len(s.frpcObs))
var deleted int64
for _, row := range s.frpcObs {
if row.CapturedAt.Before(cutoff) {
deleted++
continue
}
remaining = append(remaining, row)
}
s.frpcObs = remaining
return deleted, nil
}
func memoryFilterMetricSnapshots(rows []*OpenFlareMetricSnapshot, nodeID string, since time.Time) []*OpenFlareMetricSnapshot {
result := make([]*OpenFlareMetricSnapshot, 0, len(rows))
for _, row := range rows {
if !memoryObservabilityMatchesNodeID(row.NodeID, nodeID) {
continue
}
if !since.IsZero() && row.CapturedAt.Before(since) {
continue
}
result = append(result, row)
}
return result
}
func memoryFilterEdgeHealth(rows []*OpenFlareEdgeHealth, nodeID string, since time.Time) []*OpenFlareEdgeHealth {
result := make([]*OpenFlareEdgeHealth, 0, len(rows))
for _, row := range rows {
if !memoryObservabilityMatchesNodeID(row.NodeID, nodeID) {
continue
}
if !since.IsZero() && row.CapturedAt.Before(since) {
continue
}
result = append(result, row)
}
return result
}
func memoryFilterFrpsObservations(rows []*OpenFlareNodeObservationFrps, nodeID string, since time.Time) []*OpenFlareNodeObservationFrps {
result := make([]*OpenFlareNodeObservationFrps, 0, len(rows))
for _, row := range rows {
if !memoryObservabilityMatchesNodeID(row.NodeID, nodeID) {
continue
}
if !since.IsZero() && row.CapturedAt.Before(since) {
continue
}
result = append(result, row)
}
return result
}
func memoryFilterFrpcObservations(rows []*OpenFlareNodeObservationFrpc, nodeID string, since time.Time) []*OpenFlareNodeObservationFrpc {
result := make([]*OpenFlareNodeObservationFrpc, 0, len(rows))
for _, row := range rows {
if !memoryObservabilityMatchesNodeID(row.NodeID, nodeID) {
continue
}
if !since.IsZero() && row.CapturedAt.Before(since) {
continue
}
result = append(result, row)
}
return result
}
func memoryObservabilityMatchesNodeID(rowNodeID string, nodeID string) bool {
trimmed := strings.TrimSpace(nodeID)
if trimmed == "" {
return true
}
return rowNodeID == trimmed
}
func memoryMetricSnapshotExists(rows []*OpenFlareMetricSnapshot, nodeID string, capturedAt time.Time) bool {
capturedAt = capturedAt.UTC()
for _, row := range rows {
if row.NodeID == nodeID && row.CapturedAt.UTC().Equal(capturedAt) {
return true
}
}
return false
}
func sortOpenFlareMetricSnapshots(items []*OpenFlareMetricSnapshot) {
sort.Slice(items, func(i, j int) bool {
left := items[i]
right := items[j]
if left == nil || right == nil {
return left != nil
}
if compare := openFlareAccessLogCompareInt64(left.CapturedAt.Unix(), right.CapturedAt.Unix()); compare != 0 {
return compare > 0
}
return openFlareAccessLogCompareInt64(openFlareAccessLogUintToInt64(uint64(left.ID)), openFlareAccessLogUintToInt64(uint64(right.ID))) > 0
})
}
func sortOpenFlareEdgeHealth(items []*OpenFlareEdgeHealth) {
sort.Slice(items, func(i, j int) bool {
left := items[i]
right := items[j]
if left == nil || right == nil {
return left != nil
}
if compare := openFlareAccessLogCompareInt64(left.CapturedAt.Unix(), right.CapturedAt.Unix()); compare != 0 {
return compare > 0
}
return openFlareAccessLogCompareInt64(openFlareAccessLogUintToInt64(uint64(left.ID)), openFlareAccessLogUintToInt64(uint64(right.ID))) > 0
})
}
func sortOpenFlareNodeObservationFrps(items []*OpenFlareNodeObservationFrps) {
sort.Slice(items, func(i, j int) bool {
left := items[i]
right := items[j]
if left == nil || right == nil {
return left != nil
}
if compare := openFlareAccessLogCompareInt64(left.CapturedAt.Unix(), right.CapturedAt.Unix()); compare != 0 {
return compare > 0
}
return openFlareAccessLogCompareInt64(openFlareAccessLogUintToInt64(uint64(left.ID)), openFlareAccessLogUintToInt64(uint64(right.ID))) > 0
})
}
func sortOpenFlareNodeObservationFrpc(items []*OpenFlareNodeObservationFrpc) {
sort.Slice(items, func(i, j int) bool {
left := items[i]
right := items[j]
if left == nil || right == nil {
return left != nil
}
if compare := openFlareAccessLogCompareInt64(left.CapturedAt.Unix(), right.CapturedAt.Unix()); compare != 0 {
return compare > 0
}
return openFlareAccessLogCompareInt64(openFlareAccessLogUintToInt64(uint64(left.ID)), openFlareAccessLogUintToInt64(uint64(right.ID))) > 0
})
}
func memoryLimitObservabilityRows[T any](rows []T, limit int) []T {
if limit <= 0 || len(rows) <= limit {
result := make([]T, len(rows))
copy(result, rows)
return result
}
result := make([]T, limit)
copy(result, rows[:limit])
return result
}
func cloneOpenFlareMetricSnapshot(record *OpenFlareMetricSnapshot) *OpenFlareMetricSnapshot {
copyRecord := *record
if copyRecord.ID == 0 {
copyRecord.ID = uint(idgen.NextUint64ID())
}
now := time.Now().UTC()
if copyRecord.CreatedAt.IsZero() {
copyRecord.CreatedAt = now
}
copyRecord.CapturedAt = copyRecord.CapturedAt.UTC()
copyRecord.CreatedAt = copyRecord.CreatedAt.UTC()
return &copyRecord
}
func cloneOpenFlareEdgeHealth(record *OpenFlareEdgeHealth) *OpenFlareEdgeHealth {
copyRecord := *record
if copyRecord.ID == 0 {
copyRecord.ID = uint(idgen.NextUint64ID())
}
now := time.Now().UTC()
if copyRecord.CreatedAt.IsZero() {
copyRecord.CreatedAt = now
}
if copyRecord.CapturedAt.IsZero() {
copyRecord.CapturedAt = now
}
if strings.TrimSpace(copyRecord.Status) == "" {
copyRecord.Status = edgeHealthStatusUnknown
}
copyRecord.CapturedAt = copyRecord.CapturedAt.UTC()
copyRecord.CreatedAt = copyRecord.CreatedAt.UTC()
return &copyRecord
}
func cloneOpenFlareNodeObservationFrps(record *OpenFlareNodeObservationFrps) *OpenFlareNodeObservationFrps {
copyRecord := *record
if copyRecord.ID == 0 {
copyRecord.ID = uint(idgen.NextUint64ID())
}
now := time.Now().UTC()
if copyRecord.CreatedAt.IsZero() {
copyRecord.CreatedAt = now
}
if copyRecord.CapturedAt.IsZero() {
copyRecord.CapturedAt = now
}
copyRecord.CapturedAt = copyRecord.CapturedAt.UTC()
copyRecord.CreatedAt = copyRecord.CreatedAt.UTC()
return &copyRecord
}
func cloneOpenFlareNodeObservationFrpc(record *OpenFlareNodeObservationFrpc) *OpenFlareNodeObservationFrpc {
copyRecord := *record
if copyRecord.ID == 0 {
copyRecord.ID = uint(idgen.NextUint64ID())
}
now := time.Now().UTC()
if copyRecord.CreatedAt.IsZero() {
copyRecord.CreatedAt = now
}
if copyRecord.CapturedAt.IsZero() {
copyRecord.CapturedAt = now
}
copyRecord.CapturedAt = copyRecord.CapturedAt.UTC()
copyRecord.CreatedAt = copyRecord.CreatedAt.UTC()
return &copyRecord
}
-88
View File
@@ -4,10 +4,7 @@
package model
import (
"context"
"time"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
)
// Origin OpenFlare 源站实体。
@@ -46,88 +43,3 @@ type OriginProxyRoute struct {
func (OriginProxyRoute) TableName() string {
return tableOfProxyRoutes
}
// HasProxyRoutesTable 判断代理规则表是否已迁移。
func HasProxyRoutesTable(ctx context.Context) bool {
return db.DB(ctx).Migrator().HasTable(&OriginProxyRoute{})
}
// ListOrigins 列出全部源站。
func ListOrigins(ctx context.Context) ([]Origin, error) {
var origins []Origin
if err := db.DB(ctx).Order("id desc").Find(&origins).Error; err != nil {
return nil, err
}
return origins, nil
}
// GetOriginByID 按 ID 查询源站。
func GetOriginByID(ctx context.Context, id uint) (*Origin, error) {
var origin Origin
if err := db.DB(ctx).First(&origin, id).Error; err != nil {
return nil, err
}
return &origin, nil
}
// GetOriginByAddress 按地址查询源站。
func GetOriginByAddress(ctx context.Context, address string) (*Origin, error) {
var origin Origin
if err := db.DB(ctx).Where("address = ?", address).First(&origin).Error; err != nil {
return nil, err
}
return &origin, nil
}
// CreateOriginRecord 创建源站。
func CreateOriginRecord(ctx context.Context, origin *Origin) error {
return db.DB(ctx).Create(origin).Error
}
// SaveOrigin 保存源站。
func SaveOrigin(ctx context.Context, origin *Origin) error {
return db.DB(ctx).Save(origin).Error
}
// DeleteOriginRecord 删除源站。
func DeleteOriginRecord(ctx context.Context, id uint) error {
return db.DB(ctx).Delete(&Origin{}, id).Error
}
// ListOriginRouteCounts 统计各源站关联的代理规则数量。
func ListOriginRouteCounts(ctx context.Context) ([]OriginRouteCount, error) {
if !HasProxyRoutesTable(ctx) {
return nil, nil
}
result := make([]OriginRouteCount, 0)
err := db.DB(ctx).Model(&OriginProxyRoute{}).
Select("origin_id, COUNT(*) AS route_count").
Where("origin_id IS NOT NULL").
Group("origin_id").
Scan(&result).Error
return result, err
}
// ListProxyRoutesByOriginID 列出源站关联的代理规则。
func ListProxyRoutesByOriginID(ctx context.Context, originID uint) ([]OriginProxyRoute, error) {
if !HasProxyRoutesTable(ctx) {
return nil, nil
}
var routes []OriginProxyRoute
if err := db.DB(ctx).Where("origin_id = ?", originID).Order("id desc").Find(&routes).Error; err != nil {
return nil, err
}
return routes, nil
}
// CountProxyRoutesByOriginID 统计源站关联的代理规则数量。
func CountProxyRoutesByOriginID(ctx context.Context, originID uint) (int64, error) {
if !HasProxyRoutesTable(ctx) {
return 0, nil
}
var count int64
if err := db.DB(ctx).Model(&OriginProxyRoute{}).Where("origin_id = ?", originID).Count(&count).Error; err != nil {
return 0, err
}
return count, nil
}
-88
View File
@@ -4,10 +4,7 @@
package model
import (
"context"
"time"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
)
// Pages deployment status constants.
@@ -83,88 +80,3 @@ type PagesDeploymentFile struct {
func (PagesDeploymentFile) TableName() string {
return "of_pages_deployment_files"
}
// HasPagesProjectsTable 判断 Pages 项目表是否已迁移。
func HasPagesProjectsTable(ctx context.Context) bool {
return db.DB(ctx).Migrator().HasTable(&PagesProject{})
}
// ListPagesProjects 列出全部 Pages 项目。
func ListPagesProjects(ctx context.Context) ([]PagesProject, error) {
var projects []PagesProject
if err := db.DB(ctx).Order("id desc").Find(&projects).Error; err != nil {
return nil, err
}
return projects, nil
}
// GetPagesProjectByID 按 ID 查询 Pages 项目。
func GetPagesProjectByID(ctx context.Context, id uint) (*PagesProject, error) {
var project PagesProject
if err := db.DB(ctx).First(&project, id).Error; err != nil {
return nil, err
}
return &project, nil
}
// GetPagesProjectBySlug 按 slug 查询 Pages 项目。
func GetPagesProjectBySlug(ctx context.Context, slug string) (*PagesProject, error) {
var project PagesProject
if err := db.DB(ctx).Where("slug = ?", slug).First(&project).Error; err != nil {
return nil, err
}
return &project, nil
}
// CreatePagesProjectRecord 创建 Pages 项目。
func CreatePagesProjectRecord(ctx context.Context, project *PagesProject) error {
return db.DB(ctx).Create(project).Error
}
// ListPagesDeployments 列出项目的全部部署。
func ListPagesDeployments(ctx context.Context, projectID uint) ([]PagesDeployment, error) {
var deployments []PagesDeployment
if err := db.DB(ctx).Where("project_id = ?", projectID).Order("id desc").Find(&deployments).Error; err != nil {
return nil, err
}
return deployments, nil
}
// GetPagesDeploymentByID 按 ID 查询 Pages 部署。
func GetPagesDeploymentByID(ctx context.Context, id uint) (*PagesDeployment, error) {
var deployment PagesDeployment
if err := db.DB(ctx).First(&deployment, id).Error; err != nil {
return nil, err
}
return &deployment, nil
}
// ListPagesDeploymentFiles 列出部署文件清单。
func ListPagesDeploymentFiles(ctx context.Context, deploymentID uint) ([]PagesDeploymentFile, error) {
var files []PagesDeploymentFile
if err := db.DB(ctx).Where("deployment_id = ?", deploymentID).Order("path asc").Find(&files).Error; err != nil {
return nil, err
}
return files, nil
}
// CountPagesDeploymentsByProjectID 统计项目部署数量。
func CountPagesDeploymentsByProjectID(ctx context.Context, projectID uint) (int64, error) {
var count int64
if err := db.DB(ctx).Model(&PagesDeployment{}).Where("project_id = ?", projectID).Count(&count).Error; err != nil {
return 0, err
}
return count, nil
}
// CountProxyRoutesByPagesProjectID 统计引用 Pages 项目的代理规则数量。
func CountProxyRoutesByPagesProjectID(ctx context.Context, projectID uint) (int64, error) {
if !HasProxyRoutesTable(ctx) {
return 0, nil
}
var count int64
if err := db.DB(ctx).Model(&ProxyRoute{}).Where("pages_project_id = ?", projectID).Count(&count).Error; err != nil {
return 0, err
}
return count, nil
}
+4 -53
View File
@@ -4,19 +4,16 @@
package model
import (
"context"
"errors"
"time"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
)
const (
// PagesOrphanUploadCandidateLimit bounds one delayed Pages upload cleanup pass.
PagesOrphanUploadCandidateLimit = 100
pagesOrphanMarkerPredicatePostgres = "w_uploads.metadata #>> '{extra,pages_ingest_marker}' = ?"
pagesOrphanMarkerPredicateSQLite = "CASE WHEN json_valid(w_uploads.metadata) THEN json_extract(w_uploads.metadata, '$.extra.pages_ingest_marker') ELSE NULL END = ?"
// PagesOrphanMarkerPredicatePostgres is the Postgres JSON marker match SQL fragment.
PagesOrphanMarkerPredicatePostgres = "w_uploads.metadata #>> '{extra,pages_ingest_marker}' = ?"
// PagesOrphanMarkerPredicateSQLite is the SQLite JSON marker match SQL fragment.
PagesOrphanMarkerPredicateSQLite = "CASE WHEN json_valid(w_uploads.metadata) THEN json_extract(w_uploads.metadata, '$.extra.pages_ingest_marker') ELSE NULL END = ?"
)
// PagesOrphanUploadCandidateQuery describes the fail-closed SQL candidate set
@@ -27,49 +24,3 @@ type PagesOrphanUploadCandidateQuery struct {
Marker string
CreatedBefore time.Time
}
// ListPagesOrphanUploadCandidates returns at most 100 unreferenced, isolated
// Pages V2 upload records. Callers must still lock and recheck every condition
// before deleting a candidate.
func ListPagesOrphanUploadCandidates(
ctx context.Context,
input PagesOrphanUploadCandidateQuery,
) ([]Upload, error) {
if input.SystemUserID == 0 || input.UploadType == "" || input.Marker == "" || input.CreatedBefore.IsZero() {
return nil, errors.New("invalid pages orphan upload candidate query")
}
markerPredicate, err := pagesOrphanMarkerPredicate(db.DB(ctx).Name())
if err != nil {
return nil, err
}
deploymentTable := (PagesDeployment{}).TableName()
uploadTable := (Upload{}).TableName()
var candidates []Upload
err = db.DB(ctx).
Model(&Upload{}).
Where(uploadTable+".status = ?", UploadStatusUsed).
Where(uploadTable+".user_id = ?", input.SystemUserID).
Where(uploadTable+".type = ?", input.UploadType).
Where(uploadTable+".created_at < ?", input.CreatedBefore).
Where(markerPredicate, input.Marker).
Where("NOT EXISTS (SELECT 1 FROM " + deploymentTable + " WHERE " + deploymentTable + ".upload_id = " + uploadTable + ".id)").
Order(uploadTable + ".id ASC").
Limit(PagesOrphanUploadCandidateLimit).
Find(&candidates).Error
if err != nil {
return nil, err
}
return candidates, nil
}
func pagesOrphanMarkerPredicate(dialect string) (string, error) {
switch dialect {
case "postgres":
return pagesOrphanMarkerPredicatePostgres, nil
case "sqlite":
return pagesOrphanMarkerPredicateSQLite, nil
default:
return "", errors.New("unsupported database dialect for Pages orphan cleanup")
}
}
@@ -1,193 +0,0 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package model
import (
"context"
"strings"
"testing"
"time"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/glebarez/sqlite"
"gorm.io/gorm"
)
func TestPagesOrphanMarkerPredicate(t *testing.T) {
tests := []struct {
name string
dialect string
want string
wantErr bool
}{
{
name: "postgres jsonb path",
dialect: "postgres",
want: "metadata #>> '{extra,pages_ingest_marker}'",
},
{
name: "sqlite guarded json extract",
dialect: "sqlite",
want: "CASE WHEN json_valid(w_uploads.metadata) THEN json_extract",
},
{
name: "unknown dialect rejected",
dialect: "mysql",
wantErr: true,
},
}
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
got, err := pagesOrphanMarkerPredicate(test.dialect)
if gotErr := err != nil; gotErr != test.wantErr {
t.Fatalf("pagesOrphanMarkerPredicate(%q) error = %v, want error presence = %t", test.dialect, err, test.wantErr)
}
if test.want != "" && !strings.Contains(got, test.want) {
t.Errorf("pagesOrphanMarkerPredicate(%q) = %q, want substring %q", test.dialect, got, test.want)
}
})
}
}
func TestListPagesOrphanUploadCandidatesFiltersAndLimits(t *testing.T) {
ctx := context.Background()
gormDB := setupPagesCleanupModelTestDB(t)
cutoff := time.Now().UTC().Add(-2 * time.Hour)
old := cutoff.Add(-time.Minute)
marker := UploadMetadata{Extra: map[string]any{
"pages_ingest_marker": "pages_deployment_v2",
"pages_project_id": "1",
}}
valid := make([]Upload, 0, PagesOrphanUploadCandidateLimit+1)
for index := 0; index < PagesOrphanUploadCandidateLimit+1; index++ {
valid = append(valid, pagesCleanupModelUpload(uint64(index+100), 999, "openflare_pages_deployment", UploadStatusUsed, old, marker))
}
if err := gormDB.Create(&valid).Error; err != nil {
t.Fatalf("create valid candidates error = %v, want nil", err)
}
referenced := pagesCleanupModelUpload(1, 999, "openflare_pages_deployment", UploadStatusUsed, old, marker)
wrongOwner := pagesCleanupModelUpload(2, 1000, "openflare_pages_deployment", UploadStatusUsed, old, marker)
wrongType := pagesCleanupModelUpload(3, 999, "generic", UploadStatusUsed, old, marker)
wrongStatus := pagesCleanupModelUpload(4, 999, "openflare_pages_deployment", UploadStatusPending, old, marker)
fresh := pagesCleanupModelUpload(5, 999, "openflare_pages_deployment", UploadStatusUsed, cutoff, marker)
wrongMarker := pagesCleanupModelUpload(6, 999, "openflare_pages_deployment", UploadStatusUsed, old, UploadMetadata{Extra: map[string]any{
"pages_ingest_marker": "pages_deployment_v1",
"pages_project_id": "1",
}})
for _, upload := range []Upload{referenced, wrongOwner, wrongType, wrongStatus, fresh, wrongMarker} {
if err := gormDB.Create(&upload).Error; err != nil {
t.Fatalf("create filtered upload %d error = %v, want nil", upload.ID, err)
}
}
if err := gormDB.Create(&PagesDeployment{
ProjectID: 1,
DeploymentNumber: 1,
Checksum: "referenced",
Status: PagesDeploymentStatusUploaded,
UploadID: referenced.ID,
}).Error; err != nil {
t.Fatalf("create referenced deployment error = %v, want nil", err)
}
invalidJSON := pagesCleanupModelUpload(7, 999, "openflare_pages_deployment", UploadStatusUsed, old, marker)
if err := gormDB.Create(&invalidJSON).Error; err != nil {
t.Fatalf("create invalid JSON upload error = %v, want nil", err)
}
if err := gormDB.Table((Upload{}).TableName()).Where("id = ?", invalidJSON.ID).
UpdateColumn("metadata", "{invalid").Error; err != nil {
t.Fatalf("corrupt upload metadata error = %v, want nil", err)
}
got, err := ListPagesOrphanUploadCandidates(ctx, PagesOrphanUploadCandidateQuery{
SystemUserID: 999,
UploadType: "openflare_pages_deployment",
Marker: "pages_deployment_v2",
CreatedBefore: cutoff,
})
if err != nil {
t.Fatalf("ListPagesOrphanUploadCandidates() error = %v, want nil", err)
}
if len(got) != PagesOrphanUploadCandidateLimit {
t.Fatalf("ListPagesOrphanUploadCandidates() count = %d, want %d", len(got), PagesOrphanUploadCandidateLimit)
}
for index, candidate := range got {
wantID := uint64(index + 100)
if candidate.ID != wantID {
t.Errorf("ListPagesOrphanUploadCandidates()[%d].ID = %d, want %d", index, candidate.ID, wantID)
}
}
}
func TestListPagesOrphanUploadCandidatesSkipsInvalidSQLiteJSON(t *testing.T) {
ctx := context.Background()
gormDB := setupPagesCleanupModelTestDB(t)
cutoff := time.Now().UTC().Add(-2 * time.Hour)
upload := pagesCleanupModelUpload(1, 999, "openflare_pages_deployment", UploadStatusUsed, cutoff.Add(-time.Minute), UploadMetadata{})
if err := gormDB.Create(&upload).Error; err != nil {
t.Fatalf("create invalid JSON candidate error = %v, want nil", err)
}
if err := gormDB.Table((Upload{}).TableName()).Where("id = ?", upload.ID).
UpdateColumn("metadata", "{invalid").Error; err != nil {
t.Fatalf("corrupt upload metadata error = %v, want nil", err)
}
got, err := ListPagesOrphanUploadCandidates(ctx, PagesOrphanUploadCandidateQuery{
SystemUserID: 999,
UploadType: "openflare_pages_deployment",
Marker: "pages_deployment_v2",
CreatedBefore: cutoff,
})
if err != nil {
t.Fatalf("ListPagesOrphanUploadCandidates(invalid JSON) error = %v, want nil", err)
}
if len(got) != 0 {
t.Errorf("ListPagesOrphanUploadCandidates(invalid JSON) count = %d, want 0", len(got))
}
}
func setupPagesCleanupModelTestDB(t *testing.T) *gorm.DB {
t.Helper()
gormDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{
DisableForeignKeyConstraintWhenMigrating: true,
})
if err != nil {
t.Fatalf("open Pages cleanup model test database error = %v, want nil", err)
}
if err := gormDB.AutoMigrate(&Upload{}, &PagesDeployment{}); err != nil {
t.Fatalf("migrate Pages cleanup model test database error = %v, want nil", err)
}
db.SetDB(gormDB)
t.Cleanup(func() { db.SetDB(nil) })
return gormDB
}
func pagesCleanupModelUpload(
id uint64,
userID uint64,
uploadType string,
status UploadStatus,
createdAt time.Time,
metadata UploadMetadata,
) Upload {
return Upload{
ID: id,
UserID: userID,
FileName: "site.zip",
FilePath: "pages/site.zip",
FileSize: 10,
MimeType: "application/zip",
Extension: "zip",
Hash: "checksum",
Type: uploadType,
Status: status,
AccessMode: 0,
Metadata: metadata,
CreatedAt: createdAt,
UpdatedAt: createdAt,
}
}
+16
View File
@@ -56,3 +56,19 @@ type PagesProjectSourceRuntime struct {
func (PagesProjectSourceRuntime) TableName() string {
return "of_pages_project_source_runtime"
}
// PagesExpiredSourceLeaseCandidate is a scanner query DTO for expired runtime leases.
type PagesExpiredSourceLeaseCandidate struct {
SourceID uint
LeaseToken string
LeaseExpiresAt time.Time
SyncStatus string
SourceType string
ReleaseSelector string
}
// PagesDueGitHubSourceCandidate is a scanner query DTO for due GitHub latest checks.
type PagesDueGitHubSourceCandidate struct {
SourceID uint
ConfigVersion int
}
-61
View File
@@ -4,10 +4,7 @@
package model
import (
"context"
"time"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
)
// ProxyRoute OpenFlare 代理规则实体。
@@ -47,61 +44,3 @@ type ProxyRoute struct {
func (ProxyRoute) TableName() string {
return tableOfProxyRoutes
}
// ListProxyRoutes 列出全部代理规则。
func ListProxyRoutes(ctx context.Context) ([]*ProxyRoute, error) {
var routes []*ProxyRoute
if err := db.DB(ctx).Order("id desc").Find(&routes).Error; err != nil {
return nil, err
}
return routes, nil
}
// GetProxyRouteByID 按 ID 查询代理规则。
func GetProxyRouteByID(ctx context.Context, id uint) (*ProxyRoute, error) {
var route ProxyRoute
if err := db.DB(ctx).First(&route, id).Error; err != nil {
return nil, err
}
return &route, nil
}
// CreateProxyRouteRecord 创建代理规则。
func CreateProxyRouteRecord(ctx context.Context, route *ProxyRoute) error {
return db.DB(ctx).Create(route).Error
}
// UpdateProxyRouteRecord 更新代理规则。
func UpdateProxyRouteRecord(ctx context.Context, route *ProxyRoute) error {
return db.DB(ctx).Model(&ProxyRoute{}).Where("id = ?", route.ID).Updates(map[string]any{
"site_name": route.SiteName,
"origin_id": route.OriginID,
"origin_url": route.OriginURL,
"origin_host": route.OriginHost,
"upstreams": route.Upstreams,
colEnabled: route.Enabled,
"enable_https": route.EnableHTTPS,
"redirect_http": route.RedirectHTTP,
"limit_conn_per_server": route.LimitConnPerServer,
"limit_conn_per_ip": route.LimitConnPerIP,
"limit_rate": route.LimitRate,
"limit_req_per_ip": route.LimitReqPerIP,
"cache_enabled": route.CacheEnabled,
"cache_policy": route.CachePolicy,
"cache_rules": route.CacheRules,
"custom_headers": route.CustomHeaders,
"basic_auth_enabled": route.BasicAuthEnabled,
"basic_auth_username": route.BasicAuthUsername,
"basic_auth_password": route.BasicAuthPassword,
"upstream_type": route.UpstreamType,
"tunnel_node_id": route.TunnelNodeID,
"tunnel_target_addr": route.TunnelTargetAddr,
"tunnel_target_protocol": route.TunnelTargetProtocol,
"pages_project_id": route.PagesProjectID,
}).Error
}
// DeleteProxyRouteRecord 删除代理规则。
func DeleteProxyRouteRecord(ctx context.Context, id uint) error {
return db.DB(ctx).Delete(&ProxyRoute{}, id).Error
}
-87
View File
@@ -4,11 +4,7 @@
package model
import (
"context"
"errors"
"time"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
)
// TLSCertificate OpenFlare TLS 证书实体。
@@ -54,86 +50,3 @@ type TLSProxyRouteRef struct {
func (TLSProxyRouteRef) TableName() string {
return tableOfProxyRoutes
}
// HasTLSProxyRoutesTable 判断代理规则表是否已迁移。
func HasTLSProxyRoutesTable(ctx context.Context) bool {
return db.DB(ctx).Migrator().HasTable(&TLSProxyRouteRef{})
}
// ListTLSCertificates 列出全部证书(不含 PEM 敏感字段的 JSON 暴露由 struct tag 控制)。
func ListTLSCertificates(ctx context.Context) ([]TLSCertificate, error) {
conn := db.DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
var certificates []TLSCertificate
if err := conn.Order("id desc").Find(&certificates).Error; err != nil {
return nil, err
}
return certificates, nil
}
// GetTLSCertificateByID 按 ID 查询证书。
func GetTLSCertificateByID(ctx context.Context, id uint) (*TLSCertificate, error) {
conn := db.DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
var certificate TLSCertificate
if err := conn.First(&certificate, id).Error; err != nil {
return nil, err
}
return &certificate, nil
}
// CreateTLSCertificateRecord 创建证书记录。
func CreateTLSCertificateRecord(ctx context.Context, certificate *TLSCertificate) error {
conn := db.DB(ctx)
if conn == nil {
return errors.New(errDatabaseNotInitialized)
}
return conn.Create(certificate).Error
}
// SaveTLSCertificate 保存证书记录。
func SaveTLSCertificate(ctx context.Context, certificate *TLSCertificate) error {
conn := db.DB(ctx)
if conn == nil {
return errors.New(errDatabaseNotInitialized)
}
return conn.Save(certificate).Error
}
// DeleteTLSCertificateRecord 删除证书记录。
func DeleteTLSCertificateRecord(ctx context.Context, id uint) error {
conn := db.DB(ctx)
if conn == nil {
return errors.New(errDatabaseNotInitialized)
}
return conn.Delete(&TLSCertificate{}, id).Error
}
// CountTLSCertificatesByDNSAccountID 统计引用指定 DNS 账号的证书数量。
func CountTLSCertificatesByDNSAccountID(ctx context.Context, dnsAccountID uint) (int64, error) {
conn := db.DB(ctx)
if conn == nil {
return 0, errors.New(errDatabaseNotInitialized)
}
var count int64
if err := conn.Model(&TLSCertificate{}).Where("dns_account_id = ?", dnsAccountID).Count(&count).Error; err != nil {
return 0, err
}
return count, nil
}
// ListTLSProxyRouteRefs 列出代理规则证书引用字段。
func ListTLSProxyRouteRefs(ctx context.Context) ([]TLSProxyRouteRef, error) {
if !HasTLSProxyRoutesTable(ctx) {
return nil, nil
}
var routes []TLSProxyRouteRef
if err := db.DB(ctx).Order("id asc").Find(&routes).Error; err != nil {
return nil, err
}
return routes, nil
}
-346
View File
@@ -4,12 +4,8 @@
package model
import (
"context"
"errors"
"time"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
"gorm.io/gorm"
)
// OpenFlareWAFRuleGroup stores a WAF rule group.
@@ -71,345 +67,3 @@ var ErrWAFRuleRevisionConflict = errors.New("waf rule revision conflict")
func (OpenFlareWAFRuleGroupBinding) TableName() string {
return "of_waf_rule_group_bindings"
}
func wafDB(ctx context.Context) (*gorm.DB, error) {
conn := db.DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
return conn, nil
}
// ListOpenFlareWAFRuleGroups returns all rule groups.
func ListOpenFlareWAFRuleGroups(ctx context.Context) ([]*OpenFlareWAFRuleGroup, error) {
conn, err := wafDB(ctx)
if err != nil {
return nil, err
}
var groups []*OpenFlareWAFRuleGroup
if err = conn.Order("is_global desc").Order("id asc").Find(&groups).Error; err != nil {
return nil, err
}
return groups, nil
}
// GetOpenFlareWAFRuleGroupByID returns a rule group by id.
func GetOpenFlareWAFRuleGroupByID(ctx context.Context, id uint) (*OpenFlareWAFRuleGroup, error) {
conn, err := wafDB(ctx)
if err != nil {
return nil, err
}
var group OpenFlareWAFRuleGroup
if err = conn.First(&group, id).Error; err != nil {
return nil, err
}
return &group, nil
}
// GetGlobalOpenFlareWAFRuleGroup returns the global rule group if present.
func GetGlobalOpenFlareWAFRuleGroup(ctx context.Context) (*OpenFlareWAFRuleGroup, error) {
conn, err := wafDB(ctx)
if err != nil {
return nil, err
}
var group OpenFlareWAFRuleGroup
if err = conn.Where("is_global = ?", true).Order("id asc").First(&group).Error; err != nil {
return nil, err
}
return &group, nil
}
// CreateOpenFlareWAFRuleGroup inserts a rule group.
func CreateOpenFlareWAFRuleGroup(ctx context.Context, group *OpenFlareWAFRuleGroup) error {
conn, err := wafDB(ctx)
if err != nil {
return err
}
return conn.Create(group).Error
}
// UpdateOpenFlareWAFRuleGroup updates mutable rule group fields.
func UpdateOpenFlareWAFRuleGroup(ctx context.Context, group *OpenFlareWAFRuleGroup) error {
conn, err := wafDB(ctx)
if err != nil {
return err
}
return conn.Model(&OpenFlareWAFRuleGroup{}).Where("id = ?", group.ID).Updates(map[string]any{
"name": group.Name,
colEnabled: group.Enabled,
"is_global": group.IsGlobal,
}).Error
}
// UpdateOpenFlareWAFRuleGraph atomically replaces a graph when revision is current.
func UpdateOpenFlareWAFRuleGraph(ctx context.Context, id uint, revision uint64, graph string) (uint64, error) {
conn, err := wafDB(ctx)
if err != nil {
return 0, err
}
result := conn.Model(&OpenFlareWAFRuleGroup{}).
Where("id = ? AND revision = ?", id, revision).
Updates(map[string]any{"graph": graph, "revision": gorm.Expr("revision + 1")})
if result.Error != nil {
return 0, result.Error
}
if result.RowsAffected != 1 {
return 0, ErrWAFRuleRevisionConflict
}
return revision + 1, nil
}
// DeleteOpenFlareWAFRuleGroup removes a rule group.
func DeleteOpenFlareWAFRuleGroup(ctx context.Context, id uint) error {
conn, err := wafDB(ctx)
if err != nil {
return err
}
return conn.Delete(&OpenFlareWAFRuleGroup{}, id).Error
}
// ListOpenFlareWAFIPGroups returns all IP groups.
func ListOpenFlareWAFIPGroups(ctx context.Context) ([]*OpenFlareWAFIPGroup, error) {
conn, err := wafDB(ctx)
if err != nil {
return nil, err
}
var groups []*OpenFlareWAFIPGroup
if err = conn.Order("type asc").Order("id asc").Find(&groups).Error; err != nil {
return nil, err
}
return groups, nil
}
// ListOpenFlareWAFIPGroupsByIDs returns IP groups for the given ids.
func ListOpenFlareWAFIPGroupsByIDs(ctx context.Context, ids []uint) ([]*OpenFlareWAFIPGroup, error) {
if len(ids) == 0 {
return []*OpenFlareWAFIPGroup{}, nil
}
conn, err := wafDB(ctx)
if err != nil {
return nil, err
}
var groups []*OpenFlareWAFIPGroup
if err = conn.Where("id IN ?", ids).Find(&groups).Error; err != nil {
return nil, err
}
return groups, nil
}
// GetOpenFlareWAFIPGroupByID returns an IP group by id.
func GetOpenFlareWAFIPGroupByID(ctx context.Context, id uint) (*OpenFlareWAFIPGroup, error) {
conn, err := wafDB(ctx)
if err != nil {
return nil, err
}
var group OpenFlareWAFIPGroup
if err = conn.First(&group, id).Error; err != nil {
return nil, err
}
return &group, nil
}
// CreateOpenFlareWAFIPGroup inserts an IP group.
func CreateOpenFlareWAFIPGroup(ctx context.Context, group *OpenFlareWAFIPGroup) error {
conn, err := wafDB(ctx)
if err != nil {
return err
}
return conn.Create(group).Error
}
// UpdateOpenFlareWAFIPGroup updates mutable IP group fields.
func UpdateOpenFlareWAFIPGroup(ctx context.Context, group *OpenFlareWAFIPGroup) error {
conn, err := wafDB(ctx)
if err != nil {
return err
}
return conn.Model(&OpenFlareWAFIPGroup{}).Where("id = ?", group.ID).Updates(map[string]any{
"name": group.Name,
"type": group.Type,
"enabled": group.Enabled,
"ip_list": group.IPList,
"auto_config": group.AutoConfig,
"ext_ips": group.ExtIPs,
"subscription_url": group.SubscriptionURL,
"subscription_format": group.SubscriptionFormat,
"subscription_mapping_rule": group.SubscriptionMappingRule,
"sync_interval_minutes": group.SyncIntervalMinutes,
"next_sync_at": group.NextSyncAt,
"last_sync_status": group.LastSyncStatus,
"last_sync_message": group.LastSyncMessage,
}).Error
}
// ListDueOpenFlareWAFIPGroups returns enabled automatic/subscription groups due for sync.
func ListDueOpenFlareWAFIPGroups(ctx context.Context, now time.Time) ([]*OpenFlareWAFIPGroup, error) {
conn, err := wafDB(ctx)
if err != nil {
return nil, err
}
var groups []*OpenFlareWAFIPGroup
err = conn.Where(
"enabled = ? AND (type = ? OR (type = ? AND subscription_url <> '')) AND (next_sync_at IS NULL OR next_sync_at <= ?)",
true, "automatic", "subscription", now,
).Order("id asc").Find(&groups).Error
return groups, err
}
// UpdateOpenFlareWAFIPGroupSyncResult persists IP group sync outcome fields.
func UpdateOpenFlareWAFIPGroupSyncResult(ctx context.Context, group *OpenFlareWAFIPGroup) error {
conn, err := wafDB(ctx)
if err != nil {
return err
}
return conn.Model(&OpenFlareWAFIPGroup{}).Where("id = ?", group.ID).Updates(map[string]any{
"ip_list": group.IPList,
"ext_ips": group.ExtIPs,
"last_synced_at": group.LastSyncedAt,
"next_sync_at": group.NextSyncAt,
"last_sync_status": group.LastSyncStatus,
"last_sync_message": group.LastSyncMessage,
"subscription_format": group.SubscriptionFormat,
}).Error
}
// DeleteOpenFlareWAFIPGroup removes an IP group.
func DeleteOpenFlareWAFIPGroup(ctx context.Context, id uint) error {
conn, err := wafDB(ctx)
if err != nil {
return err
}
return conn.Delete(&OpenFlareWAFIPGroup{}, id).Error
}
// ListOpenFlareWAFRuleGroupBindings returns all bindings.
func ListOpenFlareWAFRuleGroupBindings(ctx context.Context) ([]OpenFlareWAFRuleGroupBinding, error) {
conn, err := wafDB(ctx)
if err != nil {
return nil, err
}
var bindings []OpenFlareWAFRuleGroupBinding
if err = conn.Order("sequence asc").Order("id asc").Find(&bindings).Error; err != nil {
return nil, err
}
return bindings, nil
}
// ListOpenFlareWAFRuleGroupBindingsByRouteID returns bindings for a proxy route.
func ListOpenFlareWAFRuleGroupBindingsByRouteID(ctx context.Context, routeID uint) ([]OpenFlareWAFRuleGroupBinding, error) {
conn, err := wafDB(ctx)
if err != nil {
return nil, err
}
var bindings []OpenFlareWAFRuleGroupBinding
if err = conn.Where("proxy_route_id = ?", routeID).Order("sequence asc").Order("id asc").Find(&bindings).Error; err != nil {
return nil, err
}
return bindings, nil
}
func syncWAFBindingIDSequence(tx *gorm.DB) error {
if tx == nil || tx.Dialector.Name() != "postgres" { //nolint:staticcheck // QF1008: keep explicit Dialector field access
return nil
}
return tx.Exec(`
SELECT setval(
pg_get_serial_sequence('of_waf_rule_group_bindings', 'id'),
GREATEST(COALESCE((SELECT MAX(id) FROM of_waf_rule_group_bindings), 0), 1),
COALESCE((SELECT MAX(id) FROM of_waf_rule_group_bindings), 0) > 0
)
`).Error
}
func insertOpenFlareWAFRuleGroupBindings(tx *gorm.DB, bindings []OpenFlareWAFRuleGroupBinding) error {
if len(bindings) == 0 {
return nil
}
if err := syncWAFBindingIDSequence(tx); err != nil {
return err
}
return tx.Create(&bindings).Error
}
// ReplaceOpenFlareWAFRuleGroupBindings replaces bindings for a rule group.
func ReplaceOpenFlareWAFRuleGroupBindings(ctx context.Context, groupID uint, routeIDs []uint) error {
conn, err := wafDB(ctx)
if err != nil {
return err
}
return conn.Transaction(func(tx *gorm.DB) error {
if err = tx.Where("rule_group_id = ?", groupID).Delete(&OpenFlareWAFRuleGroupBinding{}).Error; err != nil {
return err
}
bindings := make([]OpenFlareWAFRuleGroupBinding, 0, len(routeIDs))
for index, routeID := range routeIDs {
bindings = append(bindings, OpenFlareWAFRuleGroupBinding{
RuleGroupID: groupID,
ProxyRouteID: routeID,
Sequence: index,
})
}
return insertOpenFlareWAFRuleGroupBindings(tx, bindings)
})
}
// ReplaceOpenFlareWAFSiteRuleGroupBindings replaces bindings for a proxy route.
func ReplaceOpenFlareWAFSiteRuleGroupBindings(ctx context.Context, routeID uint, groupIDs []uint) error {
conn, err := wafDB(ctx)
if err != nil {
return err
}
return conn.Transaction(func(tx *gorm.DB) error {
if err = tx.Where("proxy_route_id = ?", routeID).Delete(&OpenFlareWAFRuleGroupBinding{}).Error; err != nil {
return err
}
bindings := make([]OpenFlareWAFRuleGroupBinding, 0, len(groupIDs))
for index, groupID := range groupIDs {
bindings = append(bindings, OpenFlareWAFRuleGroupBinding{
RuleGroupID: groupID,
ProxyRouteID: routeID,
Sequence: index,
})
}
return insertOpenFlareWAFRuleGroupBindings(tx, bindings)
})
}
// DeleteOpenFlareWAFRuleGroupBindingsByGroupID removes bindings for a rule group.
func DeleteOpenFlareWAFRuleGroupBindingsByGroupID(ctx context.Context, groupID uint) error {
conn, err := wafDB(ctx)
if err != nil {
return err
}
return conn.Where("rule_group_id = ?", groupID).Delete(&OpenFlareWAFRuleGroupBinding{}).Error
}
// DeleteOpenFlareWAFRuleGroupWithBindings removes a rule group and its bindings.
func DeleteOpenFlareWAFRuleGroupWithBindings(ctx context.Context, groupID uint) error {
conn, err := wafDB(ctx)
if err != nil {
return err
}
return conn.Transaction(func(tx *gorm.DB) error {
if err = tx.Where("rule_group_id = ?", groupID).Delete(&OpenFlareWAFRuleGroupBinding{}).Error; err != nil {
return err
}
return tx.Delete(&OpenFlareWAFRuleGroup{}, groupID).Error
})
}
// GetOpenFlareProxyRouteByID returns a proxy route by id when the table exists.
func GetOpenFlareProxyRouteByID(ctx context.Context, id uint) (*OriginProxyRoute, error) {
if !HasProxyRoutesTable(ctx) {
return nil, gorm.ErrRecordNotFound
}
conn, err := wafDB(ctx)
if err != nil {
return nil, err
}
var route OriginProxyRoute
if err = conn.First(&route, id).Error; err != nil {
return nil, err
}
return &route, nil
}
@@ -1,54 +0,0 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package model
import (
"context"
"testing"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/glebarez/sqlite"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"gorm.io/gorm"
)
func setupWAFBindingsTestDB(t *testing.T) func() {
t.Helper()
sqliteDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{
DisableForeignKeyConstraintWhenMigrating: true,
})
require.NoError(t, err)
require.NoError(t, sqliteDB.AutoMigrate(&OpenFlareWAFRuleGroupBinding{}))
db.SetDB(sqliteDB)
return func() {
db.SetDB(nil)
}
}
func TestReplaceOpenFlareWAFRuleGroupBindingsAfterExplicitHighID(t *testing.T) {
cleanup := setupWAFBindingsTestDB(t)
defer cleanup()
ctx := context.Background()
conn := db.DB(ctx)
require.NotNil(t, conn)
require.NoError(t, conn.Create(&OpenFlareWAFRuleGroupBinding{
ID: 50,
RuleGroupID: 1,
ProxyRouteID: 1,
}).Error)
require.NoError(t, ReplaceOpenFlareWAFRuleGroupBindings(ctx, 2, []uint{2, 3}))
var bindings []OpenFlareWAFRuleGroupBinding
require.NoError(t, conn.Where("rule_group_id = ?", 2).Order("proxy_route_id asc").Find(&bindings).Error)
require.Len(t, bindings, 2)
assert.Equal(t, uint(2), bindings[0].ProxyRouteID)
assert.Equal(t, uint(3), bindings[1].ProxyRouteID)
assert.Greater(t, bindings[0].ID, uint(50))
assert.Greater(t, bindings[1].ID, bindings[0].ID)
}
-123
View File
@@ -1,123 +0,0 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package model
import (
"context"
"io/fs"
"os"
"path/filepath"
"runtime"
"testing"
"testing/fstest"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/glebarez/sqlite"
"github.com/pressly/goose/v3"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"gorm.io/gorm"
)
const defaultWAFRuleGraph = `{"schema_version":1,"nodes":[{"id":"start","type":"start","position":{"x":0,"y":0},"config":{}},{"id":"allow","type":"allow","position":{"x":320,"y":0},"config":{}}],"edges":[{"id":"start-allow","source":"start","source_handle":"next","target":"allow"}]}`
func wafMigrationFS(t *testing.T) fs.FS {
t.Helper()
_, filename, _, ok := runtime.Caller(0)
require.True(t, ok)
dir := filepath.Join(filepath.Dir(filename), "..", "db", "migrator", "goose", "sqlite")
migrations := fstest.MapFS{}
for _, name := range []string{
"202607150001_orchestrate_waf_rules.sql",
"202607150002_reset_waf_rule_graphs.sql",
"202607150003_drop_legacy_waf_rule_fields.sql",
} {
contents, err := os.ReadFile(filepath.Join(dir, name))
require.NoError(t, err)
migrations[name] = &fstest.MapFile{Data: contents}
}
return migrations
}
func TestOpenFlareWAFGraphMigrationResetsGraphsAndOrdersBindings(t *testing.T) {
conn, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
require.NoError(t, err)
sqlDB, err := conn.DB()
require.NoError(t, err)
require.NoError(t, conn.Exec(`CREATE TABLE of_waf_rule_groups (id INTEGER PRIMARY KEY AUTOINCREMENT, name TEXT NOT NULL, block_status_code INTEGER NOT NULL DEFAULT 418, block_response_body TEXT NOT NULL DEFAULT '', ip_whitelist TEXT NOT NULL DEFAULT '[]', ip_blacklist TEXT NOT NULL DEFAULT '[]', ip_whitelist_groups TEXT NOT NULL DEFAULT '[]', ip_blacklist_groups TEXT NOT NULL DEFAULT '[]', country_whitelist TEXT NOT NULL DEFAULT '[]', country_blacklist TEXT NOT NULL DEFAULT '[]', region_whitelist TEXT NOT NULL DEFAULT '[]', region_blacklist TEXT NOT NULL DEFAULT '[]', pow_enabled BOOLEAN NOT NULL DEFAULT 0, pow_config TEXT NOT NULL DEFAULT '{}')`).Error)
require.NoError(t, conn.Exec(`CREATE TABLE of_waf_rule_group_bindings (id INTEGER PRIMARY KEY AUTOINCREMENT, rule_group_id INTEGER NOT NULL, proxy_route_id INTEGER NOT NULL)`).Error)
require.NoError(t, conn.Exec(`INSERT INTO of_waf_rule_groups (id, name) VALUES (1, 'one'), (2, 'two')`).Error)
require.NoError(t, conn.Exec(`INSERT INTO of_waf_rule_group_bindings (id, rule_group_id, proxy_route_id) VALUES (20, 2, 7), (10, 1, 7)`).Error)
goose.SetBaseFS(wafMigrationFS(t))
require.NoError(t, goose.SetDialect("sqlite3"))
require.NoError(t, goose.Up(sqlDB, "."))
var groups []OpenFlareWAFRuleGroup
require.NoError(t, conn.Order("id asc").Find(&groups).Error)
require.Len(t, groups, 2)
for _, group := range groups {
require.JSONEq(t, defaultWAFRuleGraph, group.Graph)
assert.Equal(t, uint64(1), group.Revision)
}
for _, column := range []string{"block_status_code", "ip_whitelist", "pow_enabled"} {
assert.False(t, conn.Migrator().HasColumn("of_waf_rule_groups", column))
}
require.NoError(t, conn.Exec(`INSERT INTO of_waf_rule_groups (name) VALUES ('new')`).Error)
var newGroup OpenFlareWAFRuleGroup
require.NoError(t, conn.First(&newGroup, 3).Error)
assert.Empty(t, newGroup.Graph)
assert.Equal(t, uint64(1), newGroup.Revision)
var bindings []OpenFlareWAFRuleGroupBinding
require.NoError(t, conn.Where("proxy_route_id = ?", 7).Order("sequence asc").Order("id asc").Find(&bindings).Error)
require.Len(t, bindings, 2)
assert.Equal(t, []int{0, 1}, []int{bindings[0].Sequence, bindings[1].Sequence})
assert.Equal(t, []uint{10, 20}, []uint{bindings[0].ID, bindings[1].ID})
}
func TestOpenFlareWAFGraphOptimisticUpdate(t *testing.T) {
conn, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
require.NoError(t, err)
require.NoError(t, conn.AutoMigrate(&OpenFlareWAFRuleGroup{}))
db.SetDB(conn)
t.Cleanup(func() { db.SetDB(nil) })
group := OpenFlareWAFRuleGroup{Name: "rule", Graph: defaultWAFRuleGraph, Revision: 1}
require.NoError(t, conn.Create(&group).Error)
nextRevision, err := UpdateOpenFlareWAFRuleGraph(context.Background(), group.ID, 1, `{"schema_version":1}`)
require.NoError(t, err)
assert.Equal(t, uint64(2), nextRevision)
_, err = UpdateOpenFlareWAFRuleGraph(context.Background(), group.ID, 1, defaultWAFRuleGraph)
assert.ErrorIs(t, err, ErrWAFRuleRevisionConflict)
}
func TestReplaceOpenFlareWAFRuleGroupBindingsPreservesInputOrder(t *testing.T) {
cleanup := setupWAFBindingsTestDB(t)
defer cleanup()
ctx := context.Background()
require.NoError(t, ReplaceOpenFlareWAFSiteRuleGroupBindings(ctx, 7, []uint{30, 10, 20}))
bindings, err := ListOpenFlareWAFRuleGroupBindingsByRouteID(ctx, 7)
require.NoError(t, err)
require.Len(t, bindings, 3)
assert.Equal(t, []uint{30, 10, 20}, []uint{bindings[0].RuleGroupID, bindings[1].RuleGroupID, bindings[2].RuleGroupID})
assert.Equal(t, []int{0, 1, 2}, []int{bindings[0].Sequence, bindings[1].Sequence, bindings[2].Sequence})
}
func TestLegacyWAFColumnsRemoved(t *testing.T) {
conn, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
require.NoError(t, err)
require.NoError(t, conn.AutoMigrate(&OpenFlareWAFRuleGroup{}))
legacy := []string{"block_status_code", "block_response_body", "ip_whitelist", "ip_blacklist", "ip_whitelist_groups", "ip_blacklist_groups", "country_whitelist", "country_blacklist", "region_whitelist", "region_blacklist", "pow_enabled", "pow_config"}
for _, column := range legacy {
if conn.Migrator().HasColumn(&OpenFlareWAFRuleGroup{}, column) {
t.Fatalf("legacy WAF column %s still exists", column)
}
}
}
+4 -100
View File
@@ -4,14 +4,7 @@
package model
import (
"context"
"errors"
"fmt"
"time"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
"gorm.io/gorm"
"gorm.io/gorm/clause"
)
const (
@@ -19,8 +12,6 @@ const (
tableOfZoneDomains = "of_zone_domains"
)
var errZoneDomainBoundToAnotherRoute = errors.New("zone domain is already bound to another proxy route")
// Zone OpenFlare 注册根域实体。
type Zone struct {
ID uint `json:"id" gorm:"primaryKey;autoIncrement"`
@@ -50,95 +41,8 @@ func (ZoneDomain) TableName() string {
return tableOfZoneDomains
}
// ListZoneDomainsByRouteID returns the domains bound to a proxy route.
func ListZoneDomainsByRouteID(ctx context.Context, routeID uint) ([]ZoneDomain, error) {
var domains []ZoneDomain
if err := db.DB(ctx).Where("proxy_route_id = ?", routeID).Order("id asc").Find(&domains).Error; err != nil {
return nil, err
}
return domains, nil
}
// ListZoneDomainsByIDs returns explicit domains in the requested ID order.
func ListZoneDomainsByIDs(ctx context.Context, domainIDs []uint) ([]ZoneDomain, error) {
if len(domainIDs) == 0 {
return []ZoneDomain{}, nil
}
var domains []ZoneDomain
if err := db.DB(ctx).Where("id IN ?", domainIDs).Find(&domains).Error; err != nil {
return nil, err
}
byID := make(map[uint]ZoneDomain, len(domains))
for _, domain := range domains {
byID[domain.ID] = domain
}
ordered := make([]ZoneDomain, 0, len(domainIDs))
for _, id := range domainIDs {
domain, ok := byID[id]
if !ok {
return nil, fmt.Errorf("one or more zone domains do not exist")
}
ordered = append(ordered, domain)
}
return ordered, nil
}
// CountZoneDomainsByCertificateID reports whether a certificate is assigned to a Zone domain.
func CountZoneDomainsByCertificateID(ctx context.Context, certificateID uint) (int64, error) {
var count int64
err := db.DB(ctx).Model(&ZoneDomain{}).Where("cert_id = ?", certificateID).Count(&count).Error
return count, err
}
// ReplaceZoneDomainRouteBindings replaces every ZoneDomain binding for a proxy route.
func ReplaceZoneDomainRouteBindings(ctx context.Context, routeID uint, domainIDs []uint) error {
conn := db.DB(ctx)
if conn == nil {
return errors.New("database is not initialized")
}
return conn.Transaction(func(tx *gorm.DB) error {
var requested []ZoneDomain
if len(domainIDs) > 0 {
if err := tx.Clauses(clause.Locking{Strength: "UPDATE"}).
Where("id IN ?", domainIDs).
Find(&requested).Error; err != nil {
return err
}
if len(requested) != len(uniqueZoneDomainIDs(domainIDs)) {
return fmt.Errorf("one or more zone domains do not exist")
}
for _, domain := range requested {
if domain.ProxyRouteID != nil && *domain.ProxyRouteID != routeID {
return errZoneDomainBoundToAnotherRoute
}
}
}
current := tx.Model(&ZoneDomain{}).Where("proxy_route_id = ?", routeID)
if len(domainIDs) > 0 {
current = current.Where("id NOT IN ?", domainIDs)
}
if err := current.Update("proxy_route_id", nil).Error; err != nil {
return err
}
if len(domainIDs) == 0 {
return nil
}
return tx.Model(&ZoneDomain{}).Where("id IN ?", domainIDs).Update("proxy_route_id", routeID).Error
})
}
func uniqueZoneDomainIDs(domainIDs []uint) []uint {
seen := make(map[uint]struct{}, len(domainIDs))
ids := make([]uint, 0, len(domainIDs))
for _, id := range domainIDs {
if _, ok := seen[id]; ok {
continue
}
seen[id] = struct{}{}
ids = append(ids, id)
}
return ids
// ZoneDomainCount is the per-zone explicit domain count for list queries.
type ZoneDomainCount struct {
ZoneID uint `json:"zone_id" gorm:"column:zone_id"`
Count int64 `json:"count" gorm:"column:count"`
}
-88
View File
@@ -1,88 +0,0 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package model
import (
"context"
"testing"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/glebarez/sqlite"
"github.com/stretchr/testify/require"
"gorm.io/gorm"
)
func setupZoneTestDB(t *testing.T) *gorm.DB {
t.Helper()
sqliteDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{
DisableForeignKeyConstraintWhenMigrating: true,
})
require.NoError(t, err)
require.NoError(t, sqliteDB.AutoMigrate(&Zone{}, &ZoneDomain{}))
db.SetDB(sqliteDB)
t.Cleanup(func() { db.SetDB(nil) })
return sqliteDB
}
func TestReplaceZoneDomainRouteBindingsRejectsForeignDomain(t *testing.T) {
conn := setupZoneTestDB(t)
ctx := context.Background()
zone := Zone{Domain: "example.com"}
require.NoError(t, conn.Create(&zone).Error)
foreignRouteID := uint(11)
domain := ZoneDomain{
ZoneID: zone.ID,
ProxyRouteID: &foreignRouteID,
Domain: "api.example.com",
}
require.NoError(t, conn.Create(&domain).Error)
err := ReplaceZoneDomainRouteBindings(ctx, 12, []uint{domain.ID})
require.Error(t, err)
var got ZoneDomain
require.NoError(t, conn.First(&got, domain.ID).Error)
require.Equal(t, &foreignRouteID, got.ProxyRouteID)
}
func TestReplaceZoneDomainRouteBindingsReplacesCurrentRouteBindings(t *testing.T) {
conn := setupZoneTestDB(t)
ctx := context.Background()
zone := Zone{Domain: "example.com"}
require.NoError(t, conn.Create(&zone).Error)
routeID := uint(21)
boundDomain := ZoneDomain{ZoneID: zone.ID, ProxyRouteID: &routeID, Domain: "old.example.com"}
requestedDomain := ZoneDomain{ZoneID: zone.ID, Domain: "new.example.com"}
require.NoError(t, conn.Create(&boundDomain).Error)
require.NoError(t, conn.Create(&requestedDomain).Error)
require.NoError(t, ReplaceZoneDomainRouteBindings(ctx, routeID, []uint{requestedDomain.ID}))
var domains []ZoneDomain
require.NoError(t, conn.Order("id asc").Find(&domains).Error)
require.Len(t, domains, 2)
require.Nil(t, domains[0].ProxyRouteID)
require.Equal(t, &routeID, domains[1].ProxyRouteID)
}
func TestListZoneDomainsByRouteID(t *testing.T) {
conn := setupZoneTestDB(t)
ctx := context.Background()
zone := Zone{Domain: "example.com"}
require.NoError(t, conn.Create(&zone).Error)
routeID := uint(31)
boundDomain := ZoneDomain{ZoneID: zone.ID, ProxyRouteID: &routeID, Domain: "api.example.com"}
unboundDomain := ZoneDomain{ZoneID: zone.ID, Domain: "www.example.com"}
require.NoError(t, conn.Create(&boundDomain).Error)
require.NoError(t, conn.Create(&unboundDomain).Error)
domains, err := ListZoneDomainsByRouteID(ctx, routeID)
require.NoError(t, err)
require.Len(t, domains, 1)
require.Equal(t, boundDomain.ID, domains[0].ID)
}
-45
View File
@@ -4,10 +4,7 @@
package model
import (
"context"
"time"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
)
// Schedule 定时任务配置表
@@ -26,45 +23,3 @@ type Schedule struct {
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
}
-235
View File
@@ -5,15 +5,7 @@
package model
import (
"context"
"errors"
"fmt"
"strings"
"time"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/infra/persistence/idgen"
"github.com/redis/go-redis/v9"
)
// TaskExecutionStatus 任务执行状态
@@ -25,10 +17,6 @@ const (
TaskExecutionStatusRunning TaskExecutionStatus = "running"
TaskExecutionStatusSucceeded TaskExecutionStatus = "succeeded"
TaskExecutionStatusFailed TaskExecutionStatus = "failed"
taskExecutionLogRedisKeyPrefix = "task:execution:log:"
taskExecutionLogExpiration = 24 * time.Hour
taskExecutionLogMaxLines = 1000
)
// TaskExecution 任务执行记录
@@ -64,95 +52,6 @@ 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"`
@@ -160,137 +59,3 @@ type ListTaskExecutionsRequest struct {
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
}
-484
View File
@@ -1,484 +0,0 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package model
import (
"context"
"fmt"
"testing"
"time"
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(&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)
}
-75
View File
@@ -5,16 +5,13 @@
package model
import (
"context"
"errors"
"strconv"
"strings"
"time"
"github.com/Rain-kl/Wavelet/internal/infra/persistence/idgen"
"github.com/Rain-kl/Wavelet/internal/shared"
"github.com/Rain-kl/Wavelet/pkg/util"
"gorm.io/gorm"
)
// OAuthUserInfo 用户信息结构(同时支持 OIDC ID Token claims 和 UserEndpoint 响应)
@@ -104,14 +101,6 @@ func (u *User) CheckPassword(password string) bool {
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
@@ -129,67 +118,3 @@ func (u *User) CheckActive() error {
}
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
}