fix(persistence): migrate all pkg/persistence imports to plugins/infra/database and plugins/infra/cache

- Replace db.DB(ctx) with database.DB(ctx) from plugins/infra/database
- Replace db.Redis/db.PrefixedKey/db.GetJSON/db.SetJSON with cachepkg.* from plugins/infra/cache
- Replace pkg/persistence/idgen with pkg/idgen (already exists)
- Replace pkg/persistence/batchwriter with pkg/batchwriter (already exists)
- Replace pkg/persistence/migrator with pkg/migrator (already exists)
- Replace pkg/persistence/logstore with plugins/domain/risk_control/logstore
- Delete defunct pkg/{persistence,cap,message_gateway,push,shared,task}
- Fix vet issues: db alias in domain_test.go, driver_asynq_worker.TaskHandler reference
- Update Makefile architecture guard
- Update docs and skill references
- Update go.mod: gorilla/sessions promotion to direct dependency
This commit is contained in:
ryan
2026-08-28 10:59:24 +08:00
parent fb6a3edb89
commit 416603b616
223 changed files with 1304 additions and 10057 deletions
+309
View File
@@ -0,0 +1,309 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package batchwriter
import (
"context"
"errors"
"sync"
"sync/atomic"
"testing"
"time"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
type testEvent struct {
ID int
Data string
}
func testConfig() Config {
return Config{
Name: "test-writer",
QueueSize: 100,
MaxBatchSize: 5,
FlushInterval: 20 * time.Millisecond,
}
}
func TestWriter_BatchSizeFlush(t *testing.T) {
var (
mu sync.Mutex
batches [][]testEvent
flushWg sync.WaitGroup
)
flushWg.Add(1)
cfg := testConfig()
cfg.FlushInterval = time.Hour
w, err := New(cfg, func(_ context.Context, items []testEvent) error {
mu.Lock()
defer mu.Unlock()
batches = append(batches, items)
if len(items) == 5 {
flushWg.Done()
}
return nil
})
require.NoError(t, err)
w.Start(context.Background())
defer func() { _ = w.Stop(context.Background()) }()
for i := 1; i <= 5; i++ {
ok := w.TryEnqueue(testEvent{ID: i, Data: "payload"})
assert.True(t, ok)
}
done := make(chan struct{})
go func() {
flushWg.Wait()
close(done)
}()
select {
case <-done:
case <-time.After(2 * time.Second):
t.Fatal("timed out waiting for batch flush")
}
mu.Lock()
defer mu.Unlock()
require.Len(t, batches, 1)
assert.Len(t, batches[0], 5)
for i, item := range batches[0] {
assert.Equal(t, i+1, item.ID)
}
}
func TestWriter_IntervalFlush(t *testing.T) {
var (
mu sync.Mutex
flushed []testEvent
done = make(chan struct{})
)
cfg := testConfig()
cfg.MaxBatchSize = 100
cfg.FlushInterval = 30 * time.Millisecond
w, err := New(cfg, func(_ context.Context, items []testEvent) error {
mu.Lock()
defer mu.Unlock()
flushed = append(flushed, items...)
if len(flushed) == 2 {
select {
case <-done:
default:
close(done)
}
}
return nil
})
require.NoError(t, err)
w.Start(context.Background())
defer func() { _ = w.Stop(context.Background()) }()
assert.True(t, w.TryEnqueue(testEvent{ID: 1}))
assert.True(t, w.TryEnqueue(testEvent{ID: 2}))
select {
case <-done:
case <-time.After(2 * time.Second):
t.Fatal("timed out waiting for interval flush")
}
mu.Lock()
defer mu.Unlock()
assert.Len(t, flushed, 2)
}
func TestWriter_MinBatchSizeThreshold(t *testing.T) {
var (
mu sync.Mutex
flushed []testEvent
done = make(chan struct{})
)
cfg := testConfig()
cfg.MaxBatchSize = 100
cfg.MinBatchSize = 3
cfg.FlushInterval = 20 * time.Millisecond
cfg.MaxFlushWait = 60 * time.Millisecond
w, err := New(cfg, func(_ context.Context, items []testEvent) error {
mu.Lock()
defer mu.Unlock()
flushed = append(flushed, items...)
if len(flushed) == 2 {
select {
case <-done:
default:
close(done)
}
}
return nil
})
require.NoError(t, err)
w.Start(context.Background())
defer func() { _ = w.Stop(context.Background()) }()
assert.True(t, w.TryEnqueue(testEvent{ID: 1}))
assert.True(t, w.TryEnqueue(testEvent{ID: 2}))
time.Sleep(30 * time.Millisecond)
mu.Lock()
assert.Empty(t, flushed, "items should wait until MinBatchSize or MaxFlushWait")
mu.Unlock()
select {
case <-done:
case <-time.After(2 * time.Second):
t.Fatal("timed out waiting for forced max wait flush")
}
mu.Lock()
defer mu.Unlock()
assert.Len(t, flushed, 2)
}
func TestWriter_StopDrainsRemaining(t *testing.T) {
var (
mu sync.Mutex
flushed []testEvent
)
cfg := testConfig()
cfg.FlushInterval = time.Hour
w, err := New(cfg, func(_ context.Context, items []testEvent) error {
mu.Lock()
defer mu.Unlock()
flushed = append(flushed, items...)
return nil
})
require.NoError(t, err)
w.Start(context.Background())
for i := 1; i <= 3; i++ {
assert.True(t, w.TryEnqueue(testEvent{ID: i}))
}
ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second)
defer cancel()
require.NoError(t, w.Stop(ctx))
assert.False(t, w.Running())
mu.Lock()
defer mu.Unlock()
assert.Len(t, flushed, 3)
}
func TestWriter_DropWhenFull(t *testing.T) {
var (
dropped atomic.Int64
blockCh = make(chan struct{})
)
cfg := Config{
QueueSize: 2,
MaxBatchSize: 10,
FlushInterval: time.Hour,
}
w, err := New(cfg, func(_ context.Context, _ []testEvent) error {
<-blockCh
return nil
}, WithDropHandler(func(_ testEvent) {
dropped.Add(1)
}))
require.NoError(t, err)
w.Start(context.Background())
defer func() {
close(blockCh)
_ = w.Stop(context.Background())
}()
time.Sleep(10 * time.Millisecond)
w.ch <- testEvent{ID: 1}
w.ch <- testEvent{ID: 2}
assert.False(t, w.TryEnqueue(testEvent{ID: 3}))
assert.Equal(t, int64(1), dropped.Load())
assert.Equal(t, int64(1), w.Stats().Drops)
}
func TestWriter_FlushErrorCallback(t *testing.T) {
var (
called atomic.Bool
flushErr = errors.New("clickhouse write timeout")
done = make(chan struct{})
)
cfg := testConfig()
cfg.MaxBatchSize = 1
cfg.FlushInterval = time.Hour
w, err := New(cfg, func(_ context.Context, _ []testEvent) error {
return flushErr
}, WithFlushErrorHandler(func(_ context.Context, items []testEvent, err error) {
called.Store(true)
assert.Equal(t, flushErr, err)
assert.Len(t, items, 1)
close(done)
}))
require.NoError(t, err)
w.Start(context.Background())
defer func() { _ = w.Stop(context.Background()) }()
assert.True(t, w.TryEnqueue(testEvent{ID: 1}))
select {
case <-done:
case <-time.After(2 * time.Second):
t.Fatal("timed out waiting for error callback")
}
assert.True(t, called.Load())
assert.Equal(t, int64(1), w.Stats().FlushErrors)
}
func TestWriter_ValidateConfig(t *testing.T) {
tests := []struct {
name string
cfg Config
wantErr bool
}{
{"valid", DefaultConfig(), false},
{"zero queue", Config{QueueSize: 0, MaxBatchSize: 10, FlushInterval: time.Second}, true},
{"zero max batch", Config{QueueSize: 10, MaxBatchSize: 0, FlushInterval: time.Second}, true},
{"negative min batch", Config{QueueSize: 10, MaxBatchSize: 10, MinBatchSize: -1, FlushInterval: time.Second}, true},
{"zero flush interval", Config{QueueSize: 10, MaxBatchSize: 10, FlushInterval: 0}, true},
{"negative max flush wait", Config{QueueSize: 10, MaxBatchSize: 10, FlushInterval: time.Second, MaxFlushWait: -1}, true},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
_, err := New(tt.cfg, func(_ context.Context, _ []testEvent) error { return nil })
if tt.wantErr {
assert.Error(t, err)
} else {
assert.NoError(t, err)
}
})
}
}
func TestWriter_NilFlushFunc(t *testing.T) {
_, err := New[testEvent](DefaultConfig(), nil)
assert.ErrorIs(t, err, errNilFlushFunc)
}
-259
View File
@@ -1,259 +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)
}
// 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
}
-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
}
-21
View File
@@ -1,21 +0,0 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package message_gateway
import "context"
// Handler processes one inbound message.
type Handler func(ctx context.Context, msg InboundMessage) error
// Factory constructs a Channel from decrypted config.
type Factory func(cfg ChannelConfig, onInbound Handler) (Channel, error)
// Channel is one connected messaging adapter.
type Channel interface {
Type() string
Connect(ctx context.Context) error
Disconnect(ctx context.Context) error
Send(ctx context.Context, to Recipient, msg OutboundMessage) error
Capabilities() Capability
}
-161
View File
@@ -1,161 +0,0 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package qq implements the official QQ Bot C2C adapter.
package qq
import (
"context"
"fmt"
"strings"
"sync"
"time"
"github.com/Rain-kl/Wavelet/pkg/logger"
"github.com/Rain-kl/Wavelet/pkg/message_gateway"
"github.com/tencent-connect/botgo"
"github.com/tencent-connect/botgo/dto"
"github.com/tencent-connect/botgo/event"
"github.com/tencent-connect/botgo/openapi"
"github.com/tencent-connect/botgo/token"
"golang.org/x/oauth2"
)
// qqEvent is a testable inbound envelope.
type qqEvent struct {
Kind string
UserID string
Text string
MessageID string
}
// Adapter is an official QQ Bot C2C channel.
type Adapter struct {
cfg message_gateway.ChannelConfig
onInbound message_gateway.Handler
api openapi.OpenAPI
tokenSrc oauth2.TokenSource
cancel context.CancelFunc
mu sync.Mutex
disconnected bool
}
// New constructs a QQ adapter.
func New(cfg message_gateway.ChannelConfig, onInbound message_gateway.Handler) (message_gateway.Channel, error) {
if strings.TrimSpace(cfg.Credentials["app_id"]) == "" || strings.TrimSpace(cfg.Credentials["app_secret"]) == "" {
return nil, fmt.Errorf("qq: app_id and app_secret are required")
}
return &Adapter{cfg: cfg, onInbound: onInbound}, nil
}
// Type returns qq.
func (a *Adapter) Type() string { return message_gateway.ChannelTypeQQ }
// Capabilities reports C2C text/media support.
func (a *Adapter) Capabilities() message_gateway.Capability {
return message_gateway.Capability{Text: true, Image: true, File: true, Reply: true}
}
// Connect starts the official WebSocket session (C2C intent).
func (a *Adapter) Connect(ctx context.Context) error {
credentials := &token.QQBotCredentials{
AppID: a.cfg.Credentials["app_id"],
AppSecret: a.cfg.Credentials["app_secret"],
}
tokSrc := token.NewQQBotTokenSource(credentials)
runCtx, cancel := context.WithCancel(ctx)
if err := token.StartRefreshAccessToken(runCtx, tokSrc); err != nil {
cancel()
return fmt.Errorf("qq: refresh token: %w", err)
}
var api openapi.OpenAPI
const apiTimeout = 5 * time.Second
if strings.EqualFold(strings.TrimSpace(a.cfg.Extra["sandbox"]), "true") {
api = botgo.NewSandboxOpenAPI(credentials.AppID, tokSrc).WithTimeout(apiTimeout)
} else {
api = botgo.NewOpenAPI(credentials.AppID, tokSrc).WithTimeout(apiTimeout)
}
wsAP, err := api.WS(ctx, nil, "")
if err != nil {
cancel()
return fmt.Errorf("qq: websocket ap: %w", err)
}
intent := event.RegisterHandlers(event.C2CMessageEventHandler(func(_ *dto.WSPayload, data *dto.WSC2CMessageData) error {
authorID := ""
if data != nil && data.Author != nil {
authorID = data.Author.ID
}
text := ""
id := ""
if data != nil {
text = data.Content
id = data.ID
}
a.handleEvent(runCtx, qqEvent{Kind: "c2c", UserID: authorID, Text: text, MessageID: id})
return nil
}))
a.mu.Lock()
a.api = api
a.tokenSrc = tokSrc
a.cancel = cancel
a.disconnected = false
a.mu.Unlock()
go func() {
if err := botgo.NewSessionManager().Start(wsAP, tokSrc, &intent); err != nil {
logger.ErrorF(runCtx, "qq session stopped: %v", err)
}
}()
return nil
}
// Disconnect stops token refresh and drops further inbound events.
func (a *Adapter) Disconnect(_ context.Context) error {
a.mu.Lock()
defer a.mu.Unlock()
a.disconnected = true
if a.cancel != nil {
a.cancel()
a.cancel = nil
}
return nil
}
// Send posts a C2C text reply.
func (a *Adapter) Send(ctx context.Context, to message_gateway.Recipient, msg message_gateway.OutboundMessage) error {
a.mu.Lock()
api := a.api
a.mu.Unlock()
if api == nil {
return fmt.Errorf("qq: not connected")
}
_, err := api.PostC2CMessage(ctx, to.PlatformUserID, &dto.MessageToCreate{
Content: msg.Text,
MsgID: msg.ReplyToID,
})
return err
}
func (a *Adapter) handleEvent(ctx context.Context, ev qqEvent) {
if ev.Kind != "c2c" {
return
}
a.mu.Lock()
disconnected := a.disconnected
a.mu.Unlock()
if disconnected || a.onInbound == nil {
return
}
_ = a.onInbound(ctx, message_gateway.InboundMessage{
ChannelID: a.cfg.ID,
ChannelType: message_gateway.ChannelTypeQQ,
PlatformUserID: ev.UserID,
ChatID: ev.UserID,
MessageID: ev.MessageID,
Text: ev.Text,
})
}
@@ -1,42 +0,0 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package qq
import (
"context"
"testing"
"github.com/Rain-kl/Wavelet/pkg/message_gateway"
)
func TestHandleEvent_DropsNonC2C(t *testing.T) {
var got int
a := &Adapter{onInbound: func(ctx context.Context, msg message_gateway.InboundMessage) error {
got++
return nil
}}
a.handleEvent(context.Background(), qqEvent{Kind: "group", UserID: "u1", Text: "hi"})
if got != 0 {
t.Fatal("non-C2C must be ignored")
}
}
func TestHandleEvent_C2CText(t *testing.T) {
var got message_gateway.InboundMessage
a := &Adapter{cfg: message_gateway.ChannelConfig{ID: 3}, onInbound: func(ctx context.Context, msg message_gateway.InboundMessage) error {
got = msg
return nil
}}
a.handleEvent(context.Background(), qqEvent{Kind: "c2c", UserID: "openid-1", Text: "hello", MessageID: "m1"})
if got.Text != "hello" || got.PlatformUserID != "openid-1" || got.ChannelID != 3 {
t.Fatalf("%+v", got)
}
}
func TestNew_RequiresCreds(t *testing.T) {
_, err := New(message_gateway.ChannelConfig{}, nil)
if err == nil {
t.Fatal("expected error")
}
}
@@ -1,153 +0,0 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package telegram implements the Telegram private-chat adapter.
package telegram
import (
"context"
"fmt"
"os"
"path/filepath"
"strconv"
"strings"
"github.com/Rain-kl/Wavelet/pkg/message_gateway"
tele "gopkg.in/telebot.v4"
)
// Adapter is a Telegram private-chat channel.
type Adapter struct {
cfg message_gateway.ChannelConfig
onInbound message_gateway.Handler
bot *tele.Bot
}
// New constructs a Telegram adapter. Call message_gateway.Register from the runner.
func New(cfg message_gateway.ChannelConfig, onInbound message_gateway.Handler) (message_gateway.Channel, error) {
if strings.TrimSpace(cfg.Credentials["bot_token"]) == "" {
return nil, fmt.Errorf("telegram: bot_token is required")
}
return &Adapter{cfg: cfg, onInbound: onInbound}, nil
}
// Type returns telegram.
func (a *Adapter) Type() string { return message_gateway.ChannelTypeTelegram }
// Capabilities reports private-chat media support.
func (a *Adapter) Capabilities() message_gateway.Capability {
return message_gateway.Capability{Text: true, Image: true, File: true, Reply: true}
}
// Connect starts long polling.
func (a *Adapter) Connect(ctx context.Context) error {
pref := tele.Settings{
Token: a.cfg.Credentials["bot_token"],
Poller: &tele.LongPoller{Timeout: 10},
}
if base := strings.TrimSpace(a.cfg.Extra["base_url"]); base != "" {
pref.URL = strings.TrimSuffix(base, "/")
}
bot, err := tele.NewBot(pref)
if err != nil {
return fmt.Errorf("telegram: new bot: %w", err)
}
a.bot = bot
bot.Handle(tele.OnText, func(c tele.Context) error {
a.handleTeleMessage(ctx, c.Message())
return nil
})
bot.Handle(tele.OnPhoto, func(c tele.Context) error {
a.handleTeleMessage(ctx, c.Message())
return nil
})
bot.Handle(tele.OnDocument, func(c tele.Context) error {
a.handleTeleMessage(ctx, c.Message())
return nil
})
go bot.Start()
go func() {
<-ctx.Done()
bot.Stop()
}()
return nil
}
// Disconnect stops the bot.
func (a *Adapter) Disconnect(_ context.Context) error {
if a.bot != nil {
a.bot.Stop()
}
return nil
}
// Send replies to a private chat.
func (a *Adapter) Send(_ context.Context, to message_gateway.Recipient, msg message_gateway.OutboundMessage) error {
if a.bot == nil {
return fmt.Errorf("telegram: not connected")
}
chatID, err := strconv.ParseInt(to.ChatID, 10, 64)
if err != nil {
return fmt.Errorf("telegram: chat id: %w", err)
}
_, err = a.bot.Send(tele.ChatID(chatID), msg.Text)
return err
}
func (a *Adapter) handleTeleMessage(ctx context.Context, m *tele.Message) {
if m == nil || m.Chat == nil || m.Chat.Type != tele.ChatPrivate {
return
}
if a.onInbound == nil {
return
}
msg := message_gateway.InboundMessage{
ChannelID: a.cfg.ID,
ChannelType: message_gateway.ChannelTypeTelegram,
PlatformUserID: strconv.FormatInt(m.Sender.ID, 10),
ChatID: strconv.FormatInt(m.Chat.ID, 10),
MessageID: strconv.Itoa(m.ID),
Text: m.Text,
}
if m.Caption != "" && msg.Text == "" {
msg.Text = m.Caption
}
if a.bot != nil {
msg.Attachments = a.downloadMedia(m)
}
_ = a.onInbound(ctx, msg)
}
func (a *Adapter) downloadMedia(m *tele.Message) []message_gateway.Attachment {
var files []*tele.File
var names []string
if m.Photo != nil {
files = append(files, m.Photo.MediaFile())
names = append(names, "photo.jpg")
}
if m.Document != nil {
files = append(files, &m.Document.File)
name := m.Document.FileName
if name == "" {
name = "file"
}
names = append(names, name)
}
if len(files) == 0 {
return nil
}
dir, err := os.MkdirTemp("", "wg-tg-*")
if err != nil {
return []message_gateway.Attachment{{Error: err.Error()}}
}
out := make([]message_gateway.Attachment, 0, len(files))
for i, f := range files {
path := filepath.Join(dir, names[i])
if err := a.bot.Download(f, path); err != nil {
out = append(out, message_gateway.Attachment{FileName: names[i], Error: err.Error()})
continue
}
out = append(out, message_gateway.Attachment{Path: path, FileName: names[i]})
}
return out
}
@@ -1,56 +0,0 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package telegram
import (
"context"
"testing"
"github.com/Rain-kl/Wavelet/pkg/message_gateway"
tele "gopkg.in/telebot.v4"
)
func TestHandleUpdate_DropsGroups(t *testing.T) {
var got int
a := &Adapter{onInbound: func(ctx context.Context, msg message_gateway.InboundMessage) error {
got++
return nil
}}
a.handleTeleMessage(context.Background(), &tele.Message{
ID: 1,
Text: "hi",
Chat: &tele.Chat{ID: -100, Type: tele.ChatGroup},
Sender: &tele.User{ID: 1},
})
if got != 0 {
t.Fatalf("group must be ignored")
}
}
func TestHandleUpdate_PrivateText(t *testing.T) {
var got message_gateway.InboundMessage
a := &Adapter{
cfg: message_gateway.ChannelConfig{ID: 7, Type: "telegram"},
onInbound: func(ctx context.Context, msg message_gateway.InboundMessage) error {
got = msg
return nil
},
}
a.handleTeleMessage(context.Background(), &tele.Message{
ID: 9,
Text: "hi",
Chat: &tele.Chat{ID: 42, Type: tele.ChatPrivate},
Sender: &tele.User{ID: 42},
})
if got.Text != "hi" || got.PlatformUserID != "42" || got.ChannelID != 7 {
t.Fatalf("%+v", got)
}
}
func TestNew_RequiresToken(t *testing.T) {
_, err := New(message_gateway.ChannelConfig{}, nil)
if err == nil {
t.Fatal("expected error")
}
}
-50
View File
@@ -1,50 +0,0 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package message_gateway
import (
"crypto/rand"
"strings"
"unicode"
)
// CodeAlphabet excludes easily confused runes 0/O/1/I.
const CodeAlphabet = "ABCDEFGHJKLMNPQRSTUVWXYZ23456789"
// CodeLength is the raw pairing code size.
const CodeLength = 8
// GenerateCode returns an 8-character pairing code.
func GenerateCode() (string, error) {
buf := make([]byte, CodeLength)
if _, err := rand.Read(buf); err != nil {
return "", err
}
out := make([]byte, CodeLength)
for i, b := range buf {
out[i] = CodeAlphabet[int(b)%len(CodeAlphabet)]
}
return string(out), nil
}
// NormalizeCode strips separators and uppercases.
func NormalizeCode(s string) string {
var b strings.Builder
for _, r := range s {
if r == '-' || unicode.IsSpace(r) {
continue
}
b.WriteRune(unicode.ToUpper(r))
}
return b.String()
}
// FormatCode renders ABCD-EFGH.
func FormatCode(s string) string {
s = NormalizeCode(s)
if len(s) != CodeLength {
return s
}
return s[:4] + "-" + s[4:]
}
-33
View File
@@ -1,33 +0,0 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package message_gateway
import (
"strings"
"testing"
)
func TestGenerateCode_AlphabetAndLength(t *testing.T) {
code, err := GenerateCode()
if err != nil {
t.Fatal(err)
}
if len(code) != 8 {
t.Fatalf("len=%d", len(code))
}
for _, r := range code {
if !strings.ContainsRune(CodeAlphabet, r) {
t.Fatalf("bad rune %q", r)
}
}
}
func TestNormalizeAndFormat(t *testing.T) {
if got := NormalizeCode("ab-cd-ef-gh"); got != "ABCDEFGH" {
t.Fatalf("got %q", got)
}
if got := FormatCode("ABCDEFGH"); got != "ABCD-EFGH" {
t.Fatalf("got %q", got)
}
}
-26
View File
@@ -1,26 +0,0 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package message_gateway
import "sync"
var (
factoriesMu sync.RWMutex
factories = map[string]Factory{}
)
// Register stores a channel factory under typ.
func Register(typ string, fn Factory) {
factoriesMu.Lock()
defer factoriesMu.Unlock()
factories[typ] = fn
}
// Lookup returns a previously registered factory.
func Lookup(typ string) (Factory, bool) {
factoriesMu.RLock()
defer factoriesMu.RUnlock()
fn, ok := factories[typ]
return fn, ok
}
-38
View File
@@ -1,38 +0,0 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package message_gateway
import (
"context"
"testing"
)
type stubChannel struct{}
func (stubChannel) Type() string { return "stub" }
func (stubChannel) Connect(context.Context) error {
return nil
}
func (stubChannel) Disconnect(context.Context) error { return nil }
func (stubChannel) Send(context.Context, Recipient, OutboundMessage) error {
return nil
}
func (stubChannel) Capabilities() Capability { return Capability{Text: true} }
func TestRegisterLookup(t *testing.T) {
Register("stub", func(ChannelConfig, Handler) (Channel, error) {
return stubChannel{}, nil
})
fn, ok := Lookup("stub")
if !ok {
t.Fatal("expected factory")
}
ch, err := fn(ChannelConfig{}, nil)
if err != nil {
t.Fatal(err)
}
if ch.Type() != "stub" {
t.Fatalf("type=%s", ch.Type())
}
}
-62
View File
@@ -1,62 +0,0 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package message_gateway defines channel adapters, pairing codes, and inbound types.
package message_gateway
// ChannelTypeTelegram is the Telegram private-chat adapter type.
const ChannelTypeTelegram = "telegram"
// ChannelTypeQQ is the official QQ Bot C2C adapter type.
const ChannelTypeQQ = "qq"
// Capability describes what an adapter can send and receive.
type Capability struct {
Text bool
Image bool
File bool
Reply bool
Group bool
}
// ChannelConfig is the decrypted runtime config passed to a factory.
type ChannelConfig struct {
ID uint64
Type string
Name string
Credentials map[string]string
Extra map[string]string
}
// Recipient is the outbound destination on a platform.
type Recipient struct {
ChatID string
PlatformUserID string
}
// Attachment is a downloaded inbound file sitting on local disk.
type Attachment struct {
Path string
FileName string
MIME string
Error string
}
// InboundMessage is a normalized private-chat message.
type InboundMessage struct {
ChannelID uint64
ChannelType string
PlatformUserID string
ChatID string
MessageID string
Text string
Attachments []Attachment
BindingUserID *uint64
}
// OutboundMessage is a reply or probe send.
type OutboundMessage struct {
Text string
ReplyToID string
Attachments []Attachment
}
@@ -11,8 +11,7 @@ import (
"log"
"github.com/Rain-kl/Wavelet/pkg/config"
"github.com/Rain-kl/Wavelet/pkg/persistence"
"github.com/Rain-kl/Wavelet/plugins/infra/database"
"github.com/pressly/goose/v3"
)
@@ -59,7 +58,7 @@ func migrationDir() string {
// Migrate 执行数据库迁移
func Migrate() Report {
gormDB := db.DB(context.Background())
gormDB := database.DB(context.Background())
if gormDB == nil {
log.Fatalf("[%s] database not initialized\n", dbType())
}
@@ -7,7 +7,7 @@ import (
"testing"
"github.com/Rain-kl/Wavelet/pkg/config"
db "github.com/Rain-kl/Wavelet/pkg/persistence"
db "github.com/Rain-kl/Wavelet/plugins/infra/database"
"github.com/alicebob/miniredis/v2"
"github.com/glebarez/sqlite"
"github.com/redis/go-redis/v9"
-120
View File
@@ -1,120 +0,0 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package analytics provides ClickHouse data access for analytics tables.
package analytics
import (
"context"
"fmt"
"time"
"github.com/Rain-kl/Wavelet/pkg/persistence"
"github.com/Rain-kl/Wavelet/pkg/util"
"gorm.io/gorm"
)
// CountAccessLogs returns the number of access logs matching filter.
func CountAccessLogs(ctx context.Context, filter AccessLogFilter) (uint64, error) {
ch := db.ChDB(ctx)
if ch == nil {
return 0, fmt.Errorf("clickhouse gorm connection is not initialized")
}
var count int64
query := applyFilter(ch.Model(&UserAccessLog{}), filter)
if err := query.Count(&count).Error; err != nil {
return 0, fmt.Errorf("count access logs: %w", err)
}
return safeUint64Count(count), nil
}
// ListAccessLogs returns paginated access logs and the total match count.
func ListAccessLogs(ctx context.Context, filter AccessLogFilter, page, pageSize int) ([]UserAccessLog, uint64, error) {
ch := db.ChDB(ctx)
if ch == nil {
return nil, 0, fmt.Errorf("clickhouse gorm connection is not initialized")
}
if filter.UserIDs != nil && len(filter.UserIDs) == 0 {
return []UserAccessLog{}, 0, nil
}
var total int64
baseQuery := applyFilter(ch.Model(&UserAccessLog{}), filter)
if err := baseQuery.Count(&total).Error; err != nil {
return nil, 0, fmt.Errorf("count access logs: %w", err)
}
if total == 0 {
return []UserAccessLog{}, 0, nil
}
if page < 1 {
page = 1
}
if pageSize < 1 {
pageSize = 20
}
offset := (page - 1) * pageSize
var logs []UserAccessLog
err := applyFilter(ch.Model(&UserAccessLog{}), filter).
Order("created_at DESC, id DESC").
Limit(pageSize).
Offset(offset).
Find(&logs).Error
if err != nil {
return nil, 0, fmt.Errorf("list access logs: %w", err)
}
return logs, safeUint64Count(total), nil
}
// DeleteAllUserAccessLogs hard-deletes all user access logs via TRUNCATE.
func DeleteAllUserAccessLogs(ctx context.Context) (int64, error) {
if db.ChConn == nil {
return 0, fmt.Errorf("clickhouse connection is not initialized")
}
if err := db.ChConn.Exec(ctx, "TRUNCATE TABLE "+UserAccessLog{}.TableName()); err != nil {
return 0, fmt.Errorf("truncate user access logs: %w", err)
}
return 0, nil
}
// DeleteUserAccessLogsBefore deletes user access logs older than cutoff.
func DeleteUserAccessLogsBefore(ctx context.Context, cutoff time.Time) (int64, error) {
if db.ChConn == nil {
return 0, fmt.Errorf("clickhouse connection is not initialized")
}
if err := db.ChConn.Exec(ctx, "ALTER TABLE "+UserAccessLog{}.TableName()+" DELETE WHERE created_at < ?", cutoff); err != nil {
return 0, fmt.Errorf("delete expired user access logs: %w", err)
}
return 0, nil
}
func safeUint64Count(count int64) uint64 {
if count < 0 {
return 0
}
return uint64(count)
}
func applyFilter(query *gorm.DB, filter AccessLogFilter) *gorm.DB {
if filter.UserIDs != nil {
if len(filter.UserIDs) == 0 {
return query.Where("1 = 0")
}
query = query.Where("user_id IN ?", filter.UserIDs)
}
if filter.Path != "" {
query = query.Where("path LIKE ?", "%"+util.EscapeLike(filter.Path)+"%")
}
if filter.StartTime != nil {
query = query.Where("created_at >= ?", *filter.StartTime)
}
if filter.EndTime != nil {
query = query.Where("created_at <= ?", *filter.EndTime)
}
return query
}
@@ -1,17 +0,0 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package analytics
import "time"
// AccessLogFilter scopes ClickHouse user access log queries.
type AccessLogFilter struct {
// UserIDs filters by user IDs. nil means no user filter; an empty slice means no matches.
UserIDs []uint64
Path string
// StartTime filters created_at >= StartTime when non-nil.
StartTime *time.Time
// EndTime filters created_at <= EndTime when non-nil.
EndTime *time.Time
}
@@ -1,158 +0,0 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package analytics
import (
"context"
"fmt"
"sort"
"time"
"github.com/Rain-kl/Wavelet/pkg/persistence"
)
const hoursInDay = 24
// DailyTrend is a single day's access count.
type DailyTrend struct {
Date string
Count uint64
}
// BrowserShare is a browser group's share of access logs.
type BrowserShare struct {
Browser string
Count uint64
}
// TopUser is an active user ranked by access count.
type TopUser struct {
UserID uint64
Count uint64
}
// GetDailyTrend returns per-day access counts for the last days days (inclusive of today).
func GetDailyTrend(ctx context.Context, days int) ([]DailyTrend, error) {
if days < 1 {
days = 7
}
ch := db.ChDB(ctx)
if ch == nil {
return nil, fmt.Errorf("clickhouse gorm connection is not initialized")
}
startTime := time.Now().AddDate(0, 0, -(days - 1)).Truncate(hoursInDay * time.Hour)
tableName := UserAccessLog{}.TableName()
query := fmt.Sprintf(`
SELECT toDate(created_at) AS date, count() AS count
FROM %s
WHERE created_at >= ?
GROUP BY date
ORDER BY date ASC
`, tableName)
type trendRow struct {
Date time.Time
Count uint64
}
var rows []trendRow
if err := ch.Raw(query, startTime).Scan(&rows).Error; err != nil {
return nil, fmt.Errorf("get daily trend: %w", err)
}
trendMap := make(map[string]uint64, days)
for i := 0; i < days; i++ {
dateStr := time.Now().AddDate(0, 0, -i).Format("2006-01-02")
trendMap[dateStr] = 0
}
for _, row := range rows {
dateStr := row.Date.Format("2006-01-02")
trendMap[dateStr] = row.Count
}
result := make([]DailyTrend, 0, days)
for i := days - 1; i >= 0; i-- {
dateStr := time.Now().AddDate(0, 0, -i).Format("2006-01-02")
result = append(result, DailyTrend{
Date: dateStr,
Count: trendMap[dateStr],
})
}
return result, nil
}
// GetBrowserDistribution returns browser-grouped access counts since startTime.
func GetBrowserDistribution(ctx context.Context, startTime time.Time) ([]BrowserShare, error) {
ch := db.ChDB(ctx)
if ch == nil {
return nil, fmt.Errorf("clickhouse gorm connection is not initialized")
}
tableName := UserAccessLog{}.TableName()
query := fmt.Sprintf(`
SELECT user_agent, count() AS count
FROM %s
WHERE created_at >= ?
GROUP BY user_agent
`, tableName)
type uaRow struct {
UserAgent string
Count uint64
}
var rows []uaRow
if err := ch.Raw(query, startTime).Scan(&rows).Error; err != nil {
return nil, fmt.Errorf("get browser distribution: %w", err)
}
browserCounts := make(map[string]uint64)
for _, row := range rows {
browser := ParseBrowserName(row.UserAgent)
browserCounts[browser] += row.Count
}
result := make([]BrowserShare, 0, len(browserCounts))
for browser, count := range browserCounts {
result = append(result, BrowserShare{
Browser: browser,
Count: count,
})
}
sort.Slice(result, func(i, j int) bool {
return result[i].Count > result[j].Count
})
return result, nil
}
// GetTopActiveUsers returns the most active users since startTime.
func GetTopActiveUsers(ctx context.Context, startTime time.Time, limit int) ([]TopUser, error) {
if limit < 1 {
limit = 10
}
ch := db.ChDB(ctx)
if ch == nil {
return nil, fmt.Errorf("clickhouse gorm connection is not initialized")
}
tableName := UserAccessLog{}.TableName()
query := fmt.Sprintf(`
SELECT user_id, count() AS count
FROM %s
WHERE created_at >= ? AND user_id > 0
GROUP BY user_id
ORDER BY count DESC
LIMIT ?
`, tableName)
var users []TopUser
if err := ch.Raw(query, startTime, limit).Scan(&users).Error; err != nil {
return nil, fmt.Errorf("get top active users: %w", err)
}
return users, nil
}
@@ -1,214 +0,0 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package analytics
import (
"context"
"io"
"testing"
"time"
"github.com/ClickHouse/clickhouse-go/v2/lib/column"
"github.com/ClickHouse/clickhouse-go/v2/lib/driver"
"github.com/Rain-kl/Wavelet/pkg/persistence"
"github.com/glebarez/sqlite"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"gorm.io/gorm"
)
func setupChGormDB(t *testing.T) *gorm.DB {
t.Helper()
gormDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{
DisableForeignKeyConstraintWhenMigrating: true,
})
require.NoError(t, err)
require.NoError(t, gormDB.AutoMigrate(&UserAccessLog{}))
db.SetChDBForTest(gormDB)
return gormDB
}
func TestParseBrowserName(t *testing.T) {
tests := []struct {
name string
ua string
want string
}{
{name: "chrome", ua: "Mozilla/5.0 Chrome/120.0.0.0", want: "Chrome"},
{name: "firefox", ua: "Mozilla/5.0 Firefox/121.0", want: "Firefox"},
{name: "safari", ua: "Mozilla/5.0 Safari/605.1.15", want: "Safari"},
{name: "edge", ua: "Mozilla/5.0 Edg/120.0.0.0", want: "Edge"},
{name: "wechat", ua: "MicroMessenger/8.0", want: "WeChat"},
{name: "postman", ua: "PostmanRuntime/7.36.0", want: "Postman"},
{name: "other", ua: "curl/8.0", want: "Other"},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
assert.Equal(t, tt.want, ParseBrowserName(tt.ua))
})
}
}
func TestCountAccessLogs_EmptyUserIDs(t *testing.T) {
setupChGormDB(t)
t.Cleanup(func() { db.SetChDBForTest(nil) })
count, err := CountAccessLogs(context.Background(), AccessLogFilter{UserIDs: []uint64{}})
require.NoError(t, err)
assert.Equal(t, uint64(0), count)
}
func TestListAccessLogs_EmptyUserIDs(t *testing.T) {
setupChGormDB(t)
t.Cleanup(func() { db.SetChDBForTest(nil) })
logs, total, err := ListAccessLogs(context.Background(), AccessLogFilter{UserIDs: []uint64{}}, 1, 20)
require.NoError(t, err)
assert.Equal(t, uint64(0), total)
assert.Empty(t, logs)
}
func TestListAccessLogs_WithFilters(t *testing.T) {
gormDB := setupChGormDB(t)
t.Cleanup(func() { db.SetChDBForTest(nil) })
now := time.Now().UTC().Truncate(time.Second)
logs := []UserAccessLog{
{ID: 1, UserID: 10, Path: "/api/v1/users", Method: "GET", Status: 200, CreatedAt: now},
{ID: 2, UserID: 20, Path: "/api/v1/admin/logs", Method: "GET", Status: 200, CreatedAt: now},
{ID: 3, UserID: 10, Path: "/api/v1/other", Method: "POST", Status: 201, CreatedAt: now},
}
require.NoError(t, gormDB.Create(&logs).Error)
start := now.Add(-time.Hour)
filter := AccessLogFilter{
UserIDs: []uint64{10},
Path: "users",
StartTime: &start,
}
count, err := CountAccessLogs(context.Background(), filter)
require.NoError(t, err)
assert.Equal(t, uint64(1), count)
result, total, err := ListAccessLogs(context.Background(), filter, 1, 10)
require.NoError(t, err)
assert.Equal(t, uint64(1), total)
require.Len(t, result, 1)
assert.Equal(t, uint64(1), result[0].ID)
assert.Equal(t, "/api/v1/users", result[0].Path)
}
func TestBatchInsert_Empty(t *testing.T) {
err := BatchInsert(context.Background(), nil)
require.NoError(t, err)
}
func TestBatchInsert_UsesModelBatchSQL(t *testing.T) {
ctx := context.Background()
mockBatch := &mockBatch{}
mockConn := &mockConn{
batch: mockBatch,
batchQuery: UserAccessLog{}.BatchInsertSQL(),
}
db.SetChConnForTest(mockConn)
t.Cleanup(func() { db.SetChConnForTest(nil) })
createdAt := time.Now().UTC()
err := BatchInsert(ctx, []UserAccessLog{
{
ID: 1,
UserID: 42,
Path: "/api/v1/test",
Method: "GET",
IP: "127.0.0.1",
UserAgent: "test-agent",
Headers: "{}",
Status: 200,
Latency: 12,
CreatedAt: createdAt,
},
})
require.NoError(t, err)
assert.True(t, mockConn.prepareCalled)
assert.Equal(t, UserAccessLog{}.BatchInsertSQL(), mockConn.preparedQuery)
assert.True(t, mockBatch.sendCalled)
require.Len(t, mockBatch.rows, 1)
assert.Equal(t, uint64(42), mockBatch.rows[0][1])
}
type mockConn struct {
batch driver.Batch
batchQuery string
prepareCalled bool
preparedQuery string
}
func (m *mockConn) Contributors() []string { return nil }
func (m *mockConn) ServerVersion() (*driver.ServerVersion, error) { return nil, nil }
func (m *mockConn) Select(_ context.Context, _ any, _ string, _ ...any) error { return nil }
func (m *mockConn) Query(_ context.Context, _ string, _ ...any) (driver.Rows, error) {
return nil, nil
}
func (m *mockConn) QueryRow(_ context.Context, _ string, _ ...any) driver.Row { return nil }
func (m *mockConn) PrepareBatch(_ context.Context, query string, _ ...driver.PrepareBatchOption) (driver.Batch, error) {
m.prepareCalled = true
m.preparedQuery = query
return m.batch, nil
}
func (m *mockConn) Exec(_ context.Context, _ string, _ ...any) error { return nil }
func (m *mockConn) AsyncInsert(_ context.Context, _ string, _ bool, _ ...any) error { return nil }
func (m *mockConn) InsertFormat(_ context.Context, _ string, _ string, _ io.Reader) error { return nil }
func (m *mockConn) QueryFormat(_ context.Context, _ string, _ string, _ ...any) (io.ReadCloser, error) {
return nil, nil
}
func (m *mockConn) Ping(_ context.Context) error { return nil }
func (m *mockConn) Stats() driver.Stats { return driver.Stats{} }
func (m *mockConn) Close() error { return nil }
type mockBatch struct {
rows [][]any
sendCalled bool
}
func (m *mockBatch) Abort() error { return nil }
func (m *mockBatch) Append(v ...any) error {
m.rows = append(m.rows, v)
return nil
}
func (m *mockBatch) AppendStruct(_ any) error { return nil }
func (m *mockBatch) Column(_ int) driver.BatchColumn { return nil }
func (m *mockBatch) Flush() error { return nil }
func (m *mockBatch) Send() error {
m.sendCalled = true
return nil
}
func (m *mockBatch) IsSent() bool { return m.sendCalled }
func (m *mockBatch) Rows() int { return len(m.rows) }
func (m *mockBatch) Columns() []column.Interface { return nil }
func (m *mockBatch) Close() error { return nil }
@@ -1,48 +0,0 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package analytics
import (
"context"
"fmt"
"github.com/Rain-kl/Wavelet/pkg/persistence"
)
// BatchInsert writes access logs to ClickHouse using the native batch API.
func BatchInsert(ctx context.Context, logs []UserAccessLog) error {
if len(logs) == 0 {
return nil
}
if db.ChConn == nil {
return fmt.Errorf("clickhouse connection is not initialized")
}
batch, err := db.ChConn.PrepareBatch(ctx, UserAccessLog{}.BatchInsertSQL())
if err != nil {
return fmt.Errorf("prepare clickhouse batch: %w", err)
}
for _, logItem := range logs {
if err := batch.Append(
logItem.ID,
logItem.UserID,
logItem.Path,
logItem.Method,
logItem.IP,
logItem.UserAgent,
logItem.Headers,
logItem.Status,
logItem.Latency,
logItem.CreatedAt,
); err != nil {
return fmt.Errorf("append access log to batch: %w", err)
}
}
if err := batch.Send(); err != nil {
return fmt.Errorf("send clickhouse batch: %w", err)
}
return nil
}
-30
View File
@@ -1,30 +0,0 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package analytics
import "strings"
// ParseBrowserName performs lightweight User-Agent browser identification.
func ParseBrowserName(ua string) string {
uaLower := strings.ToLower(ua)
if strings.Contains(uaLower, "micromessenger") {
return "WeChat"
}
if strings.Contains(uaLower, "postman") {
return "Postman"
}
if strings.Contains(uaLower, "edg/") || strings.Contains(uaLower, "edge") {
return "Edge"
}
if strings.Contains(uaLower, "firefox") {
return "Firefox"
}
if strings.Contains(uaLower, "chrome") {
return "Chrome"
}
if strings.Contains(uaLower, "safari") {
return "Safari"
}
return "Other"
}
-44
View File
@@ -1,44 +0,0 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package analytics defines ClickHouse analytics domain models.
package analytics
import (
"fmt"
"time"
)
const (
userAccessLogTableName = "w_user_access_logs"
userAccessLogInsertColumns = "id, user_id, path, method, ip, user_agent, headers, status, latency, created_at"
)
// UserAccessLog stores HTTP access records in ClickHouse.
type UserAccessLog struct {
ID uint64 `gorm:"column:id"`
UserID uint64 `gorm:"column:user_id"`
Path string `gorm:"column:path"`
Method string `gorm:"column:method"`
IP string `gorm:"column:ip"`
UserAgent string `gorm:"column:user_agent"`
Headers string `gorm:"column:headers"`
Status int32 `gorm:"column:status"`
Latency int64 `gorm:"column:latency"`
CreatedAt time.Time `gorm:"column:created_at"`
}
// TableName returns the ClickHouse table name.
func (UserAccessLog) TableName() string {
return userAccessLogTableName
}
// InsertColumns returns comma-separated column names for batch insert.
func (UserAccessLog) InsertColumns() string {
return userAccessLogInsertColumns
}
// BatchInsertSQL returns the INSERT prefix used by native batch writers.
func (UserAccessLog) BatchInsertSQL() string {
return fmt.Sprintf("INSERT INTO %s (%s)", userAccessLogTableName, userAccessLogInsertColumns)
}
-491
View File
@@ -1,491 +0,0 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package batchwriter
import (
"context"
"errors"
"sync"
"testing"
"time"
"github.com/google/go-cmp/cmp"
)
func TestNewRejectsInvalidConfig(t *testing.T) {
t.Parallel()
_, err := New[int](Config{}, func(context.Context, []int) error { return nil })
if err == nil {
t.Fatal("New() = nil, want validation error")
}
}
func TestNewRejectsNilFlushFunc(t *testing.T) {
t.Parallel()
cfg := DefaultConfig()
_, err := New[int](cfg, nil)
if !errors.Is(err, errNilFlushFunc) {
t.Fatalf("New() error = %v, want %v", err, errNilFlushFunc)
}
}
func TestWriterFlushesOnMaxBatchSize(t *testing.T) {
t.Parallel()
var (
mu sync.Mutex
batches [][]int
)
cfg := DefaultConfig()
cfg.MaxBatchSize = 3
cfg.FlushInterval = time.Hour
writer, err := New[int](cfg, func(_ context.Context, items []int) error {
mu.Lock()
defer mu.Unlock()
batches = append(batches, append([]int(nil), items...))
return nil
})
if err != nil {
t.Fatalf("New() error = %v", err)
}
writer.Start(context.Background())
t.Cleanup(func() {
stopCtx, cancel := context.WithTimeout(context.Background(), time.Second)
defer cancel()
if err := writer.Stop(stopCtx); err != nil {
t.Fatalf("Stop() error = %v", err)
}
})
for i := range 3 {
if !writer.TryEnqueue(i + 1) {
t.Fatalf("TryEnqueue(%d) = false, want true", i+1)
}
}
deadline := time.Now().Add(time.Second)
for {
mu.Lock()
ready := len(batches) == 1
mu.Unlock()
if ready || time.Now().After(deadline) {
break
}
time.Sleep(10 * time.Millisecond)
}
mu.Lock()
got := batches
mu.Unlock()
want := [][]int{{1, 2, 3}}
if diff := cmp.Diff(want, got); diff != "" {
t.Fatalf("flush batches mismatch (-want +got):\n%s", diff)
}
}
func TestWriterFlushesOnInterval(t *testing.T) {
t.Parallel()
var (
mu sync.Mutex
batch []int
)
cfg := DefaultConfig()
cfg.MaxBatchSize = 100
cfg.MinBatchSize = 0
cfg.FlushInterval = 20 * time.Millisecond
writer, err := New[int](cfg, func(_ context.Context, items []int) error {
mu.Lock()
defer mu.Unlock()
batch = append([]int(nil), items...)
return nil
})
if err != nil {
t.Fatalf("New() error = %v", err)
}
writer.Start(context.Background())
t.Cleanup(func() {
stopCtx, cancel := context.WithTimeout(context.Background(), time.Second)
defer cancel()
if err := writer.Stop(stopCtx); err != nil {
t.Fatalf("Stop() error = %v", err)
}
})
if !writer.TryEnqueue(42) {
t.Fatal("TryEnqueue() = false, want true")
}
deadline := time.Now().Add(time.Second)
for {
mu.Lock()
ready := len(batch) == 1
mu.Unlock()
if ready || time.Now().After(deadline) {
break
}
time.Sleep(5 * time.Millisecond)
}
mu.Lock()
got := batch
mu.Unlock()
want := []int{42}
if diff := cmp.Diff(want, got); diff != "" {
t.Fatalf("interval flush mismatch (-want +got):\n%s", diff)
}
}
func TestWriterTryEnqueueDropsWhenFull(t *testing.T) {
t.Parallel()
cfg := DefaultConfig()
cfg.QueueSize = 1
cfg.MaxBatchSize = 10
cfg.FlushInterval = time.Hour
var dropped int
writer, err := New[int](cfg, func(context.Context, []int) error { return nil }, WithDropHandler[int](func(int) {
dropped++
}))
if err != nil {
t.Fatalf("New() error = %v", err)
}
writer.Start(context.Background())
t.Cleanup(func() {
stopCtx, cancel := context.WithTimeout(context.Background(), time.Second)
defer cancel()
_ = writer.Stop(stopCtx)
})
if !writer.TryEnqueue(1) {
t.Fatal("TryEnqueue(1) = false, want true")
}
if writer.TryEnqueue(2) {
t.Fatal("TryEnqueue(2) = true, want false")
}
if !writer.IsFull() {
t.Fatal("IsFull() = false, want true")
}
if dropped != 1 {
t.Fatalf("dropped = %d, want 1", dropped)
}
}
func TestWriterStopDrainsQueuedItems(t *testing.T) {
t.Parallel()
cfg := DefaultConfig()
cfg.MaxBatchSize = 10
cfg.FlushInterval = time.Hour
var flushed []int
writer, err := New[int](cfg, func(_ context.Context, items []int) error {
flushed = append(flushed, items...)
return nil
})
if err != nil {
t.Fatalf("New() error = %v", err)
}
writer.Start(context.Background())
for i := range 2 {
if !writer.TryEnqueue(i + 1) {
t.Fatalf("TryEnqueue(%d) = false, want true", i+1)
}
}
stopCtx, cancel := context.WithTimeout(context.Background(), time.Second)
defer cancel()
if err := writer.Stop(stopCtx); err != nil {
t.Fatalf("Stop() error = %v", err)
}
want := []int{1, 2}
if diff := cmp.Diff(want, flushed); diff != "" {
t.Fatalf("Stop() drain mismatch (-want +got):\n%s", diff)
}
if writer.Running() {
t.Fatal("Running() = true after Stop(), want false")
}
}
func TestWriterInvokesFlushErrorHandler(t *testing.T) {
t.Parallel()
cfg := DefaultConfig()
cfg.MaxBatchSize = 1
cfg.FlushInterval = time.Hour
flushErr := errors.New("flush failed")
var (
mu sync.Mutex
errCount int
gotItems []int
)
writer, err := New[int](cfg, func(context.Context, []int) error {
return flushErr
}, WithFlushErrorHandler[int](func(_ context.Context, items []int, err error) {
mu.Lock()
defer mu.Unlock()
errCount++
gotItems = append([]int(nil), items...)
if !errors.Is(err, flushErr) {
t.Errorf("flush error = %v, want %v", err, flushErr)
}
}))
if err != nil {
t.Fatalf("New() error = %v", err)
}
writer.Start(context.Background())
t.Cleanup(func() {
stopCtx, cancel := context.WithTimeout(context.Background(), time.Second)
defer cancel()
_ = writer.Stop(stopCtx)
})
if !writer.TryEnqueue(7) {
t.Fatal("TryEnqueue() = false, want true")
}
deadline := time.Now().Add(time.Second)
for {
mu.Lock()
ready := errCount == 1
mu.Unlock()
if ready || time.Now().After(deadline) {
break
}
time.Sleep(5 * time.Millisecond)
}
mu.Lock()
gotCount := errCount
items := gotItems
mu.Unlock()
if gotCount != 1 {
t.Fatalf("flush error handler count = %d, want 1", gotCount)
}
if diff := cmp.Diff([]int{7}, items); diff != "" {
t.Fatalf("flush error handler items mismatch (-want +got):\n%s", diff)
}
stats := writer.Stats()
if stats.FlushErrors != 1 {
t.Fatalf("Stats().FlushErrors = %d, want 1", stats.FlushErrors)
}
}
func TestWriterSkipsIntervalFlushBelowMinBatchSize(t *testing.T) {
t.Parallel()
var (
mu sync.Mutex
batch []int
)
cfg := DefaultConfig()
cfg.MaxBatchSize = 100
cfg.MinBatchSize = 5
cfg.FlushInterval = 20 * time.Millisecond
writer, err := New[int](cfg, func(_ context.Context, items []int) error {
mu.Lock()
defer mu.Unlock()
batch = append([]int(nil), items...)
return nil
})
if err != nil {
t.Fatalf("New() error = %v", err)
}
writer.Start(context.Background())
t.Cleanup(func() {
stopCtx, cancel := context.WithTimeout(context.Background(), time.Second)
defer cancel()
if err := writer.Stop(stopCtx); err != nil {
t.Fatalf("Stop() error = %v", err)
}
})
for i := range 3 {
if !writer.TryEnqueue(i + 1) {
t.Fatalf("TryEnqueue(%d) = false, want true", i+1)
}
}
time.Sleep(100 * time.Millisecond)
mu.Lock()
got := batch
mu.Unlock()
if len(got) != 0 {
t.Fatalf("interval flush with below-min batch = %v, want no flush", got)
}
}
func TestWriterFlushesOnIntervalWhenMinBatchSizeReached(t *testing.T) {
t.Parallel()
var (
mu sync.Mutex
batch []int
)
cfg := DefaultConfig()
cfg.MaxBatchSize = 100
cfg.MinBatchSize = 3
cfg.FlushInterval = 20 * time.Millisecond
writer, err := New[int](cfg, func(_ context.Context, items []int) error {
mu.Lock()
defer mu.Unlock()
batch = append([]int(nil), items...)
return nil
})
if err != nil {
t.Fatalf("New() error = %v", err)
}
writer.Start(context.Background())
t.Cleanup(func() {
stopCtx, cancel := context.WithTimeout(context.Background(), time.Second)
defer cancel()
if err := writer.Stop(stopCtx); err != nil {
t.Fatalf("Stop() error = %v", err)
}
})
for i := range 3 {
if !writer.TryEnqueue(i + 1) {
t.Fatalf("TryEnqueue(%d) = false, want true", i+1)
}
}
deadline := time.Now().Add(time.Second)
for {
mu.Lock()
ready := len(batch) == 3
mu.Unlock()
if ready || time.Now().After(deadline) {
break
}
time.Sleep(5 * time.Millisecond)
}
mu.Lock()
got := batch
mu.Unlock()
want := []int{1, 2, 3}
if diff := cmp.Diff(want, got); diff != "" {
t.Fatalf("interval flush at min batch size mismatch (-want +got):\n%s", diff)
}
}
func TestWriterForcesFlushAfterMaxFlushWait(t *testing.T) {
t.Parallel()
var (
mu sync.Mutex
batch []int
)
cfg := DefaultConfig()
cfg.MaxBatchSize = 100
cfg.MinBatchSize = 50
cfg.FlushInterval = 20 * time.Millisecond
cfg.MaxFlushWait = 80 * time.Millisecond
writer, err := New[int](cfg, func(_ context.Context, items []int) error {
mu.Lock()
defer mu.Unlock()
batch = append([]int(nil), items...)
return nil
})
if err != nil {
t.Fatalf("New() error = %v", err)
}
writer.Start(context.Background())
t.Cleanup(func() {
stopCtx, cancel := context.WithTimeout(context.Background(), time.Second)
defer cancel()
if err := writer.Stop(stopCtx); err != nil {
t.Fatalf("Stop() error = %v", err)
}
})
if !writer.TryEnqueue(1) {
t.Fatal("TryEnqueue(1) = false, want true")
}
deadline := time.Now().Add(time.Second)
for {
mu.Lock()
ready := len(batch) == 1
mu.Unlock()
if ready || time.Now().After(deadline) {
break
}
time.Sleep(5 * time.Millisecond)
}
mu.Lock()
got := batch
mu.Unlock()
if diff := cmp.Diff([]int{1}, got); diff != "" {
t.Fatalf("max flush wait mismatch (-want +got):\n%s", diff)
}
}
func TestWriterStatsTracksDrops(t *testing.T) {
t.Parallel()
cfg := DefaultConfig()
cfg.Name = "test-drops"
cfg.QueueSize = 1
cfg.MaxBatchSize = 10
cfg.FlushInterval = time.Hour
writer, err := New[int](cfg, func(context.Context, []int) error { return nil })
if err != nil {
t.Fatalf("New() error = %v", err)
}
writer.Start(context.Background())
t.Cleanup(func() {
stopCtx, cancel := context.WithTimeout(context.Background(), time.Second)
defer cancel()
_ = writer.Stop(stopCtx)
})
if !writer.TryEnqueue(1) {
t.Fatal("TryEnqueue(1) = false, want true")
}
if writer.TryEnqueue(2) {
t.Fatal("TryEnqueue(2) = true, want false")
}
stats := writer.Stats()
if stats.Drops != 1 {
t.Fatalf("Stats().Drops = %d, want 1", stats.Drops)
}
if stats.Cap != 1 {
t.Fatalf("Stats().Cap = %d, want 1", stats.Cap)
}
if !stats.Running {
t.Fatal("Stats().Running = false, want true")
}
}
-155
View File
@@ -1,155 +0,0 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package db 提供数据库连接与基础设施
package db
import (
"context"
"fmt"
"log"
"net/url"
"strconv"
"strings"
"time"
"github.com/ClickHouse/clickhouse-go/v2"
"github.com/ClickHouse/clickhouse-go/v2/lib/driver"
"github.com/Rain-kl/Wavelet/pkg/config"
"go.opentelemetry.io/otel/attribute"
clickhouseDriver "gorm.io/driver/clickhouse"
"gorm.io/gorm"
"gorm.io/plugin/opentelemetry/tracing"
)
const (
clickhouseMaxExecTime = 60 // ClickHouse 最大执行时间(秒)
clickhouseReadTimeoutFactor = 2 // ReadTimeout 为 DialTimeout 的倍数
)
var (
// ChConn ClickHouse 原生连接实例,用于批量写入
ChConn driver.Conn
chDB *gorm.DB
)
func init() {
if !config.Config.ClickHouse.Enabled {
return
}
cfg := config.Config.ClickHouse
if cfg.Database == "" {
log.Fatalf("[ClickHouse] database name is required (expected: wavelet)\n")
}
opts := buildClickHouseOptions()
var err error
ChConn, err = clickhouse.Open(opts)
if err != nil {
log.Fatalf("[ClickHouse] init connection failed: %v\n", err)
}
if err = ChConn.Ping(context.Background()); err != nil {
log.Fatalf("[ClickHouse] ping failed: %v\n", err)
}
chDB, err = gorm.Open(clickhouseDriver.New(clickhouseDriver.Config{
DSN: buildClickHouseDSN(),
}), &gorm.Config{
SkipDefaultTransaction: true,
})
if err != nil {
log.Fatalf("[ClickHouse] init gorm connection failed: %v\n", err)
}
if err = chDB.Use(
tracing.NewPlugin(
tracing.WithoutMetrics(),
tracing.WithAttributes(
attribute.String("db.instance", cfg.Database),
attribute.String("db.system", "ClickHouse"),
),
),
); err != nil {
log.Fatalf("[ClickHouse] init trace failed: %v\n", err)
}
sqlDB, err := chDB.DB()
if err != nil {
log.Fatalf("[ClickHouse] load sql db failed: %v\n", err)
}
sqlDB.SetMaxIdleConns(cfg.MaxIdleConn)
sqlDB.SetMaxOpenConns(cfg.MaxOpenConn)
sqlDB.SetConnMaxLifetime(time.Duration(cfg.ConnMaxLifetime) * time.Second)
log.Println("[ClickHouse] connection established successfully")
}
func buildClickHouseOptions() *clickhouse.Options {
cfg := config.Config.ClickHouse
return &clickhouse.Options{
Addr: cfg.Hosts,
Auth: clickhouse.Auth{
Database: cfg.Database,
Username: cfg.Username,
Password: cfg.Password,
},
Settings: clickhouse.Settings{
"max_execution_time": clickhouseMaxExecTime,
},
Compression: &clickhouse.Compression{
Method: clickhouse.CompressionLZ4,
},
DialTimeout: time.Duration(cfg.DialTimeout) * time.Second,
MaxOpenConns: cfg.MaxOpenConn,
MaxIdleConns: cfg.MaxIdleConn,
ConnMaxLifetime: time.Duration(cfg.ConnMaxLifetime) * time.Second,
ReadTimeout: time.Duration(cfg.DialTimeout*clickhouseReadTimeoutFactor) * time.Second,
BlockBufferSize: cfg.BlockBufferSize,
}
}
func buildClickHouseDSN() string {
cfg := config.Config.ClickHouse
chURL := &url.URL{
Scheme: "clickhouse",
Host: strings.Join(cfg.Hosts, ","),
Path: "/" + cfg.Database,
}
if cfg.Username != "" || cfg.Password != "" {
chURL.User = url.UserPassword(cfg.Username, cfg.Password)
}
query := chURL.Query()
query.Set("dial_timeout", fmt.Sprintf("%ds", cfg.DialTimeout))
query.Set("read_timeout", fmt.Sprintf("%ds", cfg.DialTimeout*clickhouseReadTimeoutFactor))
query.Set("max_execution_time", strconv.Itoa(clickhouseMaxExecTime))
chURL.RawQuery = query.Encode()
return chURL.String()
}
// ChDB returns a context-aware GORM ClickHouse instance.
func ChDB(ctx context.Context) *gorm.DB {
if chDB == nil {
return nil
}
return chDB.WithContext(ctx)
}
// SetChDBForTest sets the package-level ClickHouse GORM instance for testing.
func SetChDBForTest(d *gorm.DB) {
chDB = d
}
// SetChConnForTest sets the package-level native ClickHouse connection for testing.
func SetChConnForTest(c driver.Conn) {
ChConn = c
}
-12
View File
@@ -1,12 +0,0 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package db
const (
errRedisHashSetFailed = "failed to set redis hash: %w"
errRedisHashDeleteFailed = "failed to delete redis hash field: %w"
errUnmarshalDataFailed = "failed to unmarshal data: %w"
errMarshalDataFailed = "failed to marshal data: %w"
errRedisKeySetFailed = "failed to set redis key: %w"
)
-99
View File
@@ -1,99 +0,0 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package logstore
import (
"context"
"errors"
"fmt"
"strconv"
"time"
"github.com/Rain-kl/Wavelet/pkg/logger"
)
const (
defaultLogRetentionDays = 30
partitionLeadMonths = 2
userAccessLogTable = "w_user_access_logs"
)
// CleanupSummary 汇总本次清理结果。
type CleanupSummary struct {
ActiveDatabase string `json:"active_database"`
RetentionDays int `json:"retention_days"`
Deleted int64 `json:"deleted"`
}
// CleanupExpired 按当前日志库保留天数删除过期用户访问日志,并预建 PG 分区。
func CleanupExpired(ctx context.Context) (CleanupSummary, error) {
active, err := ActiveDatabase(ctx)
if err != nil {
return CleanupSummary{}, err
}
days := retentionDaysForDatabase(ctx, active)
summary := CleanupSummary{ActiveDatabase: active, RetentionDays: days}
store, err := Active(ctx)
if err != nil {
return summary, err
}
now := time.Now().UTC()
if err := store.UserAccessLogs.EnsurePartitions(ctx, now, now.AddDate(0, partitionLeadMonths, 0)); err != nil {
logger.WarnF(ctx, "logstore: ensure partitions during cleanup failed: %v", err)
}
cutoff := now.AddDate(0, 0, -days)
// 先 DROP 完全过期的整月分区,再对边界月逐行 DeleteBefore。
if err := store.UserAccessLogs.DropExpiredPartitions(ctx, cutoff); err != nil {
return summary, fmt.Errorf("drop expired partitions: %w", err)
}
deleted, err := store.UserAccessLogs.DeleteBefore(ctx, cutoff)
if err != nil {
return summary, fmt.Errorf("delete expired user access logs: %w", err)
}
summary.Deleted = deleted
if err := store.UserAccessLogs.DropEmptyPartitions(ctx, now); err != nil {
logger.WarnF(ctx, "drop empty log partitions failed: %v", err)
}
return summary, nil
}
func retentionDaysForDatabase(ctx context.Context, dbName string) int {
key := "log_retention_days_postgres"
switch dbName {
case dbNameSQLite:
key = "log_retention_days_sqlite"
case dbNameClickHouse:
key = "log_retention_days_clickhouse"
}
v, err := getConfig(ctx, key)
if err != nil {
if !errors.Is(err, errConfigReaderNotWired) {
logger.ErrorF(ctx, "读取日志保留天数配置失败(key=%s),回退默认 %d 天: %v", key, defaultLogRetentionDays, err)
}
return defaultLogRetentionDays
}
days, perr := strconv.Atoi(v)
if perr != nil || days <= 0 {
logger.ErrorF(ctx, "日志保留天数配置非法(key=%s, value=%q),回退默认 %d 天", key, v, defaultLogRetentionDays)
return defaultLogRetentionDays
}
return days
}
func partitionStatementsRange(from, to time.Time) []string {
var out []string
start := time.Date(from.Year(), from.Month(), 1, 0, 0, 0, 0, time.UTC)
end := time.Date(to.Year(), to.Month(), 1, 0, 0, 0, 0, time.UTC).AddDate(0, 1, 0)
for ; start.Before(end); start = start.AddDate(0, 1, 0) {
monthEnd := start.AddDate(0, 1, 0)
suffix := start.Format("200601")
fromDay := start.Format("2006-01-02")
toDay := monthEnd.Format("2006-01-02")
out = append(out, fmt.Sprintf(
"CREATE TABLE IF NOT EXISTS %s_%s PARTITION OF %s FOR VALUES FROM ('%s') TO ('%s')",
userAccessLogTable, suffix, userAccessLogTable, fromDay, toDay))
}
return out
}
-153
View File
@@ -1,153 +0,0 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package logstore
import (
"context"
"fmt"
"time"
"github.com/ClickHouse/clickhouse-go/v2/lib/driver"
db "github.com/Rain-kl/Wavelet/pkg/persistence"
"github.com/Rain-kl/Wavelet/pkg/persistence/analytics"
)
type clickhouseUserAccessLogStore struct {
skipFreeze bool
}
func newClickHouseUserAccessLogStore() *clickhouseUserAccessLogStore {
return &clickhouseUserAccessLogStore{}
}
var (
_ UserAccessLogStore = (*clickhouseUserAccessLogStore)(nil)
_ StatusStore = (*clickhouseUserAccessLogStore)(nil)
)
func (s *clickhouseUserAccessLogStore) ActiveDatabase(_ context.Context) (string, error) {
return dbNameClickHouse, nil
}
func (s *clickhouseUserAccessLogStore) ensureWritable(ctx context.Context) error {
if !s.skipFreeze && Migrating(ctx) {
return ErrMigrating
}
return nil
}
func (s *clickhouseUserAccessLogStore) BatchInsert(ctx context.Context, logs []analytics.UserAccessLog) error {
if len(logs) == 0 {
return nil
}
if err := s.ensureWritable(ctx); err != nil {
return err
}
return analytics.BatchInsert(ctx, logs)
}
func (s *clickhouseUserAccessLogStore) DeleteAll(ctx context.Context) (int64, error) {
if err := s.ensureWritable(ctx); err != nil {
return 0, err
}
return analytics.DeleteAllUserAccessLogs(ctx)
}
func (s *clickhouseUserAccessLogStore) DeleteBefore(ctx context.Context, cutoff time.Time) (int64, error) {
if err := s.ensureWritable(ctx); err != nil {
return 0, err
}
return analytics.DeleteUserAccessLogsBefore(ctx, cutoff)
}
func (s *clickhouseUserAccessLogStore) Count(ctx context.Context, filter analytics.AccessLogFilter) (uint64, error) {
return analytics.CountAccessLogs(ctx, filter)
}
func (s *clickhouseUserAccessLogStore) List(ctx context.Context, filter analytics.AccessLogFilter, page, pageSize int) ([]analytics.UserAccessLog, uint64, error) {
return analytics.ListAccessLogs(ctx, filter, page, pageSize)
}
func (s *clickhouseUserAccessLogStore) GetDailyTrend(ctx context.Context, days int) ([]analytics.DailyTrend, error) {
return analytics.GetDailyTrend(ctx, days)
}
func (s *clickhouseUserAccessLogStore) GetBrowserDistribution(ctx context.Context, startTime time.Time) ([]analytics.BrowserShare, error) {
return analytics.GetBrowserDistribution(ctx, startTime)
}
func (s *clickhouseUserAccessLogStore) GetTopActiveUsers(ctx context.Context, startTime time.Time, limit int) ([]analytics.TopUser, error) {
return analytics.GetTopActiveUsers(ctx, startTime, limit)
}
func (s *clickhouseUserAccessLogStore) EnsurePartitions(_ context.Context, _, _ time.Time) error {
return nil
}
func (s *clickhouseUserAccessLogStore) DropEmptyPartitions(_ context.Context, _ time.Time) error {
return nil
}
func (s *clickhouseUserAccessLogStore) DropExpiredPartitions(_ context.Context, _ time.Time) error {
return nil
}
func (s *clickhouseUserAccessLogStore) MigrationRange(ctx context.Context) (time.Time, time.Time, error) {
if db.ChConn == nil {
return time.Time{}, time.Time{}, fmt.Errorf("clickhouse connection is not initialized")
}
table := analytics.UserAccessLog{}.TableName()
var minTime, maxTime *time.Time
if err := db.ChConn.QueryRow(ctx, "SELECT min(created_at), max(created_at) FROM "+table).Scan(&minTime, &maxTime); err != nil {
return time.Time{}, time.Time{}, fmt.Errorf("query migration range %s: %w", table, err)
}
if minTime == nil || maxTime == nil {
return time.Time{}, time.Time{}, nil
}
return minTime.UTC(), maxTime.UTC(), nil
}
func (s *clickhouseUserAccessLogStore) ListForMigration(ctx context.Context, afterID uint64, limit int) ([]analytics.UserAccessLog, error) {
if db.ChConn == nil {
return nil, fmt.Errorf("clickhouse connection is not initialized")
}
if limit <= 0 {
limit = migrationPageSize
}
table := analytics.UserAccessLog{}.TableName()
columns := analytics.UserAccessLog{}.InsertColumns()
rows, err := db.ChConn.Query(ctx, fmt.Sprintf(
"SELECT %s FROM %s WHERE id > ? ORDER BY id ASC LIMIT ?",
columns, table,
), afterID, limit)
if err != nil {
return nil, fmt.Errorf("list user access logs for migration: %w", err)
}
defer func() { _ = rows.Close() }()
return scanUserAccessLogs(rows)
}
func scanUserAccessLogs(rows driver.Rows) ([]analytics.UserAccessLog, error) {
var result []analytics.UserAccessLog
for rows.Next() {
var item analytics.UserAccessLog
if err := rows.Scan(
&item.ID,
&item.UserID,
&item.Path,
&item.Method,
&item.IP,
&item.UserAgent,
&item.Headers,
&item.Status,
&item.Latency,
&item.CreatedAt,
); err != nil {
return nil, fmt.Errorf("scan user access log row: %w", err)
}
item.CreatedAt = item.CreatedAt.UTC()
result = append(result, item)
}
return result, nil
}
-327
View File
@@ -1,327 +0,0 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package logstore
import (
"context"
"errors"
"fmt"
"sort"
"strings"
"time"
"github.com/Rain-kl/Wavelet/pkg/persistence/analytics"
"github.com/Rain-kl/Wavelet/pkg/persistence/idgen"
"gorm.io/gorm"
)
const (
insertBatchSize = 500
migrationPageSize = 100
defaultPageSize = 20
defaultTopN = 10
topUserAgents = 100
dayDuration = 24 * time.Hour
)
type gormLogStore struct {
db *gorm.DB
skipFreeze bool
}
func newGormStore(db *gorm.DB) *gormLogStore { return &gormLogStore{db: db} }
type userAccessLogGormStore struct {
*gormLogStore
}
func newUserAccessLogGormStore(db *gorm.DB) *userAccessLogGormStore {
return &userAccessLogGormStore{gormLogStore: newGormStore(db)}
}
var (
_ UserAccessLogStore = (*userAccessLogGormStore)(nil)
_ StatusStore = (*userAccessLogGormStore)(nil)
)
func (s *gormLogStore) ActiveDatabase(_ context.Context) (string, error) {
if isPostgresDialect(s.db) {
return dbNamePostgres, nil
}
return dbNameSQLite, nil
}
func (s *gormLogStore) ensureWritable(ctx context.Context) error {
if !s.skipFreeze && Migrating(ctx) {
return ErrMigrating
}
return nil
}
func (s *userAccessLogGormStore) BatchInsert(ctx context.Context, logs []analytics.UserAccessLog) error {
if len(logs) == 0 {
return nil
}
if err := s.ensureWritable(ctx); err != nil {
return err
}
for i := range logs {
if logs[i].ID == 0 {
logs[i].ID = idgen.NextUint64ID()
}
}
return s.db.WithContext(ctx).CreateInBatches(logs, insertBatchSize).Error
}
func (s *userAccessLogGormStore) DeleteAll(ctx context.Context) (int64, error) {
if err := s.ensureWritable(ctx); err != nil {
return 0, err
}
res := s.db.WithContext(ctx).Where("1 = 1").Delete(&analytics.UserAccessLog{})
return res.RowsAffected, res.Error
}
func (s *userAccessLogGormStore) DeleteBefore(ctx context.Context, cutoff time.Time) (int64, error) {
if err := s.ensureWritable(ctx); err != nil {
return 0, err
}
res := s.db.WithContext(ctx).Where("created_at < ?", cutoff).Delete(&analytics.UserAccessLog{})
if res.Error != nil && isMissingRelation(res.Error) {
return 0, nil
}
return res.RowsAffected, res.Error
}
func (s *userAccessLogGormStore) ListForMigration(ctx context.Context, afterID uint64, limit int) ([]analytics.UserAccessLog, error) {
var rows []analytics.UserAccessLog
q := s.db.WithContext(ctx).Model(&analytics.UserAccessLog{}).
Where("id > ?", afterID).
Order("id ASC").
Limit(limitOr(limit, migrationPageSize))
if err := q.Find(&rows).Error; err != nil {
return nil, err
}
return rows, nil
}
func (s *userAccessLogGormStore) MigrationRange(ctx context.Context) (time.Time, time.Time, error) {
return gormMigrationRange(ctx, s.db, "created_at", analytics.UserAccessLog{}, func(v *analytics.UserAccessLog) time.Time {
return v.CreatedAt
})
}
func (s *userAccessLogGormStore) Count(ctx context.Context, filter analytics.AccessLogFilter) (uint64, error) {
where, args, ok := buildUserAccessLogWhere(filter)
if !ok {
return 0, nil
}
var total int64
if err := s.db.WithContext(ctx).Model(&analytics.UserAccessLog{}).Where(where, args...).Count(&total).Error; err != nil {
return 0, err
}
return countToUint64(total), nil
}
func (s *userAccessLogGormStore) List(ctx context.Context, filter analytics.AccessLogFilter, page, pageSize int) ([]analytics.UserAccessLog, uint64, error) {
where, args, ok := buildUserAccessLogWhere(filter)
if !ok {
return []analytics.UserAccessLog{}, 0, nil
}
var total int64
if err := s.db.WithContext(ctx).Model(&analytics.UserAccessLog{}).Where(where, args...).Count(&total).Error; err != nil {
return nil, 0, err
}
if total == 0 {
return []analytics.UserAccessLog{}, 0, nil
}
var rows []analytics.UserAccessLog
q := s.db.WithContext(ctx).Where(where, args...).Order("created_at DESC, id DESC")
if err := q.Limit(limitOr(pageSize, defaultPageSize)).Offset(offsetOf(page, pageSize)).Find(&rows).Error; err != nil {
return nil, 0, err
}
return rows, countToUint64(total), nil
}
func buildUserAccessLogWhere(filter analytics.AccessLogFilter) (string, []any, bool) {
if filter.UserIDs != nil && len(filter.UserIDs) == 0 {
return "", nil, false
}
var parts []string
var args []any
if filter.UserIDs != nil {
parts = append(parts, "user_id IN ?")
args = append(args, filter.UserIDs)
}
if trimmed := strings.TrimSpace(filter.Path); trimmed != "" {
parts = append(parts, "path LIKE ?")
args = append(args, "%"+trimmed+"%")
}
if filter.StartTime != nil {
parts = append(parts, "created_at >= ?")
args = append(args, *filter.StartTime)
}
if filter.EndTime != nil {
parts = append(parts, "created_at <= ?")
args = append(args, *filter.EndTime)
}
if len(parts) == 0 {
return "1 = 1", args, true
}
return strings.Join(parts, " AND "), args, true
}
func (s *userAccessLogGormStore) GetDailyTrend(ctx context.Context, days int) ([]analytics.DailyTrend, error) {
if days <= 0 {
days = 7
}
start := time.Now().AddDate(0, 0, -(days - 1)).Truncate(dayDuration)
type row struct {
Date string
Cnt uint64
}
var rows []row
err := s.db.WithContext(ctx).Model(&analytics.UserAccessLog{}).
Select(dailyTrendDateSQL(s.db)+" AS date, COUNT(*) AS cnt").
Where("created_at >= ?", start).
Group("date").Order("date ASC").Scan(&rows).Error
if err != nil {
return nil, err
}
counts := make(map[string]uint64, len(rows))
for _, r := range rows {
counts[r.Date] = r.Cnt
}
out := make([]analytics.DailyTrend, 0, days)
for i := 0; i < days; i++ {
d := start.AddDate(0, 0, i).Format("2006-01-02")
out = append(out, analytics.DailyTrend{Date: d, Count: counts[d]})
}
return out, nil
}
func (s *userAccessLogGormStore) GetBrowserDistribution(ctx context.Context, startTime time.Time) ([]analytics.BrowserShare, error) {
type row struct {
UserAgent string
Cnt uint64
}
var rows []row
err := s.db.WithContext(ctx).Model(&analytics.UserAccessLog{}).
Select("user_agent, COUNT(*) AS cnt").
Where("created_at >= ?", startTime).
Group("user_agent").Order("cnt DESC").Limit(topUserAgents).Scan(&rows).Error
if err != nil {
return nil, err
}
counts := make(map[string]uint64)
for _, r := range rows {
counts[analytics.ParseBrowserName(r.UserAgent)] += r.Cnt
}
out := make([]analytics.BrowserShare, 0, len(counts))
for label, count := range counts {
out = append(out, analytics.BrowserShare{Browser: label, Count: count})
}
sort.Slice(out, func(i, j int) bool { return out[i].Count > out[j].Count })
return out, nil
}
func (s *userAccessLogGormStore) GetTopActiveUsers(ctx context.Context, startTime time.Time, limit int) ([]analytics.TopUser, error) {
type row struct {
UserID uint64
Cnt uint64
}
var rows []row
err := s.db.WithContext(ctx).Model(&analytics.UserAccessLog{}).
Select("user_id, COUNT(*) AS cnt").
Where("user_id <> 0 AND created_at >= ?", startTime).
Group("user_id").Order("cnt DESC").Limit(limitOr(limit, defaultTopN)).Scan(&rows).Error
if err != nil {
return nil, err
}
out := make([]analytics.TopUser, len(rows))
for i, r := range rows {
out[i] = analytics.TopUser{UserID: r.UserID, Count: r.Cnt}
}
return out, nil
}
func (s *userAccessLogGormStore) EnsurePartitions(ctx context.Context, from, to time.Time) error {
if !isPostgresDialect(s.db) {
return nil
}
for _, sql := range partitionStatementsRange(from, to) {
if err := s.db.WithContext(ctx).Exec(sql).Error; err != nil {
return fmt.Errorf("ensure partition: %w", err)
}
}
return nil
}
func gormMigrationRange[T any](
ctx context.Context,
gdb *gorm.DB,
column string,
model T,
timeOf func(*T) time.Time,
) (time.Time, time.Time, error) {
var first, last T
found := false
for _, order := range []string{"ASC", "DESC"} {
out := &first
if order == "DESC" {
out = &last
}
res := gdb.WithContext(ctx).Model(model).Order(column + " " + order).Limit(1).Take(out)
if res.Error != nil && !errors.Is(res.Error, gorm.ErrRecordNotFound) {
return time.Time{}, time.Time{}, fmt.Errorf("query migration range %s: %w", column, res.Error)
}
if res.Error == nil {
found = true
}
}
if !found {
return time.Time{}, time.Time{}, nil
}
return timeOf(&first).UTC(), timeOf(&last).UTC(), nil
}
func limitOr(v, def int) int {
if v <= 0 {
return def
}
return v
}
func offsetOf(page, pageSize int) int {
if page < 1 {
page = 1
}
return (page - 1) * limitOr(pageSize, defaultPageSize)
}
func countToUint64(v int64) uint64 {
if v < 0 {
return 0
}
return uint64(v)
}
func isPostgresDialect(db *gorm.DB) bool {
return db != nil && db.Dialector != nil && db.Name() == "postgres"
}
func dailyTrendDateSQL(db *gorm.DB) string {
if isPostgresDialect(db) {
return "to_char(created_at, 'YYYY-MM-DD')"
}
return "strftime('%Y-%m-%d', created_at)"
}
func isMissingRelation(err error) bool {
if err == nil {
return false
}
msg := strings.ToLower(err.Error())
return strings.Contains(msg, "no such table") || strings.Contains(msg, "does not exist")
}
-60
View File
@@ -1,60 +0,0 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package logstore
import (
"context"
"testing"
"time"
"github.com/Rain-kl/Wavelet/pkg/persistence/analytics"
"github.com/glebarez/sqlite"
"github.com/stretchr/testify/require"
"gorm.io/gorm"
)
func newTestUserAccessStore(t *testing.T) *userAccessLogGormStore {
t.Helper()
gdb, err := gorm.Open(sqlite.Open("file:logstore-"+t.Name()+"?mode=memory&cache=shared"), &gorm.Config{})
require.NoError(t, err)
require.NoError(t, gdb.AutoMigrate(&analytics.UserAccessLog{}))
return newUserAccessLogGormStore(gdb)
}
func TestGormUserAccessLogCountList(t *testing.T) {
ua := newTestUserAccessStore(t)
ctx := context.Background()
now := time.Now().UTC().Truncate(time.Second)
require.NoError(t, ua.BatchInsert(ctx, []analytics.UserAccessLog{
{UserID: 10, Path: "/api/v1/users", Method: "GET", Status: 200, CreatedAt: now},
{UserID: 20, Path: "/api/v1/admin", Method: "GET", Status: 200, CreatedAt: now},
{UserID: 10, Path: "/api/v1/other", Method: "POST", Status: 201, CreatedAt: now},
}))
count, err := ua.Count(ctx, analytics.AccessLogFilter{UserIDs: []uint64{10}, Path: "users"})
require.NoError(t, err)
require.Equal(t, uint64(1), count)
rows, total, err := ua.List(ctx, analytics.AccessLogFilter{UserIDs: []uint64{10}, Path: "users"}, 1, 10)
require.NoError(t, err)
require.Equal(t, uint64(1), total)
require.Len(t, rows, 1)
require.Equal(t, "/api/v1/users", rows[0].Path)
require.NotZero(t, rows[0].ID)
}
func TestGormUserAccessLogFreeze(t *testing.T) {
ua := newTestUserAccessStore(t)
SetConfigReader(func(_ context.Context, key string) (string, error) {
if key == logMigrationKey {
return "migrating", nil
}
return "", nil
})
t.Cleanup(ResetForTest)
err := ua.BatchInsert(context.Background(), []analytics.UserAccessLog{{UserID: 1, CreatedAt: time.Now()}})
require.ErrorIs(t, err, ErrMigrating)
}
-45
View File
@@ -1,45 +0,0 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package logstore
import (
"os/exec"
"strings"
"testing"
)
// apps 禁止直连 analytics 做日志读写;查询过滤器请用 logstore.AccessLogFilter。
var forbiddenImports = []string{
"github.com/Rain-kl/Wavelet/pkg/persistence/analytics",
}
// logstore 的 CH 实现按设计委托 analyticsrepo。
var allowedAnalyticsDelegation = map[string]bool{
"github.com/Rain-kl/Wavelet/pkg/persistence/logstore": true,
}
func TestAppsMustNotImportLogBackendDirectly(t *testing.T) {
t.Chdir("../../..")
out, err := exec.Command("go", "list", "-test", "-f", `{{.ImportPath}} {{join .Imports " "}}`, "./plugins/domain/...").Output()
if err != nil {
t.Fatalf("go list: %v", err)
}
for _, line := range strings.Split(string(out), "\n") {
fields := strings.Fields(line)
if len(fields) == 0 {
continue
}
pkg := fields[0]
if !strings.HasPrefix(pkg, "github.com/Rain-kl/Wavelet/plugins/domain") {
continue
}
for _, imp := range fields[1:] {
for _, forbidden := range forbiddenImports {
if imp == forbidden && !allowedAnalyticsDelegation[pkg] {
t.Errorf("%s must not import forbidden log backend %s", pkg, forbidden)
}
}
}
}
}
-63
View File
@@ -1,63 +0,0 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package logstore abstracts user access-log storage across ClickHouse, PostgreSQL and SQLite.
package logstore
import (
"context"
"errors"
"time"
"github.com/Rain-kl/Wavelet/pkg/persistence/analytics"
)
// ErrMigrating 表示日志数据库正在迁移,当前禁止写入。
var ErrMigrating = errors.New("log database is migrating, writes are disabled")
// UserAccessLogStore 用户访问日志(w_user_access_logs)。
type UserAccessLogStore interface {
BatchInsert(ctx context.Context, logs []analytics.UserAccessLog) error
DeleteAll(ctx context.Context) (int64, error)
DeleteBefore(ctx context.Context, cutoff time.Time) (int64, error)
Count(ctx context.Context, filter analytics.AccessLogFilter) (uint64, error)
List(ctx context.Context, filter analytics.AccessLogFilter, page, pageSize int) ([]analytics.UserAccessLog, uint64, error)
GetDailyTrend(ctx context.Context, days int) ([]analytics.DailyTrend, error)
GetBrowserDistribution(ctx context.Context, startTime time.Time) ([]analytics.BrowserShare, error)
GetTopActiveUsers(ctx context.Context, startTime time.Time, limit int) ([]analytics.TopUser, error)
ListForMigration(ctx context.Context, afterID uint64, limit int) ([]analytics.UserAccessLog, error)
MigrationRange(ctx context.Context) (from, to time.Time, err error)
EnsurePartitions(ctx context.Context, from, to time.Time) error
// DropEmptyPartitions 幂等清理 PG 空分区表:删除 before 月份之前、且无任何数据的按月分区;
// CH/SQLite 为 no-op。
DropEmptyPartitions(ctx context.Context, before time.Time) error
// DropExpiredPartitions 直接删除完全过期的 PG 整月分区(候选为月份早于 cutoff 月的分区,
// 删除前校验分区内无保留期内数据,避免时区偏移下误删;迁移冻结期间拒绝执行);CH/SQLite 为 no-op。
DropExpiredPartitions(ctx context.Context, cutoff time.Time) error
}
// UserAccessLog 用户访问日志实体类型别名
type UserAccessLog = analytics.UserAccessLog
// AccessLogFilter 访问日志查询过滤条件类型别名
type AccessLogFilter = analytics.AccessLogFilter
// DailyTrend 每日趋势数据类型别名
type DailyTrend = analytics.DailyTrend
// BrowserShare 浏览器分布数据类型别名
type BrowserShare = analytics.BrowserShare
// TopUser Top 用户活跃统计类型别名
type TopUser = analytics.TopUser
// StatusStore 日志库状态。
type StatusStore interface {
ActiveDatabase(ctx context.Context) (string, error)
}
// Store 当前生效日志库。
type Store struct {
UserAccessLogs UserAccessLogStore
Status StatusStore
}
@@ -1,115 +0,0 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package logstore
import (
"context"
"fmt"
"strings"
"time"
"gorm.io/gorm"
)
// listPartitionNames 列出 table 在当前 schema 下的全部直接分区表名(pg_inherits)。
func listPartitionNames(ctx context.Context, gdb *gorm.DB, table string) ([]string, error) {
var names []string
if err := gdb.WithContext(ctx).Raw(`
SELECT c.relname
FROM pg_inherits i
JOIN pg_class c ON c.oid = i.inhrelid
JOIN pg_class p ON p.oid = i.inhparent
JOIN pg_namespace n ON n.oid = p.relnamespace AND n.nspname = current_schema()
WHERE p.relname = ?`, table).Scan(&names).Error; err != nil {
return nil, fmt.Errorf("list partitions of %s: %w", table, err)
}
return names, nil
}
// partitionNameMonth 解析按月分区表名 <table>_YYYYMM 的所属月份;命名不匹配返回 (零值, false)。
func partitionNameMonth(table, name string) (time.Time, bool) {
suffix, ok := strings.CutPrefix(name, table+"_")
if !ok || len(suffix) != 6 {
return time.Time{}, false
}
m, err := time.Parse("200601", suffix)
if err != nil {
return time.Time{}, false
}
return m, true
}
// dropEligiblePartitionNames 返回 before 月份之前、命名合法的分区表名(是否为空由调用方校验)。
func dropEligiblePartitionNames(table string, names []string, before time.Time) []string {
beforeMonth := time.Date(before.Year(), before.Month(), 1, 0, 0, 0, 0, time.UTC)
out := make([]string, 0, len(names))
for _, name := range names {
month, ok := partitionNameMonth(table, name)
if !ok || !month.Before(beforeMonth) {
continue
}
out = append(out, name)
}
return out
}
// DropEmptyPartitions 幂等清理 PG 空分区表:仅删除 before 月份之前、且无任何数据的分区。
// 非 PG 方言为 no-op。
func (s *gormLogStore) DropEmptyPartitions(ctx context.Context, before time.Time) error {
if !isPostgresDialect(s.db) {
return nil
}
names, err := listPartitionNames(ctx, s.db, userAccessLogTable)
if err != nil {
return err
}
for _, name := range dropEligiblePartitionNames(userAccessLogTable, names, before) {
var one int
if err := s.db.WithContext(ctx).Raw("SELECT 1 FROM " + name + " LIMIT 1").Scan(&one).Error; err != nil {
return fmt.Errorf("check partition %s empty: %w", name, err)
}
if one == 1 {
continue
}
if err := s.db.WithContext(ctx).Exec("DROP TABLE IF EXISTS " + name).Error; err != nil {
return fmt.Errorf("drop empty partition %s: %w", name, err)
}
}
return nil
}
// DropExpiredPartitions 直接删除完全过期的 PG 整月分区(避免 retention 清理逐行 DELETE)。
// 候选 = 月份早于 cutoff 月(按 cutoff 的 UTC 时刻取月);删除前校验分区内不存在 created_at >= cutoff 的行。
// 迁移冻结期间返回 ErrMigrating。CH/SQLite 为 no-op。
func (s *gormLogStore) DropExpiredPartitions(ctx context.Context, cutoff time.Time) error {
if !isPostgresDialect(s.db) {
return nil
}
if err := s.ensureWritable(ctx); err != nil {
return err
}
names, err := listPartitionNames(ctx, s.db, userAccessLogTable)
if err != nil {
return err
}
cu := cutoff.UTC()
cutoffMonth := time.Date(cu.Year(), cu.Month(), 1, 0, 0, 0, 0, time.UTC)
for _, name := range names {
month, ok := partitionNameMonth(userAccessLogTable, name)
if !ok || !month.Before(cutoffMonth) {
continue
}
var hasRetained int
if err := s.db.WithContext(ctx).Raw("SELECT 1 FROM "+name+" WHERE created_at >= ? LIMIT 1", cu).Scan(&hasRetained).Error; err != nil {
return fmt.Errorf("check partition %s retained rows: %w", name, err)
}
if hasRetained == 1 {
continue
}
if err := s.db.WithContext(ctx).Exec("DROP TABLE IF EXISTS " + name).Error; err != nil {
return fmt.Errorf("drop expired partition %s: %w", name, err)
}
}
return nil
}
@@ -1,70 +0,0 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package logstore
import (
"context"
"testing"
"time"
"github.com/Rain-kl/Wavelet/pkg/persistence/analytics"
"github.com/stretchr/testify/require"
)
func TestPartitionNameMonth(t *testing.T) {
cases := []struct {
table string
name string
want string
}{
{"w_user_access_logs", "w_user_access_logs_202612", "2026-12"},
{"w_user_access_logs", "w_user_access_logs_202608", "2026-08"},
{"w_user_access_logs", "of_node_access_logs_202608", ""},
{"w_user_access_logs", "w_user_access_logs_20268", ""},
{"w_user_access_logs", "w_user_access_logs_202613", ""},
{"w_user_access_logs", "w_user_access_logs_default", ""},
}
for _, c := range cases {
got, ok := partitionNameMonth(c.table, c.name)
if c.want == "" {
if ok {
t.Fatalf("partitionNameMonth(%q, %q) ok = true, want false", c.table, c.name)
}
continue
}
if !ok || got.Format("2006-01") != c.want {
t.Fatalf("partitionNameMonth(%q, %q) = %v, want %s", c.table, c.name, got, c.want)
}
}
}
func TestDropEligiblePartitionNames(t *testing.T) {
before := time.Date(2026, 10, 15, 0, 0, 0, 0, time.UTC)
names := []string{
"w_user_access_logs_202608",
"w_user_access_logs_202609",
"w_user_access_logs_202610",
"w_user_access_logs_202611",
"w_user_access_logs_default",
}
got := dropEligiblePartitionNames(userAccessLogTable, names, before)
want := []string{"w_user_access_logs_202608", "w_user_access_logs_202609"}
require.Equal(t, want, got)
first := time.Date(2026, 10, 1, 0, 0, 0, 0, time.UTC)
require.Empty(t, dropEligiblePartitionNames(userAccessLogTable, []string{"w_user_access_logs_202610"}, first))
}
func TestDropPartitionHelpersSQLiteNoop(t *testing.T) {
ua := newTestUserAccessStore(t)
ctx := context.Background()
require.NoError(t, ua.BatchInsert(ctx, []analytics.UserAccessLog{
{UserID: 1, Path: "/x", CreatedAt: time.Now().UTC()},
}))
require.NoError(t, ua.DropExpiredPartitions(ctx, time.Now().AddDate(0, 0, -90)))
require.NoError(t, ua.DropEmptyPartitions(ctx, time.Now()))
count, err := ua.Count(ctx, AccessLogFilter{})
require.NoError(t, err)
require.Equal(t, uint64(1), count)
}
-189
View File
@@ -1,189 +0,0 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package logstore
import (
"context"
"errors"
"fmt"
"sync"
"time"
"github.com/Rain-kl/Wavelet/pkg/config"
"github.com/Rain-kl/Wavelet/pkg/logger"
db "github.com/Rain-kl/Wavelet/pkg/persistence"
)
const (
logDatabaseKey = "log_database"
logMigrationKey = "log_db_migration"
)
const (
dbNamePostgres = "postgres"
dbNameSQLite = "sqlite"
dbNameClickHouse = "clickhouse"
)
var errConfigReaderNotWired = errors.New("logstore: config reader not wired")
// ConfigReader 读取系统配置字符串值,由 bootstrap 注入(避免 logstore ↔ repository 循环依赖)。
type ConfigReader func(ctx context.Context, key string) (string, error)
const resolveCacheTTL = 1 * time.Second
var (
configReader ConfigReader
storeMu sync.RWMutex
active *Store
activeDB string
lastResolveDB string
lastResolveTime time.Time
)
// SetConfigReader 注入系统配置读取函数(bootstrap 调用,测试可注入内存实现)。
func SetConfigReader(fn ConfigReader) { configReader = fn }
func getConfig(ctx context.Context, key string) (string, error) {
if configReader == nil {
return "", errConfigReaderNotWired
}
return configReader(ctx, key)
}
// Active 返回当前生效的日志库 Store。
func Active(ctx context.Context) (*Store, error) {
current, err := resolveDatabase(ctx)
if err != nil {
return nil, err
}
storeMu.RLock()
if active != nil && activeDB == current {
s := active
storeMu.RUnlock()
return s, nil
}
storeMu.RUnlock()
storeMu.Lock()
defer storeMu.Unlock()
if active != nil && activeDB == current {
return active, nil
}
s, err := buildStore(ctx, current, false)
if err != nil {
return nil, err
}
active = s
activeDB = current
return s, nil
}
// Build 直接按目标构造 store(不经 Active 缓存)。
func Build(ctx context.Context, database string) (*Store, error) {
return buildStore(ctx, database, false)
}
// BuildForMigration 构造迁移目标 store,跳过冻结检查。
func BuildForMigration(ctx context.Context, database string) (*Store, error) {
return buildStore(ctx, database, true)
}
func buildStore(ctx context.Context, database string, skipFreeze bool) (*Store, error) {
switch database {
case dbNameClickHouse:
ual := newClickHouseUserAccessLogStore()
ual.skipFreeze = skipFreeze
return &Store{UserAccessLogs: ual, Status: ual}, nil
case dbNamePostgres, dbNameSQLite:
gdb := db.DB(ctx)
ual := newUserAccessLogGormStore(gdb)
ual.skipFreeze = skipFreeze
return &Store{UserAccessLogs: ual, Status: ual}, nil
default:
return nil, fmt.Errorf("unsupported log database: %s", database)
}
}
// Migrating 返回日志库是否处于迁移冻结状态。
func Migrating(ctx context.Context) bool {
v, err := getConfig(ctx, logMigrationKey)
if err != nil {
if !errors.Is(err, errConfigReaderNotWired) {
logger.ErrorF(ctx, "read log migration config failed: %v", err)
}
return false
}
return v == "migrating"
}
// Init 预热激活 store,并兜底预建当前月及未来分区。
func Init(ctx context.Context) {
s, err := Active(ctx)
if err != nil {
return
}
now := time.Now().UTC()
if err := s.UserAccessLogs.EnsurePartitions(ctx, now, now.AddDate(0, partitionLeadMonths, 0)); err != nil {
logger.WarnF(ctx, "logstore: ensure startup partitions failed: %v", err)
}
}
// InvalidateCache 清空日志库解析缓存。
func InvalidateCache() {
storeMu.Lock()
defer storeMu.Unlock()
lastResolveTime = time.Time{}
lastResolveDB = ""
}
// ResetForTest 清空缓存的激活 store 与 config reader。
func ResetForTest() {
storeMu.Lock()
active = nil
activeDB = ""
lastResolveDB = ""
lastResolveTime = time.Time{}
storeMu.Unlock()
configReader = nil
}
// ActiveDatabase 返回当前日志主库名。
func ActiveDatabase(ctx context.Context) (string, error) {
return resolveDatabase(ctx)
}
func resolveDatabase(ctx context.Context) (string, error) {
storeMu.RLock()
if active != nil && time.Since(lastResolveTime) < resolveCacheTTL {
name := lastResolveDB
storeMu.RUnlock()
return name, nil
}
storeMu.RUnlock()
v, err := getConfig(ctx, logDatabaseKey)
if err != nil && !errors.Is(err, errConfigReaderNotWired) {
return "", err
}
resolved := v
if resolved == "" {
resolved = dbNameSQLite
if config.Config.Database.Enabled {
resolved = dbNamePostgres
}
if config.Config.ClickHouse.Enabled {
resolved = dbNameClickHouse
}
}
storeMu.Lock()
lastResolveDB = resolved
lastResolveTime = time.Now()
storeMu.Unlock()
return resolved, nil
}
-217
View File
@@ -1,217 +0,0 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package db
import (
"context"
"log"
"net"
"net/url"
"strconv"
"time"
"github.com/Rain-kl/Wavelet/pkg/config"
"github.com/glebarez/sqlite"
"go.opentelemetry.io/otel/attribute"
"gorm.io/driver/postgres"
"gorm.io/gorm"
"gorm.io/plugin/dbresolver"
"gorm.io/plugin/opentelemetry/tracing"
)
var (
db *gorm.DB
)
func init() {
if !config.Config.Database.Enabled {
// PostgreSQL 禁用,使用 SQLite
initSQLite()
return
}
initPostgres()
}
// initSQLite 初始化 SQLite 数据库(PostgreSQL 禁用时的后备方案)
func initSQLite() {
sqlitePath := config.Config.Database.SQLitePath
if sqlitePath == "" {
sqlitePath = "./data/wavelet.db"
}
var err error
db, err = gorm.Open(sqlite.Open(sqlitePath), &gorm.Config{
DisableForeignKeyConstraintWhenMigrating: true,
Logger: &gormZapLogger{
logLevel: parseLogLevel(config.Config.Database.LogLevel),
slowThreshold: config.Config.Database.SlowThreshold,
ignoreRecordNotFoundError: config.Config.App.IsProduction(),
},
})
if err != nil {
log.Fatalf("[SQLite] init connection failed: %v\n", err)
}
// Trace 注入
if err = db.Use(
tracing.NewPlugin(
tracing.WithoutMetrics(),
tracing.WithAttributes(
attribute.String("db.instance", sqlitePath),
attribute.String("db.system", "SQLite"),
),
),
); err != nil {
log.Fatalf("[SQLite] init trace failed: %v\n", err)
}
log.Printf("[SQLite] initialized (path: %s)\n", sqlitePath)
}
// initPostgres 初始化 PostgreSQL 数据库
func initPostgres() {
var err error
dbConfig := config.Config.Database
// 构建主库 DSN 并连接
primaryDSN := buildDSN(dbConfig.Host, dbConfig.Port, dbConfig.Username, dbConfig.Password)
pgConfig := postgres.Config{
DSN: primaryDSN,
PreferSimpleProtocol: dbConfig.PreferSimpleProtocol,
}
db, err = gorm.Open(postgres.New(pgConfig), &gorm.Config{
DisableForeignKeyConstraintWhenMigrating: true,
Logger: &gormZapLogger{
logLevel: parseLogLevel(config.Config.Database.LogLevel),
slowThreshold: config.Config.Database.SlowThreshold,
ignoreRecordNotFoundError: config.Config.App.IsProduction(),
},
})
if err != nil {
log.Fatalf("[PostgreSQL] init connection failed: %v\n", err)
}
// Trace 注入
if err = db.Use(
tracing.NewPlugin(
tracing.WithoutMetrics(),
tracing.WithAttributes(
attribute.String("db.instance", dbConfig.Database),
attribute.String("db.ip", dbConfig.Host),
attribute.String("server.address", net.JoinHostPort(dbConfig.Host, strconv.Itoa(dbConfig.Port))),
attribute.String("db.system", "PostgreSQL"),
),
),
); err != nil {
log.Fatalf("[PostgreSQL] init trace failed: %v\n", err)
}
if len(dbConfig.Replicas) > 0 {
var replicaDialectors []gorm.Dialector
for _, replica := range dbConfig.Replicas {
username := replica.Username
if username == "" {
username = dbConfig.Username
}
password := replica.Password
if password == "" {
password = dbConfig.Password
}
replicaDSN := buildDSN(replica.Host, replica.Port, username, password)
replicaDialectors = append(replicaDialectors, postgres.New(postgres.Config{
DSN: replicaDSN,
PreferSimpleProtocol: dbConfig.PreferSimpleProtocol,
}))
}
resolver := dbresolver.Register(dbresolver.Config{
Replicas: replicaDialectors,
Policy: dbresolver.RandomPolicy{},
})
resolver.SetMaxIdleConns(dbConfig.MaxIdleConn).
SetMaxOpenConns(dbConfig.MaxOpenConn).
SetConnMaxLifetime(time.Duration(dbConfig.ConnMaxLifetime) * time.Second).
SetConnMaxIdleTime(time.Duration(dbConfig.ConnMaxIdleTime) * time.Second)
if err = db.Use(resolver); err != nil {
log.Fatalf("[PostgreSQL] init dbresolver failed: %v\n", err)
}
log.Printf("[PostgreSQL] initialized in Primary-Replica mode (%d replicas)\n", len(dbConfig.Replicas))
} else {
log.Println("[PostgreSQL] initialized in Standalone mode")
}
// 获取通用数据库对象设置连接池
sqlDB, err := db.DB()
if err != nil {
log.Fatalf("[PostgreSQL] load sql db failed: %v\n", err)
}
sqlDB.SetMaxIdleConns(dbConfig.MaxIdleConn)
sqlDB.SetMaxOpenConns(dbConfig.MaxOpenConn)
sqlDB.SetConnMaxLifetime(time.Duration(dbConfig.ConnMaxLifetime) * time.Second)
sqlDB.SetConnMaxIdleTime(time.Duration(dbConfig.ConnMaxIdleTime) * time.Second)
}
// buildDSN 构建 PostgreSQL DSN
func buildDSN(host string, port int, username, password string) string {
cfg := config.Config.Database
pqURL := &url.URL{
Scheme: "postgres",
Host: net.JoinHostPort(host, strconv.Itoa(port)),
Path: cfg.Database,
}
if username != "" {
pqURL.User = url.UserPassword(username, password)
}
query := pqURL.Query()
sslMode := cfg.SSLMode
if sslMode == "" {
sslMode = "disable"
}
query.Set("sslmode", sslMode)
if cfg.ApplicationName != "" {
query.Set("application_name", cfg.ApplicationName)
}
if cfg.SearchPath != "" {
query.Set("search_path", cfg.SearchPath)
}
if cfg.DefaultQueryExecMode != "" {
query.Set("default_query_exec_mode", cfg.DefaultQueryExecMode)
}
if cfg.StatementCacheCapacity > 0 {
query.Set("statement_cache_capacity", strconv.Itoa(cfg.StatementCacheCapacity))
}
rawQuery := query.Encode()
if cfg.TimeZone != "" {
if rawQuery != "" {
rawQuery += "&"
}
rawQuery += "TimeZone=" + cfg.TimeZone
}
pqURL.RawQuery = rawQuery
return pqURL.String()
}
// DB 返回带上下文追踪的 GORM 数据库实例
func DB(ctx context.Context) *gorm.DB {
if db == nil {
return nil
}
return db.WithContext(ctx)
}
// SetDB sets the package-level database instance for testing.
func SetDB(d *gorm.DB) {
db = d
}
-91
View File
@@ -1,91 +0,0 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package db
import (
"context"
"errors"
"fmt"
"strings"
"time"
"github.com/Rain-kl/Wavelet/pkg/logger"
"gorm.io/gorm"
gormLogger "gorm.io/gorm/logger"
)
// nanoToMilli 纳秒转毫秒的除数
const nanoToMilli = 1e6
type gormZapLogger struct {
logLevel gormLogger.LogLevel
ignoreRecordNotFoundError bool
slowThreshold time.Duration
}
func (l *gormZapLogger) LogMode(level gormLogger.LogLevel) gormLogger.Interface {
clone := *l
clone.logLevel = level
return &clone
}
func (l *gormZapLogger) Info(ctx context.Context, fmt string, args ...interface{}) {
if l.logLevel >= gormLogger.Info {
logger.InfoF(ctx, fmt, args...)
}
}
func (l *gormZapLogger) Warn(ctx context.Context, fmt string, args ...interface{}) {
if l.logLevel >= gormLogger.Warn {
logger.WarnF(ctx, fmt, args...)
}
}
func (l *gormZapLogger) Error(ctx context.Context, fmt string, args ...interface{}) {
if l.logLevel >= gormLogger.Error {
logger.ErrorF(ctx, fmt, args...)
}
}
func (l *gormZapLogger) Trace(ctx context.Context, begin time.Time, fc func() (sql string, rowsAffected int64), err error) {
elapsed := time.Since(begin)
switch {
case err != nil && l.logLevel >= gormLogger.Error && (!errors.Is(err, gorm.ErrRecordNotFound) || !l.ignoreRecordNotFoundError):
_, rows := fc()
logger.ErrorF(ctx, "database query failed: %s [%.3fms] [rows:%v]", err, float64(elapsed.Nanoseconds())/nanoToMilli, formatRows(rows))
case elapsed > l.slowThreshold && l.slowThreshold != 0 && l.logLevel >= gormLogger.Warn:
_, rows := fc()
slowLog := fmt.Sprintf("SLOW SQL >= %v", l.slowThreshold)
logger.WarnF(ctx, "%s [%.3fms] [rows:%v]", slowLog, float64(elapsed.Nanoseconds())/nanoToMilli, formatRows(rows))
case l.logLevel == gormLogger.Info:
sql, rows := fc()
logger.DebugF(ctx, "[%.3fms] [rows:%v] %s", float64(elapsed.Nanoseconds())/nanoToMilli, formatRows(rows), sql)
}
}
func formatRows(rows int64) interface{} {
if rows == -1 {
return "-"
}
return rows
}
func parseLogLevel(level string) gormLogger.LogLevel {
level = strings.ToLower(level)
switch level {
case "silent":
return gormLogger.Silent
case "error":
return gormLogger.Error
case "warn":
return gormLogger.Warn
case "info":
return gormLogger.Info
case "debug":
return gormLogger.Info
default:
return gormLogger.Info
}
}
-39
View File
@@ -1,39 +0,0 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package db
import (
"testing"
gormLogger "gorm.io/gorm/logger"
)
func TestParseLogLevel(t *testing.T) {
t.Parallel()
tests := []struct {
name string
configuredLevel string
want gormLogger.LogLevel
}{
{
name: "debug enables SQL trace processing",
configuredLevel: "debug",
want: gormLogger.Info,
},
{
name: "development preserves configured level",
configuredLevel: "warn",
want: gormLogger.Warn,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
if got := parseLogLevel(tt.configuredLevel); got != tt.want {
t.Fatalf("parseLogLevel() = %v, want %v", got, tt.want)
}
})
}
}
-198
View File
@@ -1,198 +0,0 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package db
import (
"context"
"encoding/json"
"fmt"
"log"
"strings"
"time"
"github.com/Rain-kl/Wavelet/pkg/config"
"github.com/redis/go-redis/extra/redisotel/v9"
"github.com/redis/go-redis/v9"
"github.com/redis/go-redis/v9/maintnotifications"
"go.opentelemetry.io/otel/attribute"
)
var (
// Redis 全局 Redis 客户端实例
Redis redis.UniversalClient
)
func init() {
cfg := config.Config.Redis
if !cfg.Enabled {
log.Println("[Redis] is disabled, skipping Redis initialization")
return
}
if cfg.ClusterMode {
// Cluster 模式
Redis = redis.NewClusterClient(&redis.ClusterOptions{
Addrs: cfg.Addrs,
Username: cfg.Username,
Password: cfg.Password,
PoolSize: cfg.PoolSize,
MinIdleConns: cfg.MinIdleConn,
DialTimeout: time.Duration(cfg.DialTimeout) * time.Second,
ReadTimeout: time.Duration(cfg.ReadTimeout) * time.Second,
WriteTimeout: time.Duration(cfg.WriteTimeout) * time.Second,
MaxRetries: cfg.MaxRetries,
PoolTimeout: time.Duration(cfg.PoolTimeout) * time.Second,
ConnMaxIdleTime: time.Duration(cfg.ConnMaxIdleTime) * time.Second,
MaintNotificationsConfig: redisMaintNotificationsConfig(cfg.MaintNotifications),
})
log.Println("[Redis] initialized in Cluster mode")
} else {
// Standalone 或 Sentinel 模式
options := &redis.UniversalOptions{
Addrs: cfg.Addrs,
MasterName: cfg.MasterName, // 非空时启用 Sentinel
Username: cfg.Username,
Password: cfg.Password,
DB: cfg.DB,
PoolSize: cfg.PoolSize,
MinIdleConns: cfg.MinIdleConn,
DialTimeout: time.Duration(cfg.DialTimeout) * time.Second,
ReadTimeout: time.Duration(cfg.ReadTimeout) * time.Second,
WriteTimeout: time.Duration(cfg.WriteTimeout) * time.Second,
MaxRetries: cfg.MaxRetries,
PoolTimeout: time.Duration(cfg.PoolTimeout) * time.Second,
ConnMaxIdleTime: time.Duration(cfg.ConnMaxIdleTime) * time.Second,
MaintNotificationsConfig: redisMaintNotificationsConfig(cfg.MaintNotifications),
}
if cfg.MasterName != "" {
client := redis.NewFailoverClient(options.Failover())
// FailoverOptions 暂不暴露该配置,在首次建连前写入客户端选项。
client.Options().MaintNotificationsConfig = redisMaintNotificationsConfig(cfg.MaintNotifications)
Redis = client
log.Println("[Redis] initialized in Sentinel mode")
} else {
Redis = redis.NewUniversalClient(options)
log.Println("[Redis] initialized in Standalone mode")
}
}
// OpenTelemetry 追踪(UniversalClient 兼容)
if err := redisotel.InstrumentTracing(
Redis,
redisotel.WithAttributes(
attribute.String("db.instance", fmt.Sprintf("%v", cfg.DB)),
attribute.String("db.ip", strings.Join(cfg.Addrs, ",")),
attribute.String("db.system", "Redis"),
),
); err != nil {
log.Fatalf("[Redis] failed to init trace: %v\n", err)
}
// 测试连接
_, err := Redis.Ping(context.Background()).Result()
if err != nil {
log.Fatalf("[Redis] failed to connect to redis: %v\n", err)
}
}
func redisMaintNotificationsConfig(enabled bool) *maintnotifications.Config {
mode := maintnotifications.ModeDisabled
if enabled {
mode = maintnotifications.ModeAuto
}
return &maintnotifications.Config{Mode: mode}
}
// PrefixedKey 返回带前缀的 Key
func PrefixedKey(key string) string {
prefix := config.Config.Redis.KeyPrefix
if prefix == "" {
return key
}
return prefix + key
}
// HSetJSON 将泛型数据序列化为 JSON 并设置到 Redis Hash
// ctx: 上下文
// hashKey: Redis Hash key
// fieldKey: Hash field key
// data: 要存储的数据(泛型)
func HSetJSON[T any](ctx context.Context, hashKey, fieldKey string, data T) error {
jsonData, err := json.Marshal(data)
if err != nil {
return err
}
if err := Redis.HSet(ctx, PrefixedKey(hashKey), fieldKey, jsonData).Err(); err != nil {
return fmt.Errorf(errRedisHashSetFailed, err)
}
return nil
}
// HDel removes one or more fields from a Redis Hash.
func HDel(ctx context.Context, hashKey string, fieldKeys ...string) error {
if Redis == nil || len(fieldKeys) == 0 {
return nil
}
if err := Redis.HDel(ctx, PrefixedKey(hashKey), fieldKeys...).Err(); err != nil {
return fmt.Errorf(errRedisHashDeleteFailed, err)
}
return nil
}
// HGetJSON 从 Redis Hash 获取数据并反序列化为泛型类型
// ctx: 上下文
// hashKey: Redis Hash key
// fieldKey: Hash field key
// data: 用于接收数据的指针(泛型)
func HGetJSON[T any](ctx context.Context, hashKey, fieldKey string, data *T) error {
val, err := Redis.HGet(ctx, PrefixedKey(hashKey), fieldKey).Result()
if err != nil {
return err
}
if err := json.Unmarshal([]byte(val), data); err != nil {
return fmt.Errorf(errUnmarshalDataFailed, err)
}
return nil
}
// GetJSON 从Redis获取数据并反序列化为泛型类型
// ctx: 上下文
// key: Redis key
// data: 用于接收数据的指针(泛型)
func GetJSON[T any](ctx context.Context, key string, data *T) error {
val, err := Redis.Get(ctx, PrefixedKey(key)).Bytes()
if err != nil {
return err
}
if err := json.Unmarshal(val, data); err != nil {
return fmt.Errorf(errUnmarshalDataFailed, err)
}
return nil
}
// SetJSON 将泛型数据序列化为JSON并设置到Redis
// ctx: 上下文
// key: Redis key
// data: 要存储的数据(泛型)
// expiration: 过期时间
func SetJSON[T any](ctx context.Context, key string, data T, expiration time.Duration) error {
jsonData, err := json.Marshal(data)
if err != nil {
return fmt.Errorf(errMarshalDataFailed, err)
}
if err := Redis.Set(ctx, PrefixedKey(key), jsonData, expiration).Err(); err != nil {
return fmt.Errorf(errRedisKeySetFailed, err)
}
return nil
}

Some files were not shown because too many files have changed in this diff Show More