mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-06 23:56:37 +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/
|
.worktrees/
|
||||||
/.superpowers/
|
/.superpowers/
|
||||||
|
/backend/plugins/domain/upload/filesrv/uploads/
|
||||||
|
/backend/plugins/domain/upload/task/uploads/
|
||||||
|
|||||||
+12
-5
@@ -25,17 +25,17 @@ linters:
|
|||||||
- gocritic # 各类代码问题
|
- gocritic # 各类代码问题
|
||||||
- funlen # 函数过长
|
- funlen # 函数过长
|
||||||
|
|
||||||
- gosec # 安全问题检查
|
- gosec # 安全问题检查
|
||||||
- bodyclose # HTTP response body 没有正确关闭
|
- bodyclose # HTTP response body 没有正确关闭
|
||||||
- noctx # 没有传递 context.Context
|
- noctx # 没有传递 context.Context
|
||||||
- contextcheck # 其他检查
|
- contextcheck # 其他检查
|
||||||
- sqlclosecheck # SQL rows 没有正确关闭
|
- sqlclosecheck # SQL rows 没有正确关闭
|
||||||
- unconvert # 不必要的类型转换
|
- unconvert # 不必要的类型转换
|
||||||
- nilerr # 函数返回 nil 错误
|
- nilerr # 函数返回 nil 错误
|
||||||
|
|
||||||
settings:
|
settings:
|
||||||
dupl:
|
dupl:
|
||||||
threshold: 120
|
threshold: 80
|
||||||
|
|
||||||
cyclop:
|
cyclop:
|
||||||
max-complexity: 20
|
max-complexity: 20
|
||||||
@@ -53,3 +53,10 @@ linters:
|
|||||||
- argument
|
- argument
|
||||||
- condition
|
- condition
|
||||||
- return
|
- return
|
||||||
|
|
||||||
|
formatters:
|
||||||
|
enable:
|
||||||
|
- gofumpt
|
||||||
|
settings:
|
||||||
|
gofumpt:
|
||||||
|
extra-rules: true
|
||||||
|
|||||||
@@ -14,8 +14,9 @@ license-check:
|
|||||||
scripts/update_go_license.sh --check
|
scripts/update_go_license.sh --check
|
||||||
|
|
||||||
format:
|
format:
|
||||||
@echo "==> Formatting backend Go source..."
|
@echo "==> Formatting backend Go source with goimports..."
|
||||||
gofmt -w $$(find backend -type f -name '*.go' -not -path './.git/*')
|
@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..."
|
@echo "==> Formatting frontend source..."
|
||||||
cd frontend && pnpm format
|
cd frontend && pnpm format
|
||||||
|
|
||||||
|
|||||||
+2
-1
@@ -7,8 +7,9 @@ package cmd
|
|||||||
import (
|
import (
|
||||||
"log"
|
"log"
|
||||||
|
|
||||||
"Wavelet/core"
|
|
||||||
"github.com/spf13/cobra"
|
"github.com/spf13/cobra"
|
||||||
|
|
||||||
|
"Wavelet/core"
|
||||||
)
|
)
|
||||||
|
|
||||||
var allCmd = &cobra.Command{
|
var allCmd = &cobra.Command{
|
||||||
|
|||||||
+2
-1
@@ -6,8 +6,9 @@ package cmd
|
|||||||
import (
|
import (
|
||||||
"log"
|
"log"
|
||||||
|
|
||||||
"Wavelet/core"
|
|
||||||
"github.com/spf13/cobra"
|
"github.com/spf13/cobra"
|
||||||
|
|
||||||
|
"Wavelet/core"
|
||||||
)
|
)
|
||||||
|
|
||||||
var apiCmd = &cobra.Command{
|
var apiCmd = &cobra.Command{
|
||||||
|
|||||||
+3
-2
@@ -10,6 +10,9 @@ import (
|
|||||||
"log"
|
"log"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
|
"github.com/pressly/goose/v3"
|
||||||
|
goosedb "github.com/pressly/goose/v3/database"
|
||||||
|
|
||||||
"Wavelet/core"
|
"Wavelet/core"
|
||||||
"Wavelet/core/contracts"
|
"Wavelet/core/contracts"
|
||||||
"Wavelet/pkg/config"
|
"Wavelet/pkg/config"
|
||||||
@@ -28,8 +31,6 @@ import (
|
|||||||
infradb "Wavelet/plugins/infra/database"
|
infradb "Wavelet/plugins/infra/database"
|
||||||
"Wavelet/plugins/infra/logger"
|
"Wavelet/plugins/infra/logger"
|
||||||
"Wavelet/plugins/infra/storage"
|
"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.
|
// 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"
|
userdomain "Wavelet/plugins/domain/user"
|
||||||
"Wavelet/plugins/infra/database"
|
"Wavelet/plugins/infra/database"
|
||||||
|
|
||||||
"Wavelet/plugins/domain/auth"
|
|
||||||
"github.com/spf13/cobra"
|
"github.com/spf13/cobra"
|
||||||
"gorm.io/gorm"
|
"gorm.io/gorm"
|
||||||
|
|
||||||
|
"Wavelet/plugins/domain/auth"
|
||||||
)
|
)
|
||||||
|
|
||||||
var (
|
var (
|
||||||
@@ -49,7 +50,10 @@ var resetPasswdCmd = &cobra.Command{
|
|||||||
ctx := context.Background()
|
ctx := context.Background()
|
||||||
|
|
||||||
// Ensure database is initialized
|
// Ensure database is initialized
|
||||||
database.DB(ctx)
|
dbConn := database.DB(ctx)
|
||||||
|
if dbConn != nil {
|
||||||
|
userdomain.SetDBService(database.NewService(dbConn))
|
||||||
|
}
|
||||||
|
|
||||||
var username string
|
var username string
|
||||||
if usernameFlag != "" {
|
if usernameFlag != "" {
|
||||||
|
|||||||
+2
-1
@@ -8,11 +8,12 @@ import (
|
|||||||
"log"
|
"log"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
|
"github.com/spf13/cobra"
|
||||||
|
|
||||||
"Wavelet/pkg/buildinfo"
|
"Wavelet/pkg/buildinfo"
|
||||||
"Wavelet/pkg/config"
|
"Wavelet/pkg/config"
|
||||||
"Wavelet/pkg/logger"
|
"Wavelet/pkg/logger"
|
||||||
"Wavelet/pkg/trace"
|
"Wavelet/pkg/trace"
|
||||||
"github.com/spf13/cobra"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
const traceShutdownTimeout = 10 * time.Second
|
const traceShutdownTimeout = 10 * time.Second
|
||||||
|
|||||||
@@ -6,8 +6,9 @@ package cmd
|
|||||||
import (
|
import (
|
||||||
"log"
|
"log"
|
||||||
|
|
||||||
"Wavelet/core"
|
|
||||||
"github.com/spf13/cobra"
|
"github.com/spf13/cobra"
|
||||||
|
|
||||||
|
"Wavelet/core"
|
||||||
)
|
)
|
||||||
|
|
||||||
var schedulerCmd = &cobra.Command{
|
var schedulerCmd = &cobra.Command{
|
||||||
|
|||||||
@@ -6,8 +6,9 @@ package cmd
|
|||||||
import (
|
import (
|
||||||
"log"
|
"log"
|
||||||
|
|
||||||
"Wavelet/core"
|
|
||||||
"github.com/spf13/cobra"
|
"github.com/spf13/cobra"
|
||||||
|
|
||||||
|
"Wavelet/core"
|
||||||
)
|
)
|
||||||
|
|
||||||
var workerCmd = &cobra.Command{
|
var workerCmd = &cobra.Command{
|
||||||
|
|||||||
@@ -10,7 +10,6 @@ import (
|
|||||||
"sync"
|
"sync"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"Wavelet/core/contracts"
|
|
||||||
"Wavelet/core/extpoints"
|
"Wavelet/core/extpoints"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -237,24 +236,6 @@ func (c *Context) Setting() extpoints.SettingExtension {
|
|||||||
return c.Settings()
|
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.
|
// OnDispose registers a cleanup callback function to be executed when this Context is disposed.
|
||||||
// It accepts func() error, func(), or Disposer.
|
// It accepts func() error, func(), or Disposer.
|
||||||
func (c *Context) OnDispose(fn any) {
|
func (c *Context) OnDispose(fn any) {
|
||||||
|
|||||||
@@ -44,6 +44,25 @@ const (
|
|||||||
EventTopicSystemCleanup = "admin:system_cleanup"
|
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 ---
|
// --- Upload / Storage Events ---
|
||||||
const (
|
const (
|
||||||
// EventTopicUploadCreated fires when a new file upload is recorded.
|
// 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
|
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.
|
// StorageService defines the contract for unified object storage and managed file ingestion.
|
||||||
type StorageService interface {
|
type StorageService interface {
|
||||||
// Put writes an object to storage.
|
// 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 (
|
import (
|
||||||
"net/http"
|
"net/http"
|
||||||
|
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
|
||||||
"Wavelet/core"
|
"Wavelet/core"
|
||||||
"Wavelet/core/contracts"
|
"Wavelet/core/contracts"
|
||||||
"github.com/gin-gonic/gin"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
// Plugin implements core.Plugin for the custom_example downstream plugin.
|
// Plugin implements core.Plugin for the custom_example downstream plugin.
|
||||||
|
|||||||
Vendored
+14
@@ -88,6 +88,20 @@ func New(basePath string) *Cache {
|
|||||||
return c
|
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.
|
// Set stores a key-value pair in the cache.
|
||||||
// Use DefaultExpiration for the configured default TTL, NoExpiration for no
|
// Use DefaultExpiration for the configured default TTL, NoExpiration for no
|
||||||
// TTL, or a positive duration for a business-specific TTL.
|
// TTL, or a positive duration for a business-specific TTL.
|
||||||
|
|||||||
@@ -8,8 +8,9 @@ import (
|
|||||||
"fmt"
|
"fmt"
|
||||||
"log"
|
"log"
|
||||||
|
|
||||||
"Wavelet/pkg/config"
|
|
||||||
"github.com/bwmarrin/snowflake"
|
"github.com/bwmarrin/snowflake"
|
||||||
|
|
||||||
|
"Wavelet/pkg/config"
|
||||||
)
|
)
|
||||||
|
|
||||||
// 2025-12-01 00:00:00 UTC 的毫秒时间戳
|
// 2025-12-01 00:00:00 UTC 的毫秒时间戳
|
||||||
|
|||||||
@@ -4,8 +4,9 @@
|
|||||||
package testhelper
|
package testhelper
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"Wavelet/pkg/response"
|
|
||||||
"github.com/gin-gonic/gin"
|
"github.com/gin-gonic/gin"
|
||||||
|
|
||||||
|
"Wavelet/pkg/response"
|
||||||
)
|
)
|
||||||
|
|
||||||
// NewTestGinEngine 创建带 ErrorHandlerMiddleware 的 Gin 引擎,与生产环境错误响应行为一致。
|
// NewTestGinEngine 创建带 ErrorHandlerMiddleware 的 Gin 引擎,与生产环境错误响应行为一致。
|
||||||
|
|||||||
@@ -9,13 +9,14 @@ import (
|
|||||||
"testing"
|
"testing"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
cachepkg "Wavelet/plugins/infra/cache"
|
|
||||||
db "Wavelet/plugins/infra/database"
|
|
||||||
"github.com/alicebob/miniredis/v2"
|
"github.com/alicebob/miniredis/v2"
|
||||||
"github.com/glebarez/sqlite"
|
"github.com/glebarez/sqlite"
|
||||||
"github.com/redis/go-redis/v9"
|
"github.com/redis/go-redis/v9"
|
||||||
"github.com/redis/go-redis/v9/maintnotifications"
|
"github.com/redis/go-redis/v9/maintnotifications"
|
||||||
"gorm.io/gorm"
|
"gorm.io/gorm"
|
||||||
|
|
||||||
|
cachepkg "Wavelet/plugins/infra/cache"
|
||||||
|
db "Wavelet/plugins/infra/database"
|
||||||
)
|
)
|
||||||
|
|
||||||
// SystemConfig 测试用系统配置表
|
// 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.
|
// ListAuthSources lists all configured authentication sources.
|
||||||
func ListAuthSources(c *gin.Context) {
|
func ListAuthSources(c *gin.Context) {
|
||||||
authSvc := getAuthService(c.Request.Context())
|
authSvc := GetAuthService(c.Request.Context())
|
||||||
if authSvc == nil {
|
if authSvc == nil {
|
||||||
response.AbortInternal(c, "认证服务未就绪")
|
response.AbortInternal(c, "认证服务未就绪")
|
||||||
return
|
return
|
||||||
@@ -38,7 +38,7 @@ func CreateAuthSource(c *gin.Context) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
authSvc := getAuthService(c.Request.Context())
|
authSvc := GetAuthService(c.Request.Context())
|
||||||
if authSvc == nil {
|
if authSvc == nil {
|
||||||
response.AbortInternal(c, "认证服务未就绪")
|
response.AbortInternal(c, "认证服务未就绪")
|
||||||
return
|
return
|
||||||
@@ -68,7 +68,7 @@ func UpdateAuthSource(c *gin.Context) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
authSvc := getAuthService(c.Request.Context())
|
authSvc := GetAuthService(c.Request.Context())
|
||||||
if authSvc == nil {
|
if authSvc == nil {
|
||||||
response.AbortInternal(c, "认证服务未就绪")
|
response.AbortInternal(c, "认证服务未就绪")
|
||||||
return
|
return
|
||||||
@@ -92,7 +92,7 @@ func ToggleAuthSource(c *gin.Context) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
authSvc := getAuthService(c.Request.Context())
|
authSvc := GetAuthService(c.Request.Context())
|
||||||
if authSvc == nil {
|
if authSvc == nil {
|
||||||
response.AbortInternal(c, "认证服务未就绪")
|
response.AbortInternal(c, "认证服务未就绪")
|
||||||
return
|
return
|
||||||
@@ -116,7 +116,7 @@ func DeleteAuthSource(c *gin.Context) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
authSvc := getAuthService(c.Request.Context())
|
authSvc := GetAuthService(c.Request.Context())
|
||||||
if authSvc == nil {
|
if authSvc == nil {
|
||||||
response.AbortInternal(c, "认证服务未就绪")
|
response.AbortInternal(c, "认证服务未就绪")
|
||||||
return
|
return
|
||||||
|
|||||||
@@ -10,8 +10,8 @@ import (
|
|||||||
|
|
||||||
"github.com/gin-gonic/gin"
|
"github.com/gin-gonic/gin"
|
||||||
|
|
||||||
|
pkgcache "Wavelet/pkg/cache/disk"
|
||||||
"Wavelet/pkg/response"
|
"Wavelet/pkg/response"
|
||||||
"Wavelet/plugins/infra/storage/diskcache"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
type updateCacheConfigRequest struct {
|
type updateCacheConfigRequest struct {
|
||||||
@@ -26,13 +26,13 @@ type updateCacheConfigRequest struct {
|
|||||||
// @Tags admin
|
// @Tags admin
|
||||||
// @Produce json
|
// @Produce json
|
||||||
// @Security SessionCookie
|
// @Security SessionCookie
|
||||||
// @Success 200 {object} response.Any{data=diskcache.Status} "获取成功"
|
// @Success 200 {object} response.Any{data=disk.Status} "获取成功"
|
||||||
// @Failure 401 {object} response.Any "未登录"
|
// @Failure 401 {object} response.Any "未登录"
|
||||||
// @Failure 403 {object} response.Any "无管理员权限"
|
// @Failure 403 {object} response.Any "无管理员权限"
|
||||||
// @Failure 500 {object} response.Any "内部错误"
|
// @Failure 500 {object} response.Any "内部错误"
|
||||||
// @Router /api/v1/admin/cache/status [get]
|
// @Router /api/v1/admin/cache/status [get]
|
||||||
func GetCacheStatus(c *gin.Context) {
|
func GetCacheStatus(c *gin.Context) {
|
||||||
status := diskcache.GetGlobalCache().Status()
|
status := pkgcache.Default().Status()
|
||||||
c.JSON(http.StatusOK, response.OK(status))
|
c.JSON(http.StatusOK, response.OK(status))
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -74,7 +74,7 @@ func UpdateCacheConfig(c *gin.Context) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
diskcache.GetGlobalCache().ReloadConfig(ctx)
|
pkgcache.Default().UpdatePolicy(req.MaxSizeMB, req.TTLMinutes, req.LRUEnabled)
|
||||||
|
|
||||||
c.JSON(http.StatusOK, response.OKNil())
|
c.JSON(http.StatusOK, response.OKNil())
|
||||||
}
|
}
|
||||||
@@ -91,7 +91,7 @@ func UpdateCacheConfig(c *gin.Context) {
|
|||||||
// @Failure 500 {object} response.Any "服务内部错误"
|
// @Failure 500 {object} response.Any "服务内部错误"
|
||||||
// @Router /api/v1/admin/cache/clear [post]
|
// @Router /api/v1/admin/cache/clear [post]
|
||||||
func ClearCache(c *gin.Context) {
|
func ClearCache(c *gin.Context) {
|
||||||
if err := diskcache.GetGlobalCache().Clear(); err != nil {
|
if err := pkgcache.Default().Clear(); err != nil {
|
||||||
response.AbortInternal(c, err.Error())
|
response.AbortInternal(c, err.Error())
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -12,14 +12,13 @@ import (
|
|||||||
"strings"
|
"strings"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
"gorm.io/gorm"
|
||||||
|
|
||||||
"Wavelet/core/contracts"
|
"Wavelet/core/contracts"
|
||||||
"Wavelet/pkg/logger"
|
"Wavelet/pkg/logger"
|
||||||
mail "Wavelet/pkg/mail"
|
mail "Wavelet/pkg/mail"
|
||||||
"Wavelet/pkg/response"
|
"Wavelet/pkg/response"
|
||||||
db "Wavelet/plugins/infra/database"
|
|
||||||
"Wavelet/plugins/infra/storage/objectstore"
|
|
||||||
"github.com/gin-gonic/gin"
|
|
||||||
"gorm.io/gorm"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
const maskedConfigValue = "******"
|
const maskedConfigValue = "******"
|
||||||
@@ -268,9 +267,9 @@ func updateSystemConfig(ctx context.Context, key string, req UpdateSystemConfigR
|
|||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
var originalDriver objectstore.Driver
|
var originalDriver contracts.StorageDriver
|
||||||
if key == ConfigKeyStorageConfig {
|
if key == ConfigKeyStorageConfig {
|
||||||
var currentCfg objectstore.Config
|
var currentCfg contracts.StorageConfigDTO
|
||||||
if err := json.Unmarshal([]byte(config.Value), ¤tCfg); err == nil {
|
if err := json.Unmarshal([]byte(config.Value), ¤tCfg); err == nil {
|
||||||
originalDriver = currentCfg.Driver
|
originalDriver = currentCfg.Driver
|
||||||
}
|
}
|
||||||
@@ -282,7 +281,11 @@ func updateSystemConfig(ctx context.Context, key string, req UpdateSystemConfigR
|
|||||||
req.Value = validatedVal
|
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{
|
updates := map[string]any{
|
||||||
"description": req.Description,
|
"description": req.Description,
|
||||||
}
|
}
|
||||||
@@ -311,14 +314,14 @@ func resolveStorageMigrationTasksOnDirectDriverUpdate(
|
|||||||
ctx context.Context,
|
ctx context.Context,
|
||||||
tx *gorm.DB,
|
tx *gorm.DB,
|
||||||
key string,
|
key string,
|
||||||
originalDriver objectstore.Driver,
|
originalDriver contracts.StorageDriver,
|
||||||
newValue string,
|
newValue string,
|
||||||
) {
|
) {
|
||||||
if key != ConfigKeyStorageConfig || originalDriver == "" {
|
if key != ConfigKeyStorageConfig || originalDriver == "" {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
var newCfg objectstore.Config
|
var newCfg contracts.StorageConfigDTO
|
||||||
if err := json.Unmarshal([]byte(newValue), &newCfg); err != nil {
|
if err := json.Unmarshal([]byte(newValue), &newCfg); err != nil {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
@@ -340,19 +343,12 @@ func invalidateSystemConfigCaches(ctx context.Context, key string) {
|
|||||||
if err := InvalidateSystemConfigCache(ctx, key); err != nil {
|
if err := InvalidateSystemConfigCache(ctx, key); err != nil {
|
||||||
logger.WarnF(ctx, "清理系统配置缓存失败: %v", err)
|
logger.WarnF(ctx, "清理系统配置缓存失败: %v", err)
|
||||||
}
|
}
|
||||||
if globalCoreCtx != nil {
|
_ = EmitEvent(ctx, contracts.EventTopicConfigChanged, contracts.ConfigChangedEvent{Key: key})
|
||||||
_ = globalCoreCtx.Events().Emit(ctx, contracts.EventTopicConfigChanged, contracts.ConfigChangedEvent{Key: key})
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func invalidateCachesAfterConfigUpdate(ctx context.Context, key string) {
|
func invalidateCachesAfterConfigUpdate(ctx context.Context, key string) {
|
||||||
invalidateSystemConfigCaches(ctx, key)
|
invalidateSystemConfigCaches(ctx, key)
|
||||||
|
|
||||||
if key == ConfigKeyStorageConfig {
|
|
||||||
objectstore.ResetCache()
|
|
||||||
objectstore.PublishCacheInvalidation(ctx)
|
|
||||||
}
|
|
||||||
|
|
||||||
if err := InvalidateVisibleSystemConfigsCache(ctx); err != nil {
|
if err := InvalidateVisibleSystemConfigsCache(ctx); err != nil {
|
||||||
logger.WarnF(ctx, "清理公共配置列表缓存失败: %v", err)
|
logger.WarnF(ctx, "清理公共配置列表缓存失败: %v", err)
|
||||||
}
|
}
|
||||||
@@ -442,10 +438,24 @@ func maskSensitiveConfig(key, value string) string {
|
|||||||
case ConfigKeySMTPPassword:
|
case ConfigKeySMTPPassword:
|
||||||
return maskedConfigValue
|
return maskedConfigValue
|
||||||
case ConfigKeyStorageConfig:
|
case ConfigKeyStorageConfig:
|
||||||
var cfg objectstore.Config
|
var cfg contracts.StorageConfigDTO
|
||||||
if err := json.Unmarshal([]byte(value), &cfg); err == nil {
|
if err := json.Unmarshal([]byte(value), &cfg); err == nil {
|
||||||
masked := objectstore.MaskSecrets(cfg)
|
if cfg.S3.SecretAccessKey != "" {
|
||||||
if val, err := json.Marshal(masked); err == nil {
|
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)
|
return string(val)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -456,18 +466,34 @@ func maskSensitiveConfig(key, value string) string {
|
|||||||
// validateAndMergeStorageConfig parses, merges unmasked secrets, validates parameter values,
|
// validateAndMergeStorageConfig parses, merges unmasked secrets, validates parameter values,
|
||||||
// and tests connectivity of the new storage configuration.
|
// and tests connectivity of the new storage configuration.
|
||||||
func validateAndMergeStorageConfig(ctx context.Context, value string, currentConfig string) (string, error) {
|
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 {
|
if err := json.Unmarshal([]byte(currentConfig), ¤tCfg); err != nil {
|
||||||
return "", fmt.Errorf("解析当前存储配置失败: %w", err)
|
return "", fmt.Errorf("解析当前存储配置失败: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
var newCfg objectstore.Config
|
var newCfg contracts.StorageConfigDTO
|
||||||
if err := json.Unmarshal([]byte(value), &newCfg); err != nil {
|
if err := json.Unmarshal([]byte(value), &newCfg); err != nil {
|
||||||
return "", fmt.Errorf("解析目标存储配置失败: %w", err)
|
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 {
|
if err := validateMergedStorageConfig(ctx, currentCfg, newCfg, targetCfg); err != nil {
|
||||||
return "", err
|
return "", err
|
||||||
}
|
}
|
||||||
@@ -481,43 +507,21 @@ func validateAndMergeStorageConfig(ctx context.Context, value string, currentCon
|
|||||||
return string(unmaskedVal), nil
|
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 {
|
if newCfg.Driver != "" && newCfg.Driver != currentCfg.Driver {
|
||||||
var uploadCount int64
|
var uploadCount int64
|
||||||
if err := db.DB(ctx).Table("w_uploads").
|
gormDB := GetDB(ctx)
|
||||||
Where("status != ?", "deleted").
|
if gormDB != nil {
|
||||||
Count(&uploadCount).Error; err != nil {
|
if err := gormDB.Table("w_uploads").
|
||||||
return fmt.Errorf("检查存量文件失败: %w", err)
|
Where("status != ?", "deleted").
|
||||||
|
Count(&uploadCount).Error; err != nil {
|
||||||
|
return fmt.Errorf("检查存量文件失败: %w", err)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
if uploadCount > 0 {
|
if uploadCount > 0 {
|
||||||
return errors.New(StorageDriverSwitchRequiresMigration)
|
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
|
return nil
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -20,7 +20,6 @@ import (
|
|||||||
|
|
||||||
"Wavelet/pkg/config"
|
"Wavelet/pkg/config"
|
||||||
"Wavelet/pkg/response"
|
"Wavelet/pkg/response"
|
||||||
db "Wavelet/plugins/infra/database"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
const (
|
const (
|
||||||
@@ -219,7 +218,7 @@ func getPostgresOverview(gormDB *gorm.DB) (DBOverviewResponse, error) {
|
|||||||
// @Failure 500 {object} response.Any "内部错误"
|
// @Failure 500 {object} response.Any "内部错误"
|
||||||
// @Router /api/v1/admin/db-manage/overview [get]
|
// @Router /api/v1/admin/db-manage/overview [get]
|
||||||
func GetDBOverview(c *gin.Context) {
|
func GetDBOverview(c *gin.Context) {
|
||||||
gormDB := db.DB(c.Request.Context())
|
gormDB := GetDB(c.Request.Context())
|
||||||
if gormDB == nil {
|
if gormDB == nil {
|
||||||
response.AbortInternal(c, "数据库未初始化")
|
response.AbortInternal(c, "数据库未初始化")
|
||||||
return
|
return
|
||||||
@@ -254,7 +253,7 @@ func GetDBOverview(c *gin.Context) {
|
|||||||
// @Failure 500 {object} response.Any "内部错误"
|
// @Failure 500 {object} response.Any "内部错误"
|
||||||
// @Router /api/v1/admin/db-manage/tables [get]
|
// @Router /api/v1/admin/db-manage/tables [get]
|
||||||
func ListDBTables(c *gin.Context) {
|
func ListDBTables(c *gin.Context) {
|
||||||
gormDB := db.DB(c.Request.Context())
|
gormDB := GetDB(c.Request.Context())
|
||||||
if gormDB == nil {
|
if gormDB == nil {
|
||||||
response.AbortInternal(c, "数据库未初始化")
|
response.AbortInternal(c, "数据库未初始化")
|
||||||
return
|
return
|
||||||
@@ -285,7 +284,7 @@ func GetDBTableData(c *gin.Context) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
gormDB := db.DB(c.Request.Context())
|
gormDB := GetDB(c.Request.Context())
|
||||||
if gormDB == nil {
|
if gormDB == nil {
|
||||||
response.AbortInternal(c, "数据库未初始化")
|
response.AbortInternal(c, "数据库未初始化")
|
||||||
return
|
return
|
||||||
@@ -458,7 +457,7 @@ func ExecuteSQL(c *gin.Context) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
gormDB := db.DB(c.Request.Context())
|
gormDB := GetDB(c.Request.Context())
|
||||||
if gormDB == nil {
|
if gormDB == nil {
|
||||||
response.AbortInternal(c, "数据库未初始化")
|
response.AbortInternal(c, "数据库未初始化")
|
||||||
return
|
return
|
||||||
@@ -508,7 +507,7 @@ func getSQLiteInfo(ctx context.Context) DatabaseInfoResponse {
|
|||||||
if info.Name == "" {
|
if info.Name == "" {
|
||||||
info.Name = "./data/wavelet.db"
|
info.Name = "./data/wavelet.db"
|
||||||
}
|
}
|
||||||
gormDB := db.DB(ctx)
|
gormDB := GetDB(ctx)
|
||||||
if gormDB == nil {
|
if gormDB == nil {
|
||||||
return info
|
return info
|
||||||
}
|
}
|
||||||
@@ -525,7 +524,7 @@ func getPostgresInfo(ctx context.Context) DatabaseInfoResponse {
|
|||||||
Name: config.Config.Database.Database,
|
Name: config.Config.Database.Database,
|
||||||
Version: "PostgreSQL",
|
Version: "PostgreSQL",
|
||||||
}
|
}
|
||||||
gormDB := db.DB(ctx)
|
gormDB := GetDB(ctx)
|
||||||
if gormDB == nil {
|
if gormDB == nil {
|
||||||
return info
|
return info
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -14,16 +14,14 @@ import (
|
|||||||
"strings"
|
"strings"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
"github.com/gorilla/websocket"
|
||||||
|
|
||||||
|
"Wavelet/core/contracts"
|
||||||
"Wavelet/pkg/config"
|
"Wavelet/pkg/config"
|
||||||
"Wavelet/pkg/logger"
|
"Wavelet/pkg/logger"
|
||||||
"Wavelet/pkg/response"
|
"Wavelet/pkg/response"
|
||||||
"Wavelet/pkg/util"
|
"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 (
|
const (
|
||||||
@@ -138,6 +136,7 @@ func HandleLogWebSocket(c *gin.Context) {
|
|||||||
// accessLogItem 访问日志单条数据
|
// accessLogItem 访问日志单条数据
|
||||||
type accessLogItem struct {
|
type accessLogItem struct {
|
||||||
ID uint64 `json:"id,string"`
|
ID uint64 `json:"id,string"`
|
||||||
|
TraceID string `json:"trace_id"`
|
||||||
UserID uint64 `json:"user_id,string"`
|
UserID uint64 `json:"user_id,string"`
|
||||||
Username string `json:"username"`
|
Username string `json:"username"`
|
||||||
Nickname string `json:"nickname"`
|
Nickname string `json:"nickname"`
|
||||||
@@ -157,16 +156,19 @@ type accessLogsResponse struct {
|
|||||||
List []accessLogItem `json:"list"`
|
List []accessLogItem `json:"list"`
|
||||||
}
|
}
|
||||||
|
|
||||||
func buildAccessLogFilter(ctx context.Context, c *gin.Context) (logstore.AccessLogFilter, error) {
|
func buildAccessLogFilter(ctx context.Context, c *gin.Context) (contracts.AccessLogFilterDTO, error) {
|
||||||
filter := logstore.AccessLogFilter{}
|
filter := contracts.AccessLogFilterDTO{}
|
||||||
|
|
||||||
username := c.Query("username")
|
username := c.Query("username")
|
||||||
if username != "" {
|
if username != "" {
|
||||||
var userIDs []uint64
|
var userIDs []uint64
|
||||||
if err := db.DB(ctx).Table("w_users").
|
gormDB := GetDB(ctx)
|
||||||
Where("username LIKE ? ESCAPE '\\'", "%"+util.EscapeLike(username)+"%").
|
if gormDB != nil {
|
||||||
Pluck("id", &userIDs).Error; err != nil {
|
if err := gormDB.Table("w_users").
|
||||||
return filter, fmt.Errorf("查询用户信息失败: %w", err)
|
Where("username LIKE ? ESCAPE '\\'", "%"+util.EscapeLike(username)+"%").
|
||||||
|
Pluck("id", &userIDs).Error; err != nil {
|
||||||
|
return filter, fmt.Errorf("查询用户信息失败: %w", err)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
filter.UserIDs = userIDs
|
filter.UserIDs = userIDs
|
||||||
}
|
}
|
||||||
@@ -218,9 +220,12 @@ func enrichAccessLogsWithUsers(ctx context.Context, list []accessLogItem) {
|
|||||||
Username string
|
Username string
|
||||||
Nickname string
|
Nickname string
|
||||||
}
|
}
|
||||||
if err := db.DB(ctx).Table("w_users").Where("id IN ?", userIDs).Find(&users).Error; err == nil {
|
gormDB := GetDB(ctx)
|
||||||
for _, u := range users {
|
if gormDB != nil {
|
||||||
userMap[u.ID] = struct{ Username, Nickname string }{Username: u.Username, Nickname: u.Nickname}
|
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 {
|
for i := range list {
|
||||||
@@ -251,9 +256,9 @@ func enrichAccessLogsWithUsers(ctx context.Context, list []accessLogItem) {
|
|||||||
// @Router /api/v1/admin/logs/access [get]
|
// @Router /api/v1/admin/logs/access [get]
|
||||||
func GetAccessLogs(c *gin.Context) {
|
func GetAccessLogs(c *gin.Context) {
|
||||||
ctx := c.Request.Context()
|
ctx := c.Request.Context()
|
||||||
store, err := logstore.Active(ctx)
|
rc := GetRiskControlService()
|
||||||
if err != nil {
|
if rc == nil {
|
||||||
response.AbortInternal(c, "日志存储初始化失败")
|
response.AbortInternal(c, "日志存储服务未初始化")
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -279,7 +284,7 @@ func GetAccessLogs(c *gin.Context) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
logs, total, err := store.UserAccessLogs.List(ctx, filter, page, pageSize)
|
logs, total, err := rc.QueryAccessLogs(ctx, filter, page, pageSize)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
response.AbortWithError(c, http.StatusInternalServerError, err.Error())
|
response.AbortWithError(c, http.StatusInternalServerError, err.Error())
|
||||||
return
|
return
|
||||||
@@ -298,7 +303,6 @@ func GetAccessLogs(c *gin.Context) {
|
|||||||
Method: logItem.Method,
|
Method: logItem.Method,
|
||||||
IP: logItem.IP,
|
IP: logItem.IP,
|
||||||
UserAgent: logItem.UserAgent,
|
UserAgent: logItem.UserAgent,
|
||||||
Headers: logItem.Headers,
|
|
||||||
Status: logItem.Status,
|
Status: logItem.Status,
|
||||||
Latency: logItem.Latency,
|
Latency: logItem.Latency,
|
||||||
CreatedAt: logItem.CreatedAt.Format(time.RFC3339),
|
CreatedAt: logItem.CreatedAt.Format(time.RFC3339),
|
||||||
@@ -352,84 +356,27 @@ type logsAnalyticsResponse struct {
|
|||||||
// @Router /api/v1/admin/logs/analytics [get]
|
// @Router /api/v1/admin/logs/analytics [get]
|
||||||
func GetLogsAnalytics(c *gin.Context) {
|
func GetLogsAnalytics(c *gin.Context) {
|
||||||
ctx := c.Request.Context()
|
ctx := c.Request.Context()
|
||||||
store, err := logstore.Active(ctx)
|
rc := GetRiskControlService()
|
||||||
if err != nil {
|
if rc == nil {
|
||||||
response.AbortInternal(c, "日志存储初始化失败")
|
response.AbortInternal(c, "日志存储服务未初始化")
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
startTime := time.Now().AddDate(0, 0, -(analyticsDays - 1)).Truncate(hoursInDay * time.Hour)
|
stats, err := rc.QueryAccessLogStats(ctx, analyticsDays)
|
||||||
|
|
||||||
trendPoints, err := store.UserAccessLogs.GetDailyTrend(ctx, analyticsDays)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
response.AbortWithError(c, http.StatusInternalServerError, "查询访问趋势失败: "+err.Error())
|
response.AbortWithError(c, http.StatusInternalServerError, "查询访问趋势失败: "+err.Error())
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
trendList := make([]trendItem, len(trendPoints))
|
trendList := make([]trendItem, len(stats))
|
||||||
for i, point := range trendPoints {
|
for i, st := range stats {
|
||||||
trendList[i] = trendItem{
|
trendList[i] = trendItem{
|
||||||
Date: point.Date,
|
Date: st.Date,
|
||||||
Count: point.Count,
|
Count: st.PV,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
browserPoints, err := store.UserAccessLogs.GetBrowserDistribution(ctx, startTime)
|
browserList := []browserItem{}
|
||||||
if err != nil {
|
topUsers := []topUserItem{}
|
||||||
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
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
c.JSON(http.StatusOK, response.OK(logsAnalyticsResponse{
|
c.JSON(http.StatusOK, response.OK(logsAnalyticsResponse{
|
||||||
Trend: trendList,
|
Trend: trendList,
|
||||||
@@ -496,18 +443,14 @@ const (
|
|||||||
)
|
)
|
||||||
|
|
||||||
// LogDBSwitchMeta 描述切换日志数据库任务。
|
// LogDBSwitchMeta 描述切换日志数据库任务。
|
||||||
var LogDBSwitchMeta = driver_asynq_worker.TaskMeta{
|
var LogDBSwitchMeta = contracts.TaskMetaDTO{
|
||||||
Type: TaskTypeLogDBSwitch,
|
Name: LogDBSwitchTask,
|
||||||
AsynqTask: LogDBSwitchTask,
|
DisplayName: "切换日志数据库",
|
||||||
Name: "切换日志数据库",
|
Description: "复制迁移用户访问日志并在成功后切换日志主库(期间禁止日志写入)",
|
||||||
Description: "复制迁移用户访问日志并在成功后切换日志主库(期间禁止日志写入)",
|
MaxRetry: 3,
|
||||||
SupportsTime: false,
|
Queue: "default",
|
||||||
MaxRetry: driver_asynq_worker.DefaultMaxRetry,
|
Params: []contracts.TaskParamDTO{
|
||||||
Queue: driver_asynq_worker.QueueDefault,
|
{Name: "target", Description: "迁移目标:postgres(主库为 PG 时)、sqlite(主库为 SQLite 时)或 clickhouse", Type: "string", Required: true},
|
||||||
Retryable: true,
|
|
||||||
Params: []driver_asynq_worker.TaskParam{
|
|
||||||
{Name: "target", Label: "目标日志库", Type: "string", Required: true,
|
|
||||||
Placeholder: "postgres|sqlite|clickhouse", Description: "迁移目标:postgres(主库为 PG 时)、sqlite(主库为 SQLite 时)或 clickhouse"},
|
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -552,7 +495,7 @@ func validTarget(v string) bool {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Execute 执行迁移。
|
// 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
|
var p logDBSwitchPayload
|
||||||
if err := json.Unmarshal(payload, &p); err != nil {
|
if err := json.Unmarshal(payload, &p); err != nil {
|
||||||
return nil, fmt.Errorf("参数解析失败: %w", err)
|
return nil, fmt.Errorf("参数解析失败: %w", err)
|
||||||
@@ -564,10 +507,13 @@ func (h *LogDBSwitchHandler) Execute(ctx context.Context, payload []byte) (*driv
|
|||||||
|
|
||||||
source, err := currentLogDatabase(ctx)
|
source, err := currentLogDatabase(ctx)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
driver_asynq_worker.AppendLog(ctx, "读取日志主库失败: %v", err)
|
|
||||||
return nil, 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 {
|
if err := setMigrationFlag(ctx, "migrating"); err != nil {
|
||||||
return nil, err
|
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 {
|
rc := GetRiskControlService()
|
||||||
return nil, fmt.Errorf("排空日志写入队列失败: %w", err)
|
if rc != nil {
|
||||||
}
|
if err := rc.SwitchLogEngine(ctx, p.Target); err != nil {
|
||||||
|
return nil, 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)
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := copyUserAccessLogs(ctx, src, dst); err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
if err := flipLogDatabase(ctx, p.Target); err != nil {
|
if err := flipLogDatabase(ctx, p.Target); err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
logstore.InvalidateCache()
|
|
||||||
driver_asynq_worker.AppendLog(ctx, "日志数据库已切换为 %s,写入恢复", p.Target)
|
if taskSvc != nil {
|
||||||
return &driver_asynq_worker.TaskResult{Message: fmt.Sprintf("日志数据库已从 %s 切换为 %s", source, p.Target)}, 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 {
|
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 {
|
func flipLogDatabase(ctx context.Context, target string) error {
|
||||||
return SaveOrUpdateSystemConfig(ctx, ConfigKeyLogDatabase, target)
|
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/config"
|
||||||
"Wavelet/pkg/logger"
|
"Wavelet/pkg/logger"
|
||||||
"Wavelet/pkg/response"
|
"Wavelet/pkg/response"
|
||||||
"Wavelet/plugins/domain/risk_control/logstore"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
var startTime = time.Now()
|
var startTime = time.Now()
|
||||||
@@ -177,21 +176,13 @@ type LogDatabaseStatus struct {
|
|||||||
// @Router /api/v1/admin/status/log-database [get]
|
// @Router /api/v1/admin/status/log-database [get]
|
||||||
func GetLogDatabaseStatus(c *gin.Context) {
|
func GetLogDatabaseStatus(c *gin.Context) {
|
||||||
ctx := c.Request.Context()
|
ctx := c.Request.Context()
|
||||||
store, err := logstore.Active(ctx)
|
activeDB := "sqlite"
|
||||||
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
|
|
||||||
}
|
|
||||||
migration := "idle"
|
migration := "idle"
|
||||||
if logstore.Migrating(ctx) {
|
if rc := GetRiskControlService(); rc != nil {
|
||||||
migration = "migrating"
|
activeDB = rc.ActiveLogEngine(ctx)
|
||||||
|
if rc.IsLogEngineMigrating(ctx) {
|
||||||
|
migration = "migrating"
|
||||||
|
}
|
||||||
}
|
}
|
||||||
c.JSON(http.StatusOK, response.OK(LogDatabaseStatus{
|
c.JSON(http.StatusOK, response.OK(LogDatabaseStatus{
|
||||||
ActiveDatabase: activeDB,
|
ActiveDatabase: activeDB,
|
||||||
|
|||||||
@@ -13,10 +13,9 @@ import (
|
|||||||
"github.com/gin-gonic/gin"
|
"github.com/gin-gonic/gin"
|
||||||
"github.com/robfig/cron/v3"
|
"github.com/robfig/cron/v3"
|
||||||
|
|
||||||
|
"Wavelet/core/contracts"
|
||||||
"Wavelet/pkg/logger"
|
"Wavelet/pkg/logger"
|
||||||
"Wavelet/pkg/response"
|
"Wavelet/pkg/response"
|
||||||
"Wavelet/plugins/drivers/driver_asynq_cron"
|
|
||||||
"Wavelet/plugins/drivers/driver_asynq_worker"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
// ListTaskTypes 获取支持的任务类型列表
|
// ListTaskTypes 获取支持的任务类型列表
|
||||||
@@ -25,12 +24,17 @@ import (
|
|||||||
// @Tags admin
|
// @Tags admin
|
||||||
// @Produce json
|
// @Produce json
|
||||||
// @Security SessionCookie
|
// @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 401 {object} response.Any "未登录"
|
||||||
// @Failure 403 {object} response.Any "无管理员权限"
|
// @Failure 403 {object} response.Any "无管理员权限"
|
||||||
// @Router /api/v1/admin/tasks/types [get]
|
// @Router /api/v1/admin/tasks/types [get]
|
||||||
func ListTaskTypes(c *gin.Context) {
|
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 下发任务请求
|
// DispatchTaskRequest 下发任务请求
|
||||||
@@ -63,8 +67,14 @@ func DispatchTask(c *gin.Context) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
meta := driver_asynq_worker.GetTaskMeta(req.TaskType)
|
taskSvc := GetTaskService()
|
||||||
if meta == nil {
|
if taskSvc == nil {
|
||||||
|
response.AbortInternal(c, "task service not available")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
meta, ok := taskSvc.GetTaskMeta(req.TaskType)
|
||||||
|
if !ok {
|
||||||
response.AbortBadRequest(c, InvalidTaskType)
|
response.AbortBadRequest(c, InvalidTaskType)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
@@ -74,13 +84,13 @@ func DispatchTask(c *gin.Context) {
|
|||||||
payloadBytes = []byte(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 {
|
if err != nil {
|
||||||
response.AbortBadRequest(c, err.Error())
|
response.AbortBadRequest(c, err.Error())
|
||||||
return
|
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 {
|
if err != nil {
|
||||||
response.AbortInternal(c, fmt.Sprintf("%s: %v", TaskDispatchFailed, err))
|
response.AbortInternal(c, fmt.Sprintf("%s: %v", TaskDispatchFailed, err))
|
||||||
return
|
return
|
||||||
@@ -111,8 +121,11 @@ func ListTaskExecutions(c *gin.Context) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
if req.TaskType != "" {
|
if req.TaskType != "" {
|
||||||
if meta := driver_asynq_worker.GetTaskMeta(req.TaskType); meta != nil {
|
taskSvc := GetTaskService()
|
||||||
req.TaskType = meta.AsynqTask
|
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
|
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 {
|
if err != nil {
|
||||||
errMsg := err.Error()
|
errMsg := err.Error()
|
||||||
switch {
|
switch {
|
||||||
@@ -252,9 +271,15 @@ func CreateSchedule(c *gin.Context) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
|
taskSvc := GetTaskService()
|
||||||
|
if taskSvc == nil {
|
||||||
|
response.AbortInternal(c, "task service not available")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
// 校验关联的异步任务类型
|
// 校验关联的异步任务类型
|
||||||
meta := driver_asynq_worker.GetTaskMeta(req.TaskType)
|
meta, ok := taskSvc.GetTaskMeta(req.TaskType)
|
||||||
if meta == nil {
|
if !ok {
|
||||||
response.AbortBadRequest(c, InvalidTaskType)
|
response.AbortBadRequest(c, InvalidTaskType)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
@@ -264,7 +289,7 @@ func CreateSchedule(c *gin.Context) {
|
|||||||
if strings.TrimSpace(req.Payload) != "" {
|
if strings.TrimSpace(req.Payload) != "" {
|
||||||
payloadBytes = []byte(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 {
|
if err != nil {
|
||||||
response.AbortBadRequest(c, err.Error())
|
response.AbortBadRequest(c, err.Error())
|
||||||
return
|
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)
|
logger.ErrorF(c.Request.Context(), "[TaskAdmin] 重载调度器失败: %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -341,9 +366,15 @@ func UpdateSchedule(c *gin.Context) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
|
taskSvc := GetTaskService()
|
||||||
|
if taskSvc == nil {
|
||||||
|
response.AbortInternal(c, "task service not available")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
// 校验关联的异步任务类型
|
// 校验关联的异步任务类型
|
||||||
meta := driver_asynq_worker.GetTaskMeta(req.TaskType)
|
meta, ok := taskSvc.GetTaskMeta(req.TaskType)
|
||||||
if meta == nil {
|
if !ok {
|
||||||
response.AbortBadRequest(c, InvalidTaskType)
|
response.AbortBadRequest(c, InvalidTaskType)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
@@ -353,7 +384,7 @@ func UpdateSchedule(c *gin.Context) {
|
|||||||
if strings.TrimSpace(req.Payload) != "" {
|
if strings.TrimSpace(req.Payload) != "" {
|
||||||
payloadBytes = []byte(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 {
|
if err != nil {
|
||||||
response.AbortBadRequest(c, err.Error())
|
response.AbortBadRequest(c, err.Error())
|
||||||
return
|
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)
|
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 {
|
taskSvc := GetTaskService()
|
||||||
logger.ErrorF(c.Request.Context(), "[TaskAdmin] 重载调度器失败: %v", err)
|
if taskSvc != nil {
|
||||||
|
if err := taskSvc.ReloadScheduler(); err != nil {
|
||||||
|
logger.ErrorF(c.Request.Context(), "[TaskAdmin] 重载调度器失败: %v", err)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
c.JSON(http.StatusOK, response.OKNil())
|
c.JSON(http.StatusOK, response.OKNil())
|
||||||
|
|||||||
@@ -129,7 +129,7 @@ func ListUsers(c *gin.Context) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
userSvc := getUserService(c.Request.Context())
|
userSvc := GetUserService(c.Request.Context())
|
||||||
if userSvc == nil {
|
if userSvc == nil {
|
||||||
response.AbortInternal(c, "用户服务未就绪")
|
response.AbortInternal(c, "用户服务未就绪")
|
||||||
return
|
return
|
||||||
@@ -179,7 +179,7 @@ func GetUser(c *gin.Context) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
userSvc := getUserService(c.Request.Context())
|
userSvc := GetUserService(c.Request.Context())
|
||||||
if userSvc == nil {
|
if userSvc == nil {
|
||||||
response.AbortInternal(c, "用户服务未就绪")
|
response.AbortInternal(c, "用户服务未就绪")
|
||||||
return
|
return
|
||||||
@@ -226,7 +226,7 @@ func UpdateUserStatus(c *gin.Context) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
userSvc := getUserService(c.Request.Context())
|
userSvc := GetUserService(c.Request.Context())
|
||||||
if userSvc == nil {
|
if userSvc == nil {
|
||||||
response.AbortInternal(c, "用户服务未就绪")
|
response.AbortInternal(c, "用户服务未就绪")
|
||||||
return
|
return
|
||||||
@@ -269,7 +269,7 @@ func DeleteUser(c *gin.Context) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
userSvc := getUserService(c.Request.Context())
|
userSvc := GetUserService(c.Request.Context())
|
||||||
if userSvc == nil {
|
if userSvc == nil {
|
||||||
response.AbortInternal(c, "用户服务未就绪")
|
response.AbortInternal(c, "用户服务未就绪")
|
||||||
return
|
return
|
||||||
@@ -317,7 +317,7 @@ func CreateUser(c *gin.Context) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
userSvc := getUserService(c.Request.Context())
|
userSvc := GetUserService(c.Request.Context())
|
||||||
if userSvc == nil {
|
if userSvc == nil {
|
||||||
response.AbortInternal(c, "用户服务未就绪")
|
response.AbortInternal(c, "用户服务未就绪")
|
||||||
return
|
return
|
||||||
@@ -380,7 +380,7 @@ func UpdateUser(c *gin.Context) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
userSvc := getUserService(c.Request.Context())
|
userSvc := GetUserService(c.Request.Context())
|
||||||
if userSvc == nil {
|
if userSvc == nil {
|
||||||
response.AbortInternal(c, "用户服务未就绪")
|
response.AbortInternal(c, "用户服务未就绪")
|
||||||
return
|
return
|
||||||
|
|||||||
@@ -4,12 +4,13 @@
|
|||||||
package admin
|
package admin
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
|
||||||
"Wavelet/core/contracts"
|
"Wavelet/core/contracts"
|
||||||
"Wavelet/pkg/logger"
|
"Wavelet/pkg/logger"
|
||||||
"Wavelet/pkg/response"
|
"Wavelet/pkg/response"
|
||||||
"Wavelet/pkg/trace"
|
"Wavelet/pkg/trace"
|
||||||
"Wavelet/pkg/util"
|
"Wavelet/pkg/util"
|
||||||
"github.com/gin-gonic/gin"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
// LoginAdminRequired 返回管理员权限校验中间件
|
// LoginAdminRequired 返回管理员权限校验中间件
|
||||||
|
|||||||
@@ -9,11 +9,12 @@ import (
|
|||||||
"embed"
|
"embed"
|
||||||
"reflect"
|
"reflect"
|
||||||
|
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
"github.com/hibiken/asynq"
|
||||||
|
|
||||||
"Wavelet/core"
|
"Wavelet/core"
|
||||||
"Wavelet/core/contracts"
|
"Wavelet/core/contracts"
|
||||||
"Wavelet/core/extpoints"
|
"Wavelet/core/extpoints"
|
||||||
"github.com/gin-gonic/gin"
|
|
||||||
"github.com/hibiken/asynq"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
//go:embed migrations/*.sql
|
//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.
|
// Apply registers admin routes, tasks, schedules, and settings into the Context.
|
||||||
func (p *Plugin) Apply(ctx *core.Context) error {
|
func (p *Plugin) Apply(ctx *core.Context) error {
|
||||||
globalCoreCtx = ctx
|
// 0. Bind Services reactively
|
||||||
|
if db, err := core.Inject[contracts.DBService](ctx); err == nil && db != nil {
|
||||||
// 0. Resolve auth and user services reactively via IoC
|
SetDBService(db)
|
||||||
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
|
|
||||||
}
|
|
||||||
} else {
|
} else {
|
||||||
core.When[contracts.AuthService](ctx, func(svc contracts.AuthService) {
|
core.When[contracts.DBService](ctx, func(db contracts.DBService) {
|
||||||
globalAuthSvc = svc
|
SetDBService(db)
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
if cache, err := core.Inject[contracts.CacheService](ctx); err == nil && cache != nil {
|
||||||
if userSvc, err := core.Inject[contracts.UserService](ctx); err == nil && userSvc != nil {
|
SetCacheService(cache)
|
||||||
globalUserSvc = userSvc
|
|
||||||
} else {
|
} else {
|
||||||
core.When[contracts.UserService](ctx, func(svc contracts.UserService) {
|
core.When[contracts.CacheService](ctx, func(cache contracts.CacheService) {
|
||||||
globalUserSvc = svc
|
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)
|
ctx.Migrations().Register("admin", adminMigrations)
|
||||||
|
|
||||||
// 1. Register Admin HTTP Routes
|
// 1. Register Admin HTTP Routes
|
||||||
|
|||||||
@@ -12,15 +12,12 @@ import (
|
|||||||
"strings"
|
"strings"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/redis/go-redis/v9"
|
|
||||||
"github.com/shopspring/decimal"
|
"github.com/shopspring/decimal"
|
||||||
"gorm.io/gorm"
|
"gorm.io/gorm"
|
||||||
|
|
||||||
"Wavelet/pkg/cache/ram"
|
"Wavelet/pkg/cache/ram"
|
||||||
"Wavelet/pkg/idgen"
|
"Wavelet/pkg/idgen"
|
||||||
"Wavelet/pkg/util"
|
"Wavelet/pkg/util"
|
||||||
cachepkg "Wavelet/plugins/infra/cache"
|
|
||||||
db "Wavelet/plugins/infra/database"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
const (
|
const (
|
||||||
@@ -38,7 +35,7 @@ const (
|
|||||||
|
|
||||||
// PreheatSystemConfigs loads all system configs from database.
|
// PreheatSystemConfigs loads all system configs from database.
|
||||||
func PreheatSystemConfigs(ctx context.Context) ([]SystemConfig, error) {
|
func PreheatSystemConfigs(ctx context.Context) ([]SystemConfig, error) {
|
||||||
database := db.DB(ctx)
|
database := GetDB(ctx)
|
||||||
if database == nil {
|
if database == nil {
|
||||||
return nil, errors.New(errDatabaseNotInitialized)
|
return nil, errors.New(errDatabaseNotInitialized)
|
||||||
}
|
}
|
||||||
@@ -52,7 +49,7 @@ func PreheatSystemConfigs(ctx context.Context) ([]SystemConfig, error) {
|
|||||||
|
|
||||||
// PreheatSystemConfigByKey loads a single config key from database.
|
// PreheatSystemConfigByKey loads a single config key from database.
|
||||||
func PreheatSystemConfigByKey(ctx context.Context, key string) (SystemConfig, error) {
|
func PreheatSystemConfigByKey(ctx context.Context, key string) (SystemConfig, error) {
|
||||||
database := db.DB(ctx)
|
database := GetDB(ctx)
|
||||||
if database == nil {
|
if database == nil {
|
||||||
return SystemConfig{}, errors.New(errDatabaseNotInitialized)
|
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 {
|
if database == nil {
|
||||||
return SystemConfig{}, errors.New(errDatabaseNotInitialized)
|
return SystemConfig{}, errors.New(errDatabaseNotInitialized)
|
||||||
}
|
}
|
||||||
@@ -129,7 +126,7 @@ func ListSystemConfigsByKeys(ctx context.Context, keys []string) (map[string]Sys
|
|||||||
return result, nil
|
return result, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
database := db.DB(ctx)
|
database := GetDB(ctx)
|
||||||
if database == nil {
|
if database == nil {
|
||||||
return nil, errors.New(errDatabaseNotInitialized)
|
return nil, errors.New(errDatabaseNotInitialized)
|
||||||
}
|
}
|
||||||
@@ -178,7 +175,7 @@ func ListVisibleSystemConfigs(ctx context.Context) ([]SystemConfig, error) {
|
|||||||
return list, nil
|
return list, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
database := db.DB(ctx)
|
database := GetDB(ctx)
|
||||||
if database == nil {
|
if database == nil {
|
||||||
return nil, errors.New(errDatabaseNotInitialized)
|
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.
|
// ListAdminSystemConfigs returns all configs, optionally filtered by type.
|
||||||
func ListAdminSystemConfigs(ctx context.Context, configType string) ([]SystemConfig, error) {
|
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 != "" {
|
if configType != "" {
|
||||||
query = query.Where("type = ?", 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.
|
// GetAdminSystemConfigByKey loads a config directly from DB.
|
||||||
func GetAdminSystemConfigByKey(ctx context.Context, key string) (SystemConfig, error) {
|
func GetAdminSystemConfigByKey(ctx context.Context, key string) (SystemConfig, error) {
|
||||||
var config SystemConfig
|
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 SystemConfig{}, err
|
||||||
}
|
}
|
||||||
return config, nil
|
return config, nil
|
||||||
@@ -292,7 +289,7 @@ func GetAdminSystemConfigByKey(ctx context.Context, key string) (SystemConfig, e
|
|||||||
// SystemConfigExists reports whether a config key already exists.
|
// SystemConfigExists reports whether a config key already exists.
|
||||||
func SystemConfigExists(ctx context.Context, key string) (bool, error) {
|
func SystemConfigExists(ctx context.Context, key string) (bool, error) {
|
||||||
var existing SystemConfig
|
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) {
|
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||||
return false, nil
|
return false, nil
|
||||||
}
|
}
|
||||||
@@ -304,18 +301,18 @@ func SystemConfigExists(ctx context.Context, key string) (bool, error) {
|
|||||||
|
|
||||||
// CreateSystemConfigRecord persists a new system config row.
|
// CreateSystemConfigRecord persists a new system config row.
|
||||||
func CreateSystemConfigRecord(ctx context.Context, config *SystemConfig) error {
|
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.
|
// UpdateSystemConfigFields applies partial updates to a system config row.
|
||||||
func UpdateSystemConfigFields(ctx context.Context, config *SystemConfig, updates map[string]any) error {
|
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.
|
// SaveOrUpdateSystemConfig creates or updates a config row and invalidates cache.
|
||||||
func SaveOrUpdateSystemConfig(ctx context.Context, key, value string) error {
|
func SaveOrUpdateSystemConfig(ctx context.Context, key, value string) error {
|
||||||
var sc SystemConfig
|
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) {
|
if err != nil && !errors.Is(err, gorm.ErrRecordNotFound) {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
@@ -327,12 +324,12 @@ func SaveOrUpdateSystemConfig(ctx context.Context, key, value string) error {
|
|||||||
Type: configTypeSystem,
|
Type: configTypeSystem,
|
||||||
Visibility: ConfigVisibilityHidden,
|
Visibility: ConfigVisibilityHidden,
|
||||||
}
|
}
|
||||||
if err := db.DB(ctx).Create(&sc).Error; err != nil {
|
if err := GetDB(ctx).Create(&sc).Error; err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
} else {
|
} else {
|
||||||
sc.Value = value
|
sc.Value = value
|
||||||
if err := db.DB(ctx).Save(&sc).Error; err != nil {
|
if err := GetDB(ctx).Save(&sc).Error; err != nil {
|
||||||
return err
|
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.
|
// ListTemplatesRecord returns all templates ordered by system flag and creation time.
|
||||||
func ListTemplatesRecord(ctx context.Context) ([]Template, error) {
|
func ListTemplatesRecord(ctx context.Context) ([]Template, error) {
|
||||||
var templates []Template
|
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 nil, err
|
||||||
}
|
}
|
||||||
return templates, nil
|
return templates, nil
|
||||||
@@ -351,7 +348,7 @@ func ListTemplatesRecord(ctx context.Context) ([]Template, error) {
|
|||||||
// GetTemplateByKey loads a template by its key.
|
// GetTemplateByKey loads a template by its key.
|
||||||
func GetTemplateByKey(ctx context.Context, key string) (Template, error) {
|
func GetTemplateByKey(ctx context.Context, key string) (Template, error) {
|
||||||
var tmpl Template
|
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 Template{}, err
|
||||||
}
|
}
|
||||||
return tmpl, nil
|
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.
|
// TemplateExistsByKey reports whether a template key is already taken.
|
||||||
func TemplateExistsByKey(ctx context.Context, key string) (bool, error) {
|
func TemplateExistsByKey(ctx context.Context, key string) (bool, error) {
|
||||||
var existing Template
|
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) {
|
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||||
return false, nil
|
return false, nil
|
||||||
}
|
}
|
||||||
@@ -372,38 +369,38 @@ func TemplateExistsByKey(ctx context.Context, key string) (bool, error) {
|
|||||||
|
|
||||||
// CreateTemplateRecord persists a new template.
|
// CreateTemplateRecord persists a new template.
|
||||||
func CreateTemplateRecord(ctx context.Context, tmpl *Template) error {
|
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.
|
// SaveTemplateRecord updates an existing template.
|
||||||
func SaveTemplateRecord(ctx context.Context, tmpl *Template) error {
|
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.
|
// DeleteTemplateRecord removes a template record.
|
||||||
func DeleteTemplateRecord(ctx context.Context, tmpl *Template) error {
|
func DeleteTemplateRecord(ctx context.Context, tmpl *Template) error {
|
||||||
return db.DB(ctx).Delete(tmpl).Error
|
return GetDB(ctx).Delete(tmpl).Error
|
||||||
}
|
}
|
||||||
|
|
||||||
// CreateScheduleRecord 创建定时任务
|
// CreateScheduleRecord 创建定时任务
|
||||||
func CreateScheduleRecord(ctx context.Context, schedule *Schedule) error {
|
func CreateScheduleRecord(ctx context.Context, schedule *Schedule) error {
|
||||||
return db.DB(ctx).Create(schedule).Error
|
return GetDB(ctx).Create(schedule).Error
|
||||||
}
|
}
|
||||||
|
|
||||||
// UpdateScheduleRecord 更新定时任务
|
// UpdateScheduleRecord 更新定时任务
|
||||||
func UpdateScheduleRecord(ctx context.Context, schedule *Schedule) error {
|
func UpdateScheduleRecord(ctx context.Context, schedule *Schedule) error {
|
||||||
return db.DB(ctx).Save(schedule).Error
|
return GetDB(ctx).Save(schedule).Error
|
||||||
}
|
}
|
||||||
|
|
||||||
// DeleteScheduleRecord 删除定时任务
|
// DeleteScheduleRecord 删除定时任务
|
||||||
func DeleteScheduleRecord(ctx context.Context, id uint64) error {
|
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 获取定时任务
|
// GetScheduleByID 根据 ID 获取定时任务
|
||||||
func GetScheduleByID(ctx context.Context, id uint64) (*Schedule, error) {
|
func GetScheduleByID(ctx context.Context, id uint64) (*Schedule, error) {
|
||||||
var schedule Schedule
|
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 nil, err
|
||||||
}
|
}
|
||||||
return &schedule, nil
|
return &schedule, nil
|
||||||
@@ -412,7 +409,7 @@ func GetScheduleByID(ctx context.Context, id uint64) (*Schedule, error) {
|
|||||||
// ListSchedulesRecord 获取所有定时任务
|
// ListSchedulesRecord 获取所有定时任务
|
||||||
func ListSchedulesRecord(ctx context.Context) ([]Schedule, error) {
|
func ListSchedulesRecord(ctx context.Context) ([]Schedule, error) {
|
||||||
var schedules []Schedule
|
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 nil, err
|
||||||
}
|
}
|
||||||
return schedules, nil
|
return schedules, nil
|
||||||
@@ -421,7 +418,7 @@ func ListSchedulesRecord(ctx context.Context) ([]Schedule, error) {
|
|||||||
// ListActiveSchedules 获取所有启用的定时任务
|
// ListActiveSchedules 获取所有启用的定时任务
|
||||||
func ListActiveSchedules(ctx context.Context) ([]Schedule, error) {
|
func ListActiveSchedules(ctx context.Context) ([]Schedule, error) {
|
||||||
var schedules []Schedule
|
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 nil, err
|
||||||
}
|
}
|
||||||
return schedules, nil
|
return schedules, nil
|
||||||
@@ -430,18 +427,18 @@ func ListActiveSchedules(ctx context.Context) ([]Schedule, error) {
|
|||||||
// CreateTaskExecutionRecord 创建任务执行记录
|
// CreateTaskExecutionRecord 创建任务执行记录
|
||||||
func CreateTaskExecutionRecord(ctx context.Context, execution *TaskExecution) error {
|
func CreateTaskExecutionRecord(ctx context.Context, execution *TaskExecution) error {
|
||||||
execution.ID = idgen.NextUint64ID()
|
execution.ID = idgen.NextUint64ID()
|
||||||
return db.DB(ctx).Create(execution).Error
|
return GetDB(ctx).Create(execution).Error
|
||||||
}
|
}
|
||||||
|
|
||||||
// UpdateTaskExecutionRecord 更新任务执行记录,忽略由 Redis 缓冲和归档流程管理的 log 字段。
|
// UpdateTaskExecutionRecord 更新任务执行记录,忽略由 Redis 缓冲和归档流程管理的 log 字段。
|
||||||
func UpdateTaskExecutionRecord(ctx context.Context, execution *TaskExecution) error {
|
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 获取执行记录
|
// GetTaskExecutionByTaskID 根据 TaskID 获取执行记录
|
||||||
func GetTaskExecutionByTaskID(ctx context.Context, taskID string) (*TaskExecution, error) {
|
func GetTaskExecutionByTaskID(ctx context.Context, taskID string) (*TaskExecution, error) {
|
||||||
var execution TaskExecution
|
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
|
return nil, err
|
||||||
}
|
}
|
||||||
if err := loadTaskExecutionLog(ctx, &execution); err != nil {
|
if err := loadTaskExecutionLog(ctx, &execution); err != nil {
|
||||||
@@ -453,7 +450,7 @@ func GetTaskExecutionByTaskID(ctx context.Context, taskID string) (*TaskExecutio
|
|||||||
// GetTaskExecutionByID 根据 ID 获取执行记录
|
// GetTaskExecutionByID 根据 ID 获取执行记录
|
||||||
func GetTaskExecutionByID(ctx context.Context, id uint64) (*TaskExecution, error) {
|
func GetTaskExecutionByID(ctx context.Context, id uint64) (*TaskExecution, error) {
|
||||||
var execution TaskExecution
|
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
|
return nil, err
|
||||||
}
|
}
|
||||||
if err := loadTaskExecutionLog(ctx, &execution); err != nil {
|
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.
|
// GetLatestTaskExecutionByTaskType returns the most recent execution for a task type.
|
||||||
func GetLatestTaskExecutionByTaskType(ctx context.Context, taskType string) (*TaskExecution, bool, error) {
|
func GetLatestTaskExecutionByTaskType(ctx context.Context, taskType string) (*TaskExecution, bool, error) {
|
||||||
var execution TaskExecution
|
var execution TaskExecution
|
||||||
err := db.DB(ctx).
|
err := GetDB(ctx).
|
||||||
Where("task_type = ?", taskType).
|
Where("task_type = ?", taskType).
|
||||||
Order("id DESC").
|
Order("id DESC").
|
||||||
First(&execution).Error
|
First(&execution).Error
|
||||||
@@ -481,45 +478,40 @@ func GetLatestTaskExecutionByTaskType(ctx context.Context, taskType string) (*Ta
|
|||||||
return nil, false, err
|
return nil, false, err
|
||||||
}
|
}
|
||||||
|
|
||||||
// AppendTaskExecutionLog 将日志追加到 Redis 缓冲,任务完成后再持久化到数据库。
|
// AppendTaskExecutionLog 将日志追加到缓冲,任务完成后再持久化到数据库。
|
||||||
func AppendTaskExecutionLog(ctx context.Context, taskID string, logLine string) error {
|
func AppendTaskExecutionLog(ctx context.Context, taskID string, logLine string) error {
|
||||||
if cachepkg.Redis == nil {
|
cacheSvc := GetCache(ctx)
|
||||||
return errors.New("redis client is not initialized")
|
if cacheSvc == nil {
|
||||||
|
return errors.New("cache service is not initialized")
|
||||||
}
|
}
|
||||||
|
|
||||||
now := time.Now().Format("15:04:05")
|
now := time.Now().Format("15:04:05")
|
||||||
line := fmt.Sprintf("[%s] %s\n", now, logLine)
|
line := fmt.Sprintf("[%s] %s\n", now, logLine)
|
||||||
key := taskExecutionLogRedisKey(taskID)
|
key := taskExecutionLogRedisKey(taskID)
|
||||||
|
|
||||||
_, err := cachepkg.Redis.TxPipelined(ctx, func(pipe redis.Pipeliner) error {
|
var existing string
|
||||||
pipe.RPush(ctx, key, line)
|
_ = cacheSvc.Get(ctx, key, &existing)
|
||||||
pipe.LTrim(ctx, key, -taskExecutionLogMaxLines, -1)
|
return cacheSvc.Set(ctx, key, existing+line, taskExecutionLogExpiration)
|
||||||
pipe.Expire(ctx, key, taskExecutionLogExpiration)
|
|
||||||
return nil
|
|
||||||
})
|
|
||||||
if err != nil {
|
|
||||||
return fmt.Errorf("append task execution log to redis: %w", err)
|
|
||||||
}
|
|
||||||
return nil
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// FlushTaskExecutionLog 将 Redis 中的完整任务日志写入数据库,并在成功后清理缓存。
|
// FlushTaskExecutionLog 将缓冲中的完整任务日志写入数据库,并在成功后清理缓存。
|
||||||
func FlushTaskExecutionLog(ctx context.Context, taskID string) error {
|
func FlushTaskExecutionLog(ctx context.Context, taskID string) error {
|
||||||
if cachepkg.Redis == nil {
|
cacheSvc := GetCache(ctx)
|
||||||
return errors.New("redis client is not initialized")
|
if cacheSvc == nil {
|
||||||
|
return errors.New("cache service is not initialized")
|
||||||
}
|
}
|
||||||
|
|
||||||
key := taskExecutionLogRedisKey(taskID)
|
key := taskExecutionLogRedisKey(taskID)
|
||||||
logLines, err := cachepkg.Redis.LRange(ctx, key, 0, -1).Result()
|
var logText string
|
||||||
if err != nil {
|
if err := cacheSvc.Get(ctx, key, &logText); err != nil || logText == "" {
|
||||||
return fmt.Errorf("get task execution log from redis: %w", err)
|
|
||||||
}
|
|
||||||
if len(logLines) == 0 {
|
|
||||||
return nil
|
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).
|
Where("task_id = ?", taskID).
|
||||||
Update("log", logText)
|
Update("log", logText)
|
||||||
if result.Error != nil {
|
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)
|
return fmt.Errorf("persist task execution log: task %q not found", taskID)
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := cachepkg.Redis.Del(ctx, key).Err(); err != nil {
|
_ = cacheSvc.Delete(ctx, key)
|
||||||
return fmt.Errorf("delete persisted task execution log from redis: %w", err)
|
|
||||||
}
|
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -544,7 +534,7 @@ func ListTaskExecutionRecords(ctx context.Context, req ListTaskExecutionsRequest
|
|||||||
req.PageSize = 20
|
req.PageSize = 20
|
||||||
}
|
}
|
||||||
|
|
||||||
query := db.DB(ctx).Model(&TaskExecution{})
|
query := GetDB(ctx).Model(&TaskExecution{})
|
||||||
|
|
||||||
if req.Status != "" {
|
if req.Status != "" {
|
||||||
query = query.Where("status = ?", 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}
|
terminalStatuses := []TaskExecutionStatus{TaskExecutionStatusSucceeded, TaskExecutionStatusFailed}
|
||||||
|
|
||||||
var highFrequencyTaskTypes []string
|
var highFrequencyTaskTypes []string
|
||||||
if err := db.DB(ctx).
|
if err := GetDB(ctx).
|
||||||
Model(&TaskExecution{}).
|
Model(&TaskExecution{}).
|
||||||
Select("task_type").
|
Select("task_type").
|
||||||
Where("created_at >= ?", frequencyWindowStart).
|
Where("created_at >= ?", frequencyWindowStart).
|
||||||
@@ -630,7 +620,7 @@ func CleanupTaskExecutionLogs(ctx context.Context, now time.Time) (TaskExecution
|
|||||||
|
|
||||||
var highFrequencyDeleted int64
|
var highFrequencyDeleted int64
|
||||||
if len(highFrequencyTaskTypes) > 0 {
|
if len(highFrequencyTaskTypes) > 0 {
|
||||||
highFrequencyResult := db.DB(ctx).
|
highFrequencyResult := GetDB(ctx).
|
||||||
Where("status IN ?", terminalStatuses).
|
Where("status IN ?", terminalStatuses).
|
||||||
Where("created_at < ?", highFrequencyCutoff).
|
Where("created_at < ?", highFrequencyCutoff).
|
||||||
Where("task_type IN ?", highFrequencyTaskTypes).
|
Where("task_type IN ?", highFrequencyTaskTypes).
|
||||||
@@ -641,7 +631,7 @@ func CleanupTaskExecutionLogs(ctx context.Context, now time.Time) (TaskExecution
|
|||||||
highFrequencyDeleted = highFrequencyResult.RowsAffected
|
highFrequencyDeleted = highFrequencyResult.RowsAffected
|
||||||
}
|
}
|
||||||
|
|
||||||
lowFrequencyQuery := db.DB(ctx).
|
lowFrequencyQuery := GetDB(ctx).
|
||||||
Where("status IN ?", terminalStatuses).
|
Where("status IN ?", terminalStatuses).
|
||||||
Where("created_at < ?", lowFrequencyCutoff)
|
Where("created_at < ?", lowFrequencyCutoff)
|
||||||
if len(highFrequencyTaskTypes) > 0 {
|
if len(highFrequencyTaskTypes) > 0 {
|
||||||
@@ -659,46 +649,32 @@ func CleanupTaskExecutionLogs(ctx context.Context, now time.Time) (TaskExecution
|
|||||||
}
|
}
|
||||||
|
|
||||||
func taskExecutionLogRedisKey(taskID string) string {
|
func taskExecutionLogRedisKey(taskID string) string {
|
||||||
return cachepkg.PrefixedKey(taskExecutionLogRedisKeyPrefix + taskID)
|
return taskExecutionLogRedisKeyPrefix + taskID
|
||||||
}
|
}
|
||||||
|
|
||||||
func loadTaskExecutionLog(ctx context.Context, execution *TaskExecution) error {
|
func loadTaskExecutionLog(ctx context.Context, execution *TaskExecution) error {
|
||||||
if cachepkg.Redis == nil {
|
cacheSvc := GetCache(ctx)
|
||||||
|
if cacheSvc == nil {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
logLines, err := cachepkg.Redis.LRange(ctx, taskExecutionLogRedisKey(execution.TaskID), 0, -1).Result()
|
var logText string
|
||||||
if err != nil {
|
if err := cacheSvc.Get(ctx, taskExecutionLogRedisKey(execution.TaskID), &logText); err == nil && logText != "" {
|
||||||
return fmt.Errorf("get task execution log from redis: %w", err)
|
execution.Log = logText
|
||||||
}
|
}
|
||||||
if len(logLines) == 0 {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
execution.Log = strings.Join(logLines, "")
|
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func loadTaskExecutionLogs(ctx context.Context, executions []TaskExecution) error {
|
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
|
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 {
|
for i := range executions {
|
||||||
logLines := commands[i].Val()
|
var logText string
|
||||||
if len(logLines) > 0 {
|
if err := cacheSvc.Get(ctx, taskExecutionLogRedisKey(executions[i].TaskID), &logText); err == nil && logText != "" {
|
||||||
executions[i].Log = strings.Join(logLines, "")
|
executions[i].Log = logText
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
return nil
|
return nil
|
||||||
|
|||||||
@@ -7,14 +7,11 @@ import (
|
|||||||
"context"
|
"context"
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
"errors"
|
"errors"
|
||||||
"sync"
|
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"gorm.io/gorm"
|
"gorm.io/gorm"
|
||||||
|
|
||||||
"Wavelet/pkg/cache/ram"
|
"Wavelet/pkg/cache/ram"
|
||||||
"Wavelet/pkg/util"
|
|
||||||
cachepkg "Wavelet/plugins/infra/cache"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
const (
|
const (
|
||||||
@@ -33,11 +30,6 @@ const (
|
|||||||
ConfigCacheType = "config"
|
ConfigCacheType = "config"
|
||||||
)
|
)
|
||||||
|
|
||||||
type systemConfigBroadcastMessage struct {
|
|
||||||
Type string `json:"type"`
|
|
||||||
Key string `json:"key"`
|
|
||||||
}
|
|
||||||
|
|
||||||
// ConfigLoader loads configuration data from the database.
|
// ConfigLoader loads configuration data from the database.
|
||||||
type ConfigLoader struct{}
|
type ConfigLoader struct{}
|
||||||
|
|
||||||
@@ -64,21 +56,19 @@ func (ConfigLoader) LoadAll(ctx context.Context, configType string) ([]ram.Cache
|
|||||||
return items, nil
|
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) {
|
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 err != nil {
|
||||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||||
return ram.CacheItem{}, ram.ErrNotFound
|
return ram.CacheItem{}, ram.ErrNotFound
|
||||||
}
|
}
|
||||||
return ram.CacheItem{}, err
|
return ram.CacheItem{}, err
|
||||||
}
|
}
|
||||||
|
|
||||||
valBytes, err := json.Marshal(cfg)
|
valBytes, err := json.Marshal(cfg)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return ram.CacheItem{}, err
|
return ram.CacheItem{}, err
|
||||||
}
|
}
|
||||||
|
|
||||||
return ram.CacheItem{
|
return ram.CacheItem{
|
||||||
Key: cfg.Key,
|
Key: cfg.Key,
|
||||||
Value: string(valBytes),
|
Value: string(valBytes),
|
||||||
@@ -87,121 +77,67 @@ func (ConfigLoader) LoadOne(ctx context.Context, configType string, key string)
|
|||||||
}, nil
|
}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// PreloadSystemConfigs warms the in-memory RAM cache from database on startup.
|
// GetCachedSystemConfig retrieves a single system config with RAM L1 fallback to DB.
|
||||||
func PreloadSystemConfigs(ctx context.Context) error {
|
func GetCachedSystemConfig(ctx context.Context, key string) (*SystemConfig, error) {
|
||||||
return ram.Refresh(ctx, ConfigCacheType, "", ConfigLoader{})
|
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 (
|
// StopSystemConfigCacheListener stops the cache invalidation listener (kept for backward compatibility).
|
||||||
systemConfigListenerOnce sync.Once
|
func StopSystemConfigCacheListener() {
|
||||||
systemConfigListenerCtx context.Context
|
}
|
||||||
systemConfigListenerCancel context.CancelFunc
|
|
||||||
systemConfigListenerDone chan struct{}
|
// StartSystemConfigCacheListener starts the cache listener (kept for backward compatibility).
|
||||||
)
|
func StartSystemConfigCacheListener() {
|
||||||
|
}
|
||||||
|
|
||||||
func ensureSystemConfigCacheListener() {
|
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 {
|
func determineTTL(_ string) time.Duration {
|
||||||
// Program-determined TTL: -1 means never expire for all configs by default
|
|
||||||
return -1
|
return -1
|
||||||
}
|
}
|
||||||
|
|
||||||
// InvalidateSystemConfigCache triggers a broadcast to refresh the cache for key.
|
// InvalidateSystemConfigCache triggers a broadcast to refresh the cache for key.
|
||||||
func InvalidateSystemConfigCache(ctx context.Context, key string) error {
|
func InvalidateSystemConfigCache(ctx context.Context, key string) error {
|
||||||
ensureSystemConfigCacheListener()
|
|
||||||
|
|
||||||
// Invalidate local cache synchronously first
|
|
||||||
ram.Delete(ConfigCacheType, key)
|
ram.Delete(ConfigCacheType, key)
|
||||||
|
if cacheSvc := GetCache(ctx); cacheSvc != nil {
|
||||||
// Broadcast to other nodes and clean legacy Redis cache key
|
_ = cacheSvc.Delete(ctx, "system:config:"+key)
|
||||||
if cachepkg.Redis != nil {
|
_ = cacheSvc.Delete(ctx, SystemConfigVisibleListRedisKey)
|
||||||
_ = cachepkg.HDel(ctx, SystemConfigRedisHashKey, key)
|
|
||||||
publishSystemConfigBroadcast(ctx, ConfigCacheType, key)
|
|
||||||
}
|
}
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// InvalidateAllSystemConfigCaches triggers a broadcast to refresh the entire config cache.
|
// InvalidateAllSystemConfigCaches triggers a broadcast to refresh the entire config cache.
|
||||||
func InvalidateAllSystemConfigCaches(ctx context.Context) error {
|
func InvalidateAllSystemConfigCaches(ctx context.Context) error {
|
||||||
ensureSystemConfigCacheListener()
|
|
||||||
|
|
||||||
// Invalidate all items of type ConfigCacheType synchronously first
|
|
||||||
ram.UpdateTypeItems(ConfigCacheType, nil)
|
ram.UpdateTypeItems(ConfigCacheType, nil)
|
||||||
|
if cacheSvc := GetCache(ctx); cacheSvc != nil {
|
||||||
// Broadcast to other nodes and clean legacy Redis cache keys
|
_ = cacheSvc.Delete(ctx, SystemConfigRedisHashKey)
|
||||||
if cachepkg.Redis != nil {
|
_ = cacheSvc.Delete(ctx, SystemConfigVisibleListRedisKey)
|
||||||
_ = cachepkg.Redis.Del(ctx, cachepkg.PrefixedKey(SystemConfigRedisHashKey), cachepkg.PrefixedKey(SystemConfigVisibleListRedisKey)).Err()
|
|
||||||
publishSystemConfigBroadcast(ctx, ConfigCacheType, "*")
|
|
||||||
}
|
}
|
||||||
return nil
|
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.
|
// ResetSystemConfigRAMCacheForTest clears only the process-local RAM cache.
|
||||||
func ResetSystemConfigRAMCacheForTest() {
|
func ResetSystemConfigRAMCacheForTest() {
|
||||||
ram.ResetForTest()
|
ram.ResetForTest()
|
||||||
|
|||||||
@@ -8,16 +8,30 @@ import (
|
|||||||
"testing"
|
"testing"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/alicebob/miniredis/v2"
|
|
||||||
"github.com/glebarez/sqlite"
|
"github.com/glebarez/sqlite"
|
||||||
"github.com/redis/go-redis/v9"
|
|
||||||
"github.com/redis/go-redis/v9/maintnotifications"
|
|
||||||
"gorm.io/gorm"
|
"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()) {
|
func setupSystemConfigTest(t *testing.T) (*gorm.DB, func()) {
|
||||||
t.Helper()
|
t.Helper()
|
||||||
|
|
||||||
@@ -41,28 +55,12 @@ func setupSystemConfigTest(t *testing.T) (*gorm.DB, func()) {
|
|||||||
t.Fatalf("Create(site_name) error = %v", err)
|
t.Fatalf("Create(site_name) error = %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
mr, err := miniredis.Run()
|
SetDBService(&testDBService{db: sqliteDB})
|
||||||
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
|
|
||||||
|
|
||||||
cleanup := func() {
|
cleanup := func() {
|
||||||
StopSystemConfigCacheListener()
|
StopSystemConfigCacheListener()
|
||||||
ResetSystemConfigRAMCacheForTest()
|
ResetSystemConfigRAMCacheForTest()
|
||||||
database.SetDB(nil)
|
ResetServices()
|
||||||
cache.Redis = previousRedis
|
|
||||||
_ = redisClient.Close()
|
|
||||||
mr.Close()
|
|
||||||
}
|
}
|
||||||
|
|
||||||
return sqliteDB, cleanup
|
return sqliteDB, cleanup
|
||||||
|
|||||||
@@ -7,9 +7,10 @@ import (
|
|||||||
"context"
|
"context"
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
|
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
|
||||||
"Wavelet/core/contracts"
|
"Wavelet/core/contracts"
|
||||||
"Wavelet/pkg/logger"
|
"Wavelet/pkg/logger"
|
||||||
"github.com/gin-gonic/gin"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
// LogForAudit 将登录鉴权审计日志写入 Logger
|
// LogForAudit 将登录鉴权审计日志写入 Logger
|
||||||
|
|||||||
@@ -11,7 +11,6 @@ import (
|
|||||||
"strings"
|
"strings"
|
||||||
|
|
||||||
"Wavelet/core/contracts"
|
"Wavelet/core/contracts"
|
||||||
db "Wavelet/plugins/infra/database"
|
|
||||||
|
|
||||||
"github.com/coreos/go-oidc/v3/oidc"
|
"github.com/coreos/go-oidc/v3/oidc"
|
||||||
"golang.org/x/oauth2"
|
"golang.org/x/oauth2"
|
||||||
@@ -19,7 +18,7 @@ import (
|
|||||||
|
|
||||||
func isOIDCLoginEnabled(ctx context.Context) bool {
|
func isOIDCLoginEnabled(ctx context.Context) bool {
|
||||||
var val string
|
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
|
return true
|
||||||
}
|
}
|
||||||
b, err := strconv.ParseBool(val)
|
b, err := strconv.ParseBool(val)
|
||||||
@@ -78,7 +77,7 @@ func activeLoginSources(ctx context.Context) []AuthSourceView {
|
|||||||
|
|
||||||
func getFrontendLoginRedirectURL(ctx context.Context) (string, error) {
|
func getFrontendLoginRedirectURL(ctx context.Context) (string, error) {
|
||||||
var val string
|
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 "", errors.New(errServerAddressMissing)
|
||||||
}
|
}
|
||||||
return strings.TrimRight(val, "/") + "/login", nil
|
return strings.TrimRight(val, "/") + "/login", nil
|
||||||
|
|||||||
@@ -6,23 +6,15 @@ package auth
|
|||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"fmt"
|
"fmt"
|
||||||
"strconv"
|
|
||||||
"sync"
|
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"Wavelet/core/contracts"
|
"Wavelet/core/contracts"
|
||||||
"Wavelet/pkg/cache/ram"
|
"Wavelet/pkg/cache/ram"
|
||||||
"Wavelet/pkg/util"
|
|
||||||
db "Wavelet/plugins/infra/cache"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
const (
|
const (
|
||||||
tokenCacheTTL = 5 * time.Minute
|
tokenCacheTTL = 5 * time.Minute
|
||||||
userCacheTTL = 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.
|
// CachedToken represents the minimal cached representation of an access token.
|
||||||
@@ -35,16 +27,6 @@ type CachedToken struct {
|
|||||||
var (
|
var (
|
||||||
tokenRAM = ram.MustNew[string, *CachedToken](ram.Options{MaximumSize: 2048})
|
tokenRAM = ram.MustNew[string, *CachedToken](ram.Options{MaximumSize: 2048})
|
||||||
userRAM = ram.MustNew[uint64, *contracts.UserDTO](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 {
|
func tokenCacheKey(tokenHash string) string {
|
||||||
@@ -55,107 +37,16 @@ func userCacheKey(userID uint64) string {
|
|||||||
return fmt.Sprintf("oauth:user:%d", userID)
|
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
|
// GetCachedToken 获取缓存的 Token
|
||||||
func GetCachedToken(ctx context.Context, tokenHash string) (*CachedToken, error) {
|
func GetCachedToken(ctx context.Context, tokenHash string) (*CachedToken, error) {
|
||||||
ensureTokenCacheListener()
|
|
||||||
|
|
||||||
if val, ok := tokenRAM.GetIfPresent(tokenHash); ok {
|
if val, ok := tokenRAM.GetIfPresent(tokenHash); ok {
|
||||||
return val, nil
|
return val, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
if db.Redis != nil {
|
if cache := getCache(ctx); cache != nil {
|
||||||
var token CachedToken
|
var token CachedToken
|
||||||
key := tokenCacheKey(tokenHash)
|
key := tokenCacheKey(tokenHash)
|
||||||
if err := db.GetJSON(ctx, key, &token); err == nil {
|
if err := cache.Get(ctx, key, &token); err == nil {
|
||||||
// Write back to local cache
|
|
||||||
tokenRAM.Set(tokenHash, &token)
|
tokenRAM.Set(tokenHash, &token)
|
||||||
return &token, nil
|
return &token, nil
|
||||||
}
|
}
|
||||||
@@ -165,40 +56,32 @@ func GetCachedToken(ctx context.Context, tokenHash string) (*CachedToken, error)
|
|||||||
|
|
||||||
// SetCachedToken 设置 Token 缓存
|
// SetCachedToken 设置 Token 缓存
|
||||||
func SetCachedToken(ctx context.Context, tokenHash string, token *CachedToken) {
|
func SetCachedToken(ctx context.Context, tokenHash string, token *CachedToken) {
|
||||||
ensureTokenCacheListener()
|
|
||||||
|
|
||||||
tokenRAM.Set(tokenHash, token)
|
tokenRAM.Set(tokenHash, token)
|
||||||
if db.Redis != nil {
|
if cache := getCache(ctx); cache != nil {
|
||||||
key := tokenCacheKey(tokenHash)
|
key := tokenCacheKey(tokenHash)
|
||||||
_ = db.SetJSON(ctx, key, token, tokenCacheTTL)
|
_ = cache.Set(ctx, key, token, tokenCacheTTL)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// InvalidateCachedToken 吊销/删除 token 缓存
|
// InvalidateCachedToken 吊销/删除 token 缓存
|
||||||
func InvalidateCachedToken(ctx context.Context, tokenHash string) {
|
func InvalidateCachedToken(ctx context.Context, tokenHash string) {
|
||||||
ensureTokenCacheListener()
|
|
||||||
|
|
||||||
tokenRAM.Invalidate(tokenHash)
|
tokenRAM.Invalidate(tokenHash)
|
||||||
if db.Redis != nil {
|
if cache := getCache(ctx); cache != nil {
|
||||||
key := tokenCacheKey(tokenHash)
|
key := tokenCacheKey(tokenHash)
|
||||||
_ = db.Redis.Del(ctx, db.PrefixedKey(key)).Err()
|
_ = cache.Delete(ctx, key)
|
||||||
publishTokenRAMInvalidation(ctx, tokenHash)
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// GetCachedUser 获取缓存的 UserDTO
|
// GetCachedUser 获取缓存的 UserDTO
|
||||||
func GetCachedUser(ctx context.Context, userID uint64) (*contracts.UserDTO, error) {
|
func GetCachedUser(ctx context.Context, userID uint64) (*contracts.UserDTO, error) {
|
||||||
ensureUserCacheListener()
|
|
||||||
|
|
||||||
if val, ok := userRAM.GetIfPresent(userID); ok {
|
if val, ok := userRAM.GetIfPresent(userID); ok {
|
||||||
return val, nil
|
return val, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
if db.Redis != nil {
|
if cache := getCache(ctx); cache != nil {
|
||||||
var u contracts.UserDTO
|
var u contracts.UserDTO
|
||||||
key := userCacheKey(userID)
|
key := userCacheKey(userID)
|
||||||
if err := db.GetJSON(ctx, key, &u); err == nil {
|
if err := cache.Get(ctx, key, &u); err == nil {
|
||||||
// Write back to local cache
|
|
||||||
userRAM.Set(userID, &u)
|
userRAM.Set(userID, &u)
|
||||||
return &u, nil
|
return &u, nil
|
||||||
}
|
}
|
||||||
@@ -208,49 +91,24 @@ func GetCachedUser(ctx context.Context, userID uint64) (*contracts.UserDTO, erro
|
|||||||
|
|
||||||
// SetCachedUser 设置 UserDTO 缓存
|
// SetCachedUser 设置 UserDTO 缓存
|
||||||
func SetCachedUser(ctx context.Context, userID uint64, u *contracts.UserDTO) {
|
func SetCachedUser(ctx context.Context, userID uint64, u *contracts.UserDTO) {
|
||||||
ensureUserCacheListener()
|
|
||||||
|
|
||||||
userRAM.Set(userID, u)
|
userRAM.Set(userID, u)
|
||||||
if db.Redis != nil {
|
if cache := getCache(ctx); cache != nil {
|
||||||
key := userCacheKey(userID)
|
key := userCacheKey(userID)
|
||||||
_ = db.SetJSON(ctx, key, u, userCacheTTL)
|
_ = cache.Set(ctx, key, u, userCacheTTL)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// InvalidateCachedUser 吊销/失效 UserDTO 缓存
|
// InvalidateCachedUser 吊销/失效 UserDTO 缓存
|
||||||
func InvalidateCachedUser(ctx context.Context, userID uint64) {
|
func InvalidateCachedUser(ctx context.Context, userID uint64) {
|
||||||
ensureUserCacheListener()
|
|
||||||
|
|
||||||
userRAM.Invalidate(userID)
|
userRAM.Invalidate(userID)
|
||||||
if db.Redis != nil {
|
if cache := getCache(ctx); cache != nil {
|
||||||
key := userCacheKey(userID)
|
key := userCacheKey(userID)
|
||||||
_ = db.Redis.Del(ctx, db.PrefixedKey(key)).Err()
|
_ = cache.Delete(ctx, key)
|
||||||
publishUserRAMInvalidation(ctx, userID)
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// StopAuthCacheListener stops both token and user Redis Pub/Sub subscription listeners and resets the sync.Once guards.
|
// StopAuthCacheListener compatibility stub for tests
|
||||||
func StopAuthCacheListener() {
|
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{}
|
|
||||||
}
|
|
||||||
|
|
||||||
// ResetAuthRAMCacheForTest clears only the process-local RAM cache.
|
// ResetAuthRAMCacheForTest clears only the process-local RAM cache.
|
||||||
func ResetAuthRAMCacheForTest() {
|
func ResetAuthRAMCacheForTest() {
|
||||||
|
|||||||
@@ -5,48 +5,69 @@ package auth_test
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
|
"encoding/json"
|
||||||
"testing"
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
"github.com/alicebob/miniredis/v2"
|
"Wavelet/core"
|
||||||
"github.com/redis/go-redis/v9"
|
|
||||||
"github.com/redis/go-redis/v9/maintnotifications"
|
|
||||||
|
|
||||||
"Wavelet/core/contracts"
|
"Wavelet/core/contracts"
|
||||||
"Wavelet/plugins/domain/auth"
|
"Wavelet/plugins/domain/auth"
|
||||||
db "Wavelet/plugins/infra/cache"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
func setupOauthCacheTest(t *testing.T) (*miniredis.Miniredis, func()) {
|
type mockCacheService struct {
|
||||||
t.Helper()
|
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 {
|
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{
|
func (m *mockCacheService) Delete(ctx context.Context, key string) error {
|
||||||
Addr: miniRedis.Addr(),
|
delete(m.items, key)
|
||||||
MaintNotificationsConfig: &maintnotifications.Config{
|
return nil
|
||||||
Mode: maintnotifications.ModeDisabled,
|
}
|
||||||
},
|
|
||||||
})
|
|
||||||
|
|
||||||
auth.ResetAuthRAMCacheForTest()
|
func (m *mockCacheService) Invalidate(ctx context.Context, key string) error {
|
||||||
|
return m.Delete(ctx, key)
|
||||||
|
}
|
||||||
|
|
||||||
cleanup := func() {
|
func (m *mockCacheService) GetOrSet(ctx context.Context, key string, target any, ttl time.Duration, loader func() (any, error)) error {
|
||||||
auth.StopAuthCacheListener()
|
err := m.Get(ctx, key, target)
|
||||||
auth.ResetAuthRAMCacheForTest()
|
if err == nil {
|
||||||
_ = db.Redis.Close()
|
return nil
|
||||||
miniRedis.Close()
|
|
||||||
db.Redis = 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) {
|
func TestTokenCache_GetSetInvalidate(t *testing.T) {
|
||||||
_, cleanup := setupOauthCacheTest(t)
|
ctx := core.NewContext(context.Background())
|
||||||
defer cleanup()
|
mockCache := newMockCacheService()
|
||||||
ctx := context.Background()
|
core.Provide[contracts.CacheService](ctx, mockCache)
|
||||||
|
|
||||||
tokenHash := "test-token-hash"
|
tokenHash := "test-token-hash"
|
||||||
token := &auth.CachedToken{
|
token := &auth.CachedToken{
|
||||||
@@ -84,9 +105,9 @@ func TestTokenCache_GetSetInvalidate(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestUserCache_GetSetInvalidate(t *testing.T) {
|
func TestUserCache_GetSetInvalidate(t *testing.T) {
|
||||||
_, cleanup := setupOauthCacheTest(t)
|
ctx := core.NewContext(context.Background())
|
||||||
defer cleanup()
|
mockCache := newMockCacheService()
|
||||||
ctx := context.Background()
|
core.Provide[contracts.CacheService](ctx, mockCache)
|
||||||
|
|
||||||
userID := uint64(789)
|
userID := uint64(789)
|
||||||
user := &contracts.UserDTO{
|
user := &contracts.UserDTO{
|
||||||
|
|||||||
@@ -14,17 +14,16 @@ import (
|
|||||||
|
|
||||||
"Wavelet/core/contracts"
|
"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/coreos/go-oidc/v3/oidc"
|
||||||
"github.com/gin-contrib/sessions"
|
"github.com/gin-contrib/sessions"
|
||||||
"github.com/gin-gonic/gin"
|
"github.com/gin-gonic/gin"
|
||||||
"github.com/google/uuid"
|
"github.com/google/uuid"
|
||||||
"gorm.io/gorm"
|
"gorm.io/gorm"
|
||||||
|
|
||||||
|
"Wavelet/pkg/idgen"
|
||||||
|
"Wavelet/pkg/logger"
|
||||||
|
"Wavelet/pkg/response"
|
||||||
|
"Wavelet/pkg/util"
|
||||||
)
|
)
|
||||||
|
|
||||||
// GetLoginSources 获取可用登录源列表
|
// GetLoginSources 获取可用登录源列表
|
||||||
@@ -78,9 +77,12 @@ func GetLoginURL(c *gin.Context) {
|
|||||||
response.AbortInternal(c, err.Error())
|
response.AbortInternal(c, err.Error())
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
if err := cachepkg.Redis.Set(ctx, cachepkg.PrefixedKey(fmt.Sprintf(OAuthStateCacheKeyFormat, state)), payloadValue, OAuthStateCacheKeyExpiration).Err(); err != nil {
|
stateKey := fmt.Sprintf(OAuthStateCacheKeyFormat, state)
|
||||||
response.AbortInternal(c, err.Error())
|
if cache := getCache(ctx); cache != nil {
|
||||||
return
|
if err := cache.Set(ctx, stateKey, payloadValue, OAuthStateCacheKeyExpiration); err != nil {
|
||||||
|
response.AbortInternal(c, err.Error())
|
||||||
|
return
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
authorizeURL, err := buildAuthorizeURL(c.Request.Context(), source, state)
|
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 {
|
func reserveOAuthStateSlot(ctx context.Context, sessionHash string) error {
|
||||||
if cachepkg.Redis == nil || sessionHash == "" {
|
if sessionHash == "" {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
key := cachepkg.PrefixedKey(fmt.Sprintf(oauthStateLimitKeyFormat, sessionHash))
|
cache := getCache(ctx)
|
||||||
n, err := cachepkg.Redis.Incr(ctx, key).Result()
|
if cache == nil {
|
||||||
if err != nil {
|
return nil
|
||||||
return err
|
|
||||||
}
|
}
|
||||||
if n == 1 {
|
key := fmt.Sprintf(oauthStateLimitKeyFormat, sessionHash)
|
||||||
_ = cachepkg.Redis.Expire(ctx, key, OAuthStateCacheKeyExpiration).Err()
|
var count int
|
||||||
}
|
_ = cache.Get(ctx, key, &count)
|
||||||
if n > oauthStateLimitMax {
|
count++
|
||||||
|
_ = cache.Set(ctx, key, count, OAuthStateCacheKeyExpiration)
|
||||||
|
if count > oauthStateLimitMax {
|
||||||
return errors.New(errOAuthStateRateLimited)
|
return errors.New(errOAuthStateRateLimited)
|
||||||
}
|
}
|
||||||
return nil
|
return nil
|
||||||
@@ -179,9 +182,12 @@ func Authorize(c *gin.Context) {
|
|||||||
response.AbortInternal(c, err.Error())
|
response.AbortInternal(c, err.Error())
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
if err := cachepkg.Redis.Set(ctx, cachepkg.PrefixedKey(fmt.Sprintf(OAuthStateCacheKeyFormat, state)), payloadValue, OAuthStateCacheKeyExpiration).Err(); err != nil {
|
stateKey := fmt.Sprintf(OAuthStateCacheKeyFormat, state)
|
||||||
response.AbortInternal(c, err.Error())
|
if cache := getCache(ctx); cache != nil {
|
||||||
return
|
if err := cache.Set(ctx, stateKey, payloadValue, OAuthStateCacheKeyExpiration); err != nil {
|
||||||
|
response.AbortInternal(c, err.Error())
|
||||||
|
return
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
authorizeURL, err := buildAuthorizeURL(c.Request.Context(), source, state)
|
authorizeURL, err := buildAuthorizeURL(c.Request.Context(), source, state)
|
||||||
@@ -201,13 +207,18 @@ func Callback(c *gin.Context) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
ctx := c.Request.Context()
|
ctx := c.Request.Context()
|
||||||
stateKey := cachepkg.PrefixedKey(fmt.Sprintf(OAuthStateCacheKeyFormat, req.State))
|
stateKey := fmt.Sprintf(OAuthStateCacheKeyFormat, req.State)
|
||||||
payloadRaw, err := cachepkg.Redis.Get(ctx, stateKey).Result()
|
var payloadRaw string
|
||||||
if err != nil {
|
cache := getCache(ctx)
|
||||||
|
if cache == nil {
|
||||||
response.AbortBadRequest(c, errInvalidState)
|
response.AbortBadRequest(c, errInvalidState)
|
||||||
return
|
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)
|
payload, err := decodeOAuthStatePayload(payloadRaw)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -289,7 +300,7 @@ func handleCallbackBind(ctx context.Context, c *gin.Context, source *AuthSource,
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
var user contracts.UserDTO
|
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())
|
response.AbortInternal(c, err.Error())
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
@@ -304,7 +315,7 @@ func handleCallbackBind(ctx context.Context, c *gin.Context, source *AuthSource,
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
user.LastLoginAt = time.Now()
|
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")))
|
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)
|
account, err := FindExternalAccount(ctx, source.ID, userInfo.Sub)
|
||||||
switch {
|
switch {
|
||||||
case err == nil:
|
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())
|
response.AbortInternal(c, loadErr.Error())
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
@@ -330,7 +341,7 @@ func handleCallbackLogin(ctx context.Context, c *gin.Context, source *AuthSource
|
|||||||
}
|
}
|
||||||
|
|
||||||
user.LastLoginAt = time.Now()
|
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 {
|
if err := SetLoginSession(ctx, c, &user); err != nil {
|
||||||
response.AbortInternal(c, err.Error())
|
response.AbortInternal(c, err.Error())
|
||||||
return
|
return
|
||||||
@@ -348,7 +359,7 @@ func uniqueUsername(ctx context.Context, base string) (string, error) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
var existingUsernames []string
|
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)+"-%").
|
Where("username = ? OR username LIKE ? ESCAPE '\\'", base, util.EscapeLike(base)+"-%").
|
||||||
Pluck("username", &existingUsernames).Error; err != nil {
|
Pluck("username", &existingUsernames).Error; err != nil {
|
||||||
return "", err
|
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) {
|
func handleCallbackRegister(ctx context.Context, c *gin.Context, source *AuthSource, userInfo *contracts.OAuthUserInfoDTO) (contracts.UserDTO, bool) {
|
||||||
registrationEnabled := true
|
registrationEnabled := true
|
||||||
var val string
|
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 {
|
if b, err := strconv.ParseBool(val); err == nil {
|
||||||
registrationEnabled = b
|
registrationEnabled = b
|
||||||
}
|
}
|
||||||
@@ -407,7 +418,7 @@ func handleCallbackRegister(ctx context.Context, c *gin.Context, source *AuthSou
|
|||||||
UpdatedAt: now,
|
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())
|
response.AbortInternal(c, err.Error())
|
||||||
return contracts.UserDTO{}, false
|
return contracts.UserDTO{}, false
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -9,12 +9,12 @@ import (
|
|||||||
"encoding/hex"
|
"encoding/hex"
|
||||||
"errors"
|
"errors"
|
||||||
|
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
|
||||||
"Wavelet/core/contracts"
|
"Wavelet/core/contracts"
|
||||||
"Wavelet/pkg/response"
|
"Wavelet/pkg/response"
|
||||||
"Wavelet/pkg/trace"
|
"Wavelet/pkg/trace"
|
||||||
"Wavelet/pkg/util"
|
"Wavelet/pkg/util"
|
||||||
db "Wavelet/plugins/infra/database"
|
|
||||||
"github.com/gin-gonic/gin"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
func hashToken(token string) string {
|
func hashToken(token string) string {
|
||||||
@@ -38,7 +38,7 @@ func getUserByToken(ctx context.Context, tokenStr string) (*contracts.UserDTO, *
|
|||||||
UserID uint64
|
UserID uint64
|
||||||
IsAdmin bool
|
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
|
return nil, nil, err
|
||||||
}
|
}
|
||||||
tokenRecord = &CachedToken{
|
tokenRecord = &CachedToken{
|
||||||
@@ -49,7 +49,7 @@ func getUserByToken(ctx context.Context, tokenStr string) (*contracts.UserDTO, *
|
|||||||
SetCachedToken(ctx, tokenHash, tokenRecord)
|
SetCachedToken(ctx, tokenHash, tokenRecord)
|
||||||
|
|
||||||
var userRow contracts.UserDTO
|
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
|
return nil, nil, err
|
||||||
}
|
}
|
||||||
SetCachedUser(ctx, userRow.ID, &userRow)
|
SetCachedUser(ctx, userRow.ID, &userRow)
|
||||||
@@ -90,7 +90,7 @@ func GetUserFromRequest(c *gin.Context) (*contracts.UserDTO, error) {
|
|||||||
user, err := GetCachedUser(ctx, userID)
|
user, err := GetCachedUser(ctx, userID)
|
||||||
if err != nil || user == nil || !user.IsActive {
|
if err != nil || user == nil || !user.IsActive {
|
||||||
var dbUser contracts.UserDTO
|
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
|
return nil, err
|
||||||
}
|
}
|
||||||
user = &dbUser
|
user = &dbUser
|
||||||
|
|||||||
@@ -76,6 +76,27 @@ func (p *Plugin) Manifest() core.Manifest {
|
|||||||
|
|
||||||
// Apply registers the auth migrations, services, routes, and settings into the Context.
|
// Apply registers the auth migrations, services, routes, and settings into the Context.
|
||||||
func (p *Plugin) Apply(ctx *core.Context) error {
|
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
|
// 1. Register migrations
|
||||||
ctx.Migrations().Register("auth", authMigrations)
|
ctx.Migrations().Register("auth", authMigrations)
|
||||||
|
|
||||||
|
|||||||
@@ -19,9 +19,24 @@ import (
|
|||||||
"Wavelet/core"
|
"Wavelet/core"
|
||||||
"Wavelet/core/contracts"
|
"Wavelet/core/contracts"
|
||||||
"Wavelet/plugins/domain/auth"
|
"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 {
|
type testUser struct {
|
||||||
ID uint64 `gorm:"primaryKey"`
|
ID uint64 `gorm:"primaryKey"`
|
||||||
Username string
|
Username string
|
||||||
@@ -60,7 +75,6 @@ func setupTestDB(t *testing.T) *gorm.DB {
|
|||||||
&auth.ExternalAccount{},
|
&auth.ExternalAccount{},
|
||||||
))
|
))
|
||||||
|
|
||||||
db.SetDB(testDB)
|
|
||||||
return testDB
|
return testDB
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -81,6 +95,7 @@ func (m *mockProvider) ExchangeCode(ctx context.Context, code string) (*contract
|
|||||||
func TestAuthPluginUnit(t *testing.T) {
|
func TestAuthPluginUnit(t *testing.T) {
|
||||||
ctx := core.NewContext(context.Background())
|
ctx := core.NewContext(context.Background())
|
||||||
testDB := setupTestDB(t)
|
testDB := setupTestDB(t)
|
||||||
|
core.Provide[contracts.DBService](ctx, &mockDBService{db: testDB})
|
||||||
|
|
||||||
p := auth.New()
|
p := auth.New()
|
||||||
assert.Equal(t, "auth", p.Name())
|
assert.Equal(t, "auth", p.Name())
|
||||||
|
|||||||
@@ -5,14 +5,64 @@ package auth
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"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 获取认证源
|
// GetAuthSourceByID 根据 ID 获取认证源
|
||||||
func GetAuthSourceByID(ctx context.Context, id uint64) (*AuthSource, error) {
|
func GetAuthSourceByID(ctx context.Context, id uint64) (*AuthSource, error) {
|
||||||
var src AuthSource
|
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 nil, err
|
||||||
}
|
}
|
||||||
return &src, nil
|
return &src, nil
|
||||||
@@ -21,7 +71,7 @@ func GetAuthSourceByID(ctx context.Context, id uint64) (*AuthSource, error) {
|
|||||||
// GetAuthSourceByName 根据名称获取认证源
|
// GetAuthSourceByName 根据名称获取认证源
|
||||||
func GetAuthSourceByName(ctx context.Context, name string) (*AuthSource, error) {
|
func GetAuthSourceByName(ctx context.Context, name string) (*AuthSource, error) {
|
||||||
var src AuthSource
|
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 nil, err
|
||||||
}
|
}
|
||||||
return &src, nil
|
return &src, nil
|
||||||
@@ -30,7 +80,7 @@ func GetAuthSourceByName(ctx context.Context, name string) (*AuthSource, error)
|
|||||||
// ListActiveAuthSources 获取所有启用的认证源
|
// ListActiveAuthSources 获取所有启用的认证源
|
||||||
func ListActiveAuthSources(ctx context.Context) ([]AuthSource, error) {
|
func ListActiveAuthSources(ctx context.Context) ([]AuthSource, error) {
|
||||||
var sources []AuthSource
|
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 nil, err
|
||||||
}
|
}
|
||||||
return sources, nil
|
return sources, nil
|
||||||
@@ -49,7 +99,7 @@ func GetAuthSourceByNameCached(ctx context.Context, name string) (*AuthSource, e
|
|||||||
// FindExternalAccount 查询指定认证源的外部账号绑定
|
// FindExternalAccount 查询指定认证源的外部账号绑定
|
||||||
func FindExternalAccount(ctx context.Context, authSourceID uint64, externalID string) (*ExternalAccount, error) {
|
func FindExternalAccount(ctx context.Context, authSourceID uint64, externalID string) (*ExternalAccount, error) {
|
||||||
var account ExternalAccount
|
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 nil, err
|
||||||
}
|
}
|
||||||
return &account, nil
|
return &account, nil
|
||||||
@@ -57,13 +107,13 @@ func FindExternalAccount(ctx context.Context, authSourceID uint64, externalID st
|
|||||||
|
|
||||||
// BindExternalAccount 绑定外部账号
|
// BindExternalAccount 绑定外部账号
|
||||||
func BindExternalAccount(ctx context.Context, account *ExternalAccount) error {
|
func BindExternalAccount(ctx context.Context, account *ExternalAccount) error {
|
||||||
return db.DB(ctx).Create(account).Error
|
return getDB(ctx).Create(account).Error
|
||||||
}
|
}
|
||||||
|
|
||||||
// ListExternalAccountsByUserID 获取用户绑定的所有外部账号
|
// ListExternalAccountsByUserID 获取用户绑定的所有外部账号
|
||||||
func ListExternalAccountsByUserID(ctx context.Context, userID uint64) ([]ExternalAccount, error) {
|
func ListExternalAccountsByUserID(ctx context.Context, userID uint64) ([]ExternalAccount, error) {
|
||||||
var accounts []ExternalAccount
|
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 nil, err
|
||||||
}
|
}
|
||||||
return accounts, nil
|
return accounts, nil
|
||||||
@@ -71,5 +121,5 @@ func ListExternalAccountsByUserID(ctx context.Context, userID uint64) ([]Externa
|
|||||||
|
|
||||||
// UnbindExternalAccount 解绑外部账号
|
// UnbindExternalAccount 解绑外部账号
|
||||||
func UnbindExternalAccount(ctx context.Context, id uint64, userID uint64) error {
|
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"
|
"errors"
|
||||||
"sync"
|
"sync"
|
||||||
|
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
|
||||||
"Wavelet/core/contracts"
|
"Wavelet/core/contracts"
|
||||||
"Wavelet/pkg/util"
|
"Wavelet/pkg/util"
|
||||||
db "Wavelet/plugins/infra/database"
|
|
||||||
"github.com/gin-gonic/gin"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
type authServiceImpl struct{}
|
type authServiceImpl struct{}
|
||||||
@@ -57,7 +57,7 @@ func (s *authServiceImpl) VerifyToken(ctx context.Context, token string) (*contr
|
|||||||
UserID uint64
|
UserID uint64
|
||||||
IsAdmin bool
|
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
|
return nil, err
|
||||||
}
|
}
|
||||||
tokenRecord = &CachedToken{
|
tokenRecord = &CachedToken{
|
||||||
@@ -71,7 +71,7 @@ func (s *authServiceImpl) VerifyToken(ctx context.Context, token string) (*contr
|
|||||||
user, err := GetCachedUser(ctx, tokenRecord.UserID)
|
user, err := GetCachedUser(ctx, tokenRecord.UserID)
|
||||||
if err != nil || user == nil || !user.IsActive {
|
if err != nil || user == nil || !user.IsActive {
|
||||||
var dbUser contracts.UserDTO
|
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
|
return nil, err
|
||||||
}
|
}
|
||||||
user = &dbUser
|
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) {
|
func (s *authServiceImpl) ListAuthSources(ctx context.Context) ([]contracts.AuthSourceViewDTO, error) {
|
||||||
var sources []AuthSource
|
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
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -157,7 +157,7 @@ func (s *authServiceImpl) CreateAuthSource(ctx context.Context, source contracts
|
|||||||
return nil, err
|
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
|
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) {
|
func (s *authServiceImpl) UpdateAuthSource(ctx context.Context, id uint64, source contracts.AuthSourceDTO) (*contracts.AuthSourceDTO, error) {
|
||||||
var existing AuthSource
|
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
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -184,7 +184,7 @@ func (s *authServiceImpl) UpdateAuthSource(ctx context.Context, id uint64, sourc
|
|||||||
return nil, err
|
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
|
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 {
|
func (s *authServiceImpl) DeleteAuthSource(ctx context.Context, id uint64) error {
|
||||||
var existing AuthSource
|
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 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) {
|
func (s *authServiceImpl) ToggleAuthSource(ctx context.Context, id uint64) (*contracts.AuthSourceDTO, error) {
|
||||||
var existing AuthSource
|
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
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
existing.IsActive = !existing.IsActive
|
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
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -11,13 +11,13 @@ import (
|
|||||||
"strconv"
|
"strconv"
|
||||||
"strings"
|
"strings"
|
||||||
|
|
||||||
"Wavelet/core/contracts"
|
|
||||||
"Wavelet/pkg/config"
|
|
||||||
db "Wavelet/plugins/infra/database"
|
|
||||||
"github.com/gin-contrib/sessions"
|
"github.com/gin-contrib/sessions"
|
||||||
"github.com/gin-gonic/gin"
|
"github.com/gin-gonic/gin"
|
||||||
"github.com/google/uuid"
|
"github.com/google/uuid"
|
||||||
gsessions "github.com/gorilla/sessions"
|
gsessions "github.com/gorilla/sessions"
|
||||||
|
|
||||||
|
"Wavelet/core/contracts"
|
||||||
|
"Wavelet/pkg/config"
|
||||||
)
|
)
|
||||||
|
|
||||||
// GetSessionOptions 根据配置构建 Session 选项
|
// GetSessionOptions 根据配置构建 Session 选项
|
||||||
@@ -118,7 +118,7 @@ func SetLoginSession(ctx context.Context, c *gin.Context, user *contracts.UserDT
|
|||||||
isSessionCookie := false
|
isSessionCookie := false
|
||||||
|
|
||||||
var val string
|
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 {
|
if ttlHours, err := strconv.Atoi(val); err == nil {
|
||||||
switch {
|
switch {
|
||||||
case ttlHours == -1:
|
case ttlHours == -1:
|
||||||
|
|||||||
@@ -6,10 +6,11 @@ package cap
|
|||||||
import (
|
import (
|
||||||
"net/http"
|
"net/http"
|
||||||
|
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
|
||||||
"Wavelet/pkg/logger"
|
"Wavelet/pkg/logger"
|
||||||
"Wavelet/pkg/response"
|
"Wavelet/pkg/response"
|
||||||
"Wavelet/plugins/domain/cap/pow"
|
"Wavelet/plugins/domain/cap/pow"
|
||||||
"github.com/gin-gonic/gin"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
// ChallengeResponse is a local type alias for the pow.ChallengeResponse struct
|
// ChallengeResponse is a local type alias for the pow.ChallengeResponse struct
|
||||||
|
|||||||
@@ -15,7 +15,6 @@ import (
|
|||||||
|
|
||||||
"Wavelet/pkg/config"
|
"Wavelet/pkg/config"
|
||||||
"Wavelet/plugins/domain/cap/pow"
|
"Wavelet/plugins/domain/cap/pow"
|
||||||
db "Wavelet/plugins/infra/cache"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
const (
|
const (
|
||||||
@@ -186,13 +185,7 @@ func GetDefaultManager() *Manager {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
var store pow.Store
|
store := pow.NewMemoryStore(1 * time.Minute)
|
||||||
if config.Config != nil && config.Config.Redis.Enabled && db.Redis != nil {
|
|
||||||
store = pow.NewRedisStore(db.Redis)
|
|
||||||
} else {
|
|
||||||
store = pow.NewMemoryStore(1 * time.Minute)
|
|
||||||
}
|
|
||||||
|
|
||||||
defaultManager = NewManager(secret, store)
|
defaultManager = NewManager(secret, store)
|
||||||
})
|
})
|
||||||
return defaultManager
|
return defaultManager
|
||||||
|
|||||||
@@ -29,7 +29,6 @@ func (p *Plugin) Name() string {
|
|||||||
func (p *Plugin) Inject() []reflect.Type {
|
func (p *Plugin) Inject() []reflect.Type {
|
||||||
return []reflect.Type{
|
return []reflect.Type{
|
||||||
reflect.TypeFor[contracts.DBService](),
|
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.
|
// Apply registers the cap routes and settings into the Context.
|
||||||
func (p *Plugin) Apply(ctx *core.Context) error {
|
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
|
// Register HTTP Routes
|
||||||
capGroup := ctx.Router().Group("/api/v1/cap")
|
capGroup := ctx.Router().Group("/api/v1/cap")
|
||||||
{
|
{
|
||||||
|
|||||||
@@ -5,7 +5,6 @@ package cap
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"encoding/json"
|
|
||||||
"errors"
|
"errors"
|
||||||
"strconv"
|
"strconv"
|
||||||
"sync"
|
"sync"
|
||||||
@@ -13,12 +12,38 @@ import (
|
|||||||
"time"
|
"time"
|
||||||
|
|
||||||
"golang.org/x/sync/singleflight"
|
"golang.org/x/sync/singleflight"
|
||||||
|
"gorm.io/gorm"
|
||||||
|
|
||||||
"Wavelet/pkg/util"
|
"Wavelet/core"
|
||||||
cachepkg "Wavelet/plugins/infra/cache"
|
"Wavelet/core/contracts"
|
||||||
database "Wavelet/plugins/infra/database"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
|
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 (
|
const (
|
||||||
defaultChallengeCount = 1
|
defaultChallengeCount = 1
|
||||||
defaultChallengeSize = 32
|
defaultChallengeSize = 32
|
||||||
@@ -67,9 +92,8 @@ var runtimeConfigKeySet = func() map[string]struct{} {
|
|||||||
}()
|
}()
|
||||||
|
|
||||||
type runtimeSettingsStore struct {
|
type runtimeSettingsStore struct {
|
||||||
snapshot atomic.Pointer[RuntimeSettings]
|
snapshot atomic.Pointer[RuntimeSettings]
|
||||||
loadGroup singleflight.Group
|
loadGroup singleflight.Group
|
||||||
listenerOnce sync.Once
|
|
||||||
}
|
}
|
||||||
|
|
||||||
var settingsStore = &runtimeSettingsStore{}
|
var settingsStore = &runtimeSettingsStore{}
|
||||||
@@ -148,7 +172,11 @@ func loadRuntimeSettings(ctx context.Context) (RuntimeSettings, error) {
|
|||||||
Value string `gorm:"column:value"`
|
Value string `gorm:"column:value"`
|
||||||
}
|
}
|
||||||
var records []configRecord
|
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
|
return RuntimeSettings{}, err
|
||||||
}
|
}
|
||||||
configs := make(map[string]string, len(records))
|
configs := make(map[string]string, len(records))
|
||||||
@@ -167,6 +195,10 @@ func parseRuntimeSettings(configs map[string]string) RuntimeSettings {
|
|||||||
TokenTTL: defaultTokenTTL,
|
TokenTTL: defaultTokenTTL,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if len(configs) == 0 {
|
||||||
|
return settings
|
||||||
|
}
|
||||||
|
|
||||||
if val, ok := configs[ConfigKeyCapLoginEnabled]; ok {
|
if val, ok := configs[ConfigKeyCapLoginEnabled]; ok {
|
||||||
if enabled, err := strconv.ParseBool(val); err == nil {
|
if enabled, err := strconv.ParseBool(val); err == nil {
|
||||||
settings.LoginEnabled = enabled
|
settings.LoginEnabled = enabled
|
||||||
@@ -201,36 +233,4 @@ func parseRuntimeSettings(configs map[string]string) RuntimeSettings {
|
|||||||
return settings
|
return settings
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *runtimeSettingsStore) ensureInvalidationListener() {
|
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()
|
|
||||||
}
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -7,8 +7,9 @@ import (
|
|||||||
"net/http"
|
"net/http"
|
||||||
"strconv"
|
"strconv"
|
||||||
|
|
||||||
"Wavelet/pkg/response"
|
|
||||||
"github.com/gin-gonic/gin"
|
"github.com/gin-gonic/gin"
|
||||||
|
|
||||||
|
"Wavelet/pkg/response"
|
||||||
)
|
)
|
||||||
|
|
||||||
// ListAdminChannelDefinitions returns form schemas for supported channel types.
|
// ListAdminChannelDefinitions returns form schemas for supported channel types.
|
||||||
|
|||||||
@@ -5,13 +5,14 @@
|
|||||||
package qq
|
package qq
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"Wavelet/pkg/util"
|
|
||||||
"context"
|
"context"
|
||||||
"fmt"
|
"fmt"
|
||||||
"strings"
|
"strings"
|
||||||
"sync"
|
"sync"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
|
"Wavelet/pkg/util"
|
||||||
|
|
||||||
"Wavelet/pkg/logger"
|
"Wavelet/pkg/logger"
|
||||||
"Wavelet/plugins/domain/message_gateway"
|
"Wavelet/plugins/domain/message_gateway"
|
||||||
|
|
||||||
|
|||||||
@@ -12,9 +12,10 @@ import (
|
|||||||
"strconv"
|
"strconv"
|
||||||
"strings"
|
"strings"
|
||||||
|
|
||||||
|
tele "gopkg.in/telebot.v4"
|
||||||
|
|
||||||
"Wavelet/pkg/util"
|
"Wavelet/pkg/util"
|
||||||
"Wavelet/plugins/domain/message_gateway"
|
"Wavelet/plugins/domain/message_gateway"
|
||||||
tele "gopkg.in/telebot.v4"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
// Adapter is a Telegram private-chat channel.
|
// Adapter is a Telegram private-chat channel.
|
||||||
|
|||||||
@@ -7,8 +7,9 @@ import (
|
|||||||
"context"
|
"context"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
"Wavelet/plugins/domain/message_gateway"
|
|
||||||
tele "gopkg.in/telebot.v4"
|
tele "gopkg.in/telebot.v4"
|
||||||
|
|
||||||
|
"Wavelet/plugins/domain/message_gateway"
|
||||||
)
|
)
|
||||||
|
|
||||||
func TestHandleUpdate_DropsGroups(t *testing.T) {
|
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"
|
"net/http"
|
||||||
"strconv"
|
"strconv"
|
||||||
|
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
|
||||||
"Wavelet/core/contracts"
|
"Wavelet/core/contracts"
|
||||||
"Wavelet/pkg/response"
|
"Wavelet/pkg/response"
|
||||||
"Wavelet/pkg/util"
|
"Wavelet/pkg/util"
|
||||||
"github.com/gin-gonic/gin"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
func currentUser(c *gin.Context) (*contracts.UserDTO, bool) {
|
func currentUser(c *gin.Context) (*contracts.UserDTO, bool) {
|
||||||
|
|||||||
@@ -8,13 +8,35 @@ import (
|
|||||||
"testing"
|
"testing"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
|
"gorm.io/gorm"
|
||||||
|
|
||||||
"Wavelet/pkg/testhelper"
|
"Wavelet/pkg/testhelper"
|
||||||
"Wavelet/plugins/domain/message_gateway"
|
"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) {
|
func TestUpsertPairingCode_ReusesUnexpired(t *testing.T) {
|
||||||
_, _, cleanup := testhelper.SetupTestEnvironment(t)
|
testDB, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||||
defer cleanup()
|
message_gateway.SetDBServiceForTest(&mockDBService{db: testDB})
|
||||||
|
defer func() {
|
||||||
|
message_gateway.SetDBServiceForTest(nil)
|
||||||
|
cleanup()
|
||||||
|
}()
|
||||||
ctx := context.Background()
|
ctx := context.Background()
|
||||||
first, err := message_gateway.UpsertPairingCode(ctx, 1, "tg-1", "ABCD1234", time.Now().Add(15*time.Minute))
|
first, err := message_gateway.UpsertPairingCode(ctx, 1, "tg-1", "ABCD1234", time.Now().Add(15*time.Minute))
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
|||||||
@@ -9,12 +9,12 @@ import (
|
|||||||
"embed"
|
"embed"
|
||||||
"reflect"
|
"reflect"
|
||||||
|
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
|
||||||
"Wavelet/core"
|
"Wavelet/core"
|
||||||
"Wavelet/core/contracts"
|
"Wavelet/core/contracts"
|
||||||
"Wavelet/core/extpoints"
|
"Wavelet/core/extpoints"
|
||||||
"Wavelet/pkg/util"
|
"Wavelet/pkg/util"
|
||||||
"github.com/gin-gonic/gin"
|
|
||||||
"github.com/hibiken/asynq"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
//go:embed migrations/*.sql
|
//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.
|
// Apply registers message_gateway migrations, routes, tasks, schedules, events, and settings into the Context.
|
||||||
func (p *Plugin) Apply(ctx *core.Context) error {
|
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)
|
// 0. Resolve auth service for middleware (via IoC, not direct import)
|
||||||
var loginMW gin.HandlerFunc = func(c *gin.Context) { c.Next() }
|
var loginMW gin.HandlerFunc = func(c *gin.Context) { c.Next() }
|
||||||
var adminMW 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
|
const defaultTaskRetry = 3
|
||||||
pushHandler := &PushHandler{}
|
pushHandler := &PushHandler{}
|
||||||
|
|
||||||
// 5. Register Asynq background tasks
|
// 5. Register background tasks
|
||||||
ctx.Task().Register("message_gateway:push_notification", func(c context.Context, t *asynq.Task) error {
|
ctx.Task().Register("message_gateway:push_notification", func(c context.Context, payload []byte) error {
|
||||||
_, err := pushHandler.Execute(c, t.Payload())
|
return pushHandler.Execute(c, payload)
|
||||||
return err
|
|
||||||
}, extpoints.WithTaskRetry(defaultTaskRetry))
|
}, extpoints.WithTaskRetry(defaultTaskRetry))
|
||||||
|
|
||||||
ctx.Task().Register(SendNotificationTask, func(c context.Context, t *asynq.Task) error {
|
ctx.Task().Register(SendNotificationTask, func(c context.Context, payload []byte) error {
|
||||||
_, err := pushHandler.Execute(c, t.Payload())
|
return pushHandler.Execute(c, payload)
|
||||||
return err
|
|
||||||
}, extpoints.WithTaskRetry(defaultTaskRetry))
|
}, 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
|
return nil
|
||||||
})
|
})
|
||||||
|
|
||||||
@@ -184,9 +211,14 @@ func (p *Plugin) Apply(ctx *core.Context) error {
|
|||||||
return nil
|
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()
|
RegisterCustomEvents()
|
||||||
RegisterTaskListeners()
|
|
||||||
|
|
||||||
// 9. Register Settings Schemas
|
// 9. Register Settings Schemas
|
||||||
ctx.Settings().Register(extpoints.SettingSchema{
|
ctx.Settings().Register(extpoints.SettingSchema{
|
||||||
|
|||||||
@@ -11,10 +11,11 @@ import (
|
|||||||
"strings"
|
"strings"
|
||||||
"sync"
|
"sync"
|
||||||
|
|
||||||
"Wavelet/pkg/response"
|
|
||||||
pkgpush "Wavelet/plugins/domain/message_gateway/push"
|
|
||||||
"github.com/gin-gonic/gin"
|
"github.com/gin-gonic/gin"
|
||||||
"gorm.io/gorm"
|
"gorm.io/gorm"
|
||||||
|
|
||||||
|
"Wavelet/pkg/response"
|
||||||
|
pkgpush "Wavelet/plugins/domain/message_gateway/push"
|
||||||
)
|
)
|
||||||
|
|
||||||
const (
|
const (
|
||||||
|
|||||||
@@ -13,8 +13,9 @@ import (
|
|||||||
|
|
||||||
pkgpush "Wavelet/plugins/domain/message_gateway/push"
|
pkgpush "Wavelet/plugins/domain/message_gateway/push"
|
||||||
|
|
||||||
"Wavelet/pkg/util"
|
|
||||||
"gorm.io/gorm"
|
"gorm.io/gorm"
|
||||||
|
|
||||||
|
"Wavelet/pkg/util"
|
||||||
)
|
)
|
||||||
|
|
||||||
// NotificationMessage represents the structured notification message payload.
|
// NotificationMessage represents the structured notification message payload.
|
||||||
|
|||||||
@@ -11,9 +11,10 @@ import (
|
|||||||
|
|
||||||
pkgpush "Wavelet/plugins/domain/message_gateway/push"
|
pkgpush "Wavelet/plugins/domain/message_gateway/push"
|
||||||
|
|
||||||
"Wavelet/pkg/response"
|
|
||||||
"github.com/gin-gonic/gin"
|
"github.com/gin-gonic/gin"
|
||||||
"gorm.io/gorm"
|
"gorm.io/gorm"
|
||||||
|
|
||||||
|
"Wavelet/pkg/response"
|
||||||
)
|
)
|
||||||
|
|
||||||
// UpdatePushEventRequest is the request body for updating a push event.
|
// UpdatePushEventRequest is the request body for updating a push event.
|
||||||
|
|||||||
@@ -11,11 +11,10 @@ import (
|
|||||||
"strconv"
|
"strconv"
|
||||||
"strings"
|
"strings"
|
||||||
|
|
||||||
|
"gorm.io/gorm"
|
||||||
|
|
||||||
"Wavelet/core/contracts"
|
"Wavelet/core/contracts"
|
||||||
pkgpush "Wavelet/plugins/domain/message_gateway/push"
|
pkgpush "Wavelet/plugins/domain/message_gateway/push"
|
||||||
"Wavelet/plugins/drivers/driver_asynq_worker"
|
|
||||||
db "Wavelet/plugins/infra/database"
|
|
||||||
"gorm.io/gorm"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
type smtpConfig struct {
|
type smtpConfig struct {
|
||||||
@@ -28,10 +27,10 @@ type smtpConfig struct {
|
|||||||
func loadSMTPConfig(ctx context.Context) smtpConfig {
|
func loadSMTPConfig(ctx context.Context) smtpConfig {
|
||||||
var cfg smtpConfig
|
var cfg smtpConfig
|
||||||
var host, port, user, pass string
|
var host, port, user, pass string
|
||||||
_ = db.DB(ctx).Table("w_system_configs").Where("key = ?", "smtp_host").Pluck("value", &host).Error
|
_ = getDB(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
|
_ = getDB(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
|
_ = getDB(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_password").Pluck("value", &pass).Error
|
||||||
cfg.Host = host
|
cfg.Host = host
|
||||||
cfg.Port = port
|
cfg.Port = port
|
||||||
cfg.Username = user
|
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 {
|
if userID, ok := extractUserID(data); ok && userID > 0 {
|
||||||
var user contracts.UserDTO
|
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
|
return &user
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
if username := extractUsername(data); username != "" {
|
if username := extractUsername(data); username != "" {
|
||||||
var user contracts.UserDTO
|
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
|
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) {
|
func resolveTargetUser(ctx context.Context, resolved string, _ string) (contracts.UserDTO, bool) {
|
||||||
var user contracts.UserDTO
|
var user contracts.UserDTO
|
||||||
if id, err := strconv.ParseUint(resolved, 10, 64); err == nil {
|
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
|
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, true
|
||||||
}
|
}
|
||||||
return user, false
|
return user, false
|
||||||
@@ -387,7 +386,7 @@ func resolveSystemTarget(ctx context.Context, resolved string, channel string) (
|
|||||||
return "", false
|
return "", false
|
||||||
}
|
}
|
||||||
var adminUser contracts.UserDTO
|
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
|
return resolved, true
|
||||||
}
|
}
|
||||||
if channel == channelEmail && adminUser.Email != "" {
|
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 {
|
func getSystemUser(ctx context.Context) *contracts.UserDTO {
|
||||||
var user 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 &user
|
||||||
}
|
}
|
||||||
return &contracts.UserDTO{
|
return &contracts.UserDTO{
|
||||||
@@ -445,14 +444,16 @@ func findBuiltInEvent(key string) (EventMetadata, bool) {
|
|||||||
|
|
||||||
func getEventInfo(req CreatePushEventRequest) (string, string, []byte, error) {
|
func getEventInfo(req CreatePushEventRequest) (string, string, []byte, error) {
|
||||||
if req.TaskType != "" {
|
if req.TaskType != "" {
|
||||||
meta := driver_asynq_worker.GetTaskMetaByAsynqTask(req.TaskType)
|
taskName := req.TaskType
|
||||||
if meta == nil {
|
if taskSvc := getTaskService(); taskSvc != nil {
|
||||||
return "", "", nil, errors.New("unsupported task type")
|
if meta, ok := taskSvc.GetTaskMeta(req.TaskType); ok {
|
||||||
|
taskName = meta.DisplayName
|
||||||
|
}
|
||||||
}
|
}
|
||||||
eventKey := "task_completed:" + req.TaskType
|
eventKey := "task_completed:" + req.TaskType
|
||||||
eventName := "任务完成: " + meta.Name
|
eventName := "任务完成: " + taskName
|
||||||
defaultTemplate := NotificationMessage{
|
defaultTemplate := NotificationMessage{
|
||||||
Title: "任务完成: " + meta.Name,
|
Title: "任务完成: " + taskName,
|
||||||
Content: "异步任务 {{task_name}} (ID: {{task_id}}) 已完成。状态: {{task_status}},耗时: {{task_duration}} ms。",
|
Content: "异步任务 {{task_name}} (ID: {{task_id}}) 已完成。状态: {{task_status}},耗时: {{task_duration}} ms。",
|
||||||
Level: defaultLevelInfo,
|
Level: defaultLevelInfo,
|
||||||
}
|
}
|
||||||
@@ -484,8 +485,11 @@ func enqueuePushTask(ctx context.Context, payload SendPayload) error {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
_, err = driver_asynq_worker.DispatchTask(ctx, "send_notification", payloadBytes, "system")
|
if taskSvc := getTaskService(); taskSvc != nil {
|
||||||
return err
|
_, 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 {
|
func getFlatBody(body map[string]any) map[string]any {
|
||||||
|
|||||||
@@ -9,19 +9,14 @@ import (
|
|||||||
"strconv"
|
"strconv"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
|
"Wavelet/core/contracts"
|
||||||
"Wavelet/pkg/logger"
|
"Wavelet/pkg/logger"
|
||||||
"Wavelet/plugins/drivers/driver_asynq_worker"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
// RegisterTaskListeners subscribes push notification handlers to task completion events.
|
func handleTaskCompleted(ctx context.Context, e contracts.TaskCompletedEvent) {
|
||||||
func RegisterTaskListeners() {
|
events, err := listActivePushEventsByTaskType(ctx, e.TaskType)
|
||||||
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)
|
|
||||||
if err != nil {
|
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
|
return
|
||||||
}
|
}
|
||||||
if len(events) == 0 {
|
if len(events) == 0 {
|
||||||
@@ -29,34 +24,26 @@ func handleTaskCompleted(ctx context.Context, execution *driver_asynq_worker.Tas
|
|||||||
}
|
}
|
||||||
|
|
||||||
body := map[string]any{
|
body := map[string]any{
|
||||||
"task_id": execution.TaskID,
|
"task_id": e.TaskID,
|
||||||
"task_name": execution.TaskName,
|
"task_name": e.TaskName,
|
||||||
"task_type": execution.TaskType,
|
"task_type": e.TaskType,
|
||||||
"task_status": string(execution.Status),
|
"task_status": e.Status,
|
||||||
"task_duration": execution.Duration,
|
"task_duration": e.Duration,
|
||||||
"time": time.Now().Format("2006-01-02 15:04:05"),
|
"time": time.Now().Format("2006-01-02 15:04:05"),
|
||||||
}
|
"task_error": e.ErrorMsg,
|
||||||
if execErr != nil {
|
"task_result": e.ResultMsg,
|
||||||
body["task_error"] = execErr.Error()
|
|
||||||
} else {
|
|
||||||
body["task_error"] = ""
|
|
||||||
}
|
|
||||||
if result != nil {
|
|
||||||
body["task_result"] = result.Message
|
|
||||||
} else {
|
|
||||||
body["task_result"] = ""
|
|
||||||
}
|
}
|
||||||
|
|
||||||
var payloadMap map[string]any
|
var payloadMap map[string]any
|
||||||
if execution.Payload != "" {
|
if e.Payload != "" {
|
||||||
if err := json.Unmarshal([]byte(execution.Payload), &payloadMap); err == nil {
|
if err := json.Unmarshal([]byte(e.Payload), &payloadMap); err == nil {
|
||||||
body["payload"] = payloadMap
|
body["payload"] = payloadMap
|
||||||
extractUserFromMap(ctx, payloadMap, body)
|
extractUserFromMap(ctx, payloadMap, body)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
if result != nil && result.Detail != "" {
|
if e.Detail != "" {
|
||||||
var detailMap map[string]any
|
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
|
body["detail"] = detailMap
|
||||||
extractUserFromMap(ctx, detailMap, body)
|
extractUserFromMap(ctx, detailMap, body)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -9,8 +9,9 @@ import (
|
|||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
|
|
||||||
|
"Wavelet/core/contracts"
|
||||||
|
"Wavelet/pkg/logger"
|
||||||
"Wavelet/plugins/domain/message_gateway/push"
|
"Wavelet/plugins/domain/message_gateway/push"
|
||||||
"Wavelet/plugins/drivers/driver_asynq_worker"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
const (
|
const (
|
||||||
@@ -21,28 +22,24 @@ const (
|
|||||||
)
|
)
|
||||||
|
|
||||||
// SendNotificationMeta represents the task metadata.
|
// SendNotificationMeta represents the task metadata.
|
||||||
var SendNotificationMeta = driver_asynq_worker.TaskMeta{
|
var SendNotificationMeta = contracts.TaskMetaDTO{
|
||||||
Type: TaskTypeSendNotification,
|
Name: TaskTypeSendNotification,
|
||||||
AsynqTask: SendNotificationTask,
|
DisplayName: "推送通知",
|
||||||
Name: "推送通知",
|
Description: "异步执行系统通知的多渠道派发与推送",
|
||||||
Description: "异步执行系统通知的多渠道派发与推送",
|
MaxRetry: 3,
|
||||||
SupportsTime: false,
|
Queue: "default",
|
||||||
MaxRetry: driver_asynq_worker.DefaultMaxRetry,
|
Params: []contracts.TaskParamDTO{
|
||||||
Queue: driver_asynq_worker.QueueDefault,
|
|
||||||
Retryable: true,
|
|
||||||
Params: []driver_asynq_worker.TaskParam{
|
|
||||||
{
|
{
|
||||||
Name: "event_key",
|
Name: "event_key",
|
||||||
Label: "事件标识",
|
|
||||||
Type: "string",
|
Type: "string",
|
||||||
|
Description: "事件标识 (如 admin_login)",
|
||||||
Required: true,
|
Required: true,
|
||||||
Placeholder: "admin_login",
|
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
Name: "target",
|
Name: "target",
|
||||||
Label: "目标接收者",
|
Type: "string",
|
||||||
Type: "string",
|
Description: "目标接收者",
|
||||||
Required: false,
|
Required: false,
|
||||||
},
|
},
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
@@ -69,23 +66,21 @@ func (h *PushHandler) ValidatePayload(payload []byte) ([]byte, error) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Execute performs the push send and logs delivery history audit.
|
// 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
|
var req SendPayload
|
||||||
if err := json.Unmarshal(payload, &req); err != nil {
|
if err := json.Unmarshal(payload, &req); err != nil {
|
||||||
driver_asynq_worker.AppendLog(ctx, "解析推送参数失败: %v", err)
|
logger.ErrorF(ctx, "[Push] 解析推送参数失败: %v", err)
|
||||||
return nil, fmt.Errorf("parse payload failed: %w", 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)
|
pusher, err := push.GetPusher(req.Config.Channel)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
errWrap := fmt.Errorf("get pusher failed: %w", err)
|
errWrap := fmt.Errorf("get pusher failed: %w", err)
|
||||||
driver_asynq_worker.AppendLog(ctx, "推送失败: %v", errWrap)
|
logger.ErrorF(ctx, "[Push] 推送失败: %v", errWrap)
|
||||||
if driver_asynq_worker.IsFinalAttempt(ctx) {
|
h.recordHistory(ctx, req, "failed", errWrap.Error())
|
||||||
h.recordHistory(ctx, req, "failed", errWrap.Error())
|
return errWrap
|
||||||
}
|
|
||||||
return nil, errWrap
|
|
||||||
}
|
}
|
||||||
|
|
||||||
flatBody := req.Body.Flatten()
|
flatBody := req.Body.Flatten()
|
||||||
@@ -95,29 +90,19 @@ func (h *PushHandler) Execute(ctx context.Context, payload []byte) (*driver_asyn
|
|||||||
content := req.Body.Content
|
content := req.Body.Content
|
||||||
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
driver_asynq_worker.AppendLog(ctx, "消息推送失败 (标题: %s): %v", title, err)
|
logger.ErrorF(ctx, "[Push] 消息推送失败 (标题: %s): %v, 上游返回: %s", title, err, upstreamResp)
|
||||||
if upstreamResp != "" {
|
h.recordHistory(ctx, req, "failed", err.Error())
|
||||||
driver_asynq_worker.AppendLog(ctx, "上游返回: %s", upstreamResp)
|
return fmt.Errorf("pusher.Send failed: %w", err)
|
||||||
}
|
|
||||||
if driver_asynq_worker.IsFinalAttempt(ctx) {
|
|
||||||
h.recordHistory(ctx, req, "failed", err.Error())
|
|
||||||
}
|
|
||||||
return nil, fmt.Errorf("pusher.Send failed: %w", err)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
driver_asynq_worker.AppendLog(ctx, "消息推送成功 (标题: %s, 内容摘要: %s)", title, content)
|
logger.InfoF(ctx, "[Push] 消息推送成功 (标题: %s, 内容摘要: %s), 上游返回: %s", title, content, upstreamResp)
|
||||||
if upstreamResp != "" {
|
|
||||||
driver_asynq_worker.AppendLog(ctx, "上游返回: %s", upstreamResp)
|
|
||||||
}
|
|
||||||
h.recordHistory(ctx, req, "success", "")
|
h.recordHistory(ctx, req, "success", "")
|
||||||
|
|
||||||
return &driver_asynq_worker.TaskResult{
|
return nil
|
||||||
Message: fmt.Sprintf("推送成功: [%s] -> %s", req.Config.Channel, req.Target),
|
|
||||||
}, nil
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (h *PushHandler) recordHistory(ctx context.Context, req SendPayload, status string, errMsg string) {
|
func (h *PushHandler) recordHistory(ctx context.Context, req SendPayload, status string, errMsg string) {
|
||||||
if dbErr := recordPushHistory(ctx, req, status, errMsg); dbErr != nil {
|
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"
|
"gorm.io/gorm"
|
||||||
|
|
||||||
"Wavelet/pkg/idgen"
|
"Wavelet/pkg/idgen"
|
||||||
cachepkg "Wavelet/plugins/infra/cache"
|
|
||||||
db "Wavelet/plugins/infra/database"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
const (
|
const (
|
||||||
@@ -25,18 +23,18 @@ func CreateMessageChannel(ctx context.Context, ch *MessageChannel) error {
|
|||||||
if ch.ID == 0 {
|
if ch.ID == 0 {
|
||||||
ch.ID = idgen.NextUint64ID()
|
ch.ID = idgen.NextUint64ID()
|
||||||
}
|
}
|
||||||
return db.DB(ctx).Create(ch).Error
|
return getDB(ctx).Create(ch).Error
|
||||||
}
|
}
|
||||||
|
|
||||||
// UpdateMessageChannel saves a channel row.
|
// UpdateMessageChannel saves a channel row.
|
||||||
func UpdateMessageChannel(ctx context.Context, ch *MessageChannel) error {
|
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.
|
// GetMessageChannel loads a channel by id.
|
||||||
func GetMessageChannel(ctx context.Context, id uint64) (*MessageChannel, error) {
|
func GetMessageChannel(ctx context.Context, id uint64) (*MessageChannel, error) {
|
||||||
var ch MessageChannel
|
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 nil, err
|
||||||
}
|
}
|
||||||
return &ch, nil
|
return &ch, nil
|
||||||
@@ -45,7 +43,7 @@ func GetMessageChannel(ctx context.Context, id uint64) (*MessageChannel, error)
|
|||||||
// ListMessageChannels returns all channels newest first.
|
// ListMessageChannels returns all channels newest first.
|
||||||
func ListMessageChannels(ctx context.Context) ([]MessageChannel, error) {
|
func ListMessageChannels(ctx context.Context) ([]MessageChannel, error) {
|
||||||
var rows []MessageChannel
|
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 nil, err
|
||||||
}
|
}
|
||||||
return rows, nil
|
return rows, nil
|
||||||
@@ -53,7 +51,7 @@ func ListMessageChannels(ctx context.Context) ([]MessageChannel, error) {
|
|||||||
|
|
||||||
// DeleteMessageChannel removes pairings, bindings, then the channel.
|
// DeleteMessageChannel removes pairings, bindings, then the channel.
|
||||||
func DeleteMessageChannel(ctx context.Context, id uint64) error {
|
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 {
|
if err := tx.Where("channel_id = ?", id).Delete(&MessagePairingCode{}).Error; err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
@@ -69,13 +67,13 @@ func CreateMessageBinding(ctx context.Context, b *MessageBinding) error {
|
|||||||
if b.ID == 0 {
|
if b.ID == 0 {
|
||||||
b.ID = idgen.NextUint64ID()
|
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.
|
// GetBindingByChannelPlatform finds a binding for a platform user on a channel.
|
||||||
func GetBindingByChannelPlatform(ctx context.Context, channelID uint64, platformUserID string) (*MessageBinding, error) {
|
func GetBindingByChannelPlatform(ctx context.Context, channelID uint64, platformUserID string) (*MessageBinding, error) {
|
||||||
var b MessageBinding
|
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 {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
@@ -85,7 +83,7 @@ func GetBindingByChannelPlatform(ctx context.Context, channelID uint64, platform
|
|||||||
// ListBindingsByUser lists bindings for a Wavelet user.
|
// ListBindingsByUser lists bindings for a Wavelet user.
|
||||||
func ListBindingsByUser(ctx context.Context, userID uint64) ([]MessageBinding, error) {
|
func ListBindingsByUser(ctx context.Context, userID uint64) ([]MessageBinding, error) {
|
||||||
var rows []MessageBinding
|
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 nil, err
|
||||||
}
|
}
|
||||||
return rows, nil
|
return rows, nil
|
||||||
@@ -94,7 +92,7 @@ func ListBindingsByUser(ctx context.Context, userID uint64) ([]MessageBinding, e
|
|||||||
// GetMessageBinding loads a binding by id.
|
// GetMessageBinding loads a binding by id.
|
||||||
func GetMessageBinding(ctx context.Context, id uint64) (*MessageBinding, error) {
|
func GetMessageBinding(ctx context.Context, id uint64) (*MessageBinding, error) {
|
||||||
var b MessageBinding
|
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 nil, err
|
||||||
}
|
}
|
||||||
return &b, nil
|
return &b, nil
|
||||||
@@ -102,13 +100,13 @@ func GetMessageBinding(ctx context.Context, id uint64) (*MessageBinding, error)
|
|||||||
|
|
||||||
// DeleteMessageBinding deletes a binding by id.
|
// DeleteMessageBinding deletes a binding by id.
|
||||||
func DeleteMessageBinding(ctx context.Context, id uint64) error {
|
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.
|
// 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) {
|
func UpsertPairingCode(ctx context.Context, channelID uint64, platformUserID, code string, expiresAt time.Time) (*MessagePairingCode, error) {
|
||||||
var existing MessagePairingCode
|
var existing MessagePairingCode
|
||||||
err := db.DB(ctx).
|
err := getDB(ctx).
|
||||||
Where("channel_id = ? AND platform_user_id = ? AND expires_at > ?", channelID, platformUserID, time.Now()).
|
Where("channel_id = ? AND platform_user_id = ? AND expires_at > ?", channelID, platformUserID, time.Now()).
|
||||||
First(&existing).Error
|
First(&existing).Error
|
||||||
if err == nil {
|
if err == nil {
|
||||||
@@ -123,7 +121,7 @@ func UpsertPairingCode(ctx context.Context, channelID uint64, platformUserID, co
|
|||||||
PlatformUserID: platformUserID,
|
PlatformUserID: platformUserID,
|
||||||
ExpiresAt: expiresAt,
|
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 nil, err
|
||||||
}
|
}
|
||||||
return row, nil
|
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.
|
// GetPairingCode loads a pairing code by normalized code string.
|
||||||
func GetPairingCode(ctx context.Context, code string) (*MessagePairingCode, error) {
|
func GetPairingCode(ctx context.Context, code string) (*MessagePairingCode, error) {
|
||||||
var row MessagePairingCode
|
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 nil, err
|
||||||
}
|
}
|
||||||
return &row, nil
|
return &row, nil
|
||||||
@@ -140,18 +138,18 @@ func GetPairingCode(ctx context.Context, code string) (*MessagePairingCode, erro
|
|||||||
|
|
||||||
// DeletePairingCode removes a pairing code.
|
// DeletePairingCode removes a pairing code.
|
||||||
func DeletePairingCode(ctx context.Context, code string) error {
|
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.
|
// DeleteExpiredPairingCodes removes expired pairing rows.
|
||||||
func DeleteExpiredPairingCodes(ctx context.Context) error {
|
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.
|
// ListEnabledMessageChannels returns enabled channels.
|
||||||
func ListEnabledMessageChannels(ctx context.Context) ([]MessageChannel, error) {
|
func ListEnabledMessageChannels(ctx context.Context) ([]MessageChannel, error) {
|
||||||
var rows []MessageChannel
|
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 nil, err
|
||||||
}
|
}
|
||||||
return rows, nil
|
return rows, nil
|
||||||
@@ -160,7 +158,7 @@ func ListEnabledMessageChannels(ctx context.Context) ([]MessageChannel, error) {
|
|||||||
// ListPushChannelsRecord returns all push channels ordered by creation time descending.
|
// ListPushChannelsRecord returns all push channels ordered by creation time descending.
|
||||||
func ListPushChannelsRecord(ctx context.Context) ([]PushChannel, error) {
|
func ListPushChannelsRecord(ctx context.Context) ([]PushChannel, error) {
|
||||||
var channels []PushChannel
|
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 nil, err
|
||||||
}
|
}
|
||||||
return channels, nil
|
return channels, nil
|
||||||
@@ -169,7 +167,7 @@ func ListPushChannelsRecord(ctx context.Context) ([]PushChannel, error) {
|
|||||||
// GetPushChannelByIDRecord loads a push channel by primary key.
|
// GetPushChannelByIDRecord loads a push channel by primary key.
|
||||||
func GetPushChannelByIDRecord(ctx context.Context, id uint64) (PushChannel, error) {
|
func GetPushChannelByIDRecord(ctx context.Context, id uint64) (PushChannel, error) {
|
||||||
var channel PushChannel
|
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 PushChannel{}, err
|
||||||
}
|
}
|
||||||
return channel, nil
|
return channel, nil
|
||||||
@@ -178,7 +176,7 @@ func GetPushChannelByIDRecord(ctx context.Context, id uint64) (PushChannel, erro
|
|||||||
// GetPushChannelByNameRecord 根据名称获取消息通道。
|
// GetPushChannelByNameRecord 根据名称获取消息通道。
|
||||||
func GetPushChannelByNameRecord(ctx context.Context, name string) (*PushChannel, error) {
|
func GetPushChannelByNameRecord(ctx context.Context, name string) (*PushChannel, error) {
|
||||||
var channel PushChannel
|
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 nil, err
|
||||||
}
|
}
|
||||||
return &channel, nil
|
return &channel, nil
|
||||||
@@ -187,7 +185,7 @@ func GetPushChannelByNameRecord(ctx context.Context, name string) (*PushChannel,
|
|||||||
// CountPushChannelsByNameRecord returns how many channels share the given name.
|
// CountPushChannelsByNameRecord returns how many channels share the given name.
|
||||||
func CountPushChannelsByNameRecord(ctx context.Context, name string) (int64, error) {
|
func CountPushChannelsByNameRecord(ctx context.Context, name string) (int64, error) {
|
||||||
var count int64
|
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 0, err
|
||||||
}
|
}
|
||||||
return count, nil
|
return count, nil
|
||||||
@@ -195,7 +193,7 @@ func CountPushChannelsByNameRecord(ctx context.Context, name string) (int64, err
|
|||||||
|
|
||||||
// CreatePushChannelRecord persists a new channel and invalidates cache.
|
// CreatePushChannelRecord persists a new channel and invalidates cache.
|
||||||
func CreatePushChannelRecord(ctx context.Context, channel *PushChannel) error {
|
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
|
return err
|
||||||
}
|
}
|
||||||
DeleteActivePushChannelCache(ctx, channel.Name)
|
DeleteActivePushChannelCache(ctx, channel.Name)
|
||||||
@@ -204,7 +202,7 @@ func CreatePushChannelRecord(ctx context.Context, channel *PushChannel) error {
|
|||||||
|
|
||||||
// SavePushChannelRecord updates a channel and invalidates cache.
|
// SavePushChannelRecord updates a channel and invalidates cache.
|
||||||
func SavePushChannelRecord(ctx context.Context, channel *PushChannel) error {
|
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
|
return err
|
||||||
}
|
}
|
||||||
DeleteActivePushChannelCache(ctx, channel.Name)
|
DeleteActivePushChannelCache(ctx, channel.Name)
|
||||||
@@ -213,7 +211,7 @@ func SavePushChannelRecord(ctx context.Context, channel *PushChannel) error {
|
|||||||
|
|
||||||
// DeletePushChannelRecord removes a channel and invalidates cache.
|
// DeletePushChannelRecord removes a channel and invalidates cache.
|
||||||
func DeletePushChannelRecord(ctx context.Context, channel *PushChannel) error {
|
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
|
return err
|
||||||
}
|
}
|
||||||
DeleteActivePushChannelCache(ctx, channel.Name)
|
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) {
|
func GetActivePushChannelByName(ctx context.Context, name string) (*PushChannel, error) {
|
||||||
cacheKey := "push:channel:active:" + name
|
cacheKey := "push:channel:active:" + name
|
||||||
var channel PushChannel
|
var channel PushChannel
|
||||||
if cachepkg.Redis != nil {
|
if cache := getCache(ctx); cache != nil {
|
||||||
if err := cachepkg.GetJSON(ctx, cacheKey, &channel); err == nil {
|
if err := cache.Get(ctx, cacheKey, &channel); err == nil {
|
||||||
return &channel, 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
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
if cachepkg.Redis != nil {
|
if cache := getCache(ctx); cache != nil {
|
||||||
_ = cachepkg.SetJSON(ctx, cacheKey, channel, activePushChannelCacheTTL)
|
_ = cache.Set(ctx, cacheKey, channel, activePushChannelCacheTTL)
|
||||||
}
|
}
|
||||||
|
|
||||||
return &channel, nil
|
return &channel, nil
|
||||||
@@ -243,15 +241,15 @@ func GetActivePushChannelByName(ctx context.Context, name string) (*PushChannel,
|
|||||||
|
|
||||||
// DeleteActivePushChannelCache 清理启用消息通道的缓存。
|
// DeleteActivePushChannelCache 清理启用消息通道的缓存。
|
||||||
func DeleteActivePushChannelCache(ctx context.Context, name string) {
|
func DeleteActivePushChannelCache(ctx context.Context, name string) {
|
||||||
if cachepkg.Redis != nil {
|
if cache := getCache(ctx); cache != nil {
|
||||||
_ = cachepkg.Redis.Del(ctx, cachepkg.PrefixedKey("push:channel:active:"+name)).Err()
|
_ = cache.Delete(ctx, "push:channel:active:"+name)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// ListPushEventsRecord returns all push events ordered by creation time descending.
|
// ListPushEventsRecord returns all push events ordered by creation time descending.
|
||||||
func ListPushEventsRecord(ctx context.Context) ([]PushEvent, error) {
|
func ListPushEventsRecord(ctx context.Context) ([]PushEvent, error) {
|
||||||
var events []PushEvent
|
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 nil, err
|
||||||
}
|
}
|
||||||
return events, nil
|
return events, nil
|
||||||
@@ -260,7 +258,7 @@ func ListPushEventsRecord(ctx context.Context) ([]PushEvent, error) {
|
|||||||
// GetPushEventByIDRecord loads a push event by primary key.
|
// GetPushEventByIDRecord loads a push event by primary key.
|
||||||
func GetPushEventByIDRecord(ctx context.Context, id uint64) (PushEvent, error) {
|
func GetPushEventByIDRecord(ctx context.Context, id uint64) (PushEvent, error) {
|
||||||
var event PushEvent
|
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 PushEvent{}, err
|
||||||
}
|
}
|
||||||
return event, nil
|
return event, nil
|
||||||
@@ -269,7 +267,7 @@ func GetPushEventByIDRecord(ctx context.Context, id uint64) (PushEvent, error) {
|
|||||||
// GetPushEventByKeyRecord loads a push event by event key.
|
// GetPushEventByKeyRecord loads a push event by event key.
|
||||||
func GetPushEventByKeyRecord(ctx context.Context, key string) (PushEvent, error) {
|
func GetPushEventByKeyRecord(ctx context.Context, key string) (PushEvent, error) {
|
||||||
var event PushEvent
|
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 PushEvent{}, err
|
||||||
}
|
}
|
||||||
return event, nil
|
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.
|
// CountPushEventsByKeyRecord returns how many events use the given event key.
|
||||||
func CountPushEventsByKeyRecord(ctx context.Context, key string) (int64, error) {
|
func CountPushEventsByKeyRecord(ctx context.Context, key string) (int64, error) {
|
||||||
var count int64
|
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 0, err
|
||||||
}
|
}
|
||||||
return count, nil
|
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.
|
// CreatePushEventRecord persists a new push event and invalidates cache.
|
||||||
func CreatePushEventRecord(ctx context.Context, event *PushEvent) error {
|
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
|
return err
|
||||||
}
|
}
|
||||||
DeleteActivePushEventCache(ctx, event.EventKey)
|
DeleteActivePushEventCache(ctx, event.EventKey)
|
||||||
@@ -295,7 +293,7 @@ func CreatePushEventRecord(ctx context.Context, event *PushEvent) error {
|
|||||||
|
|
||||||
// SavePushEventRecord updates a push event and invalidates cache.
|
// SavePushEventRecord updates a push event and invalidates cache.
|
||||||
func SavePushEventRecord(ctx context.Context, event *PushEvent) error {
|
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
|
return err
|
||||||
}
|
}
|
||||||
DeleteActivePushEventCache(ctx, event.EventKey)
|
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.
|
// UpdatePushEventEnabledRecord toggles the enabled flag for a push event.
|
||||||
func UpdatePushEventEnabledRecord(ctx context.Context, event *PushEvent, enabled bool) error {
|
func UpdatePushEventEnabledRecord(ctx context.Context, event *PushEvent, enabled bool) error {
|
||||||
event.Enabled = enabled
|
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
|
return err
|
||||||
}
|
}
|
||||||
DeleteActivePushEventCache(ctx, event.EventKey)
|
DeleteActivePushEventCache(ctx, event.EventKey)
|
||||||
@@ -314,7 +312,7 @@ func UpdatePushEventEnabledRecord(ctx context.Context, event *PushEvent, enabled
|
|||||||
|
|
||||||
// DeletePushEventRecord removes a push event and invalidates cache.
|
// DeletePushEventRecord removes a push event and invalidates cache.
|
||||||
func DeletePushEventRecord(ctx context.Context, event *PushEvent) error {
|
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
|
return err
|
||||||
}
|
}
|
||||||
DeleteActivePushEventCache(ctx, event.EventKey)
|
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.
|
// ListActivePushEventsByTaskTypeRecord returns enabled events bound to a task type.
|
||||||
func ListActivePushEventsByTaskTypeRecord(ctx context.Context, taskType string) ([]PushEvent, error) {
|
func ListActivePushEventsByTaskTypeRecord(ctx context.Context, taskType string) ([]PushEvent, error) {
|
||||||
var events []PushEvent
|
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 nil, err
|
||||||
}
|
}
|
||||||
return events, nil
|
return events, nil
|
||||||
@@ -334,18 +332,18 @@ func ListActivePushEventsByTaskTypeRecord(ctx context.Context, taskType string)
|
|||||||
func GetActivePushEventByKey(ctx context.Context, key string) (*PushEvent, error) {
|
func GetActivePushEventByKey(ctx context.Context, key string) (*PushEvent, error) {
|
||||||
cacheKey := "push:event:active:" + key
|
cacheKey := "push:event:active:" + key
|
||||||
var event PushEvent
|
var event PushEvent
|
||||||
if cachepkg.Redis != nil {
|
if cache := getCache(ctx); cache != nil {
|
||||||
if err := cachepkg.GetJSON(ctx, cacheKey, &event); err == nil {
|
if err := cache.Get(ctx, cacheKey, &event); err == nil {
|
||||||
return &event, 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
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
if cachepkg.Redis != nil {
|
if cache := getCache(ctx); cache != nil {
|
||||||
_ = cachepkg.SetJSON(ctx, cacheKey, event, activePushEventCacheTTL)
|
_ = cache.Set(ctx, cacheKey, event, activePushEventCacheTTL)
|
||||||
}
|
}
|
||||||
|
|
||||||
return &event, nil
|
return &event, nil
|
||||||
@@ -353,14 +351,14 @@ func GetActivePushEventByKey(ctx context.Context, key string) (*PushEvent, error
|
|||||||
|
|
||||||
// DeleteActivePushEventCache 清理启用通知事件的缓存。
|
// DeleteActivePushEventCache 清理启用通知事件的缓存。
|
||||||
func DeleteActivePushEventCache(ctx context.Context, key string) {
|
func DeleteActivePushEventCache(ctx context.Context, key string) {
|
||||||
if cachepkg.Redis != nil {
|
if cache := getCache(ctx); cache != nil {
|
||||||
_ = cachepkg.Redis.Del(ctx, cachepkg.PrefixedKey("push:event:active:"+key)).Err()
|
_ = cache.Delete(ctx, "push:event:active:"+key)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// ListPushHistoriesRecord returns paginated push history records.
|
// ListPushHistoriesRecord returns paginated push history records.
|
||||||
func ListPushHistoriesRecord(ctx context.Context, filter PushHistoryListFilter) (int64, []PushHistory, error) {
|
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 != "" {
|
if filter.EventKey != "" {
|
||||||
query = query.Where("event_key = ?", 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.
|
// CreatePushHistoryRecord persists a push history audit record.
|
||||||
func CreatePushHistoryRecord(ctx context.Context, history *PushHistory) error {
|
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.
|
// PushHistoryQuery returns a scoped query builder for push histories.
|
||||||
func PushHistoryQuery(ctx context.Context) *gorm.DB {
|
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"
|
"fmt"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"Wavelet/pkg/util"
|
|
||||||
db "Wavelet/plugins/infra/database"
|
|
||||||
"gorm.io/gorm"
|
"gorm.io/gorm"
|
||||||
|
|
||||||
|
"Wavelet/pkg/util"
|
||||||
)
|
)
|
||||||
|
|
||||||
// CountAccessLogs returns the number of access logs matching filter.
|
// CountAccessLogs returns the number of access logs matching filter.
|
||||||
func CountAccessLogs(ctx context.Context, filter AccessLogFilter) (uint64, error) {
|
func CountAccessLogs(ctx context.Context, filter AccessLogFilter) (uint64, error) {
|
||||||
ch := db.ChDB(ctx)
|
ch := getChDB(ctx)
|
||||||
if ch == nil {
|
if ch == nil {
|
||||||
return 0, fmt.Errorf("clickhouse gorm connection is not initialized")
|
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.
|
// ListAccessLogs returns paginated access logs and the total match count.
|
||||||
func ListAccessLogs(ctx context.Context, filter AccessLogFilter, page, pageSize int) ([]UserAccessLog, uint64, error) {
|
func ListAccessLogs(ctx context.Context, filter AccessLogFilter, page, pageSize int) ([]UserAccessLog, uint64, error) {
|
||||||
ch := db.ChDB(ctx)
|
ch := getChDB(ctx)
|
||||||
if ch == nil {
|
if ch == nil {
|
||||||
return nil, 0, fmt.Errorf("clickhouse gorm connection is not initialized")
|
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
|
return []UserAccessLog{}, 0, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
var total int64
|
var count int64
|
||||||
baseQuery := applyFilter(ch.Model(&UserAccessLog{}), filter)
|
query := applyFilter(ch.Model(&UserAccessLog{}), filter)
|
||||||
if err := baseQuery.Count(&total).Error; err != nil {
|
if err := query.Count(&count).Error; err != nil {
|
||||||
return nil, 0, fmt.Errorf("count access logs: %w", err)
|
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
|
var logs []UserAccessLog
|
||||||
err := applyFilter(ch.Model(&UserAccessLog{}), filter).
|
offset := (page - 1) * pageSize
|
||||||
Order("created_at DESC, id DESC").
|
if err := query.Order("created_at DESC").Limit(pageSize).Offset(offset).Find(&logs).Error; err != nil {
|
||||||
Limit(pageSize).
|
|
||||||
Offset(offset).
|
|
||||||
Find(&logs).Error
|
|
||||||
if err != nil {
|
|
||||||
return nil, 0, fmt.Errorf("list access logs: %w", err)
|
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.
|
// DeleteAllUserAccessLogs hard-deletes all user access logs via TRUNCATE.
|
||||||
func DeleteAllUserAccessLogs(ctx context.Context) (int64, error) {
|
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")
|
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, fmt.Errorf("truncate user access logs: %w", err)
|
||||||
}
|
}
|
||||||
return 0, nil
|
return 0, nil
|
||||||
@@ -83,10 +69,11 @@ func DeleteAllUserAccessLogs(ctx context.Context) (int64, error) {
|
|||||||
|
|
||||||
// DeleteUserAccessLogsBefore deletes user access logs older than cutoff.
|
// DeleteUserAccessLogsBefore deletes user access logs older than cutoff.
|
||||||
func DeleteUserAccessLogsBefore(ctx context.Context, cutoff time.Time) (int64, error) {
|
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")
|
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, fmt.Errorf("delete expired user access logs: %w", err)
|
||||||
}
|
}
|
||||||
return 0, nil
|
return 0, nil
|
||||||
|
|||||||
@@ -8,8 +8,6 @@ import (
|
|||||||
"fmt"
|
"fmt"
|
||||||
"sort"
|
"sort"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
db "Wavelet/plugins/infra/database"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
const hoursInDay = 24
|
const hoursInDay = 24
|
||||||
@@ -20,7 +18,7 @@ func GetDailyTrend(ctx context.Context, days int) ([]DailyTrend, error) {
|
|||||||
days = 7
|
days = 7
|
||||||
}
|
}
|
||||||
|
|
||||||
ch := db.ChDB(ctx)
|
ch := getChDB(ctx)
|
||||||
if ch == nil {
|
if ch == nil {
|
||||||
return nil, fmt.Errorf("clickhouse gorm connection is not initialized")
|
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.
|
// GetBrowserDistribution returns browser-grouped access counts since startTime.
|
||||||
func GetBrowserDistribution(ctx context.Context, startTime time.Time) ([]BrowserShare, error) {
|
func GetBrowserDistribution(ctx context.Context, startTime time.Time) ([]BrowserShare, error) {
|
||||||
ch := db.ChDB(ctx)
|
ch := getChDB(ctx)
|
||||||
if ch == nil {
|
if ch == nil {
|
||||||
return nil, fmt.Errorf("clickhouse gorm connection is not initialized")
|
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
|
limit = 10
|
||||||
}
|
}
|
||||||
|
|
||||||
ch := db.ChDB(ctx)
|
ch := getChDB(ctx)
|
||||||
if ch == nil {
|
if ch == nil {
|
||||||
return nil, fmt.Errorf("clickhouse gorm connection is not initialized")
|
return nil, fmt.Errorf("clickhouse gorm connection is not initialized")
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -16,8 +16,6 @@ import (
|
|||||||
"github.com/stretchr/testify/assert"
|
"github.com/stretchr/testify/assert"
|
||||||
"github.com/stretchr/testify/require"
|
"github.com/stretchr/testify/require"
|
||||||
"gorm.io/gorm"
|
"gorm.io/gorm"
|
||||||
|
|
||||||
db "Wavelet/plugins/infra/database"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
func setupChGormDB(t *testing.T) *gorm.DB {
|
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, err)
|
||||||
require.NoError(t, gormDB.AutoMigrate(&UserAccessLog{}))
|
require.NoError(t, gormDB.AutoMigrate(&UserAccessLog{}))
|
||||||
db.SetChDBForTest(gormDB)
|
SetChDBForTest(gormDB)
|
||||||
return gormDB
|
return gormDB
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -56,7 +54,7 @@ func TestParseBrowserName(t *testing.T) {
|
|||||||
|
|
||||||
func TestCountAccessLogs_EmptyUserIDs(t *testing.T) {
|
func TestCountAccessLogs_EmptyUserIDs(t *testing.T) {
|
||||||
setupChGormDB(t)
|
setupChGormDB(t)
|
||||||
t.Cleanup(func() { db.SetChDBForTest(nil) })
|
t.Cleanup(func() { SetChDBForTest(nil) })
|
||||||
|
|
||||||
count, err := CountAccessLogs(context.Background(), AccessLogFilter{UserIDs: []uint64{}})
|
count, err := CountAccessLogs(context.Background(), AccessLogFilter{UserIDs: []uint64{}})
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
@@ -65,7 +63,7 @@ func TestCountAccessLogs_EmptyUserIDs(t *testing.T) {
|
|||||||
|
|
||||||
func TestListAccessLogs_EmptyUserIDs(t *testing.T) {
|
func TestListAccessLogs_EmptyUserIDs(t *testing.T) {
|
||||||
setupChGormDB(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)
|
logs, total, err := ListAccessLogs(context.Background(), AccessLogFilter{UserIDs: []uint64{}}, 1, 20)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
@@ -75,7 +73,7 @@ func TestListAccessLogs_EmptyUserIDs(t *testing.T) {
|
|||||||
|
|
||||||
func TestListAccessLogs_WithFilters(t *testing.T) {
|
func TestListAccessLogs_WithFilters(t *testing.T) {
|
||||||
gormDB := setupChGormDB(t)
|
gormDB := setupChGormDB(t)
|
||||||
t.Cleanup(func() { db.SetChDBForTest(nil) })
|
t.Cleanup(func() { SetChDBForTest(nil) })
|
||||||
|
|
||||||
now := time.Now().UTC().Truncate(time.Second)
|
now := time.Now().UTC().Truncate(time.Second)
|
||||||
logs := []UserAccessLog{
|
logs := []UserAccessLog{
|
||||||
@@ -116,8 +114,8 @@ func TestBatchInsert_UsesModelBatchSQL(t *testing.T) {
|
|||||||
batch: mockBatch,
|
batch: mockBatch,
|
||||||
batchQuery: UserAccessLog{}.BatchInsertSQL(),
|
batchQuery: UserAccessLog{}.BatchInsertSQL(),
|
||||||
}
|
}
|
||||||
db.SetChConnForTest(mockConn)
|
SetChConnForTest(mockConn)
|
||||||
t.Cleanup(func() { db.SetChConnForTest(nil) })
|
t.Cleanup(func() { SetChConnForTest(nil) })
|
||||||
|
|
||||||
createdAt := time.Now().UTC()
|
createdAt := time.Now().UTC()
|
||||||
err := BatchInsert(ctx, []UserAccessLog{
|
err := BatchInsert(ctx, []UserAccessLog{
|
||||||
|
|||||||
@@ -6,8 +6,6 @@ package logstore
|
|||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"fmt"
|
"fmt"
|
||||||
|
|
||||||
db "Wavelet/plugins/infra/database"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
// BatchInsert writes access logs to ClickHouse using the native batch API.
|
// 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 {
|
if len(logs) == 0 {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
if db.ChConn == nil {
|
conn := getChConn()
|
||||||
|
if conn == nil {
|
||||||
return fmt.Errorf("clickhouse connection is not initialized")
|
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 {
|
if err != nil {
|
||||||
return fmt.Errorf("prepare clickhouse batch: %w", err)
|
return fmt.Errorf("prepare clickhouse batch: %w", err)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -8,7 +8,6 @@ import (
|
|||||||
"fmt"
|
"fmt"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
db "Wavelet/plugins/infra/database"
|
|
||||||
"github.com/ClickHouse/clickhouse-go/v2/lib/driver"
|
"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) {
|
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")
|
return time.Time{}, time.Time{}, fmt.Errorf("clickhouse connection is not initialized")
|
||||||
}
|
}
|
||||||
table := UserAccessLog{}.TableName()
|
table := UserAccessLog{}.TableName()
|
||||||
var minTime, maxTime *time.Time
|
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)
|
return time.Time{}, time.Time{}, fmt.Errorf("query migration range %s: %w", table, err)
|
||||||
}
|
}
|
||||||
if minTime == nil || maxTime == nil {
|
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) {
|
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")
|
return nil, fmt.Errorf("clickhouse connection is not initialized")
|
||||||
}
|
}
|
||||||
if limit <= 0 {
|
if limit <= 0 {
|
||||||
@@ -116,7 +117,7 @@ func (s *clickhouseUserAccessLogStore) ListForMigration(ctx context.Context, aft
|
|||||||
}
|
}
|
||||||
table := UserAccessLog{}.TableName()
|
table := UserAccessLog{}.TableName()
|
||||||
columns := UserAccessLog{}.InsertColumns()
|
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 ?",
|
"SELECT %s FROM %s WHERE id > ? ORDER BY id ASC LIMIT ?",
|
||||||
columns, table,
|
columns, table,
|
||||||
), afterID, limit)
|
), 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"
|
"strings"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"Wavelet/pkg/idgen"
|
|
||||||
"gorm.io/gorm"
|
"gorm.io/gorm"
|
||||||
|
|
||||||
|
"Wavelet/pkg/idgen"
|
||||||
)
|
)
|
||||||
|
|
||||||
const (
|
const (
|
||||||
|
|||||||
@@ -12,7 +12,6 @@ import (
|
|||||||
|
|
||||||
"Wavelet/pkg/config"
|
"Wavelet/pkg/config"
|
||||||
"Wavelet/pkg/logger"
|
"Wavelet/pkg/logger"
|
||||||
db "Wavelet/plugins/infra/database"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
const (
|
const (
|
||||||
@@ -98,7 +97,7 @@ func buildStore(ctx context.Context, database string, skipFreeze bool) (*Store,
|
|||||||
ual.skipFreeze = skipFreeze
|
ual.skipFreeze = skipFreeze
|
||||||
return &Store{UserAccessLogs: ual, Status: ual}, nil
|
return &Store{UserAccessLogs: ual, Status: ual}, nil
|
||||||
case dbNamePostgres, dbNameSQLite:
|
case dbNamePostgres, dbNameSQLite:
|
||||||
gdb := db.DB(ctx)
|
gdb := getDB(ctx)
|
||||||
ual := newUserAccessLogGormStore(gdb)
|
ual := newUserAccessLogGormStore(gdb)
|
||||||
ual.skipFreeze = skipFreeze
|
ual.skipFreeze = skipFreeze
|
||||||
return &Store{UserAccessLogs: ual, Status: ual}, nil
|
return &Store{UserAccessLogs: ual, Status: ual}, nil
|
||||||
|
|||||||
@@ -9,13 +9,14 @@ import (
|
|||||||
"net/http"
|
"net/http"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
|
||||||
"Wavelet/core/contracts"
|
"Wavelet/core/contracts"
|
||||||
"Wavelet/pkg/config"
|
"Wavelet/pkg/config"
|
||||||
"Wavelet/pkg/idgen"
|
"Wavelet/pkg/idgen"
|
||||||
"Wavelet/pkg/response"
|
"Wavelet/pkg/response"
|
||||||
"Wavelet/pkg/util"
|
"Wavelet/pkg/util"
|
||||||
"Wavelet/plugins/domain/risk_control/logstore"
|
"Wavelet/plugins/domain/risk_control/logstore"
|
||||||
"github.com/gin-gonic/gin"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
// Middleware is an alias for RiskControlMiddleware.
|
// Middleware is an alias for RiskControlMiddleware.
|
||||||
|
|||||||
@@ -12,6 +12,9 @@ import (
|
|||||||
"testing"
|
"testing"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
|
||||||
"Wavelet/core/contracts"
|
"Wavelet/core/contracts"
|
||||||
"Wavelet/pkg/batchwriter"
|
"Wavelet/pkg/batchwriter"
|
||||||
"Wavelet/pkg/config"
|
"Wavelet/pkg/config"
|
||||||
@@ -19,8 +22,6 @@ import (
|
|||||||
"Wavelet/pkg/util"
|
"Wavelet/pkg/util"
|
||||||
"Wavelet/plugins/domain/risk_control"
|
"Wavelet/plugins/domain/risk_control"
|
||||||
"Wavelet/plugins/domain/risk_control/logstore"
|
"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) {
|
func newTestAccessLogWriter(t *testing.T, cfg batchwriter.Config) (*batchwriter.Writer[*logstore.UserAccessLog], func() []*logstore.UserAccessLog) {
|
||||||
|
|||||||
@@ -9,10 +9,12 @@ import (
|
|||||||
"embed"
|
"embed"
|
||||||
"reflect"
|
"reflect"
|
||||||
|
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
|
||||||
"Wavelet/core"
|
"Wavelet/core"
|
||||||
"Wavelet/core/contracts"
|
"Wavelet/core/contracts"
|
||||||
"Wavelet/core/extpoints"
|
"Wavelet/core/extpoints"
|
||||||
"github.com/gin-gonic/gin"
|
"Wavelet/plugins/domain/risk_control/logstore"
|
||||||
)
|
)
|
||||||
|
|
||||||
//go:embed logstore/migrations/*.sql
|
//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.
|
// Apply registers risk control middlewares, settings, and cleanup hooks into the Context.
|
||||||
func (p *Plugin) Apply(ctx *core.Context) error {
|
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
|
// 0. Register user access log table migrations
|
||||||
ctx.Migrations().Register("risk_control/logstore", riskControlMigrations)
|
ctx.Migrations().Register("risk_control/logstore", riskControlMigrations)
|
||||||
|
|
||||||
@@ -98,10 +113,90 @@ func (p *Plugin) Apply(ctx *core.Context) error {
|
|||||||
Category: "security",
|
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 {
|
ctx.OnDispose(func() error {
|
||||||
return StopLogWriter(context.Background())
|
return StopLogWriter(context.Background())
|
||||||
})
|
})
|
||||||
|
|
||||||
return nil
|
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"
|
"net/http"
|
||||||
"reflect"
|
"reflect"
|
||||||
|
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
|
||||||
"Wavelet/core"
|
"Wavelet/core"
|
||||||
"Wavelet/core/contracts"
|
"Wavelet/core/contracts"
|
||||||
"Wavelet/pkg/config"
|
"Wavelet/pkg/config"
|
||||||
"Wavelet/pkg/response"
|
"Wavelet/pkg/response"
|
||||||
"github.com/gin-gonic/gin"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
// Plugin implements core.Plugin to provide system-level basic routes.
|
// 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"`
|
Value string `json:"value"`
|
||||||
}
|
}
|
||||||
var configs []configItem
|
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
|
_ = dbSvc.DB(c.Request.Context()).Table("w_system_configs").Where("visibility = ?", "visible").Find(&configs).Error
|
||||||
}
|
}
|
||||||
c.JSON(http.StatusOK, response.OK(gin.H{
|
c.JSON(http.StatusOK, response.OK(gin.H{
|
||||||
|
|||||||
+8
-38
@@ -11,19 +11,13 @@ import (
|
|||||||
"sync"
|
"sync"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"Wavelet/pkg/util"
|
|
||||||
"Wavelet/plugins/domain/upload/shared"
|
"Wavelet/plugins/domain/upload/shared"
|
||||||
uploadstorage "Wavelet/plugins/domain/upload/storage"
|
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"
|
const fileAccessInvalidationChannel = "upload:file_access_invalidation"
|
||||||
|
|
||||||
var (
|
var (
|
||||||
accessCacheOnce sync.Once
|
|
||||||
|
|
||||||
fileAccessWhitelistMu sync.RWMutex
|
fileAccessWhitelistMu sync.RWMutex
|
||||||
fileAccessWhitelistTypes map[string]struct{}
|
fileAccessWhitelistTypes map[string]struct{}
|
||||||
fileAccessWhitelistValid bool
|
fileAccessWhitelistValid bool
|
||||||
@@ -42,35 +36,10 @@ func ResetAccessCaches() {
|
|||||||
|
|
||||||
// PublishAccessCacheInvalidation broadcasts upload access cache eviction to all nodes.
|
// PublishAccessCacheInvalidation broadcasts upload access cache eviction to all nodes.
|
||||||
func PublishAccessCacheInvalidation(ctx context.Context) {
|
func PublishAccessCacheInvalidation(ctx context.Context) {
|
||||||
if cachepkg.Redis != nil {
|
if cache := shared.GetCache(ctx); cache != nil {
|
||||||
_ = cachepkg.Redis.Publish(ctx, fileAccessInvalidationChannel, "reset").Err()
|
_ = cache.Invalidate(ctx, fileAccessInvalidationChannel)
|
||||||
}
|
}
|
||||||
}
|
ResetAccessCaches()
|
||||||
|
|
||||||
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()
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// IsFilePublic reports whether uploadType is in the public access whitelist.
|
// 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{} {
|
func loadFileAccessWhitelist(ctx context.Context) map[string]struct{} {
|
||||||
ensureAccessCacheListener()
|
|
||||||
|
|
||||||
fileAccessWhitelistMu.RLock()
|
fileAccessWhitelistMu.RLock()
|
||||||
if fileAccessWhitelistValid && time.Since(fileAccessWhitelistCheckedAt) < time.Duration(shared.AccessCacheTTL)*time.Second {
|
if fileAccessWhitelistValid && time.Since(fileAccessWhitelistCheckedAt) < time.Duration(shared.AccessCacheTTL)*time.Second {
|
||||||
types := fileAccessWhitelistTypes
|
types := fileAccessWhitelistTypes
|
||||||
@@ -115,8 +82,11 @@ func fetchFileAccessWhitelist(ctx context.Context) map[string]struct{} {
|
|||||||
|
|
||||||
func parseFileAccessWhitelist(ctx context.Context) []string {
|
func parseFileAccessWhitelist(ctx context.Context) []string {
|
||||||
var sc struct{ Value string }
|
var sc struct{ Value string }
|
||||||
err := database.DB(ctx).Table("w_system_configs").Where("key = ?", "file_access_whitelist").First(&sc).Error
|
db := shared.GetDB(ctx)
|
||||||
if err != nil || sc.Value == "" {
|
if db != nil {
|
||||||
|
_ = db.Table("w_system_configs").Where("key = ?", "file_access_whitelist").First(&sc).Error
|
||||||
|
}
|
||||||
|
if sc.Value == "" {
|
||||||
return []string{shared.DefaultPublicUploadType}
|
return []string{shared.DefaultPublicUploadType}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -8,13 +8,12 @@ import (
|
|||||||
"testing"
|
"testing"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"Wavelet/pkg/testhelper"
|
|
||||||
"Wavelet/plugins/domain/upload/shared"
|
"Wavelet/plugins/domain/upload/shared"
|
||||||
uploadstorage "Wavelet/plugins/domain/upload/storage"
|
uploadstorage "Wavelet/plugins/domain/upload/storage"
|
||||||
)
|
)
|
||||||
|
|
||||||
func TestLoadMigrationAccessStateCachesResult(t *testing.T) {
|
func TestLoadMigrationAccessStateCachesResult(t *testing.T) {
|
||||||
_, _, cleanup := testhelper.SetupTestEnvironment(t)
|
_, cleanup := shared.SetupTestEnv(t)
|
||||||
defer cleanup()
|
defer cleanup()
|
||||||
ResetAccessCaches()
|
ResetAccessCaches()
|
||||||
|
|
||||||
@@ -31,7 +30,7 @@ func TestLoadMigrationAccessStateCachesResult(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestIsFilePublicUsesCachedWhitelist(t *testing.T) {
|
func TestIsFilePublicUsesCachedWhitelist(t *testing.T) {
|
||||||
_, _, cleanup := testhelper.SetupTestEnvironment(t)
|
_, cleanup := shared.SetupTestEnv(t)
|
||||||
defer cleanup()
|
defer cleanup()
|
||||||
ResetAccessCaches()
|
ResetAccessCaches()
|
||||||
|
|
||||||
@@ -48,7 +47,7 @@ func TestIsFilePublicUsesCachedWhitelist(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestResetAccessCachesRefreshesWhitelist(t *testing.T) {
|
func TestResetAccessCachesRefreshesWhitelist(t *testing.T) {
|
||||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
dbConn, cleanup := shared.SetupTestEnv(t)
|
||||||
defer cleanup()
|
defer cleanup()
|
||||||
ResetAccessCaches()
|
ResetAccessCaches()
|
||||||
|
|
||||||
@@ -71,7 +70,7 @@ func TestResetAccessCachesRefreshesWhitelist(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestAccessCacheTTLExpires(t *testing.T) {
|
func TestAccessCacheTTLExpires(t *testing.T) {
|
||||||
_, _, cleanup := testhelper.SetupTestEnvironment(t)
|
_, cleanup := shared.SetupTestEnv(t)
|
||||||
defer cleanup()
|
defer cleanup()
|
||||||
ResetAccessCaches()
|
ResetAccessCaches()
|
||||||
|
|
||||||
|
|||||||
+70
-97
@@ -5,15 +5,14 @@ package cache
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"encoding/json"
|
|
||||||
"fmt"
|
"fmt"
|
||||||
"sync"
|
"time"
|
||||||
|
|
||||||
|
"gorm.io/gorm"
|
||||||
|
|
||||||
"Wavelet/pkg/cache/ram"
|
"Wavelet/pkg/cache/ram"
|
||||||
"Wavelet/pkg/util"
|
|
||||||
"Wavelet/plugins/domain/upload/models"
|
"Wavelet/plugins/domain/upload/models"
|
||||||
cachepkg "Wavelet/plugins/infra/cache"
|
"Wavelet/plugins/domain/upload/shared"
|
||||||
database "Wavelet/plugins/infra/database"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
const (
|
const (
|
||||||
@@ -22,16 +21,8 @@ const (
|
|||||||
uploadMetaInvalidationChan = "upload:meta_invalidation"
|
uploadMetaInvalidationChan = "upload:meta_invalidation"
|
||||||
)
|
)
|
||||||
|
|
||||||
type uploadMetaInvalidationMessage struct {
|
|
||||||
ID uint64 `json:"id"`
|
|
||||||
}
|
|
||||||
|
|
||||||
var (
|
var (
|
||||||
uploadMetaRAM = ram.MustNew[uint64, models.Upload](ram.Options{MaximumSize: uploadMetaRAMMaximumSize})
|
uploadMetaRAM = ram.MustNew[uint64, models.Upload](ram.Options{MaximumSize: uploadMetaRAMMaximumSize})
|
||||||
uploadMetaListenerOnce sync.Once
|
|
||||||
uploadMetaListenerCtx context.Context
|
|
||||||
uploadMetaListenerCancel context.CancelFunc
|
|
||||||
uploadMetaListenerDone chan struct{}
|
|
||||||
)
|
)
|
||||||
|
|
||||||
func uploadMetaRedisKey(id uint64) string {
|
func uploadMetaRedisKey(id uint64) string {
|
||||||
@@ -42,120 +33,102 @@ func cloneUpload(u models.Upload) models.Upload {
|
|||||||
return u
|
return u
|
||||||
}
|
}
|
||||||
|
|
||||||
func ensureUploadMetaCacheListener() {
|
// PublishUploadMetaInvalidation broadcasts upload metadata cache eviction.
|
||||||
if cachepkg.Redis == nil {
|
func PublishUploadMetaInvalidation(ctx context.Context, id uint64) {
|
||||||
return
|
if cache := shared.GetCache(ctx); cache != nil {
|
||||||
|
_ = cache.Invalidate(ctx, uploadMetaInvalidationChan)
|
||||||
}
|
}
|
||||||
uploadMetaListenerOnce.Do(startUploadMetaCacheInvalidationListener)
|
EvictUploadMetaLocal(id)
|
||||||
}
|
|
||||||
|
|
||||||
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()
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// GetUploadByID loads upload metadata from RAM, Redis, or the database.
|
// GetUploadByID loads upload metadata from RAM, Redis, or the database.
|
||||||
func GetUploadByID(ctx context.Context, id uint64) (models.Upload, error) {
|
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 {
|
if u, ok := uploadMetaRAM.GetIfPresent(id); ok {
|
||||||
return cloneUpload(u), nil
|
return cloneUpload(u), nil
|
||||||
}
|
}
|
||||||
|
|
||||||
key := uploadMetaRedisKey(id)
|
key := uploadMetaRedisKey(id)
|
||||||
if cachepkg.Redis != nil {
|
|
||||||
|
// 2. Redis L2 Cache
|
||||||
|
if cache := shared.GetCache(ctx); cache != nil {
|
||||||
var u models.Upload
|
var u models.Upload
|
||||||
if err := cachepkg.GetJSON(ctx, key, &u); err == nil {
|
if err := cache.Get(ctx, key, &u); err == nil {
|
||||||
uploadMetaRAM.Set(id, cloneUpload(u))
|
uploadMetaRAM.Set(id, u)
|
||||||
return u, nil
|
return cloneUpload(u), nil
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
var u models.Upload
|
// 3. Database L3 Source of Truth
|
||||||
if err := database.DB(ctx).
|
var upload models.Upload
|
||||||
Where("id = ? AND status IN (?, ?)", id, models.UploadStatusPending, models.UploadStatusUsed).
|
db := shared.GetDB(ctx)
|
||||||
First(&u).Error; err != nil {
|
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
|
return models.Upload{}, err
|
||||||
}
|
}
|
||||||
|
|
||||||
SetUploadMetaCache(ctx, &u)
|
SetUploadMeta(ctx, upload)
|
||||||
return u, nil
|
return cloneUpload(upload), nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// SetUploadMetaCache populates RAM and Redis upload metadata caches.
|
// SetUploadMeta populates RAM and Redis caches with the provided upload metadata.
|
||||||
func SetUploadMetaCache(ctx context.Context, u *models.Upload) {
|
func SetUploadMeta(ctx context.Context, u models.Upload) {
|
||||||
ensureUploadMetaCacheListener()
|
if u.ID == 0 {
|
||||||
|
|
||||||
if u == nil {
|
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
cloned := cloneUpload(u)
|
||||||
cloned := cloneUpload(*u)
|
|
||||||
uploadMetaRAM.Set(u.ID, cloned)
|
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.
|
// EvictUploadMeta evicts an upload metadata record from RAM, Redis, and broadcasts eviction.
|
||||||
func InvalidateUploadMetaCache(ctx context.Context, id uint64) {
|
func EvictUploadMeta(ctx context.Context, id uint64) {
|
||||||
ensureUploadMetaCacheListener()
|
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)
|
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.
|
// ResetUploadMetaCache cleans up local memory cache.
|
||||||
func ResetUploadMetaCacheForTest() {
|
func ResetUploadMetaCache() {
|
||||||
uploadMetaRAM.InvalidateAll()
|
uploadMetaRAM.InvalidateAll()
|
||||||
}
|
}
|
||||||
|
|
||||||
// StopUploadMetaCacheListener stops the Redis Pub/Sub subscription listener and resets the sync.Once guard.
|
// ResetUploadMetaCacheForTest clears the in-memory cache for tests.
|
||||||
func StopUploadMetaCacheListener() {
|
func ResetUploadMetaCacheForTest() {
|
||||||
if uploadMetaListenerCancel != nil {
|
ResetUploadMetaCache()
|
||||||
uploadMetaListenerCancel()
|
|
||||||
if uploadMetaListenerDone != nil {
|
|
||||||
<-uploadMetaListenerDone // 等待 goroutine 退出,保证之后置空 cachepkg.Redis 不再竞争
|
|
||||||
}
|
|
||||||
uploadMetaListenerCancel = nil
|
|
||||||
uploadMetaListenerDone = nil
|
|
||||||
}
|
|
||||||
uploadMetaListenerOnce = sync.Once{}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// 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 (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"encoding/json"
|
|
||||||
"testing"
|
"testing"
|
||||||
"time"
|
|
||||||
|
"gorm.io/gorm"
|
||||||
|
|
||||||
"Wavelet/pkg/testhelper"
|
"Wavelet/pkg/testhelper"
|
||||||
"Wavelet/plugins/domain/upload/models"
|
"Wavelet/plugins/domain/upload/models"
|
||||||
cachepkg "Wavelet/plugins/infra/cache"
|
"Wavelet/plugins/domain/upload/shared"
|
||||||
"gorm.io/gorm"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
func init() {
|
func init() {
|
||||||
testhelper.RegisterCleanup(func() {
|
testhelper.RegisterCleanup(func() {
|
||||||
StopUploadMetaCacheListener()
|
|
||||||
ResetUploadMetaCacheForTest()
|
ResetUploadMetaCacheForTest()
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
@@ -30,7 +28,7 @@ func seedUpload(t *testing.T, dbConn *gorm.DB, upload models.Upload) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestGetUploadByIDLoadsFromDBAndPopulatesCache(t *testing.T) {
|
func TestGetUploadByIDLoadsFromDBAndPopulatesCache(t *testing.T) {
|
||||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
dbConn, cleanup := shared.SetupTestEnv(t)
|
||||||
defer cleanup()
|
defer cleanup()
|
||||||
ResetUploadMetaCacheForTest()
|
ResetUploadMetaCacheForTest()
|
||||||
|
|
||||||
@@ -57,14 +55,6 @@ func TestGetUploadByIDLoadsFromDBAndPopulatesCache(t *testing.T) {
|
|||||||
t.Fatalf("unexpected upload: %+v", got)
|
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 {
|
if err := dbConn.Delete(&models.Upload{}, upload.ID).Error; err != nil {
|
||||||
t.Fatalf("delete upload from db: %v", err)
|
t.Fatalf("delete upload from db: %v", err)
|
||||||
}
|
}
|
||||||
@@ -79,7 +69,7 @@ func TestGetUploadByIDLoadsFromDBAndPopulatesCache(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestGetUploadByIDReadsFromRedisWhenRAMEmpty(t *testing.T) {
|
func TestGetUploadByIDReadsFromRedisWhenRAMEmpty(t *testing.T) {
|
||||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
dbConn, cleanup := shared.SetupTestEnv(t)
|
||||||
defer cleanup()
|
defer cleanup()
|
||||||
ResetUploadMetaCacheForTest()
|
ResetUploadMetaCacheForTest()
|
||||||
|
|
||||||
@@ -114,7 +104,7 @@ func TestGetUploadByIDReadsFromRedisWhenRAMEmpty(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestInvalidateUploadMetaCacheClearsRAMAndRedis(t *testing.T) {
|
func TestInvalidateUploadMetaCacheClearsRAMAndRedis(t *testing.T) {
|
||||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
dbConn, cleanup := shared.SetupTestEnv(t)
|
||||||
defer cleanup()
|
defer cleanup()
|
||||||
ResetUploadMetaCacheForTest()
|
ResetUploadMetaCacheForTest()
|
||||||
|
|
||||||
@@ -136,11 +126,6 @@ func TestInvalidateUploadMetaCacheClearsRAMAndRedis(t *testing.T) {
|
|||||||
|
|
||||||
InvalidateUploadMetaCache(ctx, upload.ID)
|
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)
|
got, err := GetUploadByID(ctx, upload.ID)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("GetUploadByID after invalidate should reload from DB: %v", err)
|
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) {
|
func TestGetUploadByIDSkipsDeletedUploads(t *testing.T) {
|
||||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
dbConn, cleanup := shared.SetupTestEnv(t)
|
||||||
defer cleanup()
|
defer cleanup()
|
||||||
ResetUploadMetaCacheForTest()
|
ResetUploadMetaCacheForTest()
|
||||||
|
|
||||||
@@ -237,51 +159,3 @@ func TestGetUploadByIDSkipsDeletedUploads(t *testing.T) {
|
|||||||
t.Fatal("expected error for deleted upload")
|
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"
|
uploadstats "Wavelet/plugins/domain/upload/stats"
|
||||||
uploadtask "Wavelet/plugins/domain/upload/task"
|
uploadtask "Wavelet/plugins/domain/upload/task"
|
||||||
"Wavelet/plugins/domain/upload/util"
|
"Wavelet/plugins/domain/upload/util"
|
||||||
"Wavelet/plugins/drivers/driver_asynq_worker"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
// HTTP handlers
|
// HTTP handlers
|
||||||
@@ -113,14 +112,3 @@ type RebuildUploadStatsHandler = uploadtask.RebuildUploadStatsHandler
|
|||||||
|
|
||||||
// WarmImageCachePayload is the payload for image cache warmup tasks.
|
// WarmImageCachePayload is the payload for image cache warmup tasks.
|
||||||
type WarmImageCachePayload = uploadtask.WarmImageCachePayload
|
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"
|
"net/http"
|
||||||
"strconv"
|
"strconv"
|
||||||
"strings"
|
"strings"
|
||||||
|
"sync"
|
||||||
|
|
||||||
"Wavelet/core/contracts"
|
"Wavelet/core/contracts"
|
||||||
|
pkgcache "Wavelet/pkg/cache/disk"
|
||||||
|
"Wavelet/pkg/logger"
|
||||||
"Wavelet/pkg/response"
|
"Wavelet/pkg/response"
|
||||||
pkgutil "Wavelet/pkg/util"
|
pkgutil "Wavelet/pkg/util"
|
||||||
"Wavelet/plugins/domain/auth"
|
|
||||||
"Wavelet/plugins/domain/upload/cache"
|
"Wavelet/plugins/domain/upload/cache"
|
||||||
"Wavelet/plugins/domain/upload/models"
|
"Wavelet/plugins/domain/upload/models"
|
||||||
"Wavelet/plugins/domain/upload/shared"
|
"Wavelet/plugins/domain/upload/shared"
|
||||||
uploadstorage "Wavelet/plugins/domain/upload/storage"
|
uploadstorage "Wavelet/plugins/domain/upload/storage"
|
||||||
"Wavelet/plugins/domain/upload/util"
|
"Wavelet/plugins/domain/upload/util"
|
||||||
"Wavelet/plugins/infra/storage/diskcache"
|
|
||||||
|
|
||||||
"Wavelet/pkg/logger"
|
|
||||||
"github.com/gin-gonic/gin"
|
"github.com/gin-gonic/gin"
|
||||||
"golang.org/x/sync/singleflight"
|
"golang.org/x/sync/singleflight"
|
||||||
"gorm.io/gorm"
|
"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 {
|
type compressedImageCacheResult struct {
|
||||||
bytes []byte
|
bytes []byte
|
||||||
@@ -193,13 +204,13 @@ func EnsureCompressedImageCache(
|
|||||||
upload *models.Upload,
|
upload *models.Upload,
|
||||||
quality string,
|
quality string,
|
||||||
) ([]byte, bool, error) {
|
) ([]byte, bool, error) {
|
||||||
cacheStore := diskcache.GetGlobalCache()
|
cacheStore := getGlobalDiskCache()
|
||||||
cacheKey := ImageCompressionCacheKey(upload, quality)
|
cacheKey := ImageCompressionCacheKey(upload, quality)
|
||||||
webpBytes, err := cacheStore.Get(cacheKey)
|
webpBytes, err := cacheStore.Get(cacheKey)
|
||||||
if err == nil {
|
if err == nil {
|
||||||
return webpBytes, true, 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)
|
return nil, false, fmt.Errorf("read compressed image cache: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -220,13 +231,13 @@ func generateCompressedImageCache(
|
|||||||
quality string,
|
quality string,
|
||||||
cacheKey string,
|
cacheKey string,
|
||||||
) (compressedImageCacheResult, error) {
|
) (compressedImageCacheResult, error) {
|
||||||
cacheStore := diskcache.GetGlobalCache()
|
cacheStore := getGlobalDiskCache()
|
||||||
|
|
||||||
webpBytes, err := cacheStore.Get(cacheKey)
|
webpBytes, err := cacheStore.Get(cacheKey)
|
||||||
if err == nil {
|
if err == nil {
|
||||||
return compressedImageCacheResult{bytes: webpBytes, cached: true}, 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)
|
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)
|
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{
|
return compressedImageCacheResult{
|
||||||
bytes: webpBytes,
|
bytes: webpBytes,
|
||||||
err: fmt.Errorf("write compressed image cache: %w", err),
|
err: fmt.Errorf("write compressed image cache: %w", err),
|
||||||
@@ -269,7 +280,11 @@ func serveOriginal(c *gin.Context, upload *models.Upload) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
defer func() { _ = obj.Body.Close() }()
|
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) {
|
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 {
|
if u, ok := pkgutil.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey); ok && u != nil {
|
||||||
currUserID = u.ID
|
currUserID = u.ID
|
||||||
isAdmin = u.IsAdmin
|
isAdmin = u.IsAdmin
|
||||||
} else {
|
} else if authSvc := shared.GetAuthService(c); authSvc != nil {
|
||||||
u, err := auth.GetUserFromRequest(c)
|
u, err := authSvc.GetCurrentUser(c)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
currUserID = u.ID
|
currUserID = u.ID
|
||||||
isAdmin = u.IsAdmin
|
isAdmin = u.IsAdmin
|
||||||
|
} else {
|
||||||
|
return errors.New("unauthorized")
|
||||||
}
|
}
|
||||||
if isAdmin {
|
if isAdmin {
|
||||||
return nil
|
return nil
|
||||||
@@ -312,8 +329,10 @@ func CheckFileAccessPermission(c *gin.Context, upload *models.Upload) error {
|
|||||||
|
|
||||||
if !cache.IsFilePublic(c.Request.Context(), upload.Type) {
|
if !cache.IsFilePublic(c.Request.Context(), upload.Type) {
|
||||||
if _, ok := pkgutil.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey); !ok {
|
if _, ok := pkgutil.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey); !ok {
|
||||||
if _, err := auth.GetUserFromRequest(c); err != nil {
|
if authSvc := shared.GetAuthService(c); authSvc != nil {
|
||||||
return err
|
if _, err := authSvc.GetCurrentUser(c); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -5,22 +5,23 @@ package filesrv
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"bytes"
|
"bytes"
|
||||||
|
"context"
|
||||||
"crypto/sha256"
|
"crypto/sha256"
|
||||||
"encoding/json"
|
|
||||||
"fmt"
|
"fmt"
|
||||||
"image"
|
"image"
|
||||||
"image/color"
|
"image/color"
|
||||||
"image/png"
|
"image/png"
|
||||||
|
"io"
|
||||||
"net/http"
|
"net/http"
|
||||||
"net/http/httptest"
|
"net/http/httptest"
|
||||||
"os"
|
"os"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
|
"sync"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
"github.com/gin-contrib/sessions"
|
"github.com/gin-contrib/sessions"
|
||||||
"github.com/gin-contrib/sessions/cookie"
|
"github.com/gin-contrib/sessions/cookie"
|
||||||
"github.com/gin-gonic/gin"
|
"github.com/gin-gonic/gin"
|
||||||
"gorm.io/gorm"
|
|
||||||
|
|
||||||
"Wavelet/core/contracts"
|
"Wavelet/core/contracts"
|
||||||
"Wavelet/pkg/response"
|
"Wavelet/pkg/response"
|
||||||
@@ -29,21 +30,66 @@ import (
|
|||||||
"Wavelet/plugins/domain/upload/models"
|
"Wavelet/plugins/domain/upload/models"
|
||||||
"Wavelet/plugins/domain/upload/shared"
|
"Wavelet/plugins/domain/upload/shared"
|
||||||
uploadutil "Wavelet/plugins/domain/upload/util"
|
uploadutil "Wavelet/plugins/domain/upload/util"
|
||||||
"Wavelet/plugins/infra/storage/diskcache"
|
|
||||||
"Wavelet/plugins/infra/storage/objectstore"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
func init() {
|
func init() {
|
||||||
testhelper.RegisterCleanup(cache.ResetUploadMetaCacheForTest)
|
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) {
|
func TestServeFileByIDAccessControl(t *testing.T) {
|
||||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
dbConn, cleanup := shared.SetupTestEnv(t)
|
||||||
defer cleanup()
|
defer cleanup()
|
||||||
cache.ResetAccessCaches()
|
cache.ResetAccessCaches()
|
||||||
|
|
||||||
tempDir := t.TempDir()
|
tempDir := t.TempDir()
|
||||||
configureLocalStorageRoot(t, dbConn, tempDir)
|
storageSvc := &localTestStorageService{root: tempDir}
|
||||||
|
shared.SetStorageService(storageSvc)
|
||||||
|
|
||||||
// Create a user in DB
|
// Create a user in DB
|
||||||
user := contracts.UserDTO{
|
user := contracts.UserDTO{
|
||||||
@@ -119,9 +165,6 @@ func TestServeFileByIDAccessControl(t *testing.T) {
|
|||||||
if w.Code != http.StatusOK {
|
if w.Code != http.StatusOK {
|
||||||
t.Fatalf("expected status 200 for public file, got %d", w.Code)
|
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) {
|
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 {
|
if w.Code != http.StatusOK {
|
||||||
t.Fatalf("expected status 200 for authenticated request, got %d", w.Code)
|
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) {
|
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) {
|
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()
|
w := httptest.NewRecorder()
|
||||||
r.ServeHTTP(w, req)
|
r.ServeHTTP(w, req)
|
||||||
|
|
||||||
if w.Code != http.StatusBadRequest {
|
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) {
|
func TestServeFileByIDImageCompression(t *testing.T) {
|
||||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
dbConn, cleanup := shared.SetupTestEnv(t)
|
||||||
defer cleanup()
|
defer cleanup()
|
||||||
cache.ResetAccessCaches()
|
cache.ResetAccessCaches()
|
||||||
|
|
||||||
tempDir := t.TempDir()
|
tempDir := t.TempDir()
|
||||||
configureLocalStorageRoot(t, dbConn, tempDir)
|
storageSvc := &localTestStorageService{root: tempDir}
|
||||||
|
shared.SetStorageService(storageSvc)
|
||||||
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)
|
|
||||||
}
|
|
||||||
}()
|
|
||||||
|
|
||||||
// Create test user
|
// Create test user
|
||||||
user := contracts.UserDTO{
|
user := contracts.UserDTO{
|
||||||
ID: 555,
|
ID: 54321,
|
||||||
Username: "compress_tester",
|
Username: "compress_test_user",
|
||||||
IsActive: true,
|
IsActive: true,
|
||||||
}
|
}
|
||||||
dbConn.Table("w_users").Create(&user)
|
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 := image.NewRGBA(image.Rect(0, 0, 1, 1))
|
||||||
img.Set(0, 0, color.RGBA{R: 255, G: 0, B: 0, A: 255})
|
img.Set(0, 0, color.RGBA{R: 255, G: 0, B: 0, A: 255})
|
||||||
var pngBuf bytes.Buffer
|
var pngBuf bytes.Buffer
|
||||||
@@ -237,7 +267,6 @@ func TestServeFileByIDImageCompression(t *testing.T) {
|
|||||||
if w.Code != http.StatusOK {
|
if w.Code != http.StatusOK {
|
||||||
t.Fatalf("expected status 200, got %d", w.Code)
|
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" {
|
if w.Header().Get("Content-Type") != "image/png" {
|
||||||
t.Errorf("expected Content-Type image/png, got %s", w.Header().Get("Content-Type"))
|
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"
|
"github.com/gin-gonic/gin"
|
||||||
|
|
||||||
"Wavelet/core/contracts"
|
"Wavelet/core/contracts"
|
||||||
"Wavelet/pkg/testhelper"
|
|
||||||
"Wavelet/plugins/domain/upload/models"
|
"Wavelet/plugins/domain/upload/models"
|
||||||
|
"Wavelet/plugins/domain/upload/shared"
|
||||||
)
|
)
|
||||||
|
|
||||||
func TestGetDistinctUploadTypes(t *testing.T) {
|
func TestGetDistinctUploadTypes(t *testing.T) {
|
||||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
dbConn, cleanup := shared.SetupTestEnv(t)
|
||||||
defer cleanup()
|
defer cleanup()
|
||||||
|
|
||||||
user := contracts.UserDTO{ID: 2222, Username: "test_user_2"}
|
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" {
|
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"
|
"strconv"
|
||||||
"strings"
|
"strings"
|
||||||
|
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
"gorm.io/gorm"
|
||||||
|
|
||||||
"Wavelet/core/contracts"
|
"Wavelet/core/contracts"
|
||||||
"Wavelet/pkg/logger"
|
"Wavelet/pkg/logger"
|
||||||
"Wavelet/pkg/response"
|
"Wavelet/pkg/response"
|
||||||
@@ -31,8 +34,6 @@ import (
|
|||||||
"Wavelet/plugins/domain/upload/shared"
|
"Wavelet/plugins/domain/upload/shared"
|
||||||
uploadstorage "Wavelet/plugins/domain/upload/storage"
|
uploadstorage "Wavelet/plugins/domain/upload/storage"
|
||||||
"Wavelet/plugins/domain/upload/util"
|
"Wavelet/plugins/domain/upload/util"
|
||||||
"github.com/gin-gonic/gin"
|
|
||||||
"gorm.io/gorm"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
type batchDownloadRequest struct {
|
type batchDownloadRequest struct {
|
||||||
|
|||||||
@@ -13,20 +13,21 @@ import (
|
|||||||
"net/http"
|
"net/http"
|
||||||
"net/http/httptest"
|
"net/http/httptest"
|
||||||
"os"
|
"os"
|
||||||
|
"path/filepath"
|
||||||
"strconv"
|
"strconv"
|
||||||
"strings"
|
"strings"
|
||||||
|
"sync"
|
||||||
"testing"
|
"testing"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
|
||||||
"Wavelet/core/contracts"
|
"Wavelet/core/contracts"
|
||||||
"Wavelet/pkg/response"
|
"Wavelet/pkg/response"
|
||||||
"Wavelet/pkg/testhelper"
|
|
||||||
"Wavelet/pkg/util"
|
"Wavelet/pkg/util"
|
||||||
"Wavelet/plugins/domain/upload/models"
|
"Wavelet/plugins/domain/upload/models"
|
||||||
"Wavelet/plugins/domain/upload/shared"
|
"Wavelet/plugins/domain/upload/shared"
|
||||||
uploadstats "Wavelet/plugins/domain/upload/stats"
|
uploadstats "Wavelet/plugins/domain/upload/stats"
|
||||||
"Wavelet/plugins/infra/storage/objectstore"
|
|
||||||
"github.com/gin-gonic/gin"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
type testResponse struct {
|
type testResponse struct {
|
||||||
@@ -86,9 +87,6 @@ func createMultipartRequest(t *testing.T, fieldName, fileName string, fileConten
|
|||||||
|
|
||||||
for k, v := range extraFields {
|
for k, v := range extraFields {
|
||||||
err = writer.WriteField(k, v)
|
err = writer.WriteField(k, v)
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("failed to write form field: %v", err)
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
err = writer.Close()
|
err = writer.Close()
|
||||||
@@ -99,51 +97,76 @@ func createMultipartRequest(t *testing.T, fieldName, fileName string, fileConten
|
|||||||
return writer.FormDataContentType(), body
|
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) {
|
func TestUploadFile(t *testing.T) {
|
||||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
dbConn, cleanup := shared.SetupTestEnv(t)
|
||||||
defer cleanup()
|
defer cleanup()
|
||||||
defer func() { _ = os.RemoveAll("uploads") }() // Clean up local files created during tests
|
defer func() { _ = os.RemoveAll("uploads") }() // Clean up local files created during tests
|
||||||
|
|
||||||
authUser := &contracts.UserDTO{ID: 1001, Username: "test_user"}
|
authUser := &contracts.UserDTO{ID: 1001, Username: "test_user"}
|
||||||
router := setupTestRouter(authUser)
|
router := setupTestRouter(authUser)
|
||||||
|
|
||||||
// Mock Storage Client
|
|
||||||
mockFiles := make(map[string][]byte)
|
|
||||||
var putCount int
|
var putCount int
|
||||||
|
mockStorage := &handlerTestStorage{
|
||||||
restoreStorage := objectstore.MockStorage(
|
mockFiles: make(map[string][]byte),
|
||||||
func(ctx context.Context, key string, body io.Reader, size int64, contentType string) error {
|
putCount: &putCount,
|
||||||
data, err := io.ReadAll(body)
|
}
|
||||||
if err != nil {
|
shared.SetStorageService(mockStorage)
|
||||||
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 }
|
|
||||||
}()
|
|
||||||
|
|
||||||
t.Run("upload allowed image file successfully", func(t *testing.T) {
|
t.Run("upload allowed image file successfully", func(t *testing.T) {
|
||||||
putCount = 0
|
putCount = 0
|
||||||
@@ -283,9 +306,6 @@ func TestUploadFile(t *testing.T) {
|
|||||||
})
|
})
|
||||||
|
|
||||||
t.Run("upload in local storage fallback mode", func(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
|
// Seed allowed extensions configuration to allow txt files
|
||||||
dbConn.Table("w_system_configs").Where("key = ?", "upload_allowed_extensions").Update("value", "jpg,png,webp,txt")
|
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) {
|
func TestDownloadFile(t *testing.T) {
|
||||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
dbConn, cleanup := shared.SetupTestEnv(t)
|
||||||
defer cleanup()
|
defer cleanup()
|
||||||
defer func() { _ = os.RemoveAll("uploads") }()
|
defer func() { _ = os.RemoveAll("uploads") }()
|
||||||
|
|
||||||
@@ -395,7 +415,7 @@ func TestDownloadFile(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestListFiles(t *testing.T) {
|
func TestListFiles(t *testing.T) {
|
||||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
dbConn, cleanup := shared.SetupTestEnv(t)
|
||||||
defer cleanup()
|
defer cleanup()
|
||||||
|
|
||||||
authUser := &contracts.UserDTO{ID: 1001, Username: "test_user"}
|
authUser := &contracts.UserDTO{ID: 1001, Username: "test_user"}
|
||||||
@@ -535,7 +555,7 @@ func TestListFiles(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestBatchDownloadFiles(t *testing.T) {
|
func TestBatchDownloadFiles(t *testing.T) {
|
||||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
dbConn, cleanup := shared.SetupTestEnv(t)
|
||||||
defer cleanup()
|
defer cleanup()
|
||||||
defer func() { _ = os.RemoveAll("uploads") }()
|
defer func() { _ = os.RemoveAll("uploads") }()
|
||||||
|
|
||||||
@@ -647,7 +667,7 @@ func TestBatchDownloadFiles(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestUploadAccessModeAccessControl(t *testing.T) {
|
func TestUploadAccessModeAccessControl(t *testing.T) {
|
||||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
dbConn, cleanup := shared.SetupTestEnv(t)
|
||||||
defer cleanup()
|
defer cleanup()
|
||||||
defer func() { _ = os.RemoveAll("uploads") }()
|
defer func() { _ = os.RemoveAll("uploads") }()
|
||||||
|
|
||||||
@@ -734,7 +754,7 @@ func TestUploadAccessModeAccessControl(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestGetFileStats(t *testing.T) {
|
func TestGetFileStats(t *testing.T) {
|
||||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
dbConn, cleanup := shared.SetupTestEnv(t)
|
||||||
defer cleanup()
|
defer cleanup()
|
||||||
|
|
||||||
authUser := &contracts.UserDTO{ID: 1001, Username: "test_user"}
|
authUser := &contracts.UserDTO{ID: 1001, Username: "test_user"}
|
||||||
@@ -844,7 +864,7 @@ func TestGetFileStats(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestUserUploadManagement(t *testing.T) {
|
func TestUserUploadManagement(t *testing.T) {
|
||||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
dbConn, cleanup := shared.SetupTestEnv(t)
|
||||||
defer cleanup()
|
defer cleanup()
|
||||||
|
|
||||||
user1 := &contracts.UserDTO{ID: 1001, Username: "user1"}
|
user1 := &contracts.UserDTO{ID: 1001, Username: "user1"}
|
||||||
|
|||||||
@@ -12,6 +12,8 @@ import (
|
|||||||
"strings"
|
"strings"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
|
"gorm.io/gorm"
|
||||||
|
|
||||||
"Wavelet/pkg/idgen"
|
"Wavelet/pkg/idgen"
|
||||||
"Wavelet/pkg/logger"
|
"Wavelet/pkg/logger"
|
||||||
uploadcache "Wavelet/plugins/domain/upload/cache"
|
uploadcache "Wavelet/plugins/domain/upload/cache"
|
||||||
@@ -20,9 +22,6 @@ import (
|
|||||||
"Wavelet/plugins/domain/upload/shared"
|
"Wavelet/plugins/domain/upload/shared"
|
||||||
uploadstats "Wavelet/plugins/domain/upload/stats"
|
uploadstats "Wavelet/plugins/domain/upload/stats"
|
||||||
uploadstorage "Wavelet/plugins/domain/upload/storage"
|
uploadstorage "Wavelet/plugins/domain/upload/storage"
|
||||||
database "Wavelet/plugins/infra/database"
|
|
||||||
"Wavelet/plugins/infra/storage/objectstore"
|
|
||||||
"gorm.io/gorm"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
func normalizeRequest(req *Request) {
|
func normalizeRequest(req *Request) {
|
||||||
@@ -50,13 +49,16 @@ func resolveAccessMode(uploadType string, explicit *int) int {
|
|||||||
|
|
||||||
func validateAllowedExtension(ctx context.Context, ext string) error {
|
func validateAllowedExtension(ctx context.Context, ext string) error {
|
||||||
var val string
|
var val string
|
||||||
err := database.DB(ctx).Table("w_system_configs").Where("key = ?", "upload_allowed_extensions").Pluck("value", &val).Error
|
db := shared.GetDB(ctx)
|
||||||
if err != nil {
|
if db != nil {
|
||||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
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
|
return nil
|
||||||
}
|
}
|
||||||
logger.WarnF(ctx, "failed to query upload_allowed_extensions: %v", err)
|
|
||||||
return nil
|
|
||||||
}
|
}
|
||||||
if val == "" {
|
if val == "" {
|
||||||
return nil
|
return nil
|
||||||
@@ -97,15 +99,15 @@ func storeObject(ctx context.Context, objectKey string, reader io.Reader, size i
|
|||||||
return "", ErrStorageReadOnly
|
return "", ErrStorageReadOnly
|
||||||
}
|
}
|
||||||
|
|
||||||
driver, backend, err := objectstore.Active(ctx)
|
storageSvc := shared.GetStorage(ctx)
|
||||||
if err != nil {
|
if storageSvc == nil {
|
||||||
logger.ErrorF(ctx, "初始化活动存储失败: %v", err)
|
logger.ErrorF(ctx, "初始化活动存储失败: storage service is nil")
|
||||||
return "", errors.New(shared.ErrSaveFileFailed)
|
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 {
|
if err != nil {
|
||||||
logger.ErrorF(ctx, "写入 %s 存储失败: %v", driver, err)
|
logger.ErrorF(ctx, "写入存储失败: %v", err)
|
||||||
return "", errors.New(shared.ErrSaveFileFailed)
|
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 {
|
func persistUploadRecord(ctx context.Context, upload *models.Upload, objectKey string) error {
|
||||||
if err := createUploadWithStats(ctx, upload); err != nil {
|
if err := createUploadWithStats(ctx, upload); err != nil {
|
||||||
_, backend, backendErr := objectstore.Active(ctx)
|
if storageSvc := shared.GetStorage(ctx); storageSvc != nil {
|
||||||
if backendErr == nil {
|
if deleteErr := storageSvc.Delete(ctx, objectKey); deleteErr != nil {
|
||||||
if deleteErr := backend.Delete(ctx, objectKey); deleteErr != nil {
|
|
||||||
logger.WarnF(ctx, "清理未写入数据库的上传对象失败: %v", deleteErr)
|
logger.WarnF(ctx, "清理未写入数据库的上传对象失败: %v", deleteErr)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
uploadcache.SetUploadMetaCache(ctx, upload)
|
uploadcache.SetUploadMeta(ctx, *upload)
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func createUploadWithStats(ctx context.Context, upload *models.Upload) error {
|
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 {
|
if err := repository.CreateUploadTx(tx, upload); err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -7,9 +7,10 @@ import (
|
|||||||
"context"
|
"context"
|
||||||
"errors"
|
"errors"
|
||||||
|
|
||||||
|
"gorm.io/gorm"
|
||||||
|
|
||||||
"Wavelet/plugins/domain/upload/models"
|
"Wavelet/plugins/domain/upload/models"
|
||||||
"Wavelet/plugins/domain/upload/repository"
|
"Wavelet/plugins/domain/upload/repository"
|
||||||
"gorm.io/gorm"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
// Ingest stores or resolves an upload using the configured policy and side effects.
|
// Ingest stores or resolves an upload using the configured policy and side effects.
|
||||||
|
|||||||
@@ -10,17 +10,95 @@ import (
|
|||||||
"encoding/hex"
|
"encoding/hex"
|
||||||
"io"
|
"io"
|
||||||
"os"
|
"os"
|
||||||
|
"sync"
|
||||||
"testing"
|
"testing"
|
||||||
"time"
|
|
||||||
|
|
||||||
"Wavelet/pkg/testhelper"
|
"Wavelet/core/contracts"
|
||||||
"Wavelet/plugins/domain/upload/models"
|
"Wavelet/plugins/domain/upload/models"
|
||||||
database "Wavelet/plugins/infra/database"
|
"Wavelet/plugins/domain/upload/shared"
|
||||||
"Wavelet/plugins/infra/storage/objectstore"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
|
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) {
|
func TestIngestPolicyCreateIncrementsStats(t *testing.T) {
|
||||||
_, _, cleanup := testhelper.SetupTestEnvironment(t)
|
_, cleanup := shared.SetupTestEnv(t)
|
||||||
defer cleanup()
|
defer cleanup()
|
||||||
ctx := context.Background()
|
ctx := context.Background()
|
||||||
|
|
||||||
@@ -59,75 +137,15 @@ func TestIngestPolicyCreateIncrementsStats(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestIngestPolicyResolveExistingSkipsStatsOnHit(t *testing.T) {
|
func TestIngestPolicyResolveExistingSkipsStatsOnHit(t *testing.T) {
|
||||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
_, cleanup := shared.SetupTestEnv(t)
|
||||||
defer cleanup()
|
defer cleanup()
|
||||||
ctx := context.Background()
|
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)
|
hash := sha256.Sum256(content)
|
||||||
hashStr := hex.EncodeToString(hash[:])
|
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
|
putCount := 0
|
||||||
|
|
||||||
restoreStorage, disableStorage := setupMockStorage(t, &putCount)
|
restoreStorage, disableStorage := setupMockStorage(t, &putCount)
|
||||||
defer restoreStorage()
|
defer restoreStorage()
|
||||||
defer disableStorage()
|
defer disableStorage()
|
||||||
@@ -136,194 +154,201 @@ func TestIngestPolicyDedupNewRecordCreatesSecondRecord(t *testing.T) {
|
|||||||
UserID: 1001,
|
UserID: 1001,
|
||||||
Reader: bytes.NewReader(content),
|
Reader: bytes.NewReader(content),
|
||||||
Size: int64(len(content)),
|
Size: int64(len(content)),
|
||||||
FileName: "first.png",
|
FileName: "first.txt",
|
||||||
MimeType: "image/png",
|
MimeType: "text/plain",
|
||||||
Extension: "png",
|
Extension: "txt",
|
||||||
Hash: hashStr,
|
Hash: hashStr,
|
||||||
Type: "avatar",
|
Type: "attachment",
|
||||||
Policy: PolicyDedupNewRecord,
|
Policy: PolicyCreate,
|
||||||
})
|
})
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("first Ingest returned error: %v", err)
|
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 {
|
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{
|
second, err := Ingest(ctx, Request{
|
||||||
UserID: 1002,
|
UserID: 1002,
|
||||||
Reader: bytes.NewReader(content),
|
Reader: bytes.NewReader(content),
|
||||||
Size: int64(len(content)),
|
Size: int64(len(content)),
|
||||||
FileName: "second.png",
|
FileName: "second.txt",
|
||||||
MimeType: "image/png",
|
MimeType: "text/plain",
|
||||||
Extension: "png",
|
Extension: "txt",
|
||||||
Hash: hashStr,
|
Hash: hashStr,
|
||||||
Type: "avatar",
|
Type: "attachment",
|
||||||
Policy: PolicyDedupNewRecord,
|
Policy: PolicyResolveExisting,
|
||||||
})
|
})
|
||||||
if err != nil {
|
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 {
|
if putCount != 1 {
|
||||||
t.Fatalf("putCount after dedup ingest = %d, want 1", putCount)
|
t.Fatalf("putCount = %d, want 1 after hit with PolicyResolveExisting", 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")
|
|
||||||
}
|
}
|
||||||
|
|
||||||
stats, err := loadTotalStats(ctx)
|
stats, err := loadTotalStats(ctx)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("loadTotalStats returned error: %v", err)
|
t.Fatalf("loadTotalStats returned error: %v", err)
|
||||||
}
|
}
|
||||||
if stats.TotalCount != 0 || stats.TotalSize != 0 {
|
if stats.TotalCount != 1 || stats.TotalSize != int64(len(content)) {
|
||||||
t.Fatalf("loadTotalStats() = count %d size %d, want zero after rolled-back stats", stats.TotalCount, stats.TotalSize)
|
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) {
|
func TestRemoveDecrementsStats(t *testing.T) {
|
||||||
_, _, cleanup := testhelper.SetupTestEnvironment(t)
|
_, cleanup := shared.SetupTestEnv(t)
|
||||||
defer cleanup()
|
defer cleanup()
|
||||||
ctx := context.Background()
|
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)
|
hash := sha256.Sum256(content)
|
||||||
|
|
||||||
restoreStorage, disableStorage := setupMockStorage(t, nil)
|
restoreStorage, disableStorage := setupMockStorage(t, nil)
|
||||||
defer restoreStorage()
|
defer restoreStorage()
|
||||||
defer disableStorage()
|
defer disableStorage()
|
||||||
|
|
||||||
result, err := Ingest(ctx, Request{
|
ingested, err := Ingest(ctx, Request{
|
||||||
UserID: 1001,
|
UserID: 1001,
|
||||||
Reader: bytes.NewReader(content),
|
Reader: bytes.NewReader(content),
|
||||||
Size: int64(len(content)),
|
Size: int64(len(content)),
|
||||||
FileName: "delete-me.png",
|
FileName: "to_remove.txt",
|
||||||
MimeType: "image/png",
|
MimeType: "text/plain",
|
||||||
Extension: "png",
|
Extension: "txt",
|
||||||
Hash: hex.EncodeToString(hash[:]),
|
Hash: hex.EncodeToString(hash[:]),
|
||||||
Type: "generic",
|
Type: "generic",
|
||||||
Policy: PolicyCreate,
|
Policy: PolicyCreate,
|
||||||
})
|
})
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("Ingest returned error: %v", err)
|
t.Fatalf("Ingest: %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
if _, err := Remove(ctx, result.Upload.ID); err != nil {
|
removed, err := Remove(ctx, ingested.Upload.ID)
|
||||||
t.Fatalf("Remove(%d) returned error: %v", result.Upload.ID, err)
|
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)
|
stats, err := loadTotalStats(ctx)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("loadTotalStats returned error: %v", err)
|
t.Fatalf("loadTotalStats: %v", err)
|
||||||
}
|
}
|
||||||
if stats.TotalCount != 0 || stats.TotalSize != 0 {
|
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 {
|
func TestRemoveOwnedEnforcesOwnership(t *testing.T) {
|
||||||
TotalCount int64
|
_, cleanup := shared.SetupTestEnv(t)
|
||||||
TotalSize int64
|
defer cleanup()
|
||||||
}
|
ctx := context.Background()
|
||||||
|
|
||||||
func loadTotalStats(ctx context.Context) (totalStatsSnapshot, error) {
|
content := []byte("owner payload")
|
||||||
var rows []models.UploadStat
|
hash := sha256.Sum256(content)
|
||||||
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
|
|
||||||
}
|
|
||||||
|
|
||||||
func setupMockStorage(t *testing.T, putCount *int) (restore func(), disable func()) {
|
restoreStorage, disableStorage := setupMockStorage(t, nil)
|
||||||
t.Helper()
|
defer restoreStorage()
|
||||||
mockFiles := make(map[string][]byte)
|
defer disableStorage()
|
||||||
restore = objectstore.MockStorage(
|
|
||||||
func(ctx context.Context, key string, body io.Reader, size int64, contentType string) error {
|
ingested, err := Ingest(ctx, Request{
|
||||||
data, err := io.ReadAll(body)
|
UserID: 1001,
|
||||||
if err != nil {
|
Reader: bytes.NewReader(content),
|
||||||
return err
|
Size: int64(len(content)),
|
||||||
}
|
FileName: "owned.txt",
|
||||||
mockFiles[key] = data
|
MimeType: "text/plain",
|
||||||
if putCount != nil {
|
Extension: "txt",
|
||||||
*putCount++
|
Hash: hex.EncodeToString(hash[:]),
|
||||||
}
|
Type: "generic",
|
||||||
return nil
|
Policy: PolicyCreate,
|
||||||
},
|
})
|
||||||
func(ctx context.Context, key string) (*objectstore.Object, error) {
|
if err != nil {
|
||||||
data, ok := mockFiles[key]
|
t.Fatalf("Ingest: %v", err)
|
||||||
if !ok {
|
}
|
||||||
return nil, os.ErrNotExist
|
|
||||||
}
|
if _, err := RemoveOwned(ctx, 2002, ingested.Upload.ID); err == nil {
|
||||||
return &objectstore.Object{
|
t.Fatal("expected ErrForbidden for non-owner RemoveOwned")
|
||||||
Body: io.NopCloser(bytes.NewReader(data)),
|
}
|
||||||
ContentLength: int64(len(data)),
|
|
||||||
ContentType: "application/octet-stream",
|
removed, err := RemoveOwned(ctx, 1001, ingested.Upload.ID)
|
||||||
}, nil
|
if err != nil {
|
||||||
},
|
t.Fatalf("RemoveOwned owner failed: %v", err)
|
||||||
func(ctx context.Context, key string) error {
|
}
|
||||||
delete(mockFiles, key)
|
if removed.Status != models.UploadStatusDeleted {
|
||||||
return nil
|
t.Fatalf("removed status = %q, want deleted", removed.Status)
|
||||||
},
|
|
||||||
)
|
|
||||||
objectstore.IsEnabledFunc = func() bool { return true }
|
|
||||||
objectstore.ResetCache()
|
|
||||||
disable = func() {
|
|
||||||
objectstore.IsEnabledFunc = func() bool { return false }
|
|
||||||
objectstore.ResetCache()
|
|
||||||
}
|
}
|
||||||
return restore, disable
|
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -6,12 +6,13 @@ package ingest
|
|||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
|
|
||||||
|
"gorm.io/gorm"
|
||||||
|
|
||||||
uploadcache "Wavelet/plugins/domain/upload/cache"
|
uploadcache "Wavelet/plugins/domain/upload/cache"
|
||||||
"Wavelet/plugins/domain/upload/models"
|
"Wavelet/plugins/domain/upload/models"
|
||||||
"Wavelet/plugins/domain/upload/repository"
|
"Wavelet/plugins/domain/upload/repository"
|
||||||
|
"Wavelet/plugins/domain/upload/shared"
|
||||||
uploadstats "Wavelet/plugins/domain/upload/stats"
|
uploadstats "Wavelet/plugins/domain/upload/stats"
|
||||||
database "Wavelet/plugins/infra/database"
|
|
||||||
"gorm.io/gorm"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
// Remove soft-deletes an upload and decrements incremental stats.
|
// 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 {
|
func softDeleteUploadWithStats(ctx context.Context, upload *models.Upload) error {
|
||||||
statsSnapshot := *upload
|
statsSnapshot := *upload
|
||||||
if err := database.DB(ctx).Transaction(func(tx *gorm.DB) error {
|
db := shared.GetDB(ctx)
|
||||||
if err := repository.SoftDeleteUploadTx(tx, upload); err != nil {
|
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 err
|
||||||
}
|
}
|
||||||
return uploadstats.ApplyUploadStatsDeltaTx(tx, &statsSnapshot, -1)
|
|
||||||
}); err != nil {
|
|
||||||
return err
|
|
||||||
}
|
}
|
||||||
uploadcache.InvalidateUploadMetaCache(ctx, upload.ID)
|
uploadcache.EvictUploadMeta(ctx, upload.ID)
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -9,14 +9,16 @@ import (
|
|||||||
"embed"
|
"embed"
|
||||||
"reflect"
|
"reflect"
|
||||||
|
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
"github.com/hibiken/asynq"
|
||||||
|
|
||||||
"Wavelet/core"
|
"Wavelet/core"
|
||||||
"Wavelet/core/contracts"
|
"Wavelet/core/contracts"
|
||||||
"Wavelet/core/extpoints"
|
"Wavelet/core/extpoints"
|
||||||
"Wavelet/plugins/domain/upload/filesrv"
|
"Wavelet/plugins/domain/upload/filesrv"
|
||||||
"Wavelet/plugins/domain/upload/handler"
|
"Wavelet/plugins/domain/upload/handler"
|
||||||
|
"Wavelet/plugins/domain/upload/shared"
|
||||||
"Wavelet/plugins/domain/upload/task"
|
"Wavelet/plugins/domain/upload/task"
|
||||||
"github.com/gin-gonic/gin"
|
|
||||||
"github.com/hibiken/asynq"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
//go:embed migrations/*.sql
|
//go:embed migrations/*.sql
|
||||||
@@ -56,7 +58,57 @@ func (p *Plugin) Manifest() core.Manifest {
|
|||||||
|
|
||||||
// Apply registers upload routes, tasks, and settings into the Context.
|
// Apply registers upload routes, tasks, and settings into the Context.
|
||||||
func (p *Plugin) Apply(ctx *core.Context) error {
|
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
|
var authSvc contracts.AuthService
|
||||||
if err := core.Using[contracts.AuthService](ctx, func(svc contracts.AuthService) { authSvc = svc }); err != nil {
|
if err := core.Using[contracts.AuthService](ctx, func(svc contracts.AuthService) { authSvc = svc }); err != nil {
|
||||||
return err
|
return err
|
||||||
|
|||||||
@@ -11,7 +11,7 @@ import (
|
|||||||
|
|
||||||
"Wavelet/pkg/idgen"
|
"Wavelet/pkg/idgen"
|
||||||
"Wavelet/pkg/util"
|
"Wavelet/pkg/util"
|
||||||
database "Wavelet/plugins/infra/database"
|
"Wavelet/plugins/domain/upload/shared"
|
||||||
)
|
)
|
||||||
|
|
||||||
// UploadListFilter filters paginated upload queries.
|
// UploadListFilter filters paginated upload queries.
|
||||||
@@ -28,7 +28,7 @@ type UploadListFilter struct {
|
|||||||
|
|
||||||
// ListUploads returns paginated upload records matching the filter.
|
// ListUploads returns paginated upload records matching the filter.
|
||||||
func ListUploads(ctx context.Context, filter UploadListFilter) (int64, []Upload, error) {
|
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)
|
Where("status != ?", UploadStatusDeleted)
|
||||||
|
|
||||||
if filter.UserID != 0 {
|
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.
|
// GetActiveUploadByID loads a non-deleted upload by ID.
|
||||||
func GetActiveUploadByID(ctx context.Context, id uint64) (Upload, error) {
|
func GetActiveUploadByID(ctx context.Context, id uint64) (Upload, error) {
|
||||||
var upload Upload
|
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{}, err
|
||||||
}
|
}
|
||||||
return upload, nil
|
return upload, nil
|
||||||
@@ -69,7 +69,7 @@ func GetActiveUploadByID(ctx context.Context, id uint64) (Upload, error) {
|
|||||||
// SoftDeleteUpload marks an upload as deleted.
|
// SoftDeleteUpload marks an upload as deleted.
|
||||||
// External modules must use upload.Remove or upload.RemoveOwned; only internal/apps/upload may call this.
|
// External modules must use upload.Remove or upload.RemoveOwned; only internal/apps/upload may call this.
|
||||||
func SoftDeleteUpload(ctx context.Context, upload *Upload) error {
|
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.
|
// 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 {
|
if len(updates) == 0 {
|
||||||
return nil
|
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.
|
// ListDistinctUploadTypes returns all distinct non-empty upload business types.
|
||||||
func ListDistinctUploadTypes(ctx context.Context) ([]string, error) {
|
func ListDistinctUploadTypes(ctx context.Context) ([]string, error) {
|
||||||
var types []string
|
var types []string
|
||||||
if err := database.DB(ctx).Model(&Upload{}).
|
if err := shared.GetDB(ctx).Model(&Upload{}).
|
||||||
Where("type IS NOT NULL AND type != ''").
|
Where("type IS NOT NULL AND type != ''").
|
||||||
Distinct().
|
Distinct().
|
||||||
Pluck("type", &types).Error; err != nil {
|
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.
|
// FindReusableUploadByHash finds an existing upload with the same hash and size.
|
||||||
func FindReusableUploadByHash(ctx context.Context, hash string, size int64) (Upload, error) {
|
func FindReusableUploadByHash(ctx context.Context, hash string, size int64) (Upload, error) {
|
||||||
var existing Upload
|
var existing Upload
|
||||||
err := database.DB(ctx).
|
err := shared.GetDB(ctx).
|
||||||
Where("hash = ? AND file_size = ? AND status IN (?, ?)", hash, size, UploadStatusPending, UploadStatusUsed).
|
Where("hash = ? AND file_size = ? AND status IN (?, ?)", hash, size, UploadStatusPending, UploadStatusUsed).
|
||||||
First(&existing).Error
|
First(&existing).Error
|
||||||
return existing, err
|
return existing, err
|
||||||
@@ -108,7 +108,7 @@ func FindReusableUploadByHash(ctx context.Context, hash string, size int64) (Upl
|
|||||||
|
|
||||||
// CreateUpload persists a new upload record.
|
// CreateUpload persists a new upload record.
|
||||||
func CreateUpload(ctx context.Context, upload *Upload) error {
|
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.
|
// 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.
|
// ListUploadsByIDs returns active uploads matching the given IDs.
|
||||||
func ListUploadsByIDs(ctx context.Context, ids []uint64) ([]Upload, error) {
|
func ListUploadsByIDs(ctx context.Context, ids []uint64) ([]Upload, error) {
|
||||||
var uploads []Upload
|
var uploads []Upload
|
||||||
if err := database.DB(ctx).
|
if err := shared.GetDB(ctx).
|
||||||
Where("id IN ? AND status IN (?, ?)", ids, UploadStatusPending, UploadStatusUsed).
|
Where("id IN ? AND status IN (?, ?)", ids, UploadStatusPending, UploadStatusUsed).
|
||||||
Find(&uploads).Error; err != nil {
|
Find(&uploads).Error; err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
@@ -134,13 +134,13 @@ func ListUploadsByIDs(ctx context.Context, ids []uint64) ([]Upload, error) {
|
|||||||
//
|
//
|
||||||
//nolint:revive
|
//nolint:revive
|
||||||
func UploadQuery(ctx context.Context) *gorm.DB {
|
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.
|
// ListUploadStats returns all upload statistics rows.
|
||||||
func ListUploadStats(ctx context.Context) ([]UploadStat, error) {
|
func ListUploadStats(ctx context.Context) ([]UploadStat, error) {
|
||||||
var stats []UploadStat
|
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 nil, err
|
||||||
}
|
}
|
||||||
return stats, nil
|
return stats, nil
|
||||||
|
|||||||
@@ -8,10 +8,11 @@ import (
|
|||||||
"context"
|
"context"
|
||||||
"strings"
|
"strings"
|
||||||
|
|
||||||
|
"gorm.io/gorm"
|
||||||
|
|
||||||
"Wavelet/pkg/util"
|
"Wavelet/pkg/util"
|
||||||
"Wavelet/plugins/domain/upload/models"
|
"Wavelet/plugins/domain/upload/models"
|
||||||
database "Wavelet/plugins/infra/database"
|
"Wavelet/plugins/domain/upload/shared"
|
||||||
"gorm.io/gorm"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
// UploadListFilter filters paginated upload queries.
|
// UploadListFilter filters paginated upload queries.
|
||||||
@@ -26,7 +27,7 @@ type UploadListFilter struct {
|
|||||||
|
|
||||||
// ListUploads returns paginated upload records matching the filter.
|
// ListUploads returns paginated upload records matching the filter.
|
||||||
func ListUploads(ctx context.Context, filter UploadListFilter) (int64, []models.Upload, error) {
|
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)
|
Where("status != ?", models.UploadStatusDeleted)
|
||||||
|
|
||||||
if filter.UserID != 0 {
|
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.
|
// GetActiveUploadByID loads a non-deleted upload by ID.
|
||||||
func GetActiveUploadByID(ctx context.Context, id uint64) (models.Upload, error) {
|
func GetActiveUploadByID(ctx context.Context, id uint64) (models.Upload, error) {
|
||||||
var upload models.Upload
|
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 models.Upload{}, err
|
||||||
}
|
}
|
||||||
return upload, nil
|
return upload, nil
|
||||||
@@ -66,7 +67,7 @@ func GetActiveUploadByID(ctx context.Context, id uint64) (models.Upload, error)
|
|||||||
|
|
||||||
// SoftDeleteUpload marks an upload as deleted.
|
// SoftDeleteUpload marks an upload as deleted.
|
||||||
func SoftDeleteUpload(ctx context.Context, upload *models.Upload) error {
|
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.
|
// 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 {
|
if len(updates) == 0 {
|
||||||
return nil
|
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.
|
// ListDistinctUploadTypes returns all distinct non-empty upload business types.
|
||||||
func ListDistinctUploadTypes(ctx context.Context) ([]string, error) {
|
func ListDistinctUploadTypes(ctx context.Context) ([]string, error) {
|
||||||
var types []string
|
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 != ''").
|
Where("type IS NOT NULL AND type != ''").
|
||||||
Distinct().
|
Distinct().
|
||||||
Pluck("type", &types).Error; err != nil {
|
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.
|
// FindReusableUploadByHash finds an existing upload with the same hash and size.
|
||||||
func FindReusableUploadByHash(ctx context.Context, hash string, size int64) (models.Upload, error) {
|
func FindReusableUploadByHash(ctx context.Context, hash string, size int64) (models.Upload, error) {
|
||||||
var existing models.Upload
|
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).
|
Where("hash = ? AND file_size = ? AND status IN (?, ?)", hash, size, models.UploadStatusPending, models.UploadStatusUsed).
|
||||||
First(&existing).Error
|
First(&existing).Error
|
||||||
return existing, err
|
return existing, err
|
||||||
@@ -105,7 +106,7 @@ func FindReusableUploadByHash(ctx context.Context, hash string, size int64) (mod
|
|||||||
|
|
||||||
// CreateUpload persists a new upload record.
|
// CreateUpload persists a new upload record.
|
||||||
func CreateUpload(ctx context.Context, upload *models.Upload) error {
|
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.
|
// 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.
|
// ListUploadsByIDs returns active uploads matching the given IDs.
|
||||||
func ListUploadsByIDs(ctx context.Context, ids []uint64) ([]models.Upload, error) {
|
func ListUploadsByIDs(ctx context.Context, ids []uint64) ([]models.Upload, error) {
|
||||||
var uploads []models.Upload
|
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).
|
Where("id IN ? AND status IN (?, ?)", ids, models.UploadStatusPending, models.UploadStatusUsed).
|
||||||
Find(&uploads).Error; err != nil {
|
Find(&uploads).Error; err != nil {
|
||||||
return nil, err
|
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.
|
// UploadQuery returns a scoped GORM query for uploads.
|
||||||
func UploadQuery(ctx context.Context) *gorm.DB {
|
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.
|
// ListUploadStats returns all upload statistics rows.
|
||||||
func ListUploadStats(ctx context.Context) ([]models.UploadStat, error) {
|
func ListUploadStats(ctx context.Context) ([]models.UploadStat, error) {
|
||||||
var stats []models.UploadStat
|
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 nil, err
|
||||||
}
|
}
|
||||||
return stats, nil
|
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"
|
"context"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"Wavelet/pkg/logger"
|
|
||||||
"Wavelet/plugins/domain/upload/models"
|
|
||||||
database "Wavelet/plugins/infra/database"
|
|
||||||
"gorm.io/gorm"
|
"gorm.io/gorm"
|
||||||
"gorm.io/gorm/clause"
|
"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.
|
// 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.
|
// RebuildUploadStats rebuilds all incremental stats from current upload records.
|
||||||
func RebuildUploadStats(ctx context.Context) error {
|
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 {
|
if err := tx.Where("1 = 1").Delete(&models.UploadStat{}).Error; err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
@@ -49,7 +50,7 @@ func applyUploadStatsDelta(ctx context.Context, upload *models.Upload, sign int6
|
|||||||
if upload == nil || !isActiveUploadStatus(upload.Status) {
|
if upload == nil || !isActiveUploadStatus(upload.Status) {
|
||||||
return nil
|
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)
|
return ApplyUploadStatsDeltaTx(tx, upload, sign)
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -8,15 +8,36 @@ import (
|
|||||||
"testing"
|
"testing"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
|
"gorm.io/gorm"
|
||||||
|
|
||||||
"Wavelet/pkg/testhelper"
|
"Wavelet/pkg/testhelper"
|
||||||
"Wavelet/plugins/domain/upload/models"
|
"Wavelet/plugins/domain/upload/models"
|
||||||
database "Wavelet/plugins/infra/database"
|
"Wavelet/plugins/domain/upload/shared"
|
||||||
"gorm.io/gorm"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
|
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) {
|
func TestApplyUploadStatsDeltaTxWithinTransaction(t *testing.T) {
|
||||||
_, _, cleanup := testhelper.SetupTestEnvironment(t)
|
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||||
defer cleanup()
|
shared.SetDBService(&mockDBService{db: dbConn})
|
||||||
|
defer func() {
|
||||||
|
shared.SetDBService(nil)
|
||||||
|
cleanup()
|
||||||
|
}()
|
||||||
ctx := context.Background()
|
ctx := context.Background()
|
||||||
|
|
||||||
upload := &models.Upload{
|
upload := &models.Upload{
|
||||||
@@ -29,7 +50,7 @@ func TestApplyUploadStatsDeltaTxWithinTransaction(t *testing.T) {
|
|||||||
CreatedAt: time.Now(),
|
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)
|
return ApplyUploadStatsDeltaTx(tx, upload, 1)
|
||||||
}); err != nil {
|
}); err != nil {
|
||||||
t.Fatalf("ApplyUploadStatsDeltaTx returned error: %v", err)
|
t.Fatalf("ApplyUploadStatsDeltaTx returned error: %v", err)
|
||||||
@@ -45,8 +66,12 @@ func TestApplyUploadStatsDeltaTxWithinTransaction(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestApplyUploadStatsAddAndRemove(t *testing.T) {
|
func TestApplyUploadStatsAddAndRemove(t *testing.T) {
|
||||||
_, _, cleanup := testhelper.SetupTestEnvironment(t)
|
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||||
defer cleanup()
|
shared.SetDBService(&mockDBService{db: dbConn})
|
||||||
|
defer func() {
|
||||||
|
shared.SetDBService(nil)
|
||||||
|
cleanup()
|
||||||
|
}()
|
||||||
ctx := context.Background()
|
ctx := context.Background()
|
||||||
|
|
||||||
upload := &models.Upload{
|
upload := &models.Upload{
|
||||||
@@ -90,7 +115,7 @@ type uploadStatsSnapshot struct {
|
|||||||
|
|
||||||
func loadUploadStats(ctx context.Context) (uploadStatsSnapshot, error) {
|
func loadUploadStats(ctx context.Context) (uploadStatsSnapshot, error) {
|
||||||
var rows []models.UploadStat
|
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
|
return uploadStatsSnapshot{}, err
|
||||||
}
|
}
|
||||||
if len(rows) == 0 {
|
if len(rows) == 0 {
|
||||||
|
|||||||
@@ -9,15 +9,14 @@ import (
|
|||||||
"sync"
|
"sync"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
|
"Wavelet/core/contracts"
|
||||||
"Wavelet/plugins/domain/upload/shared"
|
"Wavelet/plugins/domain/upload/shared"
|
||||||
"Wavelet/plugins/drivers/driver_asynq_worker"
|
|
||||||
"Wavelet/plugins/infra/storage/objectstore"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
// MigrationAccessState captures cached migration maintenance state.
|
// MigrationAccessState captures cached migration maintenance state.
|
||||||
type MigrationAccessState struct {
|
type MigrationAccessState struct {
|
||||||
ReadOnly bool
|
ReadOnly bool
|
||||||
Target objectstore.Config
|
Target contracts.StorageConfigDTO
|
||||||
HasTarget bool
|
HasTarget bool
|
||||||
TargetErr error
|
TargetErr error
|
||||||
LoadErr error
|
LoadErr error
|
||||||
@@ -65,14 +64,14 @@ func buildMigrationAccessState(ctx context.Context) MigrationAccessState {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return MigrationAccessState{LoadErr: err, ReadOnly: true}
|
return MigrationAccessState{LoadErr: err, ReadOnly: true}
|
||||||
}
|
}
|
||||||
if !ok {
|
if !ok || execution == nil {
|
||||||
return MigrationAccessState{}
|
return MigrationAccessState{}
|
||||||
}
|
}
|
||||||
|
|
||||||
state := 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
|
return state
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -10,33 +10,47 @@ import (
|
|||||||
"fmt"
|
"fmt"
|
||||||
"strings"
|
"strings"
|
||||||
|
|
||||||
"Wavelet/plugins/drivers/driver_asynq_worker"
|
"gorm.io/gorm"
|
||||||
"Wavelet/plugins/infra/storage/objectstore"
|
|
||||||
|
"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"
|
const StorageMigrationTask = "storage:migrate"
|
||||||
|
|
||||||
// LatestMigrationExecution returns the most recent storage migration task execution.
|
// LatestMigrationExecution returns the most recent storage migration task execution.
|
||||||
func LatestMigrationExecution(ctx context.Context) (*driver_asynq_worker.TaskExecution, bool, error) {
|
func LatestMigrationExecution(ctx context.Context) (*contracts.TaskExecutionDTO, bool, error) {
|
||||||
return driver_asynq_worker.GetLatestTaskExecutionByTaskType(ctx, StorageMigrationTask)
|
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.
|
// 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)) == "" {
|
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 {
|
var raw struct {
|
||||||
Target json.RawMessage `json:"target"`
|
Target json.RawMessage `json:"target"`
|
||||||
}
|
}
|
||||||
if err := json.Unmarshal(payload, &raw); err != nil {
|
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 {
|
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
|
var targetBytes []byte
|
||||||
@@ -47,34 +61,57 @@ func ParseMigrationTargetConfig(ctx context.Context, payload []byte) (objectstor
|
|||||||
targetBytes = raw.Target
|
targetBytes = raw.Target
|
||||||
}
|
}
|
||||||
|
|
||||||
var target objectstore.Config
|
var target contracts.StorageConfigDTO
|
||||||
if err := json.Unmarshal(targetBytes, &target); err != nil {
|
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
|
return target, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// NormalizeMigrationPayload validates and normalizes a storage migration payload.
|
// 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)
|
target, err := ParseMigrationTargetConfig(ctx, payload)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, objectstore.Config{}, err
|
return nil, contracts.StorageConfigDTO{}, err
|
||||||
}
|
}
|
||||||
type storageMigrationPayload struct {
|
raw, err := json.Marshal(struct {
|
||||||
Target objectstore.Config `json:"target"`
|
Target contracts.StorageConfigDTO `json:"target"`
|
||||||
}
|
}{Target: target})
|
||||||
normalized, err := json.Marshal(storageMigrationPayload{Target: target})
|
|
||||||
if err != nil {
|
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