refactor(repo): consolidate openflare-server to root and move subprojects to internal/apps

- Merge all files inside openflare-server to the repository root directory.
- Relocate agent, relay, and flared subprojects from internal/ to internal/apps/.
- Combine docker-compose files and update build context paths to root.
- Update GitHub workflows and Dockerfiles to refer to new directories and package names.
- Rewrite Go package imports across all files.
- Resolve database renew test race condition and clean up docs.
This commit is contained in:
ryan
2026-06-19 14:23:29 +08:00
parent 19d476ed7f
commit 63cd906cfc
1064 changed files with 366 additions and 1397 deletions
+391
View File
@@ -0,0 +1,391 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package disk implements a platform-level disk-backed cache with size limit, TTL, and LRU eviction.
package disk
import (
"container/list"
"encoding/binary"
"errors"
"fmt"
"os"
"path/filepath"
"sort"
"sync"
"time"
"github.com/peterbourgon/diskv/v3"
)
// ErrCacheMiss represents a cache miss.
var ErrCacheMiss = errors.New("cache miss")
// Constants for disk cache configuration and sizing
const (
headerSize = 8 // 8 bytes metadata prefix for expiration UnixNano timestamp
defaultMaxSizeMB = 100
defaultTTLMinutes = 60
cacheDirPerm = 0750
// DefaultExpiration applies the cache-wide default TTL.
DefaultExpiration time.Duration = 0
// NoExpiration stores the item without a TTL. Size limits and LRU eviction still apply.
NoExpiration time.Duration = -1
)
// Status represents the runtime cache statistics.
type Status struct {
TotalSize int64 `json:"total_size"`
KeysCount int `json:"keys_count"`
MaxSizeMB int64 `json:"max_size_mb"`
TTLMinutes int64 `json:"ttl_minutes"`
LRUEnabled bool `json:"lru_enabled"`
BasePath string `json:"base_path"`
}
// Cache implements the disk-backed cache with size limits, TTL, and LRU eviction.
type Cache struct {
mu sync.RWMutex
d *diskv.Diskv
basePath string
maxSize int64 // in bytes
defaultTTL time.Duration
lruEnabled bool
// LRU and Size tracking
currentSize int64
items map[string]*list.Element
evictList *list.List
}
type cacheItem struct {
key string
size int64
expiredAt time.Time
}
// New creates a new Cache instance.
func New(basePath string) *Cache {
d := diskv.New(diskv.Options{
BasePath: basePath,
Transform: func(_ string) []string { return []string{} }, // flat structure for easy walk
CacheSizeMax: 1024 * 1024, // 1MB in-memory cache size for diskv itself
})
c := &Cache{
d: d,
basePath: basePath,
maxSize: defaultMaxSizeMB * 1024 * 1024, // 100MB default
defaultTTL: defaultTTLMinutes * time.Minute, // 60 minutes default
lruEnabled: true,
items: make(map[string]*list.Element),
evictList: list.New(),
}
// Scan directory on startup to rebuild LRU and size tracking
_ = c.loadTracker()
return c
}
// Set stores a key-value pair in the cache.
// Use DefaultExpiration for the configured default TTL, NoExpiration for no
// TTL, or a positive duration for a business-specific TTL.
func (c *Cache) Set(key string, value []byte, ttl time.Duration) error {
c.mu.Lock()
defer c.mu.Unlock()
if ttl == DefaultExpiration {
ttl = c.defaultTTL
}
var expiredAt time.Time
if ttl > 0 {
expiredAt = time.Now().Add(ttl)
}
// Prepare data layout: 8 bytes expiration timestamp + raw payload
buf := make([]byte, headerSize+len(value))
var expNano int64
if !expiredAt.IsZero() {
expNano = expiredAt.UnixNano()
}
binary.BigEndian.PutUint64(buf[0:headerSize], uint64(expNano))
copy(buf[headerSize:], value)
// Write to diskv
if err := c.d.Write(key, buf); err != nil {
return fmt.Errorf("failed to write key to disk: %w", err)
}
// Get file size on disk (approximate)
size := int64(len(buf))
// Update memory tracker
if elem, ok := c.items[key]; ok {
item := elem.Value.(*cacheItem)
c.currentSize += size - item.size
item.size = size
item.expiredAt = expiredAt
c.evictList.MoveToFront(elem)
} else {
item := &cacheItem{
key: key,
size: size,
expiredAt: expiredAt,
}
elem := c.evictList.PushFront(item)
c.items[key] = elem
c.currentSize += size
}
// Evict items if size limit exceeded and LRU is enabled
c.evict()
return nil
}
// Get retrieves a key's value from the cache.
func (c *Cache) Get(key string) ([]byte, error) {
c.mu.RLock()
elem, ok := c.items[key]
if !ok {
c.mu.RUnlock()
return nil, ErrCacheMiss
}
item := elem.Value.(*cacheItem)
if !item.expiredAt.IsZero() && time.Now().After(item.expiredAt) {
c.mu.RUnlock()
return c.getAndDeleteIfExpired(key)
}
c.mu.RUnlock()
// Read from disk outside the lock so concurrent cache hits do not serialize on I/O.
data, err := c.d.Read(key)
if err != nil {
c.mu.Lock()
defer c.mu.Unlock()
if _, stillExists := c.items[key]; stillExists {
_ = c.deleteUnlocked(key)
}
return nil, ErrCacheMiss
}
if len(data) < headerSize {
c.mu.Lock()
defer c.mu.Unlock()
if _, stillExists := c.items[key]; stillExists {
_ = c.deleteUnlocked(key)
}
return nil, ErrCacheMiss
}
payload := data[headerSize:]
// Brief write lock only for LRU bookkeeping.
c.mu.Lock()
defer c.mu.Unlock()
elem, ok = c.items[key]
if !ok {
return nil, ErrCacheMiss
}
c.evictList.MoveToFront(elem)
return payload, nil
}
func (c *Cache) getAndDeleteIfExpired(key string) ([]byte, error) {
c.mu.Lock()
defer c.mu.Unlock()
elem, ok := c.items[key]
if !ok {
return nil, ErrCacheMiss
}
item := elem.Value.(*cacheItem)
if !item.expiredAt.IsZero() && time.Now().After(item.expiredAt) {
_ = c.deleteUnlocked(key)
return nil, ErrCacheMiss
}
data, err := c.d.Read(key)
if err != nil {
_ = c.deleteUnlocked(key)
return nil, ErrCacheMiss
}
if len(data) < headerSize {
_ = c.deleteUnlocked(key)
return nil, ErrCacheMiss
}
c.evictList.MoveToFront(elem)
return data[headerSize:], nil
}
// Delete removes a key-value pair from the cache.
func (c *Cache) Delete(key string) error {
c.mu.Lock()
defer c.mu.Unlock()
return c.deleteUnlocked(key)
}
func (c *Cache) deleteUnlocked(key string) error {
if elem, ok := c.items[key]; ok {
item := elem.Value.(*cacheItem)
c.currentSize -= item.size
c.evictList.Remove(elem)
delete(c.items, key)
}
return c.d.Erase(key)
}
// Clear flushes all cached elements.
func (c *Cache) Clear() error {
c.mu.Lock()
defer c.mu.Unlock()
c.currentSize = 0
c.items = make(map[string]*list.Element)
c.evictList.Init()
return c.d.EraseAll()
}
// Status returns the cache status.
func (c *Cache) Status() Status {
c.mu.RLock()
defer c.mu.RUnlock()
return Status{
TotalSize: c.currentSize,
KeysCount: len(c.items),
MaxSizeMB: c.maxSize / (1024 * 1024),
TTLMinutes: int64(c.defaultTTL.Minutes()),
LRUEnabled: c.lruEnabled,
BasePath: c.basePath,
}
}
// UpdatePolicy dynamically updates policies.
func (c *Cache) UpdatePolicy(maxSizeMB int64, ttlMinutes int64, lruEnabled bool) {
c.mu.Lock()
defer c.mu.Unlock()
c.maxSize = maxSizeMB * 1024 * 1024
c.defaultTTL = time.Duration(ttlMinutes) * time.Minute
c.lruEnabled = lruEnabled
c.evict()
}
// evict evicts oldest items if current size exceeds maxSize and LRU is enabled.
func (c *Cache) evict() {
if !c.lruEnabled {
return
}
for c.currentSize > c.maxSize && c.evictList.Len() > 0 {
elem := c.evictList.Back()
item := elem.Value.(*cacheItem)
c.currentSize -= item.size
c.evictList.Remove(elem)
delete(c.items, item.key)
_ = c.d.Erase(item.key)
}
}
// loadTracker scans the cache directory on startup to rebuild memory state.
func (c *Cache) loadTracker() error {
c.mu.Lock()
defer c.mu.Unlock()
// Ensure directory exists
if err := os.MkdirAll(c.basePath, cacheDirPerm); err != nil {
return err
}
type loadedItem struct {
key string
size int64
expiredAt time.Time
modTime time.Time
}
var loadedItems []loadedItem
// Walk keys through diskv
keysChan := c.d.Keys(nil)
for key := range keysChan {
// Read raw bytes to parse expiration prefix
data, err := c.d.Read(key)
if err != nil || len(data) < headerSize {
_ = c.d.Erase(key) // corrupted file, wipe
continue
}
expNano := int64(binary.BigEndian.Uint64(data[0:headerSize])) //nolint:gosec // false positive: UnixNano fits within int64
var expiredAt time.Time
if expNano > 0 {
expiredAt = time.Unix(0, expNano)
}
// Check mod time for ordering
path := filepath.Join(c.basePath, key)
info, err := os.Stat(path)
if err != nil {
continue
}
loadedItems = append(loadedItems, loadedItem{
key: key,
size: int64(len(data)),
expiredAt: expiredAt,
modTime: info.ModTime(),
})
}
// Sort by ModTime ascending (oldest first) so we rebuild LRU correctly
sort.Slice(loadedItems, func(i, j int) bool {
return loadedItems[i].modTime.Before(loadedItems[j].modTime)
})
// Populate LRU (PushFront so that newest items are at the front, oldest at the back)
for _, item := range loadedItems {
entry := &cacheItem{
key: item.key,
size: item.size,
expiredAt: item.expiredAt,
}
element := c.evictList.PushFront(entry)
c.items[item.key] = element
c.currentSize += item.size
}
return nil
}
// StartCleanupWorker periodically cleans up expired cache items.
func (c *Cache) StartCleanupWorker(interval time.Duration) {
ticker := time.NewTicker(interval)
for range ticker.C {
c.cleanExpired()
}
}
// cleanExpired scans memory for expired items and removes them.
func (c *Cache) cleanExpired() {
c.mu.Lock()
defer c.mu.Unlock()
now := time.Now()
for key, elem := range c.items {
item := elem.Value.(*cacheItem)
if !item.expiredAt.IsZero() && now.After(item.expiredAt) {
c.currentSize -= item.size
c.evictList.Remove(elem)
delete(c.items, key)
_ = c.d.Erase(key)
}
}
}
+212
View File
@@ -0,0 +1,212 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package disk
import (
"bytes"
"os"
"testing"
"time"
)
func TestDiskCacheBasic(t *testing.T) {
testDir := "uploads/test_diskcache_basic"
defer func() { _ = os.RemoveAll(testDir) }()
_ = os.RemoveAll(testDir)
c := New(testDir)
defer func() { _ = c.Clear() }()
key := "key1"
val := []byte("value1")
// Get non-existent
_, err := c.Get(key)
if err != ErrCacheMiss {
t.Fatalf("expected ErrCacheMiss, got %v", err)
}
// Set & Get
err = c.Set(key, val, 10*time.Second)
if err != nil {
t.Fatalf("failed to set cache: %v", err)
}
got, err := c.Get(key)
if err != nil {
t.Fatalf("failed to get cache: %v", err)
}
if !bytes.Equal(got, val) {
t.Errorf("expected %s, got %s", val, got)
}
// Delete
err = c.Delete(key)
if err != nil {
t.Fatalf("failed to delete: %v", err)
}
_, err = c.Get(key)
if err != ErrCacheMiss {
t.Errorf("expected ErrCacheMiss after delete, got %v", err)
}
}
func TestDiskCacheTTL(t *testing.T) {
testDir := "uploads/test_diskcache_ttl"
defer func() { _ = os.RemoveAll(testDir) }()
_ = os.RemoveAll(testDir)
c := New(testDir)
defer func() { _ = c.Clear() }()
key := "ttlkey"
val := []byte("ttlval")
// Set with 200ms TTL
err := c.Set(key, val, 200*time.Millisecond)
if err != nil {
t.Fatalf("failed to set: %v", err)
}
// Immediate Get should succeed
got, err := c.Get(key)
if err != nil {
t.Fatalf("failed to get: %v", err)
}
if !bytes.Equal(got, val) {
t.Errorf("expected %s, got %s", val, got)
}
// Sleep 250ms to expire
time.Sleep(250 * time.Millisecond)
// Get should fail with cache miss
_, err = c.Get(key)
if err != ErrCacheMiss {
t.Errorf("expected ErrCacheMiss after TTL expiration, got %v", err)
}
}
func TestDiskCacheExpirationPolicies(t *testing.T) {
testDir := "uploads/test_diskcache_expiration_policies"
defer func() { _ = os.RemoveAll(testDir) }()
_ = os.RemoveAll(testDir)
c := New(testDir)
defer func() { _ = c.Clear() }()
c.defaultTTL = 50 * time.Millisecond
if err := c.Set("default", []byte("default"), DefaultExpiration); err != nil {
t.Fatalf("Set(default, DefaultExpiration) returned error: %v", err)
}
if err := c.Set("custom", []byte("custom"), 100*time.Millisecond); err != nil {
t.Fatalf("Set(custom, 100ms) returned error: %v", err)
}
if err := c.Set("permanent", []byte("permanent"), NoExpiration); err != nil {
t.Fatalf("Set(permanent, NoExpiration) returned error: %v", err)
}
time.Sleep(75 * time.Millisecond)
if _, err := c.Get("default"); err != ErrCacheMiss {
t.Errorf("Get(default) error = %v, want ErrCacheMiss", err)
}
if _, err := c.Get("custom"); err != nil {
t.Errorf("Get(custom) returned error before custom TTL elapsed: %v", err)
}
if _, err := c.Get("permanent"); err != nil {
t.Errorf("Get(permanent) returned error: %v", err)
}
time.Sleep(50 * time.Millisecond)
if _, err := c.Get("custom"); err != ErrCacheMiss {
t.Errorf("Get(custom) error = %v, want ErrCacheMiss", err)
}
if _, err := c.Get("permanent"); err != nil {
t.Errorf("Get(permanent) returned error after other entries expired: %v", err)
}
}
func TestDiskCacheNoExpirationSurvivesReload(t *testing.T) {
testDir := "uploads/test_diskcache_no_expiration_reload"
defer func() { _ = os.RemoveAll(testDir) }()
_ = os.RemoveAll(testDir)
c := New(testDir)
if err := c.Set("permanent", []byte("value"), NoExpiration); err != nil {
t.Fatalf("Set(permanent, NoExpiration) returned error: %v", err)
}
reloaded := New(testDir)
defer func() { _ = reloaded.Clear() }()
got, err := reloaded.Get("permanent")
if err != nil {
t.Fatalf("reloaded Get(permanent) returned error: %v", err)
}
if !bytes.Equal(got, []byte("value")) {
t.Errorf("reloaded Get(permanent) = %q, want %q", got, "value")
}
}
func TestDiskCacheLRUEviction(t *testing.T) {
testDir := "uploads/test_diskcache_lru"
defer func() { _ = os.RemoveAll(testDir) }()
_ = os.RemoveAll(testDir)
c := New(testDir)
defer func() { _ = c.Clear() }()
// Force a very small max size of 20 bytes for testing (8 bytes header + payload)
// So 2 items of 2 bytes payload = 2 * (8 + 2) = 20 bytes max.
c.maxSize = 20
c.lruEnabled = true
// Write item 1: 8 + 2 = 10 bytes
err := c.Set("k1", []byte("v1"), DefaultExpiration)
if err != nil {
t.Fatalf("failed to set k1: %v", err)
}
// Write item 2: 8 + 2 = 10 bytes
err = c.Set("k2", []byte("v2"), DefaultExpiration)
if err != nil {
t.Fatalf("failed to set k2: %v", err)
}
// Both should exist
if _, err := c.Get("k1"); err != nil {
t.Errorf("k1 should exist: %v", err)
}
if _, err := c.Get("k2"); err != nil {
t.Errorf("k2 should exist: %v", err)
}
// Write item 3: 8 + 2 = 10 bytes -> total size would be 30, exceeding 20.
// This should evict the oldest item. Since k1 was accessed, but then k2 was accessed,
// wait, let's access k1 again to make it the most recently used, so k2 becomes oldest!
_, _ = c.Get("k1") // k1 is now MRU, k2 is LRU
err = c.Set("k3", []byte("v3"), DefaultExpiration)
if err != nil {
t.Fatalf("failed to set k3: %v", err)
}
// k2 should be evicted, k1 and k3 should exist
_, err = c.Get("k2")
if err != ErrCacheMiss {
t.Errorf("expected k2 to be evicted, got error %v", err)
}
if _, err := c.Get("k1"); err != nil {
t.Errorf("k1 should still exist: %v", err)
}
if _, err := c.Get("k3"); err != nil {
t.Errorf("k3 should exist: %v", err)
}
}
+73
View File
@@ -0,0 +1,73 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package ram provides a thin wrapper around Otter v2 for process-local caching.
package ram
import (
"github.com/maypok86/otter/v2"
)
const defaultMaximumSize = 256
// Options configures a RAM cache instance.
type Options struct {
// MaximumSize bounds the number of entries. Zero uses a small default.
MaximumSize int
}
// Cache is a concurrency-safe in-memory cache backed by Otter.
type Cache[K comparable, V any] struct {
inner *otter.Cache[K, V]
}
// New creates a RAM cache from the provided options.
func New[K comparable, V any](opts Options) (*Cache[K, V], error) {
maximumSize := opts.MaximumSize
if maximumSize == 0 {
maximumSize = defaultMaximumSize
}
inner, err := otter.New(&otter.Options[K, V]{
MaximumSize: maximumSize,
})
if err != nil {
return nil, err
}
return &Cache[K, V]{inner: inner}, nil
}
// MustNew creates a RAM cache and panics when configuration is invalid.
func MustNew[K comparable, V any](opts Options) *Cache[K, V] {
cache, err := New[K, V](opts)
if err != nil {
panic(err)
}
return cache
}
// GetIfPresent returns the cached value when present.
func (c *Cache[K, V]) GetIfPresent(key K) (V, bool) {
return c.inner.GetIfPresent(key)
}
// Set stores a value in the cache.
func (c *Cache[K, V]) Set(key K, value V) {
c.inner.Set(key, value)
}
// Invalidate removes one entry from the cache.
func (c *Cache[K, V]) Invalidate(key K) {
c.inner.Invalidate(key)
}
// InvalidateAll removes every entry from the cache.
func (c *Cache[K, V]) InvalidateAll() {
c.inner.InvalidateAll()
}
// EstimatedSize returns the approximate number of cached entries.
func (c *Cache[K, V]) EstimatedSize() int {
return c.inner.EstimatedSize()
}
+38
View File
@@ -0,0 +1,38 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package ram
import "testing"
func TestCacheSetGetInvalidate(t *testing.T) {
cache := MustNew[string, int](Options{MaximumSize: 8})
cache.Set("count", 3)
got, ok := cache.GetIfPresent("count")
if !ok {
t.Fatal("GetIfPresent(count) ok = false, want true")
}
if got != 3 {
t.Fatalf("GetIfPresent(count) = %d, want %d", got, 3)
}
cache.Invalidate("count")
if _, ok := cache.GetIfPresent("count"); ok {
t.Fatal("GetIfPresent(count) after Invalidate ok = true, want false")
}
}
func TestCacheInvalidateAll(t *testing.T) {
cache := MustNew[string, string](Options{MaximumSize: 8})
cache.Set("a", "1")
cache.Set("b", "2")
cache.InvalidateAll()
if cache.EstimatedSize() != 0 {
t.Fatalf("EstimatedSize() after InvalidateAll = %d, want 0", cache.EstimatedSize())
}
}
+259
View File
@@ -0,0 +1,259 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package cap 提供人机验证(CAPTCHA)功能
package cap
import (
"crypto/hmac"
"crypto/rand"
"crypto/sha256"
"encoding/base64"
"encoding/hex"
"encoding/json"
"errors"
"strconv"
"strings"
"time"
)
const (
jwtHeaderB64 = "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9"
jwtPartsCount = 3 // JWT 三段结构
defaultChallengeCount = 50 // 默认 PoW 难题数
defaultChallengeSize = 32 // 默认盐值长度
defaultDifficulty = 4 // 默认难度
defaultNonceLength = 25 // 随机 Nonce 字节长度
defaultExpires = 10 * time.Minute // 默认过期时间
)
// ChallengeConfig holds parameters for the PoW challenge
type ChallengeConfig struct {
Count int // Number of puzzles (c)
Size int // Salt length (s)
Difficulty int // Difficulty prefix length (d)
Expires time.Duration // Challenge TTL
}
// ChallengeResponse is returned to the client
type ChallengeResponse struct {
Challenge struct {
C int `json:"c"`
S int `json:"s"`
D int `json:"d"`
} `json:"challenge"`
Token string `json:"token"`
Expires int64 `json:"expires"` // ms timestamp
}
// ChallengePayload represents the signed JWT payload
type ChallengePayload struct {
Nonce string `json:"n"`
Count int `json:"c"`
Size int `json:"s"`
Difficulty int `json:"d"`
Expires int64 `json:"exp"` // ms timestamp
IssuedAt int64 `json:"iat"` // ms timestamp
Scope string `json:"sk,omitempty"`
}
// RedeemRequest payload sent by client
type RedeemRequest struct {
Token string `json:"token"`
Solutions []int `json:"solutions"`
}
// RedeemResponse returned to client after verification
type RedeemResponse struct {
Success bool `json:"success"`
Token string `json:"token,omitempty"`
Expires int64 `json:"expires,omitempty"`
Error string `json:"error,omitempty"`
}
func b64urlEncode(data []byte) string {
return base64.RawURLEncoding.EncodeToString(data)
}
func b64urlDecode(str string) ([]byte, error) {
return base64.RawURLEncoding.DecodeString(str)
}
// RandomHex generates a cryptographically secure random hexadecimal string of the specified byte length.
func RandomHex(byteLen int) string {
bytes := make([]byte, byteLen)
if _, err := rand.Read(bytes); err != nil {
panic(err)
}
return hex.EncodeToString(bytes)
}
func jwtSign(payload []byte, secret []byte) string {
body := b64urlEncode(payload)
sigInput := jwtHeaderB64 + "." + body
mac := hmac.New(sha256.New, secret)
mac.Write([]byte(sigInput))
sig := mac.Sum(nil)
return sigInput + "." + b64urlEncode(sig)
}
func jwtVerify(token string, secret []byte) ([]byte, error) {
parts := strings.Split(token, ".")
if len(parts) != jwtPartsCount {
return nil, errors.New(errInvalidTokenFormat)
}
if parts[0] != jwtHeaderB64 {
return nil, errors.New(errInvalidHeader)
}
sigInput := parts[0] + "." + parts[1]
mac := hmac.New(sha256.New, secret)
mac.Write([]byte(sigInput))
expectedSig := mac.Sum(nil)
actualSig, err := b64urlDecode(parts[2])
if err != nil {
return nil, err
}
if !hmac.Equal(expectedSig, actualSig) {
return nil, errors.New(errSignatureMismatch)
}
payload, err := b64urlDecode(parts[1])
if err != nil {
return nil, err
}
return payload, nil
}
// JwtSigHex extracts the signature part of a JWT token and returns it as a hexadecimal string.
func JwtSigHex(token string) string {
parts := strings.Split(token, ".")
if len(parts) != jwtPartsCount {
return ""
}
sigBytes, err := b64urlDecode(parts[2])
if err != nil {
return ""
}
return hex.EncodeToString(sigBytes)
}
// GenerateChallenge produces a new challenge and signed token
func GenerateChallenge(secret []byte, conf ChallengeConfig, scope string) (*ChallengeResponse, error) {
if conf.Count <= 0 {
conf.Count = defaultChallengeCount
}
if conf.Size <= 0 {
conf.Size = defaultChallengeSize
}
if conf.Difficulty <= 0 {
conf.Difficulty = defaultDifficulty
}
if conf.Expires <= 0 {
conf.Expires = defaultExpires
}
now := time.Now().UnixNano() / int64(time.Millisecond)
expires := now + int64(conf.Expires/time.Millisecond)
payload := ChallengePayload{
Nonce: RandomHex(defaultNonceLength),
Count: conf.Count,
Size: conf.Size,
Difficulty: conf.Difficulty,
Expires: expires,
IssuedAt: now,
Scope: scope,
}
payloadBytes, err := json.Marshal(payload)
if err != nil {
return nil, err
}
token := jwtSign(payloadBytes, secret)
resp := &ChallengeResponse{
Token: token,
Expires: expires,
}
resp.Challenge.C = conf.Count
resp.Challenge.S = conf.Size
resp.Challenge.D = conf.Difficulty
return resp, nil
}
// VerifyChallengeSolutions verifies client submitted solutions
func VerifyChallengeSolutions(token string, solutions []int, secret []byte, expectedScope string) (*ChallengePayload, error) {
payloadBytes, err := jwtVerify(token, secret)
if err != nil {
return nil, errors.New(errInvalidToken)
}
var payload ChallengePayload
if err := json.Unmarshal(payloadBytes, &payload); err != nil {
return nil, errors.New(errInvalidToken)
}
if expectedScope != "" && payload.Scope != expectedScope {
return nil, errors.New(errScopeMismatch)
}
now := time.Now().UnixNano() / int64(time.Millisecond)
if payload.Expires < now {
return nil, errors.New(errExpired)
}
if len(solutions) != payload.Count {
return nil, errors.New(errInvalidSolutions)
}
tokenFnv := fnv1a(token)
for i := 0; i < payload.Count; i++ {
idxStr := strconv.Itoa(i + 1)
saltSeed := fnv1aResume(tokenFnv, idxStr)
targetSeed := fnv1aResume(saltSeed, "d")
salt := prngFromHash(saltSeed, payload.Size)
target := prngFromHash(targetSeed, payload.Difficulty)
hashInput := salt + strconv.Itoa(solutions[i])
hashBytes := sha256.Sum256([]byte(hashInput))
hashHex := hex.EncodeToString(hashBytes[:])
if !strings.HasPrefix(hashHex, target) {
return nil, errors.New(errInvalidSolution)
}
}
return &payload, nil
}
// Solve is a utility function to solve a challenge (mainly used for tests and reference implementation)
func Solve(token string, count, size, difficulty int) []int {
solutions := make([]int, count)
tokenFnv := fnv1a(token)
for i := 0; i < count; i++ {
idxStr := strconv.Itoa(i + 1)
saltSeed := fnv1aResume(tokenFnv, idxStr)
targetSeed := fnv1aResume(saltSeed, "d")
salt := prngFromHash(saltSeed, size)
target := prngFromHash(targetSeed, difficulty)
for nonce := 0; nonce < 1000000; nonce++ {
hashInput := salt + strconv.Itoa(nonce)
hashBytes := sha256.Sum256([]byte(hashInput))
hashHex := hex.EncodeToString(hashBytes[:])
if strings.HasPrefix(hashHex, target) {
solutions[i] = nonce
break
}
}
}
return solutions
}
+15
View File
@@ -0,0 +1,15 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package cap
const (
errInvalidTokenFormat = "invalid token format"
errInvalidHeader = "invalid header"
errSignatureMismatch = "signature mismatch"
errInvalidToken = "invalid_token"
errScopeMismatch = "scope_mismatch"
errExpired = "expired"
errInvalidSolutions = "invalid_solutions"
errInvalidSolution = "invalid_solution"
)
+49
View File
@@ -0,0 +1,49 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package cap
import (
"fmt"
"strings"
)
// fnv1a returns the 32-bit FNV-1a hash of a string
//
//nolint:mnd // FNV-1a 算法位移常量
func fnv1a(str string) uint32 {
var hash uint32 = 2166136261
for i := 0; i < len(str); i++ {
hash ^= uint32(str[i])
hash += (hash << 1) + (hash << 4) + (hash << 7) + (hash << 8) + (hash << 24)
}
return hash
}
// fnv1aResume resumes FNV-1a hashing from a given state
//
//nolint:mnd // FNV-1a 算法位移常量
func fnv1aResume(state uint32, str string) uint32 {
h := state
for i := 0; i < len(str); i++ {
h ^= uint32(str[i])
h += (h << 1) + (h << 4) + (h << 7) + (h << 8) + (h << 24)
}
return h
}
// prngFromHash generates a hex string of specified length using an initial hash state
//
//nolint:mnd // xorshift 算法位移常量
func prngFromHash(initialHash uint32, length int) string {
state := initialHash
var result strings.Builder
for result.Len() < length {
state ^= state << 13
state ^= state >> 17
state ^= state << 5
hexStr := fmt.Sprintf("%08x", state)
result.WriteString(hexStr)
}
return result.String()[:length]
}
+186
View File
@@ -0,0 +1,186 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package cap
import (
"context"
"sync"
"time"
"github.com/redis/go-redis/v9"
)
// Store defines the storage interface for challenge nonces and verification tokens
type Store interface {
Get(ctx context.Context, key string) (string, bool, error)
Set(ctx context.Context, key string, val string, ttl time.Duration) error
Delete(ctx context.Context, key string) error
// SetNX atomically sets key=val with the given TTL only when the key does not
// exist yet. It returns true when the key was actually written (i.e. this
// caller "won" the race), and false when the key already existed.
SetNX(ctx context.Context, key string, val string, ttl time.Duration) (bool, error)
// GetAndDelete atomically retrieves the value of key and removes it in a
// single operation. Returns ("", false, nil) when the key does not exist.
GetAndDelete(ctx context.Context, key string) (string, bool, error)
}
type memoryItem struct {
value string
expiresAt time.Time
}
// MemoryStore is a thread-safe in-memory implementation of Store
type MemoryStore struct {
items map[string]memoryItem
mu sync.Mutex // unified write-lock; promotes to exclusive for all ops
}
// NewMemoryStore creates and initializes a new MemoryStore
func NewMemoryStore(cleanupInterval time.Duration) *MemoryStore {
store := &MemoryStore{
items: make(map[string]memoryItem),
}
if cleanupInterval > 0 {
go store.startCleanupLoop(cleanupInterval)
}
return store
}
// Get 从 MemoryStore 获取指定 key 的值
func (s *MemoryStore) Get(_ context.Context, key string) (string, bool, error) {
s.mu.Lock()
defer s.mu.Unlock()
return s.getLocked(key)
}
// getLocked is the internal helper – caller must hold s.mu.
func (s *MemoryStore) getLocked(key string) (string, bool, error) {
item, found := s.items[key]
if !found {
return "", false, nil
}
if time.Now().After(item.expiresAt) {
delete(s.items, key)
return "", false, nil
}
return item.value, true, nil
}
// Set 向 MemoryStore 写入指定 key 的值
func (s *MemoryStore) Set(_ context.Context, key string, val string, ttl time.Duration) error {
s.mu.Lock()
defer s.mu.Unlock()
s.items[key] = memoryItem{
value: val,
expiresAt: time.Now().Add(ttl),
}
return nil
}
// Delete 从 MemoryStore 删除指定 key
func (s *MemoryStore) Delete(_ context.Context, key string) error {
s.mu.Lock()
defer s.mu.Unlock()
delete(s.items, key)
return nil
}
// SetNX atomically sets key only when it is absent (or expired).
// Returns true if the key was written by this call.
func (s *MemoryStore) SetNX(_ context.Context, key string, val string, ttl time.Duration) (bool, error) {
s.mu.Lock()
defer s.mu.Unlock()
_, exists, _ := s.getLocked(key)
if exists {
return false, nil
}
s.items[key] = memoryItem{
value: val,
expiresAt: time.Now().Add(ttl),
}
return true, nil
}
// GetAndDelete atomically retrieves and removes key in one critical section.
func (s *MemoryStore) GetAndDelete(_ context.Context, key string) (string, bool, error) {
s.mu.Lock()
defer s.mu.Unlock()
val, exists, err := s.getLocked(key)
if err != nil || !exists {
return "", false, err
}
delete(s.items, key)
return val, true, nil
}
func (s *MemoryStore) startCleanupLoop(interval time.Duration) {
ticker := time.NewTicker(interval)
for range ticker.C {
s.cleanupExpired()
}
}
func (s *MemoryStore) cleanupExpired() {
now := time.Now()
s.mu.Lock()
defer s.mu.Unlock()
for k, v := range s.items {
if now.After(v.expiresAt) {
delete(s.items, k)
}
}
}
// RedisStore is a GORM-compatible/standalone Redis-backed implementation of Store
type RedisStore struct {
client redis.UniversalClient
}
// NewRedisStore creates a new RedisStore wrapping a redis.UniversalClient
func NewRedisStore(client redis.UniversalClient) *RedisStore {
return &RedisStore{
client: client,
}
}
// Get 从 RedisStore 获取指定 key 的值
func (s *RedisStore) Get(ctx context.Context, key string) (string, bool, error) {
val, err := s.client.Get(ctx, key).Result()
if err == redis.Nil {
return "", false, nil
}
if err != nil {
return "", false, err
}
return val, true, nil
}
// Set 向 RedisStore 写入指定 key 的值
func (s *RedisStore) Set(ctx context.Context, key string, val string, ttl time.Duration) error {
return s.client.Set(ctx, key, val, ttl).Err()
}
// Delete 从 RedisStore 删除指定 key
func (s *RedisStore) Delete(ctx context.Context, key string) error {
return s.client.Del(ctx, key).Err()
}
// SetNX wraps Redis SET NX – returns true only when the key was newly created.
func (s *RedisStore) SetNX(ctx context.Context, key string, val string, ttl time.Duration) (bool, error) {
return s.client.SetNX(ctx, key, val, ttl).Result()
}
// GetAndDelete wraps Redis GETDEL (available since Redis 6.2).
func (s *RedisStore) GetAndDelete(ctx context.Context, key string) (string, bool, error) {
val, err := s.client.GetDel(ctx, key).Result()
if err == redis.Nil {
return "", false, nil
}
if err != nil {
return "", false, err
}
return val, true, nil
}
+25
View File
@@ -0,0 +1,25 @@
package geoip
import (
"fmt"
"net"
)
type EmptyProvider struct{}
func (e *EmptyProvider) Name() string {
return "EmptyProvider"
}
func (e *EmptyProvider) Initialize() error {
return nil
}
func (e *EmptyProvider) GetGeoInfo(ip net.IP) (*GeoInfo, error) {
return nil, fmt.Errorf("you are using an empty GeoIP provider, please set a valid provider")
}
func (e *EmptyProvider) UpdateDatabase() error {
return fmt.Errorf("you are using an empty GeoIP provider, please set a valid provider")
}
func (e *EmptyProvider) Close() error {
return nil
}
+242
View File
@@ -0,0 +1,242 @@
package geoip
import (
"fmt"
"log/slog"
"net"
"strings"
"sync"
"time"
"unicode"
ristretto "github.com/dgraph-io/ristretto/v2"
)
var CurrentProvider GeoIPService
var geoCache *providerCache
var providerMutex sync.RWMutex
var providerFactory = newProvider
const (
ProviderDisabled = "disabled"
ProviderMaxMind = "mmdb"
ProviderIPAPI = "ip-api"
ProviderGeoJS = "geojs"
ProviderIPInfo = "ipinfo"
)
type GeoInfo struct {
ISOCode string
Name string
Latitude *float64
Longitude *float64
}
func init() {
CurrentProvider = &EmptyProvider{}
geoCache = newProviderCache(48 * time.Hour)
}
// GeoIPService 接口定义了获取地理位置信息的核心方法。
type GeoIPService interface {
Name() string
GetGeoInfo(ip net.IP) (*GeoInfo, error)
UpdateDatabase() error
Close() error
}
type cachedGeoInfo struct {
info *GeoInfo
expiresAt time.Time
}
type providerCache struct {
items *ristretto.Cache[string, cachedGeoInfo]
duration time.Duration
}
func newProviderCache(duration time.Duration) *providerCache {
items, err := ristretto.NewCache(&ristretto.Config[string, cachedGeoInfo]{
NumCounters: 1e5,
MaxCost: 2e4,
BufferItems: 64,
})
if err != nil {
panic(err)
}
return &providerCache{
items: items,
duration: duration,
}
}
func (c *providerCache) Get(key string) (*GeoInfo, bool) {
entry, ok := c.items.Get(key)
if !ok {
return nil, false
}
if time.Now().After(entry.expiresAt) {
c.items.Del(key)
return nil, false
}
return entry.info, true
}
func (c *providerCache) Set(key string, info *GeoInfo) {
c.items.Set(key, cachedGeoInfo{
info: info,
expiresAt: time.Now().Add(c.duration),
}, 1)
c.items.Wait()
}
func (c *providerCache) Flush() {
c.items.Clear()
}
func GetRegionUnicodeEmoji(isoCode string) string {
if len(isoCode) != 2 {
return ""
}
isoCode = strings.ToUpper(isoCode)
if !unicode.IsLetter(rune(isoCode[0])) || !unicode.IsLetter(rune(isoCode[1])) {
return ""
}
rune1 := rune(0x1F1E6 + (rune(isoCode[0]) - 'A'))
rune2 := rune(0x1F1E6 + (rune(isoCode[1]) - 'A'))
return string(rune1) + string(rune2)
}
func InitGeoIP(provider string) {
providerName := normalizeProvider(provider)
nextProvider, err := providerFactory(providerName)
if err != nil {
slog.Error("initialize GeoIP provider failed", "provider", providerName, "error", err)
nextProvider = &EmptyProvider{}
}
setProvider(nextProvider)
if providerName == ProviderDisabled {
slog.Info("GeoIP provider disabled")
return
}
slog.Info("GeoIP provider configured", "provider", CurrentProvider.Name())
}
func GetGeoInfo(ip net.IP) (*GeoInfo, error) {
if ip == nil {
return nil, fmt.Errorf("IP address cannot be nil")
}
provider := getProvider()
cacheKey := provider.Name() + ":" + ip.String()
if cachedInfo, found := geoCache.Get(cacheKey); found {
return cachedInfo, nil
}
info, err := provider.GetGeoInfo(ip)
if err == nil && info != nil {
geoCache.Set(cacheKey, info)
}
return info, err
}
func LookupGeoInfoWithProvider(providerName string, ip net.IP) (*GeoInfo, error) {
if ip == nil {
return nil, fmt.Errorf("IP address cannot be nil")
}
provider, err := providerFactory(normalizeProvider(providerName))
if err != nil {
return nil, err
}
defer func() {
if closeErr := provider.Close(); closeErr != nil {
slog.Warn("close temporary GeoIP provider failed", "provider", provider.Name(), "error", closeErr)
}
}()
return provider.GetGeoInfo(ip)
}
func UpdateDatabase() error {
err := getProvider().UpdateDatabase()
if err == nil {
geoCache.Flush()
slog.Info("GeoIP cache cleared due to database update.")
}
return err
}
func IsValidProvider(provider string) bool {
switch normalizeProvider(provider) {
case ProviderDisabled, ProviderMaxMind, ProviderIPAPI, ProviderGeoJS, ProviderIPInfo:
return true
default:
return false
}
}
func normalizeProvider(provider string) string {
normalized := strings.TrimSpace(strings.ToLower(provider))
if normalized == "" {
return ProviderDisabled
}
return normalized
}
func newProvider(provider string) (GeoIPService, error) {
switch provider {
case ProviderDisabled:
return &EmptyProvider{}, nil
case ProviderMaxMind:
return NewMaxMindGeoIPService()
case ProviderIPAPI:
return NewIPAPIService()
case ProviderGeoJS:
return NewGeoJSService()
case ProviderIPInfo:
return NewIPInfoService()
default:
return nil, fmt.Errorf("unsupported GeoIP provider %q", provider)
}
}
func setProvider(provider GeoIPService) {
providerMutex.Lock()
previous := CurrentProvider
CurrentProvider = provider
providerMutex.Unlock()
geoCache.Flush()
if previous != nil && previous != provider {
if err := previous.Close(); err != nil {
slog.Warn("close previous GeoIP provider failed", "error", err)
}
}
}
func getProvider() GeoIPService {
providerMutex.RLock()
defer providerMutex.RUnlock()
if CurrentProvider == nil {
return &EmptyProvider{}
}
return CurrentProvider
}
func float64Pointer(value float64) *float64 {
return &value
}
func ProviderFactoryForTest() func(string) (GeoIPService, error) {
return providerFactory
}
func SetProviderFactoryForTest(factory func(string) (GeoIPService, error)) {
if factory == nil {
providerFactory = newProvider
return
}
providerFactory = factory
}
+99
View File
@@ -0,0 +1,99 @@
package geoip
import (
"net"
"testing"
)
type fakeProvider struct {
calls int
}
func (f *fakeProvider) Name() string {
return "fake"
}
func (f *fakeProvider) GetGeoInfo(ip net.IP) (*GeoInfo, error) {
f.calls++
return &GeoInfo{
ISOCode: "CN",
Name: "China",
}, nil
}
func (f *fakeProvider) UpdateDatabase() error {
return nil
}
func (f *fakeProvider) Close() error {
return nil
}
func TestGetGeoInfoCachesByProviderAndIP(t *testing.T) {
originalProvider := CurrentProvider
geoCache.Flush()
fake := &fakeProvider{}
CurrentProvider = fake
defer func() {
CurrentProvider = originalProvider
}()
ip := net.ParseIP("8.8.8.8")
record, err := GetGeoInfo(ip)
if err != nil {
t.Fatalf("expected nil error, got %v", err)
}
if record == nil || record.ISOCode != "CN" {
t.Fatalf("expected cached record, got %#v", record)
}
_, err = GetGeoInfo(ip)
if err != nil {
t.Fatalf("expected nil error on second call, got %v", err)
}
if fake.calls != 1 {
t.Fatalf("expected provider to be called once, got %d", fake.calls)
}
}
func TestUnicodeEmoji(t *testing.T) {
emoji := GetRegionUnicodeEmoji("CN")
if emoji != "🇨🇳" {
t.Errorf("expected emoji for CN, got %s", emoji)
}
}
func TestIsValidProvider(t *testing.T) {
cases := map[string]bool{
"disabled": true,
"mmdb": true,
"ip-api": true,
"geojs": true,
"ipinfo": true,
"unknown": false,
}
for provider, want := range cases {
if got := IsValidProvider(provider); got != want {
t.Fatalf("provider %s validity mismatch: want %v, got %v", provider, want, got)
}
}
}
func TestLookupGeoInfoWithProviderUsesTemporaryProvider(t *testing.T) {
previousFactory := providerFactory
providerFactory = func(provider string) (GeoIPService, error) {
return &fakeProvider{}, nil
}
defer func() {
providerFactory = previousFactory
}()
info, err := LookupGeoInfoWithProvider("ipinfo", net.ParseIP("8.8.8.8"))
if err != nil {
t.Fatalf("expected lookup to succeed, got %v", err)
}
if info == nil || info.ISOCode != "CN" || info.Name != "China" {
t.Fatalf("unexpected geo info: %#v", info)
}
}
+84
View File
@@ -0,0 +1,84 @@
package geoip
import (
"encoding/json"
"fmt"
"net"
"net/http"
"time"
)
// GeoJSService 使用 geojs.io 服务实现 GeoIPService 接口。
type GeoJSService struct {
Client *http.Client
}
// geoJSResponse 定义了 geojs.io 服务返回的 JSON 响应的结构。
// 我们只定义我们需要的字段。
type geoJSResponse struct {
Country string `json:"country"`
CountryCode string `json:"country_code"`
Latitude float64 `json:"latitude,string"`
Longitude float64 `json:"longitude,string"`
// 可以根据需要添加其他字段,例如:
// City string `json:"city"`
// Region string `json:"region"`
}
// NewGeoJSService 创建并返回一个 GeoJSService 的新实例。
func NewGeoJSService() (*GeoJSService, error) {
return &GeoJSService{
Client: &http.Client{
Timeout: 5 * time.Second, // 设置一个合理的超时时间
},
}, nil
}
// Name 返回服务的名称。
func (s *GeoJSService) Name() string {
return "geojs.io"
}
// GetGeoInfo 使用 geojs.io 服务检索给定 IP 地址的地理位置信息。
func (s *GeoJSService) GetGeoInfo(ip net.IP) (*GeoInfo, error) {
// GeoJS 的 API 端点
apiURL := fmt.Sprintf("https://get.geojs.io/v1/ip/geo/%s.json", ip.String())
resp, err := s.Client.Get(apiURL)
if err != nil {
return nil, fmt.Errorf("failed to get geo info from geojs.io: %w", err)
}
defer resp.Body.Close()
// 检查响应状态码
if resp.StatusCode != http.StatusOK {
return nil, fmt.Errorf("geojs.io returned non-200 status code: %d", resp.StatusCode)
}
var apiResp geoJSResponse
if err := json.NewDecoder(resp.Body).Decode(&apiResp); err != nil {
return nil, fmt.Errorf("failed to decode geojs.io response: %w", err)
}
// 检查国家代码是否为空,因为 geojs 对无效/私有IP可能返回200 OK但内容为空
if apiResp.CountryCode == "" {
return nil, fmt.Errorf("geojs.io returned empty geo info for ip: %s", ip.String())
}
return &GeoInfo{
ISOCode: apiResp.CountryCode,
Name: apiResp.Country,
Latitude: float64Pointer(apiResp.Latitude),
Longitude: float64Pointer(apiResp.Longitude),
}, nil
}
// UpdateDatabase 对于 geojs.io 是一个空操作,因为它是一个 Web 服务。
func (s *GeoJSService) UpdateDatabase() error {
return nil
}
// Close 对于 geojs.io 是一个空操作。
func (s *GeoJSService) Close() error {
return nil
}
+86
View File
@@ -0,0 +1,86 @@
package geoip
import (
"encoding/json"
"fmt"
"net"
"net/http"
"time"
)
// IPAPIService 使用 ip-api.com 服务实现 GeoIPService 接口。
type IPAPIService struct {
Client *http.Client
}
// ipAPIResponse 定义了 ip-api.com 服务返回的 JSON 响应的结构。
type ipAPIResponse struct {
Status string `json:"status"`
Message string `json:"message"` // 当 status 为 fail 时出现
Country string `json:"country"`
CountryCode string `json:"countryCode"`
Region string `json:"region"`
RegionName string `json:"regionName"`
City string `json:"city"`
Zip string `json:"zip"`
Lat float64 `json:"lat"`
Lon float64 `json:"lon"`
Timezone string `json:"timezone"`
ISP string `json:"isp"`
Org string `json:"org"`
As string `json:"as"`
Query string `json:"query"`
}
func (s *IPAPIService) Name() string {
return "ip-api.com"
}
// NewIPAPIService 创建并返回一个 IPAPIService 的新实例。
func NewIPAPIService() (*IPAPIService, error) {
return &IPAPIService{
Client: &http.Client{
Timeout: 5 * time.Second, // 设置请求超时
},
}, nil
}
// GetGeoInfo 使用 ip-api.com 服务检索给定 IP 地址的地理位置信息。
func (s *IPAPIService) GetGeoInfo(ip net.IP) (*GeoInfo, error) {
// API URL, 使用 fields 参数来仅请求需要的字段
apiURL := fmt.Sprintf("http://ip-api.com/json/%s?fields=status,message,country,countryCode", ip.String())
resp, err := s.Client.Get(apiURL)
if err != nil {
return nil, fmt.Errorf("failed to get geo info from ip-api.com: %w", err)
}
defer resp.Body.Close()
var apiResp ipAPIResponse
if err := json.NewDecoder(resp.Body).Decode(&apiResp); err != nil {
return nil, fmt.Errorf("failed to decode ip-api.com response: %w", err)
}
if apiResp.Status != "success" {
return nil, fmt.Errorf("ip-api.com returned an error: %s", apiResp.Message)
}
return &GeoInfo{
ISOCode: apiResp.CountryCode,
Name: apiResp.Country,
Latitude: float64Pointer(apiResp.Lat),
Longitude: float64Pointer(apiResp.Lon),
}, nil
}
// UpdateDatabase 对于 ip-api.com 是一个空操作,因为它是一个 Web 服务。
func (s *IPAPIService) UpdateDatabase() error {
// 无需执行任何操作,因为数据由外部服务提供
return nil
}
// Close 对于 ip-api.com 是一个空操作,因为没有需要关闭的持久连接。
func (s *IPAPIService) Close() error {
// 无需执行任何操作
return nil
}
+112
View File
@@ -0,0 +1,112 @@
package geoip
import (
"encoding/json"
"fmt"
"net"
"net/http"
"strconv"
"strings"
"time"
)
// IPInfoService 使用 ipinfo.io 服务实现 GeoIPService 接口。
type IPInfoService struct {
Client *http.Client
// 每天 1000 次请求,限制由 IP 地址的所有人共享。
// APIToken string
}
// ipInfoResponse 定义了 ipinfo.io 服务返回的 JSON 响应的结构,只包含免费额度可用的字段。
type ipInfoResponse struct {
IP string `json:"ip"`
Hostname string `json:"hostname"`
City string `json:"city"`
Region string `json:"region"`
Country string `json:"country"`
CountryCode string `json:"countryCode"` // ipinfo.io 返回 "country" 的 ISO 代码,这里为了与 GeoInfo 保持一致,额外添加一个 CountryCode
Loc string `json:"loc"` // Latitude,Longitude
Org string `json:"org"`
Postal string `json:"postal"`
Timezone string `json:"timezone"`
}
// NewIPInfoService 创建并返回一个 IPInfoService 的新实例。
func NewIPInfoService() (*IPInfoService, error) {
return &IPInfoService{
Client: &http.Client{
Timeout: 5 * time.Second,
},
}, nil
}
// Name 返回服务的名称。
func (s *IPInfoService) Name() string {
return "ipinfo.io"
}
// GetGeoInfo 使用 ipinfo.io 服务检索给定 IP 地址的地理位置信息。
// 免费额度主要提供国家信息。
func (s *IPInfoService) GetGeoInfo(ip net.IP) (*GeoInfo, error) {
// IPinfo 免费额度不需要 API token 就可以查询基本的 IP 信息。
// API URL: https://ipinfo.io/json (查询自身IP) 或 https://ipinfo.io/YOUR_IP/json
apiURL := fmt.Sprintf("https://ipinfo.io/%s/json", ip.String())
resp, err := s.Client.Get(apiURL)
if err != nil {
return nil, fmt.Errorf("failed to get geo info from ipinfo.io: %w", err)
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusOK {
return nil, fmt.Errorf("ipinfo.io returned non-200 status: %d %s", resp.StatusCode, resp.Status)
}
var apiResp ipInfoResponse
if err := json.NewDecoder(resp.Body).Decode(&apiResp); err != nil {
return nil, fmt.Errorf("failed to decode ipinfo.io response: %w", err)
}
latitude, longitude := parseIPInfoCoordinates(apiResp.Loc)
// IPinfo 的 "country" 字段直接返回 ISO 2-letter code,例如 "US", "CN"
// 我们需要将 "country" 字段作为 ISOCode,并尝试获取其对应的国家名称。
// IPinfo 响应中通常不直接提供完整的国家名称,但我们可以通过 CountryCode 映射。
// 为了简化并符合 GeoInfo 结构,我们直接使用 Country 作为 ISOCode,并尝试从 CountryCode 获取名称。
// 实际上,IPinfo 的 'country' 字段就是 ISO 2-letter code。
// 如果需要完整的国家名称,可能需要一个本地的 ISO 代码到名称的映射。
// 为了与 GetRegionUnicodeEmoji 函数兼容,我们直接使用 country 作为 ISOCode。
return &GeoInfo{
ISOCode: apiResp.Country,
Name: apiResp.Country,
Latitude: latitude,
Longitude: longitude,
}, nil
}
// UpdateDatabase 对于 ipinfo.io 是一个空操作,因为它是一个 Web 服务。
func (s *IPInfoService) UpdateDatabase() error {
// 无需执行任何操作,因为数据由外部服务提供
return nil
}
// Close 对于 ipinfo.io 是一个空操作,因为没有需要关闭的持久连接。
func (s *IPInfoService) Close() error {
// 无需执行任何操作
return nil
}
func parseIPInfoCoordinates(value string) (*float64, *float64) {
parts := strings.Split(strings.TrimSpace(value), ",")
if len(parts) != 2 {
return nil, nil
}
latitudeValue, latErr := strconv.ParseFloat(strings.TrimSpace(parts[0]), 64)
longitudeValue, lonErr := strconv.ParseFloat(strings.TrimSpace(parts[1]), 64)
if latErr != nil || lonErr != nil {
return nil, nil
}
return float64Pointer(latitudeValue), float64Pointer(longitudeValue)
}
+66
View File
@@ -0,0 +1,66 @@
package iputil
import (
"net"
"strings"
)
func NormalizeIP(raw string) string {
trimmed := strings.TrimSpace(raw)
if trimmed == "" {
return ""
}
ip := net.ParseIP(trimmed)
if ip == nil {
return ""
}
if ipv4 := ip.To4(); ipv4 != nil {
return ipv4.String()
}
return ip.String()
}
func NormalizeRemoteAddr(remoteAddr string) string {
trimmed := strings.TrimSpace(remoteAddr)
if trimmed == "" {
return ""
}
if host, _, err := net.SplitHostPort(trimmed); err == nil {
return NormalizeIP(host)
}
return NormalizeIP(trimmed)
}
func IsPublic(ip net.IP) bool {
if ip == nil {
return false
}
if ipv4 := ip.To4(); ipv4 != nil {
ip = ipv4
}
if !ip.IsGlobalUnicast() || ip.IsLoopback() || ip.IsPrivate() || ip.IsLinkLocalUnicast() || ip.IsMulticast() || ip.IsUnspecified() {
return false
}
return true
}
func IsPublicString(raw string) bool {
ip := net.ParseIP(strings.TrimSpace(raw))
return IsPublic(ip)
}
func Score(ip net.IP) int {
if ip == nil {
return -1
}
if ipv4 := ip.To4(); ipv4 != nil {
ip = ipv4
}
if !ip.IsGlobalUnicast() || ip.IsLoopback() || ip.IsMulticast() || ip.IsUnspecified() {
return -1
}
if IsPublic(ip) {
return 2
}
return 1
}
+45
View File
@@ -0,0 +1,45 @@
package iputil
import (
"net"
"testing"
)
func TestNormalizeIP(t *testing.T) {
if got := NormalizeIP(" 8.8.8.8 "); got != "8.8.8.8" {
t.Fatalf("unexpected normalized ipv4: %q", got)
}
if got := NormalizeIP("[::1]"); got != "" {
t.Fatalf("expected invalid bracketed host to be rejected, got %q", got)
}
}
func TestNormalizeRemoteAddr(t *testing.T) {
if got := NormalizeRemoteAddr("203.0.113.10:8443"); got != "203.0.113.10" {
t.Fatalf("unexpected remote addr normalization: %q", got)
}
}
func TestIsPublic(t *testing.T) {
if !IsPublic(net.ParseIP("8.8.8.8")) {
t.Fatal("expected public ip to be detected")
}
if IsPublic(net.ParseIP("10.0.0.8")) {
t.Fatal("expected private ip to be rejected")
}
if IsPublic(net.ParseIP("127.0.0.1")) {
t.Fatal("expected loopback ip to be rejected")
}
}
func TestScore(t *testing.T) {
if got := Score(net.ParseIP("8.8.8.8")); got != 2 {
t.Fatalf("unexpected score for public ip: %d", got)
}
if got := Score(net.ParseIP("10.0.0.8")); got != 1 {
t.Fatalf("unexpected score for private ip: %d", got)
}
if got := Score(net.ParseIP("127.0.0.1")); got != -1 {
t.Fatalf("unexpected score for loopback ip: %d", got)
}
}
+171
View File
@@ -0,0 +1,171 @@
package geoip
import (
"fmt"
"io"
"net"
"net/http"
"os"
"path/filepath"
"sync"
"github.com/oschwald/maxminddb-golang"
)
var GeoIpUrl = "https://raw.githubusercontent.com/Loyalsoldier/geoip/release/GeoLite2-Country.mmdb"
var GeoIpFilePath = "./data/GeoLite2-Country.mmdb"
type GeoIpRecord struct {
Country struct {
ISOCode string `maxminddb:"iso_code"`
Names map[string]string `maxminddb:"names"`
} `maxminddb:"country"`
}
type MaxMindGeoIPService struct {
maxMindDBReader *maxminddb.Reader
dbFilePath string
mu sync.RWMutex
}
func (s *MaxMindGeoIPService) Name() string {
return "MaxMind"
}
func NewMaxMindGeoIPService() (*MaxMindGeoIPService, error) {
return NewMaxMindGeoIPServiceWithConfig(GeoIpFilePath, GeoIpUrl)
}
func NewMaxMindGeoIPServiceWithConfig(dbFilePath string, downloadURL string) (*MaxMindGeoIPService, error) {
if dbFilePath == "" {
dbFilePath = GeoIpFilePath
}
if downloadURL == "" {
downloadURL = GeoIpUrl
}
service := &MaxMindGeoIPService{
dbFilePath: dbFilePath,
}
if err := os.MkdirAll(filepath.Dir(service.dbFilePath), os.ModePerm); err != nil {
return nil, fmt.Errorf("failed to create data directory for MaxMind database: %w", err)
}
if _, err := os.Stat(service.dbFilePath); os.IsNotExist(err) {
if err := DownloadMaxMindDatabase(service.dbFilePath, downloadURL); err != nil {
return nil, fmt.Errorf("failed to download initial MaxMind database: %w", err)
}
}
if err := service.initialize(); err != nil {
return nil, fmt.Errorf("failed to initialize MaxMind database: %w", err)
}
return service, nil
}
func (s *MaxMindGeoIPService) initialize() error {
s.mu.Lock()
defer s.mu.Unlock()
if s.maxMindDBReader != nil {
_ = s.maxMindDBReader.Close()
s.maxMindDBReader = nil
}
reader, err := maxminddb.Open(s.dbFilePath)
if err != nil {
return fmt.Errorf("error opening MaxMind database at %s: %w", s.dbFilePath, err)
}
s.maxMindDBReader = reader
return nil
}
func (s *MaxMindGeoIPService) GetGeoInfo(ip net.IP) (*GeoInfo, error) {
s.mu.RLock()
defer s.mu.RUnlock()
if s.maxMindDBReader == nil {
return nil, fmt.Errorf("MaxMind database is not initialized or failed to open")
}
if ip == nil {
return nil, fmt.Errorf("IP address cannot be nil")
}
var record GeoIpRecord
if err := s.maxMindDBReader.Lookup(ip, &record); err != nil {
return nil, fmt.Errorf("error looking up IP %s in MaxMind database: %w", ip.String(), err)
}
geoInfo := &GeoInfo{
ISOCode: record.Country.ISOCode,
Name: record.Country.Names["en"],
}
if geoInfo.Name == "" && geoInfo.ISOCode != "" {
geoInfo.Name = geoInfo.ISOCode
}
return geoInfo, nil
}
func (s *MaxMindGeoIPService) UpdateDatabase() error {
if err := DownloadMaxMindDatabase(s.dbFilePath, GeoIpUrl); err != nil {
return err
}
return s.initialize()
}
func DownloadMaxMindDatabase(dbFilePath string, downloadURL string) error {
if dbFilePath == "" {
dbFilePath = GeoIpFilePath
}
if downloadURL == "" {
downloadURL = GeoIpUrl
}
resp, err := http.Get(downloadURL)
if err != nil {
return fmt.Errorf("failed to initiate MaxMind database download: %w", err)
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusOK {
return fmt.Errorf("failed to download MaxMind database: HTTP status %s", resp.Status)
}
if err := os.MkdirAll(filepath.Dir(dbFilePath), os.ModePerm); err != nil {
return fmt.Errorf("failed to create data directory for MaxMind database update: %w", err)
}
tempPath := dbFilePath + ".download"
out, err := os.Create(tempPath)
if err != nil {
return fmt.Errorf("failed to create MaxMind database file at %s: %w", tempPath, err)
}
defer func() {
_ = out.Close()
}()
if _, err = io.Copy(out, resp.Body); err != nil {
return fmt.Errorf("failed to write MaxMind database file: %w", err)
}
if err = out.Close(); err != nil {
return fmt.Errorf("failed to close MaxMind database file: %w", err)
}
if err = os.Rename(tempPath, dbFilePath); err != nil {
return fmt.Errorf("failed to move MaxMind database file into place: %w", err)
}
return nil
}
func (s *MaxMindGeoIPService) Close() error {
s.mu.Lock()
defer s.mu.Unlock()
if s.maxMindDBReader != nil {
err := s.maxMindDBReader.Close()
s.maxMindDBReader = nil
if err != nil {
return fmt.Errorf("error closing MaxMind database: %w", err)
}
}
return nil
}
+152
View File
@@ -0,0 +1,152 @@
package geoip
import (
"context"
"encoding/json"
"errors"
"fmt"
"io"
"net"
"net/http"
"strings"
"time"
"github.com/Rain-kl/Wavelet/pkg/geoip/iputil"
)
const defaultOutboundIPLookupTimeout = 5 * time.Second
// OutboundIPStrategy defines a lookup strategy for the current public egress IP.
type OutboundIPStrategy interface {
Name() string
GetOutboundIP(ctx context.Context) (net.IP, error)
}
// OutboundIPAPIAdapter adapts a third-party HTTP API response into an IP value.
type OutboundIPAPIAdapter interface {
Name() string
Endpoint() string
DecodeIP(io.Reader) (net.IP, error)
}
type HTTPOutboundIPStrategy struct {
Client *http.Client
Adapter OutboundIPAPIAdapter
}
func NewHTTPOutboundIPStrategy(adapter OutboundIPAPIAdapter, client *http.Client) *HTTPOutboundIPStrategy {
if client == nil {
client = &http.Client{Timeout: defaultOutboundIPLookupTimeout}
}
return &HTTPOutboundIPStrategy{
Client: client,
Adapter: adapter,
}
}
func (s *HTTPOutboundIPStrategy) Name() string {
if s == nil || s.Adapter == nil {
return "http-outbound-ip"
}
return s.Adapter.Name()
}
func (s *HTTPOutboundIPStrategy) GetOutboundIP(ctx context.Context) (net.IP, error) {
if s == nil || s.Adapter == nil {
return nil, errors.New("outbound IP adapter is nil")
}
if ctx == nil {
ctx = context.Background()
}
client := s.Client
if client == nil {
client = &http.Client{Timeout: defaultOutboundIPLookupTimeout}
}
request, err := http.NewRequestWithContext(ctx, http.MethodGet, s.Adapter.Endpoint(), nil)
if err != nil {
return nil, fmt.Errorf("%s create request failed: %w", s.Name(), err)
}
response, err := client.Do(request)
if err != nil {
return nil, fmt.Errorf("%s request failed: %w", s.Name(), err)
}
defer response.Body.Close()
if response.StatusCode != http.StatusOK {
return nil, fmt.Errorf("%s returned non-200 status: %d %s", s.Name(), response.StatusCode, response.Status)
}
ip, err := s.Adapter.DecodeIP(response.Body)
if err != nil {
return nil, fmt.Errorf("%s decode response failed: %w", s.Name(), err)
}
if !iputil.IsPublic(ip) {
return nil, fmt.Errorf("%s returned non-public IP: %s", s.Name(), ip.String())
}
return ip, nil
}
type RealIPCCAdapter struct {
URL string
}
type realIPCCResponse struct {
IP string `json:"ip"`
}
func NewRealIPCCOutboundIPStrategy() *HTTPOutboundIPStrategy {
return NewHTTPOutboundIPStrategy(RealIPCCAdapter{}, nil)
}
func (a RealIPCCAdapter) Name() string {
return "realip.cc"
}
func (a RealIPCCAdapter) Endpoint() string {
if strings.TrimSpace(a.URL) != "" {
return strings.TrimSpace(a.URL)
}
return "https://realip.cc"
}
func (a RealIPCCAdapter) DecodeIP(reader io.Reader) (net.IP, error) {
var payload realIPCCResponse
if err := json.NewDecoder(reader).Decode(&payload); err != nil {
return nil, err
}
ip := net.ParseIP(strings.TrimSpace(payload.IP))
if ip == nil {
return nil, fmt.Errorf("invalid IP %q", payload.IP)
}
if ipv4 := ip.To4(); ipv4 != nil {
return ipv4, nil
}
return ip, nil
}
func DefaultOutboundIPStrategies() []OutboundIPStrategy {
return []OutboundIPStrategy{
NewRealIPCCOutboundIPStrategy(),
}
}
func GetOutboundIP(ctx context.Context, strategies ...OutboundIPStrategy) (net.IP, error) {
if len(strategies) == 0 {
strategies = DefaultOutboundIPStrategies()
}
var errs []error
for _, strategy := range strategies {
if strategy == nil {
continue
}
ip, err := strategy.GetOutboundIP(ctx)
if err == nil && ip != nil {
return ip, nil
}
if err != nil {
errs = append(errs, fmt.Errorf("%s: %w", strategy.Name(), err))
}
}
if len(errs) == 0 {
return nil, errors.New("no outbound IP lookup strategy configured")
}
return nil, errors.Join(errs...)
}
+82
View File
@@ -0,0 +1,82 @@
package geoip
import (
"context"
"errors"
"net"
"net/http"
"net/http/httptest"
"strings"
"testing"
)
type fakeOutboundIPStrategy struct {
name string
ip net.IP
err error
}
func (f fakeOutboundIPStrategy) Name() string {
return f.name
}
func (f fakeOutboundIPStrategy) GetOutboundIP(ctx context.Context) (net.IP, error) {
return f.ip, f.err
}
func TestRealIPCCAdapterDecodeIP(t *testing.T) {
ip, err := RealIPCCAdapter{}.DecodeIP(strings.NewReader(`{"ip":"8.8.8.8","country":"United States"}`))
if err != nil {
t.Fatalf("DecodeIP failed: %v", err)
}
if ip.String() != "8.8.8.8" {
t.Fatalf("unexpected IP: %s", ip.String())
}
}
func TestHTTPOutboundIPStrategyUsesAdapter(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodGet {
t.Fatalf("unexpected method: %s", r.Method)
}
w.Header().Set("Content-Type", "application/json")
_, _ = w.Write([]byte(`{"ip":"8.8.4.4"}`))
}))
defer server.Close()
strategy := NewHTTPOutboundIPStrategy(RealIPCCAdapter{URL: server.URL}, server.Client())
ip, err := strategy.GetOutboundIP(context.Background())
if err != nil {
t.Fatalf("GetOutboundIP failed: %v", err)
}
if ip.String() != "8.8.4.4" {
t.Fatalf("unexpected outbound IP: %s", ip.String())
}
}
func TestGetOutboundIPFallsBackToNextStrategy(t *testing.T) {
ip, err := GetOutboundIP(
context.Background(),
fakeOutboundIPStrategy{name: "first", err: errors.New("temporary failure")},
fakeOutboundIPStrategy{name: "second", ip: net.ParseIP("1.1.1.1")},
)
if err != nil {
t.Fatalf("GetOutboundIP failed: %v", err)
}
if ip.String() != "1.1.1.1" {
t.Fatalf("unexpected outbound IP: %s", ip.String())
}
}
func TestHTTPOutboundIPStrategyRejectsPrivateIP(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "application/json")
_, _ = w.Write([]byte(`{"ip":"172.17.0.2"}`))
}))
defer server.Close()
strategy := NewHTTPOutboundIPStrategy(RealIPCCAdapter{URL: server.URL}, server.Client())
if _, err := strategy.GetOutboundIP(context.Background()); err == nil {
t.Fatal("expected private IP to be rejected")
}
}
+66
View File
@@ -0,0 +1,66 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package httppool manages shared, optimized HTTP transports to reuse TCP connections.
package httppool
import (
"crypto/tls"
"net"
"net/http"
"sync"
"time"
"go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp"
)
const (
dialTimeout = 30 * time.Second
dialKeepAlive = 30 * time.Second
maxIdleConns = 200
maxIdleConnsPerHost = 32
idleConnTimeout = 90 * time.Second
tlsHandshakeTimeout = 10 * time.Second
expectContinueTimeout = 1 * time.Second
tlsSessionCacheSize = 100
)
var (
defaultTransport http.RoundTripper
once sync.Once
)
// DefaultTransport returns a globally shared, optimized http.RoundTripper
// with OTel instrumentation. It maintains a pool of idle TCP connections
// across hosts.
func DefaultTransport() http.RoundTripper {
once.Do(func() {
transport := &http.Transport{
Proxy: http.ProxyFromEnvironment,
DialContext: (&net.Dialer{
Timeout: dialTimeout,
KeepAlive: dialKeepAlive,
}).DialContext,
ForceAttemptHTTP2: true,
MaxIdleConns: maxIdleConns,
MaxIdleConnsPerHost: maxIdleConnsPerHost,
IdleConnTimeout: idleConnTimeout,
TLSHandshakeTimeout: tlsHandshakeTimeout,
ExpectContinueTimeout: expectContinueTimeout,
TLSClientConfig: &tls.Config{
ClientSessionCache: tls.NewLRUClientSessionCache(tlsSessionCacheSize),
},
}
defaultTransport = otelhttp.NewTransport(transport)
})
return defaultTransport
}
// NewClient returns a new http.Client that shares the global connection pool
// but has its own timeout configuration.
func NewClient(timeout time.Duration) *http.Client {
return &http.Client{
Timeout: timeout,
Transport: DefaultTransport(),
}
}
+37
View File
@@ -0,0 +1,37 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package httppool
import (
"testing"
"time"
)
func TestDefaultTransport(t *testing.T) {
tr1 := DefaultTransport()
if tr1 == nil {
t.Fatal("DefaultTransport() returned nil")
}
tr2 := DefaultTransport()
if tr1 != tr2 {
t.Error("DefaultTransport() did not return a singleton instance")
}
}
func TestNewClient(t *testing.T) {
timeout := 15 * time.Second
client := NewClient(timeout)
if client == nil {
t.Fatal("NewClient() returned nil")
}
if client.Timeout != timeout {
t.Errorf("NewClient() timeout = %v, want %v", client.Timeout, timeout)
}
if client.Transport != DefaultTransport() {
t.Error("NewClient() is not configured with the default transport")
}
}
+9
View File
@@ -0,0 +1,9 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package logger 提供结构化日志封装
package logger
const (
errCreateLogFileDirFailed = "[Logger] create log file dir err: %w"
)
+102
View File
@@ -0,0 +1,102 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package logger
import (
"context"
"fmt"
"log"
"github.com/uptrace/opentelemetry-go-extra/otelzap"
"go.uber.org/zap"
"go.uber.org/zap/zapcore"
)
// Config represents the logging configuration.
type Config struct {
Level string
Format string
Output string
FilePath string
MaxSize int
MaxAge int
MaxBackups int
Compress bool
}
var logger *otelzap.Logger
// ringBufferCapacity 环形缓冲区容量
const ringBufferCapacity = 5000
// GlobalRingBuffer 全局日志环形缓冲区,供 Admin 日志查询和 WebSocket 推送使用
var GlobalRingBuffer *LogRingBuffer
func doInit(cfg Config) {
logWriter, err := getLogWriterForConfig(cfg)
if err != nil {
log.Fatalf("[Logger] get log writer err: %v\n", err)
}
// 初始化 ring buffer(保留最近 5000 行日志),如果是多次调用 Init,不需要重复创建 GlobalRingBuffer
if GlobalRingBuffer == nil {
GlobalRingBuffer = NewLogRingBuffer(ringBufferCapacity)
}
// 使用 multi writer 同时写入原始输出和 ring buffer
multiWriter := zapcore.NewMultiWriteSyncer(
logWriter,
zapcore.AddSync(GlobalRingBuffer),
)
zapLogger := zap.New(
zapcore.NewCore(getEncoderForConfig(cfg), multiWriter, getLogLevelForConfig(cfg)),
zap.AddCaller(),
zap.AddCallerSkip(1),
)
logger = otelzap.New(
zapLogger,
otelzap.WithMinLevel(zapLogger.Level()),
)
}
func init() {
// 默认使用 console stdout INFO 日志输出,避免在 Init 前或测试中发生空指针崩溃
defaultCfg := Config{
Level: "info",
Format: "console",
Output: "stdout",
}
doInit(defaultCfg)
}
// Init initializes the logger with a custom configuration.
func Init(cfg Config) {
doInit(cfg)
}
// DebugF 输出 Debug 级别日志
func DebugF(ctx context.Context, format string, args ...interface{}) {
msg := fmt.Sprintf(format, args...)
logger.Ctx(ctx).Debug(msg, getTraceIDFields(ctx)...)
}
// InfoF 输出 Info 级别日志
func InfoF(ctx context.Context, format string, args ...interface{}) {
msg := fmt.Sprintf(format, args...)
logger.Ctx(ctx).Info(msg, getTraceIDFields(ctx)...)
}
// WarnF 输出 Warn 级别日志
func WarnF(ctx context.Context, format string, args ...interface{}) {
msg := fmt.Sprintf(format, args...)
logger.Ctx(ctx).Warn(msg, getTraceIDFields(ctx)...)
}
// ErrorF 输出 Error 级别日志
func ErrorF(ctx context.Context, format string, args ...interface{}) {
msg := fmt.Sprintf(format, args...)
logger.Ctx(ctx).Error(msg, getTraceIDFields(ctx)...)
}
+175
View File
@@ -0,0 +1,175 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package logger
import (
"io"
"sync"
)
// LogEntry 日志条目,对应 ring buffer 中的一行日志
type LogEntry struct {
Index int `json:"index"` // 全局递增序号
Data string `json:"data"` // 一行日志原文(含换行符)
}
// LogRingBuffer 固定容量的环形缓冲区,存储最近的日志行
// 支持:追加日志、按 cursor 分页查询、订阅实时推送
type LogRingBuffer struct {
mu sync.RWMutex
entries []LogEntry
cap int
head int // 下一条写入的位置
count int // 当前条目数
seq int // 全局递增序号
subscribers map[chan LogEntry]struct{}
subMu sync.RWMutex
}
// NewLogRingBuffer 创建指定容量的日志环形缓冲区
func NewLogRingBuffer(capacity int) *LogRingBuffer {
return &LogRingBuffer{
entries: make([]LogEntry, capacity),
cap: capacity,
subscribers: make(map[chan LogEntry]struct{}),
}
}
// Write 实现 io.Writer 接口,供 zapcore.WriteSyncer 调用
// 按 '\n' 分割为独立行写入 ring buffer
func (r *LogRingBuffer) Write(p []byte) (int, error) {
if len(p) == 0 {
return 0, nil
}
data := string(p)
start := 0
for i := 0; i < len(data); i++ {
if data[i] == '\n' {
line := data[start:i]
start = i + 1
if len(line) > 0 {
r.appendLine(line)
}
}
}
// 处理最后一行(没有换行符结尾的情况)
if start < len(data) && len(data[start:]) > 0 {
r.appendLine(data[start:])
}
return len(p), nil
}
// Sync 实现 zapcore.WriteSyncer 接口
func (r *LogRingBuffer) Sync() error {
return nil
}
// appendLine 追加一行日志到 ring buffer 并通知订阅者
func (r *LogRingBuffer) appendLine(line string) {
r.mu.Lock()
entry := LogEntry{
Index: r.seq,
Data: line,
}
r.entries[r.head] = entry
r.head = (r.head + 1) % r.cap
if r.count < r.cap {
r.count++
}
r.seq++
r.mu.Unlock()
// 异步通知订阅者
r.subMu.RLock()
for ch := range r.subscribers {
select {
case ch <- entry:
default:
// 订阅者消费太慢,丢弃(避免阻塞日志写入)
}
}
r.subMu.RUnlock()
}
// Query 查询历史日志
// cursor=0 表示查询最新日志,cursor>0 表示查询 index < cursor 的更早日志
// limit 为返回条数上限
// 返回日志条目(按 index 升序)和是否有更早的日志
func (r *LogRingBuffer) Query(cursor int, limit int) ([]LogEntry, bool) {
r.mu.RLock()
defer r.mu.RUnlock()
if r.count == 0 {
return nil, false
}
// 计算 ring buffer 中有效条目的范围
// oldest index in ring: head - count (wrapping)
oldestPos := (r.head - r.count + r.cap) % r.cap
// 将 ring buffer 中的有效条目按顺序收集
ordered := make([]LogEntry, 0, r.count)
for i := 0; i < r.count; i++ {
pos := (oldestPos + i) % r.cap
ordered = append(ordered, r.entries[pos])
}
if cursor == 0 {
// 查询最新日志:返回最后 limit 条
if len(ordered) <= limit {
return ordered, false
}
return ordered[len(ordered)-limit:], true
}
// 查询 index < cursor 的更早日志
// 找到 index < cursor 的条目
var cut int
for cut = len(ordered); cut > 0; cut-- {
if ordered[cut-1].Index < cursor {
break
}
}
if cut == 0 {
return nil, false
}
// 返回 cut 之前的最后 limit 条
start := cut - limit
if start < 0 {
start = 0
}
hasMore := start > 0
return ordered[start:cut], hasMore
}
// subscribeChanSize 订阅者 channel 缓冲区大小
const subscribeChanSize = 64
// Subscribe 订阅实时日志推送
// 返回一个 channel,调用者应 defer Unsubscribe
func (r *LogRingBuffer) Subscribe() chan LogEntry {
ch := make(chan LogEntry, subscribeChanSize)
r.subMu.Lock()
r.subscribers[ch] = struct{}{}
r.subMu.Unlock()
return ch
}
// Unsubscribe 取消订阅
func (r *LogRingBuffer) Unsubscribe(ch chan LogEntry) {
r.subMu.Lock()
delete(r.subscribers, ch)
r.subMu.Unlock()
close(ch)
}
// 确保 LogRingBuffer 实现 io.Writer 接口
var _ io.Writer = (*LogRingBuffer)(nil)
+192
View File
@@ -0,0 +1,192 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package logger
import (
"testing"
"github.com/stretchr/testify/assert"
)
func TestLogRingBuffer_WriteAndQuery(t *testing.T) {
rb := NewLogRingBuffer(5)
// Write some logs
_, _ = rb.Write([]byte("line1\nline2\nline3\n"))
entries, hasMore := rb.Query(0, 10)
assert.False(t, hasMore)
assert.Equal(t, 3, len(entries))
assert.Equal(t, "line1", entries[0].Data)
assert.Equal(t, "line2", entries[1].Data)
assert.Equal(t, "line3", entries[2].Data)
assert.Equal(t, 0, entries[0].Index)
assert.Equal(t, 1, entries[1].Index)
assert.Equal(t, 2, entries[2].Index)
}
func TestLogRingBuffer_CapacityOverflow(t *testing.T) {
rb := NewLogRingBuffer(3)
_, _ = rb.Write([]byte("a\nb\nc\nd\ne\n"))
entries, hasMore := rb.Query(0, 10)
assert.False(t, hasMore)
assert.Equal(t, 3, len(entries))
assert.Equal(t, "c", entries[0].Data)
assert.Equal(t, "d", entries[1].Data)
assert.Equal(t, "e", entries[2].Data)
}
func TestLogRingBuffer_QueryLatest(t *testing.T) {
rb := NewLogRingBuffer(10)
_, _ = rb.Write([]byte("a\nb\nc\nd\ne\n"))
// Query latest 2
entries, hasMore := rb.Query(0, 2)
assert.True(t, hasMore)
assert.Equal(t, 2, len(entries))
assert.Equal(t, "d", entries[0].Data)
assert.Equal(t, "e", entries[1].Data)
}
func TestLogRingBuffer_QueryByCursor(t *testing.T) {
rb := NewLogRingBuffer(10)
_, _ = rb.Write([]byte("a\nb\nc\nd\ne\n"))
// First get all to find indices
all, _ := rb.Query(0, 10)
assert.Equal(t, 5, len(all))
// Query entries before index 3
entries, hasMore := rb.Query(3, 10)
assert.False(t, hasMore)
assert.Equal(t, 3, len(entries))
assert.Equal(t, "a", entries[0].Data)
assert.Equal(t, "b", entries[1].Data)
assert.Equal(t, "c", entries[2].Data)
}
func TestLogRingBuffer_QueryByCursorWithLimit(t *testing.T) {
rb := NewLogRingBuffer(10)
_, _ = rb.Write([]byte("a\nb\nc\nd\ne\n"))
// Query 2 entries before index 4
entries, hasMore := rb.Query(4, 2)
assert.True(t, hasMore)
assert.Equal(t, 2, len(entries))
assert.Equal(t, "c", entries[0].Data)
assert.Equal(t, "d", entries[1].Data)
}
func TestLogRingBuffer_QueryEmpty(t *testing.T) {
rb := NewLogRingBuffer(5)
entries, hasMore := rb.Query(0, 10)
assert.False(t, hasMore)
assert.Nil(t, entries)
}
func TestLogRingBuffer_QueryNonExistentCursor(t *testing.T) {
rb := NewLogRingBuffer(5)
_, _ = rb.Write([]byte("a\nb\n"))
entries, hasMore := rb.Query(999, 10)
assert.False(t, hasMore)
assert.Equal(t, 2, len(entries))
assert.Equal(t, "a", entries[0].Data)
assert.Equal(t, "b", entries[1].Data)
}
func TestLogRingBuffer_Subscribe(t *testing.T) {
rb := NewLogRingBuffer(5)
ch := rb.Subscribe()
defer rb.Unsubscribe(ch)
_, _ = rb.Write([]byte("hello\n"))
entry := <-ch
assert.Equal(t, "hello", entry.Data)
assert.Equal(t, 0, entry.Index)
}
func TestLogRingBuffer_SubscribeMultiple(t *testing.T) {
rb := NewLogRingBuffer(5)
ch1 := rb.Subscribe()
defer rb.Unsubscribe(ch1)
ch2 := rb.Subscribe()
defer rb.Unsubscribe(ch2)
_, _ = rb.Write([]byte("msg\n"))
e1 := <-ch1
e2 := <-ch2
assert.Equal(t, "msg", e1.Data)
assert.Equal(t, "msg", e2.Data)
}
func TestLogRingBuffer_WriteNoNewline(t *testing.T) {
rb := NewLogRingBuffer(5)
_, _ = rb.Write([]byte("partial"))
entries, _ := rb.Query(0, 10)
assert.Equal(t, 1, len(entries))
assert.Equal(t, "partial", entries[0].Data)
}
func TestLogRingBuffer_WriteEmpty(t *testing.T) {
rb := NewLogRingBuffer(5)
n, err := rb.Write([]byte(""))
assert.Equal(t, 0, n)
assert.NoError(t, err)
entries, _ := rb.Query(0, 10)
assert.Nil(t, entries)
}
func TestLogRingBuffer_QueryAfterOverflow(t *testing.T) {
rb := NewLogRingBuffer(3)
_, _ = rb.Write([]byte("1\n2\n3\n4\n5\n6\n7\n"))
entries, hasMore := rb.Query(0, 10)
assert.False(t, hasMore)
assert.Equal(t, 3, len(entries))
assert.Equal(t, "5", entries[0].Data)
assert.Equal(t, "6", entries[1].Data)
assert.Equal(t, "7", entries[2].Data)
// Query by cursor - index 4 is "5", so cursor=4 should return index < 4
older, hasMore2 := rb.Query(4, 10)
assert.False(t, hasMore2)
assert.Nil(t, older)
}
func TestLogRingBuffer_NextCursor(t *testing.T) {
rb := NewLogRingBuffer(10)
_, _ = rb.Write([]byte("a\nb\nc\nd\ne\n"))
// Query latest 2, should return next_cursor pointing to first returned entry
entries, _ := rb.Query(0, 2)
assert.Equal(t, 2, len(entries))
// entries[0].Index = 3 ("d"), entries[1].Index = 4 ("e")
assert.Equal(t, 3, entries[0].Index)
// Now use that index as cursor to get older entries
older, hasMore := rb.Query(entries[0].Index, 10)
assert.False(t, hasMore)
assert.Equal(t, 3, len(older))
assert.Equal(t, "a", older[0].Data)
assert.Equal(t, "b", older[1].Data)
assert.Equal(t, "c", older[2].Data)
}
+99
View File
@@ -0,0 +1,99 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package logger
import (
"context"
"fmt"
"log"
"os"
"path/filepath"
"go.opentelemetry.io/otel/trace"
"go.uber.org/zap"
"go.uber.org/zap/zapcore"
"gopkg.in/natefinch/lumberjack.v2"
)
// logDirPerm 日志目录权限
const logDirPerm = 0750
func getLogWriterForConfig(cfg Config) (zapcore.WriteSyncer, error) {
if cfg.Output == "file" {
// 初始化日志目录
logPath := cfg.FilePath
logDir := filepath.Dir(logPath)
if err := os.MkdirAll(logDir, logDirPerm); err != nil {
return nil, fmt.Errorf(errCreateLogFileDirFailed, err)
}
// 配置日志轮转
logOutput := &lumberjack.Logger{
Filename: logPath,
MaxSize: cfg.MaxSize,
MaxBackups: cfg.MaxBackups,
MaxAge: cfg.MaxAge,
Compress: cfg.Compress,
}
return zapcore.AddSync(logOutput), nil
}
return zapcore.AddSync(os.Stdout), nil
}
// getEncoderForConfig 获取日志编码器
func getEncoderForConfig(cfg Config) zapcore.Encoder {
// 编码器配置
encoderConfig := zapcore.EncoderConfig{
TimeKey: "time",
LevelKey: "level",
NameKey: "logger",
CallerKey: "caller",
MessageKey: "msg",
StacktraceKey: "stacktrace",
LineEnding: zapcore.DefaultLineEnding,
EncodeLevel: zapcore.LowercaseLevelEncoder,
EncodeTime: zapcore.ISO8601TimeEncoder,
EncodeDuration: zapcore.SecondsDurationEncoder,
EncodeCaller: zapcore.ShortCallerEncoder,
}
if cfg.Format == "json" {
return zapcore.NewJSONEncoder(encoderConfig)
}
return zapcore.NewConsoleEncoder(encoderConfig)
}
// getLogLevelForConfig 获取日志级别
func getLogLevelForConfig(cfg Config) zapcore.Level {
level := cfg.Level
switch level {
case "debug":
return zapcore.DebugLevel
case "info":
return zapcore.InfoLevel
case "warn":
return zapcore.WarnLevel
case "error":
return zapcore.ErrorLevel
default:
log.Printf("[Logger] invalid log level: %s, defaulting to info\n", level)
return zapcore.InfoLevel
}
}
func getTraceIDFields(ctx context.Context) []zap.Field {
span := trace.SpanFromContext(ctx)
spanContext := span.SpanContext()
if !spanContext.IsValid() {
return nil
}
return []zap.Field{
zap.String("traceID", spanContext.TraceID().String()),
zap.String("spanID", spanContext.SpanID().String()),
}
}
+16
View File
@@ -0,0 +1,16 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package mail 提供 SMTP 邮件发送功能。
package mail
const (
errDialTLSFailed = "dial tls failed: %w"
errSMTPClientCreationFailed = "smtp client creation failed: %w"
errSMTPAuthFailed = "smtp auth failed: %w"
errSMTPMailCommandFailed = "smtp mail command failed: %w"
errSMTPRcptCommandFailed = "smtp rcpt command failed: %w"
errSMTPDataCommandFailed = "smtp data command failed: %w"
errSMTPWritingBodyFailed = "smtp writing body failed: %w" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
errSendMailFailed = "send mail failed: %w"
)
+241
View File
@@ -0,0 +1,241 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package mail
import (
"bytes"
"context"
"crypto/tls"
"fmt"
"net"
"net/smtp"
"strconv"
"time"
)
const (
smtpSSLPort = 465 // SMTP SSL 端口
smtpDialTimeout = 5 * time.Second // SMTP 连接超时
smtpSessionDeadline = 10 * time.Second // SMTP 会话截止时间
)
// Config represents SMTP mail configuration
type Config struct {
Host string
Port int
Username string
Password string
}
// SendMail sends an HTML email using the provided config and message details
func SendMail(ctx context.Context, cfg Config, to string, subject, body string) error {
return SendMailHTML(ctx, cfg, to, subject, body)
}
// SendMailHTML sends an HTML format email
func SendMailHTML(ctx context.Context, cfg Config, to string, subject, body string) error {
addr := net.JoinHostPort(cfg.Host, strconv.Itoa(cfg.Port))
// Header & MIME settings for HTML email
header := make(map[string]string)
header["From"] = cfg.Username
header["To"] = to
header["Subject"] = subject
header["MIME-Version"] = "1.0"
header["Content-Type"] = "text/html; charset=UTF-8"
message := ""
for k, v := range header {
message += fmt.Sprintf("%s: %s\r\n", k, v)
}
message += "\r\n" + body
auth := smtp.PlainAuth("", cfg.Username, cfg.Password, cfg.Host)
// If using SSL port 465, we connection via TLS dial
if cfg.Port == smtpSSLPort {
return sendMailViaSSL(ctx, addr, auth, cfg, to, message)
}
// For standard port (587 / 25), use smtp.SendMail directly (handles STARTTLS automatically if server supports it)
err := smtp.SendMail(addr, auth, cfg.Username, []string{to}, []byte(message))
if err != nil {
return fmt.Errorf(errSendMailFailed, err)
}
return nil
}
// sendMailViaSSL 通过 TLS 直接连接 SMTP SSL 端口发送邮件
func sendMailViaSSL(ctx context.Context, addr string, auth smtp.Auth, cfg Config, to, message string) error {
tlsConfig := &tls.Config{
InsecureSkipVerify: true, //nolint:gosec // SMTP servers might use self-signed certificates
ServerName: cfg.Host,
}
dialer := &net.Dialer{Timeout: smtpDialTimeout}
tlsDialer := &tls.Dialer{
NetDialer: dialer,
Config: tlsConfig,
}
conn, err := tlsDialer.DialContext(ctx, "tcp", addr)
if err != nil {
return fmt.Errorf(errDialTLSFailed, err)
}
defer func() { _ = conn.Close() }()
_ = conn.SetDeadline(time.Now().Add(smtpSessionDeadline))
client, err := smtp.NewClient(conn, cfg.Host)
if err != nil {
return fmt.Errorf(errSMTPClientCreationFailed, err)
}
defer func() { _ = client.Close() }()
if err = client.Auth(auth); err != nil {
return fmt.Errorf(errSMTPAuthFailed, err)
}
if err = client.Mail(cfg.Username); err != nil {
return fmt.Errorf(errSMTPMailCommandFailed, err)
}
if err = client.Rcpt(to); err != nil {
return fmt.Errorf(errSMTPRcptCommandFailed, err)
}
w, err := client.Data()
if err != nil {
return fmt.Errorf(errSMTPDataCommandFailed, err)
}
defer func() { _ = w.Close() }()
_, err = w.Write([]byte(message))
if err != nil {
return fmt.Errorf(errSMTPWritingBodyFailed, err)
}
return nil
}
// SendMailWithLog sends a test email and records a detailed SMTP connection log
func SendMailWithLog(ctx context.Context, cfg Config, to string, subject, body string) (string, error) {
var logBuf bytes.Buffer
logLine := func(dir string, format string, args ...interface{}) {
fmt.Fprintf(&logBuf, "[%s] %s\n", dir, fmt.Sprintf(format, args...))
}
addr := net.JoinHostPort(cfg.Host, strconv.Itoa(cfg.Port))
logLine("System", "Connecting to %s...", addr)
var conn net.Conn
var err error
dialer := &net.Dialer{Timeout: smtpDialTimeout}
if cfg.Port == smtpSSLPort {
tlsConfig := &tls.Config{
InsecureSkipVerify: true, //nolint:gosec // SMTP servers might use self-signed certificates
ServerName: cfg.Host,
}
tlsDialer := &tls.Dialer{
NetDialer: dialer,
Config: tlsConfig,
}
conn, err = tlsDialer.DialContext(ctx, "tcp", addr)
} else {
conn, err = dialer.DialContext(ctx, "tcp", addr)
}
if err != nil {
logLine("Error", "Connection failed: %v", err)
return logBuf.String(), err
}
defer func() { _ = conn.Close() }()
logLine("System", "Connected successfully.")
// Set a 10-second session deadline for read/write operations
_ = conn.SetDeadline(time.Now().Add(smtpSessionDeadline))
client, err := smtp.NewClient(conn, cfg.Host)
if err != nil {
logLine("Error", "SMTP client handshake failed: %v", err)
return logBuf.String(), err
}
defer func() { _ = client.Close() }()
// If not 465, support STARTTLS if available
if cfg.Port != smtpSSLPort {
if ok, _ := client.Extension("STARTTLS"); ok {
logLine("C", "STARTTLS")
tlsConfig := &tls.Config{
InsecureSkipVerify: true, //nolint:gosec // SMTP servers might use self-signed certificates
ServerName: cfg.Host,
}
if err = client.StartTLS(tlsConfig); err != nil {
logLine("Error", "STARTTLS failed: %v", err)
return logBuf.String(), err
}
logLine("S", "220 Ready to start TLS")
}
}
// Authentication
if cfg.Username != "" && cfg.Password != "" {
auth := smtp.PlainAuth("", cfg.Username, cfg.Password, cfg.Host)
logLine("C", "AUTH PLAIN **********")
if err = client.Auth(auth); err != nil {
logLine("Error", "Authentication failed: %v", err)
return logBuf.String(), err
}
logLine("S", "235 Authentication successful")
}
// Mail command
logLine("C", "MAIL FROM:<%s>", cfg.Username)
if err = client.Mail(cfg.Username); err != nil {
logLine("Error", "MAIL FROM command failed: %v", err)
return logBuf.String(), err
}
logLine("S", "250 OK")
// Rcpt command
logLine("C", "RCPT TO:<%s>", to)
if err = client.Rcpt(to); err != nil {
logLine("Error", "RCPT TO command failed: %v", err)
return logBuf.String(), err
}
logLine("S", "250 OK")
// Data command
logLine("C", "DATA")
w, err := client.Data()
if err != nil {
logLine("Error", "DATA command failed: %v", err)
return logBuf.String(), err
}
logLine("S", "354 Start mail input")
// Header & MIME settings for HTML email
header := make(map[string]string)
header["From"] = cfg.Username
header["To"] = to
header["Subject"] = subject
header["MIME-Version"] = "1.0"
header["Content-Type"] = "text/html; charset=UTF-8"
message := ""
for k, v := range header {
message += fmt.Sprintf("%s: %s\r\n", k, v)
}
message += "\r\n" + body
logLine("System", "Sending message body...")
if _, err = w.Write([]byte(message)); err != nil {
_ = w.Close()
logLine("Error", "Writing message body failed: %v", err)
return logBuf.String(), err
}
_ = w.Close()
logLine("S", "250 OK")
logLine("C", "QUIT")
_ = client.Quit()
logLine("System", "Mail sent successfully!")
return logBuf.String(), nil
}
+92
View File
@@ -0,0 +1,92 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package mail
import (
"bufio"
"context"
"net"
"net/textproto"
"testing"
)
func TestSendMailMock(t *testing.T) {
// Start a mock SMTP server
l, err := net.Listen("tcp", "127.0.0.1:0")
if err != nil {
t.Fatalf("failed to start mock smtp server: %v", err)
}
defer func() { _ = l.Close() }()
port := l.Addr().(*net.TCPAddr).Port
go func() {
conn, err := l.Accept()
if err != nil {
return
}
defer func() { _ = conn.Close() }()
writer := bufio.NewWriter(conn)
reader := bufio.NewReader(conn)
tp := textproto.NewReader(reader)
// 220 Ready
_, _ = writer.WriteString("220 mock.smtp.com SMTP Ready\r\n")
_ = writer.Flush()
// Read HELO/EHLO
_, _ = tp.ReadLine()
_, _ = writer.WriteString("250-mock.smtp.com\r\n250 AUTH PLAIN\r\n")
_ = writer.Flush()
// Read AUTH PLAIN
_, _ = tp.ReadLine()
_, _ = writer.WriteString("235 Authentication successful\r\n")
_ = writer.Flush()
// Read MAIL FROM
_, _ = tp.ReadLine()
_, _ = writer.WriteString("250 OK\r\n")
_ = writer.Flush()
// Read RCPT TO
_, _ = tp.ReadLine()
_, _ = writer.WriteString("250 OK\r\n")
_ = writer.Flush()
// Read DATA
_, _ = tp.ReadLine()
_, _ = writer.WriteString("354 Start mail input\r\n")
_ = writer.Flush()
// Read body lines until dot
for {
line, err := tp.ReadLine()
if err != nil || line == "." {
break
}
}
_, _ = writer.WriteString("250 OK\r\n")
_ = writer.Flush()
// Read QUIT
_, _ = tp.ReadLine()
_, _ = writer.WriteString("221 Bye\r\n")
_ = writer.Flush()
}()
cfg := Config{
Host: "127.0.0.1",
Port: port,
Username: "test@example.com",
Password: "password",
}
err = SendMail(context.Background(), cfg, "recipient@example.com", "Test Subject", "<h1>Test Body</h1>")
if err != nil {
t.Errorf("failed to send mail: %v", err)
}
}
+148
View File
@@ -0,0 +1,148 @@
package protocol
type WSMessage struct {
Type string `json:"type"`
Payload any `json:"payload,omitempty"`
}
type AgentNodeSystemProfile struct {
Hostname string `json:"hostname"`
OSName string `json:"os_name"`
OSVersion string `json:"os_version"`
KernelVersion string `json:"kernel_version"`
Architecture string `json:"architecture"`
CPUModel string `json:"cpu_model"`
CPUCores int `json:"cpu_cores"`
TotalMemoryBytes int64 `json:"total_memory_bytes"`
TotalDiskBytes int64 `json:"total_disk_bytes"`
UptimeSeconds int64 `json:"uptime_seconds"`
ReportedAtUnix int64 `json:"reported_at_unix"`
}
type AgentNodeMetricSnapshot struct {
CapturedAtUnix int64 `json:"captured_at_unix"`
CPUUsagePercent float64 `json:"cpu_usage_percent"`
MemoryUsedBytes int64 `json:"memory_used_bytes"`
MemoryTotalBytes int64 `json:"memory_total_bytes"`
StorageUsedBytes int64 `json:"storage_used_bytes"`
StorageTotalBytes int64 `json:"storage_total_bytes"`
DiskReadBytes int64 `json:"disk_read_bytes"`
DiskWriteBytes int64 `json:"disk_write_bytes"`
NetworkRxBytes int64 `json:"network_rx_bytes"`
NetworkTxBytes int64 `json:"network_tx_bytes"`
}
type AgentNodeHealthEvent struct {
EventType string `json:"event_type"`
Severity string `json:"severity"`
Message string `json:"message"`
TriggeredAtUnix int64 `json:"triggered_at_unix"`
Metadata map[string]string `json:"metadata"`
}
type RelayProxyStat struct {
Name string `json:"name"`
Type string `json:"type"`
Status string `json:"status"`
ClientVersion string `json:"client_version"`
LastStartTime string `json:"last_start_time"`
LastCloseTime string `json:"last_close_time"`
ClientAddr string `json:"client_addr"`
}
type RelayHeartbeatPayload struct {
Version string `json:"version"`
ExtVersion string `json:"frp_version"`
RelayStatus string `json:"relay_status"`
FrpsConnCount int `json:"frps_connections"`
FrpsProxyCount int `json:"frps_proxy_count"`
FrpsClientCount int `json:"frps_client_count"`
FrpsProxies []RelayProxyStat `json:"frps_proxies,omitempty"`
Name string `json:"name"`
IP string `json:"ip"`
Profile *AgentNodeSystemProfile `json:"profile,omitempty"`
Snapshot *AgentNodeMetricSnapshot `json:"snapshot,omitempty"`
HealthEvents []AgentNodeHealthEvent `json:"health_events,omitempty"`
}
type RelayConfig struct {
BindPort int `json:"bind_port"`
VhostHTTPPort int `json:"vhost_http_port"`
AuthToken string `json:"auth_token"`
LogLevel string `json:"log_level"`
WebServerEnabled bool `json:"web_server_enabled"`
}
type RelaySettings struct {
HeartbeatInterval int `json:"heartbeat_interval"`
WebsocketUpgradeEnabled bool `json:"websocket_upgrade_enabled"`
AutoUpdate bool `json:"auto_update"`
UpdateRepo string `json:"update_repo"`
UpdateNow bool `json:"update_now"`
UpdateChannel string `json:"update_channel"`
UpdateTag string `json:"update_tag"`
}
type RelayHeartbeatResponse struct {
RelayConfig *RelayConfig `json:"relay_config"`
RelaySettings *RelaySettings `json:"relay_settings"`
}
type ActiveConfigMeta struct {
Version string `json:"version"`
Checksum string `json:"checksum"`
}
type FlaredConnectedRelay struct {
RelayNodeID string `json:"relay_node_id"`
Status string `json:"status"`
ProxyCount int `json:"proxy_count"`
}
type FlaredHeartbeatPayload struct {
ClientVersion string `json:"client_version"`
FrpVersion string `json:"frp_version"`
IP string `json:"ip"`
TunnelStatus string `json:"tunnel_status"`
ConnectedRelays []FlaredConnectedRelay `json:"connected_relays"`
CurrentVersion string `json:"current_version"`
CurrentChecksum string `json:"current_checksum"`
}
type FlaredHeartbeatResponse struct {
ActiveConfig *ActiveConfigMeta `json:"active_config"`
TunnelSettings *RelaySettings `json:"tunnel_settings"`
}
type FlaredTunnelConfigResponse struct {
Version string `json:"version"`
Checksum string `json:"checksum"`
Relays []FlaredRelayInfo `json:"relays"`
Proxies []FlaredProxyEntry `json:"proxies"`
}
type FlaredRelayInfo struct {
RelayNodeID string `json:"relay_node_id"`
Address string `json:"address"`
AuthToken string `json:"auth_token"`
ProxyURL string `json:"proxy_url"`
}
type FlaredProxyEntry struct {
Name string `json:"name"`
Type string `json:"type"`
LocalAddr string `json:"local_addr"`
LocalPort int `json:"local_port"`
CustomDomains []string `json:"custom_domains"`
}
type ApplyLogPayload struct {
NodeID string `json:"node_id"`
Version string `json:"version"`
Result string `json:"result"`
Message string `json:"message"`
Checksum string `json:"checksum"`
MainConfigChecksum string `json:"main_config_checksum"`
RouteConfigChecksum string `json:"route_config_checksum"`
SupportFileCount int `json:"support_file_count"`
}
+81
View File
@@ -0,0 +1,81 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package push
import (
"bytes"
"context"
"encoding/json"
"errors"
"fmt"
"net/http"
"strings"
"github.com/Rain-kl/Wavelet/pkg/httppool"
)
func init() {
Register("custom", &CustomPusher{})
}
// CustomPusher 自定义 Webhook 发送实现
type CustomPusher struct{}
// Send 发送自定义 webhook
func (p *CustomPusher) Send(ctx context.Context, cfg Config, _ string, body map[string]any, template string, _ map[string]any) error {
if cfg.URL == "" {
return errors.New("custom: URL is required")
}
var reqBody []byte
if template != "" {
// 替换模板中的 {{key}} 占位符
rendered := ParseTemplate(template, body)
reqBody = []byte(rendered)
} else {
// 兜底:直接把 body 转为 JSON 字符串发送
var err error
reqBody, err = json.Marshal(body)
if err != nil {
return fmt.Errorf("custom: marshal body failed: %w", err)
}
}
httpReq, err := http.NewRequestWithContext(ctx, http.MethodPost, cfg.URL, bytes.NewReader(reqBody))
if err != nil {
return fmt.Errorf("custom: create http request failed: %w", err)
}
httpReq.Header.Set("Content-Type", "application/json")
// 如果配置了 Key 且格式为 "HeaderName:HeaderValue",我们可以附加测试用 Header
if cfg.Key != "" && strings.Contains(cfg.Key, ":") {
parts := strings.SplitN(cfg.Key, ":", 2) //nolint:mnd
httpReq.Header.Set(strings.TrimSpace(parts[0]), strings.TrimSpace(parts[1]))
}
client := httppool.NewClient(defaultHTTPClientTimeout)
resp, err := client.Do(httpReq)
if err != nil {
return fmt.Errorf("custom: http request failed: %w", err)
}
defer func() { _ = resp.Body.Close() }()
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
return fmt.Errorf("custom: http status %s", resp.Status)
}
return nil
}
// ValidateConfig 校验自定义配置
func (p *CustomPusher) ValidateConfig(cfg Config) error {
if cfg.URL == "" {
return errors.New("webhook URL is required")
}
if !strings.HasPrefix(cfg.URL, "http://") && !strings.HasPrefix(cfg.URL, "https://") {
return errors.New("webhook URL must start with http:// or https://")
}
return nil
}
+109
View File
@@ -0,0 +1,109 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package push
import (
"context"
"errors"
"fmt"
"net"
"net/smtp"
"strings"
)
func init() {
Register("email", &EmailPusher{})
}
// EmailPusher 极简 SMTP 邮件推送实现 (静态、解耦)
type EmailPusher struct{}
// Send 发送邮件
func (p *EmailPusher) Send(ctx context.Context, cfg Config, target string, body map[string]any, _ string, ext map[string]any) error {
if cfg.URL == "" || cfg.Key == "" || cfg.Secret == "" {
return errors.New("email: SMTP configuration (url, key, secret) is incomplete")
}
if target == "" {
return errors.New("email: target email address is required")
}
title := defaultTitle
if t, ok := body["title"].(string); ok && t != "" {
title = t
}
content := ""
if c, ok := body["content"].(string); ok && c != "" {
content = c
} else {
// 自动格式化 map
var parts []string
for k, v := range body {
parts = append(parts, fmt.Sprintf("<p><b>%s</b>: %v</p>", k, v))
}
content = strings.Join(parts, "")
}
// 邮件头和体
from := cfg.Key
to := target
// 如果 ext 中指定了 from_name,我们在 From 头部包含它
fromName := "System Notification"
if ext != nil {
if fn, ok := ext["from_name"].(string); ok && fn != "" {
fromName = fn
}
}
subjectHeader := fmt.Sprintf("Subject: %s\r\n", title)
fromHeader := fmt.Sprintf("From: %s <%s>\r\n", fromName, from)
toHeader := fmt.Sprintf("To: %s\r\n", to)
mimeHeader := "MIME-version: 1.0;\r\nContent-Type: text/html; charset=\"UTF-8\";\r\n\r\n"
// 拼装完整的邮件报文
// 简单的 HTML 正文渲染
htmlBody := fmt.Sprintf(`<html><body><h2>%s</h2><div>%s</div></body></html>`, title, content)
msg := []byte(fromHeader + toHeader + subjectHeader + mimeHeader + htmlBody + "\r\n")
// 解析 Host 和 Port
host, port, err := net.SplitHostPort(cfg.URL)
if err != nil {
host = cfg.URL
port = "25" // 默认 SMTP 端口
}
auth := smtp.PlainAuth("", cfg.Key, cfg.Secret, host)
// 异步超时处理
errChan := make(chan error, 1)
go func() {
errChan <- smtp.SendMail(host+":"+port, auth, from, []string{to}, msg)
}()
select {
case <-ctx.Done():
return ctx.Err()
case err := <-errChan:
if err != nil {
return fmt.Errorf("email: send smtp mail failed: %w", err)
}
}
return nil
}
// ValidateConfig 校验邮件 SMTP 配置
func (p *EmailPusher) ValidateConfig(cfg Config) error {
if cfg.URL == "" {
return errors.New("SMTP host:port is required")
}
if cfg.Key == "" {
return errors.New("SMTP username is required")
}
if cfg.Secret == "" {
return errors.New("SMTP password is required")
}
return nil
}
+275
View File
@@ -0,0 +1,275 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package push
import (
"bytes"
"context"
"crypto/hmac"
"crypto/sha256"
"encoding/base64"
"encoding/json"
"errors"
"fmt"
"net/http"
"strconv"
"strings"
"time"
"github.com/Rain-kl/Wavelet/pkg/httppool"
)
func init() {
Register("lark", &LarkPusher{})
}
const (
msgTypeInteractive = "interactive"
)
// LarkPusher 飞书 Webhook 机器人推送实现
type LarkPusher struct{}
type larkTextContent struct {
Text string `json:"text"`
}
type larkCardHeaderTitle struct {
Content string `json:"content"`
Tag string `json:"tag"`
}
type larkCardHeader struct {
Template string `json:"template"` // "blue", "orange", "red" etc.
Title larkCardHeaderTitle `json:"title"`
}
type larkCardElementText struct {
Content string `json:"content"`
Tag string `json:"tag"` // "lark_md"
}
type larkCardElement struct {
Tag string `json:"tag"` // "div"
Text larkCardElementText `json:"text"`
}
type larkCardContent struct {
Header larkCardHeader `json:"header"`
Elements []larkCardElement `json:"elements"`
}
type larkMessageRequest struct {
MessageType string `json:"msg_type"`
Timestamp string `json:"timestamp,omitempty"`
Sign string `json:"sign,omitempty"`
Content larkTextContent `json:"content,omitempty"`
Card *larkCardContent `json:"card,omitempty"`
}
type larkMessageResponse struct {
Code int `json:"code"`
Msg string `json:"msg"`
}
// Send 执行飞书消息发送
//
//nolint:nestif,cyclop
func (p *LarkPusher) Send(ctx context.Context, cfg Config, _ string, body map[string]any, template string, _ map[string]any) error {
if cfg.URL == "" {
return errors.New("lark: URL is required")
}
var req larkMessageRequest
// 1. 如果有自定义模板,我们尝试进行解析
if template != "" {
rendered := ParseTemplate(template, body)
// 尝试解析原生的 Lark Card
var customCard larkCardContent
var rawMap map[string]any
_ = json.Unmarshal([]byte(rendered), &rawMap)
if rawMap != nil && rawMap["elements"] != nil {
// 如果包含 elements 字段,说明是用户定制的原生飞书卡片 JSON
if err := json.Unmarshal([]byte(rendered), &customCard); err == nil {
req.MessageType = msgTypeInteractive
req.Card = &customCard
} else {
req.MessageType = "text"
req.Content.Text = rendered
}
} else {
// 说明配置的是系统统一通知消息 of JSON 模板:{"title": "...", "content": "...", "level": "..."}
type larkNotificationMessage struct {
Title string `json:"title"`
Content string `json:"content"`
Level string `json:"level"`
}
var msg larkNotificationMessage
if err := json.Unmarshal([]byte(rendered), &msg); err == nil && (msg.Title != "" || msg.Content != "") {
title := msg.Title
if title == "" {
title = defaultTitle
}
content := msg.Content
level := strings.ToUpper(msg.Level)
if level == "" {
level = levelInfo
}
headerColor := "blue"
switch level {
case "IMPORTANT":
headerColor = "orange"
case "CRITICAL":
headerColor = "red"
}
req.MessageType = msgTypeInteractive
req.Card = &larkCardContent{
Header: larkCardHeader{
Template: headerColor,
Title: larkCardHeaderTitle{
Content: title,
Tag: "plain_text",
},
},
Elements: []larkCardElement{
{
Tag: "div",
Text: larkCardElementText{
Content: content,
Tag: "lark_md",
},
},
},
}
} else {
// 兜底:如果无法按 JSON 解析出结构化字段,当做普通文本发送
req.MessageType = "text"
req.Content.Text = rendered
}
}
} else {
// 2. 如果无模板,默认生成一个精美的飞书互动卡片
title := defaultTitle
if t, ok := body["title"].(string); ok && t != "" {
title = t
}
content := ""
if c, ok := body["content"].(string); ok && c != "" {
content = c
} else {
// 兜底:如果连 content 都没有,把 body 里的所有值拼成 markdown
var parts []string
for k, v := range body {
parts = append(parts, fmt.Sprintf("**%s**: %v", k, v))
}
content = strings.Join(parts, "\n")
}
level := levelInfo
if l, ok := body["level"].(string); ok && l != "" {
level = strings.ToUpper(l)
}
// 根据级别确定飞书卡片头部的背景色模板
headerColor := "blue"
switch level {
case "IMPORTANT":
headerColor = "orange"
case "CRITICAL":
headerColor = "red"
}
req.MessageType = msgTypeInteractive
req.Card = &larkCardContent{
Header: larkCardHeader{
Template: headerColor,
Title: larkCardHeaderTitle{
Content: title,
Tag: "plain_text",
},
},
Elements: []larkCardElement{
{
Tag: "div",
Text: larkCardElementText{
Content: content,
Tag: "lark_md",
},
},
},
}
}
// 3. 计算签名 (如果配置了 secret)
if cfg.Secret != "" {
timestamp := time.Now().Unix()
sign, err := larkSign(cfg.Secret, timestamp)
if err != nil {
return fmt.Errorf("lark: sign failed: %w", err)
}
req.Timestamp = strconv.FormatInt(timestamp, 10)
req.Sign = sign
}
jsonData, err := json.Marshal(req)
if err != nil {
return fmt.Errorf("lark: marshal request failed: %w", err)
}
// 4. 发送 POST 请求
httpReq, err := http.NewRequestWithContext(ctx, http.MethodPost, cfg.URL, bytes.NewBuffer(jsonData))
if err != nil {
return fmt.Errorf("lark: create http request failed: %w", err)
}
httpReq.Header.Set("Content-Type", "application/json")
client := httppool.NewClient(defaultHTTPClientTimeout)
resp, err := client.Do(httpReq)
if err != nil {
return fmt.Errorf("lark: http request failed: %w", err)
}
defer func() { _ = resp.Body.Close() }()
if resp.StatusCode != http.StatusOK {
return fmt.Errorf("lark: http status %s", resp.Status)
}
var res larkMessageResponse
if err := json.NewDecoder(resp.Body).Decode(&res); err != nil {
return fmt.Errorf("lark: decode response failed: %w", err)
}
if res.Code != 0 {
return fmt.Errorf("lark: send message failed, code %d: %s", res.Code, res.Msg)
}
return nil
}
// ValidateConfig 校验飞书配置
func (p *LarkPusher) ValidateConfig(cfg Config) error {
if cfg.URL == "" {
return errors.New("webhook URL is required")
}
if !strings.HasPrefix(cfg.URL, "http://") && !strings.HasPrefix(cfg.URL, "https://") {
return errors.New("webhook URL must start with http:// or https://")
}
return nil
}
func larkSign(secret string, timestamp int64) (string, error) {
stringToSign := fmt.Sprintf("%v", timestamp) + "\n" + secret
h := hmac.New(sha256.New, []byte(stringToSign))
_, err := h.Write(nil)
if err != nil {
return "", err
}
return base64.StdEncoding.EncodeToString(h.Sum(nil)), nil
}
+66
View File
@@ -0,0 +1,66 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package push 提供解耦的、无外部业务依赖 of 通知推送底层实现
package push
import (
"context"
"fmt"
"sync"
"time"
)
const (
defaultTitle = "系统通知"
levelInfo = "INFO"
defaultHTTPClientTimeout = 10 * time.Second
)
// Config 基础通知渠道配置
type Config struct {
Channel string `json:"channel"` // 渠道名称,例如 "lark", "custom", "email" 等,唯一标识
URL string `json:"url,omitempty"` // Webhook 地址或 SMTP 地址
Secret string `json:"secret,omitempty"` // 签名密钥或 SMTP 密码/Token
Key string `json:"key,omitempty"` // AppID 或 SMTP 用户名
Ext map[string]any `json:"ext,omitempty"` // 预留拓展 JSON 配置
}
// Pusher 通知推送渠道接口
type Pusher interface {
// Send 发送通知消息
// target: 发送目标 (如邮箱地址或特定用户标识;若为 bot 机器人此项为空)
// body: 消息体数据 (含默认字段如 title, content, level)
// template: 消息卡片/模板 JSON (可选)
// ext: 预留的单次发送拓展数据
Send(ctx context.Context, cfg Config, target string, body map[string]any, template string, ext map[string]any) error
// ValidateConfig 校验渠道配置合法性
ValidateConfig(cfg Config) error
}
var (
pushersMu sync.RWMutex
pushers = make(map[string]Pusher)
)
// Register 注册一个推送渠道实现
func Register(channelType string, pusher Pusher) {
pushersMu.Lock()
defer pushersMu.Unlock()
if pusher == nil {
panic("push: Register pusher is nil")
}
pushers[channelType] = pusher
}
// GetPusher 获取指定类型的推送渠道实现
func GetPusher(channelType string) (Pusher, error) {
pushersMu.RLock()
defer pushersMu.RUnlock()
pusher, ok := pushers[channelType]
if !ok {
return nil, fmt.Errorf("push: unknown channel type %q", channelType)
}
return pusher, nil
}
+158
View File
@@ -0,0 +1,158 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package push
import (
"bytes"
"context"
"encoding/json"
"errors"
"fmt"
"net/http"
"strings"
"github.com/Rain-kl/Wavelet/pkg/httppool"
)
func init() {
Register("telegram", &TelegramPusher{})
}
// TelegramPusher Telegram 机器人推送实现
type TelegramPusher struct{}
type telegramMessageRequest struct {
ChatID string `json:"chat_id"`
Text string `json:"text"`
ParseMode string `json:"parse_mode,omitempty"`
}
type telegramErrorResponse struct {
Ok bool `json:"ok"`
ErrorCode int `json:"error_code"`
Description string `json:"description"`
}
// Send 执行 Telegram 消息发送
//
//nolint:cyclop
func (p *TelegramPusher) Send(ctx context.Context, cfg Config, target string, body map[string]any, template string, _ map[string]any) error {
if cfg.Secret == "" {
return errors.New("telegram: Bot Token (Secret) is required")
}
chatID := target
if chatID == "" {
chatID = cfg.Key // Use default chat ID (Key) if target is blank
}
if chatID == "" {
return errors.New("telegram: chat_id (target or default Key) is required")
}
baseURL := cfg.URL
if baseURL == "" {
baseURL = "https://api.telegram.org"
}
baseURL = strings.TrimSuffix(baseURL, "/")
title := defaultTitle
if t, ok := body["title"].(string); ok && t != "" {
title = t
}
content := ""
if c, ok := body["content"].(string); ok && c != "" {
content = c
} else {
var parts []string
for k, v := range body {
parts = append(parts, fmt.Sprintf("<b>%s</b>: %v", k, v))
}
content = strings.Join(parts, "\n")
}
level := levelInfo
if l, ok := body["level"].(string); ok && l != "" {
level = strings.ToUpper(l)
}
var text string
if template != "" {
text = ParseTemplate(template, body)
} else {
text = fmt.Sprintf("<b>[%s] %s</b>\n\n%s", escapeHTML(level), escapeHTML(title), escapeHTML(content))
}
// Try sending with HTML parse mode
err := p.sendMessage(ctx, baseURL, cfg.Secret, chatID, text, "HTML")
if err != nil {
// Fallback: send as plain text without parse mode
plainText := text
if template == "" {
plainText = fmt.Sprintf("[%s] %s\n\n%s", level, title, content)
}
fallbackErr := p.sendMessage(ctx, baseURL, cfg.Secret, chatID, plainText, "")
if fallbackErr != nil {
return fmt.Errorf("telegram: send message failed (fallback also failed): %w (original HTML error: %v)", fallbackErr, err)
}
}
return nil
}
// ValidateConfig 校验 Telegram 配置
func (p *TelegramPusher) ValidateConfig(cfg Config) error {
if cfg.Secret == "" {
return errors.New("bot Token (Secret) is required")
}
if cfg.URL != "" {
if !strings.HasPrefix(cfg.URL, "http://") && !strings.HasPrefix(cfg.URL, "https://") {
return errors.New("API base URL must start with http:// or https://")
}
}
return nil
}
func (p *TelegramPusher) sendMessage(ctx context.Context, baseURL, token, chatID, text, parseMode string) error {
apiURL := fmt.Sprintf("%s/bot%s/sendMessage", baseURL, token)
reqPayload := telegramMessageRequest{
ChatID: chatID,
Text: text,
ParseMode: parseMode,
}
jsonData, err := json.Marshal(reqPayload)
if err != nil {
return fmt.Errorf("marshal request failed: %w", err)
}
httpReq, err := http.NewRequestWithContext(ctx, http.MethodPost, apiURL, bytes.NewBuffer(jsonData))
if err != nil {
return fmt.Errorf("create http request failed: %w", err)
}
httpReq.Header.Set("Content-Type", "application/json")
client := httppool.NewClient(defaultHTTPClientTimeout)
resp, err := client.Do(httpReq)
if err != nil {
return fmt.Errorf("http request failed: %w", err)
}
defer func() { _ = resp.Body.Close() }()
if resp.StatusCode != http.StatusOK {
var errRes telegramErrorResponse
if decodeErr := json.NewDecoder(resp.Body).Decode(&errRes); decodeErr == nil {
return fmt.Errorf("http status %d: %s", resp.StatusCode, errRes.Description)
}
return fmt.Errorf("http status %s", resp.Status)
}
return nil
}
func escapeHTML(s string) string {
s = strings.ReplaceAll(s, "&", "&amp;")
s = strings.ReplaceAll(s, "<", "&lt;")
s = strings.ReplaceAll(s, ">", "&gt;")
return s
}
+116
View File
@@ -0,0 +1,116 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package push
import (
"context"
"encoding/json"
"net/http"
"net/http/httptest"
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
func TestTelegramPusher_Send(t *testing.T) {
t.Run("successful send with HTML parse mode", func(t *testing.T) {
var receivedReq telegramMessageRequest
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
assert.Equal(t, "/botmy-token/sendMessage", r.URL.Path)
assert.Equal(t, http.MethodPost, r.Method)
assert.Equal(t, "application/json", r.Header.Get("Content-Type"))
err := json.NewDecoder(r.Body).Decode(&receivedReq)
require.NoError(t, err)
w.WriteHeader(http.StatusOK)
_, _ = w.Write([]byte(`{"ok": true}`))
}))
defer server.Close()
pusher := &TelegramPusher{}
cfg := Config{
Channel: "telegram",
URL: server.URL,
Secret: "my-token",
}
body := map[string]any{
"title": "Alert",
"content": "Host down",
"level": "CRITICAL",
}
err := pusher.Send(context.Background(), cfg, "123456", body, "", nil)
require.NoError(t, err)
assert.Equal(t, "123456", receivedReq.ChatID)
assert.Contains(t, receivedReq.Text, "[CRITICAL] Alert")
assert.Contains(t, receivedReq.Text, "Host down")
assert.Equal(t, "HTML", receivedReq.ParseMode)
})
t.Run("fallback to plain text on HTML error", func(t *testing.T) {
var requests []*telegramMessageRequest
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
var req telegramMessageRequest
err := json.NewDecoder(r.Body).Decode(&req)
require.NoError(t, err)
requests = append(requests, &req)
if len(requests) == 1 {
w.WriteHeader(http.StatusBadRequest)
_, _ = w.Write([]byte(`{"ok": false, "error_code": 400, "description": "Bad Request: can't parse entities"}`))
} else {
w.WriteHeader(http.StatusOK)
_, _ = w.Write([]byte(`{"ok": true}`))
}
}))
defer server.Close()
pusher := &TelegramPusher{}
cfg := Config{
Channel: "telegram",
URL: server.URL,
Secret: "my-token",
}
body := map[string]any{
"title": "Alert & Info",
"content": "A < B comparison",
"level": "INFO",
}
err := pusher.Send(context.Background(), cfg, "123456", body, "", nil)
require.NoError(t, err)
require.Len(t, requests, 2)
assert.Equal(t, "HTML", requests[0].ParseMode)
assert.Equal(t, "", requests[1].ParseMode)
assert.Contains(t, requests[1].Text, "[INFO] Alert & Info")
assert.Contains(t, requests[1].Text, "A < B comparison")
})
t.Run("validation error", func(t *testing.T) {
pusher := &TelegramPusher{}
cfg := Config{
Channel: "telegram",
URL: "https://api.telegram.org",
}
err := pusher.ValidateConfig(cfg)
assert.Error(t, err)
cfg = Config{
Channel: "telegram",
URL: "ftp://api.telegram.org",
Secret: "token",
}
err = pusher.ValidateConfig(cfg)
assert.Error(t, err)
cfg = Config{
Channel: "telegram",
Secret: "token",
}
err = pusher.ValidateConfig(cfg)
assert.NoError(t, err)
})
}
+78
View File
@@ -0,0 +1,78 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package push
import (
"encoding/json"
"fmt"
"strconv"
"strings"
)
// ParseTemplate parses template strings by replacing {{placeholder}} structures with values from body.
// It is a single-pass parser designed for high performance and low allocations.
func ParseTemplate(template string, body map[string]any) string {
var buf strings.Builder
buf.Grow(len(template))
i := 0
for {
pos := strings.Index(template[i:], "{{")
if pos == -1 {
buf.WriteString(template[i:])
break
}
// Write prefix
buf.WriteString(template[i : i+pos])
i += pos + 2 // skip "{{"
endPos := strings.Index(template[i:], "}}")
if endPos == -1 {
// Unbalanced "{{"
buf.WriteString("{{")
buf.WriteString(template[i:])
break
}
key := template[i : i+endPos]
if val, ok := body[key]; ok {
buf.WriteString(formatValue(val))
} else {
// Keep the placeholder if key not found
buf.WriteString("{{")
buf.WriteString(key)
buf.WriteString("}}")
}
i += endPos + 2 // skip "}}"
}
return buf.String()
}
func formatValue(v any) string {
if v == nil {
return ""
}
switch val := v.(type) {
case string:
return val
case []byte:
return string(val)
case int:
return strconv.Itoa(val)
case int32:
return strconv.FormatInt(int64(val), 10)
case int64:
return strconv.FormatInt(val, 10)
case float64:
return strconv.FormatFloat(val, 'f', -1, 64)
case bool:
return strconv.FormatBool(val)
default:
// If it's a map, slice, or struct, marshal it to JSON.
b, err := json.Marshal(v)
if err == nil {
return string(b)
}
return fmt.Sprintf("%v", v)
}
}
+75
View File
@@ -0,0 +1,75 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package push
import (
"testing"
"github.com/stretchr/testify/assert"
)
func TestParseTemplate(t *testing.T) {
tests := []struct {
name string
template string
body map[string]any
expected string
}{
{
name: "simple replacement",
template: "hello {{name}}",
body: map[string]any{"name": "world"},
expected: "hello world",
},
{
name: "multiple replacements",
template: "{{greeting}} {{name}}!",
body: map[string]any{"greeting": "Hello", "name": "Alice"},
expected: "Hello Alice!",
},
{
name: "missing key preserves placeholder",
template: "hello {{name}} and {{other}}",
body: map[string]any{"name": "world"},
expected: "hello world and {{other}}",
},
{
name: "unbalanced placeholders",
template: "hello {{name",
body: map[string]any{"name": "world"},
expected: "hello {{name",
},
{
name: "nil value",
template: "val: {{val}}",
body: map[string]any{"val": nil},
expected: "val: ",
},
{
name: "basic types",
template: "int: {{i}}, float: {{f}}, bool: {{b}}",
body: map[string]any{"i": 123, "f": 45.67, "b": true},
expected: "int: 123, float: 45.67, bool: true",
},
{
name: "complex type slice",
template: "items: {{items}}",
body: map[string]any{"items": []string{"a", "b"}},
expected: `items: ["a","b"]`,
},
{
name: "complex type map",
template: "obj: {{obj}}",
body: map[string]any{"obj": map[string]any{"key": "value"}},
expected: `obj: {"key":"value"}`,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
result := ParseTemplate(tt.template, tt.body)
assert.Equal(t, tt.expected, result)
})
}
}
+994
View File
@@ -0,0 +1,994 @@
package openresty
import (
"crypto/sha256"
"crypto/x509"
"encoding/base64"
"encoding/hex"
"encoding/json"
"encoding/pem"
"errors"
"fmt"
"net/url"
"path"
"regexp"
"sort"
"strings"
)
func RenderJSON(sourceJSON string, certificateFiles []SupportFile) (*Result, error) {
var doc Document
if err := json.Unmarshal([]byte(strings.TrimSpace(sourceJSON)), &doc); err != nil {
return nil, fmt.Errorf("openresty source config json is invalid: %w", err)
}
return Render(doc, certificateFiles)
}
func Render(doc Document, certificateFiles []SupportFile) (*Result, error) {
mainConfig := RenderMainConfig(doc.OpenRestyConfig)
routeConfig, err := RenderRouteConfig(doc, certificateFiles)
if err != nil {
return nil, err
}
wafConfig, err := RenderWAFConfig(doc.WAF)
if err != nil {
return nil, err
}
files := append([]SupportFile(nil), certificateFiles...)
files = append(files, SupportFile{Path: "waf_config.json", Content: wafConfig})
files = DedupeSupportFiles(files)
return &Result{
MainConfig: mainConfig,
RouteConfig: routeConfig,
SupportFiles: files,
Checksum: ChecksumBundle(mainConfig, routeConfig, files),
}, nil
}
func RenderMainConfig(cfg ConfigSnapshot) string {
templateText := cfg.MainConfigTemplate
if strings.TrimSpace(templateText) == "" {
templateText = defaultMainConfigTemplate
}
return renderMainConfigTemplate(templateText, cfg)
}
func ValidateMainConfigTemplate(templateText string) error {
trimmed := strings.TrimSpace(templateText)
if trimmed == "" {
return errors.New("OpenRestyMainConfigTemplate 不能为空")
}
for _, placeholder := range requiredMainConfigTemplatePlaceholders {
if !strings.Contains(trimmed, placeholder) {
return fmt.Errorf("OpenRestyMainConfigTemplate 必须保留占位符 %s", placeholder)
}
}
return nil
}
func RenderRouteConfig(doc Document, certificateFiles []SupportFile) (string, error) {
var builder strings.Builder
builder.WriteString("# This file is generated by OpenFlare. Do not edit manually.\n")
certificates := certificatesByID(certificateFiles)
for _, route := range doc.Routes {
domains := normalizedRouteDomains(route)
if len(domains) == 0 {
return "", fmt.Errorf("route %s domains are invalid", route.Domain)
}
serverNames := renderServerNames(domains)
displayName := strings.TrimSpace(route.SiteName)
if displayName == "" {
displayName = domains[0]
}
cacheConfig := routeCacheConfig{Enabled: route.CacheEnabled, Policy: route.CachePolicy, Rules: route.CacheRules}
limitConfig := routeLimitConfig{LimitConnPerServer: route.LimitConnPerServer, LimitConnPerIP: route.LimitConnPerIP, LimitRate: route.LimitRate}
powEnabled, _ := getPoWConfigForRoute(route.ID, doc.WAF)
if normalizeRouteUpstreamType(route.UpstreamType) == "pages" {
if route.PagesDeployment == nil {
return "", fmt.Errorf("route %s pages deployment is missing", route.Domain)
}
if !route.EnableHTTPS {
builder.WriteString(renderHTTPPagesServer(serverNames, displayName, route.PagesDeployment, limitConfig, powEnabled, route.BasicAuthEnabled, route.BasicAuthUsername, route.BasicAuthPassword))
continue
}
certIDs := normalizeCertIDs(route.CertID, route.CertIDs)
domainCertIDs := normalizeDomainCertIDs(domains, certIDs, route.DomainCertIDs)
if len(certIDs) == 0 {
return "", fmt.Errorf("路由 %s 未配置证书", route.Domain)
}
httpOnlyDomains := make([]string, 0, len(domains))
domainsByCertID := make(map[uint][]string, len(certIDs))
for index, domain := range domains {
if index >= len(domainCertIDs) || domainCertIDs[index] == 0 {
httpOnlyDomains = append(httpOnlyDomains, domain)
continue
}
domainsByCertID[domainCertIDs[index]] = append(domainsByCertID[domainCertIDs[index]], domain)
}
for _, certID := range certIDs {
assignedDomains := domainsByCertID[certID]
if len(assignedDomains) == 0 {
continue
}
certPEM, ok := certificates[certID]
if !ok {
return "", fmt.Errorf("route %s certificate %d does not exist", route.Domain, certID)
}
if err := validateCertificateCoverage(certPEM, assignedDomains); err != nil {
return "", fmt.Errorf("site %s certificate validation failed: %w", displayName, err)
}
}
if route.RedirectHTTP {
if len(httpOnlyDomains) > 0 {
builder.WriteString(renderHTTPPagesServer(renderServerNames(httpOnlyDomains), displayName, route.PagesDeployment, limitConfig, powEnabled, route.BasicAuthEnabled, route.BasicAuthUsername, route.BasicAuthPassword))
}
for _, certID := range certIDs {
if assignedDomains := domainsByCertID[certID]; len(assignedDomains) > 0 {
builder.WriteString(renderHTTPRedirectServer(renderServerNames(assignedDomains)))
}
}
} else {
builder.WriteString(renderHTTPPagesServer(serverNames, displayName, route.PagesDeployment, limitConfig, powEnabled, route.BasicAuthEnabled, route.BasicAuthUsername, route.BasicAuthPassword))
}
for _, certID := range certIDs {
if assignedDomains := domainsByCertID[certID]; len(assignedDomains) > 0 {
builder.WriteString(renderHTTPSPagesServer(renderServerNames(assignedDomains), displayName, certID, route.PagesDeployment, limitConfig, powEnabled, route.BasicAuthEnabled, route.BasicAuthUsername, route.BasicAuthPassword, doc.OpenRestyConfig))
}
}
continue
}
upstreams := route.Upstreams
if len(upstreams) == 0 && strings.TrimSpace(route.OriginURL) != "" {
upstreams = []string{route.OriginURL}
}
upstreamConfig := buildRouteUpstreamConfig(route, upstreams)
if upstreamConfig.UsesNamedUpstream {
builder.WriteString(renderNamedUpstreamBlock(upstreamConfig))
}
if !route.EnableHTTPS {
builder.WriteString(renderHTTPProxyServer(serverNames, displayName, route.OriginURL, route.OriginHost, route.CustomHeaders, cacheConfig, limitConfig, upstreamConfig, powEnabled, route.BasicAuthEnabled, route.BasicAuthUsername, route.BasicAuthPassword, doc.OpenRestyConfig))
continue
}
certIDs := normalizeCertIDs(route.CertID, route.CertIDs)
domainCertIDs := normalizeDomainCertIDs(domains, certIDs, route.DomainCertIDs)
if len(certIDs) == 0 {
return "", fmt.Errorf("路由 %s 未配置证书", route.Domain)
}
httpOnlyDomains := make([]string, 0, len(domains))
domainsByCertID := make(map[uint][]string, len(certIDs))
for index, domain := range domains {
if index >= len(domainCertIDs) || domainCertIDs[index] == 0 {
httpOnlyDomains = append(httpOnlyDomains, domain)
continue
}
domainsByCertID[domainCertIDs[index]] = append(domainsByCertID[domainCertIDs[index]], domain)
}
for _, certID := range certIDs {
assignedDomains := domainsByCertID[certID]
if len(assignedDomains) == 0 {
continue
}
certPEM, ok := certificates[certID]
if !ok {
return "", fmt.Errorf("route %s certificate %d does not exist", route.Domain, certID)
}
if err := validateCertificateCoverage(certPEM, assignedDomains); err != nil {
return "", fmt.Errorf("site %s certificate validation failed: %w", displayName, err)
}
}
if route.RedirectHTTP {
if len(httpOnlyDomains) > 0 {
builder.WriteString(renderHTTPProxyServer(renderServerNames(httpOnlyDomains), displayName, route.OriginURL, route.OriginHost, route.CustomHeaders, cacheConfig, limitConfig, upstreamConfig, powEnabled, route.BasicAuthEnabled, route.BasicAuthUsername, route.BasicAuthPassword, doc.OpenRestyConfig))
}
for _, certID := range certIDs {
if assignedDomains := domainsByCertID[certID]; len(assignedDomains) > 0 {
builder.WriteString(renderHTTPRedirectServer(renderServerNames(assignedDomains)))
}
}
} else {
builder.WriteString(renderHTTPProxyServer(serverNames, displayName, route.OriginURL, route.OriginHost, route.CustomHeaders, cacheConfig, limitConfig, upstreamConfig, powEnabled, route.BasicAuthEnabled, route.BasicAuthUsername, route.BasicAuthPassword, doc.OpenRestyConfig))
}
for _, certID := range certIDs {
if assignedDomains := domainsByCertID[certID]; len(assignedDomains) > 0 {
builder.WriteString(renderHTTPSServer(renderServerNames(assignedDomains), displayName, route.OriginURL, route.OriginHost, certID, route.CustomHeaders, cacheConfig, limitConfig, upstreamConfig, powEnabled, route.BasicAuthEnabled, route.BasicAuthUsername, route.BasicAuthPassword, doc.OpenRestyConfig))
}
}
}
return builder.String(), nil
}
func RenderPoWConfig(doc Document) (string, error) {
type domainEntry struct {
Domains []string `json:"domains"`
Enabled bool `json:"enabled"`
Config *PoWConfig `json:"config"`
}
entries := make([]domainEntry, 0)
for _, route := range doc.Routes {
powEnabled, powConfig := getPoWConfigForRoute(route.ID, doc.WAF)
if !powEnabled {
continue
}
entries = append(entries, domainEntry{Domains: normalizedRouteDomains(route), Enabled: true, Config: powConfig})
}
if len(entries) == 0 {
return "{}", nil
}
data, err := json.Marshal(entries)
return string(data), err
}
func RenderWAFConfig(snapshot WAFDocument) (string, error) {
type wafRuntimeRuleGroup struct {
ID uint `json:"id"`
Name string `json:"name"`
IsGlobal bool `json:"is_global"`
BlockStatusCode int `json:"block_status_code"`
BlockResponseBody string `json:"block_response_body"`
IPWhitelist []string `json:"ip_whitelist"`
IPBlacklist []string `json:"ip_blacklist"`
IPWhitelistGroups []uint `json:"ip_whitelist_group_ids,omitempty"`
IPBlacklistGroups []uint `json:"ip_blacklist_group_ids,omitempty"`
CountryWhitelist []string `json:"country_whitelist"`
CountryBlacklist []string `json:"country_blacklist"`
RegionWhitelist []string `json:"region_whitelist"`
RegionBlacklist []string `json:"region_blacklist"`
PoWEnabled bool `json:"pow_enabled"`
PoWConfig *PoWConfig `json:"pow_config,omitempty"`
}
type wafRuntimeConfig struct {
DefaultBlockStatusCode int `json:"default_block_status_code"`
RuleGroups []wafRuntimeRuleGroup `json:"rule_groups"`
SiteRuleGroups map[string][]uint `json:"site_rule_groups"`
}
groups := make([]wafRuntimeRuleGroup, 0, len(snapshot.RuleGroups))
globalGroupIDs := make([]uint, 0)
enabledGroupIDs := make(map[uint]struct{}, len(snapshot.RuleGroups))
for _, group := range snapshot.RuleGroups {
if !group.Enabled {
continue
}
statusCode := group.BlockStatusCode
if statusCode == 0 {
statusCode = defaultWAFBlockStatus
}
if group.IsGlobal {
globalGroupIDs = append(globalGroupIDs, group.ID)
}
enabledGroupIDs[group.ID] = struct{}{}
powConfig := group.PoWConfig
if !group.PoWEnabled {
powConfig = nil
}
groups = append(groups, wafRuntimeRuleGroup{
ID: group.ID,
Name: group.Name,
IsGlobal: group.IsGlobal,
BlockStatusCode: statusCode,
BlockResponseBody: group.BlockResponseBody,
IPWhitelist: sortedUniqueStrings(group.IPWhitelist),
IPBlacklist: sortedUniqueStrings(group.IPBlacklist),
IPWhitelistGroups: sortedUniqueUintIDs(group.IPWhitelistGroups),
IPBlacklistGroups: sortedUniqueUintIDs(group.IPBlacklistGroups),
CountryWhitelist: group.CountryWhitelist,
CountryBlacklist: group.CountryBlacklist,
RegionWhitelist: group.RegionWhitelist,
RegionBlacklist: group.RegionBlacklist,
PoWEnabled: group.PoWEnabled,
PoWConfig: powConfig,
})
}
sort.Slice(groups, func(i, j int) bool {
if groups[i].IsGlobal != groups[j].IsGlobal {
return groups[i].IsGlobal
}
return groups[i].ID < groups[j].ID
})
sort.Slice(globalGroupIDs, func(i, j int) bool { return globalGroupIDs[i] < globalGroupIDs[j] })
siteRuleGroups := make(map[string][]uint, len(snapshot.Bindings))
for _, binding := range snapshot.Bindings {
ids := append([]uint{}, globalGroupIDs...)
for _, id := range binding.RuleGroupIDs {
if _, ok := enabledGroupIDs[id]; ok {
ids = append(ids, id)
}
}
siteRuleGroups[binding.SiteName] = uniqueUintIDs(ids)
}
data, err := json.Marshal(wafRuntimeConfig{DefaultBlockStatusCode: defaultWAFBlockStatus, RuleGroups: groups, SiteRuleGroups: siteRuleGroups})
return string(data), err
}
func sortedUniqueStrings(values []string) []string {
items := append([]string{}, values...)
items = uniqueStrings(items)
sort.Strings(items)
return items
}
func sortedUniqueUintIDs(values []uint) []uint {
items := uniqueUintIDs(values)
sort.Slice(items, func(i, j int) bool { return items[i] < items[j] })
return items
}
func ChecksumBundle(mainConfig string, routeConfig string, supportFiles []SupportFile) string {
var builder strings.Builder
builder.WriteString(mainConfig)
builder.WriteString("\n--route-config--\n")
builder.WriteString(routeConfig)
builder.WriteString("\n--support-files--\n")
files := DedupeSupportFiles(supportFiles)
sort.Slice(files, func(i int, j int) bool { return files[i].Path < files[j].Path })
for _, file := range files {
if file.Path == SourceConfigFileName {
continue
}
builder.WriteString(file.Path)
builder.WriteString("\n")
builder.WriteString(file.Content)
builder.WriteString("\n")
}
sum := sha256.Sum256([]byte(builder.String()))
return hex.EncodeToString(sum[:])
}
func DedupeSupportFiles(files []SupportFile) []SupportFile {
if len(files) == 0 {
return nil
}
unique := make(map[string]SupportFile, len(files))
for _, file := range files {
unique[file.Path] = file
}
result := make([]SupportFile, 0, len(unique))
for _, file := range unique {
result = append(result, file)
}
return result
}
func renderMainConfigTemplate(templateText string, cfg ConfigSnapshot) string {
replacer := strings.NewReplacer(
"{{OpenRestyWorkerProcesses}}", cfg.WorkerProcesses,
"{{OpenRestyWorkerConnections}}", fmt.Sprintf("%d", cfg.WorkerConnections),
"{{OpenRestyWorkerRlimitNofile}}", fmt.Sprintf("%d", cfg.WorkerRlimitNofile),
"{{OpenRestyConnectionUpgradeMap}}", renderConnectionUpgradeMap(),
"{{OpenRestyDefaultServerBlock}}", renderDefaultServerBlock(cfg.DefaultServerReturnStatus, cfg.HTTP3Enabled),
"{{OpenRestyAccessLogPath}}", AccessLogPlaceholder,
"{{OpenRestyErrorLogPath}}", ErrorLogPlaceholder,
"{{OpenRestyEventsUseDirective}}", renderTemplateDirective(cfg.EventsUse != "", fmt.Sprintf("use %s;", cfg.EventsUse)),
"{{OpenRestyEventsMultiAcceptDirective}}", renderTemplateDirective(cfg.EventsMultiAcceptEnabled, "multi_accept on;"),
"{{OpenRestyKeepaliveTimeout}}", fmt.Sprintf("%d", cfg.KeepaliveTimeout),
"{{OpenRestyKeepaliveRequests}}", fmt.Sprintf("%d", cfg.KeepaliveRequests),
"{{OpenRestyClientHeaderTimeout}}", fmt.Sprintf("%d", cfg.ClientHeaderTimeout),
"{{OpenRestyClientBodyTimeout}}", fmt.Sprintf("%d", cfg.ClientBodyTimeout),
"{{OpenRestyClientMaxBodySize}}", cfg.ClientMaxBodySize,
"{{OpenRestyLargeClientHeaderBuffers}}", cfg.LargeClientHeaderBuffers,
"{{OpenRestySendTimeout}}", fmt.Sprintf("%d", cfg.SendTimeout),
"{{OpenRestyProxyConnectTimeout}}", fmt.Sprintf("%d", cfg.ProxyConnectTimeout),
"{{OpenRestyProxySendTimeout}}", fmt.Sprintf("%d", cfg.ProxySendTimeout),
"{{OpenRestyProxyReadTimeout}}", fmt.Sprintf("%d", cfg.ProxyReadTimeout),
"{{OpenRestyProxyRequestBuffering}}", onOff(cfg.ProxyRequestBuffering),
"{{OpenRestyProxyBuffering}}", onOff(cfg.ProxyBufferingEnabled),
"{{OpenRestyProxyBuffers}}", cfg.ProxyBuffers,
"{{OpenRestyProxyBufferSize}}", cfg.ProxyBufferSize,
"{{OpenRestyProxyBusyBuffersSize}}", cfg.ProxyBusyBuffersSize,
"{{OpenRestyGzip}}", onOff(cfg.GzipEnabled),
"{{OpenRestyGzipMinLength}}", fmt.Sprintf("%d", cfg.GzipMinLength),
"{{OpenRestyGzipCompLevel}}", fmt.Sprintf("%d", cfg.GzipCompLevel),
"{{OpenRestyResolverDirective}}", renderTemplateDirective(cfg.Resolvers != "", fmt.Sprintf("resolver %s;", cfg.Resolvers)),
"{{OpenRestyCacheBlock}}", renderOpenRestyCacheTemplateBlock(cfg),
"{{OpenRestyRouteConfigInclude}}", RouteConfigPlaceholder,
)
return replacer.Replace(templateText)
}
func renderTemplateDirective(enabled bool, statement string) string {
if !enabled {
return ""
}
return fmt.Sprintf(" %s\n", statement)
}
func renderOpenRestyCacheTemplateBlock(cfg ConfigSnapshot) string {
lines := []string{renderOpenRestyLimitZoneBlock()}
if !cfg.CacheEnabled {
lines = append(lines, renderOpenRestyObservabilityTemplateBlock())
return strings.Join(lines, "")
}
lines = append(lines, strings.Join([]string{
fmt.Sprintf(" proxy_cache_path %s levels=%s keys_zone=openflare_cache:10m inactive=%s max_size=%s;", cfg.CachePath, cfg.CacheLevels, cfg.CacheInactive, cfg.CacheMaxSize),
fmt.Sprintf(" proxy_cache_key \"%s\";", cfg.CacheKeyTemplate),
fmt.Sprintf(" proxy_cache_lock %s;", onOff(cfg.CacheLockEnabled)),
fmt.Sprintf(" proxy_cache_lock_timeout %s;", cfg.CacheLockTimeout),
fmt.Sprintf(" proxy_cache_use_stale %s;", cfg.CacheUseStale),
"",
}, "\n"))
lines = append(lines, renderOpenRestyObservabilityTemplateBlock())
return strings.Join(lines, "")
}
func renderOpenRestyLimitZoneBlock() string {
return " limit_conn_zone $server_name zone=openflare_conn_per_server:10m;\n limit_conn_zone $binary_remote_addr zone=openflare_conn_per_ip:10m;\n"
}
func renderOpenRestyObservabilityTemplateBlock() string {
return fmt.Sprintf(" lua_shared_dict openflare_observability 10m;\n lua_shared_dict openflare_pow_challenges 10m;\n lua_shared_dict openflare_pow_sessions 10m;\n lua_shared_dict openflare_pow_config 1m;\n lua_shared_dict openflare_waf_config 1m;\n init_worker_by_lua_file %s/observability/init.lua;\n log_by_lua_file %s/observability/log.lua;\n\n server {\n listen %s;\n server_name openflare-observability;\n access_log off;\n\n location = /openflare/stub_status {\n stub_status;\n }\n\n location = /openflare/observability {\n default_type application/json;\n content_by_lua_file %s/observability/read.lua;\n }\n }\n\n", LuaDirPlaceholder, LuaDirPlaceholder, ObservabilityListenPlaceholder, LuaDirPlaceholder)
}
func renderHTTPProxyServer(serverNames string, siteName string, originURL string, originHost string, customHeaders []CustomHeader, cacheConfig routeCacheConfig, limitConfig routeLimitConfig, upstreamConfig routeUpstreamConfig, powEnabled bool, basicAuthEnabled bool, basicAuthUsername string, basicAuthPassword string, cfg ConfigSnapshot) string {
return fmt.Sprintf("server {\n listen 80;\n server_name %s;\n%s%s location / {\n%s%s%s%s%s }\n%s}\n\n", serverNames, renderAccessBlock(siteName, powEnabled), renderPowLocationBlocks(powEnabled), renderBasicAuthBlock(basicAuthEnabled, basicAuthUsername, basicAuthPassword), renderProxyHeaderBlock(originURL, originHost, customHeaders, upstreamConfig, cfg), renderRouteLimitBlock(limitConfig), renderRouteCacheBlock(cacheConfig, cfg), renderProxyPassBlock(originURL, upstreamConfig), renderPowStaticLocationBlock(powEnabled))
}
func renderPagesAPIProxyLocationBlock(deployment *PagesDeployment) string {
if deployment == nil || !deployment.APIProxyEnabled {
return ""
}
path := strings.TrimSpace(deployment.APIProxyPath)
pass := strings.TrimSpace(deployment.APIProxyPass)
rewrite := strings.TrimSpace(deployment.APIProxyRewrite)
if path == "" || pass == "" {
return ""
}
if !strings.HasPrefix(path, "/") {
path = "/" + path
}
cleanPath := strings.TrimSuffix(path, "/")
var builder strings.Builder
builder.WriteString(fmt.Sprintf("\n location %s {\n", cleanPath))
if rewrite != "" {
if !strings.HasPrefix(rewrite, "/") {
rewrite = "/" + rewrite
}
cleanRewrite := strings.TrimSuffix(rewrite, "/")
if cleanRewrite == "" {
builder.WriteString(fmt.Sprintf(" rewrite ^%s/(.*)$ /$1 break;\n", regexp.QuoteMeta(cleanPath)))
builder.WriteString(fmt.Sprintf(" rewrite ^%s$ / break;\n", regexp.QuoteMeta(cleanPath)))
} else {
builder.WriteString(fmt.Sprintf(" rewrite ^%s/(.*)$ %s/$1 break;\n", regexp.QuoteMeta(cleanPath), cleanRewrite))
builder.WriteString(fmt.Sprintf(" rewrite ^%s$ %s break;\n", regexp.QuoteMeta(cleanPath), cleanRewrite))
}
}
builder.WriteString(fmt.Sprintf(" proxy_pass %s;\n", pass))
builder.WriteString(" proxy_http_version 1.1;\n")
builder.WriteString(" proxy_set_header Host $http_host;\n")
builder.WriteString(" proxy_set_header X-Real-IP $remote_addr;\n")
builder.WriteString(" proxy_set_header X-Forwarded-For $proxy_add_x_forwarded_for;\n")
builder.WriteString(" proxy_set_header X-Forwarded-Proto $scheme;\n")
builder.WriteString(" proxy_set_header Upgrade $http_upgrade;\n")
builder.WriteString(" proxy_set_header Connection $connection_upgrade;\n")
builder.WriteString(" }\n")
return builder.String()
}
func renderHTTPPagesServer(serverNames string, siteName string, deployment *PagesDeployment, limitConfig routeLimitConfig, powEnabled bool, basicAuthEnabled bool, basicAuthUsername string, basicAuthPassword string) string {
return fmt.Sprintf("server {\n listen 80;\n server_name %s;\n%s%s root %s;\n index %s;%s\n\n location / {\n%s%s }\n%s}\n\n", serverNames, renderAccessBlock(siteName, powEnabled), renderPowLocationBlocks(powEnabled), quoteNginxStringLiteral(pagesDeploymentRoot(deployment)), quoteNginxStringLiteral(pagesEntryFile(deployment)), renderPagesAPIProxyLocationBlock(deployment), renderBasicAuthBlock(basicAuthEnabled, basicAuthUsername, basicAuthPassword), renderPagesLocationBlock(deployment, limitConfig), renderPowStaticLocationBlock(powEnabled))
}
func renderHTTPRedirectServer(serverNames string) string {
return fmt.Sprintf("server {\n listen 80;\n server_name %s;\n\n return 301 https://$host$request_uri;\n}\n\n", serverNames)
}
func renderHTTPSServer(serverNames string, siteName string, originURL string, originHost string, certificateID uint, customHeaders []CustomHeader, cacheConfig routeCacheConfig, limitConfig routeLimitConfig, upstreamConfig routeUpstreamConfig, powEnabled bool, basicAuthEnabled bool, basicAuthUsername string, basicAuthPassword string, cfg ConfigSnapshot) string {
certPath := fmt.Sprintf("%s/%d.crt", CertDirPlaceholder, certificateID)
keyPath := fmt.Sprintf("%s/%d.key", CertDirPlaceholder, certificateID)
var h3Listen string
var h3Header string
if cfg.HTTP3Enabled {
h3Listen = " listen 443 quic;\n"
h3Header = " add_header Alt-Svc 'h3=\":443\"; ma=86400';\n"
}
return fmt.Sprintf("server {\n listen 443 ssl;\n%s http2 on;\n server_name %s;\n ssl_certificate %s;\n ssl_certificate_key %s;\n%s%s%s location / {\n%s%s%s%s%s }\n%s}\n\n", h3Listen, serverNames, certPath, keyPath, h3Header, renderAccessBlock(siteName, powEnabled), renderPowLocationBlocks(powEnabled), renderBasicAuthBlock(basicAuthEnabled, basicAuthUsername, basicAuthPassword), renderProxyHeaderBlock(originURL, originHost, customHeaders, upstreamConfig, cfg), renderRouteLimitBlock(limitConfig), renderRouteCacheBlock(cacheConfig, cfg), renderProxyPassBlock(originURL, upstreamConfig), renderPowStaticLocationBlock(powEnabled))
}
func renderHTTPSPagesServer(serverNames string, siteName string, certificateID uint, deployment *PagesDeployment, limitConfig routeLimitConfig, powEnabled bool, basicAuthEnabled bool, basicAuthUsername string, basicAuthPassword string, cfg ConfigSnapshot) string {
certPath := fmt.Sprintf("%s/%d.crt", CertDirPlaceholder, certificateID)
keyPath := fmt.Sprintf("%s/%d.key", CertDirPlaceholder, certificateID)
var h3Listen string
var h3Header string
if cfg.HTTP3Enabled {
h3Listen = " listen 443 quic;\n"
h3Header = " add_header Alt-Svc 'h3=\":443\"; ma=86400';\n"
}
return fmt.Sprintf("server {\n listen 443 ssl;\n%s http2 on;\n server_name %s;\n ssl_certificate %s;\n ssl_certificate_key %s;\n%s%s%s root %s;\n index %s;%s\n\n location / {\n%s%s }\n%s}\n\n", h3Listen, serverNames, certPath, keyPath, h3Header, renderAccessBlock(siteName, powEnabled), renderPowLocationBlocks(powEnabled), quoteNginxStringLiteral(pagesDeploymentRoot(deployment)), quoteNginxStringLiteral(pagesEntryFile(deployment)), renderPagesAPIProxyLocationBlock(deployment), renderBasicAuthBlock(basicAuthEnabled, basicAuthUsername, basicAuthPassword), renderPagesLocationBlock(deployment, limitConfig), renderPowStaticLocationBlock(powEnabled))
}
func renderPagesLocationBlock(deployment *PagesDeployment, limitConfig routeLimitConfig) string {
var builder strings.Builder
builder.WriteString(renderRouteLimitBlock(limitConfig))
if deployment != nil && deployment.SPAFallbackEnabled {
builder.WriteString(fmt.Sprintf(" try_files $uri $uri/ %s;\n", pagesFallbackPath(deployment)))
} else {
builder.WriteString(" try_files $uri $uri/ =404;\n")
}
return builder.String()
}
func pagesDeploymentRoot(deployment *PagesDeployment) string {
if deployment == nil || strings.TrimSpace(deployment.LocalRoot) == "" {
return PagesDirPlaceholder
}
return filepathToNginxPath(deployment.LocalRoot)
}
func pagesEntryFile(deployment *PagesDeployment) string {
if deployment == nil || strings.TrimSpace(deployment.EntryFile) == "" {
return "index.html"
}
return strings.TrimPrefix(filepathToNginxPath(deployment.EntryFile), "/")
}
func pagesFallbackPath(deployment *PagesDeployment) string {
if deployment == nil || strings.TrimSpace(deployment.SPAFallbackPath) == "" {
return "/index.html"
}
value := filepathToNginxPath(strings.TrimSpace(deployment.SPAFallbackPath))
if !strings.HasPrefix(value, "/") {
value = "/" + value
}
if value == "/" || strings.HasSuffix(value, "/") || strings.Contains(value, "\\") || strings.ContainsAny(value, "\"';") || strings.ContainsAny(value, " \t\r\n") {
return "/index.html"
}
for _, segment := range strings.Split(value, "/") {
if segment == "." || segment == ".." {
return "/index.html"
}
}
cleaned := path.Clean(value)
if cleaned == "/" || strings.HasSuffix(cleaned, "/") {
return "/index.html"
}
return cleaned
}
func renderProxyHeaderBlock(originURL string, originHost string, customHeaders []CustomHeader, upstreamConfig routeUpstreamConfig, cfg ConfigSnapshot) string {
var builder strings.Builder
if strings.TrimSpace(originHost) != "" {
builder.WriteString(fmt.Sprintf(" proxy_set_header Host %s;\n", quoteNginxStringLiteral(originHost)))
} else {
builder.WriteString(" proxy_set_header Host $host;\n")
}
if upstreamServerName := resolveUpstreamServerName(originURL, originHost); upstreamServerName != "" {
builder.WriteString(" proxy_ssl_server_name on;\n")
builder.WriteString(fmt.Sprintf(" proxy_ssl_name %s;\n", quoteNginxStringLiteral(upstreamServerName)))
}
builder.WriteString(" proxy_set_header X-Real-IP $remote_addr;\n")
builder.WriteString(" proxy_set_header X-Forwarded-For $proxy_add_x_forwarded_for;\n")
builder.WriteString(" proxy_set_header X-Forwarded-Proto $scheme;\n")
if cfg.WebsocketEnabled {
builder.WriteString(" proxy_http_version 1.1;\n")
builder.WriteString(" proxy_set_header Connection $connection_upgrade;\n")
builder.WriteString(" proxy_set_header Upgrade $http_upgrade;\n")
} else if upstreamConfig.UsesNamedUpstream {
builder.WriteString(" proxy_http_version 1.1;\n")
builder.WriteString(" proxy_set_header Connection \"\";\n")
}
for _, header := range customHeaders {
builder.WriteString(fmt.Sprintf(" proxy_set_header %s %s;\n", header.Key, quoteNginxStringLiteral(header.Value)))
}
return builder.String()
}
func renderAccessBlock(siteName string, powEnabled bool) string {
escapedSiteName := escapeNginxString(siteName)
if !powEnabled {
return fmt.Sprintf(" set $openflare_waf_site \"%s\";\n access_by_lua_file %s/waf/check.lua;\n", escapedSiteName, LuaDirPlaceholder)
}
return fmt.Sprintf(` set $openflare_waf_site "%s";
access_by_lua_block {
if not string.find(package.path, "%s/?.lua", 1, true) then
package.path = "%s/?.lua;%s/?/init.lua;" .. package.path
end
require("waf.runtime").check()
if ngx.ctx.openflare_waf_blocked then
return
end
require("pow.runtime").check()
}
`, escapedSiteName, LuaDirPlaceholder, LuaDirPlaceholder, LuaDirPlaceholder)
}
func renderBasicAuthBlock(enabled bool, username, password string) string {
if !enabled || username == "" || password == "" {
return ""
}
encoded := base64.StdEncoding.EncodeToString([]byte(username + ":" + password))
return fmt.Sprintf(" rewrite_by_lua_block {\n local auth = ngx.var.http_authorization\n if auth ~= \"Basic %s\" then\n ngx.header[\"WWW-Authenticate\"] = 'Basic realm=\"Restricted\"'\n return ngx.exit(401)\n end\n }\n", encoded)
}
func renderPowLocationBlocks(powEnabled bool) string {
if !powEnabled {
return ""
}
return fmt.Sprintf("\n location = %spass-challenge {\n content_by_lua_file %s/pow/verify.lua;\n }\n\n location = %smake-challenge {\n content_by_lua_file %s/pow/challenge.lua;\n }\n\n", anubisAPIPrefix, LuaDirPlaceholder, anubisAPIPrefix, LuaDirPlaceholder)
}
func renderPowStaticLocationBlock(powEnabled bool) string {
if !powEnabled {
return ""
}
return fmt.Sprintf(" location %s {\n alias %s/;\n types {\n text/css css;\n application/javascript js mjs;\n application/json json;\n image/webp webp;\n font/woff2 woff2;\n }\n }\n\n", anubisStaticPrefix, PowStaticDirPlaceholder)
}
func renderRouteCacheBlock(cacheConfig routeCacheConfig, cfg ConfigSnapshot) string {
if !cfg.CacheEnabled || !cacheConfig.Enabled {
return ""
}
var builder strings.Builder
builder.WriteString(" set $openflare_skip_cache 0;\n")
builder.WriteString(" if ($request_method != GET) {\n set $openflare_skip_cache 1;\n }\n")
builder.WriteString(" if ($http_authorization != \"\") {\n set $openflare_skip_cache 1;\n }\n")
builder.WriteString(" if ($http_cookie ~* \"(session|sess|token|auth|jwt|logged_in|remember|laravel_session|connect\\\\.sid|_session)\") {\n set $openflare_skip_cache 1;\n }\n")
builder.WriteString(" if ($http_cache_control ~* \"(no-cache|no-store|private)\") {\n set $openflare_skip_cache 1;\n }\n")
if condition := renderRouteCachePolicyCondition(cacheConfig); condition != "" {
builder.WriteString(condition)
}
builder.WriteString(" proxy_cache openflare_cache;\n")
builder.WriteString(" proxy_cache_methods GET;\n")
builder.WriteString(" proxy_cache_bypass $openflare_skip_cache;\n")
builder.WriteString(" proxy_no_cache $openflare_skip_cache;\n")
return builder.String()
}
func renderRouteLimitBlock(limitConfig routeLimitConfig) string {
var builder strings.Builder
if limitConfig.LimitConnPerServer > 0 {
builder.WriteString(fmt.Sprintf(" limit_conn openflare_conn_per_server %d;\n", limitConfig.LimitConnPerServer))
}
if limitConfig.LimitConnPerIP > 0 {
builder.WriteString(fmt.Sprintf(" limit_conn openflare_conn_per_ip %d;\n", limitConfig.LimitConnPerIP))
}
if strings.TrimSpace(limitConfig.LimitRate) != "" {
builder.WriteString(fmt.Sprintf(" limit_rate %s;\n", limitConfig.LimitRate))
}
return builder.String()
}
func renderRouteCachePolicyCondition(cacheConfig routeCacheConfig) string {
switch cacheConfig.Policy {
case cachePolicySuffix:
return fmt.Sprintf(" if ($uri !~* %s) {\n set $openflare_skip_cache 1;\n }\n", quoteNginxStringLiteral(buildSuffixMatchPattern(cacheConfig.Rules)))
case cachePolicyPathPrefix:
return fmt.Sprintf(" if ($uri !~ %s) {\n set $openflare_skip_cache 1;\n }\n", quoteNginxStringLiteral(buildPathPrefixMatchPattern(cacheConfig.Rules)))
case cachePolicyPathExact:
return fmt.Sprintf(" if ($uri !~ %s) {\n set $openflare_skip_cache 1;\n }\n", quoteNginxStringLiteral(buildPathExactMatchPattern(cacheConfig.Rules)))
default:
return ""
}
}
func renderProxyPassBlock(originURL string, upstreamConfig routeUpstreamConfig) string {
parsed, err := url.Parse(originURL)
if err != nil || parsed.Host == "" || parsed.Scheme == "" {
return fmt.Sprintf(" proxy_pass %s;\n", originURL)
}
if upstreamConfig.UsesNamedUpstream {
return fmt.Sprintf(" proxy_pass %s://%s%s;\n", upstreamConfig.Scheme, upstreamConfig.Name, upstreamConfig.ProxyPassURI)
}
return fmt.Sprintf(" proxy_pass %s;\n", originURL)
}
func buildRouteUpstreamConfig(route Route, upstreams []string) routeUpstreamConfig {
if len(upstreams) == 0 {
return routeUpstreamConfig{}
}
if len(upstreams) == 1 {
parsed, err := url.Parse(strings.TrimSpace(upstreams[0]))
if err != nil || parsed.Host == "" || parsed.Scheme == "" {
return routeUpstreamConfig{}
}
return routeUpstreamConfig{Name: buildRouteUpstreamName(route), Scheme: parsed.Scheme, ProxyPassURI: buildUpstreamProxyPassURI(parsed), Servers: []string{parsed.Host}, UsesNamedUpstream: true}
}
servers := make([]string, 0, len(upstreams))
var scheme string
for _, upstream := range upstreams {
parsed, err := url.Parse(strings.TrimSpace(upstream))
if err != nil || parsed.Host == "" || parsed.Scheme == "" || (strings.TrimSpace(parsed.EscapedPath()) != "" && strings.TrimSpace(parsed.EscapedPath()) != "/") || parsed.RawQuery != "" {
return routeUpstreamConfig{}
}
if scheme == "" {
scheme = parsed.Scheme
} else if scheme != parsed.Scheme {
return routeUpstreamConfig{}
}
servers = append(servers, parsed.Host)
}
return routeUpstreamConfig{Name: buildRouteUpstreamName(route), Scheme: scheme, Servers: servers, UsesNamedUpstream: true}
}
func normalizeRouteUpstreamType(raw string) string {
switch strings.ToLower(strings.TrimSpace(raw)) {
case "pages":
return "pages"
default:
return "direct"
}
}
func renderNamedUpstreamBlock(upstreamConfig routeUpstreamConfig) string {
var builder strings.Builder
builder.WriteString(fmt.Sprintf("upstream %s {\n", upstreamConfig.Name))
for _, server := range upstreamConfig.Servers {
builder.WriteString(fmt.Sprintf(" server %s max_fails=3 fail_timeout=10s;\n", server))
}
builder.WriteString(" keepalive 128;\n}\n\n")
return builder.String()
}
func buildRouteUpstreamName(route Route) string {
identity := strings.TrimSpace(route.SiteName)
if identity == "" {
identity = route.Domain
}
sanitized := strings.Map(func(r rune) rune {
switch {
case r >= 'a' && r <= 'z':
return r
case r >= 'A' && r <= 'Z':
return r + ('a' - 'A')
case r >= '0' && r <= '9':
return r
default:
return '_'
}
}, identity)
sanitized = strings.Trim(sanitized, "_")
if sanitized == "" {
sanitized = "backend"
}
return fmt.Sprintf("backend_%s_%d", sanitized, route.ID)
}
func buildUpstreamProxyPassURI(parsed *url.URL) string {
path := parsed.EscapedPath()
if path == "/" {
path = ""
}
if parsed.RawQuery == "" {
return path
}
return fmt.Sprintf("%s?%s", path, parsed.RawQuery)
}
func renderConnectionUpgradeMap() string {
return " map $http_upgrade $connection_upgrade {\n default upgrade;\n '' \"\";\n }\n\n"
}
func renderDefaultServerBlock(statusCode int, http3Enabled bool) string {
if statusCode <= 0 {
statusCode = 421
}
var h3Default string
if http3Enabled {
h3Default = "\n listen 443 quic reuseport default_server;"
}
return strings.Join([]string{
" server {",
" listen 80 default_server;",
" server_name _;",
"",
fmt.Sprintf(" return %d;", statusCode),
" }",
"",
" server {",
fmt.Sprintf(" listen 443 ssl default_server;%s", h3Default),
" server_name _;",
"",
" ssl_reject_handshake on;",
" }",
"",
}, "\n")
}
func normalizedRouteDomains(route Route) []string {
if len(route.Domains) > 0 {
return route.Domains
}
if strings.TrimSpace(route.Domain) == "" {
return nil
}
return []string{strings.TrimSpace(route.Domain)}
}
func normalizeCertIDs(primaryCertID *uint, certIDs []uint) []uint {
candidates := make([]uint, 0, len(certIDs)+1)
if primaryCertID != nil && *primaryCertID != 0 {
candidates = append(candidates, *primaryCertID)
}
candidates = append(candidates, certIDs...)
seen := make(map[uint]struct{}, len(candidates))
normalized := make([]uint, 0, len(candidates))
for _, id := range candidates {
if id == 0 {
continue
}
if _, ok := seen[id]; ok {
continue
}
seen[id] = struct{}{}
normalized = append(normalized, id)
}
return normalized
}
func normalizeDomainCertIDs(domains []string, certIDs []uint, domainCertIDs []uint) []uint {
if len(domainCertIDs) > 0 {
normalized := make([]uint, len(domainCertIDs))
copy(normalized, domainCertIDs)
return normalized
}
if len(certIDs) == 1 {
normalized := make([]uint, len(domains))
for index := range normalized {
normalized[index] = certIDs[0]
}
return normalized
}
if len(certIDs) == len(domains) {
normalized := make([]uint, len(certIDs))
copy(normalized, certIDs)
return normalized
}
return []uint{}
}
func certificatesByID(files []SupportFile) map[uint]string {
result := make(map[uint]string)
for _, file := range files {
if !strings.HasSuffix(file.Path, ".crt") {
continue
}
idText := strings.TrimSuffix(file.Path, ".crt")
var id uint
if _, err := fmt.Sscanf(idText, "%d", &id); err == nil && id != 0 {
result[id] = file.Content
}
}
return result
}
func validateCertificateCoverage(certPEM string, domains []string) error {
block, _ := pem.Decode([]byte(certPEM))
if block == nil {
return errors.New("certificate PEM is invalid")
}
leaf, err := x509.ParseCertificate(block.Bytes)
if err != nil {
return err
}
for _, domain := range domains {
if err := leaf.VerifyHostname(domain); err != nil {
return fmt.Errorf("certificate does not cover domain %s", domain)
}
}
return nil
}
func getPoWConfigForRoute(routeID uint, snapshot WAFDocument) (bool, *PoWConfig) {
for _, binding := range snapshot.Bindings {
if binding.RouteID != routeID {
continue
}
for _, groupID := range binding.RuleGroupIDs {
for _, group := range snapshot.RuleGroups {
if group.ID == groupID && group.PoWEnabled {
return true, group.PoWConfig
}
}
}
break
}
for _, group := range snapshot.RuleGroups {
if group.IsGlobal && group.PoWEnabled {
return true, group.PoWConfig
}
}
return false, nil
}
func uniqueUintIDs(values []uint) []uint {
seen := make(map[uint]struct{}, len(values))
result := make([]uint, 0, len(values))
for _, value := range values {
if value == 0 {
continue
}
if _, ok := seen[value]; ok {
continue
}
seen[value] = struct{}{}
result = append(result, value)
}
return result
}
func uniqueStrings(values []string) []string {
seen := make(map[string]struct{}, len(values))
result := make([]string, 0, len(values))
for _, value := range values {
item := strings.TrimSpace(value)
if item == "" {
continue
}
if _, ok := seen[item]; ok {
continue
}
seen[item] = struct{}{}
result = append(result, item)
}
return result
}
func resolveUpstreamServerName(originURL string, originHost string) string {
parsed, err := url.Parse(originURL)
if err != nil || !strings.EqualFold(parsed.Scheme, "https") {
return ""
}
if strings.TrimSpace(originHost) != "" {
parsedHost, err := url.Parse("//" + originHost)
if err == nil && parsedHost.Hostname() != "" {
return parsedHost.Hostname()
}
return originHost
}
return parsed.Hostname()
}
func renderServerNames(domains []string) string { return strings.Join(domains, " ") }
func onOff(value bool) string {
if value {
return "on"
}
return "off"
}
func quoteNginxStringLiteral(value string) string {
escaped := strings.ReplaceAll(value, `\`, `\\`)
escaped = strings.ReplaceAll(escaped, `"`, `\"`)
return fmt.Sprintf(`"%s"`, escaped)
}
func filepathToNginxPath(value string) string {
return strings.ReplaceAll(strings.TrimSpace(value), `\`, `/`)
}
func escapeNginxString(value string) string {
escaped := strings.ReplaceAll(value, `\`, `\\`)
escaped = strings.ReplaceAll(escaped, `"`, `\"`)
return escaped
}
func buildSuffixMatchPattern(rules []string) string {
parts := make([]string, 0, len(rules))
for _, rule := range rules {
parts = append(parts, regexp.QuoteMeta(rule))
}
return fmt.Sprintf("\\.(?:%s)$", strings.Join(parts, "|"))
}
func buildPathPrefixMatchPattern(rules []string) string {
parts := make([]string, 0, len(rules))
for _, rule := range rules {
trimmed := strings.TrimRight(rule, "/")
if trimmed == "" {
trimmed = "/"
}
if trimmed == "/" {
parts = append(parts, "/")
continue
}
parts = append(parts, fmt.Sprintf("%s(?:/|$)", regexp.QuoteMeta(trimmed)))
}
return fmt.Sprintf("^(?:%s)", strings.Join(parts, "|"))
}
func buildPathExactMatchPattern(rules []string) string {
parts := make([]string, 0, len(rules))
for _, rule := range rules {
parts = append(parts, regexp.QuoteMeta(rule))
}
return fmt.Sprintf("^(?:%s)$", strings.Join(parts, "|"))
}
+100
View File
@@ -0,0 +1,100 @@
package openresty
import (
"strings"
"testing"
)
func TestRenderPagesAPIProxyLocationBlock(t *testing.T) {
tests := []struct {
name string
deployment *PagesDeployment
expected []string
unexpected []string
}{
{
name: "nil deployment",
deployment: nil,
expected: []string{""},
},
{
name: "disabled proxy",
deployment: &PagesDeployment{
APIProxyEnabled: false,
APIProxyPath: "/api",
APIProxyPass: "http://127.0.0.1:8080",
},
expected: []string{""},
},
{
name: "enabled proxy without rewrite",
deployment: &PagesDeployment{
APIProxyEnabled: true,
APIProxyPath: "/api",
APIProxyPass: "http://127.0.0.1:8080",
APIProxyRewrite: "",
},
expected: []string{
"location /api {",
"proxy_pass http://127.0.0.1:8080;",
"proxy_http_version 1.1;",
"proxy_set_header Host $http_host;",
},
unexpected: []string{
"rewrite",
},
},
{
name: "enabled proxy with rewrite to root",
deployment: &PagesDeployment{
APIProxyEnabled: true,
APIProxyPath: "/api",
APIProxyPass: "http://127.0.0.1:8080",
APIProxyRewrite: "/",
},
expected: []string{
"location /api {",
"rewrite ^/api/(.*)$ /$1 break;",
"rewrite ^/api$ / break;",
"proxy_pass http://127.0.0.1:8080;",
},
},
{
name: "enabled proxy with rewrite to subpath",
deployment: &PagesDeployment{
APIProxyEnabled: true,
APIProxyPath: "/api",
APIProxyPass: "http://127.0.0.1:8080",
APIProxyRewrite: "/v2",
},
expected: []string{
"location /api {",
"rewrite ^/api/(.*)$ /v2/$1 break;",
"rewrite ^/api$ /v2 break;",
"proxy_pass http://127.0.0.1:8080;",
},
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
got := renderPagesAPIProxyLocationBlock(tt.deployment)
if len(tt.expected) == 1 && tt.expected[0] == "" {
if got != "" {
t.Fatalf("expected empty output, got: %q", got)
}
return
}
for _, exp := range tt.expected {
if !strings.Contains(got, exp) {
t.Errorf("expected output to contain %q, but got:\n%s", exp, got)
}
}
for _, unexp := range tt.unexpected {
if strings.Contains(got, unexp) {
t.Errorf("expected output NOT to contain %q, but got:\n%s", unexp, got)
}
}
})
}
}
+282
View File
@@ -0,0 +1,282 @@
package openresty
const (
CertDirPlaceholder = "__OPENFLARE_CERT_DIR__"
RouteConfigPlaceholder = "__OPENFLARE_ROUTE_CONFIG__"
AccessLogPlaceholder = "__OPENFLARE_ACCESS_LOG__"
ErrorLogPlaceholder = "__OPENFLARE_ERROR_LOG__"
LuaDirPlaceholder = "__OPENFLARE_LUA_DIR__"
ObservabilityListenPlaceholder = "__OPENFLARE_OBSERVABILITY_LISTEN__"
ObservabilityPortPlaceholder = "__OPENFLARE_OBSERVABILITY_PORT__"
PowStaticDirPlaceholder = "__OPENFLARE_POW_STATIC_DIR__"
PagesDirPlaceholder = "__OPENFLARE_PAGES_DIR__"
SourceConfigFileName = "openresty_config.json"
)
const (
cachePolicySuffix = "suffix"
cachePolicyPathPrefix = "path_prefix"
cachePolicyPathExact = "path_exact"
defaultWAFBlockStatus = 418
anubisStaticPrefix = "/.within.website/x/cmd/anubis/static/"
anubisAPIPrefix = "/.within.website/x/cmd/anubis/api/"
)
const defaultMainConfigTemplate = `# This file is generated by OpenFlare. Do not edit manually.
worker_processes {{OpenRestyWorkerProcesses}};
worker_rlimit_nofile {{OpenRestyWorkerRlimitNofile}};
pid logs/nginx.pid;
error_log {{OpenRestyErrorLogPath}} warn;
events {
worker_connections {{OpenRestyWorkerConnections}};
{{OpenRestyEventsUseDirective}}{{OpenRestyEventsMultiAcceptDirective}}}
http {
include mime.types;
default_type application/octet-stream;
{{OpenRestyConnectionUpgradeMap}}{{OpenRestyDefaultServerBlock}} log_format openflare_json escape=json '{"ts":"$time_iso8601","host":"$host","path":"$request_uri","remote_addr":"$remote_addr","status":$status,"request_time":$request_time,"bytes_sent":$body_bytes_sent,"request_length":$request_length}';
access_log {{OpenRestyAccessLogPath}} openflare_json;
sendfile on;
tcp_nopush on;
tcp_nodelay on;
keepalive_timeout {{OpenRestyKeepaliveTimeout}};
keepalive_requests {{OpenRestyKeepaliveRequests}};
client_header_timeout {{OpenRestyClientHeaderTimeout}};
client_body_timeout {{OpenRestyClientBodyTimeout}};
client_max_body_size {{OpenRestyClientMaxBodySize}};
large_client_header_buffers {{OpenRestyLargeClientHeaderBuffers}};
send_timeout {{OpenRestySendTimeout}};
proxy_connect_timeout {{OpenRestyProxyConnectTimeout}};
proxy_send_timeout {{OpenRestyProxySendTimeout}};
proxy_read_timeout {{OpenRestyProxyReadTimeout}};
proxy_request_buffering {{OpenRestyProxyRequestBuffering}};
proxy_buffering {{OpenRestyProxyBuffering}};
proxy_buffers {{OpenRestyProxyBuffers}};
proxy_buffer_size {{OpenRestyProxyBufferSize}};
proxy_busy_buffers_size {{OpenRestyProxyBusyBuffersSize}};
gzip {{OpenRestyGzip}};
gzip_min_length {{OpenRestyGzipMinLength}};
gzip_comp_level {{OpenRestyGzipCompLevel}};
{{OpenRestyResolverDirective}}{{OpenRestyCacheBlock}} include {{OpenRestyRouteConfigInclude}};
}
`
type SupportFile struct {
Path string `json:"path"`
Content string `json:"content"`
}
type CustomHeader struct {
Key string `json:"key"`
Value string `json:"value"`
}
type PoWListConfig struct {
IPs []string `json:"ips"`
IPCidrs []string `json:"ip_cidrs"`
Paths []string `json:"paths"`
PathRegexes []string `json:"path_regexes"`
UserAgents []string `json:"user_agents"`
}
type PoWConfig struct {
Difficulty int `json:"difficulty"`
Algorithm string `json:"algorithm"`
SessionTTL int `json:"session_ttl"`
ChallengeTTL int `json:"challenge_ttl"`
Whitelist PoWListConfig `json:"whitelist"`
Blacklist PoWListConfig `json:"blacklist"`
}
type Route struct {
ID uint `json:"id,omitempty"`
SiteName string `json:"site_name,omitempty"`
Domain string `json:"domain"`
Domains []string `json:"domains,omitempty"`
OriginURL string `json:"origin_url"`
OriginHost string `json:"origin_host,omitempty"`
Upstreams []string `json:"upstreams,omitempty"`
Enabled bool `json:"enabled"`
EnableHTTPS bool `json:"enable_https"`
CertID *uint `json:"cert_id,omitempty"`
CertIDs []uint `json:"cert_ids,omitempty"`
DomainCertIDs []uint `json:"domain_cert_ids,omitempty"`
RedirectHTTP bool `json:"redirect_http"`
LimitConnPerServer int `json:"limit_conn_per_server,omitempty"`
LimitConnPerIP int `json:"limit_conn_per_ip,omitempty"`
LimitRate string `json:"limit_rate,omitempty"`
CacheEnabled bool `json:"cache_enabled"`
CachePolicy string `json:"cache_policy,omitempty"`
CacheRules []string `json:"cache_rules,omitempty"`
CustomHeaders []CustomHeader `json:"custom_headers,omitempty"`
PoWEnabled bool `json:"pow_enabled,omitempty"`
PoWConfig *PoWConfig `json:"pow_config,omitempty"`
BasicAuthEnabled bool `json:"basic_auth_enabled,omitempty"`
BasicAuthUsername string `json:"basic_auth_username,omitempty"`
BasicAuthPassword string `json:"basic_auth_password,omitempty"`
Remark string `json:"remark,omitempty"`
UpstreamType string `json:"upstream_type,omitempty"`
PagesDeployment *PagesDeployment `json:"pages_deployment,omitempty"`
}
type PagesDeployment struct {
ProjectID uint `json:"project_id"`
ProjectSlug string `json:"project_slug"`
DeploymentID uint `json:"deployment_id"`
DeploymentNumber int `json:"deployment_number"`
Checksum string `json:"checksum"`
EntryFile string `json:"entry_file"`
SPAFallbackEnabled bool `json:"spa_fallback_enabled"`
SPAFallbackPath string `json:"spa_fallback_path"`
APIProxyEnabled bool `json:"api_proxy_enabled"`
APIProxyPath string `json:"api_proxy_path"`
APIProxyPass string `json:"api_proxy_pass"`
APIProxyRewrite string `json:"api_proxy_rewrite"`
LocalRoot string `json:"local_root"`
}
type WAFRuleGroup struct {
ID uint `json:"id"`
Name string `json:"name"`
Enabled bool `json:"enabled"`
IsGlobal bool `json:"is_global"`
BlockStatusCode int `json:"block_status_code"`
BlockResponseBody string `json:"block_response_body,omitempty"`
IPWhitelist []string `json:"ip_whitelist,omitempty"`
IPBlacklist []string `json:"ip_blacklist,omitempty"`
IPWhitelistGroups []uint `json:"ip_whitelist_group_ids,omitempty"`
IPBlacklistGroups []uint `json:"ip_blacklist_group_ids,omitempty"`
CountryWhitelist []string `json:"country_whitelist,omitempty"`
CountryBlacklist []string `json:"country_blacklist,omitempty"`
RegionWhitelist []string `json:"region_whitelist,omitempty"`
RegionBlacklist []string `json:"region_blacklist,omitempty"`
PoWEnabled bool `json:"pow_enabled,omitempty"`
PoWConfig *PoWConfig `json:"pow_config,omitempty"`
}
type WAFIPGroup struct {
ID uint `json:"id"`
Name string `json:"name"`
Type string `json:"type"`
Enabled bool `json:"enabled"`
IPList []string `json:"ip_list,omitempty"`
}
type WAFBinding struct {
RouteID uint `json:"route_id"`
SiteName string `json:"site_name"`
RuleGroupIDs []uint `json:"rule_group_ids"`
}
type WAFDocument struct {
RuleGroups []WAFRuleGroup `json:"rule_groups"`
IPGroups []WAFIPGroup `json:"ip_groups,omitempty"`
Bindings []WAFBinding `json:"bindings"`
}
type ConfigSnapshot struct {
DefaultServerReturnStatus int `json:"default_server_return_status"`
WorkerProcesses string `json:"worker_processes"`
WorkerConnections int `json:"worker_connections"`
WorkerRlimitNofile int `json:"worker_rlimit_nofile"`
EventsUse string `json:"events_use,omitempty"`
EventsMultiAcceptEnabled bool `json:"events_multi_accept_enabled"`
KeepaliveTimeout int `json:"keepalive_timeout"`
KeepaliveRequests int `json:"keepalive_requests"`
ClientHeaderTimeout int `json:"client_header_timeout"`
ClientBodyTimeout int `json:"client_body_timeout"`
ClientMaxBodySize string `json:"client_max_body_size"`
LargeClientHeaderBuffers string `json:"large_client_header_buffers"`
SendTimeout int `json:"send_timeout"`
ProxyConnectTimeout int `json:"proxy_connect_timeout"`
ProxySendTimeout int `json:"proxy_send_timeout"`
ProxyReadTimeout int `json:"proxy_read_timeout"`
WebsocketEnabled bool `json:"websocket_enabled"`
HTTP3Enabled bool `json:"http3_enabled"`
ProxyRequestBuffering bool `json:"proxy_request_buffering"`
ProxyBufferingEnabled bool `json:"proxy_buffering_enabled"`
ProxyBuffers string `json:"proxy_buffers"`
ProxyBufferSize string `json:"proxy_buffer_size"`
ProxyBusyBuffersSize string `json:"proxy_busy_buffers_size"`
GzipEnabled bool `json:"gzip_enabled"`
GzipMinLength int `json:"gzip_min_length"`
GzipCompLevel int `json:"gzip_comp_level"`
Resolvers string `json:"resolvers,omitempty"`
CacheEnabled bool `json:"cache_enabled"`
CachePath string `json:"cache_path,omitempty"`
CacheLevels string `json:"cache_levels"`
CacheInactive string `json:"cache_inactive"`
CacheMaxSize string `json:"cache_max_size"`
CacheKeyTemplate string `json:"cache_key_template"`
CacheLockEnabled bool `json:"cache_lock_enabled"`
CacheLockTimeout string `json:"cache_lock_timeout"`
CacheUseStale string `json:"cache_use_stale"`
MainConfigTemplate string `json:"main_config_template,omitempty"`
}
type Document struct {
Routes []Route `json:"routes"`
OpenRestyConfig ConfigSnapshot `json:"openresty_config"`
WAF WAFDocument `json:"waf"`
}
type Result struct {
MainConfig string
RouteConfig string
SupportFiles []SupportFile
Checksum string
}
type routeCacheConfig struct {
Enabled bool
Policy string
Rules []string
}
type routeLimitConfig struct {
LimitConnPerServer int
LimitConnPerIP int
LimitRate string
}
type routeUpstreamConfig struct {
Name string
Scheme string
ProxyPassURI string
Servers []string
UsesNamedUpstream bool
}
var requiredMainConfigTemplatePlaceholders = []string{
"{{OpenRestyWorkerProcesses}}",
"{{OpenRestyWorkerConnections}}",
"{{OpenRestyWorkerRlimitNofile}}",
"{{OpenRestyConnectionUpgradeMap}}",
"{{OpenRestyDefaultServerBlock}}",
"{{OpenRestyAccessLogPath}}",
"{{OpenRestyErrorLogPath}}",
"{{OpenRestyEventsUseDirective}}",
"{{OpenRestyEventsMultiAcceptDirective}}",
"{{OpenRestyKeepaliveTimeout}}",
"{{OpenRestyKeepaliveRequests}}",
"{{OpenRestyClientHeaderTimeout}}",
"{{OpenRestyClientBodyTimeout}}",
"{{OpenRestyClientMaxBodySize}}",
"{{OpenRestyLargeClientHeaderBuffers}}",
"{{OpenRestySendTimeout}}",
"{{OpenRestyProxyConnectTimeout}}",
"{{OpenRestyProxySendTimeout}}",
"{{OpenRestyProxyReadTimeout}}",
"{{OpenRestyProxyRequestBuffering}}",
"{{OpenRestyProxyBuffering}}",
"{{OpenRestyProxyBuffers}}",
"{{OpenRestyProxyBufferSize}}",
"{{OpenRestyProxyBusyBuffersSize}}",
"{{OpenRestyGzip}}",
"{{OpenRestyGzipMinLength}}",
"{{OpenRestyGzipCompLevel}}",
"{{OpenRestyCacheBlock}}",
"{{OpenRestyRouteConfigInclude}}",
}
+15
View File
@@ -0,0 +1,15 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package trace 提供 OpenTelemetry 链路追踪封装工具
package trace
import "go.opentelemetry.io/otel/propagation"
func newPropagator() propagation.TextMapPropagator {
return propagation.NewCompositeTextMapPropagator(
propagation.TraceContext{},
propagation.Baggage{},
)
}
+19
View File
@@ -0,0 +1,19 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package trace
import (
sdktrace "go.opentelemetry.io/otel/sdk/trace"
)
// ParentBasedRatioSampler 创建父级感知的概率采样器
// - 如果父 Span 已采样,则子 Span 也采样
// - 如果父 Span 未采样,则子 Span 也不采样
// - 如果是根 Span,按 samplingRate 概率采样
func ParentBasedRatioSampler(samplingRate float64) sdktrace.Sampler {
return sdktrace.ParentBased(
sdktrace.TraceIDRatioBased(samplingRate),
)
}
+62
View File
@@ -0,0 +1,62 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package trace
import (
"context"
"log"
"go.opentelemetry.io/otel"
"go.opentelemetry.io/otel/trace"
)
// Tracer 全局 OpenTelemetry Tracer 实例
var Tracer trace.Tracer
var shutdownFuncs []func(context.Context) error
func init() {
// 初始化 Propagator
prop := newPropagator()
otel.SetTextMapPropagator(prop)
// 初始化 Tracer 实例为 No-op 默认以避免未初始化前或测试环境崩溃
Tracer = otel.GetTracerProvider().Tracer("github.com/Rain-kl/OpenFlare")
}
// Config 链路追踪配置
type Config struct {
AppName string
SamplingRate float64
TracerName string
}
// Init 初始化 Tracer Provider 并关联全局 Tracer 实例
func Init(cfg Config) {
tracerProvider, err := newTracerProvider(cfg)
if err != nil {
log.Fatalf("[Trace] init trace provider failed: %v", err)
}
shutdownFuncs = append(shutdownFuncs, tracerProvider.Shutdown)
otel.SetTracerProvider(tracerProvider)
// 更新 Tracer
tracerName := cfg.TracerName
if tracerName == "" {
tracerName = "github.com/Rain-kl/OpenFlare"
}
Tracer = tracerProvider.Tracer(tracerName)
}
// Shutdown 关闭所有 Trace Provider
func Shutdown(ctx context.Context) {
for _, fn := range shutdownFuncs {
_ = fn(ctx)
}
}
// Start 创建一个新的 Trace Span
func Start(ctx context.Context, name string, opts ...trace.SpanStartOption) (context.Context, trace.Span) {
return Tracer.Start(ctx, name, opts...)
}
+53
View File
@@ -0,0 +1,53 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package trace
import (
"context"
"os"
"go.opentelemetry.io/otel/exporters/otlp/otlptrace/otlptracegrpc"
"go.opentelemetry.io/otel/sdk/resource"
sdktrace "go.opentelemetry.io/otel/sdk/trace"
semconv "go.opentelemetry.io/otel/semconv/v1.40.0"
)
func newTracerProvider(cfg Config) (*sdktrace.TracerProvider, error) {
// 获取主机名和容器信息
hostname, err := os.Hostname()
if err != nil {
return nil, err
}
// 初始化 Resource
r, err := resource.Merge(
resource.Default(),
resource.NewWithAttributes(
semconv.SchemaURL,
semconv.ServiceName(cfg.AppName),
semconv.HostName(hostname),
semconv.K8SNamespaceName(os.Getenv("KUBERNETES_NAMESPACE")),
semconv.K8SPodName(os.Getenv("KUBERNETES_POD_NAME")),
semconv.K8SPodUID(os.Getenv("KUBERNETES_POD_UID")),
),
)
if err != nil {
return nil, err
}
// 初始化 Exporter
traceExporter, err := otlptracegrpc.New(context.Background())
if err != nil {
return nil, err
}
// 初始化 Trace
tracerProvider := sdktrace.NewTracerProvider(
sdktrace.WithBatcher(traceExporter),
sdktrace.WithResource(r),
sdktrace.WithSampler(ParentBasedRatioSampler(cfg.SamplingRate)),
)
return tracerProvider, nil
}
+160
View File
@@ -0,0 +1,160 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package util provides generic utility functions.
package util
import (
"crypto/aes"
"crypto/cipher"
"crypto/ed25519"
"crypto/rand"
"encoding/base64"
"encoding/hex"
"errors"
"fmt"
"io"
)
const (
aesKeyLength = 32
errInvalidSignKey = "invalid sign key: %w"
errSignKeyLengthInvalid = "sign key must be 32 bytes (64 hex characters)"
errCreateCipherFailed = "failed to create cipher: %w"
errCreateGCMFailed = "failed to create GCM: %w"
errGenerateNonceFailed = "failed to generate nonce: %w"
errDecodeCiphertextFailed = "failed to decode ciphertext: %w"
errCiphertextTooShort = "ciphertext too short"
errDecryptFailed = "failed to decrypt: %w"
)
// Encrypt 使用 SignKey 加密字符串数据
// signKey: 64 字符 hex 编码的密钥(对应 32 字节,用于 AES-256)
// plaintext: 要加密的明文字符串
// return: base64 编码的密文
func Encrypt(signKey string, plaintext string) (string, error) {
return encryptBytes(signKey, []byte(plaintext))
}
// Decrypt 使用 SignKey 解密字符串数据
// signKey: 64 字符 hex 编码的密钥(对应 32 字节,用于 AES-256)
// ciphertext: base64 编码的密文
// return: 解密后的明文字符串
func Decrypt(signKey string, ciphertext string) (string, error) {
plaintext, err := decryptBytes(signKey, ciphertext)
if err != nil {
return "", err
}
return string(plaintext), nil
}
// encryptBytes 加密函数,处理字节数据
func encryptBytes(signKey string, plaintext []byte) (string, error) {
// 将 hex 编码的密钥转换为字节
key, err := hex.DecodeString(signKey)
if err != nil {
return "", fmt.Errorf(errInvalidSignKey, err)
}
if len(key) != aesKeyLength {
return "", errors.New(errSignKeyLengthInvalid)
}
// 创建 AES cipher
block, err := aes.NewCipher(key)
if err != nil {
return "", fmt.Errorf(errCreateCipherFailed, err)
}
// 使用 GCM 模式(Galois/Counter Mode)
gcm, err := cipher.NewGCM(block)
if err != nil {
return "", fmt.Errorf(errCreateGCMFailed, err)
}
// 生成随机 nonce
nonce := make([]byte, gcm.NonceSize())
if _, err := io.ReadFull(rand.Reader, nonce); err != nil {
return "", fmt.Errorf(errGenerateNonceFailed, err)
}
// 加密数据
ciphertext := gcm.Seal(nonce, nonce, plaintext, nil)
// 返回 base64 编码的密文
return Base64Encode(ciphertext), nil
}
// decryptBytes 解密函数,处理字节数据
func decryptBytes(signKey string, ciphertext string) ([]byte, error) {
// 将 hex 编码的密钥转换为字节
key, err := hex.DecodeString(signKey)
if err != nil {
return nil, fmt.Errorf(errInvalidSignKey, err)
}
if len(key) != aesKeyLength {
return nil, errors.New(errSignKeyLengthInvalid)
}
// 解码 base64 密文
data, err := Base64Decode(ciphertext)
if err != nil {
return nil, fmt.Errorf(errDecodeCiphertextFailed, err)
}
// 创建 AES cipher
block, err := aes.NewCipher(key)
if err != nil {
return nil, fmt.Errorf(errCreateCipherFailed, err)
}
// 使用 GCM 模式
gcm, err := cipher.NewGCM(block)
if err != nil {
return nil, fmt.Errorf(errCreateGCMFailed, err)
}
// 提取 nonce
nonceSize := gcm.NonceSize()
if len(data) < nonceSize {
return nil, errors.New(errCiphertextTooShort)
}
nonce, ciphertextBytes := data[:nonceSize], data[nonceSize:]
// 解密数据
plaintext, err := gcm.Open(nil, nonce, ciphertextBytes, nil)
if err != nil {
return nil, fmt.Errorf(errDecryptFailed, err)
}
return plaintext, nil
}
// Base64Encode Base64编码
func Base64Encode(data []byte) string {
return base64.StdEncoding.EncodeToString(data)
}
// Base64Decode Base64解码
func Base64Decode(encoded string) ([]byte, error) {
return base64.StdEncoding.DecodeString(encoded)
}
// Ed25519Verify 验证 Ed25519 签名
// publicKey: 32 字节的公钥(已解码的二进制格式)
// message: 待验证的原始消息
// signature: 64 字节的签名(已解码的二进制格式)
// return: 签名是否有效
func Ed25519Verify(publicKey, message, signature []byte) bool {
if len(publicKey) != ed25519.PublicKeySize {
return false
}
if len(signature) != ed25519.SignatureSize {
return false
}
return ed25519.Verify(publicKey, message, signature)
}
+20
View File
@@ -0,0 +1,20 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package util
import "golang.org/x/crypto/bcrypt"
// HashPassword 使用 bcrypt 对密码进行哈希处理
func HashPassword(password string) (string, error) {
hash, err := bcrypt.GenerateFromPassword([]byte(password), bcrypt.DefaultCost)
if err != nil {
return "", err
}
return string(hash), nil
}
// CheckPasswordHash 比较 bcrypt 哈希值与明文密码是否匹配
func CheckPasswordHash(hash, password string) bool {
return bcrypt.CompareHashAndPassword([]byte(hash), []byte(password)) == nil
}
+35
View File
@@ -0,0 +1,35 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package util
import "strings"
// emailPartsCount 邮箱地址由 @ 分割为两部分
const (
emailPartsCount = 2
emailLocalMinChars = 2 // 邮箱 local 部分掩码显示的最小字符数
)
// DerefString 安全地解引用字符串指针,nil 返回空字符串
func DerefString(s *string) string {
if s == nil {
return ""
}
return *s
}
// MaskEmail 安全脱敏用户的邮箱地址(例如 us***@example.com)
func MaskEmail(email string) string {
parts := strings.Split(email, "@")
if len(parts) != emailPartsCount {
return email
}
local := parts[0]
domain := parts[1]
if len(local) <= emailLocalMinChars {
return "**@" + domain
}
return local[:2] + "***" + local[len(local)-1:] + "@" + domain
}
+29
View File
@@ -0,0 +1,29 @@
// Copyright 2025 linux.do
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package util
import (
"crypto/rand"
"crypto/sha256"
"encoding/hex"
"io"
"github.com/google/uuid"
)
// uniqueIDBytes 生成唯一 ID 所需的随机字节长度
const uniqueIDBytes = 32
// GenerateUniqueIDSimple 生成 64 位唯一标识符
func GenerateUniqueIDSimple() string {
randomBytes := make([]byte, uniqueIDBytes)
if _, err := io.ReadFull(rand.Reader, randomBytes); err != nil {
// 如果随机数生成失败,使用 UUID 作为后备
uuidBytes := []byte(uuid.NewString())
hash := sha256.Sum256(uuidBytes)
copy(randomBytes, hash[:])
}
return hex.EncodeToString(randomBytes)
}
+53
View File
@@ -0,0 +1,53 @@
package utils
import (
"fmt"
"strconv"
)
var sizeKB = 1024
var sizeMB = sizeKB * 1024
var sizeGB = sizeMB * 1024
func Bytes2Size(num int64) string {
numStr := ""
unit := "B"
if num/int64(sizeGB) > 1 {
numStr = fmt.Sprintf("%.2f", float64(num)/float64(sizeGB))
unit = "GB"
} else if num/int64(sizeMB) > 1 {
numStr = fmt.Sprintf("%d", int(float64(num)/float64(sizeMB)))
unit = "MB"
} else if num/int64(sizeKB) > 1 {
numStr = fmt.Sprintf("%d", int(float64(num)/float64(sizeKB)))
unit = "KB"
} else {
numStr = fmt.Sprintf("%d", num)
}
return numStr + " " + unit
}
func Seconds2Time(num int) (time string) {
if num/31104000 > 0 {
time += strconv.Itoa(num/31104000) + " 年 "
num %= 31104000
}
if num/2592000 > 0 {
time += strconv.Itoa(num/2592000) + " 个月 "
num %= 2592000
}
if num/86400 > 0 {
time += strconv.Itoa(num/86400) + " 天 "
num %= 86400
}
if num/3600 > 0 {
time += strconv.Itoa(num/3600) + " 小时 "
num %= 3600
}
if num/60 > 0 {
time += strconv.Itoa(num/60) + " 分钟 "
num %= 60
}
time += strconv.Itoa(num) + " 秒"
return
}
+34
View File
@@ -0,0 +1,34 @@
package utils
import (
"log/slog"
"net"
"strings"
)
func GetIp() (ip string) {
ips, err := net.InterfaceAddrs()
if err != nil {
slog.Error("get interface addresses failed", "error", err)
return ip
}
for _, a := range ips {
if ipNet, ok := a.(*net.IPNet); ok && !ipNet.IP.IsLoopback() {
if ipNet.IP.To4() != nil {
ip = ipNet.IP.String()
if strings.HasPrefix(ip, "10") {
return
}
if strings.HasPrefix(ip, "172") {
return
}
if strings.HasPrefix(ip, "192.168") {
return
}
ip = ""
}
}
}
return
}
+76
View File
@@ -0,0 +1,76 @@
package utils
import (
"sort"
"strings"
"time"
)
// Unique returns a new slice containing only the unique elements of the input slice,
// preserving their original order.
func Unique[T comparable](slice []T) []T {
if slice == nil {
return nil
}
seen := make(map[T]struct{})
result := make([]T, 0)
for _, item := range slice {
if _, ok := seen[item]; ok {
continue
}
seen[item] = struct{}{}
result = append(result, item)
}
return result
}
// UniqueAndCleanStringSlice trims spaces, removes empty elements, and returns only the unique elements
// of the input string slice. It preserves order and returns nil if the resulting slice is empty.
func UniqueAndCleanStringSlice(slice []string) []string {
if slice == nil {
return nil
}
seen := make(map[string]struct{})
result := make([]string, 0)
for _, item := range slice {
trimmed := strings.TrimSpace(item)
if trimmed == "" {
continue
}
if _, ok := seen[trimmed]; ok {
continue
}
seen[trimmed] = struct{}{}
result = append(result, trimmed)
}
if len(result) == 0 {
return nil
}
return result
}
// IdentifiableTimeRecord represents a database record that has a unique ID and a primary timestamp field.
type IdentifiableTimeRecord interface {
GetID() uint
GetTime() time.Time
}
// SortAndLimitRecords sorts a slice of IdentifiableTimeRecord descendingly by their timestamp (and ID as a tie-breaker),
// and limits the slice to the specified size if limit > 0.
func SortAndLimitRecords[T IdentifiableTimeRecord](rows []T, limit int) []T {
if len(rows) == 0 {
return rows
}
sort.Slice(rows, func(i, j int) bool {
ti := rows[i].GetTime()
tj := rows[j].GetTime()
if ti.Equal(tj) {
return rows[i].GetID() > rows[j].GetID()
}
return ti.After(tj)
})
if limit > 0 && len(rows) > limit {
rows = rows[:limit]
}
return rows
}
+12
View File
@@ -0,0 +1,12 @@
package utils
import "strings"
// TrimStringFields trims leading and trailing spaces from all provided string pointers.
func TrimStringFields(fields ...*string) {
for _, f := range fields {
if f != nil {
*f = strings.TrimSpace(*f)
}
}
}
+15
View File
@@ -0,0 +1,15 @@
package utils
import "fmt"
func Interface2String(inter interface{}) string {
switch inter.(type) {
case string:
return inter.(string)
case int:
return fmt.Sprintf("%d", inter.(int))
case float64:
return fmt.Sprintf("%f", inter.(float64))
}
return "Not Implemented"
}
+208
View File
@@ -0,0 +1,208 @@
package utils
import (
"strconv"
"strings"
)
type VersionInfo struct {
Valid bool
IsDev bool
Numbers []int
Prerelease []string
GitDescribeDistance int
GitDescribeTail []string
}
func ParseVersionInfo(version string) VersionInfo {
normalized := strings.TrimSpace(strings.TrimPrefix(version, "v"))
if normalized == "" || normalized == "dev" {
return VersionInfo{IsDev: strings.EqualFold(normalized, "dev")}
}
base := normalized
prerelease := ""
if separator := strings.IndexRune(normalized, '-'); separator >= 0 {
base = normalized[:separator]
prerelease = normalized[separator+1:]
}
segments := strings.Split(base, ".")
parts := make([]int, 0, len(segments))
for _, segment := range segments {
segment = strings.TrimSpace(segment)
if segment == "" {
parts = append(parts, 0)
continue
}
numeric := strings.Builder{}
for _, r := range segment {
if r < '0' || r > '9' {
break
}
numeric.WriteRune(r)
}
if numeric.Len() == 0 {
parts = append(parts, 0)
continue
}
value, err := strconv.Atoi(numeric.String())
if err != nil {
return VersionInfo{}
}
parts = append(parts, value)
}
info := VersionInfo{Valid: len(parts) > 0, Numbers: parts}
if prerelease != "" {
identifiers := splitPrereleaseIdentifiers(prerelease)
if distance, tail, ok := parseGitDescribeIdentifiers(identifiers); ok {
info.GitDescribeDistance = distance
info.GitDescribeTail = tail
} else {
info.Prerelease = identifiers
}
}
return info
}
func parseGitDescribeIdentifiers(identifiers []string) (int, []string, bool) {
if len(identifiers) < 2 {
return 0, nil, false
}
distance, err := strconv.Atoi(strings.TrimSpace(identifiers[0]))
if err != nil || distance <= 0 {
return 0, nil, false
}
commitToken := strings.TrimSpace(identifiers[1])
if commitToken == "" || !strings.HasPrefix(strings.ToLower(commitToken), "g") {
return 0, nil, false
}
return distance, identifiers[1:], true
}
func splitPrereleaseIdentifiers(value string) []string {
parts := strings.FieldsFunc(strings.TrimSpace(value), func(r rune) bool {
return r == '.' || r == '-'
})
filtered := make([]string, 0, len(parts))
for _, part := range parts {
part = strings.TrimSpace(part)
if part != "" {
filtered = append(filtered, part)
}
}
return filtered
}
// CompareVersions compares two version strings.
// Returns -1 if left < right, 1 if left > right, and 0 if they are equal.
func CompareVersions(local, remote string) int {
left := ParseVersionInfo(local)
right := ParseVersionInfo(remote)
if left.IsDev {
if right.Valid {
return -1
}
return 0
}
if !left.Valid || !right.Valid {
return 0
}
maxLen := len(left.Numbers)
if len(right.Numbers) > maxLen {
maxLen = len(right.Numbers)
}
for index := 0; index < maxLen; index++ {
leftValue := 0
rightValue := 0
if index < len(left.Numbers) {
leftValue = left.Numbers[index]
}
if index < len(right.Numbers) {
rightValue = right.Numbers[index]
}
if leftValue < rightValue {
return -1
}
if leftValue > rightValue {
return 1
}
}
if left.GitDescribeDistance != right.GitDescribeDistance {
if left.GitDescribeDistance < right.GitDescribeDistance {
return -1
}
return 1
}
if left.GitDescribeDistance > 0 || right.GitDescribeDistance > 0 {
maxLen = len(left.GitDescribeTail)
if len(right.GitDescribeTail) > maxLen {
maxLen = len(right.GitDescribeTail)
}
for index := 0; index < maxLen; index++ {
if index >= len(left.GitDescribeTail) {
return -1
}
if index >= len(right.GitDescribeTail) {
return 1
}
if left.GitDescribeTail[index] < right.GitDescribeTail[index] {
return -1
}
if left.GitDescribeTail[index] > right.GitDescribeTail[index] {
return 1
}
}
return 0
}
if len(left.Prerelease) == 0 && len(right.Prerelease) == 0 {
return 0
}
if len(left.Prerelease) == 0 {
return 1
}
if len(right.Prerelease) == 0 {
return -1
}
maxLen = len(left.Prerelease)
if len(right.Prerelease) > maxLen {
maxLen = len(right.Prerelease)
}
for index := 0; index < maxLen; index++ {
if index >= len(left.Prerelease) {
return -1
}
if index >= len(right.Prerelease) {
return 1
}
leftPart := left.Prerelease[index]
rightPart := right.Prerelease[index]
leftNumber, leftErr := strconv.Atoi(leftPart)
rightNumber, rightErr := strconv.Atoi(rightPart)
switch {
case leftErr == nil && rightErr == nil:
if leftNumber < rightNumber {
return -1
}
if leftNumber > rightNumber {
return 1
}
case leftErr == nil:
return -1
case rightErr == nil:
return 1
default:
if leftPart < rightPart {
return -1
}
if leftPart > rightPart {
return 1
}
}
}
return 0
}
+222
View File
@@ -0,0 +1,222 @@
package wsclient
import (
"context"
"encoding/json"
"errors"
"log/slog"
"net"
"net/http"
"net/url"
"strings"
"time"
"golang.org/x/net/websocket"
)
type Config struct {
BaseURL string
Token string
Timeout time.Duration
HeaderKey string // e.g. "X-Agent-Token", "X-Tunnel-Token"
WSPath string // e.g. "/api/relay/ws", "/api/agent/ws", "/api/flared/ws"
}
type Client struct {
cfg Config
}
type WSMessage struct {
Type string `json:"type"`
Payload json.RawMessage `json:"payload,omitempty"`
}
type MessageHandler interface {
OnConnect(ctx context.Context) error
HandleMessage(ctx context.Context, msg WSMessage) error
OnClose(err error)
}
type Connection struct {
Conn *websocket.Conn
URL string
ReadTimeout time.Duration
}
func New(cfg Config) *Client {
cfg.BaseURL = strings.TrimRight(cfg.BaseURL, "/")
cfg.Token = strings.TrimSpace(cfg.Token)
cfg.HeaderKey = strings.TrimSpace(cfg.HeaderKey)
cfg.WSPath = strings.TrimSpace(cfg.WSPath)
return &Client{
cfg: cfg,
}
}
func (c *Client) SetToken(token string) {
c.cfg.Token = strings.TrimSpace(token)
}
func (c *Client) URL() string {
wsURL, err := c.BuildWebsocketURL()
if err != nil {
return ""
}
return wsURL
}
func (c *Client) BuildWebsocketURL() (string, error) {
parsed, err := url.Parse(c.cfg.BaseURL)
if err != nil {
return "", err
}
switch parsed.Scheme {
case "http":
parsed.Scheme = "ws"
case "https":
parsed.Scheme = "wss"
case "ws", "wss":
default:
return "", errors.New("server_url scheme must be http, https, ws, or wss")
}
wsPath := c.cfg.WSPath
if !strings.HasPrefix(wsPath, "/") {
wsPath = "/" + wsPath
}
parsed.Path = strings.TrimRight(parsed.Path, "/") + wsPath
parsed.RawQuery = ""
parsed.Fragment = ""
return parsed.String(), nil
}
func (c *Client) Connect(ctx context.Context) (*Connection, error) {
wsURL, err := c.BuildWebsocketURL()
if err != nil {
return nil, err
}
if c.cfg.Token == "" {
return nil, errors.New("ws token is empty")
}
origin := c.cfg.BaseURL
if origin == "" {
origin = "http://localhost"
}
config, err := websocket.NewConfig(wsURL, origin)
if err != nil {
return nil, err
}
config.Header = http.Header{}
if c.cfg.HeaderKey != "" {
config.Header.Set(c.cfg.HeaderKey, c.cfg.Token)
}
if c.cfg.Timeout > 0 {
config.Dialer = &net.Dialer{Timeout: c.cfg.Timeout}
}
slog.Debug("ws dialing server", "url", wsURL)
conn, err := config.DialContext(ctx)
if err != nil {
return nil, err
}
slog.Debug("ws dial succeeded", "url", wsURL)
return &Connection{Conn: conn, URL: wsURL, ReadTimeout: websocketReadTimeout(c.cfg.Timeout)}, nil
}
func (conn *Connection) SendMessage(msgType string, payload any) error {
if conn == nil || conn.Conn == nil {
return errors.New("ws connection is nil")
}
slog.Debug("ws sending message", "type", msgType)
// Create the outbound message wrapper
message := struct {
Type string `json:"type"`
Payload any `json:"payload,omitempty"`
}{
Type: msgType,
Payload: payload,
}
_ = conn.Conn.SetWriteDeadline(time.Now().Add(5 * time.Second))
return websocket.JSON.Send(conn.Conn, message)
}
func (conn *Connection) Receive(target any) error {
if conn == nil || conn.Conn == nil {
return errors.New("ws connection is nil")
}
if conn.ReadTimeout > 0 {
_ = conn.Conn.SetReadDeadline(time.Now().Add(conn.ReadTimeout))
}
err := websocket.JSON.Receive(conn.Conn, target)
if err != nil {
var netErr net.Error
if errors.As(err, &netErr) && netErr.Timeout() {
slog.Debug("ws receive timeout waiting for server message", "timeout", conn.ReadTimeout)
}
return err
}
return nil
}
func websocketReadTimeout(requestTimeout time.Duration) time.Duration {
timeout := requestTimeout * 6
if timeout < 75*time.Second {
return 75 * time.Second
}
return timeout
}
func (conn *Connection) RunReceiveLoop(ctx context.Context, handler MessageHandler) error {
doneChan := make(chan struct{})
defer close(doneChan)
go func() {
select {
case <-ctx.Done():
_ = conn.Close()
case <-doneChan:
}
}()
if err := handler.OnConnect(ctx); err != nil {
handler.OnClose(err)
return err
}
for {
select {
case <-ctx.Done():
return ctx.Err()
default:
}
var raw WSMessage
if err := conn.Receive(&raw); err != nil {
handler.OnClose(err)
return err
}
switch raw.Type {
case "ping":
slog.Debug("ws received ping from server, replying with pong")
if err := conn.SendMessage("pong", nil); err != nil {
slog.Error("ws send pong response failed", "error", err)
}
case "pong":
slog.Debug("ws received pong response from server")
default:
if err := handler.HandleMessage(ctx, raw); err != nil {
slog.Error("ws handler failed to process message", "type", raw.Type, "error", err)
return err
}
}
}
}
func (conn *Connection) Close() error {
if conn == nil || conn.Conn == nil {
return nil
}
return conn.Conn.Close()
}