mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-01 22:46:38 +08:00
wavelet init
This commit is contained in:
Vendored
+391
@@ -0,0 +1,391 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package disk implements a platform-level disk-backed cache with size limit, TTL, and LRU eviction.
|
||||
package disk
|
||||
|
||||
import (
|
||||
"container/list"
|
||||
"encoding/binary"
|
||||
"errors"
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"sort"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/peterbourgon/diskv/v3"
|
||||
)
|
||||
|
||||
// ErrCacheMiss represents a cache miss.
|
||||
var ErrCacheMiss = errors.New("cache miss")
|
||||
|
||||
// Constants for disk cache configuration and sizing
|
||||
const (
|
||||
headerSize = 8 // 8 bytes metadata prefix for expiration UnixNano timestamp
|
||||
defaultMaxSizeMB = 100
|
||||
defaultTTLMinutes = 60
|
||||
cacheDirPerm = 0750
|
||||
|
||||
// DefaultExpiration applies the cache-wide default TTL.
|
||||
DefaultExpiration time.Duration = 0
|
||||
// NoExpiration stores the item without a TTL. Size limits and LRU eviction still apply.
|
||||
NoExpiration time.Duration = -1
|
||||
)
|
||||
|
||||
// Status represents the runtime cache statistics.
|
||||
type Status struct {
|
||||
TotalSize int64 `json:"total_size"`
|
||||
KeysCount int `json:"keys_count"`
|
||||
MaxSizeMB int64 `json:"max_size_mb"`
|
||||
TTLMinutes int64 `json:"ttl_minutes"`
|
||||
LRUEnabled bool `json:"lru_enabled"`
|
||||
BasePath string `json:"base_path"`
|
||||
}
|
||||
|
||||
// Cache implements the disk-backed cache with size limits, TTL, and LRU eviction.
|
||||
type Cache struct {
|
||||
mu sync.RWMutex
|
||||
d *diskv.Diskv
|
||||
basePath string
|
||||
maxSize int64 // in bytes
|
||||
defaultTTL time.Duration
|
||||
lruEnabled bool
|
||||
|
||||
// LRU and Size tracking
|
||||
currentSize int64
|
||||
items map[string]*list.Element
|
||||
evictList *list.List
|
||||
}
|
||||
|
||||
type cacheItem struct {
|
||||
key string
|
||||
size int64
|
||||
expiredAt time.Time
|
||||
}
|
||||
|
||||
// New creates a new Cache instance.
|
||||
func New(basePath string) *Cache {
|
||||
d := diskv.New(diskv.Options{
|
||||
BasePath: basePath,
|
||||
Transform: func(_ string) []string { return []string{} }, // flat structure for easy walk
|
||||
CacheSizeMax: 1024 * 1024, // 1MB in-memory cache size for diskv itself
|
||||
})
|
||||
|
||||
c := &Cache{
|
||||
d: d,
|
||||
basePath: basePath,
|
||||
maxSize: defaultMaxSizeMB * 1024 * 1024, // 100MB default
|
||||
defaultTTL: defaultTTLMinutes * time.Minute, // 60 minutes default
|
||||
lruEnabled: true,
|
||||
items: make(map[string]*list.Element),
|
||||
evictList: list.New(),
|
||||
}
|
||||
|
||||
// Scan directory on startup to rebuild LRU and size tracking
|
||||
_ = c.loadTracker()
|
||||
return c
|
||||
}
|
||||
|
||||
// Set stores a key-value pair in the cache.
|
||||
// Use DefaultExpiration for the configured default TTL, NoExpiration for no
|
||||
// TTL, or a positive duration for a business-specific TTL.
|
||||
func (c *Cache) Set(key string, value []byte, ttl time.Duration) error {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
|
||||
if ttl == DefaultExpiration {
|
||||
ttl = c.defaultTTL
|
||||
}
|
||||
|
||||
var expiredAt time.Time
|
||||
if ttl > 0 {
|
||||
expiredAt = time.Now().Add(ttl)
|
||||
}
|
||||
|
||||
// Prepare data layout: 8 bytes expiration timestamp + raw payload
|
||||
buf := make([]byte, headerSize+len(value))
|
||||
var expNano int64
|
||||
if !expiredAt.IsZero() {
|
||||
expNano = expiredAt.UnixNano()
|
||||
}
|
||||
binary.BigEndian.PutUint64(buf[0:headerSize], uint64(expNano))
|
||||
copy(buf[headerSize:], value)
|
||||
|
||||
// Write to diskv
|
||||
if err := c.d.Write(key, buf); err != nil {
|
||||
return fmt.Errorf("failed to write key to disk: %w", err)
|
||||
}
|
||||
|
||||
// Get file size on disk (approximate)
|
||||
size := int64(len(buf))
|
||||
|
||||
// Update memory tracker
|
||||
if elem, ok := c.items[key]; ok {
|
||||
item := elem.Value.(*cacheItem)
|
||||
c.currentSize += size - item.size
|
||||
item.size = size
|
||||
item.expiredAt = expiredAt
|
||||
c.evictList.MoveToFront(elem)
|
||||
} else {
|
||||
item := &cacheItem{
|
||||
key: key,
|
||||
size: size,
|
||||
expiredAt: expiredAt,
|
||||
}
|
||||
elem := c.evictList.PushFront(item)
|
||||
c.items[key] = elem
|
||||
c.currentSize += size
|
||||
}
|
||||
|
||||
// Evict items if size limit exceeded and LRU is enabled
|
||||
c.evict()
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// Get retrieves a key's value from the cache.
|
||||
func (c *Cache) Get(key string) ([]byte, error) {
|
||||
c.mu.RLock()
|
||||
elem, ok := c.items[key]
|
||||
if !ok {
|
||||
c.mu.RUnlock()
|
||||
return nil, ErrCacheMiss
|
||||
}
|
||||
|
||||
item := elem.Value.(*cacheItem)
|
||||
if !item.expiredAt.IsZero() && time.Now().After(item.expiredAt) {
|
||||
c.mu.RUnlock()
|
||||
return c.getAndDeleteIfExpired(key)
|
||||
}
|
||||
c.mu.RUnlock()
|
||||
|
||||
// Read from disk outside the lock so concurrent cache hits do not serialize on I/O.
|
||||
data, err := c.d.Read(key)
|
||||
if err != nil {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
if _, stillExists := c.items[key]; stillExists {
|
||||
_ = c.deleteUnlocked(key)
|
||||
}
|
||||
return nil, ErrCacheMiss
|
||||
}
|
||||
|
||||
if len(data) < headerSize {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
if _, stillExists := c.items[key]; stillExists {
|
||||
_ = c.deleteUnlocked(key)
|
||||
}
|
||||
return nil, ErrCacheMiss
|
||||
}
|
||||
|
||||
payload := data[headerSize:]
|
||||
|
||||
// Brief write lock only for LRU bookkeeping.
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
elem, ok = c.items[key]
|
||||
if !ok {
|
||||
return nil, ErrCacheMiss
|
||||
}
|
||||
c.evictList.MoveToFront(elem)
|
||||
|
||||
return payload, nil
|
||||
}
|
||||
|
||||
func (c *Cache) getAndDeleteIfExpired(key string) ([]byte, error) {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
|
||||
elem, ok := c.items[key]
|
||||
if !ok {
|
||||
return nil, ErrCacheMiss
|
||||
}
|
||||
|
||||
item := elem.Value.(*cacheItem)
|
||||
if !item.expiredAt.IsZero() && time.Now().After(item.expiredAt) {
|
||||
_ = c.deleteUnlocked(key)
|
||||
return nil, ErrCacheMiss
|
||||
}
|
||||
|
||||
data, err := c.d.Read(key)
|
||||
if err != nil {
|
||||
_ = c.deleteUnlocked(key)
|
||||
return nil, ErrCacheMiss
|
||||
}
|
||||
|
||||
if len(data) < headerSize {
|
||||
_ = c.deleteUnlocked(key)
|
||||
return nil, ErrCacheMiss
|
||||
}
|
||||
|
||||
c.evictList.MoveToFront(elem)
|
||||
return data[headerSize:], nil
|
||||
}
|
||||
|
||||
// Delete removes a key-value pair from the cache.
|
||||
func (c *Cache) Delete(key string) error {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
return c.deleteUnlocked(key)
|
||||
}
|
||||
|
||||
func (c *Cache) deleteUnlocked(key string) error {
|
||||
if elem, ok := c.items[key]; ok {
|
||||
item := elem.Value.(*cacheItem)
|
||||
c.currentSize -= item.size
|
||||
c.evictList.Remove(elem)
|
||||
delete(c.items, key)
|
||||
}
|
||||
return c.d.Erase(key)
|
||||
}
|
||||
|
||||
// Clear flushes all cached elements.
|
||||
func (c *Cache) Clear() error {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
|
||||
c.currentSize = 0
|
||||
c.items = make(map[string]*list.Element)
|
||||
c.evictList.Init()
|
||||
|
||||
return c.d.EraseAll()
|
||||
}
|
||||
|
||||
// Status returns the cache status.
|
||||
func (c *Cache) Status() Status {
|
||||
c.mu.RLock()
|
||||
defer c.mu.RUnlock()
|
||||
|
||||
return Status{
|
||||
TotalSize: c.currentSize,
|
||||
KeysCount: len(c.items),
|
||||
MaxSizeMB: c.maxSize / (1024 * 1024),
|
||||
TTLMinutes: int64(c.defaultTTL.Minutes()),
|
||||
LRUEnabled: c.lruEnabled,
|
||||
BasePath: c.basePath,
|
||||
}
|
||||
}
|
||||
|
||||
// UpdatePolicy dynamically updates policies.
|
||||
func (c *Cache) UpdatePolicy(maxSizeMB int64, ttlMinutes int64, lruEnabled bool) {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
|
||||
c.maxSize = maxSizeMB * 1024 * 1024
|
||||
c.defaultTTL = time.Duration(ttlMinutes) * time.Minute
|
||||
c.lruEnabled = lruEnabled
|
||||
c.evict()
|
||||
}
|
||||
|
||||
// evict evicts oldest items if current size exceeds maxSize and LRU is enabled.
|
||||
func (c *Cache) evict() {
|
||||
if !c.lruEnabled {
|
||||
return
|
||||
}
|
||||
|
||||
for c.currentSize > c.maxSize && c.evictList.Len() > 0 {
|
||||
elem := c.evictList.Back()
|
||||
item := elem.Value.(*cacheItem)
|
||||
c.currentSize -= item.size
|
||||
c.evictList.Remove(elem)
|
||||
delete(c.items, item.key)
|
||||
_ = c.d.Erase(item.key)
|
||||
}
|
||||
}
|
||||
|
||||
// loadTracker scans the cache directory on startup to rebuild memory state.
|
||||
func (c *Cache) loadTracker() error {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
|
||||
// Ensure directory exists
|
||||
if err := os.MkdirAll(c.basePath, cacheDirPerm); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
type loadedItem struct {
|
||||
key string
|
||||
size int64
|
||||
expiredAt time.Time
|
||||
modTime time.Time
|
||||
}
|
||||
var loadedItems []loadedItem
|
||||
|
||||
// Walk keys through diskv
|
||||
keysChan := c.d.Keys(nil)
|
||||
for key := range keysChan {
|
||||
// Read raw bytes to parse expiration prefix
|
||||
data, err := c.d.Read(key)
|
||||
if err != nil || len(data) < headerSize {
|
||||
_ = c.d.Erase(key) // corrupted file, wipe
|
||||
continue
|
||||
}
|
||||
|
||||
expNano := int64(binary.BigEndian.Uint64(data[0:headerSize])) //nolint:gosec // false positive: UnixNano fits within int64
|
||||
var expiredAt time.Time
|
||||
if expNano > 0 {
|
||||
expiredAt = time.Unix(0, expNano)
|
||||
}
|
||||
|
||||
// Check mod time for ordering
|
||||
path := filepath.Join(c.basePath, key)
|
||||
info, err := os.Stat(path)
|
||||
if err != nil {
|
||||
continue
|
||||
}
|
||||
|
||||
loadedItems = append(loadedItems, loadedItem{
|
||||
key: key,
|
||||
size: int64(len(data)),
|
||||
expiredAt: expiredAt,
|
||||
modTime: info.ModTime(),
|
||||
})
|
||||
}
|
||||
|
||||
// Sort by ModTime ascending (oldest first) so we rebuild LRU correctly
|
||||
sort.Slice(loadedItems, func(i, j int) bool {
|
||||
return loadedItems[i].modTime.Before(loadedItems[j].modTime)
|
||||
})
|
||||
|
||||
// Populate LRU (PushFront so that newest items are at the front, oldest at the back)
|
||||
for _, item := range loadedItems {
|
||||
entry := &cacheItem{
|
||||
key: item.key,
|
||||
size: item.size,
|
||||
expiredAt: item.expiredAt,
|
||||
}
|
||||
element := c.evictList.PushFront(entry)
|
||||
c.items[item.key] = element
|
||||
c.currentSize += item.size
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// StartCleanupWorker periodically cleans up expired cache items.
|
||||
func (c *Cache) StartCleanupWorker(interval time.Duration) {
|
||||
ticker := time.NewTicker(interval)
|
||||
for range ticker.C {
|
||||
c.cleanExpired()
|
||||
}
|
||||
}
|
||||
|
||||
// cleanExpired scans memory for expired items and removes them.
|
||||
func (c *Cache) cleanExpired() {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
|
||||
now := time.Now()
|
||||
for key, elem := range c.items {
|
||||
item := elem.Value.(*cacheItem)
|
||||
if !item.expiredAt.IsZero() && now.After(item.expiredAt) {
|
||||
c.currentSize -= item.size
|
||||
c.evictList.Remove(elem)
|
||||
delete(c.items, key)
|
||||
_ = c.d.Erase(key)
|
||||
}
|
||||
}
|
||||
}
|
||||
Vendored
+212
@@ -0,0 +1,212 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package disk
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"os"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func TestDiskCacheBasic(t *testing.T) {
|
||||
testDir := "uploads/test_diskcache_basic"
|
||||
defer func() { _ = os.RemoveAll(testDir) }()
|
||||
_ = os.RemoveAll(testDir)
|
||||
|
||||
c := New(testDir)
|
||||
defer func() { _ = c.Clear() }()
|
||||
|
||||
key := "key1"
|
||||
val := []byte("value1")
|
||||
|
||||
// Get non-existent
|
||||
_, err := c.Get(key)
|
||||
if err != ErrCacheMiss {
|
||||
t.Fatalf("expected ErrCacheMiss, got %v", err)
|
||||
}
|
||||
|
||||
// Set & Get
|
||||
err = c.Set(key, val, 10*time.Second)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to set cache: %v", err)
|
||||
}
|
||||
|
||||
got, err := c.Get(key)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to get cache: %v", err)
|
||||
}
|
||||
|
||||
if !bytes.Equal(got, val) {
|
||||
t.Errorf("expected %s, got %s", val, got)
|
||||
}
|
||||
|
||||
// Delete
|
||||
err = c.Delete(key)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to delete: %v", err)
|
||||
}
|
||||
|
||||
_, err = c.Get(key)
|
||||
if err != ErrCacheMiss {
|
||||
t.Errorf("expected ErrCacheMiss after delete, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDiskCacheTTL(t *testing.T) {
|
||||
testDir := "uploads/test_diskcache_ttl"
|
||||
defer func() { _ = os.RemoveAll(testDir) }()
|
||||
_ = os.RemoveAll(testDir)
|
||||
|
||||
c := New(testDir)
|
||||
defer func() { _ = c.Clear() }()
|
||||
|
||||
key := "ttlkey"
|
||||
val := []byte("ttlval")
|
||||
|
||||
// Set with 200ms TTL
|
||||
err := c.Set(key, val, 200*time.Millisecond)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to set: %v", err)
|
||||
}
|
||||
|
||||
// Immediate Get should succeed
|
||||
got, err := c.Get(key)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to get: %v", err)
|
||||
}
|
||||
if !bytes.Equal(got, val) {
|
||||
t.Errorf("expected %s, got %s", val, got)
|
||||
}
|
||||
|
||||
// Sleep 250ms to expire
|
||||
time.Sleep(250 * time.Millisecond)
|
||||
|
||||
// Get should fail with cache miss
|
||||
_, err = c.Get(key)
|
||||
if err != ErrCacheMiss {
|
||||
t.Errorf("expected ErrCacheMiss after TTL expiration, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDiskCacheExpirationPolicies(t *testing.T) {
|
||||
testDir := "uploads/test_diskcache_expiration_policies"
|
||||
defer func() { _ = os.RemoveAll(testDir) }()
|
||||
_ = os.RemoveAll(testDir)
|
||||
|
||||
c := New(testDir)
|
||||
defer func() { _ = c.Clear() }()
|
||||
c.defaultTTL = 50 * time.Millisecond
|
||||
|
||||
if err := c.Set("default", []byte("default"), DefaultExpiration); err != nil {
|
||||
t.Fatalf("Set(default, DefaultExpiration) returned error: %v", err)
|
||||
}
|
||||
if err := c.Set("custom", []byte("custom"), 100*time.Millisecond); err != nil {
|
||||
t.Fatalf("Set(custom, 100ms) returned error: %v", err)
|
||||
}
|
||||
if err := c.Set("permanent", []byte("permanent"), NoExpiration); err != nil {
|
||||
t.Fatalf("Set(permanent, NoExpiration) returned error: %v", err)
|
||||
}
|
||||
|
||||
time.Sleep(75 * time.Millisecond)
|
||||
|
||||
if _, err := c.Get("default"); err != ErrCacheMiss {
|
||||
t.Errorf("Get(default) error = %v, want ErrCacheMiss", err)
|
||||
}
|
||||
if _, err := c.Get("custom"); err != nil {
|
||||
t.Errorf("Get(custom) returned error before custom TTL elapsed: %v", err)
|
||||
}
|
||||
if _, err := c.Get("permanent"); err != nil {
|
||||
t.Errorf("Get(permanent) returned error: %v", err)
|
||||
}
|
||||
|
||||
time.Sleep(50 * time.Millisecond)
|
||||
|
||||
if _, err := c.Get("custom"); err != ErrCacheMiss {
|
||||
t.Errorf("Get(custom) error = %v, want ErrCacheMiss", err)
|
||||
}
|
||||
if _, err := c.Get("permanent"); err != nil {
|
||||
t.Errorf("Get(permanent) returned error after other entries expired: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDiskCacheNoExpirationSurvivesReload(t *testing.T) {
|
||||
testDir := "uploads/test_diskcache_no_expiration_reload"
|
||||
defer func() { _ = os.RemoveAll(testDir) }()
|
||||
_ = os.RemoveAll(testDir)
|
||||
|
||||
c := New(testDir)
|
||||
if err := c.Set("permanent", []byte("value"), NoExpiration); err != nil {
|
||||
t.Fatalf("Set(permanent, NoExpiration) returned error: %v", err)
|
||||
}
|
||||
|
||||
reloaded := New(testDir)
|
||||
defer func() { _ = reloaded.Clear() }()
|
||||
|
||||
got, err := reloaded.Get("permanent")
|
||||
if err != nil {
|
||||
t.Fatalf("reloaded Get(permanent) returned error: %v", err)
|
||||
}
|
||||
if !bytes.Equal(got, []byte("value")) {
|
||||
t.Errorf("reloaded Get(permanent) = %q, want %q", got, "value")
|
||||
}
|
||||
}
|
||||
|
||||
func TestDiskCacheLRUEviction(t *testing.T) {
|
||||
testDir := "uploads/test_diskcache_lru"
|
||||
defer func() { _ = os.RemoveAll(testDir) }()
|
||||
_ = os.RemoveAll(testDir)
|
||||
|
||||
c := New(testDir)
|
||||
defer func() { _ = c.Clear() }()
|
||||
|
||||
// Force a very small max size of 20 bytes for testing (8 bytes header + payload)
|
||||
// So 2 items of 2 bytes payload = 2 * (8 + 2) = 20 bytes max.
|
||||
c.maxSize = 20
|
||||
c.lruEnabled = true
|
||||
|
||||
// Write item 1: 8 + 2 = 10 bytes
|
||||
err := c.Set("k1", []byte("v1"), DefaultExpiration)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to set k1: %v", err)
|
||||
}
|
||||
|
||||
// Write item 2: 8 + 2 = 10 bytes
|
||||
err = c.Set("k2", []byte("v2"), DefaultExpiration)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to set k2: %v", err)
|
||||
}
|
||||
|
||||
// Both should exist
|
||||
if _, err := c.Get("k1"); err != nil {
|
||||
t.Errorf("k1 should exist: %v", err)
|
||||
}
|
||||
if _, err := c.Get("k2"); err != nil {
|
||||
t.Errorf("k2 should exist: %v", err)
|
||||
}
|
||||
|
||||
// Write item 3: 8 + 2 = 10 bytes -> total size would be 30, exceeding 20.
|
||||
// This should evict the oldest item. Since k1 was accessed, but then k2 was accessed,
|
||||
// wait, let's access k1 again to make it the most recently used, so k2 becomes oldest!
|
||||
_, _ = c.Get("k1") // k1 is now MRU, k2 is LRU
|
||||
|
||||
err = c.Set("k3", []byte("v3"), DefaultExpiration)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to set k3: %v", err)
|
||||
}
|
||||
|
||||
// k2 should be evicted, k1 and k3 should exist
|
||||
_, err = c.Get("k2")
|
||||
if err != ErrCacheMiss {
|
||||
t.Errorf("expected k2 to be evicted, got error %v", err)
|
||||
}
|
||||
|
||||
if _, err := c.Get("k1"); err != nil {
|
||||
t.Errorf("k1 should still exist: %v", err)
|
||||
}
|
||||
|
||||
if _, err := c.Get("k3"); err != nil {
|
||||
t.Errorf("k3 should exist: %v", err)
|
||||
}
|
||||
}
|
||||
Vendored
+73
@@ -0,0 +1,73 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package ram provides a thin wrapper around Otter v2 for process-local caching.
|
||||
package ram
|
||||
|
||||
import (
|
||||
"github.com/maypok86/otter/v2"
|
||||
)
|
||||
|
||||
const defaultMaximumSize = 256
|
||||
|
||||
// Options configures a RAM cache instance.
|
||||
type Options struct {
|
||||
// MaximumSize bounds the number of entries. Zero uses a small default.
|
||||
MaximumSize int
|
||||
}
|
||||
|
||||
// Cache is a concurrency-safe in-memory cache backed by Otter.
|
||||
type Cache[K comparable, V any] struct {
|
||||
inner *otter.Cache[K, V]
|
||||
}
|
||||
|
||||
// New creates a RAM cache from the provided options.
|
||||
func New[K comparable, V any](opts Options) (*Cache[K, V], error) {
|
||||
maximumSize := opts.MaximumSize
|
||||
if maximumSize == 0 {
|
||||
maximumSize = defaultMaximumSize
|
||||
}
|
||||
|
||||
inner, err := otter.New(&otter.Options[K, V]{
|
||||
MaximumSize: maximumSize,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return &Cache[K, V]{inner: inner}, nil
|
||||
}
|
||||
|
||||
// MustNew creates a RAM cache and panics when configuration is invalid.
|
||||
func MustNew[K comparable, V any](opts Options) *Cache[K, V] {
|
||||
cache, err := New[K, V](opts)
|
||||
if err != nil {
|
||||
panic(err)
|
||||
}
|
||||
return cache
|
||||
}
|
||||
|
||||
// GetIfPresent returns the cached value when present.
|
||||
func (c *Cache[K, V]) GetIfPresent(key K) (V, bool) {
|
||||
return c.inner.GetIfPresent(key)
|
||||
}
|
||||
|
||||
// Set stores a value in the cache.
|
||||
func (c *Cache[K, V]) Set(key K, value V) {
|
||||
c.inner.Set(key, value)
|
||||
}
|
||||
|
||||
// Invalidate removes one entry from the cache.
|
||||
func (c *Cache[K, V]) Invalidate(key K) {
|
||||
c.inner.Invalidate(key)
|
||||
}
|
||||
|
||||
// InvalidateAll removes every entry from the cache.
|
||||
func (c *Cache[K, V]) InvalidateAll() {
|
||||
c.inner.InvalidateAll()
|
||||
}
|
||||
|
||||
// EstimatedSize returns the approximate number of cached entries.
|
||||
func (c *Cache[K, V]) EstimatedSize() int {
|
||||
return c.inner.EstimatedSize()
|
||||
}
|
||||
Vendored
+38
@@ -0,0 +1,38 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package ram
|
||||
|
||||
import "testing"
|
||||
|
||||
func TestCacheSetGetInvalidate(t *testing.T) {
|
||||
cache := MustNew[string, int](Options{MaximumSize: 8})
|
||||
|
||||
cache.Set("count", 3)
|
||||
|
||||
got, ok := cache.GetIfPresent("count")
|
||||
if !ok {
|
||||
t.Fatal("GetIfPresent(count) ok = false, want true")
|
||||
}
|
||||
if got != 3 {
|
||||
t.Fatalf("GetIfPresent(count) = %d, want %d", got, 3)
|
||||
}
|
||||
|
||||
cache.Invalidate("count")
|
||||
if _, ok := cache.GetIfPresent("count"); ok {
|
||||
t.Fatal("GetIfPresent(count) after Invalidate ok = true, want false")
|
||||
}
|
||||
}
|
||||
|
||||
func TestCacheInvalidateAll(t *testing.T) {
|
||||
cache := MustNew[string, string](Options{MaximumSize: 8})
|
||||
|
||||
cache.Set("a", "1")
|
||||
cache.Set("b", "2")
|
||||
|
||||
cache.InvalidateAll()
|
||||
|
||||
if cache.EstimatedSize() != 0 {
|
||||
t.Fatalf("EstimatedSize() after InvalidateAll = %d, want 0", cache.EstimatedSize())
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,259 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package cap 提供人机验证(CAPTCHA)功能
|
||||
package cap
|
||||
|
||||
import (
|
||||
"crypto/hmac"
|
||||
"crypto/rand"
|
||||
"crypto/sha256"
|
||||
"encoding/base64"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
const (
|
||||
jwtHeaderB64 = "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9"
|
||||
jwtPartsCount = 3 // JWT 三段结构
|
||||
defaultChallengeCount = 50 // 默认 PoW 难题数
|
||||
defaultChallengeSize = 32 // 默认盐值长度
|
||||
defaultDifficulty = 4 // 默认难度
|
||||
defaultNonceLength = 25 // 随机 Nonce 字节长度
|
||||
defaultExpires = 10 * time.Minute // 默认过期时间
|
||||
)
|
||||
|
||||
// ChallengeConfig holds parameters for the PoW challenge
|
||||
type ChallengeConfig struct {
|
||||
Count int // Number of puzzles (c)
|
||||
Size int // Salt length (s)
|
||||
Difficulty int // Difficulty prefix length (d)
|
||||
Expires time.Duration // Challenge TTL
|
||||
}
|
||||
|
||||
// ChallengeResponse is returned to the client
|
||||
type ChallengeResponse struct {
|
||||
Challenge struct {
|
||||
C int `json:"c"`
|
||||
S int `json:"s"`
|
||||
D int `json:"d"`
|
||||
} `json:"challenge"`
|
||||
Token string `json:"token"`
|
||||
Expires int64 `json:"expires"` // ms timestamp
|
||||
}
|
||||
|
||||
// ChallengePayload represents the signed JWT payload
|
||||
type ChallengePayload struct {
|
||||
Nonce string `json:"n"`
|
||||
Count int `json:"c"`
|
||||
Size int `json:"s"`
|
||||
Difficulty int `json:"d"`
|
||||
Expires int64 `json:"exp"` // ms timestamp
|
||||
IssuedAt int64 `json:"iat"` // ms timestamp
|
||||
Scope string `json:"sk,omitempty"`
|
||||
}
|
||||
|
||||
// RedeemRequest payload sent by client
|
||||
type RedeemRequest struct {
|
||||
Token string `json:"token"`
|
||||
Solutions []int `json:"solutions"`
|
||||
}
|
||||
|
||||
// RedeemResponse returned to client after verification
|
||||
type RedeemResponse struct {
|
||||
Success bool `json:"success"`
|
||||
Token string `json:"token,omitempty"`
|
||||
Expires int64 `json:"expires,omitempty"`
|
||||
Error string `json:"error,omitempty"`
|
||||
}
|
||||
|
||||
func b64urlEncode(data []byte) string {
|
||||
return base64.RawURLEncoding.EncodeToString(data)
|
||||
}
|
||||
|
||||
func b64urlDecode(str string) ([]byte, error) {
|
||||
return base64.RawURLEncoding.DecodeString(str)
|
||||
}
|
||||
|
||||
// RandomHex generates a cryptographically secure random hexadecimal string of the specified byte length.
|
||||
func RandomHex(byteLen int) string {
|
||||
bytes := make([]byte, byteLen)
|
||||
if _, err := rand.Read(bytes); err != nil {
|
||||
panic(err)
|
||||
}
|
||||
return hex.EncodeToString(bytes)
|
||||
}
|
||||
|
||||
func jwtSign(payload []byte, secret []byte) string {
|
||||
body := b64urlEncode(payload)
|
||||
sigInput := jwtHeaderB64 + "." + body
|
||||
|
||||
mac := hmac.New(sha256.New, secret)
|
||||
mac.Write([]byte(sigInput))
|
||||
sig := mac.Sum(nil)
|
||||
|
||||
return sigInput + "." + b64urlEncode(sig)
|
||||
}
|
||||
|
||||
func jwtVerify(token string, secret []byte) ([]byte, error) {
|
||||
parts := strings.Split(token, ".")
|
||||
if len(parts) != jwtPartsCount {
|
||||
return nil, errors.New(errInvalidTokenFormat)
|
||||
}
|
||||
if parts[0] != jwtHeaderB64 {
|
||||
return nil, errors.New(errInvalidHeader)
|
||||
}
|
||||
|
||||
sigInput := parts[0] + "." + parts[1]
|
||||
mac := hmac.New(sha256.New, secret)
|
||||
mac.Write([]byte(sigInput))
|
||||
expectedSig := mac.Sum(nil)
|
||||
|
||||
actualSig, err := b64urlDecode(parts[2])
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if !hmac.Equal(expectedSig, actualSig) {
|
||||
return nil, errors.New(errSignatureMismatch)
|
||||
}
|
||||
|
||||
payload, err := b64urlDecode(parts[1])
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return payload, nil
|
||||
}
|
||||
|
||||
// JwtSigHex extracts the signature part of a JWT token and returns it as a hexadecimal string.
|
||||
func JwtSigHex(token string) string {
|
||||
parts := strings.Split(token, ".")
|
||||
if len(parts) != jwtPartsCount {
|
||||
return ""
|
||||
}
|
||||
sigBytes, err := b64urlDecode(parts[2])
|
||||
if err != nil {
|
||||
return ""
|
||||
}
|
||||
return hex.EncodeToString(sigBytes)
|
||||
}
|
||||
|
||||
// GenerateChallenge produces a new challenge and signed token
|
||||
func GenerateChallenge(secret []byte, conf ChallengeConfig, scope string) (*ChallengeResponse, error) {
|
||||
if conf.Count <= 0 {
|
||||
conf.Count = defaultChallengeCount
|
||||
}
|
||||
if conf.Size <= 0 {
|
||||
conf.Size = defaultChallengeSize
|
||||
}
|
||||
if conf.Difficulty <= 0 {
|
||||
conf.Difficulty = defaultDifficulty
|
||||
}
|
||||
if conf.Expires <= 0 {
|
||||
conf.Expires = defaultExpires
|
||||
}
|
||||
|
||||
now := time.Now().UnixNano() / int64(time.Millisecond)
|
||||
expires := now + int64(conf.Expires/time.Millisecond)
|
||||
|
||||
payload := ChallengePayload{
|
||||
Nonce: RandomHex(defaultNonceLength),
|
||||
Count: conf.Count,
|
||||
Size: conf.Size,
|
||||
Difficulty: conf.Difficulty,
|
||||
Expires: expires,
|
||||
IssuedAt: now,
|
||||
Scope: scope,
|
||||
}
|
||||
|
||||
payloadBytes, err := json.Marshal(payload)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
token := jwtSign(payloadBytes, secret)
|
||||
|
||||
resp := &ChallengeResponse{
|
||||
Token: token,
|
||||
Expires: expires,
|
||||
}
|
||||
resp.Challenge.C = conf.Count
|
||||
resp.Challenge.S = conf.Size
|
||||
resp.Challenge.D = conf.Difficulty
|
||||
|
||||
return resp, nil
|
||||
}
|
||||
|
||||
// VerifyChallengeSolutions verifies client submitted solutions
|
||||
func VerifyChallengeSolutions(token string, solutions []int, secret []byte, expectedScope string) (*ChallengePayload, error) {
|
||||
payloadBytes, err := jwtVerify(token, secret)
|
||||
if err != nil {
|
||||
return nil, errors.New(errInvalidToken)
|
||||
}
|
||||
|
||||
var payload ChallengePayload
|
||||
if err := json.Unmarshal(payloadBytes, &payload); err != nil {
|
||||
return nil, errors.New(errInvalidToken)
|
||||
}
|
||||
|
||||
if expectedScope != "" && payload.Scope != expectedScope {
|
||||
return nil, errors.New(errScopeMismatch)
|
||||
}
|
||||
|
||||
now := time.Now().UnixNano() / int64(time.Millisecond)
|
||||
if payload.Expires < now {
|
||||
return nil, errors.New(errExpired)
|
||||
}
|
||||
|
||||
if len(solutions) != payload.Count {
|
||||
return nil, errors.New(errInvalidSolutions)
|
||||
}
|
||||
|
||||
tokenFnv := fnv1a(token)
|
||||
for i := 0; i < payload.Count; i++ {
|
||||
idxStr := strconv.Itoa(i + 1)
|
||||
saltSeed := fnv1aResume(tokenFnv, idxStr)
|
||||
targetSeed := fnv1aResume(saltSeed, "d")
|
||||
salt := prngFromHash(saltSeed, payload.Size)
|
||||
target := prngFromHash(targetSeed, payload.Difficulty)
|
||||
|
||||
hashInput := salt + strconv.Itoa(solutions[i])
|
||||
hashBytes := sha256.Sum256([]byte(hashInput))
|
||||
hashHex := hex.EncodeToString(hashBytes[:])
|
||||
|
||||
if !strings.HasPrefix(hashHex, target) {
|
||||
return nil, errors.New(errInvalidSolution)
|
||||
}
|
||||
}
|
||||
|
||||
return &payload, nil
|
||||
}
|
||||
|
||||
// Solve is a utility function to solve a challenge (mainly used for tests and reference implementation)
|
||||
func Solve(token string, count, size, difficulty int) []int {
|
||||
solutions := make([]int, count)
|
||||
tokenFnv := fnv1a(token)
|
||||
for i := 0; i < count; i++ {
|
||||
idxStr := strconv.Itoa(i + 1)
|
||||
saltSeed := fnv1aResume(tokenFnv, idxStr)
|
||||
targetSeed := fnv1aResume(saltSeed, "d")
|
||||
salt := prngFromHash(saltSeed, size)
|
||||
target := prngFromHash(targetSeed, difficulty)
|
||||
|
||||
for nonce := 0; nonce < 1000000; nonce++ {
|
||||
hashInput := salt + strconv.Itoa(nonce)
|
||||
hashBytes := sha256.Sum256([]byte(hashInput))
|
||||
hashHex := hex.EncodeToString(hashBytes[:])
|
||||
if strings.HasPrefix(hashHex, target) {
|
||||
solutions[i] = nonce
|
||||
break
|
||||
}
|
||||
}
|
||||
}
|
||||
return solutions
|
||||
}
|
||||
@@ -0,0 +1,15 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package cap
|
||||
|
||||
const (
|
||||
errInvalidTokenFormat = "invalid token format"
|
||||
errInvalidHeader = "invalid header"
|
||||
errSignatureMismatch = "signature mismatch"
|
||||
errInvalidToken = "invalid_token"
|
||||
errScopeMismatch = "scope_mismatch"
|
||||
errExpired = "expired"
|
||||
errInvalidSolutions = "invalid_solutions"
|
||||
errInvalidSolution = "invalid_solution"
|
||||
)
|
||||
@@ -0,0 +1,49 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package cap
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// fnv1a returns the 32-bit FNV-1a hash of a string
|
||||
//
|
||||
//nolint:mnd // FNV-1a 算法位移常量
|
||||
func fnv1a(str string) uint32 {
|
||||
var hash uint32 = 2166136261
|
||||
for i := 0; i < len(str); i++ {
|
||||
hash ^= uint32(str[i])
|
||||
hash += (hash << 1) + (hash << 4) + (hash << 7) + (hash << 8) + (hash << 24)
|
||||
}
|
||||
return hash
|
||||
}
|
||||
|
||||
// fnv1aResume resumes FNV-1a hashing from a given state
|
||||
//
|
||||
//nolint:mnd // FNV-1a 算法位移常量
|
||||
func fnv1aResume(state uint32, str string) uint32 {
|
||||
h := state
|
||||
for i := 0; i < len(str); i++ {
|
||||
h ^= uint32(str[i])
|
||||
h += (h << 1) + (h << 4) + (h << 7) + (h << 8) + (h << 24)
|
||||
}
|
||||
return h
|
||||
}
|
||||
|
||||
// prngFromHash generates a hex string of specified length using an initial hash state
|
||||
//
|
||||
//nolint:mnd // xorshift 算法位移常量
|
||||
func prngFromHash(initialHash uint32, length int) string {
|
||||
state := initialHash
|
||||
var result strings.Builder
|
||||
for result.Len() < length {
|
||||
state ^= state << 13
|
||||
state ^= state >> 17
|
||||
state ^= state << 5
|
||||
hexStr := fmt.Sprintf("%08x", state)
|
||||
result.WriteString(hexStr)
|
||||
}
|
||||
return result.String()[:length]
|
||||
}
|
||||
@@ -0,0 +1,186 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package cap
|
||||
|
||||
import (
|
||||
"context"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/redis/go-redis/v9"
|
||||
)
|
||||
|
||||
// Store defines the storage interface for challenge nonces and verification tokens
|
||||
type Store interface {
|
||||
Get(ctx context.Context, key string) (string, bool, error)
|
||||
Set(ctx context.Context, key string, val string, ttl time.Duration) error
|
||||
Delete(ctx context.Context, key string) error
|
||||
// SetNX atomically sets key=val with the given TTL only when the key does not
|
||||
// exist yet. It returns true when the key was actually written (i.e. this
|
||||
// caller "won" the race), and false when the key already existed.
|
||||
SetNX(ctx context.Context, key string, val string, ttl time.Duration) (bool, error)
|
||||
// GetAndDelete atomically retrieves the value of key and removes it in a
|
||||
// single operation. Returns ("", false, nil) when the key does not exist.
|
||||
GetAndDelete(ctx context.Context, key string) (string, bool, error)
|
||||
}
|
||||
|
||||
type memoryItem struct {
|
||||
value string
|
||||
expiresAt time.Time
|
||||
}
|
||||
|
||||
// MemoryStore is a thread-safe in-memory implementation of Store
|
||||
type MemoryStore struct {
|
||||
items map[string]memoryItem
|
||||
mu sync.Mutex // unified write-lock; promotes to exclusive for all ops
|
||||
}
|
||||
|
||||
// NewMemoryStore creates and initializes a new MemoryStore
|
||||
func NewMemoryStore(cleanupInterval time.Duration) *MemoryStore {
|
||||
store := &MemoryStore{
|
||||
items: make(map[string]memoryItem),
|
||||
}
|
||||
if cleanupInterval > 0 {
|
||||
go store.startCleanupLoop(cleanupInterval)
|
||||
}
|
||||
return store
|
||||
}
|
||||
|
||||
// Get 从 MemoryStore 获取指定 key 的值
|
||||
func (s *MemoryStore) Get(_ context.Context, key string) (string, bool, error) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
return s.getLocked(key)
|
||||
}
|
||||
|
||||
// getLocked is the internal helper – caller must hold s.mu.
|
||||
func (s *MemoryStore) getLocked(key string) (string, bool, error) {
|
||||
item, found := s.items[key]
|
||||
if !found {
|
||||
return "", false, nil
|
||||
}
|
||||
if time.Now().After(item.expiresAt) {
|
||||
delete(s.items, key)
|
||||
return "", false, nil
|
||||
}
|
||||
return item.value, true, nil
|
||||
}
|
||||
|
||||
// Set 向 MemoryStore 写入指定 key 的值
|
||||
func (s *MemoryStore) Set(_ context.Context, key string, val string, ttl time.Duration) error {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
s.items[key] = memoryItem{
|
||||
value: val,
|
||||
expiresAt: time.Now().Add(ttl),
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// Delete 从 MemoryStore 删除指定 key
|
||||
func (s *MemoryStore) Delete(_ context.Context, key string) error {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
delete(s.items, key)
|
||||
return nil
|
||||
}
|
||||
|
||||
// SetNX atomically sets key only when it is absent (or expired).
|
||||
// Returns true if the key was written by this call.
|
||||
func (s *MemoryStore) SetNX(_ context.Context, key string, val string, ttl time.Duration) (bool, error) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
|
||||
_, exists, _ := s.getLocked(key)
|
||||
if exists {
|
||||
return false, nil
|
||||
}
|
||||
s.items[key] = memoryItem{
|
||||
value: val,
|
||||
expiresAt: time.Now().Add(ttl),
|
||||
}
|
||||
return true, nil
|
||||
}
|
||||
|
||||
// GetAndDelete atomically retrieves and removes key in one critical section.
|
||||
func (s *MemoryStore) GetAndDelete(_ context.Context, key string) (string, bool, error) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
|
||||
val, exists, err := s.getLocked(key)
|
||||
if err != nil || !exists {
|
||||
return "", false, err
|
||||
}
|
||||
delete(s.items, key)
|
||||
return val, true, nil
|
||||
}
|
||||
|
||||
func (s *MemoryStore) startCleanupLoop(interval time.Duration) {
|
||||
ticker := time.NewTicker(interval)
|
||||
for range ticker.C {
|
||||
s.cleanupExpired()
|
||||
}
|
||||
}
|
||||
|
||||
func (s *MemoryStore) cleanupExpired() {
|
||||
now := time.Now()
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
for k, v := range s.items {
|
||||
if now.After(v.expiresAt) {
|
||||
delete(s.items, k)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// RedisStore is a GORM-compatible/standalone Redis-backed implementation of Store
|
||||
type RedisStore struct {
|
||||
client redis.UniversalClient
|
||||
}
|
||||
|
||||
// NewRedisStore creates a new RedisStore wrapping a redis.UniversalClient
|
||||
func NewRedisStore(client redis.UniversalClient) *RedisStore {
|
||||
return &RedisStore{
|
||||
client: client,
|
||||
}
|
||||
}
|
||||
|
||||
// Get 从 RedisStore 获取指定 key 的值
|
||||
func (s *RedisStore) Get(ctx context.Context, key string) (string, bool, error) {
|
||||
val, err := s.client.Get(ctx, key).Result()
|
||||
if err == redis.Nil {
|
||||
return "", false, nil
|
||||
}
|
||||
if err != nil {
|
||||
return "", false, err
|
||||
}
|
||||
return val, true, nil
|
||||
}
|
||||
|
||||
// Set 向 RedisStore 写入指定 key 的值
|
||||
func (s *RedisStore) Set(ctx context.Context, key string, val string, ttl time.Duration) error {
|
||||
return s.client.Set(ctx, key, val, ttl).Err()
|
||||
}
|
||||
|
||||
// Delete 从 RedisStore 删除指定 key
|
||||
func (s *RedisStore) Delete(ctx context.Context, key string) error {
|
||||
return s.client.Del(ctx, key).Err()
|
||||
}
|
||||
|
||||
// SetNX wraps Redis SET NX – returns true only when the key was newly created.
|
||||
func (s *RedisStore) SetNX(ctx context.Context, key string, val string, ttl time.Duration) (bool, error) {
|
||||
return s.client.SetNX(ctx, key, val, ttl).Result()
|
||||
}
|
||||
|
||||
// GetAndDelete wraps Redis GETDEL (available since Redis 6.2).
|
||||
func (s *RedisStore) GetAndDelete(ctx context.Context, key string) (string, bool, error) {
|
||||
val, err := s.client.GetDel(ctx, key).Result()
|
||||
if err == redis.Nil {
|
||||
return "", false, nil
|
||||
}
|
||||
if err != nil {
|
||||
return "", false, err
|
||||
}
|
||||
return val, true, nil
|
||||
}
|
||||
@@ -0,0 +1,66 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package httppool manages shared, optimized HTTP transports to reuse TCP connections.
|
||||
package httppool
|
||||
|
||||
import (
|
||||
"crypto/tls"
|
||||
"net"
|
||||
"net/http"
|
||||
"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
|
||||
)
|
||||
|
||||
// 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() {
|
||||
transport := &http.Transport{
|
||||
Proxy: http.ProxyFromEnvironment,
|
||||
DialContext: (&net.Dialer{
|
||||
Timeout: dialTimeout,
|
||||
KeepAlive: dialKeepAlive,
|
||||
}).DialContext,
|
||||
ForceAttemptHTTP2: true,
|
||||
MaxIdleConns: maxIdleConns,
|
||||
MaxIdleConnsPerHost: maxIdleConnsPerHost,
|
||||
IdleConnTimeout: idleConnTimeout,
|
||||
TLSHandshakeTimeout: tlsHandshakeTimeout,
|
||||
ExpectContinueTimeout: expectContinueTimeout,
|
||||
TLSClientConfig: &tls.Config{
|
||||
ClientSessionCache: tls.NewLRUClientSessionCache(tlsSessionCacheSize),
|
||||
},
|
||||
}
|
||||
defaultTransport = otelhttp.NewTransport(transport)
|
||||
})
|
||||
return defaultTransport
|
||||
}
|
||||
|
||||
// NewClient returns a new http.Client that shares the global connection pool
|
||||
// but has its own timeout configuration.
|
||||
func NewClient(timeout time.Duration) *http.Client {
|
||||
return &http.Client{
|
||||
Timeout: timeout,
|
||||
Transport: DefaultTransport(),
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,37 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package httppool
|
||||
|
||||
import (
|
||||
"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")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,9 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package logger 提供结构化日志封装
|
||||
package logger
|
||||
|
||||
const (
|
||||
errCreateLogFileDirFailed = "[Logger] create log file dir err: %w"
|
||||
)
|
||||
@@ -0,0 +1,102 @@
|
||||
// Copyright 2025 linux.do
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package logger
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"log"
|
||||
|
||||
"github.com/uptrace/opentelemetry-go-extra/otelzap"
|
||||
"go.uber.org/zap"
|
||||
"go.uber.org/zap/zapcore"
|
||||
)
|
||||
|
||||
// Config represents the logging configuration.
|
||||
type Config struct {
|
||||
Level string
|
||||
Format string
|
||||
Output string
|
||||
FilePath string
|
||||
MaxSize int
|
||||
MaxAge int
|
||||
MaxBackups int
|
||||
Compress bool
|
||||
}
|
||||
|
||||
var logger *otelzap.Logger
|
||||
|
||||
// ringBufferCapacity 环形缓冲区容量
|
||||
const ringBufferCapacity = 5000
|
||||
|
||||
// GlobalRingBuffer 全局日志环形缓冲区,供 Admin 日志查询和 WebSocket 推送使用
|
||||
var GlobalRingBuffer *LogRingBuffer
|
||||
|
||||
func doInit(cfg Config) {
|
||||
logWriter, err := getLogWriterForConfig(cfg)
|
||||
if err != nil {
|
||||
log.Fatalf("[Logger] get log writer err: %v\n", err)
|
||||
}
|
||||
|
||||
// 初始化 ring buffer(保留最近 5000 行日志),如果是多次调用 Init,不需要重复创建 GlobalRingBuffer
|
||||
if GlobalRingBuffer == nil {
|
||||
GlobalRingBuffer = NewLogRingBuffer(ringBufferCapacity)
|
||||
}
|
||||
|
||||
// 使用 multi writer 同时写入原始输出和 ring buffer
|
||||
multiWriter := zapcore.NewMultiWriteSyncer(
|
||||
logWriter,
|
||||
zapcore.AddSync(GlobalRingBuffer),
|
||||
)
|
||||
|
||||
zapLogger := zap.New(
|
||||
zapcore.NewCore(getEncoderForConfig(cfg), multiWriter, getLogLevelForConfig(cfg)),
|
||||
zap.AddCaller(),
|
||||
zap.AddCallerSkip(1),
|
||||
)
|
||||
logger = otelzap.New(
|
||||
zapLogger,
|
||||
otelzap.WithMinLevel(zapLogger.Level()),
|
||||
)
|
||||
}
|
||||
|
||||
func init() {
|
||||
// 默认使用 console stdout INFO 日志输出,避免在 Init 前或测试中发生空指针崩溃
|
||||
defaultCfg := Config{
|
||||
Level: "info",
|
||||
Format: "console",
|
||||
Output: "stdout",
|
||||
}
|
||||
doInit(defaultCfg)
|
||||
}
|
||||
|
||||
// Init initializes the logger with a custom configuration.
|
||||
func Init(cfg Config) {
|
||||
doInit(cfg)
|
||||
}
|
||||
|
||||
// DebugF 输出 Debug 级别日志
|
||||
func DebugF(ctx context.Context, format string, args ...interface{}) {
|
||||
msg := fmt.Sprintf(format, args...)
|
||||
logger.Ctx(ctx).Debug(msg, getTraceIDFields(ctx)...)
|
||||
}
|
||||
|
||||
// InfoF 输出 Info 级别日志
|
||||
func InfoF(ctx context.Context, format string, args ...interface{}) {
|
||||
msg := fmt.Sprintf(format, args...)
|
||||
logger.Ctx(ctx).Info(msg, getTraceIDFields(ctx)...)
|
||||
}
|
||||
|
||||
// WarnF 输出 Warn 级别日志
|
||||
func WarnF(ctx context.Context, format string, args ...interface{}) {
|
||||
msg := fmt.Sprintf(format, args...)
|
||||
logger.Ctx(ctx).Warn(msg, getTraceIDFields(ctx)...)
|
||||
}
|
||||
|
||||
// ErrorF 输出 Error 级别日志
|
||||
func ErrorF(ctx context.Context, format string, args ...interface{}) {
|
||||
msg := fmt.Sprintf(format, args...)
|
||||
logger.Ctx(ctx).Error(msg, getTraceIDFields(ctx)...)
|
||||
}
|
||||
@@ -0,0 +1,175 @@
|
||||
// Copyright 2025 linux.do
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package logger
|
||||
|
||||
import (
|
||||
"io"
|
||||
"sync"
|
||||
)
|
||||
|
||||
// LogEntry 日志条目,对应 ring buffer 中的一行日志
|
||||
type LogEntry struct {
|
||||
Index int `json:"index"` // 全局递增序号
|
||||
Data string `json:"data"` // 一行日志原文(含换行符)
|
||||
}
|
||||
|
||||
// LogRingBuffer 固定容量的环形缓冲区,存储最近的日志行
|
||||
// 支持:追加日志、按 cursor 分页查询、订阅实时推送
|
||||
type LogRingBuffer struct {
|
||||
mu sync.RWMutex
|
||||
entries []LogEntry
|
||||
cap int
|
||||
head int // 下一条写入的位置
|
||||
count int // 当前条目数
|
||||
seq int // 全局递增序号
|
||||
|
||||
subscribers map[chan LogEntry]struct{}
|
||||
subMu sync.RWMutex
|
||||
}
|
||||
|
||||
// NewLogRingBuffer 创建指定容量的日志环形缓冲区
|
||||
func NewLogRingBuffer(capacity int) *LogRingBuffer {
|
||||
return &LogRingBuffer{
|
||||
entries: make([]LogEntry, capacity),
|
||||
cap: capacity,
|
||||
subscribers: make(map[chan LogEntry]struct{}),
|
||||
}
|
||||
}
|
||||
|
||||
// Write 实现 io.Writer 接口,供 zapcore.WriteSyncer 调用
|
||||
// 按 '\n' 分割为独立行写入 ring buffer
|
||||
func (r *LogRingBuffer) Write(p []byte) (int, error) {
|
||||
if len(p) == 0 {
|
||||
return 0, nil
|
||||
}
|
||||
|
||||
data := string(p)
|
||||
start := 0
|
||||
for i := 0; i < len(data); i++ {
|
||||
if data[i] == '\n' {
|
||||
line := data[start:i]
|
||||
start = i + 1
|
||||
if len(line) > 0 {
|
||||
r.appendLine(line)
|
||||
}
|
||||
}
|
||||
}
|
||||
// 处理最后一行(没有换行符结尾的情况)
|
||||
if start < len(data) && len(data[start:]) > 0 {
|
||||
r.appendLine(data[start:])
|
||||
}
|
||||
|
||||
return len(p), nil
|
||||
}
|
||||
|
||||
// Sync 实现 zapcore.WriteSyncer 接口
|
||||
func (r *LogRingBuffer) Sync() error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// appendLine 追加一行日志到 ring buffer 并通知订阅者
|
||||
func (r *LogRingBuffer) appendLine(line string) {
|
||||
r.mu.Lock()
|
||||
entry := LogEntry{
|
||||
Index: r.seq,
|
||||
Data: line,
|
||||
}
|
||||
r.entries[r.head] = entry
|
||||
r.head = (r.head + 1) % r.cap
|
||||
if r.count < r.cap {
|
||||
r.count++
|
||||
}
|
||||
r.seq++
|
||||
r.mu.Unlock()
|
||||
|
||||
// 异步通知订阅者
|
||||
r.subMu.RLock()
|
||||
for ch := range r.subscribers {
|
||||
select {
|
||||
case ch <- entry:
|
||||
default:
|
||||
// 订阅者消费太慢,丢弃(避免阻塞日志写入)
|
||||
}
|
||||
}
|
||||
r.subMu.RUnlock()
|
||||
}
|
||||
|
||||
// Query 查询历史日志
|
||||
// cursor=0 表示查询最新日志,cursor>0 表示查询 index < cursor 的更早日志
|
||||
// limit 为返回条数上限
|
||||
// 返回日志条目(按 index 升序)和是否有更早的日志
|
||||
func (r *LogRingBuffer) Query(cursor int, limit int) ([]LogEntry, bool) {
|
||||
r.mu.RLock()
|
||||
defer r.mu.RUnlock()
|
||||
|
||||
if r.count == 0 {
|
||||
return nil, false
|
||||
}
|
||||
|
||||
// 计算 ring buffer 中有效条目的范围
|
||||
// oldest index in ring: head - count (wrapping)
|
||||
oldestPos := (r.head - r.count + r.cap) % r.cap
|
||||
|
||||
// 将 ring buffer 中的有效条目按顺序收集
|
||||
ordered := make([]LogEntry, 0, r.count)
|
||||
for i := 0; i < r.count; i++ {
|
||||
pos := (oldestPos + i) % r.cap
|
||||
ordered = append(ordered, r.entries[pos])
|
||||
}
|
||||
|
||||
if cursor == 0 {
|
||||
// 查询最新日志:返回最后 limit 条
|
||||
if len(ordered) <= limit {
|
||||
return ordered, false
|
||||
}
|
||||
return ordered[len(ordered)-limit:], true
|
||||
}
|
||||
|
||||
// 查询 index < cursor 的更早日志
|
||||
// 找到 index < cursor 的条目
|
||||
var cut int
|
||||
for cut = len(ordered); cut > 0; cut-- {
|
||||
if ordered[cut-1].Index < cursor {
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
if cut == 0 {
|
||||
return nil, false
|
||||
}
|
||||
|
||||
// 返回 cut 之前的最后 limit 条
|
||||
start := cut - limit
|
||||
if start < 0 {
|
||||
start = 0
|
||||
}
|
||||
|
||||
hasMore := start > 0
|
||||
return ordered[start:cut], hasMore
|
||||
}
|
||||
|
||||
// subscribeChanSize 订阅者 channel 缓冲区大小
|
||||
const subscribeChanSize = 64
|
||||
|
||||
// Subscribe 订阅实时日志推送
|
||||
// 返回一个 channel,调用者应 defer Unsubscribe
|
||||
func (r *LogRingBuffer) Subscribe() chan LogEntry {
|
||||
ch := make(chan LogEntry, subscribeChanSize)
|
||||
r.subMu.Lock()
|
||||
r.subscribers[ch] = struct{}{}
|
||||
r.subMu.Unlock()
|
||||
return ch
|
||||
}
|
||||
|
||||
// Unsubscribe 取消订阅
|
||||
func (r *LogRingBuffer) Unsubscribe(ch chan LogEntry) {
|
||||
r.subMu.Lock()
|
||||
delete(r.subscribers, ch)
|
||||
r.subMu.Unlock()
|
||||
close(ch)
|
||||
}
|
||||
|
||||
// 确保 LogRingBuffer 实现 io.Writer 接口
|
||||
var _ io.Writer = (*LogRingBuffer)(nil)
|
||||
@@ -0,0 +1,192 @@
|
||||
// Copyright 2025 linux.do
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package logger
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
)
|
||||
|
||||
func TestLogRingBuffer_WriteAndQuery(t *testing.T) {
|
||||
rb := NewLogRingBuffer(5)
|
||||
|
||||
// Write some logs
|
||||
_, _ = rb.Write([]byte("line1\nline2\nline3\n"))
|
||||
|
||||
entries, hasMore := rb.Query(0, 10)
|
||||
assert.False(t, hasMore)
|
||||
assert.Equal(t, 3, len(entries))
|
||||
assert.Equal(t, "line1", entries[0].Data)
|
||||
assert.Equal(t, "line2", entries[1].Data)
|
||||
assert.Equal(t, "line3", entries[2].Data)
|
||||
assert.Equal(t, 0, entries[0].Index)
|
||||
assert.Equal(t, 1, entries[1].Index)
|
||||
assert.Equal(t, 2, entries[2].Index)
|
||||
}
|
||||
|
||||
func TestLogRingBuffer_CapacityOverflow(t *testing.T) {
|
||||
rb := NewLogRingBuffer(3)
|
||||
|
||||
_, _ = rb.Write([]byte("a\nb\nc\nd\ne\n"))
|
||||
|
||||
entries, hasMore := rb.Query(0, 10)
|
||||
assert.False(t, hasMore)
|
||||
assert.Equal(t, 3, len(entries))
|
||||
assert.Equal(t, "c", entries[0].Data)
|
||||
assert.Equal(t, "d", entries[1].Data)
|
||||
assert.Equal(t, "e", entries[2].Data)
|
||||
}
|
||||
|
||||
func TestLogRingBuffer_QueryLatest(t *testing.T) {
|
||||
rb := NewLogRingBuffer(10)
|
||||
|
||||
_, _ = rb.Write([]byte("a\nb\nc\nd\ne\n"))
|
||||
|
||||
// Query latest 2
|
||||
entries, hasMore := rb.Query(0, 2)
|
||||
assert.True(t, hasMore)
|
||||
assert.Equal(t, 2, len(entries))
|
||||
assert.Equal(t, "d", entries[0].Data)
|
||||
assert.Equal(t, "e", entries[1].Data)
|
||||
}
|
||||
|
||||
func TestLogRingBuffer_QueryByCursor(t *testing.T) {
|
||||
rb := NewLogRingBuffer(10)
|
||||
|
||||
_, _ = rb.Write([]byte("a\nb\nc\nd\ne\n"))
|
||||
|
||||
// First get all to find indices
|
||||
all, _ := rb.Query(0, 10)
|
||||
assert.Equal(t, 5, len(all))
|
||||
|
||||
// Query entries before index 3
|
||||
entries, hasMore := rb.Query(3, 10)
|
||||
assert.False(t, hasMore)
|
||||
assert.Equal(t, 3, len(entries))
|
||||
assert.Equal(t, "a", entries[0].Data)
|
||||
assert.Equal(t, "b", entries[1].Data)
|
||||
assert.Equal(t, "c", entries[2].Data)
|
||||
}
|
||||
|
||||
func TestLogRingBuffer_QueryByCursorWithLimit(t *testing.T) {
|
||||
rb := NewLogRingBuffer(10)
|
||||
|
||||
_, _ = rb.Write([]byte("a\nb\nc\nd\ne\n"))
|
||||
|
||||
// Query 2 entries before index 4
|
||||
entries, hasMore := rb.Query(4, 2)
|
||||
assert.True(t, hasMore)
|
||||
assert.Equal(t, 2, len(entries))
|
||||
assert.Equal(t, "c", entries[0].Data)
|
||||
assert.Equal(t, "d", entries[1].Data)
|
||||
}
|
||||
|
||||
func TestLogRingBuffer_QueryEmpty(t *testing.T) {
|
||||
rb := NewLogRingBuffer(5)
|
||||
|
||||
entries, hasMore := rb.Query(0, 10)
|
||||
assert.False(t, hasMore)
|
||||
assert.Nil(t, entries)
|
||||
}
|
||||
|
||||
func TestLogRingBuffer_QueryNonExistentCursor(t *testing.T) {
|
||||
rb := NewLogRingBuffer(5)
|
||||
_, _ = rb.Write([]byte("a\nb\n"))
|
||||
|
||||
entries, hasMore := rb.Query(999, 10)
|
||||
assert.False(t, hasMore)
|
||||
assert.Equal(t, 2, len(entries))
|
||||
assert.Equal(t, "a", entries[0].Data)
|
||||
assert.Equal(t, "b", entries[1].Data)
|
||||
}
|
||||
|
||||
func TestLogRingBuffer_Subscribe(t *testing.T) {
|
||||
rb := NewLogRingBuffer(5)
|
||||
|
||||
ch := rb.Subscribe()
|
||||
defer rb.Unsubscribe(ch)
|
||||
|
||||
_, _ = rb.Write([]byte("hello\n"))
|
||||
|
||||
entry := <-ch
|
||||
assert.Equal(t, "hello", entry.Data)
|
||||
assert.Equal(t, 0, entry.Index)
|
||||
}
|
||||
|
||||
func TestLogRingBuffer_SubscribeMultiple(t *testing.T) {
|
||||
rb := NewLogRingBuffer(5)
|
||||
|
||||
ch1 := rb.Subscribe()
|
||||
defer rb.Unsubscribe(ch1)
|
||||
ch2 := rb.Subscribe()
|
||||
defer rb.Unsubscribe(ch2)
|
||||
|
||||
_, _ = rb.Write([]byte("msg\n"))
|
||||
|
||||
e1 := <-ch1
|
||||
e2 := <-ch2
|
||||
assert.Equal(t, "msg", e1.Data)
|
||||
assert.Equal(t, "msg", e2.Data)
|
||||
}
|
||||
|
||||
func TestLogRingBuffer_WriteNoNewline(t *testing.T) {
|
||||
rb := NewLogRingBuffer(5)
|
||||
|
||||
_, _ = rb.Write([]byte("partial"))
|
||||
|
||||
entries, _ := rb.Query(0, 10)
|
||||
assert.Equal(t, 1, len(entries))
|
||||
assert.Equal(t, "partial", entries[0].Data)
|
||||
}
|
||||
|
||||
func TestLogRingBuffer_WriteEmpty(t *testing.T) {
|
||||
rb := NewLogRingBuffer(5)
|
||||
|
||||
n, err := rb.Write([]byte(""))
|
||||
assert.Equal(t, 0, n)
|
||||
assert.NoError(t, err)
|
||||
|
||||
entries, _ := rb.Query(0, 10)
|
||||
assert.Nil(t, entries)
|
||||
}
|
||||
|
||||
func TestLogRingBuffer_QueryAfterOverflow(t *testing.T) {
|
||||
rb := NewLogRingBuffer(3)
|
||||
|
||||
_, _ = rb.Write([]byte("1\n2\n3\n4\n5\n6\n7\n"))
|
||||
|
||||
entries, hasMore := rb.Query(0, 10)
|
||||
assert.False(t, hasMore)
|
||||
assert.Equal(t, 3, len(entries))
|
||||
assert.Equal(t, "5", entries[0].Data)
|
||||
assert.Equal(t, "6", entries[1].Data)
|
||||
assert.Equal(t, "7", entries[2].Data)
|
||||
|
||||
// Query by cursor - index 4 is "5", so cursor=4 should return index < 4
|
||||
older, hasMore2 := rb.Query(4, 10)
|
||||
assert.False(t, hasMore2)
|
||||
assert.Nil(t, older)
|
||||
}
|
||||
|
||||
func TestLogRingBuffer_NextCursor(t *testing.T) {
|
||||
rb := NewLogRingBuffer(10)
|
||||
|
||||
_, _ = rb.Write([]byte("a\nb\nc\nd\ne\n"))
|
||||
|
||||
// Query latest 2, should return next_cursor pointing to first returned entry
|
||||
entries, _ := rb.Query(0, 2)
|
||||
assert.Equal(t, 2, len(entries))
|
||||
// entries[0].Index = 3 ("d"), entries[1].Index = 4 ("e")
|
||||
assert.Equal(t, 3, entries[0].Index)
|
||||
|
||||
// Now use that index as cursor to get older entries
|
||||
older, hasMore := rb.Query(entries[0].Index, 10)
|
||||
assert.False(t, hasMore)
|
||||
assert.Equal(t, 3, len(older))
|
||||
assert.Equal(t, "a", older[0].Data)
|
||||
assert.Equal(t, "b", older[1].Data)
|
||||
assert.Equal(t, "c", older[2].Data)
|
||||
}
|
||||
@@ -0,0 +1,99 @@
|
||||
// Copyright 2025 linux.do
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package logger
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"log"
|
||||
"os"
|
||||
"path/filepath"
|
||||
|
||||
"go.opentelemetry.io/otel/trace"
|
||||
"go.uber.org/zap"
|
||||
"go.uber.org/zap/zapcore"
|
||||
"gopkg.in/natefinch/lumberjack.v2"
|
||||
)
|
||||
|
||||
// logDirPerm 日志目录权限
|
||||
const logDirPerm = 0750
|
||||
|
||||
func getLogWriterForConfig(cfg Config) (zapcore.WriteSyncer, error) {
|
||||
if cfg.Output == "file" {
|
||||
// 初始化日志目录
|
||||
logPath := cfg.FilePath
|
||||
logDir := filepath.Dir(logPath)
|
||||
if err := os.MkdirAll(logDir, logDirPerm); err != nil {
|
||||
return nil, fmt.Errorf(errCreateLogFileDirFailed, err)
|
||||
}
|
||||
|
||||
// 配置日志轮转
|
||||
logOutput := &lumberjack.Logger{
|
||||
Filename: logPath,
|
||||
MaxSize: cfg.MaxSize,
|
||||
MaxBackups: cfg.MaxBackups,
|
||||
MaxAge: cfg.MaxAge,
|
||||
Compress: cfg.Compress,
|
||||
}
|
||||
|
||||
return zapcore.AddSync(logOutput), nil
|
||||
}
|
||||
|
||||
return zapcore.AddSync(os.Stdout), nil
|
||||
}
|
||||
|
||||
// getEncoderForConfig 获取日志编码器
|
||||
func getEncoderForConfig(cfg Config) zapcore.Encoder {
|
||||
// 编码器配置
|
||||
encoderConfig := zapcore.EncoderConfig{
|
||||
TimeKey: "time",
|
||||
LevelKey: "level",
|
||||
NameKey: "logger",
|
||||
CallerKey: "caller",
|
||||
MessageKey: "msg",
|
||||
StacktraceKey: "stacktrace",
|
||||
LineEnding: zapcore.DefaultLineEnding,
|
||||
EncodeLevel: zapcore.LowercaseLevelEncoder,
|
||||
EncodeTime: zapcore.ISO8601TimeEncoder,
|
||||
EncodeDuration: zapcore.SecondsDurationEncoder,
|
||||
EncodeCaller: zapcore.ShortCallerEncoder,
|
||||
}
|
||||
|
||||
if cfg.Format == "json" {
|
||||
return zapcore.NewJSONEncoder(encoderConfig)
|
||||
}
|
||||
return zapcore.NewConsoleEncoder(encoderConfig)
|
||||
}
|
||||
|
||||
// getLogLevelForConfig 获取日志级别
|
||||
func getLogLevelForConfig(cfg Config) zapcore.Level {
|
||||
level := cfg.Level
|
||||
|
||||
switch level {
|
||||
case "debug":
|
||||
return zapcore.DebugLevel
|
||||
case "info":
|
||||
return zapcore.InfoLevel
|
||||
case "warn":
|
||||
return zapcore.WarnLevel
|
||||
case "error":
|
||||
return zapcore.ErrorLevel
|
||||
default:
|
||||
log.Printf("[Logger] invalid log level: %s, defaulting to info\n", level)
|
||||
return zapcore.InfoLevel
|
||||
}
|
||||
}
|
||||
|
||||
func getTraceIDFields(ctx context.Context) []zap.Field {
|
||||
span := trace.SpanFromContext(ctx)
|
||||
spanContext := span.SpanContext()
|
||||
if !spanContext.IsValid() {
|
||||
return nil
|
||||
}
|
||||
return []zap.Field{
|
||||
zap.String("traceID", spanContext.TraceID().String()),
|
||||
zap.String("spanID", spanContext.SpanID().String()),
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,16 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package mail 提供 SMTP 邮件发送功能。
|
||||
package mail
|
||||
|
||||
const (
|
||||
errDialTLSFailed = "dial tls failed: %w"
|
||||
errSMTPClientCreationFailed = "smtp client creation failed: %w"
|
||||
errSMTPAuthFailed = "smtp auth failed: %w"
|
||||
errSMTPMailCommandFailed = "smtp mail command failed: %w"
|
||||
errSMTPRcptCommandFailed = "smtp rcpt command failed: %w"
|
||||
errSMTPDataCommandFailed = "smtp data command failed: %w"
|
||||
errSMTPWritingBodyFailed = "smtp writing body failed: %w" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
|
||||
errSendMailFailed = "send mail failed: %w"
|
||||
)
|
||||
@@ -0,0 +1,241 @@
|
||||
// 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"
|
||||
"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
|
||||
}
|
||||
|
||||
// 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"] = cfg.Username
|
||||
header["To"] = to
|
||||
header["Subject"] = 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"] = cfg.Username
|
||||
header["To"] = to
|
||||
header["Subject"] = subject
|
||||
header["MIME-Version"] = "1.0"
|
||||
header["Content-Type"] = "text/html; charset=UTF-8"
|
||||
|
||||
message := ""
|
||||
for k, v := range header {
|
||||
message += fmt.Sprintf("%s: %s\r\n", k, v)
|
||||
}
|
||||
message += "\r\n" + body
|
||||
|
||||
logLine("System", "Sending message body...")
|
||||
if _, err = w.Write([]byte(message)); err != nil {
|
||||
_ = w.Close()
|
||||
logLine("Error", "Writing message body failed: %v", err)
|
||||
return logBuf.String(), err
|
||||
}
|
||||
_ = w.Close()
|
||||
logLine("S", "250 OK")
|
||||
|
||||
logLine("C", "QUIT")
|
||||
_ = client.Quit()
|
||||
logLine("System", "Mail sent successfully!")
|
||||
|
||||
return logBuf.String(), nil
|
||||
}
|
||||
@@ -0,0 +1,92 @@
|
||||
// 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)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,81 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package push
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"strings"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/pkg/httppool"
|
||||
)
|
||||
|
||||
func init() {
|
||||
Register("custom", &CustomPusher{})
|
||||
}
|
||||
|
||||
// CustomPusher 自定义 Webhook 发送实现
|
||||
type CustomPusher struct{}
|
||||
|
||||
// Send 发送自定义 webhook
|
||||
func (p *CustomPusher) Send(ctx context.Context, cfg Config, _ string, body map[string]any, template string, _ map[string]any) error {
|
||||
if cfg.URL == "" {
|
||||
return errors.New("custom: URL is required")
|
||||
}
|
||||
|
||||
var reqBody []byte
|
||||
|
||||
if template != "" {
|
||||
// 替换模板中的 {{key}} 占位符
|
||||
rendered := ParseTemplate(template, body)
|
||||
reqBody = []byte(rendered)
|
||||
} else {
|
||||
// 兜底:直接把 body 转为 JSON 字符串发送
|
||||
var err error
|
||||
reqBody, err = json.Marshal(body)
|
||||
if err != nil {
|
||||
return fmt.Errorf("custom: marshal body failed: %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
httpReq, err := http.NewRequestWithContext(ctx, http.MethodPost, cfg.URL, bytes.NewReader(reqBody))
|
||||
if err != nil {
|
||||
return fmt.Errorf("custom: create http request failed: %w", err)
|
||||
}
|
||||
httpReq.Header.Set("Content-Type", "application/json")
|
||||
|
||||
// 如果配置了 Key 且格式为 "HeaderName:HeaderValue",我们可以附加测试用 Header
|
||||
if cfg.Key != "" && strings.Contains(cfg.Key, ":") {
|
||||
parts := strings.SplitN(cfg.Key, ":", 2) //nolint:mnd
|
||||
httpReq.Header.Set(strings.TrimSpace(parts[0]), strings.TrimSpace(parts[1]))
|
||||
}
|
||||
|
||||
client := httppool.NewClient(defaultHTTPClientTimeout)
|
||||
resp, err := client.Do(httpReq)
|
||||
if err != nil {
|
||||
return fmt.Errorf("custom: http request failed: %w", err)
|
||||
}
|
||||
defer func() { _ = resp.Body.Close() }()
|
||||
|
||||
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
|
||||
return fmt.Errorf("custom: http status %s", resp.Status)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// ValidateConfig 校验自定义配置
|
||||
func (p *CustomPusher) ValidateConfig(cfg Config) error {
|
||||
if cfg.URL == "" {
|
||||
return errors.New("webhook URL is required")
|
||||
}
|
||||
if !strings.HasPrefix(cfg.URL, "http://") && !strings.HasPrefix(cfg.URL, "https://") {
|
||||
return errors.New("webhook URL must start with http:// or https://")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,109 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package push
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net"
|
||||
"net/smtp"
|
||||
"strings"
|
||||
)
|
||||
|
||||
func init() {
|
||||
Register("email", &EmailPusher{})
|
||||
}
|
||||
|
||||
// EmailPusher 极简 SMTP 邮件推送实现 (静态、解耦)
|
||||
type EmailPusher struct{}
|
||||
|
||||
// Send 发送邮件
|
||||
func (p *EmailPusher) Send(ctx context.Context, cfg Config, target string, body map[string]any, _ string, ext map[string]any) error {
|
||||
if cfg.URL == "" || cfg.Key == "" || cfg.Secret == "" {
|
||||
return errors.New("email: SMTP configuration (url, key, secret) is incomplete")
|
||||
}
|
||||
if target == "" {
|
||||
return errors.New("email: target email address is required")
|
||||
}
|
||||
|
||||
title := defaultTitle
|
||||
if t, ok := body["title"].(string); ok && t != "" {
|
||||
title = t
|
||||
}
|
||||
|
||||
content := ""
|
||||
if c, ok := body["content"].(string); ok && c != "" {
|
||||
content = c
|
||||
} else {
|
||||
// 自动格式化 map
|
||||
var parts []string
|
||||
for k, v := range body {
|
||||
parts = append(parts, fmt.Sprintf("<p><b>%s</b>: %v</p>", k, v))
|
||||
}
|
||||
content = strings.Join(parts, "")
|
||||
}
|
||||
|
||||
// 邮件头和体
|
||||
from := cfg.Key
|
||||
to := target
|
||||
|
||||
// 如果 ext 中指定了 from_name,我们在 From 头部包含它
|
||||
fromName := "System Notification"
|
||||
if ext != nil {
|
||||
if fn, ok := ext["from_name"].(string); ok && fn != "" {
|
||||
fromName = fn
|
||||
}
|
||||
}
|
||||
|
||||
subjectHeader := fmt.Sprintf("Subject: %s\r\n", title)
|
||||
fromHeader := fmt.Sprintf("From: %s <%s>\r\n", fromName, from)
|
||||
toHeader := fmt.Sprintf("To: %s\r\n", to)
|
||||
mimeHeader := "MIME-version: 1.0;\r\nContent-Type: text/html; charset=\"UTF-8\";\r\n\r\n"
|
||||
|
||||
// 拼装完整的邮件报文
|
||||
// 简单的 HTML 正文渲染
|
||||
htmlBody := fmt.Sprintf(`<html><body><h2>%s</h2><div>%s</div></body></html>`, title, content)
|
||||
msg := []byte(fromHeader + toHeader + subjectHeader + mimeHeader + htmlBody + "\r\n")
|
||||
|
||||
// 解析 Host 和 Port
|
||||
host, port, err := net.SplitHostPort(cfg.URL)
|
||||
if err != nil {
|
||||
host = cfg.URL
|
||||
port = "25" // 默认 SMTP 端口
|
||||
}
|
||||
|
||||
auth := smtp.PlainAuth("", cfg.Key, cfg.Secret, host)
|
||||
|
||||
// 异步超时处理
|
||||
errChan := make(chan error, 1)
|
||||
go func() {
|
||||
errChan <- smtp.SendMail(host+":"+port, auth, from, []string{to}, msg)
|
||||
}()
|
||||
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return ctx.Err()
|
||||
case err := <-errChan:
|
||||
if err != nil {
|
||||
return fmt.Errorf("email: send smtp mail failed: %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// ValidateConfig 校验邮件 SMTP 配置
|
||||
func (p *EmailPusher) ValidateConfig(cfg Config) error {
|
||||
if cfg.URL == "" {
|
||||
return errors.New("SMTP host:port is required")
|
||||
}
|
||||
if cfg.Key == "" {
|
||||
return errors.New("SMTP username is required")
|
||||
}
|
||||
if cfg.Secret == "" {
|
||||
return errors.New("SMTP password is required")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,275 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package push
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"crypto/hmac"
|
||||
"crypto/sha256"
|
||||
"encoding/base64"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/pkg/httppool"
|
||||
)
|
||||
|
||||
func init() {
|
||||
Register("lark", &LarkPusher{})
|
||||
}
|
||||
|
||||
const (
|
||||
msgTypeInteractive = "interactive"
|
||||
)
|
||||
|
||||
// LarkPusher 飞书 Webhook 机器人推送实现
|
||||
type LarkPusher struct{}
|
||||
|
||||
type larkTextContent struct {
|
||||
Text string `json:"text"`
|
||||
}
|
||||
|
||||
type larkCardHeaderTitle struct {
|
||||
Content string `json:"content"`
|
||||
Tag string `json:"tag"`
|
||||
}
|
||||
|
||||
type larkCardHeader struct {
|
||||
Template string `json:"template"` // "blue", "orange", "red" etc.
|
||||
Title larkCardHeaderTitle `json:"title"`
|
||||
}
|
||||
|
||||
type larkCardElementText struct {
|
||||
Content string `json:"content"`
|
||||
Tag string `json:"tag"` // "lark_md"
|
||||
}
|
||||
|
||||
type larkCardElement struct {
|
||||
Tag string `json:"tag"` // "div"
|
||||
Text larkCardElementText `json:"text"`
|
||||
}
|
||||
|
||||
type larkCardContent struct {
|
||||
Header larkCardHeader `json:"header"`
|
||||
Elements []larkCardElement `json:"elements"`
|
||||
}
|
||||
|
||||
type larkMessageRequest struct {
|
||||
MessageType string `json:"msg_type"`
|
||||
Timestamp string `json:"timestamp,omitempty"`
|
||||
Sign string `json:"sign,omitempty"`
|
||||
Content larkTextContent `json:"content,omitempty"`
|
||||
Card *larkCardContent `json:"card,omitempty"`
|
||||
}
|
||||
|
||||
type larkMessageResponse struct {
|
||||
Code int `json:"code"`
|
||||
Msg string `json:"msg"`
|
||||
}
|
||||
|
||||
// Send 执行飞书消息发送
|
||||
//
|
||||
//nolint:nestif,cyclop
|
||||
func (p *LarkPusher) Send(ctx context.Context, cfg Config, _ string, body map[string]any, template string, _ map[string]any) error {
|
||||
if cfg.URL == "" {
|
||||
return errors.New("lark: URL is required")
|
||||
}
|
||||
|
||||
var req larkMessageRequest
|
||||
|
||||
// 1. 如果有自定义模板,我们尝试进行解析
|
||||
if template != "" {
|
||||
rendered := ParseTemplate(template, body)
|
||||
|
||||
// 尝试解析原生的 Lark Card
|
||||
var customCard larkCardContent
|
||||
var rawMap map[string]any
|
||||
_ = json.Unmarshal([]byte(rendered), &rawMap)
|
||||
|
||||
if rawMap != nil && rawMap["elements"] != nil {
|
||||
// 如果包含 elements 字段,说明是用户定制的原生飞书卡片 JSON
|
||||
if err := json.Unmarshal([]byte(rendered), &customCard); err == nil {
|
||||
req.MessageType = msgTypeInteractive
|
||||
req.Card = &customCard
|
||||
} else {
|
||||
req.MessageType = "text"
|
||||
req.Content.Text = rendered
|
||||
}
|
||||
} else {
|
||||
// 说明配置的是系统统一通知消息 of JSON 模板:{"title": "...", "content": "...", "level": "..."}
|
||||
type larkNotificationMessage struct {
|
||||
Title string `json:"title"`
|
||||
Content string `json:"content"`
|
||||
Level string `json:"level"`
|
||||
}
|
||||
var msg larkNotificationMessage
|
||||
if err := json.Unmarshal([]byte(rendered), &msg); err == nil && (msg.Title != "" || msg.Content != "") {
|
||||
title := msg.Title
|
||||
if title == "" {
|
||||
title = defaultTitle
|
||||
}
|
||||
content := msg.Content
|
||||
level := strings.ToUpper(msg.Level)
|
||||
if level == "" {
|
||||
level = levelInfo
|
||||
}
|
||||
|
||||
headerColor := "blue"
|
||||
switch level {
|
||||
case "IMPORTANT":
|
||||
headerColor = "orange"
|
||||
case "CRITICAL":
|
||||
headerColor = "red"
|
||||
}
|
||||
|
||||
req.MessageType = msgTypeInteractive
|
||||
req.Card = &larkCardContent{
|
||||
Header: larkCardHeader{
|
||||
Template: headerColor,
|
||||
Title: larkCardHeaderTitle{
|
||||
Content: title,
|
||||
Tag: "plain_text",
|
||||
},
|
||||
},
|
||||
Elements: []larkCardElement{
|
||||
{
|
||||
Tag: "div",
|
||||
Text: larkCardElementText{
|
||||
Content: content,
|
||||
Tag: "lark_md",
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
} else {
|
||||
// 兜底:如果无法按 JSON 解析出结构化字段,当做普通文本发送
|
||||
req.MessageType = "text"
|
||||
req.Content.Text = rendered
|
||||
}
|
||||
}
|
||||
} else {
|
||||
// 2. 如果无模板,默认生成一个精美的飞书互动卡片
|
||||
title := defaultTitle
|
||||
if t, ok := body["title"].(string); ok && t != "" {
|
||||
title = t
|
||||
}
|
||||
|
||||
content := ""
|
||||
if c, ok := body["content"].(string); ok && c != "" {
|
||||
content = c
|
||||
} else {
|
||||
// 兜底:如果连 content 都没有,把 body 里的所有值拼成 markdown
|
||||
var parts []string
|
||||
for k, v := range body {
|
||||
parts = append(parts, fmt.Sprintf("**%s**: %v", k, v))
|
||||
}
|
||||
content = strings.Join(parts, "\n")
|
||||
}
|
||||
|
||||
level := levelInfo
|
||||
if l, ok := body["level"].(string); ok && l != "" {
|
||||
level = strings.ToUpper(l)
|
||||
}
|
||||
|
||||
// 根据级别确定飞书卡片头部的背景色模板
|
||||
headerColor := "blue"
|
||||
switch level {
|
||||
case "IMPORTANT":
|
||||
headerColor = "orange"
|
||||
case "CRITICAL":
|
||||
headerColor = "red"
|
||||
}
|
||||
|
||||
req.MessageType = msgTypeInteractive
|
||||
req.Card = &larkCardContent{
|
||||
Header: larkCardHeader{
|
||||
Template: headerColor,
|
||||
Title: larkCardHeaderTitle{
|
||||
Content: title,
|
||||
Tag: "plain_text",
|
||||
},
|
||||
},
|
||||
Elements: []larkCardElement{
|
||||
{
|
||||
Tag: "div",
|
||||
Text: larkCardElementText{
|
||||
Content: content,
|
||||
Tag: "lark_md",
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
// 3. 计算签名 (如果配置了 secret)
|
||||
if cfg.Secret != "" {
|
||||
timestamp := time.Now().Unix()
|
||||
sign, err := larkSign(cfg.Secret, timestamp)
|
||||
if err != nil {
|
||||
return fmt.Errorf("lark: sign failed: %w", err)
|
||||
}
|
||||
req.Timestamp = strconv.FormatInt(timestamp, 10)
|
||||
req.Sign = sign
|
||||
}
|
||||
|
||||
jsonData, err := json.Marshal(req)
|
||||
if err != nil {
|
||||
return fmt.Errorf("lark: marshal request failed: %w", err)
|
||||
}
|
||||
|
||||
// 4. 发送 POST 请求
|
||||
httpReq, err := http.NewRequestWithContext(ctx, http.MethodPost, cfg.URL, bytes.NewBuffer(jsonData))
|
||||
if err != nil {
|
||||
return fmt.Errorf("lark: create http request failed: %w", err)
|
||||
}
|
||||
httpReq.Header.Set("Content-Type", "application/json")
|
||||
|
||||
client := httppool.NewClient(defaultHTTPClientTimeout)
|
||||
resp, err := client.Do(httpReq)
|
||||
if err != nil {
|
||||
return fmt.Errorf("lark: http request failed: %w", err)
|
||||
}
|
||||
defer func() { _ = resp.Body.Close() }()
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return fmt.Errorf("lark: http status %s", resp.Status)
|
||||
}
|
||||
|
||||
var res larkMessageResponse
|
||||
if err := json.NewDecoder(resp.Body).Decode(&res); err != nil {
|
||||
return fmt.Errorf("lark: decode response failed: %w", err)
|
||||
}
|
||||
|
||||
if res.Code != 0 {
|
||||
return fmt.Errorf("lark: send message failed, code %d: %s", res.Code, res.Msg)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// ValidateConfig 校验飞书配置
|
||||
func (p *LarkPusher) ValidateConfig(cfg Config) error {
|
||||
if cfg.URL == "" {
|
||||
return errors.New("webhook URL is required")
|
||||
}
|
||||
if !strings.HasPrefix(cfg.URL, "http://") && !strings.HasPrefix(cfg.URL, "https://") {
|
||||
return errors.New("webhook URL must start with http:// or https://")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func larkSign(secret string, timestamp int64) (string, error) {
|
||||
stringToSign := fmt.Sprintf("%v", timestamp) + "\n" + secret
|
||||
h := hmac.New(sha256.New, []byte(stringToSign))
|
||||
_, err := h.Write(nil)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return base64.StdEncoding.EncodeToString(h.Sum(nil)), nil
|
||||
}
|
||||
@@ -0,0 +1,66 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package push 提供解耦的、无外部业务依赖 of 通知推送底层实现
|
||||
package push
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"sync"
|
||||
"time"
|
||||
)
|
||||
|
||||
const (
|
||||
defaultTitle = "系统通知"
|
||||
levelInfo = "INFO"
|
||||
defaultHTTPClientTimeout = 10 * time.Second
|
||||
)
|
||||
|
||||
// Config 基础通知渠道配置
|
||||
type Config struct {
|
||||
Channel string `json:"channel"` // 渠道名称,例如 "lark", "custom", "email" 等,唯一标识
|
||||
URL string `json:"url,omitempty"` // Webhook 地址或 SMTP 地址
|
||||
Secret string `json:"secret,omitempty"` // 签名密钥或 SMTP 密码/Token
|
||||
Key string `json:"key,omitempty"` // AppID 或 SMTP 用户名
|
||||
Ext map[string]any `json:"ext,omitempty"` // 预留拓展 JSON 配置
|
||||
}
|
||||
|
||||
// Pusher 通知推送渠道接口
|
||||
type Pusher interface {
|
||||
// Send 发送通知消息
|
||||
// target: 发送目标 (如邮箱地址或特定用户标识;若为 bot 机器人此项为空)
|
||||
// body: 消息体数据 (含默认字段如 title, content, level)
|
||||
// template: 消息卡片/模板 JSON (可选)
|
||||
// ext: 预留的单次发送拓展数据
|
||||
Send(ctx context.Context, cfg Config, target string, body map[string]any, template string, ext map[string]any) error
|
||||
|
||||
// ValidateConfig 校验渠道配置合法性
|
||||
ValidateConfig(cfg Config) error
|
||||
}
|
||||
|
||||
var (
|
||||
pushersMu sync.RWMutex
|
||||
pushers = make(map[string]Pusher)
|
||||
)
|
||||
|
||||
// Register 注册一个推送渠道实现
|
||||
func Register(channelType string, pusher Pusher) {
|
||||
pushersMu.Lock()
|
||||
defer pushersMu.Unlock()
|
||||
if pusher == nil {
|
||||
panic("push: Register pusher is nil")
|
||||
}
|
||||
pushers[channelType] = pusher
|
||||
}
|
||||
|
||||
// GetPusher 获取指定类型的推送渠道实现
|
||||
func GetPusher(channelType string) (Pusher, error) {
|
||||
pushersMu.RLock()
|
||||
defer pushersMu.RUnlock()
|
||||
pusher, ok := pushers[channelType]
|
||||
if !ok {
|
||||
return nil, fmt.Errorf("push: unknown channel type %q", channelType)
|
||||
}
|
||||
return pusher, nil
|
||||
}
|
||||
@@ -0,0 +1,158 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package push
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"strings"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/pkg/httppool"
|
||||
)
|
||||
|
||||
func init() {
|
||||
Register("telegram", &TelegramPusher{})
|
||||
}
|
||||
|
||||
// TelegramPusher Telegram 机器人推送实现
|
||||
type TelegramPusher struct{}
|
||||
|
||||
type telegramMessageRequest struct {
|
||||
ChatID string `json:"chat_id"`
|
||||
Text string `json:"text"`
|
||||
ParseMode string `json:"parse_mode,omitempty"`
|
||||
}
|
||||
|
||||
type telegramErrorResponse struct {
|
||||
Ok bool `json:"ok"`
|
||||
ErrorCode int `json:"error_code"`
|
||||
Description string `json:"description"`
|
||||
}
|
||||
|
||||
// Send 执行 Telegram 消息发送
|
||||
//
|
||||
//nolint:cyclop
|
||||
func (p *TelegramPusher) Send(ctx context.Context, cfg Config, target string, body map[string]any, template string, _ map[string]any) error {
|
||||
if cfg.Secret == "" {
|
||||
return errors.New("telegram: Bot Token (Secret) is required")
|
||||
}
|
||||
|
||||
chatID := target
|
||||
if chatID == "" {
|
||||
chatID = cfg.Key // Use default chat ID (Key) if target is blank
|
||||
}
|
||||
if chatID == "" {
|
||||
return errors.New("telegram: chat_id (target or default Key) is required")
|
||||
}
|
||||
|
||||
baseURL := cfg.URL
|
||||
if baseURL == "" {
|
||||
baseURL = "https://api.telegram.org"
|
||||
}
|
||||
baseURL = strings.TrimSuffix(baseURL, "/")
|
||||
|
||||
title := defaultTitle
|
||||
if t, ok := body["title"].(string); ok && t != "" {
|
||||
title = t
|
||||
}
|
||||
content := ""
|
||||
if c, ok := body["content"].(string); ok && c != "" {
|
||||
content = c
|
||||
} else {
|
||||
var parts []string
|
||||
for k, v := range body {
|
||||
parts = append(parts, fmt.Sprintf("<b>%s</b>: %v", k, v))
|
||||
}
|
||||
content = strings.Join(parts, "\n")
|
||||
}
|
||||
level := levelInfo
|
||||
if l, ok := body["level"].(string); ok && l != "" {
|
||||
level = strings.ToUpper(l)
|
||||
}
|
||||
|
||||
var text string
|
||||
if template != "" {
|
||||
text = ParseTemplate(template, body)
|
||||
} else {
|
||||
text = fmt.Sprintf("<b>[%s] %s</b>\n\n%s", escapeHTML(level), escapeHTML(title), escapeHTML(content))
|
||||
}
|
||||
|
||||
// Try sending with HTML parse mode
|
||||
err := p.sendMessage(ctx, baseURL, cfg.Secret, chatID, text, "HTML")
|
||||
if err != nil {
|
||||
// Fallback: send as plain text without parse mode
|
||||
plainText := text
|
||||
if template == "" {
|
||||
plainText = fmt.Sprintf("[%s] %s\n\n%s", level, title, content)
|
||||
}
|
||||
fallbackErr := p.sendMessage(ctx, baseURL, cfg.Secret, chatID, plainText, "")
|
||||
if fallbackErr != nil {
|
||||
return fmt.Errorf("telegram: send message failed (fallback also failed): %w (original HTML error: %v)", fallbackErr, err)
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// ValidateConfig 校验 Telegram 配置
|
||||
func (p *TelegramPusher) ValidateConfig(cfg Config) error {
|
||||
if cfg.Secret == "" {
|
||||
return errors.New("bot Token (Secret) is required")
|
||||
}
|
||||
if cfg.URL != "" {
|
||||
if !strings.HasPrefix(cfg.URL, "http://") && !strings.HasPrefix(cfg.URL, "https://") {
|
||||
return errors.New("API base URL must start with http:// or https://")
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (p *TelegramPusher) sendMessage(ctx context.Context, baseURL, token, chatID, text, parseMode string) error {
|
||||
apiURL := fmt.Sprintf("%s/bot%s/sendMessage", baseURL, token)
|
||||
|
||||
reqPayload := telegramMessageRequest{
|
||||
ChatID: chatID,
|
||||
Text: text,
|
||||
ParseMode: parseMode,
|
||||
}
|
||||
|
||||
jsonData, err := json.Marshal(reqPayload)
|
||||
if err != nil {
|
||||
return fmt.Errorf("marshal request failed: %w", err)
|
||||
}
|
||||
|
||||
httpReq, err := http.NewRequestWithContext(ctx, http.MethodPost, apiURL, bytes.NewBuffer(jsonData))
|
||||
if err != nil {
|
||||
return fmt.Errorf("create http request failed: %w", err)
|
||||
}
|
||||
httpReq.Header.Set("Content-Type", "application/json")
|
||||
|
||||
client := httppool.NewClient(defaultHTTPClientTimeout)
|
||||
resp, err := client.Do(httpReq)
|
||||
if err != nil {
|
||||
return fmt.Errorf("http request failed: %w", err)
|
||||
}
|
||||
defer func() { _ = resp.Body.Close() }()
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
var errRes telegramErrorResponse
|
||||
if decodeErr := json.NewDecoder(resp.Body).Decode(&errRes); decodeErr == nil {
|
||||
return fmt.Errorf("http status %d: %s", resp.StatusCode, errRes.Description)
|
||||
}
|
||||
return fmt.Errorf("http status %s", resp.Status)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func escapeHTML(s string) string {
|
||||
s = strings.ReplaceAll(s, "&", "&")
|
||||
s = strings.ReplaceAll(s, "<", "<")
|
||||
s = strings.ReplaceAll(s, ">", ">")
|
||||
return s
|
||||
}
|
||||
@@ -0,0 +1,116 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package push
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestTelegramPusher_Send(t *testing.T) {
|
||||
t.Run("successful send with HTML parse mode", func(t *testing.T) {
|
||||
var receivedReq telegramMessageRequest
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
assert.Equal(t, "/botmy-token/sendMessage", r.URL.Path)
|
||||
assert.Equal(t, http.MethodPost, r.Method)
|
||||
assert.Equal(t, "application/json", r.Header.Get("Content-Type"))
|
||||
|
||||
err := json.NewDecoder(r.Body).Decode(&receivedReq)
|
||||
require.NoError(t, err)
|
||||
|
||||
w.WriteHeader(http.StatusOK)
|
||||
_, _ = w.Write([]byte(`{"ok": true}`))
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
pusher := &TelegramPusher{}
|
||||
cfg := Config{
|
||||
Channel: "telegram",
|
||||
URL: server.URL,
|
||||
Secret: "my-token",
|
||||
}
|
||||
body := map[string]any{
|
||||
"title": "Alert",
|
||||
"content": "Host down",
|
||||
"level": "CRITICAL",
|
||||
}
|
||||
err := pusher.Send(context.Background(), cfg, "123456", body, "", nil)
|
||||
require.NoError(t, err)
|
||||
|
||||
assert.Equal(t, "123456", receivedReq.ChatID)
|
||||
assert.Contains(t, receivedReq.Text, "[CRITICAL] Alert")
|
||||
assert.Contains(t, receivedReq.Text, "Host down")
|
||||
assert.Equal(t, "HTML", receivedReq.ParseMode)
|
||||
})
|
||||
|
||||
t.Run("fallback to plain text on HTML error", func(t *testing.T) {
|
||||
var requests []*telegramMessageRequest
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
var req telegramMessageRequest
|
||||
err := json.NewDecoder(r.Body).Decode(&req)
|
||||
require.NoError(t, err)
|
||||
requests = append(requests, &req)
|
||||
|
||||
if len(requests) == 1 {
|
||||
w.WriteHeader(http.StatusBadRequest)
|
||||
_, _ = w.Write([]byte(`{"ok": false, "error_code": 400, "description": "Bad Request: can't parse entities"}`))
|
||||
} else {
|
||||
w.WriteHeader(http.StatusOK)
|
||||
_, _ = w.Write([]byte(`{"ok": true}`))
|
||||
}
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
pusher := &TelegramPusher{}
|
||||
cfg := Config{
|
||||
Channel: "telegram",
|
||||
URL: server.URL,
|
||||
Secret: "my-token",
|
||||
}
|
||||
body := map[string]any{
|
||||
"title": "Alert & Info",
|
||||
"content": "A < B comparison",
|
||||
"level": "INFO",
|
||||
}
|
||||
err := pusher.Send(context.Background(), cfg, "123456", body, "", nil)
|
||||
require.NoError(t, err)
|
||||
|
||||
require.Len(t, requests, 2)
|
||||
assert.Equal(t, "HTML", requests[0].ParseMode)
|
||||
assert.Equal(t, "", requests[1].ParseMode)
|
||||
assert.Contains(t, requests[1].Text, "[INFO] Alert & Info")
|
||||
assert.Contains(t, requests[1].Text, "A < B comparison")
|
||||
})
|
||||
|
||||
t.Run("validation error", func(t *testing.T) {
|
||||
pusher := &TelegramPusher{}
|
||||
cfg := Config{
|
||||
Channel: "telegram",
|
||||
URL: "https://api.telegram.org",
|
||||
}
|
||||
err := pusher.ValidateConfig(cfg)
|
||||
assert.Error(t, err)
|
||||
|
||||
cfg = Config{
|
||||
Channel: "telegram",
|
||||
URL: "ftp://api.telegram.org",
|
||||
Secret: "token",
|
||||
}
|
||||
err = pusher.ValidateConfig(cfg)
|
||||
assert.Error(t, err)
|
||||
|
||||
cfg = Config{
|
||||
Channel: "telegram",
|
||||
Secret: "token",
|
||||
}
|
||||
err = pusher.ValidateConfig(cfg)
|
||||
assert.NoError(t, err)
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,78 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package push
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"strconv"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// ParseTemplate parses template strings by replacing {{placeholder}} structures with values from body.
|
||||
// It is a single-pass parser designed for high performance and low allocations.
|
||||
func ParseTemplate(template string, body map[string]any) string {
|
||||
var buf strings.Builder
|
||||
buf.Grow(len(template))
|
||||
|
||||
i := 0
|
||||
for {
|
||||
pos := strings.Index(template[i:], "{{")
|
||||
if pos == -1 {
|
||||
buf.WriteString(template[i:])
|
||||
break
|
||||
}
|
||||
// Write prefix
|
||||
buf.WriteString(template[i : i+pos])
|
||||
i += pos + 2 // skip "{{"
|
||||
|
||||
endPos := strings.Index(template[i:], "}}")
|
||||
if endPos == -1 {
|
||||
// Unbalanced "{{"
|
||||
buf.WriteString("{{")
|
||||
buf.WriteString(template[i:])
|
||||
break
|
||||
}
|
||||
key := template[i : i+endPos]
|
||||
if val, ok := body[key]; ok {
|
||||
buf.WriteString(formatValue(val))
|
||||
} else {
|
||||
// Keep the placeholder if key not found
|
||||
buf.WriteString("{{")
|
||||
buf.WriteString(key)
|
||||
buf.WriteString("}}")
|
||||
}
|
||||
i += endPos + 2 // skip "}}"
|
||||
}
|
||||
return buf.String()
|
||||
}
|
||||
|
||||
func formatValue(v any) string {
|
||||
if v == nil {
|
||||
return ""
|
||||
}
|
||||
switch val := v.(type) {
|
||||
case string:
|
||||
return val
|
||||
case []byte:
|
||||
return string(val)
|
||||
case int:
|
||||
return strconv.Itoa(val)
|
||||
case int32:
|
||||
return strconv.FormatInt(int64(val), 10)
|
||||
case int64:
|
||||
return strconv.FormatInt(val, 10)
|
||||
case float64:
|
||||
return strconv.FormatFloat(val, 'f', -1, 64)
|
||||
case bool:
|
||||
return strconv.FormatBool(val)
|
||||
default:
|
||||
// If it's a map, slice, or struct, marshal it to JSON.
|
||||
b, err := json.Marshal(v)
|
||||
if err == nil {
|
||||
return string(b)
|
||||
}
|
||||
return fmt.Sprintf("%v", v)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,75 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package push
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
)
|
||||
|
||||
func TestParseTemplate(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
template string
|
||||
body map[string]any
|
||||
expected string
|
||||
}{
|
||||
{
|
||||
name: "simple replacement",
|
||||
template: "hello {{name}}",
|
||||
body: map[string]any{"name": "world"},
|
||||
expected: "hello world",
|
||||
},
|
||||
{
|
||||
name: "multiple replacements",
|
||||
template: "{{greeting}} {{name}}!",
|
||||
body: map[string]any{"greeting": "Hello", "name": "Alice"},
|
||||
expected: "Hello Alice!",
|
||||
},
|
||||
{
|
||||
name: "missing key preserves placeholder",
|
||||
template: "hello {{name}} and {{other}}",
|
||||
body: map[string]any{"name": "world"},
|
||||
expected: "hello world and {{other}}",
|
||||
},
|
||||
{
|
||||
name: "unbalanced placeholders",
|
||||
template: "hello {{name",
|
||||
body: map[string]any{"name": "world"},
|
||||
expected: "hello {{name",
|
||||
},
|
||||
{
|
||||
name: "nil value",
|
||||
template: "val: {{val}}",
|
||||
body: map[string]any{"val": nil},
|
||||
expected: "val: ",
|
||||
},
|
||||
{
|
||||
name: "basic types",
|
||||
template: "int: {{i}}, float: {{f}}, bool: {{b}}",
|
||||
body: map[string]any{"i": 123, "f": 45.67, "b": true},
|
||||
expected: "int: 123, float: 45.67, bool: true",
|
||||
},
|
||||
{
|
||||
name: "complex type slice",
|
||||
template: "items: {{items}}",
|
||||
body: map[string]any{"items": []string{"a", "b"}},
|
||||
expected: `items: ["a","b"]`,
|
||||
},
|
||||
{
|
||||
name: "complex type map",
|
||||
template: "obj: {{obj}}",
|
||||
body: map[string]any{"obj": map[string]any{"key": "value"}},
|
||||
expected: `obj: {"key":"value"}`,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
result := ParseTemplate(tt.template, tt.body)
|
||||
assert.Equal(t, tt.expected, result)
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,15 @@
|
||||
// Copyright 2025 linux.do
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package trace 提供 OpenTelemetry 链路追踪封装工具
|
||||
package trace
|
||||
|
||||
import "go.opentelemetry.io/otel/propagation"
|
||||
|
||||
func newPropagator() propagation.TextMapPropagator {
|
||||
return propagation.NewCompositeTextMapPropagator(
|
||||
propagation.TraceContext{},
|
||||
propagation.Baggage{},
|
||||
)
|
||||
}
|
||||
@@ -0,0 +1,19 @@
|
||||
// Copyright 2025 linux.do
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package trace
|
||||
|
||||
import (
|
||||
sdktrace "go.opentelemetry.io/otel/sdk/trace"
|
||||
)
|
||||
|
||||
// ParentBasedRatioSampler 创建父级感知的概率采样器
|
||||
// - 如果父 Span 已采样,则子 Span 也采样
|
||||
// - 如果父 Span 未采样,则子 Span 也不采样
|
||||
// - 如果是根 Span,按 samplingRate 概率采样
|
||||
func ParentBasedRatioSampler(samplingRate float64) sdktrace.Sampler {
|
||||
return sdktrace.ParentBased(
|
||||
sdktrace.TraceIDRatioBased(samplingRate),
|
||||
)
|
||||
}
|
||||
@@ -0,0 +1,62 @@
|
||||
// Copyright 2025 linux.do
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package trace
|
||||
|
||||
import (
|
||||
"context"
|
||||
"log"
|
||||
|
||||
"go.opentelemetry.io/otel"
|
||||
"go.opentelemetry.io/otel/trace"
|
||||
)
|
||||
|
||||
// Tracer 全局 OpenTelemetry Tracer 实例
|
||||
var Tracer trace.Tracer
|
||||
var shutdownFuncs []func(context.Context) error
|
||||
|
||||
func init() {
|
||||
// 初始化 Propagator
|
||||
prop := newPropagator()
|
||||
otel.SetTextMapPropagator(prop)
|
||||
|
||||
// 初始化 Tracer 实例为 No-op 默认以避免未初始化前或测试环境崩溃
|
||||
Tracer = otel.GetTracerProvider().Tracer("github.com/Rain-kl/Wavelet")
|
||||
}
|
||||
|
||||
// Config 链路追踪配置
|
||||
type Config struct {
|
||||
AppName string
|
||||
SamplingRate float64
|
||||
TracerName string
|
||||
}
|
||||
|
||||
// Init 初始化 Tracer Provider 并关联全局 Tracer 实例
|
||||
func Init(cfg Config) {
|
||||
tracerProvider, err := newTracerProvider(cfg)
|
||||
if err != nil {
|
||||
log.Fatalf("[Trace] init trace provider failed: %v", err)
|
||||
}
|
||||
shutdownFuncs = append(shutdownFuncs, tracerProvider.Shutdown)
|
||||
otel.SetTracerProvider(tracerProvider)
|
||||
|
||||
// 更新 Tracer
|
||||
tracerName := cfg.TracerName
|
||||
if tracerName == "" {
|
||||
tracerName = "github.com/Rain-kl/Wavelet"
|
||||
}
|
||||
Tracer = tracerProvider.Tracer(tracerName)
|
||||
}
|
||||
|
||||
// Shutdown 关闭所有 Trace Provider
|
||||
func Shutdown(ctx context.Context) {
|
||||
for _, fn := range shutdownFuncs {
|
||||
_ = fn(ctx)
|
||||
}
|
||||
}
|
||||
|
||||
// Start 创建一个新的 Trace Span
|
||||
func Start(ctx context.Context, name string, opts ...trace.SpanStartOption) (context.Context, trace.Span) {
|
||||
return Tracer.Start(ctx, name, opts...)
|
||||
}
|
||||
@@ -0,0 +1,53 @@
|
||||
// Copyright 2025 linux.do
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package trace
|
||||
|
||||
import (
|
||||
"context"
|
||||
"os"
|
||||
|
||||
"go.opentelemetry.io/otel/exporters/otlp/otlptrace/otlptracegrpc"
|
||||
"go.opentelemetry.io/otel/sdk/resource"
|
||||
sdktrace "go.opentelemetry.io/otel/sdk/trace"
|
||||
semconv "go.opentelemetry.io/otel/semconv/v1.26.0"
|
||||
)
|
||||
|
||||
func newTracerProvider(cfg Config) (*sdktrace.TracerProvider, error) {
|
||||
// 获取主机名和容器信息
|
||||
hostname, err := os.Hostname()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// 初始化 Resource
|
||||
r, err := resource.Merge(
|
||||
resource.Default(),
|
||||
resource.NewWithAttributes(
|
||||
semconv.SchemaURL,
|
||||
semconv.ServiceName(cfg.AppName),
|
||||
semconv.HostName(hostname),
|
||||
semconv.K8SNamespaceName(os.Getenv("KUBERNETES_NAMESPACE")),
|
||||
semconv.K8SPodName(os.Getenv("KUBERNETES_POD_NAME")),
|
||||
semconv.K8SPodUID(os.Getenv("KUBERNETES_POD_UID")),
|
||||
),
|
||||
)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// 初始化 Exporter
|
||||
traceExporter, err := otlptracegrpc.New(context.Background())
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// 初始化 Trace
|
||||
tracerProvider := sdktrace.NewTracerProvider(
|
||||
sdktrace.WithBatcher(traceExporter),
|
||||
sdktrace.WithResource(r),
|
||||
sdktrace.WithSampler(ParentBasedRatioSampler(cfg.SamplingRate)),
|
||||
)
|
||||
return tracerProvider, nil
|
||||
}
|
||||
@@ -0,0 +1,160 @@
|
||||
// Copyright 2025 linux.do
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package util provides generic utility functions.
|
||||
package util
|
||||
|
||||
import (
|
||||
"crypto/aes"
|
||||
"crypto/cipher"
|
||||
"crypto/ed25519"
|
||||
"crypto/rand"
|
||||
"encoding/base64"
|
||||
"encoding/hex"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
)
|
||||
|
||||
const (
|
||||
aesKeyLength = 32
|
||||
|
||||
errInvalidSignKey = "invalid sign key: %w"
|
||||
errSignKeyLengthInvalid = "sign key must be 32 bytes (64 hex characters)"
|
||||
errCreateCipherFailed = "failed to create cipher: %w"
|
||||
errCreateGCMFailed = "failed to create GCM: %w"
|
||||
errGenerateNonceFailed = "failed to generate nonce: %w"
|
||||
errDecodeCiphertextFailed = "failed to decode ciphertext: %w"
|
||||
errCiphertextTooShort = "ciphertext too short"
|
||||
errDecryptFailed = "failed to decrypt: %w"
|
||||
)
|
||||
|
||||
// Encrypt 使用 SignKey 加密字符串数据
|
||||
// signKey: 64 字符 hex 编码的密钥(对应 32 字节,用于 AES-256)
|
||||
// plaintext: 要加密的明文字符串
|
||||
// return: base64 编码的密文
|
||||
func Encrypt(signKey string, plaintext string) (string, error) {
|
||||
return encryptBytes(signKey, []byte(plaintext))
|
||||
}
|
||||
|
||||
// Decrypt 使用 SignKey 解密字符串数据
|
||||
// signKey: 64 字符 hex 编码的密钥(对应 32 字节,用于 AES-256)
|
||||
// ciphertext: base64 编码的密文
|
||||
// return: 解密后的明文字符串
|
||||
func Decrypt(signKey string, ciphertext string) (string, error) {
|
||||
plaintext, err := decryptBytes(signKey, ciphertext)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return string(plaintext), nil
|
||||
}
|
||||
|
||||
// encryptBytes 加密函数,处理字节数据
|
||||
func encryptBytes(signKey string, plaintext []byte) (string, error) {
|
||||
// 将 hex 编码的密钥转换为字节
|
||||
key, err := hex.DecodeString(signKey)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf(errInvalidSignKey, err)
|
||||
}
|
||||
if len(key) != aesKeyLength {
|
||||
return "", errors.New(errSignKeyLengthInvalid)
|
||||
}
|
||||
|
||||
// 创建 AES cipher
|
||||
block, err := aes.NewCipher(key)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf(errCreateCipherFailed, err)
|
||||
}
|
||||
|
||||
// 使用 GCM 模式(Galois/Counter Mode)
|
||||
gcm, err := cipher.NewGCM(block)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf(errCreateGCMFailed, err)
|
||||
}
|
||||
|
||||
// 生成随机 nonce
|
||||
nonce := make([]byte, gcm.NonceSize())
|
||||
if _, err := io.ReadFull(rand.Reader, nonce); err != nil {
|
||||
return "", fmt.Errorf(errGenerateNonceFailed, err)
|
||||
}
|
||||
|
||||
// 加密数据
|
||||
ciphertext := gcm.Seal(nonce, nonce, plaintext, nil)
|
||||
|
||||
// 返回 base64 编码的密文
|
||||
return Base64Encode(ciphertext), nil
|
||||
}
|
||||
|
||||
// decryptBytes 解密函数,处理字节数据
|
||||
func decryptBytes(signKey string, ciphertext string) ([]byte, error) {
|
||||
// 将 hex 编码的密钥转换为字节
|
||||
key, err := hex.DecodeString(signKey)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf(errInvalidSignKey, err)
|
||||
}
|
||||
if len(key) != aesKeyLength {
|
||||
return nil, errors.New(errSignKeyLengthInvalid)
|
||||
}
|
||||
|
||||
// 解码 base64 密文
|
||||
data, err := Base64Decode(ciphertext)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf(errDecodeCiphertextFailed, err)
|
||||
}
|
||||
|
||||
// 创建 AES cipher
|
||||
block, err := aes.NewCipher(key)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf(errCreateCipherFailed, err)
|
||||
}
|
||||
|
||||
// 使用 GCM 模式
|
||||
gcm, err := cipher.NewGCM(block)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf(errCreateGCMFailed, err)
|
||||
}
|
||||
|
||||
// 提取 nonce
|
||||
nonceSize := gcm.NonceSize()
|
||||
if len(data) < nonceSize {
|
||||
return nil, errors.New(errCiphertextTooShort)
|
||||
}
|
||||
|
||||
nonce, ciphertextBytes := data[:nonceSize], data[nonceSize:]
|
||||
|
||||
// 解密数据
|
||||
plaintext, err := gcm.Open(nil, nonce, ciphertextBytes, nil)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf(errDecryptFailed, err)
|
||||
}
|
||||
|
||||
return plaintext, nil
|
||||
}
|
||||
|
||||
// Base64Encode Base64编码
|
||||
func Base64Encode(data []byte) string {
|
||||
return base64.StdEncoding.EncodeToString(data)
|
||||
}
|
||||
|
||||
// Base64Decode Base64解码
|
||||
func Base64Decode(encoded string) ([]byte, error) {
|
||||
return base64.StdEncoding.DecodeString(encoded)
|
||||
}
|
||||
|
||||
// Ed25519Verify 验证 Ed25519 签名
|
||||
// publicKey: 32 字节的公钥(已解码的二进制格式)
|
||||
// message: 待验证的原始消息
|
||||
// signature: 64 字节的签名(已解码的二进制格式)
|
||||
// return: 签名是否有效
|
||||
func Ed25519Verify(publicKey, message, signature []byte) bool {
|
||||
if len(publicKey) != ed25519.PublicKeySize {
|
||||
return false
|
||||
}
|
||||
|
||||
if len(signature) != ed25519.SignatureSize {
|
||||
return false
|
||||
}
|
||||
|
||||
return ed25519.Verify(publicKey, message, signature)
|
||||
}
|
||||
@@ -0,0 +1,20 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package util
|
||||
|
||||
import "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
|
||||
}
|
||||
|
||||
// CheckPasswordHash 比较 bcrypt 哈希值与明文密码是否匹配
|
||||
func CheckPasswordHash(hash, password string) bool {
|
||||
return bcrypt.CompareHashAndPassword([]byte(hash), []byte(password)) == nil
|
||||
}
|
||||
@@ -0,0 +1,35 @@
|
||||
// Copyright 2025 linux.do
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package util
|
||||
|
||||
import "strings"
|
||||
|
||||
// emailPartsCount 邮箱地址由 @ 分割为两部分
|
||||
const (
|
||||
emailPartsCount = 2
|
||||
emailLocalMinChars = 2 // 邮箱 local 部分掩码显示的最小字符数
|
||||
)
|
||||
|
||||
// DerefString 安全地解引用字符串指针,nil 返回空字符串
|
||||
func DerefString(s *string) string {
|
||||
if s == nil {
|
||||
return ""
|
||||
}
|
||||
return *s
|
||||
}
|
||||
|
||||
// MaskEmail 安全脱敏用户的邮箱地址(例如 us***@example.com)
|
||||
func MaskEmail(email string) string {
|
||||
parts := strings.Split(email, "@")
|
||||
if len(parts) != emailPartsCount {
|
||||
return email
|
||||
}
|
||||
local := parts[0]
|
||||
domain := parts[1]
|
||||
if len(local) <= emailLocalMinChars {
|
||||
return "**@" + domain
|
||||
}
|
||||
return local[:2] + "***" + local[len(local)-1:] + "@" + domain
|
||||
}
|
||||
@@ -0,0 +1,29 @@
|
||||
// Copyright 2025 linux.do
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package util
|
||||
|
||||
import (
|
||||
"crypto/rand"
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"io"
|
||||
|
||||
"github.com/google/uuid"
|
||||
)
|
||||
|
||||
// uniqueIDBytes 生成唯一 ID 所需的随机字节长度
|
||||
const uniqueIDBytes = 32
|
||||
|
||||
// GenerateUniqueIDSimple 生成 64 位唯一标识符
|
||||
func GenerateUniqueIDSimple() string {
|
||||
randomBytes := make([]byte, uniqueIDBytes)
|
||||
if _, err := io.ReadFull(rand.Reader, randomBytes); err != nil {
|
||||
// 如果随机数生成失败,使用 UUID 作为后备
|
||||
uuidBytes := []byte(uuid.NewString())
|
||||
hash := sha256.Sum256(uuidBytes)
|
||||
copy(randomBytes, hash[:])
|
||||
}
|
||||
return hex.EncodeToString(randomBytes)
|
||||
}
|
||||
Reference in New Issue
Block a user