refactor(auth): merge cap domain plugin into auth

This commit is contained in:
ryan
2026-09-03 08:58:34 +08:00
parent 6b28adfacf
commit 30ab8810cc
23 changed files with 393 additions and 544 deletions
+23
View File
@@ -0,0 +1,23 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package auth
// HTTP 响应错误文案
const (
errCapTokenMissing = "验证码验证失败,缺少验证码凭证" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
errCapTokenInvalidOrExpired = "验证码校验失败或已过期,请重试" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
errCapNotConfigured = "captcha is not configured"
errChallengeGenerateFailed = "生成验证难题失败,请稍后再试"
errInvalidRequestParams = "无效的参数"
errSolutionVerifyFailed = "校验验证解答失败,请稍后再试"
)
// Redeem 结果码,属于 redeem 响应 JSON 的对外契约取值,禁止改写取值
const (
redeemErrInvalidToken = "invalid_token"
redeemErrNonceStoreFailed = "nonce_store_error"
redeemErrAlreadyRedeemed = "already_redeemed"
redeemErrSettingsLoad = "settings_load_error"
redeemErrTokenStoreFailed = "token_store_error" //nolint:gosec // error code, not hardcoded credentials
)
@@ -0,0 +1,88 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package auth
import (
"Wavelet/pkg/logger"
"Wavelet/pkg/response"
"net/http"
"github.com/gin-gonic/gin"
)
// Challenge 生成 PoW 人机验证难题
// @Summary 生成人机验证难题
// @Description 客户端获取 PoW 难题和签名的 JWT Token,并在后台计算。
// @Tags cap
// @Accept json
// @Produce json
// @Param request body challengeRequest false "可选范围限制参数"
// @Success 200 {object} response.Any{data=auth.ChallengeResponse} "成功返回 PoW 难题"
// @Failure 500 {object} response.Any "内部服务错误"
// @Router /api/v1/cap/challenge [get]
// @Router /api/v1/cap/challenge [post]
func Challenge(c *gin.Context) {
var req challengeRequest
_ = c.ShouldBind(&req) // 允许不传 body,默认使用 login scope
if req.Scope == "" {
req.Scope = "login"
}
mgr := GetDefaultCapManager()
if mgr == nil {
response.AbortInternal(c, errCapNotConfigured)
return
}
resp, err := mgr.Generate(c.Request.Context(), req.Scope)
if err != nil {
logger.ErrorF(c.Request.Context(), "Generate cap challenge failed: %v", err)
response.AbortInternal(c, errChallengeGenerateFailed)
return
}
c.JSON(http.StatusOK, response.OK(resp))
}
// Redeem 提交 PoW 解答并兑换一次性凭证 Token
// @Summary 校验人机验证解答
// @Description 提交 PoW 解答进行核销,成功后返回一次性 X-Cap-Token 凭证
// @Tags cap
// @Accept json
// @Produce json
// @Param request body redeemRequest true "难题 Token 与解答 solutions 数组"
// @Success 200 {object} response.Any{data=auth.RedeemResponse} "核销成功,返回 X-Cap-Token"
// @Failure 400 {object} response.Any "参数错误或核销失败"
// @Failure 500 {object} response.Any "内部服务错误"
// @Router /api/v1/cap/redeem [post]
func Redeem(c *gin.Context) {
var req redeemRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.AbortBadRequest(c, errInvalidRequestParams)
return
}
if req.Scope == "" {
req.Scope = "login"
}
mgr := GetDefaultCapManager()
if mgr == nil {
response.AbortInternal(c, errCapNotConfigured)
return
}
resp, err := mgr.Redeem(c.Request.Context(), req.Token, req.Solutions, req.Scope)
if err != nil {
logger.ErrorF(c.Request.Context(), "Redeem cap solutions failed: %v", err)
response.AbortInternal(c, errSolutionVerifyFailed)
return
}
if !resp.Success {
response.AbortBadRequest(c, resp.Error)
return
}
c.JSON(http.StatusOK, response.OK(resp))
}
@@ -0,0 +1,38 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package auth
import (
"Wavelet/pkg/response"
"github.com/gin-gonic/gin"
)
// VerifyCaptchaMiddleware returns a Gin middleware that checks and consumes the X-Cap-Token header.
func VerifyCaptchaMiddleware(mgr *CaptchaManager, scope string) gin.HandlerFunc {
return func(c *gin.Context) {
if !CapProtectionEnabled(c.Request.Context()) {
c.Next()
return
}
if mgr == nil {
response.AbortBadRequest(c, errCapTokenInvalidOrExpired)
return
}
token := c.GetHeader("X-Cap-Token")
if token == "" {
response.AbortBadRequest(c, errCapTokenMissing)
return
}
valid, err := mgr.VerifyToken(c.Request.Context(), token, scope)
if err != nil || !valid {
response.AbortBadRequest(c, errCapTokenInvalidOrExpired)
return
}
c.Next()
}
}
@@ -0,0 +1,32 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package auth
import (
"Wavelet/pkg/response"
"net/http"
"net/http/httptest"
"testing"
"github.com/gin-gonic/gin"
)
func TestVerifyMiddlewareMissingTokenIsBadRequest(t *testing.T) {
gin.SetMode(gin.TestMode)
restore := InstallCapTestRuntimeSettings(CapRuntimeSettings{LoginEnabled: true})
t.Cleanup(restore)
engine := gin.New()
engine.Use(response.ErrorHandlerMiddleware())
engine.POST("/register", VerifyCaptchaMiddleware(GetDefaultCapManager(), "register"), func(c *gin.Context) {
c.Status(http.StatusOK)
})
req := httptest.NewRequest(http.MethodPost, "/register", nil)
rec := httptest.NewRecorder()
engine.ServeHTTP(rec, req)
if rec.Code != http.StatusBadRequest {
t.Errorf("VerifyCaptchaMiddleware() status = %d, want %d", rec.Code, http.StatusBadRequest)
}
}
+37
View File
@@ -0,0 +1,37 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package auth
import (
"Wavelet/plugins/domain/auth/pow"
)
// ChallengeResponse is a local type alias for the pow.ChallengeResponse struct
type ChallengeResponse = pow.ChallengeResponse
// challengeRequest is the CAPTCHA challenge request payload.
type challengeRequest struct {
Scope string `json:"scope" form:"scope"`
}
// redeemRequest is the CAPTCHA redeem request payload.
type redeemRequest struct {
Token string `json:"token" binding:"required"`
Solutions []int `json:"solutions" binding:"required"`
Scope string `json:"scope" form:"scope"`
}
// RedeemResponse is returned to the client on redeem.
type RedeemResponse struct {
Success bool `json:"success"`
Token string `json:"token,omitempty"`
Expires int64 `json:"expires,omitempty"`
Error string `json:"error,omitempty"`
}
// capConfigRecord maps the columns selected from the system config table.
type capConfigRecord struct {
Key string `gorm:"column:key"`
Value string `gorm:"column:value"`
}
@@ -0,0 +1,197 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package auth
import (
"context"
"errors"
"strconv"
"sync/atomic"
"time"
"golang.org/x/sync/singleflight"
)
const (
defaultCapChallengeCount = 1
defaultCapChallengeSize = 32
defaultCapChallengeDifficulty = 4
defaultCapChallengeTTL = 10 * time.Minute
defaultCapTokenTTL = 20 * time.Minute
)
// CapRuntimeSettings is the parsed CAPTCHA runtime configuration loaded from system_configs.
type CapRuntimeSettings struct {
LoginEnabled bool
ChallengeCount int
ChallengeSize int
ChallengeDifficulty int
ChallengeTTL time.Duration
TokenTTL time.Duration
}
// CAP 动态配置键常量
const (
ConfigKeyCapLoginEnabled = "cap_login_enabled"
ConfigKeyCapChallengeCount = "cap_challenge_count"
ConfigKeyCapChallengeSize = "cap_challenge_size"
ConfigKeyCapChallengeDifficulty = "cap_challenge_difficulty"
ConfigKeyCapChallengeTTL = "cap_challenge_ttl"
// ConfigKeyCapTokenTTL 验证码 Token 过期时间键
// #nosec G101
ConfigKeyCapTokenTTL = "cap_token_ttl"
)
var capRuntimeConfigKeys = []string{
ConfigKeyCapLoginEnabled,
ConfigKeyCapChallengeCount,
ConfigKeyCapChallengeSize,
ConfigKeyCapChallengeDifficulty,
ConfigKeyCapChallengeTTL,
ConfigKeyCapTokenTTL,
}
var capRuntimeConfigKeySet = func() map[string]struct{} {
set := make(map[string]struct{}, len(capRuntimeConfigKeys))
for _, key := range capRuntimeConfigKeys {
set[key] = struct{}{}
}
return set
}()
type capRuntimeSettingsStore struct {
snapshot atomic.Pointer[CapRuntimeSettings]
loadGroup singleflight.Group
}
var capSettingsStore = &capRuntimeSettingsStore{}
// IsCapRuntimeConfigKey reports whether a system config key affects CAPTCHA runtime settings.
func IsCapRuntimeConfigKey(key string) bool {
_, ok := capRuntimeConfigKeySet[key]
return ok
}
// CurrentCapSettings returns the cached CAPTCHA runtime settings snapshot.
func CurrentCapSettings(ctx context.Context) (CapRuntimeSettings, error) {
return capSettingsStore.current(ctx)
}
// CapProtectionEnabled reports whether CAPTCHA verification is required for protected routes.
func CapProtectionEnabled(ctx context.Context) bool {
settings, err := CurrentCapSettings(ctx)
if err != nil {
return false
}
return settings.LoginEnabled
}
// InvalidateCapRuntimeSettings drops the in-process CAPTCHA settings snapshot.
func InvalidateCapRuntimeSettings() {
capSettingsStore.snapshot.Store(nil)
}
// ResetCapRuntimeSettingsForTest clears the CAPTCHA runtime snapshot.
func ResetCapRuntimeSettingsForTest() {
InvalidateCapRuntimeSettings()
}
// InstallCapTestRuntimeSettings installs a fixed snapshot for unit tests.
func InstallCapTestRuntimeSettings(settings CapRuntimeSettings) func() {
snapshot := settings
capSettingsStore.snapshot.Store(&snapshot)
return InvalidateCapRuntimeSettings
}
func (s *capRuntimeSettingsStore) current(ctx context.Context) (CapRuntimeSettings, error) {
if snapshot := s.snapshot.Load(); snapshot != nil {
return *snapshot, nil
}
loaded, err, _ := s.loadGroup.Do("cap-runtime-settings", func() (any, error) {
if snapshot := s.snapshot.Load(); snapshot != nil {
return *snapshot, nil
}
settings, loadErr := loadCapRuntimeSettings(ctx)
if loadErr != nil {
return CapRuntimeSettings{}, loadErr
}
s.snapshot.Store(&settings)
return settings, nil
})
if err != nil {
return CapRuntimeSettings{}, err
}
settings, ok := loaded.(CapRuntimeSettings)
if !ok {
return CapRuntimeSettings{}, errors.New("cap runtime settings loader returned unexpected type")
}
return settings, nil
}
func loadCapRuntimeSettings(ctx context.Context) (CapRuntimeSettings, error) {
var records []capConfigRecord
db := getDB(ctx)
if db == nil {
return parseCapRuntimeSettings(nil), nil
}
if err := db.Table("w_system_configs").Where("key IN ?", capRuntimeConfigKeys).Find(&records).Error; err != nil {
return CapRuntimeSettings{}, err
}
configs := make(map[string]string, len(records))
for _, r := range records {
configs[r.Key] = r.Value
}
return parseCapRuntimeSettings(configs), nil
}
func parseCapRuntimeSettings(configs map[string]string) CapRuntimeSettings {
settings := CapRuntimeSettings{
ChallengeCount: defaultCapChallengeCount,
ChallengeSize: defaultCapChallengeSize,
ChallengeDifficulty: defaultCapChallengeDifficulty,
ChallengeTTL: defaultCapChallengeTTL,
TokenTTL: defaultCapTokenTTL,
}
if len(configs) == 0 {
return settings
}
if val, ok := configs[ConfigKeyCapLoginEnabled]; ok {
if enabled, err := strconv.ParseBool(val); err == nil {
settings.LoginEnabled = enabled
}
}
if val, ok := configs[ConfigKeyCapChallengeCount]; ok {
if count, err := strconv.Atoi(val); err == nil && count > 0 {
settings.ChallengeCount = count
}
}
if val, ok := configs[ConfigKeyCapChallengeSize]; ok {
if size, err := strconv.Atoi(val); err == nil && size > 0 {
settings.ChallengeSize = size
}
}
if val, ok := configs[ConfigKeyCapChallengeDifficulty]; ok {
if diff, err := strconv.Atoi(val); err == nil && diff > 0 {
settings.ChallengeDifficulty = diff
}
}
if val, ok := configs[ConfigKeyCapChallengeTTL]; ok {
if ttlSeconds, err := strconv.Atoi(val); err == nil && ttlSeconds > 0 {
settings.ChallengeTTL = time.Duration(ttlSeconds) * time.Second
}
}
if val, ok := configs[ConfigKeyCapTokenTTL]; ok {
if ttlSeconds, err := strconv.Atoi(val); err == nil && ttlSeconds > 0 {
settings.TokenTTL = time.Duration(ttlSeconds) * time.Second
}
}
return settings
}
+191
View File
@@ -0,0 +1,191 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package auth
import (
"Wavelet/plugins/domain/auth/pow"
"context"
"crypto/sha256"
"encoding/hex"
"strconv"
"strings"
"sync"
"time"
)
const (
redeemTokenIDLength = 8 // 兑换 Token ID 字节长度
redeemVerTokenLength = 15 // 兑换验证 Token 字节长度
tokenPartsCount = 2 // 兑换 Token 由两部分组成
valuePartsCount = 2 // 存储值由 scope 和过期时间组成
)
// CaptchaManager orchestrates challenge generation and solution validation.
type CaptchaManager struct {
secret []byte
store pow.Store
}
// NewCaptchaManager creates a new CAPTCHA Manager.
func NewCaptchaManager(secret []byte, store pow.Store) *CaptchaManager {
return &CaptchaManager{
secret: secret,
store: store,
}
}
// Generate creates a challenge response.
func (m *CaptchaManager) Generate(ctx context.Context, scope string) (*pow.ChallengeResponse, error) {
settings, err := CurrentCapSettings(ctx)
if err != nil {
return nil, err
}
challengeConfig := pow.ChallengeConfig{
Count: settings.ChallengeCount,
Size: settings.ChallengeSize,
Difficulty: settings.ChallengeDifficulty,
Expires: settings.ChallengeTTL,
}
return pow.GenerateChallenge(m.secret, challengeConfig, scope)
}
// Redeem verifies PoW solutions and returns a one-time redeem token.
func (m *CaptchaManager) Redeem(ctx context.Context, token string, solutions []int, scope string) (*RedeemResponse, error) {
sigHex := pow.JwtSigHex(token)
if sigHex == "" {
return &RedeemResponse{Success: false, Error: redeemErrInvalidToken}, nil
}
nonceKey := "cap:nonce:" + sigHex
payload, err := pow.VerifyChallengeSolutions(token, solutions, m.secret, scope)
if err != nil {
return &RedeemResponse{Success: false, Error: err.Error()}, nil //nolint:nilerr // validation errors are returned as response, not system errors
}
now := time.Now().UnixNano() / int64(time.Millisecond)
nonceTTL := time.Duration(payload.Expires-now) * time.Millisecond
if nonceTTL < time.Second {
nonceTTL = time.Second
}
set, err := m.store.SetNX(ctx, nonceKey, "1", nonceTTL)
if err != nil {
return &RedeemResponse{Success: false, Error: redeemErrNonceStoreFailed}, err
}
if !set {
return &RedeemResponse{Success: false, Error: redeemErrAlreadyRedeemed}, nil
}
settings, err := CurrentCapSettings(ctx)
if err != nil {
return &RedeemResponse{Success: false, Error: redeemErrSettingsLoad}, err
}
id := pow.RandomHex(redeemTokenIDLength)
verToken := pow.RandomHex(redeemVerTokenLength)
verHashBytes := sha256.Sum256([]byte(verToken))
verHashHex := hex.EncodeToString(verHashBytes[:])
tokenKey := "cap:token:" + id + ":" + verHashHex
tokenExpires := time.Now().Add(settings.TokenTTL)
storeVal := strconv.FormatInt(tokenExpires.UnixNano(), 10) + "|" + scope
if err := m.store.Set(ctx, tokenKey, storeVal, settings.TokenTTL); err != nil {
return &RedeemResponse{Success: false, Error: redeemErrTokenStoreFailed}, err
}
return &RedeemResponse{
Success: true,
Token: id + ":" + verToken,
Expires: tokenExpires.UnixNano() / int64(time.Millisecond),
}, nil
}
// VerifyToken validates and consumes the redeem token (single-use).
func (m *CaptchaManager) VerifyToken(ctx context.Context, token, 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
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 // invalid format is treated as validation failure
}
tokenScope := valParts[1]
if expectedScope != "" && tokenScope != expectedScope {
return false, nil
}
if time.Now().UnixNano() > expNano {
return false, nil
}
return true, nil
}
func sGetAndDelete(ctx context.Context, store pow.Store, key string) (string, bool, error) {
if store == nil {
return "", false, nil
}
return store.GetAndDelete(ctx, key)
}
var (
defaultCapManagerMu sync.RWMutex
defaultCapManager *CaptchaManager
)
// SetCapSecret sets the shared secret used by the default CAPTCHA manager.
func SetCapSecret(secret []byte) {
defaultCapManagerMu.Lock()
defer defaultCapManagerMu.Unlock()
if len(secret) > 0 {
store := pow.NewMemoryStore(1 * time.Minute)
defaultCapManager = NewCaptchaManager(secret, store)
}
}
// GetDefaultCapManager yields the global singleton CAPTCHA manager.
func GetDefaultCapManager() *CaptchaManager {
defaultCapManagerMu.RLock()
defer defaultCapManagerMu.RUnlock()
return defaultCapManager
}
type captchaService struct{}
func (captchaService) VerifyMiddleware(scope string) any {
return VerifyCaptchaMiddleware(GetDefaultCapManager(), scope)
}
func (captchaService) ChallengeHandler() any { return Challenge }
func (captchaService) RedeemHandler() any { return Redeem }
+40 -5
View File
@@ -85,6 +85,9 @@ func (p *Plugin) Apply(ctx *core.Context) error {
var cfg SessionConfig
if err := ctx.Config().Bind("app", &cfg); err == nil {
SetSessionConfig(cfg)
if cfg.SessionSecret != "" {
SetCapSecret([]byte(cfg.SessionSecret))
}
}
core.Bind[contracts.DBService](ctx, setDBService)
@@ -100,7 +103,7 @@ func (p *Plugin) Apply(ctx *core.Context) error {
// 1. Register migrations
ctx.Migrations().Register("auth", authMigrations)
// 2. Initialize and provide AuthService & AuthRegistry
// 2. Initialize and provide AuthService, AuthRegistry & CaptchaService
if p.authSvc == nil {
p.authSvc = newAuthService()
}
@@ -110,6 +113,7 @@ func (p *Plugin) Apply(ctx *core.Context) error {
core.Provide[contracts.AuthService](ctx, p.authSvc)
core.Provide[contracts.AuthRegistry](ctx, p.authRegistry)
core.Provide[contracts.CaptchaService](ctx, captchaService{})
// 2.1 Register Public / Auth Whitelist Endpoints
publicEndpoints := []string{
@@ -144,20 +148,47 @@ func (p *Plugin) Apply(ctx *core.Context) error {
}
ctx.Router().GET("/api/v1/user-info", LoginRequired(), UserInfo)
// 3.1 Register CAPTCHA HTTP Routes
capGroup := ctx.Router().Group("/api/v1/cap")
{
capGroup.GET("/challenge", Challenge)
capGroup.POST("/challenge", Challenge)
capGroup.POST("/redeem", Redeem)
}
// 4. Register Settings Schemas
const (
settingTypeInteger = "integer"
settingCategorySecurity = "security"
)
ctx.Settings().Register(extpoints.SettingSchema{
Key: "auth.session_age",
Default: 86400 * 7,
Description: "Default session lifetime in seconds",
Type: "integer",
Category: "security",
Type: settingTypeInteger,
Category: settingCategorySecurity,
})
ctx.Settings().Register(extpoints.SettingSchema{
Key: "auth.login_rate_limit_max_attempts",
Default: 5,
Description: "Max login failure attempts before temporary IP lock",
Type: "integer",
Category: "security",
Type: settingTypeInteger,
Category: settingCategorySecurity,
})
ctx.Settings().Register(extpoints.SettingSchema{
Key: "cap.login_enabled",
Default: false,
Description: "Whether to require CAPTCHA verification for user login",
Type: "boolean",
Category: settingCategorySecurity,
})
ctx.Settings().Register(extpoints.SettingSchema{
Key: "cap.challenge_count",
Default: 1,
Description: "Number of PoW puzzle challenges to solve",
Type: settingTypeInteger,
Category: settingCategorySecurity,
})
// 5. Register Event Listeners for domain events
@@ -171,5 +202,9 @@ func (p *Plugin) Apply(ctx *core.Context) error {
return nil
})
ctx.Events().On(contracts.EventTopicConfigChanged, func(_ any) {
InvalidateCapRuntimeSettings()
})
return nil
}
@@ -161,4 +161,25 @@ func TestAuthPluginUnit(t *testing.T) {
current, err := authSvc.GetCurrentUser(userCtx)
require.NoError(t, err)
assert.Equal(t, user.ID, current.ID)
// Test CaptchaService injection
capSvc, err := core.Inject[contracts.CaptchaService](ctx)
require.NoError(t, err)
assert.NotNil(t, capSvc)
assert.NotNil(t, capSvc.ChallengeHandler())
assert.NotNil(t, capSvc.RedeemHandler())
assert.NotNil(t, capSvc.VerifyMiddleware("login"))
// Verify CAPTCHA routes registered
var foundChallenge, foundRedeem bool
for _, rd := range ctx.Router().Routes() {
if rd.Path == "/api/v1/cap/challenge" {
foundChallenge = true
}
if rd.Path == "/api/v1/cap/redeem" {
foundRedeem = true
}
}
assert.True(t, foundChallenge, "expected /api/v1/cap/challenge route")
assert.True(t, foundRedeem, "expected /api/v1/cap/redeem route")
}
+259
View File
@@ -0,0 +1,259 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package pow provides proof-of-work challenge generation and verification.
package pow
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, 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 pow
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"
)
+111
View File
@@ -0,0 +1,111 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package pow
import (
"context"
"testing"
"time"
)
func TestPowChallengeFlow(t *testing.T) {
secret := []byte("test-secret-key-1234567890123456")
conf := ChallengeConfig{
Count: 2,
Size: 16,
Difficulty: 1,
Expires: 1 * time.Minute,
}
scope := "login"
resp, err := GenerateChallenge(secret, conf, scope)
if err != nil {
t.Fatalf("GenerateChallenge failed: %v", err)
}
if resp.Token == "" {
t.Fatal("expected non-empty token")
}
if resp.Challenge.C != 2 {
t.Fatalf("expected count 2, got %d", resp.Challenge.C)
}
sigHex := JwtSigHex(resp.Token)
if sigHex == "" {
t.Fatal("expected non-empty sigHex")
}
solutions := Solve(resp.Token, resp.Challenge.C, resp.Challenge.S, resp.Challenge.D)
if len(solutions) != 2 {
t.Fatalf("expected 2 solutions, got %d", len(solutions))
}
payload, err := VerifyChallengeSolutions(resp.Token, solutions, secret, scope)
if err != nil {
t.Fatalf("VerifyChallengeSolutions failed: %v", err)
}
if payload.Scope != scope {
t.Fatalf("expected scope %s, got %s", scope, payload.Scope)
}
// Scope mismatch test
_, err = VerifyChallengeSolutions(resp.Token, solutions, secret, "other_scope")
if err == nil {
t.Fatal("expected scope mismatch error")
}
// Invalid solutions test
_, err = VerifyChallengeSolutions(resp.Token, []int{9999999, 9999999}, secret, scope)
if err == nil {
t.Fatal("expected invalid solution error")
}
// Invalid token test
_, err = VerifyChallengeSolutions("invalid.jwt.token", solutions, secret, scope)
if err == nil {
t.Fatal("expected invalid token error")
}
}
func TestMemoryStore(t *testing.T) {
ctx := context.Background()
store := NewMemoryStore(100 * time.Millisecond)
// Set and Get
err := store.Set(ctx, "k1", "v1", 200*time.Millisecond)
if err != nil {
t.Fatalf("Set failed: %v", err)
}
val, ok, err := store.Get(ctx, "k1")
if err != nil || !ok || val != "v1" {
t.Fatalf("Get failed: val=%s, ok=%v, err=%v", val, ok, err)
}
// SetNX
set, err := store.SetNX(ctx, "k1", "v2", 200*time.Millisecond)
if err != nil || set {
t.Fatalf("SetNX should have failed because key exists: set=%v, err=%v", set, err)
}
set, err = store.SetNX(ctx, "k2", "v2", 200*time.Millisecond)
if err != nil || !set {
t.Fatalf("SetNX should have succeeded: set=%v, err=%v", set, err)
}
// GetAndDelete
val, ok, err = store.GetAndDelete(ctx, "k2")
if err != nil || !ok || val != "v2" {
t.Fatalf("GetAndDelete failed: val=%s, ok=%v, err=%v", val, ok, err)
}
_, ok, _ = store.Get(ctx, "k2")
if ok {
t.Fatal("k2 should be deleted")
}
// Delete
_ = store.Delete(ctx, "k1")
_, ok, _ = store.Get(ctx, "k1")
if ok {
t.Fatal("k1 should be deleted")
}
}
+54
View File
@@ -0,0 +1,54 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package pow
import (
"fmt"
"strings"
)
// fnv1a returns the 32-bit FNV-1a hash of a string
//
// 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
//
// FNV-1a 算法位移常量
func fnv1aResume(state uint32, str string) uint32 {
h := state
for i := 0; i < len(str); i++ {
h ^= uint32(str[i])
h += (hashShift(h))
}
return h
}
// hashShift computes FNV-1a mix additions
func hashShift(h uint32) uint32 {
return (h << 1) + (h << 4) + (h << 7) + (h << 8) + (h << 24)
}
// prngFromHash generates a hex string of specified length using an initial hash state
//
// 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]
}
+188
View File
@@ -0,0 +1,188 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package pow
import (
"Wavelet/pkg/util"
"context"
"errors"
"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, 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, 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 {
util.Go(func() { 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, 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, 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 errors.Is(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, 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, 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 errors.Is(err, redis.Nil) {
return "", false, nil
}
if err != nil {
return "", false, err
}
return val, true, nil
}