mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-06 07:36:37 +08:00
refactor(plugins): restructure admin and message_gateway into standard layered sub-packages
This commit is contained in:
@@ -0,0 +1,87 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package service
|
||||
|
||||
import (
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/pkg/logger"
|
||||
"Wavelet/plugins/domain/admin/errs"
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
)
|
||||
|
||||
// ListAuthSources returns every configured authentication source.
|
||||
func ListAuthSources(ctx context.Context) ([]contracts.AuthSourceViewDTO, error) {
|
||||
authSvc, err := requireAuthService(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
views, err := authSvc.ListAuthSources(ctx)
|
||||
if err != nil {
|
||||
logger.ErrorF(ctx, "List auth sources failed: %v", err)
|
||||
return nil, errors.New(errs.ListAuthSourcesFailed)
|
||||
}
|
||||
return views, nil
|
||||
}
|
||||
|
||||
// CreateAuthSource registers a new authentication source.
|
||||
func CreateAuthSource(ctx context.Context, source contracts.AuthSourceDTO) (*contracts.AuthSourceDTO, error) {
|
||||
authSvc, err := requireAuthService(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
created, err := authSvc.CreateAuthSource(ctx, source)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("%s%w", errs.CreateAuthSourceFailed, err)
|
||||
}
|
||||
return created, nil
|
||||
}
|
||||
|
||||
// UpdateAuthSource rewrites an existing authentication source.
|
||||
func UpdateAuthSource(
|
||||
ctx context.Context,
|
||||
id uint64,
|
||||
source contracts.AuthSourceDTO,
|
||||
) (*contracts.AuthSourceDTO, error) {
|
||||
authSvc, err := requireAuthService(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
updated, err := authSvc.UpdateAuthSource(ctx, id, source)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return updated, nil
|
||||
}
|
||||
|
||||
// ToggleAuthSource flips the active state of an authentication source.
|
||||
func ToggleAuthSource(ctx context.Context, id uint64) (*contracts.AuthSourceDTO, error) {
|
||||
authSvc, err := requireAuthService(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
toggled, err := authSvc.ToggleAuthSource(ctx, id)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("%s%w", errs.ToggleAuthSourceFailed, err)
|
||||
}
|
||||
return toggled, nil
|
||||
}
|
||||
|
||||
// DeleteAuthSource removes an authentication source.
|
||||
func DeleteAuthSource(ctx context.Context, id uint64) error {
|
||||
authSvc, err := requireAuthService(ctx)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if err := authSvc.DeleteAuthSource(ctx, id); err != nil {
|
||||
return fmt.Errorf("%s%w", errs.DeleteAuthSourceFailed, err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,45 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package service
|
||||
|
||||
import (
|
||||
"context"
|
||||
"strconv"
|
||||
|
||||
pkgcache "Wavelet/pkg/cache/disk"
|
||||
"Wavelet/plugins/domain/admin/model"
|
||||
"Wavelet/plugins/domain/admin/repository"
|
||||
)
|
||||
|
||||
// DiskCacheStatus reports the disk cache usage counters.
|
||||
func DiskCacheStatus() pkgcache.Status {
|
||||
return pkgcache.Default().Status()
|
||||
}
|
||||
|
||||
// ClearDiskCache purges every cached object and resets the tracking counters.
|
||||
func ClearDiskCache() error {
|
||||
return pkgcache.Default().Clear()
|
||||
}
|
||||
|
||||
// UpdateDiskCachePolicy persists the disk cache settings and applies them hot.
|
||||
func UpdateDiskCachePolicy(ctx context.Context, req model.UpdateCacheConfigRequest) error {
|
||||
if err := saveOrUpdateCacheConfig(ctx, model.ConfigKeyDiskCacheMaxSizeMB, strconv.FormatInt(req.MaxSizeMB, 10)); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if err := saveOrUpdateCacheConfig(ctx, model.ConfigKeyDiskCacheTTLMinutes, strconv.FormatInt(req.TTLMinutes, 10)); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if err := saveOrUpdateCacheConfig(ctx, model.ConfigKeyDiskCacheLRUEnabled, strconv.FormatBool(req.LRUEnabled)); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
pkgcache.Default().UpdatePolicy(req.MaxSizeMB, req.TTLMinutes, req.LRUEnabled)
|
||||
return nil
|
||||
}
|
||||
|
||||
func saveOrUpdateCacheConfig(ctx context.Context, key, value string) error {
|
||||
return repository.SaveOrUpdateSystemConfig(ctx, key, value)
|
||||
}
|
||||
@@ -0,0 +1,299 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package service
|
||||
|
||||
import (
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/pkg/logger"
|
||||
mail "Wavelet/pkg/mail"
|
||||
"Wavelet/plugins/domain/admin/errs"
|
||||
"Wavelet/plugins/domain/admin/model"
|
||||
"Wavelet/plugins/domain/admin/repository"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
)
|
||||
|
||||
const maskedConfigValue = "******"
|
||||
|
||||
// PublicSystemConfigs returns the key/value map exposed to unauthenticated clients.
|
||||
func PublicSystemConfigs(ctx context.Context) (map[string]string, error) {
|
||||
configs, err := repository.ListVisibleSystemConfigs(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
resp := make(map[string]string, len(configs))
|
||||
for _, config := range configs {
|
||||
resp[config.Key] = config.Value
|
||||
}
|
||||
return resp, nil
|
||||
}
|
||||
|
||||
// ListAdminSystemConfigs returns every config, optionally filtered by type, with secrets masked.
|
||||
func ListAdminSystemConfigs(ctx context.Context, configType string) ([]model.SystemConfig, error) {
|
||||
configs, err := repository.ListAdminSystemConfigs(ctx, configType)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
for i := range configs {
|
||||
configs[i].Value = MaskSensitiveConfig(configs[i].Key, configs[i].Value)
|
||||
}
|
||||
return configs, nil
|
||||
}
|
||||
|
||||
// GetAdminSystemConfig loads a single config with its secrets masked.
|
||||
func GetAdminSystemConfig(ctx context.Context, key string) (model.SystemConfig, error) {
|
||||
config, err := repository.GetAdminSystemConfigByKey(ctx, key)
|
||||
if err != nil {
|
||||
return model.SystemConfig{}, translateNotFound(err, errs.ErrSystemConfigNotFound)
|
||||
}
|
||||
config.Value = MaskSensitiveConfig(config.Key, config.Value)
|
||||
return config, nil
|
||||
}
|
||||
|
||||
// CreateAdminSystemConfig persists a new config key and refreshes the cache layer.
|
||||
func CreateAdminSystemConfig(ctx context.Context, req model.CreateSystemConfigRequest) error {
|
||||
if isProtectedConfigKey(req.Key) {
|
||||
return errs.ErrProtectedConfigKey
|
||||
}
|
||||
exists, err := repository.SystemConfigExists(ctx, req.Key)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if exists {
|
||||
return errs.ErrConfigKeyExists
|
||||
}
|
||||
|
||||
config := model.SystemConfig{
|
||||
Key: req.Key,
|
||||
Value: req.Value,
|
||||
Type: req.Type,
|
||||
Visibility: req.Visibility,
|
||||
Description: req.Description,
|
||||
}
|
||||
if err := repository.CreateSystemConfigRecord(ctx, &config); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
invalidateSystemConfigCaches(ctx, req.Key)
|
||||
if err := repository.InvalidateVisibleSystemConfigsCache(ctx); err != nil {
|
||||
logger.WarnF(ctx, "清理公共配置列表缓存失败: %v", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// UpdateAdminSystemConfig applies an update to a protected-aware config key inside a transaction.
|
||||
func UpdateAdminSystemConfig(ctx context.Context, key string, req model.UpdateSystemConfigRequest) error {
|
||||
if isProtectedConfigKey(key) {
|
||||
return errs.ErrProtectedConfigKey
|
||||
}
|
||||
config, err := repository.GetAdminSystemConfigByKey(ctx, key)
|
||||
if err != nil {
|
||||
return translateNotFound(err, errs.ErrSystemConfigNotFound)
|
||||
}
|
||||
|
||||
var originalDriver contracts.StorageDriver
|
||||
resolveTaskType := ""
|
||||
resolveResult := ""
|
||||
if key == model.ConfigKeyStorageConfig {
|
||||
var currentCfg contracts.StorageConfigDTO
|
||||
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
|
||||
|
||||
var newCfg contracts.StorageConfigDTO
|
||||
if err := json.Unmarshal([]byte(req.Value), &newCfg); err == nil {
|
||||
resolveTaskType, resolveResult = storageMigrationResolutionTask(originalDriver, newCfg.Driver)
|
||||
}
|
||||
}
|
||||
|
||||
updates := map[string]any{
|
||||
"description": req.Description,
|
||||
}
|
||||
if req.Visibility != nil {
|
||||
updates["visibility"] = *req.Visibility
|
||||
config.Visibility = *req.Visibility
|
||||
}
|
||||
if key != model.ConfigKeySMTPPassword || req.Value != maskedConfigValue {
|
||||
updates["value"] = req.Value
|
||||
config.Value = req.Value
|
||||
}
|
||||
|
||||
if err := repository.UpdateSystemConfigTx(ctx, &config, updates, resolveTaskType, resolveResult); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
invalidateCachesAfterConfigUpdate(ctx, key)
|
||||
return nil
|
||||
}
|
||||
|
||||
// storageMigrationResolutionTask reports the failed-task resolution that a direct storage
|
||||
// config rewrite implies. An empty task type means nothing has to be resolved.
|
||||
func storageMigrationResolutionTask(
|
||||
originalDriver contracts.StorageDriver,
|
||||
newDriver contracts.StorageDriver,
|
||||
) (string, string) {
|
||||
if originalDriver == "" || newDriver != originalDriver {
|
||||
return "", ""
|
||||
}
|
||||
return errs.StorageMigrationTaskType, errs.StorageDriverResolvedResult
|
||||
}
|
||||
|
||||
func isProtectedConfigKey(key string) bool {
|
||||
return key == model.ConfigKeyLogDatabase || key == model.ConfigKeyLogDBMigration
|
||||
}
|
||||
|
||||
func invalidateSystemConfigCaches(ctx context.Context, key string) {
|
||||
if err := repository.InvalidateSystemConfigCache(ctx, key); err != nil {
|
||||
logger.WarnF(ctx, "清理系统配置缓存失败: %v", err)
|
||||
}
|
||||
_ = EmitEvent(ctx, contracts.EventTopicConfigChanged, contracts.ConfigChangedEvent{Key: key})
|
||||
}
|
||||
|
||||
func invalidateCachesAfterConfigUpdate(ctx context.Context, key string) {
|
||||
invalidateSystemConfigCaches(ctx, key)
|
||||
|
||||
if err := repository.InvalidateVisibleSystemConfigsCache(ctx); err != nil {
|
||||
logger.WarnF(ctx, "清理公共配置列表缓存失败: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// TestSMTP sends a probe mail, resolving a masked password from the stored config.
|
||||
func TestSMTP(ctx context.Context, req model.TestSMTPRequest) model.TestSMTPResponse {
|
||||
password := req.SMTPPassword
|
||||
if password == maskedConfigValue {
|
||||
if sc, err := repository.GetSystemConfigByKey(ctx, model.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(ctx, cfg, req.To, subject, body)
|
||||
resp := model.TestSMTPResponse{
|
||||
Success: err == nil,
|
||||
Log: logs,
|
||||
}
|
||||
if err != nil {
|
||||
resp.Error = err.Error()
|
||||
}
|
||||
return resp
|
||||
}
|
||||
|
||||
// MaskSensitiveConfig masks secret config values before exposing to clients.
|
||||
func MaskSensitiveConfig(key, value string) string {
|
||||
if value == "" {
|
||||
return value
|
||||
}
|
||||
switch key {
|
||||
case model.ConfigKeySMTPPassword:
|
||||
return maskedConfigValue
|
||||
case model.ConfigKeyStorageConfig:
|
||||
return maskStorageConfig(value)
|
||||
}
|
||||
return value
|
||||
}
|
||||
|
||||
func maskStorageConfig(value string) string {
|
||||
var cfg contracts.StorageConfigDTO
|
||||
if err := json.Unmarshal([]byte(value), &cfg); err != nil {
|
||||
return value
|
||||
}
|
||||
if cfg.S3.SecretAccessKey != "" {
|
||||
cfg.S3.SecretAccessKey = maskedConfigValue
|
||||
}
|
||||
if cfg.R2.SecretAccessKey != "" {
|
||||
cfg.R2.SecretAccessKey = maskedConfigValue
|
||||
}
|
||||
if cfg.MinIO.SecretAccessKey != "" {
|
||||
cfg.MinIO.SecretAccessKey = maskedConfigValue
|
||||
}
|
||||
if cfg.OSS.SecretAccessKey != "" {
|
||||
cfg.OSS.SecretAccessKey = maskedConfigValue
|
||||
}
|
||||
if cfg.WebDAV.Password != "" {
|
||||
cfg.WebDAV.Password = maskedConfigValue
|
||||
}
|
||||
val, err := json.Marshal(cfg)
|
||||
if err != nil {
|
||||
return value
|
||||
}
|
||||
return string(val)
|
||||
}
|
||||
|
||||
// validateAndMergeStorageConfig parses, merges unmasked secrets, validates parameter values,
|
||||
// and tests connectivity of the new storage configuration.
|
||||
func validateAndMergeStorageConfig(ctx context.Context, value, currentConfig string) (string, error) {
|
||||
var currentCfg contracts.StorageConfigDTO
|
||||
if err := json.Unmarshal([]byte(currentConfig), ¤tCfg); err != nil {
|
||||
return "", fmt.Errorf(errs.ErrParseCurrentStorageConfigFailed, err)
|
||||
}
|
||||
|
||||
var newCfg contracts.StorageConfigDTO
|
||||
if err := json.Unmarshal([]byte(value), &newCfg); err != nil {
|
||||
return "", fmt.Errorf(errs.ErrParseTargetStorageConfigFailed, err)
|
||||
}
|
||||
|
||||
// 合并被掩码屏蔽的敏感信息,获取完整的真实配置
|
||||
targetCfg := newCfg
|
||||
if targetCfg.S3.SecretAccessKey == maskedConfigValue {
|
||||
targetCfg.S3.SecretAccessKey = currentCfg.S3.SecretAccessKey
|
||||
}
|
||||
if targetCfg.R2.SecretAccessKey == maskedConfigValue {
|
||||
targetCfg.R2.SecretAccessKey = currentCfg.R2.SecretAccessKey
|
||||
}
|
||||
if targetCfg.MinIO.SecretAccessKey == maskedConfigValue {
|
||||
targetCfg.MinIO.SecretAccessKey = currentCfg.MinIO.SecretAccessKey
|
||||
}
|
||||
if targetCfg.OSS.SecretAccessKey == maskedConfigValue {
|
||||
targetCfg.OSS.SecretAccessKey = currentCfg.OSS.SecretAccessKey
|
||||
}
|
||||
if targetCfg.WebDAV.Password == maskedConfigValue {
|
||||
targetCfg.WebDAV.Password = currentCfg.WebDAV.Password
|
||||
}
|
||||
|
||||
if err := validateMergedStorageConfig(ctx, currentCfg, newCfg, targetCfg); err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
// 序列化为最终保存的真实明文配置,防止保存屏蔽的 ****** 字符
|
||||
unmaskedVal, err := json.Marshal(targetCfg)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf(errs.ErrSerializeStorageConfigFailed, err)
|
||||
}
|
||||
|
||||
return string(unmaskedVal), nil
|
||||
}
|
||||
|
||||
func validateMergedStorageConfig(ctx context.Context, currentCfg, newCfg, _ contracts.StorageConfigDTO) error {
|
||||
if newCfg.Driver != "" && newCfg.Driver != currentCfg.Driver {
|
||||
uploadCount, err := repository.CountActiveUploads(ctx)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if uploadCount > 0 {
|
||||
return errors.New(errs.StorageDriverSwitchRequiresMigration)
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,131 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package service
|
||||
|
||||
import (
|
||||
"Wavelet/pkg/config"
|
||||
"Wavelet/plugins/domain/admin/model"
|
||||
"Wavelet/plugins/domain/admin/repository"
|
||||
"context"
|
||||
"os"
|
||||
"os/exec"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
// selectSQLKeywords marks statements that return a result set instead of a row count.
|
||||
var selectSQLKeywords = []string{"select", "show", "explain", "describe", "pragma"}
|
||||
|
||||
// DatabaseOverview collects the runtime overview of the active database.
|
||||
func DatabaseOverview(ctx context.Context) (model.DBOverviewResponse, error) {
|
||||
if !config.Config.Database.Enabled {
|
||||
return repository.GetSQLiteOverview(ctx)
|
||||
}
|
||||
return repository.GetPostgresOverview(ctx)
|
||||
}
|
||||
|
||||
// DatabaseTableNames returns every user table of the active database.
|
||||
func DatabaseTableNames(ctx context.Context) ([]string, error) {
|
||||
return repository.ListDatabaseTableNames(ctx)
|
||||
}
|
||||
|
||||
// DatabaseTableData loads one page of a table with its column layout and total row count.
|
||||
func DatabaseTableData(ctx context.Context, req model.GetTableDataRequest) (model.TableDataResponse, error) {
|
||||
quotedTable := repository.QuoteTableName(req.Table)
|
||||
|
||||
total, err := repository.CountDatabaseTableRows(ctx, quotedTable)
|
||||
if err != nil {
|
||||
return model.TableDataResponse{}, err
|
||||
}
|
||||
|
||||
offset := (req.Page - 1) * req.PageSize
|
||||
if offset < 0 {
|
||||
offset = 0
|
||||
}
|
||||
limit := req.PageSize
|
||||
if limit <= 0 {
|
||||
limit = 10
|
||||
}
|
||||
|
||||
cols, results, err := repository.QueryDatabaseTableRows(ctx, quotedTable, limit, offset)
|
||||
if err != nil {
|
||||
return model.TableDataResponse{}, err
|
||||
}
|
||||
|
||||
return model.TableDataResponse{
|
||||
Columns: cols,
|
||||
Total: total,
|
||||
Results: truncateCellValues(results),
|
||||
}, nil
|
||||
}
|
||||
|
||||
// truncateCellValues caps oversized string cells before they reach the console grid.
|
||||
func truncateCellValues(rows []map[string]any) []map[string]any {
|
||||
for _, row := range rows {
|
||||
for column, value := range row {
|
||||
if str, ok := value.(string); ok {
|
||||
row[column] = model.TruncateDisplayValue(str)
|
||||
}
|
||||
}
|
||||
}
|
||||
return rows
|
||||
}
|
||||
|
||||
// ExecuteCustomSQL runs an arbitrary statement issued from the console SQL runner.
|
||||
func ExecuteCustomSQL(ctx context.Context, trimmedSQL string) (model.ExecuteSQLResponse, error) {
|
||||
startTime := time.Now()
|
||||
|
||||
if isSelectStatement(trimmedSQL) {
|
||||
cols, results, err := repository.RunSelectSQL(ctx, trimmedSQL)
|
||||
if err != nil {
|
||||
return model.ExecuteSQLResponse{}, err
|
||||
}
|
||||
return model.ExecuteSQLResponse{
|
||||
Type: "select",
|
||||
Columns: cols,
|
||||
Results: results,
|
||||
AffectedRows: int64(len(results)),
|
||||
ExecutionTimeMs: time.Since(startTime).Milliseconds(),
|
||||
}, nil
|
||||
}
|
||||
|
||||
affectedRows, err := repository.RunMutationSQL(ctx, trimmedSQL)
|
||||
if err != nil {
|
||||
return model.ExecuteSQLResponse{}, err
|
||||
}
|
||||
return model.ExecuteSQLResponse{
|
||||
Type: "exec",
|
||||
AffectedRows: affectedRows,
|
||||
ExecutionTimeMs: time.Since(startTime).Milliseconds(),
|
||||
}, nil
|
||||
}
|
||||
|
||||
// isSelectStatement reports whether the statement yields a result set.
|
||||
func isSelectStatement(trimmedSQL string) bool {
|
||||
lowerSQL := strings.ToLower(trimmedSQL)
|
||||
for _, kw := range selectSQLKeywords {
|
||||
if strings.HasPrefix(lowerSQL, kw) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// DatabaseInfo returns the active database type, name and version.
|
||||
func DatabaseInfo(ctx context.Context) model.DatabaseInfoResponse {
|
||||
if !config.Config.Database.Enabled {
|
||||
return repository.GetSQLiteInfo(ctx)
|
||||
}
|
||||
return repository.GetPostgresInfo(ctx)
|
||||
}
|
||||
|
||||
// OpenSQLiteExportFile opens the active SQLite database file together with its stat info.
|
||||
func OpenSQLiteExportFile() (*os.File, os.FileInfo, error) {
|
||||
return repository.OpenSQLiteExportFile()
|
||||
}
|
||||
|
||||
// NewPgDumpCommand builds the streaming pg_dump command for the active database.
|
||||
func NewPgDumpCommand(ctx context.Context) (*exec.Cmd, string, error) {
|
||||
return repository.NewPgDumpCommand(ctx)
|
||||
}
|
||||
@@ -0,0 +1,238 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package service
|
||||
|
||||
import (
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/pkg/logger"
|
||||
"Wavelet/plugins/domain/admin/errs"
|
||||
"Wavelet/plugins/domain/admin/model"
|
||||
"Wavelet/plugins/domain/admin/repository"
|
||||
"context"
|
||||
"fmt"
|
||||
"net/url"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
const (
|
||||
analyticsDays = 7
|
||||
|
||||
denyingRobotsFile = "User-Agent: *\nDisallow: /\n"
|
||||
allowingRobotsFile = "User-Agent: *\nAllow: /\n"
|
||||
)
|
||||
|
||||
// RecentSystemLogs reads a page of the process log ring buffer.
|
||||
func RecentSystemLogs(cursor, limit int) model.LogsResponse {
|
||||
entries, hasMore := logger.GlobalRingBuffer.Query(cursor, limit)
|
||||
|
||||
resp := model.LogsResponse{
|
||||
Lines: entries,
|
||||
HasMore: hasMore,
|
||||
}
|
||||
if len(entries) > 0 {
|
||||
resp.NextCursor = entries[0].Index
|
||||
}
|
||||
return resp
|
||||
}
|
||||
|
||||
// RobotsTxtBody resolves the robots.txt payload from the indexing setting.
|
||||
func RobotsTxtBody(ctx context.Context) string {
|
||||
enabled, err := repository.GetBoolByKey(ctx, model.ConfigKeySearchEngineIndexingEnabled)
|
||||
if err == nil && enabled {
|
||||
return allowingRobotsFile
|
||||
}
|
||||
return denyingRobotsFile
|
||||
}
|
||||
|
||||
// IsAllowedLogOrigin reports whether a WebSocket handshake origin may subscribe to logs.
|
||||
func IsAllowedLogOrigin(ctx context.Context, origin, host string) bool {
|
||||
if origin == "" {
|
||||
return true
|
||||
}
|
||||
|
||||
// 1. 同源检查 (Same-origin check)
|
||||
u, err := url.Parse(origin)
|
||||
if err == nil && strings.EqualFold(u.Host, host) {
|
||||
return true
|
||||
}
|
||||
|
||||
// 2. 检查配置的允许跨域 Origin (Check allowed origins in system config)
|
||||
sc, cfgErr := repository.GetSystemConfigByKey(ctx, model.ConfigKeyServerAddress)
|
||||
if cfgErr != nil || sc.Value == "" {
|
||||
return false
|
||||
}
|
||||
originToCheck := strings.TrimRight(strings.TrimSpace(origin), "/")
|
||||
for _, allowed := range strings.Split(sc.Value, ",") {
|
||||
allowed = strings.TrimRight(strings.TrimSpace(allowed), "/")
|
||||
if allowed != "" && strings.EqualFold(allowed, originToCheck) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// AccessLogs queries the analytical access log store and decorates rows with user names.
|
||||
func AccessLogs(ctx context.Context, q model.AccessLogQuery) (model.AccessLogsResponse, error) {
|
||||
rc := GetRiskControlService()
|
||||
if rc == nil {
|
||||
return model.AccessLogsResponse{}, errs.ErrLogStoreUnavailable
|
||||
}
|
||||
|
||||
filter, err := buildAccessLogFilter(ctx, q)
|
||||
if err != nil {
|
||||
return model.AccessLogsResponse{}, err
|
||||
}
|
||||
if filter.UserIDs != nil && len(filter.UserIDs) == 0 {
|
||||
return model.AccessLogsResponse{Total: 0, List: []model.AccessLogItem{}}, nil
|
||||
}
|
||||
|
||||
logs, total, err := rc.QueryAccessLogs(ctx, filter, q.Page, q.PageSize)
|
||||
if err != nil {
|
||||
return model.AccessLogsResponse{}, err
|
||||
}
|
||||
if total == 0 {
|
||||
return model.AccessLogsResponse{Total: 0, List: []model.AccessLogItem{}}, nil
|
||||
}
|
||||
|
||||
list := make([]model.AccessLogItem, len(logs))
|
||||
for i, logItem := range logs {
|
||||
list[i] = model.AccessLogItem{
|
||||
ID: logItem.ID,
|
||||
UserID: logItem.UserID,
|
||||
Path: logItem.Path,
|
||||
Method: logItem.Method,
|
||||
IP: logItem.IP,
|
||||
UserAgent: logItem.UserAgent,
|
||||
Status: logItem.Status,
|
||||
Latency: logItem.Latency,
|
||||
CreatedAt: logItem.CreatedAt.Format(time.RFC3339),
|
||||
}
|
||||
}
|
||||
enrichAccessLogsWithUsers(ctx, list)
|
||||
|
||||
return model.AccessLogsResponse{Total: total, List: list}, nil
|
||||
}
|
||||
|
||||
// AccessLogAnalytics aggregates the daily trend of the access log store.
|
||||
func AccessLogAnalytics(ctx context.Context) (model.LogsAnalyticsResponse, error) {
|
||||
rc := GetRiskControlService()
|
||||
if rc == nil {
|
||||
return model.LogsAnalyticsResponse{}, errs.ErrLogStoreUnavailable
|
||||
}
|
||||
|
||||
stats, err := rc.QueryAccessLogStats(ctx, analyticsDays)
|
||||
if err != nil {
|
||||
return model.LogsAnalyticsResponse{}, fmt.Errorf("%s%w", errs.ErrQueryAccessTrendFailed, err)
|
||||
}
|
||||
|
||||
trendList := make([]model.TrendItem, len(stats))
|
||||
for i, st := range stats {
|
||||
trendList[i] = model.TrendItem{
|
||||
Date: st.Date,
|
||||
Count: st.PV,
|
||||
}
|
||||
}
|
||||
|
||||
return model.LogsAnalyticsResponse{
|
||||
Trend: trendList,
|
||||
Browsers: []model.BrowserItem{},
|
||||
TopUsers: []model.TopUserItem{},
|
||||
}, nil
|
||||
}
|
||||
|
||||
// findUserIDsByUsername resolves the user id filter behind a username search term.
|
||||
func findUserIDsByUsername(ctx context.Context, username string) ([]uint64, error) {
|
||||
if userSvc := GetUserService(ctx); userSvc != nil {
|
||||
users, _, err := userSvc.ListUsers(ctx, 1, userQueryMaxLimit, username)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf(errs.ErrQueryUserFailed, err)
|
||||
}
|
||||
ids := make([]uint64, 0, len(users))
|
||||
for _, u := range users {
|
||||
ids = append(ids, u.ID)
|
||||
}
|
||||
return ids, nil
|
||||
}
|
||||
|
||||
ids, err := repository.SearchUserIDsByUsername(ctx, username)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf(errs.ErrQueryUserFailed, err)
|
||||
}
|
||||
return ids, nil
|
||||
}
|
||||
|
||||
const userQueryMaxLimit = 100
|
||||
|
||||
func buildAccessLogFilter(ctx context.Context, q model.AccessLogQuery) (contracts.AccessLogFilterDTO, error) {
|
||||
filter := contracts.AccessLogFilterDTO{}
|
||||
|
||||
if q.Username != "" {
|
||||
userIDs, err := findUserIDsByUsername(ctx, q.Username)
|
||||
if err != nil {
|
||||
return filter, err
|
||||
}
|
||||
filter.UserIDs = userIDs
|
||||
}
|
||||
|
||||
if q.Path != "" {
|
||||
filter.Path = q.Path
|
||||
}
|
||||
|
||||
if q.StartTime != "" {
|
||||
if t, err := parseAccessLogTime(q.StartTime); err == nil {
|
||||
filter.StartTime = &t
|
||||
}
|
||||
}
|
||||
|
||||
if q.EndTime != "" {
|
||||
if t, err := parseAccessLogTime(q.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)
|
||||
}
|
||||
|
||||
// enrichAccessLogsWithUsers attaches usernames and nicknames to access log rows.
|
||||
func enrichAccessLogsWithUsers(ctx context.Context, list []model.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]repository.UserDisplayName, len(userIDs))
|
||||
if userSvc := GetUserService(ctx); userSvc != nil {
|
||||
for _, uid := range userIDs {
|
||||
if u, err := userSvc.GetUserByID(ctx, uid); err == nil && u != nil {
|
||||
userMap[uid] = repository.UserDisplayName{Username: u.Username, Nickname: u.Nickname}
|
||||
}
|
||||
}
|
||||
} else if names, err := repository.LoadUserDisplayNames(ctx, userIDs); err == nil {
|
||||
userMap = names
|
||||
}
|
||||
|
||||
for i := range list {
|
||||
if info, ok := userMap[list[i].UserID]; ok {
|
||||
list[i].Username = info.Username
|
||||
list[i].Nickname = info.Nickname
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,174 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package service
|
||||
|
||||
import (
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/pkg/config"
|
||||
"Wavelet/pkg/logger"
|
||||
"Wavelet/plugins/domain/admin/errs"
|
||||
"Wavelet/plugins/domain/admin/model"
|
||||
"Wavelet/plugins/domain/admin/repository"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
)
|
||||
|
||||
const (
|
||||
// LogDBSwitchTask 切换日志数据库任务标识。
|
||||
LogDBSwitchTask = "logs:db_switch"
|
||||
// TaskTypeLogDBSwitch 管理端任务类型。
|
||||
TaskTypeLogDBSwitch = "logs_db_switch"
|
||||
|
||||
targetPostgres = "postgres"
|
||||
targetSQLite = "sqlite"
|
||||
targetClickHouse = "clickhouse"
|
||||
|
||||
errParseTaskPayloadFailed = "参数解析失败: %w"
|
||||
errInvalidLogTarget = "目标日志库不合法: %s"
|
||||
)
|
||||
|
||||
// LogDBSwitchMeta 描述切换日志数据库任务。
|
||||
var LogDBSwitchMeta = contracts.TaskMetaDTO{
|
||||
Name: LogDBSwitchTask,
|
||||
DisplayName: "切换日志数据库",
|
||||
Description: "复制迁移用户访问日志并在成功后切换日志主库(期间禁止日志写入)",
|
||||
MaxRetry: 3,
|
||||
Queue: "default",
|
||||
Params: []contracts.TaskParamDTO{
|
||||
{Name: "target", Description: "迁移目标:postgres(主库为 PG 时)、sqlite(主库为 SQLite 时)或 clickhouse", Type: "string", Required: true},
|
||||
},
|
||||
}
|
||||
|
||||
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(errParseTaskPayloadFailed, err)
|
||||
}
|
||||
p.Target = normalizeTarget(p.Target)
|
||||
if !validTarget(p.Target) {
|
||||
return nil, fmt.Errorf(errInvalidLogTarget, 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) (*contracts.TaskResultDTO, error) {
|
||||
var p logDBSwitchPayload
|
||||
if err := json.Unmarshal(payload, &p); err != nil {
|
||||
return nil, fmt.Errorf(errParseTaskPayloadFailed, err)
|
||||
}
|
||||
p.Target = normalizeTarget(p.Target)
|
||||
if err := validateSwitch(ctx, p.Target); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
source, err := currentLogDatabase(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
taskSvc := GetTaskService()
|
||||
if taskSvc != nil {
|
||||
taskSvc.AppendLog(ctx, "开始切换日志数据库:%s -> %s", source, p.Target)
|
||||
}
|
||||
|
||||
if err := setMigrationFlag(ctx, logMigrationInProgress); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer func() {
|
||||
if err := setMigrationFlag(ctx, ""); err != nil {
|
||||
logger.ErrorF(ctx, "清除日志迁移冻结标记失败: %v", err)
|
||||
}
|
||||
}()
|
||||
|
||||
rc := GetRiskControlService()
|
||||
if rc != nil {
|
||||
if err := rc.SwitchLogEngine(ctx, p.Target); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
|
||||
if err := flipLogDatabase(ctx, p.Target); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if taskSvc != nil {
|
||||
taskSvc.AppendLog(ctx, "日志数据库已切换为 %s,写入恢复", p.Target)
|
||||
}
|
||||
return &contracts.TaskResultDTO{Message: fmt.Sprintf("日志数据库已从 %s 切换为 %s", source, p.Target)}, nil
|
||||
}
|
||||
|
||||
func validateSwitch(ctx context.Context, target string) error {
|
||||
source, err := currentLogDatabase(ctx)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if source == target {
|
||||
return errors.New(errs.ErrSameLogTarget)
|
||||
}
|
||||
switch target {
|
||||
case targetClickHouse:
|
||||
if !config.Config.ClickHouse.Enabled {
|
||||
return errors.New(errs.ErrClickHouseNotEnabled)
|
||||
}
|
||||
case targetPostgres:
|
||||
if !config.Config.Database.Enabled {
|
||||
return errors.New(errs.ErrPostgresNotEnabled)
|
||||
}
|
||||
case targetSQLite:
|
||||
if config.Config.Database.Enabled {
|
||||
return errors.New(errs.ErrSQLiteNotAllowedAsLogDB)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func currentLogDatabase(ctx context.Context) (string, error) {
|
||||
cfg, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeyLogDatabase)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf(errs.ErrReadLogDatabaseFailed, err)
|
||||
}
|
||||
if cfg.Value == "" {
|
||||
return "", errors.New(errs.ErrLogDatabaseEmpty)
|
||||
}
|
||||
return cfg.Value, nil
|
||||
}
|
||||
|
||||
func setMigrationFlag(ctx context.Context, v string) error {
|
||||
return repository.SaveOrUpdateSystemConfig(ctx, model.ConfigKeyLogDBMigration, v)
|
||||
}
|
||||
|
||||
func flipLogDatabase(ctx context.Context, target string) error {
|
||||
return repository.SaveOrUpdateSystemConfig(ctx, model.ConfigKeyLogDatabase, target)
|
||||
}
|
||||
@@ -0,0 +1,50 @@
|
||||
//go:build !windows
|
||||
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package service
|
||||
|
||||
import (
|
||||
"Wavelet/pkg/logger"
|
||||
"context"
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"syscall"
|
||||
)
|
||||
|
||||
const installedBinaryMode = 0o755
|
||||
|
||||
// ReplaceAndRestart replaces the current executable binary with the staged binary and restarts via syscall.Exec.
|
||||
func ReplaceAndRestart(executable, stagedBinary string) error {
|
||||
ctx := context.Background()
|
||||
logger.InfoF(ctx, "[Updater] Swapping executable: %s -> %s", executable, stagedBinary)
|
||||
backup := executable + ".old"
|
||||
|
||||
if err := os.Remove(backup); err != nil && !os.IsNotExist(err) {
|
||||
return fmt.Errorf("删除旧备份失败: %w", err)
|
||||
}
|
||||
|
||||
if err := os.Rename(executable, backup); err != nil {
|
||||
return fmt.Errorf("备份当前程序失败: %w", err)
|
||||
}
|
||||
|
||||
if err := os.Rename(stagedBinary, executable); err != nil {
|
||||
_ = os.Rename(backup, executable)
|
||||
return fmt.Errorf("替换当前程序失败: %w", err)
|
||||
}
|
||||
|
||||
if err := os.Chmod(executable, installedBinaryMode); err != nil {
|
||||
_ = os.Remove(executable)
|
||||
_ = os.Rename(backup, executable)
|
||||
return fmt.Errorf("设置程序执行权限失败: %w", err)
|
||||
}
|
||||
|
||||
stagingDir := filepath.Dir(stagedBinary)
|
||||
_ = os.RemoveAll(stagingDir)
|
||||
|
||||
logger.InfoF(ctx, "[Updater] Executing syscall.Exec to restart service: %s %v", executable, os.Args)
|
||||
//nolint:gosec // restart process via exec with same binary and args
|
||||
return syscall.Exec(executable, os.Args, os.Environ())
|
||||
}
|
||||
@@ -0,0 +1,16 @@
|
||||
//go:build windows
|
||||
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package service
|
||||
|
||||
import (
|
||||
"Wavelet/plugins/domain/admin/errs"
|
||||
"errors"
|
||||
)
|
||||
|
||||
// ReplaceAndRestart is blocked on Windows.
|
||||
func ReplaceAndRestart(_, _ string) error {
|
||||
return errors.New(errs.ErrAutomaticUpgradeBlocked)
|
||||
}
|
||||
@@ -0,0 +1,204 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package service provides business logic and orchestration for the admin domain.
|
||||
package service
|
||||
|
||||
import (
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/plugins/domain/admin/errs"
|
||||
"Wavelet/plugins/domain/admin/repository"
|
||||
"context"
|
||||
"errors"
|
||||
"sync"
|
||||
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
var (
|
||||
servicesMu sync.RWMutex
|
||||
dbService contracts.DBService
|
||||
cacheService contracts.CacheService
|
||||
userService contracts.UserService
|
||||
authService contracts.AuthService
|
||||
taskService contracts.TaskService
|
||||
storageSvc contracts.StorageService
|
||||
riskControlService contracts.RiskControlService
|
||||
eventEmitter func(ctx context.Context, topic string, payload any) error
|
||||
)
|
||||
|
||||
// SetDBService injects the DBService contract.
|
||||
func SetDBService(s contracts.DBService) {
|
||||
servicesMu.Lock()
|
||||
defer servicesMu.Unlock()
|
||||
dbService = s
|
||||
repository.SetDBService(s)
|
||||
}
|
||||
|
||||
// SetCacheService injects the CacheService contract.
|
||||
func SetCacheService(s contracts.CacheService) {
|
||||
servicesMu.Lock()
|
||||
defer servicesMu.Unlock()
|
||||
cacheService = s
|
||||
repository.SetCacheService(s)
|
||||
}
|
||||
|
||||
// SetUserService injects the UserService contract.
|
||||
func SetUserService(s contracts.UserService) {
|
||||
servicesMu.Lock()
|
||||
defer servicesMu.Unlock()
|
||||
userService = s
|
||||
}
|
||||
|
||||
// SetAuthService injects the AuthService contract.
|
||||
func SetAuthService(s contracts.AuthService) {
|
||||
servicesMu.Lock()
|
||||
defer servicesMu.Unlock()
|
||||
authService = s
|
||||
}
|
||||
|
||||
// SetTaskService injects the TaskService contract.
|
||||
func SetTaskService(s contracts.TaskService) {
|
||||
servicesMu.Lock()
|
||||
defer servicesMu.Unlock()
|
||||
taskService = s
|
||||
}
|
||||
|
||||
// SetStorageService injects the StorageService contract.
|
||||
func SetStorageService(s contracts.StorageService) {
|
||||
servicesMu.Lock()
|
||||
defer servicesMu.Unlock()
|
||||
storageSvc = s
|
||||
}
|
||||
|
||||
// SetRiskControlService injects the RiskControlService contract.
|
||||
func SetRiskControlService(s contracts.RiskControlService) {
|
||||
servicesMu.Lock()
|
||||
defer servicesMu.Unlock()
|
||||
riskControlService = s
|
||||
}
|
||||
|
||||
// SetEventEmitter sets the event emission callback.
|
||||
func SetEventEmitter(fn func(ctx context.Context, topic string, payload any) error) {
|
||||
servicesMu.Lock()
|
||||
defer servicesMu.Unlock()
|
||||
eventEmitter = fn
|
||||
}
|
||||
|
||||
// EmitEvent publishes a domain event if an emitter is registered.
|
||||
func EmitEvent(ctx context.Context, topic string, payload any) error {
|
||||
servicesMu.RLock()
|
||||
defer servicesMu.RUnlock()
|
||||
if eventEmitter == nil {
|
||||
return nil
|
||||
}
|
||||
return eventEmitter(ctx, topic, payload)
|
||||
}
|
||||
|
||||
// ResetServices clears all injected services (used on disposal and testing).
|
||||
func ResetServices() {
|
||||
servicesMu.Lock()
|
||||
defer servicesMu.Unlock()
|
||||
dbService = nil
|
||||
cacheService = nil
|
||||
userService = nil
|
||||
authService = nil
|
||||
taskService = nil
|
||||
storageSvc = nil
|
||||
riskControlService = nil
|
||||
eventEmitter = nil
|
||||
repository.ResetServices()
|
||||
}
|
||||
|
||||
// GetDB returns the GORM DB instance bound to the context if available.
|
||||
func GetDB(ctx context.Context) *gorm.DB {
|
||||
servicesMu.RLock()
|
||||
defer servicesMu.RUnlock()
|
||||
if dbService == nil {
|
||||
return nil
|
||||
}
|
||||
return dbService.DB(ctx)
|
||||
}
|
||||
|
||||
// GetCache returns the unified CacheService instance.
|
||||
func GetCache(_ context.Context) contracts.CacheService {
|
||||
servicesMu.RLock()
|
||||
defer servicesMu.RUnlock()
|
||||
return cacheService
|
||||
}
|
||||
|
||||
// GetUserService returns the UserService instance.
|
||||
func GetUserService(_ context.Context) contracts.UserService {
|
||||
servicesMu.RLock()
|
||||
defer servicesMu.RUnlock()
|
||||
return userService
|
||||
}
|
||||
|
||||
// GetAuthService returns the AuthService instance.
|
||||
func GetAuthService(_ context.Context) contracts.AuthService {
|
||||
servicesMu.RLock()
|
||||
defer servicesMu.RUnlock()
|
||||
return authService
|
||||
}
|
||||
|
||||
// GetTaskService returns the TaskService instance.
|
||||
func GetTaskService() contracts.TaskService {
|
||||
servicesMu.RLock()
|
||||
defer servicesMu.RUnlock()
|
||||
return taskService
|
||||
}
|
||||
|
||||
// GetStorageService returns the StorageService instance.
|
||||
func GetStorageService() contracts.StorageService {
|
||||
servicesMu.RLock()
|
||||
defer servicesMu.RUnlock()
|
||||
return storageSvc
|
||||
}
|
||||
|
||||
// GetRiskControlService returns the RiskControlService instance.
|
||||
func GetRiskControlService() contracts.RiskControlService {
|
||||
servicesMu.RLock()
|
||||
defer servicesMu.RUnlock()
|
||||
return riskControlService
|
||||
}
|
||||
|
||||
// translateNotFound collapses the persistence layer's record-not-found sentinel into
|
||||
// the plugin's own domain error so that no layer above the repository has to import gorm.
|
||||
func translateNotFound(err error, notFound error) error {
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return notFound
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
// isRecordMissing reports whether err originates from a missing persistence row.
|
||||
func isRecordMissing(err error) bool {
|
||||
return errors.Is(err, gorm.ErrRecordNotFound)
|
||||
}
|
||||
|
||||
// requireUserService resolves the injected user contract service.
|
||||
func requireUserService(ctx context.Context) (contracts.UserService, error) {
|
||||
userSvc := GetUserService(ctx)
|
||||
if userSvc == nil {
|
||||
return nil, errs.ErrUserServiceUnavailable
|
||||
}
|
||||
return userSvc, nil
|
||||
}
|
||||
|
||||
// requireAuthService resolves the injected auth contract service.
|
||||
func requireAuthService(ctx context.Context) (contracts.AuthService, error) {
|
||||
authSvc := GetAuthService(ctx)
|
||||
if authSvc == nil {
|
||||
return nil, errs.ErrAuthServiceUnavailable
|
||||
}
|
||||
return authSvc, nil
|
||||
}
|
||||
|
||||
// requireTaskService resolves the injected task contract service.
|
||||
func requireTaskService() (contracts.TaskService, error) {
|
||||
taskSvc := GetTaskService()
|
||||
if taskSvc == nil {
|
||||
return nil, errs.ErrTaskServiceUnavailable
|
||||
}
|
||||
return taskSvc, nil
|
||||
}
|
||||
@@ -0,0 +1,164 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package service
|
||||
|
||||
import (
|
||||
"Wavelet/pkg/config"
|
||||
"Wavelet/pkg/logger"
|
||||
"Wavelet/plugins/domain/admin/model"
|
||||
"Wavelet/plugins/domain/admin/repository"
|
||||
"context"
|
||||
"fmt"
|
||||
"math"
|
||||
"runtime"
|
||||
"time"
|
||||
)
|
||||
|
||||
var startTime = time.Now()
|
||||
|
||||
const (
|
||||
hoursInDay = 24
|
||||
minutesInHour = 60
|
||||
secondsInMinute = 60
|
||||
nanosPerSecond = 1e9
|
||||
|
||||
logDBNamePostgres = "postgres"
|
||||
logDBNameSQLite = "sqlite"
|
||||
logDBNameClickHouse = "clickhouse"
|
||||
defaultLogRetentionDays = 30
|
||||
|
||||
logMigrationIdle = "idle"
|
||||
logMigrationInProgress = "migrating"
|
||||
|
||||
unknownGCLabel = "未知"
|
||||
noGCLabel = "无"
|
||||
)
|
||||
|
||||
// CollectSystemStatus samples the Go runtime counters for the console status page.
|
||||
func CollectSystemStatus() model.SystemStatusResponse {
|
||||
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 = unknownGCLabel
|
||||
default:
|
||||
lastGCTime = noGCLabel
|
||||
}
|
||||
|
||||
var lastPause string
|
||||
if m.NumGC > 0 {
|
||||
lastPause = fmt.Sprintf("%.3fs", float64(m.PauseNs[(m.NumGC-1)%256])/nanosPerSecond)
|
||||
} else {
|
||||
lastPause = "0.000s"
|
||||
}
|
||||
|
||||
return model.SystemStatusResponse{
|
||||
Uptime: uptime,
|
||||
NumGoroutine: numGoroutine,
|
||||
Alloc: model.FormatBytes(m.Alloc),
|
||||
TotalAlloc: model.FormatBytes(m.TotalAlloc),
|
||||
Sys: model.FormatBytes(m.Sys),
|
||||
Lookups: m.Lookups,
|
||||
Mallocs: m.Mallocs,
|
||||
Frees: m.Frees,
|
||||
HeapAlloc: model.FormatBytes(m.HeapAlloc),
|
||||
HeapSys: model.FormatBytes(m.HeapSys),
|
||||
HeapIdle: model.FormatBytes(m.HeapIdle),
|
||||
HeapInuse: model.FormatBytes(m.HeapInuse),
|
||||
HeapReleased: model.FormatBytes(m.HeapReleased),
|
||||
HeapObjects: m.HeapObjects,
|
||||
StackInuse: model.FormatBytes(m.StackInuse),
|
||||
StackSys: model.FormatBytes(m.StackSys),
|
||||
MSpanInuse: model.FormatBytes(m.MSpanInuse),
|
||||
MSpanSys: model.FormatBytes(m.MSpanSys),
|
||||
MCacheInuse: model.FormatBytes(m.MCacheInuse),
|
||||
MCacheSys: model.FormatBytes(m.MCacheSys),
|
||||
BuckHashSys: model.FormatBytes(m.BuckHashSys),
|
||||
GCSys: model.FormatBytes(m.GCSys),
|
||||
OtherSys: model.FormatBytes(m.OtherSys),
|
||||
NextGC: model.FormatBytes(m.NextGC),
|
||||
LastGCTime: lastGCTime,
|
||||
PauseTotalNs: fmt.Sprintf("%.1fs", float64(m.PauseTotalNs)/nanosPerSecond),
|
||||
LastPause: lastPause,
|
||||
NumGC: m.NumGC,
|
||||
}
|
||||
}
|
||||
|
||||
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
|
||||
}
|
||||
|
||||
// LogDatabaseStatus reports the active log engine, migration freeze state and retention.
|
||||
func LogDatabaseStatus(ctx context.Context) model.LogDatabaseStatus {
|
||||
activeDB := logDBNameSQLite
|
||||
migration := logMigrationIdle
|
||||
if rc := GetRiskControlService(); rc != nil {
|
||||
activeDB = rc.ActiveLogEngine(ctx)
|
||||
if rc.IsLogEngineMigrating(ctx) {
|
||||
migration = logMigrationInProgress
|
||||
}
|
||||
}
|
||||
return model.LogDatabaseStatus{
|
||||
ActiveDatabase: activeDB,
|
||||
Migration: migration,
|
||||
RetentionDays: map[string]int{
|
||||
logDBNamePostgres: retentionOr(ctx, model.ConfigKeyLogRetentionDaysPostgres),
|
||||
logDBNameSQLite: retentionOr(ctx, model.ConfigKeyLogRetentionDaysSQLite),
|
||||
logDBNameClickHouse: retentionOr(ctx, model.ConfigKeyLogRetentionDaysClickHouse),
|
||||
},
|
||||
AvailableTargets: availableLogTargets(activeDB),
|
||||
}
|
||||
}
|
||||
|
||||
func retentionOr(ctx context.Context, key string) int {
|
||||
v, err := repository.GetIntByKey(ctx, key)
|
||||
if err != nil {
|
||||
if !isRecordMissing(err) {
|
||||
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,159 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package service_test
|
||||
|
||||
import (
|
||||
"Wavelet/plugins/domain/admin/model"
|
||||
"Wavelet/plugins/domain/admin/repository"
|
||||
"Wavelet/plugins/domain/admin/service"
|
||||
"context"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/glebarez/sqlite"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
type testDBService struct {
|
||||
db *gorm.DB
|
||||
}
|
||||
|
||||
func (s *testDBService) DB(ctx context.Context) *gorm.DB {
|
||||
return s.db
|
||||
}
|
||||
|
||||
func (s *testDBService) MasterDB(ctx context.Context) *gorm.DB {
|
||||
return s.db
|
||||
}
|
||||
|
||||
func (s *testDBService) GORM() *gorm.DB {
|
||||
return s.db
|
||||
}
|
||||
|
||||
func (s *testDBService) Named(_ string) *gorm.DB {
|
||||
return s.db
|
||||
}
|
||||
|
||||
func setupSystemConfigTest(t *testing.T) (*gorm.DB, func()) {
|
||||
t.Helper()
|
||||
|
||||
sqliteDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{
|
||||
DisableForeignKeyConstraintWhenMigrating: true,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("gorm.Open(sqlite) error = %v", err)
|
||||
}
|
||||
if err := sqliteDB.AutoMigrate(&model.SystemConfig{}); err != nil {
|
||||
t.Fatalf("AutoMigrate(SystemConfig) error = %v", err)
|
||||
}
|
||||
|
||||
siteConfig := model.SystemConfig{
|
||||
Key: model.ConfigKeySiteName,
|
||||
Value: "Wavelet",
|
||||
Type: "system",
|
||||
Description: "系统平台的展示名称",
|
||||
}
|
||||
if err := sqliteDB.Create(&siteConfig).Error; err != nil {
|
||||
t.Fatalf("Create(site_name) error = %v", err)
|
||||
}
|
||||
|
||||
service.SetDBService(&testDBService{db: sqliteDB})
|
||||
|
||||
cleanup := func() {
|
||||
repository.StopSystemConfigCacheListener()
|
||||
repository.ResetSystemConfigRAMCacheForTest()
|
||||
service.ResetServices()
|
||||
}
|
||||
|
||||
return sqliteDB, cleanup
|
||||
}
|
||||
|
||||
func TestListSystemConfigsByKeys_EmptyKeys(t *testing.T) {
|
||||
result, err := repository.ListSystemConfigsByKeys(context.Background(), nil)
|
||||
if err != nil {
|
||||
t.Fatalf("ListSystemConfigsByKeys(nil) error = %v", err)
|
||||
}
|
||||
if len(result) != 0 {
|
||||
t.Fatalf("ListSystemConfigsByKeys(nil) = %#v, want empty map", result)
|
||||
}
|
||||
}
|
||||
|
||||
func TestListSystemConfigsByKeys_LoadsFromRAMCache(t *testing.T) {
|
||||
dbConn, cleanup := setupSystemConfigTest(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
repository.ResetSystemConfigRAMCacheForTest()
|
||||
|
||||
// Initial load
|
||||
warm, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeySiteName)
|
||||
if err != nil {
|
||||
t.Fatalf("GetSystemConfigByKey(site_name) warm error = %v", err)
|
||||
}
|
||||
if warm.Value != "Wavelet" {
|
||||
t.Fatalf("GetSystemConfigByKey(site_name).Value = %q, want %q", warm.Value, "Wavelet")
|
||||
}
|
||||
|
||||
// Update DB directly
|
||||
if err := dbConn.Model(&model.SystemConfig{}).
|
||||
Where("key = ?", model.ConfigKeySiteName).
|
||||
Update("value", "db_only_value").Error; err != nil {
|
||||
t.Fatalf("Update(site_name) error = %v", err)
|
||||
}
|
||||
|
||||
// Fetch via ListSystemConfigsByKeys should serve from local store (meaning the old value "Wavelet")
|
||||
configs, err := repository.ListSystemConfigsByKeys(ctx, []string{model.ConfigKeySiteName})
|
||||
if err != nil {
|
||||
t.Fatalf("ListSystemConfigsByKeys(site_name) error = %v", err)
|
||||
}
|
||||
|
||||
sc, ok := configs[model.ConfigKeySiteName]
|
||||
if !ok {
|
||||
t.Fatal("ListSystemConfigsByKeys(site_name) missing site_name entry")
|
||||
}
|
||||
if sc.Value != "Wavelet" {
|
||||
t.Fatalf("ListSystemConfigsByKeys(site_name).Value = %q, want cached value %q", sc.Value, "Wavelet")
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetSystemConfigByGroupAndInvalidation(t *testing.T) {
|
||||
dbConn, cleanup := setupSystemConfigTest(t)
|
||||
defer cleanup()
|
||||
ctx := context.Background()
|
||||
|
||||
repository.ResetSystemConfigRAMCacheForTest()
|
||||
|
||||
// Get via specific group/type
|
||||
cfg, err := repository.GetSystemConfigByGroup(ctx, repository.ConfigCacheType, model.ConfigKeySiteName)
|
||||
if err != nil {
|
||||
t.Fatalf("GetSystemConfigByGroup error = %v", err)
|
||||
}
|
||||
if cfg.Value != "Wavelet" {
|
||||
t.Fatalf("value = %q, want %q", cfg.Value, "Wavelet")
|
||||
}
|
||||
|
||||
// Direct DB update
|
||||
if err := dbConn.Model(&model.SystemConfig{}).
|
||||
Where("key = ?", model.ConfigKeySiteName).
|
||||
Update("value", "new_site_name").Error; err != nil {
|
||||
t.Fatalf("DB Update error = %v", err)
|
||||
}
|
||||
|
||||
// Invalidate
|
||||
if err := repository.InvalidateSystemConfigCache(ctx, model.ConfigKeySiteName); err != nil {
|
||||
t.Fatalf("InvalidateSystemConfigCache error = %v", err)
|
||||
}
|
||||
|
||||
// Wait for broadcast execution
|
||||
time.Sleep(100 * time.Millisecond)
|
||||
|
||||
// Fetch again
|
||||
updated, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeySiteName)
|
||||
if err != nil {
|
||||
t.Fatalf("GetSystemConfigByKey error = %v", err)
|
||||
}
|
||||
if updated.Value != "new_site_name" {
|
||||
t.Fatalf("value = %q, want %q", updated.Value, "new_site_name")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,218 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package service
|
||||
|
||||
import (
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/pkg/logger"
|
||||
"Wavelet/plugins/domain/admin/errs"
|
||||
"Wavelet/plugins/domain/admin/model"
|
||||
"Wavelet/plugins/domain/admin/repository"
|
||||
"context"
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"github.com/robfig/cron/v3"
|
||||
)
|
||||
|
||||
// ListTaskTypes returns every dispatchable task type declared in the task registry.
|
||||
func ListTaskTypes() []contracts.TaskMetaDTO {
|
||||
taskSvc := GetTaskService()
|
||||
if taskSvc == nil {
|
||||
return []contracts.TaskMetaDTO{}
|
||||
}
|
||||
return taskSvc.ListTasks()
|
||||
}
|
||||
|
||||
// DispatchTask validates and enqueues a manual task run, returning the new task id.
|
||||
func DispatchTask(ctx context.Context, req model.DispatchTaskRequest) (string, error) {
|
||||
taskSvc, err := requireTaskService()
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
meta, ok := taskSvc.GetTaskMeta(req.TaskType)
|
||||
if !ok {
|
||||
return "", errs.ErrInvalidTaskType
|
||||
}
|
||||
|
||||
validated, err := validateTaskPayload(taskSvc, meta.Name, req.Payload)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
taskID, err := taskSvc.Dispatch(ctx, req.TaskType, validated, "manual")
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("%s: %w", errs.TaskDispatchFailed, err)
|
||||
}
|
||||
return taskID, nil
|
||||
}
|
||||
|
||||
// validateTaskPayload normalises an optional raw payload through the task registry.
|
||||
func validateTaskPayload(taskSvc contracts.TaskService, name, payload string) ([]byte, error) {
|
||||
var payloadBytes []byte
|
||||
if strings.TrimSpace(payload) != "" {
|
||||
payloadBytes = []byte(payload)
|
||||
}
|
||||
|
||||
validated, err := taskSvc.ValidatePayload(name, payloadBytes)
|
||||
if err != nil {
|
||||
return nil, errs.NewInvalidInputError(err.Error())
|
||||
}
|
||||
return validated, nil
|
||||
}
|
||||
|
||||
// ListTaskExecutions pages task execution records for the console.
|
||||
func ListTaskExecutions(
|
||||
ctx context.Context,
|
||||
req model.ListTaskExecutionsRequest,
|
||||
) ([]model.TaskExecution, int64, error) {
|
||||
if req.TaskType != "" {
|
||||
if taskSvc := GetTaskService(); taskSvc != nil {
|
||||
if meta, ok := taskSvc.GetTaskMeta(req.TaskType); ok {
|
||||
req.TaskType = meta.Name
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
executions, total, err := repository.ListTaskExecutionRecords(ctx, req)
|
||||
if err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
return executions, total, nil
|
||||
}
|
||||
|
||||
// TaskExecution loads a single execution record including its buffered log.
|
||||
func TaskExecution(ctx context.Context, id uint64) (*model.TaskExecution, error) {
|
||||
return repository.GetTaskExecutionByID(ctx, id)
|
||||
}
|
||||
|
||||
// RetryTask re-dispatches a failed execution as a new task run.
|
||||
func RetryTask(ctx context.Context, id uint64) (string, error) {
|
||||
taskSvc, err := requireTaskService()
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
newTaskID, err := taskSvc.Retry(ctx, id)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return newTaskID, nil
|
||||
}
|
||||
|
||||
// IsRetryConflictError reports whether the task registry rejected the retry request
|
||||
// because of the record state rather than an infrastructure failure.
|
||||
func IsRetryConflictError(err error) bool {
|
||||
msg := err.Error()
|
||||
return strings.Contains(msg, errs.RemoteTaskNotFailedMsg) ||
|
||||
strings.Contains(msg, errs.RemoteTaskNotRetryableMsg) ||
|
||||
strings.Contains(msg, errs.RemoteTaskMaxRetryMsg)
|
||||
}
|
||||
|
||||
// IsRetryMissingError reports whether the referenced execution record is absent.
|
||||
func IsRetryMissingError(err error) bool {
|
||||
return strings.Contains(err.Error(), errs.RemoteTaskNotFoundMsg)
|
||||
}
|
||||
|
||||
// ListSchedules returns every dynamic schedule definition.
|
||||
func ListSchedules(ctx context.Context) ([]model.Schedule, error) {
|
||||
return repository.ListSchedulesRecord(ctx)
|
||||
}
|
||||
|
||||
// CreateSchedule validates a schedule definition, persists it and reloads the scheduler.
|
||||
func CreateSchedule(ctx context.Context, req model.CreateScheduleRequest) (*model.Schedule, error) {
|
||||
if _, err := cron.ParseStandard(req.Cron); err != nil {
|
||||
return nil, errs.ErrInvalidCronExpression
|
||||
}
|
||||
|
||||
taskSvc, err := requireTaskService()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
meta, ok := taskSvc.GetTaskMeta(req.TaskType)
|
||||
if !ok {
|
||||
return nil, errs.ErrInvalidTaskType
|
||||
}
|
||||
|
||||
validated, err := validateTaskPayload(taskSvc, meta.Name, req.Payload)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
schedule := &model.Schedule{
|
||||
Name: req.Name,
|
||||
TaskType: req.TaskType,
|
||||
Cron: req.Cron,
|
||||
Payload: string(validated),
|
||||
IsActive: *req.IsActive,
|
||||
}
|
||||
|
||||
if err := repository.CreateScheduleRecord(ctx, schedule); err != nil {
|
||||
return nil, fmt.Errorf("%s: %w", errs.ScheduleSaveFailed, err)
|
||||
}
|
||||
|
||||
reloadScheduler(ctx, taskSvc)
|
||||
return schedule, nil
|
||||
}
|
||||
|
||||
// UpdateSchedule rewrites an existing schedule definition and reloads the scheduler.
|
||||
func UpdateSchedule(ctx context.Context, id uint64, req model.UpdateScheduleRequest) (*model.Schedule, error) {
|
||||
schedule, err := repository.GetScheduleByID(ctx, id)
|
||||
if err != nil {
|
||||
return nil, errs.ErrScheduleNotFound
|
||||
}
|
||||
|
||||
if _, err := cron.ParseStandard(req.Cron); err != nil {
|
||||
return nil, errs.ErrInvalidCronExpression
|
||||
}
|
||||
|
||||
taskSvc, err := requireTaskService()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
meta, ok := taskSvc.GetTaskMeta(req.TaskType)
|
||||
if !ok {
|
||||
return nil, errs.ErrInvalidTaskType
|
||||
}
|
||||
|
||||
validated, err := validateTaskPayload(taskSvc, meta.Name, req.Payload)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
schedule.Name = req.Name
|
||||
schedule.TaskType = req.TaskType
|
||||
schedule.Cron = req.Cron
|
||||
schedule.Payload = string(validated)
|
||||
schedule.IsActive = *req.IsActive
|
||||
|
||||
if err := repository.UpdateScheduleRecord(ctx, schedule); err != nil {
|
||||
return nil, fmt.Errorf("%s: %w", errs.ScheduleSaveFailed, err)
|
||||
}
|
||||
|
||||
reloadScheduler(ctx, taskSvc)
|
||||
return schedule, nil
|
||||
}
|
||||
|
||||
// DeleteSchedule removes a schedule definition and reloads the scheduler.
|
||||
func DeleteSchedule(ctx context.Context, id uint64) error {
|
||||
if err := repository.DeleteScheduleRecord(ctx, id); err != nil {
|
||||
return fmt.Errorf("%s: %w", errs.ScheduleDeleteFailed, err)
|
||||
}
|
||||
|
||||
if taskSvc := GetTaskService(); taskSvc != nil {
|
||||
reloadScheduler(ctx, taskSvc)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// reloadScheduler triggers the hot reload, degrading gracefully when the scheduler rejects it.
|
||||
func reloadScheduler(ctx context.Context, taskSvc contracts.TaskService) {
|
||||
if err := taskSvc.ReloadScheduler(); err != nil {
|
||||
logger.ErrorF(ctx, "[TaskAdmin] 重载调度器失败: %v", err)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,86 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package service
|
||||
|
||||
import (
|
||||
"Wavelet/plugins/domain/admin/errs"
|
||||
"Wavelet/plugins/domain/admin/model"
|
||||
"Wavelet/plugins/domain/admin/repository"
|
||||
"context"
|
||||
)
|
||||
|
||||
// CreateTemplate persists a new notification template after key collision and field checks.
|
||||
func CreateTemplate(ctx context.Context, req model.CreateTemplateRequest) (model.Template, error) {
|
||||
exists, err := repository.TemplateExistsByKey(ctx, req.Key)
|
||||
if err != nil {
|
||||
return model.Template{}, err
|
||||
}
|
||||
if exists {
|
||||
return model.Template{}, errs.ErrTemplateKeyExists
|
||||
}
|
||||
|
||||
tmpl := model.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 model.Template{}, err
|
||||
}
|
||||
if err := repository.CreateTemplateRecord(ctx, &tmpl); err != nil {
|
||||
return model.Template{}, err
|
||||
}
|
||||
return tmpl, nil
|
||||
}
|
||||
|
||||
// ListTemplates returns every notification template.
|
||||
func ListTemplates(ctx context.Context) ([]model.Template, error) {
|
||||
return repository.ListTemplatesRecord(ctx)
|
||||
}
|
||||
|
||||
// GetTemplate loads a template by its identifier.
|
||||
func GetTemplate(ctx context.Context, key string) (model.Template, error) {
|
||||
tmpl, err := repository.GetTemplateByKey(ctx, key)
|
||||
if err != nil {
|
||||
return model.Template{}, translateNotFound(err, errs.ErrTemplateNotFound)
|
||||
}
|
||||
return tmpl, nil
|
||||
}
|
||||
|
||||
// UpdateTemplate rewrites the mutable fields of an existing template.
|
||||
func UpdateTemplate(ctx context.Context, key string, req model.UpdateTemplateRequest) (model.Template, error) {
|
||||
tmpl, err := GetTemplate(ctx, key)
|
||||
if err != nil {
|
||||
return model.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 model.Template{}, err
|
||||
}
|
||||
if err := repository.SaveTemplateRecord(ctx, &tmpl); err != nil {
|
||||
return model.Template{}, err
|
||||
}
|
||||
return tmpl, nil
|
||||
}
|
||||
|
||||
// DeleteTemplate removes a custom template; system presets are protected.
|
||||
func DeleteTemplate(ctx context.Context, key string) error {
|
||||
tmpl, err := GetTemplate(ctx, key)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if tmpl.IsSystem {
|
||||
return errs.ErrSystemTemplateCannotDelete
|
||||
}
|
||||
return repository.DeleteTemplateRecord(ctx, &tmpl)
|
||||
}
|
||||
@@ -0,0 +1,635 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package service
|
||||
|
||||
import (
|
||||
"Wavelet/pkg/buildinfo"
|
||||
"Wavelet/pkg/logger"
|
||||
"Wavelet/plugins/domain/admin/errs"
|
||||
"Wavelet/plugins/domain/admin/model"
|
||||
"Wavelet/plugins/domain/admin/repository"
|
||||
"archive/tar"
|
||||
"archive/zip"
|
||||
"compress/gzip"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"runtime"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"golang.org/x/mod/semver"
|
||||
)
|
||||
|
||||
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"`
|
||||
}
|
||||
|
||||
type releaseClient interface {
|
||||
Do(req *http.Request) (*http.Response, error)
|
||||
}
|
||||
|
||||
// UpdaterManager manages application binary updates from GitHub releases.
|
||||
type UpdaterManager struct {
|
||||
client releaseClient
|
||||
mu sync.Mutex
|
||||
upgrading bool
|
||||
}
|
||||
|
||||
// DefaultUpdaterManager is the default singleton update manager.
|
||||
var DefaultUpdaterManager = &UpdaterManager{
|
||||
client: &http.Client{Timeout: 10 * time.Minute},
|
||||
}
|
||||
|
||||
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(errs.ErrInvalidRepository)
|
||||
}
|
||||
|
||||
if !strings.Contains(raw, "://") {
|
||||
repo := strings.TrimSuffix(strings.Trim(raw, "/"), ".git")
|
||||
if len(strings.Split(repo, "/")) == repositoryParts {
|
||||
return repo, nil
|
||||
}
|
||||
return "", errors.New(errs.ErrInvalidRepository)
|
||||
}
|
||||
|
||||
parsed, err := url.Parse(raw)
|
||||
if err != nil || !strings.EqualFold(parsed.Hostname(), "github.com") {
|
||||
return "", errors.New(errs.ErrInvalidRepository)
|
||||
}
|
||||
repo := strings.TrimSuffix(strings.Trim(parsed.Path, "/"), ".git")
|
||||
if len(strings.Split(repo, "/")) != repositoryParts {
|
||||
return "", errors.New(errs.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(repo, tag string) []string {
|
||||
names := []string{expectedAssetName(tag)}
|
||||
if parts := strings.Split(repo, "/"); 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(repo 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(repo, 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(errs.ErrNoCompatibleRelease)
|
||||
}
|
||||
return selected, selectedAsset, nil
|
||||
}
|
||||
|
||||
func (m *UpdaterManager) fetchRelease(ctx context.Context, repo string) (githubRelease, releaseAsset, error) {
|
||||
req, err := http.NewRequestWithContext(
|
||||
ctx,
|
||||
http.MethodGet,
|
||||
fmt.Sprintf("%s/repos/%s/releases?per_page=30", githubAPIBaseURL, repo),
|
||||
nil,
|
||||
)
|
||||
if err != nil {
|
||||
return githubRelease{}, releaseAsset{}, fmt.Errorf("%s: %w", errs.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", errs.ErrReleaseRequestFailed, err)
|
||||
}
|
||||
defer func() {
|
||||
_ = resp.Body.Close()
|
||||
}()
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return githubRelease{}, releaseAsset{}, fmt.Errorf("%s: HTTP %d", errs.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", errs.ErrReleaseResponseInvalid, err)
|
||||
}
|
||||
|
||||
release, asset, err := selectLatestRelease(repo, 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) {
|
||||
cfg, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeyUpdateUpstreamRepository)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("%s: %w", errs.ErrInvalidRepository, err)
|
||||
}
|
||||
return parseRepository(cfg.Value)
|
||||
}
|
||||
|
||||
// status returns current version and update status.
|
||||
func (m *UpdaterManager) status(ctx context.Context) (model.UpdaterStatus, releaseAsset, error) {
|
||||
upstreamRepo, err := loadRepository(ctx)
|
||||
if err != nil {
|
||||
return model.UpdaterStatus{}, releaseAsset{}, err
|
||||
}
|
||||
release, asset, err := m.fetchRelease(ctx, upstreamRepo)
|
||||
if err != nil {
|
||||
return model.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 model.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
|
||||
}
|
||||
|
||||
// GetUpdateStatus returns current updater status.
|
||||
func GetUpdateStatus(ctx context.Context) (model.UpdaterStatus, error) {
|
||||
status, _, err := DefaultUpdaterManager.status(ctx)
|
||||
return status, err
|
||||
}
|
||||
|
||||
func downloadArchive(ctx context.Context, client releaseClient, asset releaseAsset, destination string) error {
|
||||
if asset.Size <= 0 || asset.Size > maxArchiveSize {
|
||||
return fmt.Errorf(errs.ErrReleaseAssetSizeInvalid, 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(errs.ErrCreateUpgradeRequestFailed, err)
|
||||
}
|
||||
req.Header.Set("User-Agent", "Wavelet-Updater")
|
||||
|
||||
resp, err := client.Do(req)
|
||||
if err != nil {
|
||||
return fmt.Errorf(errs.ErrDownloadUpgradeAssetFailed, err)
|
||||
}
|
||||
defer func() {
|
||||
_ = resp.Body.Close()
|
||||
}()
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return fmt.Errorf(errs.ErrUpgradeAssetHTTPFailed, 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(errs.ErrCreateUpgradeArchiveFailed, err)
|
||||
}
|
||||
|
||||
written, err := io.Copy(file, io.LimitReader(resp.Body, maxArchiveSize+1))
|
||||
if err != nil {
|
||||
_ = file.Close()
|
||||
return fmt.Errorf(errs.ErrWriteUpgradeArchiveFailed, err)
|
||||
}
|
||||
if err := file.Close(); err != nil {
|
||||
return fmt.Errorf(errs.ErrCloseUpgradeArchiveFailed, err)
|
||||
}
|
||||
if written > maxArchiveSize || written != asset.Size {
|
||||
return fmt.Errorf(errs.ErrUpgradeArchiveSizeMismatch, 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(errs.ErrArchiveContainsIllegalPath, 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(errs.ErrArchivePathOutOfDestination, 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, repo 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(repo, "/"); 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()&0o111 != 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(errs.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(errs.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(errs.ErrExtractedBinaryTooLarge)
|
||||
}
|
||||
logger.InfoF(ctx, "[Updater] Successfully extracted binary to %s", target)
|
||||
return target, nil
|
||||
}
|
||||
return "", errors.New(errs.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(errs.ErrExtractedBinaryTooLarge)
|
||||
}
|
||||
logger.InfoF(ctx, "[Updater] Successfully extracted binary to %s", target)
|
||||
return target, nil
|
||||
}
|
||||
return "", errors.New(errs.ErrNoCompatibleAsset)
|
||||
}
|
||||
|
||||
// PrepareUpgrade validates preconditions and downloads the newest binary.
|
||||
func (m *UpdaterManager) PrepareUpgrade(ctx context.Context) (string, string, error) {
|
||||
if runtime.GOOS == windowsOS {
|
||||
return "", "", errors.New(errs.ErrAutomaticUpgradeBlocked)
|
||||
}
|
||||
if normalizeVersion(buildinfo.Version) == "" {
|
||||
return "", "", errors.New(errs.ErrDevelopmentBuild)
|
||||
}
|
||||
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
if m.upgrading {
|
||||
return "", "", errors.New(errs.ErrUpgradeAlreadyRunning)
|
||||
}
|
||||
|
||||
status, asset, err := m.status(ctx)
|
||||
if err != nil {
|
||||
return "", "", err
|
||||
}
|
||||
if !status.UpdateAvailable {
|
||||
return "", "", errors.New(errs.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(errs.ErrLocateExecutableFailed, err)
|
||||
}
|
||||
executable, err = filepath.EvalSymlinks(executable)
|
||||
if err != nil {
|
||||
return "", "", fmt.Errorf(errs.ErrResolveExecutablePathFailed, err)
|
||||
}
|
||||
|
||||
tempDir, err := os.MkdirTemp(filepath.Dir(executable), ".wavelet-update-*")
|
||||
if err != nil {
|
||||
return "", "", fmt.Errorf(errs.ErrCreateUpgradeDirFailed, 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(errs.ErrExtractUpgradeAssetFailed, err)
|
||||
}
|
||||
logger.InfoF(ctx, "[Updater] Staged binary successfully prepared: %s", stagedBinary)
|
||||
m.upgrading = true
|
||||
return executable, stagedBinary, nil
|
||||
}
|
||||
|
||||
// FinishUpgrade resets the upgrading flag.
|
||||
func (m *UpdaterManager) FinishUpgrade() {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
m.upgrading = false
|
||||
}
|
||||
@@ -0,0 +1,120 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package service
|
||||
|
||||
import (
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/pkg/logger"
|
||||
"Wavelet/plugins/domain/admin/errs"
|
||||
"Wavelet/plugins/domain/admin/model"
|
||||
"context"
|
||||
"errors"
|
||||
)
|
||||
|
||||
// ToUserResponse projects the user contract DTO onto the console response shape.
|
||||
func ToUserResponse(u *contracts.UserDTO) model.UserResponse {
|
||||
if u == nil {
|
||||
return model.UserResponse{}
|
||||
}
|
||||
return model.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,
|
||||
}
|
||||
}
|
||||
|
||||
// AdminListUsers pages users through the user contract service.
|
||||
func AdminListUsers(
|
||||
ctx context.Context,
|
||||
filter contracts.AdminListUsersFilter,
|
||||
) (int64, []*contracts.UserDTO, error) {
|
||||
userSvc, err := requireUserService(ctx)
|
||||
if err != nil {
|
||||
return 0, nil, err
|
||||
}
|
||||
|
||||
total, dtos, err := userSvc.AdminListUsers(ctx, filter)
|
||||
if err != nil {
|
||||
logger.ErrorF(ctx, "List admin users failed: %v", err)
|
||||
return 0, nil, errors.New(errs.ListAdminUsersFailed)
|
||||
}
|
||||
return total, dtos, nil
|
||||
}
|
||||
|
||||
// AdminGetUser loads a single user profile.
|
||||
func AdminGetUser(ctx context.Context, id uint64) (*contracts.UserDTO, error) {
|
||||
userSvc, err := requireUserService(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
targetUser, err := userSvc.AdminGetUser(ctx, id)
|
||||
if err != nil {
|
||||
return nil, translateNotFound(err, errs.ErrUserNotFound)
|
||||
}
|
||||
return targetUser, nil
|
||||
}
|
||||
|
||||
// AdminUpdateUserStatus enables or disables a user account.
|
||||
func AdminUpdateUserStatus(ctx context.Context, id uint64, isActive bool) error {
|
||||
userSvc, err := requireUserService(ctx)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
err = userSvc.AdminUpdateUserStatus(ctx, id, isActive)
|
||||
return translateNotFound(err, errs.ErrUserNotFound)
|
||||
}
|
||||
|
||||
// AdminDeleteUser removes a user on behalf of the acting administrator.
|
||||
func AdminDeleteUser(ctx context.Context, operatorID, id uint64) error {
|
||||
userSvc, err := requireUserService(ctx)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
err = userSvc.AdminDeleteUser(ctx, operatorID, id)
|
||||
return translateNotFound(err, errs.ErrUserNotFound)
|
||||
}
|
||||
|
||||
// AdminCreateUser registers a local-password user.
|
||||
func AdminCreateUser(ctx context.Context, req contracts.AdminCreateUserRequest) (*contracts.UserDTO, error) {
|
||||
userSvc, err := requireUserService(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
newUser, err := userSvc.AdminCreateUser(ctx, req)
|
||||
if err != nil {
|
||||
return nil, translateNotFound(err, errs.ErrUserNotFound)
|
||||
}
|
||||
return newUser, nil
|
||||
}
|
||||
|
||||
// AdminUpdateUser rewrites a user profile and optionally resets its password.
|
||||
func AdminUpdateUser(
|
||||
ctx context.Context,
|
||||
operatorID uint64,
|
||||
req contracts.AdminUpdateUserRequest,
|
||||
) error {
|
||||
userSvc, err := requireUserService(ctx)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
err = userSvc.AdminUpdateUser(ctx, operatorID, req)
|
||||
return translateNotFound(err, errs.ErrUserNotFound)
|
||||
}
|
||||
Reference in New Issue
Block a user