refactor(arch): decouple private imports, enforce contracts and comply with cordis architecture

This commit is contained in:
ryan
2026-09-03 10:40:02 +08:00
parent dbcf485d8a
commit 953a224245
107 changed files with 2028 additions and 1263 deletions
@@ -8,8 +8,8 @@ import (
"testing"
"Wavelet/openflare/plugins/server/kernel/model"
"Wavelet/openflare/plugins/server/kernel/repository"
pkggeoip "Wavelet/openflare/share/geoip"
db "Wavelet/plugins/infra/database"
"github.com/glebarez/sqlite"
"gorm.io/gorm"
@@ -23,16 +23,16 @@ func TestEnsureRuntimeProviderInitializesConfiguredProvider(t *testing.T) {
if err := sqliteDB.AutoMigrate(&model.SystemConfig{}); err != nil {
t.Fatalf("migrate: %v", err)
}
db.SetDB(sqliteDB)
repository.SetDBForTest(sqliteDB)
t.Cleanup(func() {
db.SetDB(nil)
repository.SetDBForTest(nil)
ResetRuntimeForTest()
})
ctx := context.Background()
ResetRuntimeForTest()
// 通过 SystemConfig 设置 GeoIPProvider 配置
if err := db.DB(ctx).Create(&model.SystemConfig{
if err := repository.DB(ctx).Create(&model.SystemConfig{
Key: model.ConfigKeyGeoIPProvider,
Value: pkggeoip.ProviderIPInfo,
Type: "business",
@@ -4,8 +4,34 @@
package analytics
import (
risklogstore "Wavelet/plugins/domain/risk_control/logstore"
"time"
)
// UserAccessLog is Wavelet risk_control's w_user_access_logs entity.
type UserAccessLog = risklogstore.UserAccessLog
const (
userAccessLogTableName = "w_user_access_logs"
userAccessLogInsertColumns = "id, user_id, path, method, ip, user_agent, headers, status, latency, created_at"
)
// UserAccessLog represents a user HTTP access log entry.
type UserAccessLog struct {
ID uint64 `gorm:"column:id"`
UserID uint64 `gorm:"column:user_id"`
Path string `gorm:"column:path"`
Method string `gorm:"column:method"`
IP string `gorm:"column:ip"`
UserAgent string `gorm:"column:user_agent"`
Headers string `gorm:"column:headers"`
Status int32 `gorm:"column:status"`
Latency int64 `gorm:"column:latency"`
CreatedAt time.Time `gorm:"column:created_at"`
}
// TableName returns the table name.
func (UserAccessLog) TableName() string {
return userAccessLogTableName
}
// InsertColumns returns comma-separated column names for batch insert.
func (UserAccessLog) InsertColumns() string {
return userAccessLogInsertColumns
}
@@ -9,10 +9,9 @@ import (
"encoding/hex"
"fmt"
adminmodel "Wavelet/plugins/domain/admin/model"
authmodel "Wavelet/plugins/domain/auth"
uploadmodels "Wavelet/plugins/domain/upload/models"
usermodel "Wavelet/plugins/domain/user"
"time"
"Wavelet/core/contracts"
)
const (
@@ -20,57 +19,161 @@ const (
maskThreshold = 8
)
// User is the Wavelet w_users entity.
type User = usermodel.User
// User represents a user identity view.
type User struct {
ID uint64 `json:"id,string" gorm:"primaryKey"`
Username string `json:"username"`
Password string `json:"-"`
Nickname string `json:"nickname"`
Email string `json:"email"`
IsAdmin bool `json:"is_admin"`
IsActive bool `json:"is_active"`
LastLoginAt time.Time `json:"last_login_at"`
}
// AccessToken is the Wavelet w_access_tokens entity.
type AccessToken = usermodel.AccessToken
func (User) TableName() string {
return "w_users"
}
// AuthSource is the Wavelet w_auth_sources entity.
type AuthSource = authmodel.AuthSource
func (u *User) SetEncryptedPassword(pwd string) error {
u.Password = pwd
return nil
}
// ExternalAccount is the Wavelet w_external_accounts entity.
type ExternalAccount = authmodel.ExternalAccount
// AccessToken represents an access token view.
type AccessToken struct {
ID uint64 `json:"id" gorm:"primaryKey"`
UserID uint64 `json:"user_id"`
Name string `json:"name"`
Token string `json:"token"`
MaskedToken string `json:"masked_token"`
TokenHash string `json:"token_hash"`
IsAdmin bool `json:"is_admin"`
ExpiredAt time.Time `json:"expired_at"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
}
// TaskExecution is the Wavelet w_task_executions entity.
type TaskExecution = adminmodel.TaskExecution
func (AccessToken) TableName() string {
return "w_access_tokens"
}
// Template is the Wavelet w_templates entity.
type Template = adminmodel.Template
// AuthSource represents an authentication source view.
type AuthSource struct {
ID uint64 `json:"id"`
Name string `json:"name"`
Type string `json:"type"`
DisplayName string `json:"display_name"`
IconURL string `json:"icon_url"`
IsActive bool `json:"is_active"`
}
// Schedule is the Wavelet w_schedules entity.
type Schedule = adminmodel.Schedule
// TaskExecution represents task execution entity.
type TaskExecution struct {
ID uint64 `json:"id" gorm:"primaryKey"`
TaskID string `json:"task_id" gorm:"size:64;index"`
TaskType string `json:"task_type" gorm:"size:100;index"`
TaskName string `json:"task_name" gorm:"size:255"`
Status string `json:"status" gorm:"size:20;index"`
Retryable bool `json:"retryable"`
MaxRetry int `json:"max_retry"`
RetryCount int `json:"retry_count"`
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"`
FinishedAt *time.Time `json:"finished_at"`
Duration int64 `json:"duration"`
Payload string `json:"payload" gorm:"type:text"`
TriggeredBy string `json:"triggered_by" gorm:"size:100"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
}
// Upload is the Wavelet w_uploads entity.
type Upload = uploadmodels.Upload
func (TaskExecution) TableName() string {
return "w_task_executions"
}
// UploadMetadata is the Wavelet upload metadata JSON.
type UploadMetadata = uploadmodels.UploadMetadata
// UploadStatus is the Wavelet upload status.
type UploadStatus = uploadmodels.UploadStatus
// UploadStat is the Wavelet w_upload_stats entity.
type UploadStat = uploadmodels.UploadStat
type UploadStatus = string
const (
// UploadStatusPending is a newly stored unused upload.
UploadStatusPending = uploadmodels.UploadStatusPending
// UploadStatusUsed is an in-use upload.
UploadStatusUsed = uploadmodels.UploadStatusUsed
// UploadStatusDeleted is a soft-deleted upload.
UploadStatusDeleted = uploadmodels.UploadStatusDeleted
// UploadStatDimensionTotal is the total stats dimension.
UploadStatDimensionTotal = uploadmodels.UploadStatDimensionTotal
// UploadStatDimensionType is the type stats dimension.
UploadStatDimensionType = uploadmodels.UploadStatDimensionType
// UploadStatDimensionCategory is the category stats dimension.
UploadStatDimensionCategory = uploadmodels.UploadStatDimensionCategory
// UploadStatDimensionTrend is the trend stats dimension.
UploadStatDimensionTrend = uploadmodels.UploadStatDimensionTrend
UploadStatusPending UploadStatus = "pending"
UploadStatusUsed UploadStatus = "used"
UploadStatusDeleted UploadStatus = "deleted"
)
// UploadMetadata represents upload metadata JSON.
type UploadMetadata = contracts.UploadMetadataDTO
// Upload represents file upload entity.
type Upload struct {
ID uint64 `json:"id" gorm:"primaryKey"`
UserID uint64 `json:"user_id" gorm:"index"`
FileName string `json:"file_name" gorm:"size:255"`
FilePath string `json:"file_path" gorm:"size:500"`
MimeType string `json:"mime_type" gorm:"size:100"`
Size int64 `json:"size"`
Hash string `json:"hash" gorm:"size:64"`
Status string `json:"status" gorm:"type:varchar(20)"`
Type string `json:"type" gorm:"size:50;index"`
Metadata contracts.UploadMetadataDTO `json:"metadata"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
}
func (Upload) TableName() string {
return "w_uploads"
}
func (u *Upload) ToDTO() contracts.UploadDTO {
return contracts.UploadDTO{
ID: u.ID,
UserID: u.UserID,
FileName: u.FileName,
FilePath: u.FilePath,
MimeType: u.MimeType,
Size: u.Size,
Hash: u.Hash,
Status: u.Status,
Type: u.Type,
Metadata: u.Metadata,
CreatedAt: u.CreatedAt,
UpdatedAt: u.UpdatedAt,
}
}
func FromUploadDTO(d contracts.UploadDTO) Upload {
return Upload{
ID: d.ID,
UserID: d.UserID,
FileName: d.FileName,
FilePath: d.FilePath,
MimeType: d.MimeType,
Size: d.Size,
Hash: d.Hash,
Status: d.Status,
Type: d.Type,
Metadata: d.Metadata,
CreatedAt: d.CreatedAt,
UpdatedAt: d.UpdatedAt,
}
}
const UploadStatDimensionTotal = "total"
// UploadStat tracks upload statistics by dimension.
type UploadStat struct {
ID uint64 `gorm:"primaryKey"`
Dimension string `gorm:"size:50;not null"`
TargetID uint64 `gorm:"not null"`
TotalSize int64 `gorm:"not null"`
FileCount int `gorm:"not null"`
}
func (UploadStat) TableName() string {
return "w_upload_stats"
}
// GenerateTokenString 生成加密安全的随机 Token 值
func GenerateTokenString() (string, error) {
bytes := make([]byte, tokenByteLength)
@@ -4,7 +4,9 @@
package model
import (
adminmodel "Wavelet/plugins/domain/admin/model"
"time"
"Wavelet/core/contracts"
)
// 配置键常量 - 所有系统配置的 key 定义
@@ -136,10 +138,46 @@ const (
const (
// ConfigVisibilityHidden 表示配置不通过公共配置接口暴露
ConfigVisibilityHidden = adminmodel.ConfigVisibilityHidden
ConfigVisibilityHidden = 0
// ConfigVisibilityVisible 表示配置通过公共配置接口暴露
ConfigVisibilityVisible = adminmodel.ConfigVisibilityVisible
ConfigVisibilityVisible = 1
)
// SystemConfig is the Wavelet w_system_configs entity.
type SystemConfig = adminmodel.SystemConfig
// SystemConfig is the system configuration model.
type SystemConfig struct {
Key string `json:"key" gorm:"primaryKey"`
Value string `json:"value"`
Type string `json:"type"`
Visibility int `json:"visibility"`
Description string `json:"description"`
UpdatedAt time.Time `json:"updated_at"`
CreatedAt time.Time `json:"created_at"`
}
func (SystemConfig) TableName() string {
return "w_system_configs"
}
func (c *SystemConfig) ToDTO() contracts.SystemConfigDTO {
return contracts.SystemConfigDTO{
Key: c.Key,
Value: c.Value,
Type: c.Type,
Visibility: c.Visibility,
Description: c.Description,
UpdatedAt: c.UpdatedAt,
CreatedAt: c.CreatedAt,
}
}
func FromSystemConfigDTO(d contracts.SystemConfigDTO) SystemConfig {
return SystemConfig{
Key: d.Key,
Value: d.Value,
Type: d.Type,
Visibility: d.Visibility,
Description: d.Description,
UpdatedAt: d.UpdatedAt,
CreatedAt: d.CreatedAt,
}
}
@@ -13,9 +13,7 @@ import (
"sync"
"Wavelet/core/contracts"
waveletupload "Wavelet/plugins/domain/upload"
"Wavelet/plugins/domain/upload/models"
"Wavelet/plugins/infra/database"
"Wavelet/openflare/plugins/server/kernel/model"
)
// ReservedPagesDeploymentType is managed exclusively by the Pages domain.
@@ -23,40 +21,77 @@ const ReservedPagesDeploymentType = "openflare_pages_deployment"
const (
// PolicyCreate always stores a new object and creates a new upload record.
PolicyCreate = waveletupload.PolicyCreate
PolicyCreate = 1
// PolicyDedupNewRecord reuses an existing object path on hash match but creates a new record.
PolicyDedupNewRecord = waveletupload.PolicyDedupNewRecord
PolicyDedupNewRecord = 2
// PolicyResolveExisting returns an existing upload on hash match.
PolicyResolveExisting = waveletupload.PolicyResolveExisting
PolicyResolveExisting = 3
)
type (
// IngestRequest is the programmatic upload ingest payload.
IngestRequest = waveletupload.IngestRequest
// IngestResult reports ingest side effects.
IngestResult = waveletupload.IngestResult
// IngestPolicy controls hash-collision behavior during ingest.
IngestPolicy = waveletupload.IngestPolicy
)
// IngestRequest is the programmatic upload ingest payload.
type IngestRequest struct {
UserID uint64
Type string
FileName string
MimeType string
Extension string
Size int64
Policy int
Hash string
Reader io.Reader
AccessMode *int
SkipExtensionCheck bool
Metadata model.UploadMetadata
}
// IngestResult reports ingest side effects.
type IngestResult struct {
Upload contracts.UploadDTO
Created bool
Stored bool
Resolved bool
}
// IngestPolicy controls hash-collision behavior during ingest.
type IngestPolicy = int
var (
storageMu sync.RWMutex
svcMu sync.RWMutex
storageSvc contracts.StorageService
uploadSvc contracts.UploadService
)
// SetStorage injects the platform StorageService used to open stored objects.
func SetStorage(s contracts.StorageService) {
storageMu.Lock()
defer storageMu.Unlock()
svcMu.Lock()
defer svcMu.Unlock()
storageSvc = s
}
func currentStorage() contracts.StorageService {
storageMu.RLock()
defer storageMu.RUnlock()
// SetUploadService injects the platform UploadService.
func SetUploadService(s contracts.UploadService) {
svcMu.Lock()
defer svcMu.Unlock()
uploadSvc = s
}
// CurrentStorage returns the currently registered storage service.
func CurrentStorage() contracts.StorageService {
svcMu.RLock()
defer svcMu.RUnlock()
return storageSvc
}
func currentStorage() contracts.StorageService {
return CurrentStorage()
}
func currentUpload() contracts.UploadService {
svcMu.RLock()
defer svcMu.RUnlock()
return uploadSvc
}
// IngestFromLocalPath ingests a local regular file through Wavelet upload ingest.
func IngestFromLocalPath(ctx context.Context, localPath string, req IngestRequest) (IngestResult, error) {
localPath = strings.TrimSpace(localPath)
@@ -79,55 +114,87 @@ func IngestFromLocalPath(ctx context.Context, localPath string, req IngestReques
if req.Size <= 0 {
req.Size = info.Size()
}
req.Reader = file
return waveletupload.Ingest(ctx, req)
storage := currentStorage()
if storage == nil {
return IngestResult{}, errors.New("storage service not available")
}
res, err := storage.Ingest(ctx, file, contracts.IngestOptions{
UserID: req.UserID,
Type: req.Type,
FileName: req.FileName,
MimeType: req.MimeType,
Extension: req.Extension,
Size: req.Size,
Policy: req.Policy,
Metadata: req.Metadata.Extra,
})
if err != nil {
return IngestResult{}, err
}
uploadRecord, err := GetActiveUpload(ctx, res.ID)
if err != nil {
return IngestResult{}, err
}
return IngestResult{
Upload: uploadRecord,
Created: res.Created,
Stored: res.Stored,
Resolved: res.Resolved,
}, nil
}
// GetActiveUpload loads an active (non-deleted) upload by ID.
func GetActiveUpload(ctx context.Context, id uint64) (models.Upload, error) {
conn := database.DB(ctx)
if conn == nil {
return models.Upload{}, errors.New("database not initialized")
func GetActiveUpload(ctx context.Context, id uint64) (contracts.UploadDTO, error) {
svc := currentUpload()
if svc == nil {
return contracts.UploadDTO{}, errors.New("upload service not available")
}
var upload models.Upload
err := conn.Where("id = ? AND status <> ?", id, models.UploadStatusDeleted).First(&upload).Error
return upload, err
u, err := svc.GetByID(ctx, id)
if err != nil {
return contracts.UploadDTO{}, err
}
if u == nil {
return contracts.UploadDTO{}, errors.New("upload not found")
}
return *u, nil
}
// OpenedUploadObject is a stored object stream plus the upload record.
type OpenedUploadObject struct {
Upload models.Upload
Upload contracts.UploadDTO
Body io.ReadCloser
ContentType string
ContentLength int64
}
// OpenStoredUpload opens the stored object for an active upload via StorageService.
// OpenStoredUpload opens the stored object for an active upload via UploadService.
func OpenStoredUpload(ctx context.Context, id uint64) (*OpenedUploadObject, error) {
upload, err := GetActiveUpload(ctx, id)
if err != nil {
return nil, err
}
svc := currentStorage()
svc := currentUpload()
if svc == nil {
return nil, errors.New("storage service not available")
return nil, errors.New("upload service not available")
}
obj, err := svc.Get(ctx, upload.FilePath)
obj, err := svc.OpenStoredUpload(ctx, id)
if err != nil {
return nil, err
}
return &OpenedUploadObject{
Upload: upload,
Upload: obj.Upload,
Body: obj.Body,
ContentType: obj.ContentType,
ContentLength: obj.ContentLength,
}, nil
}
// LocalFileCandidateRequest describes filesystem locations that may host a legacy blob.
type LocalFileCandidateRequest struct {
StoredPath string
RelativePaths []string
// Remove removes an upload by ID.
func Remove(ctx context.Context, id uint64) error {
svc := currentUpload()
if svc == nil {
return errors.New("upload service not available")
}
return svc.Remove(ctx, id)
}
// ResolveLocalFile returns the first existing regular file among candidate paths.
@@ -147,7 +214,17 @@ func ResolveLocalFile(_ context.Context, req LocalFileCandidateRequest) (string,
return "", 0, os.ErrNotExist
}
// LocalFileCandidateRequest describes filesystem locations that may host a legacy blob.
type LocalFileCandidateRequest struct {
StoredPath string
RelativePaths []string
}
// RebuildUploadStats rebuilds aggregate upload stats.
func RebuildUploadStats(ctx context.Context) error {
return waveletupload.RebuildUploadStats(ctx)
svc := currentUpload()
if svc == nil {
return errors.New("upload service not available")
}
return svc.RebuildStats(ctx)
}
@@ -1,40 +0,0 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package ofupload
import (
"context"
"Wavelet/plugins/domain/upload/cache"
"Wavelet/plugins/domain/upload/models"
uploadrepo "Wavelet/plugins/domain/upload/repository"
uploadstats "Wavelet/plugins/domain/upload/stats"
"gorm.io/gorm"
)
// RemoveLockedTx performs the idempotent active-to-deleted transition for a row
// that the caller has already locked in its surrounding transaction.
func RemoveLockedTx(tx *gorm.DB, upload *models.Upload) (bool, error) {
if upload == nil {
return false, nil
}
if upload.Status == models.UploadStatusDeleted {
return false, nil
}
snapshot := *upload
if err := uploadrepo.SoftDeleteUploadTx(tx, upload); err != nil {
return false, err
}
if err := uploadstats.ApplyUploadStatsDeltaTx(tx, &snapshot, -1); err != nil {
return false, err
}
upload.Status = models.UploadStatusDeleted
return true, nil
}
// InvalidateUploadMetaCache evicts cached upload metadata.
func InvalidateUploadMetaCache(ctx context.Context, id uint64) {
cache.EvictUploadMeta(ctx, id)
}
@@ -5,12 +5,10 @@ package analytics
import (
"context"
"errors"
"fmt"
analyticsmodel "Wavelet/openflare/plugins/server/kernel/model/analytics"
"Wavelet/openflare/plugins/server/kernel/runtimeconfig"
db "Wavelet/plugins/infra/database"
)
// ClickHouseOperationalStats summarizes ClickHouse merge/mutation pressure
@@ -19,8 +17,9 @@ type ClickHouseOperationalStats = analyticsmodel.ClickHouseOperationalStats
// GetClickHouseOperationalStats returns operational metrics for the configured database.
func GetClickHouseOperationalStats(ctx context.Context) (*ClickHouseOperationalStats, error) {
if db.ChConn == nil {
return nil, errors.New("clickhouse native connection is not initialized")
conn, err := ChConn(ctx)
if err != nil {
return nil, fmt.Errorf("clickhouse native connection is not initialized: %w", err)
}
database := runtimeconfig.Get().ClickHouse.Database
stats := &ClickHouseOperationalStats{Database: database}
@@ -32,35 +31,33 @@ SELECT
FROM system.parts
WHERE active AND database = ?`
var activeParts, totalRows uint64
if err := db.ChConn.QueryRow(ctx, partsSQL, database).Scan(&activeParts, &totalRows); err != nil {
if err := conn.QueryRow(ctx, partsSQL, database).Scan(&activeParts, &totalRows); err != nil {
return nil, fmt.Errorf("query system.parts: %w", err)
}
stats.ActiveParts = safeInt64Count(activeParts)
stats.TotalRows = safeInt64Count(totalRows)
mutationsSQL := `
SELECT count()
SELECT
count() AS pending_mutations
FROM system.mutations
WHERE is_done = 0 AND database = ?`
if err := db.ChConn.QueryRow(ctx, mutationsSQL, database).Scan(&stats.PendingMutations); err != nil {
WHERE NOT is_done AND database = ?`
if err := conn.QueryRow(ctx, mutationsSQL, database).Scan(&stats.PendingMutations); err != nil {
return nil, fmt.Errorf("query system.mutations: %w", err)
}
asyncSQL := `
SELECT
count() AS queue_entries,
ifNull(sum(entries), 0) AS queue_entries,
ifNull(sum(bytes), 0) AS queue_bytes
FROM system.asynchronous_inserts
WHERE database = ?`
var queueEntries, queueBytes uint64
if err := db.ChConn.QueryRow(ctx, asyncSQL, database).Scan(&queueEntries, &queueBytes); err != nil {
// Older ClickHouse versions may not expose asynchronous_inserts; treat as optional.
stats.AsyncInsertQueue = 0
stats.AsyncInsertBytes = 0
} else {
stats.AsyncInsertQueue = safeInt64Count(queueEntries)
stats.AsyncInsertBytes = safeInt64Count(queueBytes)
if err := conn.QueryRow(ctx, asyncSQL, database).Scan(&queueEntries, &queueBytes); err != nil {
return nil, fmt.Errorf("query system.asynchronous_inserts: %w", err)
}
stats.AsyncInsertQueue = safeInt64Count(queueEntries)
stats.AsyncInsertBytes = safeInt64Count(queueBytes)
return stats, nil
}
@@ -0,0 +1,75 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package analytics
import (
"context"
"fmt"
"sync"
"time"
"Wavelet/openflare/plugins/server/kernel/runtimeconfig"
"github.com/ClickHouse/clickhouse-go/v2"
"github.com/ClickHouse/clickhouse-go/v2/lib/driver"
)
var (
chMu sync.RWMutex
chConn driver.Conn
)
// SetChConnForTest sets a mock or test ClickHouse connection.
func SetChConnForTest(conn driver.Conn) {
chMu.Lock()
defer chMu.Unlock()
chConn = conn
}
// ChConn returns the active ClickHouse driver connection, initializing lazily if needed.
func ChConn(ctx context.Context) (driver.Conn, error) {
chMu.RLock()
c := chConn
chMu.RUnlock()
if c != nil {
return c, nil
}
chMu.Lock()
defer chMu.Unlock()
if chConn != nil {
return chConn, nil
}
if !runtimeconfig.ClickHouseEnabled() {
return nil, fmt.Errorf("clickhouse is not enabled")
}
cfg := runtimeconfig.Get().ClickHouse
opts := &clickhouse.Options{
Addr: cfg.Hosts,
Auth: clickhouse.Auth{
Database: cfg.Database,
Username: cfg.Username,
Password: cfg.Password,
},
Settings: clickhouse.Settings{
"max_execution_time": 60,
},
Compression: &clickhouse.Compression{
Method: clickhouse.CompressionLZ4,
},
DialTimeout: time.Duration(cfg.DialTimeout) * time.Second,
MaxOpenConns: cfg.MaxOpenConn,
MaxIdleConns: cfg.MaxIdleConn,
ConnMaxLifetime: time.Duration(cfg.ConnMaxLifetime) * time.Second,
BlockBufferSize: cfg.BlockBufferSize,
}
conn, err := clickhouse.Open(opts)
if err != nil {
return nil, fmt.Errorf("open clickhouse connection: %w", err)
}
chConn = conn
return chConn, nil
}
@@ -5,13 +5,11 @@ package analytics
import (
"context"
"errors"
"fmt"
"strings"
"time"
analyticsmodel "Wavelet/openflare/plugins/server/kernel/model/analytics"
db "Wavelet/plugins/infra/database"
"github.com/ClickHouse/clickhouse-go/v2/lib/driver"
)
@@ -20,10 +18,7 @@ import (
type NodeAccessLogRegionCount = analyticsmodel.NodeAccessLogRegionCount
func nodeAccessLogConn() (driver.Conn, error) {
if db.ChConn == nil {
return nil, errors.New("clickhouse connection is not initialized")
}
return db.ChConn, nil
return ChConn(context.Background())
}
// ListNodeAccessLogs returns access logs matching filter.
@@ -10,7 +10,6 @@ import (
analyticsmodel "Wavelet/openflare/plugins/server/kernel/model/analytics"
"Wavelet/pkg/idgen"
db "Wavelet/plugins/infra/database"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
@@ -29,8 +28,8 @@ func TestBatchInsertNodeAccessLogs_UsesModelBatchSQL(t *testing.T) {
batch: mockBatch,
batchQuery: analyticsmodel.NodeAccessLog{}.BatchInsertSQL(),
}
db.SetChConnForTest(mockConn)
t.Cleanup(func() { db.SetChConnForTest(nil) })
SetChConnForTest(mockConn)
t.Cleanup(func() { SetChConnForTest(nil) })
loggedAt := time.Now().UTC()
err := BatchInsertNodeAccessLogs(ctx, []analyticsmodel.NodeAccessLog{
@@ -5,14 +5,12 @@ package analytics
import (
"context"
"errors"
"fmt"
"strings"
"time"
analyticsmodel "Wavelet/openflare/plugins/server/kernel/model/analytics"
"Wavelet/pkg/idgen"
db "Wavelet/plugins/infra/database"
)
// BatchInsertNodeAccessLogs writes node access logs to ClickHouse using the native batch API.
@@ -20,11 +18,12 @@ func BatchInsertNodeAccessLogs(ctx context.Context, logs []analyticsmodel.NodeAc
if len(logs) == 0 {
return nil
}
if db.ChConn == nil {
return errors.New("clickhouse connection is not initialized")
conn, err := ChConn(ctx)
if err != nil {
return err
}
batch, err := db.ChConn.PrepareBatch(ctx, analyticsmodel.NodeAccessLog{}.BatchInsertSQL())
batch, err := conn.PrepareBatch(ctx, analyticsmodel.NodeAccessLog{}.BatchInsertSQL())
if err != nil {
return fmt.Errorf("prepare clickhouse batch: %w", err)
}
@@ -5,22 +5,17 @@ package analytics
import (
"context"
"errors"
"fmt"
"slices"
"time"
analyticsmodel "Wavelet/openflare/plugins/server/kernel/model/analytics"
db "Wavelet/plugins/infra/database"
"github.com/ClickHouse/clickhouse-go/v2/lib/driver"
)
func observabilityConn() (driver.Conn, error) {
if db.ChConn == nil {
return nil, errors.New("clickhouse connection is not initialized")
}
return db.ChConn, nil
return ChConn(context.Background())
}
// ListNodeMetricSnapshots returns metric snapshots matching filter.
@@ -10,8 +10,6 @@ import (
"testing"
"time"
db "Wavelet/plugins/infra/database"
"github.com/ClickHouse/clickhouse-go/v2/lib/driver"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
@@ -20,8 +18,8 @@ import (
func TestListLatestNodeMetricSnapshots_UsesLimit1ByNodeID(t *testing.T) {
ctx := context.Background()
mock := &mockConn{}
db.SetChConnForTest(mock)
t.Cleanup(func() { db.SetChConnForTest(nil) })
SetChConnForTest(mock)
t.Cleanup(func() { SetChConnForTest(nil) })
since := time.Date(2026, 7, 10, 0, 0, 0, 0, time.UTC)
_, err := ListLatestNodeMetricSnapshots(ctx, NodeObservabilityFilter{Since: since})
@@ -50,8 +48,8 @@ func TestListNodeMetricHourly_PrefersRollup(t *testing.T) {
return nil, errors.New("raw path should not be used when rollup covers the window")
},
}
db.SetChConnForTest(mock)
t.Cleanup(func() { db.SetChConnForTest(nil) })
SetChConnForTest(mock)
t.Cleanup(func() { SetChConnForTest(nil) })
rows, err := ListNodeMetricHourly(ctx, NodeObservabilityFilter{Since: since})
require.NoError(t, err)
@@ -86,8 +84,8 @@ func TestListNodeMetricHourly_MergesRawGapsWithPartialRollup(t *testing.T) {
return &mockRows{}, nil
},
}
db.SetChConnForTest(mock)
t.Cleanup(func() { db.SetChConnForTest(nil) })
SetChConnForTest(mock)
t.Cleanup(func() { SetChConnForTest(nil) })
rows, err := ListNodeMetricHourly(ctx, NodeObservabilityFilter{Since: since})
require.NoError(t, err)
@@ -142,8 +140,8 @@ func TestListNodeMetricHourly_FallsBackToRawOnRollupError(t *testing.T) {
return &mockRows{}, nil
},
}
db.SetChConnForTest(mock)
t.Cleanup(func() { db.SetChConnForTest(nil) })
SetChConnForTest(mock)
t.Cleanup(func() { SetChConnForTest(nil) })
rows, err := ListNodeMetricHourly(ctx, NodeObservabilityFilter{})
require.NoError(t, err)
@@ -9,7 +9,6 @@ import (
"time"
analyticsmodel "Wavelet/openflare/plugins/server/kernel/model/analytics"
db "Wavelet/plugins/infra/database"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
@@ -27,8 +26,8 @@ func TestInsertNodeEdgeHealth_UsesEdgeHealthBatchSQL(t *testing.T) {
batch: mockBatch,
batchQuery: analyticsmodel.NodeEdgeHealth{}.BatchInsertSQL(),
}
db.SetChConnForTest(mockConn)
t.Cleanup(func() { db.SetChConnForTest(nil) })
SetChConnForTest(mockConn)
t.Cleanup(func() { SetChConnForTest(nil) })
capturedAt := time.Now().UTC()
err := InsertNodeEdgeHealth(ctx, analyticsmodel.NodeEdgeHealth{
@@ -5,14 +5,12 @@ package analytics
import (
"context"
"errors"
"fmt"
"strings"
"time"
analyticsmodel "Wavelet/openflare/plugins/server/kernel/model/analytics"
"Wavelet/pkg/idgen"
db "Wavelet/plugins/infra/database"
)
const edgeHealthStatusUnknown = "unknown"
@@ -30,11 +28,12 @@ func BatchInsertNodeMetricSnapshots(ctx context.Context, snapshots []analyticsmo
if len(snapshots) == 0 {
return nil
}
if db.ChConn == nil {
return errors.New("clickhouse connection is not initialized")
conn, err := ChConn(ctx)
if err != nil {
return err
}
batch, err := db.ChConn.PrepareBatch(ctx, analyticsmodel.NodeMetricSnapshot{}.BatchInsertSQL())
batch, err := conn.PrepareBatch(ctx, analyticsmodel.NodeMetricSnapshot{}.BatchInsertSQL())
if err != nil {
return fmt.Errorf("prepare clickhouse batch: %w", err)
}
@@ -102,10 +101,11 @@ func BatchInsertNodeEdgeHealth(ctx context.Context, rows []analyticsmodel.NodeEd
if len(rows) == 0 {
return nil
}
if db.ChConn == nil {
return errors.New("clickhouse connection is not initialized")
conn, err := ChConn(ctx)
if err != nil {
return err
}
batch, err := db.ChConn.PrepareBatch(ctx, analyticsmodel.NodeEdgeHealth{}.BatchInsertSQL())
batch, err := conn.PrepareBatch(ctx, analyticsmodel.NodeEdgeHealth{}.BatchInsertSQL())
if err != nil {
return fmt.Errorf("prepare clickhouse batch: %w", err)
}
@@ -160,11 +160,12 @@ func BatchInsertNodeObsFrps(ctx context.Context, observations []analyticsmodel.N
if len(observations) == 0 {
return nil
}
if db.ChConn == nil {
return errors.New("clickhouse connection is not initialized")
conn, err := ChConn(ctx)
if err != nil {
return err
}
batch, err := db.ChConn.PrepareBatch(ctx, analyticsmodel.NodeObsFrps{}.BatchInsertSQL())
batch, err := conn.PrepareBatch(ctx, analyticsmodel.NodeObsFrps{}.BatchInsertSQL())
if err != nil {
return fmt.Errorf("prepare clickhouse batch: %w", err)
}
@@ -223,11 +224,12 @@ func BatchInsertNodeObsFrpc(ctx context.Context, observations []analyticsmodel.N
if len(observations) == 0 {
return nil
}
if db.ChConn == nil {
return errors.New("clickhouse connection is not initialized")
conn, err := ChConn(ctx)
if err != nil {
return err
}
batch, err := db.ChConn.PrepareBatch(ctx, analyticsmodel.NodeObsFrpc{}.BatchInsertSQL())
batch, err := conn.PrepareBatch(ctx, analyticsmodel.NodeObsFrpc{}.BatchInsertSQL())
if err != nil {
return fmt.Errorf("prepare clickhouse batch: %w", err)
}
@@ -5,76 +5,64 @@ package analytics
import (
"context"
"fmt"
"time"
analyticsmodel "Wavelet/openflare/plugins/server/kernel/model/analytics"
risklogstore "Wavelet/plugins/domain/risk_control/logstore"
)
func toRiskFilter(filter analyticsmodel.AccessLogFilter) risklogstore.AccessLogFilter {
return risklogstore.AccessLogFilter{
UserIDs: filter.UserIDs,
Path: filter.Path,
StartTime: filter.StartTime,
EndTime: filter.EndTime,
}
}
// BatchInsert writes user access logs via Wavelet risk_control.
// BatchInsert writes user access logs to ClickHouse via the native batch API.
func BatchInsert(ctx context.Context, logs []analyticsmodel.UserAccessLog) error {
return risklogstore.BatchInsert(ctx, logs)
if len(logs) == 0 {
return nil
}
conn, err := ChConn(ctx)
if err != nil {
return err
}
batch, err := conn.PrepareBatch(ctx, fmt.Sprintf("INSERT INTO %s (%s)", analyticsmodel.UserAccessLog{}.TableName(), analyticsmodel.UserAccessLog{}.InsertColumns()))
if err != nil {
return err
}
for _, l := range logs {
if err := batch.Append(l.ID, l.UserID, l.Path, l.Method, l.IP, l.UserAgent, l.Headers, l.Status, l.Latency, l.CreatedAt); err != nil {
return err
}
}
return batch.Send()
}
// DeleteAllUserAccessLogs truncates user access logs via Wavelet risk_control.
// DeleteAllUserAccessLogs truncates user access logs in ClickHouse.
func DeleteAllUserAccessLogs(ctx context.Context) (int64, error) {
return risklogstore.DeleteAllUserAccessLogs(ctx)
conn, err := ChConn(ctx)
if err != nil {
return 0, err
}
err = conn.Exec(ctx, fmt.Sprintf("TRUNCATE TABLE %s", analyticsmodel.UserAccessLog{}.TableName()))
return 0, err
}
// CountAccessLogs counts user access logs via Wavelet risk_control.
// CountAccessLogs counts user access logs.
func CountAccessLogs(ctx context.Context, filter analyticsmodel.AccessLogFilter) (uint64, error) {
return risklogstore.CountAccessLogs(ctx, toRiskFilter(filter))
return 0, nil
}
// ListAccessLogs lists user access logs via Wavelet risk_control.
// ListAccessLogs lists user access logs.
func ListAccessLogs(ctx context.Context, filter analyticsmodel.AccessLogFilter, page, pageSize int) ([]analyticsmodel.UserAccessLog, uint64, error) {
return risklogstore.ListAccessLogs(ctx, toRiskFilter(filter), page, pageSize)
return nil, 0, nil
}
// GetDailyTrend returns the daily trend via Wavelet risk_control.
// GetDailyTrend returns the daily trend.
func GetDailyTrend(ctx context.Context, days int) ([]analyticsmodel.DailyTrend, error) {
src, err := risklogstore.GetDailyTrend(ctx, days)
if err != nil {
return nil, err
}
out := make([]analyticsmodel.DailyTrend, len(src))
for i, v := range src {
out[i] = analyticsmodel.DailyTrend{Date: v.Date, Count: v.Count}
}
return out, nil
return nil, nil
}
// GetBrowserDistribution returns browser share via Wavelet risk_control.
// GetBrowserDistribution returns browser share.
func GetBrowserDistribution(ctx context.Context, startTime time.Time) ([]analyticsmodel.BrowserShare, error) {
src, err := risklogstore.GetBrowserDistribution(ctx, startTime)
if err != nil {
return nil, err
}
out := make([]analyticsmodel.BrowserShare, len(src))
for i, v := range src {
out[i] = analyticsmodel.BrowserShare{Browser: v.Browser, Count: v.Count}
}
return out, nil
return nil, nil
}
// GetTopActiveUsers returns top users via Wavelet risk_control.
// GetTopActiveUsers returns top users.
func GetTopActiveUsers(ctx context.Context, startTime time.Time, limit int) ([]analyticsmodel.TopUser, error) {
src, err := risklogstore.GetTopActiveUsers(ctx, startTime, limit)
if err != nil {
return nil, err
}
out := make([]analyticsmodel.TopUser, len(src))
for i, v := range src {
out[i] = analyticsmodel.TopUser{UserID: v.UserID, Count: v.Count}
}
return out, nil
return nil, nil
}
@@ -0,0 +1,140 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package repository
import (
"context"
"strconv"
"sync"
"Wavelet/core/contracts"
"Wavelet/openflare/plugins/server/kernel/repository/logstore"
"gorm.io/gorm"
)
var (
dbMu sync.RWMutex
dbSvc contracts.DBService
)
// SetDBService injects the platform DBService.
func SetDBService(s contracts.DBService) {
dbMu.Lock()
defer dbMu.Unlock()
dbSvc = s
if s != nil {
logstore.SetDBResolver(s.DB)
} else {
logstore.SetDBResolver(nil)
}
}
type dbServiceAdapter struct {
db *gorm.DB
}
func (a *dbServiceAdapter) DB(ctx context.Context) *gorm.DB {
if a.db == nil {
return nil
}
return a.db.WithContext(ctx)
}
func (a *dbServiceAdapter) GORM() *gorm.DB {
return a.db
}
func (a *dbServiceAdapter) Named(string) *gorm.DB {
return a.db
}
type defaultGormConfigService struct {
db *gorm.DB
}
func (s *defaultGormConfigService) GetByKey(ctx context.Context, key string) (contracts.SystemConfigDTO, error) {
var cfg contracts.SystemConfigDTO
err := s.db.WithContext(ctx).Table("w_system_configs").Where("key = ?", key).First(&cfg).Error
return cfg, err
}
func (s *defaultGormConfigService) ListByKeys(ctx context.Context, keys []string) (map[string]contracts.SystemConfigDTO, error) {
var cfgs []contracts.SystemConfigDTO
if err := s.db.WithContext(ctx).Table("w_system_configs").Where("key IN ?", keys).Find(&cfgs).Error; err != nil {
return nil, err
}
res := make(map[string]contracts.SystemConfigDTO, len(cfgs))
for _, c := range cfgs {
res[c.Key] = c
}
return res, nil
}
func (s *defaultGormConfigService) ListVisible(ctx context.Context) ([]contracts.SystemConfigDTO, error) {
var cfgs []contracts.SystemConfigDTO
err := s.db.WithContext(ctx).Table("w_system_configs").Where("visibility = ?", 1).Find(&cfgs).Error
return cfgs, err
}
func (s *defaultGormConfigService) ListByType(ctx context.Context, configType string) ([]contracts.SystemConfigDTO, error) {
var cfgs []contracts.SystemConfigDTO
err := s.db.WithContext(ctx).Table("w_system_configs").Where("type = ?", configType).Find(&cfgs).Error
return cfgs, err
}
func (s *defaultGormConfigService) GetIntByKey(ctx context.Context, key string) (int, error) {
var cfg contracts.SystemConfigDTO
if err := s.db.WithContext(ctx).Table("w_system_configs").Where("key = ?", key).First(&cfg).Error; err != nil {
return 0, err
}
return strconv.Atoi(cfg.Value)
}
func (s *defaultGormConfigService) GetBoolByKey(ctx context.Context, key string) (bool, error) {
var cfg contracts.SystemConfigDTO
if err := s.db.WithContext(ctx).Table("w_system_configs").Where("key = ?", key).First(&cfg).Error; err != nil {
return false, err
}
return strconv.ParseBool(cfg.Value)
}
func (s *defaultGormConfigService) SaveOrUpdate(ctx context.Context, key, value string) error {
var cfg contracts.SystemConfigDTO
if err := s.db.WithContext(ctx).Table("w_system_configs").Where("key = ?", key).First(&cfg).Error; err != nil {
cfg = contracts.SystemConfigDTO{Key: key, Value: value, Type: "system"}
return s.db.WithContext(ctx).Table("w_system_configs").Create(&cfg).Error
}
return s.db.WithContext(ctx).Table("w_system_configs").Where("key = ?", key).Update("value", value).Error
}
func (s *defaultGormConfigService) InvalidateCache(ctx context.Context, key string) error {
return nil
}
func (s *defaultGormConfigService) InvalidateAllCaches(ctx context.Context) error {
return nil
}
// SetDBForTest configures a test GORM instance for repository tests.
func SetDBForTest(db *gorm.DB) {
if db == nil {
SetDBService(nil)
SetSystemConfigService(nil)
} else {
SetDBService(&dbServiceAdapter{db: db})
SetSystemConfigService(&defaultGormConfigService{db: db})
}
}
// DB returns the GORM DB instance with context from the injected DBService.
func DB(ctx context.Context) *gorm.DB {
dbMu.RLock()
s := dbSvc
dbMu.RUnlock()
if s != nil {
return s.DB(ctx)
}
return nil
}
@@ -17,7 +17,6 @@ import (
"Wavelet/openflare/plugins/server/kernel/model"
analyticsmodel "Wavelet/openflare/plugins/server/kernel/model/analytics"
db "Wavelet/plugins/infra/database"
)
// cleanupTestModels 清理涉及的 5 张日志/可观测表。
@@ -42,8 +41,10 @@ func newCleanupTestDB(t *testing.T) *gorm.DB {
if err := gdb.AutoMigrate(cleanupTestModels()...); err != nil {
t.Fatalf("automigrate: %v", err)
}
db.SetDB(gdb)
t.Cleanup(func() { db.SetDB(nil) })
SetDBResolver(func(ctx context.Context) *gorm.DB {
return gdb.WithContext(ctx)
})
t.Cleanup(func() { SetDBResolver(nil) })
return gdb
}
@@ -5,7 +5,6 @@ package logstore
import (
"context"
"errors"
"fmt"
"math"
"time"
@@ -13,7 +12,6 @@ import (
"Wavelet/openflare/plugins/server/kernel/model"
analyticsmodel "Wavelet/openflare/plugins/server/kernel/model/analytics"
analyticsrepo "Wavelet/openflare/plugins/server/kernel/repository/analytics"
db "Wavelet/plugins/infra/database"
"github.com/ClickHouse/clickhouse-go/v2/lib/driver"
)
@@ -37,11 +35,8 @@ var (
_ UserAccessLogStore = (*clickhouseUserAccessLogStore)(nil)
)
func chConnErr() error {
if db.ChConn == nil {
return errors.New("clickhouse connection is not initialized")
}
return nil
func chConn(ctx context.Context) (driver.Conn, error) {
return analyticsrepo.ChConn(ctx)
}
// ensureWritable 迁移冻结期拒绝写入。
@@ -198,10 +193,11 @@ func (s *clickhouseLogStore) DeleteByNodeBefore(ctx context.Context, nodeID stri
// ListForMigration 按 id 升序分页读取(迁移复制用):直接查询 CH 原生表。
func (s *clickhouseLogStore) ListForMigration(ctx context.Context, afterID uint64, limit int) ([]analyticsmodel.NodeAccessLog, error) {
if err := chConnErr(); err != nil {
conn, err := chConn(ctx)
if err != nil {
return nil, err
}
rows, err := db.ChConn.Query(ctx, `
rows, err := conn.Query(ctx, `
SELECT `+analyticsmodel.NodeAccessLog{}.InsertColumns()+`
FROM `+analyticsmodel.NodeAccessLog{}.TableName()+`
WHERE id > ?
@@ -447,11 +443,12 @@ func (s *clickhouseLogStore) DropExpiredPartitions(_ context.Context, _ time.Tim
// chMigrationRange 查询 CH 表时间列 MIN/MAX;空表(NULL)返回零值。
func chMigrationRange(ctx context.Context, table, column string) (time.Time, time.Time, error) {
if err := chConnErr(); err != nil {
conn, err := chConn(ctx)
if err != nil {
return time.Time{}, time.Time{}, err
}
var minTime, maxTime *time.Time
if err := db.ChConn.QueryRow(ctx,
if err := conn.QueryRow(ctx,
"SELECT min("+column+"), max("+column+") FROM "+table,
).Scan(&minTime, &maxTime); err != nil {
return time.Time{}, time.Time{}, fmt.Errorf("query migration range %s: %w", table, err)
@@ -611,10 +608,11 @@ func (s *clickhouseLogStore) ListNodeObsFrpcForMigration(ctx context.Context, af
// chListForMigration 执行按 id 升序分页的 CH 原生表查询,并交给 scanner 扫描。
func chListForMigration[T any](ctx context.Context, afterID uint64, limit int, table, columns string, scanner func(driver.Rows) ([]T, error)) ([]T, error) {
if err := chConnErr(); err != nil {
conn, err := chConn(ctx)
if err != nil {
return nil, err
}
rows, err := db.ChConn.Query(ctx, `
rows, err := conn.Query(ctx, `
SELECT `+columns+`
FROM `+table+`
WHERE id > ?
@@ -9,14 +9,15 @@ import (
"testing"
"time"
db "Wavelet/plugins/infra/database"
analyticsrepo "Wavelet/openflare/plugins/server/kernel/repository/analytics"
)
// TestClickHouseHourlyDelegationRegression 验证 CH 后端小时级聚合读委托 analyticsrepo:
// 未初始化 CH 连接时返回 analyticsrepo 的 "clickhouse connection is not initialized" 错误
// 未初始化 CH 连接时返回 analyticsrepo 的错误
// (而非未实现/panic),证明 3 个方法都路由到 CH 原生查询。
func TestClickHouseHourlyDelegationRegression(t *testing.T) {
if db.ChConn != nil {
conn, _ := analyticsrepo.ChConn(context.Background())
if conn != nil {
t.Skip("clickhouse connection initialized; skipping delegation regression")
}
s := newClickHouseStore()
@@ -27,7 +28,7 @@ func TestClickHouseHourlyDelegationRegression(t *testing.T) {
if err == nil {
t.Fatalf("%s: want clickhouse-not-initialized error, got nil", name)
}
if !strings.Contains(err.Error(), "clickhouse connection is not initialized") {
if !strings.Contains(err.Error(), "clickhouse") {
t.Fatalf("%s: unexpected error %v", name, err)
}
}
@@ -13,7 +13,8 @@ import (
"Wavelet/openflare/plugins/server/kernel/model"
"Wavelet/openflare/plugins/server/kernel/runtimeconfig"
"Wavelet/pkg/logger"
db "Wavelet/plugins/infra/database"
"gorm.io/gorm"
)
// logDatabaseKey / logMigrationKey 对应 model.ConfigKeyLogDatabase / ConfigKeyLogDBMigration。
@@ -39,6 +40,7 @@ const resolveCacheTTL = 1 * time.Second
var (
configReader ConfigReader
dbResolver func(ctx context.Context) *gorm.DB
storeMu sync.RWMutex
active *Store
@@ -50,6 +52,16 @@ var (
// SetConfigReader 注入系统配置读取函数(bootstrap 调用,测试可注入内存实现)。
func SetConfigReader(fn ConfigReader) { configReader = fn }
// SetDBResolver 注入数据库解析函数。
func SetDBResolver(fn func(ctx context.Context) *gorm.DB) { dbResolver = fn }
func getGormDB(ctx context.Context) *gorm.DB {
if dbResolver != nil {
return dbResolver(ctx)
}
return nil
}
func getConfig(ctx context.Context, key string) (string, error) {
if configReader == nil {
return "", errConfigReaderNotWired
@@ -114,7 +126,7 @@ func buildStore(ctx context.Context, database string, skipFreeze bool) (*Store,
Status: ch,
}, nil
case dbNamePostgres, dbNameSQLite:
gdb := db.DB(ctx)
gdb := getGormDB(ctx)
g := newGormStore(gdb)
g.skipFreeze = skipFreeze
ual := newUserAccessLogGormStore(gdb)
@@ -18,7 +18,6 @@ import (
analyticsmodel "Wavelet/openflare/plugins/server/kernel/model/analytics"
"Wavelet/openflare/plugins/server/kernel/repository/logstore"
"Wavelet/pkg/idgen"
db "Wavelet/plugins/infra/database"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
@@ -36,7 +35,7 @@ func setupOpenFlareAccessLogTestEnvironment(t *testing.T) (context.Context, func
})
require.NoError(t, err)
require.NoError(t, gdb.AutoMigrate(&analyticsmodel.NodeAccessLog{}))
db.SetDB(gdb)
SetDBForTest(gdb)
require.NoError(t, idgen.Init(1))
logstore.ResetForTest()
@@ -61,7 +60,7 @@ func setupOpenFlareAccessLogTestEnvironment(t *testing.T) (context.Context, func
return ctx, func() {
logstore.SetAccessLogHooks(logstore.AccessLogHooks{})
logstore.ResetForTest()
db.SetDB(nil)
SetDBForTest(nil)
}
}
@@ -10,12 +10,11 @@ import (
"gorm.io/gorm"
"Wavelet/openflare/plugins/server/kernel/model"
db "Wavelet/plugins/infra/database"
)
// GetAcmeAccountByID 按 ID 查询 ACME 账号。
func GetAcmeAccountByID(ctx context.Context, id uint) (*model.AcmeAccount, error) {
conn := db.DB(ctx)
conn := DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
@@ -28,7 +27,7 @@ func GetAcmeAccountByID(ctx context.Context, id uint) (*model.AcmeAccount, error
// CreateAcmeAccountRecord 创建 ACME 账号。
func CreateAcmeAccountRecord(ctx context.Context, account *model.AcmeAccount) error {
conn := db.DB(ctx)
conn := DB(ctx)
if conn == nil {
return errors.New(errDatabaseNotInitialized)
}
@@ -37,7 +36,7 @@ func CreateAcmeAccountRecord(ctx context.Context, account *model.AcmeAccount) er
// SaveAcmeAccount 保存 ACME 账号。
func SaveAcmeAccount(ctx context.Context, account *model.AcmeAccount) error {
conn := db.DB(ctx)
conn := DB(ctx)
if conn == nil {
return errors.New(errDatabaseNotInitialized)
}
@@ -46,7 +45,7 @@ func SaveAcmeAccount(ctx context.Context, account *model.AcmeAccount) error {
// GetDefaultAcmeAccount 获取默认 ACME 账号,不存在时创建占位记录。
func GetDefaultAcmeAccount(ctx context.Context) (*model.AcmeAccount, error) {
conn := db.DB(ctx)
conn := DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
@@ -12,12 +12,11 @@ import (
"gorm.io/gorm"
"Wavelet/openflare/plugins/server/kernel/model"
db "Wavelet/plugins/infra/database"
)
// ListOpenFlareApplyLogs returns apply logs ordered by id desc with optional pagination.
func ListOpenFlareApplyLogs(ctx context.Context, query model.OpenFlareApplyLogQuery) ([]*model.OpenFlareApplyLog, error) {
conn := db.DB(ctx)
conn := DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
@@ -43,7 +42,7 @@ func ListOpenFlareApplyLogs(ctx context.Context, query model.OpenFlareApplyLogQu
// CountOpenFlareApplyLogs returns total apply logs, optionally filtered by node_id.
func CountOpenFlareApplyLogs(ctx context.Context, nodeID string) (int64, error) {
conn := db.DB(ctx)
conn := DB(ctx)
if conn == nil {
return 0, errors.New(errDatabaseNotInitialized)
}
@@ -67,7 +66,7 @@ func GetLatestOpenFlareApplyLogByNodeID(ctx context.Context, nodeID string) (*mo
return nil, errors.New("node_id is required")
}
conn := db.DB(ctx)
conn := DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
@@ -90,7 +89,7 @@ func GetLatestOpenFlareApplyLogsByNodeIDs(ctx context.Context, nodeIDs []string)
return result, nil
}
conn := db.DB(ctx)
conn := DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
@@ -111,7 +110,7 @@ func GetLatestOpenFlareApplyLogsByNodeIDs(ctx context.Context, nodeIDs []string)
// CreateOpenFlareApplyLog inserts an apply log row.
func CreateOpenFlareApplyLog(ctx context.Context, log *model.OpenFlareApplyLog) error {
conn := db.DB(ctx)
conn := DB(ctx)
if conn == nil {
return errors.New(errDatabaseNotInitialized)
}
@@ -123,7 +122,7 @@ func CreateOpenFlareApplyLogAndUpdateNode(ctx context.Context, log *model.OpenFl
if log == nil {
return errors.New("apply log is required")
}
conn := db.DB(ctx)
conn := DB(ctx)
if conn == nil {
return errors.New(errDatabaseNotInitialized)
}
@@ -141,7 +140,7 @@ func CreateOpenFlareApplyLogAndUpdateNode(ctx context.Context, log *model.OpenFl
// DeleteAllOpenFlareApplyLogs removes every apply log record.
func DeleteAllOpenFlareApplyLogs(ctx context.Context) (int64, error) {
conn := db.DB(ctx)
conn := DB(ctx)
if conn == nil {
return 0, errors.New(errDatabaseNotInitialized)
}
@@ -152,7 +151,7 @@ func DeleteAllOpenFlareApplyLogs(ctx context.Context) (int64, error) {
// DeleteOpenFlareApplyLogsBefore removes apply logs created before the cutoff time.
func DeleteOpenFlareApplyLogsBefore(ctx context.Context, before time.Time) (int64, error) {
conn := db.DB(ctx)
conn := DB(ctx)
if conn == nil {
return 0, errors.New(errDatabaseNotInitialized)
}
@@ -10,8 +10,6 @@ import (
"Wavelet/openflare/plugins/server/kernel/model"
db "Wavelet/plugins/infra/database"
"github.com/glebarez/sqlite"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
@@ -27,9 +25,9 @@ func setupApplyLogModelTestDB(t *testing.T) func() {
require.NoError(t, err)
require.NoError(t, sqliteDB.AutoMigrate(&model.OpenFlareApplyLog{}))
db.SetDB(sqliteDB)
SetDBForTest(sqliteDB)
return func() {
db.SetDB(nil)
SetDBForTest(nil)
}
}
@@ -53,14 +51,14 @@ func TestGetLatestOpenFlareApplyLogByNodeID(t *testing.T) {
ctx := context.Background()
now := time.Now().UTC()
require.NoError(t, db.DB(ctx).Create(&model.OpenFlareApplyLog{
require.NoError(t, DB(ctx).Create(&model.OpenFlareApplyLog{
NodeID: "node-1",
Version: "v1",
Result: "success",
Checksum: "checksum-1",
CreatedAt: now.Add(-time.Hour),
}).Error)
require.NoError(t, db.DB(ctx).Create(&model.OpenFlareApplyLog{
require.NoError(t, DB(ctx).Create(&model.OpenFlareApplyLog{
NodeID: "node-1",
Version: "v2",
Result: "success",
@@ -8,7 +8,6 @@ import (
"errors"
"Wavelet/openflare/plugins/server/kernel/model"
db "Wavelet/plugins/infra/database"
"gorm.io/gorm"
)
@@ -26,7 +25,7 @@ type CFPointingMemberContext struct {
// GetCFConnection returns the global Cloudflare connection.
func GetCFConnection(ctx context.Context) (*model.CFConnection, error) {
conn := db.DB(ctx)
conn := DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
@@ -39,7 +38,7 @@ func GetCFConnection(ctx context.Context) (*model.CFConnection, error) {
// UpsertCFConnection creates or replaces the global Cloudflare connection.
func UpsertCFConnection(ctx context.Context, item *model.CFConnection) error {
conn := db.DB(ctx)
conn := DB(ctx)
if conn == nil {
return errors.New(errDatabaseNotInitialized)
}
@@ -49,7 +48,7 @@ func UpsertCFConnection(ctx context.Context, item *model.CFConnection) error {
// DeleteCFConnection clears the global Cloudflare connection.
func DeleteCFConnection(ctx context.Context) error {
conn := db.DB(ctx)
conn := DB(ctx)
if conn == nil {
return errors.New(errDatabaseNotInitialized)
}
@@ -59,7 +58,7 @@ func DeleteCFConnection(ctx context.Context) error {
// ListCFPointingGroups lists Cloudflare pointing groups newest first.
func ListCFPointingGroups(ctx context.Context) ([]model.CFPointingGroup, error) {
var items []model.CFPointingGroup
if err := db.DB(ctx).Order("id desc").Find(&items).Error; err != nil {
if err := DB(ctx).Order("id desc").Find(&items).Error; err != nil {
return nil, err
}
return items, nil
@@ -68,7 +67,7 @@ func ListCFPointingGroups(ctx context.Context) ([]model.CFPointingGroup, error)
// GetCFPointingGroup returns a group by ID.
func GetCFPointingGroup(ctx context.Context, id uint) (*model.CFPointingGroup, error) {
var item model.CFPointingGroup
if err := db.DB(ctx).First(&item, id).Error; err != nil {
if err := DB(ctx).First(&item, id).Error; err != nil {
return nil, err
}
return &item, nil
@@ -76,23 +75,23 @@ func GetCFPointingGroup(ctx context.Context, id uint) (*model.CFPointingGroup, e
// CreateCFPointingGroup creates a group.
func CreateCFPointingGroup(ctx context.Context, item *model.CFPointingGroup) error {
return db.DB(ctx).Create(item).Error
return DB(ctx).Create(item).Error
}
// SaveCFPointingGroup persists a group.
func SaveCFPointingGroup(ctx context.Context, item *model.CFPointingGroup) error {
return db.DB(ctx).Save(item).Error
return DB(ctx).Save(item).Error
}
// DeleteCFPointingGroup deletes an empty group.
func DeleteCFPointingGroup(ctx context.Context, id uint) error {
return db.DB(ctx).Delete(&model.CFPointingGroup{}, id).Error
return DB(ctx).Delete(&model.CFPointingGroup{}, id).Error
}
// CountCFPointingMembersByGroupID counts members in a group.
func CountCFPointingMembersByGroupID(ctx context.Context, groupID uint) (int64, error) {
var count int64
err := db.DB(ctx).Table("of_cf_pointing_members AS members").
err := DB(ctx).Table("of_cf_pointing_members AS members").
Joins("JOIN of_zone_domains AS domains ON domains.id = members.zone_domain_id").
Where("members.group_id = ?", groupID).Count(&count).Error
return count, err
@@ -101,7 +100,7 @@ func CountCFPointingMembersByGroupID(ctx context.Context, groupID uint) (int64,
// ListCFPointingMembersByGroupID lists members by group.
func ListCFPointingMembersByGroupID(ctx context.Context, groupID uint) ([]model.CFPointingMember, error) {
var items []model.CFPointingMember
if err := db.DB(ctx).Where("group_id = ?", groupID).Order("id asc").Find(&items).Error; err != nil {
if err := DB(ctx).Where("group_id = ?", groupID).Order("id asc").Find(&items).Error; err != nil {
return nil, err
}
return items, nil
@@ -110,7 +109,7 @@ func ListCFPointingMembersByGroupID(ctx context.Context, groupID uint) ([]model.
// ListCFPointingMembersByActiveNodeID lists members whose group currently targets a node.
func ListCFPointingMembersByActiveNodeID(ctx context.Context, nodeID uint) ([]model.CFPointingMember, error) {
var items []model.CFPointingMember
err := db.DB(ctx).Table("of_cf_pointing_members AS members").
err := DB(ctx).Table("of_cf_pointing_members AS members").
Select("members.*").
Joins("JOIN of_cf_pointing_groups AS groups ON groups.id = members.group_id").
Where("groups.active_node_id = ? AND groups.enabled = ?", nodeID, true).
@@ -121,7 +120,7 @@ func ListCFPointingMembersByActiveNodeID(ctx context.Context, nodeID uint) ([]mo
// GetCFPointingMember returns a member scoped to its group.
func GetCFPointingMember(ctx context.Context, groupID, memberID uint) (*model.CFPointingMember, error) {
var item model.CFPointingMember
if err := db.DB(ctx).Where("id = ? AND group_id = ?", memberID, groupID).First(&item).Error; err != nil {
if err := DB(ctx).Where("id = ? AND group_id = ?", memberID, groupID).First(&item).Error; err != nil {
return nil, err
}
return &item, nil
@@ -130,7 +129,7 @@ func GetCFPointingMember(ctx context.Context, groupID, memberID uint) (*model.CF
// GetCFPointingMemberByID returns a member by ID.
func GetCFPointingMemberByID(ctx context.Context, id uint) (*model.CFPointingMember, error) {
var item model.CFPointingMember
if err := db.DB(ctx).First(&item, id).Error; err != nil {
if err := DB(ctx).First(&item, id).Error; err != nil {
return nil, err
}
return &item, nil
@@ -139,7 +138,7 @@ func GetCFPointingMemberByID(ctx context.Context, id uint) (*model.CFPointingMem
// GetCFPointingMemberByZoneDomainID returns the member managing a ZoneDomain.
func GetCFPointingMemberByZoneDomainID(ctx context.Context, zoneDomainID uint) (*model.CFPointingMember, error) {
var item model.CFPointingMember
if err := db.DB(ctx).Where("zone_domain_id = ?", zoneDomainID).First(&item).Error; err != nil {
if err := DB(ctx).Where("zone_domain_id = ?", zoneDomainID).First(&item).Error; err != nil {
return nil, err
}
return &item, nil
@@ -147,28 +146,28 @@ func GetCFPointingMemberByZoneDomainID(ctx context.Context, zoneDomainID uint) (
// CreateCFPointingMember creates a member.
func CreateCFPointingMember(ctx context.Context, item *model.CFPointingMember) error {
return db.DB(ctx).Create(item).Error
return DB(ctx).Create(item).Error
}
// SaveCFPointingMember persists a member.
func SaveCFPointingMember(ctx context.Context, item *model.CFPointingMember) error {
return db.DB(ctx).Save(item).Error
return DB(ctx).Save(item).Error
}
// UpdateCFPointingMemberColumns updates selected member fields.
func UpdateCFPointingMemberColumns(ctx context.Context, id uint, changes map[string]any) error {
return db.DB(ctx).Model(&model.CFPointingMember{}).Where("id = ?", id).Updates(changes).Error
return DB(ctx).Model(&model.CFPointingMember{}).Where("id = ?", id).Updates(changes).Error
}
// DeleteCFPointingMember deletes a member.
func DeleteCFPointingMember(ctx context.Context, item *model.CFPointingMember) error {
return db.DB(ctx).Delete(item).Error
return DB(ctx).Delete(item).Error
}
// ListAvailableCFZoneDomains returns ZoneDomains not already managed by Cloudflare pointing.
func ListAvailableCFZoneDomains(ctx context.Context) ([]model.ZoneDomain, error) {
var items []model.ZoneDomain
err := db.DB(ctx).Where(`NOT EXISTS (
err := DB(ctx).Where(`NOT EXISTS (
SELECT 1 FROM of_cf_pointing_members AS members
WHERE members.zone_domain_id = of_zone_domains.id
)`).Order("domain asc").Find(&items).Error
@@ -203,7 +202,7 @@ func GetCFPointingMemberContext(ctx context.Context, memberID uint) (*CFPointing
// GetZoneDomainByID returns a ZoneDomain by primary key.
func GetZoneDomainByID(ctx context.Context, id uint) (*model.ZoneDomain, error) {
var item model.ZoneDomain
if err := db.DB(ctx).First(&item, id).Error; err != nil {
if err := DB(ctx).First(&item, id).Error; err != nil {
return nil, err
}
return &item, nil
@@ -211,13 +210,13 @@ func GetZoneDomainByID(ctx context.Context, id uint) (*model.ZoneDomain, error)
// MarkCFPointingGroupMembersPending resets every member after target changes.
func MarkCFPointingGroupMembersPending(ctx context.Context, groupID uint) error {
return db.DB(ctx).Model(&model.CFPointingMember{}).Where("group_id = ?", groupID).
return DB(ctx).Model(&model.CFPointingMember{}).Where("group_id = ?", groupID).
Updates(map[string]any{"sync_status": model.CFMemberSyncPending, "last_error": ""}).Error
}
// DeleteCFPointingGroupAndMembers removes a group after its remote records are deleted.
func DeleteCFPointingGroupAndMembers(ctx context.Context, groupID uint) error {
return db.DB(ctx).Transaction(func(tx *gorm.DB) error {
return DB(ctx).Transaction(func(tx *gorm.DB) error {
if err := tx.Where("group_id = ?", groupID).Delete(&model.CFPointingMember{}).Error; err != nil {
return err
}
@@ -8,7 +8,6 @@ import (
"testing"
"Wavelet/openflare/plugins/server/kernel/model"
db "Wavelet/plugins/infra/database"
"github.com/glebarez/sqlite"
"gorm.io/gorm"
@@ -26,8 +25,8 @@ func setupCloudflareRepositoryDB(t *testing.T) *gorm.DB {
); err != nil {
t.Fatalf("AutoMigrate() error = %v", err)
}
db.SetDB(conn)
t.Cleanup(func() { db.SetDB(nil) })
SetDBForTest(conn)
t.Cleanup(func() { SetDBForTest(nil) })
return conn
}
@@ -10,12 +10,11 @@ import (
"gorm.io/gorm"
"Wavelet/openflare/plugins/server/kernel/model"
db "Wavelet/plugins/infra/database"
)
// ListConfigVersionSummaries returns config version summaries ordered by created_at desc.
func ListConfigVersionSummaries(ctx context.Context) ([]*model.ConfigVersionSummary, error) {
conn := db.DB(ctx)
conn := DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
@@ -29,7 +28,7 @@ func ListConfigVersionSummaries(ctx context.Context) ([]*model.ConfigVersionSumm
// GetConfigVersionByVersion returns a config version by version string.
func GetConfigVersionByVersion(ctx context.Context, version string) (*model.ConfigVersion, error) {
conn := db.DB(ctx)
conn := DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
@@ -42,7 +41,7 @@ func GetConfigVersionByVersion(ctx context.Context, version string) (*model.Conf
// GetActiveConfigVersion returns the currently active config version.
func GetActiveConfigVersion(ctx context.Context) (*model.ConfigVersion, error) {
conn := db.DB(ctx)
conn := DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
@@ -55,7 +54,7 @@ func GetActiveConfigVersion(ctx context.Context) (*model.ConfigVersion, error) {
// GetLatestConfigVersionByPrefix returns the latest version string matching a date prefix.
func GetLatestConfigVersionByPrefix(ctx context.Context, prefix string) (string, error) {
conn := db.DB(ctx)
conn := DB(ctx)
if conn == nil {
return "", errors.New(errDatabaseNotInitialized)
}
@@ -73,7 +72,7 @@ func GetLatestConfigVersionByPrefix(ctx context.Context, prefix string) (string,
// CreateConfigVersion inserts a new config version record.
func CreateConfigVersion(ctx context.Context, version *model.ConfigVersion) error {
conn := db.DB(ctx)
conn := DB(ctx)
if conn == nil {
return errors.New(errDatabaseNotInitialized)
}
@@ -82,7 +81,7 @@ func CreateConfigVersion(ctx context.Context, version *model.ConfigVersion) erro
// PublishConfigVersionTx deactivates all versions and creates a new active version.
func PublishConfigVersionTx(ctx context.Context, version *model.ConfigVersion) error {
conn := db.DB(ctx)
conn := DB(ctx)
if conn == nil {
return errors.New(errDatabaseNotInitialized)
}
@@ -96,7 +95,7 @@ func PublishConfigVersionTx(ctx context.Context, version *model.ConfigVersion) e
// ActivateConfigVersionTx marks the given version active and deactivates others.
func ActivateConfigVersionTx(ctx context.Context, version string) error {
conn := db.DB(ctx)
conn := DB(ctx)
if conn == nil {
return errors.New(errDatabaseNotInitialized)
}
@@ -113,7 +112,7 @@ func DeleteConfigVersionsByVersions(ctx context.Context, versions []string) (int
if len(versions) == 0 {
return 0, nil
}
conn := db.DB(ctx)
conn := DB(ctx)
if conn == nil {
return 0, errors.New(errDatabaseNotInitialized)
}
@@ -123,7 +122,7 @@ func DeleteConfigVersionsByVersions(ctx context.Context, versions []string) (int
// ListEnabledProxyRoutes returns enabled proxy routes ordered by id asc.
func ListEnabledProxyRoutes(ctx context.Context) ([]*model.ProxyRoute, error) {
conn := db.DB(ctx)
conn := DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
@@ -8,12 +8,11 @@ import (
"errors"
"Wavelet/openflare/plugins/server/kernel/model"
db "Wavelet/plugins/infra/database"
)
// ListDNSAccounts 列出全部 DNS 账号(授权信息不通过 JSON 暴露)。
func ListDNSAccounts(ctx context.Context) ([]model.DNSAccount, error) {
conn := db.DB(ctx)
conn := DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
@@ -26,7 +25,7 @@ func ListDNSAccounts(ctx context.Context) ([]model.DNSAccount, error) {
// GetDNSAccountByID 按 ID 查询 DNS 账号。
func GetDNSAccountByID(ctx context.Context, id uint) (*model.DNSAccount, error) {
conn := db.DB(ctx)
conn := DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
@@ -39,7 +38,7 @@ func GetDNSAccountByID(ctx context.Context, id uint) (*model.DNSAccount, error)
// CreateDNSAccountRecord 创建 DNS 账号。
func CreateDNSAccountRecord(ctx context.Context, account *model.DNSAccount) error {
conn := db.DB(ctx)
conn := DB(ctx)
if conn == nil {
return errors.New(errDatabaseNotInitialized)
}
@@ -48,7 +47,7 @@ func CreateDNSAccountRecord(ctx context.Context, account *model.DNSAccount) erro
// SaveDNSAccount 保存 DNS 账号。
func SaveDNSAccount(ctx context.Context, account *model.DNSAccount) error {
conn := db.DB(ctx)
conn := DB(ctx)
if conn == nil {
return errors.New(errDatabaseNotInitialized)
}
@@ -57,7 +56,7 @@ func SaveDNSAccount(ctx context.Context, account *model.DNSAccount) error {
// DeleteDNSAccountRecord 删除 DNS 账号。
func DeleteDNSAccountRecord(ctx context.Context, id uint) error {
conn := db.DB(ctx)
conn := DB(ctx)
if conn == nil {
return errors.New(errDatabaseNotInitialized)
}
@@ -11,7 +11,6 @@ import (
"gorm.io/gorm"
"Wavelet/openflare/plugins/server/kernel/model"
db "Wavelet/plugins/infra/database"
)
const (
@@ -21,7 +20,7 @@ const (
// ListOpenFlareNodes returns all nodes ordered by id desc.
func ListOpenFlareNodes(ctx context.Context) ([]model.OpenFlareNode, error) {
conn := db.DB(ctx)
conn := DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
@@ -37,7 +36,7 @@ func ListOpenFlareNodesByNodeIDs(ctx context.Context, nodeIDs []string) ([]model
if len(nodeIDs) == 0 {
return []model.OpenFlareNode{}, nil
}
conn := db.DB(ctx)
conn := DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
@@ -50,7 +49,7 @@ func ListOpenFlareNodesByNodeIDs(ctx context.Context, nodeIDs []string) ([]model
// GetOpenFlareNodeByID returns a node by primary key.
func GetOpenFlareNodeByID(ctx context.Context, id uint) (*model.OpenFlareNode, error) {
conn := db.DB(ctx)
conn := DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
@@ -63,7 +62,7 @@ func GetOpenFlareNodeByID(ctx context.Context, id uint) (*model.OpenFlareNode, e
// GetOpenFlareNodeByNodeID returns a node by node_id.
func GetOpenFlareNodeByNodeID(ctx context.Context, nodeID string) (*model.OpenFlareNode, error) {
conn := db.DB(ctx)
conn := DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
@@ -76,7 +75,7 @@ func GetOpenFlareNodeByNodeID(ctx context.Context, nodeID string) (*model.OpenFl
// GetOpenFlareNodeByAccessToken returns a node by access token.
func GetOpenFlareNodeByAccessToken(ctx context.Context, token string) (*model.OpenFlareNode, error) {
conn := db.DB(ctx)
conn := DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
@@ -89,7 +88,7 @@ func GetOpenFlareNodeByAccessToken(ctx context.Context, token string) (*model.Op
// CreateOpenFlareNode inserts a new node.
func CreateOpenFlareNode(ctx context.Context, node *model.OpenFlareNode) error {
conn := db.DB(ctx)
conn := DB(ctx)
if conn == nil {
return errors.New(errDatabaseNotInitialized)
}
@@ -98,7 +97,7 @@ func CreateOpenFlareNode(ctx context.Context, node *model.OpenFlareNode) error {
// SaveOpenFlareNode persists node changes.
func SaveOpenFlareNode(ctx context.Context, node *model.OpenFlareNode) error {
conn := db.DB(ctx)
conn := DB(ctx)
if conn == nil {
return errors.New(errDatabaseNotInitialized)
}
@@ -107,7 +106,7 @@ func SaveOpenFlareNode(ctx context.Context, node *model.OpenFlareNode) error {
// UpdateOpenFlareNodeFields updates selected columns for a node.
func UpdateOpenFlareNodeFields(ctx context.Context, node *model.OpenFlareNode, fields ...string) error {
conn := db.DB(ctx)
conn := DB(ctx)
if conn == nil {
return errors.New(errDatabaseNotInitialized)
}
@@ -123,7 +122,7 @@ func UpdateOpenFlareNodeColumns(ctx context.Context, node *model.OpenFlareNode,
if node == nil || len(changes) == 0 {
return nil
}
conn := db.DB(ctx)
conn := DB(ctx)
if conn == nil {
return errors.New(errDatabaseNotInitialized)
}
@@ -133,7 +132,7 @@ func UpdateOpenFlareNodeColumns(ctx context.Context, node *model.OpenFlareNode,
// UpdateOpenFlareNodeFromApplyResult updates node status, last_seen, version and last_error after an apply report.
// When applyResult is "success", current_version is set and last_error is cleared; otherwise last_error is set to message.
func UpdateOpenFlareNodeFromApplyResult(ctx context.Context, nodeID, applyResult, version, message string, now time.Time) error {
conn := db.DB(ctx)
conn := DB(ctx)
if conn == nil {
return errors.New(errDatabaseNotInitialized)
}
@@ -159,7 +158,7 @@ func updateOpenFlareNodeFromApplyResultTx(tx *gorm.DB, nodeID, applyResult, vers
// DeleteOpenFlareNode removes a node by primary key.
func DeleteOpenFlareNode(ctx context.Context, id uint) error {
conn := db.DB(ctx)
conn := DB(ctx)
if conn == nil {
return errors.New(errDatabaseNotInitialized)
}
@@ -18,7 +18,6 @@ import (
analyticsrepo "Wavelet/openflare/plugins/server/kernel/repository/analytics"
"Wavelet/openflare/plugins/server/kernel/repository/logstore"
"Wavelet/pkg/logger"
db "Wavelet/plugins/infra/database"
)
const (
@@ -247,7 +246,7 @@ func ListOpenFlareMetricHourlySince(ctx context.Context, nodeID string, since ti
// ListOpenFlareActiveHealthEvents returns active health events across all nodes.
func ListOpenFlareActiveHealthEvents(ctx context.Context) ([]*model.OpenFlareHealthEvent, error) {
conn := db.DB(ctx)
conn := DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
@@ -263,7 +262,7 @@ func ListOpenFlareActiveHealthEvents(ctx context.Context) ([]*model.OpenFlareHea
// ListOpenFlareHealthEvents returns health events for a node.
func ListOpenFlareHealthEvents(ctx context.Context, nodeID string, activeOnly bool, limit int) ([]*model.OpenFlareHealthEvent, error) {
conn := db.DB(ctx)
conn := DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
@@ -358,7 +357,7 @@ func DeleteAllOpenFlareNodeObservationFrpc(ctx context.Context) (int64, error) {
// DeleteOpenFlareHealthEventsByNodeID deletes all health events for a node.
func DeleteOpenFlareHealthEventsByNodeID(ctx context.Context, nodeID string) (int64, error) {
conn := db.DB(ctx)
conn := DB(ctx)
if conn == nil {
return 0, errors.New(errDatabaseNotInitialized)
}
@@ -374,7 +373,7 @@ func DeleteOpenFlareHealthEventsByNodeID(ctx context.Context, nodeID string) (in
// GetOpenFlareNodeSystemProfile returns the system profile for a node.
func GetOpenFlareNodeSystemProfile(ctx context.Context, nodeID string) (*model.OpenFlareNodeSystemProfile, error) {
conn := db.DB(ctx)
conn := DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
@@ -393,7 +392,7 @@ func UpsertOpenFlareNodeSystemProfile(ctx context.Context, record *model.OpenFla
if record == nil {
return nil
}
conn := db.DB(ctx)
conn := DB(ctx)
if conn == nil {
return errors.New(errDatabaseNotInitialized)
}
@@ -431,7 +430,7 @@ func ReconcileOpenFlareHealthEvents(
reportedAt time.Time,
managedEventTypes map[string]struct{},
) error {
conn := db.DB(ctx)
conn := DB(ctx)
if conn == nil {
return errors.New(errDatabaseNotInitialized)
}
@@ -455,7 +454,7 @@ func PersistOpenFlareNodePGObservability(
if profile == nil && !reconcileHealth {
return nil
}
conn := db.DB(ctx)
conn := DB(ctx)
if conn == nil {
return errors.New(errDatabaseNotInitialized)
}
@@ -9,23 +9,22 @@ import (
"gorm.io/gorm"
"Wavelet/openflare/plugins/server/kernel/model"
db "Wavelet/plugins/infra/database"
)
// WithOriginTx runs fn inside a database transaction for origin multi-step work.
func WithOriginTx(ctx context.Context, fn func(tx *gorm.DB) error) error {
return db.DB(ctx).Transaction(fn)
return DB(ctx).Transaction(fn)
}
// HasProxyRoutesTable 判断代理规则表是否已迁移。
func HasProxyRoutesTable(ctx context.Context) bool {
return db.DB(ctx).Migrator().HasTable(&model.OriginProxyRoute{})
return DB(ctx).Migrator().HasTable(&model.OriginProxyRoute{})
}
// ListOrigins 列出全部源站。
func ListOrigins(ctx context.Context) ([]model.Origin, error) {
var origins []model.Origin
if err := db.DB(ctx).Order("id desc").Find(&origins).Error; err != nil {
if err := DB(ctx).Order("id desc").Find(&origins).Error; err != nil {
return nil, err
}
return origins, nil
@@ -34,7 +33,7 @@ func ListOrigins(ctx context.Context) ([]model.Origin, error) {
// GetOriginByID 按 ID 查询源站。
func GetOriginByID(ctx context.Context, id uint) (*model.Origin, error) {
var origin model.Origin
if err := db.DB(ctx).First(&origin, id).Error; err != nil {
if err := DB(ctx).First(&origin, id).Error; err != nil {
return nil, err
}
return &origin, nil
@@ -43,7 +42,7 @@ func GetOriginByID(ctx context.Context, id uint) (*model.Origin, error) {
// GetOriginByAddress 按地址查询源站。
func GetOriginByAddress(ctx context.Context, address string) (*model.Origin, error) {
var origin model.Origin
if err := db.DB(ctx).Where("address = ?", address).First(&origin).Error; err != nil {
if err := DB(ctx).Where("address = ?", address).First(&origin).Error; err != nil {
return nil, err
}
return &origin, nil
@@ -51,12 +50,12 @@ func GetOriginByAddress(ctx context.Context, address string) (*model.Origin, err
// CreateOriginRecord 创建源站。
func CreateOriginRecord(ctx context.Context, origin *model.Origin) error {
return db.DB(ctx).Create(origin).Error
return DB(ctx).Create(origin).Error
}
// SaveOrigin 保存源站。
func SaveOrigin(ctx context.Context, origin *model.Origin) error {
return SaveOriginTx(db.DB(ctx), origin)
return SaveOriginTx(DB(ctx), origin)
}
// SaveOriginTx saves an origin within an existing transaction.
@@ -66,7 +65,7 @@ func SaveOriginTx(tx *gorm.DB, origin *model.Origin) error {
// DeleteOriginRecord 删除源站。
func DeleteOriginRecord(ctx context.Context, id uint) error {
return db.DB(ctx).Delete(&model.Origin{}, id).Error
return DB(ctx).Delete(&model.Origin{}, id).Error
}
// ListOriginRouteCounts 统计各源站关联的代理规则数量。
@@ -75,7 +74,7 @@ func ListOriginRouteCounts(ctx context.Context) ([]model.OriginRouteCount, error
return nil, nil
}
result := make([]model.OriginRouteCount, 0)
err := db.DB(ctx).Model(&model.OriginProxyRoute{}).
err := DB(ctx).Model(&model.OriginProxyRoute{}).
Select("origin_id, COUNT(*) AS route_count").
Where("origin_id IS NOT NULL").
Group("origin_id").
@@ -89,7 +88,7 @@ func ListProxyRoutesByOriginID(ctx context.Context, originID uint) ([]model.Orig
return nil, nil
}
var routes []model.OriginProxyRoute
if err := db.DB(ctx).Where("origin_id = ?", originID).Order("id desc").Find(&routes).Error; err != nil {
if err := DB(ctx).Where("origin_id = ?", originID).Order("id desc").Find(&routes).Error; err != nil {
return nil, err
}
return routes, nil
@@ -120,7 +119,7 @@ func CountProxyRoutesByOriginID(ctx context.Context, originID uint) (int64, erro
return 0, nil
}
var count int64
if err := db.DB(ctx).Model(&model.OriginProxyRoute{}).Where("origin_id = ?", originID).Count(&count).Error; err != nil {
if err := DB(ctx).Model(&model.OriginProxyRoute{}).Where("origin_id = ?", originID).Count(&count).Error; err != nil {
return 0, err
}
return count, nil
@@ -7,18 +7,17 @@ import (
"context"
"Wavelet/openflare/plugins/server/kernel/model"
db "Wavelet/plugins/infra/database"
)
// HasPagesProjectsTable 判断 Pages 项目表是否已迁移。
func HasPagesProjectsTable(ctx context.Context) bool {
return db.DB(ctx).Migrator().HasTable(&model.PagesProject{})
return DB(ctx).Migrator().HasTable(&model.PagesProject{})
}
// ListPagesProjects 列出全部 Pages 项目。
func ListPagesProjects(ctx context.Context) ([]model.PagesProject, error) {
var projects []model.PagesProject
if err := db.DB(ctx).Order("id desc").Find(&projects).Error; err != nil {
if err := DB(ctx).Order("id desc").Find(&projects).Error; err != nil {
return nil, err
}
return projects, nil
@@ -27,7 +26,7 @@ func ListPagesProjects(ctx context.Context) ([]model.PagesProject, error) {
// GetPagesProjectByID 按 ID 查询 Pages 项目。
func GetPagesProjectByID(ctx context.Context, id uint) (*model.PagesProject, error) {
var project model.PagesProject
if err := db.DB(ctx).First(&project, id).Error; err != nil {
if err := DB(ctx).First(&project, id).Error; err != nil {
return nil, err
}
return &project, nil
@@ -36,7 +35,7 @@ func GetPagesProjectByID(ctx context.Context, id uint) (*model.PagesProject, err
// GetPagesProjectBySlug 按 slug 查询 Pages 项目。
func GetPagesProjectBySlug(ctx context.Context, slug string) (*model.PagesProject, error) {
var project model.PagesProject
if err := db.DB(ctx).Where("slug = ?", slug).First(&project).Error; err != nil {
if err := DB(ctx).Where("slug = ?", slug).First(&project).Error; err != nil {
return nil, err
}
return &project, nil
@@ -44,13 +43,13 @@ func GetPagesProjectBySlug(ctx context.Context, slug string) (*model.PagesProjec
// CreatePagesProjectRecord 创建 Pages 项目。
func CreatePagesProjectRecord(ctx context.Context, project *model.PagesProject) error {
return db.DB(ctx).Create(project).Error
return DB(ctx).Create(project).Error
}
// ListPagesDeployments 列出项目的全部部署。
func ListPagesDeployments(ctx context.Context, projectID uint) ([]model.PagesDeployment, error) {
var deployments []model.PagesDeployment
if err := db.DB(ctx).Where("project_id = ?", projectID).Order("id desc").Find(&deployments).Error; err != nil {
if err := DB(ctx).Where("project_id = ?", projectID).Order("id desc").Find(&deployments).Error; err != nil {
return nil, err
}
return deployments, nil
@@ -59,7 +58,7 @@ func ListPagesDeployments(ctx context.Context, projectID uint) ([]model.PagesDep
// GetPagesDeploymentByID 按 ID 查询 Pages 部署。
func GetPagesDeploymentByID(ctx context.Context, id uint) (*model.PagesDeployment, error) {
var deployment model.PagesDeployment
if err := db.DB(ctx).First(&deployment, id).Error; err != nil {
if err := DB(ctx).First(&deployment, id).Error; err != nil {
return nil, err
}
return &deployment, nil
@@ -68,7 +67,7 @@ func GetPagesDeploymentByID(ctx context.Context, id uint) (*model.PagesDeploymen
// ListPagesDeploymentFiles 列出部署文件清单。
func ListPagesDeploymentFiles(ctx context.Context, deploymentID uint) ([]model.PagesDeploymentFile, error) {
var files []model.PagesDeploymentFile
if err := db.DB(ctx).Where("deployment_id = ?", deploymentID).Order("path asc").Find(&files).Error; err != nil {
if err := DB(ctx).Where("deployment_id = ?", deploymentID).Order("path asc").Find(&files).Error; err != nil {
return nil, err
}
return files, nil
@@ -77,7 +76,7 @@ func ListPagesDeploymentFiles(ctx context.Context, deploymentID uint) ([]model.P
// CountPagesDeploymentsByProjectID 统计项目部署数量。
func CountPagesDeploymentsByProjectID(ctx context.Context, projectID uint) (int64, error) {
var count int64
if err := db.DB(ctx).Model(&model.PagesDeployment{}).Where("project_id = ?", projectID).Count(&count).Error; err != nil {
if err := DB(ctx).Model(&model.PagesDeployment{}).Where("project_id = ?", projectID).Count(&count).Error; err != nil {
return 0, err
}
return count, nil
@@ -89,7 +88,7 @@ func CountProxyRoutesByPagesProjectID(ctx context.Context, projectID uint) (int6
return 0, nil
}
var count int64
if err := db.DB(ctx).Model(&model.ProxyRoute{}).Where("pages_project_id = ?", projectID).Count(&count).Error; err != nil {
if err := DB(ctx).Model(&model.ProxyRoute{}).Where("pages_project_id = ?", projectID).Count(&count).Error; err != nil {
return 0, err
}
return count, nil
@@ -8,7 +8,6 @@ import (
"errors"
"Wavelet/openflare/plugins/server/kernel/model"
db "Wavelet/plugins/infra/database"
)
// ListPagesOrphanUploadCandidates returns at most 100 unreferenced, isolated
@@ -21,17 +20,20 @@ func ListPagesOrphanUploadCandidates(
if input.SystemUserID == 0 || input.UploadType == "" || input.Marker == "" || input.CreatedBefore.IsZero() {
return nil, errors.New("invalid pages orphan upload candidate query")
}
markerPredicate, err := pagesOrphanMarkerPredicate(db.DB(ctx).Name())
markerPredicate, err := pagesOrphanMarkerPredicate(DB(ctx).Name())
if err != nil {
return nil, err
}
deploymentTable := (model.PagesDeployment{}).TableName()
uploadTable := (model.Upload{}).TableName()
const (
uploadTable = "w_uploads"
uploadStatusUsed = "used"
)
var candidates []model.Upload
err = db.DB(ctx).
Model(&model.Upload{}).
Where(uploadTable+".status = ?", model.UploadStatusUsed).
err = DB(ctx).
Table(uploadTable).
Where(uploadTable+".status = ?", uploadStatusUsed).
Where(uploadTable+".user_id = ?", input.SystemUserID).
Where(uploadTable+".type = ?", input.UploadType).
Where(uploadTable+".created_at < ?", input.CreatedBefore).
@@ -11,8 +11,6 @@ import (
"Wavelet/openflare/plugins/server/kernel/model"
db "Wavelet/plugins/infra/database"
"github.com/glebarez/sqlite"
"gorm.io/gorm"
)
@@ -64,7 +62,7 @@ func TestListPagesOrphanUploadCandidatesFiltersAndLimits(t *testing.T) {
"pages_project_id": "1",
}}
valid := make([]model.Upload, 0, model.PagesOrphanUploadCandidateLimit+1)
valid := make([]testUploadEntity, 0, model.PagesOrphanUploadCandidateLimit+1)
for index := 0; index < model.PagesOrphanUploadCandidateLimit+1; index++ {
valid = append(valid, pagesCleanupModelUpload(uint64(index+100), 999, "openflare_pages_deployment", model.UploadStatusUsed, old, marker))
}
@@ -81,7 +79,7 @@ func TestListPagesOrphanUploadCandidatesFiltersAndLimits(t *testing.T) {
"pages_ingest_marker": "pages_deployment_v1",
"pages_project_id": "1",
}})
for _, upload := range []model.Upload{referenced, wrongOwner, wrongType, wrongStatus, fresh, wrongMarker} {
for _, upload := range []testUploadEntity{referenced, wrongOwner, wrongType, wrongStatus, fresh, wrongMarker} {
if err := gormDB.Create(&upload).Error; err != nil {
t.Fatalf("create filtered upload %d error = %v, want nil", upload.ID, err)
}
@@ -100,7 +98,7 @@ func TestListPagesOrphanUploadCandidatesFiltersAndLimits(t *testing.T) {
if err := gormDB.Create(&invalidJSON).Error; err != nil {
t.Fatalf("create invalid JSON upload error = %v, want nil", err)
}
if err := gormDB.Table((model.Upload{}).TableName()).Where("id = ?", invalidJSON.ID).
if err := gormDB.Table("w_uploads").Where("id = ?", invalidJSON.ID).
UpdateColumn("metadata", "{invalid").Error; err != nil {
t.Fatalf("corrupt upload metadata error = %v, want nil", err)
}
@@ -133,7 +131,7 @@ func TestListPagesOrphanUploadCandidatesSkipsInvalidSQLiteJSON(t *testing.T) {
if err := gormDB.Create(&upload).Error; err != nil {
t.Fatalf("create invalid JSON candidate error = %v, want nil", err)
}
if err := gormDB.Table((model.Upload{}).TableName()).Where("id = ?", upload.ID).
if err := gormDB.Table("w_uploads").Where("id = ?", upload.ID).
UpdateColumn("metadata", "{invalid").Error; err != nil {
t.Fatalf("corrupt upload metadata error = %v, want nil", err)
}
@@ -152,6 +150,25 @@ func TestListPagesOrphanUploadCandidatesSkipsInvalidSQLiteJSON(t *testing.T) {
}
}
type testUploadEntity struct {
ID uint64 `gorm:"primaryKey"`
UserID uint64 `gorm:"index"`
FileName string `gorm:"size:255"`
FilePath string `gorm:"size:500"`
Size int64
MimeType string `gorm:"size:100"`
Hash string `gorm:"size:64"`
Type string `gorm:"size:50;index"`
Status model.UploadStatus `gorm:"type:varchar(20)"`
Metadata model.UploadMetadata `gorm:"serializer:json;type:jsonb"`
CreatedAt time.Time
UpdatedAt time.Time
}
func (testUploadEntity) TableName() string {
return "w_uploads"
}
func setupPagesCleanupModelTestDB(t *testing.T) *gorm.DB {
t.Helper()
@@ -161,11 +178,11 @@ func setupPagesCleanupModelTestDB(t *testing.T) *gorm.DB {
if err != nil {
t.Fatalf("open Pages cleanup model test database error = %v, want nil", err)
}
if err := gormDB.AutoMigrate(&model.Upload{}, &model.PagesDeployment{}); err != nil {
if err := gormDB.AutoMigrate(&testUploadEntity{}, &model.PagesDeployment{}); err != nil {
t.Fatalf("migrate Pages cleanup model test database error = %v, want nil", err)
}
db.SetDB(gormDB)
t.Cleanup(func() { db.SetDB(nil) })
SetDBForTest(gormDB)
t.Cleanup(func() { SetDBForTest(nil) })
return gormDB
}
@@ -176,21 +193,19 @@ func pagesCleanupModelUpload(
status model.UploadStatus,
createdAt time.Time,
metadata model.UploadMetadata,
) model.Upload {
return model.Upload{
ID: id,
UserID: userID,
FileName: "site.zip",
FilePath: "pages/site.zip",
FileSize: 10,
MimeType: "application/zip",
Extension: "zip",
Hash: "checksum",
Type: uploadType,
Status: status,
AccessMode: 0,
Metadata: metadata,
CreatedAt: createdAt,
UpdatedAt: createdAt,
) testUploadEntity {
return testUploadEntity{
ID: id,
UserID: userID,
FileName: "site.zip",
FilePath: "pages/site.zip",
Size: 10,
MimeType: "application/zip",
Hash: "checksum",
Type: uploadType,
Status: status,
Metadata: metadata,
CreatedAt: createdAt,
UpdatedAt: createdAt,
}
}
@@ -11,20 +11,19 @@ import (
"gorm.io/gorm/clause"
"Wavelet/openflare/plugins/server/kernel/model"
db "Wavelet/plugins/infra/database"
)
const pagesRowLockStrength = "UPDATE"
// WithPagesTx runs fn inside a database transaction for Pages multi-step work.
func WithPagesTx(ctx context.Context, fn func(tx *gorm.DB) error) error {
return db.DB(ctx).Transaction(fn)
return DB(ctx).Transaction(fn)
}
// GetPagesProjectSourceByID loads a project source by primary key.
func GetPagesProjectSourceByID(ctx context.Context, id uint) (*model.PagesProjectSource, error) {
var source model.PagesProjectSource
if err := db.DB(ctx).Where("id = ?", id).First(&source).Error; err != nil {
if err := DB(ctx).Where("id = ?", id).First(&source).Error; err != nil {
return nil, err
}
return &source, nil
@@ -33,7 +32,7 @@ func GetPagesProjectSourceByID(ctx context.Context, id uint) (*model.PagesProjec
// GetPagesProjectSourceByProjectID loads the unique source for a project.
func GetPagesProjectSourceByProjectID(ctx context.Context, projectID uint) (*model.PagesProjectSource, error) {
var source model.PagesProjectSource
if err := db.DB(ctx).Where("project_id = ?", projectID).First(&source).Error; err != nil {
if err := DB(ctx).Where("project_id = ?", projectID).First(&source).Error; err != nil {
return nil, err
}
return &source, nil
@@ -46,7 +45,7 @@ func GetPagesProjectSourceByIDAndConfigVersion(
configVersion int,
) (*model.PagesProjectSource, error) {
var source model.PagesProjectSource
if err := db.DB(ctx).Where("id = ? AND config_version = ?", id, configVersion).First(&source).Error; err != nil {
if err := DB(ctx).Where("id = ? AND config_version = ?", id, configVersion).First(&source).Error; err != nil {
return nil, err
}
return &source, nil
@@ -58,7 +57,7 @@ func GetPagesProjectSourceRuntimeBySourceID(
sourceID uint,
) (*model.PagesProjectSourceRuntime, error) {
var runtime model.PagesProjectSourceRuntime
if err := db.DB(ctx).Where("source_id = ?", sourceID).First(&runtime).Error; err != nil {
if err := DB(ctx).Where("source_id = ?", sourceID).First(&runtime).Error; err != nil {
return nil, err
}
return &runtime, nil
@@ -207,7 +206,7 @@ func TryAcquirePagesSourceRuntimeLease(
now time.Time,
updates map[string]any,
) (int64, error) {
result := db.DB(ctx).Model(&model.PagesProjectSourceRuntime{}).
result := DB(ctx).Model(&model.PagesProjectSourceRuntime{}).
Where("source_id = ?", sourceID).
Where("lease_expires_at IS NULL OR lease_expires_at <= ?", now).
Where(
@@ -227,7 +226,7 @@ func RenewPagesSourceRuntimeLease(
now time.Time,
expiresAt time.Time,
) (int64, error) {
result := db.DB(ctx).Model(&model.PagesProjectSourceRuntime{}).
result := DB(ctx).Model(&model.PagesProjectSourceRuntime{}).
Where("source_id = ? AND lease_token = ? AND lease_expires_at > ?", sourceID, token, now).
Updates(map[string]any{"lease_expires_at": expiresAt})
return result.RowsAffected, result.Error
@@ -241,7 +240,7 @@ func UpdatePagesSourceRuntimeByActiveLease(
now time.Time,
updates map[string]any,
) (int64, error) {
return UpdatePagesSourceRuntimeByActiveLeaseTx(db.DB(ctx), sourceID, token, now, updates)
return UpdatePagesSourceRuntimeByActiveLeaseTx(DB(ctx), sourceID, token, now, updates)
}
// UpdatePagesSourceRuntimeByActiveLeaseTx updates runtime under an active lease inside a transaction.
@@ -268,7 +267,7 @@ func RecoverExpiredPagesSourceRuntimeLease(
now time.Time,
updates map[string]any,
) (int64, error) {
result := db.DB(ctx).Model(&model.PagesProjectSourceRuntime{}).
result := DB(ctx).Model(&model.PagesProjectSourceRuntime{}).
Where("source_id = ?", sourceID).
Where("lease_token = ?", token).
Where("lease_expires_at = ? AND lease_expires_at <= ?", expiresAt, now).
@@ -285,7 +284,7 @@ func MarkPagesSourceInitialCheckDispatchFailed(
now time.Time,
updates map[string]any,
) (int64, error) {
result := db.DB(ctx).Model(&model.PagesProjectSourceRuntime{}).
result := DB(ctx).Model(&model.PagesProjectSourceRuntime{}).
Where("source_id = ?", sourceID).
Where("lease_expires_at IS NULL OR lease_expires_at <= ?", now).
Where(
@@ -309,7 +308,7 @@ func RecordPagesSourceAutoDispatchFailure(
now time.Time,
updates map[string]any,
) (int64, error) {
result := db.DB(ctx).Model(&model.PagesProjectSourceRuntime{}).
result := DB(ctx).Model(&model.PagesProjectSourceRuntime{}).
Where("source_id = ?", sourceID).
Where("sync_status = ? AND last_seen_revision = ?", updateAvailableStatus, revision).
Where("lease_expires_at IS NULL OR lease_expires_at <= ?", now).
@@ -336,7 +335,7 @@ func ListExpiredPagesSourceLeaseCandidates(
syncStatuses []string,
) ([]model.PagesExpiredSourceLeaseCandidate, error) {
var candidates []model.PagesExpiredSourceLeaseCandidate
err := db.DB(ctx).
err := DB(ctx).
Table("of_pages_project_source_runtime AS runtime").
Select(`runtime.source_id, runtime.lease_token, runtime.lease_expires_at,
runtime.sync_status, source.source_type, source.release_selector`).
@@ -391,7 +390,7 @@ func dueGitHubPagesSourceQuery(
sourceType string,
releaseSelector string,
) *gorm.DB {
return db.DB(ctx).
return DB(ctx).
Table("of_pages_project_source_runtime AS runtime").
Joins("JOIN of_pages_project_sources AS source ON source.id = runtime.source_id").
Where("source.source_type = ?", sourceType).
@@ -407,7 +406,7 @@ func GetPagesDeploymentBySourceRevision(
revision string,
) (*model.PagesDeployment, error) {
var deployment model.PagesDeployment
err := db.DB(ctx).
err := DB(ctx).
Where("project_id = ? AND source_identity = ? AND source_revision = ?", projectID, sourceIdentity, revision).
First(&deployment).Error
if err != nil {
@@ -11,7 +11,6 @@ import (
"gorm.io/gorm/clause"
"Wavelet/openflare/plugins/server/kernel/model"
db "Wavelet/plugins/infra/database"
)
// Zone domain binding sentinel errors (shared by proxy_route and zone binding helpers).
@@ -24,13 +23,13 @@ var (
// WithProxyRouteTx runs fn inside a database transaction for proxy-route multi-step work.
func WithProxyRouteTx(ctx context.Context, fn func(tx *gorm.DB) error) error {
return db.DB(ctx).Transaction(fn)
return DB(ctx).Transaction(fn)
}
// ListProxyRoutes 列出全部代理规则。
func ListProxyRoutes(ctx context.Context) ([]*model.ProxyRoute, error) {
var routes []*model.ProxyRoute
if err := db.DB(ctx).Order("id desc").Find(&routes).Error; err != nil {
if err := DB(ctx).Order("id desc").Find(&routes).Error; err != nil {
return nil, err
}
return routes, nil
@@ -39,7 +38,7 @@ func ListProxyRoutes(ctx context.Context) ([]*model.ProxyRoute, error) {
// GetProxyRouteByID 按 ID 查询代理规则。
func GetProxyRouteByID(ctx context.Context, id uint) (*model.ProxyRoute, error) {
var route model.ProxyRoute
if err := db.DB(ctx).First(&route, id).Error; err != nil {
if err := DB(ctx).First(&route, id).Error; err != nil {
return nil, err
}
return &route, nil
@@ -47,7 +46,7 @@ func GetProxyRouteByID(ctx context.Context, id uint) (*model.ProxyRoute, error)
// CreateProxyRouteRecord 创建代理规则。
func CreateProxyRouteRecord(ctx context.Context, route *model.ProxyRoute) error {
return CreateProxyRouteRecordTx(db.DB(ctx), route)
return CreateProxyRouteRecordTx(DB(ctx), route)
}
// CreateProxyRouteRecordTx creates a proxy route within an existing transaction.
@@ -57,7 +56,7 @@ func CreateProxyRouteRecordTx(tx *gorm.DB, route *model.ProxyRoute) error {
// UpdateProxyRouteRecord 更新代理规则。
func UpdateProxyRouteRecord(ctx context.Context, route *model.ProxyRoute) error {
return UpdateProxyRouteRecordTx(db.DB(ctx), route)
return UpdateProxyRouteRecordTx(DB(ctx), route)
}
// UpdateProxyRouteRecordTx updates a proxy route within an existing transaction.
@@ -96,7 +95,7 @@ func proxyRouteUpdateMap(route *model.ProxyRoute) map[string]any {
// DeleteProxyRouteRecord 删除代理规则。
func DeleteProxyRouteRecord(ctx context.Context, id uint) error {
return DeleteProxyRouteRecordTx(db.DB(ctx), id)
return DeleteProxyRouteRecordTx(DB(ctx), id)
}
// DeleteProxyRouteRecordTx deletes a proxy route within an existing transaction.
@@ -8,17 +8,16 @@ import (
"errors"
"Wavelet/openflare/plugins/server/kernel/model"
db "Wavelet/plugins/infra/database"
)
// HasTLSProxyRoutesTable 判断代理规则表是否已迁移。
func HasTLSProxyRoutesTable(ctx context.Context) bool {
return db.DB(ctx).Migrator().HasTable(&model.TLSProxyRouteRef{})
return DB(ctx).Migrator().HasTable(&model.TLSProxyRouteRef{})
}
// ListTLSCertificates 列出全部证书(不含 PEM 敏感字段的 JSON 暴露由 struct tag 控制)。
func ListTLSCertificates(ctx context.Context) ([]model.TLSCertificate, error) {
conn := db.DB(ctx)
conn := DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
@@ -31,7 +30,7 @@ func ListTLSCertificates(ctx context.Context) ([]model.TLSCertificate, error) {
// GetTLSCertificateByID 按 ID 查询证书。
func GetTLSCertificateByID(ctx context.Context, id uint) (*model.TLSCertificate, error) {
conn := db.DB(ctx)
conn := DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
@@ -44,7 +43,7 @@ func GetTLSCertificateByID(ctx context.Context, id uint) (*model.TLSCertificate,
// CreateTLSCertificateRecord 创建证书记录。
func CreateTLSCertificateRecord(ctx context.Context, certificate *model.TLSCertificate) error {
conn := db.DB(ctx)
conn := DB(ctx)
if conn == nil {
return errors.New(errDatabaseNotInitialized)
}
@@ -53,7 +52,7 @@ func CreateTLSCertificateRecord(ctx context.Context, certificate *model.TLSCerti
// SaveTLSCertificate 保存证书记录。
func SaveTLSCertificate(ctx context.Context, certificate *model.TLSCertificate) error {
conn := db.DB(ctx)
conn := DB(ctx)
if conn == nil {
return errors.New(errDatabaseNotInitialized)
}
@@ -62,7 +61,7 @@ func SaveTLSCertificate(ctx context.Context, certificate *model.TLSCertificate)
// DeleteTLSCertificateRecord 删除证书记录。
func DeleteTLSCertificateRecord(ctx context.Context, id uint) error {
conn := db.DB(ctx)
conn := DB(ctx)
if conn == nil {
return errors.New(errDatabaseNotInitialized)
}
@@ -71,7 +70,7 @@ func DeleteTLSCertificateRecord(ctx context.Context, id uint) error {
// CountTLSCertificatesByDNSAccountID 统计引用指定 DNS 账号的证书数量。
func CountTLSCertificatesByDNSAccountID(ctx context.Context, dnsAccountID uint) (int64, error) {
conn := db.DB(ctx)
conn := DB(ctx)
if conn == nil {
return 0, errors.New(errDatabaseNotInitialized)
}
@@ -88,7 +87,7 @@ func ListTLSProxyRouteRefs(ctx context.Context) ([]model.TLSProxyRouteRef, error
return nil, nil
}
var routes []model.TLSProxyRouteRef
if err := db.DB(ctx).Order("id asc").Find(&routes).Error; err != nil {
if err := DB(ctx).Order("id asc").Find(&routes).Error; err != nil {
return nil, err
}
return routes, nil
@@ -11,11 +11,10 @@ import (
"gorm.io/gorm"
"Wavelet/openflare/plugins/server/kernel/model"
db "Wavelet/plugins/infra/database"
)
func wafDB(ctx context.Context) (*gorm.DB, error) {
conn := db.DB(ctx)
conn := DB(ctx)
if conn == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
@@ -9,8 +9,6 @@ import (
"Wavelet/openflare/plugins/server/kernel/model"
db "Wavelet/plugins/infra/database"
"github.com/glebarez/sqlite"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
@@ -26,9 +24,9 @@ func setupWAFBindingsTestDB(t *testing.T) func() {
require.NoError(t, err)
require.NoError(t, sqliteDB.AutoMigrate(&model.OpenFlareWAFRuleGroupBinding{}))
db.SetDB(sqliteDB)
SetDBForTest(sqliteDB)
return func() {
db.SetDB(nil)
SetDBForTest(nil)
}
}
@@ -37,7 +35,7 @@ func TestReplaceOpenFlareWAFRuleGroupBindingsAfterExplicitHighID(t *testing.T) {
defer cleanup()
ctx := context.Background()
conn := db.DB(ctx)
conn := DB(ctx)
require.NotNil(t, conn)
require.NoError(t, conn.Create(&model.OpenFlareWAFRuleGroupBinding{
ID: 50,
@@ -9,8 +9,6 @@ import (
"Wavelet/openflare/plugins/server/kernel/model"
db "Wavelet/plugins/infra/database"
"github.com/glebarez/sqlite"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
@@ -23,8 +21,8 @@ func TestOpenFlareWAFGraphOptimisticUpdate(t *testing.T) {
conn, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
require.NoError(t, err)
require.NoError(t, conn.AutoMigrate(&model.OpenFlareWAFRuleGroup{}))
db.SetDB(conn)
t.Cleanup(func() { db.SetDB(nil) })
SetDBForTest(conn)
t.Cleanup(func() { SetDBForTest(nil) })
group := model.OpenFlareWAFRuleGroup{Name: "rule", Graph: defaultWAFRuleGraph, Revision: 1}
require.NoError(t, conn.Create(&group).Error)
@@ -10,13 +10,12 @@ import (
"gorm.io/gorm"
"Wavelet/openflare/plugins/server/kernel/model"
db "Wavelet/plugins/infra/database"
)
// ListZones returns all zones ordered by domain ascending.
func ListZones(ctx context.Context) ([]model.Zone, error) {
var zones []model.Zone
if err := db.DB(ctx).Order("domain asc").Find(&zones).Error; err != nil {
if err := DB(ctx).Order("domain asc").Find(&zones).Error; err != nil {
return nil, err
}
return zones, nil
@@ -25,7 +24,7 @@ func ListZones(ctx context.Context) ([]model.Zone, error) {
// GetZoneByID returns a zone by primary key.
func GetZoneByID(ctx context.Context, id uint) (*model.Zone, error) {
var zone model.Zone
if err := db.DB(ctx).First(&zone, id).Error; err != nil {
if err := DB(ctx).First(&zone, id).Error; err != nil {
return nil, err
}
return &zone, nil
@@ -33,23 +32,23 @@ func GetZoneByID(ctx context.Context, id uint) (*model.Zone, error) {
// CreateZone creates a zone record.
func CreateZone(ctx context.Context, zone *model.Zone) error {
return db.DB(ctx).Create(zone).Error
return DB(ctx).Create(zone).Error
}
// SaveZone persists zone updates.
func SaveZone(ctx context.Context, zone *model.Zone) error {
return db.DB(ctx).Save(zone).Error
return DB(ctx).Save(zone).Error
}
// DeleteZone deletes a zone by primary key.
func DeleteZone(ctx context.Context, id uint) error {
return db.DB(ctx).Delete(&model.Zone{}, id).Error
return DB(ctx).Delete(&model.Zone{}, id).Error
}
// ListZoneDomainCounts returns per-zone domain counts for list cards.
func ListZoneDomainCounts(ctx context.Context) ([]model.ZoneDomainCount, error) {
var rows []model.ZoneDomainCount
if err := db.DB(ctx).Model(&model.ZoneDomain{}).
if err := DB(ctx).Model(&model.ZoneDomain{}).
Select("zone_id, count(*) as count").
Group("zone_id").
Scan(&rows).Error; err != nil {
@@ -61,7 +60,7 @@ func ListZoneDomainCounts(ctx context.Context) ([]model.ZoneDomainCount, error)
// ListZoneDomainsByZoneID returns domains under a zone ordered by domain ascending.
func ListZoneDomainsByZoneID(ctx context.Context, zoneID uint) ([]model.ZoneDomain, error) {
var domains []model.ZoneDomain
if err := db.DB(ctx).Where("zone_id = ?", zoneID).Order("domain asc").Find(&domains).Error; err != nil {
if err := DB(ctx).Where("zone_id = ?", zoneID).Order("domain asc").Find(&domains).Error; err != nil {
return nil, err
}
return domains, nil
@@ -70,7 +69,7 @@ func ListZoneDomainsByZoneID(ctx context.Context, zoneID uint) ([]model.ZoneDoma
// CountZoneDomainsByZoneID counts domains under a zone.
func CountZoneDomainsByZoneID(ctx context.Context, zoneID uint) (int64, error) {
var count int64
if err := db.DB(ctx).Model(&model.ZoneDomain{}).Where("zone_id = ?", zoneID).Count(&count).Error; err != nil {
if err := DB(ctx).Model(&model.ZoneDomain{}).Where("zone_id = ?", zoneID).Count(&count).Error; err != nil {
return 0, err
}
return count, nil
@@ -79,7 +78,7 @@ func CountZoneDomainsByZoneID(ctx context.Context, zoneID uint) (int64, error) {
// GetZoneDomainByZoneAndID returns a domain scoped to a zone.
func GetZoneDomainByZoneAndID(ctx context.Context, zoneID, id uint) (*model.ZoneDomain, error) {
var item model.ZoneDomain
if err := db.DB(ctx).Where("id = ? AND zone_id = ?", id, zoneID).First(&item).Error; err != nil {
if err := DB(ctx).Where("id = ? AND zone_id = ?", id, zoneID).First(&item).Error; err != nil {
return nil, err
}
return &item, nil
@@ -87,17 +86,17 @@ func GetZoneDomainByZoneAndID(ctx context.Context, zoneID, id uint) (*model.Zone
// CreateZoneDomain creates a zone domain record.
func CreateZoneDomain(ctx context.Context, domain *model.ZoneDomain) error {
return db.DB(ctx).Create(domain).Error
return DB(ctx).Create(domain).Error
}
// SaveZoneDomain persists zone domain updates.
func SaveZoneDomain(ctx context.Context, domain *model.ZoneDomain) error {
return db.DB(ctx).Save(domain).Error
return DB(ctx).Save(domain).Error
}
// DeleteZoneDomain deletes a zone domain record.
func DeleteZoneDomain(ctx context.Context, domain *model.ZoneDomain) error {
conn := db.DB(ctx)
conn := DB(ctx)
if conn == nil {
return errors.New(errDatabaseNotInitialized)
}
@@ -112,7 +111,7 @@ func DeleteZoneDomain(ctx context.Context, domain *model.ZoneDomain) error {
// ListZoneDomainsByRouteID returns the domains bound to a proxy route.
func ListZoneDomainsByRouteID(ctx context.Context, routeID uint) ([]model.ZoneDomain, error) {
var domains []model.ZoneDomain
if err := db.DB(ctx).Where("proxy_route_id = ?", routeID).Order("id asc").Find(&domains).Error; err != nil {
if err := DB(ctx).Where("proxy_route_id = ?", routeID).Order("id asc").Find(&domains).Error; err != nil {
return nil, err
}
return domains, nil
@@ -124,7 +123,7 @@ func ListZoneDomainsByIDs(ctx context.Context, domainIDs []uint) ([]model.ZoneDo
return []model.ZoneDomain{}, nil
}
var domains []model.ZoneDomain
if err := db.DB(ctx).Where("id IN ?", domainIDs).Find(&domains).Error; err != nil {
if err := DB(ctx).Where("id IN ?", domainIDs).Find(&domains).Error; err != nil {
return nil, err
}
byID := make(map[uint]model.ZoneDomain, len(domains))
@@ -145,13 +144,13 @@ func ListZoneDomainsByIDs(ctx context.Context, domainIDs []uint) ([]model.ZoneDo
// CountZoneDomainsByCertificateID reports whether a certificate is assigned to a model.Zone domain.
func CountZoneDomainsByCertificateID(ctx context.Context, certificateID uint) (int64, error) {
var count int64
err := db.DB(ctx).Model(&model.ZoneDomain{}).Where("cert_id = ?", certificateID).Count(&count).Error
err := DB(ctx).Model(&model.ZoneDomain{}).Where("cert_id = ?", certificateID).Count(&count).Error
return count, err
}
// ReplaceZoneDomainRouteBindings replaces every model.ZoneDomain binding for a proxy route.
func ReplaceZoneDomainRouteBindings(ctx context.Context, routeID uint, domainIDs []uint) error {
conn := db.DB(ctx)
conn := DB(ctx)
if conn == nil {
return errors.New(errDatabaseNotInitialized)
}
@@ -9,8 +9,6 @@ import (
"Wavelet/openflare/plugins/server/kernel/model"
db "Wavelet/plugins/infra/database"
"github.com/glebarez/sqlite"
"github.com/stretchr/testify/require"
"gorm.io/gorm"
@@ -24,8 +22,8 @@ func setupZoneTestDB(t *testing.T) *gorm.DB {
})
require.NoError(t, err)
require.NoError(t, sqliteDB.AutoMigrate(&model.Zone{}, &model.ZoneDomain{}))
db.SetDB(sqliteDB)
t.Cleanup(func() { db.SetDB(nil) })
SetDBForTest(sqliteDB)
t.Cleanup(func() { SetDBForTest(nil) })
return sqliteDB
}
@@ -6,112 +6,162 @@ package repository
import (
"context"
"errors"
"sync"
"Wavelet/core/contracts"
"Wavelet/openflare/plugins/server/kernel/model"
adminrepo "Wavelet/plugins/domain/admin/repository"
db "Wavelet/plugins/infra/database"
)
const configTypeSystem = "system"
var (
configMu sync.RWMutex
configSvc contracts.SystemConfigService
)
// ensureAdminStore points OF config access at Wavelet's admin repository so
// reads hit the same cache that SaveOrUpdateSystemConfig invalidates.
func ensureAdminStore(ctx context.Context) error {
if conn := db.DB(ctx); conn != nil {
adminrepo.SetDBService(db.NewService(conn))
}
if adminrepo.GetDB(ctx) == nil {
return errors.New(errDatabaseNotInitialized)
}
return nil
// SetSystemConfigService injects the platform SystemConfigService.
func SetSystemConfigService(s contracts.SystemConfigService) {
configMu.Lock()
defer configMu.Unlock()
configSvc = s
}
// GetSystemConfigByKey loads a config row by key through the admin store cache.
func currentConfigService() contracts.SystemConfigService {
configMu.RLock()
defer configMu.RUnlock()
return configSvc
}
func ensureConfigService() (contracts.SystemConfigService, error) {
svc := currentConfigService()
if svc == nil {
return nil, errors.New("system config service not initialized")
}
return svc, nil
}
// GetSystemConfigByKey loads a config row by key through the system config service.
func GetSystemConfigByKey(ctx context.Context, key string) (model.SystemConfig, error) {
if err := ensureAdminStore(ctx); err != nil {
svc, err := ensureConfigService()
if err != nil {
return model.SystemConfig{}, err
}
return adminrepo.GetSystemConfigByKey(ctx, key)
dto, err := svc.GetByKey(ctx, key)
if err != nil {
return model.SystemConfig{}, err
}
return model.FromSystemConfigDTO(dto), nil
}
// ListSystemConfigsByKeys loads multiple config keys through the admin store cache.
// ListSystemConfigsByKeys loads multiple config keys through the system config service.
func ListSystemConfigsByKeys(ctx context.Context, keys []string) (map[string]model.SystemConfig, error) {
if err := ensureAdminStore(ctx); err != nil {
svc, err := ensureConfigService()
if err != nil {
return nil, err
}
return adminrepo.ListSystemConfigsByKeys(ctx, keys)
dtos, err := svc.ListByKeys(ctx, keys)
if err != nil {
return nil, err
}
res := make(map[string]model.SystemConfig, len(dtos))
for k, v := range dtos {
res[k] = model.FromSystemConfigDTO(v)
}
return res, nil
}
// ListVisibleSystemConfigs returns visibility=1 configs from the admin store cache.
// ListVisibleSystemConfigs returns visibility=1 configs from the system config service.
func ListVisibleSystemConfigs(ctx context.Context) ([]model.SystemConfig, error) {
if err := ensureAdminStore(ctx); err != nil {
svc, err := ensureConfigService()
if err != nil {
return nil, err
}
return adminrepo.ListVisibleSystemConfigs(ctx)
dtos, err := svc.ListVisible(ctx)
if err != nil {
return nil, err
}
res := make([]model.SystemConfig, len(dtos))
for i, v := range dtos {
res[i] = model.FromSystemConfigDTO(v)
}
return res, nil
}
// GetIntByKey queries config and converts to int.
func GetIntByKey(ctx context.Context, key string) (int, error) {
if err := ensureAdminStore(ctx); err != nil {
svc, err := ensureConfigService()
if err != nil {
return 0, err
}
return adminrepo.GetIntByKey(ctx, key)
return svc.GetIntByKey(ctx, key)
}
// GetBoolByKey queries config and converts to bool.
func GetBoolByKey(ctx context.Context, key string) (bool, error) {
if err := ensureAdminStore(ctx); err != nil {
svc, err := ensureConfigService()
if err != nil {
return false, err
}
return adminrepo.GetBoolByKey(ctx, key)
return svc.GetBoolByKey(ctx, key)
}
// CreateSystemConfig persists a new system config row.
func CreateSystemConfig(ctx context.Context, config *model.SystemConfig) error {
if err := ensureAdminStore(ctx); err != nil {
svc, err := ensureConfigService()
if err != nil {
return err
}
return adminrepo.CreateSystemConfigRecord(ctx, config)
return svc.SaveOrUpdate(ctx, config.Key, config.Value)
}
// SaveOrUpdateSystemConfig creates or updates a config row and invalidates the admin cache.
// SaveOrUpdateSystemConfig creates or updates a config row.
func SaveOrUpdateSystemConfig(ctx context.Context, key, value string) error {
if err := ensureAdminStore(ctx); err != nil {
svc, err := ensureConfigService()
if err != nil {
return err
}
return adminrepo.SaveOrUpdateSystemConfig(ctx, key, value)
return svc.SaveOrUpdate(ctx, key, value)
}
// InvalidateSystemConfigCache evicts one key from Wavelet's system-config cache.
// InvalidateSystemConfigCache evicts one key from the system-config cache.
func InvalidateSystemConfigCache(ctx context.Context, key string) error {
if err := ensureAdminStore(ctx); err != nil {
svc, err := ensureConfigService()
if err != nil {
return err
}
return adminrepo.InvalidateSystemConfigCache(ctx, key)
return svc.InvalidateCache(ctx, key)
}
// InvalidateAllSystemConfigCaches evicts the whole Wavelet system-config cache.
// InvalidateAllSystemConfigCaches evicts the whole system-config cache.
func InvalidateAllSystemConfigCaches(ctx context.Context) error {
if err := ensureAdminStore(ctx); err != nil {
svc, err := ensureConfigService()
if err != nil {
return err
}
return adminrepo.InvalidateAllSystemConfigCaches(ctx)
return svc.InvalidateAllCaches(ctx)
}
// StopSystemConfigCacheListener is retained for existing tests.
func StopSystemConfigCacheListener() {
adminrepo.StopSystemConfigCacheListener()
}
// StopSystemConfigCacheListener is retained for test compatibility.
func StopSystemConfigCacheListener() {}
// ResetSystemConfigRAMCacheForTest clears the process-local admin config cache.
func ResetSystemConfigRAMCacheForTest() {
adminrepo.ResetSystemConfigRAMCacheForTest()
if svc := currentConfigService(); svc != nil {
_ = svc.InvalidateAllCaches(context.Background())
}
}
// ListAdminSystemConfigs returns configs, optionally filtered by type.
func ListAdminSystemConfigs(ctx context.Context, configType string) ([]model.SystemConfig, error) {
if err := ensureAdminStore(ctx); err != nil {
svc, err := ensureConfigService()
if err != nil {
return nil, err
}
return adminrepo.ListAdminSystemConfigs(ctx, configType)
dtos, err := svc.ListByType(ctx, configType)
if err != nil {
return nil, err
}
res := make([]model.SystemConfig, len(dtos))
for i, v := range dtos {
res[i] = model.FromSystemConfigDTO(v)
}
return res, nil
}
@@ -6,12 +6,34 @@ package repository
import (
"context"
"errors"
"sync"
"Wavelet/core/contracts"
"Wavelet/openflare/plugins/server/kernel/model"
adminrepo "Wavelet/plugins/domain/admin/repository"
)
const fallbackSystemUserID uint64 = 999
const (
fallbackSystemUserID uint64 = 999
configTypeSystem = "system"
)
var (
taskMu sync.RWMutex
taskSvc contracts.TaskService
)
// SetTaskService injects the platform TaskService.
func SetTaskService(s contracts.TaskService) {
taskMu.Lock()
defer taskMu.Unlock()
taskSvc = s
}
func currentTaskService() contracts.TaskService {
taskMu.RLock()
defer taskMu.RUnlock()
return taskSvc
}
// GetActiveAuthSources lists enabled Wavelet auth sources via AuthService.
func GetActiveAuthSources(ctx context.Context) ([]model.AuthSource, error) {
@@ -41,11 +63,12 @@ func GetActiveAuthSources(ctx context.Context) ([]model.AuthSource, error) {
}
// GetTaskExecutionByTaskID loads a task execution by public task ID.
func GetTaskExecutionByTaskID(ctx context.Context, taskID string) (*model.TaskExecution, error) {
if err := ensureAdminStore(ctx); err != nil {
return nil, err
func GetTaskExecutionByTaskID(ctx context.Context, taskID string) (*contracts.TaskExecutionDTO, error) {
svc := currentTaskService()
if svc == nil {
return nil, errors.New("task service not initialized")
}
return adminrepo.GetTaskExecutionByTaskID(ctx, taskID)
return svc.GetExecutionByTaskID(ctx, taskID)
}
// GetSystemUser loads the built-in system user via UserService, or a synthetic fallback.
@@ -5,15 +5,10 @@ package repository
import (
"context"
"errors"
"testing"
"Wavelet/core/contracts"
"Wavelet/pkg/idgen"
adminmodel "Wavelet/plugins/domain/admin/model"
"Wavelet/plugins/infra/database"
"github.com/glebarez/sqlite"
"gorm.io/gorm"
)
type stubUserService struct {
@@ -34,24 +29,6 @@ func (s stubAuthService) ListAuthSources(context.Context) ([]contracts.AuthSourc
return s.sources, nil
}
func setupRepoTestDB(t *testing.T) (*gorm.DB, func()) {
t.Helper()
sqliteDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{
DisableForeignKeyConstraintWhenMigrating: true,
})
if err != nil {
t.Fatalf("gorm.Open() error = %v", err)
}
if err := sqliteDB.AutoMigrate(&adminmodel.TaskExecution{}); err != nil {
t.Fatalf("AutoMigrate(TaskExecution) error = %v", err)
}
if err := idgen.Init(1); err != nil {
t.Fatalf("idgen.Init() error = %v", err)
}
database.SetDB(sqliteDB)
return sqliteDB, func() { database.SetDB(nil) }
}
func TestGetActiveAuthSourcesUsesAuthService(t *testing.T) {
SetAuthService(stubAuthService{})
t.Cleanup(func() { SetAuthService(nil) })
@@ -96,21 +73,27 @@ func TestGetSystemUserUsesUserService(t *testing.T) {
}
}
func TestGetTaskExecutionByTaskIDUsesAdminStore(t *testing.T) {
_, cleanup := setupRepoTestDB(t)
t.Cleanup(cleanup)
type mockTaskSvc struct {
contracts.TaskService
execution contracts.TaskExecutionDTO
}
ctx := context.Background()
row := &adminmodel.TaskExecution{
func (m *mockTaskSvc) GetExecutionByTaskID(ctx context.Context, taskID string) (*contracts.TaskExecutionDTO, error) {
if taskID == m.execution.TaskID {
return &m.execution, nil
}
return nil, errors.New("not found")
}
func TestGetTaskExecutionByTaskIDUsesAdminStore(t *testing.T) {
SetTaskService(&mockTaskSvc{execution: contracts.TaskExecutionDTO{
ID: 7,
TaskID: "task-public-id",
TaskType: "pages_source_action",
Status: adminmodel.TaskExecutionStatusPending,
}
if err := database.DB(ctx).Create(row).Error; err != nil {
t.Fatalf("Create(TaskExecution) error = %v", err)
}
}})
t.Cleanup(func() { SetTaskService(nil) })
ctx := context.Background()
got, err := GetTaskExecutionByTaskID(ctx, "task-public-id")
if err != nil {
t.Fatalf("GetTaskExecutionByTaskID(%q) error = %v", "task-public-id", err)
@@ -6,15 +6,27 @@ package runtimeconfig
import (
"sync"
"Wavelet/plugins/infra/database"
)
// ClickHouseConfig represents ClickHouse connection parameters.
type ClickHouseConfig struct {
Enabled bool `config:"enabled" env:"CLICKHOUSE_ENABLED" default:"false" autoEnable:"CLICKHOUSE_HOST"`
Hosts []string `config:"hosts" env:"CLICKHOUSE_HOST"`
Username string `config:"username" env:"CLICKHOUSE_USERNAME"`
Password string `config:"password" env:"CLICKHOUSE_PASSWORD" secret:"true"`
Database string `config:"database" env:"CLICKHOUSE_NAME" default:"wavelet"`
MaxIdleConn int `config:"max_idle_conn" env:"CLICKHOUSE_MAX_IDLE_CONN" default:"10"`
MaxOpenConn int `config:"max_open_conn" env:"CLICKHOUSE_MAX_OPEN_CONN" default:"50"`
ConnMaxLifetime int `config:"conn_max_lifetime" env:"CLICKHOUSE_CONN_MAX_LIFETIME" default:"3600"`
DialTimeout int `config:"dial_timeout" env:"CLICKHOUSE_DIAL_TIMEOUT" default:"10"`
BlockBufferSize uint8 `config:"block_buffer_size" env:"CLICKHOUSE_BLOCK_BUFFER_SIZE" default:"10"`
}
// Snapshot is the subset of host config remaining OF packages still need.
type Snapshot struct {
SessionSecret string
DatabaseEnabled bool
ClickHouse database.ClickHouseConfig
ClickHouse ClickHouseConfig
}
var (
@@ -0,0 +1,113 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package testhelper
import (
"bytes"
"context"
"io"
"sync"
"sync/atomic"
"time"
"Wavelet/core/contracts"
"Wavelet/openflare/plugins/server/kernel/repository"
"gorm.io/gorm"
)
// MockStorageService provides an in-memory contracts.StorageService for tests.
type MockStorageService struct {
mu sync.RWMutex
objects map[string][]byte
seq uint64
}
// NewMockStorageService creates an initialized MockStorageService.
func NewMockStorageService() *MockStorageService {
return &MockStorageService{
objects: make(map[string][]byte),
}
}
// Put writes an object into memory.
func (m *MockStorageService) Put(_ context.Context, key string, body io.Reader, _ int64, _ string) (contracts.StoragePutResult, error) {
m.mu.Lock()
defer m.mu.Unlock()
data, err := io.ReadAll(body)
if err != nil {
return contracts.StoragePutResult{}, err
}
m.objects[key] = data
return contracts.StoragePutResult{Key: key, Bucket: "test-bucket"}, nil
}
// Get reads an object from memory.
func (m *MockStorageService) Get(_ context.Context, key string) (*contracts.StorageObject, error) {
m.mu.RLock()
defer m.mu.RUnlock()
data, ok := m.objects[key]
if ok {
return &contracts.StorageObject{
Key: key,
Body: io.NopCloser(bytes.NewReader(data)),
ContentLength: int64(len(data)),
ContentType: "application/octet-stream",
}, nil
}
return nil, gorm.ErrRecordNotFound
}
// Delete removes an object from memory.
func (m *MockStorageService) Delete(_ context.Context, key string) error {
m.mu.Lock()
defer m.mu.Unlock()
delete(m.objects, key)
return nil
}
// Ingest ingests content into mock storage.
func (m *MockStorageService) Ingest(ctx context.Context, r io.Reader, opts contracts.IngestOptions) (*contracts.IngestResult, error) {
id := atomic.AddUint64(&m.seq, 1)
m.mu.Lock()
data, _ := io.ReadAll(r)
key := opts.FileName
if key == "" {
key = "file.dat"
}
m.objects[key] = data
m.mu.Unlock()
gdb := repository.DB(ctx)
if gdb != nil {
type testUpload struct {
ID uint64 `gorm:"primaryKey"`
UserID uint64
FileName string
FilePath string
MimeType string
Size int64
Status string
Type string
Metadata contracts.UploadMetadataDTO `gorm:"serializer:json;type:jsonb"`
CreatedAt time.Time
UpdatedAt time.Time
}
u := testUpload{
ID: id,
UserID: opts.UserID,
FileName: key,
FilePath: "mock/" + key,
MimeType: opts.MimeType,
Size: opts.Size,
Status: "used",
Type: opts.Type,
Metadata: contracts.UploadMetadataDTO{Extra: opts.Metadata},
CreatedAt: time.Now().UTC(),
UpdatedAt: time.Now().UTC(),
}
_ = gdb.Table("w_uploads").Save(&u).Error
}
return &contracts.IngestResult{ID: id, Key: "mock/" + key, Created: true, Stored: true}, nil
}
@@ -9,9 +9,8 @@ import (
"time"
"Wavelet/core/contracts"
"Wavelet/openflare/plugins/server/kernel/repository"
"Wavelet/pkg/idgen"
adminmodel "Wavelet/plugins/domain/admin/model"
"Wavelet/plugins/infra/database"
)
// NoopTaskService is a contracts.TaskService that records dispatch and ignores the rest.
@@ -22,35 +21,76 @@ type NoopTaskService struct {
var _ contracts.TaskService = (*NoopTaskService)(nil)
// Dispatch dispatches a task mock execution.
func (s *NoopTaskService) Dispatch(ctx context.Context, taskType string, payload []byte, triggeredBy string) (string, error) {
s.LastType = taskType
s.LastPayload = payload
taskID := fmt.Sprintf("test-task-%d", time.Now().UnixNano())
if conn := database.DB(ctx); conn != nil {
_ = conn.Create(&adminmodel.TaskExecution{
ID: idgen.NextUint64ID(),
gdb := repository.DB(ctx)
if gdb != nil {
var id uint64
func() {
defer func() {
if r := recover(); r != nil {
id = uint64(time.Now().UnixNano())
}
}()
id = idgen.NextUint64ID()
}()
_ = gdb.Table("w_task_executions").Create(&contracts.TaskExecutionDTO{
ID: id,
TaskID: taskID,
TaskType: taskType,
Status: adminmodel.TaskExecutionStatusPending,
TriggeredBy: triggeredBy,
Payload: string(payload),
TriggeredBy: triggeredBy,
Status: "pending",
CreatedAt: time.Now().UTC(),
UpdatedAt: time.Now().UTC(),
}).Error
}
return taskID, nil
}
// Retry retries a task mock execution.
func (s *NoopTaskService) Retry(context.Context, uint64) (string, error) { return "", nil }
func (s *NoopTaskService) ListTasks() []contracts.TaskMetaDTO { return nil }
// ListTasks lists task mock metadata.
func (s *NoopTaskService) ListTasks() []contracts.TaskMetaDTO { return nil }
// GetTaskMeta returns task mock metadata.
func (s *NoopTaskService) GetTaskMeta(string) (contracts.TaskMetaDTO, bool) {
return contracts.TaskMetaDTO{}, false
}
// ValidatePayload validates task payload.
func (s *NoopTaskService) ValidatePayload(_ string, payload []byte) ([]byte, error) {
return payload, nil
}
func (s *NoopTaskService) ReloadScheduler() error { return nil }
// ReloadScheduler reloads task scheduler.
func (s *NoopTaskService) ReloadScheduler() error { return nil }
// AppendLog appends log message.
func (s *NoopTaskService) AppendLog(context.Context, string, ...any) {}
// ListExecutions lists task executions.
func (s *NoopTaskService) ListExecutions(context.Context, string, string, int, int) ([]contracts.TaskExecutionDTO, int64, error) {
return nil, 0, nil
}
// GetExecution gets task execution by ID.
func (s *NoopTaskService) GetExecution(context.Context, uint64) (*contracts.TaskExecutionDTO, error) {
return &contracts.TaskExecutionDTO{TaskID: "test-task"}, nil
}
// GetExecutionByTaskID gets task execution by taskID.
func (s *NoopTaskService) GetExecutionByTaskID(ctx context.Context, taskID string) (*contracts.TaskExecutionDTO, error) {
gdb := repository.DB(ctx)
if gdb != nil {
var exec contracts.TaskExecutionDTO
if err := gdb.Table("w_task_executions").Where("task_id = ?", taskID).First(&exec).Error; err == nil {
return &exec, nil
}
}
return &contracts.TaskExecutionDTO{ID: 1, TaskID: taskID, Payload: string(s.LastPayload)}, nil
}
@@ -23,31 +23,51 @@ func passThrough() gin.HandlerFunc {
return func(c *gin.Context) { c.Next() }
}
func (s StubAuth) RequireAuthMiddleware() any { return passThrough() }
// RequireAuthMiddleware returns a passthrough middleware.
func (s StubAuth) RequireAuthMiddleware() any { return passThrough() }
// RequireAdminMiddleware returns a passthrough middleware.
func (s StubAuth) RequireAdminMiddleware() any { return passThrough() }
// DisallowTokenAuthMiddleware returns a passthrough middleware.
func (s StubAuth) DisallowTokenAuthMiddleware() any {
return passThrough()
}
// GetCurrentUser returns the stub user.
func (s StubAuth) GetCurrentUser(context.Context) (*contracts.UserDTO, error) {
return s.User, nil
}
// GetCurrentUserID returns the stub user ID.
func (s StubAuth) GetCurrentUserID(context.Context) (uint64, error) {
if s.User == nil {
return 0, nil
}
return s.User.ID, nil
}
// VerifyToken returns the stub user.
func (s StubAuth) VerifyToken(context.Context, string) (*contracts.UserDTO, error) {
return s.User, nil
}
// CreateSession creates a stub session.
func (s StubAuth) CreateSession(context.Context, uint64, map[string]any) (string, error) {
return "", nil
}
func (s StubAuth) RevokeToken(context.Context, string) error { return nil }
// RevokeToken revokes a stub token.
func (s StubAuth) RevokeToken(context.Context, string) error { return nil }
// RevokeUserSessions revokes stub user sessions.
func (s StubAuth) RevokeUserSessions(context.Context, uint64) error { return nil }
func (s StubAuth) InvalidateCachedUser(context.Context, uint64) {}
func (s StubAuth) InvalidateCachedToken(context.Context, string) {}
// InvalidateCachedUser invalidates stub cached user.
func (s StubAuth) InvalidateCachedUser(context.Context, uint64) {}
// InvalidateCachedToken invalidates stub cached token.
func (s StubAuth) InvalidateCachedToken(context.Context, string) {}
func (s StubAuth) ListAuthSources(context.Context) ([]contracts.AuthSourceViewDTO, error) {
return s.Sources, nil
}
@@ -6,15 +6,21 @@
package testhelper
import (
"bytes"
"context"
"io"
"strconv"
"testing"
"time"
"Wavelet/core/contracts"
"Wavelet/openflare/plugins/server/kernel/model"
analyticsmodel "Wavelet/openflare/plugins/server/kernel/model/analytics"
"Wavelet/openflare/plugins/server/kernel/ofupload"
"Wavelet/openflare/plugins/server/kernel/repository"
"Wavelet/openflare/plugins/server/kernel/repository/logstore"
oftask "Wavelet/openflare/plugins/server/kernel/task"
"Wavelet/pkg/idgen"
db "Wavelet/plugins/infra/database"
"github.com/glebarez/sqlite"
"github.com/stretchr/testify/require"
@@ -28,6 +34,87 @@ const (
configValueFalse = "false"
)
type testConfigService struct {
db *gorm.DB
}
// NewMockSystemConfigService creates a test SystemConfigService backed by GORM.
func NewMockSystemConfigService(db *gorm.DB) contracts.SystemConfigService {
return &testConfigService{db: db}
}
func (s *testConfigService) GetByKey(ctx context.Context, key string) (contracts.SystemConfigDTO, error) {
var cfg contracts.SystemConfigDTO
err := s.db.WithContext(ctx).Table("w_system_configs").Where("key = ?", key).First(&cfg).Error
return cfg, err
}
func (s *testConfigService) ListByKeys(ctx context.Context, keys []string) (map[string]contracts.SystemConfigDTO, error) {
var cfgs []contracts.SystemConfigDTO
if err := s.db.WithContext(ctx).Table("w_system_configs").Where("key IN ?", keys).Find(&cfgs).Error; err != nil {
return nil, err
}
res := make(map[string]contracts.SystemConfigDTO, len(cfgs))
for _, c := range cfgs {
res[c.Key] = c
}
return res, nil
}
func (s *testConfigService) ListVisible(ctx context.Context) ([]contracts.SystemConfigDTO, error) {
var cfgs []contracts.SystemConfigDTO
err := s.db.WithContext(ctx).Table("w_system_configs").Where("visibility = ?", 1).Find(&cfgs).Error
return cfgs, err
}
func (s *testConfigService) ListByType(ctx context.Context, configType string) ([]contracts.SystemConfigDTO, error) {
var cfgs []contracts.SystemConfigDTO
err := s.db.WithContext(ctx).Table("w_system_configs").Where("type = ?", configType).Find(&cfgs).Error
return cfgs, err
}
func (s *testConfigService) GetIntByKey(ctx context.Context, key string) (int, error) {
var cfg contracts.SystemConfigDTO
if err := s.db.WithContext(ctx).Table("w_system_configs").Where("key = ?", key).First(&cfg).Error; err != nil {
return 0, err
}
return strconv.Atoi(cfg.Value)
}
func (s *testConfigService) GetBoolByKey(ctx context.Context, key string) (bool, error) {
var cfg contracts.SystemConfigDTO
if err := s.db.WithContext(ctx).Table("w_system_configs").Where("key = ?", key).First(&cfg).Error; err != nil {
return false, err
}
return strconv.ParseBool(cfg.Value)
}
func (s *testConfigService) SaveOrUpdate(ctx context.Context, key, value string) error {
var cfg model.SystemConfig
if err := s.db.WithContext(ctx).Table("w_system_configs").Where("key = ?", key).First(&cfg).Error; err != nil {
cfg = model.SystemConfig{Key: key, Value: value, Type: "system"}
return s.db.WithContext(ctx).Table("w_system_configs").Create(&cfg).Error
}
return s.db.WithContext(ctx).Table("w_system_configs").Where("key = ?", key).Update("value", value).Error
}
func (s *testConfigService) InvalidateCache(ctx context.Context, key string) error { return nil }
func (s *testConfigService) InvalidateAllCaches(ctx context.Context) error { return nil }
type testSystemConfigEntity struct {
Key string `gorm:"primaryKey"`
Value string
Type string
Visibility int
Description string
UpdatedAt time.Time
CreatedAt time.Time
}
func (testSystemConfigEntity) TableName() string {
return "w_system_configs"
}
// SetupTestEnvironment initializes an in-memory SQLite DB and seeds default
// configurations. Redis is no longer owned by OpenFlare.
func SetupTestEnvironment(t *testing.T) (*gorm.DB, any, func()) {
@@ -44,40 +131,133 @@ func SetupTestEnvironment(t *testing.T) (*gorm.DB, any, func()) {
}
err = sqliteDB.AutoMigrate(
&testSystemConfigEntity{},
&model.User{},
&model.AuthSource{},
&model.ExternalAccount{},
&model.SystemConfig{},
&model.AccessToken{},
&model.Upload{},
&model.UploadStat{},
&model.TaskExecution{},
&model.Template{},
&model.AccessToken{},
&model.Schedule{},
)
if err != nil {
t.Fatalf("failed to auto migrate tables: %v", err)
}
db.SetDB(sqliteDB)
repository.SetDBForTest(sqliteDB)
repository.SetSystemConfigService(&testConfigService{db: sqliteDB})
mockStorage := NewMockStorageService()
ofupload.SetStorage(mockStorage)
ofupload.SetUploadService(&mockUploadService{db: sqliteDB})
noopTask := &NoopTaskService{}
repository.SetTaskService(noopTask)
oftask.SetService(noopTask)
if err := idgen.Init(1); err != nil {
t.Fatalf("idgen.Init: %v", err)
}
seedDefaultConfigs(t, sqliteDB)
repository.ResetSystemConfigRAMCacheForTest()
cleanup := func() {
runExtraCleanups()
repository.StopSystemConfigCacheListener()
repository.ResetSystemConfigRAMCacheForTest()
repository.SetAuthService(nil)
repository.SetUserService(nil)
db.SetDB(nil)
repository.SetSystemConfigService(nil)
repository.SetTaskService(nil)
repository.SetDBForTest(nil)
ofupload.SetStorage(nil)
ofupload.SetUploadService(nil)
oftask.SetService(nil)
}
return sqliteDB, nil, cleanup
}
type mockUploadService struct {
db *gorm.DB
}
// NewMockUploadService creates a mock UploadService backed by GORM.
func NewMockUploadService(db *gorm.DB) contracts.UploadService {
return &mockUploadService{db: db}
}
func (s *mockUploadService) GetByID(ctx context.Context, id uint64) (*contracts.UploadDTO, error) {
var u contracts.UploadDTO
err := s.db.WithContext(ctx).Table("w_uploads").Where("id = ?", id).First(&u).Error
if err != nil {
return &contracts.UploadDTO{
ID: id,
Status: "used",
Type: "openflare_pages_deployment",
Size: 100,
CreatedAt: time.Now().UTC(),
UpdatedAt: time.Now().UTC(),
}, nil
}
return &u, nil
}
func (s *mockUploadService) OpenStoredUpload(ctx context.Context, id uint64) (*contracts.OpenedUploadDTO, error) {
u, err := s.GetByID(ctx, id)
if err != nil {
return nil, err
}
body := io.ReadCloser(io.NopCloser(bytes.NewReader(nil)))
storage := ofupload.CurrentStorage()
if storage != nil {
if obj, err := storage.Get(ctx, u.FilePath); err == nil && obj != nil && obj.Body != nil {
body = obj.Body
} else if obj, err := storage.Get(ctx, u.FileName); err == nil && obj != nil && obj.Body != nil {
body = obj.Body
}
}
return &contracts.OpenedUploadDTO{
Upload: *u,
Body: body,
ContentType: u.MimeType,
ContentLength: u.Size,
}, nil
}
func (s *mockUploadService) Remove(ctx context.Context, id uint64) error {
if err := s.db.WithContext(ctx).Table("w_uploads").Where("id = ?", id).Update("status", "deleted").Error; err != nil {
return err
}
return s.RebuildStats(ctx)
}
func (s *mockUploadService) RemoveOwned(ctx context.Context, id uint64, userID uint64) error {
if err := s.db.WithContext(ctx).Table("w_uploads").Where("id = ? AND user_id = ?", id, userID).Update("status", "deleted").Error; err != nil {
return err
}
return s.RebuildStats(ctx)
}
func (s *mockUploadService) FindByHash(ctx context.Context, hash string, size int64) (*contracts.UploadDTO, error) {
var u contracts.UploadDTO
err := s.db.WithContext(ctx).Table("w_uploads").Where("hash = ? AND size = ?", hash, size).First(&u).Error
if err != nil {
return nil, err
}
return &u, nil
}
func (s *mockUploadService) RebuildStats(ctx context.Context) error {
var count int64
_ = s.db.WithContext(ctx).Table("w_uploads").Where("status != ?", "deleted").Count(&count).Error
var stat model.UploadStat
if err := s.db.WithContext(ctx).Table("w_upload_stats").Where("dimension = ?", model.UploadStatDimensionTotal).First(&stat).Error; err != nil {
stat = model.UploadStat{
Dimension: model.UploadStatDimensionTotal,
FileCount: int(count),
}
return s.db.WithContext(ctx).Table("w_upload_stats").Create(&stat).Error
}
stat.FileCount = int(count)
return s.db.WithContext(ctx).Table("w_upload_stats").Save(&stat).Error
}
func getSeedConfigsPart1() []model.SystemConfig {
return []model.SystemConfig{
{Key: model.ConfigKeyUploadAllowedExtensions, Value: "jpg,png,webp", Type: configTypeSystem, Description: "允许上传的图片扩展名(逗号分隔)"},
@@ -128,7 +308,7 @@ func getSeedConfigsPart2() []model.SystemConfig {
func seedDefaultConfigs(t *testing.T, tx *gorm.DB) {
t.Helper()
defaultConfigs := append(getSeedConfigsPart1(), getSeedConfigsPart2()...)
if err := tx.Create(&defaultConfigs).Error; err != nil {
if err := tx.Table("w_system_configs").Create(&defaultConfigs).Error; err != nil {
t.Fatalf("failed to seed default system configs: %v", err)
}
@@ -148,18 +328,18 @@ func seedDefaultConfigs(t *testing.T, tx *gorm.DB) {
model.ConfigKeySearchEngineIndexingEnabled,
model.ConfigKeyFileAccessWhitelist,
}
if err := tx.Model(&model.SystemConfig{}).
if err := tx.Table("w_system_configs").
Where("key IN ?", publicKeys).
Update("visibility", model.ConfigVisibilityVisible).Error; err != nil {
t.Fatalf("failed to seed public system config visibility: %v", err)
}
}
// SetupLogStoresForTest 将 logstore 指向测试已通过 db.SetDB 注入的 sqlite 库。
// SetupLogStoresForTest 将 logstore 指向测试已通过 SetDBForTest 注入的 sqlite 库。
func SetupLogStoresForTest(t *testing.T) {
t.Helper()
gdb := db.DB(context.Background())
gdb := repository.DB(context.Background())
require.NoError(t, idgen.Init(1))
require.NoError(t, gdb.AutoMigrate(
&analyticsmodel.NodeAccessLog{},