mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-08 00:26:37 +08:00
cap 与系统信息
This commit is contained in:
@@ -0,0 +1,189 @@
|
||||
/*
|
||||
Copyright 2025 linux.do
|
||||
|
||||
Licensed under the Apache License, Version 2.0 (the "License");
|
||||
you may not use this file except in compliance with the License.
|
||||
You may obtain a copy of the License at
|
||||
|
||||
http://www.apache.org/licenses/LICENSE-2.0
|
||||
|
||||
Unless required by applicable law or agreed to in writing, software
|
||||
distributed under the License is distributed on an "AS IS" BASIS,
|
||||
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
See the License for the specific language governing permissions and
|
||||
limitations under the License.
|
||||
*/
|
||||
|
||||
package status
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"math"
|
||||
"net/http"
|
||||
"runtime"
|
||||
"time"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/linux-do/credit/internal/util"
|
||||
)
|
||||
|
||||
// startTime 记录服务启动时间
|
||||
var startTime = time.Now()
|
||||
|
||||
// SystemStatusResponse 系统状态响应结构体
|
||||
type SystemStatusResponse struct {
|
||||
Uptime string `json:"uptime"`
|
||||
NumGoroutine int `json:"num_goroutine"`
|
||||
Alloc string `json:"alloc"`
|
||||
TotalAlloc string `json:"total_alloc"`
|
||||
Sys string `json:"sys"`
|
||||
Lookups uint64 `json:"lookups"`
|
||||
Mallocs uint64 `json:"mallocs"`
|
||||
Frees uint64 `json:"frees"`
|
||||
HeapAlloc string `json:"heap_alloc"`
|
||||
HeapSys string `json:"heap_sys"`
|
||||
HeapIdle string `json:"heap_idle"`
|
||||
HeapInuse string `json:"heap_inuse"`
|
||||
HeapReleased string `json:"heap_released"`
|
||||
HeapObjects uint64 `json:"heap_objects"`
|
||||
StackInuse string `json:"stack_inuse"`
|
||||
StackSys string `json:"stack_sys"`
|
||||
MSpanInuse string `json:"mspan_inuse"`
|
||||
MSpanSys string `json:"mspan_sys"`
|
||||
MCacheInuse string `json:"mcache_inuse"`
|
||||
MCacheSys string `json:"mcache_sys"`
|
||||
BuckHashSys string `json:"buck_hash_sys"`
|
||||
GCSys string `json:"gc_sys"`
|
||||
OtherSys string `json:"other_sys"`
|
||||
NextGC string `json:"next_gc"`
|
||||
LastGCTime string `json:"last_gc_time"`
|
||||
PauseTotalNs string `json:"pause_total_ns"`
|
||||
LastPause string `json:"last_pause"`
|
||||
NumGC uint32 `json:"num_gc"`
|
||||
}
|
||||
|
||||
// formatBytes 格式化字节大小
|
||||
func formatBytes(bytes uint64) string {
|
||||
const unit = 1024
|
||||
if bytes < unit {
|
||||
return fmt.Sprintf("%d B", bytes)
|
||||
}
|
||||
div, exp := int64(unit), 0
|
||||
for n := bytes / unit; n >= unit; n /= unit {
|
||||
div *= unit
|
||||
exp++
|
||||
}
|
||||
value := float64(bytes) / float64(div)
|
||||
var suffix string
|
||||
switch exp {
|
||||
case 0:
|
||||
suffix = "KiB"
|
||||
case 1:
|
||||
suffix = "MiB"
|
||||
case 2:
|
||||
suffix = "GiB"
|
||||
default:
|
||||
suffix = "TiB"
|
||||
}
|
||||
|
||||
// 格式化规则:
|
||||
// - 如果是整数(如 16, 73, 105, 986, 112):
|
||||
// - 如果 >= 10,则格式化为 "%.0f" (e.g. "16 KiB")
|
||||
// - 如果 < 10,则格式化为 "%.1f" (e.g. "9.0 KiB")
|
||||
// - 如果不是整数(如 5.8, 9.1, 7.6, 4.8):格式化为 "%.1f"
|
||||
if value == math.Trunc(value) {
|
||||
if value >= 10 {
|
||||
return fmt.Sprintf("%.0f %s", value, suffix)
|
||||
}
|
||||
return fmt.Sprintf("%.1f %s", value, suffix)
|
||||
}
|
||||
return fmt.Sprintf("%.1f %s", value, suffix)
|
||||
}
|
||||
|
||||
// formatDuration 格式化时间持续时间
|
||||
func formatDuration(d time.Duration) string {
|
||||
days := int(d.Hours()) / 24
|
||||
hours := int(d.Hours()) % 24
|
||||
minutes := int(d.Minutes()) % 60
|
||||
seconds := int(d.Seconds()) % 60
|
||||
|
||||
var res string
|
||||
if days > 0 {
|
||||
res += fmt.Sprintf("%d天", days)
|
||||
}
|
||||
if hours > 0 {
|
||||
res += fmt.Sprintf("%d小时", hours)
|
||||
}
|
||||
if minutes > 0 {
|
||||
res += fmt.Sprintf("%d分钟", minutes)
|
||||
}
|
||||
if seconds > 0 || res == "" {
|
||||
res += fmt.Sprintf("%d秒钟", seconds)
|
||||
}
|
||||
return res
|
||||
}
|
||||
|
||||
// GetSystemStatus 获取系统状态信息
|
||||
// @Summary 获取系统状态信息
|
||||
// @Description 获取后端服务运行状态、Goroutine、内存指标等详细统计数据,需要管理员权限
|
||||
// @Tags admin
|
||||
// @Produce json
|
||||
// @Security SessionCookie
|
||||
// @Success 200 {object} util.ResponseAny{data=status.SystemStatusResponse} "获取成功"
|
||||
// @Failure 401 {object} util.ResponseAny "未登录"
|
||||
// @Failure 403 {object} util.ResponseAny "无管理员权限"
|
||||
// @Router /api/v1/admin/status [get]
|
||||
func GetSystemStatus(c *gin.Context) {
|
||||
var m runtime.MemStats
|
||||
runtime.ReadMemStats(&m)
|
||||
|
||||
uptime := formatDuration(time.Since(startTime))
|
||||
numGoroutine := runtime.NumGoroutine()
|
||||
|
||||
var lastGCTime string
|
||||
if m.LastGC > 0 {
|
||||
lastGCTime = formatDuration(time.Since(time.Unix(0, int64(m.LastGC))))
|
||||
} else {
|
||||
lastGCTime = "无"
|
||||
}
|
||||
|
||||
var lastPause string
|
||||
if m.NumGC > 0 {
|
||||
lastPause = fmt.Sprintf("%.3fs", float64(m.PauseNs[(m.NumGC-1)%256])/1e9)
|
||||
} else {
|
||||
lastPause = "0.000s"
|
||||
}
|
||||
|
||||
res := SystemStatusResponse{
|
||||
Uptime: uptime,
|
||||
NumGoroutine: numGoroutine,
|
||||
Alloc: formatBytes(m.Alloc),
|
||||
TotalAlloc: formatBytes(m.TotalAlloc),
|
||||
Sys: formatBytes(m.Sys),
|
||||
Lookups: m.Lookups,
|
||||
Mallocs: m.Mallocs,
|
||||
Frees: m.Frees,
|
||||
HeapAlloc: formatBytes(m.HeapAlloc),
|
||||
HeapSys: formatBytes(m.HeapSys),
|
||||
HeapIdle: formatBytes(m.HeapIdle),
|
||||
HeapInuse: formatBytes(m.HeapInuse),
|
||||
HeapReleased: formatBytes(m.HeapReleased),
|
||||
HeapObjects: m.HeapObjects,
|
||||
StackInuse: formatBytes(m.StackInuse),
|
||||
StackSys: formatBytes(m.StackSys),
|
||||
MSpanInuse: formatBytes(m.MSpanInuse),
|
||||
MSpanSys: formatBytes(m.MSpanSys),
|
||||
MCacheInuse: formatBytes(m.MCacheInuse),
|
||||
MCacheSys: formatBytes(m.MCacheSys),
|
||||
BuckHashSys: formatBytes(m.BuckHashSys),
|
||||
GCSys: formatBytes(m.GCSys),
|
||||
OtherSys: formatBytes(m.OtherSys),
|
||||
NextGC: formatBytes(m.NextGC),
|
||||
LastGCTime: lastGCTime,
|
||||
PauseTotalNs: fmt.Sprintf("%.1fs", float64(m.PauseTotalNs)/1e9),
|
||||
LastPause: lastPause,
|
||||
NumGC: m.NumGC,
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, util.OK(res))
|
||||
}
|
||||
@@ -0,0 +1,34 @@
|
||||
package cap
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/linux-do/credit/internal/util"
|
||||
caputil "github.com/linux-do/credit/internal/util/cap"
|
||||
)
|
||||
|
||||
// VerifyMiddleware returns a Gin middleware that checks and consumes the X-Cap-Token header.
|
||||
// enabledFunc is an optional callback allowing dynamic check of whether captcha protection is turned on.
|
||||
func VerifyMiddleware(mgr *caputil.Manager, scope string, enabledFunc func() bool) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
if enabledFunc != nil && !enabledFunc() {
|
||||
c.Next()
|
||||
return
|
||||
}
|
||||
|
||||
token := c.GetHeader("X-Cap-Token")
|
||||
if token == "" {
|
||||
c.AbortWithStatusJSON(http.StatusUnauthorized, util.Err("验证码验证失败,缺少验证码凭证"))
|
||||
return
|
||||
}
|
||||
|
||||
valid, err := mgr.VerifyToken(c.Request.Context(), token, scope)
|
||||
if err != nil || !valid {
|
||||
c.AbortWithStatusJSON(http.StatusUnauthorized, util.Err("验证码校验失败或已过期,请重试"))
|
||||
return
|
||||
}
|
||||
|
||||
c.Next()
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,92 @@
|
||||
package cap
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/linux-do/credit/internal/util/cap"
|
||||
)
|
||||
|
||||
type challengeRequest struct {
|
||||
Scope string `json:"scope" form:"scope"`
|
||||
}
|
||||
|
||||
type redeemRequest struct {
|
||||
Token string `json:"token" binding:"required"`
|
||||
Solutions []int `json:"solutions" binding:"required"`
|
||||
Scope string `json:"scope" form:"scope"`
|
||||
}
|
||||
|
||||
// Challenge 生成 PoW 人机验证难题
|
||||
// @Summary 生成人机验证难题
|
||||
// @Description 客户端获取 PoW 难题和签名的 JWT Token,并在后台计算。
|
||||
// @Tags cap
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Param request body challengeRequest false "可选范围限制参数"
|
||||
// @Success 200 {object} cap.ChallengeResponse "成功返回 PoW 难题"
|
||||
// @Failure 500 {object} cap.RedeemResponse "内部服务错误"
|
||||
// @Router /api/cap/challenge [post]
|
||||
func Challenge(c *gin.Context) {
|
||||
var req challengeRequest
|
||||
_ = c.ShouldBind(&req) // 允许不传 body,默认使用 login scope
|
||||
|
||||
if req.Scope == "" {
|
||||
req.Scope = "login"
|
||||
}
|
||||
|
||||
mgr := cap.GetDefaultManager()
|
||||
resp, err := mgr.Generate(c.Request.Context(), req.Scope)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, cap.RedeemResponse{
|
||||
Success: false,
|
||||
Error: err.Error(),
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, 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} cap.RedeemResponse "核销成功,返回 X-Cap-Token"
|
||||
// @Failure 400 {object} cap.RedeemResponse "参数错误或核销失败"
|
||||
// @Failure 500 {object} cap.RedeemResponse "内部服务错误"
|
||||
// @Router /api/cap/redeem [post]
|
||||
func Redeem(c *gin.Context) {
|
||||
var req redeemRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
c.JSON(http.StatusBadRequest, cap.RedeemResponse{
|
||||
Success: false,
|
||||
Error: "无效的参数",
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
if req.Scope == "" {
|
||||
req.Scope = "login"
|
||||
}
|
||||
|
||||
mgr := cap.GetDefaultManager()
|
||||
resp, err := mgr.Redeem(c.Request.Context(), req.Token, req.Solutions, req.Scope)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, cap.RedeemResponse{
|
||||
Success: false,
|
||||
Error: err.Error(),
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
if !resp.Success {
|
||||
c.JSON(http.StatusBadRequest, resp)
|
||||
return
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, resp)
|
||||
}
|
||||
@@ -0,0 +1,132 @@
|
||||
package cap
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/linux-do/credit/internal/db"
|
||||
"github.com/linux-do/credit/internal/model"
|
||||
"github.com/linux-do/credit/internal/testhelper"
|
||||
"github.com/linux-do/credit/internal/util"
|
||||
capUtil "github.com/linux-do/credit/internal/util/cap"
|
||||
)
|
||||
|
||||
func TestCapEndpointsAndMiddleware(t *testing.T) {
|
||||
sqliteDB, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
defer cleanup()
|
||||
|
||||
gin.SetMode(gin.TestMode)
|
||||
r := gin.New()
|
||||
|
||||
// Mount CAPTCHA API endpoints
|
||||
capGroup := r.Group("/api/cap")
|
||||
{
|
||||
capGroup.POST("/challenge", Challenge)
|
||||
capGroup.POST("/redeem", Redeem)
|
||||
}
|
||||
|
||||
// Login endpoint with CAPTCHA middleware
|
||||
r.POST("/api/v1/user/login", VerifyMiddleware(capUtil.GetDefaultManager(), "login", func() bool {
|
||||
enabled, err := model.GetBoolByKey(context.Background(), model.ConfigKeyCapLoginEnabled)
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
return enabled
|
||||
}), func(c *gin.Context) {
|
||||
c.JSON(http.StatusOK, util.OK("login success"))
|
||||
})
|
||||
|
||||
// 1. Test challenge generation
|
||||
w := httptest.NewRecorder()
|
||||
req, _ := http.NewRequest("POST", "/api/cap/challenge", nil)
|
||||
r.ServeHTTP(w, req)
|
||||
|
||||
if w.Code != http.StatusOK {
|
||||
t.Fatalf("expected 200 OK, got %d. Body: %s", w.Code, w.Body.String())
|
||||
}
|
||||
|
||||
var challengeResp capUtil.ChallengeResponse
|
||||
if err := json.Unmarshal(w.Body.Bytes(), &challengeResp); err != nil {
|
||||
t.Fatalf("failed to unmarshal challenge response: %v", err)
|
||||
}
|
||||
|
||||
if challengeResp.Token == "" {
|
||||
t.Fatalf("expected token in challenge response")
|
||||
}
|
||||
|
||||
// 2. Test login with CAPTCHA disabled (should pass)
|
||||
w = httptest.NewRecorder()
|
||||
req, _ = http.NewRequest("POST", "/api/v1/user/login", nil)
|
||||
r.ServeHTTP(w, req)
|
||||
if w.Code != http.StatusOK {
|
||||
t.Fatalf("expected 200 OK when CAPTCHA is disabled, got %d. Body: %s", w.Code, w.Body.String())
|
||||
}
|
||||
|
||||
// 3. Enable CAPTCHA in DB
|
||||
err := sqliteDB.Model(&model.SystemConfig{}).Where("key = ?", model.ConfigKeyCapLoginEnabled).Update("value", "true").Error
|
||||
if err != nil {
|
||||
t.Fatalf("failed to enable cap_login_enabled in DB: %v", err)
|
||||
}
|
||||
// Update cache
|
||||
var sysCfg model.SystemConfig
|
||||
sqliteDB.Where("key = ?", model.ConfigKeyCapLoginEnabled).First(&sysCfg)
|
||||
_ = db.HSetJSON(context.Background(), model.SystemConfigRedisHashKey, model.ConfigKeyCapLoginEnabled, &sysCfg)
|
||||
|
||||
// 4. Test login with CAPTCHA enabled but no header (should be blocked)
|
||||
w = httptest.NewRecorder()
|
||||
req, _ = http.NewRequest("POST", "/api/v1/user/login", nil)
|
||||
r.ServeHTTP(w, req)
|
||||
if w.Code != http.StatusUnauthorized {
|
||||
t.Fatalf("expected 401 Unauthorized, got %d. Body: %s", w.Code, w.Body.String())
|
||||
}
|
||||
|
||||
// 5. Solve the challenge
|
||||
solutions := capUtil.Solve(challengeResp.Token, challengeResp.Challenge.C, challengeResp.Challenge.S, challengeResp.Challenge.D)
|
||||
|
||||
// 6. Redeem solutions
|
||||
redeemReqPayload := redeemRequest{
|
||||
Token: challengeResp.Token,
|
||||
Solutions: solutions,
|
||||
}
|
||||
bodyBytes, _ := json.Marshal(redeemReqPayload)
|
||||
w = httptest.NewRecorder()
|
||||
req, _ = http.NewRequest("POST", "/api/cap/redeem", bytes.NewBuffer(bodyBytes))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
r.ServeHTTP(w, req)
|
||||
|
||||
if w.Code != http.StatusOK {
|
||||
t.Fatalf("expected 200 OK for redeem, got %d. Body: %s", w.Code, w.Body.String())
|
||||
}
|
||||
|
||||
var redeemResp capUtil.RedeemResponse
|
||||
if err := json.Unmarshal(w.Body.Bytes(), &redeemResp); err != nil {
|
||||
t.Fatalf("failed to unmarshal redeem response: %v", err)
|
||||
}
|
||||
|
||||
if !redeemResp.Success || redeemResp.Token == "" {
|
||||
t.Fatalf("redeem failed or returned empty token: %+v", redeemResp)
|
||||
}
|
||||
|
||||
// 7. Login with valid redeem token (should pass)
|
||||
w = httptest.NewRecorder()
|
||||
req, _ = http.NewRequest("POST", "/api/v1/user/login", nil)
|
||||
req.Header.Set("X-Cap-Token", redeemResp.Token)
|
||||
r.ServeHTTP(w, req)
|
||||
if w.Code != http.StatusOK {
|
||||
t.Fatalf("expected 200 OK with valid cap token, got %d. Body: %s", w.Code, w.Body.String())
|
||||
}
|
||||
|
||||
// 8. Replay attack: Login with the same redeem token again (should be blocked as it is single-use)
|
||||
w = httptest.NewRecorder()
|
||||
req, _ = http.NewRequest("POST", "/api/v1/user/login", nil)
|
||||
req.Header.Set("X-Cap-Token", redeemResp.Token)
|
||||
r.ServeHTTP(w, req)
|
||||
if w.Code != http.StatusUnauthorized {
|
||||
t.Fatalf("expected 401 Unauthorized on replayed token, got %d. Body: %s", w.Code, w.Body.String())
|
||||
}
|
||||
}
|
||||
@@ -33,6 +33,8 @@ type PublicConfigResponse struct {
|
||||
PasswordRegisterEnabled bool `json:"password_register_enabled"` // 是否允许密码注册
|
||||
OIDCLoginEnabled bool `json:"oidc_login_enabled"` // 是否允许 OIDC 登录
|
||||
MaxAPIKeysPerUser int `json:"max_api_keys_per_user"` // 每个用户最大 API Key 数量
|
||||
CapLoginEnabled bool `json:"cap_login_enabled"` // 是否启用人机验证
|
||||
CapAutoSolve bool `json:"cap_auto_solve"` // 打开页面后是否自动开始计算
|
||||
}
|
||||
|
||||
// GetPublicConfig 获取公共配置
|
||||
@@ -83,6 +85,18 @@ func GetPublicConfig(c *gin.Context) {
|
||||
oidcLoginEnabled = val
|
||||
}
|
||||
|
||||
// 3.4 cap_login_enabled
|
||||
var capLoginEnabled bool
|
||||
if val, err := model.GetBoolByKey(ctx, model.ConfigKeyCapLoginEnabled); err == nil {
|
||||
capLoginEnabled = val
|
||||
}
|
||||
|
||||
// 3.5 cap_auto_solve
|
||||
capAutoSolve := true // 默认自动开始
|
||||
if val, err := model.GetBoolByKey(ctx, model.ConfigKeyCapAutoSolve); err == nil {
|
||||
capAutoSolve = val
|
||||
}
|
||||
|
||||
// 4. max_api_keys_per_user
|
||||
var maxAPIKeys int
|
||||
if val, err := model.GetIntByKey(ctx, model.ConfigKeyMaxAPIKeysPerUser); err == nil {
|
||||
@@ -97,6 +111,8 @@ func GetPublicConfig(c *gin.Context) {
|
||||
PasswordRegisterEnabled: passwordRegisterEnabled,
|
||||
OIDCLoginEnabled: oidcLoginEnabled,
|
||||
MaxAPIKeysPerUser: maxAPIKeys,
|
||||
CapLoginEnabled: capLoginEnabled,
|
||||
CapAutoSolve: capAutoSolve,
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, util.OK(response))
|
||||
|
||||
@@ -52,6 +52,25 @@ func Migrate() {
|
||||
initDefaultAdmin()
|
||||
}
|
||||
|
||||
// ensureConfigKeyExists ensures a system config key exists in the database
|
||||
func ensureConfigKeyExists(key, value, configType, description string) {
|
||||
tx := db.DB(context.Background())
|
||||
var cfg model.SystemConfig
|
||||
if err := tx.Where("key = ?", key).First(&cfg).Error; err != nil {
|
||||
newConfig := model.SystemConfig{
|
||||
Key: key,
|
||||
Value: value,
|
||||
Type: configType,
|
||||
Description: description,
|
||||
}
|
||||
if err := tx.Create(&newConfig).Error; err != nil {
|
||||
log.Printf("[PostgreSQL] failed to create system config key %s: %v\n", key, err)
|
||||
} else {
|
||||
log.Printf("[PostgreSQL] initialized system config key %s\n", key)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// initSystemConfigs 初始化系统配置数据
|
||||
func initSystemConfigs() {
|
||||
tx := db.DB(context.Background())
|
||||
@@ -63,10 +82,59 @@ func initSystemConfigs() {
|
||||
}
|
||||
|
||||
if count > 0 {
|
||||
ensureConfigKeyExists(model.ConfigKeyCapLoginEnabled, "false", "system", "是否启用登录人机验证(true/false)")
|
||||
ensureConfigKeyExists(model.ConfigKeyCapAutoSolve, "true", "system", "打开页面后是否自动开始计算,关闭则需用户手动点击触发")
|
||||
ensureConfigKeyExists(model.ConfigKeyCapChallengeCount, "1", "system", "客户端需求解的 PoW 难题总数,默认 1,推荐 1~5")
|
||||
ensureConfigKeyExists(model.ConfigKeyCapChallengeSize, "32", "system", "人机验证盐值长度")
|
||||
ensureConfigKeyExists(model.ConfigKeyCapChallengeDifficulty, "4", "system", "人机验证 PoW 难度(目标前缀长度)")
|
||||
ensureConfigKeyExists(model.ConfigKeyCapChallengeTTL, "600", "system", "人机验证难题有效时间(秒)")
|
||||
ensureConfigKeyExists(model.ConfigKeyCapTokenTTL, "1200", "system", "人机验证兑换凭证有效时间(秒)")
|
||||
return
|
||||
}
|
||||
|
||||
defaultConfigs := []model.SystemConfig{
|
||||
{
|
||||
Key: model.ConfigKeyCapLoginEnabled,
|
||||
Value: "false",
|
||||
Type: "system",
|
||||
Description: "是否启用登录人机验证(true/false)",
|
||||
},
|
||||
{
|
||||
Key: model.ConfigKeyCapAutoSolve,
|
||||
Value: "true",
|
||||
Type: "system",
|
||||
Description: "打开页面后是否自动开始计算,关闭则需用户手动点击触发",
|
||||
},
|
||||
{
|
||||
Key: model.ConfigKeyCapChallengeCount,
|
||||
Value: "1",
|
||||
Type: "system",
|
||||
Description: "客户端需求解的 PoW 难题总数,默认 1,推荐 1~5",
|
||||
},
|
||||
{
|
||||
Key: model.ConfigKeyCapChallengeSize,
|
||||
Value: "32",
|
||||
Type: "system",
|
||||
Description: "人机验证盐值长度",
|
||||
},
|
||||
{
|
||||
Key: model.ConfigKeyCapChallengeDifficulty,
|
||||
Value: "4",
|
||||
Type: "system",
|
||||
Description: "人机验证 PoW 难度(目标前缀长度)",
|
||||
},
|
||||
{
|
||||
Key: model.ConfigKeyCapChallengeTTL,
|
||||
Value: "600",
|
||||
Type: "system",
|
||||
Description: "人机验证难题有效时间(秒)",
|
||||
},
|
||||
{
|
||||
Key: model.ConfigKeyCapTokenTTL,
|
||||
Value: "1200",
|
||||
Type: "system",
|
||||
Description: "人机验证兑换凭证有效时间(秒)",
|
||||
},
|
||||
{
|
||||
Key: model.ConfigKeyUploadAllowedExtensions,
|
||||
Value: "jpg,png,webp",
|
||||
|
||||
@@ -173,6 +173,9 @@ func buildDSN(host string, port int, username, password string) string {
|
||||
}
|
||||
|
||||
func DB(ctx context.Context) *gorm.DB {
|
||||
if db == nil {
|
||||
return nil
|
||||
}
|
||||
return db.WithContext(ctx)
|
||||
}
|
||||
|
||||
|
||||
@@ -38,6 +38,13 @@ const (
|
||||
ConfigKeyPasswordRegisterEnabled = "password_register_enabled" // 是否允许密码注册
|
||||
ConfigKeyOIDCLoginEnabled = "oidc_login_enabled" // 是否允许 OIDC 登录
|
||||
ConfigKeyMaxAPIKeysPerUser = "max_api_keys_per_user" // 每个用户最大 API Key 数量
|
||||
ConfigKeyCapLoginEnabled = "cap_login_enabled" // 是否启用登录人机验证
|
||||
ConfigKeyCapAutoSolve = "cap_auto_solve" // 打开页面后是否自动开始计算(false 则需用户手动点击)
|
||||
ConfigKeyCapChallengeCount = "cap_challenge_count" // 客户端需求解的 PoW 难题总数,默认 1,推荐 1~5
|
||||
ConfigKeyCapChallengeSize = "cap_challenge_size" // 人机验证盐值长度
|
||||
ConfigKeyCapChallengeDifficulty = "cap_challenge_difficulty" // 人机验证 PoW 难度(目标前缀长度)
|
||||
ConfigKeyCapChallengeTTL = "cap_challenge_ttl_seconds" // 人机验证难题有效时间(秒)
|
||||
ConfigKeyCapTokenTTL = "cap_token_ttl_seconds" // 人机验证兑换凭证有效时间(秒)
|
||||
)
|
||||
|
||||
const (
|
||||
@@ -56,20 +63,29 @@ type SystemConfig struct {
|
||||
|
||||
// GetByKey 通过 key 查询配置(带 Redis 缓存)
|
||||
func (sc *SystemConfig) GetByKey(ctx context.Context, key string) error {
|
||||
if err := db.HGetJSON(ctx, SystemConfigRedisHashKey, key, sc); err == nil {
|
||||
return nil
|
||||
} else if !errors.Is(err, redis.Nil) {
|
||||
// Redis 服务错误,返回错误
|
||||
return err
|
||||
if db.Redis != nil {
|
||||
if err := db.HGetJSON(ctx, SystemConfigRedisHashKey, key, sc); err == nil {
|
||||
return nil
|
||||
} else if !errors.Is(err, redis.Nil) {
|
||||
// Redis 服务错误,返回错误
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
// 查数据库
|
||||
if err := db.DB(ctx).Where("key = ?", key).First(sc).Error; err != nil {
|
||||
database := db.DB(ctx)
|
||||
if database == nil {
|
||||
return errors.New("database not initialized")
|
||||
}
|
||||
|
||||
if err := database.Where("key = ?", key).First(sc).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// 更新 Redis Hash 缓存
|
||||
_ = db.HSetJSON(ctx, SystemConfigRedisHashKey, key, sc)
|
||||
if db.Redis != nil {
|
||||
_ = db.HSetJSON(ctx, SystemConfigRedisHashKey, key, sc)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -29,13 +29,17 @@ import (
|
||||
|
||||
"github.com/linux-do/credit/internal/apps/admin"
|
||||
admin_auth_source "github.com/linux-do/credit/internal/apps/admin/auth_source"
|
||||
admin_status "github.com/linux-do/credit/internal/apps/admin/status"
|
||||
admin_task "github.com/linux-do/credit/internal/apps/admin/task"
|
||||
admin_user "github.com/linux-do/credit/internal/apps/admin/user"
|
||||
capApp "github.com/linux-do/credit/internal/apps/cap"
|
||||
publicconfig "github.com/linux-do/credit/internal/apps/config"
|
||||
"github.com/linux-do/credit/internal/apps/health"
|
||||
"github.com/linux-do/credit/internal/apps/upload"
|
||||
"github.com/linux-do/credit/internal/apps/user"
|
||||
"github.com/linux-do/credit/internal/model"
|
||||
"github.com/linux-do/credit/internal/util"
|
||||
capUtil "github.com/linux-do/credit/internal/util/cap"
|
||||
|
||||
"github.com/gin-contrib/sessions"
|
||||
"github.com/gin-contrib/sessions/redis"
|
||||
@@ -104,6 +108,13 @@ func Serve() {
|
||||
apiGroup.GET("/swagger/*any", ginSwagger.WrapHandler(swaggerFiles.Handler))
|
||||
}
|
||||
|
||||
// CAPTCHA
|
||||
capGroup := apiGroup.Group("/cap")
|
||||
{
|
||||
capGroup.POST("/challenge", capApp.Challenge)
|
||||
capGroup.POST("/redeem", capApp.Redeem)
|
||||
}
|
||||
|
||||
// API V1
|
||||
apiV1Router := apiGroup.Group("/v1")
|
||||
{
|
||||
@@ -124,7 +135,13 @@ func Serve() {
|
||||
// User
|
||||
userRouter := apiV1Router.Group("/user")
|
||||
{
|
||||
userRouter.POST("/login", user.Login)
|
||||
userRouter.POST("/login", capApp.VerifyMiddleware(capUtil.GetDefaultManager(), "login", func() bool {
|
||||
enabled, err := model.GetBoolByKey(context.Background(), model.ConfigKeyCapLoginEnabled)
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
return enabled
|
||||
}), user.Login)
|
||||
userRouter.POST("/register", user.Register)
|
||||
userRouter.GET("/logout", user.Logout)
|
||||
userRouter.GET("/self", oauth.LoginRequired(), oauth.UserInfo)
|
||||
@@ -162,6 +179,9 @@ func Serve() {
|
||||
adminRouter := apiV1Router.Group("/admin")
|
||||
adminRouter.Use(oauth.LoginRequired(), admin.LoginAdminRequired())
|
||||
{
|
||||
// System status
|
||||
adminRouter.GET("/status", admin_status.GetSystemStatus)
|
||||
|
||||
// Task dispatch
|
||||
adminRouter.GET("/tasks/types", admin_task.ListTaskTypes)
|
||||
adminRouter.POST("/tasks/dispatch", admin_task.DispatchTask)
|
||||
|
||||
@@ -134,6 +134,48 @@ func seedDefaultConfigs(t *testing.T, tx *gorm.DB) {
|
||||
Type: "business",
|
||||
Description: "限制每个普通用户可以创建的 API Key 最大数量",
|
||||
},
|
||||
{
|
||||
Key: model.ConfigKeyCapLoginEnabled,
|
||||
Value: "false",
|
||||
Type: "system",
|
||||
Description: "是否启用登录人机验证(true/false)",
|
||||
},
|
||||
{
|
||||
Key: model.ConfigKeyCapAutoSolve,
|
||||
Value: "true",
|
||||
Type: "system",
|
||||
Description: "打开页面后是否自动开始计算,关闭则需用户手动点击触发",
|
||||
},
|
||||
{
|
||||
Key: model.ConfigKeyCapChallengeCount,
|
||||
Value: "1",
|
||||
Type: "system",
|
||||
Description: "客户端需求解的 PoW 难题总数,默认 1,推荐 1~5",
|
||||
},
|
||||
{
|
||||
Key: model.ConfigKeyCapChallengeSize,
|
||||
Value: "32",
|
||||
Type: "system",
|
||||
Description: "人机验证盐值长度",
|
||||
},
|
||||
{
|
||||
Key: model.ConfigKeyCapChallengeDifficulty,
|
||||
Value: "4",
|
||||
Type: "system",
|
||||
Description: "人机验证 PoW 难度(目标前缀长度)",
|
||||
},
|
||||
{
|
||||
Key: model.ConfigKeyCapChallengeTTL,
|
||||
Value: "600",
|
||||
Type: "system",
|
||||
Description: "人机验证难题有效时间(秒)",
|
||||
},
|
||||
{
|
||||
Key: model.ConfigKeyCapTokenTTL,
|
||||
Value: "1200",
|
||||
Type: "system",
|
||||
Description: "人机验证兑换凭证有效时间(秒)",
|
||||
},
|
||||
}
|
||||
|
||||
if err := tx.Create(&defaultConfigs).Error; err != nil {
|
||||
|
||||
@@ -0,0 +1,245 @@
|
||||
package cap
|
||||
|
||||
import (
|
||||
"crypto/hmac"
|
||||
"crypto/rand"
|
||||
"crypto/sha256"
|
||||
"encoding/base64"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
const jwtHeaderB64 = "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9"
|
||||
|
||||
// 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)
|
||||
ExpiresMs 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) != 3 {
|
||||
return nil, errors.New("invalid token format")
|
||||
}
|
||||
if parts[0] != jwtHeaderB64 {
|
||||
return nil, errors.New("invalid header")
|
||||
}
|
||||
|
||||
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("signature mismatch")
|
||||
}
|
||||
|
||||
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) != 3 {
|
||||
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 = 50
|
||||
}
|
||||
if conf.Size <= 0 {
|
||||
conf.Size = 32
|
||||
}
|
||||
if conf.Difficulty <= 0 {
|
||||
conf.Difficulty = 4
|
||||
}
|
||||
if conf.ExpiresMs <= 0 {
|
||||
conf.ExpiresMs = 10 * time.Minute
|
||||
}
|
||||
|
||||
now := time.Now().UnixNano() / int64(time.Millisecond)
|
||||
expires := now + int64(conf.ExpiresMs/time.Millisecond)
|
||||
|
||||
payload := ChallengePayload{
|
||||
Nonce: randomHex(25),
|
||||
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("invalid_token")
|
||||
}
|
||||
|
||||
var payload ChallengePayload
|
||||
if err := json.Unmarshal(payloadBytes, &payload); err != nil {
|
||||
return nil, errors.New("invalid_token")
|
||||
}
|
||||
|
||||
if expectedScope != "" && payload.Scope != expectedScope {
|
||||
return nil, errors.New("scope_mismatch")
|
||||
}
|
||||
|
||||
now := time.Now().UnixNano() / int64(time.Millisecond)
|
||||
if payload.Expires < now {
|
||||
return nil, errors.New("expired")
|
||||
}
|
||||
|
||||
if len(solutions) != payload.Count {
|
||||
return nil, errors.New("invalid_solutions")
|
||||
}
|
||||
|
||||
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("invalid_solution")
|
||||
}
|
||||
}
|
||||
|
||||
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
|
||||
}
|
||||
@@ -0,0 +1,167 @@
|
||||
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)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,271 @@
|
||||
package cap
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/linux-do/credit/internal/config"
|
||||
"github.com/linux-do/credit/internal/db"
|
||||
"github.com/linux-do/credit/internal/model"
|
||||
)
|
||||
|
||||
// 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 = 1
|
||||
}
|
||||
if conf.ChallengeSize <= 0 {
|
||||
conf.ChallengeSize = 32
|
||||
}
|
||||
if conf.ChallengeDifficulty <= 0 {
|
||||
conf.ChallengeDifficulty = 4
|
||||
}
|
||||
if conf.ChallengeTTL <= 0 {
|
||||
conf.ChallengeTTL = 10 * time.Minute
|
||||
}
|
||||
if conf.TokenTTL <= 0 {
|
||||
conf.TokenTTL = 20 * time.Minute
|
||||
}
|
||||
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),
|
||||
ExpiresMs: 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
|
||||
}
|
||||
|
||||
// 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(8)
|
||||
verToken := randomHex(15)
|
||||
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) != 2 {
|
||||
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) != 2 {
|
||||
return false, nil
|
||||
}
|
||||
|
||||
expNano, err := strconv.ParseInt(valParts[0], 10, 64)
|
||||
if err != nil {
|
||||
return false, nil
|
||||
}
|
||||
tokenScope := valParts[1]
|
||||
|
||||
if expectedScope != "" && tokenScope != expectedScope {
|
||||
return false, nil
|
||||
}
|
||||
|
||||
if time.Now().UnixNano() > expNano {
|
||||
return false, nil // Expired
|
||||
}
|
||||
|
||||
return true, nil
|
||||
}
|
||||
|
||||
// sGet safely calls store.Get, treating a nil store as a miss.
|
||||
func sGet(ctx context.Context, store Store, key string) (string, bool, error) {
|
||||
if store == nil {
|
||||
return "", false, nil
|
||||
}
|
||||
return store.Get(ctx, key)
|
||||
}
|
||||
|
||||
// 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 := 1
|
||||
challengeSize := 32
|
||||
challengeDifficulty := 4
|
||||
challengeTTL := 10 * time.Minute
|
||||
tokenTTL := 20 * time.Minute
|
||||
|
||||
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
|
||||
}
|
||||
@@ -0,0 +1,45 @@
|
||||
package cap
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// fnv1a returns the 32-bit FNV-1a hash of a string
|
||||
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
|
||||
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
|
||||
}
|
||||
|
||||
// prng generates a hex string of specified length using a seed
|
||||
func prng(seed string, length int) string {
|
||||
return prngFromHash(fnv1a(seed), length)
|
||||
}
|
||||
|
||||
// prngFromHash generates a hex string of specified length using an initial hash state
|
||||
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]
|
||||
}
|
||||
@@ -0,0 +1,177 @@
|
||||
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
|
||||
}
|
||||
|
||||
func (s *MemoryStore) Get(ctx 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
|
||||
}
|
||||
|
||||
func (s *MemoryStore) Set(ctx 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
|
||||
}
|
||||
|
||||
func (s *MemoryStore) Delete(ctx 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(ctx 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(ctx 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,
|
||||
}
|
||||
}
|
||||
|
||||
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
|
||||
}
|
||||
|
||||
func (s *RedisStore) Set(ctx context.Context, key string, val string, ttl time.Duration) error {
|
||||
return s.client.Set(ctx, key, val, ttl).Err()
|
||||
}
|
||||
|
||||
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
|
||||
}
|
||||
Reference in New Issue
Block a user