refactor(layout): consolidate backend codebase into backend/ package and clean root directory

- Moved cmd/, core/, plugins/, pkg/, downstream/, and main.go into backend/ directory
- Batch updated all Go source files to import github.com/Rain-kl/Wavelet/backend/...
- Updated Makefile, scripts/swagger.sh, architecture guards, and platform skills
- Passed all quality gates (100% tests, 0 lint issues, clean build)
This commit is contained in:
ryan
2026-08-28 12:56:02 +08:00
parent 33b38f8687
commit 43dc97e48c
319 changed files with 912 additions and 1031 deletions
+69
View File
@@ -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
}
+8
View File
@@ -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")
+264
View File
@@ -0,0 +1,264 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package batchwriter provides a reusable buffered batch writer for high-throughput
// append-only sinks such as ClickHouse. Each business domain should own an independent
// Writer instance with its own queue, flush callback, and tuning parameters.
package batchwriter
import (
"context"
"sync"
"sync/atomic"
"time"
)
// FlushFunc persists a batch of queued items. It is invoked from the worker goroutine.
type FlushFunc[T any] func(ctx context.Context, items []T) error
// FlushErrorHandler is called when FlushFunc returns an error after optional retries.
// The batch is discarded after the handler returns; the worker continues processing.
// Handlers receive the failed items so callers can release dedup keys or re-queue.
type FlushErrorHandler[T any] func(ctx context.Context, items []T, err error)
// Stats is a point-in-time snapshot of Writer queue and failure counters.
type Stats struct {
Name string
Depth int
Cap int
Drops int64
FlushErrors int64
Running bool
}
// Writer buffers items and flushes them by size or interval.
type Writer[T any] struct {
cfg Config
flush FlushFunc[T]
onFlushError FlushErrorHandler[T]
onDrop func(T)
startOnce sync.Once
stopOnce sync.Once
mu sync.RWMutex
ch chan T
workerCtx context.Context
done chan struct{}
drops atomic.Int64
flushErrors atomic.Int64
}
// Option configures optional Writer callbacks.
type Option[T any] func(*Writer[T])
// WithFlushErrorHandler registers a callback for flush failures.
func WithFlushErrorHandler[T any](handler FlushErrorHandler[T]) Option[T] {
return func(w *Writer[T]) {
w.onFlushError = handler
}
}
// WithDropHandler registers a callback when TryEnqueue cannot accept an item.
func WithDropHandler[T any](handler func(T)) Option[T] {
return func(w *Writer[T]) {
w.onDrop = handler
}
}
// New creates a Writer. Call Start before enqueueing items.
func New[T any](cfg Config, flush FlushFunc[T], opts ...Option[T]) (*Writer[T], error) {
if flush == nil {
return nil, errNilFlushFunc
}
if err := cfg.validate(); err != nil {
return nil, err
}
w := &Writer[T]{
cfg: cfg,
flush: flush,
done: make(chan struct{}),
}
for _, opt := range opts {
opt(w)
}
return w, nil
}
// Start launches the background worker. It is safe to call at most once.
func (w *Writer[T]) Start(parent context.Context) {
w.startOnce.Do(func() {
w.mu.Lock()
defer w.mu.Unlock()
w.ch = make(chan T, w.cfg.QueueSize)
w.workerCtx = context.WithoutCancel(parent)
go w.run()
})
}
// Stop closes the queue and waits until the worker drains pending items and exits.
func (w *Writer[T]) Stop(ctx context.Context) error {
w.mu.RLock()
ch := w.ch
done := w.done
w.mu.RUnlock()
if ch == nil {
return nil
}
w.stopOnce.Do(func() {
close(ch)
})
select {
case <-done:
return nil
case <-ctx.Done():
return ctx.Err()
}
}
// Running reports whether Start has been called and Stop has not completed.
func (w *Writer[T]) Running() bool {
w.mu.RLock()
defer w.mu.RUnlock()
if w.ch == nil {
return false
}
select {
case <-w.done:
return false
default:
return true
}
}
// TryEnqueue adds one item without blocking. It returns false when the writer is not
// running or the queue is full.
func (w *Writer[T]) TryEnqueue(item T) bool {
w.mu.RLock()
ch := w.ch
w.mu.RUnlock()
if ch == nil {
w.notifyDrop(item)
return false
}
select {
case ch <- item:
return true
default:
w.notifyDrop(item)
return false
}
}
// IsFull reports whether the queue has no remaining capacity.
func (w *Writer[T]) IsFull() bool {
w.mu.RLock()
defer w.mu.RUnlock()
if w.ch == nil {
return false
}
return len(w.ch) >= cap(w.ch)
}
// Len returns the current queue depth.
func (w *Writer[T]) Len() int {
w.mu.RLock()
defer w.mu.RUnlock()
if w.ch == nil {
return 0
}
return len(w.ch)
}
// Cap returns the queue capacity.
func (w *Writer[T]) Cap() int {
return w.cfg.QueueSize
}
// Stats returns a point-in-time snapshot of queue depth and failure counters.
func (w *Writer[T]) Stats() Stats {
return Stats{
Name: w.cfg.Name,
Depth: w.Len(),
Cap: w.Cap(),
Drops: w.drops.Load(),
FlushErrors: w.flushErrors.Load(),
Running: w.Running(),
}
}
func (w *Writer[T]) run() {
ticker := time.NewTicker(w.cfg.FlushInterval)
defer ticker.Stop()
batch := make([]T, 0, w.cfg.MaxBatchSize)
var batchStartedAt time.Time
flush := func() {
if len(batch) == 0 {
return
}
items := append([]T(nil), batch...)
if err := w.flush(w.workerCtx, items); err != nil {
w.flushErrors.Add(1)
if w.onFlushError != nil {
w.onFlushError(w.workerCtx, items, err)
}
}
batch = batch[:0]
batchStartedAt = time.Time{}
}
defer func() {
flush()
close(w.done)
}()
for {
select {
case item, ok := <-w.ch:
if !ok {
return
}
if len(batch) == 0 {
batchStartedAt = time.Now()
}
batch = append(batch, item)
if len(batch) >= w.cfg.MaxBatchSize {
flush()
}
case <-ticker.C:
if w.shouldFlushOnInterval(len(batch), batchStartedAt, time.Now()) {
flush()
}
}
}
}
func (w *Writer[T]) shouldFlushOnInterval(batchLen int, batchStartedAt time.Time, now time.Time) bool {
if batchLen == 0 {
return false
}
if w.cfg.MinBatchSize == 0 || batchLen >= w.cfg.MinBatchSize {
return true
}
if w.cfg.MaxFlushWait <= 0 || batchStartedAt.IsZero() {
return false
}
return !now.Before(batchStartedAt.Add(w.cfg.MaxFlushWait))
}
func (w *Writer[T]) notifyDrop(item T) {
w.drops.Add(1)
if w.onDrop == nil {
return
}
w.onDrop(item)
}
+317
View File
@@ -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)
}
+12
View File
@@ -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 = ""
)
+391
View File
@@ -0,0 +1,391 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package disk implements a platform-level disk-backed cache with size limit, TTL, and LRU eviction.
package disk
import (
"container/list"
"encoding/binary"
"errors"
"fmt"
"os"
"path/filepath"
"sort"
"sync"
"time"
"github.com/peterbourgon/diskv/v3"
)
// ErrCacheMiss represents a cache miss.
var ErrCacheMiss = errors.New("cache miss")
// Constants for disk cache configuration and sizing
const (
headerSize = 8 // 8 bytes metadata prefix for expiration UnixNano timestamp
defaultMaxSizeMB = 100
defaultTTLMinutes = 60
cacheDirPerm = 0750
// DefaultExpiration applies the cache-wide default TTL.
DefaultExpiration time.Duration = 0
// NoExpiration stores the item without a TTL. Size limits and LRU eviction still apply.
NoExpiration time.Duration = -1
)
// Status represents the runtime cache statistics.
type Status struct {
TotalSize int64 `json:"total_size"`
KeysCount int `json:"keys_count"`
MaxSizeMB int64 `json:"max_size_mb"`
TTLMinutes int64 `json:"ttl_minutes"`
LRUEnabled bool `json:"lru_enabled"`
BasePath string `json:"base_path"`
}
// Cache implements the disk-backed cache with size limits, TTL, and LRU eviction.
type Cache struct {
mu sync.RWMutex
d *diskv.Diskv
basePath string
maxSize int64 // in bytes
defaultTTL time.Duration
lruEnabled bool
// LRU and Size tracking
currentSize int64
items map[string]*list.Element
evictList *list.List
}
type cacheItem struct {
key string
size int64
expiredAt time.Time
}
// New creates a new Cache instance.
func New(basePath string) *Cache {
d := diskv.New(diskv.Options{
BasePath: basePath,
Transform: func(_ string) []string { return []string{} }, // flat structure for easy walk
CacheSizeMax: 1024 * 1024, // 1MB in-memory cache size for diskv itself
})
c := &Cache{
d: d,
basePath: basePath,
maxSize: defaultMaxSizeMB * 1024 * 1024, // 100MB default
defaultTTL: defaultTTLMinutes * time.Minute, // 60 minutes default
lruEnabled: true,
items: make(map[string]*list.Element),
evictList: list.New(),
}
// Scan directory on startup to rebuild LRU and size tracking
_ = c.loadTracker()
return c
}
// Set stores a key-value pair in the cache.
// Use DefaultExpiration for the configured default TTL, NoExpiration for no
// TTL, or a positive duration for a business-specific TTL.
func (c *Cache) Set(key string, value []byte, ttl time.Duration) error {
c.mu.Lock()
defer c.mu.Unlock()
if ttl == DefaultExpiration {
ttl = c.defaultTTL
}
var expiredAt time.Time
if ttl > 0 {
expiredAt = time.Now().Add(ttl)
}
// Prepare data layout: 8 bytes expiration timestamp + raw payload
buf := make([]byte, headerSize+len(value))
var expNano int64
if !expiredAt.IsZero() {
expNano = expiredAt.UnixNano()
}
binary.BigEndian.PutUint64(buf[0:headerSize], uint64(expNano))
copy(buf[headerSize:], value)
// Write to diskv
if err := c.d.Write(key, buf); err != nil {
return fmt.Errorf("failed to write key to disk: %w", err)
}
// Get file size on disk (approximate)
size := int64(len(buf))
// Update memory tracker
if elem, ok := c.items[key]; ok {
item := elem.Value.(*cacheItem)
c.currentSize += size - item.size
item.size = size
item.expiredAt = expiredAt
c.evictList.MoveToFront(elem)
} else {
item := &cacheItem{
key: key,
size: size,
expiredAt: expiredAt,
}
elem := c.evictList.PushFront(item)
c.items[key] = elem
c.currentSize += size
}
// Evict items if size limit exceeded and LRU is enabled
c.evict()
return nil
}
// Get retrieves a key's value from the cache.
func (c *Cache) Get(key string) ([]byte, error) {
c.mu.RLock()
elem, ok := c.items[key]
if !ok {
c.mu.RUnlock()
return nil, ErrCacheMiss
}
item := elem.Value.(*cacheItem)
if !item.expiredAt.IsZero() && time.Now().After(item.expiredAt) {
c.mu.RUnlock()
return c.getAndDeleteIfExpired(key)
}
c.mu.RUnlock()
// Read from disk outside the lock so concurrent cache hits do not serialize on I/O.
data, err := c.d.Read(key)
if err != nil {
c.mu.Lock()
defer c.mu.Unlock()
if _, stillExists := c.items[key]; stillExists {
_ = c.deleteUnlocked(key)
}
return nil, ErrCacheMiss
}
if len(data) < headerSize {
c.mu.Lock()
defer c.mu.Unlock()
if _, stillExists := c.items[key]; stillExists {
_ = c.deleteUnlocked(key)
}
return nil, ErrCacheMiss
}
payload := data[headerSize:]
// Brief write lock only for LRU bookkeeping.
c.mu.Lock()
defer c.mu.Unlock()
elem, ok = c.items[key]
if !ok {
return nil, ErrCacheMiss
}
c.evictList.MoveToFront(elem)
return payload, nil
}
func (c *Cache) getAndDeleteIfExpired(key string) ([]byte, error) {
c.mu.Lock()
defer c.mu.Unlock()
elem, ok := c.items[key]
if !ok {
return nil, ErrCacheMiss
}
item := elem.Value.(*cacheItem)
if !item.expiredAt.IsZero() && time.Now().After(item.expiredAt) {
_ = c.deleteUnlocked(key)
return nil, ErrCacheMiss
}
data, err := c.d.Read(key)
if err != nil {
_ = c.deleteUnlocked(key)
return nil, ErrCacheMiss
}
if len(data) < headerSize {
_ = c.deleteUnlocked(key)
return nil, ErrCacheMiss
}
c.evictList.MoveToFront(elem)
return data[headerSize:], nil
}
// Delete removes a key-value pair from the cache.
func (c *Cache) Delete(key string) error {
c.mu.Lock()
defer c.mu.Unlock()
return c.deleteUnlocked(key)
}
func (c *Cache) deleteUnlocked(key string) error {
if elem, ok := c.items[key]; ok {
item := elem.Value.(*cacheItem)
c.currentSize -= item.size
c.evictList.Remove(elem)
delete(c.items, key)
}
return c.d.Erase(key)
}
// Clear flushes all cached elements.
func (c *Cache) Clear() error {
c.mu.Lock()
defer c.mu.Unlock()
c.currentSize = 0
c.items = make(map[string]*list.Element)
c.evictList.Init()
return c.d.EraseAll()
}
// Status returns the cache status.
func (c *Cache) Status() Status {
c.mu.RLock()
defer c.mu.RUnlock()
return Status{
TotalSize: c.currentSize,
KeysCount: len(c.items),
MaxSizeMB: c.maxSize / (1024 * 1024),
TTLMinutes: int64(c.defaultTTL.Minutes()),
LRUEnabled: c.lruEnabled,
BasePath: c.basePath,
}
}
// UpdatePolicy dynamically updates policies.
func (c *Cache) UpdatePolicy(maxSizeMB int64, ttlMinutes int64, lruEnabled bool) {
c.mu.Lock()
defer c.mu.Unlock()
c.maxSize = maxSizeMB * 1024 * 1024
c.defaultTTL = time.Duration(ttlMinutes) * time.Minute
c.lruEnabled = lruEnabled
c.evict()
}
// evict evicts oldest items if current size exceeds maxSize and LRU is enabled.
func (c *Cache) evict() {
if !c.lruEnabled {
return
}
for c.currentSize > c.maxSize && c.evictList.Len() > 0 {
elem := c.evictList.Back()
item := elem.Value.(*cacheItem)
c.currentSize -= item.size
c.evictList.Remove(elem)
delete(c.items, item.key)
_ = c.d.Erase(item.key)
}
}
// loadTracker scans the cache directory on startup to rebuild memory state.
func (c *Cache) loadTracker() error {
c.mu.Lock()
defer c.mu.Unlock()
// Ensure directory exists
if err := os.MkdirAll(c.basePath, cacheDirPerm); err != nil {
return err
}
type loadedItem struct {
key string
size int64
expiredAt time.Time
modTime time.Time
}
var loadedItems []loadedItem
// Walk keys through diskv
keysChan := c.d.Keys(nil)
for key := range keysChan {
// Read raw bytes to parse expiration prefix
data, err := c.d.Read(key)
if err != nil || len(data) < headerSize {
_ = c.d.Erase(key) // corrupted file, wipe
continue
}
expNano := int64(binary.BigEndian.Uint64(data[0:headerSize])) //nolint:gosec // false positive: UnixNano fits within int64
var expiredAt time.Time
if expNano > 0 {
expiredAt = time.Unix(0, expNano)
}
// Check mod time for ordering
path := filepath.Join(c.basePath, key)
info, err := os.Stat(path)
if err != nil {
continue
}
loadedItems = append(loadedItems, loadedItem{
key: key,
size: int64(len(data)),
expiredAt: expiredAt,
modTime: info.ModTime(),
})
}
// Sort by ModTime ascending (oldest first) so we rebuild LRU correctly
sort.Slice(loadedItems, func(i, j int) bool {
return loadedItems[i].modTime.Before(loadedItems[j].modTime)
})
// Populate LRU (PushFront so that newest items are at the front, oldest at the back)
for _, item := range loadedItems {
entry := &cacheItem{
key: item.key,
size: item.size,
expiredAt: item.expiredAt,
}
element := c.evictList.PushFront(entry)
c.items[item.key] = element
c.currentSize += item.size
}
return nil
}
// StartCleanupWorker periodically cleans up expired cache items.
func (c *Cache) StartCleanupWorker(interval time.Duration) {
ticker := time.NewTicker(interval)
for range ticker.C {
c.cleanExpired()
}
}
// cleanExpired scans memory for expired items and removes them.
func (c *Cache) cleanExpired() {
c.mu.Lock()
defer c.mu.Unlock()
now := time.Now()
for key, elem := range c.items {
item := elem.Value.(*cacheItem)
if !item.expiredAt.IsZero() && now.After(item.expiredAt) {
c.currentSize -= item.size
c.evictList.Remove(elem)
delete(c.items, key)
_ = c.d.Erase(key)
}
}
}
+212
View File
@@ -0,0 +1,212 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package disk
import (
"bytes"
"os"
"testing"
"time"
)
func TestDiskCacheBasic(t *testing.T) {
testDir := "uploads/test_diskcache_basic"
defer func() { _ = os.RemoveAll(testDir) }()
_ = os.RemoveAll(testDir)
c := New(testDir)
defer func() { _ = c.Clear() }()
key := "key1"
val := []byte("value1")
// Get non-existent
_, err := c.Get(key)
if err != ErrCacheMiss {
t.Fatalf("expected ErrCacheMiss, got %v", err)
}
// Set & Get
err = c.Set(key, val, 10*time.Second)
if err != nil {
t.Fatalf("failed to set cache: %v", err)
}
got, err := c.Get(key)
if err != nil {
t.Fatalf("failed to get cache: %v", err)
}
if !bytes.Equal(got, val) {
t.Errorf("expected %s, got %s", val, got)
}
// Delete
err = c.Delete(key)
if err != nil {
t.Fatalf("failed to delete: %v", err)
}
_, err = c.Get(key)
if err != ErrCacheMiss {
t.Errorf("expected ErrCacheMiss after delete, got %v", err)
}
}
func TestDiskCacheTTL(t *testing.T) {
testDir := "uploads/test_diskcache_ttl"
defer func() { _ = os.RemoveAll(testDir) }()
_ = os.RemoveAll(testDir)
c := New(testDir)
defer func() { _ = c.Clear() }()
key := "ttlkey"
val := []byte("ttlval")
// Set with 200ms TTL
err := c.Set(key, val, 200*time.Millisecond)
if err != nil {
t.Fatalf("failed to set: %v", err)
}
// Immediate Get should succeed
got, err := c.Get(key)
if err != nil {
t.Fatalf("failed to get: %v", err)
}
if !bytes.Equal(got, val) {
t.Errorf("expected %s, got %s", val, got)
}
// Sleep 250ms to expire
time.Sleep(250 * time.Millisecond)
// Get should fail with cache miss
_, err = c.Get(key)
if err != ErrCacheMiss {
t.Errorf("expected ErrCacheMiss after TTL expiration, got %v", err)
}
}
func TestDiskCacheExpirationPolicies(t *testing.T) {
testDir := "uploads/test_diskcache_expiration_policies"
defer func() { _ = os.RemoveAll(testDir) }()
_ = os.RemoveAll(testDir)
c := New(testDir)
defer func() { _ = c.Clear() }()
c.defaultTTL = 50 * time.Millisecond
if err := c.Set("default", []byte("default"), DefaultExpiration); err != nil {
t.Fatalf("Set(default, DefaultExpiration) returned error: %v", err)
}
if err := c.Set("custom", []byte("custom"), 100*time.Millisecond); err != nil {
t.Fatalf("Set(custom, 100ms) returned error: %v", err)
}
if err := c.Set("permanent", []byte("permanent"), NoExpiration); err != nil {
t.Fatalf("Set(permanent, NoExpiration) returned error: %v", err)
}
time.Sleep(75 * time.Millisecond)
if _, err := c.Get("default"); err != ErrCacheMiss {
t.Errorf("Get(default) error = %v, want ErrCacheMiss", err)
}
if _, err := c.Get("custom"); err != nil {
t.Errorf("Get(custom) returned error before custom TTL elapsed: %v", err)
}
if _, err := c.Get("permanent"); err != nil {
t.Errorf("Get(permanent) returned error: %v", err)
}
time.Sleep(50 * time.Millisecond)
if _, err := c.Get("custom"); err != ErrCacheMiss {
t.Errorf("Get(custom) error = %v, want ErrCacheMiss", err)
}
if _, err := c.Get("permanent"); err != nil {
t.Errorf("Get(permanent) returned error after other entries expired: %v", err)
}
}
func TestDiskCacheNoExpirationSurvivesReload(t *testing.T) {
testDir := "uploads/test_diskcache_no_expiration_reload"
defer func() { _ = os.RemoveAll(testDir) }()
_ = os.RemoveAll(testDir)
c := New(testDir)
if err := c.Set("permanent", []byte("value"), NoExpiration); err != nil {
t.Fatalf("Set(permanent, NoExpiration) returned error: %v", err)
}
reloaded := New(testDir)
defer func() { _ = reloaded.Clear() }()
got, err := reloaded.Get("permanent")
if err != nil {
t.Fatalf("reloaded Get(permanent) returned error: %v", err)
}
if !bytes.Equal(got, []byte("value")) {
t.Errorf("reloaded Get(permanent) = %q, want %q", got, "value")
}
}
func TestDiskCacheLRUEviction(t *testing.T) {
testDir := "uploads/test_diskcache_lru"
defer func() { _ = os.RemoveAll(testDir) }()
_ = os.RemoveAll(testDir)
c := New(testDir)
defer func() { _ = c.Clear() }()
// Force a very small max size of 20 bytes for testing (8 bytes header + payload)
// So 2 items of 2 bytes payload = 2 * (8 + 2) = 20 bytes max.
c.maxSize = 20
c.lruEnabled = true
// Write item 1: 8 + 2 = 10 bytes
err := c.Set("k1", []byte("v1"), DefaultExpiration)
if err != nil {
t.Fatalf("failed to set k1: %v", err)
}
// Write item 2: 8 + 2 = 10 bytes
err = c.Set("k2", []byte("v2"), DefaultExpiration)
if err != nil {
t.Fatalf("failed to set k2: %v", err)
}
// Both should exist
if _, err := c.Get("k1"); err != nil {
t.Errorf("k1 should exist: %v", err)
}
if _, err := c.Get("k2"); err != nil {
t.Errorf("k2 should exist: %v", err)
}
// Write item 3: 8 + 2 = 10 bytes -> total size would be 30, exceeding 20.
// This should evict the oldest item. Since k1 was accessed, but then k2 was accessed,
// wait, let's access k1 again to make it the most recently used, so k2 becomes oldest!
_, _ = c.Get("k1") // k1 is now MRU, k2 is LRU
err = c.Set("k3", []byte("v3"), DefaultExpiration)
if err != nil {
t.Fatalf("failed to set k3: %v", err)
}
// k2 should be evicted, k1 and k3 should exist
_, err = c.Get("k2")
if err != ErrCacheMiss {
t.Errorf("expected k2 to be evicted, got error %v", err)
}
if _, err := c.Get("k1"); err != nil {
t.Errorf("k1 should still exist: %v", err)
}
if _, err := c.Get("k3"); err != nil {
t.Errorf("k3 should exist: %v", err)
}
}
+73
View File
@@ -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()
}
+38
View File
@@ -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())
}
}
+233
View File
@@ -0,0 +1,233 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package ram
import (
"context"
"errors"
"sync"
"time"
)
var (
// ErrNotFound is returned by the Loader when the requested item is not found.
ErrNotFound = errors.New("cache item not found in data source")
managerCache *Cache[string, map[string]cacheEntry]
writeLocks = make(map[string]*sync.Mutex)
writeLocksMu sync.Mutex
)
// CacheItem represents a unified cache entity.
type CacheItem struct {
Key string `json:"key"`
Value string `json:"value"`
Type string `json:"type"`
TTL time.Duration `json:"ttl"` // -1 means never expire
}
// Loader is an interface that the cache client must implement to handle database retrieval.
type Loader interface {
LoadAll(ctx context.Context, configType string) ([]CacheItem, error)
LoadOne(ctx context.Context, configType string, key string) (CacheItem, error)
}
type cacheEntry struct {
item CacheItem
expireAt time.Time
}
func init() {
// Initialize with a large maximum size since it only stores one entry per configType
managerCache = MustNew[string, map[string]cacheEntry](Options{
MaximumSize: 1000,
})
}
func getWriteLock(configType string) *sync.Mutex {
writeLocksMu.Lock()
defer writeLocksMu.Unlock()
lock, found := writeLocks[configType]
if !found {
lock = &sync.Mutex{}
writeLocks[configType] = lock
}
return lock
}
// Get retrieves a cache item from the local cache store, checking for expiration.
// Reads are completely lock-free because maps stored in Otter are immutable.
func Get(configType string, key string) (CacheItem, bool) {
m, ok := managerCache.GetIfPresent(configType)
if !ok {
return CacheItem{}, false
}
entry, found := m[key]
if !found {
return CacheItem{}, false
}
// Check expiration
if entry.item.TTL != -1 && !entry.expireAt.IsZero() && time.Now().After(entry.expireAt) {
// Asynchronously remove the expired item from the map and write back
go deleteKeyIfExpired(configType, key, entry.expireAt)
return CacheItem{}, false
}
return entry.item, true
}
func deleteKeyIfExpired(configType string, key string, expireAt time.Time) {
lock := getWriteLock(configType)
lock.Lock()
defer lock.Unlock()
currentMap, ok := managerCache.GetIfPresent(configType)
if !ok {
return
}
entry, found := currentMap[key]
if !found {
return
}
// Double-check expiration time to ensure we don't delete a newly updated key
if entry.expireAt != expireAt || !time.Now().After(entry.expireAt) {
return
}
newMap := make(map[string]cacheEntry, len(currentMap)-1)
for k, v := range currentMap {
if k != key {
newMap[k] = v
}
}
managerCache.Set(configType, newMap)
}
// Set stores a cache item in the local cache store.
// Writes are protected by a fine-grained lock per configType.
func Set(item CacheItem) {
lock := getWriteLock(item.Type)
lock.Lock()
defer lock.Unlock()
currentMap, ok := managerCache.GetIfPresent(item.Type)
newMap := make(map[string]cacheEntry)
if ok {
for k, v := range currentMap {
newMap[k] = v
}
}
var expireAt time.Time
if item.TTL != -1 {
expireAt = time.Now().Add(item.TTL)
}
newMap[item.Key] = cacheEntry{
item: item,
expireAt: expireAt,
}
managerCache.Set(item.Type, newMap)
}
// Delete removes a single item from the local cache store.
// Writes are protected by a fine-grained lock per configType.
func Delete(configType string, key string) {
lock := getWriteLock(configType)
lock.Lock()
defer lock.Unlock()
currentMap, ok := managerCache.GetIfPresent(configType)
if !ok {
return
}
newMap := make(map[string]cacheEntry, len(currentMap))
for k, v := range currentMap {
if k != key {
newMap[k] = v
}
}
managerCache.Set(configType, newMap)
}
// UpdateTypeItems replaces all cache items of a specific type atomically.
// Writes are protected by a fine-grained lock per configType.
func UpdateTypeItems(configType string, items []CacheItem) {
lock := getWriteLock(configType)
lock.Lock()
defer lock.Unlock()
newMap := make(map[string]cacheEntry, len(items))
for _, item := range items {
var expireAt time.Time
if item.TTL != -1 {
expireAt = time.Now().Add(item.TTL)
}
newMap[item.Key] = cacheEntry{
item: item,
expireAt: expireAt,
}
}
managerCache.Set(configType, newMap)
}
// GetTypeItems retrieves all unexpired cache items of a specific type.
// Reads are completely lock-free because maps stored in Otter are immutable.
func GetTypeItems(configType string) []CacheItem {
currentMap, ok := managerCache.GetIfPresent(configType)
if !ok {
return nil
}
var list []CacheItem
for _, entry := range currentMap {
if entry.item.TTL == -1 || entry.expireAt.IsZero() || time.Now().Before(entry.expireAt) {
list = append(list, entry.item)
}
}
return list
}
// Refresh reloads configuration cache from database via the Loader.
func Refresh(ctx context.Context, configType string, key string, loader Loader) error {
if configType == "" {
return errors.New("type is required")
}
if key != "" {
// Single key refresh: first fetch latest value from database
item, err := loader.LoadOne(ctx, configType, key)
if err != nil {
if errors.Is(err, ErrNotFound) {
Delete(configType, key)
return nil
}
return err
}
Set(item)
return nil
}
// All keys refresh: load all of that type from database first, then replace cache
items, err := loader.LoadAll(ctx, configType)
if err != nil {
return err
}
UpdateTypeItems(configType, items)
return nil
}
// ResetForTest clears the local store and locks.
func ResetForTest() {
writeLocksMu.Lock()
writeLocks = make(map[string]*sync.Mutex)
writeLocksMu.Unlock()
managerCache.InvalidateAll()
}
+254
View File
@@ -0,0 +1,254 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package config 负责应用配置的加载、解析与环境变量覆盖。
package config
import (
"encoding/json"
"flag"
"log"
"os"
"strconv"
"strings"
"github.com/spf13/viper"
)
// 默认队列优先级
const (
webhookQueuePriority = 10
whitelistQueuePriority = 5
defaultQueuePriority = 3
)
// Config 全局配置单例,初始化后不可变
var Config *configModel
// findConfigPath searches upward for the config file to handle tests running in subdirectories.
func findConfigPath(configPath string) string {
if _, err := os.Stat(configPath); err == nil {
return configPath
}
dir := "."
for i := 0; i < 5; i++ {
dir += "/.."
path := dir + "/" + configPath
if _, err := os.Stat(path); err == nil {
return path
}
}
return configPath
}
// isTest checks if the current execution context is within 'go test'.
func isTest() bool {
if flag.Lookup("test.v") != nil {
return true
}
for _, arg := range os.Args {
if strings.HasPrefix(arg, "-test.") || strings.HasSuffix(arg, ".test") {
return true
}
}
return false
}
func init() {
// 加载配置文件路径
configPath := os.Getenv("CONFIG_PATH")
if configPath == "" {
configPath = findConfigPath("config.yaml")
}
// 设置配置文件
viper.SetConfigFile(configPath)
viper.AutomaticEnv()
// 读取配置文件(可选:找不到文件时使用空默认值 + 环境变量)
if err := viper.ReadInConfig(); err != nil {
if _, ok := err.(viper.ConfigFileNotFoundError); !ok {
// 文件存在但读取/解析失败
if _, statErr := os.Stat(configPath); statErr == nil { //nolint:gosec // configPath is loaded from CONFIG_PATH environment variable
log.Fatalf("[Config] read config failed: %v\n", err)
}
}
log.Println("[Config] no config file found, using environment variables only")
viper.SetConfigType("yaml")
if err := viper.ReadConfig(strings.NewReader("")); err != nil {
log.Fatalf("[Config] failed to init empty config: %v\n", err)
}
}
// 解析配置到结构体
var c configModel
if err := viper.Unmarshal(&c); err != nil {
log.Fatalf("[Config] parse config failed: %v\n", err)
}
applyDefaults(&c)
// 环境变量覆盖(优先级高于 config.yaml)
applyEnvOverrides(&c)
applyDefaults(&c)
// Disable standard DB/Redis initializations during tests to prevent connection attempts.
if isTest() {
c.Database.Enabled = false
c.Database.SQLitePath = ":memory:"
c.Redis.Enabled = false
c.ClickHouse.Enabled = false
}
// 设置全局配置
Config = &c
// 打印配置
printConfig(&c)
}
func applyDefaults(c *configModel) {
if c.App.SessionAge <= 0 {
c.App.SessionAge = 86400
}
if c.Otel.TracerName == "" {
c.Otel.TracerName = "github.com/Rain-kl/Wavelet"
}
}
// ─── 环境变量覆盖层 ────────────────────────────────────────────────────────────
// 环境变量优先级高于 config.yaml,未设置则保留 yaml 中的值。
func envStr(key, fallback string) string {
if v, ok := os.LookupEnv(key); ok {
return v
}
return fallback
}
func envInt(key string, fallback int) int {
if v, ok := os.LookupEnv(key); ok {
if n, err := strconv.Atoi(v); err == nil {
return n
}
}
return fallback
}
func envInt64(key string, fallback int64) int64 {
if v, ok := os.LookupEnv(key); ok {
if n, err := strconv.ParseInt(v, 10, 64); err == nil {
return n
}
}
return fallback
}
func envFloat64(key string, fallback float64) float64 {
if v, ok := os.LookupEnv(key); ok {
if n, err := strconv.ParseFloat(v, 64); err == nil {
return n
}
}
return fallback
}
func envBool(key string, fallback bool) bool {
if v, ok := os.LookupEnv(key); ok {
if b, err := strconv.ParseBool(v); err == nil {
return b
}
}
return fallback
}
// applyEnvOverrides 将环境变量值覆盖到配置结构体上(仅当环境变量已设置时生效)
func applyEnvOverrides(c *configModel) {
// ─── App ───
c.App.AppName = envStr("APP_NAME", c.App.AppName)
c.App.Env = envStr("APP_ENV", c.App.Env)
c.App.Addr = envStr("APP_ADDR", c.App.Addr)
c.App.NodeID = envInt64("APP_NODE_ID", c.App.NodeID)
c.App.APIPrefix = envStr("APP_API_PREFIX", c.App.APIPrefix)
c.App.GracefulShutdownTimeout = envInt("APP_GRACEFUL_SHUTDOWN_TIMEOUT", c.App.GracefulShutdownTimeout)
c.App.SessionCookieName = envStr("APP_SESSION_COOKIE_NAME", c.App.SessionCookieName)
c.App.SessionSecret = envStr("APP_SESSION_SECRET", c.App.SessionSecret)
c.App.SessionDomain = envStr("APP_SESSION_DOMAIN", c.App.SessionDomain)
c.App.SessionAge = envInt("APP_SESSION_AGE", c.App.SessionAge)
c.App.SessionHTTPOnly = envBool("APP_SESSION_HTTP_ONLY", c.App.SessionHTTPOnly)
c.App.SessionSecure = envBool("APP_SESSION_SECURE", c.App.SessionSecure)
// ─── Database ───
c.Database.Host = envStr("DB_HOST", c.Database.Host)
c.Database.Port = envInt("DB_PORT", c.Database.Port)
c.Database.Username = envStr("DB_USERNAME", c.Database.Username)
c.Database.Password = envStr("DB_PASSWORD", c.Database.Password)
c.Database.Database = envStr("DB_NAME", c.Database.Database)
c.Database.SSLMode = envStr("DB_SSL_MODE", c.Database.SSLMode)
c.Database.TimeZone = envStr("DB_TIMEZONE", c.Database.TimeZone)
c.Database.LogLevel = envStr("DB_LOG_LEVEL", c.Database.LogLevel)
c.Database.MaxIdleConn = envInt("DB_MAX_IDLE_CONN", c.Database.MaxIdleConn)
c.Database.MaxOpenConn = envInt("DB_MAX_OPEN_CONN", c.Database.MaxOpenConn)
// 当 DB_HOST 环境变量已设置时自动启用数据库
if _, ok := os.LookupEnv("DB_HOST"); ok {
c.Database.Enabled = true
}
c.Database.Enabled = envBool("DB_ENABLED", c.Database.Enabled)
c.Database.SQLitePath = envStr("SQLITE_PATH", c.Database.SQLitePath)
// ─── Redis ───
if v, ok := os.LookupEnv("REDIS_ADDR"); ok {
c.Redis.Addrs = []string{v}
c.Redis.Enabled = true // 当 REDIS_ADDR 已设置时自动启用
}
c.Redis.Enabled = envBool("REDIS_ENABLED", c.Redis.Enabled)
c.Redis.Username = envStr("REDIS_USERNAME", c.Redis.Username)
c.Redis.Password = envStr("REDIS_PASSWORD", c.Redis.Password)
c.Redis.DB = envInt("REDIS_DB", c.Redis.DB)
c.Redis.KeyPrefix = envStr("REDIS_KEY_PREFIX", c.Redis.KeyPrefix)
c.Redis.PoolSize = envInt("REDIS_POOL_SIZE", c.Redis.PoolSize)
c.Redis.MaintNotifications = envBool("REDIS_MAINT_NOTIFICATIONS", c.Redis.MaintNotifications)
// ─── ClickHouse ───
if v, ok := os.LookupEnv("CLICKHOUSE_HOST"); ok {
c.ClickHouse.Hosts = []string{v}
c.ClickHouse.Enabled = true
}
c.ClickHouse.Enabled = envBool("CLICKHOUSE_ENABLED", c.ClickHouse.Enabled)
c.ClickHouse.Username = envStr("CLICKHOUSE_USERNAME", c.ClickHouse.Username)
c.ClickHouse.Password = envStr("CLICKHOUSE_PASSWORD", c.ClickHouse.Password)
c.ClickHouse.Database = envStr("CLICKHOUSE_NAME", c.ClickHouse.Database)
// ─── Log ───
c.Log.Level = envStr("LOG_LEVEL", c.Log.Level)
c.Log.Format = envStr("LOG_FORMAT", c.Log.Format)
c.Log.Output = envStr("LOG_OUTPUT", c.Log.Output)
// ─── OTel ───
c.Otel.SamplingRate = envFloat64("OTEL_SAMPLING_RATE", c.Otel.SamplingRate)
c.Otel.TracerName = envStr("OTEL_TRACER_NAME", c.Otel.TracerName)
// ─── Worker ───
c.Worker.Concurrency = envInt("WORKER_CONCURRENCY", c.Worker.Concurrency)
c.Worker.StrictPriority = envBool("WORKER_STRICT_PRIORITY", c.Worker.StrictPriority)
// 无 yaml 且无环境变量时,使用代码级默认队列
if len(c.Worker.Queues) == 0 {
c.Worker.Queues = []QueueConfig{
{Name: "webhook", Priority: webhookQueuePriority},
{Name: "whitelist_only", Priority: whitelistQueuePriority},
{Name: "default", Priority: defaultQueuePriority},
}
}
}
// printConfig 打印配置内容
func printConfig(c *configModel) {
configJSON, err := json.MarshalIndent(c, "", " ")
if err != nil {
log.Printf("[Config] failed to marshal config: %v\n", err)
return
}
log.Printf("[Config] loaded configuration:\n%s\n", string(configJSON))
}
+17
View File
@@ -0,0 +1,17 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package config
import "testing"
func TestApplyEnvOverridesRedisMaintNotifications(t *testing.T) {
t.Setenv("REDIS_MAINT_NOTIFICATIONS", "true")
cfg := &configModel{}
applyEnvOverrides(cfg)
if !cfg.Redis.MaintNotifications {
t.Fatal("REDIS_MAINT_NOTIFICATIONS=true was not applied")
}
}
+142
View File
@@ -0,0 +1,142 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package config
import "time"
type configModel struct {
App appConfig `mapstructure:"app"`
Database databaseConfig `mapstructure:"database"`
Redis redisConfig `mapstructure:"redis"`
Log logConfig `mapstructure:"log"`
Scheduler schedulerConfig `mapstructure:"scheduler"`
Worker workerConfig `mapstructure:"worker"`
ClickHouse clickHouseConfig `mapstructure:"clickhouse"`
Otel otelConfig `mapstructure:"otel"`
}
// appConfig 应用基本配置
type appConfig struct {
AppName string `mapstructure:"app_name"`
Env string `mapstructure:"env"`
Addr string `mapstructure:"addr"`
NodeID int64 `mapstructure:"node_id"`
APIPrefix string `mapstructure:"api_prefix"`
GracefulShutdownTimeout int `mapstructure:"graceful_shutdown_timeout"`
SessionCookieName string `mapstructure:"session_cookie_name"`
SessionSecret string `mapstructure:"session_secret"`
SessionDomain string `mapstructure:"session_domain"`
SessionAge int `mapstructure:"session_age"`
SessionHTTPOnly bool `mapstructure:"session_http_only"`
SessionSecure bool `mapstructure:"session_secure"`
}
// IsProduction 检查当前环境是否为生产环境
func (a *appConfig) IsProduction() bool {
return a.Env == "production"
}
// databaseConfig 数据库配置
type databaseConfig struct {
Enabled bool `mapstructure:"enabled"`
SQLitePath string `mapstructure:"sqlite_path"` // PostgreSQL 禁用时的 SQLite 文件路径
Host string `mapstructure:"host"`
Port int `mapstructure:"port"`
Username string `mapstructure:"username"`
Password string `mapstructure:"password"`
Database string `mapstructure:"database"`
MaxIdleConn int `mapstructure:"max_idle_conn"`
MaxOpenConn int `mapstructure:"max_open_conn"`
ConnMaxLifetime int `mapstructure:"conn_max_lifetime"`
ConnMaxIdleTime int `mapstructure:"conn_max_idle_time"`
LogLevel string `mapstructure:"log_level"`
SSLMode string `mapstructure:"ssl_mode"`
TimeZone string `mapstructure:"time_zone"`
ApplicationName string `mapstructure:"application_name"`
SearchPath string `mapstructure:"search_path"`
PreferSimpleProtocol bool `mapstructure:"prefer_simple_protocol"`
StatementCacheCapacity int `mapstructure:"statement_cache_capacity"`
DefaultQueryExecMode string `mapstructure:"default_query_exec_mode"`
Replicas []databaseReplicaConfig `mapstructure:"replicas"`
SlowThreshold time.Duration `mapstructure:"slow_threshold"`
}
// databaseReplicaConfig 只读副本配置
type databaseReplicaConfig struct {
Host string `mapstructure:"host"`
Port int `mapstructure:"port"`
Username string `mapstructure:"username"`
Password string `mapstructure:"password"`
}
// clickhouse 配置
type clickHouseConfig struct {
Enabled bool `mapstructure:"enabled"`
Hosts []string `mapstructure:"hosts"`
Username string `mapstructure:"username"`
Password string `mapstructure:"password"`
Database string `mapstructure:"database"`
MaxIdleConn int `mapstructure:"max_idle_conn"`
MaxOpenConn int `mapstructure:"max_open_conn"`
ConnMaxLifetime int `mapstructure:"conn_max_lifetime"`
DialTimeout int `mapstructure:"dial_timeout"`
BlockBufferSize uint8 `mapstructure:"block_buffer_size"`
}
// redisConfig Redis配置
type redisConfig struct {
Enabled bool `mapstructure:"enabled"`
Addrs []string `mapstructure:"addrs"`
Username string `mapstructure:"username"`
Password string `mapstructure:"password"`
DB int `mapstructure:"db"`
ClusterMode bool `mapstructure:"cluster_mode"`
MasterName string `mapstructure:"master_name"`
KeyPrefix string `mapstructure:"key_prefix"`
PoolSize int `mapstructure:"pool_size"`
MinIdleConn int `mapstructure:"min_idle_conn"`
DialTimeout int `mapstructure:"dial_timeout"`
ReadTimeout int `mapstructure:"read_timeout"`
WriteTimeout int `mapstructure:"write_timeout"`
MaxRetries int `mapstructure:"max_retries"`
PoolTimeout int `mapstructure:"pool_timeout"`
ConnMaxIdleTime int `mapstructure:"conn_max_idle_time"`
MaintNotifications bool `mapstructure:"maint_notifications"`
}
// logConfig 日志配置
type logConfig struct {
Level string `mapstructure:"level"`
Format string `mapstructure:"format"`
Output string `mapstructure:"output"`
FilePath string `mapstructure:"file_path"`
MaxSize int `mapstructure:"max_size"`
MaxAge int `mapstructure:"max_age"`
MaxBackups int `mapstructure:"max_backups"`
Compress bool `mapstructure:"compress"`
}
// schedulerConfig 定时任务配置
type schedulerConfig struct {
}
// workerConfig 工作配置
type workerConfig struct {
Concurrency int `mapstructure:"concurrency"`
StrictPriority bool `mapstructure:"strict_priority"`
Queues []QueueConfig `mapstructure:"queues"`
}
// QueueConfig 队列配置
type QueueConfig struct {
Name string `mapstructure:"name"`
Priority int `mapstructure:"priority"`
}
// otelConfig OpenTelemetry 配置
type otelConfig struct {
SamplingRate float64 `mapstructure:"sampling_rate"`
TracerName string `mapstructure:"tracer_name"`
}
+106
View File
@@ -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(),
}
}
+100
View File
@@ -0,0 +1,100 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package httppool
import (
"context"
"crypto/tls"
"io"
"net"
"net/http"
"net/http/httptest"
"testing"
"time"
)
func TestDefaultTransport(t *testing.T) {
tr1 := DefaultTransport()
if tr1 == nil {
t.Fatal("DefaultTransport() returned nil")
}
tr2 := DefaultTransport()
if tr1 != tr2 {
t.Error("DefaultTransport() did not return a singleton instance")
}
}
func TestNewClient(t *testing.T) {
timeout := 15 * time.Second
client := NewClient(timeout)
if client == nil {
t.Fatal("NewClient() returned nil")
}
if client.Timeout != timeout {
t.Errorf("NewClient() timeout = %v, want %v", client.Timeout, timeout)
}
if client.Transport != DefaultTransport() {
t.Error("NewClient() is not configured with the default transport")
}
}
func TestNewTransportUsesConfiguredDirectDialer(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, _ *http.Request) {
_, _ = writer.Write([]byte("ok"))
}))
t.Cleanup(server.Close)
var dialedAddress string
dialer := &net.Dialer{}
transport := NewTransport(TransportOptions{
Proxy: nil,
DialContext: func(ctx context.Context, network string, address string) (net.Conn, error) {
dialedAddress = address
return dialer.DialContext(ctx, network, server.Listener.Addr().String())
},
})
client := &http.Client{Transport: transport}
t.Cleanup(client.CloseIdleConnections)
request, err := http.NewRequestWithContext(t.Context(), http.MethodGet, "http://artifact.example/site.zip", nil)
if err != nil {
t.Fatalf("NewRequestWithContext() error = %v", err)
}
response, err := client.Do(request)
if err != nil {
t.Fatalf("client.Do() error = %v", err)
}
defer func() { _ = response.Body.Close() }()
if _, err := io.ReadAll(response.Body); err != nil {
t.Fatalf("ReadAll() error = %v", err)
}
if dialedAddress != "artifact.example:80" {
t.Fatalf("DialContext address = %q, want direct target", dialedAddress)
}
}
func TestNewTransportClonesTLSConfig(t *testing.T) {
server := httptest.NewTLSServer(http.HandlerFunc(func(writer http.ResponseWriter, _ *http.Request) {
writer.WriteHeader(http.StatusNoContent)
}))
t.Cleanup(server.Close)
tlsConfig := &tls.Config{InsecureSkipVerify: true} //nolint:gosec // test-only self-signed server
client := &http.Client{Transport: NewTransport(TransportOptions{TLSClientConfig: tlsConfig})}
t.Cleanup(client.CloseIdleConnections)
tlsConfig.InsecureSkipVerify = false
request, err := http.NewRequestWithContext(t.Context(), http.MethodGet, server.URL, nil)
if err != nil {
t.Fatalf("NewRequestWithContext() error = %v", err)
}
response, err := client.Do(request)
if err != nil {
t.Fatalf("client.Do() error = %v", err)
}
_ = response.Body.Close()
}
+46
View File
@@ -0,0 +1,46 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package idgen 提供分布式 ID 生成器
package idgen
import (
"fmt"
"log"
"github.com/Rain-kl/Wavelet/backend/pkg/config"
"github.com/bwmarrin/snowflake"
)
// 2025-12-01 00:00:00 UTC 的毫秒时间戳
const epoch int64 = 1764547200000
const maxNegativeIDRetries = 3
var node *snowflake.Node
func init() {
snowflake.Epoch = epoch
nodeID := config.Config.App.NodeID
var err error
node, err = snowflake.NewNode(nodeID)
if err != nil {
log.Fatalf("[Snowflake] init failed: %v\n", err)
}
log.Printf("[Snowflake] initialized with node ID: %d, epoch: 2025-12-01\n", nodeID)
}
// NextUint64ID 生成下一个分布式唯一 ID。
// 理论上不应出现负值;若出现则最多重试 maxNegativeIDRetries 次,仍失败则 panic。
func NextUint64ID() uint64 {
for attempt := 1; attempt <= maxNegativeIDRetries; attempt++ {
id := node.Generate().Int64()
if id >= 0 {
return uint64(id)
}
log.Printf("[Snowflake] generated negative ID: %d (attempt %d/%d)", id, attempt, maxNegativeIDRetries)
}
panic(fmt.Sprintf("[Snowflake] generated negative ID after %d attempts", maxNegativeIDRetries))
}
+15
View File
@@ -0,0 +1,15 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package idgen
import (
"testing"
"github.com/stretchr/testify/assert"
)
func TestNextUint64ID(t *testing.T) {
id := NextUint64ID()
assert.NotZero(t, id)
}
+9
View File
@@ -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"
)
+102
View File
@@ -0,0 +1,102 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package logger
import (
"context"
"fmt"
"log"
"github.com/uptrace/opentelemetry-go-extra/otelzap"
"go.uber.org/zap"
"go.uber.org/zap/zapcore"
)
// Config represents the logging configuration.
type Config struct {
Level string
Format string
Output string
FilePath string
MaxSize int
MaxAge int
MaxBackups int
Compress bool
}
var logger *otelzap.Logger
// ringBufferCapacity 环形缓冲区容量
const ringBufferCapacity = 5000
// GlobalRingBuffer 全局日志环形缓冲区,供 Admin 日志查询和 WebSocket 推送使用
var GlobalRingBuffer *LogRingBuffer
func doInit(cfg Config) {
logWriter, err := getLogWriterForConfig(cfg)
if err != nil {
log.Fatalf("[Logger] get log writer err: %v\n", err)
}
// 初始化 ring buffer(保留最近 5000 行日志),如果是多次调用 Init,不需要重复创建 GlobalRingBuffer
if GlobalRingBuffer == nil {
GlobalRingBuffer = NewLogRingBuffer(ringBufferCapacity)
}
// 使用 multi writer 同时写入原始输出和 ring buffer
multiWriter := zapcore.NewMultiWriteSyncer(
logWriter,
zapcore.AddSync(GlobalRingBuffer),
)
zapLogger := zap.New(
zapcore.NewCore(getEncoderForConfig(cfg), multiWriter, getLogLevelForConfig(cfg)),
zap.AddCaller(),
zap.AddCallerSkip(1),
)
logger = otelzap.New(
zapLogger,
otelzap.WithMinLevel(zapLogger.Level()),
)
}
func init() {
// 默认使用 console stdout INFO 日志输出,避免在 Init 前或测试中发生空指针崩溃
defaultCfg := Config{
Level: "info",
Format: "console",
Output: "stdout",
}
doInit(defaultCfg)
}
// Init initializes the logger with a custom configuration.
func Init(cfg Config) {
doInit(cfg)
}
// DebugF 输出 Debug 级别日志
func DebugF(ctx context.Context, format string, args ...interface{}) {
msg := fmt.Sprintf(format, args...)
logger.Ctx(ctx).Debug(msg, getTraceIDFields(ctx)...)
}
// InfoF 输出 Info 级别日志
func InfoF(ctx context.Context, format string, args ...interface{}) {
msg := fmt.Sprintf(format, args...)
logger.Ctx(ctx).Info(msg, getTraceIDFields(ctx)...)
}
// WarnF 输出 Warn 级别日志
func WarnF(ctx context.Context, format string, args ...interface{}) {
msg := fmt.Sprintf(format, args...)
logger.Ctx(ctx).Warn(msg, getTraceIDFields(ctx)...)
}
// ErrorF 输出 Error 级别日志
func ErrorF(ctx context.Context, format string, args ...interface{}) {
msg := fmt.Sprintf(format, args...)
logger.Ctx(ctx).Error(msg, getTraceIDFields(ctx)...)
}
+175
View File
@@ -0,0 +1,175 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package logger
import (
"io"
"sync"
)
// LogEntry 日志条目,对应 ring buffer 中的一行日志
type LogEntry struct {
Index int `json:"index"` // 全局递增序号
Data string `json:"data"` // 一行日志原文(含换行符)
}
// LogRingBuffer 固定容量的环形缓冲区,存储最近的日志行
// 支持:追加日志、按 cursor 分页查询、订阅实时推送
type LogRingBuffer struct {
mu sync.RWMutex
entries []LogEntry
cap int
head int // 下一条写入的位置
count int // 当前条目数
seq int // 全局递增序号
subscribers map[chan LogEntry]struct{}
subMu sync.RWMutex
}
// NewLogRingBuffer 创建指定容量的日志环形缓冲区
func NewLogRingBuffer(capacity int) *LogRingBuffer {
return &LogRingBuffer{
entries: make([]LogEntry, capacity),
cap: capacity,
subscribers: make(map[chan LogEntry]struct{}),
}
}
// Write 实现 io.Writer 接口,供 zapcore.WriteSyncer 调用
// 按 '\n' 分割为独立行写入 ring buffer
func (r *LogRingBuffer) Write(p []byte) (int, error) {
if len(p) == 0 {
return 0, nil
}
data := string(p)
start := 0
for i := 0; i < len(data); i++ {
if data[i] == '\n' {
line := data[start:i]
start = i + 1
if len(line) > 0 {
r.appendLine(line)
}
}
}
// 处理最后一行(没有换行符结尾的情况)
if start < len(data) && len(data[start:]) > 0 {
r.appendLine(data[start:])
}
return len(p), nil
}
// Sync 实现 zapcore.WriteSyncer 接口
func (r *LogRingBuffer) Sync() error {
return nil
}
// appendLine 追加一行日志到 ring buffer 并通知订阅者
func (r *LogRingBuffer) appendLine(line string) {
r.mu.Lock()
entry := LogEntry{
Index: r.seq,
Data: line,
}
r.entries[r.head] = entry
r.head = (r.head + 1) % r.cap
if r.count < r.cap {
r.count++
}
r.seq++
r.mu.Unlock()
// 异步通知订阅者
r.subMu.RLock()
for ch := range r.subscribers {
select {
case ch <- entry:
default:
// 订阅者消费太慢,丢弃(避免阻塞日志写入)
}
}
r.subMu.RUnlock()
}
// Query 查询历史日志
// cursor=0 表示查询最新日志,cursor>0 表示查询 index < cursor 的更早日志
// limit 为返回条数上限
// 返回日志条目(按 index 升序)和是否有更早的日志
func (r *LogRingBuffer) Query(cursor int, limit int) ([]LogEntry, bool) {
r.mu.RLock()
defer r.mu.RUnlock()
if r.count == 0 {
return nil, false
}
// 计算 ring buffer 中有效条目的范围
// oldest index in ring: head - count (wrapping)
oldestPos := (r.head - r.count + r.cap) % r.cap
// 将 ring buffer 中的有效条目按顺序收集
ordered := make([]LogEntry, 0, r.count)
for i := 0; i < r.count; i++ {
pos := (oldestPos + i) % r.cap
ordered = append(ordered, r.entries[pos])
}
if cursor == 0 {
// 查询最新日志:返回最后 limit 条
if len(ordered) <= limit {
return ordered, false
}
return ordered[len(ordered)-limit:], true
}
// 查询 index < cursor 的更早日志
// 找到 index < cursor 的条目
var cut int
for cut = len(ordered); cut > 0; cut-- {
if ordered[cut-1].Index < cursor {
break
}
}
if cut == 0 {
return nil, false
}
// 返回 cut 之前的最后 limit 条
start := cut - limit
if start < 0 {
start = 0
}
hasMore := start > 0
return ordered[start:cut], hasMore
}
// subscribeChanSize 订阅者 channel 缓冲区大小
const subscribeChanSize = 64
// Subscribe 订阅实时日志推送
// 返回一个 channel,调用者应 defer Unsubscribe
func (r *LogRingBuffer) Subscribe() chan LogEntry {
ch := make(chan LogEntry, subscribeChanSize)
r.subMu.Lock()
r.subscribers[ch] = struct{}{}
r.subMu.Unlock()
return ch
}
// Unsubscribe 取消订阅
func (r *LogRingBuffer) Unsubscribe(ch chan LogEntry) {
r.subMu.Lock()
delete(r.subscribers, ch)
r.subMu.Unlock()
close(ch)
}
// 确保 LogRingBuffer 实现 io.Writer 接口
var _ io.Writer = (*LogRingBuffer)(nil)
+192
View File
@@ -0,0 +1,192 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package logger
import (
"testing"
"github.com/stretchr/testify/assert"
)
func TestLogRingBuffer_WriteAndQuery(t *testing.T) {
rb := NewLogRingBuffer(5)
// Write some logs
_, _ = rb.Write([]byte("line1\nline2\nline3\n"))
entries, hasMore := rb.Query(0, 10)
assert.False(t, hasMore)
assert.Equal(t, 3, len(entries))
assert.Equal(t, "line1", entries[0].Data)
assert.Equal(t, "line2", entries[1].Data)
assert.Equal(t, "line3", entries[2].Data)
assert.Equal(t, 0, entries[0].Index)
assert.Equal(t, 1, entries[1].Index)
assert.Equal(t, 2, entries[2].Index)
}
func TestLogRingBuffer_CapacityOverflow(t *testing.T) {
rb := NewLogRingBuffer(3)
_, _ = rb.Write([]byte("a\nb\nc\nd\ne\n"))
entries, hasMore := rb.Query(0, 10)
assert.False(t, hasMore)
assert.Equal(t, 3, len(entries))
assert.Equal(t, "c", entries[0].Data)
assert.Equal(t, "d", entries[1].Data)
assert.Equal(t, "e", entries[2].Data)
}
func TestLogRingBuffer_QueryLatest(t *testing.T) {
rb := NewLogRingBuffer(10)
_, _ = rb.Write([]byte("a\nb\nc\nd\ne\n"))
// Query latest 2
entries, hasMore := rb.Query(0, 2)
assert.True(t, hasMore)
assert.Equal(t, 2, len(entries))
assert.Equal(t, "d", entries[0].Data)
assert.Equal(t, "e", entries[1].Data)
}
func TestLogRingBuffer_QueryByCursor(t *testing.T) {
rb := NewLogRingBuffer(10)
_, _ = rb.Write([]byte("a\nb\nc\nd\ne\n"))
// First get all to find indices
all, _ := rb.Query(0, 10)
assert.Equal(t, 5, len(all))
// Query entries before index 3
entries, hasMore := rb.Query(3, 10)
assert.False(t, hasMore)
assert.Equal(t, 3, len(entries))
assert.Equal(t, "a", entries[0].Data)
assert.Equal(t, "b", entries[1].Data)
assert.Equal(t, "c", entries[2].Data)
}
func TestLogRingBuffer_QueryByCursorWithLimit(t *testing.T) {
rb := NewLogRingBuffer(10)
_, _ = rb.Write([]byte("a\nb\nc\nd\ne\n"))
// Query 2 entries before index 4
entries, hasMore := rb.Query(4, 2)
assert.True(t, hasMore)
assert.Equal(t, 2, len(entries))
assert.Equal(t, "c", entries[0].Data)
assert.Equal(t, "d", entries[1].Data)
}
func TestLogRingBuffer_QueryEmpty(t *testing.T) {
rb := NewLogRingBuffer(5)
entries, hasMore := rb.Query(0, 10)
assert.False(t, hasMore)
assert.Nil(t, entries)
}
func TestLogRingBuffer_QueryNonExistentCursor(t *testing.T) {
rb := NewLogRingBuffer(5)
_, _ = rb.Write([]byte("a\nb\n"))
entries, hasMore := rb.Query(999, 10)
assert.False(t, hasMore)
assert.Equal(t, 2, len(entries))
assert.Equal(t, "a", entries[0].Data)
assert.Equal(t, "b", entries[1].Data)
}
func TestLogRingBuffer_Subscribe(t *testing.T) {
rb := NewLogRingBuffer(5)
ch := rb.Subscribe()
defer rb.Unsubscribe(ch)
_, _ = rb.Write([]byte("hello\n"))
entry := <-ch
assert.Equal(t, "hello", entry.Data)
assert.Equal(t, 0, entry.Index)
}
func TestLogRingBuffer_SubscribeMultiple(t *testing.T) {
rb := NewLogRingBuffer(5)
ch1 := rb.Subscribe()
defer rb.Unsubscribe(ch1)
ch2 := rb.Subscribe()
defer rb.Unsubscribe(ch2)
_, _ = rb.Write([]byte("msg\n"))
e1 := <-ch1
e2 := <-ch2
assert.Equal(t, "msg", e1.Data)
assert.Equal(t, "msg", e2.Data)
}
func TestLogRingBuffer_WriteNoNewline(t *testing.T) {
rb := NewLogRingBuffer(5)
_, _ = rb.Write([]byte("partial"))
entries, _ := rb.Query(0, 10)
assert.Equal(t, 1, len(entries))
assert.Equal(t, "partial", entries[0].Data)
}
func TestLogRingBuffer_WriteEmpty(t *testing.T) {
rb := NewLogRingBuffer(5)
n, err := rb.Write([]byte(""))
assert.Equal(t, 0, n)
assert.NoError(t, err)
entries, _ := rb.Query(0, 10)
assert.Nil(t, entries)
}
func TestLogRingBuffer_QueryAfterOverflow(t *testing.T) {
rb := NewLogRingBuffer(3)
_, _ = rb.Write([]byte("1\n2\n3\n4\n5\n6\n7\n"))
entries, hasMore := rb.Query(0, 10)
assert.False(t, hasMore)
assert.Equal(t, 3, len(entries))
assert.Equal(t, "5", entries[0].Data)
assert.Equal(t, "6", entries[1].Data)
assert.Equal(t, "7", entries[2].Data)
// Query by cursor - index 4 is "5", so cursor=4 should return index < 4
older, hasMore2 := rb.Query(4, 10)
assert.False(t, hasMore2)
assert.Nil(t, older)
}
func TestLogRingBuffer_NextCursor(t *testing.T) {
rb := NewLogRingBuffer(10)
_, _ = rb.Write([]byte("a\nb\nc\nd\ne\n"))
// Query latest 2, should return next_cursor pointing to first returned entry
entries, _ := rb.Query(0, 2)
assert.Equal(t, 2, len(entries))
// entries[0].Index = 3 ("d"), entries[1].Index = 4 ("e")
assert.Equal(t, 3, entries[0].Index)
// Now use that index as cursor to get older entries
older, hasMore := rb.Query(entries[0].Index, 10)
assert.False(t, hasMore)
assert.Equal(t, 3, len(older))
assert.Equal(t, "a", older[0].Data)
assert.Equal(t, "b", older[1].Data)
assert.Equal(t, "c", older[2].Data)
}
+99
View File
@@ -0,0 +1,99 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package logger
import (
"context"
"fmt"
"log"
"os"
"path/filepath"
"go.opentelemetry.io/otel/trace"
"go.uber.org/zap"
"go.uber.org/zap/zapcore"
"gopkg.in/natefinch/lumberjack.v2"
)
// logDirPerm 日志目录权限
const logDirPerm = 0750
func getLogWriterForConfig(cfg Config) (zapcore.WriteSyncer, error) {
if cfg.Output == "file" {
// 初始化日志目录
logPath := cfg.FilePath
logDir := filepath.Dir(logPath)
if err := os.MkdirAll(logDir, logDirPerm); err != nil {
return nil, fmt.Errorf(errCreateLogFileDirFailed, err)
}
// 配置日志轮转
logOutput := &lumberjack.Logger{
Filename: logPath,
MaxSize: cfg.MaxSize,
MaxBackups: cfg.MaxBackups,
MaxAge: cfg.MaxAge,
Compress: cfg.Compress,
}
return zapcore.AddSync(logOutput), nil
}
return zapcore.AddSync(os.Stdout), nil
}
// getEncoderForConfig 获取日志编码器
func getEncoderForConfig(cfg Config) zapcore.Encoder {
// 编码器配置
encoderConfig := zapcore.EncoderConfig{
TimeKey: "time",
LevelKey: "level",
NameKey: "logger",
CallerKey: "caller",
MessageKey: "msg",
StacktraceKey: "stacktrace",
LineEnding: zapcore.DefaultLineEnding,
EncodeLevel: zapcore.LowercaseLevelEncoder,
EncodeTime: zapcore.ISO8601TimeEncoder,
EncodeDuration: zapcore.SecondsDurationEncoder,
EncodeCaller: zapcore.ShortCallerEncoder,
}
if cfg.Format == "json" {
return zapcore.NewJSONEncoder(encoderConfig)
}
return zapcore.NewConsoleEncoder(encoderConfig)
}
// getLogLevelForConfig 获取日志级别
func getLogLevelForConfig(cfg Config) zapcore.Level {
level := cfg.Level
switch level {
case "debug":
return zapcore.DebugLevel
case "info":
return zapcore.InfoLevel
case "warn":
return zapcore.WarnLevel
case "error":
return zapcore.ErrorLevel
default:
log.Printf("[Logger] invalid log level: %s, defaulting to info\n", level)
return zapcore.InfoLevel
}
}
func getTraceIDFields(ctx context.Context) []zap.Field {
span := trace.SpanFromContext(ctx)
spanContext := span.SpanContext()
if !spanContext.IsValid() {
return nil
}
return []zap.Field{
zap.String("traceID", spanContext.TraceID().String()),
zap.String("spanID", spanContext.SpanID().String()),
}
}
+16
View File
@@ -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"
)
+250
View File
@@ -0,0 +1,250 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package mail
import (
"bytes"
"context"
"crypto/tls"
"fmt"
"net"
"net/smtp"
"strconv"
"strings"
"time"
)
const (
smtpSSLPort = 465 // SMTP SSL 端口
smtpDialTimeout = 5 * time.Second // SMTP 连接超时
smtpSessionDeadline = 10 * time.Second // SMTP 会话截止时间
)
// Config represents SMTP mail configuration
type Config struct {
Host string
Port int
Username string
Password string
}
// sanitizeHeaderValue removes CR/LF bytes so untrusted values cannot inject
// additional email headers (email header injection).
func sanitizeHeaderValue(v string) string {
v = strings.ReplaceAll(v, "\r", "")
v = strings.ReplaceAll(v, "\n", "")
return v
}
// SendMail sends an HTML email using the provided config and message details
func SendMail(ctx context.Context, cfg Config, to string, subject, body string) error {
return SendMailHTML(ctx, cfg, to, subject, body)
}
// SendMailHTML sends an HTML format email
func SendMailHTML(ctx context.Context, cfg Config, to string, subject, body string) error {
addr := net.JoinHostPort(cfg.Host, strconv.Itoa(cfg.Port))
// Header & MIME settings for HTML email
header := make(map[string]string)
header["From"] = sanitizeHeaderValue(cfg.Username)
header["To"] = sanitizeHeaderValue(to)
header["Subject"] = sanitizeHeaderValue(subject)
header["MIME-Version"] = "1.0"
header["Content-Type"] = "text/html; charset=UTF-8"
message := ""
for k, v := range header {
message += fmt.Sprintf("%s: %s\r\n", k, v)
}
message += "\r\n" + body
auth := smtp.PlainAuth("", cfg.Username, cfg.Password, cfg.Host)
// If using SSL port 465, we connection via TLS dial
if cfg.Port == smtpSSLPort {
return sendMailViaSSL(ctx, addr, auth, cfg, to, message)
}
// For standard port (587 / 25), use smtp.SendMail directly (handles STARTTLS automatically if server supports it)
err := smtp.SendMail(addr, auth, cfg.Username, []string{to}, []byte(message))
if err != nil {
return fmt.Errorf(errSendMailFailed, err)
}
return nil
}
// sendMailViaSSL 通过 TLS 直接连接 SMTP SSL 端口发送邮件
func sendMailViaSSL(ctx context.Context, addr string, auth smtp.Auth, cfg Config, to, message string) error {
tlsConfig := &tls.Config{
InsecureSkipVerify: true, //nolint:gosec // SMTP servers might use self-signed certificates
ServerName: cfg.Host,
}
dialer := &net.Dialer{Timeout: smtpDialTimeout}
tlsDialer := &tls.Dialer{
NetDialer: dialer,
Config: tlsConfig,
}
conn, err := tlsDialer.DialContext(ctx, "tcp", addr)
if err != nil {
return fmt.Errorf(errDialTLSFailed, err)
}
defer func() { _ = conn.Close() }()
_ = conn.SetDeadline(time.Now().Add(smtpSessionDeadline))
client, err := smtp.NewClient(conn, cfg.Host)
if err != nil {
return fmt.Errorf(errSMTPClientCreationFailed, err)
}
defer func() { _ = client.Close() }()
if err = client.Auth(auth); err != nil {
return fmt.Errorf(errSMTPAuthFailed, err)
}
if err = client.Mail(cfg.Username); err != nil {
return fmt.Errorf(errSMTPMailCommandFailed, err)
}
if err = client.Rcpt(to); err != nil {
return fmt.Errorf(errSMTPRcptCommandFailed, err)
}
w, err := client.Data()
if err != nil {
return fmt.Errorf(errSMTPDataCommandFailed, err)
}
defer func() { _ = w.Close() }()
_, err = w.Write([]byte(message))
if err != nil {
return fmt.Errorf(errSMTPWritingBodyFailed, err)
}
return nil
}
// SendMailWithLog sends a test email and records a detailed SMTP connection log
func SendMailWithLog(ctx context.Context, cfg Config, to string, subject, body string) (string, error) {
var logBuf bytes.Buffer
logLine := func(dir string, format string, args ...interface{}) {
fmt.Fprintf(&logBuf, "[%s] %s\n", dir, fmt.Sprintf(format, args...))
}
addr := net.JoinHostPort(cfg.Host, strconv.Itoa(cfg.Port))
logLine("System", "Connecting to %s...", addr)
var conn net.Conn
var err error
dialer := &net.Dialer{Timeout: smtpDialTimeout}
if cfg.Port == smtpSSLPort {
tlsConfig := &tls.Config{
InsecureSkipVerify: true, //nolint:gosec // SMTP servers might use self-signed certificates
ServerName: cfg.Host,
}
tlsDialer := &tls.Dialer{
NetDialer: dialer,
Config: tlsConfig,
}
conn, err = tlsDialer.DialContext(ctx, "tcp", addr)
} else {
conn, err = dialer.DialContext(ctx, "tcp", addr)
}
if err != nil {
logLine("Error", "Connection failed: %v", err)
return logBuf.String(), err
}
defer func() { _ = conn.Close() }()
logLine("System", "Connected successfully.")
// Set a 10-second session deadline for read/write operations
_ = conn.SetDeadline(time.Now().Add(smtpSessionDeadline))
client, err := smtp.NewClient(conn, cfg.Host)
if err != nil {
logLine("Error", "SMTP client handshake failed: %v", err)
return logBuf.String(), err
}
defer func() { _ = client.Close() }()
// If not 465, support STARTTLS if available
if cfg.Port != smtpSSLPort {
if ok, _ := client.Extension("STARTTLS"); ok {
logLine("C", "STARTTLS")
tlsConfig := &tls.Config{
InsecureSkipVerify: true, //nolint:gosec // SMTP servers might use self-signed certificates
ServerName: cfg.Host,
}
if err = client.StartTLS(tlsConfig); err != nil {
logLine("Error", "STARTTLS failed: %v", err)
return logBuf.String(), err
}
logLine("S", "220 Ready to start TLS")
}
}
// Authentication
if cfg.Username != "" && cfg.Password != "" {
auth := smtp.PlainAuth("", cfg.Username, cfg.Password, cfg.Host)
logLine("C", "AUTH PLAIN **********")
if err = client.Auth(auth); err != nil {
logLine("Error", "Authentication failed: %v", err)
return logBuf.String(), err
}
logLine("S", "235 Authentication successful")
}
// Mail command
logLine("C", "MAIL FROM:<%s>", cfg.Username)
if err = client.Mail(cfg.Username); err != nil {
logLine("Error", "MAIL FROM command failed: %v", err)
return logBuf.String(), err
}
logLine("S", "250 OK")
// Rcpt command
logLine("C", "RCPT TO:<%s>", to)
if err = client.Rcpt(to); err != nil {
logLine("Error", "RCPT TO command failed: %v", err)
return logBuf.String(), err
}
logLine("S", "250 OK")
// Data command
logLine("C", "DATA")
w, err := client.Data()
if err != nil {
logLine("Error", "DATA command failed: %v", err)
return logBuf.String(), err
}
logLine("S", "354 Start mail input")
// Header & MIME settings for HTML email
header := make(map[string]string)
header["From"] = sanitizeHeaderValue(cfg.Username)
header["To"] = sanitizeHeaderValue(to)
header["Subject"] = sanitizeHeaderValue(subject)
header["MIME-Version"] = "1.0"
header["Content-Type"] = "text/html; charset=UTF-8"
message := ""
for k, v := range header {
message += fmt.Sprintf("%s: %s\r\n", k, v)
}
message += "\r\n" + body
logLine("System", "Sending message body...")
if _, err = w.Write([]byte(message)); err != nil {
_ = w.Close()
logLine("Error", "Writing message body failed: %v", err)
return logBuf.String(), err
}
_ = w.Close()
logLine("S", "250 OK")
logLine("C", "QUIT")
_ = client.Quit()
logLine("System", "Mail sent successfully!")
return logBuf.String(), nil
}
+112
View File
@@ -0,0 +1,112 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package mail
import (
"bufio"
"context"
"net"
"net/textproto"
"testing"
)
func TestSendMailMock(t *testing.T) {
// Start a mock SMTP server
l, err := net.Listen("tcp", "127.0.0.1:0")
if err != nil {
t.Fatalf("failed to start mock smtp server: %v", err)
}
defer func() { _ = l.Close() }()
port := l.Addr().(*net.TCPAddr).Port
go func() {
conn, err := l.Accept()
if err != nil {
return
}
defer func() { _ = conn.Close() }()
writer := bufio.NewWriter(conn)
reader := bufio.NewReader(conn)
tp := textproto.NewReader(reader)
// 220 Ready
_, _ = writer.WriteString("220 mock.smtp.com SMTP Ready\r\n")
_ = writer.Flush()
// Read HELO/EHLO
_, _ = tp.ReadLine()
_, _ = writer.WriteString("250-mock.smtp.com\r\n250 AUTH PLAIN\r\n")
_ = writer.Flush()
// Read AUTH PLAIN
_, _ = tp.ReadLine()
_, _ = writer.WriteString("235 Authentication successful\r\n")
_ = writer.Flush()
// Read MAIL FROM
_, _ = tp.ReadLine()
_, _ = writer.WriteString("250 OK\r\n")
_ = writer.Flush()
// Read RCPT TO
_, _ = tp.ReadLine()
_, _ = writer.WriteString("250 OK\r\n")
_ = writer.Flush()
// Read DATA
_, _ = tp.ReadLine()
_, _ = writer.WriteString("354 Start mail input\r\n")
_ = writer.Flush()
// Read body lines until dot
for {
line, err := tp.ReadLine()
if err != nil || line == "." {
break
}
}
_, _ = writer.WriteString("250 OK\r\n")
_ = writer.Flush()
// Read QUIT
_, _ = tp.ReadLine()
_, _ = writer.WriteString("221 Bye\r\n")
_ = writer.Flush()
}()
cfg := Config{
Host: "127.0.0.1",
Port: port,
Username: "test@example.com",
Password: "password",
}
err = SendMail(context.Background(), cfg, "recipient@example.com", "Test Subject", "<h1>Test Body</h1>")
if err != nil {
t.Errorf("failed to send mail: %v", err)
}
}
func TestSanitizeHeaderValue(t *testing.T) {
tests := []struct {
name string
input string
want string
}{
{"plain", "System Notification", "System Notification"},
{"crlf stripped", "alert\r\nBcc: attacker@example.com", "alertBcc: attacker@example.com"},
{"cr stripped", "a\rb", "ab"},
{"lf stripped", "a\nb", "ab"},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
if got := sanitizeHeaderValue(tt.input); got != tt.want {
t.Errorf("sanitizeHeaderValue(%q) = %q, want %q", tt.input, got, tt.want)
}
})
}
}
+45
View File
@@ -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)
}
+40
View File
@@ -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("内部系统错误"))
}
}
+168
View File
@@ -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())
}
+57
View File
@@ -0,0 +1,57 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package response provides shared HTTP API response structures.
package response
import "github.com/gin-gonic/gin"
// Response 通用响应体
type Response[T any] struct {
ErrorMsg string `json:"error_msg"`
Data T `json:"data"`
}
// Any 用于 Swagger 文档的响应类型(非泛型)
// swag 不支持泛型,使用此类型替代 Response[T]
type Any struct {
ErrorMsg string `json:"error_msg" example:""`
Data interface{} `json:"data"`
}
// APIError 统一的 API 业务错误类型,可被全局错误处理中间件捕获
type APIError struct {
Code int
Msg string
}
func (e *APIError) Error() string {
return e.Msg
}
// NewError 实例化一个 APIError
func NewError(code int, msg string) *APIError {
return &APIError{Code: code, Msg: msg}
}
// AbortWithError 将 API 错误挂载到 Gin Context 并中断执行流
func AbortWithError(c *gin.Context, code int, msg string) {
_ = c.Error(NewError(code, msg))
c.Abort()
}
// OK 构造成功响应
func OK[T any](data T) Response[T] {
return Response[T]{Data: data}
}
// OKNil 构造成功响应(data 为 null)
func OKNil() Response[any] {
return Response[any]{Data: nil}
}
// Err 构造错误响应
func Err(msg string) Response[any] {
return Response[any]{ErrorMsg: msg, Data: nil}
}
+17
View File
@@ -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()
}
}
+20
View File
@@ -0,0 +1,20 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package testhelper
import (
"github.com/Rain-kl/Wavelet/backend/pkg/response"
"github.com/gin-gonic/gin"
)
// NewTestGinEngine 创建带 ErrorHandlerMiddleware 的 Gin 引擎,与生产环境错误响应行为一致。
func NewTestGinEngine(middlewares ...gin.HandlerFunc) *gin.Engine {
gin.SetMode(gin.TestMode)
r := gin.New()
r.Use(response.ErrorHandlerMiddleware())
for _, middleware := range middlewares {
r.Use(middleware)
}
return r
}
+333
View File
@@ -0,0 +1,333 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package testhelper 提供测试辅助工具
package testhelper
import (
"context"
"testing"
"time"
cachepkg "github.com/Rain-kl/Wavelet/backend/plugins/infra/cache"
db "github.com/Rain-kl/Wavelet/backend/plugins/infra/database"
"github.com/alicebob/miniredis/v2"
"github.com/glebarez/sqlite"
"github.com/redis/go-redis/v9"
"github.com/redis/go-redis/v9/maintnotifications"
"gorm.io/gorm"
)
// SystemConfig 测试用系统配置表
type SystemConfig struct {
Key string `gorm:"primaryKey;size:64;not null"`
Value string `gorm:"type:text;not null"`
Type string `gorm:"size:32;not null"`
Visibility string `gorm:"size:32;not null;default:'hidden'"`
Description string `gorm:"size:255"`
}
// TableName 返回测试配置表表名
func (SystemConfig) TableName() string {
return "w_system_configs"
}
type userHelper struct {
ID uint64 `gorm:"primaryKey;autoIncrement"`
Username string `gorm:"size:64;uniqueIndex;not null"`
Nickname string `gorm:"size:64;not null;default:''"`
Password string `gorm:"size:255;not null;default:''"`
Email string `gorm:"size:128;index;default:''"`
AvatarURL string `gorm:"size:255;default:''"`
IsAdmin bool `gorm:"default:false;not null"`
IsActive bool `gorm:"default:true;not null"`
NeedChangePassword bool `gorm:"default:false;not null"`
Bio string `gorm:"size:500;default:''"`
Phone string `gorm:"size:32;default:''"`
Gender string `gorm:"size:16;default:''"`
Website string `gorm:"size:255;default:''"`
Location string `gorm:"size:255;default:''"`
LastLoginAt time.Time
CreatedAt time.Time `gorm:"autoCreateTime"`
UpdatedAt time.Time `gorm:"autoUpdateTime"`
}
func (userHelper) TableName() string { return "w_users" }
type accessTokenHelper struct {
ID uint64 `gorm:"primaryKey;autoIncrement"`
UserID uint64 `gorm:"not null;index"`
Name string `gorm:"size:64;not null"`
TokenHash string `gorm:"size:64;uniqueIndex;not null"`
MaskedToken string `gorm:"size:32;not null"`
IsAdmin bool `gorm:"default:false;not null"`
CreatedAt time.Time `gorm:"autoCreateTime"`
UpdatedAt time.Time `gorm:"autoUpdateTime"`
}
func (accessTokenHelper) TableName() string { return "w_access_tokens" }
type authSourceHelper struct {
ID uint64 `gorm:"primaryKey;autoIncrement"`
Type string `gorm:"size:32;not null"`
Name string `gorm:"size:64;not null"`
Enabled bool `gorm:"default:false;not null"`
CreatedAt time.Time `gorm:"autoCreateTime"`
UpdatedAt time.Time `gorm:"autoUpdateTime"`
}
func (authSourceHelper) TableName() string { return "w_auth_sources" }
type externalAccountHelper struct {
ID uint64 `gorm:"primaryKey;autoIncrement"`
UserID uint64 `gorm:"not null;index"`
AuthSourceType string `gorm:"size:32;not null;index"`
ExternalID string `gorm:"size:128;not null;index"`
Username string `gorm:"size:128;default:''"`
Email string `gorm:"size:128;default:''"`
CreatedAt time.Time `gorm:"autoCreateTime"`
UpdatedAt time.Time `gorm:"autoUpdateTime"`
}
func (externalAccountHelper) TableName() string { return "w_external_accounts" }
type uploadHelper struct {
ID uint64 `gorm:"primaryKey;autoIncrement"`
UserID uint64 `gorm:"not null;index"`
FileName string `gorm:"size:255;not null"`
FilePath string `gorm:"size:500;not null"`
FileSize int64 `gorm:"not null"`
MimeType string `gorm:"size:128;not null"`
Extension string `gorm:"size:32;not null"`
Hash string `gorm:"size:64;index;not null;default:''"`
Type string `gorm:"size:50;not null;index"`
Status string `gorm:"size:20;not null;default:'pending'"`
AccessMode int `gorm:"not null;default:0"`
Metadata string `gorm:"type:text"`
CreatedAt time.Time `gorm:"autoCreateTime;index"`
UpdatedAt time.Time `gorm:"autoUpdateTime"`
}
func (uploadHelper) TableName() string { return "w_uploads" }
type uploadStatHelper struct {
ID uint64 `gorm:"primaryKey;autoIncrement"`
Dimension string `gorm:"size:32;not null;uniqueIndex:idx_stat_dimension_key"`
StatKey string `gorm:"size:100;not null;uniqueIndex:idx_stat_dimension_key"`
FileCount int64 `gorm:"not null;default:0"`
FileSize int64 `gorm:"not null;default:0"`
UpdatedAt time.Time `gorm:"autoUpdateTime"`
}
func (uploadStatHelper) TableName() string { return "w_upload_stats" }
type messageChannelHelper struct {
ID uint64 `gorm:"primaryKey;autoIncrement"`
Type string `gorm:"size:32;not null"`
Name string `gorm:"size:64;not null"`
OwnerScope string `gorm:"size:32;not null;default:'system'"`
OwnerID *uint64
Credentials string `gorm:"type:text;not null"`
Extra string `gorm:"type:text"`
Enabled bool `gorm:"default:false;not null"`
CreatedAt time.Time `gorm:"autoCreateTime"`
UpdatedAt time.Time `gorm:"autoUpdateTime"`
}
func (messageChannelHelper) TableName() string { return "w_message_channels" }
type messageBindingHelper struct {
ID uint64 `gorm:"primaryKey;autoIncrement"`
ChannelID uint64 `gorm:"not null;index"`
PlatformUserID string `gorm:"size:128;not null;index"`
UserID uint64 `gorm:"not null;index"`
CreatedAt time.Time `gorm:"autoCreateTime"`
}
func (messageBindingHelper) TableName() string { return "w_message_bindings" }
type messagePairingHelper struct {
ID uint64 `gorm:"primaryKey;autoIncrement"`
Code string `gorm:"size:32;uniqueIndex;not null"`
ChannelID uint64 `gorm:"not null;index"`
PlatformUserID string `gorm:"size:128;not null;index"`
UserID uint64 `gorm:"not null;index"`
ExpiresAt time.Time `gorm:"not null;index"`
CreatedAt time.Time `gorm:"autoCreateTime"`
}
func (messagePairingHelper) TableName() string { return "w_message_pairing_codes" }
type taskExecutionHelper struct {
ID uint64 `gorm:"primaryKey;autoIncrement"`
TaskID string `gorm:"size:128;uniqueIndex;not null"`
TaskType string `gorm:"size:64;index;not null"`
TaskName string `gorm:"size:128"`
Status string `gorm:"size:32;index;not null"`
Retryable bool `gorm:"not null;default:false"`
MaxRetry int `gorm:"not null;default:0"`
RetryCount int `gorm:"not null;default:0"`
Log string `gorm:"type:text"`
ErrorMessage string `gorm:"type:text"`
Result string `gorm:"type:text"`
StartedAt *time.Time `gorm:"index"`
FinishedAt *time.Time
Duration int64
Payload string `gorm:"type:text"`
TriggeredBy string `gorm:"size:32;not null;default:system"`
CreatedAt time.Time `gorm:"autoCreateTime;index"`
UpdatedAt time.Time `gorm:"autoUpdateTime"`
}
func (taskExecutionHelper) TableName() string {
return "w_task_executions"
}
type scheduleHelper struct {
ID uint64 `gorm:"primaryKey;autoIncrement"`
TaskType string `gorm:"size:64;uniqueIndex;not null"`
TaskName string `gorm:"size:128;not null"`
CronExpr string `gorm:"size:64;not null"`
Payload string `gorm:"type:text"`
Enabled bool `gorm:"default:true;not null"`
CreatedAt time.Time `gorm:"autoCreateTime"`
UpdatedAt time.Time `gorm:"autoUpdateTime"`
}
func (scheduleHelper) TableName() string { return "w_schedules" }
const (
configTypeSystem = "system"
configTypeBusiness = "business"
configValueTrue = "true"
configValueFalse = "false"
)
// SetupTestEnvironment initializes an in-memory SQLite DB, seeds default configurations,
// starts miniredis, and overrides the global db/Redis clients. It returns a cleanup function.
func SetupTestEnvironment(t *testing.T) (*gorm.DB, *miniredis.Miniredis, func()) {
sqliteDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{
DisableForeignKeyConstraintWhenMigrating: true,
})
if err != nil {
t.Fatalf("failed to open in-memory SQLite db: %v", err)
}
if sqlDB, err := sqliteDB.DB(); err == nil {
sqlDB.SetMaxOpenConns(1)
}
// AutoMigrate all tables via internal test helpers to completely decouple testhelper from domain plugins
err = sqliteDB.AutoMigrate(
&userHelper{},
&accessTokenHelper{},
&authSourceHelper{},
&externalAccountHelper{},
&SystemConfig{},
&uploadHelper{},
&uploadStatHelper{},
&taskExecutionHelper{},
&scheduleHelper{},
&messageChannelHelper{},
&messageBindingHelper{},
&messagePairingHelper{},
)
if err != nil {
t.Fatalf("failed to auto migrate tables: %v", err)
}
mr, err := miniredis.Run()
if err != nil {
t.Fatalf("failed to start miniredis: %v", err)
}
redisClient := redis.NewClient(&redis.Options{
Addr: mr.Addr(),
MaintNotificationsConfig: &maintnotifications.Config{
Mode: maintnotifications.ModeDisabled,
},
})
db.SetDB(sqliteDB)
cachepkg.Redis = redisClient
seedDefaultConfigs(t, sqliteDB)
cleanup := func() {
runExtraCleanups()
_ = redisClient.Close()
mr.Close()
db.SetDB(nil)
cachepkg.Redis = nil
}
return sqliteDB, mr, cleanup
}
func getSeedConfigsPart1() []SystemConfig {
return []SystemConfig{
{Key: "upload_allowed_extensions", Value: `["jpg", "jpeg", "png", "gif", "webp", "txt", "pdf", "zip"]`, Type: configTypeSystem, Description: "允许上传的文件扩展名列表(JSON 字符串数组)"},
{Key: "site_name", Value: "Wavelet", Type: configTypeSystem, Description: "站点名称"},
{Key: "site_description", Value: "Lightweight and Modular Web Application Platform", Type: configTypeSystem, Description: "站点描述"},
{Key: "password_login_enabled", Value: configValueTrue, Type: configTypeSystem, Description: "是否开启账号密码登录(true/false)"},
{Key: "registration_enabled", Value: configValueTrue, Type: configTypeSystem, Description: "是否允许新用户注册(全局总开关,true/false)"},
{Key: "password_register_enabled", Value: configValueTrue, Type: configTypeSystem, Description: "是否允许账号密码注册(true/false)"},
{Key: "oidc_login_enabled", Value: configValueFalse, Type: configTypeSystem, Description: "是否开启 OIDC 登录(true/false)"},
{Key: "server_address", Value: "http://localhost:8000", Type: configTypeSystem, Description: "服务端访问地址(用于生成绝对路径链接,多个地址用英文逗号分隔)"},
{Key: "smtp_host", Value: "", Type: configTypeSystem, Description: "SMTP 服务器主机名或 IP"},
{Key: "smtp_port", Value: "587", Type: configTypeSystem, Description: "SMTP 服务器端口(标准 STARTTLS 为 587,SMTPS 为 465)"},
{Key: "smtp_username", Value: "", Type: configTypeSystem, Description: "SMTP 账户(如 sender@example.com)"},
{Key: "smtp_password", Value: "", Type: configTypeSystem, Description: "SMTP 访问凭证(授权码/密码)"},
{Key: "email_login_verification_enabled", Value: configValueFalse, Type: configTypeSystem, Description: "是否开启邮箱登录验证(true/false)"},
{Key: "email_register_verification_enabled", Value: configValueFalse, Type: configTypeSystem, Description: "是否开启邮箱注册验证(true/false)"},
{Key: "menu_display_config", Value: "{}", Type: configTypeSystem, Description: "目录显示配置(JSON 字符串,格式为 {url: enabled})"},
{Key: "search_engine_indexing_enabled", Value: configValueFalse, Type: configTypeSystem, Description: "是否允许搜索引擎检索"},
{Key: "file_access_whitelist", Value: `["avatar"]`, Type: configTypeSystem, Description: "免登录访问的文件业务类型白名单"},
{Key: "disk_cache_max_size_mb", Value: "100", Type: configTypeSystem, Description: "磁盘缓存最大空间大小 (MB)"},
{Key: "disk_cache_ttl_minutes", Value: "60", Type: configTypeSystem, Description: "磁盘缓存默认有效期 (分钟)"},
{Key: "disk_cache_lru_enabled", Value: configValueTrue, Type: configTypeSystem, Description: "是否启用 LRU 淘汰机制"},
{Key: "login_session_ttl_hours", Value: "0", Type: configTypeSystem, Description: "登录会话过期时间 (小时,0表示浏览器关闭后自动退出,-1表示永不过期)"},
{Key: "update_upstream_repository", Value: "Rain-kl/Wavelet", Type: configTypeSystem, Description: "GitHub Actions Release 上游仓库"},
{Key: "storage_config", Value: `{"driver":"local","local":{"root":"."},"s3":{"region":"us-east-1"},"r2":{"region":"auto"},"minio":{"region":"us-east-1","path_style":true},"oss":{},"webdav":{}}`, Type: configTypeSystem, Description: "文件存储驱动及连接配置(JSON)"},
{Key: "log_database", Value: "sqlite", Type: configTypeSystem, Description: "当前日志主库"},
{Key: "log_db_migration", Value: "", Type: configTypeSystem, Description: "日志库迁移冻结标记"},
{Key: "log_retention_days_postgres", Value: "30", Type: configTypeBusiness, Description: "PostgreSQL 用户访问日志保留天数"},
{Key: "log_retention_days_sqlite", Value: "30", Type: configTypeBusiness, Description: "SQLite 用户访问日志保留天数"},
{Key: "log_retention_days_clickhouse", Value: "30", Type: configTypeBusiness, Description: "ClickHouse 用户访问日志保留天数"},
}
}
func seedDefaultConfigs(t *testing.T, tx *gorm.DB) {
defaultConfigs := getSeedConfigsPart1()
if err := tx.Create(&defaultConfigs).Error; err != nil {
t.Fatalf("failed to seed default system configs: %v", err)
}
publicKeys := map[string]struct{}{
"upload_allowed_extensions": {},
"site_name": {},
"password_login_enabled": {},
"registration_enabled": {},
"password_register_enabled": {},
"oidc_login_enabled": {},
}
keys := make([]string, 0, len(publicKeys))
for key := range publicKeys {
keys = append(keys, key)
}
if err := tx.Model(&SystemConfig{}).
Where("key IN ?", keys).
Update("visibility", "visible").Error; err != nil {
t.Fatalf("failed to seed public system config visibility: %v", err)
}
for _, config := range defaultConfigs {
if _, ok := publicKeys[config.Key]; ok {
config.Visibility = "visible"
}
_ = cachepkg.HSetJSON(context.Background(), "system_configs", config.Key, &config)
}
}
+15
View File
@@ -0,0 +1,15 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package trace 提供 OpenTelemetry 链路追踪封装工具
package trace
import "go.opentelemetry.io/otel/propagation"
func newPropagator() propagation.TextMapPropagator {
return propagation.NewCompositeTextMapPropagator(
propagation.TraceContext{},
propagation.Baggage{},
)
}
+19
View File
@@ -0,0 +1,19 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package trace
import (
sdktrace "go.opentelemetry.io/otel/sdk/trace"
)
// ParentBasedRatioSampler 创建父级感知的概率采样器
// - 如果父 Span 已采样,则子 Span 也采样
// - 如果父 Span 未采样,则子 Span 也不采样
// - 如果是根 Span,按 samplingRate 概率采样
func ParentBasedRatioSampler(samplingRate float64) sdktrace.Sampler {
return sdktrace.ParentBased(
sdktrace.TraceIDRatioBased(samplingRate),
)
}
+62
View File
@@ -0,0 +1,62 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package trace
import (
"context"
"log"
"go.opentelemetry.io/otel"
"go.opentelemetry.io/otel/trace"
)
// Tracer 全局 OpenTelemetry Tracer 实例
var Tracer trace.Tracer
var shutdownFuncs []func(context.Context) error
func init() {
// 初始化 Propagator
prop := newPropagator()
otel.SetTextMapPropagator(prop)
// 初始化 Tracer 实例为 No-op 默认以避免未初始化前或测试环境崩溃
Tracer = otel.GetTracerProvider().Tracer("github.com/Rain-kl/Wavelet")
}
// Config 链路追踪配置
type Config struct {
AppName string
SamplingRate float64
TracerName string
}
// Init 初始化 Tracer Provider 并关联全局 Tracer 实例
func Init(cfg Config) {
tracerProvider, err := newTracerProvider(cfg)
if err != nil {
log.Fatalf("[Trace] init trace provider failed: %v", err)
}
shutdownFuncs = append(shutdownFuncs, tracerProvider.Shutdown)
otel.SetTracerProvider(tracerProvider)
// 更新 Tracer
tracerName := cfg.TracerName
if tracerName == "" {
tracerName = "github.com/Rain-kl/Wavelet"
}
Tracer = tracerProvider.Tracer(tracerName)
}
// Shutdown 关闭所有 Trace Provider
func Shutdown(ctx context.Context) {
for _, fn := range shutdownFuncs {
_ = fn(ctx)
}
}
// Start 创建一个新的 Trace Span
func Start(ctx context.Context, name string, opts ...trace.SpanStartOption) (context.Context, trace.Span) {
return Tracer.Start(ctx, name, opts...)
}
+52
View File
@@ -0,0 +1,52 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package trace
import (
"context"
"os"
"go.opentelemetry.io/otel/attribute"
"go.opentelemetry.io/otel/exporters/otlp/otlptrace/otlptracegrpc"
"go.opentelemetry.io/otel/sdk/resource"
sdktrace "go.opentelemetry.io/otel/sdk/trace"
)
func newTracerProvider(cfg Config) (*sdktrace.TracerProvider, error) {
// 获取主机名和容器信息
hostname, err := os.Hostname()
if err != nil {
return nil, err
}
// 业务属性不绑定 schema URL,合并时继承 resource.Default() 的 SDK 内置版本,避免 semconv 与 otel/sdk 升级不同步。
r, err := resource.Merge(
resource.Default(),
resource.NewSchemaless(
attribute.String("service.name", cfg.AppName),
attribute.String("host.name", hostname),
attribute.String("k8s.namespace.name", os.Getenv("KUBERNETES_NAMESPACE")),
attribute.String("k8s.pod.name", os.Getenv("KUBERNETES_POD_NAME")),
attribute.String("k8s.pod.uid", os.Getenv("KUBERNETES_POD_UID")),
),
)
if err != nil {
return nil, err
}
// 初始化 Exporter
traceExporter, err := otlptracegrpc.New(context.Background())
if err != nil {
return nil, err
}
// 初始化 Trace
tracerProvider := sdktrace.NewTracerProvider(
sdktrace.WithBatcher(traceExporter),
sdktrace.WithResource(r),
sdktrace.WithSampler(ParentBasedRatioSampler(cfg.SamplingRate)),
)
return tracerProvider, nil
}
+10
View File
@@ -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"
)
+160
View File
@@ -0,0 +1,160 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package util provides generic utility functions.
package util
import (
"crypto/aes"
"crypto/cipher"
"crypto/ed25519"
"crypto/rand"
"encoding/base64"
"encoding/hex"
"errors"
"fmt"
"io"
)
const (
aesKeyLength = 32
errInvalidSignKey = "invalid sign key: %w"
errSignKeyLengthInvalid = "sign key must be 32 bytes (64 hex characters)"
errCreateCipherFailed = "failed to create cipher: %w"
errCreateGCMFailed = "failed to create GCM: %w"
errGenerateNonceFailed = "failed to generate nonce: %w"
errDecodeCiphertextFailed = "failed to decode ciphertext: %w"
errCiphertextTooShort = "ciphertext too short"
errDecryptFailed = "failed to decrypt: %w"
)
// Encrypt 使用 SignKey 加密字符串数据
// signKey: 64 字符 hex 编码的密钥(对应 32 字节,用于 AES-256)
// plaintext: 要加密的明文字符串
// return: base64 编码的密文
func Encrypt(signKey string, plaintext string) (string, error) {
return encryptBytes(signKey, []byte(plaintext))
}
// Decrypt 使用 SignKey 解密字符串数据
// signKey: 64 字符 hex 编码的密钥(对应 32 字节,用于 AES-256)
// ciphertext: base64 编码的密文
// return: 解密后的明文字符串
func Decrypt(signKey string, ciphertext string) (string, error) {
plaintext, err := decryptBytes(signKey, ciphertext)
if err != nil {
return "", err
}
return string(plaintext), nil
}
// encryptBytes 加密函数,处理字节数据
func encryptBytes(signKey string, plaintext []byte) (string, error) {
// 将 hex 编码的密钥转换为字节
key, err := hex.DecodeString(signKey)
if err != nil {
return "", fmt.Errorf(errInvalidSignKey, err)
}
if len(key) != aesKeyLength {
return "", errors.New(errSignKeyLengthInvalid)
}
// 创建 AES cipher
block, err := aes.NewCipher(key)
if err != nil {
return "", fmt.Errorf(errCreateCipherFailed, err)
}
// 使用 GCM 模式(Galois/Counter Mode)
gcm, err := cipher.NewGCM(block)
if err != nil {
return "", fmt.Errorf(errCreateGCMFailed, err)
}
// 生成随机 nonce
nonce := make([]byte, gcm.NonceSize())
if _, err := io.ReadFull(rand.Reader, nonce); err != nil {
return "", fmt.Errorf(errGenerateNonceFailed, err)
}
// 加密数据
ciphertext := gcm.Seal(nonce, nonce, plaintext, nil)
// 返回 base64 编码的密文
return Base64Encode(ciphertext), nil
}
// decryptBytes 解密函数,处理字节数据
func decryptBytes(signKey string, ciphertext string) ([]byte, error) {
// 将 hex 编码的密钥转换为字节
key, err := hex.DecodeString(signKey)
if err != nil {
return nil, fmt.Errorf(errInvalidSignKey, err)
}
if len(key) != aesKeyLength {
return nil, errors.New(errSignKeyLengthInvalid)
}
// 解码 base64 密文
data, err := Base64Decode(ciphertext)
if err != nil {
return nil, fmt.Errorf(errDecodeCiphertextFailed, err)
}
// 创建 AES cipher
block, err := aes.NewCipher(key)
if err != nil {
return nil, fmt.Errorf(errCreateCipherFailed, err)
}
// 使用 GCM 模式
gcm, err := cipher.NewGCM(block)
if err != nil {
return nil, fmt.Errorf(errCreateGCMFailed, err)
}
// 提取 nonce
nonceSize := gcm.NonceSize()
if len(data) < nonceSize {
return nil, errors.New(errCiphertextTooShort)
}
nonce, ciphertextBytes := data[:nonceSize], data[nonceSize:]
// 解密数据
plaintext, err := gcm.Open(nil, nonce, ciphertextBytes, nil)
if err != nil {
return nil, fmt.Errorf(errDecryptFailed, err)
}
return plaintext, nil
}
// Base64Encode Base64编码
func Base64Encode(data []byte) string {
return base64.StdEncoding.EncodeToString(data)
}
// Base64Decode Base64解码
func Base64Decode(encoded string) ([]byte, error) {
return base64.StdEncoding.DecodeString(encoded)
}
// Ed25519Verify 验证 Ed25519 签名
// publicKey: 32 字节的公钥(已解码的二进制格式)
// message: 待验证的原始消息
// signature: 64 字节的签名(已解码的二进制格式)
// return: 签名是否有效
func Ed25519Verify(publicKey, message, signature []byte) bool {
if len(publicKey) != ed25519.PublicKeySize {
return false
}
if len(signature) != ed25519.SignatureSize {
return false
}
return ed25519.Verify(publicKey, message, signature)
}
+29
View File
@@ -0,0 +1,29 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package util provides framework-agnostic helper types and HTTP utilities.
package util
import (
"database/sql/driver"
"encoding/json"
"fmt"
)
// StringArray custom type for handling JSON arrays
type StringArray []string
// Scan 实现 sql.Scanner 接口,从数据库读取 JSON 数组
func (sa *StringArray) Scan(value interface{}) error {
bytesValue, ok := value.([]byte)
if !ok {
return fmt.Errorf(errInvalidCustomValue, value)
}
return json.Unmarshal(bytesValue, sa)
}
// Value 实现 driver.Valuer 接口,将 JSON 数组序列化为数据库存储值
func (sa StringArray) Value() (driver.Value, error) {
return json.Marshal(sa)
}
+22
View File
@@ -0,0 +1,22 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package util
import "github.com/gin-gonic/gin"
// GetFromContext retrieves a typed value from Gin context.
func GetFromContext[T any](c *gin.Context, key string) (T, bool) {
value, exists := c.Get(key)
if !exists {
var zero T
return zero, false
}
typed, ok := value.(T)
return typed, ok
}
// SetToContext sets a typed value into Gin context.
func SetToContext[T any](c *gin.Context, key string, value T) {
c.Set(key, value)
}
+31
View File
@@ -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()
}()
}
+38
View File
@@ -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")
}
}
+68
View File
@@ -0,0 +1,68 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package util
import (
"context"
"fmt"
"io"
"net/http"
"net/url"
"time"
"github.com/Rain-kl/Wavelet/backend/pkg/httppool"
)
// IsLocalhost 检查 URL 是否为 localhost
func IsLocalhost(urlStr string) bool {
u, err := url.Parse(urlStr)
if err != nil {
return false
}
hostname := u.Hostname()
return hostname == "localhost" || hostname == "127.0.0.1" || hostname == "::1"
}
// HTTP 客户端配置常量
const (
httpClientTimeout = 10 // HTTP 客户端超时时间(秒)
httpMaxIdleConns = 100
httpMaxIdleConnsPerHost = 20
httpIdleConnTimeout = 60 // 空闲连接超时(秒)
)
// 配置HTTP客户端 使用 otelhttp 自动注入 trace span
var httpClient = &http.Client{
Timeout: httpClientTimeout * time.Second,
Transport: httppool.DefaultTransport(),
}
// SetHTTPClient 替换全局 HTTP 客户端实例
func SetHTTPClient(c *http.Client) {
httpClient = c
}
// Request 发送 HTTP 请求,支持自定义 Headers 和 Cookies
func Request(ctx context.Context, method, url string, body io.Reader, headers, cookies map[string]string) (*http.Response, error) {
req, err := http.NewRequestWithContext(ctx, method, url, body)
if err != nil {
return nil, fmt.Errorf(errCreateHTTPRequestFailed, err)
}
for key, value := range cookies {
req.AddCookie(&http.Cookie{Name: key, Value: value}) //nolint:gosec // client-side cookies do not require server attributes (Secure/HttpOnly)
}
for key, value := range headers {
req.Header.Set(key, value)
}
resp, err := httpClient.Do(req)
if err != nil {
return nil, fmt.Errorf(errHTTPRequestFailed, url, err)
}
return resp, nil
}
+16
View File
@@ -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)
}
+22
View File
@@ -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)
}
}
}
+48
View File
@@ -0,0 +1,48 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package util
import (
"sync"
"golang.org/x/crypto/bcrypt"
)
// HashPassword 使用 bcrypt 对密码进行哈希处理
func HashPassword(password string) (string, error) {
hash, err := bcrypt.GenerateFromPassword([]byte(password), bcrypt.DefaultCost)
if err != nil {
return "", err
}
return string(hash), nil
}
var dummyPasswordHashOnce sync.Once
var dummyPasswordHash string
func dummyHash() string {
dummyPasswordHashOnce.Do(func() {
hash, err := bcrypt.GenerateFromPassword([]byte("x"), bcrypt.DefaultCost)
if err != nil {
return
}
dummyPasswordHash = string(hash)
})
return dummyPasswordHash
}
// CheckPasswordHash 比较 bcrypt 哈希值与明文密码是否匹配
func CheckPasswordHash(hash, password string) bool {
return bcrypt.CompareHashAndPassword([]byte(hash), []byte(password)) == nil
}
// DummyCheckPassword runs a bcrypt compare against a dummy hash so missing-user
// login failures take a similar amount of time as a real password miss.
func DummyCheckPassword(password string) {
hash := dummyHash()
if hash == "" {
return
}
_ = CheckPasswordHash(hash, password)
}
+23
View File
@@ -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")
}
}
+35
View File
@@ -0,0 +1,35 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package util
import "strings"
// emailPartsCount 邮箱地址由 @ 分割为两部分
const (
emailPartsCount = 2
emailLocalMinChars = 2 // 邮箱 local 部分掩码显示的最小字符数
)
// DerefString 安全地解引用字符串指针,nil 返回空字符串
func DerefString(s *string) string {
if s == nil {
return ""
}
return *s
}
// MaskEmail 安全脱敏用户的邮箱地址(例如 us***@example.com)
func MaskEmail(email string) string {
parts := strings.Split(email, "@")
if len(parts) != emailPartsCount {
return email
}
local := parts[0]
domain := parts[1]
if len(local) <= emailLocalMinChars {
return "**@" + domain
}
return local[:2] + "***" + local[len(local)-1:] + "@" + domain
}
+29
View File
@@ -0,0 +1,29 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package util
import (
"crypto/rand"
"crypto/sha256"
"encoding/hex"
"io"
"github.com/google/uuid"
)
// uniqueIDBytes 生成唯一 ID 所需的随机字节长度
const uniqueIDBytes = 32
// GenerateUniqueIDSimple 生成 64 位唯一标识符
func GenerateUniqueIDSimple() string {
randomBytes := make([]byte, uniqueIDBytes)
if _, err := io.ReadFull(rand.Reader, randomBytes); err != nil {
// 如果随机数生成失败,使用 UUID 作为后备
uuidBytes := []byte(uuid.NewString())
hash := sha256.Sum256(uuidBytes)
copy(randomBytes, hash[:])
}
return hex.EncodeToString(randomBytes)
}