refactor(backend): extract to pkg/cap

This commit is contained in:
ryan
2026-06-15 16:19:18 +08:00
parent 84ae4ec27e
commit b3ed94342c
63 changed files with 770 additions and 706 deletions
+1 -1
View File
@@ -18,9 +18,9 @@ import (
"github.com/Rain-kl/Wavelet/internal/apps/admin"
"github.com/Rain-kl/Wavelet/internal/config"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/logger"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/util"
"github.com/Rain-kl/Wavelet/pkg/logger"
"github.com/gin-gonic/gin"
)
+2 -2
View File
@@ -7,10 +7,10 @@ package admin
import (
"net/http"
"github.com/Rain-kl/Wavelet/internal/logger"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/otel_trace"
"github.com/Rain-kl/Wavelet/internal/util"
"github.com/Rain-kl/Wavelet/pkg/logger"
otel_trace "github.com/Rain-kl/Wavelet/pkg/trace"
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
"github.com/gin-gonic/gin"
+1 -1
View File
@@ -14,9 +14,9 @@ import (
"sync"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/logger"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/task"
"github.com/Rain-kl/Wavelet/pkg/logger"
pkgpush "github.com/Rain-kl/Wavelet/pkg/push"
"gorm.io/gorm"
)
+1 -1
View File
@@ -10,9 +10,9 @@ import (
"time"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/logger"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/task"
"github.com/Rain-kl/Wavelet/pkg/logger"
)
func init() {
+2 -2
View File
@@ -13,11 +13,11 @@ import (
"time"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/logger"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/storage"
"github.com/Rain-kl/Wavelet/internal/util"
mail "github.com/Rain-kl/Wavelet/internal/util/mail"
"github.com/Rain-kl/Wavelet/pkg/logger"
mail "github.com/Rain-kl/Wavelet/pkg/mail"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
)
+1 -1
View File
@@ -12,12 +12,12 @@ import (
"time"
"github.com/Rain-kl/Wavelet/internal/apps/admin"
"github.com/Rain-kl/Wavelet/internal/logger"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/task"
taskhandlers "github.com/Rain-kl/Wavelet/internal/task/handlers"
"github.com/Rain-kl/Wavelet/internal/task/scheduler"
"github.com/Rain-kl/Wavelet/internal/util"
"github.com/Rain-kl/Wavelet/pkg/logger"
"github.com/gin-gonic/gin"
"github.com/robfig/cron/v3"
)
+1 -1
View File
@@ -22,8 +22,8 @@ import (
"time"
"github.com/Rain-kl/Wavelet/internal/buildinfo"
"github.com/Rain-kl/Wavelet/internal/logger"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/pkg/logger"
"golang.org/x/mod/semver"
)
+1 -1
View File
@@ -12,7 +12,7 @@ import (
"path/filepath"
"syscall"
"github.com/Rain-kl/Wavelet/internal/logger"
"github.com/Rain-kl/Wavelet/pkg/logger"
)
const installedBinaryMode = 0o755
+1 -1
View File
@@ -8,8 +8,8 @@ import (
"net/http"
"time"
"github.com/Rain-kl/Wavelet/internal/logger"
"github.com/Rain-kl/Wavelet/internal/util"
"github.com/Rain-kl/Wavelet/pkg/logger"
"github.com/gin-gonic/gin"
)
+1 -1
View File
@@ -7,7 +7,7 @@ import (
"net/http"
"github.com/Rain-kl/Wavelet/internal/util"
caputil "github.com/Rain-kl/Wavelet/internal/util/cap"
caputil "github.com/Rain-kl/Wavelet/internal/service/cap"
"github.com/gin-gonic/gin"
)
+6 -6
View File
@@ -6,7 +6,7 @@ package cap
import (
"net/http"
"github.com/Rain-kl/Wavelet/internal/util/cap"
capService "github.com/Rain-kl/Wavelet/internal/service/cap"
"github.com/gin-gonic/gin"
)
@@ -38,10 +38,10 @@ func Challenge(c *gin.Context) {
req.Scope = "login"
}
mgr := cap.GetDefaultManager()
mgr := capService.GetDefaultManager()
resp, err := mgr.Generate(c.Request.Context(), req.Scope)
if err != nil {
c.JSON(http.StatusInternalServerError, cap.RedeemResponse{
c.JSON(http.StatusInternalServerError, capService.RedeemResponse{
Success: false,
Error: err.Error(),
})
@@ -65,7 +65,7 @@ func Challenge(c *gin.Context) {
func Redeem(c *gin.Context) {
var req redeemRequest
if err := c.ShouldBindJSON(&req); err != nil {
c.JSON(http.StatusBadRequest, cap.RedeemResponse{
c.JSON(http.StatusBadRequest, capService.RedeemResponse{
Success: false,
Error: "无效的参数",
})
@@ -76,10 +76,10 @@ func Redeem(c *gin.Context) {
req.Scope = "login"
}
mgr := cap.GetDefaultManager()
mgr := capService.GetDefaultManager()
resp, err := mgr.Redeem(c.Request.Context(), req.Token, req.Solutions, req.Scope)
if err != nil {
c.JSON(http.StatusInternalServerError, cap.RedeemResponse{
c.JSON(http.StatusInternalServerError, capService.RedeemResponse{
Success: false,
Error: err.Error(),
})
+4 -3
View File
@@ -15,7 +15,8 @@ import (
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/testhelper"
"github.com/Rain-kl/Wavelet/internal/util"
capUtil "github.com/Rain-kl/Wavelet/internal/util/cap"
capUtil "github.com/Rain-kl/Wavelet/internal/service/cap"
pkgcap "github.com/Rain-kl/Wavelet/pkg/cap"
"github.com/gin-gonic/gin"
)
@@ -53,7 +54,7 @@ func TestCapEndpointsAndMiddleware(t *testing.T) {
t.Fatalf("expected 200 OK, got %d. Body: %s", w.Code, w.Body.String())
}
var challengeResp capUtil.ChallengeResponse
var challengeResp pkgcap.ChallengeResponse
if err := json.Unmarshal(w.Body.Bytes(), &challengeResp); err != nil {
t.Fatalf("failed to unmarshal challenge response: %v", err)
}
@@ -89,7 +90,7 @@ func TestCapEndpointsAndMiddleware(t *testing.T) {
}
// 5. Solve the challenge
solutions := capUtil.Solve(challengeResp.Token, challengeResp.Challenge.C, challengeResp.Challenge.S, challengeResp.Challenge.D)
solutions := pkgcap.Solve(challengeResp.Token, challengeResp.Challenge.C, challengeResp.Challenge.S, challengeResp.Challenge.D)
// 6. Redeem solutions
redeemReqPayload := redeemRequest{
+1 -1
View File
@@ -9,8 +9,8 @@ import (
"context"
"encoding/json"
"github.com/Rain-kl/Wavelet/internal/logger"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/pkg/logger"
"github.com/gin-gonic/gin"
)
+1 -1
View File
@@ -12,8 +12,8 @@ import (
"github.com/Rain-kl/Wavelet/internal/common"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/otel_trace"
"github.com/Rain-kl/Wavelet/internal/util"
otel_trace "github.com/Rain-kl/Wavelet/pkg/trace"
"github.com/gin-contrib/sessions"
"github.com/gin-gonic/gin"
)
+1 -1
View File
@@ -7,9 +7,9 @@ package oauth
import (
"net/http"
"github.com/Rain-kl/Wavelet/internal/logger"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/util"
"github.com/Rain-kl/Wavelet/pkg/logger"
"github.com/gin-contrib/sessions"
"github.com/gin-gonic/gin"
)
+1 -1
View File
@@ -18,9 +18,9 @@ import (
"github.com/Rain-kl/Wavelet/internal/common"
"github.com/Rain-kl/Wavelet/internal/config"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/logger"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/util"
"github.com/Rain-kl/Wavelet/pkg/logger"
"github.com/coreos/go-oidc/v3/oidc"
"github.com/gin-contrib/sessions"
"github.com/gin-gonic/gin"
+1 -1
View File
@@ -9,7 +9,7 @@ import (
"github.com/Rain-kl/Wavelet/internal/config"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/logger"
"github.com/Rain-kl/Wavelet/pkg/logger"
)
var logChan chan *UserAccessLog
+1 -1
View File
@@ -19,9 +19,9 @@ import (
"github.com/Rain-kl/Wavelet/internal/common"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/diskcache"
"github.com/Rain-kl/Wavelet/internal/logger"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/util"
"github.com/Rain-kl/Wavelet/pkg/logger"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
)
+1 -1
View File
@@ -26,10 +26,10 @@ import (
"github.com/Rain-kl/Wavelet/internal/common"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/db/idgen"
"github.com/Rain-kl/Wavelet/internal/logger"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/storage"
"github.com/Rain-kl/Wavelet/internal/util"
"github.com/Rain-kl/Wavelet/pkg/logger"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
)
+1 -1
View File
@@ -7,9 +7,9 @@ import (
"context"
"fmt"
"github.com/Rain-kl/Wavelet/internal/logger"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/storage"
"github.com/Rain-kl/Wavelet/pkg/logger"
)
// StorageReadOnly checks if the storage system is in read-only maintenance mode.
+2 -1
View File
@@ -20,6 +20,7 @@ import (
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/task"
"github.com/Rain-kl/Wavelet/internal/util"
pkgu "github.com/Rain-kl/Wavelet/pkg/util"
"github.com/gin-gonic/gin"
)
@@ -171,7 +172,7 @@ func handleLoginEmailVerification(ctx context.Context, c *gin.Context, req *logi
}
}
maskedEmail := util.MaskEmail(user.Email)
maskedEmail := pkgu.MaskEmail(user.Email)
c.JSON(http.StatusOK, util.Err(errNeedEmailCodePrefix+maskedEmail))
return errors.New("handled")
}
+1 -1
View File
@@ -15,9 +15,9 @@ import (
"github.com/Rain-kl/Wavelet/internal/config"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/db/idgen"
"github.com/Rain-kl/Wavelet/internal/logger"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/util"
"github.com/Rain-kl/Wavelet/pkg/logger"
"github.com/gin-contrib/sessions"
"github.com/gin-gonic/gin"
)
+1 -1
View File
@@ -13,7 +13,7 @@ import (
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/task"
"github.com/Rain-kl/Wavelet/internal/util/mail"
"github.com/Rain-kl/Wavelet/pkg/mail"
)
// 异步任务名称与管理类型定义
+19
View File
@@ -8,12 +8,31 @@ import (
"log"
"github.com/Rain-kl/Wavelet/internal/buildinfo"
"github.com/Rain-kl/Wavelet/internal/config"
"github.com/Rain-kl/Wavelet/internal/db/migrator"
"github.com/Rain-kl/Wavelet/pkg/logger"
"github.com/Rain-kl/Wavelet/pkg/trace"
"github.com/spf13/cobra"
)
var rootCmd = &cobra.Command{
Use: "wavelet",
PersistentPreRun: func(_ *cobra.Command, _ []string) {
logger.Init(logger.Config{
Level: config.Config.Log.Level,
Format: config.Config.Log.Format,
Output: config.Config.Log.Output,
FilePath: config.Config.Log.FilePath,
MaxSize: config.Config.Log.MaxSize,
MaxAge: config.Config.Log.MaxAge,
MaxBackups: config.Config.Log.MaxBackups,
Compress: config.Config.Log.Compress,
})
trace.Init(trace.Config{
AppName: config.Config.App.AppName,
SamplingRate: config.Config.Otel.SamplingRate,
})
},
PreRun: func(_ *cobra.Command, _ []string) {
migrator.Migrate()
},
+1 -1
View File
@@ -11,7 +11,7 @@ import (
"strings"
"time"
"github.com/Rain-kl/Wavelet/internal/logger"
"github.com/Rain-kl/Wavelet/pkg/logger"
"gorm.io/gorm"
gormLogger "gorm.io/gorm/logger"
)
+20 -360
View File
@@ -1,261 +1,68 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package diskcache implements a platform-level disk-backed cache with size limit, TTL, and LRU eviction.
// Package diskcache wraps the generic pkg/diskcache to provide database configuration integration.
package diskcache
import (
"container/list"
"context"
"encoding/binary"
"errors"
"fmt"
"os"
"path/filepath"
"sort"
"strconv"
"strings"
"sync"
"time"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/peterbourgon/diskv/v3"
pkgcache "github.com/Rain-kl/Wavelet/pkg/diskcache"
)
// ErrCacheMiss represents a cache miss.
var ErrCacheMiss = errors.New("cache miss")
// Status represents the runtime cache statistics.
type Status = pkgcache.Status
// Constants for disk cache configuration and sizing
const (
defaultCacheDir = "uploads/diskcache"
headerSize = 8 // 8 bytes metadata prefix for expiration UnixNano timestamp
defaultMaxSizeMB = 100
defaultTTLMinutes = 60
defaultCleanupInterval = 10
// DefaultExpiration applies the cache-wide default TTL.
DefaultExpiration time.Duration = 0
DefaultExpiration = pkgcache.DefaultExpiration
// NoExpiration stores the item without a TTL. Size limits and LRU eviction still apply.
NoExpiration time.Duration = -1
NoExpiration = pkgcache.NoExpiration
)
// ErrCacheMiss represents a cache miss.
var ErrCacheMiss = pkgcache.ErrCacheMiss
// DiskCache is a wrapper around the generic pkg/diskcache that integrates with the DB for configs.
type DiskCache struct {
*pkgcache.DiskCache
}
var (
globalCache *DiskCache
globalCacheOnce sync.Once
)
// 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"`
}
// DiskCache implements the disk-backed cache with size limits, TTL, and LRU eviction.
type DiskCache 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
}
// GetGlobalCache returns the global singleton DiskCache instance.
func GetGlobalCache() *DiskCache {
globalCacheOnce.Do(func() {
globalCache = New(defaultCacheDir)
pureCache := pkgcache.New(defaultCacheDir)
globalCache = &DiskCache{pureCache}
// Load initial configs from database
globalCache.ReloadConfig(context.Background())
// Start background routine to clean expired items every 10 minutes
go globalCache.startCleanupWorker(defaultCleanupInterval * time.Minute)
go globalCache.StartCleanupWorker(defaultCleanupInterval * time.Minute)
})
return globalCache
}
// New creates a new DiskCache instance.
// New creates a new DiskCache wrapper.
func New(basePath string) *DiskCache {
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 := &DiskCache{
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 *DiskCache) 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 *DiskCache) Get(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)
// Check expiration
if !item.expiredAt.IsZero() && time.Now().After(item.expiredAt) {
// Lazily delete expired item
_ = c.deleteUnlocked(key)
return nil, ErrCacheMiss
}
// Read from diskv
data, err := c.d.Read(key)
if err != nil {
// Key exists in memory but not on disk, sync state
_ = c.deleteUnlocked(key)
return nil, ErrCacheMiss
}
if len(data) < headerSize {
_ = c.deleteUnlocked(key)
return nil, ErrCacheMiss
}
// Update LRU access order
c.evictList.MoveToFront(elem)
// Slice off the metadata header
return data[headerSize:], nil
}
// Delete removes a key-value pair from the cache.
func (c *DiskCache) Delete(key string) error {
c.mu.Lock()
defer c.mu.Unlock()
return c.deleteUnlocked(key)
}
func (c *DiskCache) 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 *DiskCache) 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 *DiskCache) 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,
}
return &DiskCache{pkgcache.New(basePath)}
}
// ReloadConfig reloads policies from database configs dynamically.
func (c *DiskCache) ReloadConfig(ctx context.Context) {
c.mu.Lock()
defer c.mu.Unlock()
// Ensure DB is initialized before querying
if db.DB(ctx) == nil {
return
@@ -269,7 +76,6 @@ func (c *DiskCache) ReloadConfig(ctx context.Context) {
maxSizeMB = val
}
}
c.maxSize = maxSizeMB * 1024 * 1024
// 2. Default TTL
var scTTL model.SystemConfig
@@ -279,7 +85,6 @@ func (c *DiskCache) ReloadConfig(ctx context.Context) {
ttlMinutes = val
}
}
c.defaultTTL = time.Duration(ttlMinutes) * time.Minute
// 3. LRU Enabled
var scLRU model.SystemConfig
@@ -289,151 +94,6 @@ func (c *DiskCache) ReloadConfig(ctx context.Context) {
lruEnabled = val
}
}
c.lruEnabled = lruEnabled
// Apply eviction immediately under new configs
c.evict()
}
// evict evicts oldest items if current size exceeds maxSize and LRU is enabled.
func (c *DiskCache) evict() {
if !c.lruEnabled {
return
}
for c.currentSize > c.maxSize && c.evictList.Len() > 0 {
elem := c.evictList.Back()
if elem == nil {
break
}
item := elem.Value.(*cacheItem)
key := item.key
// Remove from memory
c.currentSize -= item.size
c.evictList.Remove(elem)
delete(c.items, key)
// Delete from disk
_ = c.d.Erase(key)
}
}
// loadTracker walks the directory to rebuild the LRU and size tracking structures.
func (c *DiskCache) loadTracker() error {
c.mu.Lock()
defer c.mu.Unlock()
c.currentSize = 0
c.items = make(map[string]*list.Element)
c.evictList = list.New()
files, err := os.ReadDir(c.basePath)
if err != nil {
if os.IsNotExist(err) {
return nil
}
return err
}
type tempItem struct {
key string
size int64
expiredAt time.Time
mtime time.Time
}
var loadedItems []tempItem
for _, file := range files {
if file.IsDir() {
continue
}
name := file.Name()
// Skip temporary files
if strings.HasPrefix(name, ".") || strings.Contains(name, "temp") {
continue
}
info, err := file.Info()
if err != nil {
continue
}
filePath := filepath.Join(c.basePath, name)
// #nosec G304
f, err := os.Open(filePath)
if err != nil {
continue
}
var expiredAt time.Time
var expNano int64
err = binary.Read(f, binary.BigEndian, &expNano)
_ = f.Close()
if err != nil {
// Corrupt metadata header: delete file
_ = os.Remove(filePath)
continue
}
if expNano > 0 {
expiredAt = time.Unix(0, expNano)
// Expired: delete file
if time.Now().After(expiredAt) {
_ = os.Remove(filePath)
continue
}
}
loadedItems = append(loadedItems, tempItem{
key: name,
size: info.Size(),
expiredAt: expiredAt,
mtime: info.ModTime(),
})
}
// Sort loaded items by modification time ascending (oldest first)
sort.Slice(loadedItems, func(i, j int) bool {
return loadedItems[i].mtime.Before(loadedItems[j].mtime)
})
// 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 *DiskCache) 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 *DiskCache) 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)
}
}
c.UpdatePolicy(maxSizeMB, ttlMinutes, lruEnabled)
}
-203
View File
@@ -4,218 +4,15 @@
package diskcache
import (
"bytes"
"context"
"os"
"testing"
"time"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/testhelper"
)
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)
}
}
func TestDiskCacheReloadConfig(t *testing.T) {
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
-66
View File
@@ -1,66 +0,0 @@
// 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(),
}
}
-37
View File
@@ -1,37 +0,0 @@
// 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")
}
}
-9
View File
@@ -1,9 +0,0 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package logger 提供结构化日志封装
package logger
const (
errCreateLogFileDirFailed = "[Logger] create log file dir err: %w"
)
-75
View File
@@ -1,75 +0,0 @@
// 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"
)
var logger *otelzap.Logger
// ringBufferCapacity 环形缓冲区容量
const ringBufferCapacity = 5000
// GlobalRingBuffer 全局日志环形缓冲区,供 Admin 日志查询和 WebSocket 推送使用
var GlobalRingBuffer *LogRingBuffer
func init() {
logWriter, err := GetLogWriter()
if err != nil {
log.Fatalf("[Logger] get log writer err: %v\n", err)
}
// 初始化 ring buffer(保留最近 5000 行日志)
GlobalRingBuffer = NewLogRingBuffer(ringBufferCapacity)
// 使用 multi writer 同时写入原始输出和 ring buffer
multiWriter := zapcore.NewMultiWriteSyncer(
logWriter,
zapcore.AddSync(GlobalRingBuffer),
)
zapLogger := zap.New(
zapcore.NewCore(getEncoder(), multiWriter, getLogLevel()),
zap.AddCaller(),
zap.AddCallerSkip(1),
)
logger = otelzap.New(
zapLogger,
otelzap.WithMinLevel(zapLogger.Level()),
)
fmt.Printf("[Logger] %s\n", logger.Level())
}
// DebugF 输出 Debug 级别日志
func DebugF(ctx context.Context, format string, args ...interface{}) {
msg := fmt.Sprintf(format, args...)
logger.Ctx(ctx).Debug(msg, getTraceIDFields(ctx)...)
}
// InfoF 输出 Info 级别日志
func InfoF(ctx context.Context, format string, args ...interface{}) {
msg := fmt.Sprintf(format, args...)
logger.Ctx(ctx).Info(msg, getTraceIDFields(ctx)...)
}
// WarnF 输出 Warn 级别日志
func WarnF(ctx context.Context, format string, args ...interface{}) {
msg := fmt.Sprintf(format, args...)
logger.Ctx(ctx).Warn(msg, getTraceIDFields(ctx)...)
}
// ErrorF 输出 Error 级别日志
func ErrorF(ctx context.Context, format string, args ...interface{}) {
msg := fmt.Sprintf(format, args...)
logger.Ctx(ctx).Error(msg, getTraceIDFields(ctx)...)
}
-175
View File
@@ -1,175 +0,0 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package logger
import (
"io"
"sync"
)
// LogEntry 日志条目,对应 ring buffer 中的一行日志
type LogEntry struct {
Index int `json:"index"` // 全局递增序号
Data string `json:"data"` // 一行日志原文(含换行符)
}
// LogRingBuffer 固定容量的环形缓冲区,存储最近的日志行
// 支持:追加日志、按 cursor 分页查询、订阅实时推送
type LogRingBuffer struct {
mu sync.RWMutex
entries []LogEntry
cap int
head int // 下一条写入的位置
count int // 当前条目数
seq int // 全局递增序号
subscribers map[chan LogEntry]struct{}
subMu sync.RWMutex
}
// NewLogRingBuffer 创建指定容量的日志环形缓冲区
func NewLogRingBuffer(capacity int) *LogRingBuffer {
return &LogRingBuffer{
entries: make([]LogEntry, capacity),
cap: capacity,
subscribers: make(map[chan LogEntry]struct{}),
}
}
// Write 实现 io.Writer 接口,供 zapcore.WriteSyncer 调用
// 按 '\n' 分割为独立行写入 ring buffer
func (r *LogRingBuffer) Write(p []byte) (int, error) {
if len(p) == 0 {
return 0, nil
}
data := string(p)
start := 0
for i := 0; i < len(data); i++ {
if data[i] == '\n' {
line := data[start:i]
start = i + 1
if len(line) > 0 {
r.appendLine(line)
}
}
}
// 处理最后一行(没有换行符结尾的情况)
if start < len(data) && len(data[start:]) > 0 {
r.appendLine(data[start:])
}
return len(p), nil
}
// Sync 实现 zapcore.WriteSyncer 接口
func (r *LogRingBuffer) Sync() error {
return nil
}
// appendLine 追加一行日志到 ring buffer 并通知订阅者
func (r *LogRingBuffer) appendLine(line string) {
r.mu.Lock()
entry := LogEntry{
Index: r.seq,
Data: line,
}
r.entries[r.head] = entry
r.head = (r.head + 1) % r.cap
if r.count < r.cap {
r.count++
}
r.seq++
r.mu.Unlock()
// 异步通知订阅者
r.subMu.RLock()
for ch := range r.subscribers {
select {
case ch <- entry:
default:
// 订阅者消费太慢,丢弃(避免阻塞日志写入)
}
}
r.subMu.RUnlock()
}
// Query 查询历史日志
// cursor=0 表示查询最新日志,cursor>0 表示查询 index < cursor 的更早日志
// limit 为返回条数上限
// 返回日志条目(按 index 升序)和是否有更早的日志
func (r *LogRingBuffer) Query(cursor int, limit int) ([]LogEntry, bool) {
r.mu.RLock()
defer r.mu.RUnlock()
if r.count == 0 {
return nil, false
}
// 计算 ring buffer 中有效条目的范围
// oldest index in ring: head - count (wrapping)
oldestPos := (r.head - r.count + r.cap) % r.cap
// 将 ring buffer 中的有效条目按顺序收集
ordered := make([]LogEntry, 0, r.count)
for i := 0; i < r.count; i++ {
pos := (oldestPos + i) % r.cap
ordered = append(ordered, r.entries[pos])
}
if cursor == 0 {
// 查询最新日志:返回最后 limit 条
if len(ordered) <= limit {
return ordered, false
}
return ordered[len(ordered)-limit:], true
}
// 查询 index < cursor 的更早日志
// 找到 index < cursor 的条目
var cut int
for cut = len(ordered); cut > 0; cut-- {
if ordered[cut-1].Index < cursor {
break
}
}
if cut == 0 {
return nil, false
}
// 返回 cut 之前的最后 limit 条
start := cut - limit
if start < 0 {
start = 0
}
hasMore := start > 0
return ordered[start:cut], hasMore
}
// subscribeChanSize 订阅者 channel 缓冲区大小
const subscribeChanSize = 64
// Subscribe 订阅实时日志推送
// 返回一个 channel,调用者应 defer Unsubscribe
func (r *LogRingBuffer) Subscribe() chan LogEntry {
ch := make(chan LogEntry, subscribeChanSize)
r.subMu.Lock()
r.subscribers[ch] = struct{}{}
r.subMu.Unlock()
return ch
}
// Unsubscribe 取消订阅
func (r *LogRingBuffer) Unsubscribe(ch chan LogEntry) {
r.subMu.Lock()
delete(r.subscribers, ch)
r.subMu.Unlock()
close(ch)
}
// 确保 LogRingBuffer 实现 io.Writer 接口
var _ io.Writer = (*LogRingBuffer)(nil)
-192
View File
@@ -1,192 +0,0 @@
// 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)
}
-115
View File
@@ -1,115 +0,0 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package logger
import (
"context"
"fmt"
"log"
"os"
"path/filepath"
"sync"
"github.com/Rain-kl/Wavelet/internal/config"
"go.opentelemetry.io/otel/trace"
"go.uber.org/zap"
"go.uber.org/zap/zapcore"
"gopkg.in/natefinch/lumberjack.v2"
)
var (
logWriter zapcore.WriteSyncer
initLogWriterOnce sync.Once
initLogWriterErr error
)
// GetLogWriter 获取日志输出写入器
func GetLogWriter() (zapcore.WriteSyncer, error) {
initLogWriterOnce.Do(func() {
logWriter, initLogWriterErr = initWriter()
})
return logWriter, initLogWriterErr
}
// logDirPerm 日志目录权限
const logDirPerm = 0750
func initWriter() (zapcore.WriteSyncer, error) {
logConfig := config.Config.Log
if logConfig.Output == "file" {
// 初始化日志目录
logPath := logConfig.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: logConfig.MaxSize,
MaxBackups: logConfig.MaxBackups,
MaxAge: logConfig.MaxAge,
Compress: logConfig.Compress,
}
return zapcore.AddSync(logOutput), nil
}
return zapcore.AddSync(os.Stdout), nil
}
// getEncoder 获取日志编码器
func getEncoder() 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 config.Config.Log.Format == "json" {
return zapcore.NewJSONEncoder(encoderConfig)
}
return zapcore.NewConsoleEncoder(encoderConfig)
}
// getLogLevel 获取日志级别
func getLogLevel() zapcore.Level {
level := config.Config.Log.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.Fatalf("[Logger] invalid log level: %s\n", level)
return zapcore.InfoLevel
}
}
func getTraceIDFields(ctx context.Context) []zap.Field {
span := trace.SpanFromContext(ctx)
spanContext := span.SpanContext()
return []zap.Field{
zap.String("traceID", spanContext.TraceID().String()),
zap.String("spanID", spanContext.SpanID().String()),
}
}
+1 -1
View File
@@ -12,7 +12,7 @@ import (
"time"
"github.com/Rain-kl/Wavelet/internal/common"
"github.com/Rain-kl/Wavelet/internal/util"
"github.com/Rain-kl/Wavelet/pkg/util"
"gorm.io/gorm"
)
-15
View File
@@ -1,15 +0,0 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package otel_trace 提供 OpenTelemetry 链路追踪封装工具
package otel_trace
import "go.opentelemetry.io/otel/propagation"
func newPropagator() propagation.TextMapPropagator {
return propagation.NewCompositeTextMapPropagator(
propagation.TraceContext{},
propagation.Baggage{},
)
}
-19
View File
@@ -1,19 +0,0 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package otel_trace
import (
sdktrace "go.opentelemetry.io/otel/sdk/trace"
)
// ParentBasedErrorAwareSampler 创建父级感知的概率采样器
// - 如果父 Span 已采样,则子 Span 也采样
// - 如果父 Span 未采样,则子 Span 也不采样
// - 如果是根 Span,按 samplingRate 概率采样
func ParentBasedErrorAwareSampler(samplingRate float64) sdktrace.Sampler {
return sdktrace.ParentBased(
sdktrace.TraceIDRatioBased(samplingRate),
)
}
-47
View File
@@ -1,47 +0,0 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package otel_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)
// 初始化 Trace Provider
tracerProvider, err := newTracerProvider()
if err != nil {
log.Fatalf("[Trace] init trace provider failed: %v", err)
}
shutdownFuncs = append(shutdownFuncs, tracerProvider.Shutdown)
otel.SetTracerProvider(tracerProvider)
// 初始化 Tracer
Tracer = tracerProvider.Tracer("github.com/Rain-kl/Wavelet")
}
// Shutdown 关闭所有 Trace Provider
func Shutdown(ctx context.Context) {
for _, fn := range shutdownFuncs {
_ = fn(ctx)
}
shutdownFuncs = nil
}
// Start 创建一个新的 Trace Span
func Start(ctx context.Context, name string, opts ...trace.SpanStartOption) (context.Context, trace.Span) {
return Tracer.Start(ctx, name, opts...)
}
-54
View File
@@ -1,54 +0,0 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package otel_trace
import (
"context"
"os"
"github.com/Rain-kl/Wavelet/internal/config"
"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() (*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(config.Config.App.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(ParentBasedErrorAwareSampler(config.Config.Otel.SamplingRate)),
)
return tracerProvider, nil
}
+2 -2
View File
@@ -12,9 +12,9 @@ import (
"time"
"github.com/Rain-kl/Wavelet/internal/config"
"github.com/Rain-kl/Wavelet/internal/logger"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/otel_trace"
"github.com/Rain-kl/Wavelet/pkg/logger"
otel_trace "github.com/Rain-kl/Wavelet/pkg/trace"
"github.com/gin-gonic/gin"
"go.opentelemetry.io/otel/codes"
"go.opentelemetry.io/otel/trace"
+2 -2
View File
@@ -34,7 +34,7 @@ import (
"github.com/Rain-kl/Wavelet/internal/apps/user"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/util"
capUtil "github.com/Rain-kl/Wavelet/internal/util/cap"
capUtil "github.com/Rain-kl/Wavelet/internal/service/cap"
// Swagger 文档生成
_ "github.com/Rain-kl/Wavelet/docs"
@@ -42,7 +42,7 @@ import (
"github.com/Rain-kl/Wavelet/internal/apps/admin/system_config"
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
"github.com/Rain-kl/Wavelet/internal/config"
"github.com/Rain-kl/Wavelet/internal/otel_trace"
otel_trace "github.com/Rain-kl/Wavelet/pkg/trace"
"github.com/gin-contrib/sessions"
"github.com/gin-contrib/sessions/redis"
"github.com/gin-gonic/gin"
@@ -1,6 +1,7 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package cap provides CAPTCHA and proof-of-work (PoW) verification services.
package cap
import (
@@ -15,6 +16,7 @@ import (
"github.com/Rain-kl/Wavelet/internal/config"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
pkgcap "github.com/Rain-kl/Wavelet/pkg/cap"
)
const (
@@ -42,11 +44,11 @@ type Config struct {
// Manager orchestrates challenge generation and solution validation
type Manager struct {
conf Config
store Store
store pkgcap.Store
}
// NewManager creates a new CAPTCHA Manager
func NewManager(conf Config, store Store) *Manager {
func NewManager(conf Config, store pkgcap.Store) *Manager {
if conf.ChallengeCount <= 0 {
conf.ChallengeCount = managerDefaultChallengeCount
}
@@ -69,19 +71,27 @@ func NewManager(conf Config, store Store) *Manager {
}
// Generate creates a challenge response
func (m *Manager) Generate(ctx context.Context, scope string) (*ChallengeResponse, error) {
c := ChallengeConfig{
func (m *Manager) Generate(ctx context.Context, scope string) (*pkgcap.ChallengeResponse, error) {
c := pkgcap.ChallengeConfig{
Count: m.getChallengeCount(ctx),
Size: m.getChallengeSize(ctx),
Difficulty: m.getChallengeDifficulty(ctx),
Expires: m.getChallengeTTL(ctx),
}
return GenerateChallenge(m.conf.Secret, c, scope)
return pkgcap.GenerateChallenge(m.conf.Secret, c, scope)
}
// RedeemResponse is returned to the client on redeem
type RedeemResponse struct {
Success bool `json:"success"`
Token string `json:"token,omitempty"`
Expires int64 `json:"expires,omitempty"`
Error string `json:"error,omitempty"`
}
// Redeem verifies PoW solutions and returns a one-time redeem token
func (m *Manager) Redeem(ctx context.Context, token string, solutions []int, scope string) (*RedeemResponse, error) {
sigHex := jwtSigHex(token)
sigHex := pkgcap.JwtSigHex(token)
if sigHex == "" {
return &RedeemResponse{Success: false, Error: "invalid_token"}, nil
}
@@ -89,10 +99,7 @@ func (m *Manager) Redeem(ctx context.Context, token string, solutions []int, sco
nonceKey := "cap:nonce:" + sigHex
// Atomically claim the nonce slot BEFORE verifying solutions.
// SetNX returns true only when the key did not previously exist, so two
// concurrent requests carrying the same JWT can never both succeed here.
// TTL is set to the challenge's remaining lifetime so the slot auto-expires.
payload, err := VerifyChallengeSolutions(token, solutions, m.conf.Secret, scope)
payload, err := pkgcap.VerifyChallengeSolutions(token, solutions, m.conf.Secret, scope)
if err != nil {
return &RedeemResponse{Success: false, Error: err.Error()}, nil //nolint:nilerr // expected behavior: validation error is returned as response, not system error
}
@@ -115,8 +122,8 @@ func (m *Manager) Redeem(ctx context.Context, token string, solutions []int, sco
}
// Generate a redeem token formatted as "id:verToken"
id := randomHex(redeemTokenIDLength)
verToken := randomHex(redeemVerTokenLength)
id := pkgcap.RandomHex(redeemTokenIDLength)
verToken := pkgcap.RandomHex(redeemVerTokenLength)
verHashBytes := sha256.Sum256([]byte(verToken))
verHashHex := hex.EncodeToString(verHashBytes[:])
@@ -139,8 +146,6 @@ func (m *Manager) Redeem(ctx context.Context, token string, solutions []int, sco
}
// VerifyToken validates and consumes the redeem token (single-use).
// GetAndDelete is used so that retrieval and removal happen atomically:
// two concurrent requests carrying the same token can never both see a value.
func (m *Manager) VerifyToken(ctx context.Context, token string, expectedScope string) (bool, error) {
if token == "" {
return false, nil
@@ -157,8 +162,7 @@ func (m *Manager) VerifyToken(ctx context.Context, token string, expectedScope s
tokenKey := "cap:token:" + id + ":" + verHashHex
// Atomically retrieve-and-delete: the first caller gets the value, any
// subsequent caller (even concurrent) receives (false, nil) immediately.
// Atomically retrieve-and-delete
val, exists, err := sGetAndDelete(ctx, m.store, tokenKey)
if err != nil {
return false, err
@@ -190,7 +194,7 @@ func (m *Manager) VerifyToken(ctx context.Context, token string, expectedScope s
}
// sGetAndDelete safely calls store.GetAndDelete, treating a nil store as a miss.
func sGetAndDelete(ctx context.Context, store Store, key string) (string, bool, error) {
func sGetAndDelete(ctx context.Context, store pkgcap.Store, key string) (string, bool, error) {
if store == nil {
return "", false, nil
}
@@ -258,11 +262,11 @@ func GetDefaultManager() *Manager {
challengeTTL := defaultChallengeTTL
tokenTTL := defaultTokenTTL
var store Store
var store pkgcap.Store
if config.Config != nil && config.Config.Redis.Enabled && db.Redis != nil {
store = NewRedisStore(db.Redis)
store = pkgcap.NewRedisStore(db.Redis)
} else {
store = NewMemoryStore(1 * time.Minute)
store = pkgcap.NewMemoryStore(1 * time.Minute)
}
defaultManager = NewManager(Config{
@@ -9,11 +9,13 @@ import (
"sync/atomic"
"testing"
"time"
pkgcap "github.com/Rain-kl/Wavelet/pkg/cap"
)
func TestCapFullFlow(t *testing.T) {
secret := []byte("a-very-long-secret-key-at-least-16-bytes")
store := NewMemoryStore(1 * time.Minute)
store := pkgcap.NewMemoryStore(1 * time.Minute)
manager := NewManager(Config{
Secret: secret,
@@ -36,7 +38,7 @@ func TestCapFullFlow(t *testing.T) {
}
// Solve the challenge (acting as client)
solutions := Solve(resp.Token, resp.Challenge.C, resp.Challenge.S, resp.Challenge.D)
solutions := pkgcap.Solve(resp.Token, resp.Challenge.C, resp.Challenge.S, resp.Challenge.D)
// Redeem
redeemResp, err := manager.Redeem(ctx, resp.Token, solutions, scope)
@@ -76,7 +78,7 @@ func TestRedeemConcurrentRace(t *testing.T) {
const goroutines = 50
secret := []byte("race-test-secret-key-at-least-16-bytes")
store := NewMemoryStore(1 * time.Minute)
store := pkgcap.NewMemoryStore(1 * time.Minute)
manager := NewManager(Config{
Secret: secret,
ChallengeCount: 1,
@@ -91,7 +93,7 @@ func TestRedeemConcurrentRace(t *testing.T) {
if err != nil {
t.Fatalf("Generate failed: %v", err)
}
solutions := Solve(resp.Token, resp.Challenge.C, resp.Challenge.S, resp.Challenge.D)
solutions := pkgcap.Solve(resp.Token, resp.Challenge.C, resp.Challenge.S, resp.Challenge.D)
var (
wg sync.WaitGroup
@@ -125,7 +127,7 @@ func TestVerifyTokenConcurrentRace(t *testing.T) {
const goroutines = 50
secret := []byte("race-test-secret-key-at-least-16-bytes")
store := NewMemoryStore(1 * time.Minute)
store := pkgcap.NewMemoryStore(1 * time.Minute)
manager := NewManager(Config{
Secret: secret,
ChallengeCount: 1,
@@ -137,7 +139,7 @@ func TestVerifyTokenConcurrentRace(t *testing.T) {
ctx := context.Background()
resp, _ := manager.Generate(ctx, "login")
solutions := Solve(resp.Token, resp.Challenge.C, resp.Challenge.S, resp.Challenge.D)
solutions := pkgcap.Solve(resp.Token, resp.Challenge.C, resp.Challenge.S, resp.Challenge.D)
redeemResp, err := manager.Redeem(ctx, resp.Token, solutions, "login")
if err != nil || !redeemResp.Success {
t.Fatalf("Redeem failed: %v %+v", err, redeemResp)
+1 -1
View File
@@ -11,10 +11,10 @@ import (
"time"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/logger"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/storage"
"github.com/Rain-kl/Wavelet/internal/task"
"github.com/Rain-kl/Wavelet/pkg/logger"
"gorm.io/gorm"
)
+1 -1
View File
@@ -10,7 +10,7 @@ import (
"net/url"
"time"
"github.com/Rain-kl/Wavelet/internal/httppool"
"github.com/Rain-kl/Wavelet/pkg/httppool"
)
func getHTTPObject(ctx context.Context, baseURL, key string) (*Object, error) {
+2 -2
View File
@@ -10,7 +10,7 @@ import (
"path"
"strings"
"github.com/Rain-kl/Wavelet/internal/httppool"
"github.com/Rain-kl/Wavelet/pkg/httppool"
"github.com/studio-b12/gowebdav"
)
@@ -36,7 +36,7 @@ func (b *webDAVBackend) Put(_ context.Context, key string, body io.Reader, size
}
}
if err := b.client.WriteStreamWithLength(key, body, size, storageFilePerm); err != nil {
return PutResult{}, fmt.Errorf("put WebDAV object: %w", err)
return PutResult{}, fmt.Errorf("put WebDAV object: %w", err)
}
return PutResult{Key: key}, nil
}
+2 -2
View File
@@ -11,9 +11,9 @@ import (
"time"
"github.com/Rain-kl/Wavelet/internal/db/idgen"
"github.com/Rain-kl/Wavelet/internal/logger"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/otel_trace"
"github.com/Rain-kl/Wavelet/pkg/logger"
otel_trace "github.com/Rain-kl/Wavelet/pkg/trace"
"github.com/hibiken/asynq"
"go.opentelemetry.io/otel/attribute"
"go.opentelemetry.io/otel/codes"
+1 -1
View File
@@ -11,9 +11,9 @@ import (
"syscall"
"time"
"github.com/Rain-kl/Wavelet/internal/logger"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/task"
"github.com/Rain-kl/Wavelet/pkg/logger"
"github.com/hibiken/asynq"
)
-257
View File
@@ -1,257 +0,0 @@
// 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)
}
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
}
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
}
-15
View File
@@ -1,15 +0,0 @@
// 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"
)
-49
View File
@@ -1,49 +0,0 @@
// 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]
}
-186
View File
@@ -1,186 +0,0 @@
// 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
}
-149
View File
@@ -1,149 +0,0 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package util
import (
"crypto/aes"
"crypto/cipher"
"crypto/ed25519"
"crypto/rand"
"encoding/base64"
"encoding/hex"
"errors"
"fmt"
"io"
)
// aesKeyLength AES-256 密钥字节长度
const aesKeyLength = 32
// 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)
}
-8
View File
@@ -7,12 +7,4 @@ const (
errCreateHTTPRequestFailed = "创建HTTP请求失败: %w"
errHTTPRequestFailed = "请求%s接口失败: %w"
errInvalidCustomValue = "invalid value: %v"
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"
)
+3 -7
View File
@@ -12,7 +12,7 @@ import (
"net/url"
"time"
"go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp"
"github.com/Rain-kl/Wavelet/pkg/httppool"
)
// IsLocalhost 检查 URL 是否为 localhost
@@ -35,12 +35,8 @@ const (
// 配置HTTP客户端 使用 otelhttp 自动注入 trace span
var httpClient = &http.Client{
Timeout: httpClientTimeout * time.Second,
Transport: otelhttp.NewTransport(&http.Transport{
MaxIdleConns: httpMaxIdleConns,
MaxIdleConnsPerHost: httpMaxIdleConnsPerHost,
IdleConnTimeout: httpIdleConnTimeout * time.Second,
}),
Timeout: httpClientTimeout * time.Second,
Transport: httppool.DefaultTransport(),
}
// SetHTTPClient 替换全局 HTTP 客户端实例
-16
View File
@@ -1,16 +0,0 @@
// 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"
)
-241
View File
@@ -1,241 +0,0 @@
// 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
}
-92
View File
@@ -1,92 +0,0 @@
// 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)
}
}
-20
View File
@@ -1,20 +0,0 @@
// 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
}
-35
View File
@@ -1,35 +0,0 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package util
import "strings"
// emailPartsCount 邮箱地址由 @ 分割为两部分
const (
emailPartsCount = 2
emailLocalMinChars = 2 // 邮箱 local 部分掩码显示的最小字符数
)
// DerefString 安全地解引用字符串指针,nil 返回空字符串
func DerefString(s *string) string {
if s == nil {
return ""
}
return *s
}
// MaskEmail 安全脱敏用户的邮箱地址(例如 us***@example.com)
func MaskEmail(email string) string {
parts := strings.Split(email, "@")
if len(parts) != emailPartsCount {
return email
}
local := parts[0]
domain := parts[1]
if len(local) <= emailLocalMinChars {
return "**@" + domain
}
return local[:2] + "***" + local[len(local)-1:] + "@" + domain
}
-29
View File
@@ -1,29 +0,0 @@
// 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)
}