mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-03 23:06:36 +08:00
feat(cordis): add OpenFlare Cordis 架构改造设计
docs(changelog): 修正表述笔误
refactor(cordis): 磁盘缓存改用上上游能力并清理本地副本
按上游/下游归属规约:类型断言守卫已回流 Wavelet(f3d85d5,附回归用例),
本仓库删除 OpenFlare/plugins/server/pkg/cache 整包并改 import 到
Wavelet/pkg/cache/disk,同步后与上游零漂移。
验证:go build 通过;go test ./... exit 0(137 包 ok);256 条路由对拍与
232 条 swagger 操作均零差异;make build-all 四进制;前端零改动。
docs(cordis): 记录 T1 清理结果与五个复用阻塞点
refactor(cordis): server 复用上游 pkg 能力并删除等价本地副本
按上游/下游归属规约清理重复实现,删除 7 个与上游等价的本地包并改 import:
shared/response→pkg/response、pkg/{logger,mail,trace,httppool,cache/ram}→
上游同名包、infra/persistence/batchwriter→pkg/batchwriter。逐项核过差异:
httppool 逐字节相同;logger 的 Config 字段完全一致;response 的 7 个 Abort*
一致;cache/ram 换过去顺带把裸 go 变回带 panic 恢复的 util.Go。
两处非等价差异按语义处理:
- batchwriter.Stats 与 status DTO 原为类型别名,改为消费侧逐字段转换,
避免 model 反向依赖基础设施类型;
- 上游 pkg/idgen 要求显式 Init(本地副本为懒加载自动初始化),本次保留本地
副本,待与 infra 初始化一并迁移(已登记在清理计划)。
验证:go build 通过;go test ./... exit 0(138 包 ok);256 条路由对拍零差异;
make swagger 232 条操作零增减,且归一化后与旧文档深度相等——差异仅为
response.Any / logger.LogEntry 两个定义名随包路径改名,接口形状未变。
chore(cordis): 回流内核与 pkg/util 通用能力并清理 vendoring 污染
按新增的上游/下游归属规约:HandleRaw/BasePath 与版本比较、网络、格式化助手
属通用能力,已提交到 Wavelet 分支 feat/cordis-router-raw-routes,本仓库改为
纯同步获取(pkg/util 已零漂移),补丁登记保留至上游合并。
同时修掉我此前 git add -A 造成的污染:首次 vendoring 把上游工作区里被
gitignore 的运行期产物一起提交进来(upload 的 diskcache 缓存块 650 个与
driver_http/dist 前端构建物 380 个,共 12872 行/1030 文件)。sync-upstream.sh
现显式排除 uploads/dist/data/*.db,.gitignore 补上对应兜底规则。
AGENTS.md 增加上游/下游改动归属规约,并把仍指向前 Cordis 布局的硬性约束
(internal/router + Serve、internal/repository/logstore、internal/platform/bootstrap、
internal/cmd)改到当前插件路径。
验证:go build 通过;go test ./... exit 0(144 包 ok);make swagger 232 条
操作与基线逐条一致;make build-all 四进制;gofmt 干净。
feat(cordis): server 插件化并改由内核挂载控制面路由
新增 plugins/server/plugin.go:Apply 以 ctx.Router().Group(app.api_prefix)
声明根级与 /v1 全部路由;33 个注册函数由 *gin.RouterGroup 改为
core.RouterExtension,RegisterCollection 改用内核新增的 HandleRaw 保留
尾部斜杠变体,AdminMiddlewares 返回 []any(Go 不允许把 []T 展开为 ...any)。
删除 router.Serve 与 registerRoutes,装配根改为 core.App +
driver_http.New(WithEngine(router.BuildEngine())),监听、信号与优雅退出归内核;
前端 SPA 的 NoRoute 兜底因内核暂无贡献点而保留在引擎层。
路由保真证据:plugin_parity_test 对拍 baseline/routes-engine.txt 的 256 条
(方法 路径) 零差异;go test ./... exit 0(144 包 ok,含真实 handler 的
openflare/integration 用例走同一条挂载路径);make swagger 232 条操作与基线
逐条一致;golangci-lint 0 issues;make build-all 四进制;embed_frontend
标签编译通过;前端零改动。
已知待补:带 Redis 的实机 HTTP 冒烟(本机 6379 未启动,session store 与
改造前一样在建店阶段即 fatal),以及 bootstrap 的任务/设置/迁移注册迁入 Apply。
feat(core): RouterExtension 增加 HandleRaw 与 BasePath 以保真尾部斜杠路由
server 插件化的前置:Handle 经 cleanPath 会剥掉尾部斜杠,无法表达
/resource 与 /resource/ 两条不同路由,而 OpenFlare 有 20 个历史 list
端点两者都注册且部署关闭了 RedirectTrailingSlash,缺失即 404。新增
HandleRaw 与 BasePath(作用域包装器同样登记反注册),补 extpoints 用例;
并把 router.Serve 拆出 BuildEngine 以便交给 driver_http.WithEngine 复用,
新增路由表导出 harness,固化 256 条 (方法 路径) 基线供插件化对拍。
上游补丁登记于 backend/OpenFlare/upstream-patches.md,同步脚本改为按目录
前缀输出差异并在同步后提醒确认补丁是否仍在。
验证:go build 通过;go test ./... exit 0(143 包 ok);gofmt 干净。
docs(cordis): 记录 server 插件接入内核的可行路径与内核能力缺口
feat(cordis): agent/relay/flared 落地为内核驱动插件
三个边缘守护进程各新增 plugin.go,实现 core.Plugin + core.Driver
(自定义 DriverType 与同名 profile),装配与生命周期从 main 迁入
Apply/Start/Stop:Apply 负责 JSON 配置加载、运行环境与用户确保、
openresty/frps/frpc 管理器与各服务装配;Start 以 util.Go 拉起阻塞式
runner 与 GeoIP 周期更新;Stop 收敛主循环结果并在超时时报错而非静默。
入口改为 core.NewApp(core.WithProfile(...)) + Prepare/Run,保持
-config 旗标、默认路径、退出码与启动/停止日志不变。
验证:go build 通过;go test ./... exit 0(143 包 ok,含 3 个插件身份
与配置失败路径测试);make build-all 四进制产出;三进制实跑缺失配置
均 exit 1 且错误链保留 load {agent,relay,flared} config 原因;gofmt 干净。
refactor(cordis): 按功能职责拆分为 4 个插件与 share 共享层
backend/OpenFlare 不再平铺遗留分层,改为 plugins/{server,agent,relay,flared}
加 share/:控制面业务(openflare/admin/oauth/user/upload/cap/config/health 与
repository/model/infra/router 等支撑层)归 server;三个边缘守护进程各自成插件;
被两个以上插件消费的 protocol/geoip/wsclient/render/pagesarchive/edge 归 share。
同时把 pkg/util 与 buildinfo 合并回上游 pkg(上游已覆盖全部符号,仅 8 个函数与
2 个类型为 OpenFlare 独有,已一并迁入),装配根统一到 backend/cmd(含三个 daemon
入口),Dockerfile 与 release 工作流的构建路径和 -X 注入路径同步更新。
验证:go build 通过;go test ./... exit 0(141 包 ok);make swagger exit 0 且
232 条 API 操作与基线逐条一致;make build-all 产出 4 进制;-X 注入经二进制
strings 实测生效;日志后端直连门禁改写为按 server 插件业务域扫描并在扫描数为 0
时报错(防门禁静默失效);前端零改动。
feat(cordis): 落地 backend/share 共享层与上游同步脚本
跨插件共享资源(控制消息协议、GeoIP+iputil、边缘守护进程日志)从下游包
移入 backend/share,并声明其只能依赖 core/pkg 与标准/第三方库,禁止反向
引用下游业务与具体插件实现;新增 scripts/sync-upstream.sh 只覆盖
backend/{core,pkg,plugins},同步后 --check 报告零差异,证明与上游逐字一致。
go build 通过,go test ./... exit 0(142 包 ok),前端零改动。
refactor(cordis): 采用与 Wavelet 同构的单模块布局并引入上游内核
按上游结构落位:backend/{core,pkg,plugins} 为 Wavelet 上游拷贝,OpenFlare
全部业务收拢到上游 downstream 所对应的位置 backend/OpenFlare/,模块名保持
Wavelet 以保证上游 import 路径逐字一致、同步零改写;三个 daemon 入口移至
backend/OpenFlare/cmd,backend/cmd 与 main.go 作为控制面装配根。
行为不变:go build 通过,142 个测试包全绿(含上游插件测试),232 条 API
操作与改造前逐条一致,四进制产物正常,前端零改动。swagger 暂只扫描下游代码,
待 P4 挂载上游路由后再纳入 plugins/。
style: 修正模块路径改写导致的 import 分组排序漂移
refactor(layout): Go 代码迁入 backend/ 并将模块名简化为 OpenFlare
对齐上游 Wavelet 的仓库布局,为以第二 module 形态 vendoring Cordis 内核与
平台插件做准备:模块路径整体改写为 OpenFlare,Go 目标加 cd backend,
swaggo 产物移至 backend/docs 并把 json/yaml 复制回 docs/ 供站点消费,
Dockerfile 与 release 工作流的构建目录、ldflags 模块路径同步更新。
行为保持不变:232 条路由与改造前逐条一致,95 个测试包全绿,
四进制产物正常,前端零改动。
chore(cordis): 落地改造计划与 schema/路由基线
新增 legacy_dump_test 迁移快照 harness:在临时 sqlite 库上按生产顺序
(goose.UpTo → zone 导入 → goose.Up)跑完 76 个历史迁移并导出 schema 与
版本序列,作为改造前后一致性门禁的唯一事实来源。同时记录 232 条路由清单
与 foundation 实施计划。
docs(cordis): add OpenFlare Cordis 架构改造设计
明确上游以第二 module 形态 vendoring 进 backend/Wavelet、4 个插件
(server/agent/relay/flared) 全部装载内核,并规定保留 76 个历史 goose
迁移 + 一次性版本 stamp 桥接的迁移方案,配套三方 schema 一致性门禁,
确保已部署库不重跑历史、不丢数据。
This commit is contained in:
@@ -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,265 @@
|
||||
// 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 (
|
||||
"Wavelet/pkg/util"
|
||||
"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)
|
||||
util.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, 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
+432
@@ -0,0 +1,432 @@
|
||||
// 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 (
|
||||
"Wavelet/pkg/util"
|
||||
"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 = 0o750
|
||||
|
||||
// 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
|
||||
}
|
||||
|
||||
var (
|
||||
defaultCache *Cache
|
||||
defaultCacheOnce sync.Once
|
||||
|
||||
// defaultCleanupInterval 全局磁盘缓存默认的过期清理巡检周期
|
||||
defaultCleanupInterval = 10 * time.Minute
|
||||
)
|
||||
|
||||
// Default returns the default global disk cache instance.
|
||||
func Default() *Cache {
|
||||
defaultCacheOnce.Do(func() {
|
||||
defaultCache = New("uploads/diskcache")
|
||||
util.Go(func() {
|
||||
defaultCache.StartCleanupWorker(defaultCleanupInterval)
|
||||
})
|
||||
})
|
||||
return defaultCache
|
||||
}
|
||||
|
||||
// Set stores a key-value pair in the cache.
|
||||
// Use DefaultExpiration for the configured default TTL, NoExpiration for no
|
||||
// TTL, or a positive duration for a business-specific TTL.
|
||||
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, ok := elem.Value.(*cacheItem)
|
||||
if !ok {
|
||||
return fmt.Errorf("cache: evict list entry for %q has invalid type %T", key, elem.Value)
|
||||
}
|
||||
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, ok := elem.Value.(*cacheItem)
|
||||
if !ok {
|
||||
c.mu.RUnlock()
|
||||
return nil, ErrCacheMiss
|
||||
}
|
||||
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, ok := elem.Value.(*cacheItem)
|
||||
if !ok {
|
||||
_ = c.deleteUnlocked(key)
|
||||
return nil, ErrCacheMiss
|
||||
}
|
||||
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 {
|
||||
if item, ok := elem.Value.(*cacheItem); ok {
|
||||
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, 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, ok := elem.Value.(*cacheItem)
|
||||
if !ok {
|
||||
c.evictList.Remove(elem)
|
||||
continue
|
||||
}
|
||||
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, ok := elem.Value.(*cacheItem)
|
||||
if !ok {
|
||||
c.evictList.Remove(elem)
|
||||
delete(c.items, key)
|
||||
continue
|
||||
}
|
||||
if !item.expiredAt.IsZero() && now.After(item.expiredAt) {
|
||||
c.currentSize -= item.size
|
||||
c.evictList.Remove(elem)
|
||||
delete(c.items, key)
|
||||
_ = c.d.Erase(key)
|
||||
}
|
||||
}
|
||||
}
|
||||
+63
@@ -0,0 +1,63 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package disk
|
||||
|
||||
import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
// 这些用例锁住「LRU 链表节点被污染时不得 panic」的行为:一旦 items 与 evictList
|
||||
// 的不变量被破坏(例如后续改动误写节点),缓存必须降级为未命中/跳过,
|
||||
// 而不是在读、写、删除与淘汰路径上崩掉整个进程。
|
||||
|
||||
// corruptEntry 写入一个键后把其链表节点值换成非法类型,返回缓存。
|
||||
func corruptEntry(t *testing.T, key string) *Cache {
|
||||
t.Helper()
|
||||
|
||||
c := New(t.TempDir())
|
||||
require.NoError(t, c.Set(key, []byte("payload"), time.Minute))
|
||||
|
||||
elem, ok := c.items[key]
|
||||
require.True(t, ok, "entry must be tracked after Set")
|
||||
elem.Value = "not-a-cacheItem"
|
||||
return c
|
||||
}
|
||||
|
||||
func TestGetToleratesCorruptEvictEntry(t *testing.T) {
|
||||
c := corruptEntry(t, "k")
|
||||
|
||||
got, err := c.Get("k")
|
||||
require.ErrorIs(t, err, ErrCacheMiss)
|
||||
require.Nil(t, got)
|
||||
}
|
||||
|
||||
func TestSetOverCorruptEvictEntryReportsError(t *testing.T) {
|
||||
c := corruptEntry(t, "k")
|
||||
|
||||
err := c.Set("k", []byte("second"), time.Minute)
|
||||
require.Error(t, err, "Set must report the corrupted tracker entry instead of panicking")
|
||||
require.Contains(t, err.Error(), "invalid type")
|
||||
}
|
||||
|
||||
func TestDeleteToleratesCorruptEvictEntry(t *testing.T) {
|
||||
c := corruptEntry(t, "k")
|
||||
|
||||
require.NotPanics(t, func() { _ = c.Delete("k") })
|
||||
require.NotContains(t, c.items, "k")
|
||||
}
|
||||
|
||||
func TestEvictToleratesCorruptEvictEntry(t *testing.T) {
|
||||
c := corruptEntry(t, "k")
|
||||
// 让任意写入都触发淘汰扫描:扫到被污染的节点必须跳过而非 panic。
|
||||
c.UpdatePolicy(0, 0, true)
|
||||
|
||||
require.NotPanics(t, func() {
|
||||
for i := range 4 {
|
||||
_ = c.Set(string(rune('a'+i)), []byte("x"), time.Minute)
|
||||
}
|
||||
})
|
||||
}
|
||||
Vendored
+199
@@ -0,0 +1,199 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package disk
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func TestDiskCacheBasic(t *testing.T) {
|
||||
testDir := t.TempDir()
|
||||
|
||||
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 := t.TempDir()
|
||||
|
||||
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 := t.TempDir()
|
||||
|
||||
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"), 150*time.Millisecond); err != nil {
|
||||
t.Fatalf("Set(custom, 150ms) 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(80 * 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(100 * 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 := t.TempDir()
|
||||
|
||||
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 := t.TempDir()
|
||||
|
||||
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)
|
||||
}
|
||||
|
||||
// Access k1 again to make it MRU, k2 becomes LRU
|
||||
_, _ = c.Get("k1")
|
||||
|
||||
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
+234
@@ -0,0 +1,234 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package ram
|
||||
|
||||
import (
|
||||
"Wavelet/pkg/util"
|
||||
"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, 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, 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
|
||||
util.Go(func() { deleteKeyIfExpired(configType, key, entry.expireAt) })
|
||||
return CacheItem{}, false
|
||||
}
|
||||
|
||||
return entry.item, true
|
||||
}
|
||||
|
||||
func deleteKeyIfExpired(configType, 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, 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, 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,24 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package ginutil
|
||||
|
||||
import (
|
||||
"Wavelet/pkg/response"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
// errAuthUnavailable is reported when a route's authentication guard cannot be
|
||||
// resolved, so the request is rejected instead of reaching the handler.
|
||||
const errAuthUnavailable = "authServiceUnavailable"
|
||||
|
||||
// AuthUnavailable returns a middleware that denies the request. Plugins use it
|
||||
// as the fallback when contracts.AuthService cannot be resolved or its middleware
|
||||
// has an unexpected shape: the alternative is a pass-through closure that serves
|
||||
// the request as if it were authenticated.
|
||||
func AuthUnavailable() gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
response.AbortUnauthorized(c, errAuthUnavailable)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,23 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package ginutil provides helper utilities for Gin web framework contexts.
|
||||
package ginutil
|
||||
|
||||
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,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, 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,65 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package idgen 提供分布式 ID 生成器
|
||||
package idgen
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"log"
|
||||
"sync"
|
||||
|
||||
"github.com/bwmarrin/snowflake"
|
||||
)
|
||||
|
||||
// 2025-12-01 00:00:00 UTC 的毫秒时间戳
|
||||
const epoch int64 = 1764547200000
|
||||
|
||||
const maxNegativeIDRetries = 3
|
||||
|
||||
var (
|
||||
mu sync.RWMutex
|
||||
node *snowflake.Node
|
||||
)
|
||||
|
||||
// Init initializes the snowflake ID generator with the given node ID.
|
||||
func Init(nodeID int64) error {
|
||||
snowflake.Epoch = epoch
|
||||
|
||||
n, err := snowflake.NewNode(nodeID)
|
||||
if err != nil {
|
||||
return fmt.Errorf("idgen: init node %d failed: %w", nodeID, err)
|
||||
}
|
||||
|
||||
mu.Lock()
|
||||
node = n
|
||||
mu.Unlock()
|
||||
|
||||
log.Printf("[Snowflake] initialized with node ID: %d, epoch: 2025-12-01\n", nodeID)
|
||||
return nil
|
||||
}
|
||||
|
||||
// ErrNotInitialized indicates NextUint64ID was called before Init.
|
||||
var ErrNotInitialized = errors.New("idgen: Init must be called before generating IDs")
|
||||
|
||||
// NextUint64ID 生成下一个分布式唯一 ID。
|
||||
// 理论上不应出现负值;若出现则最多重试 maxNegativeIDRetries 次,仍失败则 panic。
|
||||
func NextUint64ID() uint64 {
|
||||
mu.RLock()
|
||||
n := node
|
||||
mu.RUnlock()
|
||||
|
||||
if n == nil {
|
||||
panic(ErrNotInitialized)
|
||||
}
|
||||
|
||||
for attempt := 1; attempt <= maxNegativeIDRetries; attempt++ {
|
||||
id := n.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,30 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package idgen
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestNextUint64ID(t *testing.T) {
|
||||
require.NoError(t, Init(1))
|
||||
id := NextUint64ID()
|
||||
assert.NotZero(t, id)
|
||||
}
|
||||
|
||||
func TestNextUint64ID_PanicsWhenNotInitialized(t *testing.T) {
|
||||
mu.Lock()
|
||||
node = nil
|
||||
mu.Unlock()
|
||||
|
||||
assert.PanicsWithError(t, ErrNotInitialized.Error(), func() {
|
||||
NextUint64ID()
|
||||
})
|
||||
|
||||
// Restore initialization for other tests
|
||||
require.NoError(t, Init(1))
|
||||
}
|
||||
@@ -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,101 @@
|
||||
// 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,174 @@
|
||||
// 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, 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,191 @@
|
||||
// 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,98 @@
|
||||
// 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 = 0o750
|
||||
|
||||
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,249 @@
|
||||
// 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, 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, 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, subject, body string) (string, error) {
|
||||
var logBuf bytes.Buffer
|
||||
logLine := func(dir, 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,111 @@
|
||||
// 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,56 @@
|
||||
// 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,21 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package testhelper
|
||||
|
||||
import (
|
||||
"Wavelet/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 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package testhelper 提供测试辅助工具
|
||||
package testhelper
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/alicebob/miniredis/v2"
|
||||
"github.com/glebarez/sqlite"
|
||||
"github.com/redis/go-redis/v9"
|
||||
"github.com/redis/go-redis/v9/maintnotifications"
|
||||
"gorm.io/gorm"
|
||||
|
||||
cachepkg "Wavelet/plugins/infra/cache"
|
||||
db "Wavelet/plugins/infra/database"
|
||||
)
|
||||
|
||||
// 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,14 @@
|
||||
// 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,18 @@
|
||||
// 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,63 @@
|
||||
// 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
|
||||
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,51 @@
|
||||
// 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,159 @@
|
||||
// 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, plaintext string) (string, error) {
|
||||
return encryptBytes(signKey, []byte(plaintext))
|
||||
}
|
||||
|
||||
// Decrypt 使用 SignKey 解密字符串数据
|
||||
// signKey: 64 字符 hex 编码的密钥(对应 32 字节,用于 AES-256)
|
||||
// ciphertext: base64 编码的密文
|
||||
// return: 解密后的明文字符串
|
||||
func Decrypt(signKey, 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, 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,28 @@
|
||||
// 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,70 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package util provides shared formatting and string helper functions.
|
||||
package util
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"strconv"
|
||||
)
|
||||
|
||||
const (
|
||||
secondsPerYear = 31104000 // 360 days
|
||||
secondsPerMonth = 2592000 // 30 days
|
||||
secondsPerDay = 86400
|
||||
secondsPerHour = 3600
|
||||
secondsPerMinute = 60
|
||||
)
|
||||
|
||||
const (
|
||||
sizeKB = 1024
|
||||
sizeMB = sizeKB * 1024
|
||||
sizeGB = sizeMB * 1024
|
||||
)
|
||||
|
||||
// Bytes2Size converts a byte count to a human-readable string with unit (B, KB, MB, GB).
|
||||
func Bytes2Size(num int64) string {
|
||||
var numStr string
|
||||
unit := "B"
|
||||
switch {
|
||||
case num/int64(sizeGB) >= 1:
|
||||
numStr = fmt.Sprintf("%.2f", float64(num)/float64(sizeGB))
|
||||
unit = "GB"
|
||||
case num/int64(sizeMB) >= 1:
|
||||
numStr = strconv.Itoa(int(float64(num) / float64(sizeMB)))
|
||||
unit = "MB"
|
||||
case num/int64(sizeKB) >= 1:
|
||||
numStr = strconv.Itoa(int(float64(num) / float64(sizeKB)))
|
||||
unit = "KB"
|
||||
default:
|
||||
numStr = strconv.FormatInt(num, 10)
|
||||
}
|
||||
return numStr + " " + unit
|
||||
}
|
||||
|
||||
// Seconds2Time converts a number of seconds to a human-readable Chinese duration string.
|
||||
func Seconds2Time(num int) (time string) {
|
||||
if num/secondsPerYear > 0 {
|
||||
time += strconv.Itoa(num/secondsPerYear) + " 年 "
|
||||
num %= secondsPerYear
|
||||
}
|
||||
if num/secondsPerMonth > 0 {
|
||||
time += strconv.Itoa(num/secondsPerMonth) + " 个月 "
|
||||
num %= secondsPerMonth
|
||||
}
|
||||
if num/secondsPerDay > 0 {
|
||||
time += strconv.Itoa(num/secondsPerDay) + " 天 "
|
||||
num %= secondsPerDay
|
||||
}
|
||||
if num/secondsPerHour > 0 {
|
||||
time += strconv.Itoa(num/secondsPerHour) + " 小时 "
|
||||
num %= secondsPerHour
|
||||
}
|
||||
if num/secondsPerMinute > 0 {
|
||||
time += strconv.Itoa(num/secondsPerMinute) + " 分钟 "
|
||||
num %= secondsPerMinute
|
||||
}
|
||||
time += strconv.Itoa(num) + " 秒"
|
||||
return
|
||||
}
|
||||
@@ -0,0 +1,52 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package util
|
||||
|
||||
import (
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestBytes2Size(t *testing.T) {
|
||||
tests := []struct {
|
||||
input int64
|
||||
expected string
|
||||
}{
|
||||
{0, "0 B"},
|
||||
{500, "500 B"},
|
||||
{1023, "1023 B"},
|
||||
{1024, "1 KB"},
|
||||
{2048, "2 KB"},
|
||||
{1024 * 1024, "1 MB"},
|
||||
{1024 * 1024 * 1024, "1.00 GB"},
|
||||
{1024 * 1024 * 1024 * 2, "2.00 GB"},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
result := Bytes2Size(tt.input)
|
||||
if result != tt.expected {
|
||||
t.Errorf("Bytes2Size(%d) = %q, expected %q", tt.input, result, tt.expected)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestSeconds2Time(t *testing.T) {
|
||||
tests := []struct {
|
||||
input int
|
||||
expected string
|
||||
}{
|
||||
{0, "0 秒"},
|
||||
{30, "30 秒"},
|
||||
{60, "1 分钟 0 秒"},
|
||||
{125, "2 分钟 5 秒"},
|
||||
{3600, "1 小时 0 秒"},
|
||||
{86400, "1 天 0 秒"},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
result := Seconds2Time(tt.input)
|
||||
if result != tt.expected {
|
||||
t.Errorf("Seconds2Time(%d) = %q, expected %q", tt.input, result, tt.expected)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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,66 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package util
|
||||
|
||||
import (
|
||||
"Wavelet/pkg/httppool"
|
||||
"context"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"time"
|
||||
)
|
||||
|
||||
// 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,45 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package util
|
||||
|
||||
import (
|
||||
"log/slog"
|
||||
"net"
|
||||
)
|
||||
|
||||
// GetIP returns the first private IPv4 address found on the local network interfaces.
|
||||
func GetIP() (ip string) {
|
||||
ips, err := net.InterfaceAddrs()
|
||||
if err != nil {
|
||||
slog.Error("get interface addresses failed", "error", err)
|
||||
return ip
|
||||
}
|
||||
|
||||
for _, a := range ips {
|
||||
if candidate, ok := privateIPv4FromAddr(a); ok {
|
||||
return candidate
|
||||
}
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
func privateIPv4FromAddr(addr net.Addr) (string, bool) {
|
||||
ipNet, ok := addr.(*net.IPNet)
|
||||
if !ok || ipNet.IP.IsLoopback() || ipNet.IP.To4() == nil {
|
||||
return "", false
|
||||
}
|
||||
ip := ipNet.IP.String()
|
||||
if isPrivateIPv4(ip) {
|
||||
return ip, true
|
||||
}
|
||||
return "", false
|
||||
}
|
||||
|
||||
func isPrivateIPv4(ip string) bool {
|
||||
parsedIP := net.ParseIP(ip)
|
||||
if parsedIP == nil {
|
||||
return false
|
||||
}
|
||||
return parsedIP.IsPrivate()
|
||||
}
|
||||
@@ -0,0 +1,36 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package util
|
||||
|
||||
import (
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestIsPrivateIPv4(t *testing.T) {
|
||||
tests := []struct {
|
||||
ip string
|
||||
expected bool
|
||||
}{
|
||||
{"127.0.0.1", false}, // Loopback is not in RFC 1918 private range
|
||||
{"10.0.0.1", true},
|
||||
{"172.16.0.1", true},
|
||||
{"192.168.1.1", true},
|
||||
{"8.8.8.8", false},
|
||||
{"invalid-ip", false},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
result := isPrivateIPv4(tt.ip)
|
||||
if result != tt.expected {
|
||||
t.Errorf("isPrivateIPv4(%q) = %v, expected %v", tt.ip, result, tt.expected)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetIP(t *testing.T) {
|
||||
ip := GetIP()
|
||||
// GetIP should return empty if no private IPv4 address is configured, or a valid IP.
|
||||
// We just ensure it doesn't panic.
|
||||
t.Logf("GetIP returned: %q", ip)
|
||||
}
|
||||
@@ -0,0 +1,50 @@
|
||||
// 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
|
||||
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,79 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package util
|
||||
|
||||
import (
|
||||
"sort"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
// Unique returns a new slice containing only the unique elements of the input slice,
|
||||
// preserving their original order.
|
||||
func Unique[T comparable](slice []T) []T {
|
||||
if slice == nil {
|
||||
return nil
|
||||
}
|
||||
seen := make(map[T]struct{})
|
||||
result := make([]T, 0)
|
||||
for _, item := range slice {
|
||||
if _, ok := seen[item]; ok {
|
||||
continue
|
||||
}
|
||||
seen[item] = struct{}{}
|
||||
result = append(result, item)
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
// UniqueAndCleanStringSlice trims spaces, removes empty elements, and returns only the unique elements
|
||||
// of the input string slice. It preserves order and returns nil if the resulting slice is empty.
|
||||
func UniqueAndCleanStringSlice(slice []string) []string {
|
||||
if slice == nil {
|
||||
return nil
|
||||
}
|
||||
seen := make(map[string]struct{})
|
||||
result := make([]string, 0)
|
||||
for _, item := range slice {
|
||||
trimmed := strings.TrimSpace(item)
|
||||
if trimmed == "" {
|
||||
continue
|
||||
}
|
||||
if _, ok := seen[trimmed]; ok {
|
||||
continue
|
||||
}
|
||||
seen[trimmed] = struct{}{}
|
||||
result = append(result, trimmed)
|
||||
}
|
||||
if len(result) == 0 {
|
||||
return nil
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
// IdentifiableTimeRecord represents a database record that has a unique ID and a primary timestamp field.
|
||||
type IdentifiableTimeRecord interface {
|
||||
GetID() uint
|
||||
GetTime() time.Time
|
||||
}
|
||||
|
||||
// SortAndLimitRecords sorts a slice of IdentifiableTimeRecord descendingly by their timestamp (and ID as a tie-breaker),
|
||||
// and limits the slice to the specified size if limit > 0.
|
||||
func SortAndLimitRecords[T IdentifiableTimeRecord](rows []T, limit int) []T {
|
||||
if len(rows) == 0 {
|
||||
return rows
|
||||
}
|
||||
sort.Slice(rows, func(i, j int) bool {
|
||||
ti := rows[i].GetTime()
|
||||
tj := rows[j].GetTime()
|
||||
if ti.Equal(tj) {
|
||||
return rows[i].GetID() > rows[j].GetID()
|
||||
}
|
||||
return ti.After(tj)
|
||||
})
|
||||
if limit > 0 && len(rows) > limit {
|
||||
rows = rows[:limit]
|
||||
}
|
||||
return rows
|
||||
}
|
||||
@@ -0,0 +1,15 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package util
|
||||
|
||||
import "strings"
|
||||
|
||||
// TrimStringFields trims leading and trailing spaces from all provided string pointers.
|
||||
func TrimStringFields(fields ...*string) {
|
||||
for _, f := range fields {
|
||||
if f != nil {
|
||||
*f = strings.TrimSpace(*f)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,34 @@
|
||||
// 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,28 @@
|
||||
// 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,22 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package util
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"strconv"
|
||||
)
|
||||
|
||||
// Interface2String converts a string, int, or float64 value to its string representation.
|
||||
func Interface2String(inter any) string {
|
||||
switch v := inter.(type) {
|
||||
case string:
|
||||
return v
|
||||
case int:
|
||||
return strconv.Itoa(v)
|
||||
case float64:
|
||||
return fmt.Sprintf("%f", v)
|
||||
}
|
||||
return "Not Implemented"
|
||||
}
|
||||
@@ -0,0 +1,129 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package util
|
||||
|
||||
import (
|
||||
"strconv"
|
||||
"strings"
|
||||
)
|
||||
|
||||
const gitDescribeMinIdentifiers = 2
|
||||
|
||||
// VersionInfo holds the parsed components of a semantic version string.
|
||||
type VersionInfo struct {
|
||||
Valid bool
|
||||
IsDev bool
|
||||
Numbers []int
|
||||
Prerelease []string
|
||||
GitDescribeDistance int
|
||||
GitDescribeTail []string
|
||||
}
|
||||
|
||||
// ParseVersionInfo parses a version string into a structured VersionInfo.
|
||||
func ParseVersionInfo(version string) VersionInfo {
|
||||
normalized := strings.TrimSpace(strings.TrimPrefix(version, "v"))
|
||||
if normalized == "" || normalized == "dev" {
|
||||
return VersionInfo{IsDev: strings.EqualFold(normalized, "dev")}
|
||||
}
|
||||
base := normalized
|
||||
prerelease := ""
|
||||
if separator := strings.IndexRune(normalized, '-'); separator >= 0 {
|
||||
base = normalized[:separator]
|
||||
prerelease = normalized[separator+1:]
|
||||
}
|
||||
|
||||
segments := strings.Split(base, ".")
|
||||
parts := make([]int, 0, len(segments))
|
||||
for _, segment := range segments {
|
||||
segment = strings.TrimSpace(segment)
|
||||
if segment == "" {
|
||||
parts = append(parts, 0)
|
||||
continue
|
||||
}
|
||||
|
||||
numeric := strings.Builder{}
|
||||
for _, r := range segment {
|
||||
if r < '0' || r > '9' {
|
||||
break
|
||||
}
|
||||
numeric.WriteRune(r)
|
||||
}
|
||||
if numeric.Len() == 0 {
|
||||
parts = append(parts, 0)
|
||||
continue
|
||||
}
|
||||
value, err := strconv.Atoi(numeric.String())
|
||||
if err != nil {
|
||||
return VersionInfo{}
|
||||
}
|
||||
parts = append(parts, value)
|
||||
}
|
||||
info := VersionInfo{Valid: len(parts) > 0, Numbers: parts}
|
||||
if prerelease != "" {
|
||||
identifiers := splitPrereleaseIdentifiers(prerelease)
|
||||
if distance, tail, ok := parseGitDescribeIdentifiers(identifiers); ok {
|
||||
info.GitDescribeDistance = distance
|
||||
info.GitDescribeTail = tail
|
||||
} else {
|
||||
info.Prerelease = identifiers
|
||||
}
|
||||
}
|
||||
return info
|
||||
}
|
||||
|
||||
func parseGitDescribeIdentifiers(identifiers []string) (int, []string, bool) {
|
||||
if len(identifiers) < gitDescribeMinIdentifiers {
|
||||
return 0, nil, false
|
||||
}
|
||||
distance, err := strconv.Atoi(strings.TrimSpace(identifiers[0]))
|
||||
if err != nil || distance <= 0 {
|
||||
return 0, nil, false
|
||||
}
|
||||
commitToken := strings.TrimSpace(identifiers[1])
|
||||
if commitToken == "" || !strings.HasPrefix(strings.ToLower(commitToken), "g") {
|
||||
return 0, nil, false
|
||||
}
|
||||
return distance, identifiers[1:], true
|
||||
}
|
||||
|
||||
func splitPrereleaseIdentifiers(value string) []string {
|
||||
parts := strings.FieldsFunc(strings.TrimSpace(value), func(r rune) bool {
|
||||
return r == '.' || r == '-'
|
||||
})
|
||||
filtered := make([]string, 0, len(parts))
|
||||
for _, part := range parts {
|
||||
part = strings.TrimSpace(part)
|
||||
if part != "" {
|
||||
filtered = append(filtered, part)
|
||||
}
|
||||
}
|
||||
return filtered
|
||||
}
|
||||
|
||||
// CompareVersions compares two version strings.
|
||||
// Returns -1 if left < right, 1 if left > right, and 0 if they are equal.
|
||||
func CompareVersions(local, remote string) int {
|
||||
left := ParseVersionInfo(local)
|
||||
right := ParseVersionInfo(remote)
|
||||
if left.IsDev {
|
||||
if right.Valid {
|
||||
return -1
|
||||
}
|
||||
return 0
|
||||
}
|
||||
if !left.Valid || !right.Valid {
|
||||
return 0
|
||||
}
|
||||
|
||||
if result := compareVersionNumbers(left, right); result != 0 {
|
||||
return result
|
||||
}
|
||||
if result := compareGitDescribeDistance(left, right); result != 0 {
|
||||
return result
|
||||
}
|
||||
if left.GitDescribeDistance > 0 || right.GitDescribeDistance > 0 {
|
||||
return compareGitDescribeTails(left, right)
|
||||
}
|
||||
return comparePrereleaseIdentifiers(left, right)
|
||||
}
|
||||
@@ -0,0 +1,108 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package util
|
||||
|
||||
import "strconv"
|
||||
|
||||
func compareVersionNumbers(left, right VersionInfo) int {
|
||||
maxLen := max(len(right.Numbers), len(left.Numbers))
|
||||
for index := range maxLen {
|
||||
leftValue := 0
|
||||
rightValue := 0
|
||||
if index < len(left.Numbers) {
|
||||
leftValue = left.Numbers[index]
|
||||
}
|
||||
if index < len(right.Numbers) {
|
||||
rightValue = right.Numbers[index]
|
||||
}
|
||||
if leftValue < rightValue {
|
||||
return -1
|
||||
}
|
||||
if leftValue > rightValue {
|
||||
return 1
|
||||
}
|
||||
}
|
||||
return 0
|
||||
}
|
||||
|
||||
func compareGitDescribeDistance(left, right VersionInfo) int {
|
||||
if left.GitDescribeDistance == right.GitDescribeDistance {
|
||||
return 0
|
||||
}
|
||||
if left.GitDescribeDistance < right.GitDescribeDistance {
|
||||
return -1
|
||||
}
|
||||
return 1
|
||||
}
|
||||
|
||||
func compareGitDescribeTails(left, right VersionInfo) int {
|
||||
maxLen := max(len(right.GitDescribeTail), len(left.GitDescribeTail))
|
||||
for index := range maxLen {
|
||||
if index >= len(left.GitDescribeTail) {
|
||||
return -1
|
||||
}
|
||||
if index >= len(right.GitDescribeTail) {
|
||||
return 1
|
||||
}
|
||||
if left.GitDescribeTail[index] < right.GitDescribeTail[index] {
|
||||
return -1
|
||||
}
|
||||
if left.GitDescribeTail[index] > right.GitDescribeTail[index] {
|
||||
return 1
|
||||
}
|
||||
}
|
||||
return 0
|
||||
}
|
||||
|
||||
func comparePrereleaseIdentifiers(left, right VersionInfo) int {
|
||||
if len(left.Prerelease) == 0 && len(right.Prerelease) == 0 {
|
||||
return 0
|
||||
}
|
||||
if len(left.Prerelease) == 0 {
|
||||
return 1
|
||||
}
|
||||
if len(right.Prerelease) == 0 {
|
||||
return -1
|
||||
}
|
||||
|
||||
maxLen := max(len(right.Prerelease), len(left.Prerelease))
|
||||
for index := range maxLen {
|
||||
if index >= len(left.Prerelease) {
|
||||
return -1
|
||||
}
|
||||
if index >= len(right.Prerelease) {
|
||||
return 1
|
||||
}
|
||||
if result := comparePrereleasePart(left.Prerelease[index], right.Prerelease[index]); result != 0 {
|
||||
return result
|
||||
}
|
||||
}
|
||||
return 0
|
||||
}
|
||||
|
||||
func comparePrereleasePart(leftPart, rightPart string) int {
|
||||
leftNumber, leftErr := strconv.Atoi(leftPart)
|
||||
rightNumber, rightErr := strconv.Atoi(rightPart)
|
||||
switch {
|
||||
case leftErr == nil && rightErr == nil:
|
||||
if leftNumber < rightNumber {
|
||||
return -1
|
||||
}
|
||||
if leftNumber > rightNumber {
|
||||
return 1
|
||||
}
|
||||
case leftErr == nil:
|
||||
return -1
|
||||
case rightErr == nil:
|
||||
return 1
|
||||
default:
|
||||
if leftPart < rightPart {
|
||||
return -1
|
||||
}
|
||||
if leftPart > rightPart {
|
||||
return 1
|
||||
}
|
||||
}
|
||||
return 0
|
||||
}
|
||||
Reference in New Issue
Block a user