refactor(repo): replace legacy openflare-server with Wavelet rename

Delete the old monolithic openflare-server implementation and rename
Wavelet/ to openflare-server/ to complete the migration consolidation.
Update CI workflows, agent Dockerfiles, and deployment docs for the new
layout (frontend/, docker/Dockerfile).
This commit is contained in:
ryan
2026-06-19 11:29:17 +08:00
parent 46123d62ae
commit 88360350a0
1305 changed files with 40652 additions and 157629 deletions
@@ -0,0 +1,60 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package model 定义数据模型与 GORM 实体
package model
import (
"crypto/rand"
"crypto/sha256"
"encoding/hex"
"fmt"
"time"
)
const (
tokenByteLength = 24 // Token 随机字节长度
maskThreshold = 8 // 脱敏显示阈值
)
// AccessToken 个人访问令牌实体
type AccessToken struct {
ID uint64 `json:"id" gorm:"primaryKey;autoIncrement"`
UserID uint64 `json:"user_id" gorm:"index;not null"`
Name string `json:"name" gorm:"size:128;not null"`
TokenHash string `json:"-" gorm:"size:64;uniqueIndex;not null"`
MaskedToken string `json:"masked_token" gorm:"size:64;not null"`
IsAdmin bool `json:"is_admin" gorm:"default:false"`
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime"`
UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime"`
}
// TableName 表名
func (AccessToken) TableName() string {
return "w_access_tokens"
}
// GenerateTokenString 生成加密安全的随机 Token 值
func GenerateTokenString() (string, error) {
bytes := make([]byte, tokenByteLength)
if _, err := rand.Read(bytes); err != nil {
return "", err
}
return fmt.Sprintf("at_%s", hex.EncodeToString(bytes)), nil
}
// HashToken 计算 Token 的 SHA-256 哈希值用于数据库存储与查询
func HashToken(token string) string {
h := sha256.New()
h.Write([]byte(token))
return hex.EncodeToString(h.Sum(nil))
}
// MaskTokenString 生成脱敏显示的 Token,仅保留前缀和最后四位
func MaskTokenString(token string) string {
if len(token) <= maskThreshold {
return "at_****"
}
return fmt.Sprintf("%s...%s", token[:7], token[len(token)-4:])
}
@@ -1,41 +0,0 @@
package model
import "time"
type AcmeAccount struct {
ID uint `json:"id" gorm:"primaryKey"`
Email string `json:"email" gorm:"size:255"`
URL string `json:"url" gorm:"size:255"`
PrivateKey string `json:"-" gorm:"type:text;not null"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
}
func GetAcmeAccountByID(id uint) (*AcmeAccount, error) {
account := &AcmeAccount{}
err := DB.First(account, id).Error
return account, err
}
func GetDefaultAcmeAccount() (*AcmeAccount, error) {
account := &AcmeAccount{}
err := DB.Order("id asc").First(account).Error
if err != nil {
// Auto-create a default account placeholder if none exists
account.Email = "admin@openflare.dev"
err = DB.Create(account).Error
}
return account, err
}
func (account *AcmeAccount) Insert() error {
return DB.Create(account).Error
}
func (account *AcmeAccount) Update() error {
return DB.Save(account).Error
}
func (account *AcmeAccount) Delete() error {
return DB.Delete(account).Error
}
@@ -1,81 +0,0 @@
package model
import (
"time"
"gorm.io/gorm"
)
type ApplyLogQuery struct {
NodeID string
PageNo int
PageSize int
}
type ApplyLog struct {
ID uint `json:"id" gorm:"primaryKey"`
NodeID string `json:"node_id" gorm:"index;size:64;not null"`
Version string `json:"version" gorm:"size:32;not null"`
Result string `json:"result" gorm:"size:32;not null"`
Message string `json:"message" gorm:"type:text"`
Checksum string `json:"checksum" gorm:"size:64;not null;default:''"`
MainConfigChecksum string `json:"main_config_checksum" gorm:"size:64;not null;default:''"`
RouteConfigChecksum string `json:"route_config_checksum" gorm:"size:64;not null;default:''"`
SupportFileCount int `json:"support_file_count" gorm:"not null;default:0"`
CreatedAt time.Time `json:"created_at"`
}
func ListApplyLogs(query ApplyLogQuery) (logs []*ApplyLog, err error) {
db := DB.Order("id desc")
if query.NodeID != "" {
db = db.Where("node_id = ?", query.NodeID)
}
if query.PageSize > 0 {
offset := 0
if query.PageNo > 1 {
offset = (query.PageNo - 1) * query.PageSize
}
db = db.Limit(query.PageSize).Offset(offset)
}
err = db.Find(&logs).Error
return logs, err
}
func CountApplyLogs(nodeID string) (total int64, err error) {
query := DB.Model(&ApplyLog{})
if nodeID != "" {
query = query.Where("node_id = ?", nodeID)
}
err = query.Count(&total).Error
return total, err
}
func GetLatestApplyLogsByNodeIDs(nodeIDs []string) (map[string]*ApplyLog, error) {
result := make(map[string]*ApplyLog)
if len(nodeIDs) == 0 {
return result, nil
}
var logs []*ApplyLog
subQuery := DB.Model(&ApplyLog{}).
Select("MAX(id) AS id").
Where("node_id IN ?", nodeIDs).
Group("node_id")
if err := DB.Where("id IN (?)", subQuery).Find(&logs).Error; err != nil {
return nil, err
}
for _, log := range logs {
result[log.NodeID] = log
}
return result, nil
}
func DeleteAllApplyLogs() (deleted int64, err error) {
result := DB.Session(&gorm.Session{AllowGlobalUpdate: true}).Delete(&ApplyLog{})
return result.RowsAffected, result.Error
}
func DeleteApplyLogsBefore(before time.Time) (deleted int64, err error) {
result := DB.Where("created_at < ?", before).Delete(&ApplyLog{})
return result.RowsAffected, result.Error
}
+148 -109
View File
@@ -1,31 +1,35 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package model
import (
"context"
"errors"
"regexp"
"strings"
"time"
"github.com/rain-kl/openflare/pkg/utils"
"github.com/Rain-kl/Wavelet/internal/db"
"gorm.io/gorm"
)
// 认证源类型
const (
AuthSourceTypeGitHub = "github"
AuthSourceTypeOIDC = "oidc"
AuthSourceTypeOIDC = "oidc"
)
var authSourceNamePattern = regexp.MustCompile(`^[A-Za-z0-9][A-Za-z0-9_-]{0,79}$`)
// AuthSource 认证源实体
type AuthSource struct {
ID uint `json:"id"`
ID uint64 `json:"id" gorm:"primaryKey"`
Name string `json:"name" gorm:"uniqueIndex;size:80;not null"`
Type string `json:"type" gorm:"index;size:20;not null"`
Type string `json:"type" gorm:"size:20;not null"`
DisplayName string `json:"display_name" gorm:"size:100"`
IsActive bool `json:"is_active" gorm:"index;not null;default:false"`
ClientID string `json:"client_id" gorm:"column:client_id;size:255"`
ClientSecret string `json:"-" gorm:"column:client_secret;size:1024"`
ClientID string `json:"client_id" gorm:"size:255"`
ClientSecret string `json:"-" gorm:"size:1024"`
OpenIDDiscoveryURL string `json:"openid_discovery_url" gorm:"column:openid_discovery_url;size:1024"`
Scopes string `json:"scopes" gorm:"size:255"`
IconURL string `json:"icon_url" gorm:"size:1024"`
@@ -34,21 +38,32 @@ type AuthSource struct {
ClientSecretConfigured bool `json:"client_secret_configured" gorm:"-"`
}
type ExternalAccount struct {
ID uint `json:"id"`
AuthSourceID uint `json:"auth_source_id" gorm:"uniqueIndex:idx_external_account_source_external;index;not null"`
UserID int `json:"user_id" gorm:"index;not null"`
ExternalID string `json:"external_id" gorm:"uniqueIndex:idx_external_account_source_external;size:255;not null"`
ExternalUsername string `json:"external_username" gorm:"size:255"`
Email string `json:"email" gorm:"size:255"`
AuthSource AuthSource `json:"-" gorm:"constraint:OnDelete:CASCADE"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
// TableName 表名
func (AuthSource) TableName() string {
return "w_auth_sources"
}
// ExternalAccount 外部账号绑定实体
type ExternalAccount struct {
ID uint64 `json:"id" gorm:"primaryKey"`
AuthSourceID uint64 `json:"auth_source_id" gorm:"uniqueIndex:idx_external_accounts_source_external,priority:1;index"`
UserID uint64 `json:"user_id" gorm:"index;not null"`
ExternalID string `json:"external_id" gorm:"uniqueIndex:idx_external_accounts_source_external,priority:2;size:255;not null"`
ExternalUsername string `json:"external_username" gorm:"size:255"`
Email string `json:"email" gorm:"size:255"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
}
// TableName 表名
func (ExternalAccount) TableName() string {
return "w_external_accounts"
}
// ExternalAccountView 外部帐号绑定视图(脱敏展示用)
type ExternalAccountView struct {
ID uint `json:"id"`
AuthSourceID uint `json:"auth_source_id"`
ID uint64 `json:"id"`
AuthSourceID uint64 `json:"auth_source_id"`
AuthSourceName string `json:"auth_source_name"`
AuthSourceType string `json:"auth_source_type"`
AuthSourceLabel string `json:"auth_source_label"`
@@ -57,115 +72,117 @@ type ExternalAccountView struct {
CreatedAt time.Time `json:"created_at"`
}
// Normalize 对认证源字段进行标准化处理
func (source *AuthSource) Normalize() {
source.Type = strings.ToLower(source.Type)
utils.TrimStringFields(
&source.Name,
&source.Type,
&source.DisplayName,
&source.ClientID,
&source.ClientSecret,
&source.OpenIDDiscoveryURL,
&source.Scopes,
&source.IconURL,
)
source.Type = strings.ToLower(strings.TrimSpace(source.Type))
source.Name = strings.TrimSpace(source.Name)
source.DisplayName = strings.TrimSpace(source.DisplayName)
source.ClientID = strings.TrimSpace(source.ClientID)
source.ClientSecret = strings.TrimSpace(source.ClientSecret)
source.OpenIDDiscoveryURL = strings.TrimSpace(source.OpenIDDiscoveryURL)
source.Scopes = strings.TrimSpace(source.Scopes)
source.IconURL = strings.TrimSpace(source.IconURL)
if source.DisplayName == "" {
source.DisplayName = source.Name
}
if source.Type == AuthSourceTypeOIDC && source.Scopes == "" {
source.Scopes = "openid profile email"
}
if source.Type == AuthSourceTypeGitHub && source.Scopes == "" {
source.Scopes = "user:email"
}
}
// Validate 校验认证源字段合法性
func (source *AuthSource) Validate() error {
source.Normalize()
if source.Name == "" {
return errors.New("认证源名称不能为空")
return errors.New(errAuthSourceNameRequired)
}
if !authSourceNamePattern.MatchString(source.Name) {
return errors.New("认证源名称只能包含字母、数字、短横线或下划线,且必须以字母或数字开头")
return errors.New(errAuthSourceNameInvalid)
}
switch source.Type {
case AuthSourceTypeGitHub:
case AuthSourceTypeOIDC:
if source.OpenIDDiscoveryURL == "" {
return errors.New("OIDC 认证源必须配置 Discovery URL")
}
default:
return errors.New("认证源类型仅支持 github 或 oidc")
if source.Type != AuthSourceTypeOIDC {
return errors.New(errAuthSourceTypeUnsupported)
}
if source.IsActive {
if source.ClientID == "" || source.ClientSecret == "" {
return errors.New("启用认证源前必须配置 Client ID 和 Client Secret")
}
if source.OpenIDDiscoveryURL == "" {
return errors.New(errAuthSourceDiscoveryURLRequired)
}
if source.IsActive && (source.ClientID == "" || source.ClientSecret == "") {
return errors.New(errAuthSourceClientCredentialsRequired)
}
return nil
}
// Sanitize 脱敏处理,将 ClientSecret 清空并设置 ClientSecretConfigured 标志
func (source *AuthSource) Sanitize() {
source.ClientSecretConfigured = source.ClientSecret != ""
source.ClientSecret = ""
}
func GetAuthSources() ([]AuthSource, error) {
// GetAuthSources 获取所有认证源(已脱敏)
func GetAuthSources(ctx context.Context) ([]AuthSource, error) {
var sources []AuthSource
err := DB.Order("id asc").Find(&sources).Error
for index := range sources {
sources[index].Sanitize()
if err := db.DB(ctx).Order("id asc").Find(&sources).Error; err != nil {
return nil, err
}
return sources, err
for i := range sources {
sources[i].Sanitize()
}
return sources, nil
}
func GetActiveAuthSources() ([]AuthSource, error) {
// GetActiveAuthSources 获取所有已启用的认证源(已脱敏)
func GetActiveAuthSources(ctx context.Context) ([]AuthSource, error) {
var sources []AuthSource
err := DB.Where("is_active = ?", true).Order("id asc").Find(&sources).Error
for index := range sources {
sources[index].Sanitize()
if err := db.DB(ctx).Where("is_active = ?", true).Order("id asc").Find(&sources).Error; err != nil {
return nil, err
}
return sources, err
for i := range sources {
sources[i].Sanitize()
}
return sources, nil
}
func GetAuthSourceByID(id uint) (*AuthSource, error) {
// GetAuthSourceByID 根据 ID 获取认证源
func GetAuthSourceByID(ctx context.Context, id uint64) (*AuthSource, error) {
if id == 0 {
return nil, errors.New("认证源 ID 不能为空")
return nil, errors.New(errAuthSourceIDRequired)
}
var source AuthSource
if err := DB.First(&source, "id = ?", id).Error; err != nil {
if err := db.DB(ctx).First(&source, "id = ?", id).Error; err != nil {
return nil, err
}
source.ClientSecretConfigured = source.ClientSecret != ""
return &source, nil
}
func GetAuthSourceByName(name string) (*AuthSource, error) {
// GetAuthSourceByName 根据名称获取认证源(名称比较不区分大小写)
func GetAuthSourceByName(ctx context.Context, name string) (*AuthSource, error) {
name = strings.TrimSpace(name)
if name == "" {
return nil, errors.New("认证源名称不能为空")
return nil, errors.New(errAuthSourceNameRequired)
}
var source AuthSource
if err := DB.First(&source, "name = ?", name).Error; err != nil {
if err := db.DB(ctx).First(&source, "LOWER(name) = LOWER(?)", name).Error; err != nil {
return nil, err
}
source.ClientSecretConfigured = source.ClientSecret != ""
return &source, nil
}
func CreateAuthSource(source *AuthSource) error {
// CreateAuthSource 创建认证源
func CreateAuthSource(ctx context.Context, source *AuthSource) error {
if err := source.Validate(); err != nil {
return err
}
return DB.Create(source).Error
return db.DB(ctx).Create(source).Error
}
func UpdateAuthSource(source *AuthSource, keepSecret bool) error {
// UpdateAuthSource 更新认证源,keepSecret 为 true 时保留原密钥
func UpdateAuthSource(ctx context.Context, source *AuthSource, keepSecret bool) error {
if source.ID == 0 {
return errors.New("认证源 ID 不能为空")
return errors.New(errAuthSourceIDRequired)
}
var current AuthSource
if err := DB.First(&current, "id = ?", source.ID).Error; err != nil {
if err := db.DB(ctx).First(&current, "id = ?", source.ID).Error; err != nil {
return err
}
if keepSecret {
@@ -174,7 +191,7 @@ func UpdateAuthSource(source *AuthSource, keepSecret bool) error {
if err := source.Validate(); err != nil {
return err
}
return DB.Model(&current).Updates(map[string]any{
return db.DB(ctx).Model(&current).Updates(map[string]any{
"name": source.Name,
"type": source.Type,
"display_name": source.DisplayName,
@@ -187,8 +204,9 @@ func UpdateAuthSource(source *AuthSource, keepSecret bool) error {
}).Error
}
func ToggleAuthSource(id uint, isActive bool) error {
source, err := GetAuthSourceByID(id)
// ToggleAuthSource 切换认证源启用状态
func ToggleAuthSource(ctx context.Context, id uint64, isActive bool) error {
source, err := GetAuthSourceByID(ctx, id)
if err != nil {
return err
}
@@ -196,14 +214,15 @@ func ToggleAuthSource(id uint, isActive bool) error {
if err := source.Validate(); err != nil {
return err
}
return DB.Model(&AuthSource{}).Where("id = ?", id).Update("is_active", isActive).Error
return db.DB(ctx).Model(&AuthSource{}).Where("id = ?", id).Update("is_active", isActive).Error
}
func DeleteAuthSource(id uint) error {
// DeleteAuthSource 删除认证源及其关联的外部帐号绑定
func DeleteAuthSource(ctx context.Context, id uint64) error {
if id == 0 {
return errors.New("认证源 ID 不能为空")
return errors.New(errAuthSourceIDRequired)
}
return DB.Transaction(func(tx *gorm.DB) error {
return db.DB(ctx).Transaction(func(tx *gorm.DB) error {
if err := tx.Where("auth_source_id = ?", id).Delete(&ExternalAccount{}).Error; err != nil {
return err
}
@@ -211,47 +230,76 @@ func DeleteAuthSource(id uint) error {
})
}
func FindExternalAccount(sourceID uint, externalID string) (*ExternalAccount, error) {
// FindExternalAccount 查找外部帐号绑定记录
func FindExternalAccount(ctx context.Context, sourceID uint64, externalID string) (*ExternalAccount, error) {
var account ExternalAccount
err := DB.Where("auth_source_id = ? AND external_id = ?", sourceID, externalID).First(&account).Error
if err != nil {
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
}
func LinkExternalAccount(account *ExternalAccount) error {
if account.AuthSourceID == 0 || account.UserID == 0 || strings.TrimSpace(account.ExternalID) == "" {
return errors.New("外部账号绑定信息不完整")
// BindExternalAccount 绑定外部帐号(已存在时更新用户名和邮箱)
func BindExternalAccount(ctx context.Context, account *ExternalAccount) error {
if account.UserID == 0 || strings.TrimSpace(account.ExternalID) == "" {
return errors.New(errExternalAccountBindingIncomplete)
}
account.ExternalID = strings.TrimSpace(account.ExternalID)
account.ExternalUsername = strings.TrimSpace(account.ExternalUsername)
account.Email = strings.TrimSpace(account.Email)
return DB.Where(ExternalAccount{
AuthSourceID: account.AuthSourceID,
ExternalID: account.ExternalID,
}).FirstOrCreate(account).Error
return db.DB(ctx).Transaction(func(tx *gorm.DB) error {
var current ExternalAccount
err := tx.Where("auth_source_id = ? AND external_id = ?", account.AuthSourceID, account.ExternalID).First(&current).Error
if err == nil {
if current.UserID != account.UserID {
return errors.New(errExternalAccountAlreadyBoundToAnother)
}
return tx.Model(&current).Updates(map[string]any{
"external_username": account.ExternalUsername,
"email": account.Email,
}).Error
}
if !errors.Is(err, gorm.ErrRecordNotFound) {
return err
}
return tx.Create(account).Error
})
}
func ListExternalAccountsByUserID(userID int) ([]ExternalAccountView, error) {
if userID <= 0 {
return nil, errors.New("用户 ID 不能为空")
// ListExternalAccountsByUserID 获取指定用户的所有外部帐号绑定视图
func ListExternalAccountsByUserID(ctx context.Context, userID uint64) ([]ExternalAccountView, error) {
if userID == 0 {
return nil, errors.New(errUserIDRequired)
}
var accounts []ExternalAccount
if err := DB.Preload("AuthSource").Where("user_id = ?", userID).Order("id asc").Find(&accounts).Error; err != nil {
if err := db.DB(ctx).Where("user_id = ?", userID).Order("id asc").Find(&accounts).Error; err != nil {
return nil, err
}
views := make([]ExternalAccountView, 0, len(accounts))
for _, account := range accounts {
label := account.AuthSource.DisplayName
if label == "" {
label = account.AuthSource.Name
var name, sourceType, label string
if account.AuthSourceID == 0 {
name = "default"
sourceType = "oidc"
label = "历史认证源"
} else {
source, err := GetAuthSourceByID(ctx, account.AuthSourceID)
if err != nil {
continue
}
name = source.Name
sourceType = source.Type
label = source.DisplayName
if label == "" {
label = source.Name
}
}
views = append(views, ExternalAccountView{
ID: account.ID,
AuthSourceID: account.AuthSourceID,
AuthSourceName: account.AuthSource.Name,
AuthSourceType: account.AuthSource.Type,
AuthSourceName: name,
AuthSourceType: sourceType,
AuthSourceLabel: label,
ExternalUsername: account.ExternalUsername,
Email: account.Email,
@@ -261,19 +309,10 @@ func ListExternalAccountsByUserID(userID int) ([]ExternalAccountView, error) {
return views, nil
}
func DeleteExternalAccountForUser(id uint, userID int) error {
if id == 0 {
return errors.New("绑定记录 ID 不能为空")
// DeleteExternalAccountForUser 删除指定用户的外部帐号绑定
func DeleteExternalAccountForUser(ctx context.Context, id uint64, userID uint64) error {
if id == 0 || userID == 0 {
return errors.New(errExternalAccountBindingIDRequired)
}
if userID <= 0 {
return errors.New("用户 ID 不能为空")
}
result := DB.Where("id = ? AND user_id = ?", id, userID).Delete(&ExternalAccount{})
if result.Error != nil {
return result.Error
}
if result.RowsAffected == 0 {
return errors.New("绑定记录不存在")
}
return nil
return db.DB(ctx).Where("id = ? AND user_id = ?", id, userID).Delete(&ExternalAccount{}).Error
}
@@ -1,45 +0,0 @@
package model
import "time"
type ConfigVersionSummary struct {
ID uint `json:"id"`
Version string `json:"version"`
Checksum string `json:"checksum"`
IsActive bool `json:"is_active"`
CreatedBy string `json:"created_by"`
CreatedAt time.Time `json:"created_at"`
}
type ConfigVersion struct {
ID uint `json:"id" gorm:"primaryKey"`
Version string `json:"version" gorm:"uniqueIndex;size:32;not null"`
SnapshotJSON string `json:"snapshot_json" gorm:"type:text;not null"`
MainConfig string `json:"main_config" gorm:"type:text;not null;default:''"`
RenderedConfig string `json:"rendered_config" gorm:"type:text;not null"`
SupportFilesJSON string `json:"support_files_json" gorm:"type:text;not null;default:'[]'"`
Checksum string `json:"checksum" gorm:"size:64;not null"`
IsActive bool `json:"is_active" gorm:"not null;default:false;index"`
CreatedBy string `json:"created_by" gorm:"size:64;not null"`
CreatedAt time.Time `json:"created_at"`
}
func ListConfigVersionSummaries() (versions []*ConfigVersionSummary, err error) {
err = DB.Model(&ConfigVersion{}).
Select("id", "version", "checksum", "is_active", "created_by", "created_at").
Order("id desc").
Find(&versions).Error
return versions, err
}
func GetConfigVersionByID(id uint) (*ConfigVersion, error) {
version := &ConfigVersion{}
err := DB.First(version, id).Error
return version, err
}
func GetActiveConfigVersion() (*ConfigVersion, error) {
version := &ConfigVersion{}
err := DB.Where("is_active = ?", true).Order("id desc").First(version).Error
return version, err
}
@@ -1,27 +0,0 @@
package model
import (
"time"
"github.com/rain-kl/openflare/openflare-server/internal/model/migrate"
)
const (
legacyDatabaseSchemaVersion = migrate.BaseDatabaseSchemaVersion
legacyMigrationTerminalVersion = 17
databaseSchemaVersionRowID = 1
)
// currentDatabaseSchemaVersion tracks the current physical schema validated by the
// legacy validator set. Goose owns only post-v17 migrations, and none exist yet.
var currentDatabaseSchemaVersion = legacyMigrationTerminalVersion
type DatabaseSchemaVersion struct {
ID uint `json:"id" gorm:"primaryKey"`
Version int `json:"version" gorm:"not null"`
UpdatedAt time.Time `json:"updated_at"`
}
func (DatabaseSchemaVersion) TableName() string {
return "database_schema_versions"
}
@@ -1,35 +0,0 @@
package model
import "time"
type DnsAccount struct {
ID uint `json:"id" gorm:"primaryKey"`
Name string `json:"name" gorm:"size:255;not null"`
Type string `json:"type" gorm:"size:64;not null"`
Authorization string `json:"-" gorm:"type:text;not null"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
}
func ListDnsAccounts() (accounts []*DnsAccount, err error) {
err = DB.Order("id desc").Find(&accounts).Error
return accounts, err
}
func GetDnsAccountByID(id uint) (*DnsAccount, error) {
account := &DnsAccount{}
err := DB.First(account, id).Error
return account, err
}
func (account *DnsAccount) Insert() error {
return DB.Create(account).Error
}
func (account *DnsAccount) Update() error {
return DB.Save(account).Error
}
func (account *DnsAccount) Delete() error {
return DB.Delete(account).Error
}
+30
View File
@@ -0,0 +1,30 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package model
const (
errRegistrationDisabled = "注册已关闭"
errDatabaseNotInitialized = "database not initialized"
errUsernameExists = "用户名已存在"
errEmailAlreadyBound = "该邮箱已被其他账号绑定"
errConfigIntParseFailed = "配置 %s 的值 '%s' 无法转换为整数: %w"
errConfigDecimalParseFailed = "配置 %s 的值 '%s' 无法转换为decimal: %w"
errConfigBoolParseFailed = "配置 %s 的值 '%s' 无法转换为布尔值: %w"
errParseMenuDisplayConfigFailed = "解析目录显示配置失败: %w"
errTemplateKeyRequired = "模板标识符不能为空"
errTemplateNameRequired = "模板名称不能为空"
errTemplateContentRequired = "模板内容不能为空"
errTemplateUnavailable = "模板 %s 不存在或不可用: %w"
errTemplateRenderFailed = "模板 %s 渲染失败: %w"
errAuthSourceNameRequired = "认证源名称不能为空"
errAuthSourceNameInvalid = "认证源名称只能包含字母、数字、短横线或下划线,且必须以字母或数字开头"
errAuthSourceTypeUnsupported = "认证源类型仅支持 oidc"
errAuthSourceDiscoveryURLRequired = "OIDC 认证源必须配置 Discovery URL"
errAuthSourceClientCredentialsRequired = "启用认证源前必须配置 Client ID 和 Client Secret" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
errAuthSourceIDRequired = "认证源 ID 不能为空"
errExternalAccountBindingIncomplete = "外部账号绑定信息不完整"
errExternalAccountAlreadyBoundToAnother = "该外部账号已绑定到其他用户"
errUserIDRequired = "用户 ID 不能为空"
errExternalAccountBindingIDRequired = "绑定记录 ID 不能为空"
)
@@ -1,326 +0,0 @@
package goose
import (
"database/sql"
"errors"
"fmt"
"io"
"log/slog"
"os"
"gorm.io/gorm"
)
type BridgeContext interface {
Context
AutoMigrateLegacySchemaMetadata(db *gorm.DB) error
InitializeFreshDatabaseSchema(db *gorm.DB, backend string) error
IsDatabaseEmpty(db *gorm.DB) (bool, error)
RepairCurrentSchemaState(db *gorm.DB, backend string) error
SaveLegacyDatabaseSchemaVersion(db *gorm.DB, version int) error
UpgradeLegacyDatabaseSchema(db *gorm.DB, backend string, version int) error
ValidateCurrentDatabaseSchema(db *gorm.DB, backend string) error
}
type schemaMigrationState int
const (
schemaMigrationStateFresh schemaMigrationState = iota
schemaMigrationStateLegacyOnly
schemaMigrationStateGooseOnly
schemaMigrationStateLegacyBootstrap
schemaMigrationStateMixed
)
func detectSchemaState(db *gorm.DB, ctx BridgeContext) (schemaMigrationState, error) {
hasLegacyTable := db.Migrator().HasTable("database_schema_versions")
hasGooseTable := db.Migrator().HasTable("goose_db_version")
switch {
case hasLegacyTable && hasGooseTable:
return schemaMigrationStateMixed, nil
case hasLegacyTable:
return schemaMigrationStateLegacyOnly, nil
case hasGooseTable:
return schemaMigrationStateGooseOnly, nil
}
empty, err := ctx.IsDatabaseEmpty(db)
if err != nil {
return 0, err
}
if empty {
return schemaMigrationStateFresh, nil
}
return schemaMigrationStateLegacyBootstrap, nil
}
func LoadDatabaseVersion(db *gorm.DB) (int, bool, error) {
if db == nil || !db.Migrator().HasTable("goose_db_version") {
return 0, false, nil
}
var version int64
err := db.Table("goose_db_version").
Where("is_applied = ?", true).
Order("version_id DESC").
Select("version_id").
Limit(1).
Row().
Scan(&version)
if errors.Is(err, sql.ErrNoRows) {
return 0, false, nil
}
if err != nil {
return 0, false, err
}
return int(version), true, nil
}
func loadLegacyDatabaseSchemaVersion(db *gorm.DB) (int, bool, error) {
if db == nil || !db.Migrator().HasTable("database_schema_versions") {
return 0, false, nil
}
var version int
err := db.Table("database_schema_versions").
Where("id = ?", 1).
Select("version").
Limit(1).
Row().
Scan(&version)
if errors.Is(err, sql.ErrNoRows) {
return 0, false, nil
}
if err != nil {
return 0, false, err
}
return version, true, nil
}
func bootstrapLegacySchemaVersion(db *gorm.DB, ctx BridgeContext) error {
if err := ctx.AutoMigrateLegacySchemaMetadata(db); err != nil {
return err
}
version, exists, err := loadLegacyDatabaseSchemaVersion(db)
if err != nil {
return err
}
if exists {
if int64(version) > LegacyBridgeVersion {
return fmt.Errorf("legacy schema version %d is newer than supported terminal version %d", version, LegacyBridgeVersion)
}
return nil
}
return ctx.SaveLegacyDatabaseSchemaVersion(db, 7)
}
func upgradeLegacyToTerminal(db *gorm.DB, backend string, ctx BridgeContext) error {
if err := bootstrapLegacySchemaVersion(db, ctx); err != nil {
return err
}
version, exists, err := loadLegacyDatabaseSchemaVersion(db)
if err != nil {
return err
}
if !exists {
return fmt.Errorf("legacy schema version record is missing after bootstrap")
}
return ctx.UpgradeLegacyDatabaseSchema(db, backend, version)
}
func validateGooseBridgeState(db *gorm.DB) error {
version, exists, err := LoadDatabaseVersion(db)
if err != nil {
return err
}
if !exists {
return nil
}
if int64(version) < LegacyBridgeVersion {
return fmt.Errorf("goose schema version %d is below legacy bridge baseline %d", version, LegacyBridgeVersion)
}
if int64(version) > CurrentTargetVersion() {
return fmt.Errorf("goose schema version %d is newer than application target version %d", version, CurrentTargetVersion())
}
return nil
}
func finalizeLegacyToGooseBridge(db *gorm.DB) error {
gooseVersion, exists, err := LoadDatabaseVersion(db)
if err != nil {
return err
}
if !exists || int64(gooseVersion) < LegacyBridgeVersion {
return nil
}
if !db.Migrator().HasTable("database_schema_versions") {
return nil
}
if err := db.Exec("DROP TABLE IF EXISTS database_schema_versions").Error; err != nil {
return fmt.Errorf("drop legacy schema versions table failed: %w", err)
}
slog.Info("completed legacy-to-goose migration bridge", "goose_version", gooseVersion)
return nil
}
func ValidateRegisteredSchema(db *gorm.DB) error {
if err := validateNodeCapabilitiesJSON(db); err != nil {
return err
}
return nil
}
func EnsureDatabaseSchemaUpToDate(db *gorm.DB, backend string, ctx BridgeContext) (returnedErr error) {
if backend == "sqlite" {
backupPath, restore, err := backupSQLiteDatabase(db)
if err != nil {
slog.Warn("failed to backup sqlite database before migration", "error", err)
} else if backupPath != "" {
defer func() {
if returnedErr != nil {
restore()
} else {
os.Remove(backupPath)
}
}()
}
}
var startDesc string
legacyVer, hasLegacy, _ := loadLegacyDatabaseSchemaVersion(db)
gooseVer, hasGoose, _ := LoadDatabaseVersion(db)
if hasGoose {
startDesc = fmt.Sprintf("goose version %d", gooseVer)
} else if hasLegacy {
startDesc = fmt.Sprintf("legacy version %d", legacyVer)
} else {
startDesc = "none (fresh database)"
}
state, err := detectSchemaState(db, ctx)
if err != nil {
return err
}
switch state {
case schemaMigrationStateFresh:
if err := ctx.InitializeFreshDatabaseSchema(db, backend); err != nil {
return err
}
case schemaMigrationStateLegacyOnly:
if err := upgradeLegacyToTerminal(db, backend, ctx); err != nil {
return err
}
case schemaMigrationStateGooseOnly:
if err := validateGooseBridgeState(db); err != nil {
return err
}
case schemaMigrationStateLegacyBootstrap:
if err := upgradeLegacyToTerminal(db, backend, ctx); err != nil {
return err
}
case schemaMigrationStateMixed:
legacyVersion, exists, err := loadLegacyDatabaseSchemaVersion(db)
if err != nil {
return err
}
if exists && int64(legacyVersion) != LegacyBridgeVersion {
return fmt.Errorf("incomplete mixed migration state: legacy schema version %d does not match bridge terminal version %d", legacyVersion, LegacyBridgeVersion)
}
if err := validateGooseBridgeState(db); err != nil {
return err
}
default:
return fmt.Errorf("unknown schema migration state: %d", state)
}
if err := runMigrations(db, backend, ctx); err != nil {
return err
}
if err := finalizeLegacyToGooseBridge(db); err != nil {
return err
}
if err := ctx.RepairCurrentSchemaState(db, backend); err != nil {
return err
}
if err := ctx.ValidateCurrentDatabaseSchema(db, backend); err != nil {
return err
}
if err := ValidateRegisteredSchema(db); err != nil {
return err
}
endVer, _, _ := LoadDatabaseVersion(db)
if hasGoose && int64(gooseVer) == int64(endVer) {
slog.Info("database schema is already up to date", "version", endVer)
} else {
slog.Info("database migration completed successfully", "from", startDesc, "to", fmt.Sprintf("goose version %d", endVer))
}
return nil
}
func backupSQLiteDatabase(db *gorm.DB) (string, func(), error) {
var dbList []struct {
Seq int
Name string
File string
}
if err := db.Raw("PRAGMA database_list").Scan(&dbList).Error; err != nil {
return "", nil, err
}
var dbPath string
for _, item := range dbList {
if item.Name == "main" && item.File != "" {
dbPath = item.File
break
}
}
if dbPath == "" {
return "", nil, nil
}
backupPath := dbPath + ".bak"
src, err := os.Open(dbPath)
if err != nil {
return "", nil, err
}
defer src.Close()
dst, err := os.Create(backupPath)
if err != nil {
return "", nil, err
}
defer dst.Close()
if _, err := io.Copy(dst, src); err != nil {
return "", nil, err
}
dst.Sync()
restoreFunc := func() {
src, err := os.Open(backupPath)
if err != nil {
slog.Error("failed to open sqlite backup for restore", "error", err)
return
}
defer src.Close()
dst, err := os.OpenFile(dbPath, os.O_WRONLY|os.O_TRUNC, 0644)
if err != nil {
slog.Error("failed to open sqlite db for restore", "error", err)
return
}
defer dst.Close()
if _, err := io.Copy(dst, src); err != nil {
slog.Error("failed to restore sqlite backup", "error", err)
} else {
dst.Sync()
slog.Warn("restored sqlite database from backup due to migration failure")
}
}
return backupPath, restoreFunc, nil
}
@@ -1,50 +0,0 @@
package goose
import (
"encoding/json"
"fmt"
presslygoose "github.com/pressly/goose/v3"
"gorm.io/gorm"
)
const versionNodeCapabilitiesJSON int64 = 202606020001
// migration202606020001 adds a future-proof JSON field for node capability
// summaries after the legacy v17 migration bridge.
func migration202606020001(backend string, ctx Context) *presslygoose.Migration {
return newGORMMigration(
versionNodeCapabilitiesJSON,
"202606020001_add_node_capabilities_json.go",
backend,
ctx,
migrateNodeCapabilitiesJSON,
)
}
func migrateNodeCapabilitiesJSON(ctx Context, db *gorm.DB, backend string) error {
if err := ctx.ApplyCurrentSchema(db, backend); err != nil {
return err
}
emptyJSON, err := json.Marshal([]string{})
if err != nil {
return fmt.Errorf("marshal default node capabilities: %w", err)
}
if err := db.Exec(
`UPDATE nodes SET capabilities_json = ? WHERE capabilities_json IS NULL OR TRIM(capabilities_json) = ''`,
string(emptyJSON),
).Error; err != nil {
return fmt.Errorf("backfill nodes.capabilities_json: %w", err)
}
return validateNodeCapabilitiesJSON(db)
}
func validateNodeCapabilitiesJSON(db *gorm.DB) error {
if db == nil {
return fmt.Errorf("database handle is nil")
}
if !db.Migrator().HasColumn("nodes", "capabilities_json") {
return fmt.Errorf("column nodes.capabilities_json is missing")
}
return nil
}
@@ -1,61 +0,0 @@
package goose
import (
"fmt"
presslygoose "github.com/pressly/goose/v3"
"gorm.io/gorm"
)
const versionPagesStaticHosting int64 = 202606030001
// migration202606030001 adds OpenFlare Pages static hosting tables and the
// proxy_routes.pages_project_id binding used by the global release snapshot.
func migration202606030001(backend string, ctx Context) *presslygoose.Migration {
return newGORMMigration(
versionPagesStaticHosting,
"202606030001_add_pages_static_hosting.go",
backend,
ctx,
migratePagesStaticHosting,
)
}
func migratePagesStaticHosting(ctx Context, db *gorm.DB, backend string) error {
if err := ctx.ApplyCurrentSchema(db, backend); err != nil {
return err
}
if err := db.Exec(
`UPDATE proxy_routes SET upstream_type = 'direct' WHERE upstream_type IS NULL OR TRIM(upstream_type) = ''`,
).Error; err != nil {
return fmt.Errorf("backfill proxy_routes.upstream_type: %w", err)
}
return validatePagesStaticHosting(db)
}
func validatePagesStaticHosting(db *gorm.DB) error {
if db == nil {
return fmt.Errorf("database handle is nil")
}
for _, table := range []string{"pages_projects", "pages_deployments", "pages_deployment_files"} {
if !db.Migrator().HasTable(table) {
return fmt.Errorf("table %s is missing", table)
}
}
for _, column := range []string{"upstream_type", "pages_project_id"} {
if !db.Migrator().HasColumn("proxy_routes", column) {
return fmt.Errorf("column proxy_routes.%s is missing", column)
}
}
for _, column := range []string{"slug", "active_deployment_id", "spa_fallback_enabled", "spa_fallback_path"} {
if !db.Migrator().HasColumn("pages_projects", column) {
return fmt.Errorf("column pages_projects.%s is missing", column)
}
}
for _, column := range []string{"project_id", "checksum", "artifact_path"} {
if !db.Migrator().HasColumn("pages_deployments", column) {
return fmt.Errorf("column pages_deployments.%s is missing", column)
}
}
return nil
}
@@ -1,37 +0,0 @@
package goose
import (
"fmt"
presslygoose "github.com/pressly/goose/v3"
"gorm.io/gorm"
)
const versionPagesSPAFallbackPath int64 = 202606030002
// migration202606030002 adds a configurable SPA fallback path for Pages
// projects. Existing projects keep the previous /index.html behavior.
func migration202606030002(backend string, ctx Context) *presslygoose.Migration {
return newGORMMigration(
versionPagesSPAFallbackPath,
"202606030002_add_pages_spa_fallback_path.go",
backend,
ctx,
migratePagesSPAFallbackPath,
)
}
func migratePagesSPAFallbackPath(ctx Context, db *gorm.DB, backend string) error {
if err := ctx.ApplyCurrentSchema(db, backend); err != nil {
return err
}
if err := db.Exec(
`UPDATE pages_projects SET spa_fallback_path = '/index.html' WHERE spa_fallback_path IS NULL OR TRIM(spa_fallback_path) = ''`,
).Error; err != nil {
return fmt.Errorf("backfill pages_projects.spa_fallback_path: %w", err)
}
if !db.Migrator().HasColumn("pages_projects", "spa_fallback_path") {
return fmt.Errorf("column pages_projects.spa_fallback_path is missing")
}
return nil
}
@@ -1,41 +0,0 @@
package goose
import (
"fmt"
presslygoose "github.com/pressly/goose/v3"
"gorm.io/gorm"
)
const versionDropProxyRouteLegacyPoW int64 = 202606030003
// migration202606030003 drops the legacy pow_enabled and pow_config columns
// from proxy_routes table, since PoW is now entirely managed under WAF rule groups.
func migration202606030003(backend string, ctx Context) *presslygoose.Migration {
return newGORMMigration(
versionDropProxyRouteLegacyPoW,
"202606030003_drop_proxy_route_legacy_pow.go",
backend,
ctx,
migrateDropProxyRouteLegacyPoW,
)
}
func migrateDropProxyRouteLegacyPoW(ctx Context, db *gorm.DB, backend string) error {
if err := ctx.ApplyCurrentSchema(db, backend); err != nil {
return err
}
// Drop pow_enabled column if exists
if db.Migrator().HasColumn("proxy_routes", "pow_enabled") {
if err := db.Exec("ALTER TABLE proxy_routes DROP COLUMN pow_enabled").Error; err != nil {
return fmt.Errorf("drop proxy_routes.pow_enabled: %w", err)
}
}
// Drop pow_config column if exists
if db.Migrator().HasColumn("proxy_routes", "pow_config") {
if err := db.Exec("ALTER TABLE proxy_routes DROP COLUMN pow_config").Error; err != nil {
return fmt.Errorf("drop proxy_routes.pow_config: %w", err)
}
}
return nil
}
@@ -1,64 +0,0 @@
package goose
import (
"fmt"
presslygoose "github.com/pressly/goose/v3"
"gorm.io/gorm"
)
const versionPagesFeaturesAndCleanup int64 = 202606040004
// migration202606040004 merges migrations 202606030004, 202606040001, 202606040002, and 202606040003.
// It adds Pages API proxying fields and RootDir/EntryFile to Pages projects,
// backfills default entry_file to 'index.html', and ensures unused fields (root_dir, entry_file)
// are dropped from Pages deployments.
func migration202606040004(backend string, ctx Context) *presslygoose.Migration {
return newGORMMigration(
versionPagesFeaturesAndCleanup,
"202606040004_add_pages_features_and_cleanup.go",
backend,
ctx,
migratePagesFeaturesAndCleanup,
)
}
func migratePagesFeaturesAndCleanup(ctx Context, db *gorm.DB, backend string) error {
if err := ctx.ApplyCurrentSchema(db, backend); err != nil {
return err
}
// 1. Verify Pages projects columns
cols := []string{
"api_proxy_enabled", "api_proxy_path", "api_proxy_pass", "api_proxy_rewrite",
"root_dir", "entry_file",
}
for _, col := range cols {
if !db.Migrator().HasColumn("pages_projects", col) {
return fmt.Errorf("column pages_projects.%s is missing", col)
}
}
// 2. Backfill pages_projects.entry_file to 'index.html' if empty
type PagesProject struct {
ID uint `gorm:"primaryKey"`
EntryFile string `gorm:"size:512;not null;default:'index.html'"`
}
if err := db.Model(&PagesProject{}).Where("entry_file = '' OR entry_file IS NULL").Update("entry_file", "index.html").Error; err != nil {
return fmt.Errorf("failed to backfill pages_projects.entry_file: %w", err)
}
// 3. Drop unused fields root_dir and entry_file from pages_deployments if they exist
if db.Migrator().HasColumn("pages_deployments", "root_dir") {
if err := db.Exec("ALTER TABLE pages_deployments DROP COLUMN root_dir").Error; err != nil {
return fmt.Errorf("failed to drop pages_deployments.root_dir: %w", err)
}
}
if db.Migrator().HasColumn("pages_deployments", "entry_file") {
if err := db.Exec("ALTER TABLE pages_deployments DROP COLUMN entry_file").Error; err != nil {
return fmt.Errorf("failed to drop pages_deployments.entry_file: %w", err)
}
}
return nil
}
@@ -1,75 +0,0 @@
package goose
import (
"context"
"database/sql"
"fmt"
presslygoose "github.com/pressly/goose/v3"
"gorm.io/gorm"
)
const LegacyBridgeVersion int64 = 17
type migrationFunc func(ctx Context, db *gorm.DB, backend string) error
func newBaselineMigration() *presslygoose.Migration {
migration := presslygoose.NewGoMigration(LegacyBridgeVersion, nil, nil)
migration.Source = fmt.Sprintf("%05d_legacy_terminal_baseline.go", LegacyBridgeVersion)
return migration
}
func newGORMMigration(version int64, source string, backend string, ctx Context, up migrationFunc) *presslygoose.Migration {
migration := presslygoose.NewGoMigration(version, &presslygoose.GoFunc{
RunDB: func(_ context.Context, sqlDB *sql.DB) error {
gormDB, err := openGORMDB(ctx, sqlDB, backend)
if err != nil {
return err
}
if backend == "postgres" {
return gormDB.Transaction(func(tx *gorm.DB) error {
return up(ctx, tx, backend)
})
}
return up(ctx, gormDB, backend)
},
}, nil)
migration.Source = source
return migration
}
func registeredMigrations(backend string, ctx Context) []*presslygoose.Migration {
return []*presslygoose.Migration{
migration202606020001(backend, ctx),
migration202606030001(backend, ctx),
migration202606030002(backend, ctx),
migration202606030003(backend, ctx),
migration202606040004(backend, ctx),
}
}
func buildMigrations(backend string, ctx Context) []*presslygoose.Migration {
migrations := []*presslygoose.Migration{newBaselineMigration()}
migrations = append(migrations, registeredMigrations(backend, ctx)...)
return migrations
}
func CurrentTargetVersion() int64 {
var maxVersion int64 = LegacyBridgeVersion
for _, migration := range buildMigrations("sqlite", noopContext{}) {
if migration.Version > maxVersion {
maxVersion = migration.Version
}
}
return maxVersion
}
type noopContext struct{}
func (noopContext) ApplyCurrentSchema(db *gorm.DB, backend string) error {
return nil
}
func (noopContext) RegisterSharding(db *gorm.DB, backend string) error {
return nil
}
@@ -1,81 +0,0 @@
package goose
import (
"context"
"database/sql"
"fmt"
"github.com/glebarez/sqlite"
presslygoose "github.com/pressly/goose/v3"
"gorm.io/driver/postgres"
"gorm.io/gorm"
"gorm.io/gorm/schema"
)
type Context interface {
ApplyCurrentSchema(db *gorm.DB, backend string) error
RegisterSharding(db *gorm.DB, backend string) error
}
func dialectForBackend(backend string) (presslygoose.Dialect, error) {
switch backend {
case "postgres":
return presslygoose.DialectPostgres, nil
case "sqlite":
return presslygoose.DialectSQLite3, nil
default:
return "", fmt.Errorf("unsupported database backend: %s", backend)
}
}
func openGORMDB(ctx Context, db *sql.DB, backend string) (*gorm.DB, error) {
var dialector gorm.Dialector
switch backend {
case "postgres":
dialector = postgres.New(postgres.Config{Conn: db})
case "sqlite":
dialector = &sqlite.Dialector{Conn: db}
default:
return nil, fmt.Errorf("unsupported database backend: %s", backend)
}
gormDB, err := gorm.Open(dialector, &gorm.Config{
NamingStrategy: schema.NamingStrategy{},
})
if err != nil {
return nil, err
}
if err := ctx.RegisterSharding(gormDB, backend); err != nil {
return nil, err
}
return gormDB, nil
}
func buildProvider(db *gorm.DB, backend string, ctx Context) (*presslygoose.Provider, error) {
sqlDB, err := db.DB()
if err != nil {
return nil, err
}
dialect, err := dialectForBackend(backend)
if err != nil {
return nil, err
}
return presslygoose.NewProvider(
dialect,
sqlDB,
nil,
presslygoose.WithDisableGlobalRegistry(true),
presslygoose.WithGoMigrations(buildMigrations(backend, ctx)...),
)
}
func runMigrations(db *gorm.DB, backend string, ctx Context) error {
provider, err := buildProvider(db, backend, ctx)
if err != nil {
return fmt.Errorf("build goose provider: %w", err)
}
if _, err := provider.Up(context.Background()); err != nil {
return fmt.Errorf("goose up failed: %w", err)
}
return nil
}
-387
View File
@@ -1,387 +0,0 @@
package model
import (
"fmt"
"log/slog"
"os"
"reflect"
"strings"
"sync"
"github.com/rain-kl/openflare/openflare-server/internal/common"
"github.com/rain-kl/openflare/openflare-server/internal/utils/security"
"github.com/glebarez/sqlite"
"gorm.io/driver/postgres"
"gorm.io/gorm"
"gorm.io/gorm/schema"
)
var DB *gorm.DB
type dbModel struct {
value any
tableName string
hasIDPK bool
}
func registeredModels() []any {
return []any{
&User{},
&AuthSource{},
&ExternalAccount{},
&Option{},
&Origin{},
&ProxyRoute{},
&PagesProject{},
&PagesDeployment{},
&PagesDeploymentFile{},
&ConfigVersion{},
&Node{},
&NodeSystemProfile{},
&ApplyLog{},
&NodeMetricSnapshot{},
&NodeRequestReport{},
&NodeAccessLog{},
&NodeHealthEvent{},
&NodeObservationOpenresty{},
&NodeObservationFrps{},
&NodeObservationFrpc{},
&TLSCertificate{},
&ManagedDomain{},
&AcmeAccount{},
&DnsAccount{},
&WAFRuleGroup{},
&WAFIPGroup{},
&WAFRuleGroupBinding{},
}
}
func currentSchemaMetadataModels() []any {
return nil
}
func legacySchemaMetadataModels() []any {
return []any{
&DatabaseSchemaVersion{},
}
}
func schemaMetadataModels() []any {
models := make([]any, 0, len(currentSchemaMetadataModels())+len(legacySchemaMetadataModels()))
models = append(models, currentSchemaMetadataModels()...)
models = append(models, legacySchemaMetadataModels()...)
return models
}
func buildDBModels() ([]dbModel, error) {
models := registeredModels()
result := make([]dbModel, 0, len(models))
namer := schema.NamingStrategy{}
cache := &sync.Map{}
for _, item := range models {
parsed, err := schema.Parse(item, cache, namer)
if err != nil {
return nil, err
}
hasIDPK := len(parsed.PrimaryFields) == 1 && parsed.PrimaryFields[0].DBName == "id"
result = append(result, dbModel{
value: item,
tableName: parsed.Table,
hasIDPK: hasIDPK,
})
}
return result, nil
}
func createRootAccountIfNeed() error {
var user User
//if user.Status != common.UserStatusEnabled {
if err := DB.First(&user).Error; err != nil {
slog.Info("no user exists, create a root user", "username", "root")
hashedPassword, err := security.Password2Hash("123456")
if err != nil {
return err
}
rootUser := User{
Username: "root",
Password: hashedPassword,
Role: common.RoleRootUser,
Status: common.UserStatusEnabled,
DisplayName: "Root User",
}
DB.Create(&rootUser)
}
return nil
}
func CountTable(tableName string) (num int64) {
DB.Table(tableName).Count(&num)
return
}
func openDatabase() (*gorm.DB, string, error) {
if common.SQLDSN != "" {
db, err := gorm.Open(postgres.Open(common.SQLDSN), &gorm.Config{})
if err != nil {
return nil, "", err
}
return db, "postgres", nil
}
db, err := gorm.Open(sqlite.Open(common.SQLitePath), &gorm.Config{})
if err != nil {
return nil, "", err
}
slog.Info("database DSN not set, using SQLite as database", "sqlite_path", common.SQLitePath)
return db, "sqlite", nil
}
func autoMigrateAll(db *gorm.DB) error {
return autoMigrateAllExcept(db, nil)
}
func autoMigrateAllExcept(db *gorm.DB, excludedTables map[string]bool) error {
models := registeredModels()
for i, item := range models {
name := fmt.Sprintf("%T", item)
tableName, err := tableNameForModel(item)
if err != nil {
return fmt.Errorf("resolve table name for %s failed: %w", name, err)
}
if excludedTables[tableName] {
slog.Info("autoMigrateAll: skipped model", "index", fmt.Sprintf("%d/%d", i+1, len(models)), "model", name, "table", tableName)
continue
}
slog.Info("autoMigrateAll: migrating model", "index", fmt.Sprintf("%d/%d", i+1, len(models)), "model", name)
if err := db.AutoMigrate(item); err != nil {
return fmt.Errorf("AutoMigrate %s failed: %w", name, err)
}
slog.Info("autoMigrateAll: migrated model", "model", name)
}
return nil
}
func tableNameForModel(item any) (string, error) {
namer := schema.NamingStrategy{}
cache := &sync.Map{}
parsed, err := schema.Parse(item, cache, namer)
if err != nil {
return "", err
}
return parsed.Table, nil
}
func isDatabaseEmpty(db *gorm.DB) (bool, error) {
models, err := buildDBModels()
if err != nil {
return false, err
}
for _, item := range models {
if isShardedObservabilityTable(item.tableName) {
for _, table := range observabilityShardTables(item.tableName) {
if !db.Migrator().HasTable(table) {
continue
}
var count int64
if err := db.Table(table).Limit(1).Count(&count).Error; err != nil {
return false, err
}
if count > 0 {
return false, nil
}
}
continue
}
if !db.Migrator().HasTable(item.value) {
continue
}
var count int64
if err := db.Model(item.value).Limit(1).Count(&count).Error; err != nil {
return false, err
}
if count > 0 {
return false, nil
}
}
return true, nil
}
func sqliteSourceExists() bool {
info, err := os.Stat(common.SQLitePath)
if err != nil {
return false
}
return !info.IsDir()
}
func migrateSQLiteDataIfNeeded(target *gorm.DB, backend string) error {
if backend != "postgres" {
return nil
}
empty, err := isDatabaseEmpty(target)
if err != nil {
return err
}
if !empty {
slog.Info("skip sqlite migration because target database already has data", "backend", backend)
return nil
}
if !sqliteSourceExists() {
slog.Info("skip sqlite migration because sqlite source file was not found", "sqlite_path", common.SQLitePath)
return nil
}
source, err := gorm.Open(sqlite.Open(common.SQLitePath), &gorm.Config{
PrepareStmt: true,
})
if err != nil {
return fmt.Errorf("open sqlite source database failed: %w", err)
}
sourceSQLDB, err := source.DB()
if err != nil {
return fmt.Errorf("get sqlite source database handle failed: %w", err)
}
defer func() {
_ = sourceSQLDB.Close()
}()
models, err := buildDBModels()
if err != nil {
return err
}
slog.Info("starting sqlite to postgres database migration", "sqlite_path", common.SQLitePath)
err = target.Transaction(func(tx *gorm.DB) error {
for _, item := range models {
if err := migrateTableData(source, tx, item); err != nil {
return err
}
if item.hasIDPK {
if err := resetPostgresSequence(tx, item.tableName); err != nil {
return err
}
}
}
return nil
})
if err != nil {
return err
}
slog.Info("sqlite to postgres database migration completed", "sqlite_path", common.SQLitePath)
return nil
}
func migrateTableData(source *gorm.DB, target *gorm.DB, item dbModel) error {
if !source.Migrator().HasTable(item.value) {
slog.Info("database migration progress", "table", item.tableName, "migrated", 0, "total", 0, "status", "skipped_missing_source_table")
return nil
}
var total int64
if err := source.Model(item.value).Count(&total).Error; err != nil {
return fmt.Errorf("count sqlite table %s failed: %w", item.tableName, err)
}
slog.Info("database migration progress", "table", item.tableName, "migrated", 0, "total", total, "status", "starting")
if total == 0 {
slog.Info("database migration progress", "table", item.tableName, "migrated", 0, "total", total, "status", "completed")
return nil
}
modelType := reflect.TypeOf(item.value).Elem()
sliceType := reflect.SliceOf(modelType)
migrated := int64(0)
offset := 0
const batchSize = 200
for {
batchPtr := reflect.New(sliceType)
query := source.Model(item.value).Limit(batchSize).Offset(offset)
if item.hasIDPK {
query = query.Order("id ASC")
}
if err := query.Find(batchPtr.Interface()).Error; err != nil {
return fmt.Errorf("read sqlite table %s failed: %w", item.tableName, err)
}
batchLen := batchPtr.Elem().Len()
if batchLen == 0 {
break
}
if isShardedObservabilityTable(item.tableName) {
for index := 0; index < batchLen; index++ {
record := batchPtr.Elem().Index(index)
if err := target.Create(record.Addr().Interface()).Error; err != nil {
return fmt.Errorf("write target sharded table %s failed: %w", item.tableName, err)
}
}
} else {
if err := target.Create(batchPtr.Interface()).Error; err != nil {
return fmt.Errorf("write target table %s failed: %w", item.tableName, err)
}
}
migrated += int64(batchLen)
offset += batchLen
slog.Info("database migration progress", "table", item.tableName, "migrated", migrated, "total", total, "status", "running")
}
slog.Info("database migration progress", "table", item.tableName, "migrated", migrated, "total", total, "status", "completed")
return nil
}
func resetPostgresSequence(db *gorm.DB, tableName string) error {
sql := fmt.Sprintf(
"SELECT setval(pg_get_serial_sequence('%s', 'id'), COALESCE(MAX(id), 1), MAX(id) IS NOT NULL) FROM \"%s\"",
tableName,
tableName,
)
return db.Exec(sql).Error
}
func InitBenchmarkDB(dsn string) error {
db, err := gorm.Open(postgres.Open(dsn), &gorm.Config{})
if err != nil {
return err
}
sqlDB, err := db.DB()
if err != nil {
return err
}
sqlDB.SetMaxOpenConns(20)
sqlDB.SetMaxIdleConns(10)
DB = db
if err = registerSharding(db, "postgres"); err != nil {
return err
}
return ensureDatabaseSchemaUpToDate(db, "postgres")
}
func InitDB() (err error) {
db, backend, err := openDatabase()
if err != nil {
slog.Error("open database failed", "error", err)
os.Exit(1)
}
DB = db
if err = registerSharding(db, backend); err != nil {
return err
}
if err = ensureDatabaseSchemaUpToDate(db, backend); err != nil {
return err
}
return createRootAccountIfNeed()
}
func CloseDB() error {
sqlDB, err := DB.DB()
if err != nil {
return err
}
err = sqlDB.Close()
return err
}
func IsUniqueConstraintError(err error) bool {
if err == nil {
return false
}
return strings.Contains(strings.ToLower(err.Error()), "unique")
}
@@ -1,886 +0,0 @@
package model
import (
"encoding/json"
"go/ast"
"go/parser"
"go/token"
"os"
"path/filepath"
"reflect"
"strings"
"testing"
"time"
"github.com/glebarez/sqlite"
"gorm.io/gorm"
)
type legacyProxyRouteV7 struct {
ID uint `gorm:"primaryKey"`
SiteName string `gorm:"size:255;not null;default:''"`
Domain string `gorm:"uniqueIndex;size:255;not null"`
Domains string `gorm:"type:text;not null;default:'[]'"`
OriginID *uint `gorm:"index"`
OriginURL string `gorm:"size:2048;not null"`
OriginHost string `gorm:"size:255"`
Upstreams string `gorm:"type:text;not null;default:'[]'"`
Enabled bool `gorm:"not null;default:true"`
EnableHTTPS bool `gorm:"column:enable_https;not null;default:false"`
CertID *uint
CertIDs string `gorm:"type:text;not null;default:'[]'"`
RedirectHTTP bool `gorm:"not null;default:false"`
LimitConnPerServer int `gorm:"not null;default:0"`
LimitConnPerIP int `gorm:"not null;default:0"`
LimitRate string `gorm:"size:32;not null;default:''"`
CacheEnabled bool `gorm:"not null;default:false"`
CachePolicy string `gorm:"size:32;not null;default:''"`
CacheRules string `gorm:"type:text;not null;default:'[]'"`
CustomHeaders string `gorm:"type:text;not null;default:'[]'"`
Remark string `gorm:"size:255"`
CreatedAt time.Time
UpdatedAt time.Time
}
func (legacyProxyRouteV7) TableName() string {
return "proxy_routes"
}
func openBareTestSQLiteDB(t *testing.T, name string) *gorm.DB {
t.Helper()
db, err := gorm.Open(sqlite.Open(filepath.Join(t.TempDir(), name)), &gorm.Config{})
if err != nil {
t.Fatalf("open sqlite db: %v", err)
}
sqlDB, err := db.DB()
if err != nil {
t.Fatalf("get sql db: %v", err)
}
t.Cleanup(func() {
_ = sqlDB.Close()
})
return db
}
func openTestSQLiteDB(t *testing.T, name string) *gorm.DB {
t.Helper()
db := openBareTestSQLiteDB(t, name)
if err := autoMigrateAll(db); err != nil {
t.Fatalf("auto migrate db: %v", err)
}
return db
}
func findDBModelByTableName(t *testing.T, tableName string) dbModel {
t.Helper()
models, err := buildDBModels()
if err != nil {
t.Fatalf("build db models: %v", err)
}
for _, item := range models {
if item.tableName == tableName {
return item
}
}
t.Fatalf("db model not found for table %s", tableName)
return dbModel{}
}
func expectedCurrentDatabaseVersion() int {
return int(currentGooseTargetVersion())
}
func TestIsDatabaseEmpty(t *testing.T) {
db := openTestSQLiteDB(t, "empty.db")
empty, err := isDatabaseEmpty(db)
if err != nil {
t.Fatalf("isDatabaseEmpty returned error: %v", err)
}
if !empty {
t.Fatal("expected database to be empty")
}
if err := db.Create(&User{
Username: "alice",
Password: "secret",
DisplayName: "Alice",
Role: 1,
Status: 1,
}).Error; err != nil {
t.Fatalf("seed user: %v", err)
}
empty, err = isDatabaseEmpty(db)
if err != nil {
t.Fatalf("isDatabaseEmpty after seed returned error: %v", err)
}
if empty {
t.Fatal("expected database to be non-empty")
}
}
func TestMigrateTableDataCopiesRows(t *testing.T) {
source := openTestSQLiteDB(t, "source.db")
target := openTestSQLiteDB(t, "target.db")
user := User{
Id: 1,
Username: "root",
Password: "hashed",
DisplayName: "Root User",
Role: 100,
Status: 1,
}
option := Option{
Key: "AgentHeartbeatInterval",
Value: "10000",
}
if err := source.Create(&user).Error; err != nil {
t.Fatalf("seed source user: %v", err)
}
if err := source.Create(&option).Error; err != nil {
t.Fatalf("seed source option: %v", err)
}
if err := migrateTableData(source, target, findDBModelByTableName(t, "users")); err != nil {
t.Fatalf("migrate users: %v", err)
}
if err := migrateTableData(source, target, findDBModelByTableName(t, "options")); err != nil {
t.Fatalf("migrate options: %v", err)
}
var gotUser User
if err := target.First(&gotUser, 1).Error; err != nil {
t.Fatalf("query migrated user: %v", err)
}
if gotUser.Username != user.Username || gotUser.DisplayName != user.DisplayName {
t.Fatalf("unexpected migrated user: %+v", gotUser)
}
var gotOption Option
if err := target.First(&gotOption, "key = ?", option.Key).Error; err != nil {
t.Fatalf("query migrated option: %v", err)
}
if gotOption.Value != option.Value {
t.Fatalf("unexpected migrated option value: %s", gotOption.Value)
}
}
func TestRegisterShardingAutoMigratesShardTables(t *testing.T) {
db := openBareTestSQLiteDB(t, "sharded.db")
if err := registerSharding(db, "sqlite"); err != nil {
t.Fatalf("register sharding: %v", err)
}
if err := autoMigrateAll(db); err != nil {
t.Fatalf("auto migrate db: %v", err)
}
for _, table := range []string{
"node_metric_snapshots_00",
"node_metric_snapshots_09",
"node_request_reports_00",
"node_request_reports_09",
"node_access_logs_00",
"node_access_logs_09",
} {
if !db.Migrator().HasTable(table) {
t.Fatalf("expected sharded table %s to exist", table)
}
}
}
func TestUpgradeDatabaseSchemaV15ToV16AppliesCompressedReleaseSchema(t *testing.T) {
db := openBareTestSQLiteDB(t, "v16.db")
if err := registerSharding(db, "sqlite"); err != nil {
t.Fatalf("register sharding: %v", err)
}
if err := autoMigrateLegacySchemaMetadata(db); err != nil {
t.Fatalf("auto migrate legacy schema metadata: %v", err)
}
if err := applyCurrentSchema(db, "sqlite"); err != nil {
t.Fatalf("apply current schema: %v", err)
}
if err := db.Exec("ALTER TABLE proxy_routes ADD COLUMN pow_enabled BOOLEAN NOT NULL DEFAULT 0").Error; err != nil {
t.Fatalf("failed to add legacy pow_enabled: %v", err)
}
if err := db.Exec("ALTER TABLE proxy_routes ADD COLUMN pow_config TEXT NOT NULL DEFAULT '{}'").Error; err != nil {
t.Fatalf("failed to add legacy pow_config: %v", err)
}
if err := ensureDefaultWAFRuleGroup(db); err != nil {
t.Fatalf("ensure default waf rule group: %v", err)
}
if err := saveDatabaseSchemaVersion(db, 15); err != nil {
t.Fatalf("save schema version: %v", err)
}
if err := upgradeDatabaseSchema(db, "sqlite", 15); err != nil {
t.Fatalf("upgrade schema: %v", err)
}
if !db.Migrator().HasTable(&WAFIPGroup{}) {
t.Fatal("expected waf_ip_groups table")
}
if !db.Migrator().HasColumn(&WAFRuleGroup{}, "ip_whitelist_groups") {
t.Fatal("expected waf_rule_groups.ip_whitelist_groups column")
}
if !db.Migrator().HasColumn(&Node{}, "access_token") {
t.Fatal("expected nodes.access_token column")
}
if !db.Migrator().HasColumn(&Node{}, "version") {
t.Fatal("expected nodes.version column")
}
if !db.Migrator().HasColumn(&Node{}, "ext_version") {
t.Fatal("expected nodes.ext_version column")
}
if !db.Migrator().HasColumn(&ProxyRoute{}, "tunnel_node_id") {
t.Fatal("expected proxy_routes.tunnel_node_id column")
}
if db.Migrator().HasTable("tunnels") {
t.Fatal("expected pre-release tunnels table to be absent")
}
version, ok, err := loadDatabaseSchemaVersion(db)
if err != nil {
t.Fatalf("load schema version: %v", err)
}
if !ok || version != currentDatabaseSchemaVersion {
t.Fatalf("unexpected schema version: got %d ok=%v want %d", version, ok, currentDatabaseSchemaVersion)
}
}
func TestMigrateObservabilityLegacyColumnsBackfillsHealthEventMetadata(t *testing.T) {
db := openTestSQLiteDB(t, "legacy-health-events.db")
if err := db.Exec("ALTER TABLE node_health_events ADD COLUMN raw_json TEXT").Error; err != nil {
t.Fatalf("add raw_json column: %v", err)
}
rawJSON, err := json.Marshal(map[string]any{
"event_type": "sync_error",
"metadata": map[string]string{
"reason": "checksum_mismatch",
"scope": "routes",
},
})
if err != nil {
t.Fatalf("marshal raw json: %v", err)
}
event := &NodeHealthEvent{
NodeID: "node-legacy",
EventType: "sync_error",
Severity: "warning",
Status: "active",
Message: "checksum mismatch",
FirstTriggeredAt: time.Now().Add(-time.Minute),
LastTriggeredAt: time.Now(),
ReportedAt: time.Now(),
}
if err := db.Create(event).Error; err != nil {
t.Fatalf("create health event: %v", err)
}
if err := db.Exec("UPDATE node_health_events SET raw_json = ? WHERE id = ?", string(rawJSON), event.ID).Error; err != nil {
t.Fatalf("seed legacy raw_json: %v", err)
}
if err := migrateObservabilityLegacyColumns(db); err != nil {
t.Fatalf("migrateObservabilityLegacyColumns: %v", err)
}
var got NodeHealthEvent
if err := db.First(&got, event.ID).Error; err != nil {
t.Fatalf("query health event: %v", err)
}
if got.MetadataJSON == "" {
t.Fatal("expected metadata_json to be backfilled")
}
}
func TestEnsureDatabaseSchemaUpToDateInitializesFreshDatabase(t *testing.T) {
db := openBareTestSQLiteDB(t, "fresh-schema.db")
if err := registerSharding(db, "sqlite"); err != nil {
t.Fatalf("register sharding: %v", err)
}
if err := ensureDatabaseSchemaUpToDate(db, "sqlite"); err != nil {
t.Fatalf("ensureDatabaseSchemaUpToDate: %v", err)
}
version, exists, err := loadDatabaseSchemaVersion(db)
if err != nil {
t.Fatalf("loadDatabaseSchemaVersion: %v", err)
}
if !exists {
t.Fatal("expected database schema version to be recorded")
}
if version != expectedCurrentDatabaseVersion() {
t.Fatalf("unexpected schema version: got %d want %d", version, expectedCurrentDatabaseVersion())
}
if db.Migrator().HasTable(&DatabaseSchemaVersion{}) {
t.Fatal("expected fresh database to avoid legacy database_schema_versions table")
}
if !db.Migrator().HasTable("goose_db_version") {
t.Fatal("expected fresh database to initialize goose_db_version")
}
if !db.Migrator().HasColumn(&Node{}, "capabilities_json") {
t.Fatal("expected fresh database to apply goose migration nodes.capabilities_json")
}
}
func TestEnsureDatabaseSchemaUpToDateUpgradesLegacyDatabase(t *testing.T) {
db := openBareTestSQLiteDB(t, "legacy-schema.db")
if err := registerSharding(db, "sqlite"); err != nil {
t.Fatalf("register sharding: %v", err)
}
if err := autoMigrateAll(db); err != nil {
t.Fatalf("auto migrate db: %v", err)
}
// Add legacy PoW columns manually to proxy_routes table to simulate legacy schema v9-v17 state
if err := db.Exec("ALTER TABLE proxy_routes ADD COLUMN pow_enabled BOOLEAN NOT NULL DEFAULT 0").Error; err != nil {
t.Fatalf("failed to add legacy pow_enabled: %v", err)
}
if err := db.Exec("ALTER TABLE proxy_routes ADD COLUMN pow_config TEXT NOT NULL DEFAULT '{}'").Error; err != nil {
t.Fatalf("failed to add legacy pow_config: %v", err)
}
if err := db.Create(&User{
Username: "legacy",
Password: "secret",
DisplayName: "Legacy User",
Role: 1,
Status: 1,
}).Error; err != nil {
t.Fatalf("seed legacy user: %v", err)
}
if err := ensureDatabaseSchemaUpToDate(db, "sqlite"); err != nil {
t.Fatalf("ensureDatabaseSchemaUpToDate: %v", err)
}
version, exists, err := loadDatabaseSchemaVersion(db)
if err != nil {
t.Fatalf("loadDatabaseSchemaVersion: %v", err)
}
if !exists {
t.Fatal("expected legacy database to gain a schema version record")
}
if version != expectedCurrentDatabaseVersion() {
t.Fatalf("unexpected schema version: got %d want %d", version, expectedCurrentDatabaseVersion())
}
if db.Migrator().HasTable(&DatabaseSchemaVersion{}) {
t.Fatal("expected legacy database_schema_versions table to be removed after bridging to goose")
}
if !db.Migrator().HasTable("goose_db_version") {
t.Fatal("expected legacy upgrade to initialize goose_db_version")
}
if !db.Migrator().HasColumn(&Node{}, "capabilities_json") {
t.Fatal("expected legacy upgrade to apply goose migration nodes.capabilities_json")
}
}
func TestMigrateOriginsSchemaBackfillsOrigins(t *testing.T) {
db := openBareTestSQLiteDB(t, "legacy-origins.db")
if err := registerSharding(db, "sqlite"); err != nil {
t.Fatalf("register sharding: %v", err)
}
if err := applyCurrentSchema(db, "sqlite"); err != nil {
t.Fatalf("applyCurrentSchema: %v", err)
}
now := time.Now().UTC()
route := &ProxyRoute{
Domain: "app.example.com",
OriginURL: "https://origin-a.internal:8443/api",
Upstreams: `["https://origin-a.internal:8443/api"]`,
Enabled: true,
CreatedAt: now,
UpdatedAt: now,
}
if err := db.Create(route).Error; err != nil {
t.Fatalf("seed proxy route: %v", err)
}
if err := db.Exec(`DELETE FROM origins`).Error; err != nil {
t.Fatalf("clear origins: %v", err)
}
if err := db.Model(&ProxyRoute{}).Where("id = ?", route.ID).Update("origin_id", nil).Error; err != nil {
t.Fatalf("clear route origin_id: %v", err)
}
if err := backfillOriginsFromProxyRoutes(db); err != nil {
t.Fatalf("backfillOriginsFromProxyRoutes: %v", err)
}
if !db.Migrator().HasTable(&Origin{}) {
t.Fatal("expected origins table to exist")
}
if !db.Migrator().HasColumn(&ProxyRoute{}, "origin_id") {
t.Fatal("expected proxy_routes.origin_id column to exist")
}
reloadedRoute := &ProxyRoute{}
if err := db.First(reloadedRoute, route.ID).Error; err != nil {
t.Fatalf("query proxy route: %v", err)
}
if reloadedRoute.OriginID == nil || *reloadedRoute.OriginID == 0 {
t.Fatal("expected migrated route to be linked to a backfilled origin")
}
origin := &Origin{}
if err := db.First(origin, *reloadedRoute.OriginID).Error; err != nil {
t.Fatalf("query origin: %v", err)
}
if origin.Address != "origin-a.internal" {
t.Fatalf("unexpected backfilled origin address: %s", origin.Address)
}
}
func TestEnsureDatabaseSchemaUpToDateAddsProxyRouteDomainCertificateFields(t *testing.T) {
db := openBareTestSQLiteDB(t, "legacy-proxy-route-domain-cert-ids.db")
if err := registerSharding(db, "sqlite"); err != nil {
t.Fatalf("register sharding: %v", err)
}
if err := autoMigrateLegacySchemaMetadata(db); err != nil {
t.Fatalf("auto migrate legacy schema metadata: %v", err)
}
for _, item := range registeredModels() {
if _, ok := item.(*ProxyRoute); ok {
continue
}
if err := db.AutoMigrate(item); err != nil {
t.Fatalf("auto migrate supporting table: %v", err)
}
}
if err := db.AutoMigrate(&legacyProxyRouteV7{}); err != nil {
t.Fatalf("auto migrate legacy proxy_routes v7: %v", err)
}
// Add legacy PoW columns manually to proxy_routes table to simulate legacy schema v9-v17 state
if err := db.Exec("ALTER TABLE proxy_routes ADD COLUMN pow_enabled BOOLEAN NOT NULL DEFAULT 0").Error; err != nil {
t.Fatalf("failed to add legacy pow_enabled: %v", err)
}
if err := db.Exec("ALTER TABLE proxy_routes ADD COLUMN pow_config TEXT NOT NULL DEFAULT '{}'").Error; err != nil {
t.Fatalf("failed to add legacy pow_config: %v", err)
}
now := time.Now().UTC()
certID := uint(9)
if err := db.Create(&legacyProxyRouteV7{
SiteName: "secure-site",
Domain: "secure.example.com",
Domains: `["secure.example.com","www.secure.example.com"]`,
OriginURL: "https://origin-secure.internal:8443",
Upstreams: `["https://origin-secure.internal:8443"]`,
Enabled: true,
EnableHTTPS: true,
CertID: &certID,
CertIDs: `[9]`,
RedirectHTTP: true,
LimitConnPerServer: 120,
LimitConnPerIP: 12,
LimitRate: "512k",
CacheEnabled: false,
CachePolicy: "",
CacheRules: `[]`,
CustomHeaders: `[]`,
CreatedAt: now,
UpdatedAt: now,
}).Error; err != nil {
t.Fatalf("seed legacy proxy route v7: %v", err)
}
if err := saveDatabaseSchemaVersion(db, 7); err != nil {
t.Fatalf("save schema version: %v", err)
}
previousDB := DB
DB = db
t.Cleanup(func() {
DB = previousDB
})
if err := ensureDatabaseSchemaUpToDate(db, "sqlite"); err != nil {
t.Fatalf("ensureDatabaseSchemaUpToDate: %v", err)
}
var route ProxyRoute
if err := db.First(&route).Error; err != nil {
t.Fatalf("query migrated proxy route: %v", err)
}
var domainCertIDs []uint
if err := json.Unmarshal([]byte(route.DomainCertIDs), &domainCertIDs); err != nil {
t.Fatalf("decode migrated domain_cert_ids: %v", err)
}
if len(domainCertIDs) != 2 || domainCertIDs[0] != certID || domainCertIDs[1] != certID {
t.Fatalf("unexpected migrated domain_cert_ids: %#v", domainCertIDs)
}
}
func TestRunDatabaseSchemaMigrationDoesNotAdvanceVersionWhenValidationFails(t *testing.T) {
db := openBareTestSQLiteDB(t, "failed-validation.db")
err := runDatabaseSchemaMigration(db, "sqlite", databaseSchemaMigration{
fromVersion: legacyDatabaseSchemaVersion,
toVersion: 11,
migrate: func(tx *gorm.DB, backend string) error {
return autoMigrateLegacySchemaMetadata(tx)
},
validate: func(tx *gorm.DB, backend string) error {
return gorm.ErrInvalidDB
},
})
if err == nil {
t.Fatal("expected migration validation to fail")
}
_, exists, loadErr := loadDatabaseSchemaVersion(db)
if loadErr != nil {
t.Fatalf("loadDatabaseSchemaVersion: %v", loadErr)
}
if exists {
t.Fatal("expected schema version to remain unset after failed validation")
}
}
func TestEnsureDatabaseSchemaUpToDateAddsNodeIPManualOverride(t *testing.T) {
db := openBareTestSQLiteDB(t, "node-ip-manual-override-migration.db")
if err := registerSharding(db, "sqlite"); err != nil {
t.Fatalf("register sharding: %v", err)
}
if err := applyCurrentSchema(db, "sqlite"); err != nil {
t.Fatalf("apply current schema: %v", err)
}
if err := db.Exec("ALTER TABLE proxy_routes ADD COLUMN pow_enabled BOOLEAN NOT NULL DEFAULT 0").Error; err != nil {
t.Fatalf("failed to add legacy pow_enabled: %v", err)
}
if err := db.Exec("ALTER TABLE proxy_routes ADD COLUMN pow_config TEXT NOT NULL DEFAULT '{}'").Error; err != nil {
t.Fatalf("failed to add legacy pow_config: %v", err)
}
if err := ensureDefaultWAFRuleGroup(db); err != nil {
t.Fatalf("ensure default waf rule group: %v", err)
}
if err := db.Migrator().DropColumn(&Node{}, "ip_manual_override"); err != nil {
t.Fatalf("drop ip_manual_override column: %v", err)
}
if db.Migrator().HasColumn(&Node{}, "ip_manual_override") {
t.Fatal("expected test database to simulate schema v14 without ip_manual_override")
}
if err := saveDatabaseSchemaVersion(db, 14); err != nil {
t.Fatalf("save schema version: %v", err)
}
if err := ensureDatabaseSchemaUpToDate(db, "sqlite"); err != nil {
t.Fatalf("ensureDatabaseSchemaUpToDate: %v", err)
}
if !db.Migrator().HasColumn(&Node{}, "ip_manual_override") {
t.Fatal("expected migration to add nodes.ip_manual_override")
}
version, exists, err := loadDatabaseSchemaVersion(db)
if err != nil {
t.Fatalf("loadDatabaseSchemaVersion: %v", err)
}
if !exists {
t.Fatal("expected schema version record to exist")
}
if version != expectedCurrentDatabaseVersion() {
t.Fatalf("unexpected schema version: got %d want %d", version, expectedCurrentDatabaseVersion())
}
if !db.Migrator().HasColumn(&Node{}, "capabilities_json") {
t.Fatal("expected migration chain to include nodes.capabilities_json")
}
}
func TestEnsureDatabaseSchemaUpToDateV16BackfillsNodeColumnsWhenNewColumnsAlreadyExist(t *testing.T) {
db := openBareTestSQLiteDB(t, "node-v16-existing-target-columns.db")
if err := registerSharding(db, "sqlite"); err != nil {
t.Fatalf("register sharding: %v", err)
}
if err := applyCurrentSchema(db, "sqlite"); err != nil {
t.Fatalf("apply current schema: %v", err)
}
if err := db.Exec("ALTER TABLE proxy_routes ADD COLUMN pow_enabled BOOLEAN NOT NULL DEFAULT 0").Error; err != nil {
t.Fatalf("failed to add legacy pow_enabled: %v", err)
}
if err := db.Exec("ALTER TABLE proxy_routes ADD COLUMN pow_config TEXT NOT NULL DEFAULT '{}'").Error; err != nil {
t.Fatalf("failed to add legacy pow_config: %v", err)
}
if err := ensureDefaultWAFRuleGroup(db); err != nil {
t.Fatalf("ensure default waf rule group: %v", err)
}
for _, stmt := range []string{
`ALTER TABLE nodes ADD COLUMN agent_token text`,
`ALTER TABLE nodes ADD COLUMN agent_version text`,
`ALTER TABLE nodes ADD COLUMN nginx_version text`,
`ALTER TABLE nodes ADD COLUMN relay_version text`,
`ALTER TABLE nodes ADD COLUMN relay_frp_version text`,
`ALTER TABLE nodes ADD COLUMN relay_frps_connections integer`,
`ALTER TABLE nodes ADD COLUMN relay_frps_proxy_count integer`,
} {
if err := db.Exec(stmt).Error; err != nil {
t.Fatalf("prepare legacy node column with %q: %v", stmt, err)
}
}
now := time.Now()
if err := db.Exec(`
INSERT INTO nodes (
node_id, name, ip, access_token, version, ext_version,
agent_token, agent_version, nginx_version,
status, last_seen_at, created_at, updated_at
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
`, "node-v16", "Node v16", "127.0.0.1", "", "", "", "legacy-token", "v2.0.0", "openresty/1.25.3", "offline", now, now, now).Error; err != nil {
t.Fatalf("seed node with legacy columns: %v", err)
}
if err := saveDatabaseSchemaVersion(db, 15); err != nil {
t.Fatalf("save schema version: %v", err)
}
if err := ensureDatabaseSchemaUpToDate(db, "sqlite"); err != nil {
t.Fatalf("ensureDatabaseSchemaUpToDate: %v", err)
}
var node Node
if err := db.Where("node_id = ?", "node-v16").First(&node).Error; err != nil {
t.Fatalf("query migrated node: %v", err)
}
if node.AccessToken != "legacy-token" {
t.Fatalf("unexpected access_token: got %q", node.AccessToken)
}
if node.Version != "v2.0.0" {
t.Fatalf("unexpected version: got %q", node.Version)
}
if node.ExtVersion != "openresty/1.25.3" {
t.Fatalf("unexpected ext_version: got %q", node.ExtVersion)
}
for _, column := range []string{
"agent_token",
"agent_version",
"nginx_version",
"relay_version",
"relay_frp_version",
"relay_frps_connections",
"relay_frps_proxy_count",
} {
exists, err := databaseColumnExists(db, "nodes", column)
if err != nil {
t.Fatalf("inspect legacy nodes.%s: %v", column, err)
}
if exists {
t.Fatalf("expected migration to drop legacy nodes.%s column", column)
}
}
version, exists, err := loadDatabaseSchemaVersion(db)
if err != nil {
t.Fatalf("loadDatabaseSchemaVersion: %v", err)
}
if !exists {
t.Fatal("expected schema version record to exist")
}
if version != expectedCurrentDatabaseVersion() {
t.Fatalf("unexpected schema version: got %d want %d", version, expectedCurrentDatabaseVersion())
}
if !db.Migrator().HasColumn(&Node{}, "capabilities_json") {
t.Fatal("expected v16 upgrade path to apply goose migration nodes.capabilities_json")
}
}
func TestEnsureDatabaseSchemaUpToDateV16DropsLegacyNodeColumnsWhenAlreadyCurrent(t *testing.T) {
db := openBareTestSQLiteDB(t, "node-v16-current-legacy-columns.db")
if err := registerSharding(db, "sqlite"); err != nil {
t.Fatalf("register sharding: %v", err)
}
if err := applyCurrentSchema(db, "sqlite"); err != nil {
t.Fatalf("apply current schema: %v", err)
}
for _, stmt := range []string{
`ALTER TABLE nodes ADD COLUMN agent_token text`,
`ALTER TABLE nodes ADD COLUMN agent_version text`,
`ALTER TABLE nodes ADD COLUMN nginx_version text`,
} {
if err := db.Exec(stmt).Error; err != nil {
t.Fatalf("prepare legacy node column with %q: %v", stmt, err)
}
}
if err := saveDatabaseSchemaVersion(db, currentDatabaseSchemaVersion); err != nil {
t.Fatalf("save schema version: %v", err)
}
if err := ensureDatabaseSchemaUpToDate(db, "sqlite"); err != nil {
t.Fatalf("ensureDatabaseSchemaUpToDate: %v", err)
}
for _, column := range []string{"agent_token", "agent_version", "nginx_version"} {
exists, err := databaseColumnExists(db, "nodes", column)
if err != nil {
t.Fatalf("inspect legacy nodes.%s: %v", column, err)
}
if exists {
t.Fatalf("expected current-schema cleanup to drop legacy nodes.%s column", column)
}
}
if db.Migrator().HasTable(&DatabaseSchemaVersion{}) {
t.Fatal("expected current-schema legacy version table to be removed after goose bridge")
}
if !db.Migrator().HasTable("goose_db_version") {
t.Fatal("expected current-schema goose_db_version table to exist")
}
if !db.Migrator().HasColumn(&Node{}, "capabilities_json") {
t.Fatal("expected current-schema repair to preserve goose column nodes.capabilities_json")
}
}
func TestEnsureDatabaseSchemaUpToDateKeepsGooseOnlyDatabaseOnReentry(t *testing.T) {
db := openBareTestSQLiteDB(t, "goose-only-reentry.db")
if err := registerSharding(db, "sqlite"); err != nil {
t.Fatalf("register sharding: %v", err)
}
if err := ensureDatabaseSchemaUpToDate(db, "sqlite"); err != nil {
t.Fatalf("first ensureDatabaseSchemaUpToDate: %v", err)
}
if db.Migrator().HasTable(&DatabaseSchemaVersion{}) {
t.Fatal("expected first initialization to avoid legacy table")
}
if err := ensureDatabaseSchemaUpToDate(db, "sqlite"); err != nil {
t.Fatalf("second ensureDatabaseSchemaUpToDate: %v", err)
}
if db.Migrator().HasTable(&DatabaseSchemaVersion{}) {
t.Fatal("expected goose-only database to remain free of legacy version table")
}
version, exists, err := loadGooseDatabaseVersion(db)
if err != nil {
t.Fatalf("loadGooseDatabaseVersion: %v", err)
}
if !exists {
t.Fatal("expected goose-only database to keep goose version record")
}
if version != expectedCurrentDatabaseVersion() {
t.Fatalf("unexpected goose version: got %d want %d", version, expectedCurrentDatabaseVersion())
}
if !db.Migrator().HasColumn(&Node{}, "capabilities_json") {
t.Fatal("expected goose-only database to keep nodes.capabilities_json")
}
}
func TestAllRegisteredMigrationsHaveValidationDefined(t *testing.T) {
ctx := databaseSchemaMigrationContext{}
for _, migration := range databaseSchemaMigrations() {
err := ctx.ValidateDatabaseSchemaVersion(nil, "sqlite", migration.toVersion)
if err != nil && strings.Contains(err.Error(), "is not defined") {
t.Fatalf("Validation is not defined in migrations.go for registered migration version v%d: %v", migration.toVersion, err)
}
}
}
func TestAllGORMModelsAreRegistered(t *testing.T) {
// 1. Gather all registered model names
registeredNames := make(map[string]bool)
for _, item := range registeredModels() {
name := reflect.TypeOf(item).Elem().Name()
registeredNames[name] = true
}
for _, item := range schemaMetadataModels() {
name := reflect.TypeOf(item).Elem().Name()
registeredNames[name] = true
}
// 2. Parse all .go files in model/ package
fset := token.NewFileSet()
pkgs, err := parser.ParseDir(fset, ".", func(info os.FileInfo) bool {
// Only parse .go files, exclude _test.go files and subdirectories
return !info.IsDir() && strings.HasSuffix(info.Name(), ".go") && !strings.HasSuffix(info.Name(), "_test.go")
}, 0)
if err != nil {
t.Fatalf("failed to parse directory: %v", err)
}
for _, pkg := range pkgs {
for _, file := range pkg.Files {
for _, decl := range file.Decls {
genDecl, ok := decl.(*ast.GenDecl)
if !ok || genDecl.Tok != token.TYPE {
continue
}
for _, spec := range genDecl.Specs {
typeSpec, ok := spec.(*ast.TypeSpec)
if !ok {
continue
}
structType, ok := typeSpec.Type.(*ast.StructType)
if !ok {
continue
}
// Verify if this struct has any field with a `gorm:"..."` tag
isGORMModel := false
for _, field := range structType.Fields.List {
if field.Tag != nil && strings.Contains(field.Tag.Value, "gorm:") {
isGORMModel = true
break
}
}
if isGORMModel {
structName := typeSpec.Name.Name
if !registeredNames[structName] {
t.Errorf("Model struct %q is defined with GORM tags but is NOT registered in registeredModels() or schemaMetadataModels() in model/main.go!", structName)
}
}
}
}
}
}
}
func TestEnsureDatabaseSchemaUpToDateDropsPagesDeploymentUnusedFields(t *testing.T) {
db := openBareTestSQLiteDB(t, "drop-pages-deployment-unused-fields.db")
if err := registerSharding(db, "sqlite"); err != nil {
t.Fatalf("register sharding: %v", err)
}
if err := ensureDatabaseSchemaUpToDate(db, "sqlite"); err != nil {
t.Fatalf("first ensureDatabaseSchemaUpToDate: %v", err)
}
// Verify columns do not exist
if db.Migrator().HasColumn("pages_deployments", "root_dir") {
t.Fatal("expected root_dir column to be absent initially")
}
if db.Migrator().HasColumn("pages_deployments", "entry_file") {
t.Fatal("expected entry_file column to be absent initially")
}
// Manually add columns to simulate old state
if err := db.Exec("ALTER TABLE pages_deployments ADD COLUMN root_dir TEXT").Error; err != nil {
t.Fatalf("failed to add root_dir column: %v", err)
}
if err := db.Exec("ALTER TABLE pages_deployments ADD COLUMN entry_file TEXT").Error; err != nil {
t.Fatalf("failed to add entry_file column: %v", err)
}
// Verify columns were added
if !db.Migrator().HasColumn("pages_deployments", "root_dir") {
t.Fatal("expected root_dir column to be present after manual add")
}
if !db.Migrator().HasColumn("pages_deployments", "entry_file") {
t.Fatal("expected entry_file column to be present after manual add")
}
// Remove the migration record from goose_db_version table
const versionToRerun = 202606040004
if err := db.Exec("DELETE FROM goose_db_version WHERE version_id = ?", versionToRerun).Error; err != nil {
t.Fatalf("failed to delete migration record: %v", err)
}
// Run migration again
if err := ensureDatabaseSchemaUpToDate(db, "sqlite"); err != nil {
t.Fatalf("second ensureDatabaseSchemaUpToDate: %v", err)
}
// Verify columns were dropped successfully
if db.Migrator().HasColumn("pages_deployments", "root_dir") {
t.Fatal("expected root_dir column to be dropped after migration rerun")
}
if db.Migrator().HasColumn("pages_deployments", "entry_file") {
t.Fatal("expected entry_file column to be dropped after migration rerun")
}
}
@@ -1,41 +0,0 @@
package model
import "time"
type ManagedDomain struct {
ID uint `json:"id" gorm:"primaryKey"`
Domain string `json:"domain" gorm:"uniqueIndex;size:255;not null"`
CertID *uint `json:"cert_id"`
Enabled bool `json:"enabled" gorm:"not null;default:true"`
Remark string `json:"remark" gorm:"size:255"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
}
func ListManagedDomains() (domains []*ManagedDomain, err error) {
err = DB.Order("id desc").Find(&domains).Error
return domains, err
}
func ListEnabledManagedDomainsWithCertificate() (domains []*ManagedDomain, err error) {
err = DB.Where("enabled = ? AND cert_id IS NOT NULL", true).Order("id desc").Find(&domains).Error
return domains, err
}
func GetManagedDomainByID(id uint) (*ManagedDomain, error) {
domain := &ManagedDomain{}
err := DB.First(domain, id).Error
return domain, err
}
func (domain *ManagedDomain) Insert() error {
return DB.Create(domain).Error
}
func (domain *ManagedDomain) Update() error {
return DB.Save(domain).Error
}
func (domain *ManagedDomain) Delete() error {
return DB.Delete(domain).Error
}
@@ -1,4 +0,0 @@
package migrate
// Versions 1 through 7 are treated as the historical baseline. There are no
// supported deployments below v8, so new upgrades start from this base version.
@@ -1,54 +0,0 @@
package migrate
import (
"sort"
"gorm.io/gorm"
)
const BaseDatabaseSchemaVersion = 7
type Context interface {
ApplyCurrentSchema(db *gorm.DB, backend string) error
ApplyCurrentSchemaExcept(db *gorm.DB, backend string, excludedTables ...string) error
BackfillOriginsFromProxyRoutes(db *gorm.DB) error
BackfillProxyRouteSiteFields(db *gorm.DB) error
EnsureProxyRouteSiteNameUniqueIndex(db *gorm.DB) error
BackfillProxyRouteCertificateFields(db *gorm.DB) error
BackfillProxyRouteDomainCertificateFields(db *gorm.DB) error
EnsureDefaultGitHubAuthSource(db *gorm.DB) error
EnsureDefaultWAFRuleGroup(db *gorm.DB) error
DropLegacyNodeColumns(db *gorm.DB, backend string) error
ValidateDatabaseSchemaVersion(db *gorm.DB, backend string, version int) error
}
type Migration struct {
FromVersion int
ToVersion int
Migrate func(ctx Context, db *gorm.DB, backend string) error
Validate func(ctx Context, db *gorm.DB, backend string) error
}
var registeredMigrations []Migration
func Register(migration Migration) {
registeredMigrations = append(registeredMigrations, migration)
}
func Migrations() []Migration {
migrations := append([]Migration{}, registeredMigrations...)
sort.Slice(migrations, func(i int, j int) bool {
return migrations[i].FromVersion < migrations[j].FromVersion
})
return migrations
}
func CurrentVersion() int {
version := BaseDatabaseSchemaVersion
for _, migration := range registeredMigrations {
if migration.ToVersion > version {
version = migration.ToVersion
}
}
return version
}
@@ -1,23 +0,0 @@
package migrate
import "testing"
func TestMigrationsAreContinuousFromBaseVersion(t *testing.T) {
migrations := Migrations()
if len(migrations) == 0 {
t.Fatal("expected at least one registered migration")
}
expectedFrom := BaseDatabaseSchemaVersion
for _, migration := range migrations {
if migration.FromVersion != expectedFrom {
t.Fatalf("expected migration from v%d, got v%d -> v%d", expectedFrom, migration.FromVersion, migration.ToVersion)
}
if migration.ToVersion != migration.FromVersion+1 {
t.Fatalf("expected one-step migration, got v%d -> v%d", migration.FromVersion, migration.ToVersion)
}
expectedFrom = migration.ToVersion
}
if CurrentVersion() != expectedFrom {
t.Fatalf("unexpected current version: got %d want %d", CurrentVersion(), expectedFrom)
}
}
@@ -1,26 +0,0 @@
// v10 升级内容:新增可配置认证源与第三方账号绑定,并迁移旧 GitHub 登录配置。
// 背景说明:登录体系从固定 GitHub OAuth 字段演进为通用认证源模型,需要创建 auth_sources、external_accounts,并把旧用户 GitHub 绑定迁移到新表。
package migrate
import "gorm.io/gorm"
func init() {
Register(V10())
}
func V10() Migration {
return Migration{
FromVersion: 9,
ToVersion: 10,
Migrate: migrateV10,
Validate: validateV10,
}
}
func migrateV10(ctx Context, db *gorm.DB, backend string) error {
return ctx.EnsureDefaultGitHubAuthSource(db)
}
func validateV10(ctx Context, db *gorm.DB, backend string) error {
return ctx.ValidateDatabaseSchemaVersion(db, backend, 10)
}
@@ -1,26 +0,0 @@
// v11 升级内容:新增 ACME 账户、DNS 账户,并扩展证书 provider 字段。
// 背景说明:证书申请能力从单一手工导入扩展到自动签发,需要持久化 ACME/DNS 凭据,并标记证书来源。
package migrate
import "gorm.io/gorm"
func init() {
Register(V11())
}
func V11() Migration {
return Migration{
FromVersion: 10,
ToVersion: 11,
Migrate: migrateV11,
Validate: validateV11,
}
}
func migrateV11(ctx Context, db *gorm.DB, backend string) error {
return nil
}
func validateV11(ctx Context, db *gorm.DB, backend string) error {
return ctx.ValidateDatabaseSchemaVersion(db, backend, 11)
}
@@ -1,26 +0,0 @@
// v12 升级内容:为 proxy_routes 增加 Basic Auth 相关字段。
// 背景说明:站点级访问控制需要支持基础认证,因此在代理路由配置中持久化 Basic Auth 开关与凭据配置。
package migrate
import "gorm.io/gorm"
func init() {
Register(V12())
}
func V12() Migration {
return Migration{
FromVersion: 11,
ToVersion: 12,
Migrate: migrateV12,
Validate: validateV12,
}
}
func migrateV12(ctx Context, db *gorm.DB, backend string) error {
return nil
}
func validateV12(ctx Context, db *gorm.DB, backend string) error {
return ctx.ValidateDatabaseSchemaVersion(db, backend, 12)
}
@@ -1,26 +0,0 @@
// v13 升级内容:新增 WAF 规则组与站点绑定表,并创建默认全局规则组。
// 背景说明:WAF 配置从零散站点字段演进为可复用规则组,需要全局规则组作为默认入口,并支持站点与规则组绑定。
package migrate
import "gorm.io/gorm"
func init() {
Register(V13())
}
func V13() Migration {
return Migration{
FromVersion: 12,
ToVersion: 13,
Migrate: migrateV13,
Validate: validateV13,
}
}
func migrateV13(ctx Context, db *gorm.DB, backend string) error {
return ctx.EnsureDefaultWAFRuleGroup(db)
}
func validateV13(ctx Context, db *gorm.DB, backend string) error {
return ctx.ValidateDatabaseSchemaVersion(db, backend, 13)
}
@@ -1,26 +0,0 @@
// v14 升级内容:为 WAF 规则组增加 PoW 策略字段。
// 背景说明:PoW 能力从站点路由侧沉淀到 WAF 规则组中,便于统一按规则组管理人机挑战策略。
package migrate
import "gorm.io/gorm"
func init() {
Register(V14())
}
func V14() Migration {
return Migration{
FromVersion: 13,
ToVersion: 14,
Migrate: migrateV14,
Validate: validateV14,
}
}
func migrateV14(ctx Context, db *gorm.DB, backend string) error {
return ctx.EnsureDefaultWAFRuleGroup(db)
}
func validateV14(ctx Context, db *gorm.DB, backend string) error {
return ctx.ValidateDatabaseSchemaVersion(db, backend, 14)
}
@@ -1,52 +0,0 @@
// v15 升级内容:为 nodes 增加 ip_manual_override 字段。
// 背景说明:管理端手动指定节点 IP 后,Agent 心跳不应继续覆盖该值,因此需要在节点表中记录 IP 是否由管理端锁定。
package migrate
import (
"fmt"
"gorm.io/gorm"
)
type nodeV15 struct {
IPManualOverride bool `gorm:"column:ip_manual_override;not null;default:false"`
}
func init() {
Register(V15())
}
func V15() Migration {
return Migration{
FromVersion: 14,
ToVersion: 15,
Migrate: migrateV15,
Validate: validateV15,
}
}
func (nodeV15) TableName() string {
return "nodes"
}
func migrateV15(ctx Context, db *gorm.DB, backend string) error {
if db == nil {
return fmt.Errorf("database handle is nil")
}
if !db.Migrator().HasColumn(&nodeV15{}, "ip_manual_override") {
if err := db.Migrator().AddColumn(&nodeV15{}, "IPManualOverride"); err != nil {
return fmt.Errorf("add nodes.ip_manual_override: %w", err)
}
}
return nil
}
func validateV15(ctx Context, db *gorm.DB, backend string) error {
if err := ctx.ValidateDatabaseSchemaVersion(db, backend, 14); err != nil {
return err
}
if db == nil || !db.Migrator().HasColumn(&nodeV15{}, "ip_manual_override") {
return fmt.Errorf("column nodes.ip_manual_override is missing")
}
return nil
}
@@ -1,185 +0,0 @@
// v16 is the first database migration after the V15 formal release baseline.
// It folds the previously drafted v16-v21 schema work into a single official
// upgrade: tunnel-relay fields, WAF IP groups, current node identity/version
// columns, and split node observation tables. The migration also backfills
// legacy node columns and removes obsolete pre-release tunnel metadata when
// present, so V15 deployments can upgrade directly to the new formal schema.
package migrate
import (
"fmt"
"log/slog"
"gorm.io/gorm"
)
func (nodeV16) TableName() string {
return "nodes"
}
func (tunnelV16) TableName() string {
return "tunnels"
}
func (proxyRouteV16) TableName() string {
return "proxy_routes"
}
type nodeV16 struct{}
type tunnelV16 struct{}
type proxyRouteV16 struct{}
type wafIPGroupV16 struct{}
type wafRuleGroupV16 struct{}
func (wafIPGroupV16) TableName() string {
return "waf_ip_groups"
}
func (wafRuleGroupV16) TableName() string {
return "waf_rule_groups"
}
func init() {
Register(V16())
}
func V16() Migration {
return Migration{
FromVersion: 15,
ToVersion: 16,
Migrate: migrateV16,
Validate: validateV16,
}
}
func migrateV16(ctx Context, db *gorm.DB, backend string) error {
if err := ctx.ApplyCurrentSchema(db, backend); err != nil {
return err
}
migrator := db.Migrator()
if migrator.HasColumn(&nodeV16{}, "agent_token") {
if err := db.Exec(`UPDATE nodes SET access_token = agent_token WHERE access_token IS NULL OR access_token = ''`).Error; err != nil {
return fmt.Errorf("backfill nodes.access_token from agent_token: %w", err)
}
}
if migrator.HasColumn(&nodeV16{}, "agent_version") {
if err := db.Exec(`UPDATE nodes SET version = agent_version WHERE version = '' OR version IS NULL`).Error; err != nil {
return fmt.Errorf("backfill nodes.version from agent_version: %w", err)
}
}
if migrator.HasColumn(&nodeV16{}, "nginx_version") {
if err := db.Exec(`UPDATE nodes SET ext_version = nginx_version WHERE ext_version IS NULL OR ext_version = ''`).Error; err != nil {
return fmt.Errorf("backfill nodes.ext_version from nginx_version: %w", err)
}
}
if err := ctx.DropLegacyNodeColumns(db, backend); err != nil {
return err
}
if err := db.Exec("UPDATE nodes SET node_type = 'edge_node' WHERE node_type = '' OR node_type IS NULL").Error; err != nil {
return fmt.Errorf("backfill nodes.node_type: %w", err)
}
if err := db.Exec("UPDATE proxy_routes SET upstream_type = 'direct' WHERE upstream_type = '' OR upstream_type IS NULL").Error; err != nil {
return fmt.Errorf("backfill proxy_routes.upstream_type: %w", err)
}
if migrator.HasColumn(&proxyRouteV16{}, "tunnel_id") {
if err := db.Model(&proxyRouteV16{}).Where("upstream_type = ?", "tunnel").Update("upstream_type", "direct").Error; err != nil {
return fmt.Errorf("reset pre-release tunnel proxy routes: %w", err)
}
// Drop the legacy index idx_proxy_routes_tunnel_id if it exists, to avoid errors on dropping the tunnel_id column (especially on SQLite).
if migrator.HasIndex(&proxyRouteV16{}, "idx_proxy_routes_tunnel_id") {
if err := migrator.DropIndex(&proxyRouteV16{}, "idx_proxy_routes_tunnel_id"); err != nil {
return fmt.Errorf("drop index idx_proxy_routes_tunnel_id failed: %w", err)
}
}
if err := migrator.DropColumn(&proxyRouteV16{}, "tunnel_id"); err != nil {
return fmt.Errorf("drop pre-release proxy_routes.tunnel_id: %w", err)
}
}
if migrator.HasTable(&tunnelV16{}) {
if err := migrator.DropTable(&tunnelV16{}); err != nil {
return fmt.Errorf("drop pre-release tunnels table: %w", err)
}
slog.Info("dropped pre-release tunnels table during v16 migration")
}
return nil
}
func validateV16(ctx Context, db *gorm.DB, backend string) error {
if err := ctx.ValidateDatabaseSchemaVersion(db, backend, 15); err != nil {
return err
}
if db == nil {
return fmt.Errorf("database handle is nil")
}
migrator := db.Migrator()
for _, column := range []string{
"access_token",
"version",
"ext_version",
"node_type",
"relay_bind_port",
"relay_vhost_http_port",
"relay_auth_token",
"relay_agent_access_addr",
"relay_client_access_addr",
"relay_client_proxy_url",
"relay_status",
} {
if !migrator.HasColumn(&nodeV16{}, column) {
return fmt.Errorf("column nodes.%s is missing", column)
}
}
for _, column := range []string{
"upstream_type",
"tunnel_node_id",
"tunnel_target_addr",
"tunnel_target_protocol",
} {
if !migrator.HasColumn(&proxyRouteV16{}, column) {
return fmt.Errorf("column proxy_routes.%s is missing", column)
}
}
if migrator.HasColumn(&proxyRouteV16{}, "tunnel_id") {
return fmt.Errorf("column proxy_routes.tunnel_id should not exist in v16")
}
if migrator.HasTable(&tunnelV16{}) {
return fmt.Errorf("table tunnels should not exist in v16")
}
for _, column := range []string{
"agent_token",
"agent_version",
"nginx_version",
"relay_version",
"relay_frp_version",
"relay_frps_connections",
"relay_frps_proxy_count",
} {
if migrator.HasColumn(&nodeV16{}, column) {
return fmt.Errorf("column nodes.%s should not exist in v16", column)
}
}
if !migrator.HasTable(&wafIPGroupV16{}) {
return fmt.Errorf("table waf_ip_groups is missing")
}
for _, column := range []string{
"ip_whitelist_groups",
"ip_blacklist_groups",
} {
if !migrator.HasColumn(&wafRuleGroupV16{}, column) {
return fmt.Errorf("column waf_rule_groups.%s is missing", column)
}
}
if !migrator.HasColumn(&wafIPGroupV16{}, "ext_ips") {
return fmt.Errorf("column waf_ip_groups.ext_ips is missing")
}
return nil
}
@@ -1,58 +0,0 @@
package migrate
import (
"fmt"
"gorm.io/gorm"
)
type nodeV17 struct{}
func (nodeV17) TableName() string {
return "nodes"
}
func init() {
Register(V17())
}
func V17() Migration {
return Migration{
FromVersion: 16,
ToVersion: 17,
Migrate: migrateV17,
Validate: validateV17,
}
}
func migrateV17(ctx Context, db *gorm.DB, backend string) error {
if err := ctx.ApplyCurrentSchema(db, backend); err != nil {
return err
}
return nil
}
func validateV17(ctx Context, db *gorm.DB, backend string) error {
if err := ctx.ValidateDatabaseSchemaVersion(db, backend, 16); err != nil {
return err
}
if db == nil {
return fmt.Errorf("database handle is nil")
}
migrator := db.Migrator()
if !migrator.HasColumn(&nodeV17{}, "relay_web_server_enabled") {
return fmt.Errorf("column nodes.relay_web_server_enabled is missing")
}
// Validate columns on a sharded partition table
for _, shard := range []string{"node_observation_frps_00"} {
for _, column := range []string{"frps_client_count", "frps_proxies"} {
if !migrator.HasColumn(shard, column) {
return fmt.Errorf("column %s.%s is missing", shard, column)
}
}
}
return nil
}
@@ -1,41 +0,0 @@
// v8 升级内容:为 proxy_routes 增加域名级证书绑定字段 domain_cert_ids,并回填已有站点的证书映射。
// 背景说明:v1-v7 已作为历史初始基线合并;v8 是当前保留逐版本升级链的起点,用于把早期站点级证书列表扩展为每个域名可独立绑定证书。
package migrate
import "gorm.io/gorm"
func init() {
Register(V8())
}
func V8() Migration {
return Migration{
FromVersion: 7,
ToVersion: 8,
Migrate: migrateV8,
Validate: validateV8,
}
}
func migrateV8(ctx Context, db *gorm.DB, backend string) error {
if err := ctx.ApplyCurrentSchema(db, backend); err != nil {
return err
}
if err := ctx.BackfillOriginsFromProxyRoutes(db); err != nil {
return err
}
if err := ctx.BackfillProxyRouteSiteFields(db); err != nil {
return err
}
if err := ctx.EnsureProxyRouteSiteNameUniqueIndex(db); err != nil {
return err
}
if err := ctx.BackfillProxyRouteCertificateFields(db); err != nil {
return err
}
return ctx.BackfillProxyRouteDomainCertificateFields(db)
}
func validateV8(ctx Context, db *gorm.DB, backend string) error {
return ctx.ValidateDatabaseSchemaVersion(db, backend, 8)
}
@@ -1,29 +0,0 @@
// v9 升级内容:为 proxy_routes 增加 PoW 防护配置字段。
// 背景说明:反向代理站点需要支持 Proof-of-Work 抗机器人能力,因此在路由配置中持久化 PoW 开关与策略,并沿用 v8 的证书与站点字段回填。
package migrate
import "gorm.io/gorm"
func init() {
Register(V9())
}
func V9() Migration {
return Migration{
FromVersion: 8,
ToVersion: 9,
Migrate: migrateV9,
Validate: validateV9,
}
}
func migrateV9(ctx Context, db *gorm.DB, backend string) error {
if err := migrateV8(ctx, db, backend); err != nil {
return err
}
return nil
}
func validateV9(ctx Context, db *gorm.DB, backend string) error {
return ctx.ValidateDatabaseSchemaVersion(db, backend, 9)
}
File diff suppressed because it is too large Load Diff
-91
View File
@@ -1,91 +0,0 @@
package model
import "time"
type Node struct {
ID uint `json:"id" gorm:"primaryKey"`
NodeID string `json:"node_id" gorm:"uniqueIndex;size:64;not null"`
Name string `json:"name" gorm:"size:128;not null"`
IP string `json:"ip" gorm:"size:64;not null"`
IPManualOverride bool `json:"ip_manual_override" gorm:"not null;default:false"`
GeoName string `json:"geo_name" gorm:"size:128"`
GeoLatitude *float64 `json:"geo_latitude"`
GeoLongitude *float64 `json:"geo_longitude"`
GeoManualOverride bool `json:"geo_manual_override" gorm:"not null;default:false"`
AccessToken string `json:"-" gorm:"column:access_token;size:128;index"`
AutoUpdateEnabled bool `json:"auto_update_enabled" gorm:"not null;default:false"`
UpdateRequested bool `json:"update_requested" gorm:"not null;default:false"`
UpdateChannel string `json:"update_channel" gorm:"size:16;not null;default:'stable'"`
UpdateTag string `json:"update_tag" gorm:"size:64"`
RestartOpenrestyRequested bool `json:"restart_openresty_requested" gorm:"not null;default:false"`
Version string `json:"version" gorm:"size:64;not null;default:''"`
ExtVersion string `json:"ext_version" gorm:"size:64"`
OpenrestyStatus string `json:"openresty_status" gorm:"size:16;not null;default:'unknown'"`
OpenrestyMessage string `json:"openresty_message" gorm:"type:text"`
Status string `json:"status" gorm:"size:16;not null;default:'offline'"`
CurrentVersion string `json:"current_version" gorm:"size:32"`
LastSeenAt time.Time `json:"last_seen_at"`
LastError string `json:"last_error" gorm:"type:text"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
// Node type: edge_node (default) | tunnel_relay | tunnel_client
NodeType string `json:"node_type" gorm:"size:32;not null;default:'edge_node'"`
// TunnelRelay specific fields
RelayBindPort int `json:"relay_bind_port" gorm:"not null;default:0"`
RelayVhostHTTPPort int `json:"relay_vhost_http_port" gorm:"not null;default:0"`
RelayAuthToken string `json:"-" gorm:"size:128"`
RelayAgentAccessAddr string `json:"relay_agent_access_addr" gorm:"size:255"`
RelayClientAccessAddr string `json:"relay_client_access_addr" gorm:"size:255"`
RelayClientProxyURL string `json:"relay_client_proxy_url" gorm:"size:512"`
CapabilitiesJSON string `json:"capabilities_json" gorm:"type:text;not null;default:'[]'"`
RelayStatus string `json:"relay_status" gorm:"size:16;not null;default:'unknown'"`
RelayWebServerEnabled bool `json:"relay_web_server_enabled" gorm:"not null;default:false"`
}
func ListNodes() (nodes []*Node, err error) {
err = DB.Order("id desc").Find(&nodes).Error
return nodes, err
}
func ListNodesByNodeIDs(nodeIDs []string) (nodes []*Node, err error) {
if len(nodeIDs) == 0 {
return []*Node{}, nil
}
err = DB.Where("node_id IN ?", nodeIDs).Find(&nodes).Error
return nodes, err
}
func GetNodeByNodeID(nodeID string) (*Node, error) {
node := &Node{}
err := DB.Where("node_id = ?", nodeID).First(node).Error
return node, err
}
func GetNodeByID(id uint) (*Node, error) {
node := &Node{}
err := DB.First(node, id).Error
return node, err
}
func GetNodeByAccessToken(token string) (*Node, error) {
node := &Node{}
err := DB.Where("access_token = ?", token).First(node).Error
return node, err
}
func (node *Node) Insert() error {
return DB.Create(node).Error
}
func (node *Node) Update() error {
return DB.Save(node).Error
}
func (node *Node) Delete() error {
return DB.Delete(node).Error
}
func ListNodesByType(nodeType string) (nodes []*Node, err error) {
err = DB.Where("node_type = ?", nodeType).Order("id desc").Find(&nodes).Error
return nodes, err
}
@@ -1,664 +0,0 @@
package model
import (
"fmt"
"sort"
"strings"
"sync"
"time"
"gorm.io/gorm"
)
type NodeAccessLog struct {
ID uint `json:"id" gorm:"primaryKey"`
NodeID string `json:"node_id" gorm:"index:,composite:node_logged_at,priority:1;size:64;not null"`
LoggedAt time.Time `json:"logged_at" gorm:"index;index:,composite:node_logged_at,priority:2"`
RemoteAddr string `json:"remote_addr" gorm:"index;size:128"`
Region string `json:"region" gorm:"size:128"`
Host string `json:"host" gorm:"index;size:255"`
Path string `json:"path" gorm:"size:2048"`
StatusCode int `json:"status_code" gorm:"index"`
CreatedAt time.Time `json:"created_at"`
}
type NodeAccessLogRegionCount struct {
Region string `json:"region"`
Count int64 `json:"count"`
}
type NodeAccessLogQuery struct {
NodeID string
RemoteAddr string
Host string
Path string
Since time.Time
Until time.Time
Page int
PageSize int
SortBy string
SortOrder string
}
type NodeAccessLogBucketQuery struct {
NodeID string
RemoteAddr string
Host string
Path string
Since time.Time
Page int
PageSize int
SortBy string
SortOrder string
FoldMinutes int
}
type NodeAccessLogBucketRow struct {
BucketEpoch int64 `json:"bucket_epoch"`
RequestCount int64 `json:"request_count"`
UniqueIPCount int64 `json:"unique_ip_count"`
UniqueHostCount int64 `json:"unique_host_count"`
SuccessCount int64 `json:"success_count"`
ClientErrorCount int64 `json:"client_error_count"`
ServerErrorCount int64 `json:"server_error_count"`
}
type NodeAccessLogBucketIPQuery struct {
NodeID string
RemoteAddr string
Host string
Path string
BucketStartedAt time.Time
FoldMinutes int
Page int
PageSize int
SortBy string
SortOrder string
}
type NodeAccessLogBucketIPRow struct {
RemoteAddr string `json:"remote_addr"`
RequestCount int64 `json:"request_count"`
SuccessCount int64 `json:"success_count"`
ClientErrorCount int64 `json:"client_error_count"`
ServerErrorCount int64 `json:"server_error_count"`
LastSeenEpoch int64 `json:"last_seen_epoch"`
}
type NodeAccessLogIPSummaryQuery struct {
NodeID string
RemoteAddr string
Host string
Since time.Time
Page int
PageSize int
SortBy string
SortOrder string
}
type NodeAccessLogIPSummaryRow struct {
RemoteAddr string `json:"remote_addr"`
TotalRequests int64 `json:"total_requests"`
RecentRequests int64 `json:"recent_requests"`
LastSeenEpoch int64 `json:"last_seen_epoch"`
}
type NodeAccessLogIPTrendQuery struct {
NodeID string
RemoteAddr string
Host string
Since time.Time
BucketMinutes int
}
type NodeAccessLogTrendPointRow struct {
BucketEpoch int64 `json:"bucket_epoch"`
RequestCount int64 `json:"request_count"`
}
func (log *NodeAccessLog) BeforeCreate(*gorm.DB) error {
return assignObservabilityID(&log.ID)
}
func ListNodeAccessLogs(query NodeAccessLogQuery) (logs []*NodeAccessLog, err error) {
if query.PageSize > 0 {
return listNodeAccessLogsPaginatedAcrossShards(query)
}
return listNodeAccessLogsAcrossShards(query)
}
func ListNodeAccessLogsForWAFIPGroup(query NodeAccessLogQuery) ([]*NodeAccessLog, error) {
return listNodeAccessLogsAcrossShards(query)
}
func CountNodeAccessLogs(query NodeAccessLogQuery) (totalRecords int64, totalIPs int64, err error) {
db := normalizeShardedDB(DB)
var countErr error
var distinctErr error
var wg sync.WaitGroup
wg.Add(2)
go func() {
defer wg.Done()
totalRecords, countErr = countNodeAccessLogRecordsAcrossShards(db, query)
}()
go func() {
defer wg.Done()
totalIPs, distinctErr = countDistinctNodeAccessLogIPsAcrossShards(db, query)
}()
wg.Wait()
if countErr != nil {
return 0, 0, countErr
}
if distinctErr != nil {
return 0, 0, distinctErr
}
return totalRecords, totalIPs, nil
}
func ListNodeAccessLogRegionCounts(nodeID string, since time.Time, limit int) (items []*NodeAccessLogRegionCount, err error) {
logs, err := listNodeAccessLogsAcrossShards(NodeAccessLogQuery{
NodeID: nodeID,
Since: since,
})
if err != nil {
return nil, err
}
counts := make(map[string]int64)
for _, item := range logs {
if item == nil {
continue
}
region := strings.TrimSpace(item.Region)
if region == "" {
continue
}
counts[region]++
}
items = make([]*NodeAccessLogRegionCount, 0, len(counts))
for region, count := range counts {
items = append(items, &NodeAccessLogRegionCount{
Region: region,
Count: count,
})
}
sort.Slice(items, func(i int, j int) bool {
if items[i].Count == items[j].Count {
return items[i].Region < items[j].Region
}
return items[i].Count > items[j].Count
})
if limit > 0 && len(items) > limit {
items = items[:limit]
}
return items, nil
}
func ListNodeAccessLogBuckets(query NodeAccessLogBucketQuery) (items []*NodeAccessLogBucketRow, err error) {
rows, err := buildNodeAccessLogBucketRows(query)
if err != nil {
return nil, err
}
start, end := paginateBounds(len(rows), query.Page, query.PageSize)
if start >= len(rows) {
return []*NodeAccessLogBucketRow{}, nil
}
return rows[start:end], nil
}
func CountNodeAccessLogBuckets(query NodeAccessLogBucketQuery) (total int64, err error) {
rows, err := buildNodeAccessLogBucketRows(query)
if err != nil {
return 0, err
}
return int64(len(rows)), nil
}
func ListNodeAccessLogBucketIPs(query NodeAccessLogBucketIPQuery) (items []*NodeAccessLogBucketIPRow, err error) {
rows, err := buildNodeAccessLogBucketIPRows(query)
if err != nil {
return nil, err
}
start, end := paginateBounds(len(rows), query.Page, query.PageSize)
if start >= len(rows) {
return []*NodeAccessLogBucketIPRow{}, nil
}
return rows[start:end], nil
}
func CountNodeAccessLogBucketIPs(query NodeAccessLogBucketIPQuery) (total int64, err error) {
rows, err := buildNodeAccessLogBucketIPRows(query)
if err != nil {
return 0, err
}
return int64(len(rows)), nil
}
func ListNodeAccessLogIPSummaries(query NodeAccessLogIPSummaryQuery, recentSince time.Time) (items []*NodeAccessLogIPSummaryRow, err error) {
rows, err := buildNodeAccessLogIPSummaryRows(query, recentSince)
if err != nil {
return nil, err
}
start, end := paginateBounds(len(rows), query.Page, query.PageSize)
if start >= len(rows) {
return []*NodeAccessLogIPSummaryRow{}, nil
}
return rows[start:end], nil
}
func CountNodeAccessLogIPSummaries(query NodeAccessLogIPSummaryQuery) (total int64, err error) {
rows, err := buildNodeAccessLogIPSummaryRows(query, time.Time{})
if err != nil {
return 0, err
}
return int64(len(rows)), nil
}
func ListNodeAccessLogIPTrend(query NodeAccessLogIPTrendQuery) (items []*NodeAccessLogTrendPointRow, err error) {
return queryIPTrendRows(query)
}
func DeleteNodeAccessLogsBefore(before time.Time) (deleted int64, err error) {
return deleteAcrossShards(DB, "node_access_logs", &NodeAccessLog{}, func(tx *gorm.DB) *gorm.DB {
return tx.Where("logged_at < ?", before)
})
}
func DeleteAllNodeAccessLogs(db *gorm.DB) (deleted int64, err error) {
return deleteAcrossShards(db, "node_access_logs", &NodeAccessLog{}, nil)
}
func NodeAccessLogExists(db *gorm.DB, record *NodeAccessLog) (bool, error) {
if record == nil {
return false, nil
}
db = normalizeShardedDB(db)
for _, table := range observabilityShardTables("node_access_logs") {
var count int64
if err := db.Table(table).
Where(
"node_id = ? AND logged_at = ? AND remote_addr = ? AND host = ? AND path = ? AND status_code = ?",
record.NodeID,
record.LoggedAt,
record.RemoteAddr,
record.Host,
record.Path,
record.StatusCode,
).
Limit(1).
Count(&count).Error; err != nil {
return false, err
}
if count > 0 {
return true, nil
}
}
return false, nil
}
func DeleteNodeAccessLogsByNodeBefore(db *gorm.DB, nodeID string, before time.Time) (deleted int64, err error) {
return deleteAcrossShards(db, "node_access_logs", &NodeAccessLog{}, func(tx *gorm.DB) *gorm.DB {
return tx.Where("node_id = ? AND logged_at < ?", nodeID, before)
})
}
func buildNodeAccessLogFilterClause(query NodeAccessLogQuery) (string, []any) {
parts := make([]string, 0, 6)
args := make([]any, 0, 6)
if trimmed := strings.TrimSpace(query.NodeID); trimmed != "" {
parts = append(parts, "node_id = ?")
args = append(args, trimmed)
}
if trimmed := strings.TrimSpace(query.RemoteAddr); trimmed != "" {
parts = append(parts, "remote_addr LIKE ?")
args = append(args, trimmed+"%")
}
if trimmed := strings.TrimSpace(query.Host); trimmed != "" {
parts = append(parts, "host LIKE ?")
args = append(args, trimmed+"%")
}
if trimmed := strings.TrimSpace(query.Path); trimmed != "" {
parts = append(parts, "path LIKE ?")
args = append(args, trimmed+"%")
}
if !query.Since.IsZero() {
parts = append(parts, "logged_at >= ?")
args = append(args, query.Since)
}
if !query.Until.IsZero() {
parts = append(parts, "logged_at < ?")
args = append(args, query.Until)
}
if len(parts) == 0 {
return "TRUE", nil
}
return strings.Join(parts, " AND "), args
}
func applyNodeAccessLogFilters(db *gorm.DB, query NodeAccessLogQuery) *gorm.DB {
clause, args := buildNodeAccessLogFilterClause(query)
if clause == "TRUE" {
return db
}
return db.Where(clause, args...)
}
func countNodeAccessLogRecordsAcrossShards(db *gorm.DB, query NodeAccessLogQuery) (int64, error) {
tables := observabilityShardTables("node_access_logs")
counts := make([]int64, len(tables))
errs := make([]error, len(tables))
var wg sync.WaitGroup
for index, table := range tables {
wg.Add(1)
go func(index int, table string) {
defer wg.Done()
var count int64
errs[index] = applyNodeAccessLogFilters(db.Table(table), query).Count(&count).Error
counts[index] = count
}(index, table)
}
wg.Wait()
var total int64
for index := range tables {
if errs[index] != nil {
return 0, errs[index]
}
total += counts[index]
}
return total, nil
}
func countDistinctNodeAccessLogIPsAcrossShards(db *gorm.DB, query NodeAccessLogQuery) (int64, error) {
clause, args := buildNodeAccessLogFilterClause(query)
tables := observabilityShardTables("node_access_logs")
unionParts := make([]string, 0, len(tables))
allArgs := make([]any, 0, len(args)*len(tables))
for _, table := range tables {
unionParts = append(unionParts, fmt.Sprintf(
"SELECT TRIM(remote_addr) AS remote_addr FROM %s WHERE %s AND remote_addr <> ''",
table,
clause,
))
allArgs = append(allArgs, args...)
}
sql := fmt.Sprintf(`
SELECT COUNT(*) FROM (
SELECT remote_addr
FROM (%s) AS all_ips
GROUP BY remote_addr
) AS ips`, strings.Join(unionParts, " UNION ALL "))
var total int64
if err := db.Raw(sql, allArgs...).Scan(&total).Error; err != nil {
return 0, err
}
return total, nil
}
func listNodeAccessLogsAcrossShards(query NodeAccessLogQuery) ([]*NodeAccessLog, error) {
items, err := queryAcrossShards("node_access_logs", func(tx *gorm.DB) ([]*NodeAccessLog, error) {
var shardRows []*NodeAccessLog
if err := applyNodeAccessLogFilters(tx, query).Find(&shardRows).Error; err != nil {
return nil, err
}
return shardRows, nil
})
if err != nil {
return nil, err
}
sortNodeAccessLogs(items, query.SortBy, query.SortOrder)
return items, nil
}
func listNodeAccessLogsPaginatedAcrossShards(query NodeAccessLogQuery) ([]*NodeAccessLog, error) {
fetchLimit := nodeAccessLogFetchLimit(query.Page, query.PageSize)
orderClause := nodeAccessLogOrderClause(query.SortBy, query.SortOrder)
items := make([]*NodeAccessLog, 0, fetchLimit*observabilityShardCount)
db := normalizeShardedDB(DB)
for _, table := range observabilityShardTables("node_access_logs") {
var shardRows []*NodeAccessLog
tx := applyNodeAccessLogFilters(db.Table(table), query).Order(orderClause).Limit(fetchLimit)
if err := tx.Find(&shardRows).Error; err != nil {
return nil, err
}
items = append(items, shardRows...)
}
sortNodeAccessLogs(items, query.SortBy, query.SortOrder)
start, end := paginateBounds(len(items), query.Page, query.PageSize)
if start >= len(items) {
return []*NodeAccessLog{}, nil
}
return items[start:end], nil
}
func nodeAccessLogFetchLimit(page int, pageSize int) int {
if page < 0 {
page = 0
}
if pageSize <= 0 {
return 0
}
return (page + 1) * pageSize
}
func nodeAccessLogOrderClause(sortBy string, sortOrder string) string {
direction := "DESC"
if normalizeSortOrder(sortOrder) == "asc" {
direction = "ASC"
}
column := "logged_at"
switch strings.TrimSpace(sortBy) {
case "status_code":
column = "status_code"
case "remote_addr":
column = "remote_addr"
case "host":
column = "host"
case "path":
column = "path"
}
if column == "logged_at" {
return column + " " + direction + ", id " + direction
}
return column + " " + direction + ", logged_at " + direction + ", id " + direction
}
func sortNodeAccessLogBucketIPRows(items []*NodeAccessLogBucketIPRow, sortBy string, sortOrder string) {
desc := normalizeSortOrder(sortOrder) != "asc"
sort.Slice(items, func(i int, 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 = compareInt64(left.LastSeenEpoch, right.LastSeenEpoch)
case "remote_addr":
compare = strings.Compare(left.RemoteAddr, right.RemoteAddr)
default:
compare = compareInt64(left.RequestCount, right.RequestCount)
}
if compare == 0 {
compare = compareInt64(left.LastSeenEpoch, right.LastSeenEpoch)
}
if compare == 0 {
compare = strings.Compare(left.RemoteAddr, right.RemoteAddr)
}
if desc {
return compare > 0
}
return compare < 0
})
}
func sortNodeAccessLogs(items []*NodeAccessLog, sortBy string, sortOrder string) {
desc := normalizeSortOrder(sortOrder) != "asc"
sort.Slice(items, func(i int, 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 = compareInt(left.StatusCode, right.StatusCode)
case "remote_addr":
compare = strings.Compare(left.RemoteAddr, right.RemoteAddr)
case "host":
compare = strings.Compare(left.Host, right.Host)
case "path":
compare = strings.Compare(left.Path, right.Path)
default:
compare = compareTime(left.LoggedAt, right.LoggedAt)
}
if compare == 0 {
compare = compareTime(left.LoggedAt, right.LoggedAt)
}
if compare == 0 {
compare = compareUint(left.ID, right.ID)
}
if desc {
return compare > 0
}
return compare < 0
})
}
func sortNodeAccessLogBucketRows(items []*NodeAccessLogBucketRow, sortBy string, sortOrder string) {
desc := normalizeSortOrder(sortOrder) != "asc"
sort.Slice(items, func(i int, 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 = compareInt64(left.RequestCount, right.RequestCount)
default:
compare = compareInt64(left.BucketEpoch, right.BucketEpoch)
}
if compare == 0 {
compare = compareInt64(left.BucketEpoch, right.BucketEpoch)
}
if desc {
return compare > 0
}
return compare < 0
})
}
func sortNodeAccessLogIPSummaryRows(items []*NodeAccessLogIPSummaryRow, sortBy string, sortOrder string) {
desc := normalizeSortOrder(sortOrder) != "asc"
sort.Slice(items, func(i int, 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 "recent_requests":
compare = compareInt64(left.RecentRequests, right.RecentRequests)
case "last_seen_at":
compare = compareInt64(left.LastSeenEpoch, right.LastSeenEpoch)
case "remote_addr":
compare = strings.Compare(left.RemoteAddr, right.RemoteAddr)
default:
compare = compareInt64(left.TotalRequests, right.TotalRequests)
}
if compare == 0 {
compare = compareInt64(left.LastSeenEpoch, right.LastSeenEpoch)
}
if compare == 0 {
compare = strings.Compare(left.RemoteAddr, right.RemoteAddr)
}
if desc {
return compare > 0
}
return compare < 0
})
}
func paginateBounds(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 bucketEpochForTime(value time.Time, bucketMinutes int) int64 {
bucketSeconds := int64(bucketMinutes * 60)
if bucketSeconds <= 0 {
bucketSeconds = 180
}
return (value.UTC().Unix() / bucketSeconds) * bucketSeconds
}
func compareTime(left time.Time, right time.Time) int {
switch {
case left.After(right):
return 1
case left.Before(right):
return -1
default:
return 0
}
}
func compareInt(left int, right int) int {
switch {
case left > right:
return 1
case left < right:
return -1
default:
return 0
}
}
func compareInt64(left int64, right int64) int {
switch {
case left > right:
return 1
case left < right:
return -1
default:
return 0
}
}
func compareUint(left uint, right uint) int {
switch {
case left > right:
return 1
case left < right:
return -1
default:
return 0
}
}
func normalizeSortOrder(sortOrder string) string {
if strings.EqualFold(strings.TrimSpace(sortOrder), "asc") {
return "asc"
}
return "desc"
}
@@ -1,421 +0,0 @@
package model
import (
"fmt"
"sort"
"strings"
"time"
"gorm.io/gorm"
)
type shardBucketAggregateRow struct {
BucketEpoch int64 `gorm:"column:bucket_epoch"`
RequestCount int64 `gorm:"column:request_count"`
SuccessCount int64 `gorm:"column:success_count"`
ClientErrorCount int64 `gorm:"column:client_error_count"`
ServerErrorCount int64 `gorm:"column:server_error_count"`
}
type shardBucketDimensionRow struct {
BucketEpoch int64 `gorm:"column:bucket_epoch"`
Value string `gorm:"column:value"`
}
type shardIPAggregateRow struct {
RemoteAddr string `gorm:"column:remote_addr"`
RequestCount int64 `gorm:"column:request_count"`
SuccessCount int64 `gorm:"column:success_count"`
ClientErrorCount int64 `gorm:"column:client_error_count"`
ServerErrorCount int64 `gorm:"column:server_error_count"`
LastSeenEpoch int64 `gorm:"column:last_seen_epoch"`
}
type shardIPSummaryRow struct {
RemoteAddr string `gorm:"column:remote_addr"`
TotalRequests int64 `gorm:"column:total_requests"`
RecentRequests int64 `gorm:"column:recent_requests"`
LastSeenEpoch int64 `gorm:"column:last_seen_epoch"`
}
type shardIPTrendRow struct {
BucketEpoch int64 `gorm:"column:bucket_epoch"`
RequestCount int64 `gorm:"column:request_count"`
}
func buildNodeAccessLogBucketRows(query NodeAccessLogBucketQuery) ([]*NodeAccessLogBucketRow, error) {
db := normalizeShardedDB(DB)
filter := nodeAccessLogQueryFromBucket(query)
clause, args := buildNodeAccessLogFilterClause(filter)
bucketSeconds := int64(query.FoldMinutes * 60)
if bucketSeconds <= 0 {
bucketSeconds = 180
}
bucketExpr := accessLogBucketEpochExpr(databaseDialect(db), bucketSeconds)
type bucketAccumulator struct {
requestCount int64
uniqueIPs map[string]struct{}
uniqueHosts map[string]struct{}
successCount int64
clientErrorCount int64
serverErrorCount int64
}
accumulators := make(map[int64]*bucketAccumulator)
for _, table := range observabilityShardTables("node_access_logs") {
var partials []shardBucketAggregateRow
sql := fmt.Sprintf(`
SELECT
%s AS bucket_epoch,
COUNT(*) AS request_count,
SUM(CASE WHEN status_code < 400 THEN 1 ELSE 0 END) AS success_count,
SUM(CASE WHEN status_code >= 400 AND status_code < 500 THEN 1 ELSE 0 END) AS client_error_count,
SUM(CASE WHEN status_code >= 500 THEN 1 ELSE 0 END) AS server_error_count
FROM %s
WHERE %s
GROUP BY bucket_epoch`, bucketExpr, table, clause)
if err := db.Raw(sql, args...).Scan(&partials).Error; err != nil {
return nil, err
}
for _, partial := range partials {
accumulator := accumulators[partial.BucketEpoch]
if accumulator == nil {
accumulator = &bucketAccumulator{
uniqueIPs: make(map[string]struct{}),
uniqueHosts: make(map[string]struct{}),
}
accumulators[partial.BucketEpoch] = accumulator
}
accumulator.requestCount += partial.RequestCount
accumulator.successCount += partial.SuccessCount
accumulator.clientErrorCount += partial.ClientErrorCount
accumulator.serverErrorCount += partial.ServerErrorCount
}
for _, column := range []string{"remote_addr", "host"} {
dimensions, err := queryBucketDimensionRows(db, table, clause, args, column, bucketExpr)
if err != nil {
return nil, err
}
for _, item := range dimensions {
accumulator := accumulators[item.BucketEpoch]
if accumulator == nil {
accumulator = &bucketAccumulator{
uniqueIPs: make(map[string]struct{}),
uniqueHosts: make(map[string]struct{}),
}
accumulators[item.BucketEpoch] = accumulator
}
trimmed := strings.TrimSpace(item.Value)
if trimmed == "" {
continue
}
switch column {
case "remote_addr":
accumulator.uniqueIPs[trimmed] = struct{}{}
case "host":
accumulator.uniqueHosts[trimmed] = struct{}{}
}
}
}
}
rows := make([]*NodeAccessLogBucketRow, 0, len(accumulators))
for bucketEpoch, accumulator := range accumulators {
rows = append(rows, &NodeAccessLogBucketRow{
BucketEpoch: bucketEpoch,
RequestCount: accumulator.requestCount,
UniqueIPCount: int64(len(accumulator.uniqueIPs)),
UniqueHostCount: int64(len(accumulator.uniqueHosts)),
SuccessCount: accumulator.successCount,
ClientErrorCount: accumulator.clientErrorCount,
ServerErrorCount: accumulator.serverErrorCount,
})
}
sortNodeAccessLogBucketRows(rows, query.SortBy, query.SortOrder)
return rows, nil
}
func queryBucketDimensionRows(db *gorm.DB, table string, clause string, args []any, column string, bucketExpr string) ([]shardBucketDimensionRow, error) {
var rows []shardBucketDimensionRow
sql := fmt.Sprintf(`
SELECT
%s AS bucket_epoch,
TRIM(%s) AS value
FROM %s
WHERE %s AND TRIM(%s) <> ''
GROUP BY bucket_epoch, TRIM(%s)`, bucketExpr, column, table, clause, column, column)
if err := db.Raw(sql, args...).Scan(&rows).Error; err != nil {
return nil, err
}
return rows, nil
}
func buildNodeAccessLogBucketIPRows(query NodeAccessLogBucketIPQuery) ([]*NodeAccessLogBucketIPRow, error) {
if query.BucketStartedAt.IsZero() {
return []*NodeAccessLogBucketIPRow{}, nil
}
foldMinutes := query.FoldMinutes
if foldMinutes <= 0 {
foldMinutes = 3
}
bucketStartedAt := query.BucketStartedAt.UTC()
filter := NodeAccessLogQuery{
NodeID: query.NodeID,
RemoteAddr: query.RemoteAddr,
Host: query.Host,
Path: query.Path,
Since: bucketStartedAt,
Until: bucketStartedAt.Add(time.Duration(foldMinutes) * time.Minute),
}
rows, err := queryIPAggregateRows(filter, false)
if err != nil {
return nil, err
}
sortNodeAccessLogBucketIPRows(rows, query.SortBy, query.SortOrder)
return rows, nil
}
func buildNodeAccessLogIPSummaryRows(query NodeAccessLogIPSummaryQuery, recentSince time.Time) ([]*NodeAccessLogIPSummaryRow, error) {
filter := NodeAccessLogQuery{
NodeID: query.NodeID,
RemoteAddr: query.RemoteAddr,
Host: query.Host,
Since: query.Since,
}
db := normalizeShardedDB(DB)
clause, args := buildNodeAccessLogFilterClause(filter)
lastSeenExpr := accessLogEpochExpr(databaseDialect(db))
type accumulator struct {
totalRequests int64
recentRequests int64
lastSeenEpoch int64
}
accumulators := make(map[string]*accumulator)
for _, table := range observabilityShardTables("node_access_logs") {
recentClause := "0"
queryArgs := make([]any, 0, len(args)+1)
if !recentSince.IsZero() {
recentClause = "CASE WHEN logged_at >= ? THEN 1 ELSE 0 END"
queryArgs = append(queryArgs, recentSince)
}
queryArgs = append(queryArgs, args...)
var partials []shardIPSummaryRow
sql := fmt.Sprintf(`
SELECT
TRIM(remote_addr) AS remote_addr,
COUNT(*) AS total_requests,
SUM(%s) AS recent_requests,
MAX(%s) AS last_seen_epoch
FROM %s
WHERE %s AND TRIM(remote_addr) <> ''
GROUP BY TRIM(remote_addr)`, recentClause, lastSeenExpr, table, clause)
if err := db.Raw(sql, queryArgs...).Scan(&partials).Error; err != nil {
return nil, err
}
for _, partial := range partials {
remoteAddr := strings.TrimSpace(partial.RemoteAddr)
if remoteAddr == "" {
continue
}
acc := accumulators[remoteAddr]
if acc == nil {
acc = &accumulator{}
accumulators[remoteAddr] = acc
}
acc.totalRequests += partial.TotalRequests
acc.recentRequests += partial.RecentRequests
if partial.LastSeenEpoch > acc.lastSeenEpoch {
acc.lastSeenEpoch = partial.LastSeenEpoch
}
}
}
rows := make([]*NodeAccessLogIPSummaryRow, 0, len(accumulators))
for remoteAddr, acc := range accumulators {
rows = append(rows, &NodeAccessLogIPSummaryRow{
RemoteAddr: remoteAddr,
TotalRequests: acc.totalRequests,
RecentRequests: acc.recentRequests,
LastSeenEpoch: acc.lastSeenEpoch,
})
}
sortNodeAccessLogIPSummaryRows(rows, query.SortBy, query.SortOrder)
return rows, nil
}
func queryIPAggregateRows(filter NodeAccessLogQuery, exactRemoteAddr bool) ([]*NodeAccessLogBucketIPRow, error) {
db := normalizeShardedDB(DB)
clause, args := buildNodeAccessLogFilterClause(filter)
lastSeenExpr := accessLogEpochExpr(databaseDialect(db))
type accumulator struct {
requestCount int64
successCount int64
clientErrorCount int64
serverErrorCount int64
lastSeenEpoch int64
}
accumulators := make(map[string]*accumulator)
for _, table := range observabilityShardTables("node_access_logs") {
queryClause := clause
queryArgs := append([]any{}, args...)
if exactRemoteAddr {
trimmed := strings.TrimSpace(filter.RemoteAddr)
if trimmed == "" {
return []*NodeAccessLogBucketIPRow{}, nil
}
queryClause = combineSQLClauses(queryClause, "TRIM(remote_addr) = ?")
queryArgs = append(queryArgs, trimmed)
}
var partials []shardIPAggregateRow
sql := fmt.Sprintf(`
SELECT
TRIM(remote_addr) AS remote_addr,
COUNT(*) AS request_count,
SUM(CASE WHEN status_code < 400 THEN 1 ELSE 0 END) AS success_count,
SUM(CASE WHEN status_code >= 400 AND status_code < 500 THEN 1 ELSE 0 END) AS client_error_count,
SUM(CASE WHEN status_code >= 500 THEN 1 ELSE 0 END) AS server_error_count,
MAX(%s) AS last_seen_epoch
FROM %s
WHERE %s AND TRIM(remote_addr) <> ''
GROUP BY TRIM(remote_addr)`, lastSeenExpr, table, queryClause)
if err := db.Raw(sql, queryArgs...).Scan(&partials).Error; err != nil {
return nil, err
}
for _, partial := range partials {
remoteAddr := strings.TrimSpace(partial.RemoteAddr)
if remoteAddr == "" {
continue
}
acc := accumulators[remoteAddr]
if acc == nil {
acc = &accumulator{}
accumulators[remoteAddr] = acc
}
acc.requestCount += partial.RequestCount
acc.successCount += partial.SuccessCount
acc.clientErrorCount += partial.ClientErrorCount
acc.serverErrorCount += partial.ServerErrorCount
if partial.LastSeenEpoch > acc.lastSeenEpoch {
acc.lastSeenEpoch = partial.LastSeenEpoch
}
}
}
rows := make([]*NodeAccessLogBucketIPRow, 0, len(accumulators))
for remoteAddr, acc := range accumulators {
rows = append(rows, &NodeAccessLogBucketIPRow{
RemoteAddr: remoteAddr,
RequestCount: acc.requestCount,
SuccessCount: acc.successCount,
ClientErrorCount: acc.clientErrorCount,
ServerErrorCount: acc.serverErrorCount,
LastSeenEpoch: acc.lastSeenEpoch,
})
}
return rows, nil
}
func queryIPTrendRows(query NodeAccessLogIPTrendQuery) ([]*NodeAccessLogTrendPointRow, error) {
remoteAddr := strings.TrimSpace(query.RemoteAddr)
if remoteAddr == "" {
return []*NodeAccessLogTrendPointRow{}, nil
}
db := normalizeShardedDB(DB)
filter := NodeAccessLogQuery{
NodeID: query.NodeID,
RemoteAddr: remoteAddr,
Host: query.Host,
Since: query.Since,
}
clause, args := buildNodeAccessLogFilterClause(filter)
bucketSeconds := int64(query.BucketMinutes * 60)
if bucketSeconds <= 0 {
bucketSeconds = 1800
}
bucketExpr := accessLogBucketEpochExpr(databaseDialect(db), bucketSeconds)
queryClause := combineSQLClauses(clause, "TRIM(remote_addr) = ?")
queryArgs := append(append([]any{}, args...), remoteAddr)
buckets := make(map[int64]int64)
for _, table := range observabilityShardTables("node_access_logs") {
var partials []shardIPTrendRow
sql := fmt.Sprintf(`
SELECT
%s AS bucket_epoch,
COUNT(*) AS request_count
FROM %s
WHERE %s
GROUP BY bucket_epoch`, bucketExpr, table, queryClause)
if err := db.Raw(sql, queryArgs...).Scan(&partials).Error; err != nil {
return nil, err
}
for _, partial := range partials {
buckets[partial.BucketEpoch] += partial.RequestCount
}
}
items := make([]*NodeAccessLogTrendPointRow, 0, len(buckets))
for bucketEpoch, requestCount := range buckets {
items = append(items, &NodeAccessLogTrendPointRow{
BucketEpoch: bucketEpoch,
RequestCount: requestCount,
})
}
sort.Slice(items, func(i int, j int) bool {
return items[i].BucketEpoch < items[j].BucketEpoch
})
return items, nil
}
func nodeAccessLogQueryFromBucket(query NodeAccessLogBucketQuery) NodeAccessLogQuery {
return NodeAccessLogQuery{
NodeID: query.NodeID,
RemoteAddr: query.RemoteAddr,
Host: query.Host,
Path: query.Path,
Since: query.Since,
}
}
func databaseDialect(db *gorm.DB) string {
if db == nil || db.Dialector == nil {
return "sqlite"
}
switch db.Dialector.Name() {
case "postgres":
return "postgres"
default:
return "sqlite"
}
}
func accessLogBucketEpochExpr(dialect string, bucketSeconds int64) string {
switch dialect {
case "postgres":
return fmt.Sprintf("FLOOR(EXTRACT(EPOCH FROM logged_at AT TIME ZONE 'UTC') / %d) * %d", bucketSeconds, bucketSeconds)
default:
return fmt.Sprintf("(CAST(strftime('%%s', logged_at) AS INTEGER) / %d) * %d", bucketSeconds, bucketSeconds)
}
}
func accessLogEpochExpr(dialect string) string {
switch dialect {
case "postgres":
return "FLOOR(EXTRACT(EPOCH FROM logged_at AT TIME ZONE 'UTC'))::bigint"
default:
return "CAST((julianday(logged_at) - 2440587.5) * 86400 AS INTEGER)"
}
}
func combineSQLClauses(left string, right string) string {
if strings.TrimSpace(left) == "" || left == "TRUE" {
return right
}
return left + " AND " + right
}
@@ -1,542 +0,0 @@
package model
import (
"fmt"
"sort"
"strings"
"testing"
"time"
)
func TestListNodeAccessLogsPaginatedAcrossShards(t *testing.T) {
db := openBareTestSQLiteDB(t, "node_access_log_pagination.db")
if err := registerSharding(db, "sqlite"); err != nil {
t.Fatalf("register sharding: %v", err)
}
if err := autoMigrateAll(db); err != nil {
t.Fatalf("auto migrate db: %v", err)
}
previousDB := DB
DB = db
t.Cleanup(func() {
DB = previousDB
})
now := time.Now().UTC()
for index := range 15 {
record := &NodeAccessLog{
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,
}
if err := db.Create(record).Error; err != nil {
t.Fatalf("seed access log %d: %v", index, err)
}
}
query := NodeAccessLogQuery{
NodeID: "node-page",
Page: 1,
PageSize: 5,
SortBy: "logged_at",
SortOrder: "desc",
}
page, err := ListNodeAccessLogs(query)
if err != nil {
t.Fatalf("ListNodeAccessLogs failed: %v", err)
}
if len(page) != 5 {
t.Fatalf("expected 5 rows, got %d", len(page))
}
if page[0].Path != "/path-05" || page[4].Path != "/path-09" {
t.Fatalf("unexpected page ordering: %+v", page)
}
totalRecords, totalIPs, err := CountNodeAccessLogs(query)
if err != nil {
t.Fatalf("CountNodeAccessLogs failed: %v", err)
}
if totalRecords != 15 {
t.Fatalf("expected total_records=15, got %d", totalRecords)
}
if totalIPs != 5 {
t.Fatalf("expected total_ip=5, got %d", totalIPs)
}
}
func TestNodeAccessLogOptimizedQueriesMatchReference(t *testing.T) {
db := openBareTestSQLiteDB(t, "node_access_log_correctness.db")
if err := registerSharding(db, "sqlite"); err != nil {
t.Fatalf("register sharding: %v", err)
}
if err := autoMigrateAll(db); err != nil {
t.Fatalf("auto migrate db: %v", err)
}
previousDB := DB
DB = db
t.Cleanup(func() {
DB = previousDB
})
now := time.Now().UTC()
records := []*NodeAccessLog{
{NodeID: "node-a", LoggedAt: now.Add(-5 * time.Minute), RemoteAddr: "1.1.1.1", Host: "a.example.com", Path: "/alpha", StatusCode: 200},
{NodeID: "node-a", LoggedAt: now.Add(-4 * time.Minute), RemoteAddr: "2.2.2.2", Host: "a.example.com", Path: "/beta", StatusCode: 404},
{NodeID: "node-b", LoggedAt: now.Add(-3 * time.Minute), RemoteAddr: "1.1.1.1", Host: "b.example.com", Path: "/gamma", StatusCode: 502},
{NodeID: "node-b", LoggedAt: now.Add(-2 * time.Minute), RemoteAddr: " 3.3.3.3 ", Host: "b.example.com", Path: "/delta", StatusCode: 200},
{NodeID: "node-b", LoggedAt: now.Add(-1 * time.Minute), RemoteAddr: "", Host: "b.example.com", Path: "/empty-ip", StatusCode: 200},
}
for _, record := range records {
if err := db.Create(record).Error; err != nil {
t.Fatalf("seed access log: %v", err)
}
}
baseQuery := NodeAccessLogQuery{
Since: now.Add(-10 * time.Minute),
SortBy: "logged_at",
SortOrder: "desc",
}
reference, err := listNodeAccessLogsAcrossShards(baseQuery)
if err != nil {
t.Fatalf("reference list failed: %v", err)
}
referenceTotal, referenceIPs, err := countNodeAccessLogsReference(baseQuery)
if err != nil {
t.Fatalf("reference count failed: %v", err)
}
totalRecords, totalIPs, err := CountNodeAccessLogs(baseQuery)
if err != nil {
t.Fatalf("CountNodeAccessLogs failed: %v", err)
}
if totalRecords != referenceTotal {
t.Fatalf("total_records mismatch: got %d want %d", totalRecords, referenceTotal)
}
if totalIPs != referenceIPs {
t.Fatalf("total_ip mismatch: got %d want %d", totalIPs, referenceIPs)
}
if totalRecords != int64(len(reference)) {
t.Fatalf("total_records should equal reference rows: got %d want %d", totalRecords, len(reference))
}
for page := range 3 {
query := baseQuery
query.Page = page
query.PageSize = 2
pageRows, err := ListNodeAccessLogs(query)
if err != nil {
t.Fatalf("ListNodeAccessLogs page %d failed: %v", page, err)
}
start, end := paginateBounds(len(reference), page, query.PageSize)
if start >= len(reference) {
if len(pageRows) != 0 {
t.Fatalf("page %d expected empty slice, got %d rows", page, len(pageRows))
}
continue
}
want := reference[start:end]
if !nodeAccessLogsEqual(pageRows, want) {
t.Fatalf("page %d mismatch:\n got=%+v\nwant=%+v", page, pageRows, want)
}
}
filteredQuery := NodeAccessLogQuery{
NodeID: "node-a",
Since: baseQuery.Since,
SortBy: "status_code",
SortOrder: "asc",
Page: 0,
PageSize: 10,
}
filteredReference, err := listNodeAccessLogsAcrossShards(filteredQuery)
if err != nil {
t.Fatalf("filtered reference list failed: %v", err)
}
filteredRows, err := ListNodeAccessLogs(filteredQuery)
if err != nil {
t.Fatalf("filtered ListNodeAccessLogs failed: %v", err)
}
if !nodeAccessLogsEqual(filteredRows, filteredReference) {
t.Fatalf("filtered list mismatch:\n got=%+v\nwant=%+v", filteredRows, filteredReference)
}
filteredTotal, filteredIPs, err := CountNodeAccessLogs(filteredQuery)
if err != nil {
t.Fatalf("filtered CountNodeAccessLogs failed: %v", err)
}
wantFilteredTotal, wantFilteredIPs, err := countNodeAccessLogsReference(filteredQuery)
if err != nil {
t.Fatalf("filtered reference count failed: %v", err)
}
if filteredTotal != wantFilteredTotal || filteredIPs != wantFilteredIPs {
t.Fatalf("filtered count mismatch: got (%d,%d) want (%d,%d)", filteredTotal, filteredIPs, wantFilteredTotal, wantFilteredIPs)
}
}
func countNodeAccessLogsReference(query NodeAccessLogQuery) (int64, int64, error) {
all, err := listNodeAccessLogsAcrossShards(query)
if err != nil {
return 0, 0, err
}
ips := make(map[string]struct{})
for _, item := range all {
if item == nil {
continue
}
if trimmed := strings.TrimSpace(item.RemoteAddr); trimmed != "" {
ips[trimmed] = struct{}{}
}
}
return int64(len(all)), int64(len(ips)), nil
}
func nodeAccessLogsEqual(left []*NodeAccessLog, right []*NodeAccessLog) bool {
if len(left) != len(right) {
return false
}
for index := range left {
if left[index] == nil || right[index] == nil {
if left[index] != right[index] {
return false
}
continue
}
if left[index].ID != right[index].ID ||
left[index].NodeID != right[index].NodeID ||
!left[index].LoggedAt.Equal(right[index].LoggedAt) ||
left[index].RemoteAddr != right[index].RemoteAddr ||
left[index].Host != right[index].Host ||
left[index].Path != right[index].Path ||
left[index].StatusCode != right[index].StatusCode {
return false
}
}
return true
}
func TestNodeAccessLogAggregationsMatchReference(t *testing.T) {
db := openBareTestSQLiteDB(t, "node_access_log_agg.db")
if err := registerSharding(db, "sqlite"); err != nil {
t.Fatalf("register sharding: %v", err)
}
if err := autoMigrateAll(db); err != nil {
t.Fatalf("auto migrate db: %v", err)
}
previousDB := DB
DB = db
t.Cleanup(func() {
DB = previousDB
})
now := time.Date(2026, 3, 19, 8, 12, 30, 0, time.UTC)
records := []*NodeAccessLog{
{NodeID: "node-folded", LoggedAt: now.Add(-4 * time.Minute), RemoteAddr: "203.0.113.1", Host: "alpha.example.com", Path: "/first", StatusCode: 200},
{NodeID: "node-folded", LoggedAt: now.Add(-3 * time.Minute), RemoteAddr: "203.0.113.1", Host: "alpha.example.com", Path: "/second", StatusCode: 502},
{NodeID: "node-folded", LoggedAt: now.Add(-2 * time.Minute), RemoteAddr: "203.0.113.2", Host: "beta.example.com", Path: "/third", StatusCode: 404},
}
for _, record := range records {
if err := db.Create(record).Error; err != nil {
t.Fatalf("seed access log: %v", err)
}
}
since := now.Add(-10 * time.Minute)
bucketRows, err := buildNodeAccessLogBucketRows(NodeAccessLogBucketQuery{
NodeID: "node-folded", Since: since, FoldMinutes: 5, SortBy: "request_count", SortOrder: "desc",
})
if err != nil {
t.Fatalf("buildNodeAccessLogBucketRows failed: %v", err)
}
referenceBuckets := referenceBucketRows(records, 5, "request_count", "desc")
if !bucketRowsEqual(bucketRows, referenceBuckets) {
t.Fatalf("bucket rows mismatch:\n got=%+v\nwant=%+v", bucketRows, referenceBuckets)
}
if len(bucketRows) == 0 {
t.Fatal("expected bucket rows before bucket ip verification")
}
bucketStartedAt := time.Unix(bucketRows[0].BucketEpoch, 0).UTC()
bucketIPRows, err := buildNodeAccessLogBucketIPRows(NodeAccessLogBucketIPQuery{
NodeID: "node-folded", BucketStartedAt: bucketStartedAt, FoldMinutes: 5, SortBy: "request_count", SortOrder: "desc",
})
if err != nil {
t.Fatalf("buildNodeAccessLogBucketIPRows failed: %v", err)
}
referenceBucketIPs := referenceBucketIPRows(records, bucketStartedAt, 5, "request_count", "desc")
if !bucketIPRowsEqual(bucketIPRows, referenceBucketIPs) {
if len(bucketIPRows) > 0 && len(referenceBucketIPs) > 0 {
t.Fatalf("bucket ip rows mismatch:\n got=%+v\nwant=%+v", *bucketIPRows[0], *referenceBucketIPs[0])
}
t.Fatalf("bucket ip rows mismatch:\n got=%+v\nwant=%+v", bucketIPRows, referenceBucketIPs)
}
recentSince := now.Add(-150 * time.Minute)
summaryRows, err := buildNodeAccessLogIPSummaryRows(NodeAccessLogIPSummaryQuery{
NodeID: "node-folded", Since: since, SortBy: "total_requests", SortOrder: "desc",
}, recentSince)
if err != nil {
t.Fatalf("buildNodeAccessLogIPSummaryRows failed: %v", err)
}
referenceSummaries := referenceIPSummaryRows(records, since, recentSince, "total_requests", "desc")
if !ipSummaryRowsEqual(summaryRows, referenceSummaries) {
t.Fatalf("ip summary rows mismatch:\n got=%+v\nwant=%+v", summaryRows, referenceSummaries)
}
trendRows, err := queryIPTrendRows(NodeAccessLogIPTrendQuery{
NodeID: "node-folded", RemoteAddr: "203.0.113.1", Since: since, BucketMinutes: 5,
})
if err != nil {
t.Fatalf("queryIPTrendRows failed: %v", err)
}
referenceTrend := referenceIPTrendRows(records, "203.0.113.1", 5)
if !trendRowsEqual(trendRows, referenceTrend) {
t.Fatalf("trend rows mismatch:\n got=%+v\nwant=%+v", trendRows, referenceTrend)
}
}
func referenceBucketRows(records []*NodeAccessLog, foldMinutes int, sortBy string, sortOrder string) []*NodeAccessLogBucketRow {
type bucketAccumulator struct {
requestCount int64
uniqueIPs map[string]struct{}
uniqueHosts map[string]struct{}
successCount int64
clientErrorCount int64
serverErrorCount int64
}
accumulators := make(map[int64]*bucketAccumulator)
for _, item := range records {
if item == nil {
continue
}
bucketEpoch := bucketEpochForTime(item.LoggedAt, foldMinutes)
accumulator := accumulators[bucketEpoch]
if accumulator == nil {
accumulator = &bucketAccumulator{
uniqueIPs: make(map[string]struct{}),
uniqueHosts: make(map[string]struct{}),
}
accumulators[bucketEpoch] = accumulator
}
accumulator.requestCount++
if trimmed := strings.TrimSpace(item.RemoteAddr); trimmed != "" {
accumulator.uniqueIPs[trimmed] = struct{}{}
}
if trimmed := strings.TrimSpace(item.Host); trimmed != "" {
accumulator.uniqueHosts[trimmed] = struct{}{}
}
switch {
case item.StatusCode < 400:
accumulator.successCount++
case item.StatusCode < 500:
accumulator.clientErrorCount++
default:
accumulator.serverErrorCount++
}
}
rows := make([]*NodeAccessLogBucketRow, 0, len(accumulators))
for bucketEpoch, accumulator := range accumulators {
rows = append(rows, &NodeAccessLogBucketRow{
BucketEpoch: bucketEpoch,
RequestCount: accumulator.requestCount,
UniqueIPCount: int64(len(accumulator.uniqueIPs)),
UniqueHostCount: int64(len(accumulator.uniqueHosts)),
SuccessCount: accumulator.successCount,
ClientErrorCount: accumulator.clientErrorCount,
ServerErrorCount: accumulator.serverErrorCount,
})
}
sortNodeAccessLogBucketRows(rows, sortBy, sortOrder)
return rows
}
func referenceBucketIPRows(records []*NodeAccessLog, bucketStartedAt time.Time, foldMinutes int, sortBy string, sortOrder string) []*NodeAccessLogBucketIPRow {
type accumulator struct {
requestCount int64
successCount int64
clientErrorCount int64
serverErrorCount int64
lastSeenAt time.Time
}
accumulators := make(map[string]*accumulator)
until := bucketStartedAt.Add(time.Duration(foldMinutes) * time.Minute)
for _, item := range records {
if item == nil || item.LoggedAt.Before(bucketStartedAt) || !item.LoggedAt.Before(until) {
continue
}
remoteAddr := strings.TrimSpace(item.RemoteAddr)
if remoteAddr == "" {
continue
}
acc := accumulators[remoteAddr]
if acc == nil {
acc = &accumulator{}
accumulators[remoteAddr] = acc
}
acc.requestCount++
switch {
case item.StatusCode < 400:
acc.successCount++
case item.StatusCode < 500:
acc.clientErrorCount++
default:
acc.serverErrorCount++
}
if item.LoggedAt.After(acc.lastSeenAt) {
acc.lastSeenAt = item.LoggedAt
}
}
rows := make([]*NodeAccessLogBucketIPRow, 0, len(accumulators))
for remoteAddr, acc := range accumulators {
rows = append(rows, &NodeAccessLogBucketIPRow{
RemoteAddr: remoteAddr,
RequestCount: acc.requestCount,
SuccessCount: acc.successCount,
ClientErrorCount: acc.clientErrorCount,
ServerErrorCount: acc.serverErrorCount,
LastSeenEpoch: acc.lastSeenAt.Unix(),
})
}
sortNodeAccessLogBucketIPRows(rows, sortBy, sortOrder)
return rows
}
func referenceIPSummaryRows(records []*NodeAccessLog, since time.Time, recentSince time.Time, sortBy string, sortOrder string) []*NodeAccessLogIPSummaryRow {
type accumulator struct {
totalRequests int64
recentRequests int64
lastSeenAt time.Time
}
accumulators := make(map[string]*accumulator)
for _, item := range records {
if item == nil || item.LoggedAt.Before(since) {
continue
}
remoteAddr := strings.TrimSpace(item.RemoteAddr)
if remoteAddr == "" {
continue
}
acc := accumulators[remoteAddr]
if acc == nil {
acc = &accumulator{}
accumulators[remoteAddr] = acc
}
acc.totalRequests++
if !recentSince.IsZero() && !item.LoggedAt.Before(recentSince) {
acc.recentRequests++
}
if item.LoggedAt.After(acc.lastSeenAt) {
acc.lastSeenAt = item.LoggedAt
}
}
rows := make([]*NodeAccessLogIPSummaryRow, 0, len(accumulators))
for remoteAddr, acc := range accumulators {
rows = append(rows, &NodeAccessLogIPSummaryRow{
RemoteAddr: remoteAddr,
TotalRequests: acc.totalRequests,
RecentRequests: acc.recentRequests,
LastSeenEpoch: acc.lastSeenAt.Unix(),
})
}
sortNodeAccessLogIPSummaryRows(rows, sortBy, sortOrder)
return rows
}
func referenceIPTrendRows(records []*NodeAccessLog, remoteAddr string, bucketMinutes int) []*NodeAccessLogTrendPointRow {
buckets := make(map[int64]int64)
for _, item := range records {
if item == nil || strings.TrimSpace(item.RemoteAddr) != remoteAddr {
continue
}
buckets[bucketEpochForTime(item.LoggedAt, bucketMinutes)]++
}
rows := make([]*NodeAccessLogTrendPointRow, 0, len(buckets))
for bucketEpoch, requestCount := range buckets {
rows = append(rows, &NodeAccessLogTrendPointRow{BucketEpoch: bucketEpoch, RequestCount: requestCount})
}
sort.Slice(rows, func(i int, j int) bool { return rows[i].BucketEpoch < rows[j].BucketEpoch })
return rows
}
func bucketRowsEqual(left []*NodeAccessLogBucketRow, right []*NodeAccessLogBucketRow) bool {
if len(left) != len(right) {
return false
}
for index := range left {
if left[index] == nil || right[index] == nil {
if left[index] != right[index] {
return false
}
continue
}
if *left[index] != *right[index] {
return false
}
}
return true
}
func bucketIPRowsEqual(left []*NodeAccessLogBucketIPRow, right []*NodeAccessLogBucketIPRow) bool {
if len(left) != len(right) {
return false
}
for index := range left {
if left[index] == nil || right[index] == nil {
if left[index] != right[index] {
return false
}
continue
}
if *left[index] != *right[index] {
return false
}
}
return true
}
func ipSummaryRowsEqual(left []*NodeAccessLogIPSummaryRow, right []*NodeAccessLogIPSummaryRow) bool {
if len(left) != len(right) {
return false
}
for index := range left {
if left[index] == nil || right[index] == nil {
if left[index] != right[index] {
return false
}
continue
}
if *left[index] != *right[index] {
return false
}
}
return true
}
func trendRowsEqual(left []*NodeAccessLogTrendPointRow, right []*NodeAccessLogTrendPointRow) bool {
if len(left) != len(right) {
return false
}
for index := range left {
if left[index] == nil || right[index] == nil {
if left[index] != right[index] {
return false
}
continue
}
if *left[index] != *right[index] {
return false
}
}
return true
}
func TestNodeAccessLogOrderClauseMatchesSort(t *testing.T) {
if got := nodeAccessLogOrderClause("logged_at", "desc"); got != "logged_at DESC, id DESC" {
t.Fatalf("unexpected logged_at order clause: %q", got)
}
if got := nodeAccessLogOrderClause("status_code", "asc"); got != "status_code ASC, logged_at ASC, id ASC" {
t.Fatalf("unexpected status_code order clause: %q", got)
}
}
@@ -1,47 +0,0 @@
package model
import "time"
type NodeHealthEvent struct {
ID uint `json:"id" gorm:"primaryKey"`
NodeID string `json:"node_id" gorm:"index;size:64;not null"`
EventType string `json:"event_type" gorm:"index;size:64;not null"`
Severity string `json:"severity" gorm:"size:16;not null"`
Status string `json:"status" gorm:"index;size:16;not null"`
Message string `json:"message" gorm:"type:text"`
FirstTriggeredAt time.Time `json:"first_triggered_at" gorm:"index"`
LastTriggeredAt time.Time `json:"last_triggered_at" gorm:"index"`
ReportedAt time.Time `json:"reported_at" gorm:"index"`
ResolvedAt *time.Time `json:"resolved_at" gorm:"index"`
MetadataJSON string `json:"metadata_json" gorm:"type:text"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
}
func GetActiveNodeHealthEvent(nodeID string, eventType string) (*NodeHealthEvent, error) {
event := &NodeHealthEvent{}
err := DB.Where("node_id = ? AND event_type = ? AND status = ?", nodeID, eventType, "active").First(event).Error
return event, err
}
func ListNodeHealthEvents(nodeID string, activeOnly bool, limit int) (events []*NodeHealthEvent, err error) {
query := DB.Where("node_id = ?", nodeID).Order("last_triggered_at desc")
if activeOnly {
query = query.Where("status = ?", "active")
}
if limit > 0 {
query = query.Limit(limit)
}
err = query.Find(&events).Error
return events, err
}
func ListActiveNodeHealthEvents() (events []*NodeHealthEvent, err error) {
err = DB.Where("status = ?", "active").Order("last_triggered_at desc").Find(&events).Error
return events, err
}
func DeleteNodeHealthEvents(nodeID string) (deleted int64, err error) {
result := DB.Where("node_id = ?", nodeID).Delete(&NodeHealthEvent{})
return result.RowsAffected, result.Error
}
@@ -1,107 +0,0 @@
package model
import (
"time"
"github.com/rain-kl/openflare/pkg/utils"
"gorm.io/gorm"
)
type NodeMetricSnapshot struct {
ID uint `json:"id" gorm:"primaryKey"`
NodeID string `json:"node_id" gorm:"index;size:64;not null"`
CapturedAt time.Time `json:"captured_at" gorm:"index"`
CPUUsagePercent float64 `json:"cpu_usage_percent"`
MemoryUsedBytes int64 `json:"memory_used_bytes"`
MemoryTotalBytes int64 `json:"memory_total_bytes"`
StorageUsedBytes int64 `json:"storage_used_bytes"`
StorageTotalBytes int64 `json:"storage_total_bytes"`
DiskReadBytes int64 `json:"disk_read_bytes"`
DiskWriteBytes int64 `json:"disk_write_bytes"`
NetworkRxBytes int64 `json:"network_rx_bytes"`
NetworkTxBytes int64 `json:"network_tx_bytes"`
CreatedAt time.Time `json:"created_at"`
}
func (snapshot *NodeMetricSnapshot) GetID() uint {
return snapshot.ID
}
func (snapshot *NodeMetricSnapshot) GetTime() time.Time {
return snapshot.CapturedAt
}
func (snapshot *NodeMetricSnapshot) BeforeCreate(tx *gorm.DB) error {
return assignObservabilityID(&snapshot.ID)
}
func (snapshot *NodeMetricSnapshot) Insert() error {
return DB.Create(snapshot).Error
}
func ListNodeMetricSnapshots(nodeID string, since time.Time, limit int) (snapshots []*NodeMetricSnapshot, err error) {
rows, err := queryAcrossShards("node_metric_snapshots", func(tx *gorm.DB) ([]*NodeMetricSnapshot, error) {
var shardRows []*NodeMetricSnapshot
query := tx.Order("captured_at desc, id desc")
if nodeID != "" {
query = query.Where("node_id = ?", nodeID)
}
if !since.IsZero() {
query = query.Where("captured_at >= ?", since)
}
if err := query.Find(&shardRows).Error; err != nil {
return nil, err
}
return shardRows, nil
})
if err != nil {
return nil, err
}
return utils.SortAndLimitRecords(rows, limit), nil
}
func ListMetricSnapshotsSince(since time.Time) (snapshots []*NodeMetricSnapshot, err error) {
rows, err := queryAcrossShards("node_metric_snapshots", func(tx *gorm.DB) ([]*NodeMetricSnapshot, error) {
var shardRows []*NodeMetricSnapshot
query := tx.Order("captured_at desc")
if !since.IsZero() {
query = query.Where("captured_at >= ?", since)
}
if err := query.Find(&shardRows).Error; err != nil {
return nil, err
}
return shardRows, nil
})
if err != nil {
return nil, err
}
return utils.SortAndLimitRecords(rows, 0), nil
}
func NodeMetricSnapshotExists(db *gorm.DB, nodeID string, capturedAt time.Time) (bool, error) {
db = normalizeShardedDB(db)
for _, table := range observabilityShardTables("node_metric_snapshots") {
var count int64
if err := db.Table(table).
Where("node_id = ? AND captured_at = ?", nodeID, capturedAt).
Limit(1).
Count(&count).Error; err != nil {
return false, err
}
if count > 0 {
return true, nil
}
}
return false, nil
}
func DeleteNodeMetricSnapshotsBefore(db *gorm.DB, before time.Time) (int64, error) {
return deleteAcrossShards(db, "node_metric_snapshots", &NodeMetricSnapshot{}, func(tx *gorm.DB) *gorm.DB {
return tx.Where("captured_at < ?", before)
})
}
func DeleteAllNodeMetricSnapshots(db *gorm.DB) (int64, error) {
return deleteAcrossShards(db, "node_metric_snapshots", &NodeMetricSnapshot{}, nil)
}
@@ -1,61 +0,0 @@
package model
import (
"time"
"github.com/rain-kl/openflare/pkg/utils"
"gorm.io/gorm"
)
type NodeObservationFrpc struct {
ID uint `json:"id" gorm:"primaryKey"`
NodeID string `json:"node_id" gorm:"index;size:64;not null"`
CapturedAt time.Time `json:"captured_at" gorm:"index"`
TunnelStatus string `json:"tunnel_status" gorm:"size:16"`
ConnectedRelaysCount int `json:"connected_relays_count"`
CreatedAt time.Time `json:"created_at"`
}
func (obs *NodeObservationFrpc) GetID() uint {
return obs.ID
}
func (obs *NodeObservationFrpc) GetTime() time.Time {
return obs.CapturedAt
}
func (obs *NodeObservationFrpc) BeforeCreate(tx *gorm.DB) error {
return assignObservabilityID(&obs.ID)
}
func (obs *NodeObservationFrpc) Insert() error {
return DB.Create(obs).Error
}
func ListNodeObservationFrpcs(nodeID string, since time.Time, limit int) (observations []*NodeObservationFrpc, err error) {
rows, err := queryAcrossShards("node_observation_frpcs", func(tx *gorm.DB) ([]*NodeObservationFrpc, error) {
var shardRows []*NodeObservationFrpc
query := tx.Order("captured_at desc, id desc")
if nodeID != "" {
query = query.Where("node_id = ?", nodeID)
}
if !since.IsZero() {
query = query.Where("captured_at >= ?", since)
}
if err := query.Find(&shardRows).Error; err != nil {
return nil, err
}
return shardRows, nil
})
if err != nil {
return nil, err
}
return utils.SortAndLimitRecords(rows, limit), nil
}
func DeleteNodeObservationFrpcsBefore(db *gorm.DB, before time.Time) (int64, error) {
return deleteAcrossShards(db, "node_observation_frpcs", &NodeObservationFrpc{}, func(tx *gorm.DB) *gorm.DB {
return tx.Where("captured_at < ?", before)
})
}
@@ -1,63 +0,0 @@
package model
import (
"time"
"github.com/rain-kl/openflare/pkg/utils"
"gorm.io/gorm"
)
type NodeObservationFrps struct {
ID uint `json:"id" gorm:"primaryKey"`
NodeID string `json:"node_id" gorm:"index;size:64;not null"`
CapturedAt time.Time `json:"captured_at" gorm:"index"`
FrpsConnections int `json:"frps_connections"`
FrpsProxyCount int `json:"frps_proxy_count"`
FrpsClientCount int `json:"frps_client_count"`
FrpsProxies string `json:"frps_proxies" gorm:"type:text"`
CreatedAt time.Time `json:"created_at"`
}
func (obs *NodeObservationFrps) GetID() uint {
return obs.ID
}
func (obs *NodeObservationFrps) GetTime() time.Time {
return obs.CapturedAt
}
func (obs *NodeObservationFrps) BeforeCreate(tx *gorm.DB) error {
return assignObservabilityID(&obs.ID)
}
func (obs *NodeObservationFrps) Insert() error {
return DB.Create(obs).Error
}
func ListNodeObservationFrps(nodeID string, since time.Time, limit int) (observations []*NodeObservationFrps, err error) {
rows, err := queryAcrossShards("node_observation_frps", func(tx *gorm.DB) ([]*NodeObservationFrps, error) {
var shardRows []*NodeObservationFrps
query := tx.Order("captured_at desc, id desc")
if nodeID != "" {
query = query.Where("node_id = ?", nodeID)
}
if !since.IsZero() {
query = query.Where("captured_at >= ?", since)
}
if err := query.Find(&shardRows).Error; err != nil {
return nil, err
}
return shardRows, nil
})
if err != nil {
return nil, err
}
return utils.SortAndLimitRecords(rows, limit), nil
}
func DeleteNodeObservationFrpsBefore(db *gorm.DB, before time.Time) (int64, error) {
return deleteAcrossShards(db, "node_observation_frps", &NodeObservationFrps{}, func(tx *gorm.DB) *gorm.DB {
return tx.Where("captured_at < ?", before)
})
}
@@ -1,62 +0,0 @@
package model
import (
"time"
"github.com/rain-kl/openflare/pkg/utils"
"gorm.io/gorm"
)
type NodeObservationOpenresty struct {
ID uint `json:"id" gorm:"primaryKey"`
NodeID string `json:"node_id" gorm:"index;size:64;not null"`
CapturedAt time.Time `json:"captured_at" gorm:"index"`
OpenrestyRxBytes int64 `json:"openresty_rx_bytes"`
OpenrestyTxBytes int64 `json:"openresty_tx_bytes"`
OpenrestyConnections int64 `json:"openresty_connections"`
CreatedAt time.Time `json:"created_at"`
}
func (obs *NodeObservationOpenresty) GetID() uint {
return obs.ID
}
func (obs *NodeObservationOpenresty) GetTime() time.Time {
return obs.CapturedAt
}
func (obs *NodeObservationOpenresty) BeforeCreate(tx *gorm.DB) error {
return assignObservabilityID(&obs.ID)
}
func (obs *NodeObservationOpenresty) Insert() error {
return DB.Create(obs).Error
}
func ListNodeObservationOpenresty(nodeID string, since time.Time, limit int) (observations []*NodeObservationOpenresty, err error) {
rows, err := queryAcrossShards("node_observation_openresties", func(tx *gorm.DB) ([]*NodeObservationOpenresty, error) {
var shardRows []*NodeObservationOpenresty
query := tx.Order("captured_at desc, id desc")
if nodeID != "" {
query = query.Where("node_id = ?", nodeID)
}
if !since.IsZero() {
query = query.Where("captured_at >= ?", since)
}
if err := query.Find(&shardRows).Error; err != nil {
return nil, err
}
return shardRows, nil
})
if err != nil {
return nil, err
}
return utils.SortAndLimitRecords(rows, limit), nil
}
func DeleteNodeObservationOpenrestiesBefore(db *gorm.DB, before time.Time) (int64, error) {
return deleteAcrossShards(db, "node_observation_openresties", &NodeObservationOpenresty{}, func(tx *gorm.DB) *gorm.DB {
return tx.Where("captured_at < ?", before)
})
}
@@ -1,105 +0,0 @@
package model
import (
"time"
"github.com/rain-kl/openflare/pkg/utils"
"gorm.io/gorm"
)
type NodeRequestReport struct {
ID uint `json:"id" gorm:"primaryKey"`
NodeID string `json:"node_id" gorm:"index;size:64;not null"`
WindowStartedAt time.Time `json:"window_started_at" gorm:"index"`
WindowEndedAt time.Time `json:"window_ended_at" gorm:"index"`
RequestCount int64 `json:"request_count"`
ErrorCount int64 `json:"error_count"`
UniqueVisitorCount int64 `json:"unique_visitor_count"`
StatusCodesJSON string `json:"status_codes_json" gorm:"type:text"`
TopDomainsJSON string `json:"top_domains_json" gorm:"type:text"`
SourceCountriesJSON string `json:"source_countries_json" gorm:"type:text"`
CreatedAt time.Time `json:"created_at"`
}
func (report *NodeRequestReport) GetID() uint {
return report.ID
}
func (report *NodeRequestReport) GetTime() time.Time {
return report.WindowEndedAt
}
func (report *NodeRequestReport) BeforeCreate(tx *gorm.DB) error {
return assignObservabilityID(&report.ID)
}
func (report *NodeRequestReport) Insert() error {
return DB.Create(report).Error
}
func ListNodeRequestReports(nodeID string, since time.Time, limit int) (reports []*NodeRequestReport, err error) {
rows, err := queryAcrossShards("node_request_reports", func(tx *gorm.DB) ([]*NodeRequestReport, error) {
var shardRows []*NodeRequestReport
query := tx.Order("window_ended_at desc, id desc")
if nodeID != "" {
query = query.Where("node_id = ?", nodeID)
}
if !since.IsZero() {
query = query.Where("window_ended_at >= ?", since)
}
if err := query.Find(&shardRows).Error; err != nil {
return nil, err
}
return shardRows, nil
})
if err != nil {
return nil, err
}
return utils.SortAndLimitRecords(rows, limit), nil
}
func ListRequestReportsSince(since time.Time) (reports []*NodeRequestReport, err error) {
rows, err := queryAcrossShards("node_request_reports", func(tx *gorm.DB) ([]*NodeRequestReport, error) {
var shardRows []*NodeRequestReport
query := tx.Order("window_ended_at desc")
if !since.IsZero() {
query = query.Where("window_ended_at >= ?", since)
}
if err := query.Find(&shardRows).Error; err != nil {
return nil, err
}
return shardRows, nil
})
if err != nil {
return nil, err
}
return utils.SortAndLimitRecords(rows, 0), nil
}
func NodeRequestReportExists(db *gorm.DB, nodeID string, windowStartedAt time.Time, windowEndedAt time.Time) (bool, error) {
db = normalizeShardedDB(db)
for _, table := range observabilityShardTables("node_request_reports") {
var count int64
if err := db.Table(table).
Where("node_id = ? AND window_started_at = ? AND window_ended_at = ?", nodeID, windowStartedAt, windowEndedAt).
Limit(1).
Count(&count).Error; err != nil {
return false, err
}
if count > 0 {
return true, nil
}
}
return false, nil
}
func DeleteNodeRequestReportsBefore(db *gorm.DB, before time.Time) (int64, error) {
return deleteAcrossShards(db, "node_request_reports", &NodeRequestReport{}, func(tx *gorm.DB) *gorm.DB {
return tx.Where("window_ended_at < ?", before)
})
}
func DeleteAllNodeRequestReports(db *gorm.DB) (int64, error) {
return deleteAcrossShards(db, "node_request_reports", &NodeRequestReport{}, nil)
}
@@ -1,54 +0,0 @@
package model
import (
"time"
"gorm.io/gorm/clause"
)
type NodeSystemProfile struct {
ID uint `json:"id" gorm:"primaryKey"`
NodeID string `json:"node_id" gorm:"uniqueIndex;size:64;not null"`
Hostname string `json:"hostname" gorm:"size:255"`
OSName string `json:"os_name" gorm:"size:128"`
OSVersion string `json:"os_version" gorm:"size:128"`
KernelVersion string `json:"kernel_version" gorm:"size:128"`
Architecture string `json:"architecture" gorm:"size:64"`
CPUModel string `json:"cpu_model" gorm:"size:255"`
CPUCores int `json:"cpu_cores"`
TotalMemoryBytes int64 `json:"total_memory_bytes"`
TotalDiskBytes int64 `json:"total_disk_bytes"`
UptimeSeconds int64 `json:"uptime_seconds"`
ReportedAt time.Time `json:"reported_at" gorm:"index"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
}
func GetNodeSystemProfile(nodeID string) (*NodeSystemProfile, error) {
profile := &NodeSystemProfile{}
err := DB.Where("node_id = ?", nodeID).First(profile).Error
return profile, err
}
func UpsertNodeSystemProfile(profile *NodeSystemProfile) error {
if profile == nil {
return nil
}
return DB.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(profile).Error
}
@@ -0,0 +1,778 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package model
import (
"context"
"errors"
"fmt"
"sort"
"strings"
"time"
"github.com/Rain-kl/Wavelet/internal/db"
"gorm.io/gorm"
)
const openFlareAccessLogTable = "of_node_access_logs"
type openFlareAccessLogBucketAggregateRow struct {
BucketEpoch int64 `gorm:"column:bucket_epoch"`
RequestCount int64 `gorm:"column:request_count"`
SuccessCount int64 `gorm:"column:success_count"`
ClientErrorCount int64 `gorm:"column:client_error_count"`
ServerErrorCount int64 `gorm:"column:server_error_count"`
}
type openFlareAccessLogBucketDimensionRow struct {
BucketEpoch int64 `gorm:"column:bucket_epoch"`
Value string `gorm:"column:value"`
}
type openFlareAccessLogIPAggregateRow struct {
RemoteAddr string `gorm:"column:remote_addr"`
RequestCount int64 `gorm:"column:request_count"`
SuccessCount int64 `gorm:"column:success_count"`
ClientErrorCount int64 `gorm:"column:client_error_count"`
ServerErrorCount int64 `gorm:"column:server_error_count"`
LastSeenEpoch int64 `gorm:"column:last_seen_epoch"`
}
type openFlareAccessLogIPSummaryRow struct {
RemoteAddr string `gorm:"column:remote_addr"`
TotalRequests int64 `gorm:"column:total_requests"`
RecentRequests int64 `gorm:"column:recent_requests"`
LastSeenEpoch int64 `gorm:"column:last_seen_epoch"`
}
type openFlareAccessLogIPTrendRow struct {
BucketEpoch int64 `gorm:"column:bucket_epoch"`
RequestCount int64 `gorm:"column:request_count"`
}
// ListOpenFlareAccessLogsForWAFIPGroup lists access logs in a time window for automatic IP group rules.
func ListOpenFlareAccessLogsForWAFIPGroup(ctx context.Context, query OpenFlareAccessLogQuery) ([]*OpenFlareAccessLog, error) {
return ListOpenFlareAccessLogs(ctx, query)
}
// ListOpenFlareAccessLogs lists access logs matching the query.
func ListOpenFlareAccessLogs(ctx context.Context, query OpenFlareAccessLogQuery) ([]*OpenFlareAccessLog, error) {
conn := db.DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
tx := applyOpenFlareAccessLogFilters(conn.Model(&OpenFlareAccessLog{}), query)
tx = tx.Order(openFlareAccessLogOrderClause(query.SortBy, query.SortOrder))
if query.PageSize > 0 {
if query.Page < 0 {
query.Page = 0
}
tx = tx.Offset(query.Page * query.PageSize).Limit(query.PageSize)
}
var rows []*OpenFlareAccessLog
if err := tx.Find(&rows).Error; err != nil {
if isMissingTableError(err) {
return []*OpenFlareAccessLog{}, nil
}
return nil, err
}
return rows, nil
}
// CountOpenFlareAccessLogs counts access logs and distinct IPs matching the query.
func CountOpenFlareAccessLogs(ctx context.Context, query OpenFlareAccessLogQuery) (int64, int64, error) {
conn := db.DB(ctx)
if conn == nil {
return 0, 0, errors.New(errDatabaseNotInitialized)
}
totalRecords, err := countOpenFlareAccessLogRecords(conn, query)
if err != nil {
if isMissingTableError(err) {
return 0, 0, nil
}
return 0, 0, err
}
totalIPs, err := countDistinctOpenFlareAccessLogIPs(conn, query)
if err != nil {
if isMissingTableError(err) {
return 0, 0, nil
}
return 0, 0, err
}
return totalRecords, totalIPs, nil
}
// ListOpenFlareAccessLogRegionCounts returns region counts for access logs.
func ListOpenFlareAccessLogRegionCounts(ctx context.Context, nodeID string, since time.Time, limit int) ([]*OpenFlareAccessLogRegionCount, error) {
conn := db.DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
filter := OpenFlareAccessLogQuery{
NodeID: nodeID,
Since: since,
}
clause, args := buildOpenFlareAccessLogFilterClause(filter)
sql := fmt.Sprintf(`
SELECT TRIM(region) AS region, COUNT(*) AS count
FROM %s
WHERE %s AND TRIM(region) <> ''
GROUP BY TRIM(region)
ORDER BY count DESC, region ASC`, openFlareAccessLogTable, clause)
if limit > 0 {
sql += fmt.Sprintf(" LIMIT %d", limit)
}
var rows []*OpenFlareAccessLogRegionCount
if err := conn.Raw(sql, args...).Scan(&rows).Error; err != nil {
if isMissingTableError(err) {
return []*OpenFlareAccessLogRegionCount{}, nil
}
return nil, err
}
return rows, nil
}
// ListOpenFlareAccessLogBuckets lists folded access log buckets.
func ListOpenFlareAccessLogBuckets(ctx context.Context, query OpenFlareAccessLogBucketQuery) ([]*OpenFlareAccessLogBucketRow, error) {
rows, err := buildOpenFlareAccessLogBucketRows(ctx, query)
if err != nil {
return nil, err
}
start, end := openFlareAccessLogPaginateBounds(len(rows), query.Page, query.PageSize)
if start >= len(rows) {
return []*OpenFlareAccessLogBucketRow{}, nil
}
return rows[start:end], nil
}
// CountOpenFlareAccessLogBuckets counts folded access log buckets.
func CountOpenFlareAccessLogBuckets(ctx context.Context, query OpenFlareAccessLogBucketQuery) (int64, error) {
rows, err := buildOpenFlareAccessLogBucketRows(ctx, query)
if err != nil {
return 0, err
}
return int64(len(rows)), nil
}
// ListOpenFlareAccessLogBucketIPs lists folded IP rows for a bucket window.
func ListOpenFlareAccessLogBucketIPs(ctx context.Context, query OpenFlareAccessLogBucketIPQuery) ([]*OpenFlareAccessLogBucketIPRow, error) {
rows, err := buildOpenFlareAccessLogBucketIPRows(ctx, query)
if err != nil {
return nil, err
}
start, end := openFlareAccessLogPaginateBounds(len(rows), query.Page, query.PageSize)
if start >= len(rows) {
return []*OpenFlareAccessLogBucketIPRow{}, nil
}
return rows[start:end], nil
}
// CountOpenFlareAccessLogBucketIPs counts folded IP rows for a bucket window.
func CountOpenFlareAccessLogBucketIPs(ctx context.Context, query OpenFlareAccessLogBucketIPQuery) (int64, error) {
rows, err := buildOpenFlareAccessLogBucketIPRows(ctx, query)
if err != nil {
return 0, err
}
return int64(len(rows)), nil
}
// ListOpenFlareAccessLogIPSummaries lists IP summaries.
func ListOpenFlareAccessLogIPSummaries(ctx context.Context, query OpenFlareAccessLogIPSummaryQuery, recentSince time.Time) ([]*OpenFlareAccessLogIPSummaryRow, error) {
rows, err := buildOpenFlareAccessLogIPSummaryRows(ctx, query, recentSince)
if err != nil {
return nil, err
}
start, end := openFlareAccessLogPaginateBounds(len(rows), query.Page, query.PageSize)
if start >= len(rows) {
return []*OpenFlareAccessLogIPSummaryRow{}, nil
}
return rows[start:end], nil
}
// CountOpenFlareAccessLogIPSummaries counts IP summaries.
func CountOpenFlareAccessLogIPSummaries(ctx context.Context, query OpenFlareAccessLogIPSummaryQuery) (int64, error) {
rows, err := buildOpenFlareAccessLogIPSummaryRows(ctx, query, time.Time{})
if err != nil {
return 0, err
}
return int64(len(rows)), nil
}
// ListOpenFlareAccessLogIPTrend lists IP trend points.
func ListOpenFlareAccessLogIPTrend(ctx context.Context, query OpenFlareAccessLogIPTrendQuery) ([]*OpenFlareAccessLogIPTrendRow, error) {
conn := db.DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
remoteAddr := strings.TrimSpace(query.RemoteAddr)
if remoteAddr == "" {
return []*OpenFlareAccessLogIPTrendRow{}, nil
}
filter := OpenFlareAccessLogQuery{
NodeID: query.NodeID,
RemoteAddr: remoteAddr,
Host: query.Host,
Since: query.Since,
}
clause, args := buildOpenFlareAccessLogFilterClause(filter)
bucketSeconds := int64(query.BucketMinutes * 60)
if bucketSeconds <= 0 {
bucketSeconds = 1800
}
bucketExpr := openFlareAccessLogBucketEpochExpr(openFlareAccessLogDialect(conn), bucketSeconds)
queryClause := combineOpenFlareAccessLogSQLClauses(clause, "TRIM(remote_addr) = ?")
queryArgs := append(append([]any{}, args...), remoteAddr)
sql := fmt.Sprintf(`
SELECT
%s AS bucket_epoch,
COUNT(*) AS request_count
FROM %s
WHERE %s
GROUP BY bucket_epoch
ORDER BY bucket_epoch ASC`, bucketExpr, openFlareAccessLogTable, queryClause)
var rows []*OpenFlareAccessLogIPTrendRow
if err := conn.Raw(sql, queryArgs...).Scan(&rows).Error; err != nil {
if isMissingTableError(err) {
return []*OpenFlareAccessLogIPTrendRow{}, nil
}
return nil, err
}
return rows, nil
}
// DeleteAllOpenFlareAccessLogs deletes all access logs.
func DeleteAllOpenFlareAccessLogs(ctx context.Context) (int64, error) {
conn := db.DB(ctx)
if conn == nil {
return 0, errors.New(errDatabaseNotInitialized)
}
result := conn.Where("1 = 1").Delete(&OpenFlareAccessLog{})
if result.Error != nil {
if isMissingTableError(result.Error) {
return 0, nil
}
return 0, result.Error
}
return result.RowsAffected, nil
}
// DeleteOpenFlareAccessLogsBefore deletes access logs older than cutoff.
func DeleteOpenFlareAccessLogsBefore(ctx context.Context, cutoff time.Time) (int64, error) {
conn := db.DB(ctx)
if conn == nil {
return 0, errors.New(errDatabaseNotInitialized)
}
result := conn.Where("logged_at < ?", cutoff).Delete(&OpenFlareAccessLog{})
if result.Error != nil {
if isMissingTableError(result.Error) {
return 0, nil
}
return 0, result.Error
}
return result.RowsAffected, nil
}
func buildOpenFlareAccessLogBucketRows(ctx context.Context, query OpenFlareAccessLogBucketQuery) ([]*OpenFlareAccessLogBucketRow, error) {
conn := db.DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
filter := openFlareAccessLogQueryFromBucket(query)
clause, args := buildOpenFlareAccessLogFilterClause(filter)
bucketSeconds := int64(query.FoldMinutes * 60)
if bucketSeconds <= 0 {
bucketSeconds = 180
}
bucketExpr := openFlareAccessLogBucketEpochExpr(openFlareAccessLogDialect(conn), bucketSeconds)
type bucketAccumulator struct {
requestCount int64
uniqueIPs map[string]struct{}
uniqueHosts map[string]struct{}
successCount int64
clientErrorCount int64
serverErrorCount int64
}
accumulators := make(map[int64]*bucketAccumulator)
var partials []openFlareAccessLogBucketAggregateRow
sql := fmt.Sprintf(`
SELECT
%s AS bucket_epoch,
COUNT(*) AS request_count,
SUM(CASE WHEN status_code < 400 THEN 1 ELSE 0 END) AS success_count,
SUM(CASE WHEN status_code >= 400 AND status_code < 500 THEN 1 ELSE 0 END) AS client_error_count,
SUM(CASE WHEN status_code >= 500 THEN 1 ELSE 0 END) AS server_error_count
FROM %s
WHERE %s
GROUP BY bucket_epoch`, bucketExpr, openFlareAccessLogTable, clause)
if err := conn.Raw(sql, args...).Scan(&partials).Error; err != nil {
if isMissingTableError(err) {
return []*OpenFlareAccessLogBucketRow{}, nil
}
return nil, err
}
for _, partial := range partials {
accumulator := accumulators[partial.BucketEpoch]
if accumulator == nil {
accumulator = &bucketAccumulator{
uniqueIPs: make(map[string]struct{}),
uniqueHosts: make(map[string]struct{}),
}
accumulators[partial.BucketEpoch] = accumulator
}
accumulator.requestCount += partial.RequestCount
accumulator.successCount += partial.SuccessCount
accumulator.clientErrorCount += partial.ClientErrorCount
accumulator.serverErrorCount += partial.ServerErrorCount
}
for _, column := range []string{"remote_addr", "host"} {
dimensions, err := queryOpenFlareAccessLogBucketDimensionRows(conn, clause, args, column, bucketExpr)
if err != nil {
return nil, err
}
for _, item := range dimensions {
accumulator := accumulators[item.BucketEpoch]
if accumulator == nil {
accumulator = &bucketAccumulator{
uniqueIPs: make(map[string]struct{}),
uniqueHosts: make(map[string]struct{}),
}
accumulators[item.BucketEpoch] = accumulator
}
trimmed := strings.TrimSpace(item.Value)
if trimmed == "" {
continue
}
switch column {
case "remote_addr":
accumulator.uniqueIPs[trimmed] = struct{}{}
case "host":
accumulator.uniqueHosts[trimmed] = struct{}{}
}
}
}
rows := make([]*OpenFlareAccessLogBucketRow, 0, len(accumulators))
for bucketEpoch, accumulator := range accumulators {
rows = append(rows, &OpenFlareAccessLogBucketRow{
BucketEpoch: bucketEpoch,
RequestCount: accumulator.requestCount,
UniqueIPCount: int64(len(accumulator.uniqueIPs)),
UniqueHostCount: int64(len(accumulator.uniqueHosts)),
SuccessCount: accumulator.successCount,
ClientErrorCount: accumulator.clientErrorCount,
ServerErrorCount: accumulator.serverErrorCount,
})
}
sortOpenFlareAccessLogBucketRows(rows, query.SortBy, query.SortOrder)
return rows, nil
}
func queryOpenFlareAccessLogBucketDimensionRows(conn *gorm.DB, clause string, args []any, column string, bucketExpr string) ([]openFlareAccessLogBucketDimensionRow, error) {
var rows []openFlareAccessLogBucketDimensionRow
sql := fmt.Sprintf(`
SELECT
%s AS bucket_epoch,
TRIM(%s) AS value
FROM %s
WHERE %s AND TRIM(%s) <> ''
GROUP BY bucket_epoch, TRIM(%s)`, bucketExpr, column, openFlareAccessLogTable, clause, column, column)
if err := conn.Raw(sql, args...).Scan(&rows).Error; err != nil {
if isMissingTableError(err) {
return []openFlareAccessLogBucketDimensionRow{}, nil
}
return nil, err
}
return rows, nil
}
func buildOpenFlareAccessLogBucketIPRows(ctx context.Context, query OpenFlareAccessLogBucketIPQuery) ([]*OpenFlareAccessLogBucketIPRow, error) {
if query.BucketStartedAt.IsZero() {
return []*OpenFlareAccessLogBucketIPRow{}, nil
}
foldMinutes := query.FoldMinutes
if foldMinutes <= 0 {
foldMinutes = 3
}
bucketStartedAt := query.BucketStartedAt.UTC()
filter := OpenFlareAccessLogQuery{
NodeID: query.NodeID,
RemoteAddr: query.RemoteAddr,
Host: query.Host,
Path: query.Path,
Since: bucketStartedAt,
Until: bucketStartedAt.Add(time.Duration(foldMinutes) * time.Minute),
}
rows, err := queryOpenFlareAccessLogIPAggregateRows(ctx, filter, false)
if err != nil {
return nil, err
}
sortOpenFlareAccessLogBucketIPRows(rows, query.SortBy, query.SortOrder)
return rows, nil
}
func buildOpenFlareAccessLogIPSummaryRows(ctx context.Context, query OpenFlareAccessLogIPSummaryQuery, recentSince time.Time) ([]*OpenFlareAccessLogIPSummaryRow, error) {
conn := db.DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
filter := OpenFlareAccessLogQuery{
NodeID: query.NodeID,
RemoteAddr: query.RemoteAddr,
Host: query.Host,
Since: query.Since,
}
clause, args := buildOpenFlareAccessLogFilterClause(filter)
lastSeenExpr := openFlareAccessLogEpochExpr(openFlareAccessLogDialect(conn))
recentClause := "0"
queryArgs := make([]any, 0, len(args)+1)
if !recentSince.IsZero() {
recentClause = "CASE WHEN logged_at >= ? THEN 1 ELSE 0 END"
queryArgs = append(queryArgs, recentSince)
}
queryArgs = append(queryArgs, args...)
sql := fmt.Sprintf(`
SELECT
TRIM(remote_addr) AS remote_addr,
COUNT(*) AS total_requests,
SUM(%s) AS recent_requests,
MAX(%s) AS last_seen_epoch
FROM %s
WHERE %s AND TRIM(remote_addr) <> ''
GROUP BY TRIM(remote_addr)`, recentClause, lastSeenExpr, openFlareAccessLogTable, clause)
var partials []openFlareAccessLogIPSummaryRow
if err := conn.Raw(sql, queryArgs...).Scan(&partials).Error; err != nil {
if isMissingTableError(err) {
return []*OpenFlareAccessLogIPSummaryRow{}, nil
}
return nil, err
}
rows := make([]*OpenFlareAccessLogIPSummaryRow, 0, len(partials))
for _, partial := range partials {
remoteAddr := strings.TrimSpace(partial.RemoteAddr)
if remoteAddr == "" {
continue
}
rows = append(rows, &OpenFlareAccessLogIPSummaryRow{
RemoteAddr: remoteAddr,
TotalRequests: partial.TotalRequests,
RecentRequests: partial.RecentRequests,
LastSeenEpoch: partial.LastSeenEpoch,
})
}
sortOpenFlareAccessLogIPSummaryRows(rows, query.SortBy, query.SortOrder)
return rows, nil
}
func queryOpenFlareAccessLogIPAggregateRows(ctx context.Context, filter OpenFlareAccessLogQuery, exactRemoteAddr bool) ([]*OpenFlareAccessLogBucketIPRow, error) {
conn := db.DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
clause, args := buildOpenFlareAccessLogFilterClause(filter)
lastSeenExpr := openFlareAccessLogEpochExpr(openFlareAccessLogDialect(conn))
queryClause := clause
queryArgs := append([]any{}, args...)
if exactRemoteAddr {
trimmed := strings.TrimSpace(filter.RemoteAddr)
if trimmed == "" {
return []*OpenFlareAccessLogBucketIPRow{}, nil
}
queryClause = combineOpenFlareAccessLogSQLClauses(queryClause, "TRIM(remote_addr) = ?")
queryArgs = append(queryArgs, trimmed)
}
sql := fmt.Sprintf(`
SELECT
TRIM(remote_addr) AS remote_addr,
COUNT(*) AS request_count,
SUM(CASE WHEN status_code < 400 THEN 1 ELSE 0 END) AS success_count,
SUM(CASE WHEN status_code >= 400 AND status_code < 500 THEN 1 ELSE 0 END) AS client_error_count,
SUM(CASE WHEN status_code >= 500 THEN 1 ELSE 0 END) AS server_error_count,
MAX(%s) AS last_seen_epoch
FROM %s
WHERE %s AND TRIM(remote_addr) <> ''
GROUP BY TRIM(remote_addr)`, lastSeenExpr, openFlareAccessLogTable, queryClause)
var partials []openFlareAccessLogIPAggregateRow
if err := conn.Raw(sql, queryArgs...).Scan(&partials).Error; err != nil {
if isMissingTableError(err) {
return []*OpenFlareAccessLogBucketIPRow{}, nil
}
return nil, err
}
rows := make([]*OpenFlareAccessLogBucketIPRow, 0, len(partials))
for _, partial := range partials {
remoteAddr := strings.TrimSpace(partial.RemoteAddr)
if remoteAddr == "" {
continue
}
rows = append(rows, &OpenFlareAccessLogBucketIPRow{
RemoteAddr: remoteAddr,
RequestCount: partial.RequestCount,
SuccessCount: partial.SuccessCount,
ClientErrorCount: partial.ClientErrorCount,
ServerErrorCount: partial.ServerErrorCount,
LastSeenEpoch: partial.LastSeenEpoch,
})
}
return rows, nil
}
func openFlareAccessLogQueryFromBucket(query OpenFlareAccessLogBucketQuery) OpenFlareAccessLogQuery {
return OpenFlareAccessLogQuery{
NodeID: query.NodeID,
RemoteAddr: query.RemoteAddr,
Host: query.Host,
Path: query.Path,
Since: query.Since,
}
}
func buildOpenFlareAccessLogFilterClause(query OpenFlareAccessLogQuery) (string, []any) {
parts := make([]string, 0, 6)
args := make([]any, 0, 6)
if trimmed := strings.TrimSpace(query.NodeID); trimmed != "" {
parts = append(parts, "node_id = ?")
args = append(args, trimmed)
}
if trimmed := strings.TrimSpace(query.RemoteAddr); trimmed != "" {
parts = append(parts, "remote_addr LIKE ?")
args = append(args, trimmed+"%")
}
if trimmed := strings.TrimSpace(query.Host); trimmed != "" {
parts = append(parts, "host LIKE ?")
args = append(args, trimmed+"%")
}
if trimmed := strings.TrimSpace(query.Path); trimmed != "" {
parts = append(parts, "path LIKE ?")
args = append(args, trimmed+"%")
}
if !query.Since.IsZero() {
parts = append(parts, "logged_at >= ?")
args = append(args, query.Since)
}
if !query.Until.IsZero() {
parts = append(parts, "logged_at < ?")
args = append(args, query.Until)
}
if len(parts) == 0 {
return "TRUE", nil
}
return strings.Join(parts, " AND "), args
}
func applyOpenFlareAccessLogFilters(tx *gorm.DB, query OpenFlareAccessLogQuery) *gorm.DB {
clause, args := buildOpenFlareAccessLogFilterClause(query)
if clause == "TRUE" {
return tx
}
return tx.Where(clause, args...)
}
func countOpenFlareAccessLogRecords(conn *gorm.DB, query OpenFlareAccessLogQuery) (int64, error) {
var count int64
if err := applyOpenFlareAccessLogFilters(conn.Model(&OpenFlareAccessLog{}), query).Count(&count).Error; err != nil {
return 0, err
}
return count, nil
}
func countDistinctOpenFlareAccessLogIPs(conn *gorm.DB, query OpenFlareAccessLogQuery) (int64, error) {
clause, args := buildOpenFlareAccessLogFilterClause(query)
sql := fmt.Sprintf(`
SELECT COUNT(*) FROM (
SELECT TRIM(remote_addr) AS remote_addr
FROM %s
WHERE %s AND remote_addr <> ''
GROUP BY TRIM(remote_addr)
) AS ips`, openFlareAccessLogTable, clause)
var total int64
if err := conn.Raw(sql, args...).Scan(&total).Error; err != nil {
return 0, err
}
return total, nil
}
func openFlareAccessLogDialect(conn *gorm.DB) string {
if conn == nil || conn.Dialector == nil {
return "sqlite"
}
switch conn.Dialector.Name() {
case "postgres":
return "postgres"
default:
return "sqlite"
}
}
func openFlareAccessLogBucketEpochExpr(dialect string, bucketSeconds int64) string {
switch dialect {
case "postgres":
return fmt.Sprintf("FLOOR(EXTRACT(EPOCH FROM logged_at AT TIME ZONE 'UTC') / %d) * %d", bucketSeconds, bucketSeconds)
default:
return fmt.Sprintf("(CAST(strftime('%%s', logged_at) AS INTEGER) / %d) * %d", bucketSeconds, bucketSeconds)
}
}
func openFlareAccessLogEpochExpr(dialect string) string {
switch dialect {
case "postgres":
return "FLOOR(EXTRACT(EPOCH FROM logged_at AT TIME ZONE 'UTC'))::bigint"
default:
return "CAST((julianday(logged_at) - 2440587.5) * 86400 AS INTEGER)"
}
}
func combineOpenFlareAccessLogSQLClauses(left string, right string) string {
if strings.TrimSpace(left) == "" || left == "TRUE" {
return right
}
return left + " AND " + right
}
func openFlareAccessLogOrderClause(sortBy string, sortOrder string) string {
direction := "DESC"
if openFlareAccessLogNormalizeSortOrder(sortOrder) == "asc" {
direction = "ASC"
}
column := "logged_at"
switch strings.TrimSpace(sortBy) {
case "status_code":
column = "status_code"
case "remote_addr":
column = "remote_addr"
case "host":
column = "host"
case "path":
column = "path"
}
if column == "logged_at" {
return column + " " + direction + ", id " + direction
}
return column + " " + direction + ", logged_at " + direction + ", id " + direction
}
func sortOpenFlareAccessLogBucketIPRows(items []*OpenFlareAccessLogBucketIPRow, sortBy string, sortOrder string) {
desc := openFlareAccessLogNormalizeSortOrder(sortOrder) != "asc"
sort.Slice(items, func(i int, j int) bool {
left := items[i]
right := items[j]
if left == nil || right == nil {
return left != nil
}
var compare int
switch strings.TrimSpace(sortBy) {
case "last_seen_at":
compare = openFlareAccessLogCompareInt64(left.LastSeenEpoch, right.LastSeenEpoch)
case "remote_addr":
compare = strings.Compare(left.RemoteAddr, right.RemoteAddr)
default:
compare = openFlareAccessLogCompareInt64(left.RequestCount, right.RequestCount)
}
if compare == 0 {
compare = openFlareAccessLogCompareInt64(left.LastSeenEpoch, right.LastSeenEpoch)
}
if compare == 0 {
compare = strings.Compare(left.RemoteAddr, right.RemoteAddr)
}
if desc {
return compare > 0
}
return compare < 0
})
}
func sortOpenFlareAccessLogBucketRows(items []*OpenFlareAccessLogBucketRow, sortBy string, sortOrder string) {
desc := openFlareAccessLogNormalizeSortOrder(sortOrder) != "asc"
sort.Slice(items, func(i int, j int) bool {
left := items[i]
right := items[j]
if left == nil || right == nil {
return left != nil
}
var compare int
switch strings.TrimSpace(sortBy) {
case "request_count":
compare = openFlareAccessLogCompareInt64(left.RequestCount, right.RequestCount)
default:
compare = openFlareAccessLogCompareInt64(left.BucketEpoch, right.BucketEpoch)
}
if compare == 0 {
compare = openFlareAccessLogCompareInt64(left.BucketEpoch, right.BucketEpoch)
}
if desc {
return compare > 0
}
return compare < 0
})
}
func sortOpenFlareAccessLogIPSummaryRows(items []*OpenFlareAccessLogIPSummaryRow, sortBy string, sortOrder string) {
desc := openFlareAccessLogNormalizeSortOrder(sortOrder) != "asc"
sort.Slice(items, func(i int, 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 "recent_requests":
compare = openFlareAccessLogCompareInt64(left.RecentRequests, right.RecentRequests)
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 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), "asc") {
return "asc"
}
return "desc"
}
func openFlareAccessLogCompareInt64(left int64, right int64) int {
switch {
case left > right:
return 1
case left < right:
return -1
default:
return 0
}
}
@@ -0,0 +1,144 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package model
import (
"context"
"fmt"
"testing"
"time"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/glebarez/sqlite"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"gorm.io/gorm"
)
func setupOpenFlareAccessLogTestEnvironment(t *testing.T) (context.Context, func()) {
t.Helper()
sqliteDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{
DisableForeignKeyConstraintWhenMigrating: true,
})
require.NoError(t, err)
require.NoError(t, sqliteDB.AutoMigrate(&OpenFlareAccessLog{}))
db.SetDB(sqliteDB)
return context.Background(), func() {
db.SetDB(nil)
}
}
func seedOpenFlareAccessLogs(t *testing.T, ctx context.Context, now time.Time) {
t.Helper()
records := []*OpenFlareAccessLog{
{NodeID: "node-a", LoggedAt: now.Add(-5 * time.Minute), RemoteAddr: "1.1.1.1", Region: "US", Host: "a.example.com", Path: "/alpha", StatusCode: 200},
{NodeID: "node-a", LoggedAt: now.Add(-4 * time.Minute), RemoteAddr: "2.2.2.2", Region: "US", Host: "a.example.com", Path: "/beta", StatusCode: 404},
{NodeID: "node-b", LoggedAt: now.Add(-3 * time.Minute), RemoteAddr: "1.1.1.1", Region: "EU", Host: "b.example.com", Path: "/gamma", StatusCode: 502},
{NodeID: "node-b", LoggedAt: now.Add(-2 * time.Minute), RemoteAddr: "3.3.3.3", Region: "EU", Host: "b.example.com", Path: "/delta", StatusCode: 200},
{NodeID: "node-b", LoggedAt: now.Add(-1 * time.Minute), RemoteAddr: "", Region: "", Host: "b.example.com", Path: "/empty-ip", StatusCode: 200},
}
for index, record := range records {
require.NoError(t, db.DB(ctx).Create(record).Error, "seed access log %d", index)
}
}
func TestListOpenFlareAccessLogsPaginated(t *testing.T) {
ctx, cleanup := setupOpenFlareAccessLogTestEnvironment(t)
defer cleanup()
now := time.Now().UTC()
for index := range 15 {
record := &OpenFlareAccessLog{
NodeID: "node-page",
LoggedAt: now.Add(-time.Duration(index) * time.Minute),
RemoteAddr: fmt.Sprintf("203.0.113.%d", (index%5)+1),
Host: "example.com",
Path: fmt.Sprintf("/path-%02d", index),
StatusCode: 200,
}
require.NoError(t, db.DB(ctx).Create(record).Error)
}
query := OpenFlareAccessLogQuery{
NodeID: "node-page",
Since: now.Add(-24 * time.Hour),
Page: 1,
PageSize: 5,
SortBy: "logged_at",
SortOrder: "desc",
}
page, err := ListOpenFlareAccessLogs(ctx, query)
require.NoError(t, err)
require.Len(t, page, 5)
assert.Equal(t, "/path-05", page[0].Path)
assert.Equal(t, "/path-09", page[4].Path)
}
func TestCountOpenFlareAccessLogs(t *testing.T) {
ctx, cleanup := setupOpenFlareAccessLogTestEnvironment(t)
defer cleanup()
now := time.Now().UTC()
seedOpenFlareAccessLogs(t, ctx, now)
query := OpenFlareAccessLogQuery{
Since: now.Add(-10 * time.Minute),
}
totalRecords, totalIPs, err := CountOpenFlareAccessLogs(ctx, query)
require.NoError(t, err)
assert.Equal(t, int64(5), totalRecords)
assert.Equal(t, int64(3), totalIPs)
}
func TestListOpenFlareAccessLogsFiltersAndSort(t *testing.T) {
ctx, cleanup := setupOpenFlareAccessLogTestEnvironment(t)
defer cleanup()
now := time.Now().UTC()
seedOpenFlareAccessLogs(t, ctx, now)
query := OpenFlareAccessLogQuery{
NodeID: "node-a",
Since: now.Add(-10 * time.Minute),
SortBy: "status_code",
SortOrder: "desc",
}
rows, err := ListOpenFlareAccessLogs(ctx, query)
require.NoError(t, err)
require.Len(t, rows, 2)
assert.Equal(t, 404, rows[0].StatusCode)
assert.Equal(t, 200, rows[1].StatusCode)
}
func TestListOpenFlareAccessLogsMissingTableGraceful(t *testing.T) {
ctx, cleanup := setupOpenFlareAccessLogTestEnvironment(t)
defer cleanup()
require.NoError(t, db.DB(ctx).Migrator().DropTable(&OpenFlareAccessLog{}))
query := OpenFlareAccessLogQuery{Since: time.Now().UTC().Add(-time.Hour)}
rows, err := ListOpenFlareAccessLogs(ctx, query)
require.NoError(t, err)
assert.Empty(t, rows)
totalRecords, totalIPs, err := CountOpenFlareAccessLogs(ctx, query)
require.NoError(t, err)
assert.Zero(t, totalRecords)
assert.Zero(t, totalIPs)
}
func TestDeleteOpenFlareAccessLogsBefore(t *testing.T) {
ctx, cleanup := setupOpenFlareAccessLogTestEnvironment(t)
defer cleanup()
now := time.Now().UTC()
seedOpenFlareAccessLogs(t, ctx, now)
deleted, err := DeleteOpenFlareAccessLogsBefore(ctx, now.Add(-2*time.Minute))
require.NoError(t, err)
assert.Equal(t, int64(3), deleted)
totalRecords, _, err := CountOpenFlareAccessLogs(ctx, OpenFlareAccessLogQuery{Since: now.Add(-10 * time.Minute)})
require.NoError(t, err)
assert.Equal(t, int64(2), totalRecords)
}
@@ -0,0 +1,82 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package model
import (
"context"
"errors"
"time"
"github.com/Rain-kl/Wavelet/internal/db"
"gorm.io/gorm"
)
// AcmeAccount OpenFlare ACME 账号实体。
type AcmeAccount struct {
ID uint `json:"id" gorm:"primaryKey;autoIncrement"`
Email string `json:"email" gorm:"size:255"`
URL string `json:"url" gorm:"size:255"`
PrivateKey string `json:"-" gorm:"type:text;not null"`
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime"`
UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime"`
}
// TableName 表名。
func (AcmeAccount) TableName() string {
return "of_acme_accounts"
}
// GetAcmeAccountByID 按 ID 查询 ACME 账号。
func GetAcmeAccountByID(ctx context.Context, id uint) (*AcmeAccount, error) {
conn := db.DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
var account AcmeAccount
if err := conn.First(&account, id).Error; err != nil {
return nil, err
}
return &account, nil
}
// CreateAcmeAccountRecord 创建 ACME 账号。
func CreateAcmeAccountRecord(ctx context.Context, account *AcmeAccount) error {
conn := db.DB(ctx)
if conn == nil {
return errors.New(errDatabaseNotInitialized)
}
return conn.Create(account).Error
}
// SaveAcmeAccount 保存 ACME 账号。
func SaveAcmeAccount(ctx context.Context, account *AcmeAccount) error {
conn := db.DB(ctx)
if conn == nil {
return errors.New(errDatabaseNotInitialized)
}
return conn.Save(account).Error
}
// GetDefaultAcmeAccount 获取默认 ACME 账号,不存在时创建占位记录。
func GetDefaultAcmeAccount(ctx context.Context) (*AcmeAccount, error) {
conn := db.DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
var account AcmeAccount
err := conn.Order("id asc").First(&account).Error
if err == nil {
return &account, nil
}
if !errors.Is(err, gorm.ErrRecordNotFound) {
return nil, err
}
account = AcmeAccount{
Email: "admin@openflare.dev",
}
if err = conn.Create(&account).Error; err != nil {
return nil, err
}
return &account, nil
}
@@ -0,0 +1,132 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package model
import (
"context"
"errors"
"time"
"github.com/Rain-kl/Wavelet/internal/db"
"gorm.io/gorm"
)
// OpenFlareApplyLogQuery filters apply logs for list queries.
type OpenFlareApplyLogQuery struct {
NodeID string
PageNo int
PageSize int
}
// OpenFlareApplyLog stores node configuration apply results.
type OpenFlareApplyLog struct {
ID uint `json:"id" gorm:"primaryKey;autoIncrement"`
NodeID string `json:"node_id" gorm:"index;size:64;not null"`
Version string `json:"version" gorm:"size:32;not null"`
Result string `json:"result" gorm:"size:32;not null"`
Message string `json:"message" gorm:"type:text"`
Checksum string `json:"checksum" gorm:"size:64;not null;default:''"`
MainConfigChecksum string `json:"main_config_checksum" gorm:"size:64;not null;default:''"`
RouteConfigChecksum string `json:"route_config_checksum" gorm:"size:64;not null;default:''"`
SupportFileCount int `json:"support_file_count" gorm:"not null;default:0"`
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime;index"`
}
// TableName returns the GORM table name.
func (OpenFlareApplyLog) TableName() string {
return "of_apply_logs"
}
// ListOpenFlareApplyLogs returns apply logs ordered by id desc with optional pagination.
func ListOpenFlareApplyLogs(ctx context.Context, query OpenFlareApplyLogQuery) ([]*OpenFlareApplyLog, error) {
conn := db.DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
dbQuery := conn.Model(&OpenFlareApplyLog{}).Order("id desc")
if query.NodeID != "" {
dbQuery = dbQuery.Where("node_id = ?", query.NodeID)
}
if query.PageSize > 0 {
offset := 0
if query.PageNo > 1 {
offset = (query.PageNo - 1) * query.PageSize
}
dbQuery = dbQuery.Limit(query.PageSize).Offset(offset)
}
var logs []*OpenFlareApplyLog
if err := dbQuery.Find(&logs).Error; err != nil {
return nil, err
}
return logs, nil
}
// CountOpenFlareApplyLogs returns total apply logs, optionally filtered by node_id.
func CountOpenFlareApplyLogs(ctx context.Context, nodeID string) (int64, error) {
conn := db.DB(ctx)
if conn == nil {
return 0, errors.New(errDatabaseNotInitialized)
}
query := conn.Model(&OpenFlareApplyLog{})
if nodeID != "" {
query = query.Where("node_id = ?", nodeID)
}
var total int64
if err := query.Count(&total).Error; err != nil {
return 0, err
}
return total, nil
}
// GetLatestOpenFlareApplyLogsByNodeIDs returns the latest apply log per node id.
func GetLatestOpenFlareApplyLogsByNodeIDs(ctx context.Context, nodeIDs []string) (map[string]*OpenFlareApplyLog, error) {
result := make(map[string]*OpenFlareApplyLog)
if len(nodeIDs) == 0 {
return result, nil
}
conn := db.DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
var logs []*OpenFlareApplyLog
subQuery := conn.Model(&OpenFlareApplyLog{}).
Select("MAX(id) AS id").
Where("node_id IN ?", nodeIDs).
Group("node_id")
if err := conn.Where("id IN (?)", subQuery).Find(&logs).Error; err != nil {
return nil, err
}
for _, log := range logs {
result[log.NodeID] = log
}
return result, nil
}
// DeleteAllOpenFlareApplyLogs removes every apply log record.
func DeleteAllOpenFlareApplyLogs(ctx context.Context) (int64, error) {
conn := db.DB(ctx)
if conn == nil {
return 0, errors.New(errDatabaseNotInitialized)
}
result := conn.Session(&gorm.Session{AllowGlobalUpdate: true}).Delete(&OpenFlareApplyLog{})
return result.RowsAffected, result.Error
}
// DeleteOpenFlareApplyLogsBefore removes apply logs created before the cutoff time.
func DeleteOpenFlareApplyLogsBefore(ctx context.Context, before time.Time) (int64, error) {
conn := db.DB(ctx)
if conn == nil {
return 0, errors.New(errDatabaseNotInitialized)
}
result := conn.Where("created_at < ?", before).Delete(&OpenFlareApplyLog{})
return result.RowsAffected, result.Error
}
@@ -0,0 +1,163 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package model
import (
"context"
"errors"
"time"
"github.com/Rain-kl/Wavelet/internal/db"
"gorm.io/gorm"
)
// ConfigVersionSummary is the list view for config versions.
type ConfigVersionSummary struct {
ID uint `json:"id"`
Version string `json:"version"`
Checksum string `json:"checksum"`
IsActive bool `json:"is_active"`
CreatedBy string `json:"created_by"`
CreatedAt time.Time `json:"created_at"`
}
// ConfigVersion stores a published OpenResty configuration snapshot.
type ConfigVersion struct {
ID uint `json:"id" gorm:"primaryKey;autoIncrement"`
Version string `json:"version" gorm:"uniqueIndex;size:32;not null"`
SnapshotJSON string `json:"snapshot_json" gorm:"type:text;not null"`
MainConfig string `json:"main_config" gorm:"type:text;not null;default:''"`
RenderedConfig string `json:"rendered_config" gorm:"type:text;not null"`
SupportFilesJSON string `json:"support_files_json" gorm:"type:text;not null;default:'[]'"`
Checksum string `json:"checksum" gorm:"size:64;not null"`
IsActive bool `json:"is_active" gorm:"not null;default:false;index"`
CreatedBy string `json:"created_by" gorm:"size:64;not null"`
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime"`
}
// TableName returns the GORM table name.
func (ConfigVersion) TableName() string {
return "of_config_versions"
}
// ListConfigVersionSummaries returns config version summaries ordered by id desc.
func ListConfigVersionSummaries(ctx context.Context) ([]*ConfigVersionSummary, error) {
conn := db.DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
var versions []*ConfigVersionSummary
err := conn.Model(&ConfigVersion{}).
Select("id", "version", "checksum", "is_active", "created_by", "created_at").
Order("id desc").
Find(&versions).Error
return versions, err
}
// GetConfigVersionByID returns a config version by primary key.
func GetConfigVersionByID(ctx context.Context, id uint) (*ConfigVersion, error) {
conn := db.DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
var version ConfigVersion
if err := conn.First(&version, id).Error; err != nil {
return nil, err
}
return &version, nil
}
// GetActiveConfigVersion returns the currently active config version.
func GetActiveConfigVersion(ctx context.Context) (*ConfigVersion, error) {
conn := db.DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
var version ConfigVersion
if err := conn.Where("is_active = ?", true).Order("id desc").First(&version).Error; err != nil {
return nil, err
}
return &version, nil
}
// GetLatestConfigVersionByPrefix returns the latest version string matching a date prefix.
func GetLatestConfigVersionByPrefix(ctx context.Context, prefix string) (string, error) {
conn := db.DB(ctx)
if conn == nil {
return "", errors.New(errDatabaseNotInitialized)
}
var version ConfigVersion
err := conn.Model(&ConfigVersion{}).
Select("version").
Where("version LIKE ?", prefix+"-%").
Order("version desc").
First(&version).Error
if err != nil {
return "", err
}
return version.Version, nil
}
// CreateConfigVersion inserts a new config version record.
func CreateConfigVersion(ctx context.Context, version *ConfigVersion) error {
conn := db.DB(ctx)
if conn == nil {
return errors.New(errDatabaseNotInitialized)
}
return conn.Create(version).Error
}
// PublishConfigVersionTx deactivates all versions and creates a new active version.
func PublishConfigVersionTx(ctx context.Context, version *ConfigVersion) error {
conn := db.DB(ctx)
if conn == nil {
return errors.New(errDatabaseNotInitialized)
}
return conn.Transaction(func(tx *gorm.DB) error {
if err := tx.Model(&ConfigVersion{}).Where("is_active = ?", true).Update("is_active", false).Error; err != nil {
return err
}
return tx.Create(version).Error
})
}
// ActivateConfigVersionTx marks the given version active and deactivates others.
func ActivateConfigVersionTx(ctx context.Context, id uint) error {
conn := db.DB(ctx)
if conn == nil {
return errors.New(errDatabaseNotInitialized)
}
return conn.Transaction(func(tx *gorm.DB) error {
if err := tx.Model(&ConfigVersion{}).Where("is_active = ?", true).Update("is_active", false).Error; err != nil {
return err
}
return tx.Model(&ConfigVersion{}).Where("id = ?", id).Update("is_active", true).Error
})
}
// DeleteConfigVersionsByIDs removes config versions by ids.
func DeleteConfigVersionsByIDs(ctx context.Context, ids []uint) (int64, error) {
if len(ids) == 0 {
return 0, nil
}
conn := db.DB(ctx)
if conn == nil {
return 0, errors.New(errDatabaseNotInitialized)
}
result := conn.Where("id IN ?", ids).Delete(&ConfigVersion{})
return result.RowsAffected, result.Error
}
// ListEnabledProxyRoutes returns enabled proxy routes ordered by id asc.
func ListEnabledProxyRoutes(ctx context.Context) ([]*ProxyRoute, error) {
conn := db.DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
var routes []*ProxyRoute
if err := conn.Where("enabled = ?", true).Order("id asc").Find(&routes).Error; err != nil {
return nil, err
}
return routes, nil
}
@@ -0,0 +1,80 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package model
import (
"context"
"errors"
"time"
"github.com/Rain-kl/Wavelet/internal/db"
)
// DNSAccount OpenFlare DNS 账号实体。
type DNSAccount struct {
ID uint `json:"id" gorm:"primaryKey;autoIncrement"`
Name string `json:"name" gorm:"size:255;not null"`
Type string `json:"type" gorm:"size:64;not null"`
Authorization string `json:"-" gorm:"type:text;not null"`
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime"`
UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime"`
}
// TableName 表名。
func (DNSAccount) TableName() string {
return "of_dns_accounts"
}
// ListDNSAccounts 列出全部 DNS 账号(授权信息不通过 JSON 暴露)。
func ListDNSAccounts(ctx context.Context) ([]DNSAccount, error) {
conn := db.DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
var accounts []DNSAccount
if err := conn.Order("id desc").Find(&accounts).Error; err != nil {
return nil, err
}
return accounts, nil
}
// GetDNSAccountByID 按 ID 查询 DNS 账号。
func GetDNSAccountByID(ctx context.Context, id uint) (*DNSAccount, error) {
conn := db.DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
var account DNSAccount
if err := conn.First(&account, id).Error; err != nil {
return nil, err
}
return &account, nil
}
// CreateDNSAccountRecord 创建 DNS 账号。
func CreateDNSAccountRecord(ctx context.Context, account *DNSAccount) error {
conn := db.DB(ctx)
if conn == nil {
return errors.New(errDatabaseNotInitialized)
}
return conn.Create(account).Error
}
// SaveDNSAccount 保存 DNS 账号。
func SaveDNSAccount(ctx context.Context, account *DNSAccount) error {
conn := db.DB(ctx)
if conn == nil {
return errors.New(errDatabaseNotInitialized)
}
return conn.Save(account).Error
}
// DeleteDNSAccountRecord 删除 DNS 账号。
func DeleteDNSAccountRecord(ctx context.Context, id uint) error {
conn := db.DB(ctx)
if conn == nil {
return errors.New(errDatabaseNotInitialized)
}
return conn.Delete(&DNSAccount{}, id).Error
}
@@ -0,0 +1,94 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package model
import (
"context"
"errors"
"time"
"github.com/Rain-kl/Wavelet/internal/db"
)
// ManagedDomain OpenFlare 托管域名实体。
type ManagedDomain struct {
ID uint `json:"id" gorm:"primaryKey;autoIncrement"`
Domain string `json:"domain" gorm:"uniqueIndex;size:255;not null"`
CertID *uint `json:"cert_id"`
Enabled bool `json:"enabled" gorm:"not null;default:true"`
Remark string `json:"remark" gorm:"size:255"`
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime"`
UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime"`
}
// TableName 表名。
func (ManagedDomain) TableName() string {
return "of_managed_domains"
}
// ListManagedDomains 列出全部托管域名。
func ListManagedDomains(ctx context.Context) ([]ManagedDomain, error) {
conn := db.DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
var domains []ManagedDomain
if err := conn.Order("id desc").Find(&domains).Error; err != nil {
return nil, err
}
return domains, nil
}
// ListEnabledManagedDomainsWithCertificate 列出已启用且绑定证书的托管域名。
func ListEnabledManagedDomainsWithCertificate(ctx context.Context) ([]ManagedDomain, error) {
conn := db.DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
var domains []ManagedDomain
if err := conn.Where("enabled = ? AND cert_id IS NOT NULL", true).Order("id desc").Find(&domains).Error; err != nil {
return nil, err
}
return domains, nil
}
// GetManagedDomainByID 按 ID 查询托管域名。
func GetManagedDomainByID(ctx context.Context, id uint) (*ManagedDomain, error) {
conn := db.DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
var domain ManagedDomain
if err := conn.First(&domain, id).Error; err != nil {
return nil, err
}
return &domain, nil
}
// CreateManagedDomainRecord 创建托管域名。
func CreateManagedDomainRecord(ctx context.Context, domain *ManagedDomain) error {
conn := db.DB(ctx)
if conn == nil {
return errors.New(errDatabaseNotInitialized)
}
return conn.Create(domain).Error
}
// SaveManagedDomain 保存托管域名。
func SaveManagedDomain(ctx context.Context, domain *ManagedDomain) error {
conn := db.DB(ctx)
if conn == nil {
return errors.New(errDatabaseNotInitialized)
}
return conn.Save(domain).Error
}
// DeleteManagedDomainRecord 删除托管域名。
func DeleteManagedDomainRecord(ctx context.Context, id uint) error {
conn := db.DB(ctx)
if conn == nil {
return errors.New(errDatabaseNotInitialized)
}
return conn.Delete(&ManagedDomain{}, id).Error
}
@@ -0,0 +1,163 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package model
import (
"context"
"errors"
"time"
"github.com/Rain-kl/Wavelet/internal/db"
)
// OpenFlareNode stores an edge, relay, or tunnel client node.
type OpenFlareNode struct {
ID uint `json:"id" gorm:"primaryKey;autoIncrement"`
NodeID string `json:"node_id" gorm:"uniqueIndex;size:64;not null"`
Name string `json:"name" gorm:"size:128;not null"`
IP string `json:"ip" gorm:"size:64;not null;default:''"`
IPManualOverride bool `json:"ip_manual_override" gorm:"not null;default:false"`
GeoName string `json:"geo_name" gorm:"size:128;not null;default:''"`
GeoLatitude *float64 `json:"geo_latitude"`
GeoLongitude *float64 `json:"geo_longitude"`
GeoManualOverride bool `json:"geo_manual_override" gorm:"not null;default:false"`
AccessToken string `json:"-" gorm:"column:access_token;size:128;index"`
AutoUpdateEnabled bool `json:"auto_update_enabled" gorm:"not null;default:false"`
UpdateRequested bool `json:"update_requested" gorm:"not null;default:false"`
UpdateChannel string `json:"update_channel" gorm:"size:16;not null;default:'stable'"`
UpdateTag string `json:"update_tag" gorm:"size:64;not null;default:''"`
RestartOpenrestyRequested bool `json:"restart_openresty_requested" gorm:"not null;default:false"`
Version string `json:"version" gorm:"size:64;not null;default:''"`
ExtVersion string `json:"ext_version" gorm:"size:64;not null;default:''"`
OpenrestyStatus string `json:"openresty_status" gorm:"size:16;not null;default:'unknown'"`
OpenrestyMessage string `json:"openresty_message" gorm:"type:text"`
Status string `json:"status" gorm:"size:16;not null;default:'offline'"`
CurrentVersion string `json:"current_version" gorm:"size:32;not null;default:''"`
LastSeenAt *time.Time `json:"last_seen_at"`
LastError string `json:"last_error" gorm:"type:text"`
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime"`
UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime"`
NodeType string `json:"node_type" gorm:"size:32;not null;default:'edge_node'"`
RelayBindPort int `json:"relay_bind_port" gorm:"not null;default:0"`
RelayVhostHTTPPort int `json:"relay_vhost_http_port" gorm:"not null;default:0"`
RelayAuthToken string `json:"-" gorm:"size:128;not null;default:''"`
RelayAgentAccessAddr string `json:"relay_agent_access_addr" gorm:"size:255;not null;default:''"`
RelayClientAccessAddr string `json:"relay_client_access_addr" gorm:"size:255;not null;default:''"`
RelayClientProxyURL string `json:"relay_client_proxy_url" gorm:"size:512;not null;default:''"`
CapabilitiesJSON string `json:"capabilities_json" gorm:"type:text;not null;default:'[]'"`
RelayStatus string `json:"relay_status" gorm:"size:16;not null;default:'unknown'"`
RelayWebServerEnabled bool `json:"relay_web_server_enabled" gorm:"not null;default:false"`
}
// TableName returns the GORM table name.
func (OpenFlareNode) TableName() string {
return "of_nodes"
}
// ListOpenFlareNodes returns all nodes ordered by id desc.
func ListOpenFlareNodes(ctx context.Context) ([]OpenFlareNode, error) {
conn := db.DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
var nodes []OpenFlareNode
if err := conn.Order("id desc").Find(&nodes).Error; err != nil {
return nil, err
}
return nodes, nil
}
// ListOpenFlareNodesByNodeIDs returns nodes matching the given node ids.
func ListOpenFlareNodesByNodeIDs(ctx context.Context, nodeIDs []string) ([]OpenFlareNode, error) {
if len(nodeIDs) == 0 {
return []OpenFlareNode{}, nil
}
conn := db.DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
var nodes []OpenFlareNode
if err := conn.Where("node_id IN ?", nodeIDs).Find(&nodes).Error; err != nil {
return nil, err
}
return nodes, nil
}
// GetOpenFlareNodeByID returns a node by primary key.
func GetOpenFlareNodeByID(ctx context.Context, id uint) (*OpenFlareNode, error) {
conn := db.DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
var node OpenFlareNode
if err := conn.First(&node, id).Error; err != nil {
return nil, err
}
return &node, nil
}
// GetOpenFlareNodeByNodeID returns a node by node_id.
func GetOpenFlareNodeByNodeID(ctx context.Context, nodeID string) (*OpenFlareNode, error) {
conn := db.DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
var node OpenFlareNode
if err := conn.Where("node_id = ?", nodeID).First(&node).Error; err != nil {
return nil, err
}
return &node, nil
}
// GetOpenFlareNodeByAccessToken returns a node by access token.
func GetOpenFlareNodeByAccessToken(ctx context.Context, token string) (*OpenFlareNode, error) {
conn := db.DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
var node OpenFlareNode
if err := conn.Where("access_token = ?", token).First(&node).Error; err != nil {
return nil, err
}
return &node, nil
}
// CreateOpenFlareNode inserts a new node.
func CreateOpenFlareNode(ctx context.Context, node *OpenFlareNode) error {
conn := db.DB(ctx)
if conn == nil {
return errors.New(errDatabaseNotInitialized)
}
return conn.Create(node).Error
}
// SaveOpenFlareNode persists node changes.
func SaveOpenFlareNode(ctx context.Context, node *OpenFlareNode) error {
conn := db.DB(ctx)
if conn == nil {
return errors.New(errDatabaseNotInitialized)
}
return conn.Save(node).Error
}
// UpdateOpenFlareNodeFields updates selected columns for a node.
func UpdateOpenFlareNodeFields(ctx context.Context, node *OpenFlareNode, fields ...string) error {
conn := db.DB(ctx)
if conn == nil {
return errors.New(errDatabaseNotInitialized)
}
if len(fields) == 0 {
return conn.Save(node).Error
}
return conn.Model(node).Select(fields).Updates(node).Error
}
// DeleteOpenFlareNode removes a node by primary key.
func DeleteOpenFlareNode(ctx context.Context, id uint) error {
conn := db.DB(ctx)
if conn == nil {
return errors.New(errDatabaseNotInitialized)
}
return conn.Delete(&OpenFlareNode{}, id).Error
}
@@ -0,0 +1,550 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package model
import (
"context"
"errors"
"strings"
"time"
"github.com/Rain-kl/Wavelet/internal/db"
"gorm.io/gorm"
)
// OpenFlareMetricSnapshot stores a node capacity snapshot (v1 single table, no sharding).
type OpenFlareMetricSnapshot struct {
ID uint `json:"id" gorm:"primaryKey;autoIncrement"`
NodeID string `json:"node_id" gorm:"index;size:64;not null"`
CapturedAt time.Time `json:"captured_at" gorm:"index"`
CPUUsagePercent float64 `json:"cpu_usage_percent"`
MemoryUsedBytes int64 `json:"memory_used_bytes"`
MemoryTotalBytes int64 `json:"memory_total_bytes"`
StorageUsedBytes int64 `json:"storage_used_bytes"`
StorageTotalBytes int64 `json:"storage_total_bytes"`
DiskReadBytes int64 `json:"disk_read_bytes"`
DiskWriteBytes int64 `json:"disk_write_bytes"`
NetworkRxBytes int64 `json:"network_rx_bytes"`
NetworkTxBytes int64 `json:"network_tx_bytes"`
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime"`
}
// TableName returns the GORM table name.
func (OpenFlareMetricSnapshot) TableName() string {
return "of_node_metric_snapshots"
}
// OpenFlareRequestReport stores aggregated traffic windows per node.
type OpenFlareRequestReport struct {
ID uint `json:"id" gorm:"primaryKey;autoIncrement"`
NodeID string `json:"node_id" gorm:"index;size:64;not null"`
WindowStartedAt time.Time `json:"window_started_at" gorm:"index"`
WindowEndedAt time.Time `json:"window_ended_at" gorm:"index"`
RequestCount int64 `json:"request_count"`
ErrorCount int64 `json:"error_count"`
UniqueVisitorCount int64 `json:"unique_visitor_count"`
StatusCodesJSON string `json:"status_codes_json" gorm:"type:text"`
TopDomainsJSON string `json:"top_domains_json" gorm:"type:text"`
SourceCountriesJSON string `json:"source_countries_json" gorm:"type:text"`
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime"`
}
// TableName returns the GORM table name.
func (OpenFlareRequestReport) TableName() string {
return "of_node_request_reports"
}
// OpenFlareAccessLog stores a single access log row (v1 single table, no sharding).
type OpenFlareAccessLog struct {
ID uint `json:"id" gorm:"primaryKey;autoIncrement"`
NodeID string `json:"node_id" gorm:"index;size:64;not null"`
LoggedAt time.Time `json:"logged_at" gorm:"index"`
RemoteAddr string `json:"remote_addr" gorm:"index;size:128"`
Region string `json:"region" gorm:"size:128"`
Host string `json:"host" gorm:"index;size:255"`
Path string `json:"path" gorm:"size:2048"`
StatusCode int `json:"status_code" gorm:"index"`
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime"`
}
// TableName returns the GORM table name.
func (OpenFlareAccessLog) TableName() string {
return "of_node_access_logs"
}
// OpenFlareAccessLogRegionCount aggregates access log regions.
type OpenFlareAccessLogRegionCount struct {
Region string `json:"region"`
Count int64 `json:"count"`
}
// OpenFlareHealthEvent stores node health alert events.
type OpenFlareHealthEvent struct {
ID uint `json:"id" gorm:"primaryKey;autoIncrement"`
NodeID string `json:"node_id" gorm:"index;size:64;not null"`
EventType string `json:"event_type" gorm:"index;size:64;not null"`
Severity string `json:"severity" gorm:"size:16;not null"`
Status string `json:"status" gorm:"index;size:16;not null"`
Message string `json:"message" gorm:"type:text"`
FirstTriggeredAt time.Time `json:"first_triggered_at" gorm:"index"`
LastTriggeredAt time.Time `json:"last_triggered_at" gorm:"index"`
ReportedAt time.Time `json:"reported_at" gorm:"index"`
ResolvedAt *time.Time `json:"resolved_at" gorm:"index"`
MetadataJSON string `json:"metadata_json" gorm:"type:text"`
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime"`
UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime"`
}
// TableName returns the GORM table name.
func (OpenFlareHealthEvent) TableName() string {
return "of_node_health_events"
}
// OpenFlareNodeSystemProfile stores the latest node system profile.
type OpenFlareNodeSystemProfile struct {
ID uint `json:"id" gorm:"primaryKey;autoIncrement"`
NodeID string `json:"node_id" gorm:"uniqueIndex;size:64;not null"`
Hostname string `json:"hostname" gorm:"size:255"`
OSName string `json:"os_name" gorm:"size:128"`
OSVersion string `json:"os_version" gorm:"size:128"`
KernelVersion string `json:"kernel_version" gorm:"size:128"`
Architecture string `json:"architecture" gorm:"size:64"`
CPUModel string `json:"cpu_model" gorm:"size:255"`
CPUCores int `json:"cpu_cores"`
TotalMemoryBytes int64 `json:"total_memory_bytes"`
TotalDiskBytes int64 `json:"total_disk_bytes"`
UptimeSeconds int64 `json:"uptime_seconds"`
ReportedAt time.Time `json:"reported_at" gorm:"index"`
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime"`
UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime"`
}
// TableName returns the GORM table name.
func (OpenFlareNodeSystemProfile) TableName() string {
return "of_node_system_profiles"
}
// OpenFlareNodeObservationOpenresty stores openresty network observations.
type OpenFlareNodeObservationOpenresty struct {
ID uint `json:"id" gorm:"primaryKey;autoIncrement"`
NodeID string `json:"node_id" gorm:"index;size:64;not null"`
CapturedAt time.Time `json:"captured_at" gorm:"index"`
OpenrestyRxBytes int64 `json:"openresty_rx_bytes"`
OpenrestyTxBytes int64 `json:"openresty_tx_bytes"`
OpenrestyConnections int64 `json:"openresty_connections"`
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime"`
}
// TableName returns the GORM table name.
func (OpenFlareNodeObservationOpenresty) TableName() string {
return "of_node_obs_openresty"
}
// OpenFlareNodeObservationFrpc stores tunnel client frpc observations.
type OpenFlareNodeObservationFrpc struct {
ID uint `json:"id" gorm:"primaryKey;autoIncrement"`
NodeID string `json:"node_id" gorm:"index;size:64;not null"`
CapturedAt time.Time `json:"captured_at" gorm:"index"`
TunnelStatus string `json:"tunnel_status" gorm:"size:16"`
ConnectedRelaysCount int `json:"connected_relays_count"`
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime"`
}
// TableName returns the GORM table name.
func (OpenFlareNodeObservationFrpc) TableName() string {
return "of_node_obs_frpc"
}
// OpenFlareNodeObservationFrps stores tunnel relay frps observations.
type OpenFlareNodeObservationFrps struct {
ID uint `json:"id" gorm:"primaryKey;autoIncrement"`
NodeID string `json:"node_id" gorm:"index;size:64;not null"`
CapturedAt time.Time `json:"captured_at" gorm:"index"`
FrpsConnections int `json:"frps_connections"`
FrpsProxyCount int `json:"frps_proxy_count"`
FrpsClientCount int `json:"frps_client_count"`
FrpsProxies string `json:"frps_proxies" gorm:"type:text"`
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime"`
}
// TableName returns the GORM table name.
func (OpenFlareNodeObservationFrps) TableName() string {
return "of_node_obs_frps"
}
// OpenFlareAccessLogQuery filters access log list queries.
type OpenFlareAccessLogQuery struct {
NodeID string
RemoteAddr string
Host string
Path string
Since time.Time
Until time.Time
Page int
PageSize int
SortBy string
SortOrder string
}
// OpenFlareAccessLogBucketQuery filters folded access log queries (v1 stub).
type OpenFlareAccessLogBucketQuery struct {
NodeID string
RemoteAddr string
Host string
Path string
Since time.Time
Page int
PageSize int
SortBy string
SortOrder string
FoldMinutes int
}
// OpenFlareAccessLogBucketRow is a folded access log bucket row (v1 stub).
type OpenFlareAccessLogBucketRow struct {
BucketEpoch int64 `json:"bucket_epoch"`
RequestCount int64 `json:"request_count"`
UniqueIPCount int64 `json:"unique_ip_count"`
UniqueHostCount int64 `json:"unique_host_count"`
SuccessCount int64 `json:"success_count"`
ClientErrorCount int64 `json:"client_error_count"`
ServerErrorCount int64 `json:"server_error_count"`
}
// OpenFlareAccessLogBucketIPQuery filters folded IP summary queries (v1 stub).
type OpenFlareAccessLogBucketIPQuery struct {
NodeID string
RemoteAddr string
Host string
Path string
BucketStartedAt time.Time
FoldMinutes int
Page int
PageSize int
SortBy string
SortOrder string
}
// OpenFlareAccessLogBucketIPRow is a folded IP row (v1 stub).
type OpenFlareAccessLogBucketIPRow struct {
RemoteAddr string `json:"remote_addr"`
RequestCount int64 `json:"request_count"`
SuccessCount int64 `json:"success_count"`
ClientErrorCount int64 `json:"client_error_count"`
ServerErrorCount int64 `json:"server_error_count"`
LastSeenEpoch int64 `json:"last_seen_epoch"`
}
// OpenFlareAccessLogIPSummaryQuery filters IP summary list queries (v1 stub).
type OpenFlareAccessLogIPSummaryQuery struct {
NodeID string
RemoteAddr string
Host string
Since time.Time
Page int
PageSize int
SortBy string
SortOrder string
}
// OpenFlareAccessLogIPSummaryRow is an IP summary row (v1 stub).
type OpenFlareAccessLogIPSummaryRow struct {
RemoteAddr string `json:"remote_addr"`
TotalRequests int64 `json:"total_requests"`
RecentRequests int64 `json:"recent_requests"`
LastSeenEpoch int64 `json:"last_seen_epoch"`
}
// OpenFlareAccessLogIPTrendQuery filters IP trend queries (v1 stub).
type OpenFlareAccessLogIPTrendQuery struct {
NodeID string
RemoteAddr string
Host string
Since time.Time
BucketMinutes int
}
// OpenFlareAccessLogIPTrendRow is an IP trend bucket row (v1 stub).
type OpenFlareAccessLogIPTrendRow struct {
BucketEpoch int64 `json:"bucket_epoch"`
RequestCount int64 `json:"request_count"`
}
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")
}
// ListOpenFlareMetricSnapshotsSince returns metric snapshots since the given time.
func ListOpenFlareMetricSnapshotsSince(ctx context.Context, nodeID string, since time.Time, limit int) ([]*OpenFlareMetricSnapshot, error) {
conn := db.DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
query := conn.Model(&OpenFlareMetricSnapshot{}).Order("captured_at desc, id desc")
if nodeID != "" {
query = query.Where("node_id = ?", nodeID)
}
if !since.IsZero() {
query = query.Where("captured_at >= ?", since)
}
if limit > 0 {
query = query.Limit(limit)
}
var rows []*OpenFlareMetricSnapshot
if err := query.Find(&rows).Error; err != nil {
if isMissingTableError(err) {
return []*OpenFlareMetricSnapshot{}, nil
}
return nil, err
}
return rows, nil
}
// ListOpenFlareRequestReportsSince returns request reports since the given time.
func ListOpenFlareRequestReportsSince(ctx context.Context, nodeID string, since time.Time, limit int) ([]*OpenFlareRequestReport, error) {
conn := db.DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
query := conn.Model(&OpenFlareRequestReport{}).Order("window_ended_at desc, id desc")
if nodeID != "" {
query = query.Where("node_id = ?", nodeID)
}
if !since.IsZero() {
query = query.Where("window_ended_at >= ?", since)
}
if limit > 0 {
query = query.Limit(limit)
}
var rows []*OpenFlareRequestReport
if err := query.Find(&rows).Error; err != nil {
if isMissingTableError(err) {
return []*OpenFlareRequestReport{}, nil
}
return nil, err
}
return rows, nil
}
// ListOpenFlareActiveHealthEvents returns active health events across all nodes.
func ListOpenFlareActiveHealthEvents(ctx context.Context) ([]*OpenFlareHealthEvent, error) {
conn := db.DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
var rows []*OpenFlareHealthEvent
if err := conn.Where("status = ?", "active").Order("last_triggered_at desc").Find(&rows).Error; err != nil {
if isMissingTableError(err) {
return []*OpenFlareHealthEvent{}, nil
}
return nil, err
}
return rows, nil
}
// ListOpenFlareHealthEvents returns health events for a node.
func ListOpenFlareHealthEvents(ctx context.Context, nodeID string, activeOnly bool, limit int) ([]*OpenFlareHealthEvent, error) {
conn := db.DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
query := conn.Model(&OpenFlareHealthEvent{}).Where("node_id = ?", nodeID).Order("last_triggered_at desc")
if activeOnly {
query = query.Where("status = ?", "active")
}
if limit > 0 {
query = query.Limit(limit)
}
var rows []*OpenFlareHealthEvent
if err := query.Find(&rows).Error; err != nil {
if isMissingTableError(err) {
return []*OpenFlareHealthEvent{}, nil
}
return nil, err
}
return rows, nil
}
// DeleteOpenFlareMetricSnapshotsBefore deletes metric snapshots captured before cutoff.
func DeleteOpenFlareMetricSnapshotsBefore(ctx context.Context, cutoff time.Time) (int64, error) {
conn := db.DB(ctx)
if conn == nil {
return 0, errors.New(errDatabaseNotInitialized)
}
result := conn.Where("captured_at < ?", cutoff).Delete(&OpenFlareMetricSnapshot{})
if result.Error != nil {
if isMissingTableError(result.Error) {
return 0, nil
}
return 0, result.Error
}
return result.RowsAffected, nil
}
// DeleteAllOpenFlareMetricSnapshots deletes all metric snapshots.
func DeleteAllOpenFlareMetricSnapshots(ctx context.Context) (int64, error) {
conn := db.DB(ctx)
if conn == nil {
return 0, errors.New(errDatabaseNotInitialized)
}
result := conn.Where("1 = 1").Delete(&OpenFlareMetricSnapshot{})
if result.Error != nil {
if isMissingTableError(result.Error) {
return 0, nil
}
return 0, result.Error
}
return result.RowsAffected, nil
}
// DeleteOpenFlareRequestReportsBefore deletes request reports ending before cutoff.
func DeleteOpenFlareRequestReportsBefore(ctx context.Context, cutoff time.Time) (int64, error) {
conn := db.DB(ctx)
if conn == nil {
return 0, errors.New(errDatabaseNotInitialized)
}
result := conn.Where("window_ended_at < ?", cutoff).Delete(&OpenFlareRequestReport{})
if result.Error != nil {
if isMissingTableError(result.Error) {
return 0, nil
}
return 0, result.Error
}
return result.RowsAffected, nil
}
// DeleteAllOpenFlareRequestReports deletes all request reports.
func DeleteAllOpenFlareRequestReports(ctx context.Context) (int64, error) {
conn := db.DB(ctx)
if conn == nil {
return 0, errors.New(errDatabaseNotInitialized)
}
result := conn.Where("1 = 1").Delete(&OpenFlareRequestReport{})
if result.Error != nil {
if isMissingTableError(result.Error) {
return 0, nil
}
return 0, result.Error
}
return result.RowsAffected, nil
}
// DeleteOpenFlareHealthEventsByNodeID deletes all health events for a node.
func DeleteOpenFlareHealthEventsByNodeID(ctx context.Context, nodeID string) (int64, error) {
conn := db.DB(ctx)
if conn == nil {
return 0, errors.New(errDatabaseNotInitialized)
}
result := conn.Where("node_id = ?", nodeID).Delete(&OpenFlareHealthEvent{})
if result.Error != nil {
if isMissingTableError(result.Error) {
return 0, nil
}
return 0, result.Error
}
return result.RowsAffected, nil
}
// GetOpenFlareNodeSystemProfile returns the system profile for a node.
func GetOpenFlareNodeSystemProfile(ctx context.Context, nodeID string) (*OpenFlareNodeSystemProfile, error) {
conn := db.DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
var profile OpenFlareNodeSystemProfile
if err := conn.Where("node_id = ?", nodeID).First(&profile).Error; err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) || isMissingTableError(err) {
return nil, gorm.ErrRecordNotFound
}
return nil, err
}
return &profile, nil
}
// ListOpenFlareNodeObservationOpenresty returns openresty observations.
func ListOpenFlareNodeObservationOpenresty(ctx context.Context, nodeID string, since time.Time, limit int) ([]*OpenFlareNodeObservationOpenresty, error) {
conn := db.DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
query := conn.Model(&OpenFlareNodeObservationOpenresty{}).Order("captured_at desc, id desc")
if nodeID != "" {
query = query.Where("node_id = ?", nodeID)
}
if !since.IsZero() {
query = query.Where("captured_at >= ?", since)
}
if limit > 0 {
query = query.Limit(limit)
}
var rows []*OpenFlareNodeObservationOpenresty
if err := query.Find(&rows).Error; err != nil {
if isMissingTableError(err) {
return []*OpenFlareNodeObservationOpenresty{}, nil
}
return nil, err
}
return rows, nil
}
// ListOpenFlareNodeObservationFrpc returns frpc observations.
func ListOpenFlareNodeObservationFrpc(ctx context.Context, nodeID string, since time.Time, limit int) ([]*OpenFlareNodeObservationFrpc, error) {
conn := db.DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
query := conn.Model(&OpenFlareNodeObservationFrpc{}).Order("captured_at desc, id desc")
if nodeID != "" {
query = query.Where("node_id = ?", nodeID)
}
if !since.IsZero() {
query = query.Where("captured_at >= ?", since)
}
if limit > 0 {
query = query.Limit(limit)
}
var rows []*OpenFlareNodeObservationFrpc
if err := query.Find(&rows).Error; err != nil {
if isMissingTableError(err) {
return []*OpenFlareNodeObservationFrpc{}, nil
}
return nil, err
}
return rows, nil
}
// ListOpenFlareNodeObservationFrps returns frps observations.
func ListOpenFlareNodeObservationFrps(ctx context.Context, nodeID string, since time.Time, limit int) ([]*OpenFlareNodeObservationFrps, error) {
conn := db.DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
query := conn.Model(&OpenFlareNodeObservationFrps{}).Order("captured_at desc, id desc")
if nodeID != "" {
query = query.Where("node_id = ?", nodeID)
}
if !since.IsZero() {
query = query.Where("captured_at >= ?", since)
}
if limit > 0 {
query = query.Limit(limit)
}
var rows []*OpenFlareNodeObservationFrps
if err := query.Find(&rows).Error; err != nil {
if isMissingTableError(err) {
return []*OpenFlareNodeObservationFrps{}, nil
}
return nil, err
}
return rows, nil
}
@@ -0,0 +1,427 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package model
import (
"context"
"strconv"
"strings"
"sync"
"time"
"github.com/Rain-kl/Wavelet/internal/db"
"gorm.io/gorm"
)
// OpenFlareOption stores a hot-reloadable OpenFlare system option.
type OpenFlareOption struct {
Key string `json:"key" gorm:"column:key;primaryKey;size:128;not null"`
Value string `json:"value" gorm:"type:text;not null"`
}
// TableName returns the OpenFlare options table name.
func (OpenFlareOption) TableName() string {
return "of_options"
}
// OptionMap holds the in-memory option snapshot for fast reads and hot reload.
var (
OptionMap map[string]string
OptionMapRWMutex sync.RWMutex
// StartTime records process start time (seconds) for legacy /api/status.
StartTime = time.Now().Unix()
// Hot-reload mirrors for frequently read options (legacy OpenFlare keys).
SystemName = "OpenFlare"
ServerAddress = ""
Footer = ""
HomePageLink = ""
PasswordLoginEnabled = true
CapLoginEnabled = true
PasswordRegisterEnabled = false
EmailVerificationEnabled = false
GitHubOAuthEnabled = false
WeChatAuthEnabled = false
GitHubClientId = ""
GitHubClientSecret = ""
WeChatServerAddress = ""
WeChatServerToken = ""
WeChatAccountQRCodeImageURL = ""
SMTPServer = ""
SMTPPort = 587
SMTPAccount = ""
SMTPToken = ""
AgentDiscoveryToken = ""
AgentHeartbeatInterval = 10000
AgentWebsocketUpgradeEnabled = true
NodeOfflineThreshold = 2 * time.Minute
AgentUpdateRepo = "Rain-kl/OpenFlare"
GeoIPProvider = "ipinfo"
DatabaseAutoCleanupEnabled = false
DatabaseAutoCleanupRetentionDays = 30
UptimeKumaEnabled = false
UptimeKumaUrl = ""
UptimeKumaUsername = ""
UptimeKumaPassword = ""
UptimeKumaMonitorScope = "all"
UptimeKumaSelectedSites = ""
UptimeKumaSyncInterval = 5
UptimeKumaInterval = 60
UptimeKumaRetry = 0
UptimeKumaRetryInterval = 60
UptimeKumaTimeout = 48
OpenRestyDefaultServerReturnStatus = 421
OpenRestyWorkerProcesses = "auto"
OpenRestyWorkerConnections = 4096
OpenRestyWorkerRlimitNofile = 65535
OpenRestyEventsUse = "epoll"
OpenRestyEventsMultiAcceptEnabled = true
OpenRestyKeepaliveTimeout = 20
OpenRestyKeepaliveRequests = 1000
OpenRestyClientHeaderTimeout = 15
OpenRestyClientBodyTimeout = 15
OpenRestyClientMaxBodySize = "64m"
OpenRestyLargeClientHeaderBuffers = "4 16k"
OpenRestySendTimeout = 30
OpenRestyResolvers = ""
OpenRestyProxyConnectTimeout = 3
OpenRestyProxySendTimeout = 60
OpenRestyProxyReadTimeout = 60
OpenRestyWebsocketEnabled = true
OpenRestyHTTP3Enabled = true
OpenRestyProxyRequestBufferingEnabled = false
OpenRestyProxyBufferingEnabled = true
OpenRestyProxyBuffers = "16 16k"
OpenRestyProxyBufferSize = "8k"
OpenRestyProxyBusyBuffersSize = "64k"
OpenRestyGzipEnabled = true
OpenRestyGzipMinLength = 1024
OpenRestyGzipCompLevel = 5
OpenRestyCacheEnabled = false
OpenRestyCachePath = ""
OpenRestyCacheLevels = "1:2"
OpenRestyCacheInactive = "30m"
OpenRestyCacheMaxSize = "1g"
OpenRestyCacheKeyTemplate = "$scheme$host$request_uri"
OpenRestyCacheLockEnabled = true
OpenRestyCacheLockTimeout = "5s"
OpenRestyCacheUseStale = "error timeout updating http_500 http_502 http_503 http_504"
OpenRestyMainConfigTemplate = defaultOpenRestyMainConfigTemplate
)
const defaultOpenRestyMainConfigTemplate = `# This file is generated by OpenFlare. Do not edit manually.
worker_processes {{OpenRestyWorkerProcesses}};
worker_rlimit_nofile {{OpenRestyWorkerRlimitNofile}};
pid logs/nginx.pid;
error_log {{OpenRestyErrorLogPath}} warn;
events {
worker_connections {{OpenRestyWorkerConnections}};
{{OpenRestyEventsUseDirective}}{{OpenRestyEventsMultiAcceptDirective}}}
http {
include mime.types;
default_type application/octet-stream;
{{OpenRestyConnectionUpgradeMap}}{{OpenRestyDefaultServerBlock}} log_format openflare_json escape=json '{"ts":"$time_iso8601","host":"$host","path":"$request_uri","remote_addr":"$remote_addr","status":$status,"request_time":$request_time,"bytes_sent":$body_bytes_sent,"request_length":$request_length}';
access_log {{OpenRestyAccessLogPath}} openflare_json;
sendfile on;
tcp_nopush on;
tcp_nodelay on;
keepalive_timeout {{OpenRestyKeepaliveTimeout}};
keepalive_requests {{OpenRestyKeepaliveRequests}};
client_header_timeout {{OpenRestyClientHeaderTimeout}};
client_body_timeout {{OpenRestyClientBodyTimeout}};
client_max_body_size {{OpenRestyClientMaxBodySize}};
large_client_header_buffers {{OpenRestyLargeClientHeaderBuffers}};
send_timeout {{OpenRestySendTimeout}};
proxy_connect_timeout {{OpenRestyProxyConnectTimeout}};
proxy_send_timeout {{OpenRestyProxySendTimeout}};
proxy_read_timeout {{OpenRestyProxyReadTimeout}};
proxy_request_buffering {{OpenRestyProxyRequestBuffering}};
proxy_buffering {{OpenRestyProxyBuffering}};
proxy_buffers {{OpenRestyProxyBuffers}};
proxy_buffer_size {{OpenRestyProxyBufferSize}};
proxy_busy_buffers_size {{OpenRestyProxyBusyBuffersSize}};
gzip {{OpenRestyGzip}};
gzip_min_length {{OpenRestyGzipMinLength}};
gzip_comp_level {{OpenRestyGzipCompLevel}};
{{OpenRestyResolverDirective}}{{OpenRestyCacheBlock}} include {{OpenRestyRouteConfigInclude}};
}
`
// DefaultOpenFlareOptions returns built-in defaults keyed by legacy OpenFlare option names.
func DefaultOpenFlareOptions() map[string]string {
return map[string]string{
"PasswordLoginEnabled": strconv.FormatBool(PasswordLoginEnabled),
"CapLoginEnabled": strconv.FormatBool(CapLoginEnabled),
"PasswordRegisterEnabled": strconv.FormatBool(PasswordRegisterEnabled),
"EmailVerificationEnabled": strconv.FormatBool(EmailVerificationEnabled),
"GitHubOAuthEnabled": strconv.FormatBool(GitHubOAuthEnabled),
"WeChatAuthEnabled": strconv.FormatBool(WeChatAuthEnabled),
"SMTPServer": "",
"SMTPPort": strconv.Itoa(SMTPPort),
"SMTPAccount": "",
"SMTPToken": "",
"Notice": "",
"About": "",
"Footer": Footer,
"HomePageLink": HomePageLink,
"SystemName": SystemName,
"ServerAddress": "",
"GitHubClientId": "",
"GitHubClientSecret": "",
"WeChatServerAddress": "",
"WeChatServerToken": "",
"WeChatAccountQRCodeImageURL": "",
"AgentDiscoveryToken": "",
"AgentHeartbeatInterval": strconv.Itoa(AgentHeartbeatInterval),
"AgentWebsocketUpgradeEnabled": strconv.FormatBool(AgentWebsocketUpgradeEnabled),
"NodeOfflineThreshold": strconv.Itoa(int(NodeOfflineThreshold.Milliseconds())),
"AgentUpdateRepo": AgentUpdateRepo,
"GeoIPProvider": GeoIPProvider,
"DatabaseAutoCleanupEnabled": strconv.FormatBool(DatabaseAutoCleanupEnabled),
"UptimeKumaEnabled": strconv.FormatBool(UptimeKumaEnabled),
"UptimeKumaUrl": UptimeKumaUrl,
"UptimeKumaUsername": UptimeKumaUsername,
"UptimeKumaPassword": UptimeKumaPassword,
"UptimeKumaMonitorScope": UptimeKumaMonitorScope,
"UptimeKumaSelectedSites": UptimeKumaSelectedSites,
"UptimeKumaSyncInterval": strconv.Itoa(UptimeKumaSyncInterval),
"UptimeKumaInterval": strconv.Itoa(UptimeKumaInterval),
"UptimeKumaRetry": strconv.Itoa(UptimeKumaRetry),
"UptimeKumaRetryInterval": strconv.Itoa(UptimeKumaRetryInterval),
"UptimeKumaTimeout": strconv.Itoa(UptimeKumaTimeout),
"DatabaseAutoCleanupRetentionDays": strconv.Itoa(DatabaseAutoCleanupRetentionDays),
"OpenRestyDefaultServerReturnStatus": strconv.Itoa(OpenRestyDefaultServerReturnStatus),
"OpenRestyWorkerProcesses": OpenRestyWorkerProcesses,
"OpenRestyWorkerConnections": strconv.Itoa(OpenRestyWorkerConnections),
"OpenRestyWorkerRlimitNofile": strconv.Itoa(OpenRestyWorkerRlimitNofile),
"OpenRestyEventsUse": OpenRestyEventsUse,
"OpenRestyEventsMultiAcceptEnabled": strconv.FormatBool(OpenRestyEventsMultiAcceptEnabled),
"OpenRestyKeepaliveTimeout": strconv.Itoa(OpenRestyKeepaliveTimeout),
"OpenRestyKeepaliveRequests": strconv.Itoa(OpenRestyKeepaliveRequests),
"OpenRestyClientHeaderTimeout": strconv.Itoa(OpenRestyClientHeaderTimeout),
"OpenRestyClientBodyTimeout": strconv.Itoa(OpenRestyClientBodyTimeout),
"OpenRestyClientMaxBodySize": OpenRestyClientMaxBodySize,
"OpenRestyLargeClientHeaderBuffers": OpenRestyLargeClientHeaderBuffers,
"OpenRestySendTimeout": strconv.Itoa(OpenRestySendTimeout),
"OpenRestyProxyConnectTimeout": strconv.Itoa(OpenRestyProxyConnectTimeout),
"OpenRestyProxySendTimeout": strconv.Itoa(OpenRestyProxySendTimeout),
"OpenRestyProxyReadTimeout": strconv.Itoa(OpenRestyProxyReadTimeout),
"OpenRestyWebsocketEnabled": strconv.FormatBool(OpenRestyWebsocketEnabled),
"OpenRestyHTTP3Enabled": strconv.FormatBool(OpenRestyHTTP3Enabled),
"OpenRestyProxyRequestBufferingEnabled": strconv.FormatBool(OpenRestyProxyRequestBufferingEnabled),
"OpenRestyProxyBufferingEnabled": strconv.FormatBool(OpenRestyProxyBufferingEnabled),
"OpenRestyProxyBuffers": OpenRestyProxyBuffers,
"OpenRestyProxyBufferSize": OpenRestyProxyBufferSize,
"OpenRestyProxyBusyBuffersSize": OpenRestyProxyBusyBuffersSize,
"OpenRestyGzipEnabled": strconv.FormatBool(OpenRestyGzipEnabled),
"OpenRestyGzipMinLength": strconv.Itoa(OpenRestyGzipMinLength),
"OpenRestyGzipCompLevel": strconv.Itoa(OpenRestyGzipCompLevel),
"OpenRestyCacheEnabled": strconv.FormatBool(OpenRestyCacheEnabled),
"OpenRestyCachePath": OpenRestyCachePath,
"OpenRestyCacheLevels": OpenRestyCacheLevels,
"OpenRestyCacheInactive": OpenRestyCacheInactive,
"OpenRestyCacheMaxSize": OpenRestyCacheMaxSize,
"OpenRestyCacheKeyTemplate": OpenRestyCacheKeyTemplate,
"OpenRestyCacheLockEnabled": strconv.FormatBool(OpenRestyCacheLockEnabled),
"OpenRestyCacheLockTimeout": OpenRestyCacheLockTimeout,
"OpenRestyCacheUseStale": OpenRestyCacheUseStale,
"OpenRestyMainConfigTemplate": OpenRestyMainConfigTemplate,
}
}
// InitOptionMap seeds defaults and overlays persisted options from of_options.
func InitOptionMap(ctx context.Context) error {
OptionMapRWMutex.Lock()
OptionMap = DefaultOpenFlareOptions()
OptionMapRWMutex.Unlock()
options, err := ListOpenFlareOptions(ctx)
if err != nil {
return err
}
for _, option := range options {
applyOptionMap(option.Key, option.Value)
}
return nil
}
// ListOpenFlareOptions returns all persisted options.
func ListOpenFlareOptions(ctx context.Context) ([]OpenFlareOption, error) {
var options []OpenFlareOption
if err := db.DB(ctx).Find(&options).Error; err != nil {
return nil, err
}
return options, nil
}
// UpdateOpenFlareOption updates a single option in DB and memory.
func UpdateOpenFlareOption(ctx context.Context, key, value string) error {
return UpdateOpenFlareOptions(ctx, []OpenFlareOption{{Key: key, Value: value}})
}
// UpdateOpenFlareOptions batch-updates options in a transaction and refreshes OptionMap.
func UpdateOpenFlareOptions(ctx context.Context, options []OpenFlareOption) error {
if len(options) == 0 {
return nil
}
if err := db.DB(ctx).Transaction(func(tx *gorm.DB) error {
for _, item := range options {
if item.Key == "UptimeKumaPassword" && strings.TrimSpace(item.Value) == "" {
continue
}
option := OpenFlareOption{Key: item.Key}
if err := tx.FirstOrCreate(&option, OpenFlareOption{Key: item.Key}).Error; err != nil {
return err
}
option.Value = item.Value
if err := tx.Save(&option).Error; err != nil {
return err
}
}
return nil
}); err != nil {
return err
}
for _, item := range options {
if item.Key == "UptimeKumaPassword" && strings.TrimSpace(item.Value) == "" {
continue
}
applyOptionMap(item.Key, item.Value)
}
return nil
}
// OptionValue returns a snapshot value from OptionMap.
func OptionValue(key string) string {
OptionMapRWMutex.RLock()
defer OptionMapRWMutex.RUnlock()
if OptionMap == nil {
return ""
}
return OptionMap[key]
}
// ResetOptionMapForTest clears in-memory option state for unit tests.
func ResetOptionMapForTest() {
OptionMapRWMutex.Lock()
OptionMap = nil
OptionMapRWMutex.Unlock()
}
func applyOptionMap(key, value string) {
OptionMapRWMutex.Lock()
if OptionMap == nil {
OptionMap = make(map[string]string)
}
OptionMap[key] = value
if strings.HasSuffix(key, "Enabled") {
boolValue := value == "true"
switch key {
case "PasswordRegisterEnabled":
PasswordRegisterEnabled = boolValue
case "PasswordLoginEnabled":
PasswordLoginEnabled = boolValue
case "CapLoginEnabled":
CapLoginEnabled = boolValue
case "EmailVerificationEnabled":
EmailVerificationEnabled = boolValue
case "GitHubOAuthEnabled":
GitHubOAuthEnabled = boolValue
case "WeChatAuthEnabled":
WeChatAuthEnabled = boolValue
}
}
switch key {
case "SMTPServer":
SMTPServer = value
case "SMTPPort":
if intValue, err := strconv.Atoi(value); err == nil {
SMTPPort = intValue
}
case "SMTPAccount":
SMTPAccount = value
case "SMTPToken":
SMTPToken = value
case "ServerAddress":
ServerAddress = value
case "GitHubClientId":
GitHubClientId = value
case "GitHubClientSecret":
GitHubClientSecret = value
case "Footer":
Footer = value
case "HomePageLink":
HomePageLink = value
case "SystemName":
SystemName = value
case "WeChatServerAddress":
WeChatServerAddress = value
case "WeChatServerToken":
WeChatServerToken = value
case "WeChatAccountQRCodeImageURL":
WeChatAccountQRCodeImageURL = value
case "AgentDiscoveryToken":
AgentDiscoveryToken = value
case "AgentHeartbeatInterval":
if v, err := strconv.Atoi(value); err == nil && v > 0 {
AgentHeartbeatInterval = v
}
case "AgentWebsocketUpgradeEnabled":
AgentWebsocketUpgradeEnabled = value == "true"
case "NodeOfflineThreshold":
if v, err := strconv.Atoi(value); err == nil && v > 0 {
NodeOfflineThreshold = time.Duration(v) * time.Millisecond
}
case "AgentUpdateRepo":
if value != "" {
AgentUpdateRepo = value
}
case "GeoIPProvider":
GeoIPProvider = value
case "UptimeKumaEnabled":
UptimeKumaEnabled = value == "true"
case "UptimeKumaUrl":
UptimeKumaUrl = value
case "UptimeKumaUsername":
UptimeKumaUsername = value
case "UptimeKumaPassword":
UptimeKumaPassword = value
case "UptimeKumaMonitorScope":
UptimeKumaMonitorScope = value
case "UptimeKumaSelectedSites":
UptimeKumaSelectedSites = value
case "UptimeKumaSyncInterval":
if v, err := strconv.Atoi(value); err == nil && v > 0 {
UptimeKumaSyncInterval = v
}
case "UptimeKumaInterval":
if v, err := strconv.Atoi(value); err == nil && v > 0 {
UptimeKumaInterval = v
}
case "UptimeKumaRetry":
if v, err := strconv.Atoi(value); err == nil && v >= 0 {
UptimeKumaRetry = v
}
case "UptimeKumaRetryInterval":
if v, err := strconv.Atoi(value); err == nil && v > 0 {
UptimeKumaRetryInterval = v
}
case "UptimeKumaTimeout":
if v, err := strconv.Atoi(value); err == nil && v > 0 {
UptimeKumaTimeout = v
}
case "DatabaseAutoCleanupEnabled":
DatabaseAutoCleanupEnabled = value == "true"
case "DatabaseAutoCleanupRetentionDays":
if v, err := strconv.Atoi(value); err == nil && v >= 1 {
DatabaseAutoCleanupRetentionDays = v
}
}
OptionMapRWMutex.Unlock()
}
@@ -0,0 +1,133 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package model
import (
"context"
"time"
"github.com/Rain-kl/Wavelet/internal/db"
)
// Origin OpenFlare 源站实体。
type Origin struct {
ID uint `json:"id" gorm:"primaryKey;autoIncrement"`
Name string `json:"name" gorm:"size:255;not null"`
Address string `json:"address" gorm:"uniqueIndex;size:255;not null"`
Remark string `json:"remark" gorm:"size:255"`
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime"`
UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime"`
}
// TableName 表名。
func (Origin) TableName() string {
return "of_origins"
}
// OriginRouteCount 源站关联的代理规则数量。
type OriginRouteCount struct {
OriginID uint `json:"origin_id"`
RouteCount int64 `json:"route_count"`
}
// OriginProxyRoute 源站模块查询代理规则时使用的最小字段集。
type OriginProxyRoute struct {
ID uint `gorm:"column:id;primaryKey"`
OriginID *uint `gorm:"column:origin_id"`
Domain string `gorm:"column:domain"`
OriginURL string `gorm:"column:origin_url"`
Upstreams string `gorm:"column:upstreams"`
Enabled bool `gorm:"column:enabled"`
UpdatedAt time.Time `gorm:"column:updated_at"`
}
// TableName 表名。
func (OriginProxyRoute) TableName() string {
return "of_proxy_routes"
}
// HasProxyRoutesTable 判断代理规则表是否已迁移。
func HasProxyRoutesTable(ctx context.Context) bool {
return db.DB(ctx).Migrator().HasTable(&OriginProxyRoute{})
}
// ListOrigins 列出全部源站。
func ListOrigins(ctx context.Context) ([]Origin, error) {
var origins []Origin
if err := db.DB(ctx).Order("id desc").Find(&origins).Error; err != nil {
return nil, err
}
return origins, nil
}
// GetOriginByID 按 ID 查询源站。
func GetOriginByID(ctx context.Context, id uint) (*Origin, error) {
var origin Origin
if err := db.DB(ctx).First(&origin, id).Error; err != nil {
return nil, err
}
return &origin, nil
}
// GetOriginByAddress 按地址查询源站。
func GetOriginByAddress(ctx context.Context, address string) (*Origin, error) {
var origin Origin
if err := db.DB(ctx).Where("address = ?", address).First(&origin).Error; err != nil {
return nil, err
}
return &origin, nil
}
// CreateOriginRecord 创建源站。
func CreateOriginRecord(ctx context.Context, origin *Origin) error {
return db.DB(ctx).Create(origin).Error
}
// SaveOrigin 保存源站。
func SaveOrigin(ctx context.Context, origin *Origin) error {
return db.DB(ctx).Save(origin).Error
}
// DeleteOriginRecord 删除源站。
func DeleteOriginRecord(ctx context.Context, id uint) error {
return db.DB(ctx).Delete(&Origin{}, id).Error
}
// ListOriginRouteCounts 统计各源站关联的代理规则数量。
func ListOriginRouteCounts(ctx context.Context) ([]OriginRouteCount, error) {
if !HasProxyRoutesTable(ctx) {
return nil, nil
}
result := make([]OriginRouteCount, 0)
err := db.DB(ctx).Model(&OriginProxyRoute{}).
Select("origin_id, COUNT(*) AS route_count").
Where("origin_id IS NOT NULL").
Group("origin_id").
Scan(&result).Error
return result, err
}
// ListProxyRoutesByOriginID 列出源站关联的代理规则。
func ListProxyRoutesByOriginID(ctx context.Context, originID uint) ([]OriginProxyRoute, error) {
if !HasProxyRoutesTable(ctx) {
return nil, nil
}
var routes []OriginProxyRoute
if err := db.DB(ctx).Where("origin_id = ?", originID).Order("id desc").Find(&routes).Error; err != nil {
return nil, err
}
return routes, nil
}
// CountProxyRoutesByOriginID 统计源站关联的代理规则数量。
func CountProxyRoutesByOriginID(ctx context.Context, originID uint) (int64, error) {
if !HasProxyRoutesTable(ctx) {
return 0, nil
}
var count int64
if err := db.DB(ctx).Model(&OriginProxyRoute{}).Where("origin_id = ?", originID).Count(&count).Error; err != nil {
return 0, err
}
return count, nil
}
@@ -0,0 +1,162 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package model
import (
"context"
"time"
"github.com/Rain-kl/Wavelet/internal/db"
)
const (
PagesDeploymentStatusUploaded = "uploaded"
PagesDeploymentStatusActive = "active"
)
// PagesProject OpenFlare Pages 静态托管项目。
type PagesProject struct {
ID uint `json:"id" gorm:"primaryKey;autoIncrement"`
Name string `json:"name" gorm:"size:255;not null"`
Slug string `json:"slug" gorm:"uniqueIndex;size:128;not null"`
Description string `json:"description" gorm:"type:text;not null;default:''"`
Enabled bool `json:"enabled" gorm:"not null;default:true"`
SPAFallbackEnabled bool `json:"spa_fallback_enabled" gorm:"not null;default:false"`
SPAFallbackPath string `json:"spa_fallback_path" gorm:"size:512;not null;default:'/index.html'"`
APIProxyEnabled bool `json:"api_proxy_enabled" gorm:"not null;default:false"`
APIProxyPath string `json:"api_proxy_path" gorm:"size:255;not null;default:''"`
APIProxyPass string `json:"api_proxy_pass" gorm:"size:2048;not null;default:''"`
APIProxyRewrite string `json:"api_proxy_rewrite" gorm:"size:255;not null;default:''"`
ActiveDeploymentID *uint `json:"active_deployment_id" gorm:"index"`
RootDir string `json:"root_dir" gorm:"size:512;not null;default:''"`
EntryFile string `json:"entry_file" gorm:"size:512;not null;default:'index.html'"`
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime"`
UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime"`
}
// TableName 表名。
func (PagesProject) TableName() string {
return "of_pages_projects"
}
// PagesDeployment OpenFlare Pages 不可变部署记录。
type PagesDeployment struct {
ID uint `json:"id" gorm:"primaryKey;autoIncrement"`
ProjectID uint `json:"project_id" gorm:"not null;index"`
DeploymentNumber int `json:"deployment_number" gorm:"not null"`
Checksum string `json:"checksum" gorm:"size:64;not null;index"`
Status string `json:"status" gorm:"size:32;not null;default:'uploaded';index"`
UploadID uint64 `json:"upload_id,string" gorm:"not null;default:0;index"`
ArtifactPath string `json:"artifact_path,omitempty" gorm:"size:2048;not null;default:''"` // legacy only
FileCount int `json:"file_count" gorm:"not null;default:0"`
TotalSize int64 `json:"total_size" gorm:"not null;default:0"`
CreatedBy string `json:"created_by" gorm:"size:64;not null;default:''"`
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime"`
ActivatedAt *time.Time `json:"activated_at"`
}
// TableName 表名。
func (PagesDeployment) TableName() string {
return "of_pages_deployments"
}
// PagesDeploymentFile OpenFlare Pages 部署文件清单。
type PagesDeploymentFile struct {
ID uint `json:"id" gorm:"primaryKey;autoIncrement"`
DeploymentID uint `json:"deployment_id" gorm:"not null;index"`
Path string `json:"path" gorm:"size:2048;not null"`
Size int64 `json:"size" gorm:"not null;default:0"`
Checksum string `json:"checksum" gorm:"size:64;not null"`
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime"`
}
// TableName 表名。
func (PagesDeploymentFile) TableName() string {
return "of_pages_deployment_files"
}
// HasPagesProjectsTable 判断 Pages 项目表是否已迁移。
func HasPagesProjectsTable(ctx context.Context) bool {
return db.DB(ctx).Migrator().HasTable(&PagesProject{})
}
// ListPagesProjects 列出全部 Pages 项目。
func ListPagesProjects(ctx context.Context) ([]PagesProject, error) {
var projects []PagesProject
if err := db.DB(ctx).Order("id desc").Find(&projects).Error; err != nil {
return nil, err
}
return projects, nil
}
// GetPagesProjectByID 按 ID 查询 Pages 项目。
func GetPagesProjectByID(ctx context.Context, id uint) (*PagesProject, error) {
var project PagesProject
if err := db.DB(ctx).First(&project, id).Error; err != nil {
return nil, err
}
return &project, nil
}
// GetPagesProjectBySlug 按 slug 查询 Pages 项目。
func GetPagesProjectBySlug(ctx context.Context, slug string) (*PagesProject, error) {
var project PagesProject
if err := db.DB(ctx).Where("slug = ?", slug).First(&project).Error; err != nil {
return nil, err
}
return &project, nil
}
// CreatePagesProjectRecord 创建 Pages 项目。
func CreatePagesProjectRecord(ctx context.Context, project *PagesProject) error {
return db.DB(ctx).Create(project).Error
}
// ListPagesDeployments 列出项目的全部部署。
func ListPagesDeployments(ctx context.Context, projectID uint) ([]PagesDeployment, error) {
var deployments []PagesDeployment
if err := db.DB(ctx).Where("project_id = ?", projectID).Order("id desc").Find(&deployments).Error; err != nil {
return nil, err
}
return deployments, nil
}
// GetPagesDeploymentByID 按 ID 查询 Pages 部署。
func GetPagesDeploymentByID(ctx context.Context, id uint) (*PagesDeployment, error) {
var deployment PagesDeployment
if err := db.DB(ctx).First(&deployment, id).Error; err != nil {
return nil, err
}
return &deployment, nil
}
// ListPagesDeploymentFiles 列出部署文件清单。
func ListPagesDeploymentFiles(ctx context.Context, deploymentID uint) ([]PagesDeploymentFile, error) {
var files []PagesDeploymentFile
if err := db.DB(ctx).Where("deployment_id = ?", deploymentID).Order("path asc").Find(&files).Error; err != nil {
return nil, err
}
return files, nil
}
// CountPagesDeploymentsByProjectID 统计项目部署数量。
func CountPagesDeploymentsByProjectID(ctx context.Context, projectID uint) (int64, error) {
var count int64
if err := db.DB(ctx).Model(&PagesDeployment{}).Where("project_id = ?", projectID).Count(&count).Error; err != nil {
return 0, err
}
return count, nil
}
// CountProxyRoutesByPagesProjectID 统计引用 Pages 项目的代理规则数量。
func CountProxyRoutesByPagesProjectID(ctx context.Context, projectID uint) (int64, error) {
if !HasProxyRoutesTable(ctx) {
return 0, nil
}
var count int64
if err := db.DB(ctx).Model(&ProxyRoute{}).Where("pages_project_id = ?", projectID).Count(&count).Error; err != nil {
return 0, err
}
return count, nil
}
@@ -1,9 +1,18 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package model
import "time"
import (
"context"
"time"
"github.com/Rain-kl/Wavelet/internal/db"
)
// ProxyRoute OpenFlare 代理规则实体。
type ProxyRoute struct {
ID uint `json:"id" gorm:"primaryKey"`
ID uint `json:"id" gorm:"primaryKey;autoIncrement"`
SiteName string `json:"site_name" gorm:"size:255;not null;default:''"`
Domain string `json:"domain" gorm:"uniqueIndex;size:255;not null"`
Domains string `json:"domains" gorm:"type:text;not null;default:'[]'"`
@@ -33,37 +42,41 @@ type ProxyRoute struct {
TunnelTargetAddr string `json:"tunnel_target_addr" gorm:"size:512"`
TunnelTargetProtocol string `json:"tunnel_target_protocol" gorm:"size:16"`
PagesProjectID *uint `json:"pages_project_id" gorm:"index"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime"`
UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime"`
}
func ListProxyRoutes() (routes []*ProxyRoute, err error) {
err = DB.Order("id desc").Find(&routes).Error
return routes, err
// TableName 表名。
func (ProxyRoute) TableName() string {
return "of_proxy_routes"
}
func GetEnabledProxyRoutes() (routes []*ProxyRoute, err error) {
err = DB.Where("enabled = ?", true).Order("site_name asc").Order("domain asc").Find(&routes).Error
return routes, err
// ListProxyRoutes 列出全部代理规则。
func ListProxyRoutes(ctx context.Context) ([]*ProxyRoute, error) {
var routes []*ProxyRoute
if err := db.DB(ctx).Order("id desc").Find(&routes).Error; err != nil {
return nil, err
}
return routes, nil
}
func GetProxyRouteByID(id uint) (*ProxyRoute, error) {
route := &ProxyRoute{}
err := DB.First(route, id).Error
return route, err
// GetProxyRouteByID 按 ID 查询代理规则。
func GetProxyRouteByID(ctx context.Context, id uint) (*ProxyRoute, error) {
var route ProxyRoute
if err := db.DB(ctx).First(&route, id).Error; err != nil {
return nil, err
}
return &route, nil
}
func ListProxyRoutesByOriginID(originID uint) (routes []*ProxyRoute, err error) {
err = DB.Where("origin_id = ?", originID).Order("id desc").Find(&routes).Error
return routes, err
// CreateProxyRouteRecord 创建代理规则。
func CreateProxyRouteRecord(ctx context.Context, route *ProxyRoute) error {
return db.DB(ctx).Create(route).Error
}
func (route *ProxyRoute) Insert() error {
return DB.Create(route).Error
}
func (route *ProxyRoute) Update() error {
return DB.Model(&ProxyRoute{}).Where("id = ?", route.ID).Updates(map[string]any{
// UpdateProxyRouteRecord 更新代理规则。
func UpdateProxyRouteRecord(ctx context.Context, route *ProxyRoute) error {
return db.DB(ctx).Model(&ProxyRoute{}).Where("id = ?", route.ID).Updates(map[string]any{
"site_name": route.SiteName,
"domain": route.Domain,
"domains": route.Domains,
@@ -96,6 +109,7 @@ func (route *ProxyRoute) Update() error {
}).Error
}
func (route *ProxyRoute) Delete() error {
return DB.Delete(route).Error
// DeleteProxyRouteRecord 删除代理规则。
func DeleteProxyRouteRecord(ctx context.Context, id uint) error {
return db.DB(ctx).Delete(&ProxyRoute{}, id).Error
}
@@ -0,0 +1,139 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package model
import (
"context"
"errors"
"time"
"github.com/Rain-kl/Wavelet/internal/db"
)
// TLSCertificate OpenFlare TLS 证书实体。
type TLSCertificate struct {
ID uint `json:"id" gorm:"primaryKey;autoIncrement"`
Name string `json:"name" gorm:"uniqueIndex;size:255;not null"`
CertPEM string `json:"-" gorm:"type:text;not null"`
KeyPEM string `json:"-" gorm:"type:text;not null"`
NotBefore time.Time `json:"not_before"`
NotAfter time.Time `json:"not_after"`
Remark string `json:"remark" gorm:"size:255"`
Provider string `json:"provider" gorm:"size:64;default:upload"`
AcmeAccountID uint `json:"acme_account_id"`
DnsAccountID uint `json:"dns_account_id"`
KeyAlgorithm string `json:"key_algorithm" gorm:"size:32"`
AutoRenew bool `json:"auto_renew"`
PrimaryDomain string `json:"primary_domain" gorm:"size:255"`
OtherDomains string `json:"other_domains" gorm:"type:text"`
DisableCNAME bool `json:"disable_cname"`
SkipDNS bool `json:"skip_dns"`
DNS1 string `json:"dns1" gorm:"size:128"`
DNS2 string `json:"dns2" gorm:"size:128"`
ApplyStatus string `json:"apply_status" gorm:"size:64;default:ready"`
ApplyMessage string `json:"apply_message" gorm:"type:text"`
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime"`
UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime"`
}
// TableName 表名。
func (TLSCertificate) TableName() string {
return "of_tls_certificates"
}
// TLSProxyRouteRef 删除证书时检查代理规则引用的最小字段集。
type TLSProxyRouteRef struct {
ID uint `gorm:"column:id;primaryKey"`
CertID *uint `gorm:"column:cert_id"`
CertIDs string `gorm:"column:cert_ids"`
DomainCertIDs string `gorm:"column:domain_cert_ids"`
}
// TableName 表名。
func (TLSProxyRouteRef) TableName() string {
return "of_proxy_routes"
}
// HasTLSProxyRoutesTable 判断代理规则表是否已迁移。
func HasTLSProxyRoutesTable(ctx context.Context) bool {
return db.DB(ctx).Migrator().HasTable(&TLSProxyRouteRef{})
}
// ListTLSCertificates 列出全部证书(不含 PEM 敏感字段的 JSON 暴露由 struct tag 控制)。
func ListTLSCertificates(ctx context.Context) ([]TLSCertificate, error) {
conn := db.DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
var certificates []TLSCertificate
if err := conn.Order("id desc").Find(&certificates).Error; err != nil {
return nil, err
}
return certificates, nil
}
// GetTLSCertificateByID 按 ID 查询证书。
func GetTLSCertificateByID(ctx context.Context, id uint) (*TLSCertificate, error) {
conn := db.DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
var certificate TLSCertificate
if err := conn.First(&certificate, id).Error; err != nil {
return nil, err
}
return &certificate, nil
}
// CreateTLSCertificateRecord 创建证书记录。
func CreateTLSCertificateRecord(ctx context.Context, certificate *TLSCertificate) error {
conn := db.DB(ctx)
if conn == nil {
return errors.New(errDatabaseNotInitialized)
}
return conn.Create(certificate).Error
}
// SaveTLSCertificate 保存证书记录。
func SaveTLSCertificate(ctx context.Context, certificate *TLSCertificate) error {
conn := db.DB(ctx)
if conn == nil {
return errors.New(errDatabaseNotInitialized)
}
return conn.Save(certificate).Error
}
// DeleteTLSCertificateRecord 删除证书记录。
func DeleteTLSCertificateRecord(ctx context.Context, id uint) error {
conn := db.DB(ctx)
if conn == nil {
return errors.New(errDatabaseNotInitialized)
}
return conn.Delete(&TLSCertificate{}, id).Error
}
// CountTLSCertificatesByDNSAccountID 统计引用指定 DNS 账号的证书数量。
func CountTLSCertificatesByDNSAccountID(ctx context.Context, dnsAccountID uint) (int64, error) {
conn := db.DB(ctx)
if conn == nil {
return 0, errors.New(errDatabaseNotInitialized)
}
var count int64
if err := conn.Model(&TLSCertificate{}).Where("dns_account_id = ?", dnsAccountID).Count(&count).Error; err != nil {
return 0, err
}
return count, nil
}
// ListTLSProxyRouteRefs 列出代理规则证书引用字段。
func ListTLSProxyRouteRefs(ctx context.Context) ([]TLSProxyRouteRef, error) {
if !HasTLSProxyRoutesTable(ctx) {
return nil, nil
}
var routes []TLSProxyRouteRef
if err := db.DB(ctx).Order("id asc").Find(&routes).Error; err != nil {
return nil, err
}
return routes, nil
}
@@ -0,0 +1,392 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package model
import (
"context"
"errors"
"time"
"github.com/Rain-kl/Wavelet/internal/db"
"gorm.io/gorm"
)
// OpenFlareWAFRuleGroup stores a WAF rule group.
type OpenFlareWAFRuleGroup struct {
ID uint `json:"id" gorm:"primaryKey;autoIncrement"`
Name string `json:"name" gorm:"size:255;not null"`
Enabled bool `json:"enabled" gorm:"not null;default:true"`
IsGlobal bool `json:"is_global" gorm:"not null;default:false;index"`
BlockStatusCode int `json:"block_status_code" gorm:"not null;default:418"`
BlockResponseBody string `json:"block_response_body" gorm:"type:text;not null;default:''"`
IPWhitelist string `json:"ip_whitelist" gorm:"type:text;not null;default:'[]'"`
IPBlacklist string `json:"ip_blacklist" gorm:"type:text;not null;default:'[]'"`
IPWhitelistGroups string `json:"ip_whitelist_group_ids" gorm:"column:ip_whitelist_groups;type:text;not null;default:'[]'"`
IPBlacklistGroups string `json:"ip_blacklist_group_ids" gorm:"column:ip_blacklist_groups;type:text;not null;default:'[]'"`
CountryWhitelist string `json:"country_whitelist" gorm:"type:text;not null;default:'[]'"`
CountryBlacklist string `json:"country_blacklist" gorm:"type:text;not null;default:'[]'"`
RegionWhitelist string `json:"region_whitelist" gorm:"type:text;not null;default:'[]'"`
RegionBlacklist string `json:"region_blacklist" gorm:"type:text;not null;default:'[]'"`
PoWEnabled bool `json:"pow_enabled" gorm:"column:pow_enabled;not null;default:false"`
PoWConfig string `json:"pow_config" gorm:"column:pow_config;type:text;not null;default:'{}'"`
Remark string `json:"remark" gorm:"size:255"`
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime"`
UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime"`
}
// TableName returns the GORM table name.
func (OpenFlareWAFRuleGroup) TableName() string {
return "of_waf_rule_groups"
}
// OpenFlareWAFIPGroup stores a WAF IP group.
type OpenFlareWAFIPGroup struct {
ID uint `json:"id" gorm:"primaryKey;autoIncrement"`
Name string `json:"name" gorm:"size:255;not null"`
Type string `json:"type" gorm:"size:32;not null;index"`
Enabled bool `json:"enabled" gorm:"not null;default:true"`
IPList string `json:"ip_list" gorm:"type:text;not null;default:'[]'"`
AutoConfig string `json:"auto_config" gorm:"type:text;not null;default:'{}'"`
ExtIPs string `json:"ext_ips" gorm:"type:text;not null;default:'[]'"`
SubscriptionURL string `json:"subscription_url" gorm:"size:2048;not null;default:''"`
SubscriptionFormat string `json:"subscription_format" gorm:"size:32;not null;default:'text'"`
SubscriptionMappingRule string `json:"subscription_mapping_rule" gorm:"size:255;not null;default:''"`
SyncIntervalMinutes int `json:"sync_interval_minutes" gorm:"not null;default:1440"`
LastSyncedAt *time.Time `json:"last_synced_at"`
NextSyncAt *time.Time `json:"next_sync_at" gorm:"index"`
LastSyncStatus string `json:"last_sync_status" gorm:"size:32;not null;default:''"`
LastSyncMessage string `json:"last_sync_message" gorm:"type:text;not null;default:''"`
Remark string `json:"remark" gorm:"size:255"`
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime"`
UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime"`
}
// TableName returns the GORM table name.
func (OpenFlareWAFIPGroup) TableName() string {
return "of_waf_ip_groups"
}
// OpenFlareWAFRuleGroupBinding binds a rule group to a proxy route.
type OpenFlareWAFRuleGroupBinding struct {
ID uint `json:"id" gorm:"primaryKey;autoIncrement"`
RuleGroupID uint `json:"rule_group_id" gorm:"not null;uniqueIndex:idx_of_waf_group_route"`
ProxyRouteID uint `json:"proxy_route_id" gorm:"not null;uniqueIndex:idx_of_waf_group_route;index"`
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime"`
}
// TableName returns the GORM table name.
func (OpenFlareWAFRuleGroupBinding) TableName() string {
return "of_waf_rule_group_bindings"
}
func wafDB(ctx context.Context) (*gorm.DB, error) {
conn := db.DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
return conn, nil
}
// ListOpenFlareWAFRuleGroups returns all rule groups.
func ListOpenFlareWAFRuleGroups(ctx context.Context) ([]*OpenFlareWAFRuleGroup, error) {
conn, err := wafDB(ctx)
if err != nil {
return nil, err
}
var groups []*OpenFlareWAFRuleGroup
if err = conn.Order("is_global desc").Order("id asc").Find(&groups).Error; err != nil {
return nil, err
}
return groups, nil
}
// GetOpenFlareWAFRuleGroupByID returns a rule group by id.
func GetOpenFlareWAFRuleGroupByID(ctx context.Context, id uint) (*OpenFlareWAFRuleGroup, error) {
conn, err := wafDB(ctx)
if err != nil {
return nil, err
}
var group OpenFlareWAFRuleGroup
if err = conn.First(&group, id).Error; err != nil {
return nil, err
}
return &group, nil
}
// GetGlobalOpenFlareWAFRuleGroup returns the global rule group if present.
func GetGlobalOpenFlareWAFRuleGroup(ctx context.Context) (*OpenFlareWAFRuleGroup, error) {
conn, err := wafDB(ctx)
if err != nil {
return nil, err
}
var group OpenFlareWAFRuleGroup
if err = conn.Where("is_global = ?", true).Order("id asc").First(&group).Error; err != nil {
return nil, err
}
return &group, nil
}
// CreateOpenFlareWAFRuleGroup inserts a rule group.
func CreateOpenFlareWAFRuleGroup(ctx context.Context, group *OpenFlareWAFRuleGroup) error {
conn, err := wafDB(ctx)
if err != nil {
return err
}
return conn.Create(group).Error
}
// UpdateOpenFlareWAFRuleGroup updates mutable rule group fields.
func UpdateOpenFlareWAFRuleGroup(ctx context.Context, group *OpenFlareWAFRuleGroup) error {
conn, err := wafDB(ctx)
if err != nil {
return err
}
return conn.Model(&OpenFlareWAFRuleGroup{}).Where("id = ?", group.ID).Updates(map[string]any{
"name": group.Name,
"enabled": group.Enabled,
"is_global": group.IsGlobal,
"block_status_code": group.BlockStatusCode,
"block_response_body": group.BlockResponseBody,
"ip_whitelist": group.IPWhitelist,
"ip_blacklist": group.IPBlacklist,
"ip_whitelist_groups": group.IPWhitelistGroups,
"ip_blacklist_groups": group.IPBlacklistGroups,
"country_whitelist": group.CountryWhitelist,
"country_blacklist": group.CountryBlacklist,
"region_whitelist": group.RegionWhitelist,
"region_blacklist": group.RegionBlacklist,
"pow_enabled": group.PoWEnabled,
"pow_config": group.PoWConfig,
"remark": group.Remark,
}).Error
}
// DeleteOpenFlareWAFRuleGroup removes a rule group.
func DeleteOpenFlareWAFRuleGroup(ctx context.Context, id uint) error {
conn, err := wafDB(ctx)
if err != nil {
return err
}
return conn.Delete(&OpenFlareWAFRuleGroup{}, id).Error
}
// ListOpenFlareWAFIPGroups returns all IP groups.
func ListOpenFlareWAFIPGroups(ctx context.Context) ([]*OpenFlareWAFIPGroup, error) {
conn, err := wafDB(ctx)
if err != nil {
return nil, err
}
var groups []*OpenFlareWAFIPGroup
if err = conn.Order("type asc").Order("id asc").Find(&groups).Error; err != nil {
return nil, err
}
return groups, nil
}
// ListOpenFlareWAFIPGroupsByIDs returns IP groups for the given ids.
func ListOpenFlareWAFIPGroupsByIDs(ctx context.Context, ids []uint) ([]*OpenFlareWAFIPGroup, error) {
if len(ids) == 0 {
return []*OpenFlareWAFIPGroup{}, nil
}
conn, err := wafDB(ctx)
if err != nil {
return nil, err
}
var groups []*OpenFlareWAFIPGroup
if err = conn.Where("id IN ?", ids).Find(&groups).Error; err != nil {
return nil, err
}
return groups, nil
}
// GetOpenFlareWAFIPGroupByID returns an IP group by id.
func GetOpenFlareWAFIPGroupByID(ctx context.Context, id uint) (*OpenFlareWAFIPGroup, error) {
conn, err := wafDB(ctx)
if err != nil {
return nil, err
}
var group OpenFlareWAFIPGroup
if err = conn.First(&group, id).Error; err != nil {
return nil, err
}
return &group, nil
}
// CreateOpenFlareWAFIPGroup inserts an IP group.
func CreateOpenFlareWAFIPGroup(ctx context.Context, group *OpenFlareWAFIPGroup) error {
conn, err := wafDB(ctx)
if err != nil {
return err
}
return conn.Create(group).Error
}
// UpdateOpenFlareWAFIPGroup updates mutable IP group fields.
func UpdateOpenFlareWAFIPGroup(ctx context.Context, group *OpenFlareWAFIPGroup) error {
conn, err := wafDB(ctx)
if err != nil {
return err
}
return conn.Model(&OpenFlareWAFIPGroup{}).Where("id = ?", group.ID).Updates(map[string]any{
"name": group.Name,
"type": group.Type,
"enabled": group.Enabled,
"ip_list": group.IPList,
"auto_config": group.AutoConfig,
"ext_ips": group.ExtIPs,
"subscription_url": group.SubscriptionURL,
"subscription_format": group.SubscriptionFormat,
"subscription_mapping_rule": group.SubscriptionMappingRule,
"sync_interval_minutes": group.SyncIntervalMinutes,
"next_sync_at": group.NextSyncAt,
"last_sync_status": group.LastSyncStatus,
"last_sync_message": group.LastSyncMessage,
"remark": group.Remark,
}).Error
}
// ListDueOpenFlareWAFIPGroups returns enabled automatic/subscription groups due for sync.
func ListDueOpenFlareWAFIPGroups(ctx context.Context, now time.Time) ([]*OpenFlareWAFIPGroup, error) {
conn, err := wafDB(ctx)
if err != nil {
return nil, err
}
var groups []*OpenFlareWAFIPGroup
err = conn.Where(
"enabled = ? AND (type = ? OR (type = ? AND subscription_url <> '')) AND (next_sync_at IS NULL OR next_sync_at <= ?)",
true, "automatic", "subscription", now,
).Order("id asc").Find(&groups).Error
return groups, err
}
// UpdateOpenFlareWAFIPGroupSyncResult persists IP group sync outcome fields.
func UpdateOpenFlareWAFIPGroupSyncResult(ctx context.Context, group *OpenFlareWAFIPGroup) error {
conn, err := wafDB(ctx)
if err != nil {
return err
}
return conn.Model(&OpenFlareWAFIPGroup{}).Where("id = ?", group.ID).Updates(map[string]any{
"ip_list": group.IPList,
"ext_ips": group.ExtIPs,
"last_synced_at": group.LastSyncedAt,
"next_sync_at": group.NextSyncAt,
"last_sync_status": group.LastSyncStatus,
"last_sync_message": group.LastSyncMessage,
"subscription_format": group.SubscriptionFormat,
}).Error
}
// DeleteOpenFlareWAFIPGroup removes an IP group.
func DeleteOpenFlareWAFIPGroup(ctx context.Context, id uint) error {
conn, err := wafDB(ctx)
if err != nil {
return err
}
return conn.Delete(&OpenFlareWAFIPGroup{}, id).Error
}
// ListOpenFlareWAFRuleGroupBindings returns all bindings.
func ListOpenFlareWAFRuleGroupBindings(ctx context.Context) ([]OpenFlareWAFRuleGroupBinding, error) {
conn, err := wafDB(ctx)
if err != nil {
return nil, err
}
var bindings []OpenFlareWAFRuleGroupBinding
if err = conn.Order("rule_group_id asc").Order("proxy_route_id asc").Find(&bindings).Error; err != nil {
return nil, err
}
return bindings, nil
}
// ListOpenFlareWAFRuleGroupBindingsByRouteID returns bindings for a proxy route.
func ListOpenFlareWAFRuleGroupBindingsByRouteID(ctx context.Context, routeID uint) ([]OpenFlareWAFRuleGroupBinding, error) {
conn, err := wafDB(ctx)
if err != nil {
return nil, err
}
var bindings []OpenFlareWAFRuleGroupBinding
if err = conn.Where("proxy_route_id = ?", routeID).Order("rule_group_id asc").Find(&bindings).Error; err != nil {
return nil, err
}
return bindings, nil
}
// ReplaceOpenFlareWAFRuleGroupBindings replaces bindings for a rule group.
func ReplaceOpenFlareWAFRuleGroupBindings(ctx context.Context, groupID uint, routeIDs []uint) error {
conn, err := wafDB(ctx)
if err != nil {
return err
}
return conn.Transaction(func(tx *gorm.DB) error {
if err = tx.Where("rule_group_id = ?", groupID).Delete(&OpenFlareWAFRuleGroupBinding{}).Error; err != nil {
return err
}
for _, routeID := range routeIDs {
binding := OpenFlareWAFRuleGroupBinding{RuleGroupID: groupID, ProxyRouteID: routeID}
if err = tx.Create(&binding).Error; err != nil {
return err
}
}
return nil
})
}
// ReplaceOpenFlareWAFSiteRuleGroupBindings replaces bindings for a proxy route.
func ReplaceOpenFlareWAFSiteRuleGroupBindings(ctx context.Context, routeID uint, groupIDs []uint) error {
conn, err := wafDB(ctx)
if err != nil {
return err
}
return conn.Transaction(func(tx *gorm.DB) error {
if err = tx.Where("proxy_route_id = ?", routeID).Delete(&OpenFlareWAFRuleGroupBinding{}).Error; err != nil {
return err
}
for _, groupID := range groupIDs {
binding := OpenFlareWAFRuleGroupBinding{RuleGroupID: groupID, ProxyRouteID: routeID}
if err = tx.Create(&binding).Error; err != nil {
return err
}
}
return nil
})
}
// DeleteOpenFlareWAFRuleGroupBindingsByGroupID removes bindings for a rule group.
func DeleteOpenFlareWAFRuleGroupBindingsByGroupID(ctx context.Context, groupID uint) error {
conn, err := wafDB(ctx)
if err != nil {
return err
}
return conn.Where("rule_group_id = ?", groupID).Delete(&OpenFlareWAFRuleGroupBinding{}).Error
}
// DeleteOpenFlareWAFRuleGroupWithBindings removes a rule group and its bindings.
func DeleteOpenFlareWAFRuleGroupWithBindings(ctx context.Context, groupID uint) error {
conn, err := wafDB(ctx)
if err != nil {
return err
}
return conn.Transaction(func(tx *gorm.DB) error {
if err = tx.Where("rule_group_id = ?", groupID).Delete(&OpenFlareWAFRuleGroupBinding{}).Error; err != nil {
return err
}
return tx.Delete(&OpenFlareWAFRuleGroup{}, groupID).Error
})
}
// GetOpenFlareProxyRouteByID returns a proxy route by id when the table exists.
func GetOpenFlareProxyRouteByID(ctx context.Context, id uint) (*OriginProxyRoute, error) {
if !HasProxyRoutesTable(ctx) {
return nil, gorm.ErrRecordNotFound
}
conn, err := wafDB(ctx)
if err != nil {
return nil, err
}
var route OriginProxyRoute
if err = conn.First(&route, id).Error; err != nil {
return nil, err
}
return &route, nil
}
-426
View File
@@ -1,426 +0,0 @@
package model
import (
"strconv"
"strings"
"time"
"github.com/rain-kl/openflare/openflare-server/internal/common"
"github.com/rain-kl/openflare/pkg/geoip"
"gorm.io/gorm"
)
type Option struct {
Key string `json:"key" gorm:"primaryKey"`
Value string `json:"value"`
}
func AllOption() ([]*Option, error) {
var options []*Option
var err error
err = DB.Find(&options).Error
return options, err
}
func InitOptionMap() {
common.OptionMapRWMutex.Lock()
common.OptionMap = make(map[string]string)
common.OptionMap["PasswordLoginEnabled"] = strconv.FormatBool(common.PasswordLoginEnabled)
common.OptionMap["CapLoginEnabled"] = strconv.FormatBool(common.CapLoginEnabled)
common.OptionMap["PasswordRegisterEnabled"] = strconv.FormatBool(common.PasswordRegisterEnabled)
common.OptionMap["EmailVerificationEnabled"] = strconv.FormatBool(common.EmailVerificationEnabled)
common.OptionMap["GitHubOAuthEnabled"] = strconv.FormatBool(common.GitHubOAuthEnabled)
common.OptionMap["WeChatAuthEnabled"] = strconv.FormatBool(common.WeChatAuthEnabled)
common.OptionMap["SMTPServer"] = ""
common.OptionMap["SMTPPort"] = strconv.Itoa(common.SMTPPort)
common.OptionMap["SMTPAccount"] = ""
common.OptionMap["SMTPToken"] = ""
common.OptionMap["Notice"] = ""
common.OptionMap["About"] = ""
common.OptionMap["Footer"] = common.Footer
common.OptionMap["HomePageLink"] = common.HomePageLink
common.OptionMap["SystemName"] = common.SystemName
common.OptionMap["ServerAddress"] = ""
common.OptionMap["GitHubClientId"] = ""
common.OptionMap["GitHubClientSecret"] = ""
common.OptionMap["WeChatServerAddress"] = ""
common.OptionMap["WeChatServerToken"] = ""
common.OptionMap["WeChatAccountQRCodeImageURL"] = ""
common.OptionMap["AgentDiscoveryToken"] = ""
common.OptionMap["AgentHeartbeatInterval"] = strconv.Itoa(common.AgentHeartbeatInterval)
common.OptionMap["AgentWebsocketUpgradeEnabled"] = strconv.FormatBool(common.AgentWebsocketUpgradeEnabled)
common.OptionMap["NodeOfflineThreshold"] = strconv.Itoa(int(common.NodeOfflineThreshold.Milliseconds()))
common.OptionMap["AgentUpdateRepo"] = common.AgentUpdateRepo
common.OptionMap["GeoIPProvider"] = common.GeoIPProvider
common.OptionMap["DatabaseAutoCleanupEnabled"] = strconv.FormatBool(common.DatabaseAutoCleanupEnabled)
common.OptionMap["UptimeKumaEnabled"] = strconv.FormatBool(common.UptimeKumaEnabled)
common.OptionMap["UptimeKumaUrl"] = common.UptimeKumaUrl
common.OptionMap["UptimeKumaUsername"] = common.UptimeKumaUsername
common.OptionMap["UptimeKumaPassword"] = common.UptimeKumaPassword
common.OptionMap["UptimeKumaMonitorScope"] = common.UptimeKumaMonitorScope
common.OptionMap["UptimeKumaSelectedSites"] = common.UptimeKumaSelectedSites
common.OptionMap["UptimeKumaSyncInterval"] = strconv.Itoa(common.UptimeKumaSyncInterval)
common.OptionMap["UptimeKumaInterval"] = strconv.Itoa(common.UptimeKumaInterval)
common.OptionMap["UptimeKumaRetry"] = strconv.Itoa(common.UptimeKumaRetry)
common.OptionMap["UptimeKumaRetryInterval"] = strconv.Itoa(common.UptimeKumaRetryInterval)
common.OptionMap["UptimeKumaTimeout"] = strconv.Itoa(common.UptimeKumaTimeout)
common.OptionMap["DatabaseAutoCleanupRetentionDays"] = strconv.Itoa(common.DatabaseAutoCleanupRetentionDays)
common.OptionMap["OpenRestyDefaultServerReturnStatus"] = strconv.Itoa(common.OpenRestyDefaultServerReturnStatus)
common.OptionMap["OpenRestyWorkerProcesses"] = common.OpenRestyWorkerProcesses
common.OptionMap["OpenRestyWorkerConnections"] = strconv.Itoa(common.OpenRestyWorkerConnections)
common.OptionMap["OpenRestyWorkerRlimitNofile"] = strconv.Itoa(common.OpenRestyWorkerRlimitNofile)
common.OptionMap["OpenRestyEventsUse"] = common.OpenRestyEventsUse
common.OptionMap["OpenRestyEventsMultiAcceptEnabled"] = strconv.FormatBool(common.OpenRestyEventsMultiAcceptEnabled)
common.OptionMap["OpenRestyKeepaliveTimeout"] = strconv.Itoa(common.OpenRestyKeepaliveTimeout)
common.OptionMap["OpenRestyKeepaliveRequests"] = strconv.Itoa(common.OpenRestyKeepaliveRequests)
common.OptionMap["OpenRestyClientHeaderTimeout"] = strconv.Itoa(common.OpenRestyClientHeaderTimeout)
common.OptionMap["OpenRestyClientBodyTimeout"] = strconv.Itoa(common.OpenRestyClientBodyTimeout)
common.OptionMap["OpenRestyClientMaxBodySize"] = common.OpenRestyClientMaxBodySize
common.OptionMap["OpenRestyLargeClientHeaderBuffers"] = common.OpenRestyLargeClientHeaderBuffers
common.OptionMap["OpenRestySendTimeout"] = strconv.Itoa(common.OpenRestySendTimeout)
common.OptionMap["OpenRestyProxyConnectTimeout"] = strconv.Itoa(common.OpenRestyProxyConnectTimeout)
common.OptionMap["OpenRestyProxySendTimeout"] = strconv.Itoa(common.OpenRestyProxySendTimeout)
common.OptionMap["OpenRestyProxyReadTimeout"] = strconv.Itoa(common.OpenRestyProxyReadTimeout)
common.OptionMap["OpenRestyWebsocketEnabled"] = strconv.FormatBool(common.OpenRestyWebsocketEnabled)
common.OptionMap["OpenRestyHTTP3Enabled"] = strconv.FormatBool(common.OpenRestyHTTP3Enabled)
common.OptionMap["OpenRestyProxyRequestBufferingEnabled"] = strconv.FormatBool(common.OpenRestyProxyRequestBufferingEnabled)
common.OptionMap["OpenRestyProxyBufferingEnabled"] = strconv.FormatBool(common.OpenRestyProxyBufferingEnabled)
common.OptionMap["OpenRestyProxyBuffers"] = common.OpenRestyProxyBuffers
common.OptionMap["OpenRestyProxyBufferSize"] = common.OpenRestyProxyBufferSize
common.OptionMap["OpenRestyProxyBusyBuffersSize"] = common.OpenRestyProxyBusyBuffersSize
common.OptionMap["OpenRestyGzipEnabled"] = strconv.FormatBool(common.OpenRestyGzipEnabled)
common.OptionMap["OpenRestyGzipMinLength"] = strconv.Itoa(common.OpenRestyGzipMinLength)
common.OptionMap["OpenRestyGzipCompLevel"] = strconv.Itoa(common.OpenRestyGzipCompLevel)
common.OptionMap["OpenRestyCacheEnabled"] = strconv.FormatBool(common.OpenRestyCacheEnabled)
common.OptionMap["OpenRestyCachePath"] = common.OpenRestyCachePath
common.OptionMap["OpenRestyCacheLevels"] = common.OpenRestyCacheLevels
common.OptionMap["OpenRestyCacheInactive"] = common.OpenRestyCacheInactive
common.OptionMap["OpenRestyCacheMaxSize"] = common.OpenRestyCacheMaxSize
common.OptionMap["OpenRestyCacheKeyTemplate"] = common.OpenRestyCacheKeyTemplate
common.OptionMap["OpenRestyCacheLockEnabled"] = strconv.FormatBool(common.OpenRestyCacheLockEnabled)
common.OptionMap["OpenRestyCacheLockTimeout"] = common.OpenRestyCacheLockTimeout
common.OptionMap["OpenRestyCacheUseStale"] = common.OpenRestyCacheUseStale
common.OptionMap["OpenRestyMainConfigTemplate"] = common.OpenRestyMainConfigTemplate
common.OptionMap["GlobalApiRateLimitNum"] = strconv.Itoa(common.GlobalApiRateLimitNum)
common.OptionMap["GlobalApiRateLimitDuration"] = strconv.FormatInt(common.GlobalApiRateLimitDuration, 10)
common.OptionMap["GlobalWebRateLimitNum"] = strconv.Itoa(common.GlobalWebRateLimitNum)
common.OptionMap["GlobalWebRateLimitDuration"] = strconv.FormatInt(common.GlobalWebRateLimitDuration, 10)
common.OptionMap["CriticalRateLimitNum"] = strconv.Itoa(common.CriticalRateLimitNum)
common.OptionMap["CriticalRateLimitDuration"] = strconv.FormatInt(common.CriticalRateLimitDuration, 10)
common.OptionMapRWMutex.Unlock()
options, _ := AllOption()
for _, option := range options {
updateOptionMap(option.Key, option.Value)
}
}
func UpdateOption(key string, value string) error {
return UpdateOptions([]Option{{
Key: key,
Value: value,
}})
}
func UpdateOptions(options []Option) error {
if len(options) == 0 {
return nil
}
if err := DB.Transaction(func(tx *gorm.DB) error {
for _, item := range options {
if item.Key == "UptimeKumaPassword" && strings.TrimSpace(item.Value) == "" {
continue
}
option := Option{
Key: item.Key,
}
if err := tx.FirstOrCreate(&option, Option{Key: item.Key}).Error; err != nil {
return err
}
option.Value = item.Value
if err := tx.Save(&option).Error; err != nil {
return err
}
}
return nil
}); err != nil {
return err
}
for _, item := range options {
if item.Key == "UptimeKumaPassword" && strings.TrimSpace(item.Value) == "" {
continue
}
updateOptionMap(item.Key, item.Value)
}
return nil
}
func updateOptionMap(key string, value string) {
shouldRefreshGeoIP := false
common.OptionMapRWMutex.Lock()
if common.OptionMap == nil {
common.OptionMap = make(map[string]string)
}
common.OptionMap[key] = value
if strings.HasSuffix(key, "Enabled") {
boolValue := value == "true"
switch key {
case "PasswordRegisterEnabled":
common.PasswordRegisterEnabled = boolValue
case "PasswordLoginEnabled":
common.PasswordLoginEnabled = boolValue
case "CapLoginEnabled":
common.CapLoginEnabled = boolValue
case "EmailVerificationEnabled":
common.EmailVerificationEnabled = boolValue
case "GitHubOAuthEnabled":
common.GitHubOAuthEnabled = boolValue
case "WeChatAuthEnabled":
common.WeChatAuthEnabled = boolValue
}
}
switch key {
case "SMTPServer":
common.SMTPServer = value
case "SMTPPort":
intValue, _ := strconv.Atoi(value)
common.SMTPPort = intValue
case "SMTPAccount":
common.SMTPAccount = value
case "SMTPToken":
common.SMTPToken = value
case "ServerAddress":
common.ServerAddress = value
case "GitHubClientId":
common.GitHubClientId = value
case "GitHubClientSecret":
common.GitHubClientSecret = value
case "Footer":
common.Footer = value
case "HomePageLink":
common.HomePageLink = value
case "SystemName":
common.SystemName = value
case "WeChatServerAddress":
common.WeChatServerAddress = value
case "WeChatServerToken":
common.WeChatServerToken = value
case "WeChatAccountQRCodeImageURL":
common.WeChatAccountQRCodeImageURL = value
case "AgentDiscoveryToken":
common.AgentDiscoveryToken = value
case "AgentHeartbeatInterval":
if v, err := strconv.Atoi(value); err == nil && v > 0 {
common.AgentHeartbeatInterval = v
}
case "AgentWebsocketUpgradeEnabled":
common.AgentWebsocketUpgradeEnabled = value == "true"
case "NodeOfflineThreshold":
if v, err := strconv.Atoi(value); err == nil && v > 0 {
common.NodeOfflineThreshold = time.Duration(v) * time.Millisecond
}
case "AgentUpdateRepo":
if value != "" {
common.AgentUpdateRepo = value
}
case "GeoIPProvider":
if geoip.IsValidProvider(value) {
common.GeoIPProvider = value
shouldRefreshGeoIP = true
}
case "UptimeKumaEnabled":
common.UptimeKumaEnabled = value == "true"
case "UptimeKumaUrl":
common.UptimeKumaUrl = value
case "UptimeKumaUsername":
common.UptimeKumaUsername = value
case "UptimeKumaPassword":
common.UptimeKumaPassword = value
case "UptimeKumaMonitorScope":
common.UptimeKumaMonitorScope = value
case "UptimeKumaSelectedSites":
common.UptimeKumaSelectedSites = value
case "UptimeKumaSyncInterval":
if v, err := strconv.Atoi(value); err == nil && v > 0 {
common.UptimeKumaSyncInterval = v
}
case "UptimeKumaInterval":
if v, err := strconv.Atoi(value); err == nil && v > 0 {
common.UptimeKumaInterval = v
}
case "UptimeKumaRetry":
if v, err := strconv.Atoi(value); err == nil && v >= 0 {
common.UptimeKumaRetry = v
}
case "UptimeKumaRetryInterval":
if v, err := strconv.Atoi(value); err == nil && v > 0 {
common.UptimeKumaRetryInterval = v
}
case "UptimeKumaTimeout":
if v, err := strconv.Atoi(value); err == nil && v > 0 {
common.UptimeKumaTimeout = v
}
case "DatabaseAutoCleanupEnabled":
common.DatabaseAutoCleanupEnabled = value == "true"
case "DatabaseAutoCleanupRetentionDays":
if v, err := strconv.Atoi(value); err == nil && v >= 1 {
common.DatabaseAutoCleanupRetentionDays = v
}
case "OpenRestyDefaultServerReturnStatus":
if v, err := strconv.Atoi(value); err == nil && v >= 100 && v <= 999 {
common.OpenRestyDefaultServerReturnStatus = v
}
case "OpenRestyWorkerProcesses":
if strings.TrimSpace(value) != "" {
common.OpenRestyWorkerProcesses = value
}
case "OpenRestyWorkerConnections":
if v, err := strconv.Atoi(value); err == nil && v > 0 {
common.OpenRestyWorkerConnections = v
}
case "OpenRestyWorkerRlimitNofile":
if v, err := strconv.Atoi(value); err == nil && v > 0 {
common.OpenRestyWorkerRlimitNofile = v
}
case "OpenRestyEventsUse":
common.OpenRestyEventsUse = value
case "OpenRestyResolvers":
common.OpenRestyResolvers = value
case "OpenRestyEventsMultiAcceptEnabled":
common.OpenRestyEventsMultiAcceptEnabled = value == "true"
case "OpenRestyKeepaliveTimeout":
if v, err := strconv.Atoi(value); err == nil && v > 0 {
common.OpenRestyKeepaliveTimeout = v
}
case "OpenRestyKeepaliveRequests":
if v, err := strconv.Atoi(value); err == nil && v > 0 {
common.OpenRestyKeepaliveRequests = v
}
case "OpenRestyClientHeaderTimeout":
if v, err := strconv.Atoi(value); err == nil && v > 0 {
common.OpenRestyClientHeaderTimeout = v
}
case "OpenRestyClientBodyTimeout":
if v, err := strconv.Atoi(value); err == nil && v > 0 {
common.OpenRestyClientBodyTimeout = v
}
case "OpenRestyClientMaxBodySize":
if strings.TrimSpace(value) != "" {
common.OpenRestyClientMaxBodySize = value
}
case "OpenRestyLargeClientHeaderBuffers":
if strings.TrimSpace(value) != "" {
common.OpenRestyLargeClientHeaderBuffers = value
}
case "OpenRestySendTimeout":
if v, err := strconv.Atoi(value); err == nil && v > 0 {
common.OpenRestySendTimeout = v
}
case "OpenRestyProxyConnectTimeout":
if v, err := strconv.Atoi(value); err == nil && v > 0 {
common.OpenRestyProxyConnectTimeout = v
}
case "OpenRestyProxySendTimeout":
if v, err := strconv.Atoi(value); err == nil && v > 0 {
common.OpenRestyProxySendTimeout = v
}
case "OpenRestyProxyReadTimeout":
if v, err := strconv.Atoi(value); err == nil && v > 0 {
common.OpenRestyProxyReadTimeout = v
}
case "OpenRestyWebsocketEnabled":
common.OpenRestyWebsocketEnabled = value == "true"
case "OpenRestyHTTP3Enabled":
common.OpenRestyHTTP3Enabled = value == "true"
case "OpenRestyProxyRequestBufferingEnabled":
common.OpenRestyProxyRequestBufferingEnabled = value == "true"
case "OpenRestyProxyBufferingEnabled":
common.OpenRestyProxyBufferingEnabled = value == "true"
case "OpenRestyProxyBuffers":
if strings.TrimSpace(value) != "" {
common.OpenRestyProxyBuffers = value
}
case "OpenRestyProxyBufferSize":
if strings.TrimSpace(value) != "" {
common.OpenRestyProxyBufferSize = value
}
case "OpenRestyProxyBusyBuffersSize":
if strings.TrimSpace(value) != "" {
common.OpenRestyProxyBusyBuffersSize = value
}
case "OpenRestyGzipEnabled":
common.OpenRestyGzipEnabled = value == "true"
case "OpenRestyGzipMinLength":
if v, err := strconv.Atoi(value); err == nil && v > 0 {
common.OpenRestyGzipMinLength = v
}
case "OpenRestyGzipCompLevel":
if v, err := strconv.Atoi(value); err == nil && v > 0 {
common.OpenRestyGzipCompLevel = v
}
case "OpenRestyCacheEnabled":
common.OpenRestyCacheEnabled = value == "true"
case "OpenRestyCachePath":
common.OpenRestyCachePath = value
case "OpenRestyCacheLevels":
if strings.TrimSpace(value) != "" {
common.OpenRestyCacheLevels = value
}
case "OpenRestyCacheInactive":
if strings.TrimSpace(value) != "" {
common.OpenRestyCacheInactive = value
}
case "OpenRestyCacheMaxSize":
if strings.TrimSpace(value) != "" {
common.OpenRestyCacheMaxSize = value
}
case "OpenRestyCacheKeyTemplate":
if strings.TrimSpace(value) != "" {
common.OpenRestyCacheKeyTemplate = value
}
case "OpenRestyCacheLockEnabled":
common.OpenRestyCacheLockEnabled = value == "true"
case "OpenRestyCacheLockTimeout":
if strings.TrimSpace(value) != "" {
common.OpenRestyCacheLockTimeout = value
}
case "OpenRestyCacheUseStale":
if strings.TrimSpace(value) != "" {
common.OpenRestyCacheUseStale = value
}
case "OpenRestyMainConfigTemplate":
if strings.TrimSpace(value) != "" {
common.OpenRestyMainConfigTemplate = value
}
case "GlobalApiRateLimitNum":
if v, err := strconv.Atoi(value); err == nil && v > 0 {
common.GlobalApiRateLimitNum = v
}
case "GlobalApiRateLimitDuration":
if v, err := strconv.ParseInt(value, 10, 64); err == nil && v > 0 {
common.GlobalApiRateLimitDuration = v
}
case "GlobalWebRateLimitNum":
if v, err := strconv.Atoi(value); err == nil && v > 0 {
common.GlobalWebRateLimitNum = v
}
case "GlobalWebRateLimitDuration":
if v, err := strconv.ParseInt(value, 10, 64); err == nil && v > 0 {
common.GlobalWebRateLimitDuration = v
}
case "CriticalRateLimitNum":
if v, err := strconv.Atoi(value); err == nil && v > 0 {
common.CriticalRateLimitNum = v
}
case "CriticalRateLimitDuration":
if v, err := strconv.ParseInt(value, 10, 64); err == nil && v > 0 {
common.CriticalRateLimitDuration = v
}
}
common.OptionMapRWMutex.Unlock()
if shouldRefreshGeoIP {
geoip.InitGeoIP(common.GeoIPProvider)
}
}
-56
View File
@@ -1,56 +0,0 @@
package model
import "time"
type Origin struct {
ID uint `json:"id" gorm:"primaryKey"`
Name string `json:"name" gorm:"size:255;not null"`
Address string `json:"address" gorm:"uniqueIndex;size:255;not null"`
Remark string `json:"remark" gorm:"size:255"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
}
type OriginRouteCount struct {
OriginID uint `json:"origin_id"`
RouteCount int64 `json:"route_count"`
}
func ListOrigins() (origins []*Origin, err error) {
err = DB.Order("id desc").Find(&origins).Error
return origins, err
}
func GetOriginByID(id uint) (*Origin, error) {
origin := &Origin{}
err := DB.First(origin, id).Error
return origin, err
}
func GetOriginByAddress(address string) (*Origin, error) {
origin := &Origin{}
err := DB.Where("address = ?", address).First(origin).Error
return origin, err
}
func ListOriginRouteCounts() ([]OriginRouteCount, error) {
result := make([]OriginRouteCount, 0)
err := DB.Model(&ProxyRoute{}).
Select("origin_id, COUNT(*) AS route_count").
Where("origin_id IS NOT NULL").
Group("origin_id").
Scan(&result).Error
return result, err
}
func (origin *Origin) Insert() error {
return DB.Create(origin).Error
}
func (origin *Origin) Update() error {
return DB.Save(origin).Error
}
func (origin *Origin) Delete() error {
return DB.Delete(origin).Error
}
-83
View File
@@ -1,83 +0,0 @@
package model
import "time"
const (
PagesDeploymentStatusUploaded = "uploaded"
PagesDeploymentStatusActive = "active"
)
type PagesProject struct {
ID uint `json:"id" gorm:"primaryKey"`
Name string `json:"name" gorm:"size:255;not null"`
Slug string `json:"slug" gorm:"uniqueIndex;size:128;not null"`
Description string `json:"description" gorm:"type:text;not null;default:''"`
Enabled bool `json:"enabled" gorm:"not null;default:true"`
SPAFallbackEnabled bool `json:"spa_fallback_enabled" gorm:"not null;default:false"`
SPAFallbackPath string `json:"spa_fallback_path" gorm:"size:512;not null;default:'/index.html'"`
APIProxyEnabled bool `json:"api_proxy_enabled" gorm:"not null;default:false"`
APIProxyPath string `json:"api_proxy_path" gorm:"size:255;not null;default:''"`
APIProxyPass string `json:"api_proxy_pass" gorm:"size:2048;not null;default:''"`
APIProxyRewrite string `json:"api_proxy_rewrite" gorm:"size:255;not null;default:''"`
ActiveDeploymentID *uint `json:"active_deployment_id" gorm:"index"`
RootDir string `json:"root_dir" gorm:"size:512;not null;default:''"`
EntryFile string `json:"entry_file" gorm:"size:512;not null;default:'index.html'"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
}
type PagesDeployment struct {
ID uint `json:"id" gorm:"primaryKey"`
ProjectID uint `json:"project_id" gorm:"not null;index"`
DeploymentNumber int `json:"deployment_number" gorm:"not null"`
Checksum string `json:"checksum" gorm:"size:64;not null;index"`
Status string `json:"status" gorm:"size:32;not null;default:'uploaded';index"`
ArtifactPath string `json:"artifact_path" gorm:"size:2048;not null"`
FileCount int `json:"file_count" gorm:"not null;default:0"`
TotalSize int64 `json:"total_size" gorm:"not null;default:0"`
CreatedBy string `json:"created_by" gorm:"size:64;not null;default:''"`
CreatedAt time.Time `json:"created_at"`
ActivatedAt *time.Time `json:"activated_at"`
}
type PagesDeploymentFile struct {
ID uint `json:"id" gorm:"primaryKey"`
DeploymentID uint `json:"deployment_id" gorm:"not null;index"`
Path string `json:"path" gorm:"size:2048;not null"`
Size int64 `json:"size" gorm:"not null;default:0"`
Checksum string `json:"checksum" gorm:"size:64;not null"`
CreatedAt time.Time `json:"created_at"`
}
func ListPagesProjects() (projects []*PagesProject, err error) {
err = DB.Order("id desc").Find(&projects).Error
return projects, err
}
func GetPagesProjectByID(id uint) (*PagesProject, error) {
project := &PagesProject{}
err := DB.First(project, id).Error
return project, err
}
func GetPagesProjectBySlug(slug string) (*PagesProject, error) {
project := &PagesProject{}
err := DB.Where("slug = ?", slug).First(project).Error
return project, err
}
func ListPagesDeployments(projectID uint) (deployments []*PagesDeployment, err error) {
err = DB.Where("project_id = ?", projectID).Order("id desc").Find(&deployments).Error
return deployments, err
}
func GetPagesDeploymentByID(id uint) (*PagesDeployment, error) {
deployment := &PagesDeployment{}
err := DB.First(deployment, id).Error
return deployment, err
}
func ListPagesDeploymentFiles(deploymentID uint) (files []*PagesDeploymentFile, err error) {
err = DB.Where("deployment_id = ?", deploymentID).Order("path asc").Find(&files).Error
return files, err
}
@@ -0,0 +1,99 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package model
import (
"encoding/json"
"errors"
"regexp"
"strings"
"time"
)
const (
// TypeCustom 自定义消息通道类型
TypeCustom = "custom"
// TypeEmail 邮件推送消息通道类型
TypeEmail = "email"
// TypeTelegram 电报机器人推送消息通道类型
TypeTelegram = "telegram"
)
// PushChannel 消息通道模型
type PushChannel struct {
ID uint64 `json:"id" gorm:"primaryKey;autoIncrement"`
Name string `json:"name" gorm:"uniqueIndex;size:80;not null"` // 通道名称,仅英文字母和下划线,唯一
Description string `json:"description" gorm:"size:255"` // 备注
Type string `json:"type" gorm:"size:50;not null;default:'custom'"` // 通道类型:custom, lark, email
Token string `json:"token" gorm:"size:100"` // 鉴权令牌或发信用户名等
URL string `json:"url" gorm:"type:text;not null"` // 请求地址,HTTPS 协议或 SMTP 地址
Other string `json:"other" gorm:"type:text;not null"` // 请求体/SMTP 密码等
Enabled bool `json:"enabled" gorm:"index;not null;default:true"` // 通道是否启用
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime;index"`
UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime;index"`
}
// TableName 指定 GORM 表名
func (PushChannel) TableName() string {
return "w_push_channels"
}
var nameRegex = regexp.MustCompile(`^[a-zA-Z0-9_]+$`)
// Validate 参数合法性与 JSON 格式校验
func (pc *PushChannel) Validate() error {
pc.Name = strings.TrimSpace(pc.Name)
pc.URL = strings.TrimSpace(pc.URL)
pc.Other = strings.TrimSpace(pc.Other)
pc.Type = strings.TrimSpace(pc.Type)
if pc.Type == "" {
pc.Type = TypeCustom
}
if pc.Type == TypeTelegram && pc.URL == "" {
pc.URL = "https://api.telegram.org"
}
if pc.Name == "" {
return errors.New("channel name is required")
}
if !nameRegex.MatchString(pc.Name) {
return errors.New("channel name can only contain letters, numbers, and underscores")
}
if pc.Type != TypeEmail && pc.URL == "" {
return errors.New("request URL/address is required")
}
if pc.Type != TypeEmail && !strings.HasPrefix(pc.URL, "https://") {
return errors.New("request URL must use HTTPS protocol for security reasons")
}
switch pc.Type {
case TypeCustom:
if pc.Other == "" {
return errors.New("payload schema (request body) is required")
}
return validateJSON(pc.Other)
case TypeEmail:
// Email channel SMTP configs fall back to global settings, so they are not required to be filled.
case TypeTelegram:
if pc.Token == "" {
return errors.New("telegram bot token is required")
}
}
return nil
}
func validateJSON(s string) error {
var jsonTest map[string]any
if err := json.Unmarshal([]byte(s), &jsonTest); err == nil {
return nil
}
var jsonArr []any
if err := json.Unmarshal([]byte(s), &jsonArr); err == nil {
return nil
}
return errors.New("payload schema must be a valid JSON format")
}
@@ -0,0 +1,50 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package model
import (
"errors"
"strings"
"time"
)
// PushEvent 系统通知事件模型
type PushEvent struct {
ID uint64 `json:"id" gorm:"primaryKey;autoIncrement"`
EventKey string `json:"event_key" gorm:"uniqueIndex;size:80;not null"` // 如 admin_login
Name string `json:"name" gorm:"size:100;not null"` // 如 管理员登录
TaskType string `json:"task_type" gorm:"size:100;index;not null;default:''"` // 关联的异步任务类型
Channels []string `json:"channels" gorm:"type:text;serializer:json;not null"` // 推送渠道列表,如 ["lark"]
Targets []string `json:"targets" gorm:"type:text;serializer:json;not null"` // 推送目标用户/邮箱列表
Template string `json:"template" gorm:"type:text;not null"` // 消息模板 JSON
Enabled bool `json:"enabled" gorm:"index;not null;default:false"` // 是否启用
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime;index"`
UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime;index"`
}
// TableName 指定 GORM 表名
func (PushEvent) TableName() string {
return "w_push_events"
}
// Validate 基础校验
func (pe *PushEvent) Validate() error {
pe.EventKey = strings.TrimSpace(pe.EventKey)
pe.Name = strings.TrimSpace(pe.Name)
pe.Template = strings.TrimSpace(pe.Template)
if pe.EventKey == "" {
return errors.New("event key is required")
}
if pe.Name == "" {
return errors.New("event name is required")
}
if pe.Template == "" {
return errors.New("event template is required")
}
if pe.Enabled && len(pe.Channels) == 0 {
return errors.New("cannot enable event without any push channels configured")
}
return nil
}
@@ -0,0 +1,27 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package model
import (
"time"
)
// PushHistory 推送日志/历史实体
type PushHistory struct {
ID uint64 `json:"id" gorm:"primaryKey;autoIncrement"`
EventKey string `json:"event_key" gorm:"size:80;not null;index"`
Channel string `json:"channel" gorm:"size:50;not null"`
Target string `json:"target" gorm:"size:255;not null"`
Title string `json:"title" gorm:"size:255;not null"`
Content string `json:"content" gorm:"type:text;not null"`
Level string `json:"level" gorm:"size:20;not null"`
Status string `json:"status" gorm:"size:20;not null"` // success / failed
ErrorMsg string `json:"error_msg" gorm:"type:text"`
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime;index"`
}
// TableName 指定表名
func (PushHistory) TableName() string {
return "w_push_histories"
}
@@ -0,0 +1,70 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package model
import (
"context"
"time"
"github.com/Rain-kl/Wavelet/internal/db"
)
// Schedule 定时任务配置表
type Schedule struct {
ID uint64 `json:"id,string" gorm:"primaryKey"`
Name string `json:"name" gorm:"size:128;not null"`
TaskType string `json:"task_type" gorm:"size:64;not null"`
Cron string `json:"cron" gorm:"size:64;not null"`
Payload string `json:"payload" gorm:"type:text"`
IsActive bool `json:"is_active" gorm:"not null;default:true"`
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime;index"`
UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime"`
}
// TableName 表名
func (Schedule) TableName() string {
return "w_schedules"
}
// CreateSchedule 创建定时任务
func CreateSchedule(ctx context.Context, schedule *Schedule) error {
return db.DB(ctx).Create(schedule).Error
}
// UpdateSchedule 更新定时任务
func UpdateSchedule(ctx context.Context, schedule *Schedule) error {
return db.DB(ctx).Save(schedule).Error
}
// DeleteSchedule 删除定时任务
func DeleteSchedule(ctx context.Context, id uint64) error {
return db.DB(ctx).Delete(&Schedule{}, id).Error
}
// GetScheduleByID 根据 ID 获取定时任务
func GetScheduleByID(ctx context.Context, id uint64) (*Schedule, error) {
var schedule Schedule
if err := db.DB(ctx).Where("id = ?", id).First(&schedule).Error; err != nil {
return nil, err
}
return &schedule, nil
}
// ListSchedules 获取所有定时任务
func ListSchedules(ctx context.Context) ([]Schedule, error) {
var schedules []Schedule
if err := db.DB(ctx).Order("id DESC").Find(&schedules).Error; err != nil {
return nil, err
}
return schedules, nil
}
// ListActiveSchedules 获取所有启用的定时任务
func ListActiveSchedules(ctx context.Context) ([]Schedule, error) {
var schedules []Schedule
if err := db.DB(ctx).Where("is_active = ?", true).Find(&schedules).Error; err != nil {
return nil, err
}
return schedules, nil
}
@@ -1,60 +0,0 @@
package model
import (
schemagoose "github.com/rain-kl/openflare/openflare-server/internal/model/goose"
"gorm.io/gorm"
)
func currentGooseTargetVersion() int64 {
return schemagoose.CurrentTargetVersion()
}
func loadGooseDatabaseVersion(db *gorm.DB) (int, bool, error) {
return schemagoose.LoadDatabaseVersion(db)
}
func ensureDatabaseSchemaUpToDate(db *gorm.DB, backend string) error {
return schemagoose.EnsureDatabaseSchemaUpToDate(db, backend, databaseSchemaMigrationContext{})
}
func (databaseSchemaMigrationContext) RegisterSharding(db *gorm.DB, backend string) error {
return registerSharding(db, backend)
}
func (databaseSchemaMigrationContext) AutoMigrateLegacySchemaMetadata(db *gorm.DB) error {
return autoMigrateLegacySchemaMetadata(db)
}
func (databaseSchemaMigrationContext) InitializeFreshDatabaseSchema(db *gorm.DB, backend string) error {
return initializeFreshDatabaseSchema(db, backend)
}
func (databaseSchemaMigrationContext) IsDatabaseEmpty(db *gorm.DB) (bool, error) {
return isDatabaseEmpty(db)
}
func (databaseSchemaMigrationContext) RepairCurrentSchemaState(db *gorm.DB, backend string) error {
if err := dropLegacyNodeColumns(db, backend); err != nil {
return err
}
if err := ensureDefaultGitHubAuthSource(db); err != nil {
return err
}
if err := ensureDefaultWAFRuleGroup(db); err != nil {
return err
}
return nil
}
func (databaseSchemaMigrationContext) SaveLegacyDatabaseSchemaVersion(db *gorm.DB, version int) error {
return saveLegacyDatabaseSchemaVersion(db, version)
}
func (databaseSchemaMigrationContext) UpgradeLegacyDatabaseSchema(db *gorm.DB, backend string, version int) error {
return upgradeLegacyDatabaseSchema(db, backend, version)
}
func (databaseSchemaMigrationContext) ValidateCurrentDatabaseSchema(db *gorm.DB, backend string) error {
return validateCurrentDatabaseSchema(db, backend)
}
-208
View File
@@ -1,208 +0,0 @@
package model
import (
"fmt"
"strconv"
"strings"
"sync"
"github.com/bwmarrin/snowflake"
"gorm.io/gorm"
"gorm.io/sharding"
)
const observabilityShardCount = 10
var (
observabilityIDNode *snowflake.Node
observabilityIDNodeErr error
observabilityIDNodeOnce sync.Once
)
func registerSharding(db *gorm.DB, backend string) error {
if db == nil {
return nil
}
_ = backend
if err := db.Use(sharding.Register(sharding.Config{
ShardingKey: "id",
NumberOfShards: observabilityShardCount,
ShardingAlgorithm: func(value any) (string, error) {
return observabilityShardSuffixForValue(value)
},
ShardingAlgorithmByPrimaryKey: func(id int64) string {
return observabilityShardSuffixForInt64(id)
},
PrimaryKeyGenerator: sharding.PKCustom,
PrimaryKeyGeneratorFn: func(tableIdx int64) int64 {
return 0
},
}, shardedObservabilityTables()...)); err != nil {
return fmt.Errorf("register observability sharding failed: %w", err)
}
return nil
}
func shardedObservabilityTables() []any {
return []any{
&NodeMetricSnapshot{},
&NodeRequestReport{},
&NodeAccessLog{},
&NodeObservationOpenresty{},
&NodeObservationFrps{},
&NodeObservationFrpc{},
}
}
func shardedObservabilityBaseTables() []string {
return []string{
"node_metric_snapshots",
"node_request_reports",
"node_access_logs",
"node_observation_openresties",
"node_observation_frps",
"node_observation_frpcs",
}
}
func isShardedObservabilityTable(tableName string) bool {
switch strings.TrimSpace(tableName) {
case "node_metric_snapshots", "node_request_reports", "node_access_logs", "node_observation_openresties", "node_observation_frps", "node_observation_frpcs":
return true
default:
return false
}
}
func observabilityShardTables(baseTable string) []string {
tables := make([]string, 0, observabilityShardCount)
for _, suffix := range observabilityShardSuffixes() {
tables = append(tables, baseTable+suffix)
}
return tables
}
func observabilityShardSuffixes() []string {
suffixes := make([]string, 0, observabilityShardCount)
for index := 0; index < observabilityShardCount; index++ {
suffixes = append(suffixes, fmt.Sprintf("_%02d", index))
}
return suffixes
}
func observabilityShardSuffixForID(id uint) string {
return fmt.Sprintf("_%02d", uint64(id)%uint64(observabilityShardCount))
}
func observabilityShardSuffixForInt64(id int64) string {
if id < 0 {
id = -id
}
return fmt.Sprintf("_%02d", uint64(id)%uint64(observabilityShardCount))
}
func observabilityShardSuffixForValue(value any) (string, error) {
switch typed := value.(type) {
case int:
return observabilityShardSuffixForInt64(int64(typed)), nil
case int8:
return observabilityShardSuffixForInt64(int64(typed)), nil
case int16:
return observabilityShardSuffixForInt64(int64(typed)), nil
case int32:
return observabilityShardSuffixForInt64(int64(typed)), nil
case int64:
return observabilityShardSuffixForInt64(typed), nil
case uint:
return observabilityShardSuffixForID(typed), nil
case uint8:
return observabilityShardSuffixForID(uint(typed)), nil
case uint16:
return observabilityShardSuffixForID(uint(typed)), nil
case uint32:
return observabilityShardSuffixForID(uint(typed)), nil
case uint64:
return fmt.Sprintf("_%02d", typed%uint64(observabilityShardCount)), nil
case string:
id, err := strconv.ParseUint(strings.TrimSpace(typed), 10, 64)
if err != nil {
return "", fmt.Errorf("invalid sharding id %q", typed)
}
return fmt.Sprintf("_%02d", id%uint64(observabilityShardCount)), nil
default:
return "", fmt.Errorf("unsupported observability sharding value type %T", value)
}
}
func legacyObservabilityShardTableName(tableName string) string {
return tableName + "_legacy_v2_to_v3"
}
func normalizeShardedDB(db *gorm.DB) *gorm.DB {
if db != nil {
return db
}
return DB
}
func nextObservabilityID() (uint, error) {
observabilityIDNodeOnce.Do(func() {
observabilityIDNode, observabilityIDNodeErr = snowflake.NewNode(0)
})
if observabilityIDNodeErr != nil {
return 0, observabilityIDNodeErr
}
id := observabilityIDNode.Generate().Int64()
if id <= 0 {
return 0, fmt.Errorf("generated invalid observability id %d", id)
}
return uint(id), nil
}
func assignObservabilityID(id *uint) error {
if id == nil || *id != 0 {
return nil
}
generated, err := nextObservabilityID()
if err != nil {
return err
}
*id = generated
return nil
}
func queryAcrossShards[T any](baseTable string, query func(tx *gorm.DB) ([]T, error)) ([]T, error) {
return queryAcrossShardsWithDB(DB, baseTable, query)
}
func queryAcrossShardsWithDB[T any](db *gorm.DB, baseTable string, query func(tx *gorm.DB) ([]T, error)) ([]T, error) {
items := make([]T, 0)
db = normalizeShardedDB(db)
for _, table := range observabilityShardTables(baseTable) {
rows, err := query(db.Table(table))
if err != nil {
return nil, err
}
items = append(items, rows...)
}
return items, nil
}
func deleteAcrossShards(db *gorm.DB, baseTable string, model any, apply func(tx *gorm.DB) *gorm.DB) (int64, error) {
db = normalizeShardedDB(db)
var deleted int64
for _, table := range observabilityShardTables(baseTable) {
tx := db.Table(table)
if apply != nil {
tx = apply(tx)
} else {
tx = tx.Session(&gorm.Session{AllowGlobalUpdate: true})
}
result := tx.Delete(model)
if result.Error != nil {
return deleted, result.Error
}
deleted += result.RowsAffected
}
return deleted, nil
}
@@ -0,0 +1,63 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package model
import "time"
// 配置键常量 - 所有系统配置的 key 定义
const (
ConfigKeyUploadAllowedExtensions = "upload_allowed_extensions" // 允许上传的文件扩展名,逗号分隔
ConfigKeySiteName = "site_name" // 站点名称
ConfigKeyPasswordLoginEnabled = "password_login_enabled" // 是否允许密码登录
ConfigKeyRegistrationEnabled = "registration_enabled" // 是否允许注册
ConfigKeyPasswordRegisterEnabled = "password_register_enabled" // 是否允许密码注册
ConfigKeyOIDCLoginEnabled = "oidc_login_enabled" // 是否允许 OIDC 登录
ConfigKeyMaxAPIKeysPerUser = "max_api_keys_per_user" //nolint:gosec // false positive: config key name. 每个用户最大 API Key 数量
ConfigKeyCapLoginEnabled = "cap_login_enabled" // 是否启用登录人机验证
ConfigKeyCapAutoSolve = "cap_auto_solve" // 打开页面后是否自动开始计算(false 则需用户手动点击)
ConfigKeyCapChallengeCount = "cap_challenge_count" // 客户端需求解的 PoW 难题总数,默认 1,推荐 1~5
ConfigKeyCapChallengeSize = "cap_challenge_size" // 人机验证盐值长度
ConfigKeyCapChallengeDifficulty = "cap_challenge_difficulty" // 人机验证 PoW 难度(目标前缀长度)
ConfigKeyCapChallengeTTL = "cap_challenge_ttl_seconds" // 人机验证难题有效时间(秒)
ConfigKeyCapTokenTTL = "cap_token_ttl_seconds" //nolint:gosec // false positive: config key name. 人机验证兑换凭证有效时间(秒)
ConfigKeyServerAddress = "server_address" // 服务器地址
ConfigKeySMTPHost = "smtp_host" // SMTP 服务器地址
ConfigKeySMTPPort = "smtp_port" // SMTP 端口
ConfigKeySMTPUsername = "smtp_username" // SMTP 账户
ConfigKeySMTPPassword = "smtp_password" // SMTP 访问凭证
ConfigKeyEmailLoginVerificationEnabled = "email_login_verification_enabled" // 是否启用邮箱登录验证
ConfigKeyEmailRegisterVerificationEnabled = "email_register_verification_enabled" // 是否启用邮箱注册验证
ConfigKeyMenuDisplayConfig = "menu_display_config" // 目录显示配置 (JSON 字符串)
ConfigKeySearchEngineIndexingEnabled = "search_engine_indexing_enabled" // 是否允许搜索引擎检索
ConfigKeyFileAccessWhitelist = "file_access_whitelist" // 免登录访问的文件业务类型白名单 (JSON 数组格式)
ConfigKeyDiskCacheMaxSizeMB = "disk_cache_max_size_mb" // 磁盘缓存最大空间大小 (MB)
ConfigKeyDiskCacheTTLMinutes = "disk_cache_ttl_minutes" // 磁盘缓存默认有效期 (分钟)
ConfigKeyDiskCacheLRUEnabled = "disk_cache_lru_enabled" // 是否启用 LRU 淘汰机制
ConfigKeyLoginSessionTTLHours = "login_session_ttl_hours" // 登录会话过期时间 (小时,0表示浏览器关闭后自动退出登录,-1表示永不过期)
ConfigKeyUpdateUpstreamRepository = "update_upstream_repository" // GitHub Actions Release 上游仓库
ConfigKeyStorageConfig = "storage_config" // 文件存储配置 (JSON)
)
const (
// ConfigVisibilityHidden 表示配置不通过公共配置接口暴露
ConfigVisibilityHidden = 0
// ConfigVisibilityVisible 表示配置通过公共配置接口暴露
ConfigVisibilityVisible = 1
)
// SystemConfig 系统配置实体
type SystemConfig struct {
Key string `json:"key" gorm:"primaryKey;size:64;not null"`
Value string `json:"value" gorm:"type:text;not null"`
Type string `json:"type" gorm:"size:32;not null;default:'system'"`
Visibility int `json:"visibility" gorm:"not null;default:0"`
Description string `json:"description" gorm:"size:255"`
UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime"`
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime"`
}
// TableName 表名
func (SystemConfig) TableName() string {
return "w_system_configs"
}
@@ -0,0 +1,296 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package model
import (
"context"
"errors"
"fmt"
"strings"
"time"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/db/idgen"
"github.com/redis/go-redis/v9"
)
// TaskExecutionStatus 任务执行状态
type TaskExecutionStatus string
// 任务执行状态
const (
TaskExecutionStatusPending TaskExecutionStatus = "pending"
TaskExecutionStatusRunning TaskExecutionStatus = "running"
TaskExecutionStatusSucceeded TaskExecutionStatus = "succeeded"
TaskExecutionStatusFailed TaskExecutionStatus = "failed"
taskExecutionLogRedisKeyPrefix = "task:execution:log:"
taskExecutionLogExpiration = 24 * time.Hour
taskExecutionLogMaxLines = 1000
)
// TaskExecution 任务执行记录
type TaskExecution struct {
ID uint64 `json:"id,string" gorm:"primaryKey"`
TaskID string `json:"task_id" gorm:"size:128;uniqueIndex;not null"`
TaskType string `json:"task_type" gorm:"size:64;index;not null"`
TaskName string `json:"task_name" gorm:"size:128"`
Status TaskExecutionStatus `json:"status" gorm:"size:32;index;not null"`
Retryable bool `json:"retryable" gorm:"not null;default:false"`
MaxRetry int `json:"max_retry" gorm:"not null;default:0"`
RetryCount int `json:"retry_count" gorm:"not null;default:0"`
Log string `json:"log" gorm:"type:text"`
ErrorMessage string `json:"error_message" gorm:"type:text"`
Result string `json:"result" gorm:"type:text"`
StartedAt *time.Time `json:"started_at" gorm:"index"`
FinishedAt *time.Time `json:"finished_at"`
Duration int64 `json:"duration" gorm:"comment:耗时毫秒"`
Payload string `json:"payload" gorm:"type:text"`
TriggeredBy string `json:"triggered_by" gorm:"size:32;not null;default:system"`
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime;index"`
UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime"`
}
// TaskExecutionCleanupStats describes task execution log cleanup results.
type TaskExecutionCleanupStats struct {
HighFrequencyDeleted int64
LowFrequencyDeleted int64
}
// TableName 表名
func (TaskExecution) TableName() string {
return "w_task_executions"
}
// CreateTaskExecution 创建任务执行记录
func CreateTaskExecution(ctx context.Context, execution *TaskExecution) error {
execution.ID = idgen.NextUint64ID()
return db.DB(ctx).Create(execution).Error
}
// UpdateTaskExecution 更新任务执行记录,忽略由 Redis 缓冲和归档流程管理的 log 字段。
func UpdateTaskExecution(ctx context.Context, execution *TaskExecution) error {
return db.DB(ctx).Omit("log").Save(execution).Error
}
// GetTaskExecutionByTaskID 根据 TaskID 获取执行记录
func GetTaskExecutionByTaskID(ctx context.Context, taskID string) (*TaskExecution, error) {
var execution TaskExecution
if err := db.DB(ctx).Where("task_id = ?", taskID).First(&execution).Error; err != nil {
return nil, err
}
if err := loadTaskExecutionLog(ctx, &execution); err != nil {
return nil, err
}
return &execution, nil
}
// GetTaskExecutionByID 根据 ID 获取执行记录
func GetTaskExecutionByID(ctx context.Context, id uint64) (*TaskExecution, error) {
var execution TaskExecution
if err := db.DB(ctx).Where("id = ?", id).First(&execution).Error; err != nil {
return nil, err
}
if err := loadTaskExecutionLog(ctx, &execution); err != nil {
return nil, err
}
return &execution, nil
}
// AppendTaskExecutionLog 将日志追加到 Redis 缓冲,任务完成后再持久化到数据库。
func AppendTaskExecutionLog(ctx context.Context, taskID string, logLine string) error {
if db.Redis == nil {
return errors.New("redis client is not initialized")
}
now := time.Now().Format("15:04:05")
line := fmt.Sprintf("[%s] %s\n", now, logLine)
key := taskExecutionLogRedisKey(taskID)
_, err := db.Redis.TxPipelined(ctx, func(pipe redis.Pipeliner) error {
pipe.RPush(ctx, key, line)
pipe.LTrim(ctx, key, -taskExecutionLogMaxLines, -1)
pipe.Expire(ctx, key, taskExecutionLogExpiration)
return nil
})
if err != nil {
return fmt.Errorf("append task execution log to redis: %w", err)
}
return nil
}
// FlushTaskExecutionLog 将 Redis 中的完整任务日志写入数据库,并在成功后清理缓存。
func FlushTaskExecutionLog(ctx context.Context, taskID string) error {
if db.Redis == nil {
return errors.New("redis client is not initialized")
}
key := taskExecutionLogRedisKey(taskID)
logLines, err := db.Redis.LRange(ctx, key, 0, -1).Result()
if err != nil {
return fmt.Errorf("get task execution log from redis: %w", err)
}
if len(logLines) == 0 {
return nil
}
logText := strings.Join(logLines, "")
result := db.DB(ctx).Model(&TaskExecution{}).
Where("task_id = ?", taskID).
Update("log", logText)
if result.Error != nil {
return fmt.Errorf("persist task execution log: %w", result.Error)
}
if result.RowsAffected == 0 {
return fmt.Errorf("persist task execution log: task %q not found", taskID)
}
if err := db.Redis.Del(ctx, key).Err(); err != nil {
return fmt.Errorf("delete persisted task execution log from redis: %w", err)
}
return nil
}
// ListTaskExecutionsRequest 查询任务执行记录列表请求
type ListTaskExecutionsRequest struct {
Status string `form:"status"`
TaskType string `form:"task_type"`
Page int `form:"page"`
PageSize int `form:"page_size"`
}
// ListTaskExecutions 分页查询任务执行记录
func ListTaskExecutions(ctx context.Context, req ListTaskExecutionsRequest) ([]TaskExecution, int64, error) {
if req.Page <= 0 {
req.Page = 1
}
if req.PageSize <= 0 {
req.PageSize = 20
}
query := db.DB(ctx).Model(&TaskExecution{})
if req.Status != "" {
query = query.Where("status = ?", req.Status)
}
if req.TaskType != "" {
query = query.Where("task_type = ?", req.TaskType)
}
var total int64
if err := query.Count(&total).Error; err != nil {
return nil, 0, err
}
var executions []TaskExecution
offset := (req.Page - 1) * req.PageSize
if err := query.Order("id DESC").Offset(offset).Limit(req.PageSize).Find(&executions).Error; err != nil {
return nil, 0, err
}
if err := loadTaskExecutionLogs(ctx, executions); err != nil {
return nil, 0, err
}
return executions, total, nil
}
// CleanupTaskExecutionLogs removes finished task execution logs according to frequency-based retention.
func CleanupTaskExecutionLogs(ctx context.Context, now time.Time) (TaskExecutionCleanupStats, error) {
const (
frequencyWindowDays = 30
highFrequencyThreshold = frequencyWindowDays
)
frequencyWindowStart := now.AddDate(0, 0, -frequencyWindowDays)
highFrequencyCutoff := now.AddDate(0, 0, -3)
lowFrequencyCutoff := now.AddDate(0, 0, -30)
terminalStatuses := []TaskExecutionStatus{TaskExecutionStatusSucceeded, TaskExecutionStatusFailed}
var highFrequencyTaskTypes []string
if err := db.DB(ctx).
Model(&TaskExecution{}).
Select("task_type").
Where("created_at >= ?", frequencyWindowStart).
Group("task_type").
Having("COUNT(*) > ?", highFrequencyThreshold).
Pluck("task_type", &highFrequencyTaskTypes).Error; err != nil {
return TaskExecutionCleanupStats{}, fmt.Errorf("query high-frequency task types: %w", err)
}
var highFrequencyDeleted int64
if len(highFrequencyTaskTypes) > 0 {
highFrequencyResult := db.DB(ctx).
Where("status IN ?", terminalStatuses).
Where("created_at < ?", highFrequencyCutoff).
Where("task_type IN ?", highFrequencyTaskTypes).
Delete(&TaskExecution{})
if highFrequencyResult.Error != nil {
return TaskExecutionCleanupStats{}, fmt.Errorf("delete high-frequency task execution logs: %w", highFrequencyResult.Error)
}
highFrequencyDeleted = highFrequencyResult.RowsAffected
}
lowFrequencyQuery := db.DB(ctx).
Where("status IN ?", terminalStatuses).
Where("created_at < ?", lowFrequencyCutoff)
if len(highFrequencyTaskTypes) > 0 {
lowFrequencyQuery = lowFrequencyQuery.Where("task_type NOT IN ?", highFrequencyTaskTypes)
}
lowFrequencyResult := lowFrequencyQuery.Delete(&TaskExecution{})
if lowFrequencyResult.Error != nil {
return TaskExecutionCleanupStats{}, fmt.Errorf("delete low-frequency task execution logs: %w", lowFrequencyResult.Error)
}
return TaskExecutionCleanupStats{
HighFrequencyDeleted: highFrequencyDeleted,
LowFrequencyDeleted: lowFrequencyResult.RowsAffected,
}, nil
}
func taskExecutionLogRedisKey(taskID string) string {
return db.PrefixedKey(taskExecutionLogRedisKeyPrefix + taskID)
}
func loadTaskExecutionLog(ctx context.Context, execution *TaskExecution) error {
if db.Redis == nil {
return nil
}
logLines, err := db.Redis.LRange(ctx, taskExecutionLogRedisKey(execution.TaskID), 0, -1).Result()
if err != nil {
return fmt.Errorf("get task execution log from redis: %w", err)
}
if len(logLines) == 0 {
return nil
}
execution.Log = strings.Join(logLines, "")
return nil
}
func loadTaskExecutionLogs(ctx context.Context, executions []TaskExecution) error {
if db.Redis == nil || len(executions) == 0 {
return nil
}
commands := make([]*redis.StringSliceCmd, len(executions))
_, err := db.Redis.Pipelined(ctx, func(pipe redis.Pipeliner) error {
for i := range executions {
commands[i] = pipe.LRange(ctx, taskExecutionLogRedisKey(executions[i].TaskID), 0, -1)
}
return nil
})
if err != nil {
return fmt.Errorf("get task execution logs from redis: %w", err)
}
for i := range executions {
logLines := commands[i].Val()
if len(logLines) > 0 {
executions[i].Log = strings.Join(logLines, "")
}
}
return nil
}
@@ -0,0 +1,484 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package model
import (
"context"
"fmt"
"testing"
"time"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/alicebob/miniredis/v2"
"github.com/glebarez/sqlite"
"github.com/redis/go-redis/v9"
"github.com/redis/go-redis/v9/maintnotifications"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"gorm.io/gorm"
)
func setupTaskExecutionTestEnvironment(t *testing.T) func() {
sqliteDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{
DisableForeignKeyConstraintWhenMigrating: true,
})
require.NoError(t, err)
err = sqliteDB.AutoMigrate(&TaskExecution{})
require.NoError(t, err)
miniRedis, err := miniredis.Run()
require.NoError(t, err)
redisClient := redis.NewClient(&redis.Options{
Addr: miniRedis.Addr(),
MaintNotificationsConfig: &maintnotifications.Config{
Mode: maintnotifications.ModeDisabled,
},
})
db.SetDB(sqliteDB)
db.Redis = redisClient
return func() {
require.NoError(t, redisClient.Close())
miniRedis.Close()
db.SetDB(nil)
db.Redis = nil
}
}
func TestCreateTaskExecution(t *testing.T) {
cleanup := setupTaskExecutionTestEnvironment(t)
defer cleanup()
ctx := context.Background()
execution := &TaskExecution{
TaskID: "manual_cleanup_123",
TaskType: "system:cleanup",
TaskName: "清理未使用上传",
Status: TaskExecutionStatusPending,
Retryable: true,
MaxRetry: 3,
RetryCount: 0,
Payload: `{"test": true}`,
TriggeredBy: "manual",
}
err := CreateTaskExecution(ctx, execution)
require.NoError(t, err)
assert.NotZero(t, execution.ID, "ID should be generated")
assert.NotZero(t, execution.CreatedAt, "CreatedAt should be set")
assert.NotZero(t, execution.UpdatedAt, "UpdatedAt should be set")
}
func TestGetTaskExecutionByTaskID(t *testing.T) {
cleanup := setupTaskExecutionTestEnvironment(t)
defer cleanup()
ctx := context.Background()
// 创建记录
execution := &TaskExecution{
TaskID: "test_task_id_001",
TaskType: "system:cleanup",
TaskName: "清理未使用上传",
Status: TaskExecutionStatusPending,
Retryable: true,
MaxRetry: 3,
TriggeredBy: "manual",
}
err := CreateTaskExecution(ctx, execution)
require.NoError(t, err)
// 按 TaskID 查询
found, err := GetTaskExecutionByTaskID(ctx, "test_task_id_001")
require.NoError(t, err)
assert.Equal(t, execution.ID, found.ID)
assert.Equal(t, "test_task_id_001", found.TaskID)
assert.Equal(t, TaskExecutionStatusPending, found.Status)
assert.True(t, found.Retryable)
assert.Equal(t, 3, found.MaxRetry)
// 查询不存在的 TaskID
_, err = GetTaskExecutionByTaskID(ctx, "nonexistent")
assert.Error(t, err, "should return error for non-existent taskID")
}
func TestGetTaskExecutionByID(t *testing.T) {
cleanup := setupTaskExecutionTestEnvironment(t)
defer cleanup()
ctx := context.Background()
execution := &TaskExecution{
TaskID: "test_by_id_001",
TaskType: "system:cleanup",
TaskName: "清理未使用上传",
Status: TaskExecutionStatusPending,
TriggeredBy: "system",
}
err := CreateTaskExecution(ctx, execution)
require.NoError(t, err)
// 按主键查询
found, err := GetTaskExecutionByID(ctx, execution.ID)
require.NoError(t, err)
assert.Equal(t, execution.TaskID, found.TaskID)
}
func TestUpdateTaskExecution(t *testing.T) {
cleanup := setupTaskExecutionTestEnvironment(t)
defer cleanup()
ctx := context.Background()
// 创建记录
execution := &TaskExecution{
TaskID: "test_update_001",
TaskType: "system:cleanup",
TaskName: "清理未使用上传",
Status: TaskExecutionStatusPending,
TriggeredBy: "manual",
}
err := CreateTaskExecution(ctx, execution)
require.NoError(t, err)
// 更新状态为 running
now := time.Now()
execution.Status = TaskExecutionStatusRunning
execution.StartedAt = &now
err = UpdateTaskExecution(ctx, execution)
require.NoError(t, err)
// 验证更新
found, err := GetTaskExecutionByTaskID(ctx, "test_update_001")
require.NoError(t, err)
assert.Equal(t, TaskExecutionStatusRunning, found.Status)
assert.NotNil(t, found.StartedAt)
// 更新为 succeeded
finishTime := time.Now()
execution.Status = TaskExecutionStatusSucceeded
execution.FinishedAt = &finishTime
execution.Duration = 1500
execution.Result = "共清理 50 个文件"
err = UpdateTaskExecution(ctx, execution)
require.NoError(t, err)
found, err = GetTaskExecutionByTaskID(ctx, "test_update_001")
require.NoError(t, err)
assert.Equal(t, TaskExecutionStatusSucceeded, found.Status)
assert.Equal(t, int64(1500), found.Duration)
assert.Equal(t, "共清理 50 个文件", found.Result)
}
func TestUpdateTaskExecutionFailed(t *testing.T) {
cleanup := setupTaskExecutionTestEnvironment(t)
defer cleanup()
ctx := context.Background()
execution := &TaskExecution{
TaskID: "test_fail_001",
TaskType: "system:cleanup",
TaskName: "清理未使用上传",
Status: TaskExecutionStatusPending,
Retryable: true,
MaxRetry: 3,
TriggeredBy: "manual",
}
err := CreateTaskExecution(ctx, execution)
require.NoError(t, err)
// 标记为失败
now := time.Now()
execution.Status = TaskExecutionStatusFailed
execution.StartedAt = &now
execution.FinishedAt = &now
execution.Duration = 200
execution.ErrorMessage = "S3 连接超时"
err = UpdateTaskExecution(ctx, execution)
require.NoError(t, err)
found, err := GetTaskExecutionByTaskID(ctx, "test_fail_001")
require.NoError(t, err)
assert.Equal(t, TaskExecutionStatusFailed, found.Status)
assert.Equal(t, "S3 连接超时", found.ErrorMessage)
assert.Equal(t, int64(200), found.Duration)
}
func TestUpdateTaskExecutionDoesNotPersistBufferedLog(t *testing.T) {
cleanup := setupTaskExecutionTestEnvironment(t)
defer cleanup()
ctx := context.Background()
execution := &TaskExecution{
TaskID: "test_omit_log_001",
TaskType: "system:cleanup",
TaskName: "清理未使用上传",
Status: TaskExecutionStatusPending,
TriggeredBy: "manual",
}
err := CreateTaskExecution(ctx, execution)
require.NoError(t, err)
// 运行中的日志仅缓存在 Redis。
err = AppendTaskExecutionLog(ctx, "test_omit_log_001", "第一条执行日志")
require.NoError(t, err)
assert.Empty(t, execution.Log)
execution.Status = TaskExecutionStatusSucceeded
execution.Duration = 100
err = UpdateTaskExecution(ctx, execution)
require.NoError(t, err)
var persisted TaskExecution
err = db.DB(ctx).Where("task_id = ?", "test_omit_log_001").First(&persisted).Error
require.NoError(t, err)
assert.Equal(t, TaskExecutionStatusSucceeded, persisted.Status)
assert.Empty(t, persisted.Log)
found, err := GetTaskExecutionByTaskID(ctx, "test_omit_log_001")
require.NoError(t, err)
assert.Contains(t, found.Log, "第一条执行日志")
}
func TestAppendTaskExecutionLog(t *testing.T) {
cleanup := setupTaskExecutionTestEnvironment(t)
defer cleanup()
ctx := context.Background()
execution := &TaskExecution{
TaskID: "test_log_001",
TaskType: "system:cleanup",
TaskName: "清理未使用上传",
Status: TaskExecutionStatusPending,
TriggeredBy: "manual",
}
err := CreateTaskExecution(ctx, execution)
require.NoError(t, err)
// 追加多条日志
err = AppendTaskExecutionLog(ctx, "test_log_001", "开始扫描未使用上传文件")
require.NoError(t, err)
err = AppendTaskExecutionLog(ctx, "test_log_001", "本批次找到 42 个待清理文件")
require.NoError(t, err)
err = AppendTaskExecutionLog(ctx, "test_log_001", "清理完成,共删除 42 个文件")
require.NoError(t, err)
// 读取时优先返回 Redis 中的在途日志。
found, err := GetTaskExecutionByTaskID(ctx, "test_log_001")
require.NoError(t, err)
assert.Contains(t, found.Log, "开始扫描未使用上传文件")
assert.Contains(t, found.Log, "本批次找到 42 个待清理文件")
assert.Contains(t, found.Log, "清理完成,共删除 42 个文件")
var persisted TaskExecution
err = db.DB(ctx).Where("task_id = ?", "test_log_001").First(&persisted).Error
require.NoError(t, err)
assert.Empty(t, persisted.Log)
err = FlushTaskExecutionLog(ctx, "test_log_001")
require.NoError(t, err)
err = db.DB(ctx).Where("task_id = ?", "test_log_001").First(&persisted).Error
require.NoError(t, err)
assert.Contains(t, persisted.Log, "开始扫描未使用上传文件")
exists, err := db.Redis.Exists(ctx, taskExecutionLogRedisKey("test_log_001")).Result()
require.NoError(t, err)
assert.Zero(t, exists)
}
func TestAppendTaskExecutionLogLimitsLinesAndRefreshesTTL(t *testing.T) {
cleanup := setupTaskExecutionTestEnvironment(t)
defer cleanup()
ctx := context.Background()
const taskID = "limited_log_001"
for i := 0; i < taskExecutionLogMaxLines+5; i++ {
err := AppendTaskExecutionLog(ctx, taskID, fmt.Sprintf("日志-%04d", i))
require.NoError(t, err)
}
key := taskExecutionLogRedisKey(taskID)
logLines, err := db.Redis.LRange(ctx, key, 0, -1).Result()
require.NoError(t, err)
assert.Len(t, logLines, taskExecutionLogMaxLines)
assert.Contains(t, logLines[0], "日志-0005")
assert.Contains(t, logLines[len(logLines)-1], "日志-1004")
ttl, err := db.Redis.TTL(ctx, key).Result()
require.NoError(t, err)
assert.Equal(t, taskExecutionLogExpiration, ttl)
}
func TestAppendTaskExecutionLogNonExistent(t *testing.T) {
cleanup := setupTaskExecutionTestEnvironment(t)
defer cleanup()
ctx := context.Background()
// Redis 缓冲不依赖数据库记录是否已经创建。
err := AppendTaskExecutionLog(ctx, "nonexistent_task", "测试日志")
assert.NoError(t, err)
err = FlushTaskExecutionLog(ctx, "nonexistent_task")
assert.Error(t, err)
}
func TestGetTaskExecutionLogPrefersRedis(t *testing.T) {
cleanup := setupTaskExecutionTestEnvironment(t)
defer cleanup()
ctx := context.Background()
execution := &TaskExecution{
TaskID: "redis_priority_001",
TaskType: "system:cleanup",
TaskName: "清理未使用上传",
Status: TaskExecutionStatusRunning,
Log: "数据库旧日志",
TriggeredBy: "manual",
}
err := CreateTaskExecution(ctx, execution)
require.NoError(t, err)
err = AppendTaskExecutionLog(ctx, execution.TaskID, "Redis 最新日志")
require.NoError(t, err)
found, err := GetTaskExecutionByID(ctx, execution.ID)
require.NoError(t, err)
assert.Contains(t, found.Log, "Redis 最新日志")
assert.NotContains(t, found.Log, "数据库旧日志")
}
func TestListTaskExecutions(t *testing.T) {
cleanup := setupTaskExecutionTestEnvironment(t)
defer cleanup()
ctx := context.Background()
// 创建多条记录,包含不同状态和类型
records := []*TaskExecution{
{TaskID: "list_001", TaskType: "system:cleanup", TaskName: "系统垃圾清理", Status: TaskExecutionStatusSucceeded, TriggeredBy: "manual"},
{TaskID: "list_002", TaskType: "system:cleanup", TaskName: "系统垃圾清理", Status: TaskExecutionStatusFailed, TriggeredBy: "system"},
{TaskID: "list_003", TaskType: "other:task", TaskName: "其他任务", Status: TaskExecutionStatusPending, TriggeredBy: "manual"},
{TaskID: "list_004", TaskType: "system:cleanup", TaskName: "系统垃圾清理", Status: TaskExecutionStatusRunning, TriggeredBy: "manual"},
{TaskID: "list_005", TaskType: "other:task", TaskName: "其他任务", Status: TaskExecutionStatusSucceeded, TriggeredBy: "system"},
}
for _, r := range records {
err := CreateTaskExecution(ctx, r)
require.NoError(t, err)
}
err := AppendTaskExecutionLog(ctx, "list_004", "运行中的 Redis 日志")
require.NoError(t, err)
// 查询全部(分页)
items, total, err := ListTaskExecutions(ctx, ListTaskExecutionsRequest{Page: 1, PageSize: 10})
require.NoError(t, err)
assert.Equal(t, int64(5), total)
assert.Len(t, items, 5)
for _, item := range items {
if item.TaskID == "list_004" {
assert.Contains(t, item.Log, "运行中的 Redis 日志")
}
}
// 按状态筛选:failed
items, total, err = ListTaskExecutions(ctx, ListTaskExecutionsRequest{Status: "failed", Page: 1, PageSize: 10})
require.NoError(t, err)
assert.Equal(t, int64(1), total)
assert.Len(t, items, 1)
assert.Equal(t, "list_002", items[0].TaskID)
// 按类型筛选
_, total, err = ListTaskExecutions(ctx, ListTaskExecutionsRequest{TaskType: "other:task", Page: 1, PageSize: 10})
require.NoError(t, err)
assert.Equal(t, int64(2), total)
// 分页测试
items, total, err = ListTaskExecutions(ctx, ListTaskExecutionsRequest{Page: 1, PageSize: 2})
require.NoError(t, err)
assert.Equal(t, int64(5), total)
assert.Len(t, items, 2)
items2, total2, err := ListTaskExecutions(ctx, ListTaskExecutionsRequest{Page: 2, PageSize: 2})
require.NoError(t, err)
assert.Equal(t, int64(5), total2)
assert.Len(t, items2, 2)
// 确保分页数据不重复
assert.NotEqual(t, items[0].ID, items2[0].ID)
// 状态 + 类型组合筛选
items, total, err = ListTaskExecutions(ctx, ListTaskExecutionsRequest{Status: "succeeded", TaskType: "system:cleanup", Page: 1, PageSize: 10})
require.NoError(t, err)
assert.Equal(t, int64(1), total)
assert.Equal(t, "list_001", items[0].TaskID)
}
func TestListTaskExecutionsDefaultPaging(t *testing.T) {
cleanup := setupTaskExecutionTestEnvironment(t)
defer cleanup()
ctx := context.Background()
// 不传分页参数,应使用默认值 page=1, pageSize=20
items, total, err := ListTaskExecutions(ctx, ListTaskExecutionsRequest{})
require.NoError(t, err)
assert.Equal(t, int64(0), total)
assert.Len(t, items, 0)
}
func TestCleanupTaskExecutionLogs(t *testing.T) {
cleanup := setupTaskExecutionTestEnvironment(t)
defer cleanup()
ctx := context.Background()
now := time.Date(2026, 6, 17, 12, 0, 0, 0, time.UTC)
for i := 0; i < 31; i++ {
createTaskExecutionForCleanup(t, ctx, fmt.Sprintf("high_recent_%02d", i), "high:task", TaskExecutionStatusSucceeded, now.Add(-2*time.Hour))
}
createTaskExecutionForCleanup(t, ctx, "high_old_4d", "high:task", TaskExecutionStatusSucceeded, now.AddDate(0, 0, -4))
createTaskExecutionForCleanup(t, ctx, "high_old_40d", "high:task", TaskExecutionStatusFailed, now.AddDate(0, 0, -40))
createTaskExecutionForCleanup(t, ctx, "high_running_old", "high:task", TaskExecutionStatusRunning, now.AddDate(0, 0, -10))
createTaskExecutionForCleanup(t, ctx, "low_old_31d", "low:task", TaskExecutionStatusSucceeded, now.AddDate(0, 0, -31))
createTaskExecutionForCleanup(t, ctx, "low_recent_29d", "low:task", TaskExecutionStatusSucceeded, now.AddDate(0, 0, -29))
createTaskExecutionForCleanup(t, ctx, "low_pending_old", "low:task", TaskExecutionStatusPending, now.AddDate(0, 0, -45))
stats, err := CleanupTaskExecutionLogs(ctx, now)
require.NoError(t, err)
assert.Equal(t, int64(2), stats.HighFrequencyDeleted)
assert.Equal(t, int64(1), stats.LowFrequencyDeleted)
for _, taskID := range []string{"high_old_4d", "high_old_40d", "low_old_31d"} {
var count int64
err := db.DB(ctx).Model(&TaskExecution{}).Where("task_id = ?", taskID).Count(&count).Error
require.NoError(t, err)
assert.Equal(t, int64(0), count, "CleanupTaskExecutionLogs(%s) should delete expired log", taskID)
}
for _, taskID := range []string{"high_recent_00", "high_running_old", "low_recent_29d", "low_pending_old"} {
var count int64
err := db.DB(ctx).Model(&TaskExecution{}).Where("task_id = ?", taskID).Count(&count).Error
require.NoError(t, err)
assert.Equal(t, int64(1), count, "CleanupTaskExecutionLogs(%s) should keep retained log", taskID)
}
}
func TestTaskExecutionTableName(t *testing.T) {
execution := TaskExecution{}
assert.Equal(t, "w_task_executions", execution.TableName())
}
func createTaskExecutionForCleanup(t *testing.T, ctx context.Context, taskID string, taskType string, status TaskExecutionStatus, createdAt time.Time) {
t.Helper()
execution := &TaskExecution{
TaskID: taskID,
TaskType: taskType,
TaskName: taskType,
Status: status,
CreatedAt: createdAt,
UpdatedAt: createdAt,
TriggeredBy: "system",
}
err := CreateTaskExecution(ctx, execution)
require.NoError(t, err)
}
@@ -0,0 +1,88 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package model
import (
"bytes"
"errors"
"strings"
"text/template"
"time"
)
// Template 邮件/消息模板实体
type Template struct {
ID uint64 `json:"id" gorm:"primaryKey;autoIncrement"`
Key string `json:"key" gorm:"uniqueIndex;size:80;not null"`
Name string `json:"name" gorm:"size:100;not null"`
Type string `json:"type" gorm:"size:20;not null;default:'email'"`
Subject string `json:"subject" gorm:"size:255"`
Content string `json:"content" gorm:"type:text;not null"`
Description string `json:"description" gorm:"size:255"`
IsSystem bool `json:"is_system" gorm:"index;not null;default:false"`
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime;index"`
UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime;index"`
}
// TableName 表名
func (Template) TableName() string {
return "w_templates"
}
// Normalize 规范化模板字段
func (t *Template) Normalize() {
t.Key = strings.TrimSpace(t.Key)
t.Name = strings.TrimSpace(t.Name)
t.Type = strings.ToLower(strings.TrimSpace(t.Type))
t.Subject = strings.TrimSpace(t.Subject)
t.Content = strings.TrimSpace(t.Content)
t.Description = strings.TrimSpace(t.Description)
if t.Type == "" {
t.Type = "email"
}
}
// Validate 校验模板必填字段
func (t *Template) Validate() error {
t.Normalize()
if t.Key == "" {
return errors.New(errTemplateKeyRequired)
}
if t.Name == "" {
return errors.New(errTemplateNameRequired)
}
if t.Content == "" {
return errors.New(errTemplateContentRequired)
}
return nil
}
// Render 渲染模板的 Subject 和 Content
func (t *Template) Render(data any) (string, string, error) {
// Render Subject
var subject string
if t.Subject != "" {
tmplSubject, err := template.New(t.Key + "_subject").Parse(t.Subject)
if err != nil {
return "", "", err
}
var subBuf bytes.Buffer
if err := tmplSubject.Execute(&subBuf, data); err != nil {
return "", "", err
}
subject = subBuf.String()
}
// Render Content
tmplContent, err := template.New(t.Key + "_content").Parse(t.Content)
if err != nil {
return "", "", err
}
var bodyBuf bytes.Buffer
if err := tmplContent.Execute(&bodyBuf, data); err != nil {
return "", "", err
}
return subject, bodyBuf.String(), nil
}
@@ -1,51 +0,0 @@
package model
import "time"
type TLSCertificate struct {
ID uint `json:"id" gorm:"primaryKey"`
Name string `json:"name" gorm:"uniqueIndex;size:255;not null"`
CertPEM string `json:"-" gorm:"type:text;not null"`
KeyPEM string `json:"-" gorm:"type:text;not null"`
NotBefore time.Time `json:"not_before"`
NotAfter time.Time `json:"not_after"`
Remark string `json:"remark" gorm:"size:255"`
Provider string `json:"provider" gorm:"size:64;default:'upload'"` // upload, acme
AcmeAccountID uint `json:"acme_account_id"`
DnsAccountID uint `json:"dns_account_id"`
KeyAlgorithm string `json:"key_algorithm" gorm:"size:32"`
AutoRenew bool `json:"auto_renew"`
PrimaryDomain string `json:"primary_domain" gorm:"size:255"`
OtherDomains string `json:"other_domains" gorm:"type:text"`
DisableCNAME bool `json:"disable_cname"`
SkipDNS bool `json:"skip_dns"`
DNS1 string `json:"dns1" gorm:"size:128"`
DNS2 string `json:"dns2" gorm:"size:128"`
ApplyStatus string `json:"apply_status" gorm:"size:64;default:'ready'"`
ApplyMessage string `json:"apply_message" gorm:"type:text"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
}
func ListTLSCertificates() (certificates []*TLSCertificate, err error) {
err = DB.Order("id desc").Find(&certificates).Error
return certificates, err
}
func GetTLSCertificateByID(id uint) (*TLSCertificate, error) {
certificate := &TLSCertificate{}
err := DB.First(certificate, id).Error
return certificate, err
}
func (certificate *TLSCertificate) Insert() error {
return DB.Create(certificate).Error
}
func (certificate *TLSCertificate) Update() error {
return DB.Save(certificate).Error
}
func (certificate *TLSCertificate) Delete() error {
return DB.Delete(certificate).Error
}
@@ -0,0 +1,28 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package model
import "time"
// Upload stats dimension keys stored in w_upload_stats.dimension.
const (
UploadStatDimensionTotal = "total"
UploadStatDimensionType = "type"
UploadStatDimensionCategory = "category"
UploadStatDimensionTrend = "trend"
)
// UploadStat stores incremental upload statistics keyed by dimension and stat_key.
type UploadStat struct {
Dimension string `json:"dimension" gorm:"primaryKey;size:32;not null"`
StatKey string `json:"stat_key" gorm:"primaryKey;size:64;not null;default:''"`
FileCount int64 `json:"file_count" gorm:"not null;default:0"`
FileSize int64 `json:"file_size" gorm:"not null;default:0"`
UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime"`
}
// TableName returns the upload stats table name.
func (UploadStat) TableName() string {
return "w_upload_stats"
}
@@ -0,0 +1,54 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package model
import (
"time"
)
// UploadStatus 上传状态
type UploadStatus string
// 上传状态
const (
UploadStatusPending UploadStatus = "pending" // 待使用
UploadStatusUsed UploadStatus = "used" // 已使用
UploadStatusDeleted UploadStatus = "deleted" // 已删除
)
// UploadMetadata 自定义可扩展的 JSON 字段存储非核心或可选的文件元数据
type UploadMetadata struct {
Width int `json:"width,omitempty"` // 图像/视频宽度 (px)
Height int `json:"height,omitempty"` // 图像/视频高度 (px)
Duration float64 `json:"duration,omitempty"` // 音视频时长 (s)
OriginalMime string `json:"original_mime,omitempty"` // 原始 MIME 类型
UserAgent string `json:"user_agent,omitempty"` // 上传者的 UA
ClientIP string `json:"client_ip,omitempty"` // 上传者 IP
Bucket string `json:"bucket,omitempty"` // 存储桶名称 (适用于 S3 等)
Extra map[string]any `json:"extra,omitempty"` // 其它任意业务自定义元数据
}
// Upload 上传文件记录
type Upload struct {
ID uint64 `json:"id,string" gorm:"primaryKey"`
UserID uint64 `json:"user_id,string" gorm:"index;not null"`
FileName string `json:"file_name" gorm:"size:255;not null"` // 原始文件名 (例如: image.png)
FilePath string `json:"file_path" gorm:"size:500;not null;index"` // 文件相对路径 / S3 Key
FileSize int64 `json:"file_size" gorm:"not null"` // 文件大小(字节)
MimeType string `json:"mime_type" gorm:"size:100;not null"` // 媒体类型 (MIME, 如 image/png)
Extension string `json:"extension" gorm:"size:50;not null"` // 文件后缀名 (不含点,如 png, pdf)
Hash string `json:"hash" gorm:"size:64;index"` // 文件哈希 (SHA-256/MD5,可用于排重)
Type string `json:"type" gorm:"column:type;size:50;not null;index"` // 业务标识类型 (如 avatar, doc, attachment)
Status UploadStatus `json:"status" gorm:"type:varchar(20);not null"` // 状态
AccessMode int `json:"access_mode" gorm:"column:access_mode;not null;default:0"`
Metadata UploadMetadata `json:"metadata" gorm:"serializer:json;type:jsonb"` // 业务扩展元数据
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime"`
UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime"`
}
// TableName 表名
func (Upload) TableName() string {
return "w_uploads"
}
-193
View File
@@ -1,193 +0,0 @@
package model
import (
"errors"
"github.com/rain-kl/openflare/openflare-server/internal/common"
"github.com/rain-kl/openflare/openflare-server/internal/utils/security"
)
// User if you add sensitive fields, don't forget to clean them in setupLogin function.
// Otherwise, the sensitive information will be saved on local storage in plain text!
type User struct {
Id int `json:"id"`
Username string `json:"username" gorm:"unique;index" validate:"max=12"`
Password string `json:"password" gorm:"not null;" validate:"min=8,max=20"`
DisplayName string `json:"display_name" gorm:"index" validate:"max=20"`
Role int `json:"role" gorm:"type:int;default:1"` // admin, common
Status int `json:"status" gorm:"type:int;default:1"` // enabled, disabled
Token string `json:"token" gorm:"index"`
Email string `json:"email" gorm:"index" validate:"max=50"`
GitHubId string `json:"github_id" gorm:"column:github_id;index"`
WeChatId string `json:"wechat_id" gorm:"column:wechat_id;index"`
VerificationCode string `json:"verification_code" gorm:"-:all"` // this field is only for Email verification, don't save it to database!
}
func GetMaxUserId() int {
var user User
DB.Last(&user)
return user.Id
}
func GetAllUsers(startIdx int, num int) (users []*User, err error) {
err = DB.Order("id desc").Limit(num).Offset(startIdx).Select([]string{"id", "username", "display_name", "role", "status", "email"}).Find(&users).Error
return users, err
}
func SearchUsers(keyword string) (users []*User, err error) {
err = DB.Select([]string{"id", "username", "display_name", "role", "status", "email"}).Where("id = ? or username LIKE ? or email LIKE ? or display_name LIKE ?", keyword, keyword+"%", keyword+"%", keyword+"%").Find(&users).Error
return users, err
}
func GetUserById(id int, selectAll bool) (*User, error) {
if id == 0 {
return nil, errors.New("id 为空!")
}
user := User{Id: id}
var err error = nil
if selectAll {
err = DB.First(&user, "id = ?", id).Error
} else {
err = DB.Select([]string{"id", "username", "display_name", "role", "status", "email", "wechat_id", "github_id"}).First(&user, "id = ?", id).Error
}
return &user, err
}
func DeleteUserById(id int) (err error) {
if id == 0 {
return errors.New("id 为空!")
}
user := User{Id: id}
return user.Delete()
}
func (user *User) Insert() error {
var err error
if user.Password != "" {
user.Password, err = security.Password2Hash(user.Password)
if err != nil {
return err
}
}
err = DB.Create(user).Error
return err
}
func (user *User) Update(updatePassword bool) error {
var err error
if updatePassword {
user.Password, err = security.Password2Hash(user.Password)
if err != nil {
return err
}
}
err = DB.Model(user).Updates(user).Error
return err
}
func (user *User) Delete() error {
if user.Id == 0 {
return errors.New("id 为空!")
}
err := DB.Delete(user).Error
return err
}
// ValidateAndFill check password & user status
func (user *User) ValidateAndFill() (err error) {
// When querying with struct, GORM will only query with non-zero fields,
// that means if your field’s value is 0, '', false or other zero values,
// it won’t be used to build query conditions
password := user.Password
if user.Username == "" || password == "" {
return errors.New("用户名或密码为空")
}
DB.Where(User{Username: user.Username}).First(user)
okay := security.ValidatePasswordAndHash(password, user.Password)
if !okay || user.Status != common.UserStatusEnabled {
return errors.New("用户名或密码错误,或用户已被封禁")
}
return nil
}
func (user *User) FillUserById() error {
if user.Id == 0 {
return errors.New("id 为空!")
}
DB.Where(User{Id: user.Id}).First(user)
return nil
}
func (user *User) FillUserByEmail() error {
if user.Email == "" {
return errors.New("email 为空!")
}
DB.Where(User{Email: user.Email}).First(user)
return nil
}
func (user *User) FillUserByGitHubId() error {
if user.GitHubId == "" {
return errors.New("GitHub id 为空!")
}
DB.Where(User{GitHubId: user.GitHubId}).First(user)
return nil
}
func (user *User) FillUserByWeChatId() error {
if user.WeChatId == "" {
return errors.New("WeChat id 为空!")
}
DB.Where(User{WeChatId: user.WeChatId}).First(user)
return nil
}
func (user *User) FillUserByUsername() error {
if user.Username == "" {
return errors.New("username 为空!")
}
DB.Where(User{Username: user.Username}).First(user)
return nil
}
// ValidateUserToken looks up a user by their stored JWT token string.
// JWT signature verification is handled by middleware/auth.go; this
// function is used by Logout to find and clear the token from DB.
func ValidateUserToken(token string) (user *User) {
if token == "" {
return nil
}
user = &User{}
if DB.Where("token = ?", token).First(user).RowsAffected == 1 {
return user
}
return nil
}
func IsEmailAlreadyTaken(email string) bool {
return DB.Where("email = ?", email).Find(&User{}).RowsAffected == 1
}
func IsWeChatIdAlreadyTaken(wechatId string) bool {
return DB.Where("wechat_id = ?", wechatId).Find(&User{}).RowsAffected == 1
}
func IsGitHubIdAlreadyTaken(githubId string) bool {
return DB.Where("github_id = ?", githubId).Find(&User{}).RowsAffected == 1
}
func IsUsernameAlreadyTaken(username string) bool {
return DB.Where("username = ?", username).Find(&User{}).RowsAffected == 1
}
func ResetUserPasswordByEmail(email string, password string) error {
if email == "" || password == "" {
return errors.New("邮箱地址或密码为空!")
}
hashedPassword, err := security.Password2Hash(password)
if err != nil {
return err
}
err = DB.Model(&User{}).Where("email = ?", email).Update("password", hashedPassword).Error
return err
}
+195
View File
@@ -0,0 +1,195 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package model
import (
"context"
"errors"
"strconv"
"strings"
"time"
"github.com/Rain-kl/Wavelet/internal/common"
"github.com/Rain-kl/Wavelet/internal/db/idgen"
"github.com/Rain-kl/Wavelet/pkg/util"
"gorm.io/gorm"
)
// OAuthUserInfo 用户信息结构(同时支持 OIDC ID Token claims 和 UserEndpoint 响应)
type OAuthUserInfo struct {
ID uint64 `json:"id"`
Sub string `json:"sub"`
Username string `json:"username"`
PreferredUsername string `json:"preferred_username"`
Email string `json:"email"`
Name string `json:"name"`
Active bool `json:"active"`
AvatarURL string `json:"avatar_url"`
}
// GetID 获取用户 ID
func (u *OAuthUserInfo) GetID() uint64 {
if u.ID != 0 {
return u.ID
}
// 从 sub 解析(OIDC 格式)
if u.Sub != "" {
if id, err := strconv.ParseUint(u.Sub, 10, 64); err == nil {
return id
}
}
return 0
}
// User 用户表实体
type User struct {
ID uint64 `json:"id,string" gorm:"primaryKey;not null"`
Username string `json:"username" gorm:"size:64;uniqueIndex"`
Password string `json:"password,omitempty" gorm:"size:255"`
Nickname string `json:"nickname" gorm:"size:255"`
Email string `json:"email" gorm:"size:255;index"`
AvatarURL string `json:"avatar_url" gorm:"size:255"`
IsActive bool `json:"is_active" gorm:"default:true;index"`
IsAdmin bool `json:"is_admin" gorm:"default:false"`
Bio string `json:"bio" gorm:"size:500"`
Phone string `json:"phone" gorm:"size:32"`
Gender string `json:"gender" gorm:"size:16"`
Website string `json:"website" gorm:"size:255"`
Location string `json:"location" gorm:"size:255"`
LastLoginAt time.Time `json:"last_login_at" gorm:"index"`
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime;index"`
UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime;index"`
}
// TableName 表名
func (User) TableName() string {
return "w_users"
}
// SetPassword 设置明文密码
func (u *User) SetPassword(password string) error {
u.Password = password
return nil
}
// SetEncryptedPassword 设置加密密码
func (u *User) SetEncryptedPassword(password string) error {
if password == "" {
u.Password = ""
return nil
}
hashed, err := util.HashPassword(password)
if err != nil {
return err
}
u.Password = hashed
return nil
}
// IsPasswordEncrypted 检查密码是否已加密
func (u *User) IsPasswordEncrypted() bool {
return strings.HasPrefix(u.Password, "$2a$") || strings.HasPrefix(u.Password, "$2b$") || strings.HasPrefix(u.Password, "$2y$")
}
// CheckPassword 验证密码是否匹配
func (u *User) CheckPassword(password string) bool {
if u.Password == "" || password == "" {
return false
}
if u.IsPasswordEncrypted() {
return util.CheckPasswordHash(u.Password, password)
}
return u.Password == password
}
// GetByID 根据 ID 查询用户
func (u *User) GetByID(tx *gorm.DB, id uint64) error {
if err := tx.Where("id = ?", id).First(u).Error; err != nil {
return err
}
return nil
}
// UpdateFromOAuthInfo 根据 OAuth 信息更新用户数据
func (u *User) UpdateFromOAuthInfo(oauthInfo *OAuthUserInfo) {
u.Username = oauthInfo.Username
u.Nickname = oauthInfo.Name
u.Email = oauthInfo.Email
u.AvatarURL = oauthInfo.AvatarURL
u.IsActive = oauthInfo.Active
u.LastLoginAt = time.Now()
}
// CheckActive 检查用户账户是否激活,未激活则返回错误
func (u *User) CheckActive() error {
if !u.IsActive {
return errors.New(common.BannedAccount)
}
return nil
}
func (u *User) assignIDIfMissing() error {
if u.ID != 0 {
return nil
}
u.ID = idgen.NextUint64ID()
return nil
}
// CreateUser 创建新用户(用于 OAuth/OIDC 自动注册,含底层权限校验)
func (u *User) CreateUser(_ context.Context, tx *gorm.DB, oauthInfo *OAuthUserInfo) error {
now := time.Now()
userID := oauthInfo.GetID()
newUser := User{
ID: userID,
Username: oauthInfo.Username,
Nickname: oauthInfo.Name,
Email: oauthInfo.Email,
AvatarURL: oauthInfo.AvatarURL,
IsActive: oauthInfo.Active,
LastLoginAt: now,
IsAdmin: false,
}
if err := newUser.assignIDIfMissing(); err != nil {
return err
}
if err := tx.Create(&newUser).Error; err != nil {
return err
}
*u = newUser
return nil
}
// RegisterUser 创建新用户并注册(用于本地密码注册,含全局开关和唯一性多重底层校验)
func (u *User) RegisterUser(_ context.Context, tx *gorm.DB) error {
// 检查用户名冲突
var count int64
if err := tx.Model(&User{}).Where("username = ?", u.Username).Count(&count).Error; err != nil {
return err
}
if count > 0 {
return errors.New(errUsernameExists)
}
// 检查邮箱冲突
if u.Email != "" {
var emailCount int64
if err := tx.Model(&User{}).Where("email = ?", u.Email).Count(&emailCount).Error; err != nil {
return err
}
if emailCount > 0 {
return errors.New(errEmailAlreadyBound)
}
}
if err := u.assignIDIfMissing(); err != nil {
return err
}
if err := tx.Create(u).Error; err != nil {
return err
}
return nil
}
-168
View File
@@ -1,168 +0,0 @@
package model
import "time"
type WAFRuleGroup struct {
ID uint `json:"id" gorm:"primaryKey"`
Name string `json:"name" gorm:"size:255;not null"`
Enabled bool `json:"enabled" gorm:"not null;default:true"`
IsGlobal bool `json:"is_global" gorm:"not null;default:false;index"`
BlockStatusCode int `json:"block_status_code" gorm:"not null;default:418"`
BlockResponseBody string `json:"block_response_body" gorm:"type:text;not null;default:''"`
IPWhitelist string `json:"ip_whitelist" gorm:"type:text;not null;default:'[]'"`
IPBlacklist string `json:"ip_blacklist" gorm:"type:text;not null;default:'[]'"`
IPWhitelistGroups string `json:"ip_whitelist_group_ids" gorm:"type:text;not null;default:'[]'"`
IPBlacklistGroups string `json:"ip_blacklist_group_ids" gorm:"type:text;not null;default:'[]'"`
CountryWhitelist string `json:"country_whitelist" gorm:"type:text;not null;default:'[]'"`
CountryBlacklist string `json:"country_blacklist" gorm:"type:text;not null;default:'[]'"`
RegionWhitelist string `json:"region_whitelist" gorm:"type:text;not null;default:'[]'"`
RegionBlacklist string `json:"region_blacklist" gorm:"type:text;not null;default:'[]'"`
PoWEnabled bool `json:"pow_enabled" gorm:"column:pow_enabled;not null;default:false"`
PoWConfig string `json:"pow_config" gorm:"column:pow_config;type:text;not null;default:'{}'"`
Remark string `json:"remark" gorm:"size:255"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
}
type WAFIPGroup struct {
ID uint `json:"id" gorm:"primaryKey"`
Name string `json:"name" gorm:"size:255;not null"`
Type string `json:"type" gorm:"size:32;not null;index"`
Enabled bool `json:"enabled" gorm:"not null;default:true"`
IPList string `json:"ip_list" gorm:"type:text;not null;default:'[]'"`
AutoConfig string `json:"auto_config" gorm:"type:text;not null;default:'{}'"`
ExtIPs string `json:"ext_ips" gorm:"type:text;not null;default:'[]'"`
SubscriptionURL string `json:"subscription_url" gorm:"size:2048;not null;default:''"`
SubscriptionFormat string `json:"subscription_format" gorm:"size:32;not null;default:'text'"`
SubscriptionMappingRule string `json:"subscription_mapping_rule" gorm:"size:255;not null;default:''"`
SyncIntervalMinutes int `json:"sync_interval_minutes" gorm:"not null;default:1440"`
LastSyncedAt *time.Time `json:"last_synced_at"`
NextSyncAt *time.Time `json:"next_sync_at" gorm:"index"`
LastSyncStatus string `json:"last_sync_status" gorm:"size:32;not null;default:''"`
LastSyncMessage string `json:"last_sync_message" gorm:"type:text;not null;default:''"`
Remark string `json:"remark" gorm:"size:255"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
}
type WAFRuleGroupBinding struct {
ID uint `json:"id" gorm:"primaryKey"`
RuleGroupID uint `json:"rule_group_id" gorm:"not null;uniqueIndex:idx_waf_group_route"`
ProxyRouteID uint `json:"proxy_route_id" gorm:"not null;uniqueIndex:idx_waf_group_route;index"`
CreatedAt time.Time `json:"created_at"`
}
func ListWAFRuleGroups() ([]*WAFRuleGroup, error) {
var groups []*WAFRuleGroup
err := DB.Order("is_global desc").Order("id asc").Find(&groups).Error
return groups, err
}
func GetWAFRuleGroupByID(id uint) (*WAFRuleGroup, error) {
group := &WAFRuleGroup{}
err := DB.First(group, id).Error
return group, err
}
func GetGlobalWAFRuleGroup() (*WAFRuleGroup, error) {
group := &WAFRuleGroup{}
err := DB.Where("is_global = ?", true).Order("id asc").First(group).Error
return group, err
}
func (group *WAFRuleGroup) Insert() error {
return DB.Create(group).Error
}
func (group *WAFRuleGroup) Update() error {
return DB.Model(&WAFRuleGroup{}).Where("id = ?", group.ID).Updates(map[string]any{
"name": group.Name,
"enabled": group.Enabled,
"is_global": group.IsGlobal,
"block_status_code": group.BlockStatusCode,
"block_response_body": group.BlockResponseBody,
"ip_whitelist": group.IPWhitelist,
"ip_blacklist": group.IPBlacklist,
"ip_whitelist_groups": group.IPWhitelistGroups,
"ip_blacklist_groups": group.IPBlacklistGroups,
"country_whitelist": group.CountryWhitelist,
"country_blacklist": group.CountryBlacklist,
"region_whitelist": group.RegionWhitelist,
"region_blacklist": group.RegionBlacklist,
"pow_enabled": group.PoWEnabled,
"pow_config": group.PoWConfig,
"remark": group.Remark,
}).Error
}
func (group *WAFRuleGroup) Delete() error {
return DB.Delete(group).Error
}
func ListWAFIPGroups() ([]*WAFIPGroup, error) {
var groups []*WAFIPGroup
err := DB.Order("type asc").Order("id asc").Find(&groups).Error
return groups, err
}
func GetWAFIPGroupByID(id uint) (*WAFIPGroup, error) {
group := &WAFIPGroup{}
err := DB.First(group, id).Error
return group, err
}
func ListWAFIPGroupsByIDs(ids []uint) ([]*WAFIPGroup, error) {
if len(ids) == 0 {
return []*WAFIPGroup{}, nil
}
var groups []*WAFIPGroup
err := DB.Where("id IN ?", ids).Order("id asc").Find(&groups).Error
return groups, err
}
func ListDueWAFIPGroups(now time.Time) ([]*WAFIPGroup, error) {
var groups []*WAFIPGroup
err := DB.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
}
func (group *WAFIPGroup) Insert() error {
return DB.Create(group).Error
}
func (group *WAFIPGroup) Update() error {
return DB.Model(&WAFIPGroup{}).Where("id = ?", group.ID).Updates(map[string]any{
"name": group.Name,
"type": group.Type,
"enabled": group.Enabled,
"ip_list": group.IPList,
"auto_config": group.AutoConfig,
"ext_ips": group.ExtIPs,
"subscription_url": group.SubscriptionURL,
"subscription_format": group.SubscriptionFormat,
"subscription_mapping_rule": group.SubscriptionMappingRule,
"sync_interval_minutes": group.SyncIntervalMinutes,
"next_sync_at": group.NextSyncAt,
"last_sync_status": group.LastSyncStatus,
"last_sync_message": group.LastSyncMessage,
"remark": group.Remark,
}).Error
}
func (group *WAFIPGroup) UpdateSyncResult() error {
return DB.Model(&WAFIPGroup{}).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
}
func (group *WAFIPGroup) Delete() error {
return DB.Delete(group).Error
}