From 3c687e5e0ba6dcc97d0e5f64ebb6467cc1999cd5 Mon Sep 17 00:00:00 2001 From: wangjia <809946525@qq.com> Date: Sun, 5 Jul 2026 10:42:23 +0800 Subject: [PATCH] =?UTF-8?q?feat(backend):=20=E9=99=90=E6=B5=81/=E7=99=BB?= =?UTF-8?q?=E5=BD=95=E5=A4=B1=E8=B4=A5=E9=94=81=E7=8A=B6=E6=80=81=E5=A4=96?= =?UTF-8?q?=E7=BD=AE=20Redis=EF=BC=88todo=20#2=EF=BC=89?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 新包 internal/ratelimit:内存/redis(GCRA+Lua) 双实现 + 出错逐调用降级内存 (fail-open 到内存不 fail-closed)。REDIS_ADDR 空=内存模式,行为与既往一致; 配置后跨重启保状态、支持多实例。miniredis 全覆盖测试,零真实外部依赖。 Co-Authored-By: Claude Fable 5 --- backend/config/config.go | 13 ++ backend/go.mod | 6 + backend/go.sum | 12 + backend/internal/middleware/ratelimit.go | 76 +------ backend/internal/middleware/ratelimit_test.go | 24 +- backend/internal/ratelimit/fallback.go | 106 +++++++++ backend/internal/ratelimit/fallback_test.go | 95 ++++++++ backend/internal/ratelimit/memory.go | 212 ++++++++++++++++++ backend/internal/ratelimit/memory_test.go | 101 +++++++++ backend/internal/ratelimit/redis.go | 131 +++++++++++ backend/internal/ratelimit/redis_test.go | 101 +++++++++ backend/internal/ratelimit/store.go | 83 +++++++ backend/internal/service/auth.go | 94 ++------ backend/internal/service/auth_test.go | 4 +- backend/main.go | 5 + 15 files changed, 894 insertions(+), 169 deletions(-) create mode 100644 backend/internal/ratelimit/fallback.go create mode 100644 backend/internal/ratelimit/fallback_test.go create mode 100644 backend/internal/ratelimit/memory.go create mode 100644 backend/internal/ratelimit/memory_test.go create mode 100644 backend/internal/ratelimit/redis.go create mode 100644 backend/internal/ratelimit/redis_test.go create mode 100644 backend/internal/ratelimit/store.go diff --git a/backend/config/config.go b/backend/config/config.go index 591d8d9..ac3059b 100644 --- a/backend/config/config.go +++ b/backend/config/config.go @@ -14,6 +14,7 @@ type Config struct { Storage StorageConfig Session SessionConfig RateLimit RateLimitConfig + Redis RedisConfig Pay PayConfig } @@ -68,6 +69,14 @@ type RateLimitConfig struct { PublicReadPerMin int `mapstructure:"public_read_per_min"` // 其余公开读接口(单品/release),按 IP ShopRPS int `mapstructure:"shop_rps"` // 认证流量每店每秒,按 shop_id ShopBurst int `mapstructure:"shop_burst"` // 认证流量每店突发 + // 日配额(分钟级限流之上的第二道反爬闸,0=关闭该闸):同一 IP 每自然日累计上限 + DailyProductPerIP int `mapstructure:"daily_product_per_ip"` // 单品详情(API + OG 页同池) + DailyShopListPerIP int `mapstructure:"daily_shoplist_per_ip"` // 店铺公开商品列表 +} + +// RedisConfig 限流/登录锁状态外置。Addr 为空 = 内存模式(单实例,重启清零)。 +type RedisConfig struct { + Addr string `mapstructure:"addr"` // 如 127.0.0.1:6379 } type StorageConfig struct { @@ -99,6 +108,7 @@ func Load() { _ = viper.BindEnv("pay.base_url", "PAY_BASE_URL") _ = viper.BindEnv("pay.secret", "PAY_SECRET") _ = viper.BindEnv("pay.return_url", "PAY_RETURN_URL") + _ = viper.BindEnv("redis.addr", "REDIS_ADDR") // 默认值 viper.SetDefault("server.port", "8080") @@ -123,6 +133,9 @@ func Load() { viper.SetDefault("ratelimit.public_read_per_min", 60) viper.SetDefault("ratelimit.shop_rps", 20) viper.SetDefault("ratelimit.shop_burst", 40) + viper.SetDefault("ratelimit.daily_product_per_ip", 1000) + viper.SetDefault("ratelimit.daily_shoplist_per_ip", 300) + viper.SetDefault("redis.addr", "") // 空=内存模式 viper.SetDefault("database.max_idle_conns", 10) viper.SetDefault("database.max_open_conns", 100) viper.SetDefault("storage.upload_dir", "./uploads/images") diff --git a/backend/go.mod b/backend/go.mod index ce3a707..842b5a4 100644 --- a/backend/go.mod +++ b/backend/go.mod @@ -22,9 +22,11 @@ require ( require ( filippo.io/edwards25519 v1.1.0 // indirect + github.com/alicebob/miniredis/v2 v2.38.0 // indirect github.com/bytedance/gopkg v0.1.3 // indirect github.com/bytedance/sonic v1.15.0 // indirect github.com/bytedance/sonic/loader v0.5.0 // indirect + github.com/cespare/xxhash/v2 v2.3.0 // indirect github.com/cloudwego/base64x v0.1.6 // indirect github.com/davecgh/go-spew v1.1.1 // indirect github.com/fsnotify/fsnotify v1.9.0 // indirect @@ -33,6 +35,7 @@ require ( github.com/go-playground/locales v0.14.1 // indirect github.com/go-playground/universal-translator v0.18.1 // indirect github.com/go-playground/validator/v10 v10.30.1 // indirect + github.com/go-redis/redis_rate/v10 v10.0.1 // indirect github.com/go-sql-driver/mysql v1.8.1 // indirect github.com/go-viper/mapstructure/v2 v2.4.0 // indirect github.com/goccy/go-json v0.10.5 // indirect @@ -51,6 +54,7 @@ require ( github.com/pmezard/go-difflib v1.0.0 // indirect github.com/quic-go/qpack v0.6.0 // indirect github.com/quic-go/quic-go v0.59.0 // indirect + github.com/redis/go-redis/v9 v9.21.0 // indirect github.com/richardlehane/mscfb v1.0.6 // indirect github.com/richardlehane/msoleps v1.0.6 // indirect github.com/sagikazarmark/locafero v0.11.0 // indirect @@ -64,7 +68,9 @@ require ( github.com/ugorji/go/codec v1.3.1 // indirect github.com/xuri/efp v0.0.1 // indirect github.com/xuri/nfp v0.0.2-0.20250530014748-2ddeb826f9a9 // indirect + github.com/yuin/gopher-lua v1.1.1 // indirect go.mongodb.org/mongo-driver/v2 v2.5.0 // indirect + go.uber.org/atomic v1.11.0 // indirect go.yaml.in/yaml/v3 v3.0.4 // indirect golang.org/x/arch v0.22.0 // indirect golang.org/x/image v0.25.0 // indirect diff --git a/backend/go.sum b/backend/go.sum index 46486f8..092da39 100644 --- a/backend/go.sum +++ b/backend/go.sum @@ -1,11 +1,15 @@ filippo.io/edwards25519 v1.1.0 h1:FNf4tywRC1HmFuKW5xopWpigGjJKiJSV0Cqo0cJWDaA= filippo.io/edwards25519 v1.1.0/go.mod h1:BxyFTGdWcka3PhytdK4V28tE5sGfRvvvRV7EaN4VDT4= +github.com/alicebob/miniredis/v2 v2.38.0 h1:nZAzCR+Lj+Vxk4ZXzm2NuKq2O33RXj1XxJ2e2uP9jiw= +github.com/alicebob/miniredis/v2 v2.38.0/go.mod h1:TcL7YfarKPGDAthEtl5NBeHZfeUQj6OXMm/+iu5cLMM= github.com/bytedance/gopkg v0.1.3 h1:TPBSwH8RsouGCBcMBktLt1AymVo2TVsBVCY4b6TnZ/M= github.com/bytedance/gopkg v0.1.3/go.mod h1:576VvJ+eJgyCzdjS+c4+77QF3p7ubbtiKARP3TxducM= github.com/bytedance/sonic v1.15.0 h1:/PXeWFaR5ElNcVE84U0dOHjiMHQOwNIx3K4ymzh/uSE= github.com/bytedance/sonic v1.15.0/go.mod h1:tFkWrPz0/CUCLEF4ri4UkHekCIcdnkqXw9VduqpJh0k= github.com/bytedance/sonic/loader v0.5.0 h1:gXH3KVnatgY7loH5/TkeVyXPfESoqSBSBEiDd5VjlgE= github.com/bytedance/sonic/loader v0.5.0/go.mod h1:AR4NYCk5DdzZizZ5djGqQ92eEhCCcdf5x77udYiSJRo= +github.com/cespare/xxhash/v2 v2.3.0 h1:UL815xU9SqsFlibzuggzjXhog7bL6oX9BbNZnL2UFvs= +github.com/cespare/xxhash/v2 v2.3.0/go.mod h1:VGX0DQ3Q6kWi7AoAeZDth3/j3BFtOZR5XLFGgcrjCOs= github.com/cloudwego/base64x v0.1.6 h1:t11wG9AECkCDk5fMSoxmufanudBtJ+/HemLstXDLI2M= github.com/cloudwego/base64x v0.1.6/go.mod h1:OFcloc187FXDaYHvrNIjxSe8ncn0OOM8gEHfghB2IPU= github.com/davecgh/go-spew v1.1.0/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= @@ -31,6 +35,8 @@ github.com/go-playground/universal-translator v0.18.1 h1:Bcnm0ZwsGyWbCzImXv+pAJn github.com/go-playground/universal-translator v0.18.1/go.mod h1:xekY+UJKNuX9WP91TpwSH2VMlDf28Uj24BCp08ZFTUY= github.com/go-playground/validator/v10 v10.30.1 h1:f3zDSN/zOma+w6+1Wswgd9fLkdwy06ntQJp0BBvFG0w= github.com/go-playground/validator/v10 v10.30.1/go.mod h1:oSuBIQzuJxL//3MelwSLD5hc2Tu889bF0Idm9Dg26cM= +github.com/go-redis/redis_rate/v10 v10.0.1 h1:calPxi7tVlxojKunJwQ72kwfozdy25RjA0bCj1h0MUo= +github.com/go-redis/redis_rate/v10 v10.0.1/go.mod h1:EMiuO9+cjRkR7UvdvwMO7vbgqJkltQHtwbdIQvaBKIU= github.com/go-sql-driver/mysql v1.8.1 h1:LedoTUt/eveggdHS9qUFC1EFSa8bU2+1pZjSRpvNJ1Y= github.com/go-sql-driver/mysql v1.8.1/go.mod h1:wEBSXgmK//2ZFJyE+qWnIsVGmvmEKlqwuVSjsCm7DZg= github.com/go-viper/mapstructure/v2 v2.4.0 h1:EBsztssimR/CONLSZZ04E8qAkxNYq4Qp9LvH92wZUgs= @@ -81,6 +87,8 @@ github.com/quic-go/qpack v0.6.0 h1:g7W+BMYynC1LbYLSqRt8PBg5Tgwxn214ZZR34VIOjz8= github.com/quic-go/qpack v0.6.0/go.mod h1:lUpLKChi8njB4ty2bFLX2x4gzDqXwUpaO1DP9qMDZII= github.com/quic-go/quic-go v0.59.0 h1:OLJkp1Mlm/aS7dpKgTc6cnpynnD2Xg7C1pwL6vy/SAw= github.com/quic-go/quic-go v0.59.0/go.mod h1:upnsH4Ju1YkqpLXC305eW3yDZ4NfnNbmQRCMWS58IKU= +github.com/redis/go-redis/v9 v9.21.0 h1:FPBE4hhbAke+TLmcY3WkpbDffJEomdqPn3HYiqAtL9E= +github.com/redis/go-redis/v9 v9.21.0/go.mod h1:v/M13XI1PVCDcm01VtPFOADfZtHf8YW3baQf57KlIkA= github.com/richardlehane/mscfb v1.0.6 h1:eN3bvvZCp00bs7Zf52bxNwAx5lJDBK1tCuH19qq5aC8= github.com/richardlehane/mscfb v1.0.6/go.mod h1:pe0+IUIc0AHh0+teNzBlJCtSyZdFOGgV4ZK9bsoV+Jo= github.com/richardlehane/msoleps v1.0.6 h1:9BvkpjvD+iUBalUY4esMwv6uBkfOip/Lzvd93jvR9gg= @@ -128,8 +136,12 @@ github.com/xuri/excelize/v2 v2.10.1 h1:V62UlqopMqha3kOpnlHy2CcRVw1V8E63jFoWUmMzx github.com/xuri/excelize/v2 v2.10.1/go.mod h1:iG5tARpgaEeIhTqt3/fgXCGoBRt4hNXgCp3tfXKoOIc= github.com/xuri/nfp v0.0.2-0.20250530014748-2ddeb826f9a9 h1:+C0TIdyyYmzadGaL/HBLbf3WdLgC29pgyhTjAT/0nuE= github.com/xuri/nfp v0.0.2-0.20250530014748-2ddeb826f9a9/go.mod h1:WwHg+CVyzlv/TX9xqBFXEZAuxOPxn2k1GNHwG41IIUQ= +github.com/yuin/gopher-lua v1.1.1 h1:kYKnWBjvbNP4XLT3+bPEwAXJx262OhaHDWDVOPjL46M= +github.com/yuin/gopher-lua v1.1.1/go.mod h1:GBR0iDaNXjAgGg9zfCvksxSRnQx76gclCIb7kdAd1Pw= go.mongodb.org/mongo-driver/v2 v2.5.0 h1:yXUhImUjjAInNcpTcAlPHiT7bIXhshCTL3jVBkF3xaE= go.mongodb.org/mongo-driver/v2 v2.5.0/go.mod h1:yOI9kBsufol30iFsl1slpdq1I0eHPzybRWdyYUs8K/0= +go.uber.org/atomic v1.11.0 h1:ZvwS0R+56ePWxUNi+Atn9dWONBPp/AUETXlHW0DxSjE= +go.uber.org/atomic v1.11.0/go.mod h1:LUxbIzbOniOlMKjJjyPfpl4v+PKK2cNJn91OQbhoJI0= go.uber.org/mock v0.6.0 h1:hyF9dfmbgIX5EfOdasqLsWD6xqpNZlXblLB/Dbnwv3Y= go.uber.org/mock v0.6.0/go.mod h1:KiVJ4BqZJaMj4svdfmHM0AUx4NJYO8ZNpPnZn1Z+BBU= go.yaml.in/yaml/v3 v3.0.4 h1:tfq32ie2Jv2UxXFdLJdh3jXuOzWiL1fo0bu/FbuKpbc= diff --git a/backend/internal/middleware/ratelimit.go b/backend/internal/middleware/ratelimit.go index 7a5a97b..3f898bb 100644 --- a/backend/internal/middleware/ratelimit.go +++ b/backend/internal/middleware/ratelimit.go @@ -1,77 +1,19 @@ package middleware import ( + "math" "net/http" "strconv" - "sync" - "time" "github.com/gin-gonic/gin" "golang.org/x/time/rate" "github.com/wangjia/jiu/backend/config" + "github.com/wangjia/jiu/backend/internal/ratelimit" ) -// 限流器内存条目的清理参数:每 5 分钟扫一次,淘汰超过 10 分钟未活动的 key。 -// 保证 map 不随攻击者构造的随机 key(IP/shop)无限增长。 -const ( - rateLimitSweep = 5 * time.Minute - rateLimitIdleTTL = 10 * time.Minute - rateLimitRetryHdr = "60" // Retry-After 秒数(提示客户端退避) -) - -// keyedLimiter 按任意字符串 key(IP 或 shop_id)维护独立令牌桶,内存有界(带 janitor)。 -// 单实例进程内状态,重启即清零;多实例水平扩展时需改为 Redis(见方案「暂不做」)。 -type keyedLimiter struct { - mu sync.Mutex - entries map[string]*limiterBucket - r rate.Limit - burst int -} - -type limiterBucket struct { - lim *rate.Limiter - lastSeen time.Time -} - -func newKeyedLimiter(r rate.Limit, burst int) *keyedLimiter { - kl := &keyedLimiter{entries: map[string]*limiterBucket{}, r: r, burst: burst} - go kl.janitor() - return kl -} - -// get 取(或惰性创建)该 key 的令牌桶并刷新活动时间。 -func (kl *keyedLimiter) get(key string) *rate.Limiter { - kl.mu.Lock() - defer kl.mu.Unlock() - b := kl.entries[key] - if b == nil { - b = &limiterBucket{lim: rate.NewLimiter(kl.r, kl.burst)} - kl.entries[key] = b - } - b.lastSeen = time.Now() - return b.lim -} - -func (kl *keyedLimiter) janitor() { - t := time.NewTicker(rateLimitSweep) - defer t.Stop() - for range t.C { - kl.sweep(rateLimitIdleTTL) - } -} - -// sweep 淘汰超过 ttl 未活动的 key。拆出便于测试。 -func (kl *keyedLimiter) sweep(ttl time.Duration) { - now := time.Now() - kl.mu.Lock() - defer kl.mu.Unlock() - for k, b := range kl.entries { - if now.Sub(b.lastSeen) > ttl { - delete(kl.entries, k) - } - } -} +// 限流状态存储在 internal/ratelimit(默认内存;REDIS_ADDR 配置后外置 redis, +// 跨重启保状态、支持多实例。2026-07 外置改造,原 keyedLimiter 迁入该包)。 // PerMinute 把「每分钟 n 次」转成 rate.Limit(令牌/秒)。 func PerMinute(n int) rate.Limit { @@ -86,7 +28,7 @@ func PerSecond(n int) rate.Limit { // rateLimit 通用工厂:keyFn 抽取限流维度的 key(返回空串表示无法判定 → 放行,不误伤)。 // config.C.RateLimit.Enabled=false 时整体放行(应急/测试开关)。 func rateLimit(r rate.Limit, burst int, keyFn func(*gin.Context) string) gin.HandlerFunc { - kl := newKeyedLimiter(r, burst) + lim := ratelimit.Default().NewLimiter(r, burst) return func(c *gin.Context) { if !config.C.RateLimit.Enabled { c.Next() @@ -97,8 +39,12 @@ func rateLimit(r rate.Limit, burst int, keyFn func(*gin.Context) string) gin.Han c.Next() return } - if !kl.get(key).Allow() { - c.Header("Retry-After", rateLimitRetryHdr) + if res := lim.Allow(key); !res.Allowed { + retry := int(math.Ceil(res.RetryAfter.Seconds())) + if retry <= 0 { + retry = 60 + } + c.Header("Retry-After", strconv.Itoa(retry)) c.AbortWithStatusJSON(http.StatusTooManyRequests, gin.H{"error": "请求过于频繁,请稍后再试"}) return } diff --git a/backend/internal/middleware/ratelimit_test.go b/backend/internal/middleware/ratelimit_test.go index 632cccd..9377cd4 100644 --- a/backend/internal/middleware/ratelimit_test.go +++ b/backend/internal/middleware/ratelimit_test.go @@ -107,29 +107,7 @@ func TestRateLimitDisabledPassthrough(t *testing.T) { } } -func TestKeyedLimiterSweepEvictsIdle(t *testing.T) { - kl := newKeyedLimiter(PerMinute(60), 1) - kl.get("ip:a") - kl.get("ip:b") - if len(kl.entries) != 2 { - t.Fatalf("应有 2 个 entry,得到 %d", len(kl.entries)) - } - // 把 a 的活动时间推到很久以前,sweep 应只淘汰 a。 - kl.mu.Lock() - kl.entries["ip:a"].lastSeen = time.Now().Add(-time.Hour) - kl.mu.Unlock() - - kl.sweep(10 * time.Minute) - - kl.mu.Lock() - defer kl.mu.Unlock() - if _, ok := kl.entries["ip:a"]; ok { - t.Fatal("空闲 key a 应被淘汰") - } - if _, ok := kl.entries["ip:b"]; !ok { - t.Fatal("活跃 key b 不应被淘汰") - } -} +// keyedLimiter 内部淘汰测试已随实现迁至 internal/ratelimit/memory_test.go。 // TestTrustedProxyRealIP 复刻 main.go 的可信代理配置:只信任本机写的 X-Real-IP, // 客户端伪造的 X-Forwarded-For 不被采信 → c.ClientIP() 返回真实 IP,限流不可被请求头绕过。 diff --git a/backend/internal/ratelimit/fallback.go b/backend/internal/ratelimit/fallback.go new file mode 100644 index 0000000..587d986 --- /dev/null +++ b/backend/internal/ratelimit/fallback.go @@ -0,0 +1,106 @@ +package ratelimit + +import ( + "log" + "sync" + "time" + + "github.com/redis/go-redis/v9" + "golang.org/x/time/rate" +) + +// fallbackStore:redis 为主、内存为影子。每次调用先试 redis,出错即降级到 +// 对应的内存实现(fail-open 到内存而非 fail-closed——redis 挂掉时防护降级但 +// 不中断,业务请求不 5xx)。redis 恢复后自动回主路径。降级日志限频打印。 + +type fallbackStore struct { + redis *redisStore + mem *memoryStore +} + +func newFallbackStore(rdb *redis.Client) *fallbackStore { + return &fallbackStore{redis: newRedisStore(rdb), mem: newMemoryStore()} +} + +// warnDegraded 限频告警(每 30s 至多一条,防 redis 宕机刷爆日志)。 +var ( + warnMu sync.Mutex + warnLast time.Time +) + +func warnDegraded(err error) { + warnMu.Lock() + defer warnMu.Unlock() + if time.Since(warnLast) < 30*time.Second { + return + } + warnLast = time.Now() + log.Printf("[ratelimit] redis 不可用,已降级内存限流:%v", err) +} + +func (s *fallbackStore) NewLimiter(r rate.Limit, burst int) Limiter { + return &fallbackLimiter{r: s.redis.newLimiter(r, burst), m: s.mem.NewLimiter(r, burst)} +} +func (s *fallbackStore) Counter() Counter { + return &fallbackCounter{r: &redisCounter{rdb: s.redis.rdb}, m: s.mem.Counter()} +} +func (s *fallbackStore) FailLocker() FailLocker { + return &fallbackFailLocker{r: &redisFailLocker{rdb: s.redis.rdb}, m: s.mem.FailLocker()} +} + +type fallbackLimiter struct { + r *redisLimiter + m Limiter +} + +func (l *fallbackLimiter) Allow(key string) Result { + res, err := l.r.allow(key) + if err != nil { + warnDegraded(err) + return l.m.Allow(key) + } + return res +} + +type fallbackCounter struct { + r *redisCounter + m Counter +} + +func (c *fallbackCounter) Incr(key string, ttl time.Duration) (int64, error) { + n, err := c.r.Incr(key, ttl) + if err != nil { + warnDegraded(err) + return c.m.Incr(key, ttl) + } + return n, nil +} + +type fallbackFailLocker struct { + r *redisFailLocker + m FailLocker +} + +func (l *fallbackFailLocker) Locked(key string) bool { + locked, err := l.r.locked(key) + if err != nil { + warnDegraded(err) + return l.m.Locked(key) + } + return locked +} + +func (l *fallbackFailLocker) RecordFailure(key string, max int, lockFor time.Duration) { + if err := l.r.recordFailure(key, max, lockFor); err != nil { + warnDegraded(err) + l.m.RecordFailure(key, max, lockFor) + } +} + +func (l *fallbackFailLocker) Reset(key string) { + // 双清:降级期间可能在内存里积了计数,恢复后一并清掉。 + if err := l.r.reset(key); err != nil { + warnDegraded(err) + } + l.m.Reset(key) +} diff --git a/backend/internal/ratelimit/fallback_test.go b/backend/internal/ratelimit/fallback_test.go new file mode 100644 index 0000000..b04ac4f --- /dev/null +++ b/backend/internal/ratelimit/fallback_test.go @@ -0,0 +1,95 @@ +package ratelimit + +import ( + "testing" + "time" + + "github.com/alicebob/miniredis/v2" + "github.com/redis/go-redis/v9" + "golang.org/x/time/rate" +) + +// resetDefault 清空全局 Store(Init 测试用,避免污染其他测试)。 +func resetDefault() { + defMu.Lock() + def = nil + defMu.Unlock() +} + +func TestFallbackLimiterDegradesToMemory(t *testing.T) { + mr := miniredis.RunT(t) + rdb := redis.NewClient(&redis.Options{Addr: mr.Addr()}) + s := newFallbackStore(rdb) + l := s.NewLimiter(rate.Limit(1.0/60.0), 2) + + if !l.Allow("k").Allowed { + t.Fatal("redis 正常时应放行") + } + mr.Close() // 模拟 redis 宕机 + + // 降级内存后仍在限流:burst 2 内放行、超出拒绝(内存桶从零起算) + if !l.Allow("k").Allowed || !l.Allow("k").Allowed { + t.Fatal("降级内存后 burst 内应放行(不 5xx、不 fail-closed)") + } + if l.Allow("k").Allowed { + t.Fatal("降级内存后超出 burst 仍应限流") + } +} + +func TestFallbackFailLockerDegrades(t *testing.T) { + mr := miniredis.RunT(t) + rdb := redis.NewClient(&redis.Options{Addr: mr.Addr()}) + s := newFallbackStore(rdb) + fl := s.FailLocker() + + mr.Close() + for i := 0; i < 3; i++ { + fl.RecordFailure("k", 3, time.Minute) + } + if !fl.Locked("k") { + t.Fatal("降级内存后失败锁仍应生效") + } + fl.Reset("k") + if fl.Locked("k") { + t.Fatal("Reset 应解锁") + } +} + +func TestFallbackCounterDegrades(t *testing.T) { + mr := miniredis.RunT(t) + rdb := redis.NewClient(&redis.Options{Addr: mr.Addr()}) + s := newFallbackStore(rdb) + c := s.Counter() + + mr.Close() + n, err := c.Incr("k", time.Hour) + if err != nil || n != 1 { + t.Fatalf("降级内存后计数应可用,得到 %d err=%v", n, err) + } +} + +func TestInitBadAddrFallsToMemory(t *testing.T) { + defer resetDefault() + Init("127.0.0.1:1") // 无服务端口,Ping 必败:不 panic、落内存 + l := Default().NewLimiter(rate.Limit(1), 1) + if !l.Allow("k").Allowed { + t.Fatal("坏地址应降级内存并正常工作") + } +} + +func TestInitEmptyAddrMemory(t *testing.T) { + defer resetDefault() + Init("") + if _, ok := Default().(*memoryStore); !ok { + t.Fatal("空地址应为内存 Store") + } +} + +func TestInitGoodAddrRedis(t *testing.T) { + defer resetDefault() + mr := miniredis.RunT(t) + Init(mr.Addr()) + if _, ok := Default().(*fallbackStore); !ok { + t.Fatal("可用地址应为 redis(fallback) Store") + } +} diff --git a/backend/internal/ratelimit/memory.go b/backend/internal/ratelimit/memory.go new file mode 100644 index 0000000..b9e15ab --- /dev/null +++ b/backend/internal/ratelimit/memory.go @@ -0,0 +1,212 @@ +package ratelimit + +import ( + "sync" + "time" + + "golang.org/x/time/rate" +) + +// 内存实现:自 middleware.keyedLimiter 与 service.loginLimiter 迁入(2026-07, +// 状态外置改造),行为与迁入前一致。单实例进程内状态,重启即清零。 +// 所有 map 均带 janitor 定期淘汰,保证不随攻击者构造的随机 key 无限增长。 + +const ( + memSweepEvery = 5 * time.Minute + memLimiterIdleTTL = 10 * time.Minute + memLockerIdleTTL = 30 * time.Minute + // 内存版给固定退避提示(redis 版有精确值)。 + memRetryAfter = 60 * time.Second +) + +type memoryStore struct { + counter *memCounter + locker *memFailLocker +} + +func newMemoryStore() *memoryStore { + return &memoryStore{ + counter: newMemCounter(), + locker: newMemFailLocker(), + } +} + +func (s *memoryStore) NewLimiter(r rate.Limit, burst int) Limiter { + return newMemLimiter(r, burst) +} +func (s *memoryStore) Counter() Counter { return s.counter } +func (s *memoryStore) FailLocker() FailLocker { return s.locker } + +// ── Limiter:按 key 的令牌桶(原 keyedLimiter)────────────────────────────── + +type memLimiter struct { + mu sync.Mutex + entries map[string]*memBucket + r rate.Limit + burst int +} + +type memBucket struct { + lim *rate.Limiter + lastSeen time.Time +} + +func newMemLimiter(r rate.Limit, burst int) *memLimiter { + l := &memLimiter{entries: map[string]*memBucket{}, r: r, burst: burst} + go l.janitor() + return l +} + +func (l *memLimiter) Allow(key string) Result { + l.mu.Lock() + b := l.entries[key] + if b == nil { + b = &memBucket{lim: rate.NewLimiter(l.r, l.burst)} + l.entries[key] = b + } + b.lastSeen = time.Now() + l.mu.Unlock() + return Result{Allowed: b.lim.Allow(), RetryAfter: memRetryAfter} +} + +func (l *memLimiter) janitor() { + t := time.NewTicker(memSweepEvery) + defer t.Stop() + for range t.C { + l.sweep(memLimiterIdleTTL) + } +} + +// sweep 淘汰超过 ttl 未活动的 key。拆出便于测试。 +func (l *memLimiter) sweep(ttl time.Duration) { + now := time.Now() + l.mu.Lock() + defer l.mu.Unlock() + for k, b := range l.entries { + if now.Sub(b.lastSeen) > ttl { + delete(l.entries, k) + } + } +} + +// ── Counter:带 TTL 的累加计数(日配额用)──────────────────────────────────── + +type memCounter struct { + mu sync.Mutex + entries map[string]*memCount + janitorOnce sync.Once +} + +type memCount struct { + val int64 + expireAt time.Time +} + +func newMemCounter() *memCounter { + return &memCounter{entries: map[string]*memCount{}} +} + +func (c *memCounter) Incr(key string, ttl time.Duration) (int64, error) { + c.startJanitor() + now := time.Now() + c.mu.Lock() + defer c.mu.Unlock() + e := c.entries[key] + if e == nil || now.After(e.expireAt) { + e = &memCount{expireAt: now.Add(ttl)} + c.entries[key] = e + } + e.val++ + return e.val, nil +} + +func (c *memCounter) startJanitor() { + c.janitorOnce.Do(func() { + go func() { + t := time.NewTicker(memSweepEvery) + defer t.Stop() + for range t.C { + now := time.Now() + c.mu.Lock() + for k, e := range c.entries { + if now.After(e.expireAt) { + delete(c.entries, k) + } + } + c.mu.Unlock() + } + }() + }) +} + +// ── FailLocker:登录失败计数与锁(原 service.loginLimiter)─────────────────── + +type memFailLocker struct { + mu sync.Mutex + entries map[string]*memFailEntry + janitorOnce sync.Once +} + +type memFailEntry struct { + failures int + lockedTill time.Time + lastSeen time.Time +} + +func newMemFailLocker() *memFailLocker { + return &memFailLocker{entries: map[string]*memFailEntry{}} +} + +func (l *memFailLocker) Locked(key string) bool { + l.mu.Lock() + defer l.mu.Unlock() + e := l.entries[key] + if e == nil { + return false + } + e.lastSeen = time.Now() + return time.Now().Before(e.lockedTill) +} + +func (l *memFailLocker) RecordFailure(key string, max int, lockFor time.Duration) { + l.startJanitor() + l.mu.Lock() + defer l.mu.Unlock() + e := l.entries[key] + if e == nil { + e = &memFailEntry{} + l.entries[key] = e + } + e.lastSeen = time.Now() + e.failures++ + if max > 0 && e.failures >= max { + e.lockedTill = time.Now().Add(lockFor) + e.failures = 0 + } +} + +func (l *memFailLocker) Reset(key string) { + l.mu.Lock() + defer l.mu.Unlock() + delete(l.entries, key) +} + +// startJanitor 惰性启动后台清理:淘汰「未锁定且空闲超 TTL」的 entry。 +func (l *memFailLocker) startJanitor() { + l.janitorOnce.Do(func() { + go func() { + t := time.NewTicker(memSweepEvery) + defer t.Stop() + for range t.C { + now := time.Now() + l.mu.Lock() + for k, e := range l.entries { + if now.After(e.lockedTill) && now.Sub(e.lastSeen) > memLockerIdleTTL { + delete(l.entries, k) + } + } + l.mu.Unlock() + } + }() + }) +} diff --git a/backend/internal/ratelimit/memory_test.go b/backend/internal/ratelimit/memory_test.go new file mode 100644 index 0000000..02e2217 --- /dev/null +++ b/backend/internal/ratelimit/memory_test.go @@ -0,0 +1,101 @@ +package ratelimit + +import ( + "testing" + "time" + + "golang.org/x/time/rate" +) + +// 迁自 middleware.TestKeyedLimiterSweepEvictsIdle:空闲 key 被 sweep 淘汰、活跃保留。 +func TestMemLimiterSweepEvictsIdle(t *testing.T) { + l := newMemLimiter(rate.Limit(1), 1) + l.Allow("ip:a") + l.Allow("ip:b") + if len(l.entries) != 2 { + t.Fatalf("应有 2 个 entry,得到 %d", len(l.entries)) + } + l.mu.Lock() + l.entries["ip:a"].lastSeen = time.Now().Add(-time.Hour) + l.mu.Unlock() + + l.sweep(10 * time.Minute) + + l.mu.Lock() + defer l.mu.Unlock() + if _, ok := l.entries["ip:a"]; ok { + t.Fatal("空闲 key a 应被淘汰") + } + if _, ok := l.entries["ip:b"]; !ok { + t.Fatal("活跃 key b 不应被淘汰") + } +} + +func TestMemLimiterBurst(t *testing.T) { + l := newMemLimiter(rate.Limit(1.0/60.0), 2) // 每分钟 1 个,突发 2 + if !l.Allow("k").Allowed || !l.Allow("k").Allowed { + t.Fatal("burst 内应放行") + } + res := l.Allow("k") + if res.Allowed { + t.Fatal("超出 burst 应拒绝") + } + if res.RetryAfter <= 0 { + t.Fatal("拒绝时应给 RetryAfter") + } + if !l.Allow("other").Allowed { + t.Fatal("不同 key 独立桶") + } +} + +func TestMemFailLockerLockAndExpire(t *testing.T) { + fl := newMemFailLocker() + const key = "S|u" + for i := 0; i < 3; i++ { + if fl.Locked(key) { + t.Fatalf("第 %d 次失败前不应锁定", i+1) + } + fl.RecordFailure(key, 3, 30*time.Millisecond) + } + if !fl.Locked(key) { + t.Fatal("达阈值应锁定") + } + time.Sleep(40 * time.Millisecond) + if fl.Locked(key) { + t.Fatal("锁定到期应自动解锁") + } +} + +func TestMemFailLockerReset(t *testing.T) { + fl := newMemFailLocker() + fl.RecordFailure("k", 2, time.Minute) + fl.Reset("k") + fl.RecordFailure("k", 2, time.Minute) + if fl.Locked("k") { + t.Fatal("Reset 后重新计数,1 次失败不应锁定") + } +} + +func TestMemFailLockerMaxZeroNeverLocks(t *testing.T) { + fl := newMemFailLocker() + for i := 0; i < 10; i++ { + fl.RecordFailure("k", 0, time.Minute) + } + if fl.Locked("k") { + t.Fatal("max<=0 不应锁定") + } +} + +func TestMemCounterIncrAndTTL(t *testing.T) { + c := newMemCounter() + n1, _ := c.Incr("k", 30*time.Millisecond) + n2, _ := c.Incr("k", 30*time.Millisecond) + if n1 != 1 || n2 != 2 { + t.Fatalf("累加应为 1,2,得到 %d,%d", n1, n2) + } + time.Sleep(40 * time.Millisecond) + n3, _ := c.Incr("k", 30*time.Millisecond) + if n3 != 1 { + t.Fatalf("TTL 过期后应重新从 1 计,得到 %d", n3) + } +} diff --git a/backend/internal/ratelimit/redis.go b/backend/internal/ratelimit/redis.go new file mode 100644 index 0000000..c96e986 --- /dev/null +++ b/backend/internal/ratelimit/redis.go @@ -0,0 +1,131 @@ +package ratelimit + +import ( + "context" + "time" + + redis_rate "github.com/go-redis/redis_rate/v10" + "github.com/redis/go-redis/v9" + "golang.org/x/time/rate" +) + +// redis 实现。键统一前缀 jiu:rl:(同 redis 将来可能被同宿主其他服务复用)。 +// 所有方法返回 error 供 fallback 包装器降级判断;本文件不直接暴露给调用方。 + +const ( + redisKeyPrefix = "jiu:rl:" + redisOpTimeout = 500 * time.Millisecond + // 失败计数累计窗口(对齐内存版 janitor 的 30min 空闲清理语义)。 + redisFailTTL = 30 * time.Minute +) + +type redisStore struct { + rdb *redis.Client + rl *redis_rate.Limiter +} + +func newRedisStore(rdb *redis.Client) *redisStore { + return &redisStore{rdb: rdb, rl: redis_rate.NewLimiter(rdb)} +} + +func opCtx() (context.Context, context.CancelFunc) { + return context.WithTimeout(context.Background(), redisOpTimeout) +} + +// ── Limiter:GCRA ──────────────────────────────────────────────────────────── + +type redisLimiter struct { + rl *redis_rate.Limiter + limit redis_rate.Limit +} + +// newLimiter 把 x/time/rate 的「令牌/秒」换算成 GCRA 的每事件间隔: +// Limit{Rate:1, Period: 1/r} 与 rate.Limiter(r) 的稳态速率等价,burst 语义一致。 +func (s *redisStore) newLimiter(r rate.Limit, burst int) *redisLimiter { + if r <= 0 { + r = rate.Limit(1.0 / 60.0) // 防御:非法速率按每分钟 1 次 + } + period := time.Duration(float64(time.Second) / float64(r)) + return &redisLimiter{ + rl: s.rl, + limit: redis_rate.Limit{Rate: 1, Period: period, Burst: burst}, + } +} + +func (l *redisLimiter) allow(key string) (Result, error) { + ctx, cancel := opCtx() + defer cancel() + res, err := l.rl.Allow(ctx, redisKeyPrefix+"lim:"+key, l.limit) + if err != nil { + return Result{}, err + } + ra := res.RetryAfter + if ra <= 0 { + ra = memRetryAfter + } + return Result{Allowed: res.Allowed > 0, RetryAfter: ra}, nil +} + +// ── Counter ────────────────────────────────────────────────────────────────── + +type redisCounter struct{ rdb *redis.Client } + +func (c *redisCounter) Incr(key string, ttl time.Duration) (int64, error) { + ctx, cancel := opCtx() + defer cancel() + k := redisKeyPrefix + "cnt:" + key + n, err := c.rdb.Incr(ctx, k).Result() + if err != nil { + return 0, err + } + if n == 1 { + // 首次自增设 TTL;失败不致命(键会在下个周期被重建) + _ = c.rdb.Expire(ctx, k, ttl).Err() + } + return n, nil +} + +// ── FailLocker ─────────────────────────────────────────────────────────────── + +// recordFailScript:INCR 失败计数 + 续窗口 TTL;达阈值则置锁并清计数(原子)。 +// KEYS[1]=fail 键, KEYS[2]=lock 键;ARGV[1]=failTTL 秒, ARGV[2]=max, ARGV[3]=lock 秒。 +var recordFailScript = redis.NewScript(` +local f = redis.call('INCR', KEYS[1]) +redis.call('EXPIRE', KEYS[1], ARGV[1]) +if tonumber(ARGV[2]) > 0 and f >= tonumber(ARGV[2]) then + redis.call('SET', KEYS[2], 1, 'EX', ARGV[3]) + redis.call('DEL', KEYS[1]) +end +return f +`) + +type redisFailLocker struct{ rdb *redis.Client } + +func failKeys(key string) (string, string) { + return redisKeyPrefix + "fail:" + key, redisKeyPrefix + "lock:" + key +} + +func (l *redisFailLocker) locked(key string) (bool, error) { + ctx, cancel := opCtx() + defer cancel() + _, lockKey := failKeys(key) + n, err := l.rdb.Exists(ctx, lockKey).Result() + return n > 0, err +} + +func (l *redisFailLocker) recordFailure(key string, max int, lockFor time.Duration) error { + ctx, cancel := opCtx() + defer cancel() + failKey, lockKey := failKeys(key) + return recordFailScript.Run(ctx, l.rdb, + []string{failKey, lockKey}, + int(redisFailTTL.Seconds()), max, int(lockFor.Seconds()), + ).Err() +} + +func (l *redisFailLocker) reset(key string) error { + ctx, cancel := opCtx() + defer cancel() + failKey, lockKey := failKeys(key) + return l.rdb.Del(ctx, failKey, lockKey).Err() +} diff --git a/backend/internal/ratelimit/redis_test.go b/backend/internal/ratelimit/redis_test.go new file mode 100644 index 0000000..0a8fe42 --- /dev/null +++ b/backend/internal/ratelimit/redis_test.go @@ -0,0 +1,101 @@ +package ratelimit + +import ( + "testing" + "time" + + "github.com/alicebob/miniredis/v2" + "github.com/redis/go-redis/v9" + "golang.org/x/time/rate" +) + +func newTestRedis(t *testing.T) (*miniredis.Miniredis, *redisStore) { + t.Helper() + mr := miniredis.RunT(t) + rdb := redis.NewClient(&redis.Options{Addr: mr.Addr()}) + return mr, newRedisStore(rdb) +} + +func TestRedisLimiterBurstThenDeny(t *testing.T) { + _, s := newTestRedis(t) + l := s.newLimiter(rate.Limit(1.0/60.0), 3) // 每分钟 1 个,突发 3 + + for i := 0; i < 3; i++ { + res, err := l.allow("ip:1.1.1.1") + if err != nil { + t.Fatalf("allow 报错:%v", err) + } + if !res.Allowed { + t.Fatalf("burst 内第 %d 个应放行", i+1) + } + } + res, err := l.allow("ip:1.1.1.1") + if err != nil { + t.Fatalf("allow 报错:%v", err) + } + if res.Allowed { + t.Fatal("超出 burst 应拒绝") + } + if res.RetryAfter <= 0 { + t.Fatal("拒绝时应给正的 RetryAfter") + } + // 不同 key 独立 + if res, _ := l.allow("ip:2.2.2.2"); !res.Allowed { + t.Fatal("不同 key 应独立计") + } +} + +func TestRedisFailLockerLockExpireReset(t *testing.T) { + mr, s := newTestRedis(t) + fl := &redisFailLocker{rdb: s.rdb} + const key = "S001|admin" + + for i := 0; i < 3; i++ { + if locked, _ := fl.locked(key); locked { + t.Fatalf("第 %d 次失败前不应锁定", i+1) + } + if err := fl.recordFailure(key, 3, time.Minute); err != nil { + t.Fatalf("recordFailure 报错:%v", err) + } + } + if locked, _ := fl.locked(key); !locked { + t.Fatal("达阈值应锁定") + } + // 锁定即清零计数:fail 键应已删除 + failKey, _ := failKeys(key) + if mr.Exists(failKey) { + t.Fatal("锁定后 fail 计数键应被清零删除") + } + // TTL 到期自动解锁 + mr.FastForward(61 * time.Second) + if locked, _ := fl.locked(key); locked { + t.Fatal("锁定到期应自动解锁") + } + // Reset 清空计数 + _ = fl.recordFailure(key, 3, time.Minute) + if err := fl.reset(key); err != nil { + t.Fatalf("reset 报错:%v", err) + } + if mr.Exists(failKey) { + t.Fatal("Reset 后 fail 键应删除") + } +} + +func TestRedisCounterIncrTTL(t *testing.T) { + mr, s := newTestRedis(t) + c := &redisCounter{rdb: s.rdb} + + n1, err := c.Incr("dq:test:x", time.Hour) + if err != nil || n1 != 1 { + t.Fatalf("首次应为 1,得到 %d err=%v", n1, err) + } + n2, _ := c.Incr("dq:test:x", time.Hour) + if n2 != 2 { + t.Fatalf("第二次应为 2,得到 %d", n2) + } + mr.FastForward(time.Hour + time.Second) + n3, _ := c.Incr("dq:test:x", time.Hour) + if n3 != 1 { + t.Fatalf("TTL 过期后应重新从 1 计,得到 %d", n3) + } +} diff --git a/backend/internal/ratelimit/store.go b/backend/internal/ratelimit/store.go new file mode 100644 index 0000000..157670e --- /dev/null +++ b/backend/internal/ratelimit/store.go @@ -0,0 +1,83 @@ +// Package ratelimit 提供限流/登录失败锁/日配额计数的统一状态存储抽象。 +// +// 两个实现:内存(默认,单实例)与 Redis(REDIS_ADDR 非空时启用,跨重启保状态、 +// 支持未来多实例共享)。Redis 出错时逐调用降级到内存(fail-open 到内存而非 +// fail-closed,防护不中断,见 fallback.go)。 +package ratelimit + +import ( + "context" + "log" + "sync" + "time" + + "github.com/redis/go-redis/v9" + "golang.org/x/time/rate" +) + +// Result 一次限流判定。RetryAfter 供 429 的 Retry-After 头。 +type Result struct { + Allowed bool + RetryAfter time.Duration +} + +// Limiter 速率限流(令牌桶/GCRA 语义:速率 r,突发 burst)。 +type Limiter interface { + Allow(key string) Result +} + +// Counter 累加计数(日配额等):首次自增时设 TTL,返回累加后的值。 +type Counter interface { + Incr(key string, ttl time.Duration) (int64, error) +} + +// FailLocker 登录失败锁:max 次失败锁定 lockFor(max<=0 不锁); +// 锁定即清零计数;成功登录 Reset。 +type FailLocker interface { + Locked(key string) bool + RecordFailure(key string, max int, lockFor time.Duration) + Reset(key string) +} + +// Store 聚合工厂。 +type Store interface { + NewLimiter(r rate.Limit, burst int) Limiter + Counter() Counter + FailLocker() FailLocker +} + +var ( + defMu sync.Mutex + def Store +) + +// Init 由 main.go 在 config.Load 之后、router.Setup 之前调用。 +// addr 为空 → 内存模式;Redis Ping 失败 → 打警告并落内存(启动不因 Redis 挂而失败)。 +func Init(addr string) { + defMu.Lock() + defer defMu.Unlock() + if addr == "" { + def = newMemoryStore() + return + } + rdb := redis.NewClient(&redis.Options{Addr: addr}) + ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) + defer cancel() + if err := rdb.Ping(ctx).Err(); err != nil { + log.Printf("[ratelimit] redis %s 不可用(%v),降级为内存限流", addr, err) + def = newMemoryStore() + return + } + log.Printf("[ratelimit] 限流/登录锁状态外置 redis %s", addr) + def = newFallbackStore(rdb) +} + +// Default 返回全局 Store;未 Init 时惰性落内存(测试零配置即用内存路径)。 +func Default() Store { + defMu.Lock() + defer defMu.Unlock() + if def == nil { + def = newMemoryStore() + } + return def +} diff --git a/backend/internal/service/auth.go b/backend/internal/service/auth.go index 6ad63f7..d5201d6 100644 --- a/backend/internal/service/auth.go +++ b/backend/internal/service/auth.go @@ -6,7 +6,6 @@ import ( "errors" "fmt" "log" - "sync" "time" "github.com/golang-jwt/jwt/v5" @@ -17,6 +16,7 @@ import ( "github.com/wangjia/jiu/backend/config" "github.com/wangjia/jiu/backend/internal/middleware" "github.com/wangjia/jiu/backend/internal/model" + "github.com/wangjia/jiu/backend/internal/ratelimit" ) var ( @@ -42,81 +42,17 @@ type DeviceInfo struct { UserAgent string } -// loginLimiter 内存登录失败限流器(单实例,重启即清零)。 -// 两个维度共用同一张表:账号维度 key="|",IP 维度 key="ip|", -// 分别用不同阈值锁定。带 janitor 清理空闲 entry,避免攻击者用随机 key 灌爆内存。 -type loginLimiter struct { - mu sync.Mutex - entries map[string]*limiterEntry - janitorOnce sync.Once +// 登录失败计数与锁定的状态存储在 internal/ratelimit(默认内存;REDIS_ADDR +// 配置后外置 redis,跨重启保锁。原 loginLimiter 于 2026-07 迁入该包)。 +// 两个维度:账号维度 key="|",IP 维度 key="ip|", +// 分别用不同阈值锁定。 +func loginLocker() ratelimit.FailLocker { + return ratelimit.Default().FailLocker() } -type limiterEntry struct { - failures int - lockedTill time.Time - lastSeen time.Time -} - -// loginLimiterIdleTTL:已解锁且超过该时长未活动的 entry 会被 janitor 清理。 -const loginLimiterIdleTTL = 30 * time.Minute - -var loginLim = &loginLimiter{entries: map[string]*limiterEntry{}} - -// locked 返回该 key 是否处于锁定中。 -func (l *loginLimiter) locked(key string) bool { - l.mu.Lock() - defer l.mu.Unlock() - e := l.entries[key] - if e == nil { - return false - } - e.lastSeen = time.Now() - return time.Now().Before(e.lockedTill) -} - -// recordFailure 记一次失败,达到 max 阈值则锁定(max<=0 表示该维度不锁)。 -func (l *loginLimiter) recordFailure(key string, max int) { - l.startJanitor() - l.mu.Lock() - defer l.mu.Unlock() - e := l.entries[key] - if e == nil { - e = &limiterEntry{} - l.entries[key] = e - } - e.lastSeen = time.Now() - e.failures++ - if max > 0 && e.failures >= max { - e.lockedTill = time.Now().Add(time.Duration(config.C.Session.LockMinutes) * time.Minute) - e.failures = 0 - } -} - -// reset 登录成功后清除失败计数。 -func (l *loginLimiter) reset(key string) { - l.mu.Lock() - defer l.mu.Unlock() - delete(l.entries, key) -} - -// startJanitor 惰性启动后台清理(仅一次):每 5 分钟淘汰「未锁定且超过 TTL 未活动」的 entry。 -func (l *loginLimiter) startJanitor() { - l.janitorOnce.Do(func() { - go func() { - t := time.NewTicker(5 * time.Minute) - defer t.Stop() - for range t.C { - now := time.Now() - l.mu.Lock() - for k, e := range l.entries { - if now.After(e.lockedTill) && now.Sub(e.lastSeen) > loginLimiterIdleTTL { - delete(l.entries, k) - } - } - l.mu.Unlock() - } - }() - }) +// loginLockFor 锁定时长(读配置,调用点求值以便测试改配置生效)。 +func loginLockFor() time.Duration { + return time.Duration(config.C.Session.LockMinutes) * time.Minute } type AuthService struct { @@ -144,12 +80,12 @@ func (s *AuthService) Login(shopCode, username, password string, dev DeviceInfo) ipKey = "ip|" + dev.IP } recordFail := func() { - loginLim.recordFailure(limiterKey, config.C.Session.MaxFailures) + loginLocker().RecordFailure(limiterKey, config.C.Session.MaxFailures, loginLockFor()) if ipKey != "" { - loginLim.recordFailure(ipKey, config.C.Session.IPMaxFailures) + loginLocker().RecordFailure(ipKey, config.C.Session.IPMaxFailures, loginLockFor()) } } - if loginLim.locked(limiterKey) || (ipKey != "" && loginLim.locked(ipKey)) { + if loginLocker().Locked(limiterKey) || (ipKey != "" && loginLocker().Locked(ipKey)) { s.recordLoginAttempt(shopCode, username, dev, false, "locked") return nil, nil, ErrTooManyAttempts } @@ -258,9 +194,9 @@ func (s *AuthService) Login(shopCode, username, password string, dev DeviceInfo) return nil, nil, err } - loginLim.reset(limiterKey) + loginLocker().Reset(limiterKey) if ipKey != "" { - loginLim.reset(ipKey) + loginLocker().Reset(ipKey) } user.LastLoginAt = &now diff --git a/backend/internal/service/auth_test.go b/backend/internal/service/auth_test.go index 765ffea..2ad72a7 100644 --- a/backend/internal/service/auth_test.go +++ b/backend/internal/service/auth_test.go @@ -155,7 +155,7 @@ func TestLogin_AccountLockoutAfterMaxFailures(t *testing.T) { testutil.CreateTestUser(db, shop.ID, "admin", "password123", "admin") config.C.Session.MaxFailures = 3 svc := NewAuthService(db) - defer loginLim.reset("LOCK_ACC|admin") + defer loginLocker().Reset("LOCK_ACC|admin") for i := 0; i < 3; i++ { _, _, err := svc.Login("LOCK_ACC", "admin", "wrong", DeviceInfo{Platform: "windows", IP: "10.0.0.1"}) @@ -175,7 +175,7 @@ func TestLogin_IPLockoutAcrossAccounts(t *testing.T) { config.C.Session.IPMaxFailures = 4 const attackIP = "203.0.113.9" svc := NewAuthService(db) - defer loginLim.reset("ip|" + attackIP) + defer loginLocker().Reset("ip|" + attackIP) // 4 次不同用户名(invalid_user),账号 key 各不相同永不锁;IP key 累计到 4 → 锁 IP。 for i := 0; i < 4; i++ { diff --git a/backend/main.go b/backend/main.go index ee1190a..14f157f 100644 --- a/backend/main.go +++ b/backend/main.go @@ -11,6 +11,7 @@ import ( "github.com/wangjia/jiu/backend/config" "github.com/wangjia/jiu/backend/internal/model" + "github.com/wangjia/jiu/backend/internal/ratelimit" "github.com/wangjia/jiu/backend/internal/router" "github.com/wangjia/jiu/backend/internal/service" "github.com/wangjia/jiu/backend/internal/util" @@ -20,6 +21,10 @@ func main() { // 加载配置 config.Load() + // 限流/登录锁状态存储:REDIS_ADDR 配置后外置 redis(跨重启保状态), + // 未配置或连不上时落内存(与既往行为一致)。须在 router.Setup 之前。 + ratelimit.Init(config.C.Redis.Addr) + // 生产环境启动前置检查 if config.C.Server.Mode == "release" { if config.C.Server.CORSOrigin == "*" {