mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-03 15:06:36 +08:00
refactor(repository): 收敛 model/repository 分层为唯一持久化入口
将 OpenFlare 与平台业务的数据访问从 model 与 apps 直连迁入 repository, model 仅保留实体与无 IO 规则;补充 code-check 架构守卫与开发规范。
This commit is contained in:
@@ -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(¤t, "id = ?", source.ID).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if keepSecret {
|
||||
source.ClientSecret = current.ClientSecret
|
||||
}
|
||||
if err := source.Validate(); err != nil {
|
||||
return err
|
||||
}
|
||||
return db.DB(ctx).Model(¤t).Updates(map[string]any{
|
||||
colName: source.Name,
|
||||
"type": source.Type,
|
||||
"display_name": source.DisplayName,
|
||||
"is_active": source.IsActive,
|
||||
"client_id": source.ClientID,
|
||||
"client_secret": source.ClientSecret,
|
||||
"openid_discovery_url": source.OpenIDDiscoveryURL,
|
||||
"scopes": source.Scopes,
|
||||
"icon_url": source.IconURL,
|
||||
}).Error
|
||||
}
|
||||
|
||||
// ToggleAuthSource 切换认证源启用状态
|
||||
func ToggleAuthSource(ctx context.Context, id uint64, isActive bool) error {
|
||||
source, err := GetAuthSourceByID(ctx, id)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
source.IsActive = isActive
|
||||
if err := source.Validate(); err != nil {
|
||||
return err
|
||||
}
|
||||
return db.DB(ctx).Model(&AuthSource{}).Where("id = ?", id).Update("is_active", isActive).Error
|
||||
}
|
||||
|
||||
// DeleteAuthSource 删除认证源及其关联的外部帐号绑定
|
||||
func DeleteAuthSource(ctx context.Context, id uint64) error {
|
||||
if id == 0 {
|
||||
return errors.New(errAuthSourceIDRequired)
|
||||
}
|
||||
return db.DB(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
if err := tx.Where("auth_source_id = ?", id).Delete(&ExternalAccount{}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
return tx.Delete(&AuthSource{}, "id = ?", id).Error
|
||||
})
|
||||
}
|
||||
|
||||
// FindExternalAccount 查找外部帐号绑定记录
|
||||
func FindExternalAccount(ctx context.Context, sourceID uint64, externalID string) (*ExternalAccount, error) {
|
||||
var account ExternalAccount
|
||||
if err := db.DB(ctx).Where("auth_source_id = ? AND external_id = ?", sourceID, externalID).First(&account).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &account, nil
|
||||
}
|
||||
|
||||
// BindExternalAccount 绑定外部帐号(已存在时更新用户名和邮箱)
|
||||
func BindExternalAccount(ctx context.Context, account *ExternalAccount) error {
|
||||
if account.UserID == 0 || strings.TrimSpace(account.ExternalID) == "" {
|
||||
return errors.New(errExternalAccountBindingIncomplete)
|
||||
}
|
||||
account.ExternalID = strings.TrimSpace(account.ExternalID)
|
||||
account.ExternalUsername = strings.TrimSpace(account.ExternalUsername)
|
||||
account.Email = strings.TrimSpace(account.Email)
|
||||
|
||||
return db.DB(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
var current ExternalAccount
|
||||
err := tx.Where("auth_source_id = ? AND external_id = ?", account.AuthSourceID, account.ExternalID).First(¤t).Error
|
||||
if err == nil {
|
||||
if current.UserID != account.UserID {
|
||||
return errors.New(errExternalAccountAlreadyBoundToAnother)
|
||||
}
|
||||
return tx.Model(¤t).Updates(map[string]any{
|
||||
"external_username": account.ExternalUsername,
|
||||
"email": account.Email,
|
||||
}).Error
|
||||
}
|
||||
if !errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return err
|
||||
}
|
||||
return tx.Create(account).Error
|
||||
})
|
||||
}
|
||||
|
||||
// ListExternalAccountsByUserID 获取指定用户的所有外部帐号绑定视图
|
||||
func ListExternalAccountsByUserID(ctx context.Context, userID uint64) ([]ExternalAccountView, error) {
|
||||
if userID == 0 {
|
||||
return nil, errors.New(errUserIDRequired)
|
||||
}
|
||||
var accounts []ExternalAccount
|
||||
if err := db.DB(ctx).Where("user_id = ?", userID).Order("id asc").Find(&accounts).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
views := make([]ExternalAccountView, 0, len(accounts))
|
||||
for _, account := range accounts {
|
||||
var name, sourceType, label string
|
||||
if account.AuthSourceID == 0 {
|
||||
name = "default"
|
||||
sourceType = "oidc"
|
||||
label = "历史认证源"
|
||||
} else {
|
||||
source, err := GetAuthSourceByID(ctx, account.AuthSourceID)
|
||||
if err != nil {
|
||||
continue
|
||||
}
|
||||
name = source.Name
|
||||
sourceType = source.Type
|
||||
label = source.DisplayName
|
||||
if label == "" {
|
||||
label = source.Name
|
||||
}
|
||||
}
|
||||
views = append(views, ExternalAccountView{
|
||||
ID: account.ID,
|
||||
AuthSourceID: account.AuthSourceID,
|
||||
AuthSourceName: name,
|
||||
AuthSourceType: sourceType,
|
||||
AuthSourceLabel: label,
|
||||
ExternalUsername: account.ExternalUsername,
|
||||
Email: account.Email,
|
||||
CreatedAt: account.CreatedAt,
|
||||
})
|
||||
}
|
||||
return views, nil
|
||||
}
|
||||
|
||||
// DeleteExternalAccountForUser 删除指定用户的外部帐号绑定
|
||||
func DeleteExternalAccountForUser(ctx context.Context, id uint64, userID uint64) error {
|
||||
if id == 0 || userID == 0 {
|
||||
return errors.New(errExternalAccountBindingIDRequired)
|
||||
}
|
||||
return db.DB(ctx).Where("id = ? AND user_id = ?", id, userID).Delete(&ExternalAccount{}).Error
|
||||
}
|
||||
|
||||
+10
-24
@@ -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
|
||||
)
|
||||
|
||||
@@ -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, ©Record)
|
||||
}
|
||||
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] = ©Record
|
||||
}
|
||||
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
|
||||
})
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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 ©Record
|
||||
}
|
||||
|
||||
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 ©Record
|
||||
}
|
||||
|
||||
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 ©Record
|
||||
}
|
||||
|
||||
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 ©Record
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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,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,
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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,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"`
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user