mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-02 23:06:36 +08:00
refactor(backend): extract to pkg/cap
This commit is contained in:
@@ -1,257 +0,0 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package cap 提供人机验证(CAPTCHA)功能
|
||||
package cap
|
||||
|
||||
import (
|
||||
"crypto/hmac"
|
||||
"crypto/rand"
|
||||
"crypto/sha256"
|
||||
"encoding/base64"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
const (
|
||||
jwtHeaderB64 = "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9"
|
||||
jwtPartsCount = 3 // JWT 三段结构
|
||||
defaultChallengeCount = 50 // 默认 PoW 难题数
|
||||
defaultChallengeSize = 32 // 默认盐值长度
|
||||
defaultDifficulty = 4 // 默认难度
|
||||
defaultNonceLength = 25 // 随机 Nonce 字节长度
|
||||
defaultExpires = 10 * time.Minute // 默认过期时间
|
||||
)
|
||||
|
||||
// ChallengeConfig holds parameters for the PoW challenge
|
||||
type ChallengeConfig struct {
|
||||
Count int // Number of puzzles (c)
|
||||
Size int // Salt length (s)
|
||||
Difficulty int // Difficulty prefix length (d)
|
||||
Expires time.Duration // Challenge TTL
|
||||
}
|
||||
|
||||
// ChallengeResponse is returned to the client
|
||||
type ChallengeResponse struct {
|
||||
Challenge struct {
|
||||
C int `json:"c"`
|
||||
S int `json:"s"`
|
||||
D int `json:"d"`
|
||||
} `json:"challenge"`
|
||||
Token string `json:"token"`
|
||||
Expires int64 `json:"expires"` // ms timestamp
|
||||
}
|
||||
|
||||
// ChallengePayload represents the signed JWT payload
|
||||
type ChallengePayload struct {
|
||||
Nonce string `json:"n"`
|
||||
Count int `json:"c"`
|
||||
Size int `json:"s"`
|
||||
Difficulty int `json:"d"`
|
||||
Expires int64 `json:"exp"` // ms timestamp
|
||||
IssuedAt int64 `json:"iat"` // ms timestamp
|
||||
Scope string `json:"sk,omitempty"`
|
||||
}
|
||||
|
||||
// RedeemRequest payload sent by client
|
||||
type RedeemRequest struct {
|
||||
Token string `json:"token"`
|
||||
Solutions []int `json:"solutions"`
|
||||
}
|
||||
|
||||
// RedeemResponse returned to client after verification
|
||||
type RedeemResponse struct {
|
||||
Success bool `json:"success"`
|
||||
Token string `json:"token,omitempty"`
|
||||
Expires int64 `json:"expires,omitempty"`
|
||||
Error string `json:"error,omitempty"`
|
||||
}
|
||||
|
||||
func b64urlEncode(data []byte) string {
|
||||
return base64.RawURLEncoding.EncodeToString(data)
|
||||
}
|
||||
|
||||
func b64urlDecode(str string) ([]byte, error) {
|
||||
return base64.RawURLEncoding.DecodeString(str)
|
||||
}
|
||||
|
||||
func randomHex(byteLen int) string {
|
||||
bytes := make([]byte, byteLen)
|
||||
if _, err := rand.Read(bytes); err != nil {
|
||||
panic(err)
|
||||
}
|
||||
return hex.EncodeToString(bytes)
|
||||
}
|
||||
|
||||
func jwtSign(payload []byte, secret []byte) string {
|
||||
body := b64urlEncode(payload)
|
||||
sigInput := jwtHeaderB64 + "." + body
|
||||
|
||||
mac := hmac.New(sha256.New, secret)
|
||||
mac.Write([]byte(sigInput))
|
||||
sig := mac.Sum(nil)
|
||||
|
||||
return sigInput + "." + b64urlEncode(sig)
|
||||
}
|
||||
|
||||
func jwtVerify(token string, secret []byte) ([]byte, error) {
|
||||
parts := strings.Split(token, ".")
|
||||
if len(parts) != jwtPartsCount {
|
||||
return nil, errors.New(errInvalidTokenFormat)
|
||||
}
|
||||
if parts[0] != jwtHeaderB64 {
|
||||
return nil, errors.New(errInvalidHeader)
|
||||
}
|
||||
|
||||
sigInput := parts[0] + "." + parts[1]
|
||||
mac := hmac.New(sha256.New, secret)
|
||||
mac.Write([]byte(sigInput))
|
||||
expectedSig := mac.Sum(nil)
|
||||
|
||||
actualSig, err := b64urlDecode(parts[2])
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if !hmac.Equal(expectedSig, actualSig) {
|
||||
return nil, errors.New(errSignatureMismatch)
|
||||
}
|
||||
|
||||
payload, err := b64urlDecode(parts[1])
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return payload, nil
|
||||
}
|
||||
|
||||
func jwtSigHex(token string) string {
|
||||
parts := strings.Split(token, ".")
|
||||
if len(parts) != jwtPartsCount {
|
||||
return ""
|
||||
}
|
||||
sigBytes, err := b64urlDecode(parts[2])
|
||||
if err != nil {
|
||||
return ""
|
||||
}
|
||||
return hex.EncodeToString(sigBytes)
|
||||
}
|
||||
|
||||
// GenerateChallenge produces a new challenge and signed token
|
||||
func GenerateChallenge(secret []byte, conf ChallengeConfig, scope string) (*ChallengeResponse, error) {
|
||||
if conf.Count <= 0 {
|
||||
conf.Count = defaultChallengeCount
|
||||
}
|
||||
if conf.Size <= 0 {
|
||||
conf.Size = defaultChallengeSize
|
||||
}
|
||||
if conf.Difficulty <= 0 {
|
||||
conf.Difficulty = defaultDifficulty
|
||||
}
|
||||
if conf.Expires <= 0 {
|
||||
conf.Expires = defaultExpires
|
||||
}
|
||||
|
||||
now := time.Now().UnixNano() / int64(time.Millisecond)
|
||||
expires := now + int64(conf.Expires/time.Millisecond)
|
||||
|
||||
payload := ChallengePayload{
|
||||
Nonce: randomHex(defaultNonceLength),
|
||||
Count: conf.Count,
|
||||
Size: conf.Size,
|
||||
Difficulty: conf.Difficulty,
|
||||
Expires: expires,
|
||||
IssuedAt: now,
|
||||
Scope: scope,
|
||||
}
|
||||
|
||||
payloadBytes, err := json.Marshal(payload)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
token := jwtSign(payloadBytes, secret)
|
||||
|
||||
resp := &ChallengeResponse{
|
||||
Token: token,
|
||||
Expires: expires,
|
||||
}
|
||||
resp.Challenge.C = conf.Count
|
||||
resp.Challenge.S = conf.Size
|
||||
resp.Challenge.D = conf.Difficulty
|
||||
|
||||
return resp, nil
|
||||
}
|
||||
|
||||
// VerifyChallengeSolutions verifies client submitted solutions
|
||||
func VerifyChallengeSolutions(token string, solutions []int, secret []byte, expectedScope string) (*ChallengePayload, error) {
|
||||
payloadBytes, err := jwtVerify(token, secret)
|
||||
if err != nil {
|
||||
return nil, errors.New(errInvalidToken)
|
||||
}
|
||||
|
||||
var payload ChallengePayload
|
||||
if err := json.Unmarshal(payloadBytes, &payload); err != nil {
|
||||
return nil, errors.New(errInvalidToken)
|
||||
}
|
||||
|
||||
if expectedScope != "" && payload.Scope != expectedScope {
|
||||
return nil, errors.New(errScopeMismatch)
|
||||
}
|
||||
|
||||
now := time.Now().UnixNano() / int64(time.Millisecond)
|
||||
if payload.Expires < now {
|
||||
return nil, errors.New(errExpired)
|
||||
}
|
||||
|
||||
if len(solutions) != payload.Count {
|
||||
return nil, errors.New(errInvalidSolutions)
|
||||
}
|
||||
|
||||
tokenFnv := fnv1a(token)
|
||||
for i := 0; i < payload.Count; i++ {
|
||||
idxStr := strconv.Itoa(i + 1)
|
||||
saltSeed := fnv1aResume(tokenFnv, idxStr)
|
||||
targetSeed := fnv1aResume(saltSeed, "d")
|
||||
salt := prngFromHash(saltSeed, payload.Size)
|
||||
target := prngFromHash(targetSeed, payload.Difficulty)
|
||||
|
||||
hashInput := salt + strconv.Itoa(solutions[i])
|
||||
hashBytes := sha256.Sum256([]byte(hashInput))
|
||||
hashHex := hex.EncodeToString(hashBytes[:])
|
||||
|
||||
if !strings.HasPrefix(hashHex, target) {
|
||||
return nil, errors.New(errInvalidSolution)
|
||||
}
|
||||
}
|
||||
|
||||
return &payload, nil
|
||||
}
|
||||
|
||||
// Solve is a utility function to solve a challenge (mainly used for tests and reference implementation)
|
||||
func Solve(token string, count, size, difficulty int) []int {
|
||||
solutions := make([]int, count)
|
||||
tokenFnv := fnv1a(token)
|
||||
for i := 0; i < count; i++ {
|
||||
idxStr := strconv.Itoa(i + 1)
|
||||
saltSeed := fnv1aResume(tokenFnv, idxStr)
|
||||
targetSeed := fnv1aResume(saltSeed, "d")
|
||||
salt := prngFromHash(saltSeed, size)
|
||||
target := prngFromHash(targetSeed, difficulty)
|
||||
|
||||
for nonce := 0; nonce < 1000000; nonce++ {
|
||||
hashInput := salt + strconv.Itoa(nonce)
|
||||
hashBytes := sha256.Sum256([]byte(hashInput))
|
||||
hashHex := hex.EncodeToString(hashBytes[:])
|
||||
if strings.HasPrefix(hashHex, target) {
|
||||
solutions[i] = nonce
|
||||
break
|
||||
}
|
||||
}
|
||||
}
|
||||
return solutions
|
||||
}
|
||||
@@ -1,170 +0,0 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package cap
|
||||
|
||||
import (
|
||||
"context"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func TestCapFullFlow(t *testing.T) {
|
||||
secret := []byte("a-very-long-secret-key-at-least-16-bytes")
|
||||
store := NewMemoryStore(1 * time.Minute)
|
||||
|
||||
manager := NewManager(Config{
|
||||
Secret: secret,
|
||||
ChallengeCount: 3, // small count for fast test
|
||||
ChallengeSize: 32,
|
||||
ChallengeDifficulty: 3, // small difficulty for fast test
|
||||
ChallengeTTL: 5 * time.Second,
|
||||
TokenTTL: 10 * time.Second,
|
||||
}, store)
|
||||
|
||||
scope := "test-scope"
|
||||
ctx := context.Background()
|
||||
resp, err := manager.Generate(ctx, scope)
|
||||
if err != nil {
|
||||
t.Fatalf("Generate failed: %v", err)
|
||||
}
|
||||
|
||||
if resp.Challenge.C != 3 {
|
||||
t.Errorf("Expected count 3, got %d", resp.Challenge.C)
|
||||
}
|
||||
|
||||
// Solve the challenge (acting as client)
|
||||
solutions := Solve(resp.Token, resp.Challenge.C, resp.Challenge.S, resp.Challenge.D)
|
||||
|
||||
// Redeem
|
||||
redeemResp, err := manager.Redeem(ctx, resp.Token, solutions, scope)
|
||||
if err != nil {
|
||||
t.Fatalf("Redeem failed: %v", err)
|
||||
}
|
||||
if !redeemResp.Success {
|
||||
t.Fatalf("Redeem returned success=false: %s", redeemResp.Error)
|
||||
}
|
||||
if redeemResp.Token == "" {
|
||||
t.Fatalf("Expected token, got empty")
|
||||
}
|
||||
|
||||
// Verify the token
|
||||
valid, err := manager.VerifyToken(ctx, redeemResp.Token, scope)
|
||||
if err != nil {
|
||||
t.Fatalf("VerifyToken failed: %v", err)
|
||||
}
|
||||
if !valid {
|
||||
t.Fatalf("Expected redeem token to be valid")
|
||||
}
|
||||
|
||||
// Verify token is one-time use
|
||||
validAgain, err := manager.VerifyToken(ctx, redeemResp.Token, scope)
|
||||
if err != nil {
|
||||
t.Fatalf("VerifyToken second call failed: %v", err)
|
||||
}
|
||||
if validAgain {
|
||||
t.Fatalf("Expected redeem token to be single-use (invalidated after verification)")
|
||||
}
|
||||
}
|
||||
|
||||
// TestRedeemConcurrentRace verifies that when N goroutines simultaneously call
|
||||
// Redeem with the same challenge JWT, exactly one succeeds and the rest are
|
||||
// rejected with "already_redeemed". This guards against the TOCTOU fix.
|
||||
func TestRedeemConcurrentRace(t *testing.T) {
|
||||
const goroutines = 50
|
||||
|
||||
secret := []byte("race-test-secret-key-at-least-16-bytes")
|
||||
store := NewMemoryStore(1 * time.Minute)
|
||||
manager := NewManager(Config{
|
||||
Secret: secret,
|
||||
ChallengeCount: 1,
|
||||
ChallengeSize: 32,
|
||||
ChallengeDifficulty: 3,
|
||||
ChallengeTTL: 30 * time.Second,
|
||||
TokenTTL: 30 * time.Second,
|
||||
}, store)
|
||||
|
||||
ctx := context.Background()
|
||||
resp, err := manager.Generate(ctx, "login")
|
||||
if err != nil {
|
||||
t.Fatalf("Generate failed: %v", err)
|
||||
}
|
||||
solutions := Solve(resp.Token, resp.Challenge.C, resp.Challenge.S, resp.Challenge.D)
|
||||
|
||||
var (
|
||||
wg sync.WaitGroup
|
||||
success atomic.Int32
|
||||
barrier = make(chan struct{}) // synchronise goroutine start
|
||||
)
|
||||
|
||||
for i := 0; i < goroutines; i++ {
|
||||
wg.Add(1)
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
<-barrier // wait for the gun
|
||||
r, _ := manager.Redeem(ctx, resp.Token, solutions, "login")
|
||||
if r != nil && r.Success {
|
||||
success.Add(1)
|
||||
}
|
||||
}()
|
||||
}
|
||||
close(barrier) // fire all goroutines at once
|
||||
wg.Wait()
|
||||
|
||||
if n := success.Load(); n != 1 {
|
||||
t.Fatalf("Expected exactly 1 successful Redeem, got %d", n)
|
||||
}
|
||||
}
|
||||
|
||||
// TestVerifyTokenConcurrentRace verifies that when N goroutines simultaneously
|
||||
// call VerifyToken with the same cap token, exactly one succeeds and the rest
|
||||
// fail. This guards against the GetAndDelete fix.
|
||||
func TestVerifyTokenConcurrentRace(t *testing.T) {
|
||||
const goroutines = 50
|
||||
|
||||
secret := []byte("race-test-secret-key-at-least-16-bytes")
|
||||
store := NewMemoryStore(1 * time.Minute)
|
||||
manager := NewManager(Config{
|
||||
Secret: secret,
|
||||
ChallengeCount: 1,
|
||||
ChallengeSize: 32,
|
||||
ChallengeDifficulty: 3,
|
||||
ChallengeTTL: 30 * time.Second,
|
||||
TokenTTL: 30 * time.Second,
|
||||
}, store)
|
||||
|
||||
ctx := context.Background()
|
||||
resp, _ := manager.Generate(ctx, "login")
|
||||
solutions := Solve(resp.Token, resp.Challenge.C, resp.Challenge.S, resp.Challenge.D)
|
||||
redeemResp, err := manager.Redeem(ctx, resp.Token, solutions, "login")
|
||||
if err != nil || !redeemResp.Success {
|
||||
t.Fatalf("Redeem failed: %v %+v", err, redeemResp)
|
||||
}
|
||||
capToken := redeemResp.Token
|
||||
|
||||
var (
|
||||
wg sync.WaitGroup
|
||||
success atomic.Int32
|
||||
barrier = make(chan struct{})
|
||||
)
|
||||
|
||||
for i := 0; i < goroutines; i++ {
|
||||
wg.Add(1)
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
<-barrier
|
||||
ok, _ := manager.VerifyToken(ctx, capToken, "login")
|
||||
if ok {
|
||||
success.Add(1)
|
||||
}
|
||||
}()
|
||||
}
|
||||
close(barrier)
|
||||
wg.Wait()
|
||||
|
||||
if n := success.Load(); n != 1 {
|
||||
t.Fatalf("Expected exactly 1 successful VerifyToken, got %d", n)
|
||||
}
|
||||
}
|
||||
@@ -1,15 +0,0 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package cap
|
||||
|
||||
const (
|
||||
errInvalidTokenFormat = "invalid token format"
|
||||
errInvalidHeader = "invalid header"
|
||||
errSignatureMismatch = "signature mismatch"
|
||||
errInvalidToken = "invalid_token"
|
||||
errScopeMismatch = "scope_mismatch"
|
||||
errExpired = "expired"
|
||||
errInvalidSolutions = "invalid_solutions"
|
||||
errInvalidSolution = "invalid_solution"
|
||||
)
|
||||
@@ -1,278 +0,0 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package cap
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/config"
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
)
|
||||
|
||||
const (
|
||||
managerDefaultChallengeCount = 1
|
||||
managerDefaultChallengeSize = 32
|
||||
defaultChallengeDifficulty = 4
|
||||
defaultChallengeTTL = 10 * time.Minute
|
||||
defaultTokenTTL = 20 * time.Minute
|
||||
redeemTokenIDLength = 8 // 兑换 Token ID 字节长度
|
||||
redeemVerTokenLength = 15 // 兑换验证 Token 字节长度
|
||||
tokenPartsCount = 2 // 兑换 Token 由两部分组成
|
||||
valuePartsCount = 2 // 存储值由 scope 和过期时间组成
|
||||
)
|
||||
|
||||
// Config holds settings for the CAPTCHA manager
|
||||
type Config struct {
|
||||
Secret []byte // HMAC signing key
|
||||
ChallengeCount int // Number of PoW puzzles
|
||||
ChallengeSize int // Size of the salt string
|
||||
ChallengeDifficulty int // Length of difficulty target prefix
|
||||
ChallengeTTL time.Duration // Lifespan of the challenge JWT
|
||||
TokenTTL time.Duration // Lifespan of the redeem token
|
||||
}
|
||||
|
||||
// Manager orchestrates challenge generation and solution validation
|
||||
type Manager struct {
|
||||
conf Config
|
||||
store Store
|
||||
}
|
||||
|
||||
// NewManager creates a new CAPTCHA Manager
|
||||
func NewManager(conf Config, store Store) *Manager {
|
||||
if conf.ChallengeCount <= 0 {
|
||||
conf.ChallengeCount = managerDefaultChallengeCount
|
||||
}
|
||||
if conf.ChallengeSize <= 0 {
|
||||
conf.ChallengeSize = managerDefaultChallengeSize
|
||||
}
|
||||
if conf.ChallengeDifficulty <= 0 {
|
||||
conf.ChallengeDifficulty = defaultChallengeDifficulty
|
||||
}
|
||||
if conf.ChallengeTTL <= 0 {
|
||||
conf.ChallengeTTL = defaultChallengeTTL
|
||||
}
|
||||
if conf.TokenTTL <= 0 {
|
||||
conf.TokenTTL = defaultTokenTTL
|
||||
}
|
||||
return &Manager{
|
||||
conf: conf,
|
||||
store: store,
|
||||
}
|
||||
}
|
||||
|
||||
// Generate creates a challenge response
|
||||
func (m *Manager) Generate(ctx context.Context, scope string) (*ChallengeResponse, error) {
|
||||
c := ChallengeConfig{
|
||||
Count: m.getChallengeCount(ctx),
|
||||
Size: m.getChallengeSize(ctx),
|
||||
Difficulty: m.getChallengeDifficulty(ctx),
|
||||
Expires: m.getChallengeTTL(ctx),
|
||||
}
|
||||
return GenerateChallenge(m.conf.Secret, c, scope)
|
||||
}
|
||||
|
||||
// Redeem verifies PoW solutions and returns a one-time redeem token
|
||||
func (m *Manager) Redeem(ctx context.Context, token string, solutions []int, scope string) (*RedeemResponse, error) {
|
||||
sigHex := jwtSigHex(token)
|
||||
if sigHex == "" {
|
||||
return &RedeemResponse{Success: false, Error: "invalid_token"}, nil
|
||||
}
|
||||
|
||||
nonceKey := "cap:nonce:" + sigHex
|
||||
|
||||
// Atomically claim the nonce slot BEFORE verifying solutions.
|
||||
// SetNX returns true only when the key did not previously exist, so two
|
||||
// concurrent requests carrying the same JWT can never both succeed here.
|
||||
// TTL is set to the challenge's remaining lifetime so the slot auto-expires.
|
||||
payload, err := VerifyChallengeSolutions(token, solutions, m.conf.Secret, scope)
|
||||
if err != nil {
|
||||
return &RedeemResponse{Success: false, Error: err.Error()}, nil //nolint:nilerr // expected behavior: validation error is returned as response, not system error
|
||||
}
|
||||
|
||||
// Calculate remaining lifetime of the challenge JWT for the nonce TTL.
|
||||
now := time.Now().UnixNano() / int64(time.Millisecond)
|
||||
nonceTTL := time.Duration(payload.Expires-now) * time.Millisecond
|
||||
if nonceTTL < time.Second {
|
||||
nonceTTL = time.Second
|
||||
}
|
||||
|
||||
// Atomic claim: if another goroutine already redeemed this JWT the SetNX
|
||||
// will return false and we reject the request without issuing a token.
|
||||
set, err := m.store.SetNX(ctx, nonceKey, "1", nonceTTL)
|
||||
if err != nil {
|
||||
return &RedeemResponse{Success: false, Error: "nonce_store_error"}, err
|
||||
}
|
||||
if !set {
|
||||
return &RedeemResponse{Success: false, Error: "already_redeemed"}, nil
|
||||
}
|
||||
|
||||
// Generate a redeem token formatted as "id:verToken"
|
||||
id := randomHex(redeemTokenIDLength)
|
||||
verToken := randomHex(redeemVerTokenLength)
|
||||
verHashBytes := sha256.Sum256([]byte(verToken))
|
||||
verHashHex := hex.EncodeToString(verHashBytes[:])
|
||||
|
||||
tokenKey := "cap:token:" + id + ":" + verHashHex
|
||||
tokenTTL := m.getTokenTTL(ctx)
|
||||
tokenExpires := time.Now().Add(tokenTTL)
|
||||
|
||||
// Value stored is "expiresNano|scope"
|
||||
storeVal := strconv.FormatInt(tokenExpires.UnixNano(), 10) + "|" + scope
|
||||
|
||||
if err := m.store.Set(ctx, tokenKey, storeVal, tokenTTL); err != nil {
|
||||
return &RedeemResponse{Success: false, Error: "token_store_error"}, err
|
||||
}
|
||||
|
||||
return &RedeemResponse{
|
||||
Success: true,
|
||||
Token: id + ":" + verToken,
|
||||
Expires: tokenExpires.UnixNano() / int64(time.Millisecond),
|
||||
}, nil
|
||||
}
|
||||
|
||||
// VerifyToken validates and consumes the redeem token (single-use).
|
||||
// GetAndDelete is used so that retrieval and removal happen atomically:
|
||||
// two concurrent requests carrying the same token can never both see a value.
|
||||
func (m *Manager) VerifyToken(ctx context.Context, token string, expectedScope string) (bool, error) {
|
||||
if token == "" {
|
||||
return false, nil
|
||||
}
|
||||
parts := strings.Split(token, ":")
|
||||
if len(parts) != tokenPartsCount {
|
||||
return false, nil
|
||||
}
|
||||
id := parts[0]
|
||||
verToken := parts[1]
|
||||
|
||||
verHashBytes := sha256.Sum256([]byte(verToken))
|
||||
verHashHex := hex.EncodeToString(verHashBytes[:])
|
||||
|
||||
tokenKey := "cap:token:" + id + ":" + verHashHex
|
||||
|
||||
// Atomically retrieve-and-delete: the first caller gets the value, any
|
||||
// subsequent caller (even concurrent) receives (false, nil) immediately.
|
||||
val, exists, err := sGetAndDelete(ctx, m.store, tokenKey)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
if !exists {
|
||||
return false, nil
|
||||
}
|
||||
|
||||
valParts := strings.Split(val, "|")
|
||||
if len(valParts) != valuePartsCount {
|
||||
return false, nil
|
||||
}
|
||||
|
||||
expNano, err := strconv.ParseInt(valParts[0], 10, 64)
|
||||
if err != nil {
|
||||
return false, nil //nolint:nilerr // expected behavior: invalid format is treated as validation failure, not system error
|
||||
}
|
||||
tokenScope := valParts[1]
|
||||
|
||||
if expectedScope != "" && tokenScope != expectedScope {
|
||||
return false, nil
|
||||
}
|
||||
|
||||
if time.Now().UnixNano() > expNano {
|
||||
return false, nil // Expired
|
||||
}
|
||||
|
||||
return true, nil
|
||||
}
|
||||
|
||||
// sGetAndDelete safely calls store.GetAndDelete, treating a nil store as a miss.
|
||||
func sGetAndDelete(ctx context.Context, store Store, key string) (string, bool, error) {
|
||||
if store == nil {
|
||||
return "", false, nil
|
||||
}
|
||||
return store.GetAndDelete(ctx, key)
|
||||
}
|
||||
|
||||
func (m *Manager) getChallengeCount(ctx context.Context) int {
|
||||
val, err := model.GetIntByKey(ctx, model.ConfigKeyCapChallengeCount)
|
||||
if err != nil || val <= 0 {
|
||||
return m.conf.ChallengeCount
|
||||
}
|
||||
return val
|
||||
}
|
||||
|
||||
func (m *Manager) getChallengeSize(ctx context.Context) int {
|
||||
val, err := model.GetIntByKey(ctx, model.ConfigKeyCapChallengeSize)
|
||||
if err != nil || val <= 0 {
|
||||
return m.conf.ChallengeSize
|
||||
}
|
||||
return val
|
||||
}
|
||||
|
||||
func (m *Manager) getChallengeDifficulty(ctx context.Context) int {
|
||||
val, err := model.GetIntByKey(ctx, model.ConfigKeyCapChallengeDifficulty)
|
||||
if err != nil || val <= 0 {
|
||||
return m.conf.ChallengeDifficulty
|
||||
}
|
||||
return val
|
||||
}
|
||||
|
||||
func (m *Manager) getChallengeTTL(ctx context.Context) time.Duration {
|
||||
val, err := model.GetIntByKey(ctx, model.ConfigKeyCapChallengeTTL)
|
||||
if err != nil || val <= 0 {
|
||||
return m.conf.ChallengeTTL
|
||||
}
|
||||
return time.Duration(val) * time.Second
|
||||
}
|
||||
|
||||
func (m *Manager) getTokenTTL(ctx context.Context) time.Duration {
|
||||
val, err := model.GetIntByKey(ctx, model.ConfigKeyCapTokenTTL)
|
||||
if err != nil || val <= 0 {
|
||||
return m.conf.TokenTTL
|
||||
}
|
||||
return time.Duration(val) * time.Second
|
||||
}
|
||||
|
||||
var (
|
||||
defaultManager *Manager
|
||||
once sync.Once
|
||||
)
|
||||
|
||||
// GetDefaultManager yields the global singleton CAPTCHA manager
|
||||
func GetDefaultManager() *Manager {
|
||||
once.Do(func() {
|
||||
var secret []byte
|
||||
if config.Config != nil && config.Config.App.SessionSecret != "" {
|
||||
secret = []byte(config.Config.App.SessionSecret)
|
||||
} else {
|
||||
secret = []byte("default-captcha-secret-key-at-least-16-bytes")
|
||||
}
|
||||
|
||||
challengeCount := managerDefaultChallengeCount
|
||||
challengeSize := managerDefaultChallengeSize
|
||||
challengeDifficulty := defaultChallengeDifficulty
|
||||
challengeTTL := defaultChallengeTTL
|
||||
tokenTTL := defaultTokenTTL
|
||||
|
||||
var store Store
|
||||
if config.Config != nil && config.Config.Redis.Enabled && db.Redis != nil {
|
||||
store = NewRedisStore(db.Redis)
|
||||
} else {
|
||||
store = NewMemoryStore(1 * time.Minute)
|
||||
}
|
||||
|
||||
defaultManager = NewManager(Config{
|
||||
Secret: secret,
|
||||
ChallengeCount: challengeCount,
|
||||
ChallengeSize: challengeSize,
|
||||
ChallengeDifficulty: challengeDifficulty,
|
||||
ChallengeTTL: challengeTTL,
|
||||
TokenTTL: tokenTTL,
|
||||
}, store)
|
||||
})
|
||||
return defaultManager
|
||||
}
|
||||
@@ -1,49 +0,0 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package cap
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// fnv1a returns the 32-bit FNV-1a hash of a string
|
||||
//
|
||||
//nolint:mnd // FNV-1a 算法位移常量
|
||||
func fnv1a(str string) uint32 {
|
||||
var hash uint32 = 2166136261
|
||||
for i := 0; i < len(str); i++ {
|
||||
hash ^= uint32(str[i])
|
||||
hash += (hash << 1) + (hash << 4) + (hash << 7) + (hash << 8) + (hash << 24)
|
||||
}
|
||||
return hash
|
||||
}
|
||||
|
||||
// fnv1aResume resumes FNV-1a hashing from a given state
|
||||
//
|
||||
//nolint:mnd // FNV-1a 算法位移常量
|
||||
func fnv1aResume(state uint32, str string) uint32 {
|
||||
h := state
|
||||
for i := 0; i < len(str); i++ {
|
||||
h ^= uint32(str[i])
|
||||
h += (h << 1) + (h << 4) + (h << 7) + (h << 8) + (h << 24)
|
||||
}
|
||||
return h
|
||||
}
|
||||
|
||||
// prngFromHash generates a hex string of specified length using an initial hash state
|
||||
//
|
||||
//nolint:mnd // xorshift 算法位移常量
|
||||
func prngFromHash(initialHash uint32, length int) string {
|
||||
state := initialHash
|
||||
var result strings.Builder
|
||||
for result.Len() < length {
|
||||
state ^= state << 13
|
||||
state ^= state >> 17
|
||||
state ^= state << 5
|
||||
hexStr := fmt.Sprintf("%08x", state)
|
||||
result.WriteString(hexStr)
|
||||
}
|
||||
return result.String()[:length]
|
||||
}
|
||||
@@ -1,186 +0,0 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package cap
|
||||
|
||||
import (
|
||||
"context"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/redis/go-redis/v9"
|
||||
)
|
||||
|
||||
// Store defines the storage interface for challenge nonces and verification tokens
|
||||
type Store interface {
|
||||
Get(ctx context.Context, key string) (string, bool, error)
|
||||
Set(ctx context.Context, key string, val string, ttl time.Duration) error
|
||||
Delete(ctx context.Context, key string) error
|
||||
// SetNX atomically sets key=val with the given TTL only when the key does not
|
||||
// exist yet. It returns true when the key was actually written (i.e. this
|
||||
// caller "won" the race), and false when the key already existed.
|
||||
SetNX(ctx context.Context, key string, val string, ttl time.Duration) (bool, error)
|
||||
// GetAndDelete atomically retrieves the value of key and removes it in a
|
||||
// single operation. Returns ("", false, nil) when the key does not exist.
|
||||
GetAndDelete(ctx context.Context, key string) (string, bool, error)
|
||||
}
|
||||
|
||||
type memoryItem struct {
|
||||
value string
|
||||
expiresAt time.Time
|
||||
}
|
||||
|
||||
// MemoryStore is a thread-safe in-memory implementation of Store
|
||||
type MemoryStore struct {
|
||||
items map[string]memoryItem
|
||||
mu sync.Mutex // unified write-lock; promotes to exclusive for all ops
|
||||
}
|
||||
|
||||
// NewMemoryStore creates and initializes a new MemoryStore
|
||||
func NewMemoryStore(cleanupInterval time.Duration) *MemoryStore {
|
||||
store := &MemoryStore{
|
||||
items: make(map[string]memoryItem),
|
||||
}
|
||||
if cleanupInterval > 0 {
|
||||
go store.startCleanupLoop(cleanupInterval)
|
||||
}
|
||||
return store
|
||||
}
|
||||
|
||||
// Get 从 MemoryStore 获取指定 key 的值
|
||||
func (s *MemoryStore) Get(_ context.Context, key string) (string, bool, error) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
return s.getLocked(key)
|
||||
}
|
||||
|
||||
// getLocked is the internal helper – caller must hold s.mu.
|
||||
func (s *MemoryStore) getLocked(key string) (string, bool, error) {
|
||||
item, found := s.items[key]
|
||||
if !found {
|
||||
return "", false, nil
|
||||
}
|
||||
if time.Now().After(item.expiresAt) {
|
||||
delete(s.items, key)
|
||||
return "", false, nil
|
||||
}
|
||||
return item.value, true, nil
|
||||
}
|
||||
|
||||
// Set 向 MemoryStore 写入指定 key 的值
|
||||
func (s *MemoryStore) Set(_ context.Context, key string, val string, ttl time.Duration) error {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
s.items[key] = memoryItem{
|
||||
value: val,
|
||||
expiresAt: time.Now().Add(ttl),
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// Delete 从 MemoryStore 删除指定 key
|
||||
func (s *MemoryStore) Delete(_ context.Context, key string) error {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
delete(s.items, key)
|
||||
return nil
|
||||
}
|
||||
|
||||
// SetNX atomically sets key only when it is absent (or expired).
|
||||
// Returns true if the key was written by this call.
|
||||
func (s *MemoryStore) SetNX(_ context.Context, key string, val string, ttl time.Duration) (bool, error) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
|
||||
_, exists, _ := s.getLocked(key)
|
||||
if exists {
|
||||
return false, nil
|
||||
}
|
||||
s.items[key] = memoryItem{
|
||||
value: val,
|
||||
expiresAt: time.Now().Add(ttl),
|
||||
}
|
||||
return true, nil
|
||||
}
|
||||
|
||||
// GetAndDelete atomically retrieves and removes key in one critical section.
|
||||
func (s *MemoryStore) GetAndDelete(_ context.Context, key string) (string, bool, error) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
|
||||
val, exists, err := s.getLocked(key)
|
||||
if err != nil || !exists {
|
||||
return "", false, err
|
||||
}
|
||||
delete(s.items, key)
|
||||
return val, true, nil
|
||||
}
|
||||
|
||||
func (s *MemoryStore) startCleanupLoop(interval time.Duration) {
|
||||
ticker := time.NewTicker(interval)
|
||||
for range ticker.C {
|
||||
s.cleanupExpired()
|
||||
}
|
||||
}
|
||||
|
||||
func (s *MemoryStore) cleanupExpired() {
|
||||
now := time.Now()
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
for k, v := range s.items {
|
||||
if now.After(v.expiresAt) {
|
||||
delete(s.items, k)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// RedisStore is a GORM-compatible/standalone Redis-backed implementation of Store
|
||||
type RedisStore struct {
|
||||
client redis.UniversalClient
|
||||
}
|
||||
|
||||
// NewRedisStore creates a new RedisStore wrapping a redis.UniversalClient
|
||||
func NewRedisStore(client redis.UniversalClient) *RedisStore {
|
||||
return &RedisStore{
|
||||
client: client,
|
||||
}
|
||||
}
|
||||
|
||||
// Get 从 RedisStore 获取指定 key 的值
|
||||
func (s *RedisStore) Get(ctx context.Context, key string) (string, bool, error) {
|
||||
val, err := s.client.Get(ctx, key).Result()
|
||||
if err == redis.Nil {
|
||||
return "", false, nil
|
||||
}
|
||||
if err != nil {
|
||||
return "", false, err
|
||||
}
|
||||
return val, true, nil
|
||||
}
|
||||
|
||||
// Set 向 RedisStore 写入指定 key 的值
|
||||
func (s *RedisStore) Set(ctx context.Context, key string, val string, ttl time.Duration) error {
|
||||
return s.client.Set(ctx, key, val, ttl).Err()
|
||||
}
|
||||
|
||||
// Delete 从 RedisStore 删除指定 key
|
||||
func (s *RedisStore) Delete(ctx context.Context, key string) error {
|
||||
return s.client.Del(ctx, key).Err()
|
||||
}
|
||||
|
||||
// SetNX wraps Redis SET NX – returns true only when the key was newly created.
|
||||
func (s *RedisStore) SetNX(ctx context.Context, key string, val string, ttl time.Duration) (bool, error) {
|
||||
return s.client.SetNX(ctx, key, val, ttl).Result()
|
||||
}
|
||||
|
||||
// GetAndDelete wraps Redis GETDEL (available since Redis 6.2).
|
||||
func (s *RedisStore) GetAndDelete(ctx context.Context, key string) (string, bool, error) {
|
||||
val, err := s.client.GetDel(ctx, key).Result()
|
||||
if err == redis.Nil {
|
||||
return "", false, nil
|
||||
}
|
||||
if err != nil {
|
||||
return "", false, err
|
||||
}
|
||||
return val, true, nil
|
||||
}
|
||||
@@ -1,149 +0,0 @@
|
||||
// Copyright 2025 linux.do
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package util
|
||||
|
||||
import (
|
||||
"crypto/aes"
|
||||
"crypto/cipher"
|
||||
"crypto/ed25519"
|
||||
"crypto/rand"
|
||||
"encoding/base64"
|
||||
"encoding/hex"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
)
|
||||
|
||||
// aesKeyLength AES-256 密钥字节长度
|
||||
const aesKeyLength = 32
|
||||
|
||||
// Encrypt 使用 SignKey 加密字符串数据
|
||||
// signKey: 64 字符 hex 编码的密钥(对应 32 字节,用于 AES-256)
|
||||
// plaintext: 要加密的明文字符串
|
||||
// return: base64 编码的密文
|
||||
func Encrypt(signKey string, plaintext string) (string, error) {
|
||||
return encryptBytes(signKey, []byte(plaintext))
|
||||
}
|
||||
|
||||
// Decrypt 使用 SignKey 解密字符串数据
|
||||
// signKey: 64 字符 hex 编码的密钥(对应 32 字节,用于 AES-256)
|
||||
// ciphertext: base64 编码的密文
|
||||
// return: 解密后的明文字符串
|
||||
func Decrypt(signKey string, ciphertext string) (string, error) {
|
||||
plaintext, err := decryptBytes(signKey, ciphertext)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return string(plaintext), nil
|
||||
}
|
||||
|
||||
// encryptBytes 加密函数,处理字节数据
|
||||
func encryptBytes(signKey string, plaintext []byte) (string, error) {
|
||||
// 将 hex 编码的密钥转换为字节
|
||||
key, err := hex.DecodeString(signKey)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf(errInvalidSignKey, err)
|
||||
}
|
||||
if len(key) != aesKeyLength {
|
||||
return "", errors.New(errSignKeyLengthInvalid)
|
||||
}
|
||||
|
||||
// 创建 AES cipher
|
||||
block, err := aes.NewCipher(key)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf(errCreateCipherFailed, err)
|
||||
}
|
||||
|
||||
// 使用 GCM 模式(Galois/Counter Mode)
|
||||
gcm, err := cipher.NewGCM(block)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf(errCreateGCMFailed, err)
|
||||
}
|
||||
|
||||
// 生成随机 nonce
|
||||
nonce := make([]byte, gcm.NonceSize())
|
||||
if _, err := io.ReadFull(rand.Reader, nonce); err != nil {
|
||||
return "", fmt.Errorf(errGenerateNonceFailed, err)
|
||||
}
|
||||
|
||||
// 加密数据
|
||||
ciphertext := gcm.Seal(nonce, nonce, plaintext, nil)
|
||||
|
||||
// 返回 base64 编码的密文
|
||||
return Base64Encode(ciphertext), nil
|
||||
}
|
||||
|
||||
// decryptBytes 解密函数,处理字节数据
|
||||
func decryptBytes(signKey string, ciphertext string) ([]byte, error) {
|
||||
// 将 hex 编码的密钥转换为字节
|
||||
key, err := hex.DecodeString(signKey)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf(errInvalidSignKey, err)
|
||||
}
|
||||
if len(key) != aesKeyLength {
|
||||
return nil, errors.New(errSignKeyLengthInvalid)
|
||||
}
|
||||
|
||||
// 解码 base64 密文
|
||||
data, err := Base64Decode(ciphertext)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf(errDecodeCiphertextFailed, err)
|
||||
}
|
||||
|
||||
// 创建 AES cipher
|
||||
block, err := aes.NewCipher(key)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf(errCreateCipherFailed, err)
|
||||
}
|
||||
|
||||
// 使用 GCM 模式
|
||||
gcm, err := cipher.NewGCM(block)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf(errCreateGCMFailed, err)
|
||||
}
|
||||
|
||||
// 提取 nonce
|
||||
nonceSize := gcm.NonceSize()
|
||||
if len(data) < nonceSize {
|
||||
return nil, errors.New(errCiphertextTooShort)
|
||||
}
|
||||
|
||||
nonce, ciphertextBytes := data[:nonceSize], data[nonceSize:]
|
||||
|
||||
// 解密数据
|
||||
plaintext, err := gcm.Open(nil, nonce, ciphertextBytes, nil)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf(errDecryptFailed, err)
|
||||
}
|
||||
|
||||
return plaintext, nil
|
||||
}
|
||||
|
||||
// Base64Encode Base64编码
|
||||
func Base64Encode(data []byte) string {
|
||||
return base64.StdEncoding.EncodeToString(data)
|
||||
}
|
||||
|
||||
// Base64Decode Base64解码
|
||||
func Base64Decode(encoded string) ([]byte, error) {
|
||||
return base64.StdEncoding.DecodeString(encoded)
|
||||
}
|
||||
|
||||
// Ed25519Verify 验证 Ed25519 签名
|
||||
// publicKey: 32 字节的公钥(已解码的二进制格式)
|
||||
// message: 待验证的原始消息
|
||||
// signature: 64 字节的签名(已解码的二进制格式)
|
||||
// return: 签名是否有效
|
||||
func Ed25519Verify(publicKey, message, signature []byte) bool {
|
||||
if len(publicKey) != ed25519.PublicKeySize {
|
||||
return false
|
||||
}
|
||||
|
||||
if len(signature) != ed25519.SignatureSize {
|
||||
return false
|
||||
}
|
||||
|
||||
return ed25519.Verify(publicKey, message, signature)
|
||||
}
|
||||
@@ -7,12 +7,4 @@ const (
|
||||
errCreateHTTPRequestFailed = "创建HTTP请求失败: %w"
|
||||
errHTTPRequestFailed = "请求%s接口失败: %w"
|
||||
errInvalidCustomValue = "invalid value: %v"
|
||||
errInvalidSignKey = "invalid sign key: %w"
|
||||
errSignKeyLengthInvalid = "sign key must be 32 bytes (64 hex characters)"
|
||||
errCreateCipherFailed = "failed to create cipher: %w"
|
||||
errCreateGCMFailed = "failed to create GCM: %w"
|
||||
errGenerateNonceFailed = "failed to generate nonce: %w"
|
||||
errDecodeCiphertextFailed = "failed to decode ciphertext: %w"
|
||||
errCiphertextTooShort = "ciphertext too short"
|
||||
errDecryptFailed = "failed to decrypt: %w"
|
||||
)
|
||||
|
||||
@@ -12,7 +12,7 @@ import (
|
||||
"net/url"
|
||||
"time"
|
||||
|
||||
"go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp"
|
||||
"github.com/Rain-kl/Wavelet/pkg/httppool"
|
||||
)
|
||||
|
||||
// IsLocalhost 检查 URL 是否为 localhost
|
||||
@@ -35,12 +35,8 @@ const (
|
||||
|
||||
// 配置HTTP客户端 使用 otelhttp 自动注入 trace span
|
||||
var httpClient = &http.Client{
|
||||
Timeout: httpClientTimeout * time.Second,
|
||||
Transport: otelhttp.NewTransport(&http.Transport{
|
||||
MaxIdleConns: httpMaxIdleConns,
|
||||
MaxIdleConnsPerHost: httpMaxIdleConnsPerHost,
|
||||
IdleConnTimeout: httpIdleConnTimeout * time.Second,
|
||||
}),
|
||||
Timeout: httpClientTimeout * time.Second,
|
||||
Transport: httppool.DefaultTransport(),
|
||||
}
|
||||
|
||||
// SetHTTPClient 替换全局 HTTP 客户端实例
|
||||
|
||||
@@ -1,16 +0,0 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package mail 提供 SMTP 邮件发送功能。
|
||||
package mail
|
||||
|
||||
const (
|
||||
errDialTLSFailed = "dial tls failed: %w"
|
||||
errSMTPClientCreationFailed = "smtp client creation failed: %w"
|
||||
errSMTPAuthFailed = "smtp auth failed: %w"
|
||||
errSMTPMailCommandFailed = "smtp mail command failed: %w"
|
||||
errSMTPRcptCommandFailed = "smtp rcpt command failed: %w"
|
||||
errSMTPDataCommandFailed = "smtp data command failed: %w"
|
||||
errSMTPWritingBodyFailed = "smtp writing body failed: %w" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
|
||||
errSendMailFailed = "send mail failed: %w"
|
||||
)
|
||||
@@ -1,241 +0,0 @@
|
||||
// Copyright 2025 linux.do
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package mail
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"crypto/tls"
|
||||
"fmt"
|
||||
"net"
|
||||
"net/smtp"
|
||||
"strconv"
|
||||
"time"
|
||||
)
|
||||
|
||||
const (
|
||||
smtpSSLPort = 465 // SMTP SSL 端口
|
||||
smtpDialTimeout = 5 * time.Second // SMTP 连接超时
|
||||
smtpSessionDeadline = 10 * time.Second // SMTP 会话截止时间
|
||||
)
|
||||
|
||||
// Config represents SMTP mail configuration
|
||||
type Config struct {
|
||||
Host string
|
||||
Port int
|
||||
Username string
|
||||
Password string
|
||||
}
|
||||
|
||||
// SendMail sends an HTML email using the provided config and message details
|
||||
func SendMail(ctx context.Context, cfg Config, to string, subject, body string) error {
|
||||
return SendMailHTML(ctx, cfg, to, subject, body)
|
||||
}
|
||||
|
||||
// SendMailHTML sends an HTML format email
|
||||
func SendMailHTML(ctx context.Context, cfg Config, to string, subject, body string) error {
|
||||
addr := net.JoinHostPort(cfg.Host, strconv.Itoa(cfg.Port))
|
||||
|
||||
// Header & MIME settings for HTML email
|
||||
header := make(map[string]string)
|
||||
header["From"] = cfg.Username
|
||||
header["To"] = to
|
||||
header["Subject"] = subject
|
||||
header["MIME-Version"] = "1.0"
|
||||
header["Content-Type"] = "text/html; charset=UTF-8"
|
||||
|
||||
message := ""
|
||||
for k, v := range header {
|
||||
message += fmt.Sprintf("%s: %s\r\n", k, v)
|
||||
}
|
||||
message += "\r\n" + body
|
||||
|
||||
auth := smtp.PlainAuth("", cfg.Username, cfg.Password, cfg.Host)
|
||||
|
||||
// If using SSL port 465, we connection via TLS dial
|
||||
if cfg.Port == smtpSSLPort {
|
||||
return sendMailViaSSL(ctx, addr, auth, cfg, to, message)
|
||||
}
|
||||
|
||||
// For standard port (587 / 25), use smtp.SendMail directly (handles STARTTLS automatically if server supports it)
|
||||
err := smtp.SendMail(addr, auth, cfg.Username, []string{to}, []byte(message))
|
||||
if err != nil {
|
||||
return fmt.Errorf(errSendMailFailed, err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// sendMailViaSSL 通过 TLS 直接连接 SMTP SSL 端口发送邮件
|
||||
func sendMailViaSSL(ctx context.Context, addr string, auth smtp.Auth, cfg Config, to, message string) error {
|
||||
tlsConfig := &tls.Config{
|
||||
InsecureSkipVerify: true, //nolint:gosec // SMTP servers might use self-signed certificates
|
||||
ServerName: cfg.Host,
|
||||
}
|
||||
dialer := &net.Dialer{Timeout: smtpDialTimeout}
|
||||
tlsDialer := &tls.Dialer{
|
||||
NetDialer: dialer,
|
||||
Config: tlsConfig,
|
||||
}
|
||||
conn, err := tlsDialer.DialContext(ctx, "tcp", addr)
|
||||
if err != nil {
|
||||
return fmt.Errorf(errDialTLSFailed, err)
|
||||
}
|
||||
defer func() { _ = conn.Close() }()
|
||||
_ = conn.SetDeadline(time.Now().Add(smtpSessionDeadline))
|
||||
|
||||
client, err := smtp.NewClient(conn, cfg.Host)
|
||||
if err != nil {
|
||||
return fmt.Errorf(errSMTPClientCreationFailed, err)
|
||||
}
|
||||
defer func() { _ = client.Close() }()
|
||||
|
||||
if err = client.Auth(auth); err != nil {
|
||||
return fmt.Errorf(errSMTPAuthFailed, err)
|
||||
}
|
||||
if err = client.Mail(cfg.Username); err != nil {
|
||||
return fmt.Errorf(errSMTPMailCommandFailed, err)
|
||||
}
|
||||
if err = client.Rcpt(to); err != nil {
|
||||
return fmt.Errorf(errSMTPRcptCommandFailed, err)
|
||||
}
|
||||
|
||||
w, err := client.Data()
|
||||
if err != nil {
|
||||
return fmt.Errorf(errSMTPDataCommandFailed, err)
|
||||
}
|
||||
defer func() { _ = w.Close() }()
|
||||
|
||||
_, err = w.Write([]byte(message))
|
||||
if err != nil {
|
||||
return fmt.Errorf(errSMTPWritingBodyFailed, err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// SendMailWithLog sends a test email and records a detailed SMTP connection log
|
||||
func SendMailWithLog(ctx context.Context, cfg Config, to string, subject, body string) (string, error) {
|
||||
var logBuf bytes.Buffer
|
||||
logLine := func(dir string, format string, args ...interface{}) {
|
||||
fmt.Fprintf(&logBuf, "[%s] %s\n", dir, fmt.Sprintf(format, args...))
|
||||
}
|
||||
|
||||
addr := net.JoinHostPort(cfg.Host, strconv.Itoa(cfg.Port))
|
||||
logLine("System", "Connecting to %s...", addr)
|
||||
|
||||
var conn net.Conn
|
||||
var err error
|
||||
dialer := &net.Dialer{Timeout: smtpDialTimeout}
|
||||
if cfg.Port == smtpSSLPort {
|
||||
tlsConfig := &tls.Config{
|
||||
InsecureSkipVerify: true, //nolint:gosec // SMTP servers might use self-signed certificates
|
||||
ServerName: cfg.Host,
|
||||
}
|
||||
tlsDialer := &tls.Dialer{
|
||||
NetDialer: dialer,
|
||||
Config: tlsConfig,
|
||||
}
|
||||
conn, err = tlsDialer.DialContext(ctx, "tcp", addr)
|
||||
} else {
|
||||
conn, err = dialer.DialContext(ctx, "tcp", addr)
|
||||
}
|
||||
if err != nil {
|
||||
logLine("Error", "Connection failed: %v", err)
|
||||
return logBuf.String(), err
|
||||
}
|
||||
defer func() { _ = conn.Close() }()
|
||||
logLine("System", "Connected successfully.")
|
||||
|
||||
// Set a 10-second session deadline for read/write operations
|
||||
_ = conn.SetDeadline(time.Now().Add(smtpSessionDeadline))
|
||||
|
||||
client, err := smtp.NewClient(conn, cfg.Host)
|
||||
if err != nil {
|
||||
logLine("Error", "SMTP client handshake failed: %v", err)
|
||||
return logBuf.String(), err
|
||||
}
|
||||
defer func() { _ = client.Close() }()
|
||||
|
||||
// If not 465, support STARTTLS if available
|
||||
if cfg.Port != smtpSSLPort {
|
||||
if ok, _ := client.Extension("STARTTLS"); ok {
|
||||
logLine("C", "STARTTLS")
|
||||
tlsConfig := &tls.Config{
|
||||
InsecureSkipVerify: true, //nolint:gosec // SMTP servers might use self-signed certificates
|
||||
ServerName: cfg.Host,
|
||||
}
|
||||
if err = client.StartTLS(tlsConfig); err != nil {
|
||||
logLine("Error", "STARTTLS failed: %v", err)
|
||||
return logBuf.String(), err
|
||||
}
|
||||
logLine("S", "220 Ready to start TLS")
|
||||
}
|
||||
}
|
||||
|
||||
// Authentication
|
||||
if cfg.Username != "" && cfg.Password != "" {
|
||||
auth := smtp.PlainAuth("", cfg.Username, cfg.Password, cfg.Host)
|
||||
logLine("C", "AUTH PLAIN **********")
|
||||
if err = client.Auth(auth); err != nil {
|
||||
logLine("Error", "Authentication failed: %v", err)
|
||||
return logBuf.String(), err
|
||||
}
|
||||
logLine("S", "235 Authentication successful")
|
||||
}
|
||||
|
||||
// Mail command
|
||||
logLine("C", "MAIL FROM:<%s>", cfg.Username)
|
||||
if err = client.Mail(cfg.Username); err != nil {
|
||||
logLine("Error", "MAIL FROM command failed: %v", err)
|
||||
return logBuf.String(), err
|
||||
}
|
||||
logLine("S", "250 OK")
|
||||
|
||||
// Rcpt command
|
||||
logLine("C", "RCPT TO:<%s>", to)
|
||||
if err = client.Rcpt(to); err != nil {
|
||||
logLine("Error", "RCPT TO command failed: %v", err)
|
||||
return logBuf.String(), err
|
||||
}
|
||||
logLine("S", "250 OK")
|
||||
|
||||
// Data command
|
||||
logLine("C", "DATA")
|
||||
w, err := client.Data()
|
||||
if err != nil {
|
||||
logLine("Error", "DATA command failed: %v", err)
|
||||
return logBuf.String(), err
|
||||
}
|
||||
logLine("S", "354 Start mail input")
|
||||
|
||||
// Header & MIME settings for HTML email
|
||||
header := make(map[string]string)
|
||||
header["From"] = cfg.Username
|
||||
header["To"] = to
|
||||
header["Subject"] = subject
|
||||
header["MIME-Version"] = "1.0"
|
||||
header["Content-Type"] = "text/html; charset=UTF-8"
|
||||
|
||||
message := ""
|
||||
for k, v := range header {
|
||||
message += fmt.Sprintf("%s: %s\r\n", k, v)
|
||||
}
|
||||
message += "\r\n" + body
|
||||
|
||||
logLine("System", "Sending message body...")
|
||||
if _, err = w.Write([]byte(message)); err != nil {
|
||||
_ = w.Close()
|
||||
logLine("Error", "Writing message body failed: %v", err)
|
||||
return logBuf.String(), err
|
||||
}
|
||||
_ = w.Close()
|
||||
logLine("S", "250 OK")
|
||||
|
||||
logLine("C", "QUIT")
|
||||
_ = client.Quit()
|
||||
logLine("System", "Mail sent successfully!")
|
||||
|
||||
return logBuf.String(), nil
|
||||
}
|
||||
@@ -1,92 +0,0 @@
|
||||
// Copyright 2025 linux.do
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package mail
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"context"
|
||||
"net"
|
||||
"net/textproto"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestSendMailMock(t *testing.T) {
|
||||
// Start a mock SMTP server
|
||||
l, err := net.Listen("tcp", "127.0.0.1:0")
|
||||
if err != nil {
|
||||
t.Fatalf("failed to start mock smtp server: %v", err)
|
||||
}
|
||||
defer func() { _ = l.Close() }()
|
||||
|
||||
port := l.Addr().(*net.TCPAddr).Port
|
||||
|
||||
go func() {
|
||||
conn, err := l.Accept()
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
defer func() { _ = conn.Close() }()
|
||||
|
||||
writer := bufio.NewWriter(conn)
|
||||
reader := bufio.NewReader(conn)
|
||||
tp := textproto.NewReader(reader)
|
||||
|
||||
// 220 Ready
|
||||
_, _ = writer.WriteString("220 mock.smtp.com SMTP Ready\r\n")
|
||||
_ = writer.Flush()
|
||||
|
||||
// Read HELO/EHLO
|
||||
_, _ = tp.ReadLine()
|
||||
_, _ = writer.WriteString("250-mock.smtp.com\r\n250 AUTH PLAIN\r\n")
|
||||
_ = writer.Flush()
|
||||
|
||||
// Read AUTH PLAIN
|
||||
_, _ = tp.ReadLine()
|
||||
_, _ = writer.WriteString("235 Authentication successful\r\n")
|
||||
_ = writer.Flush()
|
||||
|
||||
// Read MAIL FROM
|
||||
_, _ = tp.ReadLine()
|
||||
_, _ = writer.WriteString("250 OK\r\n")
|
||||
_ = writer.Flush()
|
||||
|
||||
// Read RCPT TO
|
||||
_, _ = tp.ReadLine()
|
||||
_, _ = writer.WriteString("250 OK\r\n")
|
||||
_ = writer.Flush()
|
||||
|
||||
// Read DATA
|
||||
_, _ = tp.ReadLine()
|
||||
_, _ = writer.WriteString("354 Start mail input\r\n")
|
||||
_ = writer.Flush()
|
||||
|
||||
// Read body lines until dot
|
||||
for {
|
||||
line, err := tp.ReadLine()
|
||||
if err != nil || line == "." {
|
||||
break
|
||||
}
|
||||
}
|
||||
_, _ = writer.WriteString("250 OK\r\n")
|
||||
_ = writer.Flush()
|
||||
|
||||
// Read QUIT
|
||||
_, _ = tp.ReadLine()
|
||||
_, _ = writer.WriteString("221 Bye\r\n")
|
||||
_ = writer.Flush()
|
||||
}()
|
||||
|
||||
cfg := Config{
|
||||
Host: "127.0.0.1",
|
||||
Port: port,
|
||||
Username: "test@example.com",
|
||||
Password: "password",
|
||||
}
|
||||
|
||||
err = SendMail(context.Background(), cfg, "recipient@example.com", "Test Subject", "<h1>Test Body</h1>")
|
||||
if err != nil {
|
||||
t.Errorf("failed to send mail: %v", err)
|
||||
}
|
||||
}
|
||||
@@ -1,20 +0,0 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package util
|
||||
|
||||
import "golang.org/x/crypto/bcrypt"
|
||||
|
||||
// HashPassword 使用 bcrypt 对密码进行哈希处理
|
||||
func HashPassword(password string) (string, error) {
|
||||
hash, err := bcrypt.GenerateFromPassword([]byte(password), bcrypt.DefaultCost)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return string(hash), nil
|
||||
}
|
||||
|
||||
// CheckPasswordHash 比较 bcrypt 哈希值与明文密码是否匹配
|
||||
func CheckPasswordHash(hash, password string) bool {
|
||||
return bcrypt.CompareHashAndPassword([]byte(hash), []byte(password)) == nil
|
||||
}
|
||||
@@ -1,35 +0,0 @@
|
||||
// Copyright 2025 linux.do
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package util
|
||||
|
||||
import "strings"
|
||||
|
||||
// emailPartsCount 邮箱地址由 @ 分割为两部分
|
||||
const (
|
||||
emailPartsCount = 2
|
||||
emailLocalMinChars = 2 // 邮箱 local 部分掩码显示的最小字符数
|
||||
)
|
||||
|
||||
// DerefString 安全地解引用字符串指针,nil 返回空字符串
|
||||
func DerefString(s *string) string {
|
||||
if s == nil {
|
||||
return ""
|
||||
}
|
||||
return *s
|
||||
}
|
||||
|
||||
// MaskEmail 安全脱敏用户的邮箱地址(例如 us***@example.com)
|
||||
func MaskEmail(email string) string {
|
||||
parts := strings.Split(email, "@")
|
||||
if len(parts) != emailPartsCount {
|
||||
return email
|
||||
}
|
||||
local := parts[0]
|
||||
domain := parts[1]
|
||||
if len(local) <= emailLocalMinChars {
|
||||
return "**@" + domain
|
||||
}
|
||||
return local[:2] + "***" + local[len(local)-1:] + "@" + domain
|
||||
}
|
||||
@@ -1,29 +0,0 @@
|
||||
// Copyright 2025 linux.do
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package util
|
||||
|
||||
import (
|
||||
"crypto/rand"
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"io"
|
||||
|
||||
"github.com/google/uuid"
|
||||
)
|
||||
|
||||
// uniqueIDBytes 生成唯一 ID 所需的随机字节长度
|
||||
const uniqueIDBytes = 32
|
||||
|
||||
// GenerateUniqueIDSimple 生成 64 位唯一标识符
|
||||
func GenerateUniqueIDSimple() string {
|
||||
randomBytes := make([]byte, uniqueIDBytes)
|
||||
if _, err := io.ReadFull(rand.Reader, randomBytes); err != nil {
|
||||
// 如果随机数生成失败,使用 UUID 作为后备
|
||||
uuidBytes := []byte(uuid.NewString())
|
||||
hash := sha256.Sum256(uuidBytes)
|
||||
copy(randomBytes, hash[:])
|
||||
}
|
||||
return hex.EncodeToString(randomBytes)
|
||||
}
|
||||
Reference in New Issue
Block a user