mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-04 23:16:37 +08:00
refactor(layout): consolidate backend codebase into backend/ package and clean root directory
- Moved cmd/, core/, plugins/, pkg/, downstream/, and main.go into backend/ directory - Batch updated all Go source files to import github.com/Rain-kl/Wavelet/backend/... - Updated Makefile, scripts/swagger.sh, architecture guards, and platform skills - Passed all quality gates (100% tests, 0 lint issues, clean build)
This commit is contained in:
@@ -0,0 +1,24 @@
|
||||
// Copyright 2025 linux.do
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package cmd 提供 CLI 命令入口
|
||||
package cmd
|
||||
|
||||
import (
|
||||
"log"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/backend/core"
|
||||
"github.com/spf13/cobra"
|
||||
)
|
||||
|
||||
var allCmd = &cobra.Command{
|
||||
Use: "all",
|
||||
Short: "以融合模式同时启动 API、Worker 和 Scheduler",
|
||||
Run: func(_ *cobra.Command, _ []string) {
|
||||
app := newWaveletApp(core.ProfileAll)
|
||||
if err := app.Run(); err != nil {
|
||||
log.Fatalf("[All] run failed: %v\n", err)
|
||||
}
|
||||
},
|
||||
}
|
||||
@@ -0,0 +1,23 @@
|
||||
// Copyright 2025 linux.do
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package cmd
|
||||
|
||||
import (
|
||||
"log"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/backend/core"
|
||||
"github.com/spf13/cobra"
|
||||
)
|
||||
|
||||
var apiCmd = &cobra.Command{
|
||||
Use: "api",
|
||||
Short: "wavelet API",
|
||||
Run: func(_ *cobra.Command, _ []string) {
|
||||
app := newWaveletApp(core.ProfileAPI)
|
||||
if err := app.Run(); err != nil {
|
||||
log.Fatalf("[API] run failed: %v\n", err)
|
||||
}
|
||||
},
|
||||
}
|
||||
@@ -0,0 +1,253 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package cmd
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"fmt"
|
||||
"log"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/backend/core"
|
||||
"github.com/Rain-kl/Wavelet/backend/core/contracts"
|
||||
"github.com/Rain-kl/Wavelet/backend/pkg/config"
|
||||
"github.com/Rain-kl/Wavelet/backend/plugins/domain/admin"
|
||||
"github.com/Rain-kl/Wavelet/backend/plugins/domain/auth"
|
||||
"github.com/Rain-kl/Wavelet/backend/plugins/domain/cap"
|
||||
"github.com/Rain-kl/Wavelet/backend/plugins/domain/message_gateway"
|
||||
"github.com/Rain-kl/Wavelet/backend/plugins/domain/risk_control"
|
||||
"github.com/Rain-kl/Wavelet/backend/plugins/domain/system"
|
||||
"github.com/Rain-kl/Wavelet/backend/plugins/domain/upload"
|
||||
"github.com/Rain-kl/Wavelet/backend/plugins/domain/user"
|
||||
"github.com/Rain-kl/Wavelet/backend/plugins/drivers/driver_asynq_cron"
|
||||
"github.com/Rain-kl/Wavelet/backend/plugins/drivers/driver_asynq_worker"
|
||||
"github.com/Rain-kl/Wavelet/backend/plugins/drivers/driver_http"
|
||||
"github.com/Rain-kl/Wavelet/backend/plugins/infra/cache"
|
||||
infradb "github.com/Rain-kl/Wavelet/backend/plugins/infra/database"
|
||||
"github.com/Rain-kl/Wavelet/backend/plugins/infra/logger"
|
||||
"github.com/Rain-kl/Wavelet/backend/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.
|
||||
//
|
||||
//nolint:contextcheck
|
||||
func newWaveletApp(profile core.Profile) *core.App {
|
||||
app := core.NewApp(
|
||||
core.WithProfile(profile),
|
||||
core.WithShutdownTimeout(time.Duration(config.Config.App.GracefulShutdownTimeout)*time.Second),
|
||||
)
|
||||
|
||||
// 1. Register standard infrastructure plugins
|
||||
app.Use(
|
||||
infradb.New(),
|
||||
cache.New(),
|
||||
logger.New(),
|
||||
storage.New(),
|
||||
)
|
||||
|
||||
// 2. Register all 8 domain business plugins
|
||||
app.Use(
|
||||
auth.New(),
|
||||
user.New(),
|
||||
message_gateway.New(),
|
||||
risk_control.New(),
|
||||
admin.New(),
|
||||
upload.New(),
|
||||
cap.New(),
|
||||
system.New(),
|
||||
)
|
||||
|
||||
// 3. Bind Goose migration engine
|
||||
app.SetMigrationEngine(&gooseEngine{})
|
||||
|
||||
// 4. Mount runtime drivers for each aspect
|
||||
app.Use(
|
||||
driver_http.New(driver_http.WithAddr(config.Config.App.Addr)),
|
||||
driver_asynq_worker.New(),
|
||||
driver_asynq_cron.New(),
|
||||
)
|
||||
|
||||
return app
|
||||
}
|
||||
|
||||
// ─── Schema Version Store ──────────────────────────────────────────────────────
|
||||
|
||||
// sharedStore implements database.Store using a single w_schema_versions table.
|
||||
// All plugins share this table, with plugin_id as the discriminator.
|
||||
//
|
||||
// Schema:
|
||||
//
|
||||
// w_schema_versions (
|
||||
// plugin_id VARCHAR(64) NOT NULL,
|
||||
// version_id BIGINT NOT NULL,
|
||||
// applied_at TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP,
|
||||
// PRIMARY KEY (plugin_id, version_id)
|
||||
// )
|
||||
type sharedStore struct {
|
||||
pluginID string
|
||||
dialect string // "postgres" or "sqlite3"
|
||||
}
|
||||
|
||||
func (s *sharedStore) Tablename() string { return "w_schema_versions" }
|
||||
|
||||
func (s *sharedStore) CreateVersionTable(ctx context.Context, db goosedb.DBTxConn) error {
|
||||
_, err := db.ExecContext(ctx, `CREATE TABLE IF NOT EXISTS w_schema_versions (
|
||||
plugin_id VARCHAR(64) NOT NULL,
|
||||
version_id BIGINT NOT NULL,
|
||||
applied_at TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP,
|
||||
PRIMARY KEY (plugin_id, version_id)
|
||||
)`)
|
||||
return err
|
||||
}
|
||||
|
||||
//nolint:mnd
|
||||
func (s *sharedStore) Insert(ctx context.Context, db goosedb.DBTxConn, req goosedb.InsertRequest) error {
|
||||
p := s.placeholder
|
||||
_, err := db.ExecContext(ctx,
|
||||
fmt.Sprintf("INSERT INTO w_schema_versions (plugin_id, version_id) VALUES (%s, %s) ON CONFLICT (plugin_id, version_id) DO NOTHING", p(1), p(2)),
|
||||
s.pluginID, req.Version)
|
||||
return err
|
||||
}
|
||||
|
||||
//nolint:mnd
|
||||
func (s *sharedStore) Delete(ctx context.Context, db goosedb.DBTxConn, version int64) error {
|
||||
p := s.placeholder
|
||||
_, err := db.ExecContext(ctx,
|
||||
fmt.Sprintf("DELETE FROM w_schema_versions WHERE plugin_id = %s AND version_id = %s", p(1), p(2)),
|
||||
s.pluginID, version)
|
||||
return err
|
||||
}
|
||||
|
||||
//nolint:mnd
|
||||
func (s *sharedStore) GetMigration(ctx context.Context, db goosedb.DBTxConn, version int64) (*goosedb.GetMigrationResult, error) {
|
||||
p := s.placeholder
|
||||
var t time.Time
|
||||
err := db.QueryRowContext(ctx,
|
||||
fmt.Sprintf("SELECT applied_at FROM w_schema_versions WHERE plugin_id = %s AND version_id = %s", p(1), p(2)),
|
||||
s.pluginID, version).Scan(&t)
|
||||
if err == sql.ErrNoRows {
|
||||
return nil, goosedb.ErrVersionNotFound
|
||||
}
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &goosedb.GetMigrationResult{Timestamp: t, IsApplied: true}, nil
|
||||
}
|
||||
|
||||
func (s *sharedStore) GetLatestVersion(ctx context.Context, db goosedb.DBTxConn) (int64, error) {
|
||||
p := s.placeholder
|
||||
var version int64
|
||||
err := db.QueryRowContext(ctx,
|
||||
fmt.Sprintf("SELECT COALESCE(MAX(version_id), 0) FROM w_schema_versions WHERE plugin_id = %s", p(1)),
|
||||
s.pluginID).Scan(&version)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return version, nil
|
||||
}
|
||||
|
||||
func (s *sharedStore) ListMigrations(ctx context.Context, db goosedb.DBTxConn) ([]*goosedb.ListMigrationsResult, error) {
|
||||
p := s.placeholder
|
||||
rows, err := db.QueryContext(ctx,
|
||||
fmt.Sprintf("SELECT version_id, TRUE FROM w_schema_versions WHERE plugin_id = %s ORDER BY version_id DESC", p(1)),
|
||||
s.pluginID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer func() { _ = rows.Close() }()
|
||||
|
||||
var results []*goosedb.ListMigrationsResult
|
||||
for rows.Next() {
|
||||
var r goosedb.ListMigrationsResult
|
||||
r.IsApplied = true
|
||||
if err := rows.Scan(&r.Version); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
results = append(results, &r)
|
||||
}
|
||||
return results, rows.Err()
|
||||
}
|
||||
|
||||
func (s *sharedStore) placeholder(n int) string {
|
||||
if s.dialect == "postgres" {
|
||||
return fmt.Sprintf("$%d", n)
|
||||
}
|
||||
return "?"
|
||||
}
|
||||
|
||||
// ─── Migration Engine ──────────────────────────────────────────────────────────
|
||||
|
||||
// gooseEngine implements core.MigrationEngine by iterating all plugin-registered
|
||||
// migration entries and applying each plugin's migrations against the shared DB.
|
||||
//
|
||||
// Each plugin owns its own `migrations/*.sql` directory, embedded via go:embed
|
||||
// and registered via ctx.Migrations().Register(pluginID, embedFS).
|
||||
//
|
||||
// Version tracking: all plugins share a single w_schema_versions table with
|
||||
// plugin_id as the discriminator column. Querying this table shows the current
|
||||
// migration version of every plugin at a glance.
|
||||
type gooseEngine struct{}
|
||||
|
||||
func (e *gooseEngine) Migrate(ctx *core.Context, entries []core.MigrationEntry) error {
|
||||
if len(entries) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
// Resolve DBService from the IoC container.
|
||||
var dbSvc contracts.DBService
|
||||
if err := core.Using[contracts.DBService](ctx, func(svc contracts.DBService) {
|
||||
dbSvc = svc
|
||||
}); err != nil {
|
||||
return fmt.Errorf("migration: resolve DBService: %w", err)
|
||||
}
|
||||
|
||||
gormDB := dbSvc.GORM()
|
||||
if gormDB == nil {
|
||||
return fmt.Errorf("migration: DBService.GORM() returned nil")
|
||||
}
|
||||
|
||||
sqlDB, err := gormDB.DB()
|
||||
if err != nil {
|
||||
return fmt.Errorf("migration: get underlying DB from GORM: %w", err)
|
||||
}
|
||||
|
||||
dialect := gooseDialect()
|
||||
dialectStr := string(dialect)
|
||||
|
||||
for _, entry := range entries {
|
||||
store := &sharedStore{
|
||||
pluginID: entry.PluginID,
|
||||
dialect: dialectStr,
|
||||
}
|
||||
|
||||
provider, err := goose.NewProvider(dialect, sqlDB, entry.FS, goose.WithStore(store))
|
||||
if err != nil {
|
||||
return fmt.Errorf("migration %s: create provider: %w", entry.PluginID, err)
|
||||
}
|
||||
|
||||
results, err := provider.Up(context.Background())
|
||||
if err != nil {
|
||||
return fmt.Errorf("migration %s: apply %w", entry.PluginID, err)
|
||||
}
|
||||
|
||||
if len(results) > 0 {
|
||||
log.Printf("[migrate] %s: applied %d migration(s)", entry.PluginID, len(results))
|
||||
} else {
|
||||
log.Printf("[migrate] %s: up to date", entry.PluginID)
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// gooseDialect returns the goose dialect based on the configured database engine.
|
||||
func gooseDialect() goose.Dialect {
|
||||
if !config.Config.Database.Enabled {
|
||||
return goose.DialectSQLite3
|
||||
}
|
||||
return goose.DialectPostgres
|
||||
}
|
||||
@@ -0,0 +1,82 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package cmd
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/backend/core"
|
||||
)
|
||||
|
||||
func TestNewWaveletAppProfiles(t *testing.T) {
|
||||
profiles := []core.Profile{
|
||||
core.ProfileAPI,
|
||||
core.ProfileWorker,
|
||||
core.ProfileSchedule,
|
||||
core.ProfileAll,
|
||||
}
|
||||
|
||||
for _, prof := range profiles {
|
||||
t.Run(string(prof), func(t *testing.T) {
|
||||
app := newWaveletApp(prof)
|
||||
require.NotNil(t, app)
|
||||
assert.Equal(t, prof, app.Profile())
|
||||
|
||||
// Verify 4 infra plugins + 8 domain plugins + 3 driver plugins registered
|
||||
plugins := app.Plugins()
|
||||
assert.Len(t, plugins, 15)
|
||||
|
||||
// Verify each standard infra plugin is registered
|
||||
_, ok := app.Plugin("database")
|
||||
assert.True(t, ok, "database plugin missing")
|
||||
|
||||
_, ok = app.Plugin("cache")
|
||||
assert.True(t, ok, "cache plugin missing")
|
||||
|
||||
_, ok = app.Plugin("logger")
|
||||
assert.True(t, ok, "logger plugin missing")
|
||||
|
||||
_, ok = app.Plugin("storage")
|
||||
assert.True(t, ok, "storage plugin missing")
|
||||
|
||||
// Verify domain plugins
|
||||
_, ok = app.Plugin("auth")
|
||||
assert.True(t, ok, "auth plugin missing")
|
||||
|
||||
_, ok = app.Plugin("user")
|
||||
assert.True(t, ok, "user plugin missing")
|
||||
|
||||
_, ok = app.Plugin("message_gateway")
|
||||
assert.True(t, ok, "message_gateway plugin missing")
|
||||
|
||||
_, ok = app.Plugin("risk_control")
|
||||
assert.True(t, ok, "risk_control plugin missing")
|
||||
|
||||
_, ok = app.Plugin("admin")
|
||||
assert.True(t, ok, "admin plugin missing")
|
||||
|
||||
_, ok = app.Plugin("upload")
|
||||
assert.True(t, ok, "upload plugin missing")
|
||||
|
||||
_, ok = app.Plugin("cap")
|
||||
assert.True(t, ok, "cap plugin missing")
|
||||
|
||||
_, ok = app.Plugin("system")
|
||||
assert.True(t, ok, "system plugin missing")
|
||||
|
||||
// Verify driver plugins
|
||||
_, ok = app.Plugin("driver_http")
|
||||
assert.True(t, ok, "http driver missing")
|
||||
|
||||
_, ok = app.Plugin("driver_asynq_worker")
|
||||
assert.True(t, ok, "worker driver missing")
|
||||
|
||||
_, ok = app.Plugin("driver_asynq_cron")
|
||||
assert.True(t, ok, "scheduler driver missing")
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,55 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package cmd provides CLI command entry points.
|
||||
//
|
||||
//nolint:unused
|
||||
package cmd
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"log"
|
||||
"runtime"
|
||||
"strings"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/backend/pkg/buildinfo"
|
||||
"github.com/Rain-kl/Wavelet/backend/pkg/config"
|
||||
)
|
||||
|
||||
//nolint:unused // startup banner formatting utilities
|
||||
type startupState struct {
|
||||
mode string
|
||||
listensForHTTP bool
|
||||
}
|
||||
|
||||
func printStartupBanner(state startupState) {
|
||||
log.Print(formatStartupBanner(state))
|
||||
}
|
||||
|
||||
func formatStartupBanner(state startupState) string {
|
||||
lines := []string{
|
||||
"",
|
||||
"__ __ _ _ ",
|
||||
"\\ \\ / /_ ___ _____ | | ___| |_ ",
|
||||
" \\ \\ /\\ / / _` \\ \\ / / _ \\ | |/ _ \\ __|",
|
||||
" \\ V V / (_| |\\ V / __/ | | __/ |_ ",
|
||||
" \\_/\\_/ \\__,_| \\_/ \\___|_|\\___|\\__|",
|
||||
fmt.Sprintf(" Wavelet %s", buildinfo.Version),
|
||||
"",
|
||||
fmt.Sprintf(" Environment: %s", config.Config.App.Env),
|
||||
fmt.Sprintf(" Runtime: %s/%s (%s)", runtime.GOOS, runtime.GOARCH, runtime.Version()),
|
||||
fmt.Sprintf(" Build time: %s", buildTime()),
|
||||
}
|
||||
if state.listensForHTTP {
|
||||
lines = append(lines, fmt.Sprintf(" Listening: http://%s", config.Config.App.Addr))
|
||||
}
|
||||
lines = append(lines, fmt.Sprintf(" Mode: %s", state.mode), "")
|
||||
return strings.Join(lines, "\n")
|
||||
}
|
||||
|
||||
func buildTime() string {
|
||||
if buildinfo.BuildTime == "" {
|
||||
return "development build"
|
||||
}
|
||||
return buildinfo.BuildTime
|
||||
}
|
||||
@@ -0,0 +1,47 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package cmd
|
||||
|
||||
import (
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/backend/pkg/buildinfo"
|
||||
"github.com/Rain-kl/Wavelet/backend/pkg/config"
|
||||
)
|
||||
|
||||
func TestFormatStartupBanner(t *testing.T) {
|
||||
previousVersion := buildinfo.Version
|
||||
previousBuildTime := buildinfo.BuildTime
|
||||
previousEnv := config.Config.App.Env
|
||||
previousAddr := config.Config.App.Addr
|
||||
t.Cleanup(func() {
|
||||
buildinfo.Version = previousVersion
|
||||
buildinfo.BuildTime = previousBuildTime
|
||||
config.Config.App.Env = previousEnv
|
||||
config.Config.App.Addr = previousAddr
|
||||
})
|
||||
|
||||
buildinfo.Version = "v3.2.1"
|
||||
buildinfo.BuildTime = "2026-07-13T08:00:00Z"
|
||||
config.Config.App.Env = "production"
|
||||
config.Config.App.Addr = ":3000"
|
||||
|
||||
banner := formatStartupBanner(startupState{
|
||||
mode: "API",
|
||||
listensForHTTP: true,
|
||||
})
|
||||
|
||||
for _, want := range []string{
|
||||
"Wavelet v3.2.1",
|
||||
"Environment: production",
|
||||
"Build time: 2026-07-13T08:00:00Z",
|
||||
"Listening: http://:3000",
|
||||
"Mode: API",
|
||||
} {
|
||||
if !strings.Contains(banner, want) {
|
||||
t.Errorf("banner missing %q:\n%s", want, banner)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,123 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package cmd
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"context"
|
||||
"crypto/rand"
|
||||
"errors"
|
||||
"fmt"
|
||||
"log"
|
||||
"os"
|
||||
"strings"
|
||||
|
||||
userdomain "github.com/Rain-kl/Wavelet/backend/plugins/domain/user"
|
||||
"github.com/Rain-kl/Wavelet/backend/plugins/infra/database"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/backend/plugins/domain/auth"
|
||||
"github.com/spf13/cobra"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
var (
|
||||
usernameFlag string
|
||||
passwordFlag string
|
||||
)
|
||||
|
||||
const (
|
||||
passwdCharset = "abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ0123456789!@#$%^&*"
|
||||
defaultPasswordLength = 16
|
||||
)
|
||||
|
||||
func generateRandomPassword(length int) (string, error) {
|
||||
bytes := make([]byte, length)
|
||||
if _, err := rand.Read(bytes); err != nil {
|
||||
return "", err
|
||||
}
|
||||
for i, b := range bytes {
|
||||
bytes[i] = passwdCharset[int(b)%len(passwdCharset)]
|
||||
}
|
||||
return string(bytes), nil
|
||||
}
|
||||
|
||||
var resetPasswdCmd = &cobra.Command{
|
||||
Use: "reset-passwd",
|
||||
Short: "重置指定账号密码",
|
||||
Run: func(_ *cobra.Command, _ []string) {
|
||||
ctx := context.Background()
|
||||
|
||||
// Ensure database is initialized
|
||||
database.DB(ctx)
|
||||
|
||||
var username string
|
||||
if usernameFlag != "" {
|
||||
username = usernameFlag
|
||||
} else {
|
||||
fmt.Print("请输入用户名: ")
|
||||
reader := bufio.NewReader(os.Stdin)
|
||||
input, err := reader.ReadString('\n')
|
||||
if err != nil {
|
||||
log.Fatalf("读取用户名失败: %v\n", err)
|
||||
}
|
||||
username = strings.TrimSpace(input)
|
||||
if username == "" {
|
||||
log.Fatal("用户名不能为空\n")
|
||||
}
|
||||
}
|
||||
|
||||
user, err := userdomain.GetUserByUsername(ctx, username)
|
||||
if err != nil {
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
log.Fatalf("错误: 用户 '%s' 不存在\n", username)
|
||||
}
|
||||
log.Fatalf("查询用户失败: %v\n", err)
|
||||
}
|
||||
|
||||
var password string
|
||||
if passwordFlag != "" {
|
||||
password = passwordFlag
|
||||
} else {
|
||||
password, err = generateRandomPassword(defaultPasswordLength)
|
||||
if err != nil {
|
||||
log.Fatalf("生成随机密码失败: %v\n", err)
|
||||
}
|
||||
}
|
||||
|
||||
if err := user.SetEncryptedPassword(password); err != nil {
|
||||
log.Fatalf("加密密码失败: %v\n", err)
|
||||
}
|
||||
|
||||
err = database.DB(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
if err := tx.Model(&user).Update("password", user.Password).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// Invalidate existing tokens
|
||||
var tokens []userdomain.AccessToken
|
||||
if err := tx.Where("user_id = ?", user.ID).Find(&tokens).Error; err == nil {
|
||||
for _, token := range tokens {
|
||||
auth.InvalidateCachedToken(ctx, token.TokenHash)
|
||||
}
|
||||
}
|
||||
|
||||
return tx.Where("user_id = ?", user.ID).Delete(&userdomain.AccessToken{}).Error
|
||||
})
|
||||
if err != nil {
|
||||
log.Fatalf("重置密码失败: %v\n", err)
|
||||
}
|
||||
|
||||
auth.InvalidateCachedUser(ctx, user.ID)
|
||||
|
||||
fmt.Println("成功重置密码!")
|
||||
fmt.Printf("用户名: %s\n", user.Username)
|
||||
fmt.Printf("新密码: %s\n", password)
|
||||
},
|
||||
}
|
||||
|
||||
func init() {
|
||||
resetPasswdCmd.Flags().StringVar(&usernameFlag, "user", "", "重置密码的目标用户名")
|
||||
resetPasswdCmd.Flags().StringVar(&passwordFlag, "password", "", "新密码(若不指定,则自动生成随机密码)")
|
||||
rootCmd.AddCommand(resetPasswdCmd)
|
||||
}
|
||||
@@ -0,0 +1,227 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package cmd
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"io"
|
||||
"os"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/backend/pkg/testhelper"
|
||||
userdomain "github.com/Rain-kl/Wavelet/backend/plugins/domain/user"
|
||||
)
|
||||
|
||||
func TestResetPasswdCmd_WithUserAndPassword(t *testing.T) {
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
defer cleanup()
|
||||
|
||||
// Seed test user
|
||||
user := userdomain.User{
|
||||
ID: 1001,
|
||||
Username: "testuser1",
|
||||
Nickname: "Test User 1",
|
||||
Email: "test1@example.com",
|
||||
IsActive: true,
|
||||
LastLoginAt: time.Now(),
|
||||
}
|
||||
_ = user.SetEncryptedPassword("oldpassword")
|
||||
if err := dbConn.Create(&user).Error; err != nil {
|
||||
t.Fatalf("failed to create test user: %v", err)
|
||||
}
|
||||
|
||||
// Create access token to test invalidation/deletion
|
||||
token := userdomain.AccessToken{
|
||||
ID: 1,
|
||||
UserID: user.ID,
|
||||
Name: "testtoken",
|
||||
TokenHash: "somehash",
|
||||
MaskedToken: "some...",
|
||||
}
|
||||
if err := dbConn.Create(&token).Error; err != nil {
|
||||
t.Fatalf("failed to create access token: %v", err)
|
||||
}
|
||||
|
||||
// Override PreRun to bypass goose migrations in unit tests
|
||||
oldPreRun := resetPasswdCmd.PreRun
|
||||
resetPasswdCmd.PreRun = nil
|
||||
defer func() { resetPasswdCmd.PreRun = oldPreRun }()
|
||||
|
||||
// Execute command with args
|
||||
rootCmd.SetArgs([]string{"reset-passwd", "--user", "testuser1", "--password", "newpassword123"})
|
||||
|
||||
// Capture output
|
||||
oldStdout := os.Stdout
|
||||
r, w, _ := os.Pipe()
|
||||
os.Stdout = w
|
||||
|
||||
err := rootCmd.Execute()
|
||||
w.Close()
|
||||
os.Stdout = oldStdout
|
||||
|
||||
if err != nil {
|
||||
t.Fatalf("command execute failed: %v", err)
|
||||
}
|
||||
|
||||
var buf bytes.Buffer
|
||||
_, _ = io.Copy(&buf, r)
|
||||
output := buf.String()
|
||||
|
||||
if !bytes.Contains([]byte(output), []byte("成功重置密码!")) {
|
||||
t.Errorf("expected output to contain success message, got: %s", output)
|
||||
}
|
||||
|
||||
// Verify password in DB
|
||||
var dbUser userdomain.User
|
||||
if err := dbConn.Where("id = ?", user.ID).First(&dbUser).Error; err != nil {
|
||||
t.Fatalf("failed to query user from DB: %v", err)
|
||||
}
|
||||
if !dbUser.CheckPassword("newpassword123") {
|
||||
t.Errorf("password was not updated correctly in DB")
|
||||
}
|
||||
|
||||
// Verify token deleted
|
||||
var count int64
|
||||
dbConn.Model(&userdomain.AccessToken{}).Where("user_id = ?", user.ID).Count(&count)
|
||||
if count != 0 {
|
||||
t.Errorf("expected access tokens to be deleted, got %d", count)
|
||||
}
|
||||
}
|
||||
|
||||
func TestResetPasswdCmd_WithUserAndRandomPassword(t *testing.T) {
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
defer cleanup()
|
||||
|
||||
// Seed test user
|
||||
user := userdomain.User{
|
||||
ID: 1002,
|
||||
Username: "testuser2",
|
||||
Nickname: "Test User 2",
|
||||
Email: "test2@example.com",
|
||||
IsActive: true,
|
||||
LastLoginAt: time.Now(),
|
||||
}
|
||||
_ = user.SetEncryptedPassword("oldpassword")
|
||||
if err := dbConn.Create(&user).Error; err != nil {
|
||||
t.Fatalf("failed to create test user: %v", err)
|
||||
}
|
||||
|
||||
// Override PreRun
|
||||
oldPreRun := resetPasswdCmd.PreRun
|
||||
resetPasswdCmd.PreRun = nil
|
||||
defer func() { resetPasswdCmd.PreRun = oldPreRun }()
|
||||
|
||||
// Reset flags
|
||||
usernameFlag = ""
|
||||
passwordFlag = ""
|
||||
|
||||
// Execute command with args (no --password)
|
||||
rootCmd.SetArgs([]string{"reset-passwd", "--user", "testuser2"})
|
||||
|
||||
// Capture output
|
||||
oldStdout := os.Stdout
|
||||
r, w, _ := os.Pipe()
|
||||
os.Stdout = w
|
||||
|
||||
err := rootCmd.Execute()
|
||||
w.Close()
|
||||
os.Stdout = oldStdout
|
||||
|
||||
if err != nil {
|
||||
t.Fatalf("command execute failed: %v", err)
|
||||
}
|
||||
|
||||
var buf bytes.Buffer
|
||||
_, _ = io.Copy(&buf, r)
|
||||
output := buf.String()
|
||||
|
||||
if !bytes.Contains([]byte(output), []byte("成功重置密码!")) {
|
||||
t.Errorf("expected output to contain success message, got: %s", output)
|
||||
}
|
||||
|
||||
// Verify password in DB (should be updated and not equal to old one)
|
||||
var dbUser userdomain.User
|
||||
if err := dbConn.Where("id = ?", user.ID).First(&dbUser).Error; err != nil {
|
||||
t.Fatalf("failed to query user from DB: %v", err)
|
||||
}
|
||||
if dbUser.CheckPassword("oldpassword") {
|
||||
t.Errorf("expected password to change, but it matches the old one")
|
||||
}
|
||||
}
|
||||
|
||||
func TestResetPasswdCmd_InteractiveMode(t *testing.T) {
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
defer cleanup()
|
||||
|
||||
// Seed test user
|
||||
user := userdomain.User{
|
||||
ID: 1003,
|
||||
Username: "testuser3",
|
||||
Nickname: "Test User 3",
|
||||
Email: "test3@example.com",
|
||||
IsActive: true,
|
||||
LastLoginAt: time.Now(),
|
||||
}
|
||||
_ = user.SetEncryptedPassword("oldpassword")
|
||||
if err := dbConn.Create(&user).Error; err != nil {
|
||||
t.Fatalf("failed to create test user: %v", err)
|
||||
}
|
||||
|
||||
// Override PreRun
|
||||
oldPreRun := resetPasswdCmd.PreRun
|
||||
resetPasswdCmd.PreRun = nil
|
||||
defer func() { resetPasswdCmd.PreRun = oldPreRun }()
|
||||
|
||||
// Mock stdin
|
||||
inR, inW, err := os.Pipe()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer inR.Close()
|
||||
defer inW.Close()
|
||||
|
||||
oldStdin := os.Stdin
|
||||
os.Stdin = inR
|
||||
defer func() { os.Stdin = oldStdin }()
|
||||
|
||||
// Write username to stdin
|
||||
_, _ = inW.WriteString("testuser3\n")
|
||||
inW.Close()
|
||||
|
||||
// Reset flags
|
||||
usernameFlag = ""
|
||||
passwordFlag = ""
|
||||
rootCmd.SetArgs([]string{"reset-passwd"})
|
||||
|
||||
// Capture output
|
||||
oldStdout := os.Stdout
|
||||
r, w, _ := os.Pipe()
|
||||
os.Stdout = w
|
||||
|
||||
err = rootCmd.Execute()
|
||||
w.Close()
|
||||
os.Stdout = oldStdout
|
||||
|
||||
if err != nil {
|
||||
t.Fatalf("command execute failed: %v", err)
|
||||
}
|
||||
|
||||
var buf bytes.Buffer
|
||||
_, _ = io.Copy(&buf, r)
|
||||
output := buf.String()
|
||||
|
||||
if !bytes.Contains([]byte(output), []byte("成功重置密码!")) {
|
||||
t.Errorf("expected output to contain success message, got: %s", output)
|
||||
}
|
||||
|
||||
// Verify user password changed in DB
|
||||
var dbUser userdomain.User
|
||||
if err := dbConn.Where("id = ?", user.ID).First(&dbUser).Error; err != nil {
|
||||
t.Fatalf("failed to query user from DB: %v", err)
|
||||
}
|
||||
if dbUser.CheckPassword("oldpassword") {
|
||||
t.Errorf("expected password to change, but it matches the old one")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,68 @@
|
||||
// Copyright 2025 linux.do
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package cmd
|
||||
|
||||
import (
|
||||
"context"
|
||||
"log"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/backend/pkg/buildinfo"
|
||||
"github.com/Rain-kl/Wavelet/backend/pkg/config"
|
||||
"github.com/Rain-kl/Wavelet/backend/pkg/logger"
|
||||
"github.com/Rain-kl/Wavelet/backend/pkg/trace"
|
||||
"github.com/spf13/cobra"
|
||||
)
|
||||
|
||||
const traceShutdownTimeout = 10 * time.Second
|
||||
|
||||
var rootCmd = &cobra.Command{
|
||||
Use: "wavelet",
|
||||
PersistentPreRun: func(_ *cobra.Command, _ []string) {
|
||||
logger.Init(logger.Config{
|
||||
Level: config.Config.Log.Level,
|
||||
Format: config.Config.Log.Format,
|
||||
Output: config.Config.Log.Output,
|
||||
FilePath: config.Config.Log.FilePath,
|
||||
MaxSize: config.Config.Log.MaxSize,
|
||||
MaxAge: config.Config.Log.MaxAge,
|
||||
MaxBackups: config.Config.Log.MaxBackups,
|
||||
Compress: config.Config.Log.Compress,
|
||||
})
|
||||
trace.Init(trace.Config{
|
||||
AppName: config.Config.App.AppName,
|
||||
SamplingRate: config.Config.Otel.SamplingRate,
|
||||
TracerName: config.Config.Otel.TracerName,
|
||||
})
|
||||
},
|
||||
PersistentPostRun: func(_ *cobra.Command, _ []string) {
|
||||
shutdownTraceProvider()
|
||||
},
|
||||
Run: func(_ *cobra.Command, args []string) {
|
||||
// 无参数时默认以融合模式启动所有服务
|
||||
allCmd.Run(allCmd, args)
|
||||
},
|
||||
}
|
||||
|
||||
func shutdownTraceProvider() {
|
||||
ctx, cancel := context.WithTimeout(context.Background(), traceShutdownTimeout)
|
||||
defer cancel()
|
||||
trace.Shutdown(ctx)
|
||||
}
|
||||
|
||||
func init() {
|
||||
rootCmd.Version = buildinfo.Version
|
||||
rootCmd.CompletionOptions.DisableDefaultCmd = true
|
||||
|
||||
// 集中将子命令注册到根命令,以解决 Cobra 的 unknown command 校验限制
|
||||
rootCmd.AddCommand(allCmd, apiCmd, workerCmd, schedulerCmd)
|
||||
}
|
||||
|
||||
// Execute 执行根命令
|
||||
func Execute() {
|
||||
if err := rootCmd.Execute(); err != nil {
|
||||
log.Fatalf("[CMD] execute failed; %s\n", err)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,23 @@
|
||||
// Copyright 2025 linux.do
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package cmd
|
||||
|
||||
import (
|
||||
"log"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/backend/core"
|
||||
"github.com/spf13/cobra"
|
||||
)
|
||||
|
||||
var schedulerCmd = &cobra.Command{
|
||||
Use: "scheduler",
|
||||
Short: "wavelet Scheduler",
|
||||
Run: func(_ *cobra.Command, _ []string) {
|
||||
app := newWaveletApp(core.ProfileSchedule)
|
||||
if err := app.Run(); err != nil {
|
||||
log.Fatalf("[Scheduler] run failed: %v\n", err)
|
||||
}
|
||||
},
|
||||
}
|
||||
@@ -0,0 +1,23 @@
|
||||
// Copyright 2025 linux.do
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package cmd
|
||||
|
||||
import (
|
||||
"log"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/backend/core"
|
||||
"github.com/spf13/cobra"
|
||||
)
|
||||
|
||||
var workerCmd = &cobra.Command{
|
||||
Use: "worker",
|
||||
Short: "wavelet Worker",
|
||||
Run: func(_ *cobra.Command, _ []string) {
|
||||
app := newWaveletApp(core.ProfileWorker)
|
||||
if err := app.Run(); err != nil {
|
||||
log.Fatalf("[Worker] run failed: %v\n", err)
|
||||
}
|
||||
},
|
||||
}
|
||||
@@ -0,0 +1,492 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package core
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"os"
|
||||
"os/signal"
|
||||
"strings"
|
||||
"sync"
|
||||
"syscall"
|
||||
"time"
|
||||
)
|
||||
|
||||
const (
|
||||
defaultShutdownTimeout = 10 * time.Second
|
||||
)
|
||||
|
||||
// AppOption configures an App instance during construction.
|
||||
type AppOption func(*App)
|
||||
|
||||
// WithContext sets a custom root Context for the App.
|
||||
func WithContext(ctx *Context) AppOption {
|
||||
return func(a *App) {
|
||||
if ctx != nil {
|
||||
a.ctx = ctx
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// WithProfile sets the runtime profile for the App.
|
||||
func WithProfile(profile Profile) AppOption {
|
||||
return func(a *App) {
|
||||
a.profile = normalizeProfile(profile)
|
||||
}
|
||||
}
|
||||
|
||||
// WithPlugins registers initial plugins for the App.
|
||||
func WithPlugins(plugins ...Plugin) AppOption {
|
||||
return func(a *App) {
|
||||
a.Use(plugins...)
|
||||
}
|
||||
}
|
||||
|
||||
// WithMigrationEngine sets the database migration engine for the App.
|
||||
func WithMigrationEngine(engine MigrationEngine) AppOption {
|
||||
return func(a *App) {
|
||||
a.migrationEngine = engine
|
||||
}
|
||||
}
|
||||
|
||||
// WithMigrationRunner sets the migration runner function for the App.
|
||||
func WithMigrationRunner(runner MigrationRunner) AppOption {
|
||||
return func(a *App) {
|
||||
a.migrationEngine = runner
|
||||
}
|
||||
}
|
||||
|
||||
// WithShutdownTimeout sets the fallback timeout for graceful application shutdown.
|
||||
func WithShutdownTimeout(timeout time.Duration) AppOption {
|
||||
return func(a *App) {
|
||||
if timeout > 0 {
|
||||
a.shutdownTimeout = timeout
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// App is the unified assembly entrypoint and runtime aspect dispatcher of the Cordis micro-kernel.
|
||||
// It manages plugin collection, dependency mounting, migration execution, profile-based driver startup,
|
||||
// and graceful signal-driven LIFO shutdown.
|
||||
type App struct {
|
||||
mu sync.RWMutex
|
||||
ctx *Context
|
||||
profile Profile
|
||||
plugins []Plugin
|
||||
pluginMap map[string]Plugin
|
||||
applied bool
|
||||
running bool
|
||||
startedDrivers []Driver
|
||||
migrationEngine MigrationEngine
|
||||
shutdownTimeout time.Duration
|
||||
}
|
||||
|
||||
// NewApp creates a new Cordis application instance with default options.
|
||||
func NewApp(opts ...AppOption) *App {
|
||||
app := &App{
|
||||
ctx: NewContext(context.Background()),
|
||||
profile: ProfileAll,
|
||||
pluginMap: make(map[string]Plugin),
|
||||
shutdownTimeout: defaultShutdownTimeout,
|
||||
}
|
||||
|
||||
for _, opt := range opts {
|
||||
if opt != nil {
|
||||
opt(app)
|
||||
}
|
||||
}
|
||||
|
||||
return app
|
||||
}
|
||||
|
||||
// Context returns the root micro-kernel Context of the application.
|
||||
func (a *App) Context() *Context {
|
||||
return a.ctx
|
||||
}
|
||||
|
||||
// Profile returns the current runtime profile of the application.
|
||||
func (a *App) Profile() Profile {
|
||||
a.mu.RLock()
|
||||
defer a.mu.RUnlock()
|
||||
return a.profile
|
||||
}
|
||||
|
||||
// WithProfile sets the application runtime profile and returns the App for fluent chaining.
|
||||
func (a *App) WithProfile(profile Profile) *App {
|
||||
a.mu.Lock()
|
||||
defer a.mu.Unlock()
|
||||
a.profile = normalizeProfile(profile)
|
||||
return a
|
||||
}
|
||||
|
||||
// SetProfile sets the application runtime profile.
|
||||
func (a *App) SetProfile(profile Profile) *App {
|
||||
return a.WithProfile(profile)
|
||||
}
|
||||
|
||||
// Use registers one or more plugins into the application in registration order.
|
||||
// Duplicate plugins (by Name) update existing registrations in-place to preserve order.
|
||||
func (a *App) Use(plugins ...Plugin) *App {
|
||||
a.mu.Lock()
|
||||
defer a.mu.Unlock()
|
||||
|
||||
for _, p := range plugins {
|
||||
if p == nil {
|
||||
continue
|
||||
}
|
||||
name := p.Name()
|
||||
if name == "" {
|
||||
continue
|
||||
}
|
||||
|
||||
if _, exists := a.pluginMap[name]; exists {
|
||||
for i, existing := range a.plugins {
|
||||
if existing.Name() == name {
|
||||
a.plugins[i] = p
|
||||
break
|
||||
}
|
||||
}
|
||||
} else {
|
||||
a.plugins = append(a.plugins, p)
|
||||
}
|
||||
a.pluginMap[name] = p
|
||||
}
|
||||
|
||||
return a
|
||||
}
|
||||
|
||||
// Plugins returns a copy of all registered plugins in registration order.
|
||||
func (a *App) Plugins() []Plugin {
|
||||
a.mu.RLock()
|
||||
defer a.mu.RUnlock()
|
||||
|
||||
res := make([]Plugin, len(a.plugins))
|
||||
copy(res, a.plugins)
|
||||
return res
|
||||
}
|
||||
|
||||
// Plugin retrieves a registered plugin by its unique name.
|
||||
func (a *App) Plugin(name string) (Plugin, bool) {
|
||||
a.mu.RLock()
|
||||
defer a.mu.RUnlock()
|
||||
|
||||
p, ok := a.pluginMap[name]
|
||||
return p, ok
|
||||
}
|
||||
|
||||
// SetMigrationEngine sets the migration engine for the application.
|
||||
func (a *App) SetMigrationEngine(engine MigrationEngine) *App {
|
||||
a.mu.Lock()
|
||||
defer a.mu.Unlock()
|
||||
a.migrationEngine = engine
|
||||
return a
|
||||
}
|
||||
|
||||
// SetMigrationRunner sets the migration runner function for the application.
|
||||
func (a *App) SetMigrationRunner(runner MigrationRunner) *App {
|
||||
return a.SetMigrationEngine(runner)
|
||||
}
|
||||
|
||||
// ApplyPlugins applies all registered plugins on the application Context.
|
||||
// It is idempotent and only applies plugins once per App instance.
|
||||
func (a *App) ApplyPlugins() error {
|
||||
a.mu.Lock()
|
||||
if a.applied {
|
||||
a.mu.Unlock()
|
||||
return nil
|
||||
}
|
||||
a.applied = true
|
||||
plugins := make([]Plugin, len(a.plugins))
|
||||
copy(plugins, a.plugins)
|
||||
a.mu.Unlock()
|
||||
|
||||
for _, p := range plugins {
|
||||
if err := p.Apply(a.ctx); err != nil {
|
||||
return fmt.Errorf("core: apply plugin %q failed: %w", p.Name(), err)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// RunMigrations dispatches migration execution across all registered plugin migration entries.
|
||||
func (a *App) RunMigrations() error {
|
||||
entries := a.ctx.Migrations().Entries()
|
||||
if len(entries) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
a.mu.RLock()
|
||||
engine := a.migrationEngine
|
||||
a.mu.RUnlock()
|
||||
|
||||
if engine == nil {
|
||||
// Attempt to resolve from IoC container
|
||||
if resolved, err := Inject[MigrationEngine](a.ctx); err == nil && resolved != nil {
|
||||
engine = resolved
|
||||
}
|
||||
}
|
||||
|
||||
if engine == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
if err := engine.Migrate(a.ctx, entries); err != nil {
|
||||
return fmt.Errorf("core: migration failed: %w", err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// Start executes the application boot pipeline:
|
||||
// 1. Applies all registered plugins to populate services, routes, tasks, and drivers.
|
||||
// 2. Dispatches database migrations via MigrationEngine.
|
||||
// 3. Filters and starts drivers matching the active Profile.
|
||||
// 4. Emits "app:ready" on the EventBus.
|
||||
//
|
||||
//nolint:contextcheck
|
||||
func (a *App) Start(ctx ...context.Context) error {
|
||||
a.mu.Lock()
|
||||
if a.running {
|
||||
a.mu.Unlock()
|
||||
return ErrAppRunning
|
||||
}
|
||||
a.running = true
|
||||
a.mu.Unlock()
|
||||
|
||||
var baseCtx context.Context
|
||||
switch {
|
||||
case len(ctx) > 0 && ctx[0] != nil:
|
||||
baseCtx = ctx[0]
|
||||
case a.ctx != nil:
|
||||
baseCtx = a.ctx.GoContext()
|
||||
default:
|
||||
baseCtx = context.Background()
|
||||
}
|
||||
|
||||
// 1. Apply plugins
|
||||
if err := a.ApplyPlugins(); err != nil {
|
||||
a.mu.Lock()
|
||||
a.running = false
|
||||
a.mu.Unlock()
|
||||
return err
|
||||
}
|
||||
|
||||
// 2. Run migrations
|
||||
if err := a.RunMigrations(); err != nil {
|
||||
a.mu.Lock()
|
||||
a.running = false
|
||||
a.mu.Unlock()
|
||||
return err
|
||||
}
|
||||
|
||||
// 3. Filter drivers matching active profile
|
||||
a.mu.RLock()
|
||||
prof := a.profile
|
||||
a.mu.RUnlock()
|
||||
|
||||
allDrivers := a.ctx.Drivers()
|
||||
var driversToStart []Driver
|
||||
for _, d := range allDrivers {
|
||||
if matchesProfile(prof, d.Type()) {
|
||||
driversToStart = append(driversToStart, d)
|
||||
}
|
||||
}
|
||||
|
||||
// 4. Start matching drivers
|
||||
for _, d := range driversToStart {
|
||||
if err := d.Start(baseCtx); err != nil {
|
||||
// Rollback already started drivers in reverse order
|
||||
a.mu.Lock()
|
||||
started := a.startedDrivers
|
||||
a.startedDrivers = nil
|
||||
a.running = false
|
||||
a.mu.Unlock()
|
||||
|
||||
for i := len(started) - 1; i >= 0; i-- {
|
||||
_ = started[i].Stop(context.Background())
|
||||
}
|
||||
|
||||
return fmt.Errorf("core: start driver %s failed: %w", d.Type(), err)
|
||||
}
|
||||
|
||||
a.mu.Lock()
|
||||
a.startedDrivers = append(a.startedDrivers, d)
|
||||
a.mu.Unlock()
|
||||
}
|
||||
|
||||
// 5. Emit app:ready event
|
||||
_ = a.ctx.Events().Emit(baseCtx, "app:ready", a)
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// Stop gracefully shuts down the application:
|
||||
// 1. Emits "app:stopping" on the EventBus.
|
||||
// 2. Stops all started drivers in LIFO (reverse) order.
|
||||
// 3. Disposes the Context (running registered OnDispose callbacks in LIFO order).
|
||||
// 4. Emits "app:stopped" on the EventBus.
|
||||
//
|
||||
//nolint:contextcheck
|
||||
func (a *App) Stop(ctx ...context.Context) error {
|
||||
a.mu.Lock()
|
||||
if !a.running {
|
||||
a.mu.Unlock()
|
||||
return nil
|
||||
}
|
||||
a.running = false
|
||||
started := a.startedDrivers
|
||||
a.startedDrivers = nil
|
||||
timeout := a.shutdownTimeout
|
||||
a.mu.Unlock()
|
||||
|
||||
var shutdownCtx context.Context
|
||||
if len(ctx) > 0 && ctx[0] != nil {
|
||||
shutdownCtx = ctx[0]
|
||||
} else {
|
||||
var cancel context.CancelFunc
|
||||
shutdownCtx, cancel = context.WithTimeout(context.Background(), timeout)
|
||||
defer cancel()
|
||||
}
|
||||
|
||||
_ = a.ctx.Events().Emit(shutdownCtx, "app:stopping", a)
|
||||
|
||||
var errs []error
|
||||
|
||||
// 1. Stop drivers in reverse order
|
||||
for i := len(started) - 1; i >= 0; i-- {
|
||||
d := started[i]
|
||||
if err := d.Stop(shutdownCtx); err != nil {
|
||||
errs = append(errs, fmt.Errorf("core: stop driver %s failed: %w", d.Type(), err))
|
||||
}
|
||||
}
|
||||
|
||||
// 2. Dispose context
|
||||
if a.ctx != nil && !a.ctx.IsDisposed() {
|
||||
if err := a.ctx.Dispose(); err != nil {
|
||||
errs = append(errs, fmt.Errorf("core: dispose context failed: %w", err))
|
||||
}
|
||||
}
|
||||
|
||||
_ = a.ctx.Events().Emit(shutdownCtx, "app:stopped", a)
|
||||
|
||||
return errors.Join(errs...)
|
||||
}
|
||||
|
||||
// Run starts the application and blocks until an OS signal (SIGINT, SIGTERM) or context cancellation is received,
|
||||
// then executes graceful shutdown.
|
||||
//
|
||||
//nolint:contextcheck
|
||||
func (a *App) Run(ctx ...context.Context) error {
|
||||
var parent context.Context
|
||||
switch {
|
||||
case len(ctx) > 0 && ctx[0] != nil:
|
||||
parent = ctx[0]
|
||||
case a.ctx != nil:
|
||||
parent = a.ctx.GoContext()
|
||||
default:
|
||||
parent = context.Background()
|
||||
}
|
||||
|
||||
sigCtx, stopSignals := signal.NotifyContext(parent, syscall.SIGINT, syscall.SIGTERM, os.Interrupt)
|
||||
defer stopSignals()
|
||||
|
||||
if err := a.Start(sigCtx); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// Wait for OS signal or context cancellation
|
||||
<-sigCtx.Done()
|
||||
|
||||
shutdownCtx, cancel := context.WithTimeout(context.Background(), a.shutdownTimeout)
|
||||
defer cancel()
|
||||
|
||||
return a.Stop(shutdownCtx)
|
||||
}
|
||||
|
||||
// IsRunning returns whether the application is currently running.
|
||||
func (a *App) IsRunning() bool {
|
||||
a.mu.RLock()
|
||||
defer a.mu.RUnlock()
|
||||
return a.running
|
||||
}
|
||||
|
||||
// StartedDrivers returns a copy of currently running drivers.
|
||||
func (a *App) StartedDrivers() []Driver {
|
||||
a.mu.RLock()
|
||||
defer a.mu.RUnlock()
|
||||
|
||||
res := make([]Driver, len(a.startedDrivers))
|
||||
copy(res, a.startedDrivers)
|
||||
return res
|
||||
}
|
||||
|
||||
// ExecuteCLI parses CLI arguments to configure the profile and runs the application.
|
||||
func (a *App) ExecuteCLI(args ...string) error {
|
||||
var ctx context.Context
|
||||
if a.ctx != nil {
|
||||
ctx = a.ctx.GoContext()
|
||||
} else {
|
||||
ctx = context.Background()
|
||||
}
|
||||
return a.ExecuteCLIWithContext(ctx, args...)
|
||||
}
|
||||
|
||||
// ExecuteCLIWithContext parses CLI arguments, configures the profile, and runs the application with the given context.
|
||||
//
|
||||
//nolint:contextcheck
|
||||
func (a *App) ExecuteCLIWithContext(ctx context.Context, args ...string) error {
|
||||
cliArgs := args
|
||||
if len(cliArgs) == 0 {
|
||||
cliArgs = os.Args[1:]
|
||||
}
|
||||
|
||||
profile := ProfileAll
|
||||
if len(cliArgs) > 0 {
|
||||
first := strings.TrimSpace(cliArgs[0])
|
||||
switch {
|
||||
case strings.HasPrefix(first, "--profile="):
|
||||
profile = Profile(strings.TrimPrefix(first, "--profile="))
|
||||
case strings.HasPrefix(first, "-p="):
|
||||
profile = Profile(strings.TrimPrefix(first, "-p="))
|
||||
case !strings.HasPrefix(first, "-"):
|
||||
profile = Profile(first)
|
||||
}
|
||||
}
|
||||
|
||||
a.WithProfile(profile)
|
||||
return a.Run(ctx)
|
||||
}
|
||||
|
||||
func matchesProfile(profile Profile, dt DriverType) bool {
|
||||
norm := normalizeProfile(profile)
|
||||
switch norm {
|
||||
case ProfileAll, "":
|
||||
return true
|
||||
case ProfileAPI:
|
||||
return dt == DriverTypeHTTP
|
||||
case ProfileWorker:
|
||||
return dt == DriverTypeWorker
|
||||
case ProfileSchedule:
|
||||
return dt == DriverTypeScheduler
|
||||
default:
|
||||
return string(norm) == string(dt)
|
||||
}
|
||||
}
|
||||
|
||||
func normalizeProfile(p Profile) Profile {
|
||||
switch strings.ToLower(strings.TrimSpace(string(p))) {
|
||||
case "api", "http":
|
||||
return ProfileAPI
|
||||
case "worker":
|
||||
return ProfileWorker
|
||||
case "schedule", "scheduler", "cron":
|
||||
return ProfileSchedule
|
||||
case "all", "fused", "full", "":
|
||||
return ProfileAll
|
||||
default:
|
||||
return p
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,507 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package core_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"sync"
|
||||
"testing"
|
||||
"testing/fstest"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/backend/core"
|
||||
"github.com/Rain-kl/Wavelet/backend/core/extpoints"
|
||||
)
|
||||
|
||||
// appMockDriver is a test driver tracking its start/stop lifecycle.
|
||||
type appMockDriver struct {
|
||||
mu sync.Mutex
|
||||
driverType core.DriverType
|
||||
startCalled bool
|
||||
stopCalled bool
|
||||
startErr error
|
||||
stopErr error
|
||||
}
|
||||
|
||||
func newAppMockDriver(dt core.DriverType) *appMockDriver {
|
||||
return &appMockDriver{driverType: dt}
|
||||
}
|
||||
|
||||
func (m *appMockDriver) Type() core.DriverType {
|
||||
return m.driverType
|
||||
}
|
||||
|
||||
func (m *appMockDriver) Start(_ context.Context) error {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
if m.startErr != nil {
|
||||
return m.startErr
|
||||
}
|
||||
m.startCalled = true
|
||||
return nil
|
||||
}
|
||||
|
||||
func (m *appMockDriver) Stop(_ context.Context) error {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
if m.stopErr != nil {
|
||||
return m.stopErr
|
||||
}
|
||||
m.stopCalled = true
|
||||
return nil
|
||||
}
|
||||
|
||||
func (m *appMockDriver) isStarted() bool {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
return m.startCalled
|
||||
}
|
||||
|
||||
func (m *appMockDriver) isStopped() bool {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
return m.stopCalled
|
||||
}
|
||||
|
||||
// appMockPlugin is a test plugin.
|
||||
type appMockPlugin struct {
|
||||
name string
|
||||
applyFn func(ctx *core.Context) error
|
||||
}
|
||||
|
||||
func (p *appMockPlugin) Name() string {
|
||||
return p.name
|
||||
}
|
||||
|
||||
func (p *appMockPlugin) Apply(ctx *core.Context) error {
|
||||
if p.applyFn != nil {
|
||||
return p.applyFn(ctx)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func TestAppNewAndConfiguration(t *testing.T) {
|
||||
customCtx := core.NewContext(context.Background())
|
||||
p1 := &appMockPlugin{name: "plugin1"}
|
||||
p2 := &appMockPlugin{name: "plugin2"}
|
||||
|
||||
app := core.NewApp(
|
||||
core.WithContext(customCtx),
|
||||
core.WithProfile(core.ProfileAPI),
|
||||
core.WithPlugins(p1, p2),
|
||||
core.WithShutdownTimeout(5*time.Second),
|
||||
)
|
||||
|
||||
assert.Equal(t, customCtx, app.Context())
|
||||
assert.Equal(t, core.ProfileAPI, app.Profile())
|
||||
assert.Len(t, app.Plugins(), 2)
|
||||
|
||||
retrieved, ok := app.Plugin("plugin1")
|
||||
assert.True(t, ok)
|
||||
assert.Equal(t, p1, retrieved)
|
||||
|
||||
_, ok = app.Plugin("non_existent")
|
||||
assert.False(t, ok)
|
||||
|
||||
// Update existing plugin in-place
|
||||
p1Updated := &appMockPlugin{name: "plugin1"}
|
||||
app.Use(p1Updated, nil)
|
||||
assert.Len(t, app.Plugins(), 2)
|
||||
retrieved, ok = app.Plugin("plugin1")
|
||||
assert.True(t, ok)
|
||||
assert.Equal(t, p1Updated, retrieved)
|
||||
|
||||
// Test SetProfile
|
||||
app.SetProfile(core.ProfileWorker)
|
||||
assert.Equal(t, core.ProfileWorker, app.Profile())
|
||||
}
|
||||
|
||||
func TestAppProfileDispatch(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
profile core.Profile
|
||||
expectedHTTP bool
|
||||
expectedWorker bool
|
||||
expectedCron bool
|
||||
expectedCustom bool
|
||||
}{
|
||||
{
|
||||
name: "ProfileAPI only starts HTTP driver",
|
||||
profile: core.ProfileAPI,
|
||||
expectedHTTP: true,
|
||||
expectedWorker: false,
|
||||
expectedCron: false,
|
||||
expectedCustom: false,
|
||||
},
|
||||
{
|
||||
name: "ProfileWorker only starts Worker driver",
|
||||
profile: core.ProfileWorker,
|
||||
expectedHTTP: false,
|
||||
expectedWorker: true,
|
||||
expectedCron: false,
|
||||
expectedCustom: false,
|
||||
},
|
||||
{
|
||||
name: "ProfileSchedule only starts Schedule driver",
|
||||
profile: core.ProfileSchedule,
|
||||
expectedHTTP: false,
|
||||
expectedWorker: false,
|
||||
expectedCron: true,
|
||||
expectedCustom: false,
|
||||
},
|
||||
{
|
||||
name: "Profile 'scheduler' alias starts Schedule driver",
|
||||
profile: core.Profile("scheduler"),
|
||||
expectedHTTP: false,
|
||||
expectedWorker: false,
|
||||
expectedCron: true,
|
||||
expectedCustom: false,
|
||||
},
|
||||
{
|
||||
name: "ProfileAll starts all drivers",
|
||||
profile: core.ProfileAll,
|
||||
expectedHTTP: true,
|
||||
expectedWorker: true,
|
||||
expectedCron: true,
|
||||
expectedCustom: true,
|
||||
},
|
||||
{
|
||||
name: "Custom profile starts custom driver",
|
||||
profile: core.Profile("custom_rpc"),
|
||||
expectedHTTP: false,
|
||||
expectedWorker: false,
|
||||
expectedCron: false,
|
||||
expectedCustom: true,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
httpD := newAppMockDriver(core.DriverTypeHTTP)
|
||||
workerD := newAppMockDriver(core.DriverTypeWorker)
|
||||
cronD := newAppMockDriver(core.DriverTypeScheduler)
|
||||
customD := newAppMockDriver(core.DriverType("custom_rpc"))
|
||||
|
||||
p := &appMockPlugin{
|
||||
name: "drivers_plugin",
|
||||
applyFn: func(ctx *core.Context) error {
|
||||
_ = ctx.RegisterDriver(httpD)
|
||||
_ = ctx.RegisterDriver(workerD)
|
||||
_ = ctx.RegisterDriver(cronD)
|
||||
_ = ctx.RegisterDriver(customD)
|
||||
return nil
|
||||
},
|
||||
}
|
||||
|
||||
app := core.NewApp(
|
||||
core.WithProfile(tt.profile),
|
||||
core.WithPlugins(p),
|
||||
)
|
||||
|
||||
err := app.Start(context.Background())
|
||||
require.NoError(t, err)
|
||||
|
||||
assert.Equal(t, tt.expectedHTTP, httpD.isStarted(), "HTTP driver start mismatch")
|
||||
assert.Equal(t, tt.expectedWorker, workerD.isStarted(), "Worker driver start mismatch")
|
||||
assert.Equal(t, tt.expectedCron, cronD.isStarted(), "Cron driver start mismatch")
|
||||
assert.Equal(t, tt.expectedCustom, customD.isStarted(), "Custom driver start mismatch")
|
||||
|
||||
err = app.Stop(context.Background())
|
||||
require.NoError(t, err)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestAppLifecycleStartStop(t *testing.T) {
|
||||
var stopOrder []string
|
||||
var stopOrderMu sync.Mutex
|
||||
|
||||
httpD := newAppMockDriver(core.DriverTypeHTTP)
|
||||
workerD := newAppMockDriver(core.DriverTypeWorker)
|
||||
|
||||
httpD.stopErr = nil
|
||||
workerD.stopErr = nil
|
||||
|
||||
// Wrap stop to record order
|
||||
origHttpStop := httpD.Stop
|
||||
_ = origHttpStop
|
||||
|
||||
p := &appMockPlugin{
|
||||
name: "test_plugin",
|
||||
applyFn: func(ctx *core.Context) error {
|
||||
_ = ctx.RegisterDriver(httpD)
|
||||
_ = ctx.RegisterDriver(workerD)
|
||||
|
||||
ctx.OnDispose(func() error {
|
||||
stopOrderMu.Lock()
|
||||
stopOrder = append(stopOrder, "ctx_disposer")
|
||||
stopOrderMu.Unlock()
|
||||
return nil
|
||||
})
|
||||
return nil
|
||||
},
|
||||
}
|
||||
|
||||
app := core.NewApp(
|
||||
core.WithProfile(core.ProfileAll),
|
||||
core.WithPlugins(p),
|
||||
)
|
||||
|
||||
var readyReceived, stoppingReceived, stoppedReceived bool
|
||||
app.Context().Events().On("app:ready", func() {
|
||||
readyReceived = true
|
||||
})
|
||||
app.Context().Events().On("app:stopping", func() {
|
||||
stoppingReceived = true
|
||||
})
|
||||
app.Context().Events().On("app:stopped", func() {
|
||||
stoppedReceived = true
|
||||
})
|
||||
|
||||
err := app.Start(context.Background())
|
||||
require.NoError(t, err)
|
||||
assert.True(t, app.IsRunning())
|
||||
assert.Len(t, app.StartedDrivers(), 2)
|
||||
assert.True(t, readyReceived)
|
||||
|
||||
err = app.Stop(context.Background())
|
||||
require.NoError(t, err)
|
||||
assert.False(t, app.IsRunning())
|
||||
assert.Empty(t, app.StartedDrivers())
|
||||
assert.True(t, stoppingReceived)
|
||||
assert.True(t, stoppedReceived)
|
||||
|
||||
assert.True(t, httpD.isStopped())
|
||||
assert.True(t, workerD.isStopped())
|
||||
assert.True(t, app.Context().IsDisposed())
|
||||
|
||||
stopOrderMu.Lock()
|
||||
assert.Contains(t, stopOrder, "ctx_disposer")
|
||||
stopOrderMu.Unlock()
|
||||
}
|
||||
|
||||
func TestAppStartDriverFailureRollback(t *testing.T) {
|
||||
driver1 := newAppMockDriver(core.DriverTypeHTTP)
|
||||
driver2 := newAppMockDriver(core.DriverTypeWorker)
|
||||
driver2.startErr = errors.New("worker listen port conflict")
|
||||
driver3 := newAppMockDriver(core.DriverTypeScheduler)
|
||||
|
||||
p := &appMockPlugin{
|
||||
name: "fail_driver_plugin",
|
||||
applyFn: func(ctx *core.Context) error {
|
||||
_ = ctx.RegisterDriver(driver1)
|
||||
_ = ctx.RegisterDriver(driver2)
|
||||
_ = ctx.RegisterDriver(driver3)
|
||||
return nil
|
||||
},
|
||||
}
|
||||
|
||||
app := core.NewApp(
|
||||
core.WithProfile(core.ProfileAll),
|
||||
core.WithPlugins(p),
|
||||
)
|
||||
|
||||
err := app.Start(context.Background())
|
||||
require.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "worker listen port conflict")
|
||||
assert.False(t, app.IsRunning())
|
||||
|
||||
// Driver 1 was started then rolled back (stopped)
|
||||
assert.True(t, driver1.isStarted())
|
||||
assert.True(t, driver1.isStopped())
|
||||
|
||||
// Driver 3 was never started
|
||||
assert.False(t, driver3.isStarted())
|
||||
}
|
||||
|
||||
func TestAppMigrationEngineExecution(t *testing.T) {
|
||||
var migratedEntries []extpoints.MigrationEntry
|
||||
runner := core.MigrationRunner(func(ctx *core.Context, entries []extpoints.MigrationEntry) error {
|
||||
migratedEntries = entries
|
||||
return nil
|
||||
})
|
||||
|
||||
sqlFS := fstest.MapFS{
|
||||
"migrations/001_init.sql": &fstest.MapFile{Data: []byte("CREATE TABLE users(id int);")},
|
||||
}
|
||||
|
||||
p := &appMockPlugin{
|
||||
name: "auth",
|
||||
applyFn: func(ctx *core.Context) error {
|
||||
ctx.Migrations().Register("auth", sqlFS)
|
||||
return nil
|
||||
},
|
||||
}
|
||||
|
||||
app := core.NewApp(
|
||||
core.WithProfile(core.ProfileAll),
|
||||
core.WithPlugins(p),
|
||||
core.WithMigrationRunner(runner),
|
||||
)
|
||||
|
||||
err := app.Start(context.Background())
|
||||
require.NoError(t, err)
|
||||
defer func() { _ = app.Stop(context.Background()) }()
|
||||
|
||||
require.Len(t, migratedEntries, 1)
|
||||
assert.Equal(t, "auth", migratedEntries[0].PluginID)
|
||||
}
|
||||
|
||||
func TestAppMigrationEngineFromIoCContainer(t *testing.T) {
|
||||
var executed bool
|
||||
runner := core.MigrationRunner(func(ctx *core.Context, entries []extpoints.MigrationEntry) error {
|
||||
executed = true
|
||||
return nil
|
||||
})
|
||||
|
||||
sqlFS := fstest.MapFS{
|
||||
"migrations/001_init.sql": &fstest.MapFile{Data: []byte("CREATE TABLE logs(id int);")},
|
||||
}
|
||||
|
||||
p := &appMockPlugin{
|
||||
name: "logstore",
|
||||
applyFn: func(ctx *core.Context) error {
|
||||
ctx.Migrations().Register("logstore", sqlFS)
|
||||
core.Provide[core.MigrationEngine](ctx, runner)
|
||||
return nil
|
||||
},
|
||||
}
|
||||
|
||||
app := core.NewApp(
|
||||
core.WithProfile(core.ProfileAll),
|
||||
core.WithPlugins(p),
|
||||
)
|
||||
|
||||
err := app.Start(context.Background())
|
||||
require.NoError(t, err)
|
||||
defer func() { _ = app.Stop(context.Background()) }()
|
||||
|
||||
assert.True(t, executed)
|
||||
}
|
||||
|
||||
func TestAppRunContextCancellation(t *testing.T) {
|
||||
d := newAppMockDriver(core.DriverTypeHTTP)
|
||||
p := &appMockPlugin{
|
||||
name: "http_plugin",
|
||||
applyFn: func(ctx *core.Context) error {
|
||||
return ctx.RegisterDriver(d)
|
||||
},
|
||||
}
|
||||
|
||||
app := core.NewApp(
|
||||
core.WithProfile(core.ProfileAPI),
|
||||
core.WithPlugins(p),
|
||||
core.WithShutdownTimeout(1*time.Second),
|
||||
)
|
||||
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
|
||||
errCh := make(chan error, 1)
|
||||
go func() {
|
||||
errCh <- app.Run(ctx)
|
||||
}()
|
||||
|
||||
// Wait briefly then cancel
|
||||
time.Sleep(50 * time.Millisecond)
|
||||
assert.True(t, app.IsRunning())
|
||||
assert.True(t, d.isStarted())
|
||||
|
||||
cancel()
|
||||
|
||||
select {
|
||||
case err := <-errCh:
|
||||
assert.NoError(t, err)
|
||||
assert.False(t, app.IsRunning())
|
||||
assert.True(t, d.isStopped())
|
||||
case <-time.After(3 * time.Second):
|
||||
t.Fatal("app.Run did not terminate upon context cancellation")
|
||||
}
|
||||
}
|
||||
|
||||
func TestAppExecuteCLI(t *testing.T) {
|
||||
// Test CLI argument parsing logic
|
||||
tests := []struct {
|
||||
args []string
|
||||
expectedProfile core.Profile
|
||||
}{
|
||||
{args: []string{"api"}, expectedProfile: core.ProfileAPI},
|
||||
{args: []string{"worker"}, expectedProfile: core.ProfileWorker},
|
||||
{args: []string{"scheduler"}, expectedProfile: core.ProfileSchedule},
|
||||
{args: []string{"schedule"}, expectedProfile: core.ProfileSchedule},
|
||||
{args: []string{"all"}, expectedProfile: core.ProfileAll},
|
||||
{args: []string{"--profile=worker"}, expectedProfile: core.ProfileWorker},
|
||||
{args: []string{"-p=api"}, expectedProfile: core.ProfileAPI},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.args[0], func(t *testing.T) {
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
cancel() // cancel immediately
|
||||
|
||||
// Use custom root context to control cancellation
|
||||
customApp := core.NewApp(core.WithContext(core.NewContext(ctx)))
|
||||
_ = customApp.ExecuteCLI(tt.args...)
|
||||
assert.Equal(t, tt.expectedProfile, customApp.Profile())
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestAppIdempotencyAndErrorStates(t *testing.T) {
|
||||
app := core.NewApp()
|
||||
|
||||
// Double start returns error
|
||||
err := app.Start(context.Background())
|
||||
require.NoError(t, err)
|
||||
|
||||
err = app.Start(context.Background())
|
||||
assert.ErrorIs(t, err, core.ErrAppRunning)
|
||||
|
||||
// Stop clears running state
|
||||
err = app.Stop(context.Background())
|
||||
require.NoError(t, err)
|
||||
|
||||
// Double stop succeeds
|
||||
err = app.Stop(context.Background())
|
||||
require.NoError(t, err)
|
||||
|
||||
// Plugin apply failure
|
||||
failPlugin := &appMockPlugin{
|
||||
name: "failing_plugin",
|
||||
applyFn: func(ctx *core.Context) error {
|
||||
return errors.New("plugin init boom")
|
||||
},
|
||||
}
|
||||
app2 := core.NewApp(core.WithPlugins(failPlugin))
|
||||
err = app2.Start(context.Background())
|
||||
require.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "plugin init boom")
|
||||
assert.False(t, app2.IsRunning())
|
||||
|
||||
// Migration failure
|
||||
migFailRunner := core.MigrationRunner(func(ctx *core.Context, entries []extpoints.MigrationEntry) error {
|
||||
return errors.New("sql migrate error")
|
||||
})
|
||||
sqlFS := fstest.MapFS{
|
||||
"migrations/001.sql": &fstest.MapFile{Data: []byte("...")},
|
||||
}
|
||||
migPlugin := &appMockPlugin{
|
||||
name: "db_plugin",
|
||||
applyFn: func(ctx *core.Context) error {
|
||||
ctx.Migrations().Register("db_plugin", sqlFS)
|
||||
return nil
|
||||
},
|
||||
}
|
||||
app3 := core.NewApp(
|
||||
core.WithPlugins(migPlugin),
|
||||
core.WithMigrationRunner(migFailRunner),
|
||||
)
|
||||
err = app3.Start(context.Background())
|
||||
require.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "sql migrate error")
|
||||
assert.False(t, app3.IsRunning())
|
||||
}
|
||||
@@ -0,0 +1,193 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package core provides the micro-kernel service bus, generic IoC container, and runtime extensions.
|
||||
package core
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"reflect"
|
||||
"sync"
|
||||
)
|
||||
|
||||
// Container manages service registration and resolution using Go reflection and generics.
|
||||
type Container struct {
|
||||
mu sync.RWMutex
|
||||
parent *Container
|
||||
services map[reflect.Type]any
|
||||
listeners map[reflect.Type][]func(any)
|
||||
}
|
||||
|
||||
// NewContainer creates a new IoC container instance with an optional parent container.
|
||||
func NewContainer(parent *Container) *Container {
|
||||
return &Container{
|
||||
parent: parent,
|
||||
services: make(map[reflect.Type]any),
|
||||
listeners: make(map[reflect.Type][]func(any)),
|
||||
}
|
||||
}
|
||||
|
||||
func isNil(i any) bool {
|
||||
if i == nil {
|
||||
return true
|
||||
}
|
||||
v := reflect.ValueOf(i)
|
||||
switch v.Kind() {
|
||||
case reflect.Chan, reflect.Func, reflect.Map, reflect.Pointer, reflect.UnsafePointer, reflect.Interface, reflect.Slice:
|
||||
return v.IsNil()
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
// Provide registers a typed service implementation into the Context's IoC container.
|
||||
func Provide[T any](ctx *Context, service T) {
|
||||
if ctx == nil {
|
||||
panic("core: nil context provided to Provide")
|
||||
}
|
||||
if isNil(service) {
|
||||
panic("core: cannot provide nil service")
|
||||
}
|
||||
|
||||
targetType := reflect.TypeFor[T]()
|
||||
ctx.Container().provide(targetType, service)
|
||||
}
|
||||
|
||||
func (c *Container) provide(targetType reflect.Type, service any) {
|
||||
c.mu.Lock()
|
||||
c.services[targetType] = service
|
||||
|
||||
// Collect any matching listeners to invoke outside the lock
|
||||
var callbacks []func(any)
|
||||
svcType := reflect.TypeOf(service)
|
||||
for lType, cbs := range c.listeners {
|
||||
if lType == targetType || (lType.Kind() == reflect.Interface && svcType.Implements(lType)) {
|
||||
callbacks = append(callbacks, cbs...)
|
||||
}
|
||||
}
|
||||
c.mu.Unlock()
|
||||
|
||||
for _, cb := range callbacks {
|
||||
cb(service)
|
||||
}
|
||||
}
|
||||
|
||||
// Inject resolves a registered service of type T from the Context.
|
||||
func Inject[T any](ctx *Context) (T, error) {
|
||||
var zero T
|
||||
if ctx == nil {
|
||||
return zero, ErrNilContext
|
||||
}
|
||||
|
||||
targetType := reflect.TypeFor[T]()
|
||||
val, err := ctx.Container().resolve(targetType)
|
||||
if err != nil {
|
||||
return zero, err
|
||||
}
|
||||
|
||||
typedVal, ok := val.(T)
|
||||
if !ok {
|
||||
return zero, fmt.Errorf("%w: cannot cast %T to %v", ErrServiceNotFound, val, targetType)
|
||||
}
|
||||
return typedVal, nil
|
||||
}
|
||||
|
||||
func (c *Container) resolve(targetType reflect.Type) (any, error) {
|
||||
c.mu.RLock()
|
||||
// 1. Direct type match
|
||||
if val, ok := c.services[targetType]; ok {
|
||||
c.mu.RUnlock()
|
||||
return val, nil
|
||||
}
|
||||
|
||||
// 2. Interface assignment scan
|
||||
if targetType.Kind() == reflect.Interface {
|
||||
for _, val := range c.services {
|
||||
if reflect.TypeOf(val).Implements(targetType) {
|
||||
c.mu.RUnlock()
|
||||
return val, nil
|
||||
}
|
||||
}
|
||||
}
|
||||
c.mu.RUnlock()
|
||||
|
||||
// 3. Fallback to parent container
|
||||
if c.parent != nil {
|
||||
return c.parent.resolve(targetType)
|
||||
}
|
||||
|
||||
return nil, fmt.Errorf("%w: %v", ErrServiceNotFound, targetType)
|
||||
}
|
||||
|
||||
// MustInject resolves a service of type T or panics if the service is not found.
|
||||
func MustInject[T any](ctx *Context) T {
|
||||
s, err := Inject[T](ctx)
|
||||
if err != nil {
|
||||
panic(fmt.Sprintf("core: failed to inject service %v: %v", reflect.TypeFor[T](), err))
|
||||
}
|
||||
return s
|
||||
}
|
||||
|
||||
// Has returns true if a service of type T is registered and resolvable in the Context.
|
||||
func Has[T any](ctx *Context) bool {
|
||||
_, err := Inject[T](ctx)
|
||||
return err == nil
|
||||
}
|
||||
|
||||
// Using executes the given function synchronously if the required dependency is ready.
|
||||
func Using[T1 any](ctx *Context, fn func(s1 T1)) error {
|
||||
s1, err := Inject[T1](ctx)
|
||||
if err != nil {
|
||||
return fmt.Errorf("%w: %w", ErrServiceNotReady, err)
|
||||
}
|
||||
fn(s1)
|
||||
return nil
|
||||
}
|
||||
|
||||
// Using2 executes the given function synchronously if both required dependencies are ready.
|
||||
func Using2[T1, T2 any](ctx *Context, fn func(s1 T1, s2 T2)) error {
|
||||
s1, err1 := Inject[T1](ctx)
|
||||
s2, err2 := Inject[T2](ctx)
|
||||
if err1 != nil || err2 != nil {
|
||||
return fmt.Errorf("%w: (dep1: %v, dep2: %v)", ErrServiceNotReady, err1, err2)
|
||||
}
|
||||
fn(s1, s2)
|
||||
return nil
|
||||
}
|
||||
|
||||
// Using3 executes the given function synchronously if all 3 required dependencies are ready.
|
||||
func Using3[T1, T2, T3 any](ctx *Context, fn func(s1 T1, s2 T2, s3 T3)) error {
|
||||
s1, err1 := Inject[T1](ctx)
|
||||
s2, err2 := Inject[T2](ctx)
|
||||
s3, err3 := Inject[T3](ctx)
|
||||
if err1 != nil || err2 != nil || err3 != nil {
|
||||
return fmt.Errorf("%w: (dep1: %v, dep2: %v, dep3: %v)", ErrServiceNotReady, err1, err2, err3)
|
||||
}
|
||||
fn(s1, s2, s3)
|
||||
return nil
|
||||
}
|
||||
|
||||
// When registers a reactive hook that is called immediately if T is already provided,
|
||||
// or called as soon as T is provided in the future.
|
||||
func When[T any](ctx *Context, fn func(s T)) {
|
||||
if ctx == nil {
|
||||
panic("core: nil context provided to When")
|
||||
}
|
||||
|
||||
targetType := reflect.TypeFor[T]()
|
||||
c := ctx.Container()
|
||||
|
||||
// If already ready, execute immediately
|
||||
if s, err := Inject[T](ctx); err == nil {
|
||||
fn(s)
|
||||
}
|
||||
|
||||
// Also register listener for future calls / updates
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
c.listeners[targetType] = append(c.listeners[targetType], func(val any) {
|
||||
if typed, ok := val.(T); ok {
|
||||
fn(typed)
|
||||
}
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,341 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package core
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/backend/core/extpoints"
|
||||
)
|
||||
|
||||
// Context is the central micro-kernel service bus and runtime lifecycle container.
|
||||
// It embeds Go standard context.Context compatibility, hierarchical scoping,
|
||||
// service resolution, and LIFO disposer teardown.
|
||||
type Context struct {
|
||||
goCtx context.Context
|
||||
cancel context.CancelFunc
|
||||
parent *Context
|
||||
container *Container
|
||||
|
||||
events *EventBus
|
||||
router extpoints.RouterExtension
|
||||
migrations extpoints.MigrationExtension
|
||||
tasks extpoints.TaskExtension
|
||||
schedules extpoints.ScheduleExtension
|
||||
settings extpoints.SettingExtension
|
||||
|
||||
mu sync.RWMutex
|
||||
children []*Context
|
||||
disposers []Disposer
|
||||
drivers []Driver
|
||||
values map[any]any
|
||||
disposed bool
|
||||
}
|
||||
|
||||
// NewContext creates a new root Context wrapping a standard Go context.
|
||||
// If base is nil, context.Background() is used by default.
|
||||
//
|
||||
//nolint:contextcheck
|
||||
func NewContext(base context.Context) *Context {
|
||||
if base == nil {
|
||||
base = context.Background()
|
||||
}
|
||||
ctx, cancel := context.WithCancel(base)
|
||||
|
||||
return &Context{
|
||||
goCtx: ctx,
|
||||
cancel: cancel,
|
||||
container: NewContainer(nil),
|
||||
events: NewEventBus(),
|
||||
router: extpoints.NewRouterRegistry(),
|
||||
migrations: extpoints.NewMigrationRegistry(),
|
||||
tasks: extpoints.NewTaskRegistry(),
|
||||
schedules: extpoints.NewScheduleRegistry(),
|
||||
settings: extpoints.NewSettingRegistry(),
|
||||
values: make(map[any]any),
|
||||
}
|
||||
}
|
||||
|
||||
// Deadline returns the time when work done on behalf of this context should be canceled.
|
||||
func (c *Context) Deadline() (deadline time.Time, ok bool) {
|
||||
return c.goCtx.Deadline()
|
||||
}
|
||||
|
||||
// Done returns a channel that's closed when work done on behalf of this context should be canceled.
|
||||
func (c *Context) Done() <-chan struct{} {
|
||||
return c.goCtx.Done()
|
||||
}
|
||||
|
||||
// Err returns a non-nil error value after Done is closed.
|
||||
func (c *Context) Err() error {
|
||||
return c.goCtx.Err()
|
||||
}
|
||||
|
||||
// Value returns the value associated with key, searching the local values map,
|
||||
// the underlying Go context, and fallback parent Contexts.
|
||||
func (c *Context) Value(key any) any {
|
||||
c.mu.RLock()
|
||||
if v, ok := c.values[key]; ok {
|
||||
c.mu.RUnlock()
|
||||
return v
|
||||
}
|
||||
c.mu.RUnlock()
|
||||
|
||||
if v := c.goCtx.Value(key); v != nil {
|
||||
return v
|
||||
}
|
||||
|
||||
if c.parent != nil {
|
||||
return c.parent.Value(key)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// GoContext returns the underlying standard Go context.Context.
|
||||
func (c *Context) GoContext() context.Context {
|
||||
return c.goCtx
|
||||
}
|
||||
|
||||
// Set stores an arbitrary key-value pair in this Context's local storage.
|
||||
func (c *Context) Set(key any, val any) {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
if c.values == nil {
|
||||
c.values = make(map[any]any)
|
||||
}
|
||||
c.values[key] = val
|
||||
}
|
||||
|
||||
// Get retrieves a key-value pair from this Context's local storage.
|
||||
func (c *Context) Get(key any) (any, bool) {
|
||||
c.mu.RLock()
|
||||
defer c.mu.RUnlock()
|
||||
if c.values == nil {
|
||||
return nil, false
|
||||
}
|
||||
v, ok := c.values[key]
|
||||
return v, ok
|
||||
}
|
||||
|
||||
// Container returns the underlying IoC container for this Context.
|
||||
func (c *Context) Container() *Container {
|
||||
return c.container
|
||||
}
|
||||
|
||||
// Parent returns the parent Context, or nil if this is a root Context.
|
||||
func (c *Context) Parent() *Context {
|
||||
return c.parent
|
||||
}
|
||||
|
||||
// Fork creates a child Context with its own scoped IoC container and values,
|
||||
// linked to this Context for hierarchical fallback resolution and cascading teardown.
|
||||
func (c *Context) Fork() *Context {
|
||||
return c.ForkWithContext(c.goCtx)
|
||||
}
|
||||
|
||||
// ForkWithContext creates a child Context using a specific standard Go context.
|
||||
//
|
||||
//nolint:contextcheck
|
||||
func (c *Context) ForkWithContext(base context.Context) *Context {
|
||||
if base == nil {
|
||||
base = c.goCtx
|
||||
}
|
||||
ctx, cancel := context.WithCancel(base)
|
||||
|
||||
child := &Context{
|
||||
goCtx: ctx,
|
||||
cancel: cancel,
|
||||
parent: c,
|
||||
container: NewContainer(c.container),
|
||||
events: c.Events(),
|
||||
router: c.Router(),
|
||||
migrations: c.Migrations(),
|
||||
tasks: c.Tasks(),
|
||||
schedules: c.Schedules(),
|
||||
settings: c.Settings(),
|
||||
values: make(map[any]any),
|
||||
}
|
||||
|
||||
c.mu.Lock()
|
||||
c.children = append(c.children, child)
|
||||
c.mu.Unlock()
|
||||
|
||||
return child
|
||||
}
|
||||
|
||||
// Events returns the domain EventBus associated with this Context hierarchy.
|
||||
func (c *Context) Events() *EventBus {
|
||||
return c.events
|
||||
}
|
||||
|
||||
// Router returns the RouterExtension registry.
|
||||
func (c *Context) Router() extpoints.RouterExtension {
|
||||
return c.router
|
||||
}
|
||||
|
||||
// Migrations returns the MigrationExtension registry.
|
||||
func (c *Context) Migrations() extpoints.MigrationExtension {
|
||||
return c.migrations
|
||||
}
|
||||
|
||||
// Tasks returns the TaskExtension registry.
|
||||
func (c *Context) Tasks() extpoints.TaskExtension {
|
||||
return c.tasks
|
||||
}
|
||||
|
||||
// Task is an alias for Tasks().
|
||||
func (c *Context) Task() extpoints.TaskExtension {
|
||||
return c.tasks
|
||||
}
|
||||
|
||||
// Schedules returns the ScheduleExtension registry.
|
||||
func (c *Context) Schedules() extpoints.ScheduleExtension {
|
||||
return c.schedules
|
||||
}
|
||||
|
||||
// Schedule is an alias for Schedules().
|
||||
func (c *Context) Schedule() extpoints.ScheduleExtension {
|
||||
return c.schedules
|
||||
}
|
||||
|
||||
// Settings returns the SettingExtension registry.
|
||||
func (c *Context) Settings() extpoints.SettingExtension {
|
||||
return c.settings
|
||||
}
|
||||
|
||||
// Setting is an alias for Settings().
|
||||
func (c *Context) Setting() extpoints.SettingExtension {
|
||||
return c.settings
|
||||
}
|
||||
|
||||
// OnDispose registers a cleanup callback function to be executed when this Context is disposed.
|
||||
// It accepts func() error, func(), or Disposer.
|
||||
func (c *Context) OnDispose(fn any) {
|
||||
if fn == nil {
|
||||
return
|
||||
}
|
||||
|
||||
var d Disposer
|
||||
switch f := fn.(type) {
|
||||
case Disposer:
|
||||
d = f
|
||||
case func() error:
|
||||
d = f
|
||||
case func():
|
||||
d = func() error {
|
||||
f()
|
||||
return nil
|
||||
}
|
||||
default:
|
||||
panic(fmt.Sprintf("core: OnDispose expects func() error or func(), got %T", fn))
|
||||
}
|
||||
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
c.disposers = append(c.disposers, d)
|
||||
}
|
||||
|
||||
// Dispose shuts down this Context and all child Contexts, running registered disposers in LIFO order.
|
||||
func (c *Context) Dispose() error {
|
||||
c.mu.Lock()
|
||||
if c.disposed {
|
||||
c.mu.Unlock()
|
||||
return nil
|
||||
}
|
||||
c.disposed = true
|
||||
|
||||
// Copy children and disposers under lock
|
||||
children := make([]*Context, len(c.children))
|
||||
copy(children, c.children)
|
||||
|
||||
disposers := make([]Disposer, len(c.disposers))
|
||||
copy(disposers, c.disposers)
|
||||
c.mu.Unlock()
|
||||
|
||||
var errs []error
|
||||
|
||||
// 1. Dispose all child contexts in reverse order
|
||||
for i := len(children) - 1; i >= 0; i-- {
|
||||
if err := children[i].Dispose(); err != nil {
|
||||
errs = append(errs, err)
|
||||
}
|
||||
}
|
||||
|
||||
// 2. Run local disposers in LIFO order
|
||||
for i := len(disposers) - 1; i >= 0; i-- {
|
||||
if err := disposers[i](); err != nil {
|
||||
errs = append(errs, err)
|
||||
}
|
||||
}
|
||||
|
||||
// 3. Cancel the Go context
|
||||
if c.cancel != nil {
|
||||
c.cancel()
|
||||
}
|
||||
|
||||
// 4. Detach from parent
|
||||
if c.parent != nil {
|
||||
c.parent.removeChild(c)
|
||||
}
|
||||
|
||||
return errors.Join(errs...)
|
||||
}
|
||||
|
||||
func (c *Context) removeChild(target *Context) {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
for i, child := range c.children {
|
||||
if child == target {
|
||||
c.children = append(c.children[:i], c.children[i+1:]...)
|
||||
break
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// IsDisposed returns true if this Context has been disposed.
|
||||
func (c *Context) IsDisposed() bool {
|
||||
c.mu.RLock()
|
||||
defer c.mu.RUnlock()
|
||||
return c.disposed
|
||||
}
|
||||
|
||||
// RegisterDriver registers a runtime driver engine on this Context.
|
||||
func (c *Context) RegisterDriver(d Driver) error {
|
||||
if d == nil {
|
||||
return ErrNilService
|
||||
}
|
||||
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
c.drivers = append(c.drivers, d)
|
||||
return nil
|
||||
}
|
||||
|
||||
// Drivers returns a copy of all drivers registered on this Context.
|
||||
func (c *Context) Drivers() []Driver {
|
||||
c.mu.RLock()
|
||||
defer c.mu.RUnlock()
|
||||
|
||||
result := make([]Driver, len(c.drivers))
|
||||
copy(result, c.drivers)
|
||||
return result
|
||||
}
|
||||
|
||||
// Driver looks up a registered driver by its driver type.
|
||||
func (c *Context) Driver(driverType DriverType) (Driver, bool) {
|
||||
c.mu.RLock()
|
||||
defer c.mu.RUnlock()
|
||||
|
||||
for _, d := range c.drivers {
|
||||
if d.Type() == driverType {
|
||||
return d, true
|
||||
}
|
||||
}
|
||||
return nil, false
|
||||
}
|
||||
@@ -0,0 +1,502 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package core_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/backend/core"
|
||||
)
|
||||
|
||||
// Sample services for testing
|
||||
type SampleService interface {
|
||||
Greet(name string) string
|
||||
}
|
||||
|
||||
type sampleServiceImpl struct {
|
||||
prefix string
|
||||
}
|
||||
|
||||
func (s *sampleServiceImpl) Greet(name string) string {
|
||||
if s.prefix != "" {
|
||||
return s.prefix + " " + name
|
||||
}
|
||||
return "Hello, " + name
|
||||
}
|
||||
|
||||
type LogService interface {
|
||||
Log(msg string)
|
||||
}
|
||||
|
||||
type logServiceImpl struct {
|
||||
logs []string
|
||||
}
|
||||
|
||||
func (l *logServiceImpl) Log(msg string) {
|
||||
l.logs = append(l.logs, msg)
|
||||
}
|
||||
|
||||
type ConfigService interface {
|
||||
Get(key string) string
|
||||
}
|
||||
|
||||
type configServiceImpl struct {
|
||||
data map[string]string
|
||||
}
|
||||
|
||||
func (c *configServiceImpl) Get(key string) string {
|
||||
return c.data[key]
|
||||
}
|
||||
|
||||
// Sample plugin for testing
|
||||
type samplePlugin struct {
|
||||
name string
|
||||
}
|
||||
|
||||
func (p *samplePlugin) Name() string {
|
||||
return p.name
|
||||
}
|
||||
|
||||
func (p *samplePlugin) Apply(ctx *core.Context) error {
|
||||
core.Provide[SampleService](ctx, &sampleServiceImpl{prefix: "Plugin:"})
|
||||
return nil
|
||||
}
|
||||
|
||||
func (p *samplePlugin) Manifest() core.Manifest {
|
||||
return core.Manifest{
|
||||
Name: p.name,
|
||||
Version: "1.0.0",
|
||||
Description: "Sample plugin",
|
||||
}
|
||||
}
|
||||
|
||||
// Sample driver for testing
|
||||
type mockDriver struct {
|
||||
driverType core.DriverType
|
||||
started bool
|
||||
stopped bool
|
||||
}
|
||||
|
||||
func (m *mockDriver) Type() core.DriverType {
|
||||
return m.driverType
|
||||
}
|
||||
|
||||
func (m *mockDriver) Start(ctx context.Context) error {
|
||||
m.started = true
|
||||
return nil
|
||||
}
|
||||
|
||||
func (m *mockDriver) Stop(ctx context.Context) error {
|
||||
m.stopped = true
|
||||
return nil
|
||||
}
|
||||
|
||||
func TestContextProvideAndInject(t *testing.T) {
|
||||
ctx := core.NewContext(context.Background())
|
||||
|
||||
// Before providing, Inject should fail
|
||||
_, err := core.Inject[SampleService](ctx)
|
||||
require.Error(t, err)
|
||||
assert.True(t, errors.Is(err, core.ErrServiceNotFound))
|
||||
assert.False(t, core.Has[SampleService](ctx))
|
||||
|
||||
// MustInject should panic
|
||||
assert.Panics(t, func() {
|
||||
core.MustInject[SampleService](ctx)
|
||||
})
|
||||
|
||||
// Provide service
|
||||
svcImpl := &sampleServiceImpl{prefix: "Hello,"}
|
||||
core.Provide[SampleService](ctx, svcImpl)
|
||||
|
||||
// Inject should succeed
|
||||
assert.True(t, core.Has[SampleService](ctx))
|
||||
svc, err := core.Inject[SampleService](ctx)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "Hello, Wavelet", svc.Greet("Wavelet"))
|
||||
|
||||
// MustInject should succeed
|
||||
mustSvc := core.MustInject[SampleService](ctx)
|
||||
assert.Equal(t, "Hello, Cordis", mustSvc.Greet("Cordis"))
|
||||
}
|
||||
|
||||
func TestContextProvideNilPanics(t *testing.T) {
|
||||
ctx := core.NewContext(context.Background())
|
||||
|
||||
assert.Panics(t, func() {
|
||||
core.Provide[SampleService](nil, &sampleServiceImpl{})
|
||||
})
|
||||
|
||||
assert.Panics(t, func() {
|
||||
var nilSvc SampleService
|
||||
core.Provide[SampleService](ctx, nilSvc)
|
||||
})
|
||||
|
||||
assert.Panics(t, func() {
|
||||
var nilImpl *sampleServiceImpl
|
||||
core.Provide[*sampleServiceImpl](ctx, nilImpl)
|
||||
})
|
||||
|
||||
// Inject with nil context
|
||||
var nilCtx *core.Context
|
||||
_, err := core.Inject[SampleService](nilCtx)
|
||||
assert.ErrorIs(t, err, core.ErrNilContext)
|
||||
}
|
||||
|
||||
func TestContextUsing(t *testing.T) {
|
||||
ctx := core.NewContext(context.Background())
|
||||
var called bool
|
||||
|
||||
// Using when service not ready should return ErrServiceNotReady
|
||||
err := core.Using(ctx, func(s SampleService) {
|
||||
called = true
|
||||
assert.Equal(t, "Hello, Cordis", s.Greet("Cordis"))
|
||||
})
|
||||
assert.Error(t, err)
|
||||
assert.True(t, errors.Is(err, core.ErrServiceNotReady))
|
||||
assert.False(t, called)
|
||||
|
||||
// Provide service and try Using again
|
||||
core.Provide[SampleService](ctx, &sampleServiceImpl{})
|
||||
err = core.Using(ctx, func(s SampleService) {
|
||||
called = true
|
||||
assert.Equal(t, "Hello, Cordis", s.Greet("Cordis"))
|
||||
})
|
||||
assert.NoError(t, err)
|
||||
assert.True(t, called)
|
||||
}
|
||||
|
||||
func TestContextUsingMultiple(t *testing.T) {
|
||||
ctx := core.NewContext(context.Background())
|
||||
|
||||
// Using2 with missing dependencies
|
||||
var called2 bool
|
||||
err := core.Using2(ctx, func(s SampleService, l LogService) {
|
||||
called2 = true
|
||||
})
|
||||
assert.Error(t, err)
|
||||
assert.False(t, called2)
|
||||
|
||||
// Provide 1 of 2
|
||||
core.Provide[SampleService](ctx, &sampleServiceImpl{})
|
||||
err = core.Using2(ctx, func(s SampleService, l LogService) {
|
||||
called2 = true
|
||||
})
|
||||
assert.Error(t, err)
|
||||
assert.False(t, called2)
|
||||
|
||||
// Provide 2 of 2
|
||||
logSvc := &logServiceImpl{}
|
||||
core.Provide[LogService](ctx, logSvc)
|
||||
err = core.Using2(ctx, func(s SampleService, l LogService) {
|
||||
called2 = true
|
||||
l.Log(s.Greet("World"))
|
||||
})
|
||||
assert.NoError(t, err)
|
||||
assert.True(t, called2)
|
||||
assert.Equal(t, []string{"Hello, World"}, logSvc.logs)
|
||||
|
||||
// Using3 test - error condition
|
||||
err = core.Using3(ctx, func(s SampleService, l LogService, c ConfigService) {})
|
||||
assert.Error(t, err)
|
||||
|
||||
// Using3 test - success condition
|
||||
var called3 bool
|
||||
cfgSvc := &configServiceImpl{data: map[string]string{"env": "test"}}
|
||||
core.Provide[ConfigService](ctx, cfgSvc)
|
||||
|
||||
err = core.Using3(ctx, func(s SampleService, l LogService, c ConfigService) {
|
||||
called3 = true
|
||||
assert.Equal(t, "test", c.Get("env"))
|
||||
})
|
||||
assert.NoError(t, err)
|
||||
assert.True(t, called3)
|
||||
}
|
||||
|
||||
func TestContextHierarchyAndFork(t *testing.T) {
|
||||
parent := core.NewContext(nil) // nil base context test
|
||||
core.Provide[SampleService](parent, &sampleServiceImpl{prefix: "Parent:"})
|
||||
|
||||
child := parent.ForkWithContext(nil) // nil child context test
|
||||
require.NotNil(t, child)
|
||||
assert.Equal(t, parent, child.Parent())
|
||||
|
||||
// Child can resolve service from parent
|
||||
svc, err := core.Inject[SampleService](child)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "Parent: Ryan", svc.Greet("Ryan"))
|
||||
|
||||
// Child provides LogService
|
||||
childLog := &logServiceImpl{}
|
||||
core.Provide[LogService](child, childLog)
|
||||
|
||||
// Child has LogService, parent does not
|
||||
assert.True(t, core.Has[LogService](child))
|
||||
assert.False(t, core.Has[LogService](parent))
|
||||
|
||||
// Child overrides SampleService
|
||||
core.Provide[SampleService](child, &sampleServiceImpl{prefix: "Child:"})
|
||||
childSvc, err := core.Inject[SampleService](child)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "Child: Ryan", childSvc.Greet("Ryan"))
|
||||
|
||||
parentSvc, err := core.Inject[SampleService](parent)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "Parent: Ryan", parentSvc.Greet("Ryan"))
|
||||
}
|
||||
|
||||
func TestContextReactiveWhen(t *testing.T) {
|
||||
ctx := core.NewContext(context.Background())
|
||||
|
||||
assert.Panics(t, func() {
|
||||
core.When[SampleService](nil, func(s SampleService) {})
|
||||
})
|
||||
|
||||
var whenCalled atomic.Bool
|
||||
var greeted string
|
||||
|
||||
// Register When before service is provided
|
||||
core.When[SampleService](ctx, func(s SampleService) {
|
||||
whenCalled.Store(true)
|
||||
greeted = s.Greet("Reactive")
|
||||
})
|
||||
|
||||
assert.False(t, whenCalled.Load())
|
||||
|
||||
// Now Provide the service - listener should trigger
|
||||
core.Provide[SampleService](ctx, &sampleServiceImpl{})
|
||||
|
||||
assert.True(t, whenCalled.Load())
|
||||
assert.Equal(t, "Hello, Reactive", greeted)
|
||||
|
||||
// Register another When after service is already provided - should trigger immediately
|
||||
var immediateCalled bool
|
||||
core.When[SampleService](ctx, func(s SampleService) {
|
||||
immediateCalled = true
|
||||
})
|
||||
assert.True(t, immediateCalled)
|
||||
}
|
||||
|
||||
func TestContextDisposerLifecycle(t *testing.T) {
|
||||
parent := core.NewContext(context.Background())
|
||||
child := parent.Fork()
|
||||
|
||||
var order []string
|
||||
|
||||
// Test nil disposer
|
||||
parent.OnDispose(nil)
|
||||
|
||||
// Test Disposer type
|
||||
var customDisposer core.Disposer = func() error {
|
||||
order = append(order, "parent-custom")
|
||||
return nil
|
||||
}
|
||||
parent.OnDispose(customDisposer)
|
||||
|
||||
parent.OnDispose(func() error {
|
||||
order = append(order, "parent-1")
|
||||
return nil
|
||||
})
|
||||
parent.OnDispose(func() {
|
||||
order = append(order, "parent-2")
|
||||
})
|
||||
|
||||
child.OnDispose(func() error {
|
||||
order = append(order, "child-1")
|
||||
return errors.New("child-1 error")
|
||||
})
|
||||
child.OnDispose(func() {
|
||||
order = append(order, "child-2")
|
||||
})
|
||||
|
||||
assert.Panics(t, func() {
|
||||
parent.OnDispose("invalid-func")
|
||||
})
|
||||
|
||||
assert.False(t, parent.IsDisposed())
|
||||
assert.False(t, child.IsDisposed())
|
||||
|
||||
// Disposing parent should cascade to children first, and execute disposers in LIFO order
|
||||
err := parent.Dispose()
|
||||
assert.Error(t, err) // child-1 error should be joined
|
||||
assert.Contains(t, err.Error(), "child-1 error")
|
||||
|
||||
assert.True(t, parent.IsDisposed())
|
||||
assert.True(t, child.IsDisposed())
|
||||
|
||||
// Child disposers run in LIFO: child-2, child-1
|
||||
// Parent disposers run in LIFO: parent-2, parent-1, parent-custom
|
||||
expected := []string{"child-2", "child-1", "parent-2", "parent-1", "parent-custom"}
|
||||
assert.Equal(t, expected, order)
|
||||
|
||||
// Disposing again should be idempotent and return nil
|
||||
err = parent.Dispose()
|
||||
assert.NoError(t, err)
|
||||
}
|
||||
|
||||
func TestContextStandardGoContext(t *testing.T) {
|
||||
baseCtx, cancel := context.WithDeadline(context.Background(), time.Now().Add(5*time.Second))
|
||||
defer cancel()
|
||||
|
||||
parentCtx := core.NewContext(baseCtx)
|
||||
parentCtx.Set("parent_key", "parent_val")
|
||||
|
||||
childCtx := parentCtx.Fork()
|
||||
|
||||
// Deadline
|
||||
dl, ok := childCtx.Deadline()
|
||||
assert.True(t, ok)
|
||||
assert.False(t, dl.IsZero())
|
||||
|
||||
// Value fallback: child has no key, falls back to parentCtx
|
||||
assert.Equal(t, "parent_val", childCtx.Value("parent_key"))
|
||||
|
||||
// GoContext getter
|
||||
assert.NotNil(t, childCtx.GoContext())
|
||||
|
||||
// Value not found in either
|
||||
assert.Nil(t, childCtx.Value("non_existent_key"))
|
||||
|
||||
// Cancellation propagation
|
||||
select {
|
||||
case <-childCtx.Done():
|
||||
t.Fatal("ctx should not be done yet")
|
||||
default:
|
||||
}
|
||||
|
||||
cancel()
|
||||
|
||||
select {
|
||||
case <-childCtx.Done():
|
||||
assert.Equal(t, context.Canceled, childCtx.Err())
|
||||
case <-time.After(100 * time.Millisecond):
|
||||
t.Fatal("ctx should be cancelled")
|
||||
}
|
||||
}
|
||||
|
||||
func TestManifestValidation(t *testing.T) {
|
||||
mValid := core.Manifest{
|
||||
Name: "auth",
|
||||
Version: "1.0.0",
|
||||
Description: "Auth plugin",
|
||||
}
|
||||
assert.NoError(t, mValid.Validate())
|
||||
|
||||
mInvalid := core.Manifest{
|
||||
Version: "1.0.0",
|
||||
}
|
||||
assert.Error(t, mInvalid.Validate())
|
||||
}
|
||||
|
||||
func TestDriverRegistration(t *testing.T) {
|
||||
ctx := core.NewContext(context.Background())
|
||||
|
||||
// Register nil driver returns error
|
||||
assert.ErrorIs(t, ctx.RegisterDriver(nil), core.ErrNilService)
|
||||
|
||||
dHTTP := &mockDriver{driverType: core.DriverTypeHTTP}
|
||||
dWorker := &mockDriver{driverType: core.DriverTypeWorker}
|
||||
|
||||
require.NoError(t, ctx.RegisterDriver(dHTTP))
|
||||
require.NoError(t, ctx.RegisterDriver(dWorker))
|
||||
|
||||
drivers := ctx.Drivers()
|
||||
assert.Len(t, drivers, 2)
|
||||
|
||||
foundHTTP, ok := ctx.Driver(core.DriverTypeHTTP)
|
||||
assert.True(t, ok)
|
||||
assert.Equal(t, dHTTP, foundHTTP)
|
||||
|
||||
foundWorker, ok := ctx.Driver(core.DriverTypeWorker)
|
||||
assert.True(t, ok)
|
||||
assert.Equal(t, dWorker, foundWorker)
|
||||
|
||||
_, ok = ctx.Driver(core.DriverTypeScheduler)
|
||||
assert.False(t, ok)
|
||||
}
|
||||
|
||||
func TestPluginInterfaces(t *testing.T) {
|
||||
ctx := core.NewContext(context.Background())
|
||||
var p core.Plugin = &samplePlugin{name: "sample"}
|
||||
assert.Equal(t, "sample", p.Name())
|
||||
require.NoError(t, p.Apply(ctx))
|
||||
|
||||
svc, err := core.Inject[SampleService](ctx)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "Plugin: Ryan", svc.Greet("Ryan"))
|
||||
|
||||
var pwm core.PluginWithManifest = &samplePlugin{name: "sample"}
|
||||
manifest := pwm.Manifest()
|
||||
assert.Equal(t, "sample", manifest.Name)
|
||||
assert.Equal(t, "1.0.0", manifest.Version)
|
||||
}
|
||||
|
||||
func TestConcurrentAccess(t *testing.T) {
|
||||
ctx := core.NewContext(context.Background())
|
||||
var wg sync.WaitGroup
|
||||
|
||||
// Concurrently provide, inject, fork, set, and get
|
||||
for i := 0; i < 50; i++ {
|
||||
wg.Add(1)
|
||||
go func(idx int) {
|
||||
defer wg.Done()
|
||||
ctx.Set(fmt.Sprintf("key-%d", idx), idx)
|
||||
_, _ = ctx.Get(fmt.Sprintf("key-%d", idx))
|
||||
|
||||
child := ctx.Fork()
|
||||
child.Set("child_key", idx)
|
||||
}(i)
|
||||
}
|
||||
|
||||
core.Provide[SampleService](ctx, &sampleServiceImpl{})
|
||||
|
||||
for i := 0; i < 50; i++ {
|
||||
wg.Add(1)
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
svc, err := core.Inject[SampleService](ctx)
|
||||
if err == nil {
|
||||
_ = svc.Greet("Concurrency")
|
||||
}
|
||||
_ = core.Using(ctx, func(s SampleService) {
|
||||
_ = s.Greet("Safe")
|
||||
})
|
||||
}()
|
||||
}
|
||||
|
||||
wg.Wait()
|
||||
}
|
||||
|
||||
func TestContextExtensionPointsAccessors(t *testing.T) {
|
||||
ctx := core.NewContext(nil)
|
||||
assert.NotNil(t, ctx.Events())
|
||||
assert.NotNil(t, ctx.Router())
|
||||
assert.NotNil(t, ctx.Migrations())
|
||||
assert.NotNil(t, ctx.Tasks())
|
||||
assert.NotNil(t, ctx.Task())
|
||||
assert.NotNil(t, ctx.Schedules())
|
||||
assert.NotNil(t, ctx.Schedule())
|
||||
assert.NotNil(t, ctx.Settings())
|
||||
assert.NotNil(t, ctx.Setting())
|
||||
|
||||
child := ctx.Fork()
|
||||
assert.Equal(t, ctx.Events(), child.Events())
|
||||
assert.Equal(t, ctx.Router(), child.Router())
|
||||
assert.Equal(t, ctx.Migrations(), child.Migrations())
|
||||
assert.Equal(t, ctx.Tasks(), child.Tasks())
|
||||
assert.Equal(t, ctx.Task(), child.Task())
|
||||
assert.Equal(t, ctx.Schedules(), child.Schedules())
|
||||
assert.Equal(t, ctx.Schedule(), child.Schedule())
|
||||
assert.Equal(t, ctx.Settings(), child.Settings())
|
||||
assert.Equal(t, ctx.Setting(), child.Setting())
|
||||
}
|
||||
@@ -0,0 +1,100 @@
|
||||
// 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"
|
||||
)
|
||||
|
||||
// UserDTO represents a unified user data transfer object across plugins.
|
||||
type UserDTO struct {
|
||||
ID uint64 `json:"id,string"`
|
||||
Username string `json:"username"`
|
||||
Nickname string `json:"nickname"`
|
||||
Email string `json:"email"`
|
||||
AvatarURL string `json:"avatar_url"`
|
||||
IsActive bool `json:"is_active"`
|
||||
IsAdmin bool `json:"is_admin"`
|
||||
NeedChangePassword bool `json:"need_change_password,omitempty"`
|
||||
Bio string `json:"bio,omitempty"`
|
||||
Phone string `json:"phone,omitempty"`
|
||||
Gender string `json:"gender,omitempty"`
|
||||
Website string `json:"website,omitempty"`
|
||||
Location string `json:"location,omitempty"`
|
||||
LastLoginAt time.Time `json:"last_login_at"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
UpdatedAt time.Time `json:"updated_at"`
|
||||
}
|
||||
|
||||
// TableName returns w_users.
|
||||
func (UserDTO) TableName() string {
|
||||
return "w_users"
|
||||
}
|
||||
|
||||
// OAuthUserInfoDTO contains user identity claims obtained from an OAuth provider.
|
||||
type OAuthUserInfoDTO struct {
|
||||
ID uint64 `json:"id"`
|
||||
Sub string `json:"sub"`
|
||||
Username string `json:"username"`
|
||||
PreferredUsername string `json:"preferred_username"`
|
||||
Email string `json:"email"`
|
||||
Name string `json:"name"`
|
||||
Active bool `json:"active"`
|
||||
AvatarURL string `json:"avatar_url"`
|
||||
}
|
||||
|
||||
// OAuthProvider defines the pluggable OAuth provider contract.
|
||||
type OAuthProvider interface {
|
||||
Name() string
|
||||
GetAuthURL(state string) string
|
||||
ExchangeCode(ctx context.Context, code string) (*OAuthUserInfoDTO, error)
|
||||
}
|
||||
|
||||
// AuthService defines the contract for authentication, session verification, and token management.
|
||||
type AuthService interface {
|
||||
// RequireAuthMiddleware returns a middleware handler (compatible with gin.HandlerFunc or standard middleware).
|
||||
RequireAuthMiddleware() any
|
||||
|
||||
// RequireAdminMiddleware returns an admin authorization middleware.
|
||||
RequireAdminMiddleware() any
|
||||
|
||||
// GetCurrentUser retrieves the authenticated UserDTO from context.
|
||||
GetCurrentUser(ctx context.Context) (*UserDTO, error)
|
||||
|
||||
// GetCurrentUserID retrieves the authenticated user ID from session/context.
|
||||
GetCurrentUserID(ctx context.Context) (uint64, error)
|
||||
|
||||
// VerifyToken validates an access token and returns the associated user DTO.
|
||||
VerifyToken(ctx context.Context, token string) (*UserDTO, error)
|
||||
|
||||
// CreateSession establishes an authenticated session for the given user ID.
|
||||
CreateSession(ctx context.Context, userID uint64, extras map[string]any) (string, error)
|
||||
|
||||
// RevokeToken invalidates a specific access token by its hash.
|
||||
RevokeToken(ctx context.Context, tokenHash string) error
|
||||
|
||||
// RevokeUserSessions revokes all active sessions and cached tokens for a user.
|
||||
RevokeUserSessions(ctx context.Context, userID uint64) error
|
||||
|
||||
// DisallowTokenAuthMiddleware returns a middleware that rejects requests authenticated via access token.
|
||||
DisallowTokenAuthMiddleware() any
|
||||
}
|
||||
|
||||
// AuthRegistry allows downstream and domain plugins to register custom authentication providers.
|
||||
type AuthRegistry interface {
|
||||
RegisterOAuthProvider(name string, provider OAuthProvider)
|
||||
GetOAuthProvider(name string) (OAuthProvider, bool)
|
||||
ListOAuthProviders() []string
|
||||
}
|
||||
|
||||
// Auth context keys — stored in Gin context by auth middleware, consumed by domain plugins.
|
||||
const (
|
||||
AuthUserIDKey = "user_id"
|
||||
AuthUserNameKey = "username"
|
||||
AuthUserObjKey = "user_obj"
|
||||
AuthTokenAuthKey = "token_auth" // marks if request uses access token auth
|
||||
AuthTokenAdminKey = "token_admin" // whether the access token has admin privileges
|
||||
)
|
||||
@@ -0,0 +1,32 @@
|
||||
// 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"
|
||||
"errors"
|
||||
"time"
|
||||
)
|
||||
|
||||
// ErrCacheMiss is returned when an item is not found in the cache.
|
||||
var ErrCacheMiss = errors.New("contracts/cache: key not found")
|
||||
|
||||
// CacheService defines the contract for multi-layer cache operations (RAM L1 + Redis L2 + Pub/Sub invalidation).
|
||||
type CacheService interface {
|
||||
// Get retrieves an item from cache into target. Returns ErrCacheMiss if not found.
|
||||
Get(ctx context.Context, key string, target any) error
|
||||
|
||||
// Set stores an item into cache with a specified time-to-live duration.
|
||||
Set(ctx context.Context, key string, value any, ttl time.Duration) error
|
||||
|
||||
// Delete evicts a key from local and remote cache tiers and broadcasts invalidation.
|
||||
Delete(ctx context.Context, key string) error
|
||||
|
||||
// GetOrSet retrieves an item from cache, or calls loader to populate and return if missing.
|
||||
GetOrSet(ctx context.Context, key string, target any, ttl time.Duration, loader func() (any, error)) error
|
||||
|
||||
// Invalidate is a semantic alias for Delete.
|
||||
Invalidate(ctx context.Context, key string) error
|
||||
}
|
||||
@@ -0,0 +1,23 @@
|
||||
// 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"
|
||||
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
// DBService defines the standard contract for relational database access and multi-datasource routing.
|
||||
type DBService interface {
|
||||
// GORM returns the underlying GORM database instance.
|
||||
GORM() *gorm.DB
|
||||
|
||||
// DB returns the GORM database instance bound to the given context.
|
||||
DB(ctx context.Context) *gorm.DB
|
||||
|
||||
// Named returns a named database connection if multiple data sources or replicas are configured.
|
||||
Named(name string) *gorm.DB
|
||||
}
|
||||
@@ -0,0 +1,109 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package contracts defines unified service interfaces and DTOs for cross-plugin communication.
|
||||
package contracts
|
||||
|
||||
// ======================================================================
|
||||
// Domain Event Topic Constants
|
||||
// ======================================================================
|
||||
//
|
||||
// All cross-plugin domain event topics MUST be declared here so that
|
||||
// producers and consumers share the same string values without importing
|
||||
// each other's implementation packages.
|
||||
// ======================================================================
|
||||
|
||||
// --- Auth & User Events ---
|
||||
const (
|
||||
// EventTopicAdminLoggedIn fires when an admin user logs in.
|
||||
EventTopicAdminLoggedIn = "admin:logged_in"
|
||||
|
||||
// EventTopicUserCreated fires when a new user account is created.
|
||||
EventTopicUserCreated = "user:created"
|
||||
|
||||
// EventTopicUserUpdated fires when a user profile is updated.
|
||||
EventTopicUserUpdated = "user:updated"
|
||||
|
||||
// EventTopicUserDeleted fires when a user account is deleted.
|
||||
EventTopicUserDeleted = "user:deleted"
|
||||
)
|
||||
|
||||
// --- Admin & System Events ---
|
||||
const (
|
||||
// EventTopicConfigChanged fires when a system configuration value changes.
|
||||
EventTopicConfigChanged = "admin:config_changed"
|
||||
|
||||
// EventTopicSystemCleanup fires when a periodic system cleanup completes.
|
||||
EventTopicSystemCleanup = "admin:system_cleanup"
|
||||
)
|
||||
|
||||
// --- Upload / Storage Events ---
|
||||
const (
|
||||
// EventTopicUploadCreated fires when a new file upload is recorded.
|
||||
EventTopicUploadCreated = "upload:created"
|
||||
|
||||
// EventTopicUploadDeleted fires when a file upload is removed.
|
||||
EventTopicUploadDeleted = "upload:deleted"
|
||||
|
||||
// EventTopicIngestComplete fires when a programmatic file ingest finishes.
|
||||
EventTopicIngestComplete = "upload:ingest_complete"
|
||||
)
|
||||
|
||||
// --- Message Gateway Events ---
|
||||
const (
|
||||
// EventTopicNotificationSent fires when a push notification is dispatched.
|
||||
EventTopicNotificationSent = "message:notification_sent"
|
||||
|
||||
// EventTopicChannelBound fires when a user binds a messaging channel.
|
||||
EventTopicChannelBound = "message:channel_bound"
|
||||
|
||||
// EventTopicChannelUnbound fires when a user unbinds a messaging channel.
|
||||
EventTopicChannelUnbound = "message:channel_unbound"
|
||||
)
|
||||
|
||||
// --- Risk Control Events ---
|
||||
const (
|
||||
// EventTopicAccessLogRecorded fires when a user access log entry is recorded.
|
||||
EventTopicAccessLogRecorded = "risk:access_log_recorded"
|
||||
)
|
||||
|
||||
// ======================================================================
|
||||
// Domain Event Payload DTOs
|
||||
// ======================================================================
|
||||
|
||||
// AdminLoggedIn 管理员登录领域事件载荷
|
||||
type AdminLoggedIn struct {
|
||||
User *UserDTO `json:"user"`
|
||||
IP string `json:"ip"`
|
||||
}
|
||||
|
||||
// UserCreatedEvent fires when a new user account is created.
|
||||
type UserCreatedEvent struct {
|
||||
User *UserDTO `json:"user"`
|
||||
Password string `json:"-"`
|
||||
}
|
||||
|
||||
// ConfigChangedEvent fires when a system configuration value changes.
|
||||
type ConfigChangedEvent struct {
|
||||
Key string `json:"key"`
|
||||
OldVal any `json:"old_val,omitempty"`
|
||||
NewVal any `json:"new_val,omitempty"`
|
||||
}
|
||||
|
||||
// UploadCreatedEvent fires when a new file upload is recorded.
|
||||
type UploadCreatedEvent struct {
|
||||
UploadID uint64 `json:"upload_id,string"`
|
||||
UserID uint64 `json:"user_id,string"`
|
||||
FileName string `json:"file_name"`
|
||||
FileSize int64 `json:"file_size"`
|
||||
MimeType string `json:"mime_type"`
|
||||
}
|
||||
|
||||
// NotificationSentEvent fires when a push notification is dispatched.
|
||||
type NotificationSentEvent struct {
|
||||
UserID uint64 `json:"user_id,string"`
|
||||
Channel string `json:"channel"`
|
||||
Title string `json:"title"`
|
||||
Success bool `json:"success"`
|
||||
ErrorInfo string `json:"error_info,omitempty"`
|
||||
}
|
||||
@@ -0,0 +1,39 @@
|
||||
// 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"
|
||||
)
|
||||
|
||||
// LoggerService defines the contract for structured logging with trace ID and context correlation.
|
||||
type LoggerService interface {
|
||||
// Debug logs a debug message with optional key-value structured fields.
|
||||
Debug(ctx context.Context, msg string, keysAndValues ...any)
|
||||
|
||||
// Info logs an informational message with optional key-value structured fields.
|
||||
Info(ctx context.Context, msg string, keysAndValues ...any)
|
||||
|
||||
// Warn logs a warning message with optional key-value structured fields.
|
||||
Warn(ctx context.Context, msg string, keysAndValues ...any)
|
||||
|
||||
// Error logs an error message with optional key-value structured fields.
|
||||
Error(ctx context.Context, msg string, keysAndValues ...any)
|
||||
|
||||
// Debugf logs a formatted debug message.
|
||||
Debugf(ctx context.Context, format string, args ...any)
|
||||
|
||||
// Infof logs a formatted informational message.
|
||||
Infof(ctx context.Context, format string, args ...any)
|
||||
|
||||
// Warnf logs a formatted warning message.
|
||||
Warnf(ctx context.Context, format string, args ...any)
|
||||
|
||||
// Errorf logs a formatted error message.
|
||||
Errorf(ctx context.Context, format string, args ...any)
|
||||
|
||||
// With returns a child logger enriched with additional key-value attributes.
|
||||
With(keysAndValues ...any) LoggerService
|
||||
}
|
||||
@@ -0,0 +1,62 @@
|
||||
// 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"
|
||||
"io"
|
||||
)
|
||||
|
||||
// StorageObject represents a retrieved file object from the storage backend.
|
||||
type StorageObject struct {
|
||||
Key string
|
||||
CachePath string
|
||||
Body io.ReadCloser
|
||||
ContentLength int64
|
||||
ContentType string
|
||||
}
|
||||
|
||||
// StoragePutResult describes the output of a successful Put operation.
|
||||
type StoragePutResult struct {
|
||||
Key string
|
||||
Bucket string
|
||||
}
|
||||
|
||||
// IngestOptions configures programmatic ingest of files into the platform storage.
|
||||
type IngestOptions struct {
|
||||
UserID uint64
|
||||
Type string
|
||||
FileName string
|
||||
MimeType string
|
||||
Extension string
|
||||
Size int64
|
||||
Policy int
|
||||
Metadata map[string]any
|
||||
}
|
||||
|
||||
// IngestResult reports the outcome of a programmatic file ingest operation.
|
||||
type IngestResult struct {
|
||||
ID uint64
|
||||
Key string
|
||||
URL string
|
||||
Created bool
|
||||
Stored bool
|
||||
Resolved bool
|
||||
}
|
||||
|
||||
// StorageService defines the contract for unified object storage and managed file ingestion.
|
||||
type StorageService interface {
|
||||
// Put writes an object to storage.
|
||||
Put(ctx context.Context, key string, body io.Reader, size int64, contentType string) (StoragePutResult, error)
|
||||
|
||||
// Get retrieves an object from storage.
|
||||
Get(ctx context.Context, key string) (*StorageObject, error)
|
||||
|
||||
// Delete removes an object from storage.
|
||||
Delete(ctx context.Context, key string) error
|
||||
|
||||
// Ingest performs managed file ingestion into the platform storage domain with deduplication and metadata tracking.
|
||||
Ingest(ctx context.Context, reader io.Reader, opts IngestOptions) (*IngestResult, error)
|
||||
}
|
||||
@@ -0,0 +1,84 @@
|
||||
// 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"
|
||||
)
|
||||
|
||||
// CreateUserRequest contains fields to register or create a new user.
|
||||
type CreateUserRequest struct {
|
||||
Username string `json:"username"`
|
||||
Password string `json:"password"`
|
||||
Nickname string `json:"nickname"`
|
||||
Email string `json:"email"`
|
||||
IsAdmin bool `json:"is_admin"`
|
||||
}
|
||||
|
||||
// UpdateUserProfileRequest contains fields for updating a user's profile.
|
||||
type UpdateUserProfileRequest struct {
|
||||
Nickname *string `json:"nickname,omitempty"`
|
||||
Email *string `json:"email,omitempty"`
|
||||
AvatarURL *string `json:"avatar_url,omitempty"`
|
||||
Bio *string `json:"bio,omitempty"`
|
||||
Phone *string `json:"phone,omitempty"`
|
||||
Gender *string `json:"gender,omitempty"`
|
||||
Website *string `json:"website,omitempty"`
|
||||
Location *string `json:"location,omitempty"`
|
||||
}
|
||||
|
||||
// UserService defines the contract for user account management and profile queries.
|
||||
type UserService interface {
|
||||
// GetUserByID retrieves a user by ID.
|
||||
GetUserByID(ctx context.Context, id uint64) (*UserDTO, error)
|
||||
|
||||
// GetUserByUsername retrieves a user by username.
|
||||
GetUserByUsername(ctx context.Context, username string) (*UserDTO, error)
|
||||
|
||||
// GetUserByEmail retrieves a user by email.
|
||||
GetUserByEmail(ctx context.Context, email string) (*UserDTO, error)
|
||||
|
||||
// CreateUser registers or creates a new user account.
|
||||
CreateUser(ctx context.Context, req CreateUserRequest) (*UserDTO, error)
|
||||
|
||||
// UpdateProfile updates the profile of the specified user.
|
||||
UpdateProfile(ctx context.Context, id uint64, req UpdateUserProfileRequest) (*UserDTO, error)
|
||||
|
||||
// UpdatePassword updates the password for the specified user after verifying the old password.
|
||||
UpdatePassword(ctx context.Context, id uint64, oldPassword, newPassword string) error
|
||||
|
||||
// VerifyPassword verifies if the given password matches the user's password.
|
||||
VerifyPassword(ctx context.Context, id uint64, password string) bool
|
||||
|
||||
// UpdateLastLogin updates the user's last login timestamp.
|
||||
UpdateLastLogin(ctx context.Context, id uint64, ip string) error
|
||||
|
||||
// ListUsers returns a paginated list of users with optional keyword search.
|
||||
ListUsers(ctx context.Context, page, pageSize int, keyword string) ([]*UserDTO, int64, error)
|
||||
|
||||
// SetUserActive sets the active/banned status for a user.
|
||||
SetUserActive(ctx context.Context, id uint64, active bool) error
|
||||
|
||||
// SetUserAdmin sets the admin role status for a user.
|
||||
SetUserAdmin(ctx context.Context, id uint64, admin bool) error
|
||||
|
||||
// VerifyAccessToken verifies an access token hash and returns the user DTO and isAdmin flag.
|
||||
VerifyAccessToken(ctx context.Context, tokenHash string) (*UserDTO, bool, error)
|
||||
|
||||
// DeleteUser removes a user and related access tokens.
|
||||
DeleteUser(ctx context.Context, id uint64) error
|
||||
|
||||
// CountUsers returns total user count.
|
||||
CountUsers(ctx context.Context) (int64, error)
|
||||
|
||||
// CountActiveUsers returns active user count.
|
||||
CountActiveUsers(ctx context.Context) (int64, error)
|
||||
|
||||
// GetFirstAdminUser returns the earliest admin user.
|
||||
GetFirstAdminUser(ctx context.Context) (*UserDTO, error)
|
||||
|
||||
// UniqueUsername generates a unique username candidate based on base.
|
||||
UniqueUsername(ctx context.Context, base string) (string, error)
|
||||
}
|
||||
@@ -0,0 +1,278 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package core
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"reflect"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
)
|
||||
|
||||
const maxHandlerParams = 2
|
||||
|
||||
var ctxInterfaceType = reflect.TypeFor[context.Context]()
|
||||
var errInterfaceType = reflect.TypeFor[error]()
|
||||
|
||||
type eventListener struct {
|
||||
id uint64
|
||||
fnVal reflect.Value
|
||||
numIn int
|
||||
hasCtx bool
|
||||
hasPayload bool
|
||||
argType reflect.Type
|
||||
returnsErr bool
|
||||
}
|
||||
|
||||
// EventBus is a thread-safe, strongly-typed in-process domain event bus.
|
||||
type EventBus struct {
|
||||
mu sync.RWMutex
|
||||
nextID atomic.Uint64
|
||||
handlers map[string][]eventListener
|
||||
}
|
||||
|
||||
// NewEventBus creates a new EventBus instance.
|
||||
func NewEventBus() *EventBus {
|
||||
return &EventBus{
|
||||
handlers: make(map[string][]eventListener),
|
||||
}
|
||||
}
|
||||
|
||||
// On registers an event handler for the given topic.
|
||||
//
|
||||
// Supported handler signatures:
|
||||
// - func(ctx context.Context, event T) error
|
||||
// - func(ctx context.Context, event T)
|
||||
// - func(event T) error
|
||||
// - func(event T)
|
||||
// - func(ctx context.Context) error
|
||||
// - func(ctx context.Context)
|
||||
// - func() error
|
||||
// - func()
|
||||
//
|
||||
// Returns a Disposer function that unregisters the handler when called.
|
||||
func (b *EventBus) On(topic string, handler any) Disposer {
|
||||
if handler == nil {
|
||||
panic("core/events: handler cannot be nil")
|
||||
}
|
||||
|
||||
fnVal := reflect.ValueOf(handler)
|
||||
fnType := fnVal.Type()
|
||||
|
||||
if fnType.Kind() != reflect.Func {
|
||||
panic(fmt.Sprintf("core/events: expected func, got %s", fnType.Kind()))
|
||||
}
|
||||
|
||||
numIn := fnType.NumIn()
|
||||
if numIn > maxHandlerParams {
|
||||
panic(fmt.Sprintf("core/events: handler has %d parameters, maximum 2 supported (ctx, event)", numIn))
|
||||
}
|
||||
|
||||
numOut := fnType.NumOut()
|
||||
if numOut > 1 {
|
||||
panic(fmt.Sprintf("core/events: handler has %d return values, maximum 1 supported (error)", numOut))
|
||||
}
|
||||
|
||||
returnsErr := false
|
||||
if numOut == 1 {
|
||||
outType := fnType.Out(0)
|
||||
if !outType.Implements(errInterfaceType) {
|
||||
panic(fmt.Sprintf("core/events: handler return type must be error, got %v", outType))
|
||||
}
|
||||
returnsErr = true
|
||||
}
|
||||
|
||||
listener := eventListener{
|
||||
id: b.nextID.Add(1),
|
||||
fnVal: fnVal,
|
||||
numIn: numIn,
|
||||
returnsErr: returnsErr,
|
||||
}
|
||||
|
||||
switch numIn {
|
||||
case 0:
|
||||
// func() or func() error
|
||||
case 1:
|
||||
in0 := fnType.In(0)
|
||||
if in0.Implements(ctxInterfaceType) {
|
||||
listener.hasCtx = true
|
||||
} else {
|
||||
listener.hasPayload = true
|
||||
listener.argType = in0
|
||||
}
|
||||
case 2:
|
||||
in0 := fnType.In(0)
|
||||
if !in0.Implements(ctxInterfaceType) {
|
||||
panic(fmt.Sprintf("core/events: first parameter must implement context.Context, got %v", in0))
|
||||
}
|
||||
listener.hasCtx = true
|
||||
listener.hasPayload = true
|
||||
listener.argType = fnType.In(1)
|
||||
}
|
||||
|
||||
b.mu.Lock()
|
||||
b.handlers[topic] = append(b.handlers[topic], listener)
|
||||
b.mu.Unlock()
|
||||
|
||||
listenerID := listener.id
|
||||
var disposed atomic.Bool
|
||||
|
||||
return func() error {
|
||||
if disposed.Swap(true) {
|
||||
return nil
|
||||
}
|
||||
|
||||
b.mu.Lock()
|
||||
defer b.mu.Unlock()
|
||||
|
||||
list := b.handlers[topic]
|
||||
for i, l := range list {
|
||||
if l.id == listenerID {
|
||||
b.handlers[topic] = append(list[:i], list[i+1:]...)
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
if len(b.handlers[topic]) == 0 {
|
||||
delete(b.handlers, topic)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
// Subscribe registers a strongly-typed generic event listener on the given EventBus.
|
||||
func Subscribe[T any](bus *EventBus, topic string, handler func(ctx context.Context, event T) error) Disposer {
|
||||
if bus == nil {
|
||||
panic("core/events: nil EventBus provided to Subscribe")
|
||||
}
|
||||
return bus.On(topic, handler)
|
||||
}
|
||||
|
||||
// Emit publishes an event to all subscribers of the specified topic.
|
||||
// Handlers are executed synchronously. If any handler panics or returns an error,
|
||||
// the error is collected and returned via errors.Join.
|
||||
//
|
||||
//nolint:contextcheck
|
||||
func (b *EventBus) Emit(ctx context.Context, topic string, payload any) error {
|
||||
if ctx == nil {
|
||||
ctx = context.Background()
|
||||
}
|
||||
|
||||
b.mu.RLock()
|
||||
rawListeners := b.handlers[topic]
|
||||
if len(rawListeners) == 0 {
|
||||
b.mu.RUnlock()
|
||||
return nil
|
||||
}
|
||||
|
||||
listeners := make([]eventListener, len(rawListeners))
|
||||
copy(listeners, rawListeners)
|
||||
b.mu.RUnlock()
|
||||
|
||||
var payloadVal reflect.Value
|
||||
if payload != nil {
|
||||
payloadVal = reflect.ValueOf(payload)
|
||||
}
|
||||
|
||||
var errs []error
|
||||
for _, l := range listeners {
|
||||
args := b.buildArgs(ctx, l, payloadVal)
|
||||
|
||||
err := func() (resErr error) {
|
||||
defer func() {
|
||||
if r := recover(); r != nil {
|
||||
resErr = fmt.Errorf("core/events: panic in handler for topic %q: %v", topic, r)
|
||||
}
|
||||
}()
|
||||
|
||||
results := l.fnVal.Call(args)
|
||||
if l.returnsErr && len(results) > 0 && !results[0].IsNil() {
|
||||
resErr = results[0].Interface().(error)
|
||||
}
|
||||
return resErr
|
||||
}()
|
||||
|
||||
if err != nil {
|
||||
errs = append(errs, err)
|
||||
}
|
||||
}
|
||||
|
||||
return errors.Join(errs...)
|
||||
}
|
||||
|
||||
func (b *EventBus) buildArgs(ctx context.Context, l eventListener, payloadVal reflect.Value) []reflect.Value {
|
||||
if l.numIn == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
args := make([]reflect.Value, 0, l.numIn)
|
||||
if l.hasCtx {
|
||||
args = append(args, reflect.ValueOf(ctx))
|
||||
}
|
||||
|
||||
if l.hasPayload {
|
||||
arg := b.convertPayload(payloadVal, l.argType)
|
||||
args = append(args, arg)
|
||||
}
|
||||
|
||||
return args
|
||||
}
|
||||
|
||||
func (b *EventBus) convertPayload(payloadVal reflect.Value, targetType reflect.Type) reflect.Value {
|
||||
if !payloadVal.IsValid() {
|
||||
return reflect.Zero(targetType)
|
||||
}
|
||||
|
||||
valType := payloadVal.Type()
|
||||
|
||||
// 1. Direct assignable
|
||||
if valType.AssignableTo(targetType) {
|
||||
return payloadVal
|
||||
}
|
||||
|
||||
// 2. Direct convertible
|
||||
if valType.ConvertibleTo(targetType) {
|
||||
return payloadVal.Convert(targetType)
|
||||
}
|
||||
|
||||
// 3. Payload is pointer *T, target expects T
|
||||
if valType.Kind() == reflect.Pointer && valType.Elem().AssignableTo(targetType) {
|
||||
if !payloadVal.IsNil() {
|
||||
return payloadVal.Elem()
|
||||
}
|
||||
return reflect.Zero(targetType)
|
||||
}
|
||||
|
||||
// 4. Payload is value T, target expects *T
|
||||
if targetType.Kind() == reflect.Pointer && valType.AssignableTo(targetType.Elem()) {
|
||||
ptr := reflect.New(valType)
|
||||
ptr.Elem().Set(payloadVal)
|
||||
return ptr
|
||||
}
|
||||
|
||||
// Fallback to zero value of targetType
|
||||
return reflect.Zero(targetType)
|
||||
}
|
||||
|
||||
// Listeners returns the number of active listeners for a topic.
|
||||
func (b *EventBus) Listeners(topic string) int {
|
||||
b.mu.RLock()
|
||||
defer b.mu.RUnlock()
|
||||
return len(b.handlers[topic])
|
||||
}
|
||||
|
||||
// Topics returns all topics that have registered listeners.
|
||||
func (b *EventBus) Topics() []string {
|
||||
b.mu.RLock()
|
||||
defer b.mu.RUnlock()
|
||||
|
||||
topics := make([]string, 0, len(b.handlers))
|
||||
for t := range b.handlers {
|
||||
topics = append(topics, t)
|
||||
}
|
||||
return topics
|
||||
}
|
||||
@@ -0,0 +1,318 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package core_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/backend/core"
|
||||
)
|
||||
|
||||
type UserRegisteredEvent struct {
|
||||
UserID string `json:"user_id"`
|
||||
Username string `json:"username"`
|
||||
}
|
||||
|
||||
type OrderCreatedEvent struct {
|
||||
OrderID string `json:"order_id"`
|
||||
Amount float64 `json:"amount"`
|
||||
}
|
||||
|
||||
func TestEventBusPublishSubscribe(t *testing.T) {
|
||||
bus := core.NewEventBus()
|
||||
var receivedID string
|
||||
|
||||
disposer := bus.On("user:registered", func(ctx context.Context, e UserRegisteredEvent) error {
|
||||
receivedID = e.UserID
|
||||
return nil
|
||||
})
|
||||
require.NotNil(t, disposer)
|
||||
|
||||
assert.Equal(t, []string{"user:registered"}, bus.Topics())
|
||||
|
||||
err := bus.Emit(context.Background(), "user:registered", UserRegisteredEvent{UserID: "u_999", Username: "alice"})
|
||||
assert.NoError(t, err)
|
||||
assert.Equal(t, "u_999", receivedID)
|
||||
|
||||
// Emit to empty topic returns nil error
|
||||
err = bus.Emit(nil, "empty:topic", nil)
|
||||
assert.NoError(t, err)
|
||||
}
|
||||
|
||||
func TestEventBusGenericSubscribe(t *testing.T) {
|
||||
bus := core.NewEventBus()
|
||||
var receivedOrder string
|
||||
|
||||
assert.Panics(t, func() {
|
||||
core.Subscribe[OrderCreatedEvent](nil, "order:created", func(ctx context.Context, e OrderCreatedEvent) error {
|
||||
return nil
|
||||
})
|
||||
})
|
||||
|
||||
disposer := core.Subscribe(bus, "order:created", func(ctx context.Context, e OrderCreatedEvent) error {
|
||||
receivedOrder = e.OrderID
|
||||
return nil
|
||||
})
|
||||
require.NotNil(t, disposer)
|
||||
|
||||
err := bus.Emit(context.Background(), "order:created", OrderCreatedEvent{OrderID: "ord_123", Amount: 99.5})
|
||||
assert.NoError(t, err)
|
||||
assert.Equal(t, "ord_123", receivedOrder)
|
||||
|
||||
// Test unsubscribe via disposer
|
||||
err = disposer()
|
||||
assert.NoError(t, err)
|
||||
|
||||
receivedOrder = ""
|
||||
err = bus.Emit(context.Background(), "order:created", OrderCreatedEvent{OrderID: "ord_456", Amount: 100})
|
||||
assert.NoError(t, err)
|
||||
assert.Empty(t, receivedOrder, "handler should not be called after disposal")
|
||||
}
|
||||
|
||||
func TestEventBusHandlerSignatures(t *testing.T) {
|
||||
bus := core.NewEventBus()
|
||||
|
||||
var (
|
||||
calledWithCtxPayloadErr atomic.Bool
|
||||
calledWithCtxPayload atomic.Bool
|
||||
calledWithPayloadErr atomic.Bool
|
||||
calledWithPayload atomic.Bool
|
||||
calledWithCtxErr atomic.Bool
|
||||
calledWithCtx atomic.Bool
|
||||
calledWithNoArgsErr atomic.Bool
|
||||
calledWithNoArgs atomic.Bool
|
||||
)
|
||||
|
||||
bus.On("test:sig", func(ctx context.Context, e UserRegisteredEvent) error {
|
||||
calledWithCtxPayloadErr.Store(true)
|
||||
assert.Equal(t, "u_1", e.UserID)
|
||||
return nil
|
||||
})
|
||||
|
||||
bus.On("test:sig", func(ctx context.Context, e UserRegisteredEvent) {
|
||||
calledWithCtxPayload.Store(true)
|
||||
assert.Equal(t, "u_1", e.UserID)
|
||||
})
|
||||
|
||||
bus.On("test:sig", func(e UserRegisteredEvent) error {
|
||||
calledWithPayloadErr.Store(true)
|
||||
assert.Equal(t, "u_1", e.UserID)
|
||||
return nil
|
||||
})
|
||||
|
||||
bus.On("test:sig", func(e UserRegisteredEvent) {
|
||||
calledWithPayload.Store(true)
|
||||
assert.Equal(t, "u_1", e.UserID)
|
||||
})
|
||||
|
||||
bus.On("test:sig", func(ctx context.Context) error {
|
||||
calledWithCtxErr.Store(true)
|
||||
return nil
|
||||
})
|
||||
|
||||
bus.On("test:sig", func(ctx context.Context) {
|
||||
calledWithCtx.Store(true)
|
||||
})
|
||||
|
||||
bus.On("test:sig", func() error {
|
||||
calledWithNoArgsErr.Store(true)
|
||||
return nil
|
||||
})
|
||||
|
||||
bus.On("test:sig", func() {
|
||||
calledWithNoArgs.Store(true)
|
||||
})
|
||||
|
||||
err := bus.Emit(context.Background(), "test:sig", UserRegisteredEvent{UserID: "u_1", Username: "test"})
|
||||
assert.NoError(t, err)
|
||||
|
||||
assert.True(t, calledWithCtxPayloadErr.Load())
|
||||
assert.True(t, calledWithCtxPayload.Load())
|
||||
assert.True(t, calledWithPayloadErr.Load())
|
||||
assert.True(t, calledWithPayload.Load())
|
||||
assert.True(t, calledWithCtxErr.Load())
|
||||
assert.True(t, calledWithCtx.Load())
|
||||
assert.True(t, calledWithNoArgsErr.Load())
|
||||
assert.True(t, calledWithNoArgs.Load())
|
||||
}
|
||||
|
||||
func TestEventBusPointerAndValueConversion(t *testing.T) {
|
||||
bus := core.NewEventBus()
|
||||
|
||||
var (
|
||||
receivedFromValueToPtr atomic.Bool
|
||||
receivedFromPtrToValue atomic.Bool
|
||||
receivedFromNilPtr atomic.Bool
|
||||
)
|
||||
|
||||
// Handler expects pointer, payload emitted as value
|
||||
bus.On("test:ptr", func(ctx context.Context, e *UserRegisteredEvent) error {
|
||||
if e != nil && e.UserID == "u_ptr" {
|
||||
receivedFromValueToPtr.Store(true)
|
||||
}
|
||||
return nil
|
||||
})
|
||||
|
||||
err := bus.Emit(context.Background(), "test:ptr", UserRegisteredEvent{UserID: "u_ptr"})
|
||||
assert.NoError(t, err)
|
||||
assert.True(t, receivedFromValueToPtr.Load())
|
||||
|
||||
// Handler expects value, payload emitted as pointer
|
||||
bus.On("test:val", func(ctx context.Context, e UserRegisteredEvent) error {
|
||||
if e.UserID == "u_val" {
|
||||
receivedFromPtrToValue.Store(true)
|
||||
}
|
||||
return nil
|
||||
})
|
||||
|
||||
err = bus.Emit(context.Background(), "test:val", &UserRegisteredEvent{UserID: "u_val"})
|
||||
assert.NoError(t, err)
|
||||
assert.True(t, receivedFromPtrToValue.Load())
|
||||
|
||||
// Handler expects value, payload is nil pointer
|
||||
var nilEvent *UserRegisteredEvent
|
||||
bus.On("test:nil_ptr", func(ctx context.Context, e UserRegisteredEvent) error {
|
||||
assert.Equal(t, "", e.UserID)
|
||||
receivedFromNilPtr.Store(true)
|
||||
return nil
|
||||
})
|
||||
err = bus.Emit(context.Background(), "test:nil_ptr", nilEvent)
|
||||
assert.NoError(t, err)
|
||||
assert.True(t, receivedFromNilPtr.Load())
|
||||
|
||||
// Convertible type test (int to int64)
|
||||
var receivedConvert int64
|
||||
bus.On("test:conv", func(e int64) {
|
||||
receivedConvert = e
|
||||
})
|
||||
err = bus.Emit(context.Background(), "test:conv", int(42))
|
||||
assert.NoError(t, err)
|
||||
assert.Equal(t, int64(42), receivedConvert)
|
||||
}
|
||||
|
||||
func TestEventBusErrorCollectionAndPanicRecovery(t *testing.T) {
|
||||
bus := core.NewEventBus()
|
||||
|
||||
errHandler1 := errors.New("handler 1 failed")
|
||||
errHandler2 := errors.New("handler 2 failed")
|
||||
|
||||
bus.On("test:err", func() error {
|
||||
return errHandler1
|
||||
})
|
||||
|
||||
bus.On("test:err", func() {
|
||||
panic("something went horribly wrong")
|
||||
})
|
||||
|
||||
bus.On("test:err", func() error {
|
||||
return errHandler2
|
||||
})
|
||||
|
||||
err := bus.Emit(context.Background(), "test:err", nil)
|
||||
require.Error(t, err)
|
||||
assert.True(t, errors.Is(err, errHandler1) || errors.Is(err, errHandler2))
|
||||
assert.Contains(t, err.Error(), "handler 1 failed")
|
||||
assert.Contains(t, err.Error(), "handler 2 failed")
|
||||
assert.Contains(t, err.Error(), "panic")
|
||||
}
|
||||
|
||||
func TestEventBusInvalidHandlerPanics(t *testing.T) {
|
||||
bus := core.NewEventBus()
|
||||
|
||||
assert.Panics(t, func() {
|
||||
bus.On("test:invalid", nil)
|
||||
})
|
||||
|
||||
assert.Panics(t, func() {
|
||||
bus.On("test:invalid", "not-a-func")
|
||||
})
|
||||
|
||||
assert.Panics(t, func() {
|
||||
// More than 2 arguments
|
||||
bus.On("test:invalid", func(a, b, c string) {})
|
||||
})
|
||||
|
||||
assert.Panics(t, func() {
|
||||
// 2 args, but first is not context
|
||||
bus.On("test:invalid", func(a string, b int) {})
|
||||
})
|
||||
|
||||
assert.Panics(t, func() {
|
||||
// More than 1 return value
|
||||
bus.On("test:invalid", func() (int, error) { return 0, nil })
|
||||
})
|
||||
|
||||
assert.Panics(t, func() {
|
||||
// Return value is not error
|
||||
bus.On("test:invalid", func() int { return 0 })
|
||||
})
|
||||
}
|
||||
|
||||
func TestEventBusListenersCountAndDisposerIdempotence(t *testing.T) {
|
||||
bus := core.NewEventBus()
|
||||
|
||||
assert.Equal(t, 0, bus.Listeners("topic1"))
|
||||
|
||||
d1 := bus.On("topic1", func() {})
|
||||
d2 := bus.On("topic1", func() {})
|
||||
assert.Equal(t, 2, bus.Listeners("topic1"))
|
||||
|
||||
_ = d1()
|
||||
assert.Equal(t, 1, bus.Listeners("topic1"))
|
||||
|
||||
// Calling disposer again should be no-op
|
||||
_ = d1()
|
||||
assert.Equal(t, 1, bus.Listeners("topic1"))
|
||||
|
||||
_ = d2()
|
||||
assert.Equal(t, 0, bus.Listeners("topic1"))
|
||||
}
|
||||
|
||||
func TestEventBusConcurrentAccess(t *testing.T) {
|
||||
bus := core.NewEventBus()
|
||||
var wg sync.WaitGroup
|
||||
|
||||
var receivedCount atomic.Int64
|
||||
|
||||
// Concurrently subscribe and emit
|
||||
for i := 0; i < 50; i++ {
|
||||
wg.Add(1)
|
||||
go func(idx int) {
|
||||
defer wg.Done()
|
||||
topic := fmt.Sprintf("topic:%d", idx%5)
|
||||
disposer := bus.On(topic, func(ctx context.Context, e UserRegisteredEvent) error {
|
||||
receivedCount.Add(1)
|
||||
return nil
|
||||
})
|
||||
|
||||
// Emit some events
|
||||
_ = bus.Emit(context.Background(), topic, UserRegisteredEvent{UserID: fmt.Sprintf("u_%d", idx)})
|
||||
|
||||
// Randomly dispose
|
||||
if idx%2 == 0 {
|
||||
_ = disposer()
|
||||
}
|
||||
}(i)
|
||||
}
|
||||
|
||||
for i := 0; i < 50; i++ {
|
||||
wg.Add(1)
|
||||
go func(idx int) {
|
||||
defer wg.Done()
|
||||
topic := fmt.Sprintf("topic:%d", idx%5)
|
||||
_ = bus.Emit(context.Background(), topic, UserRegisteredEvent{UserID: fmt.Sprintf("u_%d", idx)})
|
||||
}(i)
|
||||
}
|
||||
|
||||
wg.Wait()
|
||||
assert.Greater(t, receivedCount.Load(), int64(0))
|
||||
}
|
||||
@@ -0,0 +1,280 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package extpoints_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
"testing/fstest"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/backend/core"
|
||||
"github.com/Rain-kl/Wavelet/backend/core/extpoints"
|
||||
)
|
||||
|
||||
func TestRouterExtension(t *testing.T) {
|
||||
r := extpoints.NewRouterRegistry()
|
||||
require.NotNil(t, r)
|
||||
|
||||
mGlobal := "global_middleware"
|
||||
r.Use(mGlobal)
|
||||
assert.Equal(t, []any{mGlobal}, r.Middlewares())
|
||||
|
||||
// Test root methods
|
||||
hRoot := "root_handler"
|
||||
r.GET("/", hRoot)
|
||||
r.POST("/root_post", hRoot)
|
||||
r.PUT("/root_put", hRoot)
|
||||
r.DELETE("/root_del", hRoot)
|
||||
r.PATCH("/root_patch", hRoot)
|
||||
r.HEAD("/root_head", hRoot)
|
||||
r.OPTIONS("/root_opt", hRoot)
|
||||
anyRootDefs := r.Any("/root_any", hRoot)
|
||||
assert.Len(t, anyRootDefs, 7)
|
||||
|
||||
// Group and Group.Use
|
||||
mAPI := "api_middleware"
|
||||
api := r.Group("/api/v1", mAPI)
|
||||
api.Use("api_extra_middleware")
|
||||
assert.Len(t, api.Middlewares(), 2)
|
||||
|
||||
hList := "list_orders_handler"
|
||||
hCreate := "create_order_handler"
|
||||
api.GET("/orders", hList)
|
||||
api.POST("/orders", hCreate)
|
||||
|
||||
mAdmin := "admin_middleware"
|
||||
admin := api.Group("admin", mAdmin)
|
||||
|
||||
hUserGet := "get_user_handler"
|
||||
hUserPut := "put_user_handler"
|
||||
hUserDel := "del_user_handler"
|
||||
hUserPatch := "patch_user_handler"
|
||||
hUserHead := "head_user_handler"
|
||||
hUserOptions := "options_user_handler"
|
||||
admin.GET("/users/:id", hUserGet)
|
||||
admin.PUT("/users/:id", hUserPut)
|
||||
admin.DELETE("/users/:id", hUserDel)
|
||||
admin.PATCH("/users/:id", hUserPatch)
|
||||
admin.HEAD("/users/:id", hUserHead)
|
||||
admin.OPTIONS("/users/:id", hUserOptions)
|
||||
|
||||
hCustom := "custom_handler"
|
||||
admin.Handle("CUSTOM", "/custom", hCustom)
|
||||
|
||||
hAny := "any_handler"
|
||||
anyRoutes := admin.Any("/all", hAny)
|
||||
assert.NotEmpty(t, anyRoutes)
|
||||
|
||||
// Group.Routes() returns root routes
|
||||
assert.Equal(t, r.Routes(), admin.Routes())
|
||||
|
||||
routes := r.Routes()
|
||||
|
||||
// Verify route paths and middlewares
|
||||
var foundOrderGet bool
|
||||
var foundUserPut bool
|
||||
for _, route := range routes {
|
||||
if route.Method == "GET" && route.Path == "/api/v1/orders" {
|
||||
foundOrderGet = true
|
||||
assert.Equal(t, []any{mGlobal, mAPI, "api_extra_middleware"}, route.Middlewares)
|
||||
assert.Equal(t, []any{hList}, route.Handlers)
|
||||
}
|
||||
if route.Method == "PUT" && route.Path == "/api/v1/admin/users/:id" {
|
||||
foundUserPut = true
|
||||
assert.Equal(t, []any{mGlobal, mAPI, "api_extra_middleware", mAdmin}, route.Middlewares)
|
||||
assert.Equal(t, []any{hUserPut}, route.Handlers)
|
||||
}
|
||||
}
|
||||
assert.True(t, foundOrderGet)
|
||||
assert.True(t, foundUserPut)
|
||||
}
|
||||
|
||||
func TestMigrationExtension(t *testing.T) {
|
||||
m := extpoints.NewMigrationRegistry()
|
||||
require.NotNil(t, m)
|
||||
|
||||
fs1 := fstest.MapFS{
|
||||
"migrations/001_init.sql": &fstest.MapFile{Data: []byte("CREATE TABLE t1(id int);")},
|
||||
}
|
||||
fs2 := fstest.MapFS{
|
||||
"custom/001_order.sql": &fstest.MapFile{Data: []byte("CREATE TABLE t2(id int);")},
|
||||
}
|
||||
|
||||
m.Register("auth", fs1)
|
||||
m.Register("order", fs2, "custom")
|
||||
|
||||
// Update existing entry
|
||||
fs1Updated := fstest.MapFS{
|
||||
"migrations/002_update.sql": &fstest.MapFile{Data: []byte("ALTER TABLE t1 ADD col int;")},
|
||||
}
|
||||
m.Register("auth", fs1Updated, "")
|
||||
|
||||
entries := m.Entries()
|
||||
require.Len(t, entries, 2)
|
||||
assert.Equal(t, "auth", entries[0].PluginID)
|
||||
assert.Equal(t, "migrations", entries[0].Dir)
|
||||
assert.Equal(t, "order", entries[1].PluginID)
|
||||
assert.Equal(t, "custom", entries[1].Dir)
|
||||
|
||||
authEntry, ok := m.Get("auth")
|
||||
assert.True(t, ok)
|
||||
assert.Equal(t, "auth", authEntry.PluginID)
|
||||
|
||||
_, ok = m.Get("non_existent")
|
||||
assert.False(t, ok)
|
||||
}
|
||||
|
||||
func TestTaskExtension(t *testing.T) {
|
||||
tr := extpoints.NewTaskRegistry()
|
||||
require.NotNil(t, tr)
|
||||
|
||||
handler := func(ctx context.Context, payload []byte) error { return nil }
|
||||
|
||||
tr.Register("order:cancel_timeout", handler,
|
||||
extpoints.WithTaskConcurrency(5),
|
||||
extpoints.WithTaskRetry(3),
|
||||
extpoints.WithTaskTimeout(10*time.Second),
|
||||
extpoints.WithTaskMetadata("queue", "critical"),
|
||||
nil, // test nil option
|
||||
)
|
||||
|
||||
// Re-register to test update
|
||||
tr.Register("order:cancel_timeout", handler,
|
||||
extpoints.WithTaskConcurrency(10),
|
||||
extpoints.WithTaskMetadata("queue", "high"),
|
||||
)
|
||||
|
||||
tasks := tr.Tasks()
|
||||
require.Len(t, tasks, 1)
|
||||
assert.Equal(t, "order:cancel_timeout", tasks[0].Pattern)
|
||||
assert.Equal(t, 10, tasks[0].Concurrency)
|
||||
assert.Equal(t, "high", tasks[0].Metadata["queue"])
|
||||
|
||||
task, ok := tr.Get("order:cancel_timeout")
|
||||
assert.True(t, ok)
|
||||
assert.Equal(t, "order:cancel_timeout", task.Pattern)
|
||||
|
||||
_, ok = tr.Get("unknown")
|
||||
assert.False(t, ok)
|
||||
}
|
||||
|
||||
func TestScheduleExtension(t *testing.T) {
|
||||
sr := extpoints.NewScheduleRegistry()
|
||||
require.NotNil(t, sr)
|
||||
|
||||
type ReportPayload struct {
|
||||
Type string `json:"type"`
|
||||
}
|
||||
|
||||
sr.RegisterCron("0 2 * * *", "report:daily_summary", ReportPayload{Type: "daily"})
|
||||
sr.Register("@every 1h", "cleanup:expired_sessions", nil,
|
||||
extpoints.WithScheduleOption("retry", 2),
|
||||
nil, // test nil option
|
||||
)
|
||||
|
||||
// Re-register to test update
|
||||
sr.RegisterCron("0 3 * * *", "report:daily_summary", ReportPayload{Type: "all"})
|
||||
|
||||
schedules := sr.Schedules()
|
||||
require.Len(t, schedules, 2)
|
||||
|
||||
assert.Equal(t, "0 3 * * *", schedules[0].Spec)
|
||||
assert.Equal(t, "report:daily_summary", schedules[0].TaskType)
|
||||
assert.Equal(t, ReportPayload{Type: "all"}, schedules[0].Payload)
|
||||
|
||||
assert.Equal(t, "@every 1h", schedules[1].Spec)
|
||||
assert.Equal(t, "cleanup:expired_sessions", schedules[1].TaskType)
|
||||
assert.Equal(t, 2, schedules[1].Options["retry"])
|
||||
|
||||
sched, ok := sr.Get("report:daily_summary")
|
||||
assert.True(t, ok)
|
||||
assert.Equal(t, "0 3 * * *", sched.Spec)
|
||||
|
||||
_, ok = sr.Get("unknown")
|
||||
assert.False(t, ok)
|
||||
}
|
||||
|
||||
func TestSettingExtension(t *testing.T) {
|
||||
sr := extpoints.NewSettingRegistry()
|
||||
require.NotNil(t, sr)
|
||||
|
||||
assert.Panics(t, func() {
|
||||
sr.Register(extpoints.SettingSchema{}) // empty key panics
|
||||
})
|
||||
|
||||
sr.Register(extpoints.SettingSchema{
|
||||
Key: "order.auto_cancel_mins",
|
||||
Default: 15,
|
||||
Description: "Order auto cancellation timeout in minutes",
|
||||
Category: "order",
|
||||
Public: true,
|
||||
})
|
||||
|
||||
// Re-register to test update
|
||||
sr.Register(extpoints.SettingSchema{
|
||||
Key: "order.auto_cancel_mins",
|
||||
Default: 30,
|
||||
Description: "Updated timeout",
|
||||
})
|
||||
|
||||
sr.Register(extpoints.SettingSchema{
|
||||
Key: "auth.jwt_secret",
|
||||
Default: "default-secret",
|
||||
Description: "JWT secret key",
|
||||
Category: "auth",
|
||||
ReadOnly: true,
|
||||
})
|
||||
|
||||
schemas := sr.Schemas()
|
||||
require.Len(t, schemas, 2)
|
||||
|
||||
schema, ok := sr.Get("order.auto_cancel_mins")
|
||||
assert.True(t, ok)
|
||||
assert.Equal(t, 30, schema.Default)
|
||||
|
||||
_, ok = sr.Get("unknown")
|
||||
assert.False(t, ok)
|
||||
}
|
||||
|
||||
func TestContextExtensionPointsIntegration(t *testing.T) {
|
||||
ctx := core.NewContext(context.Background())
|
||||
|
||||
require.NotNil(t, ctx.Events())
|
||||
require.NotNil(t, ctx.Router())
|
||||
require.NotNil(t, ctx.Migrations())
|
||||
require.NotNil(t, ctx.Tasks())
|
||||
require.NotNil(t, ctx.Task())
|
||||
require.NotNil(t, ctx.Schedules())
|
||||
require.NotNil(t, ctx.Schedule())
|
||||
require.NotNil(t, ctx.Settings())
|
||||
require.NotNil(t, ctx.Setting())
|
||||
|
||||
// Register from child context and verify shared application registry
|
||||
child := ctx.Fork()
|
||||
child.Router().GET("/ping", "pong_handler")
|
||||
child.Task().Register("sample:task", "handler")
|
||||
child.Schedule().RegisterCron("@hourly", "sample:cron", nil)
|
||||
child.Settings().Register(extpoints.SettingSchema{
|
||||
Key: "app.name",
|
||||
Default: "Wavelet",
|
||||
})
|
||||
|
||||
assert.Len(t, ctx.Router().Routes(), 1)
|
||||
assert.Len(t, ctx.Tasks().Tasks(), 1)
|
||||
assert.Len(t, ctx.Schedules().Schedules(), 1)
|
||||
assert.Len(t, ctx.Settings().Schemas(), 1)
|
||||
|
||||
// Child and root events
|
||||
var eventReceived bool
|
||||
child.Events().On("app:ready", func() {
|
||||
eventReceived = true
|
||||
})
|
||||
err := ctx.Events().Emit(context.Background(), "app:ready", nil)
|
||||
assert.NoError(t, err)
|
||||
assert.True(t, eventReceived)
|
||||
}
|
||||
@@ -0,0 +1,87 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package extpoints defines extension points for router, migrations, tasks, schedules, and settings.
|
||||
package extpoints
|
||||
|
||||
import (
|
||||
"io/fs"
|
||||
"sync"
|
||||
)
|
||||
|
||||
// MigrationEntry contains the migration filesystem and configuration for a plugin.
|
||||
type MigrationEntry struct {
|
||||
PluginID string
|
||||
FS fs.FS
|
||||
Dir string
|
||||
}
|
||||
|
||||
// MigrationExtension defines the interface for registering and querying plugin migrations.
|
||||
type MigrationExtension interface {
|
||||
Register(pluginID string, fsys fs.FS, dir ...string)
|
||||
Entries() []MigrationEntry
|
||||
Get(pluginID string) (MigrationEntry, bool)
|
||||
}
|
||||
|
||||
// MigrationRegistry collects and stores migration entries from plugins.
|
||||
type MigrationRegistry struct {
|
||||
mu sync.RWMutex
|
||||
entries []MigrationEntry
|
||||
lookup map[string]MigrationEntry
|
||||
}
|
||||
|
||||
// NewMigrationRegistry creates a new migration registry.
|
||||
func NewMigrationRegistry() *MigrationRegistry {
|
||||
return &MigrationRegistry{
|
||||
lookup: make(map[string]MigrationEntry),
|
||||
}
|
||||
}
|
||||
|
||||
// Register adds a migration entry for a plugin.
|
||||
// If dir is not specified, it defaults to "migrations".
|
||||
func (m *MigrationRegistry) Register(pluginID string, fsys fs.FS, dir ...string) {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
|
||||
migrationDir := "migrations"
|
||||
if len(dir) > 0 && dir[0] != "" {
|
||||
migrationDir = dir[0]
|
||||
}
|
||||
|
||||
entry := MigrationEntry{
|
||||
PluginID: pluginID,
|
||||
FS: fsys,
|
||||
Dir: migrationDir,
|
||||
}
|
||||
|
||||
// If entry already exists, update in-place; otherwise append
|
||||
if _, exists := m.lookup[pluginID]; exists {
|
||||
for i, e := range m.entries {
|
||||
if e.PluginID == pluginID {
|
||||
m.entries[i] = entry
|
||||
break
|
||||
}
|
||||
}
|
||||
} else {
|
||||
m.entries = append(m.entries, entry)
|
||||
}
|
||||
|
||||
m.lookup[pluginID] = entry
|
||||
}
|
||||
|
||||
// Entries returns a copy of all registered migration entries in registration order.
|
||||
func (m *MigrationRegistry) Entries() []MigrationEntry {
|
||||
m.mu.RLock()
|
||||
defer m.mu.RUnlock()
|
||||
res := make([]MigrationEntry, len(m.entries))
|
||||
copy(res, m.entries)
|
||||
return res
|
||||
}
|
||||
|
||||
// Get retrieves the migration entry for a specific plugin ID.
|
||||
func (m *MigrationRegistry) Get(pluginID string) (MigrationEntry, bool) {
|
||||
m.mu.RLock()
|
||||
defer m.mu.RUnlock()
|
||||
e, ok := m.lookup[pluginID]
|
||||
return e, ok
|
||||
}
|
||||
@@ -0,0 +1,269 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package extpoints
|
||||
|
||||
import (
|
||||
"strings"
|
||||
"sync"
|
||||
)
|
||||
|
||||
// RouteDefinition holds the metadata and handler list for a single HTTP route.
|
||||
type RouteDefinition struct {
|
||||
Method string
|
||||
Path string
|
||||
Handlers []any
|
||||
Middlewares []any
|
||||
}
|
||||
|
||||
// RouterExtension defines the interface for registering routes and middlewares.
|
||||
type RouterExtension interface {
|
||||
Use(middlewares ...any)
|
||||
Group(prefix string, middlewares ...any) RouterExtension
|
||||
Handle(method string, path string, handlers ...any) RouteDefinition
|
||||
GET(path string, handlers ...any) RouteDefinition
|
||||
POST(path string, handlers ...any) RouteDefinition
|
||||
PUT(path string, handlers ...any) RouteDefinition
|
||||
DELETE(path string, handlers ...any) RouteDefinition
|
||||
PATCH(path string, handlers ...any) RouteDefinition
|
||||
HEAD(path string, handlers ...any) RouteDefinition
|
||||
OPTIONS(path string, handlers ...any) RouteDefinition
|
||||
Any(path string, handlers ...any) []RouteDefinition
|
||||
Routes() []RouteDefinition
|
||||
Middlewares() []any
|
||||
}
|
||||
|
||||
// RouterRegistry implements RouterExtension as the root route and middleware collector.
|
||||
type RouterRegistry struct {
|
||||
mu sync.RWMutex
|
||||
routes []RouteDefinition
|
||||
middlewares []any
|
||||
}
|
||||
|
||||
// NewRouterRegistry creates a new root router collector.
|
||||
func NewRouterRegistry() *RouterRegistry {
|
||||
return &RouterRegistry{}
|
||||
}
|
||||
|
||||
// Use registers global middlewares to the router.
|
||||
func (r *RouterRegistry) Use(middlewares ...any) {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
r.middlewares = append(r.middlewares, middlewares...)
|
||||
}
|
||||
|
||||
// Middlewares returns a copy of registered root middlewares.
|
||||
func (r *RouterRegistry) Middlewares() []any {
|
||||
r.mu.RLock()
|
||||
defer r.mu.RUnlock()
|
||||
res := make([]any, len(r.middlewares))
|
||||
copy(res, r.middlewares)
|
||||
return res
|
||||
}
|
||||
|
||||
// Group creates a new RouteGroup under the router.
|
||||
func (r *RouterRegistry) Group(prefix string, middlewares ...any) RouterExtension {
|
||||
return &RouterGroup{
|
||||
registry: r,
|
||||
prefix: cleanPath(prefix),
|
||||
middlewares: middlewares,
|
||||
}
|
||||
}
|
||||
|
||||
// Handle registers a route with a custom HTTP method and handlers.
|
||||
func (r *RouterRegistry) Handle(method, path string, handlers ...any) RouteDefinition {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
|
||||
rd := RouteDefinition{
|
||||
Method: strings.ToUpper(method),
|
||||
Path: cleanPath(path),
|
||||
Handlers: handlers,
|
||||
Middlewares: append([]any(nil), r.middlewares...),
|
||||
}
|
||||
r.routes = append(r.routes, rd)
|
||||
return rd
|
||||
}
|
||||
|
||||
// GET registers a GET route.
|
||||
func (r *RouterRegistry) GET(path string, handlers ...any) RouteDefinition {
|
||||
return r.Handle("GET", path, handlers...)
|
||||
}
|
||||
|
||||
// POST registers a POST route.
|
||||
func (r *RouterRegistry) POST(path string, handlers ...any) RouteDefinition {
|
||||
return r.Handle("POST", path, handlers...)
|
||||
}
|
||||
|
||||
// PUT registers a PUT route.
|
||||
func (r *RouterRegistry) PUT(path string, handlers ...any) RouteDefinition {
|
||||
return r.Handle("PUT", path, handlers...)
|
||||
}
|
||||
|
||||
// DELETE registers a DELETE route.
|
||||
func (r *RouterRegistry) DELETE(path string, handlers ...any) RouteDefinition {
|
||||
return r.Handle("DELETE", path, handlers...)
|
||||
}
|
||||
|
||||
// PATCH registers a PATCH route.
|
||||
func (r *RouterRegistry) PATCH(path string, handlers ...any) RouteDefinition {
|
||||
return r.Handle("PATCH", path, handlers...)
|
||||
}
|
||||
|
||||
// HEAD registers a HEAD route.
|
||||
func (r *RouterRegistry) HEAD(path string, handlers ...any) RouteDefinition {
|
||||
return r.Handle("HEAD", path, handlers...)
|
||||
}
|
||||
|
||||
// OPTIONS registers an OPTIONS route.
|
||||
func (r *RouterRegistry) OPTIONS(path string, handlers ...any) RouteDefinition {
|
||||
return r.Handle("OPTIONS", path, handlers...)
|
||||
}
|
||||
|
||||
// Any registers a route for standard HTTP methods.
|
||||
func (r *RouterRegistry) Any(path string, handlers ...any) []RouteDefinition {
|
||||
methods := []string{"GET", "POST", "PUT", "DELETE", "PATCH", "HEAD", "OPTIONS"}
|
||||
defs := make([]RouteDefinition, 0, len(methods))
|
||||
for _, m := range methods {
|
||||
defs = append(defs, r.Handle(m, path, handlers...))
|
||||
}
|
||||
return defs
|
||||
}
|
||||
|
||||
// Routes returns a copy of all collected RouteDefinitions.
|
||||
func (r *RouterRegistry) Routes() []RouteDefinition {
|
||||
r.mu.RLock()
|
||||
defer r.mu.RUnlock()
|
||||
res := make([]RouteDefinition, len(r.routes))
|
||||
copy(res, r.routes)
|
||||
return res
|
||||
}
|
||||
|
||||
// RouterGroup represents a scoped route group with a path prefix and group-level middlewares.
|
||||
type RouterGroup struct {
|
||||
registry *RouterRegistry
|
||||
prefix string
|
||||
middlewares []any
|
||||
}
|
||||
|
||||
// Use adds middlewares to this group.
|
||||
func (g *RouterGroup) Use(middlewares ...any) {
|
||||
g.middlewares = append(g.middlewares, middlewares...)
|
||||
}
|
||||
|
||||
// Group creates a nested RouteGroup.
|
||||
func (g *RouterGroup) Group(prefix string, middlewares ...any) RouterExtension {
|
||||
combinedPrefix := joinPaths(g.prefix, prefix)
|
||||
combinedMiddlewares := make([]any, 0, len(g.middlewares)+len(middlewares))
|
||||
combinedMiddlewares = append(combinedMiddlewares, g.middlewares...)
|
||||
combinedMiddlewares = append(combinedMiddlewares, middlewares...)
|
||||
|
||||
return &RouterGroup{
|
||||
registry: g.registry,
|
||||
prefix: combinedPrefix,
|
||||
middlewares: combinedMiddlewares,
|
||||
}
|
||||
}
|
||||
|
||||
// Handle registers a route under this group.
|
||||
func (g *RouterGroup) Handle(method, path string, handlers ...any) RouteDefinition {
|
||||
g.registry.mu.Lock()
|
||||
defer g.registry.mu.Unlock()
|
||||
|
||||
fullPath := joinPaths(g.prefix, path)
|
||||
|
||||
allMiddlewares := make([]any, 0, len(g.registry.middlewares)+len(g.middlewares))
|
||||
allMiddlewares = append(allMiddlewares, g.registry.middlewares...)
|
||||
allMiddlewares = append(allMiddlewares, g.middlewares...)
|
||||
|
||||
rd := RouteDefinition{
|
||||
Method: strings.ToUpper(method),
|
||||
Path: fullPath,
|
||||
Handlers: handlers,
|
||||
Middlewares: allMiddlewares,
|
||||
}
|
||||
g.registry.routes = append(g.registry.routes, rd)
|
||||
return rd
|
||||
}
|
||||
|
||||
// GET registers a GET route in this group.
|
||||
func (g *RouterGroup) GET(path string, handlers ...any) RouteDefinition {
|
||||
return g.Handle("GET", path, handlers...)
|
||||
}
|
||||
|
||||
// POST registers a POST route in this group.
|
||||
func (g *RouterGroup) POST(path string, handlers ...any) RouteDefinition {
|
||||
return g.Handle("POST", path, handlers...)
|
||||
}
|
||||
|
||||
// PUT registers a PUT route in this group.
|
||||
func (g *RouterGroup) PUT(path string, handlers ...any) RouteDefinition {
|
||||
return g.Handle("PUT", path, handlers...)
|
||||
}
|
||||
|
||||
// DELETE registers a DELETE route in this group.
|
||||
func (g *RouterGroup) DELETE(path string, handlers ...any) RouteDefinition {
|
||||
return g.Handle("DELETE", path, handlers...)
|
||||
}
|
||||
|
||||
// PATCH registers a PATCH route in this group.
|
||||
func (g *RouterGroup) PATCH(path string, handlers ...any) RouteDefinition {
|
||||
return g.Handle("PATCH", path, handlers...)
|
||||
}
|
||||
|
||||
// HEAD registers a HEAD route in this group.
|
||||
func (g *RouterGroup) HEAD(path string, handlers ...any) RouteDefinition {
|
||||
return g.Handle("HEAD", path, handlers...)
|
||||
}
|
||||
|
||||
// OPTIONS registers an OPTIONS route in this group.
|
||||
func (g *RouterGroup) OPTIONS(path string, handlers ...any) RouteDefinition {
|
||||
return g.Handle("OPTIONS", path, handlers...)
|
||||
}
|
||||
|
||||
// Any registers a route in this group for standard HTTP methods.
|
||||
func (g *RouterGroup) Any(path string, handlers ...any) []RouteDefinition {
|
||||
methods := []string{"GET", "POST", "PUT", "DELETE", "PATCH", "HEAD", "OPTIONS"}
|
||||
defs := make([]RouteDefinition, 0, len(methods))
|
||||
for _, m := range methods {
|
||||
defs = append(defs, g.Handle(m, path, handlers...))
|
||||
}
|
||||
return defs
|
||||
}
|
||||
|
||||
// Routes returns all routes from the parent registry.
|
||||
func (g *RouterGroup) Routes() []RouteDefinition {
|
||||
return g.registry.Routes()
|
||||
}
|
||||
|
||||
// Middlewares returns a copy of the group's middlewares.
|
||||
func (g *RouterGroup) Middlewares() []any {
|
||||
res := make([]any, len(g.middlewares))
|
||||
copy(res, g.middlewares)
|
||||
return res
|
||||
}
|
||||
|
||||
func cleanPath(p string) string {
|
||||
if p == "" {
|
||||
return "/"
|
||||
}
|
||||
if !strings.HasPrefix(p, "/") {
|
||||
p = "/" + p
|
||||
}
|
||||
if len(p) > 1 && strings.HasSuffix(p, "/") {
|
||||
p = strings.TrimSuffix(p, "/")
|
||||
}
|
||||
return p
|
||||
}
|
||||
|
||||
func joinPaths(base, relative string) string {
|
||||
if base == "" || base == "/" {
|
||||
return cleanPath(relative)
|
||||
}
|
||||
if relative == "" || relative == "/" {
|
||||
return cleanPath(base)
|
||||
}
|
||||
base = strings.TrimSuffix(base, "/")
|
||||
relative = strings.TrimPrefix(relative, "/")
|
||||
return cleanPath(base + "/" + relative)
|
||||
}
|
||||
@@ -0,0 +1,103 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package extpoints
|
||||
|
||||
import "sync"
|
||||
|
||||
// ScheduleDefinition holds the configuration for a scheduled/cron task.
|
||||
type ScheduleDefinition struct {
|
||||
Spec string
|
||||
TaskType string
|
||||
Payload any
|
||||
Options map[string]any
|
||||
}
|
||||
|
||||
// ScheduleOption configures a ScheduleDefinition.
|
||||
type ScheduleOption func(*ScheduleDefinition)
|
||||
|
||||
// WithScheduleOption adds a custom option to the schedule definition.
|
||||
func WithScheduleOption(key string, val any) ScheduleOption {
|
||||
return func(sd *ScheduleDefinition) {
|
||||
if sd.Options == nil {
|
||||
sd.Options = make(map[string]any)
|
||||
}
|
||||
sd.Options[key] = val
|
||||
}
|
||||
}
|
||||
|
||||
// ScheduleExtension defines the interface for registering and querying cron/scheduled tasks.
|
||||
type ScheduleExtension interface {
|
||||
Register(spec string, taskType string, payload any, opts ...ScheduleOption)
|
||||
RegisterCron(spec string, taskType string, payload any, opts ...ScheduleOption)
|
||||
Schedules() []ScheduleDefinition
|
||||
Get(taskType string) (ScheduleDefinition, bool)
|
||||
}
|
||||
|
||||
// ScheduleRegistry collects and manages schedule registrations.
|
||||
type ScheduleRegistry struct {
|
||||
mu sync.RWMutex
|
||||
schedules []ScheduleDefinition
|
||||
lookup map[string]ScheduleDefinition
|
||||
}
|
||||
|
||||
// NewScheduleRegistry creates a new schedule registry.
|
||||
func NewScheduleRegistry() *ScheduleRegistry {
|
||||
return &ScheduleRegistry{
|
||||
lookup: make(map[string]ScheduleDefinition),
|
||||
}
|
||||
}
|
||||
|
||||
// Register adds a schedule definition.
|
||||
func (s *ScheduleRegistry) Register(spec string, taskType string, payload any, opts ...ScheduleOption) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
|
||||
sd := ScheduleDefinition{
|
||||
Spec: spec,
|
||||
TaskType: taskType,
|
||||
Payload: payload,
|
||||
Options: make(map[string]any),
|
||||
}
|
||||
|
||||
for _, opt := range opts {
|
||||
if opt != nil {
|
||||
opt(&sd)
|
||||
}
|
||||
}
|
||||
|
||||
if _, exists := s.lookup[taskType]; exists {
|
||||
for i, item := range s.schedules {
|
||||
if item.TaskType == taskType {
|
||||
s.schedules[i] = sd
|
||||
break
|
||||
}
|
||||
}
|
||||
} else {
|
||||
s.schedules = append(s.schedules, sd)
|
||||
}
|
||||
|
||||
s.lookup[taskType] = sd
|
||||
}
|
||||
|
||||
// RegisterCron is an alias for Register.
|
||||
func (s *ScheduleRegistry) RegisterCron(spec string, taskType string, payload any, opts ...ScheduleOption) {
|
||||
s.Register(spec, taskType, payload, opts...)
|
||||
}
|
||||
|
||||
// Schedules returns a copy of all registered ScheduleDefinitions.
|
||||
func (s *ScheduleRegistry) Schedules() []ScheduleDefinition {
|
||||
s.mu.RLock()
|
||||
defer s.mu.RUnlock()
|
||||
res := make([]ScheduleDefinition, len(s.schedules))
|
||||
copy(res, s.schedules)
|
||||
return res
|
||||
}
|
||||
|
||||
// Get retrieves a schedule definition by its task type.
|
||||
func (s *ScheduleRegistry) Get(taskType string) (ScheduleDefinition, bool) {
|
||||
s.mu.RLock()
|
||||
defer s.mu.RUnlock()
|
||||
sd, ok := s.lookup[taskType]
|
||||
return sd, ok
|
||||
}
|
||||
@@ -0,0 +1,80 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package extpoints
|
||||
|
||||
import "sync"
|
||||
|
||||
// SettingSchema defines the configuration schema and metadata for a system or plugin setting.
|
||||
type SettingSchema struct {
|
||||
Key string `json:"key"`
|
||||
Default any `json:"default"`
|
||||
Description string `json:"description"`
|
||||
Type string `json:"type,omitempty"`
|
||||
ReadOnly bool `json:"read_only,omitempty"`
|
||||
Public bool `json:"public,omitempty"`
|
||||
Category string `json:"category,omitempty"`
|
||||
Validation string `json:"validation,omitempty"`
|
||||
}
|
||||
|
||||
// SettingExtension defines the interface for registering and querying setting configuration schemas.
|
||||
type SettingExtension interface {
|
||||
Register(schema SettingSchema)
|
||||
Schemas() []SettingSchema
|
||||
Get(key string) (SettingSchema, bool)
|
||||
}
|
||||
|
||||
// SettingRegistry collects and manages setting configuration schemas.
|
||||
type SettingRegistry struct {
|
||||
mu sync.RWMutex
|
||||
schemas []SettingSchema
|
||||
lookup map[string]SettingSchema
|
||||
}
|
||||
|
||||
// NewSettingRegistry creates a new setting schema registry.
|
||||
func NewSettingRegistry() *SettingRegistry {
|
||||
return &SettingRegistry{
|
||||
lookup: make(map[string]SettingSchema),
|
||||
}
|
||||
}
|
||||
|
||||
// Register registers a SettingSchema into the registry.
|
||||
// Panics if the schema Key is empty.
|
||||
func (s *SettingRegistry) Register(schema SettingSchema) {
|
||||
if schema.Key == "" {
|
||||
panic("core/extpoints: setting schema key cannot be empty")
|
||||
}
|
||||
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
|
||||
if _, exists := s.lookup[schema.Key]; exists {
|
||||
for i, item := range s.schemas {
|
||||
if item.Key == schema.Key {
|
||||
s.schemas[i] = schema
|
||||
break
|
||||
}
|
||||
}
|
||||
} else {
|
||||
s.schemas = append(s.schemas, schema)
|
||||
}
|
||||
|
||||
s.lookup[schema.Key] = schema
|
||||
}
|
||||
|
||||
// Schemas returns a copy of all registered SettingSchemas.
|
||||
func (s *SettingRegistry) Schemas() []SettingSchema {
|
||||
s.mu.RLock()
|
||||
defer s.mu.RUnlock()
|
||||
res := make([]SettingSchema, len(s.schemas))
|
||||
copy(res, s.schemas)
|
||||
return res
|
||||
}
|
||||
|
||||
// Get retrieves a SettingSchema by its key.
|
||||
func (s *SettingRegistry) Get(key string) (SettingSchema, bool) {
|
||||
s.mu.RLock()
|
||||
defer s.mu.RUnlock()
|
||||
schema, ok := s.lookup[key]
|
||||
return schema, ok
|
||||
}
|
||||
@@ -0,0 +1,122 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package extpoints
|
||||
|
||||
import (
|
||||
"sync"
|
||||
"time"
|
||||
)
|
||||
|
||||
// TaskDefinition holds the definition and runtime options for an asynchronous background task.
|
||||
type TaskDefinition struct {
|
||||
Pattern string
|
||||
Handler any
|
||||
Concurrency int
|
||||
Retry int
|
||||
Timeout time.Duration
|
||||
Metadata map[string]any
|
||||
}
|
||||
|
||||
// TaskOption configures a TaskDefinition.
|
||||
type TaskOption func(*TaskDefinition)
|
||||
|
||||
// WithTaskConcurrency sets the concurrency limit for the task.
|
||||
func WithTaskConcurrency(concurrency int) TaskOption {
|
||||
return func(td *TaskDefinition) {
|
||||
td.Concurrency = concurrency
|
||||
}
|
||||
}
|
||||
|
||||
// WithTaskRetry sets the maximum retry count for the task.
|
||||
func WithTaskRetry(retry int) TaskOption {
|
||||
return func(td *TaskDefinition) {
|
||||
td.Retry = retry
|
||||
}
|
||||
}
|
||||
|
||||
// WithTaskTimeout sets the execution timeout for the task.
|
||||
func WithTaskTimeout(timeout time.Duration) TaskOption {
|
||||
return func(td *TaskDefinition) {
|
||||
td.Timeout = timeout
|
||||
}
|
||||
}
|
||||
|
||||
// WithTaskMetadata adds a key-value pair to the task metadata.
|
||||
func WithTaskMetadata(key string, val any) TaskOption {
|
||||
return func(td *TaskDefinition) {
|
||||
if td.Metadata == nil {
|
||||
td.Metadata = make(map[string]any)
|
||||
}
|
||||
td.Metadata[key] = val
|
||||
}
|
||||
}
|
||||
|
||||
// TaskExtension defines the interface for registering and querying background task handlers.
|
||||
type TaskExtension interface {
|
||||
Register(pattern string, handler any, opts ...TaskOption)
|
||||
Tasks() []TaskDefinition
|
||||
Get(pattern string) (TaskDefinition, bool)
|
||||
}
|
||||
|
||||
// TaskRegistry collects and manages task registrations.
|
||||
type TaskRegistry struct {
|
||||
mu sync.RWMutex
|
||||
tasks []TaskDefinition
|
||||
lookup map[string]TaskDefinition
|
||||
}
|
||||
|
||||
// NewTaskRegistry creates a new task registry.
|
||||
func NewTaskRegistry() *TaskRegistry {
|
||||
return &TaskRegistry{
|
||||
lookup: make(map[string]TaskDefinition),
|
||||
}
|
||||
}
|
||||
|
||||
// Register registers a task pattern and its handler with optional configuration.
|
||||
func (t *TaskRegistry) Register(pattern string, handler any, opts ...TaskOption) {
|
||||
t.mu.Lock()
|
||||
defer t.mu.Unlock()
|
||||
|
||||
td := TaskDefinition{
|
||||
Pattern: pattern,
|
||||
Handler: handler,
|
||||
Metadata: make(map[string]any),
|
||||
}
|
||||
|
||||
for _, opt := range opts {
|
||||
if opt != nil {
|
||||
opt(&td)
|
||||
}
|
||||
}
|
||||
|
||||
if _, exists := t.lookup[pattern]; exists {
|
||||
for i, item := range t.tasks {
|
||||
if item.Pattern == pattern {
|
||||
t.tasks[i] = td
|
||||
break
|
||||
}
|
||||
}
|
||||
} else {
|
||||
t.tasks = append(t.tasks, td)
|
||||
}
|
||||
|
||||
t.lookup[pattern] = td
|
||||
}
|
||||
|
||||
// Tasks returns a copy of all registered TaskDefinitions.
|
||||
func (t *TaskRegistry) Tasks() []TaskDefinition {
|
||||
t.mu.RLock()
|
||||
defer t.mu.RUnlock()
|
||||
res := make([]TaskDefinition, len(t.tasks))
|
||||
copy(res, t.tasks)
|
||||
return res
|
||||
}
|
||||
|
||||
// Get retrieves a task definition by its pattern.
|
||||
func (t *TaskRegistry) Get(pattern string) (TaskDefinition, bool) {
|
||||
t.mu.RLock()
|
||||
defer t.mu.RUnlock()
|
||||
td, ok := t.lookup[pattern]
|
||||
return td, ok
|
||||
}
|
||||
@@ -0,0 +1,38 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package core
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// Manifest defines the metadata and dependency declarations for a plugin.
|
||||
type Manifest struct {
|
||||
// Name is the unique identifier for the plugin (e.g. "auth", "user", "order").
|
||||
Name string `json:"name" yaml:"name"`
|
||||
|
||||
// Version is the semantic version string of the plugin (e.g. "1.0.0").
|
||||
Version string `json:"version,omitempty" yaml:"version,omitempty"`
|
||||
|
||||
// Description gives a brief summary of the plugin capabilities.
|
||||
Description string `json:"description,omitempty" yaml:"description,omitempty"`
|
||||
|
||||
// Author specifies the author or maintainer of the plugin.
|
||||
Author string `json:"author,omitempty" yaml:"author,omitempty"`
|
||||
|
||||
// Dependencies lists the plugin names that this plugin depends on.
|
||||
Dependencies []string `json:"dependencies,omitempty" yaml:"dependencies,omitempty"`
|
||||
|
||||
// Metadata holds arbitrary plugin-specific metadata.
|
||||
Metadata map[string]any `json:"metadata,omitempty" yaml:"metadata,omitempty"`
|
||||
}
|
||||
|
||||
// Validate checks whether the manifest satisfies basic integrity requirements.
|
||||
func (m Manifest) Validate() error {
|
||||
if strings.TrimSpace(m.Name) == "" {
|
||||
return fmt.Errorf("%w: %w", ErrInvalidManifest, ErrInvalidManifestName)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,150 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package core
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/backend/core/extpoints"
|
||||
)
|
||||
|
||||
// Standard sentinel errors returned by core operations.
|
||||
var (
|
||||
// ErrServiceNotFound is returned when a requested service is not registered in the IoC container.
|
||||
ErrServiceNotFound = errors.New("core: service not found")
|
||||
|
||||
// ErrServiceNotReady is returned when one or more required services are not ready in Using/UsingN.
|
||||
ErrServiceNotReady = errors.New("core: service not ready")
|
||||
|
||||
// ErrNilContext is returned when a nil Context is passed to an operation requiring a valid Context.
|
||||
ErrNilContext = errors.New("core: context is nil")
|
||||
|
||||
// ErrNilService is returned when attempting to provide a nil service implementation.
|
||||
ErrNilService = errors.New("core: service is nil")
|
||||
|
||||
// ErrInvalidManifest is returned when a plugin manifest fails validation.
|
||||
ErrInvalidManifest = errors.New("core: invalid manifest")
|
||||
|
||||
// ErrInvalidManifestName is returned when a plugin manifest has an empty name.
|
||||
ErrInvalidManifestName = errors.New("core: manifest name is required")
|
||||
|
||||
// ErrDriverNotFound is returned when a requested driver type is not registered.
|
||||
ErrDriverNotFound = errors.New("core: driver not found")
|
||||
|
||||
// ErrAppRunning is returned when attempting to start an already running App.
|
||||
ErrAppRunning = errors.New("core: app is already running")
|
||||
|
||||
// ErrAppNotRunning is returned when attempting to operate on an App that is not running.
|
||||
ErrAppNotRunning = errors.New("core: app is not running")
|
||||
)
|
||||
|
||||
// Plugin is the unified contract for all core and downstream plugins.
|
||||
type Plugin interface {
|
||||
// Name returns the globally unique identifier of the plugin (e.g. "auth", "database").
|
||||
Name() string
|
||||
// Apply is the core mounting entrypoint: provides services, registers routes, tasks, and event listeners.
|
||||
Apply(ctx *Context) error
|
||||
}
|
||||
|
||||
// PluginWithManifest is an optional extension interface for plugins that declare metadata.
|
||||
type PluginWithManifest interface {
|
||||
Plugin
|
||||
Manifest() Manifest
|
||||
}
|
||||
|
||||
// DriverType identifies the category of a runtime driver engine.
|
||||
type DriverType string
|
||||
|
||||
const (
|
||||
// DriverTypeHTTP represents HTTP web server drivers (e.g. Gin).
|
||||
DriverTypeHTTP DriverType = "http"
|
||||
|
||||
// DriverTypeWorker represents asynchronous background worker drivers (e.g. Asynq worker server).
|
||||
DriverTypeWorker DriverType = "worker"
|
||||
|
||||
// DriverTypeScheduler represents cron and timer schedule drivers (e.g. Asynq scheduler).
|
||||
DriverTypeScheduler DriverType = "schedule"
|
||||
)
|
||||
|
||||
// Driver is a runtime engine that manages an event loop or listening port.
|
||||
type Driver interface {
|
||||
// Type returns the category of this driver engine.
|
||||
Type() DriverType
|
||||
// Start starts the driver lifecycle loop.
|
||||
Start(ctx context.Context) error
|
||||
// Stop gracefully shuts down the driver.
|
||||
Stop(ctx context.Context) error
|
||||
}
|
||||
|
||||
// Profile identifies the runtime aspect or execution mode of an application.
|
||||
type Profile string
|
||||
|
||||
const (
|
||||
// ProfileAPI runs HTTP API server drivers.
|
||||
ProfileAPI Profile = "api"
|
||||
|
||||
// ProfileWorker runs asynchronous background worker drivers.
|
||||
ProfileWorker Profile = "worker"
|
||||
|
||||
// ProfileSchedule runs cron and timer schedule drivers.
|
||||
ProfileSchedule Profile = "schedule"
|
||||
|
||||
// ProfileAll runs all registered drivers concurrently in fused mode.
|
||||
ProfileAll Profile = "all"
|
||||
)
|
||||
|
||||
// MigrationEngine is the interface for executing database migrations across registered plugins.
|
||||
// The ctx parameter is the root micro-kernel Context, allowing the engine to resolve
|
||||
// services from the IoC container via core.Inject or core.Using.
|
||||
type MigrationEngine interface {
|
||||
Migrate(ctx *Context, entries []MigrationEntry) error
|
||||
}
|
||||
|
||||
// MigrationRunner is a function adapter implementing MigrationEngine.
|
||||
type MigrationRunner func(ctx *Context, entries []MigrationEntry) error
|
||||
|
||||
// Migrate calls the underlying migration function.
|
||||
func (fn MigrationRunner) Migrate(ctx *Context, entries []MigrationEntry) error {
|
||||
return fn(ctx, entries)
|
||||
}
|
||||
|
||||
// Disposer is a cleanup function executed when a Context is disposed.
|
||||
type Disposer func() error
|
||||
|
||||
// RouterExtension re-exports extpoints.RouterExtension.
|
||||
type RouterExtension = extpoints.RouterExtension
|
||||
|
||||
// RouteDefinition re-exports extpoints.RouteDefinition.
|
||||
type RouteDefinition = extpoints.RouteDefinition
|
||||
|
||||
// MigrationExtension re-exports extpoints.MigrationExtension.
|
||||
type MigrationExtension = extpoints.MigrationExtension
|
||||
|
||||
// MigrationEntry re-exports extpoints.MigrationEntry.
|
||||
type MigrationEntry = extpoints.MigrationEntry
|
||||
|
||||
// TaskExtension re-exports extpoints.TaskExtension.
|
||||
type TaskExtension = extpoints.TaskExtension
|
||||
|
||||
// TaskDefinition re-exports extpoints.TaskDefinition.
|
||||
type TaskDefinition = extpoints.TaskDefinition
|
||||
|
||||
// TaskOption re-exports extpoints.TaskOption.
|
||||
type TaskOption = extpoints.TaskOption
|
||||
|
||||
// ScheduleExtension re-exports extpoints.ScheduleExtension.
|
||||
type ScheduleExtension = extpoints.ScheduleExtension
|
||||
|
||||
// ScheduleDefinition re-exports extpoints.ScheduleDefinition.
|
||||
type ScheduleDefinition = extpoints.ScheduleDefinition
|
||||
|
||||
// ScheduleOption re-exports extpoints.ScheduleOption.
|
||||
type ScheduleOption = extpoints.ScheduleOption
|
||||
|
||||
// SettingExtension re-exports extpoints.SettingExtension.
|
||||
type SettingExtension = extpoints.SettingExtension
|
||||
|
||||
// SettingSchema re-exports extpoints.SettingSchema.
|
||||
type SettingSchema = extpoints.SettingSchema
|
||||
@@ -0,0 +1,77 @@
|
||||
# Downstream Custom Plugins
|
||||
|
||||
This directory is the designated location for downstream (deployment-specific) Cordis plugins.
|
||||
|
||||
## Architecture
|
||||
|
||||
```
|
||||
downstream/
|
||||
├── README.md
|
||||
└── plugins/
|
||||
└── custom_example/ # Example plugin — copy & rename to get started
|
||||
└── plugin.go
|
||||
```
|
||||
|
||||
Downstream plugins follow the same `core.Plugin` contract as platform plugins:
|
||||
|
||||
```go
|
||||
type Plugin interface {
|
||||
Name() string
|
||||
Apply(ctx *core.Context) error
|
||||
}
|
||||
```
|
||||
|
||||
## Rules
|
||||
|
||||
1. **Naming**: Each plugin directory name becomes its import path and plugin ID (kebab-case recommended).
|
||||
2. **Dependencies**: Downstream plugins may import `core/`, `core/contracts/`, `pkg/`, and `plugins/infra/` packages from the platform. They MUST NOT import domain plugin internal packages — use `core.Inject[contracts.XxxService](ctx)` instead.
|
||||
3. **Registration**: Add your downstream plugin to `cmd/app.go` before the platform plugins or after, depending on which services it needs:
|
||||
```go
|
||||
// newWaveletApp in cmd/app.go
|
||||
app.Use(
|
||||
database.New(),
|
||||
cache.New(),
|
||||
logger.New(),
|
||||
storage.New(),
|
||||
// ... platform domain plugins ...
|
||||
custom_hello.New(), // your downstream plugin
|
||||
driver_http.New(driver_http.WithAddr(config.Config.App.Addr)),
|
||||
driver_asynq_worker.New(),
|
||||
driver_asynq_cron.New(),
|
||||
)
|
||||
```
|
||||
4. **Migration**: If your plugin needs database tables, embed SQL files in a `migrations/` directory and register via `ctx.Migrations().Register(...)` in `Apply()`.
|
||||
|
||||
## Quick Start
|
||||
|
||||
```go
|
||||
package custom_example
|
||||
|
||||
import (
|
||||
"github.com/Rain-kl/Wavelet/core"
|
||||
"github.com/Rain-kl/Wavelet/core/contracts"
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
type Plugin struct{}
|
||||
|
||||
func New() *Plugin { return &Plugin{} }
|
||||
|
||||
func (p *Plugin) Name() string { return "custom_example" }
|
||||
|
||||
func (p *Plugin) Apply(ctx *core.Context) error {
|
||||
// Example: register a route that uses AuthService
|
||||
var authSvc contracts.AuthService
|
||||
if err := ctx.Using(func(svc contracts.AuthService) { authSvc = svc }); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
g := ctx.Router().Group("/api/v1/custom", authSvc.RequireAuthMiddleware().(gin.HandlerFunc))
|
||||
g.GET("/hello", func(c *gin.Context) {
|
||||
user, _ := authSvc.GetCurrentUser(c.Request.Context())
|
||||
c.JSON(200, gin.H{"message": "Hello " + user.Username})
|
||||
})
|
||||
|
||||
return nil
|
||||
}
|
||||
```
|
||||
@@ -0,0 +1,50 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package custom_example demonstrates how to build a downstream Cordis plugin.
|
||||
// Copy this directory to create your own plugin.
|
||||
package custom_example
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/backend/core"
|
||||
"github.com/Rain-kl/Wavelet/backend/core/contracts"
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
// Plugin implements core.Plugin for the custom_example downstream plugin.
|
||||
type Plugin struct{}
|
||||
|
||||
// New creates a new custom_example plugin.
|
||||
func New() *Plugin {
|
||||
return &Plugin{}
|
||||
}
|
||||
|
||||
// Name returns the unique identifier for this plugin.
|
||||
func (p *Plugin) Name() string {
|
||||
return "custom_example"
|
||||
}
|
||||
|
||||
// Apply registers routes and services into the Cordis micro-kernel Context.
|
||||
func (p *Plugin) Apply(ctx *core.Context) error {
|
||||
// Resolve platform services via IoC container (no direct imports of domain plugins).
|
||||
var authSvc contracts.AuthService
|
||||
if err := core.Using[contracts.AuthService](ctx, func(svc contracts.AuthService) { authSvc = svc }); err != nil {
|
||||
return err
|
||||
}
|
||||
_ = authSvc
|
||||
|
||||
// Register routes using the auth middleware obtained through the contract.
|
||||
g := ctx.Router().Group("/api/v1/custom", authSvc.RequireAuthMiddleware().(gin.HandlerFunc))
|
||||
g.GET("/hello", func(c *gin.Context) {
|
||||
user, err := authSvc.GetCurrentUser(c.Request.Context())
|
||||
if err != nil {
|
||||
c.JSON(http.StatusUnauthorized, gin.H{"error": "unauthorized"})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"message": "Hello " + user.Username})
|
||||
})
|
||||
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,22 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package main 是 Wavelet 平台的程序入口
|
||||
package main
|
||||
|
||||
import "github.com/Rain-kl/Wavelet/backend/cmd"
|
||||
|
||||
// @title Wavelet API
|
||||
// @version 1.0.0
|
||||
// @description Wavelet 平台后端 API,提供用户认证、系统配置、任务调度等通用功能。
|
||||
// @contact.name Wavelet
|
||||
// @contact.url https://github.com/Rain-kl/Wavelet
|
||||
// @license.name Apache 2.0
|
||||
// @license.url http://www.apache.org/licenses/LICENSE-2.0.html
|
||||
// @BasePath /
|
||||
// @securityDefinitions.apikey SessionCookie
|
||||
// @in cookie
|
||||
// @name session
|
||||
func main() {
|
||||
cmd.Execute()
|
||||
}
|
||||
@@ -0,0 +1,69 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package batchwriter
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"time"
|
||||
)
|
||||
|
||||
const (
|
||||
defaultQueueSize = 10_000
|
||||
defaultMaxBatchSize = 1_000
|
||||
defaultMinBatchSize = 50
|
||||
defaultFlushEvery = time.Second
|
||||
)
|
||||
|
||||
// Config controls queue capacity and flush thresholds for a Writer instance.
|
||||
type Config struct {
|
||||
// Name identifies the writer in logs and diagnostics. Optional.
|
||||
Name string
|
||||
|
||||
// QueueSize is the buffered channel capacity.
|
||||
QueueSize int
|
||||
|
||||
// MaxBatchSize triggers a flush when the in-memory batch reaches this count.
|
||||
MaxBatchSize int
|
||||
|
||||
// MinBatchSize is the minimum in-memory batch size for time-based flushes.
|
||||
// Zero disables the threshold and preserves legacy interval flush behavior.
|
||||
// When set, interval flushes below this size are skipped unless MaxFlushWait elapses.
|
||||
MinBatchSize int
|
||||
|
||||
// FlushInterval is how often the worker checks whether a time-based flush should run.
|
||||
FlushInterval time.Duration
|
||||
|
||||
// MaxFlushWait forces a flush of any non-empty batch once the oldest item has waited
|
||||
// this long, even if MinBatchSize has not been reached. Zero disables the force path.
|
||||
MaxFlushWait time.Duration
|
||||
}
|
||||
|
||||
// DefaultConfig returns production-friendly defaults aligned with audit log batching.
|
||||
func DefaultConfig() Config {
|
||||
return Config{
|
||||
QueueSize: defaultQueueSize,
|
||||
MaxBatchSize: defaultMaxBatchSize,
|
||||
MinBatchSize: defaultMinBatchSize,
|
||||
FlushInterval: defaultFlushEvery,
|
||||
}
|
||||
}
|
||||
|
||||
func (c Config) validate() error {
|
||||
if c.QueueSize <= 0 {
|
||||
return fmt.Errorf("batchwriter: queue size must be positive")
|
||||
}
|
||||
if c.MaxBatchSize <= 0 {
|
||||
return fmt.Errorf("batchwriter: max batch size must be positive")
|
||||
}
|
||||
if c.MinBatchSize < 0 {
|
||||
return fmt.Errorf("batchwriter: min batch size must be non-negative")
|
||||
}
|
||||
if c.FlushInterval <= 0 {
|
||||
return fmt.Errorf("batchwriter: flush interval must be positive")
|
||||
}
|
||||
if c.MaxFlushWait < 0 {
|
||||
return fmt.Errorf("batchwriter: max flush wait must be non-negative")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,8 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package batchwriter
|
||||
|
||||
import "errors"
|
||||
|
||||
var errNilFlushFunc = errors.New("batchwriter: flush func is required")
|
||||
@@ -0,0 +1,264 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package batchwriter provides a reusable buffered batch writer for high-throughput
|
||||
// append-only sinks such as ClickHouse. Each business domain should own an independent
|
||||
// Writer instance with its own queue, flush callback, and tuning parameters.
|
||||
package batchwriter
|
||||
|
||||
import (
|
||||
"context"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"time"
|
||||
)
|
||||
|
||||
// FlushFunc persists a batch of queued items. It is invoked from the worker goroutine.
|
||||
type FlushFunc[T any] func(ctx context.Context, items []T) error
|
||||
|
||||
// FlushErrorHandler is called when FlushFunc returns an error after optional retries.
|
||||
// The batch is discarded after the handler returns; the worker continues processing.
|
||||
// Handlers receive the failed items so callers can release dedup keys or re-queue.
|
||||
type FlushErrorHandler[T any] func(ctx context.Context, items []T, err error)
|
||||
|
||||
// Stats is a point-in-time snapshot of Writer queue and failure counters.
|
||||
type Stats struct {
|
||||
Name string
|
||||
Depth int
|
||||
Cap int
|
||||
Drops int64
|
||||
FlushErrors int64
|
||||
Running bool
|
||||
}
|
||||
|
||||
// Writer buffers items and flushes them by size or interval.
|
||||
type Writer[T any] struct {
|
||||
cfg Config
|
||||
flush FlushFunc[T]
|
||||
|
||||
onFlushError FlushErrorHandler[T]
|
||||
onDrop func(T)
|
||||
|
||||
startOnce sync.Once
|
||||
stopOnce sync.Once
|
||||
|
||||
mu sync.RWMutex
|
||||
ch chan T
|
||||
workerCtx context.Context
|
||||
done chan struct{}
|
||||
|
||||
drops atomic.Int64
|
||||
flushErrors atomic.Int64
|
||||
}
|
||||
|
||||
// Option configures optional Writer callbacks.
|
||||
type Option[T any] func(*Writer[T])
|
||||
|
||||
// WithFlushErrorHandler registers a callback for flush failures.
|
||||
func WithFlushErrorHandler[T any](handler FlushErrorHandler[T]) Option[T] {
|
||||
return func(w *Writer[T]) {
|
||||
w.onFlushError = handler
|
||||
}
|
||||
}
|
||||
|
||||
// WithDropHandler registers a callback when TryEnqueue cannot accept an item.
|
||||
func WithDropHandler[T any](handler func(T)) Option[T] {
|
||||
return func(w *Writer[T]) {
|
||||
w.onDrop = handler
|
||||
}
|
||||
}
|
||||
|
||||
// New creates a Writer. Call Start before enqueueing items.
|
||||
func New[T any](cfg Config, flush FlushFunc[T], opts ...Option[T]) (*Writer[T], error) {
|
||||
if flush == nil {
|
||||
return nil, errNilFlushFunc
|
||||
}
|
||||
if err := cfg.validate(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
w := &Writer[T]{
|
||||
cfg: cfg,
|
||||
flush: flush,
|
||||
done: make(chan struct{}),
|
||||
}
|
||||
for _, opt := range opts {
|
||||
opt(w)
|
||||
}
|
||||
return w, nil
|
||||
}
|
||||
|
||||
// Start launches the background worker. It is safe to call at most once.
|
||||
func (w *Writer[T]) Start(parent context.Context) {
|
||||
w.startOnce.Do(func() {
|
||||
w.mu.Lock()
|
||||
defer w.mu.Unlock()
|
||||
|
||||
w.ch = make(chan T, w.cfg.QueueSize)
|
||||
w.workerCtx = context.WithoutCancel(parent)
|
||||
go w.run()
|
||||
})
|
||||
}
|
||||
|
||||
// Stop closes the queue and waits until the worker drains pending items and exits.
|
||||
func (w *Writer[T]) Stop(ctx context.Context) error {
|
||||
w.mu.RLock()
|
||||
ch := w.ch
|
||||
done := w.done
|
||||
w.mu.RUnlock()
|
||||
|
||||
if ch == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
w.stopOnce.Do(func() {
|
||||
close(ch)
|
||||
})
|
||||
|
||||
select {
|
||||
case <-done:
|
||||
return nil
|
||||
case <-ctx.Done():
|
||||
return ctx.Err()
|
||||
}
|
||||
}
|
||||
|
||||
// Running reports whether Start has been called and Stop has not completed.
|
||||
func (w *Writer[T]) Running() bool {
|
||||
w.mu.RLock()
|
||||
defer w.mu.RUnlock()
|
||||
if w.ch == nil {
|
||||
return false
|
||||
}
|
||||
select {
|
||||
case <-w.done:
|
||||
return false
|
||||
default:
|
||||
return true
|
||||
}
|
||||
}
|
||||
|
||||
// TryEnqueue adds one item without blocking. It returns false when the writer is not
|
||||
// running or the queue is full.
|
||||
func (w *Writer[T]) TryEnqueue(item T) bool {
|
||||
w.mu.RLock()
|
||||
ch := w.ch
|
||||
w.mu.RUnlock()
|
||||
if ch == nil {
|
||||
w.notifyDrop(item)
|
||||
return false
|
||||
}
|
||||
|
||||
select {
|
||||
case ch <- item:
|
||||
return true
|
||||
default:
|
||||
w.notifyDrop(item)
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
// IsFull reports whether the queue has no remaining capacity.
|
||||
func (w *Writer[T]) IsFull() bool {
|
||||
w.mu.RLock()
|
||||
defer w.mu.RUnlock()
|
||||
if w.ch == nil {
|
||||
return false
|
||||
}
|
||||
return len(w.ch) >= cap(w.ch)
|
||||
}
|
||||
|
||||
// Len returns the current queue depth.
|
||||
func (w *Writer[T]) Len() int {
|
||||
w.mu.RLock()
|
||||
defer w.mu.RUnlock()
|
||||
if w.ch == nil {
|
||||
return 0
|
||||
}
|
||||
return len(w.ch)
|
||||
}
|
||||
|
||||
// Cap returns the queue capacity.
|
||||
func (w *Writer[T]) Cap() int {
|
||||
return w.cfg.QueueSize
|
||||
}
|
||||
|
||||
// Stats returns a point-in-time snapshot of queue depth and failure counters.
|
||||
func (w *Writer[T]) Stats() Stats {
|
||||
return Stats{
|
||||
Name: w.cfg.Name,
|
||||
Depth: w.Len(),
|
||||
Cap: w.Cap(),
|
||||
Drops: w.drops.Load(),
|
||||
FlushErrors: w.flushErrors.Load(),
|
||||
Running: w.Running(),
|
||||
}
|
||||
}
|
||||
|
||||
func (w *Writer[T]) run() {
|
||||
ticker := time.NewTicker(w.cfg.FlushInterval)
|
||||
defer ticker.Stop()
|
||||
|
||||
batch := make([]T, 0, w.cfg.MaxBatchSize)
|
||||
var batchStartedAt time.Time
|
||||
flush := func() {
|
||||
if len(batch) == 0 {
|
||||
return
|
||||
}
|
||||
items := append([]T(nil), batch...)
|
||||
if err := w.flush(w.workerCtx, items); err != nil {
|
||||
w.flushErrors.Add(1)
|
||||
if w.onFlushError != nil {
|
||||
w.onFlushError(w.workerCtx, items, err)
|
||||
}
|
||||
}
|
||||
batch = batch[:0]
|
||||
batchStartedAt = time.Time{}
|
||||
}
|
||||
|
||||
defer func() {
|
||||
flush()
|
||||
close(w.done)
|
||||
}()
|
||||
|
||||
for {
|
||||
select {
|
||||
case item, ok := <-w.ch:
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
if len(batch) == 0 {
|
||||
batchStartedAt = time.Now()
|
||||
}
|
||||
batch = append(batch, item)
|
||||
if len(batch) >= w.cfg.MaxBatchSize {
|
||||
flush()
|
||||
}
|
||||
case <-ticker.C:
|
||||
if w.shouldFlushOnInterval(len(batch), batchStartedAt, time.Now()) {
|
||||
flush()
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (w *Writer[T]) shouldFlushOnInterval(batchLen int, batchStartedAt time.Time, now time.Time) bool {
|
||||
if batchLen == 0 {
|
||||
return false
|
||||
}
|
||||
if w.cfg.MinBatchSize == 0 || batchLen >= w.cfg.MinBatchSize {
|
||||
return true
|
||||
}
|
||||
if w.cfg.MaxFlushWait <= 0 || batchStartedAt.IsZero() {
|
||||
return false
|
||||
}
|
||||
return !now.Before(batchStartedAt.Add(w.cfg.MaxFlushWait))
|
||||
}
|
||||
|
||||
func (w *Writer[T]) notifyDrop(item T) {
|
||||
w.drops.Add(1)
|
||||
if w.onDrop == nil {
|
||||
return
|
||||
}
|
||||
w.onDrop(item)
|
||||
}
|
||||
@@ -0,0 +1,317 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package batchwriter
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
type testEvent struct {
|
||||
ID int
|
||||
Data string
|
||||
}
|
||||
|
||||
func testConfig() Config {
|
||||
return Config{
|
||||
Name: "test-writer",
|
||||
QueueSize: 100,
|
||||
MaxBatchSize: 5,
|
||||
FlushInterval: 20 * time.Millisecond,
|
||||
}
|
||||
}
|
||||
|
||||
func TestWriter_BatchSizeFlush(t *testing.T) {
|
||||
var (
|
||||
mu sync.Mutex
|
||||
batches [][]testEvent
|
||||
flushWg sync.WaitGroup
|
||||
)
|
||||
flushWg.Add(1)
|
||||
|
||||
cfg := testConfig()
|
||||
cfg.FlushInterval = time.Hour
|
||||
|
||||
w, err := New(cfg, func(_ context.Context, items []testEvent) error {
|
||||
mu.Lock()
|
||||
defer mu.Unlock()
|
||||
batches = append(batches, items)
|
||||
if len(items) == 5 {
|
||||
flushWg.Done()
|
||||
}
|
||||
return nil
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
w.Start(context.Background())
|
||||
defer func() { _ = w.Stop(context.Background()) }()
|
||||
|
||||
for i := 1; i <= 5; i++ {
|
||||
ok := w.TryEnqueue(testEvent{ID: i, Data: "payload"})
|
||||
assert.True(t, ok)
|
||||
}
|
||||
|
||||
done := make(chan struct{})
|
||||
go func() {
|
||||
flushWg.Wait()
|
||||
close(done)
|
||||
}()
|
||||
|
||||
select {
|
||||
case <-done:
|
||||
case <-time.After(2 * time.Second):
|
||||
t.Fatal("timed out waiting for batch flush")
|
||||
}
|
||||
|
||||
mu.Lock()
|
||||
defer mu.Unlock()
|
||||
require.Len(t, batches, 1)
|
||||
assert.Len(t, batches[0], 5)
|
||||
for i, item := range batches[0] {
|
||||
assert.Equal(t, i+1, item.ID)
|
||||
}
|
||||
}
|
||||
|
||||
func TestWriter_IntervalFlush(t *testing.T) {
|
||||
var (
|
||||
mu sync.Mutex
|
||||
flushed []testEvent
|
||||
done = make(chan struct{})
|
||||
)
|
||||
|
||||
cfg := testConfig()
|
||||
cfg.MaxBatchSize = 100
|
||||
cfg.FlushInterval = 30 * time.Millisecond
|
||||
|
||||
w, err := New(cfg, func(_ context.Context, items []testEvent) error {
|
||||
mu.Lock()
|
||||
defer mu.Unlock()
|
||||
flushed = append(flushed, items...)
|
||||
if len(flushed) == 2 {
|
||||
select {
|
||||
case <-done:
|
||||
default:
|
||||
close(done)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
w.Start(context.Background())
|
||||
defer func() { _ = w.Stop(context.Background()) }()
|
||||
|
||||
assert.True(t, w.TryEnqueue(testEvent{ID: 1}))
|
||||
assert.True(t, w.TryEnqueue(testEvent{ID: 2}))
|
||||
|
||||
select {
|
||||
case <-done:
|
||||
case <-time.After(2 * time.Second):
|
||||
t.Fatal("timed out waiting for interval flush")
|
||||
}
|
||||
|
||||
mu.Lock()
|
||||
defer mu.Unlock()
|
||||
assert.Len(t, flushed, 2)
|
||||
}
|
||||
|
||||
func TestWriter_MinBatchSizeThreshold(t *testing.T) {
|
||||
var (
|
||||
mu sync.Mutex
|
||||
flushed []testEvent
|
||||
done = make(chan struct{})
|
||||
)
|
||||
|
||||
cfg := testConfig()
|
||||
cfg.MaxBatchSize = 100
|
||||
cfg.MinBatchSize = 3
|
||||
cfg.FlushInterval = 20 * time.Millisecond
|
||||
cfg.MaxFlushWait = 60 * time.Millisecond
|
||||
|
||||
w, err := New(cfg, func(_ context.Context, items []testEvent) error {
|
||||
mu.Lock()
|
||||
defer mu.Unlock()
|
||||
flushed = append(flushed, items...)
|
||||
if len(flushed) == 2 {
|
||||
select {
|
||||
case <-done:
|
||||
default:
|
||||
close(done)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
w.Start(context.Background())
|
||||
defer func() { _ = w.Stop(context.Background()) }()
|
||||
|
||||
assert.True(t, w.TryEnqueue(testEvent{ID: 1}))
|
||||
assert.True(t, w.TryEnqueue(testEvent{ID: 2}))
|
||||
|
||||
time.Sleep(30 * time.Millisecond)
|
||||
mu.Lock()
|
||||
assert.Empty(t, flushed, "items should wait until MinBatchSize or MaxFlushWait")
|
||||
mu.Unlock()
|
||||
|
||||
select {
|
||||
case <-done:
|
||||
case <-time.After(2 * time.Second):
|
||||
t.Fatal("timed out waiting for forced max wait flush")
|
||||
}
|
||||
|
||||
mu.Lock()
|
||||
defer mu.Unlock()
|
||||
assert.Len(t, flushed, 2)
|
||||
}
|
||||
|
||||
func TestWriter_StopDrainsRemaining(t *testing.T) {
|
||||
var (
|
||||
mu sync.Mutex
|
||||
flushed []testEvent
|
||||
)
|
||||
|
||||
cfg := testConfig()
|
||||
cfg.FlushInterval = time.Hour
|
||||
|
||||
w, err := New(cfg, func(_ context.Context, items []testEvent) error {
|
||||
mu.Lock()
|
||||
defer mu.Unlock()
|
||||
flushed = append(flushed, items...)
|
||||
return nil
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
w.Start(context.Background())
|
||||
for i := 1; i <= 3; i++ {
|
||||
assert.True(t, w.TryEnqueue(testEvent{ID: i}))
|
||||
}
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second)
|
||||
defer cancel()
|
||||
require.NoError(t, w.Stop(ctx))
|
||||
|
||||
assert.False(t, w.Running())
|
||||
mu.Lock()
|
||||
defer mu.Unlock()
|
||||
assert.Len(t, flushed, 3)
|
||||
}
|
||||
|
||||
func TestWriter_DropWhenFull(t *testing.T) {
|
||||
var (
|
||||
dropped atomic.Int64
|
||||
blockCh = make(chan struct{})
|
||||
)
|
||||
|
||||
cfg := Config{
|
||||
QueueSize: 2,
|
||||
MaxBatchSize: 1,
|
||||
FlushInterval: time.Hour,
|
||||
}
|
||||
|
||||
entered := make(chan struct{})
|
||||
w, err := New(cfg, func(_ context.Context, _ []testEvent) error {
|
||||
select {
|
||||
case entered <- struct{}{}:
|
||||
default:
|
||||
}
|
||||
<-blockCh
|
||||
return nil
|
||||
}, WithDropHandler(func(_ testEvent) {
|
||||
dropped.Add(1)
|
||||
}))
|
||||
require.NoError(t, err)
|
||||
|
||||
w.Start(context.Background())
|
||||
defer func() {
|
||||
close(blockCh)
|
||||
_ = w.Stop(context.Background())
|
||||
}()
|
||||
|
||||
// 1. 推入 1 个 item 触发 flush 并阻塞在 blockCh
|
||||
w.ch <- testEvent{ID: 1}
|
||||
<-entered
|
||||
|
||||
// 2. 此时 worker 阻塞,填满 channel
|
||||
w.ch <- testEvent{ID: 2}
|
||||
w.ch <- testEvent{ID: 3}
|
||||
|
||||
assert.False(t, w.TryEnqueue(testEvent{ID: 4}))
|
||||
assert.Equal(t, int64(1), dropped.Load())
|
||||
assert.Equal(t, int64(1), w.Stats().Drops)
|
||||
}
|
||||
|
||||
func TestWriter_FlushErrorCallback(t *testing.T) {
|
||||
var (
|
||||
called atomic.Bool
|
||||
flushErr = errors.New("clickhouse write timeout")
|
||||
done = make(chan struct{})
|
||||
)
|
||||
|
||||
cfg := testConfig()
|
||||
cfg.MaxBatchSize = 1
|
||||
cfg.FlushInterval = time.Hour
|
||||
|
||||
w, err := New(cfg, func(_ context.Context, _ []testEvent) error {
|
||||
return flushErr
|
||||
}, WithFlushErrorHandler(func(_ context.Context, items []testEvent, err error) {
|
||||
called.Store(true)
|
||||
assert.Equal(t, flushErr, err)
|
||||
assert.Len(t, items, 1)
|
||||
close(done)
|
||||
}))
|
||||
require.NoError(t, err)
|
||||
|
||||
w.Start(context.Background())
|
||||
defer func() { _ = w.Stop(context.Background()) }()
|
||||
|
||||
assert.True(t, w.TryEnqueue(testEvent{ID: 1}))
|
||||
|
||||
select {
|
||||
case <-done:
|
||||
case <-time.After(2 * time.Second):
|
||||
t.Fatal("timed out waiting for error callback")
|
||||
}
|
||||
|
||||
assert.True(t, called.Load())
|
||||
assert.Equal(t, int64(1), w.Stats().FlushErrors)
|
||||
}
|
||||
|
||||
func TestWriter_ValidateConfig(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
cfg Config
|
||||
wantErr bool
|
||||
}{
|
||||
{"valid", DefaultConfig(), false},
|
||||
{"zero queue", Config{QueueSize: 0, MaxBatchSize: 10, FlushInterval: time.Second}, true},
|
||||
{"zero max batch", Config{QueueSize: 10, MaxBatchSize: 0, FlushInterval: time.Second}, true},
|
||||
{"negative min batch", Config{QueueSize: 10, MaxBatchSize: 10, MinBatchSize: -1, FlushInterval: time.Second}, true},
|
||||
{"zero flush interval", Config{QueueSize: 10, MaxBatchSize: 10, FlushInterval: 0}, true},
|
||||
{"negative max flush wait", Config{QueueSize: 10, MaxBatchSize: 10, FlushInterval: time.Second, MaxFlushWait: -1}, true},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
_, err := New(tt.cfg, func(_ context.Context, _ []testEvent) error { return nil })
|
||||
if tt.wantErr {
|
||||
assert.Error(t, err)
|
||||
} else {
|
||||
assert.NoError(t, err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestWriter_NilFlushFunc(t *testing.T) {
|
||||
_, err := New[testEvent](DefaultConfig(), nil)
|
||||
assert.ErrorIs(t, err, errNilFlushFunc)
|
||||
}
|
||||
@@ -0,0 +1,12 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package buildinfo exposes metadata injected by the release workflow.
|
||||
package buildinfo
|
||||
|
||||
var (
|
||||
// Version is the application version.
|
||||
Version = "dev"
|
||||
// BuildTime is the UTC release build timestamp.
|
||||
BuildTime = ""
|
||||
)
|
||||
Vendored
+391
@@ -0,0 +1,391 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package disk implements a platform-level disk-backed cache with size limit, TTL, and LRU eviction.
|
||||
package disk
|
||||
|
||||
import (
|
||||
"container/list"
|
||||
"encoding/binary"
|
||||
"errors"
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"sort"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/peterbourgon/diskv/v3"
|
||||
)
|
||||
|
||||
// ErrCacheMiss represents a cache miss.
|
||||
var ErrCacheMiss = errors.New("cache miss")
|
||||
|
||||
// Constants for disk cache configuration and sizing
|
||||
const (
|
||||
headerSize = 8 // 8 bytes metadata prefix for expiration UnixNano timestamp
|
||||
defaultMaxSizeMB = 100
|
||||
defaultTTLMinutes = 60
|
||||
cacheDirPerm = 0750
|
||||
|
||||
// DefaultExpiration applies the cache-wide default TTL.
|
||||
DefaultExpiration time.Duration = 0
|
||||
// NoExpiration stores the item without a TTL. Size limits and LRU eviction still apply.
|
||||
NoExpiration time.Duration = -1
|
||||
)
|
||||
|
||||
// Status represents the runtime cache statistics.
|
||||
type Status struct {
|
||||
TotalSize int64 `json:"total_size"`
|
||||
KeysCount int `json:"keys_count"`
|
||||
MaxSizeMB int64 `json:"max_size_mb"`
|
||||
TTLMinutes int64 `json:"ttl_minutes"`
|
||||
LRUEnabled bool `json:"lru_enabled"`
|
||||
BasePath string `json:"base_path"`
|
||||
}
|
||||
|
||||
// Cache implements the disk-backed cache with size limits, TTL, and LRU eviction.
|
||||
type Cache struct {
|
||||
mu sync.RWMutex
|
||||
d *diskv.Diskv
|
||||
basePath string
|
||||
maxSize int64 // in bytes
|
||||
defaultTTL time.Duration
|
||||
lruEnabled bool
|
||||
|
||||
// LRU and Size tracking
|
||||
currentSize int64
|
||||
items map[string]*list.Element
|
||||
evictList *list.List
|
||||
}
|
||||
|
||||
type cacheItem struct {
|
||||
key string
|
||||
size int64
|
||||
expiredAt time.Time
|
||||
}
|
||||
|
||||
// New creates a new Cache instance.
|
||||
func New(basePath string) *Cache {
|
||||
d := diskv.New(diskv.Options{
|
||||
BasePath: basePath,
|
||||
Transform: func(_ string) []string { return []string{} }, // flat structure for easy walk
|
||||
CacheSizeMax: 1024 * 1024, // 1MB in-memory cache size for diskv itself
|
||||
})
|
||||
|
||||
c := &Cache{
|
||||
d: d,
|
||||
basePath: basePath,
|
||||
maxSize: defaultMaxSizeMB * 1024 * 1024, // 100MB default
|
||||
defaultTTL: defaultTTLMinutes * time.Minute, // 60 minutes default
|
||||
lruEnabled: true,
|
||||
items: make(map[string]*list.Element),
|
||||
evictList: list.New(),
|
||||
}
|
||||
|
||||
// Scan directory on startup to rebuild LRU and size tracking
|
||||
_ = c.loadTracker()
|
||||
return c
|
||||
}
|
||||
|
||||
// Set stores a key-value pair in the cache.
|
||||
// Use DefaultExpiration for the configured default TTL, NoExpiration for no
|
||||
// TTL, or a positive duration for a business-specific TTL.
|
||||
func (c *Cache) Set(key string, value []byte, ttl time.Duration) error {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
|
||||
if ttl == DefaultExpiration {
|
||||
ttl = c.defaultTTL
|
||||
}
|
||||
|
||||
var expiredAt time.Time
|
||||
if ttl > 0 {
|
||||
expiredAt = time.Now().Add(ttl)
|
||||
}
|
||||
|
||||
// Prepare data layout: 8 bytes expiration timestamp + raw payload
|
||||
buf := make([]byte, headerSize+len(value))
|
||||
var expNano int64
|
||||
if !expiredAt.IsZero() {
|
||||
expNano = expiredAt.UnixNano()
|
||||
}
|
||||
binary.BigEndian.PutUint64(buf[0:headerSize], uint64(expNano))
|
||||
copy(buf[headerSize:], value)
|
||||
|
||||
// Write to diskv
|
||||
if err := c.d.Write(key, buf); err != nil {
|
||||
return fmt.Errorf("failed to write key to disk: %w", err)
|
||||
}
|
||||
|
||||
// Get file size on disk (approximate)
|
||||
size := int64(len(buf))
|
||||
|
||||
// Update memory tracker
|
||||
if elem, ok := c.items[key]; ok {
|
||||
item := elem.Value.(*cacheItem)
|
||||
c.currentSize += size - item.size
|
||||
item.size = size
|
||||
item.expiredAt = expiredAt
|
||||
c.evictList.MoveToFront(elem)
|
||||
} else {
|
||||
item := &cacheItem{
|
||||
key: key,
|
||||
size: size,
|
||||
expiredAt: expiredAt,
|
||||
}
|
||||
elem := c.evictList.PushFront(item)
|
||||
c.items[key] = elem
|
||||
c.currentSize += size
|
||||
}
|
||||
|
||||
// Evict items if size limit exceeded and LRU is enabled
|
||||
c.evict()
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// Get retrieves a key's value from the cache.
|
||||
func (c *Cache) Get(key string) ([]byte, error) {
|
||||
c.mu.RLock()
|
||||
elem, ok := c.items[key]
|
||||
if !ok {
|
||||
c.mu.RUnlock()
|
||||
return nil, ErrCacheMiss
|
||||
}
|
||||
|
||||
item := elem.Value.(*cacheItem)
|
||||
if !item.expiredAt.IsZero() && time.Now().After(item.expiredAt) {
|
||||
c.mu.RUnlock()
|
||||
return c.getAndDeleteIfExpired(key)
|
||||
}
|
||||
c.mu.RUnlock()
|
||||
|
||||
// Read from disk outside the lock so concurrent cache hits do not serialize on I/O.
|
||||
data, err := c.d.Read(key)
|
||||
if err != nil {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
if _, stillExists := c.items[key]; stillExists {
|
||||
_ = c.deleteUnlocked(key)
|
||||
}
|
||||
return nil, ErrCacheMiss
|
||||
}
|
||||
|
||||
if len(data) < headerSize {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
if _, stillExists := c.items[key]; stillExists {
|
||||
_ = c.deleteUnlocked(key)
|
||||
}
|
||||
return nil, ErrCacheMiss
|
||||
}
|
||||
|
||||
payload := data[headerSize:]
|
||||
|
||||
// Brief write lock only for LRU bookkeeping.
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
elem, ok = c.items[key]
|
||||
if !ok {
|
||||
return nil, ErrCacheMiss
|
||||
}
|
||||
c.evictList.MoveToFront(elem)
|
||||
|
||||
return payload, nil
|
||||
}
|
||||
|
||||
func (c *Cache) getAndDeleteIfExpired(key string) ([]byte, error) {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
|
||||
elem, ok := c.items[key]
|
||||
if !ok {
|
||||
return nil, ErrCacheMiss
|
||||
}
|
||||
|
||||
item := elem.Value.(*cacheItem)
|
||||
if !item.expiredAt.IsZero() && time.Now().After(item.expiredAt) {
|
||||
_ = c.deleteUnlocked(key)
|
||||
return nil, ErrCacheMiss
|
||||
}
|
||||
|
||||
data, err := c.d.Read(key)
|
||||
if err != nil {
|
||||
_ = c.deleteUnlocked(key)
|
||||
return nil, ErrCacheMiss
|
||||
}
|
||||
|
||||
if len(data) < headerSize {
|
||||
_ = c.deleteUnlocked(key)
|
||||
return nil, ErrCacheMiss
|
||||
}
|
||||
|
||||
c.evictList.MoveToFront(elem)
|
||||
return data[headerSize:], nil
|
||||
}
|
||||
|
||||
// Delete removes a key-value pair from the cache.
|
||||
func (c *Cache) Delete(key string) error {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
return c.deleteUnlocked(key)
|
||||
}
|
||||
|
||||
func (c *Cache) deleteUnlocked(key string) error {
|
||||
if elem, ok := c.items[key]; ok {
|
||||
item := elem.Value.(*cacheItem)
|
||||
c.currentSize -= item.size
|
||||
c.evictList.Remove(elem)
|
||||
delete(c.items, key)
|
||||
}
|
||||
return c.d.Erase(key)
|
||||
}
|
||||
|
||||
// Clear flushes all cached elements.
|
||||
func (c *Cache) Clear() error {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
|
||||
c.currentSize = 0
|
||||
c.items = make(map[string]*list.Element)
|
||||
c.evictList.Init()
|
||||
|
||||
return c.d.EraseAll()
|
||||
}
|
||||
|
||||
// Status returns the cache status.
|
||||
func (c *Cache) Status() Status {
|
||||
c.mu.RLock()
|
||||
defer c.mu.RUnlock()
|
||||
|
||||
return Status{
|
||||
TotalSize: c.currentSize,
|
||||
KeysCount: len(c.items),
|
||||
MaxSizeMB: c.maxSize / (1024 * 1024),
|
||||
TTLMinutes: int64(c.defaultTTL.Minutes()),
|
||||
LRUEnabled: c.lruEnabled,
|
||||
BasePath: c.basePath,
|
||||
}
|
||||
}
|
||||
|
||||
// UpdatePolicy dynamically updates policies.
|
||||
func (c *Cache) UpdatePolicy(maxSizeMB int64, ttlMinutes int64, lruEnabled bool) {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
|
||||
c.maxSize = maxSizeMB * 1024 * 1024
|
||||
c.defaultTTL = time.Duration(ttlMinutes) * time.Minute
|
||||
c.lruEnabled = lruEnabled
|
||||
c.evict()
|
||||
}
|
||||
|
||||
// evict evicts oldest items if current size exceeds maxSize and LRU is enabled.
|
||||
func (c *Cache) evict() {
|
||||
if !c.lruEnabled {
|
||||
return
|
||||
}
|
||||
|
||||
for c.currentSize > c.maxSize && c.evictList.Len() > 0 {
|
||||
elem := c.evictList.Back()
|
||||
item := elem.Value.(*cacheItem)
|
||||
c.currentSize -= item.size
|
||||
c.evictList.Remove(elem)
|
||||
delete(c.items, item.key)
|
||||
_ = c.d.Erase(item.key)
|
||||
}
|
||||
}
|
||||
|
||||
// loadTracker scans the cache directory on startup to rebuild memory state.
|
||||
func (c *Cache) loadTracker() error {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
|
||||
// Ensure directory exists
|
||||
if err := os.MkdirAll(c.basePath, cacheDirPerm); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
type loadedItem struct {
|
||||
key string
|
||||
size int64
|
||||
expiredAt time.Time
|
||||
modTime time.Time
|
||||
}
|
||||
var loadedItems []loadedItem
|
||||
|
||||
// Walk keys through diskv
|
||||
keysChan := c.d.Keys(nil)
|
||||
for key := range keysChan {
|
||||
// Read raw bytes to parse expiration prefix
|
||||
data, err := c.d.Read(key)
|
||||
if err != nil || len(data) < headerSize {
|
||||
_ = c.d.Erase(key) // corrupted file, wipe
|
||||
continue
|
||||
}
|
||||
|
||||
expNano := int64(binary.BigEndian.Uint64(data[0:headerSize])) //nolint:gosec // false positive: UnixNano fits within int64
|
||||
var expiredAt time.Time
|
||||
if expNano > 0 {
|
||||
expiredAt = time.Unix(0, expNano)
|
||||
}
|
||||
|
||||
// Check mod time for ordering
|
||||
path := filepath.Join(c.basePath, key)
|
||||
info, err := os.Stat(path)
|
||||
if err != nil {
|
||||
continue
|
||||
}
|
||||
|
||||
loadedItems = append(loadedItems, loadedItem{
|
||||
key: key,
|
||||
size: int64(len(data)),
|
||||
expiredAt: expiredAt,
|
||||
modTime: info.ModTime(),
|
||||
})
|
||||
}
|
||||
|
||||
// Sort by ModTime ascending (oldest first) so we rebuild LRU correctly
|
||||
sort.Slice(loadedItems, func(i, j int) bool {
|
||||
return loadedItems[i].modTime.Before(loadedItems[j].modTime)
|
||||
})
|
||||
|
||||
// Populate LRU (PushFront so that newest items are at the front, oldest at the back)
|
||||
for _, item := range loadedItems {
|
||||
entry := &cacheItem{
|
||||
key: item.key,
|
||||
size: item.size,
|
||||
expiredAt: item.expiredAt,
|
||||
}
|
||||
element := c.evictList.PushFront(entry)
|
||||
c.items[item.key] = element
|
||||
c.currentSize += item.size
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// StartCleanupWorker periodically cleans up expired cache items.
|
||||
func (c *Cache) StartCleanupWorker(interval time.Duration) {
|
||||
ticker := time.NewTicker(interval)
|
||||
for range ticker.C {
|
||||
c.cleanExpired()
|
||||
}
|
||||
}
|
||||
|
||||
// cleanExpired scans memory for expired items and removes them.
|
||||
func (c *Cache) cleanExpired() {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
|
||||
now := time.Now()
|
||||
for key, elem := range c.items {
|
||||
item := elem.Value.(*cacheItem)
|
||||
if !item.expiredAt.IsZero() && now.After(item.expiredAt) {
|
||||
c.currentSize -= item.size
|
||||
c.evictList.Remove(elem)
|
||||
delete(c.items, key)
|
||||
_ = c.d.Erase(key)
|
||||
}
|
||||
}
|
||||
}
|
||||
Vendored
+212
@@ -0,0 +1,212 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package disk
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"os"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func TestDiskCacheBasic(t *testing.T) {
|
||||
testDir := "uploads/test_diskcache_basic"
|
||||
defer func() { _ = os.RemoveAll(testDir) }()
|
||||
_ = os.RemoveAll(testDir)
|
||||
|
||||
c := New(testDir)
|
||||
defer func() { _ = c.Clear() }()
|
||||
|
||||
key := "key1"
|
||||
val := []byte("value1")
|
||||
|
||||
// Get non-existent
|
||||
_, err := c.Get(key)
|
||||
if err != ErrCacheMiss {
|
||||
t.Fatalf("expected ErrCacheMiss, got %v", err)
|
||||
}
|
||||
|
||||
// Set & Get
|
||||
err = c.Set(key, val, 10*time.Second)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to set cache: %v", err)
|
||||
}
|
||||
|
||||
got, err := c.Get(key)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to get cache: %v", err)
|
||||
}
|
||||
|
||||
if !bytes.Equal(got, val) {
|
||||
t.Errorf("expected %s, got %s", val, got)
|
||||
}
|
||||
|
||||
// Delete
|
||||
err = c.Delete(key)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to delete: %v", err)
|
||||
}
|
||||
|
||||
_, err = c.Get(key)
|
||||
if err != ErrCacheMiss {
|
||||
t.Errorf("expected ErrCacheMiss after delete, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDiskCacheTTL(t *testing.T) {
|
||||
testDir := "uploads/test_diskcache_ttl"
|
||||
defer func() { _ = os.RemoveAll(testDir) }()
|
||||
_ = os.RemoveAll(testDir)
|
||||
|
||||
c := New(testDir)
|
||||
defer func() { _ = c.Clear() }()
|
||||
|
||||
key := "ttlkey"
|
||||
val := []byte("ttlval")
|
||||
|
||||
// Set with 200ms TTL
|
||||
err := c.Set(key, val, 200*time.Millisecond)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to set: %v", err)
|
||||
}
|
||||
|
||||
// Immediate Get should succeed
|
||||
got, err := c.Get(key)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to get: %v", err)
|
||||
}
|
||||
if !bytes.Equal(got, val) {
|
||||
t.Errorf("expected %s, got %s", val, got)
|
||||
}
|
||||
|
||||
// Sleep 250ms to expire
|
||||
time.Sleep(250 * time.Millisecond)
|
||||
|
||||
// Get should fail with cache miss
|
||||
_, err = c.Get(key)
|
||||
if err != ErrCacheMiss {
|
||||
t.Errorf("expected ErrCacheMiss after TTL expiration, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDiskCacheExpirationPolicies(t *testing.T) {
|
||||
testDir := "uploads/test_diskcache_expiration_policies"
|
||||
defer func() { _ = os.RemoveAll(testDir) }()
|
||||
_ = os.RemoveAll(testDir)
|
||||
|
||||
c := New(testDir)
|
||||
defer func() { _ = c.Clear() }()
|
||||
c.defaultTTL = 50 * time.Millisecond
|
||||
|
||||
if err := c.Set("default", []byte("default"), DefaultExpiration); err != nil {
|
||||
t.Fatalf("Set(default, DefaultExpiration) returned error: %v", err)
|
||||
}
|
||||
if err := c.Set("custom", []byte("custom"), 100*time.Millisecond); err != nil {
|
||||
t.Fatalf("Set(custom, 100ms) returned error: %v", err)
|
||||
}
|
||||
if err := c.Set("permanent", []byte("permanent"), NoExpiration); err != nil {
|
||||
t.Fatalf("Set(permanent, NoExpiration) returned error: %v", err)
|
||||
}
|
||||
|
||||
time.Sleep(75 * time.Millisecond)
|
||||
|
||||
if _, err := c.Get("default"); err != ErrCacheMiss {
|
||||
t.Errorf("Get(default) error = %v, want ErrCacheMiss", err)
|
||||
}
|
||||
if _, err := c.Get("custom"); err != nil {
|
||||
t.Errorf("Get(custom) returned error before custom TTL elapsed: %v", err)
|
||||
}
|
||||
if _, err := c.Get("permanent"); err != nil {
|
||||
t.Errorf("Get(permanent) returned error: %v", err)
|
||||
}
|
||||
|
||||
time.Sleep(50 * time.Millisecond)
|
||||
|
||||
if _, err := c.Get("custom"); err != ErrCacheMiss {
|
||||
t.Errorf("Get(custom) error = %v, want ErrCacheMiss", err)
|
||||
}
|
||||
if _, err := c.Get("permanent"); err != nil {
|
||||
t.Errorf("Get(permanent) returned error after other entries expired: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDiskCacheNoExpirationSurvivesReload(t *testing.T) {
|
||||
testDir := "uploads/test_diskcache_no_expiration_reload"
|
||||
defer func() { _ = os.RemoveAll(testDir) }()
|
||||
_ = os.RemoveAll(testDir)
|
||||
|
||||
c := New(testDir)
|
||||
if err := c.Set("permanent", []byte("value"), NoExpiration); err != nil {
|
||||
t.Fatalf("Set(permanent, NoExpiration) returned error: %v", err)
|
||||
}
|
||||
|
||||
reloaded := New(testDir)
|
||||
defer func() { _ = reloaded.Clear() }()
|
||||
|
||||
got, err := reloaded.Get("permanent")
|
||||
if err != nil {
|
||||
t.Fatalf("reloaded Get(permanent) returned error: %v", err)
|
||||
}
|
||||
if !bytes.Equal(got, []byte("value")) {
|
||||
t.Errorf("reloaded Get(permanent) = %q, want %q", got, "value")
|
||||
}
|
||||
}
|
||||
|
||||
func TestDiskCacheLRUEviction(t *testing.T) {
|
||||
testDir := "uploads/test_diskcache_lru"
|
||||
defer func() { _ = os.RemoveAll(testDir) }()
|
||||
_ = os.RemoveAll(testDir)
|
||||
|
||||
c := New(testDir)
|
||||
defer func() { _ = c.Clear() }()
|
||||
|
||||
// Force a very small max size of 20 bytes for testing (8 bytes header + payload)
|
||||
// So 2 items of 2 bytes payload = 2 * (8 + 2) = 20 bytes max.
|
||||
c.maxSize = 20
|
||||
c.lruEnabled = true
|
||||
|
||||
// Write item 1: 8 + 2 = 10 bytes
|
||||
err := c.Set("k1", []byte("v1"), DefaultExpiration)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to set k1: %v", err)
|
||||
}
|
||||
|
||||
// Write item 2: 8 + 2 = 10 bytes
|
||||
err = c.Set("k2", []byte("v2"), DefaultExpiration)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to set k2: %v", err)
|
||||
}
|
||||
|
||||
// Both should exist
|
||||
if _, err := c.Get("k1"); err != nil {
|
||||
t.Errorf("k1 should exist: %v", err)
|
||||
}
|
||||
if _, err := c.Get("k2"); err != nil {
|
||||
t.Errorf("k2 should exist: %v", err)
|
||||
}
|
||||
|
||||
// Write item 3: 8 + 2 = 10 bytes -> total size would be 30, exceeding 20.
|
||||
// This should evict the oldest item. Since k1 was accessed, but then k2 was accessed,
|
||||
// wait, let's access k1 again to make it the most recently used, so k2 becomes oldest!
|
||||
_, _ = c.Get("k1") // k1 is now MRU, k2 is LRU
|
||||
|
||||
err = c.Set("k3", []byte("v3"), DefaultExpiration)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to set k3: %v", err)
|
||||
}
|
||||
|
||||
// k2 should be evicted, k1 and k3 should exist
|
||||
_, err = c.Get("k2")
|
||||
if err != ErrCacheMiss {
|
||||
t.Errorf("expected k2 to be evicted, got error %v", err)
|
||||
}
|
||||
|
||||
if _, err := c.Get("k1"); err != nil {
|
||||
t.Errorf("k1 should still exist: %v", err)
|
||||
}
|
||||
|
||||
if _, err := c.Get("k3"); err != nil {
|
||||
t.Errorf("k3 should exist: %v", err)
|
||||
}
|
||||
}
|
||||
Vendored
+73
@@ -0,0 +1,73 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package ram provides a thin wrapper around Otter v2 for process-local caching.
|
||||
package ram
|
||||
|
||||
import (
|
||||
"github.com/maypok86/otter/v2"
|
||||
)
|
||||
|
||||
const defaultMaximumSize = 256
|
||||
|
||||
// Options configures a RAM cache instance.
|
||||
type Options struct {
|
||||
// MaximumSize bounds the number of entries. Zero uses a small default.
|
||||
MaximumSize int
|
||||
}
|
||||
|
||||
// Cache is a concurrency-safe in-memory cache backed by Otter.
|
||||
type Cache[K comparable, V any] struct {
|
||||
inner *otter.Cache[K, V]
|
||||
}
|
||||
|
||||
// New creates a RAM cache from the provided options.
|
||||
func New[K comparable, V any](opts Options) (*Cache[K, V], error) {
|
||||
maximumSize := opts.MaximumSize
|
||||
if maximumSize == 0 {
|
||||
maximumSize = defaultMaximumSize
|
||||
}
|
||||
|
||||
inner, err := otter.New(&otter.Options[K, V]{
|
||||
MaximumSize: maximumSize,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return &Cache[K, V]{inner: inner}, nil
|
||||
}
|
||||
|
||||
// MustNew creates a RAM cache and panics when configuration is invalid.
|
||||
func MustNew[K comparable, V any](opts Options) *Cache[K, V] {
|
||||
cache, err := New[K, V](opts)
|
||||
if err != nil {
|
||||
panic(err)
|
||||
}
|
||||
return cache
|
||||
}
|
||||
|
||||
// GetIfPresent returns the cached value when present.
|
||||
func (c *Cache[K, V]) GetIfPresent(key K) (V, bool) {
|
||||
return c.inner.GetIfPresent(key)
|
||||
}
|
||||
|
||||
// Set stores a value in the cache.
|
||||
func (c *Cache[K, V]) Set(key K, value V) {
|
||||
c.inner.Set(key, value)
|
||||
}
|
||||
|
||||
// Invalidate removes one entry from the cache.
|
||||
func (c *Cache[K, V]) Invalidate(key K) {
|
||||
c.inner.Invalidate(key)
|
||||
}
|
||||
|
||||
// InvalidateAll removes every entry from the cache.
|
||||
func (c *Cache[K, V]) InvalidateAll() {
|
||||
c.inner.InvalidateAll()
|
||||
}
|
||||
|
||||
// EstimatedSize returns the approximate number of cached entries.
|
||||
func (c *Cache[K, V]) EstimatedSize() int {
|
||||
return c.inner.EstimatedSize()
|
||||
}
|
||||
Vendored
+38
@@ -0,0 +1,38 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package ram
|
||||
|
||||
import "testing"
|
||||
|
||||
func TestCacheSetGetInvalidate(t *testing.T) {
|
||||
cache := MustNew[string, int](Options{MaximumSize: 8})
|
||||
|
||||
cache.Set("count", 3)
|
||||
|
||||
got, ok := cache.GetIfPresent("count")
|
||||
if !ok {
|
||||
t.Fatal("GetIfPresent(count) ok = false, want true")
|
||||
}
|
||||
if got != 3 {
|
||||
t.Fatalf("GetIfPresent(count) = %d, want %d", got, 3)
|
||||
}
|
||||
|
||||
cache.Invalidate("count")
|
||||
if _, ok := cache.GetIfPresent("count"); ok {
|
||||
t.Fatal("GetIfPresent(count) after Invalidate ok = true, want false")
|
||||
}
|
||||
}
|
||||
|
||||
func TestCacheInvalidateAll(t *testing.T) {
|
||||
cache := MustNew[string, string](Options{MaximumSize: 8})
|
||||
|
||||
cache.Set("a", "1")
|
||||
cache.Set("b", "2")
|
||||
|
||||
cache.InvalidateAll()
|
||||
|
||||
if cache.EstimatedSize() != 0 {
|
||||
t.Fatalf("EstimatedSize() after InvalidateAll = %d, want 0", cache.EstimatedSize())
|
||||
}
|
||||
}
|
||||
Vendored
+233
@@ -0,0 +1,233 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package ram
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"sync"
|
||||
"time"
|
||||
)
|
||||
|
||||
var (
|
||||
// ErrNotFound is returned by the Loader when the requested item is not found.
|
||||
ErrNotFound = errors.New("cache item not found in data source")
|
||||
|
||||
managerCache *Cache[string, map[string]cacheEntry]
|
||||
|
||||
writeLocks = make(map[string]*sync.Mutex)
|
||||
writeLocksMu sync.Mutex
|
||||
)
|
||||
|
||||
// CacheItem represents a unified cache entity.
|
||||
type CacheItem struct {
|
||||
Key string `json:"key"`
|
||||
Value string `json:"value"`
|
||||
Type string `json:"type"`
|
||||
TTL time.Duration `json:"ttl"` // -1 means never expire
|
||||
}
|
||||
|
||||
// Loader is an interface that the cache client must implement to handle database retrieval.
|
||||
type Loader interface {
|
||||
LoadAll(ctx context.Context, configType string) ([]CacheItem, error)
|
||||
LoadOne(ctx context.Context, configType string, key string) (CacheItem, error)
|
||||
}
|
||||
|
||||
type cacheEntry struct {
|
||||
item CacheItem
|
||||
expireAt time.Time
|
||||
}
|
||||
|
||||
func init() {
|
||||
// Initialize with a large maximum size since it only stores one entry per configType
|
||||
managerCache = MustNew[string, map[string]cacheEntry](Options{
|
||||
MaximumSize: 1000,
|
||||
})
|
||||
}
|
||||
|
||||
func getWriteLock(configType string) *sync.Mutex {
|
||||
writeLocksMu.Lock()
|
||||
defer writeLocksMu.Unlock()
|
||||
lock, found := writeLocks[configType]
|
||||
if !found {
|
||||
lock = &sync.Mutex{}
|
||||
writeLocks[configType] = lock
|
||||
}
|
||||
return lock
|
||||
}
|
||||
|
||||
// Get retrieves a cache item from the local cache store, checking for expiration.
|
||||
// Reads are completely lock-free because maps stored in Otter are immutable.
|
||||
func Get(configType string, key string) (CacheItem, bool) {
|
||||
m, ok := managerCache.GetIfPresent(configType)
|
||||
if !ok {
|
||||
return CacheItem{}, false
|
||||
}
|
||||
|
||||
entry, found := m[key]
|
||||
if !found {
|
||||
return CacheItem{}, false
|
||||
}
|
||||
|
||||
// Check expiration
|
||||
if entry.item.TTL != -1 && !entry.expireAt.IsZero() && time.Now().After(entry.expireAt) {
|
||||
// Asynchronously remove the expired item from the map and write back
|
||||
go deleteKeyIfExpired(configType, key, entry.expireAt)
|
||||
return CacheItem{}, false
|
||||
}
|
||||
|
||||
return entry.item, true
|
||||
}
|
||||
|
||||
func deleteKeyIfExpired(configType string, key string, expireAt time.Time) {
|
||||
lock := getWriteLock(configType)
|
||||
lock.Lock()
|
||||
defer lock.Unlock()
|
||||
|
||||
currentMap, ok := managerCache.GetIfPresent(configType)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
|
||||
entry, found := currentMap[key]
|
||||
if !found {
|
||||
return
|
||||
}
|
||||
|
||||
// Double-check expiration time to ensure we don't delete a newly updated key
|
||||
if entry.expireAt != expireAt || !time.Now().After(entry.expireAt) {
|
||||
return
|
||||
}
|
||||
|
||||
newMap := make(map[string]cacheEntry, len(currentMap)-1)
|
||||
for k, v := range currentMap {
|
||||
if k != key {
|
||||
newMap[k] = v
|
||||
}
|
||||
}
|
||||
managerCache.Set(configType, newMap)
|
||||
}
|
||||
|
||||
// Set stores a cache item in the local cache store.
|
||||
// Writes are protected by a fine-grained lock per configType.
|
||||
func Set(item CacheItem) {
|
||||
lock := getWriteLock(item.Type)
|
||||
lock.Lock()
|
||||
defer lock.Unlock()
|
||||
|
||||
currentMap, ok := managerCache.GetIfPresent(item.Type)
|
||||
newMap := make(map[string]cacheEntry)
|
||||
if ok {
|
||||
for k, v := range currentMap {
|
||||
newMap[k] = v
|
||||
}
|
||||
}
|
||||
|
||||
var expireAt time.Time
|
||||
if item.TTL != -1 {
|
||||
expireAt = time.Now().Add(item.TTL)
|
||||
}
|
||||
|
||||
newMap[item.Key] = cacheEntry{
|
||||
item: item,
|
||||
expireAt: expireAt,
|
||||
}
|
||||
managerCache.Set(item.Type, newMap)
|
||||
}
|
||||
|
||||
// Delete removes a single item from the local cache store.
|
||||
// Writes are protected by a fine-grained lock per configType.
|
||||
func Delete(configType string, key string) {
|
||||
lock := getWriteLock(configType)
|
||||
lock.Lock()
|
||||
defer lock.Unlock()
|
||||
|
||||
currentMap, ok := managerCache.GetIfPresent(configType)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
|
||||
newMap := make(map[string]cacheEntry, len(currentMap))
|
||||
for k, v := range currentMap {
|
||||
if k != key {
|
||||
newMap[k] = v
|
||||
}
|
||||
}
|
||||
managerCache.Set(configType, newMap)
|
||||
}
|
||||
|
||||
// UpdateTypeItems replaces all cache items of a specific type atomically.
|
||||
// Writes are protected by a fine-grained lock per configType.
|
||||
func UpdateTypeItems(configType string, items []CacheItem) {
|
||||
lock := getWriteLock(configType)
|
||||
lock.Lock()
|
||||
defer lock.Unlock()
|
||||
|
||||
newMap := make(map[string]cacheEntry, len(items))
|
||||
for _, item := range items {
|
||||
var expireAt time.Time
|
||||
if item.TTL != -1 {
|
||||
expireAt = time.Now().Add(item.TTL)
|
||||
}
|
||||
newMap[item.Key] = cacheEntry{
|
||||
item: item,
|
||||
expireAt: expireAt,
|
||||
}
|
||||
}
|
||||
managerCache.Set(configType, newMap)
|
||||
}
|
||||
|
||||
// GetTypeItems retrieves all unexpired cache items of a specific type.
|
||||
// Reads are completely lock-free because maps stored in Otter are immutable.
|
||||
func GetTypeItems(configType string) []CacheItem {
|
||||
currentMap, ok := managerCache.GetIfPresent(configType)
|
||||
if !ok {
|
||||
return nil
|
||||
}
|
||||
|
||||
var list []CacheItem
|
||||
for _, entry := range currentMap {
|
||||
if entry.item.TTL == -1 || entry.expireAt.IsZero() || time.Now().Before(entry.expireAt) {
|
||||
list = append(list, entry.item)
|
||||
}
|
||||
}
|
||||
return list
|
||||
}
|
||||
|
||||
// Refresh reloads configuration cache from database via the Loader.
|
||||
func Refresh(ctx context.Context, configType string, key string, loader Loader) error {
|
||||
if configType == "" {
|
||||
return errors.New("type is required")
|
||||
}
|
||||
|
||||
if key != "" {
|
||||
// Single key refresh: first fetch latest value from database
|
||||
item, err := loader.LoadOne(ctx, configType, key)
|
||||
if err != nil {
|
||||
if errors.Is(err, ErrNotFound) {
|
||||
Delete(configType, key)
|
||||
return nil
|
||||
}
|
||||
return err
|
||||
}
|
||||
Set(item)
|
||||
return nil
|
||||
}
|
||||
|
||||
// All keys refresh: load all of that type from database first, then replace cache
|
||||
items, err := loader.LoadAll(ctx, configType)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
UpdateTypeItems(configType, items)
|
||||
return nil
|
||||
}
|
||||
|
||||
// ResetForTest clears the local store and locks.
|
||||
func ResetForTest() {
|
||||
writeLocksMu.Lock()
|
||||
writeLocks = make(map[string]*sync.Mutex)
|
||||
writeLocksMu.Unlock()
|
||||
managerCache.InvalidateAll()
|
||||
}
|
||||
@@ -0,0 +1,254 @@
|
||||
// Copyright 2025 linux.do
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package config 负责应用配置的加载、解析与环境变量覆盖。
|
||||
package config
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"flag"
|
||||
"log"
|
||||
"os"
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
"github.com/spf13/viper"
|
||||
)
|
||||
|
||||
// 默认队列优先级
|
||||
const (
|
||||
webhookQueuePriority = 10
|
||||
whitelistQueuePriority = 5
|
||||
defaultQueuePriority = 3
|
||||
)
|
||||
|
||||
// Config 全局配置单例,初始化后不可变
|
||||
var Config *configModel
|
||||
|
||||
// findConfigPath searches upward for the config file to handle tests running in subdirectories.
|
||||
func findConfigPath(configPath string) string {
|
||||
if _, err := os.Stat(configPath); err == nil {
|
||||
return configPath
|
||||
}
|
||||
dir := "."
|
||||
for i := 0; i < 5; i++ {
|
||||
dir += "/.."
|
||||
path := dir + "/" + configPath
|
||||
if _, err := os.Stat(path); err == nil {
|
||||
return path
|
||||
}
|
||||
}
|
||||
return configPath
|
||||
}
|
||||
|
||||
// isTest checks if the current execution context is within 'go test'.
|
||||
func isTest() bool {
|
||||
if flag.Lookup("test.v") != nil {
|
||||
return true
|
||||
}
|
||||
for _, arg := range os.Args {
|
||||
if strings.HasPrefix(arg, "-test.") || strings.HasSuffix(arg, ".test") {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func init() {
|
||||
// 加载配置文件路径
|
||||
configPath := os.Getenv("CONFIG_PATH")
|
||||
if configPath == "" {
|
||||
configPath = findConfigPath("config.yaml")
|
||||
}
|
||||
|
||||
// 设置配置文件
|
||||
viper.SetConfigFile(configPath)
|
||||
viper.AutomaticEnv()
|
||||
|
||||
// 读取配置文件(可选:找不到文件时使用空默认值 + 环境变量)
|
||||
if err := viper.ReadInConfig(); err != nil {
|
||||
if _, ok := err.(viper.ConfigFileNotFoundError); !ok {
|
||||
// 文件存在但读取/解析失败
|
||||
if _, statErr := os.Stat(configPath); statErr == nil { //nolint:gosec // configPath is loaded from CONFIG_PATH environment variable
|
||||
log.Fatalf("[Config] read config failed: %v\n", err)
|
||||
}
|
||||
}
|
||||
log.Println("[Config] no config file found, using environment variables only")
|
||||
viper.SetConfigType("yaml")
|
||||
if err := viper.ReadConfig(strings.NewReader("")); err != nil {
|
||||
log.Fatalf("[Config] failed to init empty config: %v\n", err)
|
||||
}
|
||||
}
|
||||
|
||||
// 解析配置到结构体
|
||||
var c configModel
|
||||
if err := viper.Unmarshal(&c); err != nil {
|
||||
log.Fatalf("[Config] parse config failed: %v\n", err)
|
||||
}
|
||||
|
||||
applyDefaults(&c)
|
||||
|
||||
// 环境变量覆盖(优先级高于 config.yaml)
|
||||
applyEnvOverrides(&c)
|
||||
applyDefaults(&c)
|
||||
|
||||
// Disable standard DB/Redis initializations during tests to prevent connection attempts.
|
||||
if isTest() {
|
||||
c.Database.Enabled = false
|
||||
c.Database.SQLitePath = ":memory:"
|
||||
c.Redis.Enabled = false
|
||||
c.ClickHouse.Enabled = false
|
||||
}
|
||||
|
||||
// 设置全局配置
|
||||
Config = &c
|
||||
|
||||
// 打印配置
|
||||
printConfig(&c)
|
||||
}
|
||||
|
||||
func applyDefaults(c *configModel) {
|
||||
if c.App.SessionAge <= 0 {
|
||||
c.App.SessionAge = 86400
|
||||
}
|
||||
if c.Otel.TracerName == "" {
|
||||
c.Otel.TracerName = "github.com/Rain-kl/Wavelet"
|
||||
}
|
||||
}
|
||||
|
||||
// ─── 环境变量覆盖层 ────────────────────────────────────────────────────────────
|
||||
// 环境变量优先级高于 config.yaml,未设置则保留 yaml 中的值。
|
||||
|
||||
func envStr(key, fallback string) string {
|
||||
if v, ok := os.LookupEnv(key); ok {
|
||||
return v
|
||||
}
|
||||
return fallback
|
||||
}
|
||||
|
||||
func envInt(key string, fallback int) int {
|
||||
if v, ok := os.LookupEnv(key); ok {
|
||||
if n, err := strconv.Atoi(v); err == nil {
|
||||
return n
|
||||
}
|
||||
}
|
||||
return fallback
|
||||
}
|
||||
|
||||
func envInt64(key string, fallback int64) int64 {
|
||||
if v, ok := os.LookupEnv(key); ok {
|
||||
if n, err := strconv.ParseInt(v, 10, 64); err == nil {
|
||||
return n
|
||||
}
|
||||
}
|
||||
return fallback
|
||||
}
|
||||
|
||||
func envFloat64(key string, fallback float64) float64 {
|
||||
if v, ok := os.LookupEnv(key); ok {
|
||||
if n, err := strconv.ParseFloat(v, 64); err == nil {
|
||||
return n
|
||||
}
|
||||
}
|
||||
return fallback
|
||||
}
|
||||
|
||||
func envBool(key string, fallback bool) bool {
|
||||
if v, ok := os.LookupEnv(key); ok {
|
||||
if b, err := strconv.ParseBool(v); err == nil {
|
||||
return b
|
||||
}
|
||||
}
|
||||
return fallback
|
||||
}
|
||||
|
||||
// applyEnvOverrides 将环境变量值覆盖到配置结构体上(仅当环境变量已设置时生效)
|
||||
func applyEnvOverrides(c *configModel) {
|
||||
// ─── App ───
|
||||
c.App.AppName = envStr("APP_NAME", c.App.AppName)
|
||||
c.App.Env = envStr("APP_ENV", c.App.Env)
|
||||
c.App.Addr = envStr("APP_ADDR", c.App.Addr)
|
||||
c.App.NodeID = envInt64("APP_NODE_ID", c.App.NodeID)
|
||||
c.App.APIPrefix = envStr("APP_API_PREFIX", c.App.APIPrefix)
|
||||
c.App.GracefulShutdownTimeout = envInt("APP_GRACEFUL_SHUTDOWN_TIMEOUT", c.App.GracefulShutdownTimeout)
|
||||
c.App.SessionCookieName = envStr("APP_SESSION_COOKIE_NAME", c.App.SessionCookieName)
|
||||
c.App.SessionSecret = envStr("APP_SESSION_SECRET", c.App.SessionSecret)
|
||||
c.App.SessionDomain = envStr("APP_SESSION_DOMAIN", c.App.SessionDomain)
|
||||
c.App.SessionAge = envInt("APP_SESSION_AGE", c.App.SessionAge)
|
||||
c.App.SessionHTTPOnly = envBool("APP_SESSION_HTTP_ONLY", c.App.SessionHTTPOnly)
|
||||
c.App.SessionSecure = envBool("APP_SESSION_SECURE", c.App.SessionSecure)
|
||||
|
||||
// ─── Database ───
|
||||
c.Database.Host = envStr("DB_HOST", c.Database.Host)
|
||||
c.Database.Port = envInt("DB_PORT", c.Database.Port)
|
||||
c.Database.Username = envStr("DB_USERNAME", c.Database.Username)
|
||||
c.Database.Password = envStr("DB_PASSWORD", c.Database.Password)
|
||||
c.Database.Database = envStr("DB_NAME", c.Database.Database)
|
||||
c.Database.SSLMode = envStr("DB_SSL_MODE", c.Database.SSLMode)
|
||||
c.Database.TimeZone = envStr("DB_TIMEZONE", c.Database.TimeZone)
|
||||
c.Database.LogLevel = envStr("DB_LOG_LEVEL", c.Database.LogLevel)
|
||||
c.Database.MaxIdleConn = envInt("DB_MAX_IDLE_CONN", c.Database.MaxIdleConn)
|
||||
c.Database.MaxOpenConn = envInt("DB_MAX_OPEN_CONN", c.Database.MaxOpenConn)
|
||||
// 当 DB_HOST 环境变量已设置时自动启用数据库
|
||||
if _, ok := os.LookupEnv("DB_HOST"); ok {
|
||||
c.Database.Enabled = true
|
||||
}
|
||||
c.Database.Enabled = envBool("DB_ENABLED", c.Database.Enabled)
|
||||
c.Database.SQLitePath = envStr("SQLITE_PATH", c.Database.SQLitePath)
|
||||
|
||||
// ─── Redis ───
|
||||
if v, ok := os.LookupEnv("REDIS_ADDR"); ok {
|
||||
c.Redis.Addrs = []string{v}
|
||||
c.Redis.Enabled = true // 当 REDIS_ADDR 已设置时自动启用
|
||||
}
|
||||
c.Redis.Enabled = envBool("REDIS_ENABLED", c.Redis.Enabled)
|
||||
c.Redis.Username = envStr("REDIS_USERNAME", c.Redis.Username)
|
||||
c.Redis.Password = envStr("REDIS_PASSWORD", c.Redis.Password)
|
||||
c.Redis.DB = envInt("REDIS_DB", c.Redis.DB)
|
||||
c.Redis.KeyPrefix = envStr("REDIS_KEY_PREFIX", c.Redis.KeyPrefix)
|
||||
c.Redis.PoolSize = envInt("REDIS_POOL_SIZE", c.Redis.PoolSize)
|
||||
c.Redis.MaintNotifications = envBool("REDIS_MAINT_NOTIFICATIONS", c.Redis.MaintNotifications)
|
||||
|
||||
// ─── ClickHouse ───
|
||||
if v, ok := os.LookupEnv("CLICKHOUSE_HOST"); ok {
|
||||
c.ClickHouse.Hosts = []string{v}
|
||||
c.ClickHouse.Enabled = true
|
||||
}
|
||||
c.ClickHouse.Enabled = envBool("CLICKHOUSE_ENABLED", c.ClickHouse.Enabled)
|
||||
c.ClickHouse.Username = envStr("CLICKHOUSE_USERNAME", c.ClickHouse.Username)
|
||||
c.ClickHouse.Password = envStr("CLICKHOUSE_PASSWORD", c.ClickHouse.Password)
|
||||
c.ClickHouse.Database = envStr("CLICKHOUSE_NAME", c.ClickHouse.Database)
|
||||
|
||||
// ─── Log ───
|
||||
c.Log.Level = envStr("LOG_LEVEL", c.Log.Level)
|
||||
c.Log.Format = envStr("LOG_FORMAT", c.Log.Format)
|
||||
c.Log.Output = envStr("LOG_OUTPUT", c.Log.Output)
|
||||
|
||||
// ─── OTel ───
|
||||
c.Otel.SamplingRate = envFloat64("OTEL_SAMPLING_RATE", c.Otel.SamplingRate)
|
||||
c.Otel.TracerName = envStr("OTEL_TRACER_NAME", c.Otel.TracerName)
|
||||
|
||||
// ─── Worker ───
|
||||
c.Worker.Concurrency = envInt("WORKER_CONCURRENCY", c.Worker.Concurrency)
|
||||
c.Worker.StrictPriority = envBool("WORKER_STRICT_PRIORITY", c.Worker.StrictPriority)
|
||||
|
||||
// 无 yaml 且无环境变量时,使用代码级默认队列
|
||||
if len(c.Worker.Queues) == 0 {
|
||||
c.Worker.Queues = []QueueConfig{
|
||||
{Name: "webhook", Priority: webhookQueuePriority},
|
||||
{Name: "whitelist_only", Priority: whitelistQueuePriority},
|
||||
{Name: "default", Priority: defaultQueuePriority},
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// printConfig 打印配置内容
|
||||
func printConfig(c *configModel) {
|
||||
configJSON, err := json.MarshalIndent(c, "", " ")
|
||||
if err != nil {
|
||||
log.Printf("[Config] failed to marshal config: %v\n", err)
|
||||
return
|
||||
}
|
||||
log.Printf("[Config] loaded configuration:\n%s\n", string(configJSON))
|
||||
}
|
||||
@@ -0,0 +1,17 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package config
|
||||
|
||||
import "testing"
|
||||
|
||||
func TestApplyEnvOverridesRedisMaintNotifications(t *testing.T) {
|
||||
t.Setenv("REDIS_MAINT_NOTIFICATIONS", "true")
|
||||
|
||||
cfg := &configModel{}
|
||||
applyEnvOverrides(cfg)
|
||||
|
||||
if !cfg.Redis.MaintNotifications {
|
||||
t.Fatal("REDIS_MAINT_NOTIFICATIONS=true was not applied")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,142 @@
|
||||
// Copyright 2025 linux.do
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package config
|
||||
|
||||
import "time"
|
||||
|
||||
type configModel struct {
|
||||
App appConfig `mapstructure:"app"`
|
||||
Database databaseConfig `mapstructure:"database"`
|
||||
Redis redisConfig `mapstructure:"redis"`
|
||||
Log logConfig `mapstructure:"log"`
|
||||
Scheduler schedulerConfig `mapstructure:"scheduler"`
|
||||
Worker workerConfig `mapstructure:"worker"`
|
||||
ClickHouse clickHouseConfig `mapstructure:"clickhouse"`
|
||||
Otel otelConfig `mapstructure:"otel"`
|
||||
}
|
||||
|
||||
// appConfig 应用基本配置
|
||||
type appConfig struct {
|
||||
AppName string `mapstructure:"app_name"`
|
||||
Env string `mapstructure:"env"`
|
||||
Addr string `mapstructure:"addr"`
|
||||
NodeID int64 `mapstructure:"node_id"`
|
||||
APIPrefix string `mapstructure:"api_prefix"`
|
||||
GracefulShutdownTimeout int `mapstructure:"graceful_shutdown_timeout"`
|
||||
SessionCookieName string `mapstructure:"session_cookie_name"`
|
||||
SessionSecret string `mapstructure:"session_secret"`
|
||||
SessionDomain string `mapstructure:"session_domain"`
|
||||
SessionAge int `mapstructure:"session_age"`
|
||||
SessionHTTPOnly bool `mapstructure:"session_http_only"`
|
||||
SessionSecure bool `mapstructure:"session_secure"`
|
||||
}
|
||||
|
||||
// IsProduction 检查当前环境是否为生产环境
|
||||
func (a *appConfig) IsProduction() bool {
|
||||
return a.Env == "production"
|
||||
}
|
||||
|
||||
// databaseConfig 数据库配置
|
||||
type databaseConfig struct {
|
||||
Enabled bool `mapstructure:"enabled"`
|
||||
SQLitePath string `mapstructure:"sqlite_path"` // PostgreSQL 禁用时的 SQLite 文件路径
|
||||
Host string `mapstructure:"host"`
|
||||
Port int `mapstructure:"port"`
|
||||
Username string `mapstructure:"username"`
|
||||
Password string `mapstructure:"password"`
|
||||
Database string `mapstructure:"database"`
|
||||
MaxIdleConn int `mapstructure:"max_idle_conn"`
|
||||
MaxOpenConn int `mapstructure:"max_open_conn"`
|
||||
ConnMaxLifetime int `mapstructure:"conn_max_lifetime"`
|
||||
ConnMaxIdleTime int `mapstructure:"conn_max_idle_time"`
|
||||
LogLevel string `mapstructure:"log_level"`
|
||||
SSLMode string `mapstructure:"ssl_mode"`
|
||||
TimeZone string `mapstructure:"time_zone"`
|
||||
ApplicationName string `mapstructure:"application_name"`
|
||||
SearchPath string `mapstructure:"search_path"`
|
||||
PreferSimpleProtocol bool `mapstructure:"prefer_simple_protocol"`
|
||||
StatementCacheCapacity int `mapstructure:"statement_cache_capacity"`
|
||||
DefaultQueryExecMode string `mapstructure:"default_query_exec_mode"`
|
||||
Replicas []databaseReplicaConfig `mapstructure:"replicas"`
|
||||
SlowThreshold time.Duration `mapstructure:"slow_threshold"`
|
||||
}
|
||||
|
||||
// databaseReplicaConfig 只读副本配置
|
||||
type databaseReplicaConfig struct {
|
||||
Host string `mapstructure:"host"`
|
||||
Port int `mapstructure:"port"`
|
||||
Username string `mapstructure:"username"`
|
||||
Password string `mapstructure:"password"`
|
||||
}
|
||||
|
||||
// clickhouse 配置
|
||||
type clickHouseConfig struct {
|
||||
Enabled bool `mapstructure:"enabled"`
|
||||
Hosts []string `mapstructure:"hosts"`
|
||||
Username string `mapstructure:"username"`
|
||||
Password string `mapstructure:"password"`
|
||||
Database string `mapstructure:"database"`
|
||||
MaxIdleConn int `mapstructure:"max_idle_conn"`
|
||||
MaxOpenConn int `mapstructure:"max_open_conn"`
|
||||
ConnMaxLifetime int `mapstructure:"conn_max_lifetime"`
|
||||
DialTimeout int `mapstructure:"dial_timeout"`
|
||||
BlockBufferSize uint8 `mapstructure:"block_buffer_size"`
|
||||
}
|
||||
|
||||
// redisConfig Redis配置
|
||||
type redisConfig struct {
|
||||
Enabled bool `mapstructure:"enabled"`
|
||||
Addrs []string `mapstructure:"addrs"`
|
||||
Username string `mapstructure:"username"`
|
||||
Password string `mapstructure:"password"`
|
||||
DB int `mapstructure:"db"`
|
||||
ClusterMode bool `mapstructure:"cluster_mode"`
|
||||
MasterName string `mapstructure:"master_name"`
|
||||
KeyPrefix string `mapstructure:"key_prefix"`
|
||||
PoolSize int `mapstructure:"pool_size"`
|
||||
MinIdleConn int `mapstructure:"min_idle_conn"`
|
||||
DialTimeout int `mapstructure:"dial_timeout"`
|
||||
ReadTimeout int `mapstructure:"read_timeout"`
|
||||
WriteTimeout int `mapstructure:"write_timeout"`
|
||||
MaxRetries int `mapstructure:"max_retries"`
|
||||
PoolTimeout int `mapstructure:"pool_timeout"`
|
||||
ConnMaxIdleTime int `mapstructure:"conn_max_idle_time"`
|
||||
MaintNotifications bool `mapstructure:"maint_notifications"`
|
||||
}
|
||||
|
||||
// logConfig 日志配置
|
||||
type logConfig struct {
|
||||
Level string `mapstructure:"level"`
|
||||
Format string `mapstructure:"format"`
|
||||
Output string `mapstructure:"output"`
|
||||
FilePath string `mapstructure:"file_path"`
|
||||
MaxSize int `mapstructure:"max_size"`
|
||||
MaxAge int `mapstructure:"max_age"`
|
||||
MaxBackups int `mapstructure:"max_backups"`
|
||||
Compress bool `mapstructure:"compress"`
|
||||
}
|
||||
|
||||
// schedulerConfig 定时任务配置
|
||||
type schedulerConfig struct {
|
||||
}
|
||||
|
||||
// workerConfig 工作配置
|
||||
type workerConfig struct {
|
||||
Concurrency int `mapstructure:"concurrency"`
|
||||
StrictPriority bool `mapstructure:"strict_priority"`
|
||||
Queues []QueueConfig `mapstructure:"queues"`
|
||||
}
|
||||
|
||||
// QueueConfig 队列配置
|
||||
type QueueConfig struct {
|
||||
Name string `mapstructure:"name"`
|
||||
Priority int `mapstructure:"priority"`
|
||||
}
|
||||
|
||||
// otelConfig OpenTelemetry 配置
|
||||
type otelConfig struct {
|
||||
SamplingRate float64 `mapstructure:"sampling_rate"`
|
||||
TracerName string `mapstructure:"tracer_name"`
|
||||
}
|
||||
@@ -0,0 +1,106 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package httppool manages shared, optimized HTTP transports to reuse TCP connections.
|
||||
package httppool
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/tls"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp"
|
||||
)
|
||||
|
||||
const (
|
||||
dialTimeout = 30 * time.Second
|
||||
dialKeepAlive = 30 * time.Second
|
||||
maxIdleConns = 200
|
||||
maxIdleConnsPerHost = 32
|
||||
idleConnTimeout = 90 * time.Second
|
||||
tlsHandshakeTimeout = 10 * time.Second
|
||||
expectContinueTimeout = 1 * time.Second
|
||||
tlsSessionCacheSize = 100
|
||||
)
|
||||
|
||||
var (
|
||||
defaultTransport http.RoundTripper
|
||||
once sync.Once
|
||||
)
|
||||
|
||||
// TransportOptions configures the request-specific parts of a pooled HTTP
|
||||
// transport. Pool sizes and timeout defaults remain managed by this package.
|
||||
// A nil Proxy explicitly disables proxy use.
|
||||
type TransportOptions struct {
|
||||
Proxy func(*http.Request) (*url.URL, error)
|
||||
DialContext func(context.Context, string, string) (net.Conn, error)
|
||||
TLSClientConfig *tls.Config
|
||||
ResponseHeaderTimeout time.Duration
|
||||
TraceFilter func(*http.Request) bool
|
||||
}
|
||||
|
||||
// NewTransport returns an independently configurable pooled transport wrapped
|
||||
// with OTel instrumentation. The supplied TLS configuration is cloned before
|
||||
// use so later caller mutations cannot change an active transport.
|
||||
func NewTransport(options TransportOptions) http.RoundTripper {
|
||||
dialContext := options.DialContext
|
||||
if dialContext == nil {
|
||||
dialContext = (&net.Dialer{
|
||||
Timeout: dialTimeout,
|
||||
KeepAlive: dialKeepAlive,
|
||||
}).DialContext
|
||||
}
|
||||
|
||||
tlsConfig := options.TLSClientConfig
|
||||
if tlsConfig == nil {
|
||||
tlsConfig = &tls.Config{}
|
||||
} else {
|
||||
tlsConfig = tlsConfig.Clone()
|
||||
}
|
||||
if tlsConfig.ClientSessionCache == nil {
|
||||
tlsConfig.ClientSessionCache = tls.NewLRUClientSessionCache(tlsSessionCacheSize)
|
||||
}
|
||||
|
||||
transport := &http.Transport{
|
||||
Proxy: options.Proxy,
|
||||
DialContext: dialContext,
|
||||
ForceAttemptHTTP2: true,
|
||||
MaxIdleConns: maxIdleConns,
|
||||
MaxIdleConnsPerHost: maxIdleConnsPerHost,
|
||||
IdleConnTimeout: idleConnTimeout,
|
||||
TLSHandshakeTimeout: tlsHandshakeTimeout,
|
||||
ResponseHeaderTimeout: options.ResponseHeaderTimeout,
|
||||
ExpectContinueTimeout: expectContinueTimeout,
|
||||
TLSClientConfig: tlsConfig,
|
||||
}
|
||||
otelOptions := make([]otelhttp.Option, 0, 1)
|
||||
if options.TraceFilter != nil {
|
||||
otelOptions = append(otelOptions, otelhttp.WithFilter(options.TraceFilter))
|
||||
}
|
||||
return otelhttp.NewTransport(transport, otelOptions...)
|
||||
}
|
||||
|
||||
// DefaultTransport returns a globally shared, optimized http.RoundTripper
|
||||
// with OTel instrumentation. It maintains a pool of idle TCP connections
|
||||
// across hosts.
|
||||
func DefaultTransport() http.RoundTripper {
|
||||
once.Do(func() {
|
||||
defaultTransport = NewTransport(TransportOptions{
|
||||
Proxy: http.ProxyFromEnvironment,
|
||||
})
|
||||
})
|
||||
return defaultTransport
|
||||
}
|
||||
|
||||
// NewClient returns a new http.Client that shares the global connection pool
|
||||
// but has its own timeout configuration.
|
||||
func NewClient(timeout time.Duration) *http.Client {
|
||||
return &http.Client{
|
||||
Timeout: timeout,
|
||||
Transport: DefaultTransport(),
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,100 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package httppool
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/tls"
|
||||
"io"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func TestDefaultTransport(t *testing.T) {
|
||||
tr1 := DefaultTransport()
|
||||
if tr1 == nil {
|
||||
t.Fatal("DefaultTransport() returned nil")
|
||||
}
|
||||
|
||||
tr2 := DefaultTransport()
|
||||
if tr1 != tr2 {
|
||||
t.Error("DefaultTransport() did not return a singleton instance")
|
||||
}
|
||||
}
|
||||
|
||||
func TestNewClient(t *testing.T) {
|
||||
timeout := 15 * time.Second
|
||||
client := NewClient(timeout)
|
||||
if client == nil {
|
||||
t.Fatal("NewClient() returned nil")
|
||||
}
|
||||
|
||||
if client.Timeout != timeout {
|
||||
t.Errorf("NewClient() timeout = %v, want %v", client.Timeout, timeout)
|
||||
}
|
||||
|
||||
if client.Transport != DefaultTransport() {
|
||||
t.Error("NewClient() is not configured with the default transport")
|
||||
}
|
||||
}
|
||||
|
||||
func TestNewTransportUsesConfiguredDirectDialer(t *testing.T) {
|
||||
server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, _ *http.Request) {
|
||||
_, _ = writer.Write([]byte("ok"))
|
||||
}))
|
||||
t.Cleanup(server.Close)
|
||||
|
||||
var dialedAddress string
|
||||
dialer := &net.Dialer{}
|
||||
transport := NewTransport(TransportOptions{
|
||||
Proxy: nil,
|
||||
DialContext: func(ctx context.Context, network string, address string) (net.Conn, error) {
|
||||
dialedAddress = address
|
||||
return dialer.DialContext(ctx, network, server.Listener.Addr().String())
|
||||
},
|
||||
})
|
||||
client := &http.Client{Transport: transport}
|
||||
t.Cleanup(client.CloseIdleConnections)
|
||||
|
||||
request, err := http.NewRequestWithContext(t.Context(), http.MethodGet, "http://artifact.example/site.zip", nil)
|
||||
if err != nil {
|
||||
t.Fatalf("NewRequestWithContext() error = %v", err)
|
||||
}
|
||||
response, err := client.Do(request)
|
||||
if err != nil {
|
||||
t.Fatalf("client.Do() error = %v", err)
|
||||
}
|
||||
defer func() { _ = response.Body.Close() }()
|
||||
if _, err := io.ReadAll(response.Body); err != nil {
|
||||
t.Fatalf("ReadAll() error = %v", err)
|
||||
}
|
||||
if dialedAddress != "artifact.example:80" {
|
||||
t.Fatalf("DialContext address = %q, want direct target", dialedAddress)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNewTransportClonesTLSConfig(t *testing.T) {
|
||||
server := httptest.NewTLSServer(http.HandlerFunc(func(writer http.ResponseWriter, _ *http.Request) {
|
||||
writer.WriteHeader(http.StatusNoContent)
|
||||
}))
|
||||
t.Cleanup(server.Close)
|
||||
|
||||
tlsConfig := &tls.Config{InsecureSkipVerify: true} //nolint:gosec // test-only self-signed server
|
||||
client := &http.Client{Transport: NewTransport(TransportOptions{TLSClientConfig: tlsConfig})}
|
||||
t.Cleanup(client.CloseIdleConnections)
|
||||
tlsConfig.InsecureSkipVerify = false
|
||||
|
||||
request, err := http.NewRequestWithContext(t.Context(), http.MethodGet, server.URL, nil)
|
||||
if err != nil {
|
||||
t.Fatalf("NewRequestWithContext() error = %v", err)
|
||||
}
|
||||
response, err := client.Do(request)
|
||||
if err != nil {
|
||||
t.Fatalf("client.Do() error = %v", err)
|
||||
}
|
||||
_ = response.Body.Close()
|
||||
}
|
||||
@@ -0,0 +1,46 @@
|
||||
// Copyright 2025 linux.do
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package idgen 提供分布式 ID 生成器
|
||||
package idgen
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"log"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/backend/pkg/config"
|
||||
"github.com/bwmarrin/snowflake"
|
||||
)
|
||||
|
||||
// 2025-12-01 00:00:00 UTC 的毫秒时间戳
|
||||
const epoch int64 = 1764547200000
|
||||
|
||||
const maxNegativeIDRetries = 3
|
||||
|
||||
var node *snowflake.Node
|
||||
|
||||
func init() {
|
||||
snowflake.Epoch = epoch
|
||||
|
||||
nodeID := config.Config.App.NodeID
|
||||
var err error
|
||||
node, err = snowflake.NewNode(nodeID)
|
||||
if err != nil {
|
||||
log.Fatalf("[Snowflake] init failed: %v\n", err)
|
||||
}
|
||||
log.Printf("[Snowflake] initialized with node ID: %d, epoch: 2025-12-01\n", nodeID)
|
||||
}
|
||||
|
||||
// NextUint64ID 生成下一个分布式唯一 ID。
|
||||
// 理论上不应出现负值;若出现则最多重试 maxNegativeIDRetries 次,仍失败则 panic。
|
||||
func NextUint64ID() uint64 {
|
||||
for attempt := 1; attempt <= maxNegativeIDRetries; attempt++ {
|
||||
id := node.Generate().Int64()
|
||||
if id >= 0 {
|
||||
return uint64(id)
|
||||
}
|
||||
log.Printf("[Snowflake] generated negative ID: %d (attempt %d/%d)", id, attempt, maxNegativeIDRetries)
|
||||
}
|
||||
panic(fmt.Sprintf("[Snowflake] generated negative ID after %d attempts", maxNegativeIDRetries))
|
||||
}
|
||||
@@ -0,0 +1,15 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package idgen
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
)
|
||||
|
||||
func TestNextUint64ID(t *testing.T) {
|
||||
id := NextUint64ID()
|
||||
assert.NotZero(t, id)
|
||||
}
|
||||
@@ -0,0 +1,9 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package logger 提供结构化日志封装
|
||||
package logger
|
||||
|
||||
const (
|
||||
errCreateLogFileDirFailed = "[Logger] create log file dir err: %w"
|
||||
)
|
||||
@@ -0,0 +1,102 @@
|
||||
// Copyright 2025 linux.do
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package logger
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"log"
|
||||
|
||||
"github.com/uptrace/opentelemetry-go-extra/otelzap"
|
||||
"go.uber.org/zap"
|
||||
"go.uber.org/zap/zapcore"
|
||||
)
|
||||
|
||||
// Config represents the logging configuration.
|
||||
type Config struct {
|
||||
Level string
|
||||
Format string
|
||||
Output string
|
||||
FilePath string
|
||||
MaxSize int
|
||||
MaxAge int
|
||||
MaxBackups int
|
||||
Compress bool
|
||||
}
|
||||
|
||||
var logger *otelzap.Logger
|
||||
|
||||
// ringBufferCapacity 环形缓冲区容量
|
||||
const ringBufferCapacity = 5000
|
||||
|
||||
// GlobalRingBuffer 全局日志环形缓冲区,供 Admin 日志查询和 WebSocket 推送使用
|
||||
var GlobalRingBuffer *LogRingBuffer
|
||||
|
||||
func doInit(cfg Config) {
|
||||
logWriter, err := getLogWriterForConfig(cfg)
|
||||
if err != nil {
|
||||
log.Fatalf("[Logger] get log writer err: %v\n", err)
|
||||
}
|
||||
|
||||
// 初始化 ring buffer(保留最近 5000 行日志),如果是多次调用 Init,不需要重复创建 GlobalRingBuffer
|
||||
if GlobalRingBuffer == nil {
|
||||
GlobalRingBuffer = NewLogRingBuffer(ringBufferCapacity)
|
||||
}
|
||||
|
||||
// 使用 multi writer 同时写入原始输出和 ring buffer
|
||||
multiWriter := zapcore.NewMultiWriteSyncer(
|
||||
logWriter,
|
||||
zapcore.AddSync(GlobalRingBuffer),
|
||||
)
|
||||
|
||||
zapLogger := zap.New(
|
||||
zapcore.NewCore(getEncoderForConfig(cfg), multiWriter, getLogLevelForConfig(cfg)),
|
||||
zap.AddCaller(),
|
||||
zap.AddCallerSkip(1),
|
||||
)
|
||||
logger = otelzap.New(
|
||||
zapLogger,
|
||||
otelzap.WithMinLevel(zapLogger.Level()),
|
||||
)
|
||||
}
|
||||
|
||||
func init() {
|
||||
// 默认使用 console stdout INFO 日志输出,避免在 Init 前或测试中发生空指针崩溃
|
||||
defaultCfg := Config{
|
||||
Level: "info",
|
||||
Format: "console",
|
||||
Output: "stdout",
|
||||
}
|
||||
doInit(defaultCfg)
|
||||
}
|
||||
|
||||
// Init initializes the logger with a custom configuration.
|
||||
func Init(cfg Config) {
|
||||
doInit(cfg)
|
||||
}
|
||||
|
||||
// DebugF 输出 Debug 级别日志
|
||||
func DebugF(ctx context.Context, format string, args ...interface{}) {
|
||||
msg := fmt.Sprintf(format, args...)
|
||||
logger.Ctx(ctx).Debug(msg, getTraceIDFields(ctx)...)
|
||||
}
|
||||
|
||||
// InfoF 输出 Info 级别日志
|
||||
func InfoF(ctx context.Context, format string, args ...interface{}) {
|
||||
msg := fmt.Sprintf(format, args...)
|
||||
logger.Ctx(ctx).Info(msg, getTraceIDFields(ctx)...)
|
||||
}
|
||||
|
||||
// WarnF 输出 Warn 级别日志
|
||||
func WarnF(ctx context.Context, format string, args ...interface{}) {
|
||||
msg := fmt.Sprintf(format, args...)
|
||||
logger.Ctx(ctx).Warn(msg, getTraceIDFields(ctx)...)
|
||||
}
|
||||
|
||||
// ErrorF 输出 Error 级别日志
|
||||
func ErrorF(ctx context.Context, format string, args ...interface{}) {
|
||||
msg := fmt.Sprintf(format, args...)
|
||||
logger.Ctx(ctx).Error(msg, getTraceIDFields(ctx)...)
|
||||
}
|
||||
@@ -0,0 +1,175 @@
|
||||
// Copyright 2025 linux.do
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package logger
|
||||
|
||||
import (
|
||||
"io"
|
||||
"sync"
|
||||
)
|
||||
|
||||
// LogEntry 日志条目,对应 ring buffer 中的一行日志
|
||||
type LogEntry struct {
|
||||
Index int `json:"index"` // 全局递增序号
|
||||
Data string `json:"data"` // 一行日志原文(含换行符)
|
||||
}
|
||||
|
||||
// LogRingBuffer 固定容量的环形缓冲区,存储最近的日志行
|
||||
// 支持:追加日志、按 cursor 分页查询、订阅实时推送
|
||||
type LogRingBuffer struct {
|
||||
mu sync.RWMutex
|
||||
entries []LogEntry
|
||||
cap int
|
||||
head int // 下一条写入的位置
|
||||
count int // 当前条目数
|
||||
seq int // 全局递增序号
|
||||
|
||||
subscribers map[chan LogEntry]struct{}
|
||||
subMu sync.RWMutex
|
||||
}
|
||||
|
||||
// NewLogRingBuffer 创建指定容量的日志环形缓冲区
|
||||
func NewLogRingBuffer(capacity int) *LogRingBuffer {
|
||||
return &LogRingBuffer{
|
||||
entries: make([]LogEntry, capacity),
|
||||
cap: capacity,
|
||||
subscribers: make(map[chan LogEntry]struct{}),
|
||||
}
|
||||
}
|
||||
|
||||
// Write 实现 io.Writer 接口,供 zapcore.WriteSyncer 调用
|
||||
// 按 '\n' 分割为独立行写入 ring buffer
|
||||
func (r *LogRingBuffer) Write(p []byte) (int, error) {
|
||||
if len(p) == 0 {
|
||||
return 0, nil
|
||||
}
|
||||
|
||||
data := string(p)
|
||||
start := 0
|
||||
for i := 0; i < len(data); i++ {
|
||||
if data[i] == '\n' {
|
||||
line := data[start:i]
|
||||
start = i + 1
|
||||
if len(line) > 0 {
|
||||
r.appendLine(line)
|
||||
}
|
||||
}
|
||||
}
|
||||
// 处理最后一行(没有换行符结尾的情况)
|
||||
if start < len(data) && len(data[start:]) > 0 {
|
||||
r.appendLine(data[start:])
|
||||
}
|
||||
|
||||
return len(p), nil
|
||||
}
|
||||
|
||||
// Sync 实现 zapcore.WriteSyncer 接口
|
||||
func (r *LogRingBuffer) Sync() error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// appendLine 追加一行日志到 ring buffer 并通知订阅者
|
||||
func (r *LogRingBuffer) appendLine(line string) {
|
||||
r.mu.Lock()
|
||||
entry := LogEntry{
|
||||
Index: r.seq,
|
||||
Data: line,
|
||||
}
|
||||
r.entries[r.head] = entry
|
||||
r.head = (r.head + 1) % r.cap
|
||||
if r.count < r.cap {
|
||||
r.count++
|
||||
}
|
||||
r.seq++
|
||||
r.mu.Unlock()
|
||||
|
||||
// 异步通知订阅者
|
||||
r.subMu.RLock()
|
||||
for ch := range r.subscribers {
|
||||
select {
|
||||
case ch <- entry:
|
||||
default:
|
||||
// 订阅者消费太慢,丢弃(避免阻塞日志写入)
|
||||
}
|
||||
}
|
||||
r.subMu.RUnlock()
|
||||
}
|
||||
|
||||
// Query 查询历史日志
|
||||
// cursor=0 表示查询最新日志,cursor>0 表示查询 index < cursor 的更早日志
|
||||
// limit 为返回条数上限
|
||||
// 返回日志条目(按 index 升序)和是否有更早的日志
|
||||
func (r *LogRingBuffer) Query(cursor int, limit int) ([]LogEntry, bool) {
|
||||
r.mu.RLock()
|
||||
defer r.mu.RUnlock()
|
||||
|
||||
if r.count == 0 {
|
||||
return nil, false
|
||||
}
|
||||
|
||||
// 计算 ring buffer 中有效条目的范围
|
||||
// oldest index in ring: head - count (wrapping)
|
||||
oldestPos := (r.head - r.count + r.cap) % r.cap
|
||||
|
||||
// 将 ring buffer 中的有效条目按顺序收集
|
||||
ordered := make([]LogEntry, 0, r.count)
|
||||
for i := 0; i < r.count; i++ {
|
||||
pos := (oldestPos + i) % r.cap
|
||||
ordered = append(ordered, r.entries[pos])
|
||||
}
|
||||
|
||||
if cursor == 0 {
|
||||
// 查询最新日志:返回最后 limit 条
|
||||
if len(ordered) <= limit {
|
||||
return ordered, false
|
||||
}
|
||||
return ordered[len(ordered)-limit:], true
|
||||
}
|
||||
|
||||
// 查询 index < cursor 的更早日志
|
||||
// 找到 index < cursor 的条目
|
||||
var cut int
|
||||
for cut = len(ordered); cut > 0; cut-- {
|
||||
if ordered[cut-1].Index < cursor {
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
if cut == 0 {
|
||||
return nil, false
|
||||
}
|
||||
|
||||
// 返回 cut 之前的最后 limit 条
|
||||
start := cut - limit
|
||||
if start < 0 {
|
||||
start = 0
|
||||
}
|
||||
|
||||
hasMore := start > 0
|
||||
return ordered[start:cut], hasMore
|
||||
}
|
||||
|
||||
// subscribeChanSize 订阅者 channel 缓冲区大小
|
||||
const subscribeChanSize = 64
|
||||
|
||||
// Subscribe 订阅实时日志推送
|
||||
// 返回一个 channel,调用者应 defer Unsubscribe
|
||||
func (r *LogRingBuffer) Subscribe() chan LogEntry {
|
||||
ch := make(chan LogEntry, subscribeChanSize)
|
||||
r.subMu.Lock()
|
||||
r.subscribers[ch] = struct{}{}
|
||||
r.subMu.Unlock()
|
||||
return ch
|
||||
}
|
||||
|
||||
// Unsubscribe 取消订阅
|
||||
func (r *LogRingBuffer) Unsubscribe(ch chan LogEntry) {
|
||||
r.subMu.Lock()
|
||||
delete(r.subscribers, ch)
|
||||
r.subMu.Unlock()
|
||||
close(ch)
|
||||
}
|
||||
|
||||
// 确保 LogRingBuffer 实现 io.Writer 接口
|
||||
var _ io.Writer = (*LogRingBuffer)(nil)
|
||||
@@ -0,0 +1,192 @@
|
||||
// Copyright 2025 linux.do
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package logger
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
)
|
||||
|
||||
func TestLogRingBuffer_WriteAndQuery(t *testing.T) {
|
||||
rb := NewLogRingBuffer(5)
|
||||
|
||||
// Write some logs
|
||||
_, _ = rb.Write([]byte("line1\nline2\nline3\n"))
|
||||
|
||||
entries, hasMore := rb.Query(0, 10)
|
||||
assert.False(t, hasMore)
|
||||
assert.Equal(t, 3, len(entries))
|
||||
assert.Equal(t, "line1", entries[0].Data)
|
||||
assert.Equal(t, "line2", entries[1].Data)
|
||||
assert.Equal(t, "line3", entries[2].Data)
|
||||
assert.Equal(t, 0, entries[0].Index)
|
||||
assert.Equal(t, 1, entries[1].Index)
|
||||
assert.Equal(t, 2, entries[2].Index)
|
||||
}
|
||||
|
||||
func TestLogRingBuffer_CapacityOverflow(t *testing.T) {
|
||||
rb := NewLogRingBuffer(3)
|
||||
|
||||
_, _ = rb.Write([]byte("a\nb\nc\nd\ne\n"))
|
||||
|
||||
entries, hasMore := rb.Query(0, 10)
|
||||
assert.False(t, hasMore)
|
||||
assert.Equal(t, 3, len(entries))
|
||||
assert.Equal(t, "c", entries[0].Data)
|
||||
assert.Equal(t, "d", entries[1].Data)
|
||||
assert.Equal(t, "e", entries[2].Data)
|
||||
}
|
||||
|
||||
func TestLogRingBuffer_QueryLatest(t *testing.T) {
|
||||
rb := NewLogRingBuffer(10)
|
||||
|
||||
_, _ = rb.Write([]byte("a\nb\nc\nd\ne\n"))
|
||||
|
||||
// Query latest 2
|
||||
entries, hasMore := rb.Query(0, 2)
|
||||
assert.True(t, hasMore)
|
||||
assert.Equal(t, 2, len(entries))
|
||||
assert.Equal(t, "d", entries[0].Data)
|
||||
assert.Equal(t, "e", entries[1].Data)
|
||||
}
|
||||
|
||||
func TestLogRingBuffer_QueryByCursor(t *testing.T) {
|
||||
rb := NewLogRingBuffer(10)
|
||||
|
||||
_, _ = rb.Write([]byte("a\nb\nc\nd\ne\n"))
|
||||
|
||||
// First get all to find indices
|
||||
all, _ := rb.Query(0, 10)
|
||||
assert.Equal(t, 5, len(all))
|
||||
|
||||
// Query entries before index 3
|
||||
entries, hasMore := rb.Query(3, 10)
|
||||
assert.False(t, hasMore)
|
||||
assert.Equal(t, 3, len(entries))
|
||||
assert.Equal(t, "a", entries[0].Data)
|
||||
assert.Equal(t, "b", entries[1].Data)
|
||||
assert.Equal(t, "c", entries[2].Data)
|
||||
}
|
||||
|
||||
func TestLogRingBuffer_QueryByCursorWithLimit(t *testing.T) {
|
||||
rb := NewLogRingBuffer(10)
|
||||
|
||||
_, _ = rb.Write([]byte("a\nb\nc\nd\ne\n"))
|
||||
|
||||
// Query 2 entries before index 4
|
||||
entries, hasMore := rb.Query(4, 2)
|
||||
assert.True(t, hasMore)
|
||||
assert.Equal(t, 2, len(entries))
|
||||
assert.Equal(t, "c", entries[0].Data)
|
||||
assert.Equal(t, "d", entries[1].Data)
|
||||
}
|
||||
|
||||
func TestLogRingBuffer_QueryEmpty(t *testing.T) {
|
||||
rb := NewLogRingBuffer(5)
|
||||
|
||||
entries, hasMore := rb.Query(0, 10)
|
||||
assert.False(t, hasMore)
|
||||
assert.Nil(t, entries)
|
||||
}
|
||||
|
||||
func TestLogRingBuffer_QueryNonExistentCursor(t *testing.T) {
|
||||
rb := NewLogRingBuffer(5)
|
||||
_, _ = rb.Write([]byte("a\nb\n"))
|
||||
|
||||
entries, hasMore := rb.Query(999, 10)
|
||||
assert.False(t, hasMore)
|
||||
assert.Equal(t, 2, len(entries))
|
||||
assert.Equal(t, "a", entries[0].Data)
|
||||
assert.Equal(t, "b", entries[1].Data)
|
||||
}
|
||||
|
||||
func TestLogRingBuffer_Subscribe(t *testing.T) {
|
||||
rb := NewLogRingBuffer(5)
|
||||
|
||||
ch := rb.Subscribe()
|
||||
defer rb.Unsubscribe(ch)
|
||||
|
||||
_, _ = rb.Write([]byte("hello\n"))
|
||||
|
||||
entry := <-ch
|
||||
assert.Equal(t, "hello", entry.Data)
|
||||
assert.Equal(t, 0, entry.Index)
|
||||
}
|
||||
|
||||
func TestLogRingBuffer_SubscribeMultiple(t *testing.T) {
|
||||
rb := NewLogRingBuffer(5)
|
||||
|
||||
ch1 := rb.Subscribe()
|
||||
defer rb.Unsubscribe(ch1)
|
||||
ch2 := rb.Subscribe()
|
||||
defer rb.Unsubscribe(ch2)
|
||||
|
||||
_, _ = rb.Write([]byte("msg\n"))
|
||||
|
||||
e1 := <-ch1
|
||||
e2 := <-ch2
|
||||
assert.Equal(t, "msg", e1.Data)
|
||||
assert.Equal(t, "msg", e2.Data)
|
||||
}
|
||||
|
||||
func TestLogRingBuffer_WriteNoNewline(t *testing.T) {
|
||||
rb := NewLogRingBuffer(5)
|
||||
|
||||
_, _ = rb.Write([]byte("partial"))
|
||||
|
||||
entries, _ := rb.Query(0, 10)
|
||||
assert.Equal(t, 1, len(entries))
|
||||
assert.Equal(t, "partial", entries[0].Data)
|
||||
}
|
||||
|
||||
func TestLogRingBuffer_WriteEmpty(t *testing.T) {
|
||||
rb := NewLogRingBuffer(5)
|
||||
|
||||
n, err := rb.Write([]byte(""))
|
||||
assert.Equal(t, 0, n)
|
||||
assert.NoError(t, err)
|
||||
|
||||
entries, _ := rb.Query(0, 10)
|
||||
assert.Nil(t, entries)
|
||||
}
|
||||
|
||||
func TestLogRingBuffer_QueryAfterOverflow(t *testing.T) {
|
||||
rb := NewLogRingBuffer(3)
|
||||
|
||||
_, _ = rb.Write([]byte("1\n2\n3\n4\n5\n6\n7\n"))
|
||||
|
||||
entries, hasMore := rb.Query(0, 10)
|
||||
assert.False(t, hasMore)
|
||||
assert.Equal(t, 3, len(entries))
|
||||
assert.Equal(t, "5", entries[0].Data)
|
||||
assert.Equal(t, "6", entries[1].Data)
|
||||
assert.Equal(t, "7", entries[2].Data)
|
||||
|
||||
// Query by cursor - index 4 is "5", so cursor=4 should return index < 4
|
||||
older, hasMore2 := rb.Query(4, 10)
|
||||
assert.False(t, hasMore2)
|
||||
assert.Nil(t, older)
|
||||
}
|
||||
|
||||
func TestLogRingBuffer_NextCursor(t *testing.T) {
|
||||
rb := NewLogRingBuffer(10)
|
||||
|
||||
_, _ = rb.Write([]byte("a\nb\nc\nd\ne\n"))
|
||||
|
||||
// Query latest 2, should return next_cursor pointing to first returned entry
|
||||
entries, _ := rb.Query(0, 2)
|
||||
assert.Equal(t, 2, len(entries))
|
||||
// entries[0].Index = 3 ("d"), entries[1].Index = 4 ("e")
|
||||
assert.Equal(t, 3, entries[0].Index)
|
||||
|
||||
// Now use that index as cursor to get older entries
|
||||
older, hasMore := rb.Query(entries[0].Index, 10)
|
||||
assert.False(t, hasMore)
|
||||
assert.Equal(t, 3, len(older))
|
||||
assert.Equal(t, "a", older[0].Data)
|
||||
assert.Equal(t, "b", older[1].Data)
|
||||
assert.Equal(t, "c", older[2].Data)
|
||||
}
|
||||
@@ -0,0 +1,99 @@
|
||||
// Copyright 2025 linux.do
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package logger
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"log"
|
||||
"os"
|
||||
"path/filepath"
|
||||
|
||||
"go.opentelemetry.io/otel/trace"
|
||||
"go.uber.org/zap"
|
||||
"go.uber.org/zap/zapcore"
|
||||
"gopkg.in/natefinch/lumberjack.v2"
|
||||
)
|
||||
|
||||
// logDirPerm 日志目录权限
|
||||
const logDirPerm = 0750
|
||||
|
||||
func getLogWriterForConfig(cfg Config) (zapcore.WriteSyncer, error) {
|
||||
if cfg.Output == "file" {
|
||||
// 初始化日志目录
|
||||
logPath := cfg.FilePath
|
||||
logDir := filepath.Dir(logPath)
|
||||
if err := os.MkdirAll(logDir, logDirPerm); err != nil {
|
||||
return nil, fmt.Errorf(errCreateLogFileDirFailed, err)
|
||||
}
|
||||
|
||||
// 配置日志轮转
|
||||
logOutput := &lumberjack.Logger{
|
||||
Filename: logPath,
|
||||
MaxSize: cfg.MaxSize,
|
||||
MaxBackups: cfg.MaxBackups,
|
||||
MaxAge: cfg.MaxAge,
|
||||
Compress: cfg.Compress,
|
||||
}
|
||||
|
||||
return zapcore.AddSync(logOutput), nil
|
||||
}
|
||||
|
||||
return zapcore.AddSync(os.Stdout), nil
|
||||
}
|
||||
|
||||
// getEncoderForConfig 获取日志编码器
|
||||
func getEncoderForConfig(cfg Config) zapcore.Encoder {
|
||||
// 编码器配置
|
||||
encoderConfig := zapcore.EncoderConfig{
|
||||
TimeKey: "time",
|
||||
LevelKey: "level",
|
||||
NameKey: "logger",
|
||||
CallerKey: "caller",
|
||||
MessageKey: "msg",
|
||||
StacktraceKey: "stacktrace",
|
||||
LineEnding: zapcore.DefaultLineEnding,
|
||||
EncodeLevel: zapcore.LowercaseLevelEncoder,
|
||||
EncodeTime: zapcore.ISO8601TimeEncoder,
|
||||
EncodeDuration: zapcore.SecondsDurationEncoder,
|
||||
EncodeCaller: zapcore.ShortCallerEncoder,
|
||||
}
|
||||
|
||||
if cfg.Format == "json" {
|
||||
return zapcore.NewJSONEncoder(encoderConfig)
|
||||
}
|
||||
return zapcore.NewConsoleEncoder(encoderConfig)
|
||||
}
|
||||
|
||||
// getLogLevelForConfig 获取日志级别
|
||||
func getLogLevelForConfig(cfg Config) zapcore.Level {
|
||||
level := cfg.Level
|
||||
|
||||
switch level {
|
||||
case "debug":
|
||||
return zapcore.DebugLevel
|
||||
case "info":
|
||||
return zapcore.InfoLevel
|
||||
case "warn":
|
||||
return zapcore.WarnLevel
|
||||
case "error":
|
||||
return zapcore.ErrorLevel
|
||||
default:
|
||||
log.Printf("[Logger] invalid log level: %s, defaulting to info\n", level)
|
||||
return zapcore.InfoLevel
|
||||
}
|
||||
}
|
||||
|
||||
func getTraceIDFields(ctx context.Context) []zap.Field {
|
||||
span := trace.SpanFromContext(ctx)
|
||||
spanContext := span.SpanContext()
|
||||
if !spanContext.IsValid() {
|
||||
return nil
|
||||
}
|
||||
return []zap.Field{
|
||||
zap.String("traceID", spanContext.TraceID().String()),
|
||||
zap.String("spanID", spanContext.SpanID().String()),
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,16 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package mail 提供 SMTP 邮件发送功能。
|
||||
package mail
|
||||
|
||||
const (
|
||||
errDialTLSFailed = "dial tls failed: %w"
|
||||
errSMTPClientCreationFailed = "smtp client creation failed: %w"
|
||||
errSMTPAuthFailed = "smtp auth failed: %w"
|
||||
errSMTPMailCommandFailed = "smtp mail command failed: %w"
|
||||
errSMTPRcptCommandFailed = "smtp rcpt command failed: %w"
|
||||
errSMTPDataCommandFailed = "smtp data command failed: %w"
|
||||
errSMTPWritingBodyFailed = "smtp writing body failed: %w" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
|
||||
errSendMailFailed = "send mail failed: %w"
|
||||
)
|
||||
@@ -0,0 +1,250 @@
|
||||
// Copyright 2025 linux.do
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package mail
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"crypto/tls"
|
||||
"fmt"
|
||||
"net"
|
||||
"net/smtp"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
const (
|
||||
smtpSSLPort = 465 // SMTP SSL 端口
|
||||
smtpDialTimeout = 5 * time.Second // SMTP 连接超时
|
||||
smtpSessionDeadline = 10 * time.Second // SMTP 会话截止时间
|
||||
)
|
||||
|
||||
// Config represents SMTP mail configuration
|
||||
type Config struct {
|
||||
Host string
|
||||
Port int
|
||||
Username string
|
||||
Password string
|
||||
}
|
||||
|
||||
// sanitizeHeaderValue removes CR/LF bytes so untrusted values cannot inject
|
||||
// additional email headers (email header injection).
|
||||
func sanitizeHeaderValue(v string) string {
|
||||
v = strings.ReplaceAll(v, "\r", "")
|
||||
v = strings.ReplaceAll(v, "\n", "")
|
||||
return v
|
||||
}
|
||||
|
||||
// SendMail sends an HTML email using the provided config and message details
|
||||
func SendMail(ctx context.Context, cfg Config, to string, subject, body string) error {
|
||||
return SendMailHTML(ctx, cfg, to, subject, body)
|
||||
}
|
||||
|
||||
// SendMailHTML sends an HTML format email
|
||||
func SendMailHTML(ctx context.Context, cfg Config, to string, subject, body string) error {
|
||||
addr := net.JoinHostPort(cfg.Host, strconv.Itoa(cfg.Port))
|
||||
|
||||
// Header & MIME settings for HTML email
|
||||
header := make(map[string]string)
|
||||
header["From"] = sanitizeHeaderValue(cfg.Username)
|
||||
header["To"] = sanitizeHeaderValue(to)
|
||||
header["Subject"] = sanitizeHeaderValue(subject)
|
||||
header["MIME-Version"] = "1.0"
|
||||
header["Content-Type"] = "text/html; charset=UTF-8"
|
||||
|
||||
message := ""
|
||||
for k, v := range header {
|
||||
message += fmt.Sprintf("%s: %s\r\n", k, v)
|
||||
}
|
||||
message += "\r\n" + body
|
||||
|
||||
auth := smtp.PlainAuth("", cfg.Username, cfg.Password, cfg.Host)
|
||||
|
||||
// If using SSL port 465, we connection via TLS dial
|
||||
if cfg.Port == smtpSSLPort {
|
||||
return sendMailViaSSL(ctx, addr, auth, cfg, to, message)
|
||||
}
|
||||
|
||||
// For standard port (587 / 25), use smtp.SendMail directly (handles STARTTLS automatically if server supports it)
|
||||
err := smtp.SendMail(addr, auth, cfg.Username, []string{to}, []byte(message))
|
||||
if err != nil {
|
||||
return fmt.Errorf(errSendMailFailed, err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// sendMailViaSSL 通过 TLS 直接连接 SMTP SSL 端口发送邮件
|
||||
func sendMailViaSSL(ctx context.Context, addr string, auth smtp.Auth, cfg Config, to, message string) error {
|
||||
tlsConfig := &tls.Config{
|
||||
InsecureSkipVerify: true, //nolint:gosec // SMTP servers might use self-signed certificates
|
||||
ServerName: cfg.Host,
|
||||
}
|
||||
dialer := &net.Dialer{Timeout: smtpDialTimeout}
|
||||
tlsDialer := &tls.Dialer{
|
||||
NetDialer: dialer,
|
||||
Config: tlsConfig,
|
||||
}
|
||||
conn, err := tlsDialer.DialContext(ctx, "tcp", addr)
|
||||
if err != nil {
|
||||
return fmt.Errorf(errDialTLSFailed, err)
|
||||
}
|
||||
defer func() { _ = conn.Close() }()
|
||||
_ = conn.SetDeadline(time.Now().Add(smtpSessionDeadline))
|
||||
|
||||
client, err := smtp.NewClient(conn, cfg.Host)
|
||||
if err != nil {
|
||||
return fmt.Errorf(errSMTPClientCreationFailed, err)
|
||||
}
|
||||
defer func() { _ = client.Close() }()
|
||||
|
||||
if err = client.Auth(auth); err != nil {
|
||||
return fmt.Errorf(errSMTPAuthFailed, err)
|
||||
}
|
||||
if err = client.Mail(cfg.Username); err != nil {
|
||||
return fmt.Errorf(errSMTPMailCommandFailed, err)
|
||||
}
|
||||
if err = client.Rcpt(to); err != nil {
|
||||
return fmt.Errorf(errSMTPRcptCommandFailed, err)
|
||||
}
|
||||
|
||||
w, err := client.Data()
|
||||
if err != nil {
|
||||
return fmt.Errorf(errSMTPDataCommandFailed, err)
|
||||
}
|
||||
defer func() { _ = w.Close() }()
|
||||
|
||||
_, err = w.Write([]byte(message))
|
||||
if err != nil {
|
||||
return fmt.Errorf(errSMTPWritingBodyFailed, err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// SendMailWithLog sends a test email and records a detailed SMTP connection log
|
||||
func SendMailWithLog(ctx context.Context, cfg Config, to string, subject, body string) (string, error) {
|
||||
var logBuf bytes.Buffer
|
||||
logLine := func(dir string, format string, args ...interface{}) {
|
||||
fmt.Fprintf(&logBuf, "[%s] %s\n", dir, fmt.Sprintf(format, args...))
|
||||
}
|
||||
|
||||
addr := net.JoinHostPort(cfg.Host, strconv.Itoa(cfg.Port))
|
||||
logLine("System", "Connecting to %s...", addr)
|
||||
|
||||
var conn net.Conn
|
||||
var err error
|
||||
dialer := &net.Dialer{Timeout: smtpDialTimeout}
|
||||
if cfg.Port == smtpSSLPort {
|
||||
tlsConfig := &tls.Config{
|
||||
InsecureSkipVerify: true, //nolint:gosec // SMTP servers might use self-signed certificates
|
||||
ServerName: cfg.Host,
|
||||
}
|
||||
tlsDialer := &tls.Dialer{
|
||||
NetDialer: dialer,
|
||||
Config: tlsConfig,
|
||||
}
|
||||
conn, err = tlsDialer.DialContext(ctx, "tcp", addr)
|
||||
} else {
|
||||
conn, err = dialer.DialContext(ctx, "tcp", addr)
|
||||
}
|
||||
if err != nil {
|
||||
logLine("Error", "Connection failed: %v", err)
|
||||
return logBuf.String(), err
|
||||
}
|
||||
defer func() { _ = conn.Close() }()
|
||||
logLine("System", "Connected successfully.")
|
||||
|
||||
// Set a 10-second session deadline for read/write operations
|
||||
_ = conn.SetDeadline(time.Now().Add(smtpSessionDeadline))
|
||||
|
||||
client, err := smtp.NewClient(conn, cfg.Host)
|
||||
if err != nil {
|
||||
logLine("Error", "SMTP client handshake failed: %v", err)
|
||||
return logBuf.String(), err
|
||||
}
|
||||
defer func() { _ = client.Close() }()
|
||||
|
||||
// If not 465, support STARTTLS if available
|
||||
if cfg.Port != smtpSSLPort {
|
||||
if ok, _ := client.Extension("STARTTLS"); ok {
|
||||
logLine("C", "STARTTLS")
|
||||
tlsConfig := &tls.Config{
|
||||
InsecureSkipVerify: true, //nolint:gosec // SMTP servers might use self-signed certificates
|
||||
ServerName: cfg.Host,
|
||||
}
|
||||
if err = client.StartTLS(tlsConfig); err != nil {
|
||||
logLine("Error", "STARTTLS failed: %v", err)
|
||||
return logBuf.String(), err
|
||||
}
|
||||
logLine("S", "220 Ready to start TLS")
|
||||
}
|
||||
}
|
||||
|
||||
// Authentication
|
||||
if cfg.Username != "" && cfg.Password != "" {
|
||||
auth := smtp.PlainAuth("", cfg.Username, cfg.Password, cfg.Host)
|
||||
logLine("C", "AUTH PLAIN **********")
|
||||
if err = client.Auth(auth); err != nil {
|
||||
logLine("Error", "Authentication failed: %v", err)
|
||||
return logBuf.String(), err
|
||||
}
|
||||
logLine("S", "235 Authentication successful")
|
||||
}
|
||||
|
||||
// Mail command
|
||||
logLine("C", "MAIL FROM:<%s>", cfg.Username)
|
||||
if err = client.Mail(cfg.Username); err != nil {
|
||||
logLine("Error", "MAIL FROM command failed: %v", err)
|
||||
return logBuf.String(), err
|
||||
}
|
||||
logLine("S", "250 OK")
|
||||
|
||||
// Rcpt command
|
||||
logLine("C", "RCPT TO:<%s>", to)
|
||||
if err = client.Rcpt(to); err != nil {
|
||||
logLine("Error", "RCPT TO command failed: %v", err)
|
||||
return logBuf.String(), err
|
||||
}
|
||||
logLine("S", "250 OK")
|
||||
|
||||
// Data command
|
||||
logLine("C", "DATA")
|
||||
w, err := client.Data()
|
||||
if err != nil {
|
||||
logLine("Error", "DATA command failed: %v", err)
|
||||
return logBuf.String(), err
|
||||
}
|
||||
logLine("S", "354 Start mail input")
|
||||
|
||||
// Header & MIME settings for HTML email
|
||||
header := make(map[string]string)
|
||||
header["From"] = sanitizeHeaderValue(cfg.Username)
|
||||
header["To"] = sanitizeHeaderValue(to)
|
||||
header["Subject"] = sanitizeHeaderValue(subject)
|
||||
header["MIME-Version"] = "1.0"
|
||||
header["Content-Type"] = "text/html; charset=UTF-8"
|
||||
|
||||
message := ""
|
||||
for k, v := range header {
|
||||
message += fmt.Sprintf("%s: %s\r\n", k, v)
|
||||
}
|
||||
message += "\r\n" + body
|
||||
|
||||
logLine("System", "Sending message body...")
|
||||
if _, err = w.Write([]byte(message)); err != nil {
|
||||
_ = w.Close()
|
||||
logLine("Error", "Writing message body failed: %v", err)
|
||||
return logBuf.String(), err
|
||||
}
|
||||
_ = w.Close()
|
||||
logLine("S", "250 OK")
|
||||
|
||||
logLine("C", "QUIT")
|
||||
_ = client.Quit()
|
||||
logLine("System", "Mail sent successfully!")
|
||||
|
||||
return logBuf.String(), nil
|
||||
}
|
||||
@@ -0,0 +1,112 @@
|
||||
// Copyright 2025 linux.do
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package mail
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"context"
|
||||
"net"
|
||||
"net/textproto"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestSendMailMock(t *testing.T) {
|
||||
// Start a mock SMTP server
|
||||
l, err := net.Listen("tcp", "127.0.0.1:0")
|
||||
if err != nil {
|
||||
t.Fatalf("failed to start mock smtp server: %v", err)
|
||||
}
|
||||
defer func() { _ = l.Close() }()
|
||||
|
||||
port := l.Addr().(*net.TCPAddr).Port
|
||||
|
||||
go func() {
|
||||
conn, err := l.Accept()
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
defer func() { _ = conn.Close() }()
|
||||
|
||||
writer := bufio.NewWriter(conn)
|
||||
reader := bufio.NewReader(conn)
|
||||
tp := textproto.NewReader(reader)
|
||||
|
||||
// 220 Ready
|
||||
_, _ = writer.WriteString("220 mock.smtp.com SMTP Ready\r\n")
|
||||
_ = writer.Flush()
|
||||
|
||||
// Read HELO/EHLO
|
||||
_, _ = tp.ReadLine()
|
||||
_, _ = writer.WriteString("250-mock.smtp.com\r\n250 AUTH PLAIN\r\n")
|
||||
_ = writer.Flush()
|
||||
|
||||
// Read AUTH PLAIN
|
||||
_, _ = tp.ReadLine()
|
||||
_, _ = writer.WriteString("235 Authentication successful\r\n")
|
||||
_ = writer.Flush()
|
||||
|
||||
// Read MAIL FROM
|
||||
_, _ = tp.ReadLine()
|
||||
_, _ = writer.WriteString("250 OK\r\n")
|
||||
_ = writer.Flush()
|
||||
|
||||
// Read RCPT TO
|
||||
_, _ = tp.ReadLine()
|
||||
_, _ = writer.WriteString("250 OK\r\n")
|
||||
_ = writer.Flush()
|
||||
|
||||
// Read DATA
|
||||
_, _ = tp.ReadLine()
|
||||
_, _ = writer.WriteString("354 Start mail input\r\n")
|
||||
_ = writer.Flush()
|
||||
|
||||
// Read body lines until dot
|
||||
for {
|
||||
line, err := tp.ReadLine()
|
||||
if err != nil || line == "." {
|
||||
break
|
||||
}
|
||||
}
|
||||
_, _ = writer.WriteString("250 OK\r\n")
|
||||
_ = writer.Flush()
|
||||
|
||||
// Read QUIT
|
||||
_, _ = tp.ReadLine()
|
||||
_, _ = writer.WriteString("221 Bye\r\n")
|
||||
_ = writer.Flush()
|
||||
}()
|
||||
|
||||
cfg := Config{
|
||||
Host: "127.0.0.1",
|
||||
Port: port,
|
||||
Username: "test@example.com",
|
||||
Password: "password",
|
||||
}
|
||||
|
||||
err = SendMail(context.Background(), cfg, "recipient@example.com", "Test Subject", "<h1>Test Body</h1>")
|
||||
if err != nil {
|
||||
t.Errorf("failed to send mail: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSanitizeHeaderValue(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
input string
|
||||
want string
|
||||
}{
|
||||
{"plain", "System Notification", "System Notification"},
|
||||
{"crlf stripped", "alert\r\nBcc: attacker@example.com", "alertBcc: attacker@example.com"},
|
||||
{"cr stripped", "a\rb", "ab"},
|
||||
{"lf stripped", "a\nb", "ab"},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
if got := sanitizeHeaderValue(tt.input); got != tt.want {
|
||||
t.Errorf("sanitizeHeaderValue(%q) = %q, want %q", tt.input, got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,45 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package response
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
// AbortBadRequest 以 400 中断请求并将错误挂载到 Gin Error 链,供全局中间件统一记录 Trace 并响应。
|
||||
func AbortBadRequest(c *gin.Context, msg string) {
|
||||
AbortWithError(c, http.StatusBadRequest, msg)
|
||||
}
|
||||
|
||||
// AbortUnauthorized 以 401 中断请求并将错误挂载到 Gin Error 链。
|
||||
func AbortUnauthorized(c *gin.Context, msg string) {
|
||||
AbortWithError(c, http.StatusUnauthorized, msg)
|
||||
}
|
||||
|
||||
// AbortForbidden 以 403 中断请求并将错误挂载到 Gin Error 链。
|
||||
func AbortForbidden(c *gin.Context, msg string) {
|
||||
AbortWithError(c, http.StatusForbidden, msg)
|
||||
}
|
||||
|
||||
// AbortNotFound 以 404 中断请求并将错误挂载到 Gin Error 链。
|
||||
func AbortNotFound(c *gin.Context, msg string) {
|
||||
AbortWithError(c, http.StatusNotFound, msg)
|
||||
}
|
||||
|
||||
// AbortInternal 以 500 中断请求并将错误挂载到 Gin Error 链。
|
||||
func AbortInternal(c *gin.Context, msg string) {
|
||||
AbortWithError(c, http.StatusInternalServerError, msg)
|
||||
}
|
||||
|
||||
// AbortTooManyRequests 以 429 中断请求并将错误挂载到 Gin Error 链。
|
||||
func AbortTooManyRequests(c *gin.Context, msg string) {
|
||||
AbortWithError(c, http.StatusTooManyRequests, msg)
|
||||
}
|
||||
|
||||
// AbortConflict 以 409 中断请求并将错误挂载到 Gin Error 链。
|
||||
func AbortConflict(c *gin.Context, msg string) {
|
||||
AbortWithError(c, http.StatusConflict, msg)
|
||||
}
|
||||
@@ -0,0 +1,40 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package response
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"net/http"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"go.opentelemetry.io/otel/codes"
|
||||
"go.opentelemetry.io/otel/trace"
|
||||
)
|
||||
|
||||
// ErrorHandlerMiddleware 捕获 c.Errors 并统一格式化为 JSON 返回给客户端,同时将其记录到 Span 异常中。
|
||||
// 与 AbortWithError / AbortBadRequest 等配合使用,是全局 OTel 友好错误响应的唯一出口。
|
||||
func ErrorHandlerMiddleware() gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
c.Next()
|
||||
|
||||
if len(c.Errors) == 0 || c.Writer.Written() {
|
||||
return
|
||||
}
|
||||
|
||||
err := c.Errors.Last().Err
|
||||
span := trace.SpanFromContext(c.Request.Context())
|
||||
if span.IsRecording() {
|
||||
span.RecordError(err)
|
||||
span.SetStatus(codes.Error, err.Error())
|
||||
}
|
||||
|
||||
var apiErr *APIError
|
||||
if errors.As(err, &apiErr) {
|
||||
c.JSON(apiErr.Code, Err(apiErr.Msg))
|
||||
return
|
||||
}
|
||||
|
||||
c.JSON(http.StatusInternalServerError, Err("内部系统错误"))
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,168 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package response
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"go.opentelemetry.io/otel"
|
||||
"go.opentelemetry.io/otel/codes"
|
||||
sdktrace "go.opentelemetry.io/otel/sdk/trace"
|
||||
"go.opentelemetry.io/otel/sdk/trace/tracetest"
|
||||
"go.opentelemetry.io/otel/trace"
|
||||
)
|
||||
|
||||
func init() {
|
||||
gin.SetMode(gin.TestMode)
|
||||
}
|
||||
|
||||
func TestAbortWithError(t *testing.T) {
|
||||
w := httptest.NewRecorder()
|
||||
c, _ := gin.CreateTestContext(w)
|
||||
|
||||
AbortWithError(c, http.StatusBadRequest, "invalid input")
|
||||
|
||||
require.Len(t, c.Errors, 1)
|
||||
|
||||
var apiErr *APIError
|
||||
require.True(t, errors.As(c.Errors.Last().Err, &apiErr))
|
||||
assert.Equal(t, http.StatusBadRequest, apiErr.Code)
|
||||
assert.Equal(t, "invalid input", apiErr.Msg)
|
||||
assert.True(t, c.IsAborted())
|
||||
}
|
||||
|
||||
func TestErrorHandlerMiddleware_APIErrorStatusCodes(t *testing.T) {
|
||||
cases := []struct {
|
||||
name string
|
||||
statusCode int
|
||||
message string
|
||||
abort func(*gin.Context, string)
|
||||
}{
|
||||
{"400 Bad Request", http.StatusBadRequest, "bad request", AbortBadRequest},
|
||||
{"401 Unauthorized", http.StatusUnauthorized, "unauthorized", AbortUnauthorized},
|
||||
{"403 Forbidden", http.StatusForbidden, "forbidden", AbortForbidden},
|
||||
{"404 Not Found", http.StatusNotFound, "not found", AbortNotFound},
|
||||
{"409 Conflict", http.StatusConflict, "conflict", AbortConflict},
|
||||
{"429 Too Many Requests", http.StatusTooManyRequests, "too many requests", AbortTooManyRequests},
|
||||
{"500 Internal Server Error", http.StatusInternalServerError, "internal error", AbortInternal},
|
||||
}
|
||||
|
||||
for _, tc := range cases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
r := gin.New()
|
||||
r.Use(ErrorHandlerMiddleware())
|
||||
r.GET("/test", func(c *gin.Context) {
|
||||
tc.abort(c, tc.message)
|
||||
})
|
||||
|
||||
req := httptest.NewRequest(http.MethodGet, "/test", nil)
|
||||
w := httptest.NewRecorder()
|
||||
r.ServeHTTP(w, req)
|
||||
|
||||
assert.Equal(t, tc.statusCode, w.Code)
|
||||
assert.Equal(t, "application/json; charset=utf-8", w.Header().Get("Content-Type"))
|
||||
|
||||
var body Response[any]
|
||||
require.NoError(t, json.Unmarshal(w.Body.Bytes(), &body))
|
||||
assert.Equal(t, tc.message, body.ErrorMsg)
|
||||
assert.Nil(t, body.Data)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestErrorHandlerMiddleware_SkipsWhenNoErrors(t *testing.T) {
|
||||
r := gin.New()
|
||||
r.Use(ErrorHandlerMiddleware())
|
||||
r.GET("/ok", func(c *gin.Context) {
|
||||
c.JSON(http.StatusOK, OK("success"))
|
||||
})
|
||||
|
||||
req := httptest.NewRequest(http.MethodGet, "/ok", nil)
|
||||
w := httptest.NewRecorder()
|
||||
r.ServeHTTP(w, req)
|
||||
|
||||
assert.Equal(t, http.StatusOK, w.Code)
|
||||
|
||||
var body Response[string]
|
||||
require.NoError(t, json.Unmarshal(w.Body.Bytes(), &body))
|
||||
assert.Equal(t, "success", body.Data)
|
||||
assert.Empty(t, body.ErrorMsg)
|
||||
}
|
||||
|
||||
func TestErrorHandlerMiddleware_SkipsWhenResponseAlreadyWritten(t *testing.T) {
|
||||
r := gin.New()
|
||||
r.Use(ErrorHandlerMiddleware())
|
||||
r.GET("/written", func(c *gin.Context) {
|
||||
c.JSON(http.StatusOK, OKNil())
|
||||
_ = c.Error(NewError(http.StatusBadRequest, "should not overwrite"))
|
||||
})
|
||||
|
||||
req := httptest.NewRequest(http.MethodGet, "/written", nil)
|
||||
w := httptest.NewRecorder()
|
||||
r.ServeHTTP(w, req)
|
||||
|
||||
assert.Equal(t, http.StatusOK, w.Code)
|
||||
|
||||
var body Response[any]
|
||||
require.NoError(t, json.Unmarshal(w.Body.Bytes(), &body))
|
||||
assert.Empty(t, body.ErrorMsg)
|
||||
assert.Nil(t, body.Data)
|
||||
}
|
||||
|
||||
func TestErrorHandlerMiddleware_FallbackForNonAPIError(t *testing.T) {
|
||||
r := gin.New()
|
||||
r.Use(ErrorHandlerMiddleware())
|
||||
r.GET("/plain", func(c *gin.Context) {
|
||||
_ = c.Error(errors.New("plain error"))
|
||||
})
|
||||
|
||||
req := httptest.NewRequest(http.MethodGet, "/plain", nil)
|
||||
w := httptest.NewRecorder()
|
||||
r.ServeHTTP(w, req)
|
||||
|
||||
assert.Equal(t, http.StatusInternalServerError, w.Code)
|
||||
|
||||
var body Response[any]
|
||||
require.NoError(t, json.Unmarshal(w.Body.Bytes(), &body))
|
||||
assert.Equal(t, "内部系统错误", body.ErrorMsg)
|
||||
assert.Nil(t, body.Data)
|
||||
}
|
||||
|
||||
func TestErrorHandlerMiddleware_RecordsSpanOnAPIError(t *testing.T) {
|
||||
sr := tracetest.NewSpanRecorder()
|
||||
tp := sdktrace.NewTracerProvider(sdktrace.WithSpanProcessor(sr))
|
||||
otel.SetTracerProvider(tp)
|
||||
defer otel.SetTracerProvider(trace.NewNoopTracerProvider())
|
||||
|
||||
tracer := tp.Tracer("test")
|
||||
ctx, span := tracer.Start(context.Background(), "request")
|
||||
|
||||
r := gin.New()
|
||||
r.Use(ErrorHandlerMiddleware())
|
||||
r.GET("/err", func(c *gin.Context) {
|
||||
c.Request = c.Request.WithContext(ctx)
|
||||
AbortBadRequest(c, "bad request")
|
||||
})
|
||||
|
||||
req := httptest.NewRequest(http.MethodGet, "/err", nil)
|
||||
w := httptest.NewRecorder()
|
||||
r.ServeHTTP(w, req)
|
||||
span.End()
|
||||
|
||||
require.Equal(t, http.StatusBadRequest, w.Code)
|
||||
|
||||
spans := sr.Ended()
|
||||
require.Len(t, spans, 1)
|
||||
assert.Equal(t, codes.Error, spans[0].Status().Code)
|
||||
assert.Equal(t, "bad request", spans[0].Status().Description)
|
||||
require.NotEmpty(t, spans[0].Events())
|
||||
}
|
||||
@@ -0,0 +1,57 @@
|
||||
// Copyright 2025 linux.do
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package response provides shared HTTP API response structures.
|
||||
package response
|
||||
|
||||
import "github.com/gin-gonic/gin"
|
||||
|
||||
// Response 通用响应体
|
||||
type Response[T any] struct {
|
||||
ErrorMsg string `json:"error_msg"`
|
||||
Data T `json:"data"`
|
||||
}
|
||||
|
||||
// Any 用于 Swagger 文档的响应类型(非泛型)
|
||||
// swag 不支持泛型,使用此类型替代 Response[T]
|
||||
type Any struct {
|
||||
ErrorMsg string `json:"error_msg" example:""`
|
||||
Data interface{} `json:"data"`
|
||||
}
|
||||
|
||||
// APIError 统一的 API 业务错误类型,可被全局错误处理中间件捕获
|
||||
type APIError struct {
|
||||
Code int
|
||||
Msg string
|
||||
}
|
||||
|
||||
func (e *APIError) Error() string {
|
||||
return e.Msg
|
||||
}
|
||||
|
||||
// NewError 实例化一个 APIError
|
||||
func NewError(code int, msg string) *APIError {
|
||||
return &APIError{Code: code, Msg: msg}
|
||||
}
|
||||
|
||||
// AbortWithError 将 API 错误挂载到 Gin Context 并中断执行流
|
||||
func AbortWithError(c *gin.Context, code int, msg string) {
|
||||
_ = c.Error(NewError(code, msg))
|
||||
c.Abort()
|
||||
}
|
||||
|
||||
// OK 构造成功响应
|
||||
func OK[T any](data T) Response[T] {
|
||||
return Response[T]{Data: data}
|
||||
}
|
||||
|
||||
// OKNil 构造成功响应(data 为 null)
|
||||
func OKNil() Response[any] {
|
||||
return Response[any]{Data: nil}
|
||||
}
|
||||
|
||||
// Err 构造错误响应
|
||||
func Err(msg string) Response[any] {
|
||||
return Response[any]{ErrorMsg: msg, Data: nil}
|
||||
}
|
||||
@@ -0,0 +1,17 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package testhelper
|
||||
|
||||
// RegisterCleanup registers an extra cleanup hook invoked by SetupTestEnvironment.
|
||||
func RegisterCleanup(fn func()) {
|
||||
extraCleanups = append(extraCleanups, fn)
|
||||
}
|
||||
|
||||
var extraCleanups []func()
|
||||
|
||||
func runExtraCleanups() {
|
||||
for _, fn := range extraCleanups {
|
||||
fn()
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,20 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package testhelper
|
||||
|
||||
import (
|
||||
"github.com/Rain-kl/Wavelet/backend/pkg/response"
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
// NewTestGinEngine 创建带 ErrorHandlerMiddleware 的 Gin 引擎,与生产环境错误响应行为一致。
|
||||
func NewTestGinEngine(middlewares ...gin.HandlerFunc) *gin.Engine {
|
||||
gin.SetMode(gin.TestMode)
|
||||
r := gin.New()
|
||||
r.Use(response.ErrorHandlerMiddleware())
|
||||
for _, middleware := range middlewares {
|
||||
r.Use(middleware)
|
||||
}
|
||||
return r
|
||||
}
|
||||
@@ -0,0 +1,333 @@
|
||||
// Copyright 2025 linux.do
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package testhelper 提供测试辅助工具
|
||||
package testhelper
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
cachepkg "github.com/Rain-kl/Wavelet/backend/plugins/infra/cache"
|
||||
db "github.com/Rain-kl/Wavelet/backend/plugins/infra/database"
|
||||
"github.com/alicebob/miniredis/v2"
|
||||
"github.com/glebarez/sqlite"
|
||||
"github.com/redis/go-redis/v9"
|
||||
"github.com/redis/go-redis/v9/maintnotifications"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
// SystemConfig 测试用系统配置表
|
||||
type SystemConfig struct {
|
||||
Key string `gorm:"primaryKey;size:64;not null"`
|
||||
Value string `gorm:"type:text;not null"`
|
||||
Type string `gorm:"size:32;not null"`
|
||||
Visibility string `gorm:"size:32;not null;default:'hidden'"`
|
||||
Description string `gorm:"size:255"`
|
||||
}
|
||||
|
||||
// TableName 返回测试配置表表名
|
||||
func (SystemConfig) TableName() string {
|
||||
return "w_system_configs"
|
||||
}
|
||||
|
||||
type userHelper struct {
|
||||
ID uint64 `gorm:"primaryKey;autoIncrement"`
|
||||
Username string `gorm:"size:64;uniqueIndex;not null"`
|
||||
Nickname string `gorm:"size:64;not null;default:''"`
|
||||
Password string `gorm:"size:255;not null;default:''"`
|
||||
Email string `gorm:"size:128;index;default:''"`
|
||||
AvatarURL string `gorm:"size:255;default:''"`
|
||||
IsAdmin bool `gorm:"default:false;not null"`
|
||||
IsActive bool `gorm:"default:true;not null"`
|
||||
NeedChangePassword bool `gorm:"default:false;not null"`
|
||||
Bio string `gorm:"size:500;default:''"`
|
||||
Phone string `gorm:"size:32;default:''"`
|
||||
Gender string `gorm:"size:16;default:''"`
|
||||
Website string `gorm:"size:255;default:''"`
|
||||
Location string `gorm:"size:255;default:''"`
|
||||
LastLoginAt time.Time
|
||||
CreatedAt time.Time `gorm:"autoCreateTime"`
|
||||
UpdatedAt time.Time `gorm:"autoUpdateTime"`
|
||||
}
|
||||
|
||||
func (userHelper) TableName() string { return "w_users" }
|
||||
|
||||
type accessTokenHelper struct {
|
||||
ID uint64 `gorm:"primaryKey;autoIncrement"`
|
||||
UserID uint64 `gorm:"not null;index"`
|
||||
Name string `gorm:"size:64;not null"`
|
||||
TokenHash string `gorm:"size:64;uniqueIndex;not null"`
|
||||
MaskedToken string `gorm:"size:32;not null"`
|
||||
IsAdmin bool `gorm:"default:false;not null"`
|
||||
CreatedAt time.Time `gorm:"autoCreateTime"`
|
||||
UpdatedAt time.Time `gorm:"autoUpdateTime"`
|
||||
}
|
||||
|
||||
func (accessTokenHelper) TableName() string { return "w_access_tokens" }
|
||||
|
||||
type authSourceHelper struct {
|
||||
ID uint64 `gorm:"primaryKey;autoIncrement"`
|
||||
Type string `gorm:"size:32;not null"`
|
||||
Name string `gorm:"size:64;not null"`
|
||||
Enabled bool `gorm:"default:false;not null"`
|
||||
CreatedAt time.Time `gorm:"autoCreateTime"`
|
||||
UpdatedAt time.Time `gorm:"autoUpdateTime"`
|
||||
}
|
||||
|
||||
func (authSourceHelper) TableName() string { return "w_auth_sources" }
|
||||
|
||||
type externalAccountHelper struct {
|
||||
ID uint64 `gorm:"primaryKey;autoIncrement"`
|
||||
UserID uint64 `gorm:"not null;index"`
|
||||
AuthSourceType string `gorm:"size:32;not null;index"`
|
||||
ExternalID string `gorm:"size:128;not null;index"`
|
||||
Username string `gorm:"size:128;default:''"`
|
||||
Email string `gorm:"size:128;default:''"`
|
||||
CreatedAt time.Time `gorm:"autoCreateTime"`
|
||||
UpdatedAt time.Time `gorm:"autoUpdateTime"`
|
||||
}
|
||||
|
||||
func (externalAccountHelper) TableName() string { return "w_external_accounts" }
|
||||
|
||||
type uploadHelper struct {
|
||||
ID uint64 `gorm:"primaryKey;autoIncrement"`
|
||||
UserID uint64 `gorm:"not null;index"`
|
||||
FileName string `gorm:"size:255;not null"`
|
||||
FilePath string `gorm:"size:500;not null"`
|
||||
FileSize int64 `gorm:"not null"`
|
||||
MimeType string `gorm:"size:128;not null"`
|
||||
Extension string `gorm:"size:32;not null"`
|
||||
Hash string `gorm:"size:64;index;not null;default:''"`
|
||||
Type string `gorm:"size:50;not null;index"`
|
||||
Status string `gorm:"size:20;not null;default:'pending'"`
|
||||
AccessMode int `gorm:"not null;default:0"`
|
||||
Metadata string `gorm:"type:text"`
|
||||
CreatedAt time.Time `gorm:"autoCreateTime;index"`
|
||||
UpdatedAt time.Time `gorm:"autoUpdateTime"`
|
||||
}
|
||||
|
||||
func (uploadHelper) TableName() string { return "w_uploads" }
|
||||
|
||||
type uploadStatHelper struct {
|
||||
ID uint64 `gorm:"primaryKey;autoIncrement"`
|
||||
Dimension string `gorm:"size:32;not null;uniqueIndex:idx_stat_dimension_key"`
|
||||
StatKey string `gorm:"size:100;not null;uniqueIndex:idx_stat_dimension_key"`
|
||||
FileCount int64 `gorm:"not null;default:0"`
|
||||
FileSize int64 `gorm:"not null;default:0"`
|
||||
UpdatedAt time.Time `gorm:"autoUpdateTime"`
|
||||
}
|
||||
|
||||
func (uploadStatHelper) TableName() string { return "w_upload_stats" }
|
||||
|
||||
type messageChannelHelper struct {
|
||||
ID uint64 `gorm:"primaryKey;autoIncrement"`
|
||||
Type string `gorm:"size:32;not null"`
|
||||
Name string `gorm:"size:64;not null"`
|
||||
OwnerScope string `gorm:"size:32;not null;default:'system'"`
|
||||
OwnerID *uint64
|
||||
Credentials string `gorm:"type:text;not null"`
|
||||
Extra string `gorm:"type:text"`
|
||||
Enabled bool `gorm:"default:false;not null"`
|
||||
CreatedAt time.Time `gorm:"autoCreateTime"`
|
||||
UpdatedAt time.Time `gorm:"autoUpdateTime"`
|
||||
}
|
||||
|
||||
func (messageChannelHelper) TableName() string { return "w_message_channels" }
|
||||
|
||||
type messageBindingHelper struct {
|
||||
ID uint64 `gorm:"primaryKey;autoIncrement"`
|
||||
ChannelID uint64 `gorm:"not null;index"`
|
||||
PlatformUserID string `gorm:"size:128;not null;index"`
|
||||
UserID uint64 `gorm:"not null;index"`
|
||||
CreatedAt time.Time `gorm:"autoCreateTime"`
|
||||
}
|
||||
|
||||
func (messageBindingHelper) TableName() string { return "w_message_bindings" }
|
||||
|
||||
type messagePairingHelper struct {
|
||||
ID uint64 `gorm:"primaryKey;autoIncrement"`
|
||||
Code string `gorm:"size:32;uniqueIndex;not null"`
|
||||
ChannelID uint64 `gorm:"not null;index"`
|
||||
PlatformUserID string `gorm:"size:128;not null;index"`
|
||||
UserID uint64 `gorm:"not null;index"`
|
||||
ExpiresAt time.Time `gorm:"not null;index"`
|
||||
CreatedAt time.Time `gorm:"autoCreateTime"`
|
||||
}
|
||||
|
||||
func (messagePairingHelper) TableName() string { return "w_message_pairing_codes" }
|
||||
|
||||
type taskExecutionHelper struct {
|
||||
ID uint64 `gorm:"primaryKey;autoIncrement"`
|
||||
TaskID string `gorm:"size:128;uniqueIndex;not null"`
|
||||
TaskType string `gorm:"size:64;index;not null"`
|
||||
TaskName string `gorm:"size:128"`
|
||||
Status string `gorm:"size:32;index;not null"`
|
||||
Retryable bool `gorm:"not null;default:false"`
|
||||
MaxRetry int `gorm:"not null;default:0"`
|
||||
RetryCount int `gorm:"not null;default:0"`
|
||||
Log string `gorm:"type:text"`
|
||||
ErrorMessage string `gorm:"type:text"`
|
||||
Result string `gorm:"type:text"`
|
||||
StartedAt *time.Time `gorm:"index"`
|
||||
FinishedAt *time.Time
|
||||
Duration int64
|
||||
Payload string `gorm:"type:text"`
|
||||
TriggeredBy string `gorm:"size:32;not null;default:system"`
|
||||
CreatedAt time.Time `gorm:"autoCreateTime;index"`
|
||||
UpdatedAt time.Time `gorm:"autoUpdateTime"`
|
||||
}
|
||||
|
||||
func (taskExecutionHelper) TableName() string {
|
||||
return "w_task_executions"
|
||||
}
|
||||
|
||||
type scheduleHelper struct {
|
||||
ID uint64 `gorm:"primaryKey;autoIncrement"`
|
||||
TaskType string `gorm:"size:64;uniqueIndex;not null"`
|
||||
TaskName string `gorm:"size:128;not null"`
|
||||
CronExpr string `gorm:"size:64;not null"`
|
||||
Payload string `gorm:"type:text"`
|
||||
Enabled bool `gorm:"default:true;not null"`
|
||||
CreatedAt time.Time `gorm:"autoCreateTime"`
|
||||
UpdatedAt time.Time `gorm:"autoUpdateTime"`
|
||||
}
|
||||
|
||||
func (scheduleHelper) TableName() string { return "w_schedules" }
|
||||
|
||||
const (
|
||||
configTypeSystem = "system"
|
||||
configTypeBusiness = "business"
|
||||
configValueTrue = "true"
|
||||
configValueFalse = "false"
|
||||
)
|
||||
|
||||
// SetupTestEnvironment initializes an in-memory SQLite DB, seeds default configurations,
|
||||
// starts miniredis, and overrides the global db/Redis clients. It returns a cleanup function.
|
||||
func SetupTestEnvironment(t *testing.T) (*gorm.DB, *miniredis.Miniredis, func()) {
|
||||
sqliteDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{
|
||||
DisableForeignKeyConstraintWhenMigrating: true,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("failed to open in-memory SQLite db: %v", err)
|
||||
}
|
||||
|
||||
if sqlDB, err := sqliteDB.DB(); err == nil {
|
||||
sqlDB.SetMaxOpenConns(1)
|
||||
}
|
||||
|
||||
// AutoMigrate all tables via internal test helpers to completely decouple testhelper from domain plugins
|
||||
err = sqliteDB.AutoMigrate(
|
||||
&userHelper{},
|
||||
&accessTokenHelper{},
|
||||
&authSourceHelper{},
|
||||
&externalAccountHelper{},
|
||||
&SystemConfig{},
|
||||
&uploadHelper{},
|
||||
&uploadStatHelper{},
|
||||
&taskExecutionHelper{},
|
||||
&scheduleHelper{},
|
||||
&messageChannelHelper{},
|
||||
&messageBindingHelper{},
|
||||
&messagePairingHelper{},
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to auto migrate tables: %v", err)
|
||||
}
|
||||
|
||||
mr, err := miniredis.Run()
|
||||
if err != nil {
|
||||
t.Fatalf("failed to start miniredis: %v", err)
|
||||
}
|
||||
|
||||
redisClient := redis.NewClient(&redis.Options{
|
||||
Addr: mr.Addr(),
|
||||
MaintNotificationsConfig: &maintnotifications.Config{
|
||||
Mode: maintnotifications.ModeDisabled,
|
||||
},
|
||||
})
|
||||
|
||||
db.SetDB(sqliteDB)
|
||||
cachepkg.Redis = redisClient
|
||||
|
||||
seedDefaultConfigs(t, sqliteDB)
|
||||
|
||||
cleanup := func() {
|
||||
runExtraCleanups()
|
||||
_ = redisClient.Close()
|
||||
mr.Close()
|
||||
db.SetDB(nil)
|
||||
cachepkg.Redis = nil
|
||||
}
|
||||
|
||||
return sqliteDB, mr, cleanup
|
||||
}
|
||||
|
||||
func getSeedConfigsPart1() []SystemConfig {
|
||||
return []SystemConfig{
|
||||
{Key: "upload_allowed_extensions", Value: `["jpg", "jpeg", "png", "gif", "webp", "txt", "pdf", "zip"]`, Type: configTypeSystem, Description: "允许上传的文件扩展名列表(JSON 字符串数组)"},
|
||||
{Key: "site_name", Value: "Wavelet", Type: configTypeSystem, Description: "站点名称"},
|
||||
{Key: "site_description", Value: "Lightweight and Modular Web Application Platform", Type: configTypeSystem, Description: "站点描述"},
|
||||
{Key: "password_login_enabled", Value: configValueTrue, Type: configTypeSystem, Description: "是否开启账号密码登录(true/false)"},
|
||||
{Key: "registration_enabled", Value: configValueTrue, Type: configTypeSystem, Description: "是否允许新用户注册(全局总开关,true/false)"},
|
||||
{Key: "password_register_enabled", Value: configValueTrue, Type: configTypeSystem, Description: "是否允许账号密码注册(true/false)"},
|
||||
{Key: "oidc_login_enabled", Value: configValueFalse, Type: configTypeSystem, Description: "是否开启 OIDC 登录(true/false)"},
|
||||
{Key: "server_address", Value: "http://localhost:8000", Type: configTypeSystem, Description: "服务端访问地址(用于生成绝对路径链接,多个地址用英文逗号分隔)"},
|
||||
{Key: "smtp_host", Value: "", Type: configTypeSystem, Description: "SMTP 服务器主机名或 IP"},
|
||||
{Key: "smtp_port", Value: "587", Type: configTypeSystem, Description: "SMTP 服务器端口(标准 STARTTLS 为 587,SMTPS 为 465)"},
|
||||
{Key: "smtp_username", Value: "", Type: configTypeSystem, Description: "SMTP 账户(如 sender@example.com)"},
|
||||
{Key: "smtp_password", Value: "", Type: configTypeSystem, Description: "SMTP 访问凭证(授权码/密码)"},
|
||||
{Key: "email_login_verification_enabled", Value: configValueFalse, Type: configTypeSystem, Description: "是否开启邮箱登录验证(true/false)"},
|
||||
{Key: "email_register_verification_enabled", Value: configValueFalse, Type: configTypeSystem, Description: "是否开启邮箱注册验证(true/false)"},
|
||||
{Key: "menu_display_config", Value: "{}", Type: configTypeSystem, Description: "目录显示配置(JSON 字符串,格式为 {url: enabled})"},
|
||||
{Key: "search_engine_indexing_enabled", Value: configValueFalse, Type: configTypeSystem, Description: "是否允许搜索引擎检索"},
|
||||
{Key: "file_access_whitelist", Value: `["avatar"]`, Type: configTypeSystem, Description: "免登录访问的文件业务类型白名单"},
|
||||
{Key: "disk_cache_max_size_mb", Value: "100", Type: configTypeSystem, Description: "磁盘缓存最大空间大小 (MB)"},
|
||||
{Key: "disk_cache_ttl_minutes", Value: "60", Type: configTypeSystem, Description: "磁盘缓存默认有效期 (分钟)"},
|
||||
{Key: "disk_cache_lru_enabled", Value: configValueTrue, Type: configTypeSystem, Description: "是否启用 LRU 淘汰机制"},
|
||||
{Key: "login_session_ttl_hours", Value: "0", Type: configTypeSystem, Description: "登录会话过期时间 (小时,0表示浏览器关闭后自动退出,-1表示永不过期)"},
|
||||
{Key: "update_upstream_repository", Value: "Rain-kl/Wavelet", Type: configTypeSystem, Description: "GitHub Actions Release 上游仓库"},
|
||||
{Key: "storage_config", Value: `{"driver":"local","local":{"root":"."},"s3":{"region":"us-east-1"},"r2":{"region":"auto"},"minio":{"region":"us-east-1","path_style":true},"oss":{},"webdav":{}}`, Type: configTypeSystem, Description: "文件存储驱动及连接配置(JSON)"},
|
||||
{Key: "log_database", Value: "sqlite", Type: configTypeSystem, Description: "当前日志主库"},
|
||||
{Key: "log_db_migration", Value: "", Type: configTypeSystem, Description: "日志库迁移冻结标记"},
|
||||
{Key: "log_retention_days_postgres", Value: "30", Type: configTypeBusiness, Description: "PostgreSQL 用户访问日志保留天数"},
|
||||
{Key: "log_retention_days_sqlite", Value: "30", Type: configTypeBusiness, Description: "SQLite 用户访问日志保留天数"},
|
||||
{Key: "log_retention_days_clickhouse", Value: "30", Type: configTypeBusiness, Description: "ClickHouse 用户访问日志保留天数"},
|
||||
}
|
||||
}
|
||||
|
||||
func seedDefaultConfigs(t *testing.T, tx *gorm.DB) {
|
||||
defaultConfigs := getSeedConfigsPart1()
|
||||
|
||||
if err := tx.Create(&defaultConfigs).Error; err != nil {
|
||||
t.Fatalf("failed to seed default system configs: %v", err)
|
||||
}
|
||||
|
||||
publicKeys := map[string]struct{}{
|
||||
"upload_allowed_extensions": {},
|
||||
"site_name": {},
|
||||
"password_login_enabled": {},
|
||||
"registration_enabled": {},
|
||||
"password_register_enabled": {},
|
||||
"oidc_login_enabled": {},
|
||||
}
|
||||
|
||||
keys := make([]string, 0, len(publicKeys))
|
||||
for key := range publicKeys {
|
||||
keys = append(keys, key)
|
||||
}
|
||||
if err := tx.Model(&SystemConfig{}).
|
||||
Where("key IN ?", keys).
|
||||
Update("visibility", "visible").Error; err != nil {
|
||||
t.Fatalf("failed to seed public system config visibility: %v", err)
|
||||
}
|
||||
|
||||
for _, config := range defaultConfigs {
|
||||
if _, ok := publicKeys[config.Key]; ok {
|
||||
config.Visibility = "visible"
|
||||
}
|
||||
_ = cachepkg.HSetJSON(context.Background(), "system_configs", config.Key, &config)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,15 @@
|
||||
// Copyright 2025 linux.do
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package trace 提供 OpenTelemetry 链路追踪封装工具
|
||||
package trace
|
||||
|
||||
import "go.opentelemetry.io/otel/propagation"
|
||||
|
||||
func newPropagator() propagation.TextMapPropagator {
|
||||
return propagation.NewCompositeTextMapPropagator(
|
||||
propagation.TraceContext{},
|
||||
propagation.Baggage{},
|
||||
)
|
||||
}
|
||||
@@ -0,0 +1,19 @@
|
||||
// Copyright 2025 linux.do
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package trace
|
||||
|
||||
import (
|
||||
sdktrace "go.opentelemetry.io/otel/sdk/trace"
|
||||
)
|
||||
|
||||
// ParentBasedRatioSampler 创建父级感知的概率采样器
|
||||
// - 如果父 Span 已采样,则子 Span 也采样
|
||||
// - 如果父 Span 未采样,则子 Span 也不采样
|
||||
// - 如果是根 Span,按 samplingRate 概率采样
|
||||
func ParentBasedRatioSampler(samplingRate float64) sdktrace.Sampler {
|
||||
return sdktrace.ParentBased(
|
||||
sdktrace.TraceIDRatioBased(samplingRate),
|
||||
)
|
||||
}
|
||||
@@ -0,0 +1,62 @@
|
||||
// Copyright 2025 linux.do
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package trace
|
||||
|
||||
import (
|
||||
"context"
|
||||
"log"
|
||||
|
||||
"go.opentelemetry.io/otel"
|
||||
"go.opentelemetry.io/otel/trace"
|
||||
)
|
||||
|
||||
// Tracer 全局 OpenTelemetry Tracer 实例
|
||||
var Tracer trace.Tracer
|
||||
var shutdownFuncs []func(context.Context) error
|
||||
|
||||
func init() {
|
||||
// 初始化 Propagator
|
||||
prop := newPropagator()
|
||||
otel.SetTextMapPropagator(prop)
|
||||
|
||||
// 初始化 Tracer 实例为 No-op 默认以避免未初始化前或测试环境崩溃
|
||||
Tracer = otel.GetTracerProvider().Tracer("github.com/Rain-kl/Wavelet")
|
||||
}
|
||||
|
||||
// Config 链路追踪配置
|
||||
type Config struct {
|
||||
AppName string
|
||||
SamplingRate float64
|
||||
TracerName string
|
||||
}
|
||||
|
||||
// Init 初始化 Tracer Provider 并关联全局 Tracer 实例
|
||||
func Init(cfg Config) {
|
||||
tracerProvider, err := newTracerProvider(cfg)
|
||||
if err != nil {
|
||||
log.Fatalf("[Trace] init trace provider failed: %v", err)
|
||||
}
|
||||
shutdownFuncs = append(shutdownFuncs, tracerProvider.Shutdown)
|
||||
otel.SetTracerProvider(tracerProvider)
|
||||
|
||||
// 更新 Tracer
|
||||
tracerName := cfg.TracerName
|
||||
if tracerName == "" {
|
||||
tracerName = "github.com/Rain-kl/Wavelet"
|
||||
}
|
||||
Tracer = tracerProvider.Tracer(tracerName)
|
||||
}
|
||||
|
||||
// Shutdown 关闭所有 Trace Provider
|
||||
func Shutdown(ctx context.Context) {
|
||||
for _, fn := range shutdownFuncs {
|
||||
_ = fn(ctx)
|
||||
}
|
||||
}
|
||||
|
||||
// Start 创建一个新的 Trace Span
|
||||
func Start(ctx context.Context, name string, opts ...trace.SpanStartOption) (context.Context, trace.Span) {
|
||||
return Tracer.Start(ctx, name, opts...)
|
||||
}
|
||||
@@ -0,0 +1,52 @@
|
||||
// Copyright 2025 linux.do
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package trace
|
||||
|
||||
import (
|
||||
"context"
|
||||
"os"
|
||||
|
||||
"go.opentelemetry.io/otel/attribute"
|
||||
"go.opentelemetry.io/otel/exporters/otlp/otlptrace/otlptracegrpc"
|
||||
"go.opentelemetry.io/otel/sdk/resource"
|
||||
sdktrace "go.opentelemetry.io/otel/sdk/trace"
|
||||
)
|
||||
|
||||
func newTracerProvider(cfg Config) (*sdktrace.TracerProvider, error) {
|
||||
// 获取主机名和容器信息
|
||||
hostname, err := os.Hostname()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// 业务属性不绑定 schema URL,合并时继承 resource.Default() 的 SDK 内置版本,避免 semconv 与 otel/sdk 升级不同步。
|
||||
r, err := resource.Merge(
|
||||
resource.Default(),
|
||||
resource.NewSchemaless(
|
||||
attribute.String("service.name", cfg.AppName),
|
||||
attribute.String("host.name", hostname),
|
||||
attribute.String("k8s.namespace.name", os.Getenv("KUBERNETES_NAMESPACE")),
|
||||
attribute.String("k8s.pod.name", os.Getenv("KUBERNETES_POD_NAME")),
|
||||
attribute.String("k8s.pod.uid", os.Getenv("KUBERNETES_POD_UID")),
|
||||
),
|
||||
)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// 初始化 Exporter
|
||||
traceExporter, err := otlptracegrpc.New(context.Background())
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// 初始化 Trace
|
||||
tracerProvider := sdktrace.NewTracerProvider(
|
||||
sdktrace.WithBatcher(traceExporter),
|
||||
sdktrace.WithResource(r),
|
||||
sdktrace.WithSampler(ParentBasedRatioSampler(cfg.SamplingRate)),
|
||||
)
|
||||
return tracerProvider, nil
|
||||
}
|
||||
@@ -0,0 +1,10 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package util
|
||||
|
||||
const (
|
||||
errCreateHTTPRequestFailed = "创建HTTP请求失败: %w"
|
||||
errHTTPRequestFailed = "请求%s接口失败: %w"
|
||||
errInvalidCustomValue = "invalid value: %v"
|
||||
)
|
||||
@@ -0,0 +1,160 @@
|
||||
// Copyright 2025 linux.do
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package util provides generic utility functions.
|
||||
package util
|
||||
|
||||
import (
|
||||
"crypto/aes"
|
||||
"crypto/cipher"
|
||||
"crypto/ed25519"
|
||||
"crypto/rand"
|
||||
"encoding/base64"
|
||||
"encoding/hex"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
)
|
||||
|
||||
const (
|
||||
aesKeyLength = 32
|
||||
|
||||
errInvalidSignKey = "invalid sign key: %w"
|
||||
errSignKeyLengthInvalid = "sign key must be 32 bytes (64 hex characters)"
|
||||
errCreateCipherFailed = "failed to create cipher: %w"
|
||||
errCreateGCMFailed = "failed to create GCM: %w"
|
||||
errGenerateNonceFailed = "failed to generate nonce: %w"
|
||||
errDecodeCiphertextFailed = "failed to decode ciphertext: %w"
|
||||
errCiphertextTooShort = "ciphertext too short"
|
||||
errDecryptFailed = "failed to decrypt: %w"
|
||||
)
|
||||
|
||||
// Encrypt 使用 SignKey 加密字符串数据
|
||||
// signKey: 64 字符 hex 编码的密钥(对应 32 字节,用于 AES-256)
|
||||
// plaintext: 要加密的明文字符串
|
||||
// return: base64 编码的密文
|
||||
func Encrypt(signKey string, plaintext string) (string, error) {
|
||||
return encryptBytes(signKey, []byte(plaintext))
|
||||
}
|
||||
|
||||
// Decrypt 使用 SignKey 解密字符串数据
|
||||
// signKey: 64 字符 hex 编码的密钥(对应 32 字节,用于 AES-256)
|
||||
// ciphertext: base64 编码的密文
|
||||
// return: 解密后的明文字符串
|
||||
func Decrypt(signKey string, ciphertext string) (string, error) {
|
||||
plaintext, err := decryptBytes(signKey, ciphertext)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return string(plaintext), nil
|
||||
}
|
||||
|
||||
// encryptBytes 加密函数,处理字节数据
|
||||
func encryptBytes(signKey string, plaintext []byte) (string, error) {
|
||||
// 将 hex 编码的密钥转换为字节
|
||||
key, err := hex.DecodeString(signKey)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf(errInvalidSignKey, err)
|
||||
}
|
||||
if len(key) != aesKeyLength {
|
||||
return "", errors.New(errSignKeyLengthInvalid)
|
||||
}
|
||||
|
||||
// 创建 AES cipher
|
||||
block, err := aes.NewCipher(key)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf(errCreateCipherFailed, err)
|
||||
}
|
||||
|
||||
// 使用 GCM 模式(Galois/Counter Mode)
|
||||
gcm, err := cipher.NewGCM(block)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf(errCreateGCMFailed, err)
|
||||
}
|
||||
|
||||
// 生成随机 nonce
|
||||
nonce := make([]byte, gcm.NonceSize())
|
||||
if _, err := io.ReadFull(rand.Reader, nonce); err != nil {
|
||||
return "", fmt.Errorf(errGenerateNonceFailed, err)
|
||||
}
|
||||
|
||||
// 加密数据
|
||||
ciphertext := gcm.Seal(nonce, nonce, plaintext, nil)
|
||||
|
||||
// 返回 base64 编码的密文
|
||||
return Base64Encode(ciphertext), nil
|
||||
}
|
||||
|
||||
// decryptBytes 解密函数,处理字节数据
|
||||
func decryptBytes(signKey string, ciphertext string) ([]byte, error) {
|
||||
// 将 hex 编码的密钥转换为字节
|
||||
key, err := hex.DecodeString(signKey)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf(errInvalidSignKey, err)
|
||||
}
|
||||
if len(key) != aesKeyLength {
|
||||
return nil, errors.New(errSignKeyLengthInvalid)
|
||||
}
|
||||
|
||||
// 解码 base64 密文
|
||||
data, err := Base64Decode(ciphertext)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf(errDecodeCiphertextFailed, err)
|
||||
}
|
||||
|
||||
// 创建 AES cipher
|
||||
block, err := aes.NewCipher(key)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf(errCreateCipherFailed, err)
|
||||
}
|
||||
|
||||
// 使用 GCM 模式
|
||||
gcm, err := cipher.NewGCM(block)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf(errCreateGCMFailed, err)
|
||||
}
|
||||
|
||||
// 提取 nonce
|
||||
nonceSize := gcm.NonceSize()
|
||||
if len(data) < nonceSize {
|
||||
return nil, errors.New(errCiphertextTooShort)
|
||||
}
|
||||
|
||||
nonce, ciphertextBytes := data[:nonceSize], data[nonceSize:]
|
||||
|
||||
// 解密数据
|
||||
plaintext, err := gcm.Open(nil, nonce, ciphertextBytes, nil)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf(errDecryptFailed, err)
|
||||
}
|
||||
|
||||
return plaintext, nil
|
||||
}
|
||||
|
||||
// Base64Encode Base64编码
|
||||
func Base64Encode(data []byte) string {
|
||||
return base64.StdEncoding.EncodeToString(data)
|
||||
}
|
||||
|
||||
// Base64Decode Base64解码
|
||||
func Base64Decode(encoded string) ([]byte, error) {
|
||||
return base64.StdEncoding.DecodeString(encoded)
|
||||
}
|
||||
|
||||
// Ed25519Verify 验证 Ed25519 签名
|
||||
// publicKey: 32 字节的公钥(已解码的二进制格式)
|
||||
// message: 待验证的原始消息
|
||||
// signature: 64 字节的签名(已解码的二进制格式)
|
||||
// return: 签名是否有效
|
||||
func Ed25519Verify(publicKey, message, signature []byte) bool {
|
||||
if len(publicKey) != ed25519.PublicKeySize {
|
||||
return false
|
||||
}
|
||||
|
||||
if len(signature) != ed25519.SignatureSize {
|
||||
return false
|
||||
}
|
||||
|
||||
return ed25519.Verify(publicKey, message, signature)
|
||||
}
|
||||
@@ -0,0 +1,29 @@
|
||||
// Copyright 2025 linux.do
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package util provides framework-agnostic helper types and HTTP utilities.
|
||||
package util
|
||||
|
||||
import (
|
||||
"database/sql/driver"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
)
|
||||
|
||||
// StringArray custom type for handling JSON arrays
|
||||
type StringArray []string
|
||||
|
||||
// Scan 实现 sql.Scanner 接口,从数据库读取 JSON 数组
|
||||
func (sa *StringArray) Scan(value interface{}) error {
|
||||
bytesValue, ok := value.([]byte)
|
||||
if !ok {
|
||||
return fmt.Errorf(errInvalidCustomValue, value)
|
||||
}
|
||||
return json.Unmarshal(bytesValue, sa)
|
||||
}
|
||||
|
||||
// Value 实现 driver.Valuer 接口,将 JSON 数组序列化为数据库存储值
|
||||
func (sa StringArray) Value() (driver.Value, error) {
|
||||
return json.Marshal(sa)
|
||||
}
|
||||
@@ -0,0 +1,22 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package util
|
||||
|
||||
import "github.com/gin-gonic/gin"
|
||||
|
||||
// GetFromContext retrieves a typed value from Gin context.
|
||||
func GetFromContext[T any](c *gin.Context, key string) (T, bool) {
|
||||
value, exists := c.Get(key)
|
||||
if !exists {
|
||||
var zero T
|
||||
return zero, false
|
||||
}
|
||||
typed, ok := value.(T)
|
||||
return typed, ok
|
||||
}
|
||||
|
||||
// SetToContext sets a typed value into Gin context.
|
||||
func SetToContext[T any](c *gin.Context, key string, value T) {
|
||||
c.Set(key, value)
|
||||
}
|
||||
@@ -0,0 +1,31 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package util
|
||||
|
||||
import (
|
||||
"log/slog"
|
||||
"runtime"
|
||||
"runtime/debug"
|
||||
)
|
||||
|
||||
// Go runs fn in a new goroutine and recovers panics, so a background task
|
||||
// cannot crash the whole process. The panic is logged together with the
|
||||
// util.Go call site. Use it for every fire-and-forget / long-lived
|
||||
// background goroutine; HTTP handlers are already covered by gin.Recovery.
|
||||
func Go(fn func()) {
|
||||
pc, file, line, _ := runtime.Caller(1)
|
||||
go func() {
|
||||
defer func() {
|
||||
if r := recover(); r != nil {
|
||||
slog.Error("panic recovered in background goroutine",
|
||||
"caller", runtime.FuncForPC(pc).Name(),
|
||||
"file", file,
|
||||
"line", line,
|
||||
"panic", r,
|
||||
"stack", string(debug.Stack()))
|
||||
}
|
||||
}()
|
||||
fn()
|
||||
}()
|
||||
}
|
||||
@@ -0,0 +1,38 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package util
|
||||
|
||||
import (
|
||||
"sync"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestGoRecoversPanic(t *testing.T) {
|
||||
var wg sync.WaitGroup
|
||||
wg.Add(1)
|
||||
|
||||
// Go should swallow the panic without crashing the test process
|
||||
Go(func() {
|
||||
defer wg.Done()
|
||||
panic("boom")
|
||||
})
|
||||
|
||||
wg.Wait()
|
||||
}
|
||||
|
||||
func TestGoRunsNormally(t *testing.T) {
|
||||
var wg sync.WaitGroup
|
||||
wg.Add(1)
|
||||
ran := false
|
||||
|
||||
Go(func() {
|
||||
defer wg.Done()
|
||||
ran = true
|
||||
})
|
||||
|
||||
wg.Wait()
|
||||
if !ran {
|
||||
t.Fatal("expected fn to run")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,68 @@
|
||||
// Copyright 2025 linux.do
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package util
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/backend/pkg/httppool"
|
||||
)
|
||||
|
||||
// IsLocalhost 检查 URL 是否为 localhost
|
||||
func IsLocalhost(urlStr string) bool {
|
||||
u, err := url.Parse(urlStr)
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
hostname := u.Hostname()
|
||||
return hostname == "localhost" || hostname == "127.0.0.1" || hostname == "::1"
|
||||
}
|
||||
|
||||
// HTTP 客户端配置常量
|
||||
const (
|
||||
httpClientTimeout = 10 // HTTP 客户端超时时间(秒)
|
||||
httpMaxIdleConns = 100
|
||||
httpMaxIdleConnsPerHost = 20
|
||||
httpIdleConnTimeout = 60 // 空闲连接超时(秒)
|
||||
)
|
||||
|
||||
// 配置HTTP客户端 使用 otelhttp 自动注入 trace span
|
||||
var httpClient = &http.Client{
|
||||
Timeout: httpClientTimeout * time.Second,
|
||||
Transport: httppool.DefaultTransport(),
|
||||
}
|
||||
|
||||
// SetHTTPClient 替换全局 HTTP 客户端实例
|
||||
func SetHTTPClient(c *http.Client) {
|
||||
httpClient = c
|
||||
}
|
||||
|
||||
// Request 发送 HTTP 请求,支持自定义 Headers 和 Cookies
|
||||
func Request(ctx context.Context, method, url string, body io.Reader, headers, cookies map[string]string) (*http.Response, error) {
|
||||
req, err := http.NewRequestWithContext(ctx, method, url, body)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf(errCreateHTTPRequestFailed, err)
|
||||
}
|
||||
|
||||
for key, value := range cookies {
|
||||
req.AddCookie(&http.Cookie{Name: key, Value: value}) //nolint:gosec // client-side cookies do not require server attributes (Secure/HttpOnly)
|
||||
}
|
||||
|
||||
for key, value := range headers {
|
||||
req.Header.Set(key, value)
|
||||
}
|
||||
|
||||
resp, err := httpClient.Do(req)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf(errHTTPRequestFailed, url, err)
|
||||
}
|
||||
|
||||
return resp, nil
|
||||
}
|
||||
@@ -0,0 +1,16 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package util
|
||||
|
||||
import "strings"
|
||||
|
||||
var likeEscaper = strings.NewReplacer(`\`, `\\`, `%`, `\%`, `_`, `\_`)
|
||||
|
||||
// EscapeLike escapes SQL LIKE metacharacters (\, %, _) so a user-supplied
|
||||
// value matches literally in LIKE patterns. Pair it with an explicit
|
||||
// `ESCAPE '\'` clause where the dialect has no backslash default (SQLite);
|
||||
// PostgreSQL and ClickHouse treat backslash as the default LIKE escape.
|
||||
func EscapeLike(value string) string {
|
||||
return likeEscaper.Replace(value)
|
||||
}
|
||||
@@ -0,0 +1,22 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package util
|
||||
|
||||
import "testing"
|
||||
|
||||
func TestEscapeLike(t *testing.T) {
|
||||
cases := map[string]string{
|
||||
"": "",
|
||||
"/my_page": `/my\_page`,
|
||||
"100%": `100\%`,
|
||||
`a\b`: `a\\b`,
|
||||
`%_\`: `\%\_\\`,
|
||||
"normal/path": "normal/path",
|
||||
}
|
||||
for input, want := range cases {
|
||||
if got := EscapeLike(input); got != want {
|
||||
t.Errorf("EscapeLike(%q) = %q, want %q", input, got, want)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,48 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package util
|
||||
|
||||
import (
|
||||
"sync"
|
||||
|
||||
"golang.org/x/crypto/bcrypt"
|
||||
)
|
||||
|
||||
// HashPassword 使用 bcrypt 对密码进行哈希处理
|
||||
func HashPassword(password string) (string, error) {
|
||||
hash, err := bcrypt.GenerateFromPassword([]byte(password), bcrypt.DefaultCost)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return string(hash), nil
|
||||
}
|
||||
|
||||
var dummyPasswordHashOnce sync.Once
|
||||
var dummyPasswordHash string
|
||||
|
||||
func dummyHash() string {
|
||||
dummyPasswordHashOnce.Do(func() {
|
||||
hash, err := bcrypt.GenerateFromPassword([]byte("x"), bcrypt.DefaultCost)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
dummyPasswordHash = string(hash)
|
||||
})
|
||||
return dummyPasswordHash
|
||||
}
|
||||
|
||||
// CheckPasswordHash 比较 bcrypt 哈希值与明文密码是否匹配
|
||||
func CheckPasswordHash(hash, password string) bool {
|
||||
return bcrypt.CompareHashAndPassword([]byte(hash), []byte(password)) == nil
|
||||
}
|
||||
|
||||
// DummyCheckPassword runs a bcrypt compare against a dummy hash so missing-user
|
||||
// login failures take a similar amount of time as a real password miss.
|
||||
func DummyCheckPassword(password string) {
|
||||
hash := dummyHash()
|
||||
if hash == "" {
|
||||
return
|
||||
}
|
||||
_ = CheckPasswordHash(hash, password)
|
||||
}
|
||||
@@ -0,0 +1,23 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package util
|
||||
|
||||
import "testing"
|
||||
|
||||
func TestDummyCheckPasswordDoesNotPanic(t *testing.T) {
|
||||
DummyCheckPassword("any-password")
|
||||
}
|
||||
|
||||
func TestCheckPasswordHashRoundTrip(t *testing.T) {
|
||||
hash, err := HashPassword("secret-pass")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !CheckPasswordHash(hash, "secret-pass") {
|
||||
t.Fatal("expected matching password to succeed")
|
||||
}
|
||||
if CheckPasswordHash(hash, "other-pass") {
|
||||
t.Fatal("expected mismatched password to fail")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,35 @@
|
||||
// Copyright 2025 linux.do
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package util
|
||||
|
||||
import "strings"
|
||||
|
||||
// emailPartsCount 邮箱地址由 @ 分割为两部分
|
||||
const (
|
||||
emailPartsCount = 2
|
||||
emailLocalMinChars = 2 // 邮箱 local 部分掩码显示的最小字符数
|
||||
)
|
||||
|
||||
// DerefString 安全地解引用字符串指针,nil 返回空字符串
|
||||
func DerefString(s *string) string {
|
||||
if s == nil {
|
||||
return ""
|
||||
}
|
||||
return *s
|
||||
}
|
||||
|
||||
// MaskEmail 安全脱敏用户的邮箱地址(例如 us***@example.com)
|
||||
func MaskEmail(email string) string {
|
||||
parts := strings.Split(email, "@")
|
||||
if len(parts) != emailPartsCount {
|
||||
return email
|
||||
}
|
||||
local := parts[0]
|
||||
domain := parts[1]
|
||||
if len(local) <= emailLocalMinChars {
|
||||
return "**@" + domain
|
||||
}
|
||||
return local[:2] + "***" + local[len(local)-1:] + "@" + domain
|
||||
}
|
||||
@@ -0,0 +1,29 @@
|
||||
// Copyright 2025 linux.do
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package util
|
||||
|
||||
import (
|
||||
"crypto/rand"
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"io"
|
||||
|
||||
"github.com/google/uuid"
|
||||
)
|
||||
|
||||
// uniqueIDBytes 生成唯一 ID 所需的随机字节长度
|
||||
const uniqueIDBytes = 32
|
||||
|
||||
// GenerateUniqueIDSimple 生成 64 位唯一标识符
|
||||
func GenerateUniqueIDSimple() string {
|
||||
randomBytes := make([]byte, uniqueIDBytes)
|
||||
if _, err := io.ReadFull(rand.Reader, randomBytes); err != nil {
|
||||
// 如果随机数生成失败,使用 UUID 作为后备
|
||||
uuidBytes := []byte(uuid.NewString())
|
||||
hash := sha256.Sum256(uuidBytes)
|
||||
copy(randomBytes, hash[:])
|
||||
}
|
||||
return hex.EncodeToString(randomBytes)
|
||||
}
|
||||
@@ -0,0 +1,82 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package admin
|
||||
|
||||
// 管理后台公共错误常量
|
||||
const (
|
||||
AdminRequired = "未经授权访问"
|
||||
TokenAdminRequired = "该访问令牌没有管理员权限,无法访问管理端点" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
|
||||
InvalidAuthSourceID = "认证源 ID 无效"
|
||||
InvalidCursorParam = "无效的 cursor 参数"
|
||||
InvalidTaskExecutionID = "无效的任务执行记录 ID"
|
||||
)
|
||||
|
||||
// 系统配置错误消息常量
|
||||
const (
|
||||
SystemConfigNotFound = "系统配置不存在"
|
||||
ConfigKeyRequired = "配置键不能为空"
|
||||
ConfigValueRequired = "配置值不能为空"
|
||||
ConfigKeyExists = "配置键已存在"
|
||||
protectedConfigKeyMessage = "该配置项由系统任务管理,禁止手动修改"
|
||||
StorageDriverSwitchRequiresMigration = "存在存量文件,请通过存储迁移任务切换存储引擎"
|
||||
)
|
||||
|
||||
// 模板管理相关错误消息常量
|
||||
const (
|
||||
TemplateNotFound = "模板不存在"
|
||||
TemplateKeyRequired = "模板标识符不能为空"
|
||||
TemplateNameRequired = "模板名称不能为空"
|
||||
TemplateContentRequired = "模板内容不能为空"
|
||||
TemplateKeyExists = "模板标识符已存在"
|
||||
SystemTemplateCannotDelete = "系统预置模板不可删除"
|
||||
SystemTemplateCannotModifyKey = "系统预置模板不可修改标识符"
|
||||
)
|
||||
|
||||
// 任务调度相关错误消息常量
|
||||
const (
|
||||
InvalidTaskType = "无效的任务类型"
|
||||
InvalidTimeRange = "无效的时间范围"
|
||||
TaskDispatchFailed = "任务下发失败"
|
||||
UserIDRequired = "用户ID必填"
|
||||
TaskNotFound = "任务执行记录不存在"
|
||||
TaskNotRetryable = "该任务不支持重试"
|
||||
TaskNotFailed = "只有失败的任务才能重试"
|
||||
TaskMaxRetryExceeded = "已达到最大重试次数"
|
||||
TaskRetryFailed = "任务重试失败"
|
||||
InvalidCronExpression = "无效的 Cron 表达式"
|
||||
ScheduleNotFound = "定时任务不存在"
|
||||
ScheduleSaveFailed = "保存定时任务失败"
|
||||
ScheduleDeleteFailed = "删除定时任务失败"
|
||||
)
|
||||
|
||||
// 应用更新相关错误消息常量
|
||||
const (
|
||||
errInvalidRepository = "上游仓库地址无效"
|
||||
errReleaseRequestFailed = "获取上游版本失败"
|
||||
errReleaseResponseInvalid = "上游版本响应无效"
|
||||
errNoCompatibleRelease = "未找到兼容的 Release"
|
||||
errNoCompatibleAsset = "未找到当前系统对应的 Release 资产"
|
||||
errDevelopmentBuild = "开发版本无法执行自动升级"
|
||||
errAlreadyUpToDate = "当前已是最新版本"
|
||||
errUpgradeAlreadyRunning = "已有升级任务正在执行"
|
||||
errAutomaticUpgradeBlocked = "当前平台暂不支持自动替换二进制"
|
||||
)
|
||||
|
||||
// 用户管理(管理员视角)错误消息常量
|
||||
const (
|
||||
userNotFound = "用户不存在"
|
||||
cannotDisable = "不能禁用管理员账号"
|
||||
cannotDelete = "不能删除管理员账号"
|
||||
cannotDeleteSelf = "不能删除当前登录账号"
|
||||
usernameRequired = "用户名不能为空"
|
||||
emailRequired = "邮箱不能为空"
|
||||
//nolint:gosec // error message, not hardcoded credentials
|
||||
passwordTooShort = "密码长度不能少于 8 位"
|
||||
usernameExists = "用户名已存在"
|
||||
emailExists = "邮箱已被使用"
|
||||
cannotRevokeSelfAdmin = "不能取消自身的管理员权限"
|
||||
updateUserFailed = "更新用户状态失败"
|
||||
deleteUserFailed = "删除用户失败"
|
||||
updateUserInfoFailed = "更新用户信息失败"
|
||||
)
|
||||
@@ -0,0 +1,157 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package admin
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"strconv"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/backend/pkg/response"
|
||||
"github.com/Rain-kl/Wavelet/backend/plugins/domain/auth"
|
||||
"github.com/Rain-kl/Wavelet/backend/plugins/infra/database"
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
// ListAuthSources lists all configured authentication sources.
|
||||
func ListAuthSources(c *gin.Context) {
|
||||
var sources []auth.AuthSource
|
||||
gormDB := database.DB(c.Request.Context())
|
||||
if err := gormDB.Order("id ASC").Find(&sources).Error; err != nil {
|
||||
response.AbortInternal(c, "获取认证源列表失败")
|
||||
return
|
||||
}
|
||||
|
||||
views := make([]auth.AuthSourceView, len(sources))
|
||||
for i := range sources {
|
||||
views[i] = auth.AuthSourceView{
|
||||
ID: sources[i].ID,
|
||||
Name: sources[i].Name,
|
||||
Type: sources[i].Type,
|
||||
DisplayName: sources[i].DisplayName,
|
||||
IsActive: sources[i].IsActive,
|
||||
IconURL: sources[i].IconURL,
|
||||
ClientSecretConfigured: sources[i].ClientSecret != "",
|
||||
}
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, response.OK(views))
|
||||
}
|
||||
|
||||
// CreateAuthSource creates a new authentication source.
|
||||
func CreateAuthSource(c *gin.Context) {
|
||||
var source auth.AuthSource
|
||||
if err := c.ShouldBindJSON(&source); err != nil {
|
||||
response.AbortBadRequest(c, "无效的参数")
|
||||
return
|
||||
}
|
||||
|
||||
if err := source.Validate(); err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
gormDB := database.DB(c.Request.Context())
|
||||
if err := gormDB.Create(&source).Error; err != nil {
|
||||
response.AbortBadRequest(c, "创建认证源失败: "+err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
source.Sanitize()
|
||||
c.JSON(http.StatusOK, response.OK(source))
|
||||
}
|
||||
|
||||
// UpdateAuthSource updates an authentication source.
|
||||
func UpdateAuthSource(c *gin.Context) {
|
||||
idStr := c.Param("id")
|
||||
id, err := strconv.ParseUint(idStr, 10, 64)
|
||||
if err != nil {
|
||||
response.AbortBadRequest(c, "无效的认证源 ID")
|
||||
return
|
||||
}
|
||||
|
||||
gormDB := database.DB(c.Request.Context())
|
||||
var existing auth.AuthSource
|
||||
if err := gormDB.First(&existing, id).Error; err != nil {
|
||||
response.AbortNotFound(c, "认证源不存在")
|
||||
return
|
||||
}
|
||||
|
||||
var req auth.AuthSource
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
response.AbortBadRequest(c, "无效的参数")
|
||||
return
|
||||
}
|
||||
|
||||
existing.DisplayName = req.DisplayName
|
||||
existing.ClientID = req.ClientID
|
||||
if req.ClientSecret != "" {
|
||||
existing.ClientSecret = req.ClientSecret
|
||||
}
|
||||
existing.OpenIDDiscoveryURL = req.OpenIDDiscoveryURL
|
||||
existing.Scopes = req.Scopes
|
||||
existing.IconURL = req.IconURL
|
||||
|
||||
if err := existing.Validate(); err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
if err := gormDB.Save(&existing).Error; err != nil {
|
||||
response.AbortInternal(c, "更新认证源失败")
|
||||
return
|
||||
}
|
||||
|
||||
existing.Sanitize()
|
||||
c.JSON(http.StatusOK, response.OK(existing))
|
||||
}
|
||||
|
||||
// ToggleAuthSource toggles the active state of an auth source.
|
||||
func ToggleAuthSource(c *gin.Context) {
|
||||
idStr := c.Param("id")
|
||||
id, err := strconv.ParseUint(idStr, 10, 64)
|
||||
if err != nil {
|
||||
response.AbortBadRequest(c, "无效的认证源 ID")
|
||||
return
|
||||
}
|
||||
|
||||
gormDB := database.DB(c.Request.Context())
|
||||
var existing auth.AuthSource
|
||||
if err := gormDB.First(&existing, id).Error; err != nil {
|
||||
response.AbortNotFound(c, "认证源不存在")
|
||||
return
|
||||
}
|
||||
|
||||
existing.IsActive = !existing.IsActive
|
||||
if existing.IsActive {
|
||||
if err := existing.Validate(); err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
if err := gormDB.Model(&existing).Update("is_active", existing.IsActive).Error; err != nil {
|
||||
response.AbortInternal(c, "切换认证源状态失败")
|
||||
return
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, response.OK(gin.H{"is_active": existing.IsActive}))
|
||||
}
|
||||
|
||||
// DeleteAuthSource deletes an authentication source.
|
||||
func DeleteAuthSource(c *gin.Context) {
|
||||
idStr := c.Param("id")
|
||||
id, err := strconv.ParseUint(idStr, 10, 64)
|
||||
if err != nil {
|
||||
response.AbortBadRequest(c, "无效的认证源 ID")
|
||||
return
|
||||
}
|
||||
|
||||
gormDB := database.DB(c.Request.Context())
|
||||
if err := gormDB.Delete(&auth.AuthSource{}, id).Error; err != nil {
|
||||
response.AbortInternal(c, "删除认证源失败")
|
||||
return
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, response.OKNil())
|
||||
}
|
||||
@@ -0,0 +1,103 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package admin
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/http"
|
||||
"strconv"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/backend/pkg/response"
|
||||
"github.com/Rain-kl/Wavelet/backend/plugins/infra/storage/diskcache"
|
||||
)
|
||||
|
||||
type updateCacheConfigRequest struct {
|
||||
MaxSizeMB int64 `json:"max_size_mb" binding:"required,min=1"`
|
||||
TTLMinutes int64 `json:"ttl_minutes" binding:"required,min=0"`
|
||||
LRUEnabled bool `json:"lru_enabled"`
|
||||
}
|
||||
|
||||
// GetCacheStatus 获取磁盘缓存状态与当前统计数据
|
||||
// @Summary 获取缓存状态
|
||||
// @Description 获取当前系统磁盘缓存的使用情况(已占用字节、Key 数量等)与策略配置
|
||||
// @Tags admin
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Success 200 {object} response.Any{data=diskcache.Status} "获取成功"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 403 {object} response.Any "无管理员权限"
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/admin/cache/status [get]
|
||||
func GetCacheStatus(c *gin.Context) {
|
||||
status := diskcache.GetGlobalCache().Status()
|
||||
c.JSON(http.StatusOK, response.OK(status))
|
||||
}
|
||||
|
||||
// UpdateCacheConfig 更新磁盘缓存策略配置
|
||||
// @Summary 更新缓存配置
|
||||
// @Description 更改磁盘缓存最大容量限制、文件生存时间(TTL)以及是否启用 LRU 淘汰淘汰算法,并进行热更新
|
||||
// @Tags admin
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Param request body updateCacheConfigRequest true "缓存配置请求体"
|
||||
// @Security SessionCookie
|
||||
// @Success 200 {object} response.Any "更新成功"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 403 {object} response.Any "无管理员权限"
|
||||
// @Failure 500 {object} response.Any "服务内部错误"
|
||||
// @Router /api/v1/admin/cache/config [post]
|
||||
func UpdateCacheConfig(c *gin.Context) {
|
||||
var req updateCacheConfigRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
ctx := c.Request.Context()
|
||||
|
||||
if err := saveOrUpdateCacheConfig(ctx, ConfigKeyDiskCacheMaxSizeMB, strconv.FormatInt(req.MaxSizeMB, 10)); err != nil {
|
||||
response.AbortInternal(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
if err := saveOrUpdateCacheConfig(ctx, ConfigKeyDiskCacheTTLMinutes, strconv.FormatInt(req.TTLMinutes, 10)); err != nil {
|
||||
response.AbortInternal(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
if err := saveOrUpdateCacheConfig(ctx, ConfigKeyDiskCacheLRUEnabled, strconv.FormatBool(req.LRUEnabled)); err != nil {
|
||||
response.AbortInternal(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
diskcache.GetGlobalCache().ReloadConfig(ctx)
|
||||
|
||||
c.JSON(http.StatusOK, response.OKNil())
|
||||
}
|
||||
|
||||
// ClearCache 一键清空所有磁盘缓存数据
|
||||
// @Summary 清空缓存
|
||||
// @Description 清除系统磁盘缓存目录中的所有临时文件,并重置缓存容量和 Key 追踪数据
|
||||
// @Tags admin
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Success 200 {object} response.Any "清理成功"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 403 {object} response.Any "无管理员权限"
|
||||
// @Failure 500 {object} response.Any "服务内部错误"
|
||||
// @Router /api/v1/admin/cache/clear [post]
|
||||
func ClearCache(c *gin.Context) {
|
||||
if err := diskcache.GetGlobalCache().Clear(); err != nil {
|
||||
response.AbortInternal(c, err.Error())
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OKNil())
|
||||
}
|
||||
|
||||
func saveOrUpdateCacheConfig(ctx context.Context, key, value string) error {
|
||||
return SaveOrUpdateSystemConfig(ctx, key, value)
|
||||
}
|
||||
@@ -0,0 +1,533 @@
|
||||
// Copyright 2025 linux.do
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package admin
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/backend/pkg/logger"
|
||||
mail "github.com/Rain-kl/Wavelet/backend/pkg/mail"
|
||||
"github.com/Rain-kl/Wavelet/backend/pkg/response"
|
||||
"github.com/Rain-kl/Wavelet/backend/plugins/domain/cap"
|
||||
cachepkg "github.com/Rain-kl/Wavelet/backend/plugins/infra/cache"
|
||||
db "github.com/Rain-kl/Wavelet/backend/plugins/infra/database"
|
||||
"github.com/Rain-kl/Wavelet/backend/plugins/infra/storage/objectstore"
|
||||
"github.com/gin-gonic/gin"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
const maskedConfigValue = "******"
|
||||
|
||||
// CreateSystemConfigRequest 创建系统配置请求
|
||||
type CreateSystemConfigRequest struct {
|
||||
Key string `json:"key" binding:"required,max=64"`
|
||||
Value string `json:"value" binding:"required"`
|
||||
Type string `json:"type" binding:"required,oneof=system business"`
|
||||
Visibility int `json:"visibility" binding:"oneof=0 1"`
|
||||
Description string `json:"description" binding:"max=255"`
|
||||
}
|
||||
|
||||
// UpdateSystemConfigRequest 更新系统配置请求
|
||||
type UpdateSystemConfigRequest struct {
|
||||
Value string `json:"value" binding:"required"`
|
||||
Visibility *int `json:"visibility" binding:"omitempty,oneof=0 1"`
|
||||
Description string `json:"description" binding:"max=255"`
|
||||
}
|
||||
|
||||
// GetPublicConfig 获取公共配置
|
||||
// @Summary 获取公共配置
|
||||
// @Description 返回系统配置表中 visibility 为 1 的配置键值集合
|
||||
// @Tags config
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Success 200 {object} response.Any
|
||||
// @Router /api/v1/config/public [get]
|
||||
func GetPublicConfig(c *gin.Context) {
|
||||
ctx := c.Request.Context()
|
||||
configs, err := ListVisibleSystemConfigs(ctx)
|
||||
if err != nil {
|
||||
response.AbortInternal(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
resp := make(map[string]string, len(configs))
|
||||
for _, config := range configs {
|
||||
resp[config.Key] = config.Value
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, response.OK(resp))
|
||||
}
|
||||
|
||||
// GetRobotsTXT 动态生成 robots.txt
|
||||
// @Summary 获取 robots.txt
|
||||
// @Description 根据系统配置决定是否允许搜索引擎检索,并返回相应的 robots.txt 文件内容
|
||||
// @Tags config
|
||||
// @Produce text/plain
|
||||
// @Success 200 {string} string "robots.txt 内容"
|
||||
// @Router /robots.txt [get]
|
||||
func GetRobotsTXT(c *gin.Context) {
|
||||
ctx := c.Request.Context()
|
||||
enabled, err := GetBoolByKey(ctx, ConfigKeySearchEngineIndexingEnabled)
|
||||
content := "User-Agent: *\nDisallow: /\n"
|
||||
if err == nil && enabled {
|
||||
content = "User-Agent: *\nAllow: /\n"
|
||||
}
|
||||
c.Data(http.StatusOK, "text/plain; charset=utf-8", []byte(content))
|
||||
}
|
||||
|
||||
// CreateSystemConfig 创建系统配置
|
||||
// @Summary 创建系统配置
|
||||
// @Description 创建一条新的系统配置项,配置键不可重复,同时将新配置同步到 Redis,需要管理员权限
|
||||
// @Tags admin
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param request body CreateSystemConfigRequest true "创建请求参数"
|
||||
// @Success 200 {object} response.Any{data=string} "创建成功"
|
||||
// @Failure 400 {object} response.Any "参数错误或配置键已存在"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 403 {object} response.Any "无管理员权限"
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/admin/system-configs [post]
|
||||
func CreateSystemConfig(c *gin.Context) {
|
||||
var req CreateSystemConfigRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
if isProtectedConfigKey(req.Key) {
|
||||
response.AbortBadRequest(c, protectedConfigKeyMessage)
|
||||
return
|
||||
}
|
||||
|
||||
if err := createSystemConfig(c.Request.Context(), req); err != nil {
|
||||
if err.Error() == ConfigKeyExists {
|
||||
response.AbortBadRequest(c, ConfigKeyExists)
|
||||
return
|
||||
}
|
||||
response.AbortInternal(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, response.OKNil())
|
||||
}
|
||||
|
||||
// ListSystemConfigs 获取系统配置列表
|
||||
// @Summary 获取系统配置列表
|
||||
// @Description 返回所有系统配置列表,支持按配置类型(system/business)过滤,需要管理员权限
|
||||
// @Tags admin
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param type query string false "配置类型(system/business)"
|
||||
// @Success 200 {object} response.Any{data=[]SystemConfig} "系统配置列表"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 403 {object} response.Any "无管理员权限"
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/admin/system-configs [get]
|
||||
func ListSystemConfigs(c *gin.Context) {
|
||||
configs, err := listSystemConfigs(c.Request.Context(), c.Query("type"))
|
||||
if err != nil {
|
||||
response.AbortInternal(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
for i := range configs {
|
||||
configs[i].Value = maskSensitiveConfig(configs[i].Key, configs[i].Value)
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, response.OK(configs))
|
||||
}
|
||||
|
||||
// GetSystemConfig 获取单个系统配置
|
||||
// @Summary 获取单个系统配置
|
||||
// @Description 根据配置键获取对应的系统配置详情,需要管理员权限
|
||||
// @Tags admin
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param key path string true "配置键"
|
||||
// @Success 200 {object} response.Any{data=SystemConfig} "系统配置详情"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 403 {object} response.Any "无管理员权限"
|
||||
// @Failure 404 {object} response.Any "配置不存在"
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/admin/system-configs/{key} [get]
|
||||
func GetSystemConfig(c *gin.Context) {
|
||||
config, err := getSystemConfig(c.Request.Context(), c.Param("key"))
|
||||
if err != nil {
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
response.AbortNotFound(c, SystemConfigNotFound)
|
||||
} else {
|
||||
response.AbortInternal(c, err.Error())
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
config.Value = maskSensitiveConfig(config.Key, config.Value)
|
||||
|
||||
c.JSON(http.StatusOK, response.OK(config))
|
||||
}
|
||||
|
||||
// UpdateSystemConfig 更新系统配置
|
||||
// @Summary 更新系统配置
|
||||
// @Description 根据配置键更新对应的配置内容,同时将更新同步到 Redis,需要管理员权限
|
||||
// @Tags admin
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param key path string true "配置键"
|
||||
// @Param request body UpdateSystemConfigRequest true "更新请求参数"
|
||||
// @Success 200 {object} response.Any{data=string} "更新成功"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 403 {object} response.Any "无管理员权限"
|
||||
// @Failure 404 {object} response.Any "配置不存在"
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/admin/system-configs/{key} [put]
|
||||
func UpdateSystemConfig(c *gin.Context) {
|
||||
var req UpdateSystemConfigRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
key := c.Param("key")
|
||||
if isProtectedConfigKey(key) {
|
||||
response.AbortBadRequest(c, protectedConfigKeyMessage)
|
||||
return
|
||||
}
|
||||
if err := updateSystemConfig(c.Request.Context(), key, req); err != nil {
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
response.AbortNotFound(c, SystemConfigNotFound)
|
||||
return
|
||||
}
|
||||
if isStorageConfigValidationError(err) {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
response.AbortInternal(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, response.OKNil())
|
||||
}
|
||||
|
||||
func isProtectedConfigKey(key string) bool {
|
||||
return key == ConfigKeyLogDatabase || key == ConfigKeyLogDBMigration
|
||||
}
|
||||
|
||||
func createSystemConfig(ctx context.Context, req CreateSystemConfigRequest) error {
|
||||
if isProtectedConfigKey(req.Key) {
|
||||
return errors.New(protectedConfigKeyMessage)
|
||||
}
|
||||
exists, err := SystemConfigExists(ctx, req.Key)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if exists {
|
||||
return errors.New(ConfigKeyExists)
|
||||
}
|
||||
|
||||
config := SystemConfig{
|
||||
Key: req.Key,
|
||||
Value: req.Value,
|
||||
Type: req.Type,
|
||||
Visibility: req.Visibility,
|
||||
Description: req.Description,
|
||||
}
|
||||
if err := CreateSystemConfigRecord(ctx, &config); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
invalidateSystemConfigCaches(ctx, req.Key)
|
||||
if err := InvalidateVisibleSystemConfigsCache(ctx); err != nil {
|
||||
logger.WarnF(ctx, "清理公共配置列表缓存失败: %v", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func listSystemConfigs(ctx context.Context, configType string) ([]SystemConfig, error) {
|
||||
return ListAdminSystemConfigs(ctx, configType)
|
||||
}
|
||||
|
||||
func getSystemConfig(ctx context.Context, key string) (SystemConfig, error) {
|
||||
return GetAdminSystemConfigByKey(ctx, key)
|
||||
}
|
||||
|
||||
func updateSystemConfig(ctx context.Context, key string, req UpdateSystemConfigRequest) error {
|
||||
if isProtectedConfigKey(key) {
|
||||
return errors.New(protectedConfigKeyMessage)
|
||||
}
|
||||
config, err := GetAdminSystemConfigByKey(ctx, key)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
var originalDriver objectstore.Driver
|
||||
if key == ConfigKeyStorageConfig {
|
||||
var currentCfg objectstore.Config
|
||||
if err := json.Unmarshal([]byte(config.Value), ¤tCfg); err == nil {
|
||||
originalDriver = currentCfg.Driver
|
||||
}
|
||||
|
||||
validatedVal, err := validateAndMergeStorageConfig(ctx, req.Value, config.Value)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
req.Value = validatedVal
|
||||
}
|
||||
|
||||
if err := db.DB(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
updates := map[string]any{
|
||||
"description": req.Description,
|
||||
}
|
||||
if req.Visibility != nil {
|
||||
updates["visibility"] = *req.Visibility
|
||||
config.Visibility = *req.Visibility
|
||||
}
|
||||
if key != ConfigKeySMTPPassword || req.Value != maskedConfigValue {
|
||||
updates["value"] = req.Value
|
||||
config.Value = req.Value
|
||||
}
|
||||
if err := tx.Model(&config).Updates(updates).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
resolveStorageMigrationTasksOnDirectDriverUpdate(ctx, tx, key, originalDriver, req.Value)
|
||||
return nil
|
||||
}); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
invalidateCachesAfterConfigUpdate(ctx, key)
|
||||
return nil
|
||||
}
|
||||
|
||||
func resolveStorageMigrationTasksOnDirectDriverUpdate(
|
||||
ctx context.Context,
|
||||
tx *gorm.DB,
|
||||
key string,
|
||||
originalDriver objectstore.Driver,
|
||||
newValue string,
|
||||
) {
|
||||
if key != ConfigKeyStorageConfig || originalDriver == "" {
|
||||
return
|
||||
}
|
||||
|
||||
var newCfg objectstore.Config
|
||||
if err := json.Unmarshal([]byte(newValue), &newCfg); err != nil {
|
||||
return
|
||||
}
|
||||
if newCfg.Driver != originalDriver {
|
||||
return
|
||||
}
|
||||
|
||||
if err := MarkFailedTaskExecutionsSucceededTx(
|
||||
tx,
|
||||
"storage:migrate",
|
||||
"存储配置直接更新,故障迁移任务自动标记为已解决",
|
||||
time.Now(),
|
||||
); err != nil {
|
||||
logger.ErrorF(ctx, "自动更新迁移任务状态失败: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func invalidateSystemConfigCaches(ctx context.Context, key string) {
|
||||
if err := InvalidateSystemConfigCache(ctx, key); err != nil {
|
||||
logger.WarnF(ctx, "清理系统配置缓存失败: %v", err)
|
||||
}
|
||||
if cap.IsRuntimeConfigKey(key) {
|
||||
cap.InvalidateRuntimeSettings()
|
||||
}
|
||||
}
|
||||
|
||||
func invalidateCachesAfterConfigUpdate(ctx context.Context, key string) {
|
||||
invalidateSystemConfigCaches(ctx, key)
|
||||
|
||||
if key == ConfigKeyStorageConfig {
|
||||
if cachepkg.Redis != nil {
|
||||
_ = cachepkg.Redis.Publish(ctx, "upload:access_cache:invalidate", "reset").Err()
|
||||
}
|
||||
objectstore.ResetCache()
|
||||
objectstore.PublishCacheInvalidation(ctx)
|
||||
}
|
||||
if key == ConfigKeyFileAccessWhitelist {
|
||||
if cachepkg.Redis != nil {
|
||||
_ = cachepkg.Redis.Publish(ctx, "upload:access_cache:invalidate", "reset").Err()
|
||||
}
|
||||
}
|
||||
|
||||
if err := InvalidateVisibleSystemConfigsCache(ctx); err != nil {
|
||||
logger.WarnF(ctx, "清理公共配置列表缓存失败: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// TestSMTPRequest 测试 SMTP 配置请求
|
||||
type TestSMTPRequest struct {
|
||||
SMTPHost string `json:"smtp_host" binding:"required,max=255"`
|
||||
SMTPPort int `json:"smtp_port" binding:"required"`
|
||||
SMTPUsername string `json:"smtp_username" binding:"required,max=255"`
|
||||
SMTPPassword string `json:"smtp_password" binding:"required,max=255"`
|
||||
To string `json:"to" binding:"required,email"`
|
||||
}
|
||||
|
||||
// TestSMTPResponse 测试 SMTP 配置响应
|
||||
type TestSMTPResponse struct {
|
||||
Success bool `json:"success"`
|
||||
Log string `json:"log"`
|
||||
Error string `json:"error"`
|
||||
}
|
||||
|
||||
// TestSMTP 测试 SMTP 邮件发送
|
||||
// @Summary 测试 SMTP 邮件发送
|
||||
// @Description 使用传入的配置进行 SMTP 邮件发送测试,支持使用 ****** 占位符使用保存的数据库密码
|
||||
// @Tags admin
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param request body TestSMTPRequest true "测试请求参数"
|
||||
// @Success 200 {object} response.Any{data=TestSMTPResponse} "测试执行完毕"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Router /api/v1/admin/system-configs/smtp/test [post]
|
||||
func TestSMTP(c *gin.Context) {
|
||||
var req TestSMTPRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
password := req.SMTPPassword
|
||||
if password == maskedConfigValue {
|
||||
if sc, err := GetSystemConfigByKey(c.Request.Context(), ConfigKeySMTPPassword); err == nil {
|
||||
password = sc.Value
|
||||
}
|
||||
}
|
||||
|
||||
cfg := mail.Config{
|
||||
Host: req.SMTPHost,
|
||||
Port: req.SMTPPort,
|
||||
Username: req.SMTPUsername,
|
||||
Password: password,
|
||||
}
|
||||
|
||||
subject := "Wavelet SMTP Test Mail"
|
||||
body := `<h3>SMTP Mail Connection Test</h3>
|
||||
<p>If you received this message, your SMTP configuration is correct and mail sending is working properly.</p>
|
||||
<p>Sent from Wavelet.</p>`
|
||||
|
||||
logs, err := mail.SendMailWithLog(c.Request.Context(), cfg, req.To, subject, body)
|
||||
resp := TestSMTPResponse{
|
||||
Success: err == nil,
|
||||
Log: logs,
|
||||
}
|
||||
if err != nil {
|
||||
resp.Error = err.Error()
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, response.OK(resp))
|
||||
}
|
||||
|
||||
func isStorageConfigValidationError(err error) bool {
|
||||
msg := err.Error()
|
||||
return msg == StorageDriverSwitchRequiresMigration ||
|
||||
strings.HasPrefix(msg, "解析") ||
|
||||
strings.HasPrefix(msg, "验证") ||
|
||||
strings.HasPrefix(msg, "初始化测试") ||
|
||||
strings.HasPrefix(msg, "存储连通性") ||
|
||||
strings.HasPrefix(msg, "序列化") ||
|
||||
strings.HasPrefix(msg, "检查存量文件")
|
||||
}
|
||||
|
||||
func maskSensitiveConfig(key, value string) string {
|
||||
if value == "" {
|
||||
return value
|
||||
}
|
||||
switch key {
|
||||
case ConfigKeySMTPPassword:
|
||||
return maskedConfigValue
|
||||
case ConfigKeyStorageConfig:
|
||||
var cfg objectstore.Config
|
||||
if err := json.Unmarshal([]byte(value), &cfg); err == nil {
|
||||
masked := objectstore.MaskSecrets(cfg)
|
||||
if val, err := json.Marshal(masked); err == nil {
|
||||
return string(val)
|
||||
}
|
||||
}
|
||||
}
|
||||
return value
|
||||
}
|
||||
|
||||
// validateAndMergeStorageConfig parses, merges unmasked secrets, validates parameter values,
|
||||
// and tests connectivity of the new storage configuration.
|
||||
func validateAndMergeStorageConfig(ctx context.Context, value string, currentConfig string) (string, error) {
|
||||
var currentCfg objectstore.Config
|
||||
if err := json.Unmarshal([]byte(currentConfig), ¤tCfg); err != nil {
|
||||
return "", fmt.Errorf("解析当前存储配置失败: %w", err)
|
||||
}
|
||||
|
||||
var newCfg objectstore.Config
|
||||
if err := json.Unmarshal([]byte(value), &newCfg); err != nil {
|
||||
return "", fmt.Errorf("解析目标存储配置失败: %w", err)
|
||||
}
|
||||
|
||||
// 合并被掩码屏蔽的敏感信息,获取完整的真实配置
|
||||
targetCfg := objectstore.MergeMaskedSecrets(newCfg, currentCfg)
|
||||
if err := validateMergedStorageConfig(ctx, currentCfg, newCfg, targetCfg); err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
// 序列化为最终保存的真实明文配置,防止保存屏蔽的 ****** 字符
|
||||
unmaskedVal, err := json.Marshal(targetCfg)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("序列化存储配置失败: %w", err)
|
||||
}
|
||||
|
||||
return string(unmaskedVal), nil
|
||||
}
|
||||
|
||||
func validateMergedStorageConfig(ctx context.Context, currentCfg, newCfg, targetCfg objectstore.Config) error {
|
||||
if newCfg.Driver != "" && newCfg.Driver != currentCfg.Driver {
|
||||
var uploadCount int64
|
||||
if err := db.DB(ctx).Table("w_uploads").
|
||||
Where("status != ?", "deleted").
|
||||
Count(&uploadCount).Error; err != nil {
|
||||
return fmt.Errorf("检查存量文件失败: %w", err)
|
||||
}
|
||||
if uploadCount > 0 {
|
||||
return errors.New(StorageDriverSwitchRequiresMigration)
|
||||
}
|
||||
if err := validateDriverConfig(targetCfg, newCfg.Driver); err != nil {
|
||||
return fmt.Errorf("验证目标存储配置参数失败: %w", err)
|
||||
}
|
||||
pendingCfg := targetCfg
|
||||
pendingCfg.Driver = newCfg.Driver
|
||||
return testStorageBackend(ctx, pendingCfg, newCfg.Driver)
|
||||
}
|
||||
|
||||
if err := objectstore.ValidateConfig(targetCfg); err != nil {
|
||||
return fmt.Errorf("验证存储配置参数失败: %w", err)
|
||||
}
|
||||
return testStorageBackend(ctx, targetCfg, targetCfg.Driver)
|
||||
}
|
||||
|
||||
func validateDriverConfig(cfg objectstore.Config, driver objectstore.Driver) error {
|
||||
cfg.Driver = driver
|
||||
return objectstore.ValidateConfig(cfg)
|
||||
}
|
||||
|
||||
func testStorageBackend(ctx context.Context, cfg objectstore.Config, driver objectstore.Driver) error {
|
||||
testBackend, err := objectstore.NewBackend(ctx, cfg, driver)
|
||||
if err != nil {
|
||||
return fmt.Errorf("初始化测试存储实例失败: %w", err)
|
||||
}
|
||||
if err := testBackend.Test(ctx); err != nil {
|
||||
return fmt.Errorf("存储连通性测试失败: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,646 @@
|
||||
// Copyright 2025 linux.do
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package admin
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"fmt"
|
||||
"log"
|
||||
"math"
|
||||
"net/http"
|
||||
"os"
|
||||
"os/exec"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"gorm.io/gorm"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/backend/pkg/config"
|
||||
"github.com/Rain-kl/Wavelet/backend/pkg/response"
|
||||
db "github.com/Rain-kl/Wavelet/backend/plugins/infra/database"
|
||||
)
|
||||
|
||||
const (
|
||||
binaryKB = 0
|
||||
binaryMB = 1
|
||||
binaryGB = 2
|
||||
valueThreshold = 10
|
||||
maxStringLength = 200
|
||||
)
|
||||
|
||||
// DBOverviewResponse 数据库运行概览响应结构体
|
||||
type DBOverviewResponse struct {
|
||||
Type string `json:"type"`
|
||||
Version string `json:"version"`
|
||||
Name string `json:"name"`
|
||||
Size string `json:"size"`
|
||||
TableCount int64 `json:"table_count"`
|
||||
Connections int64 `json:"connections"`
|
||||
}
|
||||
|
||||
// GetTableDataRequest 分页拉取表数据请求结构体
|
||||
type GetTableDataRequest struct {
|
||||
Table string `form:"table" binding:"required"`
|
||||
Page int `form:"page,default=1"`
|
||||
PageSize int `form:"pageSize,default=10"`
|
||||
}
|
||||
|
||||
// TableDataResponse 动态数据表响应结构体
|
||||
type TableDataResponse struct {
|
||||
Columns []string `json:"columns"`
|
||||
Total int64 `json:"total"`
|
||||
Results []map[string]interface{} `json:"results"`
|
||||
}
|
||||
|
||||
// ExecuteSQLRequest 执行自定义 SQL 请求结构体
|
||||
type ExecuteSQLRequest struct {
|
||||
SQL string `json:"sql" binding:"required"`
|
||||
}
|
||||
|
||||
// ExecuteSQLResponse 执行自定义 SQL 响应结构体
|
||||
type ExecuteSQLResponse struct {
|
||||
Type string `json:"type"` // "select" 或 "exec"
|
||||
Columns []string `json:"columns,omitempty"`
|
||||
Results []map[string]interface{} `json:"results,omitempty"`
|
||||
AffectedRows int64 `json:"affected_rows"`
|
||||
ExecutionTimeMs int64 `json:"execution_time_ms"`
|
||||
}
|
||||
|
||||
// DatabaseInfoResponse 数据库信息响应结构体
|
||||
type DatabaseInfoResponse struct {
|
||||
Type string `json:"type"`
|
||||
Name string `json:"name"`
|
||||
Version string `json:"version"`
|
||||
}
|
||||
|
||||
func formatBytes(bytes uint64) string {
|
||||
const unit = 1024
|
||||
if bytes < unit {
|
||||
return fmt.Sprintf("%d B", bytes)
|
||||
}
|
||||
div, exp := int64(unit), 0
|
||||
for n := bytes / unit; n >= unit; n /= unit {
|
||||
div *= unit
|
||||
exp++
|
||||
}
|
||||
value := float64(bytes) / float64(div)
|
||||
var suffix string
|
||||
switch exp {
|
||||
case binaryKB:
|
||||
suffix = "KiB"
|
||||
case binaryMB:
|
||||
suffix = "MiB"
|
||||
case binaryGB:
|
||||
suffix = "GiB"
|
||||
default:
|
||||
suffix = "TiB"
|
||||
}
|
||||
|
||||
if value == math.Trunc(value) {
|
||||
if value >= valueThreshold {
|
||||
return fmt.Sprintf("%.0f %s", value, suffix)
|
||||
}
|
||||
return fmt.Sprintf("%.1f %s", value, suffix)
|
||||
}
|
||||
return fmt.Sprintf("%.1f %s", value, suffix)
|
||||
}
|
||||
|
||||
const defaultSQLiteDBPath = "./data/wavelet.db"
|
||||
|
||||
func getSQLiteOverview(gormDB *gorm.DB) (DBOverviewResponse, error) {
|
||||
name := config.Config.Database.SQLitePath
|
||||
if name == "" {
|
||||
name = defaultSQLiteDBPath
|
||||
}
|
||||
|
||||
var version string
|
||||
var ver string
|
||||
if err := gormDB.Raw("SELECT sqlite_version()").Scan(&ver).Error; err == nil {
|
||||
version = "SQLite " + ver
|
||||
} else {
|
||||
version = "SQLite"
|
||||
}
|
||||
|
||||
var sizeStr string
|
||||
if fi, err := os.Stat(name); err == nil {
|
||||
size := fi.Size()
|
||||
if size < 0 {
|
||||
size = 0
|
||||
}
|
||||
sizeStr = formatBytes(uint64(size))
|
||||
} else {
|
||||
sizeStr = "0 B"
|
||||
}
|
||||
|
||||
var tableCount int64
|
||||
if err := gormDB.Raw("SELECT count(*) FROM sqlite_master WHERE type='table' AND name NOT LIKE 'sqlite_%'").Scan(&tableCount).Error; err != nil {
|
||||
tableCount = 0
|
||||
}
|
||||
|
||||
var connCount int64
|
||||
if sqlDB, err := gormDB.DB(); err == nil {
|
||||
connCount = int64(sqlDB.Stats().OpenConnections)
|
||||
} else {
|
||||
connCount = 1
|
||||
}
|
||||
|
||||
return DBOverviewResponse{
|
||||
Type: "sqlite",
|
||||
Version: version,
|
||||
Name: name,
|
||||
Size: sizeStr,
|
||||
TableCount: tableCount,
|
||||
Connections: connCount,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func getPostgresOverview(gormDB *gorm.DB) (DBOverviewResponse, error) {
|
||||
name := config.Config.Database.Database
|
||||
|
||||
var version string
|
||||
var ver string
|
||||
if err := gormDB.Raw("SELECT version()").Scan(&ver).Error; err == nil {
|
||||
version = ver
|
||||
} else {
|
||||
version = "PostgreSQL"
|
||||
}
|
||||
|
||||
var sizeStr string
|
||||
var sizeBytes sql.NullInt64
|
||||
if err := gormDB.Raw("SELECT pg_database_size(current_database())").Scan(&sizeBytes).Error; err == nil && sizeBytes.Valid {
|
||||
size := sizeBytes.Int64
|
||||
if size < 0 {
|
||||
size = 0
|
||||
}
|
||||
sizeStr = formatBytes(uint64(size))
|
||||
} else {
|
||||
sizeStr = "0 B"
|
||||
}
|
||||
|
||||
var tableCount int64
|
||||
if err := gormDB.Raw("SELECT count(*) FROM information_schema.tables WHERE table_schema = current_schema()").Scan(&tableCount).Error; err != nil {
|
||||
tableCount = 0
|
||||
}
|
||||
|
||||
var connCount int64
|
||||
var pgc sql.NullInt64
|
||||
if err := gormDB.Raw("SELECT count(*) FROM pg_stat_activity WHERE datname = current_database()").Scan(&pgc).Error; err == nil && pgc.Valid {
|
||||
connCount = pgc.Int64
|
||||
} else {
|
||||
if sqlDB, err := gormDB.DB(); err == nil {
|
||||
connCount = int64(sqlDB.Stats().OpenConnections)
|
||||
} else {
|
||||
connCount = 1
|
||||
}
|
||||
}
|
||||
|
||||
return DBOverviewResponse{
|
||||
Type: "postgres",
|
||||
Version: version,
|
||||
Name: name,
|
||||
Size: sizeStr,
|
||||
TableCount: tableCount,
|
||||
Connections: connCount,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// GetDBOverview 获取数据库运行概览
|
||||
// @Summary 获取数据库运行概览
|
||||
// @Description 获取数据库类型、版本、名称、文件大小、表数量及当前连接数,需要管理员权限
|
||||
// @Tags admin
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Success 200 {object} response.Any{data=DBOverviewResponse} "获取成功"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 403 {object} response.Any "无管理员权限"
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/admin/db-manage/overview [get]
|
||||
func GetDBOverview(c *gin.Context) {
|
||||
gormDB := db.DB(c.Request.Context())
|
||||
if gormDB == nil {
|
||||
response.AbortInternal(c, "数据库未初始化")
|
||||
return
|
||||
}
|
||||
|
||||
var overview DBOverviewResponse
|
||||
var err error
|
||||
|
||||
if !config.Config.Database.Enabled {
|
||||
overview, err = getSQLiteOverview(gormDB)
|
||||
} else {
|
||||
overview, err = getPostgresOverview(gormDB)
|
||||
}
|
||||
|
||||
if err != nil {
|
||||
response.AbortInternal(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, response.OK(overview))
|
||||
}
|
||||
|
||||
// ListDBTables 获取数据库所有表名
|
||||
// @Summary 获取数据库所有表名
|
||||
// @Description 返回当前数据库的所有用户自定义表名称列表,需要管理员权限
|
||||
// @Tags admin
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Success 200 {object} response.Any{data=[]string} "获取成功"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 403 {object} response.Any "无管理员权限"
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/admin/db-manage/tables [get]
|
||||
func ListDBTables(c *gin.Context) {
|
||||
gormDB := db.DB(c.Request.Context())
|
||||
if gormDB == nil {
|
||||
response.AbortInternal(c, "数据库未初始化")
|
||||
return
|
||||
}
|
||||
|
||||
var tables []string
|
||||
var err error
|
||||
|
||||
if !config.Config.Database.Enabled {
|
||||
err = gormDB.Raw("SELECT name FROM sqlite_master WHERE type='table' AND name NOT LIKE 'sqlite_%' ORDER BY name").Scan(&tables).Error
|
||||
} else {
|
||||
err = gormDB.Raw("SELECT table_name FROM information_schema.tables WHERE table_schema = current_schema() ORDER BY table_name").Scan(&tables).Error
|
||||
}
|
||||
|
||||
if err != nil {
|
||||
response.AbortInternal(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, response.OK(tables))
|
||||
}
|
||||
|
||||
// GetDBTableData 获取数据表 data
|
||||
func GetDBTableData(c *gin.Context) {
|
||||
var req GetTableDataRequest
|
||||
if err := c.ShouldBindQuery(&req); err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
gormDB := db.DB(c.Request.Context())
|
||||
if gormDB == nil {
|
||||
response.AbortInternal(c, "数据库未初始化")
|
||||
return
|
||||
}
|
||||
|
||||
// 安全转义表名并拼接
|
||||
quotedTable := `"` + strings.ReplaceAll(req.Table, `"`, `""`) + `"`
|
||||
|
||||
var total int64
|
||||
if err := gormDB.Raw("SELECT count(*) FROM " + quotedTable).Scan(&total).Error; err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
offset := (req.Page - 1) * req.PageSize
|
||||
if offset < 0 {
|
||||
offset = 0
|
||||
}
|
||||
limit := req.PageSize
|
||||
if limit <= 0 {
|
||||
limit = 10
|
||||
}
|
||||
|
||||
rows, err := gormDB.Raw("SELECT * FROM "+quotedTable+" LIMIT ? OFFSET ?", limit, offset).Rows()
|
||||
if err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
defer func() {
|
||||
_ = rows.Close()
|
||||
}()
|
||||
|
||||
cols, err := rows.Columns()
|
||||
if err != nil {
|
||||
response.AbortInternal(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
results, err := scanTableRows(rows, cols)
|
||||
if err != nil {
|
||||
response.AbortInternal(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, response.OK(TableDataResponse{
|
||||
Columns: cols,
|
||||
Total: total,
|
||||
Results: results,
|
||||
}))
|
||||
}
|
||||
|
||||
func scanTableRows(rows *sql.Rows, cols []string) ([]map[string]interface{}, error) {
|
||||
results := make([]map[string]interface{}, 0)
|
||||
for rows.Next() {
|
||||
columns := make([]interface{}, len(cols))
|
||||
columnPointers := make([]interface{}, len(cols))
|
||||
for i := range columns {
|
||||
columnPointers[i] = &columns[i]
|
||||
}
|
||||
|
||||
if err := rows.Scan(columnPointers...); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
rowMap := make(map[string]interface{})
|
||||
for i, colName := range cols {
|
||||
val := columns[i]
|
||||
if b, ok := val.([]byte); ok {
|
||||
strVal := string(b)
|
||||
runes := []rune(strVal)
|
||||
if len(runes) > maxStringLength {
|
||||
strVal = string(runes[:maxStringLength]) + "..."
|
||||
}
|
||||
rowMap[colName] = strVal
|
||||
} else if str, ok := val.(string); ok {
|
||||
runes := []rune(str)
|
||||
if len(runes) > maxStringLength {
|
||||
str = string(runes[:maxStringLength]) + "..."
|
||||
}
|
||||
rowMap[colName] = str
|
||||
} else {
|
||||
rowMap[colName] = val
|
||||
}
|
||||
}
|
||||
results = append(results, rowMap)
|
||||
}
|
||||
return results, nil
|
||||
}
|
||||
|
||||
func executeSQLQuery(gormDB *gorm.DB, sqlStr string, startTime time.Time) (ExecuteSQLResponse, error) {
|
||||
rows, err := gormDB.Raw(sqlStr).Rows()
|
||||
if err != nil {
|
||||
return ExecuteSQLResponse{}, err
|
||||
}
|
||||
defer func() {
|
||||
_ = rows.Close()
|
||||
}()
|
||||
|
||||
cols, err := rows.Columns()
|
||||
if err != nil {
|
||||
return ExecuteSQLResponse{}, err
|
||||
}
|
||||
|
||||
results := make([]map[string]interface{}, 0)
|
||||
for rows.Next() {
|
||||
columns := make([]interface{}, len(cols))
|
||||
columnPointers := make([]interface{}, len(cols))
|
||||
for i := range columns {
|
||||
columnPointers[i] = &columns[i]
|
||||
}
|
||||
|
||||
if err := rows.Scan(columnPointers...); err != nil {
|
||||
return ExecuteSQLResponse{}, err
|
||||
}
|
||||
|
||||
rowMap := make(map[string]interface{})
|
||||
for i, colName := range cols {
|
||||
val := columns[i]
|
||||
if b, ok := val.([]byte); ok {
|
||||
rowMap[colName] = string(b)
|
||||
} else {
|
||||
rowMap[colName] = val
|
||||
}
|
||||
}
|
||||
results = append(results, rowMap)
|
||||
}
|
||||
|
||||
executionTime := time.Since(startTime).Milliseconds()
|
||||
return ExecuteSQLResponse{
|
||||
Type: "select",
|
||||
Columns: cols,
|
||||
Results: results,
|
||||
AffectedRows: int64(len(results)),
|
||||
ExecutionTimeMs: executionTime,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func executeSQLMutation(gormDB *gorm.DB, sqlStr string, startTime time.Time) (ExecuteSQLResponse, error) {
|
||||
tx := gormDB.Exec(sqlStr)
|
||||
if tx.Error != nil {
|
||||
return ExecuteSQLResponse{}, tx.Error
|
||||
}
|
||||
|
||||
executionTime := time.Since(startTime).Milliseconds()
|
||||
return ExecuteSQLResponse{
|
||||
Type: "exec",
|
||||
AffectedRows: tx.RowsAffected,
|
||||
ExecutionTimeMs: executionTime,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// ExecuteSQL 执行 SQL 查询
|
||||
// @Summary 执行 SQL 查询
|
||||
// @Description 在当前数据库中执行任意自定义 SQL,如果是查询语句将返回格式化后的列与数据集,否则返回受影响行数,需要管理员权限
|
||||
// @Tags admin
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param request body ExecuteSQLRequest true "SQL 请求参数"
|
||||
// @Success 200 {object} response.Any{data=ExecuteSQLResponse} "执行完毕"
|
||||
// @Failure 400 {object} response.Any "SQL 语句错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 403 {object} response.Any "无管理员权限"
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/admin/db-manage/query [post]
|
||||
func ExecuteSQL(c *gin.Context) {
|
||||
var req ExecuteSQLRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
gormDB := db.DB(c.Request.Context())
|
||||
if gormDB == nil {
|
||||
response.AbortInternal(c, "数据库未初始化")
|
||||
return
|
||||
}
|
||||
|
||||
trimmedSQL := strings.TrimSpace(req.SQL)
|
||||
if trimmedSQL == "" {
|
||||
response.AbortBadRequest(c, "SQL 语句不能为空")
|
||||
return
|
||||
}
|
||||
|
||||
startTime := time.Now()
|
||||
|
||||
isQuery := false
|
||||
lowerSQL := strings.ToLower(trimmedSQL)
|
||||
queryKeywords := []string{"select", "show", "explain", "describe", "pragma"}
|
||||
for _, kw := range queryKeywords {
|
||||
if strings.HasPrefix(lowerSQL, kw) {
|
||||
isQuery = true
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
var resp ExecuteSQLResponse
|
||||
var err error
|
||||
|
||||
if isQuery {
|
||||
resp, err = executeSQLQuery(gormDB, trimmedSQL, startTime)
|
||||
} else {
|
||||
resp, err = executeSQLMutation(gormDB, trimmedSQL, startTime)
|
||||
}
|
||||
|
||||
if err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, response.OK(resp))
|
||||
}
|
||||
|
||||
func getSQLiteInfo(ctx context.Context) DatabaseInfoResponse {
|
||||
info := DatabaseInfoResponse{
|
||||
Type: "sqlite",
|
||||
Name: config.Config.Database.SQLitePath,
|
||||
Version: "SQLite",
|
||||
}
|
||||
if info.Name == "" {
|
||||
info.Name = "./data/wavelet.db"
|
||||
}
|
||||
gormDB := db.DB(ctx)
|
||||
if gormDB == nil {
|
||||
return info
|
||||
}
|
||||
var ver string
|
||||
if err := gormDB.Raw("SELECT sqlite_version()").Scan(&ver).Error; err == nil && ver != "" {
|
||||
info.Version = "SQLite " + ver
|
||||
}
|
||||
return info
|
||||
}
|
||||
|
||||
func getPostgresInfo(ctx context.Context) DatabaseInfoResponse {
|
||||
info := DatabaseInfoResponse{
|
||||
Type: "postgres",
|
||||
Name: config.Config.Database.Database,
|
||||
Version: "PostgreSQL",
|
||||
}
|
||||
gormDB := db.DB(ctx)
|
||||
if gormDB == nil {
|
||||
return info
|
||||
}
|
||||
var ver string
|
||||
if err := gormDB.Raw("SELECT version()").Scan(&ver).Error; err == nil && ver != "" {
|
||||
info.Version = ver
|
||||
}
|
||||
return info
|
||||
}
|
||||
|
||||
// GetDatabaseInfo 获取当前数据库类型及版本信息
|
||||
// @Summary 获取数据库信息
|
||||
// @Description 返回当前使用的数据库类型(sqlite/postgres)、名称/路径及版本字符串,需要管理员权限
|
||||
// @Tags admin
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Success 200 {object} response.Any{data=DatabaseInfoResponse} "获取成功"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 403 {object} response.Any "无管理员权限"
|
||||
// @Router /api/v1/admin/db-info [get]
|
||||
func GetDatabaseInfo(c *gin.Context) {
|
||||
var info DatabaseInfoResponse
|
||||
if !config.Config.Database.Enabled {
|
||||
info = getSQLiteInfo(c.Request.Context())
|
||||
} else {
|
||||
info = getPostgresInfo(c.Request.Context())
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(info))
|
||||
}
|
||||
|
||||
// ExportDatabase 导出数据库
|
||||
// @Summary 导出数据库
|
||||
// @Description SQLite 时直接下载 .db 文件;PostgreSQL 时执行 pg_dump 并流式下载 .sql 文件,需要管理员权限
|
||||
// @Tags admin
|
||||
// @Produce application/octet-stream
|
||||
// @Security SessionCookie
|
||||
// @Success 200 {file} binary "数据库文件"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 403 {object} response.Any "无管理员权限"
|
||||
// @Failure 500 {object} response.Any "导出失败"
|
||||
// @Router /api/v1/admin/db-export [get]
|
||||
func ExportDatabase(c *gin.Context) {
|
||||
if !config.Config.Database.Enabled {
|
||||
exportSQLite(c)
|
||||
} else {
|
||||
exportPostgres(c)
|
||||
}
|
||||
}
|
||||
|
||||
func exportSQLite(c *gin.Context) {
|
||||
path := config.Config.Database.SQLitePath
|
||||
if path == "" {
|
||||
path = defaultSQLiteDBPath
|
||||
}
|
||||
|
||||
//nolint:gosec // export db file path is trusted
|
||||
f, err := os.Open(path)
|
||||
if err != nil {
|
||||
response.AbortInternal(c, "无法打开数据库文件: "+err.Error())
|
||||
return
|
||||
}
|
||||
defer func() {
|
||||
if closeErr := f.Close(); closeErr != nil {
|
||||
_ = closeErr
|
||||
}
|
||||
}()
|
||||
|
||||
fi, err := f.Stat()
|
||||
if err != nil {
|
||||
response.AbortInternal(c, "无法读取数据库文件信息: "+err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
c.Header("Content-Disposition", `attachment; filename="wavelet.db"`)
|
||||
c.Header("Content-Type", "application/octet-stream")
|
||||
c.Header("Content-Length", fmt.Sprintf("%d", fi.Size()))
|
||||
c.Status(http.StatusOK)
|
||||
http.ServeContent(c.Writer, c.Request, "wavelet.db", fi.ModTime(), f)
|
||||
}
|
||||
|
||||
func exportPostgres(c *gin.Context) {
|
||||
dbCfg := config.Config.Database
|
||||
|
||||
pgDumpPath, err := exec.LookPath("pg_dump")
|
||||
if err != nil {
|
||||
response.AbortInternal(c, "pg_dump 不可用,请确保服务器已安装 PostgreSQL 客户端工具")
|
||||
return
|
||||
}
|
||||
|
||||
args := []string{
|
||||
"--no-password",
|
||||
"-h", dbCfg.Host,
|
||||
"-p", fmt.Sprintf("%d", dbCfg.Port),
|
||||
"-U", dbCfg.Username,
|
||||
dbCfg.Database,
|
||||
}
|
||||
|
||||
//nolint:gosec // pg_dump args are constructed from validated db config
|
||||
cmd := exec.CommandContext(c.Request.Context(), pgDumpPath, args...)
|
||||
if dbCfg.Password != "" {
|
||||
cmd.Env = append(os.Environ(), "PGPASSWORD="+dbCfg.Password)
|
||||
} else {
|
||||
cmd.Env = os.Environ()
|
||||
}
|
||||
|
||||
fileName := fmt.Sprintf("wavelet_%s.sql", time.Now().Format("20060102_150405"))
|
||||
c.Header("Content-Disposition", `attachment; filename="`+fileName+`"`)
|
||||
c.Header("Content-Type", "application/octet-stream")
|
||||
c.Status(http.StatusOK)
|
||||
|
||||
cmd.Stdout = c.Writer
|
||||
cmd.Stderr = nil
|
||||
|
||||
if err := cmd.Run(); err != nil {
|
||||
log.Printf("[db-export] pg_dump failed: %v\n", err)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,685 @@
|
||||
// Copyright 2025 linux.do
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package admin
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/backend/pkg/config"
|
||||
"github.com/Rain-kl/Wavelet/backend/pkg/logger"
|
||||
"github.com/Rain-kl/Wavelet/backend/pkg/response"
|
||||
"github.com/Rain-kl/Wavelet/backend/pkg/util"
|
||||
"github.com/Rain-kl/Wavelet/backend/plugins/domain/risk_control"
|
||||
"github.com/Rain-kl/Wavelet/backend/plugins/domain/risk_control/logstore"
|
||||
"github.com/Rain-kl/Wavelet/backend/plugins/drivers/driver_asynq_worker"
|
||||
db "github.com/Rain-kl/Wavelet/backend/plugins/infra/database"
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/gorilla/websocket"
|
||||
)
|
||||
|
||||
const (
|
||||
defaultLimit = 200
|
||||
maxLimit = 500
|
||||
maxPageSize = 100
|
||||
hoursInDay = 24
|
||||
analyticsDays = 7
|
||||
topActiveLimit = 10
|
||||
)
|
||||
|
||||
// logsResponse 历史日志查询响应
|
||||
type logsResponse struct {
|
||||
Lines []logger.LogEntry `json:"lines"`
|
||||
HasMore bool `json:"has_more"`
|
||||
NextCursor int `json:"next_cursor"` // 用于加载更早日志的 cursor
|
||||
}
|
||||
|
||||
// GetLogs 获取历史日志
|
||||
// @Summary 获取系统日志
|
||||
// @Description 分页获取系统历史日志,cursor=0 获取最新日志,cursor>0 获取更早日志
|
||||
// @Tags admin
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param cursor query int false "日志游标,0=获取最新" default(0)
|
||||
// @Param limit query int false "每页条数" default(200)
|
||||
// @Success 200 {object} response.Any{data=logsResponse} "日志列表"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 403 {object} response.Any "无管理员权限"
|
||||
// @Router /api/v1/admin/logs [get]
|
||||
func GetLogs(c *gin.Context) {
|
||||
cursorStr := c.DefaultQuery("cursor", "0")
|
||||
limitStr := c.DefaultQuery("limit", "200")
|
||||
|
||||
var cursor, limit int
|
||||
if _, err := parsePositiveInt(cursorStr, &cursor); err != nil {
|
||||
response.AbortWithError(c, http.StatusBadRequest, InvalidCursorParam)
|
||||
return
|
||||
}
|
||||
if _, err := parsePositiveInt(limitStr, &limit); err != nil || limit <= 0 {
|
||||
limit = defaultLimit
|
||||
}
|
||||
if limit > maxLimit {
|
||||
limit = maxLimit
|
||||
}
|
||||
|
||||
entries, hasMore := logger.GlobalRingBuffer.Query(cursor, limit)
|
||||
|
||||
resp := logsResponse{
|
||||
Lines: entries,
|
||||
HasMore: hasMore,
|
||||
}
|
||||
if len(entries) > 0 {
|
||||
resp.NextCursor = entries[0].Index
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, response.OK(resp))
|
||||
}
|
||||
|
||||
// wsMessage WebSocket 消息格式
|
||||
type wsMessage struct {
|
||||
Type string `json:"type"` // "log" | "error"
|
||||
Data json.RawMessage `json:"data"`
|
||||
}
|
||||
|
||||
// HandleLogWebSocket WebSocket 端点,实时推送系统日志
|
||||
// @Summary 系统日志实时推送
|
||||
// @Description 通过 WebSocket 实时推送系统日志,需要管理员权限
|
||||
// @Tags admin
|
||||
// @Router /api/v1/admin/logs/ws [get]
|
||||
func HandleLogWebSocket(c *gin.Context) {
|
||||
upgrader := getUpgrader()
|
||||
|
||||
conn, err := upgrader.Upgrade(c.Writer, c.Request, nil)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
defer func() { _ = conn.Close() }()
|
||||
|
||||
ch := logger.GlobalRingBuffer.Subscribe()
|
||||
defer logger.GlobalRingBuffer.Unsubscribe(ch)
|
||||
|
||||
done := make(chan struct{})
|
||||
util.Go(func() {
|
||||
defer close(done)
|
||||
for {
|
||||
_, _, err := conn.ReadMessage()
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
}
|
||||
})
|
||||
|
||||
for {
|
||||
select {
|
||||
case <-done:
|
||||
return
|
||||
case entry, ok := <-ch:
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
data, _ := json.Marshal(entry)
|
||||
msg := wsMessage{Type: "log", Data: data}
|
||||
payload, _ := json.Marshal(msg)
|
||||
if err := conn.WriteMessage(1, payload); err != nil {
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// accessLogItem 访问日志单条数据
|
||||
type accessLogItem struct {
|
||||
ID uint64 `json:"id,string"`
|
||||
UserID uint64 `json:"user_id,string"`
|
||||
Username string `json:"username"`
|
||||
Nickname string `json:"nickname"`
|
||||
Path string `json:"path"`
|
||||
Method string `json:"method"`
|
||||
IP string `json:"ip"`
|
||||
UserAgent string `json:"user_agent"`
|
||||
Headers string `json:"headers"`
|
||||
Status int32 `json:"status"`
|
||||
Latency int64 `json:"latency"`
|
||||
CreatedAt string `json:"created_at"`
|
||||
}
|
||||
|
||||
// accessLogsResponse 访问日志查询响应
|
||||
type accessLogsResponse struct {
|
||||
Total uint64 `json:"total"`
|
||||
List []accessLogItem `json:"list"`
|
||||
}
|
||||
|
||||
func buildAccessLogFilter(ctx context.Context, c *gin.Context) (logstore.AccessLogFilter, error) {
|
||||
filter := logstore.AccessLogFilter{}
|
||||
|
||||
username := c.Query("username")
|
||||
if username != "" {
|
||||
var userIDs []uint64
|
||||
if err := db.DB(ctx).Table("w_users").
|
||||
Where("username LIKE ? ESCAPE '\\'", "%"+util.EscapeLike(username)+"%").
|
||||
Pluck("id", &userIDs).Error; err != nil {
|
||||
return filter, fmt.Errorf("查询用户信息失败: %w", err)
|
||||
}
|
||||
filter.UserIDs = userIDs
|
||||
}
|
||||
|
||||
if path := c.Query("path"); path != "" {
|
||||
filter.Path = path
|
||||
}
|
||||
|
||||
if startTime := c.Query("start_time"); startTime != "" {
|
||||
if t, err := parseAccessLogTime(startTime); err == nil {
|
||||
filter.StartTime = &t
|
||||
}
|
||||
}
|
||||
|
||||
if endTime := c.Query("end_time"); endTime != "" {
|
||||
if t, err := parseAccessLogTime(endTime); err == nil {
|
||||
filter.EndTime = &t
|
||||
}
|
||||
}
|
||||
|
||||
return filter, nil
|
||||
}
|
||||
|
||||
func parseAccessLogTime(value string) (time.Time, error) {
|
||||
if t, err := time.Parse(time.RFC3339, value); err == nil {
|
||||
return t, nil
|
||||
}
|
||||
return time.Parse("2006-01-02 15:04:05", value)
|
||||
}
|
||||
|
||||
func enrichAccessLogsWithUsers(ctx context.Context, list []accessLogItem) {
|
||||
if len(list) == 0 {
|
||||
return
|
||||
}
|
||||
|
||||
userIDs := make([]uint64, 0, len(list))
|
||||
seen := make(map[uint64]struct{}, len(list))
|
||||
for _, item := range list {
|
||||
if _, ok := seen[item.UserID]; ok {
|
||||
continue
|
||||
}
|
||||
seen[item.UserID] = struct{}{}
|
||||
userIDs = append(userIDs, item.UserID)
|
||||
}
|
||||
|
||||
userMap := make(map[uint64]struct{ Username, Nickname string })
|
||||
var users []struct {
|
||||
ID uint64
|
||||
Username string
|
||||
Nickname string
|
||||
}
|
||||
if err := db.DB(ctx).Table("w_users").Where("id IN ?", userIDs).Find(&users).Error; err == nil {
|
||||
for _, u := range users {
|
||||
userMap[u.ID] = struct{ Username, Nickname string }{Username: u.Username, Nickname: u.Nickname}
|
||||
}
|
||||
}
|
||||
for i := range list {
|
||||
if info, ok := userMap[list[i].UserID]; ok {
|
||||
list[i].Username = info.Username
|
||||
list[i].Nickname = info.Nickname
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// GetAccessLogs 获取 ClickHouse 异步采集的访问日志
|
||||
// @Summary 获取用户访问日志
|
||||
// @Description 分页并按照用户、接口路径、时间范围等维度检索用户访问日志列表(需要管理员权限)
|
||||
// @Tags admin
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param page query int false "页码" default(1)
|
||||
// @Param page_size query int false "每页条数" default(20)
|
||||
// @Param username query string false "用户名模糊搜索"
|
||||
// @Param path query string false "接口路径模糊搜索"
|
||||
// @Param start_time query string false "起始时间(RFC3339 或 YYYY-MM-DD HH:MM:SS)"
|
||||
// @Param end_time query string false "结束时间(RFC3339 或 YYYY-MM-DD HH:MM:SS)"
|
||||
// @Success 200 {object} response.Any{data=accessLogsResponse} "访问日志列表"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 403 {object} response.Any "无管理员权限"
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/admin/logs/access [get]
|
||||
func GetAccessLogs(c *gin.Context) {
|
||||
ctx := c.Request.Context()
|
||||
store, err := logstore.Active(ctx)
|
||||
if err != nil {
|
||||
response.AbortInternal(c, "日志存储初始化失败")
|
||||
return
|
||||
}
|
||||
|
||||
page, _ := strconv.Atoi(c.DefaultQuery("page", "1"))
|
||||
if page < 1 {
|
||||
page = 1
|
||||
}
|
||||
pageSize, _ := strconv.Atoi(c.DefaultQuery("page_size", "20"))
|
||||
if pageSize < 1 {
|
||||
pageSize = 20
|
||||
}
|
||||
if pageSize > maxPageSize {
|
||||
pageSize = maxPageSize
|
||||
}
|
||||
|
||||
filter, err := buildAccessLogFilter(ctx, c)
|
||||
if err != nil {
|
||||
response.AbortWithError(c, http.StatusInternalServerError, err.Error())
|
||||
return
|
||||
}
|
||||
if filter.UserIDs != nil && len(filter.UserIDs) == 0 {
|
||||
c.JSON(http.StatusOK, response.OK(accessLogsResponse{Total: 0, List: []accessLogItem{}}))
|
||||
return
|
||||
}
|
||||
|
||||
logs, total, err := store.UserAccessLogs.List(ctx, filter, page, pageSize)
|
||||
if err != nil {
|
||||
response.AbortWithError(c, http.StatusInternalServerError, err.Error())
|
||||
return
|
||||
}
|
||||
if total == 0 {
|
||||
c.JSON(http.StatusOK, response.OK(accessLogsResponse{Total: 0, List: []accessLogItem{}}))
|
||||
return
|
||||
}
|
||||
|
||||
list := make([]accessLogItem, len(logs))
|
||||
for i, logItem := range logs {
|
||||
list[i] = accessLogItem{
|
||||
ID: logItem.ID,
|
||||
UserID: logItem.UserID,
|
||||
Path: logItem.Path,
|
||||
Method: logItem.Method,
|
||||
IP: logItem.IP,
|
||||
UserAgent: logItem.UserAgent,
|
||||
Headers: logItem.Headers,
|
||||
Status: logItem.Status,
|
||||
Latency: logItem.Latency,
|
||||
CreatedAt: logItem.CreatedAt.Format(time.RFC3339),
|
||||
}
|
||||
}
|
||||
enrichAccessLogsWithUsers(ctx, list)
|
||||
|
||||
c.JSON(http.StatusOK, response.OK(accessLogsResponse{
|
||||
Total: total,
|
||||
List: list,
|
||||
}))
|
||||
}
|
||||
|
||||
// trendItem 趋势图数据点
|
||||
type trendItem struct {
|
||||
Date string `json:"date"`
|
||||
Count uint64 `json:"count"`
|
||||
}
|
||||
|
||||
// browserItem 浏览器占比排行
|
||||
type browserItem struct {
|
||||
Browser string `json:"browser"`
|
||||
Count uint64 `json:"count"`
|
||||
}
|
||||
|
||||
// topUserItem 活跃用户数据
|
||||
type topUserItem struct {
|
||||
UserID uint64 `json:"user_id,string"`
|
||||
Username string `json:"username"`
|
||||
Nickname string `json:"nickname"`
|
||||
Count uint64 `json:"count"`
|
||||
}
|
||||
|
||||
// logsAnalyticsResponse 访问日志数据分析结果
|
||||
type logsAnalyticsResponse struct {
|
||||
Trend []trendItem `json:"trend"`
|
||||
Browsers []browserItem `json:"browsers"`
|
||||
TopUsers []topUserItem `json:"top_users"`
|
||||
}
|
||||
|
||||
// GetLogsAnalytics 获取 ClickHouse 访问日志图表聚合指标
|
||||
// @Summary 获取访问日志分析数据
|
||||
// @Description 聚合统计最近 7 天的每日访问趋势、浏览器分布以及前 10 名最活跃用户排行(需要管理员权限)
|
||||
// @Tags admin
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Success 200 {object} response.Any{data=logsAnalyticsResponse} "分析统计数据"
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 403 {object} response.Any "无管理员权限"
|
||||
// @Router /api/v1/admin/logs/analytics [get]
|
||||
func GetLogsAnalytics(c *gin.Context) {
|
||||
ctx := c.Request.Context()
|
||||
store, err := logstore.Active(ctx)
|
||||
if err != nil {
|
||||
response.AbortInternal(c, "日志存储初始化失败")
|
||||
return
|
||||
}
|
||||
|
||||
startTime := time.Now().AddDate(0, 0, -(analyticsDays - 1)).Truncate(hoursInDay * time.Hour)
|
||||
|
||||
trendPoints, err := store.UserAccessLogs.GetDailyTrend(ctx, analyticsDays)
|
||||
if err != nil {
|
||||
response.AbortWithError(c, http.StatusInternalServerError, "查询访问趋势失败: "+err.Error())
|
||||
return
|
||||
}
|
||||
trendList := make([]trendItem, len(trendPoints))
|
||||
for i, point := range trendPoints {
|
||||
trendList[i] = trendItem{
|
||||
Date: point.Date,
|
||||
Count: point.Count,
|
||||
}
|
||||
}
|
||||
|
||||
browserPoints, err := store.UserAccessLogs.GetBrowserDistribution(ctx, startTime)
|
||||
if err != nil {
|
||||
response.AbortWithError(c, http.StatusInternalServerError, "查询浏览器分布失败: "+err.Error())
|
||||
return
|
||||
}
|
||||
browserList := make([]browserItem, len(browserPoints))
|
||||
for i, point := range browserPoints {
|
||||
browserList[i] = browserItem{
|
||||
Browser: point.Browser,
|
||||
Count: point.Count,
|
||||
}
|
||||
}
|
||||
|
||||
topUserPoints, err := store.UserAccessLogs.GetTopActiveUsers(ctx, startTime, topActiveLimit)
|
||||
if err != nil {
|
||||
response.AbortWithError(c, http.StatusInternalServerError, "查询活跃用户失败: "+err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
topUsers := make([]topUserItem, len(topUserPoints))
|
||||
userIDs := make([]uint64, len(topUserPoints))
|
||||
for i, point := range topUserPoints {
|
||||
topUsers[i] = topUserItem{
|
||||
UserID: point.UserID,
|
||||
Count: point.Count,
|
||||
}
|
||||
userIDs[i] = point.UserID
|
||||
}
|
||||
|
||||
if len(userIDs) > 0 {
|
||||
userProfileMap := make(map[uint64]struct {
|
||||
Username string
|
||||
Nickname string
|
||||
})
|
||||
var users []struct {
|
||||
ID uint64
|
||||
Username string
|
||||
Nickname string
|
||||
}
|
||||
if errProfile := db.DB(ctx).Table("w_users").Where("id IN ?", userIDs).Find(&users).Error; errProfile == nil {
|
||||
for _, u := range users {
|
||||
userProfileMap[u.ID] = struct {
|
||||
Username string
|
||||
Nickname string
|
||||
}{
|
||||
Username: u.Username,
|
||||
Nickname: u.Nickname,
|
||||
}
|
||||
}
|
||||
}
|
||||
for i := range topUsers {
|
||||
if profile, ok := userProfileMap[topUsers[i].UserID]; ok {
|
||||
topUsers[i].Username = profile.Username
|
||||
topUsers[i].Nickname = profile.Nickname
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, response.OK(logsAnalyticsResponse{
|
||||
Trend: trendList,
|
||||
Browsers: browserList,
|
||||
TopUsers: topUsers,
|
||||
}))
|
||||
}
|
||||
|
||||
func getUpgrader() *websocket.Upgrader {
|
||||
return &websocket.Upgrader{
|
||||
CheckOrigin: func(r *http.Request) bool {
|
||||
origin := r.Header.Get("Origin")
|
||||
if origin == "" {
|
||||
return true
|
||||
}
|
||||
|
||||
// 1. 同源检查 (Same-origin check)
|
||||
u, err := url.Parse(origin)
|
||||
if err == nil && strings.EqualFold(u.Host, r.Host) {
|
||||
return true
|
||||
}
|
||||
|
||||
// 2. 检查配置的允许跨域 Origin (Check allowed origins in system config)
|
||||
ctx := r.Context()
|
||||
if sc, err := GetSystemConfigByKey(ctx, ConfigKeyServerAddress); err == nil && sc.Value != "" {
|
||||
originToCheck := strings.TrimRight(strings.TrimSpace(origin), "/")
|
||||
allowedOrigins := strings.Split(sc.Value, ",")
|
||||
for _, allowed := range allowedOrigins {
|
||||
allowed = strings.TrimRight(strings.TrimSpace(allowed), "/")
|
||||
if allowed != "" && strings.EqualFold(allowed, originToCheck) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
}
|
||||
return false
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func parsePositiveInt(s string, result *int) (bool, error) {
|
||||
if s == "" {
|
||||
*result = 0
|
||||
return true, nil
|
||||
}
|
||||
n, err := strconv.Atoi(s)
|
||||
if err != nil || n < 0 {
|
||||
return false, err
|
||||
}
|
||||
*result = n
|
||||
return true, nil
|
||||
}
|
||||
|
||||
// Log DB Switch Task
|
||||
const (
|
||||
// LogDBSwitchTask 切换日志数据库任务标识。
|
||||
LogDBSwitchTask = "logs:db_switch"
|
||||
// TaskTypeLogDBSwitch 管理端任务类型。
|
||||
TaskTypeLogDBSwitch = "logs_db_switch"
|
||||
|
||||
copyBatchSize = 1000
|
||||
targetPostgres = "postgres"
|
||||
targetSQLite = "sqlite"
|
||||
targetClickHouse = "clickhouse"
|
||||
)
|
||||
|
||||
// LogDBSwitchMeta 描述切换日志数据库任务。
|
||||
var LogDBSwitchMeta = driver_asynq_worker.TaskMeta{
|
||||
Type: TaskTypeLogDBSwitch,
|
||||
AsynqTask: LogDBSwitchTask,
|
||||
Name: "切换日志数据库",
|
||||
Description: "复制迁移用户访问日志并在成功后切换日志主库(期间禁止日志写入)",
|
||||
SupportsTime: false,
|
||||
MaxRetry: driver_asynq_worker.DefaultMaxRetry,
|
||||
Queue: driver_asynq_worker.QueueDefault,
|
||||
Retryable: true,
|
||||
Params: []driver_asynq_worker.TaskParam{
|
||||
{Name: "target", Label: "目标日志库", Type: "string", Required: true,
|
||||
Placeholder: "postgres|sqlite|clickhouse", Description: "迁移目标:postgres(主库为 PG 时)、sqlite(主库为 SQLite 时)或 clickhouse"},
|
||||
},
|
||||
}
|
||||
|
||||
type logDBSwitchPayload struct {
|
||||
Target string `json:"target"`
|
||||
}
|
||||
|
||||
// LogDBSwitchHandler 切换日志数据库任务处理器。
|
||||
type LogDBSwitchHandler struct{}
|
||||
|
||||
// ValidatePayload 校验并规范化参数。
|
||||
func (h *LogDBSwitchHandler) ValidatePayload(payload []byte) ([]byte, error) {
|
||||
var p logDBSwitchPayload
|
||||
if err := json.Unmarshal(payload, &p); err != nil {
|
||||
return nil, fmt.Errorf("参数解析失败: %w", err)
|
||||
}
|
||||
p.Target = normalizeTarget(p.Target)
|
||||
if !validTarget(p.Target) {
|
||||
return nil, fmt.Errorf("目标日志库不合法: %s", p.Target)
|
||||
}
|
||||
out, err := json.Marshal(p)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
func normalizeTarget(v string) string {
|
||||
switch v {
|
||||
case targetPostgres, "postgresql":
|
||||
return targetPostgres
|
||||
case targetSQLite, "sqlite3":
|
||||
return targetSQLite
|
||||
case targetClickHouse, "ch":
|
||||
return targetClickHouse
|
||||
}
|
||||
return v
|
||||
}
|
||||
|
||||
func validTarget(v string) bool {
|
||||
return v == targetPostgres || v == targetSQLite || v == targetClickHouse
|
||||
}
|
||||
|
||||
// Execute 执行迁移。
|
||||
func (h *LogDBSwitchHandler) Execute(ctx context.Context, payload []byte) (*driver_asynq_worker.TaskResult, error) {
|
||||
var p logDBSwitchPayload
|
||||
if err := json.Unmarshal(payload, &p); err != nil {
|
||||
return nil, fmt.Errorf("参数解析失败: %w", err)
|
||||
}
|
||||
p.Target = normalizeTarget(p.Target)
|
||||
if err := validateSwitch(ctx, p.Target); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
source, err := currentLogDatabase(ctx)
|
||||
if err != nil {
|
||||
driver_asynq_worker.AppendLog(ctx, "读取日志主库失败: %v", err)
|
||||
return nil, err
|
||||
}
|
||||
driver_asynq_worker.AppendLog(ctx, "开始切换日志数据库:%s -> %s", source, p.Target)
|
||||
|
||||
if err := setMigrationFlag(ctx, "migrating"); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer func() {
|
||||
if err := setMigrationFlag(ctx, ""); err != nil {
|
||||
logger.ErrorF(ctx, "清除日志迁移冻结标记失败: %v", err)
|
||||
}
|
||||
}()
|
||||
|
||||
if err := risk_control.Drain(ctx); err != nil {
|
||||
return nil, fmt.Errorf("排空日志写入队列失败: %w", err)
|
||||
}
|
||||
|
||||
src, err := logstore.Active(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
dst, err := logstore.BuildForMigration(ctx, p.Target)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if _, err := dst.UserAccessLogs.DeleteAll(ctx); err != nil {
|
||||
return nil, fmt.Errorf("清空目标用户访问日志失败: %w", err)
|
||||
}
|
||||
from, to, err := src.UserAccessLogs.MigrationRange(ctx)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("读取源库时间范围失败: %w", err)
|
||||
}
|
||||
if !from.IsZero() && !to.IsZero() {
|
||||
if err := dst.UserAccessLogs.EnsurePartitions(ctx, from, to); err != nil {
|
||||
return nil, fmt.Errorf("预建目标分区失败: %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
if err := copyUserAccessLogs(ctx, src, dst); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := flipLogDatabase(ctx, p.Target); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
logstore.InvalidateCache()
|
||||
driver_asynq_worker.AppendLog(ctx, "日志数据库已切换为 %s,写入恢复", p.Target)
|
||||
return &driver_asynq_worker.TaskResult{Message: fmt.Sprintf("日志数据库已从 %s 切换为 %s", source, p.Target)}, nil
|
||||
}
|
||||
|
||||
func validateSwitch(ctx context.Context, target string) error {
|
||||
source, err := currentLogDatabase(ctx)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if source == target {
|
||||
return errors.New("目标日志库与当前日志库相同,无需迁移")
|
||||
}
|
||||
switch target {
|
||||
case targetClickHouse:
|
||||
if !config.Config.ClickHouse.Enabled {
|
||||
return errors.New("ClickHouse 未启用,无法迁移到 ClickHouse")
|
||||
}
|
||||
case targetPostgres:
|
||||
if !config.Config.Database.Enabled {
|
||||
return errors.New("PostgreSQL 未启用(当前主库为 SQLite),无法迁移到 PostgreSQL")
|
||||
}
|
||||
case targetSQLite:
|
||||
if config.Config.Database.Enabled {
|
||||
return errors.New("当前主库为 PostgreSQL,日志库不能设置为 SQLite")
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func currentLogDatabase(ctx context.Context) (string, error) {
|
||||
cfg, err := GetSystemConfigByKey(ctx, ConfigKeyLogDatabase)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("读取日志主库失败: %w", err)
|
||||
}
|
||||
if cfg.Value == "" {
|
||||
return "", errors.New("日志主库配置为空")
|
||||
}
|
||||
return cfg.Value, nil
|
||||
}
|
||||
|
||||
func setMigrationFlag(ctx context.Context, v string) error {
|
||||
return SaveOrUpdateSystemConfig(ctx, ConfigKeyLogDBMigration, v)
|
||||
}
|
||||
|
||||
func flipLogDatabase(ctx context.Context, target string) error {
|
||||
return SaveOrUpdateSystemConfig(ctx, ConfigKeyLogDatabase, target)
|
||||
}
|
||||
|
||||
func copyUserAccessLogs(ctx context.Context, src, dst *logstore.Store) error {
|
||||
var afterID uint64
|
||||
var copied int
|
||||
for {
|
||||
rows, err := src.UserAccessLogs.ListForMigration(ctx, afterID, copyBatchSize)
|
||||
if err != nil {
|
||||
return fmt.Errorf("读取源用户访问日志失败: %w", err)
|
||||
}
|
||||
if len(rows) == 0 {
|
||||
break
|
||||
}
|
||||
if err := dst.UserAccessLogs.BatchInsert(ctx, rows); err != nil {
|
||||
return fmt.Errorf("写入目标用户访问日志失败: %w", err)
|
||||
}
|
||||
afterID = rows[len(rows)-1].ID
|
||||
copied += len(rows)
|
||||
driver_asynq_worker.AppendLog(ctx, "已复制用户访问日志 %d 条", copied)
|
||||
if len(rows) < copyBatchSize {
|
||||
break
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,234 @@
|
||||
// Copyright 2025 linux.do
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package admin
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"math"
|
||||
"net/http"
|
||||
"runtime"
|
||||
"time"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"gorm.io/gorm"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/backend/pkg/config"
|
||||
"github.com/Rain-kl/Wavelet/backend/pkg/logger"
|
||||
"github.com/Rain-kl/Wavelet/backend/pkg/response"
|
||||
"github.com/Rain-kl/Wavelet/backend/plugins/domain/risk_control/logstore"
|
||||
)
|
||||
|
||||
var startTime = time.Now()
|
||||
|
||||
const (
|
||||
minutesInHour = 60
|
||||
secondsInMinute = 60
|
||||
nanosPerSecond = 1e9
|
||||
|
||||
logDBNamePostgres = "postgres"
|
||||
logDBNameSQLite = "sqlite"
|
||||
logDBNameClickHouse = "clickhouse"
|
||||
defaultLogRetentionDays = 30
|
||||
)
|
||||
|
||||
// SystemStatusResponse 系统状态响应结构体
|
||||
type SystemStatusResponse struct {
|
||||
Uptime string `json:"uptime"`
|
||||
NumGoroutine int `json:"num_goroutine"`
|
||||
Alloc string `json:"alloc"`
|
||||
TotalAlloc string `json:"total_alloc"`
|
||||
Sys string `json:"sys"`
|
||||
Lookups uint64 `json:"lookups"`
|
||||
Mallocs uint64 `json:"mallocs"`
|
||||
Frees uint64 `json:"frees"`
|
||||
HeapAlloc string `json:"heap_alloc"`
|
||||
HeapSys string `json:"heap_sys"`
|
||||
HeapIdle string `json:"heap_idle"`
|
||||
HeapInuse string `json:"heap_inuse"`
|
||||
HeapReleased string `json:"heap_released"`
|
||||
HeapObjects uint64 `json:"heap_objects"`
|
||||
StackInuse string `json:"stack_inuse"`
|
||||
StackSys string `json:"stack_sys"`
|
||||
MSpanInuse string `json:"mspan_inuse"`
|
||||
MSpanSys string `json:"mspan_sys"`
|
||||
MCacheInuse string `json:"mcache_inuse"`
|
||||
MCacheSys string `json:"mcache_sys"`
|
||||
BuckHashSys string `json:"buck_hash_sys"`
|
||||
GCSys string `json:"gc_sys"`
|
||||
OtherSys string `json:"other_sys"`
|
||||
NextGC string `json:"next_gc"`
|
||||
LastGCTime string `json:"last_gc_time"`
|
||||
PauseTotalNs string `json:"pause_total_ns"`
|
||||
LastPause string `json:"last_pause"`
|
||||
NumGC uint32 `json:"num_gc"`
|
||||
}
|
||||
|
||||
func formatDuration(d time.Duration) string {
|
||||
days := int(d.Hours()) / hoursInDay
|
||||
hours := int(d.Hours()) % hoursInDay
|
||||
minutes := int(d.Minutes()) % minutesInHour
|
||||
seconds := int(d.Seconds()) % secondsInMinute
|
||||
|
||||
var res string
|
||||
if days > 0 {
|
||||
res += fmt.Sprintf("%d天", days)
|
||||
}
|
||||
if hours > 0 {
|
||||
res += fmt.Sprintf("%d小时", hours)
|
||||
}
|
||||
if minutes > 0 {
|
||||
res += fmt.Sprintf("%d分钟", minutes)
|
||||
}
|
||||
if seconds > 0 || res == "" {
|
||||
res += fmt.Sprintf("%d秒钟", seconds)
|
||||
}
|
||||
return res
|
||||
}
|
||||
|
||||
// GetSystemStatus 获取系统状态信息
|
||||
// @Summary 获取系统状态信息
|
||||
// @Description 获取后端服务运行状态、Goroutine、内存指标等详细统计数据,需要管理员权限
|
||||
// @Tags admin
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Success 200 {object} response.Any{data=SystemStatusResponse} "获取成功"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 403 {object} response.Any "无管理员权限"
|
||||
// @Router /api/v1/admin/status [get]
|
||||
func GetSystemStatus(c *gin.Context) {
|
||||
var m runtime.MemStats
|
||||
runtime.ReadMemStats(&m)
|
||||
|
||||
uptime := formatDuration(time.Since(startTime))
|
||||
numGoroutine := runtime.NumGoroutine()
|
||||
|
||||
var lastGCTime string
|
||||
switch {
|
||||
case m.LastGC > 0 && m.LastGC <= math.MaxInt64:
|
||||
lastGCTime = formatDuration(time.Since(time.Unix(0, int64(m.LastGC))))
|
||||
case m.LastGC > 0:
|
||||
lastGCTime = "未知"
|
||||
default:
|
||||
lastGCTime = "无"
|
||||
}
|
||||
|
||||
var lastPause string
|
||||
if m.NumGC > 0 {
|
||||
lastPause = fmt.Sprintf("%.3fs", float64(m.PauseNs[(m.NumGC-1)%256])/nanosPerSecond)
|
||||
} else {
|
||||
lastPause = "0.000s"
|
||||
}
|
||||
|
||||
res := SystemStatusResponse{
|
||||
Uptime: uptime,
|
||||
NumGoroutine: numGoroutine,
|
||||
Alloc: formatBytes(m.Alloc),
|
||||
TotalAlloc: formatBytes(m.TotalAlloc),
|
||||
Sys: formatBytes(m.Sys),
|
||||
Lookups: m.Lookups,
|
||||
Mallocs: m.Mallocs,
|
||||
Frees: m.Frees,
|
||||
HeapAlloc: formatBytes(m.HeapAlloc),
|
||||
HeapSys: formatBytes(m.HeapSys),
|
||||
HeapIdle: formatBytes(m.HeapIdle),
|
||||
HeapInuse: formatBytes(m.HeapInuse),
|
||||
HeapReleased: formatBytes(m.HeapReleased),
|
||||
HeapObjects: m.HeapObjects,
|
||||
StackInuse: formatBytes(m.StackInuse),
|
||||
StackSys: formatBytes(m.StackSys),
|
||||
MSpanInuse: formatBytes(m.MSpanInuse),
|
||||
MSpanSys: formatBytes(m.MSpanSys),
|
||||
MCacheInuse: formatBytes(m.MCacheInuse),
|
||||
MCacheSys: formatBytes(m.MCacheSys),
|
||||
BuckHashSys: formatBytes(m.BuckHashSys),
|
||||
GCSys: formatBytes(m.GCSys),
|
||||
OtherSys: formatBytes(m.OtherSys),
|
||||
NextGC: formatBytes(m.NextGC),
|
||||
LastGCTime: lastGCTime,
|
||||
PauseTotalNs: fmt.Sprintf("%.1fs", float64(m.PauseTotalNs)/nanosPerSecond),
|
||||
LastPause: lastPause,
|
||||
NumGC: m.NumGC,
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, response.OK(res))
|
||||
}
|
||||
|
||||
// LogDatabaseStatus 日志库状态。
|
||||
type LogDatabaseStatus struct {
|
||||
ActiveDatabase string `json:"active_database"`
|
||||
Migration string `json:"migration"`
|
||||
RetentionDays map[string]int `json:"retention_days"`
|
||||
AvailableTargets []string `json:"available_targets"`
|
||||
}
|
||||
|
||||
// GetLogDatabaseStatus 返回当前日志库状态。
|
||||
// @Summary 获取日志数据库状态
|
||||
// @Description 返回当前日志主库、迁移状态、各库保留天数与合法迁移目标,需要管理员权限
|
||||
// @Tags admin
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Success 200 {object} response.Any{data=LogDatabaseStatus} "获取成功"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 403 {object} response.Any "无管理员权限"
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/admin/status/log-database [get]
|
||||
func GetLogDatabaseStatus(c *gin.Context) {
|
||||
ctx := c.Request.Context()
|
||||
store, err := logstore.Active(ctx)
|
||||
if err != nil {
|
||||
logger.ErrorF(ctx, "获取日志存储实例失败: %v", err)
|
||||
response.AbortInternal(c, "日志存储初始化失败")
|
||||
return
|
||||
}
|
||||
activeDB, err := store.Status.ActiveDatabase(ctx)
|
||||
if err != nil {
|
||||
logger.ErrorF(ctx, "获取日志库状态失败: %v", err)
|
||||
response.AbortInternal(c, "获取日志库状态失败")
|
||||
return
|
||||
}
|
||||
migration := "idle"
|
||||
if logstore.Migrating(ctx) {
|
||||
migration = "migrating"
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(LogDatabaseStatus{
|
||||
ActiveDatabase: activeDB,
|
||||
Migration: migration,
|
||||
RetentionDays: map[string]int{
|
||||
logDBNamePostgres: retentionOr(ctx, ConfigKeyLogRetentionDaysPostgres),
|
||||
logDBNameSQLite: retentionOr(ctx, ConfigKeyLogRetentionDaysSQLite),
|
||||
logDBNameClickHouse: retentionOr(ctx, ConfigKeyLogRetentionDaysClickHouse),
|
||||
},
|
||||
AvailableTargets: availableLogTargets(activeDB),
|
||||
}))
|
||||
}
|
||||
|
||||
func retentionOr(ctx context.Context, key string) int {
|
||||
v, err := GetIntByKey(ctx, key)
|
||||
if err != nil {
|
||||
if !errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
logger.ErrorF(ctx, "读取日志保留天数配置失败 key=%s: %v", key, err)
|
||||
}
|
||||
return defaultLogRetentionDays
|
||||
}
|
||||
if v < 1 {
|
||||
return defaultLogRetentionDays
|
||||
}
|
||||
return v
|
||||
}
|
||||
|
||||
func availableLogTargets(active string) []string {
|
||||
if active == logDBNameClickHouse {
|
||||
if config.Config.Database.Enabled {
|
||||
return []string{logDBNamePostgres}
|
||||
}
|
||||
return []string{logDBNameSQLite}
|
||||
}
|
||||
if config.Config.ClickHouse.Enabled {
|
||||
return []string{logDBNameClickHouse}
|
||||
}
|
||||
return []string{}
|
||||
}
|
||||
@@ -0,0 +1,413 @@
|
||||
// Copyright 2025 linux.do
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package admin
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"net/http"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/robfig/cron/v3"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/backend/pkg/logger"
|
||||
"github.com/Rain-kl/Wavelet/backend/pkg/response"
|
||||
"github.com/Rain-kl/Wavelet/backend/plugins/drivers/driver_asynq_cron"
|
||||
"github.com/Rain-kl/Wavelet/backend/plugins/drivers/driver_asynq_worker"
|
||||
)
|
||||
|
||||
// ListTaskTypes 获取支持的任务类型列表
|
||||
// @Summary 获取支持的任务类型
|
||||
// @Description 返回系统支持的所有可调度任务类型列表,包括任务名称、描述、是否支持时间范围等元数据,需要管理员权限
|
||||
// @Tags admin
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Success 200 {object} response.Any{data=[]driver_asynq_worker.TaskMeta} "任务类型列表"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 403 {object} response.Any "无管理员权限"
|
||||
// @Router /api/v1/admin/tasks/types [get]
|
||||
func ListTaskTypes(c *gin.Context) {
|
||||
c.JSON(http.StatusOK, response.OK(driver_asynq_worker.GetDispatchableTasks()))
|
||||
}
|
||||
|
||||
// DispatchTaskRequest 下发任务请求
|
||||
type DispatchTaskRequest struct {
|
||||
TaskType string `json:"task_type" binding:"required"`
|
||||
StartTime *time.Time `json:"start_time"`
|
||||
EndTime *time.Time `json:"end_time"`
|
||||
UserID *uint64 `json:"user_id"`
|
||||
Payload string `json:"payload"`
|
||||
}
|
||||
|
||||
// DispatchTask 下发任务
|
||||
// @Summary 下发异步任务
|
||||
// @Description 手动触发指定类型的异步任务,支持指定时间范围和用户,需要管理员权限
|
||||
// @Tags admin
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param request body DispatchTaskRequest true "任务请求参数"
|
||||
// @Success 200 {object} response.Any{data=string} "任务已入队"
|
||||
// @Failure 400 {object} response.Any "任务类型不存在或参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 403 {object} response.Any "无管理员权限"
|
||||
// @Failure 500 {object} response.Any "任务入队失败"
|
||||
// @Router /api/v1/admin/tasks/dispatch [post]
|
||||
func DispatchTask(c *gin.Context) {
|
||||
var req DispatchTaskRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
meta := driver_asynq_worker.GetTaskMeta(req.TaskType)
|
||||
if meta == nil {
|
||||
response.AbortBadRequest(c, InvalidTaskType)
|
||||
return
|
||||
}
|
||||
|
||||
var payloadBytes []byte
|
||||
if strings.TrimSpace(req.Payload) != "" {
|
||||
payloadBytes = []byte(req.Payload)
|
||||
}
|
||||
|
||||
validated, err := driver_asynq_worker.ValidateAndNormalizePayload(meta.AsynqTask, payloadBytes)
|
||||
if err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
taskID, err := driver_asynq_worker.DispatchTask(c.Request.Context(), req.TaskType, validated, "manual")
|
||||
if err != nil {
|
||||
response.AbortInternal(c, fmt.Sprintf("%s: %v", TaskDispatchFailed, err))
|
||||
return
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, response.OK(taskID))
|
||||
}
|
||||
|
||||
// ListTaskExecutions 查询任务执行记录列表
|
||||
// @Summary 查询任务执行记录
|
||||
// @Description 分页查询任务执行记录,支持按状态和任务类型筛选,需要管理员权限
|
||||
// @Tags admin
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param status query string false "状态筛选 (pending/running/succeeded/failed)"
|
||||
// @Param task_type query string false "任务类型筛选"
|
||||
// @Param page query int false "页码" default(1)
|
||||
// @Param page_size query int false "每页条数" default(20)
|
||||
// @Success 200 {object} response.Any{data=object} "任务执行记录列表"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 403 {object} response.Any "无管理员权限"
|
||||
// @Router /api/v1/admin/tasks/executions [get]
|
||||
func ListTaskExecutions(c *gin.Context) {
|
||||
var req ListTaskExecutionsRequest
|
||||
if err := c.ShouldBindQuery(&req); err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
if req.TaskType != "" {
|
||||
if meta := driver_asynq_worker.GetTaskMeta(req.TaskType); meta != nil {
|
||||
req.TaskType = meta.AsynqTask
|
||||
}
|
||||
}
|
||||
|
||||
executions, total, err := ListTaskExecutionRecords(c.Request.Context(), req)
|
||||
if err != nil {
|
||||
response.AbortInternal(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, response.OK(gin.H{
|
||||
"items": executions,
|
||||
"total": total,
|
||||
"page": req.Page,
|
||||
"page_size": req.PageSize,
|
||||
}))
|
||||
}
|
||||
|
||||
// GetTaskExecution 查询单条任务执行详情
|
||||
// @Summary 查询任务执行详情
|
||||
// @Description 根据 ID 查询任务执行记录详情,包含完整执行日志,需要管理员权限
|
||||
// @Tags admin
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param id path int true "任务执行记录 ID"
|
||||
// @Success 200 {object} response.Any{data=TaskExecution} "任务执行详情"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 403 {object} response.Any "无管理员权限"
|
||||
// @Failure 404 {object} response.Any "记录不存在"
|
||||
// @Router /api/v1/admin/tasks/executions/{id} [get]
|
||||
func GetTaskExecution(c *gin.Context) {
|
||||
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
|
||||
if err != nil {
|
||||
response.AbortBadRequest(c, InvalidTaskExecutionID)
|
||||
return
|
||||
}
|
||||
|
||||
execution, err := GetTaskExecutionByID(c.Request.Context(), id)
|
||||
if err != nil {
|
||||
response.AbortNotFound(c, TaskNotFound)
|
||||
return
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, response.OK(execution))
|
||||
}
|
||||
|
||||
// RetryTask 重试失败的任务
|
||||
// @Summary 重试失败任务
|
||||
// @Description 重新下发一条失败的任务,创建新的执行记录,需要管理员权限
|
||||
// @Tags admin
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param id path int true "任务执行记录 ID"
|
||||
// @Success 200 {object} response.Any{data=string} "新任务的 TaskID"
|
||||
// @Failure 400 {object} response.Any "任务不支持重试或参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 403 {object} response.Any "无管理员权限"
|
||||
// @Failure 404 {object} response.Any "记录不存在"
|
||||
// @Failure 500 {object} response.Any "重试失败"
|
||||
// @Router /api/v1/admin/tasks/executions/{id}/retry [post]
|
||||
func RetryTask(c *gin.Context) {
|
||||
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
|
||||
if err != nil {
|
||||
response.AbortBadRequest(c, InvalidTaskExecutionID)
|
||||
return
|
||||
}
|
||||
|
||||
newTaskID, err := driver_asynq_worker.RetryTask(c.Request.Context(), id)
|
||||
if err != nil {
|
||||
errMsg := err.Error()
|
||||
switch {
|
||||
case strings.Contains(errMsg, "不存在"):
|
||||
response.AbortNotFound(c, errMsg)
|
||||
case strings.Contains(errMsg, "只有失败的任务") || strings.Contains(errMsg, "不支持重试") || strings.Contains(errMsg, "已达到最大重试"):
|
||||
response.AbortBadRequest(c, errMsg)
|
||||
default:
|
||||
response.AbortInternal(c, fmt.Sprintf("%s: %v", TaskRetryFailed, err))
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, response.OK(newTaskID))
|
||||
}
|
||||
|
||||
// ListSchedules 获取定时任务列表
|
||||
// @Summary 获取定时任务列表
|
||||
// @Description 返回系统所有的定时任务配置列表,包括名称、关联的异步任务类型、Cron 表达式和启用状态,需要管理员权限
|
||||
// @Tags admin
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Success 200 {object} response.Any{data=[]Schedule} "定时任务列表"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 403 {object} response.Any "无管理员权限"
|
||||
// @Router /api/v1/admin/tasks/schedules [get]
|
||||
func ListSchedules(c *gin.Context) {
|
||||
schedules, err := ListSchedulesRecord(c.Request.Context())
|
||||
if err != nil {
|
||||
response.AbortInternal(c, err.Error())
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(schedules))
|
||||
}
|
||||
|
||||
// CreateScheduleRequest 创建定时任务请求
|
||||
type CreateScheduleRequest struct {
|
||||
Name string `json:"name" binding:"required"`
|
||||
TaskType string `json:"task_type" binding:"required"`
|
||||
Cron string `json:"cron" binding:"required"`
|
||||
Payload string `json:"payload"`
|
||||
IsActive *bool `json:"is_active" binding:"required"`
|
||||
}
|
||||
|
||||
// CreateSchedule 创建定时任务
|
||||
// @Summary 创建定时任务
|
||||
// @Description 新增一个动态定时任务配置,关联已有的异步任务,配置 Cron 表达式和执行参数,并触发调度器热加载,需要管理员权限
|
||||
// @Tags admin
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param request body CreateScheduleRequest true "创建定时任务请求参数"
|
||||
// @Success 200 {object} response.Any{data=Schedule} "创建成功的定时任务信息"
|
||||
// @Failure 400 {object} response.Any "Cron 表达式无效、异步任务类型不存在或参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 403 {object} response.Any "无管理员权限"
|
||||
// @Failure 500 {object} response.Any "保存定时任务失败"
|
||||
// @Router /api/v1/admin/tasks/schedules [post]
|
||||
func CreateSchedule(c *gin.Context) {
|
||||
var req CreateScheduleRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
// 校验 Cron 表达式
|
||||
if _, err := cron.ParseStandard(req.Cron); err != nil {
|
||||
response.AbortBadRequest(c, InvalidCronExpression)
|
||||
return
|
||||
}
|
||||
|
||||
// 校验关联的异步任务类型
|
||||
meta := driver_asynq_worker.GetTaskMeta(req.TaskType)
|
||||
if meta == nil {
|
||||
response.AbortBadRequest(c, InvalidTaskType)
|
||||
return
|
||||
}
|
||||
|
||||
// 校验并规范化 Payload
|
||||
var payloadBytes []byte
|
||||
if strings.TrimSpace(req.Payload) != "" {
|
||||
payloadBytes = []byte(req.Payload)
|
||||
}
|
||||
validated, err := driver_asynq_worker.ValidateAndNormalizePayload(meta.AsynqTask, payloadBytes)
|
||||
if err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
schedule := &Schedule{
|
||||
Name: req.Name,
|
||||
TaskType: req.TaskType,
|
||||
Cron: req.Cron,
|
||||
Payload: string(validated),
|
||||
IsActive: *req.IsActive,
|
||||
}
|
||||
|
||||
if err := CreateScheduleRecord(c.Request.Context(), schedule); err != nil {
|
||||
response.AbortInternal(c, fmt.Sprintf("%s: %v", ScheduleSaveFailed, err))
|
||||
return
|
||||
}
|
||||
|
||||
// 触发调度服务重载
|
||||
if err := driver_asynq_cron.ReloadScheduler(); err != nil {
|
||||
logger.ErrorF(c.Request.Context(), "[TaskAdmin] 重载调度器失败: %v", err)
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, response.OK(schedule))
|
||||
}
|
||||
|
||||
// UpdateScheduleRequest 修改定时任务请求
|
||||
type UpdateScheduleRequest struct {
|
||||
Name string `json:"name" binding:"required"`
|
||||
TaskType string `json:"task_type" binding:"required"`
|
||||
Cron string `json:"cron" binding:"required"`
|
||||
Payload string `json:"payload"`
|
||||
IsActive *bool `json:"is_active" binding:"required"`
|
||||
}
|
||||
|
||||
// UpdateSchedule 修改定时任务
|
||||
// @Summary 修改定时任务
|
||||
// @Description 修改一个定时任务的配置(名称、Cron 表达式、异步任务参数和是否启用等),并触发调度器热加载,需要管理员权限
|
||||
// @Tags admin
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param id path int true "定时任务 ID"
|
||||
// @Param request body UpdateScheduleRequest true "修改定时任务请求参数"
|
||||
// @Success 200 {object} response.Any{data=Schedule} "修改后的定时任务信息"
|
||||
// @Failure 400 {object} response.Any "Cron 表达式无效、参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 403 {object} response.Any "无管理员权限"
|
||||
// @Failure 404 {object} response.Any "定时任务不存在"
|
||||
// @Failure 500 {object} response.Any "修改定时任务失败"
|
||||
// @Router /api/v1/admin/tasks/schedules/{id} [put]
|
||||
func UpdateSchedule(c *gin.Context) {
|
||||
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
|
||||
if err != nil {
|
||||
response.AbortBadRequest(c, "无效的定时任务ID")
|
||||
return
|
||||
}
|
||||
|
||||
var req UpdateScheduleRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
schedule, err := GetScheduleByID(c.Request.Context(), id)
|
||||
if err != nil {
|
||||
response.AbortNotFound(c, ScheduleNotFound)
|
||||
return
|
||||
}
|
||||
|
||||
// 校验 Cron 表达式
|
||||
if _, err := cron.ParseStandard(req.Cron); err != nil {
|
||||
response.AbortBadRequest(c, InvalidCronExpression)
|
||||
return
|
||||
}
|
||||
|
||||
// 校验关联的异步任务类型
|
||||
meta := driver_asynq_worker.GetTaskMeta(req.TaskType)
|
||||
if meta == nil {
|
||||
response.AbortBadRequest(c, InvalidTaskType)
|
||||
return
|
||||
}
|
||||
|
||||
// 校验并规范化 Payload
|
||||
var payloadBytes []byte
|
||||
if strings.TrimSpace(req.Payload) != "" {
|
||||
payloadBytes = []byte(req.Payload)
|
||||
}
|
||||
validated, err := driver_asynq_worker.ValidateAndNormalizePayload(meta.AsynqTask, payloadBytes)
|
||||
if err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
schedule.Name = req.Name
|
||||
schedule.TaskType = req.TaskType
|
||||
schedule.Cron = req.Cron
|
||||
schedule.Payload = string(validated)
|
||||
schedule.IsActive = *req.IsActive
|
||||
|
||||
if err := UpdateScheduleRecord(c.Request.Context(), schedule); err != nil {
|
||||
response.AbortInternal(c, fmt.Sprintf("%s: %v", ScheduleSaveFailed, err))
|
||||
return
|
||||
}
|
||||
|
||||
// 触发调度服务重载
|
||||
if err := driver_asynq_cron.ReloadScheduler(); err != nil {
|
||||
logger.ErrorF(c.Request.Context(), "[TaskAdmin] 重载调度器失败: %v", err)
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, response.OK(schedule))
|
||||
}
|
||||
|
||||
// DeleteSchedule 删除定时任务
|
||||
// @Summary 删除定时任务
|
||||
// @Description 删除指定的定时任务配置,并触发调度器热加载,需要管理员权限
|
||||
// @Tags admin
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param id path int true "定时任务 ID"
|
||||
// @Success 200 {object} response.Any{data=string} "删除结果"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 403 {object} response.Any "无管理员权限"
|
||||
// @Failure 500 {object} response.Any "删除定时任务失败"
|
||||
// @Router /api/v1/admin/tasks/schedules/{id} [delete]
|
||||
func DeleteSchedule(c *gin.Context) {
|
||||
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
|
||||
if err != nil {
|
||||
response.AbortBadRequest(c, "无效的定时任务ID")
|
||||
return
|
||||
}
|
||||
|
||||
if err := DeleteScheduleRecord(c.Request.Context(), id); err != nil {
|
||||
response.AbortInternal(c, fmt.Sprintf("%s: %v", ScheduleDeleteFailed, err))
|
||||
return
|
||||
}
|
||||
|
||||
// 触发调度服务重载
|
||||
if err := driver_asynq_cron.ReloadScheduler(); err != nil {
|
||||
logger.ErrorF(c.Request.Context(), "[TaskAdmin] 重载调度器失败: %v", err)
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, response.OKNil())
|
||||
}
|
||||
@@ -0,0 +1,243 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package admin
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"net/http"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"gorm.io/gorm"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/backend/pkg/response"
|
||||
)
|
||||
|
||||
// CreateTemplateRequest 创建模板请求
|
||||
type CreateTemplateRequest struct {
|
||||
Key string `json:"key" binding:"required,max=80"`
|
||||
Name string `json:"name" binding:"required,max=100"`
|
||||
Type string `json:"type" binding:"required,max=20"`
|
||||
Subject string `json:"subject" binding:"max=255"`
|
||||
Content string `json:"content" binding:"required"`
|
||||
Description string `json:"description" binding:"max=255"`
|
||||
}
|
||||
|
||||
// UpdateTemplateRequest 更新模板请求
|
||||
type UpdateTemplateRequest struct {
|
||||
Name string `json:"name" binding:"required,max=100"`
|
||||
Type string `json:"type" binding:"required,max=20"`
|
||||
Subject string `json:"subject" binding:"max=255"`
|
||||
Content string `json:"content" binding:"required"`
|
||||
Description string `json:"description" binding:"max=255"`
|
||||
}
|
||||
|
||||
func abortTemplateLogicError(c *gin.Context, err error) bool {
|
||||
if err == nil {
|
||||
return false
|
||||
}
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
response.AbortNotFound(c, TemplateNotFound)
|
||||
return true
|
||||
}
|
||||
msg := err.Error()
|
||||
switch msg {
|
||||
case TemplateKeyExists, SystemTemplateCannotDelete:
|
||||
response.AbortBadRequest(c, msg)
|
||||
return true
|
||||
}
|
||||
response.AbortInternal(c, msg)
|
||||
return true
|
||||
}
|
||||
|
||||
// CreateTemplate 创建模板
|
||||
// @Summary 创建模板
|
||||
// @Description 创建一条新的自定义通知模板,模板标识符(Key)不可重复,需要管理员权限
|
||||
// @Tags admin
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param request body CreateTemplateRequest true "创建请求参数"
|
||||
// @Success 200 {object} response.Any{data=string} "创建成功"
|
||||
// @Failure 400 {object} response.Any "参数错误或模板标识符已存在"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 403 {object} response.Any "无管理员权限"
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/admin/templates [post]
|
||||
func CreateTemplate(c *gin.Context) {
|
||||
var req CreateTemplateRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
tmpl, err := createTemplate(c.Request.Context(), req)
|
||||
if abortTemplateLogicError(c, err) {
|
||||
return
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, response.OK(tmpl))
|
||||
}
|
||||
|
||||
// ListTemplates 获取模板列表
|
||||
// @Summary 获取模板列表
|
||||
// @Description 返回所有通知模板列表,需要管理员权限
|
||||
// @Tags admin
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Success 200 {object} response.Any{data=[]Template} "模板列表"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 403 {object} response.Any "无管理员权限"
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/admin/templates [get]
|
||||
func ListTemplates(c *gin.Context) {
|
||||
templates, err := listTemplates(c.Request.Context())
|
||||
if err != nil {
|
||||
response.AbortInternal(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, response.OK(templates))
|
||||
}
|
||||
|
||||
// GetTemplate 获取单个模板
|
||||
// @Summary 获取单个模板
|
||||
// @Description 根据模板标识符获取对应的模板详情,需要管理员权限
|
||||
// @Tags admin
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param key path string true "模板标识符"
|
||||
// @Success 200 {object} response.Any{data=Template} "模板详情"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 403 {object} response.Any "无管理员权限"
|
||||
// @Failure 404 {object} response.Any "模板不存在"
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/admin/templates/{key} [get]
|
||||
func GetTemplate(c *gin.Context) {
|
||||
tmpl, err := getTemplate(c.Request.Context(), c.Param("key"))
|
||||
if abortTemplateLogicError(c, err) {
|
||||
return
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, response.OK(tmpl))
|
||||
}
|
||||
|
||||
// UpdateTemplate 更新模板
|
||||
// @Summary 更新模板
|
||||
// @Description 根据模板标识符更新对应的模板内容,需要管理员权限
|
||||
// @Tags admin
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param key path string true "模板标识符"
|
||||
// @Param request body UpdateTemplateRequest true "更新请求参数"
|
||||
// @Success 200 {object} response.Any{data=Template} "更新成功"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 403 {object} response.Any "无管理员权限"
|
||||
// @Failure 404 {object} response.Any "模板不存在"
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/admin/templates/{key} [put]
|
||||
func UpdateTemplate(c *gin.Context) {
|
||||
var req UpdateTemplateRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
tmpl, err := updateTemplate(c.Request.Context(), c.Param("key"), req)
|
||||
if abortTemplateLogicError(c, err) {
|
||||
return
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, response.OK(tmpl))
|
||||
}
|
||||
|
||||
// DeleteTemplate 删除模板
|
||||
// @Summary 删除模板
|
||||
// @Description 根据模板标识符删除对应模板,系统预置模板不可删除,需要管理员权限
|
||||
// @Tags admin
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param key path string true "模板标识符"
|
||||
// @Success 200 {object} response.Any{data=string} "删除成功"
|
||||
// @Failure 400 {object} response.Any "不可删除系统模板"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 403 {object} response.Any "无管理员权限"
|
||||
// @Failure 404 {object} response.Any "模板不存在"
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/admin/templates/{key} [delete]
|
||||
func DeleteTemplate(c *gin.Context) {
|
||||
if err := deleteTemplate(c.Request.Context(), c.Param("key")); abortTemplateLogicError(c, err) {
|
||||
return
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, response.OKNil())
|
||||
}
|
||||
|
||||
func createTemplate(ctx context.Context, req CreateTemplateRequest) (Template, error) {
|
||||
exists, err := TemplateExistsByKey(ctx, req.Key)
|
||||
if err != nil {
|
||||
return Template{}, err
|
||||
}
|
||||
if exists {
|
||||
return Template{}, errors.New(TemplateKeyExists)
|
||||
}
|
||||
|
||||
tmpl := Template{
|
||||
Key: req.Key,
|
||||
Name: req.Name,
|
||||
Type: req.Type,
|
||||
Subject: req.Subject,
|
||||
Content: req.Content,
|
||||
Description: req.Description,
|
||||
IsSystem: false,
|
||||
}
|
||||
if err := tmpl.Validate(); err != nil {
|
||||
return Template{}, err
|
||||
}
|
||||
if err := CreateTemplateRecord(ctx, &tmpl); err != nil {
|
||||
return Template{}, err
|
||||
}
|
||||
return tmpl, nil
|
||||
}
|
||||
|
||||
func listTemplates(ctx context.Context) ([]Template, error) {
|
||||
return ListTemplatesRecord(ctx)
|
||||
}
|
||||
|
||||
func getTemplate(ctx context.Context, key string) (Template, error) {
|
||||
return GetTemplateByKey(ctx, key)
|
||||
}
|
||||
|
||||
func updateTemplate(ctx context.Context, key string, req UpdateTemplateRequest) (Template, error) {
|
||||
tmpl, err := GetTemplateByKey(ctx, key)
|
||||
if err != nil {
|
||||
return Template{}, err
|
||||
}
|
||||
|
||||
tmpl.Name = req.Name
|
||||
tmpl.Type = req.Type
|
||||
tmpl.Subject = req.Subject
|
||||
tmpl.Content = req.Content
|
||||
tmpl.Description = req.Description
|
||||
if err := tmpl.Validate(); err != nil {
|
||||
return Template{}, err
|
||||
}
|
||||
if err := SaveTemplateRecord(ctx, &tmpl); err != nil {
|
||||
return Template{}, err
|
||||
}
|
||||
return tmpl, nil
|
||||
}
|
||||
|
||||
func deleteTemplate(ctx context.Context, key string) error {
|
||||
tmpl, err := GetTemplateByKey(ctx, key)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if tmpl.IsSystem {
|
||||
return errors.New(SystemTemplateCannotDelete)
|
||||
}
|
||||
return DeleteTemplateRecord(ctx, &tmpl)
|
||||
}
|
||||
@@ -0,0 +1,695 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package admin
|
||||
|
||||
import (
|
||||
"archive/tar"
|
||||
"archive/zip"
|
||||
"compress/gzip"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"runtime"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"golang.org/x/mod/semver"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/backend/pkg/buildinfo"
|
||||
"github.com/Rain-kl/Wavelet/backend/pkg/logger"
|
||||
"github.com/Rain-kl/Wavelet/backend/pkg/response"
|
||||
"github.com/Rain-kl/Wavelet/backend/pkg/util"
|
||||
)
|
||||
|
||||
const (
|
||||
githubAPIBaseURL = "https://api.github.com"
|
||||
maxArchiveSize = int64(1024 * 1024 * 1024)
|
||||
maxReleaseSize = int64(4 * 1024 * 1024)
|
||||
repositoryParts = 2
|
||||
windowsOS = "windows"
|
||||
archiveFileMode = 0o600
|
||||
stagedBinaryMode = 0o700
|
||||
)
|
||||
|
||||
type releaseAsset struct {
|
||||
Name string `json:"name"`
|
||||
BrowserDownloadURL string `json:"browser_download_url"`
|
||||
Size int64 `json:"size"`
|
||||
State string `json:"state"`
|
||||
}
|
||||
|
||||
type githubRelease struct {
|
||||
TagName string `json:"tag_name"`
|
||||
Name string `json:"name"`
|
||||
Body string `json:"body"`
|
||||
HTMLURL string `json:"html_url"`
|
||||
Draft bool `json:"draft"`
|
||||
Prerelease bool `json:"prerelease"`
|
||||
Published time.Time `json:"published_at"`
|
||||
Assets []releaseAsset `json:"assets"`
|
||||
}
|
||||
|
||||
// UpdaterStatus describes the current build and the newest compatible upstream release.
|
||||
type UpdaterStatus struct {
|
||||
CurrentVersion string `json:"current_version"`
|
||||
BuildTime string `json:"build_time"`
|
||||
LatestVersion string `json:"latest_version"`
|
||||
UpdateAvailable bool `json:"update_available"`
|
||||
CanUpgrade bool `json:"can_upgrade"`
|
||||
Prerelease bool `json:"prerelease"`
|
||||
ReleaseName string `json:"release_name"`
|
||||
ReleaseNotes string `json:"release_notes"`
|
||||
ReleaseURL string `json:"release_url"`
|
||||
PublishedAt string `json:"published_at"`
|
||||
UpstreamRepository string `json:"upstream_repository"`
|
||||
AssetName string `json:"asset_name"`
|
||||
Platform string `json:"platform"`
|
||||
}
|
||||
|
||||
type releaseClient interface {
|
||||
Do(req *http.Request) (*http.Response, error)
|
||||
}
|
||||
|
||||
type updaterManager struct {
|
||||
client releaseClient
|
||||
mu sync.Mutex
|
||||
upgrading bool
|
||||
}
|
||||
|
||||
var defaultUpdaterManager = &updaterManager{
|
||||
client: &http.Client{Timeout: 10 * time.Minute},
|
||||
}
|
||||
|
||||
// GetUpdateStatus 获取应用更新状态
|
||||
// @Summary 获取应用更新状态
|
||||
// @Description 从系统配置指定的 GitHub 上游仓库查询最新兼容 Release,并与当前服务版本比较
|
||||
// @Tags admin
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Success 200 {object} response.Any{data=UpdaterStatus} "更新状态"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 403 {object} response.Any "无管理员权限"
|
||||
// @Failure 500 {object} response.Any "查询失败"
|
||||
// @Router /api/v1/admin/update [get]
|
||||
func GetUpdateStatus(c *gin.Context) {
|
||||
status, _, err := defaultUpdaterManager.status(c.Request.Context())
|
||||
if err != nil {
|
||||
logger.ErrorF(c.Request.Context(), "[Updater] check release failed: %v", err)
|
||||
response.AbortInternal(c, err.Error())
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, response.OK(status))
|
||||
}
|
||||
|
||||
// ApplyUpdate 下载并应用应用更新
|
||||
// @Summary 下载并应用应用更新
|
||||
// @Description 下载当前平台对应的 GitHub Actions Release 资产,替换当前二进制并重启进程
|
||||
// @Tags admin
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Success 200 {object} response.Any "升级已准备并即将重启"
|
||||
// @Failure 400 {object} response.Any "当前版本不可升级"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 403 {object} response.Any "无管理员权限"
|
||||
// @Failure 500 {object} response.Any "升级准备失败"
|
||||
// @Router /api/v1/admin/update/apply [post]
|
||||
func ApplyUpdate(c *gin.Context) {
|
||||
executable, stagedBinary, err := defaultUpdaterManager.prepareUpgrade(c.Request.Context())
|
||||
if err != nil {
|
||||
logger.ErrorF(c.Request.Context(), "[Updater] prepare upgrade failed: %v", err)
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
logger.InfoF(c.Request.Context(), "[Updater] upgrade prepared; restarting with %s", stagedBinary)
|
||||
c.JSON(http.StatusOK, response.OKNil())
|
||||
|
||||
util.Go(func() {
|
||||
time.Sleep(time.Second)
|
||||
if err := replaceAndRestart(executable, stagedBinary); err != nil {
|
||||
defaultUpdaterManager.finishUpgrade()
|
||||
logger.ErrorF(context.Background(), "[Updater] replace and restart failed: %v", err)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func normalizeVersion(version string) string {
|
||||
version = strings.TrimSpace(version)
|
||||
if version == "" || version == "dev" {
|
||||
return ""
|
||||
}
|
||||
if !strings.HasPrefix(version, "v") {
|
||||
version = "v" + version
|
||||
}
|
||||
if !semver.IsValid(version) {
|
||||
return ""
|
||||
}
|
||||
return version
|
||||
}
|
||||
|
||||
func parseRepository(raw string) (string, error) {
|
||||
raw = strings.TrimSpace(raw)
|
||||
if raw == "" {
|
||||
return "", errors.New(errInvalidRepository)
|
||||
}
|
||||
|
||||
if !strings.Contains(raw, "://") {
|
||||
repo := strings.TrimSuffix(strings.Trim(raw, "/"), ".git")
|
||||
if len(strings.Split(repo, "/")) == repositoryParts {
|
||||
return repo, nil
|
||||
}
|
||||
return "", errors.New(errInvalidRepository)
|
||||
}
|
||||
|
||||
parsed, err := url.Parse(raw)
|
||||
if err != nil || !strings.EqualFold(parsed.Hostname(), "github.com") {
|
||||
return "", errors.New(errInvalidRepository)
|
||||
}
|
||||
repo := strings.TrimSuffix(strings.Trim(parsed.Path, "/"), ".git")
|
||||
if len(strings.Split(repo, "/")) != repositoryParts {
|
||||
return "", errors.New(errInvalidRepository)
|
||||
}
|
||||
return repo, nil
|
||||
}
|
||||
|
||||
func expectedAssetName(tag string) string {
|
||||
extension := "tar.gz"
|
||||
if runtime.GOOS == windowsOS {
|
||||
extension = "zip"
|
||||
}
|
||||
return fmt.Sprintf("wavelet_%s_%s_%s.%s", tag, runtime.GOOS, runtime.GOARCH, extension)
|
||||
}
|
||||
|
||||
func expectedAssetNames(repository, tag string) []string {
|
||||
names := []string{expectedAssetName(tag)}
|
||||
if parts := strings.Split(repository, "/"); len(parts) == repositoryParts {
|
||||
repoName := parts[1]
|
||||
if repoName != "wavelet" {
|
||||
extension := "tar.gz"
|
||||
if runtime.GOOS == windowsOS {
|
||||
extension = "zip"
|
||||
}
|
||||
names = append(names, fmt.Sprintf("%s_%s_%s_%s.%s", repoName, tag, runtime.GOOS, runtime.GOARCH, extension))
|
||||
}
|
||||
}
|
||||
return names
|
||||
}
|
||||
|
||||
func selectLatestRelease(repository string, releases []githubRelease) (githubRelease, releaseAsset, error) {
|
||||
var selected githubRelease
|
||||
var selectedAsset releaseAsset
|
||||
selectedVersion := ""
|
||||
|
||||
for _, release := range releases {
|
||||
version := normalizeVersion(release.TagName)
|
||||
if release.Draft || version == "" {
|
||||
continue
|
||||
}
|
||||
expectedNames := expectedAssetNames(repository, release.TagName)
|
||||
for _, asset := range release.Assets {
|
||||
matched := false
|
||||
for _, name := range expectedNames {
|
||||
if asset.Name == name {
|
||||
matched = true
|
||||
break
|
||||
}
|
||||
}
|
||||
if !matched || asset.BrowserDownloadURL == "" || asset.State != "uploaded" {
|
||||
continue
|
||||
}
|
||||
if selectedVersion == "" || semver.Compare(version, selectedVersion) > 0 {
|
||||
selected = release
|
||||
selectedAsset = asset
|
||||
selectedVersion = version
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if selectedVersion == "" {
|
||||
return githubRelease{}, releaseAsset{}, errors.New(errNoCompatibleRelease)
|
||||
}
|
||||
return selected, selectedAsset, nil
|
||||
}
|
||||
|
||||
func (m *updaterManager) fetchRelease(ctx context.Context, repository string) (githubRelease, releaseAsset, error) {
|
||||
req, err := http.NewRequestWithContext(
|
||||
ctx,
|
||||
http.MethodGet,
|
||||
fmt.Sprintf("%s/repos/%s/releases?per_page=30", githubAPIBaseURL, repository),
|
||||
nil,
|
||||
)
|
||||
if err != nil {
|
||||
return githubRelease{}, releaseAsset{}, fmt.Errorf("%s: %w", errReleaseRequestFailed, err)
|
||||
}
|
||||
req.Header.Set("Accept", "application/vnd.github+json")
|
||||
req.Header.Set("User-Agent", "Wavelet-Updater")
|
||||
req.Header.Set("X-GitHub-Api-Version", "2022-11-28")
|
||||
|
||||
resp, err := m.client.Do(req)
|
||||
if err != nil {
|
||||
return githubRelease{}, releaseAsset{}, fmt.Errorf("%s: %w", errReleaseRequestFailed, err)
|
||||
}
|
||||
defer func() {
|
||||
_ = resp.Body.Close()
|
||||
}()
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return githubRelease{}, releaseAsset{}, fmt.Errorf("%s: HTTP %d", errReleaseRequestFailed, resp.StatusCode)
|
||||
}
|
||||
|
||||
var releases []githubRelease
|
||||
decoder := json.NewDecoder(io.LimitReader(resp.Body, maxReleaseSize))
|
||||
if err := decoder.Decode(&releases); err != nil {
|
||||
return githubRelease{}, releaseAsset{}, fmt.Errorf("%s: %w", errReleaseResponseInvalid, err)
|
||||
}
|
||||
|
||||
release, asset, err := selectLatestRelease(repository, releases)
|
||||
if err != nil {
|
||||
return githubRelease{}, releaseAsset{}, err
|
||||
}
|
||||
logger.InfoF(ctx, "[Updater] Selected latest compatible release: %s (Asset: %s)", release.TagName, asset.Name)
|
||||
return release, asset, nil
|
||||
}
|
||||
|
||||
func loadRepository(ctx context.Context) (string, error) {
|
||||
config, err := GetSystemConfigByKey(ctx, ConfigKeyUpdateUpstreamRepository)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("%s: %w", errInvalidRepository, err)
|
||||
}
|
||||
return parseRepository(config.Value)
|
||||
}
|
||||
|
||||
func (m *updaterManager) status(ctx context.Context) (UpdaterStatus, releaseAsset, error) {
|
||||
upstreamRepo, err := loadRepository(ctx)
|
||||
if err != nil {
|
||||
return UpdaterStatus{}, releaseAsset{}, err
|
||||
}
|
||||
release, asset, err := m.fetchRelease(ctx, upstreamRepo)
|
||||
if err != nil {
|
||||
return UpdaterStatus{}, releaseAsset{}, err
|
||||
}
|
||||
|
||||
currentVersion := normalizeVersion(buildinfo.Version)
|
||||
latestVersion := normalizeVersion(release.TagName)
|
||||
updateAvailable := currentVersion != "" && semver.Compare(latestVersion, currentVersion) > 0
|
||||
|
||||
logger.InfoF(ctx, "[Updater] Check update complete. current: %s, latest: %s, update_available: %t", buildinfo.Version, release.TagName, updateAvailable)
|
||||
|
||||
return UpdaterStatus{
|
||||
CurrentVersion: buildinfo.Version,
|
||||
BuildTime: buildinfo.BuildTime,
|
||||
LatestVersion: release.TagName,
|
||||
UpdateAvailable: updateAvailable,
|
||||
CanUpgrade: updateAvailable && runtime.GOOS != windowsOS,
|
||||
Prerelease: release.Prerelease,
|
||||
ReleaseName: release.Name,
|
||||
ReleaseNotes: release.Body,
|
||||
ReleaseURL: release.HTMLURL,
|
||||
PublishedAt: release.Published.Format(time.RFC3339),
|
||||
UpstreamRepository: upstreamRepo,
|
||||
AssetName: asset.Name,
|
||||
Platform: runtime.GOOS + "/" + runtime.GOARCH,
|
||||
}, asset, nil
|
||||
}
|
||||
|
||||
func downloadArchive(ctx context.Context, client releaseClient, asset releaseAsset, destination string) error {
|
||||
if asset.Size <= 0 || asset.Size > maxArchiveSize {
|
||||
return fmt.Errorf("release 资产大小无效: %d", asset.Size)
|
||||
}
|
||||
logger.InfoF(ctx, "[Updater] Downloading release asset: %s", asset.Name)
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodGet, asset.BrowserDownloadURL, nil)
|
||||
if err != nil {
|
||||
return fmt.Errorf("创建升级下载请求失败: %w", err)
|
||||
}
|
||||
req.Header.Set("User-Agent", "Wavelet-Updater")
|
||||
|
||||
resp, err := client.Do(req)
|
||||
if err != nil {
|
||||
return fmt.Errorf("下载升级资产失败: %w", err)
|
||||
}
|
||||
defer func() {
|
||||
_ = resp.Body.Close()
|
||||
}()
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return fmt.Errorf("下载升级资产失败: HTTP %d", resp.StatusCode)
|
||||
}
|
||||
|
||||
//nolint:gosec // updater download destination is validated
|
||||
file, err := os.OpenFile(destination, os.O_CREATE|os.O_EXCL|os.O_WRONLY, archiveFileMode)
|
||||
if err != nil {
|
||||
return fmt.Errorf("创建升级归档失败: %w", err)
|
||||
}
|
||||
|
||||
written, err := io.Copy(file, io.LimitReader(resp.Body, maxArchiveSize+1))
|
||||
if err != nil {
|
||||
_ = file.Close()
|
||||
return fmt.Errorf("写入升级归档失败: %w", err)
|
||||
}
|
||||
if err := file.Close(); err != nil {
|
||||
return fmt.Errorf("关闭升级归档失败: %w", err)
|
||||
}
|
||||
if written > maxArchiveSize || written != asset.Size {
|
||||
return fmt.Errorf("升级归档大小不匹配: got %d, want %d", written, asset.Size)
|
||||
}
|
||||
logger.InfoF(ctx, "[Updater] Successfully downloaded release asset to %s", destination)
|
||||
return nil
|
||||
}
|
||||
|
||||
func safeArchivePath(destination, name string) (string, error) {
|
||||
cleanName := filepath.Clean(name)
|
||||
if filepath.IsAbs(cleanName) || cleanName == "." || strings.HasPrefix(cleanName, ".."+string(filepath.Separator)) {
|
||||
return "", fmt.Errorf("归档包含非法路径: %s", name)
|
||||
}
|
||||
target := filepath.Join(destination, cleanName)
|
||||
relative, err := filepath.Rel(destination, target)
|
||||
if err != nil || relative == ".." || strings.HasPrefix(relative, ".."+string(filepath.Separator)) {
|
||||
return "", fmt.Errorf("归档路径越界: %s", name)
|
||||
}
|
||||
return target, nil
|
||||
}
|
||||
|
||||
func matchBinaryName(name string, candidates []string) bool {
|
||||
for _, candidate := range candidates {
|
||||
if runtime.GOOS == windowsOS {
|
||||
if strings.EqualFold(name, candidate) {
|
||||
return true
|
||||
}
|
||||
} else {
|
||||
if name == candidate {
|
||||
return true
|
||||
}
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func getCandidateBinaryNames(executable string, repository string) []string {
|
||||
execName := filepath.Base(executable)
|
||||
names := []string{execName}
|
||||
|
||||
addName := func(base string) {
|
||||
name := base
|
||||
if runtime.GOOS == windowsOS && !strings.HasSuffix(strings.ToLower(name), ".exe") {
|
||||
name += ".exe"
|
||||
}
|
||||
for _, existing := range names {
|
||||
if existing == name {
|
||||
return
|
||||
}
|
||||
}
|
||||
names = append(names, name)
|
||||
}
|
||||
|
||||
if parts := strings.Split(repository, "/"); len(parts) == repositoryParts {
|
||||
addName(parts[1])
|
||||
}
|
||||
addName("wavelet")
|
||||
|
||||
return names
|
||||
}
|
||||
|
||||
func isLikelyBinary(name string, isDir bool, mode os.FileMode) bool {
|
||||
if isDir {
|
||||
return false
|
||||
}
|
||||
base := strings.ToLower(filepath.Base(name))
|
||||
|
||||
exclusions := []string{
|
||||
"license", "licence", "copying", "notice", "readme", "changelog",
|
||||
}
|
||||
for _, excl := range exclusions {
|
||||
if strings.HasPrefix(base, excl) {
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
if runtime.GOOS == windowsOS {
|
||||
return filepath.Ext(base) == ".exe"
|
||||
}
|
||||
|
||||
return (mode.Perm()&0111 != 0) || (filepath.Ext(base) == "")
|
||||
}
|
||||
|
||||
func findBinaryInTarGz(archivePath string, candidates []string) (string, error) {
|
||||
//nolint:gosec // updater archivePath is verified
|
||||
file, err := os.Open(archivePath)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
defer func() {
|
||||
_ = file.Close()
|
||||
}()
|
||||
|
||||
gzipReader, err := gzip.NewReader(file)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
defer func() {
|
||||
_ = gzipReader.Close()
|
||||
}()
|
||||
|
||||
reader := tar.NewReader(gzipReader)
|
||||
var binaries []string
|
||||
for {
|
||||
header, err := reader.Next()
|
||||
if errors.Is(err, io.EOF) {
|
||||
break
|
||||
}
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
if header.Typeflag == tar.TypeReg && isLikelyBinary(header.Name, false, header.FileInfo().Mode()) {
|
||||
binaries = append(binaries, header.Name)
|
||||
}
|
||||
}
|
||||
|
||||
if len(binaries) == 1 {
|
||||
return binaries[0], nil
|
||||
}
|
||||
|
||||
for _, name := range binaries {
|
||||
if matchBinaryName(filepath.Base(name), candidates) {
|
||||
return name, nil
|
||||
}
|
||||
}
|
||||
|
||||
return "", errors.New(errNoCompatibleAsset)
|
||||
}
|
||||
|
||||
func findBinaryInZip(archivePath string, candidates []string) (string, error) {
|
||||
reader, err := zip.OpenReader(archivePath)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
defer func() {
|
||||
_ = reader.Close()
|
||||
}()
|
||||
|
||||
var binaries []string
|
||||
for _, file := range reader.File {
|
||||
if !file.FileInfo().IsDir() && isLikelyBinary(file.Name, false, file.FileInfo().Mode()) {
|
||||
binaries = append(binaries, file.Name)
|
||||
}
|
||||
}
|
||||
|
||||
if len(binaries) == 1 {
|
||||
return binaries[0], nil
|
||||
}
|
||||
|
||||
for _, name := range binaries {
|
||||
if matchBinaryName(filepath.Base(name), candidates) {
|
||||
return name, nil
|
||||
}
|
||||
}
|
||||
|
||||
return "", errors.New(errNoCompatibleAsset)
|
||||
}
|
||||
|
||||
func extractTarGz(ctx context.Context, archivePath, destination, targetName string, candidates []string) (string, error) {
|
||||
binaryPathInArchive, err := findBinaryInTarGz(archivePath, candidates)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
logger.InfoF(ctx, "[Updater] Extracting tar.gz archive: %s (extracting: %s)", archivePath, binaryPathInArchive)
|
||||
//nolint:gosec // updater archivePath is verified
|
||||
file, err := os.Open(archivePath)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
defer func() {
|
||||
_ = file.Close()
|
||||
}()
|
||||
gzipReader, err := gzip.NewReader(file)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
defer func() {
|
||||
_ = gzipReader.Close()
|
||||
}()
|
||||
|
||||
reader := tar.NewReader(gzipReader)
|
||||
for {
|
||||
header, err := reader.Next()
|
||||
if errors.Is(err, io.EOF) {
|
||||
break
|
||||
}
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
if header.Name != binaryPathInArchive {
|
||||
continue
|
||||
}
|
||||
target, err := safeArchivePath(destination, targetName)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
//nolint:gosec // updater destination is sanitized
|
||||
output, err := os.OpenFile(target, os.O_CREATE|os.O_EXCL|os.O_WRONLY, stagedBinaryMode)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
written, copyErr := io.Copy(output, io.LimitReader(reader, maxArchiveSize+1))
|
||||
closeErr := output.Close()
|
||||
if copyErr != nil {
|
||||
return "", copyErr
|
||||
}
|
||||
if closeErr != nil {
|
||||
return "", closeErr
|
||||
}
|
||||
if written > maxArchiveSize {
|
||||
return "", errors.New("解压后的程序文件超过大小限制")
|
||||
}
|
||||
logger.InfoF(ctx, "[Updater] Successfully extracted binary to %s", target)
|
||||
return target, nil
|
||||
}
|
||||
return "", errors.New(errNoCompatibleAsset)
|
||||
}
|
||||
|
||||
func extractZip(ctx context.Context, archivePath, destination, targetName string, candidates []string) (string, error) {
|
||||
binaryPathInArchive, err := findBinaryInZip(archivePath, candidates)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
logger.InfoF(ctx, "[Updater] Extracting zip archive: %s (extracting: %s)", archivePath, binaryPathInArchive)
|
||||
reader, err := zip.OpenReader(archivePath)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
defer func() {
|
||||
_ = reader.Close()
|
||||
}()
|
||||
for _, file := range reader.File {
|
||||
if file.Name != binaryPathInArchive {
|
||||
continue
|
||||
}
|
||||
target, err := safeArchivePath(destination, targetName)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
input, err := file.Open()
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
//nolint:gosec // updater extraction target is safe
|
||||
output, err := os.OpenFile(target, os.O_CREATE|os.O_EXCL|os.O_WRONLY, stagedBinaryMode)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
written, copyErr := io.Copy(output, io.LimitReader(input, maxArchiveSize+1))
|
||||
inputCloseErr := input.Close()
|
||||
outputCloseErr := output.Close()
|
||||
if copyErr != nil {
|
||||
return "", copyErr
|
||||
}
|
||||
if inputCloseErr != nil {
|
||||
return "", inputCloseErr
|
||||
}
|
||||
if outputCloseErr != nil {
|
||||
return "", outputCloseErr
|
||||
}
|
||||
if written > maxArchiveSize {
|
||||
return "", errors.New("解压后的程序文件超过大小限制")
|
||||
}
|
||||
logger.InfoF(ctx, "[Updater] Successfully extracted binary to %s", target)
|
||||
return target, nil
|
||||
}
|
||||
return "", errors.New(errNoCompatibleAsset)
|
||||
}
|
||||
|
||||
func (m *updaterManager) prepareUpgrade(ctx context.Context) (string, string, error) {
|
||||
if runtime.GOOS == windowsOS {
|
||||
return "", "", errors.New(errAutomaticUpgradeBlocked)
|
||||
}
|
||||
if normalizeVersion(buildinfo.Version) == "" {
|
||||
return "", "", errors.New(errDevelopmentBuild)
|
||||
}
|
||||
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
if m.upgrading {
|
||||
return "", "", errors.New(errUpgradeAlreadyRunning)
|
||||
}
|
||||
|
||||
status, asset, err := m.status(ctx)
|
||||
if err != nil {
|
||||
return "", "", err
|
||||
}
|
||||
if !status.UpdateAvailable {
|
||||
return "", "", errors.New(errAlreadyUpToDate)
|
||||
}
|
||||
|
||||
logger.InfoF(ctx, "[Updater] Preparing upgrade. current: %s, latest: %s", status.CurrentVersion, status.LatestVersion)
|
||||
|
||||
executable, err := os.Executable()
|
||||
if err != nil {
|
||||
return "", "", fmt.Errorf("定位当前程序失败: %w", err)
|
||||
}
|
||||
executable, err = filepath.EvalSymlinks(executable)
|
||||
if err != nil {
|
||||
return "", "", fmt.Errorf("解析当前程序路径失败: %w", err)
|
||||
}
|
||||
|
||||
tempDir, err := os.MkdirTemp(filepath.Dir(executable), ".wavelet-update-*")
|
||||
if err != nil {
|
||||
return "", "", fmt.Errorf("创建升级目录失败: %w", err)
|
||||
}
|
||||
|
||||
archivePath := filepath.Join(tempDir, asset.Name)
|
||||
if err := downloadArchive(ctx, m.client, asset, archivePath); err != nil {
|
||||
_ = os.RemoveAll(tempDir)
|
||||
return "", "", err
|
||||
}
|
||||
|
||||
targetName := filepath.Base(executable)
|
||||
candidates := getCandidateBinaryNames(executable, status.UpstreamRepository)
|
||||
|
||||
var stagedBinary string
|
||||
if strings.HasSuffix(asset.Name, ".zip") {
|
||||
stagedBinary, err = extractZip(ctx, archivePath, tempDir, targetName, candidates)
|
||||
} else {
|
||||
stagedBinary, err = extractTarGz(ctx, archivePath, tempDir, targetName, candidates)
|
||||
}
|
||||
if err != nil {
|
||||
_ = os.RemoveAll(tempDir)
|
||||
return "", "", fmt.Errorf("解压升级资产失败: %w", err)
|
||||
}
|
||||
logger.InfoF(ctx, "[Updater] Staged binary successfully prepared: %s", stagedBinary)
|
||||
m.upgrading = true
|
||||
return executable, stagedBinary, nil
|
||||
}
|
||||
|
||||
func (m *updaterManager) finishUpgrade() {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
m.upgrading = false
|
||||
}
|
||||
@@ -0,0 +1,615 @@
|
||||
// Copyright 2025 linux.do
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package admin
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"net/http"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"gorm.io/gorm"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/backend/core/contracts"
|
||||
"github.com/Rain-kl/Wavelet/backend/pkg/idgen"
|
||||
"github.com/Rain-kl/Wavelet/backend/pkg/logger"
|
||||
"github.com/Rain-kl/Wavelet/backend/pkg/response"
|
||||
"github.com/Rain-kl/Wavelet/backend/pkg/util"
|
||||
"github.com/Rain-kl/Wavelet/backend/plugins/domain/auth"
|
||||
db "github.com/Rain-kl/Wavelet/backend/plugins/infra/database"
|
||||
)
|
||||
|
||||
const minPasswordLength = 8
|
||||
|
||||
// listUsersRequest 用户列表查询请求
|
||||
type listUsersRequest struct {
|
||||
Page int `form:"page" binding:"min=1"`
|
||||
PageSize int `form:"page_size" binding:"min=1,max=100"`
|
||||
UserID *uint64 `form:"user_id" binding:"omitempty,gt=0"`
|
||||
Username string `form:"username"`
|
||||
Email string `form:"email"`
|
||||
}
|
||||
|
||||
type userResponse struct {
|
||||
ID uint64 `json:"id,string"`
|
||||
Username string `json:"username"`
|
||||
Nickname string `json:"nickname"`
|
||||
Email string `json:"email"`
|
||||
AvatarURL string `json:"avatar_url"`
|
||||
IsActive bool `json:"is_active"`
|
||||
IsAdmin bool `json:"is_admin"`
|
||||
Bio string `json:"bio"`
|
||||
Phone string `json:"phone"`
|
||||
Gender string `json:"gender"`
|
||||
Website string `json:"website"`
|
||||
Location string `json:"location"`
|
||||
LastLoginAt time.Time `json:"last_login_at"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
UpdatedAt time.Time `json:"updated_at"`
|
||||
}
|
||||
|
||||
// listUsersResponse 用户列表响应
|
||||
type listUsersResponse struct {
|
||||
Users []userResponse `json:"users"`
|
||||
Total int64 `json:"total"`
|
||||
}
|
||||
|
||||
func parseUserID(c *gin.Context) (uint64, bool) {
|
||||
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
|
||||
if err != nil || id == 0 {
|
||||
response.AbortBadRequest(c, userNotFound)
|
||||
return 0, false
|
||||
}
|
||||
return id, true
|
||||
}
|
||||
|
||||
func toUserResponse(u *contracts.UserDTO) userResponse {
|
||||
if u == nil {
|
||||
return userResponse{}
|
||||
}
|
||||
return userResponse{
|
||||
ID: u.ID,
|
||||
Username: u.Username,
|
||||
Nickname: u.Nickname,
|
||||
Email: u.Email,
|
||||
AvatarURL: u.AvatarURL,
|
||||
IsActive: u.IsActive,
|
||||
IsAdmin: u.IsAdmin,
|
||||
Bio: u.Bio,
|
||||
Phone: u.Phone,
|
||||
Gender: u.Gender,
|
||||
Website: u.Website,
|
||||
Location: u.Location,
|
||||
LastLoginAt: u.LastLoginAt,
|
||||
CreatedAt: u.CreatedAt,
|
||||
UpdatedAt: u.UpdatedAt,
|
||||
}
|
||||
}
|
||||
|
||||
func abortUserLogicError(c *gin.Context, err error, notFoundMsg string, forbiddenMsgs, badRequestMsgs []string) bool {
|
||||
if err == nil {
|
||||
return false
|
||||
}
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
response.AbortNotFound(c, notFoundMsg)
|
||||
return true
|
||||
}
|
||||
msg := err.Error()
|
||||
for _, m := range badRequestMsgs {
|
||||
if msg == m {
|
||||
response.AbortBadRequest(c, msg)
|
||||
return true
|
||||
}
|
||||
}
|
||||
for _, m := range forbiddenMsgs {
|
||||
if msg == m {
|
||||
response.AbortForbidden(c, msg)
|
||||
return true
|
||||
}
|
||||
}
|
||||
logger.ErrorF(c.Request.Context(), "Admin user error: %v", err)
|
||||
response.AbortInternal(c, "内部服务器错误")
|
||||
return true
|
||||
}
|
||||
|
||||
// ListUsers 获取用户列表
|
||||
// @Summary 获取用户列表
|
||||
// @Description 分页返回用户列表,支持按用户 ID 和用户名筛选,需要管理员权限
|
||||
// @Tags admin
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param request query listUsersRequest true "查询参数"
|
||||
// @Success 200 {object} response.Any{data=listUsersResponse} "用户列表"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 403 {object} response.Any "无管理员权限"
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/admin/users [get]
|
||||
func ListUsers(c *gin.Context) {
|
||||
var req listUsersRequest
|
||||
if err := c.ShouldBindQuery(&req); err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
total, dtos, err := listUsers(c.Request.Context(), req)
|
||||
if err != nil {
|
||||
logger.ErrorF(c.Request.Context(), "List admin users failed: %v", err)
|
||||
response.AbortInternal(c, "获取用户列表失败")
|
||||
return
|
||||
}
|
||||
|
||||
users := make([]userResponse, 0, len(dtos))
|
||||
for _, dto := range dtos {
|
||||
users = append(users, toUserResponse(dto))
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, response.OK(listUsersResponse{
|
||||
Users: users,
|
||||
Total: total,
|
||||
}))
|
||||
}
|
||||
|
||||
// GetUser 获取用户详情
|
||||
// @Summary 获取用户详情
|
||||
// @Description 返回指定用户的完整个人资料和系统状态,需要管理员权限,不返回密码等敏感字段
|
||||
// @Tags admin
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param id path int true "用户 ID"
|
||||
// @Success 200 {object} response.Any{data=userResponse} "用户详情"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 403 {object} response.Any "无管理员权限"
|
||||
// @Failure 404 {object} response.Any "用户不存在"
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/admin/users/{id} [get]
|
||||
func GetUser(c *gin.Context) {
|
||||
id, ok := parseUserID(c)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
|
||||
targetUser, err := getUserDetail(c.Request.Context(), id)
|
||||
if abortUserLogicError(c, err, userNotFound, nil, nil) {
|
||||
return
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, response.OK(toUserResponse(targetUser)))
|
||||
}
|
||||
|
||||
// updateUserStatusRequest 更新用户状态请求
|
||||
type updateUserStatusRequest struct {
|
||||
IsActive bool `json:"is_active"`
|
||||
}
|
||||
|
||||
// UpdateUserStatus 更新用户状态(启用/禁用)
|
||||
// @Summary 更新用户状态
|
||||
// @Description 启用或禁用指定用户,管理员账号无法被禁用,需要管理员权限
|
||||
// @Tags admin
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param id path int true "用户 ID"
|
||||
// @Param request body updateUserStatusRequest true "状态参数"
|
||||
// @Success 200 {object} response.Any{data=string} "更新成功"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 403 {object} response.Any "无管理员权限或尝试禁用管理员"
|
||||
// @Failure 404 {object} response.Any "用户不存在"
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/admin/users/{id}/status [put]
|
||||
func UpdateUserStatus(c *gin.Context) {
|
||||
var req updateUserStatusRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
id, ok := parseUserID(c)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
|
||||
if err := updateUserStatus(c.Request.Context(), id, req.IsActive); err != nil {
|
||||
if abortUserLogicError(c, err, userNotFound, []string{cannotDisable}, nil) {
|
||||
return
|
||||
}
|
||||
response.AbortInternal(c, updateUserFailed)
|
||||
return
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, response.OKNil())
|
||||
}
|
||||
|
||||
// DeleteUser 删除用户
|
||||
// @Summary 删除用户
|
||||
// @Description 删除指定非管理员用户,需要管理员权限,不能删除当前登录用户
|
||||
// @Tags admin
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param id path int true "用户 ID"
|
||||
// @Success 200 {object} response.Any{data=string} "删除成功"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 403 {object} response.Any "无管理员权限、尝试删除管理员或当前用户"
|
||||
// @Failure 404 {object} response.Any "用户不存在"
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/admin/users/{id} [delete]
|
||||
func DeleteUser(c *gin.Context) {
|
||||
id, ok := parseUserID(c)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
|
||||
currUser, _ := util.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey)
|
||||
if currUser == nil {
|
||||
response.AbortUnauthorized(c, AdminRequired)
|
||||
return
|
||||
}
|
||||
if err := deleteUser(c.Request.Context(), currUser.ID, id); err != nil {
|
||||
if abortUserLogicError(c, err, userNotFound, []string{cannotDelete, cannotDeleteSelf}, nil) {
|
||||
return
|
||||
}
|
||||
response.AbortInternal(c, deleteUserFailed)
|
||||
return
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, response.OKNil())
|
||||
}
|
||||
|
||||
// createUserRequest 创建用户请求
|
||||
type createUserRequest struct {
|
||||
Username string `json:"username" binding:"required,min=3,max=64"`
|
||||
Password string `json:"password" binding:"required,min=8,max=64"`
|
||||
Nickname string `json:"nickname" binding:"omitempty,max=64"`
|
||||
Email string `json:"email" binding:"required,email,max=255"`
|
||||
IsActive bool `json:"is_active"`
|
||||
IsAdmin bool `json:"is_admin"`
|
||||
}
|
||||
|
||||
// CreateUser 创建用户
|
||||
// @Summary 创建用户
|
||||
// @Description 创建一个本地密码登录的新用户,需要管理员权限
|
||||
// @Tags admin
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param request body createUserRequest true "创建用户参数"
|
||||
// @Success 200 {object} response.Any{data=userResponse} "创建成功"
|
||||
// @Failure 400 {object} response.Any "参数错误或用户名已存在"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 403 {object} response.Any "无管理员权限"
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/admin/users [post]
|
||||
func CreateUser(c *gin.Context) {
|
||||
var req createUserRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
newUser, err := createUser(c.Request.Context(), req)
|
||||
if abortUserLogicError(c, err, "", nil, []string{usernameRequired, emailRequired, passwordTooShort, usernameExists, emailExists}) {
|
||||
return
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, response.OK(toUserResponse(newUser)))
|
||||
}
|
||||
|
||||
// updateUserRequest 更新用户信息请求
|
||||
type updateUserRequest struct {
|
||||
Nickname string `json:"nickname" binding:"max=64"`
|
||||
Email string `json:"email" binding:"required,email,max=255"`
|
||||
IsAdmin bool `json:"is_admin"`
|
||||
Password string `json:"password" binding:"omitempty,min=8,max=64"`
|
||||
}
|
||||
|
||||
// UpdateUser 更新用户信息
|
||||
// @Summary 更新用户信息
|
||||
// @Description 更新指定用户的昵称、邮箱、管理员权限,并可选重置密码,需要管理员权限
|
||||
// @Tags admin
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Param id path int true "用户 ID"
|
||||
// @Param request body updateUserRequest true "更新参数"
|
||||
// @Success 200 {object} response.Any{data=string} "更新成功"
|
||||
// @Failure 400 {object} response.Any "参数错误"
|
||||
// @Failure 401 {object} response.Any "未登录"
|
||||
// @Failure 403 {object} response.Any "无管理员权限或尝试修改自身权限"
|
||||
// @Failure 404 {object} response.Any "用户不存在"
|
||||
// @Failure 500 {object} response.Any "内部错误"
|
||||
// @Router /api/v1/admin/users/{id} [put]
|
||||
func UpdateUser(c *gin.Context) {
|
||||
var req updateUserRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
response.AbortBadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
id, ok := parseUserID(c)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
|
||||
currUser, _ := util.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey)
|
||||
if currUser == nil {
|
||||
response.AbortUnauthorized(c, AdminRequired)
|
||||
return
|
||||
}
|
||||
err := updateUser(c.Request.Context(), currUser.ID, updateUserParam{
|
||||
ID: id,
|
||||
Nickname: req.Nickname,
|
||||
Email: req.Email,
|
||||
IsAdmin: req.IsAdmin,
|
||||
Password: req.Password,
|
||||
})
|
||||
|
||||
if err != nil {
|
||||
if abortUserLogicError(c, err, userNotFound, []string{cannotRevokeSelfAdmin}, []string{emailRequired, emailExists, passwordTooShort}) {
|
||||
return
|
||||
}
|
||||
response.AbortInternal(c, updateUserInfoFailed)
|
||||
return
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, response.OKNil())
|
||||
}
|
||||
|
||||
func listUsers(ctx context.Context, req listUsersRequest) (int64, []*contracts.UserDTO, error) {
|
||||
query := db.DB(ctx).Table("w_users")
|
||||
if req.UserID != nil {
|
||||
query = query.Where("id = ?", *req.UserID)
|
||||
}
|
||||
if req.Username != "" {
|
||||
query = query.Where("username LIKE ? ESCAPE '\\'", util.EscapeLike(req.Username)+"%")
|
||||
}
|
||||
if req.Email != "" {
|
||||
query = query.Where("email LIKE ? ESCAPE '\\'", util.EscapeLike(req.Email)+"%")
|
||||
}
|
||||
|
||||
var total int64
|
||||
if err := query.Count(&total).Error; err != nil {
|
||||
return 0, nil, err
|
||||
}
|
||||
|
||||
var users []*contracts.UserDTO
|
||||
offset := (req.Page - 1) * req.PageSize
|
||||
if err := query.
|
||||
Select("id, username, nickname, email, avatar_url, is_active, is_admin, last_login_at, created_at, updated_at").
|
||||
Order("id ASC").
|
||||
Offset(offset).
|
||||
Limit(req.PageSize).
|
||||
Find(&users).Error; err != nil {
|
||||
return 0, nil, err
|
||||
}
|
||||
return total, users, nil
|
||||
}
|
||||
|
||||
func getUserDetail(ctx context.Context, id uint64) (*contracts.UserDTO, error) {
|
||||
var user contracts.UserDTO
|
||||
if err := db.DB(ctx).Table("w_users").
|
||||
Select("id, username, nickname, email, avatar_url, is_active, is_admin, bio, phone, gender, website, location, last_login_at, created_at, updated_at").
|
||||
Where("id = ?", id).
|
||||
First(&user).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &user, nil
|
||||
}
|
||||
|
||||
func updateUserStatus(ctx context.Context, id uint64, active bool) error {
|
||||
var flags struct {
|
||||
ID uint64
|
||||
IsAdmin bool
|
||||
}
|
||||
if err := db.DB(ctx).Table("w_users").Select("id, is_admin").Where("id = ?", id).First(&flags).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if !active && flags.IsAdmin {
|
||||
return errors.New(cannotDisable)
|
||||
}
|
||||
|
||||
var tokenHashes []string
|
||||
if !active {
|
||||
_ = db.DB(ctx).Table("w_access_tokens").Where("user_id = ?", id).Pluck("token_hash", &tokenHashes).Error
|
||||
}
|
||||
|
||||
err := db.DB(ctx).Table("w_users").Where("id = ?", id).Update("is_active", active).Error
|
||||
if err == nil {
|
||||
auth.InvalidateCachedUser(ctx, id)
|
||||
if !active {
|
||||
for _, hash := range tokenHashes {
|
||||
auth.InvalidateCachedToken(ctx, hash)
|
||||
}
|
||||
}
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
func deleteUser(ctx context.Context, currentUserID, targetID uint64) error {
|
||||
if currentUserID == targetID {
|
||||
return errors.New(cannotDeleteSelf)
|
||||
}
|
||||
var flags struct {
|
||||
ID uint64
|
||||
IsAdmin bool
|
||||
}
|
||||
if err := db.DB(ctx).Table("w_users").Select("id, is_admin").Where("id = ?", targetID).First(&flags).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if flags.IsAdmin {
|
||||
return errors.New(cannotDelete)
|
||||
}
|
||||
|
||||
var tokenHashes []string
|
||||
_ = db.DB(ctx).Table("w_access_tokens").Where("user_id = ?", targetID).Pluck("token_hash", &tokenHashes).Error
|
||||
|
||||
err := db.DB(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
if err := tx.Table("w_access_tokens").Where("user_id = ?", targetID).Delete(map[string]any{}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if err := tx.Table("w_external_accounts").Where("user_id = ?", targetID).Delete(map[string]any{}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
return tx.Table("w_users").Where("id = ?", targetID).Delete(map[string]any{}).Error
|
||||
})
|
||||
if err == nil {
|
||||
auth.InvalidateCachedUser(ctx, targetID)
|
||||
for _, hash := range tokenHashes {
|
||||
auth.InvalidateCachedToken(ctx, hash)
|
||||
}
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
func createUser(ctx context.Context, req createUserRequest) (*contracts.UserDTO, error) {
|
||||
req.Username = strings.TrimSpace(req.Username)
|
||||
req.Nickname = strings.TrimSpace(req.Nickname)
|
||||
req.Password = strings.TrimSpace(req.Password)
|
||||
req.Email = strings.TrimSpace(req.Email)
|
||||
|
||||
if req.Username == "" {
|
||||
return nil, errors.New(usernameRequired)
|
||||
}
|
||||
if req.Email == "" {
|
||||
return nil, errors.New(emailRequired)
|
||||
}
|
||||
if len(req.Password) < minPasswordLength {
|
||||
return nil, errors.New(passwordTooShort)
|
||||
}
|
||||
|
||||
var count int64
|
||||
if err := db.DB(ctx).Table("w_users").Where("username = ?", req.Username).Count(&count).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if count > 0 {
|
||||
return nil, errors.New(usernameExists)
|
||||
}
|
||||
|
||||
var emailCount int64
|
||||
if err := db.DB(ctx).Table("w_users").Where("email = ?", req.Email).Count(&emailCount).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if emailCount > 0 {
|
||||
return nil, errors.New(emailExists)
|
||||
}
|
||||
|
||||
hash, err := util.HashPassword(req.Password)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if req.Nickname == "" {
|
||||
req.Nickname = req.Username
|
||||
}
|
||||
|
||||
now := time.Now()
|
||||
newUser := contracts.UserDTO{
|
||||
ID: idgen.NextUint64ID(),
|
||||
Username: req.Username,
|
||||
Nickname: req.Nickname,
|
||||
Email: req.Email,
|
||||
IsActive: req.IsActive,
|
||||
IsAdmin: req.IsAdmin,
|
||||
CreatedAt: now,
|
||||
UpdatedAt: now,
|
||||
}
|
||||
|
||||
row := map[string]any{
|
||||
"id": newUser.ID,
|
||||
"username": newUser.Username,
|
||||
"password": hash,
|
||||
"nickname": newUser.Nickname,
|
||||
"email": newUser.Email,
|
||||
"is_active": newUser.IsActive,
|
||||
"is_admin": newUser.IsAdmin,
|
||||
"created_at": now,
|
||||
"updated_at": now,
|
||||
}
|
||||
if err := db.DB(ctx).Table("w_users").Create(row).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &newUser, nil
|
||||
}
|
||||
|
||||
type updateUserParam struct {
|
||||
ID uint64
|
||||
Nickname string
|
||||
Email string
|
||||
IsAdmin bool
|
||||
Password string
|
||||
}
|
||||
|
||||
func updateUser(ctx context.Context, currentUserID uint64, param updateUserParam) error {
|
||||
param.Nickname = strings.TrimSpace(param.Nickname)
|
||||
param.Email = strings.TrimSpace(param.Email)
|
||||
param.Password = strings.TrimSpace(param.Password)
|
||||
|
||||
if param.Email == "" {
|
||||
return errors.New(emailRequired)
|
||||
}
|
||||
|
||||
var targetUser contracts.UserDTO
|
||||
if err := db.DB(ctx).Table("w_users").Where("id = ?", param.ID).First(&targetUser).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if currentUserID == param.ID && !param.IsAdmin && targetUser.IsAdmin {
|
||||
return errors.New(cannotRevokeSelfAdmin)
|
||||
}
|
||||
|
||||
if targetUser.Email != param.Email {
|
||||
var count int64
|
||||
if err := db.DB(ctx).Table("w_users").Where("email = ? AND id != ?", param.Email, param.ID).Count(&count).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if count > 0 {
|
||||
return errors.New(emailExists)
|
||||
}
|
||||
}
|
||||
|
||||
if param.Password != "" && len(param.Password) < minPasswordLength {
|
||||
return errors.New(passwordTooShort)
|
||||
}
|
||||
|
||||
needRevokeTokens := (param.Password != "") || (targetUser.IsAdmin && !param.IsAdmin)
|
||||
var tokenHashes []string
|
||||
if needRevokeTokens {
|
||||
_ = db.DB(ctx).Table("w_access_tokens").Where("user_id = ?", param.ID).Pluck("token_hash", &tokenHashes).Error
|
||||
}
|
||||
|
||||
if param.Nickname == "" {
|
||||
param.Nickname = targetUser.Username
|
||||
}
|
||||
|
||||
updates := map[string]any{
|
||||
"nickname": param.Nickname,
|
||||
"email": param.Email,
|
||||
"is_admin": param.IsAdmin,
|
||||
"updated_at": time.Now(),
|
||||
}
|
||||
if param.Password != "" {
|
||||
hash, err := util.HashPassword(param.Password)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
updates["password"] = hash
|
||||
}
|
||||
|
||||
err := db.DB(ctx).Table("w_users").Where("id = ?", param.ID).Updates(updates).Error
|
||||
if err == nil {
|
||||
auth.InvalidateCachedUser(ctx, param.ID)
|
||||
if needRevokeTokens {
|
||||
for _, hash := range tokenHashes {
|
||||
auth.InvalidateCachedToken(ctx, hash)
|
||||
}
|
||||
}
|
||||
}
|
||||
return err
|
||||
}
|
||||
@@ -0,0 +1,44 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package admin
|
||||
|
||||
import (
|
||||
"github.com/Rain-kl/Wavelet/backend/core/contracts"
|
||||
"github.com/Rain-kl/Wavelet/backend/pkg/logger"
|
||||
"github.com/Rain-kl/Wavelet/backend/pkg/response"
|
||||
"github.com/Rain-kl/Wavelet/backend/pkg/trace"
|
||||
"github.com/Rain-kl/Wavelet/backend/pkg/util"
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
// LoginAdminRequired 返回管理员权限校验中间件
|
||||
func LoginAdminRequired() gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
ctx, span := trace.Start(c.Request.Context(), "LoginAdminRequired")
|
||||
defer span.End()
|
||||
|
||||
user, _ := util.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey)
|
||||
if user == nil {
|
||||
response.AbortNotFound(c, AdminRequired)
|
||||
return
|
||||
}
|
||||
|
||||
// 如果是通过 Access Token 鉴权,需要检查令牌本身是否具有管理员权限
|
||||
if tokenAuth, _ := util.GetFromContext[bool](c, contracts.AuthTokenAuthKey); tokenAuth {
|
||||
tokenAdmin, _ := util.GetFromContext[bool](c, contracts.AuthTokenAdminKey)
|
||||
if !tokenAdmin {
|
||||
response.AbortNotFound(c, TokenAdminRequired)
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
if !user.IsAdmin {
|
||||
response.AbortNotFound(c, AdminRequired)
|
||||
return
|
||||
}
|
||||
|
||||
logger.InfoF(ctx, "[LoginAdminRequired] %d %s", user.ID, user.Username)
|
||||
c.Next()
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,87 @@
|
||||
-- +goose Up
|
||||
-- +goose StatementBegin
|
||||
CREATE TABLE IF NOT EXISTS w_system_configs (
|
||||
key VARCHAR(64) PRIMARY KEY,
|
||||
value TEXT NOT NULL,
|
||||
type VARCHAR(32) NOT NULL DEFAULT 'system',
|
||||
visibility INTEGER NOT NULL DEFAULT 0,
|
||||
description VARCHAR(255),
|
||||
updated_at TIMESTAMPTZ DEFAULT CURRENT_TIMESTAMP,
|
||||
created_at TIMESTAMPTZ DEFAULT CURRENT_TIMESTAMP
|
||||
);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS w_templates (
|
||||
id BIGINT PRIMARY KEY,
|
||||
key VARCHAR(80) NOT NULL UNIQUE,
|
||||
name VARCHAR(100) NOT NULL,
|
||||
type VARCHAR(20) NOT NULL DEFAULT 'email',
|
||||
subject VARCHAR(255),
|
||||
content TEXT NOT NULL,
|
||||
description VARCHAR(255),
|
||||
is_system BOOLEAN NOT NULL DEFAULT FALSE,
|
||||
created_at TIMESTAMPTZ DEFAULT CURRENT_TIMESTAMP,
|
||||
updated_at TIMESTAMPTZ DEFAULT CURRENT_TIMESTAMP
|
||||
);
|
||||
CREATE INDEX IF NOT EXISTS idx_w_templates_is_system ON w_templates (is_system);
|
||||
CREATE INDEX IF NOT EXISTS idx_w_templates_created_at ON w_templates (created_at);
|
||||
CREATE INDEX IF NOT EXISTS idx_w_templates_updated_at ON w_templates (updated_at);
|
||||
|
||||
-- Seed system configs (all default platform configs)
|
||||
INSERT INTO w_system_configs (key, value, type, visibility, description, created_at, updated_at) VALUES
|
||||
('cap_login_enabled', 'false', 'system', 1, '是否启用登录人机验证(true/false)', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
|
||||
('cap_auto_solve', 'true', 'system', 1, '打开页面后是否自动开始计算,关闭则需用户手动点击触发', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
|
||||
('cap_challenge_count', '1', 'system', 0, '客户端需求解的 PoW 难题总数,默认 1,推荐 1~5', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
|
||||
('cap_challenge_size', '32', 'system', 0, '人机验证盐值长度', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
|
||||
('cap_challenge_difficulty', '4', 'system', 0, '人机验证 PoW 难度(目标前缀长度)', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
|
||||
('cap_challenge_ttl_seconds', '600', 'system', 0, '人机验证难题有效时间(秒)', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
|
||||
('cap_token_ttl_seconds', '1200', 'system', 0, '人机验证兑换凭证有效时间(秒)', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
|
||||
('server_address', '', 'system', 0, '服务器地址(用于跨域源控制,不设定则允许任意源)', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
|
||||
('smtp_host', '', 'system', 0, 'SMTP 服务器地址(例如 smtp.example.com)', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
|
||||
('smtp_port', '587', 'system', 0, 'SMTP 端口(例如 587 或 465)', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
|
||||
('smtp_username', '', 'system', 0, 'SMTP 账户(如 sender@example.com)', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
|
||||
('smtp_password', '', 'system', 0, 'SMTP 访问凭证(授权码/密码)', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
|
||||
('upload_allowed_extensions', 'jpg,png,webp', 'system', 1, '允许上传的图片扩展名(逗号分隔)', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
|
||||
('site_name', 'Wavelet', 'system', 1, '系统平台的展示名称', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
|
||||
('password_login_enabled', 'true', 'system', 1, '是否允许使用账号密码登录', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
|
||||
('registration_enabled', 'true', 'system', 1, '控制普通用户是否可以自主注册(true/false)', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
|
||||
('password_register_enabled', 'true', 'system', 1, '是否允许通过密码创建本地账号', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
|
||||
('oidc_login_enabled', 'true', 'system', 1, '是否允许使用第三方 OIDC 认证源登录', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
|
||||
('max_api_keys_per_user', '5', 'business', 1, '限制每个普通用户可以创建的 API Key 最大数量', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
|
||||
('email_login_verification_enabled', 'false', 'system', 1, '是否开启邮箱登录验证(true/false)', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
|
||||
('email_register_verification_enabled', 'false', 'system', 1, '是否开启邮箱注册验证(true/false)', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
|
||||
('menu_display_config', '{}', 'system', 1, '目录显示配置(JSON 字符串,格式为 {url: enabled})', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
|
||||
('search_engine_indexing_enabled', 'false', 'system', 1, '是否允许搜索引擎爬取/检索该站点(true/false)', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
|
||||
('update_upstream_repository', 'Rain-kl/Wavelet', 'system', 0, 'GitHub Actions Release 上游仓库(owner/repo 或 GitHub 仓库地址)', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
|
||||
('storage_config', '{"driver":"local","local":{"root":"."},"s3":{"region":"us-east-1"},"r2":{"region":"auto"},"minio":{"region":"us-east-1","path_style":true},"oss":{},"webdav":{}}', 'system', 0, '文件存储驱动及连接配置(JSON)', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
|
||||
('disk_cache_max_size_mb', '1024', 'system', 0, '磁盘缓存最大空间大小(MB)', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
|
||||
('disk_cache_ttl_minutes', '1440', 'system', 0, '磁盘缓存默认有效期(分钟)', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
|
||||
('disk_cache_lru_enabled', 'true', 'system', 0, '是否启用 LRU 淘汰机制', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
|
||||
('file_access_whitelist', '["avatar"]', 'system', 0, '免登录访问的文件业务类型白名单 (JSON 数组)', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
|
||||
('login_session_ttl_hours', '168', 'system', 0, '登录会话过期时间(小时)', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
|
||||
('log_database', '', 'system', 0, '当前日志主库(postgres/sqlite/clickhouse),由切换任务写入', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
|
||||
('log_db_migration', '', 'system', 0, '日志库迁移冻结标记(空或 migrating)', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP)
|
||||
ON CONFLICT (key) DO NOTHING;
|
||||
|
||||
INSERT INTO w_templates (id, key, name, type, subject, content, description, is_system, created_at, updated_at) VALUES
|
||||
(1, 'login_email', '登录验证码邮件', 'email', 'Wavelet 登录验证码', '<h3>Wavelet 登录验证</h3><p>您的登录验证码为:<strong>{{.Code}}</strong>,5分钟内有效,请勿将验证码泄露给他人。</p>', '用户密码登录时发送的验证码邮件模板,支持变量:{{.Code}}', TRUE, CURRENT_TIMESTAMP, CURRENT_TIMESTAMP),
|
||||
(2, 'register_email', '注册验证码邮件', 'email', 'Wavelet 注册验证码', '<h3>Wavelet 注册验证</h3><p>您的注册验证码为:<strong>{{.Code}}</strong>,5分钟内有效,请勿泄露给他人。</p>', '用户注册时发送的验证码邮件模板,支持变量:{{.Code}}', TRUE, CURRENT_TIMESTAMP, CURRENT_TIMESTAMP)
|
||||
ON CONFLICT (key) DO NOTHING;
|
||||
-- +goose StatementEnd
|
||||
|
||||
-- +goose Down
|
||||
-- +goose StatementBegin
|
||||
DELETE FROM w_templates WHERE key IN ('login_email', 'register_email');
|
||||
DELETE FROM w_system_configs WHERE key IN (
|
||||
'cap_login_enabled', 'cap_auto_solve', 'cap_challenge_count', 'cap_challenge_size',
|
||||
'cap_challenge_difficulty', 'cap_challenge_ttl_seconds', 'cap_token_ttl_seconds',
|
||||
'server_address', 'smtp_host', 'smtp_port', 'smtp_username', 'smtp_password',
|
||||
'upload_allowed_extensions', 'site_name', 'password_login_enabled', 'registration_enabled',
|
||||
'password_register_enabled', 'oidc_login_enabled', 'max_api_keys_per_user',
|
||||
'email_login_verification_enabled', 'email_register_verification_enabled',
|
||||
'menu_display_config', 'search_engine_indexing_enabled', 'update_upstream_repository',
|
||||
'storage_config', 'disk_cache_max_size_mb', 'disk_cache_ttl_minutes', 'disk_cache_lru_enabled',
|
||||
'file_access_whitelist', 'login_session_ttl_hours', 'log_database', 'log_db_migration'
|
||||
);
|
||||
DROP TABLE IF EXISTS w_templates;
|
||||
DROP TABLE IF EXISTS w_system_configs;
|
||||
-- +goose StatementEnd
|
||||
@@ -0,0 +1,222 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package admin
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"errors"
|
||||
"strings"
|
||||
"text/template"
|
||||
"time"
|
||||
)
|
||||
|
||||
// 配置键常量 - 所有系统配置的 key 定义
|
||||
const (
|
||||
ConfigKeyUploadAllowedExtensions = "upload_allowed_extensions" // 允许上传的文件扩展名,逗号分隔
|
||||
ConfigKeySiteName = "site_name" // 站点名称
|
||||
ConfigKeyPasswordLoginEnabled = "password_login_enabled" // 是否允许密码登录
|
||||
ConfigKeyRegistrationEnabled = "registration_enabled" // 是否允许注册
|
||||
ConfigKeyPasswordRegisterEnabled = "password_register_enabled" // 是否允许密码注册
|
||||
ConfigKeyOIDCLoginEnabled = "oidc_login_enabled" // 是否允许 OIDC 登录
|
||||
ConfigKeyMaxAPIKeysPerUser = "max_api_keys_per_user" //nolint:gosec // false positive: config key name, not credentials
|
||||
ConfigKeyCapLoginEnabled = "cap_login_enabled" // 是否启用登录人机验证
|
||||
ConfigKeyCapAutoSolve = "cap_auto_solve" // 打开页面后是否自动开始计算(false 则需用户手动点击)
|
||||
ConfigKeyCapChallengeCount = "cap_challenge_count" // 客户端需求解的 PoW 难题总数,默认 1,推荐 1~5
|
||||
ConfigKeyCapChallengeSize = "cap_challenge_size" // 人机验证盐值长度
|
||||
ConfigKeyCapChallengeDifficulty = "cap_challenge_difficulty" // 人机验证 PoW 难度(目标前缀长度)
|
||||
ConfigKeyCapChallengeTTL = "cap_challenge_ttl_seconds" // 人机验证难题有效时间(秒)
|
||||
ConfigKeyCapTokenTTL = "cap_token_ttl_seconds" //nolint:gosec // false positive: config key name, not credentials
|
||||
ConfigKeyServerAddress = "server_address" // 服务器地址
|
||||
ConfigKeySMTPHost = "smtp_host" // SMTP 服务器地址
|
||||
ConfigKeySMTPPort = "smtp_port" // SMTP 端口
|
||||
ConfigKeySMTPUsername = "smtp_username" // SMTP 账户
|
||||
ConfigKeySMTPPassword = "smtp_password" // SMTP 访问凭证
|
||||
ConfigKeyEmailLoginVerificationEnabled = "email_login_verification_enabled" // 是否启用邮箱登录验证
|
||||
ConfigKeyEmailRegisterVerificationEnabled = "email_register_verification_enabled" // 是否启用邮箱注册验证
|
||||
ConfigKeyMenuDisplayConfig = "menu_display_config" // 目录显示配置 (JSON 字符串)
|
||||
ConfigKeySearchEngineIndexingEnabled = "search_engine_indexing_enabled" // 是否允许搜索引擎检索
|
||||
ConfigKeyFileAccessWhitelist = "file_access_whitelist" // 免登录访问的文件业务类型白名单 (JSON 数组格式)
|
||||
ConfigKeyDiskCacheMaxSizeMB = "disk_cache_max_size_mb" // 磁盘缓存最大空间大小 (MB)
|
||||
ConfigKeyDiskCacheTTLMinutes = "disk_cache_ttl_minutes" // 磁盘缓存默认有效期 (分钟)
|
||||
ConfigKeyDiskCacheLRUEnabled = "disk_cache_lru_enabled" // 是否启用 LRU 淘汰机制
|
||||
ConfigKeyLoginSessionTTLHours = "login_session_ttl_hours" // 登录会话过期时间 (小时)
|
||||
ConfigKeyUpdateUpstreamRepository = "update_upstream_repository" // GitHub Actions Release 上游仓库
|
||||
ConfigKeyStorageConfig = "storage_config" // 文件存储配置 (JSON)
|
||||
ConfigKeyLogDatabase = "log_database" // 当前日志主库(postgres/sqlite/clickhouse),受保护
|
||||
ConfigKeyLogDBMigration = "log_db_migration" // 日志库迁移冻结标记(空/migrating),受保护
|
||||
ConfigKeyLogRetentionDaysPostgres = "log_retention_days_postgres" // PostgreSQL 用户访问日志保留天数
|
||||
ConfigKeyLogRetentionDaysSQLite = "log_retention_days_sqlite" // SQLite 用户访问日志保留天数
|
||||
ConfigKeyLogRetentionDaysClickHouse = "log_retention_days_clickhouse" // ClickHouse 用户访问日志保留天数
|
||||
)
|
||||
|
||||
const (
|
||||
// ConfigVisibilityHidden 表示配置不通过公共配置接口暴露
|
||||
ConfigVisibilityHidden = 0
|
||||
// ConfigVisibilityVisible 表示配置通过公共配置接口暴露
|
||||
ConfigVisibilityVisible = 1
|
||||
)
|
||||
|
||||
// SystemConfig 系统配置实体
|
||||
type SystemConfig struct {
|
||||
Key string `json:"key" gorm:"primaryKey;size:64;not null"`
|
||||
Value string `json:"value" gorm:"type:text;not null"`
|
||||
Type string `json:"type" gorm:"size:32;not null;default:'system'"`
|
||||
Visibility int `json:"visibility" gorm:"not null;default:0"`
|
||||
Description string `json:"description" gorm:"size:255"`
|
||||
UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime"`
|
||||
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime"`
|
||||
}
|
||||
|
||||
// TableName 表名
|
||||
func (SystemConfig) TableName() string {
|
||||
return "w_system_configs"
|
||||
}
|
||||
|
||||
// Template 邮件/消息模板实体
|
||||
type Template struct {
|
||||
ID uint64 `json:"id" gorm:"primaryKey;autoIncrement"`
|
||||
Key string `json:"key" gorm:"uniqueIndex;size:80;not null"`
|
||||
Name string `json:"name" gorm:"size:100;not null"`
|
||||
Type string `json:"type" gorm:"size:20;not null;default:'email'"`
|
||||
Subject string `json:"subject" gorm:"size:255"`
|
||||
Content string `json:"content" gorm:"type:text;not null"`
|
||||
Description string `json:"description" gorm:"size:255"`
|
||||
IsSystem bool `json:"is_system" gorm:"index;not null;default:false"`
|
||||
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime;index"`
|
||||
UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime;index"`
|
||||
}
|
||||
|
||||
// TableName 表名
|
||||
func (Template) TableName() string {
|
||||
return "w_templates"
|
||||
}
|
||||
|
||||
// TemplateTypeEmail 邮件模板类型
|
||||
const TemplateTypeEmail = "email"
|
||||
|
||||
// Normalize 规范化模板字段
|
||||
func (t *Template) Normalize() {
|
||||
t.Key = strings.TrimSpace(t.Key)
|
||||
t.Name = strings.TrimSpace(t.Name)
|
||||
t.Type = strings.ToLower(strings.TrimSpace(t.Type))
|
||||
t.Subject = strings.TrimSpace(t.Subject)
|
||||
t.Content = strings.TrimSpace(t.Content)
|
||||
t.Description = strings.TrimSpace(t.Description)
|
||||
if t.Type == "" {
|
||||
t.Type = TemplateTypeEmail
|
||||
}
|
||||
}
|
||||
|
||||
// Validate 校验模板必填字段
|
||||
func (t *Template) Validate() error {
|
||||
t.Normalize()
|
||||
if t.Key == "" {
|
||||
return errors.New(TemplateKeyRequired)
|
||||
}
|
||||
if t.Name == "" {
|
||||
return errors.New(TemplateNameRequired)
|
||||
}
|
||||
if t.Content == "" {
|
||||
return errors.New(TemplateContentRequired)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// Render 渲染模板的 Subject 和 Content
|
||||
func (t *Template) Render(data any) (string, string, error) {
|
||||
var subject string
|
||||
if t.Subject != "" {
|
||||
tmplSubject, err := template.New(t.Key + "_subject").Parse(t.Subject)
|
||||
if err != nil {
|
||||
return "", "", err
|
||||
}
|
||||
var subBuf bytes.Buffer
|
||||
if err := tmplSubject.Execute(&subBuf, data); err != nil {
|
||||
return "", "", err
|
||||
}
|
||||
subject = subBuf.String()
|
||||
}
|
||||
|
||||
tmplContent, err := template.New(t.Key + "_content").Parse(t.Content)
|
||||
if err != nil {
|
||||
return "", "", err
|
||||
}
|
||||
var bodyBuf bytes.Buffer
|
||||
if err := tmplContent.Execute(&bodyBuf, data); err != nil {
|
||||
return "", "", err
|
||||
}
|
||||
|
||||
return subject, bodyBuf.String(), nil
|
||||
}
|
||||
|
||||
// Schedule 定时任务配置表
|
||||
type Schedule struct {
|
||||
ID uint64 `json:"id,string" gorm:"primaryKey"`
|
||||
Name string `json:"name" gorm:"size:128;not null"`
|
||||
TaskType string `json:"task_type" gorm:"size:64;not null"`
|
||||
Cron string `json:"cron" gorm:"size:64;not null"`
|
||||
Payload string `json:"payload" gorm:"type:text"`
|
||||
IsActive bool `json:"is_active" gorm:"not null;default:true"`
|
||||
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime;index"`
|
||||
UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime"`
|
||||
}
|
||||
|
||||
// TableName 表名
|
||||
func (Schedule) TableName() string {
|
||||
return "w_schedules"
|
||||
}
|
||||
|
||||
// TaskExecutionStatus 任务执行状态
|
||||
type TaskExecutionStatus string
|
||||
|
||||
// 任务执行状态
|
||||
const (
|
||||
TaskExecutionStatusPending TaskExecutionStatus = "pending"
|
||||
TaskExecutionStatusRunning TaskExecutionStatus = "running"
|
||||
TaskExecutionStatusSucceeded TaskExecutionStatus = "succeeded"
|
||||
TaskExecutionStatusFailed TaskExecutionStatus = "failed"
|
||||
)
|
||||
|
||||
// TaskExecution 任务执行记录
|
||||
type TaskExecution struct {
|
||||
ID uint64 `json:"id,string" gorm:"primaryKey"`
|
||||
TaskID string `json:"task_id" gorm:"size:128;uniqueIndex;not null"`
|
||||
TaskType string `json:"task_type" gorm:"size:64;index;not null"`
|
||||
TaskName string `json:"task_name" gorm:"size:128"`
|
||||
Status TaskExecutionStatus `json:"status" gorm:"size:32;index;not null"`
|
||||
Retryable bool `json:"retryable" gorm:"not null;default:false"`
|
||||
MaxRetry int `json:"max_retry" gorm:"not null;default:0"`
|
||||
RetryCount int `json:"retry_count" gorm:"not null;default:0"`
|
||||
Log string `json:"log" gorm:"type:text"`
|
||||
ErrorMessage string `json:"error_message" gorm:"type:text"`
|
||||
Result string `json:"result" gorm:"type:text"`
|
||||
StartedAt *time.Time `json:"started_at" gorm:"index"`
|
||||
FinishedAt *time.Time `json:"finished_at"`
|
||||
Duration int64 `json:"duration" gorm:"comment:耗时毫秒"`
|
||||
Payload string `json:"payload" gorm:"type:text"`
|
||||
TriggeredBy string `json:"triggered_by" gorm:"size:32;not null;default:system"`
|
||||
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime;index"`
|
||||
UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime"`
|
||||
}
|
||||
|
||||
// TableName 表名
|
||||
func (TaskExecution) TableName() string {
|
||||
return "w_task_executions"
|
||||
}
|
||||
|
||||
// ListTaskExecutionsRequest 分页查询任务执行记录请求参数
|
||||
type ListTaskExecutionsRequest struct {
|
||||
Page int `form:"page"`
|
||||
PageSize int `form:"page_size"`
|
||||
Status string `form:"status"`
|
||||
TaskType string `form:"task_type"`
|
||||
TaskTypes string `form:"task_types"`
|
||||
TaskTypePrefix string `form:"task_type_prefix"`
|
||||
}
|
||||
|
||||
// TaskExecutionCleanupStats 任务日志清理结果统计
|
||||
type TaskExecutionCleanupStats struct {
|
||||
HighFrequencyDeleted int64 `json:"high_frequency_deleted"`
|
||||
LowFrequencyDeleted int64 `json:"low_frequency_deleted"`
|
||||
}
|
||||
@@ -0,0 +1,202 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package admin provides the system management console, diagnostics, audit logging, and configuration hot-reloading domain plugin for Cordis.
|
||||
package admin
|
||||
|
||||
import (
|
||||
"context"
|
||||
"embed"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/backend/core"
|
||||
"github.com/Rain-kl/Wavelet/backend/core/contracts"
|
||||
"github.com/Rain-kl/Wavelet/backend/core/extpoints"
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/hibiken/asynq"
|
||||
)
|
||||
|
||||
//go:embed migrations/*.sql
|
||||
var adminMigrations embed.FS
|
||||
|
||||
// Option configures the admin plugin.
|
||||
type Option func(*Plugin)
|
||||
|
||||
// Plugin implements core.Plugin to provide system administration and management APIs.
|
||||
type Plugin struct{}
|
||||
|
||||
// New creates a new admin domain plugin.
|
||||
func New(opts ...Option) *Plugin {
|
||||
p := &Plugin{}
|
||||
for _, opt := range opts {
|
||||
if opt != nil {
|
||||
opt(p)
|
||||
}
|
||||
}
|
||||
return p
|
||||
}
|
||||
|
||||
// Name returns the unique identifier for the admin domain plugin.
|
||||
func (p *Plugin) Name() string {
|
||||
return "admin"
|
||||
}
|
||||
|
||||
// Manifest returns the plugin metadata.
|
||||
func (p *Plugin) Manifest() core.Manifest {
|
||||
return core.Manifest{
|
||||
Name: "admin",
|
||||
Version: "1.0.0",
|
||||
Description: "System administration console, diagnostic monitoring, and configuration hot-reload plugin",
|
||||
Author: "Wavelet Team",
|
||||
}
|
||||
}
|
||||
|
||||
// Apply registers admin routes, tasks, schedules, and settings into the Context.
|
||||
func (p *Plugin) Apply(ctx *core.Context) error {
|
||||
// 0. Resolve auth service for middleware (via IoC, not direct import)
|
||||
var loginMW gin.HandlerFunc = func(c *gin.Context) { c.Next() }
|
||||
var adminMW gin.HandlerFunc = func(c *gin.Context) { c.Next() }
|
||||
if authSvc, err := core.Inject[contracts.AuthService](ctx); err == nil && authSvc != nil {
|
||||
if mw, ok := authSvc.RequireAuthMiddleware().(gin.HandlerFunc); ok {
|
||||
loginMW = mw
|
||||
}
|
||||
if mw, ok := authSvc.RequireAdminMiddleware().(gin.HandlerFunc); ok {
|
||||
adminMW = mw
|
||||
}
|
||||
}
|
||||
|
||||
// 0a. Register migrations
|
||||
ctx.Migrations().Register("admin", adminMigrations)
|
||||
|
||||
// 1. Register Admin HTTP Routes
|
||||
adminRouter := ctx.Router().Group("/api/v1/admin", loginMW, adminMW)
|
||||
{
|
||||
// Status & Diagnostics
|
||||
adminRouter.GET("/status", GetSystemStatus)
|
||||
adminRouter.GET("/status/log-database", GetLogDatabaseStatus)
|
||||
adminRouter.GET("/db-info", GetDatabaseInfo)
|
||||
adminRouter.GET("/db-export", ExportDatabase)
|
||||
|
||||
// DB Management
|
||||
dbGroup := adminRouter.Group("/db-manage")
|
||||
{
|
||||
dbGroup.GET("/overview", GetDBOverview)
|
||||
dbGroup.GET("/tables", ListDBTables)
|
||||
dbGroup.GET("/table-data", GetDBTableData)
|
||||
dbGroup.POST("/query", ExecuteSQL)
|
||||
}
|
||||
|
||||
// Cache Management
|
||||
cacheGroup := adminRouter.Group("/cache")
|
||||
{
|
||||
cacheGroup.GET("/status", GetCacheStatus)
|
||||
cacheGroup.POST("/config", UpdateCacheConfig)
|
||||
cacheGroup.POST("/clear", ClearCache)
|
||||
}
|
||||
|
||||
// Updater
|
||||
updateGroup := adminRouter.Group("/update")
|
||||
{
|
||||
updateGroup.GET("", GetUpdateStatus)
|
||||
updateGroup.POST("/apply", ApplyUpdate)
|
||||
}
|
||||
|
||||
// Logs
|
||||
logsGroup := adminRouter.Group("/logs")
|
||||
{
|
||||
logsGroup.GET("", GetLogs)
|
||||
logsGroup.GET("/access", GetAccessLogs)
|
||||
logsGroup.GET("/analytics", GetLogsAnalytics)
|
||||
logsGroup.GET("/ws", HandleLogWebSocket)
|
||||
}
|
||||
|
||||
// Users
|
||||
usersGroup := adminRouter.Group("/users")
|
||||
{
|
||||
usersGroup.GET("", ListUsers)
|
||||
usersGroup.POST("", CreateUser)
|
||||
usersGroup.GET("/:id", GetUser)
|
||||
usersGroup.PUT("/:id/status", UpdateUserStatus)
|
||||
usersGroup.PUT("/:id", UpdateUser)
|
||||
usersGroup.DELETE("/:id", DeleteUser)
|
||||
}
|
||||
|
||||
// Auth Sources
|
||||
authSourcesGroup := adminRouter.Group("/auth-sources")
|
||||
{
|
||||
authSourcesGroup.GET("", ListAuthSources)
|
||||
authSourcesGroup.POST("", CreateAuthSource)
|
||||
authSourcesGroup.PUT("/:id", UpdateAuthSource)
|
||||
authSourcesGroup.PUT("/:id/toggle", ToggleAuthSource)
|
||||
authSourcesGroup.DELETE("/:id", DeleteAuthSource)
|
||||
}
|
||||
|
||||
// System Configs
|
||||
configGroup := adminRouter.Group("/system-configs")
|
||||
{
|
||||
configGroup.GET("", ListSystemConfigs)
|
||||
configGroup.POST("", CreateSystemConfig)
|
||||
configGroup.POST("/smtp/test", TestSMTP)
|
||||
|
||||
keyGroup := configGroup.Group("/:key")
|
||||
{
|
||||
keyGroup.GET("", GetSystemConfig)
|
||||
keyGroup.PUT("", UpdateSystemConfig)
|
||||
}
|
||||
}
|
||||
|
||||
// Templates
|
||||
templateGroup := adminRouter.Group("/templates")
|
||||
{
|
||||
templateGroup.GET("", ListTemplates)
|
||||
templateGroup.POST("", CreateTemplate)
|
||||
|
||||
keyGroup := templateGroup.Group("/:key")
|
||||
{
|
||||
keyGroup.GET("", GetTemplate)
|
||||
keyGroup.PUT("", UpdateTemplate)
|
||||
keyGroup.DELETE("", DeleteTemplate)
|
||||
}
|
||||
}
|
||||
|
||||
// Tasks
|
||||
taskGroup := adminRouter.Group("/tasks")
|
||||
{
|
||||
taskGroup.GET("/types", ListTaskTypes)
|
||||
taskGroup.POST("/dispatch", DispatchTask)
|
||||
|
||||
executions := taskGroup.Group("/executions")
|
||||
{
|
||||
executions.GET("", ListTaskExecutions)
|
||||
executions.GET("/:id", GetTaskExecution)
|
||||
executions.POST("/:id/retry", RetryTask)
|
||||
}
|
||||
|
||||
schedules := taskGroup.Group("/schedules")
|
||||
{
|
||||
schedules.GET("", ListSchedules)
|
||||
schedules.POST("", CreateSchedule)
|
||||
schedules.PUT("/:id", UpdateSchedule)
|
||||
schedules.DELETE("/:id", DeleteSchedule)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// 2. Register Background Tasks
|
||||
ctx.Task().Register("admin:system_cleanup", func(_ context.Context, _ *asynq.Task) error {
|
||||
return nil
|
||||
}, extpoints.WithTaskRetry(1))
|
||||
|
||||
// 3. Register Cron Schedules
|
||||
ctx.Schedule().RegisterCron("0 4 * * *", "admin:system_cleanup", map[string]string{"type": "daily"})
|
||||
|
||||
// 4. Register Settings Schemas
|
||||
ctx.Settings().Register(extpoints.SettingSchema{
|
||||
Key: "admin.system_cleanup_cron",
|
||||
Default: "0 4 * * *",
|
||||
Description: "Cron expression for nightly system logs and expired tokens cleanup",
|
||||
Type: "string",
|
||||
Category: "maintenance",
|
||||
})
|
||||
|
||||
return nil
|
||||
}
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user