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

将 OpenFlare 与平台业务的数据访问从 model 与 apps 直连迁入 repository,
model 仅保留实体与无 IO 规则;补充 code-check 架构守卫与开发规范。
This commit is contained in:
ryan
2026-07-24 17:00:17 +08:00
parent 23a5488203
commit 943818f7d4
184 changed files with 5592 additions and 4364 deletions
+69
View File
@@ -0,0 +1,69 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package repository
import (
"context"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/model"
)
// ListAccessTokensByUserID returns all access tokens for a user ordered by created_at desc.
func ListAccessTokensByUserID(ctx context.Context, userID uint64) ([]model.AccessToken, error) {
var tokens []model.AccessToken
if err := db.DB(ctx).Where("user_id = ?", userID).Order("created_at desc").Find(&tokens).Error; err != nil {
return nil, err
}
return tokens, nil
}
// CountAccessTokensByUserID returns how many access tokens a user owns.
func CountAccessTokensByUserID(ctx context.Context, userID uint64) (int64, error) {
var count int64
if err := db.DB(ctx).Model(&model.AccessToken{}).Where("user_id = ?", userID).Count(&count).Error; err != nil {
return 0, err
}
return count, nil
}
// CreateAccessToken inserts a new access token record.
func CreateAccessToken(ctx context.Context, record *model.AccessToken) error {
return db.DB(ctx).Create(record).Error
}
// GetAccessTokenByIDAndUserID loads a token owned by the given user.
func GetAccessTokenByIDAndUserID(ctx context.Context, id, userID uint64) (model.AccessToken, error) {
var token model.AccessToken
if err := db.DB(ctx).Where("id = ? AND user_id = ?", id, userID).First(&token).Error; err != nil {
return model.AccessToken{}, err
}
return token, nil
}
// DeleteAccessTokenForUser deletes a token if it belongs to the user.
// Returns the number of rows affected.
func DeleteAccessTokenForUser(ctx context.Context, id, userID uint64) (int64, error) {
tx := db.DB(ctx).Where("id = ? AND user_id = ?", id, userID).Delete(&model.AccessToken{})
return tx.RowsAffected, tx.Error
}
// GetAccessTokenByHash loads an access token by its token hash.
func GetAccessTokenByHash(ctx context.Context, tokenHash string) (model.AccessToken, error) {
var token model.AccessToken
if err := db.DB(ctx).Where("token_hash = ?", tokenHash).First(&token).Error; err != nil {
return model.AccessToken{}, err
}
return token, nil
}
// SaveAccessToken persists all fields of an existing access token.
func SaveAccessToken(ctx context.Context, record *model.AccessToken) error {
return db.DB(ctx).Save(record).Error
}
// DeleteAccessTokensByUserID deletes all access tokens for a user.
func DeleteAccessTokensByUserID(ctx context.Context, userID uint64) error {
return db.DB(ctx).Where("user_id = ?", userID).Delete(&model.AccessToken{}).Error
}
+215
View File
@@ -0,0 +1,215 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package repository
import (
"context"
"errors"
"strings"
"gorm.io/gorm"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/model"
)
// GetAuthSources 获取所有认证源(已脱敏)
func GetAuthSources(ctx context.Context) ([]model.AuthSource, error) {
var sources []model.AuthSource
if err := db.DB(ctx).Order("id asc").Find(&sources).Error; err != nil {
return nil, err
}
for i := range sources {
sources[i].Sanitize()
}
return sources, nil
}
// GetActiveAuthSources 获取所有已启用的认证源(已脱敏)
func GetActiveAuthSources(ctx context.Context) ([]model.AuthSource, error) {
var sources []model.AuthSource
if err := db.DB(ctx).Where("is_active = ?", true).Order("id asc").Find(&sources).Error; err != nil {
return nil, err
}
for i := range sources {
sources[i].Sanitize()
}
return sources, nil
}
// GetAuthSourceByID 根据 ID 获取认证源
func GetAuthSourceByID(ctx context.Context, id uint64) (*model.AuthSource, error) {
if id == 0 {
return nil, errors.New(errAuthSourceIDRequired)
}
var source model.AuthSource
if err := db.DB(ctx).First(&source, "id = ?", id).Error; err != nil {
return nil, err
}
source.ClientSecretConfigured = source.ClientSecret != ""
return &source, nil
}
// GetAuthSourceByName 根据名称获取认证源(名称比较不区分大小写)
func GetAuthSourceByName(ctx context.Context, name string) (*model.AuthSource, error) {
name = strings.TrimSpace(name)
if name == "" {
return nil, errors.New(errAuthSourceNameRequired)
}
var source model.AuthSource
if err := db.DB(ctx).First(&source, "LOWER(name) = LOWER(?)", name).Error; err != nil {
return nil, err
}
source.ClientSecretConfigured = source.ClientSecret != ""
return &source, nil
}
// CreateAuthSource 创建认证源
func CreateAuthSource(ctx context.Context, source *model.AuthSource) error {
if err := source.Validate(); err != nil {
return err
}
return db.DB(ctx).Create(source).Error
}
// UpdateAuthSource 更新认证源,keepSecret 为 true 时保留原密钥
func UpdateAuthSource(ctx context.Context, source *model.AuthSource, keepSecret bool) error {
if source.ID == 0 {
return errors.New(errAuthSourceIDRequired)
}
var current model.AuthSource
if err := db.DB(ctx).First(&current, "id = ?", source.ID).Error; err != nil {
return err
}
if keepSecret {
source.ClientSecret = current.ClientSecret
}
if err := source.Validate(); err != nil {
return err
}
return db.DB(ctx).Model(&current).Updates(map[string]any{
colName: source.Name,
"type": source.Type,
"display_name": source.DisplayName,
"is_active": source.IsActive,
"client_id": source.ClientID,
"client_secret": source.ClientSecret,
"openid_discovery_url": source.OpenIDDiscoveryURL,
"scopes": source.Scopes,
"icon_url": source.IconURL,
}).Error
}
// ToggleAuthSource 切换认证源启用状态
func ToggleAuthSource(ctx context.Context, id uint64, isActive bool) error {
source, err := GetAuthSourceByID(ctx, id)
if err != nil {
return err
}
source.IsActive = isActive
if err := source.Validate(); err != nil {
return err
}
return db.DB(ctx).Model(&model.AuthSource{}).Where("id = ?", id).Update("is_active", isActive).Error
}
// DeleteAuthSource 删除认证源及其关联的外部帐号绑定
func DeleteAuthSource(ctx context.Context, id uint64) error {
if id == 0 {
return errors.New(errAuthSourceIDRequired)
}
return db.DB(ctx).Transaction(func(tx *gorm.DB) error {
if err := tx.Where("auth_source_id = ?", id).Delete(&model.ExternalAccount{}).Error; err != nil {
return err
}
return tx.Delete(&model.AuthSource{}, "id = ?", id).Error
})
}
// FindExternalAccount 查找外部帐号绑定记录
func FindExternalAccount(ctx context.Context, sourceID uint64, externalID string) (*model.ExternalAccount, error) {
var account model.ExternalAccount
if err := db.DB(ctx).Where("auth_source_id = ? AND external_id = ?", sourceID, externalID).First(&account).Error; err != nil {
return nil, err
}
return &account, nil
}
// BindExternalAccount 绑定外部帐号(已存在时更新用户名和邮箱)
func BindExternalAccount(ctx context.Context, account *model.ExternalAccount) error {
if account.UserID == 0 || strings.TrimSpace(account.ExternalID) == "" {
return errors.New(errExternalAccountBindingIncomplete)
}
account.ExternalID = strings.TrimSpace(account.ExternalID)
account.ExternalUsername = strings.TrimSpace(account.ExternalUsername)
account.Email = strings.TrimSpace(account.Email)
return db.DB(ctx).Transaction(func(tx *gorm.DB) error {
var current model.ExternalAccount
err := tx.Where("auth_source_id = ? AND external_id = ?", account.AuthSourceID, account.ExternalID).First(&current).Error
if err == nil {
if current.UserID != account.UserID {
return errors.New(errExternalAccountAlreadyBoundToAnother)
}
return tx.Model(&current).Updates(map[string]any{
"external_username": account.ExternalUsername,
"email": account.Email,
}).Error
}
if !errors.Is(err, gorm.ErrRecordNotFound) {
return err
}
return tx.Create(account).Error
})
}
// ListExternalAccountsByUserID 获取指定用户的所有外部帐号绑定视图
func ListExternalAccountsByUserID(ctx context.Context, userID uint64) ([]model.ExternalAccountView, error) {
if userID == 0 {
return nil, errors.New(errUserIDRequired)
}
var accounts []model.ExternalAccount
if err := db.DB(ctx).Where("user_id = ?", userID).Order("id asc").Find(&accounts).Error; err != nil {
return nil, err
}
views := make([]model.ExternalAccountView, 0, len(accounts))
for _, account := range accounts {
var name, sourceType, label string
if account.AuthSourceID == 0 {
name = "default"
sourceType = "oidc"
label = "历史认证源"
} else {
source, err := GetAuthSourceByID(ctx, account.AuthSourceID)
if err != nil {
continue
}
name = source.Name
sourceType = source.Type
label = source.DisplayName
if label == "" {
label = source.Name
}
}
views = append(views, model.ExternalAccountView{
ID: account.ID,
AuthSourceID: account.AuthSourceID,
AuthSourceName: name,
AuthSourceType: sourceType,
AuthSourceLabel: label,
ExternalUsername: account.ExternalUsername,
Email: account.Email,
CreatedAt: account.CreatedAt,
})
}
return views, nil
}
// DeleteExternalAccountForUser 删除指定用户的外部帐号绑定
func DeleteExternalAccountForUser(ctx context.Context, id uint64, userID uint64) error {
if id == 0 || userID == 0 {
return errors.New(errExternalAccountBindingIDRequired)
}
return db.DB(ctx).Where("id = ? AND user_id = ?", id, userID).Delete(&model.ExternalAccount{}).Error
}
+3 -3
View File
@@ -178,7 +178,7 @@ func GetActiveAuthSourcesCached(ctx context.Context) ([]model.AuthSource, error)
}
}
sources, err := model.GetActiveAuthSources(ctx)
sources, err := GetActiveAuthSources(ctx)
if err != nil {
return nil, err
}
@@ -192,7 +192,7 @@ func GetAuthSourceByNameCached(ctx context.Context, name string) (*model.AuthSou
normalized := normalizeAuthSourceName(name)
if normalized == "" {
return model.GetAuthSourceByName(ctx, name)
return GetAuthSourceByName(ctx, name)
}
if source, ok := authSourceByNameRAM.GetIfPresent(normalized); ok {
@@ -210,7 +210,7 @@ func GetAuthSourceByNameCached(ctx context.Context, name string) (*model.AuthSou
}
}
source, err := model.GetAuthSourceByName(ctx, name)
source, err := GetAuthSourceByName(ctx, name)
if err != nil {
return nil, err
}
@@ -74,7 +74,7 @@ func TestGetActiveAuthSourcesCached_LoadsFromRedisBeforeDB(t *testing.T) {
ClientSecret: "client-secret",
OpenIDDiscoveryURL: "https://issuer.example.com",
}
if err := model.CreateAuthSource(ctx, &source); err != nil {
if err := CreateAuthSource(ctx, &source); err != nil {
t.Fatalf("CreateAuthSource() error = %v", err)
}
@@ -119,7 +119,7 @@ func TestGetAuthSourceByNameCached_LoadsFromRedisBeforeDB(t *testing.T) {
ClientSecret: "client-secret",
OpenIDDiscoveryURL: "https://issuer.example.com",
}
if err := model.CreateAuthSource(ctx, &source); err != nil {
if err := CreateAuthSource(ctx, &source); err != nil {
t.Fatalf("CreateAuthSource() error = %v", err)
}
@@ -164,7 +164,7 @@ func TestInvalidateAuthSourceCache_ClearsRedisKeys(t *testing.T) {
ClientSecret: "client-secret",
OpenIDDiscoveryURL: "https://issuer.example.com",
}
if err := model.CreateAuthSource(ctx, &source); err != nil {
if err := CreateAuthSource(ctx, &source); err != nil {
t.Fatalf("CreateAuthSource() error = %v", err)
}
if _, err := GetActiveAuthSourcesCached(ctx); err != nil {
@@ -209,7 +209,7 @@ func TestAuthSourceInvalidationPubSubClearsPeerRAM(t *testing.T) {
ClientSecret: "client-secret",
OpenIDDiscoveryURL: "https://issuer.example.com",
}
if err := model.CreateAuthSource(ctx, &source); err != nil {
if err := CreateAuthSource(ctx, &source); err != nil {
t.Fatalf("CreateAuthSource() error = %v", err)
}
+27
View File
@@ -0,0 +1,27 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package repository
// Persistence and repository-layer parameter messages live here (unexported).
// Domain field validation used by model.Validate stays in internal/model/errs.go;
// repository may call model.Validate and return those errors as-is.
// Keep wording aligned with model where the same user-facing phrase applies,
// but do not import or re-export model unexported consts (would require exporting).
const (
errDatabaseNotInitialized = "database not initialized"
errConfigIntParseFailed = "配置 %s 的值 '%s' 无法转换为整数: %w"
errConfigDecimalParseFailed = "配置 %s 的值 '%s' 无法转换为decimal: %w"
errConfigBoolParseFailed = "配置 %s 的值 '%s' 无法转换为布尔值: %w"
errParseMenuDisplayConfigFailed = "解析目录显示配置失败: %w"
errAuthSourceNameRequired = "认证源名称不能为空"
errAuthSourceIDRequired = "认证源 ID 不能为空"
errExternalAccountBindingIncomplete = "外部账号绑定信息不完整"
errExternalAccountAlreadyBoundToAnother = "该外部账号已绑定到其他用户"
errUserIDRequired = "用户 ID 不能为空"
errExternalAccountBindingIDRequired = "绑定记录 ID 不能为空"
)
const colName = "name"
const colEnabled = "enabled"
+476
View File
@@ -0,0 +1,476 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package repository
import (
"context"
"math"
"sort"
"strings"
"time"
analyticsmodel "github.com/Rain-kl/Wavelet/internal/model/analytics"
"github.com/Rain-kl/Wavelet/internal/model"
)
type openFlareAccessLogBucketAggregateRow = analyticsmodel.NodeAccessLogBucketAggregate
type openFlareAccessLogBucketDimensionRow = analyticsmodel.NodeAccessLogBucketDimension
type openFlareAccessLogIPAggregateRow = analyticsmodel.NodeAccessLogIPAggregate
type openFlareAccessLogIPSummaryRow = analyticsmodel.NodeAccessLogIPSummary
type openFlareAccessLogIPTrendRow = analyticsmodel.NodeAccessLogIPTrend
type openFlareAccessLogWAFIPAggregateRow = analyticsmodel.NodeAccessLogWAFIPAggregate
const (
sortOrderAsc = "asc"
columnRemoteAddr = "remote_addr"
columnHost = "host"
secondsPerMinute = 60
)
// ListOpenFlareAccessLogWAFIPAggregates returns per-IP aggregates for WAF automatic rules.
func ListOpenFlareAccessLogWAFIPAggregates(ctx context.Context, query model.OpenFlareAccessLogQuery) ([]*model.OpenFlareAccessLogWAFIPAggregate, error) {
rows, err := currentAccessLogStore().WAFIPAggregates(ctx, query)
if err != nil {
return nil, err
}
result := make([]*model.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, &model.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
}
// InsertOpenFlareAccessLogsBatch inserts access log rows into ClickHouse.
func InsertOpenFlareAccessLogsBatch(ctx context.Context, records []*model.OpenFlareAccessLog) error {
return currentAccessLogStore().InsertBatch(ctx, records)
}
// ListOpenFlareAccessLogs lists access logs matching the query.
func ListOpenFlareAccessLogs(ctx context.Context, query model.OpenFlareAccessLogQuery) ([]*model.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 model.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 model.OpenFlareAccessLogQuery) (model.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 model.OpenFlareAccessLogQuery, column string, limit int) ([]model.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 model.OpenFlareAccessLogQuery) ([]model.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) ([]*model.OpenFlareAccessLogRegionCount, error) {
return currentAccessLogStore().RegionCounts(ctx, nodeID, since, limit)
}
// ListOpenFlareAccessLogBuckets lists folded access log buckets.
func ListOpenFlareAccessLogBuckets(ctx context.Context, query model.OpenFlareAccessLogBucketQuery) ([]*model.OpenFlareAccessLogBucketRow, error) {
return buildOpenFlareAccessLogBucketRows(ctx, query)
}
// CountOpenFlareAccessLogBuckets counts folded access log buckets.
func CountOpenFlareAccessLogBuckets(ctx context.Context, query model.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 model.OpenFlareAccessLogBucketIPQuery) ([]*model.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 []*model.OpenFlareAccessLogBucketIPRow{}, nil
}
return rows[start:end], nil
}
// CountOpenFlareAccessLogBucketIPs counts folded IP rows for a bucket window.
func CountOpenFlareAccessLogBucketIPs(ctx context.Context, query model.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 model.OpenFlareAccessLogIPSummaryQuery, recentSince time.Time) ([]*analyticsmodel.NodeAccessLogIPSummary, error) {
return buildOpenFlareAccessLogIPSummaryRows(ctx, query, recentSince)
}
// CountOpenFlareAccessLogIPSummaries counts IP summaries.
func CountOpenFlareAccessLogIPSummaries(ctx context.Context, query model.OpenFlareAccessLogIPSummaryQuery) (int64, error) {
filter := openFlareAccessLogQueryFromIPSummary(query)
return currentAccessLogStore().CountIPSummaries(ctx, filter)
}
// ListOpenFlareAccessLogIPTrend lists IP trend points.
func ListOpenFlareAccessLogIPTrend(ctx context.Context, query model.OpenFlareAccessLogIPTrendQuery) ([]*analyticsmodel.NodeAccessLogIPTrend, error) {
remoteAddr := strings.TrimSpace(query.RemoteAddr)
if remoteAddr == "" {
return []*analyticsmodel.NodeAccessLogIPTrend{}, nil
}
filter := model.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([]*analyticsmodel.NodeAccessLogIPTrend, len(rows))
for index, row := range rows {
result[index] = &analyticsmodel.NodeAccessLogIPTrend{
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 model.OpenFlareAccessLogBucketQuery) ([]*model.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([]*model.OpenFlareAccessLogBucketRow, 0, len(partials))
for _, partial := range partials {
rows = append(rows, &model.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 model.OpenFlareAccessLogBucketIPQuery) ([]*model.OpenFlareAccessLogBucketIPRow, error) {
if query.BucketStartedAt.IsZero() {
return []*model.OpenFlareAccessLogBucketIPRow{}, nil
}
foldMinutes := query.FoldMinutes
if foldMinutes <= 0 {
foldMinutes = 3
}
bucketStartedAt := query.BucketStartedAt.UTC()
filter := model.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 model.OpenFlareAccessLogIPSummaryQuery, recentSince time.Time) ([]*analyticsmodel.NodeAccessLogIPSummary, error) {
filter := openFlareAccessLogQueryFromIPSummary(query)
partials, err := currentAccessLogStore().IPSummaries(ctx, filter, recentSince)
if err != nil {
return nil, err
}
rows := make([]*analyticsmodel.NodeAccessLogIPSummary, 0, len(partials))
for _, partial := range partials {
remoteAddr := strings.TrimSpace(partial.RemoteAddr)
if remoteAddr == "" {
continue
}
rows = append(rows, &analyticsmodel.NodeAccessLogIPSummary{
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 model.OpenFlareAccessLogQuery, exactRemoteAddr bool) ([]*model.OpenFlareAccessLogBucketIPRow, error) {
partials, err := currentAccessLogStore().IPAggregates(ctx, filter, exactRemoteAddr)
if err != nil {
return nil, err
}
rows := make([]*model.OpenFlareAccessLogBucketIPRow, 0, len(partials))
for _, partial := range partials {
remoteAddr := strings.TrimSpace(partial.RemoteAddr)
if remoteAddr == "" {
continue
}
rows = append(rows, &model.OpenFlareAccessLogBucketIPRow{
RemoteAddr: remoteAddr,
RequestCount: partial.RequestCount,
SuccessCount: partial.SuccessCount,
ClientErrorCount: partial.ClientErrorCount,
ServerErrorCount: partial.ServerErrorCount,
LastSeenEpoch: partial.LastSeenEpoch,
})
}
return rows, nil
}
func openFlareAccessLogQueryFromBucket(query model.OpenFlareAccessLogBucketQuery) model.OpenFlareAccessLogQuery {
return model.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 model.OpenFlareAccessLogIPSummaryQuery) model.OpenFlareAccessLogQuery {
return model.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 []*model.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 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 openFlareAccessLogCompareInt64(left int64, right int64) int {
switch {
case left > right:
return 1
case left < right:
return -1
default:
return 0
}
}
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 sortOpenFlareAccessLogBucketRows(items []*model.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 []*model.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 openFlareAccessLogUintToInt64(value uint64) int64 {
if value > math.MaxInt64 {
return math.MaxInt64
}
return int64(value)
}
@@ -0,0 +1,308 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package repository
import (
"context"
"math"
"sync"
"time"
"github.com/Rain-kl/Wavelet/internal/model"
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 []*model.OpenFlareAccessLog) error
List(ctx context.Context, query model.OpenFlareAccessLogQuery) ([]*model.OpenFlareAccessLog, error)
Count(ctx context.Context, query model.OpenFlareAccessLogQuery) (int64, int64, int64, error)
RegionCounts(ctx context.Context, nodeID string, since time.Time, limit int) ([]*model.OpenFlareAccessLogRegionCount, error)
BucketAggregates(ctx context.Context, filter model.OpenFlareAccessLogQuery, bucketSeconds int64) ([]openFlareAccessLogBucketAggregateRow, error)
CountBuckets(ctx context.Context, filter model.OpenFlareAccessLogQuery, bucketSeconds int64) (int64, error)
BucketDimensions(ctx context.Context, filter model.OpenFlareAccessLogQuery, column string, bucketSeconds int64) ([]openFlareAccessLogBucketDimensionRow, error)
IPAggregates(ctx context.Context, filter model.OpenFlareAccessLogQuery, exactRemoteAddr bool) ([]openFlareAccessLogIPAggregateRow, error)
WAFIPAggregates(ctx context.Context, filter model.OpenFlareAccessLogQuery) ([]openFlareAccessLogWAFIPAggregateRow, error)
IPSummaries(ctx context.Context, filter model.OpenFlareAccessLogQuery, recentSince time.Time) ([]openFlareAccessLogIPSummaryRow, error)
CountIPSummaries(ctx context.Context, filter model.OpenFlareAccessLogQuery) (int64, error)
IPTrend(ctx context.Context, filter model.OpenFlareAccessLogQuery, bucketSeconds int64) ([]openFlareAccessLogIPTrendRow, error)
TrafficSummary(ctx context.Context, filter model.OpenFlareAccessLogQuery) (model.OpenFlareAccessLogTrafficSummary, error)
ValueCounts(ctx context.Context, filter model.OpenFlareAccessLogQuery, column string, limit int) ([]model.OpenFlareAccessLogValueCount, error)
NodeAggregates(ctx context.Context, filter model.OpenFlareAccessLogQuery) ([]model.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)
}
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([]*model.OpenFlareAccessLog, 0),
}
}
type clickhouseAccessLogStore struct{}
func (clickhouseAccessLogStore) InsertBatch(_ context.Context, records []*model.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 model.OpenFlareAccessLogQuery) ([]*model.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 model.OpenFlareAccessLogQuery) (int64, int64, int64, error) {
return analyticsrepo.CountNodeAccessLogs(ctx, toNodeAccessLogFilter(query))
}
func (clickhouseAccessLogStore) RegionCounts(ctx context.Context, nodeID string, since time.Time, limit int) ([]*model.OpenFlareAccessLogRegionCount, error) {
rows, err := analyticsrepo.RegionCountsNodeAccessLogs(ctx, nodeID, since, limit)
if err != nil {
return nil, err
}
result := make([]*model.OpenFlareAccessLogRegionCount, len(rows))
for index, row := range rows {
result[index] = &model.OpenFlareAccessLogRegionCount{
Region: row.Region,
Count: row.Count,
}
}
return result, nil
}
func (clickhouseAccessLogStore) BucketAggregates(ctx context.Context, filter model.OpenFlareAccessLogQuery, bucketSeconds int64) ([]openFlareAccessLogBucketAggregateRow, error) {
return analyticsrepo.BucketAggregatesNodeAccessLogs(ctx, toNodeAccessLogFilter(filter), bucketSeconds)
}
func (clickhouseAccessLogStore) CountBuckets(ctx context.Context, filter model.OpenFlareAccessLogQuery, bucketSeconds int64) (int64, error) {
return analyticsrepo.CountBucketAggregatesNodeAccessLogs(ctx, toNodeAccessLogFilter(filter), bucketSeconds)
}
func (clickhouseAccessLogStore) BucketDimensions(ctx context.Context, filter model.OpenFlareAccessLogQuery, column string, bucketSeconds int64) ([]openFlareAccessLogBucketDimensionRow, error) {
return analyticsrepo.BucketDimensionsNodeAccessLogs(ctx, toNodeAccessLogFilter(filter), column, bucketSeconds)
}
func (clickhouseAccessLogStore) IPAggregates(ctx context.Context, filter model.OpenFlareAccessLogQuery, exactRemoteAddr bool) ([]openFlareAccessLogIPAggregateRow, error) {
return analyticsrepo.IPAggregatesNodeAccessLogs(ctx, toNodeAccessLogFilter(filter), exactRemoteAddr)
}
func (clickhouseAccessLogStore) IPSummaries(ctx context.Context, filter model.OpenFlareAccessLogQuery, recentSince time.Time) ([]openFlareAccessLogIPSummaryRow, error) {
return analyticsrepo.IPSummariesNodeAccessLogs(ctx, toNodeAccessLogFilter(filter), recentSince)
}
func (clickhouseAccessLogStore) CountIPSummaries(ctx context.Context, filter model.OpenFlareAccessLogQuery) (int64, error) {
return analyticsrepo.CountIPSummaryNodeAccessLogs(ctx, toNodeAccessLogFilter(filter))
}
func (clickhouseAccessLogStore) WAFIPAggregates(ctx context.Context, filter model.OpenFlareAccessLogQuery) ([]openFlareAccessLogWAFIPAggregateRow, error) {
return analyticsrepo.IPAggregatesForWAFNodeAccessLogs(ctx, toNodeAccessLogFilter(filter))
}
func (clickhouseAccessLogStore) IPTrend(ctx context.Context, filter model.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 model.OpenFlareAccessLogQuery) (model.OpenFlareAccessLogTrafficSummary, error) {
row, err := analyticsrepo.TrafficSummaryNodeAccessLogs(ctx, toNodeAccessLogFilter(filter))
if err != nil {
return model.OpenFlareAccessLogTrafficSummary{}, err
}
return model.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 model.OpenFlareAccessLogQuery, column string, limit int) ([]model.OpenFlareAccessLogValueCount, error) {
rows, err := analyticsrepo.ValueCountsNodeAccessLogs(ctx, toNodeAccessLogFilter(filter), column, limit)
if err != nil {
return nil, err
}
result := make([]model.OpenFlareAccessLogValueCount, len(rows))
for i, row := range rows {
result[i] = model.OpenFlareAccessLogValueCount{Value: row.Value, Count: row.Count}
}
return result, nil
}
func (clickhouseAccessLogStore) NodeAggregates(ctx context.Context, filter model.OpenFlareAccessLogQuery) ([]model.OpenFlareAccessLogNodeAggregate, error) {
rows, err := analyticsrepo.NodeAggregatesNodeAccessLogs(ctx, toNodeAccessLogFilter(filter))
if err != nil {
return nil, err
}
result := make([]model.OpenFlareAccessLogNodeAggregate, len(rows))
for i, row := range rows {
result[i] = model.OpenFlareAccessLogNodeAggregate{
NodeID: row.NodeID,
RequestCount: row.RequestCount,
ErrorCount: row.ErrorCount,
UniqueIPCount: row.UniqueIPCount,
}
}
return result, nil
}
func toNodeAccessLogFilter(query model.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 *model.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) []*model.OpenFlareAccessLog {
result := make([]*model.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] = &model.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
}
@@ -0,0 +1,710 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package repository
import (
"context"
"net"
"net/http"
"net/netip"
"sort"
"strconv"
"strings"
"sync"
"time"
"github.com/Rain-kl/Wavelet/internal/model"
"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 []*model.OpenFlareAccessLog
}
func (s *memoryAccessLogStore) InsertBatch(_ context.Context, records []*model.OpenFlareAccessLog) error {
s.mu.Lock()
defer s.mu.Unlock()
now := time.Now().UTC()
for _, record := range records {
if record == nil {
continue
}
copyRecord := *record
if copyRecord.ID == 0 {
copyRecord.ID = idgen.NextUint64ID()
}
if copyRecord.CreatedAt.IsZero() {
copyRecord.CreatedAt = now
}
copyRecord.LoggedAt = copyRecord.LoggedAt.UTC()
copyRecord.CreatedAt = copyRecord.CreatedAt.UTC()
s.records = append(s.records, &copyRecord)
}
return nil
}
func (s *memoryAccessLogStore) List(_ context.Context, query model.OpenFlareAccessLogQuery) ([]*model.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 model.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) ([]*model.OpenFlareAccessLogRegionCount, error) {
s.mu.RLock()
defer s.mu.RUnlock()
rows := s.filterRecords(model.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([]*model.OpenFlareAccessLogRegionCount, 0, len(counts))
for region, count := range counts {
result = append(result, &model.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 model.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([]*model.OpenFlareAccessLogBucketRow, len(result))
for index := range result {
bucketRows[index] = &model.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 model.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 model.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 model.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 model.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([]*model.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, &model.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 model.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 model.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 model.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([]*model.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([]*model.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 model.OpenFlareAccessLogQuery) (model.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 model.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 model.OpenFlareAccessLogQuery, column string, limit int) ([]model.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([]model.OpenFlareAccessLogValueCount, 0, len(counts))
for value, count := range counts {
result = append(result, model.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 model.OpenFlareAccessLogQuery) ([]model.OpenFlareAccessLogNodeAggregate, error) {
s.mu.RLock()
defer s.mu.RUnlock()
rows := s.filterRecords(filter)
type acc struct {
model.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: model.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([]model.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 model.OpenFlareAccessLogQuery) []*model.OpenFlareAccessLog {
result := make([]*model.OpenFlareAccessLog, 0, len(s.records))
for _, row := range s.records {
if !memoryAccessLogMatches(row, query) {
continue
}
result = append(result, row)
}
return result
}
func memoryAccessLogMatches(row *model.OpenFlareAccessLog, query model.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 []*model.OpenFlareAccessLog) []*model.OpenFlareAccessLog {
result := make([]*model.OpenFlareAccessLog, len(rows))
for index, row := range rows {
if row == nil {
continue
}
copyRecord := *row
result[index] = &copyRecord
}
return result
}
func sortOpenFlareAccessLogRows(items []*model.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
})
}
@@ -0,0 +1,121 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package repository
import (
"context"
"fmt"
"testing"
"time"
"github.com/Rain-kl/Wavelet/internal/model"
"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 := []*model.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 := &model.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, []*model.OpenFlareAccessLog{record}))
}
query := model.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 := model.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 := model.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, model.OpenFlareAccessLogQuery{Since: now.Add(-10 * time.Minute)})
require.NoError(t, err)
assert.Equal(t, int64(2), totalRecords)
}
@@ -0,0 +1,68 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package repository
import (
"context"
"errors"
"gorm.io/gorm"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/model"
)
// GetAcmeAccountByID 按 ID 查询 ACME 账号。
func GetAcmeAccountByID(ctx context.Context, id uint) (*model.AcmeAccount, error) {
conn := db.DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
var account model.AcmeAccount
if err := conn.First(&account, id).Error; err != nil {
return nil, err
}
return &account, nil
}
// CreateAcmeAccountRecord 创建 ACME 账号。
func CreateAcmeAccountRecord(ctx context.Context, account *model.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 *model.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) (*model.AcmeAccount, error) {
conn := db.DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
var account model.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 = model.AcmeAccount{
Email: "admin@openflare.dev",
}
if err = conn.Create(&account).Error; err != nil {
return nil, err
}
return &account, nil
}
+162
View File
@@ -0,0 +1,162 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package repository
import (
"context"
"errors"
"strings"
"time"
"gorm.io/gorm"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/model"
)
// ListOpenFlareApplyLogs returns apply logs ordered by id desc with optional pagination.
func ListOpenFlareApplyLogs(ctx context.Context, query model.OpenFlareApplyLogQuery) ([]*model.OpenFlareApplyLog, error) {
conn := db.DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
dbQuery := conn.Model(&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 []*model.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(&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) (*model.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 model.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
}
// GetLatestOpenFlareApplyLogsByNodeIDs returns the latest apply log per node id.
func GetLatestOpenFlareApplyLogsByNodeIDs(ctx context.Context, nodeIDs []string) (map[string]*model.OpenFlareApplyLog, error) {
result := make(map[string]*model.OpenFlareApplyLog)
if len(nodeIDs) == 0 {
return result, nil
}
conn := db.DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
var logs []*model.OpenFlareApplyLog
subQuery := conn.Model(&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
}
// CreateOpenFlareApplyLog inserts an apply log row.
func CreateOpenFlareApplyLog(ctx context.Context, log *model.OpenFlareApplyLog) error {
conn := db.DB(ctx)
if conn == nil {
return errors.New(errDatabaseNotInitialized)
}
return conn.Create(log).Error
}
// CreateOpenFlareApplyLogAndUpdateNode creates an apply log and updates the node from the apply result in one transaction.
func CreateOpenFlareApplyLogAndUpdateNode(ctx context.Context, log *model.OpenFlareApplyLog, applyResult, version, message string) error {
if log == nil {
return errors.New("apply log is required")
}
conn := db.DB(ctx)
if conn == nil {
return errors.New(errDatabaseNotInitialized)
}
now := log.CreatedAt
if now.IsZero() {
now = time.Now()
}
return conn.Transaction(func(tx *gorm.DB) error {
if err := tx.Create(log).Error; err != nil {
return err
}
return updateOpenFlareNodeFromApplyResultTx(tx, log.NodeID, applyResult, version, message, now)
})
}
// 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(&model.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(&model.OpenFlareApplyLog{})
return result.RowsAffected, result.Error
}
@@ -0,0 +1,75 @@
package repository
import (
"context"
"testing"
"time"
"github.com/Rain-kl/Wavelet/internal/model"
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(&model.OpenFlareApplyLog{}))
db.SetDB(sqliteDB)
return func() {
db.SetDB(nil)
}
}
func TestIsRepeatSuccessApplyLog(t *testing.T) {
latest := &model.OpenFlareApplyLog{
Version: "20260615-001",
Checksum: "checksum-a",
Result: "success",
}
assert.True(t, model.IsRepeatSuccessApplyLog(latest, "20260615-001", "checksum-a", "success"))
assert.False(t, model.IsRepeatSuccessApplyLog(latest, "20260615-002", "checksum-a", "success"))
assert.False(t, model.IsRepeatSuccessApplyLog(latest, "20260615-001", "checksum-b", "success"))
assert.False(t, model.IsRepeatSuccessApplyLog(latest, "20260615-001", "checksum-a", "failed"))
assert.False(t, model.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(&model.OpenFlareApplyLog{
NodeID: "node-1",
Version: "v1",
Result: "success",
Checksum: "checksum-1",
CreatedAt: now.Add(-time.Hour),
}).Error)
require.NoError(t, db.DB(ctx).Create(&model.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)
}
@@ -0,0 +1,135 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package repository
import (
"context"
"errors"
"gorm.io/gorm"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/model"
)
// ListConfigVersionSummaries returns config version summaries ordered by created_at desc.
func ListConfigVersionSummaries(ctx context.Context) ([]*model.ConfigVersionSummary, error) {
conn := db.DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
var versions []*model.ConfigVersionSummary
err := conn.Model(&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) (*model.ConfigVersion, error) {
conn := db.DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
var cv model.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) (*model.ConfigVersion, error) {
conn := db.DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
var version model.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 model.ConfigVersion
err := conn.Model(&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 *model.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 *model.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(&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(&model.ConfigVersion{}).Where("is_active = ?", true).Update("is_active", false).Error; err != nil {
return err
}
return tx.Model(&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(&model.ConfigVersion{})
return result.RowsAffected, result.Error
}
// ListEnabledProxyRoutes returns enabled proxy routes ordered by id asc.
func ListEnabledProxyRoutes(ctx context.Context) ([]*model.ProxyRoute, error) {
conn := db.DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
var routes []*model.ProxyRoute
if err := conn.Where("enabled = ?", true).Order("id asc").Find(&routes).Error; err != nil {
return nil, err
}
return routes, nil
}
@@ -0,0 +1,65 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package repository
import (
"context"
"errors"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/model"
)
// ListDNSAccounts 列出全部 DNS 账号(授权信息不通过 JSON 暴露)。
func ListDNSAccounts(ctx context.Context) ([]model.DNSAccount, error) {
conn := db.DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
var accounts []model.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) (*model.DNSAccount, error) {
conn := db.DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
var account model.DNSAccount
if err := conn.First(&account, id).Error; err != nil {
return nil, err
}
return &account, nil
}
// CreateDNSAccountRecord 创建 DNS 账号。
func CreateDNSAccountRecord(ctx context.Context, account *model.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 *model.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(&model.DNSAccount{}, id).Error
}
@@ -0,0 +1,77 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package repository
import (
"context"
"testing"
"time"
"github.com/Rain-kl/Wavelet/internal/model"
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 := &model.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 := []*model.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(), &model.OpenFlareMetricSnapshot{NodeID: "x"}); err != nil {
t.Fatalf("InsertMetricSnapshot with nil hook error = %v", err)
}
if err := (clickhouseAccessLogStore{}).InsertBatch(context.Background(), []*model.OpenFlareAccessLog{{NodeID: "x"}}); err != nil {
t.Fatalf("InsertBatch with nil hook error = %v", err)
}
}
+167
View File
@@ -0,0 +1,167 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package repository
import (
"context"
"errors"
"time"
"gorm.io/gorm"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/model"
)
const (
openFlareNodeStatusOnline = "online"
openFlareApplyResultSuccess = "success"
)
// ListOpenFlareNodes returns all nodes ordered by id desc.
func ListOpenFlareNodes(ctx context.Context) ([]model.OpenFlareNode, error) {
conn := db.DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
var nodes []model.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) ([]model.OpenFlareNode, error) {
if len(nodeIDs) == 0 {
return []model.OpenFlareNode{}, nil
}
conn := db.DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
var nodes []model.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) (*model.OpenFlareNode, error) {
conn := db.DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
var node model.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) (*model.OpenFlareNode, error) {
conn := db.DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
var node model.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) (*model.OpenFlareNode, error) {
conn := db.DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
var node model.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 *model.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 *model.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 *model.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
}
// UpdateOpenFlareNodeColumns updates node columns from a map of column values.
// Empty maps are no-ops.
func UpdateOpenFlareNodeColumns(ctx context.Context, node *model.OpenFlareNode, changes map[string]any) error {
if node == nil || len(changes) == 0 {
return nil
}
conn := db.DB(ctx)
if conn == nil {
return errors.New(errDatabaseNotInitialized)
}
return conn.Model(node).Updates(changes).Error
}
// UpdateOpenFlareNodeFromApplyResult updates node status, last_seen, version and last_error after an apply report.
// When applyResult is "success", current_version is set and last_error is cleared; otherwise last_error is set to message.
func UpdateOpenFlareNodeFromApplyResult(ctx context.Context, nodeID, applyResult, version, message string, now time.Time) error {
conn := db.DB(ctx)
if conn == nil {
return errors.New(errDatabaseNotInitialized)
}
return updateOpenFlareNodeFromApplyResultTx(conn, nodeID, applyResult, version, message, now)
}
func updateOpenFlareNodeFromApplyResultTx(tx *gorm.DB, nodeID, applyResult, version, message string, now time.Time) error {
record := &model.OpenFlareNode{}
if err := tx.Where("node_id = ?", nodeID).First(record).Error; err != nil {
return err
}
record.Status = openFlareNodeStatusOnline
lastSeen := now
record.LastSeenAt = &lastSeen
if applyResult == openFlareApplyResultSuccess {
record.CurrentVersion = version
record.LastError = ""
} else {
record.LastError = message
}
return tx.Model(record).Select("status", "last_seen_at", "current_version", "last_error").Updates(record).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(&model.OpenFlareNode{}, id).Error
}
@@ -0,0 +1,537 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package repository
import (
"context"
"encoding/json"
"errors"
"strings"
"time"
"gorm.io/gorm"
"gorm.io/gorm/clause"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/model"
analyticsrepo "github.com/Rain-kl/Wavelet/internal/repository/analytics"
)
const (
openFlareHealthEventStatusActive = "active"
openFlareHealthEventStatusResolved = "resolved"
openFlareHealthSeverityInfo = "info"
openFlareHealthSeverityWarning = "warning"
openFlareHealthSeverityCritical = "critical"
openFlareHealthEventMessageMaxLen = 4096
)
// OpenFlareHealthEventInput describes a desired active health event for reconciliation.
type OpenFlareHealthEventInput struct {
EventType string
Severity string
Message string
TriggeredAtUnix int64
Metadata map[string]string
}
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 *model.OpenFlareMetricSnapshot) error {
return currentObservabilityStore().InsertMetricSnapshot(ctx, record)
}
// InsertOpenFlareEdgeHealth inserts an L2 edge health snapshot into ClickHouse.
func InsertOpenFlareEdgeHealth(ctx context.Context, record *model.OpenFlareEdgeHealth) error {
return currentObservabilityStore().InsertEdgeHealth(ctx, record)
}
// InsertOpenFlareNodeObservationFrps inserts an FRPS observation into ClickHouse.
func InsertOpenFlareNodeObservationFrps(ctx context.Context, record *model.OpenFlareNodeObservationFrps) error {
return currentObservabilityStore().InsertNodeObservationFrps(ctx, record)
}
// InsertOpenFlareNodeObservationFrpc inserts an FRPC observation into ClickHouse.
func InsertOpenFlareNodeObservationFrpc(ctx context.Context, record *model.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) ([]*model.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) ([]*model.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 []*model.OpenFlareMetricSnapshot) []*model.OpenFlareMetricSnapshot {
latestByNode := make(map[string]*model.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([]*model.OpenFlareMetricSnapshot, 0, len(latestByNode))
for _, snapshot := range latestByNode {
result = append(result, snapshot)
}
return result
}
// 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) ([]*model.OpenFlareTrafficHourly, error) {
rows, err := analyticsrepo.ListNodeTrafficHourly(ctx, analyticsrepo.NodeObservabilityFilter{
NodeID: nodeID,
Since: since,
})
if err != nil {
return nil, err
}
result := make([]*model.OpenFlareTrafficHourly, len(rows))
for index, row := range rows {
result[index] = &model.OpenFlareTrafficHourly{
NodeID: row.NodeID,
Hour: row.Hour,
RequestCount: row.RequestCount,
ErrorCount: row.ErrorCount,
UniqueVisitorCount: row.UniqueVisitorCount,
}
}
return result, nil
}
// ListOpenFlareAccessLogHourlySince returns of_access_log_hourly rows since the given time.
func ListOpenFlareAccessLogHourlySince(ctx context.Context, nodeID string, since time.Time) ([]*model.OpenFlareAccessLogHourly, error) {
rows, err := analyticsrepo.ListAccessLogHourly(ctx, analyticsrepo.NodeObservabilityFilter{
NodeID: nodeID,
Since: since,
})
if err != nil {
return nil, err
}
result := make([]*model.OpenFlareAccessLogHourly, len(rows))
for index, row := range rows {
result[index] = &model.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
}
// ListOpenFlareMetricHourlySince returns hourly metric aggregates since the given time.
func ListOpenFlareMetricHourlySince(ctx context.Context, nodeID string, since time.Time) ([]*model.OpenFlareMetricHourly, error) {
rows, err := analyticsrepo.ListNodeMetricHourly(ctx, analyticsrepo.NodeObservabilityFilter{
NodeID: nodeID,
Since: since,
})
if err != nil {
return nil, err
}
result := make([]*model.OpenFlareMetricHourly, len(rows))
for index, row := range rows {
result[index] = &model.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) ([]*model.OpenFlareHealthEvent, error) {
conn := db.DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
var rows []*model.OpenFlareHealthEvent
if err := conn.Where("status = ?", "active").Order("last_triggered_at desc").Find(&rows).Error; err != nil {
if isMissingTableError(err) {
return []*model.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) ([]*model.OpenFlareHealthEvent, error) {
conn := db.DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
query := conn.Model(&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 []*model.OpenFlareHealthEvent
if err := query.Find(&rows).Error; err != nil {
if isMissingTableError(err) {
return []*model.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(&model.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) (*model.OpenFlareNodeSystemProfile, error) {
conn := db.DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
var profile model.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
}
// UpsertOpenFlareNodeSystemProfile inserts or updates the latest system profile for a node.
func UpsertOpenFlareNodeSystemProfile(ctx context.Context, record *model.OpenFlareNodeSystemProfile) error {
if record == nil {
return nil
}
conn := db.DB(ctx)
if conn == nil {
return errors.New(errDatabaseNotInitialized)
}
return upsertOpenFlareNodeSystemProfileTx(conn, record)
}
func upsertOpenFlareNodeSystemProfileTx(tx *gorm.DB, record *model.OpenFlareNodeSystemProfile) error {
return tx.Clauses(clause.OnConflict{
Columns: []clause.Column{{Name: "node_id"}},
DoUpdates: clause.AssignmentColumns([]string{
"hostname",
"os_name",
"os_version",
"kernel_version",
"architecture",
"cpu_model",
"cpu_cores",
"total_memory_bytes",
"total_disk_bytes",
"uptime_seconds",
"reported_at",
"updated_at",
}),
}).Create(record).Error
}
// ReconcileOpenFlareHealthEvents reconciles active health events for a node.
// Desired active events are created or updated; previously active types not present are resolved.
// When managedEventTypes is non-empty, only those event types are considered.
// Runs inside a transaction so multi-row create/update/resolve stays atomic.
func ReconcileOpenFlareHealthEvents(
ctx context.Context,
nodeID string,
events []OpenFlareHealthEventInput,
reportedAt time.Time,
managedEventTypes map[string]struct{},
) error {
conn := db.DB(ctx)
if conn == nil {
return errors.New(errDatabaseNotInitialized)
}
return conn.Transaction(func(tx *gorm.DB) error {
return reconcileOpenFlareHealthEventsTx(tx, nodeID, events, reportedAt, managedEventTypes)
})
}
// PersistOpenFlareNodePGObservability upserts an optional system profile and optionally reconciles
// health events in a single transaction (Postgres-side heartbeat observability).
// When reconcileHealth is false, health events are left untouched.
func PersistOpenFlareNodePGObservability(
ctx context.Context,
profile *model.OpenFlareNodeSystemProfile,
nodeID string,
events []OpenFlareHealthEventInput,
reconcileHealth bool,
reportedAt time.Time,
managedEventTypes map[string]struct{},
) error {
if profile == nil && !reconcileHealth {
return nil
}
conn := db.DB(ctx)
if conn == nil {
return errors.New(errDatabaseNotInitialized)
}
return conn.Transaction(func(tx *gorm.DB) error {
if profile != nil {
if err := upsertOpenFlareNodeSystemProfileTx(tx, profile); err != nil {
return err
}
}
if reconcileHealth {
if err := reconcileOpenFlareHealthEventsTx(tx, nodeID, events, reportedAt, managedEventTypes); err != nil {
return err
}
}
return nil
})
}
func reconcileOpenFlareHealthEventsTx(
tx *gorm.DB,
nodeID string,
events []OpenFlareHealthEventInput,
reportedAt time.Time,
managedEventTypes map[string]struct{},
) error {
activeTypes := make(map[string]OpenFlareHealthEventInput, len(events))
for _, event := range events {
eventType := normalizeOpenFlareHealthEventType(event.EventType)
if eventType == "" {
continue
}
if len(managedEventTypes) > 0 {
if _, ok := managedEventTypes[eventType]; !ok {
continue
}
}
event.EventType = eventType
event.Severity = normalizeOpenFlareHealthSeverity(event.Severity)
if event.TriggeredAtUnix <= 0 {
event.TriggeredAtUnix = reportedAt.Unix()
}
activeTypes[eventType] = event
}
var activeEvents []*model.OpenFlareHealthEvent
query := tx.Where("node_id = ? AND status = ?", nodeID, openFlareHealthEventStatusActive)
if len(managedEventTypes) > 0 {
scopedTypes := make([]string, 0, len(managedEventTypes))
for eventType := range managedEventTypes {
eventType = normalizeOpenFlareHealthEventType(eventType)
if eventType != "" {
scopedTypes = append(scopedTypes, eventType)
}
}
if len(scopedTypes) == 0 {
return nil
}
query = query.Where("event_type IN ?", scopedTypes)
}
if err := query.Find(&activeEvents).Error; err != nil {
return err
}
activeByType := make(map[string]*model.OpenFlareHealthEvent, len(activeEvents))
for _, event := range activeEvents {
activeByType[event.EventType] = event
}
for eventType, event := range activeTypes {
triggeredAt := timeFromUnixSeconds(event.TriggeredAtUnix, reportedAt)
if existing, ok := activeByType[eventType]; ok {
existing.Severity = event.Severity
existing.Message = normalizeOpenFlareHealthEventMessage(event.Message)
existing.LastTriggeredAt = triggeredAt
existing.ReportedAt = reportedAt
existing.MetadataJSON = marshalOpenFlareHealthMetadata(event.Metadata)
existing.ResolvedAt = nil
if err := tx.Save(existing).Error; err != nil {
return err
}
continue
}
record := &model.OpenFlareHealthEvent{
NodeID: nodeID,
EventType: eventType,
Severity: event.Severity,
Status: openFlareHealthEventStatusActive,
Message: normalizeOpenFlareHealthEventMessage(event.Message),
FirstTriggeredAt: triggeredAt,
LastTriggeredAt: triggeredAt,
ReportedAt: reportedAt,
MetadataJSON: marshalOpenFlareHealthMetadata(event.Metadata),
}
if err := tx.Create(record).Error; err != nil {
return err
}
}
for _, existing := range activeEvents {
if _, ok := activeTypes[existing.EventType]; ok {
continue
}
resolvedAt := reportedAt
existing.Status = openFlareHealthEventStatusResolved
existing.ReportedAt = reportedAt
existing.ResolvedAt = &resolvedAt
if err := tx.Save(existing).Error; err != nil {
return err
}
}
return nil
}
func normalizeOpenFlareHealthEventType(eventType string) string {
eventType = strings.TrimSpace(strings.ToLower(eventType))
eventType = strings.ReplaceAll(eventType, " ", "_")
return eventType
}
func normalizeOpenFlareHealthSeverity(severity string) string {
switch strings.ToLower(strings.TrimSpace(severity)) {
case openFlareHealthSeverityCritical:
return openFlareHealthSeverityCritical
case openFlareHealthSeverityInfo:
return openFlareHealthSeverityInfo
default:
return openFlareHealthSeverityWarning
}
}
func normalizeOpenFlareHealthEventMessage(message string) string {
if openFlareHealthEventMessageMaxLen <= 0 {
return ""
}
runes := []rune(strings.TrimSpace(message))
if len(runes) <= openFlareHealthEventMessageMaxLen {
return string(runes)
}
return string(runes[:openFlareHealthEventMessageMaxLen])
}
func timeFromUnixSeconds(unixSeconds int64, fallback time.Time) time.Time {
if unixSeconds <= 0 {
return fallback
}
return time.Unix(unixSeconds, 0).UTC()
}
func marshalOpenFlareHealthMetadata(value map[string]string) string {
if value == nil {
return ""
}
raw, err := json.Marshal(value)
if err != nil {
return ""
}
return string(raw)
}
// ListOpenFlareEdgeHealth returns L2 edge health snapshots.
func ListOpenFlareEdgeHealth(ctx context.Context, nodeID string, since time.Time, limit int) ([]*model.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) ([]*model.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) ([]*model.OpenFlareNodeObservationFrps, error) {
return currentObservabilityStore().ListNodeObservationFrps(ctx, nodeID, since, limit)
}
@@ -0,0 +1,355 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package repository
import (
"context"
"math"
"strings"
"sync"
"time"
"github.com/Rain-kl/Wavelet/internal/model"
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 *model.OpenFlareMetricSnapshot) error
ListMetricSnapshots(ctx context.Context, nodeID string, since time.Time, limit int) ([]*model.OpenFlareMetricSnapshot, error)
DeleteAllMetricSnapshots(ctx context.Context) (int64, error)
DeleteMetricSnapshotsBefore(ctx context.Context, cutoff time.Time) (int64, error)
InsertEdgeHealth(ctx context.Context, record *model.OpenFlareEdgeHealth) error
ListEdgeHealth(ctx context.Context, nodeID string, since time.Time, limit int) ([]*model.OpenFlareEdgeHealth, error)
DeleteAllEdgeHealth(ctx context.Context) (int64, error)
DeleteEdgeHealthBefore(ctx context.Context, cutoff time.Time) (int64, error)
InsertNodeObservationFrps(ctx context.Context, record *model.OpenFlareNodeObservationFrps) error
ListNodeObservationFrps(ctx context.Context, nodeID string, since time.Time, limit int) ([]*model.OpenFlareNodeObservationFrps, error)
DeleteAllNodeObservationFrps(ctx context.Context) (int64, error)
DeleteNodeObservationFrpsBefore(ctx context.Context, cutoff time.Time) (int64, error)
InsertNodeObservationFrpc(ctx context.Context, record *model.OpenFlareNodeObservationFrpc) error
ListNodeObservationFrpc(ctx context.Context, nodeID string, since time.Time, limit int) ([]*model.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 *model.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) ([]*model.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 *model.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) ([]*model.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 *model.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) ([]*model.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 *model.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) ([]*model.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 *model.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) []*model.OpenFlareMetricSnapshot {
result := make([]*model.OpenFlareMetricSnapshot, len(rows))
for index, row := range rows {
result[index] = &model.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 *model.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) []*model.OpenFlareEdgeHealth {
result := make([]*model.OpenFlareEdgeHealth, len(rows))
for index, row := range rows {
result[index] = &model.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 *model.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) []*model.OpenFlareNodeObservationFrps {
result := make([]*model.OpenFlareNodeObservationFrps, len(rows))
for index, row := range rows {
result[index] = &model.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 *model.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) []*model.OpenFlareNodeObservationFrpc {
result := make([]*model.OpenFlareNodeObservationFrpc, len(rows))
for index, row := range rows {
result[index] = &model.OpenFlareNodeObservationFrpc{
ID: uint(row.ID),
NodeID: row.NodeID,
CapturedAt: row.CapturedAt,
TunnelStatus: row.TunnelStatus,
ConnectedRelaysCount: int(row.ConnectedRelaysCount),
CreatedAt: row.CreatedAt,
}
}
return result
}
@@ -0,0 +1,409 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package repository
import (
"context"
"sort"
"strings"
"sync"
"time"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/infra/persistence/idgen"
)
type memoryObservabilityStore struct {
mu sync.RWMutex
metricSnapshots []*model.OpenFlareMetricSnapshot
edgeHealth []*model.OpenFlareEdgeHealth
frpsObs []*model.OpenFlareNodeObservationFrps
frpcObs []*model.OpenFlareNodeObservationFrpc
}
func (s *memoryObservabilityStore) InsertMetricSnapshot(_ context.Context, record *model.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) ([]*model.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([]*model.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 *model.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) ([]*model.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([]*model.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 *model.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) ([]*model.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([]*model.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 *model.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) ([]*model.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([]*model.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 []*model.OpenFlareMetricSnapshot, nodeID string, since time.Time) []*model.OpenFlareMetricSnapshot {
result := make([]*model.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 []*model.OpenFlareEdgeHealth, nodeID string, since time.Time) []*model.OpenFlareEdgeHealth {
result := make([]*model.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 []*model.OpenFlareNodeObservationFrps, nodeID string, since time.Time) []*model.OpenFlareNodeObservationFrps {
result := make([]*model.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 []*model.OpenFlareNodeObservationFrpc, nodeID string, since time.Time) []*model.OpenFlareNodeObservationFrpc {
result := make([]*model.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 []*model.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 []*model.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 []*model.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 []*model.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 []*model.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 *model.OpenFlareMetricSnapshot) *model.OpenFlareMetricSnapshot {
copyRecord := *record
if copyRecord.ID == 0 {
copyRecord.ID = uint(idgen.NextUint64ID())
}
now := time.Now().UTC()
if copyRecord.CreatedAt.IsZero() {
copyRecord.CreatedAt = now
}
copyRecord.CapturedAt = copyRecord.CapturedAt.UTC()
copyRecord.CreatedAt = copyRecord.CreatedAt.UTC()
return &copyRecord
}
func cloneOpenFlareEdgeHealth(record *model.OpenFlareEdgeHealth) *model.OpenFlareEdgeHealth {
copyRecord := *record
if copyRecord.ID == 0 {
copyRecord.ID = uint(idgen.NextUint64ID())
}
now := time.Now().UTC()
if copyRecord.CreatedAt.IsZero() {
copyRecord.CreatedAt = now
}
if copyRecord.CapturedAt.IsZero() {
copyRecord.CapturedAt = now
}
if strings.TrimSpace(copyRecord.Status) == "" {
copyRecord.Status = edgeHealthStatusUnknown
}
copyRecord.CapturedAt = copyRecord.CapturedAt.UTC()
copyRecord.CreatedAt = copyRecord.CreatedAt.UTC()
return &copyRecord
}
func cloneOpenFlareNodeObservationFrps(record *model.OpenFlareNodeObservationFrps) *model.OpenFlareNodeObservationFrps {
copyRecord := *record
if copyRecord.ID == 0 {
copyRecord.ID = uint(idgen.NextUint64ID())
}
now := time.Now().UTC()
if copyRecord.CreatedAt.IsZero() {
copyRecord.CreatedAt = now
}
if copyRecord.CapturedAt.IsZero() {
copyRecord.CapturedAt = now
}
copyRecord.CapturedAt = copyRecord.CapturedAt.UTC()
copyRecord.CreatedAt = copyRecord.CreatedAt.UTC()
return &copyRecord
}
func cloneOpenFlareNodeObservationFrpc(record *model.OpenFlareNodeObservationFrpc) *model.OpenFlareNodeObservationFrpc {
copyRecord := *record
if copyRecord.ID == 0 {
copyRecord.ID = uint(idgen.NextUint64ID())
}
now := time.Now().UTC()
if copyRecord.CreatedAt.IsZero() {
copyRecord.CreatedAt = now
}
if copyRecord.CapturedAt.IsZero() {
copyRecord.CapturedAt = now
}
copyRecord.CapturedAt = copyRecord.CapturedAt.UTC()
copyRecord.CreatedAt = copyRecord.CreatedAt.UTC()
return &copyRecord
}
+127
View File
@@ -0,0 +1,127 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package repository
import (
"context"
"gorm.io/gorm"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/model"
)
// WithOriginTx runs fn inside a database transaction for origin multi-step work.
func WithOriginTx(ctx context.Context, fn func(tx *gorm.DB) error) error {
return db.DB(ctx).Transaction(fn)
}
// HasProxyRoutesTable 判断代理规则表是否已迁移。
func HasProxyRoutesTable(ctx context.Context) bool {
return db.DB(ctx).Migrator().HasTable(&model.OriginProxyRoute{})
}
// ListOrigins 列出全部源站。
func ListOrigins(ctx context.Context) ([]model.Origin, error) {
var origins []model.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) (*model.Origin, error) {
var origin model.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) (*model.Origin, error) {
var origin model.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 *model.Origin) error {
return db.DB(ctx).Create(origin).Error
}
// SaveOrigin 保存源站。
func SaveOrigin(ctx context.Context, origin *model.Origin) error {
return SaveOriginTx(db.DB(ctx), origin)
}
// SaveOriginTx saves an origin within an existing transaction.
func SaveOriginTx(tx *gorm.DB, origin *model.Origin) error {
return tx.Save(origin).Error
}
// DeleteOriginRecord 删除源站。
func DeleteOriginRecord(ctx context.Context, id uint) error {
return db.DB(ctx).Delete(&model.Origin{}, id).Error
}
// ListOriginRouteCounts 统计各源站关联的代理规则数量。
func ListOriginRouteCounts(ctx context.Context) ([]model.OriginRouteCount, error) {
if !HasProxyRoutesTable(ctx) {
return nil, nil
}
result := make([]model.OriginRouteCount, 0)
err := db.DB(ctx).Model(&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) ([]model.OriginProxyRoute, error) {
if !HasProxyRoutesTable(ctx) {
return nil, nil
}
var routes []model.OriginProxyRoute
if err := db.DB(ctx).Where("origin_id = ?", originID).Order("id desc").Find(&routes).Error; err != nil {
return nil, err
}
return routes, nil
}
// ListProxyRoutesByOriginIDAscTx lists origin-linked proxy routes ordered by id asc within a transaction.
func ListProxyRoutesByOriginIDAscTx(tx *gorm.DB, originID uint) ([]model.OriginProxyRoute, error) {
var routes []model.OriginProxyRoute
if err := tx.Where("origin_id = ?", originID).Order("id asc").Find(&routes).Error; err != nil {
return nil, err
}
return routes, nil
}
// UpdateProxyRouteOriginAddressTx updates a proxy route's origin_url and upstreams within a transaction.
func UpdateProxyRouteOriginAddressTx(tx *gorm.DB, routeID uint, originURL, upstreamsJSON string) error {
return tx.Model(&model.OriginProxyRoute{}).
Where("id = ?", routeID).
Updates(map[string]any{
"origin_url": originURL,
"upstreams": upstreamsJSON,
}).Error
}
// 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(&model.OriginProxyRoute{}).Where("origin_id = ?", originID).Count(&count).Error; err != nil {
return 0, err
}
return count, nil
}
+96
View File
@@ -0,0 +1,96 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package repository
import (
"context"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/model"
)
// HasPagesProjectsTable 判断 Pages 项目表是否已迁移。
func HasPagesProjectsTable(ctx context.Context) bool {
return db.DB(ctx).Migrator().HasTable(&model.PagesProject{})
}
// ListPagesProjects 列出全部 Pages 项目。
func ListPagesProjects(ctx context.Context) ([]model.PagesProject, error) {
var projects []model.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) (*model.PagesProject, error) {
var project model.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) (*model.PagesProject, error) {
var project model.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 *model.PagesProject) error {
return db.DB(ctx).Create(project).Error
}
// ListPagesDeployments 列出项目的全部部署。
func ListPagesDeployments(ctx context.Context, projectID uint) ([]model.PagesDeployment, error) {
var deployments []model.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) (*model.PagesDeployment, error) {
var deployment model.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) ([]model.PagesDeploymentFile, error) {
var files []model.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(&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(&model.ProxyRoute{}).Where("pages_project_id = ?", projectID).Count(&count).Error; err != nil {
return 0, err
}
return count, nil
}
@@ -0,0 +1,58 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package repository
import (
"context"
"errors"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/model"
)
// 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 model.PagesOrphanUploadCandidateQuery,
) ([]model.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 := (model.PagesDeployment{}).TableName()
uploadTable := (model.Upload{}).TableName()
var candidates []model.Upload
err = db.DB(ctx).
Model(&model.Upload{}).
Where(uploadTable+".status = ?", model.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(model.PagesOrphanUploadCandidateLimit).
Find(&candidates).Error
if err != nil {
return nil, err
}
return candidates, nil
}
func pagesOrphanMarkerPredicate(dialect string) (string, error) {
switch dialect {
case "postgres":
return model.PagesOrphanMarkerPredicatePostgres, nil
case "sqlite":
return model.PagesOrphanMarkerPredicateSQLite, nil
default:
return "", errors.New("unsupported database dialect for Pages orphan cleanup")
}
}
@@ -0,0 +1,195 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package repository
import (
"context"
"strings"
"testing"
"time"
"github.com/Rain-kl/Wavelet/internal/model"
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 := model.UploadMetadata{Extra: map[string]any{
"pages_ingest_marker": "pages_deployment_v2",
"pages_project_id": "1",
}}
valid := make([]model.Upload, 0, model.PagesOrphanUploadCandidateLimit+1)
for index := 0; index < model.PagesOrphanUploadCandidateLimit+1; index++ {
valid = append(valid, pagesCleanupModelUpload(uint64(index+100), 999, "openflare_pages_deployment", model.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", model.UploadStatusUsed, old, marker)
wrongOwner := pagesCleanupModelUpload(2, 1000, "openflare_pages_deployment", model.UploadStatusUsed, old, marker)
wrongType := pagesCleanupModelUpload(3, 999, "generic", model.UploadStatusUsed, old, marker)
wrongStatus := pagesCleanupModelUpload(4, 999, "openflare_pages_deployment", model.UploadStatusPending, old, marker)
fresh := pagesCleanupModelUpload(5, 999, "openflare_pages_deployment", model.UploadStatusUsed, cutoff, marker)
wrongMarker := pagesCleanupModelUpload(6, 999, "openflare_pages_deployment", model.UploadStatusUsed, old, model.UploadMetadata{Extra: map[string]any{
"pages_ingest_marker": "pages_deployment_v1",
"pages_project_id": "1",
}})
for _, upload := range []model.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(&model.PagesDeployment{
ProjectID: 1,
DeploymentNumber: 1,
Checksum: "referenced",
Status: model.PagesDeploymentStatusUploaded,
UploadID: referenced.ID,
}).Error; err != nil {
t.Fatalf("create referenced deployment error = %v, want nil", err)
}
invalidJSON := pagesCleanupModelUpload(7, 999, "openflare_pages_deployment", model.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((model.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, model.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) != model.PagesOrphanUploadCandidateLimit {
t.Fatalf("ListPagesOrphanUploadCandidates() count = %d, want %d", len(got), model.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", model.UploadStatusUsed, cutoff.Add(-time.Minute), model.UploadMetadata{})
if err := gormDB.Create(&upload).Error; err != nil {
t.Fatalf("create invalid JSON candidate error = %v, want nil", err)
}
if err := gormDB.Table((model.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, model.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(&model.Upload{}, &model.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 model.UploadStatus,
createdAt time.Time,
metadata model.UploadMetadata,
) model.Upload {
return model.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,
}
}
@@ -0,0 +1,417 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package repository
import (
"context"
"time"
"gorm.io/gorm"
"gorm.io/gorm/clause"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/model"
)
const pagesRowLockStrength = "UPDATE"
// WithPagesTx runs fn inside a database transaction for Pages multi-step work.
func WithPagesTx(ctx context.Context, fn func(tx *gorm.DB) error) error {
return db.DB(ctx).Transaction(fn)
}
// GetPagesProjectSourceByID loads a project source by primary key.
func GetPagesProjectSourceByID(ctx context.Context, id uint) (*model.PagesProjectSource, error) {
var source model.PagesProjectSource
if err := db.DB(ctx).Where("id = ?", id).First(&source).Error; err != nil {
return nil, err
}
return &source, nil
}
// GetPagesProjectSourceByProjectID loads the unique source for a project.
func GetPagesProjectSourceByProjectID(ctx context.Context, projectID uint) (*model.PagesProjectSource, error) {
var source model.PagesProjectSource
if err := db.DB(ctx).Where("project_id = ?", projectID).First(&source).Error; err != nil {
return nil, err
}
return &source, nil
}
// GetPagesProjectSourceByIDAndConfigVersion loads a source matching both id and config version.
func GetPagesProjectSourceByIDAndConfigVersion(
ctx context.Context,
id uint,
configVersion int,
) (*model.PagesProjectSource, error) {
var source model.PagesProjectSource
if err := db.DB(ctx).Where("id = ? AND config_version = ?", id, configVersion).First(&source).Error; err != nil {
return nil, err
}
return &source, nil
}
// GetPagesProjectSourceRuntimeBySourceID loads runtime for a source.
func GetPagesProjectSourceRuntimeBySourceID(
ctx context.Context,
sourceID uint,
) (*model.PagesProjectSourceRuntime, error) {
var runtime model.PagesProjectSourceRuntime
if err := db.DB(ctx).Where("source_id = ?", sourceID).First(&runtime).Error; err != nil {
return nil, err
}
return &runtime, nil
}
// GetPagesProjectSourceAndRuntimeByProjectID loads source and its runtime for a project.
func GetPagesProjectSourceAndRuntimeByProjectID(
ctx context.Context,
projectID uint,
) (*model.PagesProjectSource, *model.PagesProjectSourceRuntime, error) {
source, err := GetPagesProjectSourceByProjectID(ctx, projectID)
if err != nil {
return nil, nil, err
}
runtime, err := GetPagesProjectSourceRuntimeBySourceID(ctx, source.ID)
if err != nil {
return nil, nil, err
}
return source, runtime, nil
}
// CreatePagesProjectSourceTx creates a source row inside an existing transaction.
func CreatePagesProjectSourceTx(tx *gorm.DB, source *model.PagesProjectSource) error {
return tx.Create(source).Error
}
// CreatePagesProjectSourceRuntimeTx creates a runtime row inside an existing transaction.
func CreatePagesProjectSourceRuntimeTx(tx *gorm.DB, runtime *model.PagesProjectSourceRuntime) error {
return tx.Create(runtime).Error
}
// UpdatePagesProjectSourceTx applies partial updates to a source inside a transaction.
func UpdatePagesProjectSourceTx(tx *gorm.DB, source *model.PagesProjectSource, updates map[string]any) error {
if len(updates) == 0 {
return nil
}
return tx.Model(source).Updates(updates).Error
}
// UpdatePagesProjectSourceRuntimeTx applies partial updates to a runtime inside a transaction.
func UpdatePagesProjectSourceRuntimeTx(
tx *gorm.DB,
runtime *model.PagesProjectSourceRuntime,
updates map[string]any,
) error {
if len(updates) == 0 {
return nil
}
return tx.Model(runtime).Updates(updates).Error
}
// UpdatePagesProjectSourceRuntimeFieldTx updates a single column on a runtime row.
func UpdatePagesProjectSourceRuntimeFieldTx(
tx *gorm.DB,
runtime *model.PagesProjectSourceRuntime,
column string,
value any,
) error {
return tx.Model(runtime).Update(column, value).Error
}
// DeletePagesProjectSourceRuntimeBySourceIDTx deletes runtime rows for a source.
func DeletePagesProjectSourceRuntimeBySourceIDTx(tx *gorm.DB, sourceID uint) error {
return tx.Where("source_id = ?", sourceID).Delete(&model.PagesProjectSourceRuntime{}).Error
}
// DeletePagesProjectSourceTx deletes a source row inside a transaction.
func DeletePagesProjectSourceTx(tx *gorm.DB, source *model.PagesProjectSource) error {
return tx.Delete(source).Error
}
// LockPagesProjectByIDTx locks a project row for update.
func LockPagesProjectByIDTx(tx *gorm.DB, id uint) (*model.PagesProject, error) {
var project model.PagesProject
if err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}).First(&project, id).Error; err != nil {
return nil, err
}
return &project, nil
}
// LockPagesProjectSourceByProjectIDTx locks the source for a project.
func LockPagesProjectSourceByProjectIDTx(tx *gorm.DB, projectID uint) (*model.PagesProjectSource, error) {
var source model.PagesProjectSource
if err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}).
Where("project_id = ?", projectID).
First(&source).Error; err != nil {
return nil, err
}
return &source, nil
}
// LockPagesProjectSourceByIDTx locks a source by id.
func LockPagesProjectSourceByIDTx(tx *gorm.DB, sourceID uint) (*model.PagesProjectSource, error) {
var source model.PagesProjectSource
if err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}).
Where("id = ?", sourceID).
First(&source).Error; err != nil {
return nil, err
}
return &source, nil
}
// LockPagesProjectSourceByIDAndProjectIDTx locks a source matching both identifiers.
func LockPagesProjectSourceByIDAndProjectIDTx(
tx *gorm.DB,
sourceID uint,
projectID uint,
) (*model.PagesProjectSource, error) {
var source model.PagesProjectSource
if err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}).
Where("id = ? AND project_id = ?", sourceID, projectID).
First(&source).Error; err != nil {
return nil, err
}
return &source, nil
}
// LockPagesProjectSourceRuntimeBySourceIDTx locks runtime for a source.
func LockPagesProjectSourceRuntimeBySourceIDTx(
tx *gorm.DB,
sourceID uint,
) (*model.PagesProjectSourceRuntime, error) {
var runtime model.PagesProjectSourceRuntime
if err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}).
Where("source_id = ?", sourceID).
First(&runtime).Error; err != nil {
return nil, err
}
return &runtime, nil
}
// GetPagesProjectSourceByIDTx loads a source by id without locking.
func GetPagesProjectSourceByIDTx(tx *gorm.DB, sourceID uint) (*model.PagesProjectSource, error) {
var source model.PagesProjectSource
if err := tx.Where("id = ?", sourceID).First(&source).Error; err != nil {
return nil, err
}
return &source, nil
}
// TryAcquirePagesSourceRuntimeLease conditionally claims an idle/expired lease when config matches.
func TryAcquirePagesSourceRuntimeLease(
ctx context.Context,
sourceID uint,
expectedConfigVersion int,
now time.Time,
updates map[string]any,
) (int64, error) {
result := db.DB(ctx).Model(&model.PagesProjectSourceRuntime{}).
Where("source_id = ?", sourceID).
Where("lease_expires_at IS NULL OR lease_expires_at <= ?", now).
Where(
"EXISTS (SELECT 1 FROM of_pages_project_sources source WHERE source.id = ? AND source.config_version = ?)",
sourceID,
expectedConfigVersion,
).
Updates(updates)
return result.RowsAffected, result.Error
}
// RenewPagesSourceRuntimeLease extends an active lease held by the given token.
func RenewPagesSourceRuntimeLease(
ctx context.Context,
sourceID uint,
token string,
now time.Time,
expiresAt time.Time,
) (int64, error) {
result := db.DB(ctx).Model(&model.PagesProjectSourceRuntime{}).
Where("source_id = ? AND lease_token = ? AND lease_expires_at > ?", sourceID, token, now).
Updates(map[string]any{"lease_expires_at": expiresAt})
return result.RowsAffected, result.Error
}
// UpdatePagesSourceRuntimeByActiveLease updates runtime while the caller still owns the lease.
func UpdatePagesSourceRuntimeByActiveLease(
ctx context.Context,
sourceID uint,
token string,
now time.Time,
updates map[string]any,
) (int64, error) {
return UpdatePagesSourceRuntimeByActiveLeaseTx(db.DB(ctx), sourceID, token, now, updates)
}
// UpdatePagesSourceRuntimeByActiveLeaseTx updates runtime under an active lease inside a transaction.
func UpdatePagesSourceRuntimeByActiveLeaseTx(
tx *gorm.DB,
sourceID uint,
token string,
now time.Time,
updates map[string]any,
) (int64, error) {
result := tx.Model(&model.PagesProjectSourceRuntime{}).
Where("source_id = ? AND lease_token = ? AND lease_expires_at > ?", sourceID, token, now).
Updates(updates)
return result.RowsAffected, result.Error
}
// RecoverExpiredPagesSourceRuntimeLease clears one exact expired lease owner.
func RecoverExpiredPagesSourceRuntimeLease(
ctx context.Context,
sourceID uint,
token string,
expiresAt time.Time,
status string,
now time.Time,
updates map[string]any,
) (int64, error) {
result := db.DB(ctx).Model(&model.PagesProjectSourceRuntime{}).
Where("source_id = ?", sourceID).
Where("lease_token = ?", token).
Where("lease_expires_at = ? AND lease_expires_at <= ?", expiresAt, now).
Where("sync_status = ?", status).
Updates(updates)
return result.RowsAffected, result.Error
}
// MarkPagesSourceInitialCheckDispatchFailed marks runtime failed when config still matches and lease is free.
func MarkPagesSourceInitialCheckDispatchFailed(
ctx context.Context,
sourceID uint,
configVersion int,
now time.Time,
updates map[string]any,
) (int64, error) {
result := db.DB(ctx).Model(&model.PagesProjectSourceRuntime{}).
Where("source_id = ?", sourceID).
Where("lease_expires_at IS NULL OR lease_expires_at <= ?", now).
Where(
"EXISTS (SELECT 1 FROM of_pages_project_sources source WHERE source.id = ? AND source.config_version = ?)",
sourceID,
configVersion,
).
Updates(updates)
return result.RowsAffected, result.Error
}
// RecordPagesSourceAutoDispatchFailure records a failed auto-sync dispatch while status still matches.
func RecordPagesSourceAutoDispatchFailure(
ctx context.Context,
sourceID uint,
configVersion int,
sourceType string,
releaseSelector string,
revision string,
updateAvailableStatus string,
now time.Time,
updates map[string]any,
) (int64, error) {
result := db.DB(ctx).Model(&model.PagesProjectSourceRuntime{}).
Where("source_id = ?", sourceID).
Where("sync_status = ? AND last_seen_revision = ?", updateAvailableStatus, revision).
Where("lease_expires_at IS NULL OR lease_expires_at <= ?", now).
Where(`EXISTS (
SELECT 1 FROM of_pages_project_sources AS source
WHERE source.id = ? AND source.config_version = ?
AND source.source_type = ? AND source.release_selector = ?
AND source.auto_update_enabled = ?
)`,
sourceID,
configVersion,
sourceType,
releaseSelector,
true,
).
Updates(updates)
return result.RowsAffected, result.Error
}
// ListExpiredPagesSourceLeaseCandidates returns expired checking/syncing leases for recovery.
func ListExpiredPagesSourceLeaseCandidates(
ctx context.Context,
now time.Time,
syncStatuses []string,
) ([]model.PagesExpiredSourceLeaseCandidate, error) {
var candidates []model.PagesExpiredSourceLeaseCandidate
err := db.DB(ctx).
Table("of_pages_project_source_runtime AS runtime").
Select(`runtime.source_id, runtime.lease_token, runtime.lease_expires_at,
runtime.sync_status, source.source_type, source.release_selector`).
Joins("JOIN of_pages_project_sources AS source ON source.id = runtime.source_id").
Where("runtime.lease_token <> ''").
Where("runtime.lease_expires_at IS NOT NULL AND runtime.lease_expires_at <= ?", now).
Where("runtime.sync_status IN ?", syncStatuses).
Order("runtime.source_id ASC").
Scan(&candidates).Error
if err != nil {
return nil, err
}
return candidates, nil
}
// CountDueGitHubPagesSourceChecks counts due latest GitHub sources.
func CountDueGitHubPagesSourceChecks(
ctx context.Context,
now time.Time,
sourceType string,
releaseSelector string,
) (int64, error) {
var count int64
err := dueGitHubPagesSourceQuery(ctx, now, sourceType, releaseSelector).Count(&count).Error
return count, err
}
// ListDueGitHubPagesSourceChecks lists a batch of due latest GitHub sources in stable order.
func ListDueGitHubPagesSourceChecks(
ctx context.Context,
now time.Time,
sourceType string,
releaseSelector string,
limit int,
) ([]model.PagesDueGitHubSourceCandidate, error) {
var candidates []model.PagesDueGitHubSourceCandidate
err := dueGitHubPagesSourceQuery(ctx, now, sourceType, releaseSelector).
Select("source.id AS source_id, source.config_version").
Order("runtime.next_check_at ASC").
Order("source.id ASC").
Limit(limit).
Scan(&candidates).Error
if err != nil {
return nil, err
}
return candidates, nil
}
func dueGitHubPagesSourceQuery(
ctx context.Context,
now time.Time,
sourceType string,
releaseSelector string,
) *gorm.DB {
return db.DB(ctx).
Table("of_pages_project_source_runtime AS runtime").
Joins("JOIN of_pages_project_sources AS source ON source.id = runtime.source_id").
Where("source.source_type = ?", sourceType).
Where("source.release_selector = ?", releaseSelector).
Where("runtime.next_check_at IS NOT NULL AND runtime.next_check_at <= ?", now)
}
// GetPagesDeploymentBySourceRevision loads a deployment by project source identity and revision.
func GetPagesDeploymentBySourceRevision(
ctx context.Context,
projectID uint,
sourceIdentity string,
revision string,
) (*model.PagesDeployment, error) {
var deployment model.PagesDeployment
err := db.DB(ctx).
Where("project_id = ? AND source_identity = ? AND source_revision = ?", projectID, sourceIdentity, revision).
First(&deployment).Error
if err != nil {
return nil, err
}
return &deployment, nil
}
@@ -0,0 +1,154 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package repository
import (
"context"
"errors"
"gorm.io/gorm"
"gorm.io/gorm/clause"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/model"
)
// Zone domain binding sentinel errors (shared by proxy_route and zone binding helpers).
var (
// ErrZoneDomainBoundToAnotherRoute is returned when a domain is already bound to a different route.
ErrZoneDomainBoundToAnotherRoute = errors.New("zone domain is already bound to another proxy route")
// ErrZoneDomainNotFound is returned when one or more requested domain IDs do not exist.
ErrZoneDomainNotFound = errors.New("one or more zone domains do not exist")
)
// WithProxyRouteTx runs fn inside a database transaction for proxy-route multi-step work.
func WithProxyRouteTx(ctx context.Context, fn func(tx *gorm.DB) error) error {
return db.DB(ctx).Transaction(fn)
}
// ListProxyRoutes 列出全部代理规则。
func ListProxyRoutes(ctx context.Context) ([]*model.ProxyRoute, error) {
var routes []*model.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) (*model.ProxyRoute, error) {
var route model.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 *model.ProxyRoute) error {
return CreateProxyRouteRecordTx(db.DB(ctx), route)
}
// CreateProxyRouteRecordTx creates a proxy route within an existing transaction.
func CreateProxyRouteRecordTx(tx *gorm.DB, route *model.ProxyRoute) error {
return tx.Create(route).Error
}
// UpdateProxyRouteRecord 更新代理规则。
func UpdateProxyRouteRecord(ctx context.Context, route *model.ProxyRoute) error {
return UpdateProxyRouteRecordTx(db.DB(ctx), route)
}
// UpdateProxyRouteRecordTx updates a proxy route within an existing transaction.
func UpdateProxyRouteRecordTx(tx *gorm.DB, route *model.ProxyRoute) error {
return tx.Model(&model.ProxyRoute{}).Where("id = ?", route.ID).Updates(proxyRouteUpdateMap(route)).Error
}
func proxyRouteUpdateMap(route *model.ProxyRoute) map[string]any {
return 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,
}
}
// DeleteProxyRouteRecord 删除代理规则。
func DeleteProxyRouteRecord(ctx context.Context, id uint) error {
return DeleteProxyRouteRecordTx(db.DB(ctx), id)
}
// DeleteProxyRouteRecordTx deletes a proxy route within an existing transaction.
func DeleteProxyRouteRecordTx(tx *gorm.DB, id uint) error {
return tx.Delete(&model.ProxyRoute{}, id).Error
}
// ClearZoneDomainProxyRouteBindingsTx unbinds every zone domain from a proxy route.
func ClearZoneDomainProxyRouteBindingsTx(tx *gorm.DB, routeID uint) error {
return tx.Model(&model.ZoneDomain{}).Where("proxy_route_id = ?", routeID).Update("proxy_route_id", nil).Error
}
// ReplaceZoneDomainRouteBindingsTx replaces every ZoneDomain binding for a proxy route
// inside the caller's transaction (with row locks on requested domains).
func ReplaceZoneDomainRouteBindingsTx(tx *gorm.DB, routeID uint, domainIDs []uint) error {
var requested []model.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 ErrZoneDomainNotFound
}
for _, domain := range requested {
if domain.ProxyRouteID != nil && *domain.ProxyRouteID != routeID {
return ErrZoneDomainBoundToAnotherRoute
}
}
}
current := tx.Model(&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(&model.ZoneDomain{}).Where("id IN ?", domainIDs).Update("proxy_route_id", routeID).Error
}
// DeleteProxyRouteAndUnbind clears domain bindings then deletes the proxy route in one transaction.
func DeleteProxyRouteAndUnbind(ctx context.Context, id uint) error {
return WithProxyRouteTx(ctx, func(tx *gorm.DB) error {
if err := ClearZoneDomainProxyRouteBindingsTx(tx, id); err != nil {
return err
}
return DeleteProxyRouteRecordTx(tx, id)
})
}
+95
View File
@@ -0,0 +1,95 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package repository
import (
"context"
"errors"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/model"
)
// HasTLSProxyRoutesTable 判断代理规则表是否已迁移。
func HasTLSProxyRoutesTable(ctx context.Context) bool {
return db.DB(ctx).Migrator().HasTable(&model.TLSProxyRouteRef{})
}
// ListTLSCertificates 列出全部证书(不含 PEM 敏感字段的 JSON 暴露由 struct tag 控制)。
func ListTLSCertificates(ctx context.Context) ([]model.TLSCertificate, error) {
conn := db.DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
var certificates []model.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) (*model.TLSCertificate, error) {
conn := db.DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
var certificate model.TLSCertificate
if err := conn.First(&certificate, id).Error; err != nil {
return nil, err
}
return &certificate, nil
}
// CreateTLSCertificateRecord 创建证书记录。
func CreateTLSCertificateRecord(ctx context.Context, certificate *model.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 *model.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(&model.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(&model.TLSCertificate{}).Where("dns_account_id = ?", dnsAccountID).Count(&count).Error; err != nil {
return 0, err
}
return count, nil
}
// ListTLSProxyRouteRefs 列出代理规则证书引用字段。
func ListTLSProxyRouteRefs(ctx context.Context) ([]model.TLSProxyRouteRef, error) {
if !HasTLSProxyRoutesTable(ctx) {
return nil, nil
}
var routes []model.TLSProxyRouteRef
if err := db.DB(ctx).Order("id asc").Find(&routes).Error; err != nil {
return nil, err
}
return routes, nil
}
+357
View File
@@ -0,0 +1,357 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package repository
import (
"context"
"errors"
"time"
"gorm.io/gorm"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/model"
)
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) ([]*model.OpenFlareWAFRuleGroup, error) {
conn, err := wafDB(ctx)
if err != nil {
return nil, err
}
var groups []*model.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) (*model.OpenFlareWAFRuleGroup, error) {
conn, err := wafDB(ctx)
if err != nil {
return nil, err
}
var group model.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) (*model.OpenFlareWAFRuleGroup, error) {
conn, err := wafDB(ctx)
if err != nil {
return nil, err
}
var group model.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 *model.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 *model.OpenFlareWAFRuleGroup) error {
conn, err := wafDB(ctx)
if err != nil {
return err
}
return conn.Model(&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(&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, model.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(&model.OpenFlareWAFRuleGroup{}, id).Error
}
// ListOpenFlareWAFIPGroups returns all IP groups.
func ListOpenFlareWAFIPGroups(ctx context.Context) ([]*model.OpenFlareWAFIPGroup, error) {
conn, err := wafDB(ctx)
if err != nil {
return nil, err
}
var groups []*model.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) ([]*model.OpenFlareWAFIPGroup, error) {
if len(ids) == 0 {
return []*model.OpenFlareWAFIPGroup{}, nil
}
conn, err := wafDB(ctx)
if err != nil {
return nil, err
}
var groups []*model.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) (*model.OpenFlareWAFIPGroup, error) {
conn, err := wafDB(ctx)
if err != nil {
return nil, err
}
var group model.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 *model.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 *model.OpenFlareWAFIPGroup) error {
conn, err := wafDB(ctx)
if err != nil {
return err
}
return conn.Model(&model.OpenFlareWAFIPGroup{}).Where("id = ?", group.ID).Updates(map[string]any{
"name": group.Name,
"type": group.Type,
colEnabled: 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) ([]*model.OpenFlareWAFIPGroup, error) {
conn, err := wafDB(ctx)
if err != nil {
return nil, err
}
var groups []*model.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 *model.OpenFlareWAFIPGroup) error {
conn, err := wafDB(ctx)
if err != nil {
return err
}
return conn.Model(&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(&model.OpenFlareWAFIPGroup{}, id).Error
}
// ListOpenFlareWAFRuleGroupBindings returns all bindings.
func ListOpenFlareWAFRuleGroupBindings(ctx context.Context) ([]model.OpenFlareWAFRuleGroupBinding, error) {
conn, err := wafDB(ctx)
if err != nil {
return nil, err
}
var bindings []model.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) ([]model.OpenFlareWAFRuleGroupBinding, error) {
conn, err := wafDB(ctx)
if err != nil {
return nil, err
}
var bindings []model.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 []model.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(&model.OpenFlareWAFRuleGroupBinding{}).Error; err != nil {
return err
}
bindings := make([]model.OpenFlareWAFRuleGroupBinding, 0, len(routeIDs))
for index, routeID := range routeIDs {
bindings = append(bindings, model.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(&model.OpenFlareWAFRuleGroupBinding{}).Error; err != nil {
return err
}
bindings := make([]model.OpenFlareWAFRuleGroupBinding, 0, len(groupIDs))
for index, groupID := range groupIDs {
bindings = append(bindings, model.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(&model.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(&model.OpenFlareWAFRuleGroupBinding{}).Error; err != nil {
return err
}
return tx.Delete(&model.OpenFlareWAFRuleGroup{}, groupID).Error
})
}
// GetOpenFlareProxyRouteByID returns a proxy route by id when the table exists.
func GetOpenFlareProxyRouteByID(ctx context.Context, id uint) (*model.OriginProxyRoute, error) {
if !HasProxyRoutesTable(ctx) {
return nil, gorm.ErrRecordNotFound
}
conn, err := wafDB(ctx)
if err != nil {
return nil, err
}
var route model.OriginProxyRoute
if err = conn.First(&route, id).Error; err != nil {
return nil, err
}
return &route, nil
}
@@ -0,0 +1,56 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package repository
import (
"context"
"testing"
"github.com/Rain-kl/Wavelet/internal/model"
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(&model.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(&model.OpenFlareWAFRuleGroupBinding{
ID: 50,
RuleGroupID: 1,
ProxyRouteID: 1,
}).Error)
require.NoError(t, ReplaceOpenFlareWAFRuleGroupBindings(ctx, 2, []uint{2, 3}))
var bindings []model.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)
}
@@ -0,0 +1,125 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package repository
import (
"context"
"io/fs"
"os"
"path/filepath"
"runtime"
"testing"
"testing/fstest"
"github.com/Rain-kl/Wavelet/internal/model"
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), "..", "infra", "persistence", "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 []model.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 model.OpenFlareWAFRuleGroup
require.NoError(t, conn.First(&newGroup, 3).Error)
assert.Empty(t, newGroup.Graph)
assert.Equal(t, uint64(1), newGroup.Revision)
var bindings []model.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(&model.OpenFlareWAFRuleGroup{}))
db.SetDB(conn)
t.Cleanup(func() { db.SetDB(nil) })
group := model.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, model.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(&model.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(&model.OpenFlareWAFRuleGroup{}, column) {
t.Fatalf("legacy WAF column %s still exists", column)
}
}
}
+164
View File
@@ -0,0 +1,164 @@
package repository
import (
"context"
"errors"
"fmt"
"gorm.io/gorm"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/model"
)
// ListZones returns all zones ordered by domain ascending.
func ListZones(ctx context.Context) ([]model.Zone, error) {
var zones []model.Zone
if err := db.DB(ctx).Order("domain asc").Find(&zones).Error; err != nil {
return nil, err
}
return zones, nil
}
// GetZoneByID returns a zone by primary key.
func GetZoneByID(ctx context.Context, id uint) (*model.Zone, error) {
var zone model.Zone
if err := db.DB(ctx).First(&zone, id).Error; err != nil {
return nil, err
}
return &zone, nil
}
// CreateZone creates a zone record.
func CreateZone(ctx context.Context, zone *model.Zone) error {
return db.DB(ctx).Create(zone).Error
}
// SaveZone persists zone updates.
func SaveZone(ctx context.Context, zone *model.Zone) error {
return db.DB(ctx).Save(zone).Error
}
// DeleteZone deletes a zone by primary key.
func DeleteZone(ctx context.Context, id uint) error {
return db.DB(ctx).Delete(&model.Zone{}, id).Error
}
// ListZoneDomainCounts returns per-zone domain counts for list cards.
func ListZoneDomainCounts(ctx context.Context) ([]model.ZoneDomainCount, error) {
var rows []model.ZoneDomainCount
if err := db.DB(ctx).Model(&model.ZoneDomain{}).
Select("zone_id, count(*) as count").
Group("zone_id").
Scan(&rows).Error; err != nil {
return nil, err
}
return rows, nil
}
// ListZoneDomainsByZoneID returns domains under a zone ordered by domain ascending.
func ListZoneDomainsByZoneID(ctx context.Context, zoneID uint) ([]model.ZoneDomain, error) {
var domains []model.ZoneDomain
if err := db.DB(ctx).Where("zone_id = ?", zoneID).Order("domain asc").Find(&domains).Error; err != nil {
return nil, err
}
return domains, nil
}
// CountZoneDomainsByZoneID counts domains under a zone.
func CountZoneDomainsByZoneID(ctx context.Context, zoneID uint) (int64, error) {
var count int64
if err := db.DB(ctx).Model(&model.ZoneDomain{}).Where("zone_id = ?", zoneID).Count(&count).Error; err != nil {
return 0, err
}
return count, nil
}
// GetZoneDomainByZoneAndID returns a domain scoped to a zone.
func GetZoneDomainByZoneAndID(ctx context.Context, zoneID, id uint) (*model.ZoneDomain, error) {
var item model.ZoneDomain
if err := db.DB(ctx).Where("id = ? AND zone_id = ?", id, zoneID).First(&item).Error; err != nil {
return nil, err
}
return &item, nil
}
// CreateZoneDomain creates a zone domain record.
func CreateZoneDomain(ctx context.Context, domain *model.ZoneDomain) error {
return db.DB(ctx).Create(domain).Error
}
// SaveZoneDomain persists zone domain updates.
func SaveZoneDomain(ctx context.Context, domain *model.ZoneDomain) error {
return db.DB(ctx).Save(domain).Error
}
// DeleteZoneDomain deletes a zone domain record.
func DeleteZoneDomain(ctx context.Context, domain *model.ZoneDomain) error {
return db.DB(ctx).Delete(domain).Error
}
// ListZoneDomainsByRouteID returns the domains bound to a proxy route.
func ListZoneDomainsByRouteID(ctx context.Context, routeID uint) ([]model.ZoneDomain, error) {
var domains []model.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) ([]model.ZoneDomain, error) {
if len(domainIDs) == 0 {
return []model.ZoneDomain{}, nil
}
var domains []model.ZoneDomain
if err := db.DB(ctx).Where("id IN ?", domainIDs).Find(&domains).Error; err != nil {
return nil, err
}
byID := make(map[uint]model.ZoneDomain, len(domains))
for _, domain := range domains {
byID[domain.ID] = domain
}
ordered := make([]model.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 model.Zone domain.
func CountZoneDomainsByCertificateID(ctx context.Context, certificateID uint) (int64, error) {
var count int64
err := db.DB(ctx).Model(&model.ZoneDomain{}).Where("cert_id = ?", certificateID).Count(&count).Error
return count, err
}
// ReplaceZoneDomainRouteBindings replaces every model.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(errDatabaseNotInitialized)
}
return conn.Transaction(func(tx *gorm.DB) error {
return ReplaceZoneDomainRouteBindingsTx(tx, routeID, domainIDs)
})
}
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
}
@@ -0,0 +1,90 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package repository
import (
"context"
"testing"
"github.com/Rain-kl/Wavelet/internal/model"
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(&model.Zone{}, &model.ZoneDomain{}))
db.SetDB(sqliteDB)
t.Cleanup(func() { db.SetDB(nil) })
return sqliteDB
}
func TestReplaceZoneDomainRouteBindingsRejectsForeignDomain(t *testing.T) {
conn := setupZoneTestDB(t)
ctx := context.Background()
zone := model.Zone{Domain: "example.com"}
require.NoError(t, conn.Create(&zone).Error)
foreignRouteID := uint(11)
domain := model.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 model.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 := model.Zone{Domain: "example.com"}
require.NoError(t, conn.Create(&zone).Error)
routeID := uint(21)
boundDomain := model.ZoneDomain{ZoneID: zone.ID, ProxyRouteID: &routeID, Domain: "old.example.com"}
requestedDomain := model.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 []model.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 := model.Zone{Domain: "example.com"}
require.NoError(t, conn.Create(&zone).Error)
routeID := uint(31)
boundDomain := model.ZoneDomain{ZoneID: zone.ID, ProxyRouteID: &routeID, Domain: "api.example.com"}
unboundDomain := model.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)
}
+16 -1
View File
@@ -5,10 +5,12 @@ package repository
import (
"context"
"time"
"gorm.io/gorm"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/model"
"gorm.io/gorm"
)
// PushHistoryListFilter filters push history pagination queries.
@@ -48,6 +50,19 @@ func CreatePushHistory(ctx context.Context, history *model.PushHistory) error {
return db.DB(ctx).Create(history).Error
}
// CountPushHistoriesCreatedBefore returns how many push history rows were created before cutoff.
func CountPushHistoriesCreatedBefore(ctx context.Context, cutoff time.Time) (int64, error) {
var count int64
err := db.DB(ctx).Model(&model.PushHistory{}).Where("created_at < ?", cutoff).Count(&count).Error
return count, err
}
// DeletePushHistoriesCreatedBefore deletes push history rows created before cutoff.
func DeletePushHistoriesCreatedBefore(ctx context.Context, cutoff time.Time) (int64, error) {
result := db.DB(ctx).Where("created_at < ?", cutoff).Delete(&model.PushHistory{})
return result.RowsAffected, result.Error
}
// PushHistoryQuery returns a scoped query builder for push histories.
func PushHistoryQuery(ctx context.Context) *gorm.DB {
return db.DB(ctx).Model(&model.PushHistory{})
+53
View File
@@ -0,0 +1,53 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package repository
import (
"context"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/model"
)
// CreateSchedule 创建定时任务
func CreateSchedule(ctx context.Context, schedule *model.Schedule) error {
return db.DB(ctx).Create(schedule).Error
}
// UpdateSchedule 更新定时任务
func UpdateSchedule(ctx context.Context, schedule *model.Schedule) error {
return db.DB(ctx).Save(schedule).Error
}
// DeleteSchedule 删除定时任务
func DeleteSchedule(ctx context.Context, id uint64) error {
return db.DB(ctx).Delete(&model.Schedule{}, id).Error
}
// GetScheduleByID 根据 ID 获取定时任务
func GetScheduleByID(ctx context.Context, id uint64) (*model.Schedule, error) {
var schedule model.Schedule
if err := db.DB(ctx).Where("id = ?", id).First(&schedule).Error; err != nil {
return nil, err
}
return &schedule, nil
}
// ListSchedules 获取所有定时任务
func ListSchedules(ctx context.Context) ([]model.Schedule, error) {
var schedules []model.Schedule
if err := db.DB(ctx).Order("id DESC").Find(&schedules).Error; err != nil {
return nil, err
}
return schedules, nil
}
// ListActiveSchedules 获取所有启用的定时任务
func ListActiveSchedules(ctx context.Context) ([]model.Schedule, error) {
var schedules []model.Schedule
if err := db.DB(ctx).Where("is_active = ?", true).Find(&schedules).Error; err != nil {
return nil, err
}
return schedules, nil
}
+1 -8
View File
@@ -18,14 +18,7 @@ import (
"github.com/Rain-kl/Wavelet/pkg/cache/ram"
)
const (
configTypeSystem = "system"
errDatabaseNotInitialized = "database not initialized"
errConfigIntParseFailed = "配置 %s 的值 '%s' 无法转换为整数: %w"
errConfigDecimalParseFailed = "配置 %s 的值 '%s' 无法转换为decimal: %w"
errConfigBoolParseFailed = "配置 %s 的值 '%s' 无法转换为布尔值: %w"
errParseMenuDisplayConfigFailed = "解析目录显示配置失败: %w"
)
const configTypeSystem = "system"
// PreheatSystemConfigs loads all system configs from database.
// This function strictly performs database read and does not perform any cache read or write operations.
+6 -1
View File
@@ -54,7 +54,12 @@ func CreateSystemConfig(ctx context.Context, config *model.SystemConfig) error {
// UpdateSystemConfigFields applies partial updates to a system config row.
func UpdateSystemConfigFields(ctx context.Context, config *model.SystemConfig, updates map[string]any) error {
return db.DB(ctx).Model(config).Updates(updates).Error
return UpdateSystemConfigFieldsTx(db.DB(ctx), config, updates)
}
// UpdateSystemConfigFieldsTx applies partial updates within an existing transaction.
func UpdateSystemConfigFieldsTx(tx *gorm.DB, config *model.SystemConfig, updates map[string]any) error {
return tx.Model(config).Updates(updates).Error
}
// SaveOrUpdateSystemConfig creates or updates a config row and invalidates cache.
+281
View File
@@ -0,0 +1,281 @@
package repository
import (
"context"
"errors"
"fmt"
"strings"
"time"
"github.com/redis/go-redis/v9"
"gorm.io/gorm"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/infra/persistence/idgen"
"github.com/Rain-kl/Wavelet/internal/model"
)
const (
taskExecutionLogRedisKeyPrefix = "task:execution:log:"
taskExecutionLogExpiration = 24 * time.Hour
taskExecutionLogMaxLines = 1000
)
// CreateTaskExecution 创建任务执行记录
func CreateTaskExecution(ctx context.Context, execution *model.TaskExecution) error {
execution.ID = idgen.NextUint64ID()
return db.DB(ctx).Create(execution).Error
}
// UpdateTaskExecution 更新任务执行记录,忽略由 Redis 缓冲和归档流程管理的 log 字段。
func UpdateTaskExecution(ctx context.Context, execution *model.TaskExecution) error {
return db.DB(ctx).Omit("log").Save(execution).Error
}
// GetTaskExecutionByTaskID 根据 TaskID 获取执行记录
func GetTaskExecutionByTaskID(ctx context.Context, taskID string) (*model.TaskExecution, error) {
var execution model.TaskExecution
if err := db.DB(ctx).Where("task_id = ?", taskID).First(&execution).Error; err != nil {
return nil, err
}
if err := loadTaskExecutionLog(ctx, &execution); err != nil {
return nil, err
}
return &execution, nil
}
// GetTaskExecutionByID 根据 ID 获取执行记录
func GetTaskExecutionByID(ctx context.Context, id uint64) (*model.TaskExecution, error) {
var execution model.TaskExecution
if err := db.DB(ctx).Where("id = ?", id).First(&execution).Error; err != nil {
return nil, err
}
if err := loadTaskExecutionLog(ctx, &execution); err != nil {
return nil, err
}
return &execution, nil
}
// GetLatestTaskExecutionByTaskType returns the most recent execution for a task type.
// ok is false when no row exists.
func GetLatestTaskExecutionByTaskType(ctx context.Context, taskType string) (*model.TaskExecution, bool, error) {
var execution model.TaskExecution
err := db.DB(ctx).
Where("task_type = ?", taskType).
Order("id DESC").
First(&execution).Error
if err == nil {
if loadErr := loadTaskExecutionLog(ctx, &execution); loadErr != nil {
return nil, false, loadErr
}
return &execution, true, nil
}
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, false, nil
}
return nil, false, err
}
// AppendTaskExecutionLog 将日志追加到 Redis 缓冲,任务完成后再持久化到数据库。
func AppendTaskExecutionLog(ctx context.Context, taskID string, logLine string) error {
if db.Redis == nil {
return errors.New("redis client is not initialized")
}
now := time.Now().Format("15:04:05")
line := fmt.Sprintf("[%s] %s\n", now, logLine)
key := taskExecutionLogRedisKey(taskID)
_, err := db.Redis.TxPipelined(ctx, func(pipe redis.Pipeliner) error {
pipe.RPush(ctx, key, line)
pipe.LTrim(ctx, key, -taskExecutionLogMaxLines, -1)
pipe.Expire(ctx, key, taskExecutionLogExpiration)
return nil
})
if err != nil {
return fmt.Errorf("append task execution log to redis: %w", err)
}
return nil
}
// FlushTaskExecutionLog 将 Redis 中的完整任务日志写入数据库,并在成功后清理缓存。
func FlushTaskExecutionLog(ctx context.Context, taskID string) error {
if db.Redis == nil {
return errors.New("redis client is not initialized")
}
key := taskExecutionLogRedisKey(taskID)
logLines, err := db.Redis.LRange(ctx, key, 0, -1).Result()
if err != nil {
return fmt.Errorf("get task execution log from redis: %w", err)
}
if len(logLines) == 0 {
return nil
}
logText := strings.Join(logLines, "")
result := db.DB(ctx).Model(&model.TaskExecution{}).
Where("task_id = ?", taskID).
Update("log", logText)
if result.Error != nil {
return fmt.Errorf("persist task execution log: %w", result.Error)
}
if result.RowsAffected == 0 {
return fmt.Errorf("persist task execution log: task %q not found", taskID)
}
if err := db.Redis.Del(ctx, key).Err(); err != nil {
return fmt.Errorf("delete persisted task execution log from redis: %w", err)
}
return nil
}
// ListTaskExecutions 分页查询任务执行记录
func ListTaskExecutions(ctx context.Context, req model.ListTaskExecutionsRequest) ([]model.TaskExecution, int64, error) {
if req.Page <= 0 {
req.Page = 1
}
if req.PageSize <= 0 {
req.PageSize = 20
}
query := db.DB(ctx).Model(&model.TaskExecution{})
if req.Status != "" {
query = query.Where("status = ?", req.Status)
}
if req.TaskType != "" {
query = query.Where("task_type = ?", req.TaskType)
}
var total int64
if err := query.Count(&total).Error; err != nil {
return nil, 0, err
}
var executions []model.TaskExecution
offset := (req.Page - 1) * req.PageSize
if err := query.Order("id DESC").Offset(offset).Limit(req.PageSize).Find(&executions).Error; err != nil {
return nil, 0, err
}
if err := loadTaskExecutionLogs(ctx, executions); err != nil {
return nil, 0, err
}
return executions, total, nil
}
// MarkFailedTaskExecutionsSucceededTx marks failed executions of a task type as succeeded within a transaction.
func MarkFailedTaskExecutionsSucceededTx(
tx *gorm.DB,
taskType string,
result string,
finishedAt time.Time,
) error {
return tx.Model(&model.TaskExecution{}).
Where("task_type = ? AND status = ?", taskType, model.TaskExecutionStatusFailed).
Updates(map[string]any{
"status": model.TaskExecutionStatusSucceeded,
"result": result,
"finished_at": finishedAt,
}).Error
}
// CleanupTaskExecutionLogs removes finished task execution logs according to frequency-based retention.
func CleanupTaskExecutionLogs(ctx context.Context, now time.Time) (model.TaskExecutionCleanupStats, error) {
const (
frequencyWindowDays = 30
highFrequencyThreshold = frequencyWindowDays
)
frequencyWindowStart := now.AddDate(0, 0, -frequencyWindowDays)
highFrequencyCutoff := now.AddDate(0, 0, -3)
lowFrequencyCutoff := now.AddDate(0, 0, -30)
terminalStatuses := []model.TaskExecutionStatus{model.TaskExecutionStatusSucceeded, model.TaskExecutionStatusFailed}
var highFrequencyTaskTypes []string
if err := db.DB(ctx).
Model(&model.TaskExecution{}).
Select("task_type").
Where("created_at >= ?", frequencyWindowStart).
Group("task_type").
Having("COUNT(*) > ?", highFrequencyThreshold).
Pluck("task_type", &highFrequencyTaskTypes).Error; err != nil {
return model.TaskExecutionCleanupStats{}, fmt.Errorf("query high-frequency task types: %w", err)
}
var highFrequencyDeleted int64
if len(highFrequencyTaskTypes) > 0 {
highFrequencyResult := db.DB(ctx).
Where("status IN ?", terminalStatuses).
Where("created_at < ?", highFrequencyCutoff).
Where("task_type IN ?", highFrequencyTaskTypes).
Delete(&model.TaskExecution{})
if highFrequencyResult.Error != nil {
return model.TaskExecutionCleanupStats{}, fmt.Errorf("delete high-frequency task execution logs: %w", highFrequencyResult.Error)
}
highFrequencyDeleted = highFrequencyResult.RowsAffected
}
lowFrequencyQuery := db.DB(ctx).
Where("status IN ?", terminalStatuses).
Where("created_at < ?", lowFrequencyCutoff)
if len(highFrequencyTaskTypes) > 0 {
lowFrequencyQuery = lowFrequencyQuery.Where("task_type NOT IN ?", highFrequencyTaskTypes)
}
lowFrequencyResult := lowFrequencyQuery.Delete(&model.TaskExecution{})
if lowFrequencyResult.Error != nil {
return model.TaskExecutionCleanupStats{}, fmt.Errorf("delete low-frequency task execution logs: %w", lowFrequencyResult.Error)
}
return model.TaskExecutionCleanupStats{
HighFrequencyDeleted: highFrequencyDeleted,
LowFrequencyDeleted: lowFrequencyResult.RowsAffected,
}, nil
}
func taskExecutionLogRedisKey(taskID string) string {
return db.PrefixedKey(taskExecutionLogRedisKeyPrefix + taskID)
}
func loadTaskExecutionLog(ctx context.Context, execution *model.TaskExecution) error {
if db.Redis == nil {
return nil
}
logLines, err := db.Redis.LRange(ctx, taskExecutionLogRedisKey(execution.TaskID), 0, -1).Result()
if err != nil {
return fmt.Errorf("get task execution log from redis: %w", err)
}
if len(logLines) == 0 {
return nil
}
execution.Log = strings.Join(logLines, "")
return nil
}
func loadTaskExecutionLogs(ctx context.Context, executions []model.TaskExecution) error {
if db.Redis == nil || len(executions) == 0 {
return nil
}
commands := make([]*redis.StringSliceCmd, len(executions))
_, err := db.Redis.Pipelined(ctx, func(pipe redis.Pipeliner) error {
for i := range executions {
commands[i] = pipe.LRange(ctx, taskExecutionLogRedisKey(executions[i].TaskID), 0, -1)
}
return nil
})
if err != nil {
return fmt.Errorf("get task execution logs from redis: %w", err)
}
for i := range executions {
logLines := commands[i].Val()
if len(logLines) > 0 {
executions[i].Log = strings.Join(logLines, "")
}
}
return nil
}
+486
View File
@@ -0,0 +1,486 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package repository
import (
"context"
"fmt"
"testing"
"time"
"github.com/Rain-kl/Wavelet/internal/model"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/alicebob/miniredis/v2"
"github.com/glebarez/sqlite"
"github.com/redis/go-redis/v9"
"github.com/redis/go-redis/v9/maintnotifications"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"gorm.io/gorm"
)
func setupTaskExecutionTestEnvironment(t *testing.T) func() {
sqliteDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{
DisableForeignKeyConstraintWhenMigrating: true,
})
require.NoError(t, err)
err = sqliteDB.AutoMigrate(&model.TaskExecution{})
require.NoError(t, err)
miniRedis, err := miniredis.Run()
require.NoError(t, err)
redisClient := redis.NewClient(&redis.Options{
Addr: miniRedis.Addr(),
MaintNotificationsConfig: &maintnotifications.Config{
Mode: maintnotifications.ModeDisabled,
},
})
db.SetDB(sqliteDB)
db.Redis = redisClient
return func() {
require.NoError(t, redisClient.Close())
miniRedis.Close()
db.SetDB(nil)
db.Redis = nil
}
}
func TestCreateTaskExecution(t *testing.T) {
cleanup := setupTaskExecutionTestEnvironment(t)
defer cleanup()
ctx := context.Background()
execution := &model.TaskExecution{
TaskID: "manual_cleanup_123",
TaskType: "system:cleanup",
TaskName: "清理未使用上传",
Status: model.TaskExecutionStatusPending,
Retryable: true,
MaxRetry: 3,
RetryCount: 0,
Payload: `{"test": true}`,
TriggeredBy: "manual",
}
err := CreateTaskExecution(ctx, execution)
require.NoError(t, err)
assert.NotZero(t, execution.ID, "ID should be generated")
assert.NotZero(t, execution.CreatedAt, "CreatedAt should be set")
assert.NotZero(t, execution.UpdatedAt, "UpdatedAt should be set")
}
func TestGetTaskExecutionByTaskID(t *testing.T) {
cleanup := setupTaskExecutionTestEnvironment(t)
defer cleanup()
ctx := context.Background()
// 创建记录
execution := &model.TaskExecution{
TaskID: "test_task_id_001",
TaskType: "system:cleanup",
TaskName: "清理未使用上传",
Status: model.TaskExecutionStatusPending,
Retryable: true,
MaxRetry: 3,
TriggeredBy: "manual",
}
err := CreateTaskExecution(ctx, execution)
require.NoError(t, err)
// 按 TaskID 查询
found, err := GetTaskExecutionByTaskID(ctx, "test_task_id_001")
require.NoError(t, err)
assert.Equal(t, execution.ID, found.ID)
assert.Equal(t, "test_task_id_001", found.TaskID)
assert.Equal(t, model.TaskExecutionStatusPending, found.Status)
assert.True(t, found.Retryable)
assert.Equal(t, 3, found.MaxRetry)
// 查询不存在的 TaskID
_, err = GetTaskExecutionByTaskID(ctx, "nonexistent")
assert.Error(t, err, "should return error for non-existent taskID")
}
func TestGetTaskExecutionByID(t *testing.T) {
cleanup := setupTaskExecutionTestEnvironment(t)
defer cleanup()
ctx := context.Background()
execution := &model.TaskExecution{
TaskID: "test_by_id_001",
TaskType: "system:cleanup",
TaskName: "清理未使用上传",
Status: model.TaskExecutionStatusPending,
TriggeredBy: "system",
}
err := CreateTaskExecution(ctx, execution)
require.NoError(t, err)
// 按主键查询
found, err := GetTaskExecutionByID(ctx, execution.ID)
require.NoError(t, err)
assert.Equal(t, execution.TaskID, found.TaskID)
}
func TestUpdateTaskExecution(t *testing.T) {
cleanup := setupTaskExecutionTestEnvironment(t)
defer cleanup()
ctx := context.Background()
// 创建记录
execution := &model.TaskExecution{
TaskID: "test_update_001",
TaskType: "system:cleanup",
TaskName: "清理未使用上传",
Status: model.TaskExecutionStatusPending,
TriggeredBy: "manual",
}
err := CreateTaskExecution(ctx, execution)
require.NoError(t, err)
// 更新状态为 running
now := time.Now()
execution.Status = model.TaskExecutionStatusRunning
execution.StartedAt = &now
err = UpdateTaskExecution(ctx, execution)
require.NoError(t, err)
// 验证更新
found, err := GetTaskExecutionByTaskID(ctx, "test_update_001")
require.NoError(t, err)
assert.Equal(t, model.TaskExecutionStatusRunning, found.Status)
assert.NotNil(t, found.StartedAt)
// 更新为 succeeded
finishTime := time.Now()
execution.Status = model.TaskExecutionStatusSucceeded
execution.FinishedAt = &finishTime
execution.Duration = 1500
execution.Result = "共清理 50 个文件"
err = UpdateTaskExecution(ctx, execution)
require.NoError(t, err)
found, err = GetTaskExecutionByTaskID(ctx, "test_update_001")
require.NoError(t, err)
assert.Equal(t, model.TaskExecutionStatusSucceeded, found.Status)
assert.Equal(t, int64(1500), found.Duration)
assert.Equal(t, "共清理 50 个文件", found.Result)
}
func TestUpdateTaskExecutionFailed(t *testing.T) {
cleanup := setupTaskExecutionTestEnvironment(t)
defer cleanup()
ctx := context.Background()
execution := &model.TaskExecution{
TaskID: "test_fail_001",
TaskType: "system:cleanup",
TaskName: "清理未使用上传",
Status: model.TaskExecutionStatusPending,
Retryable: true,
MaxRetry: 3,
TriggeredBy: "manual",
}
err := CreateTaskExecution(ctx, execution)
require.NoError(t, err)
// 标记为失败
now := time.Now()
execution.Status = model.TaskExecutionStatusFailed
execution.StartedAt = &now
execution.FinishedAt = &now
execution.Duration = 200
execution.ErrorMessage = "S3 连接超时"
err = UpdateTaskExecution(ctx, execution)
require.NoError(t, err)
found, err := GetTaskExecutionByTaskID(ctx, "test_fail_001")
require.NoError(t, err)
assert.Equal(t, model.TaskExecutionStatusFailed, found.Status)
assert.Equal(t, "S3 连接超时", found.ErrorMessage)
assert.Equal(t, int64(200), found.Duration)
}
func TestUpdateTaskExecutionDoesNotPersistBufferedLog(t *testing.T) {
cleanup := setupTaskExecutionTestEnvironment(t)
defer cleanup()
ctx := context.Background()
execution := &model.TaskExecution{
TaskID: "test_omit_log_001",
TaskType: "system:cleanup",
TaskName: "清理未使用上传",
Status: model.TaskExecutionStatusPending,
TriggeredBy: "manual",
}
err := CreateTaskExecution(ctx, execution)
require.NoError(t, err)
// 运行中的日志仅缓存在 Redis。
err = AppendTaskExecutionLog(ctx, "test_omit_log_001", "第一条执行日志")
require.NoError(t, err)
assert.Empty(t, execution.Log)
execution.Status = model.TaskExecutionStatusSucceeded
execution.Duration = 100
err = UpdateTaskExecution(ctx, execution)
require.NoError(t, err)
var persisted model.TaskExecution
err = db.DB(ctx).Where("task_id = ?", "test_omit_log_001").First(&persisted).Error
require.NoError(t, err)
assert.Equal(t, model.TaskExecutionStatusSucceeded, persisted.Status)
assert.Empty(t, persisted.Log)
found, err := GetTaskExecutionByTaskID(ctx, "test_omit_log_001")
require.NoError(t, err)
assert.Contains(t, found.Log, "第一条执行日志")
}
func TestAppendTaskExecutionLog(t *testing.T) {
cleanup := setupTaskExecutionTestEnvironment(t)
defer cleanup()
ctx := context.Background()
execution := &model.TaskExecution{
TaskID: "test_log_001",
TaskType: "system:cleanup",
TaskName: "清理未使用上传",
Status: model.TaskExecutionStatusPending,
TriggeredBy: "manual",
}
err := CreateTaskExecution(ctx, execution)
require.NoError(t, err)
// 追加多条日志
err = AppendTaskExecutionLog(ctx, "test_log_001", "开始扫描未使用上传文件")
require.NoError(t, err)
err = AppendTaskExecutionLog(ctx, "test_log_001", "本批次找到 42 个待清理文件")
require.NoError(t, err)
err = AppendTaskExecutionLog(ctx, "test_log_001", "清理完成,共删除 42 个文件")
require.NoError(t, err)
// 读取时优先返回 Redis 中的在途日志。
found, err := GetTaskExecutionByTaskID(ctx, "test_log_001")
require.NoError(t, err)
assert.Contains(t, found.Log, "开始扫描未使用上传文件")
assert.Contains(t, found.Log, "本批次找到 42 个待清理文件")
assert.Contains(t, found.Log, "清理完成,共删除 42 个文件")
var persisted model.TaskExecution
err = db.DB(ctx).Where("task_id = ?", "test_log_001").First(&persisted).Error
require.NoError(t, err)
assert.Empty(t, persisted.Log)
err = FlushTaskExecutionLog(ctx, "test_log_001")
require.NoError(t, err)
err = db.DB(ctx).Where("task_id = ?", "test_log_001").First(&persisted).Error
require.NoError(t, err)
assert.Contains(t, persisted.Log, "开始扫描未使用上传文件")
exists, err := db.Redis.Exists(ctx, taskExecutionLogRedisKey("test_log_001")).Result()
require.NoError(t, err)
assert.Zero(t, exists)
}
func TestAppendTaskExecutionLogLimitsLinesAndRefreshesTTL(t *testing.T) {
cleanup := setupTaskExecutionTestEnvironment(t)
defer cleanup()
ctx := context.Background()
const taskID = "limited_log_001"
for i := 0; i < taskExecutionLogMaxLines+5; i++ {
err := AppendTaskExecutionLog(ctx, taskID, fmt.Sprintf("日志-%04d", i))
require.NoError(t, err)
}
key := taskExecutionLogRedisKey(taskID)
logLines, err := db.Redis.LRange(ctx, key, 0, -1).Result()
require.NoError(t, err)
assert.Len(t, logLines, taskExecutionLogMaxLines)
assert.Contains(t, logLines[0], "日志-0005")
assert.Contains(t, logLines[len(logLines)-1], "日志-1004")
ttl, err := db.Redis.TTL(ctx, key).Result()
require.NoError(t, err)
assert.Equal(t, taskExecutionLogExpiration, ttl)
}
func TestAppendTaskExecutionLogNonExistent(t *testing.T) {
cleanup := setupTaskExecutionTestEnvironment(t)
defer cleanup()
ctx := context.Background()
// Redis 缓冲不依赖数据库记录是否已经创建。
err := AppendTaskExecutionLog(ctx, "nonexistent_task", "测试日志")
assert.NoError(t, err)
err = FlushTaskExecutionLog(ctx, "nonexistent_task")
assert.Error(t, err)
}
func TestGetTaskExecutionLogPrefersRedis(t *testing.T) {
cleanup := setupTaskExecutionTestEnvironment(t)
defer cleanup()
ctx := context.Background()
execution := &model.TaskExecution{
TaskID: "redis_priority_001",
TaskType: "system:cleanup",
TaskName: "清理未使用上传",
Status: model.TaskExecutionStatusRunning,
Log: "数据库旧日志",
TriggeredBy: "manual",
}
err := CreateTaskExecution(ctx, execution)
require.NoError(t, err)
err = AppendTaskExecutionLog(ctx, execution.TaskID, "Redis 最新日志")
require.NoError(t, err)
found, err := GetTaskExecutionByID(ctx, execution.ID)
require.NoError(t, err)
assert.Contains(t, found.Log, "Redis 最新日志")
assert.NotContains(t, found.Log, "数据库旧日志")
}
func TestListTaskExecutions(t *testing.T) {
cleanup := setupTaskExecutionTestEnvironment(t)
defer cleanup()
ctx := context.Background()
// 创建多条记录,包含不同状态和类型
records := []*model.TaskExecution{
{TaskID: "list_001", TaskType: "system:cleanup", TaskName: "系统垃圾清理", Status: model.TaskExecutionStatusSucceeded, TriggeredBy: "manual"},
{TaskID: "list_002", TaskType: "system:cleanup", TaskName: "系统垃圾清理", Status: model.TaskExecutionStatusFailed, TriggeredBy: "system"},
{TaskID: "list_003", TaskType: "other:task", TaskName: "其他任务", Status: model.TaskExecutionStatusPending, TriggeredBy: "manual"},
{TaskID: "list_004", TaskType: "system:cleanup", TaskName: "系统垃圾清理", Status: model.TaskExecutionStatusRunning, TriggeredBy: "manual"},
{TaskID: "list_005", TaskType: "other:task", TaskName: "其他任务", Status: model.TaskExecutionStatusSucceeded, TriggeredBy: "system"},
}
for _, r := range records {
err := CreateTaskExecution(ctx, r)
require.NoError(t, err)
}
err := AppendTaskExecutionLog(ctx, "list_004", "运行中的 Redis 日志")
require.NoError(t, err)
// 查询全部(分页)
items, total, err := ListTaskExecutions(ctx, model.ListTaskExecutionsRequest{Page: 1, PageSize: 10})
require.NoError(t, err)
assert.Equal(t, int64(5), total)
assert.Len(t, items, 5)
for _, item := range items {
if item.TaskID == "list_004" {
assert.Contains(t, item.Log, "运行中的 Redis 日志")
}
}
// 按状态筛选:failed
items, total, err = ListTaskExecutions(ctx, model.ListTaskExecutionsRequest{Status: "failed", Page: 1, PageSize: 10})
require.NoError(t, err)
assert.Equal(t, int64(1), total)
assert.Len(t, items, 1)
assert.Equal(t, "list_002", items[0].TaskID)
// 按类型筛选
_, total, err = ListTaskExecutions(ctx, model.ListTaskExecutionsRequest{TaskType: "other:task", Page: 1, PageSize: 10})
require.NoError(t, err)
assert.Equal(t, int64(2), total)
// 分页测试
items, total, err = ListTaskExecutions(ctx, model.ListTaskExecutionsRequest{Page: 1, PageSize: 2})
require.NoError(t, err)
assert.Equal(t, int64(5), total)
assert.Len(t, items, 2)
items2, total2, err := ListTaskExecutions(ctx, model.ListTaskExecutionsRequest{Page: 2, PageSize: 2})
require.NoError(t, err)
assert.Equal(t, int64(5), total2)
assert.Len(t, items2, 2)
// 确保分页数据不重复
assert.NotEqual(t, items[0].ID, items2[0].ID)
// 状态 + 类型组合筛选
items, total, err = ListTaskExecutions(ctx, model.ListTaskExecutionsRequest{Status: "succeeded", TaskType: "system:cleanup", Page: 1, PageSize: 10})
require.NoError(t, err)
assert.Equal(t, int64(1), total)
assert.Equal(t, "list_001", items[0].TaskID)
}
func TestListTaskExecutionsDefaultPaging(t *testing.T) {
cleanup := setupTaskExecutionTestEnvironment(t)
defer cleanup()
ctx := context.Background()
// 不传分页参数,应使用默认值 page=1, pageSize=20
items, total, err := ListTaskExecutions(ctx, model.ListTaskExecutionsRequest{})
require.NoError(t, err)
assert.Equal(t, int64(0), total)
assert.Len(t, items, 0)
}
func TestCleanupTaskExecutionLogs(t *testing.T) {
cleanup := setupTaskExecutionTestEnvironment(t)
defer cleanup()
ctx := context.Background()
now := time.Date(2026, 6, 17, 12, 0, 0, 0, time.UTC)
for i := 0; i < 31; i++ {
createTaskExecutionForCleanup(t, ctx, fmt.Sprintf("high_recent_%02d", i), "high:task", model.TaskExecutionStatusSucceeded, now.Add(-2*time.Hour))
}
createTaskExecutionForCleanup(t, ctx, "high_old_4d", "high:task", model.TaskExecutionStatusSucceeded, now.AddDate(0, 0, -4))
createTaskExecutionForCleanup(t, ctx, "high_old_40d", "high:task", model.TaskExecutionStatusFailed, now.AddDate(0, 0, -40))
createTaskExecutionForCleanup(t, ctx, "high_running_old", "high:task", model.TaskExecutionStatusRunning, now.AddDate(0, 0, -10))
createTaskExecutionForCleanup(t, ctx, "low_old_31d", "low:task", model.TaskExecutionStatusSucceeded, now.AddDate(0, 0, -31))
createTaskExecutionForCleanup(t, ctx, "low_recent_29d", "low:task", model.TaskExecutionStatusSucceeded, now.AddDate(0, 0, -29))
createTaskExecutionForCleanup(t, ctx, "low_pending_old", "low:task", model.TaskExecutionStatusPending, now.AddDate(0, 0, -45))
stats, err := CleanupTaskExecutionLogs(ctx, now)
require.NoError(t, err)
assert.Equal(t, int64(2), stats.HighFrequencyDeleted)
assert.Equal(t, int64(1), stats.LowFrequencyDeleted)
for _, taskID := range []string{"high_old_4d", "high_old_40d", "low_old_31d"} {
var count int64
err := db.DB(ctx).Model(&model.TaskExecution{}).Where("task_id = ?", taskID).Count(&count).Error
require.NoError(t, err)
assert.Equal(t, int64(0), count, "CleanupTaskExecutionLogs(%s) should delete expired log", taskID)
}
for _, taskID := range []string{"high_recent_00", "high_running_old", "low_recent_29d", "low_pending_old"} {
var count int64
err := db.DB(ctx).Model(&model.TaskExecution{}).Where("task_id = ?", taskID).Count(&count).Error
require.NoError(t, err)
assert.Equal(t, int64(1), count, "CleanupTaskExecutionLogs(%s) should keep retained log", taskID)
}
}
func TestTaskExecutionTableName(t *testing.T) {
execution := model.TaskExecution{}
assert.Equal(t, "w_task_executions", execution.TableName())
}
func createTaskExecutionForCleanup(t *testing.T, ctx context.Context, taskID string, taskType string, status model.TaskExecutionStatus, createdAt time.Time) {
t.Helper()
execution := &model.TaskExecution{
TaskID: taskID,
TaskType: taskType,
TaskName: taskType,
Status: status,
CreatedAt: createdAt,
UpdatedAt: createdAt,
TriggeredBy: "system",
}
err := CreateTaskExecution(ctx, execution)
require.NoError(t, err)
}
+144 -1
View File
@@ -6,10 +6,13 @@ package repository
import (
"context"
"strings"
"time"
"gorm.io/gorm"
"gorm.io/gorm/clause"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/model"
"gorm.io/gorm"
)
// UploadListFilter filters paginated upload queries.
@@ -22,6 +25,22 @@ type UploadListFilter struct {
PageSize int
}
// UploadStorageObject is a distinct active object path with aggregated metadata for migration.
type UploadStorageObject struct {
FilePath string `gorm:"column:file_path"`
FileSize int64 `gorm:"column:file_size"`
MimeType string `gorm:"column:mime_type"`
Hash string `gorm:"column:hash"`
}
// RunInTransaction executes fn inside a database transaction.
// Prefer domain-specific repository methods when the full operation can live in repository.
// Upload package multi-step flows (lock + soft-delete + stats) use this boundary so apps
// do not call db.DB directly.
func RunInTransaction(ctx context.Context, fn func(tx *gorm.DB) error) error {
return db.DB(ctx).Transaction(fn)
}
// ListUploads returns paginated upload records matching the filter.
func ListUploads(ctx context.Context, filter UploadListFilter) (int64, []model.Upload, error) {
query := db.DB(ctx).Model(&model.Upload{}).
@@ -62,6 +81,28 @@ func GetActiveUploadByID(ctx context.Context, id uint64) (model.Upload, error) {
return upload, nil
}
// GetCacheableUploadByID loads a pending or used upload by ID (for metadata cache DB fallback).
func GetCacheableUploadByID(ctx context.Context, id uint64) (model.Upload, error) {
var upload model.Upload
if err := db.DB(ctx).
Where("id = ? AND status IN (?, ?)", id, model.UploadStatusPending, model.UploadStatusUsed).
First(&upload).Error; err != nil {
return model.Upload{}, err
}
return upload, nil
}
// GetUploadByIDForUpdateTx loads and row-locks an upload by ID within an existing transaction.
func GetUploadByIDForUpdateTx(tx *gorm.DB, id uint64) (model.Upload, error) {
var upload model.Upload
if err := tx.Clauses(clause.Locking{Strength: "UPDATE"}).
Where("id = ?", id).
First(&upload).Error; err != nil {
return model.Upload{}, err
}
return upload, nil
}
// SoftDeleteUpload marks an active upload as deleted and reports whether the row transitioned.
// External modules must use upload.Remove or upload.RemoveOwned; only internal/apps/upload may call this.
func SoftDeleteUpload(ctx context.Context, upload *model.Upload) (int64, error) {
@@ -131,6 +172,108 @@ func ListUploadsByIDs(ctx context.Context, ids []uint64) ([]model.Upload, error)
return uploads, nil
}
// CountActiveUploads returns the number of non-deleted upload records.
func CountActiveUploads(ctx context.Context) (int64, error) {
var count int64
err := db.DB(ctx).Model(&model.Upload{}).
Where("status != ?", model.UploadStatusDeleted).
Count(&count).Error
return count, err
}
// ListPendingUploadsOlderThan returns pending uploads created before olderThan, after lastID, ordered by id.
func ListPendingUploadsOlderThan(ctx context.Context, lastID uint64, olderThan time.Time, limit int) ([]model.Upload, error) {
var uploads []model.Upload
err := db.DB(ctx).
Where("id > ? AND status = ? AND created_at < ?", lastID, model.UploadStatusPending, olderThan).
Order("id ASC").
Limit(limit).
Find(&uploads).Error
return uploads, err
}
// ListActiveImageUploadsAfterID returns non-deleted image uploads with id greater than lastID.
func ListActiveImageUploadsAfterID(ctx context.Context, lastID uint64, limit int) ([]model.Upload, error) {
var uploads []model.Upload
err := db.DB(ctx).
Where("id > ? AND status != ? AND (LOWER(mime_type) LIKE ? OR LOWER(extension) IN ?)",
lastID,
model.UploadStatusDeleted,
"image/%",
[]string{"jpg", "jpeg", "png", "webp", "gif"},
).
Order("id ASC").
Limit(limit).
Find(&uploads).Error
return uploads, err
}
// CountDistinctActiveFilePaths returns the number of distinct non-deleted upload file paths.
func CountDistinctActiveFilePaths(ctx context.Context) (int64, error) {
var count int64
err := db.DB(ctx).Model(&model.Upload{}).
Where("status != ?", model.UploadStatusDeleted).
Distinct("file_path").
Count(&count).Error
return count, err
}
// ListDistinctActiveStorageObjects returns a page of distinct active file paths ordered by path.
// When afterFilePath is non-empty, only paths strictly greater than it are returned.
func ListDistinctActiveStorageObjects(ctx context.Context, afterFilePath string, limit int) ([]UploadStorageObject, error) {
var objects []UploadStorageObject
query := db.DB(ctx).Model(&model.Upload{}).
Select("file_path, MAX(file_size) AS file_size, MAX(mime_type) AS mime_type, MAX(hash) AS hash").
Where("status != ?", model.UploadStatusDeleted)
if afterFilePath != "" {
query = query.Where("file_path > ?", afterFilePath)
}
err := query.Group("file_path").
Order("file_path ASC").
Limit(limit).
Scan(&objects).Error
return objects, err
}
// UpdateActiveUploadsFilePath rewrites file_path for all non-deleted uploads matching oldPath.
func UpdateActiveUploadsFilePath(ctx context.Context, oldPath, newPath string) error {
return db.DB(ctx).Model(&model.Upload{}).
Where("file_path = ? AND status != ?", oldPath, model.UploadStatusDeleted).
Update("file_path", newPath).Error
}
// MarkActiveUploadsDeletedByFilePath marks all non-deleted uploads with the given path as deleted
// and returns the rows that transitioned (for stats adjustment).
func MarkActiveUploadsDeletedByFilePath(ctx context.Context, filePath string) ([]model.Upload, error) {
var affected []model.Upload
err := db.DB(ctx).Transaction(func(tx *gorm.DB) error {
if err := tx.
Where("file_path = ? AND status != ?", filePath, model.UploadStatusDeleted).
Find(&affected).Error; err != nil {
return err
}
if len(affected) == 0 {
return nil
}
return tx.Model(&model.Upload{}).
Where("file_path = ?", filePath).
Update("status", model.UploadStatusDeleted).Error
})
if err != nil {
return nil, err
}
return affected, nil
}
// ListActiveUploadsTx returns all non-deleted uploads within an existing transaction.
func ListActiveUploadsTx(tx *gorm.DB) ([]model.Upload, error) {
var uploads []model.Upload
if err := tx.Where("status != ?", model.UploadStatusDeleted).Find(&uploads).Error; err != nil {
return nil, err
}
return uploads, nil
}
// UploadQuery returns a scoped GORM query for uploads.
func UploadQuery(ctx context.Context) *gorm.DB {
return db.DB(ctx).Model(&model.Upload{})
+77
View File
@@ -5,6 +5,10 @@ package repository
import (
"context"
"time"
"gorm.io/gorm"
"gorm.io/gorm/clause"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/model"
@@ -18,3 +22,76 @@ func ListUploadStats(ctx context.Context) ([]model.UploadStat, error) {
}
return stats, nil
}
// GetTotalUploadStat returns the aggregate total-dimension stats row.
func GetTotalUploadStat(ctx context.Context) (model.UploadStat, error) {
var total model.UploadStat
if err := db.DB(ctx).
Where("dimension = ? AND stat_key = ?", model.UploadStatDimensionTotal, "").
First(&total).Error; err != nil {
return model.UploadStat{}, err
}
return total, nil
}
// ListUploadStatsByDimension returns stats rows for a single dimension.
func ListUploadStatsByDimension(ctx context.Context, dimension string) ([]model.UploadStat, error) {
var rows []model.UploadStat
if err := db.DB(ctx).Where("dimension = ?", dimension).Find(&rows).Error; err != nil {
return nil, err
}
return rows, nil
}
// DeleteAllUploadStatsTx removes every row from w_upload_stats within a transaction.
func DeleteAllUploadStatsTx(tx *gorm.DB) error {
return tx.Where("1 = 1").Delete(&model.UploadStat{}).Error
}
// UpsertUploadStatDeltaTx applies an incremental count/size delta for one dimension key.
func UpsertUploadStatDeltaTx(tx *gorm.DB, dimension, key string, countDelta, sizeDelta int64) error {
return tx.Clauses(clause.OnConflict{
Columns: []clause.Column{
{Name: "dimension"},
{Name: "stat_key"},
},
DoUpdates: clause.Assignments(map[string]any{
"file_count": gorm.Expr(
"CASE WHEN w_upload_stats.file_count + ? < 0 THEN 0 ELSE w_upload_stats.file_count + ? END",
countDelta,
countDelta,
),
"file_size": gorm.Expr(
"CASE WHEN w_upload_stats.file_size + ? < 0 THEN 0 ELSE w_upload_stats.file_size + ? END",
sizeDelta,
sizeDelta,
),
"updated_at": time.Now(),
}),
}).Create(&model.UploadStat{
Dimension: dimension,
StatKey: key,
FileCount: countDelta,
FileSize: sizeDelta,
}).Error
}
// RebuildUploadStats clears w_upload_stats and re-applies deltas for every active upload
// inside a single transaction. applyDelta should apply +1 stats for one upload row.
func RebuildUploadStats(ctx context.Context, applyDelta func(tx *gorm.DB, upload *model.Upload) error) error {
return db.DB(ctx).Transaction(func(tx *gorm.DB) error {
if err := DeleteAllUploadStatsTx(tx); err != nil {
return err
}
uploads, err := ListActiveUploadsTx(tx)
if err != nil {
return err
}
for i := range uploads {
if err := applyDelta(tx, &uploads[i]); err != nil {
return err
}
}
return nil
})
}
+113
View File
@@ -5,8 +5,11 @@ package repository
import (
"context"
"errors"
"time"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/infra/persistence/idgen"
"github.com/Rain-kl/Wavelet/internal/model"
"gorm.io/gorm"
)
@@ -175,7 +178,117 @@ func ListUsersByIDs(ctx context.Context, ids []uint64) ([]model.User, error) {
return users, nil
}
// ListUserIDsByUsernameContains returns user IDs whose username contains the given fragment.
func ListUserIDsByUsernameContains(ctx context.Context, username string) ([]uint64, error) {
if username == "" {
return []uint64{}, nil
}
var userIDs []uint64
if err := db.DB(ctx).Model(&model.User{}).
Where("username LIKE ?", "%"+username+"%").
Pluck("id", &userIDs).Error; err != nil {
return nil, err
}
return userIDs, nil
}
// UpdateUser updates all fields of an existing user.
func UpdateUser(ctx context.Context, user *model.User) error {
return db.DB(ctx).Save(user).Error
}
// CreateUserFromOAuth creates a user from OAuth profile data and fills userOut.
func CreateUserFromOAuth(ctx context.Context, userOut *model.User, oauthInfo *model.OAuthUserInfo) error {
now := time.Now()
userID := oauthInfo.GetID()
newUser := model.User{
ID: userID,
Username: oauthInfo.Username,
Nickname: oauthInfo.Name,
Email: oauthInfo.Email,
AvatarURL: oauthInfo.AvatarURL,
IsActive: oauthInfo.Active,
LastLoginAt: now,
IsAdmin: false,
}
if newUser.ID == 0 {
newUser.ID = idgen.NextUint64ID()
}
if err := db.DB(ctx).Create(&newUser).Error; err != nil {
return err
}
*userOut = newUser
return nil
}
// ListUsernamesMatchingBase returns usernames equal to base or prefixed with base+"-".
func ListUsernamesMatchingBase(ctx context.Context, base string) ([]string, error) {
var names []string
if err := db.DB(ctx).Model(&model.User{}).
Where("username = ? OR username LIKE ?", base, base+"-%").
Pluck("username", &names).Error; err != nil {
return nil, err
}
return names, nil
}
// GetActiveUserByID loads a user by ID who is active.
func GetActiveUserByID(ctx context.Context, id uint64) (model.User, error) {
var user model.User
if err := db.DB(ctx).Where("id = ? AND is_active = ?", id, true).First(&user).Error; err != nil {
return model.User{}, err
}
return user, nil
}
// GetUserByUsernameOrEmail loads a user by username or email.
func GetUserByUsernameOrEmail(ctx context.Context, input string) (model.User, error) {
var user model.User
if err := db.DB(ctx).Where("username = ? OR email = ?", input, input).First(&user).Error; err != nil {
return model.User{}, err
}
return user, nil
}
// CountUsersByEmailExceptID counts users with the email excluding a given user id.
func CountUsersByEmailExceptID(ctx context.Context, email string, exceptID uint64) (int64, error) {
var count int64
if err := db.DB(ctx).Model(&model.User{}).Where("email = ? AND id != ?", email, exceptID).Count(&count).Error; err != nil {
return 0, err
}
return count, nil
}
// UpdateUserLastLoginAt updates only last_login_at for a user.
func UpdateUserLastLoginAt(ctx context.Context, userID uint64, at time.Time) error {
return db.DB(ctx).Model(&model.User{}).Where("id = ?", userID).Update("last_login_at", at).Error
}
// UpdateUserPassword updates only the password hash for a user.
func UpdateUserPassword(ctx context.Context, userID uint64, passwordHash string) error {
return db.DB(ctx).Model(&model.User{}).Where("id = ?", userID).Update("password", passwordHash).Error
}
// RegisterUserWithChecks validates username/email uniqueness then creates the user.
func RegisterUserWithChecks(ctx context.Context, user *model.User) error {
var count int64
if err := db.DB(ctx).Model(&model.User{}).Where("username = ?", user.Username).Count(&count).Error; err != nil {
return err
}
if count > 0 {
return errors.New("用户名已存在")
}
if user.Email != "" {
var emailCount int64
if err := db.DB(ctx).Model(&model.User{}).Where("email = ?", user.Email).Count(&emailCount).Error; err != nil {
return err
}
if emailCount > 0 {
return errors.New("该邮箱已被其他账号绑定")
}
}
if user.ID == 0 {
user.ID = idgen.NextUint64ID()
}
return db.DB(ctx).Create(user).Error
}