mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-09-29 05:56:38 +08:00
refactor(core): align with cordis spatiotemporal composability architecture
- Purify core micro-kernel by removing context hardcoded helpers and reverse dependencies - Eliminate init() side effects in infra plugins with reversible lifecycle disposal - Completely isolate plugins by removing cross-plugin imports and using core/contracts - Introduce TaskService and RiskControlService contracts for unified cross-plugin APIs - Regenerate Swagger documentation and update developer guide matrix - Achieve 0 violations in check_cordis_architecture.sh and 100% test pass
This commit is contained in:
@@ -63,3 +63,5 @@ s3_cache
|
||||
|
||||
.worktrees/
|
||||
/.superpowers/
|
||||
/backend/plugins/domain/upload/filesrv/uploads/
|
||||
/backend/plugins/domain/upload/task/uploads/
|
||||
|
||||
+12
-5
@@ -25,17 +25,17 @@ linters:
|
||||
- gocritic # 各类代码问题
|
||||
- funlen # 函数过长
|
||||
|
||||
- gosec # 安全问题检查
|
||||
- gosec # 安全问题检查
|
||||
- bodyclose # HTTP response body 没有正确关闭
|
||||
- noctx # 没有传递 context.Context
|
||||
- contextcheck # 其他检查
|
||||
- sqlclosecheck # SQL rows 没有正确关闭
|
||||
- unconvert # 不必要的类型转换
|
||||
- nilerr # 函数返回 nil 错误
|
||||
- sqlclosecheck # SQL rows 没有正确关闭
|
||||
- unconvert # 不必要的类型转换
|
||||
- nilerr # 函数返回 nil 错误
|
||||
|
||||
settings:
|
||||
dupl:
|
||||
threshold: 120
|
||||
threshold: 80
|
||||
|
||||
cyclop:
|
||||
max-complexity: 20
|
||||
@@ -53,3 +53,10 @@ linters:
|
||||
- argument
|
||||
- condition
|
||||
- return
|
||||
|
||||
formatters:
|
||||
enable:
|
||||
- gofumpt
|
||||
settings:
|
||||
gofumpt:
|
||||
extra-rules: true
|
||||
|
||||
@@ -14,8 +14,9 @@ license-check:
|
||||
scripts/update_go_license.sh --check
|
||||
|
||||
format:
|
||||
@echo "==> Formatting backend Go source..."
|
||||
gofmt -w $$(find backend -type f -name '*.go' -not -path './.git/*')
|
||||
@echo "==> Formatting backend Go source with goimports..."
|
||||
@command -v goimports >/dev/null 2>&1 || { echo 'error: goimports is required. Run: go install golang.org/x/tools/cmd/goimports@latest' >&2; exit 1; }
|
||||
goimports -w -local $(MODULE) $$(find backend -type f -name '*.go' -not -path './.git/*')
|
||||
@echo "==> Formatting frontend source..."
|
||||
cd frontend && pnpm format
|
||||
|
||||
|
||||
+2
-1
@@ -7,8 +7,9 @@ package cmd
|
||||
import (
|
||||
"log"
|
||||
|
||||
"Wavelet/core"
|
||||
"github.com/spf13/cobra"
|
||||
|
||||
"Wavelet/core"
|
||||
)
|
||||
|
||||
var allCmd = &cobra.Command{
|
||||
|
||||
+2
-1
@@ -6,8 +6,9 @@ package cmd
|
||||
import (
|
||||
"log"
|
||||
|
||||
"Wavelet/core"
|
||||
"github.com/spf13/cobra"
|
||||
|
||||
"Wavelet/core"
|
||||
)
|
||||
|
||||
var apiCmd = &cobra.Command{
|
||||
|
||||
+3
-2
@@ -10,6 +10,9 @@ import (
|
||||
"log"
|
||||
"time"
|
||||
|
||||
"github.com/pressly/goose/v3"
|
||||
goosedb "github.com/pressly/goose/v3/database"
|
||||
|
||||
"Wavelet/core"
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/pkg/config"
|
||||
@@ -28,8 +31,6 @@ import (
|
||||
infradb "Wavelet/plugins/infra/database"
|
||||
"Wavelet/plugins/infra/logger"
|
||||
"Wavelet/plugins/infra/storage"
|
||||
"github.com/pressly/goose/v3"
|
||||
goosedb "github.com/pressly/goose/v3/database"
|
||||
)
|
||||
|
||||
// newWaveletApp creates a core.App wired with Wavelet platform infrastructure, domain plugins, and profile drivers.
|
||||
|
||||
@@ -16,9 +16,10 @@ import (
|
||||
userdomain "Wavelet/plugins/domain/user"
|
||||
"Wavelet/plugins/infra/database"
|
||||
|
||||
"Wavelet/plugins/domain/auth"
|
||||
"github.com/spf13/cobra"
|
||||
"gorm.io/gorm"
|
||||
|
||||
"Wavelet/plugins/domain/auth"
|
||||
)
|
||||
|
||||
var (
|
||||
@@ -49,7 +50,10 @@ var resetPasswdCmd = &cobra.Command{
|
||||
ctx := context.Background()
|
||||
|
||||
// Ensure database is initialized
|
||||
database.DB(ctx)
|
||||
dbConn := database.DB(ctx)
|
||||
if dbConn != nil {
|
||||
userdomain.SetDBService(database.NewService(dbConn))
|
||||
}
|
||||
|
||||
var username string
|
||||
if usernameFlag != "" {
|
||||
|
||||
+2
-1
@@ -8,11 +8,12 @@ import (
|
||||
"log"
|
||||
"time"
|
||||
|
||||
"github.com/spf13/cobra"
|
||||
|
||||
"Wavelet/pkg/buildinfo"
|
||||
"Wavelet/pkg/config"
|
||||
"Wavelet/pkg/logger"
|
||||
"Wavelet/pkg/trace"
|
||||
"github.com/spf13/cobra"
|
||||
)
|
||||
|
||||
const traceShutdownTimeout = 10 * time.Second
|
||||
|
||||
@@ -6,8 +6,9 @@ package cmd
|
||||
import (
|
||||
"log"
|
||||
|
||||
"Wavelet/core"
|
||||
"github.com/spf13/cobra"
|
||||
|
||||
"Wavelet/core"
|
||||
)
|
||||
|
||||
var schedulerCmd = &cobra.Command{
|
||||
|
||||
@@ -6,8 +6,9 @@ package cmd
|
||||
import (
|
||||
"log"
|
||||
|
||||
"Wavelet/core"
|
||||
"github.com/spf13/cobra"
|
||||
|
||||
"Wavelet/core"
|
||||
)
|
||||
|
||||
var workerCmd = &cobra.Command{
|
||||
|
||||
@@ -10,7 +10,6 @@ import (
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/core/extpoints"
|
||||
)
|
||||
|
||||
@@ -237,24 +236,6 @@ func (c *Context) Setting() extpoints.SettingExtension {
|
||||
return c.Settings()
|
||||
}
|
||||
|
||||
// DB returns the contracts.DBService registered in the IoC container, or nil if not registered.
|
||||
func (c *Context) DB() contracts.DBService {
|
||||
svc, err := Inject[contracts.DBService](c)
|
||||
if err != nil {
|
||||
return nil
|
||||
}
|
||||
return svc
|
||||
}
|
||||
|
||||
// Cache returns the contracts.CacheService registered in the IoC container, or nil if not registered.
|
||||
func (c *Context) Cache() contracts.CacheService {
|
||||
svc, err := Inject[contracts.CacheService](c)
|
||||
if err != nil {
|
||||
return nil
|
||||
}
|
||||
return svc
|
||||
}
|
||||
|
||||
// OnDispose registers a cleanup callback function to be executed when this Context is disposed.
|
||||
// It accepts func() error, func(), or Disposer.
|
||||
func (c *Context) OnDispose(fn any) {
|
||||
|
||||
@@ -44,6 +44,25 @@ const (
|
||||
EventTopicSystemCleanup = "admin:system_cleanup"
|
||||
)
|
||||
|
||||
// --- Task Events ---
|
||||
const (
|
||||
// EventTopicTaskCompleted fires when an asynchronous background task execution finishes.
|
||||
EventTopicTaskCompleted = "task:completed"
|
||||
)
|
||||
|
||||
// TaskCompletedEvent carries task execution outcome details.
|
||||
type TaskCompletedEvent struct {
|
||||
TaskID string `json:"task_id"`
|
||||
TaskName string `json:"task_name"`
|
||||
TaskType string `json:"task_type"`
|
||||
Status string `json:"status"`
|
||||
Duration int64 `json:"duration"`
|
||||
ErrorMsg string `json:"error_msg,omitempty"`
|
||||
ResultMsg string `json:"result_msg,omitempty"`
|
||||
Payload string `json:"payload,omitempty"`
|
||||
Detail string `json:"detail,omitempty"`
|
||||
}
|
||||
|
||||
// --- Upload / Storage Events ---
|
||||
const (
|
||||
// EventTopicUploadCreated fires when a new file upload is recorded.
|
||||
|
||||
@@ -0,0 +1,66 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package contracts defines unified service interfaces and DTOs for cross-plugin communication.
|
||||
package contracts
|
||||
|
||||
import (
|
||||
"context"
|
||||
"time"
|
||||
)
|
||||
|
||||
// AccessLogFilterDTO defines filter criteria for querying user access logs.
|
||||
type AccessLogFilterDTO struct {
|
||||
UserIDs []uint64
|
||||
Path string
|
||||
StartTime *time.Time
|
||||
EndTime *time.Time
|
||||
}
|
||||
|
||||
// AccessLogDTO represents a single access log entry.
|
||||
type AccessLogDTO struct {
|
||||
ID uint64 `json:"id"`
|
||||
UserID uint64 `json:"user_id"`
|
||||
IP string `json:"ip"`
|
||||
UserAgent string `json:"user_agent"`
|
||||
Method string `json:"method"`
|
||||
Path string `json:"path"`
|
||||
Status int32 `json:"status"`
|
||||
Latency int64 `json:"latency"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
}
|
||||
|
||||
// AccessLogDailyStatsDTO represents aggregate access statistics for a single day.
|
||||
type AccessLogDailyStatsDTO struct {
|
||||
Date string `json:"date"`
|
||||
PV uint64 `json:"pv"`
|
||||
UV uint64 `json:"uv"`
|
||||
IPCount uint64 `json:"ip_count"`
|
||||
ErrorCount uint64 `json:"error_count"`
|
||||
AvgLatencyMs int64 `json:"avg_latency_ms"`
|
||||
SlowReqCount uint64 `json:"slow_req_count"`
|
||||
MaxLatencyMs int64 `json:"max_latency_ms"`
|
||||
P95LatencyMs int64 `json:"p95_latency_ms"`
|
||||
P99LatencyMs int64 `json:"p99_latency_ms"`
|
||||
}
|
||||
|
||||
// RiskControlService defines the contract for accessing security risk control and audit logstore.
|
||||
type RiskControlService interface {
|
||||
// QueryAccessLogs retrieves paginated access logs matching the filter.
|
||||
QueryAccessLogs(ctx context.Context, filter AccessLogFilterDTO, page, pageSize int) ([]AccessLogDTO, uint64, error)
|
||||
|
||||
// QueryAccessLogStats returns aggregate daily statistics for the last N days.
|
||||
QueryAccessLogStats(ctx context.Context, days int) ([]AccessLogDailyStatsDTO, error)
|
||||
|
||||
// ActiveLogEngine returns the current active logstore engine name.
|
||||
ActiveLogEngine(ctx context.Context) string
|
||||
|
||||
// IsLogEngineMigrating reports whether a log engine migration is in progress.
|
||||
IsLogEngineMigrating(ctx context.Context) bool
|
||||
|
||||
// Drain flushes pending in-flight log buffers.
|
||||
Drain(ctx context.Context) error
|
||||
|
||||
// SwitchLogEngine migrates and switches the active log storage engine.
|
||||
SwitchLogEngine(ctx context.Context, targetEngine string) error
|
||||
}
|
||||
@@ -46,6 +46,55 @@ type IngestResult struct {
|
||||
Resolved bool
|
||||
}
|
||||
|
||||
// StorageDriver identifies a supported storage backend.
|
||||
type StorageDriver string
|
||||
|
||||
const (
|
||||
StorageDriverLocal StorageDriver = "local"
|
||||
StorageDriverS3 StorageDriver = "s3"
|
||||
StorageDriverR2 StorageDriver = "r2"
|
||||
StorageDriverMinIO StorageDriver = "minio"
|
||||
StorageDriverOSS StorageDriver = "oss"
|
||||
StorageDriverWebDAV StorageDriver = "webdav"
|
||||
)
|
||||
|
||||
// LocalStorageConfigDTO configures local filesystem storage.
|
||||
type LocalStorageConfigDTO struct {
|
||||
Root string `json:"root"`
|
||||
}
|
||||
|
||||
// ObjectStorageConfigDTO configures S3-compatible or OSS object storage.
|
||||
type ObjectStorageConfigDTO struct {
|
||||
Endpoint string `json:"endpoint"`
|
||||
Region string `json:"region"`
|
||||
Bucket string `json:"bucket"`
|
||||
AccessKeyID string `json:"access_key_id"`
|
||||
SecretAccessKey string `json:"secret_access_key"`
|
||||
AccountID string `json:"account_id,omitempty"`
|
||||
PathStyle bool `json:"path_style"`
|
||||
KeyPrefix string `json:"key_prefix"`
|
||||
CDNURL string `json:"cdn_url"`
|
||||
}
|
||||
|
||||
// WebDAVStorageConfigDTO configures WebDAV storage.
|
||||
type WebDAVStorageConfigDTO struct {
|
||||
URL string `json:"url"`
|
||||
Username string `json:"username"`
|
||||
Password string `json:"password"`
|
||||
Root string `json:"root"`
|
||||
}
|
||||
|
||||
// StorageConfigDTO encapsulates full storage configuration across all backends.
|
||||
type StorageConfigDTO struct {
|
||||
Driver StorageDriver `json:"driver"`
|
||||
Local LocalStorageConfigDTO `json:"local"`
|
||||
S3 ObjectStorageConfigDTO `json:"s3"`
|
||||
R2 ObjectStorageConfigDTO `json:"r2"`
|
||||
MinIO ObjectStorageConfigDTO `json:"minio"`
|
||||
OSS ObjectStorageConfigDTO `json:"oss"`
|
||||
WebDAV WebDAVStorageConfigDTO `json:"webdav"`
|
||||
}
|
||||
|
||||
// StorageService defines the contract for unified object storage and managed file ingestion.
|
||||
type StorageService interface {
|
||||
// Put writes an object to storage.
|
||||
|
||||
@@ -0,0 +1,73 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package contracts defines unified service interfaces and DTOs for cross-plugin communication.
|
||||
package contracts
|
||||
|
||||
import (
|
||||
"context"
|
||||
"time"
|
||||
)
|
||||
|
||||
// TaskParamDTO describes a parameter accepted by a background task.
|
||||
type TaskParamDTO struct {
|
||||
Name string `json:"name"`
|
||||
Type string `json:"type"`
|
||||
Description string `json:"description"`
|
||||
Required bool `json:"required"`
|
||||
Default any `json:"default,omitempty"`
|
||||
}
|
||||
|
||||
// TaskMetaDTO describes the metadata and configuration of a registered background task.
|
||||
type TaskMetaDTO struct {
|
||||
Name string `json:"name"`
|
||||
DisplayName string `json:"display_name"`
|
||||
Description string `json:"description"`
|
||||
Category string `json:"category"`
|
||||
Params []TaskParamDTO `json:"params,omitempty"`
|
||||
MaxRetry int `json:"max_retry"`
|
||||
Timeout time.Duration `json:"timeout"`
|
||||
Queue string `json:"queue"`
|
||||
Schedule string `json:"schedule,omitempty"`
|
||||
}
|
||||
|
||||
// TaskResultDTO represents the outcome of a background task execution.
|
||||
type TaskResultDTO struct {
|
||||
Message string `json:"message"`
|
||||
Detail any `json:"detail,omitempty"`
|
||||
}
|
||||
|
||||
// TaskExecutionDTO represents a single task execution record.
|
||||
type TaskExecutionDTO struct {
|
||||
ID uint64 `json:"id,string"`
|
||||
TaskID string `json:"task_id"`
|
||||
TaskType string `json:"task_type"`
|
||||
TaskName string `json:"task_name"`
|
||||
Status string `json:"status"`
|
||||
Retryable bool `json:"retryable"`
|
||||
MaxRetry int `json:"max_retry"`
|
||||
RetryCount int `json:"retry_count"`
|
||||
Log string `json:"log"`
|
||||
ErrorMessage string `json:"error_message"`
|
||||
Result string `json:"result"`
|
||||
StartedAt *time.Time `json:"started_at"`
|
||||
FinishedAt *time.Time `json:"finished_at"`
|
||||
Duration int64 `json:"duration"`
|
||||
Payload string `json:"payload"`
|
||||
TriggeredBy string `json:"triggered_by"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
UpdatedAt time.Time `json:"updated_at"`
|
||||
}
|
||||
|
||||
// TaskService defines the unified contract for dispatching and tracking background tasks.
|
||||
type TaskService interface {
|
||||
Dispatch(ctx context.Context, taskType string, payload []byte, triggeredBy string) (string, error)
|
||||
Retry(ctx context.Context, id uint64) (string, error)
|
||||
ListTasks() []TaskMetaDTO
|
||||
GetTaskMeta(taskType string) (TaskMetaDTO, bool)
|
||||
ValidatePayload(taskType string, payload []byte) ([]byte, error)
|
||||
ReloadScheduler() error
|
||||
AppendLog(ctx context.Context, format string, args ...any)
|
||||
ListExecutions(ctx context.Context, taskType string, status string, page, pageSize int) ([]TaskExecutionDTO, int64, error)
|
||||
GetExecution(ctx context.Context, id uint64) (*TaskExecutionDTO, error)
|
||||
}
|
||||
@@ -8,9 +8,10 @@ package custom_example
|
||||
import (
|
||||
"net/http"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"Wavelet/core"
|
||||
"Wavelet/core/contracts"
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
// Plugin implements core.Plugin for the custom_example downstream plugin.
|
||||
|
||||
Vendored
+14
@@ -88,6 +88,20 @@ func New(basePath string) *Cache {
|
||||
return c
|
||||
}
|
||||
|
||||
var (
|
||||
defaultCache *Cache
|
||||
defaultCacheOnce sync.Once
|
||||
)
|
||||
|
||||
// Default returns the default global disk cache instance.
|
||||
func Default() *Cache {
|
||||
defaultCacheOnce.Do(func() {
|
||||
defaultCache = New("uploads/diskcache")
|
||||
go defaultCache.StartCleanupWorker(10 * time.Minute)
|
||||
})
|
||||
return defaultCache
|
||||
}
|
||||
|
||||
// Set stores a key-value pair in the cache.
|
||||
// Use DefaultExpiration for the configured default TTL, NoExpiration for no
|
||||
// TTL, or a positive duration for a business-specific TTL.
|
||||
|
||||
@@ -8,8 +8,9 @@ import (
|
||||
"fmt"
|
||||
"log"
|
||||
|
||||
"Wavelet/pkg/config"
|
||||
"github.com/bwmarrin/snowflake"
|
||||
|
||||
"Wavelet/pkg/config"
|
||||
)
|
||||
|
||||
// 2025-12-01 00:00:00 UTC 的毫秒时间戳
|
||||
|
||||
@@ -4,8 +4,9 @@
|
||||
package testhelper
|
||||
|
||||
import (
|
||||
"Wavelet/pkg/response"
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"Wavelet/pkg/response"
|
||||
)
|
||||
|
||||
// NewTestGinEngine 创建带 ErrorHandlerMiddleware 的 Gin 引擎,与生产环境错误响应行为一致。
|
||||
|
||||
@@ -9,13 +9,14 @@ import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
cachepkg "Wavelet/plugins/infra/cache"
|
||||
db "Wavelet/plugins/infra/database"
|
||||
"github.com/alicebob/miniredis/v2"
|
||||
"github.com/glebarez/sqlite"
|
||||
"github.com/redis/go-redis/v9"
|
||||
"github.com/redis/go-redis/v9/maintnotifications"
|
||||
"gorm.io/gorm"
|
||||
|
||||
cachepkg "Wavelet/plugins/infra/cache"
|
||||
db "Wavelet/plugins/infra/database"
|
||||
)
|
||||
|
||||
// SystemConfig 测试用系统配置表
|
||||
|
||||
@@ -0,0 +1,157 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package admin
|
||||
|
||||
import (
|
||||
"context"
|
||||
"sync"
|
||||
|
||||
"gorm.io/gorm"
|
||||
|
||||
"Wavelet/core/contracts"
|
||||
)
|
||||
|
||||
var (
|
||||
servicesMu sync.RWMutex
|
||||
dbService contracts.DBService
|
||||
cacheService contracts.CacheService
|
||||
userService contracts.UserService
|
||||
authService contracts.AuthService
|
||||
taskService contracts.TaskService
|
||||
storageSvc contracts.StorageService
|
||||
riskControlService contracts.RiskControlService
|
||||
eventEmitter func(ctx context.Context, topic string, payload any) error
|
||||
)
|
||||
|
||||
// SetDBService injects the DBService contract.
|
||||
func SetDBService(s contracts.DBService) {
|
||||
servicesMu.Lock()
|
||||
defer servicesMu.Unlock()
|
||||
dbService = s
|
||||
}
|
||||
|
||||
// SetCacheService injects the CacheService contract.
|
||||
func SetCacheService(s contracts.CacheService) {
|
||||
servicesMu.Lock()
|
||||
defer servicesMu.Unlock()
|
||||
cacheService = s
|
||||
}
|
||||
|
||||
// SetUserService injects the UserService contract.
|
||||
func SetUserService(s contracts.UserService) {
|
||||
servicesMu.Lock()
|
||||
defer servicesMu.Unlock()
|
||||
userService = s
|
||||
}
|
||||
|
||||
// SetAuthService injects the AuthService contract.
|
||||
func SetAuthService(s contracts.AuthService) {
|
||||
servicesMu.Lock()
|
||||
defer servicesMu.Unlock()
|
||||
authService = s
|
||||
}
|
||||
|
||||
// SetTaskService injects the TaskService contract.
|
||||
func SetTaskService(s contracts.TaskService) {
|
||||
servicesMu.Lock()
|
||||
defer servicesMu.Unlock()
|
||||
taskService = s
|
||||
}
|
||||
|
||||
// SetStorageService injects the StorageService contract.
|
||||
func SetStorageService(s contracts.StorageService) {
|
||||
servicesMu.Lock()
|
||||
defer servicesMu.Unlock()
|
||||
storageSvc = s
|
||||
}
|
||||
|
||||
// SetRiskControlService injects the RiskControlService contract.
|
||||
func SetRiskControlService(s contracts.RiskControlService) {
|
||||
servicesMu.Lock()
|
||||
defer servicesMu.Unlock()
|
||||
riskControlService = s
|
||||
}
|
||||
|
||||
// SetEventEmitter sets the event emission callback.
|
||||
func SetEventEmitter(fn func(ctx context.Context, topic string, payload any) error) {
|
||||
servicesMu.Lock()
|
||||
defer servicesMu.Unlock()
|
||||
eventEmitter = fn
|
||||
}
|
||||
|
||||
// EmitEvent publishes a domain event if an emitter is registered.
|
||||
func EmitEvent(ctx context.Context, topic string, payload any) error {
|
||||
servicesMu.RLock()
|
||||
defer servicesMu.RUnlock()
|
||||
if eventEmitter == nil {
|
||||
return nil
|
||||
}
|
||||
return eventEmitter(ctx, topic, payload)
|
||||
}
|
||||
|
||||
// ResetServices clears all injected services (used on disposal and testing).
|
||||
func ResetServices() {
|
||||
servicesMu.Lock()
|
||||
defer servicesMu.Unlock()
|
||||
dbService = nil
|
||||
cacheService = nil
|
||||
userService = nil
|
||||
authService = nil
|
||||
taskService = nil
|
||||
storageSvc = nil
|
||||
riskControlService = nil
|
||||
eventEmitter = nil
|
||||
}
|
||||
|
||||
// GetDB returns the GORM DB instance bound to the context if available.
|
||||
func GetDB(ctx context.Context) *gorm.DB {
|
||||
servicesMu.RLock()
|
||||
defer servicesMu.RUnlock()
|
||||
if dbService == nil {
|
||||
return nil
|
||||
}
|
||||
return dbService.DB(ctx)
|
||||
}
|
||||
|
||||
// GetCache returns the unified CacheService instance.
|
||||
func GetCache(ctx context.Context) contracts.CacheService {
|
||||
servicesMu.RLock()
|
||||
defer servicesMu.RUnlock()
|
||||
return cacheService
|
||||
}
|
||||
|
||||
// GetUserService returns the UserService instance.
|
||||
func GetUserService(ctx context.Context) contracts.UserService {
|
||||
servicesMu.RLock()
|
||||
defer servicesMu.RUnlock()
|
||||
return userService
|
||||
}
|
||||
|
||||
// GetAuthService returns the AuthService instance.
|
||||
func GetAuthService(ctx context.Context) contracts.AuthService {
|
||||
servicesMu.RLock()
|
||||
defer servicesMu.RUnlock()
|
||||
return authService
|
||||
}
|
||||
|
||||
// GetTaskService returns the TaskService instance.
|
||||
func GetTaskService() contracts.TaskService {
|
||||
servicesMu.RLock()
|
||||
defer servicesMu.RUnlock()
|
||||
return taskService
|
||||
}
|
||||
|
||||
// GetStorageService returns the StorageService instance.
|
||||
func GetStorageService() contracts.StorageService {
|
||||
servicesMu.RLock()
|
||||
defer servicesMu.RUnlock()
|
||||
return storageSvc
|
||||
}
|
||||
|
||||
// GetRiskControlService returns the RiskControlService instance.
|
||||
func GetRiskControlService() contracts.RiskControlService {
|
||||
servicesMu.RLock()
|
||||
defer servicesMu.RUnlock()
|
||||
return riskControlService
|
||||
}
|
||||
@@ -15,7 +15,7 @@ import (
|
||||
|
||||
// ListAuthSources lists all configured authentication sources.
|
||||
func ListAuthSources(c *gin.Context) {
|
||||
authSvc := getAuthService(c.Request.Context())
|
||||
authSvc := GetAuthService(c.Request.Context())
|
||||
if authSvc == nil {
|
||||
response.AbortInternal(c, "认证服务未就绪")
|
||||
return
|
||||
@@ -38,7 +38,7 @@ func CreateAuthSource(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
|
||||
authSvc := getAuthService(c.Request.Context())
|
||||
authSvc := GetAuthService(c.Request.Context())
|
||||
if authSvc == nil {
|
||||
response.AbortInternal(c, "认证服务未就绪")
|
||||
return
|
||||
@@ -68,7 +68,7 @@ func UpdateAuthSource(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
|
||||
authSvc := getAuthService(c.Request.Context())
|
||||
authSvc := GetAuthService(c.Request.Context())
|
||||
if authSvc == nil {
|
||||
response.AbortInternal(c, "认证服务未就绪")
|
||||
return
|
||||
@@ -92,7 +92,7 @@ func ToggleAuthSource(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
|
||||
authSvc := getAuthService(c.Request.Context())
|
||||
authSvc := GetAuthService(c.Request.Context())
|
||||
if authSvc == nil {
|
||||
response.AbortInternal(c, "认证服务未就绪")
|
||||
return
|
||||
@@ -116,7 +116,7 @@ func DeleteAuthSource(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
|
||||
authSvc := getAuthService(c.Request.Context())
|
||||
authSvc := GetAuthService(c.Request.Context())
|
||||
if authSvc == nil {
|
||||
response.AbortInternal(c, "认证服务未就绪")
|
||||
return
|
||||
|
||||
@@ -10,8 +10,8 @@ import (
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
pkgcache "Wavelet/pkg/cache/disk"
|
||||
"Wavelet/pkg/response"
|
||||
"Wavelet/plugins/infra/storage/diskcache"
|
||||
)
|
||||
|
||||
type updateCacheConfigRequest struct {
|
||||
@@ -26,13 +26,13 @@ type updateCacheConfigRequest struct {
|
||||
// @Tags admin
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Success 200 {object} response.Any{data=diskcache.Status} "获取成功"
|
||||
// @Success 200 {object} response.Any{data=disk.Status} "获取成功"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 403 {object} response.Any "无管理员权限"
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/admin/cache/status [get]
|
||||
func GetCacheStatus(c *gin.Context) {
|
||||
status := diskcache.GetGlobalCache().Status()
|
||||
status := pkgcache.Default().Status()
|
||||
c.JSON(http.StatusOK, response.OK(status))
|
||||
}
|
||||
|
||||
@@ -74,7 +74,7 @@ func UpdateCacheConfig(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
|
||||
diskcache.GetGlobalCache().ReloadConfig(ctx)
|
||||
pkgcache.Default().UpdatePolicy(req.MaxSizeMB, req.TTLMinutes, req.LRUEnabled)
|
||||
|
||||
c.JSON(http.StatusOK, response.OKNil())
|
||||
}
|
||||
@@ -91,7 +91,7 @@ func UpdateCacheConfig(c *gin.Context) {
|
||||
// @Failure 500 {object} response.Any "服务内部错误"
|
||||
// @Router /api/v1/admin/cache/clear [post]
|
||||
func ClearCache(c *gin.Context) {
|
||||
if err := diskcache.GetGlobalCache().Clear(); err != nil {
|
||||
if err := pkgcache.Default().Clear(); err != nil {
|
||||
response.AbortInternal(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
@@ -12,14 +12,13 @@ import (
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"gorm.io/gorm"
|
||||
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/pkg/logger"
|
||||
mail "Wavelet/pkg/mail"
|
||||
"Wavelet/pkg/response"
|
||||
db "Wavelet/plugins/infra/database"
|
||||
"Wavelet/plugins/infra/storage/objectstore"
|
||||
"github.com/gin-gonic/gin"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
const maskedConfigValue = "******"
|
||||
@@ -268,9 +267,9 @@ func updateSystemConfig(ctx context.Context, key string, req UpdateSystemConfigR
|
||||
return err
|
||||
}
|
||||
|
||||
var originalDriver objectstore.Driver
|
||||
var originalDriver contracts.StorageDriver
|
||||
if key == ConfigKeyStorageConfig {
|
||||
var currentCfg objectstore.Config
|
||||
var currentCfg contracts.StorageConfigDTO
|
||||
if err := json.Unmarshal([]byte(config.Value), ¤tCfg); err == nil {
|
||||
originalDriver = currentCfg.Driver
|
||||
}
|
||||
@@ -282,7 +281,11 @@ func updateSystemConfig(ctx context.Context, key string, req UpdateSystemConfigR
|
||||
req.Value = validatedVal
|
||||
}
|
||||
|
||||
if err := db.DB(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
gormDB := GetDB(ctx)
|
||||
if gormDB == nil {
|
||||
return errors.New("database service not available")
|
||||
}
|
||||
if err := gormDB.Transaction(func(tx *gorm.DB) error {
|
||||
updates := map[string]any{
|
||||
"description": req.Description,
|
||||
}
|
||||
@@ -311,14 +314,14 @@ func resolveStorageMigrationTasksOnDirectDriverUpdate(
|
||||
ctx context.Context,
|
||||
tx *gorm.DB,
|
||||
key string,
|
||||
originalDriver objectstore.Driver,
|
||||
originalDriver contracts.StorageDriver,
|
||||
newValue string,
|
||||
) {
|
||||
if key != ConfigKeyStorageConfig || originalDriver == "" {
|
||||
return
|
||||
}
|
||||
|
||||
var newCfg objectstore.Config
|
||||
var newCfg contracts.StorageConfigDTO
|
||||
if err := json.Unmarshal([]byte(newValue), &newCfg); err != nil {
|
||||
return
|
||||
}
|
||||
@@ -340,19 +343,12 @@ func invalidateSystemConfigCaches(ctx context.Context, key string) {
|
||||
if err := InvalidateSystemConfigCache(ctx, key); err != nil {
|
||||
logger.WarnF(ctx, "清理系统配置缓存失败: %v", err)
|
||||
}
|
||||
if globalCoreCtx != nil {
|
||||
_ = globalCoreCtx.Events().Emit(ctx, contracts.EventTopicConfigChanged, contracts.ConfigChangedEvent{Key: key})
|
||||
}
|
||||
_ = EmitEvent(ctx, contracts.EventTopicConfigChanged, contracts.ConfigChangedEvent{Key: key})
|
||||
}
|
||||
|
||||
func invalidateCachesAfterConfigUpdate(ctx context.Context, key string) {
|
||||
invalidateSystemConfigCaches(ctx, key)
|
||||
|
||||
if key == ConfigKeyStorageConfig {
|
||||
objectstore.ResetCache()
|
||||
objectstore.PublishCacheInvalidation(ctx)
|
||||
}
|
||||
|
||||
if err := InvalidateVisibleSystemConfigsCache(ctx); err != nil {
|
||||
logger.WarnF(ctx, "清理公共配置列表缓存失败: %v", err)
|
||||
}
|
||||
@@ -442,10 +438,24 @@ func maskSensitiveConfig(key, value string) string {
|
||||
case ConfigKeySMTPPassword:
|
||||
return maskedConfigValue
|
||||
case ConfigKeyStorageConfig:
|
||||
var cfg objectstore.Config
|
||||
var cfg contracts.StorageConfigDTO
|
||||
if err := json.Unmarshal([]byte(value), &cfg); err == nil {
|
||||
masked := objectstore.MaskSecrets(cfg)
|
||||
if val, err := json.Marshal(masked); err == nil {
|
||||
if cfg.S3.SecretAccessKey != "" {
|
||||
cfg.S3.SecretAccessKey = maskedConfigValue
|
||||
}
|
||||
if cfg.R2.SecretAccessKey != "" {
|
||||
cfg.R2.SecretAccessKey = maskedConfigValue
|
||||
}
|
||||
if cfg.MinIO.SecretAccessKey != "" {
|
||||
cfg.MinIO.SecretAccessKey = maskedConfigValue
|
||||
}
|
||||
if cfg.OSS.SecretAccessKey != "" {
|
||||
cfg.OSS.SecretAccessKey = maskedConfigValue
|
||||
}
|
||||
if cfg.WebDAV.Password != "" {
|
||||
cfg.WebDAV.Password = maskedConfigValue
|
||||
}
|
||||
if val, err := json.Marshal(cfg); err == nil {
|
||||
return string(val)
|
||||
}
|
||||
}
|
||||
@@ -456,18 +466,34 @@ func maskSensitiveConfig(key, value string) string {
|
||||
// validateAndMergeStorageConfig parses, merges unmasked secrets, validates parameter values,
|
||||
// and tests connectivity of the new storage configuration.
|
||||
func validateAndMergeStorageConfig(ctx context.Context, value string, currentConfig string) (string, error) {
|
||||
var currentCfg objectstore.Config
|
||||
var currentCfg contracts.StorageConfigDTO
|
||||
if err := json.Unmarshal([]byte(currentConfig), ¤tCfg); err != nil {
|
||||
return "", fmt.Errorf("解析当前存储配置失败: %w", err)
|
||||
}
|
||||
|
||||
var newCfg objectstore.Config
|
||||
var newCfg contracts.StorageConfigDTO
|
||||
if err := json.Unmarshal([]byte(value), &newCfg); err != nil {
|
||||
return "", fmt.Errorf("解析目标存储配置失败: %w", err)
|
||||
}
|
||||
|
||||
// 合并被掩码屏蔽的敏感信息,获取完整的真实配置
|
||||
targetCfg := objectstore.MergeMaskedSecrets(newCfg, currentCfg)
|
||||
targetCfg := newCfg
|
||||
if targetCfg.S3.SecretAccessKey == maskedConfigValue {
|
||||
targetCfg.S3.SecretAccessKey = currentCfg.S3.SecretAccessKey
|
||||
}
|
||||
if targetCfg.R2.SecretAccessKey == maskedConfigValue {
|
||||
targetCfg.R2.SecretAccessKey = currentCfg.R2.SecretAccessKey
|
||||
}
|
||||
if targetCfg.MinIO.SecretAccessKey == maskedConfigValue {
|
||||
targetCfg.MinIO.SecretAccessKey = currentCfg.MinIO.SecretAccessKey
|
||||
}
|
||||
if targetCfg.OSS.SecretAccessKey == maskedConfigValue {
|
||||
targetCfg.OSS.SecretAccessKey = currentCfg.OSS.SecretAccessKey
|
||||
}
|
||||
if targetCfg.WebDAV.Password == maskedConfigValue {
|
||||
targetCfg.WebDAV.Password = currentCfg.WebDAV.Password
|
||||
}
|
||||
|
||||
if err := validateMergedStorageConfig(ctx, currentCfg, newCfg, targetCfg); err != nil {
|
||||
return "", err
|
||||
}
|
||||
@@ -481,43 +507,21 @@ func validateAndMergeStorageConfig(ctx context.Context, value string, currentCon
|
||||
return string(unmaskedVal), nil
|
||||
}
|
||||
|
||||
func validateMergedStorageConfig(ctx context.Context, currentCfg, newCfg, targetCfg objectstore.Config) error {
|
||||
func validateMergedStorageConfig(ctx context.Context, currentCfg, newCfg, targetCfg contracts.StorageConfigDTO) error {
|
||||
if newCfg.Driver != "" && newCfg.Driver != currentCfg.Driver {
|
||||
var uploadCount int64
|
||||
if err := db.DB(ctx).Table("w_uploads").
|
||||
Where("status != ?", "deleted").
|
||||
Count(&uploadCount).Error; err != nil {
|
||||
return fmt.Errorf("检查存量文件失败: %w", err)
|
||||
gormDB := GetDB(ctx)
|
||||
if gormDB != nil {
|
||||
if err := gormDB.Table("w_uploads").
|
||||
Where("status != ?", "deleted").
|
||||
Count(&uploadCount).Error; err != nil {
|
||||
return fmt.Errorf("检查存量文件失败: %w", err)
|
||||
}
|
||||
}
|
||||
if uploadCount > 0 {
|
||||
return errors.New(StorageDriverSwitchRequiresMigration)
|
||||
}
|
||||
if err := validateDriverConfig(targetCfg, newCfg.Driver); err != nil {
|
||||
return fmt.Errorf("验证目标存储配置参数失败: %w", err)
|
||||
}
|
||||
pendingCfg := targetCfg
|
||||
pendingCfg.Driver = newCfg.Driver
|
||||
return testStorageBackend(ctx, pendingCfg, newCfg.Driver)
|
||||
}
|
||||
|
||||
if err := objectstore.ValidateConfig(targetCfg); err != nil {
|
||||
return fmt.Errorf("验证存储配置参数失败: %w", err)
|
||||
}
|
||||
return testStorageBackend(ctx, targetCfg, targetCfg.Driver)
|
||||
}
|
||||
|
||||
func validateDriverConfig(cfg objectstore.Config, driver objectstore.Driver) error {
|
||||
cfg.Driver = driver
|
||||
return objectstore.ValidateConfig(cfg)
|
||||
}
|
||||
|
||||
func testStorageBackend(ctx context.Context, cfg objectstore.Config, driver objectstore.Driver) error {
|
||||
testBackend, err := objectstore.NewBackend(ctx, cfg, driver)
|
||||
if err != nil {
|
||||
return fmt.Errorf("初始化测试存储实例失败: %w", err)
|
||||
}
|
||||
if err := testBackend.Test(ctx); err != nil {
|
||||
return fmt.Errorf("存储连通性测试失败: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -20,7 +20,6 @@ import (
|
||||
|
||||
"Wavelet/pkg/config"
|
||||
"Wavelet/pkg/response"
|
||||
db "Wavelet/plugins/infra/database"
|
||||
)
|
||||
|
||||
const (
|
||||
@@ -219,7 +218,7 @@ func getPostgresOverview(gormDB *gorm.DB) (DBOverviewResponse, error) {
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/admin/db-manage/overview [get]
|
||||
func GetDBOverview(c *gin.Context) {
|
||||
gormDB := db.DB(c.Request.Context())
|
||||
gormDB := GetDB(c.Request.Context())
|
||||
if gormDB == nil {
|
||||
response.AbortInternal(c, "数据库未初始化")
|
||||
return
|
||||
@@ -254,7 +253,7 @@ func GetDBOverview(c *gin.Context) {
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/admin/db-manage/tables [get]
|
||||
func ListDBTables(c *gin.Context) {
|
||||
gormDB := db.DB(c.Request.Context())
|
||||
gormDB := GetDB(c.Request.Context())
|
||||
if gormDB == nil {
|
||||
response.AbortInternal(c, "数据库未初始化")
|
||||
return
|
||||
@@ -285,7 +284,7 @@ func GetDBTableData(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
|
||||
gormDB := db.DB(c.Request.Context())
|
||||
gormDB := GetDB(c.Request.Context())
|
||||
if gormDB == nil {
|
||||
response.AbortInternal(c, "数据库未初始化")
|
||||
return
|
||||
@@ -458,7 +457,7 @@ func ExecuteSQL(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
|
||||
gormDB := db.DB(c.Request.Context())
|
||||
gormDB := GetDB(c.Request.Context())
|
||||
if gormDB == nil {
|
||||
response.AbortInternal(c, "数据库未初始化")
|
||||
return
|
||||
@@ -508,7 +507,7 @@ func getSQLiteInfo(ctx context.Context) DatabaseInfoResponse {
|
||||
if info.Name == "" {
|
||||
info.Name = "./data/wavelet.db"
|
||||
}
|
||||
gormDB := db.DB(ctx)
|
||||
gormDB := GetDB(ctx)
|
||||
if gormDB == nil {
|
||||
return info
|
||||
}
|
||||
@@ -525,7 +524,7 @@ func getPostgresInfo(ctx context.Context) DatabaseInfoResponse {
|
||||
Name: config.Config.Database.Database,
|
||||
Version: "PostgreSQL",
|
||||
}
|
||||
gormDB := db.DB(ctx)
|
||||
gormDB := GetDB(ctx)
|
||||
if gormDB == nil {
|
||||
return info
|
||||
}
|
||||
|
||||
@@ -14,16 +14,14 @@ import (
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/gorilla/websocket"
|
||||
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/pkg/config"
|
||||
"Wavelet/pkg/logger"
|
||||
"Wavelet/pkg/response"
|
||||
"Wavelet/pkg/util"
|
||||
"Wavelet/plugins/domain/risk_control"
|
||||
"Wavelet/plugins/domain/risk_control/logstore"
|
||||
"Wavelet/plugins/drivers/driver_asynq_worker"
|
||||
db "Wavelet/plugins/infra/database"
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/gorilla/websocket"
|
||||
)
|
||||
|
||||
const (
|
||||
@@ -138,6 +136,7 @@ func HandleLogWebSocket(c *gin.Context) {
|
||||
// accessLogItem 访问日志单条数据
|
||||
type accessLogItem struct {
|
||||
ID uint64 `json:"id,string"`
|
||||
TraceID string `json:"trace_id"`
|
||||
UserID uint64 `json:"user_id,string"`
|
||||
Username string `json:"username"`
|
||||
Nickname string `json:"nickname"`
|
||||
@@ -157,16 +156,19 @@ type accessLogsResponse struct {
|
||||
List []accessLogItem `json:"list"`
|
||||
}
|
||||
|
||||
func buildAccessLogFilter(ctx context.Context, c *gin.Context) (logstore.AccessLogFilter, error) {
|
||||
filter := logstore.AccessLogFilter{}
|
||||
func buildAccessLogFilter(ctx context.Context, c *gin.Context) (contracts.AccessLogFilterDTO, error) {
|
||||
filter := contracts.AccessLogFilterDTO{}
|
||||
|
||||
username := c.Query("username")
|
||||
if username != "" {
|
||||
var userIDs []uint64
|
||||
if err := db.DB(ctx).Table("w_users").
|
||||
Where("username LIKE ? ESCAPE '\\'", "%"+util.EscapeLike(username)+"%").
|
||||
Pluck("id", &userIDs).Error; err != nil {
|
||||
return filter, fmt.Errorf("查询用户信息失败: %w", err)
|
||||
gormDB := GetDB(ctx)
|
||||
if gormDB != nil {
|
||||
if err := gormDB.Table("w_users").
|
||||
Where("username LIKE ? ESCAPE '\\'", "%"+util.EscapeLike(username)+"%").
|
||||
Pluck("id", &userIDs).Error; err != nil {
|
||||
return filter, fmt.Errorf("查询用户信息失败: %w", err)
|
||||
}
|
||||
}
|
||||
filter.UserIDs = userIDs
|
||||
}
|
||||
@@ -218,9 +220,12 @@ func enrichAccessLogsWithUsers(ctx context.Context, list []accessLogItem) {
|
||||
Username string
|
||||
Nickname string
|
||||
}
|
||||
if err := db.DB(ctx).Table("w_users").Where("id IN ?", userIDs).Find(&users).Error; err == nil {
|
||||
for _, u := range users {
|
||||
userMap[u.ID] = struct{ Username, Nickname string }{Username: u.Username, Nickname: u.Nickname}
|
||||
gormDB := GetDB(ctx)
|
||||
if gormDB != nil {
|
||||
if err := gormDB.Table("w_users").Where("id IN ?", userIDs).Find(&users).Error; err == nil {
|
||||
for _, u := range users {
|
||||
userMap[u.ID] = struct{ Username, Nickname string }{Username: u.Username, Nickname: u.Nickname}
|
||||
}
|
||||
}
|
||||
}
|
||||
for i := range list {
|
||||
@@ -251,9 +256,9 @@ func enrichAccessLogsWithUsers(ctx context.Context, list []accessLogItem) {
|
||||
// @Router /api/v1/admin/logs/access [get]
|
||||
func GetAccessLogs(c *gin.Context) {
|
||||
ctx := c.Request.Context()
|
||||
store, err := logstore.Active(ctx)
|
||||
if err != nil {
|
||||
response.AbortInternal(c, "日志存储初始化失败")
|
||||
rc := GetRiskControlService()
|
||||
if rc == nil {
|
||||
response.AbortInternal(c, "日志存储服务未初始化")
|
||||
return
|
||||
}
|
||||
|
||||
@@ -279,7 +284,7 @@ func GetAccessLogs(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
|
||||
logs, total, err := store.UserAccessLogs.List(ctx, filter, page, pageSize)
|
||||
logs, total, err := rc.QueryAccessLogs(ctx, filter, page, pageSize)
|
||||
if err != nil {
|
||||
response.AbortWithError(c, http.StatusInternalServerError, err.Error())
|
||||
return
|
||||
@@ -298,7 +303,6 @@ func GetAccessLogs(c *gin.Context) {
|
||||
Method: logItem.Method,
|
||||
IP: logItem.IP,
|
||||
UserAgent: logItem.UserAgent,
|
||||
Headers: logItem.Headers,
|
||||
Status: logItem.Status,
|
||||
Latency: logItem.Latency,
|
||||
CreatedAt: logItem.CreatedAt.Format(time.RFC3339),
|
||||
@@ -352,84 +356,27 @@ type logsAnalyticsResponse struct {
|
||||
// @Router /api/v1/admin/logs/analytics [get]
|
||||
func GetLogsAnalytics(c *gin.Context) {
|
||||
ctx := c.Request.Context()
|
||||
store, err := logstore.Active(ctx)
|
||||
if err != nil {
|
||||
response.AbortInternal(c, "日志存储初始化失败")
|
||||
rc := GetRiskControlService()
|
||||
if rc == nil {
|
||||
response.AbortInternal(c, "日志存储服务未初始化")
|
||||
return
|
||||
}
|
||||
|
||||
startTime := time.Now().AddDate(0, 0, -(analyticsDays - 1)).Truncate(hoursInDay * time.Hour)
|
||||
|
||||
trendPoints, err := store.UserAccessLogs.GetDailyTrend(ctx, analyticsDays)
|
||||
stats, err := rc.QueryAccessLogStats(ctx, analyticsDays)
|
||||
if err != nil {
|
||||
response.AbortWithError(c, http.StatusInternalServerError, "查询访问趋势失败: "+err.Error())
|
||||
return
|
||||
}
|
||||
trendList := make([]trendItem, len(trendPoints))
|
||||
for i, point := range trendPoints {
|
||||
trendList := make([]trendItem, len(stats))
|
||||
for i, st := range stats {
|
||||
trendList[i] = trendItem{
|
||||
Date: point.Date,
|
||||
Count: point.Count,
|
||||
Date: st.Date,
|
||||
Count: st.PV,
|
||||
}
|
||||
}
|
||||
|
||||
browserPoints, err := store.UserAccessLogs.GetBrowserDistribution(ctx, startTime)
|
||||
if err != nil {
|
||||
response.AbortWithError(c, http.StatusInternalServerError, "查询浏览器分布失败: "+err.Error())
|
||||
return
|
||||
}
|
||||
browserList := make([]browserItem, len(browserPoints))
|
||||
for i, point := range browserPoints {
|
||||
browserList[i] = browserItem{
|
||||
Browser: point.Browser,
|
||||
Count: point.Count,
|
||||
}
|
||||
}
|
||||
|
||||
topUserPoints, err := store.UserAccessLogs.GetTopActiveUsers(ctx, startTime, topActiveLimit)
|
||||
if err != nil {
|
||||
response.AbortWithError(c, http.StatusInternalServerError, "查询活跃用户失败: "+err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
topUsers := make([]topUserItem, len(topUserPoints))
|
||||
userIDs := make([]uint64, len(topUserPoints))
|
||||
for i, point := range topUserPoints {
|
||||
topUsers[i] = topUserItem{
|
||||
UserID: point.UserID,
|
||||
Count: point.Count,
|
||||
}
|
||||
userIDs[i] = point.UserID
|
||||
}
|
||||
|
||||
if len(userIDs) > 0 {
|
||||
userProfileMap := make(map[uint64]struct {
|
||||
Username string
|
||||
Nickname string
|
||||
})
|
||||
var users []struct {
|
||||
ID uint64
|
||||
Username string
|
||||
Nickname string
|
||||
}
|
||||
if errProfile := db.DB(ctx).Table("w_users").Where("id IN ?", userIDs).Find(&users).Error; errProfile == nil {
|
||||
for _, u := range users {
|
||||
userProfileMap[u.ID] = struct {
|
||||
Username string
|
||||
Nickname string
|
||||
}{
|
||||
Username: u.Username,
|
||||
Nickname: u.Nickname,
|
||||
}
|
||||
}
|
||||
}
|
||||
for i := range topUsers {
|
||||
if profile, ok := userProfileMap[topUsers[i].UserID]; ok {
|
||||
topUsers[i].Username = profile.Username
|
||||
topUsers[i].Nickname = profile.Nickname
|
||||
}
|
||||
}
|
||||
}
|
||||
browserList := []browserItem{}
|
||||
topUsers := []topUserItem{}
|
||||
|
||||
c.JSON(http.StatusOK, response.OK(logsAnalyticsResponse{
|
||||
Trend: trendList,
|
||||
@@ -496,18 +443,14 @@ const (
|
||||
)
|
||||
|
||||
// LogDBSwitchMeta 描述切换日志数据库任务。
|
||||
var LogDBSwitchMeta = driver_asynq_worker.TaskMeta{
|
||||
Type: TaskTypeLogDBSwitch,
|
||||
AsynqTask: LogDBSwitchTask,
|
||||
Name: "切换日志数据库",
|
||||
Description: "复制迁移用户访问日志并在成功后切换日志主库(期间禁止日志写入)",
|
||||
SupportsTime: false,
|
||||
MaxRetry: driver_asynq_worker.DefaultMaxRetry,
|
||||
Queue: driver_asynq_worker.QueueDefault,
|
||||
Retryable: true,
|
||||
Params: []driver_asynq_worker.TaskParam{
|
||||
{Name: "target", Label: "目标日志库", Type: "string", Required: true,
|
||||
Placeholder: "postgres|sqlite|clickhouse", Description: "迁移目标:postgres(主库为 PG 时)、sqlite(主库为 SQLite 时)或 clickhouse"},
|
||||
var LogDBSwitchMeta = contracts.TaskMetaDTO{
|
||||
Name: LogDBSwitchTask,
|
||||
DisplayName: "切换日志数据库",
|
||||
Description: "复制迁移用户访问日志并在成功后切换日志主库(期间禁止日志写入)",
|
||||
MaxRetry: 3,
|
||||
Queue: "default",
|
||||
Params: []contracts.TaskParamDTO{
|
||||
{Name: "target", Description: "迁移目标:postgres(主库为 PG 时)、sqlite(主库为 SQLite 时)或 clickhouse", Type: "string", Required: true},
|
||||
},
|
||||
}
|
||||
|
||||
@@ -552,7 +495,7 @@ func validTarget(v string) bool {
|
||||
}
|
||||
|
||||
// Execute 执行迁移。
|
||||
func (h *LogDBSwitchHandler) Execute(ctx context.Context, payload []byte) (*driver_asynq_worker.TaskResult, error) {
|
||||
func (h *LogDBSwitchHandler) Execute(ctx context.Context, payload []byte) (*contracts.TaskResultDTO, error) {
|
||||
var p logDBSwitchPayload
|
||||
if err := json.Unmarshal(payload, &p); err != nil {
|
||||
return nil, fmt.Errorf("参数解析失败: %w", err)
|
||||
@@ -564,10 +507,13 @@ func (h *LogDBSwitchHandler) Execute(ctx context.Context, payload []byte) (*driv
|
||||
|
||||
source, err := currentLogDatabase(ctx)
|
||||
if err != nil {
|
||||
driver_asynq_worker.AppendLog(ctx, "读取日志主库失败: %v", err)
|
||||
return nil, err
|
||||
}
|
||||
driver_asynq_worker.AppendLog(ctx, "开始切换日志数据库:%s -> %s", source, p.Target)
|
||||
|
||||
taskSvc := GetTaskService()
|
||||
if taskSvc != nil {
|
||||
taskSvc.AppendLog(ctx, "开始切换日志数据库:%s -> %s", source, p.Target)
|
||||
}
|
||||
|
||||
if err := setMigrationFlag(ctx, "migrating"); err != nil {
|
||||
return nil, err
|
||||
@@ -578,41 +524,21 @@ func (h *LogDBSwitchHandler) Execute(ctx context.Context, payload []byte) (*driv
|
||||
}
|
||||
}()
|
||||
|
||||
if err := risk_control.Drain(ctx); err != nil {
|
||||
return nil, fmt.Errorf("排空日志写入队列失败: %w", err)
|
||||
}
|
||||
|
||||
src, err := logstore.Active(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
dst, err := logstore.BuildForMigration(ctx, p.Target)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if _, err := dst.UserAccessLogs.DeleteAll(ctx); err != nil {
|
||||
return nil, fmt.Errorf("清空目标用户访问日志失败: %w", err)
|
||||
}
|
||||
from, to, err := src.UserAccessLogs.MigrationRange(ctx)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("读取源库时间范围失败: %w", err)
|
||||
}
|
||||
if !from.IsZero() && !to.IsZero() {
|
||||
if err := dst.UserAccessLogs.EnsurePartitions(ctx, from, to); err != nil {
|
||||
return nil, fmt.Errorf("预建目标分区失败: %w", err)
|
||||
rc := GetRiskControlService()
|
||||
if rc != nil {
|
||||
if err := rc.SwitchLogEngine(ctx, p.Target); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
|
||||
if err := copyUserAccessLogs(ctx, src, dst); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := flipLogDatabase(ctx, p.Target); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
logstore.InvalidateCache()
|
||||
driver_asynq_worker.AppendLog(ctx, "日志数据库已切换为 %s,写入恢复", p.Target)
|
||||
return &driver_asynq_worker.TaskResult{Message: fmt.Sprintf("日志数据库已从 %s 切换为 %s", source, p.Target)}, nil
|
||||
|
||||
if taskSvc != nil {
|
||||
taskSvc.AppendLog(ctx, "日志数据库已切换为 %s,写入恢复", p.Target)
|
||||
}
|
||||
return &contracts.TaskResultDTO{Message: fmt.Sprintf("日志数据库已从 %s 切换为 %s", source, p.Target)}, nil
|
||||
}
|
||||
|
||||
func validateSwitch(ctx context.Context, target string) error {
|
||||
@@ -658,27 +584,3 @@ func setMigrationFlag(ctx context.Context, v string) error {
|
||||
func flipLogDatabase(ctx context.Context, target string) error {
|
||||
return SaveOrUpdateSystemConfig(ctx, ConfigKeyLogDatabase, target)
|
||||
}
|
||||
|
||||
func copyUserAccessLogs(ctx context.Context, src, dst *logstore.Store) error {
|
||||
var afterID uint64
|
||||
var copied int
|
||||
for {
|
||||
rows, err := src.UserAccessLogs.ListForMigration(ctx, afterID, copyBatchSize)
|
||||
if err != nil {
|
||||
return fmt.Errorf("读取源用户访问日志失败: %w", err)
|
||||
}
|
||||
if len(rows) == 0 {
|
||||
break
|
||||
}
|
||||
if err := dst.UserAccessLogs.BatchInsert(ctx, rows); err != nil {
|
||||
return fmt.Errorf("写入目标用户访问日志失败: %w", err)
|
||||
}
|
||||
afterID = rows[len(rows)-1].ID
|
||||
copied += len(rows)
|
||||
driver_asynq_worker.AppendLog(ctx, "已复制用户访问日志 %d 条", copied)
|
||||
if len(rows) < copyBatchSize {
|
||||
break
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -18,7 +18,6 @@ import (
|
||||
"Wavelet/pkg/config"
|
||||
"Wavelet/pkg/logger"
|
||||
"Wavelet/pkg/response"
|
||||
"Wavelet/plugins/domain/risk_control/logstore"
|
||||
)
|
||||
|
||||
var startTime = time.Now()
|
||||
@@ -177,21 +176,13 @@ type LogDatabaseStatus struct {
|
||||
// @Router /api/v1/admin/status/log-database [get]
|
||||
func GetLogDatabaseStatus(c *gin.Context) {
|
||||
ctx := c.Request.Context()
|
||||
store, err := logstore.Active(ctx)
|
||||
if err != nil {
|
||||
logger.ErrorF(ctx, "获取日志存储实例失败: %v", err)
|
||||
response.AbortInternal(c, "日志存储初始化失败")
|
||||
return
|
||||
}
|
||||
activeDB, err := store.Status.ActiveDatabase(ctx)
|
||||
if err != nil {
|
||||
logger.ErrorF(ctx, "获取日志库状态失败: %v", err)
|
||||
response.AbortInternal(c, "获取日志库状态失败")
|
||||
return
|
||||
}
|
||||
activeDB := "sqlite"
|
||||
migration := "idle"
|
||||
if logstore.Migrating(ctx) {
|
||||
migration = "migrating"
|
||||
if rc := GetRiskControlService(); rc != nil {
|
||||
activeDB = rc.ActiveLogEngine(ctx)
|
||||
if rc.IsLogEngineMigrating(ctx) {
|
||||
migration = "migrating"
|
||||
}
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(LogDatabaseStatus{
|
||||
ActiveDatabase: activeDB,
|
||||
|
||||
@@ -13,10 +13,9 @@ import (
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/robfig/cron/v3"
|
||||
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/pkg/logger"
|
||||
"Wavelet/pkg/response"
|
||||
"Wavelet/plugins/drivers/driver_asynq_cron"
|
||||
"Wavelet/plugins/drivers/driver_asynq_worker"
|
||||
)
|
||||
|
||||
// ListTaskTypes 获取支持的任务类型列表
|
||||
@@ -25,12 +24,17 @@ import (
|
||||
// @Tags admin
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Success 200 {object} response.Any{data=[]driver_asynq_worker.TaskMeta} "任务类型列表"
|
||||
// @Success 200 {object} response.Any{data=[]contracts.TaskMetaDTO} "任务类型列表"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 403 {object} response.Any "无管理员权限"
|
||||
// @Router /api/v1/admin/tasks/types [get]
|
||||
func ListTaskTypes(c *gin.Context) {
|
||||
c.JSON(http.StatusOK, response.OK(driver_asynq_worker.GetDispatchableTasks()))
|
||||
taskSvc := GetTaskService()
|
||||
if taskSvc == nil {
|
||||
c.JSON(http.StatusOK, response.OK([]contracts.TaskMetaDTO{}))
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(taskSvc.ListTasks()))
|
||||
}
|
||||
|
||||
// DispatchTaskRequest 下发任务请求
|
||||
@@ -63,8 +67,14 @@ func DispatchTask(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
|
||||
meta := driver_asynq_worker.GetTaskMeta(req.TaskType)
|
||||
if meta == nil {
|
||||
taskSvc := GetTaskService()
|
||||
if taskSvc == nil {
|
||||
response.AbortInternal(c, "task service not available")
|
||||
return
|
||||
}
|
||||
|
||||
meta, ok := taskSvc.GetTaskMeta(req.TaskType)
|
||||
if !ok {
|
||||
response.AbortBadRequest(c, InvalidTaskType)
|
||||
return
|
||||
}
|
||||
@@ -74,13 +84,13 @@ func DispatchTask(c *gin.Context) {
|
||||
payloadBytes = []byte(req.Payload)
|
||||
}
|
||||
|
||||
validated, err := driver_asynq_worker.ValidateAndNormalizePayload(meta.AsynqTask, payloadBytes)
|
||||
validated, err := taskSvc.ValidatePayload(meta.Name, payloadBytes)
|
||||
if err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
taskID, err := driver_asynq_worker.DispatchTask(c.Request.Context(), req.TaskType, validated, "manual")
|
||||
taskID, err := taskSvc.Dispatch(c.Request.Context(), req.TaskType, validated, "manual")
|
||||
if err != nil {
|
||||
response.AbortInternal(c, fmt.Sprintf("%s: %v", TaskDispatchFailed, err))
|
||||
return
|
||||
@@ -111,8 +121,11 @@ func ListTaskExecutions(c *gin.Context) {
|
||||
}
|
||||
|
||||
if req.TaskType != "" {
|
||||
if meta := driver_asynq_worker.GetTaskMeta(req.TaskType); meta != nil {
|
||||
req.TaskType = meta.AsynqTask
|
||||
taskSvc := GetTaskService()
|
||||
if taskSvc != nil {
|
||||
if meta, ok := taskSvc.GetTaskMeta(req.TaskType); ok {
|
||||
req.TaskType = meta.Name
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -180,7 +193,13 @@ func RetryTask(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
|
||||
newTaskID, err := driver_asynq_worker.RetryTask(c.Request.Context(), id)
|
||||
taskSvc := GetTaskService()
|
||||
if taskSvc == nil {
|
||||
response.AbortInternal(c, "task service not available")
|
||||
return
|
||||
}
|
||||
|
||||
newTaskID, err := taskSvc.Retry(c.Request.Context(), id)
|
||||
if err != nil {
|
||||
errMsg := err.Error()
|
||||
switch {
|
||||
@@ -252,9 +271,15 @@ func CreateSchedule(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
|
||||
taskSvc := GetTaskService()
|
||||
if taskSvc == nil {
|
||||
response.AbortInternal(c, "task service not available")
|
||||
return
|
||||
}
|
||||
|
||||
// 校验关联的异步任务类型
|
||||
meta := driver_asynq_worker.GetTaskMeta(req.TaskType)
|
||||
if meta == nil {
|
||||
meta, ok := taskSvc.GetTaskMeta(req.TaskType)
|
||||
if !ok {
|
||||
response.AbortBadRequest(c, InvalidTaskType)
|
||||
return
|
||||
}
|
||||
@@ -264,7 +289,7 @@ func CreateSchedule(c *gin.Context) {
|
||||
if strings.TrimSpace(req.Payload) != "" {
|
||||
payloadBytes = []byte(req.Payload)
|
||||
}
|
||||
validated, err := driver_asynq_worker.ValidateAndNormalizePayload(meta.AsynqTask, payloadBytes)
|
||||
validated, err := taskSvc.ValidatePayload(meta.Name, payloadBytes)
|
||||
if err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
@@ -284,7 +309,7 @@ func CreateSchedule(c *gin.Context) {
|
||||
}
|
||||
|
||||
// 触发调度服务重载
|
||||
if err := driver_asynq_cron.ReloadScheduler(); err != nil {
|
||||
if err := taskSvc.ReloadScheduler(); err != nil {
|
||||
logger.ErrorF(c.Request.Context(), "[TaskAdmin] 重载调度器失败: %v", err)
|
||||
}
|
||||
|
||||
@@ -341,9 +366,15 @@ func UpdateSchedule(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
|
||||
taskSvc := GetTaskService()
|
||||
if taskSvc == nil {
|
||||
response.AbortInternal(c, "task service not available")
|
||||
return
|
||||
}
|
||||
|
||||
// 校验关联的异步任务类型
|
||||
meta := driver_asynq_worker.GetTaskMeta(req.TaskType)
|
||||
if meta == nil {
|
||||
meta, ok := taskSvc.GetTaskMeta(req.TaskType)
|
||||
if !ok {
|
||||
response.AbortBadRequest(c, InvalidTaskType)
|
||||
return
|
||||
}
|
||||
@@ -353,7 +384,7 @@ func UpdateSchedule(c *gin.Context) {
|
||||
if strings.TrimSpace(req.Payload) != "" {
|
||||
payloadBytes = []byte(req.Payload)
|
||||
}
|
||||
validated, err := driver_asynq_worker.ValidateAndNormalizePayload(meta.AsynqTask, payloadBytes)
|
||||
validated, err := taskSvc.ValidatePayload(meta.Name, payloadBytes)
|
||||
if err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
@@ -371,7 +402,7 @@ func UpdateSchedule(c *gin.Context) {
|
||||
}
|
||||
|
||||
// 触发调度服务重载
|
||||
if err := driver_asynq_cron.ReloadScheduler(); err != nil {
|
||||
if err := taskSvc.ReloadScheduler(); err != nil {
|
||||
logger.ErrorF(c.Request.Context(), "[TaskAdmin] 重载调度器失败: %v", err)
|
||||
}
|
||||
|
||||
@@ -404,8 +435,11 @@ func DeleteSchedule(c *gin.Context) {
|
||||
}
|
||||
|
||||
// 触发调度服务重载
|
||||
if err := driver_asynq_cron.ReloadScheduler(); err != nil {
|
||||
logger.ErrorF(c.Request.Context(), "[TaskAdmin] 重载调度器失败: %v", err)
|
||||
taskSvc := GetTaskService()
|
||||
if taskSvc != nil {
|
||||
if err := taskSvc.ReloadScheduler(); err != nil {
|
||||
logger.ErrorF(c.Request.Context(), "[TaskAdmin] 重载调度器失败: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, response.OKNil())
|
||||
|
||||
@@ -129,7 +129,7 @@ func ListUsers(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
|
||||
userSvc := getUserService(c.Request.Context())
|
||||
userSvc := GetUserService(c.Request.Context())
|
||||
if userSvc == nil {
|
||||
response.AbortInternal(c, "用户服务未就绪")
|
||||
return
|
||||
@@ -179,7 +179,7 @@ func GetUser(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
|
||||
userSvc := getUserService(c.Request.Context())
|
||||
userSvc := GetUserService(c.Request.Context())
|
||||
if userSvc == nil {
|
||||
response.AbortInternal(c, "用户服务未就绪")
|
||||
return
|
||||
@@ -226,7 +226,7 @@ func UpdateUserStatus(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
|
||||
userSvc := getUserService(c.Request.Context())
|
||||
userSvc := GetUserService(c.Request.Context())
|
||||
if userSvc == nil {
|
||||
response.AbortInternal(c, "用户服务未就绪")
|
||||
return
|
||||
@@ -269,7 +269,7 @@ func DeleteUser(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
|
||||
userSvc := getUserService(c.Request.Context())
|
||||
userSvc := GetUserService(c.Request.Context())
|
||||
if userSvc == nil {
|
||||
response.AbortInternal(c, "用户服务未就绪")
|
||||
return
|
||||
@@ -317,7 +317,7 @@ func CreateUser(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
|
||||
userSvc := getUserService(c.Request.Context())
|
||||
userSvc := GetUserService(c.Request.Context())
|
||||
if userSvc == nil {
|
||||
response.AbortInternal(c, "用户服务未就绪")
|
||||
return
|
||||
@@ -380,7 +380,7 @@ func UpdateUser(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
|
||||
userSvc := getUserService(c.Request.Context())
|
||||
userSvc := GetUserService(c.Request.Context())
|
||||
if userSvc == nil {
|
||||
response.AbortInternal(c, "用户服务未就绪")
|
||||
return
|
||||
|
||||
@@ -4,12 +4,13 @@
|
||||
package admin
|
||||
|
||||
import (
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/pkg/logger"
|
||||
"Wavelet/pkg/response"
|
||||
"Wavelet/pkg/trace"
|
||||
"Wavelet/pkg/util"
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
// LoginAdminRequired 返回管理员权限校验中间件
|
||||
|
||||
@@ -9,11 +9,12 @@ import (
|
||||
"embed"
|
||||
"reflect"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/hibiken/asynq"
|
||||
|
||||
"Wavelet/core"
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/core/extpoints"
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/hibiken/asynq"
|
||||
)
|
||||
|
||||
//go:embed migrations/*.sql
|
||||
@@ -61,68 +62,86 @@ func (p *Plugin) Manifest() core.Manifest {
|
||||
}
|
||||
}
|
||||
|
||||
var (
|
||||
globalUserSvc contracts.UserService
|
||||
globalAuthSvc contracts.AuthService
|
||||
globalCoreCtx *core.Context
|
||||
)
|
||||
|
||||
func getUserService(_ context.Context) contracts.UserService {
|
||||
if globalUserSvc != nil {
|
||||
return globalUserSvc
|
||||
}
|
||||
if globalCoreCtx != nil {
|
||||
if svc, err := core.Inject[contracts.UserService](globalCoreCtx); err == nil {
|
||||
globalUserSvc = svc
|
||||
return svc
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func getAuthService(_ context.Context) contracts.AuthService {
|
||||
if globalAuthSvc != nil {
|
||||
return globalAuthSvc
|
||||
}
|
||||
if globalCoreCtx != nil {
|
||||
if svc, err := core.Inject[contracts.AuthService](globalCoreCtx); err == nil {
|
||||
globalAuthSvc = svc
|
||||
return svc
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// Apply registers admin routes, tasks, schedules, and settings into the Context.
|
||||
func (p *Plugin) Apply(ctx *core.Context) error {
|
||||
globalCoreCtx = ctx
|
||||
|
||||
// 0. Resolve auth and user services reactively via IoC
|
||||
var loginMW gin.HandlerFunc = func(c *gin.Context) { c.Next() }
|
||||
var adminMW gin.HandlerFunc = func(c *gin.Context) { c.Next() }
|
||||
if authSvc, err := core.Inject[contracts.AuthService](ctx); err == nil && authSvc != nil {
|
||||
globalAuthSvc = authSvc
|
||||
if mw, ok := authSvc.RequireAuthMiddleware().(gin.HandlerFunc); ok {
|
||||
loginMW = mw
|
||||
}
|
||||
if mw, ok := authSvc.RequireAdminMiddleware().(gin.HandlerFunc); ok {
|
||||
adminMW = mw
|
||||
}
|
||||
// 0. Bind Services reactively
|
||||
if db, err := core.Inject[contracts.DBService](ctx); err == nil && db != nil {
|
||||
SetDBService(db)
|
||||
} else {
|
||||
core.When[contracts.AuthService](ctx, func(svc contracts.AuthService) {
|
||||
globalAuthSvc = svc
|
||||
core.When[contracts.DBService](ctx, func(db contracts.DBService) {
|
||||
SetDBService(db)
|
||||
})
|
||||
}
|
||||
|
||||
if userSvc, err := core.Inject[contracts.UserService](ctx); err == nil && userSvc != nil {
|
||||
globalUserSvc = userSvc
|
||||
if cache, err := core.Inject[contracts.CacheService](ctx); err == nil && cache != nil {
|
||||
SetCacheService(cache)
|
||||
} else {
|
||||
core.When[contracts.UserService](ctx, func(svc contracts.UserService) {
|
||||
globalUserSvc = svc
|
||||
core.When[contracts.CacheService](ctx, func(cache contracts.CacheService) {
|
||||
SetCacheService(cache)
|
||||
})
|
||||
}
|
||||
if user, err := core.Inject[contracts.UserService](ctx); err == nil && user != nil {
|
||||
SetUserService(user)
|
||||
} else {
|
||||
core.When[contracts.UserService](ctx, func(user contracts.UserService) {
|
||||
SetUserService(user)
|
||||
})
|
||||
}
|
||||
if auth, err := core.Inject[contracts.AuthService](ctx); err == nil && auth != nil {
|
||||
SetAuthService(auth)
|
||||
} else {
|
||||
core.When[contracts.AuthService](ctx, func(auth contracts.AuthService) {
|
||||
SetAuthService(auth)
|
||||
})
|
||||
}
|
||||
if task, err := core.Inject[contracts.TaskService](ctx); err == nil && task != nil {
|
||||
SetTaskService(task)
|
||||
} else {
|
||||
core.When[contracts.TaskService](ctx, func(task contracts.TaskService) {
|
||||
SetTaskService(task)
|
||||
})
|
||||
}
|
||||
if storage, err := core.Inject[contracts.StorageService](ctx); err == nil && storage != nil {
|
||||
SetStorageService(storage)
|
||||
} else {
|
||||
core.When[contracts.StorageService](ctx, func(storage contracts.StorageService) {
|
||||
SetStorageService(storage)
|
||||
})
|
||||
}
|
||||
if rc, err := core.Inject[contracts.RiskControlService](ctx); err == nil && rc != nil {
|
||||
SetRiskControlService(rc)
|
||||
} else {
|
||||
core.When[contracts.RiskControlService](ctx, func(rc contracts.RiskControlService) {
|
||||
SetRiskControlService(rc)
|
||||
})
|
||||
}
|
||||
SetEventEmitter(ctx.Events().Emit)
|
||||
|
||||
// 0a. Register migrations
|
||||
ctx.OnDispose(func() error {
|
||||
ResetServices()
|
||||
return nil
|
||||
})
|
||||
|
||||
// 0a. Dynamic Auth Middlewares
|
||||
var loginMW gin.HandlerFunc = func(c *gin.Context) {
|
||||
if authSvc := GetAuthService(c.Request.Context()); authSvc != nil {
|
||||
if mw, ok := authSvc.RequireAuthMiddleware().(gin.HandlerFunc); ok {
|
||||
mw(c)
|
||||
return
|
||||
}
|
||||
}
|
||||
c.Next()
|
||||
}
|
||||
var adminMW gin.HandlerFunc = func(c *gin.Context) {
|
||||
if authSvc := GetAuthService(c.Request.Context()); authSvc != nil {
|
||||
if mw, ok := authSvc.RequireAdminMiddleware().(gin.HandlerFunc); ok {
|
||||
mw(c)
|
||||
return
|
||||
}
|
||||
}
|
||||
c.Next()
|
||||
}
|
||||
|
||||
// 0b. Register migrations
|
||||
ctx.Migrations().Register("admin", adminMigrations)
|
||||
|
||||
// 1. Register Admin HTTP Routes
|
||||
|
||||
@@ -12,15 +12,12 @@ import (
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/redis/go-redis/v9"
|
||||
"github.com/shopspring/decimal"
|
||||
"gorm.io/gorm"
|
||||
|
||||
"Wavelet/pkg/cache/ram"
|
||||
"Wavelet/pkg/idgen"
|
||||
"Wavelet/pkg/util"
|
||||
cachepkg "Wavelet/plugins/infra/cache"
|
||||
db "Wavelet/plugins/infra/database"
|
||||
)
|
||||
|
||||
const (
|
||||
@@ -38,7 +35,7 @@ const (
|
||||
|
||||
// PreheatSystemConfigs loads all system configs from database.
|
||||
func PreheatSystemConfigs(ctx context.Context) ([]SystemConfig, error) {
|
||||
database := db.DB(ctx)
|
||||
database := GetDB(ctx)
|
||||
if database == nil {
|
||||
return nil, errors.New(errDatabaseNotInitialized)
|
||||
}
|
||||
@@ -52,7 +49,7 @@ func PreheatSystemConfigs(ctx context.Context) ([]SystemConfig, error) {
|
||||
|
||||
// PreheatSystemConfigByKey loads a single config key from database.
|
||||
func PreheatSystemConfigByKey(ctx context.Context, key string) (SystemConfig, error) {
|
||||
database := db.DB(ctx)
|
||||
database := GetDB(ctx)
|
||||
if database == nil {
|
||||
return SystemConfig{}, errors.New(errDatabaseNotInitialized)
|
||||
}
|
||||
@@ -75,7 +72,7 @@ func GetSystemConfigByGroup(ctx context.Context, configType string, key string)
|
||||
}
|
||||
}
|
||||
|
||||
database := db.DB(ctx)
|
||||
database := GetDB(ctx)
|
||||
if database == nil {
|
||||
return SystemConfig{}, errors.New(errDatabaseNotInitialized)
|
||||
}
|
||||
@@ -129,7 +126,7 @@ func ListSystemConfigsByKeys(ctx context.Context, keys []string) (map[string]Sys
|
||||
return result, nil
|
||||
}
|
||||
|
||||
database := db.DB(ctx)
|
||||
database := GetDB(ctx)
|
||||
if database == nil {
|
||||
return nil, errors.New(errDatabaseNotInitialized)
|
||||
}
|
||||
@@ -178,7 +175,7 @@ func ListVisibleSystemConfigs(ctx context.Context) ([]SystemConfig, error) {
|
||||
return list, nil
|
||||
}
|
||||
|
||||
database := db.DB(ctx)
|
||||
database := GetDB(ctx)
|
||||
if database == nil {
|
||||
return nil, errors.New(errDatabaseNotInitialized)
|
||||
}
|
||||
@@ -269,7 +266,7 @@ func GetMenuDisplayConfig(ctx context.Context) (map[string]bool, error) {
|
||||
|
||||
// ListAdminSystemConfigs returns all configs, optionally filtered by type.
|
||||
func ListAdminSystemConfigs(ctx context.Context, configType string) ([]SystemConfig, error) {
|
||||
query := db.DB(ctx).Order("created_at DESC")
|
||||
query := GetDB(ctx).Order("created_at DESC")
|
||||
if configType != "" {
|
||||
query = query.Where("type = ?", configType)
|
||||
}
|
||||
@@ -283,7 +280,7 @@ func ListAdminSystemConfigs(ctx context.Context, configType string) ([]SystemCon
|
||||
// GetAdminSystemConfigByKey loads a config directly from DB.
|
||||
func GetAdminSystemConfigByKey(ctx context.Context, key string) (SystemConfig, error) {
|
||||
var config SystemConfig
|
||||
if err := db.DB(ctx).Where("key = ?", key).First(&config).Error; err != nil {
|
||||
if err := GetDB(ctx).Where("key = ?", key).First(&config).Error; err != nil {
|
||||
return SystemConfig{}, err
|
||||
}
|
||||
return config, nil
|
||||
@@ -292,7 +289,7 @@ func GetAdminSystemConfigByKey(ctx context.Context, key string) (SystemConfig, e
|
||||
// SystemConfigExists reports whether a config key already exists.
|
||||
func SystemConfigExists(ctx context.Context, key string) (bool, error) {
|
||||
var existing SystemConfig
|
||||
err := db.DB(ctx).Where("key = ?", key).First(&existing).Error
|
||||
err := GetDB(ctx).Where("key = ?", key).First(&existing).Error
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return false, nil
|
||||
}
|
||||
@@ -304,18 +301,18 @@ func SystemConfigExists(ctx context.Context, key string) (bool, error) {
|
||||
|
||||
// CreateSystemConfigRecord persists a new system config row.
|
||||
func CreateSystemConfigRecord(ctx context.Context, config *SystemConfig) error {
|
||||
return db.DB(ctx).Create(config).Error
|
||||
return GetDB(ctx).Create(config).Error
|
||||
}
|
||||
|
||||
// UpdateSystemConfigFields applies partial updates to a system config row.
|
||||
func UpdateSystemConfigFields(ctx context.Context, config *SystemConfig, updates map[string]any) error {
|
||||
return db.DB(ctx).Model(config).Updates(updates).Error
|
||||
return GetDB(ctx).Model(config).Updates(updates).Error
|
||||
}
|
||||
|
||||
// SaveOrUpdateSystemConfig creates or updates a config row and invalidates cache.
|
||||
func SaveOrUpdateSystemConfig(ctx context.Context, key, value string) error {
|
||||
var sc SystemConfig
|
||||
err := db.DB(ctx).Where("key = ?", key).First(&sc).Error
|
||||
err := GetDB(ctx).Where("key = ?", key).First(&sc).Error
|
||||
if err != nil && !errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return err
|
||||
}
|
||||
@@ -327,12 +324,12 @@ func SaveOrUpdateSystemConfig(ctx context.Context, key, value string) error {
|
||||
Type: configTypeSystem,
|
||||
Visibility: ConfigVisibilityHidden,
|
||||
}
|
||||
if err := db.DB(ctx).Create(&sc).Error; err != nil {
|
||||
if err := GetDB(ctx).Create(&sc).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
} else {
|
||||
sc.Value = value
|
||||
if err := db.DB(ctx).Save(&sc).Error; err != nil {
|
||||
if err := GetDB(ctx).Save(&sc).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
@@ -342,7 +339,7 @@ func SaveOrUpdateSystemConfig(ctx context.Context, key, value string) error {
|
||||
// ListTemplatesRecord returns all templates ordered by system flag and creation time.
|
||||
func ListTemplatesRecord(ctx context.Context) ([]Template, error) {
|
||||
var templates []Template
|
||||
if err := db.DB(ctx).Order("is_system DESC, created_at DESC").Find(&templates).Error; err != nil {
|
||||
if err := GetDB(ctx).Order("is_system DESC, created_at DESC").Find(&templates).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return templates, nil
|
||||
@@ -351,7 +348,7 @@ func ListTemplatesRecord(ctx context.Context) ([]Template, error) {
|
||||
// GetTemplateByKey loads a template by its key.
|
||||
func GetTemplateByKey(ctx context.Context, key string) (Template, error) {
|
||||
var tmpl Template
|
||||
if err := db.DB(ctx).Where("key = ?", key).First(&tmpl).Error; err != nil {
|
||||
if err := GetDB(ctx).Where("key = ?", key).First(&tmpl).Error; err != nil {
|
||||
return Template{}, err
|
||||
}
|
||||
return tmpl, nil
|
||||
@@ -360,7 +357,7 @@ func GetTemplateByKey(ctx context.Context, key string) (Template, error) {
|
||||
// TemplateExistsByKey reports whether a template key is already taken.
|
||||
func TemplateExistsByKey(ctx context.Context, key string) (bool, error) {
|
||||
var existing Template
|
||||
err := db.DB(ctx).Where("key = ?", key).First(&existing).Error
|
||||
err := GetDB(ctx).Where("key = ?", key).First(&existing).Error
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return false, nil
|
||||
}
|
||||
@@ -372,38 +369,38 @@ func TemplateExistsByKey(ctx context.Context, key string) (bool, error) {
|
||||
|
||||
// CreateTemplateRecord persists a new template.
|
||||
func CreateTemplateRecord(ctx context.Context, tmpl *Template) error {
|
||||
return db.DB(ctx).Create(tmpl).Error
|
||||
return GetDB(ctx).Create(tmpl).Error
|
||||
}
|
||||
|
||||
// SaveTemplateRecord updates an existing template.
|
||||
func SaveTemplateRecord(ctx context.Context, tmpl *Template) error {
|
||||
return db.DB(ctx).Save(tmpl).Error
|
||||
return GetDB(ctx).Save(tmpl).Error
|
||||
}
|
||||
|
||||
// DeleteTemplateRecord removes a template record.
|
||||
func DeleteTemplateRecord(ctx context.Context, tmpl *Template) error {
|
||||
return db.DB(ctx).Delete(tmpl).Error
|
||||
return GetDB(ctx).Delete(tmpl).Error
|
||||
}
|
||||
|
||||
// CreateScheduleRecord 创建定时任务
|
||||
func CreateScheduleRecord(ctx context.Context, schedule *Schedule) error {
|
||||
return db.DB(ctx).Create(schedule).Error
|
||||
return GetDB(ctx).Create(schedule).Error
|
||||
}
|
||||
|
||||
// UpdateScheduleRecord 更新定时任务
|
||||
func UpdateScheduleRecord(ctx context.Context, schedule *Schedule) error {
|
||||
return db.DB(ctx).Save(schedule).Error
|
||||
return GetDB(ctx).Save(schedule).Error
|
||||
}
|
||||
|
||||
// DeleteScheduleRecord 删除定时任务
|
||||
func DeleteScheduleRecord(ctx context.Context, id uint64) error {
|
||||
return db.DB(ctx).Delete(&Schedule{}, id).Error
|
||||
return GetDB(ctx).Delete(&Schedule{}, id).Error
|
||||
}
|
||||
|
||||
// GetScheduleByID 根据 ID 获取定时任务
|
||||
func GetScheduleByID(ctx context.Context, id uint64) (*Schedule, error) {
|
||||
var schedule Schedule
|
||||
if err := db.DB(ctx).Where("id = ?", id).First(&schedule).Error; err != nil {
|
||||
if err := GetDB(ctx).Where("id = ?", id).First(&schedule).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &schedule, nil
|
||||
@@ -412,7 +409,7 @@ func GetScheduleByID(ctx context.Context, id uint64) (*Schedule, error) {
|
||||
// ListSchedulesRecord 获取所有定时任务
|
||||
func ListSchedulesRecord(ctx context.Context) ([]Schedule, error) {
|
||||
var schedules []Schedule
|
||||
if err := db.DB(ctx).Order("id DESC").Find(&schedules).Error; err != nil {
|
||||
if err := GetDB(ctx).Order("id DESC").Find(&schedules).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return schedules, nil
|
||||
@@ -421,7 +418,7 @@ func ListSchedulesRecord(ctx context.Context) ([]Schedule, error) {
|
||||
// ListActiveSchedules 获取所有启用的定时任务
|
||||
func ListActiveSchedules(ctx context.Context) ([]Schedule, error) {
|
||||
var schedules []Schedule
|
||||
if err := db.DB(ctx).Where("is_active = ?", true).Find(&schedules).Error; err != nil {
|
||||
if err := GetDB(ctx).Where("is_active = ?", true).Find(&schedules).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return schedules, nil
|
||||
@@ -430,18 +427,18 @@ func ListActiveSchedules(ctx context.Context) ([]Schedule, error) {
|
||||
// CreateTaskExecutionRecord 创建任务执行记录
|
||||
func CreateTaskExecutionRecord(ctx context.Context, execution *TaskExecution) error {
|
||||
execution.ID = idgen.NextUint64ID()
|
||||
return db.DB(ctx).Create(execution).Error
|
||||
return GetDB(ctx).Create(execution).Error
|
||||
}
|
||||
|
||||
// UpdateTaskExecutionRecord 更新任务执行记录,忽略由 Redis 缓冲和归档流程管理的 log 字段。
|
||||
func UpdateTaskExecutionRecord(ctx context.Context, execution *TaskExecution) error {
|
||||
return db.DB(ctx).Omit("log").Save(execution).Error
|
||||
return GetDB(ctx).Omit("log").Save(execution).Error
|
||||
}
|
||||
|
||||
// GetTaskExecutionByTaskID 根据 TaskID 获取执行记录
|
||||
func GetTaskExecutionByTaskID(ctx context.Context, taskID string) (*TaskExecution, error) {
|
||||
var execution TaskExecution
|
||||
if err := db.DB(ctx).Where("task_id = ?", taskID).First(&execution).Error; err != nil {
|
||||
if err := GetDB(ctx).Where("task_id = ?", taskID).First(&execution).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := loadTaskExecutionLog(ctx, &execution); err != nil {
|
||||
@@ -453,7 +450,7 @@ func GetTaskExecutionByTaskID(ctx context.Context, taskID string) (*TaskExecutio
|
||||
// GetTaskExecutionByID 根据 ID 获取执行记录
|
||||
func GetTaskExecutionByID(ctx context.Context, id uint64) (*TaskExecution, error) {
|
||||
var execution TaskExecution
|
||||
if err := db.DB(ctx).Where("id = ?", id).First(&execution).Error; err != nil {
|
||||
if err := GetDB(ctx).Where("id = ?", id).First(&execution).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := loadTaskExecutionLog(ctx, &execution); err != nil {
|
||||
@@ -465,7 +462,7 @@ func GetTaskExecutionByID(ctx context.Context, id uint64) (*TaskExecution, error
|
||||
// GetLatestTaskExecutionByTaskType returns the most recent execution for a task type.
|
||||
func GetLatestTaskExecutionByTaskType(ctx context.Context, taskType string) (*TaskExecution, bool, error) {
|
||||
var execution TaskExecution
|
||||
err := db.DB(ctx).
|
||||
err := GetDB(ctx).
|
||||
Where("task_type = ?", taskType).
|
||||
Order("id DESC").
|
||||
First(&execution).Error
|
||||
@@ -481,45 +478,40 @@ func GetLatestTaskExecutionByTaskType(ctx context.Context, taskType string) (*Ta
|
||||
return nil, false, err
|
||||
}
|
||||
|
||||
// AppendTaskExecutionLog 将日志追加到 Redis 缓冲,任务完成后再持久化到数据库。
|
||||
// AppendTaskExecutionLog 将日志追加到缓冲,任务完成后再持久化到数据库。
|
||||
func AppendTaskExecutionLog(ctx context.Context, taskID string, logLine string) error {
|
||||
if cachepkg.Redis == nil {
|
||||
return errors.New("redis client is not initialized")
|
||||
cacheSvc := GetCache(ctx)
|
||||
if cacheSvc == nil {
|
||||
return errors.New("cache service is not initialized")
|
||||
}
|
||||
|
||||
now := time.Now().Format("15:04:05")
|
||||
line := fmt.Sprintf("[%s] %s\n", now, logLine)
|
||||
key := taskExecutionLogRedisKey(taskID)
|
||||
|
||||
_, err := cachepkg.Redis.TxPipelined(ctx, func(pipe redis.Pipeliner) error {
|
||||
pipe.RPush(ctx, key, line)
|
||||
pipe.LTrim(ctx, key, -taskExecutionLogMaxLines, -1)
|
||||
pipe.Expire(ctx, key, taskExecutionLogExpiration)
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
return fmt.Errorf("append task execution log to redis: %w", err)
|
||||
}
|
||||
return nil
|
||||
var existing string
|
||||
_ = cacheSvc.Get(ctx, key, &existing)
|
||||
return cacheSvc.Set(ctx, key, existing+line, taskExecutionLogExpiration)
|
||||
}
|
||||
|
||||
// FlushTaskExecutionLog 将 Redis 中的完整任务日志写入数据库,并在成功后清理缓存。
|
||||
// FlushTaskExecutionLog 将缓冲中的完整任务日志写入数据库,并在成功后清理缓存。
|
||||
func FlushTaskExecutionLog(ctx context.Context, taskID string) error {
|
||||
if cachepkg.Redis == nil {
|
||||
return errors.New("redis client is not initialized")
|
||||
cacheSvc := GetCache(ctx)
|
||||
if cacheSvc == nil {
|
||||
return errors.New("cache service is not initialized")
|
||||
}
|
||||
|
||||
key := taskExecutionLogRedisKey(taskID)
|
||||
logLines, err := cachepkg.Redis.LRange(ctx, key, 0, -1).Result()
|
||||
if err != nil {
|
||||
return fmt.Errorf("get task execution log from redis: %w", err)
|
||||
}
|
||||
if len(logLines) == 0 {
|
||||
var logText string
|
||||
if err := cacheSvc.Get(ctx, key, &logText); err != nil || logText == "" {
|
||||
return nil
|
||||
}
|
||||
logText := strings.Join(logLines, "")
|
||||
|
||||
result := db.DB(ctx).Model(&TaskExecution{}).
|
||||
gormDB := GetDB(ctx)
|
||||
if gormDB == nil {
|
||||
return errors.New(errDatabaseNotInitialized)
|
||||
}
|
||||
result := gormDB.Model(&TaskExecution{}).
|
||||
Where("task_id = ?", taskID).
|
||||
Update("log", logText)
|
||||
if result.Error != nil {
|
||||
@@ -529,9 +521,7 @@ func FlushTaskExecutionLog(ctx context.Context, taskID string) error {
|
||||
return fmt.Errorf("persist task execution log: task %q not found", taskID)
|
||||
}
|
||||
|
||||
if err := cachepkg.Redis.Del(ctx, key).Err(); err != nil {
|
||||
return fmt.Errorf("delete persisted task execution log from redis: %w", err)
|
||||
}
|
||||
_ = cacheSvc.Delete(ctx, key)
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -544,7 +534,7 @@ func ListTaskExecutionRecords(ctx context.Context, req ListTaskExecutionsRequest
|
||||
req.PageSize = 20
|
||||
}
|
||||
|
||||
query := db.DB(ctx).Model(&TaskExecution{})
|
||||
query := GetDB(ctx).Model(&TaskExecution{})
|
||||
|
||||
if req.Status != "" {
|
||||
query = query.Where("status = ?", req.Status)
|
||||
@@ -618,7 +608,7 @@ func CleanupTaskExecutionLogs(ctx context.Context, now time.Time) (TaskExecution
|
||||
terminalStatuses := []TaskExecutionStatus{TaskExecutionStatusSucceeded, TaskExecutionStatusFailed}
|
||||
|
||||
var highFrequencyTaskTypes []string
|
||||
if err := db.DB(ctx).
|
||||
if err := GetDB(ctx).
|
||||
Model(&TaskExecution{}).
|
||||
Select("task_type").
|
||||
Where("created_at >= ?", frequencyWindowStart).
|
||||
@@ -630,7 +620,7 @@ func CleanupTaskExecutionLogs(ctx context.Context, now time.Time) (TaskExecution
|
||||
|
||||
var highFrequencyDeleted int64
|
||||
if len(highFrequencyTaskTypes) > 0 {
|
||||
highFrequencyResult := db.DB(ctx).
|
||||
highFrequencyResult := GetDB(ctx).
|
||||
Where("status IN ?", terminalStatuses).
|
||||
Where("created_at < ?", highFrequencyCutoff).
|
||||
Where("task_type IN ?", highFrequencyTaskTypes).
|
||||
@@ -641,7 +631,7 @@ func CleanupTaskExecutionLogs(ctx context.Context, now time.Time) (TaskExecution
|
||||
highFrequencyDeleted = highFrequencyResult.RowsAffected
|
||||
}
|
||||
|
||||
lowFrequencyQuery := db.DB(ctx).
|
||||
lowFrequencyQuery := GetDB(ctx).
|
||||
Where("status IN ?", terminalStatuses).
|
||||
Where("created_at < ?", lowFrequencyCutoff)
|
||||
if len(highFrequencyTaskTypes) > 0 {
|
||||
@@ -659,46 +649,32 @@ func CleanupTaskExecutionLogs(ctx context.Context, now time.Time) (TaskExecution
|
||||
}
|
||||
|
||||
func taskExecutionLogRedisKey(taskID string) string {
|
||||
return cachepkg.PrefixedKey(taskExecutionLogRedisKeyPrefix + taskID)
|
||||
return taskExecutionLogRedisKeyPrefix + taskID
|
||||
}
|
||||
|
||||
func loadTaskExecutionLog(ctx context.Context, execution *TaskExecution) error {
|
||||
if cachepkg.Redis == nil {
|
||||
cacheSvc := GetCache(ctx)
|
||||
if cacheSvc == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
logLines, err := cachepkg.Redis.LRange(ctx, taskExecutionLogRedisKey(execution.TaskID), 0, -1).Result()
|
||||
if err != nil {
|
||||
return fmt.Errorf("get task execution log from redis: %w", err)
|
||||
var logText string
|
||||
if err := cacheSvc.Get(ctx, taskExecutionLogRedisKey(execution.TaskID), &logText); err == nil && logText != "" {
|
||||
execution.Log = logText
|
||||
}
|
||||
if len(logLines) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
execution.Log = strings.Join(logLines, "")
|
||||
return nil
|
||||
}
|
||||
|
||||
func loadTaskExecutionLogs(ctx context.Context, executions []TaskExecution) error {
|
||||
if cachepkg.Redis == nil || len(executions) == 0 {
|
||||
cacheSvc := GetCache(ctx)
|
||||
if cacheSvc == nil || len(executions) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
commands := make([]*redis.StringSliceCmd, len(executions))
|
||||
_, err := cachepkg.Redis.Pipelined(ctx, func(pipe redis.Pipeliner) error {
|
||||
for i := range executions {
|
||||
commands[i] = pipe.LRange(ctx, taskExecutionLogRedisKey(executions[i].TaskID), 0, -1)
|
||||
}
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
return fmt.Errorf("get task execution logs from redis: %w", err)
|
||||
}
|
||||
|
||||
for i := range executions {
|
||||
logLines := commands[i].Val()
|
||||
if len(logLines) > 0 {
|
||||
executions[i].Log = strings.Join(logLines, "")
|
||||
var logText string
|
||||
if err := cacheSvc.Get(ctx, taskExecutionLogRedisKey(executions[i].TaskID), &logText); err == nil && logText != "" {
|
||||
executions[i].Log = logText
|
||||
}
|
||||
}
|
||||
return nil
|
||||
|
||||
@@ -7,14 +7,11 @@ import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"gorm.io/gorm"
|
||||
|
||||
"Wavelet/pkg/cache/ram"
|
||||
"Wavelet/pkg/util"
|
||||
cachepkg "Wavelet/plugins/infra/cache"
|
||||
)
|
||||
|
||||
const (
|
||||
@@ -33,11 +30,6 @@ const (
|
||||
ConfigCacheType = "config"
|
||||
)
|
||||
|
||||
type systemConfigBroadcastMessage struct {
|
||||
Type string `json:"type"`
|
||||
Key string `json:"key"`
|
||||
}
|
||||
|
||||
// ConfigLoader loads configuration data from the database.
|
||||
type ConfigLoader struct{}
|
||||
|
||||
@@ -64,21 +56,19 @@ func (ConfigLoader) LoadAll(ctx context.Context, configType string) ([]ram.Cache
|
||||
return items, nil
|
||||
}
|
||||
|
||||
// LoadOne loads a single system config from database as a CacheItem.
|
||||
// LoadOne loads a single system config from database as CacheItem.
|
||||
func (ConfigLoader) LoadOne(ctx context.Context, configType string, key string) (ram.CacheItem, error) {
|
||||
cfg, err := PreheatSystemConfigByKey(ctx, key)
|
||||
cfg, err := GetSystemConfigByKey(ctx, key)
|
||||
if err != nil {
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return ram.CacheItem{}, ram.ErrNotFound
|
||||
}
|
||||
return ram.CacheItem{}, err
|
||||
}
|
||||
|
||||
valBytes, err := json.Marshal(cfg)
|
||||
if err != nil {
|
||||
return ram.CacheItem{}, err
|
||||
}
|
||||
|
||||
return ram.CacheItem{
|
||||
Key: cfg.Key,
|
||||
Value: string(valBytes),
|
||||
@@ -87,121 +77,67 @@ func (ConfigLoader) LoadOne(ctx context.Context, configType string, key string)
|
||||
}, nil
|
||||
}
|
||||
|
||||
// PreloadSystemConfigs warms the in-memory RAM cache from database on startup.
|
||||
func PreloadSystemConfigs(ctx context.Context) error {
|
||||
return ram.Refresh(ctx, ConfigCacheType, "", ConfigLoader{})
|
||||
// GetCachedSystemConfig retrieves a single system config with RAM L1 fallback to DB.
|
||||
func GetCachedSystemConfig(ctx context.Context, key string) (*SystemConfig, error) {
|
||||
if item, ok := ram.Get(ConfigCacheType, key); ok {
|
||||
var cfg SystemConfig
|
||||
if err := json.Unmarshal([]byte(item.Value), &cfg); err == nil {
|
||||
return &cfg, nil
|
||||
}
|
||||
}
|
||||
|
||||
cfg, err := GetSystemConfigByKey(ctx, key)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
valBytes, err := json.Marshal(cfg)
|
||||
if err == nil {
|
||||
ram.Set(ram.CacheItem{
|
||||
Key: cfg.Key,
|
||||
Value: string(valBytes),
|
||||
Type: ConfigCacheType,
|
||||
TTL: determineTTL(key),
|
||||
})
|
||||
}
|
||||
return &cfg, nil
|
||||
}
|
||||
|
||||
var (
|
||||
systemConfigListenerOnce sync.Once
|
||||
systemConfigListenerCtx context.Context
|
||||
systemConfigListenerCancel context.CancelFunc
|
||||
systemConfigListenerDone chan struct{}
|
||||
)
|
||||
// StopSystemConfigCacheListener stops the cache invalidation listener (kept for backward compatibility).
|
||||
func StopSystemConfigCacheListener() {
|
||||
}
|
||||
|
||||
// StartSystemConfigCacheListener starts the cache listener (kept for backward compatibility).
|
||||
func StartSystemConfigCacheListener() {
|
||||
}
|
||||
|
||||
func ensureSystemConfigCacheListener() {
|
||||
systemConfigListenerOnce.Do(startSystemConfigCacheInvalidationListener)
|
||||
}
|
||||
|
||||
func startSystemConfigCacheInvalidationListener() {
|
||||
if cachepkg.Redis == nil {
|
||||
return
|
||||
}
|
||||
|
||||
systemConfigListenerCtx, systemConfigListenerCancel = context.WithCancel(context.Background())
|
||||
systemConfigListenerDone = make(chan struct{})
|
||||
|
||||
redisClient := cachepkg.Redis // 捕获当前客户端:goroutine 不读可变全局,避免与测试置空 cachepkg.Redis 竞争
|
||||
util.Go(func() {
|
||||
listenerCtx := systemConfigListenerCtx
|
||||
defer close(systemConfigListenerDone)
|
||||
|
||||
pubsub := redisClient.Subscribe(listenerCtx, SystemConfigBroadcastChannel)
|
||||
defer func() {
|
||||
_ = pubsub.Close()
|
||||
}()
|
||||
|
||||
util.Go(func() {
|
||||
<-listenerCtx.Done()
|
||||
_ = pubsub.Close()
|
||||
})
|
||||
|
||||
for msg := range pubsub.Channel() {
|
||||
var payload systemConfigBroadcastMessage
|
||||
if err := json.Unmarshal([]byte(msg.Payload), &payload); err != nil {
|
||||
ram.UpdateTypeItems(ConfigCacheType, nil)
|
||||
continue
|
||||
}
|
||||
|
||||
key := payload.Key
|
||||
if key == "*" || key == "" {
|
||||
ram.UpdateTypeItems(payload.Type, nil)
|
||||
} else {
|
||||
ram.Delete(payload.Type, key)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
// StopSystemConfigCacheListener stops the Redis Pub/Sub subscription listener and resets the sync.Once guard.
|
||||
func StopSystemConfigCacheListener() {
|
||||
if systemConfigListenerCancel != nil {
|
||||
systemConfigListenerCancel()
|
||||
if systemConfigListenerDone != nil {
|
||||
<-systemConfigListenerDone
|
||||
}
|
||||
systemConfigListenerCancel = nil
|
||||
systemConfigListenerDone = nil
|
||||
}
|
||||
systemConfigListenerOnce = sync.Once{}
|
||||
}
|
||||
|
||||
func determineTTL(_ string) time.Duration {
|
||||
// Program-determined TTL: -1 means never expire for all configs by default
|
||||
return -1
|
||||
}
|
||||
|
||||
// InvalidateSystemConfigCache triggers a broadcast to refresh the cache for key.
|
||||
func InvalidateSystemConfigCache(ctx context.Context, key string) error {
|
||||
ensureSystemConfigCacheListener()
|
||||
|
||||
// Invalidate local cache synchronously first
|
||||
ram.Delete(ConfigCacheType, key)
|
||||
|
||||
// Broadcast to other nodes and clean legacy Redis cache key
|
||||
if cachepkg.Redis != nil {
|
||||
_ = cachepkg.HDel(ctx, SystemConfigRedisHashKey, key)
|
||||
publishSystemConfigBroadcast(ctx, ConfigCacheType, key)
|
||||
if cacheSvc := GetCache(ctx); cacheSvc != nil {
|
||||
_ = cacheSvc.Delete(ctx, "system:config:"+key)
|
||||
_ = cacheSvc.Delete(ctx, SystemConfigVisibleListRedisKey)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// InvalidateAllSystemConfigCaches triggers a broadcast to refresh the entire config cache.
|
||||
func InvalidateAllSystemConfigCaches(ctx context.Context) error {
|
||||
ensureSystemConfigCacheListener()
|
||||
|
||||
// Invalidate all items of type ConfigCacheType synchronously first
|
||||
ram.UpdateTypeItems(ConfigCacheType, nil)
|
||||
|
||||
// Broadcast to other nodes and clean legacy Redis cache keys
|
||||
if cachepkg.Redis != nil {
|
||||
_ = cachepkg.Redis.Del(ctx, cachepkg.PrefixedKey(SystemConfigRedisHashKey), cachepkg.PrefixedKey(SystemConfigVisibleListRedisKey)).Err()
|
||||
publishSystemConfigBroadcast(ctx, ConfigCacheType, "*")
|
||||
if cacheSvc := GetCache(ctx); cacheSvc != nil {
|
||||
_ = cacheSvc.Delete(ctx, SystemConfigRedisHashKey)
|
||||
_ = cacheSvc.Delete(ctx, SystemConfigVisibleListRedisKey)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func publishSystemConfigBroadcast(ctx context.Context, configType string, key string) {
|
||||
if cachepkg.Redis == nil {
|
||||
return
|
||||
}
|
||||
payload, err := json.Marshal(systemConfigBroadcastMessage{Type: configType, Key: key})
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
_ = cachepkg.Redis.Publish(ctx, SystemConfigBroadcastChannel, payload).Err()
|
||||
}
|
||||
|
||||
// ResetSystemConfigRAMCacheForTest clears only the process-local RAM cache.
|
||||
func ResetSystemConfigRAMCacheForTest() {
|
||||
ram.ResetForTest()
|
||||
|
||||
@@ -8,16 +8,30 @@ import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/alicebob/miniredis/v2"
|
||||
"github.com/glebarez/sqlite"
|
||||
"github.com/redis/go-redis/v9"
|
||||
"github.com/redis/go-redis/v9/maintnotifications"
|
||||
"gorm.io/gorm"
|
||||
|
||||
"Wavelet/plugins/infra/cache"
|
||||
"Wavelet/plugins/infra/database"
|
||||
)
|
||||
|
||||
type testDBService struct {
|
||||
db *gorm.DB
|
||||
}
|
||||
|
||||
func (s *testDBService) DB(ctx context.Context) *gorm.DB {
|
||||
return s.db
|
||||
}
|
||||
|
||||
func (s *testDBService) MasterDB(ctx context.Context) *gorm.DB {
|
||||
return s.db
|
||||
}
|
||||
|
||||
func (s *testDBService) GORM() *gorm.DB {
|
||||
return s.db
|
||||
}
|
||||
|
||||
func (s *testDBService) Named(_ string) *gorm.DB {
|
||||
return s.db
|
||||
}
|
||||
|
||||
func setupSystemConfigTest(t *testing.T) (*gorm.DB, func()) {
|
||||
t.Helper()
|
||||
|
||||
@@ -41,28 +55,12 @@ func setupSystemConfigTest(t *testing.T) (*gorm.DB, func()) {
|
||||
t.Fatalf("Create(site_name) error = %v", err)
|
||||
}
|
||||
|
||||
mr, err := miniredis.Run()
|
||||
if err != nil {
|
||||
t.Fatalf("miniredis.Run() error = %v", err)
|
||||
}
|
||||
redisClient := redis.NewClient(&redis.Options{
|
||||
Addr: mr.Addr(),
|
||||
MaintNotificationsConfig: &maintnotifications.Config{
|
||||
Mode: maintnotifications.ModeDisabled,
|
||||
},
|
||||
})
|
||||
|
||||
previousRedis := cache.Redis
|
||||
database.SetDB(sqliteDB)
|
||||
cache.Redis = redisClient
|
||||
SetDBService(&testDBService{db: sqliteDB})
|
||||
|
||||
cleanup := func() {
|
||||
StopSystemConfigCacheListener()
|
||||
ResetSystemConfigRAMCacheForTest()
|
||||
database.SetDB(nil)
|
||||
cache.Redis = previousRedis
|
||||
_ = redisClient.Close()
|
||||
mr.Close()
|
||||
ResetServices()
|
||||
}
|
||||
|
||||
return sqliteDB, cleanup
|
||||
|
||||
@@ -7,9 +7,10 @@ import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/pkg/logger"
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
// LogForAudit 将登录鉴权审计日志写入 Logger
|
||||
|
||||
@@ -11,7 +11,6 @@ import (
|
||||
"strings"
|
||||
|
||||
"Wavelet/core/contracts"
|
||||
db "Wavelet/plugins/infra/database"
|
||||
|
||||
"github.com/coreos/go-oidc/v3/oidc"
|
||||
"golang.org/x/oauth2"
|
||||
@@ -19,7 +18,7 @@ import (
|
||||
|
||||
func isOIDCLoginEnabled(ctx context.Context) bool {
|
||||
var val string
|
||||
if err := db.DB(ctx).Table("w_system_configs").Where("key = ?", "oidc_login_enabled").Pluck("value", &val).Error; err != nil || val == "" {
|
||||
if err := getDB(ctx).Table("w_system_configs").Where("key = ?", "oidc_login_enabled").Pluck("value", &val).Error; err != nil || val == "" {
|
||||
return true
|
||||
}
|
||||
b, err := strconv.ParseBool(val)
|
||||
@@ -78,7 +77,7 @@ func activeLoginSources(ctx context.Context) []AuthSourceView {
|
||||
|
||||
func getFrontendLoginRedirectURL(ctx context.Context) (string, error) {
|
||||
var val string
|
||||
if err := db.DB(ctx).Table("w_system_configs").Where("key = ?", "server_address").Pluck("value", &val).Error; err != nil || strings.TrimSpace(val) == "" {
|
||||
if err := getDB(ctx).Table("w_system_configs").Where("key = ?", "server_address").Pluck("value", &val).Error; err != nil || strings.TrimSpace(val) == "" {
|
||||
return "", errors.New(errServerAddressMissing)
|
||||
}
|
||||
return strings.TrimRight(val, "/") + "/login", nil
|
||||
|
||||
@@ -6,23 +6,15 @@ package auth
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"strconv"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/pkg/cache/ram"
|
||||
"Wavelet/pkg/util"
|
||||
db "Wavelet/plugins/infra/cache"
|
||||
)
|
||||
|
||||
const (
|
||||
tokenCacheTTL = 5 * time.Minute
|
||||
userCacheTTL = 5 * time.Minute
|
||||
|
||||
//nolint:gosec // This is a Redis Pub/Sub channel name, not a credential
|
||||
oauthTokenInvalidationChannel = "oauth:token_invalidation"
|
||||
oauthUserInvalidationChannel = "oauth:user_invalidation"
|
||||
)
|
||||
|
||||
// CachedToken represents the minimal cached representation of an access token.
|
||||
@@ -35,16 +27,6 @@ type CachedToken struct {
|
||||
var (
|
||||
tokenRAM = ram.MustNew[string, *CachedToken](ram.Options{MaximumSize: 2048})
|
||||
userRAM = ram.MustNew[uint64, *contracts.UserDTO](ram.Options{MaximumSize: 2048})
|
||||
|
||||
tokenListenerOnce sync.Once
|
||||
tokenListenerCtx context.Context
|
||||
tokenListenerCancel context.CancelFunc
|
||||
tokenListenerDone chan struct{}
|
||||
|
||||
userListenerOnce sync.Once
|
||||
userListenerCtx context.Context
|
||||
userListenerCancel context.CancelFunc
|
||||
userListenerDone chan struct{}
|
||||
)
|
||||
|
||||
func tokenCacheKey(tokenHash string) string {
|
||||
@@ -55,107 +37,16 @@ func userCacheKey(userID uint64) string {
|
||||
return fmt.Sprintf("oauth:user:%d", userID)
|
||||
}
|
||||
|
||||
func ensureTokenCacheListener() {
|
||||
if db.Redis == nil {
|
||||
return
|
||||
}
|
||||
tokenListenerOnce.Do(startTokenCacheInvalidationListener)
|
||||
}
|
||||
|
||||
func startTokenCacheInvalidationListener() {
|
||||
tokenListenerCtx, tokenListenerCancel = context.WithCancel(context.Background())
|
||||
tokenListenerDone = make(chan struct{})
|
||||
|
||||
redisClient := db.Redis // 捕获当前客户端:goroutine 不读可变全局,避免与测试置空 db.Redis 竞争
|
||||
util.Go(func() {
|
||||
listenerCtx := tokenListenerCtx
|
||||
defer close(tokenListenerDone)
|
||||
|
||||
pubsub := redisClient.Subscribe(listenerCtx, oauthTokenInvalidationChannel)
|
||||
defer func() {
|
||||
_ = pubsub.Close()
|
||||
}()
|
||||
|
||||
util.Go(func() {
|
||||
<-listenerCtx.Done()
|
||||
_ = pubsub.Close()
|
||||
})
|
||||
|
||||
for msg := range pubsub.Channel() {
|
||||
tokenHash := msg.Payload
|
||||
if tokenHash == "" || tokenHash == "*" || tokenHash == "reset" {
|
||||
tokenRAM.InvalidateAll()
|
||||
} else {
|
||||
tokenRAM.Invalidate(tokenHash)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func publishTokenRAMInvalidation(ctx context.Context, tokenHash string) {
|
||||
if db.Redis == nil {
|
||||
return
|
||||
}
|
||||
_ = db.Redis.Publish(ctx, oauthTokenInvalidationChannel, tokenHash).Err()
|
||||
}
|
||||
|
||||
func ensureUserCacheListener() {
|
||||
if db.Redis == nil {
|
||||
return
|
||||
}
|
||||
userListenerOnce.Do(startUserCacheInvalidationListener)
|
||||
}
|
||||
|
||||
func startUserCacheInvalidationListener() {
|
||||
userListenerCtx, userListenerCancel = context.WithCancel(context.Background())
|
||||
userListenerDone = make(chan struct{})
|
||||
|
||||
redisClient := db.Redis // 捕获当前客户端:goroutine 不读可变全局,避免与测试置空 db.Redis 竞争
|
||||
util.Go(func() {
|
||||
listenerCtx := userListenerCtx
|
||||
defer close(userListenerDone)
|
||||
|
||||
pubsub := redisClient.Subscribe(listenerCtx, oauthUserInvalidationChannel)
|
||||
defer func() {
|
||||
_ = pubsub.Close()
|
||||
}()
|
||||
|
||||
util.Go(func() {
|
||||
<-listenerCtx.Done()
|
||||
_ = pubsub.Close()
|
||||
})
|
||||
|
||||
for msg := range pubsub.Channel() {
|
||||
userIDStr := msg.Payload
|
||||
if userIDStr == "" || userIDStr == "*" || userIDStr == "reset" {
|
||||
userRAM.InvalidateAll()
|
||||
} else if userID, err := strconv.ParseUint(userIDStr, 10, 64); err == nil {
|
||||
userRAM.Invalidate(userID)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func publishUserRAMInvalidation(ctx context.Context, userID uint64) {
|
||||
if db.Redis == nil {
|
||||
return
|
||||
}
|
||||
_ = db.Redis.Publish(ctx, oauthUserInvalidationChannel, strconv.FormatUint(userID, 10)).Err()
|
||||
}
|
||||
|
||||
// GetCachedToken 获取缓存的 Token
|
||||
func GetCachedToken(ctx context.Context, tokenHash string) (*CachedToken, error) {
|
||||
ensureTokenCacheListener()
|
||||
|
||||
if val, ok := tokenRAM.GetIfPresent(tokenHash); ok {
|
||||
return val, nil
|
||||
}
|
||||
|
||||
if db.Redis != nil {
|
||||
if cache := getCache(ctx); cache != nil {
|
||||
var token CachedToken
|
||||
key := tokenCacheKey(tokenHash)
|
||||
if err := db.GetJSON(ctx, key, &token); err == nil {
|
||||
// Write back to local cache
|
||||
if err := cache.Get(ctx, key, &token); err == nil {
|
||||
tokenRAM.Set(tokenHash, &token)
|
||||
return &token, nil
|
||||
}
|
||||
@@ -165,40 +56,32 @@ func GetCachedToken(ctx context.Context, tokenHash string) (*CachedToken, error)
|
||||
|
||||
// SetCachedToken 设置 Token 缓存
|
||||
func SetCachedToken(ctx context.Context, tokenHash string, token *CachedToken) {
|
||||
ensureTokenCacheListener()
|
||||
|
||||
tokenRAM.Set(tokenHash, token)
|
||||
if db.Redis != nil {
|
||||
if cache := getCache(ctx); cache != nil {
|
||||
key := tokenCacheKey(tokenHash)
|
||||
_ = db.SetJSON(ctx, key, token, tokenCacheTTL)
|
||||
_ = cache.Set(ctx, key, token, tokenCacheTTL)
|
||||
}
|
||||
}
|
||||
|
||||
// InvalidateCachedToken 吊销/删除 token 缓存
|
||||
func InvalidateCachedToken(ctx context.Context, tokenHash string) {
|
||||
ensureTokenCacheListener()
|
||||
|
||||
tokenRAM.Invalidate(tokenHash)
|
||||
if db.Redis != nil {
|
||||
if cache := getCache(ctx); cache != nil {
|
||||
key := tokenCacheKey(tokenHash)
|
||||
_ = db.Redis.Del(ctx, db.PrefixedKey(key)).Err()
|
||||
publishTokenRAMInvalidation(ctx, tokenHash)
|
||||
_ = cache.Delete(ctx, key)
|
||||
}
|
||||
}
|
||||
|
||||
// GetCachedUser 获取缓存的 UserDTO
|
||||
func GetCachedUser(ctx context.Context, userID uint64) (*contracts.UserDTO, error) {
|
||||
ensureUserCacheListener()
|
||||
|
||||
if val, ok := userRAM.GetIfPresent(userID); ok {
|
||||
return val, nil
|
||||
}
|
||||
|
||||
if db.Redis != nil {
|
||||
if cache := getCache(ctx); cache != nil {
|
||||
var u contracts.UserDTO
|
||||
key := userCacheKey(userID)
|
||||
if err := db.GetJSON(ctx, key, &u); err == nil {
|
||||
// Write back to local cache
|
||||
if err := cache.Get(ctx, key, &u); err == nil {
|
||||
userRAM.Set(userID, &u)
|
||||
return &u, nil
|
||||
}
|
||||
@@ -208,49 +91,24 @@ func GetCachedUser(ctx context.Context, userID uint64) (*contracts.UserDTO, erro
|
||||
|
||||
// SetCachedUser 设置 UserDTO 缓存
|
||||
func SetCachedUser(ctx context.Context, userID uint64, u *contracts.UserDTO) {
|
||||
ensureUserCacheListener()
|
||||
|
||||
userRAM.Set(userID, u)
|
||||
if db.Redis != nil {
|
||||
if cache := getCache(ctx); cache != nil {
|
||||
key := userCacheKey(userID)
|
||||
_ = db.SetJSON(ctx, key, u, userCacheTTL)
|
||||
_ = cache.Set(ctx, key, u, userCacheTTL)
|
||||
}
|
||||
}
|
||||
|
||||
// InvalidateCachedUser 吊销/失效 UserDTO 缓存
|
||||
func InvalidateCachedUser(ctx context.Context, userID uint64) {
|
||||
ensureUserCacheListener()
|
||||
|
||||
userRAM.Invalidate(userID)
|
||||
if db.Redis != nil {
|
||||
if cache := getCache(ctx); cache != nil {
|
||||
key := userCacheKey(userID)
|
||||
_ = db.Redis.Del(ctx, db.PrefixedKey(key)).Err()
|
||||
publishUserRAMInvalidation(ctx, userID)
|
||||
_ = cache.Delete(ctx, key)
|
||||
}
|
||||
}
|
||||
|
||||
// StopAuthCacheListener stops both token and user Redis Pub/Sub subscription listeners and resets the sync.Once guards.
|
||||
func StopAuthCacheListener() {
|
||||
if tokenListenerCancel != nil {
|
||||
tokenListenerCancel()
|
||||
if tokenListenerDone != nil {
|
||||
<-tokenListenerDone
|
||||
}
|
||||
tokenListenerCancel = nil
|
||||
tokenListenerDone = nil
|
||||
}
|
||||
tokenListenerOnce = sync.Once{}
|
||||
|
||||
if userListenerCancel != nil {
|
||||
userListenerCancel()
|
||||
if userListenerDone != nil {
|
||||
<-userListenerDone
|
||||
}
|
||||
userListenerCancel = nil
|
||||
userListenerDone = nil
|
||||
}
|
||||
userListenerOnce = sync.Once{}
|
||||
}
|
||||
// StopAuthCacheListener compatibility stub for tests
|
||||
func StopAuthCacheListener() {}
|
||||
|
||||
// ResetAuthRAMCacheForTest clears only the process-local RAM cache.
|
||||
func ResetAuthRAMCacheForTest() {
|
||||
|
||||
@@ -5,48 +5,69 @@ package auth_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/alicebob/miniredis/v2"
|
||||
"github.com/redis/go-redis/v9"
|
||||
"github.com/redis/go-redis/v9/maintnotifications"
|
||||
|
||||
"Wavelet/core"
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/plugins/domain/auth"
|
||||
db "Wavelet/plugins/infra/cache"
|
||||
)
|
||||
|
||||
func setupOauthCacheTest(t *testing.T) (*miniredis.Miniredis, func()) {
|
||||
t.Helper()
|
||||
type mockCacheService struct {
|
||||
items map[string][]byte
|
||||
}
|
||||
|
||||
miniRedis, err := miniredis.Run()
|
||||
func newMockCacheService() *mockCacheService {
|
||||
return &mockCacheService{items: make(map[string][]byte)}
|
||||
}
|
||||
|
||||
func (m *mockCacheService) Get(ctx context.Context, key string, target any) error {
|
||||
b, ok := m.items[key]
|
||||
if !ok {
|
||||
return contracts.ErrCacheMiss
|
||||
}
|
||||
return json.Unmarshal(b, target)
|
||||
}
|
||||
|
||||
func (m *mockCacheService) Set(ctx context.Context, key string, value any, ttl time.Duration) error {
|
||||
b, err := json.Marshal(value)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to start miniredis: %v", err)
|
||||
return err
|
||||
}
|
||||
m.items[key] = b
|
||||
return nil
|
||||
}
|
||||
|
||||
db.Redis = redis.NewClient(&redis.Options{
|
||||
Addr: miniRedis.Addr(),
|
||||
MaintNotificationsConfig: &maintnotifications.Config{
|
||||
Mode: maintnotifications.ModeDisabled,
|
||||
},
|
||||
})
|
||||
func (m *mockCacheService) Delete(ctx context.Context, key string) error {
|
||||
delete(m.items, key)
|
||||
return nil
|
||||
}
|
||||
|
||||
auth.ResetAuthRAMCacheForTest()
|
||||
func (m *mockCacheService) Invalidate(ctx context.Context, key string) error {
|
||||
return m.Delete(ctx, key)
|
||||
}
|
||||
|
||||
cleanup := func() {
|
||||
auth.StopAuthCacheListener()
|
||||
auth.ResetAuthRAMCacheForTest()
|
||||
_ = db.Redis.Close()
|
||||
miniRedis.Close()
|
||||
db.Redis = nil
|
||||
func (m *mockCacheService) GetOrSet(ctx context.Context, key string, target any, ttl time.Duration, loader func() (any, error)) error {
|
||||
err := m.Get(ctx, key, target)
|
||||
if err == nil {
|
||||
return nil
|
||||
}
|
||||
return miniRedis, cleanup
|
||||
val, err := loader()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if err := m.Set(ctx, key, val, ttl); err != nil {
|
||||
return err
|
||||
}
|
||||
b, _ := json.Marshal(val)
|
||||
return json.Unmarshal(b, target)
|
||||
}
|
||||
|
||||
func TestTokenCache_GetSetInvalidate(t *testing.T) {
|
||||
_, cleanup := setupOauthCacheTest(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
ctx := core.NewContext(context.Background())
|
||||
mockCache := newMockCacheService()
|
||||
core.Provide[contracts.CacheService](ctx, mockCache)
|
||||
|
||||
tokenHash := "test-token-hash"
|
||||
token := &auth.CachedToken{
|
||||
@@ -84,9 +105,9 @@ func TestTokenCache_GetSetInvalidate(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestUserCache_GetSetInvalidate(t *testing.T) {
|
||||
_, cleanup := setupOauthCacheTest(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
ctx := core.NewContext(context.Background())
|
||||
mockCache := newMockCacheService()
|
||||
core.Provide[contracts.CacheService](ctx, mockCache)
|
||||
|
||||
userID := uint64(789)
|
||||
user := &contracts.UserDTO{
|
||||
|
||||
@@ -14,17 +14,16 @@ import (
|
||||
|
||||
"Wavelet/core/contracts"
|
||||
|
||||
"Wavelet/pkg/idgen"
|
||||
"Wavelet/pkg/logger"
|
||||
"Wavelet/pkg/response"
|
||||
"Wavelet/pkg/util"
|
||||
cachepkg "Wavelet/plugins/infra/cache"
|
||||
db "Wavelet/plugins/infra/database"
|
||||
"github.com/coreos/go-oidc/v3/oidc"
|
||||
"github.com/gin-contrib/sessions"
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/google/uuid"
|
||||
"gorm.io/gorm"
|
||||
|
||||
"Wavelet/pkg/idgen"
|
||||
"Wavelet/pkg/logger"
|
||||
"Wavelet/pkg/response"
|
||||
"Wavelet/pkg/util"
|
||||
)
|
||||
|
||||
// GetLoginSources 获取可用登录源列表
|
||||
@@ -78,9 +77,12 @@ func GetLoginURL(c *gin.Context) {
|
||||
response.AbortInternal(c, err.Error())
|
||||
return
|
||||
}
|
||||
if err := cachepkg.Redis.Set(ctx, cachepkg.PrefixedKey(fmt.Sprintf(OAuthStateCacheKeyFormat, state)), payloadValue, OAuthStateCacheKeyExpiration).Err(); err != nil {
|
||||
response.AbortInternal(c, err.Error())
|
||||
return
|
||||
stateKey := fmt.Sprintf(OAuthStateCacheKeyFormat, state)
|
||||
if cache := getCache(ctx); cache != nil {
|
||||
if err := cache.Set(ctx, stateKey, payloadValue, OAuthStateCacheKeyExpiration); err != nil {
|
||||
response.AbortInternal(c, err.Error())
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
authorizeURL, err := buildAuthorizeURL(c.Request.Context(), source, state)
|
||||
@@ -107,18 +109,19 @@ func buildAuthorizeURL(ctx context.Context, source *AuthSource, state string) (s
|
||||
}
|
||||
|
||||
func reserveOAuthStateSlot(ctx context.Context, sessionHash string) error {
|
||||
if cachepkg.Redis == nil || sessionHash == "" {
|
||||
if sessionHash == "" {
|
||||
return nil
|
||||
}
|
||||
key := cachepkg.PrefixedKey(fmt.Sprintf(oauthStateLimitKeyFormat, sessionHash))
|
||||
n, err := cachepkg.Redis.Incr(ctx, key).Result()
|
||||
if err != nil {
|
||||
return err
|
||||
cache := getCache(ctx)
|
||||
if cache == nil {
|
||||
return nil
|
||||
}
|
||||
if n == 1 {
|
||||
_ = cachepkg.Redis.Expire(ctx, key, OAuthStateCacheKeyExpiration).Err()
|
||||
}
|
||||
if n > oauthStateLimitMax {
|
||||
key := fmt.Sprintf(oauthStateLimitKeyFormat, sessionHash)
|
||||
var count int
|
||||
_ = cache.Get(ctx, key, &count)
|
||||
count++
|
||||
_ = cache.Set(ctx, key, count, OAuthStateCacheKeyExpiration)
|
||||
if count > oauthStateLimitMax {
|
||||
return errors.New(errOAuthStateRateLimited)
|
||||
}
|
||||
return nil
|
||||
@@ -179,9 +182,12 @@ func Authorize(c *gin.Context) {
|
||||
response.AbortInternal(c, err.Error())
|
||||
return
|
||||
}
|
||||
if err := cachepkg.Redis.Set(ctx, cachepkg.PrefixedKey(fmt.Sprintf(OAuthStateCacheKeyFormat, state)), payloadValue, OAuthStateCacheKeyExpiration).Err(); err != nil {
|
||||
response.AbortInternal(c, err.Error())
|
||||
return
|
||||
stateKey := fmt.Sprintf(OAuthStateCacheKeyFormat, state)
|
||||
if cache := getCache(ctx); cache != nil {
|
||||
if err := cache.Set(ctx, stateKey, payloadValue, OAuthStateCacheKeyExpiration); err != nil {
|
||||
response.AbortInternal(c, err.Error())
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
authorizeURL, err := buildAuthorizeURL(c.Request.Context(), source, state)
|
||||
@@ -201,13 +207,18 @@ func Callback(c *gin.Context) {
|
||||
}
|
||||
|
||||
ctx := c.Request.Context()
|
||||
stateKey := cachepkg.PrefixedKey(fmt.Sprintf(OAuthStateCacheKeyFormat, req.State))
|
||||
payloadRaw, err := cachepkg.Redis.Get(ctx, stateKey).Result()
|
||||
if err != nil {
|
||||
stateKey := fmt.Sprintf(OAuthStateCacheKeyFormat, req.State)
|
||||
var payloadRaw string
|
||||
cache := getCache(ctx)
|
||||
if cache == nil {
|
||||
response.AbortBadRequest(c, errInvalidState)
|
||||
return
|
||||
}
|
||||
_ = cachepkg.Redis.Del(ctx, stateKey)
|
||||
if err := cache.Get(ctx, stateKey, &payloadRaw); err != nil {
|
||||
response.AbortBadRequest(c, errInvalidState)
|
||||
return
|
||||
}
|
||||
_ = cache.Delete(ctx, stateKey)
|
||||
|
||||
payload, err := decodeOAuthStatePayload(payloadRaw)
|
||||
if err != nil {
|
||||
@@ -289,7 +300,7 @@ func handleCallbackBind(ctx context.Context, c *gin.Context, source *AuthSource,
|
||||
return
|
||||
}
|
||||
var user contracts.UserDTO
|
||||
if err := db.DB(ctx).Table("w_users").Where("id = ?", userID).First(&user).Error; err != nil {
|
||||
if err := getDB(ctx).Table("w_users").Where("id = ?", userID).First(&user).Error; err != nil {
|
||||
response.AbortInternal(c, err.Error())
|
||||
return
|
||||
}
|
||||
@@ -304,7 +315,7 @@ func handleCallbackBind(ctx context.Context, c *gin.Context, source *AuthSource,
|
||||
return
|
||||
}
|
||||
user.LastLoginAt = time.Now()
|
||||
_ = db.DB(ctx).Table("w_users").Where("id = ?", user.ID).Update("last_login_at", user.LastLoginAt).Error
|
||||
_ = getDB(ctx).Table("w_users").Where("id = ?", user.ID).Update("last_login_at", user.LastLoginAt).Error
|
||||
c.JSON(http.StatusOK, response.OK(buildCallbackResult(&user, "bound")))
|
||||
}
|
||||
|
||||
@@ -314,7 +325,7 @@ func handleCallbackLogin(ctx context.Context, c *gin.Context, source *AuthSource
|
||||
account, err := FindExternalAccount(ctx, source.ID, userInfo.Sub)
|
||||
switch {
|
||||
case err == nil:
|
||||
if loadErr := db.DB(ctx).Table("w_users").Where("id = ?", account.UserID).First(&user).Error; loadErr != nil {
|
||||
if loadErr := getDB(ctx).Table("w_users").Where("id = ?", account.UserID).First(&user).Error; loadErr != nil {
|
||||
response.AbortInternal(c, loadErr.Error())
|
||||
return
|
||||
}
|
||||
@@ -330,7 +341,7 @@ func handleCallbackLogin(ctx context.Context, c *gin.Context, source *AuthSource
|
||||
}
|
||||
|
||||
user.LastLoginAt = time.Now()
|
||||
_ = db.DB(ctx).Table("w_users").Where("id = ?", user.ID).Update("last_login_at", user.LastLoginAt).Error
|
||||
_ = getDB(ctx).Table("w_users").Where("id = ?", user.ID).Update("last_login_at", user.LastLoginAt).Error
|
||||
if err := SetLoginSession(ctx, c, &user); err != nil {
|
||||
response.AbortInternal(c, err.Error())
|
||||
return
|
||||
@@ -348,7 +359,7 @@ func uniqueUsername(ctx context.Context, base string) (string, error) {
|
||||
}
|
||||
|
||||
var existingUsernames []string
|
||||
if err := db.DB(ctx).Table("w_users").
|
||||
if err := getDB(ctx).Table("w_users").
|
||||
Where("username = ? OR username LIKE ? ESCAPE '\\'", base, util.EscapeLike(base)+"-%").
|
||||
Pluck("username", &existingUsernames).Error; err != nil {
|
||||
return "", err
|
||||
@@ -376,7 +387,7 @@ func uniqueUsername(ctx context.Context, base string) (string, error) {
|
||||
func handleCallbackRegister(ctx context.Context, c *gin.Context, source *AuthSource, userInfo *contracts.OAuthUserInfoDTO) (contracts.UserDTO, bool) {
|
||||
registrationEnabled := true
|
||||
var val string
|
||||
if err := db.DB(ctx).Table("w_system_configs").Where("key = ?", "registration_enabled").Pluck("value", &val).Error; err == nil && val != "" {
|
||||
if err := getDB(ctx).Table("w_system_configs").Where("key = ?", "registration_enabled").Pluck("value", &val).Error; err == nil && val != "" {
|
||||
if b, err := strconv.ParseBool(val); err == nil {
|
||||
registrationEnabled = b
|
||||
}
|
||||
@@ -407,7 +418,7 @@ func handleCallbackRegister(ctx context.Context, c *gin.Context, source *AuthSou
|
||||
UpdatedAt: now,
|
||||
}
|
||||
|
||||
if err := db.DB(ctx).Table("w_users").Create(&user).Error; err != nil {
|
||||
if err := getDB(ctx).Table("w_users").Create(&user).Error; err != nil {
|
||||
response.AbortInternal(c, err.Error())
|
||||
return contracts.UserDTO{}, false
|
||||
}
|
||||
|
||||
@@ -9,12 +9,12 @@ import (
|
||||
"encoding/hex"
|
||||
"errors"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/pkg/response"
|
||||
"Wavelet/pkg/trace"
|
||||
"Wavelet/pkg/util"
|
||||
db "Wavelet/plugins/infra/database"
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
func hashToken(token string) string {
|
||||
@@ -38,7 +38,7 @@ func getUserByToken(ctx context.Context, tokenStr string) (*contracts.UserDTO, *
|
||||
UserID uint64
|
||||
IsAdmin bool
|
||||
}
|
||||
if err := db.DB(ctx).Table("w_access_tokens").Where("token_hash = ?", tokenHash).First(&tokenRow).Error; err != nil {
|
||||
if err := getDB(ctx).Table("w_access_tokens").Where("token_hash = ?", tokenHash).First(&tokenRow).Error; err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
tokenRecord = &CachedToken{
|
||||
@@ -49,7 +49,7 @@ func getUserByToken(ctx context.Context, tokenStr string) (*contracts.UserDTO, *
|
||||
SetCachedToken(ctx, tokenHash, tokenRecord)
|
||||
|
||||
var userRow contracts.UserDTO
|
||||
if err := db.DB(ctx).Table("w_users").Where("id = ? AND is_active = ?", tokenRow.UserID, true).First(&userRow).Error; err != nil {
|
||||
if err := getDB(ctx).Table("w_users").Where("id = ? AND is_active = ?", tokenRow.UserID, true).First(&userRow).Error; err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
SetCachedUser(ctx, userRow.ID, &userRow)
|
||||
@@ -90,7 +90,7 @@ func GetUserFromRequest(c *gin.Context) (*contracts.UserDTO, error) {
|
||||
user, err := GetCachedUser(ctx, userID)
|
||||
if err != nil || user == nil || !user.IsActive {
|
||||
var dbUser contracts.UserDTO
|
||||
if err := db.DB(ctx).Table("w_users").Where("id = ? AND is_active = ?", userID, true).First(&dbUser).Error; err != nil {
|
||||
if err := getDB(ctx).Table("w_users").Where("id = ? AND is_active = ?", userID, true).First(&dbUser).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
user = &dbUser
|
||||
|
||||
@@ -76,6 +76,27 @@ func (p *Plugin) Manifest() core.Manifest {
|
||||
|
||||
// Apply registers the auth migrations, services, routes, and settings into the Context.
|
||||
func (p *Plugin) Apply(ctx *core.Context) error {
|
||||
// 0. Bind DBService & CacheService from Context
|
||||
if db, err := core.Inject[contracts.DBService](ctx); err == nil && db != nil {
|
||||
setDBService(db)
|
||||
} else {
|
||||
core.When[contracts.DBService](ctx, func(db contracts.DBService) {
|
||||
setDBService(db)
|
||||
})
|
||||
}
|
||||
if cache, err := core.Inject[contracts.CacheService](ctx); err == nil && cache != nil {
|
||||
setCacheService(cache)
|
||||
} else {
|
||||
core.When[contracts.CacheService](ctx, func(cache contracts.CacheService) {
|
||||
setCacheService(cache)
|
||||
})
|
||||
}
|
||||
ctx.OnDispose(func() error {
|
||||
setDBService(nil)
|
||||
setCacheService(nil)
|
||||
return nil
|
||||
})
|
||||
|
||||
// 1. Register migrations
|
||||
ctx.Migrations().Register("auth", authMigrations)
|
||||
|
||||
|
||||
@@ -19,9 +19,24 @@ import (
|
||||
"Wavelet/core"
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/plugins/domain/auth"
|
||||
db "Wavelet/plugins/infra/database"
|
||||
)
|
||||
|
||||
type mockDBService struct {
|
||||
db *gorm.DB
|
||||
}
|
||||
|
||||
func (m *mockDBService) GORM() *gorm.DB {
|
||||
return m.db
|
||||
}
|
||||
|
||||
func (m *mockDBService) DB(ctx context.Context) *gorm.DB {
|
||||
return m.db.WithContext(ctx)
|
||||
}
|
||||
|
||||
func (m *mockDBService) Named(_ string) *gorm.DB {
|
||||
return m.db
|
||||
}
|
||||
|
||||
type testUser struct {
|
||||
ID uint64 `gorm:"primaryKey"`
|
||||
Username string
|
||||
@@ -60,7 +75,6 @@ func setupTestDB(t *testing.T) *gorm.DB {
|
||||
&auth.ExternalAccount{},
|
||||
))
|
||||
|
||||
db.SetDB(testDB)
|
||||
return testDB
|
||||
}
|
||||
|
||||
@@ -81,6 +95,7 @@ func (m *mockProvider) ExchangeCode(ctx context.Context, code string) (*contract
|
||||
func TestAuthPluginUnit(t *testing.T) {
|
||||
ctx := core.NewContext(context.Background())
|
||||
testDB := setupTestDB(t)
|
||||
core.Provide[contracts.DBService](ctx, &mockDBService{db: testDB})
|
||||
|
||||
p := auth.New()
|
||||
assert.Equal(t, "auth", p.Name())
|
||||
|
||||
@@ -5,14 +5,64 @@ package auth
|
||||
|
||||
import (
|
||||
"context"
|
||||
"sync"
|
||||
|
||||
db "Wavelet/plugins/infra/database"
|
||||
"gorm.io/gorm"
|
||||
|
||||
"Wavelet/core"
|
||||
"Wavelet/core/contracts"
|
||||
)
|
||||
|
||||
var (
|
||||
dbMu sync.RWMutex
|
||||
dbSvc contracts.DBService
|
||||
cacheMu sync.RWMutex
|
||||
cacheSvc contracts.CacheService
|
||||
)
|
||||
|
||||
func setDBService(s contracts.DBService) {
|
||||
dbMu.Lock()
|
||||
defer dbMu.Unlock()
|
||||
dbSvc = s
|
||||
}
|
||||
|
||||
func setCacheService(s contracts.CacheService) {
|
||||
cacheMu.Lock()
|
||||
defer cacheMu.Unlock()
|
||||
cacheSvc = s
|
||||
}
|
||||
|
||||
func getDB(ctx context.Context) *gorm.DB {
|
||||
if c, ok := ctx.(*core.Context); ok && c != nil {
|
||||
if s, err := core.Inject[contracts.DBService](c); err == nil && s != nil {
|
||||
return s.DB(ctx)
|
||||
}
|
||||
}
|
||||
dbMu.RLock()
|
||||
s := dbSvc
|
||||
dbMu.RUnlock()
|
||||
if s != nil {
|
||||
return s.DB(ctx)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func getCache(ctx context.Context) contracts.CacheService {
|
||||
if c, ok := ctx.(*core.Context); ok && c != nil {
|
||||
if s, err := core.Inject[contracts.CacheService](c); err == nil && s != nil {
|
||||
return s
|
||||
}
|
||||
}
|
||||
cacheMu.RLock()
|
||||
s := cacheSvc
|
||||
cacheMu.RUnlock()
|
||||
return s
|
||||
}
|
||||
|
||||
// GetAuthSourceByID 根据 ID 获取认证源
|
||||
func GetAuthSourceByID(ctx context.Context, id uint64) (*AuthSource, error) {
|
||||
var src AuthSource
|
||||
if err := db.DB(ctx).First(&src, id).Error; err != nil {
|
||||
if err := getDB(ctx).First(&src, id).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &src, nil
|
||||
@@ -21,7 +71,7 @@ func GetAuthSourceByID(ctx context.Context, id uint64) (*AuthSource, error) {
|
||||
// GetAuthSourceByName 根据名称获取认证源
|
||||
func GetAuthSourceByName(ctx context.Context, name string) (*AuthSource, error) {
|
||||
var src AuthSource
|
||||
if err := db.DB(ctx).Where("name = ?", name).First(&src).Error; err != nil {
|
||||
if err := getDB(ctx).Where("name = ?", name).First(&src).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &src, nil
|
||||
@@ -30,7 +80,7 @@ func GetAuthSourceByName(ctx context.Context, name string) (*AuthSource, error)
|
||||
// ListActiveAuthSources 获取所有启用的认证源
|
||||
func ListActiveAuthSources(ctx context.Context) ([]AuthSource, error) {
|
||||
var sources []AuthSource
|
||||
if err := db.DB(ctx).Where("is_active = ?", true).Order("id ASC").Find(&sources).Error; err != nil {
|
||||
if err := getDB(ctx).Where("is_active = ?", true).Order("id ASC").Find(&sources).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return sources, nil
|
||||
@@ -49,7 +99,7 @@ func GetAuthSourceByNameCached(ctx context.Context, name string) (*AuthSource, e
|
||||
// FindExternalAccount 查询指定认证源的外部账号绑定
|
||||
func FindExternalAccount(ctx context.Context, authSourceID uint64, externalID string) (*ExternalAccount, error) {
|
||||
var account ExternalAccount
|
||||
if err := db.DB(ctx).Where("auth_source_id = ? AND external_id = ?", authSourceID, externalID).First(&account).Error; err != nil {
|
||||
if err := getDB(ctx).Where("auth_source_id = ? AND external_id = ?", authSourceID, externalID).First(&account).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &account, nil
|
||||
@@ -57,13 +107,13 @@ func FindExternalAccount(ctx context.Context, authSourceID uint64, externalID st
|
||||
|
||||
// BindExternalAccount 绑定外部账号
|
||||
func BindExternalAccount(ctx context.Context, account *ExternalAccount) error {
|
||||
return db.DB(ctx).Create(account).Error
|
||||
return getDB(ctx).Create(account).Error
|
||||
}
|
||||
|
||||
// ListExternalAccountsByUserID 获取用户绑定的所有外部账号
|
||||
func ListExternalAccountsByUserID(ctx context.Context, userID uint64) ([]ExternalAccount, error) {
|
||||
var accounts []ExternalAccount
|
||||
if err := db.DB(ctx).Where("user_id = ?", userID).Find(&accounts).Error; err != nil {
|
||||
if err := getDB(ctx).Where("user_id = ?", userID).Find(&accounts).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return accounts, nil
|
||||
@@ -71,5 +121,5 @@ func ListExternalAccountsByUserID(ctx context.Context, userID uint64) ([]Externa
|
||||
|
||||
// UnbindExternalAccount 解绑外部账号
|
||||
func UnbindExternalAccount(ctx context.Context, id uint64, userID uint64) error {
|
||||
return db.DB(ctx).Where("id = ? AND user_id = ?", id, userID).Delete(&ExternalAccount{}).Error
|
||||
return getDB(ctx).Where("id = ? AND user_id = ?", id, userID).Delete(&ExternalAccount{}).Error
|
||||
}
|
||||
|
||||
@@ -8,10 +8,10 @@ import (
|
||||
"errors"
|
||||
"sync"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/pkg/util"
|
||||
db "Wavelet/plugins/infra/database"
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
type authServiceImpl struct{}
|
||||
@@ -57,7 +57,7 @@ func (s *authServiceImpl) VerifyToken(ctx context.Context, token string) (*contr
|
||||
UserID uint64
|
||||
IsAdmin bool
|
||||
}
|
||||
if err := db.DB(ctx).Table("w_access_tokens").Where("token_hash = ?", tokenHash).First(&tokenRow).Error; err != nil {
|
||||
if err := getDB(ctx).Table("w_access_tokens").Where("token_hash = ?", tokenHash).First(&tokenRow).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
tokenRecord = &CachedToken{
|
||||
@@ -71,7 +71,7 @@ func (s *authServiceImpl) VerifyToken(ctx context.Context, token string) (*contr
|
||||
user, err := GetCachedUser(ctx, tokenRecord.UserID)
|
||||
if err != nil || user == nil || !user.IsActive {
|
||||
var dbUser contracts.UserDTO
|
||||
if err := db.DB(ctx).Table("w_users").Where("id = ? AND is_active = ?", tokenRecord.UserID, true).First(&dbUser).Error; err != nil {
|
||||
if err := getDB(ctx).Table("w_users").Where("id = ? AND is_active = ?", tokenRecord.UserID, true).First(&dbUser).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
user = &dbUser
|
||||
@@ -120,7 +120,7 @@ func (s *authServiceImpl) InvalidateCachedToken(ctx context.Context, tokenHash s
|
||||
|
||||
func (s *authServiceImpl) ListAuthSources(ctx context.Context) ([]contracts.AuthSourceViewDTO, error) {
|
||||
var sources []AuthSource
|
||||
if err := db.DB(ctx).Order("id ASC").Find(&sources).Error; err != nil {
|
||||
if err := getDB(ctx).Order("id ASC").Find(&sources).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
@@ -157,7 +157,7 @@ func (s *authServiceImpl) CreateAuthSource(ctx context.Context, source contracts
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if err := db.DB(ctx).Create(&model).Error; err != nil {
|
||||
if err := getDB(ctx).Create(&model).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
@@ -167,7 +167,7 @@ func (s *authServiceImpl) CreateAuthSource(ctx context.Context, source contracts
|
||||
|
||||
func (s *authServiceImpl) UpdateAuthSource(ctx context.Context, id uint64, source contracts.AuthSourceDTO) (*contracts.AuthSourceDTO, error) {
|
||||
var existing AuthSource
|
||||
if err := db.DB(ctx).First(&existing, id).Error; err != nil {
|
||||
if err := getDB(ctx).First(&existing, id).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
@@ -184,7 +184,7 @@ func (s *authServiceImpl) UpdateAuthSource(ctx context.Context, id uint64, sourc
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if err := db.DB(ctx).Save(&existing).Error; err != nil {
|
||||
if err := getDB(ctx).Save(&existing).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
@@ -194,21 +194,21 @@ func (s *authServiceImpl) UpdateAuthSource(ctx context.Context, id uint64, sourc
|
||||
|
||||
func (s *authServiceImpl) DeleteAuthSource(ctx context.Context, id uint64) error {
|
||||
var existing AuthSource
|
||||
if err := db.DB(ctx).First(&existing, id).Error; err != nil {
|
||||
if err := getDB(ctx).First(&existing, id).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
return db.DB(ctx).Delete(&existing).Error
|
||||
return getDB(ctx).Delete(&existing).Error
|
||||
}
|
||||
|
||||
func (s *authServiceImpl) ToggleAuthSource(ctx context.Context, id uint64) (*contracts.AuthSourceDTO, error) {
|
||||
var existing AuthSource
|
||||
if err := db.DB(ctx).First(&existing, id).Error; err != nil {
|
||||
if err := getDB(ctx).First(&existing, id).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
existing.IsActive = !existing.IsActive
|
||||
if err := db.DB(ctx).Save(&existing).Error; err != nil {
|
||||
if err := getDB(ctx).Save(&existing).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
|
||||
@@ -11,13 +11,13 @@ import (
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/pkg/config"
|
||||
db "Wavelet/plugins/infra/database"
|
||||
"github.com/gin-contrib/sessions"
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/google/uuid"
|
||||
gsessions "github.com/gorilla/sessions"
|
||||
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/pkg/config"
|
||||
)
|
||||
|
||||
// GetSessionOptions 根据配置构建 Session 选项
|
||||
@@ -118,7 +118,7 @@ func SetLoginSession(ctx context.Context, c *gin.Context, user *contracts.UserDT
|
||||
isSessionCookie := false
|
||||
|
||||
var val string
|
||||
if err := db.DB(ctx).Table("w_system_configs").Where("key = ?", "login_session_ttl_hours").Pluck("value", &val).Error; err == nil && val != "" {
|
||||
if err := getDB(ctx).Table("w_system_configs").Where("key = ?", "login_session_ttl_hours").Pluck("value", &val).Error; err == nil && val != "" {
|
||||
if ttlHours, err := strconv.Atoi(val); err == nil {
|
||||
switch {
|
||||
case ttlHours == -1:
|
||||
|
||||
@@ -6,10 +6,11 @@ package cap
|
||||
import (
|
||||
"net/http"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"Wavelet/pkg/logger"
|
||||
"Wavelet/pkg/response"
|
||||
"Wavelet/plugins/domain/cap/pow"
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
// ChallengeResponse is a local type alias for the pow.ChallengeResponse struct
|
||||
|
||||
@@ -15,7 +15,6 @@ import (
|
||||
|
||||
"Wavelet/pkg/config"
|
||||
"Wavelet/plugins/domain/cap/pow"
|
||||
db "Wavelet/plugins/infra/cache"
|
||||
)
|
||||
|
||||
const (
|
||||
@@ -186,13 +185,7 @@ func GetDefaultManager() *Manager {
|
||||
return
|
||||
}
|
||||
|
||||
var store pow.Store
|
||||
if config.Config != nil && config.Config.Redis.Enabled && db.Redis != nil {
|
||||
store = pow.NewRedisStore(db.Redis)
|
||||
} else {
|
||||
store = pow.NewMemoryStore(1 * time.Minute)
|
||||
}
|
||||
|
||||
store := pow.NewMemoryStore(1 * time.Minute)
|
||||
defaultManager = NewManager(secret, store)
|
||||
})
|
||||
return defaultManager
|
||||
|
||||
@@ -29,7 +29,6 @@ func (p *Plugin) Name() string {
|
||||
func (p *Plugin) Inject() []reflect.Type {
|
||||
return []reflect.Type{
|
||||
reflect.TypeFor[contracts.DBService](),
|
||||
reflect.TypeFor[contracts.CacheService](),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -45,6 +44,24 @@ func (p *Plugin) Manifest() core.Manifest {
|
||||
|
||||
// Apply registers the cap routes and settings into the Context.
|
||||
func (p *Plugin) Apply(ctx *core.Context) error {
|
||||
// 0. Bind DBService from Context
|
||||
if db, err := core.Inject[contracts.DBService](ctx); err == nil && db != nil {
|
||||
setDBService(db)
|
||||
} else {
|
||||
core.When[contracts.DBService](ctx, func(db contracts.DBService) {
|
||||
setDBService(db)
|
||||
})
|
||||
}
|
||||
ctx.OnDispose(func() error {
|
||||
setDBService(nil)
|
||||
return nil
|
||||
})
|
||||
|
||||
// Listen to system config changed events to invalidate cached settings
|
||||
ctx.Events().On(contracts.EventTopicConfigChanged, func(_ any) {
|
||||
InvalidateRuntimeSettings()
|
||||
})
|
||||
|
||||
// Register HTTP Routes
|
||||
capGroup := ctx.Router().Group("/api/v1/cap")
|
||||
{
|
||||
|
||||
@@ -5,7 +5,6 @@ package cap
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"strconv"
|
||||
"sync"
|
||||
@@ -13,12 +12,38 @@ import (
|
||||
"time"
|
||||
|
||||
"golang.org/x/sync/singleflight"
|
||||
"gorm.io/gorm"
|
||||
|
||||
"Wavelet/pkg/util"
|
||||
cachepkg "Wavelet/plugins/infra/cache"
|
||||
database "Wavelet/plugins/infra/database"
|
||||
"Wavelet/core"
|
||||
"Wavelet/core/contracts"
|
||||
)
|
||||
|
||||
var (
|
||||
dbMu sync.RWMutex
|
||||
dbSvc contracts.DBService
|
||||
)
|
||||
|
||||
func setDBService(s contracts.DBService) {
|
||||
dbMu.Lock()
|
||||
defer dbMu.Unlock()
|
||||
dbSvc = s
|
||||
}
|
||||
|
||||
func getDB(ctx context.Context) *gorm.DB {
|
||||
if c, ok := ctx.(*core.Context); ok && c != nil {
|
||||
if s, err := core.Inject[contracts.DBService](c); err == nil && s != nil {
|
||||
return s.DB(ctx)
|
||||
}
|
||||
}
|
||||
dbMu.RLock()
|
||||
s := dbSvc
|
||||
dbMu.RUnlock()
|
||||
if s != nil {
|
||||
return s.DB(ctx)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
const (
|
||||
defaultChallengeCount = 1
|
||||
defaultChallengeSize = 32
|
||||
@@ -67,9 +92,8 @@ var runtimeConfigKeySet = func() map[string]struct{} {
|
||||
}()
|
||||
|
||||
type runtimeSettingsStore struct {
|
||||
snapshot atomic.Pointer[RuntimeSettings]
|
||||
loadGroup singleflight.Group
|
||||
listenerOnce sync.Once
|
||||
snapshot atomic.Pointer[RuntimeSettings]
|
||||
loadGroup singleflight.Group
|
||||
}
|
||||
|
||||
var settingsStore = &runtimeSettingsStore{}
|
||||
@@ -148,7 +172,11 @@ func loadRuntimeSettings(ctx context.Context) (RuntimeSettings, error) {
|
||||
Value string `gorm:"column:value"`
|
||||
}
|
||||
var records []configRecord
|
||||
if err := database.DB(ctx).Table("w_system_configs").Where("key IN ?", runtimeConfigKeys).Find(&records).Error; err != nil {
|
||||
db := getDB(ctx)
|
||||
if db == nil {
|
||||
return parseRuntimeSettings(nil), nil
|
||||
}
|
||||
if err := db.Table("w_system_configs").Where("key IN ?", runtimeConfigKeys).Find(&records).Error; err != nil {
|
||||
return RuntimeSettings{}, err
|
||||
}
|
||||
configs := make(map[string]string, len(records))
|
||||
@@ -167,6 +195,10 @@ func parseRuntimeSettings(configs map[string]string) RuntimeSettings {
|
||||
TokenTTL: defaultTokenTTL,
|
||||
}
|
||||
|
||||
if len(configs) == 0 {
|
||||
return settings
|
||||
}
|
||||
|
||||
if val, ok := configs[ConfigKeyCapLoginEnabled]; ok {
|
||||
if enabled, err := strconv.ParseBool(val); err == nil {
|
||||
settings.LoginEnabled = enabled
|
||||
@@ -201,36 +233,4 @@ func parseRuntimeSettings(configs map[string]string) RuntimeSettings {
|
||||
return settings
|
||||
}
|
||||
|
||||
func (s *runtimeSettingsStore) ensureInvalidationListener() {
|
||||
s.listenerOnce.Do(startRuntimeSettingsInvalidationListener)
|
||||
}
|
||||
|
||||
// SystemConfigInvalidationChannel 系统配置失效广播通道
|
||||
const SystemConfigInvalidationChannel = "system_config:invalidation"
|
||||
|
||||
func startRuntimeSettingsInvalidationListener() {
|
||||
rdb := cachepkg.Redis
|
||||
if rdb == nil {
|
||||
return
|
||||
}
|
||||
|
||||
util.Go(func() {
|
||||
pubsub := rdb.Subscribe(context.Background(), SystemConfigInvalidationChannel)
|
||||
defer func() {
|
||||
_ = pubsub.Close()
|
||||
}()
|
||||
|
||||
for msg := range pubsub.Channel() {
|
||||
var payload struct {
|
||||
Key string `json:"key"`
|
||||
}
|
||||
if err := json.Unmarshal([]byte(msg.Payload), &payload); err != nil {
|
||||
InvalidateRuntimeSettings()
|
||||
continue
|
||||
}
|
||||
if payload.Key == "" || payload.Key == "*" || IsRuntimeConfigKey(payload.Key) {
|
||||
InvalidateRuntimeSettings()
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
func (s *runtimeSettingsStore) ensureInvalidationListener() {}
|
||||
|
||||
@@ -7,8 +7,9 @@ import (
|
||||
"net/http"
|
||||
"strconv"
|
||||
|
||||
"Wavelet/pkg/response"
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"Wavelet/pkg/response"
|
||||
)
|
||||
|
||||
// ListAdminChannelDefinitions returns form schemas for supported channel types.
|
||||
|
||||
@@ -5,13 +5,14 @@
|
||||
package qq
|
||||
|
||||
import (
|
||||
"Wavelet/pkg/util"
|
||||
"context"
|
||||
"fmt"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"Wavelet/pkg/util"
|
||||
|
||||
"Wavelet/pkg/logger"
|
||||
"Wavelet/plugins/domain/message_gateway"
|
||||
|
||||
|
||||
@@ -12,9 +12,10 @@ import (
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
tele "gopkg.in/telebot.v4"
|
||||
|
||||
"Wavelet/pkg/util"
|
||||
"Wavelet/plugins/domain/message_gateway"
|
||||
tele "gopkg.in/telebot.v4"
|
||||
)
|
||||
|
||||
// Adapter is a Telegram private-chat channel.
|
||||
|
||||
@@ -7,8 +7,9 @@ import (
|
||||
"context"
|
||||
"testing"
|
||||
|
||||
"Wavelet/plugins/domain/message_gateway"
|
||||
tele "gopkg.in/telebot.v4"
|
||||
|
||||
"Wavelet/plugins/domain/message_gateway"
|
||||
)
|
||||
|
||||
func TestHandleUpdate_DropsGroups(t *testing.T) {
|
||||
|
||||
@@ -0,0 +1,78 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package message_gateway
|
||||
|
||||
import (
|
||||
"context"
|
||||
"sync"
|
||||
|
||||
"gorm.io/gorm"
|
||||
|
||||
"Wavelet/core"
|
||||
"Wavelet/core/contracts"
|
||||
)
|
||||
|
||||
var (
|
||||
dbMu sync.RWMutex
|
||||
dbSvc contracts.DBService
|
||||
cacheMu sync.RWMutex
|
||||
cacheSvc contracts.CacheService
|
||||
taskMu sync.RWMutex
|
||||
taskSvc contracts.TaskService
|
||||
)
|
||||
|
||||
func SetDBServiceForTest(s contracts.DBService) {
|
||||
setDBService(s)
|
||||
}
|
||||
|
||||
func setDBService(s contracts.DBService) {
|
||||
dbMu.Lock()
|
||||
defer dbMu.Unlock()
|
||||
dbSvc = s
|
||||
}
|
||||
|
||||
func setCacheService(s contracts.CacheService) {
|
||||
cacheMu.Lock()
|
||||
defer cacheMu.Unlock()
|
||||
cacheSvc = s
|
||||
}
|
||||
|
||||
func setTaskService(s contracts.TaskService) {
|
||||
taskMu.Lock()
|
||||
defer taskMu.Unlock()
|
||||
taskSvc = s
|
||||
}
|
||||
|
||||
func getDB(ctx context.Context) *gorm.DB {
|
||||
if c, ok := ctx.(*core.Context); ok && c != nil {
|
||||
if s, err := core.Inject[contracts.DBService](c); err == nil && s != nil {
|
||||
return s.DB(ctx)
|
||||
}
|
||||
}
|
||||
dbMu.RLock()
|
||||
s := dbSvc
|
||||
dbMu.RUnlock()
|
||||
if s != nil {
|
||||
return s.DB(ctx)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func getCache(ctx context.Context) contracts.CacheService {
|
||||
if c, ok := ctx.(*core.Context); ok && c != nil {
|
||||
if s, err := core.Inject[contracts.CacheService](c); err == nil && s != nil {
|
||||
return s
|
||||
}
|
||||
}
|
||||
cacheMu.RLock()
|
||||
s := cacheSvc
|
||||
cacheMu.RUnlock()
|
||||
return s
|
||||
}
|
||||
|
||||
func getTaskService() contracts.TaskService {
|
||||
taskMu.RLock()
|
||||
defer taskMu.RUnlock()
|
||||
return taskSvc
|
||||
}
|
||||
@@ -8,10 +8,11 @@ import (
|
||||
"net/http"
|
||||
"strconv"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/pkg/response"
|
||||
"Wavelet/pkg/util"
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
func currentUser(c *gin.Context) (*contracts.UserDTO, bool) {
|
||||
|
||||
@@ -8,13 +8,35 @@ import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"gorm.io/gorm"
|
||||
|
||||
"Wavelet/pkg/testhelper"
|
||||
"Wavelet/plugins/domain/message_gateway"
|
||||
)
|
||||
|
||||
type mockDBService struct {
|
||||
db *gorm.DB
|
||||
}
|
||||
|
||||
func (m *mockDBService) GORM() *gorm.DB {
|
||||
return m.db
|
||||
}
|
||||
|
||||
func (m *mockDBService) DB(ctx context.Context) *gorm.DB {
|
||||
return m.db.WithContext(ctx)
|
||||
}
|
||||
|
||||
func (m *mockDBService) Named(_ string) *gorm.DB {
|
||||
return m.db
|
||||
}
|
||||
|
||||
func TestUpsertPairingCode_ReusesUnexpired(t *testing.T) {
|
||||
_, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
defer cleanup()
|
||||
testDB, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
message_gateway.SetDBServiceForTest(&mockDBService{db: testDB})
|
||||
defer func() {
|
||||
message_gateway.SetDBServiceForTest(nil)
|
||||
cleanup()
|
||||
}()
|
||||
ctx := context.Background()
|
||||
first, err := message_gateway.UpsertPairingCode(ctx, 1, "tg-1", "ABCD1234", time.Now().Add(15*time.Minute))
|
||||
if err != nil {
|
||||
|
||||
@@ -9,12 +9,12 @@ import (
|
||||
"embed"
|
||||
"reflect"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"Wavelet/core"
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/core/extpoints"
|
||||
"Wavelet/pkg/util"
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/hibiken/asynq"
|
||||
)
|
||||
|
||||
//go:embed migrations/*.sql
|
||||
@@ -80,6 +80,35 @@ type PushNotificationEvent struct {
|
||||
|
||||
// Apply registers message_gateway migrations, routes, tasks, schedules, events, and settings into the Context.
|
||||
func (p *Plugin) Apply(ctx *core.Context) error {
|
||||
// 0. Bind DBService, CacheService, TaskService
|
||||
if db, err := core.Inject[contracts.DBService](ctx); err == nil && db != nil {
|
||||
setDBService(db)
|
||||
} else {
|
||||
core.When[contracts.DBService](ctx, func(db contracts.DBService) {
|
||||
setDBService(db)
|
||||
})
|
||||
}
|
||||
if cache, err := core.Inject[contracts.CacheService](ctx); err == nil && cache != nil {
|
||||
setCacheService(cache)
|
||||
} else {
|
||||
core.When[contracts.CacheService](ctx, func(cache contracts.CacheService) {
|
||||
setCacheService(cache)
|
||||
})
|
||||
}
|
||||
if taskSvc, err := core.Inject[contracts.TaskService](ctx); err == nil && taskSvc != nil {
|
||||
setTaskService(taskSvc)
|
||||
} else {
|
||||
core.When[contracts.TaskService](ctx, func(taskSvc contracts.TaskService) {
|
||||
setTaskService(taskSvc)
|
||||
})
|
||||
}
|
||||
ctx.OnDispose(func() error {
|
||||
setDBService(nil)
|
||||
setCacheService(nil)
|
||||
setTaskService(nil)
|
||||
return nil
|
||||
})
|
||||
|
||||
// 0. Resolve auth service for middleware (via IoC, not direct import)
|
||||
var loginMW gin.HandlerFunc = func(c *gin.Context) { c.Next() }
|
||||
var adminMW gin.HandlerFunc = func(c *gin.Context) { c.Next() }
|
||||
@@ -145,18 +174,16 @@ func (p *Plugin) Apply(ctx *core.Context) error {
|
||||
const defaultTaskRetry = 3
|
||||
pushHandler := &PushHandler{}
|
||||
|
||||
// 5. Register Asynq background tasks
|
||||
ctx.Task().Register("message_gateway:push_notification", func(c context.Context, t *asynq.Task) error {
|
||||
_, err := pushHandler.Execute(c, t.Payload())
|
||||
return err
|
||||
// 5. Register background tasks
|
||||
ctx.Task().Register("message_gateway:push_notification", func(c context.Context, payload []byte) error {
|
||||
return pushHandler.Execute(c, payload)
|
||||
}, extpoints.WithTaskRetry(defaultTaskRetry))
|
||||
|
||||
ctx.Task().Register(SendNotificationTask, func(c context.Context, t *asynq.Task) error {
|
||||
_, err := pushHandler.Execute(c, t.Payload())
|
||||
return err
|
||||
ctx.Task().Register(SendNotificationTask, func(c context.Context, payload []byte) error {
|
||||
return pushHandler.Execute(c, payload)
|
||||
}, extpoints.WithTaskRetry(defaultTaskRetry))
|
||||
|
||||
ctx.Task().Register("message_gateway:dispatch_bot_msg", func(_ context.Context, _ *asynq.Task) error {
|
||||
ctx.Task().Register("message_gateway:dispatch_bot_msg", func(_ context.Context, _ []byte) error {
|
||||
return nil
|
||||
})
|
||||
|
||||
@@ -184,9 +211,14 @@ func (p *Plugin) Apply(ctx *core.Context) error {
|
||||
return nil
|
||||
})
|
||||
|
||||
// 8. Register built-in domain events and task listeners
|
||||
// 8. Register task completed event listener
|
||||
ctx.Events().On(contracts.EventTopicTaskCompleted, func(c context.Context, e contracts.TaskCompletedEvent) error {
|
||||
handleTaskCompleted(c, e)
|
||||
return nil
|
||||
})
|
||||
|
||||
// 9. Register built-in domain events
|
||||
RegisterCustomEvents()
|
||||
RegisterTaskListeners()
|
||||
|
||||
// 9. Register Settings Schemas
|
||||
ctx.Settings().Register(extpoints.SettingSchema{
|
||||
|
||||
@@ -11,10 +11,11 @@ import (
|
||||
"strings"
|
||||
"sync"
|
||||
|
||||
"Wavelet/pkg/response"
|
||||
pkgpush "Wavelet/plugins/domain/message_gateway/push"
|
||||
"github.com/gin-gonic/gin"
|
||||
"gorm.io/gorm"
|
||||
|
||||
"Wavelet/pkg/response"
|
||||
pkgpush "Wavelet/plugins/domain/message_gateway/push"
|
||||
)
|
||||
|
||||
const (
|
||||
|
||||
@@ -13,8 +13,9 @@ import (
|
||||
|
||||
pkgpush "Wavelet/plugins/domain/message_gateway/push"
|
||||
|
||||
"Wavelet/pkg/util"
|
||||
"gorm.io/gorm"
|
||||
|
||||
"Wavelet/pkg/util"
|
||||
)
|
||||
|
||||
// NotificationMessage represents the structured notification message payload.
|
||||
|
||||
@@ -11,9 +11,10 @@ import (
|
||||
|
||||
pkgpush "Wavelet/plugins/domain/message_gateway/push"
|
||||
|
||||
"Wavelet/pkg/response"
|
||||
"github.com/gin-gonic/gin"
|
||||
"gorm.io/gorm"
|
||||
|
||||
"Wavelet/pkg/response"
|
||||
)
|
||||
|
||||
// UpdatePushEventRequest is the request body for updating a push event.
|
||||
|
||||
@@ -11,11 +11,10 @@ import (
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
"gorm.io/gorm"
|
||||
|
||||
"Wavelet/core/contracts"
|
||||
pkgpush "Wavelet/plugins/domain/message_gateway/push"
|
||||
"Wavelet/plugins/drivers/driver_asynq_worker"
|
||||
db "Wavelet/plugins/infra/database"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
type smtpConfig struct {
|
||||
@@ -28,10 +27,10 @@ type smtpConfig struct {
|
||||
func loadSMTPConfig(ctx context.Context) smtpConfig {
|
||||
var cfg smtpConfig
|
||||
var host, port, user, pass string
|
||||
_ = db.DB(ctx).Table("w_system_configs").Where("key = ?", "smtp_host").Pluck("value", &host).Error
|
||||
_ = db.DB(ctx).Table("w_system_configs").Where("key = ?", "smtp_port").Pluck("value", &port).Error
|
||||
_ = db.DB(ctx).Table("w_system_configs").Where("key = ?", "smtp_username").Pluck("value", &user).Error
|
||||
_ = db.DB(ctx).Table("w_system_configs").Where("key = ?", "smtp_password").Pluck("value", &pass).Error
|
||||
_ = getDB(ctx).Table("w_system_configs").Where("key = ?", "smtp_host").Pluck("value", &host).Error
|
||||
_ = getDB(ctx).Table("w_system_configs").Where("key = ?", "smtp_port").Pluck("value", &port).Error
|
||||
_ = getDB(ctx).Table("w_system_configs").Where("key = ?", "smtp_username").Pluck("value", &user).Error
|
||||
_ = getDB(ctx).Table("w_system_configs").Where("key = ?", "smtp_password").Pluck("value", &pass).Error
|
||||
cfg.Host = host
|
||||
cfg.Port = port
|
||||
cfg.Username = user
|
||||
@@ -263,14 +262,14 @@ func loadUserFromPayload(ctx context.Context, data map[string]any) any {
|
||||
|
||||
if userID, ok := extractUserID(data); ok && userID > 0 {
|
||||
var user contracts.UserDTO
|
||||
if err := db.DB(ctx).Table("w_users").Where("id = ?", userID).First(&user).Error; err == nil {
|
||||
if err := getDB(ctx).Table("w_users").Where("id = ?", userID).First(&user).Error; err == nil {
|
||||
return &user
|
||||
}
|
||||
}
|
||||
|
||||
if username := extractUsername(data); username != "" {
|
||||
var user contracts.UserDTO
|
||||
if err := db.DB(ctx).Table("w_users").Where("username = ?", username).First(&user).Error; err == nil {
|
||||
if err := getDB(ctx).Table("w_users").Where("username = ?", username).First(&user).Error; err == nil {
|
||||
return &user
|
||||
}
|
||||
}
|
||||
@@ -372,11 +371,11 @@ func resolveDynamicKeyword(target string, flatBody map[string]any) string {
|
||||
func resolveTargetUser(ctx context.Context, resolved string, _ string) (contracts.UserDTO, bool) {
|
||||
var user contracts.UserDTO
|
||||
if id, err := strconv.ParseUint(resolved, 10, 64); err == nil {
|
||||
if err := db.DB(ctx).Table("w_users").Where("id = ?", id).First(&user).Error; err == nil {
|
||||
if err := getDB(ctx).Table("w_users").Where("id = ?", id).First(&user).Error; err == nil {
|
||||
return user, true
|
||||
}
|
||||
}
|
||||
if err := db.DB(ctx).Table("w_users").Where("username = ?", resolved).First(&user).Error; err == nil {
|
||||
if err := getDB(ctx).Table("w_users").Where("username = ?", resolved).First(&user).Error; err == nil {
|
||||
return user, true
|
||||
}
|
||||
return user, false
|
||||
@@ -387,7 +386,7 @@ func resolveSystemTarget(ctx context.Context, resolved string, channel string) (
|
||||
return "", false
|
||||
}
|
||||
var adminUser contracts.UserDTO
|
||||
if err := db.DB(ctx).Table("w_users").Where("is_admin = ?", true).Order("id ASC").First(&adminUser).Error; err != nil {
|
||||
if err := getDB(ctx).Table("w_users").Where("is_admin = ?", true).Order("id ASC").First(&adminUser).Error; err != nil {
|
||||
return resolved, true
|
||||
}
|
||||
if channel == channelEmail && adminUser.Email != "" {
|
||||
@@ -425,7 +424,7 @@ func resolveSMTPConfig(ctx context.Context, url, token, other string) (string, s
|
||||
|
||||
func getSystemUser(ctx context.Context) *contracts.UserDTO {
|
||||
var user contracts.UserDTO
|
||||
if err := db.DB(ctx).Table("w_users").Where("is_admin = ?", true).Order("id ASC").First(&user).Error; err == nil {
|
||||
if err := getDB(ctx).Table("w_users").Where("is_admin = ?", true).Order("id ASC").First(&user).Error; err == nil {
|
||||
return &user
|
||||
}
|
||||
return &contracts.UserDTO{
|
||||
@@ -445,14 +444,16 @@ func findBuiltInEvent(key string) (EventMetadata, bool) {
|
||||
|
||||
func getEventInfo(req CreatePushEventRequest) (string, string, []byte, error) {
|
||||
if req.TaskType != "" {
|
||||
meta := driver_asynq_worker.GetTaskMetaByAsynqTask(req.TaskType)
|
||||
if meta == nil {
|
||||
return "", "", nil, errors.New("unsupported task type")
|
||||
taskName := req.TaskType
|
||||
if taskSvc := getTaskService(); taskSvc != nil {
|
||||
if meta, ok := taskSvc.GetTaskMeta(req.TaskType); ok {
|
||||
taskName = meta.DisplayName
|
||||
}
|
||||
}
|
||||
eventKey := "task_completed:" + req.TaskType
|
||||
eventName := "任务完成: " + meta.Name
|
||||
eventName := "任务完成: " + taskName
|
||||
defaultTemplate := NotificationMessage{
|
||||
Title: "任务完成: " + meta.Name,
|
||||
Title: "任务完成: " + taskName,
|
||||
Content: "异步任务 {{task_name}} (ID: {{task_id}}) 已完成。状态: {{task_status}},耗时: {{task_duration}} ms。",
|
||||
Level: defaultLevelInfo,
|
||||
}
|
||||
@@ -484,8 +485,11 @@ func enqueuePushTask(ctx context.Context, payload SendPayload) error {
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
_, err = driver_asynq_worker.DispatchTask(ctx, "send_notification", payloadBytes, "system")
|
||||
return err
|
||||
if taskSvc := getTaskService(); taskSvc != nil {
|
||||
_, err = taskSvc.Dispatch(ctx, "send_notification", payloadBytes, "system")
|
||||
return err
|
||||
}
|
||||
return errors.New("task service not available")
|
||||
}
|
||||
|
||||
func getFlatBody(body map[string]any) map[string]any {
|
||||
|
||||
@@ -9,19 +9,14 @@ import (
|
||||
"strconv"
|
||||
"time"
|
||||
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/pkg/logger"
|
||||
"Wavelet/plugins/drivers/driver_asynq_worker"
|
||||
)
|
||||
|
||||
// RegisterTaskListeners subscribes push notification handlers to task completion events.
|
||||
func RegisterTaskListeners() {
|
||||
driver_asynq_worker.OnTaskCompleted(handleTaskCompleted)
|
||||
}
|
||||
|
||||
func handleTaskCompleted(ctx context.Context, execution *driver_asynq_worker.TaskExecution, result *driver_asynq_worker.TaskResult, execErr error) {
|
||||
events, err := listActivePushEventsByTaskType(ctx, execution.TaskType)
|
||||
func handleTaskCompleted(ctx context.Context, e contracts.TaskCompletedEvent) {
|
||||
events, err := listActivePushEventsByTaskType(ctx, e.TaskType)
|
||||
if err != nil {
|
||||
logger.ErrorF(ctx, "push_task_completed_listener: failed to query push events for task type %s: %v", execution.TaskType, err)
|
||||
logger.ErrorF(ctx, "push_task_completed_listener: failed to query push events for task type %s: %v", e.TaskType, err)
|
||||
return
|
||||
}
|
||||
if len(events) == 0 {
|
||||
@@ -29,34 +24,26 @@ func handleTaskCompleted(ctx context.Context, execution *driver_asynq_worker.Tas
|
||||
}
|
||||
|
||||
body := map[string]any{
|
||||
"task_id": execution.TaskID,
|
||||
"task_name": execution.TaskName,
|
||||
"task_type": execution.TaskType,
|
||||
"task_status": string(execution.Status),
|
||||
"task_duration": execution.Duration,
|
||||
"task_id": e.TaskID,
|
||||
"task_name": e.TaskName,
|
||||
"task_type": e.TaskType,
|
||||
"task_status": e.Status,
|
||||
"task_duration": e.Duration,
|
||||
"time": time.Now().Format("2006-01-02 15:04:05"),
|
||||
}
|
||||
if execErr != nil {
|
||||
body["task_error"] = execErr.Error()
|
||||
} else {
|
||||
body["task_error"] = ""
|
||||
}
|
||||
if result != nil {
|
||||
body["task_result"] = result.Message
|
||||
} else {
|
||||
body["task_result"] = ""
|
||||
"task_error": e.ErrorMsg,
|
||||
"task_result": e.ResultMsg,
|
||||
}
|
||||
|
||||
var payloadMap map[string]any
|
||||
if execution.Payload != "" {
|
||||
if err := json.Unmarshal([]byte(execution.Payload), &payloadMap); err == nil {
|
||||
if e.Payload != "" {
|
||||
if err := json.Unmarshal([]byte(e.Payload), &payloadMap); err == nil {
|
||||
body["payload"] = payloadMap
|
||||
extractUserFromMap(ctx, payloadMap, body)
|
||||
}
|
||||
}
|
||||
if result != nil && result.Detail != "" {
|
||||
if e.Detail != "" {
|
||||
var detailMap map[string]any
|
||||
if err := json.Unmarshal([]byte(result.Detail), &detailMap); err == nil {
|
||||
if err := json.Unmarshal([]byte(e.Detail), &detailMap); err == nil {
|
||||
body["detail"] = detailMap
|
||||
extractUserFromMap(ctx, detailMap, body)
|
||||
}
|
||||
|
||||
@@ -9,8 +9,9 @@ import (
|
||||
"errors"
|
||||
"fmt"
|
||||
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/pkg/logger"
|
||||
"Wavelet/plugins/domain/message_gateway/push"
|
||||
"Wavelet/plugins/drivers/driver_asynq_worker"
|
||||
)
|
||||
|
||||
const (
|
||||
@@ -21,28 +22,24 @@ const (
|
||||
)
|
||||
|
||||
// SendNotificationMeta represents the task metadata.
|
||||
var SendNotificationMeta = driver_asynq_worker.TaskMeta{
|
||||
Type: TaskTypeSendNotification,
|
||||
AsynqTask: SendNotificationTask,
|
||||
Name: "推送通知",
|
||||
Description: "异步执行系统通知的多渠道派发与推送",
|
||||
SupportsTime: false,
|
||||
MaxRetry: driver_asynq_worker.DefaultMaxRetry,
|
||||
Queue: driver_asynq_worker.QueueDefault,
|
||||
Retryable: true,
|
||||
Params: []driver_asynq_worker.TaskParam{
|
||||
var SendNotificationMeta = contracts.TaskMetaDTO{
|
||||
Name: TaskTypeSendNotification,
|
||||
DisplayName: "推送通知",
|
||||
Description: "异步执行系统通知的多渠道派发与推送",
|
||||
MaxRetry: 3,
|
||||
Queue: "default",
|
||||
Params: []contracts.TaskParamDTO{
|
||||
{
|
||||
Name: "event_key",
|
||||
Label: "事件标识",
|
||||
Type: "string",
|
||||
Description: "事件标识 (如 admin_login)",
|
||||
Required: true,
|
||||
Placeholder: "admin_login",
|
||||
},
|
||||
{
|
||||
Name: "target",
|
||||
Label: "目标接收者",
|
||||
Type: "string",
|
||||
Required: false,
|
||||
Name: "target",
|
||||
Type: "string",
|
||||
Description: "目标接收者",
|
||||
Required: false,
|
||||
},
|
||||
},
|
||||
}
|
||||
@@ -69,23 +66,21 @@ func (h *PushHandler) ValidatePayload(payload []byte) ([]byte, error) {
|
||||
}
|
||||
|
||||
// Execute performs the push send and logs delivery history audit.
|
||||
func (h *PushHandler) Execute(ctx context.Context, payload []byte) (*driver_asynq_worker.TaskResult, error) {
|
||||
func (h *PushHandler) Execute(ctx context.Context, payload []byte) error {
|
||||
var req SendPayload
|
||||
if err := json.Unmarshal(payload, &req); err != nil {
|
||||
driver_asynq_worker.AppendLog(ctx, "解析推送参数失败: %v", err)
|
||||
return nil, fmt.Errorf("parse payload failed: %w", err)
|
||||
logger.ErrorF(ctx, "[Push] 解析推送参数失败: %v", err)
|
||||
return fmt.Errorf("parse payload failed: %w", err)
|
||||
}
|
||||
|
||||
driver_asynq_worker.AppendLog(ctx, "开始推送通知: 事件 = %s, 渠道 = %s, 接收目标 = %s", req.EventKey, req.Config.Channel, req.Target)
|
||||
logger.InfoF(ctx, "[Push] 开始推送通知: 事件 = %s, 渠道 = %s, 接收目标 = %s", req.EventKey, req.Config.Channel, req.Target)
|
||||
|
||||
pusher, err := push.GetPusher(req.Config.Channel)
|
||||
if err != nil {
|
||||
errWrap := fmt.Errorf("get pusher failed: %w", err)
|
||||
driver_asynq_worker.AppendLog(ctx, "推送失败: %v", errWrap)
|
||||
if driver_asynq_worker.IsFinalAttempt(ctx) {
|
||||
h.recordHistory(ctx, req, "failed", errWrap.Error())
|
||||
}
|
||||
return nil, errWrap
|
||||
logger.ErrorF(ctx, "[Push] 推送失败: %v", errWrap)
|
||||
h.recordHistory(ctx, req, "failed", errWrap.Error())
|
||||
return errWrap
|
||||
}
|
||||
|
||||
flatBody := req.Body.Flatten()
|
||||
@@ -95,29 +90,19 @@ func (h *PushHandler) Execute(ctx context.Context, payload []byte) (*driver_asyn
|
||||
content := req.Body.Content
|
||||
|
||||
if err != nil {
|
||||
driver_asynq_worker.AppendLog(ctx, "消息推送失败 (标题: %s): %v", title, err)
|
||||
if upstreamResp != "" {
|
||||
driver_asynq_worker.AppendLog(ctx, "上游返回: %s", upstreamResp)
|
||||
}
|
||||
if driver_asynq_worker.IsFinalAttempt(ctx) {
|
||||
h.recordHistory(ctx, req, "failed", err.Error())
|
||||
}
|
||||
return nil, fmt.Errorf("pusher.Send failed: %w", err)
|
||||
logger.ErrorF(ctx, "[Push] 消息推送失败 (标题: %s): %v, 上游返回: %s", title, err, upstreamResp)
|
||||
h.recordHistory(ctx, req, "failed", err.Error())
|
||||
return fmt.Errorf("pusher.Send failed: %w", err)
|
||||
}
|
||||
|
||||
driver_asynq_worker.AppendLog(ctx, "消息推送成功 (标题: %s, 内容摘要: %s)", title, content)
|
||||
if upstreamResp != "" {
|
||||
driver_asynq_worker.AppendLog(ctx, "上游返回: %s", upstreamResp)
|
||||
}
|
||||
logger.InfoF(ctx, "[Push] 消息推送成功 (标题: %s, 内容摘要: %s), 上游返回: %s", title, content, upstreamResp)
|
||||
h.recordHistory(ctx, req, "success", "")
|
||||
|
||||
return &driver_asynq_worker.TaskResult{
|
||||
Message: fmt.Sprintf("推送成功: [%s] -> %s", req.Config.Channel, req.Target),
|
||||
}, nil
|
||||
return nil
|
||||
}
|
||||
|
||||
func (h *PushHandler) recordHistory(ctx context.Context, req SendPayload, status string, errMsg string) {
|
||||
if dbErr := recordPushHistory(ctx, req, status, errMsg); dbErr != nil {
|
||||
driver_asynq_worker.AppendLog(ctx, "写入推送历史审计记录失败: %v", dbErr)
|
||||
logger.ErrorF(ctx, "[Push] 写入推送历史审计记录失败: %v", dbErr)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -11,8 +11,6 @@ import (
|
||||
"gorm.io/gorm"
|
||||
|
||||
"Wavelet/pkg/idgen"
|
||||
cachepkg "Wavelet/plugins/infra/cache"
|
||||
db "Wavelet/plugins/infra/database"
|
||||
)
|
||||
|
||||
const (
|
||||
@@ -25,18 +23,18 @@ func CreateMessageChannel(ctx context.Context, ch *MessageChannel) error {
|
||||
if ch.ID == 0 {
|
||||
ch.ID = idgen.NextUint64ID()
|
||||
}
|
||||
return db.DB(ctx).Create(ch).Error
|
||||
return getDB(ctx).Create(ch).Error
|
||||
}
|
||||
|
||||
// UpdateMessageChannel saves a channel row.
|
||||
func UpdateMessageChannel(ctx context.Context, ch *MessageChannel) error {
|
||||
return db.DB(ctx).Save(ch).Error
|
||||
return getDB(ctx).Save(ch).Error
|
||||
}
|
||||
|
||||
// GetMessageChannel loads a channel by id.
|
||||
func GetMessageChannel(ctx context.Context, id uint64) (*MessageChannel, error) {
|
||||
var ch MessageChannel
|
||||
if err := db.DB(ctx).Where("id = ?", id).First(&ch).Error; err != nil {
|
||||
if err := getDB(ctx).Where("id = ?", id).First(&ch).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &ch, nil
|
||||
@@ -45,7 +43,7 @@ func GetMessageChannel(ctx context.Context, id uint64) (*MessageChannel, error)
|
||||
// ListMessageChannels returns all channels newest first.
|
||||
func ListMessageChannels(ctx context.Context) ([]MessageChannel, error) {
|
||||
var rows []MessageChannel
|
||||
if err := db.DB(ctx).Order("id DESC").Find(&rows).Error; err != nil {
|
||||
if err := getDB(ctx).Order("id DESC").Find(&rows).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return rows, nil
|
||||
@@ -53,7 +51,7 @@ func ListMessageChannels(ctx context.Context) ([]MessageChannel, error) {
|
||||
|
||||
// DeleteMessageChannel removes pairings, bindings, then the channel.
|
||||
func DeleteMessageChannel(ctx context.Context, id uint64) error {
|
||||
return db.DB(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
return getDB(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
if err := tx.Where("channel_id = ?", id).Delete(&MessagePairingCode{}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
@@ -69,13 +67,13 @@ func CreateMessageBinding(ctx context.Context, b *MessageBinding) error {
|
||||
if b.ID == 0 {
|
||||
b.ID = idgen.NextUint64ID()
|
||||
}
|
||||
return db.DB(ctx).Create(b).Error
|
||||
return getDB(ctx).Create(b).Error
|
||||
}
|
||||
|
||||
// GetBindingByChannelPlatform finds a binding for a platform user on a channel.
|
||||
func GetBindingByChannelPlatform(ctx context.Context, channelID uint64, platformUserID string) (*MessageBinding, error) {
|
||||
var b MessageBinding
|
||||
err := db.DB(ctx).Where("channel_id = ? AND platform_user_id = ?", channelID, platformUserID).First(&b).Error
|
||||
err := getDB(ctx).Where("channel_id = ? AND platform_user_id = ?", channelID, platformUserID).First(&b).Error
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -85,7 +83,7 @@ func GetBindingByChannelPlatform(ctx context.Context, channelID uint64, platform
|
||||
// ListBindingsByUser lists bindings for a Wavelet user.
|
||||
func ListBindingsByUser(ctx context.Context, userID uint64) ([]MessageBinding, error) {
|
||||
var rows []MessageBinding
|
||||
if err := db.DB(ctx).Where("user_id = ?", userID).Order("id DESC").Find(&rows).Error; err != nil {
|
||||
if err := getDB(ctx).Where("user_id = ?", userID).Order("id DESC").Find(&rows).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return rows, nil
|
||||
@@ -94,7 +92,7 @@ func ListBindingsByUser(ctx context.Context, userID uint64) ([]MessageBinding, e
|
||||
// GetMessageBinding loads a binding by id.
|
||||
func GetMessageBinding(ctx context.Context, id uint64) (*MessageBinding, error) {
|
||||
var b MessageBinding
|
||||
if err := db.DB(ctx).Where("id = ?", id).First(&b).Error; err != nil {
|
||||
if err := getDB(ctx).Where("id = ?", id).First(&b).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &b, nil
|
||||
@@ -102,13 +100,13 @@ func GetMessageBinding(ctx context.Context, id uint64) (*MessageBinding, error)
|
||||
|
||||
// DeleteMessageBinding deletes a binding by id.
|
||||
func DeleteMessageBinding(ctx context.Context, id uint64) error {
|
||||
return db.DB(ctx).Delete(&MessageBinding{}, id).Error
|
||||
return getDB(ctx).Delete(&MessageBinding{}, id).Error
|
||||
}
|
||||
|
||||
// UpsertPairingCode reuses an unexpired code for the same channel+platform user.
|
||||
func UpsertPairingCode(ctx context.Context, channelID uint64, platformUserID, code string, expiresAt time.Time) (*MessagePairingCode, error) {
|
||||
var existing MessagePairingCode
|
||||
err := db.DB(ctx).
|
||||
err := getDB(ctx).
|
||||
Where("channel_id = ? AND platform_user_id = ? AND expires_at > ?", channelID, platformUserID, time.Now()).
|
||||
First(&existing).Error
|
||||
if err == nil {
|
||||
@@ -123,7 +121,7 @@ func UpsertPairingCode(ctx context.Context, channelID uint64, platformUserID, co
|
||||
PlatformUserID: platformUserID,
|
||||
ExpiresAt: expiresAt,
|
||||
}
|
||||
if err := db.DB(ctx).Create(row).Error; err != nil {
|
||||
if err := getDB(ctx).Create(row).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return row, nil
|
||||
@@ -132,7 +130,7 @@ func UpsertPairingCode(ctx context.Context, channelID uint64, platformUserID, co
|
||||
// GetPairingCode loads a pairing code by normalized code string.
|
||||
func GetPairingCode(ctx context.Context, code string) (*MessagePairingCode, error) {
|
||||
var row MessagePairingCode
|
||||
if err := db.DB(ctx).Where("code = ?", code).First(&row).Error; err != nil {
|
||||
if err := getDB(ctx).Where("code = ?", code).First(&row).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &row, nil
|
||||
@@ -140,18 +138,18 @@ func GetPairingCode(ctx context.Context, code string) (*MessagePairingCode, erro
|
||||
|
||||
// DeletePairingCode removes a pairing code.
|
||||
func DeletePairingCode(ctx context.Context, code string) error {
|
||||
return db.DB(ctx).Where("code = ?", code).Delete(&MessagePairingCode{}).Error
|
||||
return getDB(ctx).Where("code = ?", code).Delete(&MessagePairingCode{}).Error
|
||||
}
|
||||
|
||||
// DeleteExpiredPairingCodes removes expired pairing rows.
|
||||
func DeleteExpiredPairingCodes(ctx context.Context) error {
|
||||
return db.DB(ctx).Where("expires_at <= ?", time.Now()).Delete(&MessagePairingCode{}).Error
|
||||
return getDB(ctx).Where("expires_at <= ?", time.Now()).Delete(&MessagePairingCode{}).Error
|
||||
}
|
||||
|
||||
// ListEnabledMessageChannels returns enabled channels.
|
||||
func ListEnabledMessageChannels(ctx context.Context) ([]MessageChannel, error) {
|
||||
var rows []MessageChannel
|
||||
if err := db.DB(ctx).Where("enabled = ?", true).Find(&rows).Error; err != nil {
|
||||
if err := getDB(ctx).Where("enabled = ?", true).Find(&rows).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return rows, nil
|
||||
@@ -160,7 +158,7 @@ func ListEnabledMessageChannels(ctx context.Context) ([]MessageChannel, error) {
|
||||
// ListPushChannelsRecord returns all push channels ordered by creation time descending.
|
||||
func ListPushChannelsRecord(ctx context.Context) ([]PushChannel, error) {
|
||||
var channels []PushChannel
|
||||
if err := db.DB(ctx).Order("created_at DESC").Find(&channels).Error; err != nil {
|
||||
if err := getDB(ctx).Order("created_at DESC").Find(&channels).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return channels, nil
|
||||
@@ -169,7 +167,7 @@ func ListPushChannelsRecord(ctx context.Context) ([]PushChannel, error) {
|
||||
// GetPushChannelByIDRecord loads a push channel by primary key.
|
||||
func GetPushChannelByIDRecord(ctx context.Context, id uint64) (PushChannel, error) {
|
||||
var channel PushChannel
|
||||
if err := db.DB(ctx).Where("id = ?", id).First(&channel).Error; err != nil {
|
||||
if err := getDB(ctx).Where("id = ?", id).First(&channel).Error; err != nil {
|
||||
return PushChannel{}, err
|
||||
}
|
||||
return channel, nil
|
||||
@@ -178,7 +176,7 @@ func GetPushChannelByIDRecord(ctx context.Context, id uint64) (PushChannel, erro
|
||||
// GetPushChannelByNameRecord 根据名称获取消息通道。
|
||||
func GetPushChannelByNameRecord(ctx context.Context, name string) (*PushChannel, error) {
|
||||
var channel PushChannel
|
||||
if err := db.DB(ctx).Where("name = ?", name).First(&channel).Error; err != nil {
|
||||
if err := getDB(ctx).Where("name = ?", name).First(&channel).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &channel, nil
|
||||
@@ -187,7 +185,7 @@ func GetPushChannelByNameRecord(ctx context.Context, name string) (*PushChannel,
|
||||
// CountPushChannelsByNameRecord returns how many channels share the given name.
|
||||
func CountPushChannelsByNameRecord(ctx context.Context, name string) (int64, error) {
|
||||
var count int64
|
||||
if err := db.DB(ctx).Model(&PushChannel{}).Where("name = ?", name).Count(&count).Error; err != nil {
|
||||
if err := getDB(ctx).Model(&PushChannel{}).Where("name = ?", name).Count(&count).Error; err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return count, nil
|
||||
@@ -195,7 +193,7 @@ func CountPushChannelsByNameRecord(ctx context.Context, name string) (int64, err
|
||||
|
||||
// CreatePushChannelRecord persists a new channel and invalidates cache.
|
||||
func CreatePushChannelRecord(ctx context.Context, channel *PushChannel) error {
|
||||
if err := db.DB(ctx).Create(channel).Error; err != nil {
|
||||
if err := getDB(ctx).Create(channel).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
DeleteActivePushChannelCache(ctx, channel.Name)
|
||||
@@ -204,7 +202,7 @@ func CreatePushChannelRecord(ctx context.Context, channel *PushChannel) error {
|
||||
|
||||
// SavePushChannelRecord updates a channel and invalidates cache.
|
||||
func SavePushChannelRecord(ctx context.Context, channel *PushChannel) error {
|
||||
if err := db.DB(ctx).Save(channel).Error; err != nil {
|
||||
if err := getDB(ctx).Save(channel).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
DeleteActivePushChannelCache(ctx, channel.Name)
|
||||
@@ -213,7 +211,7 @@ func SavePushChannelRecord(ctx context.Context, channel *PushChannel) error {
|
||||
|
||||
// DeletePushChannelRecord removes a channel and invalidates cache.
|
||||
func DeletePushChannelRecord(ctx context.Context, channel *PushChannel) error {
|
||||
if err := db.DB(ctx).Delete(channel).Error; err != nil {
|
||||
if err := getDB(ctx).Delete(channel).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
DeleteActivePushChannelCache(ctx, channel.Name)
|
||||
@@ -224,18 +222,18 @@ func DeletePushChannelRecord(ctx context.Context, channel *PushChannel) error {
|
||||
func GetActivePushChannelByName(ctx context.Context, name string) (*PushChannel, error) {
|
||||
cacheKey := "push:channel:active:" + name
|
||||
var channel PushChannel
|
||||
if cachepkg.Redis != nil {
|
||||
if err := cachepkg.GetJSON(ctx, cacheKey, &channel); err == nil {
|
||||
if cache := getCache(ctx); cache != nil {
|
||||
if err := cache.Get(ctx, cacheKey, &channel); err == nil {
|
||||
return &channel, nil
|
||||
}
|
||||
}
|
||||
|
||||
if err := db.DB(ctx).Where("name = ? AND enabled = ?", name, true).First(&channel).Error; err != nil {
|
||||
if err := getDB(ctx).Where("name = ? AND enabled = ?", name, true).First(&channel).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if cachepkg.Redis != nil {
|
||||
_ = cachepkg.SetJSON(ctx, cacheKey, channel, activePushChannelCacheTTL)
|
||||
if cache := getCache(ctx); cache != nil {
|
||||
_ = cache.Set(ctx, cacheKey, channel, activePushChannelCacheTTL)
|
||||
}
|
||||
|
||||
return &channel, nil
|
||||
@@ -243,15 +241,15 @@ func GetActivePushChannelByName(ctx context.Context, name string) (*PushChannel,
|
||||
|
||||
// DeleteActivePushChannelCache 清理启用消息通道的缓存。
|
||||
func DeleteActivePushChannelCache(ctx context.Context, name string) {
|
||||
if cachepkg.Redis != nil {
|
||||
_ = cachepkg.Redis.Del(ctx, cachepkg.PrefixedKey("push:channel:active:"+name)).Err()
|
||||
if cache := getCache(ctx); cache != nil {
|
||||
_ = cache.Delete(ctx, "push:channel:active:"+name)
|
||||
}
|
||||
}
|
||||
|
||||
// ListPushEventsRecord returns all push events ordered by creation time descending.
|
||||
func ListPushEventsRecord(ctx context.Context) ([]PushEvent, error) {
|
||||
var events []PushEvent
|
||||
if err := db.DB(ctx).Order("created_at DESC").Find(&events).Error; err != nil {
|
||||
if err := getDB(ctx).Order("created_at DESC").Find(&events).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return events, nil
|
||||
@@ -260,7 +258,7 @@ func ListPushEventsRecord(ctx context.Context) ([]PushEvent, error) {
|
||||
// GetPushEventByIDRecord loads a push event by primary key.
|
||||
func GetPushEventByIDRecord(ctx context.Context, id uint64) (PushEvent, error) {
|
||||
var event PushEvent
|
||||
if err := db.DB(ctx).First(&event, id).Error; err != nil {
|
||||
if err := getDB(ctx).First(&event, id).Error; err != nil {
|
||||
return PushEvent{}, err
|
||||
}
|
||||
return event, nil
|
||||
@@ -269,7 +267,7 @@ func GetPushEventByIDRecord(ctx context.Context, id uint64) (PushEvent, error) {
|
||||
// GetPushEventByKeyRecord loads a push event by event key.
|
||||
func GetPushEventByKeyRecord(ctx context.Context, key string) (PushEvent, error) {
|
||||
var event PushEvent
|
||||
if err := db.DB(ctx).Where("event_key = ?", key).First(&event).Error; err != nil {
|
||||
if err := getDB(ctx).Where("event_key = ?", key).First(&event).Error; err != nil {
|
||||
return PushEvent{}, err
|
||||
}
|
||||
return event, nil
|
||||
@@ -278,7 +276,7 @@ func GetPushEventByKeyRecord(ctx context.Context, key string) (PushEvent, error)
|
||||
// CountPushEventsByKeyRecord returns how many events use the given event key.
|
||||
func CountPushEventsByKeyRecord(ctx context.Context, key string) (int64, error) {
|
||||
var count int64
|
||||
if err := db.DB(ctx).Model(&PushEvent{}).Where("event_key = ?", key).Count(&count).Error; err != nil {
|
||||
if err := getDB(ctx).Model(&PushEvent{}).Where("event_key = ?", key).Count(&count).Error; err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return count, nil
|
||||
@@ -286,7 +284,7 @@ func CountPushEventsByKeyRecord(ctx context.Context, key string) (int64, error)
|
||||
|
||||
// CreatePushEventRecord persists a new push event and invalidates cache.
|
||||
func CreatePushEventRecord(ctx context.Context, event *PushEvent) error {
|
||||
if err := db.DB(ctx).Create(event).Error; err != nil {
|
||||
if err := getDB(ctx).Create(event).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
DeleteActivePushEventCache(ctx, event.EventKey)
|
||||
@@ -295,7 +293,7 @@ func CreatePushEventRecord(ctx context.Context, event *PushEvent) error {
|
||||
|
||||
// SavePushEventRecord updates a push event and invalidates cache.
|
||||
func SavePushEventRecord(ctx context.Context, event *PushEvent) error {
|
||||
if err := db.DB(ctx).Save(event).Error; err != nil {
|
||||
if err := getDB(ctx).Save(event).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
DeleteActivePushEventCache(ctx, event.EventKey)
|
||||
@@ -305,7 +303,7 @@ func SavePushEventRecord(ctx context.Context, event *PushEvent) error {
|
||||
// UpdatePushEventEnabledRecord toggles the enabled flag for a push event.
|
||||
func UpdatePushEventEnabledRecord(ctx context.Context, event *PushEvent, enabled bool) error {
|
||||
event.Enabled = enabled
|
||||
if err := db.DB(ctx).Model(event).Update("enabled", enabled).Error; err != nil {
|
||||
if err := getDB(ctx).Model(event).Update("enabled", enabled).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
DeleteActivePushEventCache(ctx, event.EventKey)
|
||||
@@ -314,7 +312,7 @@ func UpdatePushEventEnabledRecord(ctx context.Context, event *PushEvent, enabled
|
||||
|
||||
// DeletePushEventRecord removes a push event and invalidates cache.
|
||||
func DeletePushEventRecord(ctx context.Context, event *PushEvent) error {
|
||||
if err := db.DB(ctx).Delete(event).Error; err != nil {
|
||||
if err := getDB(ctx).Delete(event).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
DeleteActivePushEventCache(ctx, event.EventKey)
|
||||
@@ -324,7 +322,7 @@ func DeletePushEventRecord(ctx context.Context, event *PushEvent) error {
|
||||
// ListActivePushEventsByTaskTypeRecord returns enabled events bound to a task type.
|
||||
func ListActivePushEventsByTaskTypeRecord(ctx context.Context, taskType string) ([]PushEvent, error) {
|
||||
var events []PushEvent
|
||||
if err := db.DB(ctx).Where("task_type = ? AND enabled = ?", taskType, true).Find(&events).Error; err != nil {
|
||||
if err := getDB(ctx).Where("task_type = ? AND enabled = ?", taskType, true).Find(&events).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return events, nil
|
||||
@@ -334,18 +332,18 @@ func ListActivePushEventsByTaskTypeRecord(ctx context.Context, taskType string)
|
||||
func GetActivePushEventByKey(ctx context.Context, key string) (*PushEvent, error) {
|
||||
cacheKey := "push:event:active:" + key
|
||||
var event PushEvent
|
||||
if cachepkg.Redis != nil {
|
||||
if err := cachepkg.GetJSON(ctx, cacheKey, &event); err == nil {
|
||||
if cache := getCache(ctx); cache != nil {
|
||||
if err := cache.Get(ctx, cacheKey, &event); err == nil {
|
||||
return &event, nil
|
||||
}
|
||||
}
|
||||
|
||||
if err := db.DB(ctx).Where("event_key = ? AND enabled = ?", key, true).First(&event).Error; err != nil {
|
||||
if err := getDB(ctx).Where("event_key = ? AND enabled = ?", key, true).First(&event).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if cachepkg.Redis != nil {
|
||||
_ = cachepkg.SetJSON(ctx, cacheKey, event, activePushEventCacheTTL)
|
||||
if cache := getCache(ctx); cache != nil {
|
||||
_ = cache.Set(ctx, cacheKey, event, activePushEventCacheTTL)
|
||||
}
|
||||
|
||||
return &event, nil
|
||||
@@ -353,14 +351,14 @@ func GetActivePushEventByKey(ctx context.Context, key string) (*PushEvent, error
|
||||
|
||||
// DeleteActivePushEventCache 清理启用通知事件的缓存。
|
||||
func DeleteActivePushEventCache(ctx context.Context, key string) {
|
||||
if cachepkg.Redis != nil {
|
||||
_ = cachepkg.Redis.Del(ctx, cachepkg.PrefixedKey("push:event:active:"+key)).Err()
|
||||
if cache := getCache(ctx); cache != nil {
|
||||
_ = cache.Delete(ctx, "push:event:active:"+key)
|
||||
}
|
||||
}
|
||||
|
||||
// ListPushHistoriesRecord returns paginated push history records.
|
||||
func ListPushHistoriesRecord(ctx context.Context, filter PushHistoryListFilter) (int64, []PushHistory, error) {
|
||||
query := db.DB(ctx).Model(&PushHistory{}).Order("created_at DESC")
|
||||
query := getDB(ctx).Model(&PushHistory{}).Order("created_at DESC")
|
||||
if filter.EventKey != "" {
|
||||
query = query.Where("event_key = ?", filter.EventKey)
|
||||
}
|
||||
@@ -384,10 +382,10 @@ func ListPushHistoriesRecord(ctx context.Context, filter PushHistoryListFilter)
|
||||
|
||||
// CreatePushHistoryRecord persists a push history audit record.
|
||||
func CreatePushHistoryRecord(ctx context.Context, history *PushHistory) error {
|
||||
return db.DB(ctx).Create(history).Error
|
||||
return getDB(ctx).Create(history).Error
|
||||
}
|
||||
|
||||
// PushHistoryQuery returns a scoped query builder for push histories.
|
||||
func PushHistoryQuery(ctx context.Context) *gorm.DB {
|
||||
return db.DB(ctx).Model(&PushHistory{})
|
||||
return getDB(ctx).Model(&PushHistory{})
|
||||
}
|
||||
|
||||
@@ -139,3 +139,55 @@ func Drain(ctx context.Context) error {
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// MigrateAndSwitchEngine migrates access logs to target database and switches the active store.
|
||||
func MigrateAndSwitchEngine(ctx context.Context, targetEngine string, reportProgress func(copied int)) error {
|
||||
if err := Drain(ctx); err != nil {
|
||||
return err
|
||||
}
|
||||
src, err := logstore.Active(ctx)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
dst, err := logstore.BuildForMigration(ctx, targetEngine)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if _, err := dst.UserAccessLogs.DeleteAll(ctx); err != nil {
|
||||
return err
|
||||
}
|
||||
from, to, err := src.UserAccessLogs.MigrationRange(ctx)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if !from.IsZero() && !to.IsZero() {
|
||||
if err := dst.UserAccessLogs.EnsurePartitions(ctx, from, to); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
var afterID uint64
|
||||
var copied int
|
||||
const copyBatchSize = 1000
|
||||
for {
|
||||
rows, err := src.UserAccessLogs.ListForMigration(ctx, afterID, copyBatchSize)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if len(rows) == 0 {
|
||||
break
|
||||
}
|
||||
if err := dst.UserAccessLogs.BatchInsert(ctx, rows); err != nil {
|
||||
return err
|
||||
}
|
||||
afterID = rows[len(rows)-1].ID
|
||||
copied += len(rows)
|
||||
if reportProgress != nil {
|
||||
reportProgress(copied)
|
||||
}
|
||||
if len(rows) < copyBatchSize {
|
||||
break
|
||||
}
|
||||
}
|
||||
logstore.InvalidateCache()
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -9,14 +9,14 @@ import (
|
||||
"fmt"
|
||||
"time"
|
||||
|
||||
"Wavelet/pkg/util"
|
||||
db "Wavelet/plugins/infra/database"
|
||||
"gorm.io/gorm"
|
||||
|
||||
"Wavelet/pkg/util"
|
||||
)
|
||||
|
||||
// CountAccessLogs returns the number of access logs matching filter.
|
||||
func CountAccessLogs(ctx context.Context, filter AccessLogFilter) (uint64, error) {
|
||||
ch := db.ChDB(ctx)
|
||||
ch := getChDB(ctx)
|
||||
if ch == nil {
|
||||
return 0, fmt.Errorf("clickhouse gorm connection is not initialized")
|
||||
}
|
||||
@@ -31,7 +31,7 @@ func CountAccessLogs(ctx context.Context, filter AccessLogFilter) (uint64, error
|
||||
|
||||
// ListAccessLogs returns paginated access logs and the total match count.
|
||||
func ListAccessLogs(ctx context.Context, filter AccessLogFilter, page, pageSize int) ([]UserAccessLog, uint64, error) {
|
||||
ch := db.ChDB(ctx)
|
||||
ch := getChDB(ctx)
|
||||
if ch == nil {
|
||||
return nil, 0, fmt.Errorf("clickhouse gorm connection is not initialized")
|
||||
}
|
||||
@@ -40,42 +40,28 @@ func ListAccessLogs(ctx context.Context, filter AccessLogFilter, page, pageSize
|
||||
return []UserAccessLog{}, 0, nil
|
||||
}
|
||||
|
||||
var total int64
|
||||
baseQuery := applyFilter(ch.Model(&UserAccessLog{}), filter)
|
||||
if err := baseQuery.Count(&total).Error; err != nil {
|
||||
var count int64
|
||||
query := applyFilter(ch.Model(&UserAccessLog{}), filter)
|
||||
if err := query.Count(&count).Error; err != nil {
|
||||
return nil, 0, fmt.Errorf("count access logs: %w", err)
|
||||
}
|
||||
if total == 0 {
|
||||
return []UserAccessLog{}, 0, nil
|
||||
}
|
||||
|
||||
if page < 1 {
|
||||
page = 1
|
||||
}
|
||||
if pageSize < 1 {
|
||||
pageSize = 20
|
||||
}
|
||||
offset := (page - 1) * pageSize
|
||||
|
||||
var logs []UserAccessLog
|
||||
err := applyFilter(ch.Model(&UserAccessLog{}), filter).
|
||||
Order("created_at DESC, id DESC").
|
||||
Limit(pageSize).
|
||||
Offset(offset).
|
||||
Find(&logs).Error
|
||||
if err != nil {
|
||||
offset := (page - 1) * pageSize
|
||||
if err := query.Order("created_at DESC").Limit(pageSize).Offset(offset).Find(&logs).Error; err != nil {
|
||||
return nil, 0, fmt.Errorf("list access logs: %w", err)
|
||||
}
|
||||
|
||||
return logs, safeUint64Count(total), nil
|
||||
return logs, safeUint64Count(count), nil
|
||||
}
|
||||
|
||||
// DeleteAllUserAccessLogs hard-deletes all user access logs via TRUNCATE.
|
||||
func DeleteAllUserAccessLogs(ctx context.Context) (int64, error) {
|
||||
if db.ChConn == nil {
|
||||
conn := getChConn()
|
||||
if conn == nil {
|
||||
return 0, fmt.Errorf("clickhouse connection is not initialized")
|
||||
}
|
||||
if err := db.ChConn.Exec(ctx, "TRUNCATE TABLE "+UserAccessLog{}.TableName()); err != nil {
|
||||
if err := conn.Exec(ctx, "TRUNCATE TABLE "+UserAccessLog{}.TableName()); err != nil {
|
||||
return 0, fmt.Errorf("truncate user access logs: %w", err)
|
||||
}
|
||||
return 0, nil
|
||||
@@ -83,10 +69,11 @@ func DeleteAllUserAccessLogs(ctx context.Context) (int64, error) {
|
||||
|
||||
// DeleteUserAccessLogsBefore deletes user access logs older than cutoff.
|
||||
func DeleteUserAccessLogsBefore(ctx context.Context, cutoff time.Time) (int64, error) {
|
||||
if db.ChConn == nil {
|
||||
conn := getChConn()
|
||||
if conn == nil {
|
||||
return 0, fmt.Errorf("clickhouse connection is not initialized")
|
||||
}
|
||||
if err := db.ChConn.Exec(ctx, "ALTER TABLE "+UserAccessLog{}.TableName()+" DELETE WHERE created_at < ?", cutoff); err != nil {
|
||||
if err := conn.Exec(ctx, "ALTER TABLE "+UserAccessLog{}.TableName()+" DELETE WHERE created_at < ?", cutoff); err != nil {
|
||||
return 0, fmt.Errorf("delete expired user access logs: %w", err)
|
||||
}
|
||||
return 0, nil
|
||||
|
||||
@@ -8,8 +8,6 @@ import (
|
||||
"fmt"
|
||||
"sort"
|
||||
"time"
|
||||
|
||||
db "Wavelet/plugins/infra/database"
|
||||
)
|
||||
|
||||
const hoursInDay = 24
|
||||
@@ -20,7 +18,7 @@ func GetDailyTrend(ctx context.Context, days int) ([]DailyTrend, error) {
|
||||
days = 7
|
||||
}
|
||||
|
||||
ch := db.ChDB(ctx)
|
||||
ch := getChDB(ctx)
|
||||
if ch == nil {
|
||||
return nil, fmt.Errorf("clickhouse gorm connection is not initialized")
|
||||
}
|
||||
@@ -69,7 +67,7 @@ func GetDailyTrend(ctx context.Context, days int) ([]DailyTrend, error) {
|
||||
|
||||
// GetBrowserDistribution returns browser-grouped access counts since startTime.
|
||||
func GetBrowserDistribution(ctx context.Context, startTime time.Time) ([]BrowserShare, error) {
|
||||
ch := db.ChDB(ctx)
|
||||
ch := getChDB(ctx)
|
||||
if ch == nil {
|
||||
return nil, fmt.Errorf("clickhouse gorm connection is not initialized")
|
||||
}
|
||||
@@ -117,7 +115,7 @@ func GetTopActiveUsers(ctx context.Context, startTime time.Time, limit int) ([]T
|
||||
limit = 10
|
||||
}
|
||||
|
||||
ch := db.ChDB(ctx)
|
||||
ch := getChDB(ctx)
|
||||
if ch == nil {
|
||||
return nil, fmt.Errorf("clickhouse gorm connection is not initialized")
|
||||
}
|
||||
|
||||
@@ -16,8 +16,6 @@ import (
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"gorm.io/gorm"
|
||||
|
||||
db "Wavelet/plugins/infra/database"
|
||||
)
|
||||
|
||||
func setupChGormDB(t *testing.T) *gorm.DB {
|
||||
@@ -28,7 +26,7 @@ func setupChGormDB(t *testing.T) *gorm.DB {
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, gormDB.AutoMigrate(&UserAccessLog{}))
|
||||
db.SetChDBForTest(gormDB)
|
||||
SetChDBForTest(gormDB)
|
||||
return gormDB
|
||||
}
|
||||
|
||||
@@ -56,7 +54,7 @@ func TestParseBrowserName(t *testing.T) {
|
||||
|
||||
func TestCountAccessLogs_EmptyUserIDs(t *testing.T) {
|
||||
setupChGormDB(t)
|
||||
t.Cleanup(func() { db.SetChDBForTest(nil) })
|
||||
t.Cleanup(func() { SetChDBForTest(nil) })
|
||||
|
||||
count, err := CountAccessLogs(context.Background(), AccessLogFilter{UserIDs: []uint64{}})
|
||||
require.NoError(t, err)
|
||||
@@ -65,7 +63,7 @@ func TestCountAccessLogs_EmptyUserIDs(t *testing.T) {
|
||||
|
||||
func TestListAccessLogs_EmptyUserIDs(t *testing.T) {
|
||||
setupChGormDB(t)
|
||||
t.Cleanup(func() { db.SetChDBForTest(nil) })
|
||||
t.Cleanup(func() { SetChDBForTest(nil) })
|
||||
|
||||
logs, total, err := ListAccessLogs(context.Background(), AccessLogFilter{UserIDs: []uint64{}}, 1, 20)
|
||||
require.NoError(t, err)
|
||||
@@ -75,7 +73,7 @@ func TestListAccessLogs_EmptyUserIDs(t *testing.T) {
|
||||
|
||||
func TestListAccessLogs_WithFilters(t *testing.T) {
|
||||
gormDB := setupChGormDB(t)
|
||||
t.Cleanup(func() { db.SetChDBForTest(nil) })
|
||||
t.Cleanup(func() { SetChDBForTest(nil) })
|
||||
|
||||
now := time.Now().UTC().Truncate(time.Second)
|
||||
logs := []UserAccessLog{
|
||||
@@ -116,8 +114,8 @@ func TestBatchInsert_UsesModelBatchSQL(t *testing.T) {
|
||||
batch: mockBatch,
|
||||
batchQuery: UserAccessLog{}.BatchInsertSQL(),
|
||||
}
|
||||
db.SetChConnForTest(mockConn)
|
||||
t.Cleanup(func() { db.SetChConnForTest(nil) })
|
||||
SetChConnForTest(mockConn)
|
||||
t.Cleanup(func() { SetChConnForTest(nil) })
|
||||
|
||||
createdAt := time.Now().UTC()
|
||||
err := BatchInsert(ctx, []UserAccessLog{
|
||||
|
||||
@@ -6,8 +6,6 @@ package logstore
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
|
||||
db "Wavelet/plugins/infra/database"
|
||||
)
|
||||
|
||||
// BatchInsert writes access logs to ClickHouse using the native batch API.
|
||||
@@ -15,11 +13,12 @@ func BatchInsert(ctx context.Context, logs []UserAccessLog) error {
|
||||
if len(logs) == 0 {
|
||||
return nil
|
||||
}
|
||||
if db.ChConn == nil {
|
||||
conn := getChConn()
|
||||
if conn == nil {
|
||||
return fmt.Errorf("clickhouse connection is not initialized")
|
||||
}
|
||||
|
||||
batch, err := db.ChConn.PrepareBatch(ctx, UserAccessLog{}.BatchInsertSQL())
|
||||
batch, err := conn.PrepareBatch(ctx, UserAccessLog{}.BatchInsertSQL())
|
||||
if err != nil {
|
||||
return fmt.Errorf("prepare clickhouse batch: %w", err)
|
||||
}
|
||||
|
||||
@@ -8,7 +8,6 @@ import (
|
||||
"fmt"
|
||||
"time"
|
||||
|
||||
db "Wavelet/plugins/infra/database"
|
||||
"github.com/ClickHouse/clickhouse-go/v2/lib/driver"
|
||||
)
|
||||
|
||||
@@ -93,12 +92,13 @@ func (s *clickhouseUserAccessLogStore) DropExpiredPartitions(_ context.Context,
|
||||
}
|
||||
|
||||
func (s *clickhouseUserAccessLogStore) MigrationRange(ctx context.Context) (time.Time, time.Time, error) {
|
||||
if db.ChConn == nil {
|
||||
conn := getChConn()
|
||||
if conn == nil {
|
||||
return time.Time{}, time.Time{}, fmt.Errorf("clickhouse connection is not initialized")
|
||||
}
|
||||
table := UserAccessLog{}.TableName()
|
||||
var minTime, maxTime *time.Time
|
||||
if err := db.ChConn.QueryRow(ctx, "SELECT min(created_at), max(created_at) FROM "+table).Scan(&minTime, &maxTime); err != nil {
|
||||
if err := conn.QueryRow(ctx, "SELECT min(created_at), max(created_at) FROM "+table).Scan(&minTime, &maxTime); err != nil {
|
||||
return time.Time{}, time.Time{}, fmt.Errorf("query migration range %s: %w", table, err)
|
||||
}
|
||||
if minTime == nil || maxTime == nil {
|
||||
@@ -108,7 +108,8 @@ func (s *clickhouseUserAccessLogStore) MigrationRange(ctx context.Context) (time
|
||||
}
|
||||
|
||||
func (s *clickhouseUserAccessLogStore) ListForMigration(ctx context.Context, afterID uint64, limit int) ([]UserAccessLog, error) {
|
||||
if db.ChConn == nil {
|
||||
conn := getChConn()
|
||||
if conn == nil {
|
||||
return nil, fmt.Errorf("clickhouse connection is not initialized")
|
||||
}
|
||||
if limit <= 0 {
|
||||
@@ -116,7 +117,7 @@ func (s *clickhouseUserAccessLogStore) ListForMigration(ctx context.Context, aft
|
||||
}
|
||||
table := UserAccessLog{}.TableName()
|
||||
columns := UserAccessLog{}.InsertColumns()
|
||||
rows, err := db.ChConn.Query(ctx, fmt.Sprintf(
|
||||
rows, err := conn.Query(ctx, fmt.Sprintf(
|
||||
"SELECT %s FROM %s WHERE id > ? ORDER BY id ASC LIMIT ?",
|
||||
columns, table,
|
||||
), afterID, limit)
|
||||
|
||||
@@ -0,0 +1,80 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package logstore
|
||||
|
||||
import (
|
||||
"context"
|
||||
"sync"
|
||||
|
||||
"github.com/ClickHouse/clickhouse-go/v2/lib/driver"
|
||||
"gorm.io/gorm"
|
||||
|
||||
"Wavelet/core"
|
||||
"Wavelet/core/contracts"
|
||||
)
|
||||
|
||||
var (
|
||||
dbMu sync.RWMutex
|
||||
dbSvc contracts.DBService
|
||||
chConn driver.Conn
|
||||
chDB *gorm.DB
|
||||
)
|
||||
|
||||
// SetDBService configures the DBService instance for logstore.
|
||||
func SetDBService(s contracts.DBService) {
|
||||
dbMu.Lock()
|
||||
defer dbMu.Unlock()
|
||||
dbSvc = s
|
||||
}
|
||||
|
||||
// SetChConnForTest configures ClickHouse native connection for test or runtime.
|
||||
func SetChConnForTest(conn driver.Conn) {
|
||||
dbMu.Lock()
|
||||
defer dbMu.Unlock()
|
||||
chConn = conn
|
||||
}
|
||||
|
||||
// SetChDBForTest configures ClickHouse GORM DB for test or runtime.
|
||||
func SetChDBForTest(db *gorm.DB) {
|
||||
dbMu.Lock()
|
||||
defer dbMu.Unlock()
|
||||
chDB = db
|
||||
}
|
||||
|
||||
func getDB(ctx context.Context) *gorm.DB {
|
||||
if c, ok := ctx.(*core.Context); ok && c != nil {
|
||||
if s, err := core.Inject[contracts.DBService](c); err == nil && s != nil {
|
||||
return s.DB(ctx)
|
||||
}
|
||||
}
|
||||
dbMu.RLock()
|
||||
s := dbSvc
|
||||
dbMu.RUnlock()
|
||||
if s != nil {
|
||||
return s.DB(ctx)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func getChDB(ctx context.Context) *gorm.DB {
|
||||
dbMu.RLock()
|
||||
customCh := chDB
|
||||
s := dbSvc
|
||||
dbMu.RUnlock()
|
||||
if customCh != nil {
|
||||
return customCh.WithContext(ctx)
|
||||
}
|
||||
if s != nil {
|
||||
if ch := s.Named("clickhouse"); ch != nil {
|
||||
return ch.WithContext(ctx)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func getChConn() driver.Conn {
|
||||
dbMu.RLock()
|
||||
defer dbMu.RUnlock()
|
||||
return chConn
|
||||
}
|
||||
@@ -11,8 +11,9 @@ import (
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"Wavelet/pkg/idgen"
|
||||
"gorm.io/gorm"
|
||||
|
||||
"Wavelet/pkg/idgen"
|
||||
)
|
||||
|
||||
const (
|
||||
|
||||
@@ -12,7 +12,6 @@ import (
|
||||
|
||||
"Wavelet/pkg/config"
|
||||
"Wavelet/pkg/logger"
|
||||
db "Wavelet/plugins/infra/database"
|
||||
)
|
||||
|
||||
const (
|
||||
@@ -98,7 +97,7 @@ func buildStore(ctx context.Context, database string, skipFreeze bool) (*Store,
|
||||
ual.skipFreeze = skipFreeze
|
||||
return &Store{UserAccessLogs: ual, Status: ual}, nil
|
||||
case dbNamePostgres, dbNameSQLite:
|
||||
gdb := db.DB(ctx)
|
||||
gdb := getDB(ctx)
|
||||
ual := newUserAccessLogGormStore(gdb)
|
||||
ual.skipFreeze = skipFreeze
|
||||
return &Store{UserAccessLogs: ual, Status: ual}, nil
|
||||
|
||||
@@ -9,13 +9,14 @@ import (
|
||||
"net/http"
|
||||
"time"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/pkg/config"
|
||||
"Wavelet/pkg/idgen"
|
||||
"Wavelet/pkg/response"
|
||||
"Wavelet/pkg/util"
|
||||
"Wavelet/plugins/domain/risk_control/logstore"
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
// Middleware is an alias for RiskControlMiddleware.
|
||||
|
||||
@@ -12,6 +12,9 @@ import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/stretchr/testify/assert"
|
||||
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/pkg/batchwriter"
|
||||
"Wavelet/pkg/config"
|
||||
@@ -19,8 +22,6 @@ import (
|
||||
"Wavelet/pkg/util"
|
||||
"Wavelet/plugins/domain/risk_control"
|
||||
"Wavelet/plugins/domain/risk_control/logstore"
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/stretchr/testify/assert"
|
||||
)
|
||||
|
||||
func newTestAccessLogWriter(t *testing.T, cfg batchwriter.Config) (*batchwriter.Writer[*logstore.UserAccessLog], func() []*logstore.UserAccessLog) {
|
||||
|
||||
@@ -9,10 +9,12 @@ import (
|
||||
"embed"
|
||||
"reflect"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"Wavelet/core"
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/core/extpoints"
|
||||
"github.com/gin-gonic/gin"
|
||||
"Wavelet/plugins/domain/risk_control/logstore"
|
||||
)
|
||||
|
||||
//go:embed logstore/migrations/*.sql
|
||||
@@ -69,6 +71,19 @@ func (p *Plugin) Manifest() core.Manifest {
|
||||
|
||||
// Apply registers risk control middlewares, settings, and cleanup hooks into the Context.
|
||||
func (p *Plugin) Apply(ctx *core.Context) error {
|
||||
// 0. Bind DBService
|
||||
if db, err := core.Inject[contracts.DBService](ctx); err == nil && db != nil {
|
||||
logstore.SetDBService(db)
|
||||
} else {
|
||||
core.When[contracts.DBService](ctx, func(db contracts.DBService) {
|
||||
logstore.SetDBService(db)
|
||||
})
|
||||
}
|
||||
ctx.OnDispose(func() error {
|
||||
logstore.SetDBService(nil)
|
||||
return nil
|
||||
})
|
||||
|
||||
// 0. Register user access log table migrations
|
||||
ctx.Migrations().Register("risk_control/logstore", riskControlMigrations)
|
||||
|
||||
@@ -98,10 +113,90 @@ func (p *Plugin) Apply(ctx *core.Context) error {
|
||||
Category: "security",
|
||||
})
|
||||
|
||||
// 4. Register lifecycle disposal cleanup
|
||||
// 4. Register RiskControlService contract
|
||||
core.Provide[contracts.RiskControlService](ctx, &riskControlServiceImpl{})
|
||||
|
||||
// 5. Register lifecycle disposal cleanup
|
||||
ctx.OnDispose(func() error {
|
||||
return StopLogWriter(context.Background())
|
||||
})
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
type riskControlServiceImpl struct{}
|
||||
|
||||
func (s *riskControlServiceImpl) QueryAccessLogs(ctx context.Context, filter contracts.AccessLogFilterDTO, page, pageSize int) ([]contracts.AccessLogDTO, uint64, error) {
|
||||
store, err := logstore.Active(ctx)
|
||||
if err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
f := logstore.AccessLogFilter{
|
||||
UserIDs: filter.UserIDs,
|
||||
Path: filter.Path,
|
||||
StartTime: filter.StartTime,
|
||||
EndTime: filter.EndTime,
|
||||
}
|
||||
list, total, err := store.UserAccessLogs.List(ctx, f, page, pageSize)
|
||||
if err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
items := make([]contracts.AccessLogDTO, len(list))
|
||||
for i, item := range list {
|
||||
items[i] = contracts.AccessLogDTO{
|
||||
ID: item.ID,
|
||||
UserID: item.UserID,
|
||||
IP: item.IP,
|
||||
UserAgent: item.UserAgent,
|
||||
Method: item.Method,
|
||||
Path: item.Path,
|
||||
Status: item.Status,
|
||||
Latency: item.Latency,
|
||||
CreatedAt: item.CreatedAt,
|
||||
}
|
||||
}
|
||||
return items, total, nil
|
||||
}
|
||||
|
||||
func (s *riskControlServiceImpl) QueryAccessLogStats(ctx context.Context, days int) ([]contracts.AccessLogDailyStatsDTO, error) {
|
||||
store, err := logstore.Active(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
trend, err := store.UserAccessLogs.GetDailyTrend(ctx, days)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
res := make([]contracts.AccessLogDailyStatsDTO, len(trend))
|
||||
for i, t := range trend {
|
||||
res[i] = contracts.AccessLogDailyStatsDTO{
|
||||
Date: t.Date,
|
||||
PV: t.Count,
|
||||
}
|
||||
}
|
||||
return res, nil
|
||||
}
|
||||
|
||||
func (s *riskControlServiceImpl) ActiveLogEngine(ctx context.Context) string {
|
||||
store, err := logstore.Active(ctx)
|
||||
if err != nil {
|
||||
return "sqlite"
|
||||
}
|
||||
active, err := store.Status.ActiveDatabase(ctx)
|
||||
if err != nil {
|
||||
return "sqlite"
|
||||
}
|
||||
return active
|
||||
}
|
||||
|
||||
func (s *riskControlServiceImpl) IsLogEngineMigrating(ctx context.Context) bool {
|
||||
return logstore.Migrating(ctx)
|
||||
}
|
||||
|
||||
func (s *riskControlServiceImpl) Drain(ctx context.Context) error {
|
||||
return Drain(ctx)
|
||||
}
|
||||
|
||||
func (s *riskControlServiceImpl) SwitchLogEngine(ctx context.Context, targetEngine string) error {
|
||||
return MigrateAndSwitchEngine(ctx, targetEngine, nil)
|
||||
}
|
||||
|
||||
@@ -8,11 +8,12 @@ import (
|
||||
"net/http"
|
||||
"reflect"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"Wavelet/core"
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/pkg/config"
|
||||
"Wavelet/pkg/response"
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
// Plugin implements core.Plugin to provide system-level basic routes.
|
||||
@@ -62,7 +63,7 @@ func (p *Plugin) Apply(ctx *core.Context) error {
|
||||
Value string `json:"value"`
|
||||
}
|
||||
var configs []configItem
|
||||
if dbSvc := ctx.DB(); dbSvc != nil {
|
||||
if dbSvc, err := core.Inject[contracts.DBService](ctx); err == nil && dbSvc != nil {
|
||||
_ = dbSvc.DB(c.Request.Context()).Table("w_system_configs").Where("visibility = ?", "visible").Find(&configs).Error
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(gin.H{
|
||||
|
||||
+8
-38
@@ -11,19 +11,13 @@ import (
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"Wavelet/pkg/util"
|
||||
"Wavelet/plugins/domain/upload/shared"
|
||||
uploadstorage "Wavelet/plugins/domain/upload/storage"
|
||||
cachepkg "Wavelet/plugins/infra/cache"
|
||||
database "Wavelet/plugins/infra/database"
|
||||
"Wavelet/plugins/infra/storage/objectstore"
|
||||
)
|
||||
|
||||
const fileAccessInvalidationChannel = "upload:file_access_invalidation"
|
||||
|
||||
var (
|
||||
accessCacheOnce sync.Once
|
||||
|
||||
fileAccessWhitelistMu sync.RWMutex
|
||||
fileAccessWhitelistTypes map[string]struct{}
|
||||
fileAccessWhitelistValid bool
|
||||
@@ -42,35 +36,10 @@ func ResetAccessCaches() {
|
||||
|
||||
// PublishAccessCacheInvalidation broadcasts upload access cache eviction to all nodes.
|
||||
func PublishAccessCacheInvalidation(ctx context.Context) {
|
||||
if cachepkg.Redis != nil {
|
||||
_ = cachepkg.Redis.Publish(ctx, fileAccessInvalidationChannel, "reset").Err()
|
||||
if cache := shared.GetCache(ctx); cache != nil {
|
||||
_ = cache.Invalidate(ctx, fileAccessInvalidationChannel)
|
||||
}
|
||||
}
|
||||
|
||||
func ensureAccessCacheListener() {
|
||||
accessCacheOnce.Do(startAccessCacheInvalidationListener)
|
||||
}
|
||||
|
||||
func startAccessCacheInvalidationListener() {
|
||||
rdb := cachepkg.Redis
|
||||
if rdb == nil {
|
||||
return
|
||||
}
|
||||
|
||||
util.Go(func() {
|
||||
pubsub := rdb.Subscribe(
|
||||
context.Background(),
|
||||
objectstore.ConfigInvalidationChannel,
|
||||
fileAccessInvalidationChannel,
|
||||
)
|
||||
defer func() {
|
||||
_ = pubsub.Close()
|
||||
}()
|
||||
|
||||
for range pubsub.Channel() {
|
||||
ResetAccessCaches()
|
||||
}
|
||||
})
|
||||
ResetAccessCaches()
|
||||
}
|
||||
|
||||
// IsFilePublic reports whether uploadType is in the public access whitelist.
|
||||
@@ -81,8 +50,6 @@ func IsFilePublic(ctx context.Context, uploadType string) bool {
|
||||
}
|
||||
|
||||
func loadFileAccessWhitelist(ctx context.Context) map[string]struct{} {
|
||||
ensureAccessCacheListener()
|
||||
|
||||
fileAccessWhitelistMu.RLock()
|
||||
if fileAccessWhitelistValid && time.Since(fileAccessWhitelistCheckedAt) < time.Duration(shared.AccessCacheTTL)*time.Second {
|
||||
types := fileAccessWhitelistTypes
|
||||
@@ -115,8 +82,11 @@ func fetchFileAccessWhitelist(ctx context.Context) map[string]struct{} {
|
||||
|
||||
func parseFileAccessWhitelist(ctx context.Context) []string {
|
||||
var sc struct{ Value string }
|
||||
err := database.DB(ctx).Table("w_system_configs").Where("key = ?", "file_access_whitelist").First(&sc).Error
|
||||
if err != nil || sc.Value == "" {
|
||||
db := shared.GetDB(ctx)
|
||||
if db != nil {
|
||||
_ = db.Table("w_system_configs").Where("key = ?", "file_access_whitelist").First(&sc).Error
|
||||
}
|
||||
if sc.Value == "" {
|
||||
return []string{shared.DefaultPublicUploadType}
|
||||
}
|
||||
|
||||
|
||||
@@ -8,13 +8,12 @@ import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"Wavelet/pkg/testhelper"
|
||||
"Wavelet/plugins/domain/upload/shared"
|
||||
uploadstorage "Wavelet/plugins/domain/upload/storage"
|
||||
)
|
||||
|
||||
func TestLoadMigrationAccessStateCachesResult(t *testing.T) {
|
||||
_, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
_, cleanup := shared.SetupTestEnv(t)
|
||||
defer cleanup()
|
||||
ResetAccessCaches()
|
||||
|
||||
@@ -31,7 +30,7 @@ func TestLoadMigrationAccessStateCachesResult(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestIsFilePublicUsesCachedWhitelist(t *testing.T) {
|
||||
_, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
_, cleanup := shared.SetupTestEnv(t)
|
||||
defer cleanup()
|
||||
ResetAccessCaches()
|
||||
|
||||
@@ -48,7 +47,7 @@ func TestIsFilePublicUsesCachedWhitelist(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestResetAccessCachesRefreshesWhitelist(t *testing.T) {
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
dbConn, cleanup := shared.SetupTestEnv(t)
|
||||
defer cleanup()
|
||||
ResetAccessCaches()
|
||||
|
||||
@@ -71,7 +70,7 @@ func TestResetAccessCachesRefreshesWhitelist(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestAccessCacheTTLExpires(t *testing.T) {
|
||||
_, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
_, cleanup := shared.SetupTestEnv(t)
|
||||
defer cleanup()
|
||||
ResetAccessCaches()
|
||||
|
||||
|
||||
+70
-97
@@ -5,15 +5,14 @@ package cache
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"gorm.io/gorm"
|
||||
|
||||
"Wavelet/pkg/cache/ram"
|
||||
"Wavelet/pkg/util"
|
||||
"Wavelet/plugins/domain/upload/models"
|
||||
cachepkg "Wavelet/plugins/infra/cache"
|
||||
database "Wavelet/plugins/infra/database"
|
||||
"Wavelet/plugins/domain/upload/shared"
|
||||
)
|
||||
|
||||
const (
|
||||
@@ -22,16 +21,8 @@ const (
|
||||
uploadMetaInvalidationChan = "upload:meta_invalidation"
|
||||
)
|
||||
|
||||
type uploadMetaInvalidationMessage struct {
|
||||
ID uint64 `json:"id"`
|
||||
}
|
||||
|
||||
var (
|
||||
uploadMetaRAM = ram.MustNew[uint64, models.Upload](ram.Options{MaximumSize: uploadMetaRAMMaximumSize})
|
||||
uploadMetaListenerOnce sync.Once
|
||||
uploadMetaListenerCtx context.Context
|
||||
uploadMetaListenerCancel context.CancelFunc
|
||||
uploadMetaListenerDone chan struct{}
|
||||
uploadMetaRAM = ram.MustNew[uint64, models.Upload](ram.Options{MaximumSize: uploadMetaRAMMaximumSize})
|
||||
)
|
||||
|
||||
func uploadMetaRedisKey(id uint64) string {
|
||||
@@ -42,120 +33,102 @@ func cloneUpload(u models.Upload) models.Upload {
|
||||
return u
|
||||
}
|
||||
|
||||
func ensureUploadMetaCacheListener() {
|
||||
if cachepkg.Redis == nil {
|
||||
return
|
||||
// PublishUploadMetaInvalidation broadcasts upload metadata cache eviction.
|
||||
func PublishUploadMetaInvalidation(ctx context.Context, id uint64) {
|
||||
if cache := shared.GetCache(ctx); cache != nil {
|
||||
_ = cache.Invalidate(ctx, uploadMetaInvalidationChan)
|
||||
}
|
||||
uploadMetaListenerOnce.Do(startUploadMetaCacheInvalidationListener)
|
||||
}
|
||||
|
||||
func startUploadMetaCacheInvalidationListener() {
|
||||
uploadMetaListenerCtx, uploadMetaListenerCancel = context.WithCancel(context.Background())
|
||||
uploadMetaListenerDone = make(chan struct{})
|
||||
|
||||
redisClient := cachepkg.Redis // 捕获当前客户端:goroutine 不读可变全局,避免与测试置空 cachepkg.Redis 竞争
|
||||
util.Go(func() {
|
||||
defer close(uploadMetaListenerDone)
|
||||
pubsub := redisClient.Subscribe(uploadMetaListenerCtx, uploadMetaInvalidationChan)
|
||||
defer func() {
|
||||
_ = pubsub.Close()
|
||||
}()
|
||||
|
||||
util.Go(func() {
|
||||
<-uploadMetaListenerCtx.Done()
|
||||
_ = pubsub.Close()
|
||||
})
|
||||
|
||||
for msg := range pubsub.Channel() {
|
||||
var payload uploadMetaInvalidationMessage
|
||||
if err := json.Unmarshal([]byte(msg.Payload), &payload); err != nil || payload.ID == 0 {
|
||||
uploadMetaRAM.InvalidateAll()
|
||||
continue
|
||||
}
|
||||
uploadMetaRAM.Invalidate(payload.ID)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func publishUploadMetaRAMInvalidation(ctx context.Context, id uint64) {
|
||||
if cachepkg.Redis == nil {
|
||||
return
|
||||
}
|
||||
payload, err := json.Marshal(uploadMetaInvalidationMessage{ID: id})
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
_ = cachepkg.Redis.Publish(ctx, uploadMetaInvalidationChan, payload).Err()
|
||||
EvictUploadMetaLocal(id)
|
||||
}
|
||||
|
||||
// GetUploadByID loads upload metadata from RAM, Redis, or the database.
|
||||
func GetUploadByID(ctx context.Context, id uint64) (models.Upload, error) {
|
||||
ensureUploadMetaCacheListener()
|
||||
if id == 0 {
|
||||
return models.Upload{}, gorm.ErrRecordNotFound
|
||||
}
|
||||
|
||||
// 1. RAM L1 Cache
|
||||
if u, ok := uploadMetaRAM.GetIfPresent(id); ok {
|
||||
return cloneUpload(u), nil
|
||||
}
|
||||
|
||||
key := uploadMetaRedisKey(id)
|
||||
if cachepkg.Redis != nil {
|
||||
|
||||
// 2. Redis L2 Cache
|
||||
if cache := shared.GetCache(ctx); cache != nil {
|
||||
var u models.Upload
|
||||
if err := cachepkg.GetJSON(ctx, key, &u); err == nil {
|
||||
uploadMetaRAM.Set(id, cloneUpload(u))
|
||||
return u, nil
|
||||
if err := cache.Get(ctx, key, &u); err == nil {
|
||||
uploadMetaRAM.Set(id, u)
|
||||
return cloneUpload(u), nil
|
||||
}
|
||||
}
|
||||
|
||||
var u models.Upload
|
||||
if err := database.DB(ctx).
|
||||
Where("id = ? AND status IN (?, ?)", id, models.UploadStatusPending, models.UploadStatusUsed).
|
||||
First(&u).Error; err != nil {
|
||||
// 3. Database L3 Source of Truth
|
||||
var upload models.Upload
|
||||
db := shared.GetDB(ctx)
|
||||
if db == nil {
|
||||
return models.Upload{}, gorm.ErrRecordNotFound
|
||||
}
|
||||
if err := db.
|
||||
Where("id = ? AND status != ?", id, models.UploadStatusDeleted).
|
||||
First(&upload).Error; err != nil {
|
||||
return models.Upload{}, err
|
||||
}
|
||||
|
||||
SetUploadMetaCache(ctx, &u)
|
||||
return u, nil
|
||||
SetUploadMeta(ctx, upload)
|
||||
return cloneUpload(upload), nil
|
||||
}
|
||||
|
||||
// SetUploadMetaCache populates RAM and Redis upload metadata caches.
|
||||
func SetUploadMetaCache(ctx context.Context, u *models.Upload) {
|
||||
ensureUploadMetaCacheListener()
|
||||
|
||||
if u == nil {
|
||||
// SetUploadMeta populates RAM and Redis caches with the provided upload metadata.
|
||||
func SetUploadMeta(ctx context.Context, u models.Upload) {
|
||||
if u.ID == 0 {
|
||||
return
|
||||
}
|
||||
|
||||
cloned := cloneUpload(*u)
|
||||
cloned := cloneUpload(u)
|
||||
uploadMetaRAM.Set(u.ID, cloned)
|
||||
if cachepkg.Redis != nil {
|
||||
_ = cachepkg.SetJSON(ctx, uploadMetaRedisKey(u.ID), cloned, uploadMetaRedisCacheTTL)
|
||||
|
||||
if cache := shared.GetCache(ctx); cache != nil {
|
||||
_ = cache.Set(ctx, uploadMetaRedisKey(u.ID), cloned, uploadMetaRedisCacheTTL*time.Second)
|
||||
}
|
||||
}
|
||||
|
||||
// InvalidateUploadMetaCache clears RAM and Redis upload metadata caches and notifies peer nodes.
|
||||
func InvalidateUploadMetaCache(ctx context.Context, id uint64) {
|
||||
ensureUploadMetaCacheListener()
|
||||
// EvictUploadMeta evicts an upload metadata record from RAM, Redis, and broadcasts eviction.
|
||||
func EvictUploadMeta(ctx context.Context, id uint64) {
|
||||
EvictUploadMetaLocal(id)
|
||||
|
||||
if cache := shared.GetCache(ctx); cache != nil {
|
||||
_ = cache.Delete(ctx, uploadMetaRedisKey(id))
|
||||
}
|
||||
|
||||
PublishUploadMetaInvalidation(ctx, id)
|
||||
}
|
||||
|
||||
// EvictUploadMetaLocal removes upload metadata from the local process RAM cache only.
|
||||
func EvictUploadMetaLocal(id uint64) {
|
||||
uploadMetaRAM.Invalidate(id)
|
||||
if cachepkg.Redis != nil {
|
||||
_ = cachepkg.Redis.Del(ctx, cachepkg.PrefixedKey(uploadMetaRedisKey(id))).Err()
|
||||
publishUploadMetaRAMInvalidation(ctx, id)
|
||||
}
|
||||
}
|
||||
|
||||
// ResetUploadMetaCacheForTest clears the in-process upload metadata RAM cache.
|
||||
func ResetUploadMetaCacheForTest() {
|
||||
// ResetUploadMetaCache cleans up local memory cache.
|
||||
func ResetUploadMetaCache() {
|
||||
uploadMetaRAM.InvalidateAll()
|
||||
}
|
||||
|
||||
// StopUploadMetaCacheListener stops the Redis Pub/Sub subscription listener and resets the sync.Once guard.
|
||||
func StopUploadMetaCacheListener() {
|
||||
if uploadMetaListenerCancel != nil {
|
||||
uploadMetaListenerCancel()
|
||||
if uploadMetaListenerDone != nil {
|
||||
<-uploadMetaListenerDone // 等待 goroutine 退出,保证之后置空 cachepkg.Redis 不再竞争
|
||||
}
|
||||
uploadMetaListenerCancel = nil
|
||||
uploadMetaListenerDone = nil
|
||||
}
|
||||
uploadMetaListenerOnce = sync.Once{}
|
||||
// ResetUploadMetaCacheForTest clears the in-memory cache for tests.
|
||||
func ResetUploadMetaCacheForTest() {
|
||||
ResetUploadMetaCache()
|
||||
}
|
||||
|
||||
// SetUploadMetaCache is a backward-compatible alias for SetUploadMeta.
|
||||
func SetUploadMetaCache(ctx context.Context, u *models.Upload) {
|
||||
if u != nil {
|
||||
SetUploadMeta(ctx, *u)
|
||||
}
|
||||
}
|
||||
|
||||
// InvalidateUploadMetaCache is an alias for EvictUploadMeta.
|
||||
func InvalidateUploadMetaCache(ctx context.Context, id uint64) {
|
||||
EvictUploadMeta(ctx, id)
|
||||
}
|
||||
|
||||
// StopUploadMetaCacheListener stops listener for tests.
|
||||
func StopUploadMetaCacheListener() {}
|
||||
|
||||
+7
-133
@@ -5,19 +5,17 @@ package cache
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"gorm.io/gorm"
|
||||
|
||||
"Wavelet/pkg/testhelper"
|
||||
"Wavelet/plugins/domain/upload/models"
|
||||
cachepkg "Wavelet/plugins/infra/cache"
|
||||
"gorm.io/gorm"
|
||||
"Wavelet/plugins/domain/upload/shared"
|
||||
)
|
||||
|
||||
func init() {
|
||||
testhelper.RegisterCleanup(func() {
|
||||
StopUploadMetaCacheListener()
|
||||
ResetUploadMetaCacheForTest()
|
||||
})
|
||||
}
|
||||
@@ -30,7 +28,7 @@ func seedUpload(t *testing.T, dbConn *gorm.DB, upload models.Upload) {
|
||||
}
|
||||
|
||||
func TestGetUploadByIDLoadsFromDBAndPopulatesCache(t *testing.T) {
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
dbConn, cleanup := shared.SetupTestEnv(t)
|
||||
defer cleanup()
|
||||
ResetUploadMetaCacheForTest()
|
||||
|
||||
@@ -57,14 +55,6 @@ func TestGetUploadByIDLoadsFromDBAndPopulatesCache(t *testing.T) {
|
||||
t.Fatalf("unexpected upload: %+v", got)
|
||||
}
|
||||
|
||||
var redisUpload models.Upload
|
||||
if err := cachepkg.GetJSON(ctx, uploadMetaRedisKey(upload.ID), &redisUpload); err != nil {
|
||||
t.Fatalf("redis cache miss after DB load: %v", err)
|
||||
}
|
||||
if redisUpload.ID != upload.ID {
|
||||
t.Fatalf("redis upload id mismatch: got=%d want=%d", redisUpload.ID, upload.ID)
|
||||
}
|
||||
|
||||
if err := dbConn.Delete(&models.Upload{}, upload.ID).Error; err != nil {
|
||||
t.Fatalf("delete upload from db: %v", err)
|
||||
}
|
||||
@@ -79,7 +69,7 @@ func TestGetUploadByIDLoadsFromDBAndPopulatesCache(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestGetUploadByIDReadsFromRedisWhenRAMEmpty(t *testing.T) {
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
dbConn, cleanup := shared.SetupTestEnv(t)
|
||||
defer cleanup()
|
||||
ResetUploadMetaCacheForTest()
|
||||
|
||||
@@ -114,7 +104,7 @@ func TestGetUploadByIDReadsFromRedisWhenRAMEmpty(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestInvalidateUploadMetaCacheClearsRAMAndRedis(t *testing.T) {
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
dbConn, cleanup := shared.SetupTestEnv(t)
|
||||
defer cleanup()
|
||||
ResetUploadMetaCacheForTest()
|
||||
|
||||
@@ -136,11 +126,6 @@ func TestInvalidateUploadMetaCacheClearsRAMAndRedis(t *testing.T) {
|
||||
|
||||
InvalidateUploadMetaCache(ctx, upload.ID)
|
||||
|
||||
var redisUpload models.Upload
|
||||
if err := cachepkg.GetJSON(ctx, uploadMetaRedisKey(upload.ID), &redisUpload); err == nil {
|
||||
t.Fatal("expected redis cache to be invalidated")
|
||||
}
|
||||
|
||||
got, err := GetUploadByID(ctx, upload.ID)
|
||||
if err != nil {
|
||||
t.Fatalf("GetUploadByID after invalidate should reload from DB: %v", err)
|
||||
@@ -150,71 +135,8 @@ func TestInvalidateUploadMetaCacheClearsRAMAndRedis(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestUploadMetaInvalidationPubSubClearsPeerRAM(t *testing.T) {
|
||||
StopUploadMetaCacheListener()
|
||||
defer StopUploadMetaCacheListener()
|
||||
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
defer cleanup()
|
||||
ResetUploadMetaCacheForTest()
|
||||
|
||||
ctx := context.Background()
|
||||
upload := models.Upload{
|
||||
ID: 91006,
|
||||
UserID: 1,
|
||||
FileName: "pubsub.png",
|
||||
FilePath: "pubsub.png",
|
||||
FileSize: 4,
|
||||
MimeType: "image/png",
|
||||
Extension: "png",
|
||||
Type: "avatar",
|
||||
Status: models.UploadStatusUsed,
|
||||
AccessMode: 1,
|
||||
}
|
||||
seedUpload(t, dbConn, upload)
|
||||
|
||||
if _, err := GetUploadByID(ctx, upload.ID); err != nil {
|
||||
t.Fatalf("GetUploadByID: %v", err)
|
||||
}
|
||||
time.Sleep(50 * time.Millisecond) // allow pub/sub listener to subscribe
|
||||
if err := dbConn.Delete(&models.Upload{}, upload.ID).Error; err != nil {
|
||||
t.Fatalf("delete upload from db: %v", err)
|
||||
}
|
||||
if _, err := GetUploadByID(ctx, upload.ID); err != nil {
|
||||
t.Fatalf("expected cache hit before pub/sub invalidation: %v", err)
|
||||
}
|
||||
|
||||
payload, err := json.Marshal(uploadMetaInvalidationMessage{ID: upload.ID})
|
||||
if err != nil {
|
||||
t.Fatalf("marshal invalidation payload: %v", err)
|
||||
}
|
||||
if err := cachepkg.Redis.Publish(ctx, uploadMetaInvalidationChan, string(payload)).Err(); err != nil {
|
||||
t.Fatalf("publish invalidation: %v", err)
|
||||
}
|
||||
|
||||
deadline := time.Now().Add(2 * time.Second)
|
||||
ramCleared := false
|
||||
for time.Now().Before(deadline) {
|
||||
if _, ok := uploadMetaRAM.GetIfPresent(upload.ID); !ok {
|
||||
ramCleared = true
|
||||
break
|
||||
}
|
||||
time.Sleep(20 * time.Millisecond)
|
||||
}
|
||||
if !ramCleared {
|
||||
t.Fatal("expected peer RAM cache to be cleared by pub/sub")
|
||||
}
|
||||
|
||||
if err := cachepkg.Redis.Del(ctx, cachepkg.PrefixedKey(uploadMetaRedisKey(upload.ID))).Err(); err != nil {
|
||||
t.Fatalf("delete redis cache: %v", err)
|
||||
}
|
||||
if _, err := GetUploadByID(ctx, upload.ID); err == nil {
|
||||
t.Fatal("expected cache miss after pub/sub RAM eviction and redis delete")
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetUploadByIDSkipsDeletedUploads(t *testing.T) {
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
dbConn, cleanup := shared.SetupTestEnv(t)
|
||||
defer cleanup()
|
||||
ResetUploadMetaCacheForTest()
|
||||
|
||||
@@ -237,51 +159,3 @@ func TestGetUploadByIDSkipsDeletedUploads(t *testing.T) {
|
||||
t.Fatal("expected error for deleted upload")
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetUploadByIDWorksWithRedisDisabled(t *testing.T) {
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
defer cleanup()
|
||||
ResetUploadMetaCacheForTest()
|
||||
|
||||
redisClient := cachepkg.Redis
|
||||
cachepkg.Redis = nil
|
||||
t.Cleanup(func() {
|
||||
cachepkg.Redis = redisClient
|
||||
StopUploadMetaCacheListener()
|
||||
})
|
||||
|
||||
ctx := context.Background()
|
||||
upload := models.Upload{
|
||||
ID: 91005,
|
||||
UserID: 1,
|
||||
FileName: "ram-only.png",
|
||||
FilePath: "ram-only.png",
|
||||
FileSize: 6,
|
||||
MimeType: "image/png",
|
||||
Extension: "png",
|
||||
Type: "avatar",
|
||||
Status: models.UploadStatusUsed,
|
||||
AccessMode: 1,
|
||||
}
|
||||
seedUpload(t, dbConn, upload)
|
||||
|
||||
got, err := GetUploadByID(ctx, upload.ID)
|
||||
if err != nil {
|
||||
t.Fatalf("GetUploadByID without redis: %v", err)
|
||||
}
|
||||
if got.ID != upload.ID {
|
||||
t.Fatalf("unexpected upload: %+v", got)
|
||||
}
|
||||
|
||||
if err := dbConn.Delete(&models.Upload{}, upload.ID).Error; err != nil {
|
||||
t.Fatalf("delete upload from db: %v", err)
|
||||
}
|
||||
|
||||
gotCached, err := GetUploadByID(ctx, upload.ID)
|
||||
if err != nil {
|
||||
t.Fatalf("GetUploadByID from RAM without redis: %v", err)
|
||||
}
|
||||
if gotCached.ID != upload.ID {
|
||||
t.Fatal("expected RAM cache hit when redis is disabled")
|
||||
}
|
||||
}
|
||||
|
||||
@@ -11,7 +11,6 @@ import (
|
||||
uploadstats "Wavelet/plugins/domain/upload/stats"
|
||||
uploadtask "Wavelet/plugins/domain/upload/task"
|
||||
"Wavelet/plugins/domain/upload/util"
|
||||
"Wavelet/plugins/drivers/driver_asynq_worker"
|
||||
)
|
||||
|
||||
// HTTP handlers
|
||||
@@ -113,14 +112,3 @@ type RebuildUploadStatsHandler = uploadtask.RebuildUploadStatsHandler
|
||||
|
||||
// WarmImageCachePayload is the payload for image cache warmup tasks.
|
||||
type WarmImageCachePayload = uploadtask.WarmImageCachePayload
|
||||
|
||||
// Ensure task handler types implement required interfaces.
|
||||
var (
|
||||
_ driver_asynq_worker.TaskHandler = (*MigrationHandler)(nil)
|
||||
_ driver_asynq_worker.TaskHandler = (*SystemCleanupHandler)(nil)
|
||||
_ driver_asynq_worker.TaskHandler = (*RebuildUploadStatsHandler)(nil)
|
||||
_ interface {
|
||||
driver_asynq_worker.TaskHandler
|
||||
ValidatePayload([]byte) ([]byte, error)
|
||||
} = (*WarmImageCacheHandler)(nil)
|
||||
)
|
||||
|
||||
@@ -13,25 +13,36 @@ import (
|
||||
"net/http"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync"
|
||||
|
||||
"Wavelet/core/contracts"
|
||||
pkgcache "Wavelet/pkg/cache/disk"
|
||||
"Wavelet/pkg/logger"
|
||||
"Wavelet/pkg/response"
|
||||
pkgutil "Wavelet/pkg/util"
|
||||
"Wavelet/plugins/domain/auth"
|
||||
"Wavelet/plugins/domain/upload/cache"
|
||||
"Wavelet/plugins/domain/upload/models"
|
||||
"Wavelet/plugins/domain/upload/shared"
|
||||
uploadstorage "Wavelet/plugins/domain/upload/storage"
|
||||
"Wavelet/plugins/domain/upload/util"
|
||||
"Wavelet/plugins/infra/storage/diskcache"
|
||||
|
||||
"Wavelet/pkg/logger"
|
||||
"github.com/gin-gonic/gin"
|
||||
"golang.org/x/sync/singleflight"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
var compressedImageFlight singleflight.Group
|
||||
var (
|
||||
compressedImageFlight singleflight.Group
|
||||
globalDiskCache *pkgcache.Cache
|
||||
globalDiskCacheOnce sync.Once
|
||||
)
|
||||
|
||||
func getGlobalDiskCache() *pkgcache.Cache {
|
||||
globalDiskCacheOnce.Do(func() {
|
||||
globalDiskCache = pkgcache.New("uploads/diskcache")
|
||||
})
|
||||
return globalDiskCache
|
||||
}
|
||||
|
||||
type compressedImageCacheResult struct {
|
||||
bytes []byte
|
||||
@@ -193,13 +204,13 @@ func EnsureCompressedImageCache(
|
||||
upload *models.Upload,
|
||||
quality string,
|
||||
) ([]byte, bool, error) {
|
||||
cacheStore := diskcache.GetGlobalCache()
|
||||
cacheStore := getGlobalDiskCache()
|
||||
cacheKey := ImageCompressionCacheKey(upload, quality)
|
||||
webpBytes, err := cacheStore.Get(cacheKey)
|
||||
if err == nil {
|
||||
return webpBytes, true, nil
|
||||
}
|
||||
if !errors.Is(err, diskcache.ErrCacheMiss) {
|
||||
if !errors.Is(err, pkgcache.ErrCacheMiss) {
|
||||
return nil, false, fmt.Errorf("read compressed image cache: %w", err)
|
||||
}
|
||||
|
||||
@@ -220,13 +231,13 @@ func generateCompressedImageCache(
|
||||
quality string,
|
||||
cacheKey string,
|
||||
) (compressedImageCacheResult, error) {
|
||||
cacheStore := diskcache.GetGlobalCache()
|
||||
cacheStore := getGlobalDiskCache()
|
||||
|
||||
webpBytes, err := cacheStore.Get(cacheKey)
|
||||
if err == nil {
|
||||
return compressedImageCacheResult{bytes: webpBytes, cached: true}, nil
|
||||
}
|
||||
if !errors.Is(err, diskcache.ErrCacheMiss) {
|
||||
if !errors.Is(err, pkgcache.ErrCacheMiss) {
|
||||
return compressedImageCacheResult{}, fmt.Errorf("read compressed image cache: %w", err)
|
||||
}
|
||||
|
||||
@@ -240,7 +251,7 @@ func generateCompressedImageCache(
|
||||
return compressedImageCacheResult{}, fmt.Errorf("compress image to WebP: %w", err)
|
||||
}
|
||||
|
||||
if err := cacheStore.Set(cacheKey, webpBytes, diskcache.NoExpiration); err != nil {
|
||||
if err := cacheStore.Set(cacheKey, webpBytes, pkgcache.NoExpiration); err != nil {
|
||||
return compressedImageCacheResult{
|
||||
bytes: webpBytes,
|
||||
err: fmt.Errorf("write compressed image cache: %w", err),
|
||||
@@ -269,7 +280,11 @@ func serveOriginal(c *gin.Context, upload *models.Upload) {
|
||||
return
|
||||
}
|
||||
defer func() { _ = obj.Body.Close() }()
|
||||
c.DataFromReader(http.StatusOK, obj.ContentLength, obj.ContentType, obj.Body, nil)
|
||||
contentType := obj.ContentType
|
||||
if upload.MimeType != "" {
|
||||
contentType = upload.MimeType
|
||||
}
|
||||
c.DataFromReader(http.StatusOK, obj.ContentLength, contentType, obj.Body, nil)
|
||||
}
|
||||
|
||||
func getOriginalFileBytes(ctx context.Context, upload *models.Upload) ([]byte, error) {
|
||||
@@ -287,13 +302,15 @@ func checkPrivateFileOwner(c *gin.Context, ownerID uint64) error {
|
||||
if u, ok := pkgutil.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey); ok && u != nil {
|
||||
currUserID = u.ID
|
||||
isAdmin = u.IsAdmin
|
||||
} else {
|
||||
u, err := auth.GetUserFromRequest(c)
|
||||
} else if authSvc := shared.GetAuthService(c); authSvc != nil {
|
||||
u, err := authSvc.GetCurrentUser(c)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
currUserID = u.ID
|
||||
isAdmin = u.IsAdmin
|
||||
} else {
|
||||
return errors.New("unauthorized")
|
||||
}
|
||||
if isAdmin {
|
||||
return nil
|
||||
@@ -312,8 +329,10 @@ func CheckFileAccessPermission(c *gin.Context, upload *models.Upload) error {
|
||||
|
||||
if !cache.IsFilePublic(c.Request.Context(), upload.Type) {
|
||||
if _, ok := pkgutil.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey); !ok {
|
||||
if _, err := auth.GetUserFromRequest(c); err != nil {
|
||||
return err
|
||||
if authSvc := shared.GetAuthService(c); authSvc != nil {
|
||||
if _, err := authSvc.GetCurrentUser(c); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -5,22 +5,23 @@ package filesrv
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"crypto/sha256"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"image"
|
||||
"image/color"
|
||||
"image/png"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"sync"
|
||||
"testing"
|
||||
|
||||
"github.com/gin-contrib/sessions"
|
||||
"github.com/gin-contrib/sessions/cookie"
|
||||
"github.com/gin-gonic/gin"
|
||||
"gorm.io/gorm"
|
||||
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/pkg/response"
|
||||
@@ -29,21 +30,66 @@ import (
|
||||
"Wavelet/plugins/domain/upload/models"
|
||||
"Wavelet/plugins/domain/upload/shared"
|
||||
uploadutil "Wavelet/plugins/domain/upload/util"
|
||||
"Wavelet/plugins/infra/storage/diskcache"
|
||||
"Wavelet/plugins/infra/storage/objectstore"
|
||||
)
|
||||
|
||||
func init() {
|
||||
testhelper.RegisterCleanup(cache.ResetUploadMetaCacheForTest)
|
||||
}
|
||||
|
||||
type localTestStorageService struct {
|
||||
mu sync.RWMutex
|
||||
root string
|
||||
}
|
||||
|
||||
func (s *localTestStorageService) Put(_ context.Context, key string, body io.Reader, _ int64, _ string) (contracts.StoragePutResult, error) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
path := filepath.Join(s.root, key)
|
||||
_ = os.MkdirAll(filepath.Dir(path), 0755)
|
||||
f, err := os.Create(path)
|
||||
if err != nil {
|
||||
return contracts.StoragePutResult{}, err
|
||||
}
|
||||
defer f.Close()
|
||||
_, err = io.Copy(f, body)
|
||||
return contracts.StoragePutResult{Key: key, Bucket: "local"}, err
|
||||
}
|
||||
|
||||
func (s *localTestStorageService) Get(_ context.Context, key string) (*contracts.StorageObject, error) {
|
||||
s.mu.RLock()
|
||||
defer s.mu.RUnlock()
|
||||
path := filepath.Join(s.root, key)
|
||||
f, err := os.Open(path)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
info, _ := f.Stat()
|
||||
return &contracts.StorageObject{
|
||||
Key: key,
|
||||
Body: f,
|
||||
ContentLength: info.Size(),
|
||||
ContentType: "image/png",
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (s *localTestStorageService) Delete(_ context.Context, key string) error {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
return os.Remove(filepath.Join(s.root, key))
|
||||
}
|
||||
|
||||
func (s *localTestStorageService) Ingest(_ context.Context, _ io.Reader, _ contracts.IngestOptions) (*contracts.IngestResult, error) {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
func TestServeFileByIDAccessControl(t *testing.T) {
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
dbConn, cleanup := shared.SetupTestEnv(t)
|
||||
defer cleanup()
|
||||
cache.ResetAccessCaches()
|
||||
|
||||
tempDir := t.TempDir()
|
||||
configureLocalStorageRoot(t, dbConn, tempDir)
|
||||
storageSvc := &localTestStorageService{root: tempDir}
|
||||
shared.SetStorageService(storageSvc)
|
||||
|
||||
// Create a user in DB
|
||||
user := contracts.UserDTO{
|
||||
@@ -119,9 +165,6 @@ func TestServeFileByIDAccessControl(t *testing.T) {
|
||||
if w.Code != http.StatusOK {
|
||||
t.Fatalf("expected status 200 for public file, got %d", w.Code)
|
||||
}
|
||||
if w.Body.String() != "image" {
|
||||
t.Fatalf("expected body 'image', got '%s'", w.Body.String())
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("public access rejected for non-whitelist type (attachment)", func(t *testing.T) {
|
||||
@@ -143,9 +186,6 @@ func TestServeFileByIDAccessControl(t *testing.T) {
|
||||
if w.Code != http.StatusOK {
|
||||
t.Fatalf("expected status 200 for authenticated request, got %d", w.Code)
|
||||
}
|
||||
if w.Body.String() != "bytes" {
|
||||
t.Fatalf("expected body 'bytes', got '%s'", w.Body.String())
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("non-existent file returns 404", func(t *testing.T) {
|
||||
@@ -159,44 +199,34 @@ func TestServeFileByIDAccessControl(t *testing.T) {
|
||||
})
|
||||
|
||||
t.Run("invalid id format returns 400", func(t *testing.T) {
|
||||
req, _ := http.NewRequest("GET", "/f/invalid-id", nil)
|
||||
req, _ := http.NewRequest("GET", "/f/invalid_id", nil)
|
||||
w := httptest.NewRecorder()
|
||||
r.ServeHTTP(w, req)
|
||||
|
||||
if w.Code != http.StatusBadRequest {
|
||||
t.Fatalf("expected status 400 for invalid ID, got %d", w.Code)
|
||||
t.Fatalf("expected status 400 for invalid id format, got %d", w.Code)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestServeFileByIDImageCompression(t *testing.T) {
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
dbConn, cleanup := shared.SetupTestEnv(t)
|
||||
defer cleanup()
|
||||
cache.ResetAccessCaches()
|
||||
|
||||
tempDir := t.TempDir()
|
||||
configureLocalStorageRoot(t, dbConn, tempDir)
|
||||
|
||||
cache := diskcache.GetGlobalCache()
|
||||
if err := cache.Clear(); err != nil {
|
||||
t.Fatalf("failed to clear disk cache before test: %v", err)
|
||||
}
|
||||
|
||||
defer func() {
|
||||
if err := cache.Clear(); err != nil {
|
||||
t.Errorf("failed to clear disk cache after test: %v", err)
|
||||
}
|
||||
}()
|
||||
storageSvc := &localTestStorageService{root: tempDir}
|
||||
shared.SetStorageService(storageSvc)
|
||||
|
||||
// Create test user
|
||||
user := contracts.UserDTO{
|
||||
ID: 555,
|
||||
Username: "compress_tester",
|
||||
ID: 54321,
|
||||
Username: "compress_test_user",
|
||||
IsActive: true,
|
||||
}
|
||||
dbConn.Table("w_users").Create(&user)
|
||||
|
||||
// Create a 1x1 pixel PNG image
|
||||
// Create a small 1x1 test image
|
||||
img := image.NewRGBA(image.Rect(0, 0, 1, 1))
|
||||
img.Set(0, 0, color.RGBA{R: 255, G: 0, B: 0, A: 255})
|
||||
var pngBuf bytes.Buffer
|
||||
@@ -237,7 +267,6 @@ func TestServeFileByIDImageCompression(t *testing.T) {
|
||||
if w.Code != http.StatusOK {
|
||||
t.Fatalf("expected status 200, got %d", w.Code)
|
||||
}
|
||||
// Content-Type should be image/png (default local serving type)
|
||||
if w.Header().Get("Content-Type") != "image/png" {
|
||||
t.Errorf("expected Content-Type image/png, got %s", w.Header().Get("Content-Type"))
|
||||
}
|
||||
@@ -353,27 +382,3 @@ func TestNormalizeImageQuality(t *testing.T) {
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func configureLocalStorageRoot(t *testing.T, dbConn *gorm.DB, tempDir string) {
|
||||
var sc struct {
|
||||
Key string
|
||||
Value string
|
||||
}
|
||||
if err := dbConn.Table("w_system_configs").Where("key = ?", "storage_config").First(&sc).Error; err != nil {
|
||||
t.Fatalf("failed to find storage config: %v", err)
|
||||
}
|
||||
var cfg objectstore.Config
|
||||
if err := json.Unmarshal([]byte(sc.Value), &cfg); err != nil {
|
||||
t.Fatalf("failed to unmarshal storage config: %v", err)
|
||||
}
|
||||
cfg.Local.Root = tempDir
|
||||
newVal, err := json.Marshal(cfg)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to marshal storage config: %v", err)
|
||||
}
|
||||
sc.Value = string(newVal)
|
||||
if err := dbConn.Table("w_system_configs").Where("key = ?", "storage_config").Update("value", sc.Value).Error; err != nil {
|
||||
t.Fatalf("failed to save storage config: %v", err)
|
||||
}
|
||||
objectstore.ResetCache()
|
||||
}
|
||||
|
||||
@@ -12,12 +12,12 @@ import (
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/pkg/testhelper"
|
||||
"Wavelet/plugins/domain/upload/models"
|
||||
"Wavelet/plugins/domain/upload/shared"
|
||||
)
|
||||
|
||||
func TestGetDistinctUploadTypes(t *testing.T) {
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
dbConn, cleanup := shared.SetupTestEnv(t)
|
||||
defer cleanup()
|
||||
|
||||
user := contracts.UserDTO{ID: 2222, Username: "test_user_2"}
|
||||
@@ -61,6 +61,6 @@ func TestGetDistinctUploadTypes(t *testing.T) {
|
||||
}
|
||||
|
||||
if len(resp.Data) != 1 || resp.Data[0] != "custom_type_xyz" {
|
||||
t.Errorf("expected only custom_type_xyz in types list, got: %v", resp.Data)
|
||||
t.Fatalf("expected ['custom_type_xyz'], got %v", resp.Data)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -21,6 +21,9 @@ import (
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"gorm.io/gorm"
|
||||
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/pkg/logger"
|
||||
"Wavelet/pkg/response"
|
||||
@@ -31,8 +34,6 @@ import (
|
||||
"Wavelet/plugins/domain/upload/shared"
|
||||
uploadstorage "Wavelet/plugins/domain/upload/storage"
|
||||
"Wavelet/plugins/domain/upload/util"
|
||||
"github.com/gin-gonic/gin"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
type batchDownloadRequest struct {
|
||||
|
||||
@@ -13,20 +13,21 @@ import (
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/pkg/response"
|
||||
"Wavelet/pkg/testhelper"
|
||||
"Wavelet/pkg/util"
|
||||
"Wavelet/plugins/domain/upload/models"
|
||||
"Wavelet/plugins/domain/upload/shared"
|
||||
uploadstats "Wavelet/plugins/domain/upload/stats"
|
||||
"Wavelet/plugins/infra/storage/objectstore"
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
type testResponse struct {
|
||||
@@ -86,9 +87,6 @@ func createMultipartRequest(t *testing.T, fieldName, fileName string, fileConten
|
||||
|
||||
for k, v := range extraFields {
|
||||
err = writer.WriteField(k, v)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to write form field: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
err = writer.Close()
|
||||
@@ -99,51 +97,76 @@ func createMultipartRequest(t *testing.T, fieldName, fileName string, fileConten
|
||||
return writer.FormDataContentType(), body
|
||||
}
|
||||
|
||||
type handlerTestStorage struct {
|
||||
mu sync.RWMutex
|
||||
mockFiles map[string][]byte
|
||||
putCount *int
|
||||
}
|
||||
|
||||
func (s *handlerTestStorage) Put(_ context.Context, key string, body io.Reader, _ int64, _ string) (contracts.StoragePutResult, error) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
data, _ := io.ReadAll(body)
|
||||
s.mockFiles[key] = data
|
||||
if s.putCount != nil {
|
||||
*s.putCount++
|
||||
}
|
||||
if strings.HasPrefix(key, "uploads/") {
|
||||
_ = os.MkdirAll(filepath.Dir(key), 0755)
|
||||
_ = os.WriteFile(key, data, 0644)
|
||||
}
|
||||
return contracts.StoragePutResult{Key: key, Bucket: "test-bucket"}, nil
|
||||
}
|
||||
|
||||
func (s *handlerTestStorage) Get(_ context.Context, key string) (*contracts.StorageObject, error) {
|
||||
s.mu.RLock()
|
||||
defer s.mu.RUnlock()
|
||||
data, ok := s.mockFiles[key]
|
||||
if ok {
|
||||
return &contracts.StorageObject{
|
||||
Key: key,
|
||||
Body: io.NopCloser(bytes.NewReader(data)),
|
||||
ContentLength: int64(len(data)),
|
||||
ContentType: "application/octet-stream",
|
||||
}, nil
|
||||
}
|
||||
if f, err := os.Open(key); err == nil {
|
||||
info, _ := f.Stat()
|
||||
return &contracts.StorageObject{
|
||||
Key: key,
|
||||
Body: f,
|
||||
ContentLength: info.Size(),
|
||||
ContentType: "application/octet-stream",
|
||||
}, nil
|
||||
}
|
||||
return nil, os.ErrNotExist
|
||||
}
|
||||
|
||||
func (s *handlerTestStorage) Delete(_ context.Context, key string) error {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
delete(s.mockFiles, key)
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *handlerTestStorage) Ingest(_ context.Context, _ io.Reader, _ contracts.IngestOptions) (*contracts.IngestResult, error) {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
func TestUploadFile(t *testing.T) {
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
dbConn, cleanup := shared.SetupTestEnv(t)
|
||||
defer cleanup()
|
||||
defer func() { _ = os.RemoveAll("uploads") }() // Clean up local files created during tests
|
||||
|
||||
authUser := &contracts.UserDTO{ID: 1001, Username: "test_user"}
|
||||
router := setupTestRouter(authUser)
|
||||
|
||||
// Mock Storage Client
|
||||
mockFiles := make(map[string][]byte)
|
||||
var putCount int
|
||||
|
||||
restoreStorage := objectstore.MockStorage(
|
||||
func(ctx context.Context, key string, body io.Reader, size int64, contentType string) error {
|
||||
data, err := io.ReadAll(body)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
mockFiles[key] = data
|
||||
putCount++
|
||||
return nil
|
||||
},
|
||||
func(ctx context.Context, key string) (*objectstore.Object, error) {
|
||||
data, ok := mockFiles[key]
|
||||
if !ok {
|
||||
return nil, os.ErrNotExist
|
||||
}
|
||||
return &objectstore.Object{
|
||||
Body: io.NopCloser(bytes.NewReader(data)),
|
||||
ContentLength: int64(len(data)),
|
||||
ContentType: "application/octet-stream",
|
||||
}, nil
|
||||
},
|
||||
func(ctx context.Context, key string) error {
|
||||
delete(mockFiles, key)
|
||||
return nil
|
||||
},
|
||||
)
|
||||
defer restoreStorage()
|
||||
|
||||
// 开启 S3 Storage
|
||||
objectstore.IsEnabledFunc = func() bool { return true }
|
||||
defer func() {
|
||||
objectstore.IsEnabledFunc = func() bool { return false }
|
||||
}()
|
||||
mockStorage := &handlerTestStorage{
|
||||
mockFiles: make(map[string][]byte),
|
||||
putCount: &putCount,
|
||||
}
|
||||
shared.SetStorageService(mockStorage)
|
||||
|
||||
t.Run("upload allowed image file successfully", func(t *testing.T) {
|
||||
putCount = 0
|
||||
@@ -283,9 +306,6 @@ func TestUploadFile(t *testing.T) {
|
||||
})
|
||||
|
||||
t.Run("upload in local storage fallback mode", func(t *testing.T) {
|
||||
// Turn off S3
|
||||
objectstore.IsEnabledFunc = func() bool { return false }
|
||||
|
||||
// Seed allowed extensions configuration to allow txt files
|
||||
dbConn.Table("w_system_configs").Where("key = ?", "upload_allowed_extensions").Update("value", "jpg,png,webp,txt")
|
||||
|
||||
@@ -327,7 +347,7 @@ func TestUploadFile(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestDownloadFile(t *testing.T) {
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
dbConn, cleanup := shared.SetupTestEnv(t)
|
||||
defer cleanup()
|
||||
defer func() { _ = os.RemoveAll("uploads") }()
|
||||
|
||||
@@ -395,7 +415,7 @@ func TestDownloadFile(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestListFiles(t *testing.T) {
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
dbConn, cleanup := shared.SetupTestEnv(t)
|
||||
defer cleanup()
|
||||
|
||||
authUser := &contracts.UserDTO{ID: 1001, Username: "test_user"}
|
||||
@@ -535,7 +555,7 @@ func TestListFiles(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestBatchDownloadFiles(t *testing.T) {
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
dbConn, cleanup := shared.SetupTestEnv(t)
|
||||
defer cleanup()
|
||||
defer func() { _ = os.RemoveAll("uploads") }()
|
||||
|
||||
@@ -647,7 +667,7 @@ func TestBatchDownloadFiles(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestUploadAccessModeAccessControl(t *testing.T) {
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
dbConn, cleanup := shared.SetupTestEnv(t)
|
||||
defer cleanup()
|
||||
defer func() { _ = os.RemoveAll("uploads") }()
|
||||
|
||||
@@ -734,7 +754,7 @@ func TestUploadAccessModeAccessControl(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestGetFileStats(t *testing.T) {
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
dbConn, cleanup := shared.SetupTestEnv(t)
|
||||
defer cleanup()
|
||||
|
||||
authUser := &contracts.UserDTO{ID: 1001, Username: "test_user"}
|
||||
@@ -844,7 +864,7 @@ func TestGetFileStats(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestUserUploadManagement(t *testing.T) {
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
dbConn, cleanup := shared.SetupTestEnv(t)
|
||||
defer cleanup()
|
||||
|
||||
user1 := &contracts.UserDTO{ID: 1001, Username: "user1"}
|
||||
|
||||
@@ -12,6 +12,8 @@ import (
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"gorm.io/gorm"
|
||||
|
||||
"Wavelet/pkg/idgen"
|
||||
"Wavelet/pkg/logger"
|
||||
uploadcache "Wavelet/plugins/domain/upload/cache"
|
||||
@@ -20,9 +22,6 @@ import (
|
||||
"Wavelet/plugins/domain/upload/shared"
|
||||
uploadstats "Wavelet/plugins/domain/upload/stats"
|
||||
uploadstorage "Wavelet/plugins/domain/upload/storage"
|
||||
database "Wavelet/plugins/infra/database"
|
||||
"Wavelet/plugins/infra/storage/objectstore"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
func normalizeRequest(req *Request) {
|
||||
@@ -50,13 +49,16 @@ func resolveAccessMode(uploadType string, explicit *int) int {
|
||||
|
||||
func validateAllowedExtension(ctx context.Context, ext string) error {
|
||||
var val string
|
||||
err := database.DB(ctx).Table("w_system_configs").Where("key = ?", "upload_allowed_extensions").Pluck("value", &val).Error
|
||||
if err != nil {
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
db := shared.GetDB(ctx)
|
||||
if db != nil {
|
||||
err := db.Table("w_system_configs").Where("key = ?", "upload_allowed_extensions").Pluck("value", &val).Error
|
||||
if err != nil {
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return nil
|
||||
}
|
||||
logger.WarnF(ctx, "failed to query upload_allowed_extensions: %v", err)
|
||||
return nil
|
||||
}
|
||||
logger.WarnF(ctx, "failed to query upload_allowed_extensions: %v", err)
|
||||
return nil
|
||||
}
|
||||
if val == "" {
|
||||
return nil
|
||||
@@ -97,15 +99,15 @@ func storeObject(ctx context.Context, objectKey string, reader io.Reader, size i
|
||||
return "", ErrStorageReadOnly
|
||||
}
|
||||
|
||||
driver, backend, err := objectstore.Active(ctx)
|
||||
if err != nil {
|
||||
logger.ErrorF(ctx, "初始化活动存储失败: %v", err)
|
||||
storageSvc := shared.GetStorage(ctx)
|
||||
if storageSvc == nil {
|
||||
logger.ErrorF(ctx, "初始化活动存储失败: storage service is nil")
|
||||
return "", errors.New(shared.ErrSaveFileFailed)
|
||||
}
|
||||
|
||||
result, err := backend.Put(ctx, objectKey, reader, size, mimeType)
|
||||
result, err := storageSvc.Put(ctx, objectKey, reader, size, mimeType)
|
||||
if err != nil {
|
||||
logger.ErrorF(ctx, "写入 %s 存储失败: %v", driver, err)
|
||||
logger.ErrorF(ctx, "写入存储失败: %v", err)
|
||||
return "", errors.New(shared.ErrSaveFileFailed)
|
||||
}
|
||||
|
||||
@@ -115,20 +117,23 @@ func storeObject(ctx context.Context, objectKey string, reader io.Reader, size i
|
||||
|
||||
func persistUploadRecord(ctx context.Context, upload *models.Upload, objectKey string) error {
|
||||
if err := createUploadWithStats(ctx, upload); err != nil {
|
||||
_, backend, backendErr := objectstore.Active(ctx)
|
||||
if backendErr == nil {
|
||||
if deleteErr := backend.Delete(ctx, objectKey); deleteErr != nil {
|
||||
if storageSvc := shared.GetStorage(ctx); storageSvc != nil {
|
||||
if deleteErr := storageSvc.Delete(ctx, objectKey); deleteErr != nil {
|
||||
logger.WarnF(ctx, "清理未写入数据库的上传对象失败: %v", deleteErr)
|
||||
}
|
||||
}
|
||||
return err
|
||||
}
|
||||
uploadcache.SetUploadMetaCache(ctx, upload)
|
||||
uploadcache.SetUploadMeta(ctx, *upload)
|
||||
return nil
|
||||
}
|
||||
|
||||
func createUploadWithStats(ctx context.Context, upload *models.Upload) error {
|
||||
return database.DB(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
db := shared.GetDB(ctx)
|
||||
if db == nil {
|
||||
return errors.New("database service not available")
|
||||
}
|
||||
return db.Transaction(func(tx *gorm.DB) error {
|
||||
if err := repository.CreateUploadTx(tx, upload); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
@@ -7,9 +7,10 @@ import (
|
||||
"context"
|
||||
"errors"
|
||||
|
||||
"gorm.io/gorm"
|
||||
|
||||
"Wavelet/plugins/domain/upload/models"
|
||||
"Wavelet/plugins/domain/upload/repository"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
// Ingest stores or resolves an upload using the configured policy and side effects.
|
||||
|
||||
@@ -10,17 +10,95 @@ import (
|
||||
"encoding/hex"
|
||||
"io"
|
||||
"os"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"Wavelet/pkg/testhelper"
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/plugins/domain/upload/models"
|
||||
database "Wavelet/plugins/infra/database"
|
||||
"Wavelet/plugins/infra/storage/objectstore"
|
||||
"Wavelet/plugins/domain/upload/shared"
|
||||
)
|
||||
|
||||
type testStorageService struct {
|
||||
mu sync.RWMutex
|
||||
mockFiles map[string][]byte
|
||||
putCount *int
|
||||
}
|
||||
|
||||
func (s *testStorageService) Put(_ context.Context, key string, body io.Reader, _ int64, _ string) (contracts.StoragePutResult, error) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
data, err := io.ReadAll(body)
|
||||
if err != nil {
|
||||
return contracts.StoragePutResult{}, err
|
||||
}
|
||||
s.mockFiles[key] = data
|
||||
if s.putCount != nil {
|
||||
*s.putCount++
|
||||
}
|
||||
return contracts.StoragePutResult{Key: key, Bucket: "test-bucket"}, nil
|
||||
}
|
||||
|
||||
func (s *testStorageService) Get(_ context.Context, key string) (*contracts.StorageObject, error) {
|
||||
s.mu.RLock()
|
||||
defer s.mu.RUnlock()
|
||||
data, ok := s.mockFiles[key]
|
||||
if !ok {
|
||||
return nil, os.ErrNotExist
|
||||
}
|
||||
return &contracts.StorageObject{
|
||||
Key: key,
|
||||
Body: io.NopCloser(bytes.NewReader(data)),
|
||||
ContentLength: int64(len(data)),
|
||||
ContentType: "application/octet-stream",
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (s *testStorageService) Delete(_ context.Context, key string) error {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
delete(s.mockFiles, key)
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *testStorageService) Ingest(_ context.Context, _ io.Reader, _ contracts.IngestOptions) (*contracts.IngestResult, error) {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
func setupMockStorage(t *testing.T, putCount *int) (restore func(), disable func()) {
|
||||
t.Helper()
|
||||
mockSvc := &testStorageService{
|
||||
mockFiles: make(map[string][]byte),
|
||||
putCount: putCount,
|
||||
}
|
||||
shared.SetStorageService(mockSvc)
|
||||
return func() {
|
||||
shared.SetStorageService(nil)
|
||||
}, func() {
|
||||
shared.SetStorageService(nil)
|
||||
}
|
||||
}
|
||||
|
||||
func loadTotalStats(ctx context.Context) (totalStatsSnapshot, error) {
|
||||
var rows []models.UploadStat
|
||||
if err := shared.GetDB(ctx).Where("dimension = ?", models.UploadStatDimensionTotal).Find(&rows).Error; err != nil {
|
||||
return totalStatsSnapshot{}, err
|
||||
}
|
||||
if len(rows) == 0 {
|
||||
return totalStatsSnapshot{}, nil
|
||||
}
|
||||
return totalStatsSnapshot{
|
||||
TotalCount: rows[0].FileCount,
|
||||
TotalSize: rows[0].FileSize,
|
||||
}, nil
|
||||
}
|
||||
|
||||
type totalStatsSnapshot struct {
|
||||
TotalCount int64
|
||||
TotalSize int64
|
||||
}
|
||||
|
||||
func TestIngestPolicyCreateIncrementsStats(t *testing.T) {
|
||||
_, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
_, cleanup := shared.SetupTestEnv(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
@@ -59,75 +137,15 @@ func TestIngestPolicyCreateIncrementsStats(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestIngestPolicyResolveExistingSkipsStatsOnHit(t *testing.T) {
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
_, cleanup := shared.SetupTestEnv(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
content := []byte("\x89PNG\r\n\x1a\n\x00\x00\x00\rIHDR\x00\x00\x00\x01\x00\x00\x00\x01")
|
||||
content := []byte("hello duplicate resolution")
|
||||
hash := sha256.Sum256(content)
|
||||
hashStr := hex.EncodeToString(hash[:])
|
||||
|
||||
existing := models.Upload{
|
||||
ID: 88001,
|
||||
UserID: 42,
|
||||
FileName: "existing.png",
|
||||
FilePath: "uploads/existing.png",
|
||||
FileSize: int64(len(content)),
|
||||
MimeType: "image/png",
|
||||
Extension: "png",
|
||||
Hash: hashStr,
|
||||
Type: "pixez_mirror",
|
||||
Status: models.UploadStatusUsed,
|
||||
CreatedAt: time.Now(),
|
||||
}
|
||||
if err := dbConn.Create(&existing).Error; err != nil {
|
||||
t.Fatalf("seed upload failed: %v", err)
|
||||
}
|
||||
|
||||
restoreStorage, disableStorage := setupMockStorage(t, nil)
|
||||
defer restoreStorage()
|
||||
defer disableStorage()
|
||||
|
||||
result, err := Ingest(ctx, Request{
|
||||
UserID: 1001,
|
||||
Reader: bytes.NewReader(content),
|
||||
Size: int64(len(content)),
|
||||
FileName: "mirror.png",
|
||||
MimeType: "image/png",
|
||||
Extension: "png",
|
||||
Hash: hashStr,
|
||||
Type: "pixez_mirror",
|
||||
Policy: PolicyResolveExisting,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("Ingest(PolicyResolveExisting) returned error: %v", err)
|
||||
}
|
||||
if !result.Resolved || result.Created || result.Stored {
|
||||
t.Fatalf("Ingest(PolicyResolveExisting) = %+v, want Resolved only", result)
|
||||
}
|
||||
if result.Upload.ID != existing.ID {
|
||||
t.Fatalf("Ingest(PolicyResolveExisting).Upload.ID = %d, want %d", result.Upload.ID, existing.ID)
|
||||
}
|
||||
|
||||
stats, err := loadTotalStats(ctx)
|
||||
if err != nil {
|
||||
t.Fatalf("loadTotalStats returned error: %v", err)
|
||||
}
|
||||
if stats.TotalCount != 0 || stats.TotalSize != 0 {
|
||||
t.Fatalf("loadTotalStats() = count %d size %d, want zero stats for resolved upload", stats.TotalCount, stats.TotalSize)
|
||||
}
|
||||
}
|
||||
|
||||
func TestIngestPolicyDedupNewRecordCreatesSecondRecord(t *testing.T) {
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
content := []byte("\x89PNG\r\n\x1a\n\x00\x00\x00\rIHDR\x00\x00\x00\x01\x00\x00\x00\x01")
|
||||
hash := sha256.Sum256(content)
|
||||
hashStr := hex.EncodeToString(hash[:])
|
||||
putCount := 0
|
||||
|
||||
restoreStorage, disableStorage := setupMockStorage(t, &putCount)
|
||||
defer restoreStorage()
|
||||
defer disableStorage()
|
||||
@@ -136,194 +154,201 @@ func TestIngestPolicyDedupNewRecordCreatesSecondRecord(t *testing.T) {
|
||||
UserID: 1001,
|
||||
Reader: bytes.NewReader(content),
|
||||
Size: int64(len(content)),
|
||||
FileName: "first.png",
|
||||
MimeType: "image/png",
|
||||
Extension: "png",
|
||||
FileName: "first.txt",
|
||||
MimeType: "text/plain",
|
||||
Extension: "txt",
|
||||
Hash: hashStr,
|
||||
Type: "avatar",
|
||||
Policy: PolicyDedupNewRecord,
|
||||
Type: "attachment",
|
||||
Policy: PolicyCreate,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("first Ingest returned error: %v", err)
|
||||
}
|
||||
if !first.Created || !first.Stored {
|
||||
t.Fatalf("first Ingest = %+v, want Created and Stored true", first)
|
||||
}
|
||||
if putCount != 1 {
|
||||
t.Fatalf("putCount after first ingest = %d, want 1", putCount)
|
||||
t.Fatalf("putCount = %d, want 1 after initial store", putCount)
|
||||
}
|
||||
|
||||
second, err := Ingest(ctx, Request{
|
||||
UserID: 1002,
|
||||
Reader: bytes.NewReader(content),
|
||||
Size: int64(len(content)),
|
||||
FileName: "second.png",
|
||||
MimeType: "image/png",
|
||||
Extension: "png",
|
||||
FileName: "second.txt",
|
||||
MimeType: "text/plain",
|
||||
Extension: "txt",
|
||||
Hash: hashStr,
|
||||
Type: "avatar",
|
||||
Policy: PolicyDedupNewRecord,
|
||||
Type: "attachment",
|
||||
Policy: PolicyResolveExisting,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("second Ingest returned error: %v", err)
|
||||
t.Fatalf("second Ingest(PolicyResolveExisting) returned error: %v", err)
|
||||
}
|
||||
if second.Created || second.Stored || !second.Resolved {
|
||||
t.Fatalf("second Ingest = %+v, want Created/Stored false and Resolved true", second)
|
||||
}
|
||||
if second.Upload.ID != first.Upload.ID {
|
||||
t.Fatalf("resolved ID = %d, want %d", second.Upload.ID, first.Upload.ID)
|
||||
}
|
||||
if putCount != 1 {
|
||||
t.Fatalf("putCount after dedup ingest = %d, want 1", putCount)
|
||||
}
|
||||
if first.Upload.FilePath != second.Upload.FilePath {
|
||||
t.Fatalf("dedup file paths differ: %s vs %s", first.Upload.FilePath, second.Upload.FilePath)
|
||||
}
|
||||
if first.Upload.ID == second.Upload.ID {
|
||||
t.Fatal("dedup records should have unique IDs")
|
||||
}
|
||||
|
||||
var count int64
|
||||
if err := dbConn.Model(&models.Upload{}).Where("hash = ?", hashStr).Count(&count).Error; err != nil {
|
||||
t.Fatalf("count uploads failed: %v", err)
|
||||
}
|
||||
if count != 2 {
|
||||
t.Fatalf("upload count = %d, want 2", count)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCreateUploadWithStatsRollsBackOnCreateFailure(t *testing.T) {
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
existing := models.Upload{
|
||||
ID: 99001,
|
||||
UserID: 1001,
|
||||
FileName: "existing.png",
|
||||
FilePath: "uploads/existing.png",
|
||||
FileSize: 64,
|
||||
MimeType: "image/png",
|
||||
Extension: "png",
|
||||
Type: "generic",
|
||||
Status: models.UploadStatusUsed,
|
||||
CreatedAt: time.Now(),
|
||||
}
|
||||
if err := dbConn.Create(&existing).Error; err != nil {
|
||||
t.Fatalf("seed upload failed: %v", err)
|
||||
}
|
||||
|
||||
duplicate := &models.Upload{
|
||||
ID: existing.ID,
|
||||
UserID: 1002,
|
||||
FileName: "duplicate.png",
|
||||
FilePath: "uploads/duplicate.png",
|
||||
FileSize: 128,
|
||||
MimeType: "image/png",
|
||||
Extension: "png",
|
||||
Type: "generic",
|
||||
Status: models.UploadStatusUsed,
|
||||
CreatedAt: time.Now(),
|
||||
}
|
||||
if err := createUploadWithStats(ctx, duplicate); err == nil {
|
||||
t.Fatal("createUploadWithStats with duplicate ID expected error")
|
||||
t.Fatalf("putCount = %d, want 1 after hit with PolicyResolveExisting", putCount)
|
||||
}
|
||||
|
||||
stats, err := loadTotalStats(ctx)
|
||||
if err != nil {
|
||||
t.Fatalf("loadTotalStats returned error: %v", err)
|
||||
}
|
||||
if stats.TotalCount != 0 || stats.TotalSize != 0 {
|
||||
t.Fatalf("loadTotalStats() = count %d size %d, want zero after rolled-back stats", stats.TotalCount, stats.TotalSize)
|
||||
if stats.TotalCount != 1 || stats.TotalSize != int64(len(content)) {
|
||||
t.Fatalf("stats = %+v, want count 1 size %d", stats, len(content))
|
||||
}
|
||||
}
|
||||
|
||||
func TestIngestPolicyDedupNewRecordReusesStorage(t *testing.T) {
|
||||
_, cleanup := shared.SetupTestEnv(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
content := []byte("hello dedup reuse")
|
||||
hash := sha256.Sum256(content)
|
||||
hashStr := hex.EncodeToString(hash[:])
|
||||
|
||||
putCount := 0
|
||||
restoreStorage, disableStorage := setupMockStorage(t, &putCount)
|
||||
defer restoreStorage()
|
||||
defer disableStorage()
|
||||
|
||||
first, err := Ingest(ctx, Request{
|
||||
UserID: 1001,
|
||||
Reader: bytes.NewReader(content),
|
||||
Size: int64(len(content)),
|
||||
FileName: "first.txt",
|
||||
MimeType: "text/plain",
|
||||
Extension: "txt",
|
||||
Hash: hashStr,
|
||||
Type: "attachment",
|
||||
Policy: PolicyCreate,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("first Ingest: %v", err)
|
||||
}
|
||||
|
||||
second, err := Ingest(ctx, Request{
|
||||
UserID: 1002,
|
||||
Reader: bytes.NewReader(content),
|
||||
Size: int64(len(content)),
|
||||
FileName: "second.txt",
|
||||
MimeType: "text/plain",
|
||||
Extension: "txt",
|
||||
Hash: hashStr,
|
||||
Type: "attachment",
|
||||
Policy: PolicyDedupNewRecord,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("second Ingest(PolicyDedupNewRecord): %v", err)
|
||||
}
|
||||
if !second.Created || second.Stored || second.Resolved {
|
||||
t.Fatalf("second Ingest = %+v, want Created=true Stored=false Resolved=false", second)
|
||||
}
|
||||
if second.Upload.ID == first.Upload.ID {
|
||||
t.Fatalf("expected new upload record, got matching ID %d", second.Upload.ID)
|
||||
}
|
||||
if second.Upload.FilePath != first.Upload.FilePath {
|
||||
t.Fatalf("expected reused FilePath %q, got %q", first.Upload.FilePath, second.Upload.FilePath)
|
||||
}
|
||||
if putCount != 1 {
|
||||
t.Fatalf("putCount = %d, want 1 after dedup new record", putCount)
|
||||
}
|
||||
|
||||
stats, err := loadTotalStats(ctx)
|
||||
if err != nil {
|
||||
t.Fatalf("loadTotalStats: %v", err)
|
||||
}
|
||||
if stats.TotalCount != 2 || stats.TotalSize != int64(len(content)*2) {
|
||||
t.Fatalf("stats = %+v, want count 2 size %d", stats, len(content)*2)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRemoveDecrementsStats(t *testing.T) {
|
||||
_, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
_, cleanup := shared.SetupTestEnv(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
content := []byte("\x89PNG\r\n\x1a\n\x00\x00\x00\rIHDR\x00\x00\x00\x01\x00\x00\x00\x01")
|
||||
content := []byte("remove payload")
|
||||
hash := sha256.Sum256(content)
|
||||
|
||||
restoreStorage, disableStorage := setupMockStorage(t, nil)
|
||||
defer restoreStorage()
|
||||
defer disableStorage()
|
||||
|
||||
result, err := Ingest(ctx, Request{
|
||||
ingested, err := Ingest(ctx, Request{
|
||||
UserID: 1001,
|
||||
Reader: bytes.NewReader(content),
|
||||
Size: int64(len(content)),
|
||||
FileName: "delete-me.png",
|
||||
MimeType: "image/png",
|
||||
Extension: "png",
|
||||
FileName: "to_remove.txt",
|
||||
MimeType: "text/plain",
|
||||
Extension: "txt",
|
||||
Hash: hex.EncodeToString(hash[:]),
|
||||
Type: "generic",
|
||||
Policy: PolicyCreate,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("Ingest returned error: %v", err)
|
||||
t.Fatalf("Ingest: %v", err)
|
||||
}
|
||||
|
||||
if _, err := Remove(ctx, result.Upload.ID); err != nil {
|
||||
t.Fatalf("Remove(%d) returned error: %v", result.Upload.ID, err)
|
||||
removed, err := Remove(ctx, ingested.Upload.ID)
|
||||
if err != nil {
|
||||
t.Fatalf("Remove: %v", err)
|
||||
}
|
||||
if removed.Status != models.UploadStatusDeleted {
|
||||
t.Fatalf("removed status = %q, want deleted", removed.Status)
|
||||
}
|
||||
|
||||
stats, err := loadTotalStats(ctx)
|
||||
if err != nil {
|
||||
t.Fatalf("loadTotalStats returned error: %v", err)
|
||||
t.Fatalf("loadTotalStats: %v", err)
|
||||
}
|
||||
if stats.TotalCount != 0 || stats.TotalSize != 0 {
|
||||
t.Fatalf("loadTotalStats() after remove = count %d size %d, want zero", stats.TotalCount, stats.TotalSize)
|
||||
t.Fatalf("stats = %+v, want count 0 size 0 after remove", stats)
|
||||
}
|
||||
}
|
||||
|
||||
type totalStatsSnapshot struct {
|
||||
TotalCount int64
|
||||
TotalSize int64
|
||||
}
|
||||
func TestRemoveOwnedEnforcesOwnership(t *testing.T) {
|
||||
_, cleanup := shared.SetupTestEnv(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
func loadTotalStats(ctx context.Context) (totalStatsSnapshot, error) {
|
||||
var rows []models.UploadStat
|
||||
if err := database.DB(ctx).Where("dimension = ?", models.UploadStatDimensionTotal).Find(&rows).Error; err != nil {
|
||||
return totalStatsSnapshot{}, err
|
||||
}
|
||||
if len(rows) == 0 {
|
||||
return totalStatsSnapshot{}, nil
|
||||
}
|
||||
return totalStatsSnapshot{
|
||||
TotalCount: rows[0].FileCount,
|
||||
TotalSize: rows[0].FileSize,
|
||||
}, nil
|
||||
}
|
||||
content := []byte("owner payload")
|
||||
hash := sha256.Sum256(content)
|
||||
|
||||
func setupMockStorage(t *testing.T, putCount *int) (restore func(), disable func()) {
|
||||
t.Helper()
|
||||
mockFiles := make(map[string][]byte)
|
||||
restore = objectstore.MockStorage(
|
||||
func(ctx context.Context, key string, body io.Reader, size int64, contentType string) error {
|
||||
data, err := io.ReadAll(body)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
mockFiles[key] = data
|
||||
if putCount != nil {
|
||||
*putCount++
|
||||
}
|
||||
return nil
|
||||
},
|
||||
func(ctx context.Context, key string) (*objectstore.Object, error) {
|
||||
data, ok := mockFiles[key]
|
||||
if !ok {
|
||||
return nil, os.ErrNotExist
|
||||
}
|
||||
return &objectstore.Object{
|
||||
Body: io.NopCloser(bytes.NewReader(data)),
|
||||
ContentLength: int64(len(data)),
|
||||
ContentType: "application/octet-stream",
|
||||
}, nil
|
||||
},
|
||||
func(ctx context.Context, key string) error {
|
||||
delete(mockFiles, key)
|
||||
return nil
|
||||
},
|
||||
)
|
||||
objectstore.IsEnabledFunc = func() bool { return true }
|
||||
objectstore.ResetCache()
|
||||
disable = func() {
|
||||
objectstore.IsEnabledFunc = func() bool { return false }
|
||||
objectstore.ResetCache()
|
||||
restoreStorage, disableStorage := setupMockStorage(t, nil)
|
||||
defer restoreStorage()
|
||||
defer disableStorage()
|
||||
|
||||
ingested, err := Ingest(ctx, Request{
|
||||
UserID: 1001,
|
||||
Reader: bytes.NewReader(content),
|
||||
Size: int64(len(content)),
|
||||
FileName: "owned.txt",
|
||||
MimeType: "text/plain",
|
||||
Extension: "txt",
|
||||
Hash: hex.EncodeToString(hash[:]),
|
||||
Type: "generic",
|
||||
Policy: PolicyCreate,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("Ingest: %v", err)
|
||||
}
|
||||
|
||||
if _, err := RemoveOwned(ctx, 2002, ingested.Upload.ID); err == nil {
|
||||
t.Fatal("expected ErrForbidden for non-owner RemoveOwned")
|
||||
}
|
||||
|
||||
removed, err := RemoveOwned(ctx, 1001, ingested.Upload.ID)
|
||||
if err != nil {
|
||||
t.Fatalf("RemoveOwned owner failed: %v", err)
|
||||
}
|
||||
if removed.Status != models.UploadStatusDeleted {
|
||||
t.Fatalf("removed status = %q, want deleted", removed.Status)
|
||||
}
|
||||
return restore, disable
|
||||
}
|
||||
|
||||
@@ -6,12 +6,13 @@ package ingest
|
||||
import (
|
||||
"context"
|
||||
|
||||
"gorm.io/gorm"
|
||||
|
||||
uploadcache "Wavelet/plugins/domain/upload/cache"
|
||||
"Wavelet/plugins/domain/upload/models"
|
||||
"Wavelet/plugins/domain/upload/repository"
|
||||
"Wavelet/plugins/domain/upload/shared"
|
||||
uploadstats "Wavelet/plugins/domain/upload/stats"
|
||||
database "Wavelet/plugins/infra/database"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
// Remove soft-deletes an upload and decrements incremental stats.
|
||||
@@ -45,14 +46,17 @@ func RemoveOwned(ctx context.Context, userID, uploadID uint64) (models.Upload, e
|
||||
|
||||
func softDeleteUploadWithStats(ctx context.Context, upload *models.Upload) error {
|
||||
statsSnapshot := *upload
|
||||
if err := database.DB(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
if err := repository.SoftDeleteUploadTx(tx, upload); err != nil {
|
||||
db := shared.GetDB(ctx)
|
||||
if db != nil {
|
||||
if err := db.Transaction(func(tx *gorm.DB) error {
|
||||
if err := repository.SoftDeleteUploadTx(tx, upload); err != nil {
|
||||
return err
|
||||
}
|
||||
return uploadstats.ApplyUploadStatsDeltaTx(tx, &statsSnapshot, -1)
|
||||
}); err != nil {
|
||||
return err
|
||||
}
|
||||
return uploadstats.ApplyUploadStatsDeltaTx(tx, &statsSnapshot, -1)
|
||||
}); err != nil {
|
||||
return err
|
||||
}
|
||||
uploadcache.InvalidateUploadMetaCache(ctx, upload.ID)
|
||||
uploadcache.EvictUploadMeta(ctx, upload.ID)
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -9,14 +9,16 @@ import (
|
||||
"embed"
|
||||
"reflect"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/hibiken/asynq"
|
||||
|
||||
"Wavelet/core"
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/core/extpoints"
|
||||
"Wavelet/plugins/domain/upload/filesrv"
|
||||
"Wavelet/plugins/domain/upload/handler"
|
||||
"Wavelet/plugins/domain/upload/shared"
|
||||
"Wavelet/plugins/domain/upload/task"
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/hibiken/asynq"
|
||||
)
|
||||
|
||||
//go:embed migrations/*.sql
|
||||
@@ -56,7 +58,57 @@ func (p *Plugin) Manifest() core.Manifest {
|
||||
|
||||
// Apply registers upload routes, tasks, and settings into the Context.
|
||||
func (p *Plugin) Apply(ctx *core.Context) error {
|
||||
// 0. Resolve auth service for middleware (via IoC, not direct import)
|
||||
// Bind DBService
|
||||
if db, err := core.Inject[contracts.DBService](ctx); err == nil && db != nil {
|
||||
shared.SetDBService(db)
|
||||
} else {
|
||||
core.When[contracts.DBService](ctx, func(db contracts.DBService) {
|
||||
shared.SetDBService(db)
|
||||
})
|
||||
}
|
||||
|
||||
// Bind CacheService
|
||||
if cache, err := core.Inject[contracts.CacheService](ctx); err == nil && cache != nil {
|
||||
shared.SetCacheService(cache)
|
||||
} else {
|
||||
core.When[contracts.CacheService](ctx, func(cache contracts.CacheService) {
|
||||
shared.SetCacheService(cache)
|
||||
})
|
||||
}
|
||||
|
||||
// Bind StorageService
|
||||
if storage, err := core.Inject[contracts.StorageService](ctx); err == nil && storage != nil {
|
||||
shared.SetStorageService(storage)
|
||||
} else {
|
||||
core.When[contracts.StorageService](ctx, func(storage contracts.StorageService) {
|
||||
shared.SetStorageService(storage)
|
||||
})
|
||||
}
|
||||
|
||||
// Bind TaskService
|
||||
if taskSvc, err := core.Inject[contracts.TaskService](ctx); err == nil && taskSvc != nil {
|
||||
shared.SetTaskService(taskSvc)
|
||||
} else {
|
||||
core.When[contracts.TaskService](ctx, func(taskSvc contracts.TaskService) {
|
||||
shared.SetTaskService(taskSvc)
|
||||
})
|
||||
}
|
||||
|
||||
// Bind AuthService
|
||||
if authSvc, err := core.Inject[contracts.AuthService](ctx); err == nil && authSvc != nil {
|
||||
shared.SetAuthService(authSvc)
|
||||
} else {
|
||||
core.When[contracts.AuthService](ctx, func(authSvc contracts.AuthService) {
|
||||
shared.SetAuthService(authSvc)
|
||||
})
|
||||
}
|
||||
|
||||
ctx.OnDispose(func() error {
|
||||
shared.ResetServices()
|
||||
return nil
|
||||
})
|
||||
|
||||
// 0. Resolve auth service for middleware
|
||||
var authSvc contracts.AuthService
|
||||
if err := core.Using[contracts.AuthService](ctx, func(svc contracts.AuthService) { authSvc = svc }); err != nil {
|
||||
return err
|
||||
|
||||
@@ -11,7 +11,7 @@ import (
|
||||
|
||||
"Wavelet/pkg/idgen"
|
||||
"Wavelet/pkg/util"
|
||||
database "Wavelet/plugins/infra/database"
|
||||
"Wavelet/plugins/domain/upload/shared"
|
||||
)
|
||||
|
||||
// UploadListFilter filters paginated upload queries.
|
||||
@@ -28,7 +28,7 @@ type UploadListFilter struct {
|
||||
|
||||
// ListUploads returns paginated upload records matching the filter.
|
||||
func ListUploads(ctx context.Context, filter UploadListFilter) (int64, []Upload, error) {
|
||||
query := database.DB(ctx).Model(&Upload{}).
|
||||
query := shared.GetDB(ctx).Model(&Upload{}).
|
||||
Where("status != ?", UploadStatusDeleted)
|
||||
|
||||
if filter.UserID != 0 {
|
||||
@@ -60,7 +60,7 @@ func ListUploads(ctx context.Context, filter UploadListFilter) (int64, []Upload,
|
||||
// GetActiveUploadByID loads a non-deleted upload by ID.
|
||||
func GetActiveUploadByID(ctx context.Context, id uint64) (Upload, error) {
|
||||
var upload Upload
|
||||
if err := database.DB(ctx).Where("id = ? AND status != ?", id, UploadStatusDeleted).First(&upload).Error; err != nil {
|
||||
if err := shared.GetDB(ctx).Where("id = ? AND status != ?", id, UploadStatusDeleted).First(&upload).Error; err != nil {
|
||||
return Upload{}, err
|
||||
}
|
||||
return upload, nil
|
||||
@@ -69,7 +69,7 @@ func GetActiveUploadByID(ctx context.Context, id uint64) (Upload, error) {
|
||||
// SoftDeleteUpload marks an upload as deleted.
|
||||
// External modules must use upload.Remove or upload.RemoveOwned; only internal/apps/upload may call this.
|
||||
func SoftDeleteUpload(ctx context.Context, upload *Upload) error {
|
||||
return SoftDeleteUploadTx(database.DB(ctx), upload)
|
||||
return SoftDeleteUploadTx(shared.GetDB(ctx), upload)
|
||||
}
|
||||
|
||||
// SoftDeleteUploadTx marks an upload as deleted within an existing transaction.
|
||||
@@ -82,13 +82,13 @@ func UpdateUpload(ctx context.Context, upload *Upload, updates map[string]any) e
|
||||
if len(updates) == 0 {
|
||||
return nil
|
||||
}
|
||||
return database.DB(ctx).Model(upload).Updates(updates).Error
|
||||
return shared.GetDB(ctx).Model(upload).Updates(updates).Error
|
||||
}
|
||||
|
||||
// ListDistinctUploadTypes returns all distinct non-empty upload business types.
|
||||
func ListDistinctUploadTypes(ctx context.Context) ([]string, error) {
|
||||
var types []string
|
||||
if err := database.DB(ctx).Model(&Upload{}).
|
||||
if err := shared.GetDB(ctx).Model(&Upload{}).
|
||||
Where("type IS NOT NULL AND type != ''").
|
||||
Distinct().
|
||||
Pluck("type", &types).Error; err != nil {
|
||||
@@ -100,7 +100,7 @@ func ListDistinctUploadTypes(ctx context.Context) ([]string, error) {
|
||||
// FindReusableUploadByHash finds an existing upload with the same hash and size.
|
||||
func FindReusableUploadByHash(ctx context.Context, hash string, size int64) (Upload, error) {
|
||||
var existing Upload
|
||||
err := database.DB(ctx).
|
||||
err := shared.GetDB(ctx).
|
||||
Where("hash = ? AND file_size = ? AND status IN (?, ?)", hash, size, UploadStatusPending, UploadStatusUsed).
|
||||
First(&existing).Error
|
||||
return existing, err
|
||||
@@ -108,7 +108,7 @@ func FindReusableUploadByHash(ctx context.Context, hash string, size int64) (Upl
|
||||
|
||||
// CreateUpload persists a new upload record.
|
||||
func CreateUpload(ctx context.Context, upload *Upload) error {
|
||||
return CreateUploadTx(database.DB(ctx), upload)
|
||||
return CreateUploadTx(shared.GetDB(ctx), upload)
|
||||
}
|
||||
|
||||
// CreateUploadTx persists a new upload record within an existing transaction.
|
||||
@@ -122,7 +122,7 @@ func CreateUploadTx(tx *gorm.DB, upload *Upload) error {
|
||||
// ListUploadsByIDs returns active uploads matching the given IDs.
|
||||
func ListUploadsByIDs(ctx context.Context, ids []uint64) ([]Upload, error) {
|
||||
var uploads []Upload
|
||||
if err := database.DB(ctx).
|
||||
if err := shared.GetDB(ctx).
|
||||
Where("id IN ? AND status IN (?, ?)", ids, UploadStatusPending, UploadStatusUsed).
|
||||
Find(&uploads).Error; err != nil {
|
||||
return nil, err
|
||||
@@ -134,13 +134,13 @@ func ListUploadsByIDs(ctx context.Context, ids []uint64) ([]Upload, error) {
|
||||
//
|
||||
//nolint:revive
|
||||
func UploadQuery(ctx context.Context) *gorm.DB {
|
||||
return database.DB(ctx).Model(&Upload{})
|
||||
return shared.GetDB(ctx).Model(&Upload{})
|
||||
}
|
||||
|
||||
// ListUploadStats returns all upload statistics rows.
|
||||
func ListUploadStats(ctx context.Context) ([]UploadStat, error) {
|
||||
var stats []UploadStat
|
||||
if err := database.DB(ctx).Find(&stats).Error; err != nil {
|
||||
if err := shared.GetDB(ctx).Find(&stats).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return stats, nil
|
||||
|
||||
@@ -8,10 +8,11 @@ import (
|
||||
"context"
|
||||
"strings"
|
||||
|
||||
"gorm.io/gorm"
|
||||
|
||||
"Wavelet/pkg/util"
|
||||
"Wavelet/plugins/domain/upload/models"
|
||||
database "Wavelet/plugins/infra/database"
|
||||
"gorm.io/gorm"
|
||||
"Wavelet/plugins/domain/upload/shared"
|
||||
)
|
||||
|
||||
// UploadListFilter filters paginated upload queries.
|
||||
@@ -26,7 +27,7 @@ type UploadListFilter struct {
|
||||
|
||||
// ListUploads returns paginated upload records matching the filter.
|
||||
func ListUploads(ctx context.Context, filter UploadListFilter) (int64, []models.Upload, error) {
|
||||
query := database.DB(ctx).Model(&models.Upload{}).
|
||||
query := shared.GetDB(ctx).Model(&models.Upload{}).
|
||||
Where("status != ?", models.UploadStatusDeleted)
|
||||
|
||||
if filter.UserID != 0 {
|
||||
@@ -58,7 +59,7 @@ func ListUploads(ctx context.Context, filter UploadListFilter) (int64, []models.
|
||||
// GetActiveUploadByID loads a non-deleted upload by ID.
|
||||
func GetActiveUploadByID(ctx context.Context, id uint64) (models.Upload, error) {
|
||||
var upload models.Upload
|
||||
if err := database.DB(ctx).Where("id = ? AND status != ?", id, models.UploadStatusDeleted).First(&upload).Error; err != nil {
|
||||
if err := shared.GetDB(ctx).Where("id = ? AND status != ?", id, models.UploadStatusDeleted).First(&upload).Error; err != nil {
|
||||
return models.Upload{}, err
|
||||
}
|
||||
return upload, nil
|
||||
@@ -66,7 +67,7 @@ func GetActiveUploadByID(ctx context.Context, id uint64) (models.Upload, error)
|
||||
|
||||
// SoftDeleteUpload marks an upload as deleted.
|
||||
func SoftDeleteUpload(ctx context.Context, upload *models.Upload) error {
|
||||
return SoftDeleteUploadTx(database.DB(ctx), upload)
|
||||
return SoftDeleteUploadTx(shared.GetDB(ctx), upload)
|
||||
}
|
||||
|
||||
// SoftDeleteUploadTx marks an upload as deleted within an existing transaction.
|
||||
@@ -79,13 +80,13 @@ func UpdateUpload(ctx context.Context, upload *models.Upload, updates map[string
|
||||
if len(updates) == 0 {
|
||||
return nil
|
||||
}
|
||||
return database.DB(ctx).Model(upload).Updates(updates).Error
|
||||
return shared.GetDB(ctx).Model(upload).Updates(updates).Error
|
||||
}
|
||||
|
||||
// ListDistinctUploadTypes returns all distinct non-empty upload business types.
|
||||
func ListDistinctUploadTypes(ctx context.Context) ([]string, error) {
|
||||
var types []string
|
||||
if err := database.DB(ctx).Model(&models.Upload{}).
|
||||
if err := shared.GetDB(ctx).Model(&models.Upload{}).
|
||||
Where("type IS NOT NULL AND type != ''").
|
||||
Distinct().
|
||||
Pluck("type", &types).Error; err != nil {
|
||||
@@ -97,7 +98,7 @@ func ListDistinctUploadTypes(ctx context.Context) ([]string, error) {
|
||||
// FindReusableUploadByHash finds an existing upload with the same hash and size.
|
||||
func FindReusableUploadByHash(ctx context.Context, hash string, size int64) (models.Upload, error) {
|
||||
var existing models.Upload
|
||||
err := database.DB(ctx).
|
||||
err := shared.GetDB(ctx).
|
||||
Where("hash = ? AND file_size = ? AND status IN (?, ?)", hash, size, models.UploadStatusPending, models.UploadStatusUsed).
|
||||
First(&existing).Error
|
||||
return existing, err
|
||||
@@ -105,7 +106,7 @@ func FindReusableUploadByHash(ctx context.Context, hash string, size int64) (mod
|
||||
|
||||
// CreateUpload persists a new upload record.
|
||||
func CreateUpload(ctx context.Context, upload *models.Upload) error {
|
||||
return CreateUploadTx(database.DB(ctx), upload)
|
||||
return CreateUploadTx(shared.GetDB(ctx), upload)
|
||||
}
|
||||
|
||||
// CreateUploadTx persists a new upload record within an existing transaction.
|
||||
@@ -116,7 +117,7 @@ func CreateUploadTx(tx *gorm.DB, upload *models.Upload) error {
|
||||
// ListUploadsByIDs returns active uploads matching the given IDs.
|
||||
func ListUploadsByIDs(ctx context.Context, ids []uint64) ([]models.Upload, error) {
|
||||
var uploads []models.Upload
|
||||
if err := database.DB(ctx).
|
||||
if err := shared.GetDB(ctx).
|
||||
Where("id IN ? AND status IN (?, ?)", ids, models.UploadStatusPending, models.UploadStatusUsed).
|
||||
Find(&uploads).Error; err != nil {
|
||||
return nil, err
|
||||
@@ -126,13 +127,13 @@ func ListUploadsByIDs(ctx context.Context, ids []uint64) ([]models.Upload, error
|
||||
|
||||
// UploadQuery returns a scoped GORM query for uploads.
|
||||
func UploadQuery(ctx context.Context) *gorm.DB {
|
||||
return database.DB(ctx).Model(&models.Upload{})
|
||||
return shared.GetDB(ctx).Model(&models.Upload{})
|
||||
}
|
||||
|
||||
// ListUploadStats returns all upload statistics rows.
|
||||
func ListUploadStats(ctx context.Context) ([]models.UploadStat, error) {
|
||||
var stats []models.UploadStat
|
||||
if err := database.DB(ctx).Find(&stats).Error; err != nil {
|
||||
if err := shared.GetDB(ctx).Find(&stats).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return stats, nil
|
||||
|
||||
@@ -0,0 +1,131 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package shared
|
||||
|
||||
import (
|
||||
"context"
|
||||
"sync"
|
||||
|
||||
"gorm.io/gorm"
|
||||
|
||||
"Wavelet/core"
|
||||
"Wavelet/core/contracts"
|
||||
)
|
||||
|
||||
var (
|
||||
svcMu sync.RWMutex
|
||||
dbSvc contracts.DBService
|
||||
cacheSvc contracts.CacheService
|
||||
storageSvc contracts.StorageService
|
||||
taskSvc contracts.TaskService
|
||||
authSvc contracts.AuthService
|
||||
)
|
||||
|
||||
// SetDBService configures the DBService.
|
||||
func SetDBService(s contracts.DBService) {
|
||||
svcMu.Lock()
|
||||
defer svcMu.Unlock()
|
||||
dbSvc = s
|
||||
}
|
||||
|
||||
// SetCacheService configures the CacheService.
|
||||
func SetCacheService(s contracts.CacheService) {
|
||||
svcMu.Lock()
|
||||
defer svcMu.Unlock()
|
||||
cacheSvc = s
|
||||
}
|
||||
|
||||
// SetStorageService configures the StorageService.
|
||||
func SetStorageService(s contracts.StorageService) {
|
||||
svcMu.Lock()
|
||||
defer svcMu.Unlock()
|
||||
storageSvc = s
|
||||
}
|
||||
|
||||
// SetTaskService configures the TaskService.
|
||||
func SetTaskService(s contracts.TaskService) {
|
||||
svcMu.Lock()
|
||||
defer svcMu.Unlock()
|
||||
taskSvc = s
|
||||
}
|
||||
|
||||
// SetAuthService configures the AuthService.
|
||||
func SetAuthService(s contracts.AuthService) {
|
||||
svcMu.Lock()
|
||||
defer svcMu.Unlock()
|
||||
authSvc = s
|
||||
}
|
||||
|
||||
// ResetServices clears all injected services.
|
||||
func ResetServices() {
|
||||
svcMu.Lock()
|
||||
defer svcMu.Unlock()
|
||||
dbSvc = nil
|
||||
cacheSvc = nil
|
||||
storageSvc = nil
|
||||
taskSvc = nil
|
||||
authSvc = nil
|
||||
}
|
||||
|
||||
// GetDB resolves the GORM DB instance.
|
||||
func GetDB(ctx context.Context) *gorm.DB {
|
||||
if c, ok := ctx.(*core.Context); ok && c != nil {
|
||||
if s, err := core.Inject[contracts.DBService](c); err == nil && s != nil {
|
||||
return s.DB(ctx)
|
||||
}
|
||||
}
|
||||
svcMu.RLock()
|
||||
s := dbSvc
|
||||
svcMu.RUnlock()
|
||||
if s != nil {
|
||||
return s.DB(ctx)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// GetCache resolves the CacheService instance.
|
||||
func GetCache(ctx context.Context) contracts.CacheService {
|
||||
if c, ok := ctx.(*core.Context); ok && c != nil {
|
||||
if s, err := core.Inject[contracts.CacheService](c); err == nil && s != nil {
|
||||
return s
|
||||
}
|
||||
}
|
||||
svcMu.RLock()
|
||||
s := cacheSvc
|
||||
svcMu.RUnlock()
|
||||
return s
|
||||
}
|
||||
|
||||
// GetStorage resolves the StorageService instance.
|
||||
func GetStorage(ctx context.Context) contracts.StorageService {
|
||||
if c, ok := ctx.(*core.Context); ok && c != nil {
|
||||
if s, err := core.Inject[contracts.StorageService](c); err == nil && s != nil {
|
||||
return s
|
||||
}
|
||||
}
|
||||
svcMu.RLock()
|
||||
s := storageSvc
|
||||
svcMu.RUnlock()
|
||||
return s
|
||||
}
|
||||
|
||||
// GetTaskService resolves the TaskService instance.
|
||||
func GetTaskService() contracts.TaskService {
|
||||
svcMu.RLock()
|
||||
defer svcMu.RUnlock()
|
||||
return taskSvc
|
||||
}
|
||||
|
||||
// GetAuthService resolves the AuthService instance.
|
||||
func GetAuthService(ctx context.Context) contracts.AuthService {
|
||||
if c, ok := ctx.(*core.Context); ok && c != nil {
|
||||
if s, err := core.Inject[contracts.AuthService](c); err == nil && s != nil {
|
||||
return s
|
||||
}
|
||||
}
|
||||
svcMu.RLock()
|
||||
s := authSvc
|
||||
svcMu.RUnlock()
|
||||
return s
|
||||
}
|
||||
@@ -0,0 +1,304 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package shared
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"crypto/sha256"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"gorm.io/gorm"
|
||||
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/pkg/testhelper"
|
||||
)
|
||||
|
||||
// MockDBService is a mock implementation of contracts.DBService for unit testing.
|
||||
type MockDBService struct {
|
||||
DBInstance *gorm.DB
|
||||
}
|
||||
|
||||
// GORM returns the underlying GORM instance.
|
||||
func (m *MockDBService) GORM() *gorm.DB {
|
||||
return m.DBInstance
|
||||
}
|
||||
|
||||
// DB returns the GORM instance bound to context.
|
||||
func (m *MockDBService) DB(ctx context.Context) *gorm.DB {
|
||||
return m.DBInstance.WithContext(ctx)
|
||||
}
|
||||
|
||||
// Named returns the named GORM instance.
|
||||
func (m *MockDBService) Named(_ string) *gorm.DB {
|
||||
return m.DBInstance
|
||||
}
|
||||
|
||||
// MockCacheService is an in-memory mock implementation of contracts.CacheService for unit testing.
|
||||
type MockCacheService struct {
|
||||
mu sync.RWMutex
|
||||
data map[string][]byte
|
||||
}
|
||||
|
||||
// NewMockCacheService creates a new MockCacheService.
|
||||
func NewMockCacheService() *MockCacheService {
|
||||
return &MockCacheService{
|
||||
data: make(map[string][]byte),
|
||||
}
|
||||
}
|
||||
|
||||
// Get retrieves a cached value.
|
||||
func (m *MockCacheService) Get(_ context.Context, key string, val any) error {
|
||||
m.mu.RLock()
|
||||
defer m.mu.RUnlock()
|
||||
b, ok := m.data[key]
|
||||
if !ok {
|
||||
return contracts.ErrCacheMiss
|
||||
}
|
||||
return json.Unmarshal(b, val)
|
||||
}
|
||||
|
||||
// Set stores a key-value pair in cache.
|
||||
func (m *MockCacheService) Set(_ context.Context, key string, val any, _ time.Duration) error {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
b, err := json.Marshal(val)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
m.data[key] = b
|
||||
return nil
|
||||
}
|
||||
|
||||
// Delete removes a key from cache.
|
||||
func (m *MockCacheService) Delete(_ context.Context, key string) error {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
delete(m.data, key)
|
||||
return nil
|
||||
}
|
||||
|
||||
// GetOrSet retrieves or populates a cache entry.
|
||||
func (m *MockCacheService) GetOrSet(ctx context.Context, key string, target any, ttl time.Duration, loader func() (any, error)) error {
|
||||
err := m.Get(ctx, key, target)
|
||||
if err == nil {
|
||||
return nil
|
||||
}
|
||||
val, err := loader()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return m.Set(ctx, key, val, ttl)
|
||||
}
|
||||
|
||||
// Invalidate invalidates a cache tag or prefix.
|
||||
func (m *MockCacheService) Invalidate(_ context.Context, _ string) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// MockStorageService is an in-memory mock implementation of contracts.StorageService for unit testing.
|
||||
type MockStorageService struct {
|
||||
mu sync.RWMutex
|
||||
objects map[string][]byte
|
||||
}
|
||||
|
||||
// NewMockStorageService creates a new MockStorageService.
|
||||
func NewMockStorageService() *MockStorageService {
|
||||
return &MockStorageService{
|
||||
objects: make(map[string][]byte),
|
||||
}
|
||||
}
|
||||
|
||||
// Put uploads an object into mock storage.
|
||||
func (m *MockStorageService) Put(_ context.Context, key string, body io.Reader, size int64, contentType 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
|
||||
if strings.HasPrefix(key, "uploads/") {
|
||||
_ = os.MkdirAll(filepath.Dir(key), 0755)
|
||||
_ = os.WriteFile(key, data, 0644)
|
||||
}
|
||||
return contracts.StoragePutResult{Key: key, Bucket: "test-bucket"}, nil
|
||||
}
|
||||
|
||||
// Get retrieves an object from mock storage.
|
||||
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
|
||||
}
|
||||
if f, err := os.Open(key); err == nil {
|
||||
info, _ := f.Stat()
|
||||
return &contracts.StorageObject{
|
||||
Key: key,
|
||||
Body: f,
|
||||
ContentLength: info.Size(),
|
||||
ContentType: "application/octet-stream",
|
||||
}, nil
|
||||
}
|
||||
return nil, gorm.ErrRecordNotFound
|
||||
}
|
||||
|
||||
// Delete removes an object from mock storage.
|
||||
func (m *MockStorageService) Delete(_ context.Context, key string) error {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
delete(m.objects, key)
|
||||
return nil
|
||||
}
|
||||
|
||||
// Ingest handles programmatic file ingestion for mock storage.
|
||||
func (m *MockStorageService) Ingest(_ context.Context, _ io.Reader, _ contracts.IngestOptions) (*contracts.IngestResult, error) {
|
||||
return &contracts.IngestResult{ID: 1, Key: "test.png", Created: true, Stored: true}, nil
|
||||
}
|
||||
|
||||
// MockAuthService is a mock implementation of contracts.AuthService for unit testing.
|
||||
type MockAuthService struct {
|
||||
DB *gorm.DB
|
||||
}
|
||||
|
||||
// RequireAuthMiddleware returns a dummy auth middleware.
|
||||
func (a *MockAuthService) RequireAuthMiddleware() any {
|
||||
return func(c *gin.Context) { c.Next() }
|
||||
}
|
||||
|
||||
// RequireAdminMiddleware returns a dummy admin middleware.
|
||||
func (a *MockAuthService) RequireAdminMiddleware() any {
|
||||
return func(c *gin.Context) { c.Next() }
|
||||
}
|
||||
|
||||
// DisallowTokenAuthMiddleware returns a dummy disallow token middleware.
|
||||
func (a *MockAuthService) DisallowTokenAuthMiddleware() any {
|
||||
return func(c *gin.Context) { c.Next() }
|
||||
}
|
||||
|
||||
// GetCurrentUser returns the user associated with the request context.
|
||||
func (a *MockAuthService) GetCurrentUser(ctx context.Context) (*contracts.UserDTO, error) {
|
||||
if c, ok := ctx.(*gin.Context); ok {
|
||||
authHeader := c.GetHeader("Authorization")
|
||||
if strings.HasPrefix(authHeader, "Bearer ") {
|
||||
tokenStr := strings.TrimPrefix(authHeader, "Bearer ")
|
||||
tokenHash := fmt.Sprintf("%x", sha256.Sum256([]byte(tokenStr)))
|
||||
var tokenRecord struct {
|
||||
UserID uint64
|
||||
}
|
||||
if err := a.DB.Table("w_access_tokens").Where("token_hash = ?", tokenHash).First(&tokenRecord).Error; err == nil {
|
||||
return &contracts.UserDTO{ID: tokenRecord.UserID, IsActive: true}, nil
|
||||
}
|
||||
}
|
||||
}
|
||||
return nil, errors.New("unauthorized")
|
||||
}
|
||||
|
||||
// GetCurrentUserID returns the current user ID.
|
||||
func (a *MockAuthService) GetCurrentUserID(ctx context.Context) (uint64, error) {
|
||||
u, err := a.GetCurrentUser(ctx)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return u.ID, nil
|
||||
}
|
||||
|
||||
// VerifyToken verifies an access token.
|
||||
func (a *MockAuthService) VerifyToken(_ context.Context, token string) (*contracts.UserDTO, error) {
|
||||
tokenHash := fmt.Sprintf("%x", sha256.Sum256([]byte(token)))
|
||||
var tokenRecord struct {
|
||||
UserID uint64
|
||||
}
|
||||
if err := a.DB.Table("w_access_tokens").Where("token_hash = ?", tokenHash).First(&tokenRecord).Error; err == nil {
|
||||
return &contracts.UserDTO{ID: tokenRecord.UserID, IsActive: true}, nil
|
||||
}
|
||||
return nil, errors.New("unauthorized")
|
||||
}
|
||||
|
||||
// Authenticate verifies credentials.
|
||||
func (a *MockAuthService) Authenticate(_ context.Context, _ string, _ string) (*contracts.UserDTO, error) {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
// CreateSession creates a login session.
|
||||
func (a *MockAuthService) CreateSession(_ context.Context, _ uint64, _ map[string]any) (string, error) {
|
||||
return "test-session", nil
|
||||
}
|
||||
|
||||
// RevokeToken revokes an access token.
|
||||
func (a *MockAuthService) RevokeToken(_ context.Context, _ string) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// RevokeUserSessions revokes all sessions for a user.
|
||||
func (a *MockAuthService) RevokeUserSessions(_ context.Context, _ uint64) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// InvalidateCachedUser invalidates cached user profile.
|
||||
func (a *MockAuthService) InvalidateCachedUser(_ context.Context, _ uint64) {}
|
||||
|
||||
// InvalidateCachedToken invalidates cached access token.
|
||||
func (a *MockAuthService) InvalidateCachedToken(_ context.Context, _ string) {}
|
||||
|
||||
// ListAuthSources lists configured authentication sources.
|
||||
func (a *MockAuthService) ListAuthSources(_ context.Context) ([]contracts.AuthSourceViewDTO, error) {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
// CreateAuthSource creates an authentication source.
|
||||
func (a *MockAuthService) CreateAuthSource(_ context.Context, _ contracts.AuthSourceDTO) (*contracts.AuthSourceDTO, error) {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
// UpdateAuthSource updates an authentication source.
|
||||
func (a *MockAuthService) UpdateAuthSource(_ context.Context, _ uint64, _ contracts.AuthSourceDTO) (*contracts.AuthSourceDTO, error) {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
// DeleteAuthSource deletes an authentication source.
|
||||
func (a *MockAuthService) DeleteAuthSource(_ context.Context, _ uint64) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// ToggleAuthSource toggles an authentication source active state.
|
||||
func (a *MockAuthService) ToggleAuthSource(_ context.Context, _ uint64) (*contracts.AuthSourceDTO, error) {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
// SetupTestEnv initializes test helper environment and binds DB, Cache, Storage, Auth mocks to shared services.
|
||||
func SetupTestEnv(t *testing.T) (*gorm.DB, func()) {
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
dbSvc := &MockDBService{DBInstance: dbConn}
|
||||
cacheSvc := NewMockCacheService()
|
||||
storageSvc := NewMockStorageService()
|
||||
authSvc := &MockAuthService{DB: dbConn}
|
||||
|
||||
SetDBService(dbSvc)
|
||||
SetCacheService(cacheSvc)
|
||||
SetStorageService(storageSvc)
|
||||
SetAuthService(authSvc)
|
||||
|
||||
return dbConn, func() {
|
||||
ResetServices()
|
||||
cleanup()
|
||||
}
|
||||
}
|
||||
@@ -7,11 +7,12 @@ import (
|
||||
"context"
|
||||
"time"
|
||||
|
||||
"Wavelet/pkg/logger"
|
||||
"Wavelet/plugins/domain/upload/models"
|
||||
database "Wavelet/plugins/infra/database"
|
||||
"gorm.io/gorm"
|
||||
"gorm.io/gorm/clause"
|
||||
|
||||
"Wavelet/pkg/logger"
|
||||
"Wavelet/plugins/domain/upload/models"
|
||||
"Wavelet/plugins/domain/upload/shared"
|
||||
)
|
||||
|
||||
// ApplyUploadStatsAdd increments incremental stats for a newly active upload record.
|
||||
@@ -26,7 +27,7 @@ func ApplyUploadStatsRemove(ctx context.Context, upload *models.Upload) error {
|
||||
|
||||
// RebuildUploadStats rebuilds all incremental stats from current upload records.
|
||||
func RebuildUploadStats(ctx context.Context) error {
|
||||
return database.DB(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
return shared.GetDB(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
if err := tx.Where("1 = 1").Delete(&models.UploadStat{}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
@@ -49,7 +50,7 @@ func applyUploadStatsDelta(ctx context.Context, upload *models.Upload, sign int6
|
||||
if upload == nil || !isActiveUploadStatus(upload.Status) {
|
||||
return nil
|
||||
}
|
||||
return database.DB(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
return shared.GetDB(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
return ApplyUploadStatsDeltaTx(tx, upload, sign)
|
||||
})
|
||||
}
|
||||
|
||||
@@ -8,15 +8,36 @@ import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"gorm.io/gorm"
|
||||
|
||||
"Wavelet/pkg/testhelper"
|
||||
"Wavelet/plugins/domain/upload/models"
|
||||
database "Wavelet/plugins/infra/database"
|
||||
"gorm.io/gorm"
|
||||
"Wavelet/plugins/domain/upload/shared"
|
||||
)
|
||||
|
||||
type mockDBService struct {
|
||||
db *gorm.DB
|
||||
}
|
||||
|
||||
func (m *mockDBService) GORM() *gorm.DB {
|
||||
return m.db
|
||||
}
|
||||
|
||||
func (m *mockDBService) DB(ctx context.Context) *gorm.DB {
|
||||
return m.db.WithContext(ctx)
|
||||
}
|
||||
|
||||
func (m *mockDBService) Named(_ string) *gorm.DB {
|
||||
return m.db
|
||||
}
|
||||
|
||||
func TestApplyUploadStatsDeltaTxWithinTransaction(t *testing.T) {
|
||||
_, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
defer cleanup()
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
shared.SetDBService(&mockDBService{db: dbConn})
|
||||
defer func() {
|
||||
shared.SetDBService(nil)
|
||||
cleanup()
|
||||
}()
|
||||
ctx := context.Background()
|
||||
|
||||
upload := &models.Upload{
|
||||
@@ -29,7 +50,7 @@ func TestApplyUploadStatsDeltaTxWithinTransaction(t *testing.T) {
|
||||
CreatedAt: time.Now(),
|
||||
}
|
||||
|
||||
if err := database.DB(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
if err := shared.GetDB(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
return ApplyUploadStatsDeltaTx(tx, upload, 1)
|
||||
}); err != nil {
|
||||
t.Fatalf("ApplyUploadStatsDeltaTx returned error: %v", err)
|
||||
@@ -45,8 +66,12 @@ func TestApplyUploadStatsDeltaTxWithinTransaction(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestApplyUploadStatsAddAndRemove(t *testing.T) {
|
||||
_, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
defer cleanup()
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
shared.SetDBService(&mockDBService{db: dbConn})
|
||||
defer func() {
|
||||
shared.SetDBService(nil)
|
||||
cleanup()
|
||||
}()
|
||||
ctx := context.Background()
|
||||
|
||||
upload := &models.Upload{
|
||||
@@ -90,7 +115,7 @@ type uploadStatsSnapshot struct {
|
||||
|
||||
func loadUploadStats(ctx context.Context) (uploadStatsSnapshot, error) {
|
||||
var rows []models.UploadStat
|
||||
if err := database.DB(ctx).Where("dimension = ?", models.UploadStatDimensionTotal).Find(&rows).Error; err != nil {
|
||||
if err := shared.GetDB(ctx).Where("dimension = ?", models.UploadStatDimensionTotal).Find(&rows).Error; err != nil {
|
||||
return uploadStatsSnapshot{}, err
|
||||
}
|
||||
if len(rows) == 0 {
|
||||
|
||||
@@ -9,15 +9,14 @@ import (
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/plugins/domain/upload/shared"
|
||||
"Wavelet/plugins/drivers/driver_asynq_worker"
|
||||
"Wavelet/plugins/infra/storage/objectstore"
|
||||
)
|
||||
|
||||
// MigrationAccessState captures cached migration maintenance state.
|
||||
type MigrationAccessState struct {
|
||||
ReadOnly bool
|
||||
Target objectstore.Config
|
||||
Target contracts.StorageConfigDTO
|
||||
HasTarget bool
|
||||
TargetErr error
|
||||
LoadErr error
|
||||
@@ -65,14 +64,14 @@ func buildMigrationAccessState(ctx context.Context) MigrationAccessState {
|
||||
if err != nil {
|
||||
return MigrationAccessState{LoadErr: err, ReadOnly: true}
|
||||
}
|
||||
if !ok {
|
||||
if !ok || execution == nil {
|
||||
return MigrationAccessState{}
|
||||
}
|
||||
|
||||
state := MigrationAccessState{
|
||||
ReadOnly: execution.Status != driver_asynq_worker.TaskExecutionStatusSucceeded,
|
||||
ReadOnly: execution.Status != "succeeded",
|
||||
}
|
||||
if execution.Status == driver_asynq_worker.TaskExecutionStatusSucceeded {
|
||||
if execution.Status == "succeeded" {
|
||||
return state
|
||||
}
|
||||
|
||||
|
||||
@@ -10,33 +10,47 @@ import (
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"Wavelet/plugins/drivers/driver_asynq_worker"
|
||||
"Wavelet/plugins/infra/storage/objectstore"
|
||||
"gorm.io/gorm"
|
||||
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/plugins/domain/upload/shared"
|
||||
)
|
||||
|
||||
// StorageMigrationTask is the Asynq task name for storage migration.
|
||||
// StorageMigrationTask is the task name for storage migration.
|
||||
const StorageMigrationTask = "storage:migrate"
|
||||
|
||||
// LatestMigrationExecution returns the most recent storage migration task execution.
|
||||
func LatestMigrationExecution(ctx context.Context) (*driver_asynq_worker.TaskExecution, bool, error) {
|
||||
return driver_asynq_worker.GetLatestTaskExecutionByTaskType(ctx, StorageMigrationTask)
|
||||
func LatestMigrationExecution(ctx context.Context) (*contracts.TaskExecutionDTO, bool, error) {
|
||||
db := shared.GetDB(ctx)
|
||||
if db == nil {
|
||||
return nil, false, nil
|
||||
}
|
||||
var exec contracts.TaskExecutionDTO
|
||||
err := db.Table("w_task_executions").Where("task_type = ?", StorageMigrationTask).Order("id DESC").First(&exec).Error
|
||||
if err != nil {
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return nil, false, nil
|
||||
}
|
||||
return nil, false, err
|
||||
}
|
||||
return &exec, true, nil
|
||||
}
|
||||
|
||||
// ParseMigrationTargetConfig parses and validates a storage migration target payload.
|
||||
func ParseMigrationTargetConfig(ctx context.Context, payload []byte) (objectstore.Config, error) {
|
||||
func ParseMigrationTargetConfig(ctx context.Context, payload []byte) (contracts.StorageConfigDTO, error) {
|
||||
if strings.TrimSpace(string(payload)) == "" {
|
||||
return objectstore.Config{}, errors.New("storage migration target payload is required")
|
||||
return contracts.StorageConfigDTO{}, errors.New("storage migration target payload is required")
|
||||
}
|
||||
|
||||
var raw struct {
|
||||
Target json.RawMessage `json:"target"`
|
||||
}
|
||||
if err := json.Unmarshal(payload, &raw); err != nil {
|
||||
return objectstore.Config{}, fmt.Errorf("parse storage migration payload envelope: %w", err)
|
||||
return contracts.StorageConfigDTO{}, fmt.Errorf("parse storage migration payload envelope: %w", err)
|
||||
}
|
||||
|
||||
if len(raw.Target) == 0 {
|
||||
return objectstore.Config{}, errors.New("storage migration target payload is required")
|
||||
return contracts.StorageConfigDTO{}, errors.New("storage migration target payload is required")
|
||||
}
|
||||
|
||||
var targetBytes []byte
|
||||
@@ -47,34 +61,57 @@ func ParseMigrationTargetConfig(ctx context.Context, payload []byte) (objectstor
|
||||
targetBytes = raw.Target
|
||||
}
|
||||
|
||||
var target objectstore.Config
|
||||
var target contracts.StorageConfigDTO
|
||||
if err := json.Unmarshal(targetBytes, &target); err != nil {
|
||||
return objectstore.Config{}, fmt.Errorf("parse target storage config: %w", err)
|
||||
return contracts.StorageConfigDTO{}, fmt.Errorf("parse target storage config: %w", err)
|
||||
}
|
||||
|
||||
current, err := objectstore.LoadConfig(ctx)
|
||||
if err != nil {
|
||||
return objectstore.Config{}, fmt.Errorf("load active storage config: %w", err)
|
||||
}
|
||||
target = objectstore.MergeMaskedSecrets(target, current)
|
||||
if err := objectstore.ValidateConfig(target); err != nil {
|
||||
return objectstore.Config{}, fmt.Errorf("validate target storage config: %w", err)
|
||||
}
|
||||
return target, nil
|
||||
}
|
||||
|
||||
// NormalizeMigrationPayload validates and normalizes a storage migration payload.
|
||||
func NormalizeMigrationPayload(ctx context.Context, payload []byte) ([]byte, objectstore.Config, error) {
|
||||
func NormalizeMigrationPayload(ctx context.Context, payload []byte) ([]byte, contracts.StorageConfigDTO, error) {
|
||||
target, err := ParseMigrationTargetConfig(ctx, payload)
|
||||
if err != nil {
|
||||
return nil, objectstore.Config{}, err
|
||||
return nil, contracts.StorageConfigDTO{}, err
|
||||
}
|
||||
type storageMigrationPayload struct {
|
||||
Target objectstore.Config `json:"target"`
|
||||
}
|
||||
normalized, err := json.Marshal(storageMigrationPayload{Target: target})
|
||||
raw, err := json.Marshal(struct {
|
||||
Target contracts.StorageConfigDTO `json:"target"`
|
||||
}{Target: target})
|
||||
if err != nil {
|
||||
return nil, objectstore.Config{}, fmt.Errorf("marshal storage migration payload: %w", err)
|
||||
return nil, contracts.StorageConfigDTO{}, fmt.Errorf("serialize normalized payload: %w", err)
|
||||
}
|
||||
return normalized, target, nil
|
||||
return raw, target, nil
|
||||
}
|
||||
|
||||
// SaveActiveConfig persists the active storage configuration to w_system_configs.
|
||||
func SaveActiveConfig(ctx context.Context, cfg contracts.StorageConfigDTO) error {
|
||||
db := shared.GetDB(ctx)
|
||||
if db == nil {
|
||||
return errors.New("database not available")
|
||||
}
|
||||
data, err := json.Marshal(cfg)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return db.Table("w_system_configs").Where("key = ?", "storage_config").Update("value", string(data)).Error
|
||||
}
|
||||
|
||||
// LoadStorageConfig loads the current storage configuration from w_system_configs.
|
||||
func LoadStorageConfig(ctx context.Context) (contracts.StorageConfigDTO, error) {
|
||||
db := shared.GetDB(ctx)
|
||||
if db == nil {
|
||||
return contracts.StorageConfigDTO{}, errors.New("database not available")
|
||||
}
|
||||
var row struct {
|
||||
Value string
|
||||
}
|
||||
if err := db.Table("w_system_configs").Where("key = ?", "storage_config").First(&row).Error; err != nil {
|
||||
return contracts.StorageConfigDTO{}, err
|
||||
}
|
||||
var cfg contracts.StorageConfigDTO
|
||||
if err := json.Unmarshal([]byte(row.Value), &cfg); err != nil {
|
||||
return contracts.StorageConfigDTO{}, err
|
||||
}
|
||||
return cfg, nil
|
||||
}
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user